diff --git a/.github/workflows/api-tests.yaml b/.github/workflows/api-tests.yaml index d8f69899f..7381b9323 100644 --- a/.github/workflows/api-tests.yaml +++ b/.github/workflows/api-tests.yaml @@ -10,7 +10,7 @@ on: paths: - ".github/workflows/api-tests.yaml" - "api/**" - - "auth/api/http/**" + - "internal/atom/**" - "channels/api/http/**" - "clients/api/http/**" - "domains/api/http/**" @@ -20,9 +20,9 @@ on: - "bootstrap/api/**" - "certs/api/http/**" - "readers/api/http/**" - - "re/api/**" - - "alarms/api/**" - - "reports/api/**" + - "re/**" + - "alarms/**" + - "reports/**" - "apidocs/openapi/**" pull_request: branches: @@ -30,7 +30,7 @@ on: paths: - ".github/workflows/api-tests.yaml" - "api/**" - - "auth/api/http/**" + - "internal/atom/**" - "channels/api/http/**" - "clients/api/http/**" - "domains/api/http/**" @@ -40,9 +40,9 @@ on: - "bootstrap/api/**" - "certs/api/http/**" - "readers/api/http/**" - - "re/api/**" - - "alarms/api/**" - - "reports/api/**" + - "re/**" + - "alarms/**" + - "reports/**" - "apidocs/openapi/**" concurrency: @@ -50,17 +50,15 @@ concurrency: cancel-in-progress: true env: - TOKENS_URL: http://localhost:9002/users/tokens/issue - CREATE_DOMAINS_URL: http://localhost:9003/domains - USER_IDENTITY: admin@example.com + ATOM_LOGIN_URL: http://localhost/auth/login + USER_IDENTITY: admin USER_SECRET: 12345678 DOMAIN_NAME: demo-test - USERS_URL: http://localhost:9002 - DOMAIN_URL: http://localhost:9003 - CLIENTS_URL: http://localhost:9006 - CHANNELS_URL: http://localhost:9005 - GROUPS_URL: http://localhost:9004 - AUTH_URL: http://localhost:9001 + USERS_URL: http://localhost + DOMAIN_URL: http://localhost + CLIENTS_URL: http://localhost + CHANNELS_URL: http://localhost + GROUPS_URL: http://localhost JOURNAL_URL: http://localhost:9021 BOOTSTRAP_URL: http://localhost:9013 CERTS_URL: http://localhost:9019 @@ -96,28 +94,29 @@ jobs: - "apidocs/openapi/journal.yaml" - "journal/api/**" - auth: - - "apidocs/openapi/auth.yaml" - - "auth/api/http/**" - domains: - "apidocs/openapi/domains.yaml" + - "internal/atom/**" - "domains/api/http/**" clients: - "apidocs/openapi/clients.yaml" + - "internal/atom/**" - "clients/api/http/**" channels: - "apidocs/openapi/channels.yaml" + - "internal/atom/**" - "channels/api/http/**" groups: - "apidocs/openapi/groups.yaml" + - "internal/atom/**" - "groups/api/http/**" users: - "apidocs/openapi/users.yaml" + - "internal/atom/**" - "users/api/**" bootstrap: @@ -134,21 +133,30 @@ jobs: re: - "apidocs/openapi/rules.yaml" - - "re/api/**" + - "re/**" + - "cmd/re/**" + - "internal/atom/**" alarms: - "apidocs/openapi/alarms.yaml" - - "alarms/api/**" + - "alarms/**" + - "cmd/alarms/**" + - "internal/atom/**" reports: - "apidocs/openapi/reports.yaml" - - "reports/api/**" + - "reports/**" + - "cmd/reports/**" + - "internal/atom/**" - name: Build images run: make all -j $(nproc) && make dockers_dev -j $(nproc) + - name: Provision Atom service tokens + run: make provision_atom_tokens + - name: Start containers - run: make run_latest up args="-d" && make run_addons up args="-d" + run: make run_latest_ci up args="-d" && make run_addons up args="-d" - name: Wait for services to be ready run: | @@ -157,24 +165,28 @@ jobs: # Check if services are responding for i in {1..30}; do - if curl -f -s http://localhost:9002/health > /dev/null 2>&1; then + if curl -f -s http://localhost/health > /dev/null 2>&1; then echo "Services are ready!" - break + exit 0 fi echo "Waiting for services... ($i/30)" sleep 2 done + echo "Services failed to become ready" >&2 + docker compose -f docker/docker-compose.yaml -f docker/docker-compose-ci.yaml --env-file docker/.env --env-file docker/.env.tokens -p "${USER_REPO:-absmach_magistrala}" ps || true + docker logs --tail 100 magistrala-nginx || true + docker logs --tail 100 magistrala-atom || true + docker logs --tail 100 magistrala-atom-bootstrap || true + exit 1 + - name: Set access token run: | - export USER_TOKEN=$(curl -sSX POST $TOKENS_URL -H "Content-Type: application/json" -d "{\"identity\": \"$USER_IDENTITY\",\"secret\": \"$USER_SECRET\"}" | jq -r .access_token) - export DOMAIN_ID=$(curl -sSX POST $CREATE_DOMAINS_URL -H "Content-Type: application/json" -H "Authorization: Bearer $USER_TOKEN" -d "{\"name\":\"$DOMAIN_NAME\",\"route\":\"$DOMAIN_NAME\"}" | jq -r .id) + export USER_TOKEN=$(curl -fsS -X POST "$ATOM_LOGIN_URL" -H "Content-Type: application/json" -d "{\"identifier\":\"$USER_IDENTITY\",\"secret\":\"$USER_SECRET\",\"kind\":\"password\"}" | jq -er .token) echo "USER_TOKEN=$USER_TOKEN" >> $GITHUB_ENV - export CLIENT_SECRET=$(magistrala-cli provision test | /usr/bin/grep -Eo '"secret": "[^"]+"' | awk 'NR % 2 == 0' | sed 's/"secret": "\(.*\)"/\1/') - echo "CLIENT_SECRET=$CLIENT_SECRET" >> $GITHUB_ENV - name: Run Users API tests - if: steps.changes.outputs.users == 'true' || steps.changes.outputs.workflow == 'true' + if: (steps.changes.outputs.users == 'true' || steps.changes.outputs.workflow == 'true') && hashFiles('users/api/**') != '' uses: schemathesis/action@v3.0.0 with: schema: apidocs/openapi/users.yaml @@ -183,7 +195,7 @@ jobs: args: '--header "Authorization: Bearer ${{ env.USER_TOKEN }}" --suppress-health-check=filter_too_much --exclude-checks=positive_data_acceptance --exclude-operation-id=requestPasswordReset --phases=examples' - name: Run Groups API tests - if: steps.changes.outputs.groups == 'true' || steps.changes.outputs.workflow == 'true' + if: (steps.changes.outputs.groups == 'true' || steps.changes.outputs.workflow == 'true') && hashFiles('groups/api/http/**') != '' uses: schemathesis/action@v3.0.0 with: schema: apidocs/openapi/groups.yaml @@ -192,7 +204,7 @@ jobs: args: '--header "Authorization: Bearer ${{ env.USER_TOKEN }}" --suppress-health-check=filter_too_much --exclude-checks=positive_data_acceptance --phases=examples' - name: Run Clients API tests - if: steps.changes.outputs.clients == 'true' || steps.changes.outputs.workflow == 'true' + if: (steps.changes.outputs.clients == 'true' || steps.changes.outputs.workflow == 'true') && hashFiles('clients/api/http/**') != '' uses: schemathesis/action@v3.0.0 with: schema: apidocs/openapi/clients.yaml @@ -201,7 +213,7 @@ jobs: args: '--header "Authorization: Bearer ${{ env.USER_TOKEN }}" --suppress-health-check=filter_too_much --exclude-checks=positive_data_acceptance --phases=examples' - name: Run Channels API tests - if: steps.changes.outputs.channels == 'true' || steps.changes.outputs.workflow == 'true' + if: (steps.changes.outputs.channels == 'true' || steps.changes.outputs.workflow == 'true') && hashFiles('channels/api/http/**') != '' uses: schemathesis/action@v3.0.0 with: schema: apidocs/openapi/channels.yaml @@ -209,17 +221,8 @@ jobs: checks: all args: '--header "Authorization: Bearer ${{ env.USER_TOKEN }}" --suppress-health-check=filter_too_much --exclude-checks=positive_data_acceptance --phases=examples' - - name: Run Auth API tests - if: steps.changes.outputs.auth == 'true' || steps.changes.outputs.workflow == 'true' - uses: schemathesis/action@v3.0.0 - with: - schema: apidocs/openapi/auth.yaml - base-url: ${{ env.AUTH_URL }} - checks: all - args: '--header "Authorization: Bearer ${{ env.USER_TOKEN }}" --suppress-health-check=filter_too_much --exclude-checks=positive_data_acceptance --phases=examples' - - name: Run Domains API tests - if: steps.changes.outputs.domains == 'true' || steps.changes.outputs.workflow == 'true' + if: (steps.changes.outputs.domains == 'true' || steps.changes.outputs.workflow == 'true') && hashFiles('domains/api/http/**') != '' uses: schemathesis/action@v3.0.0 with: schema: apidocs/openapi/domains.yaml @@ -237,7 +240,7 @@ jobs: args: '--header "Authorization: Bearer ${{ env.USER_TOKEN }}" --suppress-health-check=filter_too_much --exclude-checks=positive_data_acceptance --phases=examples' - name: Run Bootstrap API tests - if: steps.changes.outputs.bootstrap == 'true' || steps.changes.outputs.workflow == 'true' + if: (steps.changes.outputs.bootstrap == 'true' || steps.changes.outputs.workflow == 'true') && hashFiles('bootstrap/api/**') != '' uses: schemathesis/action@v3.0.0 with: schema: apidocs/openapi/bootstrap.yaml @@ -246,7 +249,7 @@ jobs: args: '--header "Authorization: Bearer ${{ env.USER_TOKEN }}" --suppress-health-check=filter_too_much --exclude-checks=positive_data_acceptance --phases=examples' - name: Run Certs API tests - if: steps.changes.outputs.certs == 'true' || steps.changes.outputs.workflow == 'true' + if: (steps.changes.outputs.certs == 'true' || steps.changes.outputs.workflow == 'true') && hashFiles('docker/addons/certs/docker-compose.yaml') != '' uses: schemathesis/action@v3.0.0 with: schema: apidocs/openapi/certs.yaml @@ -292,4 +295,4 @@ jobs: - name: Stop containers if: always() - run: make run_latest down args="-v" && make run_addons down args="-v" + run: make run_latest_ci down args="-v" && make run_addons down args="-v" diff --git a/.github/workflows/check-generated-files.yaml b/.github/workflows/check-generated-files.yaml index b02f79ac6..55b83765e 100644 --- a/.github/workflows/check-generated-files.yaml +++ b/.github/workflows/check-generated-files.yaml @@ -50,6 +50,7 @@ jobs: - "pkg/messaging/*.pb.go" mocks: + - "tools/config/.mockery.yaml" - ".github/workflows/check-generated-files.yaml" - "pkg/sdk/sdk.go" - "users/postgres/clients.go" diff --git a/.github/workflows/lint-and-build.yaml b/.github/workflows/lint-and-build.yaml index 53d6a22aa..4773c4f5f 100644 --- a/.github/workflows/lint-and-build.yaml +++ b/.github/workflows/lint-and-build.yaml @@ -45,17 +45,9 @@ jobs: make all -j $(nproc) compile-check: - name: Compile Check ${{ matrix.variant.name }} + name: Compile Check Redis Event Store runs-on: ubuntu-latest needs: lint - strategy: - fail-fast: true - matrix: - variant: - - name: redis - env: MG_ES_TYPE=es_redis - target: fluxmq - steps: - name: Checkout code uses: actions/checkout@v7 @@ -66,6 +58,6 @@ jobs: go-version-file: go.mod cache-dependency-path: "go.sum" - - name: Compile check for ${{ matrix.variant.name }} + - name: Compile check run: | - ${{ matrix.variant.env }} make ${{ matrix.variant.target }} + MG_ES_TYPE=es_redis make all -j $(nproc) diff --git a/.github/workflows/tests.yaml b/.github/workflows/tests.yaml index 6e8bb4a5a..e1463d9bf 100644 --- a/.github/workflows/tests.yaml +++ b/.github/workflows/tests.yaml @@ -66,28 +66,8 @@ jobs: workflow: - ".github/workflows/tests.yaml" - auth: - - "auth/**" - - "cmd/auth/**" - - "auth.proto" - - "auth.pb.go" - - "auth_grpc.pb.go" - - "pkg/ulid/**" - - "pkg/uuid/**" - - bootstrap: - - "bootstrap/**" - - "cmd/bootstrap/**" - - "pkg/bootstrap/**" - - "provision/**" - - "pkg/sdk/**" - channels: - "channels/**" - - "cmd/channels/**" - - "auth.pb.go" - - "auth_grpc.pb.go" - - "auth/**" - "pkg/sdk/**" - "clients/api/grpc/**" - "groups/api/grpc/**" @@ -101,10 +81,6 @@ jobs: clients: - "clients/**" - - "cmd/clients/**" - - "auth.pb.go" - - "auth_grpc.pb.go" - - "auth/**" - "pkg/ulid/**" - "pkg/uuid/**" - "pkg/events/**" @@ -115,18 +91,10 @@ jobs: domains: - "domains/**" - - "cmd/domains/**" - - "auth.pb.go" - - "auth_grpc.pb.go" - - "auth/**" - "internal/grpc/**" groups: - "groups/**" - - "cmd/groups/**" - - "auth.pb.go" - - "auth_grpc.pb.go" - - "auth/**" - "pkg/ulid/**" - "pkg/uuid/**" - "clients/api/grpc/**" @@ -140,9 +108,6 @@ jobs: journal: - "journal/**" - "cmd/journal/**" - - "auth.pb.go" - - "auth_grpc.pb.go" - - "auth/**" - "pkg/events/**" logger: @@ -165,7 +130,6 @@ jobs: - "pkg/sdk/**" - "pkg/errors/**" - "pkg/groups/**" - - "auth/**" - "internal/*" - "clients/**" - "users/**" @@ -189,10 +153,6 @@ jobs: users: - "users/**" - - "cmd/users/**" - - "auth.pb.go" - - "auth_grpc.pb.go" - - "auth/**" - "pkg/ulid/**" - "pkg/uuid/**" - "pkg/events/**" @@ -200,8 +160,6 @@ jobs: notifications: - "notifications/**" - "cmd/notifications/**" - - "auth.pb.go" - - "auth_grpc.pb.go" - "consumers/notifier.go" - "pkg/events/**" @@ -240,11 +198,10 @@ jobs: if [[ "${{ steps.changes.outputs.workflow }}" == "true" || "${{ steps.changes.outputs.pkg-errors }}" == "true" ]]; then # If workflow or pkg/errors changed, test everything - modules=("auth" "bootstrap" "channels" "cli" "clients" "domains" "groups" "internal" "journal" "logger" "pkg-errors" "pkg-events" "pkg-grpcclient" "pkg-messaging" "pkg-sdk" "pkg-transformers" "pkg-ulid" "pkg-uuid" "users" "notifications" "api" "consumers" "readers" "re" "alarms" "reports") + modules=("auth" "channels" "cli" "clients" "domains" "groups" "internal" "journal" "logger" "pkg-errors" "pkg-events" "pkg-grpcclient" "pkg-messaging" "pkg-sdk" "pkg-transformers" "pkg-ulid" "pkg-uuid" "users" "notifications" "api" "consumers" "readers" "re" "alarms" "reports") else # Add only changed modules [[ "${{ steps.changes.outputs.auth }}" == "true" ]] && modules+=("auth") - [[ "${{ steps.changes.outputs.bootstrap }}" == "true" ]] && modules+=("bootstrap") [[ "${{ steps.changes.outputs.channels }}" == "true" ]] && modules+=("channels") [[ "${{ steps.changes.outputs.cli }}" == "true" ]] && modules+=("cli") [[ "${{ steps.changes.outputs.clients }}" == "true" ]] && modules+=("clients") @@ -314,9 +271,19 @@ jobs: pkg-transformers) dir="pkg/transformers" ;; pkg-ulid) dir="pkg/ulid" ;; pkg-uuid) dir="pkg/uuid" ;; + channels) dir="pkg/channels" ;; + clients) dir="pkg/clients" ;; + domains) dir="pkg/domains" ;; + groups) dir="pkg/groups" ;; *) dir="${{ matrix.module }}" ;; esac + if [[ ! -d "$dir" ]]; then + echo "Skipping ${{ matrix.module }}; ./$dir is not present in this branch" + echo "mode: atomic" > coverage-${{ matrix.module }}.out + exit 0 + fi + go test -mod=readonly --race -v -count=1 -failfast -coverprofile=coverage-${{ matrix.module }}.out ./$dir/... - name: Upload coverage diff --git a/.gitignore b/.gitignore index d55d99b2f..72f5953cd 100644 --- a/.gitignore +++ b/.gitignore @@ -19,5 +19,11 @@ coverage # Ignore Openbao data directory as it contains runtime-generated data docker/addons/certs/openbao/ +# Ignore generated local Atom service tokens. +docker/.env.tokens + # Ignore SeaweedFS data directory as it contains runtime-generated data docker/data/* + +demo-ui +node_modules diff --git a/Makefile b/Makefile index 5e2a8bdfc..9fc347c12 100644 --- a/Makefile +++ b/Makefile @@ -4,8 +4,8 @@ override MG_DOCKER_IMAGE_NAME_PREFIX := ghcr.io/absmach/magistrala MG_DOCKER_VOLUME_NAME_PREFIX ?= magistrala BUILD_DIR ?= build -SERVICES = auth users clients groups channels domains notifications certs re postgres-writer postgres-reader timescale-writer timescale-reader cli alarms reports bootstrap provision journal fluxmq -TEST_API_SERVICES = journal auth certs clients users channels groups domains +SERVICES = atom-bootstrap notifications certs re postgres-writer postgres-reader timescale-writer timescale-reader alarms reports journal fluxmq +TEST_API_SERVICES = journal certs clients users channels groups domains TEST_API = $(addprefix test_api_,$(TEST_API_SERVICES)) DOCKERS = $(addprefix docker_,$(SERVICES)) DOCKERS_DEV = $(addprefix docker_dev_,$(SERVICES)) @@ -23,6 +23,15 @@ space:= $(empty) $(empty) DOCKER_PROJECT ?= $(shell echo $(subst $(space),,$(USER_REPO)) | sed -E 's/[^a-zA-Z0-9]/_/g' | tr '[:upper:]' '[:lower:]') DOCKER_COMPOSE_COMMANDS_SUPPORTED := up down config restart DEFAULT_DOCKER_COMPOSE_COMMAND := up +ATOM_TOKENS_ENV ?= docker/.env.tokens +REQUIRED_ATOM_TOKEN_ENVS := MG_ATOM_TOKEN_FLUXMQ_AUTH MG_ATOM_TOKEN_FLUXMQ_NODE1 MG_ATOM_TOKEN_FLUXMQ_NODE2 MG_ATOM_TOKEN_FLUXMQ_NODE3 MG_ATOM_TOKEN_JOURNAL MG_ATOM_TOKEN_NOTIFICATIONS MG_ATOM_TOKEN_TIMESCALE_READER MG_ATOM_TOKEN_RE MG_ATOM_TOKEN_ALARMS MG_ATOM_TOKEN_REPORTS MG_ATOM_TOKEN_POSTGRES_READER +PROVISION_ATOM_TOKENS ?= false +PROVISION_ATOM_TOKEN_GOALS := provision-atom-tokens +DOCKER_BASE_ENV_FILES := --env-file docker/.env +DOCKER_ENV_FILES = $(if $(filter down,$(DOCKER_COMPOSE_COMMAND)),$(DOCKER_BASE_ENV_FILES),$(DOCKER_BASE_ENV_FILES) --env-file $(ATOM_TOKENS_ENV)) +DOCKER_PROVISION_ENV_FILES = $(DOCKER_BASE_ENV_FILES) $(if $(wildcard $(ATOM_TOKENS_ENV)),--env-file $(ATOM_TOKENS_ENV)) +HOST_UID := $(shell id -u) +HOST_GID := $(shell id -g) GRPC_MTLS_CERT_FILES_EXISTS = 0 MOCKERY = $(GOBIN)/mockery MOCKERY_VERSION=3.6.4 @@ -79,7 +88,48 @@ define make_docker_dev -f docker/Dockerfile.dev ./build endef +define require_atom_tokens_env + @if [ -z "$(filter down,$(DOCKER_COMPOSE_COMMAND))" ]; then \ + if [ ! -f "$(ATOM_TOKENS_ENV)" ]; then \ + echo "Missing $(ATOM_TOKENS_ENV). Run 'make provision_atom_tokens' before starting the Docker Compose stack."; \ + exit 2; \ + fi; \ + missing=""; \ + for env_name in $(REQUIRED_ATOM_TOKEN_ENVS); do \ + if ! grep -q "^$${env_name}=" "$(ATOM_TOKENS_ENV)"; then \ + missing="$${missing} $${env_name}"; \ + fi; \ + done; \ + if [ -n "$${missing}" ]; then \ + echo "Missing Atom service token(s) in $(ATOM_TOKENS_ENV):$${missing}. Run 'make provision_atom_tokens' before starting the Docker Compose stack."; \ + exit 2; \ + fi; \ + fi +endef + +define ensure_atom_tokens_env + @if [ "$(PROVISION_ATOM_TOKENS)" = "true" ] && [ -z "$(filter down,$(DOCKER_COMPOSE_COMMAND))" ]; then \ + $(MAKE) provision_atom_tokens; \ + elif [ -z "$(filter down,$(DOCKER_COMPOSE_COMMAND))" ]; then \ + if [ ! -f "$(ATOM_TOKENS_ENV)" ]; then \ + echo "Missing $(ATOM_TOKENS_ENV). Run 'make provision_atom_tokens' before starting the Docker Compose stack."; \ + exit 2; \ + fi; \ + missing=""; \ + for env_name in $(REQUIRED_ATOM_TOKEN_ENVS); do \ + if ! grep -q "^$${env_name}=" "$(ATOM_TOKENS_ENV)"; then \ + missing="$${missing} $${env_name}"; \ + fi; \ + done; \ + if [ -n "$${missing}" ]; then \ + echo "Missing Atom service token(s) in $(ATOM_TOKENS_ENV):$${missing}. Run 'make provision_atom_tokens' before starting the Docker Compose stack."; \ + exit 2; \ + fi; \ + fi +endef + define run_with_arch_detection + $(call require_atom_tokens_env) @echo "Detecting architecture..." @if [ "$(DETECTED_ARCH)" = "arm64" ] || [ "$(DETECTED_ARCH)" = "aarch64" ]; then \ echo "ARM64 architecture detected."; \ @@ -89,12 +139,12 @@ define run_with_arch_detection docker tag $(MG_DOCKER_IMAGE_NAME_PREFIX)/$$svc $(MG_DOCKER_IMAGE_NAME_PREFIX)/$$svc:latest; \ done; \ sed -i.bak 's/^MG_RELEASE_TAG=.*/MG_RELEASE_TAG=latest/' docker/.env && rm -f docker/.env.bak; \ - docker compose -f docker/docker-compose.yaml --env-file docker/.env -p $(DOCKER_PROJECT) $(DOCKER_COMPOSE_COMMAND) $(args); \ + docker compose -f docker/docker-compose.yaml $(DOCKER_ENV_FILES) -p $(DOCKER_PROJECT) $(DOCKER_COMPOSE_COMMAND) $(args); \ else \ echo "x86_64 architecture detected."; \ git checkout $(1); \ sed -i.bak 's/^MG_RELEASE_TAG=.*/MG_RELEASE_TAG=$(2)/' docker/.env && rm -f docker/.env.bak; \ - docker compose -f docker/docker-compose.yaml --env-file docker/.env -p $(DOCKER_PROJECT) $(DOCKER_COMPOSE_COMMAND) $(args); \ + docker compose -f docker/docker-compose.yaml $(DOCKER_ENV_FILES) -p $(DOCKER_PROJECT) $(DOCKER_COMPOSE_COMMAND) $(args); \ fi endef @@ -125,6 +175,9 @@ DOCKER_PLATFORM ?= ifneq ($(filter run%,$(firstword $(MAKECMDGOALS))),) temp_args := $(wordlist 2,$(words $(MAKECMDGOALS)),$(MAKECMDGOALS)) DOCKER_COMPOSE_COMMAND := $(if $(filter $(DOCKER_COMPOSE_COMMANDS_SUPPORTED),$(temp_args)), $(filter $(DOCKER_COMPOSE_COMMANDS_SUPPORTED),$(temp_args)), $(DEFAULT_DOCKER_COMPOSE_COMMAND)) + ifneq ($(filter $(PROVISION_ATOM_TOKEN_GOALS),$(temp_args)),) + override PROVISION_ATOM_TOKENS := true + endif $(eval $(DOCKER_COMPOSE_COMMAND):;@) endif @@ -144,7 +197,7 @@ FILTERED_SERVICES = $(filter-out $(RUN_ADDON_ARGS), $(SERVICES)) all: $(SERVICES) -.PHONY: all $(SERVICES) dockers dockers_dev latest release run_latest run_tls run_stable run_addons grpc_mtls_certs check_mtls check_certs test_api mocks +.PHONY: all $(SERVICES) dockers dockers_dev latest release provision_atom_tokens provision-atom-tokens run_latest run_latest_ci run_tls run_stable run_addons grpc_mtls_certs check_mtls check_certs test_api mocks clean: rm -rf ${BUILD_DIR} @@ -199,12 +252,11 @@ define test_api_service --phases=examples,stateful endef -test_api_users: TEST_API_URL := http://localhost:9002 -test_api_clients: TEST_API_URL := http://localhost:9006 -test_api_domains: TEST_API_URL := http://localhost:9003 -test_api_channels: TEST_API_URL := http://localhost:9005 -test_api_groups: TEST_API_URL := http://localhost:9004 -test_api_auth: TEST_API_URL := http://localhost:9001 +test_api_users: TEST_API_URL := http://localhost:9000 +test_api_clients: TEST_API_URL := http://localhost:9000 +test_api_domains: TEST_API_URL := http://localhost:9000 +test_api_channels: TEST_API_URL := http://localhost:9000 +test_api_groups: TEST_API_URL := http://localhost:9000 test_api_certs: TEST_API_URL := http://localhost:9019 test_api_journal: TEST_API_URL := http://localhost:9021 @@ -262,7 +314,15 @@ rundev: cd scripts && ./run.sh grpc_mtls_certs: - $(MAKE) -C docker/ssl auth_grpc_certs clients_grpc_certs + $(MAKE) -C docker/ssl clients_grpc_certs + +provision_atom_tokens: + $(DOCKER_PLATFORM) docker compose -f docker/docker-compose.yaml $(DOCKER_PROVISION_ENV_FILES) -p $(DOCKER_PROJECT) up -d --wait --wait-timeout 120 atom + $(MAKE) docker_atom-bootstrap + $(DOCKER_PLATFORM) docker compose -f docker/docker-compose.yaml $(DOCKER_PROVISION_ENV_FILES) -p $(DOCKER_PROJECT) run --rm --no-deps --user "$(HOST_UID):$(HOST_GID)" -v "$(PWD)/docker:/host/docker" atom-bootstrap provision-tokens --output /host/docker/.env.tokens + +provision-atom-tokens: + @: check_tls: ifeq ($(GRPC_TLS),true) @@ -284,14 +344,20 @@ check_certs: check_mtls check_tls ifeq ($(GRPC_MTLS_CERT_FILES_EXISTS),0) ifeq ($(filter true,$(GRPC_MTLS) $(GRPC_TLS)),true) ifeq ($(filter $(DEFAULT_DOCKER_COMPOSE_COMMAND),$(DOCKER_COMPOSE_COMMAND)),$(DEFAULT_DOCKER_COMPOSE_COMMAND)) - $(MAKE) -C docker/ssl auth_grpc_certs clients_grpc_certs + $(MAKE) -C docker/ssl clients_grpc_certs endif endif endif run_latest: check_certs $(SED_INPLACE) 's/^MG_RELEASE_TAG=.*/MG_RELEASE_TAG=latest/' docker/.env - $(DOCKER_PLATFORM) docker compose -f docker/docker-compose.yaml --env-file docker/.env -p $(DOCKER_PROJECT) $(DOCKER_COMPOSE_COMMAND) $(args) + $(call ensure_atom_tokens_env) + $(DOCKER_PLATFORM) docker compose -f docker/docker-compose.yaml $(DOCKER_ENV_FILES) -p $(DOCKER_PROJECT) $(DOCKER_COMPOSE_COMMAND) $(args) + +run_latest_ci: check_certs + $(call require_atom_tokens_env) + $(SED_INPLACE) 's/^MG_RELEASE_TAG=.*/MG_RELEASE_TAG=latest/' docker/.env + $(DOCKER_PLATFORM) docker compose -f docker/docker-compose.yaml -f docker/docker-compose-ci.yaml $(DOCKER_ENV_FILES) -p $(DOCKER_PROJECT) $(DOCKER_COMPOSE_COMMAND) $(args) run_tls: @test -n "$(host)" || (echo "Usage: make run_tls host=example.com [email=admin@example.com] [letsencrypt=false] [staging=true] [force=true]" && exit 2) @@ -305,17 +371,20 @@ run_tls: ./docker/setup-tls.sh run_stable: check_certs + $(call require_atom_tokens_env) $(eval version = $(shell git describe --abbrev=0 --tags)) git checkout $(version) $(SED_INPLACE) 's/^MG_RELEASE_TAG=.*/MG_RELEASE_TAG=$(version)/' docker/.env - $(DOCKER_PLATFORM) docker compose -f docker/docker-compose.yaml --env-file docker/.env -p $(DOCKER_PROJECT) $(DOCKER_COMPOSE_COMMAND) $(args) + $(DOCKER_PLATFORM) docker compose -f docker/docker-compose.yaml $(DOCKER_ENV_FILES) -p $(DOCKER_PROJECT) $(DOCKER_COMPOSE_COMMAND) $(args) run_addons: check_certs + $(call require_atom_tokens_env) $(foreach SVC,$(RUN_ADDON_ARGS),$(if $(filter $(SVC),$(ADDON_SERVICES) $(EXTERNAL_SERVICES)),,$(error Invalid Service $(SVC)))) - @$(DOCKER_PLATFORM) docker compose -f docker/docker-compose.yaml --env-file ./docker/.env -p $(DOCKER_PROJECT) up -d auth domains jaeger + @$(DOCKER_PLATFORM) docker compose -f docker/docker-compose.yaml $(DOCKER_ENV_FILES) -p $(DOCKER_PROJECT) up -d atom jaeger @for SVC in $(RUN_ADDON_ARGS); do \ - MG_ADDONS_CERTS_PATH_PREFIX="../" $(DOCKER_PLATFORM) docker compose -f docker/addons/$$SVC/docker-compose.yaml -p $(DOCKER_PROJECT) --env-file ./docker/.env $(DOCKER_COMPOSE_COMMAND) $(args) & \ + MG_ADDONS_CERTS_PATH_PREFIX="../" $(DOCKER_PLATFORM) docker compose -f docker/addons/$$SVC/docker-compose.yaml -p $(DOCKER_PROJECT) $(DOCKER_ENV_FILES) $(DOCKER_COMPOSE_COMMAND) $(args) & \ done run_live: check_certs - GOPATH=$(go env GOPATH) $(DOCKER_PLATFORM) docker compose -f docker/docker-compose.yaml -f docker/docker-compose-live.yaml --env-file docker/.env -p $(DOCKER_PROJECT) $(DOCKER_COMPOSE_COMMAND) $(args) + $(call require_atom_tokens_env) + GOPATH=$(go env GOPATH) $(DOCKER_PLATFORM) docker compose -f docker/docker-compose.yaml -f docker/docker-compose-live.yaml $(DOCKER_ENV_FILES) -p $(DOCKER_PROJECT) $(DOCKER_COMPOSE_COMMAND) $(args) diff --git a/README.md b/README.md index 7da4e6003..27fd7606b 100644 --- a/README.md +++ b/README.md @@ -46,7 +46,7 @@ It is extremely flexible and lets you build systems the way you want — from si At the same time, it avoids the typical complexity of many IoT platforms, where you need to learn an entirely new set of concepts before you can even get started. -Magistrala is built around a small number of core concepts: +Magistrala is built around a small number of main concepts: - users - clients (devices) - channels @@ -141,6 +141,138 @@ Magistrala provides a complete set of building blocks for IoT systems — from d - Documentation focused on getting you running quickly --- +## Atom Integration Model + +Magistrala uses **Atom** as the backend for identity, authorization, and the catalog. + +Atom is the source of truth for: +- domains +- users +- clients +- channels +- groups +- roles +- access policies + +Magistrala services such as rules, alarms, and reports remain Magistrala services, but they use Atom for identity and authorization. + +### Core Entity Mapping + +| Magistrala concept | Atom concept | Meaning | +|--------------------|--------------|---------| +| Domain | Tenant | Isolation boundary for one organization, project, or environment | +| User | Entity with kind `human` | A person who logs in and uses the UI/API | +| Client | Entity with kind `device` | A device or application that sends/receives data | +| Channel | Resource with kind `channel` | A messaging/data path that clients can publish or subscribe to | +| Group | Group | A collection of users, clients, channels, or other grouped objects | + +In simple terms: + +```text +MG Domain = Atom Tenant +MG User = Atom Human Entity +MG Client = Atom Device Entity +MG Channel = Atom Channel Resource +MG Group = Atom Group +``` + +### Actions, Permission Blocks, Roles, and Assignments + +Atom access control has these basic parts: + +| Atom word | Simple meaning | Example | +|-----------|----------------|---------| +| Action | One permission verb | `read`, `write`, `delete`, `role.manage`, `policy.manage` | +| Permission Block | Where actions apply | all channels in domain `d1` can `read`, `publish` | +| Role | A bundle of permission blocks | `tenant-admin` bundles domain, role, and member access | +| Role Assignment | Who gets a role | give `user1` the `tenant-admin` role | + +Read an assignment like this: + +```text +Give this . +The role contains permission blocks that say where and what. +``` + +Example: + +```text +Give user1 the tenant-admin role on domain d1. +``` + +That means: + +```text +user1 can use the tenant-admin permissions inside domain d1. +``` + +### How MG Roles Work With Atom + +MG UI shows actions such as: +- read +- update +- delete +- manage roles +- add/remove members +- publish +- subscribe + +These are mapped to Atom actions: + +| MG action | Atom action | +|-----------|-----------------| +| view/read | `read` | +| create/update/edit/connect | `write` | +| delete/remove | `delete` | +| manage roles | `role.manage` | +| add/remove members or access | `policy.manage` | +| channel publish | `publish` | +| channel subscribe | `subscribe` | + +So when MG UI checks: + +```text +Can user1 manage roles for client1? +``` + +Atom checks: + +```text +Does user1 have role.manage on client1, or on the domain that contains client1? +``` + +When MG UI checks: + +```text +Can user1 add a member to channel1? +``` + +Atom checks: + +```text +Does user1 have policy.manage on channel1, or on the domain that contains channel1? +``` + +### Practical Rule + +If a user is domain admin, they usually receive a tenant-scoped role in Atom. + +That tenant-scoped role can allow them to manage objects inside the domain: +- clients +- channels +- groups +- rules +- alarms +- reports + +For narrower access, create object-scoped roles. For example: + +```text +Give user2 a reader role only on channel1. +``` + +Then user2 can read only that channel, not the whole domain. + ## Installation ```bash @@ -151,54 +283,6 @@ make run_latest --- -## Upgrade from v0.19.0 to v0.20.0 - -Before upgrading, back up the Domains, Rules Engine, Reports, Alarms, Auth, and SpiceDB databases. - -v0.20.0 adds new domain admin actions for alarms and reports, and it requires existing rules and reports to have their built-in admin roles backfilled. The service database migrations run when the v0.20.0 services start, then the role backfill scripts must be run once. - -For the default Docker Compose setup: - -```bash -cd docker - -docker compose up -d \ - spicedb-db spicedb-migrate spicedb \ - auth-db auth \ - domains-db domains \ - re-db re \ - reports-db reports \ - alarms-db alarms -``` - -Wait until the services are running. The `auth` service must start successfully because it loads the SpiceDB schema. - -From the repository root, run the backfills: - -```bash -go run ./scripts/re-backfill-roles/ -go run ./scripts/reports-backfill-roles/ -``` - -The scripts are idempotent. If they are interrupted, fix the issue and run them again. - -Expected successful summaries: - -```text -backfill finished processed= skipped= failed=0 -``` - -After the backfills finish, verify that the services are still running: - -```bash -cd docker -docker compose ps re reports alarms domains auth spicedb -``` - -For non-default deployments, make sure the database and SpiceDB connection settings used by the backfill scripts match your environment before running them. - ---- - ## Usage ```bash diff --git a/alarms/README.md b/alarms/README.md index 226fb39ef..1c9ac19b2 100644 --- a/alarms/README.md +++ b/alarms/README.md @@ -26,16 +26,11 @@ The service is configured using the following environment variables (values show | `MG_MESSAGE_BROKER_URL` | Message broker URL for alarm ingestion | `nats://nats:4222` | | `MG_JAEGER_URL` | Jaeger collector endpoint | `http://jaeger:4318/v1/traces` | | `MG_JAEGER_TRACE_RATIO` | Trace sampling ratio | `1.0` | -| `MG_AUTH_GRPC_URL` | Auth gRPC endpoint | `auth:7001` | -| `MG_AUTH_GRPC_TIMEOUT` | Auth gRPC timeout | `300s` | -| `MG_AUTH_GRPC_CLIENT_CERT` | Auth gRPC client cert path | `${GRPC_MTLS:+./ssl/certs/auth-grpc-client.crt}` | -| `MG_AUTH_GRPC_CLIENT_KEY` | Auth gRPC client key path | `${GRPC_MTLS:+./ssl/certs/auth-grpc-client.key}` | -| `MG_AUTH_GRPC_SERVER_CA_CERTS` | Auth gRPC server CA path | `${GRPC_MTLS:+./ssl/certs/ca.crt}` | -| `MG_DOMAINS_GRPC_URL` | Domains gRPC endpoint | `domains:7003` | -| `MG_DOMAINS_GRPC_TIMEOUT` | Domains gRPC timeout | `300s` | -| `MG_DOMAINS_GRPC_CLIENT_CERT` | Domains gRPC client cert path | `${GRPC_MTLS:+./ssl/certs/domains-grpc-client.crt}` | -| `MG_DOMAINS_GRPC_CLIENT_KEY` | Domains gRPC client key path | `${GRPC_MTLS:+./ssl/certs/domains-grpc-client.key}` | -| `MG_DOMAINS_GRPC_SERVER_CA_CERTS` | Domains gRPC server CA path | `${GRPC_MTLS:+./ssl/certs/ca.crt}` | +| `ATOM_URL` | Atom HTTP endpoint | `http://atom:8080` | +| `ATOM_JWKS_URL` | Atom JWKS endpoint for JWT verification | `http://atom:8080/.well-known/jwks.json` | +| `ATOM_ADMIN_USERNAME` | Atom admin login for service projections | `atom-admin` | +| `ATOM_ADMIN_SECRET` | Atom admin secret for service projections | `change-me` | +| `ATOM_TIMEOUT` | Atom request timeout | `5s` | | `MG_ALLOW_UNVERIFIED_USER` | Allow unverified users to access | `true` | ## Features @@ -44,7 +39,7 @@ The service is configured using the following environment variables (values show - **Stateful updates**: Updates assignee, acknowledgment, resolution, and metadata fields. - **Filtering and paging**: Lists alarms by domain, rule, channel, client, subtopic, status, severity, and time range. - **Observability**: `/metrics` Prometheus endpoint and Jaeger tracing support. -- **Auth and authorization**: Authn/authz enforced via gRPC auth and domains services. +- **Auth and authorization**: Authn/authz enforced through Atom JWT verification and PDP checks. ## Architecture diff --git a/alarms/alarms.go b/alarms/alarms.go index d5f1ac5a3..ff49260e8 100644 --- a/alarms/alarms.go +++ b/alarms/alarms.go @@ -106,7 +106,7 @@ func (a Alarm) Validate() error { // Service specifies an API that must be fulfilled by the domain service. type Service interface { - CreateAlarm(ctx context.Context, alarm Alarm) error + CreateAlarm(ctx context.Context, alarm Alarm) (Alarm, error) UpdateAlarm(ctx context.Context, session authn.Session, alarm Alarm) (Alarm, error) ViewAlarm(ctx context.Context, session authn.Session, id string) (Alarm, error) ListAlarms(ctx context.Context, session authn.Session, pm PageMetadata) (AlarmsPage, error) @@ -118,6 +118,5 @@ type Repository interface { UpdateAlarm(ctx context.Context, alarm Alarm) (Alarm, error) ViewAlarm(ctx context.Context, alarmID, domainID string) (Alarm, error) ListAllAlarms(ctx context.Context, pm PageMetadata) (AlarmsPage, error) - ListUserAlarms(ctx context.Context, userID string, pm PageMetadata) (AlarmsPage, error) DeleteAlarm(ctx context.Context, id string) error } diff --git a/alarms/atom.go b/alarms/atom.go new file mode 100644 index 000000000..53add748f --- /dev/null +++ b/alarms/atom.go @@ -0,0 +1,97 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package alarms + +import ( + "context" + "time" + + "github.com/absmach/magistrala/internal/atom" + "github.com/absmach/magistrala/pkg/authn" +) + +type atomService struct { + Service + projector atom.Projector +} + +func WithAtom(svc Service, projector atom.Projector) Service { + if projector == nil { + return svc + } + return atomService{Service: svc, projector: projector} +} + +func (svc atomService) CreateAlarm(ctx context.Context, alarm Alarm) (Alarm, error) { + created, err := svc.Service.CreateAlarm(ctx, alarm) + if err != nil { + return created, err + } + if created.ID == "" { + return created, nil + } + if err := svc.projector.UpsertResource(ctx, alarmProjection(created)); err != nil { + return created, nil + } + return created, nil +} + +func (svc atomService) UpdateAlarm(ctx context.Context, session authn.Session, alarm Alarm) (Alarm, error) { + updated, err := svc.Service.UpdateAlarm(ctx, session, alarm) + if err != nil { + return updated, err + } + if err := svc.projector.UpsertResource(ctx, alarmProjection(updated)); err != nil { + return updated, nil + } + return updated, nil +} + +func (svc atomService) DeleteAlarm(ctx context.Context, session authn.Session, id string) error { + if err := svc.Service.DeleteAlarm(ctx, session, id); err != nil { + return err + } + _ = svc.projector.DeleteResource(ctx, id) + return nil +} + +func alarmProjection(a Alarm) atom.Resource { + res := atom.ResourceFromFields(atom.ObjectFields{ + ID: a.ID, + Kind: atom.KindAlarm, + Name: a.Cause, + TenantID: a.DomainID, + OwnerID: a.AssigneeID, + Status: a.Status.String(), + Metadata: map[string]any(a.Metadata), + UpdatedBy: a.UpdatedBy, + CreatedAt: a.CreatedAt, + UpdatedAt: a.UpdatedAt, + }) + res.Attributes["rule_id"] = a.RuleID + res.Attributes["channel_id"] = a.ChannelID + res.Attributes["client_id"] = a.ClientID + res.Attributes["subtopic"] = a.Subtopic + res.Attributes["severity"] = a.Severity + res.Attributes["measurement"] = a.Measurement + res.Attributes["value"] = a.Value + res.Attributes["unit"] = a.Unit + res.Attributes["threshold"] = a.Threshold + res.Attributes["cause"] = a.Cause + res.Attributes["assignee_id"] = a.AssigneeID + res.Attributes["assigned_at"] = alarmTimeString(a.AssignedAt) + res.Attributes["assigned_by"] = a.AssignedBy + res.Attributes["acknowledged_at"] = alarmTimeString(a.AcknowledgedAt) + res.Attributes["acknowledged_by"] = a.AcknowledgedBy + res.Attributes["resolved_at"] = alarmTimeString(a.ResolvedAt) + res.Attributes["resolved_by"] = a.ResolvedBy + return res +} + +func alarmTimeString(ts time.Time) string { + if ts.IsZero() { + return "" + } + return ts.Format(time.RFC3339Nano) +} diff --git a/alarms/atom_test.go b/alarms/atom_test.go new file mode 100644 index 000000000..426cd099f --- /dev/null +++ b/alarms/atom_test.go @@ -0,0 +1,83 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package alarms + +import ( + "context" + "testing" + + "github.com/absmach/magistrala/internal/atom" + "github.com/absmach/magistrala/pkg/authn" +) + +func TestAtomServiceCreateAlarmProjectsCreatedAlarm(t *testing.T) { + projector := &alarmProjector{} + svc := WithAtom(alarmService{ + create: Alarm{ + ID: "alarm-1", + RuleID: "rule-1", + DomainID: "domain-1", + ChannelID: "channel-1", + ClientID: "client-1", + Cause: "high temperature", + Measurement: "temperature", + Value: "92.4", + Unit: "C", + Threshold: "80", + Severity: 90, + Status: ActiveStatus, + }, + }, projector) + + created, err := svc.CreateAlarm(context.Background(), Alarm{RuleID: "rule-1"}) + if err != nil { + t.Fatalf("create alarm: %v", err) + } + if created.ID != "alarm-1" { + t.Fatalf("unexpected created alarm: %#v", created) + } + if projector.resource.ID != "alarm-1" || projector.resource.Kind != atom.KindAlarm { + t.Fatalf("unexpected projection: %#v", projector.resource) + } + if projector.resource.Attributes["rule_id"] != "rule-1" { + t.Fatalf("missing rule projection: %#v", projector.resource.Attributes) + } + if projector.resource.Attributes["value"] != "92.4" || projector.resource.Attributes["threshold"] != "80" { + t.Fatalf("missing alarm value projection: %#v", projector.resource.Attributes) + } +} + +type alarmService struct { + create Alarm +} + +func (svc alarmService) CreateAlarm(context.Context, Alarm) (Alarm, error) { + return svc.create, nil +} + +func (svc alarmService) UpdateAlarm(context.Context, authn.Session, Alarm) (Alarm, error) { + return Alarm{}, nil +} + +func (svc alarmService) ViewAlarm(context.Context, authn.Session, string) (Alarm, error) { + return Alarm{}, nil +} + +func (svc alarmService) ListAlarms(context.Context, authn.Session, PageMetadata) (AlarmsPage, error) { + return AlarmsPage{}, nil +} + +func (svc alarmService) DeleteAlarm(context.Context, authn.Session, string) error { + return nil +} + +type alarmProjector struct { + atom.Projector + resource atom.Resource +} + +func (p *alarmProjector) UpsertResource(_ context.Context, resource atom.Resource) error { + p.resource = resource + return nil +} diff --git a/alarms/consumer/consumer.go b/alarms/consumer/consumer.go index 42dd829f0..216e39e61 100644 --- a/alarms/consumer/consumer.go +++ b/alarms/consumer/consumer.go @@ -48,7 +48,8 @@ func (h handler) Handle(msg *messaging.Message) (err error) { return err } - return h.svc.CreateAlarm(context.Background(), alarm) + _, err = h.svc.CreateAlarm(context.Background(), alarm) + return err } func (h handler) Cancel() error { diff --git a/alarms/middleware/authorization.go b/alarms/middleware/authorization.go index cb725660e..f63f76ffa 100644 --- a/alarms/middleware/authorization.go +++ b/alarms/middleware/authorization.go @@ -9,6 +9,7 @@ import ( "github.com/absmach/magistrala/alarms" "github.com/absmach/magistrala/alarms/operations" "github.com/absmach/magistrala/auth" + "github.com/absmach/magistrala/internal/atom" "github.com/absmach/magistrala/pkg/authn" smqauthz "github.com/absmach/magistrala/pkg/authz" "github.com/absmach/magistrala/pkg/errors" @@ -26,6 +27,7 @@ var ( type authorizationMiddleware struct { svc alarms.Service authz smqauthz.Authorization + atomAuthz atom.Authorizer entitiesOps permissions.EntitiesOperations[permissions.Operation] } @@ -43,7 +45,19 @@ func NewAuthorizationMiddleware(svc alarms.Service, authz smqauthz.Authorization }, nil } -func (am *authorizationMiddleware) CreateAlarm(ctx context.Context, alarm alarms.Alarm) error { +func NewAtomAuthorizationMiddleware(svc alarms.Service, authz atom.Authorizer, entitiesOps permissions.EntitiesOperations[permissions.Operation]) (alarms.Service, error) { + if err := entitiesOps.Validate(); err != nil { + return nil, err + } + + return &authorizationMiddleware{ + svc: svc, + atomAuthz: authz, + entitiesOps: entitiesOps, + }, nil +} + +func (am *authorizationMiddleware) CreateAlarm(ctx context.Context, alarm alarms.Alarm) (alarms.Alarm, error) { return am.svc.CreateAlarm(ctx, alarm) } @@ -58,17 +72,19 @@ func (am *authorizationMiddleware) UpdateAlarm(ctx context.Context, session auth if err := am.authorize(ctx, operations.OpAssignAlarm, session, policies.DomainType, session.DomainID); err != nil { return alarms.Alarm{}, errors.Wrap(errDomainUpdateAlarms, err) } - domainUserID := auth.EncodeDomainUserID(session.DomainID, alarm.AssigneeID) - if err := am.authz.Authorize(ctx, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - SubjectKind: policies.UsersKind, - Subject: domainUserID, - Permission: policies.MembershipPermission, - ObjectType: policies.DomainType, - Object: session.DomainID, - }, nil); err != nil { - return alarms.Alarm{}, err + if am.atomAuthz == nil { + domainUserID := auth.EncodeDomainUserID(session.DomainID, alarm.AssigneeID) + if err := am.authz.Authorize(ctx, smqauthz.PolicyReq{ + Domain: session.DomainID, + SubjectType: policies.UserType, + SubjectKind: policies.UsersKind, + Subject: domainUserID, + Permission: policies.MembershipPermission, + ObjectType: policies.DomainType, + Object: session.DomainID, + }, nil); err != nil { + return alarms.Alarm{}, err + } } } @@ -104,6 +120,9 @@ func (am *authorizationMiddleware) ListAlarms(ctx context.Context, session authn case err == nil: session.SuperAdmin = true case errors.Contains(err, svcerr.ErrSuperAdminAction): + if err := am.authorize(ctx, operations.OpListAlarms, session, operations.EntityType, auth.AnyIDs); err != nil { + return alarms.AlarmsPage{}, errors.Wrap(errDomainViewAlarms, err) + } default: return alarms.AlarmsPage{}, err } @@ -124,6 +143,9 @@ func (am *authorizationMiddleware) authorize(ctx context.Context, op permissions if err != nil { return err } + if am.atomAuthz != nil { + return atom.Authorize(ctx, am.atomAuthz, session, perm.String(), objType, obj, atom.KindAlarm) + } pr := smqauthz.PolicyReq{ Domain: session.DomainID, @@ -159,6 +181,9 @@ func (am *authorizationMiddleware) checkSuperAdmin(ctx context.Context, session if session.Role != authn.SuperAdminRole { return svcerr.ErrSuperAdminAction } + if am.atomAuthz != nil { + return atom.Authorize(ctx, am.atomAuthz, session, policies.AdminPermission, policies.PlatformType, policies.MagistralaObject, policies.PlatformType) + } if err := am.authz.Authorize(ctx, smqauthz.PolicyReq{ SubjectType: policies.UserType, Subject: session.UserID, diff --git a/alarms/middleware/authorization_test.go b/alarms/middleware/authorization_test.go new file mode 100644 index 000000000..b21d9fc67 --- /dev/null +++ b/alarms/middleware/authorization_test.go @@ -0,0 +1,109 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package middleware + +import ( + "context" + "testing" + + "github.com/absmach/magistrala/alarms" + "github.com/absmach/magistrala/alarms/mocks" + "github.com/absmach/magistrala/alarms/operations" + "github.com/absmach/magistrala/auth" + "github.com/absmach/magistrala/internal/atom" + "github.com/absmach/magistrala/pkg/authn" + pkgerrors "github.com/absmach/magistrala/pkg/errors" + "github.com/absmach/magistrala/pkg/permissions" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/mock" + "github.com/stretchr/testify/require" +) + +type recordingAtomAuthorizer struct { + allowed bool + reqs []atom.AuthzRequest +} + +func (a *recordingAtomAuthorizer) CheckAuthz(_ context.Context, req atom.AuthzRequest) (atom.AuthzResponse, error) { + a.reqs = append(a.reqs, req) + return atom.AuthzResponse{Allowed: a.allowed}, nil +} + +func TestListAlarmsAuthorizesRegularUser(t *testing.T) { + svc := mocks.NewService(t) + pm := alarms.PageMetadata{Limit: 10} + expectedPM := pm + expectedPM.DomainID = "domain-1" + session := authn.Session{UserID: "user-1", DomainID: "domain-1", DomainUserID: "domain-1_user-1"} + authz := &recordingAtomAuthorizer{allowed: true} + wrapped, err := NewAtomAuthorizationMiddleware(svc, authz, testEntitiesOps(t)) + require.NoError(t, err) + + svc.On("ListAlarms", mock.Anything, session, expectedPM).Return(alarms.AlarmsPage{Limit: 10}, nil).Once() + page, err := wrapped.ListAlarms(context.Background(), session, pm) + + require.NoError(t, err) + assert.Equal(t, uint64(10), page.Limit) + require.Len(t, authz.reqs, 1) + assert.Equal(t, atom.AuthzRequest{ + SubjectID: "user-1", + Action: "list", + ResourceID: auth.AnyIDs, + ObjectKind: "resource", + ObjectID: auth.AnyIDs, + Context: map[string]any{ + "domain_id": "domain-1", + "legacy_object_type": operations.EntityType, + }, + }, authz.reqs[0]) +} + +func TestListAlarmsDeniedRegularUserDoesNotDelegate(t *testing.T) { + svc := mocks.NewService(t) + authz := &recordingAtomAuthorizer{allowed: false} + wrapped, err := NewAtomAuthorizationMiddleware(svc, authz, testEntitiesOps(t)) + require.NoError(t, err) + + _, err = wrapped.ListAlarms(context.Background(), authn.Session{UserID: "user-1", DomainID: "domain-1"}, alarms.PageMetadata{}) + + assert.True(t, pkgerrors.Contains(err, pkgerrors.ErrAuthorization)) + require.Len(t, authz.reqs, 1) +} + +func TestListAlarmsSuperAdminSkipsListAuthorization(t *testing.T) { + svc := mocks.NewService(t) + pm := alarms.PageMetadata{Limit: 10} + expectedPM := pm + expectedPM.DomainID = "domain-1" + session := authn.Session{UserID: "admin-1", DomainID: "domain-1", Role: authn.SuperAdminRole} + authz := &recordingAtomAuthorizer{allowed: true} + wrapped, err := NewAtomAuthorizationMiddleware(svc, authz, testEntitiesOps(t)) + require.NoError(t, err) + + svc.On("ListAlarms", mock.Anything, mock.MatchedBy(func(s authn.Session) bool { + return s.SuperAdmin + }), expectedPM).Return(alarms.AlarmsPage{Limit: 10}, nil).Once() + _, err = wrapped.ListAlarms(context.Background(), session, pm) + + require.NoError(t, err) + require.Len(t, authz.reqs, 1) + assert.Equal(t, "manage", authz.reqs[0].Action) +} + +func testEntitiesOps(t *testing.T) permissions.EntitiesOperations[permissions.Operation] { + t.Helper() + details := operations.OperationDetails() + perms := make(map[string]permissions.Permission, len(details)) + for _, detail := range details { + if detail.PermissionRequired { + perms[detail.Name] = permissions.Permission(detail.Name) + } + } + entitiesOps, err := permissions.NewEntitiesOperations( + permissions.EntitiesPermission{operations.EntityType: perms}, + permissions.EntitiesOperationDetails[permissions.Operation]{operations.EntityType: details}, + ) + require.NoError(t, err) + return entitiesOps +} diff --git a/alarms/middleware/logging.go b/alarms/middleware/logging.go index d3375e639..0f47f57c9 100644 --- a/alarms/middleware/logging.go +++ b/alarms/middleware/logging.go @@ -27,7 +27,7 @@ func NewLoggingMiddleware(logger *slog.Logger, service alarms.Service) alarms.Se } } -func (lm *loggingMiddleware) CreateAlarm(ctx context.Context, alarm alarms.Alarm) (err error) { +func (lm *loggingMiddleware) CreateAlarm(ctx context.Context, alarm alarms.Alarm) (created alarms.Alarm, err error) { defer func(begin time.Time) { args := []any{ slog.String("duration", time.Since(begin).String()), @@ -52,7 +52,7 @@ func (lm *loggingMiddleware) CreateAlarm(ctx context.Context, alarm alarms.Alarm lm.logger.Warn("Create alarm failed", args...) return } - if alarm.ID != "" { + if created.ID != "" { lm.logger.Info("Create alarm completed successfully", args...) } }(time.Now()) diff --git a/alarms/middleware/metrics.go b/alarms/middleware/metrics.go index 07ff5961e..cacb8e12e 100644 --- a/alarms/middleware/metrics.go +++ b/alarms/middleware/metrics.go @@ -28,7 +28,7 @@ func NewMetricsMiddleware(counter metrics.Counter, latency metrics.Histogram, se } } -func (mm *metricsMiddleware) CreateAlarm(ctx context.Context, alarm alarms.Alarm) error { +func (mm *metricsMiddleware) CreateAlarm(ctx context.Context, alarm alarms.Alarm) (alarms.Alarm, error) { defer func(begin time.Time) { mm.counter.With("method", "create_alarm").Add(1) mm.latency.With("method", "create_alarm").Observe(time.Since(begin).Seconds()) diff --git a/alarms/middleware/tracing.go b/alarms/middleware/tracing.go index 930bb2637..a6b6a19d3 100644 --- a/alarms/middleware/tracing.go +++ b/alarms/middleware/tracing.go @@ -27,7 +27,7 @@ func NewTracingMiddleware(tracer trace.Tracer, svc alarms.Service) alarms.Servic } } -func (tm *tracingMiddleware) CreateAlarm(ctx context.Context, alarm alarms.Alarm) error { +func (tm *tracingMiddleware) CreateAlarm(ctx context.Context, alarm alarms.Alarm) (alarms.Alarm, error) { ctx, span := smqTracing.StartSpan(ctx, tm.tracer, "create_alarm", trace.WithAttributes( attribute.String("rule_id", alarm.RuleID), attribute.String("measurement", alarm.Measurement), diff --git a/alarms/mocks/repository.go b/alarms/mocks/repository.go index c5fecb0fa..f44c3c6ff 100644 --- a/alarms/mocks/repository.go +++ b/alarms/mocks/repository.go @@ -231,78 +231,6 @@ func (_c *Repository_ListAllAlarms_Call) RunAndReturn(run func(ctx context.Conte return _c } -// ListUserAlarms provides a mock function for the type Repository -func (_mock *Repository) ListUserAlarms(ctx context.Context, userID string, pm alarms.PageMetadata) (alarms.AlarmsPage, error) { - ret := _mock.Called(ctx, userID, pm) - - if len(ret) == 0 { - panic("no return value specified for ListUserAlarms") - } - - var r0 alarms.AlarmsPage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, alarms.PageMetadata) (alarms.AlarmsPage, error)); ok { - return returnFunc(ctx, userID, pm) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, alarms.PageMetadata) alarms.AlarmsPage); ok { - r0 = returnFunc(ctx, userID, pm) - } else { - r0 = ret.Get(0).(alarms.AlarmsPage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, alarms.PageMetadata) error); ok { - r1 = returnFunc(ctx, userID, pm) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_ListUserAlarms_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ListUserAlarms' -type Repository_ListUserAlarms_Call struct { - *mock.Call -} - -// ListUserAlarms is a helper method to define mock.On call -// - ctx context.Context -// - userID string -// - pm alarms.PageMetadata -func (_e *Repository_Expecter) ListUserAlarms(ctx interface{}, userID interface{}, pm interface{}) *Repository_ListUserAlarms_Call { - return &Repository_ListUserAlarms_Call{Call: _e.mock.On("ListUserAlarms", ctx, userID, pm)} -} - -func (_c *Repository_ListUserAlarms_Call) Run(run func(ctx context.Context, userID string, pm alarms.PageMetadata)) *Repository_ListUserAlarms_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 alarms.PageMetadata - if args[2] != nil { - arg2 = args[2].(alarms.PageMetadata) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_ListUserAlarms_Call) Return(alarmsPage alarms.AlarmsPage, err error) *Repository_ListUserAlarms_Call { - _c.Call.Return(alarmsPage, err) - return _c -} - -func (_c *Repository_ListUserAlarms_Call) RunAndReturn(run func(ctx context.Context, userID string, pm alarms.PageMetadata) (alarms.AlarmsPage, error)) *Repository_ListUserAlarms_Call { - _c.Call.Return(run) - return _c -} - // UpdateAlarm provides a mock function for the type Repository func (_mock *Repository) UpdateAlarm(ctx context.Context, alarm alarms.Alarm) (alarms.Alarm, error) { ret := _mock.Called(ctx, alarm) diff --git a/alarms/mocks/service.go b/alarms/mocks/service.go index 89c24fdad..0d6df07b4 100644 --- a/alarms/mocks/service.go +++ b/alarms/mocks/service.go @@ -44,20 +44,29 @@ func (_m *Service) EXPECT() *Service_Expecter { } // CreateAlarm provides a mock function for the type Service -func (_mock *Service) CreateAlarm(ctx context.Context, alarm alarms.Alarm) error { +func (_mock *Service) CreateAlarm(ctx context.Context, alarm alarms.Alarm) (alarms.Alarm, error) { ret := _mock.Called(ctx, alarm) if len(ret) == 0 { panic("no return value specified for CreateAlarm") } - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, alarms.Alarm) error); ok { + var r0 alarms.Alarm + var r1 error + if returnFunc, ok := ret.Get(0).(func(context.Context, alarms.Alarm) (alarms.Alarm, error)); ok { + return returnFunc(ctx, alarm) + } + if returnFunc, ok := ret.Get(0).(func(context.Context, alarms.Alarm) alarms.Alarm); ok { r0 = returnFunc(ctx, alarm) } else { - r0 = ret.Error(0) + r0 = ret.Get(0).(alarms.Alarm) } - return r0 + if returnFunc, ok := ret.Get(1).(func(context.Context, alarms.Alarm) error); ok { + r1 = returnFunc(ctx, alarm) + } else { + r1 = ret.Error(1) + } + return r0, r1 } // Service_CreateAlarm_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'CreateAlarm' @@ -90,12 +99,12 @@ func (_c *Service_CreateAlarm_Call) Run(run func(ctx context.Context, alarm alar return _c } -func (_c *Service_CreateAlarm_Call) Return(err error) *Service_CreateAlarm_Call { - _c.Call.Return(err) +func (_c *Service_CreateAlarm_Call) Return(alarm1 alarms.Alarm, err error) *Service_CreateAlarm_Call { + _c.Call.Return(alarm1, err) return _c } -func (_c *Service_CreateAlarm_Call) RunAndReturn(run func(ctx context.Context, alarm alarms.Alarm) error) *Service_CreateAlarm_Call { +func (_c *Service_CreateAlarm_Call) RunAndReturn(run func(ctx context.Context, alarm alarms.Alarm) (alarms.Alarm, error)) *Service_CreateAlarm_Call { _c.Call.Return(run) return _c } diff --git a/alarms/postgres/alarms.go b/alarms/postgres/alarms.go index 37433058f..41fd2e029 100644 --- a/alarms/postgres/alarms.go +++ b/alarms/postgres/alarms.go @@ -198,35 +198,6 @@ func (r *repository) ListAllAlarms(ctx context.Context, pm alarms.PageMetadata) return r.alarmsPage(ctx, comQuery, pm) } -func (r *repository) ListUserAlarms(ctx context.Context, userID string, pm alarms.PageMetadata) (alarms.AlarmsPage, error) { - clauses := []string{ - `( - EXISTS ( - SELECT 1 - FROM rules_roles rr - JOIN rules_role_members rrm ON rrm.role_id = rr.id - WHERE rr.entity_id = alarms.rule_id AND rrm.member_id = :user_id - ) - OR EXISTS ( - SELECT 1 - FROM domains_roles dr - JOIN domains_role_members drm ON drm.role_id = dr.id - JOIN domains_role_actions dra ON dra.role_id = dr.id - WHERE dr.entity_id = alarms.domain_id - AND drm.member_id = :user_id - AND dra.action LIKE 'alarm%' - ) - )`, - } - - clauses = append(clauses, pageQueryConditions(pm)...) - query := fmt.Sprintf("WHERE %s", strings.Join(clauses, " AND ")) - pm.UserID = userID - comQuery := fmt.Sprintf(`SELECT DISTINCT %s FROM alarms %s`, alarmColumns, query) - - return r.alarmsPage(ctx, comQuery, pm) -} - func (r *repository) alarmsPage(ctx context.Context, comQuery string, pm alarms.PageMetadata) (alarms.AlarmsPage, error) { dir := api.DescDir if pm.Dir == api.AscDir { diff --git a/alarms/postgres/alarms_test.go b/alarms/postgres/alarms_test.go index 48cc7fcab..d663fc7df 100644 --- a/alarms/postgres/alarms_test.go +++ b/alarms/postgres/alarms_test.go @@ -415,215 +415,6 @@ func TestListAlarms(t *testing.T) { } } -func TestListUserAlarms(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM domains_role_actions") - require.Nil(t, err, fmt.Sprintf("clean domains_role_actions unexpected error: %s", err)) - _, err = db.Exec("DELETE FROM domains_role_members") - require.Nil(t, err, fmt.Sprintf("clean domains_role_members unexpected error: %s", err)) - _, err = db.Exec("DELETE FROM domains_roles") - require.Nil(t, err, fmt.Sprintf("clean domains_roles unexpected error: %s", err)) - _, err = db.Exec("DELETE FROM domains") - require.Nil(t, err, fmt.Sprintf("clean domains unexpected error: %s", err)) - _, err = db.Exec("DELETE FROM alarms") - require.Nil(t, err, fmt.Sprintf("clean alarms unexpected error: %s", err)) - _, err = db.Exec("DELETE FROM rules") - require.Nil(t, err, fmt.Sprintf("clean rules unexpected error: %s", err)) - }) - - repo := postgres.NewAlarmsRepo(db) - - domainID := generateUUID(t) - domainRoute := generateUUID(t) - userID := generateUUID(t) - otherUserID := generateUUID(t) - adminUserID := generateUUID(t) - domainUserID := generateUUID(t) - - _, err := db.Exec(`INSERT INTO domains (id, name, route, status) VALUES ($1, $2, $3, $4)`, domainID, namegen.Generate(), domainRoute, 0) - require.Nil(t, err, fmt.Sprintf("insert domains unexpected error: %s", err)) - - // Create 10 rules and 10 alarms referencing them. - // Assign userID to the first 6 rules via role membership. - var ruleIDs []string - var createdAlarms []alarms.Alarm - for i := range 10 { - ruleID := generateUUID(t) - _, err := db.Exec(`INSERT INTO rules (id, name, domain_id, status, logic_type, logic_value) VALUES ($1, $2, $3, 0, 0, '')`, - ruleID, fmt.Sprintf("rule-%d", i), domainID) - require.Nil(t, err, fmt.Sprintf("insert rule unexpected error: %s", err)) - ruleIDs = append(ruleIDs, ruleID) - - alarm := alarms.Alarm{ - ID: generateUUID(t), - RuleID: ruleID, - DomainID: domainID, - ChannelID: generateUUID(t), - ClientID: generateUUID(t), - Measurement: namegen.Generate(), - Value: namegen.Generate(), - Unit: namegen.Generate(), - Threshold: namegen.Generate(), - Cause: namegen.Generate(), - Status: 0, - AssigneeID: generateUUID(t), - CreatedAt: time.Now().UTC().Add(time.Duration(i) * time.Minute), - } - alarm, err = repo.CreateAlarm(context.Background(), alarm) - require.Nil(t, err, fmt.Sprintf("unexpected error: %s", err)) - createdAlarms = append(createdAlarms, alarm) - } - - // Assign userID to the first 6 rules via rules_roles + rules_role_members. - userRoleIDs := make([]string, 6) - for i := range 6 { - roleID := generateUUID(t) - userRoleIDs[i] = roleID - _, err := db.Exec(`INSERT INTO rules_roles (id, name, entity_id) VALUES ($1, $2, $3)`, roleID, "admin", ruleIDs[i]) - require.Nil(t, err, fmt.Sprintf("insert rules_roles unexpected error: %s", err)) - _, err = db.Exec(`INSERT INTO rules_role_members (role_id, member_id, entity_id) VALUES ($1, $2, $3)`, roleID, userID, ruleIDs[i]) - require.Nil(t, err, fmt.Sprintf("insert rules_role_members unexpected error: %s", err)) - } - - for i := range 10 { - var roleID string - if i < 6 { - roleID = userRoleIDs[i] - } else { - roleID = generateUUID(t) - _, err := db.Exec(`INSERT INTO rules_roles (id, name, entity_id) VALUES ($1, $2, $3)`, roleID, "admin", ruleIDs[i]) - require.Nil(t, err, fmt.Sprintf("insert rules_roles unexpected error: %s", err)) - } - _, err := db.Exec(`INSERT INTO rules_role_members (role_id, member_id, entity_id) VALUES ($1, $2, $3)`, roleID, adminUserID, ruleIDs[i]) - require.Nil(t, err, fmt.Sprintf("insert rules_role_members unexpected error: %s", err)) - } - - domainRoleID := generateUUID(t) - _, err = db.Exec(`INSERT INTO domains_roles (id, name, entity_id) VALUES ($1, $2, $3)`, domainRoleID, "admin", domainID) - require.Nil(t, err, fmt.Sprintf("insert domains_roles unexpected error: %s", err)) - _, err = db.Exec(`INSERT INTO domains_role_members (role_id, member_id, entity_id) VALUES ($1, $2, $3)`, domainRoleID, domainUserID, domainID) - require.Nil(t, err, fmt.Sprintf("insert domains_role_members unexpected error: %s", err)) - _, err = db.Exec(`INSERT INTO domains_role_actions (role_id, action) VALUES ($1, $2)`, domainRoleID, "alarm_read") - require.Nil(t, err, fmt.Sprintf("insert domains_role_actions unexpected error: %s", err)) - - _ = createdAlarms - - cases := []struct { - desc string - userID string - pm alarms.PageMetadata - count int - err error - }{ - { - desc: "list user alarms returns only accessible alarms", - userID: userID, - pm: alarms.PageMetadata{ - Offset: 0, - Limit: 100, - }, - count: 6, - err: nil, - }, - { - desc: "list user alarms with limit", - userID: userID, - pm: alarms.PageMetadata{ - Offset: 0, - Limit: 3, - }, - count: 3, - err: nil, - }, - { - desc: "list user alarms with offset", - userID: userID, - pm: alarms.PageMetadata{ - Offset: 4, - Limit: 100, - }, - count: 2, - err: nil, - }, - { - desc: "list user alarms with domain filter", - userID: userID, - pm: alarms.PageMetadata{ - DomainID: domainID, - Offset: 0, - Limit: 100, - }, - count: 6, - err: nil, - }, - { - desc: "list user alarms with non-existing domain returns 0", - userID: userID, - pm: alarms.PageMetadata{ - DomainID: generateUUID(t), - Offset: 0, - Limit: 100, - }, - count: 0, - err: nil, - }, - { - desc: "list alarms for user with no role assignments returns 0", - userID: otherUserID, - pm: alarms.PageMetadata{ - Offset: 0, - Limit: 100, - }, - count: 0, - err: nil, - }, - { - desc: "list alarms for admin user with role on all rules returns all alarms", - userID: adminUserID, - pm: alarms.PageMetadata{ - Offset: 0, - Limit: 100, - }, - count: 10, - err: nil, - }, - { - desc: "list alarms for user with domain-level rule access returns all alarms", - userID: domainUserID, - pm: alarms.PageMetadata{ - Offset: 0, - Limit: 100, - }, - count: 10, - err: nil, - }, - { - desc: "list user alarms ordered by created_at ascending", - userID: userID, - pm: alarms.PageMetadata{ - Offset: 0, - Limit: 100, - Order: "created_at", - Dir: "asc", - }, - count: 6, - err: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - page, err := repo.ListUserAlarms(context.Background(), tc.userID, tc.pm) - if tc.err != nil { - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - return - } - require.Nil(t, err, fmt.Sprintf("unexpected error: %s", err)) - assert.Equal(t, tc.count, len(page.Alarms), fmt.Sprintf("%s: expected %d alarms, got %d", tc.desc, tc.count, len(page.Alarms))) - }) - } -} - func TestDeleteAlarm(t *testing.T) { t.Cleanup(func() { _, err := db.Exec("DELETE FROM alarms") diff --git a/alarms/postgres/init.go b/alarms/postgres/init.go index 66845220f..93f217531 100644 --- a/alarms/postgres/init.go +++ b/alarms/postgres/init.go @@ -4,9 +4,6 @@ package postgres import ( - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - rpostgres "github.com/absmach/magistrala/re/postgres" _ "github.com/jackc/pgx/v5/stdlib" // required for SQL access migrate "github.com/rubenv/sql-migrate" ) @@ -54,12 +51,5 @@ func Migration() (*migrate.MemoryMigrationSource, error) { }, } - rulesMigration, err := rpostgres.Migration() - if err != nil { - return &migrate.MemoryMigrationSource{}, errors.Wrap(repoerr.ErrRoleMigration, err) - } - - alarmsMigration.Migrations = append(alarmsMigration.Migrations, rulesMigration.Migrations...) - return alarmsMigration, nil } diff --git a/alarms/service.go b/alarms/service.go index 219e78ad4..9fd8c41d7 100644 --- a/alarms/service.go +++ b/alarms/service.go @@ -26,10 +26,10 @@ func NewService(idp magistrala.IDProvider, repo Repository) Service { } } -func (s *service) CreateAlarm(ctx context.Context, alarm Alarm) error { +func (s *service) CreateAlarm(ctx context.Context, alarm Alarm) (Alarm, error) { id, err := s.idp.ID() if err != nil { - return err + return Alarm{}, err } alarm.ID = id if alarm.CreatedAt.IsZero() { @@ -37,14 +37,18 @@ func (s *service) CreateAlarm(ctx context.Context, alarm Alarm) error { } if err := alarm.Validate(); err != nil { - return err + return Alarm{}, err } - if _, err = s.repo.CreateAlarm(ctx, alarm); err != nil && err != repoerr.ErrNotFound { - return err + created, err := s.repo.CreateAlarm(ctx, alarm) + if err != nil && err != repoerr.ErrNotFound { + return Alarm{}, err + } + if err == repoerr.ErrNotFound { + return Alarm{}, nil } - return nil + return created, nil } func (s *service) ViewAlarm(ctx context.Context, session authn.Session, alarmID string) (Alarm, error) { @@ -52,10 +56,8 @@ func (s *service) ViewAlarm(ctx context.Context, session authn.Session, alarmID } func (s *service) ListAlarms(ctx context.Context, session authn.Session, pm PageMetadata) (AlarmsPage, error) { - if session.SuperAdmin { - return s.repo.ListAllAlarms(ctx, pm) - } - return s.repo.ListUserAlarms(ctx, session.UserID, pm) + pm.DomainID = session.DomainID + return s.repo.ListAllAlarms(ctx, pm) } func (s *service) DeleteAlarm(ctx context.Context, session authn.Session, alarmID string) error { diff --git a/alarms/service_test.go b/alarms/service_test.go index 834058d34..f10a735fe 100644 --- a/alarms/service_test.go +++ b/alarms/service_test.go @@ -72,7 +72,7 @@ func TestCreateAlarm(t *testing.T) { for _, tc := range cases { t.Run(tc.desc, func(t *testing.T) { repoCall := repo.On("CreateAlarm", context.Background(), mock.Anything).Return(tc.alarm, tc.err) - err := svc.CreateAlarm(context.Background(), tc.alarm) + _, err := svc.CreateAlarm(context.Background(), tc.alarm) assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) repoCall.Unset() }) @@ -205,7 +205,7 @@ func TestListAlarms(t *testing.T) { for _, tc := range cases { t.Run(tc.desc, func(t *testing.T) { s := authn.Session{DomainID: tc.pm.DomainID} - repoCall := repo.On("ListUserAlarms", context.Background(), s.UserID, tc.pm).Return(tc.page, tc.err) + repoCall := repo.On("ListAllAlarms", context.Background(), tc.pm).Return(tc.page, tc.err) _, err := svc.ListAlarms(context.Background(), s, tc.pm) if tc.err != nil { assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) diff --git a/api/http/common.go b/api/http/common.go index a069b7f9c..5319920e5 100644 --- a/api/http/common.go +++ b/api/http/common.go @@ -13,10 +13,7 @@ import ( "github.com/absmach/magistrala" apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/clients" - "github.com/absmach/magistrala/groups" "github.com/absmach/magistrala/pkg/errors" - "github.com/absmach/magistrala/users" "github.com/gofrs/uuid/v5" ) @@ -80,9 +77,9 @@ const ( DefStartLevel = 1 DefEndLevel = 0 DefStatus = "enabled" - DefClientStatus = clients.Enabled - DefUserStatus = users.Enabled - DefGroupStatus = groups.Enabled + DefClientStatus = "enabled" + DefUserStatus = "enabled" + DefGroupStatus = "enabled" // ContentType represents JSON content type. ContentType = "application/json" @@ -184,6 +181,21 @@ func EncodeError(_ context.Context, err error, w http.ResponseWriter) { return } + if errors.Contains(err, errors.ErrAuthentication) { + w.WriteHeader(http.StatusUnauthorized) + if err := json.NewEncoder(w).Encode(err); err != nil { + w.WriteHeader(http.StatusInternalServerError) + } + return + } + if errors.Contains(err, errors.ErrAuthorization) { + w.WriteHeader(http.StatusForbidden) + if err := json.NewEncoder(w).Encode(err); err != nil { + w.WriteHeader(http.StatusInternalServerError) + } + return + } + switch retErr := err.(type) { case *errors.RequestError: w.WriteHeader(http.StatusBadRequest) diff --git a/api/http/common_test.go b/api/http/common_test.go index 83597b37e..854e80fe6 100644 --- a/api/http/common_test.go +++ b/api/http/common_test.go @@ -260,12 +260,24 @@ func TestEncodeError(t *testing.T) { code: http.StatusUnauthorized, hasBody: true, }, + { + desc: "Generic Authentication Failed", + err: errors.ErrAuthentication, + code: http.StatusUnauthorized, + hasBody: true, + }, { desc: "AuthZError - Authorization Failed", err: svcerr.ErrAuthorization, code: http.StatusForbidden, hasBody: true, }, + { + desc: "Generic Authorization Failed", + err: errors.Wrap(errors.New("not authorized"), errors.ErrAuthorization), + code: http.StatusForbidden, + hasBody: true, + }, { desc: "AuthZError - Domain Authorization Failed", err: svcerr.ErrDomainAuthorization, diff --git a/apidocs/openapi/readers.yaml b/apidocs/openapi/readers.yaml index 7c1debe1b..855e52453 100644 --- a/apidocs/openapi/readers.yaml +++ b/apidocs/openapi/readers.yaml @@ -66,6 +66,8 @@ paths: description: Failed due to malformed query parameters. "401": description: Missing or invalid access token provided. + "403": + description: Failed to perform authorization over the entity. "500": $ref: "#/components/responses/ServiceError" /health: diff --git a/auth/api/grpc/auth/endpoint_test.go b/auth/api/grpc/auth/endpoint_test.go deleted file mode 100644 index 7e982bf4d..000000000 --- a/auth/api/grpc/auth/endpoint_test.go +++ /dev/null @@ -1,401 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package auth_test - -import ( - "context" - "fmt" - "net" - "testing" - "time" - - grpcAuthV1 "github.com/absmach/magistrala/api/grpc/auth/v1" - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/auth" - grpcapi "github.com/absmach/magistrala/auth/api/grpc/auth" - "github.com/absmach/magistrala/internal/testsutil" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/pkg/policies" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" - "google.golang.org/grpc" - "google.golang.org/grpc/credentials/insecure" -) - -const ( - port = 8081 - id = "testID" - usersType = "users" - adminPermission = "admin" - authoritiesObj = "authorities" - memberRelation = "member" - validToken = "valid" - inValidToken = "invalid" - validPATToken = "valid" -) - -var ( - domainID = testsutil.GenerateUUID(&testing.T{}) - authAddr = fmt.Sprintf("localhost:%d", port) - clientID = testsutil.GenerateUUID(&testing.T{}) -) - -func startGRPCServer(svc auth.Service, port int) *grpc.Server { - listener, _ := net.Listen("tcp", fmt.Sprintf(":%d", port)) - server := grpc.NewServer() - grpcAuthV1.RegisterAuthServiceServer(server, grpcapi.NewAuthServer(svc)) - go func() { - err := server.Serve(listener) - assert.Nil(&testing.T{}, err, fmt.Sprintf(`"Unexpected error creating auth server %s"`, err)) - }() - - return server -} - -func TestIdentify(t *testing.T) { - conn, err := grpc.NewClient(authAddr, grpc.WithTransportCredentials(insecure.NewCredentials())) - assert.Nil(t, err, fmt.Sprintf("Unexpected error creating client connection %s", err)) - defer conn.Close() - grpcClient := grpcapi.NewAuthClient(conn, time.Second) - - cases := []struct { - desc string - token string - key auth.Key - idt *grpcAuthV1.AuthNRes - svcErr error - err error - }{ - { - desc: "authenticate user with valid user token", - token: validToken, - key: auth.Key{ID: "", Subject: id, Role: auth.UserRole}, - idt: &grpcAuthV1.AuthNRes{UserId: id, UserRole: uint32(auth.UserRole)}, - err: nil, - }, - { - desc: "authenticate user with invalid user token", - token: "invalid", - key: auth.Key{}, - idt: &grpcAuthV1.AuthNRes{}, - svcErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "authenticate user with empty token", - token: "", - idt: &grpcAuthV1.AuthNRes{}, - err: apiutil.ErrBearerToken, - }, - { - desc: "authenticate user with valid PAT token", - token: "pat_" + validPATToken, - key: auth.Key{ID: id, Type: auth.PersonalAccessToken, Subject: clientID, Role: auth.UserRole}, - idt: &grpcAuthV1.AuthNRes{Id: id, UserId: clientID, UserRole: uint32(auth.UserRole), TokenType: uint32(auth.PersonalAccessToken)}, - err: nil, - }, - { - desc: "authenticate user with invalid PAT token", - token: "pat_invalid", - key: auth.Key{}, - idt: &grpcAuthV1.AuthNRes{}, - svcErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("Identify", mock.Anything, tc.token).Return(tc.key, tc.svcErr) - idt, err := grpcClient.Authenticate(context.Background(), &grpcAuthV1.AuthNReq{Token: tc.token}) - if idt != nil { - assert.Equal(t, tc.idt, idt, fmt.Sprintf("%s: expected %v got %v", tc.desc, tc.idt, idt)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - svcCall.Unset() - }) - } -} - -func TestAuthorize(t *testing.T) { - conn, err := grpc.NewClient(authAddr, grpc.WithTransportCredentials(insecure.NewCredentials())) - assert.Nil(t, err, fmt.Sprintf("Unexpected error creating client connection %s", err)) - defer conn.Close() - - grpcClient := grpcapi.NewAuthClient(conn, time.Second) - - cases := []struct { - desc string - token string - authRequest *grpcAuthV1.AuthZReq - authResponse *grpcAuthV1.AuthZRes - expectedReq *policies.Policy - expectedPAT *auth.PATAuthz - err error - }{ - { - desc: "authorize user with authorized token", - token: validToken, - authRequest: &grpcAuthV1.AuthZReq{ - PolicyReq: &grpcAuthV1.PolicyReq{ - Subject: id, - SubjectType: usersType, - Object: authoritiesObj, - ObjectType: usersType, - Relation: memberRelation, - Permission: adminPermission, - }, - }, - authResponse: &grpcAuthV1.AuthZRes{Authorized: true}, - err: nil, - }, - { - desc: "authorize user with unauthorized token", - token: inValidToken, - authRequest: &grpcAuthV1.AuthZReq{ - PolicyReq: &grpcAuthV1.PolicyReq{ - Subject: id, - SubjectType: usersType, - Object: authoritiesObj, - ObjectType: usersType, - Relation: memberRelation, - Permission: adminPermission, - }, - }, - authResponse: &grpcAuthV1.AuthZRes{Authorized: false}, - err: svcerr.ErrAuthorization, - }, - { - desc: "authorize user with empty subject", - token: validToken, - authRequest: &grpcAuthV1.AuthZReq{ - PolicyReq: &grpcAuthV1.PolicyReq{ - Subject: "", - SubjectType: usersType, - Object: authoritiesObj, - ObjectType: usersType, - Relation: memberRelation, - Permission: adminPermission, - }, - }, - authResponse: &grpcAuthV1.AuthZRes{Authorized: false}, - err: apiutil.ErrMissingPolicySub, - }, - { - desc: "authorize user with empty subject type", - token: validToken, - authRequest: &grpcAuthV1.AuthZReq{ - PolicyReq: &grpcAuthV1.PolicyReq{ - Subject: id, - SubjectType: "", - Object: authoritiesObj, - ObjectType: usersType, - Relation: memberRelation, - Permission: adminPermission, - }, - }, - authResponse: &grpcAuthV1.AuthZRes{Authorized: false}, - err: apiutil.ErrMissingPolicySub, - }, - { - desc: "authorize user with empty object", - token: validToken, - authRequest: &grpcAuthV1.AuthZReq{ - PolicyReq: &grpcAuthV1.PolicyReq{ - Subject: id, - SubjectType: usersType, - Object: "", - ObjectType: usersType, - Relation: memberRelation, - Permission: adminPermission, - }, - }, - authResponse: &grpcAuthV1.AuthZRes{Authorized: false}, - err: apiutil.ErrMissingPolicyObj, - }, - { - desc: "authorize user with empty object type", - token: validToken, - authRequest: &grpcAuthV1.AuthZReq{ - PolicyReq: &grpcAuthV1.PolicyReq{ - Subject: id, - SubjectType: usersType, - Object: authoritiesObj, - ObjectType: "", - Relation: memberRelation, - Permission: adminPermission, - }, - }, - authResponse: &grpcAuthV1.AuthZRes{Authorized: false}, - err: apiutil.ErrMissingPolicyObj, - }, - { - desc: "authorize user with empty permission", - token: validToken, - authRequest: &grpcAuthV1.AuthZReq{ - PolicyReq: &grpcAuthV1.PolicyReq{ - Subject: id, - SubjectType: usersType, - Object: authoritiesObj, - ObjectType: usersType, - Relation: memberRelation, - Permission: "", - }, - }, - authResponse: &grpcAuthV1.AuthZRes{Authorized: false}, - err: apiutil.ErrMalformedPolicyPer, - }, - { - desc: "authorize user with valid PAT token", - token: validPATToken, - authRequest: &grpcAuthV1.AuthZReq{ - PolicyReq: &grpcAuthV1.PolicyReq{ - Subject: id, - SubjectType: policies.UserType, - SubjectKind: policies.UsersKind, - Permission: policies.ViewPermission, - ObjectType: policies.ClientType, - Domain: domainID, - Object: clientID, - }, - PatReq: &grpcAuthV1.PATReq{ - PatId: id, - Domain: domainID, - Operation: "view", - UserId: id, - EntityId: clientID, - EntityType: auth.ClientsScopeStr, - }, - }, - authResponse: &grpcAuthV1.AuthZRes{Authorized: true}, - err: nil, - }, - { - desc: "authorize bootstrap PAT keeps PAT domain when policy domain is empty", - token: validPATToken, - authRequest: &grpcAuthV1.AuthZReq{ - PolicyReq: &grpcAuthV1.PolicyReq{ - Subject: id, - SubjectType: policies.UserType, - SubjectKind: policies.UsersKind, - Permission: policies.MembershipPermission, - ObjectType: policies.DomainType, - Object: domainID, - }, - PatReq: &grpcAuthV1.PATReq{ - PatId: id, - Domain: domainID, - Operation: "create", - UserId: id, - EntityId: auth.AnyIDs, - EntityType: auth.BootstrapStr, - }, - }, - authResponse: &grpcAuthV1.AuthZRes{Authorized: true}, - expectedReq: &policies.Policy{ - Domain: domainID, - SubjectType: policies.UserType, - SubjectKind: policies.UsersKind, - Subject: id, - Permission: policies.MembershipPermission, - ObjectType: policies.DomainType, - Object: domainID, - }, - expectedPAT: &auth.PATAuthz{ - PatID: id, - UserID: id, - EntityType: auth.BootstrapType, - EntityID: auth.AnyIDs, - Operation: "create", - Domain: domainID, - }, - err: nil, - }, - { - desc: "authorize user with unauthorized PAT token", - token: inValidToken, - authRequest: &grpcAuthV1.AuthZReq{ - PolicyReq: &grpcAuthV1.PolicyReq{ - Subject: id, - SubjectType: policies.UserType, - SubjectKind: policies.UsersKind, - Permission: policies.ViewPermission, - ObjectType: policies.ClientType, - Domain: domainID, - Object: clientID, - }, - PatReq: &grpcAuthV1.PATReq{ - PatId: id, - Domain: domainID, - Operation: "view", - UserId: id, - EntityId: clientID, - EntityType: auth.ClientsScopeStr, - }, - }, - authResponse: &grpcAuthV1.AuthZRes{Authorized: false}, - err: svcerr.ErrAuthorization, - }, - { - desc: "authorize PAT with missing user id", - token: validPATToken, - authRequest: &grpcAuthV1.AuthZReq{ - PolicyReq: &grpcAuthV1.PolicyReq{ - Subject: id, - SubjectType: policies.UserType, - SubjectKind: policies.UsersKind, - Permission: policies.ViewPermission, - ObjectType: policies.ClientType, - Domain: domainID, - Object: clientID, - }, - PatReq: &grpcAuthV1.PATReq{ - PatId: id, - Domain: domainID, - Operation: "view", - EntityId: clientID, - EntityType: auth.ClientsScopeStr, - }, - }, - authResponse: &grpcAuthV1.AuthZRes{Authorized: false}, - err: apiutil.ErrMissingUserID, - }, - { - desc: "authorize PAT with missing entity id", - token: validPATToken, - authRequest: &grpcAuthV1.AuthZReq{ - PolicyReq: &grpcAuthV1.PolicyReq{ - Subject: id, - SubjectType: policies.UserType, - SubjectKind: policies.UsersKind, - Permission: policies.ViewPermission, - ObjectType: policies.ClientType, - Domain: domainID, - Object: clientID, - }, - PatReq: &grpcAuthV1.PATReq{ - PatId: id, - Domain: domainID, - Operation: "view", - UserId: id, - EntityType: auth.ClientsScopeStr, - }, - }, - authResponse: &grpcAuthV1.AuthZRes{Authorized: false}, - err: apiutil.ErrMissingID, - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("Authorize", mock.Anything, mock.Anything, mock.Anything).Return(tc.err) - ar, err := grpcClient.Authorize(context.Background(), tc.authRequest) - if ar != nil { - assert.Equal(t, tc.authResponse, ar, fmt.Sprintf("%s: expected %v got %v", tc.desc, tc.authResponse, ar)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - svcCall.Unset() - }) - } -} diff --git a/auth/api/grpc/auth/setup_test.go b/auth/api/grpc/auth/setup_test.go deleted file mode 100644 index b6ff6bdfd..000000000 --- a/auth/api/grpc/auth/setup_test.go +++ /dev/null @@ -1,24 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package auth_test - -import ( - "os" - "testing" - - "github.com/absmach/magistrala/auth/mocks" -) - -var svc *mocks.Service - -func TestMain(m *testing.M) { - svc = new(mocks.Service) - server := startGRPCServer(svc, port) - - code := m.Run() - - server.GracefulStop() - - os.Exit(code) -} diff --git a/auth/api/grpc/token/endpoint_test.go b/auth/api/grpc/token/endpoint_test.go deleted file mode 100644 index 6d2e21099..000000000 --- a/auth/api/grpc/token/endpoint_test.go +++ /dev/null @@ -1,245 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package token_test - -import ( - "context" - "fmt" - "net" - "testing" - "time" - - grpcTokenV1 "github.com/absmach/magistrala/api/grpc/token/v1" - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/auth" - grpcapi "github.com/absmach/magistrala/auth/api/grpc/token" - "github.com/absmach/magistrala/internal/testsutil" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" - "google.golang.org/grpc" - "google.golang.org/grpc/credentials/insecure" -) - -const ( - port = 8082 - validToken = "valid" - inValidToken = "invalid" - invalidID = "invalid" -) - -var ( - validID = testsutil.GenerateUUID(&testing.T{}) - authAddr = fmt.Sprintf("localhost:%d", port) -) - -func startGRPCServer(svc auth.Service, port int) *grpc.Server { - listener, _ := net.Listen("tcp", fmt.Sprintf(":%d", port)) - server := grpc.NewServer() - grpcTokenV1.RegisterTokenServiceServer(server, grpcapi.NewTokenServer(svc)) - go func() { - err := server.Serve(listener) - assert.Nil(&testing.T{}, err, fmt.Sprintf(`"Unexpected error creating auth server %s"`, err)) - }() - - return server -} - -func TestIssue(t *testing.T) { - conn, err := grpc.NewClient(authAddr, grpc.WithTransportCredentials(insecure.NewCredentials())) - assert.Nil(t, err, fmt.Sprintf("Unexpected error creating client connection %s", err)) - grpcClient := grpcapi.NewTokenClient(conn, time.Second) - defer conn.Close() - - cases := []struct { - desc string - userId string - kind auth.KeyType - issueResponse auth.Token - err error - }{ - { - desc: "issue for user with valid token", - userId: validID, - kind: auth.AccessKey, - issueResponse: auth.Token{ - AccessToken: validToken, - RefreshToken: validToken, - }, - err: nil, - }, - { - desc: "issue recovery key", - userId: validID, - kind: auth.RecoveryKey, - issueResponse: auth.Token{ - AccessToken: validToken, - RefreshToken: validToken, - }, - err: nil, - }, - { - desc: "issue API key unauthenticated", - userId: validID, - kind: auth.APIKey, - issueResponse: auth.Token{}, - err: svcerr.ErrAuthentication, - }, - { - desc: "issue for invalid key type", - userId: validID, - kind: 32, - issueResponse: auth.Token{}, - err: errors.ErrMalformedEntity, - }, - { - desc: "issue for user that does notexist", - userId: "", - kind: auth.APIKey, - issueResponse: auth.Token{}, - err: svcerr.ErrAuthentication, - }, - } - - for _, tc := range cases { - svcCall := svc.On("Issue", mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(tc.issueResponse, tc.err) - _, err := grpcClient.Issue(context.Background(), &grpcTokenV1.IssueReq{UserId: tc.userId, Type: uint32(tc.kind)}) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - svcCall.Unset() - } -} - -func TestRefresh(t *testing.T) { - conn, err := grpc.NewClient(authAddr, grpc.WithTransportCredentials(insecure.NewCredentials())) - assert.Nil(t, err, fmt.Sprintf("Unexpected error creating client connection %s", err)) - grpcClient := grpcapi.NewTokenClient(conn, time.Second) - defer conn.Close() - - cases := []struct { - desc string - token string - issueResponse auth.Token - err error - }{ - { - desc: "refresh token with valid token", - token: validToken, - issueResponse: auth.Token{ - AccessToken: validToken, - RefreshToken: validToken, - }, - err: nil, - }, - { - desc: "refresh token with invalid token", - token: inValidToken, - issueResponse: auth.Token{}, - err: svcerr.ErrAuthentication, - }, - { - desc: "refresh token with empty token", - token: "", - issueResponse: auth.Token{}, - err: apiutil.ErrMissingSecret, - }, - } - - for _, tc := range cases { - svcCall := svc.On("Issue", mock.Anything, mock.Anything, mock.Anything).Return(tc.issueResponse, tc.err) - _, err := grpcClient.Refresh(context.Background(), &grpcTokenV1.RefreshReq{RefreshToken: tc.token}) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - svcCall.Unset() - } -} - -func TestRevoke(t *testing.T) { - conn, err := grpc.NewClient(authAddr, grpc.WithTransportCredentials(insecure.NewCredentials())) - assert.Nil(t, err, fmt.Sprintf("Unexpected error creating client connection %s", err)) - grpcClient := grpcapi.NewTokenClient(conn, time.Second) - defer conn.Close() - - cases := []struct { - desc string - id string - err error - }{ - { - desc: "revoke token with valid id", - id: validID, - err: nil, - }, - { - desc: "revoke token with invalid id", - id: invalidID, - err: svcerr.ErrAuthentication, - }, - { - desc: "revoke token with empty id", - id: "", - err: apiutil.ErrMissingID, - }, - { - desc: "revoke already revoked token", - id: validID, - err: svcerr.ErrConflict, - }, - } - - for _, tc := range cases { - svcCall := svc.On("RevokeToken", mock.Anything, mock.Anything, tc.id).Return(tc.err) - _, err := grpcClient.Revoke(context.Background(), &grpcTokenV1.RevokeReq{TokenId: tc.id}) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - svcCall.Unset() - } -} - -func TestListUserRefreshTokens(t *testing.T) { - conn, err := grpc.NewClient(authAddr, grpc.WithTransportCredentials(insecure.NewCredentials())) - assert.Nil(t, err, fmt.Sprintf("Unexpected error creating client connection %s", err)) - grpcClient := grpcapi.NewTokenClient(conn, time.Second) - defer conn.Close() - - cases := []struct { - desc string - userID string - listResponse []auth.TokenInfo - err error - }{ - { - desc: "list tokens for user with valid id", - userID: validID, - listResponse: []auth.TokenInfo{ - {ID: testsutil.GenerateUUID(&testing.T{}), Description: "Token 1"}, - {ID: testsutil.GenerateUUID(&testing.T{}), Description: "Token 2"}, - }, - err: nil, - }, - { - desc: "list tokens for user with empty list", - userID: validID, - listResponse: []auth.TokenInfo{}, - err: nil, - }, - { - desc: "list tokens with invalid user id", - userID: invalidID, - listResponse: nil, - err: svcerr.ErrAuthentication, - }, - { - desc: "list tokens with empty user id", - userID: "", - listResponse: nil, - err: apiutil.ErrMissingID, - }, - } - - for _, tc := range cases { - svcCall := svc.On("ListUserRefreshTokens", mock.Anything, tc.userID).Return(tc.listResponse, tc.err) - _, err := grpcClient.ListUserRefreshTokens(context.Background(), &grpcTokenV1.ListUserRefreshTokensReq{UserId: tc.userID}) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - svcCall.Unset() - } -} diff --git a/auth/api/grpc/token/setup_test.go b/auth/api/grpc/token/setup_test.go deleted file mode 100644 index 8a8c2e0c4..000000000 --- a/auth/api/grpc/token/setup_test.go +++ /dev/null @@ -1,24 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package token_test - -import ( - "os" - "testing" - - "github.com/absmach/magistrala/auth/mocks" -) - -var svc *mocks.Service - -func TestMain(m *testing.M) { - svc = new(mocks.Service) - server := startGRPCServer(svc, port) - - code := m.Run() - - server.GracefulStop() - - os.Exit(code) -} diff --git a/auth/middleware/logging.go b/auth/middleware/logging.go index e7b92c91b..2a7970f5c 100644 --- a/auth/middleware/logging.go +++ b/auth/middleware/logging.go @@ -21,7 +21,7 @@ type loggingMiddleware struct { svc auth.Service } -// NewLogging adds logging facilities to the core service. +// NewLogging adds logging facilities to the service. func NewLogging(svc auth.Service, logger *slog.Logger) auth.Service { return &loggingMiddleware{logger, svc} } diff --git a/auth/middleware/metrics.go b/auth/middleware/metrics.go index e2e27b1b0..bd82acabe 100644 --- a/auth/middleware/metrics.go +++ b/auth/middleware/metrics.go @@ -22,7 +22,7 @@ type metricsMiddleware struct { svc auth.Service } -// NewMetrics instruments core service by tracking request count and latency. +// NewMetrics instruments service by tracking request count and latency. func NewMetrics(svc auth.Service, counter metrics.Counter, latency metrics.Histogram) auth.Service { return &metricsMiddleware{ counter: counter, diff --git a/auth/scope_test.go b/auth/scope_test.go deleted file mode 100644 index 66f40417e..000000000 --- a/auth/scope_test.go +++ /dev/null @@ -1,376 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package auth_test - -import ( - "testing" - "time" - - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/auth" - channelsOps "github.com/absmach/magistrala/channels/operations" - clientsOps "github.com/absmach/magistrala/clients/operations" - groupsOps "github.com/absmach/magistrala/groups/operations" - "github.com/stretchr/testify/assert" -) - -func TestScopeAuthorized(t *testing.T) { - cases := []struct { - desc string - scope *auth.Scope - entityType auth.EntityType - domainID string - operation string - entityID string - expected bool - }{ - { - desc: "Authorized with matching entity type, domain, operation and entity ID", - scope: &auth.Scope{ - EntityType: auth.GroupsType, - DomainID: "domain1", - Operation: "view", - EntityID: "entity1", - }, - entityType: auth.GroupsType, - domainID: "domain1", - operation: "view", - entityID: "entity1", - expected: true, - }, - { - desc: "Authorized with wildcard entity ID", - scope: &auth.Scope{ - EntityType: auth.GroupsType, - DomainID: "domain1", - Operation: "view", - EntityID: "*", - }, - entityType: auth.GroupsType, - domainID: "domain1", - operation: "view", - entityID: "any-entity", - expected: true, - }, - { - desc: "Authorized without domain ID", - scope: &auth.Scope{ - EntityType: auth.ClientsType, - DomainID: "", - Operation: "view", - EntityID: "client1", - }, - entityType: auth.ClientsType, - domainID: "domain1", - operation: "view", - entityID: "client1", - expected: true, - }, - { - desc: "Not authorized with different entity type", - scope: &auth.Scope{ - EntityType: auth.GroupsType, - DomainID: "domain1", - Operation: "view", - EntityID: "entity1", - }, - entityType: auth.ChannelsType, - domainID: "domain1", - operation: "view", - entityID: "entity1", - expected: false, - }, - { - desc: "Not authorized with different domain ID", - scope: &auth.Scope{ - EntityType: auth.GroupsType, - DomainID: "domain1", - Operation: "view", - EntityID: "entity1", - }, - entityType: auth.GroupsType, - domainID: "domain2", - operation: "view", - entityID: "entity1", - expected: false, - }, - { - desc: "Not authorized with different operation", - scope: &auth.Scope{ - EntityType: auth.GroupsType, - DomainID: "domain1", - Operation: "view", - EntityID: "entity1", - }, - entityType: auth.GroupsType, - domainID: "domain1", - operation: "delete", - entityID: "entity1", - expected: false, - }, - { - desc: "Not authorized with different entity ID", - scope: &auth.Scope{ - EntityType: auth.GroupsType, - DomainID: "domain1", - Operation: "view", - EntityID: "entity1", - }, - entityType: auth.GroupsType, - domainID: "domain1", - operation: "view", - entityID: "entity2", - expected: false, - }, - { - desc: "Not authorized with nil scope", - scope: nil, - entityType: auth.GroupsType, - domainID: "domain1", - operation: "view", - entityID: "entity1", - expected: false, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - result := tc.scope.Authorized(tc.entityType, tc.domainID, tc.operation, tc.entityID) - assert.Equal(t, tc.expected, result, "Authorized() = %v, expected %v", result, tc.expected) - }) - } -} - -func TestScopeValidate(t *testing.T) { - cases := []struct { - desc string - scope *auth.Scope - err error - }{ - { - desc: "Valid scope for groups with domain ID", - scope: &auth.Scope{ - EntityType: auth.GroupsType, - DomainID: "domain1", - Operation: "view", - EntityID: "entity1", - }, - err: nil, - }, - { - desc: "Valid scope for channels with domain ID", - scope: &auth.Scope{ - EntityType: auth.ChannelsType, - DomainID: "domain1", - Operation: "view", - EntityID: "channel1", - }, - err: nil, - }, - { - desc: "Valid scope for clients with domain ID", - scope: &auth.Scope{ - EntityType: auth.ClientsType, - DomainID: "domain1", - Operation: "update", - EntityID: "client1", - }, - err: nil, - }, - { - desc: "Valid scope for messages with domain ID", - scope: &auth.Scope{ - EntityType: auth.MessagesType, - DomainID: "domain1", - Operation: "message_publish", - EntityID: "message1", - }, - err: nil, - }, - { - desc: "Valid scope for dashboard with domain ID", - scope: &auth.Scope{ - EntityType: auth.DashboardType, - DomainID: "domain1", - Operation: "dashboard_share", - EntityID: "dashboard1", - }, - err: nil, - }, - { - desc: "Valid scope with wildcard entity ID", - scope: &auth.Scope{ - EntityType: auth.GroupsType, - DomainID: "domain1", - Operation: "view", - EntityID: "*", - }, - err: nil, - }, - { - desc: "Invalid nil scope", - scope: nil, - err: assert.AnError, // Will be checked with Contains - }, - { - desc: "Invalid scope without entity ID", - scope: &auth.Scope{ - EntityType: auth.GroupsType, - DomainID: "domain1", - Operation: groupsOps.OperationDetails()[groupsOps.OpViewGroup].Name, - EntityID: "", - }, - err: apiutil.ErrMissingEntityID, - }, - { - desc: "Invalid scope for groups without domain ID", - scope: &auth.Scope{ - EntityType: auth.GroupsType, - DomainID: "", - Operation: groupsOps.OperationDetails()[groupsOps.OpViewGroup].Name, - EntityID: "entity1", - }, - err: apiutil.ErrMissingDomainID, - }, - { - desc: "Invalid scope for channels without domain ID", - scope: &auth.Scope{ - EntityType: auth.ChannelsType, - DomainID: "", - Operation: channelsOps.OperationDetails()[channelsOps.OpViewChannel].Name, - EntityID: "channel1", - }, - err: apiutil.ErrMissingDomainID, - }, - { - desc: "Invalid scope for clients without domain ID", - scope: &auth.Scope{ - EntityType: auth.ClientsType, - DomainID: "", - Operation: clientsOps.OperationDetails()[clientsOps.OpViewClient].Name, - EntityID: "client1", - }, - err: apiutil.ErrMissingDomainID, - }, - { - desc: "Invalid scope for dashboard without domain ID", - scope: &auth.Scope{ - EntityType: auth.DashboardType, - DomainID: "", - Operation: auth.OpShare, - EntityID: "dashboard1", - }, - err: apiutil.ErrMissingDomainID, - }, - { - desc: "Invalid scope for messages without domain ID", - scope: &auth.Scope{ - EntityType: auth.MessagesType, - DomainID: "", - Operation: auth.OpPublish, - EntityID: "message1", - }, - err: apiutil.ErrMissingDomainID, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := tc.scope.Validate() - if tc.err != nil { - assert.Error(t, err, "Validate() should return error") - if tc.err != assert.AnError { - assert.Equal(t, tc.err, err, "Validate() error = %v, expected %v", err, tc.err) - } - } else { - assert.NoError(t, err, "Validate() should not return error") - } - }) - } -} - -func TestPATValidate(t *testing.T) { - cases := []struct { - desc string - pat *auth.PAT - err bool - }{ - { - desc: "Valid PAT", - pat: &auth.PAT{ - ID: "pat-id", - User: "user-id", - Name: "test-pat", - Description: "test description", - }, - err: false, - }, - { - desc: "Invalid nil PAT", - pat: nil, - err: true, - }, - { - desc: "Invalid PAT without name", - pat: &auth.PAT{ - ID: "pat-id", - User: "user-id", - Name: "", - Description: "test description", - }, - err: true, - }, - { - desc: "Invalid PAT without user", - pat: &auth.PAT{ - ID: "pat-id", - User: "", - Name: "test-pat", - Description: "test description", - }, - err: true, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := tc.pat.Validate() - if tc.err { - assert.Error(t, err, "Validate() should return error") - } else { - assert.NoError(t, err, "Validate() should not return error") - } - }) - } -} - -func TestPATMarshalUnmarshalBinary(t *testing.T) { - pat := auth.PAT{ - ID: "pat-id", - User: "user-id", - Name: "test-pat", - Description: "test description", - Secret: "secret", - IssuedAt: time.Now().UTC().Round(time.Second), - ExpiresAt: time.Now().UTC().Add(24 * time.Hour).Round(time.Second), - Status: auth.ActiveStatus, - } - - // Marshal - data, err := pat.MarshalBinary() - assert.NoError(t, err, "MarshalBinary() should not return error") - assert.NotNil(t, data, "MarshalBinary() should return data") - - // Unmarshal - var newPAT auth.PAT - err = newPAT.UnmarshalBinary(data) - assert.NoError(t, err, "UnmarshalBinary() should not return error") - - assert.Equal(t, pat.ID, newPAT.ID, "ID mismatch") - assert.Equal(t, pat.User, newPAT.User, "User mismatch") - assert.Equal(t, pat.Name, newPAT.Name, "Name mismatch") - assert.Equal(t, pat.Description, newPAT.Description, "Description mismatch") - assert.Equal(t, pat.Secret, newPAT.Secret, "Secret mismatch") - assert.Equal(t, pat.Status, newPAT.Status, "Status mismatch") -} diff --git a/bootstrap/README.md b/bootstrap/README.md deleted file mode 100644 index 17035f7d1..000000000 --- a/bootstrap/README.md +++ /dev/null @@ -1,122 +0,0 @@ -# BOOTSTRAP SERVICE - -New devices need to be configured properly and connected to the Magistrala. Bootstrap service is used in order to accomplish that. This service provides the following features: - -1. Creating new Magistrala Clients -2. Providing basic configuration for the newly created Clients -3. Enabling/disabling bootstrap enrollments - -Pre-provisioning a new Client is as simple as sending Configuration data to the Bootstrap service. Once the Client is online, it sends a request for initial config to Bootstrap service. Bootstrap service provides an API for enabling and disabling bootstrap enrollments. Bootstrapping does not implicitly enable an enrollment; it has to be done manually. - -In order to bootstrap successfully, the Client needs to send bootstrapping request to the specific URL, as well as a secret key. This key and URL are pre-provisioned during the manufacturing process. If the Client is provisioned on the Bootstrap service side, the corresponding configuration will be sent as a response. Otherwise, the Client will be saved so that it can be provisioned later. - -## Client Configuration Entity - -Client Configuration consists of two logical parts: the custom configuration that can be interpreted by the Client itself and Magistrala-related configuration. Magistrala config contains: - -1. corresponding Magistrala Client ID -2. corresponding Magistrala Client key -3. list of the Magistrala channels the Client is connected to - -> Note: list of channels contains IDs of the Magistrala channels. These channels are _pre-provisioned_ on the Magistrala side and, unlike corresponding Magistrala Client, Bootstrap service is not able to create Magistrala Channels. - -Enabling and disabling a bootstrap enrollment is an enrollment toggle. Configuration keeps a _status_: - -| Status | What it means | -| -------- | ----------------------------------------------------------- | -| disabled | Enrollment exists, but bootstrap is not allowed | -| enabled | Enrollment can be used to fetch bootstrap configuration | - -Switching between statuses `enabled` and `disabled` enables and disables the enrollment, respectively. - -Client configuration also contains the so-called `external ID` and `external key`. An external ID is a unique identifier of corresponding Client. For example, a device MAC address is a good choice for external ID. External key is a secret key that is used for authentication during the bootstrapping procedure. - -## Configuration - -The service is configured using the environment variables presented in the following table. Note that any unset variables will be replaced with their default values. - -| Variable | Description | Default | -| ------------------------------ | -------------------------------------------------------------------------------- | --------------------------------- | -| MG_BOOTSTRAP_LOG_LEVEL | Log level for Bootstrap (debug, info, warn, error) | info | -| MG_BOOTSTRAP_DB_HOST | Database host address | localhost | -| MG_BOOTSTRAP_DB_PORT | Database host port | 5432 | -| MG_BOOTSTRAP_DB_USER | Database user | magistrala | -| MG_BOOTSTRAP_DB_PASS | Database password | magistrala | -| MG_BOOTSTRAP_DB_NAME | Name of the database used by the service | bootstrap | -| MG_BOOTSTRAP_DB_SSL_MODE | Database connection SSL mode (disable, require, verify-ca, verify-full) | disable | -| MG_BOOTSTRAP_DB_SSL_CERT | Path to the PEM encoded certificate file | "" | -| MG_BOOTSTRAP_DB_SSL_KEY | Path to the PEM encoded key file | "" | -| MG_BOOTSTRAP_DB_SSL_ROOT_CERT | Path to the PEM encoded root certificate file | "" | -| MG_BOOTSTRAP_ENCRYPT_KEY | Secret key for secure bootstrapping encryption | 12345678910111213141516171819202 | -| MG_BOOTSTRAP_HTTP_HOST | Bootstrap service HTTP host | "" | -| MG_BOOTSTRAP_HTTP_PORT | Bootstrap service HTTP port | 9013 | -| MG_BOOTSTRAP_HTTP_SERVER_CERT | Path to server certificate in pem format | "" | -| MG_BOOTSTRAP_HTTP_SERVER_KEY | Path to server key in pem format | "" | -| MG_BOOTSTRAP_EVENT_CONSUMER | Bootstrap service event source consumer name | bootstrap | -| MG_ES_URL | Event store URL | | -| MG_AUTH_GRPC_URL | Auth service Auth gRPC URL | | -| MG_AUTH_GRPC_TIMEOUT | Auth service Auth gRPC request timeout in seconds | 1s | -| MG_AUTH_GRPC_CLIENT_CERT | Path to the PEM encoded auth service Auth gRPC client certificate file | "" | -| MG_AUTH_GRPC_CLIENT_KEY | Path to the PEM encoded auth service Auth gRPC client key file | "" | -| MG_AUTH_GRPC_SERVER_CERTS | Path to the PEM encoded auth server Auth gRPC server trusted CA certificate file | "" | -| MG_CLIENTS_URL | Base URL for Magistrala Clients | | -| MG_JAEGER_URL | Jaeger server URL | | -| MG_JAEGER_TRACE_RATIO | Jaeger sampling ratio | 1.0 | -| MG_SEND_TELEMETRY | Send telemetry to magistrala call home server | true | -| MG_BOOTSTRAP_INSTANCE_ID | Bootstrap service instance ID | "" | - -## Deployment - -The service itself is distributed as Docker container. Check the [`bootstrap`](https://github.com/absmach/magistrala/blob/main/docker/addons/bootstrap/docker-compose.yaml) service section in docker-compose file to see how service is deployed. - -To start the service outside of the container, execute the following shell script: - -```bash -# download the latest version of the service -git clone https://github.com/absmach/magistrala - -cd magistrala - -# compile the servic e -make bootstrap - -# copy binary to bin -make install - -# set the environment variables and run the service -MG_BOOTSTRAP_LOG_LEVEL=info \ -MG_BOOTSTRAP_DB_HOST=localhost \ -MG_BOOTSTRAP_DB_PORT=5432 \ -MG_BOOTSTRAP_DB_USER=magistrala \ -MG_BOOTSTRAP_DB_PASS=magistrala \ -MG_BOOTSTRAP_DB_NAME=bootstrap \ -MG_BOOTSTRAP_DB_SSL_MODE=disable \ -MG_BOOTSTRAP_DB_SSL_CERT="" \ -MG_BOOTSTRAP_DB_SSL_KEY="" \ -MG_BOOTSTRAP_DB_SSL_ROOT_CERT="" \ -MG_BOOTSTRAP_HTTP_HOST=localhost \ -MG_BOOTSTRAP_HTTP_PORT=9013 \ -MG_BOOTSTRAP_HTTP_SERVER_CERT="" \ -MG_BOOTSTRAP_HTTP_SERVER_KEY="" \ -MG_BOOTSTRAP_EVENT_CONSUMER=bootstrap \ -MG_ES_URL=nats://localhost:4222 \ -MG_AUTH_GRPC_URL=localhost:8181 \ -MG_AUTH_GRPC_TIMEOUT=1s \ -MG_AUTH_GRPC_CLIENT_CERT="" \ -MG_AUTH_GRPC_CLIENT_KEY="" \ -MG_AUTH_GRPC_SERVER_CERTS="" \ -MG_CLIENTS_URL=http://localhost:9000 \ -MG_JAEGER_URL=http://localhost:14268/api/traces \ -MG_JAEGER_TRACE_RATIO=1.0 \ -MG_SEND_TELEMETRY=true \ -MG_BOOTSTRAP_INSTANCE_ID="" \ -$GOBIN/magistrala-bootstrap -``` - -Setting `MG_BOOTSTRAP_HTTP_SERVER_CERT` and `MG_BOOTSTRAP_HTTP_SERVER_KEY` will enable TLS against the service. The service expects a file in PEM format for both the certificate and the key. - -Setting `MG_AUTH_GRPC_CLIENT_CERT` and `MG_AUTH_GRPC_CLIENT_KEY` will enable TLS against the auth service. The service expects a file in PEM format for both the certificate and the key. Setting `MG_AUTH_GRPC_SERVER_CERTS` will enable TLS against the auth service trusting only those CAs that are provided. The service expects a file in PEM format of trusted CAs. - -## Usage - -For more information about service capabilities and its usage, please check out the [API documentation](https://docs.api.magistrala.absmach.eu/?urls.primaryName=bootstrap.yaml). diff --git a/bootstrap/api/doc.go b/bootstrap/api/doc.go deleted file mode 100644 index 1e8268ee6..000000000 --- a/bootstrap/api/doc.go +++ /dev/null @@ -1,5 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package api contains implementation of bootstrap service HTTP API. -package api diff --git a/bootstrap/api/endpoint.go b/bootstrap/api/endpoint.go deleted file mode 100644 index 95248640e..000000000 --- a/bootstrap/api/endpoint.go +++ /dev/null @@ -1,506 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package api - -import ( - "context" - - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/bootstrap" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/go-kit/kit/endpoint" -) - -func addEndpoint(svc bootstrap.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(addReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthorization - } - - config := bootstrap.Config{ - ExternalID: req.ExternalID, - ExternalKey: req.ExternalKey, - Name: req.Name, - ClientCert: req.ClientCert, - ClientKey: req.ClientKey, - CACert: req.CACert, - Content: req.Content, - ProfileID: req.ProfileID, - RenderContext: req.RenderContext, - } - - saved, err := svc.Add(ctx, session, req.token, config) - if err != nil { - return nil, err - } - - res := configRes{ - ID: saved.ID, - ExternalID: saved.ExternalID, - Name: saved.Name, - Content: saved.Content, - Status: saved.Status, - ProfileID: saved.ProfileID, - RenderContext: saved.RenderContext, - ClientCert: saved.ClientCert, - CACert: saved.CACert, - ClientKey: saved.ClientKey, - created: true, - } - - return res, nil - } -} - -func updateCertEndpoint(svc bootstrap.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(updateCertReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthorization - } - - cfg, err := svc.UpdateCert(ctx, session, req.configID, req.ClientCert, req.ClientKey, req.CACert) - if err != nil { - return nil, err - } - - res := updateConfigRes{ - ID: cfg.ID, - ClientCert: cfg.ClientCert, - CACert: cfg.CACert, - ClientKey: cfg.ClientKey, - } - - return res, nil - } -} - -func viewEndpoint(svc bootstrap.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(entityReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthorization - } - - config, err := svc.View(ctx, session, req.id) - if err != nil { - return nil, err - } - - res := viewRes{ - ID: config.ID, - ExternalID: config.ExternalID, - Name: config.Name, - Content: config.Content, - Status: config.Status, - ProfileID: config.ProfileID, - RenderContext: config.RenderContext, - } - - return res, nil - } -} - -func updateEndpoint(svc bootstrap.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(updateReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthorization - } - - config := bootstrap.Config{ - ID: req.id, - Name: req.Name, - Content: req.Content, - RenderContext: req.RenderContext, - } - - if err := svc.Update(ctx, session, config); err != nil { - return nil, err - } - - return updateRes{}, nil - } -} - -func listEndpoint(svc bootstrap.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(listReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthorization - } - - page, err := svc.List(ctx, session, req.filter, req.offset, req.limit) - if err != nil { - return nil, err - } - res := listRes{ - Total: page.Total, - Offset: page.Offset, - Limit: page.Limit, - Configs: []viewRes{}, - } - - for _, cfg := range page.Configs { - view := viewRes{ - ID: cfg.ID, - ExternalID: cfg.ExternalID, - Name: cfg.Name, - Content: cfg.Content, - Status: cfg.Status, - ProfileID: cfg.ProfileID, - RenderContext: cfg.RenderContext, - } - res.Configs = append(res.Configs, view) - } - - return res, nil - } -} - -func removeEndpoint(svc bootstrap.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(entityReq) - if err := req.validate(); err != nil { - return removeRes{}, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthorization - } - - if err := svc.Remove(ctx, session, req.id); err != nil { - return nil, err - } - - return removeRes{}, nil - } -} - -func bootstrapEndpoint(svc bootstrap.Service, reader bootstrap.ConfigReader, secure bool) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(bootstrapReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - cfg, err := svc.Bootstrap(ctx, req.key, req.id, secure) - if err != nil { - return nil, err - } - - return reader.ReadConfig(cfg, secure) - } -} - -func enableConfigEndpoint(svc bootstrap.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(changeConfigStatusReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthorization - } - - cfg, err := svc.EnableConfig(ctx, session, req.id) - if err != nil { - return nil, err - } - - return changeConfigStatusRes{Config: cfg}, nil - } -} - -func disableConfigEndpoint(svc bootstrap.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(changeConfigStatusReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthorization - } - - cfg, err := svc.DisableConfig(ctx, session, req.id) - if err != nil { - return nil, err - } - - return changeConfigStatusRes{Config: cfg}, nil - } -} - -func createProfileEndpoint(svc bootstrap.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(createProfileReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthorization - } - saved, err := svc.CreateProfile(ctx, session, req.Profile) - if err != nil { - return nil, err - } - return profileRes{Profile: saved, created: true}, nil - } -} - -func uploadProfileEndpoint(svc bootstrap.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(uploadProfileReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthorization - } - saved, err := svc.CreateProfile(ctx, session, req.Profile) - if err != nil { - return nil, err - } - return profileRes{Profile: saved, created: true}, nil - } -} - -func viewProfileEndpoint(svc bootstrap.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(viewProfileReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthorization - } - p, err := svc.ViewProfile(ctx, session, req.profileID) - if err != nil { - return nil, err - } - return profileRes{Profile: p}, nil - } -} - -func profileSlotsEndpoint(svc bootstrap.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(viewProfileReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthorization - } - p, err := svc.ViewProfile(ctx, session, req.profileID) - if err != nil { - return nil, err - } - return profileSlotsRes{BindingSlots: p.BindingSlots}, nil - } -} - -func renderPreviewEndpoint(svc bootstrap.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(renderPreviewReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthorization - } - p, err := svc.ViewProfile(ctx, session, req.profileID) - if err != nil { - return nil, err - } - - cfg := req.Config - bindings := req.Bindings - - if req.ConfigID != "" { - stored, err := svc.View(ctx, session, req.ConfigID) - if err != nil { - return nil, err - } - cfg = stored - bindings, err = svc.ListBindings(ctx, session, req.ConfigID) - if err != nil { - return nil, err - } - } - - cfg.DomainID = session.DomainID - cfg.ProfileID = p.ID - if cfg.RenderContext == nil { - cfg.RenderContext = req.RenderContext - } - - rendered, err := bootstrap.NewRenderer().Render(p, cfg, bindings) - if err != nil { - return nil, err - } - - return renderPreviewRes{Content: string(rendered)}, nil - } -} - -func updateProfileEndpoint(svc bootstrap.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(updateProfileReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthorization - } - req.Profile.ID = req.profileID - updated, err := svc.UpdateProfile(ctx, session, req.Profile) - if err != nil { - return nil, err - } - return profileRes{Profile: updated}, nil - } -} - -func deleteProfileEndpoint(svc bootstrap.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(deleteProfileReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthorization - } - if err := svc.DeleteProfile(ctx, session, req.profileID); err != nil { - return nil, err - } - return removeRes{}, nil - } -} - -func listProfilesEndpoint(svc bootstrap.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(listProfilesReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthorization - } - page, err := svc.ListProfiles(ctx, session, req.offset, req.limit, req.name) - if err != nil { - return nil, err - } - return profilesPageRes{ProfilesPage: page}, nil - } -} - -func assignProfileEndpoint(svc bootstrap.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(assignProfileReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthorization - } - if err := svc.AssignProfile(ctx, session, req.configID, req.ProfileID); err != nil { - return nil, err - } - return removeRes{}, nil - } -} - -func bindResourcesEndpoint(svc bootstrap.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(bindResourcesReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthorization - } - if err := svc.BindResources(ctx, session, req.token, req.configID, req.Bindings); err != nil { - return nil, err - } - return removeRes{}, nil - } -} - -func listBindingsEndpoint(svc bootstrap.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(listBindingsReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthorization - } - snapshots, err := svc.ListBindings(ctx, session, req.configID) - if err != nil { - return nil, err - } - return bindingsRes{Bindings: snapshots}, nil - } -} - -func refreshBindingsEndpoint(svc bootstrap.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(refreshBindingsReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthorization - } - if err := svc.RefreshBindings(ctx, session, req.token, req.configID); err != nil { - return nil, err - } - return removeRes{}, nil - } -} diff --git a/bootstrap/api/endpoint_test.go b/bootstrap/api/endpoint_test.go deleted file mode 100644 index 18a997a22..000000000 --- a/bootstrap/api/endpoint_test.go +++ /dev/null @@ -1,1533 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package api_test - -import ( - "context" - "crypto/aes" - "crypto/cipher" - "crypto/rand" - "encoding/hex" - "encoding/json" - "fmt" - "io" - "net/http" - "net/http/httptest" - "strconv" - "strings" - "testing" - - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/bootstrap" - bsapi "github.com/absmach/magistrala/bootstrap/api" - "github.com/absmach/magistrala/bootstrap/mocks" - "github.com/absmach/magistrala/internal/testsutil" - mglog "github.com/absmach/magistrala/logger" - smqauthn "github.com/absmach/magistrala/pkg/authn" - authnmocks "github.com/absmach/magistrala/pkg/authn/mocks" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" -) - -const ( - validToken = "validToken" - domainID = "b4d7d79e-fd99-4c2b-ac09-524e43df6888" - invalidToken = "invalid" - email = "test@example.com" - unknown = "unknown" - contentType = "application/json" - wrongID = "wrong_id" - - addName = "name" - addContent = "config" - instanceID = "5de9b29a-feb9-11ed-be56-0242ac120002" - validID = "d4ebb847-5d0e-4e46-bdd9-b6aceaaa3a22" -) - -var ( - encKey = []byte("1234567891011121") - addExternalID = testsutil.GenerateUUID(&testing.T{}) - addExternalKey = testsutil.GenerateUUID(&testing.T{}) - addID = testsutil.GenerateUUID(&testing.T{}) - addReq = struct { - ExternalID string `json:"external_id"` - ExternalKey string `json:"external_key"` - Name string `json:"name"` - Content string `json:"content"` - }{ - ExternalID: addExternalID, - ExternalKey: addExternalKey, - Name: "name", - Content: "config", - } - - updateReq = struct { - Content string `json:"content,omitempty"` - Status bootstrap.Status `json:"status,omitempty"` - ClientCert string `json:"client_cert,omitempty"` - CACert string `json:"ca_cert,omitempty"` - RenderContext map[string]any `json:"render_context,omitempty"` - }{ - Content: "config update", - Status: bootstrap.EnabledStatus, - ClientCert: "newcert", - CACert: "newca", - RenderContext: map[string]any{"site": "warehouse-2", "region": "mombasa"}, - } - - missingIDRes = toJSON(apiutil.ErrMissingID) - missingKeyRes = toJSON(apiutil.ErrBearerKey) - unknownExternalIDErrorRes = toJSON(svcerr.ErrNotFound) - extKeyRes = toJSON(bootstrap.ErrExternalKey) - extSecKeyRes = toJSON(bootstrap.ErrExternalKeySecure) -) - -type testRequest struct { - client *http.Client - method string - url string - contentType string - token string - key string - body io.Reader -} - -func newConfig() bootstrap.Config { - return bootstrap.Config{ - ID: addID, - ExternalID: addExternalID, - ExternalKey: addExternalKey, - Name: addName, - Content: addContent, - ClientCert: "newcert", - ClientKey: "newkey", - CACert: "newca", - } -} - -func (tr testRequest) make() (*http.Response, error) { - req, err := http.NewRequest(tr.method, tr.url, tr.body) - if err != nil { - return nil, err - } - - if tr.token != "" { - req.Header.Set("Authorization", apiutil.BearerPrefix+tr.token) - } - if tr.key != "" { - req.Header.Set("Authorization", apiutil.ClientPrefix+tr.key) - } - - if tr.contentType != "" { - req.Header.Set("Content-Type", tr.contentType) - } - - return tr.client.Do(req) -} - -func enc(in []byte) ([]byte, error) { - block, err := aes.NewCipher(encKey) - if err != nil { - return nil, err - } - ciphertext := make([]byte, aes.BlockSize+len(in)) - iv := ciphertext[:aes.BlockSize] - if _, err := io.ReadFull(rand.Reader, iv); err != nil { - return nil, err - } - stream := cipher.NewCFBEncrypter(block, iv) - stream.XORKeyStream(ciphertext[aes.BlockSize:], in) - return ciphertext, nil -} - -func dec(in []byte) ([]byte, error) { - block, err := aes.NewCipher(encKey) - if err != nil { - return nil, err - } - if len(in) < aes.BlockSize { - return nil, errors.ErrMalformedEntity - } - iv := in[:aes.BlockSize] - in = in[aes.BlockSize:] - stream := cipher.NewCFBDecrypter(block, iv) - stream.XORKeyStream(in, in) - return in, nil -} - -func newBootstrapServer() (*httptest.Server, *mocks.Service, *authnmocks.Authentication) { - logger := mglog.NewMock() - svc := new(mocks.Service) - authn := new(authnmocks.Authentication) - am := smqauthn.NewAuthNMiddleware(authn, smqauthn.WithAllowUnverifiedUser(true)) - mux := bsapi.MakeHandler(svc, am, bootstrap.NewConfigReader(encKey), logger, instanceID) - return httptest.NewServer(mux), svc, authn -} - -func toJSON(data any) string { - jsonData, err := json.Marshal(data) - if err != nil { - return "" - } - return string(jsonData) -} - -func TestAdd(t *testing.T) { - bs, svc, auth := newBootstrapServer() - defer bs.Close() - c := newConfig() - - data := toJSON(addReq) - - cases := []struct { - desc string - req string - domainID string - token string - session smqauthn.Session - contentType string - status int - location string - authenticateErr error - err error - }{ - { - desc: "add a config with invalid token", - req: data, - domainID: domainID, - token: invalidToken, - contentType: contentType, - status: http.StatusUnauthorized, - location: "", - authenticateErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "add a valid config", - req: data, - domainID: domainID, - token: validToken, - contentType: contentType, - status: http.StatusCreated, - location: "/clients/configs/" + c.ID, - err: nil, - }, - { - desc: "add a config with wrong content type", - req: data, - domainID: domainID, - token: validToken, - contentType: "", - status: http.StatusUnsupportedMediaType, - location: "", - err: apiutil.ErrUnsupportedContentType, - }, - { - desc: "add an existing config", - req: data, - domainID: domainID, - token: validToken, - contentType: contentType, - status: http.StatusBadRequest, - location: "", - err: svcerr.ErrConflict, - }, - { - desc: "add a config with wrong JSON", - req: "{\"external_id\": 5}", - domainID: domainID, - token: validToken, - contentType: contentType, - status: http.StatusBadRequest, - err: svcerr.ErrMalformedEntity, - }, - { - desc: "add a config with invalid request format", - req: "}", - domainID: domainID, - token: validToken, - contentType: contentType, - status: http.StatusBadRequest, - location: "", - err: svcerr.ErrMalformedEntity, - }, - { - desc: "add a config with empty JSON", - req: "{}", - domainID: domainID, - token: validToken, - contentType: contentType, - status: http.StatusBadRequest, - location: "", - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "add a config with an empty request", - req: "", - domainID: domainID, - token: validToken, - contentType: contentType, - status: http.StatusBadRequest, - location: "", - err: svcerr.ErrMalformedEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - - svcCall := svc.On("Add", mock.Anything, tc.session, tc.token, mock.Anything).Return(c, tc.err) - req := testRequest{ - client: bs.Client(), - method: http.MethodPost, - url: fmt.Sprintf("%s/%s/clients/configs", bs.URL, tc.domainID), - contentType: tc.contentType, - token: tc.token, - body: strings.NewReader(tc.req), - } - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - location := res.Header.Get("Location") - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - assert.Equal(t, tc.location, location, fmt.Sprintf("%s: expected location '%s' got '%s'", tc.desc, tc.location, location)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestView(t *testing.T) { - bs, svc, auth := newBootstrapServer() - defer bs.Close() - c := newConfig() - - data := config{ - ID: c.ID, - Status: c.Status, - ExternalID: c.ExternalID, - Name: c.Name, - Content: c.Content, - } - - cases := []struct { - desc string - token string - session smqauthn.Session - id string - status int - res config - authenticateErr error - err error - }{ - { - desc: "view a config with invalid token", - token: invalidToken, - id: c.ID, - status: http.StatusUnauthorized, - res: config{}, - authenticateErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "view a config", - token: validToken, - id: c.ID, - status: http.StatusOK, - res: data, - err: nil, - }, - { - desc: "view a non-existing config", - token: validToken, - id: wrongID, - status: http.StatusNotFound, - res: config{}, - err: svcerr.ErrNotFound, - }, - { - desc: "view a config with an empty token", - token: "", - id: c.ID, - status: http.StatusUnauthorized, - res: config{}, - err: apiutil.ErrBearerToken, - }, - { - desc: "view config without authorization", - token: validToken, - id: c.ID, - status: http.StatusForbidden, - res: config{}, - err: svcerr.ErrAuthorization, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := svc.On("View", mock.Anything, tc.session, tc.id).Return(c, tc.err) - req := testRequest{ - client: bs.Client(), - method: http.MethodGet, - url: fmt.Sprintf("%s/%s/clients/configs/%s", bs.URL, domainID, tc.id), - token: tc.token, - } - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - var view config - if err := json.NewDecoder(res.Body).Decode(&view); err != io.EOF { - assert.Nil(t, err, fmt.Sprintf("Decoding expected to succeed %s: %s", tc.desc, err)) - } - - assert.Equal(t, tc.res, view, fmt.Sprintf("%s: expected response '%s' got '%s'", tc.desc, tc.res, view)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestUpdate(t *testing.T) { - bs, svc, auth := newBootstrapServer() - defer bs.Close() - c := newConfig() - - data := toJSON(updateReq) - - cases := []struct { - desc string - req string - id string - token string - session smqauthn.Session - contentType string - status int - authenticateErr error - err error - }{ - { - desc: "update with invalid token", - req: data, - id: c.ID, - token: invalidToken, - contentType: contentType, - status: http.StatusUnauthorized, - authenticateErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "update with an empty token", - req: data, - id: c.ID, - token: "", - contentType: contentType, - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "update a valid config", - req: data, - id: c.ID, - token: validToken, - contentType: contentType, - status: http.StatusOK, - err: nil, - }, - { - desc: "update a config with wrong content type", - req: data, - id: c.ID, - token: validToken, - contentType: "", - status: http.StatusUnsupportedMediaType, - err: apiutil.ErrUnsupportedContentType, - }, - { - desc: "update a non-existing config", - req: data, - id: wrongID, - token: validToken, - contentType: contentType, - status: http.StatusNotFound, - err: svcerr.ErrNotFound, - }, - { - desc: "update a config with invalid request format", - req: "}", - id: c.ID, - token: validToken, - contentType: contentType, - status: http.StatusBadRequest, - err: svcerr.ErrMalformedEntity, - }, - { - desc: "update a config with an empty request", - id: c.ID, - req: "", - token: validToken, - contentType: contentType, - status: http.StatusBadRequest, - err: svcerr.ErrMalformedEntity, - }, - { - desc: "update a config render_context", - req: toJSON(struct { - RenderContext map[string]any `json:"render_context"` - }{RenderContext: map[string]any{"site": "warehouse-2", "region": "mombasa"}}), - id: c.ID, - token: validToken, - contentType: contentType, - status: http.StatusOK, - err: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := svc.On("Update", mock.Anything, tc.session, mock.Anything).Return(tc.err) - req := testRequest{ - client: bs.Client(), - method: http.MethodPatch, - url: fmt.Sprintf("%s/%s/clients/configs/%s", bs.URL, domainID, tc.id), - contentType: tc.contentType, - token: tc.token, - body: strings.NewReader(tc.req), - } - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestUpdateCert(t *testing.T) { - bs, svc, auth := newBootstrapServer() - defer bs.Close() - c := newConfig() - - data := toJSON(updateReq) - - cases := []struct { - desc string - req string - id string - token string - session smqauthn.Session - contentType string - status int - authenticateErr error - err error - }{ - { - desc: "update with invalid token", - req: data, - id: c.ID, - token: invalidToken, - contentType: contentType, - status: http.StatusUnauthorized, - authenticateErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "update with an empty token", - req: data, - id: c.ID, - token: "", - contentType: contentType, - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "update a valid config", - req: data, - id: c.ID, - token: validToken, - contentType: contentType, - status: http.StatusOK, - err: nil, - }, - { - desc: "update a config with wrong content type", - req: data, - id: c.ID, - token: validToken, - contentType: "", - status: http.StatusUnsupportedMediaType, - err: apiutil.ErrUnsupportedContentType, - }, - { - desc: "update a non-existing config", - req: data, - id: wrongID, - token: validToken, - contentType: contentType, - status: http.StatusNotFound, - err: svcerr.ErrNotFound, - }, - { - desc: "update a config with invalid request format", - req: "}", - id: c.ID, - token: validToken, - contentType: contentType, - status: http.StatusBadRequest, - err: svcerr.ErrMalformedEntity, - }, - { - desc: "update a config with an empty request", - id: c.ID, - req: "", - token: validToken, - contentType: contentType, - status: http.StatusBadRequest, - err: svcerr.ErrMalformedEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := svc.On("UpdateCert", mock.Anything, mock.Anything, mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(c, tc.err) - req := testRequest{ - client: bs.Client(), - method: http.MethodPatch, - url: fmt.Sprintf("%s/%s/clients/configs/certs/%s", bs.URL, domainID, tc.id), - contentType: tc.contentType, - token: tc.token, - body: strings.NewReader(tc.req), - } - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestList(t *testing.T) { - configNum := 101 - changedStatusNum := 20 - var active, inactive []config - list := make([]config, configNum) - - bs, svc, auth := newBootstrapServer() - defer bs.Close() - path := fmt.Sprintf("%s/%s/%s", bs.URL, domainID, "clients/configs") - - c := newConfig() - - for i := 0; i < configNum; i++ { - c.ExternalID = strconv.Itoa(i) - c.Name = fmt.Sprintf("%s-%d", addName, i) - c.ExternalKey = fmt.Sprintf("%s%s", addExternalKey, strconv.Itoa(i)) - - s := config{ - ID: c.ID, - ExternalID: c.ExternalID, - Name: c.Name, - Content: c.Content, - Status: c.Status, - } - list[i] = s - } - // Change status of first 20 elements for filtering tests. - for i := 0; i < changedStatusNum; i++ { - if i%2 == 0 { - // Even elements remain inactive (default status). - inactive = append(inactive, list[i]) - continue - } - // Odd elements are enabled (active). - enabledCfg := bootstrap.Config{ID: list[i].ID, Status: bootstrap.Active} - svcCall := svc.On("EnableConfig", context.Background(), mock.Anything, mock.Anything).Return(enabledCfg, nil) - _, err := svc.EnableConfig(context.Background(), smqauthn.Session{}, list[i].ID) - assert.Nil(t, err, fmt.Sprintf("Enabling config expected to succeed: %s.\n", err)) - svcCall.Unset() - list[i].Status = bootstrap.Active - active = append(active, list[i]) - } - - cases := []struct { - desc string - token string - session smqauthn.Session - url string - status int - res configPage - authenticateErr error - err error - }{ - { - desc: "view list with invalid token", - token: invalidToken, - url: fmt.Sprintf("%s?offset=%d&limit=%d", path, 0, 10), - status: http.StatusUnauthorized, - res: configPage{}, - authenticateErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "view list with an empty token", - token: "", - url: fmt.Sprintf("%s?offset=%d&limit=%d", path, 0, 10), - status: http.StatusUnauthorized, - res: configPage{}, - err: apiutil.ErrBearerToken, - }, - { - desc: "view list", - token: validToken, - url: fmt.Sprintf("%s?offset=%d&limit=%d", path, 0, 1), - status: http.StatusOK, - res: configPage{ - Total: uint64(len(list)), - Offset: 0, - Limit: 1, - Configs: list[0:1], - }, - err: nil, - }, - { - desc: "view list searching by name", - token: validToken, - url: fmt.Sprintf("%s?offset=%d&limit=%d&name=%s", path, 0, 100, "95"), - status: http.StatusOK, - res: configPage{ - Total: 1, - Offset: 0, - Limit: 100, - Configs: list[95:96], - }, - err: nil, - }, - { - desc: "view last page", - token: validToken, - url: fmt.Sprintf("%s?offset=%d&limit=%d", path, 100, 10), - status: http.StatusOK, - res: configPage{ - Total: uint64(len(list)), - Offset: 100, - Limit: 10, - Configs: list[100:], - }, - err: nil, - }, - { - desc: "view with limit greater than allowed", - token: validToken, - url: fmt.Sprintf("%s?offset=%d&limit=%d", path, 0, 1000), - status: http.StatusBadRequest, - res: configPage{}, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "view list with no specified limit and offset", - token: validToken, - url: path, - status: http.StatusOK, - res: configPage{ - Total: uint64(len(list)), - Offset: 0, - Limit: 10, - Configs: list[0:10], - }, - err: nil, - }, - { - desc: "view list with no specified limit", - token: validToken, - url: fmt.Sprintf("%s?offset=%d", path, 10), - status: http.StatusOK, - res: configPage{ - Total: uint64(len(list)), - Offset: 10, - Limit: 10, - Configs: list[10:20], - }, - err: nil, - }, - { - desc: "view list with no specified offset", - token: validToken, - url: fmt.Sprintf("%s?limit=%d", path, 10), - status: http.StatusOK, - res: configPage{ - Total: uint64(len(list)), - Offset: 0, - Limit: 10, - Configs: list[0:10], - }, - err: nil, - }, - { - desc: "view list with limit < 0", - token: validToken, - url: fmt.Sprintf("%s?limit=%d", path, -10), - status: http.StatusBadRequest, - res: configPage{}, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "view list with offset < 0", - token: validToken, - url: fmt.Sprintf("%s?offset=%d", path, -10), - status: http.StatusBadRequest, - res: configPage{}, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "view list with invalid query parameters", - token: validToken, - url: fmt.Sprintf("%s?offset=%d&limit=%d&status=%s&key=%%", path, 10, 10, bootstrap.Disabled), - status: http.StatusBadRequest, - res: configPage{}, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "view first 10 active", - token: validToken, - url: fmt.Sprintf("%s?offset=%d&limit=%d&status=%s", path, 0, 20, bootstrap.Enabled), - status: http.StatusOK, - res: configPage{ - Total: uint64(len(active)), - Offset: 0, - Limit: 20, - Configs: active, - }, - err: nil, - }, - { - desc: "view first 10 inactive", - token: validToken, - url: fmt.Sprintf("%s?offset=%d&limit=%d&status=%s", path, 0, 20, bootstrap.Disabled), - status: http.StatusOK, - res: configPage{ - Total: uint64(len(list) - len(inactive)), - Offset: 0, - Limit: 20, - Configs: inactive, - }, - err: nil, - }, - { - desc: "view first 5 active", - token: validToken, - url: fmt.Sprintf("%s?offset=%d&limit=%d&status=%s", path, 0, 10, bootstrap.Enabled), - status: http.StatusOK, - res: configPage{ - Total: uint64(len(active)), - Offset: 0, - Limit: 10, - Configs: active[:5], - }, - err: nil, - }, - { - desc: "view last 5 inactive", - token: validToken, - url: fmt.Sprintf("%s?offset=%d&limit=%d&status=%s", path, 10, 10, bootstrap.Disabled), - status: http.StatusOK, - res: configPage{ - Total: uint64(len(list) - len(active)), - Offset: 10, - Limit: 10, - Configs: inactive[5:], - }, - err: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := svc.On("List", mock.Anything, mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(bootstrap.ConfigsPage{Total: tc.res.Total, Offset: tc.res.Offset, Limit: tc.res.Limit}, tc.err) - req := testRequest{ - client: bs.Client(), - method: http.MethodGet, - url: tc.url, - token: tc.token, - } - - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - var body configPage - - err = json.NewDecoder(res.Body).Decode(&body) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - - assert.Equal(t, tc.res.Total, body.Total, fmt.Sprintf("%s: expected response total '%d' got '%d'", tc.desc, tc.res.Total, body.Total)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestRemove(t *testing.T) { - bs, svc, auth := newBootstrapServer() - defer bs.Close() - c := newConfig() - - cases := []struct { - desc string - id string - token string - session smqauthn.Session - status int - authenticateErr error - err error - }{ - { - desc: "remove with invalid token", - id: c.ID, - token: invalidToken, - status: http.StatusUnauthorized, - authenticateErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "remove with an empty token", - id: c.ID, - token: "", - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "remove non-existing config", - id: "non-existing", - token: validToken, - status: http.StatusNoContent, - err: nil, - }, - { - desc: "remove config", - id: c.ID, - token: validToken, - status: http.StatusNoContent, - err: nil, - }, - { - desc: "remove removed config", - id: wrongID, - token: validToken, - status: http.StatusNoContent, - err: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := svc.On("Remove", mock.Anything, mock.Anything, mock.Anything).Return(tc.err) - req := testRequest{ - client: bs.Client(), - method: http.MethodDelete, - url: fmt.Sprintf("%s/%s/clients/configs/%s", bs.URL, domainID, tc.id), - token: tc.token, - } - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestBootstrap(t *testing.T) { - bs, svc, _ := newBootstrapServer() - defer bs.Close() - c := newConfig() - - encExternKey, err := enc([]byte(c.ExternalKey)) - assert.Nil(t, err, fmt.Sprintf("Encrypting config expected to succeed: %s.\n", err)) - - s := struct { - ID string `json:"id"` - Content string `json:"content"` - ClientCert string `json:"client_cert"` - ClientKey string `json:"client_key"` - CACert string `json:"ca_cert"` - }{ - ID: c.ID, - Content: c.Content, - ClientCert: c.ClientCert, - ClientKey: c.ClientKey, - CACert: c.CACert, - } - - data := toJSON(s) - - cases := []struct { - desc string - externalID string - externalKey string - status int - res string - secure bool - err error - }{ - { - desc: "bootstrap a Client with unknown ID", - externalID: unknown, - externalKey: c.ExternalKey, - status: http.StatusNotFound, - res: unknownExternalIDErrorRes, - secure: false, - err: svcerr.ErrNotFound, - }, - { - desc: "bootstrap a Client with an empty ID", - externalID: "", - externalKey: c.ExternalKey, - status: http.StatusBadRequest, - res: missingIDRes, - secure: false, - err: apiutil.ErrMissingID, - }, - { - desc: "bootstrap a Client with unknown key", - externalID: c.ExternalID, - externalKey: unknown, - status: http.StatusForbidden, - res: extKeyRes, - secure: false, - err: bootstrap.ErrExternalKey, - }, - { - desc: "bootstrap a Client with an empty key", - externalID: c.ExternalID, - externalKey: "", - status: http.StatusUnauthorized, - res: missingKeyRes, - secure: false, - err: apiutil.ErrBearerKey, - }, - { - desc: "bootstrap known Client", - externalID: c.ExternalID, - externalKey: c.ExternalKey, - status: http.StatusOK, - res: data, - secure: false, - err: nil, - }, - { - desc: "bootstrap secure", - externalID: fmt.Sprintf("secure/%s", c.ExternalID), - externalKey: hex.EncodeToString(encExternKey), - status: http.StatusOK, - res: data, - secure: true, - err: nil, - }, - { - desc: "bootstrap secure with unencrypted key", - externalID: fmt.Sprintf("secure/%s", c.ExternalID), - externalKey: c.ExternalKey, - status: http.StatusForbidden, - res: extSecKeyRes, - secure: true, - err: bootstrap.ErrExternalKeySecure, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("Bootstrap", mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(c, tc.err) - req := testRequest{ - client: bs.Client(), - method: http.MethodGet, - url: fmt.Sprintf("%s/clients/bootstrap/%s", bs.URL, tc.externalID), - key: tc.externalKey, - } - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - body, err := io.ReadAll(res.Body) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - if tc.secure && tc.status == http.StatusOK { - body, err = dec(body) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding body: %s", tc.desc, err)) - } - data := strings.Trim(string(body), "\n") - assert.Equal(t, tc.res, data, fmt.Sprintf("%s: expected response '%s' got '%s'", tc.desc, tc.res, data)) - svcCall.Unset() - }) - } -} - -func TestChangeStatus(t *testing.T) { - bs, svc, auth := newBootstrapServer() - defer bs.Close() - c := newConfig() - - activeCfg := bootstrap.Config{ID: c.ID, Status: bootstrap.Active} - inactiveCfg := bootstrap.Config{ID: c.ID, Status: bootstrap.Inactive} - - cases := []struct { - desc string - id string - token string - session smqauthn.Session - action string - status int - authenticateErr error - svcCfg bootstrap.Config - svcErr error - }{ - { - desc: "enable with invalid token", - id: c.ID, - token: invalidToken, - action: "enable", - status: http.StatusUnauthorized, - authenticateErr: svcerr.ErrAuthentication, - }, - { - desc: "enable with empty token", - id: c.ID, - token: "", - action: "enable", - status: http.StatusUnauthorized, - }, - { - desc: "enable config", - id: c.ID, - token: validToken, - action: "enable", - status: http.StatusOK, - svcCfg: activeCfg, - }, - { - desc: "disable config", - id: c.ID, - token: validToken, - action: "disable", - status: http.StatusOK, - svcCfg: inactiveCfg, - }, - { - desc: "enable non-existing config", - id: wrongID, - token: validToken, - action: "enable", - status: http.StatusNotFound, - svcErr: svcerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - methodName := "EnableConfig" - if tc.action == "disable" { - methodName = "DisableConfig" - } - svcCall := svc.On(methodName, mock.Anything, tc.session, mock.Anything).Return(tc.svcCfg, tc.svcErr) - req := testRequest{ - client: bs.Client(), - method: http.MethodPost, - url: fmt.Sprintf("%s/%s/clients/configs/%s/%s", bs.URL, domainID, tc.id, tc.action), - token: tc.token, - } - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestUploadProfile(t *testing.T) { - bs, svc, auth := newBootstrapServer() - defer bs.Close() - - session := smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - saved := bootstrap.Profile{ - ID: testsutil.GenerateUUID(t), - Name: "gateway", - ContentFormat: bootstrap.ContentFormatGoTemplate, - ContentTemplate: "{{ .Device.ID }}", - } - - cases := []struct { - desc string - contentType string - body string - profile bootstrap.Profile - }{ - { - desc: "upload JSON profile", - contentType: "application/json", - body: `{"name":"gateway","content_format":"go-template","content_template":"{{ .Device.ID }}"}`, - profile: bootstrap.Profile{ - Name: "gateway", - ContentFormat: bootstrap.ContentFormatGoTemplate, - ContentTemplate: "{{ .Device.ID }}", - }, - }, - { - desc: "upload YAML profile", - contentType: "application/yaml", - body: "name: gateway\ncontent_format: go-template\ncontent_template: '{{ .Device.ID }}'\n", - profile: bootstrap.Profile{ - Name: "gateway", - ContentFormat: bootstrap.ContentFormatGoTemplate, - ContentTemplate: "{{ .Device.ID }}", - }, - }, - { - desc: "upload TOML profile", - contentType: "application/toml", - body: "name = 'gateway'\ncontent_format = 'go-template'\ncontent_template = '{{ .Device.ID }}'\n", - profile: bootstrap.Profile{ - Name: "gateway", - ContentFormat: bootstrap.ContentFormatGoTemplate, - ContentTemplate: "{{ .Device.ID }}", - }, - }, - { - desc: "upload JSON profile without content_format infers json", - contentType: "application/json", - body: `{"name":"gateway","content_template":"{{ .Device.ID }}"}`, - profile: bootstrap.Profile{ - Name: "gateway", - ContentFormat: bootstrap.ContentFormatJSON, - ContentTemplate: "{{ .Device.ID }}", - }, - }, - { - desc: "upload YAML profile without content_format infers yaml", - contentType: "application/yaml", - body: "name: gateway\ncontent_template: '{{ .Device.ID }}'\n", - profile: bootstrap.Profile{ - Name: "gateway", - ContentFormat: bootstrap.ContentFormatYAML, - ContentTemplate: "{{ .Device.ID }}", - }, - }, - { - desc: "upload TOML profile without content_format infers toml", - contentType: "application/toml", - body: "name = 'gateway'\ncontent_template = '{{ .Device.ID }}'\n", - profile: bootstrap.Profile{ - Name: "gateway", - ContentFormat: bootstrap.ContentFormatTOML, - ContentTemplate: "{{ .Device.ID }}", - }, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - authCall := auth.On("Authenticate", mock.Anything, validToken).Return(session, nil) - svcCall := svc.On("CreateProfile", mock.Anything, session, tc.profile).Return(saved, nil) - req := testRequest{ - client: bs.Client(), - method: http.MethodPost, - url: fmt.Sprintf("%s/%s/clients/bootstrap/profiles/upload", bs.URL, domainID), - contentType: tc.contentType, - token: validToken, - body: strings.NewReader(tc.body), - } - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - assert.Equal(t, http.StatusCreated, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, http.StatusCreated, res.StatusCode)) - assert.Equal(t, "/bootstrap/profiles/"+saved.ID, res.Header.Get("Location")) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestListProfiles(t *testing.T) { - bs, svc, auth := newBootstrapServer() - defer bs.Close() - - session := smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - path := fmt.Sprintf("%s/%s/clients/bootstrap/profiles", bs.URL, domainID) - - profiles := []bootstrap.Profile{ - {ID: testsutil.GenerateUUID(t), DomainID: domainID, Name: "gateway-profile"}, - {ID: testsutil.GenerateUUID(t), DomainID: domainID, Name: "sensor-profile"}, - } - fullPage := bootstrap.ProfilesPage{Total: 2, Offset: 0, Limit: 10, Profiles: profiles} - filteredPage := bootstrap.ProfilesPage{Total: 1, Offset: 0, Limit: 10, Profiles: profiles[:1]} - - cases := []struct { - desc string - token string - session smqauthn.Session - url string - name string - svcPage bootstrap.ProfilesPage - svcErr error - authenticateErr error - status int - }{ - { - desc: "list profiles successfully", - token: validToken, - session: session, - url: fmt.Sprintf("%s?offset=0&limit=10", path), - svcPage: fullPage, - status: http.StatusOK, - }, - { - desc: "list profiles filtered by name", - token: validToken, - session: session, - url: fmt.Sprintf("%s?offset=0&limit=10&name=gateway-profile", path), - name: "gateway-profile", - svcPage: filteredPage, - status: http.StatusOK, - }, - { - desc: "list profiles with invalid token", - token: invalidToken, - url: fmt.Sprintf("%s?offset=0&limit=10", path), - authenticateErr: svcerr.ErrAuthentication, - status: http.StatusUnauthorized, - }, - { - desc: "list profiles with limit exceeding max", - token: validToken, - session: session, - url: fmt.Sprintf("%s?offset=0&limit=101", path), - status: http.StatusBadRequest, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := svc.On("ListProfiles", mock.Anything, tc.session, mock.Anything, mock.Anything, tc.name).Return(tc.svcPage, tc.svcErr) - req := testRequest{ - client: bs.Client(), - method: http.MethodGet, - url: tc.url, - token: tc.token, - } - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status %d got %d", tc.desc, tc.status, res.StatusCode)) - authCall.Unset() - svcCall.Unset() - }) - } -} - -func TestProfileSlots(t *testing.T) { - bs, svc, auth := newBootstrapServer() - defer bs.Close() - - session := smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - profileID := testsutil.GenerateUUID(t) - slots := []bootstrap.BindingSlot{ - {Name: "mqtt_client", Type: "client", Required: true, Fields: []string{"id", "secret"}}, - {Name: "telemetry", Type: "channel", Required: true, Fields: []string{"id", "topic"}}, - } - profile := bootstrap.Profile{ - ID: profileID, - Name: "gateway", - BindingSlots: slots, - } - authCall := auth.On("Authenticate", mock.Anything, validToken).Return(session, nil) - svcCall := svc.On("ViewProfile", mock.Anything, session, profileID).Return(profile, nil) - - req := testRequest{ - client: bs.Client(), - method: http.MethodGet, - url: fmt.Sprintf("%s/%s/clients/bootstrap/profiles/%s/slots", bs.URL, domainID, profileID), - token: validToken, - } - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("profile slots unexpected error %s", err)) - assert.Equal(t, http.StatusOK, res.StatusCode, fmt.Sprintf("expected status code %d got %d", http.StatusOK, res.StatusCode)) - - var got struct { - BindingSlots []bootstrap.BindingSlot `json:"binding_slots"` - } - err = json.NewDecoder(res.Body).Decode(&got) - assert.Nil(t, err, fmt.Sprintf("decoding profile slots expected to succeed: %s", err)) - assert.ElementsMatch(t, slots, got.BindingSlots) - - svcCall.Unset() - authCall.Unset() -} - -func TestRenderPreview(t *testing.T) { - bs, svc, auth := newBootstrapServer() - defer bs.Close() - - session := smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - profileID := testsutil.GenerateUUID(t) - configID := testsutil.GenerateUUID(t) - profile := bootstrap.Profile{ - ID: profileID, - Name: "gateway", - ContentFormat: bootstrap.ContentFormatGoTemplate, - ContentTemplate: `device={{ .Device.ID }} site={{ .Vars.site }} topic={{ index (index .Bindings "telemetry").Snapshot "topic" }}`, - } - - storedConfig := bootstrap.Config{ - ID: configID, - ExternalID: "gw-001", - DomainID: domainID, - RenderContext: map[string]any{ - "site": "warehouse-1", - }, - } - storedBindings := []bootstrap.BindingSnapshot{ - { - Slot: "telemetry", - Type: "channel", - ResourceID: "ch-1", - Snapshot: map[string]any{"topic": "devices/gw-001/telemetry"}, - }, - } - - inlineReqBody := struct { - Config bootstrap.Config `json:"config"` - Bindings []bootstrap.BindingSnapshot `json:"bindings"` - }{ - Config: bootstrap.Config{ - ID: configID, - ExternalID: "gw-001", - RenderContext: map[string]any{"site": "warehouse-1"}, - }, - Bindings: storedBindings, - } - - configIDReqBody := struct { - ConfigID string `json:"config_id"` - }{ - ConfigID: configID, - } - - expectedContent := "device=" + configID + " site=warehouse-1 topic=devices/gw-001/telemetry" - - cases := []struct { - desc string - body string - profileErr error - configErr error - bindingsErr error - status int - }{ - { - desc: "render preview with inline config and bindings", - body: toJSON(inlineReqBody), - status: http.StatusOK, - }, - { - desc: "render preview with config_id loads from db", - body: toJSON(configIDReqBody), - status: http.StatusOK, - }, - { - desc: "render preview with config_id and config not found", - body: toJSON(configIDReqBody), - configErr: svcerr.ErrNotFound, - status: http.StatusNotFound, - }, - { - desc: "render preview with config_id and bindings error", - body: toJSON(configIDReqBody), - bindingsErr: svcerr.ErrViewEntity, - status: http.StatusUnprocessableEntity, - }, - { - desc: "render preview with profile not found", - body: toJSON(inlineReqBody), - profileErr: svcerr.ErrNotFound, - status: http.StatusNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - authCall := auth.On("Authenticate", mock.Anything, validToken).Return(session, nil) - svcCall := svc.On("ViewProfile", mock.Anything, session, profileID).Return(profile, tc.profileErr) - svcCall2 := svc.On("View", mock.Anything, session, configID).Return(storedConfig, tc.configErr) - svcCall3 := svc.On("ListBindings", mock.Anything, session, configID).Return(storedBindings, tc.bindingsErr) - - req := testRequest{ - client: bs.Client(), - method: http.MethodPost, - url: fmt.Sprintf("%s/%s/clients/bootstrap/profiles/%s/render-preview", bs.URL, domainID, profileID), - contentType: contentType, - token: validToken, - body: strings.NewReader(tc.body), - } - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - - if tc.status == http.StatusOK { - var got struct { - Content string `json:"content"` - } - err = json.NewDecoder(res.Body).Decode(&got) - assert.Nil(t, err, fmt.Sprintf("%s: decoding expected to succeed: %s", tc.desc, err)) - assert.Equal(t, expectedContent, got.Content, fmt.Sprintf("%s: expected content %q got %q", tc.desc, expectedContent, got.Content)) - } - - svcCall3.Unset() - svcCall2.Unset() - svcCall.Unset() - authCall.Unset() - }) - } -} - -type config struct { - ID string `json:"id,omitempty"` - ExternalID string `json:"external_id"` - Content string `json:"content,omitempty"` - Name string `json:"name"` - Status bootstrap.Status `json:"status"` -} - -type configPage struct { - Total uint64 `json:"total"` - Offset uint64 `json:"offset"` - Limit uint64 `json:"limit"` - Configs []config `json:"configs"` -} diff --git a/bootstrap/api/requests.go b/bootstrap/api/requests.go deleted file mode 100644 index 1863539ec..000000000 --- a/bootstrap/api/requests.go +++ /dev/null @@ -1,280 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package api - -import ( - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/bootstrap" -) - -const maxLimitSize = 100 - -type addReq struct { - token string - ExternalID string `json:"external_id"` - ExternalKey string `json:"external_key"` - Name string `json:"name"` - Content string `json:"content"` - ClientCert string `json:"client_cert"` - ClientKey string `json:"client_key"` - CACert string `json:"ca_cert"` - ProfileID string `json:"profile_id"` - RenderContext map[string]any `json:"render_context"` -} - -func (req addReq) validate() error { - if req.token == "" { - return apiutil.ErrBearerToken - } - - if req.ExternalID == "" { - return apiutil.ErrMissingID - } - - if req.ExternalKey == "" { - return apiutil.ErrBearerKey - } - - return nil -} - -type entityReq struct { - id string -} - -func (req entityReq) validate() error { - if req.id == "" { - return apiutil.ErrMissingID - } - - return nil -} - -type updateReq struct { - id string - Name string `json:"name"` - Content string `json:"content"` - RenderContext map[string]any `json:"render_context"` -} - -func (req updateReq) validate() error { - if req.id == "" { - return apiutil.ErrMissingID - } - - return nil -} - -type updateCertReq struct { - configID string - ClientCert string `json:"client_cert"` - ClientKey string `json:"client_key"` - CACert string `json:"ca_cert"` -} - -func (req updateCertReq) validate() error { - if req.configID == "" { - return apiutil.ErrMissingID - } - - return nil -} - -type listReq struct { - filter bootstrap.Filter - offset uint64 - limit uint64 -} - -func (req listReq) validate() error { - if req.limit > maxLimitSize { - return apiutil.ErrLimitSize - } - - return nil -} - -type bootstrapReq struct { - key string - id string -} - -func (req bootstrapReq) validate() error { - if req.key == "" { - return apiutil.ErrBearerKey - } - - if req.id == "" { - return apiutil.ErrMissingID - } - - return nil -} - -type changeConfigStatusReq struct { - token string - id string -} - -func (req changeConfigStatusReq) validate() error { - if req.token == "" { - return apiutil.ErrBearerToken - } - - if req.id == "" { - return apiutil.ErrMissingID - } - - return nil -} - -// --- Profile requests --- - -type createProfileReq struct { - bootstrap.Profile -} - -func (req createProfileReq) validate() error { - if req.Name == "" { - return apiutil.ErrMissingName - } - return nil -} - -type uploadProfileReq struct { - bootstrap.Profile -} - -func (req uploadProfileReq) validate() error { - if req.Name == "" { - return apiutil.ErrMissingName - } - return nil -} - -type viewProfileReq struct { - profileID string -} - -func (req viewProfileReq) validate() error { - if req.profileID == "" { - return apiutil.ErrMissingID - } - return nil -} - -type updateProfileReq struct { - profileID string - bootstrap.Profile -} - -func (req updateProfileReq) validate() error { - if req.profileID == "" { - return apiutil.ErrMissingID - } - return nil -} - -type renderPreviewReq struct { - profileID string - ConfigID string `json:"config_id,omitempty"` - Config bootstrap.Config `json:"config"` - RenderContext map[string]any `json:"render_context,omitempty"` - Bindings []bootstrap.BindingSnapshot `json:"bindings,omitempty"` -} - -func (req renderPreviewReq) validate() error { - if req.profileID == "" { - return apiutil.ErrMissingID - } - return nil -} - -type deleteProfileReq struct { - profileID string -} - -func (req deleteProfileReq) validate() error { - if req.profileID == "" { - return apiutil.ErrMissingID - } - return nil -} - -type listProfilesReq struct { - offset uint64 - limit uint64 - name string -} - -func (req listProfilesReq) validate() error { - if req.limit == 0 || req.limit > maxLimitSize { - return apiutil.ErrLimitSize - } - return nil -} - -// --- Enrollment binding requests --- - -type assignProfileReq struct { - configID string - ProfileID string `json:"profile_id"` -} - -func (req assignProfileReq) validate() error { - if req.configID == "" || req.ProfileID == "" { - return apiutil.ErrMissingID - } - return nil -} - -type bindResourcesReq struct { - token string - configID string - Bindings []bootstrap.BindingRequest `json:"bindings"` -} - -func (req bindResourcesReq) validate() error { - if req.token == "" { - return apiutil.ErrBearerToken - } - if req.configID == "" { - return apiutil.ErrMissingID - } - if len(req.Bindings) == 0 { - return apiutil.ErrEmptyList - } - for _, b := range req.Bindings { - if b.Slot == "" || b.Type == "" || b.ResourceID == "" { - return apiutil.ErrMissingID - } - } - return nil -} - -type listBindingsReq struct { - configID string -} - -func (req listBindingsReq) validate() error { - if req.configID == "" { - return apiutil.ErrMissingID - } - return nil -} - -type refreshBindingsReq struct { - token string - configID string -} - -func (req refreshBindingsReq) validate() error { - if req.token == "" { - return apiutil.ErrBearerToken - } - if req.configID == "" { - return apiutil.ErrMissingID - } - return nil -} diff --git a/bootstrap/api/requests_test.go b/bootstrap/api/requests_test.go deleted file mode 100644 index fcbf82aa4..000000000 --- a/bootstrap/api/requests_test.go +++ /dev/null @@ -1,245 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package api - -import ( - "fmt" - "testing" - - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/stretchr/testify/assert" -) - -func TestAddReqValidation(t *testing.T) { - cases := []struct { - desc string - token string - externalID string - externalKey string - err error - }{ - { - desc: "valid request", - token: "token", - externalID: "external-id", - externalKey: "external-key", - err: nil, - }, - { - desc: "empty token", - token: "", - externalID: "external-id", - externalKey: "external-key", - err: apiutil.ErrBearerToken, - }, - { - desc: "empty external ID", - token: "token", - externalID: "", - externalKey: "external-key", - err: apiutil.ErrMissingID, - }, - { - desc: "empty external key", - token: "token", - externalID: "external-id", - externalKey: "", - err: apiutil.ErrBearerKey, - }, - { - desc: "empty external key and external ID", - token: "token", - externalID: "", - externalKey: "", - err: apiutil.ErrMissingID, - }, - } - - for _, tc := range cases { - req := addReq{ - token: tc.token, - ExternalID: tc.externalID, - ExternalKey: tc.externalKey, - } - - err := req.validate() - assert.Equal(t, tc.err, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - } -} - -func TestEntityReqValidation(t *testing.T) { - cases := []struct { - desc string - id string - err error - }{ - { - desc: "empty id", - id: "", - err: apiutil.ErrMissingID, - }, - } - - for _, tc := range cases { - req := entityReq{ - id: tc.id, - } - - err := req.validate() - assert.Equal(t, tc.err, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - } -} - -func TestUpdateReqValidation(t *testing.T) { - cases := []struct { - desc string - id string - err error - }{ - { - desc: "valid request", - id: "id", - err: nil, - }, - { - desc: "empty id", - id: "", - err: apiutil.ErrMissingID, - }, - } - - for _, tc := range cases { - req := updateReq{ - id: tc.id, - } - - err := req.validate() - assert.Equal(t, tc.err, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - } -} - -func TestUpdateCertReqValidation(t *testing.T) { - cases := []struct { - desc string - configID string - err error - }{ - { - desc: "empty config id", - configID: "", - err: apiutil.ErrMissingID, - }, - } - - for _, tc := range cases { - req := updateCertReq{ - configID: tc.configID, - } - - err := req.validate() - assert.Equal(t, tc.err, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - } -} - -func TestListReqValidation(t *testing.T) { - cases := []struct { - desc string - offset uint64 - limit uint64 - err error - }{ - { - desc: "too large limit", - offset: 0, - limit: maxLimitSize + 1, - err: apiutil.ErrLimitSize, - }, - { - desc: "default limit", - offset: 0, - limit: defLimit, - err: nil, - }, - } - - for _, tc := range cases { - req := listReq{ - offset: tc.offset, - limit: tc.limit, - } - - err := req.validate() - assert.Equal(t, tc.err, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - } -} - -func TestBootstrapReqValidation(t *testing.T) { - cases := []struct { - desc string - externKey string - externID string - err error - }{ - { - desc: "empty external key", - externKey: "", - externID: "id", - err: apiutil.ErrBearerKey, - }, - { - desc: "empty external id", - externKey: "key", - externID: "", - err: apiutil.ErrMissingID, - }, - } - - for _, tc := range cases { - req := bootstrapReq{ - id: tc.externID, - key: tc.externKey, - } - - err := req.validate() - assert.Equal(t, tc.err, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - } -} - -func TestChangeConfigStatusReqValidation(t *testing.T) { - cases := []struct { - desc string - token string - id string - err error - }{ - { - desc: "empty token", - token: "", - id: "id", - err: apiutil.ErrBearerToken, - }, - { - desc: "empty id", - token: "token", - id: "", - err: apiutil.ErrMissingID, - }, - { - desc: "valid request", - token: "token", - id: "id", - err: nil, - }, - } - - for _, tc := range cases { - req := changeConfigStatusReq{ - token: tc.token, - id: tc.id, - } - - err := req.validate() - assert.Equal(t, tc.err, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - } -} diff --git a/bootstrap/api/responses.go b/bootstrap/api/responses.go deleted file mode 100644 index 3b0c5031e..000000000 --- a/bootstrap/api/responses.go +++ /dev/null @@ -1,223 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package api - -import ( - "fmt" - "net/http" - - "github.com/absmach/magistrala" - "github.com/absmach/magistrala/bootstrap" -) - -var ( - _ magistrala.Response = (*removeRes)(nil) - _ magistrala.Response = (*configRes)(nil) - _ magistrala.Response = (*changeConfigStatusRes)(nil) - _ magistrala.Response = (*viewRes)(nil) - _ magistrala.Response = (*listRes)(nil) -) - -type removeRes struct{} - -func (res removeRes) Code() int { - return http.StatusNoContent -} - -func (res removeRes) Headers() map[string]string { - return map[string]string{} -} - -func (res removeRes) Empty() bool { - return true -} - -type updateRes struct{} - -func (res updateRes) Code() int { - return http.StatusOK -} - -func (res updateRes) Headers() map[string]string { - return map[string]string{} -} - -func (res updateRes) Empty() bool { - return true -} - -type configRes struct { - ID string `json:"id"` - ExternalID string `json:"external_id"` - Name string `json:"name,omitempty"` - Content string `json:"content,omitempty"` - Status bootstrap.Status `json:"status"` - ProfileID string `json:"profile_id,omitempty"` - RenderContext map[string]any `json:"render_context,omitempty"` - ClientCert string `json:"client_cert,omitempty"` - CACert string `json:"ca_cert,omitempty"` - ClientKey string `json:"client_key,omitempty"` - created bool -} - -func (res configRes) Code() int { - if res.created { - return http.StatusCreated - } - - return http.StatusOK -} - -func (res configRes) Headers() map[string]string { - if res.created { - return map[string]string{ - "Location": fmt.Sprintf("/clients/configs/%s", res.ID), - } - } - - return map[string]string{} -} - -func (res configRes) Empty() bool { - return false -} - -type viewRes struct { - ID string `json:"id,omitempty"` - ExternalID string `json:"external_id"` - Content string `json:"content,omitempty"` - Name string `json:"name,omitempty"` - Status bootstrap.Status `json:"status"` - ProfileID string `json:"profile_id,omitempty"` - RenderContext map[string]any `json:"render_context,omitempty"` - ClientCert string `json:"client_cert,omitempty"` - CACert string `json:"ca_cert,omitempty"` - ClientKey string `json:"client_key,omitempty"` -} - -func (res viewRes) Code() int { - return http.StatusOK -} - -func (res viewRes) Headers() map[string]string { - return map[string]string{} -} - -func (res viewRes) Empty() bool { - return false -} - -type listRes struct { - Total uint64 `json:"total"` - Offset uint64 `json:"offset"` - Limit uint64 `json:"limit"` - Configs []viewRes `json:"configs"` -} - -func (res listRes) Code() int { - return http.StatusOK -} - -func (res listRes) Headers() map[string]string { - return map[string]string{} -} - -func (res listRes) Empty() bool { - return false -} - -type changeConfigStatusRes struct { - bootstrap.Config -} - -func (res changeConfigStatusRes) Code() int { - return http.StatusOK -} - -func (res changeConfigStatusRes) Headers() map[string]string { - return map[string]string{} -} - -func (res changeConfigStatusRes) Empty() bool { - return false -} - -type updateConfigRes struct { - ID string `json:"id,omitempty"` - CACert string `json:"ca_cert,omitempty"` - ClientCert string `json:"client_cert,omitempty"` - ClientKey string `json:"client_key,omitempty"` -} - -func (res updateConfigRes) Code() int { - return http.StatusOK -} - -func (res updateConfigRes) Headers() map[string]string { - return map[string]string{} -} - -func (res updateConfigRes) Empty() bool { - return false -} - -// profileRes is returned on create (201) or update (200). -type profileRes struct { - bootstrap.Profile - created bool -} - -func (res profileRes) Code() int { - if res.created { - return http.StatusCreated - } - return http.StatusOK -} - -func (res profileRes) Headers() map[string]string { - if res.created { - return map[string]string{ - "Location": fmt.Sprintf("/bootstrap/profiles/%s", res.ID), - } - } - return map[string]string{} -} - -func (res profileRes) Empty() bool { return false } - -// profilesPageRes is returned by ListProfiles. -type profilesPageRes struct { - bootstrap.ProfilesPage -} - -func (res profilesPageRes) Code() int { return http.StatusOK } -func (res profilesPageRes) Headers() map[string]string { return map[string]string{} } -func (res profilesPageRes) Empty() bool { return false } - -// profileSlotsRes is returned by profile slots endpoint. -type profileSlotsRes struct { - BindingSlots []bootstrap.BindingSlot `json:"binding_slots"` -} - -func (res profileSlotsRes) Code() int { return http.StatusOK } -func (res profileSlotsRes) Headers() map[string]string { return map[string]string{} } -func (res profileSlotsRes) Empty() bool { return false } - -// renderPreviewRes is returned by profile render-preview endpoint. -type renderPreviewRes struct { - Content string `json:"content"` -} - -func (res renderPreviewRes) Code() int { return http.StatusOK } -func (res renderPreviewRes) Headers() map[string]string { return map[string]string{} } -func (res renderPreviewRes) Empty() bool { return false } - -// bindingsRes is returned by ListBindings. -type bindingsRes struct { - Bindings []bootstrap.BindingSnapshot `json:"bindings"` -} - -func (res bindingsRes) Code() int { return http.StatusOK } -func (res bindingsRes) Headers() map[string]string { return map[string]string{} } -func (res bindingsRes) Empty() bool { return false } diff --git a/bootstrap/api/transport.go b/bootstrap/api/transport.go deleted file mode 100644 index 9d6a1ed0d..000000000 --- a/bootstrap/api/transport.go +++ /dev/null @@ -1,512 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package api - -import ( - "context" - "encoding/json" - "io" - "log/slog" - "net/http" - "net/url" - "strings" - - "github.com/absmach/magistrala" - api "github.com/absmach/magistrala/api/http" - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/bootstrap" - smqauthn "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/errors" - "github.com/go-chi/chi/v5" - kithttp "github.com/go-kit/kit/transport/http" - "github.com/pelletier/go-toml/v2" - "github.com/prometheus/client_golang/prometheus/promhttp" - "go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp" - "gopkg.in/yaml.v3" -) - -const ( - contentType = "application/json" - yamlContentType = "yaml" - tomlContentType = "toml" - byteContentType = "application/octet-stream" - offsetKey = "offset" - limitKey = "limit" - defOffset = 0 - defLimit = 10 -) - -var ( - fullMatch = []string{"status", "external_id", "id"} - partialMatch = []string{"name"} - // ErrBootstrap indicates error in getting bootstrap configuration. - ErrBootstrap = errors.New("failed to read bootstrap configuration") -) - -// MakeHandler returns a HTTP handler for API endpoints. -func MakeHandler(svc bootstrap.Service, authn smqauthn.AuthNMiddleware, reader bootstrap.ConfigReader, logger *slog.Logger, instanceID string) http.Handler { - opts := []kithttp.ServerOption{ - kithttp.ServerErrorEncoder(apiutil.LoggingErrorEncoder(logger, api.EncodeError)), - } - - r := chi.NewRouter() - - r.Route("/{domainID}/clients", func(r chi.Router) { - r.Group(func(r chi.Router) { - r.Use(authn.WithOptions(smqauthn.WithDomainCheck(true)).Middleware()) - r.Route("/configs", func(r chi.Router) { - r.Post("/", otelhttp.NewHandler(kithttp.NewServer( - addEndpoint(svc), - decodeAddRequest, - api.EncodeResponse, - opts...), "add").ServeHTTP) - - r.Get("/", otelhttp.NewHandler(kithttp.NewServer( - listEndpoint(svc), - decodeListRequest, - api.EncodeResponse, - opts...), "list").ServeHTTP) - - r.Get("/{configID}", otelhttp.NewHandler(kithttp.NewServer( - viewEndpoint(svc), - decodeEntityRequest, - api.EncodeResponse, - opts...), "view").ServeHTTP) - - r.Patch("/{configID}", otelhttp.NewHandler(kithttp.NewServer( - updateEndpoint(svc), - decodeUpdateRequest, - api.EncodeResponse, - opts...), "update").ServeHTTP) - - r.Delete("/{configID}", otelhttp.NewHandler(kithttp.NewServer( - removeEndpoint(svc), - decodeEntityRequest, - api.EncodeResponse, - opts...), "remove").ServeHTTP) - - r.Patch("/certs/{configID}", otelhttp.NewHandler(kithttp.NewServer( - updateCertEndpoint(svc), - decodeUpdateCertRequest, - api.EncodeResponse, - opts...), "update_cert").ServeHTTP) - - r.Post("/{configID}/enable", otelhttp.NewHandler(kithttp.NewServer( - enableConfigEndpoint(svc), - decodeChangeConfigStatusRequest, - api.EncodeResponse, - opts...), "enable_config").ServeHTTP) - - r.Post("/{configID}/disable", otelhttp.NewHandler(kithttp.NewServer( - disableConfigEndpoint(svc), - decodeChangeConfigStatusRequest, - api.EncodeResponse, - opts...), "disable_config").ServeHTTP) - }) - }) - - // Profile and enrollment binding endpoints. - r.Route("/bootstrap", func(r chi.Router) { - r.Use(authn.WithOptions(smqauthn.WithDomainCheck(true)).Middleware()) - - r.Route("/profiles", func(r chi.Router) { - r.Post("/", otelhttp.NewHandler(kithttp.NewServer( - createProfileEndpoint(svc), - decodeCreateProfileRequest, - api.EncodeResponse, - opts...), "create_profile").ServeHTTP) - - r.Post("/upload", otelhttp.NewHandler(kithttp.NewServer( - uploadProfileEndpoint(svc), - decodeUploadProfileRequest, - api.EncodeResponse, - opts...), "upload_profile").ServeHTTP) - - r.Get("/", otelhttp.NewHandler(kithttp.NewServer( - listProfilesEndpoint(svc), - decodeListProfilesRequest, - api.EncodeResponse, - opts...), "list_profiles").ServeHTTP) - - r.Get("/{profileID}", otelhttp.NewHandler(kithttp.NewServer( - viewProfileEndpoint(svc), - decodeProfileEntityRequest, - api.EncodeResponse, - opts...), "view_profile").ServeHTTP) - - r.Get("/{profileID}/slots", otelhttp.NewHandler(kithttp.NewServer( - profileSlotsEndpoint(svc), - decodeProfileEntityRequest, - api.EncodeResponse, - opts...), "profile_slots").ServeHTTP) - - r.Post("/{profileID}/render-preview", otelhttp.NewHandler(kithttp.NewServer( - renderPreviewEndpoint(svc), - decodeRenderPreviewRequest, - api.EncodeResponse, - opts...), "render_preview").ServeHTTP) - - r.Patch("/{profileID}", otelhttp.NewHandler(kithttp.NewServer( - updateProfileEndpoint(svc), - decodeUpdateProfileRequest, - api.EncodeResponse, - opts...), "update_profile").ServeHTTP) - - r.Delete("/{profileID}", otelhttp.NewHandler(kithttp.NewServer( - deleteProfileEndpoint(svc), - decodeDeleteProfileRequest, - api.EncodeResponse, - opts...), "delete_profile").ServeHTTP) - }) - - r.Route("/enrollments", func(r chi.Router) { - r.Patch("/{configID}/profile", otelhttp.NewHandler(kithttp.NewServer( - assignProfileEndpoint(svc), - decodeAssignProfileRequest, - api.EncodeResponse, - opts...), "assign_profile").ServeHTTP) - - r.Put("/{configID}/bindings", otelhttp.NewHandler(kithttp.NewServer( - bindResourcesEndpoint(svc), - decodeBindResourcesRequest, - api.EncodeResponse, - opts...), "bind_resources").ServeHTTP) - - r.Get("/{configID}/bindings", otelhttp.NewHandler(kithttp.NewServer( - listBindingsEndpoint(svc), - decodeEnrollmentEntityRequest, - api.EncodeResponse, - opts...), "list_bindings").ServeHTTP) - - r.Post("/{configID}/bindings/refresh", otelhttp.NewHandler(kithttp.NewServer( - refreshBindingsEndpoint(svc), - decodeRefreshBindingsRequest, - api.EncodeResponse, - opts...), "refresh_bindings").ServeHTTP) - }) - }) - }) - - r.Route("/clients/bootstrap", func(r chi.Router) { - r.Get("/", otelhttp.NewHandler(kithttp.NewServer( - bootstrapEndpoint(svc, reader, false), - decodeBootstrapRequest, - api.EncodeResponse, - opts...), "bootstrap").ServeHTTP) - r.Get("/{externalID}", otelhttp.NewHandler(kithttp.NewServer( - bootstrapEndpoint(svc, reader, false), - decodeBootstrapRequest, - api.EncodeResponse, - opts...), "bootstrap").ServeHTTP) - r.Get("/secure/{externalID}", otelhttp.NewHandler(kithttp.NewServer( - bootstrapEndpoint(svc, reader, true), - decodeBootstrapRequest, - encodeSecureRes, - opts...), "bootstrap_secure").ServeHTTP) - }) - - r.Get("/health", magistrala.Health("bootstrap", instanceID)) - r.Handle("/metrics", promhttp.Handler()) - - return r -} - -func decodeAddRequest(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), contentType) { - return nil, apiutil.ErrUnsupportedContentType - } - - req := addReq{ - token: apiutil.ExtractBearerToken(r), - } - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - - return req, nil -} - -func decodeUpdateRequest(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), contentType) { - return nil, apiutil.ErrUnsupportedContentType - } - - req := updateReq{ - id: chi.URLParam(r, "configID"), - } - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - - return req, nil -} - -func decodeUpdateCertRequest(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), contentType) { - return nil, apiutil.ErrUnsupportedContentType - } - - req := updateCertReq{ - configID: chi.URLParam(r, "configID"), - } - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - - return req, nil -} - -func decodeListRequest(_ context.Context, r *http.Request) (any, error) { - o, err := apiutil.ReadNumQuery[uint64](r, offsetKey, defOffset) - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - l, err := apiutil.ReadNumQuery[uint64](r, limitKey, defLimit) - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - q, err := url.ParseQuery(r.URL.RawQuery) - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrInvalidQueryParams) - } - - req := listReq{ - filter: parseFilter(q), - offset: o, - limit: l, - } - - rawStatus := q.Get("status") - parsed, err := bootstrap.ToStatus(rawStatus) - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrInvalidQueryParams) - } - if parsed == bootstrap.AllStatus { - delete(req.filter.FullMatch, "status") - } else { - req.filter.FullMatch["status"] = parsed.String() - } - - return req, nil -} - -func decodeBootstrapRequest(_ context.Context, r *http.Request) (any, error) { - req := bootstrapReq{ - id: chi.URLParam(r, "externalID"), - key: apiutil.ExtractClientSecret(r), - } - - return req, nil -} - -func decodeChangeConfigStatusRequest(_ context.Context, r *http.Request) (any, error) { - return changeConfigStatusReq{ - token: apiutil.ExtractBearerToken(r), - id: chi.URLParam(r, "configID"), - }, nil -} - -func decodeEntityRequest(_ context.Context, r *http.Request) (any, error) { - req := entityReq{ - id: chi.URLParam(r, "configID"), - } - - return req, nil -} - -func encodeSecureRes(_ context.Context, w http.ResponseWriter, response any) error { - w.Header().Set("Content-Type", byteContentType) - w.WriteHeader(http.StatusOK) - if b, ok := response.([]byte); ok { - if _, err := w.Write(b); err != nil { - return err - } - } - return nil -} - -func parseFilter(values url.Values) bootstrap.Filter { - ret := bootstrap.Filter{ - FullMatch: make(map[string]string), - PartialMatch: make(map[string]string), - } - for k := range values { - if contains(fullMatch, k) { - ret.FullMatch[k] = values.Get(k) - } - if contains(partialMatch, k) { - ret.PartialMatch[k] = strings.ToLower(values.Get(k)) - } - } - - return ret -} - -func contains(l []string, s string) bool { - for _, v := range l { - if v == s { - return true - } - } - return false -} - -func decodeCreateProfileRequest(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), contentType) { - return nil, apiutil.ErrUnsupportedContentType - } - var req createProfileReq - if err := json.NewDecoder(r.Body).Decode(&req.Profile); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - return req, nil -} - -func decodeUploadProfileRequest(_ context.Context, r *http.Request) (any, error) { - contentType := r.Header.Get("Content-Type") - var req uploadProfileReq - var inferredFormat bootstrap.ContentFormat - - switch { - case strings.Contains(contentType, "json"): - inferredFormat = bootstrap.ContentFormatJSON - if err := json.NewDecoder(r.Body).Decode(&req.Profile); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - case strings.Contains(contentType, yamlContentType): - inferredFormat = bootstrap.ContentFormatYAML - body, err := io.ReadAll(r.Body) - if err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - if err := decodeYAMLProfile(body, &req.Profile); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - case strings.Contains(contentType, tomlContentType): - inferredFormat = bootstrap.ContentFormatTOML - body, err := io.ReadAll(r.Body) - if err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - if err := decodeTOMLProfile(body, &req.Profile); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - default: - return nil, apiutil.ErrUnsupportedContentType - } - - if req.Profile.ContentFormat == "" { - req.Profile.ContentFormat = inferredFormat - } - - return req, nil -} - -func decodeYAMLProfile(body []byte, profile *bootstrap.Profile) error { - var raw map[string]any - if err := yaml.Unmarshal(body, &raw); err != nil { - return err - } - return decodeProfileMap(raw, profile) -} - -func decodeTOMLProfile(body []byte, profile *bootstrap.Profile) error { - var raw map[string]any - if err := toml.Unmarshal(body, &raw); err != nil { - return err - } - return decodeProfileMap(raw, profile) -} - -func decodeProfileMap(raw map[string]any, profile *bootstrap.Profile) error { - body, err := json.Marshal(raw) - if err != nil { - return err - } - return json.Unmarshal(body, profile) -} - -func decodeListProfilesRequest(_ context.Context, r *http.Request) (any, error) { - o, err := apiutil.ReadNumQuery[uint64](r, offsetKey, defOffset) - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - l, err := apiutil.ReadNumQuery[uint64](r, limitKey, defLimit) - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - n, err := apiutil.ReadStringQuery(r, api.NameKey, "") - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - return listProfilesReq{offset: o, limit: l, name: n}, nil -} - -func decodeProfileEntityRequest(_ context.Context, r *http.Request) (any, error) { - return viewProfileReq{profileID: chi.URLParam(r, "profileID")}, nil -} - -func decodeDeleteProfileRequest(_ context.Context, r *http.Request) (any, error) { - return deleteProfileReq{profileID: chi.URLParam(r, "profileID")}, nil -} - -func decodeUpdateProfileRequest(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), contentType) { - return nil, apiutil.ErrUnsupportedContentType - } - req := updateProfileReq{profileID: chi.URLParam(r, "profileID")} - if err := json.NewDecoder(r.Body).Decode(&req.Profile); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - return req, nil -} - -func decodeRenderPreviewRequest(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), contentType) { - return nil, apiutil.ErrUnsupportedContentType - } - req := renderPreviewReq{profileID: chi.URLParam(r, "profileID")} - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - return req, nil -} - -func decodeAssignProfileRequest(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), contentType) { - return nil, apiutil.ErrUnsupportedContentType - } - req := assignProfileReq{configID: chi.URLParam(r, "configID")} - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - return req, nil -} - -func decodeBindResourcesRequest(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), contentType) { - return nil, apiutil.ErrUnsupportedContentType - } - req := bindResourcesReq{ - token: apiutil.ExtractBearerToken(r), - configID: chi.URLParam(r, "configID"), - } - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - return req, nil -} - -func decodeEnrollmentEntityRequest(_ context.Context, r *http.Request) (any, error) { - return listBindingsReq{configID: chi.URLParam(r, "configID")}, nil -} - -func decodeRefreshBindingsRequest(_ context.Context, r *http.Request) (any, error) { - return refreshBindingsReq{ - token: apiutil.ExtractBearerToken(r), - configID: chi.URLParam(r, "configID"), - }, nil -} diff --git a/bootstrap/binding_validation.go b/bootstrap/binding_validation.go deleted file mode 100644 index 34bcbe0fd..000000000 --- a/bootstrap/binding_validation.go +++ /dev/null @@ -1,109 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package bootstrap - -import ( - "fmt" - "text/template" - - "github.com/absmach/magistrala/pkg/errors" -) - -var errBindingSlot = errors.New("invalid binding slot") - -func validateProfileBindingSlots(profile Profile) error { - seen := make(map[string]struct{}, len(profile.BindingSlots)) - for _, slot := range profile.BindingSlots { - if slot.Name == "" { - return fmt.Errorf("%w: slot name is required", errBindingSlot) - } - if slot.Type == "" { - return fmt.Errorf("%w: slot %q type is required", errBindingSlot, slot.Name) - } - if _, ok := seen[slot.Name]; ok { - return fmt.Errorf("%w: duplicate slot %q", errBindingSlot, slot.Name) - } - seen[slot.Name] = struct{}{} - } - return nil -} - -func validateRequestedBindings(profile Profile, requested []BindingRequest) error { - if len(profile.BindingSlots) == 0 { - return nil - } - - slots := make(map[string]BindingSlot, len(profile.BindingSlots)) - for _, slot := range profile.BindingSlots { - slots[slot.Name] = slot - } - - seen := make(map[string]struct{}, len(requested)) - for _, binding := range requested { - slot, ok := slots[binding.Slot] - if !ok { - return fmt.Errorf("%w: unknown slot %q", errBindingSlot, binding.Slot) - } - if slot.Type != binding.Type { - return fmt.Errorf("%w: slot %q expects %q, got %q", errBindingSlot, binding.Slot, slot.Type, binding.Type) - } - if _, ok := seen[binding.Slot]; ok { - return fmt.Errorf("%w: duplicate binding for slot %q", errBindingSlot, binding.Slot) - } - seen[binding.Slot] = struct{}{} - } - return nil -} - -func validateRequiredBindings(profile Profile, bindings []BindingSnapshot) error { - if len(profile.BindingSlots) == 0 { - return nil - } - - bound := make(map[string]BindingSnapshot, len(bindings)) - for _, binding := range bindings { - bound[binding.Slot] = binding - } - - for _, slot := range profile.BindingSlots { - binding, ok := bound[slot.Name] - if !slot.Required && !ok { - continue - } - if slot.Required && !ok { - return fmt.Errorf("%w: required slot %q is not bound", errBindingSlot, slot.Name) - } - if binding.Type != slot.Type { - return fmt.Errorf("%w: slot %q expects %q, got %q", errBindingSlot, slot.Name, slot.Type, binding.Type) - } - } - return nil -} - -func mergeBindingSnapshots(existing, updated []BindingSnapshot) []BindingSnapshot { - merged := make(map[string]BindingSnapshot, len(existing)+len(updated)) - for _, binding := range existing { - merged[binding.Slot] = binding - } - for _, binding := range updated { - merged[binding.Slot] = binding - } - - bindings := make([]BindingSnapshot, 0, len(merged)) - for _, binding := range merged { - bindings = append(bindings, binding) - } - return bindings -} - -func validateProfileTemplate(p Profile) error { - if p.ContentTemplate == "" || p.ContentFormat == ContentFormatRaw { - return nil - } - _, err := template.New("bootstrap").Funcs(allowlistedFuncs()).Parse(p.ContentTemplate) - if err != nil { - return errors.Wrap(ErrRenderFailed, err) - } - return nil -} diff --git a/bootstrap/bindings.go b/bootstrap/bindings.go deleted file mode 100644 index be5eb034b..000000000 --- a/bootstrap/bindings.go +++ /dev/null @@ -1,81 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package bootstrap - -import ( - "context" - "time" -) - -// BindingRequest carries a user's intent to bind a named profile slot to -// a concrete resource. -type BindingRequest struct { - Slot string `json:"slot"` - Type string `json:"type"` // "client" | "channel" | "cert" - ResourceID string `json:"resource_id"` // ID of the resource in its owning service -} - -// BindingSnapshot is a Bootstrap-owned point-in-time copy of the resource -// fields needed for template rendering. It is populated at binding time so -// that the render path never calls external services. -type BindingSnapshot struct { - ConfigID string `json:"config_id"` - Slot string `json:"slot"` - Type string `json:"type"` - ResourceID string `json:"resource_id"` - Snapshot map[string]any `json:"snapshot,omitempty"` - SecretSnapshot map[string]any `json:"secret_snapshot,omitempty"` // encrypted at rest - UpdatedAt time.Time `json:"updated_at,omitempty"` -} - -// BindingStore is the persistence interface for BindingSnapshots. -type BindingStore interface { - // Save upserts all given snapshots for the config. - Save(ctx context.Context, configID string, bindings []BindingSnapshot) error - - // Retrieve returns all snapshots for the given config. - Retrieve(ctx context.Context, configID string) ([]BindingSnapshot, error) - - // Delete removes the snapshot for a specific slot of a config. - Delete(ctx context.Context, configID, slot string) error -} - -// ResolveRequest carries everything the BindingResolver needs to snapshot a -// set of resource bindings. -type ResolveRequest struct { - Enrollment Config - Token string - Requested []BindingRequest -} - -// BindingResolver validates that requested resources exist in their owning -// services, verifies type and slot compatibility, and returns snapshots ready -// for storage. It is called at binding time only; the render path must not -// call it. -type BindingResolver interface { - Resolve(ctx context.Context, req ResolveRequest) ([]BindingSnapshot, error) -} - -// RenderContext is the typed value injected into Go templates during rendering. -type RenderContext struct { - Device DeviceContext - Vars map[string]any - Bindings map[string]BindingContext -} - -// DeviceContext holds enrollment identity fields available inside templates. -type DeviceContext struct { - ID string - ExternalID string - DomainID string -} - -// BindingContext holds the resolved resource data available inside templates -// for a specific slot. -type BindingContext struct { - Type string - ID string - Snapshot map[string]any - Secret map[string]any -} diff --git a/bootstrap/configs.go b/bootstrap/configs.go deleted file mode 100644 index d2f8b8f5b..000000000 --- a/bootstrap/configs.go +++ /dev/null @@ -1,73 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package bootstrap - -import "context" - -// Config represents a bootstrap enrollment. -type Config struct { - ID string `json:"id"` - DomainID string `json:"domain_id,omitempty"` - Name string `json:"name,omitempty"` - ClientCert string `json:"client_cert,omitempty"` - ClientKey string `json:"client_key,omitempty"` - CACert string `json:"ca_cert,omitempty"` - ExternalID string `json:"external_id"` - ExternalKey string `json:"external_key"` - Content string `json:"content,omitempty"` - Status Status `json:"status"` - ProfileID string `json:"profile_id,omitempty"` - RenderContext map[string]any `json:"render_context,omitempty"` -} - -// Filter is used for the search filters. -type Filter struct { - FullMatch map[string]string - PartialMatch map[string]string -} - -// ConfigsPage contains page related metadata as well as list of Configs that -// belong to this page. -type ConfigsPage struct { - Total uint64 `json:"total"` - Offset uint64 `json:"offset"` - Limit uint64 `json:"limit"` - Configs []Config `json:"configs"` -} - -// ConfigRepository specifies a Config persistence API. -type ConfigRepository interface { - // Save persists the Config. Successful operation is indicated by non-nil - // error response. - Save(ctx context.Context, cfg Config) (string, error) - - // RetrieveByID retrieves the Config having the provided identifier, that is owned - // by the specified user. - RetrieveByID(ctx context.Context, domainID, id string) (Config, error) - - // RetrieveAll retrieves a subset of Configs that belong to the given domain, - // with given filter parameters. - RetrieveAll(ctx context.Context, domainID string, filter Filter, offset, limit uint64) ConfigsPage - - // RetrieveByExternalID returns Config for given external ID. - RetrieveByExternalID(ctx context.Context, externalID string) (Config, error) - - // Update updates an existing Config. A non-nil error is returned - // to indicate operation failure. - Update(ctx context.Context, cfg Config) error - - // AssignProfile sets the profile reference for the given Config. - AssignProfile(ctx context.Context, domainID, id, profileID string) error - - // UpdateCerts updates and returns an existing Config certificate and domainID. - // A non-nil error is returned to indicate operation failure. - UpdateCert(ctx context.Context, domainID, id, clientCert, clientKey, caCert string) (Config, error) - - // Remove removes the Config having the provided identifier, that is owned - // by the specified user. - Remove(ctx context.Context, domainID, id string) error - - // ChangeStatus changes the Status of the Config owned by the specific user. - ChangeStatus(ctx context.Context, domainID, id string, status Status) error -} diff --git a/bootstrap/doc.go b/bootstrap/doc.go deleted file mode 100644 index 606c44a9e..000000000 --- a/bootstrap/doc.go +++ /dev/null @@ -1,6 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package bootstrap contains the domain concept definitions needed to support -// Magistrala bootstrap service functionality. -package bootstrap diff --git a/bootstrap/events/doc.go b/bootstrap/events/doc.go deleted file mode 100644 index fa65f5af2..000000000 --- a/bootstrap/events/doc.go +++ /dev/null @@ -1,6 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package events provides the domain concept definitions needed to support -// bootstrap events functionality. -package events diff --git a/bootstrap/events/producer/doc.go b/bootstrap/events/producer/doc.go deleted file mode 100644 index ab1537514..000000000 --- a/bootstrap/events/producer/doc.go +++ /dev/null @@ -1,6 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package producer contains the domain events needed to support -// event sourcing of Bootstrap service actions. -package producer diff --git a/bootstrap/events/producer/events.go b/bootstrap/events/producer/events.go deleted file mode 100644 index e6893dce5..000000000 --- a/bootstrap/events/producer/events.go +++ /dev/null @@ -1,288 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package producer - -import ( - "github.com/absmach/magistrala/bootstrap" - "github.com/absmach/magistrala/pkg/events" -) - -const ( - configPrefix = "bootstrap.config." - configCreate = configPrefix + "create" - configUpdate = configPrefix + "update" - configRemove = configPrefix + "remove" - configView = configPrefix + "view" - configList = configPrefix + "list" - clientPrefix = "bootstrap.client." - clientBootstrap = clientPrefix + "bootstrap" - configEnable = configPrefix + "enable" - configDisable = configPrefix + "disable" - certUpdate = "bootstrap.cert.update" - - profilePrefix = "bootstrap.profile." - profileCreate = profilePrefix + "create" - profileView = profilePrefix + "view" - profileUpdate = profilePrefix + "update" - profileList = profilePrefix + "list" - profileDelete = profilePrefix + "delete" - profileAssign = profilePrefix + "assign" - bindingsPrefix = "bootstrap.bindings." - bindingsBind = bindingsPrefix + "bind" - bindingsList = bindingsPrefix + "list" - bindingsRefresh = bindingsPrefix + "refresh" -) - -var ( - _ events.Event = (*configEvent)(nil) - _ events.Event = (*removeConfigEvent)(nil) - _ events.Event = (*bootstrapEvent)(nil) - _ events.Event = (*enableConfigEvent)(nil) - _ events.Event = (*disableConfigEvent)(nil) - _ events.Event = (*updateCertEvent)(nil) - _ events.Event = (*listConfigsEvent)(nil) - _ events.Event = (*profileEvent)(nil) - _ events.Event = (*deleteProfileEvent)(nil) - _ events.Event = (*assignProfileEvent)(nil) - _ events.Event = (*bindResourcesEvent)(nil) - _ events.Event = (*listBindingsEvent)(nil) - _ events.Event = (*refreshBindingsEvent)(nil) -) - -type configEvent struct { - bootstrap.Config - operation string -} - -func (ce configEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "status": ce.Status.String(), - "operation": ce.operation, - } - if ce.ID != "" { - val["config_id"] = ce.ID - } - if ce.Content != "" { - val["content"] = ce.Content - } - if ce.DomainID != "" { - val["domain_id"] = ce.DomainID - } - if ce.Name != "" { - val["name"] = ce.Name - } - if ce.ExternalID != "" { - val["external_id"] = ce.ExternalID - } - if ce.ClientCert != "" { - val["client_cert"] = ce.ClientCert - } - if ce.ClientKey != "" { - val["client_key"] = ce.ClientKey - } - if ce.CACert != "" { - val["ca_cert"] = ce.CACert - } - if ce.Content != "" { - val["content"] = ce.Content - } - - return val, nil -} - -type removeConfigEvent struct { - config string -} - -func (rce removeConfigEvent) Encode() (map[string]any, error) { - return map[string]any{ - "config_id": rce.config, - "operation": configRemove, - }, nil -} - -type listConfigsEvent struct { - offset uint64 - limit uint64 - fullMatch map[string]string - partialMatch map[string]string -} - -func (rce listConfigsEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "offset": rce.offset, - "limit": rce.limit, - "operation": configList, - } - if len(rce.fullMatch) > 0 { - val["full_match"] = rce.fullMatch - } - - if len(rce.partialMatch) > 0 { - val["full_match"] = rce.partialMatch - } - return val, nil -} - -type bootstrapEvent struct { - bootstrap.Config - externalID string - success bool -} - -func (be bootstrapEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "external_id": be.externalID, - "success": be.success, - "operation": clientBootstrap, - } - - if be.ID != "" { - val["config_id"] = be.ID - } - if be.Content != "" { - val["content"] = be.Content - } - if be.DomainID != "" { - val["domain_id"] = be.DomainID - } - if be.Name != "" { - val["name"] = be.Name - } - if be.ExternalID != "" { - val["external_id"] = be.ExternalID - } - if be.ClientCert != "" { - val["client_cert"] = be.ClientCert - } - if be.ClientKey != "" { - val["client_key"] = be.ClientKey - } - if be.CACert != "" { - val["ca_cert"] = be.CACert - } - if be.Content != "" { - val["content"] = be.Content - } - return val, nil -} - -type enableConfigEvent struct { - configID string -} - -func (e enableConfigEvent) Encode() (map[string]any, error) { - return map[string]any{ - "config_id": e.configID, - "operation": configEnable, - }, nil -} - -type disableConfigEvent struct { - configID string -} - -func (e disableConfigEvent) Encode() (map[string]any, error) { - return map[string]any{ - "config_id": e.configID, - "operation": configDisable, - }, nil -} - -type updateCertEvent struct { - configID string - clientCert string - clientKey string - caCert string -} - -func (uce updateCertEvent) Encode() (map[string]any, error) { - return map[string]any{ - "config_id": uce.configID, - "client_cert": uce.clientCert, - "client_key": uce.clientKey, - "ca_cert": uce.caCert, - "operation": certUpdate, - }, nil -} - -type profileEvent struct { - bootstrap.Profile - operation string -} - -func (pe profileEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": pe.operation, - } - if pe.ID != "" { - val["profile_id"] = pe.ID - } - if pe.DomainID != "" { - val["domain_id"] = pe.DomainID - } - if pe.Name != "" { - val["name"] = pe.Name - } - return val, nil -} - -type deleteProfileEvent struct { - profileID string -} - -func (dpe deleteProfileEvent) Encode() (map[string]any, error) { - return map[string]any{ - "profile_id": dpe.profileID, - "operation": profileDelete, - }, nil -} - -type assignProfileEvent struct { - configID string - profileID string -} - -func (ape assignProfileEvent) Encode() (map[string]any, error) { - return map[string]any{ - "config_id": ape.configID, - "profile_id": ape.profileID, - "operation": profileAssign, - }, nil -} - -type bindResourcesEvent struct { - configID string - slots []string -} - -func (bre bindResourcesEvent) Encode() (map[string]any, error) { - return map[string]any{ - "config_id": bre.configID, - "slots": bre.slots, - "operation": bindingsBind, - }, nil -} - -type listBindingsEvent struct { - configID string -} - -func (lbe listBindingsEvent) Encode() (map[string]any, error) { - return map[string]any{ - "config_id": lbe.configID, - "operation": bindingsList, - }, nil -} - -type refreshBindingsEvent struct { - configID string -} - -func (rbe refreshBindingsEvent) Encode() (map[string]any, error) { - return map[string]any{ - "config_id": rbe.configID, - "operation": bindingsRefresh, - }, nil -} diff --git a/bootstrap/events/producer/setup_test.go b/bootstrap/events/producer/setup_test.go deleted file mode 100644 index 517cd652d..000000000 --- a/bootstrap/events/producer/setup_test.go +++ /dev/null @@ -1,61 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package producer_test - -import ( - "context" - "fmt" - "log" - "os" - "testing" - - "github.com/ory/dockertest/v3" - "github.com/ory/dockertest/v3/docker" - "github.com/redis/go-redis/v9" -) - -var ( - redisClient *redis.Client - redisURL string -) - -func TestMain(m *testing.M) { - pool, err := dockertest.NewPool("") - if err != nil { - log.Fatalf("Could not connect to docker: %s", err) - } - - container, err := pool.RunWithOptions(&dockertest.RunOptions{ - Repository: "redis", - Tag: "7.2.4-alpine", - }, func(config *docker.HostConfig) { - config.AutoRemove = true - config.RestartPolicy = docker.RestartPolicy{Name: "no"} - }) - if err != nil { - log.Fatalf("Could not start container: %s", err) - } - - redisURL = fmt.Sprintf("redis://localhost:%s/0", container.GetPort("6379/tcp")) - opts, err := redis.ParseURL(redisURL) - if err != nil { - log.Fatalf("Could not parse redis URL: %s", err) - } - - if err := pool.Retry(func() error { - redisClient = redis.NewClient(opts) - - return redisClient.Ping(context.Background()).Err() - }); err != nil { - log.Fatalf("Could not connect to docker: %s", err) - } - - code := m.Run() - - if err := pool.Purge(container); err != nil { - log.Fatalf("Could not purge container: %s", err) - } - - os.Exit(code) -} diff --git a/bootstrap/events/producer/streams.go b/bootstrap/events/producer/streams.go deleted file mode 100644 index f024068f6..000000000 --- a/bootstrap/events/producer/streams.go +++ /dev/null @@ -1,284 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package producer - -import ( - "context" - - "github.com/absmach/magistrala/bootstrap" - smqauthn "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/events" -) - -var _ bootstrap.Service = (*eventStore)(nil) - -const ( - magistralaPrefix = "magistrala." - createStream = magistralaPrefix + configCreate - listStream = magistralaPrefix + configList - removeStream = magistralaPrefix + configRemove - updateCertStream = magistralaPrefix + certUpdate - bootstrapStream = magistralaPrefix + clientBootstrap - enableConfigStream = magistralaPrefix + configEnable - disableConfigStream = magistralaPrefix + configDisable - createProfileStream = magistralaPrefix + profileCreate - viewProfileStream = magistralaPrefix + profileView - updateProfileStream = magistralaPrefix + profileUpdate - listProfilesStream = magistralaPrefix + profileList - deleteProfileStream = magistralaPrefix + profileDelete - assignProfileStream = magistralaPrefix + profileAssign - bindResourcesStream = magistralaPrefix + bindingsBind - listBindingsStream = magistralaPrefix + bindingsList - refreshBindingsStream = magistralaPrefix + bindingsRefresh -) - -type eventStore struct { - events.Publisher - svc bootstrap.Service -} - -// NewEventStoreMiddleware returns wrapper around bootstrap service that sends -// events to event store. -func NewEventStoreMiddleware(svc bootstrap.Service, publisher events.Publisher) bootstrap.Service { - return &eventStore{ - svc: svc, - Publisher: publisher, - } -} - -func (es *eventStore) Add(ctx context.Context, session smqauthn.Session, token string, cfg bootstrap.Config) (bootstrap.Config, error) { - saved, err := es.svc.Add(ctx, session, token, cfg) - if err != nil { - return saved, err - } - - ev := configEvent{ - saved, configCreate, - } - - if err := es.Publish(ctx, createStream, ev); err != nil { - return saved, err - } - - return saved, err -} - -func (es *eventStore) View(ctx context.Context, session smqauthn.Session, id string) (bootstrap.Config, error) { - cfg, err := es.svc.View(ctx, session, id) - if err != nil { - return cfg, err - } - ev := configEvent{ - cfg, configView, - } - - if err := es.Publish(ctx, magistralaPrefix+configView, ev); err != nil { - return cfg, err - } - - return cfg, err -} - -func (es *eventStore) Update(ctx context.Context, session smqauthn.Session, cfg bootstrap.Config) error { - if err := es.svc.Update(ctx, session, cfg); err != nil { - return err - } - - ev := configEvent{ - cfg, configUpdate, - } - - return es.Publish(ctx, magistralaPrefix+configUpdate, ev) -} - -func (es eventStore) UpdateCert(ctx context.Context, session smqauthn.Session, id, clientCert, clientKey, caCert string) (bootstrap.Config, error) { - cfg, err := es.svc.UpdateCert(ctx, session, id, clientCert, clientKey, caCert) - if err != nil { - return cfg, err - } - - ev := updateCertEvent{ - configID: id, - clientCert: clientCert, - clientKey: clientKey, - caCert: caCert, - } - - if err := es.Publish(ctx, updateCertStream, ev); err != nil { - return cfg, err - } - - return cfg, nil -} - -func (es *eventStore) List(ctx context.Context, session smqauthn.Session, filter bootstrap.Filter, offset, limit uint64) (bootstrap.ConfigsPage, error) { - bp, err := es.svc.List(ctx, session, filter, offset, limit) - if err != nil { - return bp, err - } - - ev := listConfigsEvent{ - offset: offset, - limit: limit, - fullMatch: filter.FullMatch, - partialMatch: filter.PartialMatch, - } - - if err := es.Publish(ctx, listStream, ev); err != nil { - return bp, err - } - - return bp, nil -} - -func (es *eventStore) Remove(ctx context.Context, session smqauthn.Session, id string) error { - if err := es.svc.Remove(ctx, session, id); err != nil { - return err - } - - ev := removeConfigEvent{ - config: id, - } - - return es.Publish(ctx, removeStream, ev) -} - -func (es *eventStore) Bootstrap(ctx context.Context, externalKey, externalID string, secure bool) (bootstrap.Config, error) { - cfg, err := es.svc.Bootstrap(ctx, externalKey, externalID, secure) - - ev := bootstrapEvent{ - cfg, - externalID, - true, - } - - if err != nil { - ev.success = false - } - - if err := es.Publish(ctx, bootstrapStream, ev); err != nil { - return cfg, err - } - - return cfg, err -} - -func (es *eventStore) EnableConfig(ctx context.Context, session smqauthn.Session, id string) (bootstrap.Config, error) { - cfg, err := es.svc.EnableConfig(ctx, session, id) - if err != nil { - return cfg, err - } - - ev := enableConfigEvent{configID: id} - if err := es.Publish(ctx, enableConfigStream, ev); err != nil { - return cfg, err - } - return cfg, nil -} - -func (es *eventStore) DisableConfig(ctx context.Context, session smqauthn.Session, id string) (bootstrap.Config, error) { - cfg, err := es.svc.DisableConfig(ctx, session, id) - if err != nil { - return cfg, err - } - - ev := disableConfigEvent{configID: id} - if err := es.Publish(ctx, disableConfigStream, ev); err != nil { - return cfg, err - } - return cfg, nil -} - -func (es *eventStore) CreateProfile(ctx context.Context, session smqauthn.Session, p bootstrap.Profile) (bootstrap.Profile, error) { - saved, err := es.svc.CreateProfile(ctx, session, p) - if err != nil { - return saved, err - } - ev := profileEvent{saved, profileCreate} - if err := es.Publish(ctx, createProfileStream, ev); err != nil { - return saved, err - } - return saved, nil -} - -func (es *eventStore) ViewProfile(ctx context.Context, session smqauthn.Session, profileID string) (bootstrap.Profile, error) { - p, err := es.svc.ViewProfile(ctx, session, profileID) - if err != nil { - return p, err - } - ev := profileEvent{p, profileView} - if err := es.Publish(ctx, viewProfileStream, ev); err != nil { - return p, err - } - return p, nil -} - -func (es *eventStore) UpdateProfile(ctx context.Context, session smqauthn.Session, p bootstrap.Profile) (bootstrap.Profile, error) { - updated, err := es.svc.UpdateProfile(ctx, session, p) - if err != nil { - return bootstrap.Profile{}, err - } - ev := profileEvent{updated, profileUpdate} - return updated, es.Publish(ctx, updateProfileStream, ev) -} - -func (es *eventStore) ListProfiles(ctx context.Context, session smqauthn.Session, offset, limit uint64, name string) (bootstrap.ProfilesPage, error) { - pp, err := es.svc.ListProfiles(ctx, session, offset, limit, name) - if err != nil { - return pp, err - } - ev := profileEvent{operation: profileList} - if err := es.Publish(ctx, listProfilesStream, ev); err != nil { - return pp, err - } - return pp, nil -} - -func (es *eventStore) DeleteProfile(ctx context.Context, session smqauthn.Session, profileID string) error { - if err := es.svc.DeleteProfile(ctx, session, profileID); err != nil { - return err - } - ev := deleteProfileEvent{profileID: profileID} - return es.Publish(ctx, deleteProfileStream, ev) -} - -func (es *eventStore) AssignProfile(ctx context.Context, session smqauthn.Session, configID, profileID string) error { - if err := es.svc.AssignProfile(ctx, session, configID, profileID); err != nil { - return err - } - ev := assignProfileEvent{configID: configID, profileID: profileID} - return es.Publish(ctx, assignProfileStream, ev) -} - -func (es *eventStore) BindResources(ctx context.Context, session smqauthn.Session, token, configID string, bindings []bootstrap.BindingRequest) error { - if err := es.svc.BindResources(ctx, session, token, configID, bindings); err != nil { - return err - } - slots := make([]string, len(bindings)) - for i, b := range bindings { - slots[i] = b.Slot - } - ev := bindResourcesEvent{configID: configID, slots: slots} - return es.Publish(ctx, bindResourcesStream, ev) -} - -func (es *eventStore) ListBindings(ctx context.Context, session smqauthn.Session, configID string) ([]bootstrap.BindingSnapshot, error) { - bs, err := es.svc.ListBindings(ctx, session, configID) - if err != nil { - return bs, err - } - ev := listBindingsEvent{configID: configID} - if err := es.Publish(ctx, listBindingsStream, ev); err != nil { - return bs, err - } - return bs, nil -} - -func (es *eventStore) RefreshBindings(ctx context.Context, session smqauthn.Session, token, configID string) error { - if err := es.svc.RefreshBindings(ctx, session, token, configID); err != nil { - return err - } - ev := refreshBindingsEvent{configID: configID} - return es.Publish(ctx, refreshBindingsStream, ev) -} diff --git a/bootstrap/events/producer/streams_test.go b/bootstrap/events/producer/streams_test.go deleted file mode 100644 index 5763d22db..000000000 --- a/bootstrap/events/producer/streams_test.go +++ /dev/null @@ -1,914 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package producer_test - -import ( - "context" - "fmt" - "strconv" - "strings" - "testing" - "time" - - "github.com/absmach/magistrala/bootstrap" - "github.com/absmach/magistrala/bootstrap/events/producer" - bootstraphasher "github.com/absmach/magistrala/bootstrap/hasher" - "github.com/absmach/magistrala/bootstrap/mocks" - "github.com/absmach/magistrala/internal/testsutil" - smqauthn "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/pkg/events/store" - sdkmocks "github.com/absmach/magistrala/pkg/sdk/mocks" - "github.com/absmach/magistrala/pkg/uuid" - "github.com/redis/go-redis/v9" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" - "github.com/stretchr/testify/require" -) - -const ( - streamID = "magistrala.bootstrap" - validToken = "validToken" - unknownID = "unknown" - - configPrefix = "config." - configCreate = configPrefix + "create" - configView = configPrefix + "view" - configUpdate = configPrefix + "update" - configRemove = configPrefix + "remove" - configList = configPrefix + "list" - clientPrefix = "client." - clientBootstrap = clientPrefix + "bootstrap" - configEnable = configPrefix + "enable" - configDisable = configPrefix + "disable" - - certUpdate = "cert.update" -) - -var ( - encKey = []byte("1234567891011121") - - domainID = testsutil.GenerateUUID(&testing.T{}) - validID = testsutil.GenerateUUID(&testing.T{}) - - config = bootstrap.Config{ - ID: testsutil.GenerateUUID(&testing.T{}), - ExternalID: testsutil.GenerateUUID(&testing.T{}), - ExternalKey: testsutil.GenerateUUID(&testing.T{}), - Content: "config", - Status: bootstrap.EnabledStatus, - } -) - -type testVariable struct { - svc bootstrap.Service - boot *mocks.ConfigRepository - sdk *sdkmocks.SDK -} - -func newTestVariable(t *testing.T, redisURL string) testVariable { - boot := new(mocks.ConfigRepository) - sdk := new(sdkmocks.SDK) - idp := uuid.NewMock() - svc := bootstrap.New(boot, nil, nil, nil, nil, sdk, bootstraphasher.New(), encKey, idp) - publisher, err := store.NewPublisher(context.Background(), redisURL, "bootstrap-es-pub-test") - require.Nil(t, err, fmt.Sprintf("got unexpected error: %s", err)) - svc = producer.NewEventStoreMiddleware(svc, publisher) - return testVariable{ - svc: svc, - boot: boot, - sdk: sdk, - } -} - -func TestAdd(t *testing.T) { - err := redisClient.FlushAll(context.Background()).Err() - assert.Nil(t, err, fmt.Sprintf("got unexpected error: %s", err)) - - tv := newTestVariable(t, redisURL) - - cases := []struct { - desc string - config bootstrap.Config - token string - session smqauthn.Session - id string - domainID string - saveErr error - err error - event map[string]any - }{ - { - desc: "create config successfully", - config: config, - token: validToken, - id: validID, - domainID: domainID, - event: map[string]any{ - "config_id": "1", - "domain_id": domainID, - "name": config.Name, - "external_id": config.ExternalID, - "content": config.Content, - "timestamp": time.Now().Unix(), - "operation": configCreate, - }, - err: nil, - }, - { - desc: "create config with failed to save", - config: config, - token: validToken, - id: validID, - domainID: domainID, - event: nil, - saveErr: svcerr.ErrCreateEntity, - err: svcerr.ErrCreateEntity, - }, - } - - lastID := "0" - for _, tc := range cases { - tc.session = smqauthn.Session{UserID: validID, DomainID: tc.domainID, DomainUserID: validID} - repoCall := tv.boot.On("Save", context.Background(), mock.Anything).Return(mock.Anything, tc.saveErr) - - _, err := tv.svc.Add(context.Background(), tc.session, tc.token, tc.config) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - - streams := redisClient.XRead(context.Background(), &redis.XReadArgs{ - Streams: []string{streamID, lastID}, - Count: 1, - Block: time.Second, - }).Val() - - var event map[string]any - if len(streams) > 0 && len(streams[0].Messages) > 0 { - event := streams[0].Messages - lastID = event[0].ID - } - - test(t, tc.event, event, tc.desc) - - repoCall.Unset() - } -} - -func TestView(t *testing.T) { - err := redisClient.FlushAll(context.Background()).Err() - assert.Nil(t, err, fmt.Sprintf("got unexpected error: %s", err)) - - tv := newTestVariable(t, redisURL) - - nonExisting := config - nonExisting.ID = unknownID - - cases := []struct { - desc string - config bootstrap.Config - token string - session smqauthn.Session - id string - domainID string - retrieveErr error - err error - event map[string]any - }{ - { - desc: "view successfully", - config: config, - token: validToken, - id: validID, - domainID: domainID, - err: nil, - event: map[string]any{ - "config_id": config.ID, - "domain_id": config.DomainID, - "name": config.Name, - "external_id": config.ExternalID, - "content": config.Content, - "timestamp": time.Now().Unix(), - "operation": configView, - }, - }, - { - desc: "view with failed retrieve", - config: nonExisting, - token: validToken, - id: validID, - domainID: domainID, - retrieveErr: svcerr.ErrViewEntity, - err: svcerr.ErrViewEntity, - event: nil, - }, - } - - lastID := "0" - for _, tc := range cases { - tc.session = smqauthn.Session{UserID: validID, DomainID: tc.domainID, DomainUserID: validID} - repoCall := tv.boot.On("RetrieveByID", context.Background(), tc.domainID, tc.config.ID).Return(config, tc.retrieveErr) - _, err := tv.svc.View(context.Background(), tc.session, tc.config.ID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - - streams := redisClient.XRead(context.Background(), &redis.XReadArgs{ - Streams: []string{streamID, lastID}, - Count: 1, - Block: time.Second, - }).Val() - - var event map[string]any - if len(streams) > 0 && len(streams[0].Messages) > 0 { - msg := streams[0].Messages[0] - event = msg.Values - event["timestamp"] = msg.ID - lastID = msg.ID - } - - test(t, tc.event, event, tc.desc) - repoCall.Unset() - } -} - -func TestUpdate(t *testing.T) { - err := redisClient.FlushAll(context.Background()).Err() - assert.Nil(t, err, fmt.Sprintf("got unexpected error: %s", err)) - - tv := newTestVariable(t, redisURL) - - modified := config - modified.Content = "new-config" - modified.Name = "new name" - - nonExisting := config - nonExisting.ID = unknownID - - cases := []struct { - desc string - config bootstrap.Config - token string - session smqauthn.Session - id string - domainID string - updateErr error - err error - event map[string]any - }{ - { - desc: "update config successfully", - config: modified, - token: validToken, - id: validID, - domainID: domainID, - err: nil, - event: map[string]any{ - "name": modified.Name, - "content": modified.Content, - "timestamp": time.Now().UnixNano(), - "operation": configUpdate, - "external_id": modified.ExternalID, - "config_id": modified.ID, - "domain_id": domainID, - "status": bootstrap.Disabled, - "occurred_at": time.Now().UnixNano(), - }, - }, - { - desc: "update with failed update", - config: nonExisting, - token: validToken, - id: validID, - domainID: domainID, - updateErr: svcerr.ErrNotFound, - err: svcerr.ErrNotFound, - event: nil, - }, - } - - lastID := "0" - for _, tc := range cases { - tc.session = smqauthn.Session{UserID: validID, DomainID: tc.domainID, DomainUserID: validID} - repoCall := tv.boot.On("Update", context.Background(), mock.Anything).Return(tc.updateErr) - err := tv.svc.Update(context.Background(), tc.session, tc.config) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - - streams := redisClient.XRead(context.Background(), &redis.XReadArgs{ - Streams: []string{streamID, lastID}, - Count: 1, - Block: time.Second, - }).Val() - - var event map[string]any - if len(streams) > 0 && len(streams[0].Messages) > 0 { - msg := streams[0].Messages[0] - event = msg.Values - event["timestamp"] = msg.ID - lastID = msg.ID - } - - test(t, tc.event, event, tc.desc) - repoCall.Unset() - } -} - -func TestUpdateCert(t *testing.T) { - err := redisClient.FlushAll(context.Background()).Err() - assert.Nil(t, err, fmt.Sprintf("got unexpected error: %s", err)) - - tv := newTestVariable(t, redisURL) - - cases := []struct { - desc string - configID string - userID string - domainID string - token string - session smqauthn.Session - clientCert string - clientKey string - caCert string - updateErr error - err error - event map[string]any - }{ - { - desc: "update cert successfully", - configID: config.ID, - userID: validID, - domainID: domainID, - token: validToken, - clientCert: "clientCert", - clientKey: "clientKey", - caCert: "caCert", - err: nil, - event: map[string]any{ - "client_cert": "clientCert", - "client_key": "clientKey", - "ca_cert": "caCert", - "operation": certUpdate, - }, - }, - { - desc: "update cert with failed update", - configID: "clientID", - token: validToken, - userID: validID, - domainID: domainID, - clientCert: "clientCert", - clientKey: "clientKey", - caCert: "caCert", - updateErr: svcerr.ErrNotFound, - err: svcerr.ErrNotFound, - event: nil, - }, - { - desc: "update cert with empty client certificate", - configID: config.ID, - token: validToken, - userID: validID, - domainID: domainID, - clientCert: "", - clientKey: "clientKey", - caCert: "caCert", - err: nil, - event: nil, - }, - { - desc: "update cert with empty client key", - configID: config.ID, - token: validToken, - userID: validID, - domainID: domainID, - clientCert: "clientCert", - clientKey: "", - caCert: "caCert", - err: nil, - event: nil, - }, - { - desc: "update cert with empty CA certificate", - configID: config.ID, - token: validToken, - userID: validID, - domainID: domainID, - clientCert: "clientCert", - clientKey: "clientKey", - caCert: "", - err: nil, - event: nil, - }, - } - - lastID := "0" - for _, tc := range cases { - tc.session = smqauthn.Session{UserID: tc.userID, DomainID: tc.domainID, DomainUserID: validID} - repoCall := tv.boot.On("UpdateCert", context.Background(), tc.domainID, tc.configID, tc.clientCert, tc.clientKey, tc.caCert).Return(config, tc.updateErr) - _, err := tv.svc.UpdateCert(context.Background(), tc.session, tc.configID, tc.clientCert, tc.clientKey, tc.caCert) - - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - - streams := redisClient.XRead(context.Background(), &redis.XReadArgs{ - Streams: []string{streamID, lastID}, - Count: 1, - Block: time.Second, - }).Val() - - var event map[string]any - if len(streams) > 0 && len(streams[0].Messages) > 0 { - event := streams[0].Messages - lastID = event[0].ID - } - - test(t, tc.event, event, tc.desc) - - repoCall.Unset() - } -} - -func TestList(t *testing.T) { - tv := newTestVariable(t, redisURL) - - numClients := 101 - var c bootstrap.Config - saved := make([]bootstrap.Config, 0) - for i := 0; i < numClients; i++ { - c = config - c.ExternalID = testsutil.GenerateUUID(t) - c.ExternalKey = testsutil.GenerateUUID(t) - c.Name = fmt.Sprintf("%s-%d", config.Name, i) - if i == 41 { - c.Status = bootstrap.Active - } - saved = append(saved, c) - } - - cases := []struct { - desc string - token string - session smqauthn.Session - userID string - domainID string - config bootstrap.ConfigsPage - filter bootstrap.Filter - offset uint64 - limit uint64 - retrieveErr error - err error - event map[string]any - }{ - { - desc: "list successfully as super admin", - token: validToken, - userID: validID, - domainID: domainID, - session: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: validID, SuperAdmin: true}, - config: bootstrap.ConfigsPage{ - Total: uint64(len(saved)), - Offset: 0, - Limit: 10, - Configs: saved[0:10], - }, - filter: bootstrap.Filter{}, - offset: 0, - limit: 10, - err: nil, - event: map[string]any{ - "config_id": c.ID, - "domain_id": c.DomainID, - "name": c.Name, - "external_id": c.ExternalID, - "content": c.Content, - "timestamp": time.Now().Unix(), - "operation": configList, - }, - }, - { - desc: "list successfully as domain admin", - token: validToken, - userID: validID, - domainID: domainID, - session: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: validID, SuperAdmin: true}, - config: bootstrap.ConfigsPage{ - Total: uint64(len(saved)), - Offset: 0, - Limit: 10, - Configs: saved[0:10], - }, - filter: bootstrap.Filter{}, - offset: 0, - limit: 10, - err: nil, - event: map[string]any{ - "config_id": c.ID, - "domain_id": c.DomainID, - "name": c.Name, - "external_id": c.ExternalID, - "content": c.Content, - "timestamp": time.Now().Unix(), - "operation": configList, - }, - }, - { - desc: "list successfully as non admin", - token: validToken, - userID: validID, - domainID: domainID, - session: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: validID}, - config: bootstrap.ConfigsPage{ - Total: uint64(len(saved)), - Offset: 0, - Limit: 10, - Configs: saved[0:10], - }, - filter: bootstrap.Filter{}, - offset: 0, - limit: 10, - err: nil, - event: map[string]any{ - "config_id": c.ID, - "domain_id": c.DomainID, - "name": c.Name, - "external_id": c.ExternalID, - "content": c.Content, - "timestamp": time.Now().Unix(), - "operation": configList, - }, - }, - { - desc: "list as super admin with failed retrieve all", - token: validToken, - userID: validID, - domainID: domainID, - session: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: validID, SuperAdmin: true}, - filter: bootstrap.Filter{}, - offset: 0, - limit: 10, - retrieveErr: nil, - err: nil, - event: nil, - }, - { - desc: "list as domain admin with failed retrieve all", - token: validToken, - userID: validID, - domainID: domainID, - session: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: validID, SuperAdmin: true}, - filter: bootstrap.Filter{}, - offset: 0, - limit: 10, - retrieveErr: nil, - err: nil, - event: nil, - }, - { - desc: "list as non admin with failed retrieve all", - token: validToken, - userID: validID, - domainID: domainID, - session: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: validID}, - filter: bootstrap.Filter{}, - offset: 0, - limit: 10, - retrieveErr: nil, - err: nil, - event: nil, - }, - } - - lastID := "0" - for _, tc := range cases { - repoCall := tv.boot.On("RetrieveAll", context.Background(), mock.Anything, tc.filter, tc.offset, tc.limit).Return(tc.config, tc.retrieveErr) - - _, err := tv.svc.List(context.Background(), tc.session, tc.filter, tc.offset, tc.limit) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - - streams := redisClient.XRead(context.Background(), &redis.XReadArgs{ - Streams: []string{streamID, lastID}, - Count: 1, - Block: time.Second, - }).Val() - - var event map[string]any - if len(streams) > 0 && len(streams[0].Messages) > 0 { - event := streams[0].Messages - lastID = event[0].ID - } - - test(t, tc.event, event, tc.desc) - - repoCall.Unset() - } -} - -func TestRemove(t *testing.T) { - err := redisClient.FlushAll(context.Background()).Err() - assert.Nil(t, err, fmt.Sprintf("got unexpected error: %s", err)) - - tv := newTestVariable(t, redisURL) - - nonExisting := config - nonExisting.ID = unknownID - - cases := []struct { - desc string - configID string - userID string - domainID string - token string - session smqauthn.Session - removeErr error - err error - event map[string]any - }{ - { - desc: "remove config successfully", - configID: config.ID, - token: validToken, - userID: validID, - domainID: domainID, - err: nil, - event: map[string]any{ - "config_id": config.ID, - "timestamp": time.Now().Unix(), - "operation": configRemove, - }, - }, - { - desc: "remove config with failed removal", - configID: nonExisting.ID, - token: validToken, - userID: validID, - domainID: domainID, - removeErr: svcerr.ErrNotFound, - err: svcerr.ErrNotFound, - event: nil, - }, - } - - lastID := "0" - for _, tc := range cases { - tc.session = smqauthn.Session{UserID: validID, DomainID: tc.domainID, DomainUserID: validID} - repoCall := tv.boot.On("Remove", context.Background(), mock.Anything, mock.Anything).Return(tc.removeErr) - err := tv.svc.Remove(context.Background(), tc.session, tc.configID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - - streams := redisClient.XRead(context.Background(), &redis.XReadArgs{ - Streams: []string{streamID, lastID}, - Count: 1, - Block: time.Second, - }).Val() - - var event map[string]any - if len(streams) > 0 && len(streams[0].Messages) > 0 { - event := streams[0].Messages - lastID = event[0].ID - } - - test(t, tc.event, event, tc.desc) - repoCall.Unset() - } -} - -func TestBootstrap(t *testing.T) { - err := redisClient.FlushAll(context.Background()).Err() - assert.Nil(t, err, fmt.Sprintf("got unexpected error: %s", err)) - - tv := newTestVariable(t, redisURL) - - cases := []struct { - desc string - externalID string - externalKey string - err error - retrieveErr error - event map[string]any - }{ - { - desc: "bootstrap successfully", - externalID: config.ExternalID, - externalKey: config.ExternalKey, - err: nil, - event: map[string]any{ - "external_id": config.ExternalID, - "success": "1", - "timestamp": time.Now().Unix(), - "operation": clientBootstrap, - }, - }, - { - desc: "bootstrap with an error", - externalID: "external_id1", - externalKey: "external_id", - retrieveErr: bootstrap.ErrBootstrap, - err: bootstrap.ErrBootstrap, - event: map[string]any{ - "external_id": "external_id", - "success": "0", - "timestamp": time.Now().Unix(), - "operation": clientBootstrap, - }, - }, - } - - lastID := "0" - for _, tc := range cases { - repoCall := tv.boot.On("RetrieveByExternalID", context.Background(), mock.Anything).Return(config, tc.retrieveErr) - _, err = tv.svc.Bootstrap(context.Background(), tc.externalKey, tc.externalID, false) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - - streams := redisClient.XRead(context.Background(), &redis.XReadArgs{ - Streams: []string{streamID, lastID}, - Count: 1, - Block: time.Second, - }).Val() - - var event map[string]any - if len(streams) > 0 && len(streams[0].Messages) > 0 { - event := streams[0].Messages - lastID = event[0].ID - } - test(t, tc.event, event, tc.desc) - repoCall.Unset() - } -} - -func TestEnableConfig(t *testing.T) { - err := redisClient.FlushAll(context.Background()).Err() - assert.Nil(t, err, fmt.Sprintf("got unexpected error: %s", err)) - - tv := newTestVariable(t, redisURL) - - cases := []struct { - desc string - id string - userID string - domainID string - session smqauthn.Session - retrieveErr error - statusErr error - err error - event map[string]any - }{ - { - desc: "enable config", - id: config.ID, - userID: validID, - domainID: domainID, - err: nil, - event: map[string]any{ - "config_id": config.ID, - "timestamp": time.Now().Unix(), - "operation": configEnable, - }, - }, - { - desc: "enable with failed retrieve by ID", - id: "", - userID: validID, - domainID: domainID, - retrieveErr: svcerr.ErrNotFound, - err: svcerr.ErrNotFound, - event: nil, - }, - { - desc: "enable with repo status error", - id: config.ID, - userID: validID, - domainID: domainID, - statusErr: svcerr.ErrUpdateEntity, - err: svcerr.ErrUpdateEntity, - event: nil, - }, - } - - disabledConfig := config - disabledConfig.Status = bootstrap.DisabledStatus - - lastID := "0" - for _, tc := range cases { - tc.session = smqauthn.Session{UserID: validID, DomainID: tc.domainID, DomainUserID: validID} - repoCall := tv.boot.On("RetrieveByID", context.Background(), tc.domainID, tc.id).Return(disabledConfig, tc.retrieveErr) - repoCall1 := tv.boot.On("ChangeStatus", context.Background(), mock.Anything, mock.Anything, mock.Anything).Return(tc.statusErr) - _, err := tv.svc.EnableConfig(context.Background(), tc.session, tc.id) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - - streams := redisClient.XRead(context.Background(), &redis.XReadArgs{ - Streams: []string{streamID, lastID}, - Count: 1, - Block: time.Second, - }).Val() - - var event map[string]any - if len(streams) > 0 && len(streams[0].Messages) > 0 { - event := streams[0].Messages - lastID = event[0].ID - } - - test(t, tc.event, event, tc.desc) - repoCall.Unset() - repoCall1.Unset() - } -} - -func TestDisableConfig(t *testing.T) { - err := redisClient.FlushAll(context.Background()).Err() - assert.Nil(t, err, fmt.Sprintf("got unexpected error: %s", err)) - - tv := newTestVariable(t, redisURL) - - cases := []struct { - desc string - id string - userID string - domainID string - session smqauthn.Session - retrieveErr error - statusErr error - err error - event map[string]any - }{ - { - desc: "disable config", - id: config.ID, - userID: validID, - domainID: domainID, - err: nil, - event: map[string]any{ - "config_id": config.ID, - "timestamp": time.Now().Unix(), - "operation": configDisable, - }, - }, - { - desc: "disable with failed retrieve by ID", - id: "", - userID: validID, - domainID: domainID, - retrieveErr: svcerr.ErrNotFound, - err: svcerr.ErrNotFound, - event: nil, - }, - { - desc: "disable with repo status error", - id: config.ID, - userID: validID, - domainID: domainID, - statusErr: svcerr.ErrUpdateEntity, - err: svcerr.ErrUpdateEntity, - event: nil, - }, - } - - lastID := "0" - for _, tc := range cases { - tc.session = smqauthn.Session{UserID: validID, DomainID: tc.domainID, DomainUserID: validID} - repoCall := tv.boot.On("RetrieveByID", context.Background(), tc.domainID, tc.id).Return(config, tc.retrieveErr) - repoCall1 := tv.boot.On("ChangeStatus", context.Background(), mock.Anything, mock.Anything, mock.Anything).Return(tc.statusErr) - _, err := tv.svc.DisableConfig(context.Background(), tc.session, tc.id) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - - streams := redisClient.XRead(context.Background(), &redis.XReadArgs{ - Streams: []string{streamID, lastID}, - Count: 1, - Block: time.Second, - }).Val() - - var event map[string]any - if len(streams) > 0 && len(streams[0].Messages) > 0 { - event := streams[0].Messages - lastID = event[0].ID - } - - test(t, tc.event, event, tc.desc) - repoCall.Unset() - repoCall1.Unset() - } -} - -func test(t *testing.T, expected, actual map[string]any, description string) { - if expected != nil && actual != nil { - ts1 := expected["timestamp"].(int64) - ats := actual["timestamp"].(string) - ts2, err := strconv.ParseInt(strings.Split(ats, "-")[0], 10, 64) - require.Nil(t, err, fmt.Sprintf("%s: expected to get a valid timestamp, got %s", description, err)) - ts1 = ts1 / 1e9 - ts2 = ts2 / 1e3 - if assert.WithinDuration(t, time.Unix(ts1, 0), time.Unix(ts2, 0), time.Second, fmt.Sprintf("%s: timestamp is not in valid range of 1 second", description)) { - delete(expected, "timestamp") - delete(actual, "timestamp") - } - - oa1 := expected["occurred_at"].(int64) - aoa := actual["occurred_at"].(string) - oa2, err := strconv.ParseInt(aoa, 10, 64) - require.Nil(t, err, fmt.Sprintf("%s: expected to get a valid occurred_at, got %s", description, err)) - oa1 = oa1 / 1e9 - oa2 = oa2 / 1e9 - if assert.WithinDuration(t, time.Unix(oa1, 0), time.Unix(oa2, 0), time.Second, fmt.Sprintf("%s: occurred_at is not in valid range of 1 second", description)) { - delete(expected, "occurred_at") - delete(actual, "occurred_at") - } - - assert.Equal(t, expected, actual, fmt.Sprintf("%s: got incorrect event\n", description)) - } -} diff --git a/bootstrap/hasher.go b/bootstrap/hasher.go deleted file mode 100644 index 3460ce2f9..000000000 --- a/bootstrap/hasher.go +++ /dev/null @@ -1,14 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package bootstrap - -// Hasher specifies an API for generating hashes of arbitrary textual content. -type Hasher interface { - // Hash generates the hashed string from plain-text. - Hash(string) (string, error) - - // Compare compares plain-text version to the hashed one. An error should - // indicate failed comparison. - Compare(string, string) error -} diff --git a/bootstrap/hasher/hasher.go b/bootstrap/hasher/hasher.go deleted file mode 100644 index 51bb70037..000000000 --- a/bootstrap/hasher/hasher.go +++ /dev/null @@ -1,94 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package hasher - -import ( - "crypto/subtle" - "encoding/base64" - "strings" - - "github.com/absmach/magistrala/bootstrap" - "github.com/absmach/magistrala/pkg/errors" - "golang.org/x/crypto/bcrypt" - "golang.org/x/crypto/scrypt" -) - -const ( - cost = 10 - legacyScryptPrefix = "scrypt$" - legacyScryptKeyN = 16384 - legacyScryptKeyR = 8 - legacyScryptKeyP = 1 - legacyScryptKeySize = 32 -) - -var ( - errHashExternalKey = errors.NewServiceError("generate hash from external key failed") - errCompareExternalKey = errors.NewServiceError("compare external key and hash failed") - errInvalidHashStore = errors.New("invalid stored external key hash format") - errDecode = errors.New("failed to decode external key hash") -) - -var _ bootstrap.Hasher = (*bcryptHasher)(nil) - -type bcryptHasher struct{} - -// New instantiates a bcrypt-based hasher implementation. -func New() bootstrap.Hasher { - return &bcryptHasher{} -} - -func (*bcryptHasher) Hash(key string) (string, error) { - hash, err := bcrypt.GenerateFromPassword([]byte(key), cost) - if err != nil { - return "", errors.Wrap(errHashExternalKey, err) - } - - return string(hash), nil -} - -func (*bcryptHasher) Compare(plain, hashed string) error { - if strings.HasPrefix(hashed, legacyScryptPrefix) { - return compareLegacyScryptHash(plain, hashed) - } - - if err := bcrypt.CompareHashAndPassword([]byte(hashed), []byte(plain)); err == nil { - return nil - } - - // Legacy rows may still contain plaintext external keys. - if subtle.ConstantTimeCompare([]byte(plain), []byte(hashed)) == 1 { - return nil - } - - return bootstrap.ErrExternalKey -} - -func compareLegacyScryptHash(plain, hashed string) error { - parts := strings.Split(strings.TrimPrefix(hashed, legacyScryptPrefix), ".") - if len(parts) != 2 { - return errInvalidHashStore - } - - actualHash, err := base64.StdEncoding.DecodeString(parts[0]) - if err != nil { - return errors.Wrap(errDecode, err) - } - - salt, err := base64.StdEncoding.DecodeString(parts[1]) - if err != nil { - return errors.Wrap(errDecode, err) - } - - derivedHash, err := scrypt.Key([]byte(plain), salt, legacyScryptKeyN, legacyScryptKeyR, legacyScryptKeyP, legacyScryptKeySize) - if err != nil { - return errors.Wrap(errCompareExternalKey, err) - } - - if subtle.ConstantTimeCompare(derivedHash, actualHash) == 1 { - return nil - } - - return bootstrap.ErrExternalKey -} diff --git a/bootstrap/middleware/authorization.go b/bootstrap/middleware/authorization.go deleted file mode 100644 index 22cedf5a8..000000000 --- a/bootstrap/middleware/authorization.go +++ /dev/null @@ -1,219 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package middleware - -import ( - "context" - - "github.com/absmach/magistrala/auth" - "github.com/absmach/magistrala/bootstrap" - smqauthn "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/authz" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/pkg/policies" -) - -const ( - createOperation = "create" - viewOperation = "view" - updateOperation = "update" - updateCertOperation = "update_cert" - listOperation = "list" - removeOperation = "remove" - changeStateOperation = "change_state" -) - -var _ bootstrap.Service = (*authorizationMiddleware)(nil) - -type authorizationMiddleware struct { - svc bootstrap.Service - authz authz.Authorization -} - -// AuthorizationMiddleware adds authorization to the clients service. -func AuthorizationMiddleware(svc bootstrap.Service, authz authz.Authorization) bootstrap.Service { - return &authorizationMiddleware{ - svc: svc, - authz: authz, - } -} - -func (am *authorizationMiddleware) Add(ctx context.Context, session smqauthn.Session, token string, cfg bootstrap.Config) (bootstrap.Config, error) { - if err := am.authorize(ctx, session, "", policies.UserType, policies.UsersKind, session.DomainUserID, policies.MembershipPermission, policies.DomainType, session.DomainID, createOperation, auth.AnyIDs); err != nil { - return bootstrap.Config{}, err - } - - return am.svc.Add(ctx, session, token, cfg) -} - -func (am *authorizationMiddleware) View(ctx context.Context, session smqauthn.Session, id string) (bootstrap.Config, error) { - if err := am.authorize(ctx, session, "", policies.UserType, policies.UsersKind, session.DomainUserID, policies.MembershipPermission, policies.DomainType, session.DomainID, viewOperation, id); err != nil { - return bootstrap.Config{}, err - } - - return am.svc.View(ctx, session, id) -} - -func (am *authorizationMiddleware) Update(ctx context.Context, session smqauthn.Session, cfg bootstrap.Config) error { - if err := am.authorize(ctx, session, "", policies.UserType, policies.UsersKind, session.DomainUserID, policies.MembershipPermission, policies.DomainType, session.DomainID, updateOperation, cfg.ID); err != nil { - return err - } - - return am.svc.Update(ctx, session, cfg) -} - -func (am *authorizationMiddleware) UpdateCert(ctx context.Context, session smqauthn.Session, id, clientCert, clientKey, caCert string) (bootstrap.Config, error) { - if err := am.authorize(ctx, session, "", policies.UserType, policies.UsersKind, session.DomainUserID, policies.MembershipPermission, policies.DomainType, session.DomainID, updateCertOperation, id); err != nil { - return bootstrap.Config{}, err - } - - return am.svc.UpdateCert(ctx, session, id, clientCert, clientKey, caCert) -} - -func (am *authorizationMiddleware) List(ctx context.Context, session smqauthn.Session, filter bootstrap.Filter, offset, limit uint64) (bootstrap.ConfigsPage, error) { - if err := am.checkSuperAdmin(ctx, session); err == nil { - session.SuperAdmin = true - } - if err := am.authorize(ctx, session, "", policies.UserType, policies.UsersKind, session.DomainUserID, policies.AdminPermission, policies.DomainType, session.DomainID, listOperation, auth.AnyIDs); err == nil { - session.SuperAdmin = true - } - - return am.svc.List(ctx, session, filter, offset, limit) -} - -func (am *authorizationMiddleware) Remove(ctx context.Context, session smqauthn.Session, id string) error { - if err := am.authorize(ctx, session, "", policies.UserType, policies.UsersKind, session.DomainUserID, policies.MembershipPermission, policies.DomainType, session.DomainID, removeOperation, id); err != nil { - return err - } - - return am.svc.Remove(ctx, session, id) -} - -func (am *authorizationMiddleware) Bootstrap(ctx context.Context, externalKey, externalID string, secure bool) (bootstrap.Config, error) { - return am.svc.Bootstrap(ctx, externalKey, externalID, secure) -} - -func (am *authorizationMiddleware) EnableConfig(ctx context.Context, session smqauthn.Session, id string) (bootstrap.Config, error) { - if err := am.authorize(ctx, session, "", policies.UserType, policies.UsersKind, session.DomainUserID, policies.MembershipPermission, policies.DomainType, session.DomainID, changeStateOperation, id); err != nil { - return bootstrap.Config{}, err - } - - return am.svc.EnableConfig(ctx, session, id) -} - -func (am *authorizationMiddleware) DisableConfig(ctx context.Context, session smqauthn.Session, id string) (bootstrap.Config, error) { - if err := am.authorize(ctx, session, "", policies.UserType, policies.UsersKind, session.DomainUserID, policies.MembershipPermission, policies.DomainType, session.DomainID, changeStateOperation, id); err != nil { - return bootstrap.Config{}, err - } - - return am.svc.DisableConfig(ctx, session, id) -} - -func (am *authorizationMiddleware) CreateProfile(ctx context.Context, session smqauthn.Session, p bootstrap.Profile) (bootstrap.Profile, error) { - if err := am.authorize(ctx, session, "", policies.UserType, policies.UsersKind, session.DomainUserID, policies.MembershipPermission, policies.DomainType, session.DomainID, createOperation, auth.AnyIDs); err != nil { - return bootstrap.Profile{}, err - } - return am.svc.CreateProfile(ctx, session, p) -} - -func (am *authorizationMiddleware) ViewProfile(ctx context.Context, session smqauthn.Session, profileID string) (bootstrap.Profile, error) { - if err := am.authorize(ctx, session, "", policies.UserType, policies.UsersKind, session.DomainUserID, policies.MembershipPermission, policies.DomainType, session.DomainID, viewOperation, auth.AnyIDs); err != nil { - return bootstrap.Profile{}, err - } - return am.svc.ViewProfile(ctx, session, profileID) -} - -func (am *authorizationMiddleware) UpdateProfile(ctx context.Context, session smqauthn.Session, p bootstrap.Profile) (bootstrap.Profile, error) { - if err := am.authorize(ctx, session, "", policies.UserType, policies.UsersKind, session.DomainUserID, policies.MembershipPermission, policies.DomainType, session.DomainID, updateOperation, auth.AnyIDs); err != nil { - return bootstrap.Profile{}, err - } - return am.svc.UpdateProfile(ctx, session, p) -} - -func (am *authorizationMiddleware) ListProfiles(ctx context.Context, session smqauthn.Session, offset, limit uint64, name string) (bootstrap.ProfilesPage, error) { - if err := am.authorize(ctx, session, "", policies.UserType, policies.UsersKind, session.DomainUserID, policies.MembershipPermission, policies.DomainType, session.DomainID, listOperation, auth.AnyIDs); err != nil { - return bootstrap.ProfilesPage{}, err - } - return am.svc.ListProfiles(ctx, session, offset, limit, name) -} - -func (am *authorizationMiddleware) DeleteProfile(ctx context.Context, session smqauthn.Session, profileID string) error { - if err := am.authorize(ctx, session, "", policies.UserType, policies.UsersKind, session.DomainUserID, policies.MembershipPermission, policies.DomainType, session.DomainID, removeOperation, auth.AnyIDs); err != nil { - return err - } - return am.svc.DeleteProfile(ctx, session, profileID) -} - -func (am *authorizationMiddleware) AssignProfile(ctx context.Context, session smqauthn.Session, configID, profileID string) error { - if err := am.authorize(ctx, session, "", policies.UserType, policies.UsersKind, session.DomainUserID, policies.MembershipPermission, policies.DomainType, session.DomainID, updateOperation, configID); err != nil { - return err - } - return am.svc.AssignProfile(ctx, session, configID, profileID) -} - -func (am *authorizationMiddleware) BindResources(ctx context.Context, session smqauthn.Session, token, configID string, bindings []bootstrap.BindingRequest) error { - if err := am.authorize(ctx, session, "", policies.UserType, policies.UsersKind, session.DomainUserID, policies.MembershipPermission, policies.DomainType, session.DomainID, updateOperation, configID); err != nil { - return err - } - return am.svc.BindResources(ctx, session, token, configID, bindings) -} - -func (am *authorizationMiddleware) ListBindings(ctx context.Context, session smqauthn.Session, configID string) ([]bootstrap.BindingSnapshot, error) { - if err := am.authorize(ctx, session, "", policies.UserType, policies.UsersKind, session.DomainUserID, policies.MembershipPermission, policies.DomainType, session.DomainID, viewOperation, configID); err != nil { - return nil, err - } - return am.svc.ListBindings(ctx, session, configID) -} - -func (am *authorizationMiddleware) RefreshBindings(ctx context.Context, session smqauthn.Session, token, configID string) error { - if err := am.authorize(ctx, session, "", policies.UserType, policies.UsersKind, session.DomainUserID, policies.MembershipPermission, policies.DomainType, session.DomainID, updateOperation, configID); err != nil { - return err - } - return am.svc.RefreshBindings(ctx, session, token, configID) -} - -func (am *authorizationMiddleware) checkSuperAdmin(ctx context.Context, session smqauthn.Session) error { - if session.Role != smqauthn.SuperAdminRole { - return svcerr.ErrSuperAdminAction - } - if err := am.authz.Authorize(ctx, authz.PolicyReq{ - SubjectType: policies.UserType, - Subject: session.UserID, - Permission: policies.AdminPermission, - ObjectType: policies.PlatformType, - Object: policies.MagistralaObject, - }, nil); err != nil { - return err - } - return nil -} - -func (am *authorizationMiddleware) authorize(ctx context.Context, session smqauthn.Session, domain, subjType, subjKind, subj, perm, objType, obj, operation, entityID string) error { - req := authz.PolicyReq{ - Domain: domain, - SubjectType: subjType, - SubjectKind: subjKind, - Subject: subj, - Permission: perm, - ObjectType: objType, - Object: obj, - } - - var pat *authz.PATReq - if session.PatID != "" { - pat = &authz.PATReq{ - UserID: session.UserID, - PatID: session.PatID, - EntityID: entityID, - EntityType: auth.BootstrapType.String(), - Operation: operation, - Domain: session.DomainID, - } - } - - if err := am.authz.Authorize(ctx, req, pat); err != nil { - return err - } - return nil -} diff --git a/bootstrap/middleware/logging.go b/bootstrap/middleware/logging.go deleted file mode 100644 index 5be17d747..000000000 --- a/bootstrap/middleware/logging.go +++ /dev/null @@ -1,355 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -//go:build !test - -package middleware - -import ( - "context" - "log/slog" - "time" - - "github.com/absmach/magistrala/bootstrap" - smqauthn "github.com/absmach/magistrala/pkg/authn" -) - -var _ bootstrap.Service = (*loggingMiddleware)(nil) - -type loggingMiddleware struct { - logger *slog.Logger - svc bootstrap.Service -} - -// LoggingMiddleware adds logging facilities to the bootstrap service. -func LoggingMiddleware(svc bootstrap.Service, logger *slog.Logger) bootstrap.Service { - return &loggingMiddleware{logger, svc} -} - -// Add logs the add request. It logs the client ID and the time it took to complete the request. -// If the request fails, it logs the error. -func (lm *loggingMiddleware) Add(ctx context.Context, session smqauthn.Session, token string, cfg bootstrap.Config) (saved bootstrap.Config, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("config_id", saved.ID), - } - if err != nil { - args = append(args, slog.Any("error", err)) - lm.logger.Warn("Add new bootstrap failed", args...) - return - } - lm.logger.Info("Add new bootstrap completed successfully", args...) - }(time.Now()) - - return lm.svc.Add(ctx, session, token, cfg) -} - -// View logs the view request. It logs the client ID and the time it took to complete the request. -// If the request fails, it logs the error. -func (lm *loggingMiddleware) View(ctx context.Context, session smqauthn.Session, id string) (saved bootstrap.Config, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("config_id", id), - } - if err != nil { - args = append(args, slog.Any("error", err)) - lm.logger.Warn("View client config failed", args...) - return - } - lm.logger.Info("View client config completed successfully", args...) - }(time.Now()) - - return lm.svc.View(ctx, session, id) -} - -// Update logs the update request. It logs bootstrap client ID and the time it took to complete the request. -// If the request fails, it logs the error. -func (lm *loggingMiddleware) Update(ctx context.Context, session smqauthn.Session, cfg bootstrap.Config) (err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.Group("config", - slog.String("config_id", cfg.ID), - slog.String("name", cfg.Name), - ), - } - if err != nil { - args = append(args, slog.Any("error", err)) - lm.logger.Warn("Update bootstrap config failed", args...) - return - } - lm.logger.Info("Update bootstrap config completed successfully", args...) - }(time.Now()) - - return lm.svc.Update(ctx, session, cfg) -} - -// UpdateCert logs the update_cert request. It logs config ID and the time it took to complete the request. -// If the request fails, it logs the error. -func (lm *loggingMiddleware) UpdateCert(ctx context.Context, session smqauthn.Session, id, clientCert, clientKey, caCert string) (cfg bootstrap.Config, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("config_id", cfg.ID), - } - if err != nil { - args = append(args, slog.Any("error", err)) - lm.logger.Warn("Update bootstrap config certificate failed", args...) - return - } - lm.logger.Info("Update bootstrap config certificate completed successfully", args...) - }(time.Now()) - - return lm.svc.UpdateCert(ctx, session, id, clientCert, clientKey, caCert) -} - -// List logs the list request. It logs offset, limit and the time it took to complete the request. -// If the request fails, it logs the error. -func (lm *loggingMiddleware) List(ctx context.Context, session smqauthn.Session, filter bootstrap.Filter, offset, limit uint64) (res bootstrap.ConfigsPage, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.Group("page", - slog.Any("filter", filter), - slog.Uint64("offset", offset), - slog.Uint64("limit", limit), - slog.Uint64("total", res.Total), - ), - } - if err != nil { - args = append(args, slog.Any("error", err)) - lm.logger.Warn("List configs failed", args...) - return - } - lm.logger.Info("List configs completed successfully", args...) - }(time.Now()) - - return lm.svc.List(ctx, session, filter, offset, limit) -} - -// Remove logs the remove request. It logs bootstrap ID and the time it took to complete the request. -// If the request fails, it logs the error. -func (lm *loggingMiddleware) Remove(ctx context.Context, session smqauthn.Session, id string) (err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("config_id", id), - } - if err != nil { - args = append(args, slog.Any("error", err)) - lm.logger.Warn("Remove bootstrap config failed", args...) - return - } - lm.logger.Info("Remove bootstrap config completed successfully", args...) - }(time.Now()) - - return lm.svc.Remove(ctx, session, id) -} - -func (lm *loggingMiddleware) Bootstrap(ctx context.Context, externalKey, externalID string, secure bool) (cfg bootstrap.Config, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("external_id", externalID), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("View bootstrap config failed", args...) - return - } - lm.logger.Info("View bootstrap completed successfully", args...) - }(time.Now()) - - return lm.svc.Bootstrap(ctx, externalKey, externalID, secure) -} - -func (lm *loggingMiddleware) EnableConfig(ctx context.Context, session smqauthn.Session, id string) (cfg bootstrap.Config, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("id", id), - } - if err != nil { - args = append(args, slog.Any("error", err)) - lm.logger.Warn("Enable config failed", args...) - return - } - lm.logger.Info("Enable config completed successfully", args...) - }(time.Now()) - - return lm.svc.EnableConfig(ctx, session, id) -} - -func (lm *loggingMiddleware) DisableConfig(ctx context.Context, session smqauthn.Session, id string) (cfg bootstrap.Config, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("id", id), - } - if err != nil { - args = append(args, slog.Any("error", err)) - lm.logger.Warn("Disable config failed", args...) - return - } - lm.logger.Info("Disable config completed successfully", args...) - }(time.Now()) - - return lm.svc.DisableConfig(ctx, session, id) -} - -func (lm *loggingMiddleware) CreateProfile(ctx context.Context, session smqauthn.Session, p bootstrap.Profile) (saved bootstrap.Profile, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("profile_id", saved.ID), - } - if err != nil { - args = append(args, slog.Any("error", err)) - lm.logger.Warn("Create profile failed", args...) - return - } - lm.logger.Info("Create profile completed successfully", args...) - }(time.Now()) - - return lm.svc.CreateProfile(ctx, session, p) -} - -func (lm *loggingMiddleware) ViewProfile(ctx context.Context, session smqauthn.Session, profileID string) (p bootstrap.Profile, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("profile_id", profileID), - } - if err != nil { - args = append(args, slog.Any("error", err)) - lm.logger.Warn("View profile failed", args...) - return - } - lm.logger.Info("View profile completed successfully", args...) - }(time.Now()) - - return lm.svc.ViewProfile(ctx, session, profileID) -} - -func (lm *loggingMiddleware) UpdateProfile(ctx context.Context, session smqauthn.Session, p bootstrap.Profile) (updated bootstrap.Profile, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("profile_id", p.ID), - } - if err != nil { - args = append(args, slog.Any("error", err)) - lm.logger.Warn("Update profile failed", args...) - return - } - lm.logger.Info("Update profile completed successfully", args...) - }(time.Now()) - - return lm.svc.UpdateProfile(ctx, session, p) -} - -func (lm *loggingMiddleware) ListProfiles(ctx context.Context, session smqauthn.Session, offset, limit uint64, name string) (page bootstrap.ProfilesPage, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.Uint64("offset", offset), - slog.Uint64("limit", limit), - } - if err != nil { - args = append(args, slog.Any("error", err)) - lm.logger.Warn("List profiles failed", args...) - return - } - lm.logger.Info("List profiles completed successfully", args...) - }(time.Now()) - - return lm.svc.ListProfiles(ctx, session, offset, limit, name) -} - -func (lm *loggingMiddleware) DeleteProfile(ctx context.Context, session smqauthn.Session, profileID string) (err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("profile_id", profileID), - } - if err != nil { - args = append(args, slog.Any("error", err)) - lm.logger.Warn("Delete profile failed", args...) - return - } - lm.logger.Info("Delete profile completed successfully", args...) - }(time.Now()) - - return lm.svc.DeleteProfile(ctx, session, profileID) -} - -func (lm *loggingMiddleware) AssignProfile(ctx context.Context, session smqauthn.Session, configID, profileID string) (err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("config_id", configID), - slog.String("profile_id", profileID), - } - if err != nil { - args = append(args, slog.Any("error", err)) - lm.logger.Warn("Assign profile failed", args...) - return - } - lm.logger.Info("Assign profile completed successfully", args...) - }(time.Now()) - - return lm.svc.AssignProfile(ctx, session, configID, profileID) -} - -func (lm *loggingMiddleware) BindResources(ctx context.Context, session smqauthn.Session, token, configID string, bindings []bootstrap.BindingRequest) (err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("config_id", configID), - } - if err != nil { - args = append(args, slog.Any("error", err)) - lm.logger.Warn("Bind resources failed", args...) - return - } - lm.logger.Info("Bind resources completed successfully", args...) - }(time.Now()) - - return lm.svc.BindResources(ctx, session, token, configID, bindings) -} - -func (lm *loggingMiddleware) ListBindings(ctx context.Context, session smqauthn.Session, configID string) (snapshots []bootstrap.BindingSnapshot, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("config_id", configID), - } - if err != nil { - args = append(args, slog.Any("error", err)) - lm.logger.Warn("List bindings failed", args...) - return - } - lm.logger.Info("List bindings completed successfully", args...) - }(time.Now()) - - return lm.svc.ListBindings(ctx, session, configID) -} - -func (lm *loggingMiddleware) RefreshBindings(ctx context.Context, session smqauthn.Session, token, configID string) (err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("config_id", configID), - } - if err != nil { - args = append(args, slog.Any("error", err)) - lm.logger.Warn("Refresh bindings failed", args...) - return - } - lm.logger.Info("Refresh bindings completed successfully", args...) - }(time.Now()) - - return lm.svc.RefreshBindings(ctx, session, token, configID) -} diff --git a/bootstrap/middleware/metrics.go b/bootstrap/middleware/metrics.go deleted file mode 100644 index 801b5eb1b..000000000 --- a/bootstrap/middleware/metrics.go +++ /dev/null @@ -1,192 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -//go:build !test - -package middleware - -import ( - "context" - "time" - - "github.com/absmach/magistrala/bootstrap" - smqauthn "github.com/absmach/magistrala/pkg/authn" - "github.com/go-kit/kit/metrics" -) - -var _ bootstrap.Service = (*metricsMiddleware)(nil) - -type metricsMiddleware struct { - counter metrics.Counter - latency metrics.Histogram - svc bootstrap.Service -} - -// MetricsMiddleware instruments core service by tracking request count and latency. -func MetricsMiddleware(svc bootstrap.Service, counter metrics.Counter, latency metrics.Histogram) bootstrap.Service { - return &metricsMiddleware{ - counter: counter, - latency: latency, - svc: svc, - } -} - -// Add instruments Add method with metrics. -func (mm *metricsMiddleware) Add(ctx context.Context, session smqauthn.Session, token string, cfg bootstrap.Config) (saved bootstrap.Config, err error) { - defer func(begin time.Time) { - mm.counter.With("method", "add").Add(1) - mm.latency.With("method", "add").Observe(time.Since(begin).Seconds()) - }(time.Now()) - - return mm.svc.Add(ctx, session, token, cfg) -} - -// View instruments View method with metrics. -func (mm *metricsMiddleware) View(ctx context.Context, session smqauthn.Session, id string) (saved bootstrap.Config, err error) { - defer func(begin time.Time) { - mm.counter.With("method", "view").Add(1) - mm.latency.With("method", "view").Observe(time.Since(begin).Seconds()) - }(time.Now()) - - return mm.svc.View(ctx, session, id) -} - -// Update instruments Update method with metrics. -func (mm *metricsMiddleware) Update(ctx context.Context, session smqauthn.Session, cfg bootstrap.Config) (err error) { - defer func(begin time.Time) { - mm.counter.With("method", "update").Add(1) - mm.latency.With("method", "update").Observe(time.Since(begin).Seconds()) - }(time.Now()) - - return mm.svc.Update(ctx, session, cfg) -} - -// UpdateCert instruments UpdateCert method with metrics. -func (mm *metricsMiddleware) UpdateCert(ctx context.Context, session smqauthn.Session, id, clientCert, clientKey, caCert string) (cfg bootstrap.Config, err error) { - defer func(begin time.Time) { - mm.counter.With("method", "update_cert").Add(1) - mm.latency.With("method", "update_cert").Observe(time.Since(begin).Seconds()) - }(time.Now()) - - return mm.svc.UpdateCert(ctx, session, id, clientCert, clientKey, caCert) -} - -// List instruments List method with metrics. -func (mm *metricsMiddleware) List(ctx context.Context, session smqauthn.Session, filter bootstrap.Filter, offset, limit uint64) (saved bootstrap.ConfigsPage, err error) { - defer func(begin time.Time) { - mm.counter.With("method", "list").Add(1) - mm.latency.With("method", "list").Observe(time.Since(begin).Seconds()) - }(time.Now()) - - return mm.svc.List(ctx, session, filter, offset, limit) -} - -// Remove instruments Remove method with metrics. -func (mm *metricsMiddleware) Remove(ctx context.Context, session smqauthn.Session, id string) (err error) { - defer func(begin time.Time) { - mm.counter.With("method", "remove").Add(1) - mm.latency.With("method", "remove").Observe(time.Since(begin).Seconds()) - }(time.Now()) - - return mm.svc.Remove(ctx, session, id) -} - -// Bootstrap instruments Bootstrap method with metrics. -func (mm *metricsMiddleware) Bootstrap(ctx context.Context, externalKey, externalID string, secure bool) (cfg bootstrap.Config, err error) { - defer func(begin time.Time) { - mm.counter.With("method", "bootstrap").Add(1) - mm.latency.With("method", "bootstrap").Observe(time.Since(begin).Seconds()) - }(time.Now()) - - return mm.svc.Bootstrap(ctx, externalKey, externalID, secure) -} - -func (mm *metricsMiddleware) EnableConfig(ctx context.Context, session smqauthn.Session, id string) (bootstrap.Config, error) { - defer func(begin time.Time) { - mm.counter.With("method", "enable_config").Add(1) - mm.latency.With("method", "enable_config").Observe(time.Since(begin).Seconds()) - }(time.Now()) - - return mm.svc.EnableConfig(ctx, session, id) -} - -func (mm *metricsMiddleware) DisableConfig(ctx context.Context, session smqauthn.Session, id string) (bootstrap.Config, error) { - defer func(begin time.Time) { - mm.counter.With("method", "disable_config").Add(1) - mm.latency.With("method", "disable_config").Observe(time.Since(begin).Seconds()) - }(time.Now()) - - return mm.svc.DisableConfig(ctx, session, id) -} - -func (mm *metricsMiddleware) CreateProfile(ctx context.Context, session smqauthn.Session, p bootstrap.Profile) (bootstrap.Profile, error) { - defer func(begin time.Time) { - mm.counter.With("method", "create_profile").Add(1) - mm.latency.With("method", "create_profile").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return mm.svc.CreateProfile(ctx, session, p) -} - -func (mm *metricsMiddleware) ViewProfile(ctx context.Context, session smqauthn.Session, profileID string) (bootstrap.Profile, error) { - defer func(begin time.Time) { - mm.counter.With("method", "view_profile").Add(1) - mm.latency.With("method", "view_profile").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return mm.svc.ViewProfile(ctx, session, profileID) -} - -func (mm *metricsMiddleware) UpdateProfile(ctx context.Context, session smqauthn.Session, p bootstrap.Profile) (bootstrap.Profile, error) { - defer func(begin time.Time) { - mm.counter.With("method", "update_profile").Add(1) - mm.latency.With("method", "update_profile").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return mm.svc.UpdateProfile(ctx, session, p) -} - -func (mm *metricsMiddleware) ListProfiles(ctx context.Context, session smqauthn.Session, offset, limit uint64, name string) (bootstrap.ProfilesPage, error) { - defer func(begin time.Time) { - mm.counter.With("method", "list_profiles").Add(1) - mm.latency.With("method", "list_profiles").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return mm.svc.ListProfiles(ctx, session, offset, limit, name) -} - -func (mm *metricsMiddleware) DeleteProfile(ctx context.Context, session smqauthn.Session, profileID string) error { - defer func(begin time.Time) { - mm.counter.With("method", "delete_profile").Add(1) - mm.latency.With("method", "delete_profile").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return mm.svc.DeleteProfile(ctx, session, profileID) -} - -func (mm *metricsMiddleware) AssignProfile(ctx context.Context, session smqauthn.Session, configID, profileID string) error { - defer func(begin time.Time) { - mm.counter.With("method", "assign_profile").Add(1) - mm.latency.With("method", "assign_profile").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return mm.svc.AssignProfile(ctx, session, configID, profileID) -} - -func (mm *metricsMiddleware) BindResources(ctx context.Context, session smqauthn.Session, token, configID string, bindings []bootstrap.BindingRequest) error { - defer func(begin time.Time) { - mm.counter.With("method", "bind_resources").Add(1) - mm.latency.With("method", "bind_resources").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return mm.svc.BindResources(ctx, session, token, configID, bindings) -} - -func (mm *metricsMiddleware) ListBindings(ctx context.Context, session smqauthn.Session, configID string) ([]bootstrap.BindingSnapshot, error) { - defer func(begin time.Time) { - mm.counter.With("method", "list_bindings").Add(1) - mm.latency.With("method", "list_bindings").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return mm.svc.ListBindings(ctx, session, configID) -} - -func (mm *metricsMiddleware) RefreshBindings(ctx context.Context, session smqauthn.Session, token, configID string) error { - defer func(begin time.Time) { - mm.counter.With("method", "refresh_bindings").Add(1) - mm.latency.With("method", "refresh_bindings").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return mm.svc.RefreshBindings(ctx, session, token, configID) -} diff --git a/bootstrap/mocks/binding_resolver.go b/bootstrap/mocks/binding_resolver.go deleted file mode 100644 index 87c7122f4..000000000 --- a/bootstrap/mocks/binding_resolver.go +++ /dev/null @@ -1,111 +0,0 @@ -// Copyright (c) Abstract Machines - -// SPDX-License-Identifier: Apache-2.0 - -// Code generated by mockery; DO NOT EDIT. -// github.com/vektra/mockery -// template: testify - -package mocks - -import ( - "context" - - "github.com/absmach/magistrala/bootstrap" - mock "github.com/stretchr/testify/mock" -) - -// NewBindingResolver creates a new instance of BindingResolver. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. -// The first argument is typically a *testing.T value. -func NewBindingResolver(t interface { - mock.TestingT - Cleanup(func()) -}) *BindingResolver { - mock := &BindingResolver{} - mock.Mock.Test(t) - - t.Cleanup(func() { mock.AssertExpectations(t) }) - - return mock -} - -// BindingResolver is an autogenerated mock type for the BindingResolver type -type BindingResolver struct { - mock.Mock -} - -type BindingResolver_Expecter struct { - mock *mock.Mock -} - -func (_m *BindingResolver) EXPECT() *BindingResolver_Expecter { - return &BindingResolver_Expecter{mock: &_m.Mock} -} - -// Resolve provides a mock function for the type BindingResolver -func (_mock *BindingResolver) Resolve(ctx context.Context, req bootstrap.ResolveRequest) ([]bootstrap.BindingSnapshot, error) { - ret := _mock.Called(ctx, req) - - if len(ret) == 0 { - panic("no return value specified for Resolve") - } - - var r0 []bootstrap.BindingSnapshot - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, bootstrap.ResolveRequest) ([]bootstrap.BindingSnapshot, error)); ok { - return returnFunc(ctx, req) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, bootstrap.ResolveRequest) []bootstrap.BindingSnapshot); ok { - r0 = returnFunc(ctx, req) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]bootstrap.BindingSnapshot) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, bootstrap.ResolveRequest) error); ok { - r1 = returnFunc(ctx, req) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// BindingResolver_Resolve_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Resolve' -type BindingResolver_Resolve_Call struct { - *mock.Call -} - -// Resolve is a helper method to define mock.On call -// - ctx context.Context -// - req bootstrap.ResolveRequest -func (_e *BindingResolver_Expecter) Resolve(ctx interface{}, req interface{}) *BindingResolver_Resolve_Call { - return &BindingResolver_Resolve_Call{Call: _e.mock.On("Resolve", ctx, req)} -} - -func (_c *BindingResolver_Resolve_Call) Run(run func(ctx context.Context, req bootstrap.ResolveRequest)) *BindingResolver_Resolve_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 bootstrap.ResolveRequest - if args[1] != nil { - arg1 = args[1].(bootstrap.ResolveRequest) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *BindingResolver_Resolve_Call) Return(bindingSnapshots []bootstrap.BindingSnapshot, err error) *BindingResolver_Resolve_Call { - _c.Call.Return(bindingSnapshots, err) - return _c -} - -func (_c *BindingResolver_Resolve_Call) RunAndReturn(run func(ctx context.Context, req bootstrap.ResolveRequest) ([]bootstrap.BindingSnapshot, error)) *BindingResolver_Resolve_Call { - _c.Call.Return(run) - return _c -} diff --git a/bootstrap/mocks/binding_store.go b/bootstrap/mocks/binding_store.go deleted file mode 100644 index fb0b5857d..000000000 --- a/bootstrap/mocks/binding_store.go +++ /dev/null @@ -1,237 +0,0 @@ -// Copyright (c) Abstract Machines - -// SPDX-License-Identifier: Apache-2.0 - -// Code generated by mockery; DO NOT EDIT. -// github.com/vektra/mockery -// template: testify - -package mocks - -import ( - "context" - - "github.com/absmach/magistrala/bootstrap" - mock "github.com/stretchr/testify/mock" -) - -// NewBindingStore creates a new instance of BindingStore. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. -// The first argument is typically a *testing.T value. -func NewBindingStore(t interface { - mock.TestingT - Cleanup(func()) -}) *BindingStore { - mock := &BindingStore{} - mock.Mock.Test(t) - - t.Cleanup(func() { mock.AssertExpectations(t) }) - - return mock -} - -// BindingStore is an autogenerated mock type for the BindingStore type -type BindingStore struct { - mock.Mock -} - -type BindingStore_Expecter struct { - mock *mock.Mock -} - -func (_m *BindingStore) EXPECT() *BindingStore_Expecter { - return &BindingStore_Expecter{mock: &_m.Mock} -} - -// Delete provides a mock function for the type BindingStore -func (_mock *BindingStore) Delete(ctx context.Context, configID string, slot string) error { - ret := _mock.Called(ctx, configID, slot) - - if len(ret) == 0 { - panic("no return value specified for Delete") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) error); ok { - r0 = returnFunc(ctx, configID, slot) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// BindingStore_Delete_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Delete' -type BindingStore_Delete_Call struct { - *mock.Call -} - -// Delete is a helper method to define mock.On call -// - ctx context.Context -// - configID string -// - slot string -func (_e *BindingStore_Expecter) Delete(ctx interface{}, configID interface{}, slot interface{}) *BindingStore_Delete_Call { - return &BindingStore_Delete_Call{Call: _e.mock.On("Delete", ctx, configID, slot)} -} - -func (_c *BindingStore_Delete_Call) Run(run func(ctx context.Context, configID string, slot string)) *BindingStore_Delete_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *BindingStore_Delete_Call) Return(err error) *BindingStore_Delete_Call { - _c.Call.Return(err) - return _c -} - -func (_c *BindingStore_Delete_Call) RunAndReturn(run func(ctx context.Context, configID string, slot string) error) *BindingStore_Delete_Call { - _c.Call.Return(run) - return _c -} - -// Retrieve provides a mock function for the type BindingStore -func (_mock *BindingStore) Retrieve(ctx context.Context, configID string) ([]bootstrap.BindingSnapshot, error) { - ret := _mock.Called(ctx, configID) - - if len(ret) == 0 { - panic("no return value specified for Retrieve") - } - - var r0 []bootstrap.BindingSnapshot - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) ([]bootstrap.BindingSnapshot, error)); ok { - return returnFunc(ctx, configID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string) []bootstrap.BindingSnapshot); ok { - r0 = returnFunc(ctx, configID) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]bootstrap.BindingSnapshot) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = returnFunc(ctx, configID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// BindingStore_Retrieve_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Retrieve' -type BindingStore_Retrieve_Call struct { - *mock.Call -} - -// Retrieve is a helper method to define mock.On call -// - ctx context.Context -// - configID string -func (_e *BindingStore_Expecter) Retrieve(ctx interface{}, configID interface{}) *BindingStore_Retrieve_Call { - return &BindingStore_Retrieve_Call{Call: _e.mock.On("Retrieve", ctx, configID)} -} - -func (_c *BindingStore_Retrieve_Call) Run(run func(ctx context.Context, configID string)) *BindingStore_Retrieve_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *BindingStore_Retrieve_Call) Return(bindingSnapshots []bootstrap.BindingSnapshot, err error) *BindingStore_Retrieve_Call { - _c.Call.Return(bindingSnapshots, err) - return _c -} - -func (_c *BindingStore_Retrieve_Call) RunAndReturn(run func(ctx context.Context, configID string) ([]bootstrap.BindingSnapshot, error)) *BindingStore_Retrieve_Call { - _c.Call.Return(run) - return _c -} - -// Save provides a mock function for the type BindingStore -func (_mock *BindingStore) Save(ctx context.Context, configID string, bindings []bootstrap.BindingSnapshot) error { - ret := _mock.Called(ctx, configID, bindings) - - if len(ret) == 0 { - panic("no return value specified for Save") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, []bootstrap.BindingSnapshot) error); ok { - r0 = returnFunc(ctx, configID, bindings) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// BindingStore_Save_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Save' -type BindingStore_Save_Call struct { - *mock.Call -} - -// Save is a helper method to define mock.On call -// - ctx context.Context -// - configID string -// - bindings []bootstrap.BindingSnapshot -func (_e *BindingStore_Expecter) Save(ctx interface{}, configID interface{}, bindings interface{}) *BindingStore_Save_Call { - return &BindingStore_Save_Call{Call: _e.mock.On("Save", ctx, configID, bindings)} -} - -func (_c *BindingStore_Save_Call) Run(run func(ctx context.Context, configID string, bindings []bootstrap.BindingSnapshot)) *BindingStore_Save_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 []bootstrap.BindingSnapshot - if args[2] != nil { - arg2 = args[2].([]bootstrap.BindingSnapshot) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *BindingStore_Save_Call) Return(err error) *BindingStore_Save_Call { - _c.Call.Return(err) - return _c -} - -func (_c *BindingStore_Save_Call) RunAndReturn(run func(ctx context.Context, configID string, bindings []bootstrap.BindingSnapshot) error) *BindingStore_Save_Call { - _c.Call.Return(run) - return _c -} diff --git a/bootstrap/mocks/config_reader.go b/bootstrap/mocks/config_reader.go deleted file mode 100644 index 6cf14aaa3..000000000 --- a/bootstrap/mocks/config_reader.go +++ /dev/null @@ -1,109 +0,0 @@ -// Copyright (c) Abstract Machines - -// SPDX-License-Identifier: Apache-2.0 - -// Code generated by mockery; DO NOT EDIT. -// github.com/vektra/mockery -// template: testify - -package mocks - -import ( - "github.com/absmach/magistrala/bootstrap" - mock "github.com/stretchr/testify/mock" -) - -// NewConfigReader creates a new instance of ConfigReader. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. -// The first argument is typically a *testing.T value. -func NewConfigReader(t interface { - mock.TestingT - Cleanup(func()) -}) *ConfigReader { - mock := &ConfigReader{} - mock.Mock.Test(t) - - t.Cleanup(func() { mock.AssertExpectations(t) }) - - return mock -} - -// ConfigReader is an autogenerated mock type for the ConfigReader type -type ConfigReader struct { - mock.Mock -} - -type ConfigReader_Expecter struct { - mock *mock.Mock -} - -func (_m *ConfigReader) EXPECT() *ConfigReader_Expecter { - return &ConfigReader_Expecter{mock: &_m.Mock} -} - -// ReadConfig provides a mock function for the type ConfigReader -func (_mock *ConfigReader) ReadConfig(config bootstrap.Config, b bool) (any, error) { - ret := _mock.Called(config, b) - - if len(ret) == 0 { - panic("no return value specified for ReadConfig") - } - - var r0 any - var r1 error - if returnFunc, ok := ret.Get(0).(func(bootstrap.Config, bool) (any, error)); ok { - return returnFunc(config, b) - } - if returnFunc, ok := ret.Get(0).(func(bootstrap.Config, bool) any); ok { - r0 = returnFunc(config, b) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).(any) - } - } - if returnFunc, ok := ret.Get(1).(func(bootstrap.Config, bool) error); ok { - r1 = returnFunc(config, b) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// ConfigReader_ReadConfig_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ReadConfig' -type ConfigReader_ReadConfig_Call struct { - *mock.Call -} - -// ReadConfig is a helper method to define mock.On call -// - config bootstrap.Config -// - b bool -func (_e *ConfigReader_Expecter) ReadConfig(config interface{}, b interface{}) *ConfigReader_ReadConfig_Call { - return &ConfigReader_ReadConfig_Call{Call: _e.mock.On("ReadConfig", config, b)} -} - -func (_c *ConfigReader_ReadConfig_Call) Run(run func(config bootstrap.Config, b bool)) *ConfigReader_ReadConfig_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 bootstrap.Config - if args[0] != nil { - arg0 = args[0].(bootstrap.Config) - } - var arg1 bool - if args[1] != nil { - arg1 = args[1].(bool) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *ConfigReader_ReadConfig_Call) Return(v any, err error) *ConfigReader_ReadConfig_Call { - _c.Call.Return(v, err) - return _c -} - -func (_c *ConfigReader_ReadConfig_Call) RunAndReturn(run func(config bootstrap.Config, b bool) (any, error)) *ConfigReader_ReadConfig_Call { - _c.Call.Return(run) - return _c -} diff --git a/bootstrap/mocks/config_repository.go b/bootstrap/mocks/config_repository.go deleted file mode 100644 index 53dcdea97..000000000 --- a/bootstrap/mocks/config_repository.go +++ /dev/null @@ -1,670 +0,0 @@ -// Copyright (c) Abstract Machines - -// SPDX-License-Identifier: Apache-2.0 - -// Code generated by mockery; DO NOT EDIT. -// github.com/vektra/mockery -// template: testify - -package mocks - -import ( - "context" - - "github.com/absmach/magistrala/bootstrap" - mock "github.com/stretchr/testify/mock" -) - -// NewConfigRepository creates a new instance of ConfigRepository. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. -// The first argument is typically a *testing.T value. -func NewConfigRepository(t interface { - mock.TestingT - Cleanup(func()) -}) *ConfigRepository { - mock := &ConfigRepository{} - mock.Mock.Test(t) - - t.Cleanup(func() { mock.AssertExpectations(t) }) - - return mock -} - -// ConfigRepository is an autogenerated mock type for the ConfigRepository type -type ConfigRepository struct { - mock.Mock -} - -type ConfigRepository_Expecter struct { - mock *mock.Mock -} - -func (_m *ConfigRepository) EXPECT() *ConfigRepository_Expecter { - return &ConfigRepository_Expecter{mock: &_m.Mock} -} - -// AssignProfile provides a mock function for the type ConfigRepository -func (_mock *ConfigRepository) AssignProfile(ctx context.Context, domainID string, id string, profileID string) error { - ret := _mock.Called(ctx, domainID, id, profileID) - - if len(ret) == 0 { - panic("no return value specified for AssignProfile") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string, string) error); ok { - r0 = returnFunc(ctx, domainID, id, profileID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// ConfigRepository_AssignProfile_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'AssignProfile' -type ConfigRepository_AssignProfile_Call struct { - *mock.Call -} - -// AssignProfile is a helper method to define mock.On call -// - ctx context.Context -// - domainID string -// - id string -// - profileID string -func (_e *ConfigRepository_Expecter) AssignProfile(ctx interface{}, domainID interface{}, id interface{}, profileID interface{}) *ConfigRepository_AssignProfile_Call { - return &ConfigRepository_AssignProfile_Call{Call: _e.mock.On("AssignProfile", ctx, domainID, id, profileID)} -} - -func (_c *ConfigRepository_AssignProfile_Call) Run(run func(ctx context.Context, domainID string, id string, profileID string)) *ConfigRepository_AssignProfile_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *ConfigRepository_AssignProfile_Call) Return(err error) *ConfigRepository_AssignProfile_Call { - _c.Call.Return(err) - return _c -} - -func (_c *ConfigRepository_AssignProfile_Call) RunAndReturn(run func(ctx context.Context, domainID string, id string, profileID string) error) *ConfigRepository_AssignProfile_Call { - _c.Call.Return(run) - return _c -} - -// ChangeStatus provides a mock function for the type ConfigRepository -func (_mock *ConfigRepository) ChangeStatus(ctx context.Context, domainID string, id string, status bootstrap.Status) error { - ret := _mock.Called(ctx, domainID, id, status) - - if len(ret) == 0 { - panic("no return value specified for ChangeStatus") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string, bootstrap.Status) error); ok { - r0 = returnFunc(ctx, domainID, id, status) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// ConfigRepository_ChangeStatus_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ChangeStatus' -type ConfigRepository_ChangeStatus_Call struct { - *mock.Call -} - -// ChangeStatus is a helper method to define mock.On call -// - ctx context.Context -// - domainID string -// - id string -// - status bootstrap.Status -func (_e *ConfigRepository_Expecter) ChangeStatus(ctx interface{}, domainID interface{}, id interface{}, status interface{}) *ConfigRepository_ChangeStatus_Call { - return &ConfigRepository_ChangeStatus_Call{Call: _e.mock.On("ChangeStatus", ctx, domainID, id, status)} -} - -func (_c *ConfigRepository_ChangeStatus_Call) Run(run func(ctx context.Context, domainID string, id string, status bootstrap.Status)) *ConfigRepository_ChangeStatus_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 bootstrap.Status - if args[3] != nil { - arg3 = args[3].(bootstrap.Status) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *ConfigRepository_ChangeStatus_Call) Return(err error) *ConfigRepository_ChangeStatus_Call { - _c.Call.Return(err) - return _c -} - -func (_c *ConfigRepository_ChangeStatus_Call) RunAndReturn(run func(ctx context.Context, domainID string, id string, status bootstrap.Status) error) *ConfigRepository_ChangeStatus_Call { - _c.Call.Return(run) - return _c -} - -// Remove provides a mock function for the type ConfigRepository -func (_mock *ConfigRepository) Remove(ctx context.Context, domainID string, id string) error { - ret := _mock.Called(ctx, domainID, id) - - if len(ret) == 0 { - panic("no return value specified for Remove") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) error); ok { - r0 = returnFunc(ctx, domainID, id) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// ConfigRepository_Remove_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Remove' -type ConfigRepository_Remove_Call struct { - *mock.Call -} - -// Remove is a helper method to define mock.On call -// - ctx context.Context -// - domainID string -// - id string -func (_e *ConfigRepository_Expecter) Remove(ctx interface{}, domainID interface{}, id interface{}) *ConfigRepository_Remove_Call { - return &ConfigRepository_Remove_Call{Call: _e.mock.On("Remove", ctx, domainID, id)} -} - -func (_c *ConfigRepository_Remove_Call) Run(run func(ctx context.Context, domainID string, id string)) *ConfigRepository_Remove_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *ConfigRepository_Remove_Call) Return(err error) *ConfigRepository_Remove_Call { - _c.Call.Return(err) - return _c -} - -func (_c *ConfigRepository_Remove_Call) RunAndReturn(run func(ctx context.Context, domainID string, id string) error) *ConfigRepository_Remove_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveAll provides a mock function for the type ConfigRepository -func (_mock *ConfigRepository) RetrieveAll(ctx context.Context, domainID string, filter bootstrap.Filter, offset uint64, limit uint64) bootstrap.ConfigsPage { - ret := _mock.Called(ctx, domainID, filter, offset, limit) - - if len(ret) == 0 { - panic("no return value specified for RetrieveAll") - } - - var r0 bootstrap.ConfigsPage - if returnFunc, ok := ret.Get(0).(func(context.Context, string, bootstrap.Filter, uint64, uint64) bootstrap.ConfigsPage); ok { - r0 = returnFunc(ctx, domainID, filter, offset, limit) - } else { - r0 = ret.Get(0).(bootstrap.ConfigsPage) - } - return r0 -} - -// ConfigRepository_RetrieveAll_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveAll' -type ConfigRepository_RetrieveAll_Call struct { - *mock.Call -} - -// RetrieveAll is a helper method to define mock.On call -// - ctx context.Context -// - domainID string -// - filter bootstrap.Filter -// - offset uint64 -// - limit uint64 -func (_e *ConfigRepository_Expecter) RetrieveAll(ctx interface{}, domainID interface{}, filter interface{}, offset interface{}, limit interface{}) *ConfigRepository_RetrieveAll_Call { - return &ConfigRepository_RetrieveAll_Call{Call: _e.mock.On("RetrieveAll", ctx, domainID, filter, offset, limit)} -} - -func (_c *ConfigRepository_RetrieveAll_Call) Run(run func(ctx context.Context, domainID string, filter bootstrap.Filter, offset uint64, limit uint64)) *ConfigRepository_RetrieveAll_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 bootstrap.Filter - if args[2] != nil { - arg2 = args[2].(bootstrap.Filter) - } - var arg3 uint64 - if args[3] != nil { - arg3 = args[3].(uint64) - } - var arg4 uint64 - if args[4] != nil { - arg4 = args[4].(uint64) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *ConfigRepository_RetrieveAll_Call) Return(configsPage bootstrap.ConfigsPage) *ConfigRepository_RetrieveAll_Call { - _c.Call.Return(configsPage) - return _c -} - -func (_c *ConfigRepository_RetrieveAll_Call) RunAndReturn(run func(ctx context.Context, domainID string, filter bootstrap.Filter, offset uint64, limit uint64) bootstrap.ConfigsPage) *ConfigRepository_RetrieveAll_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveByExternalID provides a mock function for the type ConfigRepository -func (_mock *ConfigRepository) RetrieveByExternalID(ctx context.Context, externalID string) (bootstrap.Config, error) { - ret := _mock.Called(ctx, externalID) - - if len(ret) == 0 { - panic("no return value specified for RetrieveByExternalID") - } - - var r0 bootstrap.Config - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) (bootstrap.Config, error)); ok { - return returnFunc(ctx, externalID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string) bootstrap.Config); ok { - r0 = returnFunc(ctx, externalID) - } else { - r0 = ret.Get(0).(bootstrap.Config) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = returnFunc(ctx, externalID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// ConfigRepository_RetrieveByExternalID_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveByExternalID' -type ConfigRepository_RetrieveByExternalID_Call struct { - *mock.Call -} - -// RetrieveByExternalID is a helper method to define mock.On call -// - ctx context.Context -// - externalID string -func (_e *ConfigRepository_Expecter) RetrieveByExternalID(ctx interface{}, externalID interface{}) *ConfigRepository_RetrieveByExternalID_Call { - return &ConfigRepository_RetrieveByExternalID_Call{Call: _e.mock.On("RetrieveByExternalID", ctx, externalID)} -} - -func (_c *ConfigRepository_RetrieveByExternalID_Call) Run(run func(ctx context.Context, externalID string)) *ConfigRepository_RetrieveByExternalID_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *ConfigRepository_RetrieveByExternalID_Call) Return(config bootstrap.Config, err error) *ConfigRepository_RetrieveByExternalID_Call { - _c.Call.Return(config, err) - return _c -} - -func (_c *ConfigRepository_RetrieveByExternalID_Call) RunAndReturn(run func(ctx context.Context, externalID string) (bootstrap.Config, error)) *ConfigRepository_RetrieveByExternalID_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveByID provides a mock function for the type ConfigRepository -func (_mock *ConfigRepository) RetrieveByID(ctx context.Context, domainID string, id string) (bootstrap.Config, error) { - ret := _mock.Called(ctx, domainID, id) - - if len(ret) == 0 { - panic("no return value specified for RetrieveByID") - } - - var r0 bootstrap.Config - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) (bootstrap.Config, error)); ok { - return returnFunc(ctx, domainID, id) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) bootstrap.Config); ok { - r0 = returnFunc(ctx, domainID, id) - } else { - r0 = ret.Get(0).(bootstrap.Config) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, string) error); ok { - r1 = returnFunc(ctx, domainID, id) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// ConfigRepository_RetrieveByID_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveByID' -type ConfigRepository_RetrieveByID_Call struct { - *mock.Call -} - -// RetrieveByID is a helper method to define mock.On call -// - ctx context.Context -// - domainID string -// - id string -func (_e *ConfigRepository_Expecter) RetrieveByID(ctx interface{}, domainID interface{}, id interface{}) *ConfigRepository_RetrieveByID_Call { - return &ConfigRepository_RetrieveByID_Call{Call: _e.mock.On("RetrieveByID", ctx, domainID, id)} -} - -func (_c *ConfigRepository_RetrieveByID_Call) Run(run func(ctx context.Context, domainID string, id string)) *ConfigRepository_RetrieveByID_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *ConfigRepository_RetrieveByID_Call) Return(config bootstrap.Config, err error) *ConfigRepository_RetrieveByID_Call { - _c.Call.Return(config, err) - return _c -} - -func (_c *ConfigRepository_RetrieveByID_Call) RunAndReturn(run func(ctx context.Context, domainID string, id string) (bootstrap.Config, error)) *ConfigRepository_RetrieveByID_Call { - _c.Call.Return(run) - return _c -} - -// Save provides a mock function for the type ConfigRepository -func (_mock *ConfigRepository) Save(ctx context.Context, cfg bootstrap.Config) (string, error) { - ret := _mock.Called(ctx, cfg) - - if len(ret) == 0 { - panic("no return value specified for Save") - } - - var r0 string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, bootstrap.Config) (string, error)); ok { - return returnFunc(ctx, cfg) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, bootstrap.Config) string); ok { - r0 = returnFunc(ctx, cfg) - } else { - r0 = ret.Get(0).(string) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, bootstrap.Config) error); ok { - r1 = returnFunc(ctx, cfg) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// ConfigRepository_Save_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Save' -type ConfigRepository_Save_Call struct { - *mock.Call -} - -// Save is a helper method to define mock.On call -// - ctx context.Context -// - cfg bootstrap.Config -func (_e *ConfigRepository_Expecter) Save(ctx interface{}, cfg interface{}) *ConfigRepository_Save_Call { - return &ConfigRepository_Save_Call{Call: _e.mock.On("Save", ctx, cfg)} -} - -func (_c *ConfigRepository_Save_Call) Run(run func(ctx context.Context, cfg bootstrap.Config)) *ConfigRepository_Save_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 bootstrap.Config - if args[1] != nil { - arg1 = args[1].(bootstrap.Config) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *ConfigRepository_Save_Call) Return(s string, err error) *ConfigRepository_Save_Call { - _c.Call.Return(s, err) - return _c -} - -func (_c *ConfigRepository_Save_Call) RunAndReturn(run func(ctx context.Context, cfg bootstrap.Config) (string, error)) *ConfigRepository_Save_Call { - _c.Call.Return(run) - return _c -} - -// Update provides a mock function for the type ConfigRepository -func (_mock *ConfigRepository) Update(ctx context.Context, cfg bootstrap.Config) error { - ret := _mock.Called(ctx, cfg) - - if len(ret) == 0 { - panic("no return value specified for Update") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, bootstrap.Config) error); ok { - r0 = returnFunc(ctx, cfg) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// ConfigRepository_Update_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Update' -type ConfigRepository_Update_Call struct { - *mock.Call -} - -// Update is a helper method to define mock.On call -// - ctx context.Context -// - cfg bootstrap.Config -func (_e *ConfigRepository_Expecter) Update(ctx interface{}, cfg interface{}) *ConfigRepository_Update_Call { - return &ConfigRepository_Update_Call{Call: _e.mock.On("Update", ctx, cfg)} -} - -func (_c *ConfigRepository_Update_Call) Run(run func(ctx context.Context, cfg bootstrap.Config)) *ConfigRepository_Update_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 bootstrap.Config - if args[1] != nil { - arg1 = args[1].(bootstrap.Config) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *ConfigRepository_Update_Call) Return(err error) *ConfigRepository_Update_Call { - _c.Call.Return(err) - return _c -} - -func (_c *ConfigRepository_Update_Call) RunAndReturn(run func(ctx context.Context, cfg bootstrap.Config) error) *ConfigRepository_Update_Call { - _c.Call.Return(run) - return _c -} - -// UpdateCert provides a mock function for the type ConfigRepository -func (_mock *ConfigRepository) UpdateCert(ctx context.Context, domainID string, id string, clientCert string, clientKey string, caCert string) (bootstrap.Config, error) { - ret := _mock.Called(ctx, domainID, id, clientCert, clientKey, caCert) - - if len(ret) == 0 { - panic("no return value specified for UpdateCert") - } - - var r0 bootstrap.Config - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string, string, string, string) (bootstrap.Config, error)); ok { - return returnFunc(ctx, domainID, id, clientCert, clientKey, caCert) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string, string, string, string) bootstrap.Config); ok { - r0 = returnFunc(ctx, domainID, id, clientCert, clientKey, caCert) - } else { - r0 = ret.Get(0).(bootstrap.Config) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, string, string, string, string) error); ok { - r1 = returnFunc(ctx, domainID, id, clientCert, clientKey, caCert) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// ConfigRepository_UpdateCert_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateCert' -type ConfigRepository_UpdateCert_Call struct { - *mock.Call -} - -// UpdateCert is a helper method to define mock.On call -// - ctx context.Context -// - domainID string -// - id string -// - clientCert string -// - clientKey string -// - caCert string -func (_e *ConfigRepository_Expecter) UpdateCert(ctx interface{}, domainID interface{}, id interface{}, clientCert interface{}, clientKey interface{}, caCert interface{}) *ConfigRepository_UpdateCert_Call { - return &ConfigRepository_UpdateCert_Call{Call: _e.mock.On("UpdateCert", ctx, domainID, id, clientCert, clientKey, caCert)} -} - -func (_c *ConfigRepository_UpdateCert_Call) Run(run func(ctx context.Context, domainID string, id string, clientCert string, clientKey string, caCert string)) *ConfigRepository_UpdateCert_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 string - if args[4] != nil { - arg4 = args[4].(string) - } - var arg5 string - if args[5] != nil { - arg5 = args[5].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - arg5, - ) - }) - return _c -} - -func (_c *ConfigRepository_UpdateCert_Call) Return(config bootstrap.Config, err error) *ConfigRepository_UpdateCert_Call { - _c.Call.Return(config, err) - return _c -} - -func (_c *ConfigRepository_UpdateCert_Call) RunAndReturn(run func(ctx context.Context, domainID string, id string, clientCert string, clientKey string, caCert string) (bootstrap.Config, error)) *ConfigRepository_UpdateCert_Call { - _c.Call.Return(run) - return _c -} diff --git a/bootstrap/mocks/profile_repository.go b/bootstrap/mocks/profile_repository.go deleted file mode 100644 index 0fcec9c0b..000000000 --- a/bootstrap/mocks/profile_repository.go +++ /dev/null @@ -1,394 +0,0 @@ -// Copyright (c) Abstract Machines - -// SPDX-License-Identifier: Apache-2.0 - -// Code generated by mockery; DO NOT EDIT. -// github.com/vektra/mockery -// template: testify - -package mocks - -import ( - "context" - - "github.com/absmach/magistrala/bootstrap" - mock "github.com/stretchr/testify/mock" -) - -// NewProfileRepository creates a new instance of ProfileRepository. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. -// The first argument is typically a *testing.T value. -func NewProfileRepository(t interface { - mock.TestingT - Cleanup(func()) -}) *ProfileRepository { - mock := &ProfileRepository{} - mock.Mock.Test(t) - - t.Cleanup(func() { mock.AssertExpectations(t) }) - - return mock -} - -// ProfileRepository is an autogenerated mock type for the ProfileRepository type -type ProfileRepository struct { - mock.Mock -} - -type ProfileRepository_Expecter struct { - mock *mock.Mock -} - -func (_m *ProfileRepository) EXPECT() *ProfileRepository_Expecter { - return &ProfileRepository_Expecter{mock: &_m.Mock} -} - -// Delete provides a mock function for the type ProfileRepository -func (_mock *ProfileRepository) Delete(ctx context.Context, domainID string, id string) error { - ret := _mock.Called(ctx, domainID, id) - - if len(ret) == 0 { - panic("no return value specified for Delete") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) error); ok { - r0 = returnFunc(ctx, domainID, id) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// ProfileRepository_Delete_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Delete' -type ProfileRepository_Delete_Call struct { - *mock.Call -} - -// Delete is a helper method to define mock.On call -// - ctx context.Context -// - domainID string -// - id string -func (_e *ProfileRepository_Expecter) Delete(ctx interface{}, domainID interface{}, id interface{}) *ProfileRepository_Delete_Call { - return &ProfileRepository_Delete_Call{Call: _e.mock.On("Delete", ctx, domainID, id)} -} - -func (_c *ProfileRepository_Delete_Call) Run(run func(ctx context.Context, domainID string, id string)) *ProfileRepository_Delete_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *ProfileRepository_Delete_Call) Return(err error) *ProfileRepository_Delete_Call { - _c.Call.Return(err) - return _c -} - -func (_c *ProfileRepository_Delete_Call) RunAndReturn(run func(ctx context.Context, domainID string, id string) error) *ProfileRepository_Delete_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveAll provides a mock function for the type ProfileRepository -func (_mock *ProfileRepository) RetrieveAll(ctx context.Context, domainID string, offset uint64, limit uint64, name string) (bootstrap.ProfilesPage, error) { - ret := _mock.Called(ctx, domainID, offset, limit, name) - - if len(ret) == 0 { - panic("no return value specified for RetrieveAll") - } - - var r0 bootstrap.ProfilesPage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, uint64, uint64, string) (bootstrap.ProfilesPage, error)); ok { - return returnFunc(ctx, domainID, offset, limit, name) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, uint64, uint64, string) bootstrap.ProfilesPage); ok { - r0 = returnFunc(ctx, domainID, offset, limit, name) - } else { - r0 = ret.Get(0).(bootstrap.ProfilesPage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, uint64, uint64, string) error); ok { - r1 = returnFunc(ctx, domainID, offset, limit, name) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// ProfileRepository_RetrieveAll_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveAll' -type ProfileRepository_RetrieveAll_Call struct { - *mock.Call -} - -// RetrieveAll is a helper method to define mock.On call -// - ctx context.Context -// - domainID string -// - offset uint64 -// - limit uint64 -// - name string -func (_e *ProfileRepository_Expecter) RetrieveAll(ctx interface{}, domainID interface{}, offset interface{}, limit interface{}, name interface{}) *ProfileRepository_RetrieveAll_Call { - return &ProfileRepository_RetrieveAll_Call{Call: _e.mock.On("RetrieveAll", ctx, domainID, offset, limit, name)} -} - -func (_c *ProfileRepository_RetrieveAll_Call) Run(run func(ctx context.Context, domainID string, offset uint64, limit uint64, name string)) *ProfileRepository_RetrieveAll_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 uint64 - if args[2] != nil { - arg2 = args[2].(uint64) - } - var arg3 uint64 - if args[3] != nil { - arg3 = args[3].(uint64) - } - var arg4 string - if args[4] != nil { - arg4 = args[4].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *ProfileRepository_RetrieveAll_Call) Return(profilesPage bootstrap.ProfilesPage, err error) *ProfileRepository_RetrieveAll_Call { - _c.Call.Return(profilesPage, err) - return _c -} - -func (_c *ProfileRepository_RetrieveAll_Call) RunAndReturn(run func(ctx context.Context, domainID string, offset uint64, limit uint64, name string) (bootstrap.ProfilesPage, error)) *ProfileRepository_RetrieveAll_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveByID provides a mock function for the type ProfileRepository -func (_mock *ProfileRepository) RetrieveByID(ctx context.Context, domainID string, id string) (bootstrap.Profile, error) { - ret := _mock.Called(ctx, domainID, id) - - if len(ret) == 0 { - panic("no return value specified for RetrieveByID") - } - - var r0 bootstrap.Profile - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) (bootstrap.Profile, error)); ok { - return returnFunc(ctx, domainID, id) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) bootstrap.Profile); ok { - r0 = returnFunc(ctx, domainID, id) - } else { - r0 = ret.Get(0).(bootstrap.Profile) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, string) error); ok { - r1 = returnFunc(ctx, domainID, id) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// ProfileRepository_RetrieveByID_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveByID' -type ProfileRepository_RetrieveByID_Call struct { - *mock.Call -} - -// RetrieveByID is a helper method to define mock.On call -// - ctx context.Context -// - domainID string -// - id string -func (_e *ProfileRepository_Expecter) RetrieveByID(ctx interface{}, domainID interface{}, id interface{}) *ProfileRepository_RetrieveByID_Call { - return &ProfileRepository_RetrieveByID_Call{Call: _e.mock.On("RetrieveByID", ctx, domainID, id)} -} - -func (_c *ProfileRepository_RetrieveByID_Call) Run(run func(ctx context.Context, domainID string, id string)) *ProfileRepository_RetrieveByID_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *ProfileRepository_RetrieveByID_Call) Return(profile bootstrap.Profile, err error) *ProfileRepository_RetrieveByID_Call { - _c.Call.Return(profile, err) - return _c -} - -func (_c *ProfileRepository_RetrieveByID_Call) RunAndReturn(run func(ctx context.Context, domainID string, id string) (bootstrap.Profile, error)) *ProfileRepository_RetrieveByID_Call { - _c.Call.Return(run) - return _c -} - -// Save provides a mock function for the type ProfileRepository -func (_mock *ProfileRepository) Save(ctx context.Context, p bootstrap.Profile) (bootstrap.Profile, error) { - ret := _mock.Called(ctx, p) - - if len(ret) == 0 { - panic("no return value specified for Save") - } - - var r0 bootstrap.Profile - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, bootstrap.Profile) (bootstrap.Profile, error)); ok { - return returnFunc(ctx, p) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, bootstrap.Profile) bootstrap.Profile); ok { - r0 = returnFunc(ctx, p) - } else { - r0 = ret.Get(0).(bootstrap.Profile) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, bootstrap.Profile) error); ok { - r1 = returnFunc(ctx, p) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// ProfileRepository_Save_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Save' -type ProfileRepository_Save_Call struct { - *mock.Call -} - -// Save is a helper method to define mock.On call -// - ctx context.Context -// - p bootstrap.Profile -func (_e *ProfileRepository_Expecter) Save(ctx interface{}, p interface{}) *ProfileRepository_Save_Call { - return &ProfileRepository_Save_Call{Call: _e.mock.On("Save", ctx, p)} -} - -func (_c *ProfileRepository_Save_Call) Run(run func(ctx context.Context, p bootstrap.Profile)) *ProfileRepository_Save_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 bootstrap.Profile - if args[1] != nil { - arg1 = args[1].(bootstrap.Profile) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *ProfileRepository_Save_Call) Return(profile bootstrap.Profile, err error) *ProfileRepository_Save_Call { - _c.Call.Return(profile, err) - return _c -} - -func (_c *ProfileRepository_Save_Call) RunAndReturn(run func(ctx context.Context, p bootstrap.Profile) (bootstrap.Profile, error)) *ProfileRepository_Save_Call { - _c.Call.Return(run) - return _c -} - -// Update provides a mock function for the type ProfileRepository -func (_mock *ProfileRepository) Update(ctx context.Context, p bootstrap.Profile) (bootstrap.Profile, error) { - ret := _mock.Called(ctx, p) - - if len(ret) == 0 { - panic("no return value specified for Update") - } - - var r0 bootstrap.Profile - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, bootstrap.Profile) (bootstrap.Profile, error)); ok { - return returnFunc(ctx, p) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, bootstrap.Profile) bootstrap.Profile); ok { - r0 = returnFunc(ctx, p) - } else { - r0 = ret.Get(0).(bootstrap.Profile) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, bootstrap.Profile) error); ok { - r1 = returnFunc(ctx, p) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// ProfileRepository_Update_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Update' -type ProfileRepository_Update_Call struct { - *mock.Call -} - -// Update is a helper method to define mock.On call -// - ctx context.Context -// - p bootstrap.Profile -func (_e *ProfileRepository_Expecter) Update(ctx interface{}, p interface{}) *ProfileRepository_Update_Call { - return &ProfileRepository_Update_Call{Call: _e.mock.On("Update", ctx, p)} -} - -func (_c *ProfileRepository_Update_Call) Run(run func(ctx context.Context, p bootstrap.Profile)) *ProfileRepository_Update_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 bootstrap.Profile - if args[1] != nil { - arg1 = args[1].(bootstrap.Profile) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *ProfileRepository_Update_Call) Return(profile bootstrap.Profile, err error) *ProfileRepository_Update_Call { - _c.Call.Return(profile, err) - return _c -} - -func (_c *ProfileRepository_Update_Call) RunAndReturn(run func(ctx context.Context, p bootstrap.Profile) (bootstrap.Profile, error)) *ProfileRepository_Update_Call { - _c.Call.Return(run) - return _c -} diff --git a/bootstrap/mocks/renderer.go b/bootstrap/mocks/renderer.go deleted file mode 100644 index aae3f7e96..000000000 --- a/bootstrap/mocks/renderer.go +++ /dev/null @@ -1,115 +0,0 @@ -// Copyright (c) Abstract Machines - -// SPDX-License-Identifier: Apache-2.0 - -// Code generated by mockery; DO NOT EDIT. -// github.com/vektra/mockery -// template: testify - -package mocks - -import ( - "github.com/absmach/magistrala/bootstrap" - mock "github.com/stretchr/testify/mock" -) - -// NewRenderer creates a new instance of Renderer. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. -// The first argument is typically a *testing.T value. -func NewRenderer(t interface { - mock.TestingT - Cleanup(func()) -}) *Renderer { - mock := &Renderer{} - mock.Mock.Test(t) - - t.Cleanup(func() { mock.AssertExpectations(t) }) - - return mock -} - -// Renderer is an autogenerated mock type for the Renderer type -type Renderer struct { - mock.Mock -} - -type Renderer_Expecter struct { - mock *mock.Mock -} - -func (_m *Renderer) EXPECT() *Renderer_Expecter { - return &Renderer_Expecter{mock: &_m.Mock} -} - -// Render provides a mock function for the type Renderer -func (_mock *Renderer) Render(profile bootstrap.Profile, enrollment bootstrap.Config, bindings []bootstrap.BindingSnapshot) ([]byte, error) { - ret := _mock.Called(profile, enrollment, bindings) - - if len(ret) == 0 { - panic("no return value specified for Render") - } - - var r0 []byte - var r1 error - if returnFunc, ok := ret.Get(0).(func(bootstrap.Profile, bootstrap.Config, []bootstrap.BindingSnapshot) ([]byte, error)); ok { - return returnFunc(profile, enrollment, bindings) - } - if returnFunc, ok := ret.Get(0).(func(bootstrap.Profile, bootstrap.Config, []bootstrap.BindingSnapshot) []byte); ok { - r0 = returnFunc(profile, enrollment, bindings) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]byte) - } - } - if returnFunc, ok := ret.Get(1).(func(bootstrap.Profile, bootstrap.Config, []bootstrap.BindingSnapshot) error); ok { - r1 = returnFunc(profile, enrollment, bindings) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Renderer_Render_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Render' -type Renderer_Render_Call struct { - *mock.Call -} - -// Render is a helper method to define mock.On call -// - profile bootstrap.Profile -// - enrollment bootstrap.Config -// - bindings []bootstrap.BindingSnapshot -func (_e *Renderer_Expecter) Render(profile interface{}, enrollment interface{}, bindings interface{}) *Renderer_Render_Call { - return &Renderer_Render_Call{Call: _e.mock.On("Render", profile, enrollment, bindings)} -} - -func (_c *Renderer_Render_Call) Run(run func(profile bootstrap.Profile, enrollment bootstrap.Config, bindings []bootstrap.BindingSnapshot)) *Renderer_Render_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 bootstrap.Profile - if args[0] != nil { - arg0 = args[0].(bootstrap.Profile) - } - var arg1 bootstrap.Config - if args[1] != nil { - arg1 = args[1].(bootstrap.Config) - } - var arg2 []bootstrap.BindingSnapshot - if args[2] != nil { - arg2 = args[2].([]bootstrap.BindingSnapshot) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Renderer_Render_Call) Return(bytes []byte, err error) *Renderer_Render_Call { - _c.Call.Return(bytes, err) - return _c -} - -func (_c *Renderer_Render_Call) RunAndReturn(run func(profile bootstrap.Profile, enrollment bootstrap.Config, bindings []bootstrap.BindingSnapshot) ([]byte, error)) *Renderer_Render_Call { - _c.Call.Return(run) - return _c -} diff --git a/bootstrap/mocks/service.go b/bootstrap/mocks/service.go deleted file mode 100644 index 43ebaa760..000000000 --- a/bootstrap/mocks/service.go +++ /dev/null @@ -1,1366 +0,0 @@ -// Copyright (c) Abstract Machines - -// SPDX-License-Identifier: Apache-2.0 - -// Code generated by mockery; DO NOT EDIT. -// github.com/vektra/mockery -// template: testify - -package mocks - -import ( - "context" - - "github.com/absmach/magistrala/bootstrap" - "github.com/absmach/magistrala/pkg/authn" - mock "github.com/stretchr/testify/mock" -) - -// NewService creates a new instance of Service. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. -// The first argument is typically a *testing.T value. -func NewService(t interface { - mock.TestingT - Cleanup(func()) -}) *Service { - mock := &Service{} - mock.Mock.Test(t) - - t.Cleanup(func() { mock.AssertExpectations(t) }) - - return mock -} - -// Service is an autogenerated mock type for the Service type -type Service struct { - mock.Mock -} - -type Service_Expecter struct { - mock *mock.Mock -} - -func (_m *Service) EXPECT() *Service_Expecter { - return &Service_Expecter{mock: &_m.Mock} -} - -// Add provides a mock function for the type Service -func (_mock *Service) Add(ctx context.Context, session authn.Session, token string, cfg bootstrap.Config) (bootstrap.Config, error) { - ret := _mock.Called(ctx, session, token, cfg) - - if len(ret) == 0 { - panic("no return value specified for Add") - } - - var r0 bootstrap.Config - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, bootstrap.Config) (bootstrap.Config, error)); ok { - return returnFunc(ctx, session, token, cfg) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, bootstrap.Config) bootstrap.Config); ok { - r0 = returnFunc(ctx, session, token, cfg) - } else { - r0 = ret.Get(0).(bootstrap.Config) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, bootstrap.Config) error); ok { - r1 = returnFunc(ctx, session, token, cfg) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_Add_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Add' -type Service_Add_Call struct { - *mock.Call -} - -// Add is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - token string -// - cfg bootstrap.Config -func (_e *Service_Expecter) Add(ctx interface{}, session interface{}, token interface{}, cfg interface{}) *Service_Add_Call { - return &Service_Add_Call{Call: _e.mock.On("Add", ctx, session, token, cfg)} -} - -func (_c *Service_Add_Call) Run(run func(ctx context.Context, session authn.Session, token string, cfg bootstrap.Config)) *Service_Add_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 bootstrap.Config - if args[3] != nil { - arg3 = args[3].(bootstrap.Config) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_Add_Call) Return(config bootstrap.Config, err error) *Service_Add_Call { - _c.Call.Return(config, err) - return _c -} - -func (_c *Service_Add_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, token string, cfg bootstrap.Config) (bootstrap.Config, error)) *Service_Add_Call { - _c.Call.Return(run) - return _c -} - -// AssignProfile provides a mock function for the type Service -func (_mock *Service) AssignProfile(ctx context.Context, session authn.Session, configID string, profileID string) error { - ret := _mock.Called(ctx, session, configID, profileID) - - if len(ret) == 0 { - panic("no return value specified for AssignProfile") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) error); ok { - r0 = returnFunc(ctx, session, configID, profileID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_AssignProfile_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'AssignProfile' -type Service_AssignProfile_Call struct { - *mock.Call -} - -// AssignProfile is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - configID string -// - profileID string -func (_e *Service_Expecter) AssignProfile(ctx interface{}, session interface{}, configID interface{}, profileID interface{}) *Service_AssignProfile_Call { - return &Service_AssignProfile_Call{Call: _e.mock.On("AssignProfile", ctx, session, configID, profileID)} -} - -func (_c *Service_AssignProfile_Call) Run(run func(ctx context.Context, session authn.Session, configID string, profileID string)) *Service_AssignProfile_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_AssignProfile_Call) Return(err error) *Service_AssignProfile_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_AssignProfile_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, configID string, profileID string) error) *Service_AssignProfile_Call { - _c.Call.Return(run) - return _c -} - -// BindResources provides a mock function for the type Service -func (_mock *Service) BindResources(ctx context.Context, session authn.Session, token string, configID string, bindings []bootstrap.BindingRequest) error { - ret := _mock.Called(ctx, session, token, configID, bindings) - - if len(ret) == 0 { - panic("no return value specified for BindResources") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []bootstrap.BindingRequest) error); ok { - r0 = returnFunc(ctx, session, token, configID, bindings) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_BindResources_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'BindResources' -type Service_BindResources_Call struct { - *mock.Call -} - -// BindResources is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - token string -// - configID string -// - bindings []bootstrap.BindingRequest -func (_e *Service_Expecter) BindResources(ctx interface{}, session interface{}, token interface{}, configID interface{}, bindings interface{}) *Service_BindResources_Call { - return &Service_BindResources_Call{Call: _e.mock.On("BindResources", ctx, session, token, configID, bindings)} -} - -func (_c *Service_BindResources_Call) Run(run func(ctx context.Context, session authn.Session, token string, configID string, bindings []bootstrap.BindingRequest)) *Service_BindResources_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []bootstrap.BindingRequest - if args[4] != nil { - arg4 = args[4].([]bootstrap.BindingRequest) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_BindResources_Call) Return(err error) *Service_BindResources_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_BindResources_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, token string, configID string, bindings []bootstrap.BindingRequest) error) *Service_BindResources_Call { - _c.Call.Return(run) - return _c -} - -// Bootstrap provides a mock function for the type Service -func (_mock *Service) Bootstrap(ctx context.Context, externalKey string, externalID string, secure bool) (bootstrap.Config, error) { - ret := _mock.Called(ctx, externalKey, externalID, secure) - - if len(ret) == 0 { - panic("no return value specified for Bootstrap") - } - - var r0 bootstrap.Config - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string, bool) (bootstrap.Config, error)); ok { - return returnFunc(ctx, externalKey, externalID, secure) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string, bool) bootstrap.Config); ok { - r0 = returnFunc(ctx, externalKey, externalID, secure) - } else { - r0 = ret.Get(0).(bootstrap.Config) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, string, bool) error); ok { - r1 = returnFunc(ctx, externalKey, externalID, secure) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_Bootstrap_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Bootstrap' -type Service_Bootstrap_Call struct { - *mock.Call -} - -// Bootstrap is a helper method to define mock.On call -// - ctx context.Context -// - externalKey string -// - externalID string -// - secure bool -func (_e *Service_Expecter) Bootstrap(ctx interface{}, externalKey interface{}, externalID interface{}, secure interface{}) *Service_Bootstrap_Call { - return &Service_Bootstrap_Call{Call: _e.mock.On("Bootstrap", ctx, externalKey, externalID, secure)} -} - -func (_c *Service_Bootstrap_Call) Run(run func(ctx context.Context, externalKey string, externalID string, secure bool)) *Service_Bootstrap_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 bool - if args[3] != nil { - arg3 = args[3].(bool) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_Bootstrap_Call) Return(config bootstrap.Config, err error) *Service_Bootstrap_Call { - _c.Call.Return(config, err) - return _c -} - -func (_c *Service_Bootstrap_Call) RunAndReturn(run func(ctx context.Context, externalKey string, externalID string, secure bool) (bootstrap.Config, error)) *Service_Bootstrap_Call { - _c.Call.Return(run) - return _c -} - -// CreateProfile provides a mock function for the type Service -func (_mock *Service) CreateProfile(ctx context.Context, session authn.Session, p bootstrap.Profile) (bootstrap.Profile, error) { - ret := _mock.Called(ctx, session, p) - - if len(ret) == 0 { - panic("no return value specified for CreateProfile") - } - - var r0 bootstrap.Profile - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, bootstrap.Profile) (bootstrap.Profile, error)); ok { - return returnFunc(ctx, session, p) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, bootstrap.Profile) bootstrap.Profile); ok { - r0 = returnFunc(ctx, session, p) - } else { - r0 = ret.Get(0).(bootstrap.Profile) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, bootstrap.Profile) error); ok { - r1 = returnFunc(ctx, session, p) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_CreateProfile_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'CreateProfile' -type Service_CreateProfile_Call struct { - *mock.Call -} - -// CreateProfile is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - p bootstrap.Profile -func (_e *Service_Expecter) CreateProfile(ctx interface{}, session interface{}, p interface{}) *Service_CreateProfile_Call { - return &Service_CreateProfile_Call{Call: _e.mock.On("CreateProfile", ctx, session, p)} -} - -func (_c *Service_CreateProfile_Call) Run(run func(ctx context.Context, session authn.Session, p bootstrap.Profile)) *Service_CreateProfile_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 bootstrap.Profile - if args[2] != nil { - arg2 = args[2].(bootstrap.Profile) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_CreateProfile_Call) Return(profile bootstrap.Profile, err error) *Service_CreateProfile_Call { - _c.Call.Return(profile, err) - return _c -} - -func (_c *Service_CreateProfile_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, p bootstrap.Profile) (bootstrap.Profile, error)) *Service_CreateProfile_Call { - _c.Call.Return(run) - return _c -} - -// DeleteProfile provides a mock function for the type Service -func (_mock *Service) DeleteProfile(ctx context.Context, session authn.Session, profileID string) error { - ret := _mock.Called(ctx, session, profileID) - - if len(ret) == 0 { - panic("no return value specified for DeleteProfile") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) error); ok { - r0 = returnFunc(ctx, session, profileID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_DeleteProfile_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'DeleteProfile' -type Service_DeleteProfile_Call struct { - *mock.Call -} - -// DeleteProfile is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - profileID string -func (_e *Service_Expecter) DeleteProfile(ctx interface{}, session interface{}, profileID interface{}) *Service_DeleteProfile_Call { - return &Service_DeleteProfile_Call{Call: _e.mock.On("DeleteProfile", ctx, session, profileID)} -} - -func (_c *Service_DeleteProfile_Call) Run(run func(ctx context.Context, session authn.Session, profileID string)) *Service_DeleteProfile_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_DeleteProfile_Call) Return(err error) *Service_DeleteProfile_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_DeleteProfile_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, profileID string) error) *Service_DeleteProfile_Call { - _c.Call.Return(run) - return _c -} - -// DisableConfig provides a mock function for the type Service -func (_mock *Service) DisableConfig(ctx context.Context, session authn.Session, id string) (bootstrap.Config, error) { - ret := _mock.Called(ctx, session, id) - - if len(ret) == 0 { - panic("no return value specified for DisableConfig") - } - - var r0 bootstrap.Config - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) (bootstrap.Config, error)); ok { - return returnFunc(ctx, session, id) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) bootstrap.Config); ok { - r0 = returnFunc(ctx, session, id) - } else { - r0 = ret.Get(0).(bootstrap.Config) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string) error); ok { - r1 = returnFunc(ctx, session, id) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_DisableConfig_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'DisableConfig' -type Service_DisableConfig_Call struct { - *mock.Call -} - -// DisableConfig is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - id string -func (_e *Service_Expecter) DisableConfig(ctx interface{}, session interface{}, id interface{}) *Service_DisableConfig_Call { - return &Service_DisableConfig_Call{Call: _e.mock.On("DisableConfig", ctx, session, id)} -} - -func (_c *Service_DisableConfig_Call) Run(run func(ctx context.Context, session authn.Session, id string)) *Service_DisableConfig_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_DisableConfig_Call) Return(config bootstrap.Config, err error) *Service_DisableConfig_Call { - _c.Call.Return(config, err) - return _c -} - -func (_c *Service_DisableConfig_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, id string) (bootstrap.Config, error)) *Service_DisableConfig_Call { - _c.Call.Return(run) - return _c -} - -// EnableConfig provides a mock function for the type Service -func (_mock *Service) EnableConfig(ctx context.Context, session authn.Session, id string) (bootstrap.Config, error) { - ret := _mock.Called(ctx, session, id) - - if len(ret) == 0 { - panic("no return value specified for EnableConfig") - } - - var r0 bootstrap.Config - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) (bootstrap.Config, error)); ok { - return returnFunc(ctx, session, id) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) bootstrap.Config); ok { - r0 = returnFunc(ctx, session, id) - } else { - r0 = ret.Get(0).(bootstrap.Config) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string) error); ok { - r1 = returnFunc(ctx, session, id) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_EnableConfig_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'EnableConfig' -type Service_EnableConfig_Call struct { - *mock.Call -} - -// EnableConfig is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - id string -func (_e *Service_Expecter) EnableConfig(ctx interface{}, session interface{}, id interface{}) *Service_EnableConfig_Call { - return &Service_EnableConfig_Call{Call: _e.mock.On("EnableConfig", ctx, session, id)} -} - -func (_c *Service_EnableConfig_Call) Run(run func(ctx context.Context, session authn.Session, id string)) *Service_EnableConfig_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_EnableConfig_Call) Return(config bootstrap.Config, err error) *Service_EnableConfig_Call { - _c.Call.Return(config, err) - return _c -} - -func (_c *Service_EnableConfig_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, id string) (bootstrap.Config, error)) *Service_EnableConfig_Call { - _c.Call.Return(run) - return _c -} - -// List provides a mock function for the type Service -func (_mock *Service) List(ctx context.Context, session authn.Session, filter bootstrap.Filter, offset uint64, limit uint64) (bootstrap.ConfigsPage, error) { - ret := _mock.Called(ctx, session, filter, offset, limit) - - if len(ret) == 0 { - panic("no return value specified for List") - } - - var r0 bootstrap.ConfigsPage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, bootstrap.Filter, uint64, uint64) (bootstrap.ConfigsPage, error)); ok { - return returnFunc(ctx, session, filter, offset, limit) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, bootstrap.Filter, uint64, uint64) bootstrap.ConfigsPage); ok { - r0 = returnFunc(ctx, session, filter, offset, limit) - } else { - r0 = ret.Get(0).(bootstrap.ConfigsPage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, bootstrap.Filter, uint64, uint64) error); ok { - r1 = returnFunc(ctx, session, filter, offset, limit) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_List_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'List' -type Service_List_Call struct { - *mock.Call -} - -// List is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - filter bootstrap.Filter -// - offset uint64 -// - limit uint64 -func (_e *Service_Expecter) List(ctx interface{}, session interface{}, filter interface{}, offset interface{}, limit interface{}) *Service_List_Call { - return &Service_List_Call{Call: _e.mock.On("List", ctx, session, filter, offset, limit)} -} - -func (_c *Service_List_Call) Run(run func(ctx context.Context, session authn.Session, filter bootstrap.Filter, offset uint64, limit uint64)) *Service_List_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 bootstrap.Filter - if args[2] != nil { - arg2 = args[2].(bootstrap.Filter) - } - var arg3 uint64 - if args[3] != nil { - arg3 = args[3].(uint64) - } - var arg4 uint64 - if args[4] != nil { - arg4 = args[4].(uint64) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_List_Call) Return(configsPage bootstrap.ConfigsPage, err error) *Service_List_Call { - _c.Call.Return(configsPage, err) - return _c -} - -func (_c *Service_List_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, filter bootstrap.Filter, offset uint64, limit uint64) (bootstrap.ConfigsPage, error)) *Service_List_Call { - _c.Call.Return(run) - return _c -} - -// ListBindings provides a mock function for the type Service -func (_mock *Service) ListBindings(ctx context.Context, session authn.Session, configID string) ([]bootstrap.BindingSnapshot, error) { - ret := _mock.Called(ctx, session, configID) - - if len(ret) == 0 { - panic("no return value specified for ListBindings") - } - - var r0 []bootstrap.BindingSnapshot - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) ([]bootstrap.BindingSnapshot, error)); ok { - return returnFunc(ctx, session, configID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) []bootstrap.BindingSnapshot); ok { - r0 = returnFunc(ctx, session, configID) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]bootstrap.BindingSnapshot) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string) error); ok { - r1 = returnFunc(ctx, session, configID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_ListBindings_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ListBindings' -type Service_ListBindings_Call struct { - *mock.Call -} - -// ListBindings is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - configID string -func (_e *Service_Expecter) ListBindings(ctx interface{}, session interface{}, configID interface{}) *Service_ListBindings_Call { - return &Service_ListBindings_Call{Call: _e.mock.On("ListBindings", ctx, session, configID)} -} - -func (_c *Service_ListBindings_Call) Run(run func(ctx context.Context, session authn.Session, configID string)) *Service_ListBindings_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_ListBindings_Call) Return(bindingSnapshots []bootstrap.BindingSnapshot, err error) *Service_ListBindings_Call { - _c.Call.Return(bindingSnapshots, err) - return _c -} - -func (_c *Service_ListBindings_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, configID string) ([]bootstrap.BindingSnapshot, error)) *Service_ListBindings_Call { - _c.Call.Return(run) - return _c -} - -// ListProfiles provides a mock function for the type Service -func (_mock *Service) ListProfiles(ctx context.Context, session authn.Session, offset uint64, limit uint64, name string) (bootstrap.ProfilesPage, error) { - ret := _mock.Called(ctx, session, offset, limit, name) - - if len(ret) == 0 { - panic("no return value specified for ListProfiles") - } - - var r0 bootstrap.ProfilesPage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, uint64, uint64, string) (bootstrap.ProfilesPage, error)); ok { - return returnFunc(ctx, session, offset, limit, name) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, uint64, uint64, string) bootstrap.ProfilesPage); ok { - r0 = returnFunc(ctx, session, offset, limit, name) - } else { - r0 = ret.Get(0).(bootstrap.ProfilesPage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, uint64, uint64, string) error); ok { - r1 = returnFunc(ctx, session, offset, limit, name) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_ListProfiles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ListProfiles' -type Service_ListProfiles_Call struct { - *mock.Call -} - -// ListProfiles is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - offset uint64 -// - limit uint64 -// - name string -func (_e *Service_Expecter) ListProfiles(ctx interface{}, session interface{}, offset interface{}, limit interface{}, name interface{}) *Service_ListProfiles_Call { - return &Service_ListProfiles_Call{Call: _e.mock.On("ListProfiles", ctx, session, offset, limit, name)} -} - -func (_c *Service_ListProfiles_Call) Run(run func(ctx context.Context, session authn.Session, offset uint64, limit uint64, name string)) *Service_ListProfiles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 uint64 - if args[2] != nil { - arg2 = args[2].(uint64) - } - var arg3 uint64 - if args[3] != nil { - arg3 = args[3].(uint64) - } - var arg4 string - if args[4] != nil { - arg4 = args[4].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_ListProfiles_Call) Return(profilesPage bootstrap.ProfilesPage, err error) *Service_ListProfiles_Call { - _c.Call.Return(profilesPage, err) - return _c -} - -func (_c *Service_ListProfiles_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, offset uint64, limit uint64, name string) (bootstrap.ProfilesPage, error)) *Service_ListProfiles_Call { - _c.Call.Return(run) - return _c -} - -// RefreshBindings provides a mock function for the type Service -func (_mock *Service) RefreshBindings(ctx context.Context, session authn.Session, token string, configID string) error { - ret := _mock.Called(ctx, session, token, configID) - - if len(ret) == 0 { - panic("no return value specified for RefreshBindings") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) error); ok { - r0 = returnFunc(ctx, session, token, configID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RefreshBindings_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RefreshBindings' -type Service_RefreshBindings_Call struct { - *mock.Call -} - -// RefreshBindings is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - token string -// - configID string -func (_e *Service_Expecter) RefreshBindings(ctx interface{}, session interface{}, token interface{}, configID interface{}) *Service_RefreshBindings_Call { - return &Service_RefreshBindings_Call{Call: _e.mock.On("RefreshBindings", ctx, session, token, configID)} -} - -func (_c *Service_RefreshBindings_Call) Run(run func(ctx context.Context, session authn.Session, token string, configID string)) *Service_RefreshBindings_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_RefreshBindings_Call) Return(err error) *Service_RefreshBindings_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RefreshBindings_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, token string, configID string) error) *Service_RefreshBindings_Call { - _c.Call.Return(run) - return _c -} - -// Remove provides a mock function for the type Service -func (_mock *Service) Remove(ctx context.Context, session authn.Session, id string) error { - ret := _mock.Called(ctx, session, id) - - if len(ret) == 0 { - panic("no return value specified for Remove") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) error); ok { - r0 = returnFunc(ctx, session, id) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_Remove_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Remove' -type Service_Remove_Call struct { - *mock.Call -} - -// Remove is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - id string -func (_e *Service_Expecter) Remove(ctx interface{}, session interface{}, id interface{}) *Service_Remove_Call { - return &Service_Remove_Call{Call: _e.mock.On("Remove", ctx, session, id)} -} - -func (_c *Service_Remove_Call) Run(run func(ctx context.Context, session authn.Session, id string)) *Service_Remove_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_Remove_Call) Return(err error) *Service_Remove_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_Remove_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, id string) error) *Service_Remove_Call { - _c.Call.Return(run) - return _c -} - -// Update provides a mock function for the type Service -func (_mock *Service) Update(ctx context.Context, session authn.Session, cfg bootstrap.Config) error { - ret := _mock.Called(ctx, session, cfg) - - if len(ret) == 0 { - panic("no return value specified for Update") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, bootstrap.Config) error); ok { - r0 = returnFunc(ctx, session, cfg) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_Update_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Update' -type Service_Update_Call struct { - *mock.Call -} - -// Update is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - cfg bootstrap.Config -func (_e *Service_Expecter) Update(ctx interface{}, session interface{}, cfg interface{}) *Service_Update_Call { - return &Service_Update_Call{Call: _e.mock.On("Update", ctx, session, cfg)} -} - -func (_c *Service_Update_Call) Run(run func(ctx context.Context, session authn.Session, cfg bootstrap.Config)) *Service_Update_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 bootstrap.Config - if args[2] != nil { - arg2 = args[2].(bootstrap.Config) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_Update_Call) Return(err error) *Service_Update_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_Update_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, cfg bootstrap.Config) error) *Service_Update_Call { - _c.Call.Return(run) - return _c -} - -// UpdateCert provides a mock function for the type Service -func (_mock *Service) UpdateCert(ctx context.Context, session authn.Session, id string, clientCert string, clientKey string, caCert string) (bootstrap.Config, error) { - ret := _mock.Called(ctx, session, id, clientCert, clientKey, caCert) - - if len(ret) == 0 { - panic("no return value specified for UpdateCert") - } - - var r0 bootstrap.Config - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, string, string) (bootstrap.Config, error)); ok { - return returnFunc(ctx, session, id, clientCert, clientKey, caCert) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, string, string) bootstrap.Config); ok { - r0 = returnFunc(ctx, session, id, clientCert, clientKey, caCert) - } else { - r0 = ret.Get(0).(bootstrap.Config) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, string, string) error); ok { - r1 = returnFunc(ctx, session, id, clientCert, clientKey, caCert) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_UpdateCert_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateCert' -type Service_UpdateCert_Call struct { - *mock.Call -} - -// UpdateCert is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - id string -// - clientCert string -// - clientKey string -// - caCert string -func (_e *Service_Expecter) UpdateCert(ctx interface{}, session interface{}, id interface{}, clientCert interface{}, clientKey interface{}, caCert interface{}) *Service_UpdateCert_Call { - return &Service_UpdateCert_Call{Call: _e.mock.On("UpdateCert", ctx, session, id, clientCert, clientKey, caCert)} -} - -func (_c *Service_UpdateCert_Call) Run(run func(ctx context.Context, session authn.Session, id string, clientCert string, clientKey string, caCert string)) *Service_UpdateCert_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 string - if args[4] != nil { - arg4 = args[4].(string) - } - var arg5 string - if args[5] != nil { - arg5 = args[5].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - arg5, - ) - }) - return _c -} - -func (_c *Service_UpdateCert_Call) Return(config bootstrap.Config, err error) *Service_UpdateCert_Call { - _c.Call.Return(config, err) - return _c -} - -func (_c *Service_UpdateCert_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, id string, clientCert string, clientKey string, caCert string) (bootstrap.Config, error)) *Service_UpdateCert_Call { - _c.Call.Return(run) - return _c -} - -// UpdateProfile provides a mock function for the type Service -func (_mock *Service) UpdateProfile(ctx context.Context, session authn.Session, p bootstrap.Profile) (bootstrap.Profile, error) { - ret := _mock.Called(ctx, session, p) - - if len(ret) == 0 { - panic("no return value specified for UpdateProfile") - } - - var r0 bootstrap.Profile - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, bootstrap.Profile) (bootstrap.Profile, error)); ok { - return returnFunc(ctx, session, p) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, bootstrap.Profile) bootstrap.Profile); ok { - r0 = returnFunc(ctx, session, p) - } else { - r0 = ret.Get(0).(bootstrap.Profile) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, bootstrap.Profile) error); ok { - r1 = returnFunc(ctx, session, p) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_UpdateProfile_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateProfile' -type Service_UpdateProfile_Call struct { - *mock.Call -} - -// UpdateProfile is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - p bootstrap.Profile -func (_e *Service_Expecter) UpdateProfile(ctx interface{}, session interface{}, p interface{}) *Service_UpdateProfile_Call { - return &Service_UpdateProfile_Call{Call: _e.mock.On("UpdateProfile", ctx, session, p)} -} - -func (_c *Service_UpdateProfile_Call) Run(run func(ctx context.Context, session authn.Session, p bootstrap.Profile)) *Service_UpdateProfile_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 bootstrap.Profile - if args[2] != nil { - arg2 = args[2].(bootstrap.Profile) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_UpdateProfile_Call) Return(profile bootstrap.Profile, err error) *Service_UpdateProfile_Call { - _c.Call.Return(profile, err) - return _c -} - -func (_c *Service_UpdateProfile_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, p bootstrap.Profile) (bootstrap.Profile, error)) *Service_UpdateProfile_Call { - _c.Call.Return(run) - return _c -} - -// View provides a mock function for the type Service -func (_mock *Service) View(ctx context.Context, session authn.Session, id string) (bootstrap.Config, error) { - ret := _mock.Called(ctx, session, id) - - if len(ret) == 0 { - panic("no return value specified for View") - } - - var r0 bootstrap.Config - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) (bootstrap.Config, error)); ok { - return returnFunc(ctx, session, id) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) bootstrap.Config); ok { - r0 = returnFunc(ctx, session, id) - } else { - r0 = ret.Get(0).(bootstrap.Config) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string) error); ok { - r1 = returnFunc(ctx, session, id) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_View_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'View' -type Service_View_Call struct { - *mock.Call -} - -// View is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - id string -func (_e *Service_Expecter) View(ctx interface{}, session interface{}, id interface{}) *Service_View_Call { - return &Service_View_Call{Call: _e.mock.On("View", ctx, session, id)} -} - -func (_c *Service_View_Call) Run(run func(ctx context.Context, session authn.Session, id string)) *Service_View_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_View_Call) Return(config bootstrap.Config, err error) *Service_View_Call { - _c.Call.Return(config, err) - return _c -} - -func (_c *Service_View_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, id string) (bootstrap.Config, error)) *Service_View_Call { - _c.Call.Return(run) - return _c -} - -// ViewProfile provides a mock function for the type Service -func (_mock *Service) ViewProfile(ctx context.Context, session authn.Session, profileID string) (bootstrap.Profile, error) { - ret := _mock.Called(ctx, session, profileID) - - if len(ret) == 0 { - panic("no return value specified for ViewProfile") - } - - var r0 bootstrap.Profile - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) (bootstrap.Profile, error)); ok { - return returnFunc(ctx, session, profileID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) bootstrap.Profile); ok { - r0 = returnFunc(ctx, session, profileID) - } else { - r0 = ret.Get(0).(bootstrap.Profile) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string) error); ok { - r1 = returnFunc(ctx, session, profileID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_ViewProfile_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ViewProfile' -type Service_ViewProfile_Call struct { - *mock.Call -} - -// ViewProfile is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - profileID string -func (_e *Service_Expecter) ViewProfile(ctx interface{}, session interface{}, profileID interface{}) *Service_ViewProfile_Call { - return &Service_ViewProfile_Call{Call: _e.mock.On("ViewProfile", ctx, session, profileID)} -} - -func (_c *Service_ViewProfile_Call) Run(run func(ctx context.Context, session authn.Session, profileID string)) *Service_ViewProfile_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_ViewProfile_Call) Return(profile bootstrap.Profile, err error) *Service_ViewProfile_Call { - _c.Call.Return(profile, err) - return _c -} - -func (_c *Service_ViewProfile_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, profileID string) (bootstrap.Profile, error)) *Service_ViewProfile_Call { - _c.Call.Return(run) - return _c -} diff --git a/bootstrap/postgres/bindings.go b/bootstrap/postgres/bindings.go deleted file mode 100644 index 77bacba07..000000000 --- a/bootstrap/postgres/bindings.go +++ /dev/null @@ -1,146 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package postgres - -import ( - "context" - "encoding/json" - "fmt" - "log/slog" - "time" - - "github.com/absmach/magistrala/bootstrap" - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - "github.com/absmach/magistrala/pkg/postgres" -) - -var _ bootstrap.BindingStore = (*bindingRepository)(nil) - -type bindingRepository struct { - db postgres.Database - log *slog.Logger -} - -// NewBindingRepository instantiates a PostgreSQL implementation of BindingStore. -func NewBindingRepository(db postgres.Database, log *slog.Logger) bootstrap.BindingStore { - return &bindingRepository{db: db, log: log} -} - -func (br bindingRepository) Save(ctx context.Context, configID string, bindings []bootstrap.BindingSnapshot) error { - if len(bindings) == 0 { - return nil - } - q := `INSERT INTO bindings (config_id, slot, type, resource_id, snapshot, secret_snapshot, updated_at) - VALUES (:config_id, :slot, :type, :resource_id, :snapshot, :secret_snapshot, :updated_at) - ON CONFLICT (config_id, slot) DO UPDATE SET - type = EXCLUDED.type, - resource_id = EXCLUDED.resource_id, - snapshot = EXCLUDED.snapshot, - secret_snapshot = EXCLUDED.secret_snapshot, - updated_at = EXCLUDED.updated_at` - - now := time.Now().UTC() - dbBindings := make([]dbBindingSnapshot, 0, len(bindings)) - for _, b := range bindings { - b.ConfigID = configID - b.UpdatedAt = now - dbb, err := toDBBindingSnapshot(b) - if err != nil { - return errors.Wrap(repoerr.ErrCreateEntity, err) - } - dbBindings = append(dbBindings, dbb) - } - - if _, err := br.db.NamedExecContext(ctx, q, dbBindings); err != nil { - return errors.Wrap(repoerr.ErrCreateEntity, err) - } - return nil -} - -func (br bindingRepository) Retrieve(ctx context.Context, configID string) ([]bootstrap.BindingSnapshot, error) { - q := `SELECT config_id, slot, type, resource_id, snapshot, secret_snapshot, updated_at - FROM bindings WHERE config_id = $1 ORDER BY slot` - - rows, err := br.db.QueryxContext(ctx, q, configID) - if err != nil { - return nil, errors.Wrap(repoerr.ErrViewEntity, err) - } - defer rows.Close() - - var snapshots []bootstrap.BindingSnapshot - for rows.Next() { - var dbb dbBindingSnapshot - if err := rows.StructScan(&dbb); err != nil { - br.log.Error(fmt.Sprintf("failed to scan binding snapshot: %s", err)) - return nil, errors.Wrap(repoerr.ErrViewEntity, err) - } - b, err := toBindingSnapshot(dbb) - if err != nil { - return nil, errors.Wrap(repoerr.ErrViewEntity, err) - } - snapshots = append(snapshots, b) - } - return snapshots, nil -} - -func (br bindingRepository) Delete(ctx context.Context, configID, slot string) error { - q := `DELETE FROM bindings WHERE config_id = $1 AND slot = $2` - if _, err := br.db.ExecContext(ctx, q, configID, slot); err != nil { - return errors.Wrap(repoerr.ErrRemoveEntity, err) - } - return nil -} - -// dbBindingSnapshot is the database representation of a BindingSnapshot. -type dbBindingSnapshot struct { - ConfigID string `db:"config_id"` - Slot string `db:"slot"` - Type string `db:"type"` - ResourceID string `db:"resource_id"` - Snapshot []byte `db:"snapshot"` - SecretSnapshot []byte `db:"secret_snapshot"` - UpdatedAt time.Time `db:"updated_at"` -} - -func toDBBindingSnapshot(b bootstrap.BindingSnapshot) (dbBindingSnapshot, error) { - snap, err := json.Marshal(b.Snapshot) - if err != nil { - return dbBindingSnapshot{}, err - } - secret, err := json.Marshal(b.SecretSnapshot) - if err != nil { - return dbBindingSnapshot{}, err - } - return dbBindingSnapshot{ - ConfigID: b.ConfigID, - Slot: b.Slot, - Type: b.Type, - ResourceID: b.ResourceID, - Snapshot: snap, - SecretSnapshot: secret, - UpdatedAt: b.UpdatedAt, - }, nil -} - -func toBindingSnapshot(dbb dbBindingSnapshot) (bootstrap.BindingSnapshot, error) { - b := bootstrap.BindingSnapshot{ - ConfigID: dbb.ConfigID, - Slot: dbb.Slot, - Type: dbb.Type, - ResourceID: dbb.ResourceID, - UpdatedAt: dbb.UpdatedAt, - } - if len(dbb.Snapshot) > 0 && string(dbb.Snapshot) != jsonNull { - if err := json.Unmarshal(dbb.Snapshot, &b.Snapshot); err != nil { - return bootstrap.BindingSnapshot{}, err - } - } - if len(dbb.SecretSnapshot) > 0 && string(dbb.SecretSnapshot) != jsonNull { - if err := json.Unmarshal(dbb.SecretSnapshot, &b.SecretSnapshot); err != nil { - return bootstrap.BindingSnapshot{}, err - } - } - return b, nil -} diff --git a/bootstrap/postgres/configs.go b/bootstrap/postgres/configs.go deleted file mode 100644 index 9c588fb25..000000000 --- a/bootstrap/postgres/configs.go +++ /dev/null @@ -1,420 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package postgres - -import ( - "context" - "database/sql" - "encoding/json" - "fmt" - "log/slog" - "strings" - - "github.com/absmach/magistrala/bootstrap" - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - "github.com/absmach/magistrala/pkg/postgres" - "github.com/jackc/pgerrcode" - "github.com/jackc/pgx/v5/pgconn" -) - -const jsonNull = "null" - -var _ bootstrap.ConfigRepository = (*configRepository)(nil) - -type configRepository struct { - db postgres.Database - log *slog.Logger -} - -// NewConfigRepository instantiates a PostgreSQL implementation of config -// repository. -func NewConfigRepository(db postgres.Database, log *slog.Logger) bootstrap.ConfigRepository { - return &configRepository{db: db, log: log} -} - -func (cr configRepository) Save(ctx context.Context, cfg bootstrap.Config) (string, error) { - q := `INSERT INTO configs (id, domain_id, name, client_cert, client_key, ca_cert, external_id, external_key, content, status, profile_id, render_context) - VALUES (:id, :domain_id, :name, :client_cert, :client_key, :ca_cert, :external_id, :external_key, :content, :status, :profile_id, :render_context)` - - dbcfg, err := toDBConfig(cfg) - if err != nil { - return "", errors.Wrap(repoerr.ErrCreateEntity, err) - } - if _, err := cr.db.NamedExecContext(ctx, q, dbcfg); err != nil { - switch pgErr := err.(type) { - case *pgconn.PgError: - if pgErr.Code == pgerrcode.UniqueViolation { - return "", repoerr.ErrConflict - } - } - return "", errors.Wrap(repoerr.ErrCreateEntity, err) - } - - return cfg.ID, nil -} - -func (cr configRepository) RetrieveByID(ctx context.Context, domainID, id string) (bootstrap.Config, error) { - q := `SELECT id, external_id, name, content, status, client_cert, client_key, ca_cert, profile_id, render_context - FROM configs - WHERE id = :id AND domain_id = :domain_id` - - dbcfg := dbConfig{ - ID: id, - DomainID: domainID, - } - row, err := cr.db.NamedQueryContext(ctx, q, dbcfg) - if err != nil { - return bootstrap.Config{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - - if !row.Next() { - return bootstrap.Config{}, repoerr.ErrNotFound - } - - if err := row.StructScan(&dbcfg); err != nil { - return bootstrap.Config{}, err - } - - cfg, err := toConfig(dbcfg) - if err != nil { - return bootstrap.Config{}, err - } - return cfg, nil -} - -func (cr configRepository) RetrieveAll(ctx context.Context, domainID string, filter bootstrap.Filter, offset, limit uint64) bootstrap.ConfigsPage { - search, params := buildRetrieveQueryParams(domainID, filter) - n := len(params) - - q := `SELECT id, external_id, name, content, status, profile_id, render_context - FROM configs %s ORDER BY id LIMIT $%d OFFSET $%d` - q = fmt.Sprintf(q, search, n+1, n+2) - - rows, err := cr.db.QueryContext(ctx, q, append(params, limit, offset)...) - if err != nil { - cr.log.Error(fmt.Sprintf("Failed to retrieve configs due to %s", err)) - return bootstrap.ConfigsPage{} - } - defer rows.Close() - - var name, content, profileID sql.NullString - var renderContext []byte - configs := []bootstrap.Config{} - - for rows.Next() { - c := bootstrap.Config{DomainID: domainID} - if err := rows.Scan(&c.ID, &c.ExternalID, &name, &content, &c.Status, &profileID, &renderContext); err != nil { - cr.log.Error(fmt.Sprintf("Failed to read retrieved config due to %s", err)) - return bootstrap.ConfigsPage{} - } - - c.Name = name.String - c.Content = content.String - if profileID.Valid { - c.ProfileID = profileID.String - } - if len(renderContext) > 0 && string(renderContext) != jsonNull { - if err := json.Unmarshal(renderContext, &c.RenderContext); err != nil { - cr.log.Error(fmt.Sprintf("Failed to decode render context due to %s", err)) - return bootstrap.ConfigsPage{} - } - } - configs = append(configs, c) - } - - q = fmt.Sprintf(`SELECT COUNT(*) FROM configs %s`, search) - - var total uint64 - if err := cr.db.QueryRowxContext(ctx, q, params...).Scan(&total); err != nil { - cr.log.Error(fmt.Sprintf("Failed to count configs due to %s", err)) - return bootstrap.ConfigsPage{} - } - - return bootstrap.ConfigsPage{ - Total: total, - Limit: limit, - Offset: offset, - Configs: configs, - } -} - -func (cr configRepository) RetrieveByExternalID(ctx context.Context, externalID string) (bootstrap.Config, error) { - q := `SELECT id, external_key, domain_id, name, client_cert, client_key, ca_cert, content, status, profile_id, render_context - FROM configs - WHERE external_id = :external_id` - dbcfg := dbConfig{ - ExternalID: externalID, - } - - row, err := cr.db.NamedQueryContext(ctx, q, dbcfg) - if err != nil { - return bootstrap.Config{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - - if !row.Next() { - return bootstrap.Config{}, repoerr.ErrNotFound - } - - if err := row.StructScan(&dbcfg); err != nil { - return bootstrap.Config{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - - cfg, err := toConfig(dbcfg) - if err != nil { - return bootstrap.Config{}, err - } - return cfg, nil -} - -func (cr configRepository) Update(ctx context.Context, cfg bootstrap.Config) error { - q := `UPDATE configs SET name = :name, content = :content, render_context = :render_context WHERE id = :id AND domain_id = :domain_id ` - - renderContext, err := json.Marshal(cfg.RenderContext) - if err != nil { - return errors.Wrap(repoerr.ErrUpdateEntity, err) - } - - dbcfg := dbConfig{ - Name: nullString(cfg.Name), - Content: nullString(cfg.Content), - RenderContext: renderContext, - ID: cfg.ID, - DomainID: cfg.DomainID, - } - - res, err := cr.db.NamedExecContext(ctx, q, dbcfg) - if err != nil { - return errors.Wrap(repoerr.ErrUpdateEntity, err) - } - - cnt, err := res.RowsAffected() - if err != nil { - return errors.Wrap(repoerr.ErrUpdateEntity, err) - } - - if cnt == 0 { - return repoerr.ErrNotFound - } - - return nil -} - -func (cr configRepository) AssignProfile(ctx context.Context, domainID, id, profileID string) error { - q := `UPDATE configs SET profile_id = :profile_id WHERE id = :id AND domain_id = :domain_id` - - dbcfg := dbConfig{ - ID: id, - DomainID: domainID, - ProfileID: nullString(profileID), - } - - res, err := cr.db.NamedExecContext(ctx, q, dbcfg) - if err != nil { - return errors.Wrap(repoerr.ErrUpdateEntity, err) - } - - cnt, err := res.RowsAffected() - if err != nil { - return errors.Wrap(repoerr.ErrUpdateEntity, err) - } - - if cnt == 0 { - return repoerr.ErrNotFound - } - - return nil -} - -func (cr configRepository) UpdateCert(ctx context.Context, domainID, id, clientCert, clientKey, caCert string) (bootstrap.Config, error) { - q := `UPDATE configs SET client_cert = :client_cert, client_key = :client_key, ca_cert = :ca_cert WHERE id = :id AND domain_id = :domain_id - RETURNING id, client_cert, client_key, ca_cert, domain_id` - - dbcfg := dbConfig{ - ID: id, - ClientCert: nullString(clientCert), - DomainID: domainID, - ClientKey: nullString(clientKey), - CaCert: nullString(caCert), - } - - row, err := cr.db.NamedQueryContext(ctx, q, dbcfg) - if err != nil { - return bootstrap.Config{}, errors.Wrap(repoerr.ErrUpdateEntity, err) - } - defer row.Close() - - if ok := row.Next(); !ok { - return bootstrap.Config{}, errors.Wrap(repoerr.ErrNotFound, row.Err()) - } - - if err := row.StructScan(&dbcfg); err != nil { - return bootstrap.Config{}, err - } - - cfg, err := toConfig(dbcfg) - if err != nil { - return bootstrap.Config{}, err - } - return cfg, nil -} - -func (cr configRepository) Remove(ctx context.Context, domainID, id string) error { - q := `DELETE FROM configs WHERE id = :id AND domain_id = :domain_id` - dbcfg := dbConfig{ - ID: id, - DomainID: domainID, - } - - if _, err := cr.db.NamedExecContext(ctx, q, dbcfg); err != nil { - return errors.Wrap(repoerr.ErrRemoveEntity, err) - } - - return nil -} - -func (cr configRepository) ChangeStatus(ctx context.Context, domainID, id string, status bootstrap.Status) error { - q := `UPDATE configs SET status = :status WHERE id = :id AND domain_id = :domain_id;` - - dbcfg := dbConfig{ - ID: id, - Status: status, - DomainID: domainID, - } - - res, err := cr.db.NamedExecContext(ctx, q, dbcfg) - if err != nil { - return errors.Wrap(repoerr.ErrUpdateEntity, err) - } - - cnt, err := res.RowsAffected() - if err != nil { - return errors.Wrap(repoerr.ErrUpdateEntity, err) - } - - if cnt == 0 { - return repoerr.ErrNotFound - } - - return nil -} - -func buildRetrieveQueryParams(domainID string, filter bootstrap.Filter) (string, []any) { - params := []any{} - queries := []string{} - - if domainID != "" { - params = append(params, domainID) - queries = append(queries, fmt.Sprintf("domain_id = $%d", len(params))) - } - - counter := len(params) + 1 - for k, v := range filter.FullMatch { - if k == "status" { - status, err := bootstrap.ToStatus(v) - if err != nil { - return "", nil - } - if status == bootstrap.AllStatus { - continue - } - params = append(params, status) - queries = append(queries, fmt.Sprintf("%s = $%d", k, counter)) - counter++ - continue - } - params = append(params, v) - queries = append(queries, fmt.Sprintf("%s = $%d", k, counter)) - counter++ - } - for k, v := range filter.PartialMatch { - params = append(params, v) - queries = append(queries, fmt.Sprintf("LOWER(%s) LIKE '%%' || $%d || '%%'", k, counter)) - counter++ - } - - if len(queries) > 0 { - return "WHERE " + strings.Join(queries, " AND "), params - } - return "", params -} - -func nullString(s string) sql.NullString { - if s == "" { - return sql.NullString{} - } - return sql.NullString{String: s, Valid: true} -} - -type dbConfig struct { - DomainID string `db:"domain_id"` - ID string `db:"id"` - Name sql.NullString `db:"name"` - ClientCert sql.NullString `db:"client_cert"` - ClientKey sql.NullString `db:"client_key"` - CaCert sql.NullString `db:"ca_cert"` - ExternalID string `db:"external_id"` - ExternalKey string `db:"external_key"` - Content sql.NullString `db:"content"` - Status bootstrap.Status `db:"status"` - ProfileID sql.NullString `db:"profile_id"` - RenderContext []byte `db:"render_context"` -} - -func toDBConfig(cfg bootstrap.Config) (dbConfig, error) { - renderContext, err := json.Marshal(cfg.RenderContext) - if err != nil { - return dbConfig{}, err - } - - return dbConfig{ - ID: cfg.ID, - DomainID: cfg.DomainID, - Name: nullString(cfg.Name), - ClientCert: nullString(cfg.ClientCert), - ClientKey: nullString(cfg.ClientKey), - CaCert: nullString(cfg.CACert), - ExternalID: cfg.ExternalID, - ExternalKey: cfg.ExternalKey, - Content: nullString(cfg.Content), - Status: cfg.Status, - ProfileID: nullString(cfg.ProfileID), - RenderContext: renderContext, - }, nil -} - -func toConfig(dbcfg dbConfig) (bootstrap.Config, error) { - cfg := bootstrap.Config{ - ID: dbcfg.ID, - DomainID: dbcfg.DomainID, - ExternalID: dbcfg.ExternalID, - ExternalKey: dbcfg.ExternalKey, - Status: dbcfg.Status, - } - if dbcfg.ProfileID.Valid { - cfg.ProfileID = dbcfg.ProfileID.String - } - - if dbcfg.Name.Valid { - cfg.Name = dbcfg.Name.String - } - if dbcfg.Content.Valid { - cfg.Content = dbcfg.Content.String - } - if len(dbcfg.RenderContext) > 0 && string(dbcfg.RenderContext) != jsonNull { - if err := json.Unmarshal(dbcfg.RenderContext, &cfg.RenderContext); err != nil { - return bootstrap.Config{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - } - if dbcfg.ClientCert.Valid { - cfg.ClientCert = dbcfg.ClientCert.String - } - if dbcfg.ClientKey.Valid { - cfg.ClientKey = dbcfg.ClientKey.String - } - if dbcfg.CaCert.Valid { - cfg.CACert = dbcfg.CaCert.String - } - return cfg, nil -} diff --git a/bootstrap/postgres/configs_test.go b/bootstrap/postgres/configs_test.go deleted file mode 100644 index 064bcb95c..000000000 --- a/bootstrap/postgres/configs_test.go +++ /dev/null @@ -1,471 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package postgres_test - -import ( - "context" - "fmt" - "strconv" - "testing" - - "github.com/absmach/magistrala/bootstrap" - "github.com/absmach/magistrala/bootstrap/postgres" - "github.com/absmach/magistrala/internal/testsutil" - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - "github.com/gofrs/uuid/v5" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" -) - -const numConfigs = 10 - -var config = bootstrap.Config{ - ID: "mg-client", - ExternalID: "external-id", - ExternalKey: "external-key", - DomainID: testsutil.GenerateUUID(&testing.T{}), - Content: "content", - Status: bootstrap.Inactive, -} - -func TestSave(t *testing.T) { - repo := postgres.NewConfigRepository(db, testLog) - - diff := "different" - - duplicateClient := config - duplicateClient.ExternalID = diff - - duplicateExternal := config - duplicateExternal.ID = diff - - cases := []struct { - desc string - config bootstrap.Config - err error - }{ - { - desc: "save a config", - config: config, - err: nil, - }, - { - desc: "save config with same Client ID", - config: duplicateClient, - err: repoerr.ErrConflict, - }, - { - desc: "save config with same external ID", - config: duplicateExternal, - err: repoerr.ErrConflict, - }, - } - for _, tc := range cases { - id, err := repo.Save(context.Background(), tc.config) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - if err == nil { - assert.Equal(t, id, tc.config.ID, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.config.ID, id)) - } - } -} - -func TestRetrieveByID(t *testing.T) { - repo := postgres.NewConfigRepository(db, testLog) - - c := config - // Use UUID to prevent conflicts. - uid, err := uuid.NewV4() - require.Nil(t, err, fmt.Sprintf("Got unexpected error: %s.\n", err)) - c.ID = uid.String() - c.ExternalID = uid.String() - c.ExternalKey = uid.String() - id, err := repo.Save(context.Background(), c) - require.Nil(t, err, fmt.Sprintf("Saving config expected to succeed: %s.\n", err)) - - nonexistentConfID, err := uuid.NewV4() - require.Nil(t, err, fmt.Sprintf("Got unexpected error: %s.\n", err)) - - cases := []struct { - desc string - domainID string - id string - err error - }{ - { - desc: "retrieve config", - domainID: c.DomainID, - id: id, - err: nil, - }, - { - desc: "retrieve config with wrong domain ID ", - domainID: "2", - id: id, - err: repoerr.ErrNotFound, - }, - { - desc: "retrieve a non-existing config", - domainID: c.DomainID, - id: nonexistentConfID.String(), - err: repoerr.ErrNotFound, - }, - { - desc: "retrieve a config with invalid ID", - domainID: c.DomainID, - id: "invalid", - err: repoerr.ErrNotFound, - }, - } - for _, tc := range cases { - _, err := repo.RetrieveByID(context.Background(), tc.domainID, tc.id) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - } -} - -func TestRetrieveAll(t *testing.T) { - repo := postgres.NewConfigRepository(db, testLog) - - for i := 0; i < numConfigs; i++ { - c := config - - // Use UUID to prevent conflict errors. - uid, err := uuid.NewV4() - require.Nil(t, err, fmt.Sprintf("Got unexpected error: %s.\n", err)) - c.ExternalID = uid.String() - c.Name = fmt.Sprintf("name %d", i) - c.ID = uid.String() - - if i%2 == 0 { - c.Status = bootstrap.Active - } - - _, err = repo.Save(context.Background(), c) - require.Nil(t, err, fmt.Sprintf("Saving config expected to succeed: %s.\n", err)) - } - cases := []struct { - desc string - domainID string - offset uint64 - limit uint64 - filter bootstrap.Filter - size int - }{ - { - desc: "retrieve all configs", - domainID: config.DomainID, - offset: 0, - limit: uint64(numConfigs), - size: numConfigs, - }, - { - desc: "retrieve a subset of configs", - domainID: config.DomainID, - offset: 5, - limit: uint64(numConfigs - 5), - size: numConfigs - 5, - }, - { - desc: "retrieve with wrong domain ID ", - domainID: "2", - offset: 0, - limit: uint64(numConfigs), - size: 0, - }, - { - desc: "retrieve all active configs ", - domainID: config.DomainID, - offset: 0, - limit: uint64(numConfigs), - filter: bootstrap.Filter{FullMatch: map[string]string{"status": bootstrap.Active.String()}}, - size: numConfigs / 2, - }, - { - desc: "retrieve all with partial match filter", - domainID: config.DomainID, - offset: 0, - limit: uint64(numConfigs), - filter: bootstrap.Filter{PartialMatch: map[string]string{"name": "1"}}, - size: 1, - }, - { - desc: "retrieve search by name", - domainID: config.DomainID, - offset: 0, - limit: uint64(numConfigs), - filter: bootstrap.Filter{PartialMatch: map[string]string{"name": "1"}}, - size: 1, - }, - } - for _, tc := range cases { - ret := repo.RetrieveAll(context.Background(), tc.domainID, tc.filter, tc.offset, tc.limit) - size := len(ret.Configs) - assert.Equal(t, tc.size, size, fmt.Sprintf("%s: expected %d got %d\n", tc.desc, tc.size, size)) - } -} - -func TestRetrieveByExternalID(t *testing.T) { - repo := postgres.NewConfigRepository(db, testLog) - - c := config - // Use UUID to prevent conflicts. - uid, err := uuid.NewV4() - assert.Nil(t, err, fmt.Sprintf("Got unexpected error: %s.\n", err)) - c.ID = uid.String() - c.ExternalID = uid.String() - c.ExternalKey = uid.String() - _, err = repo.Save(context.Background(), c) - assert.Nil(t, err, fmt.Sprintf("Saving config expected to succeed: %s.\n", err)) - - cases := []struct { - desc string - externalID string - err error - }{ - { - desc: "retrieve with invalid external ID", - externalID: strconv.Itoa(numConfigs + 1), - err: repoerr.ErrNotFound, - }, - { - desc: "retrieve with external key", - externalID: c.ExternalID, - err: nil, - }, - } - for _, tc := range cases { - _, err := repo.RetrieveByExternalID(context.Background(), tc.externalID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - } -} - -func TestUpdate(t *testing.T) { - repo := postgres.NewConfigRepository(db, testLog) - - c := config - // Use UUID to prevent conflicts. - uid, err := uuid.NewV4() - assert.Nil(t, err, fmt.Sprintf("Got unexpected error: %s.\n", err)) - c.ID = uid.String() - c.ExternalID = uid.String() - c.ExternalKey = uid.String() - _, err = repo.Save(context.Background(), c) - assert.Nil(t, err, fmt.Sprintf("Saving config expected to succeed: %s.\n", err)) - - c.Content = "new content" - c.Name = "new name" - - withRenderContext := c - withRenderContext.RenderContext = map[string]any{ - "site": "warehouse-2", - "region": "mombasa", - } - - wrongDomainID := c - wrongDomainID.DomainID = "3" - - cases := []struct { - desc string - config bootstrap.Config - renderContext map[string]any - err error - }{ - { - desc: "update with wrong domainID", - config: wrongDomainID, - err: repoerr.ErrNotFound, - }, - { - desc: "update a config", - config: c, - err: nil, - }, - { - desc: "update a config render_context", - config: withRenderContext, - renderContext: map[string]any{"site": "warehouse-2", "region": "mombasa"}, - err: nil, - }, - } - for _, tc := range cases { - err := repo.Update(context.Background(), tc.config) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - if tc.err == nil && tc.renderContext != nil { - saved, err := repo.RetrieveByID(context.Background(), tc.config.DomainID, tc.config.ID) - require.Nil(t, err, fmt.Sprintf("%s: unexpected retrieve error: %s\n", tc.desc, err)) - assert.Equal(t, tc.renderContext, saved.RenderContext, fmt.Sprintf("%s: expected render_context %v got %v\n", tc.desc, tc.renderContext, saved.RenderContext)) - } - } -} - -func TestUpdateCert(t *testing.T) { - repo := postgres.NewConfigRepository(db, testLog) - - c := config - // Use UUID to prevent conflicts. - uid, err := uuid.NewV4() - assert.Nil(t, err, fmt.Sprintf("Got unexpected error: %s.\n", err)) - c.ID = uid.String() - c.ExternalID = uid.String() - c.ExternalKey = uid.String() - _, err = repo.Save(context.Background(), c) - assert.Nil(t, err, fmt.Sprintf("Saving config expected to succeed: %s.\n", err)) - - c.Content = "new content" - c.Name = "new name" - - wrongDomainID := c - wrongDomainID.DomainID = "3" - - cases := []struct { - desc string - configID string - domainID string - cert string - certKey string - ca string - expectedConfig bootstrap.Config - err error - }{ - { - desc: "update with wrong domain ID ", - configID: "", - cert: "cert", - certKey: "certKey", - ca: "", - domainID: wrongDomainID.DomainID, - expectedConfig: bootstrap.Config{}, - err: repoerr.ErrNotFound, - }, - { - desc: "update a config", - configID: c.ID, - cert: "cert", - certKey: "certKey", - ca: "ca", - domainID: c.DomainID, - expectedConfig: bootstrap.Config{ - ID: c.ID, - ClientCert: "cert", - CACert: "ca", - ClientKey: "certKey", - DomainID: c.DomainID, - }, - err: nil, - }, - } - for _, tc := range cases { - cfg, err := repo.UpdateCert(context.Background(), tc.domainID, tc.configID, tc.cert, tc.certKey, tc.ca) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.expectedConfig, cfg, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.expectedConfig, cfg)) - } -} - -func TestRemove(t *testing.T) { - repo := postgres.NewConfigRepository(db, testLog) - - c := config - // Use UUID to prevent conflicts. - uid, err := uuid.NewV4() - assert.Nil(t, err, fmt.Sprintf("Got unexpected error: %s.\n", err)) - c.ID = uid.String() - c.ExternalID = uid.String() - c.ExternalKey = uid.String() - id, err := repo.Save(context.Background(), c) - assert.Nil(t, err, fmt.Sprintf("Saving config expected to succeed: %s.\n", err)) - - // Removal works the same for both existing and non-existing - // (removed) config - for i := 0; i < 2; i++ { - err := repo.Remove(context.Background(), c.DomainID, id) - assert.Nil(t, err, fmt.Sprintf("%d: failed to remove config due to: %s", i, err)) - - _, err = repo.RetrieveByID(context.Background(), c.DomainID, id) - assert.True(t, errors.Contains(err, repoerr.ErrNotFound), fmt.Sprintf("%d: expected %s got %s", i, repoerr.ErrNotFound, err)) - } -} - -func TestChangeStatus(t *testing.T) { - repo := postgres.NewConfigRepository(db, testLog) - - c := config - // Use UUID to prevent conflicts. - uid, err := uuid.NewV4() - assert.Nil(t, err, fmt.Sprintf("Got unexpected error: %s.\n", err)) - c.ID = uid.String() - c.ExternalID = uid.String() - c.ExternalKey = uid.String() - saved, err := repo.Save(context.Background(), c) - assert.Nil(t, err, fmt.Sprintf("Saving config expected to succeed: %s.\n", err)) - - cases := []struct { - desc string - domainID string - id string - status bootstrap.Status - err error - }{ - { - desc: "change status with wrong domain ID ", - id: saved, - domainID: "2", - err: repoerr.ErrNotFound, - }, - { - desc: "change status with wrong id", - id: "wrong", - domainID: c.DomainID, - err: repoerr.ErrNotFound, - }, - { - desc: "change status to Active", - id: saved, - domainID: c.DomainID, - status: bootstrap.Active, - err: nil, - }, - { - desc: "change status to Inactive", - id: saved, - domainID: c.DomainID, - status: bootstrap.Inactive, - err: nil, - }, - } - for _, tc := range cases { - err := repo.ChangeStatus(context.Background(), tc.domainID, tc.id, tc.status) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - } -} - -func TestAssignProfile(t *testing.T) { - configRepo := postgres.NewConfigRepository(db, testLog) - profileRepo := postgres.NewProfileRepository(db, testLog) - - c := config - uid, err := uuid.NewV4() - require.Nil(t, err, fmt.Sprintf("Got unexpected error: %s.\n", err)) - c.ID = uid.String() - c.ExternalID = uid.String() - c.ExternalKey = uid.String() - saved, err := configRepo.Save(context.Background(), c) - require.Nil(t, err, fmt.Sprintf("Saving config expected to succeed: %s.\n", err)) - - profileID := testsutil.GenerateUUID(t) - _, err = profileRepo.Save(context.Background(), bootstrap.Profile{ - ID: profileID, - DomainID: c.DomainID, - Name: "edge-gateway", - ContentFormat: bootstrap.ContentFormatGoTemplate, - Version: 1, - }) - require.Nil(t, err, fmt.Sprintf("Saving profile expected to succeed: %s.\n", err)) - - err = configRepo.AssignProfile(context.Background(), c.DomainID, saved, profileID) - require.Nil(t, err, fmt.Sprintf("Assigning profile expected to succeed: %s.\n", err)) - - stored, err := configRepo.RetrieveByID(context.Background(), c.DomainID, saved) - require.Nil(t, err, fmt.Sprintf("Retrieving config expected to succeed: %s.\n", err)) - assert.Equal(t, profileID, stored.ProfileID, "expected profile assignment to round-trip through the repository") -} diff --git a/bootstrap/postgres/doc.go b/bootstrap/postgres/doc.go deleted file mode 100644 index 73a678477..000000000 --- a/bootstrap/postgres/doc.go +++ /dev/null @@ -1,6 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package postgres contains repository implementations using PostgreSQL as -// the underlying database. -package postgres diff --git a/bootstrap/postgres/init.go b/bootstrap/postgres/init.go deleted file mode 100644 index 928e1dfc4..000000000 --- a/bootstrap/postgres/init.go +++ /dev/null @@ -1,329 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package postgres - -import migrate "github.com/rubenv/sql-migrate" - -// Migration of bootstrap service. -func Migration() *migrate.MemoryMigrationSource { - return &migrate.MemoryMigrationSource{ - Migrations: []*migrate.Migration{ - { - Id: "configs_1", - Up: []string{ - `CREATE TABLE IF NOT EXISTS configs ( - mainflux_client TEXT UNIQUE NOT NULL, - owner VARCHAR(254), - name TEXT, - mainflux_key CHAR(36) UNIQUE NOT NULL, - external_id TEXT UNIQUE NOT NULL, - external_key TEXT NOT NULL, - content TEXT, - client_cert TEXT, - client_key TEXT, - ca_cert TEXT, - state BIGINT NOT NULL, - PRIMARY KEY (mainflux_client, owner) - )`, - `CREATE TABLE IF NOT EXISTS unknown_configs ( - external_id TEXT UNIQUE NOT NULL, - external_key TEXT NOT NULL, - PRIMARY KEY (external_id, external_key) - )`, - `CREATE TABLE IF NOT EXISTS channels ( - mainflux_channel TEXT UNIQUE NOT NULL, - owner VARCHAR(254), - name TEXT, - metadata JSON, - PRIMARY KEY (mainflux_channel, owner) - )`, - `CREATE TABLE IF NOT EXISTS connections ( - channel_id TEXT, - channel_owner VARCHAR(256), - config_id TEXT, - config_owner VARCHAR(256), - FOREIGN KEY (channel_id, channel_owner) REFERENCES channels (mainflux_channel, owner) ON DELETE CASCADE ON UPDATE CASCADE, - FOREIGN KEY (config_id, config_owner) REFERENCES configs (mainflux_client, owner) ON DELETE CASCADE ON UPDATE CASCADE, - PRIMARY KEY (channel_id, channel_owner, config_id, config_owner) - )`, - }, - Down: []string{ - "DROP TABLE connections", - "DROP TABLE configs", - "DROP TABLE channels", - "DROP TABLE unknown_configs", - }, - }, - { - Id: "configs_2", - Up: []string{ - "DROP TABLE IF EXISTS unknown_configs", - }, - Down: []string{ - "CREATE TABLE IF NOT EXISTS unknown_configs", - }, - }, - { - Id: "configs_3", - Up: []string{ - `ALTER TABLE IF EXISTS channels ADD COLUMN IF NOT EXISTS parent_id VARCHAR(36)`, - `ALTER TABLE IF EXISTS channels ADD COLUMN IF NOT EXISTS description VARCHAR(1024)`, - `ALTER TABLE IF EXISTS channels ADD COLUMN IF NOT EXISTS created_at TIMESTAMP`, - `ALTER TABLE IF EXISTS channels ADD COLUMN IF NOT EXISTS updated_at TIMESTAMP`, - `ALTER TABLE IF EXISTS channels ADD COLUMN IF NOT EXISTS updated_by VARCHAR(254)`, - `ALTER TABLE IF EXISTS channels ADD COLUMN IF NOT EXISTS status SMALLINT NOT NULL DEFAULT 0 CHECK (status >= 0)`, - }, - }, - { - Id: "configs_4", - Up: []string{ - `ALTER TABLE IF EXISTS configs RENAME COLUMN mainflux_client TO magistrala_client`, - `ALTER TABLE IF EXISTS configs RENAME COLUMN mainflux_key TO magistrala_secret`, - `ALTER TABLE IF EXISTS channels RENAME COLUMN mainflux_channel TO magistrala_channel`, - }, - }, - { - Id: "configs_5", - Up: []string{ - `ALTER TABLE IF EXISTS configs RENAME COLUMN owner TO domain_id`, - `ALTER TABLE IF EXISTS channels RENAME COLUMN owner TO domain_id`, - `ALTER TABLE IF EXISTS configs ADD CONSTRAINT configs_name_domain_id_key UNIQUE (name, domain_id)`, - }, - }, - { - Id: "configs_6", - Up: []string{ - `ALTER TABLE IF EXISTS connections DROP CONSTRAINT IF EXISTS connections_pkey`, - `ALTER TABLE IF EXISTS connections DROP COLUMN IF EXISTS channel_owner`, - `ALTER TABLE IF EXISTS connections DROP COLUMN IF EXISTS config_owner`, - `ALTER TABLE IF EXISTS connections ADD COLUMN IF NOT EXISTS domain_id VARCHAR(256) NOT NULL`, - `ALTER TABLE IF EXISTS connections ADD CONSTRAINT connections_pkey PRIMARY KEY (channel_id, config_id, domain_id)`, - `ALTER TABLE IF EXISTS connections ADD FOREIGN KEY (channel_id, domain_id) REFERENCES channels (magistrala_channel, domain_id) ON DELETE CASCADE ON UPDATE CASCADE`, - `ALTER TABLE IF EXISTS connections ADD FOREIGN KEY (config_id, domain_id) REFERENCES configs (magistrala_client, domain_id) ON DELETE CASCADE ON UPDATE CASCADE`, - }, - }, - { - Id: "configs_7", - Up: []string{ - `ALTER TABLE IF EXISTS configs RENAME COLUMN magistrala_client TO client_id`, - `ALTER TABLE IF EXISTS configs RENAME COLUMN magistrala_secret TO client_secret`, - `CREATE UNIQUE INDEX IF NOT EXISTS configs_client_id_key ON configs (client_id)`, - `CREATE UNIQUE INDEX IF NOT EXISTS configs_client_id_domain_id_key ON configs (client_id, domain_id)`, - `DROP TABLE IF EXISTS connections`, - `DROP TABLE IF EXISTS channels`, - }, - Down: []string{ - `ALTER TABLE IF EXISTS configs RENAME COLUMN client_id TO magistrala_client`, - `ALTER TABLE IF EXISTS configs RENAME COLUMN client_secret TO magistrala_secret`, - }, - }, - { - Id: "configs_8", - Up: []string{ - `DO $$ - BEGIN - IF EXISTS ( - SELECT 1 - FROM information_schema.columns - WHERE table_name = 'configs' AND column_name = 'client_id' - ) AND NOT EXISTS ( - SELECT 1 - FROM information_schema.columns - WHERE table_name = 'configs' AND column_name = 'id' - ) THEN - ALTER TABLE configs RENAME COLUMN client_id TO id; - END IF; - END $$`, - `ALTER TABLE IF EXISTS configs DROP COLUMN IF EXISTS client_secret`, - }, - Down: []string{ - `ALTER TABLE IF EXISTS configs ADD COLUMN IF NOT EXISTS client_secret TEXT`, - `DO $$ - BEGIN - IF EXISTS ( - SELECT 1 - FROM information_schema.columns - WHERE table_name = 'configs' AND column_name = 'id' - ) AND NOT EXISTS ( - SELECT 1 - FROM information_schema.columns - WHERE table_name = 'configs' AND column_name = 'client_id' - ) THEN - ALTER TABLE configs RENAME COLUMN id TO client_id; - END IF; - END $$`, - }, - }, - { - Id: "configs_10", - Up: []string{ - `CREATE TABLE IF NOT EXISTS profiles ( - id VARCHAR(36) PRIMARY KEY, - domain_id VARCHAR(36) NOT NULL, - name VARCHAR(1024) NOT NULL, - description TEXT, - template_format VARCHAR(64) NOT NULL DEFAULT 'go-template', - content_template TEXT, - defaults JSONB, - binding_slots JSONB, - version INT NOT NULL DEFAULT 1, - created_at TIMESTAMP NOT NULL DEFAULT NOW(), - updated_at TIMESTAMP NOT NULL DEFAULT NOW(), - UNIQUE (domain_id, name) - )`, - `CREATE INDEX IF NOT EXISTS idx_profiles_domain_id ON profiles (domain_id)`, - }, - Down: []string{ - `DROP TABLE IF EXISTS profiles`, - }, - }, - { - Id: "configs_11", - Up: []string{ - `ALTER TABLE IF EXISTS configs ADD COLUMN IF NOT EXISTS profile_id VARCHAR(36) REFERENCES profiles (id) ON DELETE SET NULL`, - `ALTER TABLE IF EXISTS configs ADD COLUMN IF NOT EXISTS render_context JSONB`, - }, - Down: []string{ - `ALTER TABLE IF EXISTS configs DROP COLUMN IF EXISTS render_context`, - `ALTER TABLE IF EXISTS configs DROP COLUMN IF EXISTS profile_id`, - }, - }, - { - Id: "configs_12", - Up: []string{ - `CREATE TABLE IF NOT EXISTS bindings ( - config_id TEXT NOT NULL, - slot VARCHAR(256) NOT NULL, - type VARCHAR(64) NOT NULL, - resource_id TEXT NOT NULL, - snapshot JSONB, - secret_snapshot BYTEA, - updated_at TIMESTAMP NOT NULL DEFAULT NOW(), - PRIMARY KEY (config_id, slot) - )`, - `CREATE INDEX IF NOT EXISTS idx_bindings_config_id ON bindings (config_id)`, - }, - Down: []string{ - `DROP TABLE IF EXISTS bindings`, - }, - }, - { - Id: "configs_13", - Up: []string{ - `DO $$ - BEGIN - IF EXISTS ( - SELECT 1 - FROM information_schema.columns - WHERE table_name = 'configs' AND column_name = 'state' - ) AND NOT EXISTS ( - SELECT 1 - FROM information_schema.columns - WHERE table_name = 'configs' AND column_name = 'status' - ) THEN - ALTER TABLE configs RENAME COLUMN state TO status; - END IF; - END $$`, - }, - Down: []string{ - `DO $$ - BEGIN - IF EXISTS ( - SELECT 1 - FROM information_schema.columns - WHERE table_name = 'configs' AND column_name = 'status' - ) AND NOT EXISTS ( - SELECT 1 - FROM information_schema.columns - WHERE table_name = 'configs' AND column_name = 'state' - ) THEN - ALTER TABLE configs RENAME COLUMN status TO state; - END IF; - END $$`, - }, - }, - { - Id: "configs_14", - Up: []string{ - `DO $$ - BEGIN - IF EXISTS ( - SELECT 1 - FROM information_schema.tables - WHERE table_name = 'binding_snapshots' - ) AND NOT EXISTS ( - SELECT 1 - FROM information_schema.tables - WHERE table_name = 'bindings' - ) THEN - ALTER TABLE binding_snapshots RENAME TO bindings; - END IF; - END $$`, - `DO $$ - BEGIN - IF EXISTS ( - SELECT 1 - FROM pg_class - WHERE relname = 'idx_binding_snapshots_config_id' - ) AND NOT EXISTS ( - SELECT 1 - FROM pg_class - WHERE relname = 'idx_bindings_config_id' - ) THEN - ALTER INDEX idx_binding_snapshots_config_id RENAME TO idx_bindings_config_id; - END IF; - END $$`, - }, - Down: []string{ - `DO $$ - BEGIN - IF EXISTS ( - SELECT 1 - FROM information_schema.tables - WHERE table_name = 'bindings' - ) AND NOT EXISTS ( - SELECT 1 - FROM information_schema.tables - WHERE table_name = 'binding_snapshots' - ) THEN - ALTER TABLE bindings RENAME TO binding_snapshots; - END IF; - END $$`, - `DO $$ - BEGIN - IF EXISTS ( - SELECT 1 - FROM pg_class - WHERE relname = 'idx_bindings_config_id' - ) AND NOT EXISTS ( - SELECT 1 - FROM pg_class - WHERE relname = 'idx_binding_snapshots_config_id' - ) THEN - ALTER INDEX idx_bindings_config_id RENAME TO idx_binding_snapshots_config_id; - END IF; - END $$`, - }, - }, - { - Id: "configs_15", - Up: []string{ - `ALTER TABLE IF EXISTS profiles ADD COLUMN IF NOT EXISTS binding_slots JSONB`, - }, - Down: []string{ - `ALTER TABLE IF EXISTS profiles DROP COLUMN IF EXISTS binding_slots`, - }, - }, - { - Id: "configs_16", - Up: []string{ - `ALTER TABLE IF EXISTS profiles RENAME COLUMN template_format TO content_format`, - }, - Down: []string{ - `ALTER TABLE IF EXISTS profiles RENAME COLUMN content_format TO template_format`, - }, - }, - }, - } -} diff --git a/bootstrap/postgres/profiles.go b/bootstrap/postgres/profiles.go deleted file mode 100644 index 68be8def5..000000000 --- a/bootstrap/postgres/profiles.go +++ /dev/null @@ -1,263 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package postgres - -import ( - "context" - "database/sql" - "encoding/json" - "fmt" - "log/slog" - "strings" - "time" - - "github.com/absmach/magistrala/bootstrap" - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - "github.com/absmach/magistrala/pkg/postgres" -) - -var _ bootstrap.ProfileRepository = (*profileRepository)(nil) - -type profileRepository struct { - db postgres.Database - log *slog.Logger -} - -// NewProfileRepository instantiates a PostgreSQL implementation of ProfileRepository. -func NewProfileRepository(db postgres.Database, log *slog.Logger) bootstrap.ProfileRepository { - return &profileRepository{db: db, log: log} -} - -func (pr profileRepository) Save(ctx context.Context, p bootstrap.Profile) (bootstrap.Profile, error) { - q := `INSERT INTO profiles (id, domain_id, name, description, content_format, content_template, defaults, binding_slots, version, created_at, updated_at) - VALUES (:id, :domain_id, :name, :description, :content_format, :content_template, :defaults, :binding_slots, :version, :created_at, :updated_at)` - - now := time.Now().UTC() - p.CreatedAt = now - p.UpdatedAt = now - - dbp, err := toDBProfile(p) - if err != nil { - return bootstrap.Profile{}, errors.Wrap(repoerr.ErrCreateEntity, err) - } - - if _, err = pr.db.NamedExecContext(ctx, q, dbp); err != nil { - return bootstrap.Profile{}, postgres.HandleError(repoerr.ErrCreateEntity, err) - } - - return p, nil -} - -func (pr profileRepository) RetrieveByID(ctx context.Context, domainID, id string) (bootstrap.Profile, error) { - q := `SELECT id, domain_id, name, description, content_format, content_template, defaults, binding_slots, version, created_at, updated_at - FROM profiles WHERE id = :id AND domain_id = :domain_id` - - rows, err := pr.db.NamedQueryContext(ctx, q, dbProfile{ID: id, DomainID: domainID}) - if err != nil { - return bootstrap.Profile{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - defer rows.Close() - - if !rows.Next() { - return bootstrap.Profile{}, repoerr.ErrNotFound - } - var dbp dbProfile - if err := rows.StructScan(&dbp); err != nil { - return bootstrap.Profile{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - - return toProfile(dbp) -} - -func (pr profileRepository) RetrieveAll(ctx context.Context, domainID string, offset, limit uint64, name string) (bootstrap.ProfilesPage, error) { - dbPage := dbProfilesPage{DomainID: domainID, Offset: offset, Limit: limit, Name: name} - pageQuery := profilesPageQuery(dbPage) - q := fmt.Sprintf(`SELECT id, domain_id, name, description, content_format, content_template, defaults, binding_slots, version, created_at, updated_at - FROM profiles %s`, pageQuery) - q = applyProfilesOrdering(q) - q = fmt.Sprintf(`%s LIMIT :limit OFFSET :offset`, q) - - rows, err := pr.db.NamedQueryContext(ctx, q, dbPage) - if err != nil { - return bootstrap.ProfilesPage{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - defer rows.Close() - - var profiles []bootstrap.Profile - for rows.Next() { - var dbp dbProfile - if err := rows.StructScan(&dbp); err != nil { - pr.log.Error(fmt.Sprintf("failed to scan profile row: %s", err)) - return bootstrap.ProfilesPage{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - p, err := toProfile(dbp) - if err != nil { - return bootstrap.ProfilesPage{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - profiles = append(profiles, p) - } - - cq := fmt.Sprintf(`SELECT COUNT(*) FROM profiles %s`, pageQuery) - total, err := postgres.Total(ctx, pr.db, cq, dbPage) - if err != nil { - return bootstrap.ProfilesPage{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - - return bootstrap.ProfilesPage{ - Total: total, - Offset: offset, - Limit: limit, - Profiles: profiles, - }, nil -} - -type dbProfilesPage struct { - DomainID string `db:"domain_id"` - Offset uint64 `db:"offset"` - Limit uint64 `db:"limit"` - Name string `db:"name"` -} - -func profilesPageQuery(pm dbProfilesPage) string { - var query []string - query = append(query, "domain_id = :domain_id") - if pm.Name != "" { - query = append(query, "name ILIKE '%' || :name || '%'") - } - return fmt.Sprintf("WHERE %s", strings.Join(query, " AND ")) -} - -func applyProfilesOrdering(q string) string { - return fmt.Sprintf("%s ORDER BY created_at DESC", q) -} - -func (pr profileRepository) Update(ctx context.Context, p bootstrap.Profile) (bootstrap.Profile, error) { - var query []string - var upq string - if p.Name != "" { - query = append(query, "name = :name,") - } - if p.Description != "" { - query = append(query, "description = :description,") - } - if p.ContentFormat != "" { - query = append(query, "content_format = :content_format,") - } - if p.ContentTemplate != "" { - query = append(query, "content_template = :content_template,") - } - if p.Defaults != nil { - query = append(query, "defaults = :defaults,") - } - if p.BindingSlots != nil { - query = append(query, "binding_slots = :binding_slots,") - } - if len(query) > 0 { - upq = strings.Join(query, " ") - } - - q := fmt.Sprintf(`UPDATE profiles SET %s version = version + 1, updated_at = :updated_at - WHERE id = :id AND domain_id = :domain_id - RETURNING id, domain_id, name, description, content_format, content_template, defaults, binding_slots, version, created_at, updated_at`, - upq) - - p.UpdatedAt = time.Now().UTC() - dbp, err := toDBProfile(p) - if err != nil { - return bootstrap.Profile{}, errors.Wrap(repoerr.ErrUpdateEntity, err) - } - - rows, err := pr.db.NamedQueryContext(ctx, q, dbp) - if err != nil { - return bootstrap.Profile{}, postgres.HandleError(repoerr.ErrUpdateEntity, err) - } - defer rows.Close() - - if !rows.Next() { - return bootstrap.Profile{}, repoerr.ErrNotFound - } - var updated dbProfile - if err := rows.StructScan(&updated); err != nil { - return bootstrap.Profile{}, errors.Wrap(repoerr.ErrUpdateEntity, err) - } - - return toProfile(updated) -} - -func (pr profileRepository) Delete(ctx context.Context, domainID, id string) error { - q := `DELETE FROM profiles WHERE id = :id AND domain_id = :domain_id` - if _, err := pr.db.NamedExecContext(ctx, q, dbProfile{ID: id, DomainID: domainID}); err != nil { - return errors.Wrap(repoerr.ErrRemoveEntity, err) - } - return nil -} - -// dbProfile is the database representation of a Profile. -type dbProfile struct { - ID string `db:"id"` - DomainID string `db:"domain_id"` - Name string `db:"name"` - Description sql.NullString `db:"description"` - ContentFormat string `db:"content_format"` - ContentTemplate sql.NullString `db:"content_template"` - Defaults []byte `db:"defaults"` - BindingSlots []byte `db:"binding_slots"` - Version int `db:"version"` - CreatedAt time.Time `db:"created_at"` - UpdatedAt time.Time `db:"updated_at"` -} - -func toDBProfile(p bootstrap.Profile) (dbProfile, error) { - defaults, err := json.Marshal(p.Defaults) - if err != nil { - return dbProfile{}, err - } - bindingSlots, err := json.Marshal(p.BindingSlots) - if err != nil { - return dbProfile{}, err - } - return dbProfile{ - ID: p.ID, - DomainID: p.DomainID, - Name: p.Name, - Description: nullString(p.Description), - ContentFormat: string(p.ContentFormat), - ContentTemplate: nullString(p.ContentTemplate), - Defaults: defaults, - BindingSlots: bindingSlots, - Version: p.Version, - CreatedAt: p.CreatedAt, - UpdatedAt: p.UpdatedAt, - }, nil -} - -func toProfile(dbp dbProfile) (bootstrap.Profile, error) { - p := bootstrap.Profile{ - ID: dbp.ID, - DomainID: dbp.DomainID, - Name: dbp.Name, - ContentFormat: bootstrap.ContentFormat(dbp.ContentFormat), - Version: dbp.Version, - CreatedAt: dbp.CreatedAt, - UpdatedAt: dbp.UpdatedAt, - } - if dbp.Description.Valid { - p.Description = dbp.Description.String - } - if dbp.ContentTemplate.Valid { - p.ContentTemplate = dbp.ContentTemplate.String - } - if len(dbp.Defaults) > 0 && string(dbp.Defaults) != jsonNull { - if err := json.Unmarshal(dbp.Defaults, &p.Defaults); err != nil { - return bootstrap.Profile{}, err - } - } - if len(dbp.BindingSlots) > 0 && string(dbp.BindingSlots) != jsonNull { - if err := json.Unmarshal(dbp.BindingSlots, &p.BindingSlots); err != nil { - return bootstrap.Profile{}, err - } - } - return p, nil -} diff --git a/bootstrap/postgres/setup_test.go b/bootstrap/postgres/setup_test.go deleted file mode 100644 index 9c7d3a983..000000000 --- a/bootstrap/postgres/setup_test.go +++ /dev/null @@ -1,88 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package postgres_test - -import ( - "fmt" - "log" - "os" - "testing" - - "github.com/absmach/magistrala/bootstrap/postgres" - mglog "github.com/absmach/magistrala/logger" - pgclient "github.com/absmach/magistrala/pkg/postgres" - "github.com/jmoiron/sqlx" - "github.com/ory/dockertest/v3" - "github.com/ory/dockertest/v3/docker" -) - -var ( - testLog, _ = mglog.New(os.Stdout, "info") - db *sqlx.DB -) - -func TestMain(m *testing.M) { - pool, err := dockertest.NewPool("") - if err != nil { - testLog.Error(fmt.Sprintf("Could not connect to docker: %s", err)) - } - - container, err := pool.RunWithOptions(&dockertest.RunOptions{ - Repository: "postgres", - Tag: "16.2-alpine", - Env: []string{ - "POSTGRES_USER=test", - "POSTGRES_PASSWORD=test", - "POSTGRES_DB=test", - "listen_addresses = '*'", - }, - }, func(config *docker.HostConfig) { - config.AutoRemove = true - config.RestartPolicy = docker.RestartPolicy{Name: "no"} - }) - if err != nil { - log.Fatalf("Could not start container: %s", err) - } - - port := container.GetPort("5432/tcp") - - if err := pool.Retry(func() error { - url := fmt.Sprintf("host=localhost port=%s user=test dbname=test password=test sslmode=disable", port) - db, err = sqlx.Open("pgx", url) - if err != nil { - return err - } - return db.Ping() - }); err != nil { - testLog.Error(fmt.Sprintf("Could not connect to docker: %s", err)) - } - - dbConfig := pgclient.Config{ - Host: "localhost", - Port: port, - User: "test", - Pass: "test", - Name: "test", - SSLMode: "disable", - SSLCert: "", - SSLKey: "", - SSLRootCert: "", - } - - migration := postgres.Migration() - - if db, err = pgclient.Setup(dbConfig, *migration); err != nil { - testLog.Error(fmt.Sprintf("Could not setup test DB connection: %s", err)) - } - - code := m.Run() - - // Defers will not be run when using os.Exit - db.Close() - if err := pool.Purge(container); err != nil { - testLog.Error(fmt.Sprintf("Could not purge container: %s", err)) - } - - os.Exit(code) -} diff --git a/bootstrap/profiles.go b/bootstrap/profiles.go deleted file mode 100644 index c4b2356c8..000000000 --- a/bootstrap/profiles.go +++ /dev/null @@ -1,69 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package bootstrap - -import ( - "context" - "time" -) - -// ContentFormat enumerates the supported output formats for rendered profile templates. -type ContentFormat string - -const ( - ContentFormatGoTemplate ContentFormat = "go-template" - ContentFormatRaw ContentFormat = "raw" - ContentFormatJSON ContentFormat = "json" - ContentFormatYAML ContentFormat = "yaml" - ContentFormatTOML ContentFormat = "toml" -) - -// Profile is a user-managed device configuration template. -type Profile struct { - ID string `json:"id"` - DomainID string `json:"domain_id,omitempty"` - Name string `json:"name"` - Description string `json:"description,omitempty"` - ContentFormat ContentFormat `json:"content_format"` - ContentTemplate string `json:"content_template,omitempty"` - Defaults map[string]any `json:"defaults,omitempty"` - BindingSlots []BindingSlot `json:"binding_slots,omitempty"` - Version int `json:"version,omitempty"` - CreatedAt time.Time `json:"created_at,omitempty"` - UpdatedAt time.Time `json:"updated_at,omitempty"` -} - -// BindingSlot declares a named resource placeholder that a profile template can use. -type BindingSlot struct { - Name string `json:"name"` - Type string `json:"type"` - Required bool `json:"required"` - Fields []string `json:"fields,omitempty"` -} - -// ProfilesPage contains pagination metadata and a slice of Profiles. -type ProfilesPage struct { - Total uint64 `json:"total"` - Offset uint64 `json:"offset"` - Limit uint64 `json:"limit"` - Profiles []Profile `json:"profiles"` -} - -// ProfileRepository specifies the persistence API for Profiles. -type ProfileRepository interface { - // Save persists a new Profile and returns it with server-assigned fields set. - Save(ctx context.Context, p Profile) (Profile, error) - - // RetrieveByID returns the Profile with the given ID inside the given domain. - RetrieveByID(ctx context.Context, domainID, id string) (Profile, error) - - // RetrieveAll returns a page of Profiles belonging to the given domain, optionally filtered by name. - RetrieveAll(ctx context.Context, domainID string, offset, limit uint64, name string) (ProfilesPage, error) - - // Update updates editable fields of the given Profile and returns the updated Profile. - Update(ctx context.Context, p Profile) (Profile, error) - - // Delete removes the Profile with the given ID from the given domain. - Delete(ctx context.Context, domainID, id string) error -} diff --git a/bootstrap/reader.go b/bootstrap/reader.go deleted file mode 100644 index 1fd5a237b..000000000 --- a/bootstrap/reader.go +++ /dev/null @@ -1,80 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package bootstrap - -import ( - "crypto/aes" - "crypto/cipher" - "crypto/rand" - "encoding/json" - "io" - "net/http" -) - -// bootstrapRes represent Magistrala Response to the Bootstrap request. -// This is used as a response from ConfigReader and can easily be -// replaced with any other response format. -type bootstrapRes struct { - ID string `json:"id,omitempty"` - Content string `json:"content,omitempty"` - ClientCert string `json:"client_cert,omitempty"` - ClientKey string `json:"client_key,omitempty"` - CACert string `json:"ca_cert,omitempty"` -} - -func (res bootstrapRes) Code() int { - return http.StatusOK -} - -func (res bootstrapRes) Headers() map[string]string { - return map[string]string{} -} - -func (res bootstrapRes) Empty() bool { - return false -} - -type reader struct { - encKey []byte -} - -// NewConfigReader return new reader which is used to generate response -// from the config. -func NewConfigReader(encKey []byte) ConfigReader { - return reader{encKey: encKey} -} - -func (r reader) ReadConfig(cfg Config, secure bool) (any, error) { - res := bootstrapRes{ - ID: cfg.ID, - Content: cfg.Content, - ClientCert: cfg.ClientCert, - ClientKey: cfg.ClientKey, - CACert: cfg.CACert, - } - if secure { - b, err := json.Marshal(res) - if err != nil { - return nil, err - } - return r.encrypt(b) - } - - return res, nil -} - -func (r reader) encrypt(in []byte) ([]byte, error) { - block, err := aes.NewCipher(r.encKey) - if err != nil { - return nil, err - } - ciphertext := make([]byte, aes.BlockSize+len(in)) - iv := ciphertext[:aes.BlockSize] - if _, err := io.ReadFull(rand.Reader, iv); err != nil { - return nil, err - } - stream := cipher.NewCFBEncrypter(block, iv) - stream.XORKeyStream(ciphertext[aes.BlockSize:], in) - return ciphertext, nil -} diff --git a/bootstrap/reader_test.go b/bootstrap/reader_test.go deleted file mode 100644 index e61c6122f..000000000 --- a/bootstrap/reader_test.go +++ /dev/null @@ -1,102 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package bootstrap_test - -import ( - "crypto/aes" - "crypto/cipher" - "encoding/json" - "fmt" - "net/http" - "testing" - - "github.com/absmach/magistrala" - "github.com/absmach/magistrala/bootstrap" - "github.com/absmach/magistrala/pkg/errors" - "github.com/stretchr/testify/assert" -) - -type readResp struct { - ID string `json:"id"` - Content string `json:"content,omitempty"` - ClientCert string `json:"client_cert,omitempty"` - ClientKey string `json:"client_key,omitempty"` - CACert string `json:"ca_cert,omitempty"` -} - -func dec(in []byte) ([]byte, error) { - block, err := aes.NewCipher(encKey) - if err != nil { - return nil, err - } - if len(in) < aes.BlockSize { - return nil, errors.ErrMalformedEntity - } - iv := in[:aes.BlockSize] - in = in[aes.BlockSize:] - stream := cipher.NewCFBDecrypter(block, iv) - stream.XORKeyStream(in, in) - return in, nil -} - -func TestReadConfig(t *testing.T) { - cfg := bootstrap.Config{ - ID: "smq_id", - ClientCert: "client_cert", - ClientKey: "client_key", - CACert: "ca_cert", - Content: "content", - } - ret := readResp{ - ID: "smq_id", - Content: "content", - ClientCert: "client_cert", - ClientKey: "client_key", - CACert: "ca_cert", - } - - bin, err := json.Marshal(ret) - assert.Nil(t, err, fmt.Sprintf("Marshalling expected to succeed: %s.\n", err)) - - reader := bootstrap.NewConfigReader(encKey) - cases := []struct { - desc string - config bootstrap.Config - enc []byte - secret bool - err error - }{ - { - desc: "read a config", - config: cfg, - enc: bin, - secret: false, - }, - { - desc: "read encrypted config", - config: cfg, - enc: bin, - secret: true, - }, - } - - for _, tc := range cases { - res, err := reader.ReadConfig(tc.config, tc.secret) - assert.Nil(t, err, fmt.Sprintf("Reading config to succeed: %s.\n", err)) - - if tc.secret { - d, err := dec(res.([]byte)) - assert.Nil(t, err, fmt.Sprintf("Decrypting expected to succeed: %s.\n", err)) - assert.Equal(t, tc.enc, d, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.enc, d)) - continue - } - b, err := json.Marshal(res) - assert.Nil(t, err, fmt.Sprintf("Marshalling expected to succeed: %s.\n", err)) - assert.Equal(t, tc.enc, b, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.enc, b)) - resp, ok := res.(magistrala.Response) - assert.True(t, ok, "If not encrypted, reader should return response.") - assert.False(t, resp.Empty(), fmt.Sprintf("Response should not be empty %s.", err)) - assert.Equal(t, http.StatusOK, resp.Code(), "Default config response code should be 200.") - } -} diff --git a/bootstrap/renderer.go b/bootstrap/renderer.go deleted file mode 100644 index 5689885bb..000000000 --- a/bootstrap/renderer.go +++ /dev/null @@ -1,174 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package bootstrap - -import ( - "bytes" - "encoding/json" - "fmt" - "text/template" - - "github.com/absmach/magistrala/pkg/errors" - "github.com/pelletier/go-toml/v2" - "gopkg.in/yaml.v3" -) - -// Renderer renders a Profile's content template into a concrete device -// configuration. All input data must already be stored in Bootstrap — no -// external service calls are allowed inside Render. -type Renderer interface { - Render(profile Profile, enrollment Config, bindings []BindingSnapshot) ([]byte, error) -} - -// ErrRenderFailed is returned when template execution or output validation fails. -var ErrRenderFailed = errors.New("failed to render profile template") - -type renderer struct{} - -// NewRenderer returns the default Renderer implementation using Go text/template. -func NewRenderer() Renderer { - return renderer{} -} - -func (r renderer) Render(profile Profile, enrollment Config, bindings []BindingSnapshot) ([]byte, error) { - rctx := buildRenderContext(profile, enrollment, bindings) - - switch profile.ContentFormat { - case ContentFormatRaw: - return []byte(profile.ContentTemplate), nil - case ContentFormatGoTemplate, ContentFormatJSON, ContentFormatYAML, ContentFormatTOML, "": - return r.renderTemplate(profile, rctx) - default: - return nil, fmt.Errorf("%w: unsupported template format %q", ErrRenderFailed, profile.ContentFormat) - } -} - -func (r renderer) renderTemplate(profile Profile, rctx RenderContext) ([]byte, error) { - t, err := template.New("bootstrap"). - Option("missingkey=error"). - Funcs(allowlistedFuncs()). - Parse(profile.ContentTemplate) - if err != nil { - return nil, fmt.Errorf("%w: %w", ErrRenderFailed, err) - } - - var buf bytes.Buffer - if err := t.Execute(&buf, rctx); err != nil { - return nil, fmt.Errorf("%w: %w", ErrRenderFailed, err) - } - - return convertOutput(buf.Bytes(), profile.ContentFormat) -} - -// convertOutput parses the rendered bytes as any structured format (JSON, YAML, -// or TOML) and re-marshals them into the declared target format. For go-template -// or empty format the raw bytes are returned unchanged. -func convertOutput(out []byte, format ContentFormat) ([]byte, error) { - switch format { - case ContentFormatGoTemplate, "": - return out, nil - case ContentFormatJSON, ContentFormatYAML, ContentFormatTOML: - var v any - if err := parseStructured(out, &v); err != nil { - return nil, fmt.Errorf("%w: %w", ErrRenderFailed, err) - } - result, err := marshalAs(v, format) - if err != nil { - return nil, fmt.Errorf("%w: %w", ErrRenderFailed, err) - } - return result, nil - default: - return nil, fmt.Errorf("%w: unsupported format %q", ErrRenderFailed, format) - } -} - -// parseStructured tries JSON, then YAML, then TOML and unmarshals into v. -func parseStructured(out []byte, v any) error { - if err := json.Unmarshal(out, v); err == nil { - return nil - } - if err := yaml.Unmarshal(out, v); err == nil { - return nil - } - if err := toml.Unmarshal(out, v); err == nil { - return nil - } - return fmt.Errorf("template output is not valid JSON, YAML, or TOML") -} - -// marshalAs re-marshals v into the requested format. -func marshalAs(v any, format ContentFormat) ([]byte, error) { - switch format { - case ContentFormatJSON: - return json.MarshalIndent(v, "", " ") - case ContentFormatYAML: - return yaml.Marshal(v) - case ContentFormatTOML: - var buf bytes.Buffer - if err := toml.NewEncoder(&buf).Encode(v); err != nil { - return nil, err - } - return buf.Bytes(), nil - default: - return nil, fmt.Errorf("unsupported format %q", format) - } -} - -// buildRenderContext constructs the typed RenderContext from stored data. -// No external calls are made here. -func buildRenderContext(profile Profile, enrollment Config, bindings []BindingSnapshot) RenderContext { - vars := make(map[string]any) - for k, v := range profile.Defaults { - vars[k] = v - } - for k, v := range enrollment.RenderContext { - vars[k] = v - } - - bctx := make(map[string]BindingContext, len(bindings)) - for _, b := range bindings { - bctx[b.Slot] = BindingContext{ - Type: b.Type, - ID: b.ResourceID, - Snapshot: b.Snapshot, - Secret: b.SecretSnapshot, - } - } - - return RenderContext{ - Device: DeviceContext{ - ID: enrollment.ID, - ExternalID: enrollment.ExternalID, - DomainID: enrollment.DomainID, - }, - Vars: vars, - Bindings: bctx, - } -} - -// allowlistedFuncs returns the safe set of template helper functions. -// No function in this map may call an external service or perform I/O. -func allowlistedFuncs() template.FuncMap { - return template.FuncMap{ - "toJSON": func(v any) (string, error) { - b, err := json.Marshal(v) - if err != nil { - return "", err - } - return string(b), nil - }, - "default": func(def, val any) any { - if val == nil || val == "" { - return def - } - return val - }, - "required": func(key string, val any) (any, error) { - if val == nil || val == "" { - return nil, fmt.Errorf("required value %q is missing", key) - } - return val, nil - }, - } -} diff --git a/bootstrap/renderer_test.go b/bootstrap/renderer_test.go deleted file mode 100644 index 72b023352..000000000 --- a/bootstrap/renderer_test.go +++ /dev/null @@ -1,88 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package bootstrap_test - -import ( - "fmt" - "testing" - - "github.com/absmach/magistrala/bootstrap" - "github.com/absmach/magistrala/pkg/errors" - "github.com/stretchr/testify/assert" -) - -func TestRendererStructuredOutputValidation(t *testing.T) { - renderer := bootstrap.NewRenderer() - - cases := []struct { - desc string - format bootstrap.ContentFormat - template string - err error - }{ - { - desc: "valid JSON output", - format: bootstrap.ContentFormatJSON, - template: `{"device_id":"{{ .Device.ID }}"}`, - }, - { - desc: "invalid output for JSON format", - format: bootstrap.ContentFormatJSON, - template: `[unclosed bracket`, - err: bootstrap.ErrRenderFailed, - }, - { - desc: "valid YAML output", - format: bootstrap.ContentFormatYAML, - template: "device_id: {{ .Device.ID }}", - }, - { - desc: "invalid output for YAML format", - format: bootstrap.ContentFormatYAML, - template: "[unclosed bracket", - err: bootstrap.ErrRenderFailed, - }, - { - desc: "valid TOML output", - format: bootstrap.ContentFormatTOML, - template: `[device] - device_id = "{{ .Device.ID }}"`, - }, - { - desc: "invalid output for TOML format", - format: bootstrap.ContentFormatTOML, - template: `[unclosed bracket`, - err: bootstrap.ErrRenderFailed, - }, - { - desc: "JSON template auto-converted to TOML", - format: bootstrap.ContentFormatTOML, - template: `{"device_id":"{{ .Device.ID }}"}`, - }, - { - desc: "TOML template auto-converted to JSON", - format: bootstrap.ContentFormatJSON, - template: `device_id = "{{ .Device.ID }}"`, - }, - { - desc: "YAML template auto-converted to TOML", - format: bootstrap.ContentFormatTOML, - template: "device_id: {{ .Device.ID }}", - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - _, err := renderer.Render( - bootstrap.Profile{ - ContentFormat: tc.format, - ContentTemplate: tc.template, - }, - bootstrap.Config{ID: "config-id"}, - nil, - ) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %v got %v", tc.desc, tc.err, err)) - }) - } -} diff --git a/bootstrap/sdk_resolver.go b/bootstrap/sdk_resolver.go deleted file mode 100644 index 35d8cd903..000000000 --- a/bootstrap/sdk_resolver.go +++ /dev/null @@ -1,115 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package bootstrap - -import ( - "context" - "fmt" - "time" - - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - mgsdk "github.com/absmach/magistrala/pkg/sdk" -) - -var _ BindingResolver = (*sdkResolver)(nil) - -type sdkResolver struct { - sdk mgsdk.SDK -} - -// NewSDKResolver returns a BindingResolver that validates resources against -// the Magistrala clients and channels services using the SDK. This resolver -// is called only at binding time; the render path must never call it. -func NewSDKResolver(sdk mgsdk.SDK) BindingResolver { - return &sdkResolver{sdk: sdk} -} - -func (r *sdkResolver) Resolve(ctx context.Context, req ResolveRequest) ([]BindingSnapshot, error) { - var snapshots []BindingSnapshot - - for _, br := range req.Requested { - snap, err := r.resolveOne(ctx, req.Enrollment.DomainID, req.Token, br) - if err != nil { - return nil, err - } - snapshots = append(snapshots, snap) - } - - return snapshots, nil -} - -func (r *sdkResolver) resolveOne(ctx context.Context, domainID, token string, br BindingRequest) (BindingSnapshot, error) { - switch br.Type { - case "client": - return r.resolveClient(ctx, domainID, token, br) - case "channel": - return r.resolveChannel(ctx, domainID, token, br) - default: - return BindingSnapshot{}, fmt.Errorf("unsupported binding type %q", br.Type) - } -} - -func (r *sdkResolver) resolveClient(ctx context.Context, domainID, token string, br BindingRequest) (BindingSnapshot, error) { - client, sdkErr := r.sdk.Client(ctx, br.ResourceID, domainID, token) - if sdkErr != nil { - return BindingSnapshot{}, errors.Wrap(svcerr.ErrNotFound, - fmt.Errorf("client %q not found: %s", br.ResourceID, sdkErr)) - } - - snapshot := map[string]any{ - "id": client.ID, - "name": client.Name, - } - if client.Credentials.Identity != "" { - snapshot["identity"] = client.Credentials.Identity - } - if client.DomainID != "" { - snapshot["domain_id"] = client.DomainID - } - - secret := map[string]any{} - if client.Credentials.Secret != "" { - secret["secret"] = client.Credentials.Secret - } - - return BindingSnapshot{ - Slot: br.Slot, - Type: br.Type, - ResourceID: br.ResourceID, - Snapshot: snapshot, - SecretSnapshot: secret, - UpdatedAt: time.Now().UTC(), - }, nil -} - -func (r *sdkResolver) resolveChannel(ctx context.Context, domainID, token string, br BindingRequest) (BindingSnapshot, error) { - channel, sdkErr := r.sdk.Channel(ctx, br.ResourceID, domainID, token) - if sdkErr != nil { - return BindingSnapshot{}, errors.Wrap(svcerr.ErrNotFound, - fmt.Errorf("channel %q not found: %s", br.ResourceID, sdkErr)) - } - - snapshot := map[string]any{ - "id": channel.ID, - "name": channel.Name, - } - if channel.Route != "" { - snapshot["topic"] = channel.Route - } - if channel.DomainID != "" { - snapshot["domain_id"] = channel.DomainID - } - if channel.Metadata != nil { - snapshot["metadata"] = channel.Metadata - } - - return BindingSnapshot{ - Slot: br.Slot, - Type: br.Type, - ResourceID: br.ResourceID, - Snapshot: snapshot, - UpdatedAt: time.Now().UTC(), - }, nil -} diff --git a/bootstrap/secret_snapshots.go b/bootstrap/secret_snapshots.go deleted file mode 100644 index 01d57a8af..000000000 --- a/bootstrap/secret_snapshots.go +++ /dev/null @@ -1,100 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package bootstrap - -import ( - "crypto/aes" - "crypto/cipher" - "crypto/rand" - "encoding/hex" - "encoding/json" -) - -const secretSnapshotCiphertextKey = "ciphertext" - -func (bs bootstrapService) encryptSecretSnapshots(bindings []BindingSnapshot) ([]BindingSnapshot, error) { - encrypted := make([]BindingSnapshot, len(bindings)) - for i, binding := range bindings { - encrypted[i] = binding - if len(binding.SecretSnapshot) == 0 { - continue - } - secret, err := json.Marshal(binding.SecretSnapshot) - if err != nil { - return nil, err - } - ciphertext, err := bs.encrypt(secret) - if err != nil { - return nil, err - } - encrypted[i].SecretSnapshot = map[string]any{ - secretSnapshotCiphertextKey: ciphertext, - } - } - return encrypted, nil -} - -func (bs bootstrapService) decryptSecretSnapshots(bindings []BindingSnapshot) ([]BindingSnapshot, error) { - decrypted := make([]BindingSnapshot, len(bindings)) - for i, binding := range bindings { - decrypted[i] = binding - ciphertext, ok := binding.SecretSnapshot[secretSnapshotCiphertextKey].(string) - if !ok { - continue - } - plain, err := bs.decrypt(ciphertext) - if err != nil { - return nil, err - } - var secret map[string]any - if err := json.Unmarshal(plain, &secret); err != nil { - return nil, err - } - decrypted[i].SecretSnapshot = secret - } - return decrypted, nil -} - -func hideSecretSnapshots(bindings []BindingSnapshot) []BindingSnapshot { - hidden := make([]BindingSnapshot, len(bindings)) - for i, binding := range bindings { - hidden[i] = binding - hidden[i].SecretSnapshot = nil - } - return hidden -} - -func (bs bootstrapService) encrypt(plain []byte) (string, error) { - block, err := aes.NewCipher(bs.encKey) - if err != nil { - return "", err - } - ciphertext := make([]byte, aes.BlockSize+len(plain)) - iv := ciphertext[:aes.BlockSize] - if _, err := rand.Read(iv); err != nil { - return "", err - } - stream := cipher.NewCFBEncrypter(block, iv) - stream.XORKeyStream(ciphertext[aes.BlockSize:], plain) - return hex.EncodeToString(ciphertext), nil -} - -func (bs bootstrapService) decrypt(in string) ([]byte, error) { - ciphertext, err := hex.DecodeString(in) - if err != nil { - return nil, err - } - block, err := aes.NewCipher(bs.encKey) - if err != nil { - return nil, err - } - if len(ciphertext) < aes.BlockSize { - return nil, ErrExternalKeySecure - } - iv := ciphertext[:aes.BlockSize] - ciphertext = ciphertext[aes.BlockSize:] - stream := cipher.NewCFBDecrypter(block, iv) - stream.XORKeyStream(ciphertext, ciphertext) - return ciphertext, nil -} diff --git a/bootstrap/service.go b/bootstrap/service.go deleted file mode 100644 index a66f2d164..000000000 --- a/bootstrap/service.go +++ /dev/null @@ -1,532 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package bootstrap - -import ( - "context" - "crypto/aes" - "crypto/cipher" - "encoding/hex" - - "github.com/absmach/magistrala" - smqauthn "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - mgsdk "github.com/absmach/magistrala/pkg/sdk" -) - -var ( - // ErrExternalKey indicates a non-existent bootstrap configuration for given external key. - ErrExternalKey = errors.NewAuthZError("failed to get bootstrap configuration for given external key") - - // ErrExternalKeySecure indicates error in getting bootstrap configuration for given encrypted external key. - ErrExternalKeySecure = errors.NewAuthZError("failed to get bootstrap configuration for given encrypted external key") - - // ErrBootstrap indicates error in getting bootstrap configuration. - ErrBootstrap = errors.New("failed to read bootstrap configuration") - - // ErrAddBootstrap indicates error in adding bootstrap configuration. - ErrAddBootstrap = errors.NewServiceError("failed to add bootstrap configuration") - - // ErrBootstrapStatus indicates an invalid bootstrap status. - ErrBootstrapStatus = errors.NewRequestError("invalid bootstrap status") - - errRemoveBootstrap = errors.New("failed to remove bootstrap configuration") - errEnableConfig = errors.New("failed to enable bootstrap configuration") - errDisableConfig = errors.New("failed to disable bootstrap configuration") - errUpdateCert = errors.New("failed to update cert") - - errCreateProfile = errors.New("failed to create profile") - errViewProfile = errors.New("failed to view profile") - errUpdateProfile = errors.New("failed to update profile") - errDeleteProfile = errors.New("failed to delete profile") - errListProfiles = errors.New("failed to list profiles") - errAssignProfile = errors.New("failed to assign profile to enrollment") - errBindResources = errors.New("failed to bind resources") - errListBindings = errors.New("failed to list bindings") - errRefreshBinding = errors.New("failed to refresh bindings") - errRenderBootstrap = errors.New("failed to render bootstrap configuration") -) - -var _ Service = (*bootstrapService)(nil) - -// Service specifies an API that must be fulfilled by the domain service -// implementation, and all of its decorators (e.g. logging & metrics). -type Service interface { - // Add adds new Client Config to the user identified by the provided token. - Add(ctx context.Context, session smqauthn.Session, token string, cfg Config) (Config, error) - - // View returns Client Config with given ID belonging to the user identified by the given token. - View(ctx context.Context, session smqauthn.Session, id string) (Config, error) - - // Update updates editable fields of the provided Config. - Update(ctx context.Context, session smqauthn.Session, cfg Config) error - - // UpdateCert updates an existing Config certificate and token. - // A non-nil error is returned to indicate operation failure. - UpdateCert(ctx context.Context, session smqauthn.Session, id, clientCert, clientKey, caCert string) (Config, error) - - // List returns subset of Configs with given search params that belong to the - // user identified by the given token. - List(ctx context.Context, session smqauthn.Session, filter Filter, offset, limit uint64) (ConfigsPage, error) - - // Remove removes Config with specified token that belongs to the user identified by the given token. - Remove(ctx context.Context, session smqauthn.Session, id string) error - - // Bootstrap returns Config to the Client with provided external ID using external key. - Bootstrap(ctx context.Context, externalKey, externalID string, secure bool) (Config, error) - - // EnableConfig enables the Config so its device can successfully bootstrap. - EnableConfig(ctx context.Context, session smqauthn.Session, id string) (Config, error) - - // DisableConfig disables the Config, preventing its device from bootstrapping. - DisableConfig(ctx context.Context, session smqauthn.Session, id string) (Config, error) - - // CreateProfile persists a new device Profile. - CreateProfile(ctx context.Context, session smqauthn.Session, p Profile) (Profile, error) - - // ViewProfile returns the Profile with the given ID. - ViewProfile(ctx context.Context, session smqauthn.Session, profileID string) (Profile, error) - - // UpdateProfile updates editable fields of the given Profile and returns the updated Profile. - UpdateProfile(ctx context.Context, session smqauthn.Session, p Profile) (Profile, error) - - // ListProfiles returns a page of Profiles belonging to the domain. - ListProfiles(ctx context.Context, session smqauthn.Session, offset, limit uint64, name string) (ProfilesPage, error) - - // DeleteProfile removes the Profile with the given ID. - DeleteProfile(ctx context.Context, session smqauthn.Session, profileID string) error - - // AssignProfile sets the ProfileID on an existing enrollment (Config). - AssignProfile(ctx context.Context, session smqauthn.Session, configID, profileID string) error - - // BindResources resolves the requested bindings through their owning services, - // stores snapshots, and marks the enrollment renderable when all required slots - // are satisfied. - BindResources(ctx context.Context, session smqauthn.Session, token, configID string, bindings []BindingRequest) error - - // ListBindings returns all stored binding snapshots for an enrollment. - ListBindings(ctx context.Context, session smqauthn.Session, configID string) ([]BindingSnapshot, error) - - // RefreshBindings re-resolves all existing bindings for an enrollment and - // updates the stored snapshots. - RefreshBindings(ctx context.Context, session smqauthn.Session, token, configID string) error -} - -// ConfigReader is used to parse Config into format which will be encoded -// as a JSON and consumed from the client side. The purpose of this interface -// is to provide convenient way to generate custom configuration response -// based on the specific Config which will be consumed by the client. -type ConfigReader interface { - ReadConfig(Config, bool) (any, error) -} - -type bootstrapService struct { - configs ConfigRepository - profiles ProfileRepository - bindings BindingStore - resolver BindingResolver - renderer Renderer - hasher Hasher - sdk mgsdk.SDK - encKey []byte - idProvider magistrala.IDProvider -} - -// New returns new Bootstrap service. -func New( - configs ConfigRepository, - profiles ProfileRepository, - bindings BindingStore, - resolver BindingResolver, - renderer Renderer, - sdk mgsdk.SDK, - hasher Hasher, - encKey []byte, - idp magistrala.IDProvider, -) Service { - return &bootstrapService{ - configs: configs, - profiles: profiles, - bindings: bindings, - resolver: resolver, - renderer: renderer, - hasher: hasher, - sdk: sdk, - encKey: encKey, - idProvider: idp, - } -} - -func (bs bootstrapService) Add(ctx context.Context, session smqauthn.Session, token string, cfg Config) (Config, error) { - id, err := bs.idProvider.ID() - if err != nil { - return Config{}, errors.Wrap(ErrAddBootstrap, err) - } - - hashedKey, err := bs.hasher.Hash(cfg.ExternalKey) - if err != nil { - return Config{}, errors.Wrap(ErrAddBootstrap, err) - } - - cfg.ID = id - cfg.DomainID = session.DomainID - cfg.Status = Active - cfg.ExternalKey = hashedKey - - saved, err := bs.configs.Save(ctx, cfg) - if err != nil { - if errors.Contains(err, repoerr.ErrConflict) { - return Config{}, errors.Wrap(svcerr.ErrConflict, err) - } - return Config{}, errors.Wrap(ErrAddBootstrap, err) - } - - cfg.ID = saved - return cfg, nil -} - -func (bs bootstrapService) View(ctx context.Context, session smqauthn.Session, id string) (Config, error) { - cfg, err := bs.configs.RetrieveByID(ctx, session.DomainID, id) - if err != nil { - return Config{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - return cfg, nil -} - -func (bs bootstrapService) Update(ctx context.Context, session smqauthn.Session, cfg Config) error { - cfg.DomainID = session.DomainID - if err := bs.configs.Update(ctx, cfg); err != nil { - return errors.Wrap(svcerr.ErrUpdateEntity, err) - } - return nil -} - -func (bs bootstrapService) UpdateCert(ctx context.Context, session smqauthn.Session, id, clientCert, clientKey, caCert string) (Config, error) { - cfg, err := bs.configs.UpdateCert(ctx, session.DomainID, id, clientCert, clientKey, caCert) - if err != nil { - return Config{}, errors.Wrap(errUpdateCert, err) - } - return cfg, nil -} - -func (bs bootstrapService) List(ctx context.Context, session smqauthn.Session, filter Filter, offset, limit uint64) (ConfigsPage, error) { - return bs.configs.RetrieveAll(ctx, session.DomainID, filter, offset, limit), nil -} - -func (bs bootstrapService) Remove(ctx context.Context, session smqauthn.Session, id string) error { - if err := bs.configs.Remove(ctx, session.DomainID, id); err != nil { - return errors.Wrap(errRemoveBootstrap, err) - } - return nil -} - -func (bs bootstrapService) Bootstrap(ctx context.Context, externalKey, externalID string, secure bool) (Config, error) { - cfg, err := bs.configs.RetrieveByExternalID(ctx, externalID) - if err != nil { - return cfg, errors.Wrap(ErrBootstrap, err) - } - if secure { - dec, err := bs.dec(externalKey) - if err != nil { - return Config{}, errors.Wrap(ErrExternalKeySecure, err) - } - externalKey = dec - } - - if err := bs.hasher.Compare(externalKey, cfg.ExternalKey); err != nil { - return Config{}, ErrExternalKey - } - if cfg.Status == DisabledStatus { - return Config{}, ErrBootstrap - } - - cfg, err = bs.renderBootstrapConfig(ctx, cfg) - if err != nil { - return Config{}, errors.Wrap(ErrBootstrap, err) - } - - return cfg, nil -} - -func (bs bootstrapService) renderBootstrapConfig(ctx context.Context, cfg Config) (Config, error) { - if cfg.ProfileID == "" { - return cfg, nil - } - if bs.profiles == nil || bs.bindings == nil || bs.renderer == nil { - return Config{}, errors.Wrap(errRenderBootstrap, errors.New("profile rendering support not configured")) - } - - profile, err := bs.profiles.RetrieveByID(ctx, cfg.DomainID, cfg.ProfileID) - if err != nil { - return Config{}, errors.Wrap(errRenderBootstrap, err) - } - - bindings, err := bs.bindings.Retrieve(ctx, cfg.ID) - if err != nil { - return Config{}, errors.Wrap(errRenderBootstrap, err) - } - if err := validateRequiredBindings(profile, bindings); err != nil { - return Config{}, errors.Wrap(errRenderBootstrap, err) - } - bindings, err = bs.decryptSecretSnapshots(bindings) - if err != nil { - return Config{}, errors.Wrap(errRenderBootstrap, err) - } - - rendered, err := bs.renderer.Render(profile, cfg, bindings) - if err != nil { - return Config{}, errors.Wrap(errRenderBootstrap, err) - } - - cfg.Content = string(rendered) - return cfg, nil -} - -func (bs bootstrapService) EnableConfig(ctx context.Context, session smqauthn.Session, id string) (Config, error) { - cfg, err := bs.changeConfigStatus(ctx, session.DomainID, id, EnabledStatus) - if err != nil { - return Config{}, errors.Wrap(errEnableConfig, err) - } - return cfg, nil -} - -func (bs bootstrapService) DisableConfig(ctx context.Context, session smqauthn.Session, id string) (Config, error) { - cfg, err := bs.changeConfigStatus(ctx, session.DomainID, id, DisabledStatus) - if err != nil { - return Config{}, errors.Wrap(errDisableConfig, err) - } - return cfg, nil -} - -func (bs bootstrapService) changeConfigStatus(ctx context.Context, domainID, id string, status Status) (Config, error) { - cfg, err := bs.configs.RetrieveByID(ctx, domainID, id) - if err != nil { - return Config{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - if cfg.Status == status { - return cfg, nil - } - if err := bs.configs.ChangeStatus(ctx, domainID, id, status); err != nil { - return Config{}, errors.Wrap(svcerr.ErrUpdateEntity, err) - } - cfg.Status = status - return cfg, nil -} - -// --- Profile management --- - -func (bs bootstrapService) CreateProfile(ctx context.Context, session smqauthn.Session, p Profile) (Profile, error) { - if bs.profiles == nil { - return Profile{}, errors.Wrap(errCreateProfile, errors.New("profile repository not configured")) - } - id, err := bs.idProvider.ID() - if err != nil { - return Profile{}, errors.Wrap(errCreateProfile, err) - } - p.ID = id - p.DomainID = session.DomainID - if p.ContentFormat == "" { - p.ContentFormat = ContentFormatJSON - } - p.Version = 1 - if err := validateProfileBindingSlots(p); err != nil { - return Profile{}, errors.Wrap(errCreateProfile, err) - } - if err := validateProfileTemplate(p); err != nil { - return Profile{}, errors.Wrap(errCreateProfile, err) - } - saved, err := bs.profiles.Save(ctx, p) - if err != nil { - return Profile{}, errors.Wrap(errCreateProfile, err) - } - return saved, nil -} - -func (bs bootstrapService) ViewProfile(ctx context.Context, session smqauthn.Session, profileID string) (Profile, error) { - if bs.profiles == nil { - return Profile{}, errors.Wrap(errViewProfile, errors.New("profile repository not configured")) - } - p, err := bs.profiles.RetrieveByID(ctx, session.DomainID, profileID) - if err != nil { - return Profile{}, errors.Wrap(errViewProfile, err) - } - return p, nil -} - -func (bs bootstrapService) UpdateProfile(ctx context.Context, session smqauthn.Session, p Profile) (Profile, error) { - if bs.profiles == nil { - return Profile{}, errors.Wrap(errUpdateProfile, errors.New("profile repository not configured")) - } - p.DomainID = session.DomainID - if err := validateProfileBindingSlots(p); err != nil { - return Profile{}, errors.Wrap(errUpdateProfile, err) - } - if err := validateProfileTemplate(p); err != nil { - return Profile{}, errors.Wrap(errUpdateProfile, err) - } - updated, err := bs.profiles.Update(ctx, p) - if err != nil { - return Profile{}, errors.Wrap(errUpdateProfile, err) - } - return updated, nil -} - -func (bs bootstrapService) ListProfiles(ctx context.Context, session smqauthn.Session, offset, limit uint64, name string) (ProfilesPage, error) { - if bs.profiles == nil { - return ProfilesPage{}, errors.Wrap(errListProfiles, errors.New("profile repository not configured")) - } - page, err := bs.profiles.RetrieveAll(ctx, session.DomainID, offset, limit, name) - if err != nil { - return ProfilesPage{}, errors.Wrap(errListProfiles, err) - } - return page, nil -} - -func (bs bootstrapService) DeleteProfile(ctx context.Context, session smqauthn.Session, profileID string) error { - if bs.profiles == nil { - return errors.Wrap(errDeleteProfile, errors.New("profile repository not configured")) - } - if err := bs.profiles.Delete(ctx, session.DomainID, profileID); err != nil { - return errors.Wrap(errDeleteProfile, err) - } - return nil -} - -// --- Enrollment-profile assignment --- - -func (bs bootstrapService) AssignProfile(ctx context.Context, session smqauthn.Session, configID, profileID string) error { - if bs.profiles == nil { - return errors.Wrap(errAssignProfile, errors.New("profile repository not configured")) - } - // Validate profile exists in domain. - if _, err := bs.profiles.RetrieveByID(ctx, session.DomainID, profileID); err != nil { - return errors.Wrap(errAssignProfile, err) - } - if err := bs.configs.AssignProfile(ctx, session.DomainID, configID, profileID); err != nil { - return errors.Wrap(errAssignProfile, err) - } - return nil -} - -// --- Binding management --- - -func (bs bootstrapService) BindResources(ctx context.Context, session smqauthn.Session, token, configID string, requested []BindingRequest) error { - if bs.profiles == nil || bs.bindings == nil || bs.resolver == nil { - return errors.Wrap(errBindResources, errors.New("binding support not configured")) - } - cfg, err := bs.configs.RetrieveByID(ctx, session.DomainID, configID) - if err != nil { - return errors.Wrap(errBindResources, err) - } - profile, err := bs.profiles.RetrieveByID(ctx, session.DomainID, cfg.ProfileID) - if err != nil { - return errors.Wrap(errBindResources, err) - } - if err := validateRequestedBindings(profile, requested); err != nil { - return errors.Wrap(errBindResources, err) - } - snapshots, err := bs.resolver.Resolve(ctx, ResolveRequest{ - Enrollment: cfg, - Token: token, - Requested: requested, - }) - if err != nil { - return errors.Wrap(errBindResources, err) - } - existing, err := bs.bindings.Retrieve(ctx, configID) - if err != nil { - return errors.Wrap(errBindResources, err) - } - if err := validateRequiredBindings(profile, mergeBindingSnapshots(existing, snapshots)); err != nil { - return errors.Wrap(errBindResources, err) - } - snapshots, err = bs.encryptSecretSnapshots(snapshots) - if err != nil { - return errors.Wrap(errBindResources, err) - } - if err := bs.bindings.Save(ctx, configID, snapshots); err != nil { - return errors.Wrap(errBindResources, err) - } - return nil -} - -func (bs bootstrapService) ListBindings(ctx context.Context, session smqauthn.Session, configID string) ([]BindingSnapshot, error) { - if bs.bindings == nil { - return nil, errors.Wrap(errListBindings, errors.New("binding support not configured")) - } - if _, err := bs.configs.RetrieveByID(ctx, session.DomainID, configID); err != nil { - return nil, errors.Wrap(errListBindings, err) - } - snapshots, err := bs.bindings.Retrieve(ctx, configID) - if err != nil { - return nil, errors.Wrap(errListBindings, err) - } - return hideSecretSnapshots(snapshots), nil -} - -func (bs bootstrapService) RefreshBindings(ctx context.Context, session smqauthn.Session, token, configID string) error { - if bs.profiles == nil || bs.bindings == nil || bs.resolver == nil { - return errors.Wrap(errRefreshBinding, errors.New("binding support not configured")) - } - cfg, err := bs.configs.RetrieveByID(ctx, session.DomainID, configID) - if err != nil { - return errors.Wrap(errRefreshBinding, err) - } - profile, err := bs.profiles.RetrieveByID(ctx, session.DomainID, cfg.ProfileID) - if err != nil { - return errors.Wrap(errRefreshBinding, err) - } - existing, err := bs.bindings.Retrieve(ctx, configID) - if err != nil { - return errors.Wrap(errRefreshBinding, err) - } - if len(existing) == 0 { - return nil - } - // Re-resolve every existing binding to refresh its snapshot. - requested := make([]BindingRequest, len(existing)) - for i, b := range existing { - requested[i] = BindingRequest{Slot: b.Slot, Type: b.Type, ResourceID: b.ResourceID} - } - if err := validateRequestedBindings(profile, requested); err != nil { - return errors.Wrap(errRefreshBinding, err) - } - refreshed, err := bs.resolver.Resolve(ctx, ResolveRequest{ - Enrollment: cfg, - Token: token, - Requested: requested, - }) - if err != nil { - return errors.Wrap(errRefreshBinding, err) - } - if err := validateRequiredBindings(profile, refreshed); err != nil { - return errors.Wrap(errRefreshBinding, err) - } - refreshed, err = bs.encryptSecretSnapshots(refreshed) - if err != nil { - return errors.Wrap(errRefreshBinding, err) - } - return bs.bindings.Save(ctx, configID, refreshed) -} - -func (bs bootstrapService) dec(in string) (string, error) { - ciphertext, err := hex.DecodeString(in) - if err != nil { - return "", err - } - block, err := aes.NewCipher(bs.encKey) - if err != nil { - return "", err - } - if len(ciphertext) < aes.BlockSize { - return "", err - } - iv := ciphertext[:aes.BlockSize] - ciphertext = ciphertext[aes.BlockSize:] - stream := cipher.NewCFBDecrypter(block, iv) - stream.XORKeyStream(ciphertext, ciphertext) - return string(ciphertext), nil -} diff --git a/bootstrap/service_test.go b/bootstrap/service_test.go deleted file mode 100644 index d99986e7b..000000000 --- a/bootstrap/service_test.go +++ /dev/null @@ -1,1639 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package bootstrap_test - -import ( - "context" - "crypto/aes" - "crypto/cipher" - "crypto/rand" - "encoding/hex" - "fmt" - "io" - "testing" - - "github.com/absmach/magistrala/bootstrap" - bootstraphasher "github.com/absmach/magistrala/bootstrap/hasher" - mocks "github.com/absmach/magistrala/bootstrap/mocks" - "github.com/absmach/magistrala/internal/testsutil" - smqauthn "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - sdkmocks "github.com/absmach/magistrala/pkg/sdk/mocks" - "github.com/absmach/magistrala/pkg/uuid" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" -) - -const ( - validToken = "validToken" - invalidDomainID = "invalid" - unknown = "unknown" - validID = "d4ebb847-5d0e-4e46-bdd9-b6aceaaa3a22" -) - -var ( - encKey = []byte("1234567891011121") - domainID = testsutil.GenerateUUID(&testing.T{}) - - config = bootstrap.Config{ - ID: testsutil.GenerateUUID(&testing.T{}), - ExternalID: testsutil.GenerateUUID(&testing.T{}), - ExternalKey: testsutil.GenerateUUID(&testing.T{}), - Content: "config", - } -) - -var ( - boot *mocks.ConfigRepository - sdk *sdkmocks.SDK - profileRepo *mocks.ProfileRepository - bindingStore *mocks.BindingStore - resolver *mocks.BindingResolver - renderer *mocks.Renderer -) - -func newService() bootstrap.Service { - boot = new(mocks.ConfigRepository) - sdk = new(sdkmocks.SDK) - profileRepo = new(mocks.ProfileRepository) - bindingStore = new(mocks.BindingStore) - resolver = new(mocks.BindingResolver) - renderer = new(mocks.Renderer) - idp := uuid.NewMock() - return bootstrap.New(boot, profileRepo, bindingStore, resolver, renderer, sdk, bootstraphasher.New(), encKey, idp) -} - -func enc(in []byte) ([]byte, error) { - block, err := aes.NewCipher(encKey) - if err != nil { - return nil, err - } - ciphertext := make([]byte, aes.BlockSize+len(in)) - iv := ciphertext[:aes.BlockSize] - if _, err := io.ReadFull(rand.Reader, iv); err != nil { - return nil, err - } - stream := cipher.NewCFBEncrypter(block, iv) - stream.XORKeyStream(ciphertext[aes.BlockSize:], in) - return ciphertext, nil -} - -func TestAdd(t *testing.T) { - svc := newService() - - neID := config - neID.ID = "non-existent" - - cases := []struct { - desc string - config bootstrap.Config - token string - session smqauthn.Session - userID string - domainID string - saveErr error - err error - }{ - { - desc: "add a new config", - config: config, - token: validToken, - userID: validID, - domainID: domainID, - err: nil, - }, - { - desc: "add a config with an invalid ID", - config: neID, - token: validToken, - userID: validID, - domainID: domainID, - err: nil, - }, - { - desc: "add empty config", - config: bootstrap.Config{}, - token: validToken, - userID: validID, - domainID: domainID, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - tc.session = smqauthn.Session{UserID: tc.userID, DomainID: tc.domainID, DomainUserID: validID} - repoCall3 := boot.On("Save", context.Background(), mock.Anything).Return(mock.Anything, tc.saveErr) - _, err := svc.Add(context.Background(), tc.session, tc.token, tc.config) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - repoCall3.Unset() - }) - } -} - -func TestView(t *testing.T) { - svc := newService() - - cases := []struct { - desc string - configID string - userID string - domain string - clientDomain string - token string - session smqauthn.Session - retrieveErr error - clientErr error - channelErr error - err error - }{ - { - desc: "view an existing config", - configID: config.ID, - userID: validID, - clientDomain: domainID, - domain: domainID, - token: validToken, - err: nil, - }, - { - desc: "view a non-existing config", - configID: unknown, - userID: validID, - clientDomain: domainID, - domain: domainID, - token: validToken, - retrieveErr: svcerr.ErrNotFound, - err: svcerr.ErrNotFound, - }, - { - desc: "view a config with invalid domain", - configID: config.ID, - userID: validID, - clientDomain: invalidDomainID, - domain: invalidDomainID, - token: validToken, - retrieveErr: svcerr.ErrNotFound, - err: svcerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - tc.session = smqauthn.Session{UserID: tc.userID, DomainID: tc.domain, DomainUserID: validID} - repoCall := boot.On("RetrieveByID", context.Background(), tc.clientDomain, tc.configID).Return(config, tc.retrieveErr) - _, err := svc.View(context.Background(), tc.session, tc.configID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - repoCall.Unset() - }) - } -} - -func TestUpdate(t *testing.T) { - svc := newService() - - c := config - - modifiedCreated := c - modifiedCreated.Content = "new-config" - modifiedCreated.Name = "new name" - - modifiedRenderContext := c - modifiedRenderContext.RenderContext = map[string]any{ - "site": "warehouse-2", - "region": "mombasa", - } - - nonExisting := c - nonExisting.ID = unknown - - cases := []struct { - desc string - config bootstrap.Config - token string - session smqauthn.Session - userID string - domainID string - updateErr error - err error - }{ - { - desc: "update a config with status Created", - config: modifiedCreated, - token: validToken, - userID: validID, - domainID: domainID, - err: nil, - }, - { - desc: "update a config render_context", - config: modifiedRenderContext, - token: validToken, - userID: validID, - domainID: domainID, - err: nil, - }, - { - desc: "update a non-existing config", - config: nonExisting, - token: validToken, - userID: validID, - domainID: domainID, - updateErr: svcerr.ErrNotFound, - err: svcerr.ErrNotFound, - }, - { - desc: "update a config with update error", - config: c, - token: validToken, - userID: validID, - domainID: domainID, - updateErr: svcerr.ErrUpdateEntity, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - tc.session = smqauthn.Session{UserID: tc.userID, DomainID: tc.domainID, DomainUserID: validID} - repoCall := boot.On("Update", context.Background(), mock.Anything).Return(tc.updateErr) - err := svc.Update(context.Background(), tc.session, tc.config) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - repoCall.Unset() - }) - } -} - -func TestUpdateCert(t *testing.T) { - svc := newService() - - c := config - - cases := []struct { - desc string - token string - session smqauthn.Session - userID string - domainID string - configID string - clientCert string - clientKey string - caCert string - expectedConfig bootstrap.Config - authorizeErr error - authenticateErr error - updateErr error - err error - }{ - { - desc: "update certs for the valid config", - userID: validID, - domainID: domainID, - configID: c.ID, - clientCert: "newCert", - clientKey: "newKey", - caCert: "newCert", - token: validToken, - expectedConfig: bootstrap.Config{ - Name: c.Name, - ExternalID: c.ExternalID, - ExternalKey: c.ExternalKey, - Content: c.Content, - Status: c.Status, - DomainID: c.DomainID, - ID: c.ID, - ClientCert: "newCert", - CACert: "newCert", - ClientKey: "newKey", - }, - err: nil, - }, - { - desc: "update cert for a non-existing config", - userID: validID, - domainID: domainID, - configID: "empty", - clientCert: "newCert", - clientKey: "newKey", - caCert: "newCert", - token: validToken, - expectedConfig: bootstrap.Config{}, - updateErr: svcerr.ErrNotFound, - err: svcerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - tc.session = smqauthn.Session{UserID: tc.userID, DomainID: tc.domainID, DomainUserID: validID} - repoCall := boot.On("UpdateCert", context.Background(), mock.Anything, mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(tc.expectedConfig, tc.updateErr) - cfg, err := svc.UpdateCert(context.Background(), tc.session, tc.configID, tc.clientCert, tc.clientKey, tc.caCert) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.expectedConfig, cfg, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.expectedConfig, cfg)) - repoCall.Unset() - }) - } -} - -func TestList(t *testing.T) { - svc := newService() - - numClients := 101 - var saved []bootstrap.Config - for i := 0; i < numClients; i++ { - c := config - c.ExternalID = testsutil.GenerateUUID(t) - c.ExternalKey = testsutil.GenerateUUID(t) - c.Name = fmt.Sprintf("%s-%d", config.Name, i) - if i == 41 { - c.Status = bootstrap.Active - } - saved = append(saved, c) - } - cases := []struct { - desc string - config bootstrap.ConfigsPage - filter bootstrap.Filter - offset uint64 - limit uint64 - token string - session smqauthn.Session - userID string - domainID string - retrieveErr error - err error - }{ - { - desc: "list configs successfully as super admin", - config: bootstrap.ConfigsPage{ - Total: uint64(len(saved)), - Offset: 0, - Limit: 10, - Configs: saved[0:10], - }, - filter: bootstrap.Filter{}, - token: validToken, - session: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: validID, SuperAdmin: true}, - userID: validID, - domainID: domainID, - offset: 0, - limit: 10, - err: nil, - }, - { - desc: "list configs with failed super admin check", - config: bootstrap.ConfigsPage{}, - filter: bootstrap.Filter{}, - token: validID, - session: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: validID}, - userID: validID, - domainID: domainID, - offset: 0, - limit: 10, - err: nil, - }, - { - desc: "list configs successfully as domain admin", - config: bootstrap.ConfigsPage{ - Total: uint64(len(saved)), - Offset: 0, - Limit: 10, - Configs: saved[0:10], - }, - filter: bootstrap.Filter{}, - token: validToken, - userID: validID, - domainID: domainID, - session: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: validID, SuperAdmin: true}, - offset: 0, - limit: 10, - err: nil, - }, - { - desc: "list configs successfully as non admin", - config: bootstrap.ConfigsPage{ - Total: uint64(len(saved)), - Offset: 0, - Limit: 10, - Configs: saved[0:10], - }, - filter: bootstrap.Filter{}, - token: validToken, - userID: validID, - domainID: domainID, - session: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: validID}, - offset: 0, - limit: 10, - err: nil, - }, - { - desc: "list configs with specified name as super admin", - config: bootstrap.ConfigsPage{ - Total: 1, - Offset: 0, - Limit: 100, - Configs: saved[95:96], - }, - filter: bootstrap.Filter{PartialMatch: map[string]string{"name": "95"}}, - token: validToken, - session: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: validID, SuperAdmin: true}, - userID: validID, - domainID: domainID, - offset: 0, - limit: 100, - err: nil, - }, - { - desc: "list configs with specified name as domain admin", - config: bootstrap.ConfigsPage{ - Total: 1, - Offset: 0, - Limit: 100, - Configs: saved[95:96], - }, - filter: bootstrap.Filter{PartialMatch: map[string]string{"name": "95"}}, - token: validToken, - userID: validID, - domainID: domainID, - session: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: validID, SuperAdmin: true}, - offset: 0, - limit: 100, - err: nil, - }, - { - desc: "list configs with specified name as non admin", - config: bootstrap.ConfigsPage{ - Total: 1, - Offset: 0, - Limit: 100, - Configs: saved[95:96], - }, - filter: bootstrap.Filter{PartialMatch: map[string]string{"name": "95"}}, - token: validToken, - userID: validID, - domainID: domainID, - session: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: validID}, - offset: 0, - limit: 100, - err: nil, - }, - { - desc: "list last page as super admin", - config: bootstrap.ConfigsPage{ - Total: uint64(len(saved)), - Offset: 95, - Limit: 10, - Configs: saved[95:], - }, - filter: bootstrap.Filter{}, - token: validToken, - userID: validID, - domainID: domainID, - session: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: validID, SuperAdmin: true}, - offset: 95, - limit: 10, - err: nil, - }, - { - desc: "list last page as domain admin", - config: bootstrap.ConfigsPage{ - Total: uint64(len(saved)), - Offset: 95, - Limit: 10, - Configs: saved[95:], - }, - filter: bootstrap.Filter{}, - token: validToken, - userID: validID, - domainID: domainID, - session: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: validID, SuperAdmin: true}, - offset: 95, - limit: 10, - err: nil, - }, - { - desc: "list last page as non admin", - config: bootstrap.ConfigsPage{ - Total: uint64(len(saved)), - Offset: 95, - Limit: 10, - Configs: saved[95:], - }, - filter: bootstrap.Filter{}, - token: validToken, - userID: validID, - domainID: domainID, - session: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: validID}, - offset: 95, - limit: 10, - err: nil, - }, - { - desc: "list configs with Active status as super admin", - config: bootstrap.ConfigsPage{ - Total: 1, - Offset: 35, - Limit: 20, - Configs: []bootstrap.Config{saved[41]}, - }, - filter: bootstrap.Filter{FullMatch: map[string]string{"status": bootstrap.Active.String()}}, - token: validToken, - userID: validID, - domainID: domainID, - session: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: validID, SuperAdmin: true}, - offset: 35, - limit: 20, - err: nil, - }, - { - desc: "list configs with Active status as domain admin", - config: bootstrap.ConfigsPage{ - Total: 1, - Offset: 35, - Limit: 20, - Configs: []bootstrap.Config{saved[41]}, - }, - filter: bootstrap.Filter{FullMatch: map[string]string{"status": bootstrap.Active.String()}}, - token: validToken, - userID: validID, - domainID: domainID, - session: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: validID, SuperAdmin: true}, - offset: 35, - limit: 20, - err: nil, - }, - { - desc: "list configs with Active status as non admin", - config: bootstrap.ConfigsPage{ - Total: 1, - Offset: 35, - Limit: 20, - Configs: []bootstrap.Config{saved[41]}, - }, - filter: bootstrap.Filter{FullMatch: map[string]string{"status": bootstrap.Active.String()}}, - token: validToken, - userID: validID, - domainID: domainID, - session: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: validID}, - offset: 35, - limit: 20, - err: nil, - }, - { - desc: "list configs with empty result", - config: bootstrap.ConfigsPage{}, - filter: bootstrap.Filter{}, - offset: 0, - limit: 10, - token: validToken, - userID: validID, - domainID: domainID, - session: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: validID}, - err: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := boot.On("RetrieveAll", context.Background(), mock.Anything, tc.filter, tc.offset, tc.limit).Return(tc.config, tc.retrieveErr) - - result, err := svc.List(context.Background(), tc.session, tc.filter, tc.offset, tc.limit) - assert.ElementsMatch(t, tc.config.Configs, result.Configs, fmt.Sprintf("%s: expected %v got %v", tc.desc, tc.config.Configs, result.Configs)) - assert.Equal(t, tc.config.Total, result.Total, fmt.Sprintf("%s: expected %v got %v", tc.desc, tc.config.Total, result.Total)) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - repoCall.Unset() - }) - } -} - -func TestRemove(t *testing.T) { - svc := newService() - - c := config - cases := []struct { - desc string - id string - token string - session smqauthn.Session - userID string - domainID string - removeErr error - err error - }{ - { - desc: "remove an existing config", - id: c.ID, - token: validToken, - userID: validID, - domainID: domainID, - err: nil, - }, - { - desc: "remove removed config", - id: c.ID, - token: validToken, - userID: validID, - domainID: domainID, - err: nil, - }, - { - desc: "remove a config with failed remove", - id: c.ID, - token: validToken, - userID: validID, - domainID: domainID, - removeErr: svcerr.ErrRemoveEntity, - err: svcerr.ErrRemoveEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - tc.session = smqauthn.Session{UserID: tc.userID, DomainID: tc.domainID, DomainUserID: validID} - repoCall := boot.On("Remove", context.Background(), mock.Anything, mock.Anything).Return(tc.removeErr) - err := svc.Remove(context.Background(), tc.session, tc.id) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - repoCall.Unset() - }) - } -} - -func TestBootstrap(t *testing.T) { - svc := newService() - - c := config - c.Status = bootstrap.Active - e, err := enc([]byte(c.ExternalKey)) - assert.Nil(t, err, fmt.Sprintf("Encrypting external key expected to succeed: %s.\n", err)) - - cases := []struct { - desc string - config bootstrap.Config - externalKey string - externalID string - userID string - domainID string - err error - encrypted bool - }{ - { - desc: "bootstrap using invalid external id", - config: bootstrap.Config{}, - externalID: "invalid", - externalKey: c.ExternalKey, - userID: validID, - domainID: invalidDomainID, - err: svcerr.ErrNotFound, - encrypted: false, - }, - { - desc: "bootstrap using invalid external key", - config: bootstrap.Config{}, - externalID: c.ExternalID, - externalKey: "invalid", - userID: validID, - domainID: domainID, - err: bootstrap.ErrExternalKey, - encrypted: false, - }, - { - desc: "bootstrap an existing config", - config: c, - externalID: c.ExternalID, - externalKey: c.ExternalKey, - userID: validID, - domainID: domainID, - err: nil, - encrypted: false, - }, - { - desc: "bootstrap encrypted", - config: c, - externalID: c.ExternalID, - externalKey: hex.EncodeToString(e), - userID: validID, - domainID: domainID, - err: nil, - encrypted: true, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := boot.On("RetrieveByExternalID", context.Background(), mock.Anything).Return(tc.config, tc.err) - config, err := svc.Bootstrap(context.Background(), tc.externalKey, tc.externalID, tc.encrypted) - assert.Equal(t, tc.config, config, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.config, config)) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - repoCall.Unset() - }) - } -} - -func TestBootstrapRender(t *testing.T) { - profile := bootstrap.Profile{ - ID: testsutil.GenerateUUID(&testing.T{}), - DomainID: domainID, - Name: "gateway-profile", - ContentFormat: bootstrap.ContentFormatGoTemplate, - ContentTemplate: `{"mode":"profile"}`, - } - bindings := []bootstrap.BindingSnapshot{ - { - ConfigID: config.ID, - Slot: "mqtt_client", - Type: "client", - ResourceID: config.ID, - Snapshot: map[string]any{ - "id": config.ID, - }, - }, - } - - cases := []struct { - desc string - cfg bootstrap.Config - rendererOut []byte - rendererErr error - rendered string - err error - }{ - { - desc: "bootstrap renders assigned profile content", - cfg: func() bootstrap.Config { - cfg := config - cfg.DomainID = domainID - cfg.ProfileID = profile.ID - cfg.Status = bootstrap.Active - cfg.Content = "legacy" - return cfg - }(), - rendererOut: []byte(`{"mode":"profile"}`), - rendered: `{"mode":"profile"}`, - }, - { - desc: "bootstrap falls back to legacy content when no profile is assigned", - cfg: func() bootstrap.Config { - cfg := config - cfg.DomainID = domainID - cfg.Status = bootstrap.Active - cfg.Content = "legacy" - return cfg - }(), - rendered: "legacy", - }, - { - desc: "bootstrap fails when renderer fails", - cfg: func() bootstrap.Config { - cfg := config - cfg.DomainID = domainID - cfg.ProfileID = profile.ID - cfg.Status = bootstrap.Active - return cfg - }(), - rendererErr: errors.New("render failed"), - err: errors.New("render failed"), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svc := newService() - repoCall := boot.On("RetrieveByExternalID", context.Background(), tc.cfg.ExternalID).Return(tc.cfg, nil) - - var prCall, bsCall, rndCall *mock.Call - if tc.cfg.ProfileID != "" { - prCall = profileRepo.On("RetrieveByID", context.Background(), tc.cfg.DomainID, tc.cfg.ProfileID).Return(profile, nil) - bsCall = bindingStore.On("Retrieve", context.Background(), tc.cfg.ID).Return(bindings, nil) - rndCall = renderer.On("Render", mock.Anything, mock.Anything, mock.Anything).Return(tc.rendererOut, tc.rendererErr) - } - - res, err := svc.Bootstrap(context.Background(), tc.cfg.ExternalKey, tc.cfg.ExternalID, false) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.err, err)) - if tc.err == nil { - assert.Equal(t, tc.rendered, res.Content, fmt.Sprintf("%s: expected rendered content %q got %q\n", tc.desc, tc.rendered, res.Content)) - } - - repoCall.Unset() - if prCall != nil { - prCall.Unset() - } - if bsCall != nil { - bsCall.Unset() - } - if rndCall != nil { - rndCall.Unset() - } - }) - } -} - -func TestEnableConfig(t *testing.T) { - svc := newService() - - c := config - activeConfig := config - activeConfig.Status = bootstrap.Active - inactiveConfig := config - inactiveConfig.Status = bootstrap.Inactive - - cases := []struct { - desc string - config bootstrap.Config - id string - session smqauthn.Session - userID string - domainID string - retrieveErr error - statusErr error - err error - }{ - { - desc: "enable non-existing config", - config: c, - id: unknown, - userID: validID, - domainID: domainID, - retrieveErr: svcerr.ErrNotFound, - err: svcerr.ErrNotFound, - }, - { - desc: "enable inactive config", - config: inactiveConfig, - id: c.ID, - userID: validID, - domainID: domainID, - err: nil, - }, - { - desc: "enable already active config", - config: activeConfig, - id: c.ID, - userID: validID, - domainID: domainID, - err: nil, - }, - { - desc: "enable with repo error", - config: inactiveConfig, - id: c.ID, - userID: validID, - domainID: domainID, - statusErr: svcerr.ErrUpdateEntity, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - tc.session = smqauthn.Session{UserID: tc.userID, DomainID: tc.domainID, DomainUserID: validID} - repoCall := boot.On("RetrieveByID", context.Background(), tc.domainID, tc.id).Return(tc.config, tc.retrieveErr) - repoCall1 := boot.On("ChangeStatus", context.Background(), mock.Anything, mock.Anything, mock.Anything).Return(tc.statusErr) - _, err := svc.EnableConfig(context.Background(), tc.session, tc.id) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - repoCall.Unset() - repoCall1.Unset() - }) - } -} - -func TestDisableConfig(t *testing.T) { - svc := newService() - - c := config - activeConfig := config - activeConfig.Status = bootstrap.Active - inactiveConfig := config - inactiveConfig.Status = bootstrap.Inactive - - cases := []struct { - desc string - config bootstrap.Config - id string - session smqauthn.Session - userID string - domainID string - retrieveErr error - statusErr error - err error - }{ - { - desc: "disable non-existing config", - config: c, - id: unknown, - userID: validID, - domainID: domainID, - retrieveErr: svcerr.ErrNotFound, - err: svcerr.ErrNotFound, - }, - { - desc: "disable active config", - config: activeConfig, - id: c.ID, - userID: validID, - domainID: domainID, - err: nil, - }, - { - desc: "disable already inactive config", - config: inactiveConfig, - id: c.ID, - userID: validID, - domainID: domainID, - err: nil, - }, - { - desc: "disable with repo error", - config: activeConfig, - id: c.ID, - userID: validID, - domainID: domainID, - statusErr: svcerr.ErrUpdateEntity, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - tc.session = smqauthn.Session{UserID: tc.userID, DomainID: tc.domainID, DomainUserID: validID} - repoCall := boot.On("RetrieveByID", context.Background(), tc.domainID, tc.id).Return(tc.config, tc.retrieveErr) - repoCall1 := boot.On("ChangeStatus", context.Background(), mock.Anything, mock.Anything, mock.Anything).Return(tc.statusErr) - _, err := svc.DisableConfig(context.Background(), tc.session, tc.id) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - repoCall.Unset() - repoCall1.Unset() - }) - } -} - -func TestAssignProfile(t *testing.T) { - profile := bootstrap.Profile{ - ID: testsutil.GenerateUUID(t), - DomainID: domainID, - Name: "gateway-profile", - ContentFormat: bootstrap.ContentFormatGoTemplate, - Version: 1, - } - - cases := []struct { - desc string - configID string - profileID string - retrieveErr error - assignErr error - expectedErr error - expectAssign bool - }{ - { - desc: "assign profile to enrollment", - configID: config.ID, - profileID: profile.ID, - expectAssign: true, - }, - { - desc: "assign profile with missing profile", - configID: config.ID, - profileID: profile.ID, - retrieveErr: svcerr.ErrNotFound, - expectedErr: svcerr.ErrNotFound, - }, - { - desc: "assign profile with repository error", - configID: config.ID, - profileID: profile.ID, - assignErr: svcerr.ErrUpdateEntity, - expectedErr: svcerr.ErrUpdateEntity, - expectAssign: true, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svc := newService() - session := smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: validID} - - prCall := profileRepo.On("RetrieveByID", context.Background(), domainID, tc.profileID).Return(profile, tc.retrieveErr) - - var assignCall *mock.Call - if tc.expectAssign { - assignCall = boot.On("AssignProfile", context.Background(), domainID, tc.configID, tc.profileID).Return(tc.assignErr) - } - - err := svc.AssignProfile(context.Background(), session, tc.configID, tc.profileID) - assert.True(t, errors.Contains(err, tc.expectedErr), fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.expectedErr, err)) - - prCall.Unset() - if assignCall != nil { - assignCall.Unset() - } - }) - } -} - -func TestCreateProfile(t *testing.T) { - session := smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: validID} - - validProfile := bootstrap.Profile{ - Name: "test-profile", - ContentFormat: bootstrap.ContentFormatGoTemplate, - } - - cases := []struct { - desc string - profile bootstrap.Profile - saveErr error - err error - wantFormat bootstrap.ContentFormat - }{ - { - desc: "create profile successfully", - profile: validProfile, - wantFormat: bootstrap.ContentFormatGoTemplate, - }, - { - desc: "create profile defaults to json format", - profile: bootstrap.Profile{Name: "no-format"}, - wantFormat: bootstrap.ContentFormatJSON, - }, - { - desc: "create profile with invalid slot: empty name", - profile: bootstrap.Profile{ - Name: "test", - BindingSlots: []bootstrap.BindingSlot{{Name: "", Type: "client"}}, - }, - err: errors.New("invalid binding slot: slot name is required"), - }, - { - desc: "create profile with invalid slot: empty type", - profile: bootstrap.Profile{ - Name: "test", - BindingSlots: []bootstrap.BindingSlot{{Name: "mqtt", Type: ""}}, - }, - err: errors.New("invalid binding slot: slot \"mqtt\" type is required"), - }, - { - desc: "create profile with duplicate slot names", - profile: bootstrap.Profile{ - Name: "test", - BindingSlots: []bootstrap.BindingSlot{ - {Name: "mqtt", Type: "client"}, - {Name: "mqtt", Type: "channel"}, - }, - }, - err: errors.New("invalid binding slot: duplicate slot \"mqtt\""), - }, - { - desc: "create profile with invalid template syntax", - profile: bootstrap.Profile{ - Name: "test", - ContentTemplate: `{{ index .Vars \"mqtt_url\" }}`, - }, - err: bootstrap.ErrRenderFailed, - }, - { - desc: "create profile with repository save error", - profile: validProfile, - saveErr: svcerr.ErrCreateEntity, - err: svcerr.ErrCreateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svc := newService() - saveCall := profileRepo.EXPECT().Save(mock.Anything, mock.Anything).RunAndReturn( - func(_ context.Context, p bootstrap.Profile) (bootstrap.Profile, error) { - return p, tc.saveErr - }) - saved, err := svc.CreateProfile(context.Background(), session, tc.profile) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.err, err)) - if tc.err == nil { - assert.NotEmpty(t, saved.ID, fmt.Sprintf("%s: expected non-empty profile ID\n", tc.desc)) - assert.Equal(t, domainID, saved.DomainID, fmt.Sprintf("%s: expected domain ID %s got %s\n", tc.desc, domainID, saved.DomainID)) - assert.Equal(t, tc.wantFormat, saved.ContentFormat, fmt.Sprintf("%s: expected %s format\n", tc.desc, tc.wantFormat)) - assert.Equal(t, 1, saved.Version, fmt.Sprintf("%s: expected version 1\n", tc.desc)) - } - saveCall.Unset() - }) - } -} - -func TestViewProfile(t *testing.T) { - session := smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: validID} - - profile := bootstrap.Profile{ - ID: testsutil.GenerateUUID(t), - DomainID: domainID, - Name: "view-profile", - ContentFormat: bootstrap.ContentFormatGoTemplate, - Version: 1, - } - - cases := []struct { - desc string - profileID string - retrieveErr error - err error - }{ - { - desc: "view profile successfully", - profileID: profile.ID, - }, - { - desc: "view non-existing profile", - profileID: unknown, - retrieveErr: svcerr.ErrNotFound, - err: svcerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svc := newService() - prCall := profileRepo.On("RetrieveByID", context.Background(), domainID, tc.profileID).Return(profile, tc.retrieveErr) - got, err := svc.ViewProfile(context.Background(), session, tc.profileID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.err, err)) - if tc.err == nil { - assert.Equal(t, profile, got, fmt.Sprintf("%s: expected profile %v got %v\n", tc.desc, profile, got)) - } - prCall.Unset() - }) - } -} - -func TestUpdateProfile(t *testing.T) { - session := smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: validID} - - validProfile := bootstrap.Profile{ - ID: testsutil.GenerateUUID(t), - DomainID: domainID, - Name: "updated-profile", - ContentFormat: bootstrap.ContentFormatGoTemplate, - } - - cases := []struct { - desc string - profile bootstrap.Profile - updateErr error - err error - }{ - { - desc: "update profile successfully", - profile: validProfile, - }, - { - desc: "update profile with only name", - profile: bootstrap.Profile{ID: validProfile.ID, Name: "no-format"}, - }, - { - desc: "update profile with invalid slot: empty type", - profile: bootstrap.Profile{ - ID: validProfile.ID, - Name: "test", - BindingSlots: []bootstrap.BindingSlot{{Name: "mqtt", Type: ""}}, - }, - err: errors.New("invalid binding slot: slot \"mqtt\" type is required"), - }, - { - desc: "update profile with duplicate slot names", - profile: bootstrap.Profile{ - ID: validProfile.ID, - Name: "test", - BindingSlots: []bootstrap.BindingSlot{ - {Name: "slot1", Type: "client"}, - {Name: "slot1", Type: "channel"}, - }, - }, - err: errors.New("invalid binding slot: duplicate slot \"slot1\""), - }, - { - desc: "update profile with invalid template syntax", - profile: bootstrap.Profile{ - ID: validProfile.ID, - Name: "test", - ContentTemplate: `{{ index .Vars \"mqtt_url\" }}`, - }, - err: bootstrap.ErrRenderFailed, - }, - { - desc: "update profile with repository error", - profile: validProfile, - updateErr: svcerr.ErrUpdateEntity, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svc := newService() - updateCall := profileRepo.On("Update", context.Background(), mock.Anything).Return(tc.profile, tc.updateErr) - _, err := svc.UpdateProfile(context.Background(), session, tc.profile) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.err, err)) - updateCall.Unset() - }) - } -} - -func TestListProfiles(t *testing.T) { - session := smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: validID} - - profiles := []bootstrap.Profile{ - {ID: testsutil.GenerateUUID(t), DomainID: domainID, Name: "p1", ContentFormat: bootstrap.ContentFormatGoTemplate, Version: 1}, - {ID: testsutil.GenerateUUID(t), DomainID: domainID, Name: "p2", ContentFormat: bootstrap.ContentFormatGoTemplate, Version: 1}, - } - page := bootstrap.ProfilesPage{Total: 2, Offset: 0, Limit: 10, Profiles: profiles} - filteredPage := bootstrap.ProfilesPage{Total: 1, Offset: 0, Limit: 10, Profiles: profiles[:1]} - - cases := []struct { - desc string - offset uint64 - limit uint64 - name string - page bootstrap.ProfilesPage - listErr error - err error - }{ - { - desc: "list profiles successfully", - limit: 10, - page: page, - }, - { - desc: "list profiles filtered by name", - limit: 10, - name: "p1", - page: filteredPage, - }, - { - desc: "list profiles with repository error", - limit: 10, - listErr: svcerr.ErrViewEntity, - err: svcerr.ErrViewEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svc := newService() - listCall := profileRepo.On("RetrieveAll", context.Background(), domainID, tc.offset, tc.limit, tc.name).Return(tc.page, tc.listErr) - got, err := svc.ListProfiles(context.Background(), session, tc.offset, tc.limit, tc.name) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.err, err)) - if tc.err == nil { - assert.Equal(t, tc.page, got, fmt.Sprintf("%s: expected page %v got %v\n", tc.desc, tc.page, got)) - } - listCall.Unset() - }) - } -} - -func TestDeleteProfile(t *testing.T) { - session := smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: validID} - profileID := testsutil.GenerateUUID(t) - - cases := []struct { - desc string - profileID string - deleteErr error - err error - }{ - { - desc: "delete profile successfully", - profileID: profileID, - }, - { - desc: "delete profile with repository error", - profileID: profileID, - deleteErr: svcerr.ErrRemoveEntity, - err: svcerr.ErrRemoveEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svc := newService() - deleteCall := profileRepo.On("Delete", context.Background(), domainID, tc.profileID).Return(tc.deleteErr) - err := svc.DeleteProfile(context.Background(), session, tc.profileID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.err, err)) - deleteCall.Unset() - }) - } -} - -func TestBindResources(t *testing.T) { - session := smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: validID} - - profile := bootstrap.Profile{ - ID: testsutil.GenerateUUID(t), - DomainID: domainID, - Name: "bind-profile", - BindingSlots: []bootstrap.BindingSlot{ - {Name: "mqtt", Type: "client", Required: true}, - }, - } - - channelProfile := bootstrap.Profile{ - ID: testsutil.GenerateUUID(t), - DomainID: domainID, - Name: "channel-profile", - BindingSlots: []bootstrap.BindingSlot{ - {Name: "data", Type: "channel", Required: true}, - }, - } - - cfg := bootstrap.Config{ - ID: config.ID, - DomainID: domainID, - ProfileID: profile.ID, - } - - channelCfg := bootstrap.Config{ - ID: config.ID, - DomainID: domainID, - ProfileID: channelProfile.ID, - } - - snapshot := bootstrap.BindingSnapshot{ - ConfigID: config.ID, - Slot: "mqtt", - Type: "client", - ResourceID: validID, - Snapshot: map[string]any{"id": validID}, - } - - channelSnapshot := bootstrap.BindingSnapshot{ - ConfigID: config.ID, - Slot: "data", - Type: "channel", - ResourceID: validID, - Snapshot: map[string]any{"id": validID}, - } - - requested := []bootstrap.BindingRequest{ - {Slot: "mqtt", Type: "client", ResourceID: validID}, - } - - channelRequested := []bootstrap.BindingRequest{ - {Slot: "data", Type: "channel", ResourceID: validID}, - } - - cases := []struct { - desc string - configID string - bindings []bootstrap.BindingRequest - cfgErr error - prErr error - resolveErr error - retrieveErr error - saveErr error - snapshots []bootstrap.BindingSnapshot - useChannel bool - err error - }{ - { - desc: "bind resources with config not found", - configID: config.ID, - bindings: requested, - cfgErr: svcerr.ErrNotFound, - err: svcerr.ErrNotFound, - }, - { - desc: "bind resources with profile not found", - configID: config.ID, - bindings: requested, - prErr: svcerr.ErrNotFound, - err: svcerr.ErrNotFound, - }, - { - desc: "bind resources with unknown slot", - configID: config.ID, - bindings: []bootstrap.BindingRequest{{Slot: "unknown", Type: "client", ResourceID: validID}}, - err: errors.New("invalid binding slot: unknown slot \"unknown\""), - }, - { - desc: "bind resources with wrong slot type", - configID: config.ID, - bindings: []bootstrap.BindingRequest{{Slot: "mqtt", Type: "channel", ResourceID: validID}}, - err: errors.New("invalid binding slot: slot \"mqtt\" expects \"client\", got \"channel\""), - }, - { - desc: "bind resources with resolver error", - configID: config.ID, - bindings: requested, - resolveErr: errors.New("resolve failed"), - err: errors.New("resolve failed"), - }, - { - desc: "bind resources with binding store retrieve error", - configID: config.ID, - bindings: requested, - snapshots: []bootstrap.BindingSnapshot{snapshot}, - retrieveErr: svcerr.ErrViewEntity, - err: svcerr.ErrViewEntity, - }, - { - desc: "bind resources with required slot not satisfied", - configID: config.ID, - bindings: requested, - snapshots: []bootstrap.BindingSnapshot{ - {ConfigID: config.ID, Slot: "mqtt", Type: "channel", ResourceID: validID}, - }, - err: errors.New("invalid binding slot: slot \"mqtt\" expects \"client\", got \"channel\""), - }, - { - desc: "bind resources with save error", - configID: config.ID, - bindings: requested, - snapshots: []bootstrap.BindingSnapshot{snapshot}, - saveErr: svcerr.ErrCreateEntity, - err: svcerr.ErrCreateEntity, - }, - { - desc: "bind resources successfully", - configID: config.ID, - bindings: requested, - snapshots: []bootstrap.BindingSnapshot{snapshot}, - }, - { - desc: "bind channel resource successfully", - configID: config.ID, - bindings: channelRequested, - snapshots: []bootstrap.BindingSnapshot{channelSnapshot}, - useChannel: true, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svc := newService() - activeCfg, activeProfile := cfg, profile - if tc.useChannel { - activeCfg = channelCfg - activeProfile = channelProfile - } - boot.On("RetrieveByID", context.Background(), domainID, tc.configID).Return(activeCfg, tc.cfgErr) - profileRepo.On("RetrieveByID", context.Background(), domainID, activeProfile.ID).Return(activeProfile, tc.prErr) - resolver.On("Resolve", context.Background(), mock.Anything).Return(tc.snapshots, tc.resolveErr) - bindingStore.On("Retrieve", context.Background(), tc.configID).Return([]bootstrap.BindingSnapshot{}, tc.retrieveErr) - bindingStore.On("Save", context.Background(), tc.configID, mock.Anything).Return(tc.saveErr) - - err := svc.BindResources(context.Background(), session, validToken, tc.configID, tc.bindings) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.err, err)) - }) - } -} - -func TestListBindings(t *testing.T) { - session := smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: validID} - - snapshots := []bootstrap.BindingSnapshot{ - {ConfigID: config.ID, Slot: "mqtt", Type: "client", ResourceID: validID, Snapshot: map[string]any{"id": validID}}, - } - - cases := []struct { - desc string - configID string - cfgErr error - bindings []bootstrap.BindingSnapshot - retrieveErr error - err error - }{ - { - desc: "list bindings with config not found", - configID: config.ID, - cfgErr: svcerr.ErrNotFound, - err: svcerr.ErrNotFound, - }, - { - desc: "list bindings with retrieve error", - configID: config.ID, - retrieveErr: svcerr.ErrViewEntity, - err: svcerr.ErrViewEntity, - }, - { - desc: "list bindings successfully", - configID: config.ID, - bindings: snapshots, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svc := newService() - boot.On("RetrieveByID", context.Background(), domainID, tc.configID).Return(config, tc.cfgErr) - bindingStore.On("Retrieve", context.Background(), tc.configID).Return(tc.bindings, tc.retrieveErr) - - got, err := svc.ListBindings(context.Background(), session, tc.configID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.err, err)) - if tc.err == nil { - assert.Len(t, got, len(tc.bindings), fmt.Sprintf("%s: expected %d bindings got %d\n", tc.desc, len(tc.bindings), len(got))) - } - }) - } -} - -func TestRefreshBindings(t *testing.T) { - session := smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: validID} - - profile := bootstrap.Profile{ - ID: testsutil.GenerateUUID(t), - DomainID: domainID, - Name: "refresh-profile", - BindingSlots: []bootstrap.BindingSlot{ - {Name: "mqtt", Type: "client", Required: true}, - }, - } - - cfg := bootstrap.Config{ - ID: config.ID, - DomainID: domainID, - ProfileID: profile.ID, - } - - existing := []bootstrap.BindingSnapshot{ - {ConfigID: config.ID, Slot: "mqtt", Type: "client", ResourceID: validID}, - } - - refreshed := []bootstrap.BindingSnapshot{ - {ConfigID: config.ID, Slot: "mqtt", Type: "client", ResourceID: validID, Snapshot: map[string]any{"id": validID}}, - } - - cases := []struct { - desc string - configID string - cfgErr error - prErr error - existing []bootstrap.BindingSnapshot - retrieveErr error - snapshots []bootstrap.BindingSnapshot - resolveErr error - saveErr error - err error - }{ - { - desc: "refresh bindings with config not found", - configID: config.ID, - cfgErr: svcerr.ErrNotFound, - err: svcerr.ErrNotFound, - }, - { - desc: "refresh bindings with profile not found", - configID: config.ID, - prErr: svcerr.ErrNotFound, - err: svcerr.ErrNotFound, - }, - { - desc: "refresh bindings with retrieve error", - configID: config.ID, - retrieveErr: svcerr.ErrViewEntity, - err: svcerr.ErrViewEntity, - }, - { - desc: "refresh bindings with no existing bindings is a no-op", - configID: config.ID, - }, - { - desc: "refresh bindings with resolver error", - configID: config.ID, - existing: existing, - resolveErr: errors.New("resolve failed"), - err: errors.New("resolve failed"), - }, - { - desc: "refresh bindings with required binding missing after refresh", - configID: config.ID, - existing: existing, - snapshots: []bootstrap.BindingSnapshot{ - {ConfigID: config.ID, Slot: "mqtt", Type: "channel", ResourceID: validID}, - }, - err: errors.New("invalid binding slot: slot \"mqtt\" expects \"client\", got \"channel\""), - }, - { - desc: "refresh bindings with save error", - configID: config.ID, - existing: existing, - snapshots: refreshed, - saveErr: svcerr.ErrCreateEntity, - err: svcerr.ErrCreateEntity, - }, - { - desc: "refresh bindings successfully", - configID: config.ID, - existing: existing, - snapshots: refreshed, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svc := newService() - boot.On("RetrieveByID", context.Background(), domainID, tc.configID).Return(cfg, tc.cfgErr) - profileRepo.On("RetrieveByID", context.Background(), domainID, profile.ID).Return(profile, tc.prErr) - bindingStore.On("Retrieve", context.Background(), tc.configID).Return(tc.existing, tc.retrieveErr) - resolver.On("Resolve", context.Background(), mock.Anything).Return(tc.snapshots, tc.resolveErr) - bindingStore.On("Save", context.Background(), tc.configID, mock.Anything).Return(tc.saveErr) - - err := svc.RefreshBindings(context.Background(), session, validToken, tc.configID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.err, err)) - }) - } -} diff --git a/bootstrap/status.go b/bootstrap/status.go deleted file mode 100644 index 9ef4a3c24..000000000 --- a/bootstrap/status.go +++ /dev/null @@ -1,101 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package bootstrap - -import ( - "encoding/json" - "strconv" - "strings" - - svcerr "github.com/absmach/magistrala/pkg/errors/service" -) - -// Status represents bootstrap enrollment availability. -type Status uint8 - -// Possible bootstrap enrollment statuses. -const ( - EnabledStatus Status = iota - DisabledStatus - // AllStatus is used for querying purposes to list configs irrespective - // of their status. It is never stored in the database. - AllStatus -) - -// String representation of bootstrap status values. -const ( - Disabled = "disabled" - Enabled = "enabled" - All = "all" - Unknown = "unknown" -) - -// Backward-compatible aliases kept while callers move off the old names. -const ( - Inactive = DisabledStatus - Active = EnabledStatus -) - -// String returns string representation of Status. -func (s Status) String() string { - switch s { - case DisabledStatus: - return Disabled - case EnabledStatus: - return Enabled - case AllStatus: - return All - default: - return Unknown - } -} - -// ToStatus converts a string or legacy numeric string value to Status. -func ToStatus(status string) (Status, error) { - switch strings.ToLower(status) { - case "", Enabled, "0": - return EnabledStatus, nil - case Disabled, "1": - return DisabledStatus, nil - case All: - return AllStatus, nil - } - return Status(0), svcerr.ErrInvalidStatus -} - -// MarshalJSON renders bootstrap status as a string literal. -func (s Status) MarshalJSON() ([]byte, error) { - return json.Marshal(s.String()) -} - -// UnmarshalJSON accepts both string and legacy numeric bootstrap statuses. -func (s *Status) UnmarshalJSON(data []byte) error { - if len(data) == 0 || string(data) == "null" { - return nil - } - - if data[0] != '"' { - var n int - if err := json.Unmarshal(data, &n); err != nil { - return err - } - parsed, err := ToStatus(strconv.Itoa(n)) - if err != nil { - return err - } - *s = parsed - return nil - } - - var status string - if err := json.Unmarshal(data, &status); err != nil { - return err - } - parsed, err := ToStatus(status) - if err != nil { - return err - } - *s = parsed - return nil -} diff --git a/bootstrap/tracing/doc.go b/bootstrap/tracing/doc.go deleted file mode 100644 index 7c6a64c5b..000000000 --- a/bootstrap/tracing/doc.go +++ /dev/null @@ -1,12 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package tracing provides tracing instrumentation for Magistrala Users service. -// -// This package provides tracing middleware for Magistrala Users service. -// It can be used to trace incoming requests and add tracing capabilities to -// Magistrala Users service. -// -// For more details about tracing instrumentation for Magistrala messaging refer -// to the documentation at https://magistrala.absmach.eu/docs/. -package tracing diff --git a/bootstrap/tracing/tracing.go b/bootstrap/tracing/tracing.go deleted file mode 100644 index c26207542..000000000 --- a/bootstrap/tracing/tracing.go +++ /dev/null @@ -1,198 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package tracing - -import ( - "context" - - "github.com/absmach/magistrala/bootstrap" - smqauthn "github.com/absmach/magistrala/pkg/authn" - "go.opentelemetry.io/otel/attribute" - "go.opentelemetry.io/otel/trace" -) - -var _ bootstrap.Service = (*tracingMiddleware)(nil) - -type tracingMiddleware struct { - tracer trace.Tracer - svc bootstrap.Service -} - -// New returns a new bootstrap service with tracing capabilities. -func New(svc bootstrap.Service, tracer trace.Tracer) bootstrap.Service { - return &tracingMiddleware{tracer, svc} -} - -// Add traces the "Add" operation of the wrapped bootstrap.Service. -func (tm *tracingMiddleware) Add(ctx context.Context, session smqauthn.Session, token string, cfg bootstrap.Config) (bootstrap.Config, error) { - ctx, span := tm.tracer.Start(ctx, "svc_register_user", trace.WithAttributes( - attribute.String("config_id", cfg.ID), - attribute.String("domain_id ", cfg.DomainID), - attribute.String("name", cfg.Name), - attribute.String("external_id", cfg.ExternalID), - attribute.String("content", cfg.Content), - attribute.String("status", cfg.Status.String()), - )) - defer span.End() - - return tm.svc.Add(ctx, session, token, cfg) -} - -// View traces the "View" operation of the wrapped bootstrap.Service. -func (tm *tracingMiddleware) View(ctx context.Context, session smqauthn.Session, id string) (bootstrap.Config, error) { - ctx, span := tm.tracer.Start(ctx, "svc_view_user", trace.WithAttributes( - attribute.String("id", id), - )) - defer span.End() - - return tm.svc.View(ctx, session, id) -} - -// Update traces the "Update" operation of the wrapped bootstrap.Service. -func (tm *tracingMiddleware) Update(ctx context.Context, session smqauthn.Session, cfg bootstrap.Config) error { - ctx, span := tm.tracer.Start(ctx, "svc_update_user", trace.WithAttributes( - attribute.String("name", cfg.Name), - attribute.String("content", cfg.Content), - attribute.String("config_id", cfg.ID), - attribute.String("domain_id ", cfg.DomainID), - )) - defer span.End() - - return tm.svc.Update(ctx, session, cfg) -} - -// UpdateCert traces the "UpdateCert" operation of the wrapped bootstrap.Service. -func (tm *tracingMiddleware) UpdateCert(ctx context.Context, session smqauthn.Session, id, clientCert, clientKey, caCert string) (bootstrap.Config, error) { - ctx, span := tm.tracer.Start(ctx, "svc_update_cert", trace.WithAttributes( - attribute.String("config_id", id), - )) - defer span.End() - - return tm.svc.UpdateCert(ctx, session, id, clientCert, clientKey, caCert) -} - -// List traces the "List" operation of the wrapped bootstrap.Service. -func (tm *tracingMiddleware) List(ctx context.Context, session smqauthn.Session, filter bootstrap.Filter, offset, limit uint64) (bootstrap.ConfigsPage, error) { - ctx, span := tm.tracer.Start(ctx, "svc_list_users", trace.WithAttributes( - attribute.Int64("offset", int64(offset)), - attribute.Int64("limit", int64(limit)), - )) - defer span.End() - - return tm.svc.List(ctx, session, filter, offset, limit) -} - -// Remove traces the "Remove" operation of the wrapped bootstrap.Service. -func (tm *tracingMiddleware) Remove(ctx context.Context, session smqauthn.Session, id string) error { - ctx, span := tm.tracer.Start(ctx, "svc_remove_user", trace.WithAttributes( - attribute.String("id", id), - )) - defer span.End() - - return tm.svc.Remove(ctx, session, id) -} - -// Bootstrap traces the "Bootstrap" operation of the wrapped bootstrap.Service. -func (tm *tracingMiddleware) Bootstrap(ctx context.Context, externalKey, externalID string, secure bool) (bootstrap.Config, error) { - ctx, span := tm.tracer.Start(ctx, "svc_bootstrap_user", trace.WithAttributes( - attribute.String("external_id", externalID), - attribute.Bool("secure", secure), - )) - defer span.End() - - return tm.svc.Bootstrap(ctx, externalKey, externalID, secure) -} - -func (tm *tracingMiddleware) EnableConfig(ctx context.Context, session smqauthn.Session, id string) (bootstrap.Config, error) { - ctx, span := tm.tracer.Start(ctx, "svc_enable_config", trace.WithAttributes( - attribute.String("id", id), - )) - defer span.End() - - return tm.svc.EnableConfig(ctx, session, id) -} - -func (tm *tracingMiddleware) DisableConfig(ctx context.Context, session smqauthn.Session, id string) (bootstrap.Config, error) { - ctx, span := tm.tracer.Start(ctx, "svc_disable_config", trace.WithAttributes( - attribute.String("id", id), - )) - defer span.End() - - return tm.svc.DisableConfig(ctx, session, id) -} - -func (tm *tracingMiddleware) CreateProfile(ctx context.Context, session smqauthn.Session, p bootstrap.Profile) (bootstrap.Profile, error) { - ctx, span := tm.tracer.Start(ctx, "svc_create_profile", trace.WithAttributes( - attribute.String("name", p.Name), - attribute.String("domain_id", p.DomainID), - )) - defer span.End() - return tm.svc.CreateProfile(ctx, session, p) -} - -func (tm *tracingMiddleware) ViewProfile(ctx context.Context, session smqauthn.Session, profileID string) (bootstrap.Profile, error) { - ctx, span := tm.tracer.Start(ctx, "svc_view_profile", trace.WithAttributes( - attribute.String("profile_id", profileID), - )) - defer span.End() - return tm.svc.ViewProfile(ctx, session, profileID) -} - -func (tm *tracingMiddleware) UpdateProfile(ctx context.Context, session smqauthn.Session, p bootstrap.Profile) (bootstrap.Profile, error) { - ctx, span := tm.tracer.Start(ctx, "svc_update_profile", trace.WithAttributes( - attribute.String("profile_id", p.ID), - )) - defer span.End() - return tm.svc.UpdateProfile(ctx, session, p) -} - -func (tm *tracingMiddleware) ListProfiles(ctx context.Context, session smqauthn.Session, offset, limit uint64, name string) (bootstrap.ProfilesPage, error) { - ctx, span := tm.tracer.Start(ctx, "svc_list_profiles", trace.WithAttributes( - attribute.Int64("offset", int64(offset)), - attribute.Int64("limit", int64(limit)), - )) - defer span.End() - return tm.svc.ListProfiles(ctx, session, offset, limit, name) -} - -func (tm *tracingMiddleware) DeleteProfile(ctx context.Context, session smqauthn.Session, profileID string) error { - ctx, span := tm.tracer.Start(ctx, "svc_delete_profile", trace.WithAttributes( - attribute.String("profile_id", profileID), - )) - defer span.End() - return tm.svc.DeleteProfile(ctx, session, profileID) -} - -func (tm *tracingMiddleware) AssignProfile(ctx context.Context, session smqauthn.Session, configID, profileID string) error { - ctx, span := tm.tracer.Start(ctx, "svc_assign_profile", trace.WithAttributes( - attribute.String("config_id", configID), - attribute.String("profile_id", profileID), - )) - defer span.End() - return tm.svc.AssignProfile(ctx, session, configID, profileID) -} - -func (tm *tracingMiddleware) BindResources(ctx context.Context, session smqauthn.Session, token, configID string, bindings []bootstrap.BindingRequest) error { - ctx, span := tm.tracer.Start(ctx, "svc_bind_resources", trace.WithAttributes( - attribute.String("config_id", configID), - )) - defer span.End() - return tm.svc.BindResources(ctx, session, token, configID, bindings) -} - -func (tm *tracingMiddleware) ListBindings(ctx context.Context, session smqauthn.Session, configID string) ([]bootstrap.BindingSnapshot, error) { - ctx, span := tm.tracer.Start(ctx, "svc_list_bindings", trace.WithAttributes( - attribute.String("config_id", configID), - )) - defer span.End() - return tm.svc.ListBindings(ctx, session, configID) -} - -func (tm *tracingMiddleware) RefreshBindings(ctx context.Context, session smqauthn.Session, token, configID string) error { - ctx, span := tm.tracer.Start(ctx, "svc_refresh_bindings", trace.WithAttributes( - attribute.String("config_id", configID), - )) - defer span.End() - return tm.svc.RefreshBindings(ctx, session, token, configID) -} diff --git a/certs/middleware/logging.go b/certs/middleware/logging.go index 02c3d6ebf..9f21cbcd8 100644 --- a/certs/middleware/logging.go +++ b/certs/middleware/logging.go @@ -20,7 +20,7 @@ type loggingMiddleware struct { svc certs.Service } -// LoggingMiddleware adds logging facilities to the core service. +// LoggingMiddleware adds logging facilities to the service. func LoggingMiddleware(svc certs.Service, logger *slog.Logger) certs.Service { return &loggingMiddleware{logger, svc} } diff --git a/certs/middleware/metrics.go b/certs/middleware/metrics.go index 84edc3edb..0dae1d693 100644 --- a/certs/middleware/metrics.go +++ b/certs/middleware/metrics.go @@ -20,7 +20,7 @@ type metricsMiddleware struct { svc certs.Service } -// MetricsMiddleware instruments core service by tracking request count and latency. +// MetricsMiddleware instruments service by tracking request count and latency. func MetricsMiddleware(svc certs.Service, counter metrics.Counter, latency metrics.Histogram) certs.Service { return &metricsMiddleware{ counter: counter, diff --git a/channels/README.md b/channels/README.md deleted file mode 100644 index 45cd9ba34..000000000 --- a/channels/README.md +++ /dev/null @@ -1,214 +0,0 @@ -# Channels - -The Channels service is a core component of Magistrala that manages communication channels between devices and applications. It handles channel creation, configuration, access control and message routing within the Magistrala ecosystem. - -## Configuration - -The service is configured using the following environment variables (unset variables use default values): - -| Variable | Description | Default | -| ------------------------- | --------------------------------------------- | ------------------------------ | -| `MG_CHANNELS_LOG_LEVEL` | Log level (debug, info, warn, error) | info | -| `MG_CHANNELS_HTTP_HOST` | HTTP host for Channels service | localhost | -| `MG_CHANNELS_HTTP_PORT` | HTTP port for Channels service | 9005 | -| `MG_CHANNELS_SERVER_CERT` | Path to PEM encoded server certificate | "" | -| `MG_CHANNELS_SERVER_KEY` | Path to PEM encoded server key file | "" | -| `MG_CHANNELS_GRPC_HOST` | gRPC host for Channels service | localhost | -| `MG_CHANNELS_GRPC_PORT` | gRPC port for Channels service | 7005 | -| `MG_CHANNELS_DB_HOST` | Database host address | localhost | -| `MG_CHANNELS_DB_PORT` | Database port | 5432 | -| `MG_CHANNELS_DB_USER` | Database user | magistrala | -| `MG_CHANNELS_DB_PASS` | Database password | magistrala | -| `MG_CHANNELS_DB_NAME` | Name of the database used by the service | channels | -| `MG_CHANNELS_DB_SSL_MODE` | Database connection SSL mode | disable | -| `MG_CHANNELS_CACHE_URL` | Cache database URL | | -| `MG_JAEGER_URL` | Jaeger tracing server URL | | -| `MG_SEND_TELEMETRY` | Send telemetry to Magistrala call-home server | true | - -## Features - -- **Channel Management**: Create, update, delete and list channels -- **Access Control**: Manage channel permissions and user access -- **Message Routing**: Route messages between connected devices and services -- **Channel Groups**: Organize channels into logical groups -- **Metadata Support**: Attach custom metadata to channels -- **Real-time Updates**: Live channel state synchronization - -## Architecture - -The service is built using: - -- **Go**: Core service implementation -- **gRPC**: Inter-service communication -- **PostgreSQL**: Primary data storage -- **Redis**: Caching and pub/sub messaging -- **Docker**: Containerized deployment - -### Channels Table - -| Column | Type | Description | -| ----------------- | ------------- | ----------------------------------------------------- | -| `id` | VARCHAR(36) | UUID of the channel (primary key) | -| `name` | VARCHAR(1024) | Human-readable name | -| `domain_id` | VARCHAR(36) | Domain to which the channel belongs | -| `parent_group_id` | VARCHAR(36) | Optional group parent | -| `tags` | TEXT[] | Array of tags | -| `metadata` | JSONB | Free-form structured metadata | -| `created_by` | VARCHAR(254) | User that created the channel | -| `created_at` | TIMESTAMPTZ | Timestamp of creation | -| `updated_at` | TIMESTAMPTZ | Timestamp of last update | -| `updated_by` | VARCHAR(254) | User that performed last update | -| `status` | SMALLINT | 0 = enabled, 1 = disabled | -| `route` | VARCHAR(36) | Optional route identifier unique within domain if set | - -### Connections Table - -| Column | Type | Description | -| ------------ | ----------- | ----------------------------------------------- | -| `channel_id` | VARCHAR(36) | Channel UUID | -| `domain_id` | VARCHAR(36) | Domain of channel and client | -| `client_id` | VARCHAR(36) | Client UUID | -| `type` | SMALLINT | Connection type: `1 = Publish`, `2 = Subscribe` | - -## Deployment - -The service is available as a Docker container. Refer to the Docker Compose section for the `channels` service in `docker-compose.yaml` for deployment configuration. - -To build and run locally: - -```bash -# download the latest version of the service -git clone https://github.com/absmach/magistrala -cd magistrala - -# compile the channels -make channels - -make install - -MG_CHANNELS_HTTP_HOST=localhost \ -MG_CHANNELS_HTTP_PORT=9005 \ -MG_CHANNELS_DB_HOST=localhost \ -MG_CHANNELS_DB_PORT=5432 \ -MG_CHANNELS_DB_USER=magistrala \MG_CHANNELS_DB_PASS=magistrala \MG_CHANNELS_DB_NAME=channels \ -$GOBIN/magistrala-channels -``` - -### Running the Service - -```bash -# Set environment variables -export MQ_CHANNELS_DB_HOST=localhost -export MQ_CHANNELS_DB_PORT=5432 - -# Run the service -go run cmd/main.go -``` - -### Docker Deployment - -```bash -docker run -p 8180:8180 magistrala/channels -``` - -## Testing - -```bash -# Run unit tests -go test ./... - -# Run integration tests -make test-integration -``` - -## Usage - -The Channels service supports the following operations: - -| Operation | Description | -| --------------- | -------------------------------------------- | -| `create` | Create a new channel | -| `list` | Retrieve all channels (paged) | -| `get` | Retrieve a single channel by ID | -| `update` | Update a channel’s name & metadata | -| `delete` | Permanently delete a channel | -| `enable` | Enable a previously disabled channel | -| `disable` | Disable an active channel | -| `set-parent` | Assign a parent group to a channel | -| `remove-parent` | Remove parent group from a channel | -| `connect` | Connect one or more clients to channels | -| `disconnect` | Disconnect one or more clients from channels | - -### Example: Create a Channel - -```bash -curl -X POST http://localhost:9005//channels \ - -H "Authorization: Bearer " \ - -H "Content-Type: application/json" \ - -d '{ - "name": "myChannel", - "metadata": { "location": "lab" }, - "route": "sensor-data", - "tags": ["sensor","edge"], - "status": "enabled" - }' -``` - -### Example: Connect Clients & Channels - -```bash -curl -X POST http://localhost:9005//channels/connect \ - -H "Authorization: Bearer " \ - -H "Content-Type: application/json" \ - -d '{ - "channel_ids": ["", ""], - "client_ids": ["", ""], - "types": ["publish", "subscribe"] - }' -``` - -### Example: Disconnect Clients from a Channel - -```bash -curl -X POST http://localhost:9005//channels/disconnect \ - -H "Authorization: Bearer " \ - -H "Content-Type: application/json" \ - -d '{ - "channel_ids": [""], - "client_ids": [""], - "types": ["publish"] - }' -``` - -## Best Practices - -- Use tags and metadata to manage and categorize channels (e.g., environment, region, purpose). -- Assign `route` thoughtfully when channels need a predictable identifier. -- Keep channel hierarchies shallow for easier navigation (avoid deep nesting unless required). -- Use `disable` rather than immediate delete when you want to suspend a channel temporarily. -- Clean up unused connections: regularly review which clients are connected to channels and remove stale links. -- Enforce minimal privileges: only allow clients to connect to channels they truly need. -- Monitoring: use the `/health` endpoint and version metadata for service stability. - -## Versioning & Health Check - -The Channels service exposes a `/health` endpoint to provide operational status and version info. - -### Health Check Request - -```bash -curl -X GET http://localhost:9005/health \ - -H "accept: application/health+json" -``` - -### Example Response - -```json -{ - "status": "pass", - "version": "0.18.0", - "commit": "", - "description": "channels service", - "build_time": "2025-11-19T..." -} -``` diff --git a/channels/api/grpc/client.go b/channels/api/grpc/client.go deleted file mode 100644 index 7ce1ac02d..000000000 --- a/channels/api/grpc/client.go +++ /dev/null @@ -1,224 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package grpc - -import ( - "context" - "fmt" - "time" - - grpcChannelsV1 "github.com/absmach/magistrala/api/grpc/channels/v1" - grpcCommonV1 "github.com/absmach/magistrala/api/grpc/common/v1" - "github.com/absmach/magistrala/pkg/connections" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/go-kit/kit/endpoint" - kitgrpc "github.com/go-kit/kit/transport/grpc" - "google.golang.org/grpc" - "google.golang.org/grpc/codes" - "google.golang.org/grpc/status" -) - -const svcName = "channels.v1.ChannelsService" - -var _ grpcChannelsV1.ChannelsServiceClient = (*grpcClient)(nil) - -type grpcClient struct { - timeout time.Duration - authorize endpoint.Endpoint - removeClientConnections endpoint.Endpoint - unsetParentGroupFromChannels endpoint.Endpoint - retrieveEntity endpoint.Endpoint - retrieveIDByRoute endpoint.Endpoint -} - -// NewClient returns new gRPC client instance. -func NewClient(conn *grpc.ClientConn, timeout time.Duration) grpcChannelsV1.ChannelsServiceClient { - return &grpcClient{ - authorize: kitgrpc.NewClient( - conn, - svcName, - "Authorize", - encodeAuthorizeRequest, - decodeAuthorizeResponse, - grpcChannelsV1.AuthzRes{}, - ).Endpoint(), - removeClientConnections: kitgrpc.NewClient( - conn, - svcName, - "RemoveClientConnections", - encodeRemoveClientConnectionsRequest, - decodeRemoveClientConnectionsResponse, - grpcChannelsV1.RemoveClientConnectionsRes{}, - ).Endpoint(), - unsetParentGroupFromChannels: kitgrpc.NewClient( - conn, - svcName, - "UnsetParentGroupFromChannels", - encodeUnsetParentGroupFromChannelsRequest, - decodeUnsetParentGroupFromChannelsResponse, - grpcChannelsV1.UnsetParentGroupFromChannelsRes{}, - ).Endpoint(), - retrieveEntity: kitgrpc.NewClient( - conn, - svcName, - "RetrieveEntity", - encodeRetrieveEntityRequest, - decodeRetrieveEntityResponse, - grpcCommonV1.RetrieveEntityRes{}, - ).Endpoint(), - retrieveIDByRoute: kitgrpc.NewClient( - conn, - svcName, - "RetrieveIDByRoute", - encodeRetrieveIDByRouteRequest, - decodeRetrieveIDByRouteResponse, - grpcCommonV1.RetrieveEntityRes{}, - ).Endpoint(), - timeout: timeout, - } -} - -func (client grpcClient) Authorize(ctx context.Context, req *grpcChannelsV1.AuthzReq, _ ...grpc.CallOption) (r *grpcChannelsV1.AuthzRes, err error) { - ctx, cancel := context.WithTimeout(ctx, client.timeout) - defer cancel() - - res, err := client.authorize(ctx, authorizeReq{ - domainID: req.GetDomainId(), - clientID: req.GetClientId(), - clientType: req.GetClientType(), - channelID: req.GetChannelId(), - connType: connections.ConnType(req.GetType()), - }) - if err != nil { - return &grpcChannelsV1.AuthzRes{}, decodeError(err) - } - - ar := res.(authorizeRes) - - return &grpcChannelsV1.AuthzRes{Authorized: ar.authorized}, nil -} - -func encodeAuthorizeRequest(_ context.Context, grpcReq any) (any, error) { - req := grpcReq.(authorizeReq) - - return &grpcChannelsV1.AuthzReq{ - DomainId: req.domainID, - ClientId: req.clientID, - ClientType: req.clientType, - ChannelId: req.channelID, - Type: uint32(req.connType), - }, nil -} - -func decodeAuthorizeResponse(_ context.Context, grpcRes any) (any, error) { - res := grpcRes.(*grpcChannelsV1.AuthzRes) - - return authorizeRes{authorized: res.GetAuthorized()}, nil -} - -func (client grpcClient) RemoveClientConnections(ctx context.Context, req *grpcChannelsV1.RemoveClientConnectionsReq, _ ...grpc.CallOption) (r *grpcChannelsV1.RemoveClientConnectionsRes, err error) { - ctx, cancel := context.WithTimeout(ctx, client.timeout) - defer cancel() - - if _, err := client.removeClientConnections(ctx, req); err != nil { - return &grpcChannelsV1.RemoveClientConnectionsRes{}, decodeError(err) - } - - return &grpcChannelsV1.RemoveClientConnectionsRes{}, nil -} - -func encodeRemoveClientConnectionsRequest(_ context.Context, grpcReq any) (any, error) { - return grpcReq.(*grpcChannelsV1.RemoveClientConnectionsReq), nil -} - -func decodeRemoveClientConnectionsResponse(_ context.Context, grpcRes any) (any, error) { - return grpcRes.(*grpcChannelsV1.RemoveClientConnectionsRes), nil -} - -func (client grpcClient) UnsetParentGroupFromChannels(ctx context.Context, req *grpcChannelsV1.UnsetParentGroupFromChannelsReq, _ ...grpc.CallOption) (r *grpcChannelsV1.UnsetParentGroupFromChannelsRes, err error) { - ctx, cancel := context.WithTimeout(ctx, client.timeout) - defer cancel() - - if _, err := client.unsetParentGroupFromChannels(ctx, req); err != nil { - return &grpcChannelsV1.UnsetParentGroupFromChannelsRes{}, decodeError(err) - } - - return &grpcChannelsV1.UnsetParentGroupFromChannelsRes{}, nil -} - -func encodeUnsetParentGroupFromChannelsRequest(_ context.Context, grpcReq any) (any, error) { - return grpcReq.(*grpcChannelsV1.UnsetParentGroupFromChannelsReq), nil -} - -func decodeUnsetParentGroupFromChannelsResponse(_ context.Context, grpcRes any) (any, error) { - return grpcRes.(*grpcChannelsV1.UnsetParentGroupFromChannelsRes), nil -} - -func (client grpcClient) RetrieveEntity(ctx context.Context, req *grpcCommonV1.RetrieveEntityReq, _ ...grpc.CallOption) (r *grpcCommonV1.RetrieveEntityRes, err error) { - ctx, cancel := context.WithTimeout(ctx, client.timeout) - defer cancel() - - res, err := client.retrieveEntity(ctx, req) - if err != nil { - return &grpcCommonV1.RetrieveEntityRes{}, decodeError(err) - } - - return res.(*grpcCommonV1.RetrieveEntityRes), nil -} - -func encodeRetrieveEntityRequest(_ context.Context, grpcReq any) (any, error) { - return grpcReq.(*grpcCommonV1.RetrieveEntityReq), nil -} - -func decodeRetrieveEntityResponse(_ context.Context, grpcRes any) (any, error) { - return grpcRes.(*grpcCommonV1.RetrieveEntityRes), nil -} - -func (client grpcClient) RetrieveIDByRoute(ctx context.Context, req *grpcCommonV1.RetrieveIDByRouteReq, _ ...grpc.CallOption) (r *grpcCommonV1.RetrieveEntityRes, err error) { - ctx, cancel := context.WithTimeout(ctx, client.timeout) - defer cancel() - - res, err := client.retrieveIDByRoute(ctx, req) - if err != nil { - return &grpcCommonV1.RetrieveEntityRes{}, decodeError(err) - } - - return res.(*grpcCommonV1.RetrieveEntityRes), nil -} - -func encodeRetrieveIDByRouteRequest(_ context.Context, grpcReq any) (any, error) { - return grpcReq.(*grpcCommonV1.RetrieveIDByRouteReq), nil -} - -func decodeRetrieveIDByRouteResponse(_ context.Context, grpcRes any) (any, error) { - return grpcRes.(*grpcCommonV1.RetrieveEntityRes), nil -} - -func decodeError(err error) error { - if st, ok := status.FromError(err); ok { - switch st.Code() { - case codes.Unauthenticated: - return errors.Wrap(svcerr.ErrAuthentication, errors.New(st.Message())) - case codes.PermissionDenied: - return errors.Wrap(svcerr.ErrAuthorization, errors.New(st.Message())) - case codes.InvalidArgument: - return errors.Wrap(errors.ErrMalformedEntity, errors.New(st.Message())) - case codes.FailedPrecondition: - return errors.Wrap(errors.ErrMalformedEntity, errors.New(st.Message())) - case codes.NotFound: - return errors.Wrap(svcerr.ErrNotFound, errors.New(st.Message())) - case codes.AlreadyExists: - return errors.Wrap(svcerr.ErrConflict, errors.New(st.Message())) - case codes.OK: - if msg := st.Message(); msg != "" { - return errors.Wrap(errors.ErrUnidentified, errors.New(msg)) - } - return nil - default: - return errors.Wrap(fmt.Errorf("unexpected gRPC status: %s (status code:%v)", st.Code().String(), st.Code()), errors.New(st.Message())) - } - } - return err -} diff --git a/channels/api/grpc/doc.go b/channels/api/grpc/doc.go deleted file mode 100644 index 20956ee50..000000000 --- a/channels/api/grpc/doc.go +++ /dev/null @@ -1,5 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package grpc contains implementation of Auth service gRPC API. -package grpc diff --git a/channels/api/grpc/endpoint.go b/channels/api/grpc/endpoint.go deleted file mode 100644 index 2f944e161..000000000 --- a/channels/api/grpc/endpoint.go +++ /dev/null @@ -1,85 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package grpc - -import ( - "context" - - ch "github.com/absmach/magistrala/channels" - channels "github.com/absmach/magistrala/channels/private" - "github.com/go-kit/kit/endpoint" -) - -func authorizeEndpoint(svc channels.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(authorizeReq) - if err := req.validate(); err != nil { - return authorizeRes{}, err - } - - if err := svc.Authorize(ctx, ch.AuthzReq{ - DomainID: req.domainID, - ClientID: req.clientID, - ClientType: req.clientType, - ChannelID: req.channelID, - Type: req.connType, - }); err != nil { - return authorizeRes{}, err - } - - return authorizeRes{authorized: true}, nil - } -} - -func removeClientConnectionsEndpoint(svc channels.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(removeClientConnectionsReq) - - if err := svc.RemoveClientConnections(ctx, req.clientID); err != nil { - return removeClientConnectionsRes{}, err - } - - return removeClientConnectionsRes{}, nil - } -} - -func unsetParentGroupFromChannelsEndpoint(svc channels.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(unsetParentGroupFromChannelsReq) - - if err := svc.UnsetParentGroupFromChannels(ctx, req.parentGroupID); err != nil { - return unsetParentGroupFromChannelsRes{}, err - } - - return unsetParentGroupFromChannelsRes{}, nil - } -} - -func retrieveEntityEndpoint(svc channels.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(retrieveEntityReq) - channel, err := svc.RetrieveByID(ctx, req.Id) - if err != nil { - return retrieveEntityRes{}, err - } - - return retrieveEntityRes{id: channel.ID, domain: channel.Domain, parentGroup: channel.ParentGroup, status: uint8(channel.Status)}, nil - } -} - -func retrieveIDByRouteEndpoint(svc channels.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(retrieveIDByRouteReq) - if err := req.validate(); err != nil { - return retrieveIDByRouteRes{}, err - } - - id, err := svc.RetrieveIDByRoute(ctx, req.route, req.domainID) - if err != nil { - return retrieveIDByRouteRes{}, err - } - - return retrieveIDByRouteRes{id: id}, nil - } -} diff --git a/channels/api/grpc/endpoint_test.go b/channels/api/grpc/endpoint_test.go deleted file mode 100644 index c050b7986..000000000 --- a/channels/api/grpc/endpoint_test.go +++ /dev/null @@ -1,345 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package grpc_test - -import ( - "context" - "fmt" - "net" - "testing" - "time" - - grpcChannelsV1 "github.com/absmach/magistrala/api/grpc/channels/v1" - grpcCommonV1 "github.com/absmach/magistrala/api/grpc/common/v1" - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/channels" - ch "github.com/absmach/magistrala/channels" - grpcapi "github.com/absmach/magistrala/channels/api/grpc" - "github.com/absmach/magistrala/channels/private/mocks" - "github.com/absmach/magistrala/internal/testsutil" - "github.com/absmach/magistrala/pkg/connections" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/pkg/policies" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" - "google.golang.org/grpc" - "google.golang.org/grpc/codes" - "google.golang.org/grpc/credentials/insecure" -) - -const port = 7005 - -var ( - validID = testsutil.GenerateUUID(&testing.T{}) - validChannel = ch.Channel{ - ID: validID, - Domain: testsutil.GenerateUUID(&testing.T{}), - Status: channels.EnabledStatus, - } -) - -func startGRPCServer(svc *mocks.Service, port int) *grpc.Server { - listener, err := net.Listen("tcp", fmt.Sprintf(":%d", port)) - if err != nil { - panic(fmt.Sprintf("failed to obtain port: %s", err)) - } - server := grpc.NewServer() - grpcChannelsV1.RegisterChannelsServiceServer(server, grpcapi.NewServer(svc)) - go func() { - if err := server.Serve(listener); err != nil { - panic(fmt.Sprintf("failed to serve: %s", err)) - } - }() - return server -} - -func TestAuthorize(t *testing.T) { - svc := new(mocks.Service) - server := startGRPCServer(svc, port) - defer server.GracefulStop() - authAddr := fmt.Sprintf("localhost:%d", port) - conn, _ := grpc.NewClient(authAddr, grpc.WithTransportCredentials(insecure.NewCredentials())) - client := grpcapi.NewClient(conn, time.Second) - - cases := []struct { - desc string - domainID string - clientID string - clientType string - channelID string - connType connections.ConnType - err error - authzErr error - res *grpcChannelsV1.AuthzRes - code codes.Code - }{ - { - desc: "authorize successfully", - domainID: validID, - clientID: validID, - clientType: policies.UserType, - channelID: validID, - connType: connections.Publish, - res: &grpcChannelsV1.AuthzRes{Authorized: true}, - err: nil, - }, - { - desc: "authorize with authorization error", - domainID: validID, - clientID: validID, - clientType: policies.UserType, - channelID: validID, - connType: connections.Publish, - res: &grpcChannelsV1.AuthzRes{Authorized: false}, - authzErr: svcerr.ErrAuthorization, - err: svcerr.ErrAuthorization, - }, - { - desc: "authorize withnot found error", - domainID: validID, - clientID: validID, - clientType: policies.UserType, - channelID: validID, - connType: connections.Publish, - res: &grpcChannelsV1.AuthzRes{Authorized: false}, - authzErr: svcerr.ErrNotFound, - err: svcerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - authReq := ch.AuthzReq{ - DomainID: tc.domainID, - ClientID: tc.clientID, - ClientType: tc.clientType, - ChannelID: tc.channelID, - Type: tc.connType, - } - svcCall := svc.On("Authorize", mock.Anything, authReq).Return(tc.authzErr) - res, err := client.Authorize(context.Background(), &grpcChannelsV1.AuthzReq{ - DomainId: tc.domainID, - ClientId: tc.clientID, - ClientType: tc.clientType, - ChannelId: tc.channelID, - Type: uint32(tc.connType), - }) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s", tc.desc, tc.err, err)) - assert.Equal(t, tc.res, res, fmt.Sprintf("%s: expected %s got %s", tc.desc, tc.res, res)) - svcCall.Unset() - }) - } -} - -func TestRemoveClientConnections(t *testing.T) { - svc := new(mocks.Service) - server := startGRPCServer(svc, port) - defer server.GracefulStop() - authAddr := fmt.Sprintf("localhost:%d", port) - conn, _ := grpc.NewClient(authAddr, grpc.WithTransportCredentials(insecure.NewCredentials())) - client := grpcapi.NewClient(conn, time.Second) - - cases := []struct { - desc string - clientID string - err error - code codes.Code - }{ - { - desc: "remove client connections successfully", - clientID: validID, - err: nil, - }, - { - desc: "remove client connections with error", - clientID: validID, - err: svcerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("RemoveClientConnections", mock.Anything, tc.clientID).Return(tc.err) - res, err := client.RemoveClientConnections(context.Background(), &grpcChannelsV1.RemoveClientConnectionsReq{ - ClientId: tc.clientID, - }) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s", tc.desc, tc.err, err)) - assert.Equal(t, &grpcChannelsV1.RemoveClientConnectionsRes{}, res) - svcCall.Unset() - }) - } -} - -func TestUnsetParentGroupFromChannelsEndpoint(t *testing.T) { - svc := new(mocks.Service) - server := startGRPCServer(svc, port) - defer server.GracefulStop() - authAddr := fmt.Sprintf("localhost:%d", port) - conn, _ := grpc.NewClient(authAddr, grpc.WithTransportCredentials(insecure.NewCredentials())) - client := grpcapi.NewClient(conn, time.Second) - - cases := []struct { - desc string - parentGroupID string - err error - code codes.Code - }{ - { - desc: "unset parent group from channels successfully", - parentGroupID: validID, - err: nil, - }, - { - desc: "unset parent group from channels authorization error", - parentGroupID: validID, - err: svcerr.ErrAuthorization, - }, - { - desc: "unset parent group from channels with not found error", - parentGroupID: validID, - err: svcerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("UnsetParentGroupFromChannels", mock.Anything, tc.parentGroupID).Return(tc.err) - res, err := client.UnsetParentGroupFromChannels(context.Background(), &grpcChannelsV1.UnsetParentGroupFromChannelsReq{ - ParentGroupId: tc.parentGroupID, - }) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s", tc.desc, tc.err, err)) - assert.Equal(t, &grpcChannelsV1.UnsetParentGroupFromChannelsRes{}, res) - svcCall.Unset() - }) - } -} - -func TestRetrieveEntity(t *testing.T) { - svc := new(mocks.Service) - server := startGRPCServer(svc, port) - defer server.GracefulStop() - authAddr := fmt.Sprintf("localhost:%d", port) - conn, _ := grpc.NewClient(authAddr, grpc.WithTransportCredentials(insecure.NewCredentials())) - client := grpcapi.NewClient(conn, time.Second) - - cases := []struct { - desc string - id string - svcRes ch.Channel - resp *grpcCommonV1.RetrieveEntityRes - code codes.Code - err error - }{ - { - desc: "retrieve entity successfully", - id: validID, - svcRes: validChannel, - resp: &grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Id: validChannel.ID, - DomainId: validChannel.Domain, - ParentGroupId: validChannel.ParentGroup, - Status: uint32(validChannel.Status), - }, - }, - err: nil, - }, - { - desc: "retrieve entity with error", - id: validID, - resp: &grpcCommonV1.RetrieveEntityRes{}, - err: svcerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("RetrieveByID", mock.Anything, tc.id).Return(tc.svcRes, tc.err) - res, err := client.RetrieveEntity(context.Background(), &grpcCommonV1.RetrieveEntityReq{ - Id: tc.id, - }) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp.Entity, res.Entity) - svcCall.Unset() - }) - } -} - -func TestRetrieveIDByRoute(t *testing.T) { - svc := new(mocks.Service) - server := startGRPCServer(svc, port) - defer server.GracefulStop() - authAddr := fmt.Sprintf("localhost:%d", port) - conn, _ := grpc.NewClient(authAddr, grpc.WithTransportCredentials(insecure.NewCredentials())) - client := grpcapi.NewClient(conn, time.Second) - - validRoute := "validRoute" - domainID := testsutil.GenerateUUID(t) - - cases := []struct { - desc string - retrieveReq *grpcCommonV1.RetrieveIDByRouteReq - svcRes string - svcErr error - retrieveRes *grpcCommonV1.RetrieveEntityRes - err error - }{ - { - desc: "retrieve entity by route successfully", - retrieveReq: &grpcCommonV1.RetrieveIDByRouteReq{ - Route: validRoute, - DomainId: domainID, - }, - svcRes: validID, - retrieveRes: &grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Id: validID, - }, - }, - err: nil, - }, - { - desc: "retrieve entity by route with empty route", - retrieveReq: &grpcCommonV1.RetrieveIDByRouteReq{ - Route: "", - DomainId: domainID, - }, - svcRes: "", - retrieveRes: &grpcCommonV1.RetrieveEntityRes{}, - err: apiutil.ErrMissingRoute, - }, - { - desc: "retrieve entity by route with empty domain ID", - retrieveReq: &grpcCommonV1.RetrieveIDByRouteReq{ - Route: validRoute, - DomainId: "", - }, - svcRes: "", - retrieveRes: &grpcCommonV1.RetrieveEntityRes{}, - err: apiutil.ErrMissingDomainID, - }, - { - desc: "retrieve entity by route with invalid route", - retrieveReq: &grpcCommonV1.RetrieveIDByRouteReq{ - Route: "invalidRoute", - DomainId: domainID, - }, - svcRes: "", - svcErr: svcerr.ErrNotFound, - retrieveRes: &grpcCommonV1.RetrieveEntityRes{}, - err: svcerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("RetrieveIDByRoute", mock.Anything, tc.retrieveReq.Route, tc.retrieveReq.DomainId).Return(tc.svcRes, tc.svcErr) - res, err := client.RetrieveIDByRoute(context.Background(), tc.retrieveReq) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s", tc.desc, tc.err, err)) - assert.Equal(t, tc.retrieveRes.Entity, res.Entity) - svcCall.Unset() - }) - } -} diff --git a/channels/api/grpc/request.go b/channels/api/grpc/request.go deleted file mode 100644 index e4db109ad..000000000 --- a/channels/api/grpc/request.go +++ /dev/null @@ -1,56 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package grpc - -import ( - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/pkg/connections" - "github.com/absmach/magistrala/pkg/errors" - "github.com/absmach/magistrala/pkg/policies" -) - -var errDomainID = errors.New("domain id required for users") - -type authorizeReq struct { - domainID string - channelID string - clientID string - clientType string - connType connections.ConnType -} - -func (req authorizeReq) validate() error { - if req.clientType == policies.UserType && req.domainID == "" { - return errDomainID - } - return nil -} - -type removeClientConnectionsReq struct { - clientID string -} - -type unsetParentGroupFromChannelsReq struct { - parentGroupID string -} - -type retrieveEntityReq struct { - Id string -} - -type retrieveIDByRouteReq struct { - route string - domainID string -} - -func (req retrieveIDByRouteReq) validate() error { - if req.route == "" { - return apiutil.ErrMissingRoute - } - if req.domainID == "" { - return apiutil.ErrMissingDomainID - } - - return nil -} diff --git a/channels/api/grpc/responses.go b/channels/api/grpc/responses.go deleted file mode 100644 index 3aadb9d71..000000000 --- a/channels/api/grpc/responses.go +++ /dev/null @@ -1,25 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package grpc - -type authorizeRes struct { - authorized bool -} - -type removeClientConnectionsRes struct{} - -type unsetParentGroupFromChannelsRes struct{} - -type channelBasic struct { - id string - domain string - parentGroup string - status uint8 -} - -type retrieveEntityRes channelBasic - -type retrieveIDByRouteRes struct { - id string -} diff --git a/channels/api/grpc/server.go b/channels/api/grpc/server.go deleted file mode 100644 index a285f7558..000000000 --- a/channels/api/grpc/server.go +++ /dev/null @@ -1,211 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package grpc - -import ( - "context" - - grpcChannelsV1 "github.com/absmach/magistrala/api/grpc/channels/v1" - grpcCommonV1 "github.com/absmach/magistrala/api/grpc/common/v1" - apiutil "github.com/absmach/magistrala/api/http/util" - smqauth "github.com/absmach/magistrala/auth" - channels "github.com/absmach/magistrala/channels/private" - "github.com/absmach/magistrala/pkg/connections" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - kitgrpc "github.com/go-kit/kit/transport/grpc" - "google.golang.org/grpc/codes" - "google.golang.org/grpc/status" -) - -var _ grpcChannelsV1.ChannelsServiceServer = (*grpcServer)(nil) - -type grpcServer struct { - grpcChannelsV1.UnimplementedChannelsServiceServer - authorize kitgrpc.Handler - removeClientConnections kitgrpc.Handler - unsetParentGroupFromChannels kitgrpc.Handler - retrieveEntity kitgrpc.Handler - retrieveIDByRoute kitgrpc.Handler -} - -// NewServer returns new AuthServiceServer instance. -func NewServer(svc channels.Service) grpcChannelsV1.ChannelsServiceServer { - return &grpcServer{ - authorize: kitgrpc.NewServer( - authorizeEndpoint(svc), - decodeAuthorizeRequest, - encodeAuthorizeResponse, - ), - removeClientConnections: kitgrpc.NewServer( - removeClientConnectionsEndpoint(svc), - decodeRemoveClientConnectionsRequest, - encodeRemoveClientConnectionsResponse, - ), - unsetParentGroupFromChannels: kitgrpc.NewServer( - unsetParentGroupFromChannelsEndpoint(svc), - decodeUnsetParentGroupFromChannelsRequest, - encodeUnsetParentGroupFromChannelsResponse, - ), - retrieveEntity: kitgrpc.NewServer( - retrieveEntityEndpoint(svc), - decodeRetrieveEntityRequest, - encodeRetrieveEntityResponse, - ), - retrieveIDByRoute: kitgrpc.NewServer( - retrieveIDByRouteEndpoint(svc), - decodeRetrieveIDByRouteRequest, - encodeRetrieveIDByRouteResponse, - ), - } -} - -func (s *grpcServer) Authorize(ctx context.Context, req *grpcChannelsV1.AuthzReq) (*grpcChannelsV1.AuthzRes, error) { - _, res, err := s.authorize.ServeGRPC(ctx, req) - if err != nil { - return nil, encodeError(err) - } - return res.(*grpcChannelsV1.AuthzRes), nil -} - -func decodeAuthorizeRequest(_ context.Context, grpcReq any) (any, error) { - req := grpcReq.(*grpcChannelsV1.AuthzReq) - - connType := connections.ConnType(req.GetType()) - if err := connections.CheckConnType(connType); err != nil { - return nil, err - } - return authorizeReq{ - domainID: req.GetDomainId(), - clientID: req.GetClientId(), - clientType: req.GetClientType(), - channelID: req.GetChannelId(), - connType: connType, - }, nil -} - -func encodeAuthorizeResponse(_ context.Context, grpcRes any) (any, error) { - res := grpcRes.(authorizeRes) - return &grpcChannelsV1.AuthzRes{Authorized: res.authorized}, nil -} - -func (s *grpcServer) RemoveClientConnections(ctx context.Context, req *grpcChannelsV1.RemoveClientConnectionsReq) (*grpcChannelsV1.RemoveClientConnectionsRes, error) { - _, res, err := s.removeClientConnections.ServeGRPC(ctx, req) - if err != nil { - return nil, encodeError(err) - } - return res.(*grpcChannelsV1.RemoveClientConnectionsRes), nil -} - -func decodeRemoveClientConnectionsRequest(_ context.Context, grpcReq any) (any, error) { - req := grpcReq.(*grpcChannelsV1.RemoveClientConnectionsReq) - - return removeClientConnectionsReq{ - clientID: req.GetClientId(), - }, nil -} - -func encodeRemoveClientConnectionsResponse(_ context.Context, grpcRes any) (any, error) { - _ = grpcRes.(removeClientConnectionsRes) - return &grpcChannelsV1.RemoveClientConnectionsRes{}, nil -} - -func (s *grpcServer) UnsetParentGroupFromChannels(ctx context.Context, req *grpcChannelsV1.UnsetParentGroupFromChannelsReq) (*grpcChannelsV1.UnsetParentGroupFromChannelsRes, error) { - _, res, err := s.unsetParentGroupFromChannels.ServeGRPC(ctx, req) - if err != nil { - return nil, encodeError(err) - } - return res.(*grpcChannelsV1.UnsetParentGroupFromChannelsRes), nil -} - -func decodeUnsetParentGroupFromChannelsRequest(_ context.Context, grpcReq any) (any, error) { - req := grpcReq.(*grpcChannelsV1.UnsetParentGroupFromChannelsReq) - - return unsetParentGroupFromChannelsReq{ - parentGroupID: req.GetParentGroupId(), - }, nil -} - -func encodeUnsetParentGroupFromChannelsResponse(_ context.Context, grpcRes any) (any, error) { - _ = grpcRes.(unsetParentGroupFromChannelsRes) - return &grpcChannelsV1.UnsetParentGroupFromChannelsRes{}, nil -} - -func (s *grpcServer) RetrieveEntity(ctx context.Context, req *grpcCommonV1.RetrieveEntityReq) (*grpcCommonV1.RetrieveEntityRes, error) { - _, res, err := s.retrieveEntity.ServeGRPC(ctx, req) - if err != nil { - return nil, encodeError(err) - } - return res.(*grpcCommonV1.RetrieveEntityRes), nil -} - -func decodeRetrieveEntityRequest(_ context.Context, grpcReq any) (any, error) { - req := grpcReq.(*grpcCommonV1.RetrieveEntityReq) - return retrieveEntityReq{ - Id: req.GetId(), - }, nil -} - -func encodeRetrieveEntityResponse(_ context.Context, grpcRes any) (any, error) { - res := grpcRes.(retrieveEntityRes) - - return &grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Id: res.id, - DomainId: res.domain, - ParentGroupId: res.parentGroup, - Status: uint32(res.status), - }, - }, nil -} - -func decodeRetrieveIDByRouteRequest(_ context.Context, grpcReq any) (any, error) { - req := grpcReq.(*grpcCommonV1.RetrieveIDByRouteReq) - return retrieveIDByRouteReq{ - route: req.GetRoute(), - domainID: req.GetDomainId(), - }, nil -} - -func encodeRetrieveIDByRouteResponse(_ context.Context, grpcRes any) (any, error) { - res := grpcRes.(retrieveIDByRouteRes) - - return &grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Id: res.id, - }, - }, nil -} - -func (s *grpcServer) RetrieveIDByRoute(ctx context.Context, req *grpcCommonV1.RetrieveIDByRouteReq) (*grpcCommonV1.RetrieveEntityRes, error) { - _, res, err := s.retrieveIDByRoute.ServeGRPC(ctx, req) - if err != nil { - return nil, encodeError(err) - } - return res.(*grpcCommonV1.RetrieveEntityRes), nil -} - -func encodeError(err error) error { - switch { - case errors.Contains(err, nil): - return nil - case errors.Contains(err, errors.ErrMalformedEntity), - err == apiutil.ErrInvalidAuthKey, - err == apiutil.ErrMissingID, - err == apiutil.ErrMissingMemberType, - err == apiutil.ErrMissingPolicySub, - err == apiutil.ErrMissingPolicyObj, - err == apiutil.ErrMalformedPolicyAct: - return status.Error(codes.InvalidArgument, err.Error()) - case errors.Contains(err, svcerr.ErrAuthentication), - errors.Contains(err, smqauth.ErrKeyExpired), - err == apiutil.ErrMissingEmail, - err == apiutil.ErrBearerToken: - return status.Error(codes.Unauthenticated, err.Error()) - case errors.Contains(err, svcerr.ErrAuthorization): - return status.Error(codes.PermissionDenied, err.Error()) - default: - return status.Error(codes.Internal, err.Error()) - } -} diff --git a/channels/api/http/decode.go b/channels/api/http/decode.go deleted file mode 100644 index 481c2f2ad..000000000 --- a/channels/api/http/decode.go +++ /dev/null @@ -1,329 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package http - -import ( - "context" - "encoding/json" - "net/http" - "strings" - "time" - - api "github.com/absmach/magistrala/api/http" - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/channels" - "github.com/absmach/magistrala/internal/nullable" - "github.com/absmach/magistrala/pkg/errors" - "github.com/go-chi/chi/v5" -) - -func decodeViewChannel(_ context.Context, r *http.Request) (any, error) { - roles, err := apiutil.ReadBoolQuery(r, api.RolesKey, false) - if err != nil { - return viewChannelReq{}, errors.Wrap(apiutil.ErrValidation, err) - } - - req := viewChannelReq{ - id: chi.URLParam(r, "channelID"), - roles: roles, - } - - return req, nil -} - -func decodeCreateChannelReq(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - - req := createChannelReq{} - if err := json.NewDecoder(r.Body).Decode(&req.Channel); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - - return req, nil -} - -func decodeCreateChannelsReq(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - - req := createChannelsReq{} - if err := json.NewDecoder(r.Body).Decode(&req.Channels); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - - return req, nil -} - -func decodeListChannels(_ context.Context, r *http.Request) (any, error) { - name, err := apiutil.ReadStringQuery(r, api.NameKey, "") - if err != nil { - return listChannelsReq{}, errors.Wrap(apiutil.ErrValidation, err) - } - - tags, err := apiutil.ReadStringQuery(r, api.TagsKey, "") - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - var tq channels.TagsQuery - if tags != "" { - tq = channels.ToTagsQuery(tags) - } - - s, err := apiutil.ReadStringQuery(r, api.StatusKey, api.DefGroupStatus) - if err != nil { - return listChannelsReq{}, errors.Wrap(apiutil.ErrValidation, err) - } - status, err := channels.ToStatus(s) - if err != nil { - return listChannelsReq{}, errors.Wrap(apiutil.ErrValidation, err) - } - - meta, err := apiutil.ReadMetadataQuery(r, api.MetadataKey, nil) - if err != nil { - return listChannelsReq{}, errors.Wrap(apiutil.ErrValidation, err) - } - - offset, err := apiutil.ReadNumQuery[uint64](r, api.OffsetKey, api.DefOffset) - if err != nil { - return listChannelsReq{}, errors.Wrap(apiutil.ErrValidation, err) - } - limit, err := apiutil.ReadNumQuery[uint64](r, api.LimitKey, api.DefLimit) - if err != nil { - return listChannelsReq{}, errors.Wrap(apiutil.ErrValidation, err) - } - - dir, err := apiutil.ReadStringQuery(r, api.DirKey, api.DefDir) - if err != nil { - return listChannelsReq{}, errors.Wrap(apiutil.ErrValidation, err) - } - - order, err := apiutil.ReadStringQuery(r, api.OrderKey, api.DefOrder) - if err != nil { - return listChannelsReq{}, errors.Wrap(apiutil.ErrValidation, err) - } - - allActions, err := apiutil.ReadStringQuery(r, api.ActionsKey, "") - if err != nil { - return listChannelsReq{}, errors.Wrap(apiutil.ErrValidation, err) - } - - actions := []string{} - - allActions = strings.TrimSpace(allActions) - if allActions != "" { - actions = strings.Split(allActions, ",") - } - roleID, err := apiutil.ReadStringQuery(r, api.RoleIDKey, "") - if err != nil { - return listChannelsReq{}, errors.Wrap(apiutil.ErrValidation, err) - } - - roleName, err := apiutil.ReadStringQuery(r, api.RoleNameKey, "") - if err != nil { - return listChannelsReq{}, errors.Wrap(apiutil.ErrValidation, err) - } - - accessType, err := apiutil.ReadStringQuery(r, api.AccessTypeKey, "") - if err != nil { - return listChannelsReq{}, errors.Wrap(apiutil.ErrValidation, err) - } - - userID, err := apiutil.ReadStringQuery(r, api.UserKey, "") - if err != nil { - return listChannelsReq{}, errors.Wrap(apiutil.ErrValidation, err) - } - groupID, err := nullable.Parse(r.URL.Query(), api.GroupKey, nullable.ParseString) - if err != nil { - return listChannelsReq{}, errors.Wrap(apiutil.ErrValidation, err) - } - - clientID, err := apiutil.ReadStringQuery(r, api.ClientKey, "") - if err != nil { - return listChannelsReq{}, errors.Wrap(apiutil.ErrValidation, err) - } - - id, err := apiutil.ReadStringQuery(r, api.IDOrder, "") - if err != nil { - return listChannelsReq{}, errors.Wrap(apiutil.ErrValidation, err) - } - ot, err := apiutil.ReadBoolQuery(r, api.OnlyTotal, false) - if err != nil { - return listChannelsReq{}, errors.Wrap(apiutil.ErrValidation, err) - } - connectionType, err := apiutil.ReadStringQuery(r, api.ConnTypeKey, "") - if err != nil { - return listChannelsReq{}, errors.Wrap(apiutil.ErrValidation, err) - } - - cfrom, err := apiutil.ReadStringQuery(r, "created_from", "") - if err != nil { - return listChannelsReq{}, errors.Wrap(apiutil.ErrValidation, err) - } - cto, err := apiutil.ReadStringQuery(r, "created_to", "") - if err != nil { - return listChannelsReq{}, errors.Wrap(apiutil.ErrValidation, err) - } - - var createdFrom, createdTo time.Time - if cfrom != "" { - if createdFrom, err = time.Parse(time.RFC3339, cfrom); err != nil { - return listChannelsReq{}, errors.Wrap(apiutil.ErrInvalidQueryParams, err) - } - } - if cto != "" { - if createdTo, err = time.Parse(time.RFC3339, cto); err != nil { - return listChannelsReq{}, errors.Wrap(apiutil.ErrInvalidQueryParams, err) - } - } - - req := listChannelsReq{ - Page: channels.Page{ - Name: name, - Tags: tq, - Status: status, - Metadata: meta, - RoleName: roleName, - RoleID: roleID, - Actions: actions, - AccessType: accessType, - Order: order, - Dir: dir, - Offset: offset, - Limit: limit, - Group: groupID, - Client: clientID, - ConnectionType: connectionType, - ID: id, - OnlyTotal: ot, - CreatedFrom: createdFrom, - CreatedTo: createdTo, - }, - userID: userID, - } - return req, nil -} - -func decodeUpdateChannel(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - - req := updateChannelReq{ - id: chi.URLParam(r, "channelID"), - } - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - - return req, nil -} - -func decodeUpdateChannelTags(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - - req := updateChannelTagsReq{ - id: chi.URLParam(r, "channelID"), - } - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - - return req, nil -} - -func decodeSetChannelParentGroupStatus(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - - req := setChannelParentGroupReq{ - id: chi.URLParam(r, "channelID"), - } - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - return req, nil -} - -func decodeRemoveChannelParentGroupStatus(_ context.Context, r *http.Request) (any, error) { - req := removeChannelParentGroupReq{ - id: chi.URLParam(r, "channelID"), - } - - return req, nil -} - -func decodeChangeChannelStatus(_ context.Context, r *http.Request) (any, error) { - req := changeChannelStatusReq{ - id: chi.URLParam(r, "channelID"), - } - - return req, nil -} - -func decodeDeleteChannelReq(_ context.Context, r *http.Request) (any, error) { - req := deleteChannelReq{ - id: chi.URLParam(r, "channelID"), - } - return req, nil -} - -func decodeConnectChannelClientRequest(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - req := connectChannelClientsRequest{ - channelID: chi.URLParam(r, "channelID"), - } - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - - return req, nil -} - -func decodeDisconnectChannelClientsRequest(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - req := disconnectChannelClientsRequest{ - channelID: chi.URLParam(r, "channelID"), - } - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - - return req, nil -} - -func decodeConnectRequest(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - - req := connectRequest{} - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - - return req, nil -} - -func decodeDisconnectRequest(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - - req := disconnectRequest{} - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - - return req, nil -} diff --git a/channels/api/http/endpoint_test.go b/channels/api/http/endpoint_test.go deleted file mode 100644 index 2485aeaeb..000000000 --- a/channels/api/http/endpoint_test.go +++ /dev/null @@ -1,2331 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package http - -import ( - "encoding/json" - "fmt" - "io" - "net/http" - "net/http/httptest" - "net/url" - "strings" - "testing" - "time" - - api "github.com/absmach/magistrala/api/http" - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/channels" - "github.com/absmach/magistrala/channels/mocks" - "github.com/absmach/magistrala/internal/testsutil" - mglog "github.com/absmach/magistrala/logger" - smqauthn "github.com/absmach/magistrala/pkg/authn" - authnmocks "github.com/absmach/magistrala/pkg/authn/mocks" - "github.com/absmach/magistrala/pkg/connections" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/pkg/roles" - "github.com/absmach/magistrala/pkg/uuid" - "github.com/go-chi/chi/v5" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" -) - -var ( - valid = "valid" - validChannelResp = channels.Channel{ - ID: testsutil.GenerateUUID(&testing.T{}), - Name: valid, - Domain: testsutil.GenerateUUID(&testing.T{}), - ParentGroup: testsutil.GenerateUUID(&testing.T{}), - Metadata: channels.Metadata{ - "name": "test", - }, - CreatedAt: time.Now().Add(-1 * time.Second), - UpdatedAt: time.Now(), - UpdatedBy: testsutil.GenerateUUID(&testing.T{}), - Status: channels.EnabledStatus, - } - validID = testsutil.GenerateUUID(&testing.T{}) - validToken = "validToken" - invalidToken = "invalidToken" - contentType = "application/json" - validTimeStamp = time.Date(2024, 1, 1, 0, 0, 0, 0, time.UTC) -) - -func newChannelsServer() (*httptest.Server, *mocks.Service, *authnmocks.Authentication) { - authn := new(authnmocks.Authentication) - svc := new(mocks.Service) - mux := chi.NewRouter() - idp := uuid.NewMock() - logger := mglog.NewMock() - am := smqauthn.NewAuthNMiddleware(authn, smqauthn.WithAllowUnverifiedUser(true)) - mux = MakeHandler(svc, am, mux, logger, "", idp) - - return httptest.NewServer(mux), svc, authn -} - -func TestCreateChannelEndpoint(t *testing.T) { - gs, svc, authn := newChannelsServer() - defer gs.Close() - - reqChannel := channels.Channel{ - Name: valid, - Metadata: map[string]any{ - "name": "test", - }, - Route: valid, - } - reqWithRoute := reqChannel - reqWithRoute.Route = valid - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - req channels.Channel - contentType string - svcResp []channels.Channel - svcErr error - authnErr error - status int - err error - }{ - { - desc: "create channel successfully", - token: validToken, - domainID: validID, - req: reqChannel, - contentType: contentType, - svcResp: []channels.Channel{validChannelResp}, - status: http.StatusCreated, - err: nil, - }, - { - desc: "create channel with route", - token: validToken, - domainID: validID, - req: reqWithRoute, - contentType: contentType, - svcResp: []channels.Channel{validChannelResp}, - status: http.StatusCreated, - err: nil, - }, - { - desc: "create channel with invalid token", - token: invalidToken, - session: smqauthn.Session{}, - domainID: validID, - req: reqChannel, - contentType: contentType, - authnErr: svcerr.ErrAuthentication, - status: http.StatusUnauthorized, - err: svcerr.ErrAuthentication, - }, - { - desc: "create channel with empty token", - token: "", - session: smqauthn.Session{}, - domainID: validID, - req: reqChannel, - contentType: contentType, - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "create channel with empty domainID", - token: validToken, - req: reqChannel, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrMissingDomainID, - }, - { - desc: "create channel with name that is too long", - token: validToken, - domainID: validID, - req: channels.Channel{ - Name: strings.Repeat("a", 1025), - Metadata: map[string]any{ - "name": "test", - }, - }, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrNameSize, - }, - { - desc: "create channel with invalid route format", - token: validToken, - domainID: validID, - req: channels.Channel{ - Name: valid, - Route: "__invalid", - }, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrInvalidRouteFormat, - }, - { - desc: "create channel with UUID route", - token: validToken, domainID: validID, - req: channels.Channel{ - Name: valid, - Route: testsutil.GenerateUUID(t), - }, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrInvalidRouteFormat, - }, - { - desc: "create channel with invalid content type", - token: validToken, - domainID: validID, - req: reqChannel, - contentType: "application/xml", - svcResp: []channels.Channel{validChannelResp}, - status: http.StatusUnsupportedMediaType, - err: apiutil.ErrUnsupportedContentType, - }, - { - desc: "create channel with service error", - token: validToken, - domainID: validID, - req: reqChannel, - contentType: contentType, - svcResp: []channels.Channel{}, - svcErr: svcerr.ErrAuthorization, - status: http.StatusForbidden, - err: svcerr.ErrAuthorization, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - data := toJSON(tc.req) - req := testRequest{ - client: gs.Client(), - method: http.MethodPost, - url: fmt.Sprintf("%s/%s/channels/", gs.URL, tc.domainID), - contentType: tc.contentType, - token: tc.token, - body: strings.NewReader(data), - } - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID + "_" + validID, UserID: validID, DomainID: validID} - } - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - svcCall := svc.On("CreateChannels", mock.Anything, tc.session, []channels.Channel{tc.req}).Return(tc.svcResp, []roles.RoleProvision{}, tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - var errRes respBody - err = json.NewDecoder(res.Body).Decode(&errRes) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if errRes.Err != "" || errRes.Message != "" { - err = errors.Wrap(errors.New(errRes.Err), errors.New(errRes.Message)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestCreateChannelsEndpoint(t *testing.T) { - gs, svc, authn := newChannelsServer() - defer gs.Close() - - reqChannels := []channels.Channel{ - { - Name: valid, - Metadata: map[string]any{ - "name": "test", - }, - Route: valid, - }, - } - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - req []channels.Channel - contentType string - svcResp []channels.Channel - svcErr error - authnErr error - status int - err error - }{ - { - desc: "create channels successfully", - token: validToken, - domainID: validID, - req: reqChannels, - contentType: contentType, - svcResp: []channels.Channel{validChannelResp}, - status: http.StatusOK, - err: nil, - }, - { - desc: "create channels with invalid token", - token: invalidToken, - session: smqauthn.Session{}, - domainID: validID, - req: reqChannels, - contentType: contentType, - authnErr: svcerr.ErrAuthentication, - status: http.StatusUnauthorized, - err: svcerr.ErrAuthentication, - }, - { - desc: "create channels with empty token", - token: "", - session: smqauthn.Session{}, - domainID: validID, - req: reqChannels, - contentType: contentType, - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "create channels with empty domainID", - token: validToken, - req: reqChannels, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrMissingDomainID, - }, - { - desc: "create channels with name that is too long", - token: validToken, - domainID: validID, - req: []channels.Channel{ - { - Name: strings.Repeat("a", 1025), - Metadata: map[string]any{ - "name": "test", - }, - }, - }, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrNameSize, - }, - { - desc: "create channels with invalid route format", - token: validToken, - domainID: validID, - req: []channels.Channel{ - { - Name: valid, - Route: "__invalid", - }, - }, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrInvalidRouteFormat, - }, - { - desc: "create channel with UUID route", - token: validToken, domainID: validID, - req: []channels.Channel{ - { - Name: valid, - Route: testsutil.GenerateUUID(t), - }, - }, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrInvalidRouteFormat, - }, - { - desc: "create channels with invalid content type", - token: validToken, - domainID: validID, - req: reqChannels, - contentType: "application/xml", - svcResp: []channels.Channel{validChannelResp}, - status: http.StatusUnsupportedMediaType, - err: apiutil.ErrUnsupportedContentType, - }, - { - desc: "create channels with service error", - token: validToken, - domainID: validID, - req: reqChannels, - contentType: contentType, - svcResp: []channels.Channel{}, - svcErr: svcerr.ErrAuthorization, - status: http.StatusForbidden, - err: svcerr.ErrAuthorization, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - data := toJSON(tc.req) - req := testRequest{ - client: gs.Client(), - method: http.MethodPost, - url: fmt.Sprintf("%s/%s/channels/bulk", gs.URL, tc.domainID), - contentType: tc.contentType, - token: tc.token, - body: strings.NewReader(data), - } - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID + "_" + validID, UserID: validID, DomainID: validID} - } - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - svcCall := svc.On("CreateChannels", mock.Anything, tc.session, tc.req).Return(tc.svcResp, []roles.RoleProvision{}, tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - var errRes respBody - err = json.NewDecoder(res.Body).Decode(&errRes) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if errRes.Err != "" || errRes.Message != "" { - err = errors.Wrap(errors.New(errRes.Err), errors.New(errRes.Message)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestViewChannelEndpoint(t *testing.T) { - gs, svc, authn := newChannelsServer() - defer gs.Close() - - cases := []struct { - desc string - token string - id string - domainID string - withRoles bool - session smqauthn.Session - svcResp channels.Channel - svcErr error - resp channels.Channel - status int - authnErr error - err error - }{ - { - desc: "view channel successfully", - token: validToken, - domainID: validID, - id: validID, - withRoles: false, - svcResp: validChannelResp, - svcErr: nil, - resp: validChannelResp, - status: http.StatusOK, - err: nil, - }, - { - desc: "view channel successfully with roles", - token: validToken, - domainID: validID, - id: validID, - withRoles: true, - svcResp: validChannelResp, - svcErr: nil, - resp: validChannelResp, - status: http.StatusOK, - err: nil, - }, - { - desc: "view channel with invalid token", - token: invalidToken, - session: smqauthn.Session{}, - domainID: validID, - id: validID, - withRoles: false, - svcResp: validChannelResp, - svcErr: nil, - authnErr: svcerr.ErrAuthentication, - status: http.StatusUnauthorized, - err: svcerr.ErrAuthentication, - }, - { - desc: "view channel with empty token", - token: "", - session: smqauthn.Session{}, - domainID: validID, - id: validID, - withRoles: false, - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "view channel with empty domainID", - token: validToken, - id: validID, - withRoles: false, - status: http.StatusBadRequest, - err: apiutil.ErrMissingDomainID, - }, - { - desc: "view channel with service error", - token: validToken, - id: validID, - domainID: validID, - withRoles: false, - svcResp: validChannelResp, - svcErr: svcerr.ErrAuthorization, - status: http.StatusForbidden, - err: svcerr.ErrAuthorization, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - client: gs.Client(), - method: http.MethodGet, - url: fmt.Sprintf("%s/%s/channels/%s?roles=%v", gs.URL, tc.domainID, tc.id, tc.withRoles), - token: tc.token, - } - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID + "_" + validID, UserID: validID, DomainID: validID} - } - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - svcCall := svc.On("ViewChannel", mock.Anything, tc.session, tc.id, tc.withRoles).Return(tc.svcResp, tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - var errRes respBody - err = json.NewDecoder(res.Body).Decode(&errRes) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if errRes.Err != "" || errRes.Message != "" { - err = errors.Wrap(errors.New(errRes.Err), errors.New(errRes.Message)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestListChannels(t *testing.T) { - gs, svc, authn := newChannelsServer() - defer gs.Close() - - cases := []struct { - desc string - query string - domainID string - token string - session smqauthn.Session - pageMeta channels.Page - listChannelsResponse channels.ChannelsPage - status int - authnErr error - err error - }{ - { - desc: "list channels successfully", - domainID: validID, - token: validToken, - status: http.StatusOK, - pageMeta: channels.Page{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Actions: []string{}, - }, - listChannelsResponse: channels.ChannelsPage{ - Page: channels.Page{ - Total: 1, - }, - Channels: []channels.Channel{validChannelResp}, - }, - err: nil, - }, - { - desc: "list channels with empty token", - domainID: validID, - token: "", - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "list channels with invalid token", - domainID: validID, - token: invalidToken, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "list channels with offset", - domainID: validID, - token: validToken, - pageMeta: channels.Page{ - Offset: 1, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Actions: []string{}, - }, - listChannelsResponse: channels.ChannelsPage{ - Page: channels.Page{ - Total: 1, - }, - Channels: []channels.Channel{validChannelResp}, - }, - query: "offset=1", - status: http.StatusOK, - err: nil, - }, - { - desc: "list channels with invalid offset", - domainID: validID, - token: validToken, - query: "offset=invalid", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list channels with limit", - domainID: validID, - token: validToken, - pageMeta: channels.Page{ - Offset: 0, - Limit: 1, - Order: api.DefOrder, - Dir: api.DefDir, - Actions: []string{}, - }, - listChannelsResponse: channels.ChannelsPage{ - Page: channels.Page{ - Total: 1, - }, - Channels: []channels.Channel{validChannelResp}, - }, - query: "limit=1", - status: http.StatusOK, - err: nil, - }, - { - desc: "list channels with invalid limit", - domainID: validID, - token: validToken, - query: "limit=invalid", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list channels with limit greater than max", - token: validToken, - domainID: validID, - query: fmt.Sprintf("limit=%d", api.MaxLimitSize+1), - status: http.StatusBadRequest, - err: apiutil.ErrLimitSize, - }, - { - desc: "list channels with name", - domainID: validID, - token: validToken, - pageMeta: channels.Page{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Actions: []string{}, - Name: "clientname", - }, - listChannelsResponse: channels.ChannelsPage{ - Page: channels.Page{ - Total: 1, - }, - Channels: []channels.Channel{validChannelResp}, - }, - query: "name=clientname", - status: http.StatusOK, - err: nil, - }, - { - desc: "list channels with duplicate name", - domainID: validID, - token: validToken, - query: "name=1&name=2", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list channels with status", - domainID: validID, - token: validToken, - pageMeta: channels.Page{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Actions: []string{}, - Status: channels.EnabledStatus, - }, - listChannelsResponse: channels.ChannelsPage{ - Page: channels.Page{ - Total: 1, - }, - Channels: []channels.Channel{validChannelResp}, - }, - query: "status=enabled", - status: http.StatusOK, - err: nil, - }, - { - desc: "list channels with invalid status", - domainID: validID, - token: validToken, - query: "status=invalid", - status: http.StatusBadRequest, - err: svcerr.ErrInvalidStatus, - }, - { - desc: "list channels with duplicate status", - domainID: validID, - token: validToken, - query: "status=enabled&status=disabled", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list channels with single tag", - domainID: validID, - token: validToken, - pageMeta: channels.Page{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Actions: []string{}, - Tags: channels.TagsQuery{Elements: []string{"tag1"}, Operator: channels.OrOp}, - }, - listChannelsResponse: channels.ChannelsPage{ - Page: channels.Page{ - Total: 1, - }, - Channels: []channels.Channel{validChannelResp}, - }, - query: "tags=tag1", - status: http.StatusOK, - err: nil, - }, - { - desc: "list channels with multiple tags and OR operator", - domainID: validID, - token: validToken, - pageMeta: channels.Page{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Actions: []string{}, - Tags: channels.TagsQuery{Elements: []string{"tag1", "tag2", "tag3"}, Operator: channels.OrOp}, - }, - listChannelsResponse: channels.ChannelsPage{ - Page: channels.Page{ - Total: 1, - }, - Channels: []channels.Channel{validChannelResp}, - }, - query: "tags=tag1,tag2,tag3", - status: http.StatusOK, - err: nil, - }, - { - desc: "list channels with multiple tags and AND operator", - domainID: validID, - token: validToken, - pageMeta: channels.Page{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Actions: []string{}, - Tags: channels.TagsQuery{Elements: []string{"tag1", "tag2", "tag3"}, Operator: channels.AndOp}, - }, - listChannelsResponse: channels.ChannelsPage{ - Page: channels.Page{ - Total: 1, - }, - Channels: []channels.Channel{validChannelResp}, - }, - query: "tags=tag1%2Btag2%2Btag3", - status: http.StatusOK, - err: nil, - }, - { - desc: "list channels with duplicate tags", - domainID: validID, - token: validToken, - query: "tags=tag1&tags=tag2", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list channels with metadata", - domainID: validID, - token: validToken, - pageMeta: channels.Page{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Actions: []string{}, - Metadata: channels.Metadata{"domain": "example.com"}, - }, - listChannelsResponse: channels.ChannelsPage{ - Page: channels.Page{ - Total: 1, - }, - Channels: []channels.Channel{validChannelResp}, - }, - query: fmt.Sprintf("metadata=%s", url.PathEscape(`{"domain": "example.com"}`)), - status: http.StatusOK, - err: nil, - }, - { - desc: "list channels with invalid metadata", - domainID: validID, - token: validToken, - query: "metadata=invalid", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list channels with duplicate metadata", - domainID: validID, - token: validToken, - query: fmt.Sprintf("metadata=%s&metadata=%s", url.PathEscape(`{"domain": "example.com"}`), url.PathEscape(`{"domain": "example.com"}`)), - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list channels with client ID", - domainID: validID, - token: validToken, - pageMeta: channels.Page{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Actions: []string{}, - Client: validID, - }, - listChannelsResponse: channels.ChannelsPage{ - Page: channels.Page{ - Total: 1, - }, - Channels: []channels.Channel{validChannelResp}, - }, - query: "client=" + validID, - status: http.StatusOK, - err: nil, - }, - { - desc: "list channels with client ID and connection type publish", - domainID: validID, - token: validToken, - pageMeta: channels.Page{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Actions: []string{}, - Client: validID, - ConnectionType: "publish", - }, - listChannelsResponse: channels.ChannelsPage{ - Page: channels.Page{ - Total: 1, - }, - Channels: []channels.Channel{validChannelResp}, - }, - query: "client=" + validID + "&connection_type=publish", - status: http.StatusOK, - err: nil, - }, - { - desc: "list channels with client ID and connection type subscribe", - domainID: validID, - token: validToken, - pageMeta: channels.Page{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Actions: []string{}, - Client: validID, - ConnectionType: "subscribe", - }, - listChannelsResponse: channels.ChannelsPage{ - Page: channels.Page{ - Total: 1, - }, - Channels: []channels.Channel{validChannelResp}, - }, - query: "client=" + validID + "&connection_type=subscribe", - status: http.StatusOK, - err: nil, - }, - { - desc: "list channels with invalid connection type", - domainID: validID, - token: validToken, - query: "client=" + validID + "&connection_type=invalid", - status: http.StatusBadRequest, - err: apiutil.ErrValidation, - }, - { - desc: "list channels with duplicate connection type", - domainID: validID, - token: validToken, - query: "connection_type=publish&connection_type=subscribe", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list channels with created_from", - domainID: validID, - token: validToken, - pageMeta: channels.Page{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Actions: []string{}, - CreatedFrom: validTimeStamp, - }, - listChannelsResponse: channels.ChannelsPage{ - Page: channels.Page{ - Total: 1, - }, - Channels: []channels.Channel{validChannelResp}, - }, - query: "created_from=2024-01-01T00:00:00Z", - status: http.StatusOK, - err: nil, - }, - { - desc: "list channels with created_to", - domainID: validID, - token: validToken, - pageMeta: channels.Page{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Actions: []string{}, - CreatedTo: validTimeStamp, - }, - listChannelsResponse: channels.ChannelsPage{ - Page: channels.Page{ - Total: 1, - }, - Channels: []channels.Channel{validChannelResp}, - }, - query: "created_to=2024-01-01T00:00:00Z", - status: http.StatusOK, - err: nil, - }, - { - desc: "list channels with both created_from and created_to", - domainID: validID, - token: validToken, - pageMeta: channels.Page{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Actions: []string{}, - CreatedFrom: validTimeStamp, - CreatedTo: validTimeStamp, - }, - listChannelsResponse: channels.ChannelsPage{ - Page: channels.Page{ - Total: 1, - }, - Channels: []channels.Channel{validChannelResp}, - }, - query: "created_from=2024-01-01T00:00:00Z&created_to=2024-01-01T00:00:00Z", - status: http.StatusOK, - err: nil, - }, - { - desc: "list channels with invalid created_from", - domainID: validID, - token: validToken, - query: "created_from=invalid-timestamp", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list channels with duplicate created_from", - domainID: validID, - token: validToken, - query: "created_from=2024-01-01T00:00:00Z&created_from=2024-01-02T00:00:00Z", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list channels with invalid created_to", - domainID: validID, - token: validToken, - query: "created_to=invalid-timestamp", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list channels with duplicate created_to", - domainID: validID, - token: validToken, - query: "created_to=2024-12-31T23:59:59Z&created_to=2024-12-30T23:59:59Z", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - client: gs.Client(), - method: http.MethodGet, - url: gs.URL + "/" + tc.domainID + "/channels?" + tc.query, - contentType: contentType, - token: tc.token, - } - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID + "_" + validID, UserID: validID, DomainID: validID} - } - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - svcCall := svc.On("ListChannels", mock.Anything, tc.session, tc.pageMeta).Return(tc.listChannelsResponse, tc.err) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - var bodyRes respBody - err = json.NewDecoder(res.Body).Decode(&bodyRes) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if bodyRes.Err != "" || bodyRes.Message != "" { - err = errors.Wrap(errors.New(bodyRes.Err), errors.New(bodyRes.Message)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestUpdateChannelEndpoint(t *testing.T) { - gs, svc, authn := newChannelsServer() - defer gs.Close() - - updateChannelReq := channels.Channel{ - ID: validID, - Name: valid, - Metadata: map[string]any{ - "name": "test", - }, - } - - cases := []struct { - desc string - token string - id string - domainID string - updateReq channels.Channel - contentType string - session smqauthn.Session - svcResp channels.Channel - svcErr error - resp channels.Channel - status int - authnErr error - err error - }{ - { - desc: "update channel successfully", - token: validToken, - domainID: validID, - id: validID, - updateReq: updateChannelReq, - contentType: contentType, - svcResp: validChannelResp, - status: http.StatusOK, - err: nil, - }, - { - desc: "update channel with invalid token", - token: invalidToken, - session: smqauthn.Session{}, - domainID: validID, - id: validID, - updateReq: updateChannelReq, - contentType: contentType, - authnErr: svcerr.ErrAuthentication, - status: http.StatusUnauthorized, - err: svcerr.ErrAuthentication, - }, - { - desc: "update channel with empty token", - token: "", - session: smqauthn.Session{}, - domainID: validID, - id: validID, - updateReq: updateChannelReq, - contentType: contentType, - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "update channel with empty domainID", - token: validToken, - id: validID, - updateReq: updateChannelReq, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrMissingDomainID, - }, - { - desc: "update channel with name that is too long", - token: validToken, - id: validID, - domainID: validID, - updateReq: channels.Channel{ - ID: validID, - Name: strings.Repeat("a", 1025), - Metadata: map[string]any{ - "name": "test", - }, - }, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrNameSize, - }, - { - desc: "update channel with invalid content type", - token: validToken, - id: validID, - domainID: validID, - updateReq: updateChannelReq, - contentType: "application/xml", - svcResp: validChannelResp, - status: http.StatusUnsupportedMediaType, - err: apiutil.ErrUnsupportedContentType, - }, - { - desc: "update channel with service error", - token: validToken, - id: validID, - domainID: validID, - updateReq: updateChannelReq, - contentType: contentType, - svcResp: channels.Channel{}, - svcErr: svcerr.ErrAuthorization, - status: http.StatusForbidden, - err: svcerr.ErrAuthorization, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - data := toJSON(tc.updateReq) - req := testRequest{ - client: gs.Client(), - method: http.MethodPatch, - url: fmt.Sprintf("%s/%s/channels/%s", gs.URL, tc.domainID, tc.id), - contentType: tc.contentType, - token: tc.token, - body: strings.NewReader(data), - } - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID + "_" + validID, UserID: validID, DomainID: validID} - } - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - svcCall := svc.On("UpdateChannel", mock.Anything, tc.session, tc.updateReq).Return(tc.svcResp, tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - var errRes respBody - err = json.NewDecoder(res.Body).Decode(&errRes) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if errRes.Err != "" || errRes.Message != "" { - err = errors.Wrap(errors.New(errRes.Err), errors.New(errRes.Message)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestUpdateChannelTagsEndpoint(t *testing.T) { - gs, svc, authn := newChannelsServer() - defer gs.Close() - - newTag := "newtag" - - cases := []struct { - desc string - token string - id string - domainID string - data string - contentType string - session smqauthn.Session - svcResp channels.Channel - svcErr error - resp channels.Channel - status int - authnErr error - err error - }{ - { - desc: "update channel tags successfully", - token: validToken, - domainID: validID, - id: validID, - data: fmt.Sprintf(`{"tags":["%s"]}`, newTag), - contentType: contentType, - svcResp: validChannelResp, - status: http.StatusOK, - err: nil, - }, - { - desc: "update channel tags with invalid token", - token: invalidToken, - session: smqauthn.Session{}, - domainID: validID, - id: validID, - data: fmt.Sprintf(`{"tags":["%s"]}`, newTag), - contentType: contentType, - authnErr: svcerr.ErrAuthentication, - status: http.StatusUnauthorized, - err: svcerr.ErrAuthentication, - }, - { - desc: "update channel tags with empty token", - token: "", - session: smqauthn.Session{}, - domainID: validID, - id: validID, - data: fmt.Sprintf(`{"tags":["%s"]}`, newTag), - contentType: contentType, - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "update channel tags with empty domainID", - token: validToken, - id: validID, - data: fmt.Sprintf(`{"tags":["%s"]}`, newTag), - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrMissingDomainID, - }, - { - desc: "update channel tags with invalid content type", - token: validToken, - id: validID, - domainID: validID, - data: fmt.Sprintf(`{"tags":["%s"]}`, newTag), - contentType: "application/xml", - svcResp: validChannelResp, - status: http.StatusUnsupportedMediaType, - err: apiutil.ErrUnsupportedContentType, - }, - { - desc: "update channel tags with service error", - token: validToken, - id: validID, - domainID: validID, - data: fmt.Sprintf(`{"tags":["%s"]}`, newTag), - contentType: contentType, - svcResp: channels.Channel{}, - svcErr: svcerr.ErrAuthorization, - status: http.StatusForbidden, - err: svcerr.ErrAuthorization, - }, - { - desc: "update channel with malformed request", - token: validToken, - id: validID, - domainID: validID, - contentType: contentType, - data: fmt.Sprintf(`{"tags":["%s"}`, newTag), - status: http.StatusBadRequest, - err: apiutil.ErrMalformedRequestBody, - }, - { - desc: "update channel with empty id", - token: validToken, - id: "", - domainID: validID, - contentType: contentType, - data: fmt.Sprintf(`{"tags":["%s"]}`, newTag), - status: http.StatusBadRequest, - err: apiutil.ErrMissingID, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - client: gs.Client(), - method: http.MethodPatch, - url: fmt.Sprintf("%s/%s/channels/%s/tags", gs.URL, tc.domainID, tc.id), - contentType: tc.contentType, - token: tc.token, - body: strings.NewReader(tc.data), - } - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID + "_" + validID, UserID: validID, DomainID: validID} - } - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - svcCall := svc.On("UpdateChannelTags", mock.Anything, tc.session, channels.Channel{ID: tc.id, Tags: []string{newTag}}).Return(tc.svcResp, tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - var errRes respBody - err = json.NewDecoder(res.Body).Decode(&errRes) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if errRes.Err != "" || errRes.Message != "" { - err = errors.Wrap(errors.New(errRes.Err), errors.New(errRes.Message)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestSetChannelParentGroupEndpoint(t *testing.T) { - gs, svc, authn := newChannelsServer() - defer gs.Close() - - cases := []struct { - desc string - token string - id string - domainID string - data string - contentType string - session smqauthn.Session - svcErr error - resp channels.Channel - status int - authnErr error - err error - }{ - { - desc: "set channel parent group successfully", - token: validToken, - domainID: validID, - id: validID, - data: fmt.Sprintf(`{"parent_group_id":"%s"}`, validID), - contentType: contentType, - status: http.StatusOK, - err: nil, - }, - { - desc: "set channel parent group with invalid token", - token: invalidToken, - domainID: validID, - id: validID, - data: fmt.Sprintf(`{"parent_group_id":"%s"}`, validID), - contentType: contentType, - authnErr: svcerr.ErrAuthentication, - status: http.StatusUnauthorized, - err: svcerr.ErrAuthentication, - }, - { - desc: "set channel parent group with empty token", - token: "", - domainID: validID, - id: validID, - data: fmt.Sprintf(`{"parent_group_id":"%s"}`, validID), - contentType: contentType, - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "set channel parent group with empty domainID", - token: validToken, - id: validID, - data: fmt.Sprintf(`{"parent_group_id":"%s"}`, validID), - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrMissingDomainID, - }, - { - desc: "set channel parent group with invalid content type", - token: validToken, - id: validID, - domainID: validID, - data: fmt.Sprintf(`{"parent_group_id":"%s"}`, validID), - contentType: "application/xml", - status: http.StatusUnsupportedMediaType, - err: apiutil.ErrUnsupportedContentType, - }, - { - desc: "set channel parent group with empty id", - token: validToken, - id: "", - domainID: validID, - data: fmt.Sprintf(`{"parent_group_id":"%s"}`, validID), - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrMissingID, - }, - { - desc: "set channel parent group with empty parent group id", - token: validToken, - id: validID, - domainID: validID, - data: `{"parent_group_id":""}`, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrMissingParentGroupID, - }, - { - desc: "set channel parent group with malformed request", - token: validToken, - id: validID, - domainID: validID, - data: fmt.Sprintf(`{"parent_group_id":"%s"`, validID), - contentType: contentType, - status: http.StatusBadRequest, - err: errors.ErrMalformedEntity, - }, - { - desc: "set channel parent group with service error", - token: validToken, - id: validID, - domainID: validID, - data: fmt.Sprintf(`{"parent_group_id":"%s"}`, validID), - contentType: contentType, - svcErr: svcerr.ErrAuthorization, - status: http.StatusForbidden, - err: svcerr.ErrAuthorization, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - client: gs.Client(), - method: http.MethodPost, - url: fmt.Sprintf("%s/%s/channels/%s/parent", gs.URL, tc.domainID, tc.id), - contentType: tc.contentType, - token: tc.token, - body: strings.NewReader(tc.data), - } - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID + "_" + validID, UserID: validID, DomainID: validID} - } - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - svcCall := svc.On("SetParentGroup", mock.Anything, tc.session, validID, tc.id).Return(tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestRemoveChannelParentGroupEndpoint(t *testing.T) { - gs, svc, authn := newChannelsServer() - defer gs.Close() - - cases := []struct { - desc string - token string - id string - domainID string - session smqauthn.Session - svcErr error - resp channels.Channel - status int - authnErr error - err error - }{ - { - desc: "remove channel parent group successfully", - token: validToken, - id: validID, - domainID: validID, - status: http.StatusNoContent, - err: nil, - }, - { - desc: "remove channel parent group with invalid token", - token: invalidToken, - session: smqauthn.Session{}, - id: validID, - domainID: validID, - authnErr: svcerr.ErrAuthentication, - status: http.StatusUnauthorized, - err: svcerr.ErrAuthentication, - }, - { - desc: "remove channel parent group with empty token", - token: "", - id: validID, - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "remove channel parent group with empty domainID", - token: validToken, - id: validID, - status: http.StatusBadRequest, - err: apiutil.ErrMissingDomainID, - }, - { - desc: "remove channel parent group with empty id", - token: validToken, - id: "", - domainID: validID, - status: http.StatusBadRequest, - err: apiutil.ErrMissingID, - }, - { - desc: "remove channel parent group with service error", - token: validToken, - id: validID, - domainID: validID, - svcErr: svcerr.ErrAuthorization, - status: http.StatusForbidden, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - client: gs.Client(), - method: http.MethodDelete, - url: fmt.Sprintf("%s/%s/channels/%s/parent", gs.URL, tc.domainID, tc.id), - token: tc.token, - } - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID + "_" + validID, UserID: validID, DomainID: validID} - } - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - svcCall := svc.On("RemoveParentGroup", mock.Anything, tc.session, tc.id).Return(tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestEnableChannelEndpoint(t *testing.T) { - gs, svc, authn := newChannelsServer() - defer gs.Close() - - cases := []struct { - desc string - token string - id string - domainID string - session smqauthn.Session - svcResp channels.Channel - svcErr error - resp channels.Channel - status int - authnErr error - err error - }{ - { - desc: "enable channel successfully", - token: validToken, - domainID: validID, - id: validID, - svcResp: validChannelResp, - svcErr: nil, - resp: validChannelResp, - status: http.StatusOK, - err: nil, - }, - { - desc: "enable channel with invalid token", - token: invalidToken, - session: smqauthn.Session{}, - domainID: validID, - id: validID, - authnErr: svcerr.ErrAuthentication, - status: http.StatusUnauthorized, - err: svcerr.ErrAuthentication, - }, - { - desc: "enable channel with empty token", - token: "", - session: smqauthn.Session{}, - domainID: validID, - id: validID, - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "enable channel with empty domainID", - token: validToken, - id: validID, - status: http.StatusBadRequest, - err: apiutil.ErrMissingDomainID, - }, - { - desc: "enable channel with service error", - token: validToken, - id: validID, - domainID: validID, - svcResp: channels.Channel{}, - svcErr: svcerr.ErrAuthorization, - status: http.StatusForbidden, - err: svcerr.ErrAuthorization, - }, - { - desc: "enable channel with empty id", - token: validToken, - id: "", - domainID: validID, - status: http.StatusBadRequest, - err: apiutil.ErrMissingID, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - client: gs.Client(), - method: http.MethodPost, - url: fmt.Sprintf("%s/%s/channels/%s/enable", gs.URL, tc.domainID, tc.id), - token: tc.token, - } - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID + "_" + validID, UserID: validID, DomainID: validID} - } - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - svcCall := svc.On("EnableChannel", mock.Anything, tc.session, tc.id).Return(tc.svcResp, tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - var errRes respBody - err = json.NewDecoder(res.Body).Decode(&errRes) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if errRes.Err != "" || errRes.Message != "" { - err = errors.Wrap(errors.New(errRes.Err), errors.New(errRes.Message)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestDisableChannelEndpoint(t *testing.T) { - gs, svc, authn := newChannelsServer() - defer gs.Close() - - cases := []struct { - desc string - token string - id string - domainID string - session smqauthn.Session - svcResp channels.Channel - svcErr error - resp channels.Channel - status int - authnErr error - err error - }{ - { - desc: "disable channel successfully", - token: validToken, - domainID: validID, - id: validID, - svcResp: validChannelResp, - svcErr: nil, - resp: validChannelResp, - status: http.StatusOK, - err: nil, - }, - { - desc: "disable channel with invalid token", - token: invalidToken, - session: smqauthn.Session{}, - domainID: validID, - id: validID, - authnErr: svcerr.ErrAuthentication, - status: http.StatusUnauthorized, - err: svcerr.ErrAuthentication, - }, - { - desc: "disable channel with empty token", - token: "", - session: smqauthn.Session{}, - domainID: validID, - id: validID, - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "disable channel with empty domainID", - token: validToken, - id: validID, - status: http.StatusBadRequest, - err: apiutil.ErrMissingDomainID, - }, - { - desc: "disable channel with service error", - token: validToken, - id: validID, - domainID: validID, - svcResp: channels.Channel{}, - svcErr: svcerr.ErrAuthorization, - status: http.StatusForbidden, - err: svcerr.ErrAuthorization, - }, - { - desc: "disable channel with empty id", - token: validToken, - id: "", - domainID: validID, - status: http.StatusBadRequest, - err: apiutil.ErrMissingID, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - client: gs.Client(), - method: http.MethodPost, - url: fmt.Sprintf("%s/%s/channels/%s/disable", gs.URL, tc.domainID, tc.id), - token: tc.token, - } - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID + "_" + validID, UserID: validID, DomainID: validID} - } - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - svcCall := svc.On("DisableChannel", mock.Anything, tc.session, tc.id).Return(tc.svcResp, tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - var errRes respBody - err = json.NewDecoder(res.Body).Decode(&errRes) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if errRes.Err != "" || errRes.Message != "" { - err = errors.Wrap(errors.New(errRes.Err), errors.New(errRes.Message)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestConnectChannelClientEndpoint(t *testing.T) { - gs, svc, authn := newChannelsServer() - defer gs.Close() - - cases := []struct { - desc string - token string - id string - domainID string - data string - session smqauthn.Session - contentType string - svcErr error - status int - authnErr error - err error - }{ - { - desc: "connect channel client successfully", - token: validToken, - domainID: validID, - id: validID, - data: fmt.Sprintf(`{"client_ids": ["%s"], "types": ["Publish"]}`, validID), - contentType: contentType, - svcErr: nil, - status: http.StatusCreated, - err: nil, - }, - { - desc: "connect channel client with invalid token", - token: invalidToken, - domainID: validID, - id: validID, - data: fmt.Sprintf(`{"client_ids": ["%s"], "types": ["Publish"]}`, validID), - contentType: contentType, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "connect channel client with empty token", - token: "", - session: smqauthn.Session{}, - domainID: validID, - id: validID, - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "connect channel client with empty domainID", - token: validToken, - id: validID, - status: http.StatusBadRequest, - err: apiutil.ErrMissingDomainID, - }, - { - desc: "connect channel client with service error", - token: validToken, - id: validID, - domainID: validID, - data: fmt.Sprintf(`{"client_ids": ["%s"], "types": ["Publish"]}`, validID), - contentType: contentType, - svcErr: svcerr.ErrAuthorization, - status: http.StatusForbidden, - err: svcerr.ErrAuthorization, - }, - { - desc: "connect channel client with empty id", - token: validToken, - id: "", - domainID: validID, - data: fmt.Sprintf(`{"client_ids": ["%s"], "types": ["Publish"]}`, validID), - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrMissingID, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - client: gs.Client(), - method: http.MethodPost, - url: fmt.Sprintf("%s/%s/channels/%s/connect", gs.URL, tc.domainID, tc.id), - token: tc.token, - contentType: tc.contentType, - body: strings.NewReader(tc.data), - } - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID + "_" + validID, UserID: validID, DomainID: validID} - } - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - svcCall := svc.On("Connect", mock.Anything, tc.session, []string{tc.id}, []string{validID}, []connections.ConnType{1}).Return(tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestDisconnectChannelClientEndpoint(t *testing.T) { - gs, svc, authn := newChannelsServer() - defer gs.Close() - - cases := []struct { - desc string - token string - id string - domainID string - data string - session smqauthn.Session - contentType string - svcErr error - status int - authnErr error - err error - }{ - { - desc: "disconnect channel client successfully", - token: validToken, - domainID: validID, - id: validID, - data: fmt.Sprintf(`{"client_ids": ["%s"], "types": ["Publish"]}`, validID), - contentType: contentType, - svcErr: nil, - status: http.StatusNoContent, - err: nil, - }, - { - desc: "disconnect channel client with invalid token", - token: invalidToken, - domainID: validID, - id: validID, - data: fmt.Sprintf(`{"client_ids": ["%s"], "types": ["Publish"]}`, validID), - contentType: contentType, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "disconnect channel client with empty token", - token: "", - session: smqauthn.Session{}, - domainID: validID, - id: validID, - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "disconnect channel client with empty domainID", - token: validToken, - id: validID, - status: http.StatusBadRequest, - err: apiutil.ErrMissingDomainID, - }, - { - desc: "disconnect channel client with service error", - token: validToken, - id: validID, - domainID: validID, - data: fmt.Sprintf(`{"client_ids": ["%s"], "types": ["Publish"]}`, validID), - contentType: contentType, - svcErr: svcerr.ErrAuthorization, - status: http.StatusForbidden, - err: svcerr.ErrAuthorization, - }, - { - desc: "disconnect channel client with empty id", - token: validToken, - id: "", - domainID: validID, - data: fmt.Sprintf(`{"client_ids": ["%s"], "types": ["Publish"]}`, validID), - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrMissingID, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - client: gs.Client(), - method: http.MethodPost, - url: fmt.Sprintf("%s/%s/channels/%s/disconnect", gs.URL, tc.domainID, tc.id), - token: tc.token, - contentType: tc.contentType, - body: strings.NewReader(tc.data), - } - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID + "_" + validID, UserID: validID, DomainID: validID} - } - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - svcCall := svc.On("Disconnect", mock.Anything, tc.session, []string{tc.id}, []string{validID}, []connections.ConnType{1}).Return(tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestConnectEndpoint(t *testing.T) { - gs, svc, authn := newChannelsServer() - defer gs.Close() - - cases := []struct { - desc string - token string - channelIDs []string - domainID string - clientIDs []string - types []connections.ConnType - session smqauthn.Session - svcErr error - status int - authnErr error - err error - }{ - { - desc: "connect successfully", - token: validToken, - domainID: validID, - channelIDs: []string{validID}, - clientIDs: []string{validID}, - types: []connections.ConnType{1}, - svcErr: nil, - status: http.StatusCreated, - err: nil, - }, - { - desc: "connect with invalid token", - token: invalidToken, - domainID: validID, - channelIDs: []string{validID}, - clientIDs: []string{validID}, - types: []connections.ConnType{1}, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "connect with empty token", - token: "", - session: smqauthn.Session{}, - domainID: validID, - channelIDs: []string{validID}, - clientIDs: []string{validID}, - types: []connections.ConnType{1}, - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "connect with empty domainID", - token: validToken, - channelIDs: []string{validID}, - clientIDs: []string{validID}, - types: []connections.ConnType{1}, - status: http.StatusBadRequest, - err: apiutil.ErrMissingDomainID, - }, - { - desc: "connect with service error", - token: validToken, - channelIDs: []string{validID}, - domainID: validID, - clientIDs: []string{validID}, - types: []connections.ConnType{1}, - svcErr: svcerr.ErrAuthorization, - status: http.StatusForbidden, - err: svcerr.ErrAuthorization, - }, - { - desc: "connect with empty channel ids", - token: validToken, - channelIDs: []string{}, - domainID: validID, - clientIDs: []string{validID}, - types: []connections.ConnType{1}, - status: http.StatusBadRequest, - err: apiutil.ErrMissingID, - }, - { - desc: "connect with empty client ids", - token: validToken, - channelIDs: []string{validID}, - domainID: validID, - clientIDs: []string{}, - types: []connections.ConnType{1}, - status: http.StatusBadRequest, - err: apiutil.ErrMissingID, - }, - { - desc: "connect with empty types", - token: validToken, - channelIDs: []string{validID}, - domainID: validID, - clientIDs: []string{validID}, - types: []connections.ConnType{}, - status: http.StatusBadRequest, - err: apiutil.ErrMissingID, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - client: gs.Client(), - method: http.MethodPost, - url: fmt.Sprintf("%s/%s/channels/connect", gs.URL, tc.domainID), - token: tc.token, - contentType: contentType, - body: strings.NewReader(toJSON(map[string]any{ - "channel_ids": tc.channelIDs, - "client_ids": tc.clientIDs, - "types": tc.types, - })), - } - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID + "_" + validID, UserID: validID, DomainID: validID} - } - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - svcCall := svc.On("Connect", mock.Anything, tc.session, tc.channelIDs, tc.clientIDs, tc.types).Return(tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestDisconnectEndpoint(t *testing.T) { - gs, svc, authn := newChannelsServer() - defer gs.Close() - - cases := []struct { - desc string - token string - channelIDs []string - domainID string - clientIDs []string - types []connections.ConnType - session smqauthn.Session - svcErr error - status int - authnErr error - err error - }{ - { - desc: "disconnect successfully", - token: validToken, - domainID: validID, - channelIDs: []string{validID}, - clientIDs: []string{validID}, - types: []connections.ConnType{1}, - svcErr: nil, - status: http.StatusNoContent, - err: nil, - }, - { - desc: "disconnect with invalid token", - token: invalidToken, - domainID: validID, - channelIDs: []string{validID}, - clientIDs: []string{validID}, - types: []connections.ConnType{1}, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "disconnect with empty token", - token: "", - session: smqauthn.Session{}, - domainID: validID, - channelIDs: []string{validID}, - clientIDs: []string{validID}, - types: []connections.ConnType{1}, - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "disconnect with empty domainID", - token: validToken, - channelIDs: []string{validID}, - clientIDs: []string{validID}, - types: []connections.ConnType{1}, - status: http.StatusBadRequest, - err: apiutil.ErrMissingDomainID, - }, - { - desc: "disconnect with service error", - token: validToken, - channelIDs: []string{validID}, - domainID: validID, - clientIDs: []string{validID}, - types: []connections.ConnType{1}, - svcErr: svcerr.ErrAuthorization, - status: http.StatusForbidden, - err: svcerr.ErrAuthorization, - }, - { - desc: "disconnect with empty channel ids", - token: validToken, - channelIDs: []string{}, - domainID: validID, - clientIDs: []string{validID}, - types: []connections.ConnType{1}, - status: http.StatusBadRequest, - err: apiutil.ErrMissingID, - }, - { - desc: "disconnect with empty client ids", - token: validToken, - channelIDs: []string{validID}, - domainID: validID, - clientIDs: []string{}, - types: []connections.ConnType{1}, - status: http.StatusBadRequest, - err: apiutil.ErrMissingID, - }, - { - desc: "disconnect with empty types", - token: validToken, - channelIDs: []string{validID}, - domainID: validID, - clientIDs: []string{validID}, - types: []connections.ConnType{}, - status: http.StatusBadRequest, - err: apiutil.ErrMissingID, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - client: gs.Client(), - method: http.MethodPost, - url: fmt.Sprintf("%s/%s/channels/disconnect", gs.URL, tc.domainID), - token: tc.token, - contentType: contentType, - body: strings.NewReader(toJSON(map[string]any{ - "channel_ids": tc.channelIDs, - "client_ids": tc.clientIDs, - "types": tc.types, - })), - } - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID + "_" + validID, UserID: validID, DomainID: validID} - } - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - svcCall := svc.On("Disconnect", mock.Anything, tc.session, tc.channelIDs, tc.clientIDs, tc.types).Return(tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestDeleteChannelEndpoint(t *testing.T) { - gs, svc, authn := newChannelsServer() - defer gs.Close() - - cases := []struct { - desc string - token string - id string - domainID string - session smqauthn.Session - svcErr error - status int - authnErr error - err error - }{ - { - desc: "delete channel successfully", - token: validToken, - domainID: validID, - id: validID, - svcErr: nil, - status: http.StatusNoContent, - err: nil, - }, - { - desc: "delete channel with invalid token", - token: invalidToken, - session: smqauthn.Session{}, - domainID: validID, - id: validID, - authnErr: svcerr.ErrAuthentication, - status: http.StatusUnauthorized, - err: svcerr.ErrAuthentication, - }, - { - desc: "delete channel with empty token", - token: "", - session: smqauthn.Session{}, - domainID: validID, - id: validID, - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "delete channel with empty domainID", - token: validToken, - id: validID, - status: http.StatusBadRequest, - err: apiutil.ErrMissingDomainID, - }, - { - desc: "delete channel with service error", - token: validToken, - id: validID, - domainID: validID, - svcErr: svcerr.ErrAuthorization, - status: http.StatusForbidden, - err: svcerr.ErrAuthorization, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - client: gs.Client(), - method: http.MethodDelete, - url: fmt.Sprintf("%s/%s/channels/%s", gs.URL, tc.domainID, tc.id), - token: tc.token, - } - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID + "_" + validID, UserID: validID, DomainID: validID} - } - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - svcCall := svc.On("RemoveChannel", mock.Anything, tc.session, tc.id).Return(tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -type testRequest struct { - client *http.Client - method string - url string - contentType string - token string - body io.Reader -} - -func (tr testRequest) make() (*http.Response, error) { - req, err := http.NewRequest(tr.method, tr.url, tr.body) - if err != nil { - return nil, err - } - - if tr.token != "" { - req.Header.Set("Authorization", apiutil.BearerPrefix+tr.token) - } - - if tr.contentType != "" { - req.Header.Set("Content-Type", tr.contentType) - } - - req.Header.Set("Referer", "http://localhost") - - return tr.client.Do(req) -} - -func toJSON(data any) string { - jsonData, err := json.Marshal(data) - if err != nil { - return "" - } - return string(jsonData) -} - -type respBody struct { - Err string `json:"error"` - Message string `json:"message"` - Total int `json:"total"` - Permissions []string `json:"permissions"` - ID string `json:"id"` - Tags []string `json:"tags"` - Status channels.Status `json:"status"` -} diff --git a/channels/api/http/endpoints.go b/channels/api/http/endpoints.go deleted file mode 100644 index 6f659073f..000000000 --- a/channels/api/http/endpoints.go +++ /dev/null @@ -1,364 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package http - -import ( - "context" - - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/channels" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/go-kit/kit/endpoint" -) - -func createChannelEndpoint(svc channels.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(createChannelReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - channels, _, err := svc.CreateChannels(ctx, session, req.Channel) - if err != nil { - return nil, err - } - - return createChannelRes{ - Channel: channels[0], - created: true, - }, nil - } -} - -func createChannelsEndpoint(svc channels.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(createChannelsReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - channels, _, err := svc.CreateChannels(ctx, session, req.Channels...) - if err != nil { - return nil, err - } - - res := channelsPageRes{ - pageRes: pageRes{ - Total: uint64(len(channels)), - }, - Channels: []viewChannelRes{}, - } - for _, c := range channels { - res.Channels = append(res.Channels, viewChannelRes{Channel: c}) - } - - return res, nil - } -} - -func viewChannelEndpoint(svc channels.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(viewChannelReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - c, err := svc.ViewChannel(ctx, session, req.id, req.roles) - if err != nil { - return nil, err - } - - return viewChannelRes{Channel: c}, nil - } -} - -func listChannelsEndpoint(svc channels.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(listChannelsReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - var page channels.ChannelsPage - var err error - switch req.userID != "" { - case true: - page, err = svc.ListUserChannels(ctx, session, req.userID, req.Page) - default: - page, err = svc.ListChannels(ctx, session, req.Page) - } - if err != nil { - return channelsPageRes{}, err - } - - res := channelsPageRes{ - pageRes: pageRes{ - Total: page.Total, - Offset: page.Offset, - Limit: page.Limit, - }, - Channels: []viewChannelRes{}, - } - for _, c := range page.Channels { - res.Channels = append(res.Channels, viewChannelRes{Channel: c}) - } - - return res, nil - } -} - -func updateChannelEndpoint(svc channels.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(updateChannelReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - ch := channels.Channel{ - ID: req.id, - Name: req.Name, - Metadata: req.Metadata, - } - ch, err := svc.UpdateChannel(ctx, session, ch) - if err != nil { - return nil, err - } - - return updateChannelRes{Channel: ch}, nil - } -} - -func updateChannelTagsEndpoint(svc channels.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(updateChannelTagsReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - ch := channels.Channel{ - ID: req.id, - Tags: req.Tags, - } - ch, err := svc.UpdateChannelTags(ctx, session, ch) - if err != nil { - return nil, err - } - - return updateChannelRes{Channel: ch}, nil - } -} - -func setChannelParentGroupEndpoint(svc channels.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(setChannelParentGroupReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - if err := svc.SetParentGroup(ctx, session, req.ParentGroupID, req.id); err != nil { - return nil, err - } - - return setChannelParentGroupRes{}, nil - } -} - -func removeChannelParentGroupEndpoint(svc channels.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(removeChannelParentGroupReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - if err := svc.RemoveParentGroup(ctx, session, req.id); err != nil { - return nil, err - } - - return removeChannelParentGroupRes{}, nil - } -} - -func enableChannelEndpoint(svc channels.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(changeChannelStatusReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - ch, err := svc.EnableChannel(ctx, session, req.id) - if err != nil { - return nil, err - } - - return changeChannelStatusRes{Channel: ch}, nil - } -} - -func disableChannelEndpoint(svc channels.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(changeChannelStatusReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - ch, err := svc.DisableChannel(ctx, session, req.id) - if err != nil { - return nil, err - } - - return changeChannelStatusRes{Channel: ch}, nil - } -} - -func connectChannelClientEndpoint(svc channels.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(connectChannelClientsRequest) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - if err := svc.Connect(ctx, session, []string{req.channelID}, req.ClientIDs, req.Types); err != nil { - return nil, err - } - - return connectChannelClientsRes{}, nil - } -} - -func disconnectChannelClientsEndpoint(svc channels.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(disconnectChannelClientsRequest) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - if err := svc.Disconnect(ctx, session, []string{req.channelID}, req.ClientIds, req.Types); err != nil { - return nil, err - } - - return disconnectChannelClientsRes{}, nil - } -} - -func connectEndpoint(svc channels.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(connectRequest) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - if err := svc.Connect(ctx, session, req.ChannelIds, req.ClientIds, req.Types); err != nil { - return nil, err - } - - return connectRes{}, nil - } -} - -func disconnectEndpoint(svc channels.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(disconnectRequest) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - if err := svc.Disconnect(ctx, session, req.ChannelIds, req.ClientIds, req.Types); err != nil { - return nil, err - } - - return disconnectRes{}, nil - } -} - -func deleteChannelEndpoint(svc channels.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(deleteChannelReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - if err := svc.RemoveChannel(ctx, session, req.id); err != nil { - return nil, err - } - - return deleteChannelRes{}, nil - } -} diff --git a/channels/api/http/requests.go b/channels/api/http/requests.go deleted file mode 100644 index 5dbd8256f..000000000 --- a/channels/api/http/requests.go +++ /dev/null @@ -1,321 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package http - -import ( - "strings" - - api "github.com/absmach/magistrala/api/http" - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/channels" - "github.com/absmach/magistrala/pkg/connections" -) - -type createChannelReq struct { - Channel channels.Channel -} - -func (req createChannelReq) validate() error { - if len(req.Channel.Name) > api.MaxNameSize { - return apiutil.ErrNameSize - } - if req.Channel.ID != "" { - if strings.TrimSpace(req.Channel.ID) == "" { - return apiutil.ErrMissingChannelID - } - } - if req.Channel.Route != "" { - if err := api.ValidateRoute(req.Channel.Route); err != nil { - return err - } - if err := api.ValidateUUID(req.Channel.Route); err == nil { - return apiutil.ErrInvalidRouteFormat - } - } - - return nil -} - -type createChannelsReq struct { - Channels []channels.Channel -} - -func (req createChannelsReq) validate() error { - if len(req.Channels) == 0 { - return apiutil.ErrEmptyList - } - for _, channel := range req.Channels { - if channel.ID != "" { - if strings.TrimSpace(channel.ID) == "" { - return apiutil.ErrMissingChannelID - } - } - if len(channel.Name) > api.MaxNameSize { - return apiutil.ErrNameSize - } - if channel.Route != "" { - if err := api.ValidateRoute(channel.Route); err != nil { - return err - } - if err := api.ValidateUUID(channel.Route); err == nil { - return apiutil.ErrInvalidRouteFormat - } - } - } - - return nil -} - -type viewChannelReq struct { - id string - roles bool -} - -func (req viewChannelReq) validate() error { - if req.id == "" { - return apiutil.ErrMissingID - } - return nil -} - -type listChannelsReq struct { - channels.Page - userID string -} - -func (req listChannelsReq) validate() error { - if req.Limit > api.MaxLimitSize || req.Limit < 1 { - return apiutil.ErrLimitSize - } - - if len(req.Name) > api.MaxNameSize { - return apiutil.ErrNameSize - } - - switch req.Order { - case "", api.NameOrder, api.CreatedAtOrder, api.UpdatedAtOrder: - default: - return apiutil.ErrInvalidOrder - } - - if req.Dir != "" && (req.Dir != api.DescDir && req.Dir != api.AscDir) { - return apiutil.ErrInvalidDirection - } - - if req.ConnectionType != "" { - if _, err := connections.ParseConnType(req.ConnectionType); err != nil { - return apiutil.ErrValidation - } - } - - return nil -} - -type updateChannelReq struct { - id string - Name string `json:"name,omitempty"` - Metadata map[string]any `json:"metadata,omitempty"` - Tags []string `json:"tags,omitempty"` -} - -func (req updateChannelReq) validate() error { - if req.id == "" { - return apiutil.ErrMissingID - } - if len(req.Name) > api.MaxNameSize { - return apiutil.ErrNameSize - } - - return nil -} - -type updateChannelTagsReq struct { - id string - Tags []string `json:"tags,omitempty"` -} - -func (req updateChannelTagsReq) validate() error { - if req.id == "" { - return apiutil.ErrMissingID - } - - return nil -} - -type setChannelParentGroupReq struct { - id string - ParentGroupID string `json:"parent_group_id"` -} - -func (req setChannelParentGroupReq) validate() error { - if req.id == "" { - return apiutil.ErrMissingID - } - if req.ParentGroupID == "" { - return apiutil.ErrMissingParentGroupID - } - - return nil -} - -type removeChannelParentGroupReq struct { - id string -} - -func (req removeChannelParentGroupReq) validate() error { - if req.id == "" { - return apiutil.ErrMissingID - } - - return nil -} - -type changeChannelStatusReq struct { - id string -} - -func (req changeChannelStatusReq) validate() error { - if req.id == "" { - return apiutil.ErrMissingID - } - - return nil -} - -type connectChannelClientsRequest struct { - channelID string - ClientIDs []string `json:"client_ids,omitempty"` - Types []connections.ConnType `json:"types,omitempty"` -} - -func (req *connectChannelClientsRequest) validate() error { - if req.channelID == "" || strings.TrimSpace(req.channelID) == "" { - return apiutil.ErrMissingID - } - - if len(req.ClientIDs) == 0 { - return apiutil.ErrMissingID - } - - for _, tid := range req.ClientIDs { - if err := api.ValidateUUID(tid); err != nil { - return err - } - } - - if len(req.Types) == 0 { - return apiutil.ErrMissingConnectionType - } - - return nil -} - -type disconnectChannelClientsRequest struct { - channelID string - ClientIds []string `json:"client_ids,omitempty"` - Types []connections.ConnType `json:"types,omitempty"` -} - -func (req *disconnectChannelClientsRequest) validate() error { - if req.channelID == "" { - return apiutil.ErrMissingID - } - - if err := api.ValidateUUID(req.channelID); err != nil { - return err - } - - if len(req.ClientIds) == 0 { - return apiutil.ErrMissingID - } - - for _, tid := range req.ClientIds { - if err := api.ValidateUUID(tid); err != nil { - return err - } - } - - if len(req.Types) == 0 { - return apiutil.ErrMissingConnectionType - } - - return nil -} - -type connectRequest struct { - ChannelIds []string `json:"channel_ids,omitempty"` - ClientIds []string `json:"client_ids,omitempty"` - Types []connections.ConnType `json:"types,omitempty"` -} - -func (req *connectRequest) validate() error { - if len(req.ChannelIds) == 0 { - return apiutil.ErrMissingID - } - for _, cid := range req.ChannelIds { - if strings.TrimSpace(cid) == "" { - return apiutil.ErrMissingChannelID - } - } - - if len(req.ClientIds) == 0 { - return apiutil.ErrMissingID - } - - for _, tid := range req.ClientIds { - if strings.TrimSpace(tid) == "" { - return apiutil.ErrMissingChannelID - } - } - - if len(req.Types) == 0 { - return apiutil.ErrMissingConnectionType - } - - return nil -} - -type disconnectRequest struct { - ChannelIds []string `json:"channel_ids,omitempty"` - ClientIds []string `json:"client_ids,omitempty"` - Types []connections.ConnType `json:"types,omitempty"` -} - -func (req *disconnectRequest) validate() error { - if len(req.ChannelIds) == 0 { - return apiutil.ErrMissingID - } - for _, cid := range req.ChannelIds { - if err := api.ValidateUUID(cid); err != nil { - return err - } - } - - if len(req.ClientIds) == 0 { - return apiutil.ErrMissingID - } - - for _, tid := range req.ClientIds { - if err := api.ValidateUUID(tid); err != nil { - return err - } - } - - if len(req.Types) == 0 { - return apiutil.ErrMissingConnectionType - } - - return nil -} - -type deleteChannelReq struct { - id string -} - -func (req deleteChannelReq) validate() error { - if req.id == "" { - return apiutil.ErrMissingID - } - return nil -} diff --git a/channels/api/http/requests_test.go b/channels/api/http/requests_test.go deleted file mode 100644 index 33dee3693..000000000 --- a/channels/api/http/requests_test.go +++ /dev/null @@ -1,628 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package http - -import ( - "fmt" - "strings" - "testing" - - api "github.com/absmach/magistrala/api/http" - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/channels" - "github.com/absmach/magistrala/internal/testsutil" - "github.com/absmach/magistrala/pkg/connections" - "github.com/stretchr/testify/assert" -) - -func TestCreateChannelReqValidation(t *testing.T) { - cases := []struct { - desc string - req createChannelReq - err error - }{ - { - desc: "valid request", - req: createChannelReq{ - Channel: channels.Channel{ - Name: valid, - Route: valid, - }, - }, - err: nil, - }, - { - desc: "long name", - req: createChannelReq{ - Channel: channels.Channel{ - Name: strings.Repeat("a", api.MaxNameSize+1), - Route: valid, - }, - }, - err: apiutil.ErrNameSize, - }, - { - desc: "invalid route", - req: createChannelReq{ - Channel: channels.Channel{ - Name: valid, - Route: "__invalid", - }, - }, - err: apiutil.ErrInvalidRouteFormat, - }, - { - desc: "uuid as route", - req: createChannelReq{ - Channel: channels.Channel{ - Name: valid, - Route: testsutil.GenerateUUID(t), - }, - }, - err: apiutil.ErrInvalidRouteFormat, - }, - { - desc: "missing channel ID", - req: createChannelReq{ - Channel: channels.Channel{ - ID: " ", - }, - }, - err: apiutil.ErrMissingChannelID, - }, - } - - for _, tc := range cases { - err := tc.req.validate() - assert.Equal(t, tc.err, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - } -} - -func TestCreateChannelsReqValidation(t *testing.T) { - cases := []struct { - desc string - req createChannelsReq - err error - }{ - { - desc: "valid request", - req: createChannelsReq{ - Channels: []channels.Channel{ - { - Name: valid, - Route: valid, - }, - }, - }, - err: nil, - }, - { - desc: "long name", - req: createChannelsReq{ - Channels: []channels.Channel{ - { - Name: strings.Repeat("a", api.MaxNameSize+1), - Route: valid, - }, - }, - }, - err: apiutil.ErrNameSize, - }, - { - desc: "missing channel ID", - req: createChannelsReq{ - Channels: []channels.Channel{ - { - ID: " ", - }, - }, - }, - err: apiutil.ErrMissingChannelID, - }, - { - desc: "empty list", - req: createChannelsReq{ - Channels: []channels.Channel{}, - }, - err: apiutil.ErrEmptyList, - }, - } - - for _, tc := range cases { - err := tc.req.validate() - assert.Equal(t, tc.err, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - } -} - -func TestViewChannelReqValidation(t *testing.T) { - cases := []struct { - desc string - req viewChannelReq - err error - }{ - { - desc: "valid request", - req: viewChannelReq{ - id: valid, - }, - err: nil, - }, - { - desc: "missing ID", - req: viewChannelReq{ - id: "", - }, - err: apiutil.ErrMissingID, - }, - } - - for _, tc := range cases { - err := tc.req.validate() - assert.Equal(t, tc.err, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - } -} - -func TestListChannelsReqValidation(t *testing.T) { - cases := []struct { - desc string - req listChannelsReq - err error - }{ - { - desc: "valid request", - req: listChannelsReq{ - Page: channels.Page{Limit: 10}, - }, - err: nil, - }, - { - desc: "limit is 0", - req: listChannelsReq{ - Page: channels.Page{Limit: 0}, - }, - err: apiutil.ErrLimitSize, - }, - { - desc: "limit is greater than max limit", - req: listChannelsReq{ - Page: channels.Page{Limit: api.MaxLimitSize + 1}, - }, - err: apiutil.ErrLimitSize, - }, - { - desc: "name is too long", - req: listChannelsReq{ - Page: channels.Page{Limit: 10, Name: strings.Repeat("a", api.MaxNameSize+1)}, - }, - err: apiutil.ErrNameSize, - }, - } - for _, tc := range cases { - err := tc.req.validate() - assert.Equal(t, tc.err, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - } -} - -func TestUpdateChannelReqValidate(t *testing.T) { - cases := []struct { - desc string - req updateChannelReq - err error - }{ - { - desc: "valid request", - req: updateChannelReq{ - id: valid, - }, - err: nil, - }, - { - desc: "missing ID", - req: updateChannelReq{ - id: "", - }, - err: apiutil.ErrMissingID, - }, - { - desc: "name is too long", - req: updateChannelReq{ - id: valid, - Name: strings.Repeat("a", api.MaxNameSize+1), - }, - err: apiutil.ErrNameSize, - }, - } - for _, tc := range cases { - err := tc.req.validate() - assert.Equal(t, tc.err, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - } -} - -func TestUpdateChannelTagsReqValidate(t *testing.T) { - cases := []struct { - desc string - req updateChannelTagsReq - err error - }{ - { - desc: "valid request", - req: updateChannelTagsReq{ - id: valid, - Tags: []string{"tag1", "tag2"}, - }, - err: nil, - }, - { - desc: "missing ID", - req: updateChannelTagsReq{ - id: "", - Tags: []string{"tag1", "tag2"}, - }, - err: apiutil.ErrMissingID, - }, - } - for _, tc := range cases { - err := tc.req.validate() - assert.Equal(t, tc.err, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - } -} - -func TestSetChannelsParentGroupReqValidate(t *testing.T) { - cases := []struct { - desc string - req setChannelParentGroupReq - err error - }{ - { - desc: "valid request", - req: setChannelParentGroupReq{ - id: valid, - ParentGroupID: valid, - }, - err: nil, - }, - { - desc: "missing ID", - req: setChannelParentGroupReq{ - id: "", - ParentGroupID: valid, - }, - err: apiutil.ErrMissingID, - }, - { - desc: "missing parent group ID", - req: setChannelParentGroupReq{ - id: valid, - ParentGroupID: "", - }, - err: apiutil.ErrMissingParentGroupID, - }, - } - for _, tc := range cases { - err := tc.req.validate() - assert.Equal(t, tc.err, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - } -} - -func TestRemoveChannelParentGroupReqValidate(t *testing.T) { - cases := []struct { - desc string - req removeChannelParentGroupReq - err error - }{ - { - desc: "valid request", - req: removeChannelParentGroupReq{ - id: valid, - }, - err: nil, - }, - { - desc: "missing ID", - req: removeChannelParentGroupReq{ - id: "", - }, - err: apiutil.ErrMissingID, - }, - } - for _, tc := range cases { - err := tc.req.validate() - assert.Equal(t, tc.err, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - } -} - -func TestChangeChannelStatusReqValidate(t *testing.T) { - cases := []struct { - desc string - req changeChannelStatusReq - err error - }{ - { - desc: "valid request", - req: changeChannelStatusReq{ - id: valid, - }, - err: nil, - }, - { - desc: "missing ID", - req: changeChannelStatusReq{ - id: "", - }, - err: apiutil.ErrMissingID, - }, - } - for _, tc := range cases { - err := tc.req.validate() - assert.Equal(t, tc.err, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - } -} - -func TestConnectChannelClientsReqValidate(t *testing.T) { - cases := []struct { - desc string - req connectChannelClientsRequest - err error - }{ - { - desc: "valid request", - req: connectChannelClientsRequest{ - channelID: valid, - ClientIDs: []string{testsutil.GenerateUUID(t), testsutil.GenerateUUID(t)}, - Types: []connections.ConnType{connections.Publish}, - }, - err: nil, - }, - { - desc: "missing channel ID", - req: connectChannelClientsRequest{ - channelID: "", - ClientIDs: []string{testsutil.GenerateUUID(t), testsutil.GenerateUUID(t)}, - Types: []connections.ConnType{connections.Publish}, - }, - err: apiutil.ErrMissingID, - }, - { - desc: "missing client IDs", - req: connectChannelClientsRequest{ - channelID: valid, - ClientIDs: []string{}, - Types: []connections.ConnType{connections.Publish}, - }, - err: apiutil.ErrMissingID, - }, - { - desc: "missing connection types", - req: connectChannelClientsRequest{ - channelID: valid, - ClientIDs: []string{testsutil.GenerateUUID(t), testsutil.GenerateUUID(t)}, - Types: []connections.ConnType{}, - }, - err: apiutil.ErrMissingConnectionType, - }, - { - desc: "invalid client ID", - req: connectChannelClientsRequest{ - channelID: valid, - ClientIDs: []string{"client1", "invalid"}, - Types: []connections.ConnType{connections.Publish}, - }, - err: apiutil.ErrInvalidIDFormat, - }, - } - for _, tc := range cases { - err := tc.req.validate() - assert.Equal(t, tc.err, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - } -} - -func TestDisconnectChannelClientReqValidate(t *testing.T) { - cases := []struct { - desc string - req disconnectChannelClientsRequest - err error - }{ - { - desc: "valid request", - req: disconnectChannelClientsRequest{ - channelID: testsutil.GenerateUUID(t), - ClientIds: []string{testsutil.GenerateUUID(t), testsutil.GenerateUUID(t)}, - Types: []connections.ConnType{connections.Publish}, - }, - err: nil, - }, - { - desc: "missing channel ID", - req: disconnectChannelClientsRequest{ - channelID: "", - ClientIds: []string{testsutil.GenerateUUID(t), testsutil.GenerateUUID(t)}, - Types: []connections.ConnType{connections.Publish}, - }, - err: apiutil.ErrMissingID, - }, - { - desc: "invalid channel ID", - req: disconnectChannelClientsRequest{ - channelID: "invalid", - ClientIds: []string{testsutil.GenerateUUID(t), testsutil.GenerateUUID(t)}, - Types: []connections.ConnType{connections.Publish}, - }, - err: apiutil.ErrInvalidIDFormat, - }, - { - desc: "missing client IDs", - req: disconnectChannelClientsRequest{ - channelID: testsutil.GenerateUUID(t), - ClientIds: []string{}, - Types: []connections.ConnType{connections.Publish}, - }, - err: apiutil.ErrMissingID, - }, - { - desc: "missing connection types", - req: disconnectChannelClientsRequest{ - channelID: testsutil.GenerateUUID(t), - ClientIds: []string{testsutil.GenerateUUID(t), testsutil.GenerateUUID(t)}, - Types: []connections.ConnType{}, - }, - err: apiutil.ErrMissingConnectionType, - }, - { - desc: "invalid client ID", - req: disconnectChannelClientsRequest{ - channelID: testsutil.GenerateUUID(t), - ClientIds: []string{"client1", "invalid"}, - Types: []connections.ConnType{connections.Publish}, - }, - err: apiutil.ErrInvalidIDFormat, - }, - } - for _, tc := range cases { - err := tc.req.validate() - assert.Equal(t, tc.err, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - } -} - -func TestConnectReqValidate(t *testing.T) { - cases := []struct { - desc string - req connectRequest - err error - }{ - { - desc: "valid request", - req: connectRequest{ - ChannelIds: []string{testsutil.GenerateUUID(t), testsutil.GenerateUUID(t)}, - ClientIds: []string{testsutil.GenerateUUID(t), testsutil.GenerateUUID(t)}, - Types: []connections.ConnType{connections.Publish}, - }, - err: nil, - }, - { - desc: "missing channel IDs", - req: connectRequest{ - ChannelIds: []string{}, - ClientIds: []string{testsutil.GenerateUUID(t), testsutil.GenerateUUID(t)}, - Types: []connections.ConnType{connections.Publish}, - }, - err: apiutil.ErrMissingID, - }, - { - desc: "missing client IDs", - req: connectRequest{ - ChannelIds: []string{testsutil.GenerateUUID(t), testsutil.GenerateUUID(t)}, - ClientIds: []string{}, - Types: []connections.ConnType{connections.Publish}, - }, - err: apiutil.ErrMissingID, - }, - { - desc: "missing connection types", - req: connectRequest{ - ChannelIds: []string{testsutil.GenerateUUID(t), testsutil.GenerateUUID(t)}, - ClientIds: []string{testsutil.GenerateUUID(t), testsutil.GenerateUUID(t)}, - Types: []connections.ConnType{}, - }, - err: apiutil.ErrMissingConnectionType, - }, - } - for _, tc := range cases { - err := tc.req.validate() - assert.Equal(t, tc.err, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - } -} - -func TestDisconnectReqValidate(t *testing.T) { - cases := []struct { - desc string - req disconnectRequest - err error - }{ - { - desc: "valid request", - req: disconnectRequest{ - ChannelIds: []string{testsutil.GenerateUUID(t), testsutil.GenerateUUID(t)}, - ClientIds: []string{testsutil.GenerateUUID(t), testsutil.GenerateUUID(t)}, - Types: []connections.ConnType{connections.Publish}, - }, - err: nil, - }, - { - desc: "missing channel IDs", - req: disconnectRequest{ - ChannelIds: []string{}, - ClientIds: []string{testsutil.GenerateUUID(t), testsutil.GenerateUUID(t)}, - Types: []connections.ConnType{connections.Publish}, - }, - err: apiutil.ErrMissingID, - }, - { - desc: "missing client IDs", - req: disconnectRequest{ - ChannelIds: []string{testsutil.GenerateUUID(t), testsutil.GenerateUUID(t)}, - ClientIds: []string{}, - Types: []connections.ConnType{connections.Publish}, - }, - err: apiutil.ErrMissingID, - }, - { - desc: "missing connection types", - req: disconnectRequest{ - ChannelIds: []string{testsutil.GenerateUUID(t), testsutil.GenerateUUID(t)}, - ClientIds: []string{testsutil.GenerateUUID(t), testsutil.GenerateUUID(t)}, - Types: []connections.ConnType{}, - }, - err: apiutil.ErrMissingConnectionType, - }, - { - desc: "invalid client ID", - req: disconnectRequest{ - ChannelIds: []string{testsutil.GenerateUUID(t), testsutil.GenerateUUID(t)}, - ClientIds: []string{"client1", "invalid"}, - Types: []connections.ConnType{connections.Publish}, - }, - err: apiutil.ErrInvalidIDFormat, - }, - { - desc: "invalid channel ID", - req: disconnectRequest{ - ChannelIds: []string{"invalid", testsutil.GenerateUUID(t)}, - ClientIds: []string{testsutil.GenerateUUID(t), testsutil.GenerateUUID(t)}, - Types: []connections.ConnType{connections.Publish}, - }, - err: apiutil.ErrInvalidIDFormat, - }, - } - for _, tc := range cases { - err := tc.req.validate() - assert.Equal(t, tc.err, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - } -} - -func TestDeleteChannelReqValidate(t *testing.T) { - cases := []struct { - desc string - req deleteChannelReq - err error - }{ - { - desc: "valid request", - req: deleteChannelReq{ - id: valid, - }, - err: nil, - }, - { - desc: "missing ID", - req: deleteChannelReq{ - id: "", - }, - err: apiutil.ErrMissingID, - }, - } - for _, tc := range cases { - err := tc.req.validate() - assert.Equal(t, tc.err, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - } -} diff --git a/channels/api/http/responses.go b/channels/api/http/responses.go deleted file mode 100644 index f6038c7c4..000000000 --- a/channels/api/http/responses.go +++ /dev/null @@ -1,221 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package http - -import ( - "fmt" - "net/http" - - "github.com/absmach/magistrala" - "github.com/absmach/magistrala/channels" -) - -var ( - _ magistrala.Response = (*createChannelRes)(nil) - _ magistrala.Response = (*viewChannelRes)(nil) - _ magistrala.Response = (*channelsPageRes)(nil) - _ magistrala.Response = (*updateChannelRes)(nil) - _ magistrala.Response = (*deleteChannelRes)(nil) - _ magistrala.Response = (*connectChannelClientsRes)(nil) - _ magistrala.Response = (*disconnectChannelClientsRes)(nil) - _ magistrala.Response = (*connectRes)(nil) - _ magistrala.Response = (*disconnectRes)(nil) - _ magistrala.Response = (*changeChannelStatusRes)(nil) -) - -type pageRes struct { - Limit uint64 `json:"limit,omitempty"` - Offset uint64 `json:"offset,omitempty"` - Total uint64 `json:"total"` -} - -type createChannelRes struct { - channels.Channel - created bool -} - -func (res createChannelRes) Code() int { - if res.created { - return http.StatusCreated - } - - return http.StatusOK -} - -func (res createChannelRes) Headers() map[string]string { - if res.created { - return map[string]string{ - "Location": fmt.Sprintf("/channels/%s", res.ID), - } - } - - return map[string]string{} -} - -func (res createChannelRes) Empty() bool { - return false -} - -type viewChannelRes struct { - channels.Channel -} - -func (res viewChannelRes) Code() int { - return http.StatusOK -} - -func (res viewChannelRes) Headers() map[string]string { - return map[string]string{} -} - -func (res viewChannelRes) Empty() bool { - return false -} - -type channelsPageRes struct { - pageRes - Channels []viewChannelRes `json:"channels,omitempty"` -} - -func (res channelsPageRes) Code() int { - return http.StatusOK -} - -func (res channelsPageRes) Headers() map[string]string { - return map[string]string{} -} - -func (res channelsPageRes) Empty() bool { - return false -} - -type changeChannelStatusRes struct { - channels.Channel -} - -func (res changeChannelStatusRes) Code() int { - return http.StatusOK -} - -func (res changeChannelStatusRes) Headers() map[string]string { - return map[string]string{} -} - -func (res changeChannelStatusRes) Empty() bool { - return false -} - -type updateChannelRes struct { - channels.Channel -} - -func (res updateChannelRes) Code() int { - return http.StatusOK -} - -func (res updateChannelRes) Headers() map[string]string { - return map[string]string{} -} - -func (res updateChannelRes) Empty() bool { - return false -} - -type setChannelParentGroupRes struct{} - -func (res setChannelParentGroupRes) Code() int { - return http.StatusOK -} - -func (res setChannelParentGroupRes) Headers() map[string]string { - return map[string]string{} -} - -func (res setChannelParentGroupRes) Empty() bool { - return true -} - -type removeChannelParentGroupRes struct{} - -func (res removeChannelParentGroupRes) Code() int { - return http.StatusNoContent -} - -func (res removeChannelParentGroupRes) Headers() map[string]string { - return map[string]string{} -} - -func (res removeChannelParentGroupRes) Empty() bool { - return true -} - -type deleteChannelRes struct{} - -func (res deleteChannelRes) Code() int { - return http.StatusNoContent -} - -func (res deleteChannelRes) Headers() map[string]string { - return map[string]string{} -} - -func (res deleteChannelRes) Empty() bool { - return true -} - -type connectChannelClientsRes struct{} - -func (res connectChannelClientsRes) Code() int { - return http.StatusCreated -} - -func (res connectChannelClientsRes) Headers() map[string]string { - return map[string]string{} -} - -func (res connectChannelClientsRes) Empty() bool { - return true -} - -type disconnectChannelClientsRes struct{} - -func (res disconnectChannelClientsRes) Code() int { - return http.StatusNoContent -} - -func (res disconnectChannelClientsRes) Headers() map[string]string { - return map[string]string{} -} - -func (res disconnectChannelClientsRes) Empty() bool { - return true -} - -type connectRes struct{} - -func (res connectRes) Code() int { - return http.StatusCreated -} - -func (res connectRes) Headers() map[string]string { - return map[string]string{} -} - -func (res connectRes) Empty() bool { - return true -} - -type disconnectRes struct{} - -func (res disconnectRes) Code() int { - return http.StatusNoContent -} - -func (res disconnectRes) Headers() map[string]string { - return map[string]string{} -} - -func (res disconnectRes) Empty() bool { - return true -} diff --git a/channels/api/http/transport.go b/channels/api/http/transport.go deleted file mode 100644 index 33afe4a83..000000000 --- a/channels/api/http/transport.go +++ /dev/null @@ -1,149 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package http - -import ( - "log/slog" - - "github.com/absmach/magistrala" - api "github.com/absmach/magistrala/api/http" - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/channels" - smqauthn "github.com/absmach/magistrala/pkg/authn" - roleManagerHttp "github.com/absmach/magistrala/pkg/roles/rolemanager/api" - "github.com/go-chi/chi/v5" - kithttp "github.com/go-kit/kit/transport/http" - "github.com/prometheus/client_golang/prometheus/promhttp" - "go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp" -) - -// MakeHandler returns a HTTP handler for Channels API endpoints. -func MakeHandler(svc channels.Service, authn smqauthn.AuthNMiddleware, mux *chi.Mux, logger *slog.Logger, instanceID string, idp magistrala.IDProvider) *chi.Mux { - opts := []kithttp.ServerOption{ - kithttp.ServerErrorEncoder(apiutil.LoggingErrorEncoder(logger, api.EncodeError)), - } - - d := roleManagerHttp.NewDecoder("channelID") - - mux.Route("/{domainID}/channels", func(r chi.Router) { - r.Use(authn.Middleware()) - r.Use(api.RequestIDMiddleware(idp)) - - r.Post("/", otelhttp.NewHandler(kithttp.NewServer( - createChannelEndpoint(svc), - decodeCreateChannelReq, - api.EncodeResponse, - opts..., - ), "create_channel").ServeHTTP) - - r.Post("/bulk", otelhttp.NewHandler(kithttp.NewServer( - createChannelsEndpoint(svc), - decodeCreateChannelsReq, - api.EncodeResponse, - opts..., - ), "create_channels").ServeHTTP) - - r.Get("/", otelhttp.NewHandler(kithttp.NewServer( - listChannelsEndpoint(svc), - decodeListChannels, - api.EncodeResponse, - opts..., - ), "list_channels").ServeHTTP) - - r.Post("/connect", otelhttp.NewHandler(kithttp.NewServer( - connectEndpoint(svc), - decodeConnectRequest, - api.EncodeResponse, - opts..., - ), "connect").ServeHTTP) - - r.Post("/disconnect", otelhttp.NewHandler(kithttp.NewServer( - disconnectEndpoint(svc), - decodeDisconnectRequest, - api.EncodeResponse, - opts..., - ), "disconnect").ServeHTTP) - - r = roleManagerHttp.EntityAvailableActionsRouter(svc, d, r, opts) - - r.Route("/{channelID}", func(r chi.Router) { - r.Get("/", otelhttp.NewHandler(kithttp.NewServer( - viewChannelEndpoint(svc), - decodeViewChannel, - api.EncodeResponse, - opts..., - ), "view_channel").ServeHTTP) - - r.Patch("/", otelhttp.NewHandler(kithttp.NewServer( - updateChannelEndpoint(svc), - decodeUpdateChannel, - api.EncodeResponse, - opts..., - ), "update_channel_name_and_metadata").ServeHTTP) - - r.Patch("/tags", otelhttp.NewHandler(kithttp.NewServer( - updateChannelTagsEndpoint(svc), - decodeUpdateChannelTags, - api.EncodeResponse, - opts..., - ), "update_channel_tag").ServeHTTP) - - r.Delete("/", otelhttp.NewHandler(kithttp.NewServer( - deleteChannelEndpoint(svc), - decodeDeleteChannelReq, - api.EncodeResponse, - opts..., - ), "delete_channel").ServeHTTP) - - r.Post("/enable", otelhttp.NewHandler(kithttp.NewServer( - enableChannelEndpoint(svc), - decodeChangeChannelStatus, - api.EncodeResponse, - opts..., - ), "enable_channel").ServeHTTP) - - r.Post("/disable", otelhttp.NewHandler(kithttp.NewServer( - disableChannelEndpoint(svc), - decodeChangeChannelStatus, - api.EncodeResponse, - opts..., - ), "disable_channel").ServeHTTP) - - r.Post("/parent", otelhttp.NewHandler(kithttp.NewServer( - setChannelParentGroupEndpoint(svc), - decodeSetChannelParentGroupStatus, - api.EncodeResponse, - opts..., - ), "set_channel_parent_group").ServeHTTP) - - r.Delete("/parent", otelhttp.NewHandler(kithttp.NewServer( - removeChannelParentGroupEndpoint(svc), - decodeRemoveChannelParentGroupStatus, - api.EncodeResponse, - opts..., - ), "remove_channel_parent_group").ServeHTTP) - - r.Post("/connect", otelhttp.NewHandler(kithttp.NewServer( - connectChannelClientEndpoint(svc), - decodeConnectChannelClientRequest, - api.EncodeResponse, - opts..., - ), "connect_channel_client").ServeHTTP) - - r.Post("/disconnect", otelhttp.NewHandler(kithttp.NewServer( - disconnectChannelClientsEndpoint(svc), - decodeDisconnectChannelClientsRequest, - api.EncodeResponse, - opts..., - ), "disconnect_channel_client").ServeHTTP) - - roleManagerHttp.EntityRoleMangerRouter(svc, d, r, opts) - }) - }) - - mux.Get("/health", magistrala.Health("channels", instanceID)) - mux.Handle("/metrics", promhttp.Handler()) - - return mux -} diff --git a/channels/builtinroles.go b/channels/builtinroles.go deleted file mode 100644 index fe32f9ba8..000000000 --- a/channels/builtinroles.go +++ /dev/null @@ -1,7 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 -package channels - -import "github.com/absmach/magistrala/pkg/roles" - -const BuiltInRoleAdmin roles.BuiltInRoleName = "admin" diff --git a/channels/cache/channels.go b/channels/cache/channels.go deleted file mode 100644 index ddfb27578..000000000 --- a/channels/cache/channels.go +++ /dev/null @@ -1,82 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package cache - -import ( - "context" - "time" - - "github.com/absmach/magistrala/channels" - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - "github.com/redis/go-redis/v9" -) - -var ( - ErrEmptyDomainID = errors.New("domain ID is empty") - ErrEmptyChannelID = errors.New("channel ID is empty") - ErrEmptyChannelRoute = errors.New("channel route is empty") -) - -type channelsCache struct { - client *redis.Client - duration time.Duration -} - -func NewChannelsCache(client *redis.Client, duration time.Duration) channels.Cache { - return &channelsCache{ - client: client, - duration: duration, - } -} - -func (cc *channelsCache) Save(ctx context.Context, route, domainID, channelID string) error { - key, err := encodeKey(domainID, route) - if err != nil { - return errors.Wrap(repoerr.ErrCreateEntity, err) - } - if channelID == "" { - return errors.Wrap(repoerr.ErrCreateEntity, ErrEmptyChannelID) - } - if err := cc.client.Set(ctx, key, channelID, cc.duration).Err(); err != nil { - return errors.Wrap(repoerr.ErrCreateEntity, err) - } - - return nil -} - -func (cc *channelsCache) ID(ctx context.Context, channelRoute, domainID string) (string, error) { - key, err := encodeKey(domainID, channelRoute) - if err != nil { - return "", errors.Wrap(repoerr.ErrNotFound, err) - } - id, err := cc.client.Get(ctx, key).Result() - if err != nil { - return "", errors.Wrap(repoerr.ErrNotFound, err) - } - - return id, nil -} - -func (cc *channelsCache) Remove(ctx context.Context, channelRoute, domainID string) error { - key, err := encodeKey(domainID, channelRoute) - if err != nil { - return errors.Wrap(repoerr.ErrRemoveEntity, err) - } - if err := cc.client.Del(ctx, key).Err(); err != nil { - return errors.Wrap(repoerr.ErrRemoveEntity, err) - } - - return nil -} - -func encodeKey(domainID, channelRoute string) (string, error) { - if domainID == "" { - return "", ErrEmptyDomainID - } - if channelRoute == "" { - return "", ErrEmptyChannelRoute - } - return domainID + ":" + channelRoute, nil -} diff --git a/channels/cache/channels_test.go b/channels/cache/channels_test.go deleted file mode 100644 index 51931b827..000000000 --- a/channels/cache/channels_test.go +++ /dev/null @@ -1,186 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package cache_test - -import ( - "context" - "fmt" - "testing" - "time" - - "github.com/absmach/magistrala/channels" - "github.com/absmach/magistrala/channels/cache" - "github.com/absmach/magistrala/internal/testsutil" - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - "github.com/redis/go-redis/v9" - "github.com/stretchr/testify/assert" -) - -var ( - testRoute = "test-route" - nonExistent = "non-existing" -) - -func setupChannelsClient(t *testing.T) channels.Cache { - opts, err := redis.ParseURL(redisURL) - assert.Nil(t, err, fmt.Sprintf("got unexpected error on parsing redis URL: %s", err)) - redisClient := redis.NewClient(opts) - - return cache.NewChannelsCache(redisClient, 10*time.Minute) -} - -func TestSave(t *testing.T) { - cc := setupChannelsClient(t) - - route := testRoute - domainID := testsutil.GenerateUUID(t) - - cases := []struct { - desc string - domainID string - channelID string - channelRoute string - err error - }{ - { - desc: "Save successfully", - domainID: domainID, - channelID: testsutil.GenerateUUID(t), - channelRoute: route, - err: nil, - }, - { - desc: "Save with empty domain ID", - domainID: "", - channelID: testsutil.GenerateUUID(t), - channelRoute: route, - err: cache.ErrEmptyDomainID, - }, - { - desc: "Save with empty channel ID", - domainID: domainID, - channelID: "", - channelRoute: route, - err: cache.ErrEmptyChannelID, - }, - { - desc: "Save with empty channel route", - domainID: domainID, - channelID: testsutil.GenerateUUID(t), - channelRoute: "", - err: cache.ErrEmptyChannelRoute, - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := cc.Save(context.Background(), tc.channelRoute, tc.domainID, tc.channelID) - assert.True(t, errors.Contains(err, tc.err)) - }) - } -} - -func TestID(t *testing.T) { - cc := setupChannelsClient(t) - - domainID := testsutil.GenerateUUID(t) - route := testRoute - id := testsutil.GenerateUUID(t) - - err := cc.Save(context.Background(), route, domainID, id) - assert.Nil(t, err, fmt.Sprintf("got unexpected error on saving channel ID: %s", err)) - - cases := []struct { - desc string - domainID string - channelRoute string - channelID string - err error - }{ - { - desc: "Retrieve existing channel", - domainID: domainID, - channelRoute: route, - channelID: id, - err: nil, - }, - { - desc: "Retrieve non-existing channel", - domainID: domainID, - channelRoute: nonExistent, - channelID: "", - err: repoerr.ErrNotFound, - }, - { - desc: "Retrieve with empty domain ID", - domainID: "", - channelRoute: route, - channelID: "", - err: cache.ErrEmptyDomainID, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - id, err := cc.ID(context.Background(), tc.channelRoute, tc.domainID) - assert.Equal(t, tc.channelID, id, fmt.Sprintf("expected channel ID '%s' got '%s'", tc.channelID, id)) - assert.True(t, errors.Contains(err, tc.err)) - }) - } -} - -func TestRemove(t *testing.T) { - cc := setupChannelsClient(t) - - domainID := testsutil.GenerateUUID(t) - route := testRoute - id := testsutil.GenerateUUID(t) - - err := cc.Save(context.Background(), domainID, route, id) - assert.Nil(t, err, fmt.Sprintf("got unexpected error on saving channel ID: %s", err)) - - cases := []struct { - desc string - domainID string - channelRoute string - err error - }{ - { - desc: "Remove existing channel", - domainID: domainID, - channelRoute: route, - err: nil, - }, - { - desc: "Remove non-existing channel", - domainID: domainID, - channelRoute: nonExistent, - err: nil, - }, - { - desc: "Remove with empty domain ID", - domainID: "", - channelRoute: route, - err: cache.ErrEmptyDomainID, - }, - { - desc: "Remove with empty channel route", - domainID: domainID, - channelRoute: "", - err: cache.ErrEmptyChannelRoute, - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := cc.Remove(context.Background(), tc.channelRoute, tc.domainID) - assert.True(t, errors.Contains(err, tc.err)) - - if tc.err == nil { - id, err := cc.ID(context.Background(), tc.channelRoute, tc.domainID) - assert.Equal(t, "", id, fmt.Sprintf("expected channel ID to be empty after removal, got '%s'", id)) - assert.True(t, errors.Contains(err, repoerr.ErrNotFound)) - } - }) - } -} diff --git a/channels/cache/doc.go b/channels/cache/doc.go deleted file mode 100644 index ca08f7e1c..000000000 --- a/channels/cache/doc.go +++ /dev/null @@ -1,6 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package cache contains the domain concept definitions needed to -// support Magistrala Channels cache service functionality. -package cache diff --git a/channels/cache/setup_test.go b/channels/cache/setup_test.go deleted file mode 100644 index 716f0672c..000000000 --- a/channels/cache/setup_test.go +++ /dev/null @@ -1,61 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package cache_test - -import ( - "context" - "fmt" - "log" - "os" - "testing" - - "github.com/ory/dockertest/v3" - "github.com/ory/dockertest/v3/docker" - "github.com/redis/go-redis/v9" -) - -var ( - redisClient *redis.Client - redisURL string -) - -func TestMain(m *testing.M) { - pool, err := dockertest.NewPool("") - if err != nil { - log.Fatalf("Could not connect to docker: %s", err) - } - - container, err := pool.RunWithOptions(&dockertest.RunOptions{ - Repository: "redis", - Tag: "7.2.4-alpine", - }, func(config *docker.HostConfig) { - config.AutoRemove = true - config.RestartPolicy = docker.RestartPolicy{Name: "no"} - }) - if err != nil { - log.Fatalf("Could not start container: %s", err) - } - - redisURL = fmt.Sprintf("redis://localhost:%s/0", container.GetPort("6379/tcp")) - opts, err := redis.ParseURL(redisURL) - if err != nil { - log.Fatalf("Could not parse redis URL: %s", err) - } - - if err := pool.Retry(func() error { - redisClient = redis.NewClient(opts) - - return redisClient.Ping(context.Background()).Err() - }); err != nil { - log.Fatalf("Could not connect to docker: %s", err) - } - - code := m.Run() - - if err := pool.Purge(container); err != nil { - log.Fatalf("Could not purge container: %s", err) - } - - os.Exit(code) -} diff --git a/channels/channels.go b/channels/channels.go deleted file mode 100644 index ae1a316ef..000000000 --- a/channels/channels.go +++ /dev/null @@ -1,240 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package channels - -import ( - "context" - "strings" - "time" - - "github.com/absmach/magistrala/internal/nullable" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/connections" - "github.com/absmach/magistrala/pkg/roles" -) - -// Metadata represents arbitrary JSON. -type Metadata map[string]any - -// Channel represents a Magistrala "communication topic". This topic -// contains the clients that can exchange messages between each other. -type Channel struct { - ID string `json:"id"` - Name string `json:"name,omitempty"` - Tags []string `json:"tags,omitempty"` - ParentGroup string `json:"parent_group_id,omitempty"` - Domain string `json:"domain_id,omitempty"` - Route string `json:"route,omitempty"` - Metadata Metadata `json:"metadata,omitempty"` - CreatedBy string `json:"created_by,omitempty"` - CreatedAt time.Time `json:"created_at,omitempty"` - UpdatedAt time.Time `json:"updated_at,omitempty"` - UpdatedBy string `json:"updated_by,omitempty"` - Status Status `json:"status,omitempty"` // 1 for enabled, 0 for disabled - // Extended - ParentGroupPath string `json:"parent_group_path,omitempty"` - RoleID string `json:"role_id,omitempty"` - RoleName string `json:"role_name,omitempty"` - Actions []string `json:"actions,omitempty"` - AccessType string `json:"access_type,omitempty"` - AccessProviderId string `json:"access_provider_id,omitempty"` - AccessProviderRoleId string `json:"access_provider_role_id,omitempty"` - AccessProviderRoleName string `json:"access_provider_role_name,omitempty"` - AccessProviderRoleActions []string `json:"access_provider_role_actions,omitempty"` - ConnectionTypes []connections.ConnType `json:"connection_types,omitempty"` - MemberId string `json:"member_id,omitempty"` - Roles []roles.MemberRoleActions `json:"roles,omitempty"` -} - -type Operator uint8 - -const ( - OrOp Operator = iota - AndOp -) - -type TagsQuery struct { - Elements []string - Operator Operator -} - -func ToTagsQuery(s string) TagsQuery { - switch { - case strings.Contains(s, "+"): - elements := strings.Split(s, "+") - for i := range elements { - elements[i] = strings.TrimSpace(elements[i]) - } - return TagsQuery{Elements: elements, Operator: AndOp} - case strings.Contains(s, ","): - elements := strings.Split(s, ",") - for i := range elements { - elements[i] = strings.TrimSpace(elements[i]) - } - return TagsQuery{Elements: elements, Operator: OrOp} - default: - return TagsQuery{Elements: []string{s}, Operator: OrOp} - } -} - -type Page struct { - Total uint64 `json:"total"` - Offset uint64 `json:"offset"` - Limit uint64 `json:"limit"` - OnlyTotal bool `json:"only_total"` - Order string `json:"order,omitempty"` - Dir string `json:"dir,omitempty"` - ID string `json:"id,omitempty"` - Name string `json:"name,omitempty"` - Metadata Metadata `json:"metadata,omitempty"` - Domain string `json:"domain,omitempty"` - Tags TagsQuery `json:"tags,omitempty"` - Status Status `json:"status,omitempty"` - Group nullable.Value[string] `json:"group,omitempty"` - Client string `json:"client,omitempty"` - ConnectionType string `json:"connection_type,omitempty"` - RoleName string `json:"role_name,omitempty"` - RoleID string `json:"role_id,omitempty"` - Actions []string `json:"actions,omitempty"` - AccessType string `json:"access_type,omitempty"` - IDs []string `json:"-"` - CreatedFrom time.Time `json:"created_from,omitempty"` - CreatedTo time.Time `json:"created_to,omitempty"` -} - -// ChannelsPage contains page related metadata as well as list of channels that -// belong to this page. -type ChannelsPage struct { - Page - Channels []Channel -} - -type Connection struct { - ClientID string - ChannelID string - DomainID string - Type connections.ConnType -} - -type AuthzReq struct { - DomainID string - ChannelID string - ClientID string - ClientType string - Type connections.ConnType -} - -type Service interface { - // CreateChannels adds channels to the user. - CreateChannels(ctx context.Context, session authn.Session, channels ...Channel) ([]Channel, []roles.RoleProvision, error) - - // ViewChannel retrieves data about the channel identified by the provided - // ID, that belongs to the user. - ViewChannel(ctx context.Context, session authn.Session, id string, withRoles bool) (Channel, error) - - // UpdateChannel updates the channel identified by the provided ID, that - // belongs to the user. - UpdateChannel(ctx context.Context, session authn.Session, channel Channel) (Channel, error) - - // UpdateChannelTags updates the channel's tags. - UpdateChannelTags(ctx context.Context, session authn.Session, channel Channel) (Channel, error) - - EnableChannel(ctx context.Context, session authn.Session, id string) (Channel, error) - - DisableChannel(ctx context.Context, session authn.Session, id string) (Channel, error) - - // ListChannels retrieves data about subset of channels that belongs to the user. - ListChannels(ctx context.Context, session authn.Session, pm Page) (ChannelsPage, error) - - // ListUserChannels retrieves data about subset of channels that belong to the specified user. - ListUserChannels(ctx context.Context, session authn.Session, userID string, pm Page) (ChannelsPage, error) - - // RemoveChannel removes the client identified by the provided ID, that - // belongs to the user. - RemoveChannel(ctx context.Context, session authn.Session, id string) error - - // Connect adds clients to the channels list of connected clients. - Connect(ctx context.Context, session authn.Session, chIDs, clIDs []string, connType []connections.ConnType) error - - // Disconnect removes clients from the channels list of connected clients. - Disconnect(ctx context.Context, session authn.Session, chIDs, clIDs []string, connType []connections.ConnType) error - - SetParentGroup(ctx context.Context, session authn.Session, parentGroupID string, id string) error - - RemoveParentGroup(ctx context.Context, session authn.Session, id string) error - - roles.RoleManager -} - -// ChannelRepository specifies a channel persistence API. -type Repository interface { - // Save persists multiple channels. Channels are saved using a transaction. If one channel - // fails then none will be saved. Successful operation is indicated by non-nil error response. - Save(ctx context.Context, chs ...Channel) ([]Channel, error) - - // Update performs an update to the existing channel. - Update(ctx context.Context, c Channel) (Channel, error) - - UpdateTags(ctx context.Context, ch Channel) (Channel, error) - - ChangeStatus(ctx context.Context, channel Channel) (Channel, error) - - // RetrieveUserChannels retrieves the channel of given domainID and userID. - RetrieveUserChannels(ctx context.Context, domainID, userID string, pm Page) (ChannelsPage, error) - - // RetrieveByID retrieves the channel having the provided identifier - RetrieveByID(ctx context.Context, id string) (Channel, error) - - // RetrieveByRoute retrieves the channel having the provided route - RetrieveByRoute(ctx context.Context, route, domainID string) (Channel, error) - - // RetrieveByIDWithRoles retrieves channel by its unique ID along with member roles. - RetrieveByIDWithRoles(ctx context.Context, id, memberID string) (Channel, error) - - // RetrieveAll retrieves the subset of channels. - RetrieveAll(ctx context.Context, pm Page) (ChannelsPage, error) - - // Remove removes the channel having the provided identifier - Remove(ctx context.Context, ids ...string) error - - // SetParentGroup set parent group id to a given channel id - SetParentGroup(ctx context.Context, ch Channel) error - - // RemoveParentGroup remove parent group id fr given chanel id - RemoveParentGroup(ctx context.Context, ch Channel) error - - AddConnections(ctx context.Context, conns []Connection) error - - RemoveConnections(ctx context.Context, conns []Connection) error - - CheckConnection(ctx context.Context, conn Connection) error - - ClientAuthorize(ctx context.Context, conn Connection) error - - ChannelConnectionsCount(ctx context.Context, id string) (uint64, error) - - DoesChannelHaveConnections(ctx context.Context, id string) (bool, error) - - RemoveClientConnections(ctx context.Context, clientID string) error - - RemoveChannelConnections(ctx context.Context, channelID string) error - - RetrieveParentGroupChannels(ctx context.Context, parentGroupID string) ([]Channel, error) - - UnsetParentGroupFromChannels(ctx context.Context, parentGroupID string) error - - roles.Repository -} - -// Cache contains channels caching interface. -type Cache interface { - // Save stores the channelID for the given domain ID and channel route. - Save(ctx context.Context, channelRoute, domainID, channelID string) error - - // ID retrieves the channelID for the given domain ID and channel route. - ID(ctx context.Context, channelRoute, domainID string) (string, error) - - // Remove removes the channel ID for the given domain ID and channel route. - Remove(ctx context.Context, channelRoute, domainID string) error -} diff --git a/channels/errors.go b/channels/errors.go deleted file mode 100644 index b8bde6385..000000000 --- a/channels/errors.go +++ /dev/null @@ -1,17 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package channels - -import "errors" - -var ( - // ErrInvalidStatus indicates invalid status. - ErrInvalidStatus = errors.New("invalid channels status") - - // ErrEnableChannel indicates error in enabling channel. - ErrEnableChannel = errors.New("failed to enable channel") - - // ErrDisableChannel indicates error in disabling channel. - ErrDisableChannel = errors.New("failed to disable channel") -) diff --git a/channels/events/doc.go b/channels/events/doc.go deleted file mode 100644 index d32b58f28..000000000 --- a/channels/events/doc.go +++ /dev/null @@ -1,6 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package events provides the domain concept definitions -// needed to support clients events functionality. -package events diff --git a/channels/events/events.go b/channels/events/events.go deleted file mode 100644 index 2fb0abf55..000000000 --- a/channels/events/events.go +++ /dev/null @@ -1,384 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package events - -import ( - "time" - - "github.com/absmach/magistrala/channels" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/connections" - "github.com/absmach/magistrala/pkg/events" - "github.com/absmach/magistrala/pkg/roles" -) - -const ( - channelPrefix = "channel." - channelCreate = channelPrefix + "create" - channelUpdate = channelPrefix + "update" - channelUpdateTags = channelPrefix + "update_tags" - channelEnable = channelPrefix + "enable" - channelDisable = channelPrefix + "disable" - channelRemove = channelPrefix + "remove" - channelView = channelPrefix + "view" - channelList = channelPrefix + "list" - channelListByUser = channelPrefix + "list_by_user" - channelConnect = channelPrefix + "connect" - channelDisconnect = channelPrefix + "disconnect" - channelSetParent = channelPrefix + "set_parent" - channelRemoveParent = channelPrefix + "remove_parent" -) - -var ( - _ events.Event = (*createChannelEvent)(nil) - _ events.Event = (*updateChannelEvent)(nil) - _ events.Event = (*changeChannelStatusEvent)(nil) - _ events.Event = (*viewChannelEvent)(nil) - _ events.Event = (*listChannelEvent)(nil) - _ events.Event = (*removeChannelEvent)(nil) - _ events.Event = (*connectEvent)(nil) - _ events.Event = (*disconnectEvent)(nil) -) - -type createChannelEvent struct { - channels.Channel - rolesProvisioned []roles.RoleProvision - authn.Session - requestID string -} - -func (cce createChannelEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": channelCreate, - "id": cce.ID, - "roles_provisioned": cce.rolesProvisioned, - "route": cce.Route, - "status": cce.Status.String(), - "created_at": cce.CreatedAt, - "domain": cce.DomainID, - "user_id": cce.UserID, - "token_type": cce.Type.String(), - "super_admin": cce.SuperAdmin, - "request_id": cce.requestID, - } - - if cce.Name != "" { - val["name"] = cce.Name - } - if len(cce.Tags) > 0 { - val["tags"] = cce.Tags - } - if cce.Metadata != nil { - val["metadata"] = cce.Metadata - } - - return val, nil -} - -type updateChannelEvent struct { - channels.Channel - authn.Session - operation string - requestID string -} - -func (uce updateChannelEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": uce.operation, - "updated_at": uce.UpdatedAt, - "updated_by": uce.UpdatedBy, - "domain": uce.DomainID, - "user_id": uce.UserID, - "token_type": uce.Type.String(), - "super_admin": uce.SuperAdmin, - "request_id": uce.requestID, - } - - if uce.ID != "" { - val["id"] = uce.ID - } - if uce.Route != "" { - val["route"] = uce.Route - } - if uce.Name != "" { - val["name"] = uce.Name - } - if len(uce.Tags) > 0 { - val["tags"] = uce.Tags - } - if uce.Metadata != nil { - val["metadata"] = uce.Metadata - } - if !uce.CreatedAt.IsZero() { - val["created_at"] = uce.CreatedAt - } - if uce.Status.String() != "" { - val["status"] = uce.Status.String() - } - - return val, nil -} - -type changeChannelStatusEvent struct { - id string - operation string - status string - updatedAt time.Time - updatedBy string - authn.Session - requestID string -} - -func (cse changeChannelStatusEvent) Encode() (map[string]any, error) { - return map[string]any{ - "operation": cse.operation, - "id": cse.id, - "status": cse.status, - "updated_at": cse.updatedAt, - "updated_by": cse.updatedBy, - "domain": cse.DomainID, - "user_id": cse.UserID, - "token_type": cse.Type.String(), - "super_admin": cse.SuperAdmin, - "request_id": cse.requestID, - }, nil -} - -type viewChannelEvent struct { - channels.Channel - authn.Session - requestID string -} - -func (vce viewChannelEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": channelView, - "id": vce.ID, - "domain": vce.DomainID, - "user_id": vce.UserID, - "token_type": vce.Type.String(), - "super_admin": vce.SuperAdmin, - "request_id": vce.requestID, - } - - if vce.Name != "" { - val["name"] = vce.Name - } - if vce.Route != "" { - val["route"] = vce.Route - } - if len(vce.Tags) > 0 { - val["tags"] = vce.Tags - } - if vce.Metadata != nil { - val["metadata"] = vce.Metadata - } - if !vce.CreatedAt.IsZero() { - val["created_at"] = vce.CreatedAt - } - if !vce.UpdatedAt.IsZero() { - val["updated_at"] = vce.UpdatedAt - } - if vce.UpdatedBy != "" { - val["updated_by"] = vce.UpdatedBy - } - if vce.Status.String() != "" { - val["status"] = vce.Status.String() - } - - return val, nil -} - -type listChannelEvent struct { - channels.Page - authn.Session - requestID string -} - -func (lce listChannelEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": channelList, - "total": lce.Total, - "offset": lce.Offset, - "limit": lce.Limit, - "domain": lce.DomainID, - "user_id": lce.UserID, - "token_type": lce.Type.String(), - "super_admin": lce.SuperAdmin, - "request_id": lce.requestID, - } - - if lce.Name != "" { - val["name"] = lce.Name - } - if lce.Order != "" { - val["order"] = lce.Order - } - if lce.Dir != "" { - val["dir"] = lce.Dir - } - if lce.Metadata != nil { - val["metadata"] = lce.Metadata - } - if len(lce.Tags.Elements) > 0 { - val["tag"] = lce.Tags.Elements - } - if lce.Status.String() != "" { - val["status"] = lce.Status.String() - } - if len(lce.IDs) > 0 { - val["ids"] = lce.IDs - } - - return val, nil -} - -type listUserChannelsEvent struct { - userID string - channels.Page - authn.Session - requestID string -} - -func (luce listUserChannelsEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": channelListByUser, - "req_user_id": luce.userID, - "total": luce.Total, - "offset": luce.Offset, - "limit": luce.Limit, - "domain": luce.DomainID, - "user_id": luce.UserID, - "token_type": luce.Type.String(), - "super_admin": luce.SuperAdmin, - "request_id": luce.requestID, - } - - if luce.Name != "" { - val["name"] = luce.Name - } - if luce.Order != "" { - val["order"] = luce.Order - } - if luce.Dir != "" { - val["dir"] = luce.Dir - } - if luce.Metadata != nil { - val["metadata"] = luce.Metadata - } - if luce.Domain != "" { - val["domain"] = luce.Domain - } - if len(luce.Tags.Elements) > 0 { - val["tag"] = luce.Tags.Elements - } - if luce.Status.String() != "" { - val["status"] = luce.Status.String() - } - if len(luce.IDs) > 0 { - val["ids"] = luce.IDs - } - - return val, nil -} - -type removeChannelEvent struct { - id string - authn.Session - requestID string -} - -func (dce removeChannelEvent) Encode() (map[string]any, error) { - return map[string]any{ - "operation": channelRemove, - "id": dce.id, - "domain": dce.DomainID, - "user_id": dce.UserID, - "token_type": dce.Type.String(), - "super_admin": dce.SuperAdmin, - "request_id": dce.requestID, - }, nil -} - -type connectEvent struct { - chIDs []string - thIDs []string - types []connections.ConnType - authn.Session - requestID string -} - -func (ce connectEvent) Encode() (map[string]any, error) { - return map[string]any{ - "operation": channelConnect, - "client_ids": ce.thIDs, - "channel_ids": ce.chIDs, - "types": ce.types, - "domain": ce.DomainID, - "user_id": ce.UserID, - "token_type": ce.Type.String(), - "super_admin": ce.SuperAdmin, - "request_id": ce.requestID, - }, nil -} - -type disconnectEvent struct { - chIDs []string - thIDs []string - types []connections.ConnType - authn.Session - requestID string -} - -func (de disconnectEvent) Encode() (map[string]any, error) { - return map[string]any{ - "operation": channelDisconnect, - "client_ids": de.thIDs, - "channel_ids": de.chIDs, - "types": de.types, - "domain": de.DomainID, - "user_id": de.UserID, - "token_type": de.Type.String(), - "super_admin": de.SuperAdmin, - "request_id": de.requestID, - }, nil -} - -type setParentGroupEvent struct { - id string - parentGroupID string - authn.Session - requestID string -} - -func (spge setParentGroupEvent) Encode() (map[string]any, error) { - return map[string]any{ - "operation": channelSetParent, - "id": spge.id, - "parent_group_id": spge.parentGroupID, - "domain": spge.DomainID, - "user_id": spge.UserID, - "token_type": spge.Type.String(), - "super_admin": spge.SuperAdmin, - "request_id": spge.requestID, - }, nil -} - -type removeParentGroupEvent struct { - id string - authn.Session - requestID string -} - -func (rpge removeParentGroupEvent) Encode() (map[string]any, error) { - return map[string]any{ - "operation": channelRemoveParent, - "id": rpge.id, - "domain": rpge.DomainID, - "user_id": rpge.UserID, - "token_type": rpge.Type.String(), - "super_admin": rpge.SuperAdmin, - "request_id": rpge.requestID, - }, nil -} diff --git a/channels/events/streams.go b/channels/events/streams.go deleted file mode 100644 index c15a8584c..000000000 --- a/channels/events/streams.go +++ /dev/null @@ -1,300 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package events - -import ( - "context" - - "github.com/absmach/magistrala/channels" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/connections" - "github.com/absmach/magistrala/pkg/events" - "github.com/absmach/magistrala/pkg/events/store" - "github.com/absmach/magistrala/pkg/roles" - rmEvents "github.com/absmach/magistrala/pkg/roles/rolemanager/events" - "github.com/go-chi/chi/v5/middleware" -) - -const ( - magistralaPrefix = "magistrala." - createStream = magistralaPrefix + channelCreate - updateStream = magistralaPrefix + channelUpdate - updateTagsStream = magistralaPrefix + channelUpdateTags - enableStream = magistralaPrefix + channelEnable - disableStream = magistralaPrefix + channelDisable - removeStream = magistralaPrefix + channelRemove - viewStream = magistralaPrefix + channelView - listStream = magistralaPrefix + channelList - listByUserStream = magistralaPrefix + channelListByUser - connectStream = magistralaPrefix + channelConnect - disconnectStream = magistralaPrefix + channelDisconnect - setParentStream = magistralaPrefix + channelSetParent - removeParentStream = magistralaPrefix + channelRemoveParent -) - -var _ channels.Service = (*eventStore)(nil) - -type eventStore struct { - events.Publisher - svc channels.Service - rmEvents.RoleManagerEventStore -} - -// NewEventStoreMiddleware returns wrapper around clients service that sends -// events to event store. -func NewEventStoreMiddleware(ctx context.Context, svc channels.Service, url string) (channels.Service, error) { - publisher, err := store.NewPublisher(ctx, url, "channels-es-pub") - if err != nil { - return nil, err - } - - rolesSvcEventStoreMiddleware := rmEvents.NewRoleManagerEventStore("channels", channelPrefix, svc, publisher) - return &eventStore{ - svc: svc, - Publisher: publisher, - RoleManagerEventStore: rolesSvcEventStoreMiddleware, - }, nil -} - -func (es *eventStore) CreateChannels(ctx context.Context, session authn.Session, chs ...channels.Channel) ([]channels.Channel, []roles.RoleProvision, error) { - chs, rps, err := es.svc.CreateChannels(ctx, session, chs...) - if err != nil { - return chs, rps, err - } - - for _, ch := range chs { - event := createChannelEvent{ - Channel: ch, - rolesProvisioned: rps, - Session: session, - requestID: middleware.GetReqID(ctx), - } - if err := es.Publish(ctx, createStream, event); err != nil { - return chs, rps, err - } - } - - return chs, rps, nil -} - -func (es *eventStore) UpdateChannel(ctx context.Context, session authn.Session, ch channels.Channel) (channels.Channel, error) { - ch, err := es.svc.UpdateChannel(ctx, session, ch) - if err != nil { - return ch, err - } - - event := updateChannelEvent{ - Channel: ch, - Session: session, - operation: channelUpdate, - requestID: middleware.GetReqID(ctx), - } - if err := es.Publish(ctx, updateStream, event); err != nil { - return ch, err - } - - return ch, nil -} - -func (es *eventStore) UpdateChannelTags(ctx context.Context, session authn.Session, ch channels.Channel) (channels.Channel, error) { - ch, err := es.svc.UpdateChannelTags(ctx, session, ch) - if err != nil { - return ch, err - } - - event := updateChannelEvent{ - Channel: ch, - Session: session, - operation: channelUpdateTags, - requestID: middleware.GetReqID(ctx), - } - if err := es.Publish(ctx, updateTagsStream, event); err != nil { - return ch, err - } - - return ch, nil -} - -func (es *eventStore) ViewChannel(ctx context.Context, session authn.Session, id string, withRoles bool) (channels.Channel, error) { - chann, err := es.svc.ViewChannel(ctx, session, id, withRoles) - if err != nil { - return chann, err - } - - event := viewChannelEvent{ - Channel: chann, - Session: session, - requestID: middleware.GetReqID(ctx), - } - if err := es.Publish(ctx, viewStream, event); err != nil { - return chann, err - } - - return chann, nil -} - -func (es *eventStore) ListChannels(ctx context.Context, session authn.Session, pm channels.Page) (channels.ChannelsPage, error) { - cp, err := es.svc.ListChannels(ctx, session, pm) - if err != nil { - return cp, err - } - event := listChannelEvent{ - Page: pm, - Session: session, - requestID: middleware.GetReqID(ctx), - } - if err := es.Publish(ctx, listStream, event); err != nil { - return cp, err - } - - return cp, nil -} - -func (es *eventStore) ListUserChannels(ctx context.Context, session authn.Session, userID string, pm channels.Page) (channels.ChannelsPage, error) { - cp, err := es.svc.ListUserChannels(ctx, session, userID, pm) - if err != nil { - return cp, err - } - event := listUserChannelsEvent{ - userID: userID, - Page: pm, - Session: session, - requestID: middleware.GetReqID(ctx), - } - if err := es.Publish(ctx, listByUserStream, event); err != nil { - return cp, err - } - - return cp, nil -} - -func (es *eventStore) EnableChannel(ctx context.Context, session authn.Session, id string) (channels.Channel, error) { - ch, err := es.svc.EnableChannel(ctx, session, id) - if err != nil { - return ch, err - } - - return es.changeStatus(ctx, session, channelEnable, enableStream, ch) -} - -func (es *eventStore) DisableChannel(ctx context.Context, session authn.Session, id string) (channels.Channel, error) { - ch, err := es.svc.DisableChannel(ctx, session, id) - if err != nil { - return ch, err - } - - return es.changeStatus(ctx, session, channelDisable, disableStream, ch) -} - -func (es *eventStore) changeStatus(ctx context.Context, session authn.Session, operation, stream string, ch channels.Channel) (channels.Channel, error) { - event := changeChannelStatusEvent{ - id: ch.ID, - operation: operation, - updatedAt: ch.UpdatedAt, - updatedBy: ch.UpdatedBy, - status: ch.Status.String(), - Session: session, - requestID: middleware.GetReqID(ctx), - } - if err := es.Publish(ctx, stream, event); err != nil { - return ch, err - } - - return ch, nil -} - -func (es *eventStore) RemoveChannel(ctx context.Context, session authn.Session, id string) error { - if err := es.svc.RemoveChannel(ctx, session, id); err != nil { - return err - } - - event := removeChannelEvent{ - id: id, - Session: session, - requestID: middleware.GetReqID(ctx), - } - - if err := es.Publish(ctx, removeStream, event); err != nil { - return err - } - - return nil -} - -func (es *eventStore) Connect(ctx context.Context, session authn.Session, chIDs, thIDs []string, connTypes []connections.ConnType) error { - if err := es.svc.Connect(ctx, session, chIDs, thIDs, connTypes); err != nil { - return err - } - - event := connectEvent{ - chIDs: chIDs, - thIDs: thIDs, - types: connTypes, - Session: session, - requestID: middleware.GetReqID(ctx), - } - - if err := es.Publish(ctx, connectStream, event); err != nil { - return err - } - - return nil -} - -func (es *eventStore) Disconnect(ctx context.Context, session authn.Session, chIDs, thIDs []string, connTypes []connections.ConnType) error { - if err := es.svc.Disconnect(ctx, session, chIDs, thIDs, connTypes); err != nil { - return err - } - - event := disconnectEvent{ - chIDs: chIDs, - thIDs: thIDs, - types: connTypes, - Session: session, - requestID: middleware.GetReqID(ctx), - } - - if err := es.Publish(ctx, disconnectStream, event); err != nil { - return err - } - - return nil -} - -func (es *eventStore) SetParentGroup(ctx context.Context, session authn.Session, parentGroupID string, id string) (err error) { - if err := es.svc.SetParentGroup(ctx, session, parentGroupID, id); err != nil { - return err - } - - event := setParentGroupEvent{ - parentGroupID: parentGroupID, - id: id, - Session: session, - requestID: middleware.GetReqID(ctx), - } - - if err := es.Publish(ctx, setParentStream, event); err != nil { - return err - } - - return nil -} - -func (es *eventStore) RemoveParentGroup(ctx context.Context, session authn.Session, id string) (err error) { - if err := es.svc.RemoveParentGroup(ctx, session, id); err != nil { - return err - } - - event := removeParentGroupEvent{ - id: id, - Session: session, - requestID: middleware.GetReqID(ctx), - } - - if err := es.Publish(ctx, removeParentStream, event); err != nil { - return err - } - - return nil -} diff --git a/channels/events/streams_test.go b/channels/events/streams_test.go deleted file mode 100644 index 7e9cf3ad6..000000000 --- a/channels/events/streams_test.go +++ /dev/null @@ -1,669 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package events_test - -import ( - "context" - "fmt" - "os" - "testing" - "time" - - "github.com/absmach/magistrala/channels" - "github.com/absmach/magistrala/channels/events" - "github.com/absmach/magistrala/channels/mocks" - "github.com/absmach/magistrala/internal/testsutil" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/connections" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/pkg/roles" - "github.com/go-chi/chi/v5/middleware" - "github.com/redis/go-redis/v9" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" -) - -var ( - storeClient *redis.Client - storeURL string - validSession = authn.Session{ - DomainID: testsutil.GenerateUUID(&testing.T{}), - UserID: testsutil.GenerateUUID(&testing.T{}), - } - validChannel = generateTestChannel(&testing.T{}) - validChannelsPage = channels.ChannelsPage{ - Page: channels.Page{ - Limit: 10, - Offset: 0, - Total: 1, - }, - Channels: []channels.Channel{validChannel}, - } -) - -func newEventStoreMiddleware(t *testing.T) (*mocks.Service, channels.Service) { - svc := new(mocks.Service) - nsvc, err := events.NewEventStoreMiddleware(context.Background(), svc, storeURL) - require.Nil(t, err, fmt.Sprintf("create events store middleware failed with unexpected error: %s", err)) - - return svc, nsvc -} - -func TestMain(m *testing.M) { - code := testsutil.RunRedisTest(m, &storeClient, &storeURL) - os.Exit(code) -} - -func TestCreateChannels(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validID := testsutil.GenerateUUID(t) - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, validID) - - cases := []struct { - desc string - session authn.Session - channels []channels.Channel - svcRes []channels.Channel - svcRoleRes []roles.RoleProvision - svcErr error - resp []channels.Channel - respRoleRes []roles.RoleProvision - err error - }{ - { - desc: "publish successfully", - session: validSession, - channels: []channels.Channel{validChannel}, - svcRes: []channels.Channel{validChannel}, - svcRoleRes: []roles.RoleProvision{}, - svcErr: nil, - resp: []channels.Channel{validChannel}, - respRoleRes: []roles.RoleProvision{}, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - channels: []channels.Channel{validChannel}, - svcRes: []channels.Channel{}, - svcRoleRes: []roles.RoleProvision{}, - svcErr: svcerr.ErrCreateEntity, - resp: []channels.Channel{}, - respRoleRes: []roles.RoleProvision{}, - err: svcerr.ErrCreateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("CreateChannels", validCtx, tc.session, tc.channels).Return(tc.svcRes, tc.svcRoleRes, tc.svcErr) - resp, respRoleRes, err := nsvc.CreateChannels(validCtx, tc.session, tc.channels...) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - assert.Equal(t, tc.respRoleRes, respRoleRes, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.respRoleRes, respRoleRes)) - svcCall.Unset() - }) - } -} - -func TestViewChannel(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - channelID string - withRoles bool - svcRes channels.Channel - svcErr error - resp channels.Channel - err error - }{ - { - desc: "publish successfully", - session: validSession, - channelID: validChannel.ID, - withRoles: false, - svcRes: validChannel, - svcErr: nil, - resp: validChannel, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - channelID: validChannel.ID, - withRoles: false, - svcRes: channels.Channel{}, - svcErr: svcerr.ErrViewEntity, - resp: channels.Channel{}, - err: svcerr.ErrViewEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("ViewChannel", validCtx, tc.session, tc.channelID, tc.withRoles).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.ViewChannel(validCtx, tc.session, tc.channelID, tc.withRoles) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestUpdateChannel(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - updatedChannel := validChannel - updatedChannel.Name = "updatedName" - - cases := []struct { - desc string - session authn.Session - channel channels.Channel - svcRes channels.Channel - svcErr error - resp channels.Channel - err error - }{ - { - desc: "publish successfully", - session: validSession, - channel: updatedChannel, - svcRes: updatedChannel, - svcErr: nil, - resp: updatedChannel, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - channel: updatedChannel, - svcRes: channels.Channel{}, - svcErr: svcerr.ErrUpdateEntity, - resp: channels.Channel{}, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("UpdateChannel", validCtx, tc.session, tc.channel).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.UpdateChannel(validCtx, tc.session, tc.channel) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestUpdateChannelTags(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - updatedChannel := validChannel - updatedChannel.Tags = []string{"newTag1", "newTag2"} - - cases := []struct { - desc string - session authn.Session - channel channels.Channel - svcRes channels.Channel - svcErr error - resp channels.Channel - err error - }{ - { - desc: "publish successfully", - session: validSession, - channel: updatedChannel, - svcRes: updatedChannel, - svcErr: nil, - resp: updatedChannel, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - channel: updatedChannel, - svcRes: channels.Channel{}, - svcErr: svcerr.ErrUpdateEntity, - resp: channels.Channel{}, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("UpdateChannelTags", validCtx, tc.session, tc.channel).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.UpdateChannelTags(validCtx, tc.session, tc.channel) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestEnableChannel(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - channelID string - svcRes channels.Channel - svcErr error - resp channels.Channel - err error - }{ - { - desc: "publish successfully", - session: validSession, - channelID: validChannel.ID, - svcRes: validChannel, - svcErr: nil, - resp: validChannel, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - channelID: validChannel.ID, - svcRes: channels.Channel{}, - svcErr: svcerr.ErrUpdateEntity, - resp: channels.Channel{}, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("EnableChannel", validCtx, tc.session, tc.channelID).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.EnableChannel(validCtx, tc.session, tc.channelID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestDisableChannel(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - channelID string - svcRes channels.Channel - svcErr error - resp channels.Channel - err error - }{ - { - desc: "publish successfully", - session: validSession, - channelID: validChannel.ID, - svcRes: validChannel, - svcErr: nil, - resp: validChannel, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - channelID: validChannel.ID, - svcRes: channels.Channel{}, - svcErr: svcerr.ErrUpdateEntity, - resp: channels.Channel{}, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("DisableChannel", validCtx, tc.session, tc.channelID).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.DisableChannel(validCtx, tc.session, tc.channelID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestListChannels(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - pageMeta channels.Page - svcRes channels.ChannelsPage - svcErr error - resp channels.ChannelsPage - err error - }{ - { - desc: "publish successfully", - session: validSession, - pageMeta: channels.Page{ - Limit: 10, - Offset: 0, - }, - svcRes: validChannelsPage, - svcErr: nil, - resp: validChannelsPage, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - pageMeta: channels.Page{ - Limit: 10, - Offset: 0, - }, - svcRes: channels.ChannelsPage{}, - svcErr: svcerr.ErrViewEntity, - resp: channels.ChannelsPage{}, - err: svcerr.ErrViewEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("ListChannels", validCtx, tc.session, tc.pageMeta).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.ListChannels(validCtx, tc.session, tc.pageMeta) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestListUserChannels(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - userID string - pageMeta channels.Page - svcRes channels.ChannelsPage - svcErr error - resp channels.ChannelsPage - err error - }{ - { - desc: "publish successfully", - session: validSession, - userID: validSession.UserID, - pageMeta: channels.Page{ - Limit: 10, - Offset: 0, - }, - svcRes: validChannelsPage, - svcErr: nil, - resp: validChannelsPage, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - userID: validSession.UserID, - pageMeta: channels.Page{ - Limit: 10, - Offset: 0, - }, - svcRes: channels.ChannelsPage{}, - svcErr: svcerr.ErrViewEntity, - resp: channels.ChannelsPage{}, - err: svcerr.ErrViewEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("ListUserChannels", validCtx, tc.session, tc.userID, tc.pageMeta).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.ListUserChannels(validCtx, tc.session, tc.userID, tc.pageMeta) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestRemoveChannel(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - channelID string - svcErr error - err error - }{ - { - desc: "publish successfully", - session: validSession, - channelID: validChannel.ID, - svcErr: nil, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - channelID: validChannel.ID, - svcErr: svcerr.ErrRemoveEntity, - err: svcerr.ErrRemoveEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("RemoveChannel", validCtx, tc.session, tc.channelID).Return(tc.svcErr) - err := nsvc.RemoveChannel(validCtx, tc.session, tc.channelID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - svcCall.Unset() - }) - } -} - -func TestConnect(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - chIDs []string - clIDs []string - connTypes []connections.ConnType - svcErr error - err error - }{ - { - desc: "publish successfully", - session: validSession, - chIDs: []string{validChannel.ID}, - clIDs: []string{testsutil.GenerateUUID(t)}, - connTypes: []connections.ConnType{connections.Publish}, - svcErr: nil, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - chIDs: []string{validChannel.ID}, - clIDs: []string{testsutil.GenerateUUID(t)}, - connTypes: []connections.ConnType{connections.Publish}, - svcErr: svcerr.ErrCreateEntity, - err: svcerr.ErrCreateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("Connect", validCtx, tc.session, tc.chIDs, tc.clIDs, tc.connTypes).Return(tc.svcErr) - err := nsvc.Connect(validCtx, tc.session, tc.chIDs, tc.clIDs, tc.connTypes) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - svcCall.Unset() - }) - } -} - -func TestDisconnect(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - chIDs []string - clIDs []string - connTypes []connections.ConnType - svcErr error - err error - }{ - { - desc: "publish successfully", - session: validSession, - chIDs: []string{validChannel.ID}, - clIDs: []string{testsutil.GenerateUUID(t)}, - connTypes: []connections.ConnType{connections.Publish}, - svcErr: nil, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - chIDs: []string{validChannel.ID}, - clIDs: []string{testsutil.GenerateUUID(t)}, - connTypes: []connections.ConnType{connections.Publish}, - svcErr: svcerr.ErrRemoveEntity, - err: svcerr.ErrRemoveEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("Disconnect", validCtx, tc.session, tc.chIDs, tc.clIDs, tc.connTypes).Return(tc.svcErr) - err := nsvc.Disconnect(validCtx, tc.session, tc.chIDs, tc.clIDs, tc.connTypes) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - svcCall.Unset() - }) - } -} - -func TestSetParentGroup(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - parentGroupID string - channelID string - svcErr error - err error - }{ - { - desc: "publish successfully", - session: validSession, - parentGroupID: testsutil.GenerateUUID(t), - channelID: validChannel.ID, - svcErr: nil, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - parentGroupID: testsutil.GenerateUUID(t), - channelID: validChannel.ID, - svcErr: svcerr.ErrUpdateEntity, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("SetParentGroup", validCtx, tc.session, tc.parentGroupID, tc.channelID).Return(tc.svcErr) - err := nsvc.SetParentGroup(validCtx, tc.session, tc.parentGroupID, tc.channelID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - svcCall.Unset() - }) - } -} - -func TestRemoveParentGroup(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - channelID string - svcErr error - err error - }{ - { - desc: "publish successfully", - session: validSession, - channelID: validChannel.ID, - svcErr: nil, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - channelID: validChannel.ID, - svcErr: svcerr.ErrUpdateEntity, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("RemoveParentGroup", validCtx, tc.session, tc.channelID).Return(tc.svcErr) - err := nsvc.RemoveParentGroup(validCtx, tc.session, tc.channelID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - svcCall.Unset() - }) - } -} - -func generateTestChannel(t *testing.T) channels.Channel { - createdAt, err := time.Parse(time.RFC3339, "2024-01-01T00:00:00Z") - assert.Nil(t, err, fmt.Sprintf("Unexpected error parsing time: %v", err)) - return channels.Channel{ - ID: testsutil.GenerateUUID(t), - Name: "channelname", - Domain: testsutil.GenerateUUID(t), - Tags: []string{"tag1", "tag2"}, - Metadata: channels.Metadata{"key1": "value1"}, - CreatedAt: createdAt, - UpdatedAt: createdAt, - Status: channels.EnabledStatus, - } -} diff --git a/channels/middleware/authorization.go b/channels/middleware/authorization.go deleted file mode 100644 index f3deea458..000000000 --- a/channels/middleware/authorization.go +++ /dev/null @@ -1,383 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package middleware - -import ( - "context" - "fmt" - - "github.com/absmach/magistrala/auth" - "github.com/absmach/magistrala/channels" - "github.com/absmach/magistrala/channels/operations" - cOperations "github.com/absmach/magistrala/clients/operations" - dOperations "github.com/absmach/magistrala/domains/operations" - gOperations "github.com/absmach/magistrala/groups/operations" - "github.com/absmach/magistrala/pkg/authn" - smqauthz "github.com/absmach/magistrala/pkg/authz" - "github.com/absmach/magistrala/pkg/connections" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/pkg/permissions" - "github.com/absmach/magistrala/pkg/policies" - "github.com/absmach/magistrala/pkg/roles" - rolemgr "github.com/absmach/magistrala/pkg/roles/rolemanager/middleware" -) - -var ( - errView = errors.New("not authorized to view channel") - errList = errors.New("not authorized to list user channels") - errUpdate = errors.New("not authorized to update channel") - errUpdateTags = errors.New("not authorized to update channel tags") - errEnable = errors.New("not authorized to enable channel") - errDisable = errors.New("not authorized to disable channel") - errDelete = errors.New("not authorized to delete channel") - errConnect = errors.New("not authorized to connect to channel") - errDisconnect = errors.New("not authorized to disconnect from channel") - errSetParentGroup = errors.New("not authorized to set parent group to channel") - errRemoveParentGroup = errors.New("not authorized to remove parent group from channel") - errDomainCreateChannels = errors.New("not authorized to create channel in domain") - errGroupSetChildChannels = errors.New("not authorized to set child channel for group") - errGroupRemoveChildChannels = errors.New("not authorized to remove child channel for group") - errClientDisConnectChannels = errors.New("not authorized to disconnect channel for client") - errClientConnectChannels = errors.New("not authorized to connect channel for client") -) - -var _ channels.Service = (*authorizationMiddleware)(nil) - -type authorizationMiddleware struct { - svc channels.Service - repo channels.Repository - authz smqauthz.Authorization - entitiesOps permissions.EntitiesOperations[permissions.Operation] - rolemgr.RoleManagerAuthorizationMiddleware -} - -// NewAuthorization adds authorization to the channels service. -func NewAuthorization( - entityType string, - svc channels.Service, - authz smqauthz.Authorization, - repo channels.Repository, - entitiesOps permissions.EntitiesOperations[permissions.Operation], - roleOps permissions.Operations[permissions.RoleOperation], -) (channels.Service, error) { - if err := entitiesOps.Validate(); err != nil { - return nil, err - } - ram, err := rolemgr.NewAuthorization(policies.ChannelType, svc, authz, roleOps) - if err != nil { - return nil, err - } - - return &authorizationMiddleware{ - svc: svc, - authz: authz, - repo: repo, - entitiesOps: entitiesOps, - RoleManagerAuthorizationMiddleware: ram, - }, nil -} - -func (am *authorizationMiddleware) CreateChannels(ctx context.Context, session authn.Session, chs ...channels.Channel) ([]channels.Channel, []roles.RoleProvision, error) { - if err := am.authorize(ctx, session, policies.DomainType, dOperations.OpCreateDomainChannels, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - Subject: session.DomainUserID, - ObjectType: policies.DomainType, - Object: session.DomainID, - }); err != nil { - return []channels.Channel{}, []roles.RoleProvision{}, errors.Wrap(err, errDomainCreateChannels) - } - - for _, ch := range chs { - if ch.ParentGroup != "" { - if err := am.authorize(ctx, session, policies.GroupType, gOperations.OpGroupSetChildChannel, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - Subject: session.DomainUserID, - ObjectType: policies.GroupType, - Object: ch.ParentGroup, - }); err != nil { - return []channels.Channel{}, []roles.RoleProvision{}, errors.Wrap(err, errors.Wrap(errGroupSetChildChannels, fmt.Errorf("channel name %s parent group id %s", ch.Name, ch.ParentGroup))) - } - } - } - - return am.svc.CreateChannels(ctx, session, chs...) -} - -func (am *authorizationMiddleware) ViewChannel(ctx context.Context, session authn.Session, id string, withRoles bool) (channels.Channel, error) { - if err := am.authorize(ctx, session, policies.ChannelType, operations.OpViewChannel, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - Subject: session.DomainUserID, - ObjectType: policies.ChannelType, - Object: id, - }); err != nil { - return channels.Channel{}, errors.Wrap(err, errView) - } - - return am.svc.ViewChannel(ctx, session, id, withRoles) -} - -func (am *authorizationMiddleware) ListChannels(ctx context.Context, session authn.Session, pm channels.Page) (channels.ChannelsPage, error) { - switch err := am.checkSuperAdmin(ctx, session); { - case err == nil: - session.SuperAdmin = true - case errors.Contains(err, svcerr.ErrSuperAdminAction): - default: - return channels.ChannelsPage{}, err - } - - return am.svc.ListChannels(ctx, session, pm) -} - -func (am *authorizationMiddleware) ListUserChannels(ctx context.Context, session authn.Session, userID string, pm channels.Page) (channels.ChannelsPage, error) { - if err := am.checkSuperAdmin(ctx, session); err != nil { - return channels.ChannelsPage{}, errors.Wrap(err, errList) - } - - return am.svc.ListUserChannels(ctx, session, userID, pm) -} - -func (am *authorizationMiddleware) UpdateChannel(ctx context.Context, session authn.Session, channel channels.Channel) (channels.Channel, error) { - if err := am.authorize(ctx, session, policies.ChannelType, operations.OpUpdateChannel, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - Subject: session.DomainUserID, - ObjectType: policies.ChannelType, - Object: channel.ID, - }); err != nil { - return channels.Channel{}, errors.Wrap(err, errUpdate) - } - - return am.svc.UpdateChannel(ctx, session, channel) -} - -func (am *authorizationMiddleware) UpdateChannelTags(ctx context.Context, session authn.Session, channel channels.Channel) (channels.Channel, error) { - if err := am.authorize(ctx, session, policies.ChannelType, operations.OpUpdateChannelTags, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - Subject: session.DomainUserID, - ObjectType: policies.ChannelType, - Object: channel.ID, - }); err != nil { - return channels.Channel{}, errors.Wrap(err, errUpdateTags) - } - - return am.svc.UpdateChannelTags(ctx, session, channel) -} - -func (am *authorizationMiddleware) EnableChannel(ctx context.Context, session authn.Session, id string) (channels.Channel, error) { - if err := am.authorize(ctx, session, policies.ChannelType, operations.OpEnableChannel, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - Subject: session.DomainUserID, - ObjectType: policies.ChannelType, - Object: id, - }); err != nil { - return channels.Channel{}, errors.Wrap(err, errEnable) - } - - return am.svc.EnableChannel(ctx, session, id) -} - -func (am *authorizationMiddleware) DisableChannel(ctx context.Context, session authn.Session, id string) (channels.Channel, error) { - if err := am.authorize(ctx, session, policies.ChannelType, operations.OpDisableChannel, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - Subject: session.DomainUserID, - ObjectType: policies.ChannelType, - Object: id, - }); err != nil { - return channels.Channel{}, errors.Wrap(err, errDisable) - } - - return am.svc.DisableChannel(ctx, session, id) -} - -func (am *authorizationMiddleware) RemoveChannel(ctx context.Context, session authn.Session, id string) error { - if err := am.authorize(ctx, session, policies.ChannelType, operations.OpDeleteChannel, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - Subject: session.DomainUserID, - ObjectType: policies.ChannelType, - Object: id, - }); err != nil { - return errors.Wrap(err, errDelete) - } - - return am.svc.RemoveChannel(ctx, session, id) -} - -func (am *authorizationMiddleware) Connect(ctx context.Context, session authn.Session, chIDs, thIDs []string, connTypes []connections.ConnType) error { - for _, chID := range chIDs { - if err := am.authorize(ctx, session, policies.ChannelType, operations.OpConnectClient, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - Subject: session.DomainUserID, - ObjectType: policies.ChannelType, - Object: chID, - }); err != nil { - return errors.Wrap(err, errConnect) - } - } - - for _, thID := range thIDs { - if err := am.authorize(ctx, session, policies.ClientType, cOperations.OpConnectToChannel, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - Subject: session.DomainUserID, - ObjectType: policies.ClientType, - Object: thID, - }); err != nil { - return errors.Wrap(err, errClientConnectChannels) - } - } - - return am.svc.Connect(ctx, session, chIDs, thIDs, connTypes) -} - -func (am *authorizationMiddleware) Disconnect(ctx context.Context, session authn.Session, chIDs, thIDs []string, connTypes []connections.ConnType) error { - for _, chID := range chIDs { - if err := am.authorize(ctx, session, policies.ChannelType, operations.OpDisconnectClient, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - Subject: session.DomainUserID, - ObjectType: policies.ChannelType, - Object: chID, - }); err != nil { - return errors.Wrap(err, errDisconnect) - } - } - - for _, thID := range thIDs { - if err := am.authorize(ctx, session, policies.ClientType, cOperations.OpDisconnectFromChannel, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - Subject: session.DomainUserID, - ObjectType: policies.ClientType, - Object: thID, - }); err != nil { - return errors.Wrap(err, errClientDisConnectChannels) - } - } - - return am.svc.Disconnect(ctx, session, chIDs, thIDs, connTypes) -} - -func (am *authorizationMiddleware) SetParentGroup(ctx context.Context, session authn.Session, parentGroupID string, id string) error { - if err := am.authorize(ctx, session, policies.ChannelType, operations.OpSetParentGroup, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - Subject: session.DomainUserID, - ObjectType: policies.ChannelType, - Object: id, - }); err != nil { - return errors.Wrap(err, errSetParentGroup) - } - - if err := am.authorize(ctx, session, policies.GroupType, gOperations.OpGroupSetChildChannel, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - Subject: session.DomainUserID, - ObjectType: policies.GroupType, - Object: parentGroupID, - }); err != nil { - return errors.Wrap(err, errGroupSetChildChannels) - } - - return am.svc.SetParentGroup(ctx, session, parentGroupID, id) -} - -func (am *authorizationMiddleware) RemoveParentGroup(ctx context.Context, session authn.Session, id string) error { - if err := am.authorize(ctx, session, policies.ChannelType, operations.OpSetParentGroup, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - Subject: session.DomainUserID, - ObjectType: policies.ChannelType, - Object: id, - }); err != nil { - return errors.Wrap(err, errRemoveParentGroup) - } - - ch, err := am.repo.RetrieveByID(ctx, id) - if err != nil { - return errors.Wrap(svcerr.ErrRemoveEntity, err) - } - - if ch.ParentGroup != "" { - if err := am.authorize(ctx, session, policies.GroupType, gOperations.OpGroupRemoveChildChannel, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - Subject: session.DomainUserID, - ObjectType: policies.GroupType, - Object: ch.ParentGroup, - }); err != nil { - return errors.Wrap(err, errGroupRemoveChildChannels) - } - - return am.svc.RemoveParentGroup(ctx, session, id) - } - return nil -} - -func (am *authorizationMiddleware) authorize(ctx context.Context, session authn.Session, entityType string, op permissions.Operation, req smqauthz.PolicyReq) error { - req.Domain = session.DomainID - - perm, err := am.entitiesOps.GetPermission(entityType, op) - if err != nil { - return err - } - - req.Permission = perm.String() - - var pat *smqauthz.PATReq - if session.PatID != "" { - entityID := req.Object - opName := am.entitiesOps.OperationName(entityType, op) - if op == operations.OpListUserChannels || op == dOperations.OpCreateDomainChannels || op == dOperations.OpListDomainChannels { - entityID = auth.AnyIDs - } - pat = &smqauthz.PATReq{ - UserID: session.UserID, - PatID: session.PatID, - EntityID: entityID, - EntityType: patEntityType(entityType), - Operation: opName, - Domain: session.DomainID, - } - } - - if err := am.authz.Authorize(ctx, req, pat); err != nil { - return err - } - - return nil -} - -func patEntityType(entityType string) string { - switch entityType { - case policies.ClientType: - return auth.ClientsType.String() - default: - return auth.ChannelsType.String() - } -} - -func (am *authorizationMiddleware) checkSuperAdmin(ctx context.Context, session authn.Session) error { - if session.Role != authn.SuperAdminRole { - return svcerr.ErrSuperAdminAction - } - if err := am.authz.Authorize(ctx, smqauthz.PolicyReq{ - SubjectType: policies.UserType, - Subject: session.UserID, - Permission: policies.AdminPermission, - ObjectType: policies.PlatformType, - Object: policies.MagistralaObject, - }, nil); err != nil { - return err - } - return nil -} diff --git a/channels/middleware/callout.go b/channels/middleware/callout.go deleted file mode 100644 index 94da72a16..000000000 --- a/channels/middleware/callout.go +++ /dev/null @@ -1,247 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package middleware - -import ( - "context" - "time" - - "github.com/absmach/magistrala/channels" - "github.com/absmach/magistrala/channels/operations" - dOperations "github.com/absmach/magistrala/domains/operations" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/callout" - "github.com/absmach/magistrala/pkg/connections" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/pkg/permissions" - "github.com/absmach/magistrala/pkg/policies" - "github.com/absmach/magistrala/pkg/roles" - rolemw "github.com/absmach/magistrala/pkg/roles/rolemanager/middleware" -) - -var _ channels.Service = (*calloutMiddleware)(nil) - -type calloutMiddleware struct { - svc channels.Service - repo channels.Repository - callout callout.Callout - entitiesOps permissions.EntitiesOperations[permissions.Operation] - rolemw.RoleManagerCalloutMiddleware -} - -func NewCallout(svc channels.Service, repo channels.Repository, entitiesOps permissions.EntitiesOperations[permissions.Operation], roleOps permissions.Operations[permissions.RoleOperation], callout callout.Callout) (channels.Service, error) { - call, err := rolemw.NewCallout(policies.ChannelType, svc, callout, roleOps) - if err != nil { - return nil, err - } - - if err := entitiesOps.Validate(); err != nil { - return nil, err - } - - return &calloutMiddleware{ - svc: svc, - repo: repo, - callout: callout, - entitiesOps: entitiesOps, - RoleManagerCalloutMiddleware: call, - }, nil -} - -func (cm *calloutMiddleware) CreateChannels(ctx context.Context, session authn.Session, chs ...channels.Channel) ([]channels.Channel, []roles.RoleProvision, error) { - params := map[string]any{ - "entities": chs, - "count": len(chs), - } - - if err := cm.callOut(ctx, session, policies.DomainType, dOperations.OpCreateDomainChannels, params); err != nil { - return []channels.Channel{}, []roles.RoleProvision{}, err - } - - return cm.svc.CreateChannels(ctx, session, chs...) -} - -func (cm *calloutMiddleware) ViewChannel(ctx context.Context, session authn.Session, id string, withRoles bool) (channels.Channel, error) { - params := map[string]any{ - "entity_id": id, - } - - if err := cm.callOut(ctx, session, policies.ChannelType, operations.OpViewChannel, params); err != nil { - return channels.Channel{}, err - } - - return cm.svc.ViewChannel(ctx, session, id, withRoles) -} - -func (cm *calloutMiddleware) ListChannels(ctx context.Context, session authn.Session, pm channels.Page) (channels.ChannelsPage, error) { - params := map[string]any{ - "pagemeta": pm, - } - - if err := cm.callOut(ctx, session, policies.DomainType, dOperations.OpListDomainChannels, params); err != nil { - return channels.ChannelsPage{}, err - } - - return cm.svc.ListChannels(ctx, session, pm) -} - -func (cm *calloutMiddleware) ListUserChannels(ctx context.Context, session authn.Session, userID string, pm channels.Page) (channels.ChannelsPage, error) { - params := map[string]any{ - "user_id": userID, - "pagemeta": pm, - } - - if err := cm.callOut(ctx, session, policies.ChannelType, operations.OpListUserChannels, params); err != nil { - return channels.ChannelsPage{}, err - } - - return cm.svc.ListUserChannels(ctx, session, userID, pm) -} - -func (cm *calloutMiddleware) UpdateChannel(ctx context.Context, session authn.Session, channel channels.Channel) (channels.Channel, error) { - params := map[string]any{ - "entity_id": channel.ID, - } - - if err := cm.callOut(ctx, session, policies.ChannelType, operations.OpUpdateChannel, params); err != nil { - return channels.Channel{}, err - } - - return cm.svc.UpdateChannel(ctx, session, channel) -} - -func (cm *calloutMiddleware) UpdateChannelTags(ctx context.Context, session authn.Session, channel channels.Channel) (channels.Channel, error) { - params := map[string]any{ - "entity_id": channel.ID, - } - - if err := cm.callOut(ctx, session, policies.ChannelType, operations.OpUpdateChannelTags, params); err != nil { - return channels.Channel{}, err - } - - return cm.svc.UpdateChannelTags(ctx, session, channel) -} - -func (cm *calloutMiddleware) EnableChannel(ctx context.Context, session authn.Session, id string) (channels.Channel, error) { - params := map[string]any{ - "entity_id": id, - } - - if err := cm.callOut(ctx, session, policies.ChannelType, operations.OpEnableChannel, params); err != nil { - return channels.Channel{}, err - } - - return cm.svc.EnableChannel(ctx, session, id) -} - -func (cm *calloutMiddleware) DisableChannel(ctx context.Context, session authn.Session, id string) (channels.Channel, error) { - params := map[string]any{ - "entity_id": id, - } - - if err := cm.callOut(ctx, session, policies.ChannelType, operations.OpDisableChannel, params); err != nil { - return channels.Channel{}, err - } - - return cm.svc.DisableChannel(ctx, session, id) -} - -func (cm *calloutMiddleware) RemoveChannel(ctx context.Context, session authn.Session, id string) error { - params := map[string]any{ - "entity_id": id, - } - - if err := cm.callOut(ctx, session, policies.ChannelType, operations.OpDeleteChannel, params); err != nil { - return err - } - - return cm.svc.RemoveChannel(ctx, session, id) -} - -func (cm *calloutMiddleware) Connect(ctx context.Context, session authn.Session, chIDs, thIDs []string, connTypes []connections.ConnType) error { - params := map[string]any{ - "channel_ids": chIDs, - "client_ids": thIDs, - "connection_types": connTypes, - } - - if err := cm.callOut(ctx, session, policies.ChannelType, operations.OpConnectClient, params); err != nil { - return err - } - - return cm.svc.Connect(ctx, session, chIDs, thIDs, connTypes) -} - -func (cm *calloutMiddleware) Disconnect(ctx context.Context, session authn.Session, chIDs, thIDs []string, connTypes []connections.ConnType) error { - params := map[string]any{ - "channel_ids": chIDs, - "client_ids": thIDs, - "connection_types": connTypes, - } - - if err := cm.callOut(ctx, session, policies.ChannelType, operations.OpDisconnectClient, params); err != nil { - return err - } - - return cm.svc.Disconnect(ctx, session, chIDs, thIDs, connTypes) -} - -func (cm *calloutMiddleware) SetParentGroup(ctx context.Context, session authn.Session, parentGroupID string, id string) error { - params := map[string]any{ - "entity_id": id, - "parent_group_id": parentGroupID, - } - - if err := cm.callOut(ctx, session, policies.ChannelType, operations.OpSetParentGroup, params); err != nil { - return err - } - - return cm.svc.SetParentGroup(ctx, session, parentGroupID, id) -} - -func (cm *calloutMiddleware) RemoveParentGroup(ctx context.Context, session authn.Session, id string) error { - ch, err := cm.repo.RetrieveByID(ctx, id) - if err != nil { - return errors.Wrap(svcerr.ErrRemoveEntity, err) - } - if ch.ParentGroup != "" { - params := map[string]any{ - "entity_id": id, - "parent_group_id": ch.ParentGroup, - } - - if err := cm.callOut(ctx, session, policies.ChannelType, operations.OpRemoveParentGroup, params); err != nil { - return err - } - } - - return cm.svc.RemoveParentGroup(ctx, session, id) -} - -func (cm *calloutMiddleware) callOut(ctx context.Context, session authn.Session, entityType string, op permissions.Operation, pld map[string]any) error { - var entityID string - if id, ok := pld["entity_id"].(string); ok { - entityID = id - } - - req := callout.Request{ - BaseRequest: callout.BaseRequest{ - Operation: cm.entitiesOps.OperationName(entityType, op), - EntityType: entityType, - EntityID: entityID, - CallerID: session.UserID, - CallerType: policies.UserType, - DomainID: session.DomainID, - Time: time.Now().UTC(), - }, - Payload: pld, - } - - if err := cm.callout.Callout(ctx, req); err != nil { - return err - } - - return nil -} diff --git a/channels/middleware/doc.go b/channels/middleware/doc.go deleted file mode 100644 index 25aba3a8e..000000000 --- a/channels/middleware/doc.go +++ /dev/null @@ -1,9 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package middleware provides authorization, logging, metrics and tracing middleware -// for Magistrala Channels Service. -// -// For more details about tracing instrumentation for Magistrala refer to the -// documentation at https://magistrala.absmach.eu/docs/. -package middleware diff --git a/channels/middleware/logging.go b/channels/middleware/logging.go deleted file mode 100644 index 4d0904745..000000000 --- a/channels/middleware/logging.go +++ /dev/null @@ -1,294 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package middleware - -import ( - "context" - "fmt" - "log/slog" - "time" - - "github.com/absmach/magistrala/channels" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/connections" - "github.com/absmach/magistrala/pkg/roles" - rolemw "github.com/absmach/magistrala/pkg/roles/rolemanager/middleware" - "github.com/go-chi/chi/v5/middleware" -) - -var _ channels.Service = (*loggingMiddleware)(nil) - -type loggingMiddleware struct { - logger *slog.Logger - svc channels.Service - rolemw.RoleManagerLoggingMiddleware -} - -// NewLogging adds logging facilities to the channels service. -func NewLogging(svc channels.Service, logger *slog.Logger) channels.Service { - return &loggingMiddleware{logger, svc, rolemw.NewLogging("channels", svc, logger)} -} - -func (lm *loggingMiddleware) CreateChannels(ctx context.Context, session authn.Session, clients ...channels.Channel) (cs []channels.Channel, rps []roles.RoleProvision, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("domain_id", session.DomainID), - slog.String("request_id", middleware.GetReqID(ctx)), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn(fmt.Sprintf("Create %d channels failed", len(clients)), args...) - return - } - lm.logger.Info(fmt.Sprintf("Create %d channel completed successfully", len(clients)), args...) - }(time.Now()) - return lm.svc.CreateChannels(ctx, session, clients...) -} - -func (lm *loggingMiddleware) ViewChannel(ctx context.Context, session authn.Session, id string, withRoles bool) (c channels.Channel, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("domain_id", session.DomainID), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.Group("channel", - slog.String("id", c.ID), - slog.String("name", c.Name), - slog.Bool("with_roles", withRoles), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("View channel failed", args...) - return - } - lm.logger.Info("View channel completed successfully", args...) - }(time.Now()) - return lm.svc.ViewChannel(ctx, session, id, withRoles) -} - -func (lm *loggingMiddleware) ListChannels(ctx context.Context, session authn.Session, pm channels.Page) (cp channels.ChannelsPage, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("domain_id", session.DomainID), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.Group("page", - slog.Uint64("limit", pm.Limit), - slog.Uint64("offset", pm.Offset), - slog.Uint64("total", cp.Total), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("List channels failed", args...) - return - } - lm.logger.Info("List channels completed successfully", args...) - }(time.Now()) - return lm.svc.ListChannels(ctx, session, pm) -} - -func (lm *loggingMiddleware) ListUserChannels(ctx context.Context, session authn.Session, userID string, pm channels.Page) (cp channels.ChannelsPage, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("domain_id", session.DomainID), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.String("user_id", userID), - slog.Group("page", - slog.Uint64("limit", pm.Limit), - slog.Uint64("offset", pm.Offset), - slog.Uint64("total", cp.Total), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("List user channels failed", args...) - return - } - lm.logger.Info("List user channels completed successfully", args...) - }(time.Now()) - return lm.svc.ListUserChannels(ctx, session, userID, pm) -} - -func (lm *loggingMiddleware) UpdateChannel(ctx context.Context, session authn.Session, client channels.Channel) (c channels.Channel, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("domain_id", session.DomainID), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.Group("channel", - slog.String("id", client.ID), - slog.String("name", client.Name), - slog.Any("metadata", client.Metadata), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Update channel failed", args...) - return - } - lm.logger.Info("Update channel completed successfully", args...) - }(time.Now()) - return lm.svc.UpdateChannel(ctx, session, client) -} - -func (lm *loggingMiddleware) UpdateChannelTags(ctx context.Context, session authn.Session, client channels.Channel) (c channels.Channel, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("domain_id", session.DomainID), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.Group("channel", - slog.String("id", c.ID), - slog.String("name", c.Name), - slog.Any("tags", c.Tags), - ), - } - if err != nil { - args := append(args, slog.String("error", err.Error())) - lm.logger.Warn("Update channel tags failed", args...) - return - } - lm.logger.Info("Update channel tags completed successfully", args...) - }(time.Now()) - return lm.svc.UpdateChannelTags(ctx, session, client) -} - -func (lm *loggingMiddleware) EnableChannel(ctx context.Context, session authn.Session, id string) (c channels.Channel, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("domain_id", session.DomainID), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.Group("channel", - slog.String("id", id), - slog.String("name", c.Name), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Enable channel failed", args...) - return - } - lm.logger.Info("Enable channel completed successfully", args...) - }(time.Now()) - return lm.svc.EnableChannel(ctx, session, id) -} - -func (lm *loggingMiddleware) DisableChannel(ctx context.Context, session authn.Session, id string) (c channels.Channel, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("domain_id", session.DomainID), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.Group("channel", - slog.String("id", id), - slog.String("name", c.Name), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Disable channel failed", args...) - return - } - lm.logger.Info("Disable channel completed successfully", args...) - }(time.Now()) - return lm.svc.DisableChannel(ctx, session, id) -} - -func (lm *loggingMiddleware) RemoveChannel(ctx context.Context, session authn.Session, id string) (err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("domain_id", session.DomainID), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.String("channel_id", id), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Delete channel failed", args...) - return - } - lm.logger.Info("Delete channel completed successfully", args...) - }(time.Now()) - return lm.svc.RemoveChannel(ctx, session, id) -} - -func (lm *loggingMiddleware) Connect(ctx context.Context, session authn.Session, chIDs, clIDs []string, connTypes []connections.ConnType) (err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("domain_id", session.DomainID), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.Any("channel_ids", chIDs), - slog.Any("client_ids", clIDs), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Connect channels and clients failed", args...) - return - } - lm.logger.Info("Connect channels and clients completed successfully", args...) - }(time.Now()) - return lm.svc.Connect(ctx, session, chIDs, clIDs, connTypes) -} - -func (lm *loggingMiddleware) Disconnect(ctx context.Context, session authn.Session, chIDs, clIDs []string, connTypes []connections.ConnType) (err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("domain_id", session.DomainID), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.Any("channel_ids", chIDs), - slog.Any("client_ids", clIDs), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Disconnect channels and clients failed", args...) - return - } - lm.logger.Info("Disconnect channels and clients completed successfully", args...) - }(time.Now()) - return lm.svc.Disconnect(ctx, session, chIDs, clIDs, connTypes) -} - -func (lm *loggingMiddleware) SetParentGroup(ctx context.Context, session authn.Session, parentGroupID string, id string) (err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("domain_id", session.DomainID), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.String("parent_group_id", parentGroupID), - slog.String("channel_id", id), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Set parent group to channel failed", args...) - return - } - lm.logger.Info("Set parent group to channel completed successfully", args...) - }(time.Now()) - return lm.svc.SetParentGroup(ctx, session, parentGroupID, id) -} - -func (lm *loggingMiddleware) RemoveParentGroup(ctx context.Context, session authn.Session, id string) (err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("domain_id", session.DomainID), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.String("channel_id", id), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Remove parent group from channel failed", args...) - return - } - lm.logger.Info("Remove parent group from channel completed successfully", args...) - }(time.Now()) - return lm.svc.RemoveParentGroup(ctx, session, id) -} diff --git a/channels/middleware/metrics.go b/channels/middleware/metrics.go deleted file mode 100644 index a9d5ed6fd..000000000 --- a/channels/middleware/metrics.go +++ /dev/null @@ -1,139 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package middleware - -import ( - "context" - "time" - - "github.com/absmach/magistrala/channels" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/connections" - "github.com/absmach/magistrala/pkg/roles" - rolemw "github.com/absmach/magistrala/pkg/roles/rolemanager/middleware" - "github.com/go-kit/kit/metrics" -) - -var _ channels.Service = (*metricsMiddleware)(nil) - -type metricsMiddleware struct { - counter metrics.Counter - latency metrics.Histogram - svc channels.Service - rolemw.RoleManagerMetricsMiddleware -} - -// NewMetrics returns a new metrics middleware wrapper. -func NewMetrics(svc channels.Service, counter metrics.Counter, latency metrics.Histogram) channels.Service { - return &metricsMiddleware{ - counter: counter, - latency: latency, - svc: svc, - RoleManagerMetricsMiddleware: rolemw.NewMetrics("channels", svc, counter, latency), - } -} - -func (ms *metricsMiddleware) CreateChannels(ctx context.Context, session authn.Session, chs ...channels.Channel) ([]channels.Channel, []roles.RoleProvision, error) { - defer func(begin time.Time) { - ms.counter.With("method", "register_channels").Add(1) - ms.latency.With("method", "register_channels").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.CreateChannels(ctx, session, chs...) -} - -func (ms *metricsMiddleware) ViewChannel(ctx context.Context, session authn.Session, id string, withRoles bool) (channels.Channel, error) { - defer func(begin time.Time) { - ms.counter.With("method", "view_channel").Add(1) - ms.latency.With("method", "view_channel").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.ViewChannel(ctx, session, id, withRoles) -} - -func (ms *metricsMiddleware) ListChannels(ctx context.Context, session authn.Session, pm channels.Page) (channels.ChannelsPage, error) { - defer func(begin time.Time) { - ms.counter.With("method", "list_channels").Add(1) - ms.latency.With("method", "list_channels").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.ListChannels(ctx, session, pm) -} - -func (ms *metricsMiddleware) ListUserChannels(ctx context.Context, session authn.Session, userID string, pm channels.Page) (channels.ChannelsPage, error) { - defer func(begin time.Time) { - ms.counter.With("method", "list_user_channels").Add(1) - ms.latency.With("method", "list_user_channels").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.ListUserChannels(ctx, session, userID, pm) -} - -func (ms *metricsMiddleware) UpdateChannel(ctx context.Context, session authn.Session, channel channels.Channel) (channels.Channel, error) { - defer func(begin time.Time) { - ms.counter.With("method", "update_channel").Add(1) - ms.latency.With("method", "update_channel").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.UpdateChannel(ctx, session, channel) -} - -func (ms *metricsMiddleware) UpdateChannelTags(ctx context.Context, session authn.Session, channel channels.Channel) (channels.Channel, error) { - defer func(begin time.Time) { - ms.counter.With("method", "update_channel_tags").Add(1) - ms.latency.With("method", "update_channel_tags").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.UpdateChannelTags(ctx, session, channel) -} - -func (ms *metricsMiddleware) EnableChannel(ctx context.Context, session authn.Session, id string) (channels.Channel, error) { - defer func(begin time.Time) { - ms.counter.With("method", "enable_channel").Add(1) - ms.latency.With("method", "enable_channel").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.EnableChannel(ctx, session, id) -} - -func (ms *metricsMiddleware) DisableChannel(ctx context.Context, session authn.Session, id string) (channels.Channel, error) { - defer func(begin time.Time) { - ms.counter.With("method", "disable_channel").Add(1) - ms.latency.With("method", "disable_channel").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.DisableChannel(ctx, session, id) -} - -func (ms *metricsMiddleware) RemoveChannel(ctx context.Context, session authn.Session, id string) error { - defer func(begin time.Time) { - ms.counter.With("method", "delete_channel").Add(1) - ms.latency.With("method", "delete_channel").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.RemoveChannel(ctx, session, id) -} - -func (ms *metricsMiddleware) Connect(ctx context.Context, session authn.Session, chIDs, thIDs []string, connTypes []connections.ConnType) error { - defer func(begin time.Time) { - ms.counter.With("method", "connect").Add(1) - ms.latency.With("method", "connect").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.Connect(ctx, session, chIDs, thIDs, connTypes) -} - -func (ms *metricsMiddleware) Disconnect(ctx context.Context, session authn.Session, chIDs, thIDs []string, connTypes []connections.ConnType) error { - defer func(begin time.Time) { - ms.counter.With("method", "disconnect").Add(1) - ms.latency.With("method", "disconnect").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.Disconnect(ctx, session, chIDs, thIDs, connTypes) -} - -func (ms *metricsMiddleware) SetParentGroup(ctx context.Context, session authn.Session, parentGroupID string, id string) (err error) { - defer func(begin time.Time) { - ms.counter.With("method", "set_parent_group").Add(1) - ms.latency.With("method", "set_parent_group").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.SetParentGroup(ctx, session, parentGroupID, id) -} - -func (ms *metricsMiddleware) RemoveParentGroup(ctx context.Context, session authn.Session, id string) (err error) { - defer func(begin time.Time) { - ms.counter.With("method", "remove_parent_group").Add(1) - ms.latency.With("method", "remove_parent_group").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.RemoveParentGroup(ctx, session, id) -} diff --git a/channels/middleware/tracing.go b/channels/middleware/tracing.go deleted file mode 100644 index 39fcf57f3..000000000 --- a/channels/middleware/tracing.go +++ /dev/null @@ -1,135 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package middleware - -import ( - "context" - - "github.com/absmach/magistrala/channels" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/connections" - "github.com/absmach/magistrala/pkg/roles" - rolemw "github.com/absmach/magistrala/pkg/roles/rolemanager/middleware" - "github.com/absmach/magistrala/pkg/tracing" - "go.opentelemetry.io/otel/attribute" - "go.opentelemetry.io/otel/trace" -) - -var _ channels.Service = (*tracingMiddleware)(nil) - -type tracingMiddleware struct { - tracer trace.Tracer - svc channels.Service - rolemw.RoleManagerTracing -} - -// NewTracing returns a new channels service with tracing capabilities. -func NewTracing(svc channels.Service, tracer trace.Tracer) channels.Service { - return &tracingMiddleware{tracer, svc, rolemw.NewTracing("channels", svc, tracer)} -} - -// CreateChannels traces the "CreateChannels" operation of the wrapped policies.Service. -func (tm *tracingMiddleware) CreateChannels(ctx context.Context, session authn.Session, chs ...channels.Channel) ([]channels.Channel, []roles.RoleProvision, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_create_channel") - defer span.End() - - return tm.svc.CreateChannels(ctx, session, chs...) -} - -// ViewChannel traces the "ViewChannel" operation of the wrapped policies.Service. -func (tm *tracingMiddleware) ViewChannel(ctx context.Context, session authn.Session, id string, withRoles bool) (channels.Channel, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_view_channel", trace.WithAttributes(attribute.String("id", id))) - defer span.End() - return tm.svc.ViewChannel(ctx, session, id, withRoles) -} - -// ListChannels traces the "ListChannels" operation of the wrapped policies.Service. -func (tm *tracingMiddleware) ListChannels(ctx context.Context, session authn.Session, pm channels.Page) (channels.ChannelsPage, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_list_channels") - defer span.End() - return tm.svc.ListChannels(ctx, session, pm) -} - -func (tm *tracingMiddleware) ListUserChannels(ctx context.Context, session authn.Session, userID string, pm channels.Page) (channels.ChannelsPage, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_list_user_channels") - defer span.End() - return tm.svc.ListUserChannels(ctx, session, userID, pm) -} - -// UpdateChannel traces the "UpdateChannel" operation of the wrapped policies.Service. -func (tm *tracingMiddleware) UpdateChannel(ctx context.Context, session authn.Session, cli channels.Channel) (channels.Channel, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_update_channel", trace.WithAttributes(attribute.String("id", cli.ID))) - defer span.End() - - return tm.svc.UpdateChannel(ctx, session, cli) -} - -// UpdateChannelTags traces the "UpdateChannelTags" operation of the wrapped policies.Service. -func (tm *tracingMiddleware) UpdateChannelTags(ctx context.Context, session authn.Session, cli channels.Channel) (channels.Channel, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_update_channel_tags", trace.WithAttributes( - attribute.String("id", cli.ID), - attribute.StringSlice("tags", cli.Tags), - )) - defer span.End() - - return tm.svc.UpdateChannelTags(ctx, session, cli) -} - -// EnableChannel traces the "EnableChannel" operation of the wrapped policies.Service. -func (tm *tracingMiddleware) EnableChannel(ctx context.Context, session authn.Session, id string) (channels.Channel, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_enable_channel", trace.WithAttributes(attribute.String("id", id))) - defer span.End() - - return tm.svc.EnableChannel(ctx, session, id) -} - -// DisableChannel traces the "DisableChannel" operation of the wrapped policies.Service. -func (tm *tracingMiddleware) DisableChannel(ctx context.Context, session authn.Session, id string) (channels.Channel, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_disable_channel", trace.WithAttributes(attribute.String("id", id))) - defer span.End() - - return tm.svc.DisableChannel(ctx, session, id) -} - -// DeleteChannel traces the "DeleteChannel" operation of the wrapped channels.Service. -func (tm *tracingMiddleware) RemoveChannel(ctx context.Context, session authn.Session, id string) error { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "delete_channel", trace.WithAttributes(attribute.String("id", id))) - defer span.End() - return tm.svc.RemoveChannel(ctx, session, id) -} - -func (tm *tracingMiddleware) Connect(ctx context.Context, session authn.Session, chIDs, thIDs []string, connTypes []connections.ConnType) error { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "connect", trace.WithAttributes( - attribute.StringSlice("channel_ids", chIDs), - attribute.StringSlice("client_ids", thIDs), - )) - defer span.End() - return tm.svc.Connect(ctx, session, chIDs, thIDs, connTypes) -} - -func (tm *tracingMiddleware) Disconnect(ctx context.Context, session authn.Session, chIDs, thIDs []string, connTypes []connections.ConnType) error { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "disconnect", trace.WithAttributes( - attribute.StringSlice("channel_ids", chIDs), - attribute.StringSlice("client_ids", thIDs), - )) - defer span.End() - return tm.svc.Disconnect(ctx, session, chIDs, thIDs, connTypes) -} - -func (tm *tracingMiddleware) SetParentGroup(ctx context.Context, session authn.Session, parentGroupID string, id string) error { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "set_parent_group", trace.WithAttributes( - attribute.String("parent_group_id", parentGroupID), - attribute.String("id", id), - )) - defer span.End() - return tm.svc.SetParentGroup(ctx, session, parentGroupID, id) -} - -func (tm *tracingMiddleware) RemoveParentGroup(ctx context.Context, session authn.Session, id string) error { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "remove_parent_group", trace.WithAttributes( - attribute.String("id", id), - )) - defer span.End() - return tm.svc.RemoveParentGroup(ctx, session, id) -} diff --git a/channels/mocks/cache.go b/channels/mocks/cache.go deleted file mode 100644 index 3f480b456..000000000 --- a/channels/mocks/cache.go +++ /dev/null @@ -1,246 +0,0 @@ -// Copyright (c) Abstract Machines - -// SPDX-License-Identifier: Apache-2.0 - -// Code generated by mockery; DO NOT EDIT. -// github.com/vektra/mockery -// template: testify - -package mocks - -import ( - "context" - - mock "github.com/stretchr/testify/mock" -) - -// NewCache creates a new instance of Cache. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. -// The first argument is typically a *testing.T value. -func NewCache(t interface { - mock.TestingT - Cleanup(func()) -}) *Cache { - mock := &Cache{} - mock.Mock.Test(t) - - t.Cleanup(func() { mock.AssertExpectations(t) }) - - return mock -} - -// Cache is an autogenerated mock type for the Cache type -type Cache struct { - mock.Mock -} - -type Cache_Expecter struct { - mock *mock.Mock -} - -func (_m *Cache) EXPECT() *Cache_Expecter { - return &Cache_Expecter{mock: &_m.Mock} -} - -// ID provides a mock function for the type Cache -func (_mock *Cache) ID(ctx context.Context, channelRoute string, domainID string) (string, error) { - ret := _mock.Called(ctx, channelRoute, domainID) - - if len(ret) == 0 { - panic("no return value specified for ID") - } - - var r0 string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) (string, error)); ok { - return returnFunc(ctx, channelRoute, domainID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) string); ok { - r0 = returnFunc(ctx, channelRoute, domainID) - } else { - r0 = ret.Get(0).(string) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, string) error); ok { - r1 = returnFunc(ctx, channelRoute, domainID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Cache_ID_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ID' -type Cache_ID_Call struct { - *mock.Call -} - -// ID is a helper method to define mock.On call -// - ctx context.Context -// - channelRoute string -// - domainID string -func (_e *Cache_Expecter) ID(ctx interface{}, channelRoute interface{}, domainID interface{}) *Cache_ID_Call { - return &Cache_ID_Call{Call: _e.mock.On("ID", ctx, channelRoute, domainID)} -} - -func (_c *Cache_ID_Call) Run(run func(ctx context.Context, channelRoute string, domainID string)) *Cache_ID_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Cache_ID_Call) Return(s string, err error) *Cache_ID_Call { - _c.Call.Return(s, err) - return _c -} - -func (_c *Cache_ID_Call) RunAndReturn(run func(ctx context.Context, channelRoute string, domainID string) (string, error)) *Cache_ID_Call { - _c.Call.Return(run) - return _c -} - -// Remove provides a mock function for the type Cache -func (_mock *Cache) Remove(ctx context.Context, channelRoute string, domainID string) error { - ret := _mock.Called(ctx, channelRoute, domainID) - - if len(ret) == 0 { - panic("no return value specified for Remove") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) error); ok { - r0 = returnFunc(ctx, channelRoute, domainID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Cache_Remove_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Remove' -type Cache_Remove_Call struct { - *mock.Call -} - -// Remove is a helper method to define mock.On call -// - ctx context.Context -// - channelRoute string -// - domainID string -func (_e *Cache_Expecter) Remove(ctx interface{}, channelRoute interface{}, domainID interface{}) *Cache_Remove_Call { - return &Cache_Remove_Call{Call: _e.mock.On("Remove", ctx, channelRoute, domainID)} -} - -func (_c *Cache_Remove_Call) Run(run func(ctx context.Context, channelRoute string, domainID string)) *Cache_Remove_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Cache_Remove_Call) Return(err error) *Cache_Remove_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Cache_Remove_Call) RunAndReturn(run func(ctx context.Context, channelRoute string, domainID string) error) *Cache_Remove_Call { - _c.Call.Return(run) - return _c -} - -// Save provides a mock function for the type Cache -func (_mock *Cache) Save(ctx context.Context, channelRoute string, domainID string, channelID string) error { - ret := _mock.Called(ctx, channelRoute, domainID, channelID) - - if len(ret) == 0 { - panic("no return value specified for Save") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string, string) error); ok { - r0 = returnFunc(ctx, channelRoute, domainID, channelID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Cache_Save_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Save' -type Cache_Save_Call struct { - *mock.Call -} - -// Save is a helper method to define mock.On call -// - ctx context.Context -// - channelRoute string -// - domainID string -// - channelID string -func (_e *Cache_Expecter) Save(ctx interface{}, channelRoute interface{}, domainID interface{}, channelID interface{}) *Cache_Save_Call { - return &Cache_Save_Call{Call: _e.mock.On("Save", ctx, channelRoute, domainID, channelID)} -} - -func (_c *Cache_Save_Call) Run(run func(ctx context.Context, channelRoute string, domainID string, channelID string)) *Cache_Save_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Cache_Save_Call) Return(err error) *Cache_Save_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Cache_Save_Call) RunAndReturn(run func(ctx context.Context, channelRoute string, domainID string, channelID string) error) *Cache_Save_Call { - _c.Call.Return(run) - return _c -} diff --git a/channels/mocks/channels_client.go b/channels/mocks/channels_client.go deleted file mode 100644 index 644c9fedf..000000000 --- a/channels/mocks/channels_client.go +++ /dev/null @@ -1,460 +0,0 @@ -// Copyright (c) Abstract Machines - -// SPDX-License-Identifier: Apache-2.0 - -// Code generated by mockery; DO NOT EDIT. -// github.com/vektra/mockery -// template: testify - -package mocks - -import ( - "context" - - "github.com/absmach/magistrala/api/grpc/channels/v1" - v10 "github.com/absmach/magistrala/api/grpc/common/v1" - mock "github.com/stretchr/testify/mock" - "google.golang.org/grpc" -) - -// NewChannelsServiceClient creates a new instance of ChannelsServiceClient. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. -// The first argument is typically a *testing.T value. -func NewChannelsServiceClient(t interface { - mock.TestingT - Cleanup(func()) -}) *ChannelsServiceClient { - mock := &ChannelsServiceClient{} - mock.Mock.Test(t) - - t.Cleanup(func() { mock.AssertExpectations(t) }) - - return mock -} - -// ChannelsServiceClient is an autogenerated mock type for the ChannelsServiceClient type -type ChannelsServiceClient struct { - mock.Mock -} - -type ChannelsServiceClient_Expecter struct { - mock *mock.Mock -} - -func (_m *ChannelsServiceClient) EXPECT() *ChannelsServiceClient_Expecter { - return &ChannelsServiceClient_Expecter{mock: &_m.Mock} -} - -// Authorize provides a mock function for the type ChannelsServiceClient -func (_mock *ChannelsServiceClient) Authorize(ctx context.Context, in *v1.AuthzReq, opts ...grpc.CallOption) (*v1.AuthzRes, error) { - var tmpRet mock.Arguments - if len(opts) > 0 { - tmpRet = _mock.Called(ctx, in, opts) - } else { - tmpRet = _mock.Called(ctx, in) - } - ret := tmpRet - - if len(ret) == 0 { - panic("no return value specified for Authorize") - } - - var r0 *v1.AuthzRes - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, *v1.AuthzReq, ...grpc.CallOption) (*v1.AuthzRes, error)); ok { - return returnFunc(ctx, in, opts...) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, *v1.AuthzReq, ...grpc.CallOption) *v1.AuthzRes); ok { - r0 = returnFunc(ctx, in, opts...) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).(*v1.AuthzRes) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, *v1.AuthzReq, ...grpc.CallOption) error); ok { - r1 = returnFunc(ctx, in, opts...) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// ChannelsServiceClient_Authorize_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Authorize' -type ChannelsServiceClient_Authorize_Call struct { - *mock.Call -} - -// Authorize is a helper method to define mock.On call -// - ctx context.Context -// - in *v1.AuthzReq -// - opts ...grpc.CallOption -func (_e *ChannelsServiceClient_Expecter) Authorize(ctx interface{}, in interface{}, opts ...interface{}) *ChannelsServiceClient_Authorize_Call { - return &ChannelsServiceClient_Authorize_Call{Call: _e.mock.On("Authorize", - append([]interface{}{ctx, in}, opts...)...)} -} - -func (_c *ChannelsServiceClient_Authorize_Call) Run(run func(ctx context.Context, in *v1.AuthzReq, opts ...grpc.CallOption)) *ChannelsServiceClient_Authorize_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 *v1.AuthzReq - if args[1] != nil { - arg1 = args[1].(*v1.AuthzReq) - } - var arg2 []grpc.CallOption - var variadicArgs []grpc.CallOption - if len(args) > 2 { - variadicArgs = args[2].([]grpc.CallOption) - } - arg2 = variadicArgs - run( - arg0, - arg1, - arg2..., - ) - }) - return _c -} - -func (_c *ChannelsServiceClient_Authorize_Call) Return(authzRes *v1.AuthzRes, err error) *ChannelsServiceClient_Authorize_Call { - _c.Call.Return(authzRes, err) - return _c -} - -func (_c *ChannelsServiceClient_Authorize_Call) RunAndReturn(run func(ctx context.Context, in *v1.AuthzReq, opts ...grpc.CallOption) (*v1.AuthzRes, error)) *ChannelsServiceClient_Authorize_Call { - _c.Call.Return(run) - return _c -} - -// RemoveClientConnections provides a mock function for the type ChannelsServiceClient -func (_mock *ChannelsServiceClient) RemoveClientConnections(ctx context.Context, in *v1.RemoveClientConnectionsReq, opts ...grpc.CallOption) (*v1.RemoveClientConnectionsRes, error) { - var tmpRet mock.Arguments - if len(opts) > 0 { - tmpRet = _mock.Called(ctx, in, opts) - } else { - tmpRet = _mock.Called(ctx, in) - } - ret := tmpRet - - if len(ret) == 0 { - panic("no return value specified for RemoveClientConnections") - } - - var r0 *v1.RemoveClientConnectionsRes - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, *v1.RemoveClientConnectionsReq, ...grpc.CallOption) (*v1.RemoveClientConnectionsRes, error)); ok { - return returnFunc(ctx, in, opts...) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, *v1.RemoveClientConnectionsReq, ...grpc.CallOption) *v1.RemoveClientConnectionsRes); ok { - r0 = returnFunc(ctx, in, opts...) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).(*v1.RemoveClientConnectionsRes) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, *v1.RemoveClientConnectionsReq, ...grpc.CallOption) error); ok { - r1 = returnFunc(ctx, in, opts...) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// ChannelsServiceClient_RemoveClientConnections_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveClientConnections' -type ChannelsServiceClient_RemoveClientConnections_Call struct { - *mock.Call -} - -// RemoveClientConnections is a helper method to define mock.On call -// - ctx context.Context -// - in *v1.RemoveClientConnectionsReq -// - opts ...grpc.CallOption -func (_e *ChannelsServiceClient_Expecter) RemoveClientConnections(ctx interface{}, in interface{}, opts ...interface{}) *ChannelsServiceClient_RemoveClientConnections_Call { - return &ChannelsServiceClient_RemoveClientConnections_Call{Call: _e.mock.On("RemoveClientConnections", - append([]interface{}{ctx, in}, opts...)...)} -} - -func (_c *ChannelsServiceClient_RemoveClientConnections_Call) Run(run func(ctx context.Context, in *v1.RemoveClientConnectionsReq, opts ...grpc.CallOption)) *ChannelsServiceClient_RemoveClientConnections_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 *v1.RemoveClientConnectionsReq - if args[1] != nil { - arg1 = args[1].(*v1.RemoveClientConnectionsReq) - } - var arg2 []grpc.CallOption - var variadicArgs []grpc.CallOption - if len(args) > 2 { - variadicArgs = args[2].([]grpc.CallOption) - } - arg2 = variadicArgs - run( - arg0, - arg1, - arg2..., - ) - }) - return _c -} - -func (_c *ChannelsServiceClient_RemoveClientConnections_Call) Return(removeClientConnectionsRes *v1.RemoveClientConnectionsRes, err error) *ChannelsServiceClient_RemoveClientConnections_Call { - _c.Call.Return(removeClientConnectionsRes, err) - return _c -} - -func (_c *ChannelsServiceClient_RemoveClientConnections_Call) RunAndReturn(run func(ctx context.Context, in *v1.RemoveClientConnectionsReq, opts ...grpc.CallOption) (*v1.RemoveClientConnectionsRes, error)) *ChannelsServiceClient_RemoveClientConnections_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveEntity provides a mock function for the type ChannelsServiceClient -func (_mock *ChannelsServiceClient) RetrieveEntity(ctx context.Context, in *v10.RetrieveEntityReq, opts ...grpc.CallOption) (*v10.RetrieveEntityRes, error) { - var tmpRet mock.Arguments - if len(opts) > 0 { - tmpRet = _mock.Called(ctx, in, opts) - } else { - tmpRet = _mock.Called(ctx, in) - } - ret := tmpRet - - if len(ret) == 0 { - panic("no return value specified for RetrieveEntity") - } - - var r0 *v10.RetrieveEntityRes - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, *v10.RetrieveEntityReq, ...grpc.CallOption) (*v10.RetrieveEntityRes, error)); ok { - return returnFunc(ctx, in, opts...) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, *v10.RetrieveEntityReq, ...grpc.CallOption) *v10.RetrieveEntityRes); ok { - r0 = returnFunc(ctx, in, opts...) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).(*v10.RetrieveEntityRes) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, *v10.RetrieveEntityReq, ...grpc.CallOption) error); ok { - r1 = returnFunc(ctx, in, opts...) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// ChannelsServiceClient_RetrieveEntity_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveEntity' -type ChannelsServiceClient_RetrieveEntity_Call struct { - *mock.Call -} - -// RetrieveEntity is a helper method to define mock.On call -// - ctx context.Context -// - in *v10.RetrieveEntityReq -// - opts ...grpc.CallOption -func (_e *ChannelsServiceClient_Expecter) RetrieveEntity(ctx interface{}, in interface{}, opts ...interface{}) *ChannelsServiceClient_RetrieveEntity_Call { - return &ChannelsServiceClient_RetrieveEntity_Call{Call: _e.mock.On("RetrieveEntity", - append([]interface{}{ctx, in}, opts...)...)} -} - -func (_c *ChannelsServiceClient_RetrieveEntity_Call) Run(run func(ctx context.Context, in *v10.RetrieveEntityReq, opts ...grpc.CallOption)) *ChannelsServiceClient_RetrieveEntity_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 *v10.RetrieveEntityReq - if args[1] != nil { - arg1 = args[1].(*v10.RetrieveEntityReq) - } - var arg2 []grpc.CallOption - var variadicArgs []grpc.CallOption - if len(args) > 2 { - variadicArgs = args[2].([]grpc.CallOption) - } - arg2 = variadicArgs - run( - arg0, - arg1, - arg2..., - ) - }) - return _c -} - -func (_c *ChannelsServiceClient_RetrieveEntity_Call) Return(retrieveEntityRes *v10.RetrieveEntityRes, err error) *ChannelsServiceClient_RetrieveEntity_Call { - _c.Call.Return(retrieveEntityRes, err) - return _c -} - -func (_c *ChannelsServiceClient_RetrieveEntity_Call) RunAndReturn(run func(ctx context.Context, in *v10.RetrieveEntityReq, opts ...grpc.CallOption) (*v10.RetrieveEntityRes, error)) *ChannelsServiceClient_RetrieveEntity_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveIDByRoute provides a mock function for the type ChannelsServiceClient -func (_mock *ChannelsServiceClient) RetrieveIDByRoute(ctx context.Context, in *v10.RetrieveIDByRouteReq, opts ...grpc.CallOption) (*v10.RetrieveEntityRes, error) { - var tmpRet mock.Arguments - if len(opts) > 0 { - tmpRet = _mock.Called(ctx, in, opts) - } else { - tmpRet = _mock.Called(ctx, in) - } - ret := tmpRet - - if len(ret) == 0 { - panic("no return value specified for RetrieveIDByRoute") - } - - var r0 *v10.RetrieveEntityRes - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, *v10.RetrieveIDByRouteReq, ...grpc.CallOption) (*v10.RetrieveEntityRes, error)); ok { - return returnFunc(ctx, in, opts...) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, *v10.RetrieveIDByRouteReq, ...grpc.CallOption) *v10.RetrieveEntityRes); ok { - r0 = returnFunc(ctx, in, opts...) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).(*v10.RetrieveEntityRes) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, *v10.RetrieveIDByRouteReq, ...grpc.CallOption) error); ok { - r1 = returnFunc(ctx, in, opts...) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// ChannelsServiceClient_RetrieveIDByRoute_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveIDByRoute' -type ChannelsServiceClient_RetrieveIDByRoute_Call struct { - *mock.Call -} - -// RetrieveIDByRoute is a helper method to define mock.On call -// - ctx context.Context -// - in *v10.RetrieveIDByRouteReq -// - opts ...grpc.CallOption -func (_e *ChannelsServiceClient_Expecter) RetrieveIDByRoute(ctx interface{}, in interface{}, opts ...interface{}) *ChannelsServiceClient_RetrieveIDByRoute_Call { - return &ChannelsServiceClient_RetrieveIDByRoute_Call{Call: _e.mock.On("RetrieveIDByRoute", - append([]interface{}{ctx, in}, opts...)...)} -} - -func (_c *ChannelsServiceClient_RetrieveIDByRoute_Call) Run(run func(ctx context.Context, in *v10.RetrieveIDByRouteReq, opts ...grpc.CallOption)) *ChannelsServiceClient_RetrieveIDByRoute_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 *v10.RetrieveIDByRouteReq - if args[1] != nil { - arg1 = args[1].(*v10.RetrieveIDByRouteReq) - } - var arg2 []grpc.CallOption - var variadicArgs []grpc.CallOption - if len(args) > 2 { - variadicArgs = args[2].([]grpc.CallOption) - } - arg2 = variadicArgs - run( - arg0, - arg1, - arg2..., - ) - }) - return _c -} - -func (_c *ChannelsServiceClient_RetrieveIDByRoute_Call) Return(retrieveEntityRes *v10.RetrieveEntityRes, err error) *ChannelsServiceClient_RetrieveIDByRoute_Call { - _c.Call.Return(retrieveEntityRes, err) - return _c -} - -func (_c *ChannelsServiceClient_RetrieveIDByRoute_Call) RunAndReturn(run func(ctx context.Context, in *v10.RetrieveIDByRouteReq, opts ...grpc.CallOption) (*v10.RetrieveEntityRes, error)) *ChannelsServiceClient_RetrieveIDByRoute_Call { - _c.Call.Return(run) - return _c -} - -// UnsetParentGroupFromChannels provides a mock function for the type ChannelsServiceClient -func (_mock *ChannelsServiceClient) UnsetParentGroupFromChannels(ctx context.Context, in *v1.UnsetParentGroupFromChannelsReq, opts ...grpc.CallOption) (*v1.UnsetParentGroupFromChannelsRes, error) { - var tmpRet mock.Arguments - if len(opts) > 0 { - tmpRet = _mock.Called(ctx, in, opts) - } else { - tmpRet = _mock.Called(ctx, in) - } - ret := tmpRet - - if len(ret) == 0 { - panic("no return value specified for UnsetParentGroupFromChannels") - } - - var r0 *v1.UnsetParentGroupFromChannelsRes - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, *v1.UnsetParentGroupFromChannelsReq, ...grpc.CallOption) (*v1.UnsetParentGroupFromChannelsRes, error)); ok { - return returnFunc(ctx, in, opts...) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, *v1.UnsetParentGroupFromChannelsReq, ...grpc.CallOption) *v1.UnsetParentGroupFromChannelsRes); ok { - r0 = returnFunc(ctx, in, opts...) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).(*v1.UnsetParentGroupFromChannelsRes) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, *v1.UnsetParentGroupFromChannelsReq, ...grpc.CallOption) error); ok { - r1 = returnFunc(ctx, in, opts...) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// ChannelsServiceClient_UnsetParentGroupFromChannels_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UnsetParentGroupFromChannels' -type ChannelsServiceClient_UnsetParentGroupFromChannels_Call struct { - *mock.Call -} - -// UnsetParentGroupFromChannels is a helper method to define mock.On call -// - ctx context.Context -// - in *v1.UnsetParentGroupFromChannelsReq -// - opts ...grpc.CallOption -func (_e *ChannelsServiceClient_Expecter) UnsetParentGroupFromChannels(ctx interface{}, in interface{}, opts ...interface{}) *ChannelsServiceClient_UnsetParentGroupFromChannels_Call { - return &ChannelsServiceClient_UnsetParentGroupFromChannels_Call{Call: _e.mock.On("UnsetParentGroupFromChannels", - append([]interface{}{ctx, in}, opts...)...)} -} - -func (_c *ChannelsServiceClient_UnsetParentGroupFromChannels_Call) Run(run func(ctx context.Context, in *v1.UnsetParentGroupFromChannelsReq, opts ...grpc.CallOption)) *ChannelsServiceClient_UnsetParentGroupFromChannels_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 *v1.UnsetParentGroupFromChannelsReq - if args[1] != nil { - arg1 = args[1].(*v1.UnsetParentGroupFromChannelsReq) - } - var arg2 []grpc.CallOption - var variadicArgs []grpc.CallOption - if len(args) > 2 { - variadicArgs = args[2].([]grpc.CallOption) - } - arg2 = variadicArgs - run( - arg0, - arg1, - arg2..., - ) - }) - return _c -} - -func (_c *ChannelsServiceClient_UnsetParentGroupFromChannels_Call) Return(unsetParentGroupFromChannelsRes *v1.UnsetParentGroupFromChannelsRes, err error) *ChannelsServiceClient_UnsetParentGroupFromChannels_Call { - _c.Call.Return(unsetParentGroupFromChannelsRes, err) - return _c -} - -func (_c *ChannelsServiceClient_UnsetParentGroupFromChannels_Call) RunAndReturn(run func(ctx context.Context, in *v1.UnsetParentGroupFromChannelsReq, opts ...grpc.CallOption) (*v1.UnsetParentGroupFromChannelsRes, error)) *ChannelsServiceClient_UnsetParentGroupFromChannels_Call { - _c.Call.Return(run) - return _c -} diff --git a/channels/mocks/repository.go b/channels/mocks/repository.go deleted file mode 100644 index bb0f18f7f..000000000 --- a/channels/mocks/repository.go +++ /dev/null @@ -1,2805 +0,0 @@ -// Copyright (c) Abstract Machines - -// SPDX-License-Identifier: Apache-2.0 - -// Code generated by mockery; DO NOT EDIT. -// github.com/vektra/mockery -// template: testify - -package mocks - -import ( - "context" - - "github.com/absmach/magistrala/channels" - "github.com/absmach/magistrala/pkg/roles" - mock "github.com/stretchr/testify/mock" -) - -// NewRepository creates a new instance of Repository. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. -// The first argument is typically a *testing.T value. -func NewRepository(t interface { - mock.TestingT - Cleanup(func()) -}) *Repository { - mock := &Repository{} - mock.Mock.Test(t) - - t.Cleanup(func() { mock.AssertExpectations(t) }) - - return mock -} - -// Repository is an autogenerated mock type for the Repository type -type Repository struct { - mock.Mock -} - -type Repository_Expecter struct { - mock *mock.Mock -} - -func (_m *Repository) EXPECT() *Repository_Expecter { - return &Repository_Expecter{mock: &_m.Mock} -} - -// AddConnections provides a mock function for the type Repository -func (_mock *Repository) AddConnections(ctx context.Context, conns []channels.Connection) error { - ret := _mock.Called(ctx, conns) - - if len(ret) == 0 { - panic("no return value specified for AddConnections") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, []channels.Connection) error); ok { - r0 = returnFunc(ctx, conns) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_AddConnections_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'AddConnections' -type Repository_AddConnections_Call struct { - *mock.Call -} - -// AddConnections is a helper method to define mock.On call -// - ctx context.Context -// - conns []channels.Connection -func (_e *Repository_Expecter) AddConnections(ctx interface{}, conns interface{}) *Repository_AddConnections_Call { - return &Repository_AddConnections_Call{Call: _e.mock.On("AddConnections", ctx, conns)} -} - -func (_c *Repository_AddConnections_Call) Run(run func(ctx context.Context, conns []channels.Connection)) *Repository_AddConnections_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 []channels.Connection - if args[1] != nil { - arg1 = args[1].([]channels.Connection) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_AddConnections_Call) Return(err error) *Repository_AddConnections_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_AddConnections_Call) RunAndReturn(run func(ctx context.Context, conns []channels.Connection) error) *Repository_AddConnections_Call { - _c.Call.Return(run) - return _c -} - -// AddRoles provides a mock function for the type Repository -func (_mock *Repository) AddRoles(ctx context.Context, rps []roles.RoleProvision) ([]roles.RoleProvision, error) { - ret := _mock.Called(ctx, rps) - - if len(ret) == 0 { - panic("no return value specified for AddRoles") - } - - var r0 []roles.RoleProvision - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, []roles.RoleProvision) ([]roles.RoleProvision, error)); ok { - return returnFunc(ctx, rps) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, []roles.RoleProvision) []roles.RoleProvision); ok { - r0 = returnFunc(ctx, rps) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]roles.RoleProvision) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, []roles.RoleProvision) error); ok { - r1 = returnFunc(ctx, rps) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_AddRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'AddRoles' -type Repository_AddRoles_Call struct { - *mock.Call -} - -// AddRoles is a helper method to define mock.On call -// - ctx context.Context -// - rps []roles.RoleProvision -func (_e *Repository_Expecter) AddRoles(ctx interface{}, rps interface{}) *Repository_AddRoles_Call { - return &Repository_AddRoles_Call{Call: _e.mock.On("AddRoles", ctx, rps)} -} - -func (_c *Repository_AddRoles_Call) Run(run func(ctx context.Context, rps []roles.RoleProvision)) *Repository_AddRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 []roles.RoleProvision - if args[1] != nil { - arg1 = args[1].([]roles.RoleProvision) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_AddRoles_Call) Return(roleProvisions []roles.RoleProvision, err error) *Repository_AddRoles_Call { - _c.Call.Return(roleProvisions, err) - return _c -} - -func (_c *Repository_AddRoles_Call) RunAndReturn(run func(ctx context.Context, rps []roles.RoleProvision) ([]roles.RoleProvision, error)) *Repository_AddRoles_Call { - _c.Call.Return(run) - return _c -} - -// ChangeStatus provides a mock function for the type Repository -func (_mock *Repository) ChangeStatus(ctx context.Context, channel channels.Channel) (channels.Channel, error) { - ret := _mock.Called(ctx, channel) - - if len(ret) == 0 { - panic("no return value specified for ChangeStatus") - } - - var r0 channels.Channel - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, channels.Channel) (channels.Channel, error)); ok { - return returnFunc(ctx, channel) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, channels.Channel) channels.Channel); ok { - r0 = returnFunc(ctx, channel) - } else { - r0 = ret.Get(0).(channels.Channel) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, channels.Channel) error); ok { - r1 = returnFunc(ctx, channel) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_ChangeStatus_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ChangeStatus' -type Repository_ChangeStatus_Call struct { - *mock.Call -} - -// ChangeStatus is a helper method to define mock.On call -// - ctx context.Context -// - channel channels.Channel -func (_e *Repository_Expecter) ChangeStatus(ctx interface{}, channel interface{}) *Repository_ChangeStatus_Call { - return &Repository_ChangeStatus_Call{Call: _e.mock.On("ChangeStatus", ctx, channel)} -} - -func (_c *Repository_ChangeStatus_Call) Run(run func(ctx context.Context, channel channels.Channel)) *Repository_ChangeStatus_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 channels.Channel - if args[1] != nil { - arg1 = args[1].(channels.Channel) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_ChangeStatus_Call) Return(channel1 channels.Channel, err error) *Repository_ChangeStatus_Call { - _c.Call.Return(channel1, err) - return _c -} - -func (_c *Repository_ChangeStatus_Call) RunAndReturn(run func(ctx context.Context, channel channels.Channel) (channels.Channel, error)) *Repository_ChangeStatus_Call { - _c.Call.Return(run) - return _c -} - -// ChannelConnectionsCount provides a mock function for the type Repository -func (_mock *Repository) ChannelConnectionsCount(ctx context.Context, id string) (uint64, error) { - ret := _mock.Called(ctx, id) - - if len(ret) == 0 { - panic("no return value specified for ChannelConnectionsCount") - } - - var r0 uint64 - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) (uint64, error)); ok { - return returnFunc(ctx, id) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string) uint64); ok { - r0 = returnFunc(ctx, id) - } else { - r0 = ret.Get(0).(uint64) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = returnFunc(ctx, id) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_ChannelConnectionsCount_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ChannelConnectionsCount' -type Repository_ChannelConnectionsCount_Call struct { - *mock.Call -} - -// ChannelConnectionsCount is a helper method to define mock.On call -// - ctx context.Context -// - id string -func (_e *Repository_Expecter) ChannelConnectionsCount(ctx interface{}, id interface{}) *Repository_ChannelConnectionsCount_Call { - return &Repository_ChannelConnectionsCount_Call{Call: _e.mock.On("ChannelConnectionsCount", ctx, id)} -} - -func (_c *Repository_ChannelConnectionsCount_Call) Run(run func(ctx context.Context, id string)) *Repository_ChannelConnectionsCount_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_ChannelConnectionsCount_Call) Return(v uint64, err error) *Repository_ChannelConnectionsCount_Call { - _c.Call.Return(v, err) - return _c -} - -func (_c *Repository_ChannelConnectionsCount_Call) RunAndReturn(run func(ctx context.Context, id string) (uint64, error)) *Repository_ChannelConnectionsCount_Call { - _c.Call.Return(run) - return _c -} - -// CheckConnection provides a mock function for the type Repository -func (_mock *Repository) CheckConnection(ctx context.Context, conn channels.Connection) error { - ret := _mock.Called(ctx, conn) - - if len(ret) == 0 { - panic("no return value specified for CheckConnection") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, channels.Connection) error); ok { - r0 = returnFunc(ctx, conn) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_CheckConnection_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'CheckConnection' -type Repository_CheckConnection_Call struct { - *mock.Call -} - -// CheckConnection is a helper method to define mock.On call -// - ctx context.Context -// - conn channels.Connection -func (_e *Repository_Expecter) CheckConnection(ctx interface{}, conn interface{}) *Repository_CheckConnection_Call { - return &Repository_CheckConnection_Call{Call: _e.mock.On("CheckConnection", ctx, conn)} -} - -func (_c *Repository_CheckConnection_Call) Run(run func(ctx context.Context, conn channels.Connection)) *Repository_CheckConnection_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 channels.Connection - if args[1] != nil { - arg1 = args[1].(channels.Connection) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_CheckConnection_Call) Return(err error) *Repository_CheckConnection_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_CheckConnection_Call) RunAndReturn(run func(ctx context.Context, conn channels.Connection) error) *Repository_CheckConnection_Call { - _c.Call.Return(run) - return _c -} - -// ClientAuthorize provides a mock function for the type Repository -func (_mock *Repository) ClientAuthorize(ctx context.Context, conn channels.Connection) error { - ret := _mock.Called(ctx, conn) - - if len(ret) == 0 { - panic("no return value specified for ClientAuthorize") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, channels.Connection) error); ok { - r0 = returnFunc(ctx, conn) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_ClientAuthorize_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ClientAuthorize' -type Repository_ClientAuthorize_Call struct { - *mock.Call -} - -// ClientAuthorize is a helper method to define mock.On call -// - ctx context.Context -// - conn channels.Connection -func (_e *Repository_Expecter) ClientAuthorize(ctx interface{}, conn interface{}) *Repository_ClientAuthorize_Call { - return &Repository_ClientAuthorize_Call{Call: _e.mock.On("ClientAuthorize", ctx, conn)} -} - -func (_c *Repository_ClientAuthorize_Call) Run(run func(ctx context.Context, conn channels.Connection)) *Repository_ClientAuthorize_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 channels.Connection - if args[1] != nil { - arg1 = args[1].(channels.Connection) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_ClientAuthorize_Call) Return(err error) *Repository_ClientAuthorize_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_ClientAuthorize_Call) RunAndReturn(run func(ctx context.Context, conn channels.Connection) error) *Repository_ClientAuthorize_Call { - _c.Call.Return(run) - return _c -} - -// DoesChannelHaveConnections provides a mock function for the type Repository -func (_mock *Repository) DoesChannelHaveConnections(ctx context.Context, id string) (bool, error) { - ret := _mock.Called(ctx, id) - - if len(ret) == 0 { - panic("no return value specified for DoesChannelHaveConnections") - } - - var r0 bool - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) (bool, error)); ok { - return returnFunc(ctx, id) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string) bool); ok { - r0 = returnFunc(ctx, id) - } else { - r0 = ret.Get(0).(bool) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = returnFunc(ctx, id) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_DoesChannelHaveConnections_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'DoesChannelHaveConnections' -type Repository_DoesChannelHaveConnections_Call struct { - *mock.Call -} - -// DoesChannelHaveConnections is a helper method to define mock.On call -// - ctx context.Context -// - id string -func (_e *Repository_Expecter) DoesChannelHaveConnections(ctx interface{}, id interface{}) *Repository_DoesChannelHaveConnections_Call { - return &Repository_DoesChannelHaveConnections_Call{Call: _e.mock.On("DoesChannelHaveConnections", ctx, id)} -} - -func (_c *Repository_DoesChannelHaveConnections_Call) Run(run func(ctx context.Context, id string)) *Repository_DoesChannelHaveConnections_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_DoesChannelHaveConnections_Call) Return(b bool, err error) *Repository_DoesChannelHaveConnections_Call { - _c.Call.Return(b, err) - return _c -} - -func (_c *Repository_DoesChannelHaveConnections_Call) RunAndReturn(run func(ctx context.Context, id string) (bool, error)) *Repository_DoesChannelHaveConnections_Call { - _c.Call.Return(run) - return _c -} - -// ListEntityMembers provides a mock function for the type Repository -func (_mock *Repository) ListEntityMembers(ctx context.Context, entityID string, pageQuery roles.MembersRolePageQuery) (roles.MembersRolePage, error) { - ret := _mock.Called(ctx, entityID, pageQuery) - - if len(ret) == 0 { - panic("no return value specified for ListEntityMembers") - } - - var r0 roles.MembersRolePage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, roles.MembersRolePageQuery) (roles.MembersRolePage, error)); ok { - return returnFunc(ctx, entityID, pageQuery) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, roles.MembersRolePageQuery) roles.MembersRolePage); ok { - r0 = returnFunc(ctx, entityID, pageQuery) - } else { - r0 = ret.Get(0).(roles.MembersRolePage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, roles.MembersRolePageQuery) error); ok { - r1 = returnFunc(ctx, entityID, pageQuery) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_ListEntityMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ListEntityMembers' -type Repository_ListEntityMembers_Call struct { - *mock.Call -} - -// ListEntityMembers is a helper method to define mock.On call -// - ctx context.Context -// - entityID string -// - pageQuery roles.MembersRolePageQuery -func (_e *Repository_Expecter) ListEntityMembers(ctx interface{}, entityID interface{}, pageQuery interface{}) *Repository_ListEntityMembers_Call { - return &Repository_ListEntityMembers_Call{Call: _e.mock.On("ListEntityMembers", ctx, entityID, pageQuery)} -} - -func (_c *Repository_ListEntityMembers_Call) Run(run func(ctx context.Context, entityID string, pageQuery roles.MembersRolePageQuery)) *Repository_ListEntityMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 roles.MembersRolePageQuery - if args[2] != nil { - arg2 = args[2].(roles.MembersRolePageQuery) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_ListEntityMembers_Call) Return(membersRolePage roles.MembersRolePage, err error) *Repository_ListEntityMembers_Call { - _c.Call.Return(membersRolePage, err) - return _c -} - -func (_c *Repository_ListEntityMembers_Call) RunAndReturn(run func(ctx context.Context, entityID string, pageQuery roles.MembersRolePageQuery) (roles.MembersRolePage, error)) *Repository_ListEntityMembers_Call { - _c.Call.Return(run) - return _c -} - -// Remove provides a mock function for the type Repository -func (_mock *Repository) Remove(ctx context.Context, ids ...string) error { - var tmpRet mock.Arguments - if len(ids) > 0 { - tmpRet = _mock.Called(ctx, ids) - } else { - tmpRet = _mock.Called(ctx) - } - ret := tmpRet - - if len(ret) == 0 { - panic("no return value specified for Remove") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, ...string) error); ok { - r0 = returnFunc(ctx, ids...) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_Remove_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Remove' -type Repository_Remove_Call struct { - *mock.Call -} - -// Remove is a helper method to define mock.On call -// - ctx context.Context -// - ids ...string -func (_e *Repository_Expecter) Remove(ctx interface{}, ids ...interface{}) *Repository_Remove_Call { - return &Repository_Remove_Call{Call: _e.mock.On("Remove", - append([]interface{}{ctx}, ids...)...)} -} - -func (_c *Repository_Remove_Call) Run(run func(ctx context.Context, ids ...string)) *Repository_Remove_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 []string - var variadicArgs []string - if len(args) > 1 { - variadicArgs = args[1].([]string) - } - arg1 = variadicArgs - run( - arg0, - arg1..., - ) - }) - return _c -} - -func (_c *Repository_Remove_Call) Return(err error) *Repository_Remove_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_Remove_Call) RunAndReturn(run func(ctx context.Context, ids ...string) error) *Repository_Remove_Call { - _c.Call.Return(run) - return _c -} - -// RemoveChannelConnections provides a mock function for the type Repository -func (_mock *Repository) RemoveChannelConnections(ctx context.Context, channelID string) error { - ret := _mock.Called(ctx, channelID) - - if len(ret) == 0 { - panic("no return value specified for RemoveChannelConnections") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) error); ok { - r0 = returnFunc(ctx, channelID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RemoveChannelConnections_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveChannelConnections' -type Repository_RemoveChannelConnections_Call struct { - *mock.Call -} - -// RemoveChannelConnections is a helper method to define mock.On call -// - ctx context.Context -// - channelID string -func (_e *Repository_Expecter) RemoveChannelConnections(ctx interface{}, channelID interface{}) *Repository_RemoveChannelConnections_Call { - return &Repository_RemoveChannelConnections_Call{Call: _e.mock.On("RemoveChannelConnections", ctx, channelID)} -} - -func (_c *Repository_RemoveChannelConnections_Call) Run(run func(ctx context.Context, channelID string)) *Repository_RemoveChannelConnections_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RemoveChannelConnections_Call) Return(err error) *Repository_RemoveChannelConnections_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RemoveChannelConnections_Call) RunAndReturn(run func(ctx context.Context, channelID string) error) *Repository_RemoveChannelConnections_Call { - _c.Call.Return(run) - return _c -} - -// RemoveClientConnections provides a mock function for the type Repository -func (_mock *Repository) RemoveClientConnections(ctx context.Context, clientID string) error { - ret := _mock.Called(ctx, clientID) - - if len(ret) == 0 { - panic("no return value specified for RemoveClientConnections") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) error); ok { - r0 = returnFunc(ctx, clientID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RemoveClientConnections_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveClientConnections' -type Repository_RemoveClientConnections_Call struct { - *mock.Call -} - -// RemoveClientConnections is a helper method to define mock.On call -// - ctx context.Context -// - clientID string -func (_e *Repository_Expecter) RemoveClientConnections(ctx interface{}, clientID interface{}) *Repository_RemoveClientConnections_Call { - return &Repository_RemoveClientConnections_Call{Call: _e.mock.On("RemoveClientConnections", ctx, clientID)} -} - -func (_c *Repository_RemoveClientConnections_Call) Run(run func(ctx context.Context, clientID string)) *Repository_RemoveClientConnections_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RemoveClientConnections_Call) Return(err error) *Repository_RemoveClientConnections_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RemoveClientConnections_Call) RunAndReturn(run func(ctx context.Context, clientID string) error) *Repository_RemoveClientConnections_Call { - _c.Call.Return(run) - return _c -} - -// RemoveConnections provides a mock function for the type Repository -func (_mock *Repository) RemoveConnections(ctx context.Context, conns []channels.Connection) error { - ret := _mock.Called(ctx, conns) - - if len(ret) == 0 { - panic("no return value specified for RemoveConnections") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, []channels.Connection) error); ok { - r0 = returnFunc(ctx, conns) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RemoveConnections_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveConnections' -type Repository_RemoveConnections_Call struct { - *mock.Call -} - -// RemoveConnections is a helper method to define mock.On call -// - ctx context.Context -// - conns []channels.Connection -func (_e *Repository_Expecter) RemoveConnections(ctx interface{}, conns interface{}) *Repository_RemoveConnections_Call { - return &Repository_RemoveConnections_Call{Call: _e.mock.On("RemoveConnections", ctx, conns)} -} - -func (_c *Repository_RemoveConnections_Call) Run(run func(ctx context.Context, conns []channels.Connection)) *Repository_RemoveConnections_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 []channels.Connection - if args[1] != nil { - arg1 = args[1].([]channels.Connection) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RemoveConnections_Call) Return(err error) *Repository_RemoveConnections_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RemoveConnections_Call) RunAndReturn(run func(ctx context.Context, conns []channels.Connection) error) *Repository_RemoveConnections_Call { - _c.Call.Return(run) - return _c -} - -// RemoveEntityMembers provides a mock function for the type Repository -func (_mock *Repository) RemoveEntityMembers(ctx context.Context, entityID string, members []string) error { - ret := _mock.Called(ctx, entityID, members) - - if len(ret) == 0 { - panic("no return value specified for RemoveEntityMembers") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, []string) error); ok { - r0 = returnFunc(ctx, entityID, members) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RemoveEntityMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveEntityMembers' -type Repository_RemoveEntityMembers_Call struct { - *mock.Call -} - -// RemoveEntityMembers is a helper method to define mock.On call -// - ctx context.Context -// - entityID string -// - members []string -func (_e *Repository_Expecter) RemoveEntityMembers(ctx interface{}, entityID interface{}, members interface{}) *Repository_RemoveEntityMembers_Call { - return &Repository_RemoveEntityMembers_Call{Call: _e.mock.On("RemoveEntityMembers", ctx, entityID, members)} -} - -func (_c *Repository_RemoveEntityMembers_Call) Run(run func(ctx context.Context, entityID string, members []string)) *Repository_RemoveEntityMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RemoveEntityMembers_Call) Return(err error) *Repository_RemoveEntityMembers_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RemoveEntityMembers_Call) RunAndReturn(run func(ctx context.Context, entityID string, members []string) error) *Repository_RemoveEntityMembers_Call { - _c.Call.Return(run) - return _c -} - -// RemoveMemberFromAllRoles provides a mock function for the type Repository -func (_mock *Repository) RemoveMemberFromAllRoles(ctx context.Context, memberID string) error { - ret := _mock.Called(ctx, memberID) - - if len(ret) == 0 { - panic("no return value specified for RemoveMemberFromAllRoles") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) error); ok { - r0 = returnFunc(ctx, memberID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RemoveMemberFromAllRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveMemberFromAllRoles' -type Repository_RemoveMemberFromAllRoles_Call struct { - *mock.Call -} - -// RemoveMemberFromAllRoles is a helper method to define mock.On call -// - ctx context.Context -// - memberID string -func (_e *Repository_Expecter) RemoveMemberFromAllRoles(ctx interface{}, memberID interface{}) *Repository_RemoveMemberFromAllRoles_Call { - return &Repository_RemoveMemberFromAllRoles_Call{Call: _e.mock.On("RemoveMemberFromAllRoles", ctx, memberID)} -} - -func (_c *Repository_RemoveMemberFromAllRoles_Call) Run(run func(ctx context.Context, memberID string)) *Repository_RemoveMemberFromAllRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RemoveMemberFromAllRoles_Call) Return(err error) *Repository_RemoveMemberFromAllRoles_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RemoveMemberFromAllRoles_Call) RunAndReturn(run func(ctx context.Context, memberID string) error) *Repository_RemoveMemberFromAllRoles_Call { - _c.Call.Return(run) - return _c -} - -// RemoveParentGroup provides a mock function for the type Repository -func (_mock *Repository) RemoveParentGroup(ctx context.Context, ch channels.Channel) error { - ret := _mock.Called(ctx, ch) - - if len(ret) == 0 { - panic("no return value specified for RemoveParentGroup") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, channels.Channel) error); ok { - r0 = returnFunc(ctx, ch) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RemoveParentGroup_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveParentGroup' -type Repository_RemoveParentGroup_Call struct { - *mock.Call -} - -// RemoveParentGroup is a helper method to define mock.On call -// - ctx context.Context -// - ch channels.Channel -func (_e *Repository_Expecter) RemoveParentGroup(ctx interface{}, ch interface{}) *Repository_RemoveParentGroup_Call { - return &Repository_RemoveParentGroup_Call{Call: _e.mock.On("RemoveParentGroup", ctx, ch)} -} - -func (_c *Repository_RemoveParentGroup_Call) Run(run func(ctx context.Context, ch channels.Channel)) *Repository_RemoveParentGroup_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 channels.Channel - if args[1] != nil { - arg1 = args[1].(channels.Channel) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RemoveParentGroup_Call) Return(err error) *Repository_RemoveParentGroup_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RemoveParentGroup_Call) RunAndReturn(run func(ctx context.Context, ch channels.Channel) error) *Repository_RemoveParentGroup_Call { - _c.Call.Return(run) - return _c -} - -// RemoveRoles provides a mock function for the type Repository -func (_mock *Repository) RemoveRoles(ctx context.Context, roleIDs []string) error { - ret := _mock.Called(ctx, roleIDs) - - if len(ret) == 0 { - panic("no return value specified for RemoveRoles") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, []string) error); ok { - r0 = returnFunc(ctx, roleIDs) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RemoveRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveRoles' -type Repository_RemoveRoles_Call struct { - *mock.Call -} - -// RemoveRoles is a helper method to define mock.On call -// - ctx context.Context -// - roleIDs []string -func (_e *Repository_Expecter) RemoveRoles(ctx interface{}, roleIDs interface{}) *Repository_RemoveRoles_Call { - return &Repository_RemoveRoles_Call{Call: _e.mock.On("RemoveRoles", ctx, roleIDs)} -} - -func (_c *Repository_RemoveRoles_Call) Run(run func(ctx context.Context, roleIDs []string)) *Repository_RemoveRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 []string - if args[1] != nil { - arg1 = args[1].([]string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RemoveRoles_Call) Return(err error) *Repository_RemoveRoles_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RemoveRoles_Call) RunAndReturn(run func(ctx context.Context, roleIDs []string) error) *Repository_RemoveRoles_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveAll provides a mock function for the type Repository -func (_mock *Repository) RetrieveAll(ctx context.Context, pm channels.Page) (channels.ChannelsPage, error) { - ret := _mock.Called(ctx, pm) - - if len(ret) == 0 { - panic("no return value specified for RetrieveAll") - } - - var r0 channels.ChannelsPage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, channels.Page) (channels.ChannelsPage, error)); ok { - return returnFunc(ctx, pm) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, channels.Page) channels.ChannelsPage); ok { - r0 = returnFunc(ctx, pm) - } else { - r0 = ret.Get(0).(channels.ChannelsPage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, channels.Page) error); ok { - r1 = returnFunc(ctx, pm) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveAll_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveAll' -type Repository_RetrieveAll_Call struct { - *mock.Call -} - -// RetrieveAll is a helper method to define mock.On call -// - ctx context.Context -// - pm channels.Page -func (_e *Repository_Expecter) RetrieveAll(ctx interface{}, pm interface{}) *Repository_RetrieveAll_Call { - return &Repository_RetrieveAll_Call{Call: _e.mock.On("RetrieveAll", ctx, pm)} -} - -func (_c *Repository_RetrieveAll_Call) Run(run func(ctx context.Context, pm channels.Page)) *Repository_RetrieveAll_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 channels.Page - if args[1] != nil { - arg1 = args[1].(channels.Page) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RetrieveAll_Call) Return(channelsPage channels.ChannelsPage, err error) *Repository_RetrieveAll_Call { - _c.Call.Return(channelsPage, err) - return _c -} - -func (_c *Repository_RetrieveAll_Call) RunAndReturn(run func(ctx context.Context, pm channels.Page) (channels.ChannelsPage, error)) *Repository_RetrieveAll_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveAllRoles provides a mock function for the type Repository -func (_mock *Repository) RetrieveAllRoles(ctx context.Context, entityID string, limit uint64, offset uint64) (roles.RolePage, error) { - ret := _mock.Called(ctx, entityID, limit, offset) - - if len(ret) == 0 { - panic("no return value specified for RetrieveAllRoles") - } - - var r0 roles.RolePage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, uint64, uint64) (roles.RolePage, error)); ok { - return returnFunc(ctx, entityID, limit, offset) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, uint64, uint64) roles.RolePage); ok { - r0 = returnFunc(ctx, entityID, limit, offset) - } else { - r0 = ret.Get(0).(roles.RolePage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, uint64, uint64) error); ok { - r1 = returnFunc(ctx, entityID, limit, offset) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveAllRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveAllRoles' -type Repository_RetrieveAllRoles_Call struct { - *mock.Call -} - -// RetrieveAllRoles is a helper method to define mock.On call -// - ctx context.Context -// - entityID string -// - limit uint64 -// - offset uint64 -func (_e *Repository_Expecter) RetrieveAllRoles(ctx interface{}, entityID interface{}, limit interface{}, offset interface{}) *Repository_RetrieveAllRoles_Call { - return &Repository_RetrieveAllRoles_Call{Call: _e.mock.On("RetrieveAllRoles", ctx, entityID, limit, offset)} -} - -func (_c *Repository_RetrieveAllRoles_Call) Run(run func(ctx context.Context, entityID string, limit uint64, offset uint64)) *Repository_RetrieveAllRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 uint64 - if args[2] != nil { - arg2 = args[2].(uint64) - } - var arg3 uint64 - if args[3] != nil { - arg3 = args[3].(uint64) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Repository_RetrieveAllRoles_Call) Return(rolePage roles.RolePage, err error) *Repository_RetrieveAllRoles_Call { - _c.Call.Return(rolePage, err) - return _c -} - -func (_c *Repository_RetrieveAllRoles_Call) RunAndReturn(run func(ctx context.Context, entityID string, limit uint64, offset uint64) (roles.RolePage, error)) *Repository_RetrieveAllRoles_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveByID provides a mock function for the type Repository -func (_mock *Repository) RetrieveByID(ctx context.Context, id string) (channels.Channel, error) { - ret := _mock.Called(ctx, id) - - if len(ret) == 0 { - panic("no return value specified for RetrieveByID") - } - - var r0 channels.Channel - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) (channels.Channel, error)); ok { - return returnFunc(ctx, id) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string) channels.Channel); ok { - r0 = returnFunc(ctx, id) - } else { - r0 = ret.Get(0).(channels.Channel) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = returnFunc(ctx, id) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveByID_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveByID' -type Repository_RetrieveByID_Call struct { - *mock.Call -} - -// RetrieveByID is a helper method to define mock.On call -// - ctx context.Context -// - id string -func (_e *Repository_Expecter) RetrieveByID(ctx interface{}, id interface{}) *Repository_RetrieveByID_Call { - return &Repository_RetrieveByID_Call{Call: _e.mock.On("RetrieveByID", ctx, id)} -} - -func (_c *Repository_RetrieveByID_Call) Run(run func(ctx context.Context, id string)) *Repository_RetrieveByID_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RetrieveByID_Call) Return(channel channels.Channel, err error) *Repository_RetrieveByID_Call { - _c.Call.Return(channel, err) - return _c -} - -func (_c *Repository_RetrieveByID_Call) RunAndReturn(run func(ctx context.Context, id string) (channels.Channel, error)) *Repository_RetrieveByID_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveByIDWithRoles provides a mock function for the type Repository -func (_mock *Repository) RetrieveByIDWithRoles(ctx context.Context, id string, memberID string) (channels.Channel, error) { - ret := _mock.Called(ctx, id, memberID) - - if len(ret) == 0 { - panic("no return value specified for RetrieveByIDWithRoles") - } - - var r0 channels.Channel - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) (channels.Channel, error)); ok { - return returnFunc(ctx, id, memberID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) channels.Channel); ok { - r0 = returnFunc(ctx, id, memberID) - } else { - r0 = ret.Get(0).(channels.Channel) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, string) error); ok { - r1 = returnFunc(ctx, id, memberID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveByIDWithRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveByIDWithRoles' -type Repository_RetrieveByIDWithRoles_Call struct { - *mock.Call -} - -// RetrieveByIDWithRoles is a helper method to define mock.On call -// - ctx context.Context -// - id string -// - memberID string -func (_e *Repository_Expecter) RetrieveByIDWithRoles(ctx interface{}, id interface{}, memberID interface{}) *Repository_RetrieveByIDWithRoles_Call { - return &Repository_RetrieveByIDWithRoles_Call{Call: _e.mock.On("RetrieveByIDWithRoles", ctx, id, memberID)} -} - -func (_c *Repository_RetrieveByIDWithRoles_Call) Run(run func(ctx context.Context, id string, memberID string)) *Repository_RetrieveByIDWithRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RetrieveByIDWithRoles_Call) Return(channel channels.Channel, err error) *Repository_RetrieveByIDWithRoles_Call { - _c.Call.Return(channel, err) - return _c -} - -func (_c *Repository_RetrieveByIDWithRoles_Call) RunAndReturn(run func(ctx context.Context, id string, memberID string) (channels.Channel, error)) *Repository_RetrieveByIDWithRoles_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveByRoute provides a mock function for the type Repository -func (_mock *Repository) RetrieveByRoute(ctx context.Context, route string, domainID string) (channels.Channel, error) { - ret := _mock.Called(ctx, route, domainID) - - if len(ret) == 0 { - panic("no return value specified for RetrieveByRoute") - } - - var r0 channels.Channel - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) (channels.Channel, error)); ok { - return returnFunc(ctx, route, domainID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) channels.Channel); ok { - r0 = returnFunc(ctx, route, domainID) - } else { - r0 = ret.Get(0).(channels.Channel) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, string) error); ok { - r1 = returnFunc(ctx, route, domainID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveByRoute_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveByRoute' -type Repository_RetrieveByRoute_Call struct { - *mock.Call -} - -// RetrieveByRoute is a helper method to define mock.On call -// - ctx context.Context -// - route string -// - domainID string -func (_e *Repository_Expecter) RetrieveByRoute(ctx interface{}, route interface{}, domainID interface{}) *Repository_RetrieveByRoute_Call { - return &Repository_RetrieveByRoute_Call{Call: _e.mock.On("RetrieveByRoute", ctx, route, domainID)} -} - -func (_c *Repository_RetrieveByRoute_Call) Run(run func(ctx context.Context, route string, domainID string)) *Repository_RetrieveByRoute_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RetrieveByRoute_Call) Return(channel channels.Channel, err error) *Repository_RetrieveByRoute_Call { - _c.Call.Return(channel, err) - return _c -} - -func (_c *Repository_RetrieveByRoute_Call) RunAndReturn(run func(ctx context.Context, route string, domainID string) (channels.Channel, error)) *Repository_RetrieveByRoute_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveEntitiesRolesActionsMembers provides a mock function for the type Repository -func (_mock *Repository) RetrieveEntitiesRolesActionsMembers(ctx context.Context, entityIDs []string) ([]roles.EntityActionRole, []roles.EntityMemberRole, error) { - ret := _mock.Called(ctx, entityIDs) - - if len(ret) == 0 { - panic("no return value specified for RetrieveEntitiesRolesActionsMembers") - } - - var r0 []roles.EntityActionRole - var r1 []roles.EntityMemberRole - var r2 error - if returnFunc, ok := ret.Get(0).(func(context.Context, []string) ([]roles.EntityActionRole, []roles.EntityMemberRole, error)); ok { - return returnFunc(ctx, entityIDs) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, []string) []roles.EntityActionRole); ok { - r0 = returnFunc(ctx, entityIDs) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]roles.EntityActionRole) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, []string) []roles.EntityMemberRole); ok { - r1 = returnFunc(ctx, entityIDs) - } else { - if ret.Get(1) != nil { - r1 = ret.Get(1).([]roles.EntityMemberRole) - } - } - if returnFunc, ok := ret.Get(2).(func(context.Context, []string) error); ok { - r2 = returnFunc(ctx, entityIDs) - } else { - r2 = ret.Error(2) - } - return r0, r1, r2 -} - -// Repository_RetrieveEntitiesRolesActionsMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveEntitiesRolesActionsMembers' -type Repository_RetrieveEntitiesRolesActionsMembers_Call struct { - *mock.Call -} - -// RetrieveEntitiesRolesActionsMembers is a helper method to define mock.On call -// - ctx context.Context -// - entityIDs []string -func (_e *Repository_Expecter) RetrieveEntitiesRolesActionsMembers(ctx interface{}, entityIDs interface{}) *Repository_RetrieveEntitiesRolesActionsMembers_Call { - return &Repository_RetrieveEntitiesRolesActionsMembers_Call{Call: _e.mock.On("RetrieveEntitiesRolesActionsMembers", ctx, entityIDs)} -} - -func (_c *Repository_RetrieveEntitiesRolesActionsMembers_Call) Run(run func(ctx context.Context, entityIDs []string)) *Repository_RetrieveEntitiesRolesActionsMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 []string - if args[1] != nil { - arg1 = args[1].([]string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RetrieveEntitiesRolesActionsMembers_Call) Return(entityActionRoles []roles.EntityActionRole, entityMemberRoles []roles.EntityMemberRole, err error) *Repository_RetrieveEntitiesRolesActionsMembers_Call { - _c.Call.Return(entityActionRoles, entityMemberRoles, err) - return _c -} - -func (_c *Repository_RetrieveEntitiesRolesActionsMembers_Call) RunAndReturn(run func(ctx context.Context, entityIDs []string) ([]roles.EntityActionRole, []roles.EntityMemberRole, error)) *Repository_RetrieveEntitiesRolesActionsMembers_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveEntityRole provides a mock function for the type Repository -func (_mock *Repository) RetrieveEntityRole(ctx context.Context, entityID string, roleID string) (roles.Role, error) { - ret := _mock.Called(ctx, entityID, roleID) - - if len(ret) == 0 { - panic("no return value specified for RetrieveEntityRole") - } - - var r0 roles.Role - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) (roles.Role, error)); ok { - return returnFunc(ctx, entityID, roleID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) roles.Role); ok { - r0 = returnFunc(ctx, entityID, roleID) - } else { - r0 = ret.Get(0).(roles.Role) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, string) error); ok { - r1 = returnFunc(ctx, entityID, roleID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveEntityRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveEntityRole' -type Repository_RetrieveEntityRole_Call struct { - *mock.Call -} - -// RetrieveEntityRole is a helper method to define mock.On call -// - ctx context.Context -// - entityID string -// - roleID string -func (_e *Repository_Expecter) RetrieveEntityRole(ctx interface{}, entityID interface{}, roleID interface{}) *Repository_RetrieveEntityRole_Call { - return &Repository_RetrieveEntityRole_Call{Call: _e.mock.On("RetrieveEntityRole", ctx, entityID, roleID)} -} - -func (_c *Repository_RetrieveEntityRole_Call) Run(run func(ctx context.Context, entityID string, roleID string)) *Repository_RetrieveEntityRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RetrieveEntityRole_Call) Return(role roles.Role, err error) *Repository_RetrieveEntityRole_Call { - _c.Call.Return(role, err) - return _c -} - -func (_c *Repository_RetrieveEntityRole_Call) RunAndReturn(run func(ctx context.Context, entityID string, roleID string) (roles.Role, error)) *Repository_RetrieveEntityRole_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveParentGroupChannels provides a mock function for the type Repository -func (_mock *Repository) RetrieveParentGroupChannels(ctx context.Context, parentGroupID string) ([]channels.Channel, error) { - ret := _mock.Called(ctx, parentGroupID) - - if len(ret) == 0 { - panic("no return value specified for RetrieveParentGroupChannels") - } - - var r0 []channels.Channel - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) ([]channels.Channel, error)); ok { - return returnFunc(ctx, parentGroupID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string) []channels.Channel); ok { - r0 = returnFunc(ctx, parentGroupID) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]channels.Channel) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = returnFunc(ctx, parentGroupID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveParentGroupChannels_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveParentGroupChannels' -type Repository_RetrieveParentGroupChannels_Call struct { - *mock.Call -} - -// RetrieveParentGroupChannels is a helper method to define mock.On call -// - ctx context.Context -// - parentGroupID string -func (_e *Repository_Expecter) RetrieveParentGroupChannels(ctx interface{}, parentGroupID interface{}) *Repository_RetrieveParentGroupChannels_Call { - return &Repository_RetrieveParentGroupChannels_Call{Call: _e.mock.On("RetrieveParentGroupChannels", ctx, parentGroupID)} -} - -func (_c *Repository_RetrieveParentGroupChannels_Call) Run(run func(ctx context.Context, parentGroupID string)) *Repository_RetrieveParentGroupChannels_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RetrieveParentGroupChannels_Call) Return(channels1 []channels.Channel, err error) *Repository_RetrieveParentGroupChannels_Call { - _c.Call.Return(channels1, err) - return _c -} - -func (_c *Repository_RetrieveParentGroupChannels_Call) RunAndReturn(run func(ctx context.Context, parentGroupID string) ([]channels.Channel, error)) *Repository_RetrieveParentGroupChannels_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveRole provides a mock function for the type Repository -func (_mock *Repository) RetrieveRole(ctx context.Context, roleID string) (roles.Role, error) { - ret := _mock.Called(ctx, roleID) - - if len(ret) == 0 { - panic("no return value specified for RetrieveRole") - } - - var r0 roles.Role - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) (roles.Role, error)); ok { - return returnFunc(ctx, roleID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string) roles.Role); ok { - r0 = returnFunc(ctx, roleID) - } else { - r0 = ret.Get(0).(roles.Role) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = returnFunc(ctx, roleID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveRole' -type Repository_RetrieveRole_Call struct { - *mock.Call -} - -// RetrieveRole is a helper method to define mock.On call -// - ctx context.Context -// - roleID string -func (_e *Repository_Expecter) RetrieveRole(ctx interface{}, roleID interface{}) *Repository_RetrieveRole_Call { - return &Repository_RetrieveRole_Call{Call: _e.mock.On("RetrieveRole", ctx, roleID)} -} - -func (_c *Repository_RetrieveRole_Call) Run(run func(ctx context.Context, roleID string)) *Repository_RetrieveRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RetrieveRole_Call) Return(role roles.Role, err error) *Repository_RetrieveRole_Call { - _c.Call.Return(role, err) - return _c -} - -func (_c *Repository_RetrieveRole_Call) RunAndReturn(run func(ctx context.Context, roleID string) (roles.Role, error)) *Repository_RetrieveRole_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveUserChannels provides a mock function for the type Repository -func (_mock *Repository) RetrieveUserChannels(ctx context.Context, domainID string, userID string, pm channels.Page) (channels.ChannelsPage, error) { - ret := _mock.Called(ctx, domainID, userID, pm) - - if len(ret) == 0 { - panic("no return value specified for RetrieveUserChannels") - } - - var r0 channels.ChannelsPage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string, channels.Page) (channels.ChannelsPage, error)); ok { - return returnFunc(ctx, domainID, userID, pm) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string, channels.Page) channels.ChannelsPage); ok { - r0 = returnFunc(ctx, domainID, userID, pm) - } else { - r0 = ret.Get(0).(channels.ChannelsPage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, string, channels.Page) error); ok { - r1 = returnFunc(ctx, domainID, userID, pm) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveUserChannels_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveUserChannels' -type Repository_RetrieveUserChannels_Call struct { - *mock.Call -} - -// RetrieveUserChannels is a helper method to define mock.On call -// - ctx context.Context -// - domainID string -// - userID string -// - pm channels.Page -func (_e *Repository_Expecter) RetrieveUserChannels(ctx interface{}, domainID interface{}, userID interface{}, pm interface{}) *Repository_RetrieveUserChannels_Call { - return &Repository_RetrieveUserChannels_Call{Call: _e.mock.On("RetrieveUserChannels", ctx, domainID, userID, pm)} -} - -func (_c *Repository_RetrieveUserChannels_Call) Run(run func(ctx context.Context, domainID string, userID string, pm channels.Page)) *Repository_RetrieveUserChannels_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 channels.Page - if args[3] != nil { - arg3 = args[3].(channels.Page) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Repository_RetrieveUserChannels_Call) Return(channelsPage channels.ChannelsPage, err error) *Repository_RetrieveUserChannels_Call { - _c.Call.Return(channelsPage, err) - return _c -} - -func (_c *Repository_RetrieveUserChannels_Call) RunAndReturn(run func(ctx context.Context, domainID string, userID string, pm channels.Page) (channels.ChannelsPage, error)) *Repository_RetrieveUserChannels_Call { - _c.Call.Return(run) - return _c -} - -// RoleAddActions provides a mock function for the type Repository -func (_mock *Repository) RoleAddActions(ctx context.Context, role roles.Role, actions []string) ([]string, error) { - ret := _mock.Called(ctx, role, actions) - - if len(ret) == 0 { - panic("no return value specified for RoleAddActions") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role, []string) ([]string, error)); ok { - return returnFunc(ctx, role, actions) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role, []string) []string); ok { - r0 = returnFunc(ctx, role, actions) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, roles.Role, []string) error); ok { - r1 = returnFunc(ctx, role, actions) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RoleAddActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleAddActions' -type Repository_RoleAddActions_Call struct { - *mock.Call -} - -// RoleAddActions is a helper method to define mock.On call -// - ctx context.Context -// - role roles.Role -// - actions []string -func (_e *Repository_Expecter) RoleAddActions(ctx interface{}, role interface{}, actions interface{}) *Repository_RoleAddActions_Call { - return &Repository_RoleAddActions_Call{Call: _e.mock.On("RoleAddActions", ctx, role, actions)} -} - -func (_c *Repository_RoleAddActions_Call) Run(run func(ctx context.Context, role roles.Role, actions []string)) *Repository_RoleAddActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RoleAddActions_Call) Return(ops []string, err error) *Repository_RoleAddActions_Call { - _c.Call.Return(ops, err) - return _c -} - -func (_c *Repository_RoleAddActions_Call) RunAndReturn(run func(ctx context.Context, role roles.Role, actions []string) ([]string, error)) *Repository_RoleAddActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleAddMembers provides a mock function for the type Repository -func (_mock *Repository) RoleAddMembers(ctx context.Context, role roles.Role, members []string) ([]string, error) { - ret := _mock.Called(ctx, role, members) - - if len(ret) == 0 { - panic("no return value specified for RoleAddMembers") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role, []string) ([]string, error)); ok { - return returnFunc(ctx, role, members) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role, []string) []string); ok { - r0 = returnFunc(ctx, role, members) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, roles.Role, []string) error); ok { - r1 = returnFunc(ctx, role, members) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RoleAddMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleAddMembers' -type Repository_RoleAddMembers_Call struct { - *mock.Call -} - -// RoleAddMembers is a helper method to define mock.On call -// - ctx context.Context -// - role roles.Role -// - members []string -func (_e *Repository_Expecter) RoleAddMembers(ctx interface{}, role interface{}, members interface{}) *Repository_RoleAddMembers_Call { - return &Repository_RoleAddMembers_Call{Call: _e.mock.On("RoleAddMembers", ctx, role, members)} -} - -func (_c *Repository_RoleAddMembers_Call) Run(run func(ctx context.Context, role roles.Role, members []string)) *Repository_RoleAddMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RoleAddMembers_Call) Return(strings []string, err error) *Repository_RoleAddMembers_Call { - _c.Call.Return(strings, err) - return _c -} - -func (_c *Repository_RoleAddMembers_Call) RunAndReturn(run func(ctx context.Context, role roles.Role, members []string) ([]string, error)) *Repository_RoleAddMembers_Call { - _c.Call.Return(run) - return _c -} - -// RoleCheckActionsExists provides a mock function for the type Repository -func (_mock *Repository) RoleCheckActionsExists(ctx context.Context, roleID string, actions []string) (bool, error) { - ret := _mock.Called(ctx, roleID, actions) - - if len(ret) == 0 { - panic("no return value specified for RoleCheckActionsExists") - } - - var r0 bool - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, []string) (bool, error)); ok { - return returnFunc(ctx, roleID, actions) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, []string) bool); ok { - r0 = returnFunc(ctx, roleID, actions) - } else { - r0 = ret.Get(0).(bool) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, []string) error); ok { - r1 = returnFunc(ctx, roleID, actions) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RoleCheckActionsExists_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleCheckActionsExists' -type Repository_RoleCheckActionsExists_Call struct { - *mock.Call -} - -// RoleCheckActionsExists is a helper method to define mock.On call -// - ctx context.Context -// - roleID string -// - actions []string -func (_e *Repository_Expecter) RoleCheckActionsExists(ctx interface{}, roleID interface{}, actions interface{}) *Repository_RoleCheckActionsExists_Call { - return &Repository_RoleCheckActionsExists_Call{Call: _e.mock.On("RoleCheckActionsExists", ctx, roleID, actions)} -} - -func (_c *Repository_RoleCheckActionsExists_Call) Run(run func(ctx context.Context, roleID string, actions []string)) *Repository_RoleCheckActionsExists_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RoleCheckActionsExists_Call) Return(b bool, err error) *Repository_RoleCheckActionsExists_Call { - _c.Call.Return(b, err) - return _c -} - -func (_c *Repository_RoleCheckActionsExists_Call) RunAndReturn(run func(ctx context.Context, roleID string, actions []string) (bool, error)) *Repository_RoleCheckActionsExists_Call { - _c.Call.Return(run) - return _c -} - -// RoleCheckMembersExists provides a mock function for the type Repository -func (_mock *Repository) RoleCheckMembersExists(ctx context.Context, roleID string, members []string) (bool, error) { - ret := _mock.Called(ctx, roleID, members) - - if len(ret) == 0 { - panic("no return value specified for RoleCheckMembersExists") - } - - var r0 bool - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, []string) (bool, error)); ok { - return returnFunc(ctx, roleID, members) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, []string) bool); ok { - r0 = returnFunc(ctx, roleID, members) - } else { - r0 = ret.Get(0).(bool) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, []string) error); ok { - r1 = returnFunc(ctx, roleID, members) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RoleCheckMembersExists_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleCheckMembersExists' -type Repository_RoleCheckMembersExists_Call struct { - *mock.Call -} - -// RoleCheckMembersExists is a helper method to define mock.On call -// - ctx context.Context -// - roleID string -// - members []string -func (_e *Repository_Expecter) RoleCheckMembersExists(ctx interface{}, roleID interface{}, members interface{}) *Repository_RoleCheckMembersExists_Call { - return &Repository_RoleCheckMembersExists_Call{Call: _e.mock.On("RoleCheckMembersExists", ctx, roleID, members)} -} - -func (_c *Repository_RoleCheckMembersExists_Call) Run(run func(ctx context.Context, roleID string, members []string)) *Repository_RoleCheckMembersExists_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RoleCheckMembersExists_Call) Return(b bool, err error) *Repository_RoleCheckMembersExists_Call { - _c.Call.Return(b, err) - return _c -} - -func (_c *Repository_RoleCheckMembersExists_Call) RunAndReturn(run func(ctx context.Context, roleID string, members []string) (bool, error)) *Repository_RoleCheckMembersExists_Call { - _c.Call.Return(run) - return _c -} - -// RoleListActions provides a mock function for the type Repository -func (_mock *Repository) RoleListActions(ctx context.Context, roleID string) ([]string, error) { - ret := _mock.Called(ctx, roleID) - - if len(ret) == 0 { - panic("no return value specified for RoleListActions") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) ([]string, error)); ok { - return returnFunc(ctx, roleID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string) []string); ok { - r0 = returnFunc(ctx, roleID) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = returnFunc(ctx, roleID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RoleListActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleListActions' -type Repository_RoleListActions_Call struct { - *mock.Call -} - -// RoleListActions is a helper method to define mock.On call -// - ctx context.Context -// - roleID string -func (_e *Repository_Expecter) RoleListActions(ctx interface{}, roleID interface{}) *Repository_RoleListActions_Call { - return &Repository_RoleListActions_Call{Call: _e.mock.On("RoleListActions", ctx, roleID)} -} - -func (_c *Repository_RoleListActions_Call) Run(run func(ctx context.Context, roleID string)) *Repository_RoleListActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RoleListActions_Call) Return(strings []string, err error) *Repository_RoleListActions_Call { - _c.Call.Return(strings, err) - return _c -} - -func (_c *Repository_RoleListActions_Call) RunAndReturn(run func(ctx context.Context, roleID string) ([]string, error)) *Repository_RoleListActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleListMembers provides a mock function for the type Repository -func (_mock *Repository) RoleListMembers(ctx context.Context, roleID string, limit uint64, offset uint64) (roles.MembersPage, error) { - ret := _mock.Called(ctx, roleID, limit, offset) - - if len(ret) == 0 { - panic("no return value specified for RoleListMembers") - } - - var r0 roles.MembersPage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, uint64, uint64) (roles.MembersPage, error)); ok { - return returnFunc(ctx, roleID, limit, offset) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, uint64, uint64) roles.MembersPage); ok { - r0 = returnFunc(ctx, roleID, limit, offset) - } else { - r0 = ret.Get(0).(roles.MembersPage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, uint64, uint64) error); ok { - r1 = returnFunc(ctx, roleID, limit, offset) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RoleListMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleListMembers' -type Repository_RoleListMembers_Call struct { - *mock.Call -} - -// RoleListMembers is a helper method to define mock.On call -// - ctx context.Context -// - roleID string -// - limit uint64 -// - offset uint64 -func (_e *Repository_Expecter) RoleListMembers(ctx interface{}, roleID interface{}, limit interface{}, offset interface{}) *Repository_RoleListMembers_Call { - return &Repository_RoleListMembers_Call{Call: _e.mock.On("RoleListMembers", ctx, roleID, limit, offset)} -} - -func (_c *Repository_RoleListMembers_Call) Run(run func(ctx context.Context, roleID string, limit uint64, offset uint64)) *Repository_RoleListMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 uint64 - if args[2] != nil { - arg2 = args[2].(uint64) - } - var arg3 uint64 - if args[3] != nil { - arg3 = args[3].(uint64) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Repository_RoleListMembers_Call) Return(membersPage roles.MembersPage, err error) *Repository_RoleListMembers_Call { - _c.Call.Return(membersPage, err) - return _c -} - -func (_c *Repository_RoleListMembers_Call) RunAndReturn(run func(ctx context.Context, roleID string, limit uint64, offset uint64) (roles.MembersPage, error)) *Repository_RoleListMembers_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveActions provides a mock function for the type Repository -func (_mock *Repository) RoleRemoveActions(ctx context.Context, role roles.Role, actions []string) error { - ret := _mock.Called(ctx, role, actions) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveActions") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role, []string) error); ok { - r0 = returnFunc(ctx, role, actions) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RoleRemoveActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveActions' -type Repository_RoleRemoveActions_Call struct { - *mock.Call -} - -// RoleRemoveActions is a helper method to define mock.On call -// - ctx context.Context -// - role roles.Role -// - actions []string -func (_e *Repository_Expecter) RoleRemoveActions(ctx interface{}, role interface{}, actions interface{}) *Repository_RoleRemoveActions_Call { - return &Repository_RoleRemoveActions_Call{Call: _e.mock.On("RoleRemoveActions", ctx, role, actions)} -} - -func (_c *Repository_RoleRemoveActions_Call) Run(run func(ctx context.Context, role roles.Role, actions []string)) *Repository_RoleRemoveActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RoleRemoveActions_Call) Return(err error) *Repository_RoleRemoveActions_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RoleRemoveActions_Call) RunAndReturn(run func(ctx context.Context, role roles.Role, actions []string) error) *Repository_RoleRemoveActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveAllActions provides a mock function for the type Repository -func (_mock *Repository) RoleRemoveAllActions(ctx context.Context, role roles.Role) error { - ret := _mock.Called(ctx, role) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveAllActions") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role) error); ok { - r0 = returnFunc(ctx, role) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RoleRemoveAllActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveAllActions' -type Repository_RoleRemoveAllActions_Call struct { - *mock.Call -} - -// RoleRemoveAllActions is a helper method to define mock.On call -// - ctx context.Context -// - role roles.Role -func (_e *Repository_Expecter) RoleRemoveAllActions(ctx interface{}, role interface{}) *Repository_RoleRemoveAllActions_Call { - return &Repository_RoleRemoveAllActions_Call{Call: _e.mock.On("RoleRemoveAllActions", ctx, role)} -} - -func (_c *Repository_RoleRemoveAllActions_Call) Run(run func(ctx context.Context, role roles.Role)) *Repository_RoleRemoveAllActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RoleRemoveAllActions_Call) Return(err error) *Repository_RoleRemoveAllActions_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RoleRemoveAllActions_Call) RunAndReturn(run func(ctx context.Context, role roles.Role) error) *Repository_RoleRemoveAllActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveAllMembers provides a mock function for the type Repository -func (_mock *Repository) RoleRemoveAllMembers(ctx context.Context, role roles.Role) error { - ret := _mock.Called(ctx, role) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveAllMembers") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role) error); ok { - r0 = returnFunc(ctx, role) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RoleRemoveAllMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveAllMembers' -type Repository_RoleRemoveAllMembers_Call struct { - *mock.Call -} - -// RoleRemoveAllMembers is a helper method to define mock.On call -// - ctx context.Context -// - role roles.Role -func (_e *Repository_Expecter) RoleRemoveAllMembers(ctx interface{}, role interface{}) *Repository_RoleRemoveAllMembers_Call { - return &Repository_RoleRemoveAllMembers_Call{Call: _e.mock.On("RoleRemoveAllMembers", ctx, role)} -} - -func (_c *Repository_RoleRemoveAllMembers_Call) Run(run func(ctx context.Context, role roles.Role)) *Repository_RoleRemoveAllMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RoleRemoveAllMembers_Call) Return(err error) *Repository_RoleRemoveAllMembers_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RoleRemoveAllMembers_Call) RunAndReturn(run func(ctx context.Context, role roles.Role) error) *Repository_RoleRemoveAllMembers_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveMembers provides a mock function for the type Repository -func (_mock *Repository) RoleRemoveMembers(ctx context.Context, role roles.Role, members []string) error { - ret := _mock.Called(ctx, role, members) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveMembers") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role, []string) error); ok { - r0 = returnFunc(ctx, role, members) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RoleRemoveMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveMembers' -type Repository_RoleRemoveMembers_Call struct { - *mock.Call -} - -// RoleRemoveMembers is a helper method to define mock.On call -// - ctx context.Context -// - role roles.Role -// - members []string -func (_e *Repository_Expecter) RoleRemoveMembers(ctx interface{}, role interface{}, members interface{}) *Repository_RoleRemoveMembers_Call { - return &Repository_RoleRemoveMembers_Call{Call: _e.mock.On("RoleRemoveMembers", ctx, role, members)} -} - -func (_c *Repository_RoleRemoveMembers_Call) Run(run func(ctx context.Context, role roles.Role, members []string)) *Repository_RoleRemoveMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RoleRemoveMembers_Call) Return(err error) *Repository_RoleRemoveMembers_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RoleRemoveMembers_Call) RunAndReturn(run func(ctx context.Context, role roles.Role, members []string) error) *Repository_RoleRemoveMembers_Call { - _c.Call.Return(run) - return _c -} - -// Save provides a mock function for the type Repository -func (_mock *Repository) Save(ctx context.Context, chs ...channels.Channel) ([]channels.Channel, error) { - var tmpRet mock.Arguments - if len(chs) > 0 { - tmpRet = _mock.Called(ctx, chs) - } else { - tmpRet = _mock.Called(ctx) - } - ret := tmpRet - - if len(ret) == 0 { - panic("no return value specified for Save") - } - - var r0 []channels.Channel - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, ...channels.Channel) ([]channels.Channel, error)); ok { - return returnFunc(ctx, chs...) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, ...channels.Channel) []channels.Channel); ok { - r0 = returnFunc(ctx, chs...) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]channels.Channel) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, ...channels.Channel) error); ok { - r1 = returnFunc(ctx, chs...) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_Save_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Save' -type Repository_Save_Call struct { - *mock.Call -} - -// Save is a helper method to define mock.On call -// - ctx context.Context -// - chs ...channels.Channel -func (_e *Repository_Expecter) Save(ctx interface{}, chs ...interface{}) *Repository_Save_Call { - return &Repository_Save_Call{Call: _e.mock.On("Save", - append([]interface{}{ctx}, chs...)...)} -} - -func (_c *Repository_Save_Call) Run(run func(ctx context.Context, chs ...channels.Channel)) *Repository_Save_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 []channels.Channel - var variadicArgs []channels.Channel - if len(args) > 1 { - variadicArgs = args[1].([]channels.Channel) - } - arg1 = variadicArgs - run( - arg0, - arg1..., - ) - }) - return _c -} - -func (_c *Repository_Save_Call) Return(channels1 []channels.Channel, err error) *Repository_Save_Call { - _c.Call.Return(channels1, err) - return _c -} - -func (_c *Repository_Save_Call) RunAndReturn(run func(ctx context.Context, chs ...channels.Channel) ([]channels.Channel, error)) *Repository_Save_Call { - _c.Call.Return(run) - return _c -} - -// SetParentGroup provides a mock function for the type Repository -func (_mock *Repository) SetParentGroup(ctx context.Context, ch channels.Channel) error { - ret := _mock.Called(ctx, ch) - - if len(ret) == 0 { - panic("no return value specified for SetParentGroup") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, channels.Channel) error); ok { - r0 = returnFunc(ctx, ch) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_SetParentGroup_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'SetParentGroup' -type Repository_SetParentGroup_Call struct { - *mock.Call -} - -// SetParentGroup is a helper method to define mock.On call -// - ctx context.Context -// - ch channels.Channel -func (_e *Repository_Expecter) SetParentGroup(ctx interface{}, ch interface{}) *Repository_SetParentGroup_Call { - return &Repository_SetParentGroup_Call{Call: _e.mock.On("SetParentGroup", ctx, ch)} -} - -func (_c *Repository_SetParentGroup_Call) Run(run func(ctx context.Context, ch channels.Channel)) *Repository_SetParentGroup_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 channels.Channel - if args[1] != nil { - arg1 = args[1].(channels.Channel) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_SetParentGroup_Call) Return(err error) *Repository_SetParentGroup_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_SetParentGroup_Call) RunAndReturn(run func(ctx context.Context, ch channels.Channel) error) *Repository_SetParentGroup_Call { - _c.Call.Return(run) - return _c -} - -// UnsetParentGroupFromChannels provides a mock function for the type Repository -func (_mock *Repository) UnsetParentGroupFromChannels(ctx context.Context, parentGroupID string) error { - ret := _mock.Called(ctx, parentGroupID) - - if len(ret) == 0 { - panic("no return value specified for UnsetParentGroupFromChannels") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) error); ok { - r0 = returnFunc(ctx, parentGroupID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_UnsetParentGroupFromChannels_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UnsetParentGroupFromChannels' -type Repository_UnsetParentGroupFromChannels_Call struct { - *mock.Call -} - -// UnsetParentGroupFromChannels is a helper method to define mock.On call -// - ctx context.Context -// - parentGroupID string -func (_e *Repository_Expecter) UnsetParentGroupFromChannels(ctx interface{}, parentGroupID interface{}) *Repository_UnsetParentGroupFromChannels_Call { - return &Repository_UnsetParentGroupFromChannels_Call{Call: _e.mock.On("UnsetParentGroupFromChannels", ctx, parentGroupID)} -} - -func (_c *Repository_UnsetParentGroupFromChannels_Call) Run(run func(ctx context.Context, parentGroupID string)) *Repository_UnsetParentGroupFromChannels_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_UnsetParentGroupFromChannels_Call) Return(err error) *Repository_UnsetParentGroupFromChannels_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_UnsetParentGroupFromChannels_Call) RunAndReturn(run func(ctx context.Context, parentGroupID string) error) *Repository_UnsetParentGroupFromChannels_Call { - _c.Call.Return(run) - return _c -} - -// Update provides a mock function for the type Repository -func (_mock *Repository) Update(ctx context.Context, c channels.Channel) (channels.Channel, error) { - ret := _mock.Called(ctx, c) - - if len(ret) == 0 { - panic("no return value specified for Update") - } - - var r0 channels.Channel - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, channels.Channel) (channels.Channel, error)); ok { - return returnFunc(ctx, c) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, channels.Channel) channels.Channel); ok { - r0 = returnFunc(ctx, c) - } else { - r0 = ret.Get(0).(channels.Channel) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, channels.Channel) error); ok { - r1 = returnFunc(ctx, c) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_Update_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Update' -type Repository_Update_Call struct { - *mock.Call -} - -// Update is a helper method to define mock.On call -// - ctx context.Context -// - c channels.Channel -func (_e *Repository_Expecter) Update(ctx interface{}, c interface{}) *Repository_Update_Call { - return &Repository_Update_Call{Call: _e.mock.On("Update", ctx, c)} -} - -func (_c *Repository_Update_Call) Run(run func(ctx context.Context, c channels.Channel)) *Repository_Update_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 channels.Channel - if args[1] != nil { - arg1 = args[1].(channels.Channel) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_Update_Call) Return(channel channels.Channel, err error) *Repository_Update_Call { - _c.Call.Return(channel, err) - return _c -} - -func (_c *Repository_Update_Call) RunAndReturn(run func(ctx context.Context, c channels.Channel) (channels.Channel, error)) *Repository_Update_Call { - _c.Call.Return(run) - return _c -} - -// UpdateRole provides a mock function for the type Repository -func (_mock *Repository) UpdateRole(ctx context.Context, ro roles.Role) (roles.Role, error) { - ret := _mock.Called(ctx, ro) - - if len(ret) == 0 { - panic("no return value specified for UpdateRole") - } - - var r0 roles.Role - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role) (roles.Role, error)); ok { - return returnFunc(ctx, ro) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role) roles.Role); ok { - r0 = returnFunc(ctx, ro) - } else { - r0 = ret.Get(0).(roles.Role) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, roles.Role) error); ok { - r1 = returnFunc(ctx, ro) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_UpdateRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateRole' -type Repository_UpdateRole_Call struct { - *mock.Call -} - -// UpdateRole is a helper method to define mock.On call -// - ctx context.Context -// - ro roles.Role -func (_e *Repository_Expecter) UpdateRole(ctx interface{}, ro interface{}) *Repository_UpdateRole_Call { - return &Repository_UpdateRole_Call{Call: _e.mock.On("UpdateRole", ctx, ro)} -} - -func (_c *Repository_UpdateRole_Call) Run(run func(ctx context.Context, ro roles.Role)) *Repository_UpdateRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_UpdateRole_Call) Return(role roles.Role, err error) *Repository_UpdateRole_Call { - _c.Call.Return(role, err) - return _c -} - -func (_c *Repository_UpdateRole_Call) RunAndReturn(run func(ctx context.Context, ro roles.Role) (roles.Role, error)) *Repository_UpdateRole_Call { - _c.Call.Return(run) - return _c -} - -// UpdateTags provides a mock function for the type Repository -func (_mock *Repository) UpdateTags(ctx context.Context, ch channels.Channel) (channels.Channel, error) { - ret := _mock.Called(ctx, ch) - - if len(ret) == 0 { - panic("no return value specified for UpdateTags") - } - - var r0 channels.Channel - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, channels.Channel) (channels.Channel, error)); ok { - return returnFunc(ctx, ch) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, channels.Channel) channels.Channel); ok { - r0 = returnFunc(ctx, ch) - } else { - r0 = ret.Get(0).(channels.Channel) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, channels.Channel) error); ok { - r1 = returnFunc(ctx, ch) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_UpdateTags_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateTags' -type Repository_UpdateTags_Call struct { - *mock.Call -} - -// UpdateTags is a helper method to define mock.On call -// - ctx context.Context -// - ch channels.Channel -func (_e *Repository_Expecter) UpdateTags(ctx interface{}, ch interface{}) *Repository_UpdateTags_Call { - return &Repository_UpdateTags_Call{Call: _e.mock.On("UpdateTags", ctx, ch)} -} - -func (_c *Repository_UpdateTags_Call) Run(run func(ctx context.Context, ch channels.Channel)) *Repository_UpdateTags_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 channels.Channel - if args[1] != nil { - arg1 = args[1].(channels.Channel) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_UpdateTags_Call) Return(channel channels.Channel, err error) *Repository_UpdateTags_Call { - _c.Call.Return(channel, err) - return _c -} - -func (_c *Repository_UpdateTags_Call) RunAndReturn(run func(ctx context.Context, ch channels.Channel) (channels.Channel, error)) *Repository_UpdateTags_Call { - _c.Call.Return(run) - return _c -} diff --git a/channels/mocks/service.go b/channels/mocks/service.go deleted file mode 100644 index b1f390943..000000000 --- a/channels/mocks/service.go +++ /dev/null @@ -1,2479 +0,0 @@ -// Copyright (c) Abstract Machines - -// SPDX-License-Identifier: Apache-2.0 - -// Code generated by mockery; DO NOT EDIT. -// github.com/vektra/mockery -// template: testify - -package mocks - -import ( - "context" - - "github.com/absmach/magistrala/channels" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/connections" - "github.com/absmach/magistrala/pkg/roles" - mock "github.com/stretchr/testify/mock" -) - -// NewService creates a new instance of Service. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. -// The first argument is typically a *testing.T value. -func NewService(t interface { - mock.TestingT - Cleanup(func()) -}) *Service { - mock := &Service{} - mock.Mock.Test(t) - - t.Cleanup(func() { mock.AssertExpectations(t) }) - - return mock -} - -// Service is an autogenerated mock type for the Service type -type Service struct { - mock.Mock -} - -type Service_Expecter struct { - mock *mock.Mock -} - -func (_m *Service) EXPECT() *Service_Expecter { - return &Service_Expecter{mock: &_m.Mock} -} - -// AddRole provides a mock function for the type Service -func (_mock *Service) AddRole(ctx context.Context, session authn.Session, entityID string, roleName string, optionalActions []string, optionalMembers []string) (roles.RoleProvision, error) { - ret := _mock.Called(ctx, session, entityID, roleName, optionalActions, optionalMembers) - - if len(ret) == 0 { - panic("no return value specified for AddRole") - } - - var r0 roles.RoleProvision - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string, []string) (roles.RoleProvision, error)); ok { - return returnFunc(ctx, session, entityID, roleName, optionalActions, optionalMembers) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string, []string) roles.RoleProvision); ok { - r0 = returnFunc(ctx, session, entityID, roleName, optionalActions, optionalMembers) - } else { - r0 = ret.Get(0).(roles.RoleProvision) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, []string, []string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleName, optionalActions, optionalMembers) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_AddRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'AddRole' -type Service_AddRole_Call struct { - *mock.Call -} - -// AddRole is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleName string -// - optionalActions []string -// - optionalMembers []string -func (_e *Service_Expecter) AddRole(ctx interface{}, session interface{}, entityID interface{}, roleName interface{}, optionalActions interface{}, optionalMembers interface{}) *Service_AddRole_Call { - return &Service_AddRole_Call{Call: _e.mock.On("AddRole", ctx, session, entityID, roleName, optionalActions, optionalMembers)} -} - -func (_c *Service_AddRole_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleName string, optionalActions []string, optionalMembers []string)) *Service_AddRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - var arg5 []string - if args[5] != nil { - arg5 = args[5].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - arg5, - ) - }) - return _c -} - -func (_c *Service_AddRole_Call) Return(roleProvision roles.RoleProvision, err error) *Service_AddRole_Call { - _c.Call.Return(roleProvision, err) - return _c -} - -func (_c *Service_AddRole_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleName string, optionalActions []string, optionalMembers []string) (roles.RoleProvision, error)) *Service_AddRole_Call { - _c.Call.Return(run) - return _c -} - -// Connect provides a mock function for the type Service -func (_mock *Service) Connect(ctx context.Context, session authn.Session, chIDs []string, clIDs []string, connType []connections.ConnType) error { - ret := _mock.Called(ctx, session, chIDs, clIDs, connType) - - if len(ret) == 0 { - panic("no return value specified for Connect") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, []string, []string, []connections.ConnType) error); ok { - r0 = returnFunc(ctx, session, chIDs, clIDs, connType) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_Connect_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Connect' -type Service_Connect_Call struct { - *mock.Call -} - -// Connect is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - chIDs []string -// - clIDs []string -// - connType []connections.ConnType -func (_e *Service_Expecter) Connect(ctx interface{}, session interface{}, chIDs interface{}, clIDs interface{}, connType interface{}) *Service_Connect_Call { - return &Service_Connect_Call{Call: _e.mock.On("Connect", ctx, session, chIDs, clIDs, connType)} -} - -func (_c *Service_Connect_Call) Run(run func(ctx context.Context, session authn.Session, chIDs []string, clIDs []string, connType []connections.ConnType)) *Service_Connect_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - var arg3 []string - if args[3] != nil { - arg3 = args[3].([]string) - } - var arg4 []connections.ConnType - if args[4] != nil { - arg4 = args[4].([]connections.ConnType) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_Connect_Call) Return(err error) *Service_Connect_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_Connect_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, chIDs []string, clIDs []string, connType []connections.ConnType) error) *Service_Connect_Call { - _c.Call.Return(run) - return _c -} - -// CreateChannels provides a mock function for the type Service -func (_mock *Service) CreateChannels(ctx context.Context, session authn.Session, channels1 ...channels.Channel) ([]channels.Channel, []roles.RoleProvision, error) { - var tmpRet mock.Arguments - if len(channels1) > 0 { - tmpRet = _mock.Called(ctx, session, channels1) - } else { - tmpRet = _mock.Called(ctx, session) - } - ret := tmpRet - - if len(ret) == 0 { - panic("no return value specified for CreateChannels") - } - - var r0 []channels.Channel - var r1 []roles.RoleProvision - var r2 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, ...channels.Channel) ([]channels.Channel, []roles.RoleProvision, error)); ok { - return returnFunc(ctx, session, channels1...) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, ...channels.Channel) []channels.Channel); ok { - r0 = returnFunc(ctx, session, channels1...) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]channels.Channel) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, ...channels.Channel) []roles.RoleProvision); ok { - r1 = returnFunc(ctx, session, channels1...) - } else { - if ret.Get(1) != nil { - r1 = ret.Get(1).([]roles.RoleProvision) - } - } - if returnFunc, ok := ret.Get(2).(func(context.Context, authn.Session, ...channels.Channel) error); ok { - r2 = returnFunc(ctx, session, channels1...) - } else { - r2 = ret.Error(2) - } - return r0, r1, r2 -} - -// Service_CreateChannels_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'CreateChannels' -type Service_CreateChannels_Call struct { - *mock.Call -} - -// CreateChannels is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - channels1 ...channels.Channel -func (_e *Service_Expecter) CreateChannels(ctx interface{}, session interface{}, channels1 ...interface{}) *Service_CreateChannels_Call { - return &Service_CreateChannels_Call{Call: _e.mock.On("CreateChannels", - append([]interface{}{ctx, session}, channels1...)...)} -} - -func (_c *Service_CreateChannels_Call) Run(run func(ctx context.Context, session authn.Session, channels1 ...channels.Channel)) *Service_CreateChannels_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 []channels.Channel - var variadicArgs []channels.Channel - if len(args) > 2 { - variadicArgs = args[2].([]channels.Channel) - } - arg2 = variadicArgs - run( - arg0, - arg1, - arg2..., - ) - }) - return _c -} - -func (_c *Service_CreateChannels_Call) Return(channels11 []channels.Channel, roleProvisions []roles.RoleProvision, err error) *Service_CreateChannels_Call { - _c.Call.Return(channels11, roleProvisions, err) - return _c -} - -func (_c *Service_CreateChannels_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, channels1 ...channels.Channel) ([]channels.Channel, []roles.RoleProvision, error)) *Service_CreateChannels_Call { - _c.Call.Return(run) - return _c -} - -// DisableChannel provides a mock function for the type Service -func (_mock *Service) DisableChannel(ctx context.Context, session authn.Session, id string) (channels.Channel, error) { - ret := _mock.Called(ctx, session, id) - - if len(ret) == 0 { - panic("no return value specified for DisableChannel") - } - - var r0 channels.Channel - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) (channels.Channel, error)); ok { - return returnFunc(ctx, session, id) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) channels.Channel); ok { - r0 = returnFunc(ctx, session, id) - } else { - r0 = ret.Get(0).(channels.Channel) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string) error); ok { - r1 = returnFunc(ctx, session, id) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_DisableChannel_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'DisableChannel' -type Service_DisableChannel_Call struct { - *mock.Call -} - -// DisableChannel is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - id string -func (_e *Service_Expecter) DisableChannel(ctx interface{}, session interface{}, id interface{}) *Service_DisableChannel_Call { - return &Service_DisableChannel_Call{Call: _e.mock.On("DisableChannel", ctx, session, id)} -} - -func (_c *Service_DisableChannel_Call) Run(run func(ctx context.Context, session authn.Session, id string)) *Service_DisableChannel_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_DisableChannel_Call) Return(channel channels.Channel, err error) *Service_DisableChannel_Call { - _c.Call.Return(channel, err) - return _c -} - -func (_c *Service_DisableChannel_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, id string) (channels.Channel, error)) *Service_DisableChannel_Call { - _c.Call.Return(run) - return _c -} - -// Disconnect provides a mock function for the type Service -func (_mock *Service) Disconnect(ctx context.Context, session authn.Session, chIDs []string, clIDs []string, connType []connections.ConnType) error { - ret := _mock.Called(ctx, session, chIDs, clIDs, connType) - - if len(ret) == 0 { - panic("no return value specified for Disconnect") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, []string, []string, []connections.ConnType) error); ok { - r0 = returnFunc(ctx, session, chIDs, clIDs, connType) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_Disconnect_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Disconnect' -type Service_Disconnect_Call struct { - *mock.Call -} - -// Disconnect is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - chIDs []string -// - clIDs []string -// - connType []connections.ConnType -func (_e *Service_Expecter) Disconnect(ctx interface{}, session interface{}, chIDs interface{}, clIDs interface{}, connType interface{}) *Service_Disconnect_Call { - return &Service_Disconnect_Call{Call: _e.mock.On("Disconnect", ctx, session, chIDs, clIDs, connType)} -} - -func (_c *Service_Disconnect_Call) Run(run func(ctx context.Context, session authn.Session, chIDs []string, clIDs []string, connType []connections.ConnType)) *Service_Disconnect_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - var arg3 []string - if args[3] != nil { - arg3 = args[3].([]string) - } - var arg4 []connections.ConnType - if args[4] != nil { - arg4 = args[4].([]connections.ConnType) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_Disconnect_Call) Return(err error) *Service_Disconnect_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_Disconnect_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, chIDs []string, clIDs []string, connType []connections.ConnType) error) *Service_Disconnect_Call { - _c.Call.Return(run) - return _c -} - -// EnableChannel provides a mock function for the type Service -func (_mock *Service) EnableChannel(ctx context.Context, session authn.Session, id string) (channels.Channel, error) { - ret := _mock.Called(ctx, session, id) - - if len(ret) == 0 { - panic("no return value specified for EnableChannel") - } - - var r0 channels.Channel - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) (channels.Channel, error)); ok { - return returnFunc(ctx, session, id) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) channels.Channel); ok { - r0 = returnFunc(ctx, session, id) - } else { - r0 = ret.Get(0).(channels.Channel) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string) error); ok { - r1 = returnFunc(ctx, session, id) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_EnableChannel_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'EnableChannel' -type Service_EnableChannel_Call struct { - *mock.Call -} - -// EnableChannel is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - id string -func (_e *Service_Expecter) EnableChannel(ctx interface{}, session interface{}, id interface{}) *Service_EnableChannel_Call { - return &Service_EnableChannel_Call{Call: _e.mock.On("EnableChannel", ctx, session, id)} -} - -func (_c *Service_EnableChannel_Call) Run(run func(ctx context.Context, session authn.Session, id string)) *Service_EnableChannel_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_EnableChannel_Call) Return(channel channels.Channel, err error) *Service_EnableChannel_Call { - _c.Call.Return(channel, err) - return _c -} - -func (_c *Service_EnableChannel_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, id string) (channels.Channel, error)) *Service_EnableChannel_Call { - _c.Call.Return(run) - return _c -} - -// ListAvailableActions provides a mock function for the type Service -func (_mock *Service) ListAvailableActions(ctx context.Context, session authn.Session) ([]string, error) { - ret := _mock.Called(ctx, session) - - if len(ret) == 0 { - panic("no return value specified for ListAvailableActions") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session) ([]string, error)); ok { - return returnFunc(ctx, session) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session) []string); ok { - r0 = returnFunc(ctx, session) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session) error); ok { - r1 = returnFunc(ctx, session) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_ListAvailableActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ListAvailableActions' -type Service_ListAvailableActions_Call struct { - *mock.Call -} - -// ListAvailableActions is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -func (_e *Service_Expecter) ListAvailableActions(ctx interface{}, session interface{}) *Service_ListAvailableActions_Call { - return &Service_ListAvailableActions_Call{Call: _e.mock.On("ListAvailableActions", ctx, session)} -} - -func (_c *Service_ListAvailableActions_Call) Run(run func(ctx context.Context, session authn.Session)) *Service_ListAvailableActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Service_ListAvailableActions_Call) Return(strings []string, err error) *Service_ListAvailableActions_Call { - _c.Call.Return(strings, err) - return _c -} - -func (_c *Service_ListAvailableActions_Call) RunAndReturn(run func(ctx context.Context, session authn.Session) ([]string, error)) *Service_ListAvailableActions_Call { - _c.Call.Return(run) - return _c -} - -// ListChannels provides a mock function for the type Service -func (_mock *Service) ListChannels(ctx context.Context, session authn.Session, pm channels.Page) (channels.ChannelsPage, error) { - ret := _mock.Called(ctx, session, pm) - - if len(ret) == 0 { - panic("no return value specified for ListChannels") - } - - var r0 channels.ChannelsPage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, channels.Page) (channels.ChannelsPage, error)); ok { - return returnFunc(ctx, session, pm) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, channels.Page) channels.ChannelsPage); ok { - r0 = returnFunc(ctx, session, pm) - } else { - r0 = ret.Get(0).(channels.ChannelsPage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, channels.Page) error); ok { - r1 = returnFunc(ctx, session, pm) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_ListChannels_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ListChannels' -type Service_ListChannels_Call struct { - *mock.Call -} - -// ListChannels is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - pm channels.Page -func (_e *Service_Expecter) ListChannels(ctx interface{}, session interface{}, pm interface{}) *Service_ListChannels_Call { - return &Service_ListChannels_Call{Call: _e.mock.On("ListChannels", ctx, session, pm)} -} - -func (_c *Service_ListChannels_Call) Run(run func(ctx context.Context, session authn.Session, pm channels.Page)) *Service_ListChannels_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 channels.Page - if args[2] != nil { - arg2 = args[2].(channels.Page) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_ListChannels_Call) Return(channelsPage channels.ChannelsPage, err error) *Service_ListChannels_Call { - _c.Call.Return(channelsPage, err) - return _c -} - -func (_c *Service_ListChannels_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, pm channels.Page) (channels.ChannelsPage, error)) *Service_ListChannels_Call { - _c.Call.Return(run) - return _c -} - -// ListEntityMembers provides a mock function for the type Service -func (_mock *Service) ListEntityMembers(ctx context.Context, session authn.Session, entityID string, pq roles.MembersRolePageQuery) (roles.MembersRolePage, error) { - ret := _mock.Called(ctx, session, entityID, pq) - - if len(ret) == 0 { - panic("no return value specified for ListEntityMembers") - } - - var r0 roles.MembersRolePage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, roles.MembersRolePageQuery) (roles.MembersRolePage, error)); ok { - return returnFunc(ctx, session, entityID, pq) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, roles.MembersRolePageQuery) roles.MembersRolePage); ok { - r0 = returnFunc(ctx, session, entityID, pq) - } else { - r0 = ret.Get(0).(roles.MembersRolePage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, roles.MembersRolePageQuery) error); ok { - r1 = returnFunc(ctx, session, entityID, pq) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_ListEntityMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ListEntityMembers' -type Service_ListEntityMembers_Call struct { - *mock.Call -} - -// ListEntityMembers is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - pq roles.MembersRolePageQuery -func (_e *Service_Expecter) ListEntityMembers(ctx interface{}, session interface{}, entityID interface{}, pq interface{}) *Service_ListEntityMembers_Call { - return &Service_ListEntityMembers_Call{Call: _e.mock.On("ListEntityMembers", ctx, session, entityID, pq)} -} - -func (_c *Service_ListEntityMembers_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, pq roles.MembersRolePageQuery)) *Service_ListEntityMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 roles.MembersRolePageQuery - if args[3] != nil { - arg3 = args[3].(roles.MembersRolePageQuery) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_ListEntityMembers_Call) Return(membersRolePage roles.MembersRolePage, err error) *Service_ListEntityMembers_Call { - _c.Call.Return(membersRolePage, err) - return _c -} - -func (_c *Service_ListEntityMembers_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, pq roles.MembersRolePageQuery) (roles.MembersRolePage, error)) *Service_ListEntityMembers_Call { - _c.Call.Return(run) - return _c -} - -// ListUserChannels provides a mock function for the type Service -func (_mock *Service) ListUserChannels(ctx context.Context, session authn.Session, userID string, pm channels.Page) (channels.ChannelsPage, error) { - ret := _mock.Called(ctx, session, userID, pm) - - if len(ret) == 0 { - panic("no return value specified for ListUserChannels") - } - - var r0 channels.ChannelsPage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, channels.Page) (channels.ChannelsPage, error)); ok { - return returnFunc(ctx, session, userID, pm) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, channels.Page) channels.ChannelsPage); ok { - r0 = returnFunc(ctx, session, userID, pm) - } else { - r0 = ret.Get(0).(channels.ChannelsPage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, channels.Page) error); ok { - r1 = returnFunc(ctx, session, userID, pm) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_ListUserChannels_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ListUserChannels' -type Service_ListUserChannels_Call struct { - *mock.Call -} - -// ListUserChannels is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - userID string -// - pm channels.Page -func (_e *Service_Expecter) ListUserChannels(ctx interface{}, session interface{}, userID interface{}, pm interface{}) *Service_ListUserChannels_Call { - return &Service_ListUserChannels_Call{Call: _e.mock.On("ListUserChannels", ctx, session, userID, pm)} -} - -func (_c *Service_ListUserChannels_Call) Run(run func(ctx context.Context, session authn.Session, userID string, pm channels.Page)) *Service_ListUserChannels_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 channels.Page - if args[3] != nil { - arg3 = args[3].(channels.Page) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_ListUserChannels_Call) Return(channelsPage channels.ChannelsPage, err error) *Service_ListUserChannels_Call { - _c.Call.Return(channelsPage, err) - return _c -} - -func (_c *Service_ListUserChannels_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, userID string, pm channels.Page) (channels.ChannelsPage, error)) *Service_ListUserChannels_Call { - _c.Call.Return(run) - return _c -} - -// RemoveChannel provides a mock function for the type Service -func (_mock *Service) RemoveChannel(ctx context.Context, session authn.Session, id string) error { - ret := _mock.Called(ctx, session, id) - - if len(ret) == 0 { - panic("no return value specified for RemoveChannel") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) error); ok { - r0 = returnFunc(ctx, session, id) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RemoveChannel_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveChannel' -type Service_RemoveChannel_Call struct { - *mock.Call -} - -// RemoveChannel is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - id string -func (_e *Service_Expecter) RemoveChannel(ctx interface{}, session interface{}, id interface{}) *Service_RemoveChannel_Call { - return &Service_RemoveChannel_Call{Call: _e.mock.On("RemoveChannel", ctx, session, id)} -} - -func (_c *Service_RemoveChannel_Call) Run(run func(ctx context.Context, session authn.Session, id string)) *Service_RemoveChannel_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_RemoveChannel_Call) Return(err error) *Service_RemoveChannel_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RemoveChannel_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, id string) error) *Service_RemoveChannel_Call { - _c.Call.Return(run) - return _c -} - -// RemoveEntityMembers provides a mock function for the type Service -func (_mock *Service) RemoveEntityMembers(ctx context.Context, session authn.Session, entityID string, members []string) error { - ret := _mock.Called(ctx, session, entityID, members) - - if len(ret) == 0 { - panic("no return value specified for RemoveEntityMembers") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, []string) error); ok { - r0 = returnFunc(ctx, session, entityID, members) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RemoveEntityMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveEntityMembers' -type Service_RemoveEntityMembers_Call struct { - *mock.Call -} - -// RemoveEntityMembers is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - members []string -func (_e *Service_Expecter) RemoveEntityMembers(ctx interface{}, session interface{}, entityID interface{}, members interface{}) *Service_RemoveEntityMembers_Call { - return &Service_RemoveEntityMembers_Call{Call: _e.mock.On("RemoveEntityMembers", ctx, session, entityID, members)} -} - -func (_c *Service_RemoveEntityMembers_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, members []string)) *Service_RemoveEntityMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 []string - if args[3] != nil { - arg3 = args[3].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_RemoveEntityMembers_Call) Return(err error) *Service_RemoveEntityMembers_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RemoveEntityMembers_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, members []string) error) *Service_RemoveEntityMembers_Call { - _c.Call.Return(run) - return _c -} - -// RemoveMemberFromAllRoles provides a mock function for the type Service -func (_mock *Service) RemoveMemberFromAllRoles(ctx context.Context, session authn.Session, memberID string) error { - ret := _mock.Called(ctx, session, memberID) - - if len(ret) == 0 { - panic("no return value specified for RemoveMemberFromAllRoles") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) error); ok { - r0 = returnFunc(ctx, session, memberID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RemoveMemberFromAllRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveMemberFromAllRoles' -type Service_RemoveMemberFromAllRoles_Call struct { - *mock.Call -} - -// RemoveMemberFromAllRoles is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - memberID string -func (_e *Service_Expecter) RemoveMemberFromAllRoles(ctx interface{}, session interface{}, memberID interface{}) *Service_RemoveMemberFromAllRoles_Call { - return &Service_RemoveMemberFromAllRoles_Call{Call: _e.mock.On("RemoveMemberFromAllRoles", ctx, session, memberID)} -} - -func (_c *Service_RemoveMemberFromAllRoles_Call) Run(run func(ctx context.Context, session authn.Session, memberID string)) *Service_RemoveMemberFromAllRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_RemoveMemberFromAllRoles_Call) Return(err error) *Service_RemoveMemberFromAllRoles_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RemoveMemberFromAllRoles_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, memberID string) error) *Service_RemoveMemberFromAllRoles_Call { - _c.Call.Return(run) - return _c -} - -// RemoveParentGroup provides a mock function for the type Service -func (_mock *Service) RemoveParentGroup(ctx context.Context, session authn.Session, id string) error { - ret := _mock.Called(ctx, session, id) - - if len(ret) == 0 { - panic("no return value specified for RemoveParentGroup") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) error); ok { - r0 = returnFunc(ctx, session, id) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RemoveParentGroup_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveParentGroup' -type Service_RemoveParentGroup_Call struct { - *mock.Call -} - -// RemoveParentGroup is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - id string -func (_e *Service_Expecter) RemoveParentGroup(ctx interface{}, session interface{}, id interface{}) *Service_RemoveParentGroup_Call { - return &Service_RemoveParentGroup_Call{Call: _e.mock.On("RemoveParentGroup", ctx, session, id)} -} - -func (_c *Service_RemoveParentGroup_Call) Run(run func(ctx context.Context, session authn.Session, id string)) *Service_RemoveParentGroup_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_RemoveParentGroup_Call) Return(err error) *Service_RemoveParentGroup_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RemoveParentGroup_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, id string) error) *Service_RemoveParentGroup_Call { - _c.Call.Return(run) - return _c -} - -// RemoveRole provides a mock function for the type Service -func (_mock *Service) RemoveRole(ctx context.Context, session authn.Session, entityID string, roleID string) error { - ret := _mock.Called(ctx, session, entityID, roleID) - - if len(ret) == 0 { - panic("no return value specified for RemoveRole") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) error); ok { - r0 = returnFunc(ctx, session, entityID, roleID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RemoveRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveRole' -type Service_RemoveRole_Call struct { - *mock.Call -} - -// RemoveRole is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -func (_e *Service_Expecter) RemoveRole(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}) *Service_RemoveRole_Call { - return &Service_RemoveRole_Call{Call: _e.mock.On("RemoveRole", ctx, session, entityID, roleID)} -} - -func (_c *Service_RemoveRole_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string)) *Service_RemoveRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_RemoveRole_Call) Return(err error) *Service_RemoveRole_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RemoveRole_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string) error) *Service_RemoveRole_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveAllRoles provides a mock function for the type Service -func (_mock *Service) RetrieveAllRoles(ctx context.Context, session authn.Session, entityID string, limit uint64, offset uint64) (roles.RolePage, error) { - ret := _mock.Called(ctx, session, entityID, limit, offset) - - if len(ret) == 0 { - panic("no return value specified for RetrieveAllRoles") - } - - var r0 roles.RolePage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, uint64, uint64) (roles.RolePage, error)); ok { - return returnFunc(ctx, session, entityID, limit, offset) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, uint64, uint64) roles.RolePage); ok { - r0 = returnFunc(ctx, session, entityID, limit, offset) - } else { - r0 = ret.Get(0).(roles.RolePage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, uint64, uint64) error); ok { - r1 = returnFunc(ctx, session, entityID, limit, offset) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RetrieveAllRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveAllRoles' -type Service_RetrieveAllRoles_Call struct { - *mock.Call -} - -// RetrieveAllRoles is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - limit uint64 -// - offset uint64 -func (_e *Service_Expecter) RetrieveAllRoles(ctx interface{}, session interface{}, entityID interface{}, limit interface{}, offset interface{}) *Service_RetrieveAllRoles_Call { - return &Service_RetrieveAllRoles_Call{Call: _e.mock.On("RetrieveAllRoles", ctx, session, entityID, limit, offset)} -} - -func (_c *Service_RetrieveAllRoles_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, limit uint64, offset uint64)) *Service_RetrieveAllRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 uint64 - if args[3] != nil { - arg3 = args[3].(uint64) - } - var arg4 uint64 - if args[4] != nil { - arg4 = args[4].(uint64) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_RetrieveAllRoles_Call) Return(rolePage roles.RolePage, err error) *Service_RetrieveAllRoles_Call { - _c.Call.Return(rolePage, err) - return _c -} - -func (_c *Service_RetrieveAllRoles_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, limit uint64, offset uint64) (roles.RolePage, error)) *Service_RetrieveAllRoles_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveRole provides a mock function for the type Service -func (_mock *Service) RetrieveRole(ctx context.Context, session authn.Session, entityID string, roleID string) (roles.Role, error) { - ret := _mock.Called(ctx, session, entityID, roleID) - - if len(ret) == 0 { - panic("no return value specified for RetrieveRole") - } - - var r0 roles.Role - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) (roles.Role, error)); ok { - return returnFunc(ctx, session, entityID, roleID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) roles.Role); ok { - r0 = returnFunc(ctx, session, entityID, roleID) - } else { - r0 = ret.Get(0).(roles.Role) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RetrieveRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveRole' -type Service_RetrieveRole_Call struct { - *mock.Call -} - -// RetrieveRole is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -func (_e *Service_Expecter) RetrieveRole(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}) *Service_RetrieveRole_Call { - return &Service_RetrieveRole_Call{Call: _e.mock.On("RetrieveRole", ctx, session, entityID, roleID)} -} - -func (_c *Service_RetrieveRole_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string)) *Service_RetrieveRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_RetrieveRole_Call) Return(role roles.Role, err error) *Service_RetrieveRole_Call { - _c.Call.Return(role, err) - return _c -} - -func (_c *Service_RetrieveRole_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string) (roles.Role, error)) *Service_RetrieveRole_Call { - _c.Call.Return(run) - return _c -} - -// RoleAddActions provides a mock function for the type Service -func (_mock *Service) RoleAddActions(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string) ([]string, error) { - ret := _mock.Called(ctx, session, entityID, roleID, actions) - - if len(ret) == 0 { - panic("no return value specified for RoleAddActions") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) ([]string, error)); ok { - return returnFunc(ctx, session, entityID, roleID, actions) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) []string); ok { - r0 = returnFunc(ctx, session, entityID, roleID, actions) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, []string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID, actions) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RoleAddActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleAddActions' -type Service_RoleAddActions_Call struct { - *mock.Call -} - -// RoleAddActions is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - actions []string -func (_e *Service_Expecter) RoleAddActions(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, actions interface{}) *Service_RoleAddActions_Call { - return &Service_RoleAddActions_Call{Call: _e.mock.On("RoleAddActions", ctx, session, entityID, roleID, actions)} -} - -func (_c *Service_RoleAddActions_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string)) *Service_RoleAddActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_RoleAddActions_Call) Return(ops []string, err error) *Service_RoleAddActions_Call { - _c.Call.Return(ops, err) - return _c -} - -func (_c *Service_RoleAddActions_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string) ([]string, error)) *Service_RoleAddActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleAddMembers provides a mock function for the type Service -func (_mock *Service) RoleAddMembers(ctx context.Context, session authn.Session, entityID string, roleID string, members []string) ([]string, error) { - ret := _mock.Called(ctx, session, entityID, roleID, members) - - if len(ret) == 0 { - panic("no return value specified for RoleAddMembers") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) ([]string, error)); ok { - return returnFunc(ctx, session, entityID, roleID, members) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) []string); ok { - r0 = returnFunc(ctx, session, entityID, roleID, members) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, []string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID, members) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RoleAddMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleAddMembers' -type Service_RoleAddMembers_Call struct { - *mock.Call -} - -// RoleAddMembers is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - members []string -func (_e *Service_Expecter) RoleAddMembers(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, members interface{}) *Service_RoleAddMembers_Call { - return &Service_RoleAddMembers_Call{Call: _e.mock.On("RoleAddMembers", ctx, session, entityID, roleID, members)} -} - -func (_c *Service_RoleAddMembers_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, members []string)) *Service_RoleAddMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_RoleAddMembers_Call) Return(strings []string, err error) *Service_RoleAddMembers_Call { - _c.Call.Return(strings, err) - return _c -} - -func (_c *Service_RoleAddMembers_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, members []string) ([]string, error)) *Service_RoleAddMembers_Call { - _c.Call.Return(run) - return _c -} - -// RoleCheckActionsExists provides a mock function for the type Service -func (_mock *Service) RoleCheckActionsExists(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string) (bool, error) { - ret := _mock.Called(ctx, session, entityID, roleID, actions) - - if len(ret) == 0 { - panic("no return value specified for RoleCheckActionsExists") - } - - var r0 bool - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) (bool, error)); ok { - return returnFunc(ctx, session, entityID, roleID, actions) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) bool); ok { - r0 = returnFunc(ctx, session, entityID, roleID, actions) - } else { - r0 = ret.Get(0).(bool) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, []string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID, actions) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RoleCheckActionsExists_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleCheckActionsExists' -type Service_RoleCheckActionsExists_Call struct { - *mock.Call -} - -// RoleCheckActionsExists is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - actions []string -func (_e *Service_Expecter) RoleCheckActionsExists(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, actions interface{}) *Service_RoleCheckActionsExists_Call { - return &Service_RoleCheckActionsExists_Call{Call: _e.mock.On("RoleCheckActionsExists", ctx, session, entityID, roleID, actions)} -} - -func (_c *Service_RoleCheckActionsExists_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string)) *Service_RoleCheckActionsExists_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_RoleCheckActionsExists_Call) Return(b bool, err error) *Service_RoleCheckActionsExists_Call { - _c.Call.Return(b, err) - return _c -} - -func (_c *Service_RoleCheckActionsExists_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string) (bool, error)) *Service_RoleCheckActionsExists_Call { - _c.Call.Return(run) - return _c -} - -// RoleCheckMembersExists provides a mock function for the type Service -func (_mock *Service) RoleCheckMembersExists(ctx context.Context, session authn.Session, entityID string, roleID string, members []string) (bool, error) { - ret := _mock.Called(ctx, session, entityID, roleID, members) - - if len(ret) == 0 { - panic("no return value specified for RoleCheckMembersExists") - } - - var r0 bool - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) (bool, error)); ok { - return returnFunc(ctx, session, entityID, roleID, members) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) bool); ok { - r0 = returnFunc(ctx, session, entityID, roleID, members) - } else { - r0 = ret.Get(0).(bool) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, []string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID, members) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RoleCheckMembersExists_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleCheckMembersExists' -type Service_RoleCheckMembersExists_Call struct { - *mock.Call -} - -// RoleCheckMembersExists is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - members []string -func (_e *Service_Expecter) RoleCheckMembersExists(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, members interface{}) *Service_RoleCheckMembersExists_Call { - return &Service_RoleCheckMembersExists_Call{Call: _e.mock.On("RoleCheckMembersExists", ctx, session, entityID, roleID, members)} -} - -func (_c *Service_RoleCheckMembersExists_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, members []string)) *Service_RoleCheckMembersExists_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_RoleCheckMembersExists_Call) Return(b bool, err error) *Service_RoleCheckMembersExists_Call { - _c.Call.Return(b, err) - return _c -} - -func (_c *Service_RoleCheckMembersExists_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, members []string) (bool, error)) *Service_RoleCheckMembersExists_Call { - _c.Call.Return(run) - return _c -} - -// RoleListActions provides a mock function for the type Service -func (_mock *Service) RoleListActions(ctx context.Context, session authn.Session, entityID string, roleID string) ([]string, error) { - ret := _mock.Called(ctx, session, entityID, roleID) - - if len(ret) == 0 { - panic("no return value specified for RoleListActions") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) ([]string, error)); ok { - return returnFunc(ctx, session, entityID, roleID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) []string); ok { - r0 = returnFunc(ctx, session, entityID, roleID) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RoleListActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleListActions' -type Service_RoleListActions_Call struct { - *mock.Call -} - -// RoleListActions is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -func (_e *Service_Expecter) RoleListActions(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}) *Service_RoleListActions_Call { - return &Service_RoleListActions_Call{Call: _e.mock.On("RoleListActions", ctx, session, entityID, roleID)} -} - -func (_c *Service_RoleListActions_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string)) *Service_RoleListActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_RoleListActions_Call) Return(strings []string, err error) *Service_RoleListActions_Call { - _c.Call.Return(strings, err) - return _c -} - -func (_c *Service_RoleListActions_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string) ([]string, error)) *Service_RoleListActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleListMembers provides a mock function for the type Service -func (_mock *Service) RoleListMembers(ctx context.Context, session authn.Session, entityID string, roleID string, limit uint64, offset uint64) (roles.MembersPage, error) { - ret := _mock.Called(ctx, session, entityID, roleID, limit, offset) - - if len(ret) == 0 { - panic("no return value specified for RoleListMembers") - } - - var r0 roles.MembersPage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, uint64, uint64) (roles.MembersPage, error)); ok { - return returnFunc(ctx, session, entityID, roleID, limit, offset) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, uint64, uint64) roles.MembersPage); ok { - r0 = returnFunc(ctx, session, entityID, roleID, limit, offset) - } else { - r0 = ret.Get(0).(roles.MembersPage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, uint64, uint64) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID, limit, offset) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RoleListMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleListMembers' -type Service_RoleListMembers_Call struct { - *mock.Call -} - -// RoleListMembers is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - limit uint64 -// - offset uint64 -func (_e *Service_Expecter) RoleListMembers(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, limit interface{}, offset interface{}) *Service_RoleListMembers_Call { - return &Service_RoleListMembers_Call{Call: _e.mock.On("RoleListMembers", ctx, session, entityID, roleID, limit, offset)} -} - -func (_c *Service_RoleListMembers_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, limit uint64, offset uint64)) *Service_RoleListMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 uint64 - if args[4] != nil { - arg4 = args[4].(uint64) - } - var arg5 uint64 - if args[5] != nil { - arg5 = args[5].(uint64) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - arg5, - ) - }) - return _c -} - -func (_c *Service_RoleListMembers_Call) Return(membersPage roles.MembersPage, err error) *Service_RoleListMembers_Call { - _c.Call.Return(membersPage, err) - return _c -} - -func (_c *Service_RoleListMembers_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, limit uint64, offset uint64) (roles.MembersPage, error)) *Service_RoleListMembers_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveActions provides a mock function for the type Service -func (_mock *Service) RoleRemoveActions(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string) error { - ret := _mock.Called(ctx, session, entityID, roleID, actions) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveActions") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) error); ok { - r0 = returnFunc(ctx, session, entityID, roleID, actions) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RoleRemoveActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveActions' -type Service_RoleRemoveActions_Call struct { - *mock.Call -} - -// RoleRemoveActions is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - actions []string -func (_e *Service_Expecter) RoleRemoveActions(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, actions interface{}) *Service_RoleRemoveActions_Call { - return &Service_RoleRemoveActions_Call{Call: _e.mock.On("RoleRemoveActions", ctx, session, entityID, roleID, actions)} -} - -func (_c *Service_RoleRemoveActions_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string)) *Service_RoleRemoveActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_RoleRemoveActions_Call) Return(err error) *Service_RoleRemoveActions_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RoleRemoveActions_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string) error) *Service_RoleRemoveActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveAllActions provides a mock function for the type Service -func (_mock *Service) RoleRemoveAllActions(ctx context.Context, session authn.Session, entityID string, roleID string) error { - ret := _mock.Called(ctx, session, entityID, roleID) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveAllActions") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) error); ok { - r0 = returnFunc(ctx, session, entityID, roleID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RoleRemoveAllActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveAllActions' -type Service_RoleRemoveAllActions_Call struct { - *mock.Call -} - -// RoleRemoveAllActions is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -func (_e *Service_Expecter) RoleRemoveAllActions(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}) *Service_RoleRemoveAllActions_Call { - return &Service_RoleRemoveAllActions_Call{Call: _e.mock.On("RoleRemoveAllActions", ctx, session, entityID, roleID)} -} - -func (_c *Service_RoleRemoveAllActions_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string)) *Service_RoleRemoveAllActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_RoleRemoveAllActions_Call) Return(err error) *Service_RoleRemoveAllActions_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RoleRemoveAllActions_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string) error) *Service_RoleRemoveAllActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveAllMembers provides a mock function for the type Service -func (_mock *Service) RoleRemoveAllMembers(ctx context.Context, session authn.Session, entityID string, roleID string) error { - ret := _mock.Called(ctx, session, entityID, roleID) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveAllMembers") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) error); ok { - r0 = returnFunc(ctx, session, entityID, roleID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RoleRemoveAllMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveAllMembers' -type Service_RoleRemoveAllMembers_Call struct { - *mock.Call -} - -// RoleRemoveAllMembers is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -func (_e *Service_Expecter) RoleRemoveAllMembers(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}) *Service_RoleRemoveAllMembers_Call { - return &Service_RoleRemoveAllMembers_Call{Call: _e.mock.On("RoleRemoveAllMembers", ctx, session, entityID, roleID)} -} - -func (_c *Service_RoleRemoveAllMembers_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string)) *Service_RoleRemoveAllMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_RoleRemoveAllMembers_Call) Return(err error) *Service_RoleRemoveAllMembers_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RoleRemoveAllMembers_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string) error) *Service_RoleRemoveAllMembers_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveMembers provides a mock function for the type Service -func (_mock *Service) RoleRemoveMembers(ctx context.Context, session authn.Session, entityID string, roleID string, members []string) error { - ret := _mock.Called(ctx, session, entityID, roleID, members) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveMembers") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) error); ok { - r0 = returnFunc(ctx, session, entityID, roleID, members) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RoleRemoveMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveMembers' -type Service_RoleRemoveMembers_Call struct { - *mock.Call -} - -// RoleRemoveMembers is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - members []string -func (_e *Service_Expecter) RoleRemoveMembers(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, members interface{}) *Service_RoleRemoveMembers_Call { - return &Service_RoleRemoveMembers_Call{Call: _e.mock.On("RoleRemoveMembers", ctx, session, entityID, roleID, members)} -} - -func (_c *Service_RoleRemoveMembers_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, members []string)) *Service_RoleRemoveMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_RoleRemoveMembers_Call) Return(err error) *Service_RoleRemoveMembers_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RoleRemoveMembers_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, members []string) error) *Service_RoleRemoveMembers_Call { - _c.Call.Return(run) - return _c -} - -// SetParentGroup provides a mock function for the type Service -func (_mock *Service) SetParentGroup(ctx context.Context, session authn.Session, parentGroupID string, id string) error { - ret := _mock.Called(ctx, session, parentGroupID, id) - - if len(ret) == 0 { - panic("no return value specified for SetParentGroup") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) error); ok { - r0 = returnFunc(ctx, session, parentGroupID, id) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_SetParentGroup_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'SetParentGroup' -type Service_SetParentGroup_Call struct { - *mock.Call -} - -// SetParentGroup is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - parentGroupID string -// - id string -func (_e *Service_Expecter) SetParentGroup(ctx interface{}, session interface{}, parentGroupID interface{}, id interface{}) *Service_SetParentGroup_Call { - return &Service_SetParentGroup_Call{Call: _e.mock.On("SetParentGroup", ctx, session, parentGroupID, id)} -} - -func (_c *Service_SetParentGroup_Call) Run(run func(ctx context.Context, session authn.Session, parentGroupID string, id string)) *Service_SetParentGroup_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_SetParentGroup_Call) Return(err error) *Service_SetParentGroup_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_SetParentGroup_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, parentGroupID string, id string) error) *Service_SetParentGroup_Call { - _c.Call.Return(run) - return _c -} - -// UpdateChannel provides a mock function for the type Service -func (_mock *Service) UpdateChannel(ctx context.Context, session authn.Session, channel channels.Channel) (channels.Channel, error) { - ret := _mock.Called(ctx, session, channel) - - if len(ret) == 0 { - panic("no return value specified for UpdateChannel") - } - - var r0 channels.Channel - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, channels.Channel) (channels.Channel, error)); ok { - return returnFunc(ctx, session, channel) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, channels.Channel) channels.Channel); ok { - r0 = returnFunc(ctx, session, channel) - } else { - r0 = ret.Get(0).(channels.Channel) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, channels.Channel) error); ok { - r1 = returnFunc(ctx, session, channel) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_UpdateChannel_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateChannel' -type Service_UpdateChannel_Call struct { - *mock.Call -} - -// UpdateChannel is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - channel channels.Channel -func (_e *Service_Expecter) UpdateChannel(ctx interface{}, session interface{}, channel interface{}) *Service_UpdateChannel_Call { - return &Service_UpdateChannel_Call{Call: _e.mock.On("UpdateChannel", ctx, session, channel)} -} - -func (_c *Service_UpdateChannel_Call) Run(run func(ctx context.Context, session authn.Session, channel channels.Channel)) *Service_UpdateChannel_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 channels.Channel - if args[2] != nil { - arg2 = args[2].(channels.Channel) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_UpdateChannel_Call) Return(channel1 channels.Channel, err error) *Service_UpdateChannel_Call { - _c.Call.Return(channel1, err) - return _c -} - -func (_c *Service_UpdateChannel_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, channel channels.Channel) (channels.Channel, error)) *Service_UpdateChannel_Call { - _c.Call.Return(run) - return _c -} - -// UpdateChannelTags provides a mock function for the type Service -func (_mock *Service) UpdateChannelTags(ctx context.Context, session authn.Session, channel channels.Channel) (channels.Channel, error) { - ret := _mock.Called(ctx, session, channel) - - if len(ret) == 0 { - panic("no return value specified for UpdateChannelTags") - } - - var r0 channels.Channel - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, channels.Channel) (channels.Channel, error)); ok { - return returnFunc(ctx, session, channel) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, channels.Channel) channels.Channel); ok { - r0 = returnFunc(ctx, session, channel) - } else { - r0 = ret.Get(0).(channels.Channel) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, channels.Channel) error); ok { - r1 = returnFunc(ctx, session, channel) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_UpdateChannelTags_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateChannelTags' -type Service_UpdateChannelTags_Call struct { - *mock.Call -} - -// UpdateChannelTags is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - channel channels.Channel -func (_e *Service_Expecter) UpdateChannelTags(ctx interface{}, session interface{}, channel interface{}) *Service_UpdateChannelTags_Call { - return &Service_UpdateChannelTags_Call{Call: _e.mock.On("UpdateChannelTags", ctx, session, channel)} -} - -func (_c *Service_UpdateChannelTags_Call) Run(run func(ctx context.Context, session authn.Session, channel channels.Channel)) *Service_UpdateChannelTags_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 channels.Channel - if args[2] != nil { - arg2 = args[2].(channels.Channel) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_UpdateChannelTags_Call) Return(channel1 channels.Channel, err error) *Service_UpdateChannelTags_Call { - _c.Call.Return(channel1, err) - return _c -} - -func (_c *Service_UpdateChannelTags_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, channel channels.Channel) (channels.Channel, error)) *Service_UpdateChannelTags_Call { - _c.Call.Return(run) - return _c -} - -// UpdateRoleName provides a mock function for the type Service -func (_mock *Service) UpdateRoleName(ctx context.Context, session authn.Session, entityID string, roleID string, newRoleName string) (roles.Role, error) { - ret := _mock.Called(ctx, session, entityID, roleID, newRoleName) - - if len(ret) == 0 { - panic("no return value specified for UpdateRoleName") - } - - var r0 roles.Role - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, string) (roles.Role, error)); ok { - return returnFunc(ctx, session, entityID, roleID, newRoleName) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, string) roles.Role); ok { - r0 = returnFunc(ctx, session, entityID, roleID, newRoleName) - } else { - r0 = ret.Get(0).(roles.Role) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID, newRoleName) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_UpdateRoleName_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateRoleName' -type Service_UpdateRoleName_Call struct { - *mock.Call -} - -// UpdateRoleName is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - newRoleName string -func (_e *Service_Expecter) UpdateRoleName(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, newRoleName interface{}) *Service_UpdateRoleName_Call { - return &Service_UpdateRoleName_Call{Call: _e.mock.On("UpdateRoleName", ctx, session, entityID, roleID, newRoleName)} -} - -func (_c *Service_UpdateRoleName_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, newRoleName string)) *Service_UpdateRoleName_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 string - if args[4] != nil { - arg4 = args[4].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_UpdateRoleName_Call) Return(role roles.Role, err error) *Service_UpdateRoleName_Call { - _c.Call.Return(role, err) - return _c -} - -func (_c *Service_UpdateRoleName_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, newRoleName string) (roles.Role, error)) *Service_UpdateRoleName_Call { - _c.Call.Return(run) - return _c -} - -// ViewChannel provides a mock function for the type Service -func (_mock *Service) ViewChannel(ctx context.Context, session authn.Session, id string, withRoles bool) (channels.Channel, error) { - ret := _mock.Called(ctx, session, id, withRoles) - - if len(ret) == 0 { - panic("no return value specified for ViewChannel") - } - - var r0 channels.Channel - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, bool) (channels.Channel, error)); ok { - return returnFunc(ctx, session, id, withRoles) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, bool) channels.Channel); ok { - r0 = returnFunc(ctx, session, id, withRoles) - } else { - r0 = ret.Get(0).(channels.Channel) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, bool) error); ok { - r1 = returnFunc(ctx, session, id, withRoles) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_ViewChannel_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ViewChannel' -type Service_ViewChannel_Call struct { - *mock.Call -} - -// ViewChannel is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - id string -// - withRoles bool -func (_e *Service_Expecter) ViewChannel(ctx interface{}, session interface{}, id interface{}, withRoles interface{}) *Service_ViewChannel_Call { - return &Service_ViewChannel_Call{Call: _e.mock.On("ViewChannel", ctx, session, id, withRoles)} -} - -func (_c *Service_ViewChannel_Call) Run(run func(ctx context.Context, session authn.Session, id string, withRoles bool)) *Service_ViewChannel_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 bool - if args[3] != nil { - arg3 = args[3].(bool) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_ViewChannel_Call) Return(channel channels.Channel, err error) *Service_ViewChannel_Call { - _c.Call.Return(channel, err) - return _c -} - -func (_c *Service_ViewChannel_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, id string, withRoles bool) (channels.Channel, error)) *Service_ViewChannel_Call { - _c.Call.Return(run) - return _c -} diff --git a/channels/operations/operations.go b/channels/operations/operations.go deleted file mode 100644 index a7a675870..000000000 --- a/channels/operations/operations.go +++ /dev/null @@ -1,72 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package operations - -import ( - "github.com/absmach/magistrala/pkg/permissions" -) - -// Channel Operations. -const ( - OpViewChannel permissions.Operation = iota - OpUpdateChannel - OpUpdateChannelTags - OpEnableChannel - OpDisableChannel - OpDeleteChannel - OpSetParentGroup - OpRemoveParentGroup - OpConnectClient - OpDisconnectClient - OpListUserChannels -) - -func OperationDetails() map[permissions.Operation]permissions.OperationDetails { - return map[permissions.Operation]permissions.OperationDetails{ - OpViewChannel: { - Name: "view", - PermissionRequired: true, - }, - OpUpdateChannel: { - Name: "update", - PermissionRequired: true, - }, - OpUpdateChannelTags: { - Name: "update_tags", - PermissionRequired: true, - }, - OpEnableChannel: { - Name: "enable", - PermissionRequired: true, - }, - OpDisableChannel: { - Name: "disable", - PermissionRequired: true, - }, - OpDeleteChannel: { - Name: "delete", - PermissionRequired: true, - }, - OpSetParentGroup: { - Name: "set_parent_group", - PermissionRequired: true, - }, - OpRemoveParentGroup: { - Name: "remove_parent_group", - PermissionRequired: true, - }, - OpConnectClient: { - Name: "connect_client", - PermissionRequired: true, - }, - OpDisconnectClient: { - Name: "disconnect_client", - PermissionRequired: true, - }, - OpListUserChannels: { - Name: "list_user_channels", - PermissionRequired: false, // hardcoded to superadmin - }, - } -} diff --git a/channels/postgres/channels.go b/channels/postgres/channels.go deleted file mode 100644 index 887ba4a0e..000000000 --- a/channels/postgres/channels.go +++ /dev/null @@ -1,1476 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package postgres - -import ( - "context" - "database/sql" - "encoding/json" - "fmt" - "strings" - "time" - - api "github.com/absmach/magistrala/api/http" - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/channels" - "github.com/absmach/magistrala/pkg/connections" - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - "github.com/absmach/magistrala/pkg/policies" - "github.com/absmach/magistrala/pkg/postgres" - "github.com/absmach/magistrala/pkg/roles" - rolesPostgres "github.com/absmach/magistrala/pkg/roles/repo/postgres" - "github.com/jackc/pgtype" - "github.com/lib/pq" -) - -const ( - rolesTableNamePrefix = "channels" - entityTableName = "channels" - entityIDColumnName = "id" -) - -var _ channels.Repository = (*channelRepository)(nil) - -type channelRepository struct { - db postgres.Database - eh errors.Handler - rolesPostgres.Repository -} - -// NewChannelRepository instantiates a PostgreSQL implementation of channel -// repository. -func NewRepository(db postgres.Database) channels.Repository { - rolesRepo := rolesPostgres.NewRepository(db, policies.ChannelType, rolesTableNamePrefix, entityTableName, entityIDColumnName) - errHandlerOptions := []errors.HandlerOption{ - postgres.WithDuplicateErrors(NewDuplicateErrors()), - } - return &channelRepository{ - db: db, - eh: postgres.NewErrorHandler(errHandlerOptions...), - Repository: rolesRepo, - } -} - -func (cr *channelRepository) Save(ctx context.Context, chs ...channels.Channel) ([]channels.Channel, error) { - var dbchs []dbChannel - for _, ch := range chs { - dbch, err := toDBChannel(ch) - if err != nil { - return []channels.Channel{}, errors.Wrap(repoerr.ErrCreateEntity, err) - } - dbchs = append(dbchs, dbch) - } - - q := `INSERT INTO channels (id, name, tags, domain_id, parent_group_id, route, metadata, created_at, updated_at, updated_by, status) - VALUES (:id, :name, :tags, :domain_id, :parent_group_id, :route, :metadata, :created_at, :updated_at, :updated_by, :status) - RETURNING id, name, tags, metadata, COALESCE(domain_id, '') AS domain_id, COALESCE(parent_group_id, '') AS parent_group_id, route, status, created_at, updated_at, updated_by` - - row, err := cr.db.NamedQueryContext(ctx, q, dbchs) - if err != nil { - return []channels.Channel{}, cr.eh.HandleError(repoerr.ErrCreateEntity, err) - } - - defer row.Close() - - var reChs []channels.Channel - - for row.Next() { - dbch := dbChannel{} - if err := row.StructScan(&dbch); err != nil { - return []channels.Channel{}, cr.eh.HandleError(repoerr.ErrFailedOpDB, err) - } - - ch, err := toChannel(dbch) - if err != nil { - return []channels.Channel{}, errors.Wrap(repoerr.ErrFailedOpDB, err) - } - reChs = append(reChs, ch) - } - return reChs, nil -} - -func (cr *channelRepository) Update(ctx context.Context, channel channels.Channel) (channels.Channel, error) { - var query []string - var upq string - if channel.Name != "" { - query = append(query, "name = :name,") - } - if channel.Metadata != nil { - query = append(query, "metadata = :metadata,") - } - if len(query) > 0 { - upq = strings.Join(query, " ") - } - q := fmt.Sprintf(`UPDATE channels SET %s updated_at = :updated_at, updated_by = :updated_by - WHERE id = :id AND status = :status - RETURNING id, name, tags, metadata, COALESCE(domain_id, '') AS domain_id, COALESCE(parent_group_id, '') AS parent_group_id, route, status, created_at, updated_at, updated_by`, - upq) - channel.Status = channels.EnabledStatus - return cr.update(ctx, channel, q) -} - -func (cr *channelRepository) UpdateTags(ctx context.Context, channel channels.Channel) (channels.Channel, error) { - q := `UPDATE channels SET tags = :tags, updated_at = :updated_at, updated_by = :updated_by - WHERE id = :id AND status = :status - RETURNING id, name, tags, metadata, COALESCE(domain_id, '') AS domain_id, COALESCE(parent_group_id, '') AS parent_group_id, route, status, created_at, updated_at, updated_by` - channel.Status = channels.EnabledStatus - return cr.update(ctx, channel, q) -} - -func (cr *channelRepository) ChangeStatus(ctx context.Context, channel channels.Channel) (channels.Channel, error) { - q := `UPDATE channels SET status = :status, updated_at = :updated_at, updated_by = :updated_by - WHERE id = :id - RETURNING id, name, tags, metadata, COALESCE(domain_id, '') AS domain_id, COALESCE(parent_group_id, '') AS parent_group_id, route, status, created_at, updated_at, updated_by` - - return cr.update(ctx, channel, q) -} - -func (cr *channelRepository) RetrieveByID(ctx context.Context, id string) (channels.Channel, error) { - q := `SELECT id, name, tags, COALESCE(domain_id, '') AS domain_id, COALESCE(parent_group_id, '') AS parent_group_id, route, metadata, created_at, updated_at, updated_by, status FROM channels WHERE id = :id` - - dbch := dbChannel{ - ID: id, - } - - row, err := cr.db.NamedQueryContext(ctx, q, dbch) - if err != nil { - return channels.Channel{}, cr.eh.HandleError(repoerr.ErrViewEntity, err) - } - defer row.Close() - - dbch = dbChannel{} - if row.Next() { - if err := row.StructScan(&dbch); err != nil { - return channels.Channel{}, cr.eh.HandleError(repoerr.ErrViewEntity, err) - } - return toChannel(dbch) - } - - return channels.Channel{}, repoerr.ErrNotFound -} - -func (cr *channelRepository) RetrieveByRoute(ctx context.Context, route, domainID string) (channels.Channel, error) { - q := `SELECT id, name, tags, COALESCE(domain_id, '') AS domain_id, COALESCE(parent_group_id, '') AS parent_group_id, route, metadata, created_at, updated_at, updated_by, status - FROM channels WHERE route = :route AND domain_id = :domain_id` - - dbch := dbChannel{ - Route: toNullString(route), - Domain: domainID, - } - - row, err := cr.db.NamedQueryContext(ctx, q, dbch) - if err != nil { - return channels.Channel{}, cr.eh.HandleError(repoerr.ErrViewEntity, err) - } - defer row.Close() - - dbch = dbChannel{} - if row.Next() { - if err := row.StructScan(&dbch); err != nil { - return channels.Channel{}, cr.eh.HandleError(repoerr.ErrViewEntity, err) - } - return toChannel(dbch) - } - - return channels.Channel{}, repoerr.ErrNotFound -} - -func (cr *channelRepository) RetrieveByIDWithRoles(ctx context.Context, id, memberID string) (channels.Channel, error) { - query := ` - WITH selected_channel AS ( - SELECT - c.id, - c.parent_group_id, - COALESCE(g."path", CAST('' AS ltree)) AS parent_group_path, - c.domain_id - FROM - channels c - LEFT JOIN - "groups" g ON c.parent_group_id = g.id - WHERE - c.id = :id - LIMIT 1 - ), - selected_channel_roles AS ( - SELECT - cr.entity_id AS channel_id, - crm.member_id AS member_id, - cr.id AS role_id, - cr."name" AS role_name, - jsonb_agg(DISTINCT cra."action") AS actions, - 'direct' AS access_type, - CAST('' AS ltree) AS access_provider_path, - '' AS access_provider_id - FROM - channels_roles cr - JOIN - channels_role_members crm ON cr.id = crm.role_id - JOIN - channels_role_actions cra ON cr.id = cra.role_id - JOIN - selected_channel sc ON sc.id = cr.entity_id - AND crm.member_id = :member_id - GROUP BY - cr.entity_id, cr.id, cr.name, crm.member_id - ), - selected_group_roles AS ( - SELECT - sc.id AS channel_id, - grm.member_id AS member_id, - gr.id AS role_id, - gr."name" AS role_name, - jsonb_agg(DISTINCT all_actions."action") AS actions, - gr.entity_id AS access_provider_id, - g."path" AS access_provider_path, - CASE - WHEN gr.entity_id = sc.parent_group_id - THEN 'direct_group' - ELSE 'indirect_group' - END AS access_type - FROM - "groups" g - JOIN - groups_roles gr ON gr.entity_id = g.id - JOIN - groups_role_members grm ON gr.id = grm.role_id - JOIN - groups_role_actions gra ON gr.id = gra.role_id - JOIN - groups_role_actions all_actions ON gr.id = all_actions.role_id - JOIN - selected_channel sc ON TRUE - WHERE - g."path" @> sc.parent_group_path - AND grm.member_id = :member_id - AND ( - (g.id = sc.parent_group_id AND gra."action" LIKE 'channel%%') - OR - (g.id <> sc.parent_group_id AND gra."action" LIKE 'subgroup_channel%%') - ) - GROUP BY - sc.id, sc.parent_group_id, gr.entity_id, gr.id, gr."name", g."path", grm.member_id - ), - selected_domain_roles AS ( - SELECT - sc.id AS channel_id, - drm.member_id AS member_id, - dr.entity_id AS group_id, - dr.id AS role_id, - dr."name" AS role_name, - jsonb_agg(DISTINCT all_actions."action") AS actions, - CAST('' AS ltree) access_provider_path, - 'domain' AS access_type, - dr.entity_id AS access_provider_id - FROM - domains d - JOIN - selected_channel sc ON sc.domain_id = d.id - JOIN - domains_roles dr ON dr.entity_id = d.id - JOIN - domains_role_members drm ON dr.id = drm.role_id - JOIN - domains_role_actions dra ON dr.id = dra.role_id - JOIN - domains_role_actions all_actions ON dr.id = all_actions.role_id - WHERE - drm.member_id = :member_id - AND dra."action" LIKE 'channel%%' - GROUP BY - sc.id, dr.entity_id, dr.id, dr."name", drm.member_id - ), - all_roles AS ( - SELECT - scr.channel_id, - scr.member_id, - scr.role_id AS role_id, - scr.role_name AS role_name, - scr.actions AS actions, - scr.access_type AS access_type, - scr.access_provider_path AS access_provider_path, - scr.access_provider_id AS access_provider_id - FROM - selected_channel_roles scr - UNION - SELECT - sgr.channel_id, - sgr.member_id, - sgr.role_id AS role_id, - sgr.role_name AS role_name, - sgr.actions AS actions, - sgr.access_type AS access_type, - sgr.access_provider_path AS access_provider_path, - sgr.access_provider_id AS access_provider_id - FROM - selected_group_roles sgr - UNION - SELECT - sdr.channel_id, - sdr.member_id, - sdr.role_id AS role_id, - sdr.role_name AS role_name, - sdr.actions AS actions, - sdr.access_type AS access_type, - sdr.access_provider_path AS access_provider_path, - sdr.access_provider_id AS access_provider_id - FROM - selected_domain_roles sdr - ), - final_roles AS ( - SELECT - ar.channel_id, - ar.member_id, - jsonb_agg( - jsonb_build_object( - 'role_id', ar.role_id, - 'role_name', ar.role_name, - 'actions', ar.actions, - 'access_type', ar.access_type, - 'access_provider_path', ar.access_provider_path, - 'access_provider_id', ar.access_provider_id - ) - ) AS roles - FROM all_roles ar - GROUP BY - ar.channel_id, ar.member_id - ) - SELECT - c2.id, - c2."name", - c2.tags, - COALESCE(c2.domain_id, '') AS domain_id, - COALESCE(c2.parent_group_id, '') AS parent_group_id, - c2.route, - c2.metadata, - c2.created_at, - c2.created_by, - c2.updated_at, - c2.updated_by, - c2.status, - fr.member_id, - fr.roles - FROM channels c2 - JOIN final_roles fr ON fr.channel_id = c2.id - ` - parameters := map[string]any{ - "id": id, - "member_id": memberID, - } - row, err := cr.db.NamedQueryContext(ctx, query, parameters) - if err != nil { - return channels.Channel{}, cr.eh.HandleError(repoerr.ErrViewEntity, err) - } - defer row.Close() - - dbch := dbChannel{} - if !row.Next() { - return channels.Channel{}, repoerr.ErrNotFound - } - - if err := row.StructScan(&dbch); err != nil { - return channels.Channel{}, cr.eh.HandleError(repoerr.ErrViewEntity, err) - } - - return toChannel(dbch) -} - -func (cr *channelRepository) RetrieveAll(ctx context.Context, pm channels.Page) (channels.ChannelsPage, error) { - pageQuery, err := PageQuery(pm) - if err != nil { - return channels.ChannelsPage{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - - connJoinQuery := ` - FROM - channels c - ` - - if pm.Client != "" { - connJoinQuery = ` - ,conn.connection_types - FROM - channels c - LEFT JOIN ( - SELECT - conn.client_id, - conn.channel_id, - array_agg(conn."type") AS connection_types - FROM - connections AS conn - GROUP BY - conn.client_id, conn.channel_id - ) conn ON c.id = conn.channel_id - ` - } - - comQuery := fmt.Sprintf(`WITH channels AS ( - SELECT - c.id, - c.name, - c.tags, - c.metadata, - COALESCE(c.domain_id, '') AS domain_id, - COALESCE(parent_group_id, '') AS parent_group_id, - c.route, - COALESCE(g.path, CAST('' AS ltree)) AS parent_group_path, - c.status, - c.created_by, - c.created_at, - c.updated_at, - COALESCE(c.updated_by, '') AS updated_by - FROM - channels c - LEFT JOIN - groups g ON g.id = c.parent_group_id - ) - SELECT - c.* - %s - %s - `, connJoinQuery, pageQuery) - - q := applyOrdering(comQuery, pm) - - q = applyLimitOffset(q) - - dbPage, err := toDBChannelsPage(pm) - if err != nil { - return channels.ChannelsPage{}, errors.Wrap(repoerr.ErrFailedToRetrieveAllGroups, err) - } - - var items []channels.Channel - if !pm.OnlyTotal { - rows, err := cr.db.NamedQueryContext(ctx, q, dbPage) - if err != nil { - return channels.ChannelsPage{}, cr.eh.HandleError(repoerr.ErrFailedToRetrieveAllGroups, err) - } - defer rows.Close() - - for rows.Next() { - dbch := dbChannel{} - if err := rows.StructScan(&dbch); err != nil { - return channels.ChannelsPage{}, cr.eh.HandleError(repoerr.ErrViewEntity, err) - } - - ch, err := toChannel(dbch) - if err != nil { - return channels.ChannelsPage{}, err - } - - items = append(items, ch) - } - } - cq := fmt.Sprintf(`SELECT COUNT(*) AS total_count - FROM ( - %s - ) AS sub_query; - `, comQuery) - - total, err := postgres.Total(ctx, cr.db, cq, dbPage) - if err != nil { - return channels.ChannelsPage{}, cr.eh.HandleError(repoerr.ErrViewEntity, err) - } - - page := channels.ChannelsPage{ - Channels: items, - Page: channels.Page{ - Total: total, - Offset: pm.Offset, - Limit: pm.Limit, - }, - } - return page, nil -} - -func (repo *channelRepository) RetrieveUserChannels(ctx context.Context, domainID, userID string, pm channels.Page) (channels.ChannelsPage, error) { - return repo.retrieveChannels(ctx, domainID, userID, pm) -} - -func (repo *channelRepository) retrieveChannels(ctx context.Context, domainID, userID string, pm channels.Page) (channels.ChannelsPage, error) { - pageQuery, err := PageQuery(pm) - if err != nil { - return channels.ChannelsPage{}, err - } - - bq := userChannelsBaseQuery - - connJoinQuery := ` - FROM - final_channels c - ` - connCountJoinQuery := connJoinQuery - - if pm.Client != "" { - connCountJoinQuery = ` - FROM - final_channels c - LEFT JOIN ( - SELECT - conn.client_id, - conn.channel_id, - array_agg(conn."type") AS connection_types - FROM - connections AS conn - GROUP BY - conn.client_id, conn.channel_id - ) conn ON c.id = conn.channel_id - ` - connJoinQuery = ` - ,conn.connection_types - ` + connCountJoinQuery - } - - dbPage, err := toDBChannelsPage(pm) - if err != nil { - return channels.ChannelsPage{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - dbPage.UserID = userID - dbPage.DomainID = domainID - - if pm.OnlyTotal { - cq := fmt.Sprintf(`%s - SELECT COUNT(*) AS total_count - %s - %s; - `, bq, connCountJoinQuery, pageQuery) - - total, err := postgres.Total(ctx, repo.db, cq, dbPage) - if err != nil { - return channels.ChannelsPage{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - - return channels.ChannelsPage{ - Page: channels.Page{ - Total: total, - Offset: pm.Offset, - Limit: pm.Limit, - }, - }, nil - } - - q := fmt.Sprintf(` - %s - SELECT - c.id, - c.name, - c.domain_id, - c.parent_group_id, - c.route, - c.tags, - c.metadata, - c.created_by, - c.created_at, - c.updated_at, - c.updated_by, - c.status, - c.parent_group_path, - c.role_id, - c.role_name, - c.actions, - c.access_type, - c.access_provider_id, - c.access_provider_role_id, - c.access_provider_role_name, - c.access_provider_role_actions, - COUNT(*) OVER() AS total_count - %s - %s - `, bq, connJoinQuery, pageQuery) - - q = applyOrdering(q, pm) - - q = applyLimitOffset(q) - - rows, err := repo.db.NamedQueryContext(ctx, q, dbPage) - if err != nil { - return channels.ChannelsPage{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - defer rows.Close() - - var total uint64 - var items []channels.Channel - for rows.Next() { - dbc := dbChannel{} - if err := rows.StructScan(&dbc); err != nil { - return channels.ChannelsPage{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - - total = dbc.TotalCount - - c, err := toChannel(dbc) - if err != nil { - return channels.ChannelsPage{}, err - } - - items = append(items, c) - } - - if len(items) == 0 { - cq := fmt.Sprintf(`%s - SELECT COUNT(*) AS total_count - %s - %s; - `, bq, connCountJoinQuery, pageQuery) - - total, err = postgres.Total(ctx, repo.db, cq, dbPage) - if err != nil { - return channels.ChannelsPage{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - } - - return channels.ChannelsPage{ - Channels: items, - Page: channels.Page{ - Total: total, - Offset: pm.Offset, - Limit: pm.Limit, - }, - }, nil -} - -const userChannelsBaseQuery = ` -WITH direct_channels AS ( - select - c.id, - c.name, - c.domain_id, - c.parent_group_id, - c.route, - c.tags, - c.metadata, - c.created_by, - c.created_at, - c.updated_at, - c.updated_by, - c.status, - COALESCE(pg.path, CAST('' AS ltree)) AS parent_group_path, - cr.id AS role_id, - cr."name" AS role_name, - array_agg(cra."action") AS actions, - 'direct' as access_type, - '' AS access_provider_id, - '' AS access_provider_role_id, - '' AS access_provider_role_name, - CAST(array[] AS text[]) AS access_provider_role_actions - FROM - channels_role_members crm - JOIN - channels_role_actions cra ON cra.role_id = crm.role_id - JOIN - channels_roles cr ON cr.id = crm.role_id - JOIN - channels c ON c.id = cr.entity_id - LEFT JOIN - groups pg ON pg.id = c.parent_group_id - WHERE - crm.member_id = :user_id - AND c.domain_id = :domain_id_param - GROUP BY - cr.entity_id, crm.member_id, cr.id, cr."name", c.id, pg.path -), -direct_groups AS ( - SELECT - g.*, - gr.entity_id AS entity_id, - grm.member_id AS member_id, - gr.id AS role_id, - gr."name" AS role_name, - array_agg(DISTINCT all_actions."action") AS actions - FROM - groups_role_members grm - JOIN - groups_role_actions gra ON gra.role_id = grm.role_id - JOIN - groups_roles gr ON gr.id = grm.role_id - JOIN - "groups" g ON g.id = gr.entity_id - JOIN - groups_role_actions all_actions ON all_actions.role_id = grm.role_id - WHERE - grm.member_id = :user_id - AND g.domain_id = :domain_id_param - AND gra."action" LIKE 'channel%' - GROUP BY - gr.entity_id, grm.member_id, gr.id, gr."name", g."path", g.id -), -direct_groups_with_subgroup AS ( - SELECT - g.*, - gr.entity_id AS entity_id, - grm.member_id AS member_id, - gr.id AS role_id, - gr."name" AS role_name, - array_agg(DISTINCT all_actions."action") AS actions - FROM - groups_role_members grm - JOIN - groups_role_actions gra ON gra.role_id = grm.role_id - JOIN - groups_roles gr ON gr.id = grm.role_id - JOIN - "groups" g ON g.id = gr.entity_id - JOIN - groups_role_actions all_actions ON all_actions.role_id = grm.role_id - WHERE - grm.member_id = :user_id - AND g.domain_id = :domain_id_param - AND gra."action" LIKE 'subgroup_channel%' - GROUP BY - gr.entity_id, grm.member_id, gr.id, gr."name", g."path", g.id -), -direct_leaf_groups_with_subgroup AS ( - SELECT dgws.* - FROM direct_groups_with_subgroup dgws - WHERE NOT EXISTS ( - SELECT 1 - FROM direct_groups_with_subgroup dgws2 - WHERE - dgws2.path @> dgws.path - AND dgws2.id != dgws.id - ) -), -indirect_child_groups AS ( - SELECT - DISTINCT indirect_child_groups.id as child_id, - indirect_child_groups.*, - dlgws.id as access_provider_id, - dlgws.role_id as access_provider_role_id, - dlgws.role_name as access_provider_role_name, - dlgws.actions as access_provider_role_actions - FROM - direct_leaf_groups_with_subgroup dlgws - JOIN - groups indirect_child_groups ON indirect_child_groups.path <@ dlgws.path - WHERE - indirect_child_groups.domain_id = :domain_id_param - AND NOT EXISTS ( - SELECT 1 - FROM direct_groups_with_subgroup dgws - WHERE dgws.id = indirect_child_groups.id - ) -), -final_groups AS ( - SELECT - id, - parent_id, - domain_id, - "name", - description, - metadata, - created_at, - updated_at, - updated_by, - status, - "path", - '' AS role_id, - '' AS role_name, - CAST(array[] AS text[]) AS actions, - 'direct_group' AS access_type, - id AS access_provider_id, - role_id AS access_provider_role_id, - role_name AS access_provider_role_name, - actions AS access_provider_role_actions - FROM - direct_groups - UNION - SELECT - id, - parent_id, - domain_id, - "name", - description, - metadata, - created_at, - updated_at, - updated_by, - status, - "path", - '' AS role_id, - '' AS role_name, - CAST(array[] AS text[]) AS actions, - 'indirect_group' AS access_type, - access_provider_id, - access_provider_role_id, - access_provider_role_name, - access_provider_role_actions - FROM - indirect_child_groups -), -groups_channels AS ( - SELECT - c.id, - c.name, - c.domain_id, - c.parent_group_id, - c.route, - c.tags, - c.metadata, - c.created_by, - c.created_at, - c.updated_at, - c.updated_by, - c.status, - g.path AS parent_group_path, - g.role_id, - g.role_name, - g.actions, - g.access_type, - g.access_provider_id, - g.access_provider_role_id, - g.access_provider_role_name, - g.access_provider_role_actions - FROM - final_groups g - JOIN - channels c ON c.parent_group_id = g.id - WHERE - NOT EXISTS (SELECT 1 FROM direct_channels dc WHERE dc.id = c.id) - UNION - SELECT * FROM direct_channels -), -final_channels AS ( - SELECT - gc.id, - gc."name", - gc.domain_id, - gc.parent_group_id, - gc.route, - gc.tags, - gc.metadata, - gc.created_by, - gc.created_at, - gc.updated_at, - gc.updated_by, - gc.status, - gc.parent_group_path, - gc.role_id, - gc.role_name, - gc.actions, - gc.access_type, - gc.access_provider_id, - gc.access_provider_role_id, - gc.access_provider_role_name, - gc.access_provider_role_actions - FROM - groups_channels AS gc - UNION - SELECT - dc.id, - dc."name", - dc.domain_id, - dc.parent_group_id, - dc.route, - dc.tags, - dc.metadata, - dc.created_by, - dc.created_at, - dc.updated_at, - dc.updated_by, - dc.status, - g."path" AS parent_group_path, - '' AS role_id, - '' AS role_name, - CAST(array[] AS text[]) AS actions, - 'domain' AS access_type, - d.id AS access_provider_id, - dr.id AS access_provider_role_id, - dr."name" AS access_provider_role_name, - array_agg(dra."action") as access_provider_role_actions - FROM - domains_role_members drm - JOIN - domains_role_actions dra ON dra.role_id = drm.role_id - JOIN - domains_roles dr ON dr.id = drm.role_id - JOIN - domains d ON d.id = dr.entity_id - JOIN - channels dc ON dc.domain_id = d.id - LEFT JOIN - groups g ON dc.parent_group_id = g.id - WHERE - drm.member_id = :user_id - AND d.id = :domain_id_param - AND dra."action" LIKE 'channel_%' - AND NOT EXISTS ( - SELECT 1 FROM groups_channels gc - WHERE gc.id = dc.id - ) - GROUP BY - dc.id, d.id, dr.id, g."path" -) - ` - -func (cr *channelRepository) Remove(ctx context.Context, ids ...string) error { - q := "DELETE FROM channels AS c WHERE c.id = ANY(:channel_ids) ;" - params := map[string]any{ - "channel_ids": ids, - } - result, err := cr.db.NamedExecContext(ctx, q, params) - if err != nil { - return cr.eh.HandleError(repoerr.ErrRemoveEntity, err) - } - if rows, _ := result.RowsAffected(); rows == 0 { - return repoerr.ErrNotFound - } - return nil -} - -func (cr *channelRepository) SetParentGroup(ctx context.Context, ch channels.Channel) error { - q := "UPDATE channels SET parent_group_id = :parent_group_id, updated_at = :updated_at, updated_by = :updated_by WHERE id = :id" - dbCh, err := toDBChannel(ch) - if err != nil { - return errors.Wrap(repoerr.ErrUpdateEntity, err) - } - result, err := cr.db.NamedExecContext(ctx, q, dbCh) - if err != nil { - return cr.eh.HandleError(repoerr.ErrUpdateEntity, err) - } - if rows, _ := result.RowsAffected(); rows == 0 { - return repoerr.ErrNotFound - } - return nil -} - -func (cr *channelRepository) RemoveParentGroup(ctx context.Context, ch channels.Channel) error { - q := "UPDATE channels SET parent_group_id = NULL, updated_at = :updated_at, updated_by = :updated_by WHERE id = :id" - dbCh, err := toDBChannel(ch) - if err != nil { - return errors.Wrap(repoerr.ErrUpdateEntity, err) - } - result, err := cr.db.NamedExecContext(ctx, q, dbCh) - if err != nil { - return cr.eh.HandleError(repoerr.ErrRemoveEntity, err) - } - if rows, _ := result.RowsAffected(); rows == 0 { - return repoerr.ErrNotFound - } - return nil -} - -func (cr *channelRepository) AddConnections(ctx context.Context, conns []channels.Connection) error { - dbConns := toDBConnections(conns) - q := `INSERT INTO connections (channel_id, domain_id, client_id, type) - VALUES (:channel_id, :domain_id, :client_id, :type );` - - if _, err := cr.db.NamedExecContext(ctx, q, dbConns); err != nil { - return cr.eh.HandleError(repoerr.ErrCreateEntity, err) - } - - return nil -} - -func (cr *channelRepository) RemoveConnections(ctx context.Context, conns []channels.Connection) (retErr error) { - tx, err := cr.db.BeginTxx(ctx, nil) - if err != nil { - return cr.eh.HandleError(repoerr.ErrRemoveEntity, err) - } - defer func() { - if retErr != nil { - if errRollBack := tx.Rollback(); errRollBack != nil { - retErr = errors.Wrap(retErr, errors.Wrap(apiutil.ErrRollbackTx, errRollBack)) - } - } - }() - - query := `DELETE FROM connections WHERE channel_id = :channel_id AND domain_id = :domain_id AND client_id = :client_id` - - for _, conn := range conns { - if uint8(conn.Type) > 0 { - query = query + " AND type = :type " - } - dbConn := toDBConnection(conn) - if _, err := tx.NamedExec(query, dbConn); err != nil { - return cr.eh.HandleError(repoerr.ErrRemoveEntity, errors.Wrap(fmt.Errorf("failed to delete connection for channel_id: %s, domain_id: %s client_id %s", conn.ChannelID, conn.DomainID, conn.ClientID), err)) - } - } - if err := tx.Commit(); err != nil { - return cr.eh.HandleError(repoerr.ErrRemoveEntity, err) - } - return nil -} - -func (cr *channelRepository) CheckConnection(ctx context.Context, conn channels.Connection) error { - query := `SELECT 1 FROM connections WHERE channel_id = :channel_id AND domain_id = :domain_id AND client_id = :client_id AND type = :type LIMIT 1` - dbConn := toDBConnection(conn) - rows, err := cr.db.NamedQueryContext(ctx, query, dbConn) - if err != nil { - return cr.eh.HandleError(repoerr.ErrViewEntity, err) - } - defer rows.Close() - - if !rows.Next() { - return repoerr.ErrNotFound - } - return nil -} - -func (cr *channelRepository) ClientAuthorize(ctx context.Context, conn channels.Connection) error { - query := `SELECT 1 FROM connections WHERE channel_id = :channel_id AND client_id = :client_id AND domain_id = :domain_id AND type = :type LIMIT 1` - dbConn := toDBConnection(conn) - rows, err := cr.db.NamedQueryContext(ctx, query, dbConn) - if err != nil { - return cr.eh.HandleError(repoerr.ErrViewEntity, err) - } - defer rows.Close() - - if !rows.Next() { - return repoerr.ErrNotFound - } - return nil -} - -func (cr *channelRepository) ChannelConnectionsCount(ctx context.Context, id string) (uint64, error) { - query := `SELECT COUNT(*) FROM connections WHERE channel_id = :channel_id` - dbConn := dbConnection{ChannelID: id} - - total, err := postgres.Total(ctx, cr.db, query, dbConn) - if err != nil { - return 0, cr.eh.HandleError(repoerr.ErrViewEntity, err) - } - return total, nil -} - -func (cr *channelRepository) DoesChannelHaveConnections(ctx context.Context, id string) (bool, error) { - query := `SELECT 1 FROM connections WHERE channel_id = :channel_id` - dbConn := dbConnection{ChannelID: id} - - rows, err := cr.db.NamedQueryContext(ctx, query, dbConn) - if err != nil { - return false, cr.eh.HandleError(repoerr.ErrViewEntity, err) - } - defer rows.Close() - - return rows.Next(), nil -} - -func (cr *channelRepository) RemoveClientConnections(ctx context.Context, clientID string) error { - query := `DELETE FROM connections WHERE client_id = :client_id` - - dbConn := dbConnection{ClientID: clientID} - if _, err := cr.db.NamedExecContext(ctx, query, dbConn); err != nil { - return cr.eh.HandleError(repoerr.ErrRemoveEntity, err) - } - return nil -} - -func (cr *channelRepository) RemoveChannelConnections(ctx context.Context, channelID string) error { - query := `DELETE FROM connections WHERE channel_id = :channel_id` - - dbConn := dbConnection{ChannelID: channelID} - if _, err := cr.db.NamedExecContext(ctx, query, dbConn); err != nil { - return cr.eh.HandleError(repoerr.ErrRemoveEntity, err) - } - return nil -} - -func (cr *channelRepository) RetrieveParentGroupChannels(ctx context.Context, parentGroupID string) ([]channels.Channel, error) { - query := `SELECT c.id, c.name, c.tags, c.metadata, COALESCE(c.domain_id, '') AS domain_id, COALESCE(parent_group_id, '') AS parent_group_id, c.status, - c.created_by, c.created_at, c.updated_at, COALESCE(c.updated_by, '') AS updated_by FROM channels c WHERE c.parent_group_id = :parent_group_id ;` - - rows, err := cr.db.NamedQueryContext(ctx, query, dbChannel{ParentGroup: toNullString(parentGroupID)}) - if err != nil { - return []channels.Channel{}, cr.eh.HandleError(repoerr.ErrViewEntity, err) - } - defer rows.Close() - - var chs []channels.Channel - for rows.Next() { - dbch := dbChannel{} - if err := rows.StructScan(&dbch); err != nil { - return []channels.Channel{}, cr.eh.HandleError(repoerr.ErrViewEntity, err) - } - - ch, err := toChannel(dbch) - if err != nil { - return []channels.Channel{}, err - } - - chs = append(chs, ch) - } - return chs, nil -} - -func (cr *channelRepository) UnsetParentGroupFromChannels(ctx context.Context, parentGroupID string) error { - query := "UPDATE channels SET parent_group_id = NULL WHERE parent_group_id = :parent_group_id" - - if _, err := cr.db.NamedExecContext(ctx, query, dbChannel{ParentGroup: toNullString(parentGroupID)}); err != nil { - return cr.eh.HandleError(repoerr.ErrRemoveEntity, err) - } - return nil -} - -func (cr *channelRepository) update(ctx context.Context, ch channels.Channel, query string) (channels.Channel, error) { - dbch, err := toDBChannel(ch) - if err != nil { - return channels.Channel{}, errors.Wrap(repoerr.ErrUpdateEntity, err) - } - - row, err := cr.db.NamedQueryContext(ctx, query, dbch) - if err != nil { - return channels.Channel{}, cr.eh.HandleError(repoerr.ErrUpdateEntity, err) - } - defer row.Close() - - dbch = dbChannel{} - if row.Next() { - if err := row.StructScan(&dbch); err != nil { - return channels.Channel{}, cr.eh.HandleError(repoerr.ErrUpdateEntity, err) - } - - return toChannel(dbch) - } - - return channels.Channel{}, repoerr.ErrNotFound -} - -type dbChannel struct { - ID string `db:"id"` - Name string `db:"name,omitempty"` - ParentGroup sql.NullString `db:"parent_group_id,omitempty"` - Tags pgtype.TextArray `db:"tags,omitempty"` - Domain string `db:"domain_id"` - Route sql.NullString `db:"route,omitempty"` - Metadata []byte `db:"metadata,omitempty"` - CreatedBy *string `db:"created_by,omitempty"` - CreatedAt time.Time `db:"created_at,omitempty"` - UpdatedAt sql.NullTime `db:"updated_at,omitempty"` - UpdatedBy *string `db:"updated_by,omitempty"` - Status channels.Status `db:"status,omitempty"` - ParentGroupPath sql.NullString `db:"parent_group_path,omitempty"` - RoleID string `db:"role_id,omitempty"` - RoleName string `db:"role_name,omitempty"` - Actions pq.StringArray `db:"actions,omitempty"` - AccessType string `db:"access_type,omitempty"` - AccessProviderId string `db:"access_provider_id,omitempty"` - AccessProviderRoleId string `db:"access_provider_role_id,omitempty"` - AccessProviderRoleName string `db:"access_provider_role_name,omitempty"` - AccessProviderRoleActions pq.StringArray `db:"access_provider_role_actions,omitempty"` - ConnectionTypes pq.Int32Array `db:"connection_types,omitempty"` - MemberID string `db:"member_id,omitempty"` - Roles json.RawMessage `db:"roles,omitempty"` - TotalCount uint64 `db:"total_count"` -} - -func toDBChannel(ch channels.Channel) (dbChannel, error) { - data := []byte("{}") - if len(ch.Metadata) > 0 { - b, err := json.Marshal(ch.Metadata) - if err != nil { - return dbChannel{}, errors.Wrap(repoerr.ErrMalformedEntity, err) - } - data = b - } - var tags pgtype.TextArray - if err := tags.Set(ch.Tags); err != nil { - return dbChannel{}, err - } - var createdBy *string - if ch.CreatedBy != "" { - createdBy = &ch.CreatedBy - } - var updatedBy *string - if ch.UpdatedBy != "" { - updatedBy = &ch.UpdatedBy - } - var updatedAt sql.NullTime - if ch.UpdatedAt != (time.Time{}) { - updatedAt = sql.NullTime{Time: ch.UpdatedAt, Valid: true} - } - return dbChannel{ - ID: ch.ID, - Name: ch.Name, - ParentGroup: toNullString(ch.ParentGroup), - Domain: ch.Domain, - Route: toNullString(ch.Route), - Tags: tags, - Metadata: data, - CreatedBy: createdBy, - CreatedAt: ch.CreatedAt, - UpdatedAt: updatedAt, - UpdatedBy: updatedBy, - Status: ch.Status, - }, nil -} - -func toNullString(s string) sql.NullString { - if s == "" { - return sql.NullString{} - } - - return sql.NullString{ - String: s, - Valid: true, - } -} - -func toString(s sql.NullString) string { - if s.Valid { - return s.String - } - return "" -} - -func toChannel(ch dbChannel) (channels.Channel, error) { - var metadata channels.Metadata - if ch.Metadata != nil { - if err := json.Unmarshal([]byte(ch.Metadata), &metadata); err != nil { - return channels.Channel{}, errors.Wrap(errors.ErrMalformedEntity, err) - } - } - var tags []string - for _, e := range ch.Tags.Elements { - tags = append(tags, e.String) - } - var createdBy string - if ch.CreatedBy != nil { - createdBy = *ch.CreatedBy - } - var updatedBy string - if ch.UpdatedBy != nil { - updatedBy = *ch.UpdatedBy - } - var updatedAt time.Time - if ch.UpdatedAt.Valid { - updatedAt = ch.UpdatedAt.Time.UTC() - } - - connTypes := []connections.ConnType{} - for _, ct := range ch.ConnectionTypes { - connType, err := connections.NewType(uint(ct)) - if err != nil { - return channels.Channel{}, err - } - connTypes = append(connTypes, connType) - } - - var roles []roles.MemberRoleActions - if ch.Roles != nil { - if err := json.Unmarshal(ch.Roles, &roles); err != nil { - return channels.Channel{}, errors.Wrap(errors.ErrMalformedEntity, err) - } - } - - newCh := channels.Channel{ - ID: ch.ID, - Name: ch.Name, - Tags: tags, - Domain: ch.Domain, - Route: toString(ch.Route), - ParentGroup: toString(ch.ParentGroup), - Metadata: metadata, - CreatedBy: createdBy, - CreatedAt: ch.CreatedAt.UTC(), - UpdatedAt: updatedAt, - UpdatedBy: updatedBy, - Status: ch.Status, - ParentGroupPath: toString(ch.ParentGroupPath), - RoleID: ch.RoleID, - RoleName: ch.RoleName, - Actions: ch.Actions, - AccessType: ch.AccessType, - AccessProviderId: ch.AccessProviderId, - AccessProviderRoleId: ch.AccessProviderRoleId, - AccessProviderRoleName: ch.AccessProviderRoleName, - AccessProviderRoleActions: ch.AccessProviderRoleActions, - ConnectionTypes: connTypes, - Roles: roles, - } - - return newCh, nil -} - -func PageQuery(pm channels.Page) (string, error) { - mq, _, err := postgres.CreateMetadataQuery("", pm.Metadata) - if err != nil { - return "", errors.Wrap(errors.ErrMalformedEntity, err) - } - - var query []string - if pm.Name != "" { - query = append(query, "c.name ILIKE '%' || :name || '%'") - } - - if pm.ID != "" { - query = append(query, "c.id = :id") - } - if len(pm.Tags.Elements) > 0 { - switch pm.Tags.Operator { - case channels.AndOp: - query = append(query, "tags @> :tags") - default: // OR - query = append(query, "tags && :tags") - } - } - - if mq != "" { - query = append(query, mq) - } - - if len(pm.IDs) != 0 { - query = append(query, "id = ANY(:ids)") - } - if pm.Status != channels.AllStatus { - query = append(query, "c.status = :status") - } - if pm.Domain != "" { - query = append(query, "c.domain_id = :domain_id") - } - if pm.Group.Valid { - switch { - case pm.Group.Value != "": - query = append(query, "c.parent_group_path <@ (SELECT path from groups where id = :group_id) ") - default: - query = append(query, "c.parent_group_id = '' ") - } - } - - if pm.Client != "" { - query = append(query, "conn.client_id = :client_id ") - if pm.ConnectionType != "" { - query = append(query, ":conn_type = ANY(conn.connection_types) ") - } - } - if pm.AccessType != "" { - query = append(query, "c.access_type = :access_type") - } - if pm.RoleID != "" { - query = append(query, "c.role_id = :role_id") - } - if pm.RoleName != "" { - query = append(query, "c.role_name = :role_name") - } - if len(pm.Actions) != 0 { - query = append(query, "c.actions @> :actions") - } - if len(pm.Metadata) > 0 { - query = append(query, "c.metadata @> :metadata") - } - - if !pm.CreatedFrom.IsZero() { - query = append(query, "c.created_at >= :created_from") - } - if !pm.CreatedTo.IsZero() { - query = append(query, "c.created_at <= :created_to") - } - - var emq string - if len(query) > 0 { - emq = fmt.Sprintf("WHERE %s", strings.Join(query, " AND ")) - } - return emq, nil -} - -func applyOrdering(emq string, pm channels.Page) string { - col := "COALESCE(c.updated_at, c.created_at)" - switch pm.Order { - case "name": - col = "c.name" - case "created_at": - col = "c.created_at" - case "updated_at", "": - col = "COALESCE(c.updated_at, c.created_at)" - } - - dir := pm.Dir - if dir != api.AscDir && dir != api.DescDir { - dir = api.DescDir - } - - return fmt.Sprintf("%s ORDER BY %s %s, c.id %s", emq, col, dir, dir) -} - -func applyLimitOffset(query string) string { - return fmt.Sprintf(`%s - LIMIT :limit OFFSET :offset`, query) -} - -func toDBChannelsPage(pm channels.Page) (dbChannelsPage, error) { - _, data, err := postgres.CreateMetadataQuery("", pm.Metadata) - if err != nil { - return dbChannelsPage{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - var tags pgtype.TextArray - if err := tags.Set(pm.Tags.Elements); err != nil { - return dbChannelsPage{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - - var connType uint8 - if pm.ConnectionType != "" { - ct, err := connections.ParseConnType(pm.ConnectionType) - if err != nil { - return dbChannelsPage{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - connType = uint8(ct) - } - - return dbChannelsPage{ - Limit: pm.Limit, - Offset: pm.Offset, - Name: pm.Name, - Id: pm.ID, - Domain: pm.Domain, - Metadata: data, - Tags: tags, - Status: pm.Status, - GroupID: sql.NullString{Valid: pm.Group.Valid, String: pm.Group.Value}, - ClientID: pm.Client, - ConnType: connType, - RoleName: pm.RoleName, - RoleID: pm.RoleID, - Actions: pm.Actions, - AccessType: pm.AccessType, - IDs: pq.StringArray(pm.IDs), - CreatedFrom: pm.CreatedFrom, - CreatedTo: pm.CreatedTo, - }, nil -} - -type dbChannelsPage struct { - Limit uint64 `db:"limit"` - Offset uint64 `db:"offset"` - Name string `db:"name"` - Id string `db:"id"` - Domain string `db:"domain_id"` - Metadata []byte `db:"metadata"` - Tags pgtype.TextArray `db:"tags"` - Status channels.Status `db:"status"` - GroupID sql.NullString `db:"group_id"` - ClientID string `db:"client_id"` - ConnType uint8 `db:"conn_type"` - RoleName string `db:"role_name"` - RoleID string `db:"role_id"` - Actions pq.StringArray `db:"actions"` - AccessType string `db:"access_type"` - CreatedFrom time.Time `db:"created_from"` - CreatedTo time.Time `db:"created_to"` - IDs pq.StringArray `db:"ids"` - UserID string `db:"user_id"` - DomainID string `db:"domain_id_param"` -} - -type dbConnection struct { - ChannelID string `db:"channel_id"` - DomainID string `db:"domain_id"` - ClientID string `db:"client_id"` - Type connections.ConnType `db:"type"` -} - -func toDBConnections(conns []channels.Connection) []dbConnection { - var dbconns []dbConnection - for _, conn := range conns { - dbconns = append(dbconns, toDBConnection(conn)) - } - return dbconns -} - -func toDBConnection(conn channels.Connection) dbConnection { - return dbConnection{ - ClientID: conn.ClientID, - ChannelID: conn.ChannelID, - DomainID: conn.DomainID, - Type: conn.Type, - } -} diff --git a/channels/postgres/channels_test.go b/channels/postgres/channels_test.go deleted file mode 100644 index 113669a00..000000000 --- a/channels/postgres/channels_test.go +++ /dev/null @@ -1,3333 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package postgres_test - -import ( - "context" - "fmt" - "strconv" - "strings" - "testing" - "time" - - "github.com/0x6flab/namegenerator" - "github.com/absmach/magistrala/channels" - "github.com/absmach/magistrala/channels/postgres" - "github.com/absmach/magistrala/domains" - dpostgres "github.com/absmach/magistrala/domains/postgres" - "github.com/absmach/magistrala/groups" - gpostgres "github.com/absmach/magistrala/groups/postgres" - "github.com/absmach/magistrala/internal/nullable" - "github.com/absmach/magistrala/internal/testsutil" - "github.com/absmach/magistrala/pkg/connections" - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - "github.com/absmach/magistrala/pkg/roles" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" -) - -var ( - namegen = namegenerator.NewGenerator() - invalidID = strings.Repeat("a", 37) - validChannel = channels.Channel{ - ID: testsutil.GenerateUUID(&testing.T{}), - Domain: testsutil.GenerateUUID(&testing.T{}), - ParentGroup: testsutil.GenerateUUID(&testing.T{}), - Name: namegen.Generate(), - Route: testsutil.GenerateUUID(&testing.T{}), - Tags: []string{"tag1", "tag2"}, - Metadata: map[string]any{"key": "value"}, - CreatedAt: time.Now().UTC().Truncate(time.Microsecond), - Status: channels.EnabledStatus, - ConnectionTypes: []connections.ConnType{}, - } - validConnection = channels.Connection{ - ClientID: testsutil.GenerateUUID(&testing.T{}), - ChannelID: validChannel.ID, - DomainID: validChannel.Domain, - Type: connections.Publish, - } - validTimestamp = time.Now().UTC().Truncate(time.Millisecond) - directAccess = "direct" - directGroupAccess = "direct_group" - domainAccess = "domain" - defOrder = "created_at" - ascDir = "asc" - descDir = "desc" - availableActions = []string{ - "delete", - "membership", - "read", - "update", - } - domainAvailableActions = []string{ - "channel_add_role_users", - "channel_connect_to_client", - "channel_create", - "channel_delete", - "channel_manage_role", - "channel_read", - "channel_remove_role_users", - "channel_set_parent_group", - "channel_update", - "channel_view_role_users", - } - groupAvailableActions = []string{ - "channel_add_role_users", - "channel_connect_to_client", - "channel_create", - "channel_delete", - "channel_manage_role", - "channel_read", - "channel_remove_role_users", - "channel_set_parent_group", - "channel_update", - "channel_view_role_users", - "subgroup_channel_add_role_users", - "subgroup_channel_connect_to_client", - "subgroup_channel_create", - "subgroup_channel_delete", - "subgroup_channel_manage_role", - "subgroup_channel_read", - "subgroup_channel_remove_role_users", - "subgroup_channel_set_parent_group", - "subgroup_channel_update", - "subgroup_channel_view_role_users", - "subgroup_manage_role", - "subgroup_membership", - "subgroup_read", - "subgroup_remove_role_users", - "subgroup_set_child", - "subgroup_set_parent", - "subgroup_update", - } - errChannelExists = errors.New("channel id already exists") -) - -func TestSave(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM channels") - require.Nil(t, err, fmt.Sprintf("clean channels unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - duplicateChannelID := testsutil.GenerateUUID(t) - - duplicateRoute := testsutil.GenerateUUID(t) - duplicateDomain := testsutil.GenerateUUID(t) - - duplicateChannel := channels.Channel{ - ID: testsutil.GenerateUUID(t), - Domain: duplicateDomain, - Name: namegen.Generate(), - Route: duplicateRoute, - } - - _, err := repo.Save(context.Background(), duplicateChannel) - require.Nil(t, err, fmt.Sprintf("save channel unexpected error: %s", err)) - - cases := []struct { - desc string - channel channels.Channel - resp []channels.Channel - err error - }{ - { - desc: "add new channel successfully", - channel: validChannel, - resp: []channels.Channel{validChannel}, - err: nil, - }, - { - desc: "add duplicate channel", - channel: validChannel, - resp: []channels.Channel{}, - err: errChannelExists, - }, - { - desc: "add channel with invalid ID", - channel: channels.Channel{ - ID: invalidID, - Domain: testsutil.GenerateUUID(t), - Name: namegen.Generate(), - Metadata: map[string]any{"key": "value"}, - CreatedAt: time.Now().UTC().Truncate(time.Microsecond), - Status: channels.EnabledStatus, - }, - resp: []channels.Channel{}, - err: repoerr.ErrCreateEntity, - }, - { - desc: "add channel with invalid domain", - channel: channels.Channel{ - ID: testsutil.GenerateUUID(t), - Domain: invalidID, - Name: namegen.Generate(), - Metadata: map[string]any{"key": "value"}, - CreatedAt: time.Now().UTC().Truncate(time.Microsecond), - Status: channels.EnabledStatus, - }, - resp: []channels.Channel{}, - err: repoerr.ErrCreateEntity, - }, - { - desc: "add channel with invalid name", - channel: channels.Channel{ - ID: testsutil.GenerateUUID(t), - Domain: testsutil.GenerateUUID(t), - Name: strings.Repeat("a", 1025), - Metadata: map[string]any{"key": "value"}, - CreatedAt: time.Now().UTC().Truncate(time.Microsecond), - Status: channels.EnabledStatus, - }, - resp: []channels.Channel{}, - err: repoerr.ErrCreateEntity, - }, - { - desc: "add channel with invalid metadata", - channel: channels.Channel{ - ID: testsutil.GenerateUUID(t), - Domain: testsutil.GenerateUUID(t), - Name: namegen.Generate(), - Metadata: map[string]any{ - "key": make(chan int), - }, - CreatedAt: time.Now().UTC().Truncate(time.Microsecond), - Status: channels.EnabledStatus, - }, - resp: []channels.Channel{}, - err: repoerr.ErrCreateEntity, - }, - { - desc: "add channel with duplicate name", - channel: channels.Channel{ - ID: duplicateChannelID, - Domain: validChannel.Domain, - Name: validChannel.Name, - Metadata: map[string]any{"key": "different_value"}, - CreatedAt: validTimestamp, - Status: channels.EnabledStatus, - }, - resp: []channels.Channel{ - { - ID: duplicateChannelID, - Domain: validChannel.Domain, - Name: validChannel.Name, - Metadata: map[string]any{"key": "different_value"}, - CreatedAt: validTimestamp, - Status: channels.EnabledStatus, - ConnectionTypes: []connections.ConnType{}, - }, - }, - err: nil, - }, - { - desc: "add channel with duplicate route", - channel: channels.Channel{ - ID: testsutil.GenerateUUID(t), - Domain: duplicateDomain, - Name: namegen.Generate(), - Route: duplicateRoute, - }, - resp: []channels.Channel{}, - err: errors.ErrRouteNotAvailable, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - channels, err := repo.Save(context.Background(), tc.channel) - assert.Equal(t, tc.resp, channels, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, channels)) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - }) - } -} - -func TestUpdate(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM channels") - require.Nil(t, err, fmt.Sprintf("clean channels unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - _, err := repo.Save(context.Background(), validChannel) - require.Nil(t, err, fmt.Sprintf("save channel unexpected error: %s", err)) - - cases := []struct { - desc string - update string - channel channels.Channel - err error - }{ - { - desc: "update channel successfully", - update: "all", - channel: channels.Channel{ - ID: validChannel.ID, - Name: namegen.Generate(), - Route: testsutil.GenerateUUID(t), - Metadata: map[string]any{"key": "value"}, - UpdatedAt: validTimestamp, - UpdatedBy: testsutil.GenerateUUID(t), - }, - err: nil, - }, - { - desc: "update channel name", - update: "name", - channel: channels.Channel{ - ID: validChannel.ID, - Name: namegen.Generate(), - UpdatedAt: validTimestamp, - UpdatedBy: testsutil.GenerateUUID(t), - }, - err: nil, - }, - { - desc: "update channel metadata", - update: "metadata", - channel: channels.Channel{ - ID: validChannel.ID, - Metadata: map[string]any{"key1": "value1"}, - UpdatedAt: validTimestamp, - UpdatedBy: testsutil.GenerateUUID(t), - }, - err: nil, - }, - { - desc: "update channel with invalid ID", - update: "all", - channel: channels.Channel{ - ID: testsutil.GenerateUUID(t), - Name: namegen.Generate(), - Metadata: map[string]any{"key": "value"}, - UpdatedAt: validTimestamp, - UpdatedBy: testsutil.GenerateUUID(t), - }, - err: repoerr.ErrNotFound, - }, - { - desc: "update channel with empty ID", - update: "all", - channel: channels.Channel{ - Name: namegen.Generate(), - Metadata: map[string]any{"key": "value"}, - UpdatedAt: validTimestamp, - UpdatedBy: testsutil.GenerateUUID(t), - }, - err: repoerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - channel, err := repo.Update(context.Background(), tc.channel) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - if err == nil { - assert.Equal(t, tc.channel.ID, channel.ID, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.channel.ID, channel.ID)) - assert.Equal(t, tc.channel.UpdatedAt, channel.UpdatedAt, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.channel.UpdatedAt, channel.UpdatedAt)) - assert.Equal(t, tc.channel.UpdatedBy, channel.UpdatedBy, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.channel.UpdatedBy, channel.UpdatedBy)) - switch tc.update { - case "all": - assert.Equal(t, tc.channel.Name, channel.Name, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.channel.Name, channel.Name)) - assert.Equal(t, tc.channel.Metadata, channel.Metadata, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.channel.Metadata, channel.Metadata)) - case "name": - assert.Equal(t, tc.channel.Name, channel.Name, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.channel.Name, channel.Name)) - case "metadata": - assert.Equal(t, tc.channel.Metadata, channel.Metadata, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.channel.Metadata, channel.Metadata)) - } - } - }) - } -} - -func TestUpdateTags(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM channels") - require.Nil(t, err, fmt.Sprintf("clean channels unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - _, err := repo.Save(context.Background(), validChannel) - require.Nil(t, err, fmt.Sprintf("save channel unexpected error: %s", err)) - - cases := []struct { - desc string - channel channels.Channel - err error - }{ - { - desc: "update channel tags", - channel: channels.Channel{ - ID: validChannel.ID, - Tags: []string{"tag3", "tag4"}, - UpdatedAt: validTimestamp, - UpdatedBy: testsutil.GenerateUUID(t), - }, - err: nil, - }, - { - desc: "update channel with invalid ID", - channel: channels.Channel{ - ID: testsutil.GenerateUUID(t), - Tags: []string{"tag3", "tag4"}, - UpdatedAt: validTimestamp, - UpdatedBy: testsutil.GenerateUUID(t), - }, - err: repoerr.ErrNotFound, - }, - { - desc: "update channel with empty ID", - channel: channels.Channel{ - Tags: []string{"tag3", "tag4"}, - UpdatedAt: validTimestamp, - UpdatedBy: testsutil.GenerateUUID(t), - }, - err: repoerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - channel, err := repo.UpdateTags(context.Background(), tc.channel) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - if err == nil { - assert.Equal(t, tc.channel.ID, channel.ID, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.channel.ID, channel.ID)) - assert.Equal(t, tc.channel.UpdatedAt, channel.UpdatedAt, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.channel.UpdatedAt, channel.UpdatedAt)) - assert.Equal(t, tc.channel.UpdatedBy, channel.UpdatedBy, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.channel.UpdatedBy, channel.UpdatedBy)) - assert.Equal(t, tc.channel.Tags, channel.Tags, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.channel.Tags, channel.Tags)) - } - }) - } -} - -func TestChangeStatus(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM channels") - require.Nil(t, err, fmt.Sprintf("clean channels unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - disabledChannel := validChannel - disabledChannel.ID = testsutil.GenerateUUID(t) - disabledChannel.Name = namegen.Generate() - disabledChannel.Route = testsutil.GenerateUUID(t) - disabledChannel.Status = channels.DisabledStatus - - _, err := repo.Save(context.Background(), validChannel, disabledChannel) - require.Nil(t, err, fmt.Sprintf("save channel unexpected error: %s", err)) - - cases := []struct { - desc string - channel channels.Channel - err error - }{ - { - desc: "disable channel successfully", - channel: channels.Channel{ - ID: validChannel.ID, - Status: channels.DisabledStatus, - UpdatedAt: validTimestamp, - UpdatedBy: testsutil.GenerateUUID(t), - }, - err: nil, - }, - { - desc: "enable channel successfully", - channel: channels.Channel{ - ID: disabledChannel.ID, - Status: channels.EnabledStatus, - UpdatedAt: validTimestamp, - UpdatedBy: testsutil.GenerateUUID(t), - }, - err: nil, - }, - { - desc: "change status channel with invalid ID", - channel: channels.Channel{ - ID: testsutil.GenerateUUID(t), - Status: channels.DisabledStatus, - UpdatedAt: validTimestamp, - UpdatedBy: testsutil.GenerateUUID(t), - }, - err: repoerr.ErrNotFound, - }, - { - desc: "change status channel with empty ID", - channel: channels.Channel{ - Status: channels.DisabledStatus, - UpdatedAt: validTimestamp, - UpdatedBy: testsutil.GenerateUUID(t), - }, - err: repoerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - channel, err := repo.ChangeStatus(context.Background(), tc.channel) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - if err == nil { - assert.Equal(t, tc.channel.ID, channel.ID, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.channel.ID, channel.ID)) - assert.Equal(t, tc.channel.UpdatedAt, channel.UpdatedAt, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.channel.UpdatedAt, channel.UpdatedAt)) - assert.Equal(t, tc.channel.UpdatedBy, channel.UpdatedBy, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.channel.UpdatedBy, channel.UpdatedBy)) - assert.Equal(t, tc.channel.Status, channel.Status, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.channel.Status, channel.Status)) - } - }) - } -} - -func TestRetrieveByID(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM channels") - require.Nil(t, err, fmt.Sprintf("clean channels unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - _, err := repo.Save(context.Background(), validChannel) - require.Nil(t, err, fmt.Sprintf("save channel unexpected error: %s", err)) - - cases := []struct { - desc string - id string - resp channels.Channel - err error - }{ - { - desc: "retrieve channel by id successfully", - id: validChannel.ID, - resp: validChannel, - err: nil, - }, - { - desc: "retrieve channel by id with invalid ID", - id: invalidID, - err: repoerr.ErrNotFound, - }, - { - desc: "retrieve channel by id with empty ID", - id: "", - err: repoerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - channel, err := repo.RetrieveByID(context.Background(), tc.id) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - if err == nil { - assert.Nil(t, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, channel, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, channel)) - } - }) - } -} - -func TestRetrieveByRoute(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM channels") - require.Nil(t, err, fmt.Sprintf("clean channels unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - _, err := repo.Save(context.Background(), validChannel) - require.Nil(t, err, fmt.Sprintf("save channel unexpected error: %s", err)) - - cases := []struct { - desc string - route string - domainID string - resp channels.Channel - err error - }{ - { - desc: "retrieve channel by route successfully", - route: validChannel.Route, - domainID: validChannel.Domain, - resp: validChannel, - err: nil, - }, - { - desc: "retrieve channel by id with invalid route", - route: "invalid-route", - domainID: validChannel.Domain, - err: repoerr.ErrNotFound, - }, - { - desc: "retrieve channel by id with empty route", - route: "", - domainID: validChannel.Domain, - err: repoerr.ErrNotFound, - }, - { - desc: "retrieve channel by id with invalid domain", - route: validChannel.Route, - domainID: "invalid-domain", - err: repoerr.ErrNotFound, - }, - { - desc: "retrieve channel by id with empty domain", - route: validChannel.Route, - domainID: "", - err: repoerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - channel, err := repo.RetrieveByRoute(context.Background(), tc.route, tc.domainID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - if err == nil { - assert.Nil(t, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, channel, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, channel)) - } - }) - } -} - -func TestRetrieveAll(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM channels") - require.Nil(t, err, fmt.Sprintf("clean channels unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - num := 200 - - var items []channels.Channel - parentID := "" - baseTime := time.Now().UTC().Truncate(time.Millisecond) - for i := 0; i < num; i++ { - name := namegen.Generate() - channel := channels.Channel{ - ID: testsutil.GenerateUUID(t), - Domain: testsutil.GenerateUUID(t), - ParentGroup: parentID, - Name: name, - Route: testsutil.GenerateUUID(t), - Metadata: map[string]any{"name": name}, - CreatedAt: baseTime.Add(time.Duration(i) * time.Millisecond), - UpdatedAt: baseTime.Add(time.Duration(i) * time.Millisecond), - Status: channels.EnabledStatus, - ConnectionTypes: []connections.ConnType{}, - Tags: []string{"tag1", "tag2"}, - } - if i%99 == 0 { - channel.Tags = []string{"tag1", "tag3"} - } - _, err := repo.Save(context.Background(), channel) - require.Nil(t, err, fmt.Sprintf("create channel unexpected error: %s", err)) - items = append(items, channel) - if i%20 == 0 { - parentID = channel.ID - } - } - - reversedChannels := []channels.Channel{} - for i := len(items) - 1; i >= 0; i-- { - reversedChannels = append(reversedChannels, items[i]) - } - - cases := []struct { - desc string - page channels.ChannelsPage - response channels.ChannelsPage - err error - }{ - { - desc: "retrieve channels successfully", - page: channels.ChannelsPage{ - Page: channels.Page{ - Offset: 0, - Limit: 10, - Order: defOrder, - Dir: ascDir, - }, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: uint64(num), - Offset: 0, - Limit: 10, - }, - Channels: items[:10], - }, - err: nil, - }, - { - desc: "retrieve channels with offset", - page: channels.ChannelsPage{ - Page: channels.Page{ - Offset: 10, - Limit: 10, - Order: defOrder, - Dir: ascDir, - }, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: uint64(num), - Offset: 10, - Limit: 10, - }, - Channels: items[10:20], - }, - err: nil, - }, - { - desc: "retrieve channels with limit", - page: channels.ChannelsPage{ - Page: channels.Page{ - Offset: 0, - Limit: 50, - Order: defOrder, - Dir: ascDir, - }, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: uint64(num), - Offset: 0, - Limit: 50, - }, - Channels: items[:50], - }, - err: nil, - }, - { - desc: "retrieve channels with offset and limit", - page: channels.ChannelsPage{ - Page: channels.Page{ - Offset: 50, - Limit: 50, - Order: defOrder, - Dir: ascDir, - }, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: uint64(num), - Offset: 50, - Limit: 50, - }, - Channels: items[50:100], - }, - err: nil, - }, - { - desc: "retrieve channels with offset out of range", - page: channels.ChannelsPage{ - Page: channels.Page{ - Offset: 1000, - Limit: 50, - Order: defOrder, - Dir: descDir, - }, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: uint64(num), - Offset: 1000, - Limit: 50, - }, - Channels: []channels.Channel(nil), - }, - err: nil, - }, - { - desc: "retrieve channels with offset and limit out of range", - page: channels.ChannelsPage{ - Page: channels.Page{ - Offset: 170, - Limit: 50, - Order: defOrder, - Dir: ascDir, - }, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: uint64(num), - Offset: 170, - Limit: 50, - }, - Channels: items[170:200], - }, - err: nil, - }, - { - desc: "retrieve channels with limit out of range", - page: channels.ChannelsPage{ - Page: channels.Page{ - Offset: 0, - Limit: 1000, - }, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: uint64(num), - Offset: 0, - Limit: 1000, - }, - Channels: items, - }, - err: nil, - }, - { - desc: "retrieve channels with empty page", - page: channels.ChannelsPage{}, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: uint64(num), - Offset: 0, - Limit: 0, - }, - Channels: []channels.Channel(nil), - }, - err: nil, - }, - { - desc: "retrieve channels with name", - page: channels.ChannelsPage{ - Page: channels.Page{ - Offset: 0, - Limit: 10, - Name: items[0].Name, - }, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 1, - Offset: 0, - Limit: 10, - }, - Channels: []channels.Channel{items[0]}, - }, - err: nil, - }, - { - desc: "retrieve channels with IDs filter", - page: channels.ChannelsPage{ - Page: channels.Page{ - Offset: 0, - Limit: 10, - IDs: []string{items[0].ID, items[1].ID, items[2].ID}, - }, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 3, - Offset: 0, - Limit: 10, - }, - Channels: []channels.Channel{items[0], items[1], items[2]}, - }, - err: nil, - }, - { - desc: "retrieve channels with non-existing IDs", - page: channels.ChannelsPage{ - Page: channels.Page{ - Offset: 0, - Limit: 10, - IDs: []string{testsutil.GenerateUUID(t), testsutil.GenerateUUID(t)}, - }, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 0, - Offset: 0, - Limit: 10, - }, - Channels: []channels.Channel(nil), - }, - err: nil, - }, - { - desc: "retrieve channels with domain", - page: channels.ChannelsPage{ - Page: channels.Page{ - Offset: 0, - Limit: 10, - Domain: items[0].Domain, - }, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 1, - Offset: 0, - Limit: 10, - }, - Channels: []channels.Channel{items[0]}, - }, - err: nil, - }, - { - desc: "retrieve channels with metadata", - page: channels.ChannelsPage{ - Page: channels.Page{ - Offset: 0, - Limit: 10, - Metadata: items[0].Metadata, - }, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 1, - Offset: 0, - Limit: 10, - }, - Channels: []channels.Channel{items[0]}, - }, - err: nil, - }, - { - desc: "retrieve channels with invalid metadata", - page: channels.ChannelsPage{ - Page: channels.Page{ - Offset: 0, - Limit: 10, - Metadata: map[string]any{ - "key": make(chan int), - }, - }, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 0, - Offset: 0, - Limit: 10, - }, - Channels: []channels.Channel(nil), - }, - err: errors.ErrMalformedEntity, - }, - { - desc: "retrieve channels with id", - page: channels.ChannelsPage{ - Page: channels.Page{ - Offset: 0, - Limit: 10, - ID: items[0].ID, - }, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 1, - Offset: 0, - Limit: 10, - }, - Channels: []channels.Channel{items[0]}, - }, - err: nil, - }, - { - desc: "retrieve channels with wrong id", - page: channels.ChannelsPage{ - Page: channels.Page{ - Offset: 0, - Limit: 10, - ID: "wrong", - }, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 0, - Offset: 0, - Limit: 10, - }, - Channels: []channels.Channel(nil), - }, - err: nil, - }, - { - desc: "retrieve channels with single tag", - page: channels.ChannelsPage{ - Page: channels.Page{ - Offset: 0, - Limit: uint64(num), - Tags: channels.TagsQuery{Elements: []string{"tag1"}, Operator: channels.OrOp}, - Status: channels.AllStatus, - }, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 200, - Offset: 0, - Limit: uint64(num), - }, - Channels: items, - }, - }, - { - desc: "retrieve channel with multiple tags and OR operator", - page: channels.ChannelsPage{ - Page: channels.Page{ - Offset: 0, - Limit: uint64(num), - Tags: channels.TagsQuery{Elements: []string{"tag2", "tag3"}, Operator: channels.OrOp}, - Status: channels.AllStatus, - }, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 200, - Offset: 0, - Limit: uint64(num), - }, - Channels: items, - }, - }, - { - desc: "retrieve channel with multiple tags and AND operator", - page: channels.ChannelsPage{ - Page: channels.Page{ - Offset: 0, - Limit: uint64(num), - Tags: channels.TagsQuery{Elements: []string{"tag1", "tag3"}, Operator: channels.AndOp}, - Status: channels.AllStatus, - }, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 3, - Offset: 0, - Limit: uint64(num), - }, - Channels: []channels.Channel{items[0], items[99], items[198]}, - }, - }, - { - desc: "retrieve channel with invalid tags", - page: channels.ChannelsPage{ - Page: channels.Page{ - Offset: 0, - Limit: uint64(num), - Tags: channels.TagsQuery{Elements: []string{namegen.Generate(), namegen.Generate()}, Operator: channels.OrOp}, - Status: channels.AllStatus, - }, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 0, - Offset: 0, - Limit: uint64(num), - }, - Channels: []channels.Channel(nil), - }, - }, - { - desc: "retrieve channels with order by name ascending", - page: channels.ChannelsPage{ - Page: channels.Page{ - Offset: 0, - Limit: 10, - Order: "name", - Dir: ascDir, - }, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: uint64(num), - Offset: 0, - Limit: 10, - }, - }, - err: nil, - }, - { - desc: "retrieve channels with order by name descending", - page: channels.ChannelsPage{ - Page: channels.Page{ - Offset: 0, - Limit: 10, - Order: "name", - Dir: descDir, - }, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: uint64(num), - Offset: 0, - Limit: 10, - }, - }, - err: nil, - }, - { - desc: "retrieve channels with order by created_at ascending", - page: channels.ChannelsPage{ - Page: channels.Page{ - Offset: 0, - Limit: 10, - Order: defOrder, - Dir: ascDir, - }, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: uint64(num), - Offset: 0, - Limit: 10, - }, - Channels: items[:10], - }, - err: nil, - }, - { - desc: "retrieve channels with order by created_at descending", - page: channels.ChannelsPage{ - Page: channels.Page{ - Offset: 0, - Limit: 10, - Order: defOrder, - Dir: descDir, - }, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: uint64(num), - Offset: 0, - Limit: 10, - }, - Channels: reversedChannels[:10], - }, - err: nil, - }, - { - desc: "retrieve channels with order by updated_at ascending", - page: channels.ChannelsPage{ - Page: channels.Page{ - Offset: 0, - Limit: 10, - Order: "updated_at", - Dir: ascDir, - }, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: uint64(num), - Offset: 0, - Limit: 10, - }, - }, - err: nil, - }, - { - desc: "retrieve channels with order by updated_at descending", - page: channels.ChannelsPage{ - Page: channels.Page{ - Offset: 0, - Limit: 10, - Order: "updated_at", - Dir: descDir, - }, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: uint64(num), - Offset: 0, - Limit: 10, - }, - }, - err: nil, - }, - { - desc: "retrieve channels with created_from", - page: channels.ChannelsPage{ - Page: channels.Page{ - Offset: 0, - Limit: 200, - Order: "created_at", - Dir: ascDir, - CreatedFrom: baseTime.Add(100 * time.Millisecond), - }, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 100, - Offset: 0, - Limit: 200, - }, - Channels: items[100:], - }, - err: nil, - }, - { - desc: "retrieve channels with created_to", - page: channels.ChannelsPage{ - Page: channels.Page{ - Offset: 0, - Limit: 200, - Order: "created_at", - Dir: ascDir, - CreatedTo: baseTime.Add(99 * time.Millisecond), - }, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 100, - Offset: 0, - Limit: 200, - }, - Channels: items[:100], - }, - err: nil, - }, - { - desc: "retrieve channels with both created_from and created_to", - page: channels.ChannelsPage{ - Page: channels.Page{ - Offset: 0, - Limit: 200, - Order: "created_at", - Dir: ascDir, - CreatedFrom: baseTime.Add(50 * time.Millisecond), - CreatedTo: baseTime.Add(149 * time.Millisecond), - }, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 100, - Offset: 0, - Limit: 200, - }, - Channels: items[50:150], - }, - err: nil, - }, - { - desc: "retrieve channels with created_from returning no results", - page: channels.ChannelsPage{ - Page: channels.Page{ - Offset: 0, - Limit: 10, - Order: "created_at", - Dir: ascDir, - CreatedFrom: baseTime.Add(1000 * time.Millisecond), - }, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 0, - Offset: 0, - Limit: 10, - }, - Channels: []channels.Channel{}, - }, - err: nil, - }, - { - desc: "retrieve channels with created_to returning no results", - page: channels.ChannelsPage{ - Page: channels.Page{ - Offset: 0, - Limit: 10, - Order: "created_at", - Dir: ascDir, - CreatedTo: baseTime.Add(-1 * time.Millisecond), - }, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 0, - Offset: 0, - Limit: 10, - }, - Channels: []channels.Channel{}, - }, - err: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - switch channels, err := repo.RetrieveAll(context.Background(), tc.page.Page); { - case err == nil: - assert.Nil(t, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.response.Total, channels.Total, fmt.Sprintf("%s: expected %d got %d\n", tc.desc, tc.response.Total, channels.Total)) - assert.Equal(t, tc.response.Limit, channels.Limit, fmt.Sprintf("%s: expected %d got %d\n", tc.desc, tc.response.Limit, channels.Limit)) - assert.Equal(t, tc.response.Offset, channels.Offset, fmt.Sprintf("%s: expected %d got %d\n", tc.desc, tc.response.Offset, channels.Offset)) - if len(tc.response.Channels) > 0 { - got := updateTimestamp(channels.Channels) - resp := updateTimestamp(tc.response.Channels) - assert.ElementsMatch(t, resp, got, fmt.Sprintf("%s: expected %+v got %+v\n", tc.desc, resp, got)) - } - verifyChannelsOrdering(t, channels.Channels, tc.page.Page.Order, tc.page.Page.Dir) - default: - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - } - }) - } -} - -func TestRemove(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM channels") - require.Nil(t, err, fmt.Sprintf("clean channels unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - _, err := repo.Save(context.Background(), validChannel) - require.Nil(t, err, fmt.Sprintf("save channel unexpected error: %s", err)) - - cases := []struct { - desc string - id string - err error - }{ - { - desc: "remove channel successfully", - id: validChannel.ID, - err: nil, - }, - { - desc: "remove channel with invalid ID", - id: invalidID, - err: repoerr.ErrNotFound, - }, - { - desc: "remove channel with empty ID", - id: "", - err: repoerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := repo.Remove(context.Background(), tc.id) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - }) - } -} - -func TestSetParentGroup(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM channels") - require.Nil(t, err, fmt.Sprintf("clean channels unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - _, err := repo.Save(context.Background(), validChannel) - require.Nil(t, err, fmt.Sprintf("save channel unexpected error: %s", err)) - - cases := []struct { - desc string - id string - parentGroupID string - err error - }{ - { - desc: "set parent group successfully", - id: validChannel.ID, - parentGroupID: testsutil.GenerateUUID(t), - err: nil, - }, - { - desc: "set parent group with invalid ID", - id: invalidID, - parentGroupID: testsutil.GenerateUUID(t), - err: repoerr.ErrNotFound, - }, - { - desc: "set parent group with empty ID", - id: "", - parentGroupID: testsutil.GenerateUUID(t), - err: repoerr.ErrNotFound, - }, - { - desc: "set parent group with invalid parent group ID", - id: validChannel.ID, - parentGroupID: invalidID, - err: repoerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := repo.SetParentGroup(context.Background(), channels.Channel{ - ID: tc.id, - ParentGroup: tc.parentGroupID, - }) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - if err == nil { - resp, err := repo.RetrieveByID(context.Background(), tc.id) - require.Nil(t, err, fmt.Sprintf("retrieve channel unexpected error: %s", err)) - assert.Equal(t, tc.id, resp.ID, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.id, resp.ID)) - assert.Equal(t, tc.parentGroupID, resp.ParentGroup, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.parentGroupID, resp.ParentGroup)) - } - }) - } -} - -func TestRemoveParentGroup(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM channels") - require.Nil(t, err, fmt.Sprintf("clean channels unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - _, err := repo.Save(context.Background(), validChannel) - require.Nil(t, err, fmt.Sprintf("save channel unexpected error: %s", err)) - - cases := []struct { - desc string - id string - err error - }{ - { - desc: "remove parent group successfully", - id: validChannel.ID, - err: nil, - }, - { - desc: "remove parent group with invalid ID", - id: invalidID, - err: repoerr.ErrNotFound, - }, - { - desc: "remove parent group with empty ID", - id: "", - err: repoerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := repo.RemoveParentGroup(context.Background(), channels.Channel{ - ID: tc.id, - }) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - if err == nil { - resp, err := repo.RetrieveByID(context.Background(), tc.id) - require.Nil(t, err, fmt.Sprintf("retrieve channel unexpected error: %s", err)) - assert.Equal(t, tc.id, resp.ID, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.id, resp.ID)) - assert.Equal(t, "", resp.ParentGroup, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, "", resp.ParentGroup)) - } - }) - } -} - -func TestAddConnection(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM connections") - require.Nil(t, err, fmt.Sprintf("clean connections unexpected error: %s", err)) - _, err = db.Exec("DELETE FROM channels") - require.Nil(t, err, fmt.Sprintf("clean channels unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - _, err := repo.Save(context.Background(), validChannel) - require.Nil(t, err, fmt.Sprintf("save channel unexpected error: %s", err)) - - cases := []struct { - desc string - connection channels.Connection - err error - }{ - { - desc: "add connection successfully", - connection: validConnection, - err: nil, - }, - { - desc: "add connection with non-existent channel", - connection: channels.Connection{ - ClientID: testsutil.GenerateUUID(t), - ChannelID: testsutil.GenerateUUID(t), - DomainID: validChannel.Domain, - Type: connections.Publish, - }, - err: repoerr.ErrCreateEntity, - }, - { - desc: "add connection with non-existent domain", - connection: channels.Connection{ - ClientID: testsutil.GenerateUUID(t), - ChannelID: validChannel.ID, - DomainID: testsutil.GenerateUUID(t), - Type: connections.Publish, - }, - err: repoerr.ErrCreateEntity, - }, - - { - desc: "add connection with invalid client ID", - connection: channels.Connection{ - ClientID: invalidID, - ChannelID: testsutil.GenerateUUID(t), - DomainID: testsutil.GenerateUUID(t), - Type: connections.Publish, - }, - err: repoerr.ErrCreateEntity, - }, - { - desc: "add connection with invalid channel ID", - connection: channels.Connection{ - ClientID: testsutil.GenerateUUID(t), - ChannelID: invalidID, - DomainID: testsutil.GenerateUUID(t), - Type: connections.Publish, - }, - err: repoerr.ErrCreateEntity, - }, - { - desc: "add connection with invalid domain ID", - connection: channels.Connection{ - ClientID: testsutil.GenerateUUID(t), - ChannelID: testsutil.GenerateUUID(t), - DomainID: invalidID, - Type: connections.Publish, - }, - err: repoerr.ErrCreateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := repo.AddConnections(context.Background(), []channels.Connection{tc.connection}) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - }) - } -} - -func TestRemoveConnection(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM connections") - require.Nil(t, err, fmt.Sprintf("clean connections unexpected error: %s", err)) - _, err = db.Exec("DELETE FROM channels") - require.Nil(t, err, fmt.Sprintf("clean channels unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - _, err := repo.Save(context.Background(), validChannel) - require.Nil(t, err, fmt.Sprintf("save channel unexpected error: %s", err)) - - err = repo.AddConnections(context.Background(), []channels.Connection{validConnection}) - require.Nil(t, err, fmt.Sprintf("add connection unexpected error: %s", err)) - - conn1 := channels.Connection{ - ClientID: testsutil.GenerateUUID(t), - ChannelID: validChannel.ID, - DomainID: validChannel.Domain, - Type: connections.Publish, - } - conn2 := channels.Connection{ - ClientID: testsutil.GenerateUUID(t), - ChannelID: validChannel.ID, - DomainID: validChannel.Domain, - Type: connections.Subscribe, - } - err = repo.AddConnections(context.Background(), []channels.Connection{conn1, conn2}) - require.Nil(t, err, fmt.Sprintf("add connections unexpected error: %s", err)) - - cases := []struct { - desc string - connections []channels.Connection - err error - }{ - { - desc: "remove connection successfully", - connections: []channels.Connection{validConnection}, - err: nil, - }, - { - desc: "remove connection with non-existent channel", - connections: []channels.Connection{ - { - ClientID: testsutil.GenerateUUID(t), - ChannelID: testsutil.GenerateUUID(t), - DomainID: validChannel.Domain, - Type: connections.Publish, - }, - }, - err: nil, - }, - { - desc: "remove connection with non-existent domain", - connections: []channels.Connection{ - { - ClientID: testsutil.GenerateUUID(t), - ChannelID: validChannel.ID, - DomainID: testsutil.GenerateUUID(t), - Type: connections.Publish, - }, - }, - err: nil, - }, - { - desc: "remove connection with non-existent client", - connections: []channels.Connection{ - { - ClientID: testsutil.GenerateUUID(t), - ChannelID: validChannel.ID, - DomainID: validChannel.Domain, - Type: connections.Publish, - }, - }, - err: nil, - }, - { - desc: "remove connection with invalid type", - connections: []channels.Connection{ - { - ClientID: validConnection.ClientID, - ChannelID: validConnection.ChannelID, - DomainID: validConnection.DomainID, - Type: connections.Invalid, - }, - }, - err: nil, - }, - { - desc: "remove multiple connections", - connections: []channels.Connection{conn1, conn2}, - err: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := repo.RemoveConnections(context.Background(), tc.connections) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - }) - } -} - -func TestCheckConnection(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM connections") - require.Nil(t, err, fmt.Sprintf("clean connections unexpected error: %s", err)) - _, err = db.Exec("DELETE FROM channels") - require.Nil(t, err, fmt.Sprintf("clean channels unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - _, err := repo.Save(context.Background(), validChannel) - require.Nil(t, err, fmt.Sprintf("save channel unexpected error: %s", err)) - - err = repo.AddConnections(context.Background(), []channels.Connection{validConnection}) - require.Nil(t, err, fmt.Sprintf("add connection unexpected error: %s", err)) - - cases := []struct { - desc string - connection channels.Connection - err error - }{ - { - desc: "check connection successfully", - connection: validConnection, - err: nil, - }, - { - desc: "check connection with non-existent channel", - connection: channels.Connection{ - ClientID: testsutil.GenerateUUID(t), - ChannelID: testsutil.GenerateUUID(t), - DomainID: validChannel.Domain, - Type: connections.Publish, - }, - err: repoerr.ErrNotFound, - }, - { - desc: "check connection with non-existent domain", - connection: channels.Connection{ - ClientID: testsutil.GenerateUUID(t), - ChannelID: validChannel.ID, - DomainID: testsutil.GenerateUUID(t), - Type: connections.Publish, - }, - err: repoerr.ErrNotFound, - }, - { - desc: "check connection with non-existent client", - connection: channels.Connection{ - ClientID: testsutil.GenerateUUID(t), - ChannelID: validChannel.ID, - DomainID: validChannel.Domain, - Type: connections.Publish, - }, - err: repoerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := repo.CheckConnection(context.Background(), tc.connection) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - }) - } -} - -func TestClientAuthorize(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM connections") - require.Nil(t, err, fmt.Sprintf("clean connections unexpected error: %s", err)) - _, err = db.Exec("DELETE FROM channels") - require.Nil(t, err, fmt.Sprintf("clean channels unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - _, err := repo.Save(context.Background(), validChannel) - require.Nil(t, err, fmt.Sprintf("save channel unexpected error: %s", err)) - - err = repo.AddConnections(context.Background(), []channels.Connection{validConnection}) - require.Nil(t, err, fmt.Sprintf("add connection unexpected error: %s", err)) - - cases := []struct { - desc string - connection channels.Connection - err error - }{ - { - desc: "authorize successfully", - connection: validConnection, - err: nil, - }, - { - desc: "authorize with non-existent channel", - connection: channels.Connection{ - ClientID: testsutil.GenerateUUID(t), - ChannelID: testsutil.GenerateUUID(t), - DomainID: validChannel.Domain, - Type: connections.Publish, - }, - err: repoerr.ErrNotFound, - }, - { - desc: "authorize with non-existent client", - connection: channels.Connection{ - ClientID: testsutil.GenerateUUID(t), - ChannelID: validChannel.ID, - DomainID: validChannel.Domain, - Type: connections.Publish, - }, - err: repoerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := repo.ClientAuthorize(context.Background(), tc.connection) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - }) - } -} - -func TestChannelConnectionsCount(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM connections") - require.Nil(t, err, fmt.Sprintf("clean connections unexpected error: %s", err)) - _, err = db.Exec("DELETE FROM channels") - require.Nil(t, err, fmt.Sprintf("clean channels unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - rConnections := []channels.Connection{} - for i := 0; i < 10; i++ { - connection := channels.Connection{ - ClientID: testsutil.GenerateUUID(t), - ChannelID: validChannel.ID, - DomainID: validChannel.Domain, - Type: connections.Publish, - } - rConnections = append(rConnections, connection) - } - - _, err := repo.Save(context.Background(), validChannel) - require.Nil(t, err, fmt.Sprintf("save channel unexpected error: %s", err)) - - err = repo.AddConnections(context.Background(), rConnections) - require.Nil(t, err, fmt.Sprintf("add connection unexpected error: %s", err)) - - cases := []struct { - desc string - channelID string - count uint64 - err error - }{ - { - desc: "get channel connections count successfully", - channelID: validChannel.ID, - count: 10, - err: nil, - }, - { - desc: "get channel connections count with non-existent channel", - channelID: testsutil.GenerateUUID(t), - count: 0, - err: nil, - }, - { - desc: "get channel connections count with empty channel ID", - channelID: "", - count: 0, - err: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - count, err := repo.ChannelConnectionsCount(context.Background(), tc.channelID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.count, count, fmt.Sprintf("%s: expected %d got %d\n", tc.desc, tc.count, count)) - }) - } -} - -func TestDoesChannelHaveConnections(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM connections") - require.Nil(t, err, fmt.Sprintf("clean connections unexpected error: %s", err)) - _, err = db.Exec("DELETE FROM channels") - require.Nil(t, err, fmt.Sprintf("clean channels unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - _, err := repo.Save(context.Background(), validChannel) - require.Nil(t, err, fmt.Sprintf("save channel unexpected error: %s", err)) - - err = repo.AddConnections(context.Background(), []channels.Connection{validConnection}) - require.Nil(t, err, fmt.Sprintf("add connection unexpected error: %s", err)) - - cases := []struct { - desc string - channelID string - has bool - err error - }{ - { - desc: "check if channel has connections successfully", - channelID: validChannel.ID, - has: true, - err: nil, - }, - { - desc: "check if channel has connections with non-existent channel", - channelID: testsutil.GenerateUUID(t), - has: false, - err: nil, - }, - { - desc: "check if channel has connections with empty channel ID", - channelID: "", - has: false, - err: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - has, err := repo.DoesChannelHaveConnections(context.Background(), tc.channelID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.has, has, fmt.Sprintf("%s: expected %t got %t\n", tc.desc, tc.has, has)) - }) - } -} - -func TestRemoveClientConnections(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM connections") - require.Nil(t, err, fmt.Sprintf("clean connections unexpected error: %s", err)) - _, err = db.Exec("DELETE FROM channels") - require.Nil(t, err, fmt.Sprintf("clean channels unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - _, err := repo.Save(context.Background(), validChannel) - require.Nil(t, err, fmt.Sprintf("save channel unexpected error: %s", err)) - - err = repo.AddConnections(context.Background(), []channels.Connection{validConnection}) - require.Nil(t, err, fmt.Sprintf("add connection unexpected error: %s", err)) - - cases := []struct { - desc string - clientID string - err error - }{ - { - desc: "remove client connections successfully", - clientID: validConnection.ClientID, - err: nil, - }, - { - desc: "remove client connections with non-existent client", - clientID: testsutil.GenerateUUID(t), - err: nil, - }, - { - desc: "remove client connections with empty client ID", - clientID: "", - err: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := repo.RemoveClientConnections(context.Background(), tc.clientID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - }) - } -} - -func TestRemoveChannelConnections(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM connections") - require.Nil(t, err, fmt.Sprintf("clean connections unexpected error: %s", err)) - _, err = db.Exec("DELETE FROM channels") - require.Nil(t, err, fmt.Sprintf("clean channels unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - _, err := repo.Save(context.Background(), validChannel) - require.Nil(t, err, fmt.Sprintf("save channel unexpected error: %s", err)) - - err = repo.AddConnections(context.Background(), []channels.Connection{validConnection}) - require.Nil(t, err, fmt.Sprintf("add connection unexpected error: %s", err)) - - cases := []struct { - desc string - channelID string - err error - }{ - { - desc: "remove channel connections successfully", - channelID: validConnection.ChannelID, - err: nil, - }, - { - desc: "remove channel connections with non-existent channel", - channelID: testsutil.GenerateUUID(t), - err: nil, - }, - { - desc: "remove channel connections with empty channel ID", - channelID: "", - err: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := repo.RemoveChannelConnections(context.Background(), tc.channelID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - }) - } -} - -func TestRetrieveParentGroupChannels(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM channels") - require.Nil(t, err, fmt.Sprintf("clean channels unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - var items []channels.Channel - parentID := testsutil.GenerateUUID(t) - for i := 0; i < 10; i++ { - name := namegen.Generate() - channel := channels.Channel{ - ID: testsutil.GenerateUUID(t), - Domain: testsutil.GenerateUUID(t), - ParentGroup: parentID, - Name: name, - Metadata: map[string]any{"name": name}, - CreatedAt: time.Now().UTC().Truncate(time.Microsecond), - Status: channels.EnabledStatus, - ConnectionTypes: []connections.ConnType{}, - } - items = append(items, channel) - } - - _, err := repo.Save(context.Background(), items...) - require.Nil(t, err, fmt.Sprintf("create channel unexpected error: %s", err)) - - cases := []struct { - desc string - parentGroupID string - resp []channels.Channel - err error - }{ - { - desc: "retrieve parent group channels successfully", - parentGroupID: parentID, - resp: items[:10], - err: nil, - }, - { - desc: "retrieve parent group channels with non-existent channel", - parentGroupID: testsutil.GenerateUUID(t), - resp: []channels.Channel(nil), - err: nil, - }, - { - desc: "retrieve parent group channels with empty channel ID", - parentGroupID: "", - resp: []channels.Channel(nil), - err: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - channels, err := repo.RetrieveParentGroupChannels(context.Background(), tc.parentGroupID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - if err == nil { - got := updateTimestamp(channels) - resp := updateTimestamp(tc.resp) - assert.Equal(t, len(tc.resp), len(channels), fmt.Sprintf("%s: expected %d got %d\n", tc.desc, len(tc.resp), len(channels))) - assert.ElementsMatch(t, resp, got, fmt.Sprintf("%s: expected %+v got %+v\n", tc.desc, resp, got)) - } - }) - } -} - -func TestUnsetParentGroupFromChannels(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM channels") - require.Nil(t, err, fmt.Sprintf("clean channels unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - var items []channels.Channel - parentID := testsutil.GenerateUUID(t) - for i := 0; i < 10; i++ { - name := namegen.Generate() - channel := channels.Channel{ - ID: testsutil.GenerateUUID(t), - Domain: testsutil.GenerateUUID(t), - ParentGroup: parentID, - Name: name, - Metadata: map[string]any{"name": name}, - CreatedAt: time.Now().UTC().Truncate(time.Microsecond), - Status: channels.EnabledStatus, - } - items = append(items, channel) - } - - _, err := repo.Save(context.Background(), items...) - require.Nil(t, err, fmt.Sprintf("create channel unexpected error: %s", err)) - - cases := []struct { - desc string - parentGroupID string - err error - }{ - { - desc: "unset parent group from channels successfully", - parentGroupID: parentID, - err: nil, - }, - { - desc: "unset parent group from channels with non-existent id", - parentGroupID: testsutil.GenerateUUID(t), - err: nil, - }, - { - desc: "unset parent group from channels with empty channel ID", - parentGroupID: "", - err: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := repo.UnsetParentGroupFromChannels(context.Background(), tc.parentGroupID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - }) - } -} - -func TestRetrieveByIDWithRoles(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM channels") - require.Nil(t, err, fmt.Sprintf("clean channels unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - nChannels := uint64(10) - - domainID := testsutil.GenerateUUID(t) - userID := testsutil.GenerateUUID(t) - expectedChannels := []channels.Channel{} - for range nChannels { - channel := channels.Channel{ - ID: testsutil.GenerateUUID(t), - Domain: domainID, - Name: namegen.Generate(), - Route: testsutil.GenerateUUID(t), - Tags: namegen.GenerateMultiple(5), - Metadata: map[string]any{ - "department": namegen.Generate(), - }, - Status: channels.EnabledStatus, - CreatedAt: time.Now().UTC().Truncate(time.Microsecond), - } - _, err := repo.Save(context.Background(), channel) - require.Nil(t, err, fmt.Sprintf("add new channel: expected nil got %s\n", err)) - newRolesProvision := []roles.RoleProvision{ - { - Role: roles.Role{ - ID: testsutil.GenerateUUID(t) + "_" + channel.ID, - Name: "admin", - EntityID: channel.ID, - CreatedAt: validTimestamp, - CreatedBy: userID, - }, - OptionalActions: availableActions, - OptionalMembers: []string{userID}, - }, - } - npr, err := repo.AddRoles(context.Background(), newRolesProvision) - require.Nil(t, err, fmt.Sprintf("add roles unexpected error: %s", err)) - expectedChannel := channel - expectedChannel.ConnectionTypes = []connections.ConnType{} - expectedChannel.Roles = []roles.MemberRoleActions{ - { - RoleID: npr[0].Role.ID, - RoleName: npr[0].Role.Name, - Actions: npr[0].OptionalActions, - AccessType: directAccess, - }, - } - expectedChannels = append(expectedChannels, expectedChannel) - } - - cases := []struct { - desc string - channelID string - userID string - response channels.Channel - err error - }{ - { - desc: "retrieve channel with role successfully", - channelID: expectedChannels[0].ID, - userID: userID, - response: expectedChannels[0], - err: nil, - }, - { - desc: "retrieve another channel with role successfully", - channelID: expectedChannels[1].ID, - userID: userID, - response: expectedChannels[1], - err: nil, - }, - { - desc: "retrieve channel with invalid channel id", - channelID: testsutil.GenerateUUID(t), - userID: userID, - response: channels.Channel{}, - err: repoerr.ErrNotFound, - }, - { - desc: "retrieve channel with empty channel id", - channelID: "", - userID: userID, - response: channels.Channel{}, - err: repoerr.ErrNotFound, - }, - { - desc: "retrieve channel with invalid user id", - channelID: expectedChannels[0].ID, - userID: testsutil.GenerateUUID(t), - response: channels.Channel{}, - err: repoerr.ErrNotFound, - }, - { - desc: "retrieve channel with empty user id", - channelID: expectedChannels[0].ID, - userID: "", - response: channels.Channel{}, - err: repoerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - channel, err := repo.RetrieveByIDWithRoles(context.Background(), tc.channelID, tc.userID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("expected %s to contain %s\n", err, tc.err)) - if err == nil { - assert.Equal(t, tc.response, channel, fmt.Sprintf("expected %v got %v\n", tc.response, channel)) - } - }) - } -} - -func TestRetrieveUserChannels(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM channels") - require.Nil(t, err, fmt.Sprintf("clean channels unexpected error: %s", err)) - _, err = db.Exec("DELETE FROM groups") - require.Nil(t, err, fmt.Sprintf("clean groups unexpected error: %s", err)) - _, err = db.Exec("DELETE FROM domains") - require.Nil(t, err, fmt.Sprintf("clean domains unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - nChannels := uint64(10) - - emptyGroupParam := "" - userID := testsutil.GenerateUUID(t) - domainMemberID := testsutil.GenerateUUID(t) - groupMemberID := testsutil.GenerateUUID(t) - clientID := testsutil.GenerateUUID(t) - domain := generateDomain(t, userID, domainMemberID) - group := generateGroup(t, userID, groupMemberID, domain.ID) - groupChannel := channels.Channel{} - parentGroupChannel := channels.Channel{} - connectedChannel := channels.Channel{} - directChannels := []channels.Channel{} - domainChannels := []channels.Channel{} - for i := range nChannels { - channel := channels.Channel{ - ID: testsutil.GenerateUUID(t), - Domain: domain.ID, - Name: namegen.Generate(), - Route: testsutil.GenerateUUID(t), - Tags: namegen.GenerateMultiple(5), - Metadata: map[string]any{ - "department": namegen.Generate(), - }, - Status: channels.EnabledStatus, - CreatedAt: time.Now().UTC().Truncate(time.Microsecond), - } - if i == 1 { - channel.ParentGroup = group.ID - } - _, err := repo.Save(context.Background(), channel) - require.Nil(t, err, fmt.Sprintf("add new channel: expected nil got %s\n", err)) - newRolesProvision := []roles.RoleProvision{ - { - Role: roles.Role{ - ID: testsutil.GenerateUUID(t) + "_" + channel.ID, - Name: "admin", - EntityID: channel.ID, - CreatedAt: validTimestamp, - CreatedBy: userID, - }, - OptionalActions: availableActions, - OptionalMembers: []string{userID}, - }, - } - npr, err := repo.AddRoles(context.Background(), newRolesProvision) - require.Nil(t, err, fmt.Sprintf("add roles unexpected error: %s", err)) - directChannel := channel - directChannel.RoleID = npr[0].Role.ID - directChannel.RoleName = npr[0].Role.Name - directChannel.AccessType = directAccess - directChannel.AccessProviderRoleActions = []string{} - if i == 1 { - directChannel.ParentGroupPath = group.ID - } - directChannels = append(directChannels, directChannel) - if i == 1 { - parentGroupChannel = directChannel - parentGroupChannel.ParentGroupPath = group.ID - channel.ParentGroupPath = group.ID - groupChannel = channel - groupChannel.AccessType = directGroupAccess - groupChannel.AccessProviderId = group.ID - groupChannel.AccessProviderRoleId = group.Roles[0].RoleID - groupChannel.AccessProviderRoleName = group.Roles[0].RoleName - groupChannel.AccessProviderRoleActions = groupAvailableActions - } - if i == 2 { - conn := channels.Connection{ - ClientID: clientID, - ChannelID: channel.ID, - DomainID: channel.Domain, - Type: connections.Publish, - } - err = repo.AddConnections(context.Background(), []channels.Connection{conn}) - assert.Nil(t, err, fmt.Sprintf("add connection unexpected error: %s", err)) - connectedChannel = channel - connectedChannel.RoleID = npr[0].Role.ID - connectedChannel.RoleName = npr[0].Role.Name - connectedChannel.AccessType = directAccess - connectedChannel.AccessProviderRoleActions = []string{} - connectedChannel.ConnectionTypes = []connections.ConnType{connections.Publish} - } - domainChannel := channel - domainChannel.AccessType = domainAccess - domainChannel.AccessProviderId = domain.ID - domainChannel.AccessProviderRoleId = domain.Roles[0].RoleID - domainChannel.AccessProviderRoleName = domain.Roles[0].RoleName - domainChannel.AccessProviderRoleActions = domainAvailableActions - domainChannels = append(domainChannels, domainChannel) - } - - cases := []struct { - desc string - domainID string - userID string - pm channels.Page - response channels.ChannelsPage - err error - }{ - { - desc: "retrieve channels with empty page", - domainID: domain.ID, - userID: userID, - pm: channels.Page{}, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 10, - Offset: 0, - Limit: 0, - }, - Channels: []channels.Channel(nil), - }, - }, - { - desc: "retrieve channels with offset and limit", - domainID: domain.ID, - userID: userID, - pm: channels.Page{ - Offset: 5, - Limit: 10, - Status: channels.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: nChannels, - Offset: 5, - Limit: 10, - }, - Channels: directChannels[5:10], - }, - }, - { - desc: "retrieve channels with member id of parent group with direct group access", - domainID: domain.ID, - userID: groupMemberID, - pm: channels.Page{ - Offset: 0, - Limit: 10, - Status: channels.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 1, - Offset: 0, - Limit: 10, - }, - Channels: []channels.Channel{groupChannel}, - }, - }, - { - desc: "retrieve channels with member id of domain with domain access", - domainID: domain.ID, - userID: domainMemberID, - pm: channels.Page{ - Offset: 0, - Limit: 10, - Status: channels.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 10, - Offset: 0, - Limit: 10, - }, - Channels: domainChannels, - }, - }, - { - desc: "retrieve channels connected to a client", - domainID: domain.ID, - userID: userID, - pm: channels.Page{ - Offset: 0, - Limit: 10, - Client: clientID, - Status: channels.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 1, - Offset: 0, - Limit: 10, - }, - Channels: []channels.Channel{connectedChannel}, - }, - }, - { - desc: "retrieve channels with offset out of range and limit", - domainID: domain.ID, - userID: userID, - pm: channels.Page{ - Offset: 1000, - Limit: 50, - Status: channels.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: nChannels, - Offset: 1000, - Limit: 50, - }, - Channels: []channels.Channel(nil), - }, - }, - { - desc: "retrieve channels with metadata", - domainID: domain.ID, - userID: userID, - pm: channels.Page{ - Offset: 0, - Limit: nChannels, - Metadata: directChannels[0].Metadata, - Status: channels.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 1, - Offset: 0, - Limit: nChannels, - }, - Channels: []channels.Channel{directChannels[0]}, - }, - }, - { - desc: "retrieve channels with wrong metadata", - domainID: domain.ID, - userID: userID, - pm: channels.Page{ - Offset: 0, - Limit: nChannels, - Metadata: map[string]any{ - "faculty": namegen.Generate(), - }, - Status: channels.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 0, - Offset: 0, - Limit: nChannels, - }, - Channels: []channels.Channel(nil), - }, - }, - { - desc: "retrieve channels with invalid metadata", - domainID: domain.ID, - userID: userID, - pm: channels.Page{ - Offset: 0, - Limit: nChannels, - Metadata: map[string]any{ - "faculty": make(chan int), - }, - Status: channels.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: uint64(nChannels), - Offset: 0, - Limit: nChannels, - }, - Channels: []channels.Channel(nil), - }, - err: repoerr.ErrMalformedEntity, - }, - { - desc: "retrieve channels with name", - domainID: domain.ID, - userID: userID, - pm: channels.Page{ - Offset: 0, - Limit: nChannels, - Name: directChannels[0].Name, - Status: channels.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 1, - Offset: 0, - Limit: nChannels, - }, - Channels: []channels.Channel{directChannels[0]}, - }, - }, - { - desc: "retrieve channels with wrong name", - domainID: domain.ID, - userID: userID, - pm: channels.Page{ - Offset: 0, - Limit: nChannels, - Name: namegen.Generate(), - Status: channels.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 0, - Offset: 0, - Limit: nChannels, - }, - Channels: []channels.Channel(nil), - }, - }, - { - desc: "retrieve channels with tag", - domainID: domain.ID, - userID: userID, - pm: channels.Page{ - Offset: 0, - Limit: nChannels, - Tags: channels.TagsQuery{Elements: []string{directChannels[0].Tags[0]}, Operator: channels.OrOp}, - Status: channels.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 1, - Offset: 0, - Limit: uint64(nChannels), - }, - Channels: []channels.Channel{directChannels[0]}, - }, - }, - { - desc: "retrieve channels with wrong tags", - domainID: domain.ID, - userID: userID, - pm: channels.Page{ - Offset: 0, - Limit: nChannels, - Tags: channels.TagsQuery{Elements: []string{namegen.Generate()}, Operator: channels.OrOp}, - Status: channels.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 0, - Offset: 0, - Limit: nChannels, - }, - Channels: []channels.Channel(nil), - }, - }, - { - desc: "retrieve channels with multiple parameters", - domainID: domain.ID, - userID: userID, - pm: channels.Page{ - Offset: 0, - Limit: nChannels, - Metadata: directChannels[0].Metadata, - Name: directChannels[0].Name, - Tags: channels.TagsQuery{Elements: []string{directChannels[0].Tags[0]}, Operator: channels.OrOp}, - Status: channels.AllStatus, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 1, - Offset: 0, - Limit: nChannels, - }, - Channels: []channels.Channel{directChannels[0]}, - }, - }, - { - desc: "retrieve channels with id", - domainID: domain.ID, - userID: userID, - pm: channels.Page{ - Offset: 0, - Limit: nChannels, - ID: directChannels[0].ID, - Status: channels.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 1, - Offset: 0, - Limit: nChannels, - }, - Channels: []channels.Channel{directChannels[0]}, - }, - }, - { - desc: "retrieve channels with wrong id", - domainID: domain.ID, - userID: userID, - pm: channels.Page{ - Offset: 0, - Limit: nChannels, - ID: testsutil.GenerateUUID(t), - Status: channels.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 0, - Offset: 0, - Limit: nChannels, - }, - Channels: []channels.Channel(nil), - }, - }, - { - desc: "retrieve channels with wrong domain id", - domainID: testsutil.GenerateUUID(t), - userID: userID, - pm: channels.Page{ - Offset: 0, - Limit: nChannels, - Status: channels.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 0, - Offset: 0, - Limit: nChannels, - }, - Channels: []channels.Channel(nil), - }, - }, - { - desc: "retrieve channels with wrong user id", - domainID: domain.ID, - userID: testsutil.GenerateUUID(t), - pm: channels.Page{ - Offset: 0, - Limit: nChannels, - Status: channels.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 0, - Offset: 0, - Limit: nChannels, - }, - Channels: []channels.Channel(nil), - }, - }, - { - desc: "retrieve channels with parent group", - domainID: domain.ID, - userID: userID, - pm: channels.Page{ - Offset: 0, - Limit: nChannels, - Group: nullable.Value[string]{ - Value: group.ID, - Valid: true, - }, - Status: channels.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 1, - Offset: 0, - Limit: nChannels, - }, - Channels: []channels.Channel{parentGroupChannel}, - }, - err: nil, - }, - { - desc: "retrieve channels with no parent group", - domainID: domain.ID, - userID: userID, - pm: channels.Page{ - Offset: 0, - Limit: nChannels, - Group: nullable.Value[string]{ - Value: emptyGroupParam, - Valid: true, - }, - Status: channels.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 0, - Offset: 0, - Limit: nChannels, - }, - Channels: []channels.Channel{}, - }, - }, - { - desc: "retrieve channels with access type", - domainID: domain.ID, - userID: domainMemberID, - pm: channels.Page{ - Offset: 0, - Limit: 10, - AccessType: domainAccess, - Status: channels.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 10, - Offset: 0, - Limit: 10, - }, - Channels: domainChannels, - }, - }, - { - desc: "retrieve channels with wrong access type", - domainID: domain.ID, - userID: userID, - pm: channels.Page{ - Offset: 0, - Limit: nChannels, - AccessType: domainAccess, - Status: channels.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 0, - Offset: 0, - Limit: nChannels, - }, - Channels: []channels.Channel{}, - }, - }, - { - desc: "retrieve channels with role ID", - domainID: domain.ID, - userID: userID, - pm: channels.Page{ - Offset: 0, - Limit: nChannels, - RoleID: directChannels[0].RoleID, - Status: channels.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 1, - Offset: 0, - Limit: nChannels, - }, - Channels: []channels.Channel{directChannels[0]}, - }, - }, - { - desc: "retrieve channels with wrong role ID", - domainID: domain.ID, - userID: userID, - pm: channels.Page{ - Offset: 0, - Limit: nChannels, - RoleID: testsutil.GenerateUUID(t), - Status: channels.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 0, - Offset: 0, - Limit: nChannels, - }, - Channels: []channels.Channel(nil), - }, - }, - { - desc: "retrieve channels with role name", - domainID: domain.ID, - userID: userID, - pm: channels.Page{ - Offset: 0, - Limit: 1, - RoleName: directChannels[0].RoleName, - Status: channels.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 10, - Offset: 0, - Limit: 1, - }, - Channels: directChannels[0:1], - }, - }, - { - desc: "retrieve channels with wrong role name", - domainID: domain.ID, - userID: userID, - pm: channels.Page{ - Offset: 0, - Limit: nChannels, - RoleName: namegen.Generate(), - Status: channels.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 0, - Offset: 0, - Limit: nChannels, - }, - Channels: []channels.Channel(nil), - }, - }, - { - desc: "retrieve channels with actions", - domainID: domain.ID, - userID: userID, - pm: channels.Page{ - Offset: 0, - Limit: nChannels, - Actions: availableActions, - Status: channels.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 10, - Offset: 0, - Limit: nChannels, - }, - Channels: directChannels, - }, - }, - { - desc: "retrieve channels with non-matching actions", - domainID: domain.ID, - userID: userID, - pm: channels.Page{ - Offset: 0, - Limit: nChannels, - Actions: []string{"non_existent_action"}, - Status: channels.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 0, - Offset: 0, - Limit: nChannels, - }, - Channels: []channels.Channel(nil), - }, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - page, err := repo.RetrieveUserChannels(context.Background(), tc.domainID, tc.userID, tc.pm) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("expected %s to contain %s\n", err, tc.err)) - if err == nil { - assert.Equal(t, tc.response.Total, page.Total) - assert.Equal(t, tc.response.Offset, page.Offset) - assert.Equal(t, tc.response.Limit, page.Limit) - expected := stripChannelDetails(tc.response.Channels) - got := stripChannelDetails(page.Channels) - assert.ElementsMatch(t, expected, got, fmt.Sprintf("expected %+v got %+v\n", expected, got)) - } - }) - } -} - -func TestSearchChannels(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM channels") - require.Nil(t, err, fmt.Sprintf("clean channels unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - name := namegen.Generate() - - nChannels := uint64(200) - expectedChannels := []channels.Channel{} - baseTime := time.Now().UTC().Truncate(time.Microsecond) - for i := 0; i < int(nChannels); i++ { - channelName := name + strconv.Itoa(i) - channel := channels.Channel{ - ID: testsutil.GenerateUUID(t), - Name: channelName, - Route: testsutil.GenerateUUID(t), - Metadata: map[string]any{}, - Status: channels.EnabledStatus, - CreatedAt: baseTime.Add(time.Duration(i) * time.Microsecond), - } - _, err := repo.Save(context.Background(), channel) - require.Nil(t, err, fmt.Sprintf("save channel unexpected error: %s", err)) - - expectedChannels = append(expectedChannels, channels.Channel{ - ID: channel.ID, - Name: channel.Name, - CreatedAt: channel.CreatedAt, - }) - } - - page, err := repo.RetrieveAll(context.Background(), channels.Page{Offset: 0, Limit: nChannels}) - require.Nil(t, err, fmt.Sprintf("retrieve all channels unexpected error: %s", err)) - assert.Equal(t, nChannels, page.Total) - - cases := []struct { - desc string - page channels.Page - response channels.ChannelsPage - err error - }{ - { - desc: "with empty page", - page: channels.Page{}, - response: channels.ChannelsPage{ - Channels: []channels.Channel(nil), - Page: channels.Page{ - Total: nChannels, - Offset: 0, - Limit: 0, - }, - }, - err: nil, - }, - { - desc: "with offset only", - page: channels.Page{ - Offset: 50, - }, - response: channels.ChannelsPage{ - Channels: []channels.Channel(nil), - Page: channels.Page{ - Total: nChannels, - Offset: 50, - Limit: 0, - }, - }, - err: nil, - }, - { - desc: "with limit only", - page: channels.Page{ - Limit: 10, - Order: "name", - Dir: ascDir, - }, - response: channels.ChannelsPage{ - Channels: expectedChannels[0:10], - Page: channels.Page{ - Total: nChannels, - Offset: 0, - Limit: 10, - }, - }, - err: nil, - }, - { - desc: "retrieve all channels", - page: channels.Page{ - Offset: 0, - Limit: nChannels, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: nChannels, - Offset: 0, - Limit: nChannels, - }, - Channels: expectedChannels, - }, - }, - { - desc: "with offset and limit", - page: channels.Page{ - Offset: 10, - Limit: 10, - Order: "name", - Dir: ascDir, - }, - response: channels.ChannelsPage{ - Channels: expectedChannels[10:20], - Page: channels.Page{ - Total: nChannels, - Offset: 10, - Limit: 10, - }, - }, - err: nil, - }, - { - desc: "with offset out of range and limit", - page: channels.Page{ - Offset: 1000, - Limit: 50, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: nChannels, - Offset: 1000, - Limit: 50, - }, - Channels: []channels.Channel(nil), - }, - }, - { - desc: "with offset and limit out of range", - page: channels.Page{ - Offset: 190, - Limit: 50, - Order: "name", - Dir: ascDir, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: nChannels, - Offset: 190, - Limit: 50, - }, - Channels: expectedChannels[190:200], - }, - }, - { - desc: "with shorter name", - page: channels.Page{ - Name: expectedChannels[0].Name[:4], - Offset: 0, - Limit: 10, - Order: "name", - Dir: ascDir, - }, - response: channels.ChannelsPage{ - Channels: findChannels(expectedChannels, expectedChannels[0].Name[:4], 0, 10), - Page: channels.Page{ - Total: nChannels, - Offset: 0, - Limit: 10, - }, - }, - err: nil, - }, - { - desc: "with longer name", - page: channels.Page{ - Name: expectedChannels[0].Name, - Offset: 0, - Limit: 10, - }, - response: channels.ChannelsPage{ - Channels: []channels.Channel{expectedChannels[0]}, - Page: channels.Page{ - Total: 1, - Offset: 0, - Limit: 10, - }, - }, - err: nil, - }, - { - desc: "with name SQL injected", - page: channels.Page{ - Name: fmt.Sprintf("%s' OR '1'='1", expectedChannels[0].Name[:1]), - Offset: 0, - Limit: 10, - }, - response: channels.ChannelsPage{ - Channels: []channels.Channel(nil), - Page: channels.Page{ - Total: 0, - Offset: 0, - Limit: 10, - }, - }, - err: nil, - }, - { - desc: "with unknown name", - page: channels.Page{ - Name: namegen.Generate(), - Offset: 0, - Limit: 10, - }, - response: channels.ChannelsPage{ - Channels: []channels.Channel(nil), - Page: channels.Page{ - Total: 0, - Offset: 0, - Limit: 10, - }, - }, - err: nil, - }, - { - desc: "with unknown name SQL injected", - page: channels.Page{ - Name: fmt.Sprintf("%s' OR '1'='1", namegen.Generate()), - Offset: 0, - Limit: 10, - }, - response: channels.ChannelsPage{ - Channels: []channels.Channel(nil), - Page: channels.Page{ - Total: 0, - Offset: 0, - Limit: 10, - }, - }, - err: nil, - }, - { - desc: "with name in asc order", - page: channels.Page{ - Order: "name", - Dir: ascDir, - Name: expectedChannels[0].Name[:1], - Offset: 0, - Limit: 10, - }, - response: channels.ChannelsPage{}, - err: nil, - }, - { - desc: "with name in desc order", - page: channels.Page{ - Order: "name", - Dir: descDir, - Name: expectedChannels[0].Name[:1], - Offset: 0, - Limit: 10, - }, - response: channels.ChannelsPage{}, - err: nil, - }, - } - for _, c := range cases { - t.Run(c.desc, func(t *testing.T) { - switch response, err := repo.RetrieveAll(context.Background(), c.page); { - case err == nil: - if c.page.Order != "" && c.page.Dir != "" { - c.response = response - } - assert.Nil(t, err) - assert.Equal(t, c.response.Total, response.Total) - assert.Equal(t, c.response.Limit, response.Limit) - assert.Equal(t, c.response.Offset, response.Offset) - expected := stripChannelDetails(c.response.Channels) - got := stripChannelDetails(response.Channels) - assert.ElementsMatch(t, expected, got) - default: - assert.True(t, errors.Contains(err, c.err), fmt.Sprintf("expected %s to contain %s\n", err, c.err)) - } - }) - } -} - -func updateTimestamp(channels []channels.Channel) []channels.Channel { - for i := range channels { - channels[i].CreatedAt = validTimestamp - } - - return channels -} - -func generateDomain(t *testing.T, userID, memberID string) domains.Domain { - domain := domains.Domain{ - ID: testsutil.GenerateUUID(t), - Route: namegen.Generate(), - Status: domains.EnabledStatus, - CreatedAt: time.Now().UTC().Truncate(time.Microsecond), - CreatedBy: userID, - } - - drepo := dpostgres.NewRepository(database) - _, err := drepo.SaveDomain(context.Background(), domain) - require.Nil(t, err, fmt.Sprintf("add new domain: expected nil got %s\n", err)) - newRolesProvision := []roles.RoleProvision{ - { - Role: roles.Role{ - ID: testsutil.GenerateUUID(t) + "_" + domain.ID, - Name: "admin", - EntityID: domain.ID, - CreatedAt: validTimestamp, - CreatedBy: userID, - }, - OptionalActions: domainAvailableActions, - OptionalMembers: []string{userID, memberID}, - }, - } - _, err = drepo.AddRoles(context.Background(), newRolesProvision) - require.Nil(t, err, fmt.Sprintf("add new role: expected nil got %s\n", err)) - domain.Roles = []roles.MemberRoleActions{ - { - RoleID: newRolesProvision[0].Role.ID, - RoleName: newRolesProvision[0].Role.Name, - Actions: newRolesProvision[0].OptionalActions, - }, - } - - return domain -} - -func generateGroup(t *testing.T, userID, memberID, domainID string) groups.Group { - group := groups.Group{ - ID: testsutil.GenerateUUID(t), - Name: namegen.Generate(), - Domain: domainID, - Status: groups.EnabledStatus, - CreatedAt: time.Now().UTC().Truncate(time.Microsecond), - } - - grepo := gpostgres.New(database) - _, err := grepo.Save(context.Background(), group) - require.Nil(t, err, fmt.Sprintf("add new group: expected nil got %s\n", err)) - newRolesProvision := []roles.RoleProvision{ - { - Role: roles.Role{ - ID: testsutil.GenerateUUID(t) + "_" + group.ID, - Name: "admin", - EntityID: group.ID, - CreatedAt: validTimestamp, - CreatedBy: userID, - }, - OptionalActions: groupAvailableActions, - OptionalMembers: []string{userID, memberID}, - }, - } - _, err = grepo.AddRoles(context.Background(), newRolesProvision) - require.Nil(t, err, fmt.Sprintf("add new role: expected nil got %s\n", err)) - group.Roles = []roles.MemberRoleActions{ - { - RoleID: newRolesProvision[0].Role.ID, - RoleName: newRolesProvision[0].Role.Name, - Actions: newRolesProvision[0].OptionalActions, - }, - } - - return group -} - -func stripChannelDetails(channels []channels.Channel) []channels.Channel { - for i := range channels { - channels[i].CreatedAt = validTimestamp - channels[i].Actions = []string{} - channels[i].Route = "" - if channels[i].Metadata != nil && len(channels[i].Metadata) == 0 { - channels[i].Metadata = nil - } - if channels[i].ConnectionTypes != nil && len(channels[i].ConnectionTypes) == 0 { - channels[i].ConnectionTypes = nil - } - channels[i].AccessProviderRoleActions = []string{} - } - - return channels -} - -func findChannels(chs []channels.Channel, query string, offset, limit uint64) []channels.Channel { - rchannels := []channels.Channel{} - for _, channel := range chs { - if strings.Contains(channel.Name, query) { - rchannels = append(rchannels, channel) - } - } - - if offset > uint64(len(rchannels)) { - return []channels.Channel{} - } - - if limit > uint64(len(rchannels)) { - return rchannels[offset:] - } - - return rchannels[offset:limit] -} - -func verifyChannelsOrdering(t *testing.T, chs []channels.Channel, order, dir string) { - if order == "" || len(chs) <= 1 { - return - } - - switch order { - case "name": - for i := 1; i < len(chs); i++ { - if dir == ascDir { - assert.LessOrEqual(t, chs[i-1].Name, chs[i].Name) - continue - } - assert.GreaterOrEqual(t, chs[i-1].Name, chs[i].Name) - } - case "created_at": - for i := 1; i < len(chs); i++ { - if dir == ascDir { - assert.True(t, !chs[i-1].CreatedAt.After(chs[i].CreatedAt)) - continue - } - assert.True(t, !chs[i-1].CreatedAt.Before(chs[i].CreatedAt)) - } - case "updated_at": - for i := 1; i < len(chs); i++ { - if dir == ascDir { - assert.True(t, !chs[i-1].UpdatedAt.After(chs[i].UpdatedAt)) - continue - } - assert.True(t, !chs[i-1].UpdatedAt.Before(chs[i].UpdatedAt)) - } - } -} diff --git a/channels/postgres/errors.go b/channels/postgres/errors.go deleted file mode 100644 index 00ab0b749..000000000 --- a/channels/postgres/errors.go +++ /dev/null @@ -1,26 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package postgres - -import "github.com/absmach/magistrala/pkg/errors" - -var _ errors.Mapper = (*duplicateErrors)(nil) - -type duplicateErrors struct{} - -// GetError maps constraint names to known errors. -func (d duplicateErrors) GetError(constraint string) (error, bool) { - switch constraint { - case "unique_domain_route_not_null": - return errors.ErrRouteNotAvailable, true - case "channels_pkey": - return errors.NewRequestError("channel id already exists"), true - default: - return nil, false - } -} - -func NewDuplicateErrors() errors.Mapper { - return duplicateErrors{} -} diff --git a/channels/postgres/init.go b/channels/postgres/init.go deleted file mode 100644 index bf99b4b41..000000000 --- a/channels/postgres/init.go +++ /dev/null @@ -1,121 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package postgres - -import ( - gpostgres "github.com/absmach/magistrala/groups/postgres" - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - rolesPostgres "github.com/absmach/magistrala/pkg/roles/repo/postgres" - _ "github.com/jackc/pgx/v5/stdlib" // required for SQL access - migrate "github.com/rubenv/sql-migrate" -) - -func Migration() (*migrate.MemoryMigrationSource, error) { - rolesMigration, err := rolesPostgres.Migration(rolesTableNamePrefix, entityTableName, entityIDColumnName) - if err != nil { - return &migrate.MemoryMigrationSource{}, errors.Wrap(repoerr.ErrRoleMigration, err) - } - channelsMigration := &migrate.MemoryMigrationSource{ - Migrations: []*migrate.Migration{ - { - Id: "channels_01", - // VARCHAR(36) for colums with IDs as UUIDS have a maximum of 36 characters - // STATUS 0 to imply enabled and 1 to imply disabled - Up: []string{ - `CREATE TABLE IF NOT EXISTS channels ( - id VARCHAR(36) PRIMARY KEY, - name VARCHAR(1024), - domain_id VARCHAR(36) NOT NULL, - parent_group_id VARCHAR(36) DEFAULT NULL, - tags TEXT[], - metadata JSONB, - created_by VARCHAR(254), - created_at TIMESTAMP, - updated_at TIMESTAMP, - updated_by VARCHAR(254), - status SMALLINT NOT NULL DEFAULT 0 CHECK (status >= 0), - UNIQUE (id, domain_id), - UNIQUE (domain_id, name) - )`, - `CREATE TABLE IF NOT EXISTS connections ( - channel_id VARCHAR(36), - domain_id VARCHAR(36), - client_id VARCHAR(36), - type SMALLINT NOT NULL CHECK (type IN (1, 2)), - FOREIGN KEY (channel_id, domain_id) REFERENCES channels (id, domain_id) ON DELETE CASCADE ON UPDATE CASCADE, - PRIMARY KEY (channel_id, domain_id, client_id, type) - )`, - }, - Down: []string{ - `DROP TABLE IF EXISTS channels`, - `DROP TABLE IF EXISTS connections`, - }, - }, - { - Id: "channels_02", - Up: []string{ - `ALTER TABLE channels DROP CONSTRAINT IF EXISTS channels_domain_id_name_key`, - }, - Down: []string{ - `ALTER TABLE channels ADD CONSTRAINT channels_domain_id_name_key UNIQUE (domain_id, name)`, - }, - }, - { - Id: "channels_03", - Up: []string{ - `ALTER TABLE channels ADD COLUMN IF NOT EXISTS route VARCHAR(36);`, - `CREATE UNIQUE INDEX IF NOT EXISTS unique_domain_route_not_null ON channels (domain_id, route) WHERE route IS NOT NULL;`, - }, - Down: []string{ - `DROP INDEX IF EXISTS unique_domain_route_not_null;`, - `ALTER TABLE channels DROP COLUMN IF EXISTS route;`, - }, - }, - { - Id: "channels_04", - Up: []string{ - `ALTER TABLE channels ALTER COLUMN created_at TYPE TIMESTAMPTZ;`, - `ALTER TABLE channels ALTER COLUMN updated_at TYPE TIMESTAMPTZ;`, - }, - Down: []string{ - `ALTER TABLE channels ALTER COLUMN created_at TYPE TIMESTAMP;`, - `ALTER TABLE channels ALTER COLUMN updated_at TYPE TIMESTAMP;`, - }, - }, - { - Id: "channels_05", - Up: []string{ - `UPDATE channels - SET metadata = (COALESCE(metadata, '{}'::jsonb) || COALESCE(metadata->'ui', '{}'::jsonb)) - 'ui' - WHERE metadata ? 'ui' AND jsonb_typeof(metadata->'ui') = 'object'`, - }, - Down: []string{ - `SELECT 1`, - }, - }, - { - Id: "channels_06", - Up: []string{ - `CREATE INDEX IF NOT EXISTS idx_channels_domain_id_status ON channels(domain_id, status);`, - `CREATE INDEX IF NOT EXISTS idx_channels_parent_group_id ON channels(parent_group_id);`, - }, - Down: []string{ - `DROP INDEX IF EXISTS idx_channels_domain_id_status;`, - `DROP INDEX IF EXISTS idx_channels_parent_group_id;`, - }, - }, - }, - } - channelsMigration.Migrations = append(channelsMigration.Migrations, rolesMigration.Migrations...) - - groupsMigration, err := gpostgres.Migration() - if err != nil { - return &migrate.MemoryMigrationSource{}, err - } - - channelsMigration.Migrations = append(channelsMigration.Migrations, groupsMigration.Migrations...) - - return channelsMigration, nil -} diff --git a/channels/postgres/setup_test.go b/channels/postgres/setup_test.go deleted file mode 100644 index 24407f190..000000000 --- a/channels/postgres/setup_test.go +++ /dev/null @@ -1,97 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package postgres_test - -import ( - "database/sql" - "fmt" - "log" - "os" - "testing" - "time" - - chpostgres "github.com/absmach/magistrala/channels/postgres" - pgclient "github.com/absmach/magistrala/pkg/postgres" - "github.com/jmoiron/sqlx" - "github.com/ory/dockertest/v3" - "github.com/ory/dockertest/v3/docker" - "go.opentelemetry.io/otel" -) - -var ( - db *sqlx.DB - database pgclient.Database - tracer = otel.Tracer("repo_tests") -) - -func TestMain(m *testing.M) { - pool, err := dockertest.NewPool("") - if err != nil { - log.Fatalf("Could not connect to docker: %s", err) - } - - container, err := pool.RunWithOptions(&dockertest.RunOptions{ - Repository: "postgres", - Tag: "16.2-alpine", - Env: []string{ - "POSTGRES_USER=test", - "POSTGRES_PASSWORD=test", - "POSTGRES_DB=test", - "listen_addresses = '*'", - }, - }, func(config *docker.HostConfig) { - config.AutoRemove = true - config.RestartPolicy = docker.RestartPolicy{Name: "no"} - }) - if err != nil { - log.Fatalf("Could not start container: %s", err) - } - - port := container.GetPort("5432/tcp") - - // exponential backoff-retry, because the application in the container might not be ready to accept connections yet - pool.MaxWait = 120 * time.Second - if err := pool.Retry(func() error { - url := fmt.Sprintf("host=localhost port=%s user=test dbname=test password=test sslmode=disable", port) - db, err := sql.Open("pgx", url) - if err != nil { - return err - } - return db.Ping() - }); err != nil { - log.Fatalf("Could not connect to docker: %s", err) - } - - dbConfig := pgclient.Config{ - Host: "localhost", - Port: port, - User: "test", - Pass: "test", - Name: "test", - SSLMode: "disable", - SSLCert: "", - SSLKey: "", - SSLRootCert: "", - } - - mig, err := chpostgres.Migration() - if err != nil { - log.Fatalf("Could not get groups migration : %s", err) - } - if db, err = pgclient.Setup(dbConfig, *mig); err != nil { - log.Fatalf("Could not setup test DB connection: %s", err) - } - - database = pgclient.NewDatabase(db, dbConfig, tracer) - - code := m.Run() - - // Defers will not be run when using os.Exit - db.Close() - if err := pool.Purge(container); err != nil { - log.Fatalf("Could not purge container: %s", err) - } - - os.Exit(code) -} diff --git a/channels/private/mocks/service.go b/channels/private/mocks/service.go deleted file mode 100644 index f2b1fbba4..000000000 --- a/channels/private/mocks/service.go +++ /dev/null @@ -1,352 +0,0 @@ -// Copyright (c) Abstract Machines - -// SPDX-License-Identifier: Apache-2.0 - -// Code generated by mockery; DO NOT EDIT. -// github.com/vektra/mockery -// template: testify - -package mocks - -import ( - "context" - - "github.com/absmach/magistrala/channels" - mock "github.com/stretchr/testify/mock" -) - -// NewService creates a new instance of Service. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. -// The first argument is typically a *testing.T value. -func NewService(t interface { - mock.TestingT - Cleanup(func()) -}) *Service { - mock := &Service{} - mock.Mock.Test(t) - - t.Cleanup(func() { mock.AssertExpectations(t) }) - - return mock -} - -// Service is an autogenerated mock type for the Service type -type Service struct { - mock.Mock -} - -type Service_Expecter struct { - mock *mock.Mock -} - -func (_m *Service) EXPECT() *Service_Expecter { - return &Service_Expecter{mock: &_m.Mock} -} - -// Authorize provides a mock function for the type Service -func (_mock *Service) Authorize(ctx context.Context, req channels.AuthzReq) error { - ret := _mock.Called(ctx, req) - - if len(ret) == 0 { - panic("no return value specified for Authorize") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, channels.AuthzReq) error); ok { - r0 = returnFunc(ctx, req) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_Authorize_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Authorize' -type Service_Authorize_Call struct { - *mock.Call -} - -// Authorize is a helper method to define mock.On call -// - ctx context.Context -// - req channels.AuthzReq -func (_e *Service_Expecter) Authorize(ctx interface{}, req interface{}) *Service_Authorize_Call { - return &Service_Authorize_Call{Call: _e.mock.On("Authorize", ctx, req)} -} - -func (_c *Service_Authorize_Call) Run(run func(ctx context.Context, req channels.AuthzReq)) *Service_Authorize_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 channels.AuthzReq - if args[1] != nil { - arg1 = args[1].(channels.AuthzReq) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Service_Authorize_Call) Return(err error) *Service_Authorize_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_Authorize_Call) RunAndReturn(run func(ctx context.Context, req channels.AuthzReq) error) *Service_Authorize_Call { - _c.Call.Return(run) - return _c -} - -// RemoveClientConnections provides a mock function for the type Service -func (_mock *Service) RemoveClientConnections(ctx context.Context, clientID string) error { - ret := _mock.Called(ctx, clientID) - - if len(ret) == 0 { - panic("no return value specified for RemoveClientConnections") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) error); ok { - r0 = returnFunc(ctx, clientID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RemoveClientConnections_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveClientConnections' -type Service_RemoveClientConnections_Call struct { - *mock.Call -} - -// RemoveClientConnections is a helper method to define mock.On call -// - ctx context.Context -// - clientID string -func (_e *Service_Expecter) RemoveClientConnections(ctx interface{}, clientID interface{}) *Service_RemoveClientConnections_Call { - return &Service_RemoveClientConnections_Call{Call: _e.mock.On("RemoveClientConnections", ctx, clientID)} -} - -func (_c *Service_RemoveClientConnections_Call) Run(run func(ctx context.Context, clientID string)) *Service_RemoveClientConnections_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Service_RemoveClientConnections_Call) Return(err error) *Service_RemoveClientConnections_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RemoveClientConnections_Call) RunAndReturn(run func(ctx context.Context, clientID string) error) *Service_RemoveClientConnections_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveByID provides a mock function for the type Service -func (_mock *Service) RetrieveByID(ctx context.Context, id string) (channels.Channel, error) { - ret := _mock.Called(ctx, id) - - if len(ret) == 0 { - panic("no return value specified for RetrieveByID") - } - - var r0 channels.Channel - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) (channels.Channel, error)); ok { - return returnFunc(ctx, id) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string) channels.Channel); ok { - r0 = returnFunc(ctx, id) - } else { - r0 = ret.Get(0).(channels.Channel) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = returnFunc(ctx, id) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RetrieveByID_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveByID' -type Service_RetrieveByID_Call struct { - *mock.Call -} - -// RetrieveByID is a helper method to define mock.On call -// - ctx context.Context -// - id string -func (_e *Service_Expecter) RetrieveByID(ctx interface{}, id interface{}) *Service_RetrieveByID_Call { - return &Service_RetrieveByID_Call{Call: _e.mock.On("RetrieveByID", ctx, id)} -} - -func (_c *Service_RetrieveByID_Call) Run(run func(ctx context.Context, id string)) *Service_RetrieveByID_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Service_RetrieveByID_Call) Return(channel channels.Channel, err error) *Service_RetrieveByID_Call { - _c.Call.Return(channel, err) - return _c -} - -func (_c *Service_RetrieveByID_Call) RunAndReturn(run func(ctx context.Context, id string) (channels.Channel, error)) *Service_RetrieveByID_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveIDByRoute provides a mock function for the type Service -func (_mock *Service) RetrieveIDByRoute(ctx context.Context, route string, domainID string) (string, error) { - ret := _mock.Called(ctx, route, domainID) - - if len(ret) == 0 { - panic("no return value specified for RetrieveIDByRoute") - } - - var r0 string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) (string, error)); ok { - return returnFunc(ctx, route, domainID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) string); ok { - r0 = returnFunc(ctx, route, domainID) - } else { - r0 = ret.Get(0).(string) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, string) error); ok { - r1 = returnFunc(ctx, route, domainID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RetrieveIDByRoute_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveIDByRoute' -type Service_RetrieveIDByRoute_Call struct { - *mock.Call -} - -// RetrieveIDByRoute is a helper method to define mock.On call -// - ctx context.Context -// - route string -// - domainID string -func (_e *Service_Expecter) RetrieveIDByRoute(ctx interface{}, route interface{}, domainID interface{}) *Service_RetrieveIDByRoute_Call { - return &Service_RetrieveIDByRoute_Call{Call: _e.mock.On("RetrieveIDByRoute", ctx, route, domainID)} -} - -func (_c *Service_RetrieveIDByRoute_Call) Run(run func(ctx context.Context, route string, domainID string)) *Service_RetrieveIDByRoute_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_RetrieveIDByRoute_Call) Return(s string, err error) *Service_RetrieveIDByRoute_Call { - _c.Call.Return(s, err) - return _c -} - -func (_c *Service_RetrieveIDByRoute_Call) RunAndReturn(run func(ctx context.Context, route string, domainID string) (string, error)) *Service_RetrieveIDByRoute_Call { - _c.Call.Return(run) - return _c -} - -// UnsetParentGroupFromChannels provides a mock function for the type Service -func (_mock *Service) UnsetParentGroupFromChannels(ctx context.Context, parentGroupID string) error { - ret := _mock.Called(ctx, parentGroupID) - - if len(ret) == 0 { - panic("no return value specified for UnsetParentGroupFromChannels") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) error); ok { - r0 = returnFunc(ctx, parentGroupID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_UnsetParentGroupFromChannels_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UnsetParentGroupFromChannels' -type Service_UnsetParentGroupFromChannels_Call struct { - *mock.Call -} - -// UnsetParentGroupFromChannels is a helper method to define mock.On call -// - ctx context.Context -// - parentGroupID string -func (_e *Service_Expecter) UnsetParentGroupFromChannels(ctx interface{}, parentGroupID interface{}) *Service_UnsetParentGroupFromChannels_Call { - return &Service_UnsetParentGroupFromChannels_Call{Call: _e.mock.On("UnsetParentGroupFromChannels", ctx, parentGroupID)} -} - -func (_c *Service_UnsetParentGroupFromChannels_Call) Run(run func(ctx context.Context, parentGroupID string)) *Service_UnsetParentGroupFromChannels_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Service_UnsetParentGroupFromChannels_Call) Return(err error) *Service_UnsetParentGroupFromChannels_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_UnsetParentGroupFromChannels_Call) RunAndReturn(run func(ctx context.Context, parentGroupID string) error) *Service_UnsetParentGroupFromChannels_Call { - _c.Call.Return(run) - return _c -} diff --git a/channels/private/service.go b/channels/private/service.go deleted file mode 100644 index c7de6ac1b..000000000 --- a/channels/private/service.go +++ /dev/null @@ -1,140 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package private - -import ( - "context" - - "github.com/absmach/magistrala/channels" - dom "github.com/absmach/magistrala/domains" - pkgDomains "github.com/absmach/magistrala/pkg/domains" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/pkg/policies" -) - -var errDisabledDomain = errors.New("domain is disabled or frozen") - -type Service interface { - Authorize(ctx context.Context, req channels.AuthzReq) error - UnsetParentGroupFromChannels(ctx context.Context, parentGroupID string) error - RemoveClientConnections(ctx context.Context, clientID string) error - RetrieveByID(ctx context.Context, id string) (channels.Channel, error) - RetrieveIDByRoute(ctx context.Context, route, domainID string) (string, error) -} - -type service struct { - repo channels.Repository - cache channels.Cache - evaluator policies.Evaluator - policy policies.Service - domains pkgDomains.Authorization -} - -var _ Service = (*service)(nil) - -func New(repo channels.Repository, cache channels.Cache, evaluator policies.Evaluator, policy policies.Service, domains pkgDomains.Authorization) Service { - return service{repo, cache, evaluator, policy, domains} -} - -func (svc service) Authorize(ctx context.Context, req channels.AuthzReq) error { - status, err := svc.domains.RetrieveStatus(ctx, req.DomainID) - if err != nil { - return errors.Wrap(svcerr.ErrAuthorization, err) - } - if status != dom.EnabledStatus { - return errors.Wrap(svcerr.ErrAuthorization, errDisabledDomain) - } - switch req.ClientType { - case policies.UserType: - permission, err := req.Type.Permission() - if err != nil { - return err - } - pr := policies.Policy{ - Subject: req.ClientID, - SubjectType: policies.UserType, - Object: req.ChannelID, - Permission: permission, - ObjectType: policies.ChannelType, - } - if err := svc.evaluator.CheckPolicy(ctx, pr); err != nil { - return errors.Wrap(svcerr.ErrAuthorization, err) - } - return nil - case policies.ClientType: - // Optimization: Add cache - if err := svc.repo.ClientAuthorize(ctx, channels.Connection{ - DomainID: req.DomainID, - ChannelID: req.ChannelID, - ClientID: req.ClientID, - Type: req.Type, - }); err != nil { - return errors.Wrap(svcerr.ErrAuthorization, err) - } - return nil - default: - return svcerr.ErrAuthentication - } -} - -func (svc service) RemoveClientConnections(ctx context.Context, clientID string) error { - return svc.repo.RemoveClientConnections(ctx, clientID) -} - -func (svc service) UnsetParentGroupFromChannels(ctx context.Context, parentGroupID string) (retErr error) { - chs, err := svc.repo.RetrieveParentGroupChannels(ctx, parentGroupID) - if err != nil { - return errors.Wrap(svcerr.ErrViewEntity, err) - } - - if len(chs) > 0 { - prs := []policies.Policy{} - for _, ch := range chs { - prs = append(prs, policies.Policy{ - SubjectType: policies.GroupType, - Subject: ch.ParentGroup, - Relation: policies.ParentGroupRelation, - ObjectType: policies.ChannelType, - Object: ch.ID, - }) - } - - if err := svc.policy.DeletePolicies(ctx, prs); err != nil { - return errors.Wrap(svcerr.ErrDeletePolicies, err) - } - defer func() { - if retErr != nil { - if errRollback := svc.policy.AddPolicies(ctx, prs); err != nil { - retErr = errors.Wrap(retErr, errors.Wrap(errors.ErrRollbackTx, errRollback)) - } - } - }() - - if err := svc.repo.UnsetParentGroupFromChannels(ctx, parentGroupID); err != nil { - return errors.Wrap(svcerr.ErrRemoveEntity, err) - } - } - return nil -} - -func (svc service) RetrieveByID(ctx context.Context, id string) (channels.Channel, error) { - return svc.repo.RetrieveByID(ctx, id) -} - -func (svc service) RetrieveIDByRoute(ctx context.Context, route, domainID string) (string, error) { - id, err := svc.cache.ID(ctx, route, domainID) - if err == nil { - return id, nil - } - chn, err := svc.repo.RetrieveByRoute(ctx, route, domainID) - if err != nil { - return "", errors.Wrap(svcerr.ErrViewEntity, err) - } - if err := svc.cache.Save(ctx, route, domainID, chn.ID); err != nil { - return "", errors.Wrap(svcerr.ErrUpdateEntity, err) - } - - return chn.ID, nil -} diff --git a/channels/service.go b/channels/service.go deleted file mode 100644 index beff1a204..000000000 --- a/channels/service.go +++ /dev/null @@ -1,514 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package channels - -import ( - "context" - "fmt" - "time" - - "github.com/absmach/magistrala" - grpcClientsV1 "github.com/absmach/magistrala/api/grpc/clients/v1" - grpcCommonV1 "github.com/absmach/magistrala/api/grpc/common/v1" - grpcGroupsV1 "github.com/absmach/magistrala/api/grpc/groups/v1" - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/connections" - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/pkg/policies" - "github.com/absmach/magistrala/pkg/roles" -) - -var ( - errAddConnectionsClients = errors.New("failed to add connections in clients service") - errRemoveConnectionsClients = errors.New("failed to remove connections from clients service") - errSetParentGroup = errors.New("channel already have parent") - errSetSameParentGroup = errors.New("channel already assigned to the parent group") -) - -type service struct { - repo Repository - cache Cache - policy policies.Service - idProvider magistrala.IDProvider - clients grpcClientsV1.ClientsServiceClient - groups grpcGroupsV1.GroupsServiceClient - roles.ProvisionManageService -} - -var _ Service = (*service)(nil) - -func New(repo Repository, cache Cache, policy policies.Service, idProvider magistrala.IDProvider, clients grpcClientsV1.ClientsServiceClient, groups grpcGroupsV1.GroupsServiceClient, sidProvider magistrala.IDProvider, availableActions []roles.Action, builtInRoles map[roles.BuiltInRoleName][]roles.Action) (Service, error) { - rpms, err := roles.NewProvisionManageService(policies.ChannelType, repo, policy, sidProvider, availableActions, builtInRoles) - if err != nil { - return nil, err - } - - return service{ - repo: repo, - cache: cache, - policy: policy, - idProvider: idProvider, - clients: clients, - groups: groups, - ProvisionManageService: rpms, - }, nil -} - -func (svc service) CreateChannels(ctx context.Context, session authn.Session, chs ...Channel) (retChs []Channel, retRps []roles.RoleProvision, retErr error) { - var reChs []Channel - for _, c := range chs { - if c.ID == "" { - clientID, err := svc.idProvider.ID() - if err != nil { - return []Channel{}, []roles.RoleProvision{}, err - } - c.ID = clientID - } - - if c.Status != DisabledStatus && c.Status != EnabledStatus { - return []Channel{}, []roles.RoleProvision{}, svcerr.ErrInvalidStatus - } - c.Domain = session.DomainID - c.CreatedAt = time.Now().UTC() - reChs = append(reChs, c) - } - - savedChs, err := svc.repo.Save(ctx, reChs...) - if err != nil { - if errors.Contains(err, errors.ErrRouteNotAvailable) { - return []Channel{}, []roles.RoleProvision{}, errors.ErrRouteNotAvailable - } - return []Channel{}, []roles.RoleProvision{}, errors.Wrap(svcerr.ErrCreateEntity, err) - } - chIDs := []string{} - for _, c := range savedChs { - chIDs = append(chIDs, c.ID) - } - - defer func() { - if retErr != nil { - if errRollBack := svc.repo.Remove(ctx, chIDs...); errRollBack != nil { - retErr = errors.Wrap(retErr, errors.Wrap(svcerr.ErrRollbackRepo, errRollBack)) - } - } - }() - - newBuiltInRoleMembers := map[roles.BuiltInRoleName][]roles.Member{ - BuiltInRoleAdmin: {roles.Member(session.UserID)}, - } - - optionalPolicies := []policies.Policy{} - - for _, chID := range chIDs { - optionalPolicies = append(optionalPolicies, - policies.Policy{ - SubjectType: policies.DomainType, - Subject: session.DomainID, - Relation: policies.DomainRelation, - ObjectType: policies.ChannelType, - Object: chID, - }, - ) - } - rp, err := svc.AddNewEntitiesRoles(ctx, session.DomainID, session.UserID, chIDs, optionalPolicies, newBuiltInRoleMembers) - if err != nil { - return []Channel{}, []roles.RoleProvision{}, errors.Wrap(svcerr.ErrAddPolicies, err) - } - return savedChs, rp, nil -} - -func (svc service) UpdateChannel(ctx context.Context, session authn.Session, ch Channel) (Channel, error) { - channel := Channel{ - ID: ch.ID, - Name: ch.Name, - Metadata: ch.Metadata, - UpdatedAt: time.Now().UTC(), - UpdatedBy: session.UserID, - } - channel, err := svc.repo.Update(ctx, channel) - if err != nil { - return Channel{}, errors.Wrap(svcerr.ErrUpdateEntity, err) - } - return channel, nil -} - -func (svc service) UpdateChannelTags(ctx context.Context, session authn.Session, ch Channel) (Channel, error) { - channel := Channel{ - ID: ch.ID, - Tags: ch.Tags, - UpdatedAt: time.Now().UTC(), - UpdatedBy: session.UserID, - } - channel, err := svc.repo.UpdateTags(ctx, channel) - if err != nil { - return Channel{}, errors.Wrap(svcerr.ErrUpdateEntity, err) - } - return channel, nil -} - -func (svc service) EnableChannel(ctx context.Context, session authn.Session, id string) (Channel, error) { - channel := Channel{ - ID: id, - Status: EnabledStatus, - UpdatedAt: time.Now().UTC(), - } - ch, err := svc.changeChannelStatus(ctx, session.UserID, channel) - if err != nil { - return Channel{}, errors.Wrap(ErrEnableChannel, err) - } - - return ch, nil -} - -func (svc service) DisableChannel(ctx context.Context, session authn.Session, id string) (Channel, error) { - channel := Channel{ - ID: id, - Status: DisabledStatus, - UpdatedAt: time.Now().UTC(), - } - ch, err := svc.changeChannelStatus(ctx, session.UserID, channel) - if err != nil { - return Channel{}, errors.Wrap(ErrDisableChannel, err) - } - - return ch, nil -} - -func (svc service) ViewChannel(ctx context.Context, session authn.Session, id string, withRoles bool) (Channel, error) { - var ch Channel - var err error - switch withRoles { - case true: - ch, err = svc.repo.RetrieveByIDWithRoles(ctx, id, session.UserID) - default: - ch, err = svc.repo.RetrieveByID(ctx, id) - } - if err != nil { - return Channel{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - return ch, nil -} - -func (svc service) ListChannels(ctx context.Context, session authn.Session, pm Page) (ChannelsPage, error) { - switch session.SuperAdmin { - case true: - pm.Domain = session.DomainID - cp, err := svc.repo.RetrieveAll(ctx, pm) - if err != nil { - return ChannelsPage{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - return cp, nil - default: - cp, err := svc.repo.RetrieveUserChannels(ctx, session.DomainID, session.UserID, pm) - if err != nil { - return ChannelsPage{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - return cp, nil - } -} - -func (svc service) ListUserChannels(ctx context.Context, session authn.Session, userID string, pm Page) (ChannelsPage, error) { - cp, err := svc.repo.RetrieveUserChannels(ctx, session.DomainID, userID, pm) - if err != nil { - return ChannelsPage{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - return cp, nil -} - -func (svc service) RemoveChannel(ctx context.Context, session authn.Session, id string) error { - ok, err := svc.repo.DoesChannelHaveConnections(ctx, id) - if err != nil { - return errors.Wrap(svcerr.ErrRemoveEntity, err) - } - - if ok { - if _, err := svc.clients.RemoveChannelConnections(ctx, &grpcClientsV1.RemoveChannelConnectionsReq{ChannelId: id}); err != nil { - return errors.Wrap(svcerr.ErrRemoveEntity, err) - } - } - ch, err := svc.repo.ChangeStatus(ctx, Channel{ID: id, Status: DeletedStatus}) - if err != nil { - return errors.Wrap(svcerr.ErrRemoveEntity, err) - } - if ch.Route != "" { - if err := svc.cache.Remove(ctx, ch.Route, ch.Domain); err != nil { - return errors.Wrap(svcerr.ErrRemoveEntity, err) - } - } - - deletePolicies := []policies.Policy{ - { - SubjectType: policies.DomainType, - Subject: session.DomainID, - Relation: policies.DomainRelation, - ObjectType: policies.ChannelType, - Object: id, - }, - } - - if ch.ParentGroup != "" { - deletePolicies = append(deletePolicies, policies.Policy{ - SubjectType: policies.GroupType, - Subject: ch.ParentGroup, - Relation: policies.ParentGroupRelation, - ObjectType: policies.ChannelType, - Object: id, - }) - } - - filterDeletePolicies := []policies.Policy{ - { - SubjectType: policies.ChannelType, - Subject: id, - }, - { - ObjectType: policies.ChannelType, - Object: id, - }, - } - - if err := svc.RemoveEntitiesRoles(ctx, session.DomainID, session.DomainUserID, []string{id}, filterDeletePolicies, deletePolicies); err != nil { - return errors.Wrap(svcerr.ErrDeletePolicies, err) - } - - if err := svc.repo.Remove(ctx, id); err != nil { - return errors.Wrap(svcerr.ErrRemoveEntity, err) - } - - return nil -} - -func (svc service) Connect(ctx context.Context, session authn.Session, chIDs, thIDs []string, connTypes []connections.ConnType) (retErr error) { - for _, chID := range chIDs { - c, err := svc.repo.RetrieveByID(ctx, chID) - if err != nil { - return errors.Wrap(svcerr.ErrCreateEntity, err) - } - if c.Status != EnabledStatus { - return errors.Wrap(svcerr.ErrCreateEntity, fmt.Errorf("channel id %s is not in enabled state", chID)) - } - if c.Domain != session.DomainID { - return errors.Wrap(svcerr.ErrCreateEntity, fmt.Errorf("channel id %s has invalid domain id", chID)) - } - } - - for _, thID := range thIDs { - resp, err := svc.clients.RetrieveEntity(ctx, &grpcCommonV1.RetrieveEntityReq{Id: thID}) - if err != nil { - return errors.Wrap(svcerr.ErrCreateEntity, err) - } - if resp.GetEntity().GetStatus() != uint32(EnabledStatus) { - return errors.Wrap(svcerr.ErrCreateEntity, fmt.Errorf("client id %s is not in enabled state", thID)) - } - if resp.GetEntity().GetDomainId() != session.DomainID { - return errors.Wrap(svcerr.ErrCreateEntity, fmt.Errorf("client id %s has invalid domain id", thID)) - } - } - - conns := []Connection{} - cliConns := []*grpcCommonV1.Connection{} - for _, chID := range chIDs { - for _, thID := range thIDs { - for _, connType := range connTypes { - conns = append(conns, Connection{ - ClientID: thID, - ChannelID: chID, - DomainID: session.DomainID, - Type: connType, - }) - cliConns = append(cliConns, &grpcCommonV1.Connection{ - ClientId: thID, - ChannelId: chID, - DomainId: session.DomainID, - Type: uint32(connType), - }) - } - } - } - for _, conn := range conns { - err := svc.repo.CheckConnection(ctx, conn) - - switch { - case err == nil: - return errors.Wrap(svcerr.ErrConflict, fmt.Errorf("channel %s and client %s are already connected for type %s in domain %s ", conn.ChannelID, conn.ClientID, conn.Type.String(), conn.DomainID)) - case err != repoerr.ErrNotFound: - return errors.Wrap(svcerr.ErrCreateEntity, err) - } - } - if _, err := svc.clients.AddConnections(ctx, &grpcCommonV1.AddConnectionsReq{Connections: cliConns}); err != nil { - return errors.Wrap(svcerr.ErrCreateEntity, errors.Wrap(errAddConnectionsClients, err)) - } - - if err := svc.repo.AddConnections(ctx, conns); err != nil { - return errors.Wrap(svcerr.ErrCreateEntity, err) - } - - return nil -} - -func (svc service) Disconnect(ctx context.Context, session authn.Session, chIDs, thIDs []string, connTypes []connections.ConnType) (retErr error) { - for _, chID := range chIDs { - c, err := svc.repo.RetrieveByID(ctx, chID) - if err != nil { - return errors.Wrap(svcerr.ErrRemoveEntity, err) - } - if c.Domain != session.DomainID { - return errors.Wrap(svcerr.ErrRemoveEntity, fmt.Errorf("channel id %s has invalid domain id", chID)) - } - } - - for _, thID := range thIDs { - resp, err := svc.clients.RetrieveEntity(ctx, &grpcCommonV1.RetrieveEntityReq{Id: thID}) - if err != nil { - return errors.Wrap(svcerr.ErrRemoveEntity, err) - } - - if resp.GetEntity().GetDomainId() != session.DomainID { - return errors.Wrap(svcerr.ErrRemoveEntity, fmt.Errorf("client id %s has invalid domain id", thID)) - } - } - - conns := []Connection{} - thConns := []*grpcCommonV1.Connection{} - for _, chID := range chIDs { - for _, thID := range thIDs { - for _, connType := range connTypes { - conns = append(conns, Connection{ - ClientID: thID, - ChannelID: chID, - DomainID: session.DomainID, - Type: connType, - }) - thConns = append(thConns, &grpcCommonV1.Connection{ - ClientId: thID, - ChannelId: chID, - DomainId: session.DomainID, - Type: uint32(connType), - }) - } - } - } - - if _, err := svc.clients.RemoveConnections(ctx, &grpcCommonV1.RemoveConnectionsReq{Connections: thConns}); err != nil { - return errors.Wrap(svcerr.ErrRemoveEntity, errors.Wrap(errRemoveConnectionsClients, err)) - } - - if err := svc.repo.RemoveConnections(ctx, conns); err != nil { - return errors.Wrap(svcerr.ErrRemoveEntity, err) - } - - return nil -} - -func (svc service) SetParentGroup(ctx context.Context, session authn.Session, parentGroupID string, id string) (retErr error) { - ch, err := svc.repo.RetrieveByID(ctx, id) - if err != nil { - return errors.Wrap(svcerr.ErrUpdateEntity, err) - } - - resp, err := svc.groups.RetrieveEntity(ctx, &grpcCommonV1.RetrieveEntityReq{Id: parentGroupID}) - if err != nil { - return errors.Wrap(svcerr.ErrUpdateEntity, err) - } - if resp.GetEntity().GetDomainId() != session.DomainID { - return errors.Wrap(svcerr.ErrUpdateEntity, fmt.Errorf("parent group id %s has invalid domain id", parentGroupID)) - } - if resp.GetEntity().GetStatus() != uint32(EnabledStatus) { - return errors.Wrap(svcerr.ErrUpdateEntity, fmt.Errorf("parent group id %s is not in enabled state", parentGroupID)) - } - - var pols []policies.Policy - switch ch.ParentGroup { - case parentGroupID: - return errors.Wrap(svcerr.ErrConflict, errSetSameParentGroup) - case "": - // No action needed, proceed to next code after switch - default: - return errors.Wrap(svcerr.ErrConflict, errSetParentGroup) - } - pols = append(pols, policies.Policy{ - Domain: session.DomainID, - SubjectType: policies.GroupType, - Subject: parentGroupID, - Relation: policies.ParentGroupRelation, - ObjectType: policies.ChannelType, - Object: id, - }) - - if err := svc.policy.AddPolicies(ctx, pols); err != nil { - return errors.Wrap(svcerr.ErrAddPolicies, err) - } - defer func() { - if retErr != nil { - if errRollback := svc.policy.DeletePolicies(ctx, pols); errRollback != nil { - retErr = errors.Wrap(retErr, errors.Wrap(apiutil.ErrRollbackTx, errRollback)) - } - } - }() - ch = Channel{ID: id, ParentGroup: parentGroupID, UpdatedBy: session.UserID, UpdatedAt: time.Now().UTC()} - - if err := svc.repo.SetParentGroup(ctx, ch); err != nil { - return errors.Wrap(svcerr.ErrUpdateEntity, err) - } - return nil -} - -func (svc service) RemoveParentGroup(ctx context.Context, session authn.Session, id string) (retErr error) { - ch, err := svc.repo.RetrieveByID(ctx, id) - if err != nil { - return errors.Wrap(svcerr.ErrUpdateEntity, err) - } - - if ch.ParentGroup != "" { - var pols []policies.Policy - pols = append(pols, policies.Policy{ - Domain: session.DomainID, - SubjectType: policies.GroupType, - Subject: ch.ParentGroup, - Relation: policies.ParentGroupRelation, - ObjectType: policies.ChannelType, - Object: id, - }) - - if err := svc.policy.DeletePolicies(ctx, pols); err != nil { - return errors.Wrap(svcerr.ErrDeletePolicies, err) - } - defer func() { - if retErr != nil { - if errRollback := svc.policy.AddPolicies(ctx, pols); errRollback != nil { - retErr = errors.Wrap(retErr, errors.Wrap(apiutil.ErrRollbackTx, errRollback)) - } - } - }() - - ch := Channel{ID: id, UpdatedBy: session.UserID, UpdatedAt: time.Now().UTC()} - - if err := svc.repo.RemoveParentGroup(ctx, ch); err != nil { - return err - } - } - - return nil -} - -func (svc service) changeChannelStatus(ctx context.Context, userID string, channel Channel) (Channel, error) { - dbchannel, err := svc.repo.RetrieveByID(ctx, channel.ID) - if err != nil { - return Channel{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - if dbchannel.Status == channel.Status { - return Channel{}, svcerr.ErrStatusAlreadyAssigned - } - - channel.UpdatedBy = userID - - channel, err = svc.repo.ChangeStatus(ctx, channel) - if err != nil { - return Channel{}, errors.Wrap(svcerr.ErrUpdateEntity, err) - } - return channel, nil -} diff --git a/channels/service_test.go b/channels/service_test.go deleted file mode 100644 index 2fd81973a..000000000 --- a/channels/service_test.go +++ /dev/null @@ -1,1454 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package channels_test - -import ( - "context" - "fmt" - "testing" - "time" - - "github.com/0x6flab/namegenerator" - grpcClientsV1 "github.com/absmach/magistrala/api/grpc/clients/v1" - grpcCommonV1 "github.com/absmach/magistrala/api/grpc/common/v1" - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/channels" - "github.com/absmach/magistrala/channels/mocks" - clmocks "github.com/absmach/magistrala/clients/mocks" - gpmocks "github.com/absmach/magistrala/groups/mocks" - "github.com/absmach/magistrala/internal/testsutil" - "github.com/absmach/magistrala/pkg/authn" - smqauthn "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/connections" - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - policysvc "github.com/absmach/magistrala/pkg/policies" - policymocks "github.com/absmach/magistrala/pkg/policies/mocks" - "github.com/absmach/magistrala/pkg/roles" - "github.com/absmach/magistrala/pkg/uuid" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" -) - -var ( - idProvider = uuid.New() - namegen = namegenerator.NewGenerator() - validChannel = channels.Channel{ - ID: testsutil.GenerateUUID(&testing.T{}), - Name: namegen.Generate(), - Route: namegen.Generate(), - Metadata: map[string]any{ - "key": "value", - }, - Tags: []string{"tag1", "tag2"}, - Domain: testsutil.GenerateUUID(&testing.T{}), - Status: channels.EnabledStatus, - } - validChannelWithRoles = channels.Channel{ - ID: testsutil.GenerateUUID(&testing.T{}), - Name: namegen.Generate(), - Route: namegen.Generate(), - Metadata: map[string]any{ - "key": "value", - }, - Tags: []string{"tag1", "tag2"}, - Domain: testsutil.GenerateUUID(&testing.T{}), - Status: channels.EnabledStatus, - Roles: []roles.MemberRoleActions{ - { - RoleID: "test-id", - RoleName: "test-name", - }, - }, - } - parentGroupID = testsutil.GenerateUUID(&testing.T{}) - validID = testsutil.GenerateUUID(&testing.T{}) - validSession = authn.Session{UserID: validID, DomainID: validID, DomainUserID: validID} -) - -var ( - repo *mocks.Repository - cache *mocks.Cache - policies *policymocks.Service - clientsSvc *clmocks.ClientsServiceClient - groupsSvc *gpmocks.GroupsServiceClient -) - -func newService(t *testing.T) channels.Service { - repo = new(mocks.Repository) - cache = new(mocks.Cache) - policies = new(policymocks.Service) - clientsSvc = new(clmocks.ClientsServiceClient) - groupsSvc = new(gpmocks.GroupsServiceClient) - availableActions := []roles.Action{} - builtInRoles := map[roles.BuiltInRoleName][]roles.Action{ - channels.BuiltInRoleAdmin: availableActions, - } - svc, err := channels.New(repo, cache, policies, idProvider, clientsSvc, groupsSvc, idProvider, availableActions, builtInRoles) - assert.Nil(t, err, fmt.Sprintf(" Unexpected error while creating service %v", err)) - return svc -} - -func TestCreateChannel(t *testing.T) { - svc := newService(t) - - etChan := validChannel - etChan.Route = "" - - cases := []struct { - desc string - channel channels.Channel - saveResp []channels.Channel - saveErr error - deleteErr error - addPoliciesErr error - deletePoliciesErr error - addRoleErr error - err error - }{ - { - desc: "create channel successfully", - channel: validChannel, - saveResp: []channels.Channel{{ - ID: testsutil.GenerateUUID(t), - CreatedAt: time.Now(), - Domain: validID, - }}, - err: nil, - }, - { - desc: "create channel with invalid status", - channel: channels.Channel{ - Name: namegen.Generate(), - Status: channels.Status(100), - }, - err: svcerr.ErrInvalidStatus, - }, - { - desc: "create channel successfully with parent", - channel: channels.Channel{ - Name: namegen.Generate(), - Status: channels.EnabledStatus, - ParentGroup: testsutil.GenerateUUID(t), - }, - saveResp: []channels.Channel{ - { - ID: testsutil.GenerateUUID(t), - CreatedAt: time.Now(), - Domain: testsutil.GenerateUUID(t), - ParentGroup: testsutil.GenerateUUID(t), - }, - }, - err: nil, - }, - { - desc: "create channel with failed to save", - channel: validChannel, - saveResp: []channels.Channel{}, - saveErr: errors.ErrMalformedEntity, - err: errors.ErrMalformedEntity, - }, - { - desc: " create channel with failed to add policies", - channel: validChannel, - saveResp: []channels.Channel{ - { - ID: testsutil.GenerateUUID(t), - CreatedAt: time.Now(), - Domain: validID, - }, - }, - addPoliciesErr: svcerr.ErrAuthorization, - err: svcerr.ErrAddPolicies, - }, - { - desc: " create channel with failed to add policies and failed rollback", - channel: validChannel, - saveResp: []channels.Channel{ - { - ID: testsutil.GenerateUUID(t), - CreatedAt: time.Now(), - Domain: validID, - }, - }, - addPoliciesErr: svcerr.ErrAuthorization, - deleteErr: svcerr.ErrRemoveEntity, - err: svcerr.ErrRollbackRepo, - }, - { - desc: "create channel with failed to add roles", - channel: validChannel, - saveResp: []channels.Channel{ - { - ID: testsutil.GenerateUUID(t), - CreatedAt: time.Now(), - Domain: validID, - }, - }, - addRoleErr: svcerr.ErrCreateEntity, - err: svcerr.ErrAddPolicies, - }, - { - desc: "create channels with failed to add roles and failed to delete policies", - channel: validChannel, - saveResp: []channels.Channel{ - { - ID: testsutil.GenerateUUID(t), - CreatedAt: time.Now(), - Domain: validID, - }, - }, - addRoleErr: svcerr.ErrCreateEntity, - deletePoliciesErr: svcerr.ErrRemoveEntity, - err: svcerr.ErrRemoveEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := repo.On("Save", context.Background(), mock.Anything).Return(tc.saveResp, tc.saveErr) - policyCall := policies.On("AddPolicies", context.Background(), mock.Anything).Return(tc.addPoliciesErr) - policyCall1 := policies.On("DeletePolicies", context.Background(), mock.Anything).Return(tc.deletePoliciesErr) - repoCall1 := repo.On("AddRoles", context.Background(), mock.Anything).Return([]roles.RoleProvision{}, tc.addRoleErr) - repoCall2 := repo.On("Remove", context.Background(), mock.Anything).Return(tc.deleteErr) - _, _, err := svc.CreateChannels(context.Background(), validSession, tc.channel) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("expected error %v but got %v", tc.err, err)) - if err == nil { - ok := repoCall.Parent.AssertCalled(t, "Save", context.Background(), mock.Anything) - assert.True(t, ok, fmt.Sprintf("Save was not called on %s", tc.desc)) - } - repoCall.Unset() - policyCall.Unset() - policyCall1.Unset() - repoCall1.Unset() - repoCall2.Unset() - }) - } -} - -func TestViewChannel(t *testing.T) { - svc := newService(t) - - cases := []struct { - desc string - id string - withRoles bool - repoResp channels.Channel - repoErr error - err error - }{ - { - desc: "view channel successfully", - id: validChannel.ID, - withRoles: false, - repoResp: validChannel, - }, - { - desc: "view channel successfully with roles", - id: validChannelWithRoles.ID, - withRoles: true, - repoResp: validChannelWithRoles, - }, - { - desc: "view channel with failed to retrieve", - id: testsutil.GenerateUUID(t), - withRoles: true, - repoErr: repoerr.ErrNotFound, - err: svcerr.ErrViewEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := repo.On("RetrieveByID", context.Background(), tc.id).Return(tc.repoResp, tc.repoErr) - repoCall1 := repo.On("RetrieveByIDWithRoles", context.Background(), tc.id, validSession.UserID).Return(tc.repoResp, tc.repoErr) - got, err := svc.ViewChannel(context.Background(), validSession, tc.id, tc.withRoles) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("expected error %v to contain %v", err, tc.err)) - if err == nil { - switch tc.withRoles { - case true: - assert.Equal(t, tc.repoResp, got) - assert.NotEmpty(t, got.Roles) - ok := repo.AssertCalled(t, "RetrieveByIDWithRoles", context.Background(), tc.id, validSession.UserID) - assert.True(t, ok, fmt.Sprintf("RetrieveByIDWithRoles was not called on %s", tc.desc)) - default: - assert.Equal(t, tc.repoResp, got) - ok := repo.AssertCalled(t, "RetrieveByID", context.Background(), tc.id) - assert.True(t, ok, fmt.Sprintf("RetrieveByID was not called on %s", tc.desc)) - } - } - repoCall.Unset() - repoCall1.Unset() - }) - } -} - -func TestUpdateChannel(t *testing.T) { - svc := newService(t) - - cases := []struct { - desc string - channel channels.Channel - repoResp channels.Channel - repoErr error - err error - }{ - { - desc: "update channel successfully", - channel: channels.Channel{ - ID: testsutil.GenerateUUID(t), - Name: namegen.Generate(), - Route: namegen.Generate(), - }, - repoResp: validChannel, - }, - { - desc: "update channel with repo error", - channel: channels.Channel{ - ID: testsutil.GenerateUUID(t), - Name: namegen.Generate(), - }, - repoErr: repoerr.ErrNotFound, - err: svcerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := repo.On("Update", context.Background(), mock.Anything).Return(tc.repoResp, tc.repoErr) - got, err := svc.UpdateChannel(context.Background(), validSession, tc.channel) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("expected error %v to contain %v", err, tc.err)) - if err == nil { - assert.Equal(t, tc.repoResp, got) - ok := repo.AssertCalled(t, "Update", context.Background(), mock.Anything) - assert.True(t, ok, fmt.Sprintf("Update was not called on %s", tc.desc)) - } - repoCall.Unset() - }) - } -} - -func TestUpdateChannelTags(t *testing.T) { - svc := newService(t) - - cases := []struct { - desc string - updateReq channels.Channel - repoResp channels.Channel - repoErr error - err error - }{ - { - desc: "update channel tags successfully", - updateReq: channels.Channel{ - ID: testsutil.GenerateUUID(t), - Tags: []string{"tag1", "tag2"}, - }, - repoResp: channels.Channel{ - ID: testsutil.GenerateUUID(t), - Tags: []string{"tag1", "tag2"}, - }, - }, - { - desc: "update channel tags with repo error", - updateReq: channels.Channel{ - ID: testsutil.GenerateUUID(t), - Tags: []string{"tag1", "tag2"}, - }, - repoErr: repoerr.ErrNotFound, - err: svcerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := repo.On("UpdateTags", context.Background(), mock.Anything).Return(tc.repoResp, tc.repoErr) - got, err := svc.UpdateChannelTags(context.Background(), validSession, tc.updateReq) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("expected error %v to contain %v", err, tc.err)) - if err == nil { - assert.Equal(t, tc.repoResp, got) - ok := repo.AssertCalled(t, "UpdateTags", context.Background(), mock.Anything) - assert.True(t, ok, fmt.Sprintf("UpdateTags was not called on %s", tc.desc)) - } - repoCall.Unset() - }) - } -} - -func TestEnableChannel(t *testing.T) { - svc := newService(t) - - cases := []struct { - desc string - id string - retrieveResp channels.Channel - retrieveErr error - changeResp channels.Channel - changeErr error - err error - }{ - { - desc: "enable channel successfully", - id: testsutil.GenerateUUID(t), - retrieveResp: channels.Channel{ - Status: channels.DisabledStatus, - }, - changeResp: validChannel, - }, - { - desc: "enable channel with enabled channel", - id: testsutil.GenerateUUID(t), - retrieveResp: channels.Channel{ - Status: channels.EnabledStatus, - }, - err: svcerr.ErrStatusAlreadyAssigned, - }, - { - desc: "enable channel with retrieve error", - id: testsutil.GenerateUUID(t), - retrieveResp: channels.Channel{}, - retrieveErr: repoerr.ErrNotFound, - err: repoerr.ErrNotFound, - }, - { - desc: "enable channel with change status error", - id: testsutil.GenerateUUID(t), - retrieveResp: channels.Channel{ - Status: channels.DisabledStatus, - }, - changeErr: repoerr.ErrNotFound, - err: repoerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := repo.On("RetrieveByID", context.Background(), tc.id).Return(tc.retrieveResp, tc.retrieveErr) - repoCall1 := repo.On("ChangeStatus", context.Background(), mock.Anything).Return(tc.changeResp, tc.changeErr) - got, err := svc.EnableChannel(context.Background(), validSession, tc.id) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("expected error %v to contain %v", err, tc.err)) - if err == nil { - assert.Equal(t, tc.changeResp, got) - ok := repo.AssertCalled(t, "RetrieveByID", context.Background(), tc.id) - assert.True(t, ok, fmt.Sprintf("RetrieveByID was not called on %s", tc.desc)) - } - repoCall.Unset() - repoCall1.Unset() - }) - } -} - -func TestDisableChannel(t *testing.T) { - svc := newService(t) - - cases := []struct { - desc string - id string - retrieveResp channels.Channel - retrieveErr error - changeResp channels.Channel - changeErr error - err error - }{ - { - desc: "disable channel successfully", - id: testsutil.GenerateUUID(t), - retrieveResp: channels.Channel{ - Status: channels.EnabledStatus, - }, - changeResp: validChannel, - }, - { - desc: "disable channel with disabled channel", - id: testsutil.GenerateUUID(t), - retrieveResp: channels.Channel{ - Status: channels.DisabledStatus, - }, - err: svcerr.ErrStatusAlreadyAssigned, - }, - { - desc: "disable channel with retrieve error", - id: testsutil.GenerateUUID(t), - retrieveResp: channels.Channel{}, - retrieveErr: repoerr.ErrNotFound, - err: repoerr.ErrNotFound, - }, - { - desc: "disable channel with change status error", - id: testsutil.GenerateUUID(t), - retrieveResp: channels.Channel{Status: channels.EnabledStatus}, - changeErr: repoerr.ErrNotFound, - err: repoerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := repo.On("RetrieveByID", context.Background(), tc.id).Return(tc.retrieveResp, tc.retrieveErr) - repoCall1 := repo.On("ChangeStatus", context.Background(), mock.Anything).Return(tc.changeResp, tc.changeErr) - got, err := svc.DisableChannel(context.Background(), validSession, tc.id) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("expected error %v to contain %v", err, tc.err)) - if err == nil { - assert.Equal(t, tc.changeResp, got) - ok := repo.AssertCalled(t, "RetrieveByID", context.Background(), tc.id) - assert.True(t, ok, fmt.Sprintf("RetrieveByID was not called on %s", tc.desc)) - } - repoCall.Unset() - repoCall1.Unset() - }) - } -} - -func TestListChannels(t *testing.T) { - svc := newService(t) - - adminID := testsutil.GenerateUUID(t) - domainID := testsutil.GenerateUUID(t) - nonAdminID := testsutil.GenerateUUID(t) - - cases := []struct { - desc string - userKind string - session smqauthn.Session - page channels.Page - retrieveAllResponse channels.ChannelsPage - response channels.ChannelsPage - id string - size uint64 - listObjectsErr error - retrieveAllErr error - listPermissionsErr error - err error - }{ - { - desc: "list all channels successfully as non admin", - userKind: "non-admin", - session: smqauthn.Session{UserID: nonAdminID, DomainID: domainID, SuperAdmin: false}, - id: nonAdminID, - page: channels.Page{ - Offset: 0, - Limit: 100, - }, - retrieveAllResponse: channels.ChannelsPage{ - Page: channels.Page{ - Total: 2, - Offset: 0, - Limit: 100, - }, - Channels: []channels.Channel{validChannel, validChannel}, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 2, - Offset: 0, - Limit: 100, - }, - Channels: []channels.Channel{validChannel, validChannel}, - }, - err: nil, - }, - { - desc: "list all channels as non admin with failed to retrieve all", - userKind: "non-admin", - session: smqauthn.Session{UserID: nonAdminID, DomainID: domainID, SuperAdmin: false}, - id: nonAdminID, - page: channels.Page{ - Offset: 0, - Limit: 100, - }, - retrieveAllResponse: channels.ChannelsPage{}, - response: channels.ChannelsPage{}, - retrieveAllErr: repoerr.ErrNotFound, - err: svcerr.ErrNotFound, - }, - { - desc: "list all channels as non admin with failed super admin", - userKind: "non-admin", - session: smqauthn.Session{UserID: nonAdminID, DomainID: domainID, SuperAdmin: false}, - id: nonAdminID, - page: channels.Page{ - Offset: 0, - Limit: 100, - }, - response: channels.ChannelsPage{}, - err: nil, - }, - { - desc: "list all channels as non admin with failed to list objects", - userKind: "non-admin", - id: nonAdminID, - page: channels.Page{ - Offset: 0, - Limit: 100, - }, - retrieveAllErr: repoerr.ErrNotFound, - response: channels.ChannelsPage{}, - listObjectsErr: svcerr.ErrNotFound, - err: svcerr.ErrNotFound, - }, - } - - for _, tc := range cases { - retrieveAllCall := repo.On("RetrieveAll", mock.Anything, mock.Anything).Return(tc.retrieveAllResponse, tc.retrieveAllErr) - retrieveUserClientsCall := repo.On("RetrieveUserChannels", mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(tc.retrieveAllResponse, tc.retrieveAllErr) - page, err := svc.ListChannels(context.Background(), tc.session, tc.page) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.response, page, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.response, page)) - retrieveAllCall.Unset() - retrieveUserClientsCall.Unset() - } - - cases2 := []struct { - desc string - userKind string - session smqauthn.Session - page channels.Page - retrieveAllResponse channels.ChannelsPage - response channels.ChannelsPage - id string - size uint64 - listObjectsErr error - retrieveAllErr error - listPermissionsErr error - err error - }{ - { - desc: "list all clients as admin successfully", - userKind: "admin", - id: adminID, - session: smqauthn.Session{UserID: adminID, DomainID: domainID, SuperAdmin: true}, - page: channels.Page{ - Offset: 0, - Limit: 100, - Domain: domainID, - }, - retrieveAllResponse: channels.ChannelsPage{ - Page: channels.Page{ - Total: 2, - Offset: 0, - Limit: 100, - }, - Channels: []channels.Channel{validChannel, validChannel}, - }, - response: channels.ChannelsPage{ - Page: channels.Page{ - Total: 2, - Offset: 0, - Limit: 100, - }, - Channels: []channels.Channel{validChannel, validChannel}, - }, - err: nil, - }, - { - desc: "list all clients as admin with failed to retrieve all", - userKind: "admin", - id: adminID, - session: smqauthn.Session{UserID: adminID, DomainID: domainID, SuperAdmin: true}, - page: channels.Page{ - Offset: 0, - Limit: 100, - Domain: domainID, - }, - retrieveAllResponse: channels.ChannelsPage{}, - retrieveAllErr: repoerr.ErrNotFound, - err: svcerr.ErrNotFound, - }, - { - desc: "list all clients as admin with failed to list clients", - userKind: "admin", - id: adminID, - session: smqauthn.Session{UserID: adminID, DomainID: domainID, SuperAdmin: true}, - page: channels.Page{ - Offset: 0, - Limit: 100, - Domain: domainID, - }, - retrieveAllResponse: channels.ChannelsPage{}, - retrieveAllErr: repoerr.ErrNotFound, - err: svcerr.ErrNotFound, - }, - } - - for _, tc := range cases2 { - retrieveAllCall := repo.On("RetrieveAll", mock.Anything, mock.Anything).Return(tc.retrieveAllResponse, tc.retrieveAllErr) - page, err := svc.ListChannels(context.Background(), tc.session, tc.page) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.response, page, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.response, page)) - retrieveAllCall.Unset() - } -} - -func TestRemoveChannel(t *testing.T) { - svc := newService(t) - - deletedChannel := validChannel - deletedChannel.Status = channels.DeletedStatus - - channelWithParent := deletedChannel - channelWithParent.ParentGroup = testsutil.GenerateUUID(t) - - cases := []struct { - desc string - id string - connectionsRes bool - connectionsErr error - removeConnectionsErr error - changeStatusRes channels.Channel - changeStatusErr error - deletePoliciesErr error - deletePolicyFilterErr error - removeErr error - err error - }{ - { - desc: "remove channel without connections successfully", - id: validChannel.ID, - connectionsRes: false, - changeStatusRes: deletedChannel, - err: nil, - }, - { - desc: "remove channel with connections successfully", - id: validChannel.ID, - connectionsRes: true, - err: nil, - }, - { - desc: "remove channel with parent group successfully", - id: channelWithParent.ID, - connectionsRes: false, - changeStatusRes: channelWithParent, - err: nil, - }, - { - desc: "remove channel with failed check on connections", - id: validChannel.ID, - connectionsErr: repoerr.ErrNotFound, - err: svcerr.ErrRemoveEntity, - }, - { - desc: "remove channel with failed to remove connections", - id: validChannel.ID, - connectionsRes: true, - removeConnectionsErr: svcerr.ErrAuthorization, - err: svcerr.ErrRemoveEntity, - }, - { - desc: "remove channel with failed to change status", - id: validChannel.ID, - connectionsRes: false, - changeStatusErr: repoerr.ErrNotFound, - err: repoerr.ErrNotFound, - }, - { - desc: "remove channel with failed to delete policies", - id: validChannel.ID, - connectionsRes: false, - changeStatusRes: deletedChannel, - deletePoliciesErr: svcerr.ErrAuthorization, - err: svcerr.ErrDeletePolicies, - }, - { - desc: "remove channel with failed to delete policy filter", - id: validChannel.ID, - connectionsRes: false, - changeStatusRes: deletedChannel, - deletePolicyFilterErr: svcerr.ErrAuthorization, - err: svcerr.ErrDeletePolicies, - }, - { - desc: "remove channel with failed to remove", - id: validChannel.ID, - connectionsRes: false, - changeStatusRes: deletedChannel, - removeErr: repoerr.ErrNotFound, - err: svcerr.ErrRemoveEntity, - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := repo.On("DoesChannelHaveConnections", context.Background(), validChannel.ID).Return(tc.connectionsRes, tc.connectionsErr) - clientsCall := clientsSvc.On("RemoveChannelConnections", context.Background(), &grpcClientsV1.RemoveChannelConnectionsReq{ChannelId: tc.id}).Return(&grpcClientsV1.RemoveChannelConnectionsRes{}, tc.removeConnectionsErr) - repoCall1 := repo.On("ChangeStatus", context.Background(), channels.Channel{ID: tc.id, Status: channels.DeletedStatus}).Return(tc.changeStatusRes, tc.changeStatusErr) - cacheCall := cache.On("Remove", context.Background(), tc.changeStatusRes.Route, tc.changeStatusRes.Domain).Return(nil) - repoCall2 := repo.On("RetrieveEntitiesRolesActionsMembers", context.Background(), []string{tc.id}).Return([]roles.EntityActionRole{}, []roles.EntityMemberRole{}, nil) - policyCall := policies.On("DeletePolicies", context.Background(), mock.Anything).Return(tc.deletePoliciesErr) - policyCall1 := policies.On("DeletePolicyFilter", context.Background(), mock.Anything).Return(tc.deletePolicyFilterErr) - repoCall3 := repoCall.On("Remove", context.Background(), []string{tc.id}).Return(tc.removeErr) - err := svc.RemoveChannel(context.Background(), validSession, tc.id) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("expected error %v to contain %v", err, tc.err)) - repoCall.Unset() - clientsCall.Unset() - repoCall1.Unset() - policyCall.Unset() - policyCall1.Unset() - repoCall2.Unset() - repoCall3.Unset() - cacheCall.Unset() - }) - } -} - -func TestConnect(t *testing.T) { - svc := newService(t) - - validDomainChannel := validChannel - validDomainChannel.Domain = validID - - disabledChannel := validChannel - disabledChannel.Status = channels.DisabledStatus - - cases := []struct { - desc string - channelIDs []string - thingIDs []string - connTypes []connections.ConnType - repoConn channels.Connection - clientsConn []*grpcCommonV1.Connection - retrieveByIDRes channels.Channel - retrieveByIDErr error - retrieveEntityRes *grpcCommonV1.RetrieveEntityRes - retrieveEntityErr error - checkConnErr error - addClientConnectionsErr error - addChannelConnectionsErr error - err error - }{ - { - desc: "connect successfully", - channelIDs: []string{validChannel.ID}, - thingIDs: []string{validID}, - connTypes: []connections.ConnType{connections.Publish}, - retrieveByIDRes: validDomainChannel, - retrieveEntityRes: &grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Id: validID, - DomainId: validID, - Status: uint32(channels.EnabledStatus), - }, - }, - checkConnErr: repoerr.ErrNotFound, - repoConn: channels.Connection{ - ClientID: validID, - ChannelID: validChannel.ID, - DomainID: validID, - Type: connections.Publish, - }, - clientsConn: []*grpcCommonV1.Connection{ - { - ClientId: validID, - ChannelId: validChannel.ID, - DomainId: validID, - Type: uint32(connections.Publish), - }, - }, - err: nil, - }, - { - desc: "connect with failed to retrieve channel", - channelIDs: []string{validChannel.ID}, - thingIDs: []string{validID}, - retrieveByIDRes: channels.Channel{}, - retrieveByIDErr: repoerr.ErrNotFound, - err: svcerr.ErrNotFound, - }, - { - desc: "connect to disabled channel", - channelIDs: []string{disabledChannel.ID}, - thingIDs: []string{validID}, - retrieveByIDRes: disabledChannel, - err: svcerr.ErrCreateEntity, - }, - { - desc: "connect with different domain", - channelIDs: []string{validChannel.ID}, - thingIDs: []string{validID}, - retrieveByIDRes: validChannel, - err: svcerr.ErrCreateEntity, - }, - { - desc: "connect with failed to retrieve entity", - channelIDs: []string{validChannel.ID}, - thingIDs: []string{validID}, - retrieveByIDRes: validDomainChannel, - retrieveEntityRes: &grpcCommonV1.RetrieveEntityRes{}, - retrieveEntityErr: svcerr.ErrAuthorization, - err: svcerr.ErrAuthorization, - }, - { - desc: "connect with disabled client", - channelIDs: []string{validChannel.ID}, - thingIDs: []string{validID}, - retrieveByIDRes: validDomainChannel, - retrieveEntityRes: &grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Id: validID, - DomainId: validID, - Status: uint32(channels.DisabledStatus), - }, - }, - err: svcerr.ErrCreateEntity, - }, - { - desc: "connect with client from different domain", - channelIDs: []string{validChannel.ID}, - thingIDs: []string{validID}, - retrieveByIDRes: validDomainChannel, - retrieveEntityRes: &grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Id: validID, - DomainId: testsutil.GenerateUUID(t), - Status: uint32(channels.EnabledStatus), - }, - }, - err: svcerr.ErrCreateEntity, - }, - { - desc: "connect with existing connection", - channelIDs: []string{validChannel.ID}, - thingIDs: []string{validID}, - connTypes: []connections.ConnType{connections.Publish}, - retrieveByIDRes: validDomainChannel, - retrieveEntityRes: &grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Id: validID, - DomainId: validID, - Status: uint32(channels.EnabledStatus), - }, - }, - repoConn: channels.Connection{ - ClientID: validID, - ChannelID: validChannel.ID, - DomainID: validID, - Type: connections.Publish, - }, - checkConnErr: nil, - err: svcerr.ErrConflict, - }, - { - desc: "connect with failed to check connection", - channelIDs: []string{validChannel.ID}, - thingIDs: []string{validID}, - connTypes: []connections.ConnType{connections.Publish}, - retrieveByIDRes: validDomainChannel, - retrieveEntityRes: &grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Id: validID, - DomainId: validID, - Status: uint32(channels.EnabledStatus), - }, - }, - repoConn: channels.Connection{ - ClientID: validID, - ChannelID: validChannel.ID, - DomainID: validID, - Type: connections.Publish, - }, - checkConnErr: repoerr.ErrMalformedEntity, - err: svcerr.ErrCreateEntity, - }, - { - desc: "connect with failed to add client connections", - channelIDs: []string{validChannel.ID}, - thingIDs: []string{validID}, - connTypes: []connections.ConnType{connections.Publish}, - retrieveByIDRes: validDomainChannel, - retrieveEntityRes: &grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Id: validID, - DomainId: validID, - Status: uint32(channels.EnabledStatus), - }, - }, - repoConn: channels.Connection{ - ClientID: validID, - ChannelID: validChannel.ID, - DomainID: validID, - Type: connections.Publish, - }, - checkConnErr: repoerr.ErrNotFound, - clientsConn: []*grpcCommonV1.Connection{ - { - ClientId: validID, - ChannelId: validChannel.ID, - DomainId: validID, - Type: uint32(connections.Publish), - }, - }, - addClientConnectionsErr: svcerr.ErrAuthorization, - err: svcerr.ErrCreateEntity, - }, - { - desc: "connect with failed to add channel connections", - channelIDs: []string{validChannel.ID}, - thingIDs: []string{validID}, - connTypes: []connections.ConnType{connections.Publish}, - retrieveByIDRes: validDomainChannel, - retrieveEntityRes: &grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Id: validID, - DomainId: validID, - Status: uint32(channels.EnabledStatus), - }, - }, - repoConn: channels.Connection{ - ClientID: validID, - ChannelID: validChannel.ID, - DomainID: validID, - Type: connections.Publish, - }, - checkConnErr: repoerr.ErrNotFound, - clientsConn: []*grpcCommonV1.Connection{ - { - ClientId: validID, - ChannelId: validChannel.ID, - DomainId: validID, - Type: uint32(connections.Publish), - }, - }, - addChannelConnectionsErr: svcerr.ErrAuthorization, - err: svcerr.ErrCreateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := repo.On("RetrieveByID", context.Background(), validChannel.ID).Return(tc.retrieveByIDRes, tc.retrieveByIDErr) - clientsCall := clientsSvc.On("RetrieveEntity", context.Background(), &grpcCommonV1.RetrieveEntityReq{Id: validID}).Return(tc.retrieveEntityRes, tc.retrieveEntityErr) - repoCall1 := repo.On("CheckConnection", context.Background(), tc.repoConn).Return(tc.checkConnErr) - clientsCall1 := clientsSvc.On("AddConnections", context.Background(), &grpcCommonV1.AddConnectionsReq{Connections: tc.clientsConn}).Return(&grpcCommonV1.AddConnectionsRes{}, tc.addClientConnectionsErr) - repoCall2 := repo.On("AddConnections", context.Background(), []channels.Connection{tc.repoConn}).Return(tc.addChannelConnectionsErr) - err := svc.Connect(context.Background(), validSession, tc.channelIDs, tc.thingIDs, tc.connTypes) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("expected error %v to contain %v", tc.err, err)) - repoCall.Unset() - clientsCall.Unset() - repoCall1.Unset() - clientsCall1.Unset() - repoCall2.Unset() - }) - } -} - -func TestDisconnect(t *testing.T) { - svc := newService(t) - - validDomainChannel := validChannel - validDomainChannel.Domain = validID - - cases := []struct { - desc string - channelIDs []string - thingIDs []string - connTypes []connections.ConnType - repoConn channels.Connection - clientsConn []*grpcCommonV1.Connection - retrieveByIDRes channels.Channel - retrieveByIDErr error - retrieveEntityRes *grpcCommonV1.RetrieveEntityRes - retrieveEntityErr error - removeClientConnectionsErr error - removeChannelConnectionsErr error - err error - }{ - { - desc: "disconnect successfully", - channelIDs: []string{validChannel.ID}, - thingIDs: []string{validID}, - connTypes: []connections.ConnType{connections.Publish}, - retrieveByIDRes: validDomainChannel, - retrieveEntityRes: &grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Id: validID, - DomainId: validID, - Status: uint32(channels.EnabledStatus), - }, - }, - repoConn: channels.Connection{ - ClientID: validID, - ChannelID: validChannel.ID, - DomainID: validID, - Type: connections.Publish, - }, - clientsConn: []*grpcCommonV1.Connection{ - { - ClientId: validID, - ChannelId: validChannel.ID, - DomainId: validID, - Type: uint32(connections.Publish), - }, - }, - err: nil, - }, - { - desc: "disconnect with failed to retrieve channel", - channelIDs: []string{validChannel.ID}, - thingIDs: []string{validID}, - retrieveByIDRes: channels.Channel{}, - retrieveByIDErr: repoerr.ErrNotFound, - err: svcerr.ErrNotFound, - }, - { - desc: "disconnect with different domain", - channelIDs: []string{validChannel.ID}, - thingIDs: []string{validID}, - retrieveByIDRes: validChannel, - err: svcerr.ErrRemoveEntity, - }, - { - desc: "disconnect with failed to retrieve entity", - channelIDs: []string{validChannel.ID}, - thingIDs: []string{validID}, - retrieveByIDRes: validDomainChannel, - retrieveEntityRes: &grpcCommonV1.RetrieveEntityRes{}, - retrieveEntityErr: svcerr.ErrAuthorization, - err: svcerr.ErrAuthorization, - }, - { - desc: "disconnect with client from different domain", - channelIDs: []string{validChannel.ID}, - thingIDs: []string{validID}, - retrieveByIDRes: validDomainChannel, - retrieveEntityRes: &grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Id: validID, - DomainId: testsutil.GenerateUUID(t), - Status: uint32(channels.EnabledStatus), - }, - }, - err: svcerr.ErrRemoveEntity, - }, - { - desc: "disconnect with failed to remove client connections", - channelIDs: []string{validChannel.ID}, - thingIDs: []string{validID}, - connTypes: []connections.ConnType{connections.Publish}, - retrieveByIDRes: validDomainChannel, - retrieveEntityRes: &grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Id: validID, - DomainId: validID, - Status: uint32(channels.EnabledStatus), - }, - }, - repoConn: channels.Connection{ - ClientID: validID, - ChannelID: validChannel.ID, - DomainID: validID, - Type: connections.Publish, - }, - clientsConn: []*grpcCommonV1.Connection{ - { - ClientId: validID, - ChannelId: validChannel.ID, - DomainId: validID, - Type: uint32(connections.Publish), - }, - }, - removeClientConnectionsErr: svcerr.ErrAuthorization, - err: svcerr.ErrRemoveEntity, - }, - { - desc: "disconnect with failed to remove channel connections", - channelIDs: []string{validChannel.ID}, - thingIDs: []string{validID}, - connTypes: []connections.ConnType{connections.Publish}, - retrieveByIDRes: validDomainChannel, - retrieveEntityRes: &grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Id: validID, - DomainId: validID, - Status: uint32(channels.EnabledStatus), - }, - }, - repoConn: channels.Connection{ - ClientID: validID, - ChannelID: validChannel.ID, - DomainID: validID, - Type: connections.Publish, - }, - clientsConn: []*grpcCommonV1.Connection{ - { - ClientId: validID, - ChannelId: validChannel.ID, - DomainId: validID, - Type: uint32(connections.Publish), - }, - }, - removeChannelConnectionsErr: svcerr.ErrAuthorization, - err: svcerr.ErrRemoveEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := repo.On("RetrieveByID", context.Background(), validChannel.ID).Return(tc.retrieveByIDRes, tc.retrieveByIDErr) - clientsCall := clientsSvc.On("RetrieveEntity", context.Background(), &grpcCommonV1.RetrieveEntityReq{Id: validID}).Return(tc.retrieveEntityRes, tc.retrieveEntityErr) - clientsCall1 := clientsSvc.On("RemoveConnections", context.Background(), &grpcCommonV1.RemoveConnectionsReq{Connections: tc.clientsConn}).Return(&grpcCommonV1.RemoveConnectionsRes{}, tc.removeClientConnectionsErr) - repoCall1 := repo.On("RemoveConnections", context.Background(), []channels.Connection{tc.repoConn}).Return(tc.removeChannelConnectionsErr) - err := svc.Disconnect(context.Background(), validSession, tc.channelIDs, tc.thingIDs, tc.connTypes) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("expected error %v to contain %v", tc.err, err)) - repoCall.Unset() - clientsCall.Unset() - clientsCall1.Unset() - repoCall1.Unset() - }) - } -} - -func TestSetParentGroup(t *testing.T) { - svc := newService(t) - - validDomainChannel := validChannel - validDomainChannel.Domain = validID - - parentedChannel := validChannel - parentedChannel.ParentGroup = testsutil.GenerateUUID(t) - - cases := []struct { - desc string - session authn.Session - parentGroupID string - channelID string - retrieveByIDRes channels.Channel - retrieveByIDErr error - retrieveEntityRes *grpcCommonV1.RetrieveEntityRes - retrieveEntityErr error - addPoliciesErr error - setParentGroupErr error - deletePoliciesErr error - err error - }{ - { - desc: "set parent group successfully", - parentGroupID: parentGroupID, - channelID: validChannel.ID, - retrieveByIDRes: validDomainChannel, - retrieveEntityRes: &grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Id: parentGroupID, - DomainId: validID, - Status: uint32(channels.EnabledStatus), - }, - }, - err: nil, - }, - { - desc: "set parent group with failed to retrieve channel", - parentGroupID: parentGroupID, - channelID: testsutil.GenerateUUID(t), - retrieveByIDRes: channels.Channel{}, - retrieveByIDErr: repoerr.ErrNotFound, - err: svcerr.ErrUpdateEntity, - }, - { - desc: "set parent group with failed to retrieve entity", - parentGroupID: parentGroupID, - channelID: validChannel.ID, - retrieveByIDRes: validDomainChannel, - retrieveEntityRes: &grpcCommonV1.RetrieveEntityRes{}, - retrieveEntityErr: svcerr.ErrAuthorization, - err: svcerr.ErrAuthorization, - }, - { - desc: "set parent group with parent of different domain", - parentGroupID: testsutil.GenerateUUID(t), - channelID: validChannel.ID, - retrieveByIDRes: validDomainChannel, - retrieveEntityRes: &grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Id: parentGroupID, - DomainId: testsutil.GenerateUUID(t), - Status: uint32(channels.EnabledStatus), - }, - }, - err: svcerr.ErrUpdateEntity, - }, - { - desc: "set parent groups with disabled domain", - parentGroupID: parentGroupID, - channelID: validChannel.ID, - retrieveByIDRes: validDomainChannel, - retrieveEntityRes: &grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Id: parentGroupID, - DomainId: validID, - Status: uint32(channels.DisabledStatus), - }, - }, - err: svcerr.ErrUpdateEntity, - }, - { - desc: "set parent group of channel with parent group", - parentGroupID: parentGroupID, - channelID: parentedChannel.ID, - retrieveByIDRes: parentedChannel, - retrieveEntityRes: &grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Id: parentGroupID, - DomainId: validID, - Status: uint32(channels.EnabledStatus), - }, - }, - err: svcerr.ErrConflict, - }, - { - desc: "set parent group with failed to add policies", - parentGroupID: parentGroupID, - channelID: validChannel.ID, - retrieveByIDRes: validDomainChannel, - retrieveEntityRes: &grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Id: parentGroupID, - DomainId: validID, - Status: uint32(channels.EnabledStatus), - }, - }, - addPoliciesErr: svcerr.ErrAuthorization, - err: svcerr.ErrAddPolicies, - }, - { - desc: "set parent group with failed to set parent group", - parentGroupID: parentGroupID, - channelID: validChannel.ID, - retrieveByIDRes: validDomainChannel, - retrieveEntityRes: &grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Id: parentGroupID, - DomainId: validID, - Status: uint32(channels.EnabledStatus), - }, - }, - setParentGroupErr: repoerr.ErrNotFound, - err: repoerr.ErrNotFound, - }, - { - desc: "set parent group with failed to delete policies", - parentGroupID: parentGroupID, - channelID: validChannel.ID, - retrieveByIDRes: validDomainChannel, - retrieveEntityRes: &grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Id: parentGroupID, - DomainId: validID, - Status: uint32(channels.EnabledStatus), - }, - }, - setParentGroupErr: repoerr.ErrNotFound, - deletePoliciesErr: svcerr.ErrAuthorization, - err: apiutil.ErrRollbackTx, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - pols := []policysvc.Policy{ - { - Domain: validSession.DomainID, - SubjectType: policysvc.GroupType, - Subject: tc.parentGroupID, - Relation: policysvc.ParentGroupRelation, - ObjectType: policysvc.ChannelType, - Object: tc.channelID, - }, - } - repoCall := repo.On("RetrieveByID", context.Background(), tc.channelID).Return(tc.retrieveByIDRes, tc.retrieveByIDErr) - groupsCall := groupsSvc.On("RetrieveEntity", context.Background(), &grpcCommonV1.RetrieveEntityReq{Id: tc.parentGroupID}).Return(tc.retrieveEntityRes, tc.retrieveEntityErr) - policyCall := policies.On("AddPolicies", context.Background(), pols).Return(tc.addPoliciesErr) - repoCall1 := repo.On("SetParentGroup", context.Background(), mock.Anything).Return(tc.setParentGroupErr) - policyCall1 := policies.On("DeletePolicies", context.Background(), pols).Return(tc.deletePoliciesErr) - err := svc.SetParentGroup(context.Background(), validSession, tc.parentGroupID, tc.channelID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("expected error %v to contain %v", err, tc.err)) - repoCall.Unset() - groupsCall.Unset() - policyCall.Unset() - repoCall1.Unset() - policyCall1.Unset() - }) - } -} - -func TestRemoveParentGroup(t *testing.T) { - svc := newService(t) - - validDomainChannel := validChannel - validDomainChannel.Domain = validID - - parentedChannel := validChannel - parentedChannel.ParentGroup = testsutil.GenerateUUID(t) - - cases := []struct { - desc string - session authn.Session - channelID string - retrieveByIDRes channels.Channel - retrieveByIDErr error - deletePoliciesErr error - removeParentGroupErr error - addPoliciesErr error - err error - }{ - { - desc: "remove parent group successfully", - channelID: validChannel.ID, - retrieveByIDRes: validDomainChannel, - err: nil, - }, - { - desc: "remove parent group with failed to retrieve channel", - channelID: testsutil.GenerateUUID(t), - retrieveByIDRes: channels.Channel{}, - retrieveByIDErr: repoerr.ErrNotFound, - err: svcerr.ErrUpdateEntity, - }, - { - desc: "remove parent group with failed to delete policies", - channelID: validChannel.ID, - retrieveByIDRes: parentedChannel, - deletePoliciesErr: svcerr.ErrAuthorization, - err: svcerr.ErrDeletePolicies, - }, - { - desc: "remove parent group with failed to remove parent group", - channelID: validChannel.ID, - retrieveByIDRes: parentedChannel, - removeParentGroupErr: repoerr.ErrNotFound, - err: repoerr.ErrNotFound, - }, - { - desc: "remove parent group with failed to add policies", - channelID: validChannel.ID, - retrieveByIDRes: parentedChannel, - removeParentGroupErr: repoerr.ErrNotFound, - addPoliciesErr: svcerr.ErrAuthorization, - err: apiutil.ErrRollbackTx, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - pols := []policysvc.Policy{ - { - Domain: validSession.DomainID, - SubjectType: policysvc.GroupType, - Subject: tc.retrieveByIDRes.ParentGroup, - Relation: policysvc.ParentGroupRelation, - ObjectType: policysvc.ChannelType, - Object: tc.channelID, - }, - } - repoCall := repo.On("RetrieveByID", context.Background(), tc.channelID).Return(tc.retrieveByIDRes, tc.retrieveByIDErr) - policyCall := policies.On("DeletePolicies", context.Background(), pols).Return(tc.deletePoliciesErr) - repoCall1 := repo.On("RemoveParentGroup", context.Background(), mock.Anything).Return(tc.removeParentGroupErr) - policyCall1 := policies.On("AddPolicies", context.Background(), pols).Return(tc.addPoliciesErr) - err := svc.RemoveParentGroup(context.Background(), validSession, tc.channelID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("expected error %v to contain %v", err, tc.err)) - repoCall.Unset() - policyCall.Unset() - repoCall1.Unset() - policyCall1.Unset() - }) - } -} diff --git a/channels/status.go b/channels/status.go deleted file mode 100644 index 5943bb2b4..000000000 --- a/channels/status.go +++ /dev/null @@ -1,94 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package channels - -import ( - "encoding/json" - "strings" - - svcerr "github.com/absmach/magistrala/pkg/errors/service" -) - -// Status represents Channel status. -type Status uint8 - -// Possible Channel status values. -const ( - // EnabledStatus represents enabled Channel. - EnabledStatus Status = iota - // DisabledStatus represents disabled Channel. - DisabledStatus - // DeletedStatus represents deleted Channel. - DeletedStatus - - // AllStatus is used for querying purposes to list channels irrespective - // of their status - both active and inactive. It is never stored in the - // database as the actual Channel status and should always be the largest - // value in this enumeration. - AllStatus -) - -// String representation of the possible status values. -const ( - Disabled = "disabled" - Enabled = "enabled" - Deleted = "deleted" - All = "all" - Unknown = "unknown" -) - -// String converts Channel status to string literal. -func (s Status) String() string { - switch s { - case DisabledStatus: - return Disabled - case EnabledStatus: - return Enabled - case DeletedStatus: - return Deleted - case AllStatus: - return All - default: - return Unknown - } -} - -// ToStatus converts string value to a valid Channel status. -func ToStatus(status string) (Status, error) { - switch status { - case Disabled: - return DisabledStatus, nil - case Enabled: - return EnabledStatus, nil - case Deleted: - return DeletedStatus, nil - case All: - return AllStatus, nil - } - return Status(0), svcerr.ErrInvalidStatus -} - -// Custom Marshaller for Status. -func (s Status) MarshalJSON() ([]byte, error) { - return json.Marshal(s.String()) -} - -func (channel Channel) MarshalJSON() ([]byte, error) { - type Alias Channel - return json.Marshal(&struct { - Alias - Status string `json:"status,omitempty"` - }{ - Alias: (Alias)(channel), - Status: channel.Status.String(), - }) -} - -// Custom Unmarshaler for Status. -func (s *Status) UnmarshalJSON(data []byte) error { - str := strings.Trim(string(data), "\"") - val, err := ToStatus(str) - *s = val - return err -} diff --git a/channels/status_test.go b/channels/status_test.go deleted file mode 100644 index d3de3f6a4..000000000 --- a/channels/status_test.go +++ /dev/null @@ -1,246 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package channels_test - -import ( - "testing" - - "github.com/absmach/magistrala/channels" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/stretchr/testify/assert" -) - -func TestStatusString(t *testing.T) { - cases := []struct { - desc string - status channels.Status - expected string - }{ - { - desc: "Enabled", - status: channels.EnabledStatus, - expected: "enabled", - }, - { - desc: "Disabled", - status: channels.DisabledStatus, - expected: "disabled", - }, - { - desc: "Deleted", - status: channels.DeletedStatus, - expected: "deleted", - }, - { - desc: "All", - status: channels.AllStatus, - expected: "all", - }, - { - desc: "Unknown", - status: channels.Status(100), - expected: "unknown", - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - got := tc.status.String() - assert.Equal(t, tc.expected, got, "String() = %v, expected %v", got, tc.expected) - }) - } -} - -func TestToStatus(t *testing.T) { - cases := []struct { - desc string - status string - expetcted channels.Status - err error - }{ - { - desc: "Enabled", - status: "enabled", - expetcted: channels.EnabledStatus, - err: nil, - }, - { - desc: "Disabled", - status: "disabled", - expetcted: channels.DisabledStatus, - err: nil, - }, - { - desc: "Deleted", - status: "deleted", - expetcted: channels.DeletedStatus, - err: nil, - }, - { - desc: "All", - status: "all", - expetcted: channels.AllStatus, - err: nil, - }, - { - desc: "Unknown", - status: "unknown", - expetcted: channels.Status(0), - err: svcerr.ErrInvalidStatus, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - got, err := channels.ToStatus(tc.status) - assert.Equal(t, tc.err, err, "ToStatus() error = %v, expected %v", err, tc.err) - assert.Equal(t, tc.expetcted, got, "ToStatus() = %v, expected %v", got, tc.expetcted) - }) - } -} - -func TestStatusMarshalJSON(t *testing.T) { - cases := []struct { - desc string - expected []byte - status channels.Status - err error - }{ - { - desc: "Enabled", - expected: []byte(`"enabled"`), - status: channels.EnabledStatus, - err: nil, - }, - { - desc: "Disabled", - expected: []byte(`"disabled"`), - status: channels.DisabledStatus, - err: nil, - }, - { - desc: "Deleted", - expected: []byte(`"deleted"`), - status: channels.DeletedStatus, - err: nil, - }, - { - desc: "All", - expected: []byte(`"all"`), - status: channels.AllStatus, - err: nil, - }, - { - desc: "Unknown", - expected: []byte(`"unknown"`), - status: channels.Status(100), - err: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - got, err := tc.status.MarshalJSON() - assert.Equal(t, tc.err, err, "MarshalJSON() error = %v, expected %v", err, tc.err) - assert.Equal(t, tc.expected, got, "MarshalJSON() = %v, expected %v", got, tc.expected) - }) - } -} - -func TestStatusUnmarshalJSON(t *testing.T) { - cases := []struct { - desc string - expected channels.Status - status []byte - err error - }{ - { - desc: "Enabled", - expected: channels.EnabledStatus, - status: []byte(`"enabled"`), - err: nil, - }, - { - desc: "Disabled", - expected: channels.DisabledStatus, - status: []byte(`"disabled"`), - err: nil, - }, - { - desc: "Deleted", - expected: channels.DeletedStatus, - status: []byte(`"deleted"`), - err: nil, - }, - { - desc: "All", - expected: channels.AllStatus, - status: []byte(`"all"`), - err: nil, - }, - { - desc: "Unknown", - expected: channels.Status(0), - status: []byte(`"unknown"`), - err: svcerr.ErrInvalidStatus, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - var s channels.Status - err := s.UnmarshalJSON(tc.status) - assert.Equal(t, tc.err, err, "UnmarshalJSON() error = %v, expected %v", err, tc.err) - assert.Equal(t, tc.expected, s, "UnmarshalJSON() = %v, expected %v", s, tc.expected) - }) - } -} - -func TestChannelMarshalJSON(t *testing.T) { - cases := []struct { - desc string - expected []byte - user channels.Channel - err error - }{ - { - desc: "Enabled", - expected: []byte(`{"id":"","created_at":"0001-01-01T00:00:00Z","updated_at":"0001-01-01T00:00:00Z","status":"enabled"}`), - user: channels.Channel{Status: channels.EnabledStatus}, - err: nil, - }, - { - desc: "Disabled", - expected: []byte(`{"id":"","created_at":"0001-01-01T00:00:00Z","updated_at":"0001-01-01T00:00:00Z","status":"disabled"}`), - user: channels.Channel{Status: channels.DisabledStatus}, - err: nil, - }, - { - desc: "Deleted", - expected: []byte(`{"id":"","created_at":"0001-01-01T00:00:00Z","updated_at":"0001-01-01T00:00:00Z","status":"deleted"}`), - user: channels.Channel{Status: channels.DeletedStatus}, - err: nil, - }, - { - desc: "All", - expected: []byte(`{"id":"","created_at":"0001-01-01T00:00:00Z","updated_at":"0001-01-01T00:00:00Z","status":"all"}`), - user: channels.Channel{Status: channels.AllStatus}, - err: nil, - }, - { - desc: "Unknown", - expected: []byte(`{"id":"","created_at":"0001-01-01T00:00:00Z","updated_at":"0001-01-01T00:00:00Z","status":"unknown"}`), - user: channels.Channel{Status: channels.Status(100)}, - err: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - got, err := tc.user.MarshalJSON() - assert.Equal(t, tc.err, err, "MarshalJSON() error = %v, expected %v", err, tc.err) - assert.Equal(t, tc.expected, got, "MarshalJSON() = %v, expected %v", string(got), string(tc.expected)) - }) - } -} diff --git a/cli/bootstrap_test.go b/cli/bootstrap_test.go deleted file mode 100644 index 57fbcee77..000000000 --- a/cli/bootstrap_test.go +++ /dev/null @@ -1,872 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package cli_test - -import ( - "encoding/json" - "fmt" - "net/http" - "strings" - "testing" - - "github.com/absmach/magistrala/cli" - "github.com/absmach/magistrala/internal/testsutil" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - mgsdk "github.com/absmach/magistrala/pkg/sdk" - sdkmocks "github.com/absmach/magistrala/pkg/sdk/mocks" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" -) - -var ( - clientID = testsutil.GenerateUUID(&testing.T{}) - channelID = testsutil.GenerateUUID(&testing.T{}) - domainID = testsutil.GenerateUUID(&testing.T{}) - profileID = testsutil.GenerateUUID(&testing.T{}) - bootConfig = mgsdk.BootstrapConfig{ - ID: clientID, - Name: "Test Bootstrap", - ExternalID: "09:6:0:sb:sa", - ExternalKey: "key", - } - bootProfile = mgsdk.BootstrapProfile{ - ID: profileID, - Name: "Test Profile", - Description: "Test profile", - ContentFormat: "go-template", - ContentTemplate: "{\"device_id\":\"{{ .Device.ID }}\"}", - Version: 1, - } - validToken = "validToken" - invalidToken = "invalidToken" - extraArg = "extra-arg" - invalidID = "invalidID" - all = "all" -) - -func TestCreateBootstrapConfigCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - bootCmd := cli.NewBootstrapCmd() - rootCmd := setFlags(bootCmd) - - jsonConfig := fmt.Sprintf("{\"external_id\":\"09:6:0:sb:sa\", \"external_key\":\"key\", \"name\": \"%s\"}", "Test Bootstrap") - invalidJson := fmt.Sprintf("{\"external_id\":\"09:6:0:sb:sa\", \"external_key\":\"key\", \"name\": \"%s\"", "Test Bootstrap") - cases := []struct { - desc string - args []string - logType outputLog - response string - sdkErr errors.SDKError - errLogMessage string - id string - }{ - { - desc: "create bootstrap config successfully", - args: []string{ - jsonConfig, - domainID, - validToken, - }, - logType: createLog, - id: clientID, - response: fmt.Sprintf("\ncreated: %s\n\n", clientID), - }, - { - desc: "create bootstrap config with invald args", - args: []string{ - jsonConfig, - domainID, - validToken, - extraArg, - }, - logType: usageLog, - }, - { - desc: "create bootstrap config with invald json", - args: []string{ - invalidJson, - domainID, - validToken, - }, - sdkErr: errors.NewSDKError(errors.New("unexpected end of JSON input")), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.New("unexpected end of JSON input")), - logType: errLog, - }, - { - desc: "create bootstrap config with invald token", - args: []string{ - jsonConfig, - domainID, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusUnauthorized), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusUnauthorized)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("AddBootstrap", mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(tc.id, tc.sdkErr) - out := executeCommand(t, rootCmd, append([]string{createCmd}, tc.args...)...) - - switch tc.logType { - case createLog: - assert.Equal(t, tc.response, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.response, out)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - sdkCall.Unset() - }) - } -} - -func TestGetBootstrapConfigCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - bootCmd := cli.NewBootstrapCmd() - rootCmd := setFlags(bootCmd) - - var boot mgsdk.BootstrapConfig - var page mgsdk.BootstrapPage - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - page mgsdk.BootstrapPage - boot mgsdk.BootstrapConfig - logType outputLog - errLogMessage string - }{ - { - desc: "get all bootstrap config successfully", - args: []string{ - all, - domainID, - validToken, - }, - page: mgsdk.BootstrapPage{ - PageRes: mgsdk.PageRes{ - Total: 1, - Offset: 0, - Limit: 10, - }, - Configs: []mgsdk.BootstrapConfig{bootConfig}, - }, - logType: entityLog, - }, - { - desc: "get bootstrap config with id", - args: []string{ - channelID, - domainID, - validToken, - }, - logType: entityLog, - boot: bootConfig, - }, - { - desc: "get bootstrap config with invalid args", - args: []string{ - all, - domainID, - validToken, - extraArg, - }, - logType: usageLog, - }, - { - desc: "get all bootstrap config with invalid token", - args: []string{ - all, - domainID, - invalidToken, - }, - logType: errLog, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - }, - { - desc: "get bootstrap config with invalid id", - args: []string{ - invalidID, - domainID, - validToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("ViewBootstrap", mock.Anything, tc.args[0], tc.args[1], tc.args[2]).Return(tc.boot, tc.sdkErr) - sdkCall1 := sdkMock.On("Bootstraps", mock.Anything, mock.Anything, tc.args[1], tc.args[2]).Return(tc.page, tc.sdkErr) - - out := executeCommand(t, rootCmd, append([]string{getCmd}, tc.args...)...) - - switch tc.logType { - case entityLog: - if tc.args[0] == all { - err := json.Unmarshal([]byte(out), &page) - assert.Nil(t, err) - assert.Equal(t, tc.page, page, fmt.Sprintf("%v unexpected response, expected: %v, got: %v", tc.desc, tc.page, page)) - } else { - err := json.Unmarshal([]byte(out), &boot) - assert.Nil(t, err) - assert.Equal(t, tc.boot, boot, fmt.Sprintf("%v unexpected response, expected: %v, got: %v", tc.desc, tc.boot, boot)) - } - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - sdkCall.Unset() - sdkCall1.Unset() - }) - } -} - -func TestRemoveBootstrapConfigCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - bootCmd := cli.NewBootstrapCmd() - rootCmd := setFlags(bootCmd) - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - logType outputLog - errLogMessage string - }{ - { - desc: "remove bootstrap config successfully", - args: []string{ - clientID, - domainID, - validToken, - }, - logType: okLog, - }, - { - desc: "remove bootstrap config with invalid args", - args: []string{ - clientID, - domainID, - validToken, - extraArg, - }, - logType: usageLog, - }, - { - desc: "remove bootstrap config with invalid client id", - args: []string{ - invalidID, - domainID, - validToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - { - desc: "remove bootstrap config with invalid token", - args: []string{ - clientID, - domainID, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("RemoveBootstrap", mock.Anything, tc.args[0], tc.args[1], tc.args[2]).Return(tc.sdkErr) - out := executeCommand(t, rootCmd, append([]string{rmCmd}, tc.args...)...) - - switch tc.logType { - case okLog: - assert.True(t, strings.Contains(out, "ok"), fmt.Sprintf("%s unexpected response: expected success message, got: %v", tc.desc, out)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - sdkCall.Unset() - }) - } -} - -func TestUpdateBootstrapConfigCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - bootCmd := cli.NewBootstrapCmd() - rootCmd := setFlags(bootCmd) - - config := "config" - connection := "connection" - - newConfigJson := "{\"name\" : \"New Bootstrap\"}" - chanIDsJson := fmt.Sprintf("[\"%s\"]", channelID) - cases := []struct { - desc string - args []string - boot mgsdk.BootstrapConfig - sdkErr errors.SDKError - errLogMessage string - logType outputLog - }{ - { - desc: "update bootstrap config successfully", - args: []string{ - config, - newConfigJson, - domainID, - validToken, - }, - logType: okLog, - }, - { - desc: "update bootstrap config with invalid token", - args: []string{ - config, - newConfigJson, - domainID, - invalidToken, - }, - logType: errLog, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - }, - { - desc: "update bootstrap connections successfully", - args: []string{ - connection, - clientID, - chanIDsJson, - domainID, - validToken, - }, - logType: okLog, - }, - { - desc: "update bootstrap connections with invalid json", - args: []string{ - connection, - clientID, - fmt.Sprintf("[\"%s\"", clientID), - domainID, - validToken, - }, - sdkErr: errors.NewSDKError(errors.New("unexpected end of JSON input")), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.New("unexpected end of JSON input")), - logType: errLog, - }, - { - desc: "update bootstrap connections with invalid token", - args: []string{ - connection, - clientID, - chanIDsJson, - domainID, - invalidToken, - }, - logType: errLog, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - }, - { - desc: "update bootstrap certs successfully", - args: []string{ - "certs", - clientID, - "client cert", - "client key", - "ca", - domainID, - validToken, - }, - boot: bootConfig, - logType: entityLog, - }, - { - desc: "update bootstrap certs with invalid token", - args: []string{ - "certs", - clientID, - "client cert", - "client key", - "ca", - domainID, - invalidToken, - }, - logType: errLog, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - }, - { - desc: "update bootstrap config with invalid args", - args: []string{ - newConfigJson, - domainID, - validToken, - }, - logType: usageLog, - }, - { - desc: "update bootstrap config with invalid json", - args: []string{ - config, - "{\"name\" : \"New Bootstrap\"", - domainID, - validToken, - }, - sdkErr: errors.NewSDKError(errors.New("unexpected end of JSON input")), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.New("unexpected end of JSON input")), - logType: errLog, - }, - { - desc: "update bootstrap with invalid args", - args: []string{ - extraArg, - extraArg, - extraArg, - extraArg, - extraArg, - }, - logType: usageLog, - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - var boot mgsdk.BootstrapConfig - sdkCall := sdkMock.On("UpdateBootstrap", mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(tc.sdkErr) - sdkCall1 := sdkMock.On("UpdateBootstrapConnection", mock.Anything, mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(tc.sdkErr) - sdkCall2 := sdkMock.On("UpdateBootstrapCerts", mock.Anything, mock.Anything, mock.Anything, mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(tc.boot, tc.sdkErr) - out := executeCommand(t, rootCmd, append([]string{updCmd}, tc.args...)...) - - switch tc.logType { - case entityLog: - err := json.Unmarshal([]byte(out), &boot) - assert.Nil(t, err) - assert.Equal(t, tc.boot, boot, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.boot, boot)) - case okLog: - assert.True(t, strings.Contains(out, "ok"), fmt.Sprintf("%s unexpected response: expected success message, got: %v", tc.desc, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - } - sdkCall.Unset() - sdkCall1.Unset() - sdkCall2.Unset() - }) - } -} - -func TestWhitelistConfigCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - bootCmd := cli.NewBootstrapCmd() - rootCmd := setFlags(bootCmd) - - jsonConfig := fmt.Sprintf("{\"client_id\": \"%s\", \"status\":%d}", clientID, 1) - - cases := []struct { - desc string - args []string - logType outputLog - errLogMessage string - sdkErr errors.SDKError - }{ - { - desc: "whitelist config successfully", - args: []string{ - jsonConfig, - domainID, - validToken, - }, - logType: okLog, - }, - { - desc: "whitelist config with invalid args", - args: []string{ - jsonConfig, - domainID, - validToken, - extraArg, - }, - logType: usageLog, - }, - { - desc: "whitelist config with invalid json", - args: []string{ - fmt.Sprintf("{\"client_id\": \"%s\", \"status\":%d", clientID, 1), - domainID, - validToken, - }, - sdkErr: errors.NewSDKError(errors.New("unexpected end of JSON input")), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.New("unexpected end of JSON input")), - logType: errLog, - }, - { - desc: "whitelist config with invalid token", - args: []string{ - jsonConfig, - domainID, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusUnauthorized), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusUnauthorized)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("Whitelist", mock.Anything, mock.Anything, mock.Anything, tc.args[1], tc.args[2]).Return(tc.sdkErr) - out := executeCommand(t, rootCmd, append([]string{whitelistCmd}, tc.args...)...) - switch tc.logType { - case okLog: - assert.True(t, strings.Contains(out, "ok"), fmt.Sprintf("%s unexpected response: expected success message, got: %v", tc.desc, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - } - sdkCall.Unset() - }) - } -} - -func TestBootstrapConfigCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - bootCmd := cli.NewBootstrapCmd() - rootCmd := setFlags(bootCmd) - - var boot mgsdk.BootstrapConfig - cryptoKey := "v7aT0HGxJxt2gULzr3RHwf4WIf6DusPp" - invalidKey := "invalid key" - cases := []struct { - desc string - args []string - logType outputLog - errLogMessage string - sdkErr errors.SDKError - boot mgsdk.BootstrapConfig - }{ - { - desc: "bootstrap secure config successfully", - args: []string{ - "secure", - bootConfig.ExternalID, - bootConfig.ExternalKey, - cryptoKey, - }, - boot: bootConfig, - logType: entityLog, - }, - { - desc: "bootstrap config successfully", - args: []string{ - bootConfig.ExternalID, - bootConfig.ExternalKey, - }, - boot: bootConfig, - logType: entityLog, - }, - { - desc: "bootstrap secure config with invalid args", - args: []string{ - cryptoKey, - }, - - logType: usageLog, - }, - { - desc: "bootstrap secure config with invalid key", - args: []string{ - "secure", - bootConfig.ExternalID, - invalidKey, - cryptoKey, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusUnauthorized), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusUnauthorized)), - logType: errLog, - }, - { - desc: "bootstrap config with invalid key", - args: []string{ - bootConfig.ExternalID, - invalidKey, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusUnauthorized), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusUnauthorized)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("BootstrapSecure", mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(tc.boot, tc.sdkErr) - sdkCall1 := sdkMock.On("Bootstrap", mock.Anything, mock.Anything, mock.Anything).Return(tc.boot, tc.sdkErr) - out := executeCommand(t, rootCmd, append([]string{bootStrapCmd}, tc.args...)...) - switch tc.logType { - case entityLog: - err := json.Unmarshal([]byte(out), &boot) - assert.Nil(t, err) - assert.Equal(t, tc.boot, boot, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.boot, boot)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - } - sdkCall.Unset() - sdkCall1.Unset() - }) - } -} - -func TestBootstrapProfilesCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - bootCmd := cli.NewBootstrapCmd() - rootCmd := setFlags(bootCmd) - - profilePayload, err := json.Marshal(bootProfile) - assert.Nil(t, err) - jsonProfile := string(profilePayload) - - cases := []struct { - desc string - args []string - profile mgsdk.BootstrapProfile - page mgsdk.BootstrapProfilesPage - sdkErr errors.SDKError - logType outputLog - errLogMessage string - }{ - { - desc: "create bootstrap profile successfully", - args: []string{ - "create", - jsonProfile, - domainID, - validToken, - }, - profile: bootProfile, - logType: entityLog, - }, - { - desc: "get all bootstrap profiles successfully", - args: []string{ - "get", - all, - domainID, - validToken, - }, - page: mgsdk.BootstrapProfilesPage{ - PageRes: mgsdk.PageRes{ - Total: 1, - Offset: 0, - Limit: 10, - }, - Profiles: []mgsdk.BootstrapProfile{bootProfile}, - }, - logType: entityLog, - }, - { - desc: "view bootstrap profile successfully", - args: []string{ - "get", - profileID, - domainID, - validToken, - }, - profile: bootProfile, - logType: entityLog, - }, - { - desc: "update bootstrap profile successfully", - args: []string{ - "update", - jsonProfile, - domainID, - validToken, - }, - profile: bootProfile, - logType: entityLog, - }, - { - desc: "remove bootstrap profile successfully", - args: []string{ - "remove", - profileID, - domainID, - validToken, - }, - logType: okLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - var gotProfile mgsdk.BootstrapProfile - var gotPage mgsdk.BootstrapProfilesPage - - createCall := sdkMock.On("CreateBootstrapProfile", mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(tc.profile, tc.sdkErr) - listCall := sdkMock.On("BootstrapProfiles", mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(tc.page, tc.sdkErr) - viewCall := sdkMock.On("ViewBootstrapProfile", mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(tc.profile, tc.sdkErr) - updateCall := sdkMock.On("UpdateBootstrapProfile", mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(tc.profile, tc.sdkErr) - removeCall := sdkMock.On("RemoveBootstrapProfile", mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(tc.sdkErr) - - out := executeCommand(t, rootCmd, append([]string{"profiles"}, tc.args...)...) - - switch tc.logType { - case entityLog: - if tc.args[0] == "get" && tc.args[1] == all { - err := json.Unmarshal([]byte(out), &gotPage) - assert.Nil(t, err) - assert.Equal(t, tc.page, gotPage, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.page, gotPage)) - } else { - err := json.Unmarshal([]byte(out), &gotProfile) - assert.Nil(t, err) - assert.Equal(t, tc.profile, gotProfile, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.profile, gotProfile)) - } - case okLog: - assert.True(t, strings.Contains(out, "ok"), fmt.Sprintf("%s unexpected response: expected success message, got: %v", tc.desc, out)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - - createCall.Unset() - listCall.Unset() - viewCall.Unset() - updateCall.Unset() - removeCall.Unset() - }) - } -} - -func TestBootstrapEnrollmentsCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - bootCmd := cli.NewBootstrapCmd() - rootCmd := setFlags(bootCmd) - - bindings := []mgsdk.BootstrapBindingRequest{ - { - Slot: "mqtt_client", - Type: "client", - ResourceID: clientID, - }, - } - snapshots := []mgsdk.BootstrapBindingSnapshot{ - { - ConfigID: clientID, - Slot: "mqtt_client", - Type: "client", - ResourceID: clientID, - }, - } - jsonBindings := fmt.Sprintf("[{\"slot\":\"%s\",\"type\":\"%s\",\"resource_id\":\"%s\"}]", bindings[0].Slot, bindings[0].Type, bindings[0].ResourceID) - - cases := []struct { - desc string - args []string - snapshots []mgsdk.BootstrapBindingSnapshot - sdkErr errors.SDKError - logType outputLog - errLogMessage string - }{ - { - desc: "assign bootstrap profile successfully", - args: []string{ - "assign-profile", - clientID, - profileID, - domainID, - validToken, - }, - logType: okLog, - }, - { - desc: "bind bootstrap resources successfully", - args: []string{ - "bind", - clientID, - jsonBindings, - domainID, - validToken, - }, - logType: okLog, - }, - { - desc: "get bootstrap bindings successfully", - args: []string{ - "get-bindings", - clientID, - domainID, - validToken, - }, - snapshots: snapshots, - logType: entityLog, - }, - { - desc: "refresh bootstrap bindings successfully", - args: []string{ - "refresh-bindings", - clientID, - domainID, - validToken, - }, - logType: okLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - var gotSnapshots []mgsdk.BootstrapBindingSnapshot - - assignCall := sdkMock.On("AssignBootstrapProfile", mock.Anything, mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(tc.sdkErr) - bindCall := sdkMock.On("BindBootstrapResources", mock.Anything, mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(tc.sdkErr) - listCall := sdkMock.On("BootstrapBindings", mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(tc.snapshots, tc.sdkErr) - refreshCall := sdkMock.On("RefreshBootstrapBindings", mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(tc.sdkErr) - - out := executeCommand(t, rootCmd, append([]string{"enrollments"}, tc.args...)...) - - switch tc.logType { - case entityLog: - err := json.Unmarshal([]byte(out), &gotSnapshots) - assert.Nil(t, err) - assert.Equal(t, tc.snapshots, gotSnapshots, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.snapshots, gotSnapshots)) - case okLog: - assert.True(t, strings.Contains(out, "ok"), fmt.Sprintf("%s unexpected response: expected success message, got: %v", tc.desc, out)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - - assignCall.Unset() - bindCall.Unset() - listCall.Unset() - refreshCall.Unset() - }) - } -} diff --git a/cli/certs_test.go b/cli/certs_test.go deleted file mode 100644 index 275618b46..000000000 --- a/cli/certs_test.go +++ /dev/null @@ -1,905 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package cli_test - -import ( - "encoding/json" - "fmt" - "net/http" - "os" - "strings" - "testing" - - "github.com/absmach/magistrala/certs" - "github.com/absmach/magistrala/cli" - "github.com/absmach/magistrala/pkg/errors" - "github.com/absmach/magistrala/pkg/sdk" - sdkmocks "github.com/absmach/magistrala/pkg/sdk/mocks" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" -) - -const ( - revokeCmd = "revoke" - deleteCmd = "delete" - issueCmd = "issue" - renewCmd = "renew" - certsListCmd = "get" - downloadCACmd = "download-ca" - CATokenCmd = "certsToken-ca" - viewCACmd = "view-ca" - filePermission = 0o644 -) - -var ( - serialNumber = "39054620502613157373429341617471746606" - id = "5b4c9ee3-e719-4a0a-9ee5-354932c5e6a4" - commonName = "test-name" - certsToken = "certsToken" - certsDomainID = "domain-id" -) - -func TestIssueCertCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - certCmd := cli.NewCertsCmd() - rootCmd := setFlags(certCmd) - - ipAddrs := "[\"192.168.100.22\"]" - - var cert sdk.Certificate - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - logType outputLog - cert sdk.Certificate - }{ - { - desc: "issue cert successfully", - args: []string{ - id, - commonName, - ipAddrs, - certsDomainID, - certsToken, - }, - logType: entityLog, - cert: sdk.Certificate{SerialNumber: serialNumber}, - }, - { - desc: "issue cert with invalid args", - args: []string{ - id, - ipAddrs, - }, - logType: usageLog, - }, - { - desc: "issue cert failed", - args: []string{ - id, - commonName, - ipAddrs, - certsDomainID, - certsToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(certs.ErrCreateEntity, http.StatusUnprocessableEntity), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(certs.ErrCreateEntity, http.StatusUnprocessableEntity)), - logType: errLog, - }, - { - desc: "issue cert with 6 args", - args: []string{ - id, - commonName, - ipAddrs, - "{\"organization\":[\"organization_name\"]}", - certsDomainID, - certsToken, - }, - logType: entityLog, - cert: sdk.Certificate{SerialNumber: serialNumber}, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - defer func() { - cleanupFiles(t, []string{"cert.pem", "key.pem"}) - }() - sdkCall := sdkMock.On("IssueCert", mock.Anything, mock.Anything, mock.Anything, mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(tc.cert, tc.sdkErr) - out := executeCommand(t, rootCmd, append([]string{issueCmd}, tc.args...)...) - switch tc.logType { - case entityLog: - lines := strings.Split(out, "\n") - var jsonLines []string - var inJSON bool - - for _, line := range lines { - line = strings.TrimSpace(line) - if strings.HasPrefix(line, "{") { - inJSON = true - jsonLines = append(jsonLines, line) - } else if inJSON && strings.HasSuffix(line, "}") { - jsonLines = append(jsonLines, line) - break - } else if inJSON { - jsonLines = append(jsonLines, line) - } - } - - if len(jsonLines) == 0 { - t.Fatalf("No JSON found in output: %s", out) - } - - jsonPart := strings.Join(jsonLines, "") - - err := json.Unmarshal([]byte(jsonPart), &cert) - assert.Nil(t, err) - assert.Equal(t, tc.cert, cert, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.cert, cert)) - assert.True(t, strings.Contains(out, "All certificate files have been saved successfully"), fmt.Sprintf("%s should save files", tc.desc)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - sdkCall.Unset() - }) - } -} - -func TestRevokeCertCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - certCmd := cli.NewCertsCmd() - rootCmd := setFlags(certCmd) - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - logType outputLog - }{ - { - desc: "revoke cert successfully", - args: []string{ - serialNumber, - certsDomainID, - certsToken, - }, - logType: okLog, - }, - { - desc: "revoke cert with invalid args", - args: []string{ - serialNumber, - extraArg, - }, - logType: usageLog, - }, - { - desc: "revoke cert failed", - args: []string{ - serialNumber, - certsDomainID, - certsToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(certs.ErrUpdateEntity, http.StatusUnprocessableEntity), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(certs.ErrUpdateEntity, http.StatusUnprocessableEntity)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("RevokeCert", mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(tc.sdkErr) - out := executeCommand(t, rootCmd, append([]string{revokeCmd}, tc.args...)...) - switch tc.logType { - case okLog: - assert.True(t, strings.Contains(out, "ok"), fmt.Sprintf("%s unexpected response: expected success message, got: %v", tc.desc, out)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - sdkCall.Unset() - }) - } -} - -func TestDeleteCertCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - certCmd := cli.NewCertsCmd() - rootCmd := setFlags(certCmd) - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - logType outputLog - }{ - { - desc: "delete certs successfully", - args: []string{ - id, - certsDomainID, - certsToken, - }, - logType: okLog, - }, - { - desc: "delete certs with invalid args", - args: []string{ - id, - extraArg, - }, - logType: usageLog, - }, - { - desc: "delete certs failed", - args: []string{ - id, - certsDomainID, - certsToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(certs.ErrUpdateEntity, http.StatusUnprocessableEntity), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(certs.ErrUpdateEntity, http.StatusUnprocessableEntity)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("DeleteCert", mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(tc.sdkErr) - out := executeCommand(t, rootCmd, append([]string{deleteCmd}, tc.args...)...) - switch tc.logType { - case okLog: - assert.True(t, strings.Contains(out, "ok"), fmt.Sprintf("%s unexpected response: expected success message, got: %v", tc.desc, out)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - sdkCall.Unset() - }) - } -} - -func TestRenewCertCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - certCmd := cli.NewCertsCmd() - rootCmd := setFlags(certCmd) - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - logType outputLog - }{ - { - desc: "renew cert successfully", - args: []string{ - serialNumber, - certsDomainID, - certsToken, - }, - logType: okLog, - }, - { - desc: "renew cert with invalid args", - args: []string{ - serialNumber, - extraArg, - }, - logType: usageLog, - }, - { - desc: "renew cert failed", - args: []string{ - serialNumber, - certsDomainID, - certsToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(certs.ErrUpdateEntity, http.StatusUnprocessableEntity), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(certs.ErrUpdateEntity, http.StatusUnprocessableEntity)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("RenewCert", mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(sdk.Certificate{}, tc.sdkErr) - out := executeCommand(t, rootCmd, append([]string{renewCmd}, tc.args...)...) - switch tc.logType { - case okLog: - assert.True(t, strings.Contains(out, "ok"), fmt.Sprintf("%s unexpected response: expected success message, got: %v", tc.desc, out)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - sdkCall.Unset() - }) - } -} - -func TestListCertsCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - certCmd := cli.NewCertsCmd() - rootCmd := setFlags(certCmd) - - var page sdk.CertificatePage - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - logType outputLog - page sdk.CertificatePage - }{ - { - desc: "list certs successfully", - args: []string{ - all, - certsDomainID, - certsToken, - }, - logType: entityLog, - page: sdk.CertificatePage{ - Total: 1, - Offset: 0, - Limit: 10, - Certificates: []sdk.Certificate{ - {SerialNumber: serialNumber}, - }, - }, - }, - { - desc: "list certs successfully with entity ID", - args: []string{ - id, - certsDomainID, - certsToken, - }, - logType: entityLog, - page: sdk.CertificatePage{ - Total: 1, - Offset: 0, - Limit: 10, - Certificates: []sdk.Certificate{ - {SerialNumber: serialNumber}, - }, - }, - }, - { - desc: "list certs with invalid args", - args: []string{ - all, - extraArg, - }, - logType: usageLog, - }, - { - desc: "failed list certs with all", - args: []string{ - all, - certsDomainID, - certsToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(certs.ErrViewEntity, http.StatusUnprocessableEntity), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(certs.ErrViewEntity, http.StatusUnprocessableEntity)), - logType: errLog, - }, - { - desc: "failed list certs with entity ID", - args: []string{ - id, - certsDomainID, - certsToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(certs.ErrViewEntity, http.StatusUnprocessableEntity), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(certs.ErrViewEntity, http.StatusUnprocessableEntity)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("ListCerts", mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(tc.page, tc.sdkErr) - out := executeCommand(t, rootCmd, append([]string{certsListCmd}, tc.args...)...) - - switch tc.logType { - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - case entityLog: - err := json.Unmarshal([]byte(out), &page) - if err != nil { - t.Fatalf("Failed to unmarshal JSON: %v", err) - } - assert.Equal(t, tc.page, page, fmt.Sprintf("%v unexpected response, expected: %v, got: %v", tc.desc, tc.page, page)) - } - - sdkCall.Unset() - }) - } -} - -func TestDownloadCACmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - certCmd := cli.NewCertsCmd() - rootCmd := setFlags(certCmd) - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - logMessage string - logType outputLog - certBundle sdk.CertificateBundle - }{ - { - desc: "download CA successfully", - args: []string{}, - logType: entityLog, - certBundle: sdk.CertificateBundle{ - Certificate: []byte("certificate"), - }, - logMessage: "Saved ca.crt\n\nAll certificate files have been saved successfully.\n", - }, - { - desc: "download CA with invalid args", - args: []string{ - extraArg, - }, - logType: usageLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - defer func() { - cleanupFiles(t, []string{"ca.crt"}) - }() - sdkCall := sdkMock.On("DownloadCA", mock.Anything).Return(tc.certBundle, tc.sdkErr) - out := executeCommand(t, rootCmd, append([]string{downloadCACmd}, tc.args...)...) - switch tc.logType { - case entityLog: - assert.True(t, strings.Contains(out, "Saved ca.crt"), fmt.Sprintf("%s invalid output: %s", tc.desc, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - } - sdkCall.Unset() - }) - } -} - -func TestViewCACmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - certCmd := cli.NewCertsCmd() - rootCmd := setFlags(certCmd) - - var cert sdk.Certificate - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - logType outputLog - cert sdk.Certificate - }{ - { - desc: "view cert successfully", - args: []string{}, - logType: entityLog, - cert: sdk.Certificate{ - Certificate: "certificate", - Key: "privatekey", - }, - }, - { - desc: "view cert failed", - args: []string{}, - sdkErr: errors.NewSDKErrorWithStatus(certs.ErrUpdateEntity, http.StatusUnprocessableEntity), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(certs.ErrUpdateEntity, http.StatusUnprocessableEntity)), - logType: errLog, - cert: sdk.Certificate{}, - }, - { - desc: "view cert with invalid args", - args: []string{extraArg}, - logType: usageLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("ViewCA", mock.Anything).Return(tc.cert, tc.sdkErr) - out := executeCommand(t, rootCmd, append([]string{viewCACmd}, tc.args...)...) - switch tc.logType { - case entityLog: - err := json.Unmarshal([]byte(out), &cert) - assert.Nil(t, err) - assert.Equal(t, tc.cert, cert, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.cert, cert)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - } - sdkCall.Unset() - }) - } -} - -func TestGenerateCRLCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - certCmd := cli.NewCertsCmd() - rootCmd := setFlags(certCmd) - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - logType outputLog - crlBytes []byte - }{ - { - desc: "generate CRL successfully", - args: []string{}, - logType: entityLog, - crlBytes: []byte("crl-data"), - }, - { - desc: "generate CRL failed", - args: []string{}, - sdkErr: errors.NewSDKErrorWithStatus(certs.ErrFailedCertCreation, http.StatusUnprocessableEntity), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(certs.ErrFailedCertCreation, http.StatusUnprocessableEntity)), - logType: errLog, - }, - { - desc: "generate CRL with invalid args", - args: []string{"invalid"}, - logType: usageLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - defer func() { - cleanupFiles(t, []string{"ca.crl"}) - }() - - sdkCall := sdkMock.On("GenerateCRL", mock.Anything).Return(tc.crlBytes, tc.sdkErr) - defer sdkCall.Unset() - - out := executeCommand(t, rootCmd, append([]string{"crl"}, tc.args...)...) - - switch tc.logType { - case entityLog: - assert.True(t, strings.Contains(out, "CRL file has been saved successfully"), fmt.Sprintf("%s invalid output: %s", tc.desc, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - } - }) - } -} - -func TestGetEntityIDCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - certCmd := cli.NewCertsCmd() - rootCmd := setFlags(certCmd) - - entityID := "test-entity-id" - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - logType outputLog - entityID string - }{ - { - desc: "get entity ID successfully", - args: []string{serialNumber, certsDomainID, certsToken}, - logType: entityLog, - entityID: entityID, - }, - { - desc: "get entity ID with invalid args", - args: []string{serialNumber, extraArg}, - logType: usageLog, - }, - { - desc: "get entity ID failed", - args: []string{serialNumber, certsDomainID, certsToken}, - sdkErr: errors.NewSDKErrorWithStatus(certs.ErrViewEntity, http.StatusUnprocessableEntity), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(certs.ErrViewEntity, http.StatusUnprocessableEntity)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("EntityID", mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(tc.entityID, tc.sdkErr) - out := executeCommand(t, rootCmd, append([]string{"entity-id"}, tc.args...)...) - - switch tc.logType { - case entityLog: - assert.True(t, strings.Contains(out, tc.entityID), fmt.Sprintf("%s invalid output: %s", tc.desc, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - } - sdkCall.Unset() - }) - } -} - -func cleanupFiles(t *testing.T, filenames []string) { - for _, filename := range filenames { - err := os.Remove(filename) - if err != nil && !os.IsNotExist(err) { - t.Logf("Failed to remove file %s: %v", filename, err) - } - } -} - -func TestIssueFromCSRInternalCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - certCmd := cli.NewCertsCmd() - rootCmd := setFlags(certCmd) - - agentToken := "agent-certsToken-123" - csrPath := "test.csr" - bytes := []byte("-----BEGIN CERTIFICATE REQUEST-----\n-csr-content\n-----END CERTIFICATE REQUEST-----") - - err := os.WriteFile(csrPath, bytes, filePermission) - if err != nil { - t.Fatalf("Failed to create test CSR file: %v", err) - } - defer os.Remove(csrPath) - - var cert sdk.Certificate - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - logType outputLog - cert sdk.Certificate - }{ - { - desc: "issue cert from CSR internal successfully", - args: []string{ - id, - "10h", - csrPath, - agentToken, - }, - logType: entityLog, - cert: sdk.Certificate{SerialNumber: serialNumber}, - }, - { - desc: "issue cert from CSR internal with invalid args", - args: []string{ - id, - extraArg, - }, - logType: usageLog, - }, - { - desc: "issue cert from CSR internal failed", - args: []string{ - id, - "10h", - csrPath, - agentToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(certs.ErrFailedCertCreation, http.StatusUnprocessableEntity), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(certs.ErrFailedCertCreation, http.StatusUnprocessableEntity)), - logType: errLog, - }, - { - desc: "issue cert from CSR internal with non-existent file", - args: []string{ - id, - "10h", - "non-existent.csr", - agentToken, - }, - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - defer func() { - cleanupFiles(t, []string{"cert.pem", "key.pem"}) - }() - sdkCall := sdkMock.On("IssueFromCSRInternal", mock.Anything, mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(tc.cert, tc.sdkErr) - out := executeCommand(t, rootCmd, append([]string{"issue-csr-internal"}, tc.args...)...) - switch tc.logType { - case entityLog: - lines := strings.Split(out, "\n") - var jsonLines []string - var inJSON bool - - for _, line := range lines { - line = strings.TrimSpace(line) - if strings.HasPrefix(line, "{") { - inJSON = true - jsonLines = append(jsonLines, line) - } else if inJSON && strings.HasSuffix(line, "}") { - jsonLines = append(jsonLines, line) - break - } else if inJSON { - jsonLines = append(jsonLines, line) - } - } - - if len(jsonLines) == 0 { - t.Fatalf("No JSON found in output: %s", out) - } - - jsonPart := strings.Join(jsonLines, "") - - err := json.Unmarshal([]byte(jsonPart), &cert) - assert.Nil(t, err) - assert.Equal(t, tc.cert, cert, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.cert, cert)) - assert.True(t, strings.Contains(out, "All certificate files have been saved successfully"), fmt.Sprintf("%s should save files", tc.desc)) - case errLog: - if tc.errLogMessage != "" { - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - } else { - assert.True(t, strings.Contains(out, "error"), fmt.Sprintf("%s should contain error message: %s", tc.desc, out)) - } - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - sdkCall.Unset() - }) - } -} - -func TestIssueFromCSRCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - certCmd := cli.NewCertsCmd() - rootCmd := setFlags(certCmd) - - csrPath := "test.csr" - bytes := []byte("-----BEGIN CERTIFICATE REQUEST-----\n-csr-content\n-----END CERTIFICATE REQUEST-----") - - err := os.WriteFile(csrPath, bytes, filePermission) - if err != nil { - t.Fatalf("Failed to create test CSR file: %v", err) - } - defer os.Remove(csrPath) - - var cert sdk.Certificate - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - logType outputLog - cert sdk.Certificate - }{ - { - desc: "issue cert from CSR successfully", - args: []string{ - id, - "10h", - csrPath, - certsDomainID, - certsToken, - }, - logType: entityLog, - cert: sdk.Certificate{SerialNumber: serialNumber}, - }, - { - desc: "issue cert from CSR with invalid args", - args: []string{ - id, - extraArg, - }, - logType: usageLog, - }, - { - desc: "issue cert from CSR failed", - args: []string{ - id, - "10h", - csrPath, - certsDomainID, - certsToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(certs.ErrFailedCertCreation, http.StatusUnprocessableEntity), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(certs.ErrFailedCertCreation, http.StatusUnprocessableEntity)), - logType: errLog, - }, - { - desc: "issue cert from CSR with non-existent file", - args: []string{ - id, - "10h", - "non-existent.csr", - certsDomainID, - certsToken, - }, - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - defer func() { - cleanupFiles(t, []string{"cert.pem", "key.pem"}) - }() - sdkCall := sdkMock.On("IssueFromCSR", mock.Anything, mock.Anything, mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(tc.cert, tc.sdkErr) - out := executeCommand(t, rootCmd, append([]string{"issue-csr"}, tc.args...)...) - switch tc.logType { - case entityLog: - lines := strings.Split(out, "\n") - var jsonLines []string - var inJSON bool - - for _, line := range lines { - line = strings.TrimSpace(line) - if strings.HasPrefix(line, "{") { - inJSON = true - jsonLines = append(jsonLines, line) - } else if inJSON && strings.HasSuffix(line, "}") { - jsonLines = append(jsonLines, line) - break - } else if inJSON { - jsonLines = append(jsonLines, line) - } - } - - if len(jsonLines) == 0 { - t.Fatalf("No JSON found in output: %s", out) - } - - jsonPart := strings.Join(jsonLines, "") - - err := json.Unmarshal([]byte(jsonPart), &cert) - assert.Nil(t, err) - assert.Equal(t, tc.cert, cert, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.cert, cert)) - assert.True(t, strings.Contains(out, "All certificate files have been saved successfully"), fmt.Sprintf("%s should save files", tc.desc)) - case errLog: - if tc.errLogMessage != "" { - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - } else { - assert.True(t, strings.Contains(out, "error"), fmt.Sprintf("%s should contain error message: %s", tc.desc, out)) - } - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - sdkCall.Unset() - }) - } -} diff --git a/cli/channels_test.go b/cli/channels_test.go deleted file mode 100644 index 4bbf87a5f..000000000 --- a/cli/channels_test.go +++ /dev/null @@ -1,657 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package cli_test - -import ( - "encoding/json" - "fmt" - "net/http" - "strings" - "testing" - - "github.com/absmach/magistrala/cli" - "github.com/absmach/magistrala/internal/testsutil" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - mgsdk "github.com/absmach/magistrala/pkg/sdk" - sdkmocks "github.com/absmach/magistrala/pkg/sdk/mocks" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" -) - -var channel = mgsdk.Channel{ - ID: testsutil.GenerateUUID(&testing.T{}), - Name: "testchannel", -} - -func TestCreateChannelCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - channelJson := "{\"name\":\"testchannel\", \"metadata\":{\"key1\":\"value1\"}}" - channelCmd := cli.NewChannelsCmd() - rootCmd := setFlags(channelCmd) - - cp := mgsdk.Channel{} - cases := []struct { - desc string - args []string - logType outputLog - channel mgsdk.Channel - sdkErr errors.SDKError - errLogMessage string - }{ - { - desc: "create channel successfully", - args: []string{ - createCmd, - channelJson, - domainID, - token, - }, - channel: channel, - logType: entityLog, - }, - { - desc: "create channel with invalid args", - args: []string{ - createCmd, - channelJson, - domainID, - token, - extraArg, - }, - logType: usageLog, - }, - { - desc: "create channel with invalid json", - args: []string{ - createCmd, - "{\"name\":\"testchannel\", \"metadata\":{\"key1\":\"value1\"}", - domainID, - token, - }, - sdkErr: errors.NewSDKError(errors.New("unexpected end of JSON input")), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.New("unexpected end of JSON input")), - logType: errLog, - }, - { - desc: "create channel with invalid token", - args: []string{ - createCmd, - channelJson, - domainID, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusUnauthorized), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusUnauthorized)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - var sdkCall *mock.Call - if len(tc.args) >= 4 { - sdkCall = sdkMock.On("CreateChannel", mock.Anything, mock.Anything, tc.args[2], tc.args[3]).Return(tc.channel, tc.sdkErr) - } - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case entityLog: - err := json.Unmarshal([]byte(out), &cp) - assert.Nil(t, err) - assert.Equal(t, tc.channel, cp, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.channel, cp)) - case usageLog: - assert.True(t, strings.Contains(out, "cli channels create"), fmt.Sprintf("%s invalid usage: expected to contain create usage, got: %s", tc.desc, out)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - } - if sdkCall != nil { - sdkCall.Unset() - } - }) - } -} - -func TestGetChannelsCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - channelCmd := cli.NewChannelsCmd() - rootCmd := setFlags(channelCmd) - - var ch mgsdk.Channel - var page mgsdk.ChannelsPage - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - page mgsdk.ChannelsPage - channel mgsdk.Channel - logType outputLog - errLogMessage string - }{ - { - desc: "get all channels successfully", - args: []string{ - all, - getCmd, - domainID, - token, - }, - page: mgsdk.ChannelsPage{ - Channels: []mgsdk.Channel{channel}, - }, - logType: entityLog, - }, - { - desc: "get channel with id", - args: []string{ - channel.ID, - getCmd, - domainID, - token, - }, - logType: entityLog, - channel: channel, - }, - { - desc: "get channels with invalid args", - args: []string{ - all, - getCmd, - domainID, - token, - extraArg, - }, - logType: usageLog, - }, - { - desc: "get all channels with invalid token", - args: []string{ - all, - getCmd, - domainID, - invalidToken, - }, - logType: errLog, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - }, - { - desc: "get channel with invalid id", - args: []string{ - invalidID, - getCmd, - domainID, - token, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("Channel", mock.Anything, tc.args[0], tc.args[2], tc.args[3]).Return(tc.channel, tc.sdkErr) - sdkCall1 := sdkMock.On("Channels", mock.Anything, mock.Anything, tc.args[2], tc.args[3]).Return(tc.page, tc.sdkErr) - - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case entityLog: - if tc.args[0] == all { - err := json.Unmarshal([]byte(out), &page) - assert.Nil(t, err) - assert.Equal(t, tc.page, page, fmt.Sprintf("%v unexpected response, expected: %v, got: %v", tc.desc, tc.page, page)) - } else { - err := json.Unmarshal([]byte(out), &ch) - assert.Nil(t, err) - assert.Equal(t, tc.channel, ch, fmt.Sprintf("%v unexpected response, expected: %v, got: %v", tc.desc, tc.channel, ch)) - } - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - sdkCall.Unset() - sdkCall1.Unset() - }) - } -} - -func TestDeleteChannelCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - channelCmd := cli.NewChannelsCmd() - rootCmd := setFlags(channelCmd) - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - logType outputLog - errLogMessage string - }{ - { - desc: "delete channel successfully", - args: []string{ - channel.ID, - delCmd, - domainID, - token, - }, - logType: okLog, - }, - { - desc: "delete channel with invalid args", - args: []string{ - channel.ID, - delCmd, - domainID, - token, - extraArg, - }, - logType: usageLog, - }, - { - desc: "delete channel with invalid channel id", - args: []string{ - invalidID, - delCmd, - domainID, - token, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - { - desc: "delete channel with invalid token", - args: []string{ - channel.ID, - delCmd, - domainID, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("DeleteChannel", mock.Anything, tc.args[0], tc.args[2], tc.args[3]).Return(tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case okLog: - assert.True(t, strings.Contains(out, "ok"), fmt.Sprintf("%s unexpected response: expected success message, got: %v", tc.desc, out)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - sdkCall.Unset() - }) - } -} - -func TestUpdateChannelCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - channelCmd := cli.NewChannelsCmd() - rootCmd := setFlags(channelCmd) - - newChannelJson := "{\"name\" : \"channel1\"}" - cases := []struct { - desc string - args []string - channel mgsdk.Channel - sdkErr errors.SDKError - errLogMessage string - logType outputLog - }{ - { - desc: "update channel successfully", - args: []string{ - channel.ID, - updateCmd, - newChannelJson, - domainID, - token, - }, - channel: mgsdk.Channel{ - Name: "newchannel1", - ID: channel.ID, - }, - logType: entityLog, - }, - { - desc: "update channel with invalid args", - args: []string{ - channel.ID, - updateCmd, - newChannelJson, - domainID, - token, - extraArg, - }, - logType: usageLog, - }, - { - desc: "update channel with invalid channel id", - args: []string{ - invalidID, - updateCmd, - newChannelJson, - domainID, - token, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - { - desc: "update channel with invalid json syntax", - args: []string{ - channel.ID, - updateCmd, - "{\"name\" : \"channel1\"", - domainID, - token, - }, - sdkErr: errors.NewSDKError(errors.New("unexpected end of JSON input")), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.New("unexpected end of JSON input")), - logType: errLog, - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - var ch mgsdk.Channel - sdkCall := sdkMock.On("UpdateChannel", mock.Anything, mock.Anything, tc.args[3], tc.args[4]).Return(tc.channel, tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case entityLog: - err := json.Unmarshal([]byte(out), &ch) - assert.Nil(t, err) - assert.Equal(t, tc.channel, ch, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.channel, ch)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - } - sdkCall.Unset() - }) - } -} - -func TestEnableChannelCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - channelCmd := cli.NewChannelsCmd() - rootCmd := setFlags(channelCmd) - var ch mgsdk.Channel - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - channel mgsdk.Channel - logType outputLog - }{ - { - desc: "enable channel successfully", - args: []string{ - channel.ID, - enableCmd, - domainID, - validToken, - }, - channel: channel, - logType: entityLog, - }, - { - desc: "delete channel with invalid token", - args: []string{ - channel.ID, - enableCmd, - domainID, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - { - desc: "delete channel with invalid channel ID", - args: []string{ - invalidID, - enableCmd, - domainID, - token, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - { - desc: "enable channel with invalid args", - args: []string{ - channel.ID, - enableCmd, - domainID, - validToken, - extraArg, - }, - logType: usageLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("EnableChannel", mock.Anything, tc.args[0], tc.args[2], tc.args[3]).Return(tc.channel, tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - case entityLog: - err := json.Unmarshal([]byte(out), &ch) - assert.Nil(t, err) - assert.Equal(t, tc.channel, ch, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.channel, ch)) - } - - sdkCall.Unset() - }) - } -} - -func TestDisableChannelCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - channelsCmd := cli.NewChannelsCmd() - rootCmd := setFlags(channelsCmd) - - var ch mgsdk.Channel - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - channel mgsdk.Channel - logType outputLog - }{ - { - desc: "disable channel successfully", - args: []string{ - channel.ID, - disableCmd, - domainID, - validToken, - }, - logType: entityLog, - channel: channel, - }, - { - desc: "disable channel with invalid token", - args: []string{ - channel.ID, - disableCmd, - domainID, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - { - desc: "disable channel with invalid id", - args: []string{ - invalidID, - disableCmd, - domainID, - token, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - { - desc: "disable client with invalid args", - args: []string{ - channel.ID, - disableCmd, - domainID, - validToken, - extraArg, - }, - logType: usageLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("DisableChannel", mock.Anything, tc.args[0], tc.args[2], tc.args[3]).Return(tc.channel, tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - case entityLog: - err := json.Unmarshal([]byte(out), &ch) - if err != nil { - t.Fatalf("json.Unmarshal failed: %v", err) - } - assert.Equal(t, tc.channel, ch, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.channel, ch)) - } - - sdkCall.Unset() - }) - } -} - -func TestChannelUsersCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - channelsCmd := cli.NewChannelsCmd() - rootCmd := setFlags(channelsCmd) - - var mp mgsdk.EntityMembersPage - - memberRole := mgsdk.MemberRoles{ - MemberID: testsutil.GenerateUUID(t), - Roles: []mgsdk.MemberRole{}, - } - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - usersPage mgsdk.EntityMembersPage - logType outputLog - }{ - { - desc: "list channel users successfully", - args: []string{ - channel.ID, - usersCmd, - domainID, - validToken, - }, - usersPage: mgsdk.EntityMembersPage{ - Members: []mgsdk.MemberRoles{memberRole}, - }, - logType: entityLog, - }, - { - desc: "list channel users with invalid token", - args: []string{ - channel.ID, - usersCmd, - domainID, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - { - desc: "list channel users with invalid channel id", - args: []string{ - invalidID, - usersCmd, - domainID, - token, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - { - desc: "list channel users with invalid args", - args: []string{ - channel.ID, - usersCmd, - domainID, - validToken, - extraArg, - }, - errLogMessage: rootCmd.Use, - logType: usageLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("ListChannelMembers", mock.Anything, tc.args[0], tc.args[2], mock.Anything, tc.args[3]).Return(tc.usersPage, tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - case entityLog: - err := json.Unmarshal([]byte(out), &mp) - if err != nil { - t.Fatalf("json.Unmarshal failed: %v", err) - } - assert.Equal(t, tc.usersPage, mp, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.usersPage, mp)) - } - - sdkCall.Unset() - }) - } -} diff --git a/cli/clients.go b/cli/clients.go deleted file mode 100644 index fa7c43ea4..000000000 --- a/cli/clients.go +++ /dev/null @@ -1,698 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package cli - -import ( - "encoding/json" - "fmt" - - "github.com/absmach/magistrala/clients" - smqsdk "github.com/absmach/magistrala/pkg/sdk" - "github.com/spf13/cobra" -) - -const ( - connect = "connect" - disconnect = "disconnect" - roles = "roles" - actions = "actions" - members = "members" - secret = "secret" - - // Usage strings for client operations. - usageClientCreate = "cli clients create " - usageClientGet = "cli clients get " - usageClientDelete = "cli clients delete " - usageClientUpdate = "cli clients update " - usageClientUpdateTags = "cli clients update tags " - usageClientUpdateSecret = "cli clients update secret " - usageClientEnable = "cli clients enable " - usageClientDisable = "cli clients disable " - usageClientConnect = "cli clients connect " - usageClientDisconnect = "cli clients disconnect " - usageClientUsers = "cli clients users " - - // Usage strings for client roles operations. - usageClientRolesCreate = "cli clients roles create " - usageClientRolesGet = "cli clients roles get " - usageClientRolesUpdate = "cli clients roles update " - usageClientRolesDelete = "cli clients roles delete " - - // Usage strings for client role actions operations. - usageClientRoleActionsAdd = "cli clients roles actions add " - usageClientRoleActionsList = "cli clients roles actions list " - usageClientRoleActionsDelete = "cli clients roles actions delete " - usageClientRoleActionsAvailable = "cli clients roles actions available-actions " - - // Usage strings for client role members operations. - usageClientRoleMembersAdd = "cli clients roles members add " - usageClientRoleMembersList = "cli clients roles members list " - usageClientRoleMembersDelete = "cli clients roles members delete " -) - -func NewClientsCmd() *cobra.Command { - cmd := &cobra.Command{ - Use: "clients [operation] [args...]", - Short: "Clients management", - Long: `Format: - clients create [args...] - clients [args...] - -Operations (require client_id/all): get, update, delete, enable, disable, connect, disconnect, users, roles - -Examples: - clients create - clients all get - clients get - clients update - clients delete - clients enable - clients disable - clients connect - clients users `, - - Run: func(cmd *cobra.Command, args []string) { - if len(args) == 0 { - logUsageCmd(*cmd, cmd.Use) - return - } - - if args[0] == create { - handleClientCreate(cmd, args[1:]) - return - } - - if len(args) < 2 { - logUsageCmd(*cmd, "clients [args...]") - return - } - - clientParams := args[0] - operation := args[1] - opArgs := args[2:] - - switch operation { - case get: - handleClientGet(cmd, clientParams, opArgs) - case update: - handleClientUpdate(cmd, clientParams, opArgs) - case delete: - handleClientDelete(cmd, clientParams, opArgs) - case enable: - handleClientEnable(cmd, clientParams, opArgs) - case disable: - handleClientDisable(cmd, clientParams, opArgs) - case connect: - handleClientConnect(cmd, clientParams, opArgs) - case disconnect: - handleClientDisconnect(cmd, clientParams, opArgs) - case users: - handleClientUsers(cmd, clientParams, opArgs) - case roles: - handleClientRoles(cmd, clientParams, opArgs) - default: - logErrorCmd(*cmd, fmt.Errorf("unknown operation: %s", operation)) - } - }, - } - - return cmd -} - -func handleClientCreate(cmd *cobra.Command, args []string) { - if len(args) != 3 { - logUsageCmd(*cmd, usageClientCreate) - return - } - - var client smqsdk.Client - if err := json.Unmarshal([]byte(args[0]), &client); err != nil { - logErrorCmd(*cmd, err) - return - } - - client.Status = clients.EnabledStatus.String() - client, err := sdk.CreateClient(cmd.Context(), client, args[1], args[2]) - if err != nil { - logErrorCmd(*cmd, err) - return - } - - logJSONCmd(*cmd, client) -} - -func handleClientGet(cmd *cobra.Command, clientParams string, args []string) { - if len(args) != 2 { - logUsageCmd(*cmd, usageClientGet) - return - } - - if clientParams == all { - metadata, err := convertMetadata(Metadata) - if err != nil { - logErrorCmd(*cmd, err) - return - } - - pageMetadata := smqsdk.PageMetadata{ - Name: Name, - Offset: Offset, - Limit: Limit, - Metadata: metadata, - } - - l, err := sdk.Clients(cmd.Context(), pageMetadata, args[0], args[1]) - if err != nil { - logErrorCmd(*cmd, err) - return - } - logJSONCmd(*cmd, l) - return - } - - t, err := sdk.Client(cmd.Context(), clientParams, args[0], args[1]) - if err != nil { - logErrorCmd(*cmd, err) - return - } - - logJSONCmd(*cmd, t) -} - -func handleClientUpdate(cmd *cobra.Command, clientID string, args []string) { - if len(args) < 3 || len(args) > 4 { - if args[0] == tags { - logUsageCmd(*cmd, usageClientUpdateTags) - return - } - if args[0] == secret { - logUsageCmd(*cmd, usageClientUpdateSecret) - return - } - logUsageCmd(*cmd, usageClientUpdate) - return - } - - if len(args) == 4 && args[0] == "tags" { - var client smqsdk.Client - if err := json.Unmarshal([]byte(args[1]), &client.Tags); err != nil { - logErrorCmd(*cmd, err) - return - } - client.ID = clientID - client, err := sdk.UpdateClientTags(cmd.Context(), client, args[2], args[3]) - if err != nil { - logErrorCmd(*cmd, err) - return - } - logJSONCmd(*cmd, client) - return - } - - if len(args) == 4 && args[0] == "secret" { - client, err := sdk.UpdateClientSecret(cmd.Context(), clientID, args[1], args[2], args[3]) - if err != nil { - logErrorCmd(*cmd, err) - return - } - logJSONCmd(*cmd, client) - return - } - - if len(args) != 3 { - logUsageCmd(*cmd, usageClientUpdate) - return - } - - var client smqsdk.Client - if err := json.Unmarshal([]byte(args[0]), &client); err != nil { - logErrorCmd(*cmd, err) - return - } - - client.ID = clientID - client, err := sdk.UpdateClient(cmd.Context(), client, args[1], args[2]) - if err != nil { - logErrorCmd(*cmd, err) - return - } - - logJSONCmd(*cmd, client) -} - -func handleClientDelete(cmd *cobra.Command, clientID string, args []string) { - if len(args) != 2 { - logUsageCmd(*cmd, usageClientDelete) - return - } - - if err := sdk.DeleteClient(cmd.Context(), clientID, args[0], args[1]); err != nil { - logErrorCmd(*cmd, err) - return - } - logOKCmd(*cmd) -} - -func handleClientEnable(cmd *cobra.Command, clientID string, args []string) { - if len(args) != 2 { - logUsageCmd(*cmd, usageClientEnable) - return - } - - client, err := sdk.EnableClient(cmd.Context(), clientID, args[0], args[1]) - if err != nil { - logErrorCmd(*cmd, err) - return - } - - logJSONCmd(*cmd, client) -} - -func handleClientDisable(cmd *cobra.Command, clientID string, args []string) { - if len(args) != 2 { - logUsageCmd(*cmd, usageClientDisable) - return - } - - client, err := sdk.DisableClient(cmd.Context(), clientID, args[0], args[1]) - if err != nil { - logErrorCmd(*cmd, err) - return - } - - logJSONCmd(*cmd, client) -} - -func handleClientConnect(cmd *cobra.Command, clientID string, args []string) { - if len(args) != 4 { - logUsageCmd(*cmd, usageClientConnect) - return - } - - var conn_types []string - err := json.Unmarshal([]byte(args[1]), &conn_types) - if err != nil { - logErrorCmd(*cmd, err) - return - } - - connIDs := smqsdk.Connection{ - ChannelIDs: []string{args[0]}, - ClientIDs: []string{clientID}, - Types: conn_types, - } - if err := sdk.Connect(cmd.Context(), connIDs, args[2], args[3]); err != nil { - logErrorCmd(*cmd, err) - return - } - - logOKCmd(*cmd) -} - -func handleClientDisconnect(cmd *cobra.Command, clientID string, args []string) { - if len(args) != 4 { - logUsageCmd(*cmd, usageClientDisconnect) - return - } - - var conn_types []string - err := json.Unmarshal([]byte(args[1]), &conn_types) - if err != nil { - logErrorCmd(*cmd, err) - return - } - - connIDs := smqsdk.Connection{ - ClientIDs: []string{clientID}, - ChannelIDs: []string{args[0]}, - Types: conn_types, - } - if err := sdk.Disconnect(cmd.Context(), connIDs, args[2], args[3]); err != nil { - logErrorCmd(*cmd, err) - return - } - - logOKCmd(*cmd) -} - -func handleClientUsers(cmd *cobra.Command, clientID string, args []string) { - if len(args) != 2 { - logUsageCmd(*cmd, usageClientUsers) - return - } - - pm := smqsdk.PageMetadata{ - Offset: Offset, - Limit: Limit, - } - ul, err := sdk.ListClientMembers(cmd.Context(), clientID, args[0], pm, args[1]) - if err != nil { - logErrorCmd(*cmd, err) - return - } - - logJSONCmd(*cmd, ul) -} - -func handleClientRoles(cmd *cobra.Command, clientID string, args []string) { - if len(args) < 1 { - logUsageCmd(*cmd, "cli clients roles [args...]") - return - } - - operation := args[0] - opArgs := args[1:] - - switch operation { - case create: - handleClientRoleCreate(cmd, clientID, opArgs) - case get: - handleClientRoleGet(cmd, clientID, opArgs) - case update: - handleClientRoleUpdate(cmd, clientID, opArgs) - case delete: - handleClientRoleDelete(cmd, clientID, opArgs) - case actions: - handleClientRoleActions(cmd, clientID, opArgs) - case members: - handleClientRoleMembers(cmd, clientID, opArgs) - default: - logErrorCmd(*cmd, fmt.Errorf("unknown roles operation: %s", operation)) - } -} - -func handleClientRoleCreate(cmd *cobra.Command, clientID string, args []string) { - if len(args) != 3 { - logUsageCmd(*cmd, usageClientRolesCreate) - return - } - - var roleReq smqsdk.RoleReq - if err := json.Unmarshal([]byte(args[0]), &roleReq); err != nil { - logErrorCmd(*cmd, err) - return - } - - r, err := sdk.CreateClientRole(cmd.Context(), clientID, args[1], roleReq, args[2]) - if err != nil { - logErrorCmd(*cmd, err) - return - } - - logJSONCmd(*cmd, r) -} - -func handleClientRoleGet(cmd *cobra.Command, clientID string, args []string) { - if len(args) != 3 { - logUsageCmd(*cmd, usageClientRolesGet) - return - } - - roleID := args[0] - domainID := args[1] - token := args[2] - - if roleID == all { - pageMetadata := smqsdk.PageMetadata{ - Offset: Offset, - Limit: Limit, - } - rs, err := sdk.ClientRoles(cmd.Context(), clientID, domainID, pageMetadata, token) - if err != nil { - logErrorCmd(*cmd, err) - return - } - logJSONCmd(*cmd, rs) - return - } - - r, err := sdk.ClientRole(cmd.Context(), clientID, roleID, domainID, token) - if err != nil { - logErrorCmd(*cmd, err) - return - } - logJSONCmd(*cmd, r) -} - -func handleClientRoleUpdate(cmd *cobra.Command, clientID string, args []string) { - if len(args) != 4 { - logUsageCmd(*cmd, usageClientRolesUpdate) - return - } - - roleID := args[0] - newName := args[1] - domainID := args[2] - token := args[3] - - r, err := sdk.UpdateClientRole(cmd.Context(), clientID, roleID, newName, domainID, token) - if err != nil { - logErrorCmd(*cmd, err) - return - } - logJSONCmd(*cmd, r) -} - -func handleClientRoleDelete(cmd *cobra.Command, clientID string, args []string) { - if len(args) != 3 { - logUsageCmd(*cmd, usageClientRolesDelete) - return - } - - roleID := args[0] - domainID := args[1] - token := args[2] - - if err := sdk.DeleteClientRole(cmd.Context(), clientID, roleID, domainID, token); err != nil { - logErrorCmd(*cmd, err) - return - } - logOKCmd(*cmd) -} - -func handleClientRoleActions(cmd *cobra.Command, clientID string, args []string) { - if len(args) < 1 { - logUsageCmd(*cmd, "cli clients roles actions [args...]") - return - } - - operation := args[0] - opArgs := args[1:] - - switch operation { - case add: - handleClientRoleActionsAdd(cmd, clientID, opArgs) - case list: - handleClientRoleActionsList(cmd, clientID, opArgs) - case delete: - handleClientRoleActionsDelete(cmd, clientID, opArgs) - case availableActions: - handleClientRoleActionsAvailable(cmd, opArgs) - default: - logErrorCmd(*cmd, fmt.Errorf("unknown actions operation: %s", operation)) - } -} - -func handleClientRoleActionsAdd(cmd *cobra.Command, clientID string, args []string) { - if len(args) != 4 { - logUsageCmd(*cmd, usageClientRoleActionsAdd) - return - } - - roleID := args[0] - actionsJSON := args[1] - domainID := args[2] - token := args[3] - - actions := struct { - Actions []string `json:"actions"` - }{} - if err := json.Unmarshal([]byte(actionsJSON), &actions); err != nil { - logErrorCmd(*cmd, err) - return - } - - acts, err := sdk.AddClientRoleActions(cmd.Context(), clientID, roleID, domainID, actions.Actions, token) - if err != nil { - logErrorCmd(*cmd, err) - return - } - logJSONCmd(*cmd, acts) -} - -func handleClientRoleActionsList(cmd *cobra.Command, clientID string, args []string) { - if len(args) != 3 { - logUsageCmd(*cmd, usageClientRoleActionsList) - return - } - - roleID := args[0] - domainID := args[1] - token := args[2] - - l, err := sdk.ClientRoleActions(cmd.Context(), clientID, roleID, domainID, token) - if err != nil { - logErrorCmd(*cmd, err) - return - } - logJSONCmd(*cmd, l) -} - -func handleClientRoleActionsDelete(cmd *cobra.Command, clientID string, args []string) { - if len(args) != 4 { - logUsageCmd(*cmd, usageClientRoleActionsDelete) - return - } - - roleID := args[0] - actionsJSON := args[1] - domainID := args[2] - token := args[3] - - if actionsJSON == all { - if err := sdk.RemoveAllClientRoleActions(cmd.Context(), clientID, roleID, domainID, token); err != nil { - logErrorCmd(*cmd, err) - return - } - logOKCmd(*cmd) - return - } - - actions := struct { - Actions []string `json:"actions"` - }{} - if err := json.Unmarshal([]byte(actionsJSON), &actions); err != nil { - logErrorCmd(*cmd, err) - return - } - - if err := sdk.RemoveClientRoleActions(cmd.Context(), clientID, roleID, domainID, actions.Actions, token); err != nil { - logErrorCmd(*cmd, err) - return - } - logOKCmd(*cmd) -} - -func handleClientRoleActionsAvailable(cmd *cobra.Command, args []string) { - if len(args) != 2 { - logUsageCmd(*cmd, usageClientRoleActionsAvailable) - return - } - - domainID := args[0] - token := args[1] - - acts, err := sdk.AvailableClientRoleActions(cmd.Context(), domainID, token) - if err != nil { - logErrorCmd(*cmd, err) - return - } - logJSONCmd(*cmd, acts) -} - -func handleClientRoleMembers(cmd *cobra.Command, clientID string, args []string) { - if len(args) < 1 { - logUsageCmd(*cmd, "cli clients roles members [args...]") - return - } - - operation := args[0] - opArgs := args[1:] - - switch operation { - case add: - handleClientRoleMembersAdd(cmd, clientID, opArgs) - case list: - handleClientRoleMembersList(cmd, clientID, opArgs) - case delete: - handleClientRoleMembersDelete(cmd, clientID, opArgs) - default: - logErrorCmd(*cmd, fmt.Errorf("unknown members operation: %s", operation)) - } -} - -func handleClientRoleMembersAdd(cmd *cobra.Command, clientID string, args []string) { - if len(args) != 4 { - logUsageCmd(*cmd, usageClientRoleMembersAdd) - return - } - - roleID := args[0] - membersJSON := args[1] - domainID := args[2] - token := args[3] - - members := struct { - Members []string `json:"members"` - }{} - if err := json.Unmarshal([]byte(membersJSON), &members); err != nil { - logErrorCmd(*cmd, err) - return - } - - memb, err := sdk.AddClientRoleMembers(cmd.Context(), clientID, roleID, domainID, members.Members, token) - if err != nil { - logErrorCmd(*cmd, err) - return - } - logJSONCmd(*cmd, memb) -} - -func handleClientRoleMembersList(cmd *cobra.Command, clientID string, args []string) { - if len(args) != 3 { - logUsageCmd(*cmd, usageClientRoleMembersList) - return - } - - roleID := args[0] - domainID := args[1] - token := args[2] - - pageMetadata := smqsdk.PageMetadata{ - Offset: Offset, - Limit: Limit, - } - - l, err := sdk.ClientRoleMembers(cmd.Context(), clientID, roleID, domainID, pageMetadata, token) - if err != nil { - logErrorCmd(*cmd, err) - return - } - logJSONCmd(*cmd, l) -} - -func handleClientRoleMembersDelete(cmd *cobra.Command, clientID string, args []string) { - if len(args) != 4 { - logUsageCmd(*cmd, usageClientRoleMembersDelete) - return - } - - roleID := args[0] - membersJSON := args[1] - domainID := args[2] - token := args[3] - - if membersJSON == all { - if err := sdk.RemoveAllClientRoleMembers(cmd.Context(), clientID, roleID, domainID, token); err != nil { - logErrorCmd(*cmd, err) - return - } - logOKCmd(*cmd) - return - } - - members := struct { - Members []string `json:"members"` - }{} - if err := json.Unmarshal([]byte(membersJSON), &members); err != nil { - logErrorCmd(*cmd, err) - return - } - - if err := sdk.RemoveClientRoleMembers(cmd.Context(), clientID, roleID, domainID, members.Members, token); err != nil { - logErrorCmd(*cmd, err) - return - } - logOKCmd(*cmd) -} diff --git a/cli/clients_test.go b/cli/clients_test.go deleted file mode 100644 index 26e7d227e..000000000 --- a/cli/clients_test.go +++ /dev/null @@ -1,1940 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package cli_test - -import ( - "encoding/json" - "fmt" - "net/http" - "strings" - "testing" - - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/cli" - "github.com/absmach/magistrala/clients" - "github.com/absmach/magistrala/internal/testsutil" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - smqsdk "github.com/absmach/magistrala/pkg/sdk" - sdkmocks "github.com/absmach/magistrala/pkg/sdk/mocks" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" -) - -var ( - token = "valid" + "domaintoken" - relation = "administrator" - conntype = `["publish","subscribe"]` - - errEndJSONInput = errors.New("unexpected end of JSON input") -) - -var client = smqsdk.Client{ - ID: testsutil.GenerateUUID(&testing.T{}), - Name: "testclient", - Credentials: smqsdk.ClientCredentials{ - Secret: "secret", - }, - DomainID: testsutil.GenerateUUID(&testing.T{}), - Status: clients.EnabledStatus.String(), -} - -func TestCreateClientsCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - clientJson := "{\"name\":\"testclient\", \"metadata\":{\"key1\":\"value1\"}}" - clientsCmd := cli.NewClientsCmd() - rootCmd := setFlags(clientsCmd) - - var tg smqsdk.Client - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - client smqsdk.Client - logType outputLog - }{ - { - desc: "create client successfully with token", - args: []string{ - createCmd, - clientJson, - domainID, - token, - }, - client: client, - logType: entityLog, - }, - { - desc: "create client without token", - args: []string{ - createCmd, - clientJson, - domainID, - }, - logType: usageLog, - }, - { - desc: "create client with invalid token", - args: []string{ - createCmd, - clientJson, - domainID, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusUnauthorized), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusUnauthorized)), - logType: errLog, - }, - { - desc: "failed to create client", - args: []string{ - createCmd, - clientJson, - domainID, - token, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrCreateEntity, http.StatusUnprocessableEntity), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrCreateEntity, http.StatusUnprocessableEntity)), - logType: errLog, - }, - { - desc: "create client with invalid metadata", - args: []string{ - createCmd, - "{\"name\":\"testclient\", \"metadata\":{\"key1\":value1}}", - domainID, - token, - }, - sdkErr: errors.NewSDKErrorWithStatus(errors.New("invalid character 'v' looking for beginning of value"), 306), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.New("invalid character 'v' looking for beginning of value")), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - var sdkCall *mock.Call - if len(tc.args) >= 4 { - sdkCall = sdkMock.On("CreateClient", mock.Anything, mock.Anything, tc.args[2], tc.args[3]).Return(tc.client, tc.sdkErr) - } - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case entityLog: - err := json.Unmarshal([]byte(out), &tg) - assert.Nil(t, err) - assert.Equal(t, tc.client, tg, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.client, tg)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.True(t, strings.Contains(out, "cli clients create"), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - - if sdkCall != nil { - sdkCall.Unset() - } - }) - } -} - -func TestGetClientssCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - clientsCmd := cli.NewClientsCmd() - rootCmd := setFlags(clientsCmd) - - var tg smqsdk.Client - var page smqsdk.ClientsPage - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - client smqsdk.Client - page smqsdk.ClientsPage - logType outputLog - }{ - { - desc: "get all clients successfully", - args: []string{ - all, - getCmd, - domainID, - token, - }, - logType: entityLog, - page: smqsdk.ClientsPage{ - Clients: []smqsdk.Client{client}, - }, - }, - { - desc: "get client successfully with id", - args: []string{ - client.ID, - getCmd, - domainID, - token, - }, - logType: entityLog, - client: client, - }, - { - desc: "get clients with invalid token", - args: []string{ - all, - getCmd, - domainID, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - page: smqsdk.ClientsPage{}, - logType: errLog, - }, - { - desc: "get clients with invalid args", - args: []string{ - all, - getCmd, - invalidToken, - all, - invalidToken, - all, - invalidToken, - all, - invalidToken, - }, - logType: usageLog, - }, - { - desc: "get client without token", - args: []string{ - all, - getCmd, - domainID, - }, - logType: usageLog, - }, - { - desc: "get client with invalid client id", - args: []string{ - invalidID, - getCmd, - domainID, - token, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - var sdkCall, sdkCall1 *mock.Call - if len(tc.args) >= 4 { - sdkCall = sdkMock.On("Clients", mock.Anything, mock.Anything, tc.args[2], tc.args[3]).Return(tc.page, tc.sdkErr) - sdkCall1 = sdkMock.On("Client", mock.Anything, tc.args[0], tc.args[2], tc.args[3]).Return(tc.client, tc.sdkErr) - } - - out := executeCommand(t, rootCmd, tc.args...) - - if tc.logType == entityLog { - switch { - case tc.args[0] == all: - err := json.Unmarshal([]byte(out), &page) - if err != nil { - t.Fatalf("Failed to unmarshal JSON: %v", err) - } - default: - err := json.Unmarshal([]byte(out), &tg) - if err != nil { - t.Fatalf("Failed to unmarshal JSON: %v", err) - } - } - } - - switch tc.logType { - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - - if tc.logType == entityLog { - if tc.args[1] != all { - assert.Equal(t, tc.client, tg, fmt.Sprintf("%v unexpected response, expected: %v, got: %v", tc.desc, tc.client, tg)) - } else { - assert.Equal(t, tc.page, page, fmt.Sprintf("%v unexpected response, expected: %v, got: %v", tc.desc, tc.page, page)) - } - } - - if sdkCall != nil { - sdkCall.Unset() - } - if sdkCall1 != nil { - sdkCall1.Unset() - } - }) - } -} - -func TestUpdateClientCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - clientsCmd := cli.NewClientsCmd() - rootCmd := setFlags(clientsCmd) - - tagUpdateType := "tags" - secretUpdateType := "secret" - newTagsJson := "[\"tag1\", \"tag2\"]" - newTagString := []string{"tag1", "tag2"} - newNameandMeta := "{\"name\": \"clientName\", \"metadata\": {\"role\": \"general\"}}" - newMetadata := "{\"metadata\": {\"role\": \"general\"}}" - newPrivateMeta := "{\"private_metadata\": {\"role\": \"general\"}}" - newSecret := "secret" - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - client smqsdk.Client - logType outputLog - }{ - { - desc: "update client name and metadata successfully", - args: []string{ - client.ID, - updateCmd, - newNameandMeta, - domainID, - token, - }, - client: smqsdk.Client{ - Name: "clientName", - Metadata: map[string]any{ - "role": "general", - }, - ID: client.ID, - DomainID: client.DomainID, - Status: client.Status, - }, - logType: entityLog, - }, - { - desc: "update client name and metadata successfully", - args: []string{ - client.ID, - updateCmd, - newNameandMeta, - domainID, - token, - }, - client: smqsdk.Client{ - Name: "clientName", - Metadata: map[string]any{ - "role": "general", - }, - ID: client.ID, - DomainID: client.DomainID, - Status: client.Status, - }, - logType: entityLog, - }, - { - desc: "update client private metadata successfully", - args: []string{ - client.ID, - updateCmd, - newPrivateMeta, - domainID, - token, - }, - client: smqsdk.Client{ - PrivateMetadata: map[string]any{ - "role": "general", - }, - ID: client.ID, - DomainID: client.DomainID, - Status: client.Status, - }, - logType: entityLog, - }, - { - desc: "update client metadata successfully", - args: []string{ - client.ID, - updateCmd, - newMetadata, - domainID, - token, - }, - client: smqsdk.Client{ - Metadata: map[string]any{ - "role": "general", - }, - ID: client.ID, - DomainID: client.DomainID, - Status: client.Status, - }, - logType: entityLog, - }, - { - desc: "update client private metadata with invalid json", - args: []string{ - client.ID, - updateCmd, - "{\"private_metadata\": {\"role\": \"general\"}", - domainID, - token, - }, - sdkErr: errors.NewSDKError(errEndJSONInput), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errEndJSONInput), - logType: errLog, - }, - { - desc: "update client metadata with invalid json", - args: []string{ - client.ID, - updateCmd, - "{\"metadata\": {\"role\": \"general\"}", - domainID, - token, - }, - sdkErr: errors.NewSDKError(errEndJSONInput), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errEndJSONInput), - logType: errLog, - }, - { - desc: "update client name and metadata with invalid json", - args: []string{ - client.ID, - updateCmd, - "{\"name\": \"clientName\", \"metadata\": {\"role\": \"general\"}", - domainID, - token, - }, - sdkErr: errors.NewSDKError(errEndJSONInput), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errEndJSONInput), - logType: errLog, - }, - { - desc: "update client name and metadata with invalid client id", - args: []string{ - invalidID, - updateCmd, - newNameandMeta, - domainID, - token, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - { - desc: "update client tags successfully", - args: []string{ - client.ID, - updateCmd, - tagUpdateType, - newTagsJson, - domainID, - token, - }, - client: smqsdk.Client{ - Name: client.Name, - ID: client.ID, - DomainID: client.DomainID, - Status: client.Status, - Tags: newTagString, - }, - logType: entityLog, - }, - { - desc: "update client with invalid tags", - args: []string{ - client.ID, - updateCmd, - tagUpdateType, - "[\"tag1\", \"tag2\"", - domainID, - token, - }, - logType: errLog, - sdkErr: errors.NewSDKError(errEndJSONInput), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errEndJSONInput), - }, - { - desc: "update client tags with invalid client id", - args: []string{ - invalidID, - updateCmd, - tagUpdateType, - newTagsJson, - domainID, - token, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - { - desc: "update client secret successfully", - args: []string{ - client.ID, - updateCmd, - secretUpdateType, - newSecret, - domainID, - token, - }, - client: smqsdk.Client{ - Name: client.Name, - ID: client.ID, - DomainID: client.DomainID, - Status: client.Status, - Credentials: smqsdk.ClientCredentials{ - Secret: newSecret, - }, - }, - logType: entityLog, - }, - { - desc: "update client with invalid secret", - args: []string{ - client.ID, - updateCmd, - secretUpdateType, - "", - domainID, - token, - }, - sdkErr: errors.NewSDKErrorWithStatus(errors.Wrap(apiutil.ErrValidation, apiutil.ErrMissingSecret), http.StatusBadRequest), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(errors.Wrap(apiutil.ErrValidation, apiutil.ErrMissingSecret), http.StatusBadRequest)), - logType: errLog, - }, - { - desc: "update client with invalid token", - args: []string{ - client.ID, - updateCmd, - secretUpdateType, - newSecret, - domainID, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - { - desc: "update client with invalid args", - args: []string{ - client.ID, - updateCmd, - secretUpdateType, - newSecret, - domainID, - token, - extraArg, - }, - logType: usageLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - var tg smqsdk.Client - sdkCall := sdkMock.On("UpdateClient", mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(tc.client, tc.sdkErr) - sdkCall1 := sdkMock.On("UpdateClientTags", mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(tc.client, tc.sdkErr) - sdkCall2 := sdkMock.On("UpdateClientSecret", mock.Anything, mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(tc.client, tc.sdkErr) - - switch { - case len(tc.args) > 2 && tc.args[2] == tagUpdateType: - var th smqsdk.Client - th.Tags = []string{"tag1", "tag2"} - th.ID = tc.args[0] - - sdkCall1 = sdkMock.On("UpdateClientTags", th, tc.args[5]).Return(tc.client, tc.sdkErr) - case len(tc.args) > 2 && tc.args[2] == secretUpdateType: - var th smqsdk.Client - th.Credentials.Secret = tc.args[3] - th.ID = tc.args[0] - - sdkCall2 = sdkMock.On("UpdateClientSecret", th, tc.args[3], tc.args[5]).Return(tc.client, tc.sdkErr) - } - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case entityLog: - err := json.Unmarshal([]byte(out), &tg) - assert.Nil(t, err) - assert.Equal(t, tc.client, tg, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.client, tg)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - - sdkCall.Unset() - sdkCall1.Unset() - sdkCall2.Unset() - }) - } -} - -func TestDeleteClientCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - clientdCmd := cli.NewClientsCmd() - rootCmd := setFlags(clientdCmd) - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - logType outputLog - }{ - { - desc: "delete client successfully", - args: []string{ - client.ID, - delCmd, - domainID, - token, - }, - logType: okLog, - }, - { - desc: "delete client with invalid token", - args: []string{ - client.ID, - delCmd, - domainID, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - { - desc: "delete client with invalid client id", - args: []string{ - invalidID, - delCmd, - domainID, - token, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - { - desc: "delete client with invalid args", - args: []string{ - client.ID, - delCmd, - domainID, - token, - extraArg, - }, - logType: usageLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("DeleteClient", mock.Anything, tc.args[0], tc.args[2], tc.args[3]).Return(tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case okLog: - assert.True(t, strings.Contains(out, "ok"), fmt.Sprintf("%s unexpected response: expected success message, got: %v", tc.desc, out)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - sdkCall.Unset() - }) - } -} - -func TestEnableClientCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - clientsCmd := cli.NewClientsCmd() - rootCmd := setFlags(clientsCmd) - var tg smqsdk.Client - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - client smqsdk.Client - logType outputLog - }{ - { - desc: "enable client successfully", - args: []string{ - client.ID, - enableCmd, - domainID, - validToken, - }, - sdkErr: nil, - client: client, - logType: entityLog, - }, - { - desc: "delete client with invalid token", - args: []string{ - client.ID, - enableCmd, - domainID, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - { - desc: "delete client with invalid client ID", - args: []string{ - invalidID, - enableCmd, - domainID, - token, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - { - desc: "enable client with invalid args", - args: []string{ - client.ID, - enableCmd, - domainID, - validToken, - extraArg, - }, - logType: usageLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("EnableClient", mock.Anything, tc.args[0], tc.args[2], tc.args[3]).Return(tc.client, tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - case entityLog: - err := json.Unmarshal([]byte(out), &tg) - assert.Nil(t, err) - assert.Equal(t, tc.client, tg, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.client, tg)) - } - - sdkCall.Unset() - }) - } -} - -func TestDisableclientCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - clientsCmd := cli.NewClientsCmd() - rootCmd := setFlags(clientsCmd) - - var tg smqsdk.Client - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - client smqsdk.Client - logType outputLog - }{ - { - desc: "disable client successfully", - args: []string{ - client.ID, - disableCmd, - domainID, - validToken, - }, - logType: entityLog, - client: client, - }, - { - desc: "delete client with invalid token", - args: []string{ - client.ID, - disableCmd, - domainID, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - { - desc: "delete client with invalid client ID", - args: []string{ - invalidID, - disableCmd, - domainID, - token, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - { - desc: "disable client with invalid args", - args: []string{ - client.ID, - disableCmd, - domainID, - validToken, - extraArg, - }, - logType: usageLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("DisableClient", mock.Anything, tc.args[0], tc.args[2], tc.args[3]).Return(tc.client, tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - case entityLog: - err := json.Unmarshal([]byte(out), &tg) - if err != nil { - t.Fatalf("json.Unmarshal failed: %v", err) - } - assert.Equal(t, tc.client, tg, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.client, tg)) - } - - sdkCall.Unset() - }) - } -} - -func TestConnectClientCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - clientsCmd := cli.NewClientsCmd() - rootCmd := setFlags(clientsCmd) - - cases := []struct { - desc string - args []string - logType outputLog - sdkErr errors.SDKError - errLogMessage string - }{ - { - desc: "Connect client to channel successfully", - args: []string{ - client.ID, - connCmd, - channel.ID, - conntype, - domainID, - token, - }, - logType: okLog, - }, - { - desc: "connect with invalid args", - args: []string{ - client.ID, - connCmd, - channel.ID, - conntype, - domainID, - token, - extraArg, - }, - logType: usageLog, - }, - { - desc: "connect with invalid client id", - args: []string{ - invalidID, - connCmd, - channel.ID, - conntype, - domainID, - token, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAddPolicies, http.StatusBadRequest), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAddPolicies, http.StatusBadRequest)), - logType: errLog, - }, - { - desc: "connect with invalid channel id", - args: []string{ - client.ID, - connCmd, - invalidID, - conntype, - domainID, - token, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - { - desc: "list client users' with invalid domain", - args: []string{ - client.ID, - connCmd, - channel.ID, - conntype, - invalidID, - token, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrDomainAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrDomainAuthorization, http.StatusForbidden)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("Connect", mock.Anything, mock.Anything, tc.args[4], tc.args[5]).Return(tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case okLog: - assert.True(t, strings.Contains(out, "ok"), fmt.Sprintf("%s unexpected response: expected success message, got: %v", tc.desc, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - } - sdkCall.Unset() - }) - } -} - -func TestDisconnectClientCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - clientsCmd := cli.NewClientsCmd() - rootCmd := setFlags(clientsCmd) - - cases := []struct { - desc string - args []string - logType outputLog - sdkErr errors.SDKError - errLogMessage string - }{ - { - desc: "Disconnect client to channel successfully", - args: []string{ - client.ID, - disconnCmd, - channel.ID, - conntype, - domainID, - token, - }, - logType: okLog, - }, - { - desc: "Disconnect with invalid args", - args: []string{ - client.ID, - disconnCmd, - channel.ID, - conntype, - domainID, - token, - extraArg, - }, - logType: usageLog, - }, - { - desc: "disconnect with invalid client id", - args: []string{ - invalidID, - disconnCmd, - channel.ID, - conntype, - domainID, - token, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAddPolicies, http.StatusBadRequest), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAddPolicies, http.StatusBadRequest)), - logType: errLog, - }, - { - desc: "disconnect with invalid channel id", - args: []string{ - client.ID, - disconnCmd, - invalidID, - conntype, - domainID, - token, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - { - desc: "disconnect client with invalid domain", - args: []string{ - client.ID, - disconnCmd, - channel.ID, - conntype, - invalidID, - token, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrDomainAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrDomainAuthorization, http.StatusForbidden)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("Disconnect", mock.Anything, mock.Anything, tc.args[4], tc.args[5]).Return(tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case okLog: - assert.True(t, strings.Contains(out, "ok"), fmt.Sprintf("%s unexpected response: expected success message, got: %v", tc.desc, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - } - sdkCall.Unset() - }) - } -} - -func TestCreateClientRoleCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - clientsCmd := cli.NewClientsCmd() - rootCmd := setFlags(clientsCmd) - - roleReq := smqsdk.RoleReq{ - RoleName: "admin", - OptionalActions: []string{"read", "update"}, - } - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - role smqsdk.Role - logType outputLog - }{ - { - desc: "create client role successfully", - args: []string{ - client.ID, - rolesCmd, - createCmd, - `{"role_name":"admin","optional_actions":["read","update"]}`, - domainID, - token, - }, - role: smqsdk.Role{ - ID: testsutil.GenerateUUID(&testing.T{}), - Name: "admin", - OptionalActions: []string{"read", "update"}, - }, - logType: entityLog, - }, - { - desc: "create client role with invalid JSON", - args: []string{ - client.ID, - rolesCmd, - createCmd, - `{"role_name":"admin","optional_actions":["read","update"}`, - domainID, - token, - }, - sdkErr: errors.NewSDKError(errors.New("invalid character '}' after array element")), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.New("invalid character '}' after array element")), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("CreateClientRole", mock.Anything, tc.args[0], tc.args[4], roleReq, tc.args[5]).Return(tc.role, tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case entityLog: - var role smqsdk.Role - err := json.Unmarshal([]byte(out), &role) - assert.Nil(t, err) - assert.Equal(t, tc.role, role, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.role, role)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - } - - sdkCall.Unset() - }) - } -} - -func TestGetClientRolesCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - clientsCmd := cli.NewClientsCmd() - rootCmd := setFlags(clientsCmd) - - role := smqsdk.Role{ - ID: testsutil.GenerateUUID(&testing.T{}), - Name: "admin", - OptionalActions: []string{"read", "update"}, - } - rolesPage := smqsdk.RolesPage{ - Total: 1, - Offset: 0, - Limit: 10, - Roles: []smqsdk.Role{role}, - } - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - roles smqsdk.RolesPage - logType outputLog - }{ - { - desc: "get all client roles successfully", - args: []string{ - client.ID, - rolesCmd, - getCmd, - all, - domainID, - token, - }, - roles: rolesPage, - logType: entityLog, - }, - { - desc: "get client roles with invalid token", - args: []string{ - client.ID, - rolesCmd, - getCmd, - all, - domainID, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("ClientRoles", mock.Anything, tc.args[0], tc.args[4], mock.Anything, tc.args[5]).Return(tc.roles, tc.sdkErr) - if tc.args[3] != all { - sdkCall = sdkMock.On("ClientRole", mock.Anything, tc.args[0], tc.args[3], tc.args[4], tc.args[5]).Return(role, tc.sdkErr) - } - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case entityLog: - var roles smqsdk.RolesPage - err := json.Unmarshal([]byte(out), &roles) - assert.Nil(t, err) - assert.Equal(t, tc.roles, roles, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.roles, roles)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - } - - sdkCall.Unset() - }) - } -} - -func TestUpdateClientRoleCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - clientsCmd := cli.NewClientsCmd() - rootCmd := setFlags(clientsCmd) - - role := smqsdk.Role{ - ID: testsutil.GenerateUUID(&testing.T{}), - Name: "new_name", - OptionalActions: []string{"read", "update"}, - } - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - role smqsdk.Role - logType outputLog - }{ - { - desc: "update client role name successfully", - args: []string{ - client.ID, - rolesCmd, - updateCmd, - role.ID, - "new_name", - domainID, - token, - }, - role: role, - logType: entityLog, - }, - { - desc: "update client role name with invalid token", - args: []string{ - client.ID, - rolesCmd, - updateCmd, - role.ID, - "new_name", - domainID, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("UpdateClientRole", mock.Anything, tc.args[0], tc.args[3], tc.args[4], tc.args[5], tc.args[6]).Return(tc.role, tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case entityLog: - var role smqsdk.Role - err := json.Unmarshal([]byte(out), &role) - assert.Nil(t, err) - assert.Equal(t, tc.role, role, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.role, role)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - } - - sdkCall.Unset() - }) - } -} - -func TestDeleteClientRoleCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - clientsCmd := cli.NewClientsCmd() - rootCmd := setFlags(clientsCmd) - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - logType outputLog - }{ - { - desc: "delete client role successfully", - args: []string{ - client.ID, - rolesCmd, - delCmd, - roleID, - domainID, - token, - }, - logType: okLog, - }, - { - desc: "delete client role with invalid token", - args: []string{ - client.ID, - rolesCmd, - delCmd, - roleID, - domainID, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("DeleteClientRole", mock.Anything, tc.args[0], tc.args[3], tc.args[4], tc.args[5]).Return(tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case okLog: - assert.True(t, strings.Contains(out, "ok"), fmt.Sprintf("%s unexpected response: expected success message, got: %v", tc.desc, out)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - } - - sdkCall.Unset() - }) - } -} - -func TestAddClientRoleActionsCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - clientsCmd := cli.NewClientsCmd() - rootCmd := setFlags(clientsCmd) - - actions := struct { - Actions []string `json:"actions"` - }{ - Actions: []string{"read", "write"}, - } - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - actions []string - logType outputLog - }{ - { - desc: "add actions to role successfully", - args: []string{ - client.ID, - rolesCmd, - actionsCmd, - addCmd, - roleID, - `{"actions":["read","write"]}`, - domainID, - token, - }, - actions: actions.Actions, - logType: entityLog, - }, - { - desc: "add actions to role with invalid JSON", - args: []string{ - client.ID, - rolesCmd, - actionsCmd, - addCmd, - roleID, - `{"actions":["read","write"}`, - domainID, - token, - }, - sdkErr: errors.NewSDKError(errors.New("invalid character '}' after array element")), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.New("invalid character '}' after array element")), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("AddClientRoleActions", mock.Anything, tc.args[0], tc.args[4], tc.args[6], tc.actions, tc.args[7]).Return(tc.actions, tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case entityLog: - var acts []string - err := json.Unmarshal([]byte(out), &acts) - assert.Nil(t, err) - assert.Equal(t, tc.actions, acts, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.actions, acts)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - } - - sdkCall.Unset() - }) - } -} - -func TestListClientRoleActionsCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - clientsCmd := cli.NewClientsCmd() - rootCmd := setFlags(clientsCmd) - - actions := []string{"read", "write"} - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - actions []string - logType outputLog - }{ - { - desc: "list actions of role successfully", - args: []string{ - client.ID, - rolesCmd, - actionsCmd, - listCmd, - roleID, - domainID, - token, - }, - actions: actions, - logType: entityLog, - }, - { - desc: "list actions of role with invalid token", - args: []string{ - client.ID, - rolesCmd, - actionsCmd, - listCmd, - roleID, - domainID, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("ClientRoleActions", mock.Anything, tc.args[0], tc.args[4], tc.args[5], tc.args[6]).Return(tc.actions, tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case entityLog: - var acts []string - err := json.Unmarshal([]byte(out), &acts) - assert.Nil(t, err) - assert.Equal(t, tc.actions, acts, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.actions, acts)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - } - - sdkCall.Unset() - }) - } -} - -func TestDeleteClientRoleActionsCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - clientsCmd := cli.NewClientsCmd() - rootCmd := setFlags(clientsCmd) - - actions := struct { - Actions []string `json:"actions"` - }{ - Actions: []string{"read", "write"}, - } - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - logType outputLog - }{ - { - desc: "delete actions from role successfully", - args: []string{ - client.ID, - rolesCmd, - actionsCmd, - delCmd, - roleID, - `{"actions":["read","write"]}`, - domainID, - token, - }, - logType: okLog, - }, - { - desc: "delete all actions from role successfully", - args: []string{ - client.ID, - rolesCmd, - actionsCmd, - delCmd, - roleID, - all, - domainID, - token, - }, - logType: okLog, - }, - { - desc: "delete actions from role with invalid JSON", - args: []string{ - client.ID, - rolesCmd, - actionsCmd, - delCmd, - roleID, - `{"actions":["read","write"}`, - domainID, - token, - }, - sdkErr: errors.NewSDKError(errors.New("invalid character '}' after array element")), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.New("invalid character '}' after array element")), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - var sdkCall *mock.Call - if tc.args[5] == all { - sdkCall = sdkMock.On("RemoveAllClientRoleActions", mock.Anything, tc.args[0], tc.args[4], tc.args[6], tc.args[7]).Return(tc.sdkErr) - } else { - sdkCall = sdkMock.On("RemoveClientRoleActions", mock.Anything, tc.args[0], tc.args[4], tc.args[6], actions.Actions, tc.args[7]).Return(tc.sdkErr) - } - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case okLog: - assert.True(t, strings.Contains(out, "ok"), fmt.Sprintf("%s unexpected response: expected success message, got: %v", tc.desc, out)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - } - - sdkCall.Unset() - }) - } -} - -func TestAvailableClientRoleActionsCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - clientsCmd := cli.NewClientsCmd() - rootCmd := setFlags(clientsCmd) - - actions := []string{"read", "write", "update"} - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - actions []string - logType outputLog - }{ - { - desc: "list available actions successfully", - args: []string{ - client.ID, - rolesCmd, - actionsCmd, - availableActionsCmd, - domainID, - token, - }, - actions: actions, - logType: entityLog, - }, - { - desc: "list available actions with invalid token", - args: []string{ - client.ID, - rolesCmd, - actionsCmd, - availableActionsCmd, - domainID, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("AvailableClientRoleActions", mock.Anything, tc.args[4], tc.args[5]).Return(tc.actions, tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case entityLog: - var acts []string - err := json.Unmarshal([]byte(out), &acts) - assert.Nil(t, err) - assert.Equal(t, tc.actions, acts, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.actions, acts)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - } - - sdkCall.Unset() - }) - } -} - -func TestAddClientRoleMembersCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - clientsCmd := cli.NewClientsCmd() - rootCmd := setFlags(clientsCmd) - - members := struct { - Members []string `json:"members"` - }{ - Members: []string{"5dc1ce4b-7cc9-4f12-98a6-9d74cc4980bb", "5dc1ce4b-7cc9-4f12-98a6-9d74cc4980bb"}, - } - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - members []string - logType outputLog - }{ - { - desc: "add members to role successfully", - args: []string{ - client.ID, - rolesCmd, - membersCmd, - addCmd, - roleID, - `{"members":["5dc1ce4b-7cc9-4f12-98a6-9d74cc4980bb", "5dc1ce4b-7cc9-4f12-98a6-9d74cc4980bb"]}`, - domainID, - token, - }, - members: members.Members, - logType: entityLog, - }, - { - desc: "add members to role with invalid JSON", - args: []string{ - client.ID, - rolesCmd, - membersCmd, - addCmd, - roleID, - `{"members":["5dc1ce4b-7cc9-4f12-98a6-9d74cc4980bb", "5dc1ce4b-7cc9-4f12-98a6-9d74cc4980bb"}`, - domainID, - token, - }, - sdkErr: errors.NewSDKError(errors.New("invalid character '}' after array element")), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.New("invalid character '}' after array element")), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("AddClientRoleMembers", mock.Anything, tc.args[0], tc.args[4], tc.args[6], tc.members, tc.args[7]).Return(tc.members, tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case entityLog: - var members []string - err := json.Unmarshal([]byte(out), &members) - assert.Nil(t, err) - assert.Equal(t, tc.members, members, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.members, members)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - } - - sdkCall.Unset() - }) - } -} - -func TestListClientRoleMembersCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - clientsCmd := cli.NewClientsCmd() - rootCmd := setFlags(clientsCmd) - - membersPage := smqsdk.RoleMembersPage{ - Total: 1, - Offset: 0, - Limit: 10, - Members: []string{ - "5dc1ce4b-7cc9-4f12-98a6-9d74cc4980bb", - }, - } - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - members smqsdk.RoleMembersPage - logType outputLog - }{ - { - desc: "list members of role successfully", - args: []string{ - client.ID, - rolesCmd, - membersCmd, - listCmd, - roleID, - domainID, - token, - }, - members: membersPage, - logType: entityLog, - }, - { - desc: "list members of role with invalid token", - args: []string{ - client.ID, - rolesCmd, - membersCmd, - listCmd, - roleID, - domainID, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("ClientRoleMembers", mock.Anything, tc.args[0], tc.args[4], tc.args[5], mock.Anything, tc.args[6]).Return(tc.members, tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case entityLog: - var members smqsdk.RoleMembersPage - err := json.Unmarshal([]byte(out), &members) - assert.Nil(t, err) - assert.Equal(t, tc.members, members, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.members, members)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - } - - sdkCall.Unset() - }) - } -} - -func TestDeleteClientRoleMembersCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - clientsCmd := cli.NewClientsCmd() - rootCmd := setFlags(clientsCmd) - - members := struct { - Members []string `json:"members"` - }{ - Members: []string{"5dc1ce4b-7cc9-4f12-98a6-9d74cc4980bb", "5dc1ce4b-7cc9-4f12-98a6-9d74cc4980bb"}, - } - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - logType outputLog - }{ - { - desc: "delete members from role successfully", - args: []string{ - client.ID, - rolesCmd, - membersCmd, - delCmd, - roleID, - `{"members":["5dc1ce4b-7cc9-4f12-98a6-9d74cc4980bb", "5dc1ce4b-7cc9-4f12-98a6-9d74cc4980bb"]}`, - domainID, - token, - }, - logType: okLog, - }, - { - desc: "delete all members from role successfully", - args: []string{ - client.ID, - rolesCmd, - membersCmd, - delCmd, - roleID, - all, - domainID, - token, - }, - logType: okLog, - }, - { - desc: "delete members from role with invalid JSON", - args: []string{ - client.ID, - rolesCmd, - membersCmd, - delCmd, - roleID, - `{"members":["5dc1ce4b-7cc9-4f12-98a6-9d74cc4980bb", "5dc1ce4b-7cc9-4f12-98a6-9d74cc4980bb"}`, - domainID, - token, - }, - sdkErr: errors.NewSDKError(errors.New("invalid character '}' after array element")), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.New("invalid character '}' after array element")), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - var sdkCall *mock.Call - if tc.args[5] == all { - sdkCall = sdkMock.On("RemoveAllClientRoleMembers", mock.Anything, tc.args[0], tc.args[4], tc.args[6], tc.args[7]).Return(tc.sdkErr) - } else { - sdkCall = sdkMock.On("RemoveClientRoleMembers", mock.Anything, tc.args[0], tc.args[4], tc.args[6], members.Members, tc.args[7]).Return(tc.sdkErr) - } - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case okLog: - assert.True(t, strings.Contains(out, "ok"), fmt.Sprintf("%s unexpected response: expected success message, got: %v", tc.desc, out)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - } - - sdkCall.Unset() - }) - } -} - -func TestClientUsersCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - clientsCmd := cli.NewClientsCmd() - rootCmd := setFlags(clientsCmd) - - var mp smqsdk.EntityMembersPage - - memberRole := smqsdk.MemberRoles{ - MemberID: testsutil.GenerateUUID(t), - Roles: []smqsdk.MemberRole{}, - } - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - usersPage smqsdk.EntityMembersPage - logType outputLog - }{ - { - desc: "list client users successfully", - args: []string{ - client.ID, - usersCmd, - domainID, - validToken, - }, - usersPage: smqsdk.EntityMembersPage{ - Members: []smqsdk.MemberRoles{memberRole}, - }, - logType: entityLog, - }, - { - desc: "list client users with invalid token", - args: []string{ - client.ID, - usersCmd, - domainID, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - { - desc: "list client users with invalid client id", - args: []string{ - invalidID, - usersCmd, - domainID, - token, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - { - desc: "list client users with invalid args", - args: []string{ - client.ID, - usersCmd, - domainID, - validToken, - extraArg, - }, - errLogMessage: rootCmd.Use, - logType: usageLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("ListClientMembers", mock.Anything, tc.args[0], tc.args[2], mock.Anything, tc.args[3]).Return(tc.usersPage, tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - case entityLog: - err := json.Unmarshal([]byte(out), &mp) - if err != nil { - t.Fatalf("json.Unmarshal failed: %v", err) - } - assert.Equal(t, tc.usersPage, mp, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.usersPage, mp)) - } - - sdkCall.Unset() - }) - } -} diff --git a/cli/commands_test.go b/cli/commands_test.go deleted file mode 100644 index 16ecfb259..000000000 --- a/cli/commands_test.go +++ /dev/null @@ -1,61 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package cli_test - -// CRUD and common commands -const ( - createCmd = "create" - updateCmd = "update" - getCmd = "get" - enableCmd = "enable" - disableCmd = "disable" - freezeCmd = "freeze" - delCmd = "delete" -) - -// Users commands -const ( - tokCmd = "token" - refTokCmd = "refreshtoken" - profCmd = "profile" - resPassReqCmd = "resetpasswordrequest" - resPassCmd = "resetpassword" - passCmd = "password" -) - -// Clients commands -const ( - connCmd = "connect" - disconnCmd = "disconnect" - usersCmd = "users" -) - -// Messages commands -const sendCmd = "send" - -// Invitations commands -const ( - acceptCmd = "accept" - rejectCmd = "reject" - userCmd = "user" - domainCmd = "domain" -) - -// Role commands -const ( - rolesCmd = "roles" - actionsCmd = "actions" - availableActionsCmd = "available-actions" - addCmd = "add" - listCmd = "list" - membersCmd = "members" -) - -// Bootstrap commands -const ( - updCmd = "update" - rmCmd = "remove" - whitelistCmd = "whitelist" - bootStrapCmd = "bootstrap" -) diff --git a/cli/consumers_test.go b/cli/consumers_test.go deleted file mode 100644 index 5ac037508..000000000 --- a/cli/consumers_test.go +++ /dev/null @@ -1,266 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package cli_test - -import ( - "encoding/json" - "fmt" - "net/http" - "strings" - "testing" - - "github.com/absmach/magistrala/cli" - "github.com/absmach/magistrala/internal/testsutil" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - mgsdk "github.com/absmach/magistrala/pkg/sdk" - sdkmocks "github.com/absmach/magistrala/pkg/sdk/mocks" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" -) - -var ( - userID = testsutil.GenerateUUID(&testing.T{}) - subscription = mgsdk.Subscription{ - ID: testsutil.GenerateUUID(&testing.T{}), - OwnerID: userID, - Topic: "topic", - Contact: "identity@example.com", - } -) - -func TestCreateSubscriptionCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - subCmd := cli.NewSubscriptionCmd() - rootCmd := setFlags(subCmd) - - cases := []struct { - desc string - args []string - logType outputLog - errLogMessage string - sdkErr errors.SDKError - response string - id string - }{ - { - desc: "create subscription successfully", - args: []string{ - subscription.Topic, - subscription.Contact, - validToken, - }, - id: userID, - response: fmt.Sprintf("\ncreated: %s\n\n", userID), - logType: createLog, - }, - { - desc: "create subscription with invalid args", - args: []string{ - subscription.Topic, - subscription.Contact, - validToken, - extraArg, - }, - logType: usageLog, - }, - { - desc: "create subscription with invalid token", - args: []string{ - subscription.Topic, - subscription.Contact, - invalidToken, - }, - logType: errLog, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("CreateSubscription", mock.Anything, tc.args[0], tc.args[1], tc.args[2]).Return(tc.id, tc.sdkErr) - out := executeCommand(t, rootCmd, append([]string{createCmd}, tc.args...)...) - - switch tc.logType { - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case createLog: - assert.Equal(t, tc.response, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.response, out)) - } - sdkCall.Unset() - }) - } -} - -func TestGetSubscriptionsCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - subCmd := cli.NewSubscriptionCmd() - rootCmd := setFlags(subCmd) - - var sub mgsdk.Subscription - var page mgsdk.SubscriptionPage - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - page mgsdk.SubscriptionPage - subscription mgsdk.Subscription - logType outputLog - errLogMessage string - }{ - { - desc: "get all subscriptions successfully", - args: []string{ - all, - validToken, - }, - page: mgsdk.SubscriptionPage{ - Subscriptions: []mgsdk.Subscription{subscription}, - }, - logType: entityLog, - }, - { - desc: "get subscription with id", - args: []string{ - subscription.ID, - validToken, - }, - logType: entityLog, - subscription: subscription, - }, - { - desc: "get subscriptions with invalid args", - args: []string{ - all, - validToken, - extraArg, - }, - logType: usageLog, - }, - { - desc: "get all subscriptions with invalid token", - args: []string{ - all, - invalidToken, - }, - logType: errLog, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - }, - { - desc: "get subscription with invalid id", - args: []string{ - invalidID, - validToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("ViewSubscription", mock.Anything, tc.args[0], tc.args[1]).Return(tc.subscription, tc.sdkErr) - sdkCall1 := sdkMock.On("ListSubscriptions", mock.Anything, mock.Anything, tc.args[1]).Return(tc.page, tc.sdkErr) - - out := executeCommand(t, rootCmd, append([]string{getCmd}, tc.args...)...) - - switch tc.logType { - case entityLog: - if tc.args[1] == all { - err := json.Unmarshal([]byte(out), &page) - assert.Nil(t, err) - assert.Equal(t, tc.page, page, fmt.Sprintf("%v unexpected response, expected: %v, got: %v", tc.desc, tc.page, page)) - } else { - err := json.Unmarshal([]byte(out), &sub) - assert.Nil(t, err) - assert.Equal(t, tc.subscription, sub, fmt.Sprintf("%v unexpected response, expected: %v, got: %v", tc.desc, tc.subscription, sub)) - } - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - sdkCall.Unset() - sdkCall1.Unset() - }) - } -} - -func TestRemoveSubscriptionCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - subCmd := cli.NewSubscriptionCmd() - rootCmd := setFlags(subCmd) - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - logType outputLog - errLogMessage string - }{ - { - desc: "remove subscription successfully", - args: []string{ - subscription.ID, - validToken, - }, - logType: okLog, - }, - { - desc: "remove subscription with invalid args", - args: []string{ - subscription.ID, - validToken, - extraArg, - }, - logType: usageLog, - }, - { - desc: "remove subscription with invalid subscription id", - args: []string{ - invalidID, - validToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - { - desc: "remove subscription with invalid token", - args: []string{ - subscription.ID, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("DeleteSubscription", mock.Anything, tc.args[0], tc.args[1]).Return(tc.sdkErr) - out := executeCommand(t, rootCmd, append([]string{rmCmd}, tc.args...)...) - - switch tc.logType { - case okLog: - assert.True(t, strings.Contains(out, "ok"), fmt.Sprintf("%s unexpected response: expected success message, got: %v", tc.desc, out)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - sdkCall.Unset() - }) - } -} diff --git a/cli/domains.go b/cli/domains.go deleted file mode 100644 index f044a9a44..000000000 --- a/cli/domains.go +++ /dev/null @@ -1,580 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package cli - -import ( - "encoding/json" - "fmt" - - smqsdk "github.com/absmach/magistrala/pkg/sdk" - "github.com/spf13/cobra" -) - -const ( - freeze = "freeze" - - // Usage strings for domain operations. - usageDomainCreate = "cli domains create " - usageDomainGet = "cli domains get " - usageDomainUpdate = "cli domains update " - usageDomainEnable = "cli domains enable " - usageDomainDisable = "cli domains disable " - usageDomainFreeze = "cli domains freeze " - usageDomainUsers = "cli domains users " - - // Usage strings for domain roles operations. - usageDomainRolesCreate = "cli domains roles create " - usageDomainRolesGet = "cli domains roles get " - usageDomainRolesUpdate = "cli domains roles update " - usageDomainRolesDelete = "cli domains roles delete " - - // Usage strings for domain role actions operations. - usageDomainRoleActionsAdd = "cli domains roles actions add " - usageDomainRoleActionsList = "cli domains roles actions list " - usageDomainRoleActionsDelete = "cli domains roles actions delete " - usageDomainRoleActionsAvailable = "cli domains roles actions available-actions " - - // Usage strings for domain role members operations. - usageDomainRoleMembersAdd = "cli domains roles members add " - usageDomainRoleMembersList = "cli domains roles members list " - usageDomainRoleMembersDelete = "cli domains roles members delete " -) - -func NewDomainsCmd() *cobra.Command { - cmd := &cobra.Command{ - Use: "domains [operation] [args...]", - Short: "Domains management", - Long: `Format: - domains create [args...] - domains [args...] - -Operations (require domain_id/all): get, update, enable, disable, freeze, users, roles - -Examples: - domains create - domains all get - domains get - domains update - domains enable - domains disable - domains freeze - domains users `, - - Run: func(cmd *cobra.Command, args []string) { - if len(args) == 0 { - logUsageCmd(*cmd, cmd.Use) - return - } - - if args[0] == create { - handleDomainCreate(cmd, args[1:]) - return - } - - if len(args) < 2 { - logUsageCmd(*cmd, "domains [args...]") - return - } - - domainParams := args[0] - operation := args[1] - opArgs := args[2:] - - switch operation { - case get: - handleDomainGet(cmd, domainParams, opArgs) - case update: - handleDomainUpdate(cmd, domainParams, opArgs) - case enable: - handleDomainEnable(cmd, domainParams, opArgs) - case disable: - handleDomainDisable(cmd, domainParams, opArgs) - case freeze: - handleDomainFreeze(cmd, domainParams, opArgs) - case users: - handleDomainUsers(cmd, domainParams, opArgs) - case roles: - handleDomainRoles(cmd, domainParams, opArgs) - default: - logErrorCmd(*cmd, fmt.Errorf("unknown operation: %s", operation)) - } - }, - } - - return cmd -} - -func handleDomainCreate(cmd *cobra.Command, args []string) { - if len(args) != 3 { - logUsageCmd(*cmd, usageDomainCreate) - return - } - - dom := smqsdk.Domain{ - Name: args[0], - Route: args[1], - } - d, err := sdk.CreateDomain(cmd.Context(), dom, args[2]) - if err != nil { - logErrorCmd(*cmd, err) - return - } - logJSONCmd(*cmd, d) -} - -func handleDomainGet(cmd *cobra.Command, domainParams string, args []string) { - if len(args) != 1 { - logUsageCmd(*cmd, usageDomainGet) - return - } - - if domainParams == all { - metadata, err := convertMetadata(Metadata) - if err != nil { - logErrorCmd(*cmd, err) - return - } - - pageMetadata := smqsdk.PageMetadata{ - Name: Name, - Offset: Offset, - Limit: Limit, - Metadata: metadata, - Status: Status, - } - - l, err := sdk.Domains(cmd.Context(), pageMetadata, args[0]) - if err != nil { - logErrorCmd(*cmd, err) - return - } - logJSONCmd(*cmd, l) - return - } - - d, err := sdk.Domain(cmd.Context(), domainParams, args[0]) - if err != nil { - logErrorCmd(*cmd, err) - return - } - - logJSONCmd(*cmd, d) -} - -func handleDomainUpdate(cmd *cobra.Command, domainID string, args []string) { - if len(args) != 2 { - logUsageCmd(*cmd, usageDomainUpdate) - return - } - - var d smqsdk.Domain - if err := json.Unmarshal([]byte(args[0]), &d); err != nil { - logErrorCmd(*cmd, err) - return - } - d.ID = domainID - d, err := sdk.UpdateDomain(cmd.Context(), d, args[1]) - if err != nil { - logErrorCmd(*cmd, err) - return - } - logJSONCmd(*cmd, d) -} - -func handleDomainEnable(cmd *cobra.Command, domainID string, args []string) { - if len(args) != 1 { - logUsageCmd(*cmd, usageDomainEnable) - return - } - - if err := sdk.EnableDomain(cmd.Context(), domainID, args[0]); err != nil { - logErrorCmd(*cmd, err) - return - } - logOKCmd(*cmd) -} - -func handleDomainDisable(cmd *cobra.Command, domainID string, args []string) { - if len(args) != 1 { - logUsageCmd(*cmd, usageDomainDisable) - return - } - - if err := sdk.DisableDomain(cmd.Context(), domainID, args[0]); err != nil { - logErrorCmd(*cmd, err) - return - } - logOKCmd(*cmd) -} - -func handleDomainFreeze(cmd *cobra.Command, domainID string, args []string) { - if len(args) != 1 { - logUsageCmd(*cmd, usageDomainFreeze) - return - } - - if err := sdk.FreezeDomain(cmd.Context(), domainID, args[0]); err != nil { - logErrorCmd(*cmd, err) - return - } - logOKCmd(*cmd) -} - -func handleDomainUsers(cmd *cobra.Command, domainID string, args []string) { - if len(args) != 1 { - logUsageCmd(*cmd, usageDomainUsers) - return - } - - metadata, err := convertMetadata(Metadata) - if err != nil { - logErrorCmd(*cmd, err) - return - } - pageMetadata := smqsdk.PageMetadata{ - Offset: Offset, - Limit: Limit, - Metadata: metadata, - Status: Status, - } - - l, err := sdk.ListDomainMembers(cmd.Context(), domainID, pageMetadata, args[0]) - if err != nil { - logErrorCmd(*cmd, err) - return - } - logJSONCmd(*cmd, l) -} - -func handleDomainRoles(cmd *cobra.Command, domainID string, args []string) { - if len(args) < 1 { - logUsageCmd(*cmd, "cli domains roles [args...]") - return - } - - operation := args[0] - opArgs := args[1:] - - switch operation { - case create: - handleDomainRoleCreate(cmd, domainID, opArgs) - case get: - handleDomainRoleGet(cmd, domainID, opArgs) - case update: - handleDomainRoleUpdate(cmd, domainID, opArgs) - case delete: - handleDomainRoleDelete(cmd, domainID, opArgs) - case actions: - handleDomainRoleActions(cmd, domainID, opArgs) - case members: - handleDomainRoleMembers(cmd, domainID, opArgs) - default: - logErrorCmd(*cmd, fmt.Errorf("unknown roles operation: %s", operation)) - } -} - -func handleDomainRoleCreate(cmd *cobra.Command, domainID string, args []string) { - if len(args) != 2 { - logUsageCmd(*cmd, usageDomainRolesCreate) - return - } - - var roleReq smqsdk.RoleReq - if err := json.Unmarshal([]byte(args[0]), &roleReq); err != nil { - logErrorCmd(*cmd, err) - return - } - - r, err := sdk.CreateDomainRole(cmd.Context(), domainID, roleReq, args[1]) - if err != nil { - logErrorCmd(*cmd, err) - return - } - - logJSONCmd(*cmd, r) -} - -func handleDomainRoleGet(cmd *cobra.Command, domainID string, args []string) { - if len(args) != 2 { - logUsageCmd(*cmd, usageDomainRolesGet) - return - } - - roleID := args[0] - token := args[1] - - if roleID == all { - pageMetadata := smqsdk.PageMetadata{ - Offset: Offset, - Limit: Limit, - } - rs, err := sdk.DomainRoles(cmd.Context(), domainID, pageMetadata, token) - if err != nil { - logErrorCmd(*cmd, err) - return - } - logJSONCmd(*cmd, rs) - return - } - - r, err := sdk.DomainRole(cmd.Context(), domainID, roleID, token) - if err != nil { - logErrorCmd(*cmd, err) - return - } - logJSONCmd(*cmd, r) -} - -func handleDomainRoleUpdate(cmd *cobra.Command, domainID string, args []string) { - if len(args) != 3 { - logUsageCmd(*cmd, usageDomainRolesUpdate) - return - } - - roleID := args[0] - newName := args[1] - token := args[2] - - r, err := sdk.UpdateDomainRole(cmd.Context(), domainID, roleID, newName, token) - if err != nil { - logErrorCmd(*cmd, err) - return - } - logJSONCmd(*cmd, r) -} - -func handleDomainRoleDelete(cmd *cobra.Command, domainID string, args []string) { - if len(args) != 2 { - logUsageCmd(*cmd, usageDomainRolesDelete) - return - } - - roleID := args[0] - token := args[1] - - if err := sdk.DeleteDomainRole(cmd.Context(), domainID, roleID, token); err != nil { - logErrorCmd(*cmd, err) - return - } - logOKCmd(*cmd) -} - -func handleDomainRoleActions(cmd *cobra.Command, domainID string, args []string) { - if len(args) < 1 { - logUsageCmd(*cmd, "cli domains roles actions [args...]") - return - } - - operation := args[0] - opArgs := args[1:] - - switch operation { - case add: - handleDomainRoleActionsAdd(cmd, domainID, opArgs) - case list: - handleDomainRoleActionsList(cmd, domainID, opArgs) - case delete: - handleDomainRoleActionsDelete(cmd, domainID, opArgs) - case availableActions: - handleDomainRoleActionsAvailable(cmd, opArgs) - default: - logErrorCmd(*cmd, fmt.Errorf("unknown actions operation: %s", operation)) - } -} - -func handleDomainRoleActionsAdd(cmd *cobra.Command, domainID string, args []string) { - if len(args) != 3 { - logUsageCmd(*cmd, usageDomainRoleActionsAdd) - return - } - - roleID := args[0] - actionsJSON := args[1] - token := args[2] - - actions := struct { - Actions []string `json:"actions"` - }{} - if err := json.Unmarshal([]byte(actionsJSON), &actions); err != nil { - logErrorCmd(*cmd, err) - return - } - - acts, err := sdk.AddDomainRoleActions(cmd.Context(), domainID, roleID, actions.Actions, token) - if err != nil { - logErrorCmd(*cmd, err) - return - } - logJSONCmd(*cmd, acts) -} - -func handleDomainRoleActionsList(cmd *cobra.Command, domainID string, args []string) { - if len(args) != 2 { - logUsageCmd(*cmd, usageDomainRoleActionsList) - return - } - - roleID := args[0] - token := args[1] - - l, err := sdk.DomainRoleActions(cmd.Context(), domainID, roleID, token) - if err != nil { - logErrorCmd(*cmd, err) - return - } - logJSONCmd(*cmd, l) -} - -func handleDomainRoleActionsDelete(cmd *cobra.Command, domainID string, args []string) { - if len(args) != 3 { - logUsageCmd(*cmd, usageDomainRoleActionsDelete) - return - } - - roleID := args[0] - actionsJSON := args[1] - token := args[2] - - if actionsJSON == all { - if err := sdk.RemoveAllDomainRoleActions(cmd.Context(), domainID, roleID, token); err != nil { - logErrorCmd(*cmd, err) - return - } - logOKCmd(*cmd) - return - } - - actions := struct { - Actions []string `json:"actions"` - }{} - if err := json.Unmarshal([]byte(actionsJSON), &actions); err != nil { - logErrorCmd(*cmd, err) - return - } - - if err := sdk.RemoveDomainRoleActions(cmd.Context(), domainID, roleID, actions.Actions, token); err != nil { - logErrorCmd(*cmd, err) - return - } - logOKCmd(*cmd) -} - -func handleDomainRoleActionsAvailable(cmd *cobra.Command, args []string) { - if len(args) != 1 { - logUsageCmd(*cmd, usageDomainRoleActionsAvailable) - return - } - - token := args[0] - - acts, err := sdk.AvailableDomainRoleActions(cmd.Context(), token) - if err != nil { - logErrorCmd(*cmd, err) - return - } - logJSONCmd(*cmd, acts) -} - -func handleDomainRoleMembers(cmd *cobra.Command, domainID string, args []string) { - if len(args) < 1 { - logUsageCmd(*cmd, "cli domains roles members [args...]") - return - } - - operation := args[0] - opArgs := args[1:] - - switch operation { - case add: - handleDomainRoleMembersAdd(cmd, domainID, opArgs) - case list: - handleDomainRoleMembersList(cmd, domainID, opArgs) - case delete: - handleDomainRoleMembersDelete(cmd, domainID, opArgs) - default: - logErrorCmd(*cmd, fmt.Errorf("unknown members operation: %s", operation)) - } -} - -func handleDomainRoleMembersAdd(cmd *cobra.Command, domainID string, args []string) { - if len(args) != 3 { - logUsageCmd(*cmd, usageDomainRoleMembersAdd) - return - } - - roleID := args[0] - membersJSON := args[1] - token := args[2] - - members := struct { - Members []string `json:"members"` - }{} - if err := json.Unmarshal([]byte(membersJSON), &members); err != nil { - logErrorCmd(*cmd, err) - return - } - - memb, err := sdk.AddDomainRoleMembers(cmd.Context(), domainID, roleID, members.Members, token) - if err != nil { - logErrorCmd(*cmd, err) - return - } - logJSONCmd(*cmd, memb) -} - -func handleDomainRoleMembersList(cmd *cobra.Command, domainID string, args []string) { - if len(args) != 2 { - logUsageCmd(*cmd, usageDomainRoleMembersList) - return - } - - roleID := args[0] - token := args[1] - - pageMetadata := smqsdk.PageMetadata{ - Offset: Offset, - Limit: Limit, - } - - l, err := sdk.DomainRoleMembers(cmd.Context(), domainID, roleID, pageMetadata, token) - if err != nil { - logErrorCmd(*cmd, err) - return - } - logJSONCmd(*cmd, l) -} - -func handleDomainRoleMembersDelete(cmd *cobra.Command, domainID string, args []string) { - if len(args) != 3 { - logUsageCmd(*cmd, usageDomainRoleMembersDelete) - return - } - - roleID := args[0] - membersJSON := args[1] - token := args[2] - - if membersJSON == all { - if err := sdk.RemoveAllDomainRoleMembers(cmd.Context(), domainID, roleID, token); err != nil { - logErrorCmd(*cmd, err) - return - } - logOKCmd(*cmd) - return - } - - members := struct { - Members []string `json:"members"` - }{} - if err := json.Unmarshal([]byte(membersJSON), &members); err != nil { - logErrorCmd(*cmd, err) - return - } - - if err := sdk.RemoveDomainRoleMembers(cmd.Context(), domainID, roleID, members.Members, token); err != nil { - logErrorCmd(*cmd, err) - return - } - logOKCmd(*cmd) -} diff --git a/cli/domains_test.go b/cli/domains_test.go deleted file mode 100644 index 24fdda1bb..000000000 --- a/cli/domains_test.go +++ /dev/null @@ -1,1432 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package cli_test - -import ( - "encoding/json" - "fmt" - "net/http" - "strings" - "testing" - - "github.com/absmach/magistrala/cli" - "github.com/absmach/magistrala/internal/testsutil" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - smqsdk "github.com/absmach/magistrala/pkg/sdk" - sdkmocks "github.com/absmach/magistrala/pkg/sdk/mocks" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" -) - -var ( - domain = smqsdk.Domain{ - ID: testsutil.GenerateUUID(&testing.T{}), - Name: "Test domain", - Route: "route", - } - roleID = testsutil.GenerateUUID(&testing.T{}) -) - -func TestCreateDomainsCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - domainCmd := cli.NewDomainsCmd() - rootCmd := setFlags(domainCmd) - - var dom smqsdk.Domain - - cases := []struct { - desc string - args []string - domain smqsdk.Domain - errLogMessage string - sdkErr errors.SDKError - logType outputLog - }{ - { - desc: "create domain successfully", - args: []string{ - createCmd, - dom.Name, - dom.Route, - validToken, - }, - logType: entityLog, - domain: domain, - }, - { - desc: "create domain with invalid args", - args: []string{ - createCmd, - dom.Name, - dom.Route, - validToken, - extraArg, - }, - logType: usageLog, - }, - { - desc: "create domain with invalid token", - args: []string{ - createCmd, - dom.Name, - dom.Route, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusUnauthorized), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusUnauthorized)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("CreateDomain", mock.Anything, mock.Anything, mock.Anything).Return(tc.domain, tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case entityLog: - err := json.Unmarshal([]byte(out), &dom) - assert.Nil(t, err) - assert.Equal(t, tc.domain, dom, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.domain, dom)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.True(t, strings.Contains(out, "cli domains create"), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - sdkCall.Unset() - }) - } -} - -func TestGetDomainsCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - all := "all" - domainCmd := cli.NewDomainsCmd() - rootCmd := setFlags(domainCmd) - - var dom smqsdk.Domain - var page smqsdk.DomainsPage - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - page smqsdk.DomainsPage - domain smqsdk.Domain - logType outputLog - errLogMessage string - }{ - { - desc: "get all domains successfully", - args: []string{ - all, - getCmd, - validToken, - }, - page: smqsdk.DomainsPage{ - Domains: []smqsdk.Domain{domain}, - }, - logType: entityLog, - }, - { - desc: "get domain with id", - args: []string{ - domain.ID, - getCmd, - validToken, - }, - logType: entityLog, - domain: domain, - }, - { - desc: "get domains with invalid args", - args: []string{ - all, - getCmd, - validToken, - extraArg, - }, - logType: usageLog, - }, - { - desc: "get all domains with invalid token", - args: []string{ - all, - getCmd, - invalidToken, - }, - logType: errLog, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - }, - { - desc: "get domain with invalid id", - args: []string{ - invalidID, - getCmd, - validToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("Domain", mock.Anything, tc.args[0], tc.args[2]).Return(tc.domain, tc.sdkErr) - sdkCall1 := sdkMock.On("Domains", mock.Anything, mock.Anything, tc.args[2]).Return(tc.page, tc.sdkErr) - - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case entityLog: - if tc.args[0] == all { - err := json.Unmarshal([]byte(out), &page) - assert.Nil(t, err) - assert.Equal(t, tc.page, page, fmt.Sprintf("%v unexpected response, expected: %v, got: %v", tc.desc, tc.page, page)) - } else { - err := json.Unmarshal([]byte(out), &dom) - assert.Nil(t, err) - assert.Equal(t, tc.domain, dom, fmt.Sprintf("%v unexpected response, expected: %v, got: %v", tc.desc, tc.domain, dom)) - } - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - sdkCall.Unset() - sdkCall1.Unset() - }) - } -} - -func TestUpdateDomainCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - domainsCmd := cli.NewDomainsCmd() - rootCmd := setFlags(domainsCmd) - - newDomainJson := "{\"name\" : \"New domain\"}" - cases := []struct { - desc string - args []string - domain smqsdk.Domain - sdkErr errors.SDKError - errLogMessage string - logType outputLog - }{ - { - desc: "update domain successfully", - args: []string{ - domain.ID, - updateCmd, - newDomainJson, - token, - }, - domain: smqsdk.Domain{ - Name: "New domain", - ID: domain.ID, - }, - logType: entityLog, - }, - { - desc: "update domain with invalid args", - args: []string{ - domain.ID, - updateCmd, - newDomainJson, - token, - extraArg, - extraArg, - }, - logType: usageLog, - }, - { - desc: "update domain with invalid id", - args: []string{ - invalidID, - updateCmd, - newDomainJson, - token, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - { - desc: "update domain with invalid json syntax", - args: []string{ - domain.ID, - updateCmd, - "{\"name\" : \"New domain\"", - token, - }, - sdkErr: errors.NewSDKError(errors.New("unexpected end of JSON input")), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.New("unexpected end of JSON input")), - logType: errLog, - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - var dom smqsdk.Domain - sdkCall := sdkMock.On("UpdateDomain", mock.Anything, mock.Anything, tc.args[3]).Return(tc.domain, tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case entityLog: - err := json.Unmarshal([]byte(out), &dom) - assert.Nil(t, err) - assert.Equal(t, tc.domain, dom, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.domain, dom)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - } - sdkCall.Unset() - }) - } -} - -func TestEnableDomainCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - domainsCmd := cli.NewDomainsCmd() - rootCmd := setFlags(domainsCmd) - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - logType outputLog - }{ - { - desc: "enable domain successfully", - args: []string{ - domain.ID, - enableCmd, - validToken, - }, - logType: entityLog, - }, - { - desc: "enable domain with invalid token", - args: []string{ - domain.ID, - enableCmd, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - { - desc: "enable domain with invalid domain id", - args: []string{ - invalidID, - enableCmd, - token, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - { - desc: "enable domain with invalid args", - args: []string{ - domain.ID, - enableCmd, - validToken, - extraArg, - }, - logType: usageLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("EnableDomain", mock.Anything, tc.args[0], tc.args[2]).Return(tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - case okLog: - assert.True(t, strings.Contains(out, "ok"), fmt.Sprintf("%s unexpected response: expected success message, got: %v", tc.desc, out)) - } - - sdkCall.Unset() - }) - } -} - -func TestDisableDomainCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - domainsCmd := cli.NewDomainsCmd() - rootCmd := setFlags(domainsCmd) - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - logType outputLog - }{ - { - desc: "disable domain successfully", - args: []string{ - domain.ID, - disableCmd, - validToken, - }, - logType: okLog, - }, - { - desc: "disable domain with invalid token", - args: []string{ - domain.ID, - disableCmd, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - { - desc: "disable domain with invalid id", - args: []string{ - invalidID, - disableCmd, - token, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - { - desc: "disable domain with invalid args", - args: []string{ - domain.ID, - disableCmd, - validToken, - extraArg, - }, - logType: usageLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("DisableDomain", mock.Anything, tc.args[0], tc.args[2]).Return(tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - case okLog: - assert.True(t, strings.Contains(out, "ok"), fmt.Sprintf("%s unexpected response: expected success message, got: %v", tc.desc, out)) - } - - sdkCall.Unset() - }) - } -} - -func TestFreezeDomainCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - domainsCmd := cli.NewDomainsCmd() - rootCmd := setFlags(domainsCmd) - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - logType outputLog - }{ - { - desc: "freeze domain successfully", - args: []string{ - domain.ID, - freezeCmd, - validToken, - }, - logType: okLog, - }, - { - desc: "freeze domain with invalid token", - args: []string{ - domain.ID, - freezeCmd, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - { - desc: "freeze domain with invalid id", - args: []string{ - invalidID, - freezeCmd, - token, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - { - desc: "freeze domain with invalid args", - args: []string{ - domain.ID, - freezeCmd, - validToken, - extraArg, - }, - logType: usageLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("FreezeDomain", mock.Anything, tc.args[0], tc.args[2]).Return(tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - case okLog: - assert.True(t, strings.Contains(out, "ok"), fmt.Sprintf("%s unexpected response: expected success message, got: %v", tc.desc, out)) - } - - sdkCall.Unset() - }) - } -} - -func TestCreateDomainRoleCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - domainsCmd := cli.NewDomainsCmd() - rootCmd := setFlags(domainsCmd) - - roleReq := smqsdk.RoleReq{ - RoleName: "admin", - OptionalActions: []string{"read", "update"}, - } - roleReqJson, err := json.Marshal(roleReq) - assert.Nil(t, err, fmt.Sprintf("unexpected error: %v", err)) - - role := smqsdk.Role{ - ID: roleID, - Name: "admin", - } - - cases := []struct { - desc string - args []string - roleReq smqsdk.RoleReq - role smqsdk.Role - sdkErr errors.SDKError - errLogMessage string - logType outputLog - }{ - { - desc: "create role successfully", - args: []string{ - domain.ID, - rolesCmd, - createCmd, - string(roleReqJson), - token, - }, - role: role, - roleReq: roleReq, - logType: entityLog, - }, - { - desc: "create role with invalid args", - args: []string{ - domain.ID, - rolesCmd, - createCmd, - string(roleReqJson), - token, - extraArg, - }, - logType: usageLog, - }, - { - desc: "create role with invalid token", - args: []string{ - domain.ID, - rolesCmd, - createCmd, - string(roleReqJson), - invalidToken, - }, - roleReq: roleReq, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("CreateDomainRole", mock.Anything, tc.args[0], tc.roleReq, tc.args[4]).Return(tc.role, tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case entityLog: - var resp smqsdk.Role - err := json.Unmarshal([]byte(out), &resp) - assert.Nil(t, err) - assert.Equal(t, tc.role, resp, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.roleReq, role)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - sdkCall.Unset() - }) - } -} - -func TestGetDomainRoleCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - domainsCmd := cli.NewDomainsCmd() - rootCmd := setFlags(domainsCmd) - - role := smqsdk.Role{ - ID: roleID, - Name: "admin", - } - - cases := []struct { - desc string - args []string - role smqsdk.Role - sdkErr errors.SDKError - errLogMessage string - logType outputLog - }{ - { - desc: "get role successfully", - args: []string{ - domain.ID, - rolesCmd, - getCmd, - roleID, - token, - }, - role: role, - logType: entityLog, - }, - { - desc: "get role with invalid args", - args: []string{ - domain.ID, - rolesCmd, - getCmd, - roleID, - token, - extraArg, - }, - logType: usageLog, - }, - { - desc: "get role with invalid token", - args: []string{ - domain.ID, - rolesCmd, - getCmd, - roleID, - invalidToken, - }, - role: role, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("DomainRole", mock.Anything, tc.args[0], tc.args[3], tc.args[4]).Return(tc.role, tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case entityLog: - var role smqsdk.Role - err := json.Unmarshal([]byte(out), &role) - assert.Nil(t, err) - assert.Equal(t, tc.role, role, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.role, role)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - sdkCall.Unset() - }) - } -} - -func TestUpdateDomainRoleCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - domainsCmd := cli.NewDomainsCmd() - rootCmd := setFlags(domainsCmd) - - newRoleName := "new_name" - role := smqsdk.Role{ - ID: roleID, - Name: newRoleName, - } - - cases := []struct { - desc string - args []string - role smqsdk.Role - sdkErr errors.SDKError - errLogMessage string - logType outputLog - }{ - { - desc: "update role successfully", - args: []string{ - domain.ID, - rolesCmd, - updateCmd, - roleID, - newRoleName, - token, - }, - role: role, - logType: entityLog, - }, - { - desc: "update role with invalid args", - args: []string{ - domain.ID, - rolesCmd, - updateCmd, - roleID, - newRoleName, - token, - extraArg, - }, - logType: usageLog, - }, - { - desc: "update role with invalid token", - args: []string{ - domain.ID, - rolesCmd, - updateCmd, - roleID, - newRoleName, - invalidToken, - }, - role: role, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("UpdateDomainRole", mock.Anything, tc.args[0], tc.args[3], tc.args[4], tc.args[5]).Return(tc.role, tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case entityLog: - var role smqsdk.Role - err := json.Unmarshal([]byte(out), &role) - assert.Nil(t, err) - assert.Equal(t, tc.role, role, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.role, role)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - sdkCall.Unset() - }) - } -} - -func TestDeleteDomainRoleCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - domainsCmd := cli.NewDomainsCmd() - rootCmd := setFlags(domainsCmd) - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - logType outputLog - }{ - { - desc: "delete role successfully", - args: []string{ - domain.ID, - rolesCmd, - delCmd, - roleID, - token, - }, - logType: okLog, - }, - { - desc: "delete role with invalid token", - args: []string{ - domain.ID, - rolesCmd, - delCmd, - roleID, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - { - desc: "delete role with invalid args", - args: []string{ - domain.ID, - rolesCmd, - delCmd, - roleID, - token, - extraArg, - }, - logType: usageLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("DeleteDomainRole", mock.Anything, tc.args[0], tc.args[3], tc.args[4]).Return(tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case okLog: - assert.True(t, strings.Contains(out, "ok"), fmt.Sprintf("%s unexpected response: expected success message, got: %v", tc.desc, out)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - sdkCall.Unset() - }) - } -} - -func TestAddDomainRoleActionsCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - domainsCmd := cli.NewDomainsCmd() - rootCmd := setFlags(domainsCmd) - - cases := []struct { - desc string - args []string - actions []string - sdkErr errors.SDKError - errLogMessage string - logType outputLog - }{ - { - desc: "add actions to role successfully", - args: []string{ - domain.ID, - rolesCmd, - actionsCmd, - addCmd, - roleID, - `{"actions":["read","write"]}`, - token, - }, - actions: []string{"read", "write"}, - logType: entityLog, - }, - { - desc: "add actions to role with invalid args", - args: []string{ - domain.ID, - rolesCmd, - actionsCmd, - addCmd, - roleID, - `{"actions":["read","write"]}`, - token, - extraArg, - }, - logType: usageLog, - }, - { - desc: "add actions to role with invalid token", - args: []string{ - domain.ID, - rolesCmd, - actionsCmd, - addCmd, - roleID, - `{"actions":["read","write"]}`, - invalidToken, - }, - actions: []string{"read", "write"}, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("AddDomainRoleActions", mock.Anything, tc.args[0], tc.args[4], tc.actions, tc.args[6]).Return(tc.actions, tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case entityLog: - var actions []string - err := json.Unmarshal([]byte(out), &actions) - assert.Nil(t, err) - assert.Equal(t, tc.actions, actions, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.actions, actions)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - sdkCall.Unset() - }) - } -} - -func TestListDomainRoleActionsCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - domainsCmd := cli.NewDomainsCmd() - rootCmd := setFlags(domainsCmd) - - cases := []struct { - desc string - args []string - actions []string - sdkErr errors.SDKError - errLogMessage string - logType outputLog - }{ - { - desc: "list actions of role successfully", - args: []string{ - domain.ID, - rolesCmd, - actionsCmd, - listCmd, - roleID, - token, - }, - actions: []string{"read", "write"}, - logType: entityLog, - }, - { - desc: "list actions of role with invalid args", - args: []string{ - domain.ID, - rolesCmd, - actionsCmd, - listCmd, - roleID, - token, - extraArg, - }, - logType: usageLog, - }, - { - desc: "list actions of role with invalid token", - args: []string{ - domain.ID, - rolesCmd, - actionsCmd, - listCmd, - roleID, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("DomainRoleActions", mock.Anything, tc.args[0], tc.args[4], tc.args[5]).Return(tc.actions, tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case entityLog: - var actions []string - err := json.Unmarshal([]byte(out), &actions) - assert.Nil(t, err) - assert.Equal(t, tc.actions, actions, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.actions, actions)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - sdkCall.Unset() - }) - } -} - -func TestDeleteDomainRoleActionsCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - domainsCmd := cli.NewDomainsCmd() - rootCmd := setFlags(domainsCmd) - - cases := []struct { - desc string - args []string - actions []string - sdkErr errors.SDKError - errLogMessage string - logType outputLog - }{ - { - desc: "delete actions from role successfully", - args: []string{ - domain.ID, - rolesCmd, - actionsCmd, - delCmd, - roleID, - `{"actions":["read","write"]}`, - token, - }, - actions: []string{"read", "write"}, - logType: okLog, - }, - { - desc: "delete all actions from role successfully", - args: []string{ - domain.ID, - rolesCmd, - actionsCmd, - delCmd, - roleID, - all, - token, - }, - logType: okLog, - }, - { - desc: "delete actions from role with invalid args", - args: []string{ - domain.ID, - rolesCmd, - actionsCmd, - delCmd, - roleID, - `{"actions":["read","write"]}`, - token, - extraArg, - }, - logType: usageLog, - }, - { - desc: "delete actions from role with invalid token", - args: []string{ - domain.ID, - rolesCmd, - actionsCmd, - delCmd, - roleID, - `{"actions":["read","write"]}`, - invalidToken, - }, - actions: []string{"read", "write"}, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - var sdkCall *mock.Call - if tc.args[5] == all { - sdkCall = sdkMock.On("RemoveAllDomainRoleActions", mock.Anything, tc.args[0], tc.args[4], tc.args[6]).Return(tc.sdkErr) - } else { - sdkCall = sdkMock.On("RemoveDomainRoleActions", mock.Anything, tc.args[0], tc.args[4], tc.actions, tc.args[6]).Return(tc.sdkErr) - } - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case okLog: - assert.True(t, strings.Contains(out, "ok"), fmt.Sprintf("%s unexpected response: expected success message, got: %v", tc.desc, out)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - sdkCall.Unset() - }) - } -} - -func TestAvailableDomainRoleActionsCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - domainsCmd := cli.NewDomainsCmd() - rootCmd := setFlags(domainsCmd) - - cases := []struct { - desc string - args []string - actions []string - sdkErr errors.SDKError - errLogMessage string - logType outputLog - }{ - { - desc: "list available actions successfully", - args: []string{ - domain.ID, - rolesCmd, - actionsCmd, - availableActionsCmd, - token, - }, - actions: []string{"read", "write", "update"}, - logType: entityLog, - }, - { - desc: "list available actions with invalid args", - args: []string{ - domain.ID, - rolesCmd, - actionsCmd, - availableActionsCmd, - token, - extraArg, - }, - logType: usageLog, - }, - { - desc: "list available actions with invalid token", - args: []string{ - domain.ID, - rolesCmd, - actionsCmd, - availableActionsCmd, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("AvailableDomainRoleActions", mock.Anything, tc.args[4]).Return(tc.actions, tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case entityLog: - var actions []string - err := json.Unmarshal([]byte(out), &actions) - assert.Nil(t, err) - assert.Equal(t, tc.actions, actions, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.actions, actions)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - sdkCall.Unset() - }) - } -} - -func TestAddDomainRoleMembersCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - domainsCmd := cli.NewDomainsCmd() - rootCmd := setFlags(domainsCmd) - - members := []string{"5dc1ce4b-7cc9-4f12-98a6-9d74cc4980bb", "5dc1ce4b-7cc9-4f12-98a6-9d74cc4980bb"} - membersJson := `{"members":["5dc1ce4b-7cc9-4f12-98a6-9d74cc4980bb", "5dc1ce4b-7cc9-4f12-98a6-9d74cc4980bb"]}` - - cases := []struct { - desc string - args []string - members []string - sdkErr errors.SDKError - errLogMessage string - logType outputLog - }{ - { - desc: "add members to role successfully", - args: []string{ - domain.ID, - rolesCmd, - membersCmd, - addCmd, - roleID, - membersJson, - token, - }, - members: members, - logType: entityLog, - }, - { - desc: "add members to role with invalid args", - args: []string{ - domain.ID, - rolesCmd, - membersCmd, - addCmd, - roleID, - membersJson, - token, - extraArg, - }, - logType: usageLog, - }, - { - desc: "add members to role with invalid token", - args: []string{ - domain.ID, - rolesCmd, - membersCmd, - addCmd, - roleID, - membersJson, - invalidToken, - }, - members: members, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("AddDomainRoleMembers", mock.Anything, tc.args[0], tc.args[4], tc.members, tc.args[6]).Return(tc.members, tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case entityLog: - var members []string - err := json.Unmarshal([]byte(out), &members) - assert.Nil(t, err) - assert.Equal(t, tc.members, members, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.members, members)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - sdkCall.Unset() - }) - } -} - -func TestListDomainRoleMembersCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - domainsCmd := cli.NewDomainsCmd() - rootCmd := setFlags(domainsCmd) - - page := smqsdk.RoleMembersPage{ - Total: 1, - Offset: 0, - Limit: 10, - Members: []string{"5dc1ce4b-7cc9-4f12-98a6-9d74cc4980bb"}, - } - - cases := []struct { - desc string - args []string - page smqsdk.RoleMembersPage - sdkErr errors.SDKError - errLogMessage string - logType outputLog - }{ - { - desc: "list members of role successfully", - args: []string{ - domain.ID, - rolesCmd, - membersCmd, - listCmd, - roleID, - token, - }, - page: page, - logType: entityLog, - }, - { - desc: "list members of role with invalid args", - args: []string{ - domain.ID, - rolesCmd, - membersCmd, - listCmd, - roleID, - token, - extraArg, - }, - logType: usageLog, - }, - { - desc: "list members of role with invalid token", - args: []string{ - domain.ID, - rolesCmd, - membersCmd, - listCmd, - roleID, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("DomainRoleMembers", mock.Anything, tc.args[0], tc.args[4], mock.Anything, tc.args[5]).Return(tc.page, tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case entityLog: - var page smqsdk.RoleMembersPage - err := json.Unmarshal([]byte(out), &page) - assert.Nil(t, err) - assert.Equal(t, tc.page, page, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.page, page)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - sdkCall.Unset() - }) - } -} - -func TestDeleteDomainRoleMembersCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - domainsCmd := cli.NewDomainsCmd() - rootCmd := setFlags(domainsCmd) - - members := []string{"5dc1ce4b-7cc9-4f12-98a6-9d74cc4980bb", "5dc1ce4b-7cc9-4f12-98a6-9d74cc4980bb"} - membersJson := `{"members":["5dc1ce4b-7cc9-4f12-98a6-9d74cc4980bb", "5dc1ce4b-7cc9-4f12-98a6-9d74cc4980bb"]}` - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - logType outputLog - }{ - { - desc: "delete members from role successfully", - args: []string{ - domain.ID, - rolesCmd, - membersCmd, - delCmd, - roleID, - membersJson, - token, - }, - logType: okLog, - }, - { - desc: "delete all members from role successfully", - args: []string{ - domain.ID, - rolesCmd, - membersCmd, - delCmd, - roleID, - all, - token, - }, - logType: okLog, - }, - { - desc: "delete members from role with invalid args", - args: []string{ - domain.ID, - rolesCmd, - membersCmd, - delCmd, - roleID, - membersJson, - token, - extraArg, - }, - logType: usageLog, - }, - { - desc: "delete members from role with invalid token", - args: []string{ - domain.ID, - rolesCmd, - membersCmd, - delCmd, - roleID, - membersJson, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - var sdkCall *mock.Call - if tc.args[5] == all { - sdkCall = sdkMock.On("RemoveAllDomainRoleMembers", mock.Anything, tc.args[0], tc.args[4], tc.args[6]).Return(tc.sdkErr) - } else { - sdkCall = sdkMock.On("RemoveDomainRoleMembers", mock.Anything, tc.args[0], tc.args[4], members, tc.args[6]).Return(tc.sdkErr) - } - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case okLog: - assert.True(t, strings.Contains(out, "ok"), fmt.Sprintf("%s unexpected response: expected success message, got: %v", tc.desc, out)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - sdkCall.Unset() - }) - } -} diff --git a/cli/groups.go b/cli/groups.go deleted file mode 100644 index a30adc2af..000000000 --- a/cli/groups.go +++ /dev/null @@ -1,598 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package cli - -import ( - "encoding/json" - "fmt" - - "github.com/absmach/magistrala/groups" - smqsdk "github.com/absmach/magistrala/pkg/sdk" - "github.com/spf13/cobra" -) - -const ( - tags = "tags" - add = "add" - list = "list" - availableActions = "available-actions" - - // Usage strings for group operations. - usageGroupCreate = "cli groups create " - usageGroupGet = "cli groups get " - usageGroupUpdate = "cli groups update " - usageGroupUpdateTags = "cli groups update tags " - usageGroupDelete = "cli groups delete " - usageGroupEnable = "cli groups enable " - usageGroupDisable = "cli groups disable " - - // Usage strings for group roles operations. - usageGroupRolesCreate = "cli groups roles create " - usageGroupRolesGet = "cli groups roles get " - usageGroupRolesUpdate = "cli groups roles update " - usageGroupRolesDelete = "cli groups roles delete " - - // Usage strings for group role actions operations. - usageGroupRoleActionsAdd = "cli groups roles actions add " - usageGroupRoleActionsList = "cli groups roles actions list " - usageGroupRoleActionsDelete = "cli groups roles actions delete " - usageGroupRoleActionsAvailable = "cli groups roles actions available-actions " - - // Usage strings for group role members operations. - usageGroupRoleMembersAdd = "cli groups roles members add " - usageGroupRoleMembersList = "cli groups roles members list " - usageGroupRoleMembersDelete = "cli groups roles members delete " -) - -func NewGroupsCmd() *cobra.Command { - cmd := &cobra.Command{ - Use: "groups [operation] [args...]", - Short: "Groups management", - Long: `Format: - groups create [args...] - groups [args...] - -Operations (require group_id/all): get, update, delete, enable, disable, roles - -Examples: - groups create - groups all get - groups get - groups update - groups update tags - groups delete - groups enable - groups disable `, - - Run: func(cmd *cobra.Command, args []string) { - if len(args) == 0 { - logUsageCmd(*cmd, cmd.Use) - return - } - - if args[0] == create { - handleGroupCreate(cmd, args[1:]) - return - } - - if len(args) < 2 { - logUsageCmd(*cmd, "groups [args...]") - return - } - - groupParams := args[0] - operation := args[1] - opArgs := args[2:] - - switch operation { - case get: - handleGroupGet(cmd, groupParams, opArgs) - case update: - handleGroupUpdate(cmd, groupParams, opArgs) - case delete: - handleGroupDelete(cmd, groupParams, opArgs) - case enable: - handleGroupEnable(cmd, groupParams, opArgs) - case disable: - handleGroupDisable(cmd, groupParams, opArgs) - case roles: - handleGroupRoles(cmd, groupParams, opArgs) - default: - logErrorCmd(*cmd, fmt.Errorf("unknown operation: %s", operation)) - } - }, - } - - return cmd -} - -func handleGroupCreate(cmd *cobra.Command, args []string) { - if len(args) != 3 { - logUsageCmd(*cmd, usageGroupCreate) - return - } - - var group smqsdk.Group - if err := json.Unmarshal([]byte(args[0]), &group); err != nil { - logErrorCmd(*cmd, err) - return - } - group.Status = groups.EnabledStatus.String() - group, err := sdk.CreateGroup(cmd.Context(), group, args[1], args[2]) - if err != nil { - logErrorCmd(*cmd, err) - return - } - logJSONCmd(*cmd, group) -} - -func handleGroupGet(cmd *cobra.Command, groupParams string, args []string) { - if len(args) != 2 { - logUsageCmd(*cmd, usageGroupGet) - return - } - - if groupParams == all { - metadata, err := convertMetadata(Metadata) - if err != nil { - logErrorCmd(*cmd, err) - return - } - - pageMetadata := smqsdk.PageMetadata{ - Name: Name, - Offset: Offset, - Limit: Limit, - Metadata: metadata, - } - - l, err := sdk.Groups(cmd.Context(), pageMetadata, args[0], args[1]) - if err != nil { - logErrorCmd(*cmd, err) - return - } - logJSONCmd(*cmd, l) - return - } - - g, err := sdk.Group(cmd.Context(), groupParams, args[0], args[1]) - if err != nil { - logErrorCmd(*cmd, err) - return - } - - logJSONCmd(*cmd, g) -} - -func handleGroupUpdate(cmd *cobra.Command, groupID string, args []string) { - if len(args) < 3 || len(args) > 4 { - if args[0] == tags { - logUsageCmd(*cmd, usageGroupUpdateTags) - return - } - logUsageCmd(*cmd, usageGroupUpdate) - return - } - - if len(args) == 4 && args[0] == tags { - var group smqsdk.Group - if err := json.Unmarshal([]byte(args[1]), &group.Tags); err != nil { - logErrorCmd(*cmd, err) - return - } - group.ID = groupID - group, err := sdk.UpdateGroupTags(cmd.Context(), group, args[2], args[3]) - if err != nil { - logErrorCmd(*cmd, err) - return - } - logJSONCmd(*cmd, group) - return - } - - if len(args) != 3 { - logUsageCmd(*cmd, usageGroupUpdate) - return - } - - var group smqsdk.Group - if err := json.Unmarshal([]byte(args[0]), &group); err != nil { - logErrorCmd(*cmd, err) - return - } - - group.ID = groupID - group, err := sdk.UpdateGroup(cmd.Context(), group, args[1], args[2]) - if err != nil { - logErrorCmd(*cmd, err) - return - } - - logJSONCmd(*cmd, group) -} - -func handleGroupDelete(cmd *cobra.Command, groupID string, args []string) { - if len(args) != 2 { - logUsageCmd(*cmd, usageGroupDelete) - return - } - - if err := sdk.DeleteGroup(cmd.Context(), groupID, args[0], args[1]); err != nil { - logErrorCmd(*cmd, err) - return - } - logOKCmd(*cmd) -} - -func handleGroupEnable(cmd *cobra.Command, groupID string, args []string) { - if len(args) != 2 { - logUsageCmd(*cmd, usageGroupEnable) - return - } - - group, err := sdk.EnableGroup(cmd.Context(), groupID, args[0], args[1]) - if err != nil { - logErrorCmd(*cmd, err) - return - } - - logJSONCmd(*cmd, group) -} - -func handleGroupDisable(cmd *cobra.Command, groupID string, args []string) { - if len(args) != 2 { - logUsageCmd(*cmd, usageGroupDisable) - return - } - - group, err := sdk.DisableGroup(cmd.Context(), groupID, args[0], args[1]) - if err != nil { - logErrorCmd(*cmd, err) - return - } - - logJSONCmd(*cmd, group) -} - -func handleGroupRoles(cmd *cobra.Command, groupID string, args []string) { - if len(args) < 1 { - logUsageCmd(*cmd, "cli groups roles [args...]") - return - } - - operation := args[0] - opArgs := args[1:] - - switch operation { - case create: - handleGroupRoleCreate(cmd, groupID, opArgs) - case get: - handleGroupRoleGet(cmd, groupID, opArgs) - case update: - handleGroupRoleUpdate(cmd, groupID, opArgs) - case delete: - handleGroupRoleDelete(cmd, groupID, opArgs) - case actions: - handleGroupRoleActions(cmd, groupID, opArgs) - case members: - handleGroupRoleMembers(cmd, groupID, opArgs) - default: - logErrorCmd(*cmd, fmt.Errorf("unknown roles operation: %s", operation)) - } -} - -func handleGroupRoleCreate(cmd *cobra.Command, groupID string, args []string) { - if len(args) != 3 { - logUsageCmd(*cmd, usageGroupRolesCreate) - return - } - - var roleReq smqsdk.RoleReq - if err := json.Unmarshal([]byte(args[0]), &roleReq); err != nil { - logErrorCmd(*cmd, err) - return - } - - r, err := sdk.CreateGroupRole(cmd.Context(), groupID, args[1], roleReq, args[2]) - if err != nil { - logErrorCmd(*cmd, err) - return - } - - logJSONCmd(*cmd, r) -} - -func handleGroupRoleGet(cmd *cobra.Command, groupID string, args []string) { - if len(args) != 3 { - logUsageCmd(*cmd, usageGroupRolesGet) - return - } - - roleID := args[0] - domainID := args[1] - token := args[2] - - if roleID == all { - pageMetadata := smqsdk.PageMetadata{ - Offset: Offset, - Limit: Limit, - } - rs, err := sdk.GroupRoles(cmd.Context(), groupID, domainID, pageMetadata, token) - if err != nil { - logErrorCmd(*cmd, err) - return - } - logJSONCmd(*cmd, rs) - return - } - - r, err := sdk.GroupRole(cmd.Context(), groupID, roleID, domainID, token) - if err != nil { - logErrorCmd(*cmd, err) - return - } - logJSONCmd(*cmd, r) -} - -func handleGroupRoleUpdate(cmd *cobra.Command, groupID string, args []string) { - if len(args) != 4 { - logUsageCmd(*cmd, usageGroupRolesUpdate) - return - } - - roleID := args[0] - newName := args[1] - domainID := args[2] - token := args[3] - - r, err := sdk.UpdateGroupRole(cmd.Context(), groupID, roleID, newName, domainID, token) - if err != nil { - logErrorCmd(*cmd, err) - return - } - logJSONCmd(*cmd, r) -} - -func handleGroupRoleDelete(cmd *cobra.Command, groupID string, args []string) { - if len(args) != 3 { - logUsageCmd(*cmd, usageGroupRolesDelete) - return - } - - roleID := args[0] - domainID := args[1] - token := args[2] - - if err := sdk.DeleteGroupRole(cmd.Context(), groupID, roleID, domainID, token); err != nil { - logErrorCmd(*cmd, err) - return - } - logOKCmd(*cmd) -} - -func handleGroupRoleActions(cmd *cobra.Command, groupID string, args []string) { - if len(args) < 1 { - logUsageCmd(*cmd, "cli groups roles actions [args...]") - return - } - - operation := args[0] - opArgs := args[1:] - - switch operation { - case add: - handleGroupRoleActionsAdd(cmd, groupID, opArgs) - case list: - handleGroupRoleActionsList(cmd, groupID, opArgs) - case delete: - handleGroupRoleActionsDelete(cmd, groupID, opArgs) - case availableActions: - handleGroupRoleActionsAvailable(cmd, opArgs) - default: - logErrorCmd(*cmd, fmt.Errorf("unknown actions operation: %s", operation)) - } -} - -func handleGroupRoleActionsAdd(cmd *cobra.Command, groupID string, args []string) { - if len(args) != 4 { - logUsageCmd(*cmd, usageGroupRoleActionsAdd) - return - } - - roleID := args[0] - actionsJSON := args[1] - domainID := args[2] - token := args[3] - - actions := struct { - Actions []string `json:"actions"` - }{} - if err := json.Unmarshal([]byte(actionsJSON), &actions); err != nil { - logErrorCmd(*cmd, err) - return - } - - acts, err := sdk.AddGroupRoleActions(cmd.Context(), groupID, roleID, domainID, actions.Actions, token) - if err != nil { - logErrorCmd(*cmd, err) - return - } - logJSONCmd(*cmd, acts) -} - -func handleGroupRoleActionsList(cmd *cobra.Command, groupID string, args []string) { - if len(args) != 3 { - logUsageCmd(*cmd, usageGroupRoleActionsList) - return - } - - roleID := args[0] - domainID := args[1] - token := args[2] - - l, err := sdk.GroupRoleActions(cmd.Context(), groupID, roleID, domainID, token) - if err != nil { - logErrorCmd(*cmd, err) - return - } - logJSONCmd(*cmd, l) -} - -func handleGroupRoleActionsDelete(cmd *cobra.Command, groupID string, args []string) { - if len(args) != 4 { - logUsageCmd(*cmd, usageGroupRoleActionsDelete) - return - } - - roleID := args[0] - actionsJSON := args[1] - domainID := args[2] - token := args[3] - - if actionsJSON == all { - if err := sdk.RemoveAllGroupRoleActions(cmd.Context(), groupID, roleID, domainID, token); err != nil { - logErrorCmd(*cmd, err) - return - } - logOKCmd(*cmd) - return - } - - actions := struct { - Actions []string `json:"actions"` - }{} - if err := json.Unmarshal([]byte(actionsJSON), &actions); err != nil { - logErrorCmd(*cmd, err) - return - } - - if err := sdk.RemoveGroupRoleActions(cmd.Context(), groupID, roleID, domainID, actions.Actions, token); err != nil { - logErrorCmd(*cmd, err) - return - } - logOKCmd(*cmd) -} - -func handleGroupRoleActionsAvailable(cmd *cobra.Command, args []string) { - if len(args) != 2 { - logUsageCmd(*cmd, usageGroupRoleActionsAvailable) - return - } - - domainID := args[0] - token := args[1] - - acts, err := sdk.AvailableGroupRoleActions(cmd.Context(), domainID, token) - if err != nil { - logErrorCmd(*cmd, err) - return - } - logJSONCmd(*cmd, acts) -} - -func handleGroupRoleMembers(cmd *cobra.Command, groupID string, args []string) { - if len(args) < 1 { - logUsageCmd(*cmd, "cli groups roles members [args...]") - return - } - - operation := args[0] - opArgs := args[1:] - - switch operation { - case add: - handleGroupRoleMembersAdd(cmd, groupID, opArgs) - case list: - handleGroupRoleMembersList(cmd, groupID, opArgs) - case delete: - handleGroupRoleMembersDelete(cmd, groupID, opArgs) - default: - logErrorCmd(*cmd, fmt.Errorf("unknown members operation: %s", operation)) - } -} - -func handleGroupRoleMembersAdd(cmd *cobra.Command, groupID string, args []string) { - if len(args) != 4 { - logUsageCmd(*cmd, usageGroupRoleMembersAdd) - return - } - - roleID := args[0] - membersJSON := args[1] - domainID := args[2] - token := args[3] - - members := struct { - Members []string `json:"members"` - }{} - if err := json.Unmarshal([]byte(membersJSON), &members); err != nil { - logErrorCmd(*cmd, err) - return - } - - memb, err := sdk.AddGroupRoleMembers(cmd.Context(), groupID, roleID, domainID, members.Members, token) - if err != nil { - logErrorCmd(*cmd, err) - return - } - logJSONCmd(*cmd, memb) -} - -func handleGroupRoleMembersList(cmd *cobra.Command, groupID string, args []string) { - if len(args) != 3 { - logUsageCmd(*cmd, usageGroupRoleMembersList) - return - } - - roleID := args[0] - domainID := args[1] - token := args[2] - - pageMetadata := smqsdk.PageMetadata{ - Offset: Offset, - Limit: Limit, - } - - l, err := sdk.GroupRoleMembers(cmd.Context(), groupID, roleID, domainID, pageMetadata, token) - if err != nil { - logErrorCmd(*cmd, err) - return - } - logJSONCmd(*cmd, l) -} - -func handleGroupRoleMembersDelete(cmd *cobra.Command, groupID string, args []string) { - if len(args) != 4 { - logUsageCmd(*cmd, usageGroupRoleMembersDelete) - return - } - - roleID := args[0] - membersJSON := args[1] - domainID := args[2] - token := args[3] - - if membersJSON == all { - if err := sdk.RemoveAllGroupRoleMembers(cmd.Context(), groupID, roleID, domainID, token); err != nil { - logErrorCmd(*cmd, err) - return - } - logOKCmd(*cmd) - return - } - - members := struct { - Members []string `json:"members"` - }{} - if err := json.Unmarshal([]byte(membersJSON), &members); err != nil { - logErrorCmd(*cmd, err) - return - } - - if err := sdk.RemoveGroupRoleMembers(cmd.Context(), groupID, roleID, domainID, members.Members, token); err != nil { - logErrorCmd(*cmd, err) - return - } - logOKCmd(*cmd) -} diff --git a/cli/groups_test.go b/cli/groups_test.go deleted file mode 100644 index c7e9ab948..000000000 --- a/cli/groups_test.go +++ /dev/null @@ -1,1322 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package cli_test - -import ( - "encoding/json" - "fmt" - "net/http" - "strings" - "testing" - - "github.com/absmach/magistrala/cli" - "github.com/absmach/magistrala/internal/testsutil" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - smqsdk "github.com/absmach/magistrala/pkg/sdk" - sdkmocks "github.com/absmach/magistrala/pkg/sdk/mocks" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" -) - -const ( - tagUpdateType = "tags" - newTagsJson = "[\"tag1\", \"tag2\"]" -) - -var group = smqsdk.Group{ - ID: testsutil.GenerateUUID(&testing.T{}), - Name: "testgroup", -} - -func TestCreateGroupCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - groupJson := "{\"name\":\"testgroup\", \"metadata\":{\"key1\":\"value1\"}}" - groupCmd := cli.NewGroupsCmd() - rootCmd := setFlags(groupCmd) - - gp := smqsdk.Group{} - cases := []struct { - desc string - args []string - logType outputLog - group smqsdk.Group - sdkErr errors.SDKError - errLogMessage string - }{ - { - desc: "create group successfully", - args: []string{ - createCmd, - groupJson, - domainID, - token, - }, - group: group, - logType: entityLog, - }, - { - desc: "create group with invalid args", - args: []string{ - createCmd, - groupJson, - domainID, - token, - extraArg, - }, - logType: usageLog, - }, - { - desc: "create group with invalid json", - args: []string{ - createCmd, - "{\"name\":\"testgroup\", \"metadata\":{\"key1\":\"value1\"}", - domainID, - token, - }, - sdkErr: errors.NewSDKError(errors.New("unexpected end of JSON input")), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.New("unexpected end of JSON input")), - logType: errLog, - }, - { - desc: "create group with invalid token", - args: []string{ - createCmd, - groupJson, - domainID, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusUnauthorized), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusUnauthorized)), - logType: errLog, - }, - { - desc: "create group with invalid domain", - args: []string{ - createCmd, - groupJson, - domainID, - token, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrDomainAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrDomainAuthorization, http.StatusForbidden)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("CreateGroup", mock.Anything, mock.Anything, tc.args[2], tc.args[3]).Return(tc.group, tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case entityLog: - err := json.Unmarshal([]byte(out), &gp) - assert.Nil(t, err) - assert.Equal(t, tc.group, gp, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.group, gp)) - case usageLog: - assert.True(t, strings.Contains(out, "cli groups create"), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - } - sdkCall.Unset() - }) - } -} - -func TestDeletegroupCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - groupCmd := cli.NewGroupsCmd() - rootCmd := setFlags(groupCmd) - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - logType outputLog - errLogMessage string - }{ - { - desc: "delete group successfully", - args: []string{ - group.ID, - delCmd, - domainID, - token, - }, - logType: okLog, - }, - { - desc: "delete group with invalid args", - args: []string{ - group.ID, - delCmd, - domainID, - token, - extraArg, - }, - logType: usageLog, - }, - { - desc: "delete group with invalid id", - args: []string{ - invalidID, - delCmd, - domainID, - token, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - { - desc: "delete group with invalid token", - args: []string{ - group.ID, - delCmd, - domainID, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("DeleteGroup", mock.Anything, tc.args[0], tc.args[2], tc.args[3]).Return(tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case okLog: - assert.True(t, strings.Contains(out, "ok"), fmt.Sprintf("%s unexpected response: expected success message, got: %v", tc.desc, out)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - sdkCall.Unset() - }) - } -} - -func TestUpdategroupCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - groupCmd := cli.NewGroupsCmd() - rootCmd := setFlags(groupCmd) - - newTagString := []string{"tag1", "tag2"} - - newGroupJson := fmt.Sprintf("{\"id\":\"%s\",\"name\" : \"newgroup\"}", group.ID) - cases := []struct { - desc string - args []string - group smqsdk.Group - sdkErr errors.SDKError - errLogMessage string - logType outputLog - }{ - { - desc: "update group successfully", - args: []string{ - group.ID, - updateCmd, - newGroupJson, - domainID, - token, - }, - group: smqsdk.Group{ - Name: "newgroup1", - ID: group.ID, - }, - logType: entityLog, - }, - { - desc: "update group with invalid args", - args: []string{ - group.ID, - updateCmd, - newGroupJson, - domainID, - token, - extraArg, - }, - logType: usageLog, - }, - { - desc: "update group with invalid group id", - args: []string{ - invalidID, - updateCmd, - fmt.Sprintf("{\"id\":\"%s\",\"name\" : \"group1\"}", invalidID), - domainID, - token, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - { - desc: "update group with invalid json syntax", - args: []string{ - group.ID, - updateCmd, - fmt.Sprintf("{\"id\":\"%s\",\"name\" : \"group1\"", group.ID), - domainID, - token, - }, - sdkErr: errors.NewSDKError(errors.New("unexpected end of JSON input")), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.New("unexpected end of JSON input")), - logType: errLog, - }, - { - desc: "update group tags successfully", - args: []string{ - group.ID, - updateCmd, - tagUpdateType, - newTagsJson, - domainID, - token, - }, - group: smqsdk.Group{ - Name: group.Name, - ID: group.ID, - DomainID: group.DomainID, - Status: group.Status, - Tags: newTagString, - }, - logType: entityLog, - }, - { - desc: "update group with invalid tags", - args: []string{ - group.ID, - updateCmd, - tagUpdateType, - "[\"tag1\", \"tag2\"", - domainID, - token, - }, - logType: errLog, - sdkErr: errors.NewSDKError(errors.New("unexpected end of JSON input")), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.New("unexpected end of JSON input")), - }, - { - desc: "update group tags with invalid group id", - args: []string{ - invalidID, - updateCmd, - tagUpdateType, - newTagsJson, - domainID, - token, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - var ch smqsdk.Group - sdkCall := sdkMock.On("UpdateGroup", mock.Anything, mock.Anything, tc.args[3], tc.args[4]).Return(tc.group, tc.sdkErr) - sdkCall1 := sdkMock.On("UpdateGroupTags", mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(tc.group, tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case entityLog: - err := json.Unmarshal([]byte(out), &ch) - assert.Nil(t, err) - assert.Equal(t, tc.group, ch, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.group, ch)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - } - sdkCall.Unset() - sdkCall1.Unset() - }) - } -} - -func TestEnablegroupCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - groupCmd := cli.NewGroupsCmd() - rootCmd := setFlags(groupCmd) - var ch smqsdk.Group - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - group smqsdk.Group - logType outputLog - }{ - { - desc: "enable group successfully", - args: []string{ - group.ID, - enableCmd, - domainID, - validToken, - }, - group: group, - logType: entityLog, - }, - { - desc: "delete group with invalid token", - args: []string{ - group.ID, - enableCmd, - domainID, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - { - desc: "delete group with invalid group ID", - args: []string{ - invalidID, - enableCmd, - domainID, - token, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - { - desc: "enable group with invalid args", - args: []string{ - group.ID, - enableCmd, - domainID, - validToken, - extraArg, - }, - logType: usageLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("EnableGroup", mock.Anything, tc.args[0], tc.args[2], tc.args[3]).Return(tc.group, tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - case entityLog: - err := json.Unmarshal([]byte(out), &ch) - assert.Nil(t, err) - assert.Equal(t, tc.group, ch, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.group, ch)) - } - - sdkCall.Unset() - }) - } -} - -func TestDisablegroupCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - groupsCmd := cli.NewGroupsCmd() - rootCmd := setFlags(groupsCmd) - - var ch smqsdk.Group - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - group smqsdk.Group - logType outputLog - }{ - { - desc: "disable group successfully", - args: []string{ - group.ID, - disableCmd, - domainID, - validToken, - }, - logType: entityLog, - group: group, - }, - { - desc: "disable group with invalid token", - args: []string{ - group.ID, - disableCmd, - domainID, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - { - desc: "disable group with invalid id", - args: []string{ - invalidID, - disableCmd, - domainID, - token, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - { - desc: "disable group with invalid args", - args: []string{ - group.ID, - disableCmd, - domainID, - validToken, - extraArg, - }, - logType: usageLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("DisableGroup", mock.Anything, tc.args[0], tc.args[2], tc.args[3]).Return(tc.group, tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - case entityLog: - err := json.Unmarshal([]byte(out), &ch) - if err != nil { - t.Fatalf("json.Unmarshal failed: %v", err) - } - assert.Equal(t, tc.group, ch, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.group, ch)) - } - - sdkCall.Unset() - }) - } -} - -func TestCreateGroupRoleCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - groupsCmd := cli.NewGroupsCmd() - rootCmd := setFlags(groupsCmd) - - roleReq := smqsdk.RoleReq{ - RoleName: "admin", - OptionalActions: []string{"read", "update"}, - } - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - role smqsdk.Role - logType outputLog - }{ - { - desc: "create group role successfully", - args: []string{ - group.ID, - rolesCmd, - createCmd, - `{"role_name":"admin","optional_actions":["read","update"]}`, - domainID, - token, - }, - role: smqsdk.Role{ - ID: testsutil.GenerateUUID(&testing.T{}), - Name: "admin", - OptionalActions: []string{"read", "update"}, - }, - logType: entityLog, - }, - { - desc: "create group role with invalid JSON", - args: []string{ - group.ID, - rolesCmd, - createCmd, - `{"role_name":"admin","optional_actions":["read","update"}`, - domainID, - token, - }, - sdkErr: errors.NewSDKError(errors.New("invalid character '}' after array element")), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.New("invalid character '}' after array element")), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("CreateGroupRole", mock.Anything, tc.args[0], tc.args[4], roleReq, tc.args[5]).Return(tc.role, tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case entityLog: - var role smqsdk.Role - err := json.Unmarshal([]byte(out), &role) - assert.Nil(t, err) - assert.Equal(t, tc.role, role, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.role, role)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - } - - sdkCall.Unset() - }) - } -} - -func TestGetGroupRolesCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - groupsCmd := cli.NewGroupsCmd() - rootCmd := setFlags(groupsCmd) - - role := smqsdk.Role{ - ID: testsutil.GenerateUUID(&testing.T{}), - Name: "admin", - OptionalActions: []string{"read", "update"}, - } - rolesPage := smqsdk.RolesPage{ - Total: 1, - Offset: 0, - Limit: 10, - Roles: []smqsdk.Role{role}, - } - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - roles smqsdk.RolesPage - logType outputLog - }{ - { - desc: "get all group roles successfully", - args: []string{ - group.ID, - rolesCmd, - getCmd, - all, - domainID, - token, - }, - roles: rolesPage, - logType: entityLog, - }, - { - desc: "get group roles with invalid token", - args: []string{ - group.ID, - rolesCmd, - getCmd, - all, - domainID, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("GroupRoles", mock.Anything, tc.args[0], tc.args[4], mock.Anything, tc.args[5]).Return(tc.roles, tc.sdkErr) - if tc.args[3] != all { - sdkCall = sdkMock.On("GroupRole", mock.Anything, tc.args[0], tc.args[3], tc.args[4], tc.args[5]).Return(role, tc.sdkErr) - } - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case entityLog: - var roles smqsdk.RolesPage - err := json.Unmarshal([]byte(out), &roles) - assert.Nil(t, err) - assert.Equal(t, tc.roles, roles, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.roles, roles)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - } - - sdkCall.Unset() - }) - } -} - -func TestUpdateGroupRoleCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - groupsCmd := cli.NewGroupsCmd() - rootCmd := setFlags(groupsCmd) - - role := smqsdk.Role{ - ID: testsutil.GenerateUUID(&testing.T{}), - Name: "new_name", - OptionalActions: []string{"read", "update"}, - } - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - role smqsdk.Role - logType outputLog - }{ - { - desc: "update group role name successfully", - args: []string{ - group.ID, - rolesCmd, - updateCmd, - role.ID, - "new_name", - domainID, - token, - }, - role: role, - logType: entityLog, - }, - { - desc: "update group role name with invalid token", - args: []string{ - group.ID, - rolesCmd, - updateCmd, - role.ID, - "new_name", - domainID, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("UpdateGroupRole", mock.Anything, tc.args[0], tc.args[3], tc.args[4], tc.args[5], tc.args[6]).Return(tc.role, tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case entityLog: - var role smqsdk.Role - err := json.Unmarshal([]byte(out), &role) - assert.Nil(t, err) - assert.Equal(t, tc.role, role, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.role, role)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - } - - sdkCall.Unset() - }) - } -} - -func TestDeleteGroupRoleCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - groupsCmd := cli.NewGroupsCmd() - rootCmd := setFlags(groupsCmd) - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - logType outputLog - }{ - { - desc: "delete group role successfully", - args: []string{ - group.ID, - rolesCmd, - delCmd, - roleID, - domainID, - token, - }, - logType: okLog, - }, - { - desc: "delete group role with invalid token", - args: []string{ - group.ID, - rolesCmd, - delCmd, - roleID, - domainID, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("DeleteGroupRole", mock.Anything, tc.args[0], tc.args[3], tc.args[4], tc.args[5]).Return(tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case okLog: - assert.True(t, strings.Contains(out, "ok"), fmt.Sprintf("%s unexpected response: expected success message, got: %v", tc.desc, out)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - } - - sdkCall.Unset() - }) - } -} - -func TestAddGroupRoleActionsCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - groupsCmd := cli.NewGroupsCmd() - rootCmd := setFlags(groupsCmd) - - actions := struct { - Actions []string `json:"actions"` - }{ - Actions: []string{"read", "write"}, - } - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - actions []string - logType outputLog - }{ - { - desc: "add actions to role successfully", - args: []string{ - group.ID, - rolesCmd, - actionsCmd, - addCmd, - roleID, - `{"actions":["read","write"]}`, - domainID, - token, - }, - actions: actions.Actions, - logType: entityLog, - }, - { - desc: "add actions to role with invalid JSON", - args: []string{ - group.ID, - rolesCmd, - actionsCmd, - addCmd, - roleID, - `{"actions":["read","write"}`, - domainID, - token, - }, - sdkErr: errors.NewSDKError(errors.New("invalid character '}' after array element")), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.New("invalid character '}' after array element")), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("AddGroupRoleActions", mock.Anything, tc.args[0], tc.args[4], tc.args[6], tc.actions, tc.args[7]).Return(tc.actions, tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case entityLog: - var acts []string - err := json.Unmarshal([]byte(out), &acts) - assert.Nil(t, err) - assert.Equal(t, tc.actions, acts, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.actions, acts)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - } - - sdkCall.Unset() - }) - } -} - -func TestListGroupRoleActionsCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - groupsCmd := cli.NewGroupsCmd() - rootCmd := setFlags(groupsCmd) - - actions := []string{"read", "write"} - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - actions []string - logType outputLog - }{ - { - desc: "list actions of role successfully", - args: []string{ - group.ID, - rolesCmd, - actionsCmd, - listCmd, - roleID, - domainID, - token, - }, - actions: actions, - logType: entityLog, - }, - { - desc: "list actions of role with invalid token", - args: []string{ - group.ID, - rolesCmd, - actionsCmd, - listCmd, - roleID, - domainID, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("GroupRoleActions", mock.Anything, tc.args[0], tc.args[4], tc.args[5], tc.args[6]).Return(tc.actions, tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case entityLog: - var acts []string - err := json.Unmarshal([]byte(out), &acts) - assert.Nil(t, err) - assert.Equal(t, tc.actions, acts, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.actions, acts)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - } - - sdkCall.Unset() - }) - } -} - -func TestDeleteGroupRoleActionsCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - groupsCmd := cli.NewGroupsCmd() - rootCmd := setFlags(groupsCmd) - - actions := struct { - Actions []string `json:"actions"` - }{ - Actions: []string{"read", "write"}, - } - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - logType outputLog - }{ - { - desc: "delete actions from role successfully", - args: []string{ - group.ID, - rolesCmd, - actionsCmd, - delCmd, - roleID, - `{"actions":["read","write"]}`, - domainID, - token, - }, - logType: okLog, - }, - { - desc: "delete all actions from role successfully", - args: []string{ - group.ID, - rolesCmd, - actionsCmd, - delCmd, - roleID, - all, - domainID, - token, - }, - logType: okLog, - }, - { - desc: "delete actions from role with invalid JSON", - args: []string{ - group.ID, - rolesCmd, - actionsCmd, - delCmd, - roleID, - `{"actions":["read","write"}`, - domainID, - token, - }, - sdkErr: errors.NewSDKError(errors.New("invalid character '}' after array element")), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.New("invalid character '}' after array element")), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - var sdkCall *mock.Call - if tc.args[5] == all { - sdkCall = sdkMock.On("RemoveAllGroupRoleActions", mock.Anything, tc.args[0], tc.args[4], tc.args[6], tc.args[7]).Return(tc.sdkErr) - } else { - sdkCall = sdkMock.On("RemoveGroupRoleActions", mock.Anything, tc.args[0], tc.args[4], tc.args[6], actions.Actions, tc.args[7]).Return(tc.sdkErr) - } - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case okLog: - assert.True(t, strings.Contains(out, "ok"), fmt.Sprintf("%s unexpected response: expected success message, got: %v", tc.desc, out)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - } - - sdkCall.Unset() - }) - } -} - -func TestAvailableGroupRoleActionsCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - groupsCmd := cli.NewGroupsCmd() - rootCmd := setFlags(groupsCmd) - - actions := []string{"read", "write", "update"} - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - actions []string - logType outputLog - }{ - { - desc: "list available actions successfully", - args: []string{ - group.ID, - rolesCmd, - actionsCmd, - availableActionsCmd, - domainID, - token, - }, - actions: actions, - logType: entityLog, - }, - { - desc: "list available actions with invalid token", - args: []string{ - group.ID, - rolesCmd, - actionsCmd, - availableActionsCmd, - domainID, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("AvailableGroupRoleActions", mock.Anything, tc.args[4], tc.args[5]).Return(tc.actions, tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case entityLog: - var acts []string - err := json.Unmarshal([]byte(out), &acts) - assert.Nil(t, err) - assert.Equal(t, tc.actions, acts, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.actions, acts)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - } - - sdkCall.Unset() - }) - } -} - -func TestAddGroupRoleMembersCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - groupsCmd := cli.NewGroupsCmd() - rootCmd := setFlags(groupsCmd) - - members := struct { - Members []string `json:"members"` - }{ - Members: []string{"5dc1ce4b-7cc9-4f12-98a6-9d74cc4980bb", "5dc1ce4b-7cc9-4f12-98a6-9d74cc4980bb"}, - } - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - members []string - logType outputLog - }{ - { - desc: "add members to role successfully", - args: []string{ - group.ID, - rolesCmd, - membersCmd, - addCmd, - roleID, - `{"members":["5dc1ce4b-7cc9-4f12-98a6-9d74cc4980bb", "5dc1ce4b-7cc9-4f12-98a6-9d74cc4980bb"]}`, - domainID, - token, - }, - members: members.Members, - logType: entityLog, - }, - { - desc: "add members to role with invalid JSON", - args: []string{ - group.ID, - rolesCmd, - membersCmd, - addCmd, - roleID, - `{"members":["5dc1ce4b-7cc9-4f12-98a6-9d74cc4980bb", "5dc1ce4b-7cc9-4f12-98a6-9d74cc4980bb"}`, - domainID, - token, - }, - sdkErr: errors.NewSDKError(errors.New("invalid character '}' after array element")), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.New("invalid character '}' after array element")), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("AddGroupRoleMembers", mock.Anything, tc.args[0], tc.args[4], tc.args[6], tc.members, tc.args[7]).Return(tc.members, tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case entityLog: - var members []string - err := json.Unmarshal([]byte(out), &members) - assert.Nil(t, err) - assert.Equal(t, tc.members, members, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.members, members)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - } - - sdkCall.Unset() - }) - } -} - -func TestListGroupRoleMembersCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - groupsCmd := cli.NewGroupsCmd() - rootCmd := setFlags(groupsCmd) - - membersPage := smqsdk.RoleMembersPage{ - Total: 1, - Offset: 0, - Limit: 10, - Members: []string{ - "5dc1ce4b-7cc9-4f12-98a6-9d74cc4980bb", - }, - } - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - members smqsdk.RoleMembersPage - logType outputLog - }{ - { - desc: "list members of role successfully", - args: []string{ - group.ID, - rolesCmd, - membersCmd, - listCmd, - roleID, - domainID, - token, - }, - members: membersPage, - logType: entityLog, - }, - { - desc: "list members of role with invalid token", - args: []string{ - group.ID, - rolesCmd, - membersCmd, - listCmd, - roleID, - domainID, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("GroupRoleMembers", mock.Anything, tc.args[0], tc.args[4], tc.args[5], mock.Anything, tc.args[6]).Return(tc.members, tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case entityLog: - var members smqsdk.RoleMembersPage - err := json.Unmarshal([]byte(out), &members) - assert.Nil(t, err) - assert.Equal(t, tc.members, members, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.members, members)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - } - - sdkCall.Unset() - }) - } -} - -func TestDeleteGroupRoleMembersCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - groupsCmd := cli.NewGroupsCmd() - rootCmd := setFlags(groupsCmd) - - members := struct { - Members []string `json:"members"` - }{ - Members: []string{"5dc1ce4b-7cc9-4f12-98a6-9d74cc4980bb", "5dc1ce4b-7cc9-4f12-98a6-9d74cc4980bb"}, - } - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - logType outputLog - }{ - { - desc: "delete members from role successfully", - args: []string{ - group.ID, - rolesCmd, - membersCmd, - delCmd, - roleID, - `{"members":["5dc1ce4b-7cc9-4f12-98a6-9d74cc4980bb", "5dc1ce4b-7cc9-4f12-98a6-9d74cc4980bb"]}`, - domainID, - token, - }, - logType: okLog, - }, - { - desc: "delete all members from role successfully", - args: []string{ - group.ID, - rolesCmd, - membersCmd, - delCmd, - roleID, - all, - domainID, - token, - }, - logType: okLog, - }, - { - desc: "delete members from role with invalid JSON", - args: []string{ - group.ID, - rolesCmd, - membersCmd, - delCmd, - roleID, - `{"members":["5dc1ce4b-7cc9-4f12-98a6-9d74cc4980bb", "5dc1ce4b-7cc9-4f12-98a6-9d74cc4980bb"}`, - domainID, - token, - }, - sdkErr: errors.NewSDKError(errors.New("invalid character '}' after array element")), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.New("invalid character '}' after array element")), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - var sdkCall *mock.Call - if tc.args[5] == all { - sdkCall = sdkMock.On("RemoveAllGroupRoleMembers", mock.Anything, tc.args[0], tc.args[4], tc.args[6], tc.args[7]).Return(tc.sdkErr) - } else { - sdkCall = sdkMock.On("RemoveGroupRoleMembers", mock.Anything, tc.args[0], tc.args[4], tc.args[6], members.Members, tc.args[7]).Return(tc.sdkErr) - } - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case okLog: - assert.True(t, strings.Contains(out, "ok"), fmt.Sprintf("%s unexpected response: expected success message, got: %v", tc.desc, out)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - } - - sdkCall.Unset() - }) - } -} diff --git a/cli/health_test.go b/cli/health_test.go deleted file mode 100644 index f938bfd1d..000000000 --- a/cli/health_test.go +++ /dev/null @@ -1,84 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package cli_test - -import ( - "encoding/json" - "fmt" - "strings" - "testing" - - "github.com/absmach/magistrala/cli" - "github.com/absmach/magistrala/pkg/errors" - mgsdk "github.com/absmach/magistrala/pkg/sdk" - sdkmocks "github.com/absmach/magistrala/pkg/sdk/mocks" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" -) - -func TestHealthCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - healthCmd := cli.NewHealthCmd() - rootCmd := setFlags(healthCmd) - service := "users" - - var health mgsdk.HealthInfo - cases := []struct { - desc string - args []string - logType outputLog - errLogMessage string - health mgsdk.HealthInfo - sdkErr errors.SDKError - }{ - { - desc: "Check health successfully", - args: []string{ - service, - }, - logType: entityLog, - health: mgsdk.HealthInfo{ - Status: "pass", - Description: "users service", - }, - }, - { - desc: "Check health with invalid args", - args: []string{ - service, - extraArg, - }, - logType: usageLog, - }, - { - desc: "Check health with invalid service", - args: []string{ - "invalid", - }, - sdkErr: errors.NewSDKErrorWithStatus(errors.New("unsupported protocol scheme"), 306), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(errors.New("unsupported protocol scheme"), 306)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("Health", mock.Anything).Return(tc.health, tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case entityLog: - err := json.Unmarshal([]byte(out), &health) - assert.Nil(t, err) - assert.Equal(t, tc.health, health, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.health, health)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.True(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - sdkCall.Unset() - }) - } -} diff --git a/cli/invitations_test.go b/cli/invitations_test.go deleted file mode 100644 index ea0e3181d..000000000 --- a/cli/invitations_test.go +++ /dev/null @@ -1,413 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package cli_test - -import ( - "encoding/json" - "fmt" - "net/http" - "strings" - "testing" - - "github.com/absmach/magistrala/cli" - "github.com/absmach/magistrala/internal/testsutil" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - mgsdk "github.com/absmach/magistrala/pkg/sdk" - sdkmocks "github.com/absmach/magistrala/pkg/sdk/mocks" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" -) - -var invitation = mgsdk.Invitation{ - InvitedBy: testsutil.GenerateUUID(&testing.T{}), - InviteeUserID: user.ID, - DomainID: domain.ID, -} - -func TestSendDomainInvitationCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - invCmd := cli.NewInvitationsCmd() - rootCmd := setFlags(invCmd) - - cases := []struct { - desc string - args []string - logType outputLog - errLogMessage string - sdkErr errors.SDKError - }{ - { - desc: "send domain invitation successfully", - args: []string{ - user.ID, - domain.ID, - relation, - validToken, - }, - logType: okLog, - }, - { - desc: "send domain invitation with invalid args", - args: []string{ - user.ID, - domain.ID, - relation, - validToken, - extraArg, - }, - logType: usageLog, - }, - { - desc: "send domain invitation with invalid token", - args: []string{ - user.ID, - domain.ID, - relation, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusUnauthorized), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusUnauthorized)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("SendInvitation", mock.Anything, mock.Anything, mock.Anything).Return(tc.sdkErr) - out := executeCommand(t, rootCmd, append([]string{domainCmd, sendCmd}, tc.args...)...) - switch tc.logType { - case okLog: - assert.True(t, strings.Contains(out, "ok"), fmt.Sprintf("%s unexpected response: expected success message, got: %v", tc.desc, out)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - sdkCall.Unset() - }) - } -} - -func TestGetUserInvitationsCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - invCmd := cli.NewInvitationsCmd() - rootCmd := setFlags(invCmd) - - var page mgsdk.InvitationPage - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - page mgsdk.InvitationPage - logType outputLog - errLogMessage string - }{ - { - desc: "get user invitations successfully", - args: []string{ - token, - }, - page: mgsdk.InvitationPage{ - Total: 1, - Offset: 0, - Limit: 10, - Invitations: []mgsdk.Invitation{invitation}, - }, - logType: entityLog, - }, - { - desc: "get user invitations with invalid args", - args: []string{ - token, - extraArg, - }, - logType: usageLog, - }, - { - desc: "get user invitations with invalid token", - args: []string{ - invalidToken, - }, - logType: errLog, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("Invitations", mock.Anything, mock.Anything, tc.args[0]).Return(tc.page, tc.sdkErr) - - out := executeCommand(t, rootCmd, append([]string{userCmd, getCmd}, tc.args...)...) - - switch tc.logType { - case entityLog: - err := json.Unmarshal([]byte(out), &page) - assert.Nil(t, err) - assert.Equal(t, tc.page, page, fmt.Sprintf("%v unexpected response, expected: %v, got: %v", tc.desc, tc.page, page)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - sdkCall.Unset() - }) - } -} - -func TestGetDomainInvitationsCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - invCmd := cli.NewInvitationsCmd() - rootCmd := setFlags(invCmd) - - var page mgsdk.InvitationPage - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - page mgsdk.InvitationPage - logType outputLog - errLogMessage string - }{ - { - desc: "get domain invitations successfully", - args: []string{ - domain.ID, - token, - }, - page: mgsdk.InvitationPage{ - Total: 1, - Offset: 0, - Limit: 10, - Invitations: []mgsdk.Invitation{invitation}, - }, - logType: entityLog, - }, - { - desc: "get domain invitations with invalid args", - args: []string{ - domain.ID, - token, - extraArg, - }, - logType: usageLog, - }, - { - desc: "get domain invitations with invalid token", - args: []string{ - domain.ID, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("DomainInvitations", mock.Anything, mock.Anything, tc.args[1], tc.args[0]).Return(tc.page, tc.sdkErr) - - out := executeCommand(t, rootCmd, append([]string{domainCmd, getCmd}, tc.args...)...) - - switch tc.logType { - case entityLog: - err := json.Unmarshal([]byte(out), &page) - assert.Nil(t, err) - assert.Equal(t, tc.page, page, fmt.Sprintf("%v unexpected response, expected: %v, got: %v", tc.desc, tc.page, page)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - sdkCall.Unset() - }) - } -} - -func TestAcceptUserInvitationCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - invCmd := cli.NewInvitationsCmd() - rootCmd := setFlags(invCmd) - - cases := []struct { - desc string - args []string - logType outputLog - errLogMessage string - sdkErr errors.SDKError - }{ - { - desc: "accept user invitation successfully", - args: []string{ - domain.ID, - validToken, - }, - logType: okLog, - }, - { - desc: "accept user invitation with invalid args", - args: []string{ - domain.ID, - validToken, - extraArg, - }, - logType: usageLog, - }, - { - desc: "accept user invitation with invalid token", - args: []string{ - domain.ID, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusUnauthorized), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusUnauthorized)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("AcceptInvitation", mock.Anything, mock.Anything, mock.Anything).Return(tc.sdkErr) - out := executeCommand(t, rootCmd, append([]string{userCmd, acceptCmd}, tc.args...)...) - switch tc.logType { - case okLog: - assert.True(t, strings.Contains(out, "ok"), fmt.Sprintf("%s unexpected response: expected success message, got: %v", tc.desc, out)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - sdkCall.Unset() - }) - } -} - -func TestRejectUserInvitationCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - invCmd := cli.NewInvitationsCmd() - rootCmd := setFlags(invCmd) - - cases := []struct { - desc string - args []string - logType outputLog - errLogMessage string - sdkErr errors.SDKError - }{ - { - desc: "reject user invitation successfully", - args: []string{ - domain.ID, - validToken, - }, - logType: okLog, - }, - { - desc: "reject user invitation with invalid args", - args: []string{ - domain.ID, - validToken, - extraArg, - }, - logType: usageLog, - }, - { - desc: "reject user invitation with invalid token", - args: []string{ - domain.ID, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusUnauthorized), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusUnauthorized)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("RejectInvitation", mock.Anything, mock.Anything, mock.Anything).Return(tc.sdkErr) - out := executeCommand(t, rootCmd, append([]string{userCmd, rejectCmd}, tc.args...)...) - switch tc.logType { - case okLog: - assert.True(t, strings.Contains(out, "ok"), fmt.Sprintf("%s unexpected response: expected success message, got: %v", tc.desc, out)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - sdkCall.Unset() - }) - } -} - -func TestDeleteDomainInvitationCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - invCmd := cli.NewInvitationsCmd() - rootCmd := setFlags(invCmd) - - cases := []struct { - desc string - args []string - logType outputLog - errLogMessage string - sdkErr errors.SDKError - }{ - { - desc: "delete domain invitation successfully", - args: []string{ - user.ID, - domain.ID, - validToken, - }, - logType: okLog, - }, - { - desc: "delete domain invitation with invalid args", - args: []string{ - user.ID, - domain.ID, - validToken, - extraArg, - }, - logType: usageLog, - }, - { - desc: "delete domain invitation with invalid token", - args: []string{ - user.ID, - domain.ID, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusUnauthorized), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusUnauthorized)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("DeleteInvitation", mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(tc.sdkErr) - out := executeCommand(t, rootCmd, append([]string{domainCmd, delCmd}, tc.args...)...) - switch tc.logType { - case okLog: - assert.True(t, strings.Contains(out, "ok"), fmt.Sprintf("%s unexpected response: expected success message, got: %v", tc.desc, out)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - sdkCall.Unset() - }) - } -} diff --git a/cli/journal_test.go b/cli/journal_test.go deleted file mode 100644 index f2cb0e426..000000000 --- a/cli/journal_test.go +++ /dev/null @@ -1,123 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package cli_test - -import ( - "encoding/json" - "fmt" - "net/http" - "strings" - "testing" - - "github.com/absmach/magistrala/cli" - "github.com/absmach/magistrala/internal/testsutil" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - mgsdk "github.com/absmach/magistrala/pkg/sdk" - sdkmocks "github.com/absmach/magistrala/pkg/sdk/mocks" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" -) - -var journal = mgsdk.Journal{ - ID: testsutil.GenerateUUID(&testing.T{}), -} - -func TestGetJournalCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - invCmd := cli.NewJournalCmd() - rootCmd := setFlags(invCmd) - - var page mgsdk.JournalsPage - entityType := "group" - entityId := testsutil.GenerateUUID(t) - domainId := testsutil.GenerateUUID(t) - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - page mgsdk.JournalsPage - logType outputLog - errLogMessage string - }{ - { - desc: "get user journal", - args: []string{ - "user", - entityId, - token, - }, - logType: entityLog, - page: mgsdk.JournalsPage{ - Total: 1, - Offset: 0, - Limit: 10, - Journals: []mgsdk.Journal{journal}, - }, - }, - { - desc: "get group journal", - args: []string{ - entityType, - entityId, - domainId, - token, - }, - logType: entityLog, - page: mgsdk.JournalsPage{ - Total: 1, - Offset: 0, - Limit: 10, - Journals: []mgsdk.Journal{journal}, - }, - }, - { - desc: "get journal with invalid args", - args: []string{ - entityType, - entityId, - token, - domainId, - extraArg, - }, - logType: usageLog, - }, - { - desc: "get journal with invalid token", - args: []string{ - entityType, - entityId, - domainId, - invalidToken, - }, - logType: errLog, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("Journal", mock.Anything, tc.args[0], tc.args[1], "", mock.Anything, tc.args[2]).Return(tc.page, tc.sdkErr) - if tc.args[0] != "user" { - sdkCall = sdkMock.On("Journal", mock.Anything, tc.args[0], tc.args[1], tc.args[2], mock.Anything, tc.args[3]).Return(tc.page, tc.sdkErr) - } - out := executeCommand(t, rootCmd, append([]string{getCmd}, tc.args...)...) - - switch tc.logType { - case entityLog: - err := json.Unmarshal([]byte(out), &page) - assert.Nil(t, err) - assert.Equal(t, tc.page, page, fmt.Sprintf("%v unexpected response, expected: %v, got: %v", tc.desc, tc.page, page)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - sdkCall.Unset() - }) - } -} diff --git a/cli/message_test.go b/cli/message_test.go deleted file mode 100644 index 9d8f8e023..000000000 --- a/cli/message_test.go +++ /dev/null @@ -1,86 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package cli_test - -import ( - "fmt" - "net/http" - "strings" - "testing" - - "github.com/absmach/magistrala/cli" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - sdkmocks "github.com/absmach/magistrala/pkg/sdk/mocks" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" -) - -func TestSendMesageCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - messageCmd := cli.NewMessagesCmd() - rootCmd := setFlags(messageCmd) - - message := "[{\"bn\":\"Dev1\",\"n\":\"temp\",\"v\":20}, {\"n\":\"hum\",\"v\":40}, {\"bn\":\"Dev2\", \"n\":\"temp\",\"v\":20}, {\"n\":\"hum\",\"v\":40}]" - - cases := []struct { - desc string - args []string - logType outputLog - errLogMessage string - sdkErr errors.SDKError - }{ - { - desc: "send message successfully", - args: []string{ - domainID, - channel.ID, - message, - client.Credentials.Secret, - }, - logType: okLog, - }, - { - desc: "send message with invalid args", - args: []string{ - domainID, - channel.ID, - message, - client.Credentials.Secret, - extraArg, - }, - logType: usageLog, - }, - { - desc: "send message with invalid client secret", - args: []string{ - domainID, - channel.ID, - message, - "invalid_secret", - }, - sdkErr: errors.NewSDKErrorWithStatus(errors.Wrap(svcerr.ErrAuthentication, errors.Wrap(svcerr.ErrAuthorization, svcerr.ErrNotFound)), http.StatusBadRequest), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(errors.Wrap(svcerr.ErrAuthentication, errors.Wrap(svcerr.ErrAuthorization, svcerr.ErrNotFound)), http.StatusBadRequest)), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("SendMessage", mock.Anything, tc.args[0], tc.args[1], tc.args[2], tc.args[3]).Return(tc.sdkErr) - out := executeCommand(t, rootCmd, append([]string{sendCmd}, tc.args...)...) - - switch tc.logType { - case okLog: - assert.True(t, strings.Contains(out, "ok"), fmt.Sprintf("%s unexpected response: expected success message, got: %v", tc.desc, out)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - sdkCall.Unset() - }) - } -} diff --git a/cli/setup_test.go b/cli/setup_test.go deleted file mode 100644 index d9461a6ab..000000000 --- a/cli/setup_test.go +++ /dev/null @@ -1,112 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package cli_test - -import ( - "bytes" - "testing" - - "github.com/absmach/magistrala/cli" - "github.com/spf13/cobra" - "github.com/stretchr/testify/assert" -) - -type outputLog uint8 - -const ( - usageLog outputLog = iota - errLog - entityLog - okLog - createLog - revokeLog -) - -func executeCommand(t *testing.T, root *cobra.Command, args ...string) string { - buffer := new(bytes.Buffer) - root.SetOut(buffer) - root.SetErr(buffer) - root.SetArgs(args) - err := root.Execute() - assert.NoError(t, err, "Error executing command") - return buffer.String() -} - -func setFlags(rootCmd *cobra.Command) *cobra.Command { - // Root Flags - rootCmd.PersistentFlags().BoolVarP( - &cli.RawOutput, - "raw", - "r", - cli.RawOutput, - "Enables raw output mode for easier parsing of output", - ) - - // Client and Channels Flags - rootCmd.PersistentFlags().Uint64VarP( - &cli.Limit, - "limit", - "l", - 10, - "Limit query parameter", - ) - - rootCmd.PersistentFlags().Uint64VarP( - &cli.Offset, - "offset", - "o", - 0, - "Offset query parameter", - ) - - rootCmd.PersistentFlags().StringVarP( - &cli.Name, - "name", - "n", - "", - "Name query parameter", - ) - - rootCmd.PersistentFlags().StringVarP( - &cli.Identity, - "identity", - "I", - "", - "User identity query parameter", - ) - - rootCmd.PersistentFlags().StringVarP( - &cli.Metadata, - "metadata", - "m", - "", - "Metadata query parameter", - ) - - rootCmd.PersistentFlags().StringVarP( - &cli.Status, - "status", - "S", - "", - "User status query parameter", - ) - - rootCmd.PersistentFlags().StringVarP( - &cli.Topic, - "topic", - "T", - "", - "Subscription topic query parameter", - ) - - rootCmd.PersistentFlags().StringVarP( - &cli.Contact, - "contact", - "C", - "", - "Subscription contact query parameter", - ) - - return rootCmd -} diff --git a/cli/users.go b/cli/users.go deleted file mode 100644 index 9123611d8..000000000 --- a/cli/users.go +++ /dev/null @@ -1,592 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package cli - -import ( - "encoding/json" - "fmt" - "net/url" - "strconv" - - smqsdk "github.com/absmach/magistrala/pkg/sdk" - smqusers "github.com/absmach/magistrala/users" - "github.com/spf13/cobra" -) - -const ( - token = "token" - refreshToken = "refreshtoken" - profile = "profile" - resetPasswordRequest = "resetpasswordrequest" - resetPassword = "resetpassword" - password = "password" - search = "search" - username = "username" - email = "email" - role = "role" - - // Usage strings for user operations. - usageUserCreate = "cli users create [user_auth_token]" - usageUserGet = "cli users get " - usageUserToken = "cli users token " - usageUserRefreshToken = "cli users refreshtoken " - usageUserUpdate = "cli users update " - usageUserUpdateTags = "cli users update tags " - usageUserUpdateUsername = "cli users update username " - usageUserUpdateEmail = "cli users update email " - usageUserUpdateRole = "cli users update role " - usageUserUpdateAll = `cli users update [args...] -Available update options: - cli users update - cli users update tags - cli users update username - cli users update email - cli users update role ` - usageUserProfile = "cli users profile " - usageUserResetPasswordReq = "cli users resetpasswordrequest " - usageUserResetPassword = "cli users resetpassword " - usageUserPassword = "cli users password " - usageUserEnable = "cli users enable " - usageUserDisable = "cli users disable " - usageUserDelete = "cli users delete " - usageUserSearch = "cli users search \nQuery format: username=|firstname=|lastname=|id=[&offset=][&limit=]\nExample: cli users search \"username=john_doe\" " - usageUserSendVerification = "cli users sendverification " - usageUserVerifyEmail = "cli users verifyemail " -) - -func NewUsersCmd() *cobra.Command { - cmd := &cobra.Command{ - Use: "users [operation] [args...]", - Short: "Users management", - Long: `Format: - users [args...] - users [args...] - -Operations (require user_id/all): get, update, enable, disable, delete - -Examples: - users create [user_auth_token] - users token - users refreshtoken - users profile - users resetpasswordrequest - users resetpassword - users password - users search "username=john_doe" - users search "firstname=john&limit=10" - users sendverification - users verifyemail - users all get - users get - users update - users update tags - users update username - users update email - users enable - users disable - users delete `, - - Run: func(cmd *cobra.Command, args []string) { - if len(args) == 0 { - logUsageCmd(*cmd, cmd.Use) - return - } - - switch args[0] { - case create: - handleUserCreate(cmd, args[1:]) - return - case sendVerification: - if len(args) < 2 { - logUsageCmd(*cmd, usageUserSendVerification) - return - } - handleSendVerification(cmd, args[1]) - return - case verifyEmail: - if len(args) < 2 { - logUsageCmd(*cmd, usageUserVerifyEmail) - return - } - handleVerify(cmd, args[1]) - return - case token: - if len(args) < 2 { - logUsageCmd(*cmd, usageUserToken) - return - } - if len(args) < 3 { - logUsageCmd(*cmd, usageUserToken) - return - } - handleUserToken(cmd, args[1], args[2:]) - return - case refreshToken: - if len(args) < 2 { - logUsageCmd(*cmd, usageUserRefreshToken) - return - } - handleUserRefreshToken(cmd, args[1], args[2:]) - return - case profile: - if len(args) < 2 { - logUsageCmd(*cmd, usageUserProfile) - return - } - handleUserProfile(cmd, args[1], args[2:]) - return - case resetPasswordRequest: - if len(args) < 2 { - logUsageCmd(*cmd, usageUserResetPasswordReq) - return - } - handleUserResetPasswordRequest(cmd, args[1], args[2:]) - return - case resetPassword: - if len(args) < 2 { - logUsageCmd(*cmd, usageUserResetPassword) - return - } - if len(args) < 4 { - logUsageCmd(*cmd, usageUserResetPassword) - return - } - handleUserResetPassword(cmd, args[1], args[2:]) - return - case password: - if len(args) < 2 { - logUsageCmd(*cmd, usageUserPassword) - return - } - if len(args) < 4 { - logUsageCmd(*cmd, usageUserPassword) - return - } - handleUserPassword(cmd, args[1], args[2:]) - return - case search: - if len(args) < 2 { - logUsageCmd(*cmd, usageUserSearch) - return - } - if len(args) < 3 { - logUsageCmd(*cmd, usageUserSearch) - return - } - handleUserSearch(cmd, args[1], args[2:]) - return - } - - if len(args) < 2 { - logUsageCmd(*cmd, "users [args...]") - return - } - - userParams := args[0] - operation := args[1] - opArgs := args[2:] - - switch operation { - case get: - handleUserGet(cmd, userParams, opArgs) - case update: - handleUserUpdate(cmd, userParams, opArgs) - case enable: - handleUserEnable(cmd, userParams, opArgs) - case disable: - handleUserDisable(cmd, userParams, opArgs) - case delete: - handleUserDelete(cmd, userParams, opArgs) - default: - logErrorCmd(*cmd, fmt.Errorf("unknown operation: %s", operation)) - } - }, - } - - return cmd -} - -func handleUserCreate(cmd *cobra.Command, args []string) { - if len(args) < 5 || len(args) > 6 { - logUsageCmd(*cmd, usageUserCreate) - return - } - if len(args) == 5 { - args = append(args, "") - } - - user := smqsdk.User{ - FirstName: args[0], - LastName: args[1], - Email: args[2], - Credentials: smqsdk.Credentials{ - Username: args[3], - Secret: args[4], - }, - Status: smqusers.EnabledStatus.String(), - } - user, err := sdk.CreateUser(cmd.Context(), user, args[5]) - if err != nil { - logErrorCmd(*cmd, err) - return - } - - logJSONCmd(*cmd, user) -} - -func handleSendVerification(cmd *cobra.Command, token string) { - if token == "" { - logUsageCmd(*cmd, usageUserToken) - return - } - - if err := sdk.SendVerification(cmd.Context(), token); err != nil { - logErrorCmd(*cmd, err) - return - } - - logJSONCmd(*cmd, "sent verification successfully") -} - -func handleVerify(cmd *cobra.Command, token string) { - if token == "" { - logUsageCmd(*cmd, usageUserToken) - return - } - - if err := sdk.VerifyEmail(cmd.Context(), token); err != nil { - logErrorCmd(*cmd, err) - return - } - - logJSONCmd(*cmd, "verified successfully") -} - -func handleUserGet(cmd *cobra.Command, userParams string, args []string) { - if len(args) != 1 { - logUsageCmd(*cmd, usageUserGet) - return - } - - if userParams == all { - metadata, err := convertMetadata(Metadata) - if err != nil { - logErrorCmd(*cmd, err) - return - } - - pageMetadata := smqsdk.PageMetadata{ - Username: Username, - Identity: Identity, - Offset: Offset, - Limit: Limit, - Metadata: metadata, - Status: Status, - } - - l, err := sdk.Users(cmd.Context(), pageMetadata, args[0]) - if err != nil { - logErrorCmd(*cmd, err) - return - } - logJSONCmd(*cmd, l) - return - } - - u, err := sdk.User(cmd.Context(), userParams, args[0]) - if err != nil { - logErrorCmd(*cmd, err) - return - } - - logJSONCmd(*cmd, u) -} - -func handleUserUpdate(cmd *cobra.Command, userID string, args []string) { - if len(args) < 1 { - logUsageCmd(*cmd, usageUserUpdateAll) - return - } - - if len(args) < 2 || len(args) > 3 { - if len(args) >= 1 { - switch args[0] { - case tags: - logUsageCmd(*cmd, usageUserUpdateTags) - return - case username: - logUsageCmd(*cmd, usageUserUpdateUsername) - return - case email: - logUsageCmd(*cmd, usageUserUpdateEmail) - return - case role: - logUsageCmd(*cmd, usageUserUpdateRole) - return - } - } - logUsageCmd(*cmd, usageUserUpdateAll) - return - } - - var user smqsdk.User - if args[0] == "tags" { - if len(args) != 3 { - logUsageCmd(*cmd, usageUserUpdateTags) - return - } - if err := json.Unmarshal([]byte(args[1]), &user.Tags); err != nil { - logErrorCmd(*cmd, err) - return - } - user.ID = userID - user, err := sdk.UpdateUserTags(cmd.Context(), user, args[2]) - if err != nil { - logErrorCmd(*cmd, err) - return - } - logJSONCmd(*cmd, user) - return - } - - if args[0] == "email" { - if len(args) != 3 { - logUsageCmd(*cmd, usageUserUpdateEmail) - return - } - user.ID = userID - user.Email = args[1] - user, err := sdk.UpdateUserEmail(cmd.Context(), user, args[2]) - if err != nil { - logErrorCmd(*cmd, err) - return - } - logJSONCmd(*cmd, user) - return - } - - if args[0] == "username" { - if len(args) != 3 { - logUsageCmd(*cmd, usageUserUpdateUsername) - return - } - user.ID = userID - user.Credentials.Username = args[1] - user, err := sdk.UpdateUsername(cmd.Context(), user, args[2]) - if err != nil { - logErrorCmd(*cmd, err) - return - } - logJSONCmd(*cmd, user) - return - } - - if args[0] == "role" { - if len(args) != 3 { - logUsageCmd(*cmd, usageUserUpdateRole) - return - } - user.ID = userID - user.Role = args[1] - user, err := sdk.UpdateUserRole(cmd.Context(), user, args[2]) - if err != nil { - logErrorCmd(*cmd, err) - return - } - logJSONCmd(*cmd, user) - return - } - - if len(args) != 2 { - logUsageCmd(*cmd, usageUserUpdate) - return - } - - if err := json.Unmarshal([]byte(args[0]), &user); err != nil { - logErrorCmd(*cmd, err) - return - } - user.ID = userID - user, err := sdk.UpdateUser(cmd.Context(), user, args[1]) - if err != nil { - logErrorCmd(*cmd, err) - return - } - - logJSONCmd(*cmd, user) -} - -func handleUserEnable(cmd *cobra.Command, userID string, args []string) { - if len(args) != 1 { - logUsageCmd(*cmd, usageUserEnable) - return - } - - user, err := sdk.EnableUser(cmd.Context(), userID, args[0]) - if err != nil { - logErrorCmd(*cmd, err) - return - } - - logJSONCmd(*cmd, user) -} - -func handleUserDisable(cmd *cobra.Command, userID string, args []string) { - if len(args) != 1 { - logUsageCmd(*cmd, usageUserDisable) - return - } - - user, err := sdk.DisableUser(cmd.Context(), userID, args[0]) - if err != nil { - logErrorCmd(*cmd, err) - return - } - - logJSONCmd(*cmd, user) -} - -func handleUserDelete(cmd *cobra.Command, userID string, args []string) { - if len(args) != 1 { - logUsageCmd(*cmd, usageUserDelete) - return - } - - if err := sdk.DeleteUser(cmd.Context(), userID, args[0]); err != nil { - logErrorCmd(*cmd, err) - return - } - logOKCmd(*cmd) -} - -func handleUserToken(cmd *cobra.Command, username string, args []string) { - if len(args) != 1 { - logUsageCmd(*cmd, usageUserToken) - return - } - - loginReq := smqsdk.Login{ - Username: username, - Password: args[0], - } - - token, err := sdk.CreateToken(cmd.Context(), loginReq) - if err != nil { - logErrorCmd(*cmd, err) - return - } - - logJSONCmd(*cmd, token) -} - -func handleUserRefreshToken(cmd *cobra.Command, refreshToken string, args []string) { - if len(args) != 0 { - logUsageCmd(*cmd, usageUserRefreshToken) - return - } - - token, err := sdk.RefreshToken(cmd.Context(), refreshToken) - if err != nil { - logErrorCmd(*cmd, err) - return - } - - logJSONCmd(*cmd, token) -} - -func handleUserProfile(cmd *cobra.Command, token string, args []string) { - if len(args) != 0 { - logUsageCmd(*cmd, usageUserProfile) - return - } - - user, err := sdk.UserProfile(cmd.Context(), token) - if err != nil { - logErrorCmd(*cmd, err) - return - } - - logJSONCmd(*cmd, user) -} - -func handleUserResetPasswordRequest(cmd *cobra.Command, email string, args []string) { - if len(args) != 0 { - logUsageCmd(*cmd, usageUserResetPasswordReq) - return - } - - if err := sdk.ResetPasswordRequest(cmd.Context(), email); err != nil { - logErrorCmd(*cmd, err) - return - } - - logOKCmd(*cmd) -} - -func handleUserResetPassword(cmd *cobra.Command, password string, args []string) { - if len(args) != 2 { - logUsageCmd(*cmd, usageUserResetPassword) - return - } - - if err := sdk.ResetPassword(cmd.Context(), password, args[0], args[1]); err != nil { - logErrorCmd(*cmd, err) - return - } - - logOKCmd(*cmd) -} - -func handleUserPassword(cmd *cobra.Command, oldPassword string, args []string) { - if len(args) != 2 { - logUsageCmd(*cmd, usageUserPassword) - return - } - - user, err := sdk.UpdatePassword(cmd.Context(), oldPassword, args[0], args[1]) - if err != nil { - logErrorCmd(*cmd, err) - return - } - - logJSONCmd(*cmd, user) -} - -func handleUserSearch(cmd *cobra.Command, query string, args []string) { - if len(args) != 1 { - logUsageCmd(*cmd, usageUserSearch) - return - } - - values, err := url.ParseQuery(query) - if err != nil { - logErrorCmd(*cmd, fmt.Errorf("failed to parse query: %s", err)) - return - } - - pm := smqsdk.PageMetadata{ - Offset: Offset, - Limit: Limit, - ID: values.Get("id"), - Username: values.Get("username"), - FirstName: values.Get("firstname"), - LastName: values.Get("lastname"), - } - - if off, err := strconv.Atoi(values.Get("offset")); err == nil { - pm.Offset = uint64(off) - } - - if lim, err := strconv.Atoi(values.Get("limit")); err == nil { - pm.Limit = uint64(lim) - } - - users, err := sdk.SearchUsers(cmd.Context(), pm, args[0]) - if err != nil { - logErrorCmd(*cmd, err) - return - } - - logJSONCmd(*cmd, users) -} diff --git a/cli/users_test.go b/cli/users_test.go deleted file mode 100644 index 0cc59500b..000000000 --- a/cli/users_test.go +++ /dev/null @@ -1,1403 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package cli_test - -import ( - "encoding/json" - "fmt" - "net/http" - "strings" - "testing" - - "github.com/absmach/magistrala/cli" - "github.com/absmach/magistrala/internal/testsutil" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - mgsdk "github.com/absmach/magistrala/pkg/sdk" - sdkmocks "github.com/absmach/magistrala/pkg/sdk/mocks" - "github.com/absmach/magistrala/users" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" -) - -var user = mgsdk.User{ - ID: testsutil.GenerateUUID(&testing.T{}), - FirstName: "testuserfirstname", - LastName: "testuserfirstname", - Email: "testuser@example.com", - Credentials: mgsdk.Credentials{ - Secret: "testpassword", - Username: "testusername", - }, - Status: users.EnabledStatus.String(), -} - -func TestCreateUsersCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - usersCmd := cli.NewUsersCmd() - rootCmd := setFlags(usersCmd) - - var usr mgsdk.User - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - user mgsdk.User - logType outputLog - }{ - { - desc: "create user successfully with token", - args: []string{ - createCmd, - user.FirstName, - user.LastName, - user.Email, - user.Credentials.Username, - user.Credentials.Secret, - validToken, - }, - user: user, - logType: entityLog, - }, - { - desc: "create user successfully without token", - args: []string{ - createCmd, - user.FirstName, - user.LastName, - user.Email, - user.Credentials.Username, - user.Credentials.Secret, - }, - user: user, - logType: entityLog, - }, - { - desc: "failed to create user", - args: []string{ - createCmd, - user.FirstName, - user.LastName, - user.Email, - user.Credentials.Username, - user.Credentials.Secret, - validToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrCreateEntity, http.StatusUnprocessableEntity), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrCreateEntity, http.StatusUnprocessableEntity).Error()), - logType: errLog, - }, - { - desc: "create user with invalid args", - args: []string{ - createCmd, - user.FirstName, - user.Credentials.Username, - }, - errLogMessage: rootCmd.Use, - logType: usageLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("CreateUser", mock.Anything, mock.Anything, mock.Anything).Return(tc.user, tc.sdkErr) - if len(tc.args) == 6 { - sdkUser := mgsdk.User{ - FirstName: tc.args[1], - LastName: tc.args[2], - Email: tc.args[3], - Credentials: mgsdk.Credentials{ - Username: tc.args[4], - Secret: tc.args[5], - }, - Status: users.EnabledStatus.String(), - } - sdkCall = sdkMock.On("CreateUser", mock.Anything, sdkUser, "").Return(tc.user, tc.sdkErr) - } else if len(tc.args) == 7 { - sdkUser := mgsdk.User{ - FirstName: tc.args[1], - LastName: tc.args[2], - Email: tc.args[3], - Credentials: mgsdk.Credentials{ - Username: tc.args[4], - Secret: tc.args[5], - }, - Status: users.EnabledStatus.String(), - } - sdkCall = sdkMock.On("CreateUser", mock.Anything, sdkUser, tc.args[6]).Return(tc.user, tc.sdkErr) - } - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case entityLog: - err := json.Unmarshal([]byte(out), &usr) - assert.Nil(t, err) - assert.Equal(t, tc.user, usr, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.user, usr)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - - sdkCall.Unset() - }) - } -} - -func TestGetUsersCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - usersCmd := cli.NewUsersCmd() - rootCmd := setFlags(usersCmd) - - var page mgsdk.UsersPage - var usr mgsdk.User - out := "" - userID := testsutil.GenerateUUID(t) - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - user mgsdk.User - page mgsdk.UsersPage - logType outputLog - }{ - { - desc: "get users successfully", - args: []string{ - all, - getCmd, - validToken, - }, - sdkErr: nil, - page: mgsdk.UsersPage{ - Users: []mgsdk.User{user}, - }, - logType: entityLog, - }, - { - desc: "get user successfully with id", - args: []string{ - userID, - getCmd, - validToken, - }, - sdkErr: nil, - user: user, - logType: entityLog, - }, - { - desc: "get user with invalid id", - args: []string{ - invalidID, - getCmd, - validToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrViewEntity, http.StatusBadRequest), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrViewEntity, http.StatusBadRequest).Error()), - user: mgsdk.User{}, - logType: errLog, - }, - { - desc: "get users successfully with offset and limit", - args: []string{ - all, - getCmd, - validToken, - "--offset=2", - "--limit=5", - }, - sdkErr: nil, - page: mgsdk.UsersPage{ - Users: []mgsdk.User{user}, - }, - logType: entityLog, - }, - { - desc: "get users with invalid token", - args: []string{ - all, - getCmd, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden).Error()), - page: mgsdk.UsersPage{}, - logType: errLog, - }, - { - desc: "get users with invalid args", - args: []string{ - all, - getCmd, - validToken, - extraArg, - }, - errLogMessage: "cli users get ", - logType: usageLog, - }, - { - desc: "get user with failed get operation", - args: []string{ - userID, - getCmd, - validToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrViewEntity, http.StatusInternalServerError), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrViewEntity, http.StatusInternalServerError).Error()), - user: mgsdk.User{}, - logType: errLog, - }, - { - desc: "get user without operation", - args: []string{ - userID, - }, - errLogMessage: "users [args...]", - logType: usageLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("Users", mock.Anything, mock.Anything, mock.Anything).Return(tc.page, tc.sdkErr) - var sdkCall1 *mock.Call - if len(tc.args) >= 3 { - sdkCall1 = sdkMock.On("User", mock.Anything, tc.args[0], tc.args[2]).Return(tc.user, tc.sdkErr) - } - - out = executeCommand(t, rootCmd, tc.args...) - - if tc.logType == entityLog { - switch { - case tc.args[0] == all: - err := json.Unmarshal([]byte(out), &page) - if err != nil { - t.Fatalf("Failed to unmarshal JSON: %v", err) - } - default: - err := json.Unmarshal([]byte(out), &usr) - if err != nil { - t.Fatalf("Failed to unmarshal JSON: %v", err) - } - } - } - - switch tc.logType { - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.True(t, strings.Contains(out, tc.errLogMessage), fmt.Sprintf("%s invalid usage: expected to contain %s, got: %s", tc.desc, tc.errLogMessage, out)) - } - - if tc.logType == entityLog { - if tc.args[0] != all { - assert.Equal(t, tc.user, usr, fmt.Sprintf("%v unexpected response, expected: %v, got: %v", tc.desc, tc.user, usr)) - } else { - assert.Equal(t, tc.page, page, fmt.Sprintf("%v unexpected response, expected: %v, got: %v", tc.desc, tc.page, page)) - } - } - - sdkCall.Unset() - if sdkCall1 != nil { - sdkCall1.Unset() - } - }) - } -} - -func TestIssueTokenCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - usersCmd := cli.NewUsersCmd() - rootCmd := setFlags(usersCmd) - - var tkn mgsdk.Token - invalidPassword := "wrong_password" - - token := mgsdk.Token{ - AccessToken: testsutil.GenerateUUID(t), - RefreshToken: testsutil.GenerateUUID(t), - } - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - token mgsdk.Token - logType outputLog - }{ - { - desc: "issue token successfully", - args: []string{ - tokCmd, - user.Email, - user.Credentials.Secret, - }, - sdkErr: nil, - logType: entityLog, - token: token, - }, - { - desc: "issue token with failed authentication", - args: []string{ - tokCmd, - user.Email, - invalidPassword, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden).Error()), - logType: errLog, - token: mgsdk.Token{}, - }, - { - desc: "issue token with invalid args", - args: []string{ - tokCmd, - user.Email, - user.Credentials.Secret, - extraArg, - }, - errLogMessage: "cli users token ", - logType: usageLog, - }, - { - desc: "issue token with missing password", - args: []string{ - tokCmd, - user.Email, - }, - errLogMessage: "cli users token ", - logType: usageLog, - }, - { - desc: "issue token with missing username", - args: []string{ - tokCmd, - }, - errLogMessage: "cli users token ", - logType: usageLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - var sdkCall *mock.Call - if len(tc.args) >= 3 { - lg := mgsdk.Login{ - Username: tc.args[1], - Password: tc.args[2], - } - sdkCall = sdkMock.On("CreateToken", mock.Anything, lg).Return(tc.token, tc.sdkErr) - } - - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case entityLog: - err := json.Unmarshal([]byte(out), &tkn) - assert.Nil(t, err) - assert.Equal(t, tc.token, tkn, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.token, tkn)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.True(t, strings.Contains(out, tc.errLogMessage), fmt.Sprintf("%s invalid usage: expected to contain %s, got: %s", tc.desc, tc.errLogMessage, out)) - } - - if sdkCall != nil { - sdkCall.Unset() - } - }) - } -} - -func TestRefreshIssueTokenCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - usersCmd := cli.NewUsersCmd() - rootCmd := setFlags(usersCmd) - - var tkn mgsdk.Token - - token := mgsdk.Token{ - AccessToken: testsutil.GenerateUUID(t), - RefreshToken: testsutil.GenerateUUID(t), - } - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - token mgsdk.Token - logType outputLog - }{ - { - desc: "issue refresh token successfully without domain id", - args: []string{ - refTokCmd, - "token", - }, - sdkErr: nil, - logType: entityLog, - token: token, - }, - { - desc: "issue refresh token with invalid args", - args: []string{ - refTokCmd, - "token", - extraArg, - }, - errLogMessage: rootCmd.Use, - logType: usageLog, - }, - { - desc: "issue refresh token with invalid Username", - args: []string{ - refTokCmd, - "invalidToken", - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden).Error()), - logType: errLog, - token: mgsdk.Token{}, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("RefreshToken", mock.Anything, mock.Anything).Return(tc.token, tc.sdkErr) - - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case entityLog: - err := json.Unmarshal([]byte(out), &tkn) - assert.Nil(t, err) - assert.Equal(t, tc.token, tkn, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.token, tkn)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - - sdkCall.Unset() - }) - } -} - -func TestUpdateUserCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - usersCmd := cli.NewUsersCmd() - rootCmd := setFlags(usersCmd) - - var usr mgsdk.User - - userID := testsutil.GenerateUUID(t) - - tagUpdateType := "tags" - emailUpdateType := "email" - roleUpdateType := "role" - newEmail := "newemail@example.com" - newRole := "administrator" - newTagsJSON := "[\"tag1\", \"tag2\"]" - newNameMetadataJSON := "{\"name\":\"new name\", \"metadata\":{\"key\": \"value\"}}" - newMetadataJSON := "{\"metadata\":{\"key\": \"value\"}}" - newPrivateMetadataJSON := "{\"private_metadata\":{\"key\": \"value\"}}" - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - user mgsdk.User - logType outputLog - }{ - { - desc: "update user tags successfully", - args: []string{ - userID, - updateCmd, - tagUpdateType, - newTagsJSON, - validToken, - }, - sdkErr: nil, - logType: entityLog, - user: user, - }, - { - desc: "update user tags with invalid json", - args: []string{ - userID, - updateCmd, - tagUpdateType, - "[\"tag1\", \"tag2\"", - validToken, - }, - sdkErr: errors.NewSDKError(errEndJSONInput), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errEndJSONInput), - logType: errLog, - }, - { - desc: "update user tags with invalid token", - args: []string{ - userID, - updateCmd, - tagUpdateType, - newTagsJSON, - invalidToken, - }, - logType: errLog, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - }, - { - desc: "update user public metadata successfully", - args: []string{ - userID, - updateCmd, - newPrivateMetadataJSON, - validToken, - }, - logType: entityLog, - user: user, - }, - { - desc: "update user public metadata with invalid json", - args: []string{ - userID, - updateCmd, - "{\"private_metadata\":{\"key\": \"value\"", - validToken, - }, - sdkErr: errors.NewSDKError(errEndJSONInput), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errEndJSONInput), - logType: errLog, - }, - { - desc: "update user metadata successfully", - args: []string{ - userID, - updateCmd, - newMetadataJSON, - validToken, - }, - logType: entityLog, - user: user, - }, - { - desc: "update user metadata with invalid json", - args: []string{ - userID, - updateCmd, - "{\"metadata\":{\"key\": \"value\"", - validToken, - }, - sdkErr: errors.NewSDKError(errEndJSONInput), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errEndJSONInput), - logType: errLog, - }, - { - desc: "update user email successfully", - args: []string{ - userID, - updateCmd, - emailUpdateType, - newEmail, - validToken, - }, - logType: entityLog, - user: user, - }, - { - desc: "update user email with invalid token", - args: []string{ - userID, - updateCmd, - emailUpdateType, - newEmail, - invalidToken, - }, - logType: errLog, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - }, - { - desc: "update user successfully", - args: []string{ - userID, - updateCmd, - newNameMetadataJSON, - validToken, - }, - logType: entityLog, - user: user, - }, - { - desc: "update user with invalid token", - args: []string{ - userID, - updateCmd, - newNameMetadataJSON, - invalidToken, - }, - logType: errLog, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - }, - { - desc: "update user with invalid json", - args: []string{ - userID, - updateCmd, - "{\"name\":\"new name\", \"metadata\":{\"key\": \"value\"}", - validToken, - }, - sdkErr: errors.NewSDKError(errEndJSONInput), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errEndJSONInput), - logType: errLog, - }, - { - desc: "update user role successfully", - args: []string{ - userID, - updateCmd, - roleUpdateType, - newRole, - validToken, - }, - logType: entityLog, - user: user, - }, - { - desc: "update user role with invalid token", - args: []string{ - userID, - updateCmd, - roleUpdateType, - newRole, - invalidToken, - }, - logType: errLog, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - }, - { - desc: "update user with invalid args", - args: []string{ - userID, - updateCmd, - roleUpdateType, - newRole, - validToken, - extraArg, - }, - errLogMessage: "cli users update role ", - logType: usageLog, - }, - { - desc: "update user without specifying what to update", - args: []string{ - userID, - updateCmd, - }, - errLogMessage: `cli users update [args...] -Available update options: - cli users update - cli users update tags - cli users update username - cli users update email - cli users update role `, - logType: usageLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("UpdateUser", mock.Anything, mock.Anything, mock.Anything).Return(tc.user, tc.sdkErr) - sdkCall1 := sdkMock.On("UpdateUserTags", mock.Anything, mock.Anything, mock.Anything).Return(tc.user, tc.sdkErr) - sdkCall2 := sdkMock.On("UpdateUserIdentity", mock.Anything, mock.Anything, mock.Anything).Return(tc.user, tc.sdkErr) - sdkCall3 := sdkMock.On("UpdateUserRole", mock.Anything, mock.Anything, mock.Anything).Return(tc.user, tc.sdkErr) - switch { - case len(tc.args) > 2 && tc.args[2] == tagUpdateType: - var u mgsdk.User - u.Tags = []string{"tag1", "tag2"} - u.ID = tc.args[0] - - sdkCall1 = sdkMock.On("UpdateUserTags", mock.Anything, u, tc.args[4]).Return(tc.user, tc.sdkErr) - case len(tc.args) > 2 && tc.args[2] == emailUpdateType: - var u mgsdk.User - u.Email = tc.args[3] - u.ID = tc.args[0] - - sdkCall2 = sdkMock.On("UpdateUserEmail", mock.Anything, u, tc.args[4]).Return(tc.user, tc.sdkErr) - case len(tc.args) > 2 && tc.args[2] == roleUpdateType && len(tc.args) >= 5: - sdkCall3 = sdkMock.On("UpdateUserRole", mock.Anything, mgsdk.User{ - Role: tc.args[3], - }, tc.args[4]).Return(tc.user, tc.sdkErr) - case len(tc.args) == 4: // Basic user update - sdkCall = sdkMock.On("UpdateUser", mock.Anything, mgsdk.User{ - FirstName: "new name", - PrivateMetadata: mgsdk.Metadata{ - "key": "value", - }, - }, tc.args[3]).Return(tc.user, tc.sdkErr) - } - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case entityLog: - err := json.Unmarshal([]byte(out), &usr) - assert.Nil(t, err) - assert.Equal(t, tc.user, usr, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.user, usr)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.True(t, strings.Contains(out, tc.errLogMessage), fmt.Sprintf("%s invalid usage: expected to contain %s, got: %s", tc.desc, tc.errLogMessage, out)) - } - - sdkCall.Unset() - sdkCall1.Unset() - sdkCall2.Unset() - sdkCall3.Unset() - }) - } -} - -func TestGetUserProfileCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - usersCmd := cli.NewUsersCmd() - rootCmd := setFlags(usersCmd) - - var usr mgsdk.User - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - user mgsdk.User - logType outputLog - }{ - { - desc: "get user profile successfully", - args: []string{ - profCmd, - validToken, - }, - sdkErr: nil, - logType: entityLog, - }, - { - desc: "get user profile with invalid args", - args: []string{ - profCmd, - validToken, - extraArg, - }, - errLogMessage: "cli users profile ", - logType: usageLog, - }, - { - desc: "get user profile with invalid token", - args: []string{ - profCmd, - "invalid_token_string", - }, - logType: errLog, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - }, - { - desc: "get user profile with missing token", - args: []string{ - profCmd, - }, - errLogMessage: "cli users profile ", - logType: usageLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - var sdkCall *mock.Call - if len(tc.args) >= 2 { - sdkCall = sdkMock.On("UserProfile", mock.Anything, tc.args[1]).Return(tc.user, tc.sdkErr) - } - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.True(t, strings.Contains(out, tc.errLogMessage), fmt.Sprintf("%s invalid usage: expected to contain %s, got: %s", tc.desc, tc.errLogMessage, out)) - case entityLog: - err := json.Unmarshal([]byte(out), &usr) - assert.Nil(t, err) - assert.Equal(t, tc.user, usr, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.user, usr)) - } - if sdkCall != nil { - sdkCall.Unset() - } - }) - } -} - -func TestResetPasswordRequestCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - usersCmd := cli.NewUsersCmd() - rootCmd := setFlags(usersCmd) - exampleEmail := "example@mail.com" - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - logType outputLog - }{ - { - desc: "request password reset successfully", - args: []string{ - resPassReqCmd, - exampleEmail, - }, - sdkErr: nil, - logType: okLog, - }, - { - desc: "request password reset with invalid args", - args: []string{ - resPassReqCmd, - exampleEmail, - extraArg, - }, - errLogMessage: rootCmd.Use, - logType: usageLog, - }, - { - desc: "failed request password reset", - args: []string{ - resPassReqCmd, - exampleEmail, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrUpdateEntity, http.StatusUnprocessableEntity), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrUpdateEntity, http.StatusUnprocessableEntity).Error()), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("ResetPasswordRequest", mock.Anything, tc.args[1]).Return(tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - case okLog: - assert.True(t, strings.Contains(out, "ok"), fmt.Sprintf("%s unexpected response: expected success message, got: %v", tc.desc, out)) - } - sdkCall.Unset() - }) - } -} - -func TestResetPasswordCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - usersCmd := cli.NewUsersCmd() - rootCmd := setFlags(usersCmd) - newPassword := "new-password" - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - logType outputLog - }{ - { - desc: "reset password successfully", - args: []string{ - resPassCmd, - newPassword, - newPassword, - validToken, - }, - sdkErr: nil, - logType: okLog, - }, - { - desc: "reset password with invalid args", - args: []string{ - resPassCmd, - newPassword, - newPassword, - validToken, - extraArg, - }, - errLogMessage: rootCmd.Use, - logType: usageLog, - }, - { - desc: "reset password with invalid token", - args: []string{ - resPassCmd, - newPassword, - newPassword, - invalidToken, - }, - logType: errLog, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("ResetPassword", mock.Anything, tc.args[1], tc.args[2], tc.args[3]).Return(tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - - sdkCall.Unset() - }) - } -} - -func TestUpdatePasswordCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - usersCmd := cli.NewUsersCmd() - rootCmd := setFlags(usersCmd) - oldPassword := "old-password" - newPassword := "new-password" - - var usr mgsdk.User - var err error - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - user mgsdk.User - logType outputLog - }{ - { - desc: "update password successfully", - args: []string{ - passCmd, - oldPassword, - newPassword, - validToken, - }, - sdkErr: nil, - logType: entityLog, - user: user, - }, - { - desc: "reset password with invalid args", - args: []string{ - passCmd, - oldPassword, - newPassword, - validToken, - extraArg, - }, - errLogMessage: rootCmd.Use, - sdkErr: nil, - logType: usageLog, - user: user, - }, - { - desc: "update password with invalid token", - args: []string{ - passCmd, - oldPassword, - newPassword, - invalidToken, - }, - logType: errLog, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("UpdatePassword", mock.Anything, tc.args[1], tc.args[2], tc.args[3]).Return(tc.user, tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - case entityLog: - err = json.Unmarshal([]byte(out), &usr) - assert.Nil(t, err) - assert.Equal(t, tc.user, usr, fmt.Sprintf("%s user mismatch: expected %+v got %+v", tc.desc, tc.user, usr)) - } - - sdkCall.Unset() - }) - } -} - -func TestEnableUserCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - usersCmd := cli.NewUsersCmd() - rootCmd := setFlags(usersCmd) - var usr mgsdk.User - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - user mgsdk.User - logType outputLog - }{ - { - desc: "enable user successfully", - args: []string{ - user.ID, - enableCmd, - validToken, - }, - sdkErr: nil, - user: user, - logType: entityLog, - }, - { - desc: "enable user with invalid args", - args: []string{ - user.ID, - enableCmd, - validToken, - extraArg, - }, - errLogMessage: rootCmd.Use, - logType: usageLog, - }, - { - desc: "enable user with invalid token", - args: []string{ - user.ID, - enableCmd, - invalidToken, - }, - logType: errLog, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("EnableUser", mock.Anything, tc.args[0], tc.args[2]).Return(tc.user, tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - case entityLog: - err := json.Unmarshal([]byte(out), &usr) - assert.Nil(t, err) - assert.Equal(t, tc.user, usr, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.user, usr)) - } - - sdkCall.Unset() - }) - } -} - -func TestDisableUserCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - usersCmd := cli.NewUsersCmd() - rootCmd := setFlags(usersCmd) - - var usr mgsdk.User - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - user mgsdk.User - logType outputLog - }{ - { - desc: "disable user successfully", - args: []string{ - user.ID, - disableCmd, - validToken, - }, - sdkErr: nil, - logType: entityLog, - user: user, - }, - { - desc: "disable user with invalid args", - args: []string{ - user.ID, - disableCmd, - validToken, - extraArg, - }, - errLogMessage: rootCmd.Use, - logType: usageLog, - }, - { - desc: "disable user with invalid token", - args: []string{ - user.ID, - disableCmd, - invalidToken, - }, - logType: errLog, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden)), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("DisableUser", mock.Anything, tc.args[0], tc.args[2]).Return(tc.user, tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - case entityLog: - err := json.Unmarshal([]byte(out), &usr) - if err != nil { - t.Fatalf("json.Unmarshal failed: %v", err) - } - assert.Equal(t, tc.user, usr, fmt.Sprintf("%s unexpected response: expected: %v, got: %v", tc.desc, tc.user, usr)) - } - - sdkCall.Unset() - }) - } -} - -func TestDeleteUserCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - usersCmd := cli.NewUsersCmd() - rootCmd := setFlags(usersCmd) - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - logType outputLog - }{ - { - desc: "delete user successfully", - args: []string{ - user.ID, - delCmd, - validToken, - }, - logType: okLog, - }, - { - desc: "delete user with invalid args", - args: []string{ - user.ID, - delCmd, - validToken, - extraArg, - }, - errLogMessage: rootCmd.Use, - logType: usageLog, - }, - { - desc: "delete user with invalid token", - args: []string{ - user.ID, - delCmd, - invalidToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden).Error()), - logType: errLog, - }, - { - desc: "delete user with invalid user ID", - args: []string{ - invalidID, - delCmd, - validToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden).Error()), - logType: errLog, - }, - { - desc: "delete user with failed to delete", - args: []string{ - user.ID, - delCmd, - validToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrUpdateEntity, http.StatusUnprocessableEntity), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrUpdateEntity, http.StatusUnprocessableEntity).Error()), - logType: errLog, - }, - { - desc: "delete user with invalid args", - args: []string{ - user.ID, - delCmd, - extraArg, - }, - errLogMessage: rootCmd.Use, - logType: usageLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("DeleteUser", mock.Anything, mock.Anything, mock.Anything).Return(tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case okLog: - assert.True(t, strings.Contains(out, "ok"), fmt.Sprintf("%s unexpected response: expected success message, got: %v", tc.desc, out)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - - sdkCall.Unset() - }) - } -} - -func TestSearchUsersCmd(t *testing.T) { - sdkMock := new(sdkmocks.SDK) - cli.SetSDK(sdkMock) - usersCmd := cli.NewUsersCmd() - rootCmd := setFlags(usersCmd) - - usersPage := mgsdk.UsersPage{ - Users: []mgsdk.User{user}, - PageRes: mgsdk.PageRes{ - Total: 1, - Offset: 0, - Limit: 10, - }, - } - - cases := []struct { - desc string - args []string - sdkErr errors.SDKError - errLogMessage string - usersPage mgsdk.UsersPage - logType outputLog - }{ - { - desc: "search users by username successfully", - args: []string{ - "search", - "username=testuser", - validToken, - }, - usersPage: usersPage, - logType: entityLog, - }, - { - desc: "search users with missing token", - args: []string{ - "search", - "username=testuser", - }, - logType: usageLog, - }, - { - desc: "search users with missing query", - args: []string{ - "search", - validToken, - }, - logType: usageLog, - }, - { - desc: "search users with extra arguments", - args: []string{ - "search", - "username=testuser", - validToken, - extraArg, - }, - logType: usageLog, - }, - { - desc: "search users with service error", - args: []string{ - "search", - "username=testuser", - validToken, - }, - sdkErr: errors.NewSDKErrorWithStatus(svcerr.ErrViewEntity, http.StatusBadRequest), - errLogMessage: fmt.Sprintf("\nerror: %s\n\n", errors.NewSDKErrorWithStatus(svcerr.ErrViewEntity, http.StatusBadRequest).Error()), - logType: errLog, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - sdkCall := sdkMock.On("SearchUsers", mock.Anything, mock.Anything, mock.Anything).Return(tc.usersPage, tc.sdkErr) - out := executeCommand(t, rootCmd, tc.args...) - - switch tc.logType { - case entityLog: - var page mgsdk.UsersPage - err := json.Unmarshal([]byte(out), &page) - assert.Nil(t, err, fmt.Sprintf("unexpected error: %v", err)) - assert.Equal(t, tc.usersPage, page, fmt.Sprintf("%s unexpected response: expected %v got %v", tc.desc, tc.usersPage, page)) - case errLog: - assert.Equal(t, tc.errLogMessage, out, fmt.Sprintf("%s unexpected error response: expected %s got errLogMessage:%s", tc.desc, tc.errLogMessage, out)) - case usageLog: - assert.False(t, strings.Contains(out, rootCmd.Use), fmt.Sprintf("%s invalid usage: %s", tc.desc, out)) - } - - sdkCall.Unset() - }) - } -} diff --git a/clients/README.md b/clients/README.md deleted file mode 100644 index 8e2465400..000000000 --- a/clients/README.md +++ /dev/null @@ -1,371 +0,0 @@ -# Clients - -Clients service provides an HTTP API for managing platform resources: `clients` and `channels`. -Through this API clients are able to do the following actions: - -- provision new clients -- create new channels -- "connect" clients into the channels - -For an in-depth explanation of the aforementioned scenarios, as well as thorough -understanding of Magistrala, please check out the [official documentation][doc]. - -## Configuration - -The service is configured using the environment variables presented in the -following table. Note that any unset variables will be replaced with their -default values. - -| Variable | Description | Default | -| ----------------------------- | ----------------------------------------------------------------------- | ------------------------------ | -| MG_CLIENTS_LOG_LEVEL | Log level for Clients (debug, info, warn, error) | info | -| MG_CLIENTS_HTTP_HOST | Clients service HTTP host | localhost | -| MG_CLIENTS_HTTP_PORT | Clients service HTTP port | 9000 | -| MG_CLIENTS_SERVER_CERT | Path to the PEM encoded server certificate file | "" | -| MG_CLIENTS_SERVER_KEY | Path to the PEM encoded server key file | "" | -| MG_CLIENTS_GRPC_HOST | Clients service gRPC host | localhost | -| MG_CLIENTS_GRPC_PORT | Clients service gRPC port | 7000 | -| MG_CLIENTS_GRPC_SERVER_CERT | Path to the PEM encoded server certificate file | "" | -| MG_CLIENTS_GRPC_SERVER_KEY | Path to the PEM encoded server key file | "" | -| MG_CLIENTS_DB_HOST | Database host address | localhost | -| MG_CLIENTS_DB_PORT | Database host port | 5432 | -| MG_CLIENTS_DB_USER | Database user | magistrala | -| MG_CLIENTS_DB_PASS | Database password | magistrala | -| MG_CLIENTS_DB_NAME | Name of the database used by the service | clients | -| MG_CLIENTS_DB_SSL_MODE | Database connection SSL mode (disable, require, verify-ca, verify-full) | disable | -| MG_CLIENTS_DB_SSL_CERT | Path to the PEM encoded certificate file | "" | -| MG_CLIENTS_DB_SSL_KEY | Path to the PEM encoded key file | "" | -| MG_CLIENTS_DB_SSL_ROOT_CERT | Path to the PEM encoded root certificate file | "" | -| MG_CLIENTS_CACHE_URL | Cache database URL | | -| MG_CLIENTS_CACHE_KEY_DURATION | Cache key duration in seconds | 3600 | -| MG_CLIENTS_ES_URL | Event store URL | | -| MG_CLIENTS_ES_PASS | Event store password | "" | -| MG_CLIENTS_ES_DB | Event store instance name | 0 | -| MG_CLIENTS_STANDALONE_ID | User ID for standalone mode (no gRPC communication with Auth) | "" | -| MG_CLIENTS_STANDALONE_TOKEN | User token for standalone mode that should be passed in auth header | "" | -| MG_JAEGER_URL | Jaeger server URL | | -| MG_AUTH_GRPC_URL | Auth service gRPC URL | localhost:7001 | -| MG_AUTH_GRPC_TIMEOUT | Auth service gRPC request timeout in seconds | 1s | -| MG_AUTH_GRPC_CLIENT_TLS | Enable TLS for gRPC client | false | -| MG_AUTH_GRPC_CA_CERT | Path to the CA certificate file | "" | -| MG_SEND_TELEMETRY | Send telemetry to magistrala call home server. | true | -| Clients_INSTANCE_ID | Clients instance ID | "" | - -**Note** that if you want `clients` service to have only one user locally, you should use `CLIENTS_STANDALONE` env vars. By specifying these, you don't need `auth` service in your deployment for users' authorization. - -## Deployment - -The service itself is distributed as Docker container. Check the [`clients`](https://github.com/absmach/magistrala/blob/main/docker/docker-compose.yaml#L167-L194) service section in -docker-compose file to see how service is deployed. - -To start the service outside of the container, execute the following shell script: - -```bash -# download the latest version of the service -git clone https://github.com/absmach/magistrala - -cd magistrala - -# compile the clients -make clients - -# copy binary to bin -make install - -# set the environment variables and run the service -Clients_LOG_LEVEL=[Clients log level] \ -Clients_STANDALONE_ID=[User ID for standalone mode (no gRPC communication with auth)] \ -Clients_STANDALONE_TOKEN=[User token for standalone mode that should be passed in auth header] \ -Clients_CACHE_KEY_DURATION=[Cache key duration in seconds] \ -Clients_HTTP_HOST=[Clients service HTTP host] \ -Clients_HTTP_PORT=[Clients service HTTP port] \ -Clients_HTTP_SERVER_CERT=[Path to server certificate in pem format] \ -Clients_HTTP_SERVER_KEY=[Path to server key in pem format] \ -Clients_AUTH_GRPC_HOST=[Clients service gRPC host] \ -Clients_AUTH_GRPC_PORT=[Clients service gRPC port] \ -Clients_AUTH_GRPC_SERVER_CERT=[Path to server certificate in pem format] \ -Clients_AUTH_GRPC_SERVER_KEY=[Path to server key in pem format] \ -Clients_DB_HOST=[Database host address] \ -Clients_DB_PORT=[Database host port] \ -Clients_DB_USER=[Database user] \ -Clients_DB_PASS=[Database password] \ -Clients_DB_NAME=[Name of the database used by the service] \ -Clients_DB_SSL_MODE=[SSL mode to connect to the database with] \ -Clients_DB_SSL_CERT=[Path to the PEM encoded certificate file] \ -Clients_DB_SSL_KEY=[Path to the PEM encoded key file] \ -Clients_DB_SSL_ROOT_CERT=[Path to the PEM encoded root certificate file] \ -Clients_CACHE_URL=[Cache database URL] \ -Clients_ES_URL=[Event store URL] \ -Clients_ES_PASS=[Event store password] \ -Clients_ES_DB=[Event store instance name] \ -MG_AUTH_GRPC_URL=[Auth service gRPC URL] \ -MG_AUTH_GRPC_TIMEOUT=[Auth service gRPC request timeout in seconds] \ -MG_AUTH_GRPC_CLIENT_TLS=[Enable TLS for gRPC client] \ -MG_AUTH_GRPC_CA_CERT=[Path to trusted CA certificate file] \ -MG_JAEGER_URL=[Jaeger server URL] \ -MG_SEND_TELEMETRY=[Send telemetry to magistrala call home server] \ -Clients_INSTANCE_ID=[Clients instance ID] \ -$GOBIN/magistrala-clients -``` - -Setting `Clients_CA_CERTS` expects a file in PEM format of trusted CAs. This will enable TLS against the Auth gRPC endpoint trusting only those CAs that are provided. - -In constrained environments, sometimes it makes sense to run Clients service as a standalone to reduce network traffic and simplify deployment. This means that Clients service -operates only using a single user and is able to authorize it without gRPC communication with Auth service. -To run service in a standalone mode, set `Clients_STANDALONE_EMAIL` and `Clients_STANDALONE_TOKEN`. - -## Usage - -Magistrala supports the following operations for Clients: - -| Operation | Description | -| ------------------------- | -------------------------------------------- | -| `create` | Create a new client | -| `get` | Retrieve a single client or list all clients | -| `update` | Update a client’s name and metadata | -| `delete` | Permanently delete a client | -| `enable` | Enable a previously disabled client | -| `disable` | Disable an active client | -| `setClientParentGroup` | Add a Parent Group to a client | -| `removeClientParentGroup` | Remove a Parent Group from a client | - -### API Examples - -#### Create a Client - -```bash -curl -X POST http://localhost:9006//clients \ - -H "Authorization: Bearer " \ - -H "Content-Type: application/json" \ - -d '{ - "name": "clientName", - "tags": [ - "tag1", - "tag2" - ], - "credentials": { - "identity": "clientIDentity", - "secret": "bb7edb32-2eac-4aad-aebe-ed96fe073879" - }, - "metadata": { - "model": "example" - }, - "status": "enabled" -}' -``` - -The expected response should be: - -```bash -{ - "id": "bb7edb32-2eac-4aad-aebe-ed96fe073879", - "name": "clientName", - "tags": [ - "tag1", - "tag2" - ], - "domain_id": "bb7edb32-2eac-4aad-aebe-ed96fe073879", - "credentials": { - "identity": "clientIDentity", - "secret": "bb7edb32-2eac-4aad-aebe-ed96fe073879" - }, - "metadata": { - "model": "example" - }, - "status": "enabled", - "created_at": "2019-11-26 13:31:52", - "updated_at": "2019-11-26 13:31:52" -} -``` - -#### Get Clients - -List all clients: - -```bash -curl -X GET "http://localhost:9006//clients?limit=10" \ - -H "Authorization: Bearer " -``` - -List a singular client: - -```bash -curl -X GET http://localhost:9006//clients/ \ - -H "Authorization: Bearer " -``` - -#### Update a Client - -Update is performed by replacing the current resource data with values provided in a request payload. Note that the client's type and ID cannot be changed. - -```bash -curl -X PATCH http://localhost:9006//clients/ \ - -H "Authorization: Bearer " \ - -H "Content-Type: application/json" \ - -d '{ - "name": "clientName", - "metadata": {"role": "general"} - }' -``` - -The expected response is - -```bash -{ - "id": "bb7edb32-2eac-4aad-aebe-ed96fe073879", - "name": "clientName", - "tags": [ - "tag1", - "tag2" - ], - "domain_id": "bb7edb32-2eac-4aad-aebe-ed96fe073879", - "credentials": { - "identity": "clientIDentity", - "secret": "bb7edb32-2eac-4aad-aebe-ed96fe073879" - }, - "metadata": { "model": "example" }, - "status": "enabled", - "created_at": "2019-11-26 13:31:52", - "updated_at": "2019-11-26 13:31:52" -} -``` - -#### Delete a Client - -Delete client removes a client with the given id from repo and removes all the policies related to this client. - -```bash -curl -X DELETE http://localhost:9006//clients/ \ - -H "Authorization: Bearer " -``` - -#### Disable a Client - -Disables a specific client that is identified by the client ID. - -```bash -curl -X POST http://localhost:9006//clients//disable \ - -H "Authorization: Bearer " -``` - -#### Enable a Client - -Enable logically enables the client identified with the provided ID - -```bash -curl -X POST http://localhost:9006//clients//enable \ - -H "Authorization: Bearer " -``` - -## Roles Management for Clients - -In addition to standard client lifecycle operations (create, get, update, delete, enable, disable), the Clients service supports robust role‑based operations for managing permissions and associations for each client. - -### Supported Role Operations - -| Operation | Description | -| ------------------------- | ------------------------------------------------------------------- | -| `create-role` | Create a new role for a client | -| `list-roles` | List all roles assigned to a client | -| `get-role` | Retrieve details for a specific client role | -| `update-role` | Update a specific client role | -| `delete-role` | Delete a specific client role | -| `add-role-action` | Add one or more actions (permissions) to a client role | -| `list-role-actions` | List all actions associated with a client role | -| `delete-role-action` | Remove a specific action from a client role | -| `delete-all-role-actions` | Remove all actions from a client role | -| `add-role-member` | Associate one or more users or entities to a client role | -| `list-role-members` | List all members of a client role | -| `delete-role-member` | Remove one or more members from a client role | -| `delete-all-role-members` | Remove all members from a client role | -| `list-available-actions` | Retrieve the global list of available actions key for role creation | - -### Example: Create a Client Role - -```bash -curl -X POST http://localhost:9006//clients//roles \ - -H "Authorization: Bearer " \ - -H "Content-Type: application/json" \ - -d '{ - "name": "publisher", - "actions": ["publish"], - "members": [] - }' -``` - -## Implementation Details - -Clients in Magistrala are persisted in PostgreSQL using a schema optimized for identity management, authorization, and relationship tracking (channels, groups, and users). - -### Clients Table Structure - -The main `clients` table tracks all metadata, identity, and lifecycle information for each client: - -| Column | Type | Description | -| ----------------- | ------------- | ------------------------------------------------------ | -| `id` | VARCHAR(36) | UUID of the client (primary key). | -| `name` | VARCHAR(1024) | Human‑readable name. | -| `domain_id` | VARCHAR(36) | Domain to which the client belongs. | -| `parent_group_id` | VARCHAR(36) | Optional group parent (for inheritance/scoping). | -| `identity` | VARCHAR(254) | Login identity (often an email or unique ID). | -| `secret` | VARCHAR(4096) | Hashed authentication secret. | -| `tags` | TEXT[] | Arbitrary list of client tags. | -| `metadata` | JSONB | Free‑form structured metadata. | -| `created_at` | TIMESTAMPTZ | Timestamp when the client was created. | -| `updated_at` | TIMESTAMPTZ | Timestamp when the client was last updated. | -| `updated_by` | VARCHAR(254) | Identifier of the actor who performed the last update. | -| `status` | SMALLINT | 0 = enabled, 1 = disabled. | - -#### Connections Table Structure - -Client ↔ Channel relationships are stored in the `connections` table: - -| Column | Type | Description | -| ------------ | ----------- | ------------------------------------------------ | -| `channel_id` | VARCHAR(36) | Channel UUID. | -| `domain_id` | VARCHAR(36) | Domain of the client & channel. | -| `client_id` | VARCHAR(36) | Client UUID. | -| `type` | SMALLINT | Connection type: `1 = Publish`, `2 = Subscribe`. | - -This guarantees that when a client is deleted, all channel connections are automatically removed. - -## Best Practices - -To ensure robust and secure usage of the Clients service, consider the following recommendations: - -- **Use metadata and tags meaningfully**: Store useful attributes like model, location, environment (e.g., `production`, `test`) to filter and manage clients efficiently. -- **Keep credentials secure**: Rotate client secrets periodically. Avoid using guessable strings. -- **Disable unused clients**: Use the `disable` operation to revoke access instead of deleting clients when deactivation is preferred. -- **Audit regularly**: Periodically list client roles and connections to ensure expected configuration. -- **Prefer standalone mode for edge deployments**: Use environment variables to configure standalone mode in isolated environments without needing the Auth service. - -## Versioning and Health Check - -The Clients service exposes a `/health` endpoint to verify operational status and version information. - -### Health Check Request - -```bash -curl -X 'GET' \ - 'http://localhost:9006/health' \ - -H 'accept: application/health+json' -``` - -The expected response is: - -```bash -{ - "status": "pass", - "version": "0.14.0", - "commit": "7d6f4dc4f7f0c1fa3dc24eddfb18bb5073ff4f62", - "description": "clients service", - "build_time": "1970-01-01_00:00:00" -} -``` - -This endpoint can be used for monitoring, CI/CD readiness checks, or basic diagnostics. - -For more information about service capabilities and its usage, please check out -the [API documentation](https://docs.api.magistrala.absmach.eu/?urls.primaryName=api%2Fclients.yaml). - -[doc]: https://magistrala.absmach.eu/docs/ \ No newline at end of file diff --git a/clients/api/doc.go b/clients/api/doc.go deleted file mode 100644 index 2424852cc..000000000 --- a/clients/api/doc.go +++ /dev/null @@ -1,6 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package api contains API-related concerns: endpoint definitions, middlewares -// and all resource representations. -package api diff --git a/clients/api/grpc/client.go b/clients/api/grpc/client.go deleted file mode 100644 index 5c11332f5..000000000 --- a/clients/api/grpc/client.go +++ /dev/null @@ -1,372 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package grpc - -import ( - "context" - "fmt" - "time" - - grpcClientsV1 "github.com/absmach/magistrala/api/grpc/clients/v1" - grpcCommonV1 "github.com/absmach/magistrala/api/grpc/common/v1" - "github.com/absmach/magistrala/clients" - "github.com/absmach/magistrala/pkg/connections" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/go-kit/kit/endpoint" - kitgrpc "github.com/go-kit/kit/transport/grpc" - "google.golang.org/grpc" - "google.golang.org/grpc/codes" - "google.golang.org/grpc/status" -) - -const svcName = "clients.v1.ClientsService" - -var _ grpcClientsV1.ClientsServiceClient = (*grpcClient)(nil) - -type grpcClient struct { - timeout time.Duration - authenticate endpoint.Endpoint - retrieveEntity endpoint.Endpoint - retrieveEntities endpoint.Endpoint - addConnections endpoint.Endpoint - removeConnections endpoint.Endpoint - removeChannelConnections endpoint.Endpoint - unsetParentGroupFromClient endpoint.Endpoint -} - -// NewClient returns new gRPC client instance. -func NewClient(conn *grpc.ClientConn, timeout time.Duration) grpcClientsV1.ClientsServiceClient { - return &grpcClient{ - authenticate: kitgrpc.NewClient( - conn, - svcName, - "Authenticate", - encodeAuthenticateRequest, - decodeAuthenticateResponse, - grpcClientsV1.AuthnRes{}, - ).Endpoint(), - - retrieveEntity: kitgrpc.NewClient( - conn, - svcName, - "RetrieveEntity", - encodeRetrieveEntityRequest, - decodeRetrieveEntityResponse, - grpcCommonV1.RetrieveEntityRes{}, - ).Endpoint(), - - retrieveEntities: kitgrpc.NewClient( - conn, - svcName, - "RetrieveEntities", - encodeRetrieveEntitiesRequest, - decodeRetrieveEntitiesResponse, - grpcCommonV1.RetrieveEntitiesRes{}, - ).Endpoint(), - - addConnections: kitgrpc.NewClient( - conn, - svcName, - "AddConnections", - encodeAddConnectionsRequest, - decodeAddConnectionsResponse, - grpcCommonV1.AddConnectionsRes{}, - ).Endpoint(), - - removeConnections: kitgrpc.NewClient( - conn, - svcName, - "RemoveConnections", - encodeRemoveConnectionsRequest, - decodeRemoveConnectionsResponse, - grpcCommonV1.RemoveConnectionsRes{}, - ).Endpoint(), - - removeChannelConnections: kitgrpc.NewClient( - conn, - svcName, - "RemoveChannelConnections", - encodeRemoveChannelConnectionsRequest, - decodeRemoveChannelConnectionsResponse, - grpcClientsV1.RemoveChannelConnectionsRes{}, - ).Endpoint(), - - unsetParentGroupFromClient: kitgrpc.NewClient( - conn, - svcName, - "UnsetParentGroupFromClient", - encodeUnsetParentGroupFromClientRequest, - decodeUnsetParentGroupFromClientResponse, - grpcClientsV1.UnsetParentGroupFromClientRes{}, - ).Endpoint(), - - timeout: timeout, - } -} - -func (client grpcClient) Authenticate(ctx context.Context, req *grpcClientsV1.AuthnReq, _ ...grpc.CallOption) (r *grpcClientsV1.AuthnRes, err error) { - ctx, cancel := context.WithTimeout(ctx, client.timeout) - defer cancel() - - res, err := client.authenticate(ctx, authenticateReq{ - Token: req.GetToken(), - }) - if err != nil { - return &grpcClientsV1.AuthnRes{}, decodeError(err) - } - - ar := res.(authenticateRes) - return &grpcClientsV1.AuthnRes{Authenticated: ar.authenticated, Id: ar.id}, nil -} - -func encodeAuthenticateRequest(_ context.Context, grpcReq any) (any, error) { - req := grpcReq.(authenticateReq) - return &grpcClientsV1.AuthnReq{ - Token: req.Token, - }, nil -} - -func decodeAuthenticateResponse(_ context.Context, grpcRes any) (any, error) { - res := grpcRes.(*grpcClientsV1.AuthnRes) - return authenticateRes{authenticated: res.GetAuthenticated(), id: res.GetId()}, nil -} - -func (client grpcClient) RetrieveEntity(ctx context.Context, req *grpcCommonV1.RetrieveEntityReq, _ ...grpc.CallOption) (r *grpcCommonV1.RetrieveEntityRes, err error) { - ctx, cancel := context.WithTimeout(ctx, client.timeout) - defer cancel() - - res, err := client.retrieveEntity(ctx, req.GetId()) - if err != nil { - return &grpcCommonV1.RetrieveEntityRes{}, decodeError(err) - } - - ebr := res.(retrieveEntityRes) - - return &grpcCommonV1.RetrieveEntityRes{Entity: &grpcCommonV1.EntityBasic{Id: ebr.id, DomainId: ebr.domain, Status: uint32(ebr.status)}}, nil -} - -func encodeRetrieveEntityRequest(_ context.Context, grpcReq any) (any, error) { - req := grpcReq.(string) - return &grpcCommonV1.RetrieveEntityReq{ - Id: req, - }, nil -} - -func decodeRetrieveEntityResponse(_ context.Context, grpcRes any) (any, error) { - res := grpcRes.(*grpcCommonV1.RetrieveEntityRes) - - return retrieveEntityRes{ - id: res.Entity.GetId(), - domain: res.Entity.GetDomainId(), - parentGroup: res.Entity.GetParentGroupId(), - status: uint8(res.Entity.GetStatus()), - }, nil -} - -func (client grpcClient) RetrieveEntities(ctx context.Context, req *grpcCommonV1.RetrieveEntitiesReq, _ ...grpc.CallOption) (r *grpcCommonV1.RetrieveEntitiesRes, err error) { - ctx, cancel := context.WithTimeout(ctx, client.timeout) - defer cancel() - - res, err := client.retrieveEntities(ctx, req.GetIds()) - if err != nil { - return &grpcCommonV1.RetrieveEntitiesRes{}, decodeError(err) - } - - ep := res.(retrieveEntitiesRes) - - entities := []*grpcCommonV1.EntityBasic{} - for _, c := range ep.clients { - entities = append(entities, &grpcCommonV1.EntityBasic{ - Id: c.id, - DomainId: c.domain, - Status: uint32(c.status), - }) - } - return &grpcCommonV1.RetrieveEntitiesRes{Total: ep.total, Limit: ep.limit, Offset: ep.offset, Entities: entities}, nil -} - -func encodeRetrieveEntitiesRequest(_ context.Context, grpcReq any) (any, error) { - req := grpcReq.([]string) - return &grpcCommonV1.RetrieveEntitiesReq{ - Ids: req, - }, nil -} - -func decodeRetrieveEntitiesResponse(_ context.Context, grpcRes any) (any, error) { - res := grpcRes.(*grpcCommonV1.RetrieveEntitiesRes) - - clis := []entity{} - - for _, e := range res.Entities { - clis = append(clis, entity{ - id: e.GetId(), - domain: e.GetDomainId(), - parentGroup: e.GetParentGroupId(), - status: uint8(e.GetStatus()), - }) - } - return retrieveEntitiesRes{total: res.GetTotal(), limit: res.GetLimit(), offset: res.GetOffset(), clients: clis}, nil -} - -func (client grpcClient) AddConnections(ctx context.Context, req *grpcCommonV1.AddConnectionsReq, _ ...grpc.CallOption) (r *grpcCommonV1.AddConnectionsRes, err error) { - ctx, cancel := context.WithTimeout(ctx, client.timeout) - defer cancel() - - conns := []clients.Connection{} - for _, c := range req.Connections { - conns = append(conns, clients.Connection{ - ClientID: c.GetClientId(), - ChannelID: c.GetChannelId(), - DomainID: c.GetDomainId(), - Type: connections.ConnType(c.GetType()), - }) - } - - res, err := client.addConnections(ctx, conns) - if err != nil { - return &grpcCommonV1.AddConnectionsRes{}, decodeError(err) - } - - cr := res.(connectionsRes) - - return &grpcCommonV1.AddConnectionsRes{Ok: cr.ok}, nil -} - -func encodeAddConnectionsRequest(_ context.Context, grpcReq any) (any, error) { - req := grpcReq.([]clients.Connection) - - conns := []*grpcCommonV1.Connection{} - - for _, r := range req { - conns = append(conns, &grpcCommonV1.Connection{ - ClientId: r.ClientID, - ChannelId: r.ChannelID, - DomainId: r.DomainID, - Type: uint32(r.Type), - }) - } - return &grpcCommonV1.AddConnectionsReq{ - Connections: conns, - }, nil -} - -func decodeAddConnectionsResponse(_ context.Context, grpcRes any) (any, error) { - res := grpcRes.(*grpcCommonV1.AddConnectionsRes) - - return connectionsRes{ok: res.GetOk()}, nil -} - -func (client grpcClient) RemoveConnections(ctx context.Context, req *grpcCommonV1.RemoveConnectionsReq, _ ...grpc.CallOption) (r *grpcCommonV1.RemoveConnectionsRes, err error) { - ctx, cancel := context.WithTimeout(ctx, client.timeout) - defer cancel() - - conns := []clients.Connection{} - for _, c := range req.Connections { - conns = append(conns, clients.Connection{ - ClientID: c.GetClientId(), - ChannelID: c.GetChannelId(), - DomainID: c.GetDomainId(), - Type: connections.ConnType(c.GetType()), - }) - } - - res, err := client.removeConnections(ctx, conns) - if err != nil { - return &grpcCommonV1.RemoveConnectionsRes{}, decodeError(err) - } - - cr := res.(connectionsRes) - - return &grpcCommonV1.RemoveConnectionsRes{Ok: cr.ok}, nil -} - -func encodeRemoveConnectionsRequest(_ context.Context, grpcReq any) (any, error) { - req := grpcReq.([]clients.Connection) - - conns := []*grpcCommonV1.Connection{} - - for _, r := range req { - conns = append(conns, &grpcCommonV1.Connection{ - ClientId: r.ClientID, - ChannelId: r.ChannelID, - DomainId: r.DomainID, - Type: uint32(r.Type), - }) - } - return &grpcCommonV1.RemoveConnectionsReq{ - Connections: conns, - }, nil -} - -func decodeRemoveConnectionsResponse(_ context.Context, grpcRes any) (any, error) { - res := grpcRes.(*grpcCommonV1.RemoveConnectionsRes) - - return connectionsRes{ok: res.GetOk()}, nil -} - -func (client grpcClient) RemoveChannelConnections(ctx context.Context, req *grpcClientsV1.RemoveChannelConnectionsReq, _ ...grpc.CallOption) (r *grpcClientsV1.RemoveChannelConnectionsRes, err error) { - ctx, cancel := context.WithTimeout(ctx, client.timeout) - defer cancel() - - if _, err := client.removeChannelConnections(ctx, req); err != nil { - return &grpcClientsV1.RemoveChannelConnectionsRes{}, decodeError(err) - } - - return &grpcClientsV1.RemoveChannelConnectionsRes{}, nil -} - -func encodeRemoveChannelConnectionsRequest(_ context.Context, grpcReq any) (any, error) { - return grpcReq.(*grpcClientsV1.RemoveChannelConnectionsReq), nil -} - -func decodeRemoveChannelConnectionsResponse(_ context.Context, grpcRes any) (any, error) { - return grpcRes.(*grpcClientsV1.RemoveChannelConnectionsRes), nil -} - -func (client grpcClient) UnsetParentGroupFromClient(ctx context.Context, req *grpcClientsV1.UnsetParentGroupFromClientReq, _ ...grpc.CallOption) (r *grpcClientsV1.UnsetParentGroupFromClientRes, err error) { - ctx, cancel := context.WithTimeout(ctx, client.timeout) - defer cancel() - - if _, err := client.unsetParentGroupFromClient(ctx, req); err != nil { - return &grpcClientsV1.UnsetParentGroupFromClientRes{}, decodeError(err) - } - - return &grpcClientsV1.UnsetParentGroupFromClientRes{}, nil -} - -func encodeUnsetParentGroupFromClientRequest(_ context.Context, grpcReq any) (any, error) { - return grpcReq.(*grpcClientsV1.UnsetParentGroupFromClientReq), nil -} - -func decodeUnsetParentGroupFromClientResponse(_ context.Context, grpcRes any) (any, error) { - return grpcRes.(*grpcClientsV1.UnsetParentGroupFromClientRes), nil -} - -func decodeError(err error) error { - if st, ok := status.FromError(err); ok { - switch st.Code() { - case codes.Unauthenticated: - return errors.Wrap(svcerr.ErrAuthentication, errors.New(st.Message())) - case codes.PermissionDenied: - return errors.Wrap(svcerr.ErrAuthorization, errors.New(st.Message())) - case codes.InvalidArgument: - return errors.Wrap(errors.ErrMalformedEntity, errors.New(st.Message())) - case codes.FailedPrecondition: - return errors.Wrap(errors.ErrMalformedEntity, errors.New(st.Message())) - case codes.NotFound: - return errors.Wrap(svcerr.ErrNotFound, errors.New(st.Message())) - case codes.AlreadyExists: - return errors.Wrap(svcerr.ErrConflict, errors.New(st.Message())) - case codes.OK: - if msg := st.Message(); msg != "" { - return errors.Wrap(errors.ErrUnidentified, errors.New(msg)) - } - return nil - default: - return errors.Wrap(fmt.Errorf("unexpected gRPC status: %s (status code:%v)", st.Code().String(), st.Code()), errors.New(st.Message())) - } - } - return err -} diff --git a/clients/api/grpc/doc.go b/clients/api/grpc/doc.go deleted file mode 100644 index 20956ee50..000000000 --- a/clients/api/grpc/doc.go +++ /dev/null @@ -1,5 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package grpc contains implementation of Auth service gRPC API. -package grpc diff --git a/clients/api/grpc/endpoint.go b/clients/api/grpc/endpoint.go deleted file mode 100644 index 25b7840ac..000000000 --- a/clients/api/grpc/endpoint.go +++ /dev/null @@ -1,127 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package grpc - -import ( - "context" - - "github.com/absmach/magistrala/clients" - pClients "github.com/absmach/magistrala/clients/private" - "github.com/go-kit/kit/endpoint" -) - -func authenticateEndpoint(svc pClients.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(authenticateReq) - id, err := svc.Authenticate(ctx, req.Token) - if err != nil { - return authenticateRes{}, err - } - return authenticateRes{ - authenticated: true, - id: id, - }, err - } -} - -func retrieveEntityEndpoint(svc pClients.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(retrieveEntityReq) - client, err := svc.RetrieveById(ctx, req.Id) - if err != nil { - return retrieveEntityRes{}, err - } - - return retrieveEntityRes{id: client.ID, domain: client.Domain, parentGroup: client.ParentGroup, status: uint8(client.Status)}, nil - } -} - -func retrieveEntitiesEndpoint(svc pClients.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(retrieveEntitiesReq) - tp, err := svc.RetrieveByIds(ctx, req.Ids) - if err != nil { - return retrieveEntitiesRes{}, err - } - clientsBasic := []entity{} - for _, client := range tp.Clients { - clientsBasic = append(clientsBasic, entity{id: client.ID, domain: client.Domain, parentGroup: client.ParentGroup, status: uint8(client.Status)}) - } - return retrieveEntitiesRes{ - total: tp.Total, - limit: tp.Limit, - offset: tp.Offset, - clients: clientsBasic, - }, nil - } -} - -func addConnectionsEndpoint(svc pClients.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(connectionsReq) - - var conns []clients.Connection - - for _, c := range req.connections { - conns = append(conns, clients.Connection{ - ClientID: c.clientID, - ChannelID: c.channelID, - DomainID: c.domainID, - Type: c.connType, - }) - } - - if err := svc.AddConnections(ctx, conns); err != nil { - return connectionsRes{ok: false}, err - } - - return connectionsRes{ok: true}, nil - } -} - -func removeConnectionsEndpoint(svc pClients.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(connectionsReq) - - var conns []clients.Connection - - for _, c := range req.connections { - conns = append(conns, clients.Connection{ - ClientID: c.clientID, - ChannelID: c.channelID, - DomainID: c.domainID, - Type: c.connType, - }) - } - if err := svc.RemoveConnections(ctx, conns); err != nil { - return connectionsRes{ok: false}, err - } - - return connectionsRes{ok: true}, nil - } -} - -func removeChannelConnectionsEndpoint(svc pClients.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(removeChannelConnectionsReq) - - if err := svc.RemoveChannelConnections(ctx, req.channelID); err != nil { - return removeChannelConnectionsRes{}, err - } - - return removeChannelConnectionsRes{}, nil - } -} - -func UnsetParentGroupFromClientEndpoint(svc pClients.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(UnsetParentGroupFromClientReq) - - if err := svc.UnsetParentGroupFromClient(ctx, req.parentGroupID); err != nil { - return UnsetParentGroupFromClientRes{}, err - } - - return UnsetParentGroupFromClientRes{}, nil - } -} diff --git a/clients/api/grpc/endpoint_test.go b/clients/api/grpc/endpoint_test.go deleted file mode 100644 index 7140ea616..000000000 --- a/clients/api/grpc/endpoint_test.go +++ /dev/null @@ -1,421 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package grpc_test - -import ( - "context" - "fmt" - "net" - "testing" - "time" - - grpcClientsV1 "github.com/absmach/magistrala/api/grpc/clients/v1" - grpcCommonV1 "github.com/absmach/magistrala/api/grpc/common/v1" - "github.com/absmach/magistrala/clients" - grpcapi "github.com/absmach/magistrala/clients/api/grpc" - "github.com/absmach/magistrala/clients/private/mocks" - "github.com/absmach/magistrala/internal/testsutil" - "github.com/absmach/magistrala/pkg/connections" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" - "google.golang.org/grpc" - "google.golang.org/grpc/credentials/insecure" -) - -const port = 7006 - -var ( - validID = testsutil.GenerateUUID(&testing.T{}) - validSecret = "validSecret" - invalidSecret = "invalidSecret" - validClient = clients.Client{ - ID: validID, - Domain: validID, - Status: clients.EnabledStatus, - } -) - -func startGRPCServer(svc *mocks.Service, port int) *grpc.Server { - listener, err := net.Listen("tcp", fmt.Sprintf(":%d", port)) - if err != nil { - panic(fmt.Sprintf("failed to obtain port: %s", err)) - } - server := grpc.NewServer() - grpcClientsV1.RegisterClientsServiceServer(server, grpcapi.NewServer(svc)) - go func() { - if err := server.Serve(listener); err != nil { - panic(fmt.Sprintf("failed to serve: %s", err)) - } - }() - - return server -} - -func TestAuthenticate(t *testing.T) { - svc := new(mocks.Service) - server := startGRPCServer(svc, port) - defer server.GracefulStop() - authAddr := fmt.Sprintf("localhost:%d", port) - conn, _ := grpc.NewClient(authAddr, grpc.WithTransportCredentials(insecure.NewCredentials())) - client := grpcapi.NewClient(conn, time.Second) - - cases := []struct { - desc string - clientSecret string - clientID string - resp *grpcClientsV1.AuthnRes - svcErr error - err error - }{ - { - desc: "authenticate successfully", - clientSecret: validSecret, - resp: &grpcClientsV1.AuthnRes{ - Authenticated: true, - Id: validID, - }, - clientID: validID, - svcErr: nil, - err: nil, - }, - { - desc: "failed to authenticate", - clientSecret: invalidSecret, - resp: &grpcClientsV1.AuthnRes{ - Authenticated: false, - Id: "", - }, - clientID: "", - svcErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("Authenticate", mock.Anything, tc.clientSecret).Return(tc.clientID, tc.svcErr) - res, err := client.Authenticate(context.Background(), &grpcClientsV1.AuthnReq{Token: tc.clientSecret}) - assert.True(t, errors.Contains(err, tc.err)) - assert.Equal(t, tc.resp, res) - svcCall.Unset() - }) - } -} - -func TestRetrieveEntity(t *testing.T) { - svc := new(mocks.Service) - server := startGRPCServer(svc, port) - defer server.GracefulStop() - authAddr := fmt.Sprintf("localhost:%d", port) - conn, _ := grpc.NewClient(authAddr, grpc.WithTransportCredentials(insecure.NewCredentials())) - client := grpcapi.NewClient(conn, time.Second) - - cases := []struct { - desc string - id string - svcRes clients.Client - resp *grpcCommonV1.RetrieveEntityRes - svcErr error - err error - }{ - { - desc: "retrieve entity successfully", - id: validID, - svcRes: validClient, - resp: &grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Id: validID, - DomainId: validID, - Status: uint32(clients.EnabledStatus), - }, - }, - err: nil, - }, - { - desc: "retrieve entity with empty ID", - id: "", - resp: &grpcCommonV1.RetrieveEntityRes{}, - svcErr: svcerr.ErrNotFound, - err: svcerr.ErrNotFound, - }, - { - desc: "retrieve entity with invalid ID", - id: "invalidID", - resp: &grpcCommonV1.RetrieveEntityRes{}, - svcErr: svcerr.ErrNotFound, - err: svcerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("RetrieveById", mock.Anything, tc.id).Return(tc.svcRes, tc.svcErr) - res, err := client.RetrieveEntity(context.Background(), &grpcCommonV1.RetrieveEntityReq{Id: tc.id}) - assert.True(t, errors.Contains(err, tc.err)) - assert.Equal(t, tc.resp, res) - svcCall.Unset() - }) - } -} - -func TestRetrieveEntities(t *testing.T) { - svc := new(mocks.Service) - server := startGRPCServer(svc, port) - defer server.GracefulStop() - authAddr := fmt.Sprintf("localhost:%d", port) - conn, _ := grpc.NewClient(authAddr, grpc.WithTransportCredentials(insecure.NewCredentials())) - client := grpcapi.NewClient(conn, time.Second) - - cases := []struct { - desc string - ids []string - svcRes clients.ClientsPage - resp *grpcCommonV1.RetrieveEntitiesRes - svcErr error - err error - }{ - { - desc: "retrieve entities successfully", - ids: []string{validID}, - svcRes: clients.ClientsPage{ - Page: clients.Page{ - Total: 1, - Limit: 1, - }, - Clients: []clients.Client{validClient}, - }, - resp: &grpcCommonV1.RetrieveEntitiesRes{ - Total: 1, - Limit: 1, - Offset: 0, - Entities: []*grpcCommonV1.EntityBasic{ - { - Id: validID, - DomainId: validID, - Status: uint32(clients.EnabledStatus), - }, - }, - }, - err: nil, - }, - { - desc: "retrieve entities with empty IDs", - ids: []string(nil), - resp: &grpcCommonV1.RetrieveEntitiesRes{}, - svcErr: svcerr.ErrNotFound, - err: svcerr.ErrNotFound, - }, - { - desc: "retrieve entities with invalid IDs", - ids: []string{"invalidID"}, - resp: &grpcCommonV1.RetrieveEntitiesRes{}, - svcErr: svcerr.ErrNotFound, - err: svcerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("RetrieveByIds", mock.Anything, tc.ids).Return(tc.svcRes, tc.svcErr) - res, err := client.RetrieveEntities(context.Background(), &grpcCommonV1.RetrieveEntitiesReq{Ids: tc.ids}) - assert.True(t, errors.Contains(err, tc.err)) - assert.Equal(t, tc.resp, res) - svcCall.Unset() - }) - } -} - -func TestAddConnections(t *testing.T) { - svc := new(mocks.Service) - server := startGRPCServer(svc, port) - defer server.GracefulStop() - authAddr := fmt.Sprintf("localhost:%d", port) - conn, _ := grpc.NewClient(authAddr, grpc.WithTransportCredentials(insecure.NewCredentials())) - client := grpcapi.NewClient(conn, time.Second) - - cases := []struct { - desc string - req *grpcCommonV1.AddConnectionsReq - svcErr error - err error - }{ - { - desc: "add connections successfully", - req: &grpcCommonV1.AddConnectionsReq{ - Connections: []*grpcCommonV1.Connection{ - { - ClientId: validID, - ChannelId: validID, - DomainId: validID, - Type: uint32(connections.Publish), - }, - }, - }, - err: nil, - }, - { - desc: "add connections with invalid request", - req: &grpcCommonV1.AddConnectionsReq{ - Connections: []*grpcCommonV1.Connection{ - { - ClientId: "", - ChannelId: "", - DomainId: "", - Type: uint32(connections.Publish), - }, - }, - }, - svcErr: svcerr.ErrCreateEntity, - err: svcerr.ErrCreateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("AddConnections", mock.Anything, mock.Anything).Return(tc.svcErr) - _, err := client.AddConnections(context.Background(), tc.req) - assert.True(t, errors.Contains(err, tc.err)) - svcCall.Unset() - }) - } -} - -func TestRemoveConnections(t *testing.T) { - svc := new(mocks.Service) - server := startGRPCServer(svc, port) - defer server.GracefulStop() - authAddr := fmt.Sprintf("localhost:%d", port) - conn, _ := grpc.NewClient(authAddr, grpc.WithTransportCredentials(insecure.NewCredentials())) - client := grpcapi.NewClient(conn, time.Second) - - cases := []struct { - desc string - req *grpcCommonV1.RemoveConnectionsReq - svcErr error - err error - }{ - { - desc: "remove connections successfully", - req: &grpcCommonV1.RemoveConnectionsReq{ - Connections: []*grpcCommonV1.Connection{ - { - ClientId: validID, - ChannelId: validID, - DomainId: validID, - Type: uint32(connections.Publish), - }, - }, - }, - err: nil, - }, - { - desc: "remove connections with invalid request", - req: &grpcCommonV1.RemoveConnectionsReq{ - Connections: []*grpcCommonV1.Connection{ - { - ClientId: "", - ChannelId: "", - DomainId: "", - Type: uint32(connections.Publish), - }, - }, - }, - svcErr: svcerr.ErrRemoveEntity, - err: svcerr.ErrRemoveEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("RemoveConnections", mock.Anything, mock.Anything).Return(tc.svcErr) - _, err := client.RemoveConnections(context.Background(), tc.req) - assert.True(t, errors.Contains(err, tc.err)) - svcCall.Unset() - }) - } -} - -func TestRemoveChannelConnections(t *testing.T) { - svc := new(mocks.Service) - server := startGRPCServer(svc, port) - defer server.GracefulStop() - authAddr := fmt.Sprintf("localhost:%d", port) - conn, _ := grpc.NewClient(authAddr, grpc.WithTransportCredentials(insecure.NewCredentials())) - client := grpcapi.NewClient(conn, time.Second) - - cases := []struct { - desc string - req *grpcClientsV1.RemoveChannelConnectionsReq - svcErr error - err error - }{ - { - desc: "remove channel connections successfully", - req: &grpcClientsV1.RemoveChannelConnectionsReq{ - ChannelId: validID, - }, - err: nil, - }, - { - desc: "remove channel connections with invalid request", - req: &grpcClientsV1.RemoveChannelConnectionsReq{ - ChannelId: "", - }, - svcErr: svcerr.ErrRemoveEntity, - err: svcerr.ErrRemoveEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("RemoveChannelConnections", mock.Anything, tc.req.ChannelId).Return(tc.svcErr) - _, err := client.RemoveChannelConnections(context.Background(), tc.req) - assert.True(t, errors.Contains(err, tc.err)) - svcCall.Unset() - }) - } -} - -func TestUnsetParentGroupFromClient(t *testing.T) { - svc := new(mocks.Service) - server := startGRPCServer(svc, port) - defer server.GracefulStop() - authAddr := fmt.Sprintf("localhost:%d", port) - conn, _ := grpc.NewClient(authAddr, grpc.WithTransportCredentials(insecure.NewCredentials())) - client := grpcapi.NewClient(conn, time.Second) - - cases := []struct { - desc string - req *grpcClientsV1.UnsetParentGroupFromClientReq - svcErr error - err error - }{ - { - desc: "unset parent group successfully", - req: &grpcClientsV1.UnsetParentGroupFromClientReq{ - ParentGroupId: validID, - }, - err: nil, - }, - { - desc: "unset parent group with invalid request", - req: &grpcClientsV1.UnsetParentGroupFromClientReq{ - ParentGroupId: "", - }, - svcErr: svcerr.ErrRemoveEntity, - err: svcerr.ErrRemoveEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("UnsetParentGroupFromClient", mock.Anything, tc.req.ParentGroupId).Return(tc.svcErr) - _, err := client.UnsetParentGroupFromClient(context.Background(), tc.req) - assert.True(t, errors.Contains(err, tc.err)) - svcCall.Unset() - }) - } -} diff --git a/clients/api/grpc/request.go b/clients/api/grpc/request.go deleted file mode 100644 index a0c40a7b8..000000000 --- a/clients/api/grpc/request.go +++ /dev/null @@ -1,24 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package grpc - -type authenticateReq struct { - Token string -} - -type retrieveEntitiesReq struct { - Ids []string -} - -type retrieveEntityReq struct { - Id string -} - -type removeChannelConnectionsReq struct { - channelID string -} - -type UnsetParentGroupFromClientReq struct { - parentGroupID string -} diff --git a/clients/api/grpc/responses.go b/clients/api/grpc/responses.go deleted file mode 100644 index 0e014455a..000000000 --- a/clients/api/grpc/responses.go +++ /dev/null @@ -1,45 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package grpc - -import "github.com/absmach/magistrala/pkg/connections" - -type entity struct { - id string - domain string - parentGroup string - status uint8 -} - -type authenticateRes struct { - id string - authenticated bool -} - -type retrieveEntitiesRes struct { - total uint64 - limit uint64 - offset uint64 - clients []entity -} - -type retrieveEntityRes entity - -type connectionsReq struct { - connections []connection -} - -type connection struct { - clientID string - channelID string - domainID string - connType connections.ConnType -} -type connectionsRes struct { - ok bool -} - -type removeChannelConnectionsRes struct{} - -type UnsetParentGroupFromClientRes struct{} diff --git a/clients/api/grpc/server.go b/clients/api/grpc/server.go deleted file mode 100644 index b8dd81bca..000000000 --- a/clients/api/grpc/server.go +++ /dev/null @@ -1,288 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package grpc - -import ( - "context" - - grpcClientsV1 "github.com/absmach/magistrala/api/grpc/clients/v1" - grpcCommonV1 "github.com/absmach/magistrala/api/grpc/common/v1" - apiutil "github.com/absmach/magistrala/api/http/util" - smqauth "github.com/absmach/magistrala/auth" - clients "github.com/absmach/magistrala/clients/private" - "github.com/absmach/magistrala/pkg/connections" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - kitgrpc "github.com/go-kit/kit/transport/grpc" - "google.golang.org/grpc/codes" - "google.golang.org/grpc/status" -) - -var _ grpcClientsV1.ClientsServiceServer = (*grpcServer)(nil) - -type grpcServer struct { - grpcClientsV1.UnimplementedClientsServiceServer - authenticate kitgrpc.Handler - retrieveEntity kitgrpc.Handler - retrieveEntities kitgrpc.Handler - addConnections kitgrpc.Handler - removeConnections kitgrpc.Handler - removeChannelConnections kitgrpc.Handler - unsetParentGroupFromClient kitgrpc.Handler -} - -// NewServer returns new AuthServiceServer instance. -func NewServer(svc clients.Service) grpcClientsV1.ClientsServiceServer { - return &grpcServer{ - authenticate: kitgrpc.NewServer( - authenticateEndpoint(svc), - decodeAuthorizeRequest, - encodeAuthorizeResponse, - ), - retrieveEntity: kitgrpc.NewServer( - retrieveEntityEndpoint(svc), - decodeRetrieveEntityRequest, - encodeRetrieveEntityResponse, - ), - retrieveEntities: kitgrpc.NewServer( - retrieveEntitiesEndpoint(svc), - decodeRetrieveEntitiesRequest, - encodeRetrieveEntitiesResponse, - ), - addConnections: kitgrpc.NewServer( - addConnectionsEndpoint(svc), - decodeAddConnectionsRequest, - encodeAddConnectionsResponse, - ), - removeConnections: kitgrpc.NewServer( - removeConnectionsEndpoint(svc), - decodeRemoveConnectionsRequest, - encodeRemoveConnectionsResponse, - ), - removeChannelConnections: kitgrpc.NewServer( - removeChannelConnectionsEndpoint(svc), - decodeRemoveChannelConnectionsRequest, - encodeRemoveChannelConnectionsResponse, - ), - unsetParentGroupFromClient: kitgrpc.NewServer( - UnsetParentGroupFromClientEndpoint(svc), - decodeUnsetParentGroupFromClientRequest, - encodeUnsetParentGroupFromClientResponse, - ), - } -} - -func (s *grpcServer) Authenticate(ctx context.Context, req *grpcClientsV1.AuthnReq) (*grpcClientsV1.AuthnRes, error) { - _, res, err := s.authenticate.ServeGRPC(ctx, req) - if err != nil { - return nil, encodeError(err) - } - return res.(*grpcClientsV1.AuthnRes), nil -} - -func decodeAuthorizeRequest(_ context.Context, grpcReq any) (any, error) { - req := grpcReq.(*grpcClientsV1.AuthnReq) - return authenticateReq{ - Token: req.GetToken(), - }, nil -} - -func encodeAuthorizeResponse(_ context.Context, grpcRes any) (any, error) { - res := grpcRes.(authenticateRes) - return &grpcClientsV1.AuthnRes{Authenticated: res.authenticated, Id: res.id}, nil -} - -func (s *grpcServer) RetrieveEntity(ctx context.Context, req *grpcCommonV1.RetrieveEntityReq) (*grpcCommonV1.RetrieveEntityRes, error) { - _, res, err := s.retrieveEntity.ServeGRPC(ctx, req) - if err != nil { - return nil, encodeError(err) - } - return res.(*grpcCommonV1.RetrieveEntityRes), nil -} - -func decodeRetrieveEntityRequest(_ context.Context, grpcReq any) (any, error) { - req := grpcReq.(*grpcCommonV1.RetrieveEntityReq) - return retrieveEntityReq{ - Id: req.GetId(), - }, nil -} - -func encodeRetrieveEntityResponse(_ context.Context, grpcRes any) (any, error) { - res := grpcRes.(retrieveEntityRes) - - return &grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Id: res.id, - DomainId: res.domain, - ParentGroupId: res.parentGroup, - Status: uint32(res.status), - }, - }, nil -} - -func (s *grpcServer) RetrieveEntities(ctx context.Context, req *grpcCommonV1.RetrieveEntitiesReq) (*grpcCommonV1.RetrieveEntitiesRes, error) { - _, res, err := s.retrieveEntities.ServeGRPC(ctx, req) - if err != nil { - return nil, encodeError(err) - } - return res.(*grpcCommonV1.RetrieveEntitiesRes), nil -} - -func decodeRetrieveEntitiesRequest(_ context.Context, grpcReq any) (any, error) { - req := grpcReq.(*grpcCommonV1.RetrieveEntitiesReq) - return retrieveEntitiesReq{ - Ids: req.GetIds(), - }, nil -} - -func encodeRetrieveEntitiesResponse(_ context.Context, grpcRes any) (any, error) { - res := grpcRes.(retrieveEntitiesRes) - - entities := []*grpcCommonV1.EntityBasic{} - for _, c := range res.clients { - entities = append(entities, &grpcCommonV1.EntityBasic{ - Id: c.id, - DomainId: c.domain, - ParentGroupId: c.parentGroup, - Status: uint32(c.status), - }) - } - return &grpcCommonV1.RetrieveEntitiesRes{Total: res.total, Limit: res.limit, Offset: res.offset, Entities: entities}, nil -} - -func (s *grpcServer) AddConnections(ctx context.Context, req *grpcCommonV1.AddConnectionsReq) (*grpcCommonV1.AddConnectionsRes, error) { - _, res, err := s.addConnections.ServeGRPC(ctx, req) - if err != nil { - return nil, encodeError(err) - } - return res.(*grpcCommonV1.AddConnectionsRes), nil -} - -func decodeAddConnectionsRequest(_ context.Context, grpcReq any) (any, error) { - req := grpcReq.(*grpcCommonV1.AddConnectionsReq) - - conns := []connection{} - for _, c := range req.Connections { - connType := connections.ConnType(c.GetType()) - if err := connections.CheckConnType(connType); err != nil { - return nil, err - } - conns = append(conns, connection{ - clientID: c.GetClientId(), - channelID: c.GetChannelId(), - domainID: c.GetDomainId(), - connType: connType, - }) - } - return connectionsReq{ - connections: conns, - }, nil -} - -func encodeAddConnectionsResponse(_ context.Context, grpcRes any) (any, error) { - res := grpcRes.(connectionsRes) - - return &grpcCommonV1.AddConnectionsRes{Ok: res.ok}, nil -} - -func (s *grpcServer) RemoveConnections(ctx context.Context, req *grpcCommonV1.RemoveConnectionsReq) (*grpcCommonV1.RemoveConnectionsRes, error) { - _, res, err := s.removeConnections.ServeGRPC(ctx, req) - if err != nil { - return nil, encodeError(err) - } - return res.(*grpcCommonV1.RemoveConnectionsRes), nil -} - -func decodeRemoveConnectionsRequest(_ context.Context, grpcReq any) (any, error) { - req := grpcReq.(*grpcCommonV1.RemoveConnectionsReq) - - conns := []connection{} - for _, c := range req.Connections { - connType := connections.ConnType(c.GetType()) - if err := connections.CheckConnType(connType); err != nil { - return nil, err - } - conns = append(conns, connection{ - clientID: c.GetClientId(), - channelID: c.GetChannelId(), - domainID: c.GetDomainId(), - connType: connType, - }) - } - return connectionsReq{ - connections: conns, - }, nil -} - -func encodeRemoveConnectionsResponse(_ context.Context, grpcRes any) (any, error) { - res := grpcRes.(connectionsRes) - - return &grpcCommonV1.RemoveConnectionsRes{Ok: res.ok}, nil -} - -func (s *grpcServer) RemoveChannelConnections(ctx context.Context, req *grpcClientsV1.RemoveChannelConnectionsReq) (*grpcClientsV1.RemoveChannelConnectionsRes, error) { - _, res, err := s.removeChannelConnections.ServeGRPC(ctx, req) - if err != nil { - return nil, encodeError(err) - } - return res.(*grpcClientsV1.RemoveChannelConnectionsRes), nil -} - -func decodeRemoveChannelConnectionsRequest(_ context.Context, grpcReq any) (any, error) { - req := grpcReq.(*grpcClientsV1.RemoveChannelConnectionsReq) - - return removeChannelConnectionsReq{ - channelID: req.GetChannelId(), - }, nil -} - -func encodeRemoveChannelConnectionsResponse(_ context.Context, grpcRes any) (any, error) { - _ = grpcRes.(removeChannelConnectionsRes) - return &grpcClientsV1.RemoveChannelConnectionsRes{}, nil -} - -func (s *grpcServer) UnsetParentGroupFromClient(ctx context.Context, req *grpcClientsV1.UnsetParentGroupFromClientReq) (*grpcClientsV1.UnsetParentGroupFromClientRes, error) { - _, res, err := s.unsetParentGroupFromClient.ServeGRPC(ctx, req) - if err != nil { - return nil, encodeError(err) - } - return res.(*grpcClientsV1.UnsetParentGroupFromClientRes), nil -} - -func decodeUnsetParentGroupFromClientRequest(_ context.Context, grpcReq any) (any, error) { - req := grpcReq.(*grpcClientsV1.UnsetParentGroupFromClientReq) - - return UnsetParentGroupFromClientReq{ - parentGroupID: req.GetParentGroupId(), - }, nil -} - -func encodeUnsetParentGroupFromClientResponse(_ context.Context, grpcRes any) (any, error) { - _ = grpcRes.(UnsetParentGroupFromClientRes) - return &grpcClientsV1.UnsetParentGroupFromClientRes{}, nil -} - -func encodeError(err error) error { - switch { - case errors.Contains(err, nil): - return nil - case errors.Contains(err, errors.ErrMalformedEntity), - err == apiutil.ErrInvalidAuthKey, - err == apiutil.ErrMissingID, - err == apiutil.ErrMissingMemberType, - err == apiutil.ErrMissingPolicySub, - err == apiutil.ErrMissingPolicyObj, - err == apiutil.ErrMalformedPolicyAct: - return status.Error(codes.InvalidArgument, err.Error()) - case errors.Contains(err, svcerr.ErrAuthentication), - errors.Contains(err, smqauth.ErrKeyExpired), - err == apiutil.ErrMissingEmail, - err == apiutil.ErrBearerToken: - return status.Error(codes.Unauthenticated, err.Error()) - case errors.Contains(err, svcerr.ErrAuthorization): - return status.Error(codes.PermissionDenied, err.Error()) - default: - return status.Error(codes.Internal, err.Error()) - } -} diff --git a/clients/api/http/clients.go b/clients/api/http/clients.go deleted file mode 100644 index 101814d6c..000000000 --- a/clients/api/http/clients.go +++ /dev/null @@ -1,123 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package http - -import ( - "log/slog" - - "github.com/absmach/magistrala" - api "github.com/absmach/magistrala/api/http" - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/clients" - smqauthn "github.com/absmach/magistrala/pkg/authn" - roleManagerHttp "github.com/absmach/magistrala/pkg/roles/rolemanager/api" - "github.com/go-chi/chi/v5" - kithttp "github.com/go-kit/kit/transport/http" - "go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp" -) - -func clientsHandler(svc clients.Service, authn smqauthn.AuthNMiddleware, r *chi.Mux, logger *slog.Logger, idp magistrala.IDProvider) *chi.Mux { - opts := []kithttp.ServerOption{ - kithttp.ServerErrorEncoder(apiutil.LoggingErrorEncoder(logger, api.EncodeError)), - } - d := roleManagerHttp.NewDecoder("clientID") - - r.Group(func(r chi.Router) { - r.Use(authn.Middleware()) - r.Use(api.RequestIDMiddleware(idp)) - - r.Route("/{domainID}/clients", func(r chi.Router) { - r.Post("/", otelhttp.NewHandler(kithttp.NewServer( - createClientEndpoint(svc), - decodeCreateClientReq, - api.EncodeResponse, - opts..., - ), "create_client").ServeHTTP) - - r.Get("/", otelhttp.NewHandler(kithttp.NewServer( - listClientsEndpoint(svc), - decodeListClients, - api.EncodeResponse, - opts..., - ), "list_clients").ServeHTTP) - - r.Post("/bulk", otelhttp.NewHandler(kithttp.NewServer( - createClientsEndpoint(svc), - decodeCreateClientsReq, - api.EncodeResponse, - opts..., - ), "create_clients").ServeHTTP) - - r = roleManagerHttp.EntityAvailableActionsRouter(svc, d, r, opts) - - r.Route("/{clientID}", func(r chi.Router) { - r.Get("/", otelhttp.NewHandler(kithttp.NewServer( - viewClientEndpoint(svc), - decodeViewClient, - api.EncodeResponse, - opts..., - ), "view_client").ServeHTTP) - - r.Patch("/", otelhttp.NewHandler(kithttp.NewServer( - updateClientEndpoint(svc), - decodeUpdateClient, - api.EncodeResponse, - opts..., - ), "update_client").ServeHTTP) - - r.Patch("/tags", otelhttp.NewHandler(kithttp.NewServer( - updateClientTagsEndpoint(svc), - decodeUpdateClientTags, - api.EncodeResponse, - opts..., - ), "update_client_tags").ServeHTTP) - - r.Patch("/secret", otelhttp.NewHandler(kithttp.NewServer( - updateClientSecretEndpoint(svc), - decodeUpdateClientCredentials, - api.EncodeResponse, - opts..., - ), "update_client_credentials").ServeHTTP) - - r.Post("/enable", otelhttp.NewHandler(kithttp.NewServer( - enableClientEndpoint(svc), - decodeChangeClientStatus, - api.EncodeResponse, - opts..., - ), "enable_client").ServeHTTP) - - r.Post("/disable", otelhttp.NewHandler(kithttp.NewServer( - disableClientEndpoint(svc), - decodeChangeClientStatus, - api.EncodeResponse, - opts..., - ), "disable_client").ServeHTTP) - - r.Post("/parent", otelhttp.NewHandler(kithttp.NewServer( - setClientParentGroupEndpoint(svc), - decodeSetClientParentGroupStatus, - api.EncodeResponse, - opts..., - ), "set_client_parent_group").ServeHTTP) - - r.Delete("/parent", otelhttp.NewHandler(kithttp.NewServer( - removeClientParentGroupEndpoint(svc), - decodeRemoveClientParentGroupStatus, - api.EncodeResponse, - opts..., - ), "remove_client_parent_group").ServeHTTP) - - r.Delete("/", otelhttp.NewHandler(kithttp.NewServer( - deleteClientEndpoint(svc), - decodeDeleteClientReq, - api.EncodeResponse, - opts..., - ), "delete_client").ServeHTTP) - - roleManagerHttp.EntityRoleMangerRouter(svc, d, r, opts) - }) - }) - }) - return r -} diff --git a/clients/api/http/decode.go b/clients/api/http/decode.go deleted file mode 100644 index de35216a0..000000000 --- a/clients/api/http/decode.go +++ /dev/null @@ -1,302 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package http - -import ( - "context" - "encoding/json" - "net/http" - "strings" - "time" - - api "github.com/absmach/magistrala/api/http" - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/clients" - "github.com/absmach/magistrala/pkg/errors" - "github.com/go-chi/chi/v5" -) - -const clientID = "clientID" - -func decodeViewClient(_ context.Context, r *http.Request) (any, error) { - roles, err := apiutil.ReadBoolQuery(r, api.RolesKey, false) - if err != nil { - return listClientsReq{}, errors.Wrap(apiutil.ErrValidation, err) - } - - req := viewClientReq{ - id: chi.URLParam(r, clientID), - roles: roles, - } - - return req, nil -} - -func decodeListClients(_ context.Context, r *http.Request) (any, error) { - name, err := apiutil.ReadStringQuery(r, api.NameKey, "") - if err != nil { - return listClientsReq{}, errors.Wrap(apiutil.ErrValidation, err) - } - - tags, err := apiutil.ReadStringQuery(r, api.TagsKey, "") - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - var tq clients.TagsQuery - if tags != "" { - tq = clients.ToTagsQuery(tags) - } - - s, err := apiutil.ReadStringQuery(r, api.StatusKey, api.DefGroupStatus) - if err != nil { - return listClientsReq{}, errors.Wrap(apiutil.ErrValidation, err) - } - status, err := clients.ToStatus(s) - if err != nil { - return listClientsReq{}, errors.Wrap(apiutil.ErrValidation, err) - } - - meta, err := apiutil.ReadMetadataQuery(r, api.MetadataKey, nil) - if err != nil { - return listClientsReq{}, errors.Wrap(apiutil.ErrValidation, err) - } - - offset, err := apiutil.ReadNumQuery[uint64](r, api.OffsetKey, api.DefOffset) - if err != nil { - return listClientsReq{}, errors.Wrap(apiutil.ErrValidation, err) - } - limit, err := apiutil.ReadNumQuery[uint64](r, api.LimitKey, api.DefLimit) - if err != nil { - return listClientsReq{}, errors.Wrap(apiutil.ErrValidation, err) - } - - dir, err := apiutil.ReadStringQuery(r, api.DirKey, api.DefDir) - if err != nil { - return listClientsReq{}, errors.Wrap(apiutil.ErrValidation, err) - } - - order, err := apiutil.ReadStringQuery(r, api.OrderKey, api.DefOrder) - if err != nil { - return listClientsReq{}, errors.Wrap(apiutil.ErrValidation, err) - } - - allActions, err := apiutil.ReadStringQuery(r, api.ActionsKey, "") - if err != nil { - return listClientsReq{}, errors.Wrap(apiutil.ErrValidation, err) - } - - actions := []string{} - - allActions = strings.TrimSpace(allActions) - if allActions != "" { - actions = strings.Split(allActions, ",") - } - roleID, err := apiutil.ReadStringQuery(r, api.RoleIDKey, "") - if err != nil { - return listClientsReq{}, errors.Wrap(apiutil.ErrValidation, err) - } - - roleName, err := apiutil.ReadStringQuery(r, api.RoleNameKey, "") - if err != nil { - return listClientsReq{}, errors.Wrap(apiutil.ErrValidation, err) - } - - accessType, err := apiutil.ReadStringQuery(r, api.AccessTypeKey, "") - if err != nil { - return listClientsReq{}, errors.Wrap(apiutil.ErrValidation, err) - } - - userID, err := apiutil.ReadStringQuery(r, api.UserKey, "") - if err != nil { - return listClientsReq{}, errors.Wrap(apiutil.ErrValidation, err) - } - - var groupPtr *string - groupID, err := apiutil.ReadStringQuery(r, api.GroupKey, "") - if err != nil { - return listClientsReq{}, errors.Wrap(apiutil.ErrValidation, err) - } - - if r.URL.Query().Has(api.GroupKey) { - groupPtr = &groupID - } - - channelID, err := apiutil.ReadStringQuery(r, api.ChannelKey, "") - if err != nil { - return listClientsReq{}, errors.Wrap(apiutil.ErrValidation, err) - } - - connType, err := apiutil.ReadStringQuery(r, api.ConnTypeKey, "") - if err != nil { - return listClientsReq{}, errors.Wrap(apiutil.ErrValidation, err) - } - - id, err := apiutil.ReadStringQuery(r, api.IDOrder, "") - if err != nil { - return listClientsReq{}, errors.Wrap(apiutil.ErrValidation, err) - } - ot, err := apiutil.ReadBoolQuery(r, api.OnlyTotal, false) - if err != nil { - return listClientsReq{}, errors.Wrap(apiutil.ErrValidation, err) - } - - cfrom, err := apiutil.ReadStringQuery(r, "created_from", "") - if err != nil { - return listClientsReq{}, errors.Wrap(apiutil.ErrValidation, err) - } - cto, err := apiutil.ReadStringQuery(r, "created_to", "") - if err != nil { - return listClientsReq{}, errors.Wrap(apiutil.ErrValidation, err) - } - - var createdFrom, createdTo time.Time - if cfrom != "" { - if createdFrom, err = time.Parse(time.RFC3339, cfrom); err != nil { - return listClientsReq{}, errors.Wrap(apiutil.ErrInvalidQueryParams, err) - } - } - if cto != "" { - if createdTo, err = time.Parse(time.RFC3339, cto); err != nil { - return listClientsReq{}, errors.Wrap(apiutil.ErrInvalidQueryParams, err) - } - } - - req := listClientsReq{ - Page: clients.Page{ - Name: name, - Tags: tq, - Status: status, - Metadata: meta, - RoleName: roleName, - RoleID: roleID, - Actions: actions, - AccessType: accessType, - Order: order, - Dir: dir, - Offset: offset, - Limit: limit, - Group: groupPtr, - Channel: channelID, - ConnectionType: connType, - ID: id, - OnlyTotal: ot, - CreatedFrom: createdFrom, - CreatedTo: createdTo, - }, - userID: userID, - } - return req, nil -} - -func decodeUpdateClient(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - - req := updateClientReq{ - id: chi.URLParam(r, clientID), - } - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - - return req, nil -} - -func decodeUpdateClientTags(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - - req := updateClientTagsReq{ - id: chi.URLParam(r, clientID), - } - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - - return req, nil -} - -func decodeUpdateClientCredentials(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - - req := updateClientCredentialsReq{ - id: chi.URLParam(r, clientID), - } - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - - return req, nil -} - -func decodeCreateClientReq(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - - var c clients.Client - if err := json.NewDecoder(r.Body).Decode(&c); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - req := createClientReq{ - client: c, - } - - return req, nil -} - -func decodeCreateClientsReq(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - - c := createClientsReq{} - if err := json.NewDecoder(r.Body).Decode(&c.Clients); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - - return c, nil -} - -func decodeChangeClientStatus(_ context.Context, r *http.Request) (any, error) { - req := changeClientStatusReq{ - id: chi.URLParam(r, clientID), - } - - return req, nil -} - -func decodeSetClientParentGroupStatus(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - - req := setClientParentGroupReq{ - id: chi.URLParam(r, clientID), - } - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - return req, nil -} - -func decodeRemoveClientParentGroupStatus(_ context.Context, r *http.Request) (any, error) { - req := removeClientParentGroupReq{ - id: chi.URLParam(r, clientID), - } - - return req, nil -} - -func decodeDeleteClientReq(_ context.Context, r *http.Request) (any, error) { - req := deleteClientReq{ - id: chi.URLParam(r, clientID), - } - - return req, nil -} diff --git a/clients/api/http/endpoints.go b/clients/api/http/endpoints.go deleted file mode 100644 index b913d7ec7..000000000 --- a/clients/api/http/endpoints.go +++ /dev/null @@ -1,304 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package http - -import ( - "context" - - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/clients" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/go-kit/kit/endpoint" -) - -func createClientEndpoint(svc clients.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(createClientReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - clients, _, err := svc.CreateClients(ctx, session, req.client) - if err != nil { - return nil, err - } - - return createClientRes{ - Client: clients[0], - created: true, - }, nil - } -} - -func createClientsEndpoint(svc clients.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(createClientsReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - clients, _, err := svc.CreateClients(ctx, session, req.Clients...) - if err != nil { - return nil, err - } - - res := clientsPageRes{ - clientsPageMetaRes: clientsPageMetaRes{ - Total: uint64(len(clients)), - }, - Clients: []viewClientRes{}, - } - for _, c := range clients { - res.Clients = append(res.Clients, viewClientRes{Client: c}) - } - - return res, nil - } -} - -func viewClientEndpoint(svc clients.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(viewClientReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - c, err := svc.View(ctx, session, req.id, req.roles) - if err != nil { - return nil, err - } - - return viewClientRes{Client: c}, nil - } -} - -func listClientsEndpoint(svc clients.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(listClientsReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - var page clients.ClientsPage - var err error - switch req.userID != "" { - case true: - page, err = svc.ListUserClients(ctx, session, req.userID, req.Page) - default: - page, err = svc.ListClients(ctx, session, req.Page) - } - if err != nil { - return clientsPageRes{}, err - } - - res := clientsPageRes{ - clientsPageMetaRes: clientsPageMetaRes{ - Total: page.Total, - Offset: page.Offset, - Limit: page.Limit, - }, - Clients: []viewClientRes{}, - } - for _, c := range page.Clients { - res.Clients = append(res.Clients, viewClientRes{Client: c}) - } - - return res, nil - } -} - -func updateClientEndpoint(svc clients.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(updateClientReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - cli := clients.Client{ - ID: req.id, - Name: req.Name, - Metadata: req.Metadata, - PrivateMetadata: req.PrivateMetadata, - } - client, err := svc.Update(ctx, session, cli) - if err != nil { - return nil, err - } - - return updateClientRes{Client: client}, nil - } -} - -func updateClientTagsEndpoint(svc clients.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(updateClientTagsReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - cli := clients.Client{ - ID: req.id, - Tags: req.Tags, - } - client, err := svc.UpdateTags(ctx, session, cli) - if err != nil { - return nil, err - } - - return updateClientRes{Client: client}, nil - } -} - -func updateClientSecretEndpoint(svc clients.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(updateClientCredentialsReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - client, err := svc.UpdateSecret(ctx, session, req.id, req.Secret) - if err != nil { - return nil, err - } - - return updateClientRes{Client: client}, nil - } -} - -func enableClientEndpoint(svc clients.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(changeClientStatusReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - client, err := svc.Enable(ctx, session, req.id) - if err != nil { - return nil, err - } - - return changeClientStatusRes{Client: client}, nil - } -} - -func disableClientEndpoint(svc clients.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(changeClientStatusReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - client, err := svc.Disable(ctx, session, req.id) - if err != nil { - return nil, err - } - - return changeClientStatusRes{Client: client}, nil - } -} - -func setClientParentGroupEndpoint(svc clients.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(setClientParentGroupReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - if err := svc.SetParentGroup(ctx, session, req.ParentGroupID, req.id); err != nil { - return nil, err - } - - return setParentGroupRes{}, nil - } -} - -func removeClientParentGroupEndpoint(svc clients.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(removeClientParentGroupReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - if err := svc.RemoveParentGroup(ctx, session, req.id); err != nil { - return nil, err - } - - return removeParentGroupRes{}, nil - } -} - -func deleteClientEndpoint(svc clients.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(deleteClientReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - if err := svc.Delete(ctx, session, req.id); err != nil { - return nil, err - } - - return deleteClientRes{}, nil - } -} diff --git a/clients/api/http/endpoints_test.go b/clients/api/http/endpoints_test.go deleted file mode 100644 index 355a79ddd..000000000 --- a/clients/api/http/endpoints_test.go +++ /dev/null @@ -1,1887 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package http_test - -import ( - "encoding/json" - "fmt" - "io" - "net/http" - "net/http/httptest" - "net/url" - "strings" - "testing" - "time" - - "github.com/0x6flab/namegenerator" - api "github.com/absmach/magistrala/api/http" - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/clients" - clientsapi "github.com/absmach/magistrala/clients/api/http" - "github.com/absmach/magistrala/clients/mocks" - "github.com/absmach/magistrala/internal/testsutil" - mglog "github.com/absmach/magistrala/logger" - smqauthn "github.com/absmach/magistrala/pkg/authn" - authnmocks "github.com/absmach/magistrala/pkg/authn/mocks" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/pkg/roles" - "github.com/absmach/magistrala/pkg/uuid" - "github.com/go-chi/chi/v5" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" -) - -const contentType = "application/json" - -var ( - secret = "strongsecret" - validMetadata = clients.Metadata{"role": "client"} - ID = testsutil.GenerateUUID(&testing.T{}) - client = clients.Client{ - ID: ID, - Name: "clientname", - Tags: []string{"tag1", "tag2"}, - Credentials: clients.Credentials{Identity: "clientidentity", Secret: secret}, - PrivateMetadata: validMetadata, - Metadata: validMetadata, - Status: clients.EnabledStatus, - } - validToken = "token" - inValidToken = "invalid" - inValid = "invalid" - validID = testsutil.GenerateUUID(&testing.T{}) - domainID = testsutil.GenerateUUID(&testing.T{}) - namesgen = namegenerator.NewGenerator() - validTimeStamp = time.Date(2024, 1, 1, 0, 0, 0, 0, time.UTC) -) - -type testRequest struct { - client *http.Client - method string - url string - contentType string - token string - body io.Reader -} - -func (tr testRequest) make() (*http.Response, error) { - req, err := http.NewRequest(tr.method, tr.url, tr.body) - if err != nil { - return nil, err - } - - if tr.token != "" { - req.Header.Set("Authorization", apiutil.BearerPrefix+tr.token) - } - - if tr.contentType != "" { - req.Header.Set("Content-Type", tr.contentType) - } - - req.Header.Set("Referer", "http://localhost") - - return tr.client.Do(req) -} - -func toJSON(data any) string { - jsonData, err := json.Marshal(data) - if err != nil { - return "" - } - return string(jsonData) -} - -func newClientsServer() (*httptest.Server, *mocks.Service, *authnmocks.Authentication) { - svc := new(mocks.Service) - authn := new(authnmocks.Authentication) - - logger := mglog.NewMock() - mux := chi.NewRouter() - idp := uuid.NewMock() - am := smqauthn.NewAuthNMiddleware(authn, smqauthn.WithAllowUnverifiedUser(true)) - clientsapi.MakeHandler(svc, am, mux, logger, "", idp) - - return httptest.NewServer(mux), svc, authn -} - -func TestCreateClient(t *testing.T) { - ts, svc, authn := newClientsServer() - defer ts.Close() - - cases := []struct { - desc string - client clients.Client - domainID string - token string - contentType string - status int - authnRes smqauthn.Session - authnErr error - err error - }{ - { - desc: "register a new client with a valid token", - client: client, - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID}, - contentType: contentType, - status: http.StatusCreated, - err: nil, - }, - { - desc: "register an existing client", - client: client, - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID}, - contentType: contentType, - status: http.StatusBadRequest, - err: svcerr.ErrConflict, - }, - { - desc: "register a new client with an empty token", - client: client, - domainID: domainID, - token: "", - contentType: contentType, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: apiutil.ErrBearerToken, - }, - { - desc: "register a client with an invalid ID", - client: clients.Client{ - ID: inValid, - Credentials: clients.Credentials{ - Identity: "user@example.com", - Secret: "12345678", - }, - }, - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID}, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrInvalidIDFormat, - }, - { - desc: "register a client that can't be marshalled", - client: clients.Client{ - Credentials: clients.Credentials{ - Identity: "user@example.com", - Secret: "12345678", - }, - Metadata: map[string]any{ - "test": make(chan int), - }, - }, - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID}, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrMalformedRequestBody, - }, - { - desc: "register client with invalid status", - client: clients.Client{ - ID: testsutil.GenerateUUID(t), - Credentials: clients.Credentials{ - Identity: "newclientwithinvalidstatus@example.com", - Secret: secret, - }, - Status: clients.AllStatus, - }, - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID}, - contentType: contentType, - status: http.StatusBadRequest, - err: svcerr.ErrInvalidStatus, - }, - { - desc: "create client with invalid contentype", - client: clients.Client{ - ID: testsutil.GenerateUUID(t), - Credentials: clients.Credentials{ - Identity: "example@example.com", - Secret: secret, - }, - }, - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID}, - contentType: "application/xml", - status: http.StatusUnsupportedMediaType, - err: apiutil.ErrUnsupportedContentType, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - data := toJSON(tc.client) - req := testRequest{ - client: ts.Client(), - method: http.MethodPost, - url: fmt.Sprintf("%s/%s/clients/", ts.URL, tc.domainID), - contentType: tc.contentType, - token: tc.token, - body: strings.NewReader(data), - } - - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.authnRes, tc.authnErr) - svcCall := svc.On("CreateClients", mock.Anything, tc.authnRes, []clients.Client{tc.client}).Return([]clients.Client{tc.client}, []roles.RoleProvision{}, tc.err) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - var errRes respBody - err = json.NewDecoder(res.Body).Decode(&errRes) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if errRes.Err != "" || errRes.Message != "" { - err = errors.Wrap(errors.New(errRes.Err), errors.New(errRes.Message)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestCreateClients(t *testing.T) { - ts, svc, authn := newClientsServer() - defer ts.Close() - - num := 3 - var items []clients.Client - for i := 0; i < num; i++ { - client := clients.Client{ - ID: testsutil.GenerateUUID(t), - Name: namesgen.Generate(), - Credentials: clients.Credentials{ - Identity: fmt.Sprintf("%s@example.com", namesgen.Generate()), - Secret: secret, - }, - PrivateMetadata: clients.Metadata{}, - Metadata: clients.Metadata{}, - Status: clients.EnabledStatus, - } - items = append(items, client) - } - - cases := []struct { - desc string - client []clients.Client - domainID string - token string - contentType string - status int - authnRes smqauthn.Session - authnErr error - err error - len int - }{ - { - desc: "create clients with valid token", - client: items, - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID}, - contentType: contentType, - status: http.StatusOK, - err: nil, - len: 3, - }, - { - desc: "create clients with invalid token", - client: items, - token: inValidToken, - contentType: contentType, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - len: 0, - }, - { - desc: "create clients with empty token", - client: items, - token: "", - contentType: contentType, - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - len: 0, - }, - { - desc: "create clients with empty request", - client: []clients.Client{}, - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID}, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrEmptyList, - len: 0, - }, - { - desc: "create clients with invalid IDs", - client: []clients.Client{ - { - ID: inValid, - }, - { - ID: validID, - }, - { - ID: validID, - }, - }, - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID}, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrInvalidIDFormat, - }, - { - desc: "create clients with invalid contentype", - client: []clients.Client{ - { - ID: testsutil.GenerateUUID(t), - }, - }, - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID}, - contentType: "application/xml", - status: http.StatusUnsupportedMediaType, - err: apiutil.ErrUnsupportedContentType, - }, - { - desc: "create a client that can't be marshalled", - client: []clients.Client{ - { - ID: testsutil.GenerateUUID(t), - Credentials: clients.Credentials{ - Identity: "user@example.com", - Secret: "12345678", - }, - Metadata: map[string]any{ - "test": make(chan int), - }, - }, - }, - contentType: contentType, - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID}, - status: http.StatusBadRequest, - err: apiutil.ErrMalformedRequestBody, - }, - { - desc: "create clients with service error", - client: items, - contentType: contentType, - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID}, - status: http.StatusUnprocessableEntity, - err: svcerr.ErrCreateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - data := toJSON(tc.client) - req := testRequest{ - client: ts.Client(), - method: http.MethodPost, - url: fmt.Sprintf("%s/%s/clients/bulk", ts.URL, domainID), - contentType: tc.contentType, - token: tc.token, - body: strings.NewReader(data), - } - - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.authnRes, tc.authnErr) - svcCall := svc.On("CreateClients", mock.Anything, tc.authnRes, mock.Anything, mock.Anything, mock.Anything).Return(tc.client, []roles.RoleProvision{}, tc.err) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - - var bodyRes respBody - err = json.NewDecoder(res.Body).Decode(&bodyRes) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if bodyRes.Err != "" || bodyRes.Message != "" { - err = errors.Wrap(errors.New(bodyRes.Err), errors.New(bodyRes.Message)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.len, bodyRes.Total, fmt.Sprintf("%s: expected %d got %d", tc.desc, tc.len, bodyRes.Total)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestListClients(t *testing.T) { - ts, svc, authn := newClientsServer() - defer ts.Close() - - cases := []struct { - desc string - query string - domainID string - token string - pageMeta clients.Page - listClientsResponse clients.ClientsPage - status int - authnRes smqauthn.Session - authnErr error - err error - }{ - { - desc: "list clients as admin with valid token", - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: domainID + "_" + validID, SuperAdmin: false}, - status: http.StatusOK, - pageMeta: clients.Page{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Actions: []string{}, - }, - listClientsResponse: clients.ClientsPage{ - Page: clients.Page{ - Total: 1, - }, - Clients: []clients.Client{client}, - }, - err: nil, - }, - { - desc: "list clients as non admin with valid token", - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: domainID + "_" + validID, SuperAdmin: false}, - status: http.StatusOK, - pageMeta: clients.Page{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Actions: []string{}, - }, - listClientsResponse: clients.ClientsPage{ - Page: clients.Page{ - Total: 1, - }, - Clients: []clients.Client{client}, - }, - err: nil, - }, - { - desc: "list clients with empty token", - domainID: domainID, - token: "", - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "list clients with invalid token", - domainID: domainID, - token: inValidToken, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "list clients with offset", - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: domainID + "_" + validID, SuperAdmin: false}, - pageMeta: clients.Page{ - Offset: 1, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Actions: []string{}, - }, - listClientsResponse: clients.ClientsPage{ - Page: clients.Page{ - Offset: 1, - Total: 1, - }, - Clients: []clients.Client{client}, - }, - query: "offset=1", - status: http.StatusOK, - err: nil, - }, - { - desc: "list clients with invalid offset", - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: domainID + "_" + validID, SuperAdmin: false}, - query: "offset=invalid", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list clients with limit", - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: domainID + "_" + validID, SuperAdmin: false}, - pageMeta: clients.Page{ - Offset: 0, - Limit: 1, - Order: api.DefOrder, - Dir: api.DefDir, - Actions: []string{}, - }, - listClientsResponse: clients.ClientsPage{ - Page: clients.Page{ - Limit: 1, - Total: 1, - }, - Clients: []clients.Client{client}, - }, - query: "limit=1", - status: http.StatusOK, - err: nil, - }, - { - desc: "list clients with invalid limit", - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: domainID + "_" + validID, SuperAdmin: false}, - query: "limit=invalid", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list clients with limit greater than max", - token: validToken, - domainID: domainID, - authnRes: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: domainID + "_" + validID, SuperAdmin: false}, - query: fmt.Sprintf("limit=%d", api.MaxLimitSize+1), - status: http.StatusBadRequest, - err: apiutil.ErrLimitSize, - }, - { - desc: "list clients with name", - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: domainID + "_" + validID, SuperAdmin: false}, - pageMeta: clients.Page{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Actions: []string{}, - Name: "clientname", - }, - listClientsResponse: clients.ClientsPage{ - Page: clients.Page{ - Total: 1, - }, - Clients: []clients.Client{client}, - }, - query: "name=clientname", - status: http.StatusOK, - err: nil, - }, - { - desc: "list clients with duplicate name", - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: domainID + "_" + validID, SuperAdmin: false}, - query: "name=1&name=2", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list clients with status", - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: domainID + "_" + validID, SuperAdmin: false}, - pageMeta: clients.Page{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Actions: []string{}, - Status: clients.EnabledStatus, - }, - listClientsResponse: clients.ClientsPage{ - Page: clients.Page{ - Total: 1, - }, - Clients: []clients.Client{client}, - }, - query: "status=enabled", - status: http.StatusOK, - err: nil, - }, - { - desc: "list clients with invalid status", - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: domainID + "_" + validID, SuperAdmin: false}, - query: "status=invalid", - status: http.StatusBadRequest, - err: svcerr.ErrInvalidStatus, - }, - { - desc: "list clients with duplicate status", - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: domainID + "_" + validID, SuperAdmin: false}, - query: "status=enabled&status=disabled", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list clients with single tag", - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: domainID + "_" + validID, SuperAdmin: false}, - pageMeta: clients.Page{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Actions: []string{}, - Tags: clients.TagsQuery{Elements: []string{"tag1"}, Operator: clients.OrOp}, - }, - listClientsResponse: clients.ClientsPage{ - Page: clients.Page{ - Total: 1, - }, - Clients: []clients.Client{client}, - }, - query: "tags=tag1", - status: http.StatusOK, - err: nil, - }, - { - desc: "list clients with multiple tags and OR operator", - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: domainID + "_" + validID, SuperAdmin: false}, - pageMeta: clients.Page{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Actions: []string{}, - Tags: clients.TagsQuery{Elements: []string{"tag1", "tag2", "tag3"}, Operator: clients.OrOp}, - }, - listClientsResponse: clients.ClientsPage{ - Page: clients.Page{ - Total: 1, - }, - Clients: []clients.Client{client}, - }, - query: "tags=tag1,tag2,tag3", - status: http.StatusOK, - err: nil, - }, - { - desc: "list clients with multiple tags and AND operator", - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: domainID + "_" + validID, SuperAdmin: false}, - pageMeta: clients.Page{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Actions: []string{}, - Tags: clients.TagsQuery{Elements: []string{"tag1", "tag2", "tag3"}, Operator: clients.AndOp}, - }, - listClientsResponse: clients.ClientsPage{ - Page: clients.Page{ - Total: 1, - }, - Clients: []clients.Client{client}, - }, - query: "tags=tag1%2Btag2%2Btag3", - status: http.StatusOK, - err: nil, - }, - { - desc: "list clients with duplicate tags", - domainID: domainID, - authnRes: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: domainID + "_" + validID, SuperAdmin: false}, - token: validToken, - query: "tags=tag1&tags=tag2", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list clients with metadata", - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: domainID + "_" + validID, SuperAdmin: false}, - pageMeta: clients.Page{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Actions: []string{}, - Metadata: clients.Metadata{"domain": "example.com"}, - }, - listClientsResponse: clients.ClientsPage{ - Page: clients.Page{ - Total: 1, - }, - Clients: []clients.Client{client}, - }, - query: fmt.Sprintf("metadata=%s", url.PathEscape(`{"domain": "example.com"}`)), - status: http.StatusOK, - err: nil, - }, - { - desc: "list clients with invalid metadata", - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: domainID + "_" + validID, SuperAdmin: false}, - query: "metadata=invalid", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list clients with duplicate metadata", - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: domainID + "_" + validID, SuperAdmin: false}, - query: fmt.Sprintf("metadata=%s&metadata=%s", url.PathEscape(`{"domain": "example.com"}`), url.PathEscape(`{"domain": "example.com"}`)), - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list clients with created_from parameter", - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: domainID + "_" + validID, SuperAdmin: false}, - pageMeta: clients.Page{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Actions: []string{}, - CreatedFrom: validTimeStamp, - }, - listClientsResponse: clients.ClientsPage{ - Page: clients.Page{ - Total: 1, - }, - Clients: []clients.Client{client}, - }, - query: "created_from=2024-01-01T00:00:00Z", - status: http.StatusOK, - err: nil, - }, - { - desc: "list clients with created_to parameter", - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: domainID + "_" + validID, SuperAdmin: false}, - pageMeta: clients.Page{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Actions: []string{}, - CreatedTo: validTimeStamp, - }, - listClientsResponse: clients.ClientsPage{ - Page: clients.Page{ - Total: 1, - }, - Clients: []clients.Client{client}, - }, - query: "created_to=2024-01-01T00:00:00Z", - status: http.StatusOK, - err: nil, - }, - { - desc: "list clients with both created_from and created_to parameters", - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: domainID + "_" + validID, SuperAdmin: false}, - pageMeta: clients.Page{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Actions: []string{}, - CreatedFrom: validTimeStamp, - CreatedTo: validTimeStamp, - }, - listClientsResponse: clients.ClientsPage{ - Page: clients.Page{ - Total: 1, - }, - Clients: []clients.Client{client}, - }, - query: "created_from=2024-01-01T00:00:00Z&created_to=2024-01-01T00:00:00Z", - status: http.StatusOK, - err: nil, - }, - { - desc: "list clients with invalid created_from", - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: domainID + "_" + validID, SuperAdmin: false}, - query: "created_from=invalid-timestamp", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list clients with duplicate created_from", - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: domainID + "_" + validID, SuperAdmin: false}, - query: "created_from=2024-01-01T00:00:00Z&created_from=2024-01-01T00:00:00Z", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list clients with invalid created_to", - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: domainID + "_" + validID, SuperAdmin: false}, - query: "created_to=invalid-timestamp", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list clients with duplicate created_to", - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{UserID: validID, DomainID: domainID, DomainUserID: domainID + "_" + validID, SuperAdmin: false}, - query: "created_to=2024-12-31T23:59:59Z&created_to=2024-12-31T23:59:59Z", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - client: ts.Client(), - method: http.MethodGet, - url: ts.URL + "/" + tc.domainID + "/clients?" + tc.query, - contentType: contentType, - token: tc.token, - } - - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.authnRes, tc.authnErr) - svcCall := svc.On("ListClients", mock.Anything, tc.authnRes, tc.pageMeta).Return(tc.listClientsResponse, tc.err) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - - var bodyRes respBody - err = json.NewDecoder(res.Body).Decode(&bodyRes) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if bodyRes.Err != "" || bodyRes.Message != "" { - err = errors.Wrap(errors.New(bodyRes.Err), errors.New(bodyRes.Message)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestViewClient(t *testing.T) { - ts, svc, authn := newClientsServer() - defer ts.Close() - - cases := []struct { - desc string - domainID string - token string - id string - status int - authnRes smqauthn.Session - authnErr error - err error - }{ - { - desc: "view client with valid token", - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID}, - id: client.ID, - status: http.StatusOK, - - err: nil, - }, - { - desc: "view client with invalid token", - domainID: domainID, - token: inValidToken, - id: client.ID, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "view client with empty token", - domainID: domainID, - token: "", - id: client.ID, - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "view client with invalid id", - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID}, - id: inValid, - status: http.StatusForbidden, - - err: svcerr.ErrAuthorization, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - client: ts.Client(), - method: http.MethodGet, - url: fmt.Sprintf("%s/%s/clients/%s", ts.URL, tc.domainID, tc.id), - token: tc.token, - } - - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.authnRes, tc.authnErr) - svcCall := svc.On("View", mock.Anything, tc.authnRes, tc.id, false).Return(clients.Client{}, tc.err) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - var errRes respBody - err = json.NewDecoder(res.Body).Decode(&errRes) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if errRes.Err != "" || errRes.Message != "" { - err = errors.Wrap(errors.New(errRes.Err), errors.New(errRes.Message)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestUpdateClient(t *testing.T) { - ts, svc, authn := newClientsServer() - defer ts.Close() - - newName := "newname" - newTag := "newtag" - newMetadata := clients.Metadata{"newkey": "newvalue"} - - cases := []struct { - desc string - id string - data string - clientResponse clients.Client - domainID string - token string - contentType string - status int - authnRes smqauthn.Session - authnErr error - err error - }{ - { - desc: "update client with valid token", - domainID: domainID, - id: client.ID, - authnRes: smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID}, - data: fmt.Sprintf(`{"name":"%s","tags":["%s"],"metadata":%s}`, newName, newTag, toJSON(newMetadata)), - token: validToken, - contentType: contentType, - clientResponse: clients.Client{ - ID: client.ID, - Name: newName, - Tags: []string{newTag}, - Metadata: newMetadata, - }, - status: http.StatusOK, - - err: nil, - }, - { - desc: "update client with invalid token", - id: client.ID, - data: fmt.Sprintf(`{"name":"%s","tags":["%s"],"metadata":%s}`, newName, newTag, toJSON(newMetadata)), - domainID: domainID, - token: inValidToken, - contentType: contentType, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "update client with empty token", - id: client.ID, - data: fmt.Sprintf(`{"name":"%s","tags":["%s"],"metadata":%s}`, newName, newTag, toJSON(newMetadata)), - domainID: domainID, - token: "", - contentType: contentType, - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "update client with invalid contentype", - id: client.ID, - data: fmt.Sprintf(`{"name":"%s","tags":["%s"],"metadata":%s}`, newName, newTag, toJSON(newMetadata)), - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID}, - contentType: "application/xml", - status: http.StatusUnsupportedMediaType, - err: apiutil.ErrUnsupportedContentType, - }, - { - desc: "update client with malformed data", - id: client.ID, - data: fmt.Sprintf(`{"name":%s}`, "invalid"), - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID}, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrMalformedRequestBody, - }, - { - desc: "update client with empty id", - id: " ", - data: fmt.Sprintf(`{"name":"%s","tags":["%s"],"metadata":%s}`, newName, newTag, toJSON(newMetadata)), - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID}, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrMissingID, - }, - { - desc: "update client with name that is too long", - id: client.ID, - authnRes: smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID}, - data: fmt.Sprintf(`{"name":"%s","tags":["%s"],"metadata":%s}`, strings.Repeat("a", api.MaxNameSize+1), newTag, toJSON(newMetadata)), - domainID: domainID, - token: validToken, - contentType: contentType, - clientResponse: clients.Client{}, - status: http.StatusBadRequest, - err: apiutil.ErrNameSize, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - client: ts.Client(), - method: http.MethodPatch, - url: fmt.Sprintf("%s/%s/clients/%s", ts.URL, tc.domainID, tc.id), - contentType: tc.contentType, - token: tc.token, - body: strings.NewReader(tc.data), - } - - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.authnRes, tc.authnErr) - svcCall := svc.On("Update", mock.Anything, tc.authnRes, mock.Anything).Return(tc.clientResponse, tc.err) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - - var resBody respBody - err = json.NewDecoder(res.Body).Decode(&resBody) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if resBody.Err != "" || resBody.Message != "" { - err = errors.Wrap(errors.New(resBody.Err), errors.New(resBody.Message)) - } - - if err == nil { - assert.Equal(t, tc.clientResponse.ID, resBody.ID, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.clientResponse, resBody.ID)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestUpdateClientsTags(t *testing.T) { - ts, svc, authn := newClientsServer() - defer ts.Close() - - newTag := "newtag" - - cases := []struct { - desc string - id string - data string - contentType string - clientResponse clients.Client - domainID string - token string - status int - authnRes smqauthn.Session - authnErr error - err error - }{ - { - desc: "update client tags with valid token", - id: client.ID, - data: fmt.Sprintf(`{"tags":["%s"]}`, newTag), - contentType: contentType, - clientResponse: clients.Client{ - ID: client.ID, - Tags: []string{newTag}, - }, - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID}, - status: http.StatusOK, - err: nil, - }, - { - desc: "update client tags with empty token", - id: client.ID, - data: fmt.Sprintf(`{"tags":["%s"]}`, newTag), - contentType: contentType, - domainID: domainID, - token: "", - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "update client tags with invalid token", - id: client.ID, - data: fmt.Sprintf(`{"tags":["%s"]}`, newTag), - contentType: contentType, - domainID: domainID, - token: inValidToken, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "update client tags with invalid id", - id: client.ID, - data: fmt.Sprintf(`{"tags":["%s"]}`, newTag), - contentType: contentType, - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID}, - status: http.StatusForbidden, - err: svcerr.ErrAuthorization, - }, - { - desc: "update client tags with invalid contentype", - id: client.ID, - data: fmt.Sprintf(`{"tags":["%s"]}`, newTag), - contentType: "application/xml", - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID}, - status: http.StatusUnsupportedMediaType, - err: apiutil.ErrUnsupportedContentType, - }, - { - desc: "update clients tags with empty id", - id: "", - data: fmt.Sprintf(`{"tags":["%s"]}`, newTag), - contentType: contentType, - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID}, - status: http.StatusBadRequest, - err: apiutil.ErrMissingID, - }, - { - desc: "update clients with malfomed data", - id: client.ID, - data: fmt.Sprintf(`{"tags":[%s]}`, newTag), - contentType: contentType, - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID}, - status: http.StatusBadRequest, - err: apiutil.ErrMalformedRequestBody, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - client: ts.Client(), - method: http.MethodPatch, - url: fmt.Sprintf("%s/%s/clients/%s/tags", ts.URL, tc.domainID, tc.id), - contentType: tc.contentType, - token: tc.token, - body: strings.NewReader(tc.data), - } - - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.authnRes, tc.authnErr) - svcCall := svc.On("UpdateTags", mock.Anything, tc.authnRes, mock.Anything).Return(tc.clientResponse, tc.err) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - var resBody respBody - err = json.NewDecoder(res.Body).Decode(&resBody) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if resBody.Err != "" || resBody.Message != "" { - err = errors.Wrap(errors.New(resBody.Err), errors.New(resBody.Message)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestUpdateClientSecret(t *testing.T) { - ts, svc, authn := newClientsServer() - defer ts.Close() - - cases := []struct { - desc string - data string - client clients.Client - contentType string - domainID string - token string - status int - authnRes smqauthn.Session - authnErr error - err error - }{ - { - desc: "update client secret with valid token", - data: fmt.Sprintf(`{"secret": "%s"}`, "strongersecret"), - client: clients.Client{ - ID: client.ID, - Credentials: clients.Credentials{ - Identity: "clientname", - Secret: "strongersecret", - }, - }, - contentType: contentType, - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID}, - status: http.StatusOK, - err: nil, - }, - { - desc: "update client secret with empty token", - data: fmt.Sprintf(`{"secret": "%s"}`, "strongersecret"), - client: clients.Client{ - ID: client.ID, - Credentials: clients.Credentials{ - Identity: "clientname", - Secret: "strongersecret", - }, - }, - contentType: contentType, - domainID: domainID, - token: "", - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "update client secret with invalid token", - data: fmt.Sprintf(`{"secret": "%s"}`, "strongersecret"), - client: clients.Client{ - ID: client.ID, - Credentials: clients.Credentials{ - Identity: "clientname", - Secret: "strongersecret", - }, - }, - contentType: contentType, - domainID: domainID, - token: inValid, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "update client secret with empty id", - data: fmt.Sprintf(`{"secret": "%s"}`, "strongersecret"), - client: clients.Client{ - ID: "", - Credentials: clients.Credentials{ - Identity: "clientname", - Secret: "strongersecret", - }, - }, - contentType: contentType, - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID}, - status: http.StatusBadRequest, - err: apiutil.ErrMissingID, - }, - { - desc: "update client secret with empty secret", - data: fmt.Sprintf(`{"secret": "%s"}`, ""), - client: clients.Client{ - ID: client.ID, - Credentials: clients.Credentials{ - Identity: "clientname", - Secret: "", - }, - }, - contentType: contentType, - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID}, - status: http.StatusBadRequest, - err: apiutil.ErrMissingSecret, - }, - { - desc: "update client secret with invalid contentype", - data: fmt.Sprintf(`{"secret": "%s"}`, ""), - client: clients.Client{ - ID: client.ID, - Credentials: clients.Credentials{ - Identity: "clientname", - Secret: "", - }, - }, - contentType: "application/xml", - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID}, - status: http.StatusUnsupportedMediaType, - err: apiutil.ErrUnsupportedContentType, - }, - { - desc: "update client secret with malformed data", - data: fmt.Sprintf(`{"secret": %s}`, "invalid"), - client: clients.Client{ - ID: client.ID, - Credentials: clients.Credentials{ - Identity: "clientname", - Secret: "", - }, - }, - contentType: contentType, - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID}, - status: http.StatusBadRequest, - err: apiutil.ErrMalformedRequestBody, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - client: ts.Client(), - method: http.MethodPatch, - url: fmt.Sprintf("%s/%s/clients/%s/secret", ts.URL, tc.domainID, tc.client.ID), - contentType: tc.contentType, - token: tc.token, - body: strings.NewReader(tc.data), - } - - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.authnRes, tc.authnErr) - svcCall := svc.On("UpdateSecret", mock.Anything, tc.authnRes, tc.client.ID, mock.Anything).Return(tc.client, tc.err) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - var resBody respBody - err = json.NewDecoder(res.Body).Decode(&resBody) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if resBody.Err != "" || resBody.Message != "" { - err = errors.Wrap(errors.New(resBody.Err), errors.New(resBody.Message)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestEnableClient(t *testing.T) { - ts, svc, authn := newClientsServer() - defer ts.Close() - - cases := []struct { - desc string - client clients.Client - response clients.Client - domainID string - token string - status int - authnRes smqauthn.Session - authnErr error - err error - }{ - { - desc: "enable client with valid token", - client: client, - response: clients.Client{ - ID: client.ID, - Status: clients.EnabledStatus, - }, - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID}, - status: http.StatusOK, - err: nil, - }, - { - desc: "enable client with invalid token", - client: client, - domainID: domainID, - token: inValidToken, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "enable client with empty id", - client: clients.Client{ - ID: "", - }, - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID}, - status: http.StatusBadRequest, - err: apiutil.ErrMissingID, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - data := toJSON(tc.client) - req := testRequest{ - client: ts.Client(), - method: http.MethodPost, - url: fmt.Sprintf("%s/%s/clients/%s/enable", ts.URL, tc.domainID, tc.client.ID), - contentType: contentType, - token: tc.token, - body: strings.NewReader(data), - } - - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.authnRes, tc.authnErr) - svcCall := svc.On("Enable", mock.Anything, tc.authnRes, tc.client.ID).Return(tc.response, tc.err) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - var resBody respBody - err = json.NewDecoder(res.Body).Decode(&resBody) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if resBody.Err != "" || resBody.Message != "" { - err = errors.Wrap(errors.New(resBody.Err), errors.New(resBody.Message)) - } - if err == nil { - assert.Equal(t, tc.response.Status, resBody.Status, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.response.Status, resBody.Status)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestDisableClient(t *testing.T) { - ts, svc, authn := newClientsServer() - defer ts.Close() - - cases := []struct { - desc string - client clients.Client - response clients.Client - domainID string - token string - status int - authnRes smqauthn.Session - authnErr error - err error - }{ - { - desc: "disable client with valid token", - client: client, - response: clients.Client{ - ID: client.ID, - Status: clients.DisabledStatus, - }, - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID}, - status: http.StatusOK, - err: nil, - }, - { - desc: "disable client with invalid token", - client: client, - domainID: domainID, - token: inValidToken, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "disable client with empty id", - client: clients.Client{ - ID: "", - }, - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID}, - status: http.StatusBadRequest, - err: apiutil.ErrMissingID, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - data := toJSON(tc.client) - req := testRequest{ - client: ts.Client(), - method: http.MethodPost, - url: fmt.Sprintf("%s/%s/clients/%s/disable", ts.URL, tc.domainID, tc.client.ID), - contentType: contentType, - token: tc.token, - body: strings.NewReader(data), - } - - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.authnRes, tc.authnErr) - svcCall := svc.On("Disable", mock.Anything, tc.authnRes, tc.client.ID).Return(tc.response, tc.err) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - var resBody respBody - err = json.NewDecoder(res.Body).Decode(&resBody) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if resBody.Err != "" || resBody.Message != "" { - err = errors.Wrap(errors.New(resBody.Err), errors.New(resBody.Message)) - } - if err == nil { - assert.Equal(t, tc.response.Status, resBody.Status, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.response.Status, resBody.Status)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestDeleteClient(t *testing.T) { - ts, svc, authn := newClientsServer() - defer ts.Close() - - cases := []struct { - desc string - id string - domainID string - token string - status int - authnRes smqauthn.Session - authnErr error - err error - }{ - { - desc: "delete client with valid token", - id: client.ID, - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID}, - status: http.StatusNoContent, - err: nil, - }, - { - desc: "delete client with invalid token", - id: client.ID, - domainID: domainID, - token: inValidToken, - authnRes: smqauthn.Session{}, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "delete client with empty token", - id: client.ID, - domainID: domainID, - token: "", - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "delete client with empty id", - id: " ", - domainID: domainID, - token: validToken, - authnRes: smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID}, - status: http.StatusBadRequest, - - err: apiutil.ErrMissingID, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - client: ts.Client(), - method: http.MethodDelete, - url: fmt.Sprintf("%s/%s/clients/%s", ts.URL, tc.domainID, tc.id), - token: tc.token, - } - - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.authnRes, tc.authnErr) - svcCall := svc.On("Delete", mock.Anything, tc.authnRes, tc.id).Return(tc.err) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestSetClientParentGroupEndpoint(t *testing.T) { - gs, svc, authn := newClientsServer() - defer gs.Close() - - cases := []struct { - desc string - token string - id string - domainID string - data string - contentType string - session smqauthn.Session - svcErr error - resp clients.Client - status int - authnErr error - err error - }{ - { - desc: "set client parent group successfully", - token: validToken, - domainID: validID, - id: validID, - data: fmt.Sprintf(`{"parent_group_id":"%s"}`, validID), - contentType: contentType, - status: http.StatusOK, - err: nil, - }, - { - desc: "set client parent group with invalid token", - token: inValidToken, - domainID: validID, - id: validID, - data: fmt.Sprintf(`{"parent_group_id":"%s"}`, validID), - contentType: contentType, - authnErr: svcerr.ErrAuthentication, - status: http.StatusUnauthorized, - err: svcerr.ErrAuthentication, - }, - { - desc: "set client parent group with empty token", - token: "", - domainID: validID, - id: validID, - data: fmt.Sprintf(`{"parent_group_id":"%s"}`, validID), - contentType: contentType, - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "set client parent group with empty domainID", - token: validToken, - id: validID, - data: fmt.Sprintf(`{"parent_group_id":"%s"}`, validID), - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrMissingDomainID, - }, - { - desc: "set client parent group with invalid content type", - token: validToken, - id: validID, - domainID: validID, - data: fmt.Sprintf(`{"parent_group_id":"%s"}`, validID), - contentType: "application/xml", - status: http.StatusUnsupportedMediaType, - err: apiutil.ErrUnsupportedContentType, - }, - { - desc: "set client parent group with empty id", - token: validToken, - id: "", - domainID: validID, - data: fmt.Sprintf(`{"parent_group_id":"%s"}`, validID), - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrMissingID, - }, - { - desc: "set client parent group with empty parent group id", - token: validToken, - id: validID, - domainID: validID, - data: `{"parent_group_id":""}`, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrMissingParentGroupID, - }, - { - desc: "set client parent group with malformed request", - token: validToken, - id: validID, - domainID: validID, - data: fmt.Sprintf(`{"parent_group_id":"%s"`, validID), - contentType: contentType, - status: http.StatusBadRequest, - err: errors.ErrMalformedEntity, - }, - { - desc: "set client parent group with service error", - token: validToken, - id: validID, - domainID: validID, - data: fmt.Sprintf(`{"parent_group_id":"%s"}`, validID), - contentType: contentType, - svcErr: svcerr.ErrAuthorization, - status: http.StatusForbidden, - err: svcerr.ErrAuthorization, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - client: gs.Client(), - method: http.MethodPost, - url: fmt.Sprintf("%s/%s/clients/%s/parent", gs.URL, tc.domainID, tc.id), - contentType: tc.contentType, - token: tc.token, - body: strings.NewReader(tc.data), - } - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID + "_" + validID, UserID: validID, DomainID: validID} - } - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - svcCall := svc.On("SetParentGroup", mock.Anything, tc.session, validID, tc.id).Return(tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestRemoveClientParentGroupEndpoint(t *testing.T) { - gs, svc, authn := newClientsServer() - defer gs.Close() - - cases := []struct { - desc string - token string - id string - domainID string - session smqauthn.Session - svcErr error - resp clients.Client - status int - authnErr error - err error - }{ - { - desc: "remove client parent group successfully", - token: validToken, - id: validID, - domainID: validID, - status: http.StatusNoContent, - err: nil, - }, - { - desc: "remove client parent group with invalid token", - token: inValidToken, - session: smqauthn.Session{}, - id: validID, - domainID: validID, - authnErr: svcerr.ErrAuthentication, - status: http.StatusUnauthorized, - err: svcerr.ErrAuthentication, - }, - { - desc: "remove client parent group with empty token", - token: "", - id: validID, - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "remove client parent group with empty domainID", - token: validToken, - id: validID, - status: http.StatusBadRequest, - err: apiutil.ErrMissingDomainID, - }, - { - desc: "remove client parent group with empty id", - token: validToken, - id: "", - domainID: validID, - status: http.StatusBadRequest, - err: apiutil.ErrMissingID, - }, - { - desc: "remove client parent group with service error", - token: validToken, - id: validID, - domainID: validID, - svcErr: svcerr.ErrAuthorization, - status: http.StatusForbidden, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - client: gs.Client(), - method: http.MethodDelete, - url: fmt.Sprintf("%s/%s/clients/%s/parent", gs.URL, tc.domainID, tc.id), - token: tc.token, - } - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID + "_" + validID, UserID: validID, DomainID: validID} - } - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - svcCall := svc.On("RemoveParentGroup", mock.Anything, tc.session, tc.id).Return(tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -type respBody struct { - Err string `json:"error"` - Message string `json:"message"` - Total int `json:"total"` - Permissions []string `json:"permissions"` - ID string `json:"id"` - Tags []string `json:"tags"` - Status clients.Status `json:"status"` -} diff --git a/clients/api/http/requests.go b/clients/api/http/requests.go deleted file mode 100644 index 2457078db..000000000 --- a/clients/api/http/requests.go +++ /dev/null @@ -1,212 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package http - -import ( - api "github.com/absmach/magistrala/api/http" - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/clients" -) - -type createClientReq struct { - client clients.Client -} - -func (req createClientReq) validate() error { - if len(req.client.Name) > api.MaxNameSize { - return apiutil.ErrNameSize - } - if req.client.ID != "" { - return api.ValidateUUID(req.client.ID) - } - - return nil -} - -type createClientsReq struct { - Clients []clients.Client -} - -func (req createClientsReq) validate() error { - if len(req.Clients) == 0 { - return apiutil.ErrEmptyList - } - for _, c := range req.Clients { - if c.ID != "" { - if err := api.ValidateUUID(c.ID); err != nil { - return err - } - } - if len(c.Name) > api.MaxNameSize { - return apiutil.ErrNameSize - } - } - - return nil -} - -type viewClientReq struct { - id string - roles bool -} - -func (req viewClientReq) validate() error { - if req.id == "" { - return apiutil.ErrMissingID - } - - return nil -} - -type viewClientPermsReq struct { - id string -} - -func (req viewClientPermsReq) validate() error { - if req.id == "" { - return apiutil.ErrMissingID - } - - return nil -} - -type listClientsReq struct { - clients.Page - userID string -} - -func (req listClientsReq) validate() error { - if req.Limit > api.MaxLimitSize || req.Limit < 1 { - return apiutil.ErrLimitSize - } - - if len(req.Name) > api.MaxNameSize { - return apiutil.ErrNameSize - } - - switch req.Order { - case "", api.NameOrder, api.CreatedAtOrder, api.UpdatedAtOrder: - default: - return apiutil.ErrInvalidOrder - } - - if req.Dir != "" && (req.Dir != api.DescDir && req.Dir != api.AscDir) { - return apiutil.ErrInvalidDirection - } - - return nil -} - -type listMembersReq struct { - clients.Page - groupID string -} - -func (req listMembersReq) validate() error { - if req.groupID == "" { - return apiutil.ErrMissingID - } - - return nil -} - -type updateClientReq struct { - id string - Name string `json:"name,omitempty"` - Metadata map[string]any `json:"metadata,omitempty"` - PrivateMetadata map[string]any `json:"private_metadata,omitempty"` - Tags []string `json:"tags,omitempty"` -} - -func (req updateClientReq) validate() error { - if req.id == "" { - return apiutil.ErrMissingID - } - - if len(req.Name) > api.MaxNameSize { - return apiutil.ErrNameSize - } - - return nil -} - -type updateClientTagsReq struct { - id string - Tags []string `json:"tags,omitempty"` -} - -func (req updateClientTagsReq) validate() error { - if req.id == "" { - return apiutil.ErrMissingID - } - - return nil -} - -type updateClientCredentialsReq struct { - id string - Secret string `json:"secret,omitempty"` -} - -func (req updateClientCredentialsReq) validate() error { - if req.id == "" { - return apiutil.ErrMissingID - } - - if req.Secret == "" { - return apiutil.ErrMissingSecret - } - - return nil -} - -type changeClientStatusReq struct { - id string -} - -func (req changeClientStatusReq) validate() error { - if req.id == "" { - return apiutil.ErrMissingID - } - - return nil -} - -type setClientParentGroupReq struct { - id string - ParentGroupID string `json:"parent_group_id"` -} - -func (req setClientParentGroupReq) validate() error { - if req.id == "" { - return apiutil.ErrMissingID - } - if req.ParentGroupID == "" { - return apiutil.ErrMissingParentGroupID - } - return nil -} - -type removeClientParentGroupReq struct { - id string -} - -func (req removeClientParentGroupReq) validate() error { - if req.id == "" { - return apiutil.ErrMissingID - } - return nil -} - -type deleteClientReq struct { - id string -} - -func (req deleteClientReq) validate() error { - if req.id == "" { - return apiutil.ErrMissingID - } - - return nil -} diff --git a/clients/api/http/requests_test.go b/clients/api/http/requests_test.go deleted file mode 100644 index 162780f65..000000000 --- a/clients/api/http/requests_test.go +++ /dev/null @@ -1,425 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package http - -import ( - "strings" - "testing" - - api "github.com/absmach/magistrala/api/http" - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/clients" - "github.com/absmach/magistrala/internal/testsutil" - "github.com/stretchr/testify/assert" -) - -const ( - valid = "valid" - invalid = "invalid" -) - -var validID = testsutil.GenerateUUID(&testing.T{}) - -func TestCreateClientReqValidate(t *testing.T) { - cases := []struct { - desc string - req createClientReq - err error - }{ - { - desc: "valid request", - req: createClientReq{ - client: clients.Client{ - ID: validID, - Name: valid, - }, - }, - err: nil, - }, - { - desc: "name too long", - req: createClientReq{ - client: clients.Client{ - ID: validID, - Name: strings.Repeat("a", api.MaxNameSize+1), - }, - }, - err: apiutil.ErrNameSize, - }, - { - desc: "invalid id", - req: createClientReq{ - client: clients.Client{ - ID: invalid, - Name: valid, - }, - }, - err: apiutil.ErrInvalidIDFormat, - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := tc.req.validate() - assert.Equal(t, tc.err, err) - }) - } -} - -func TestCreateClientsReqValidate(t *testing.T) { - cases := []struct { - desc string - req createClientsReq - err error - }{ - { - desc: "valid request", - req: createClientsReq{ - Clients: []clients.Client{ - { - ID: validID, - Name: valid, - }, - }, - }, - err: nil, - }, - { - desc: "empty list", - req: createClientsReq{ - Clients: []clients.Client{}, - }, - err: apiutil.ErrEmptyList, - }, - { - desc: "name too long", - req: createClientsReq{ - Clients: []clients.Client{ - { - ID: validID, - Name: strings.Repeat("a", api.MaxNameSize+1), - }, - }, - }, - err: apiutil.ErrNameSize, - }, - { - desc: "invalid id", - req: createClientsReq{ - Clients: []clients.Client{ - { - ID: invalid, - Name: valid, - }, - }, - }, - err: apiutil.ErrInvalidIDFormat, - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := tc.req.validate() - assert.Equal(t, tc.err, err, "%s: expected %s got %s\n", tc.desc, tc.err, err) - }) - } -} - -func TestViewClientReqValidate(t *testing.T) { - cases := []struct { - desc string - req viewClientReq - err error - }{ - { - desc: "valid request", - req: viewClientReq{ - id: validID, - }, - err: nil, - }, - { - desc: "empty id", - req: viewClientReq{ - id: "", - }, - err: apiutil.ErrMissingID, - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := tc.req.validate() - assert.Equal(t, tc.err, err, "%s: expected %s got %s\n", tc.desc, tc.err, err) - }) - } -} - -func TestViewClientPermsReq(t *testing.T) { - cases := []struct { - desc string - req viewClientPermsReq - err error - }{ - { - desc: "valid request", - req: viewClientPermsReq{ - id: validID, - }, - err: nil, - }, - { - desc: "empty id", - req: viewClientPermsReq{ - id: "", - }, - err: apiutil.ErrMissingID, - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := tc.req.validate() - assert.Equal(t, tc.err, err, "%s: expected %s got %s\n", tc.desc, tc.err, err) - }) - } -} - -func TestListClientsReqValidate(t *testing.T) { - cases := []struct { - desc string - req listClientsReq - err error - }{ - { - desc: "valid request", - req: listClientsReq{ - Page: clients.Page{Limit: 10}, - }, - err: nil, - }, - { - desc: "limit too big", - req: listClientsReq{ - Page: clients.Page{Limit: api.MaxLimitSize + 1}, - }, - err: apiutil.ErrLimitSize, - }, - { - desc: "limit too small", - req: listClientsReq{ - Page: clients.Page{Limit: 0}, - }, - err: apiutil.ErrLimitSize, - }, - { - desc: "name too long", - req: listClientsReq{ - Page: clients.Page{ - Limit: 10, - Name: strings.Repeat("a", api.MaxNameSize+1), - }, - }, - err: apiutil.ErrNameSize, - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := tc.req.validate() - assert.Equal(t, tc.err, err, "%s: expected %s got %s\n", tc.desc, tc.err, err) - }) - } -} - -func TestListMembersReqValidate(t *testing.T) { - cases := []struct { - desc string - req listMembersReq - err error - }{ - { - desc: "valid request", - req: listMembersReq{ - groupID: validID, - }, - err: nil, - }, - { - desc: "empty id", - req: listMembersReq{ - groupID: "", - }, - err: apiutil.ErrMissingID, - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := tc.req.validate() - assert.Equal(t, tc.err, err, "%s: expected %s got %s\n", tc.desc, tc.err, err) - }) - } -} - -func TestUpdateClientReqValidate(t *testing.T) { - cases := []struct { - desc string - req updateClientReq - err error - }{ - { - desc: "valid request", - req: updateClientReq{ - id: validID, - Name: valid, - }, - err: nil, - }, - { - desc: "empty id", - req: updateClientReq{ - id: "", - Name: valid, - }, - err: apiutil.ErrMissingID, - }, - { - desc: "name too long", - req: updateClientReq{ - id: validID, - Name: strings.Repeat("a", api.MaxNameSize+1), - }, - err: apiutil.ErrNameSize, - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := tc.req.validate() - assert.Equal(t, tc.err, err, "%s: expected %s got %s\n", tc.desc, tc.err, err) - }) - } -} - -func TestUpdateClientTagsReqValidate(t *testing.T) { - cases := []struct { - desc string - req updateClientTagsReq - err error - }{ - { - desc: "valid request", - req: updateClientTagsReq{ - id: validID, - Tags: []string{"tag1", "tag2"}, - }, - err: nil, - }, - { - desc: "empty id", - req: updateClientTagsReq{ - id: "", - Tags: []string{"tag1", "tag2"}, - }, - err: apiutil.ErrMissingID, - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := tc.req.validate() - assert.Equal(t, tc.err, err, "%s: expected %s got %s\n", tc.desc, tc.err, err) - }) - } -} - -func TestUpdateClientCredentialsReqValidate(t *testing.T) { - cases := []struct { - desc string - req updateClientCredentialsReq - err error - }{ - { - desc: "valid request", - req: updateClientCredentialsReq{ - id: validID, - Secret: valid, - }, - err: nil, - }, - { - desc: "empty id", - req: updateClientCredentialsReq{ - id: "", - Secret: valid, - }, - err: apiutil.ErrMissingID, - }, - { - desc: "empty secret", - req: updateClientCredentialsReq{ - id: validID, - Secret: "", - }, - err: apiutil.ErrMissingSecret, - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := tc.req.validate() - assert.Equal(t, tc.err, err, "%s: expected %s got %s\n", tc.desc, tc.err, err) - }) - } -} - -func TestChangeClientStatusReqValidate(t *testing.T) { - cases := []struct { - desc string - req changeClientStatusReq - err error - }{ - { - desc: "valid request", - req: changeClientStatusReq{ - id: validID, - }, - err: nil, - }, - { - desc: "empty id", - req: changeClientStatusReq{ - id: "", - }, - err: apiutil.ErrMissingID, - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := tc.req.validate() - assert.Equal(t, tc.err, err, "%s: expected %s got %s\n", tc.desc, tc.err, err) - }) - } -} - -func TestDeleteClientReqValidate(t *testing.T) { - cases := []struct { - desc string - req deleteClientReq - err error - }{ - { - desc: "valid request", - req: deleteClientReq{ - id: validID, - }, - err: nil, - }, - { - desc: "empty id", - req: deleteClientReq{ - id: "", - }, - err: apiutil.ErrMissingID, - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := tc.req.validate() - assert.Equal(t, tc.err, err, "%s: expected %s got %s\n", tc.desc, tc.err, err) - }) - } -} diff --git a/clients/api/http/responses.go b/clients/api/http/responses.go deleted file mode 100644 index 51ff21844..000000000 --- a/clients/api/http/responses.go +++ /dev/null @@ -1,177 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package http - -import ( - "fmt" - "net/http" - - "github.com/absmach/magistrala" - "github.com/absmach/magistrala/clients" -) - -var ( - _ magistrala.Response = (*createClientRes)(nil) - _ magistrala.Response = (*viewClientRes)(nil) - _ magistrala.Response = (*viewClientPermsRes)(nil) - _ magistrala.Response = (*clientsPageRes)(nil) - _ magistrala.Response = (*changeClientStatusRes)(nil) - _ magistrala.Response = (*deleteClientRes)(nil) -) - -type clientsPageMetaRes struct { - Limit uint64 `json:"limit,omitempty"` - Offset uint64 `json:"offset,omitempty"` - Total uint64 `json:"total"` -} - -type createClientRes struct { - clients.Client - created bool -} - -func (res createClientRes) Code() int { - if res.created { - return http.StatusCreated - } - - return http.StatusOK -} - -func (res createClientRes) Headers() map[string]string { - if res.created { - return map[string]string{ - "Location": fmt.Sprintf("/clients/%s", res.ID), - } - } - - return map[string]string{} -} - -func (res createClientRes) Empty() bool { - return false -} - -type updateClientRes struct { - clients.Client -} - -func (res updateClientRes) Code() int { - return http.StatusOK -} - -func (res updateClientRes) Headers() map[string]string { - return map[string]string{} -} - -func (res updateClientRes) Empty() bool { - return false -} - -type viewClientRes struct { - clients.Client -} - -func (res viewClientRes) Code() int { - return http.StatusOK -} - -func (res viewClientRes) Headers() map[string]string { - return map[string]string{} -} - -func (res viewClientRes) Empty() bool { - return false -} - -type viewClientPermsRes struct { - Permissions []string `json:"permissions"` -} - -func (res viewClientPermsRes) Code() int { - return http.StatusOK -} - -func (res viewClientPermsRes) Headers() map[string]string { - return map[string]string{} -} - -func (res viewClientPermsRes) Empty() bool { - return false -} - -type clientsPageRes struct { - clientsPageMetaRes - Clients []viewClientRes `json:"clients,omitempty"` -} - -func (res clientsPageRes) Code() int { - return http.StatusOK -} - -func (res clientsPageRes) Headers() map[string]string { - return map[string]string{} -} - -func (res clientsPageRes) Empty() bool { - return false -} - -type changeClientStatusRes struct { - clients.Client -} - -func (res changeClientStatusRes) Code() int { - return http.StatusOK -} - -func (res changeClientStatusRes) Headers() map[string]string { - return map[string]string{} -} - -func (res changeClientStatusRes) Empty() bool { - return false -} - -type setParentGroupRes struct{} - -func (res setParentGroupRes) Code() int { - return http.StatusOK -} - -func (res setParentGroupRes) Headers() map[string]string { - return map[string]string{} -} - -func (res setParentGroupRes) Empty() bool { - return true -} - -type removeParentGroupRes struct{} - -func (res removeParentGroupRes) Code() int { - return http.StatusNoContent -} - -func (res removeParentGroupRes) Headers() map[string]string { - return map[string]string{} -} - -func (res removeParentGroupRes) Empty() bool { - return true -} - -type deleteClientRes struct{} - -func (res deleteClientRes) Code() int { - return http.StatusNoContent -} - -func (res deleteClientRes) Headers() map[string]string { - return map[string]string{} -} - -func (res deleteClientRes) Empty() bool { - return true -} diff --git a/clients/api/http/transport.go b/clients/api/http/transport.go deleted file mode 100644 index d27e680f2..000000000 --- a/clients/api/http/transport.go +++ /dev/null @@ -1,25 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package http - -import ( - "log/slog" - "net/http" - - "github.com/absmach/magistrala" - "github.com/absmach/magistrala/clients" - smqauthn "github.com/absmach/magistrala/pkg/authn" - "github.com/go-chi/chi/v5" - "github.com/prometheus/client_golang/prometheus/promhttp" -) - -// MakeHandler returns a HTTP handler for clients and Groups API endpoints. -func MakeHandler(tsvc clients.Service, authn smqauthn.AuthNMiddleware, mux *chi.Mux, logger *slog.Logger, instanceID string, idp magistrala.IDProvider) http.Handler { - mux = clientsHandler(tsvc, authn, mux, logger, idp) - - mux.Get("/health", magistrala.Health("clients", instanceID)) - mux.Handle("/metrics", promhttp.Handler()) - - return mux -} diff --git a/clients/builtinroles.go b/clients/builtinroles.go deleted file mode 100644 index 6f12d5ad9..000000000 --- a/clients/builtinroles.go +++ /dev/null @@ -1,7 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 -package clients - -import "github.com/absmach/magistrala/pkg/roles" - -const BuiltInRoleAdmin roles.BuiltInRoleName = "admin" diff --git a/clients/cache/clients.go b/clients/cache/clients.go deleted file mode 100644 index 19969fdd4..000000000 --- a/clients/cache/clients.go +++ /dev/null @@ -1,85 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package cache - -import ( - "context" - "fmt" - "time" - - "github.com/absmach/magistrala/clients" - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - "github.com/redis/go-redis/v9" -) - -const ( - keyPrefix = "client_key" - idPrefix = "client_id" -) - -var _ clients.Cache = (*clientCache)(nil) - -type clientCache struct { - client *redis.Client - keyDuration time.Duration -} - -// NewCache returns redis client cache implementation. -func NewCache(client *redis.Client, duration time.Duration) clients.Cache { - return &clientCache{ - client: client, - keyDuration: duration, - } -} - -func (tc *clientCache) Save(ctx context.Context, clientKey, clientID string) error { - if clientKey == "" || clientID == "" { - return errors.Wrap(repoerr.ErrCreateEntity, errors.New("client key or client id is empty")) - } - ckey := fmt.Sprintf("%s:%s", keyPrefix, clientKey) - if err := tc.client.Set(ctx, ckey, clientID, tc.keyDuration).Err(); err != nil { - return errors.Wrap(repoerr.ErrCreateEntity, err) - } - - tid := fmt.Sprintf("%s:%s", idPrefix, clientID) - if err := tc.client.Set(ctx, tid, clientKey, tc.keyDuration).Err(); err != nil { - return errors.Wrap(repoerr.ErrCreateEntity, err) - } - - return nil -} - -func (tc *clientCache) ID(ctx context.Context, clientKey string) (string, error) { - if clientKey == "" { - return "", repoerr.ErrNotFound - } - - ckey := fmt.Sprintf("%s:%s", keyPrefix, clientKey) - clientID, err := tc.client.Get(ctx, ckey).Result() - if err != nil { - return "", errors.Wrap(repoerr.ErrNotFound, err) - } - - return clientID, nil -} - -func (tc *clientCache) Remove(ctx context.Context, clientID string) error { - tid := fmt.Sprintf("%s:%s", idPrefix, clientID) - key, err := tc.client.Get(ctx, tid).Result() - // Redis returns Nil Reply when key does not exist. - if err == redis.Nil { - return nil - } - if err != nil { - return errors.Wrap(repoerr.ErrRemoveEntity, err) - } - - tkey := fmt.Sprintf("%s:%s", keyPrefix, key) - if err := tc.client.Del(ctx, tkey, tid).Err(); err != nil { - return errors.Wrap(repoerr.ErrRemoveEntity, err) - } - - return nil -} diff --git a/clients/cache/clients_test.go b/clients/cache/clients_test.go deleted file mode 100644 index 679c4130b..000000000 --- a/clients/cache/clients_test.go +++ /dev/null @@ -1,179 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package cache_test - -import ( - "context" - "fmt" - "strings" - "testing" - "time" - - "github.com/absmach/magistrala/clients/cache" - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - "github.com/stretchr/testify/assert" -) - -const ( - testKey = "testKey" - testID = "testID" - testKey2 = "testKey2" - testID2 = "testID2" -) - -func TestSave(t *testing.T) { - redisClient.FlushAll(context.Background()) - tscache := cache.NewCache(redisClient, 1*time.Minute) - ctx := context.Background() - - cases := []struct { - desc string - key string - id string - err error - }{ - { - desc: "Save client to cache", - key: testKey, - id: testID, - err: nil, - }, - { - desc: "Save already cached client to cache", - key: testKey, - id: testID, - err: nil, - }, - { - desc: "Save another client to cache", - key: testKey2, - id: testID2, - err: nil, - }, - { - desc: "Save client with long key ", - key: strings.Repeat("a", 513*1024*1024), - id: testID, - err: repoerr.ErrCreateEntity, - }, - { - desc: "Save client with long id ", - key: testKey, - id: strings.Repeat("a", 513*1024*1024), - err: repoerr.ErrCreateEntity, - }, - { - desc: "Save client with empty key", - key: "", - id: testID, - err: repoerr.ErrCreateEntity, - }, - { - desc: "Save client with empty id", - key: testKey, - id: "", - err: repoerr.ErrCreateEntity, - }, - { - desc: "Save client with empty key and id", - key: "", - id: "", - err: repoerr.ErrCreateEntity, - }, - } - - for _, tc := range cases { - err := tscache.Save(ctx, tc.key, tc.id) - if err == nil { - id, _ := tscache.ID(ctx, tc.key) - assert.Equal(t, tc.id, id, fmt.Sprintf("%s: expected %s got %s", tc.desc, tc.id, id)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s", tc.desc, tc.err, err)) - } -} - -func TestID(t *testing.T) { - redisClient.FlushAll(context.Background()) - tscache := cache.NewCache(redisClient, 1*time.Minute) - ctx := context.Background() - - err := tscache.Save(ctx, testKey, testID) - assert.Nil(t, err, fmt.Sprintf("Unexpected error while trying to save: %s", err)) - - cases := []struct { - desc string - key string - id string - err error - }{ - { - desc: "Get client ID from cache", - key: testKey, - id: testID, - err: nil, - }, - { - desc: "Get client ID from cache for non existing client", - key: "nonExistingKey", - id: "", - err: repoerr.ErrNotFound, - }, - { - desc: "Get client ID from cache for empty key", - key: "", - id: "", - err: repoerr.ErrNotFound, - }, - } - - for _, tc := range cases { - id, err := tscache.ID(ctx, tc.key) - if err == nil { - assert.Equal(t, tc.id, id, fmt.Sprintf("%s: expected %s got %s", tc.desc, tc.id, id)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - } -} - -func TestRemove(t *testing.T) { - redisClient.FlushAll(context.Background()) - tscache := cache.NewCache(redisClient, 1*time.Minute) - ctx := context.Background() - - err := tscache.Save(ctx, testKey, testID) - assert.Nil(t, err, fmt.Sprintf("Unexpected error while trying to save: %s", err)) - - cases := []struct { - desc string - key string - err error - }{ - { - desc: "Remove existing client from cache", - key: testID, - err: nil, - }, - { - desc: "Remove non existing client from cache", - key: testID2, - err: nil, - }, - { - desc: "Remove client with empty ID from cache", - key: "", - err: nil, - }, - { - desc: "Remove client with long id from cache", - key: strings.Repeat("a", 513*1024*1024), - err: repoerr.ErrRemoveEntity, - }, - } - - for _, tc := range cases { - err := tscache.Remove(ctx, tc.key) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - } -} diff --git a/clients/cache/doc.go b/clients/cache/doc.go deleted file mode 100644 index 62e81d59a..000000000 --- a/clients/cache/doc.go +++ /dev/null @@ -1,6 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package cache contains the domain concept definitions needed to -// support Magistrala clients cache service functionality. -package cache diff --git a/clients/cache/setup_test.go b/clients/cache/setup_test.go deleted file mode 100644 index 716f0672c..000000000 --- a/clients/cache/setup_test.go +++ /dev/null @@ -1,61 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package cache_test - -import ( - "context" - "fmt" - "log" - "os" - "testing" - - "github.com/ory/dockertest/v3" - "github.com/ory/dockertest/v3/docker" - "github.com/redis/go-redis/v9" -) - -var ( - redisClient *redis.Client - redisURL string -) - -func TestMain(m *testing.M) { - pool, err := dockertest.NewPool("") - if err != nil { - log.Fatalf("Could not connect to docker: %s", err) - } - - container, err := pool.RunWithOptions(&dockertest.RunOptions{ - Repository: "redis", - Tag: "7.2.4-alpine", - }, func(config *docker.HostConfig) { - config.AutoRemove = true - config.RestartPolicy = docker.RestartPolicy{Name: "no"} - }) - if err != nil { - log.Fatalf("Could not start container: %s", err) - } - - redisURL = fmt.Sprintf("redis://localhost:%s/0", container.GetPort("6379/tcp")) - opts, err := redis.ParseURL(redisURL) - if err != nil { - log.Fatalf("Could not parse redis URL: %s", err) - } - - if err := pool.Retry(func() error { - redisClient = redis.NewClient(opts) - - return redisClient.Ping(context.Background()).Err() - }); err != nil { - log.Fatalf("Could not connect to docker: %s", err) - } - - code := m.Run() - - if err := pool.Purge(container); err != nil { - log.Fatalf("Could not purge container: %s", err) - } - - os.Exit(code) -} diff --git a/clients/clients.go b/clients/clients.go deleted file mode 100644 index acb331615..000000000 --- a/clients/clients.go +++ /dev/null @@ -1,267 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package clients - -import ( - "context" - "strings" - "time" - - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/connections" - "github.com/absmach/magistrala/pkg/postgres" - "github.com/absmach/magistrala/pkg/roles" -) - -type Connection struct { - ClientID string - ChannelID string - DomainID string - Type connections.ConnType -} - -type ClientRepository struct { - DB postgres.Database -} - -// Repository is the interface that wraps the basic methods for -// a client repository. -type Repository interface { - // RetrieveByID retrieves client by its unique ID. - RetrieveByID(ctx context.Context, id string) (Client, error) - - // RetrieveByIDWithRoles retrieves client by its unique ID along with member roles. - RetrieveByIDWithRoles(ctx context.Context, id, memberID string) (Client, error) - - // RetrieveAll retrieves all clients. - RetrieveAll(ctx context.Context, pm Page) (ClientsPage, error) - - // RetrieveUserClients retrieve all clients of a given user id. - RetrieveUserClients(ctx context.Context, domainID, userID string, pm Page) (ClientsPage, error) - - // SearchClients retrieves clients based on search criteria. - SearchClients(ctx context.Context, pm Page) (ClientsPage, error) - - // RetrieveByIds - RetrieveByIds(ctx context.Context, ids []string) (ClientsPage, error) - - // Update updates the client name and metadata. - Update(ctx context.Context, client Client) (Client, error) - - // UpdateTags updates the client tags. - UpdateTags(ctx context.Context, client Client) (Client, error) - - // UpdateIdentity updates identity for client with given id. - UpdateIdentity(ctx context.Context, client Client) (Client, error) - - // UpdateSecret updates secret for client with given identity. - UpdateSecret(ctx context.Context, client Client) (Client, error) - - // ChangeStatus changes client status to enabled or disabled - ChangeStatus(ctx context.Context, client Client) (Client, error) - - // Delete deletes client with given id - Delete(ctx context.Context, clientIDs ...string) error - - // Save persists the client account. A non-nil error is returned to indicate - // operation failure. - Save(ctx context.Context, client ...Client) ([]Client, error) - - // RetrieveBySecret retrieves a client based on the secret (key) and domainID. - // Domain ID is required because the key is not globally unique, - // but unique on the level of Domain. - RetrieveBySecret(ctx context.Context, key, id string, prefix authn.AuthPrefix) (Client, error) - - AddConnections(ctx context.Context, conns []Connection) error - - RemoveConnections(ctx context.Context, conns []Connection) error - - ClientConnectionsCount(ctx context.Context, id string) (uint64, error) - - DoesClientHaveConnections(ctx context.Context, id string) (bool, error) - - RemoveChannelConnections(ctx context.Context, channelID string) error - - RemoveClientConnections(ctx context.Context, clientID string) error - - // SetParentGroup set parent group id to a given channel id - SetParentGroup(ctx context.Context, cli Client) error - - // RemoveParentGroup remove parent group id fr given chanel id - RemoveParentGroup(ctx context.Context, cli Client) error - - RetrieveParentGroupClients(ctx context.Context, parentGroupID string) ([]Client, error) - - UnsetParentGroupFromClient(ctx context.Context, parentGroupID string) error - - roles.Repository -} - -// Service specifies an API that must be fulfilled by the domain service -// implementation, and all of its decorators (e.g. logging & metrics). -type Service interface { - // CreateClients creates new client. In case of the failed registration, a - // non-nil error value is returned. - CreateClients(ctx context.Context, session authn.Session, client ...Client) ([]Client, []roles.RoleProvision, error) - - // View retrieves client info for a given client ID and an authorized token. - View(ctx context.Context, session authn.Session, id string, withRoles bool) (Client, error) - - // ListClients retrieves clients list for given page query. - ListClients(ctx context.Context, session authn.Session, pm Page) (ClientsPage, error) - - // ListUserClients retrieves clients list for a given user id and page query. - ListUserClients(ctx context.Context, session authn.Session, userID string, pm Page) (ClientsPage, error) - - // Update updates the client's name and metadata. - Update(ctx context.Context, session authn.Session, client Client) (Client, error) - - // UpdateTags updates the client's tags. - UpdateTags(ctx context.Context, session authn.Session, client Client) (Client, error) - - // UpdateSecret updates the client's secret - UpdateSecret(ctx context.Context, session authn.Session, id, key string) (Client, error) - - // Enable logically enableds the client identified with the provided ID - Enable(ctx context.Context, session authn.Session, id string) (Client, error) - - // Disable logically disables the client identified with the provided ID - Disable(ctx context.Context, session authn.Session, id string) (Client, error) - - // Delete deletes client with given ID. - Delete(ctx context.Context, session authn.Session, id string) error - - SetParentGroup(ctx context.Context, session authn.Session, parentGroupID string, id string) error - - RemoveParentGroup(ctx context.Context, session authn.Session, id string) error - - roles.RoleManager -} - -// Cache contains client caching interface. -type Cache interface { - // Save stores pair client secret, client id. - Save(ctx context.Context, clientSecret, clientID string) error - - // ID returns client ID for given client secret. - ID(ctx context.Context, clientSecret string) (string, error) - - // Removes client from cache. - Remove(ctx context.Context, clientID string) error -} - -// Client Struct represents a client. - -type Client struct { - ID string `json:"id"` - Name string `json:"name,omitempty"` - Tags []string `json:"tags,omitempty"` - Domain string `json:"domain_id,omitempty"` - ParentGroup string `json:"parent_group_id,omitempty"` - Credentials Credentials `json:"credentials,omitempty"` - Metadata Metadata `json:"metadata,omitempty"` - PrivateMetadata Metadata `json:"private_metadata,omitempty"` - CreatedAt time.Time `json:"created_at,omitempty"` - UpdatedAt time.Time `json:"updated_at,omitempty"` - UpdatedBy string `json:"updated_by,omitempty"` - Status Status `json:"status,omitempty"` // 1 for enabled, 0 for disabled - Identity string `json:"identity,omitempty"` - // Extended - ParentGroupPath string `json:"parent_group_path,omitempty"` - RoleID string `json:"role_id,omitempty"` - RoleName string `json:"role_name,omitempty"` - Actions []string `json:"actions,omitempty"` - AccessType string `json:"access_type,omitempty"` - AccessProviderId string `json:"access_provider_id,omitempty"` - AccessProviderRoleId string `json:"access_provider_role_id,omitempty"` - AccessProviderRoleName string `json:"access_provider_role_name,omitempty"` - AccessProviderRoleActions []string `json:"access_provider_role_actions,omitempty"` - ConnectionTypes []connections.ConnType `json:"connection_types,omitempty"` - MemberId string `json:"member_id,omitempty"` - Roles []roles.MemberRoleActions `json:"roles,omitempty"` -} - -// ClientsPage contains page related metadata as well as list. -type ClientsPage struct { - Page - Clients []Client -} - -// MembersPage contains page related metadata as well as list of members that -// belong to this page. - -type MembersPage struct { - Page - Members []Client -} - -type Operator uint8 - -const ( - OrOp Operator = iota - AndOp -) - -type TagsQuery struct { - Elements []string - Operator Operator -} - -func ToTagsQuery(s string) TagsQuery { - switch { - case strings.Contains(s, "+"): - elements := strings.Split(s, "+") - for i := range elements { - elements[i] = strings.TrimSpace(elements[i]) - } - return TagsQuery{Elements: elements, Operator: AndOp} - case strings.Contains(s, ","): - elements := strings.Split(s, ",") - for i := range elements { - elements[i] = strings.TrimSpace(elements[i]) - } - return TagsQuery{Elements: elements, Operator: OrOp} - default: - return TagsQuery{Elements: []string{s}, Operator: OrOp} - } -} - -// Page contains the page metadata that helps navigation. - -type Page struct { - Total uint64 `json:"total"` - Offset uint64 `json:"offset"` - Limit uint64 `json:"limit"` - OnlyTotal bool `json:"only_total"` - Order string `json:"order,omitempty"` - Dir string `json:"dir,omitempty"` - ID string `json:"id,omitempty"` - Name string `json:"name,omitempty"` - Metadata Metadata `json:"metadata,omitempty"` - Domain string `json:"domain,omitempty"` - Tags TagsQuery `json:"tags,omitempty"` - Status Status `json:"status,omitempty"` - Identity string `json:"identity,omitempty"` - Group *string `json:"group,omitempty"` - Channel string `json:"channel,omitempty"` - ConnectionType string `json:"connection_type,omitempty"` - RoleName string `json:"role_name,omitempty"` - RoleID string `json:"role_id,omitempty"` - Actions []string `json:"actions,omitempty"` - AccessType string `json:"access_type,omitempty"` - IDs []string `json:"-"` - CreatedFrom time.Time `json:"created_from,omitempty"` - CreatedTo time.Time `json:"created_to,omitempty"` -} - -// Metadata represents arbitrary JSON. -type Metadata map[string]any - -// Credentials represent client credentials: its -// "identity" which can be a username, email, generated name; -// and "secret" which can be a password or access token. -type Credentials struct { - Identity string `json:"identity,omitempty"` // username or generated login ID - Secret string `json:"secret,omitempty"` // password or token -} diff --git a/clients/doc.go b/clients/doc.go deleted file mode 100644 index 6295edfb7..000000000 --- a/clients/doc.go +++ /dev/null @@ -1,11 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package clients contains the domain concept definitions needed to -// support Magistrala clients service functionality. -// -// This package defines the core domain concepts and types necessary to -// handle clients in the context of a Magistrala clients service. It abstracts -// the underlying complexities of user management and provides a structured -// approach to working with clients. -package clients diff --git a/clients/errors.go b/clients/errors.go deleted file mode 100644 index e5ee29516..000000000 --- a/clients/errors.go +++ /dev/null @@ -1,14 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package clients - -import "errors" - -var ( - // ErrEnableClient indicates error in enabling client. - ErrEnableClient = errors.New("failed to enable client") - - // ErrDisableClient indicates error in disabling client. - ErrDisableClient = errors.New("failed to disable client") -) diff --git a/clients/events/doc.go b/clients/events/doc.go deleted file mode 100644 index 720686489..000000000 --- a/clients/events/doc.go +++ /dev/null @@ -1,6 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package events provides the domain concept definitions needed to support -// clients events functionality. -package events diff --git a/clients/events/events.go b/clients/events/events.go deleted file mode 100644 index e8a64dce3..000000000 --- a/clients/events/events.go +++ /dev/null @@ -1,351 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package events - -import ( - "time" - - "github.com/absmach/magistrala/clients" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/events" - "github.com/absmach/magistrala/pkg/roles" -) - -const ( - clientPrefix = "client." - clientCreate = clientPrefix + "create" - clientUpdate = clientPrefix + "update" - clientUpdateTags = clientPrefix + "update_tags" - clientUpdateSecret = clientPrefix + "update_secret" - clientEnable = clientPrefix + "enable" - clientDisable = clientPrefix + "disable" - clientRemove = clientPrefix + "remove" - clientView = clientPrefix + "view" - clientList = clientPrefix + "list" - clientListByUser = clientPrefix + "list_by_user" - clientSetParent = clientPrefix + "set_parent" - clientRemoveParent = clientPrefix + "remove_parent" -) - -var ( - _ events.Event = (*createClientEvent)(nil) - _ events.Event = (*updateClientEvent)(nil) - _ events.Event = (*changeClientStatusEvent)(nil) - _ events.Event = (*viewClientEvent)(nil) - _ events.Event = (*listClientEvent)(nil) - _ events.Event = (*listUserClientEvent)(nil) - _ events.Event = (*removeClientEvent)(nil) - _ events.Event = (*setParentGroupEvent)(nil) - _ events.Event = (*removeParentGroupEvent)(nil) -) - -type createClientEvent struct { - clients.Client - rolesProvisioned []roles.RoleProvision - authn.Session - requestID string -} - -func (cce createClientEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": clientCreate, - "id": cce.ID, - "roles_provisioned": cce.rolesProvisioned, - "status": cce.Status.String(), - "created_at": cce.CreatedAt, - "domain": cce.DomainID, - "user_id": cce.UserID, - "token_type": cce.Type.String(), - "super_admin": cce.SuperAdmin, - "request_id": cce.requestID, - } - - if cce.Name != "" { - val["name"] = cce.Name - } - if len(cce.Tags) > 0 { - val["tags"] = cce.Tags - } - if cce.Metadata != nil { - val["metadata"] = cce.Metadata - } - if cce.PrivateMetadata != nil { - val["private_metadata"] = cce.PrivateMetadata - } - if cce.Credentials.Identity != "" { - val["identity"] = cce.Credentials.Identity - } - - return val, nil -} - -type updateClientEvent struct { - clients.Client - operation string - authn.Session - requestID string -} - -func (uce updateClientEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": uce.operation, - "updated_at": uce.UpdatedAt, - "updated_by": uce.UpdatedBy, - "domain": uce.DomainID, - "user_id": uce.UserID, - "token_type": uce.Type.String(), - "super_admin": uce.SuperAdmin, - "request_id": uce.requestID, - } - if uce.ID != "" { - val["id"] = uce.ID - } - if uce.Name != "" { - val["name"] = uce.Name - } - if len(uce.Tags) > 0 { - val["tags"] = uce.Tags - } - if uce.Credentials.Identity != "" { - val["identity"] = uce.Credentials.Identity - } - if uce.Metadata != nil { - val["metadata"] = uce.Metadata - } - if uce.PrivateMetadata != nil { - val["private_metadata"] = uce.PrivateMetadata - } - if !uce.CreatedAt.IsZero() { - val["created_at"] = uce.CreatedAt - } - if uce.Status.String() != "" { - val["status"] = uce.Status.String() - } - - return val, nil -} - -type changeClientStatusEvent struct { - id string - operation string - status string - updatedAt time.Time - updatedBy string - authn.Session - requestID string -} - -func (cse changeClientStatusEvent) Encode() (map[string]any, error) { - return map[string]any{ - "operation": cse.operation, - "id": cse.id, - "status": cse.status, - "updated_at": cse.updatedAt, - "updated_by": cse.updatedBy, - "domain": cse.DomainID, - "user_id": cse.UserID, - "token_type": cse.Type.String(), - "super_admin": cse.SuperAdmin, - "request_id": cse.requestID, - }, nil -} - -type viewClientEvent struct { - clients.Client - authn.Session - requestID string -} - -func (vce viewClientEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": clientView, - "id": vce.ID, - "domain": vce.DomainID, - "user_id": vce.UserID, - "token_type": vce.Type.String(), - "super_admin": vce.SuperAdmin, - "request_id": vce.requestID, - } - - if vce.Name != "" { - val["name"] = vce.Name - } - if len(vce.Tags) > 0 { - val["tags"] = vce.Tags - } - if vce.Credentials.Identity != "" { - val["identity"] = vce.Credentials.Identity - } - if vce.Metadata != nil { - val["metadata"] = vce.Metadata - } - if vce.PrivateMetadata != nil { - val["private_metadata"] = vce.PrivateMetadata - } - if !vce.CreatedAt.IsZero() { - val["created_at"] = vce.CreatedAt - } - if !vce.UpdatedAt.IsZero() { - val["updated_at"] = vce.UpdatedAt - } - if vce.UpdatedBy != "" { - val["updated_by"] = vce.UpdatedBy - } - if vce.Status.String() != "" { - val["status"] = vce.Status.String() - } - - return val, nil -} - -type listClientEvent struct { - clients.Page - authn.Session - requestID string -} - -func (lce listClientEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": clientList, - "total": lce.Total, - "offset": lce.Offset, - "limit": lce.Limit, - "domain": lce.DomainID, - "user_id": lce.UserID, - "token_type": lce.Type.String(), - "super_admin": lce.SuperAdmin, - "request_id": lce.requestID, - } - - if lce.Name != "" { - val["name"] = lce.Name - } - if lce.Order != "" { - val["order"] = lce.Order - } - if lce.Dir != "" { - val["dir"] = lce.Dir - } - if lce.Metadata != nil { - val["metadata"] = lce.Metadata - } - if len(lce.Tags.Elements) > 0 { - val["tag"] = lce.Tags.Elements - } - if lce.Status.String() != "" { - val["status"] = lce.Status.String() - } - if len(lce.IDs) > 0 { - val["ids"] = lce.IDs - } - if lce.Identity != "" { - val["identity"] = lce.Identity - } - return val, nil -} - -type listUserClientEvent struct { - userID string - clients.Page - authn.Session - requestID string -} - -func (lce listUserClientEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": clientList, - "req_user_id": lce.userID, - "total": lce.Total, - "offset": lce.Offset, - "limit": lce.Limit, - "domain": lce.DomainID, - "user_id": lce.UserID, - "token_type": lce.Type.String(), - "super_admin": lce.SuperAdmin, - "request_id": lce.requestID, - } - - if lce.Name != "" { - val["name"] = lce.Name - } - if lce.Order != "" { - val["order"] = lce.Order - } - if lce.Dir != "" { - val["dir"] = lce.Dir - } - if lce.Metadata != nil { - val["metadata"] = lce.Metadata - } - if len(lce.Tags.Elements) > 0 { - val["tag"] = lce.Tags.Elements - } - if lce.Status.String() != "" { - val["status"] = lce.Status.String() - } - if len(lce.IDs) > 0 { - val["ids"] = lce.IDs - } - if lce.Identity != "" { - val["identity"] = lce.Identity - } - - return val, nil -} - -type removeClientEvent struct { - id string - authn.Session - requestID string -} - -func (dce removeClientEvent) Encode() (map[string]any, error) { - return map[string]any{ - "operation": clientRemove, - "id": dce.id, - "domain": dce.DomainID, - "user_id": dce.UserID, - "token_type": dce.Type.String(), - "super_admin": dce.SuperAdmin, - "request_id": dce.requestID, - }, nil -} - -type setParentGroupEvent struct { - id string - parentGroupID string - authn.Session - requestID string -} - -func (spge setParentGroupEvent) Encode() (map[string]any, error) { - return map[string]any{ - "operation": clientSetParent, - "id": spge.id, - "parent_group_id": spge.parentGroupID, - "domain": spge.DomainID, - "user_id": spge.UserID, - "token_type": spge.Type.String(), - "super_admin": spge.SuperAdmin, - "request_id": spge.requestID, - }, nil -} - -type removeParentGroupEvent struct { - id string - authn.Session - requestID string -} - -func (rpge removeParentGroupEvent) Encode() (map[string]any, error) { - return map[string]any{ - "operation": clientRemoveParent, - "id": rpge.id, - "domain": rpge.DomainID, - "user_id": rpge.UserID, - "token_type": rpge.Type.String(), - "super_admin": rpge.SuperAdmin, - "request_id": rpge.requestID, - }, nil -} diff --git a/clients/events/streams.go b/clients/events/streams.go deleted file mode 100644 index 9504e31d3..000000000 --- a/clients/events/streams.go +++ /dev/null @@ -1,262 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package events - -import ( - "context" - - "github.com/absmach/magistrala/clients" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/events" - "github.com/absmach/magistrala/pkg/events/store" - "github.com/absmach/magistrala/pkg/roles" - rmEvents "github.com/absmach/magistrala/pkg/roles/rolemanager/events" - "github.com/go-chi/chi/v5/middleware" -) - -const ( - magistralaPrefix = "magistrala." - createStream = magistralaPrefix + clientCreate - updateStream = magistralaPrefix + clientUpdate - updateTagsStream = magistralaPrefix + clientUpdateTags - updateSecretStream = magistralaPrefix + clientUpdateSecret - enableStream = magistralaPrefix + clientEnable - disableStream = magistralaPrefix + clientDisable - removeStream = magistralaPrefix + clientRemove - viewStream = magistralaPrefix + clientView - listStream = magistralaPrefix + clientList - listByUserStream = magistralaPrefix + clientListByUser - setParentStream = magistralaPrefix + clientSetParent - removeParentStream = magistralaPrefix + clientRemoveParent -) - -var _ clients.Service = (*eventStore)(nil) - -type eventStore struct { - events.Publisher - svc clients.Service - rmEvents.RoleManagerEventStore -} - -// NewEventStoreMiddleware returns wrapper around clients service that sends -// events to event store. -func NewEventStoreMiddleware(ctx context.Context, svc clients.Service, url string) (clients.Service, error) { - publisher, err := store.NewPublisher(ctx, url, "clients-es-pub") - if err != nil { - return nil, err - } - res := rmEvents.NewRoleManagerEventStore("clients", clientPrefix, svc, publisher) - - return &eventStore{ - svc: svc, - Publisher: publisher, - RoleManagerEventStore: res, - }, nil -} - -func (es *eventStore) CreateClients(ctx context.Context, session authn.Session, clients ...clients.Client) ([]clients.Client, []roles.RoleProvision, error) { - clis, rps, err := es.svc.CreateClients(ctx, session, clients...) - if err != nil { - return clis, rps, err - } - - for _, cli := range clis { - event := createClientEvent{ - Client: cli, - rolesProvisioned: rps, - Session: session, - requestID: middleware.GetReqID(ctx), - } - if err := es.Publish(ctx, createStream, event); err != nil { - return clis, rps, err - } - } - - return clis, rps, nil -} - -func (es *eventStore) Update(ctx context.Context, session authn.Session, client clients.Client) (clients.Client, error) { - cli, err := es.svc.Update(ctx, session, client) - if err != nil { - return cli, err - } - - return es.update(ctx, session, clientUpdate, updateStream, cli) -} - -func (es *eventStore) UpdateTags(ctx context.Context, session authn.Session, client clients.Client) (clients.Client, error) { - cli, err := es.svc.UpdateTags(ctx, session, client) - if err != nil { - return cli, err - } - - return es.update(ctx, session, clientUpdateTags, updateTagsStream, cli) -} - -func (es *eventStore) UpdateSecret(ctx context.Context, session authn.Session, id, key string) (clients.Client, error) { - cli, err := es.svc.UpdateSecret(ctx, session, id, key) - if err != nil { - return cli, err - } - - return es.update(ctx, session, clientUpdateSecret, updateSecretStream, cli) -} - -func (es *eventStore) update(ctx context.Context, session authn.Session, operation, stream string, client clients.Client) (clients.Client, error) { - event := updateClientEvent{ - Client: client, - operation: operation, - Session: session, - requestID: middleware.GetReqID(ctx), - } - - if err := es.Publish(ctx, stream, event); err != nil { - return client, err - } - - return client, nil -} - -func (es *eventStore) View(ctx context.Context, session authn.Session, id string, withRoles bool) (clients.Client, error) { - cli, err := es.svc.View(ctx, session, id, withRoles) - if err != nil { - return cli, err - } - - event := viewClientEvent{ - Client: cli, - Session: session, - requestID: middleware.GetReqID(ctx), - } - if err := es.Publish(ctx, viewStream, event); err != nil { - return cli, err - } - - return cli, nil -} - -func (es *eventStore) ListClients(ctx context.Context, session authn.Session, pm clients.Page) (clients.ClientsPage, error) { - cp, err := es.svc.ListClients(ctx, session, pm) - if err != nil { - return cp, err - } - event := listClientEvent{ - pm, - session, - middleware.GetReqID(ctx), - } - if err := es.Publish(ctx, listStream, event); err != nil { - return cp, err - } - - return cp, nil -} - -func (es *eventStore) ListUserClients(ctx context.Context, session authn.Session, userID string, pm clients.Page) (clients.ClientsPage, error) { - cp, err := es.svc.ListUserClients(ctx, session, userID, pm) - if err != nil { - return cp, err - } - event := listUserClientEvent{ - userID, - pm, - session, - middleware.GetReqID(ctx), - } - if err := es.Publish(ctx, listByUserStream, event); err != nil { - return cp, err - } - - return cp, nil -} - -func (es *eventStore) Enable(ctx context.Context, session authn.Session, id string) (clients.Client, error) { - cli, err := es.svc.Enable(ctx, session, id) - if err != nil { - return cli, err - } - - return es.changeStatus(ctx, session, clientEnable, enableStream, cli) -} - -func (es *eventStore) Disable(ctx context.Context, session authn.Session, id string) (clients.Client, error) { - cli, err := es.svc.Disable(ctx, session, id) - if err != nil { - return cli, err - } - - return es.changeStatus(ctx, session, clientDisable, disableStream, cli) -} - -func (es *eventStore) changeStatus(ctx context.Context, session authn.Session, operation, stream string, cli clients.Client) (clients.Client, error) { - event := changeClientStatusEvent{ - id: cli.ID, - operation: operation, - updatedAt: cli.UpdatedAt, - updatedBy: cli.UpdatedBy, - status: cli.Status.String(), - Session: session, - requestID: middleware.GetReqID(ctx), - } - if err := es.Publish(ctx, stream, event); err != nil { - return cli, err - } - - return cli, nil -} - -func (es *eventStore) Delete(ctx context.Context, session authn.Session, id string) error { - if err := es.svc.Delete(ctx, session, id); err != nil { - return err - } - - event := removeClientEvent{ - id: id, - Session: session, - requestID: middleware.GetReqID(ctx), - } - - if err := es.Publish(ctx, removeStream, event); err != nil { - return err - } - - return nil -} - -func (es *eventStore) SetParentGroup(ctx context.Context, session authn.Session, parentGroupID string, id string) (err error) { - if err := es.svc.SetParentGroup(ctx, session, parentGroupID, id); err != nil { - return err - } - - event := setParentGroupEvent{ - parentGroupID: parentGroupID, - id: id, - Session: session, - requestID: middleware.GetReqID(ctx), - } - - if err := es.Publish(ctx, setParentStream, event); err != nil { - return err - } - - return nil -} - -func (es *eventStore) RemoveParentGroup(ctx context.Context, session authn.Session, id string) (err error) { - if err := es.svc.RemoveParentGroup(ctx, session, id); err != nil { - return err - } - - event := removeParentGroupEvent{ - id: id, - Session: session, - requestID: middleware.GetReqID(ctx), - } - - if err := es.Publish(ctx, removeParentStream, event); err != nil { - return err - } - - return nil -} diff --git a/clients/events/streams_test.go b/clients/events/streams_test.go deleted file mode 100644 index ad6d228f8..000000000 --- a/clients/events/streams_test.go +++ /dev/null @@ -1,635 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package events_test - -import ( - "context" - "fmt" - "os" - "testing" - "time" - - "github.com/absmach/magistrala/clients" - "github.com/absmach/magistrala/clients/events" - "github.com/absmach/magistrala/clients/mocks" - "github.com/absmach/magistrala/internal/testsutil" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/pkg/roles" - "github.com/go-chi/chi/v5/middleware" - "github.com/redis/go-redis/v9" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" -) - -var ( - storeClient *redis.Client - storeURL string - validSession = authn.Session{ - DomainID: testsutil.GenerateUUID(&testing.T{}), - UserID: testsutil.GenerateUUID(&testing.T{}), - } - validClient = generateTestClient(&testing.T{}) - validClientsPage = clients.ClientsPage{ - Page: clients.Page{ - Limit: 10, - Offset: 0, - Total: 1, - }, - Clients: []clients.Client{validClient}, - } -) - -func newEventStoreMiddleware(t *testing.T) (*mocks.Service, clients.Service) { - svc := new(mocks.Service) - nsvc, err := events.NewEventStoreMiddleware(context.Background(), svc, storeURL) - require.Nil(t, err, fmt.Sprintf("create events store middleware failed with unexpected error: %s", err)) - - return svc, nsvc -} - -func TestMain(m *testing.M) { - code := testsutil.RunRedisTest(m, &storeClient, &storeURL) - os.Exit(code) -} - -func TestCreateClients(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validID := testsutil.GenerateUUID(t) - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, validID) - - cases := []struct { - desc string - session authn.Session - clients []clients.Client - svcRes []clients.Client - svcRoleRes []roles.RoleProvision - svcErr error - resp []clients.Client - respRoleRes []roles.RoleProvision - err error - }{ - { - desc: "publish successfully", - session: validSession, - clients: []clients.Client{validClient}, - svcRes: []clients.Client{validClient}, - svcRoleRes: []roles.RoleProvision{}, - svcErr: nil, - resp: []clients.Client{validClient}, - respRoleRes: []roles.RoleProvision{}, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - clients: []clients.Client{validClient}, - svcRes: []clients.Client{}, - svcRoleRes: []roles.RoleProvision{}, - svcErr: svcerr.ErrCreateEntity, - resp: []clients.Client{}, - respRoleRes: []roles.RoleProvision{}, - err: svcerr.ErrCreateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("CreateClients", validCtx, tc.session, tc.clients).Return(tc.svcRes, tc.svcRoleRes, tc.svcErr) - resp, respRoleRes, err := nsvc.CreateClients(validCtx, tc.session, tc.clients...) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - assert.Equal(t, tc.respRoleRes, respRoleRes, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.respRoleRes, respRoleRes)) - svcCall.Unset() - }) - } -} - -func TestView(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - clientID string - withRoles bool - svcRes clients.Client - svcErr error - resp clients.Client - err error - }{ - { - desc: "publish successfully", - session: validSession, - clientID: validClient.ID, - withRoles: false, - svcRes: validClient, - svcErr: nil, - resp: validClient, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - clientID: validClient.ID, - withRoles: false, - svcRes: clients.Client{}, - svcErr: svcerr.ErrViewEntity, - resp: clients.Client{}, - err: svcerr.ErrViewEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("View", validCtx, tc.session, tc.clientID, tc.withRoles).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.View(validCtx, tc.session, tc.clientID, tc.withRoles) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestUpdate(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - updatedClient := validClient - updatedClient.Name = "updatedName" - - cases := []struct { - desc string - session authn.Session - client clients.Client - svcRes clients.Client - svcErr error - resp clients.Client - err error - }{ - { - desc: "publish successfully", - session: validSession, - client: updatedClient, - svcRes: updatedClient, - svcErr: nil, - resp: updatedClient, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - client: updatedClient, - svcRes: clients.Client{}, - svcErr: svcerr.ErrUpdateEntity, - resp: clients.Client{}, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("Update", validCtx, tc.session, tc.client).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.Update(validCtx, tc.session, tc.client) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestUpdateTags(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - updatedClient := validClient - updatedClient.Tags = []string{"newTag1", "newTag2"} - - cases := []struct { - desc string - session authn.Session - client clients.Client - svcRes clients.Client - svcErr error - resp clients.Client - err error - }{ - { - desc: "publish successfully", - session: validSession, - client: updatedClient, - svcRes: updatedClient, - svcErr: nil, - resp: updatedClient, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - client: updatedClient, - svcRes: clients.Client{}, - svcErr: svcerr.ErrUpdateEntity, - resp: clients.Client{}, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("UpdateTags", validCtx, tc.session, tc.client).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.UpdateTags(validCtx, tc.session, tc.client) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestUpdateSecret(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - updatedClient := validClient - updatedClient.Credentials.Secret = "newSecret" - - cases := []struct { - desc string - session authn.Session - clientID string - newSecret string - svcRes clients.Client - svcErr error - resp clients.Client - err error - }{ - { - desc: "publish successfully", - session: validSession, - clientID: validClient.ID, - newSecret: "newSecret", - svcRes: updatedClient, - svcErr: nil, - resp: updatedClient, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - clientID: validClient.ID, - newSecret: "newSecret", - svcRes: clients.Client{}, - svcErr: svcerr.ErrUpdateEntity, - resp: clients.Client{}, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("UpdateSecret", validCtx, tc.session, tc.clientID, tc.newSecret).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.UpdateSecret(validCtx, tc.session, tc.clientID, tc.newSecret) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestEnable(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - clientID string - svcRes clients.Client - svcErr error - resp clients.Client - err error - }{ - { - desc: "publish successfully", - session: validSession, - clientID: validClient.ID, - svcRes: validClient, - svcErr: nil, - resp: validClient, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - clientID: validClient.ID, - svcRes: clients.Client{}, - svcErr: svcerr.ErrUpdateEntity, - resp: clients.Client{}, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("Enable", validCtx, tc.session, tc.clientID).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.Enable(validCtx, tc.session, tc.clientID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestDisable(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - clientID string - svcRes clients.Client - svcErr error - resp clients.Client - err error - }{ - { - desc: "publish successfully", - session: validSession, - clientID: validClient.ID, - svcRes: validClient, - svcErr: nil, - resp: validClient, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - clientID: validClient.ID, - svcRes: clients.Client{}, - svcErr: svcerr.ErrUpdateEntity, - resp: clients.Client{}, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("Disable", validCtx, tc.session, tc.clientID).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.Disable(validCtx, tc.session, tc.clientID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestListClients(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - pageMeta clients.Page - svcRes clients.ClientsPage - svcErr error - resp clients.ClientsPage - err error - }{ - { - desc: "publish successfully", - session: validSession, - pageMeta: clients.Page{ - Limit: 10, - Offset: 0, - }, - svcRes: validClientsPage, - svcErr: nil, - resp: validClientsPage, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - pageMeta: clients.Page{ - Limit: 10, - Offset: 0, - }, - svcRes: clients.ClientsPage{}, - svcErr: svcerr.ErrViewEntity, - resp: clients.ClientsPage{}, - err: svcerr.ErrViewEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("ListClients", validCtx, tc.session, tc.pageMeta).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.ListClients(validCtx, tc.session, tc.pageMeta) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestListUserClients(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - userID string - pageMeta clients.Page - svcRes clients.ClientsPage - svcErr error - resp clients.ClientsPage - err error - }{ - { - desc: "publish successfully", - session: validSession, - userID: validSession.UserID, - pageMeta: clients.Page{ - Limit: 10, - Offset: 0, - }, - svcRes: validClientsPage, - svcErr: nil, - resp: validClientsPage, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - userID: validSession.UserID, - pageMeta: clients.Page{ - Limit: 10, - Offset: 0, - }, - svcRes: clients.ClientsPage{}, - svcErr: svcerr.ErrViewEntity, - resp: clients.ClientsPage{}, - err: svcerr.ErrViewEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("ListUserClients", validCtx, tc.session, tc.userID, tc.pageMeta).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.ListUserClients(validCtx, tc.session, tc.userID, tc.pageMeta) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestDelete(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - clientID string - svcErr error - err error - }{ - { - desc: "publish successfully", - session: validSession, - clientID: validClient.ID, - svcErr: nil, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - clientID: validClient.ID, - svcErr: svcerr.ErrRemoveEntity, - err: svcerr.ErrRemoveEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("Delete", validCtx, tc.session, tc.clientID).Return(tc.svcErr) - err := nsvc.Delete(validCtx, tc.session, tc.clientID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - svcCall.Unset() - }) - } -} - -func TestSetParentGroup(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - parentGroupID string - clientID string - svcErr error - err error - }{ - { - desc: "publish successfully", - session: validSession, - parentGroupID: testsutil.GenerateUUID(t), - clientID: validClient.ID, - svcErr: nil, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - parentGroupID: testsutil.GenerateUUID(t), - clientID: validClient.ID, - svcErr: svcerr.ErrUpdateEntity, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("SetParentGroup", validCtx, tc.session, tc.parentGroupID, tc.clientID).Return(tc.svcErr) - err := nsvc.SetParentGroup(validCtx, tc.session, tc.parentGroupID, tc.clientID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - svcCall.Unset() - }) - } -} - -func TestRemoveParentGroup(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - clientID string - svcErr error - err error - }{ - { - desc: "publish successfully", - session: validSession, - clientID: validClient.ID, - svcErr: nil, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - clientID: validClient.ID, - svcErr: svcerr.ErrUpdateEntity, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("RemoveParentGroup", validCtx, tc.session, tc.clientID).Return(tc.svcErr) - err := nsvc.RemoveParentGroup(validCtx, tc.session, tc.clientID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - svcCall.Unset() - }) - } -} - -func generateTestClient(t *testing.T) clients.Client { - createdAt, err := time.Parse(time.RFC3339, "2024-01-01T00:00:00Z") - assert.Nil(t, err, fmt.Sprintf("Unexpected error parsing time: %v", err)) - return clients.Client{ - ID: testsutil.GenerateUUID(t), - Name: "clientname", - Domain: testsutil.GenerateUUID(t), - Tags: []string{"tag1", "tag2"}, - Credentials: clients.Credentials{ - Identity: "clientidentity", - Secret: "clientsecret", - }, - Metadata: clients.Metadata{"key1": "value1"}, - CreatedAt: createdAt, - UpdatedAt: createdAt, - Status: clients.EnabledStatus, - } -} diff --git a/clients/middleware/authorization.go b/clients/middleware/authorization.go deleted file mode 100644 index 3cbfd9b82..000000000 --- a/clients/middleware/authorization.go +++ /dev/null @@ -1,307 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package middleware - -import ( - "context" - - "github.com/absmach/magistrala/auth" - "github.com/absmach/magistrala/clients" - "github.com/absmach/magistrala/clients/operations" - dOperations "github.com/absmach/magistrala/domains/operations" - gOperations "github.com/absmach/magistrala/groups/operations" - "github.com/absmach/magistrala/pkg/authn" - smqauthz "github.com/absmach/magistrala/pkg/authz" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/pkg/permissions" - "github.com/absmach/magistrala/pkg/policies" - "github.com/absmach/magistrala/pkg/roles" - rolemgr "github.com/absmach/magistrala/pkg/roles/rolemanager/middleware" -) - -var ( - errView = errors.New("not authorized to view client") - errUpdate = errors.New("not authorized to update client") - errUpdateTags = errors.New("not authorized to update client tags") - errUpdateSecret = errors.New("not authorized to update client secret") - errEnable = errors.New("not authorized to enable client") - errDisable = errors.New("not authorized to disable client") - errDelete = errors.New("not authorized to delete client") - errSetParentGroup = errors.New("not authorized to set parent group to client") - errRemoveParentGroup = errors.New("not authorized to remove parent group from client") - errDomainCreateClients = errors.New("not authorized to create client in domain") - errGroupSetChildClients = errors.New("not authorized to set child client for group") - errGroupRemoveChildClients = errors.New("not authorized to remove child client for group") -) - -var _ clients.Service = (*authorizationMiddleware)(nil) - -type authorizationMiddleware struct { - svc clients.Service - repo clients.Repository - authz smqauthz.Authorization - entitiesOps permissions.EntitiesOperations[permissions.Operation] - rolemgr.RoleManagerAuthorizationMiddleware -} - -// NewAuthorization adds authorization to the clients service. -func NewAuthorization( - entityType string, - svc clients.Service, - authz smqauthz.Authorization, - repo clients.Repository, - entitiesOps permissions.EntitiesOperations[permissions.Operation], - roleOps permissions.Operations[permissions.RoleOperation], -) (clients.Service, error) { - if err := entitiesOps.Validate(); err != nil { - return nil, err - } - ram, err := rolemgr.NewAuthorization(policies.ClientType, svc, authz, roleOps) - if err != nil { - return nil, err - } - - return &authorizationMiddleware{ - svc: svc, - authz: authz, - repo: repo, - entitiesOps: entitiesOps, - RoleManagerAuthorizationMiddleware: ram, - }, nil -} - -func (am *authorizationMiddleware) CreateClients(ctx context.Context, session authn.Session, client ...clients.Client) ([]clients.Client, []roles.RoleProvision, error) { - if err := am.authorize(ctx, session, policies.DomainType, dOperations.OpCreateDomainClients, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - Subject: session.DomainUserID, - ObjectType: policies.DomainType, - Object: session.DomainID, - }); err != nil { - return []clients.Client{}, []roles.RoleProvision{}, errors.Wrap(err, errDomainCreateClients) - } - - return am.svc.CreateClients(ctx, session, client...) -} - -func (am *authorizationMiddleware) View(ctx context.Context, session authn.Session, id string, withRoles bool) (clients.Client, error) { - if err := am.authorize(ctx, session, policies.ClientType, operations.OpViewClient, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - Subject: session.DomainUserID, - ObjectType: policies.ClientType, - Object: id, - }); err != nil { - return clients.Client{}, errors.Wrap(err, errView) - } - - return am.svc.View(ctx, session, id, withRoles) -} - -func (am *authorizationMiddleware) ListClients(ctx context.Context, session authn.Session, pm clients.Page) (clients.ClientsPage, error) { - if err := am.checkSuperAdmin(ctx, session); err == nil { - session.SuperAdmin = true - } - - return am.svc.ListClients(ctx, session, pm) -} - -func (am *authorizationMiddleware) ListUserClients(ctx context.Context, session authn.Session, userID string, pm clients.Page) (clients.ClientsPage, error) { - if err := am.checkSuperAdmin(ctx, session); err != nil { - return clients.ClientsPage{}, err - } - - return am.svc.ListUserClients(ctx, session, userID, pm) -} - -func (am *authorizationMiddleware) Update(ctx context.Context, session authn.Session, client clients.Client) (clients.Client, error) { - if err := am.authorize(ctx, session, policies.ClientType, operations.OpUpdateClient, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - Subject: session.DomainUserID, - ObjectType: policies.ClientType, - Object: client.ID, - }); err != nil { - return clients.Client{}, errors.Wrap(err, errUpdate) - } - - return am.svc.Update(ctx, session, client) -} - -func (am *authorizationMiddleware) UpdateTags(ctx context.Context, session authn.Session, client clients.Client) (clients.Client, error) { - if err := am.authorize(ctx, session, policies.ClientType, operations.OpUpdateClientTags, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - Subject: session.DomainUserID, - ObjectType: policies.ClientType, - Object: client.ID, - }); err != nil { - return clients.Client{}, errors.Wrap(err, errUpdateTags) - } - - return am.svc.UpdateTags(ctx, session, client) -} - -func (am *authorizationMiddleware) UpdateSecret(ctx context.Context, session authn.Session, id, key string) (clients.Client, error) { - if err := am.authorize(ctx, session, policies.ClientType, operations.OpUpdateClientSecret, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - Subject: session.DomainUserID, - ObjectType: policies.ClientType, - Object: id, - }); err != nil { - return clients.Client{}, errors.Wrap(err, errUpdateSecret) - } - - return am.svc.UpdateSecret(ctx, session, id, key) -} - -func (am *authorizationMiddleware) Enable(ctx context.Context, session authn.Session, id string) (clients.Client, error) { - if err := am.authorize(ctx, session, policies.ClientType, operations.OpEnableClient, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - Subject: session.DomainUserID, - ObjectType: policies.ClientType, - Object: id, - }); err != nil { - return clients.Client{}, errors.Wrap(err, errEnable) - } - - return am.svc.Enable(ctx, session, id) -} - -func (am *authorizationMiddleware) Disable(ctx context.Context, session authn.Session, id string) (clients.Client, error) { - if err := am.authorize(ctx, session, policies.ClientType, operations.OpDisableClient, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - Subject: session.DomainUserID, - ObjectType: policies.ClientType, - Object: id, - }); err != nil { - return clients.Client{}, errors.Wrap(err, errDisable) - } - - return am.svc.Disable(ctx, session, id) -} - -func (am *authorizationMiddleware) Delete(ctx context.Context, session authn.Session, id string) error { - if err := am.authorize(ctx, session, policies.ClientType, operations.OpDeleteClient, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - Subject: session.DomainUserID, - ObjectType: policies.ClientType, - Object: id, - }); err != nil { - return errors.Wrap(err, errDelete) - } - - return am.svc.Delete(ctx, session, id) -} - -func (am *authorizationMiddleware) SetParentGroup(ctx context.Context, session authn.Session, parentGroupID string, id string) error { - if err := am.authorize(ctx, session, policies.ClientType, operations.OpSetParentGroup, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - Subject: session.DomainUserID, - ObjectType: policies.ClientType, - Object: id, - }); err != nil { - return errors.Wrap(err, errSetParentGroup) - } - - if err := am.authorize(ctx, session, policies.GroupType, gOperations.OpGroupSetChildClient, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - Subject: session.DomainUserID, - ObjectType: policies.GroupType, - Object: parentGroupID, - }); err != nil { - return errors.Wrap(err, errGroupSetChildClients) - } - - return am.svc.SetParentGroup(ctx, session, parentGroupID, id) -} - -func (am *authorizationMiddleware) RemoveParentGroup(ctx context.Context, session authn.Session, id string) error { - if err := am.authorize(ctx, session, policies.ClientType, operations.OpRemoveParentGroup, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - Subject: session.DomainUserID, - ObjectType: policies.ClientType, - Object: id, - }); err != nil { - return errors.Wrap(err, errRemoveParentGroup) - } - - th, err := am.repo.RetrieveByID(ctx, id) - if err != nil { - return errors.Wrap(svcerr.ErrRemoveEntity, err) - } - - if th.ParentGroup != "" { - if err := am.authorize(ctx, session, policies.GroupType, gOperations.OpGroupRemoveChildClient, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - Subject: session.DomainUserID, - ObjectType: policies.GroupType, - Object: th.ParentGroup, - }); err != nil { - return errors.Wrap(err, errGroupRemoveChildClients) - } - - return am.svc.RemoveParentGroup(ctx, session, id) - } - return nil -} - -func (am *authorizationMiddleware) authorize(ctx context.Context, session authn.Session, entityType string, op permissions.Operation, req smqauthz.PolicyReq) error { - req.Domain = session.DomainID - - perm, err := am.entitiesOps.GetPermission(entityType, op) - if err != nil { - return err - } - - req.Permission = perm.String() - - var pat *smqauthz.PATReq - if session.PatID != "" { - entityID := req.Object - opName := am.entitiesOps.OperationName(entityType, op) - if op == operations.OpListUserClients || op == dOperations.OpCreateDomainClients || op == dOperations.OpListDomainClients { - entityID = auth.AnyIDs - } - pat = &smqauthz.PATReq{ - UserID: session.UserID, - PatID: session.PatID, - EntityID: entityID, - EntityType: auth.ClientsType.String(), - Operation: opName, - Domain: session.DomainID, - } - } - - if err := am.authz.Authorize(ctx, req, pat); err != nil { - return err - } - - return nil -} - -func (am *authorizationMiddleware) checkSuperAdmin(ctx context.Context, session authn.Session) error { - if session.Role != authn.SuperAdminRole { - return svcerr.ErrSuperAdminAction - } - if err := am.authz.Authorize(ctx, smqauthz.PolicyReq{ - SubjectType: policies.UserType, - Subject: session.UserID, - Permission: policies.AdminPermission, - ObjectType: policies.PlatformType, - Object: policies.MagistralaObject, - }, nil); err != nil { - return err - } - return nil -} diff --git a/clients/middleware/callout.go b/clients/middleware/callout.go deleted file mode 100644 index 4270b0102..000000000 --- a/clients/middleware/callout.go +++ /dev/null @@ -1,228 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package middleware - -import ( - "context" - "time" - - "github.com/absmach/magistrala/clients" - "github.com/absmach/magistrala/clients/operations" - dOperations "github.com/absmach/magistrala/domains/operations" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/callout" - "github.com/absmach/magistrala/pkg/permissions" - "github.com/absmach/magistrala/pkg/policies" - "github.com/absmach/magistrala/pkg/roles" - rolemgr "github.com/absmach/magistrala/pkg/roles/rolemanager/middleware" -) - -var _ clients.Service = (*calloutMiddleware)(nil) - -type calloutMiddleware struct { - svc clients.Service - repo clients.Repository - callout callout.Callout - entitiesOps permissions.EntitiesOperations[permissions.Operation] - rolemgr.RoleManagerCalloutMiddleware -} - -func NewCallout(svc clients.Service, repo clients.Repository, entitiesOps permissions.EntitiesOperations[permissions.Operation], roleOps permissions.Operations[permissions.RoleOperation], callout callout.Callout) (clients.Service, error) { - call, err := rolemgr.NewCallout(policies.ClientType, svc, callout, roleOps) - if err != nil { - return nil, err - } - - if err := entitiesOps.Validate(); err != nil { - return nil, err - } - - return &calloutMiddleware{ - svc: svc, - repo: repo, - callout: callout, - entitiesOps: entitiesOps, - RoleManagerCalloutMiddleware: call, - }, nil -} - -func (cm *calloutMiddleware) CreateClients(ctx context.Context, session authn.Session, client ...clients.Client) ([]clients.Client, []roles.RoleProvision, error) { - params := map[string]any{ - "entities": client, - "count": len(client), - } - - if err := cm.callOut(ctx, session, policies.DomainType, dOperations.OpCreateDomainClients, params); err != nil { - return []clients.Client{}, []roles.RoleProvision{}, err - } - - return cm.svc.CreateClients(ctx, session, client...) -} - -func (cm *calloutMiddleware) View(ctx context.Context, session authn.Session, id string, withRoles bool) (clients.Client, error) { - params := map[string]any{ - "entity_id": id, - } - if err := cm.callOut(ctx, session, policies.ClientType, operations.OpViewClient, params); err != nil { - return clients.Client{}, err - } - - return cm.svc.View(ctx, session, id, withRoles) -} - -func (cm *calloutMiddleware) ListClients(ctx context.Context, session authn.Session, pm clients.Page) (clients.ClientsPage, error) { - params := map[string]any{ - "pagemeta": pm, - } - - if err := cm.callOut(ctx, session, policies.DomainType, dOperations.OpListDomainClients, params); err != nil { - return clients.ClientsPage{}, err - } - - return cm.svc.ListClients(ctx, session, pm) -} - -func (cm *calloutMiddleware) ListUserClients(ctx context.Context, session authn.Session, userID string, pm clients.Page) (clients.ClientsPage, error) { - params := map[string]any{ - "user_id": userID, - "pagemeta": pm, - } - - if err := cm.callOut(ctx, session, policies.ClientType, operations.OpListUserClients, params); err != nil { - return clients.ClientsPage{}, err - } - - return cm.svc.ListUserClients(ctx, session, userID, pm) -} - -func (cm *calloutMiddleware) Update(ctx context.Context, session authn.Session, client clients.Client) (clients.Client, error) { - params := map[string]any{ - "entity_id": client.ID, - } - - if err := cm.callOut(ctx, session, policies.ClientType, operations.OpUpdateClient, params); err != nil { - return clients.Client{}, err - } - - return cm.svc.Update(ctx, session, client) -} - -func (cm *calloutMiddleware) UpdateTags(ctx context.Context, session authn.Session, client clients.Client) (clients.Client, error) { - params := map[string]any{ - "entity_id": client.ID, - } - - if err := cm.callOut(ctx, session, policies.ClientType, operations.OpUpdateClientTags, params); err != nil { - return clients.Client{}, err - } - - return cm.svc.UpdateTags(ctx, session, client) -} - -func (cm *calloutMiddleware) UpdateSecret(ctx context.Context, session authn.Session, id, key string) (clients.Client, error) { - params := map[string]any{ - "entity_id": id, - } - - if err := cm.callOut(ctx, session, policies.ClientType, operations.OpUpdateClientSecret, params); err != nil { - return clients.Client{}, err - } - - return cm.svc.UpdateSecret(ctx, session, id, key) -} - -func (cm *calloutMiddleware) Enable(ctx context.Context, session authn.Session, id string) (clients.Client, error) { - params := map[string]any{ - "entity_id": id, - } - - if err := cm.callOut(ctx, session, policies.ClientType, operations.OpEnableClient, params); err != nil { - return clients.Client{}, err - } - - return cm.svc.Enable(ctx, session, id) -} - -func (cm *calloutMiddleware) Disable(ctx context.Context, session authn.Session, id string) (clients.Client, error) { - params := map[string]any{ - "entity_id": id, - } - - if err := cm.callOut(ctx, session, policies.ClientType, operations.OpDisableClient, params); err != nil { - return clients.Client{}, err - } - - return cm.svc.Disable(ctx, session, id) -} - -func (cm *calloutMiddleware) Delete(ctx context.Context, session authn.Session, id string) error { - params := map[string]any{ - "entity_id": id, - } - - if err := cm.callOut(ctx, session, policies.ClientType, operations.OpDeleteClient, params); err != nil { - return err - } - - return cm.svc.Delete(ctx, session, id) -} - -func (cm *calloutMiddleware) SetParentGroup(ctx context.Context, session authn.Session, parentGroupID string, id string) error { - params := map[string]any{ - "entity_id": id, - "parent_id": parentGroupID, - } - - if err := cm.callOut(ctx, session, policies.ClientType, operations.OpSetParentGroup, params); err != nil { - return err - } - - return cm.svc.SetParentGroup(ctx, session, parentGroupID, id) -} - -func (cm *calloutMiddleware) RemoveParentGroup(ctx context.Context, session authn.Session, id string) error { - th, err := cm.repo.RetrieveByID(ctx, id) - if err != nil { - return err - } - - if th.ParentGroup != "" { - params := map[string]any{ - "entity_id": id, - "parent_id": th.ParentGroup, - } - - if err := cm.callOut(ctx, session, policies.ClientType, operations.OpRemoveParentGroup, params); err != nil { - return err - } - } - - return cm.svc.RemoveParentGroup(ctx, session, id) -} - -func (cm *calloutMiddleware) callOut(ctx context.Context, session authn.Session, entityType string, op permissions.Operation, pld map[string]any) error { - var entityID string - if id, ok := pld["entity_id"].(string); ok { - entityID = id - } - - req := callout.Request{ - BaseRequest: callout.BaseRequest{ - Operation: cm.entitiesOps.OperationName(entityType, op), - EntityType: entityType, - EntityID: entityID, - CallerID: session.UserID, - CallerType: policies.UserType, - DomainID: session.DomainID, - Time: time.Now().UTC(), - }, - Payload: pld, - } - - if err := cm.callout.Callout(ctx, req); err != nil { - return err - } - - return nil -} diff --git a/clients/middleware/doc.go b/clients/middleware/doc.go deleted file mode 100644 index fd23f416c..000000000 --- a/clients/middleware/doc.go +++ /dev/null @@ -1,9 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package middleware provides authorization, logging, metrics and tracing middleware -// for Magistrala Clients Service. -// -// For more details about tracing instrumentation for Magistrala refer to the -// documentation at https://magistrala.absmach.eu/docs/tracing/. -package middleware diff --git a/clients/middleware/logging.go b/clients/middleware/logging.go deleted file mode 100644 index b474b3c47..000000000 --- a/clients/middleware/logging.go +++ /dev/null @@ -1,280 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package middleware - -import ( - "context" - "fmt" - "log/slog" - "time" - - "github.com/absmach/magistrala/clients" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/roles" - rolemw "github.com/absmach/magistrala/pkg/roles/rolemanager/middleware" - "github.com/go-chi/chi/v5/middleware" -) - -var _ clients.Service = (*loggingMiddleware)(nil) - -type loggingMiddleware struct { - logger *slog.Logger - svc clients.Service - rolemw.RoleManagerLoggingMiddleware -} - -// NewLogging adds logging facilities to the core service. -func NewLogging(svc clients.Service, logger *slog.Logger) clients.Service { - return &loggingMiddleware{ - logger: logger, - svc: svc, - RoleManagerLoggingMiddleware: rolemw.NewLogging("clients", svc, logger), - } -} - -func (lm *loggingMiddleware) CreateClients(ctx context.Context, session authn.Session, clients ...clients.Client) (cs []clients.Client, rps []roles.RoleProvision, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("domain_id", session.DomainID), - slog.String("request_id", middleware.GetReqID(ctx)), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn(fmt.Sprintf("Create %d clients failed", len(clients)), args...) - return - } - lm.logger.Info(fmt.Sprintf("Create %d clients completed successfully", len(clients)), args...) - }(time.Now()) - return lm.svc.CreateClients(ctx, session, clients...) -} - -func (lm *loggingMiddleware) View(ctx context.Context, session authn.Session, id string, withRoles bool) (c clients.Client, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("domain_id", session.DomainID), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.Group("client", - slog.String("id", c.ID), - slog.String("name", c.Name), - slog.Bool("with_roles", withRoles), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("View client failed", args...) - return - } - lm.logger.Info("View client completed successfully", args...) - }(time.Now()) - return lm.svc.View(ctx, session, id, withRoles) -} - -func (lm *loggingMiddleware) ListClients(ctx context.Context, session authn.Session, pm clients.Page) (cp clients.ClientsPage, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("domain_id", session.DomainID), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.Group("page", - slog.Uint64("limit", pm.Limit), - slog.Uint64("offset", pm.Offset), - slog.Uint64("total", cp.Total), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("List clients failed", args...) - return - } - lm.logger.Info("List clients completed successfully", args...) - }(time.Now()) - return lm.svc.ListClients(ctx, session, pm) -} - -func (lm *loggingMiddleware) ListUserClients(ctx context.Context, session authn.Session, userID string, pm clients.Page) (cp clients.ClientsPage, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("domain_id", session.DomainID), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.String("user_id", userID), - slog.Group("page", - slog.Uint64("limit", pm.Limit), - slog.Uint64("offset", pm.Offset), - slog.Uint64("total", cp.Total), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("List clients failed", args...) - return - } - lm.logger.Info("List clients completed successfully", args...) - }(time.Now()) - return lm.svc.ListUserClients(ctx, session, userID, pm) -} - -func (lm *loggingMiddleware) Update(ctx context.Context, session authn.Session, client clients.Client) (c clients.Client, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("domain_id", session.DomainID), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.Group("client", - slog.String("id", client.ID), - slog.String("name", client.Name), - slog.Any("metadata", client.Metadata), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Update client failed", args...) - return - } - lm.logger.Info("Update client completed successfully", args...) - }(time.Now()) - return lm.svc.Update(ctx, session, client) -} - -func (lm *loggingMiddleware) UpdateTags(ctx context.Context, session authn.Session, client clients.Client) (c clients.Client, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("domain_id", session.DomainID), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.Group("client", - slog.String("id", c.ID), - slog.String("name", c.Name), - slog.Any("tags", c.Tags), - ), - } - if err != nil { - args := append(args, slog.String("error", err.Error())) - lm.logger.Warn("Update client tags failed", args...) - return - } - lm.logger.Info("Update client tags completed successfully", args...) - }(time.Now()) - return lm.svc.UpdateTags(ctx, session, client) -} - -func (lm *loggingMiddleware) UpdateSecret(ctx context.Context, session authn.Session, oldSecret, newSecret string) (c clients.Client, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("domain_id", session.DomainID), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.Group("client", - slog.String("id", c.ID), - slog.String("name", c.Name), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Update client secret failed", args...) - return - } - lm.logger.Info("Update client secret completed successfully", args...) - }(time.Now()) - return lm.svc.UpdateSecret(ctx, session, oldSecret, newSecret) -} - -func (lm *loggingMiddleware) Enable(ctx context.Context, session authn.Session, id string) (c clients.Client, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("domain_id", session.DomainID), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.Group("client", - slog.String("id", id), - slog.String("name", c.Name), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Enable client failed", args...) - return - } - lm.logger.Info("Enable client completed successfully", args...) - }(time.Now()) - return lm.svc.Enable(ctx, session, id) -} - -func (lm *loggingMiddleware) Disable(ctx context.Context, session authn.Session, id string) (c clients.Client, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("domain_id", session.DomainID), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.Group("client", - slog.String("id", id), - slog.String("name", c.Name), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Disable client failed", args...) - return - } - lm.logger.Info("Disable client completed successfully", args...) - }(time.Now()) - return lm.svc.Disable(ctx, session, id) -} - -func (lm *loggingMiddleware) Delete(ctx context.Context, session authn.Session, id string) (err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("domain_id", session.DomainID), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.String("client_id", id), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Delete client failed", args...) - return - } - lm.logger.Info("Delete client completed successfully", args...) - }(time.Now()) - return lm.svc.Delete(ctx, session, id) -} - -func (lm *loggingMiddleware) SetParentGroup(ctx context.Context, session authn.Session, parentGroupID string, id string) (err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("domain_id", session.DomainID), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.String("parent_group_id", parentGroupID), - slog.String("client_id", id), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Set parent group to client failed", args...) - return - } - lm.logger.Info("Set parent group to client completed successfully", args...) - }(time.Now()) - return lm.svc.SetParentGroup(ctx, session, parentGroupID, id) -} - -func (lm *loggingMiddleware) RemoveParentGroup(ctx context.Context, session authn.Session, id string) (err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("domain_id", session.DomainID), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.String("client_id", id), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Remove parent group from client failed", args...) - return - } - lm.logger.Info("Remove parent group from client completed successfully", args...) - }(time.Now()) - return lm.svc.RemoveParentGroup(ctx, session, id) -} diff --git a/clients/middleware/metrics.go b/clients/middleware/metrics.go deleted file mode 100644 index fe501a2d3..000000000 --- a/clients/middleware/metrics.go +++ /dev/null @@ -1,130 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package middleware - -import ( - "context" - "time" - - "github.com/absmach/magistrala/clients" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/roles" - rolemw "github.com/absmach/magistrala/pkg/roles/rolemanager/middleware" - "github.com/go-kit/kit/metrics" -) - -var _ clients.Service = (*metricsMiddleware)(nil) - -type metricsMiddleware struct { - counter metrics.Counter - latency metrics.Histogram - svc clients.Service - rolemw.RoleManagerMetricsMiddleware -} - -// NewMetrics returns a new metrics middleware wrapper. -func NewMetrics(svc clients.Service, counter metrics.Counter, latency metrics.Histogram) clients.Service { - return &metricsMiddleware{ - counter: counter, - latency: latency, - svc: svc, - RoleManagerMetricsMiddleware: rolemw.NewMetrics("clients", svc, counter, latency), - } -} - -func (ms *metricsMiddleware) CreateClients(ctx context.Context, session authn.Session, clients ...clients.Client) ([]clients.Client, []roles.RoleProvision, error) { - defer func(begin time.Time) { - ms.counter.With("method", "register_clients").Add(1) - ms.latency.With("method", "register_clients").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.CreateClients(ctx, session, clients...) -} - -func (ms *metricsMiddleware) View(ctx context.Context, session authn.Session, id string, withRoles bool) (clients.Client, error) { - defer func(begin time.Time) { - ms.counter.With("method", "view_client").Add(1) - ms.latency.With("method", "view_client").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.View(ctx, session, id, withRoles) -} - -func (ms *metricsMiddleware) ListClients(ctx context.Context, session authn.Session, pm clients.Page) (clients.ClientsPage, error) { - defer func(begin time.Time) { - ms.counter.With("method", "list_clients").Add(1) - ms.latency.With("method", "list_clients").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.ListClients(ctx, session, pm) -} - -func (ms *metricsMiddleware) ListUserClients(ctx context.Context, session authn.Session, userID string, pm clients.Page) (clients.ClientsPage, error) { - defer func(begin time.Time) { - ms.counter.With("method", "list_user_clients").Add(1) - ms.latency.With("method", "list_user_clients").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.ListUserClients(ctx, session, userID, pm) -} - -func (ms *metricsMiddleware) Update(ctx context.Context, session authn.Session, client clients.Client) (clients.Client, error) { - defer func(begin time.Time) { - ms.counter.With("method", "update_client").Add(1) - ms.latency.With("method", "update_client").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.Update(ctx, session, client) -} - -func (ms *metricsMiddleware) UpdateTags(ctx context.Context, session authn.Session, client clients.Client) (clients.Client, error) { - defer func(begin time.Time) { - ms.counter.With("method", "update_client_tags").Add(1) - ms.latency.With("method", "update_client_tags").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.UpdateTags(ctx, session, client) -} - -func (ms *metricsMiddleware) UpdateSecret(ctx context.Context, session authn.Session, oldSecret, newSecret string) (clients.Client, error) { - defer func(begin time.Time) { - ms.counter.With("method", "update_client_secret").Add(1) - ms.latency.With("method", "update_client_secret").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.UpdateSecret(ctx, session, oldSecret, newSecret) -} - -func (ms *metricsMiddleware) Enable(ctx context.Context, session authn.Session, id string) (clients.Client, error) { - defer func(begin time.Time) { - ms.counter.With("method", "enable_client").Add(1) - ms.latency.With("method", "enable_client").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.Enable(ctx, session, id) -} - -func (ms *metricsMiddleware) Disable(ctx context.Context, session authn.Session, id string) (clients.Client, error) { - defer func(begin time.Time) { - ms.counter.With("method", "disable_client").Add(1) - ms.latency.With("method", "disable_client").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.Disable(ctx, session, id) -} - -func (ms *metricsMiddleware) Delete(ctx context.Context, session authn.Session, id string) error { - defer func(begin time.Time) { - ms.counter.With("method", "delete_client").Add(1) - ms.latency.With("method", "delete_client").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.Delete(ctx, session, id) -} - -func (ms *metricsMiddleware) SetParentGroup(ctx context.Context, session authn.Session, parentGroupID string, id string) (err error) { - defer func(begin time.Time) { - ms.counter.With("method", "set_parent_group").Add(1) - ms.latency.With("method", "set_parent_group").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.SetParentGroup(ctx, session, parentGroupID, id) -} - -func (ms *metricsMiddleware) RemoveParentGroup(ctx context.Context, session authn.Session, id string) (err error) { - defer func(begin time.Time) { - ms.counter.With("method", "remove_parent_group").Add(1) - ms.latency.With("method", "remove_parent_group").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.RemoveParentGroup(ctx, session, id) -} diff --git a/clients/middleware/tracing.go b/clients/middleware/tracing.go deleted file mode 100644 index 471d92f8a..000000000 --- a/clients/middleware/tracing.go +++ /dev/null @@ -1,126 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package middleware - -import ( - "context" - - "github.com/absmach/magistrala/clients" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/roles" - rolemw "github.com/absmach/magistrala/pkg/roles/rolemanager/middleware" - "github.com/absmach/magistrala/pkg/tracing" - "go.opentelemetry.io/otel/attribute" - "go.opentelemetry.io/otel/trace" -) - -var _ clients.Service = (*tracingMiddleware)(nil) - -type tracingMiddleware struct { - tracer trace.Tracer - svc clients.Service - rolemw.RoleManagerTracing -} - -// NewTracing returns a new clients service with tracing capabilities. -func NewTracing(svc clients.Service, tracer trace.Tracer) clients.Service { - return &tracingMiddleware{ - tracer: tracer, - svc: svc, - RoleManagerTracing: rolemw.NewTracing("group", svc, tracer), - } -} - -// CreateClients traces the "CreateClients" operation of the wrapped clients.Service. -func (tm *tracingMiddleware) CreateClients(ctx context.Context, session authn.Session, cli ...clients.Client) ([]clients.Client, []roles.RoleProvision, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_create_client") - defer span.End() - - return tm.svc.CreateClients(ctx, session, cli...) -} - -// View traces the "View" operation of the wrapped clients.Service. -func (tm *tracingMiddleware) View(ctx context.Context, session authn.Session, id string, withRoles bool) (clients.Client, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_view_client", trace.WithAttributes(attribute.String("id", id), attribute.Bool("with_roles", withRoles))) - defer span.End() - return tm.svc.View(ctx, session, id, withRoles) -} - -// ListClients traces the "ListClients" operation of the wrapped clients.Service. -func (tm *tracingMiddleware) ListClients(ctx context.Context, session authn.Session, pm clients.Page) (clients.ClientsPage, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_list_clients") - defer span.End() - return tm.svc.ListClients(ctx, session, pm) -} - -func (tm *tracingMiddleware) ListUserClients(ctx context.Context, session authn.Session, userID string, pm clients.Page) (clients.ClientsPage, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_list_clients") - defer span.End() - return tm.svc.ListUserClients(ctx, session, userID, pm) -} - -// Update traces the "Update" operation of the wrapped clients.Service. -func (tm *tracingMiddleware) Update(ctx context.Context, session authn.Session, cli clients.Client) (clients.Client, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_update_client", trace.WithAttributes(attribute.String("id", cli.ID))) - defer span.End() - - return tm.svc.Update(ctx, session, cli) -} - -// UpdateTags traces the "UpdateTags" operation of the wrapped clients.Service. -func (tm *tracingMiddleware) UpdateTags(ctx context.Context, session authn.Session, cli clients.Client) (clients.Client, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_update_client_tags", trace.WithAttributes( - attribute.String("id", cli.ID), - attribute.StringSlice("tags", cli.Tags), - )) - defer span.End() - - return tm.svc.UpdateTags(ctx, session, cli) -} - -// UpdateSecret traces the "UpdateSecret" operation of the wrapped clients.Service. -func (tm *tracingMiddleware) UpdateSecret(ctx context.Context, session authn.Session, oldSecret, newSecret string) (clients.Client, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_update_client_secret") - defer span.End() - - return tm.svc.UpdateSecret(ctx, session, oldSecret, newSecret) -} - -// Enable traces the "Enable" operation of the wrapped clients.Service. -func (tm *tracingMiddleware) Enable(ctx context.Context, session authn.Session, id string) (clients.Client, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_enable_client", trace.WithAttributes(attribute.String("id", id))) - defer span.End() - - return tm.svc.Enable(ctx, session, id) -} - -// Disable traces the "Disable" operation of the wrapped clients.Service. -func (tm *tracingMiddleware) Disable(ctx context.Context, session authn.Session, id string) (clients.Client, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_disable_client", trace.WithAttributes(attribute.String("id", id))) - defer span.End() - - return tm.svc.Disable(ctx, session, id) -} - -// Delete traces the "Delete" operation of the wrapped clients.Service. -func (tm *tracingMiddleware) Delete(ctx context.Context, session authn.Session, id string) error { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "delete_client", trace.WithAttributes(attribute.String("id", id))) - defer span.End() - return tm.svc.Delete(ctx, session, id) -} - -func (tm *tracingMiddleware) SetParentGroup(ctx context.Context, session authn.Session, parentGroupID string, id string) error { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "set_parent_group", trace.WithAttributes( - attribute.String("id", id), - attribute.String("parent_group_id", parentGroupID), - )) - defer span.End() - return tm.svc.SetParentGroup(ctx, session, parentGroupID, id) -} - -func (tm *tracingMiddleware) RemoveParentGroup(ctx context.Context, session authn.Session, id string) error { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "remove_parent_group", trace.WithAttributes(attribute.String("id", id))) - defer span.End() - return tm.svc.RemoveParentGroup(ctx, session, id) -} diff --git a/clients/mocks/cache.go b/clients/mocks/cache.go deleted file mode 100644 index 73db55558..000000000 --- a/clients/mocks/cache.go +++ /dev/null @@ -1,228 +0,0 @@ -// Copyright (c) Abstract Machines - -// SPDX-License-Identifier: Apache-2.0 - -// Code generated by mockery; DO NOT EDIT. -// github.com/vektra/mockery -// template: testify - -package mocks - -import ( - "context" - - mock "github.com/stretchr/testify/mock" -) - -// NewCache creates a new instance of Cache. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. -// The first argument is typically a *testing.T value. -func NewCache(t interface { - mock.TestingT - Cleanup(func()) -}) *Cache { - mock := &Cache{} - mock.Mock.Test(t) - - t.Cleanup(func() { mock.AssertExpectations(t) }) - - return mock -} - -// Cache is an autogenerated mock type for the Cache type -type Cache struct { - mock.Mock -} - -type Cache_Expecter struct { - mock *mock.Mock -} - -func (_m *Cache) EXPECT() *Cache_Expecter { - return &Cache_Expecter{mock: &_m.Mock} -} - -// ID provides a mock function for the type Cache -func (_mock *Cache) ID(ctx context.Context, clientSecret string) (string, error) { - ret := _mock.Called(ctx, clientSecret) - - if len(ret) == 0 { - panic("no return value specified for ID") - } - - var r0 string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) (string, error)); ok { - return returnFunc(ctx, clientSecret) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string) string); ok { - r0 = returnFunc(ctx, clientSecret) - } else { - r0 = ret.Get(0).(string) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = returnFunc(ctx, clientSecret) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Cache_ID_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ID' -type Cache_ID_Call struct { - *mock.Call -} - -// ID is a helper method to define mock.On call -// - ctx context.Context -// - clientSecret string -func (_e *Cache_Expecter) ID(ctx interface{}, clientSecret interface{}) *Cache_ID_Call { - return &Cache_ID_Call{Call: _e.mock.On("ID", ctx, clientSecret)} -} - -func (_c *Cache_ID_Call) Run(run func(ctx context.Context, clientSecret string)) *Cache_ID_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Cache_ID_Call) Return(s string, err error) *Cache_ID_Call { - _c.Call.Return(s, err) - return _c -} - -func (_c *Cache_ID_Call) RunAndReturn(run func(ctx context.Context, clientSecret string) (string, error)) *Cache_ID_Call { - _c.Call.Return(run) - return _c -} - -// Remove provides a mock function for the type Cache -func (_mock *Cache) Remove(ctx context.Context, clientID string) error { - ret := _mock.Called(ctx, clientID) - - if len(ret) == 0 { - panic("no return value specified for Remove") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) error); ok { - r0 = returnFunc(ctx, clientID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Cache_Remove_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Remove' -type Cache_Remove_Call struct { - *mock.Call -} - -// Remove is a helper method to define mock.On call -// - ctx context.Context -// - clientID string -func (_e *Cache_Expecter) Remove(ctx interface{}, clientID interface{}) *Cache_Remove_Call { - return &Cache_Remove_Call{Call: _e.mock.On("Remove", ctx, clientID)} -} - -func (_c *Cache_Remove_Call) Run(run func(ctx context.Context, clientID string)) *Cache_Remove_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Cache_Remove_Call) Return(err error) *Cache_Remove_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Cache_Remove_Call) RunAndReturn(run func(ctx context.Context, clientID string) error) *Cache_Remove_Call { - _c.Call.Return(run) - return _c -} - -// Save provides a mock function for the type Cache -func (_mock *Cache) Save(ctx context.Context, clientSecret string, clientID string) error { - ret := _mock.Called(ctx, clientSecret, clientID) - - if len(ret) == 0 { - panic("no return value specified for Save") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) error); ok { - r0 = returnFunc(ctx, clientSecret, clientID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Cache_Save_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Save' -type Cache_Save_Call struct { - *mock.Call -} - -// Save is a helper method to define mock.On call -// - ctx context.Context -// - clientSecret string -// - clientID string -func (_e *Cache_Expecter) Save(ctx interface{}, clientSecret interface{}, clientID interface{}) *Cache_Save_Call { - return &Cache_Save_Call{Call: _e.mock.On("Save", ctx, clientSecret, clientID)} -} - -func (_c *Cache_Save_Call) Run(run func(ctx context.Context, clientSecret string, clientID string)) *Cache_Save_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Cache_Save_Call) Return(err error) *Cache_Save_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Cache_Save_Call) RunAndReturn(run func(ctx context.Context, clientSecret string, clientID string) error) *Cache_Save_Call { - _c.Call.Return(run) - return _c -} diff --git a/clients/mocks/clients_client.go b/clients/mocks/clients_client.go deleted file mode 100644 index a5fa66de6..000000000 --- a/clients/mocks/clients_client.go +++ /dev/null @@ -1,626 +0,0 @@ -// Copyright (c) Abstract Machines - -// SPDX-License-Identifier: Apache-2.0 - -// Code generated by mockery; DO NOT EDIT. -// github.com/vektra/mockery -// template: testify - -package mocks - -import ( - "context" - - v10 "github.com/absmach/magistrala/api/grpc/clients/v1" - "github.com/absmach/magistrala/api/grpc/common/v1" - mock "github.com/stretchr/testify/mock" - "google.golang.org/grpc" -) - -// NewClientsServiceClient creates a new instance of ClientsServiceClient. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. -// The first argument is typically a *testing.T value. -func NewClientsServiceClient(t interface { - mock.TestingT - Cleanup(func()) -}) *ClientsServiceClient { - mock := &ClientsServiceClient{} - mock.Mock.Test(t) - - t.Cleanup(func() { mock.AssertExpectations(t) }) - - return mock -} - -// ClientsServiceClient is an autogenerated mock type for the ClientsServiceClient type -type ClientsServiceClient struct { - mock.Mock -} - -type ClientsServiceClient_Expecter struct { - mock *mock.Mock -} - -func (_m *ClientsServiceClient) EXPECT() *ClientsServiceClient_Expecter { - return &ClientsServiceClient_Expecter{mock: &_m.Mock} -} - -// AddConnections provides a mock function for the type ClientsServiceClient -func (_mock *ClientsServiceClient) AddConnections(ctx context.Context, in *v1.AddConnectionsReq, opts ...grpc.CallOption) (*v1.AddConnectionsRes, error) { - var tmpRet mock.Arguments - if len(opts) > 0 { - tmpRet = _mock.Called(ctx, in, opts) - } else { - tmpRet = _mock.Called(ctx, in) - } - ret := tmpRet - - if len(ret) == 0 { - panic("no return value specified for AddConnections") - } - - var r0 *v1.AddConnectionsRes - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, *v1.AddConnectionsReq, ...grpc.CallOption) (*v1.AddConnectionsRes, error)); ok { - return returnFunc(ctx, in, opts...) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, *v1.AddConnectionsReq, ...grpc.CallOption) *v1.AddConnectionsRes); ok { - r0 = returnFunc(ctx, in, opts...) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).(*v1.AddConnectionsRes) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, *v1.AddConnectionsReq, ...grpc.CallOption) error); ok { - r1 = returnFunc(ctx, in, opts...) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// ClientsServiceClient_AddConnections_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'AddConnections' -type ClientsServiceClient_AddConnections_Call struct { - *mock.Call -} - -// AddConnections is a helper method to define mock.On call -// - ctx context.Context -// - in *v1.AddConnectionsReq -// - opts ...grpc.CallOption -func (_e *ClientsServiceClient_Expecter) AddConnections(ctx interface{}, in interface{}, opts ...interface{}) *ClientsServiceClient_AddConnections_Call { - return &ClientsServiceClient_AddConnections_Call{Call: _e.mock.On("AddConnections", - append([]interface{}{ctx, in}, opts...)...)} -} - -func (_c *ClientsServiceClient_AddConnections_Call) Run(run func(ctx context.Context, in *v1.AddConnectionsReq, opts ...grpc.CallOption)) *ClientsServiceClient_AddConnections_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 *v1.AddConnectionsReq - if args[1] != nil { - arg1 = args[1].(*v1.AddConnectionsReq) - } - var arg2 []grpc.CallOption - var variadicArgs []grpc.CallOption - if len(args) > 2 { - variadicArgs = args[2].([]grpc.CallOption) - } - arg2 = variadicArgs - run( - arg0, - arg1, - arg2..., - ) - }) - return _c -} - -func (_c *ClientsServiceClient_AddConnections_Call) Return(addConnectionsRes *v1.AddConnectionsRes, err error) *ClientsServiceClient_AddConnections_Call { - _c.Call.Return(addConnectionsRes, err) - return _c -} - -func (_c *ClientsServiceClient_AddConnections_Call) RunAndReturn(run func(ctx context.Context, in *v1.AddConnectionsReq, opts ...grpc.CallOption) (*v1.AddConnectionsRes, error)) *ClientsServiceClient_AddConnections_Call { - _c.Call.Return(run) - return _c -} - -// Authenticate provides a mock function for the type ClientsServiceClient -func (_mock *ClientsServiceClient) Authenticate(ctx context.Context, in *v10.AuthnReq, opts ...grpc.CallOption) (*v10.AuthnRes, error) { - var tmpRet mock.Arguments - if len(opts) > 0 { - tmpRet = _mock.Called(ctx, in, opts) - } else { - tmpRet = _mock.Called(ctx, in) - } - ret := tmpRet - - if len(ret) == 0 { - panic("no return value specified for Authenticate") - } - - var r0 *v10.AuthnRes - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, *v10.AuthnReq, ...grpc.CallOption) (*v10.AuthnRes, error)); ok { - return returnFunc(ctx, in, opts...) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, *v10.AuthnReq, ...grpc.CallOption) *v10.AuthnRes); ok { - r0 = returnFunc(ctx, in, opts...) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).(*v10.AuthnRes) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, *v10.AuthnReq, ...grpc.CallOption) error); ok { - r1 = returnFunc(ctx, in, opts...) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// ClientsServiceClient_Authenticate_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Authenticate' -type ClientsServiceClient_Authenticate_Call struct { - *mock.Call -} - -// Authenticate is a helper method to define mock.On call -// - ctx context.Context -// - in *v10.AuthnReq -// - opts ...grpc.CallOption -func (_e *ClientsServiceClient_Expecter) Authenticate(ctx interface{}, in interface{}, opts ...interface{}) *ClientsServiceClient_Authenticate_Call { - return &ClientsServiceClient_Authenticate_Call{Call: _e.mock.On("Authenticate", - append([]interface{}{ctx, in}, opts...)...)} -} - -func (_c *ClientsServiceClient_Authenticate_Call) Run(run func(ctx context.Context, in *v10.AuthnReq, opts ...grpc.CallOption)) *ClientsServiceClient_Authenticate_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 *v10.AuthnReq - if args[1] != nil { - arg1 = args[1].(*v10.AuthnReq) - } - var arg2 []grpc.CallOption - var variadicArgs []grpc.CallOption - if len(args) > 2 { - variadicArgs = args[2].([]grpc.CallOption) - } - arg2 = variadicArgs - run( - arg0, - arg1, - arg2..., - ) - }) - return _c -} - -func (_c *ClientsServiceClient_Authenticate_Call) Return(authnRes *v10.AuthnRes, err error) *ClientsServiceClient_Authenticate_Call { - _c.Call.Return(authnRes, err) - return _c -} - -func (_c *ClientsServiceClient_Authenticate_Call) RunAndReturn(run func(ctx context.Context, in *v10.AuthnReq, opts ...grpc.CallOption) (*v10.AuthnRes, error)) *ClientsServiceClient_Authenticate_Call { - _c.Call.Return(run) - return _c -} - -// RemoveChannelConnections provides a mock function for the type ClientsServiceClient -func (_mock *ClientsServiceClient) RemoveChannelConnections(ctx context.Context, in *v10.RemoveChannelConnectionsReq, opts ...grpc.CallOption) (*v10.RemoveChannelConnectionsRes, error) { - var tmpRet mock.Arguments - if len(opts) > 0 { - tmpRet = _mock.Called(ctx, in, opts) - } else { - tmpRet = _mock.Called(ctx, in) - } - ret := tmpRet - - if len(ret) == 0 { - panic("no return value specified for RemoveChannelConnections") - } - - var r0 *v10.RemoveChannelConnectionsRes - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, *v10.RemoveChannelConnectionsReq, ...grpc.CallOption) (*v10.RemoveChannelConnectionsRes, error)); ok { - return returnFunc(ctx, in, opts...) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, *v10.RemoveChannelConnectionsReq, ...grpc.CallOption) *v10.RemoveChannelConnectionsRes); ok { - r0 = returnFunc(ctx, in, opts...) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).(*v10.RemoveChannelConnectionsRes) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, *v10.RemoveChannelConnectionsReq, ...grpc.CallOption) error); ok { - r1 = returnFunc(ctx, in, opts...) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// ClientsServiceClient_RemoveChannelConnections_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveChannelConnections' -type ClientsServiceClient_RemoveChannelConnections_Call struct { - *mock.Call -} - -// RemoveChannelConnections is a helper method to define mock.On call -// - ctx context.Context -// - in *v10.RemoveChannelConnectionsReq -// - opts ...grpc.CallOption -func (_e *ClientsServiceClient_Expecter) RemoveChannelConnections(ctx interface{}, in interface{}, opts ...interface{}) *ClientsServiceClient_RemoveChannelConnections_Call { - return &ClientsServiceClient_RemoveChannelConnections_Call{Call: _e.mock.On("RemoveChannelConnections", - append([]interface{}{ctx, in}, opts...)...)} -} - -func (_c *ClientsServiceClient_RemoveChannelConnections_Call) Run(run func(ctx context.Context, in *v10.RemoveChannelConnectionsReq, opts ...grpc.CallOption)) *ClientsServiceClient_RemoveChannelConnections_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 *v10.RemoveChannelConnectionsReq - if args[1] != nil { - arg1 = args[1].(*v10.RemoveChannelConnectionsReq) - } - var arg2 []grpc.CallOption - var variadicArgs []grpc.CallOption - if len(args) > 2 { - variadicArgs = args[2].([]grpc.CallOption) - } - arg2 = variadicArgs - run( - arg0, - arg1, - arg2..., - ) - }) - return _c -} - -func (_c *ClientsServiceClient_RemoveChannelConnections_Call) Return(removeChannelConnectionsRes *v10.RemoveChannelConnectionsRes, err error) *ClientsServiceClient_RemoveChannelConnections_Call { - _c.Call.Return(removeChannelConnectionsRes, err) - return _c -} - -func (_c *ClientsServiceClient_RemoveChannelConnections_Call) RunAndReturn(run func(ctx context.Context, in *v10.RemoveChannelConnectionsReq, opts ...grpc.CallOption) (*v10.RemoveChannelConnectionsRes, error)) *ClientsServiceClient_RemoveChannelConnections_Call { - _c.Call.Return(run) - return _c -} - -// RemoveConnections provides a mock function for the type ClientsServiceClient -func (_mock *ClientsServiceClient) RemoveConnections(ctx context.Context, in *v1.RemoveConnectionsReq, opts ...grpc.CallOption) (*v1.RemoveConnectionsRes, error) { - var tmpRet mock.Arguments - if len(opts) > 0 { - tmpRet = _mock.Called(ctx, in, opts) - } else { - tmpRet = _mock.Called(ctx, in) - } - ret := tmpRet - - if len(ret) == 0 { - panic("no return value specified for RemoveConnections") - } - - var r0 *v1.RemoveConnectionsRes - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, *v1.RemoveConnectionsReq, ...grpc.CallOption) (*v1.RemoveConnectionsRes, error)); ok { - return returnFunc(ctx, in, opts...) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, *v1.RemoveConnectionsReq, ...grpc.CallOption) *v1.RemoveConnectionsRes); ok { - r0 = returnFunc(ctx, in, opts...) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).(*v1.RemoveConnectionsRes) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, *v1.RemoveConnectionsReq, ...grpc.CallOption) error); ok { - r1 = returnFunc(ctx, in, opts...) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// ClientsServiceClient_RemoveConnections_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveConnections' -type ClientsServiceClient_RemoveConnections_Call struct { - *mock.Call -} - -// RemoveConnections is a helper method to define mock.On call -// - ctx context.Context -// - in *v1.RemoveConnectionsReq -// - opts ...grpc.CallOption -func (_e *ClientsServiceClient_Expecter) RemoveConnections(ctx interface{}, in interface{}, opts ...interface{}) *ClientsServiceClient_RemoveConnections_Call { - return &ClientsServiceClient_RemoveConnections_Call{Call: _e.mock.On("RemoveConnections", - append([]interface{}{ctx, in}, opts...)...)} -} - -func (_c *ClientsServiceClient_RemoveConnections_Call) Run(run func(ctx context.Context, in *v1.RemoveConnectionsReq, opts ...grpc.CallOption)) *ClientsServiceClient_RemoveConnections_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 *v1.RemoveConnectionsReq - if args[1] != nil { - arg1 = args[1].(*v1.RemoveConnectionsReq) - } - var arg2 []grpc.CallOption - var variadicArgs []grpc.CallOption - if len(args) > 2 { - variadicArgs = args[2].([]grpc.CallOption) - } - arg2 = variadicArgs - run( - arg0, - arg1, - arg2..., - ) - }) - return _c -} - -func (_c *ClientsServiceClient_RemoveConnections_Call) Return(removeConnectionsRes *v1.RemoveConnectionsRes, err error) *ClientsServiceClient_RemoveConnections_Call { - _c.Call.Return(removeConnectionsRes, err) - return _c -} - -func (_c *ClientsServiceClient_RemoveConnections_Call) RunAndReturn(run func(ctx context.Context, in *v1.RemoveConnectionsReq, opts ...grpc.CallOption) (*v1.RemoveConnectionsRes, error)) *ClientsServiceClient_RemoveConnections_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveEntities provides a mock function for the type ClientsServiceClient -func (_mock *ClientsServiceClient) RetrieveEntities(ctx context.Context, in *v1.RetrieveEntitiesReq, opts ...grpc.CallOption) (*v1.RetrieveEntitiesRes, error) { - var tmpRet mock.Arguments - if len(opts) > 0 { - tmpRet = _mock.Called(ctx, in, opts) - } else { - tmpRet = _mock.Called(ctx, in) - } - ret := tmpRet - - if len(ret) == 0 { - panic("no return value specified for RetrieveEntities") - } - - var r0 *v1.RetrieveEntitiesRes - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, *v1.RetrieveEntitiesReq, ...grpc.CallOption) (*v1.RetrieveEntitiesRes, error)); ok { - return returnFunc(ctx, in, opts...) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, *v1.RetrieveEntitiesReq, ...grpc.CallOption) *v1.RetrieveEntitiesRes); ok { - r0 = returnFunc(ctx, in, opts...) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).(*v1.RetrieveEntitiesRes) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, *v1.RetrieveEntitiesReq, ...grpc.CallOption) error); ok { - r1 = returnFunc(ctx, in, opts...) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// ClientsServiceClient_RetrieveEntities_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveEntities' -type ClientsServiceClient_RetrieveEntities_Call struct { - *mock.Call -} - -// RetrieveEntities is a helper method to define mock.On call -// - ctx context.Context -// - in *v1.RetrieveEntitiesReq -// - opts ...grpc.CallOption -func (_e *ClientsServiceClient_Expecter) RetrieveEntities(ctx interface{}, in interface{}, opts ...interface{}) *ClientsServiceClient_RetrieveEntities_Call { - return &ClientsServiceClient_RetrieveEntities_Call{Call: _e.mock.On("RetrieveEntities", - append([]interface{}{ctx, in}, opts...)...)} -} - -func (_c *ClientsServiceClient_RetrieveEntities_Call) Run(run func(ctx context.Context, in *v1.RetrieveEntitiesReq, opts ...grpc.CallOption)) *ClientsServiceClient_RetrieveEntities_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 *v1.RetrieveEntitiesReq - if args[1] != nil { - arg1 = args[1].(*v1.RetrieveEntitiesReq) - } - var arg2 []grpc.CallOption - var variadicArgs []grpc.CallOption - if len(args) > 2 { - variadicArgs = args[2].([]grpc.CallOption) - } - arg2 = variadicArgs - run( - arg0, - arg1, - arg2..., - ) - }) - return _c -} - -func (_c *ClientsServiceClient_RetrieveEntities_Call) Return(retrieveEntitiesRes *v1.RetrieveEntitiesRes, err error) *ClientsServiceClient_RetrieveEntities_Call { - _c.Call.Return(retrieveEntitiesRes, err) - return _c -} - -func (_c *ClientsServiceClient_RetrieveEntities_Call) RunAndReturn(run func(ctx context.Context, in *v1.RetrieveEntitiesReq, opts ...grpc.CallOption) (*v1.RetrieveEntitiesRes, error)) *ClientsServiceClient_RetrieveEntities_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveEntity provides a mock function for the type ClientsServiceClient -func (_mock *ClientsServiceClient) RetrieveEntity(ctx context.Context, in *v1.RetrieveEntityReq, opts ...grpc.CallOption) (*v1.RetrieveEntityRes, error) { - var tmpRet mock.Arguments - if len(opts) > 0 { - tmpRet = _mock.Called(ctx, in, opts) - } else { - tmpRet = _mock.Called(ctx, in) - } - ret := tmpRet - - if len(ret) == 0 { - panic("no return value specified for RetrieveEntity") - } - - var r0 *v1.RetrieveEntityRes - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, *v1.RetrieveEntityReq, ...grpc.CallOption) (*v1.RetrieveEntityRes, error)); ok { - return returnFunc(ctx, in, opts...) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, *v1.RetrieveEntityReq, ...grpc.CallOption) *v1.RetrieveEntityRes); ok { - r0 = returnFunc(ctx, in, opts...) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).(*v1.RetrieveEntityRes) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, *v1.RetrieveEntityReq, ...grpc.CallOption) error); ok { - r1 = returnFunc(ctx, in, opts...) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// ClientsServiceClient_RetrieveEntity_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveEntity' -type ClientsServiceClient_RetrieveEntity_Call struct { - *mock.Call -} - -// RetrieveEntity is a helper method to define mock.On call -// - ctx context.Context -// - in *v1.RetrieveEntityReq -// - opts ...grpc.CallOption -func (_e *ClientsServiceClient_Expecter) RetrieveEntity(ctx interface{}, in interface{}, opts ...interface{}) *ClientsServiceClient_RetrieveEntity_Call { - return &ClientsServiceClient_RetrieveEntity_Call{Call: _e.mock.On("RetrieveEntity", - append([]interface{}{ctx, in}, opts...)...)} -} - -func (_c *ClientsServiceClient_RetrieveEntity_Call) Run(run func(ctx context.Context, in *v1.RetrieveEntityReq, opts ...grpc.CallOption)) *ClientsServiceClient_RetrieveEntity_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 *v1.RetrieveEntityReq - if args[1] != nil { - arg1 = args[1].(*v1.RetrieveEntityReq) - } - var arg2 []grpc.CallOption - var variadicArgs []grpc.CallOption - if len(args) > 2 { - variadicArgs = args[2].([]grpc.CallOption) - } - arg2 = variadicArgs - run( - arg0, - arg1, - arg2..., - ) - }) - return _c -} - -func (_c *ClientsServiceClient_RetrieveEntity_Call) Return(retrieveEntityRes *v1.RetrieveEntityRes, err error) *ClientsServiceClient_RetrieveEntity_Call { - _c.Call.Return(retrieveEntityRes, err) - return _c -} - -func (_c *ClientsServiceClient_RetrieveEntity_Call) RunAndReturn(run func(ctx context.Context, in *v1.RetrieveEntityReq, opts ...grpc.CallOption) (*v1.RetrieveEntityRes, error)) *ClientsServiceClient_RetrieveEntity_Call { - _c.Call.Return(run) - return _c -} - -// UnsetParentGroupFromClient provides a mock function for the type ClientsServiceClient -func (_mock *ClientsServiceClient) UnsetParentGroupFromClient(ctx context.Context, in *v10.UnsetParentGroupFromClientReq, opts ...grpc.CallOption) (*v10.UnsetParentGroupFromClientRes, error) { - var tmpRet mock.Arguments - if len(opts) > 0 { - tmpRet = _mock.Called(ctx, in, opts) - } else { - tmpRet = _mock.Called(ctx, in) - } - ret := tmpRet - - if len(ret) == 0 { - panic("no return value specified for UnsetParentGroupFromClient") - } - - var r0 *v10.UnsetParentGroupFromClientRes - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, *v10.UnsetParentGroupFromClientReq, ...grpc.CallOption) (*v10.UnsetParentGroupFromClientRes, error)); ok { - return returnFunc(ctx, in, opts...) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, *v10.UnsetParentGroupFromClientReq, ...grpc.CallOption) *v10.UnsetParentGroupFromClientRes); ok { - r0 = returnFunc(ctx, in, opts...) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).(*v10.UnsetParentGroupFromClientRes) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, *v10.UnsetParentGroupFromClientReq, ...grpc.CallOption) error); ok { - r1 = returnFunc(ctx, in, opts...) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// ClientsServiceClient_UnsetParentGroupFromClient_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UnsetParentGroupFromClient' -type ClientsServiceClient_UnsetParentGroupFromClient_Call struct { - *mock.Call -} - -// UnsetParentGroupFromClient is a helper method to define mock.On call -// - ctx context.Context -// - in *v10.UnsetParentGroupFromClientReq -// - opts ...grpc.CallOption -func (_e *ClientsServiceClient_Expecter) UnsetParentGroupFromClient(ctx interface{}, in interface{}, opts ...interface{}) *ClientsServiceClient_UnsetParentGroupFromClient_Call { - return &ClientsServiceClient_UnsetParentGroupFromClient_Call{Call: _e.mock.On("UnsetParentGroupFromClient", - append([]interface{}{ctx, in}, opts...)...)} -} - -func (_c *ClientsServiceClient_UnsetParentGroupFromClient_Call) Run(run func(ctx context.Context, in *v10.UnsetParentGroupFromClientReq, opts ...grpc.CallOption)) *ClientsServiceClient_UnsetParentGroupFromClient_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 *v10.UnsetParentGroupFromClientReq - if args[1] != nil { - arg1 = args[1].(*v10.UnsetParentGroupFromClientReq) - } - var arg2 []grpc.CallOption - var variadicArgs []grpc.CallOption - if len(args) > 2 { - variadicArgs = args[2].([]grpc.CallOption) - } - arg2 = variadicArgs - run( - arg0, - arg1, - arg2..., - ) - }) - return _c -} - -func (_c *ClientsServiceClient_UnsetParentGroupFromClient_Call) Return(unsetParentGroupFromClientRes *v10.UnsetParentGroupFromClientRes, err error) *ClientsServiceClient_UnsetParentGroupFromClient_Call { - _c.Call.Return(unsetParentGroupFromClientRes, err) - return _c -} - -func (_c *ClientsServiceClient_UnsetParentGroupFromClient_Call) RunAndReturn(run func(ctx context.Context, in *v10.UnsetParentGroupFromClientReq, opts ...grpc.CallOption) (*v10.UnsetParentGroupFromClientRes, error)) *ClientsServiceClient_UnsetParentGroupFromClient_Call { - _c.Call.Return(run) - return _c -} diff --git a/clients/mocks/doc.go b/clients/mocks/doc.go deleted file mode 100644 index 16ed198af..000000000 --- a/clients/mocks/doc.go +++ /dev/null @@ -1,5 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package mocks contains mocks for testing purposes. -package mocks diff --git a/clients/mocks/repository.go b/clients/mocks/repository.go deleted file mode 100644 index af6f80c54..000000000 --- a/clients/mocks/repository.go +++ /dev/null @@ -1,2962 +0,0 @@ -// Copyright (c) Abstract Machines - -// SPDX-License-Identifier: Apache-2.0 - -// Code generated by mockery; DO NOT EDIT. -// github.com/vektra/mockery -// template: testify - -package mocks - -import ( - "context" - - "github.com/absmach/magistrala/clients" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/roles" - mock "github.com/stretchr/testify/mock" -) - -// NewRepository creates a new instance of Repository. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. -// The first argument is typically a *testing.T value. -func NewRepository(t interface { - mock.TestingT - Cleanup(func()) -}) *Repository { - mock := &Repository{} - mock.Mock.Test(t) - - t.Cleanup(func() { mock.AssertExpectations(t) }) - - return mock -} - -// Repository is an autogenerated mock type for the Repository type -type Repository struct { - mock.Mock -} - -type Repository_Expecter struct { - mock *mock.Mock -} - -func (_m *Repository) EXPECT() *Repository_Expecter { - return &Repository_Expecter{mock: &_m.Mock} -} - -// AddConnections provides a mock function for the type Repository -func (_mock *Repository) AddConnections(ctx context.Context, conns []clients.Connection) error { - ret := _mock.Called(ctx, conns) - - if len(ret) == 0 { - panic("no return value specified for AddConnections") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, []clients.Connection) error); ok { - r0 = returnFunc(ctx, conns) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_AddConnections_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'AddConnections' -type Repository_AddConnections_Call struct { - *mock.Call -} - -// AddConnections is a helper method to define mock.On call -// - ctx context.Context -// - conns []clients.Connection -func (_e *Repository_Expecter) AddConnections(ctx interface{}, conns interface{}) *Repository_AddConnections_Call { - return &Repository_AddConnections_Call{Call: _e.mock.On("AddConnections", ctx, conns)} -} - -func (_c *Repository_AddConnections_Call) Run(run func(ctx context.Context, conns []clients.Connection)) *Repository_AddConnections_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 []clients.Connection - if args[1] != nil { - arg1 = args[1].([]clients.Connection) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_AddConnections_Call) Return(err error) *Repository_AddConnections_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_AddConnections_Call) RunAndReturn(run func(ctx context.Context, conns []clients.Connection) error) *Repository_AddConnections_Call { - _c.Call.Return(run) - return _c -} - -// AddRoles provides a mock function for the type Repository -func (_mock *Repository) AddRoles(ctx context.Context, rps []roles.RoleProvision) ([]roles.RoleProvision, error) { - ret := _mock.Called(ctx, rps) - - if len(ret) == 0 { - panic("no return value specified for AddRoles") - } - - var r0 []roles.RoleProvision - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, []roles.RoleProvision) ([]roles.RoleProvision, error)); ok { - return returnFunc(ctx, rps) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, []roles.RoleProvision) []roles.RoleProvision); ok { - r0 = returnFunc(ctx, rps) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]roles.RoleProvision) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, []roles.RoleProvision) error); ok { - r1 = returnFunc(ctx, rps) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_AddRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'AddRoles' -type Repository_AddRoles_Call struct { - *mock.Call -} - -// AddRoles is a helper method to define mock.On call -// - ctx context.Context -// - rps []roles.RoleProvision -func (_e *Repository_Expecter) AddRoles(ctx interface{}, rps interface{}) *Repository_AddRoles_Call { - return &Repository_AddRoles_Call{Call: _e.mock.On("AddRoles", ctx, rps)} -} - -func (_c *Repository_AddRoles_Call) Run(run func(ctx context.Context, rps []roles.RoleProvision)) *Repository_AddRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 []roles.RoleProvision - if args[1] != nil { - arg1 = args[1].([]roles.RoleProvision) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_AddRoles_Call) Return(roleProvisions []roles.RoleProvision, err error) *Repository_AddRoles_Call { - _c.Call.Return(roleProvisions, err) - return _c -} - -func (_c *Repository_AddRoles_Call) RunAndReturn(run func(ctx context.Context, rps []roles.RoleProvision) ([]roles.RoleProvision, error)) *Repository_AddRoles_Call { - _c.Call.Return(run) - return _c -} - -// ChangeStatus provides a mock function for the type Repository -func (_mock *Repository) ChangeStatus(ctx context.Context, client clients.Client) (clients.Client, error) { - ret := _mock.Called(ctx, client) - - if len(ret) == 0 { - panic("no return value specified for ChangeStatus") - } - - var r0 clients.Client - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, clients.Client) (clients.Client, error)); ok { - return returnFunc(ctx, client) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, clients.Client) clients.Client); ok { - r0 = returnFunc(ctx, client) - } else { - r0 = ret.Get(0).(clients.Client) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, clients.Client) error); ok { - r1 = returnFunc(ctx, client) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_ChangeStatus_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ChangeStatus' -type Repository_ChangeStatus_Call struct { - *mock.Call -} - -// ChangeStatus is a helper method to define mock.On call -// - ctx context.Context -// - client clients.Client -func (_e *Repository_Expecter) ChangeStatus(ctx interface{}, client interface{}) *Repository_ChangeStatus_Call { - return &Repository_ChangeStatus_Call{Call: _e.mock.On("ChangeStatus", ctx, client)} -} - -func (_c *Repository_ChangeStatus_Call) Run(run func(ctx context.Context, client clients.Client)) *Repository_ChangeStatus_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 clients.Client - if args[1] != nil { - arg1 = args[1].(clients.Client) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_ChangeStatus_Call) Return(client1 clients.Client, err error) *Repository_ChangeStatus_Call { - _c.Call.Return(client1, err) - return _c -} - -func (_c *Repository_ChangeStatus_Call) RunAndReturn(run func(ctx context.Context, client clients.Client) (clients.Client, error)) *Repository_ChangeStatus_Call { - _c.Call.Return(run) - return _c -} - -// ClientConnectionsCount provides a mock function for the type Repository -func (_mock *Repository) ClientConnectionsCount(ctx context.Context, id string) (uint64, error) { - ret := _mock.Called(ctx, id) - - if len(ret) == 0 { - panic("no return value specified for ClientConnectionsCount") - } - - var r0 uint64 - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) (uint64, error)); ok { - return returnFunc(ctx, id) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string) uint64); ok { - r0 = returnFunc(ctx, id) - } else { - r0 = ret.Get(0).(uint64) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = returnFunc(ctx, id) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_ClientConnectionsCount_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ClientConnectionsCount' -type Repository_ClientConnectionsCount_Call struct { - *mock.Call -} - -// ClientConnectionsCount is a helper method to define mock.On call -// - ctx context.Context -// - id string -func (_e *Repository_Expecter) ClientConnectionsCount(ctx interface{}, id interface{}) *Repository_ClientConnectionsCount_Call { - return &Repository_ClientConnectionsCount_Call{Call: _e.mock.On("ClientConnectionsCount", ctx, id)} -} - -func (_c *Repository_ClientConnectionsCount_Call) Run(run func(ctx context.Context, id string)) *Repository_ClientConnectionsCount_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_ClientConnectionsCount_Call) Return(v uint64, err error) *Repository_ClientConnectionsCount_Call { - _c.Call.Return(v, err) - return _c -} - -func (_c *Repository_ClientConnectionsCount_Call) RunAndReturn(run func(ctx context.Context, id string) (uint64, error)) *Repository_ClientConnectionsCount_Call { - _c.Call.Return(run) - return _c -} - -// Delete provides a mock function for the type Repository -func (_mock *Repository) Delete(ctx context.Context, clientIDs ...string) error { - var tmpRet mock.Arguments - if len(clientIDs) > 0 { - tmpRet = _mock.Called(ctx, clientIDs) - } else { - tmpRet = _mock.Called(ctx) - } - ret := tmpRet - - if len(ret) == 0 { - panic("no return value specified for Delete") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, ...string) error); ok { - r0 = returnFunc(ctx, clientIDs...) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_Delete_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Delete' -type Repository_Delete_Call struct { - *mock.Call -} - -// Delete is a helper method to define mock.On call -// - ctx context.Context -// - clientIDs ...string -func (_e *Repository_Expecter) Delete(ctx interface{}, clientIDs ...interface{}) *Repository_Delete_Call { - return &Repository_Delete_Call{Call: _e.mock.On("Delete", - append([]interface{}{ctx}, clientIDs...)...)} -} - -func (_c *Repository_Delete_Call) Run(run func(ctx context.Context, clientIDs ...string)) *Repository_Delete_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 []string - var variadicArgs []string - if len(args) > 1 { - variadicArgs = args[1].([]string) - } - arg1 = variadicArgs - run( - arg0, - arg1..., - ) - }) - return _c -} - -func (_c *Repository_Delete_Call) Return(err error) *Repository_Delete_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_Delete_Call) RunAndReturn(run func(ctx context.Context, clientIDs ...string) error) *Repository_Delete_Call { - _c.Call.Return(run) - return _c -} - -// DoesClientHaveConnections provides a mock function for the type Repository -func (_mock *Repository) DoesClientHaveConnections(ctx context.Context, id string) (bool, error) { - ret := _mock.Called(ctx, id) - - if len(ret) == 0 { - panic("no return value specified for DoesClientHaveConnections") - } - - var r0 bool - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) (bool, error)); ok { - return returnFunc(ctx, id) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string) bool); ok { - r0 = returnFunc(ctx, id) - } else { - r0 = ret.Get(0).(bool) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = returnFunc(ctx, id) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_DoesClientHaveConnections_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'DoesClientHaveConnections' -type Repository_DoesClientHaveConnections_Call struct { - *mock.Call -} - -// DoesClientHaveConnections is a helper method to define mock.On call -// - ctx context.Context -// - id string -func (_e *Repository_Expecter) DoesClientHaveConnections(ctx interface{}, id interface{}) *Repository_DoesClientHaveConnections_Call { - return &Repository_DoesClientHaveConnections_Call{Call: _e.mock.On("DoesClientHaveConnections", ctx, id)} -} - -func (_c *Repository_DoesClientHaveConnections_Call) Run(run func(ctx context.Context, id string)) *Repository_DoesClientHaveConnections_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_DoesClientHaveConnections_Call) Return(b bool, err error) *Repository_DoesClientHaveConnections_Call { - _c.Call.Return(b, err) - return _c -} - -func (_c *Repository_DoesClientHaveConnections_Call) RunAndReturn(run func(ctx context.Context, id string) (bool, error)) *Repository_DoesClientHaveConnections_Call { - _c.Call.Return(run) - return _c -} - -// ListEntityMembers provides a mock function for the type Repository -func (_mock *Repository) ListEntityMembers(ctx context.Context, entityID string, pageQuery roles.MembersRolePageQuery) (roles.MembersRolePage, error) { - ret := _mock.Called(ctx, entityID, pageQuery) - - if len(ret) == 0 { - panic("no return value specified for ListEntityMembers") - } - - var r0 roles.MembersRolePage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, roles.MembersRolePageQuery) (roles.MembersRolePage, error)); ok { - return returnFunc(ctx, entityID, pageQuery) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, roles.MembersRolePageQuery) roles.MembersRolePage); ok { - r0 = returnFunc(ctx, entityID, pageQuery) - } else { - r0 = ret.Get(0).(roles.MembersRolePage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, roles.MembersRolePageQuery) error); ok { - r1 = returnFunc(ctx, entityID, pageQuery) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_ListEntityMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ListEntityMembers' -type Repository_ListEntityMembers_Call struct { - *mock.Call -} - -// ListEntityMembers is a helper method to define mock.On call -// - ctx context.Context -// - entityID string -// - pageQuery roles.MembersRolePageQuery -func (_e *Repository_Expecter) ListEntityMembers(ctx interface{}, entityID interface{}, pageQuery interface{}) *Repository_ListEntityMembers_Call { - return &Repository_ListEntityMembers_Call{Call: _e.mock.On("ListEntityMembers", ctx, entityID, pageQuery)} -} - -func (_c *Repository_ListEntityMembers_Call) Run(run func(ctx context.Context, entityID string, pageQuery roles.MembersRolePageQuery)) *Repository_ListEntityMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 roles.MembersRolePageQuery - if args[2] != nil { - arg2 = args[2].(roles.MembersRolePageQuery) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_ListEntityMembers_Call) Return(membersRolePage roles.MembersRolePage, err error) *Repository_ListEntityMembers_Call { - _c.Call.Return(membersRolePage, err) - return _c -} - -func (_c *Repository_ListEntityMembers_Call) RunAndReturn(run func(ctx context.Context, entityID string, pageQuery roles.MembersRolePageQuery) (roles.MembersRolePage, error)) *Repository_ListEntityMembers_Call { - _c.Call.Return(run) - return _c -} - -// RemoveChannelConnections provides a mock function for the type Repository -func (_mock *Repository) RemoveChannelConnections(ctx context.Context, channelID string) error { - ret := _mock.Called(ctx, channelID) - - if len(ret) == 0 { - panic("no return value specified for RemoveChannelConnections") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) error); ok { - r0 = returnFunc(ctx, channelID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RemoveChannelConnections_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveChannelConnections' -type Repository_RemoveChannelConnections_Call struct { - *mock.Call -} - -// RemoveChannelConnections is a helper method to define mock.On call -// - ctx context.Context -// - channelID string -func (_e *Repository_Expecter) RemoveChannelConnections(ctx interface{}, channelID interface{}) *Repository_RemoveChannelConnections_Call { - return &Repository_RemoveChannelConnections_Call{Call: _e.mock.On("RemoveChannelConnections", ctx, channelID)} -} - -func (_c *Repository_RemoveChannelConnections_Call) Run(run func(ctx context.Context, channelID string)) *Repository_RemoveChannelConnections_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RemoveChannelConnections_Call) Return(err error) *Repository_RemoveChannelConnections_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RemoveChannelConnections_Call) RunAndReturn(run func(ctx context.Context, channelID string) error) *Repository_RemoveChannelConnections_Call { - _c.Call.Return(run) - return _c -} - -// RemoveClientConnections provides a mock function for the type Repository -func (_mock *Repository) RemoveClientConnections(ctx context.Context, clientID string) error { - ret := _mock.Called(ctx, clientID) - - if len(ret) == 0 { - panic("no return value specified for RemoveClientConnections") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) error); ok { - r0 = returnFunc(ctx, clientID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RemoveClientConnections_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveClientConnections' -type Repository_RemoveClientConnections_Call struct { - *mock.Call -} - -// RemoveClientConnections is a helper method to define mock.On call -// - ctx context.Context -// - clientID string -func (_e *Repository_Expecter) RemoveClientConnections(ctx interface{}, clientID interface{}) *Repository_RemoveClientConnections_Call { - return &Repository_RemoveClientConnections_Call{Call: _e.mock.On("RemoveClientConnections", ctx, clientID)} -} - -func (_c *Repository_RemoveClientConnections_Call) Run(run func(ctx context.Context, clientID string)) *Repository_RemoveClientConnections_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RemoveClientConnections_Call) Return(err error) *Repository_RemoveClientConnections_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RemoveClientConnections_Call) RunAndReturn(run func(ctx context.Context, clientID string) error) *Repository_RemoveClientConnections_Call { - _c.Call.Return(run) - return _c -} - -// RemoveConnections provides a mock function for the type Repository -func (_mock *Repository) RemoveConnections(ctx context.Context, conns []clients.Connection) error { - ret := _mock.Called(ctx, conns) - - if len(ret) == 0 { - panic("no return value specified for RemoveConnections") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, []clients.Connection) error); ok { - r0 = returnFunc(ctx, conns) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RemoveConnections_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveConnections' -type Repository_RemoveConnections_Call struct { - *mock.Call -} - -// RemoveConnections is a helper method to define mock.On call -// - ctx context.Context -// - conns []clients.Connection -func (_e *Repository_Expecter) RemoveConnections(ctx interface{}, conns interface{}) *Repository_RemoveConnections_Call { - return &Repository_RemoveConnections_Call{Call: _e.mock.On("RemoveConnections", ctx, conns)} -} - -func (_c *Repository_RemoveConnections_Call) Run(run func(ctx context.Context, conns []clients.Connection)) *Repository_RemoveConnections_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 []clients.Connection - if args[1] != nil { - arg1 = args[1].([]clients.Connection) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RemoveConnections_Call) Return(err error) *Repository_RemoveConnections_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RemoveConnections_Call) RunAndReturn(run func(ctx context.Context, conns []clients.Connection) error) *Repository_RemoveConnections_Call { - _c.Call.Return(run) - return _c -} - -// RemoveEntityMembers provides a mock function for the type Repository -func (_mock *Repository) RemoveEntityMembers(ctx context.Context, entityID string, members []string) error { - ret := _mock.Called(ctx, entityID, members) - - if len(ret) == 0 { - panic("no return value specified for RemoveEntityMembers") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, []string) error); ok { - r0 = returnFunc(ctx, entityID, members) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RemoveEntityMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveEntityMembers' -type Repository_RemoveEntityMembers_Call struct { - *mock.Call -} - -// RemoveEntityMembers is a helper method to define mock.On call -// - ctx context.Context -// - entityID string -// - members []string -func (_e *Repository_Expecter) RemoveEntityMembers(ctx interface{}, entityID interface{}, members interface{}) *Repository_RemoveEntityMembers_Call { - return &Repository_RemoveEntityMembers_Call{Call: _e.mock.On("RemoveEntityMembers", ctx, entityID, members)} -} - -func (_c *Repository_RemoveEntityMembers_Call) Run(run func(ctx context.Context, entityID string, members []string)) *Repository_RemoveEntityMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RemoveEntityMembers_Call) Return(err error) *Repository_RemoveEntityMembers_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RemoveEntityMembers_Call) RunAndReturn(run func(ctx context.Context, entityID string, members []string) error) *Repository_RemoveEntityMembers_Call { - _c.Call.Return(run) - return _c -} - -// RemoveMemberFromAllRoles provides a mock function for the type Repository -func (_mock *Repository) RemoveMemberFromAllRoles(ctx context.Context, memberID string) error { - ret := _mock.Called(ctx, memberID) - - if len(ret) == 0 { - panic("no return value specified for RemoveMemberFromAllRoles") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) error); ok { - r0 = returnFunc(ctx, memberID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RemoveMemberFromAllRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveMemberFromAllRoles' -type Repository_RemoveMemberFromAllRoles_Call struct { - *mock.Call -} - -// RemoveMemberFromAllRoles is a helper method to define mock.On call -// - ctx context.Context -// - memberID string -func (_e *Repository_Expecter) RemoveMemberFromAllRoles(ctx interface{}, memberID interface{}) *Repository_RemoveMemberFromAllRoles_Call { - return &Repository_RemoveMemberFromAllRoles_Call{Call: _e.mock.On("RemoveMemberFromAllRoles", ctx, memberID)} -} - -func (_c *Repository_RemoveMemberFromAllRoles_Call) Run(run func(ctx context.Context, memberID string)) *Repository_RemoveMemberFromAllRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RemoveMemberFromAllRoles_Call) Return(err error) *Repository_RemoveMemberFromAllRoles_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RemoveMemberFromAllRoles_Call) RunAndReturn(run func(ctx context.Context, memberID string) error) *Repository_RemoveMemberFromAllRoles_Call { - _c.Call.Return(run) - return _c -} - -// RemoveParentGroup provides a mock function for the type Repository -func (_mock *Repository) RemoveParentGroup(ctx context.Context, cli clients.Client) error { - ret := _mock.Called(ctx, cli) - - if len(ret) == 0 { - panic("no return value specified for RemoveParentGroup") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, clients.Client) error); ok { - r0 = returnFunc(ctx, cli) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RemoveParentGroup_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveParentGroup' -type Repository_RemoveParentGroup_Call struct { - *mock.Call -} - -// RemoveParentGroup is a helper method to define mock.On call -// - ctx context.Context -// - cli clients.Client -func (_e *Repository_Expecter) RemoveParentGroup(ctx interface{}, cli interface{}) *Repository_RemoveParentGroup_Call { - return &Repository_RemoveParentGroup_Call{Call: _e.mock.On("RemoveParentGroup", ctx, cli)} -} - -func (_c *Repository_RemoveParentGroup_Call) Run(run func(ctx context.Context, cli clients.Client)) *Repository_RemoveParentGroup_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 clients.Client - if args[1] != nil { - arg1 = args[1].(clients.Client) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RemoveParentGroup_Call) Return(err error) *Repository_RemoveParentGroup_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RemoveParentGroup_Call) RunAndReturn(run func(ctx context.Context, cli clients.Client) error) *Repository_RemoveParentGroup_Call { - _c.Call.Return(run) - return _c -} - -// RemoveRoles provides a mock function for the type Repository -func (_mock *Repository) RemoveRoles(ctx context.Context, roleIDs []string) error { - ret := _mock.Called(ctx, roleIDs) - - if len(ret) == 0 { - panic("no return value specified for RemoveRoles") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, []string) error); ok { - r0 = returnFunc(ctx, roleIDs) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RemoveRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveRoles' -type Repository_RemoveRoles_Call struct { - *mock.Call -} - -// RemoveRoles is a helper method to define mock.On call -// - ctx context.Context -// - roleIDs []string -func (_e *Repository_Expecter) RemoveRoles(ctx interface{}, roleIDs interface{}) *Repository_RemoveRoles_Call { - return &Repository_RemoveRoles_Call{Call: _e.mock.On("RemoveRoles", ctx, roleIDs)} -} - -func (_c *Repository_RemoveRoles_Call) Run(run func(ctx context.Context, roleIDs []string)) *Repository_RemoveRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 []string - if args[1] != nil { - arg1 = args[1].([]string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RemoveRoles_Call) Return(err error) *Repository_RemoveRoles_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RemoveRoles_Call) RunAndReturn(run func(ctx context.Context, roleIDs []string) error) *Repository_RemoveRoles_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveAll provides a mock function for the type Repository -func (_mock *Repository) RetrieveAll(ctx context.Context, pm clients.Page) (clients.ClientsPage, error) { - ret := _mock.Called(ctx, pm) - - if len(ret) == 0 { - panic("no return value specified for RetrieveAll") - } - - var r0 clients.ClientsPage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, clients.Page) (clients.ClientsPage, error)); ok { - return returnFunc(ctx, pm) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, clients.Page) clients.ClientsPage); ok { - r0 = returnFunc(ctx, pm) - } else { - r0 = ret.Get(0).(clients.ClientsPage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, clients.Page) error); ok { - r1 = returnFunc(ctx, pm) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveAll_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveAll' -type Repository_RetrieveAll_Call struct { - *mock.Call -} - -// RetrieveAll is a helper method to define mock.On call -// - ctx context.Context -// - pm clients.Page -func (_e *Repository_Expecter) RetrieveAll(ctx interface{}, pm interface{}) *Repository_RetrieveAll_Call { - return &Repository_RetrieveAll_Call{Call: _e.mock.On("RetrieveAll", ctx, pm)} -} - -func (_c *Repository_RetrieveAll_Call) Run(run func(ctx context.Context, pm clients.Page)) *Repository_RetrieveAll_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 clients.Page - if args[1] != nil { - arg1 = args[1].(clients.Page) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RetrieveAll_Call) Return(clientsPage clients.ClientsPage, err error) *Repository_RetrieveAll_Call { - _c.Call.Return(clientsPage, err) - return _c -} - -func (_c *Repository_RetrieveAll_Call) RunAndReturn(run func(ctx context.Context, pm clients.Page) (clients.ClientsPage, error)) *Repository_RetrieveAll_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveAllRoles provides a mock function for the type Repository -func (_mock *Repository) RetrieveAllRoles(ctx context.Context, entityID string, limit uint64, offset uint64) (roles.RolePage, error) { - ret := _mock.Called(ctx, entityID, limit, offset) - - if len(ret) == 0 { - panic("no return value specified for RetrieveAllRoles") - } - - var r0 roles.RolePage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, uint64, uint64) (roles.RolePage, error)); ok { - return returnFunc(ctx, entityID, limit, offset) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, uint64, uint64) roles.RolePage); ok { - r0 = returnFunc(ctx, entityID, limit, offset) - } else { - r0 = ret.Get(0).(roles.RolePage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, uint64, uint64) error); ok { - r1 = returnFunc(ctx, entityID, limit, offset) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveAllRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveAllRoles' -type Repository_RetrieveAllRoles_Call struct { - *mock.Call -} - -// RetrieveAllRoles is a helper method to define mock.On call -// - ctx context.Context -// - entityID string -// - limit uint64 -// - offset uint64 -func (_e *Repository_Expecter) RetrieveAllRoles(ctx interface{}, entityID interface{}, limit interface{}, offset interface{}) *Repository_RetrieveAllRoles_Call { - return &Repository_RetrieveAllRoles_Call{Call: _e.mock.On("RetrieveAllRoles", ctx, entityID, limit, offset)} -} - -func (_c *Repository_RetrieveAllRoles_Call) Run(run func(ctx context.Context, entityID string, limit uint64, offset uint64)) *Repository_RetrieveAllRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 uint64 - if args[2] != nil { - arg2 = args[2].(uint64) - } - var arg3 uint64 - if args[3] != nil { - arg3 = args[3].(uint64) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Repository_RetrieveAllRoles_Call) Return(rolePage roles.RolePage, err error) *Repository_RetrieveAllRoles_Call { - _c.Call.Return(rolePage, err) - return _c -} - -func (_c *Repository_RetrieveAllRoles_Call) RunAndReturn(run func(ctx context.Context, entityID string, limit uint64, offset uint64) (roles.RolePage, error)) *Repository_RetrieveAllRoles_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveByID provides a mock function for the type Repository -func (_mock *Repository) RetrieveByID(ctx context.Context, id string) (clients.Client, error) { - ret := _mock.Called(ctx, id) - - if len(ret) == 0 { - panic("no return value specified for RetrieveByID") - } - - var r0 clients.Client - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) (clients.Client, error)); ok { - return returnFunc(ctx, id) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string) clients.Client); ok { - r0 = returnFunc(ctx, id) - } else { - r0 = ret.Get(0).(clients.Client) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = returnFunc(ctx, id) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveByID_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveByID' -type Repository_RetrieveByID_Call struct { - *mock.Call -} - -// RetrieveByID is a helper method to define mock.On call -// - ctx context.Context -// - id string -func (_e *Repository_Expecter) RetrieveByID(ctx interface{}, id interface{}) *Repository_RetrieveByID_Call { - return &Repository_RetrieveByID_Call{Call: _e.mock.On("RetrieveByID", ctx, id)} -} - -func (_c *Repository_RetrieveByID_Call) Run(run func(ctx context.Context, id string)) *Repository_RetrieveByID_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RetrieveByID_Call) Return(client clients.Client, err error) *Repository_RetrieveByID_Call { - _c.Call.Return(client, err) - return _c -} - -func (_c *Repository_RetrieveByID_Call) RunAndReturn(run func(ctx context.Context, id string) (clients.Client, error)) *Repository_RetrieveByID_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveByIDWithRoles provides a mock function for the type Repository -func (_mock *Repository) RetrieveByIDWithRoles(ctx context.Context, id string, memberID string) (clients.Client, error) { - ret := _mock.Called(ctx, id, memberID) - - if len(ret) == 0 { - panic("no return value specified for RetrieveByIDWithRoles") - } - - var r0 clients.Client - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) (clients.Client, error)); ok { - return returnFunc(ctx, id, memberID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) clients.Client); ok { - r0 = returnFunc(ctx, id, memberID) - } else { - r0 = ret.Get(0).(clients.Client) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, string) error); ok { - r1 = returnFunc(ctx, id, memberID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveByIDWithRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveByIDWithRoles' -type Repository_RetrieveByIDWithRoles_Call struct { - *mock.Call -} - -// RetrieveByIDWithRoles is a helper method to define mock.On call -// - ctx context.Context -// - id string -// - memberID string -func (_e *Repository_Expecter) RetrieveByIDWithRoles(ctx interface{}, id interface{}, memberID interface{}) *Repository_RetrieveByIDWithRoles_Call { - return &Repository_RetrieveByIDWithRoles_Call{Call: _e.mock.On("RetrieveByIDWithRoles", ctx, id, memberID)} -} - -func (_c *Repository_RetrieveByIDWithRoles_Call) Run(run func(ctx context.Context, id string, memberID string)) *Repository_RetrieveByIDWithRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RetrieveByIDWithRoles_Call) Return(client clients.Client, err error) *Repository_RetrieveByIDWithRoles_Call { - _c.Call.Return(client, err) - return _c -} - -func (_c *Repository_RetrieveByIDWithRoles_Call) RunAndReturn(run func(ctx context.Context, id string, memberID string) (clients.Client, error)) *Repository_RetrieveByIDWithRoles_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveByIds provides a mock function for the type Repository -func (_mock *Repository) RetrieveByIds(ctx context.Context, ids []string) (clients.ClientsPage, error) { - ret := _mock.Called(ctx, ids) - - if len(ret) == 0 { - panic("no return value specified for RetrieveByIds") - } - - var r0 clients.ClientsPage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, []string) (clients.ClientsPage, error)); ok { - return returnFunc(ctx, ids) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, []string) clients.ClientsPage); ok { - r0 = returnFunc(ctx, ids) - } else { - r0 = ret.Get(0).(clients.ClientsPage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, []string) error); ok { - r1 = returnFunc(ctx, ids) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveByIds_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveByIds' -type Repository_RetrieveByIds_Call struct { - *mock.Call -} - -// RetrieveByIds is a helper method to define mock.On call -// - ctx context.Context -// - ids []string -func (_e *Repository_Expecter) RetrieveByIds(ctx interface{}, ids interface{}) *Repository_RetrieveByIds_Call { - return &Repository_RetrieveByIds_Call{Call: _e.mock.On("RetrieveByIds", ctx, ids)} -} - -func (_c *Repository_RetrieveByIds_Call) Run(run func(ctx context.Context, ids []string)) *Repository_RetrieveByIds_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 []string - if args[1] != nil { - arg1 = args[1].([]string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RetrieveByIds_Call) Return(clientsPage clients.ClientsPage, err error) *Repository_RetrieveByIds_Call { - _c.Call.Return(clientsPage, err) - return _c -} - -func (_c *Repository_RetrieveByIds_Call) RunAndReturn(run func(ctx context.Context, ids []string) (clients.ClientsPage, error)) *Repository_RetrieveByIds_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveBySecret provides a mock function for the type Repository -func (_mock *Repository) RetrieveBySecret(ctx context.Context, key string, id string, prefix authn.AuthPrefix) (clients.Client, error) { - ret := _mock.Called(ctx, key, id, prefix) - - if len(ret) == 0 { - panic("no return value specified for RetrieveBySecret") - } - - var r0 clients.Client - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string, authn.AuthPrefix) (clients.Client, error)); ok { - return returnFunc(ctx, key, id, prefix) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string, authn.AuthPrefix) clients.Client); ok { - r0 = returnFunc(ctx, key, id, prefix) - } else { - r0 = ret.Get(0).(clients.Client) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, string, authn.AuthPrefix) error); ok { - r1 = returnFunc(ctx, key, id, prefix) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveBySecret_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveBySecret' -type Repository_RetrieveBySecret_Call struct { - *mock.Call -} - -// RetrieveBySecret is a helper method to define mock.On call -// - ctx context.Context -// - key string -// - id string -// - prefix authn.AuthPrefix -func (_e *Repository_Expecter) RetrieveBySecret(ctx interface{}, key interface{}, id interface{}, prefix interface{}) *Repository_RetrieveBySecret_Call { - return &Repository_RetrieveBySecret_Call{Call: _e.mock.On("RetrieveBySecret", ctx, key, id, prefix)} -} - -func (_c *Repository_RetrieveBySecret_Call) Run(run func(ctx context.Context, key string, id string, prefix authn.AuthPrefix)) *Repository_RetrieveBySecret_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 authn.AuthPrefix - if args[3] != nil { - arg3 = args[3].(authn.AuthPrefix) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Repository_RetrieveBySecret_Call) Return(client clients.Client, err error) *Repository_RetrieveBySecret_Call { - _c.Call.Return(client, err) - return _c -} - -func (_c *Repository_RetrieveBySecret_Call) RunAndReturn(run func(ctx context.Context, key string, id string, prefix authn.AuthPrefix) (clients.Client, error)) *Repository_RetrieveBySecret_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveEntitiesRolesActionsMembers provides a mock function for the type Repository -func (_mock *Repository) RetrieveEntitiesRolesActionsMembers(ctx context.Context, entityIDs []string) ([]roles.EntityActionRole, []roles.EntityMemberRole, error) { - ret := _mock.Called(ctx, entityIDs) - - if len(ret) == 0 { - panic("no return value specified for RetrieveEntitiesRolesActionsMembers") - } - - var r0 []roles.EntityActionRole - var r1 []roles.EntityMemberRole - var r2 error - if returnFunc, ok := ret.Get(0).(func(context.Context, []string) ([]roles.EntityActionRole, []roles.EntityMemberRole, error)); ok { - return returnFunc(ctx, entityIDs) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, []string) []roles.EntityActionRole); ok { - r0 = returnFunc(ctx, entityIDs) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]roles.EntityActionRole) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, []string) []roles.EntityMemberRole); ok { - r1 = returnFunc(ctx, entityIDs) - } else { - if ret.Get(1) != nil { - r1 = ret.Get(1).([]roles.EntityMemberRole) - } - } - if returnFunc, ok := ret.Get(2).(func(context.Context, []string) error); ok { - r2 = returnFunc(ctx, entityIDs) - } else { - r2 = ret.Error(2) - } - return r0, r1, r2 -} - -// Repository_RetrieveEntitiesRolesActionsMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveEntitiesRolesActionsMembers' -type Repository_RetrieveEntitiesRolesActionsMembers_Call struct { - *mock.Call -} - -// RetrieveEntitiesRolesActionsMembers is a helper method to define mock.On call -// - ctx context.Context -// - entityIDs []string -func (_e *Repository_Expecter) RetrieveEntitiesRolesActionsMembers(ctx interface{}, entityIDs interface{}) *Repository_RetrieveEntitiesRolesActionsMembers_Call { - return &Repository_RetrieveEntitiesRolesActionsMembers_Call{Call: _e.mock.On("RetrieveEntitiesRolesActionsMembers", ctx, entityIDs)} -} - -func (_c *Repository_RetrieveEntitiesRolesActionsMembers_Call) Run(run func(ctx context.Context, entityIDs []string)) *Repository_RetrieveEntitiesRolesActionsMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 []string - if args[1] != nil { - arg1 = args[1].([]string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RetrieveEntitiesRolesActionsMembers_Call) Return(entityActionRoles []roles.EntityActionRole, entityMemberRoles []roles.EntityMemberRole, err error) *Repository_RetrieveEntitiesRolesActionsMembers_Call { - _c.Call.Return(entityActionRoles, entityMemberRoles, err) - return _c -} - -func (_c *Repository_RetrieveEntitiesRolesActionsMembers_Call) RunAndReturn(run func(ctx context.Context, entityIDs []string) ([]roles.EntityActionRole, []roles.EntityMemberRole, error)) *Repository_RetrieveEntitiesRolesActionsMembers_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveEntityRole provides a mock function for the type Repository -func (_mock *Repository) RetrieveEntityRole(ctx context.Context, entityID string, roleID string) (roles.Role, error) { - ret := _mock.Called(ctx, entityID, roleID) - - if len(ret) == 0 { - panic("no return value specified for RetrieveEntityRole") - } - - var r0 roles.Role - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) (roles.Role, error)); ok { - return returnFunc(ctx, entityID, roleID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) roles.Role); ok { - r0 = returnFunc(ctx, entityID, roleID) - } else { - r0 = ret.Get(0).(roles.Role) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, string) error); ok { - r1 = returnFunc(ctx, entityID, roleID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveEntityRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveEntityRole' -type Repository_RetrieveEntityRole_Call struct { - *mock.Call -} - -// RetrieveEntityRole is a helper method to define mock.On call -// - ctx context.Context -// - entityID string -// - roleID string -func (_e *Repository_Expecter) RetrieveEntityRole(ctx interface{}, entityID interface{}, roleID interface{}) *Repository_RetrieveEntityRole_Call { - return &Repository_RetrieveEntityRole_Call{Call: _e.mock.On("RetrieveEntityRole", ctx, entityID, roleID)} -} - -func (_c *Repository_RetrieveEntityRole_Call) Run(run func(ctx context.Context, entityID string, roleID string)) *Repository_RetrieveEntityRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RetrieveEntityRole_Call) Return(role roles.Role, err error) *Repository_RetrieveEntityRole_Call { - _c.Call.Return(role, err) - return _c -} - -func (_c *Repository_RetrieveEntityRole_Call) RunAndReturn(run func(ctx context.Context, entityID string, roleID string) (roles.Role, error)) *Repository_RetrieveEntityRole_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveParentGroupClients provides a mock function for the type Repository -func (_mock *Repository) RetrieveParentGroupClients(ctx context.Context, parentGroupID string) ([]clients.Client, error) { - ret := _mock.Called(ctx, parentGroupID) - - if len(ret) == 0 { - panic("no return value specified for RetrieveParentGroupClients") - } - - var r0 []clients.Client - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) ([]clients.Client, error)); ok { - return returnFunc(ctx, parentGroupID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string) []clients.Client); ok { - r0 = returnFunc(ctx, parentGroupID) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]clients.Client) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = returnFunc(ctx, parentGroupID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveParentGroupClients_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveParentGroupClients' -type Repository_RetrieveParentGroupClients_Call struct { - *mock.Call -} - -// RetrieveParentGroupClients is a helper method to define mock.On call -// - ctx context.Context -// - parentGroupID string -func (_e *Repository_Expecter) RetrieveParentGroupClients(ctx interface{}, parentGroupID interface{}) *Repository_RetrieveParentGroupClients_Call { - return &Repository_RetrieveParentGroupClients_Call{Call: _e.mock.On("RetrieveParentGroupClients", ctx, parentGroupID)} -} - -func (_c *Repository_RetrieveParentGroupClients_Call) Run(run func(ctx context.Context, parentGroupID string)) *Repository_RetrieveParentGroupClients_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RetrieveParentGroupClients_Call) Return(clients1 []clients.Client, err error) *Repository_RetrieveParentGroupClients_Call { - _c.Call.Return(clients1, err) - return _c -} - -func (_c *Repository_RetrieveParentGroupClients_Call) RunAndReturn(run func(ctx context.Context, parentGroupID string) ([]clients.Client, error)) *Repository_RetrieveParentGroupClients_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveRole provides a mock function for the type Repository -func (_mock *Repository) RetrieveRole(ctx context.Context, roleID string) (roles.Role, error) { - ret := _mock.Called(ctx, roleID) - - if len(ret) == 0 { - panic("no return value specified for RetrieveRole") - } - - var r0 roles.Role - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) (roles.Role, error)); ok { - return returnFunc(ctx, roleID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string) roles.Role); ok { - r0 = returnFunc(ctx, roleID) - } else { - r0 = ret.Get(0).(roles.Role) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = returnFunc(ctx, roleID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveRole' -type Repository_RetrieveRole_Call struct { - *mock.Call -} - -// RetrieveRole is a helper method to define mock.On call -// - ctx context.Context -// - roleID string -func (_e *Repository_Expecter) RetrieveRole(ctx interface{}, roleID interface{}) *Repository_RetrieveRole_Call { - return &Repository_RetrieveRole_Call{Call: _e.mock.On("RetrieveRole", ctx, roleID)} -} - -func (_c *Repository_RetrieveRole_Call) Run(run func(ctx context.Context, roleID string)) *Repository_RetrieveRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RetrieveRole_Call) Return(role roles.Role, err error) *Repository_RetrieveRole_Call { - _c.Call.Return(role, err) - return _c -} - -func (_c *Repository_RetrieveRole_Call) RunAndReturn(run func(ctx context.Context, roleID string) (roles.Role, error)) *Repository_RetrieveRole_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveUserClients provides a mock function for the type Repository -func (_mock *Repository) RetrieveUserClients(ctx context.Context, domainID string, userID string, pm clients.Page) (clients.ClientsPage, error) { - ret := _mock.Called(ctx, domainID, userID, pm) - - if len(ret) == 0 { - panic("no return value specified for RetrieveUserClients") - } - - var r0 clients.ClientsPage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string, clients.Page) (clients.ClientsPage, error)); ok { - return returnFunc(ctx, domainID, userID, pm) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string, clients.Page) clients.ClientsPage); ok { - r0 = returnFunc(ctx, domainID, userID, pm) - } else { - r0 = ret.Get(0).(clients.ClientsPage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, string, clients.Page) error); ok { - r1 = returnFunc(ctx, domainID, userID, pm) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveUserClients_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveUserClients' -type Repository_RetrieveUserClients_Call struct { - *mock.Call -} - -// RetrieveUserClients is a helper method to define mock.On call -// - ctx context.Context -// - domainID string -// - userID string -// - pm clients.Page -func (_e *Repository_Expecter) RetrieveUserClients(ctx interface{}, domainID interface{}, userID interface{}, pm interface{}) *Repository_RetrieveUserClients_Call { - return &Repository_RetrieveUserClients_Call{Call: _e.mock.On("RetrieveUserClients", ctx, domainID, userID, pm)} -} - -func (_c *Repository_RetrieveUserClients_Call) Run(run func(ctx context.Context, domainID string, userID string, pm clients.Page)) *Repository_RetrieveUserClients_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 clients.Page - if args[3] != nil { - arg3 = args[3].(clients.Page) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Repository_RetrieveUserClients_Call) Return(clientsPage clients.ClientsPage, err error) *Repository_RetrieveUserClients_Call { - _c.Call.Return(clientsPage, err) - return _c -} - -func (_c *Repository_RetrieveUserClients_Call) RunAndReturn(run func(ctx context.Context, domainID string, userID string, pm clients.Page) (clients.ClientsPage, error)) *Repository_RetrieveUserClients_Call { - _c.Call.Return(run) - return _c -} - -// RoleAddActions provides a mock function for the type Repository -func (_mock *Repository) RoleAddActions(ctx context.Context, role roles.Role, actions []string) ([]string, error) { - ret := _mock.Called(ctx, role, actions) - - if len(ret) == 0 { - panic("no return value specified for RoleAddActions") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role, []string) ([]string, error)); ok { - return returnFunc(ctx, role, actions) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role, []string) []string); ok { - r0 = returnFunc(ctx, role, actions) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, roles.Role, []string) error); ok { - r1 = returnFunc(ctx, role, actions) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RoleAddActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleAddActions' -type Repository_RoleAddActions_Call struct { - *mock.Call -} - -// RoleAddActions is a helper method to define mock.On call -// - ctx context.Context -// - role roles.Role -// - actions []string -func (_e *Repository_Expecter) RoleAddActions(ctx interface{}, role interface{}, actions interface{}) *Repository_RoleAddActions_Call { - return &Repository_RoleAddActions_Call{Call: _e.mock.On("RoleAddActions", ctx, role, actions)} -} - -func (_c *Repository_RoleAddActions_Call) Run(run func(ctx context.Context, role roles.Role, actions []string)) *Repository_RoleAddActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RoleAddActions_Call) Return(ops []string, err error) *Repository_RoleAddActions_Call { - _c.Call.Return(ops, err) - return _c -} - -func (_c *Repository_RoleAddActions_Call) RunAndReturn(run func(ctx context.Context, role roles.Role, actions []string) ([]string, error)) *Repository_RoleAddActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleAddMembers provides a mock function for the type Repository -func (_mock *Repository) RoleAddMembers(ctx context.Context, role roles.Role, members []string) ([]string, error) { - ret := _mock.Called(ctx, role, members) - - if len(ret) == 0 { - panic("no return value specified for RoleAddMembers") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role, []string) ([]string, error)); ok { - return returnFunc(ctx, role, members) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role, []string) []string); ok { - r0 = returnFunc(ctx, role, members) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, roles.Role, []string) error); ok { - r1 = returnFunc(ctx, role, members) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RoleAddMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleAddMembers' -type Repository_RoleAddMembers_Call struct { - *mock.Call -} - -// RoleAddMembers is a helper method to define mock.On call -// - ctx context.Context -// - role roles.Role -// - members []string -func (_e *Repository_Expecter) RoleAddMembers(ctx interface{}, role interface{}, members interface{}) *Repository_RoleAddMembers_Call { - return &Repository_RoleAddMembers_Call{Call: _e.mock.On("RoleAddMembers", ctx, role, members)} -} - -func (_c *Repository_RoleAddMembers_Call) Run(run func(ctx context.Context, role roles.Role, members []string)) *Repository_RoleAddMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RoleAddMembers_Call) Return(strings []string, err error) *Repository_RoleAddMembers_Call { - _c.Call.Return(strings, err) - return _c -} - -func (_c *Repository_RoleAddMembers_Call) RunAndReturn(run func(ctx context.Context, role roles.Role, members []string) ([]string, error)) *Repository_RoleAddMembers_Call { - _c.Call.Return(run) - return _c -} - -// RoleCheckActionsExists provides a mock function for the type Repository -func (_mock *Repository) RoleCheckActionsExists(ctx context.Context, roleID string, actions []string) (bool, error) { - ret := _mock.Called(ctx, roleID, actions) - - if len(ret) == 0 { - panic("no return value specified for RoleCheckActionsExists") - } - - var r0 bool - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, []string) (bool, error)); ok { - return returnFunc(ctx, roleID, actions) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, []string) bool); ok { - r0 = returnFunc(ctx, roleID, actions) - } else { - r0 = ret.Get(0).(bool) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, []string) error); ok { - r1 = returnFunc(ctx, roleID, actions) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RoleCheckActionsExists_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleCheckActionsExists' -type Repository_RoleCheckActionsExists_Call struct { - *mock.Call -} - -// RoleCheckActionsExists is a helper method to define mock.On call -// - ctx context.Context -// - roleID string -// - actions []string -func (_e *Repository_Expecter) RoleCheckActionsExists(ctx interface{}, roleID interface{}, actions interface{}) *Repository_RoleCheckActionsExists_Call { - return &Repository_RoleCheckActionsExists_Call{Call: _e.mock.On("RoleCheckActionsExists", ctx, roleID, actions)} -} - -func (_c *Repository_RoleCheckActionsExists_Call) Run(run func(ctx context.Context, roleID string, actions []string)) *Repository_RoleCheckActionsExists_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RoleCheckActionsExists_Call) Return(b bool, err error) *Repository_RoleCheckActionsExists_Call { - _c.Call.Return(b, err) - return _c -} - -func (_c *Repository_RoleCheckActionsExists_Call) RunAndReturn(run func(ctx context.Context, roleID string, actions []string) (bool, error)) *Repository_RoleCheckActionsExists_Call { - _c.Call.Return(run) - return _c -} - -// RoleCheckMembersExists provides a mock function for the type Repository -func (_mock *Repository) RoleCheckMembersExists(ctx context.Context, roleID string, members []string) (bool, error) { - ret := _mock.Called(ctx, roleID, members) - - if len(ret) == 0 { - panic("no return value specified for RoleCheckMembersExists") - } - - var r0 bool - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, []string) (bool, error)); ok { - return returnFunc(ctx, roleID, members) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, []string) bool); ok { - r0 = returnFunc(ctx, roleID, members) - } else { - r0 = ret.Get(0).(bool) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, []string) error); ok { - r1 = returnFunc(ctx, roleID, members) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RoleCheckMembersExists_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleCheckMembersExists' -type Repository_RoleCheckMembersExists_Call struct { - *mock.Call -} - -// RoleCheckMembersExists is a helper method to define mock.On call -// - ctx context.Context -// - roleID string -// - members []string -func (_e *Repository_Expecter) RoleCheckMembersExists(ctx interface{}, roleID interface{}, members interface{}) *Repository_RoleCheckMembersExists_Call { - return &Repository_RoleCheckMembersExists_Call{Call: _e.mock.On("RoleCheckMembersExists", ctx, roleID, members)} -} - -func (_c *Repository_RoleCheckMembersExists_Call) Run(run func(ctx context.Context, roleID string, members []string)) *Repository_RoleCheckMembersExists_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RoleCheckMembersExists_Call) Return(b bool, err error) *Repository_RoleCheckMembersExists_Call { - _c.Call.Return(b, err) - return _c -} - -func (_c *Repository_RoleCheckMembersExists_Call) RunAndReturn(run func(ctx context.Context, roleID string, members []string) (bool, error)) *Repository_RoleCheckMembersExists_Call { - _c.Call.Return(run) - return _c -} - -// RoleListActions provides a mock function for the type Repository -func (_mock *Repository) RoleListActions(ctx context.Context, roleID string) ([]string, error) { - ret := _mock.Called(ctx, roleID) - - if len(ret) == 0 { - panic("no return value specified for RoleListActions") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) ([]string, error)); ok { - return returnFunc(ctx, roleID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string) []string); ok { - r0 = returnFunc(ctx, roleID) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = returnFunc(ctx, roleID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RoleListActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleListActions' -type Repository_RoleListActions_Call struct { - *mock.Call -} - -// RoleListActions is a helper method to define mock.On call -// - ctx context.Context -// - roleID string -func (_e *Repository_Expecter) RoleListActions(ctx interface{}, roleID interface{}) *Repository_RoleListActions_Call { - return &Repository_RoleListActions_Call{Call: _e.mock.On("RoleListActions", ctx, roleID)} -} - -func (_c *Repository_RoleListActions_Call) Run(run func(ctx context.Context, roleID string)) *Repository_RoleListActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RoleListActions_Call) Return(strings []string, err error) *Repository_RoleListActions_Call { - _c.Call.Return(strings, err) - return _c -} - -func (_c *Repository_RoleListActions_Call) RunAndReturn(run func(ctx context.Context, roleID string) ([]string, error)) *Repository_RoleListActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleListMembers provides a mock function for the type Repository -func (_mock *Repository) RoleListMembers(ctx context.Context, roleID string, limit uint64, offset uint64) (roles.MembersPage, error) { - ret := _mock.Called(ctx, roleID, limit, offset) - - if len(ret) == 0 { - panic("no return value specified for RoleListMembers") - } - - var r0 roles.MembersPage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, uint64, uint64) (roles.MembersPage, error)); ok { - return returnFunc(ctx, roleID, limit, offset) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, uint64, uint64) roles.MembersPage); ok { - r0 = returnFunc(ctx, roleID, limit, offset) - } else { - r0 = ret.Get(0).(roles.MembersPage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, uint64, uint64) error); ok { - r1 = returnFunc(ctx, roleID, limit, offset) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RoleListMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleListMembers' -type Repository_RoleListMembers_Call struct { - *mock.Call -} - -// RoleListMembers is a helper method to define mock.On call -// - ctx context.Context -// - roleID string -// - limit uint64 -// - offset uint64 -func (_e *Repository_Expecter) RoleListMembers(ctx interface{}, roleID interface{}, limit interface{}, offset interface{}) *Repository_RoleListMembers_Call { - return &Repository_RoleListMembers_Call{Call: _e.mock.On("RoleListMembers", ctx, roleID, limit, offset)} -} - -func (_c *Repository_RoleListMembers_Call) Run(run func(ctx context.Context, roleID string, limit uint64, offset uint64)) *Repository_RoleListMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 uint64 - if args[2] != nil { - arg2 = args[2].(uint64) - } - var arg3 uint64 - if args[3] != nil { - arg3 = args[3].(uint64) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Repository_RoleListMembers_Call) Return(membersPage roles.MembersPage, err error) *Repository_RoleListMembers_Call { - _c.Call.Return(membersPage, err) - return _c -} - -func (_c *Repository_RoleListMembers_Call) RunAndReturn(run func(ctx context.Context, roleID string, limit uint64, offset uint64) (roles.MembersPage, error)) *Repository_RoleListMembers_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveActions provides a mock function for the type Repository -func (_mock *Repository) RoleRemoveActions(ctx context.Context, role roles.Role, actions []string) error { - ret := _mock.Called(ctx, role, actions) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveActions") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role, []string) error); ok { - r0 = returnFunc(ctx, role, actions) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RoleRemoveActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveActions' -type Repository_RoleRemoveActions_Call struct { - *mock.Call -} - -// RoleRemoveActions is a helper method to define mock.On call -// - ctx context.Context -// - role roles.Role -// - actions []string -func (_e *Repository_Expecter) RoleRemoveActions(ctx interface{}, role interface{}, actions interface{}) *Repository_RoleRemoveActions_Call { - return &Repository_RoleRemoveActions_Call{Call: _e.mock.On("RoleRemoveActions", ctx, role, actions)} -} - -func (_c *Repository_RoleRemoveActions_Call) Run(run func(ctx context.Context, role roles.Role, actions []string)) *Repository_RoleRemoveActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RoleRemoveActions_Call) Return(err error) *Repository_RoleRemoveActions_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RoleRemoveActions_Call) RunAndReturn(run func(ctx context.Context, role roles.Role, actions []string) error) *Repository_RoleRemoveActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveAllActions provides a mock function for the type Repository -func (_mock *Repository) RoleRemoveAllActions(ctx context.Context, role roles.Role) error { - ret := _mock.Called(ctx, role) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveAllActions") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role) error); ok { - r0 = returnFunc(ctx, role) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RoleRemoveAllActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveAllActions' -type Repository_RoleRemoveAllActions_Call struct { - *mock.Call -} - -// RoleRemoveAllActions is a helper method to define mock.On call -// - ctx context.Context -// - role roles.Role -func (_e *Repository_Expecter) RoleRemoveAllActions(ctx interface{}, role interface{}) *Repository_RoleRemoveAllActions_Call { - return &Repository_RoleRemoveAllActions_Call{Call: _e.mock.On("RoleRemoveAllActions", ctx, role)} -} - -func (_c *Repository_RoleRemoveAllActions_Call) Run(run func(ctx context.Context, role roles.Role)) *Repository_RoleRemoveAllActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RoleRemoveAllActions_Call) Return(err error) *Repository_RoleRemoveAllActions_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RoleRemoveAllActions_Call) RunAndReturn(run func(ctx context.Context, role roles.Role) error) *Repository_RoleRemoveAllActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveAllMembers provides a mock function for the type Repository -func (_mock *Repository) RoleRemoveAllMembers(ctx context.Context, role roles.Role) error { - ret := _mock.Called(ctx, role) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveAllMembers") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role) error); ok { - r0 = returnFunc(ctx, role) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RoleRemoveAllMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveAllMembers' -type Repository_RoleRemoveAllMembers_Call struct { - *mock.Call -} - -// RoleRemoveAllMembers is a helper method to define mock.On call -// - ctx context.Context -// - role roles.Role -func (_e *Repository_Expecter) RoleRemoveAllMembers(ctx interface{}, role interface{}) *Repository_RoleRemoveAllMembers_Call { - return &Repository_RoleRemoveAllMembers_Call{Call: _e.mock.On("RoleRemoveAllMembers", ctx, role)} -} - -func (_c *Repository_RoleRemoveAllMembers_Call) Run(run func(ctx context.Context, role roles.Role)) *Repository_RoleRemoveAllMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RoleRemoveAllMembers_Call) Return(err error) *Repository_RoleRemoveAllMembers_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RoleRemoveAllMembers_Call) RunAndReturn(run func(ctx context.Context, role roles.Role) error) *Repository_RoleRemoveAllMembers_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveMembers provides a mock function for the type Repository -func (_mock *Repository) RoleRemoveMembers(ctx context.Context, role roles.Role, members []string) error { - ret := _mock.Called(ctx, role, members) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveMembers") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role, []string) error); ok { - r0 = returnFunc(ctx, role, members) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RoleRemoveMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveMembers' -type Repository_RoleRemoveMembers_Call struct { - *mock.Call -} - -// RoleRemoveMembers is a helper method to define mock.On call -// - ctx context.Context -// - role roles.Role -// - members []string -func (_e *Repository_Expecter) RoleRemoveMembers(ctx interface{}, role interface{}, members interface{}) *Repository_RoleRemoveMembers_Call { - return &Repository_RoleRemoveMembers_Call{Call: _e.mock.On("RoleRemoveMembers", ctx, role, members)} -} - -func (_c *Repository_RoleRemoveMembers_Call) Run(run func(ctx context.Context, role roles.Role, members []string)) *Repository_RoleRemoveMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RoleRemoveMembers_Call) Return(err error) *Repository_RoleRemoveMembers_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RoleRemoveMembers_Call) RunAndReturn(run func(ctx context.Context, role roles.Role, members []string) error) *Repository_RoleRemoveMembers_Call { - _c.Call.Return(run) - return _c -} - -// Save provides a mock function for the type Repository -func (_mock *Repository) Save(ctx context.Context, client ...clients.Client) ([]clients.Client, error) { - var tmpRet mock.Arguments - if len(client) > 0 { - tmpRet = _mock.Called(ctx, client) - } else { - tmpRet = _mock.Called(ctx) - } - ret := tmpRet - - if len(ret) == 0 { - panic("no return value specified for Save") - } - - var r0 []clients.Client - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, ...clients.Client) ([]clients.Client, error)); ok { - return returnFunc(ctx, client...) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, ...clients.Client) []clients.Client); ok { - r0 = returnFunc(ctx, client...) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]clients.Client) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, ...clients.Client) error); ok { - r1 = returnFunc(ctx, client...) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_Save_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Save' -type Repository_Save_Call struct { - *mock.Call -} - -// Save is a helper method to define mock.On call -// - ctx context.Context -// - client ...clients.Client -func (_e *Repository_Expecter) Save(ctx interface{}, client ...interface{}) *Repository_Save_Call { - return &Repository_Save_Call{Call: _e.mock.On("Save", - append([]interface{}{ctx}, client...)...)} -} - -func (_c *Repository_Save_Call) Run(run func(ctx context.Context, client ...clients.Client)) *Repository_Save_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 []clients.Client - var variadicArgs []clients.Client - if len(args) > 1 { - variadicArgs = args[1].([]clients.Client) - } - arg1 = variadicArgs - run( - arg0, - arg1..., - ) - }) - return _c -} - -func (_c *Repository_Save_Call) Return(clients1 []clients.Client, err error) *Repository_Save_Call { - _c.Call.Return(clients1, err) - return _c -} - -func (_c *Repository_Save_Call) RunAndReturn(run func(ctx context.Context, client ...clients.Client) ([]clients.Client, error)) *Repository_Save_Call { - _c.Call.Return(run) - return _c -} - -// SearchClients provides a mock function for the type Repository -func (_mock *Repository) SearchClients(ctx context.Context, pm clients.Page) (clients.ClientsPage, error) { - ret := _mock.Called(ctx, pm) - - if len(ret) == 0 { - panic("no return value specified for SearchClients") - } - - var r0 clients.ClientsPage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, clients.Page) (clients.ClientsPage, error)); ok { - return returnFunc(ctx, pm) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, clients.Page) clients.ClientsPage); ok { - r0 = returnFunc(ctx, pm) - } else { - r0 = ret.Get(0).(clients.ClientsPage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, clients.Page) error); ok { - r1 = returnFunc(ctx, pm) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_SearchClients_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'SearchClients' -type Repository_SearchClients_Call struct { - *mock.Call -} - -// SearchClients is a helper method to define mock.On call -// - ctx context.Context -// - pm clients.Page -func (_e *Repository_Expecter) SearchClients(ctx interface{}, pm interface{}) *Repository_SearchClients_Call { - return &Repository_SearchClients_Call{Call: _e.mock.On("SearchClients", ctx, pm)} -} - -func (_c *Repository_SearchClients_Call) Run(run func(ctx context.Context, pm clients.Page)) *Repository_SearchClients_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 clients.Page - if args[1] != nil { - arg1 = args[1].(clients.Page) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_SearchClients_Call) Return(clientsPage clients.ClientsPage, err error) *Repository_SearchClients_Call { - _c.Call.Return(clientsPage, err) - return _c -} - -func (_c *Repository_SearchClients_Call) RunAndReturn(run func(ctx context.Context, pm clients.Page) (clients.ClientsPage, error)) *Repository_SearchClients_Call { - _c.Call.Return(run) - return _c -} - -// SetParentGroup provides a mock function for the type Repository -func (_mock *Repository) SetParentGroup(ctx context.Context, cli clients.Client) error { - ret := _mock.Called(ctx, cli) - - if len(ret) == 0 { - panic("no return value specified for SetParentGroup") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, clients.Client) error); ok { - r0 = returnFunc(ctx, cli) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_SetParentGroup_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'SetParentGroup' -type Repository_SetParentGroup_Call struct { - *mock.Call -} - -// SetParentGroup is a helper method to define mock.On call -// - ctx context.Context -// - cli clients.Client -func (_e *Repository_Expecter) SetParentGroup(ctx interface{}, cli interface{}) *Repository_SetParentGroup_Call { - return &Repository_SetParentGroup_Call{Call: _e.mock.On("SetParentGroup", ctx, cli)} -} - -func (_c *Repository_SetParentGroup_Call) Run(run func(ctx context.Context, cli clients.Client)) *Repository_SetParentGroup_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 clients.Client - if args[1] != nil { - arg1 = args[1].(clients.Client) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_SetParentGroup_Call) Return(err error) *Repository_SetParentGroup_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_SetParentGroup_Call) RunAndReturn(run func(ctx context.Context, cli clients.Client) error) *Repository_SetParentGroup_Call { - _c.Call.Return(run) - return _c -} - -// UnsetParentGroupFromClient provides a mock function for the type Repository -func (_mock *Repository) UnsetParentGroupFromClient(ctx context.Context, parentGroupID string) error { - ret := _mock.Called(ctx, parentGroupID) - - if len(ret) == 0 { - panic("no return value specified for UnsetParentGroupFromClient") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) error); ok { - r0 = returnFunc(ctx, parentGroupID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_UnsetParentGroupFromClient_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UnsetParentGroupFromClient' -type Repository_UnsetParentGroupFromClient_Call struct { - *mock.Call -} - -// UnsetParentGroupFromClient is a helper method to define mock.On call -// - ctx context.Context -// - parentGroupID string -func (_e *Repository_Expecter) UnsetParentGroupFromClient(ctx interface{}, parentGroupID interface{}) *Repository_UnsetParentGroupFromClient_Call { - return &Repository_UnsetParentGroupFromClient_Call{Call: _e.mock.On("UnsetParentGroupFromClient", ctx, parentGroupID)} -} - -func (_c *Repository_UnsetParentGroupFromClient_Call) Run(run func(ctx context.Context, parentGroupID string)) *Repository_UnsetParentGroupFromClient_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_UnsetParentGroupFromClient_Call) Return(err error) *Repository_UnsetParentGroupFromClient_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_UnsetParentGroupFromClient_Call) RunAndReturn(run func(ctx context.Context, parentGroupID string) error) *Repository_UnsetParentGroupFromClient_Call { - _c.Call.Return(run) - return _c -} - -// Update provides a mock function for the type Repository -func (_mock *Repository) Update(ctx context.Context, client clients.Client) (clients.Client, error) { - ret := _mock.Called(ctx, client) - - if len(ret) == 0 { - panic("no return value specified for Update") - } - - var r0 clients.Client - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, clients.Client) (clients.Client, error)); ok { - return returnFunc(ctx, client) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, clients.Client) clients.Client); ok { - r0 = returnFunc(ctx, client) - } else { - r0 = ret.Get(0).(clients.Client) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, clients.Client) error); ok { - r1 = returnFunc(ctx, client) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_Update_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Update' -type Repository_Update_Call struct { - *mock.Call -} - -// Update is a helper method to define mock.On call -// - ctx context.Context -// - client clients.Client -func (_e *Repository_Expecter) Update(ctx interface{}, client interface{}) *Repository_Update_Call { - return &Repository_Update_Call{Call: _e.mock.On("Update", ctx, client)} -} - -func (_c *Repository_Update_Call) Run(run func(ctx context.Context, client clients.Client)) *Repository_Update_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 clients.Client - if args[1] != nil { - arg1 = args[1].(clients.Client) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_Update_Call) Return(client1 clients.Client, err error) *Repository_Update_Call { - _c.Call.Return(client1, err) - return _c -} - -func (_c *Repository_Update_Call) RunAndReturn(run func(ctx context.Context, client clients.Client) (clients.Client, error)) *Repository_Update_Call { - _c.Call.Return(run) - return _c -} - -// UpdateIdentity provides a mock function for the type Repository -func (_mock *Repository) UpdateIdentity(ctx context.Context, client clients.Client) (clients.Client, error) { - ret := _mock.Called(ctx, client) - - if len(ret) == 0 { - panic("no return value specified for UpdateIdentity") - } - - var r0 clients.Client - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, clients.Client) (clients.Client, error)); ok { - return returnFunc(ctx, client) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, clients.Client) clients.Client); ok { - r0 = returnFunc(ctx, client) - } else { - r0 = ret.Get(0).(clients.Client) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, clients.Client) error); ok { - r1 = returnFunc(ctx, client) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_UpdateIdentity_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateIdentity' -type Repository_UpdateIdentity_Call struct { - *mock.Call -} - -// UpdateIdentity is a helper method to define mock.On call -// - ctx context.Context -// - client clients.Client -func (_e *Repository_Expecter) UpdateIdentity(ctx interface{}, client interface{}) *Repository_UpdateIdentity_Call { - return &Repository_UpdateIdentity_Call{Call: _e.mock.On("UpdateIdentity", ctx, client)} -} - -func (_c *Repository_UpdateIdentity_Call) Run(run func(ctx context.Context, client clients.Client)) *Repository_UpdateIdentity_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 clients.Client - if args[1] != nil { - arg1 = args[1].(clients.Client) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_UpdateIdentity_Call) Return(client1 clients.Client, err error) *Repository_UpdateIdentity_Call { - _c.Call.Return(client1, err) - return _c -} - -func (_c *Repository_UpdateIdentity_Call) RunAndReturn(run func(ctx context.Context, client clients.Client) (clients.Client, error)) *Repository_UpdateIdentity_Call { - _c.Call.Return(run) - return _c -} - -// UpdateRole provides a mock function for the type Repository -func (_mock *Repository) UpdateRole(ctx context.Context, ro roles.Role) (roles.Role, error) { - ret := _mock.Called(ctx, ro) - - if len(ret) == 0 { - panic("no return value specified for UpdateRole") - } - - var r0 roles.Role - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role) (roles.Role, error)); ok { - return returnFunc(ctx, ro) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role) roles.Role); ok { - r0 = returnFunc(ctx, ro) - } else { - r0 = ret.Get(0).(roles.Role) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, roles.Role) error); ok { - r1 = returnFunc(ctx, ro) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_UpdateRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateRole' -type Repository_UpdateRole_Call struct { - *mock.Call -} - -// UpdateRole is a helper method to define mock.On call -// - ctx context.Context -// - ro roles.Role -func (_e *Repository_Expecter) UpdateRole(ctx interface{}, ro interface{}) *Repository_UpdateRole_Call { - return &Repository_UpdateRole_Call{Call: _e.mock.On("UpdateRole", ctx, ro)} -} - -func (_c *Repository_UpdateRole_Call) Run(run func(ctx context.Context, ro roles.Role)) *Repository_UpdateRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_UpdateRole_Call) Return(role roles.Role, err error) *Repository_UpdateRole_Call { - _c.Call.Return(role, err) - return _c -} - -func (_c *Repository_UpdateRole_Call) RunAndReturn(run func(ctx context.Context, ro roles.Role) (roles.Role, error)) *Repository_UpdateRole_Call { - _c.Call.Return(run) - return _c -} - -// UpdateSecret provides a mock function for the type Repository -func (_mock *Repository) UpdateSecret(ctx context.Context, client clients.Client) (clients.Client, error) { - ret := _mock.Called(ctx, client) - - if len(ret) == 0 { - panic("no return value specified for UpdateSecret") - } - - var r0 clients.Client - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, clients.Client) (clients.Client, error)); ok { - return returnFunc(ctx, client) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, clients.Client) clients.Client); ok { - r0 = returnFunc(ctx, client) - } else { - r0 = ret.Get(0).(clients.Client) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, clients.Client) error); ok { - r1 = returnFunc(ctx, client) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_UpdateSecret_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateSecret' -type Repository_UpdateSecret_Call struct { - *mock.Call -} - -// UpdateSecret is a helper method to define mock.On call -// - ctx context.Context -// - client clients.Client -func (_e *Repository_Expecter) UpdateSecret(ctx interface{}, client interface{}) *Repository_UpdateSecret_Call { - return &Repository_UpdateSecret_Call{Call: _e.mock.On("UpdateSecret", ctx, client)} -} - -func (_c *Repository_UpdateSecret_Call) Run(run func(ctx context.Context, client clients.Client)) *Repository_UpdateSecret_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 clients.Client - if args[1] != nil { - arg1 = args[1].(clients.Client) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_UpdateSecret_Call) Return(client1 clients.Client, err error) *Repository_UpdateSecret_Call { - _c.Call.Return(client1, err) - return _c -} - -func (_c *Repository_UpdateSecret_Call) RunAndReturn(run func(ctx context.Context, client clients.Client) (clients.Client, error)) *Repository_UpdateSecret_Call { - _c.Call.Return(run) - return _c -} - -// UpdateTags provides a mock function for the type Repository -func (_mock *Repository) UpdateTags(ctx context.Context, client clients.Client) (clients.Client, error) { - ret := _mock.Called(ctx, client) - - if len(ret) == 0 { - panic("no return value specified for UpdateTags") - } - - var r0 clients.Client - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, clients.Client) (clients.Client, error)); ok { - return returnFunc(ctx, client) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, clients.Client) clients.Client); ok { - r0 = returnFunc(ctx, client) - } else { - r0 = ret.Get(0).(clients.Client) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, clients.Client) error); ok { - r1 = returnFunc(ctx, client) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_UpdateTags_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateTags' -type Repository_UpdateTags_Call struct { - *mock.Call -} - -// UpdateTags is a helper method to define mock.On call -// - ctx context.Context -// - client clients.Client -func (_e *Repository_Expecter) UpdateTags(ctx interface{}, client interface{}) *Repository_UpdateTags_Call { - return &Repository_UpdateTags_Call{Call: _e.mock.On("UpdateTags", ctx, client)} -} - -func (_c *Repository_UpdateTags_Call) Run(run func(ctx context.Context, client clients.Client)) *Repository_UpdateTags_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 clients.Client - if args[1] != nil { - arg1 = args[1].(clients.Client) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_UpdateTags_Call) Return(client1 clients.Client, err error) *Repository_UpdateTags_Call { - _c.Call.Return(client1, err) - return _c -} - -func (_c *Repository_UpdateTags_Call) RunAndReturn(run func(ctx context.Context, client clients.Client) (clients.Client, error)) *Repository_UpdateTags_Call { - _c.Call.Return(run) - return _c -} diff --git a/clients/mocks/service.go b/clients/mocks/service.go deleted file mode 100644 index 04ed9e28d..000000000 --- a/clients/mocks/service.go +++ /dev/null @@ -1,2406 +0,0 @@ -// Copyright (c) Abstract Machines - -// SPDX-License-Identifier: Apache-2.0 - -// Code generated by mockery; DO NOT EDIT. -// github.com/vektra/mockery -// template: testify - -package mocks - -import ( - "context" - - "github.com/absmach/magistrala/clients" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/roles" - mock "github.com/stretchr/testify/mock" -) - -// NewService creates a new instance of Service. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. -// The first argument is typically a *testing.T value. -func NewService(t interface { - mock.TestingT - Cleanup(func()) -}) *Service { - mock := &Service{} - mock.Mock.Test(t) - - t.Cleanup(func() { mock.AssertExpectations(t) }) - - return mock -} - -// Service is an autogenerated mock type for the Service type -type Service struct { - mock.Mock -} - -type Service_Expecter struct { - mock *mock.Mock -} - -func (_m *Service) EXPECT() *Service_Expecter { - return &Service_Expecter{mock: &_m.Mock} -} - -// AddRole provides a mock function for the type Service -func (_mock *Service) AddRole(ctx context.Context, session authn.Session, entityID string, roleName string, optionalActions []string, optionalMembers []string) (roles.RoleProvision, error) { - ret := _mock.Called(ctx, session, entityID, roleName, optionalActions, optionalMembers) - - if len(ret) == 0 { - panic("no return value specified for AddRole") - } - - var r0 roles.RoleProvision - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string, []string) (roles.RoleProvision, error)); ok { - return returnFunc(ctx, session, entityID, roleName, optionalActions, optionalMembers) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string, []string) roles.RoleProvision); ok { - r0 = returnFunc(ctx, session, entityID, roleName, optionalActions, optionalMembers) - } else { - r0 = ret.Get(0).(roles.RoleProvision) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, []string, []string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleName, optionalActions, optionalMembers) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_AddRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'AddRole' -type Service_AddRole_Call struct { - *mock.Call -} - -// AddRole is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleName string -// - optionalActions []string -// - optionalMembers []string -func (_e *Service_Expecter) AddRole(ctx interface{}, session interface{}, entityID interface{}, roleName interface{}, optionalActions interface{}, optionalMembers interface{}) *Service_AddRole_Call { - return &Service_AddRole_Call{Call: _e.mock.On("AddRole", ctx, session, entityID, roleName, optionalActions, optionalMembers)} -} - -func (_c *Service_AddRole_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleName string, optionalActions []string, optionalMembers []string)) *Service_AddRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - var arg5 []string - if args[5] != nil { - arg5 = args[5].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - arg5, - ) - }) - return _c -} - -func (_c *Service_AddRole_Call) Return(roleProvision roles.RoleProvision, err error) *Service_AddRole_Call { - _c.Call.Return(roleProvision, err) - return _c -} - -func (_c *Service_AddRole_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleName string, optionalActions []string, optionalMembers []string) (roles.RoleProvision, error)) *Service_AddRole_Call { - _c.Call.Return(run) - return _c -} - -// CreateClients provides a mock function for the type Service -func (_mock *Service) CreateClients(ctx context.Context, session authn.Session, client ...clients.Client) ([]clients.Client, []roles.RoleProvision, error) { - var tmpRet mock.Arguments - if len(client) > 0 { - tmpRet = _mock.Called(ctx, session, client) - } else { - tmpRet = _mock.Called(ctx, session) - } - ret := tmpRet - - if len(ret) == 0 { - panic("no return value specified for CreateClients") - } - - var r0 []clients.Client - var r1 []roles.RoleProvision - var r2 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, ...clients.Client) ([]clients.Client, []roles.RoleProvision, error)); ok { - return returnFunc(ctx, session, client...) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, ...clients.Client) []clients.Client); ok { - r0 = returnFunc(ctx, session, client...) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]clients.Client) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, ...clients.Client) []roles.RoleProvision); ok { - r1 = returnFunc(ctx, session, client...) - } else { - if ret.Get(1) != nil { - r1 = ret.Get(1).([]roles.RoleProvision) - } - } - if returnFunc, ok := ret.Get(2).(func(context.Context, authn.Session, ...clients.Client) error); ok { - r2 = returnFunc(ctx, session, client...) - } else { - r2 = ret.Error(2) - } - return r0, r1, r2 -} - -// Service_CreateClients_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'CreateClients' -type Service_CreateClients_Call struct { - *mock.Call -} - -// CreateClients is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - client ...clients.Client -func (_e *Service_Expecter) CreateClients(ctx interface{}, session interface{}, client ...interface{}) *Service_CreateClients_Call { - return &Service_CreateClients_Call{Call: _e.mock.On("CreateClients", - append([]interface{}{ctx, session}, client...)...)} -} - -func (_c *Service_CreateClients_Call) Run(run func(ctx context.Context, session authn.Session, client ...clients.Client)) *Service_CreateClients_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 []clients.Client - var variadicArgs []clients.Client - if len(args) > 2 { - variadicArgs = args[2].([]clients.Client) - } - arg2 = variadicArgs - run( - arg0, - arg1, - arg2..., - ) - }) - return _c -} - -func (_c *Service_CreateClients_Call) Return(clients1 []clients.Client, roleProvisions []roles.RoleProvision, err error) *Service_CreateClients_Call { - _c.Call.Return(clients1, roleProvisions, err) - return _c -} - -func (_c *Service_CreateClients_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, client ...clients.Client) ([]clients.Client, []roles.RoleProvision, error)) *Service_CreateClients_Call { - _c.Call.Return(run) - return _c -} - -// Delete provides a mock function for the type Service -func (_mock *Service) Delete(ctx context.Context, session authn.Session, id string) error { - ret := _mock.Called(ctx, session, id) - - if len(ret) == 0 { - panic("no return value specified for Delete") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) error); ok { - r0 = returnFunc(ctx, session, id) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_Delete_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Delete' -type Service_Delete_Call struct { - *mock.Call -} - -// Delete is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - id string -func (_e *Service_Expecter) Delete(ctx interface{}, session interface{}, id interface{}) *Service_Delete_Call { - return &Service_Delete_Call{Call: _e.mock.On("Delete", ctx, session, id)} -} - -func (_c *Service_Delete_Call) Run(run func(ctx context.Context, session authn.Session, id string)) *Service_Delete_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_Delete_Call) Return(err error) *Service_Delete_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_Delete_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, id string) error) *Service_Delete_Call { - _c.Call.Return(run) - return _c -} - -// Disable provides a mock function for the type Service -func (_mock *Service) Disable(ctx context.Context, session authn.Session, id string) (clients.Client, error) { - ret := _mock.Called(ctx, session, id) - - if len(ret) == 0 { - panic("no return value specified for Disable") - } - - var r0 clients.Client - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) (clients.Client, error)); ok { - return returnFunc(ctx, session, id) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) clients.Client); ok { - r0 = returnFunc(ctx, session, id) - } else { - r0 = ret.Get(0).(clients.Client) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string) error); ok { - r1 = returnFunc(ctx, session, id) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_Disable_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Disable' -type Service_Disable_Call struct { - *mock.Call -} - -// Disable is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - id string -func (_e *Service_Expecter) Disable(ctx interface{}, session interface{}, id interface{}) *Service_Disable_Call { - return &Service_Disable_Call{Call: _e.mock.On("Disable", ctx, session, id)} -} - -func (_c *Service_Disable_Call) Run(run func(ctx context.Context, session authn.Session, id string)) *Service_Disable_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_Disable_Call) Return(client clients.Client, err error) *Service_Disable_Call { - _c.Call.Return(client, err) - return _c -} - -func (_c *Service_Disable_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, id string) (clients.Client, error)) *Service_Disable_Call { - _c.Call.Return(run) - return _c -} - -// Enable provides a mock function for the type Service -func (_mock *Service) Enable(ctx context.Context, session authn.Session, id string) (clients.Client, error) { - ret := _mock.Called(ctx, session, id) - - if len(ret) == 0 { - panic("no return value specified for Enable") - } - - var r0 clients.Client - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) (clients.Client, error)); ok { - return returnFunc(ctx, session, id) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) clients.Client); ok { - r0 = returnFunc(ctx, session, id) - } else { - r0 = ret.Get(0).(clients.Client) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string) error); ok { - r1 = returnFunc(ctx, session, id) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_Enable_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Enable' -type Service_Enable_Call struct { - *mock.Call -} - -// Enable is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - id string -func (_e *Service_Expecter) Enable(ctx interface{}, session interface{}, id interface{}) *Service_Enable_Call { - return &Service_Enable_Call{Call: _e.mock.On("Enable", ctx, session, id)} -} - -func (_c *Service_Enable_Call) Run(run func(ctx context.Context, session authn.Session, id string)) *Service_Enable_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_Enable_Call) Return(client clients.Client, err error) *Service_Enable_Call { - _c.Call.Return(client, err) - return _c -} - -func (_c *Service_Enable_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, id string) (clients.Client, error)) *Service_Enable_Call { - _c.Call.Return(run) - return _c -} - -// ListAvailableActions provides a mock function for the type Service -func (_mock *Service) ListAvailableActions(ctx context.Context, session authn.Session) ([]string, error) { - ret := _mock.Called(ctx, session) - - if len(ret) == 0 { - panic("no return value specified for ListAvailableActions") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session) ([]string, error)); ok { - return returnFunc(ctx, session) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session) []string); ok { - r0 = returnFunc(ctx, session) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session) error); ok { - r1 = returnFunc(ctx, session) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_ListAvailableActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ListAvailableActions' -type Service_ListAvailableActions_Call struct { - *mock.Call -} - -// ListAvailableActions is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -func (_e *Service_Expecter) ListAvailableActions(ctx interface{}, session interface{}) *Service_ListAvailableActions_Call { - return &Service_ListAvailableActions_Call{Call: _e.mock.On("ListAvailableActions", ctx, session)} -} - -func (_c *Service_ListAvailableActions_Call) Run(run func(ctx context.Context, session authn.Session)) *Service_ListAvailableActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Service_ListAvailableActions_Call) Return(strings []string, err error) *Service_ListAvailableActions_Call { - _c.Call.Return(strings, err) - return _c -} - -func (_c *Service_ListAvailableActions_Call) RunAndReturn(run func(ctx context.Context, session authn.Session) ([]string, error)) *Service_ListAvailableActions_Call { - _c.Call.Return(run) - return _c -} - -// ListClients provides a mock function for the type Service -func (_mock *Service) ListClients(ctx context.Context, session authn.Session, pm clients.Page) (clients.ClientsPage, error) { - ret := _mock.Called(ctx, session, pm) - - if len(ret) == 0 { - panic("no return value specified for ListClients") - } - - var r0 clients.ClientsPage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, clients.Page) (clients.ClientsPage, error)); ok { - return returnFunc(ctx, session, pm) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, clients.Page) clients.ClientsPage); ok { - r0 = returnFunc(ctx, session, pm) - } else { - r0 = ret.Get(0).(clients.ClientsPage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, clients.Page) error); ok { - r1 = returnFunc(ctx, session, pm) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_ListClients_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ListClients' -type Service_ListClients_Call struct { - *mock.Call -} - -// ListClients is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - pm clients.Page -func (_e *Service_Expecter) ListClients(ctx interface{}, session interface{}, pm interface{}) *Service_ListClients_Call { - return &Service_ListClients_Call{Call: _e.mock.On("ListClients", ctx, session, pm)} -} - -func (_c *Service_ListClients_Call) Run(run func(ctx context.Context, session authn.Session, pm clients.Page)) *Service_ListClients_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 clients.Page - if args[2] != nil { - arg2 = args[2].(clients.Page) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_ListClients_Call) Return(clientsPage clients.ClientsPage, err error) *Service_ListClients_Call { - _c.Call.Return(clientsPage, err) - return _c -} - -func (_c *Service_ListClients_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, pm clients.Page) (clients.ClientsPage, error)) *Service_ListClients_Call { - _c.Call.Return(run) - return _c -} - -// ListEntityMembers provides a mock function for the type Service -func (_mock *Service) ListEntityMembers(ctx context.Context, session authn.Session, entityID string, pq roles.MembersRolePageQuery) (roles.MembersRolePage, error) { - ret := _mock.Called(ctx, session, entityID, pq) - - if len(ret) == 0 { - panic("no return value specified for ListEntityMembers") - } - - var r0 roles.MembersRolePage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, roles.MembersRolePageQuery) (roles.MembersRolePage, error)); ok { - return returnFunc(ctx, session, entityID, pq) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, roles.MembersRolePageQuery) roles.MembersRolePage); ok { - r0 = returnFunc(ctx, session, entityID, pq) - } else { - r0 = ret.Get(0).(roles.MembersRolePage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, roles.MembersRolePageQuery) error); ok { - r1 = returnFunc(ctx, session, entityID, pq) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_ListEntityMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ListEntityMembers' -type Service_ListEntityMembers_Call struct { - *mock.Call -} - -// ListEntityMembers is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - pq roles.MembersRolePageQuery -func (_e *Service_Expecter) ListEntityMembers(ctx interface{}, session interface{}, entityID interface{}, pq interface{}) *Service_ListEntityMembers_Call { - return &Service_ListEntityMembers_Call{Call: _e.mock.On("ListEntityMembers", ctx, session, entityID, pq)} -} - -func (_c *Service_ListEntityMembers_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, pq roles.MembersRolePageQuery)) *Service_ListEntityMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 roles.MembersRolePageQuery - if args[3] != nil { - arg3 = args[3].(roles.MembersRolePageQuery) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_ListEntityMembers_Call) Return(membersRolePage roles.MembersRolePage, err error) *Service_ListEntityMembers_Call { - _c.Call.Return(membersRolePage, err) - return _c -} - -func (_c *Service_ListEntityMembers_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, pq roles.MembersRolePageQuery) (roles.MembersRolePage, error)) *Service_ListEntityMembers_Call { - _c.Call.Return(run) - return _c -} - -// ListUserClients provides a mock function for the type Service -func (_mock *Service) ListUserClients(ctx context.Context, session authn.Session, userID string, pm clients.Page) (clients.ClientsPage, error) { - ret := _mock.Called(ctx, session, userID, pm) - - if len(ret) == 0 { - panic("no return value specified for ListUserClients") - } - - var r0 clients.ClientsPage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, clients.Page) (clients.ClientsPage, error)); ok { - return returnFunc(ctx, session, userID, pm) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, clients.Page) clients.ClientsPage); ok { - r0 = returnFunc(ctx, session, userID, pm) - } else { - r0 = ret.Get(0).(clients.ClientsPage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, clients.Page) error); ok { - r1 = returnFunc(ctx, session, userID, pm) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_ListUserClients_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ListUserClients' -type Service_ListUserClients_Call struct { - *mock.Call -} - -// ListUserClients is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - userID string -// - pm clients.Page -func (_e *Service_Expecter) ListUserClients(ctx interface{}, session interface{}, userID interface{}, pm interface{}) *Service_ListUserClients_Call { - return &Service_ListUserClients_Call{Call: _e.mock.On("ListUserClients", ctx, session, userID, pm)} -} - -func (_c *Service_ListUserClients_Call) Run(run func(ctx context.Context, session authn.Session, userID string, pm clients.Page)) *Service_ListUserClients_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 clients.Page - if args[3] != nil { - arg3 = args[3].(clients.Page) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_ListUserClients_Call) Return(clientsPage clients.ClientsPage, err error) *Service_ListUserClients_Call { - _c.Call.Return(clientsPage, err) - return _c -} - -func (_c *Service_ListUserClients_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, userID string, pm clients.Page) (clients.ClientsPage, error)) *Service_ListUserClients_Call { - _c.Call.Return(run) - return _c -} - -// RemoveEntityMembers provides a mock function for the type Service -func (_mock *Service) RemoveEntityMembers(ctx context.Context, session authn.Session, entityID string, members []string) error { - ret := _mock.Called(ctx, session, entityID, members) - - if len(ret) == 0 { - panic("no return value specified for RemoveEntityMembers") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, []string) error); ok { - r0 = returnFunc(ctx, session, entityID, members) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RemoveEntityMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveEntityMembers' -type Service_RemoveEntityMembers_Call struct { - *mock.Call -} - -// RemoveEntityMembers is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - members []string -func (_e *Service_Expecter) RemoveEntityMembers(ctx interface{}, session interface{}, entityID interface{}, members interface{}) *Service_RemoveEntityMembers_Call { - return &Service_RemoveEntityMembers_Call{Call: _e.mock.On("RemoveEntityMembers", ctx, session, entityID, members)} -} - -func (_c *Service_RemoveEntityMembers_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, members []string)) *Service_RemoveEntityMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 []string - if args[3] != nil { - arg3 = args[3].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_RemoveEntityMembers_Call) Return(err error) *Service_RemoveEntityMembers_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RemoveEntityMembers_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, members []string) error) *Service_RemoveEntityMembers_Call { - _c.Call.Return(run) - return _c -} - -// RemoveMemberFromAllRoles provides a mock function for the type Service -func (_mock *Service) RemoveMemberFromAllRoles(ctx context.Context, session authn.Session, memberID string) error { - ret := _mock.Called(ctx, session, memberID) - - if len(ret) == 0 { - panic("no return value specified for RemoveMemberFromAllRoles") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) error); ok { - r0 = returnFunc(ctx, session, memberID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RemoveMemberFromAllRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveMemberFromAllRoles' -type Service_RemoveMemberFromAllRoles_Call struct { - *mock.Call -} - -// RemoveMemberFromAllRoles is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - memberID string -func (_e *Service_Expecter) RemoveMemberFromAllRoles(ctx interface{}, session interface{}, memberID interface{}) *Service_RemoveMemberFromAllRoles_Call { - return &Service_RemoveMemberFromAllRoles_Call{Call: _e.mock.On("RemoveMemberFromAllRoles", ctx, session, memberID)} -} - -func (_c *Service_RemoveMemberFromAllRoles_Call) Run(run func(ctx context.Context, session authn.Session, memberID string)) *Service_RemoveMemberFromAllRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_RemoveMemberFromAllRoles_Call) Return(err error) *Service_RemoveMemberFromAllRoles_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RemoveMemberFromAllRoles_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, memberID string) error) *Service_RemoveMemberFromAllRoles_Call { - _c.Call.Return(run) - return _c -} - -// RemoveParentGroup provides a mock function for the type Service -func (_mock *Service) RemoveParentGroup(ctx context.Context, session authn.Session, id string) error { - ret := _mock.Called(ctx, session, id) - - if len(ret) == 0 { - panic("no return value specified for RemoveParentGroup") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) error); ok { - r0 = returnFunc(ctx, session, id) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RemoveParentGroup_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveParentGroup' -type Service_RemoveParentGroup_Call struct { - *mock.Call -} - -// RemoveParentGroup is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - id string -func (_e *Service_Expecter) RemoveParentGroup(ctx interface{}, session interface{}, id interface{}) *Service_RemoveParentGroup_Call { - return &Service_RemoveParentGroup_Call{Call: _e.mock.On("RemoveParentGroup", ctx, session, id)} -} - -func (_c *Service_RemoveParentGroup_Call) Run(run func(ctx context.Context, session authn.Session, id string)) *Service_RemoveParentGroup_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_RemoveParentGroup_Call) Return(err error) *Service_RemoveParentGroup_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RemoveParentGroup_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, id string) error) *Service_RemoveParentGroup_Call { - _c.Call.Return(run) - return _c -} - -// RemoveRole provides a mock function for the type Service -func (_mock *Service) RemoveRole(ctx context.Context, session authn.Session, entityID string, roleID string) error { - ret := _mock.Called(ctx, session, entityID, roleID) - - if len(ret) == 0 { - panic("no return value specified for RemoveRole") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) error); ok { - r0 = returnFunc(ctx, session, entityID, roleID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RemoveRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveRole' -type Service_RemoveRole_Call struct { - *mock.Call -} - -// RemoveRole is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -func (_e *Service_Expecter) RemoveRole(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}) *Service_RemoveRole_Call { - return &Service_RemoveRole_Call{Call: _e.mock.On("RemoveRole", ctx, session, entityID, roleID)} -} - -func (_c *Service_RemoveRole_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string)) *Service_RemoveRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_RemoveRole_Call) Return(err error) *Service_RemoveRole_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RemoveRole_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string) error) *Service_RemoveRole_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveAllRoles provides a mock function for the type Service -func (_mock *Service) RetrieveAllRoles(ctx context.Context, session authn.Session, entityID string, limit uint64, offset uint64) (roles.RolePage, error) { - ret := _mock.Called(ctx, session, entityID, limit, offset) - - if len(ret) == 0 { - panic("no return value specified for RetrieveAllRoles") - } - - var r0 roles.RolePage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, uint64, uint64) (roles.RolePage, error)); ok { - return returnFunc(ctx, session, entityID, limit, offset) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, uint64, uint64) roles.RolePage); ok { - r0 = returnFunc(ctx, session, entityID, limit, offset) - } else { - r0 = ret.Get(0).(roles.RolePage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, uint64, uint64) error); ok { - r1 = returnFunc(ctx, session, entityID, limit, offset) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RetrieveAllRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveAllRoles' -type Service_RetrieveAllRoles_Call struct { - *mock.Call -} - -// RetrieveAllRoles is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - limit uint64 -// - offset uint64 -func (_e *Service_Expecter) RetrieveAllRoles(ctx interface{}, session interface{}, entityID interface{}, limit interface{}, offset interface{}) *Service_RetrieveAllRoles_Call { - return &Service_RetrieveAllRoles_Call{Call: _e.mock.On("RetrieveAllRoles", ctx, session, entityID, limit, offset)} -} - -func (_c *Service_RetrieveAllRoles_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, limit uint64, offset uint64)) *Service_RetrieveAllRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 uint64 - if args[3] != nil { - arg3 = args[3].(uint64) - } - var arg4 uint64 - if args[4] != nil { - arg4 = args[4].(uint64) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_RetrieveAllRoles_Call) Return(rolePage roles.RolePage, err error) *Service_RetrieveAllRoles_Call { - _c.Call.Return(rolePage, err) - return _c -} - -func (_c *Service_RetrieveAllRoles_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, limit uint64, offset uint64) (roles.RolePage, error)) *Service_RetrieveAllRoles_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveRole provides a mock function for the type Service -func (_mock *Service) RetrieveRole(ctx context.Context, session authn.Session, entityID string, roleID string) (roles.Role, error) { - ret := _mock.Called(ctx, session, entityID, roleID) - - if len(ret) == 0 { - panic("no return value specified for RetrieveRole") - } - - var r0 roles.Role - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) (roles.Role, error)); ok { - return returnFunc(ctx, session, entityID, roleID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) roles.Role); ok { - r0 = returnFunc(ctx, session, entityID, roleID) - } else { - r0 = ret.Get(0).(roles.Role) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RetrieveRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveRole' -type Service_RetrieveRole_Call struct { - *mock.Call -} - -// RetrieveRole is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -func (_e *Service_Expecter) RetrieveRole(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}) *Service_RetrieveRole_Call { - return &Service_RetrieveRole_Call{Call: _e.mock.On("RetrieveRole", ctx, session, entityID, roleID)} -} - -func (_c *Service_RetrieveRole_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string)) *Service_RetrieveRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_RetrieveRole_Call) Return(role roles.Role, err error) *Service_RetrieveRole_Call { - _c.Call.Return(role, err) - return _c -} - -func (_c *Service_RetrieveRole_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string) (roles.Role, error)) *Service_RetrieveRole_Call { - _c.Call.Return(run) - return _c -} - -// RoleAddActions provides a mock function for the type Service -func (_mock *Service) RoleAddActions(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string) ([]string, error) { - ret := _mock.Called(ctx, session, entityID, roleID, actions) - - if len(ret) == 0 { - panic("no return value specified for RoleAddActions") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) ([]string, error)); ok { - return returnFunc(ctx, session, entityID, roleID, actions) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) []string); ok { - r0 = returnFunc(ctx, session, entityID, roleID, actions) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, []string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID, actions) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RoleAddActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleAddActions' -type Service_RoleAddActions_Call struct { - *mock.Call -} - -// RoleAddActions is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - actions []string -func (_e *Service_Expecter) RoleAddActions(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, actions interface{}) *Service_RoleAddActions_Call { - return &Service_RoleAddActions_Call{Call: _e.mock.On("RoleAddActions", ctx, session, entityID, roleID, actions)} -} - -func (_c *Service_RoleAddActions_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string)) *Service_RoleAddActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_RoleAddActions_Call) Return(ops []string, err error) *Service_RoleAddActions_Call { - _c.Call.Return(ops, err) - return _c -} - -func (_c *Service_RoleAddActions_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string) ([]string, error)) *Service_RoleAddActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleAddMembers provides a mock function for the type Service -func (_mock *Service) RoleAddMembers(ctx context.Context, session authn.Session, entityID string, roleID string, members []string) ([]string, error) { - ret := _mock.Called(ctx, session, entityID, roleID, members) - - if len(ret) == 0 { - panic("no return value specified for RoleAddMembers") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) ([]string, error)); ok { - return returnFunc(ctx, session, entityID, roleID, members) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) []string); ok { - r0 = returnFunc(ctx, session, entityID, roleID, members) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, []string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID, members) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RoleAddMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleAddMembers' -type Service_RoleAddMembers_Call struct { - *mock.Call -} - -// RoleAddMembers is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - members []string -func (_e *Service_Expecter) RoleAddMembers(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, members interface{}) *Service_RoleAddMembers_Call { - return &Service_RoleAddMembers_Call{Call: _e.mock.On("RoleAddMembers", ctx, session, entityID, roleID, members)} -} - -func (_c *Service_RoleAddMembers_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, members []string)) *Service_RoleAddMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_RoleAddMembers_Call) Return(strings []string, err error) *Service_RoleAddMembers_Call { - _c.Call.Return(strings, err) - return _c -} - -func (_c *Service_RoleAddMembers_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, members []string) ([]string, error)) *Service_RoleAddMembers_Call { - _c.Call.Return(run) - return _c -} - -// RoleCheckActionsExists provides a mock function for the type Service -func (_mock *Service) RoleCheckActionsExists(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string) (bool, error) { - ret := _mock.Called(ctx, session, entityID, roleID, actions) - - if len(ret) == 0 { - panic("no return value specified for RoleCheckActionsExists") - } - - var r0 bool - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) (bool, error)); ok { - return returnFunc(ctx, session, entityID, roleID, actions) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) bool); ok { - r0 = returnFunc(ctx, session, entityID, roleID, actions) - } else { - r0 = ret.Get(0).(bool) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, []string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID, actions) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RoleCheckActionsExists_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleCheckActionsExists' -type Service_RoleCheckActionsExists_Call struct { - *mock.Call -} - -// RoleCheckActionsExists is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - actions []string -func (_e *Service_Expecter) RoleCheckActionsExists(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, actions interface{}) *Service_RoleCheckActionsExists_Call { - return &Service_RoleCheckActionsExists_Call{Call: _e.mock.On("RoleCheckActionsExists", ctx, session, entityID, roleID, actions)} -} - -func (_c *Service_RoleCheckActionsExists_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string)) *Service_RoleCheckActionsExists_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_RoleCheckActionsExists_Call) Return(b bool, err error) *Service_RoleCheckActionsExists_Call { - _c.Call.Return(b, err) - return _c -} - -func (_c *Service_RoleCheckActionsExists_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string) (bool, error)) *Service_RoleCheckActionsExists_Call { - _c.Call.Return(run) - return _c -} - -// RoleCheckMembersExists provides a mock function for the type Service -func (_mock *Service) RoleCheckMembersExists(ctx context.Context, session authn.Session, entityID string, roleID string, members []string) (bool, error) { - ret := _mock.Called(ctx, session, entityID, roleID, members) - - if len(ret) == 0 { - panic("no return value specified for RoleCheckMembersExists") - } - - var r0 bool - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) (bool, error)); ok { - return returnFunc(ctx, session, entityID, roleID, members) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) bool); ok { - r0 = returnFunc(ctx, session, entityID, roleID, members) - } else { - r0 = ret.Get(0).(bool) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, []string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID, members) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RoleCheckMembersExists_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleCheckMembersExists' -type Service_RoleCheckMembersExists_Call struct { - *mock.Call -} - -// RoleCheckMembersExists is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - members []string -func (_e *Service_Expecter) RoleCheckMembersExists(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, members interface{}) *Service_RoleCheckMembersExists_Call { - return &Service_RoleCheckMembersExists_Call{Call: _e.mock.On("RoleCheckMembersExists", ctx, session, entityID, roleID, members)} -} - -func (_c *Service_RoleCheckMembersExists_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, members []string)) *Service_RoleCheckMembersExists_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_RoleCheckMembersExists_Call) Return(b bool, err error) *Service_RoleCheckMembersExists_Call { - _c.Call.Return(b, err) - return _c -} - -func (_c *Service_RoleCheckMembersExists_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, members []string) (bool, error)) *Service_RoleCheckMembersExists_Call { - _c.Call.Return(run) - return _c -} - -// RoleListActions provides a mock function for the type Service -func (_mock *Service) RoleListActions(ctx context.Context, session authn.Session, entityID string, roleID string) ([]string, error) { - ret := _mock.Called(ctx, session, entityID, roleID) - - if len(ret) == 0 { - panic("no return value specified for RoleListActions") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) ([]string, error)); ok { - return returnFunc(ctx, session, entityID, roleID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) []string); ok { - r0 = returnFunc(ctx, session, entityID, roleID) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RoleListActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleListActions' -type Service_RoleListActions_Call struct { - *mock.Call -} - -// RoleListActions is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -func (_e *Service_Expecter) RoleListActions(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}) *Service_RoleListActions_Call { - return &Service_RoleListActions_Call{Call: _e.mock.On("RoleListActions", ctx, session, entityID, roleID)} -} - -func (_c *Service_RoleListActions_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string)) *Service_RoleListActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_RoleListActions_Call) Return(strings []string, err error) *Service_RoleListActions_Call { - _c.Call.Return(strings, err) - return _c -} - -func (_c *Service_RoleListActions_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string) ([]string, error)) *Service_RoleListActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleListMembers provides a mock function for the type Service -func (_mock *Service) RoleListMembers(ctx context.Context, session authn.Session, entityID string, roleID string, limit uint64, offset uint64) (roles.MembersPage, error) { - ret := _mock.Called(ctx, session, entityID, roleID, limit, offset) - - if len(ret) == 0 { - panic("no return value specified for RoleListMembers") - } - - var r0 roles.MembersPage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, uint64, uint64) (roles.MembersPage, error)); ok { - return returnFunc(ctx, session, entityID, roleID, limit, offset) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, uint64, uint64) roles.MembersPage); ok { - r0 = returnFunc(ctx, session, entityID, roleID, limit, offset) - } else { - r0 = ret.Get(0).(roles.MembersPage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, uint64, uint64) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID, limit, offset) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RoleListMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleListMembers' -type Service_RoleListMembers_Call struct { - *mock.Call -} - -// RoleListMembers is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - limit uint64 -// - offset uint64 -func (_e *Service_Expecter) RoleListMembers(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, limit interface{}, offset interface{}) *Service_RoleListMembers_Call { - return &Service_RoleListMembers_Call{Call: _e.mock.On("RoleListMembers", ctx, session, entityID, roleID, limit, offset)} -} - -func (_c *Service_RoleListMembers_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, limit uint64, offset uint64)) *Service_RoleListMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 uint64 - if args[4] != nil { - arg4 = args[4].(uint64) - } - var arg5 uint64 - if args[5] != nil { - arg5 = args[5].(uint64) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - arg5, - ) - }) - return _c -} - -func (_c *Service_RoleListMembers_Call) Return(membersPage roles.MembersPage, err error) *Service_RoleListMembers_Call { - _c.Call.Return(membersPage, err) - return _c -} - -func (_c *Service_RoleListMembers_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, limit uint64, offset uint64) (roles.MembersPage, error)) *Service_RoleListMembers_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveActions provides a mock function for the type Service -func (_mock *Service) RoleRemoveActions(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string) error { - ret := _mock.Called(ctx, session, entityID, roleID, actions) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveActions") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) error); ok { - r0 = returnFunc(ctx, session, entityID, roleID, actions) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RoleRemoveActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveActions' -type Service_RoleRemoveActions_Call struct { - *mock.Call -} - -// RoleRemoveActions is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - actions []string -func (_e *Service_Expecter) RoleRemoveActions(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, actions interface{}) *Service_RoleRemoveActions_Call { - return &Service_RoleRemoveActions_Call{Call: _e.mock.On("RoleRemoveActions", ctx, session, entityID, roleID, actions)} -} - -func (_c *Service_RoleRemoveActions_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string)) *Service_RoleRemoveActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_RoleRemoveActions_Call) Return(err error) *Service_RoleRemoveActions_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RoleRemoveActions_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string) error) *Service_RoleRemoveActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveAllActions provides a mock function for the type Service -func (_mock *Service) RoleRemoveAllActions(ctx context.Context, session authn.Session, entityID string, roleID string) error { - ret := _mock.Called(ctx, session, entityID, roleID) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveAllActions") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) error); ok { - r0 = returnFunc(ctx, session, entityID, roleID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RoleRemoveAllActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveAllActions' -type Service_RoleRemoveAllActions_Call struct { - *mock.Call -} - -// RoleRemoveAllActions is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -func (_e *Service_Expecter) RoleRemoveAllActions(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}) *Service_RoleRemoveAllActions_Call { - return &Service_RoleRemoveAllActions_Call{Call: _e.mock.On("RoleRemoveAllActions", ctx, session, entityID, roleID)} -} - -func (_c *Service_RoleRemoveAllActions_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string)) *Service_RoleRemoveAllActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_RoleRemoveAllActions_Call) Return(err error) *Service_RoleRemoveAllActions_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RoleRemoveAllActions_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string) error) *Service_RoleRemoveAllActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveAllMembers provides a mock function for the type Service -func (_mock *Service) RoleRemoveAllMembers(ctx context.Context, session authn.Session, entityID string, roleID string) error { - ret := _mock.Called(ctx, session, entityID, roleID) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveAllMembers") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) error); ok { - r0 = returnFunc(ctx, session, entityID, roleID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RoleRemoveAllMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveAllMembers' -type Service_RoleRemoveAllMembers_Call struct { - *mock.Call -} - -// RoleRemoveAllMembers is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -func (_e *Service_Expecter) RoleRemoveAllMembers(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}) *Service_RoleRemoveAllMembers_Call { - return &Service_RoleRemoveAllMembers_Call{Call: _e.mock.On("RoleRemoveAllMembers", ctx, session, entityID, roleID)} -} - -func (_c *Service_RoleRemoveAllMembers_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string)) *Service_RoleRemoveAllMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_RoleRemoveAllMembers_Call) Return(err error) *Service_RoleRemoveAllMembers_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RoleRemoveAllMembers_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string) error) *Service_RoleRemoveAllMembers_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveMembers provides a mock function for the type Service -func (_mock *Service) RoleRemoveMembers(ctx context.Context, session authn.Session, entityID string, roleID string, members []string) error { - ret := _mock.Called(ctx, session, entityID, roleID, members) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveMembers") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) error); ok { - r0 = returnFunc(ctx, session, entityID, roleID, members) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RoleRemoveMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveMembers' -type Service_RoleRemoveMembers_Call struct { - *mock.Call -} - -// RoleRemoveMembers is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - members []string -func (_e *Service_Expecter) RoleRemoveMembers(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, members interface{}) *Service_RoleRemoveMembers_Call { - return &Service_RoleRemoveMembers_Call{Call: _e.mock.On("RoleRemoveMembers", ctx, session, entityID, roleID, members)} -} - -func (_c *Service_RoleRemoveMembers_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, members []string)) *Service_RoleRemoveMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_RoleRemoveMembers_Call) Return(err error) *Service_RoleRemoveMembers_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RoleRemoveMembers_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, members []string) error) *Service_RoleRemoveMembers_Call { - _c.Call.Return(run) - return _c -} - -// SetParentGroup provides a mock function for the type Service -func (_mock *Service) SetParentGroup(ctx context.Context, session authn.Session, parentGroupID string, id string) error { - ret := _mock.Called(ctx, session, parentGroupID, id) - - if len(ret) == 0 { - panic("no return value specified for SetParentGroup") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) error); ok { - r0 = returnFunc(ctx, session, parentGroupID, id) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_SetParentGroup_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'SetParentGroup' -type Service_SetParentGroup_Call struct { - *mock.Call -} - -// SetParentGroup is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - parentGroupID string -// - id string -func (_e *Service_Expecter) SetParentGroup(ctx interface{}, session interface{}, parentGroupID interface{}, id interface{}) *Service_SetParentGroup_Call { - return &Service_SetParentGroup_Call{Call: _e.mock.On("SetParentGroup", ctx, session, parentGroupID, id)} -} - -func (_c *Service_SetParentGroup_Call) Run(run func(ctx context.Context, session authn.Session, parentGroupID string, id string)) *Service_SetParentGroup_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_SetParentGroup_Call) Return(err error) *Service_SetParentGroup_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_SetParentGroup_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, parentGroupID string, id string) error) *Service_SetParentGroup_Call { - _c.Call.Return(run) - return _c -} - -// Update provides a mock function for the type Service -func (_mock *Service) Update(ctx context.Context, session authn.Session, client clients.Client) (clients.Client, error) { - ret := _mock.Called(ctx, session, client) - - if len(ret) == 0 { - panic("no return value specified for Update") - } - - var r0 clients.Client - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, clients.Client) (clients.Client, error)); ok { - return returnFunc(ctx, session, client) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, clients.Client) clients.Client); ok { - r0 = returnFunc(ctx, session, client) - } else { - r0 = ret.Get(0).(clients.Client) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, clients.Client) error); ok { - r1 = returnFunc(ctx, session, client) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_Update_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Update' -type Service_Update_Call struct { - *mock.Call -} - -// Update is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - client clients.Client -func (_e *Service_Expecter) Update(ctx interface{}, session interface{}, client interface{}) *Service_Update_Call { - return &Service_Update_Call{Call: _e.mock.On("Update", ctx, session, client)} -} - -func (_c *Service_Update_Call) Run(run func(ctx context.Context, session authn.Session, client clients.Client)) *Service_Update_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 clients.Client - if args[2] != nil { - arg2 = args[2].(clients.Client) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_Update_Call) Return(client1 clients.Client, err error) *Service_Update_Call { - _c.Call.Return(client1, err) - return _c -} - -func (_c *Service_Update_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, client clients.Client) (clients.Client, error)) *Service_Update_Call { - _c.Call.Return(run) - return _c -} - -// UpdateRoleName provides a mock function for the type Service -func (_mock *Service) UpdateRoleName(ctx context.Context, session authn.Session, entityID string, roleID string, newRoleName string) (roles.Role, error) { - ret := _mock.Called(ctx, session, entityID, roleID, newRoleName) - - if len(ret) == 0 { - panic("no return value specified for UpdateRoleName") - } - - var r0 roles.Role - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, string) (roles.Role, error)); ok { - return returnFunc(ctx, session, entityID, roleID, newRoleName) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, string) roles.Role); ok { - r0 = returnFunc(ctx, session, entityID, roleID, newRoleName) - } else { - r0 = ret.Get(0).(roles.Role) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID, newRoleName) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_UpdateRoleName_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateRoleName' -type Service_UpdateRoleName_Call struct { - *mock.Call -} - -// UpdateRoleName is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - newRoleName string -func (_e *Service_Expecter) UpdateRoleName(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, newRoleName interface{}) *Service_UpdateRoleName_Call { - return &Service_UpdateRoleName_Call{Call: _e.mock.On("UpdateRoleName", ctx, session, entityID, roleID, newRoleName)} -} - -func (_c *Service_UpdateRoleName_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, newRoleName string)) *Service_UpdateRoleName_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 string - if args[4] != nil { - arg4 = args[4].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_UpdateRoleName_Call) Return(role roles.Role, err error) *Service_UpdateRoleName_Call { - _c.Call.Return(role, err) - return _c -} - -func (_c *Service_UpdateRoleName_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, newRoleName string) (roles.Role, error)) *Service_UpdateRoleName_Call { - _c.Call.Return(run) - return _c -} - -// UpdateSecret provides a mock function for the type Service -func (_mock *Service) UpdateSecret(ctx context.Context, session authn.Session, id string, key string) (clients.Client, error) { - ret := _mock.Called(ctx, session, id, key) - - if len(ret) == 0 { - panic("no return value specified for UpdateSecret") - } - - var r0 clients.Client - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) (clients.Client, error)); ok { - return returnFunc(ctx, session, id, key) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) clients.Client); ok { - r0 = returnFunc(ctx, session, id, key) - } else { - r0 = ret.Get(0).(clients.Client) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string) error); ok { - r1 = returnFunc(ctx, session, id, key) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_UpdateSecret_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateSecret' -type Service_UpdateSecret_Call struct { - *mock.Call -} - -// UpdateSecret is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - id string -// - key string -func (_e *Service_Expecter) UpdateSecret(ctx interface{}, session interface{}, id interface{}, key interface{}) *Service_UpdateSecret_Call { - return &Service_UpdateSecret_Call{Call: _e.mock.On("UpdateSecret", ctx, session, id, key)} -} - -func (_c *Service_UpdateSecret_Call) Run(run func(ctx context.Context, session authn.Session, id string, key string)) *Service_UpdateSecret_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_UpdateSecret_Call) Return(client clients.Client, err error) *Service_UpdateSecret_Call { - _c.Call.Return(client, err) - return _c -} - -func (_c *Service_UpdateSecret_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, id string, key string) (clients.Client, error)) *Service_UpdateSecret_Call { - _c.Call.Return(run) - return _c -} - -// UpdateTags provides a mock function for the type Service -func (_mock *Service) UpdateTags(ctx context.Context, session authn.Session, client clients.Client) (clients.Client, error) { - ret := _mock.Called(ctx, session, client) - - if len(ret) == 0 { - panic("no return value specified for UpdateTags") - } - - var r0 clients.Client - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, clients.Client) (clients.Client, error)); ok { - return returnFunc(ctx, session, client) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, clients.Client) clients.Client); ok { - r0 = returnFunc(ctx, session, client) - } else { - r0 = ret.Get(0).(clients.Client) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, clients.Client) error); ok { - r1 = returnFunc(ctx, session, client) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_UpdateTags_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateTags' -type Service_UpdateTags_Call struct { - *mock.Call -} - -// UpdateTags is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - client clients.Client -func (_e *Service_Expecter) UpdateTags(ctx interface{}, session interface{}, client interface{}) *Service_UpdateTags_Call { - return &Service_UpdateTags_Call{Call: _e.mock.On("UpdateTags", ctx, session, client)} -} - -func (_c *Service_UpdateTags_Call) Run(run func(ctx context.Context, session authn.Session, client clients.Client)) *Service_UpdateTags_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 clients.Client - if args[2] != nil { - arg2 = args[2].(clients.Client) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_UpdateTags_Call) Return(client1 clients.Client, err error) *Service_UpdateTags_Call { - _c.Call.Return(client1, err) - return _c -} - -func (_c *Service_UpdateTags_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, client clients.Client) (clients.Client, error)) *Service_UpdateTags_Call { - _c.Call.Return(run) - return _c -} - -// View provides a mock function for the type Service -func (_mock *Service) View(ctx context.Context, session authn.Session, id string, withRoles bool) (clients.Client, error) { - ret := _mock.Called(ctx, session, id, withRoles) - - if len(ret) == 0 { - panic("no return value specified for View") - } - - var r0 clients.Client - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, bool) (clients.Client, error)); ok { - return returnFunc(ctx, session, id, withRoles) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, bool) clients.Client); ok { - r0 = returnFunc(ctx, session, id, withRoles) - } else { - r0 = ret.Get(0).(clients.Client) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, bool) error); ok { - r1 = returnFunc(ctx, session, id, withRoles) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_View_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'View' -type Service_View_Call struct { - *mock.Call -} - -// View is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - id string -// - withRoles bool -func (_e *Service_Expecter) View(ctx interface{}, session interface{}, id interface{}, withRoles interface{}) *Service_View_Call { - return &Service_View_Call{Call: _e.mock.On("View", ctx, session, id, withRoles)} -} - -func (_c *Service_View_Call) Run(run func(ctx context.Context, session authn.Session, id string, withRoles bool)) *Service_View_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 bool - if args[3] != nil { - arg3 = args[3].(bool) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_View_Call) Return(client clients.Client, err error) *Service_View_Call { - _c.Call.Return(client, err) - return _c -} - -func (_c *Service_View_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, id string, withRoles bool) (clients.Client, error)) *Service_View_Call { - _c.Call.Return(run) - return _c -} diff --git a/clients/operations/operations.go b/clients/operations/operations.go deleted file mode 100644 index aeb0f9434..000000000 --- a/clients/operations/operations.go +++ /dev/null @@ -1,77 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package operations - -import ( - "github.com/absmach/magistrala/pkg/permissions" -) - -// Client Operations. -const ( - OpViewClient permissions.Operation = iota - OpUpdateClient - OpUpdateClientTags - OpUpdateClientSecret - OpEnableClient - OpDisableClient - OpDeleteClient - OpSetParentGroup - OpRemoveParentGroup - OpConnectToChannel - OpDisconnectFromChannel - OpListUserClients -) - -func OperationDetails() map[permissions.Operation]permissions.OperationDetails { - return map[permissions.Operation]permissions.OperationDetails{ - OpViewClient: { - Name: "view", - PermissionRequired: true, - }, - OpUpdateClient: { - Name: "update", - PermissionRequired: true, - }, - OpUpdateClientTags: { - Name: "update_tags", - PermissionRequired: true, - }, - OpUpdateClientSecret: { - Name: "update_secret", - PermissionRequired: true, - }, - OpEnableClient: { - Name: "enable", - PermissionRequired: true, - }, - OpDisableClient: { - Name: "disable", - PermissionRequired: true, - }, - OpDeleteClient: { - Name: "delete", - PermissionRequired: true, - }, - OpSetParentGroup: { - Name: "set_parent_group", - PermissionRequired: true, - }, - OpRemoveParentGroup: { - Name: "remove_parent_group", - PermissionRequired: true, - }, - OpConnectToChannel: { - Name: "connect_to_channel", - PermissionRequired: true, - }, - OpDisconnectFromChannel: { - Name: "disconnect_from_channel", - PermissionRequired: true, - }, - OpListUserClients: { - Name: "list_user_clients", - PermissionRequired: false, // hardcoded to superadmin - }, - } -} diff --git a/clients/postgres/clients.go b/clients/postgres/clients.go deleted file mode 100644 index 5bda14728..000000000 --- a/clients/postgres/clients.go +++ /dev/null @@ -1,1597 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package postgres - -import ( - "context" - "database/sql" - "encoding/json" - "fmt" - "strings" - "time" - - api "github.com/absmach/magistrala/api/http" - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/clients" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/connections" - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - "github.com/absmach/magistrala/pkg/policies" - "github.com/absmach/magistrala/pkg/postgres" - "github.com/absmach/magistrala/pkg/roles" - rolesPostgres "github.com/absmach/magistrala/pkg/roles/repo/postgres" - "github.com/jackc/pgtype" - "github.com/lib/pq" -) - -const ( - entityTableName = "clients" - entityIDColumnName = "id" - rolesTableNamePrefix = "clients" -) - -var _ clients.Repository = (*clientRepo)(nil) - -type clientRepo struct { - DB postgres.Database - eh errors.Handler - rolesPostgres.Repository -} - -// NewRepository instantiates a PostgreSQL -// implementation of Clients repository. -func NewRepository(db postgres.Database) clients.Repository { - repo := rolesPostgres.NewRepository(db, policies.ClientType, rolesTableNamePrefix, entityTableName, entityIDColumnName) - errHandlerOptions := []errors.HandlerOption{ - postgres.WithDuplicateErrors(NewDuplicateErrors()), - } - return &clientRepo{ - DB: db, - eh: postgres.NewErrorHandler(errHandlerOptions...), - Repository: repo, - } -} - -func (repo *clientRepo) Save(ctx context.Context, cls ...clients.Client) ([]clients.Client, error) { - var dbClients []DBClient - - for _, client := range cls { - dbcli, err := ToDBClient(client) - if err != nil { - return []clients.Client{}, errors.Wrap(repoerr.ErrCreateEntity, err) - } - dbClients = append(dbClients, dbcli) - } - q := `INSERT INTO clients (id, name, tags, domain_id, parent_group_id, identity, secret, metadata, private_metadata, created_at, updated_at, updated_by, status) - VALUES (:id, :name, :tags, :domain_id, :parent_group_id, :identity, :secret, :metadata, :private_metadata, :created_at, :updated_at, :updated_by, :status) - RETURNING id, name, tags, identity, secret, metadata, private_metadata, COALESCE(domain_id, '') AS domain_id, COALESCE(parent_group_id, '') AS parent_group_id, status, created_at, updated_at, updated_by` - - row, err := repo.DB.NamedQueryContext(ctx, q, dbClients) - if err != nil { - return []clients.Client{}, repo.eh.HandleError(repoerr.ErrCreateEntity, err) - } - - defer row.Close() - - var reClients []clients.Client - for row.Next() { - dbcli := DBClient{} - if err := row.StructScan(&dbcli); err != nil { - return []clients.Client{}, repo.eh.HandleError(repoerr.ErrFailedOpDB, err) - } - - client, err := ToClient(dbcli) - if err != nil { - return []clients.Client{}, errors.Wrap(repoerr.ErrFailedOpDB, err) - } - reClients = append(reClients, client) - } - return reClients, nil -} - -func (repo *clientRepo) RetrieveBySecret(ctx context.Context, key, id string, prefix authn.AuthPrefix) (clients.Client, error) { - q := fmt.Sprintf(`SELECT id, name, tags, COALESCE(domain_id, '') AS domain_id, COALESCE(parent_group_id, '') AS parent_group_id, identity, secret, metadata, private_metadata, created_at, updated_at, updated_by, status - FROM clients - WHERE secret = :secret AND status = %d`, clients.EnabledStatus) - switch prefix { - case authn.DomainAuth: - q += " AND domain_id = :domain_id" - case authn.BasicAuth: - q += " AND id = :id" - default: - return clients.Client{}, repoerr.ErrNotFound - } - - dbc := DBClient{ - Secret: key, - Domain: id, - ID: id, - } - - rows, err := repo.DB.NamedQueryContext(ctx, q, dbc) - if err != nil { - return clients.Client{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - defer rows.Close() - - dbc = DBClient{} - if rows.Next() { - if err = rows.StructScan(&dbc); err != nil { - return clients.Client{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - - client, err := ToClient(dbc) - if err != nil { - return clients.Client{}, errors.Wrap(repoerr.ErrFailedOpDB, err) - } - - return client, nil - } - - return clients.Client{}, repoerr.ErrNotFound -} - -func (repo *clientRepo) Update(ctx context.Context, client clients.Client) (clients.Client, error) { - var query []string - var upq string - if client.Name != "" { - query = append(query, "name = :name,") - } - if client.Metadata != nil { - query = append(query, "metadata = :metadata,") - } - if client.PrivateMetadata != nil { - query = append(query, "private_metadata = :private_metadata,") - } - if len(query) > 0 { - upq = strings.Join(query, " ") - } - - q := fmt.Sprintf(`UPDATE clients SET %s updated_at = :updated_at, updated_by = :updated_by - WHERE id = :id AND status = :status - RETURNING id, name, tags, identity, secret, metadata, private_metadata, COALESCE(domain_id, '') AS domain_id, COALESCE(parent_group_id, '') AS parent_group_id, status, created_at, updated_at, updated_by`, - upq) - client.Status = clients.EnabledStatus - return repo.update(ctx, client, q) -} - -func (repo *clientRepo) UpdateTags(ctx context.Context, client clients.Client) (clients.Client, error) { - q := `UPDATE clients SET tags = :tags, updated_at = :updated_at, updated_by = :updated_by - WHERE id = :id AND status = :status - RETURNING id, name, tags, identity, metadata, private_metadata, COALESCE(domain_id, '') AS domain_id, COALESCE(parent_group_id, '') AS parent_group_id, status, created_at, updated_at, updated_by` - client.Status = clients.EnabledStatus - return repo.update(ctx, client, q) -} - -func (repo *clientRepo) UpdateIdentity(ctx context.Context, client clients.Client) (clients.Client, error) { - q := `UPDATE clients SET identity = :identity, updated_at = :updated_at, updated_by = :updated_by - WHERE id = :id AND status = :status - RETURNING id, name, tags, identity, metadata, private_metadata, COALESCE(domain_id, '') AS domain_id, status, COALESCE(parent_group_id, '') AS parent_group_id, created_at, updated_at, updated_by` - client.Status = clients.EnabledStatus - return repo.update(ctx, client, q) -} - -func (repo *clientRepo) UpdateSecret(ctx context.Context, client clients.Client) (clients.Client, error) { - q := `UPDATE clients SET secret = :secret, updated_at = :updated_at, updated_by = :updated_by - WHERE id = :id AND status = :status - RETURNING id, name, tags, identity, metadata, private_metadata, COALESCE(domain_id, '') AS domain_id, COALESCE(parent_group_id, '') AS parent_group_id, status, created_at, updated_at, updated_by` - client.Status = clients.EnabledStatus - return repo.update(ctx, client, q) -} - -func (repo *clientRepo) ChangeStatus(ctx context.Context, client clients.Client) (clients.Client, error) { - q := `UPDATE clients SET status = :status, updated_at = :updated_at, updated_by = :updated_by - WHERE id = :id - RETURNING id, name, tags, identity, metadata, private_metadata, COALESCE(domain_id, '') AS domain_id, COALESCE(parent_group_id, '') AS parent_group_id, status, created_at, updated_at, updated_by` - - return repo.update(ctx, client, q) -} - -func (repo *clientRepo) RetrieveByIDWithRoles(ctx context.Context, id, memberID string) (clients.Client, error) { - query := ` - WITH selected_client AS ( - SELECT - c.id, - c.parent_group_id, - COALESCE(g."path", CAST('' AS ltree)) AS parent_group_path, - c.domain_id - FROM - clients c - LEFT JOIN - "groups" g ON c.parent_group_id = g.id - WHERE - c.id = :id - LIMIT 1 - ), - selected_client_roles AS ( - SELECT - cr.entity_id AS client_id, - crm.member_id AS member_id, - cr.id AS role_id, - cr."name" AS role_name, - jsonb_agg(DISTINCT cra."action") AS actions, - 'direct' AS access_type, - CAST('' AS ltree) AS access_provider_path, - '' AS access_provider_id - FROM - clients_roles cr - JOIN - clients_role_members crm ON cr.id = crm.role_id - JOIN - clients_role_actions cra ON cr.id = cra.role_id - JOIN - selected_client sc ON sc.id = cr.entity_id - AND crm.member_id = :member_id - GROUP BY - cr.entity_id, cr.id, cr.name, crm.member_id - ), - selected_group_roles AS ( - SELECT - sc.id AS client_id, - grm.member_id AS member_id, - gr.id AS role_id, - gr."name" AS role_name, - jsonb_agg(DISTINCT all_actions."action") AS actions, - gr.entity_id AS access_provider_id, - g."path" AS access_provider_path, - CASE - WHEN gr.entity_id = sc.parent_group_id - THEN 'direct_group' - ELSE 'indirect_group' - END AS access_type - FROM - "groups" g - JOIN - groups_roles gr ON gr.entity_id = g.id - JOIN - groups_role_members grm ON gr.id = grm.role_id - JOIN - groups_role_actions gra ON gr.id = gra.role_id - JOIN - groups_role_actions all_actions ON gr.id = all_actions.role_id - JOIN - selected_client sc ON TRUE - WHERE - g."path" @> sc.parent_group_path - AND grm.member_id = :member_id - AND ( - (g.id = sc.parent_group_id AND gra."action" LIKE 'client%%') - OR - (g.id <> sc.parent_group_id AND gra."action" LIKE 'subgroup_client%%') - ) - GROUP BY - sc.id, sc.parent_group_id, gr.entity_id, gr.id, gr."name", g."path", grm.member_id - ), - selected_domain_roles AS ( - SELECT - sc.id AS client_id, - drm.member_id AS member_id, - dr.entity_id AS group_id, - dr.id AS role_id, - dr."name" AS role_name, - jsonb_agg(DISTINCT all_actions."action") AS actions, - CAST('' AS ltree) access_provider_path, - 'domain' AS access_type, - dr.entity_id AS access_provider_id - FROM - domains d - JOIN - selected_client sc ON sc.domain_id = d.id - JOIN - domains_roles dr ON dr.entity_id = d.id - JOIN - domains_role_members drm ON dr.id = drm.role_id - JOIN - domains_role_actions dra ON dr.id = dra.role_id - JOIN - domains_role_actions all_actions ON dr.id = all_actions.role_id - WHERE - drm.member_id = :member_id - AND dra."action" LIKE 'client%%' - GROUP BY - sc.id, dr.entity_id, dr.id, dr."name", drm.member_id - ), - all_roles AS ( - SELECT - scr.client_id, - scr.member_id, - scr.role_id AS role_id, - scr.role_name AS role_name, - scr.actions AS actions, - scr.access_type AS access_type, - scr.access_provider_path AS access_provider_path, - scr.access_provider_id AS access_provider_id - FROM - selected_client_roles scr - UNION - SELECT - sgr.client_id, - sgr.member_id, - sgr.role_id AS role_id, - sgr.role_name AS role_name, - sgr.actions AS actions, - sgr.access_type AS access_type, - sgr.access_provider_path AS access_provider_path, - sgr.access_provider_id AS access_provider_id - FROM - selected_group_roles sgr - UNION - SELECT - sdr.client_id, - sdr.member_id, - sdr.role_id AS role_id, - sdr.role_name AS role_name, - sdr.actions AS actions, - sdr.access_type AS access_type, - sdr.access_provider_path AS access_provider_path, - sdr.access_provider_id AS access_provider_id - FROM - selected_domain_roles sdr - ), - final_roles AS ( - SELECT - ar.client_id, - ar.member_id, - jsonb_agg( - jsonb_build_object( - 'role_id', ar.role_id, - 'role_name', ar.role_name, - 'actions', ar.actions, - 'access_type', ar.access_type, - 'access_provider_path', ar.access_provider_path, - 'access_provider_id', ar.access_provider_id - ) - ) AS roles - FROM all_roles ar - GROUP BY - ar.client_id, ar.member_id - ) - SELECT - c2.id, - c2."name", - c2.tags, - COALESCE(c2.domain_id, '') AS domain_id, - COALESCE(c2.parent_group_id, '') AS parent_group_id, - c2."identity", - c2.secret, - c2.metadata, - c2.created_at, - c2.updated_at, - c2.updated_by, - c2.status, - fr.member_id, - fr.roles - FROM clients c2 - JOIN final_roles fr ON fr.client_id = c2.id - ` - parameters := map[string]any{ - "id": id, - "member_id": memberID, - } - row, err := repo.DB.NamedQueryContext(ctx, query, parameters) - if err != nil { - return clients.Client{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - defer row.Close() - - dbc := DBClient{} - if !row.Next() { - return clients.Client{}, repoerr.ErrNotFound - } - - if err := row.StructScan(&dbc); err != nil { - return clients.Client{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - - return ToClient(dbc) -} - -func (repo *clientRepo) RetrieveByID(ctx context.Context, id string) (clients.Client, error) { - q := `SELECT id, name, tags, COALESCE(domain_id, '') AS domain_id, COALESCE(parent_group_id, '') AS parent_group_id, identity, secret, metadata, private_metadata, created_at, updated_at, updated_by, status - FROM clients WHERE id = :id` - - dbc := DBClient{ - ID: id, - } - - row, err := repo.DB.NamedQueryContext(ctx, q, dbc) - if err != nil { - return clients.Client{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - defer row.Close() - - dbc = DBClient{} - if row.Next() { - if err := row.StructScan(&dbc); err != nil { - return clients.Client{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - - return ToClient(dbc) - } - - return clients.Client{}, repoerr.ErrNotFound -} - -func (repo *clientRepo) RetrieveAll(ctx context.Context, pm clients.Page) (clients.ClientsPage, error) { - pageQuery, err := PageQuery(pm) - if err != nil { - return clients.ClientsPage{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - - connJoinQuery := ` - FROM - clients c - ` - - if pm.Channel != "" { - connJoinQuery = ` - ,conn.connection_types - FROM - clients c - LEFT JOIN ( - SELECT - conn.client_id, - conn.channel_id, - array_agg(conn."type") AS connection_types - FROM - connections AS conn - GROUP BY - conn.client_id, conn.channel_id - ) conn ON c.id = conn.client_id - ` - } - - comQuery := fmt.Sprintf(`WITH clients AS ( - SELECT - c.id, - c.name, - c.tags, - c.identity, - c.metadata, - COALESCE(c.domain_id, '') AS domain_id, - COALESCE(parent_group_id, '') AS parent_group_id, - COALESCE(g.path, CAST('' AS ltree)) AS parent_group_path, - c.status, - c.created_at, - c.updated_at, - COALESCE(c.updated_by, '') AS updated_by - FROM - clients c - LEFT JOIN - groups g ON g.id = c.parent_group_id - ) - SELECT - c.* - %s - %s - `, connJoinQuery, pageQuery) - - q := applyOrdering(comQuery, pm) - - q = applyLimitOffset(q) - - dbPage, err := ToDBClientsPage(pm) - if err != nil { - return clients.ClientsPage{}, errors.Wrap(repoerr.ErrFailedToRetrieveAllGroups, err) - } - var items []clients.Client - if !pm.OnlyTotal { - rows, err := repo.DB.NamedQueryContext(ctx, q, dbPage) - if err != nil { - return clients.ClientsPage{}, repo.eh.HandleError(repoerr.ErrFailedToRetrieveAllGroups, err) - } - defer rows.Close() - - for rows.Next() { - dbc := DBClient{} - if err := rows.StructScan(&dbc); err != nil { - return clients.ClientsPage{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - - c, err := ToClient(dbc) - if err != nil { - return clients.ClientsPage{}, err - } - - items = append(items, c) - } - } - cq := fmt.Sprintf(`SELECT COUNT(*) AS total_count - FROM ( - %s - ) AS sub_query; - `, comQuery) - - total, err := postgres.Total(ctx, repo.DB, cq, dbPage) - if err != nil { - return clients.ClientsPage{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - - page := clients.ClientsPage{ - Clients: items, - Page: clients.Page{ - Total: total, - Offset: pm.Offset, - Limit: pm.Limit, - }, - } - - return page, nil -} - -func (repo *clientRepo) RetrieveUserClients(ctx context.Context, domainID, userID string, pm clients.Page) (clients.ClientsPage, error) { - return repo.retrieveClients(ctx, domainID, userID, pm) -} - -func (repo *clientRepo) retrieveClients(ctx context.Context, domainID, userID string, pm clients.Page) (clients.ClientsPage, error) { - pageQuery, err := PageQuery(pm) - if err != nil { - return clients.ClientsPage{}, err - } - - bq := userClientBaseQuery - - connJoinQuery := ` - FROM - final_clients c - ` - - connCountJoinQuery := connJoinQuery - - if pm.Channel != "" { - connCountJoinQuery = ` - FROM - final_clients c - LEFT JOIN ( - SELECT - conn.client_id, - conn.channel_id, - array_agg(conn."type") AS connection_types - FROM - connections AS conn - GROUP BY - conn.client_id, conn.channel_id - ) conn ON c.id = conn.client_id - ` - connJoinQuery = ` - ,conn.connection_types` + connCountJoinQuery - } - - dbPage, err := ToDBClientsPage(pm) - if err != nil { - return clients.ClientsPage{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - dbPage.UserID = userID - dbPage.DomainID = domainID - - if pm.OnlyTotal { - cq := fmt.Sprintf(`%s - SELECT COUNT(*) AS total_count - %s - %s; - `, bq, connCountJoinQuery, pageQuery) - - total, err := postgres.Total(ctx, repo.DB, cq, dbPage) - if err != nil { - return clients.ClientsPage{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - - return clients.ClientsPage{ - Page: clients.Page{ - Total: total, - Offset: pm.Offset, - Limit: pm.Limit, - }, - }, nil - } - - q := fmt.Sprintf(` - %s - SELECT - c.id, - c.name, - c.domain_id, - c.parent_group_id, - c.identity, - c.secret, - c.tags, - c.metadata, - c.created_at, - c.updated_at, - c.updated_by, - c.status, - c.parent_group_path, - c.role_id, - c.role_name, - c.actions, - c.access_type, - c.access_provider_id, - c.access_provider_role_id, - c.access_provider_role_name, - c.access_provider_role_actions, - COUNT(*) OVER() AS total_count - %s - %s - `, bq, connJoinQuery, pageQuery) - - q = applyOrdering(q, pm) - - q = applyLimitOffset(q) - - rows, err := repo.DB.NamedQueryContext(ctx, q, dbPage) - if err != nil { - return clients.ClientsPage{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - defer rows.Close() - - var total uint64 - var items []clients.Client - for rows.Next() { - dbc := DBClient{} - if err := rows.StructScan(&dbc); err != nil { - return clients.ClientsPage{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - - total = dbc.TotalCount - - c, err := ToClient(dbc) - if err != nil { - return clients.ClientsPage{}, err - } - - items = append(items, c) - } - - if len(items) == 0 { - cq := fmt.Sprintf(`%s - SELECT COUNT(*) AS total_count - %s - %s; - `, bq, connCountJoinQuery, pageQuery) - - total, err = postgres.Total(ctx, repo.DB, cq, dbPage) - if err != nil { - return clients.ClientsPage{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - } - - return clients.ClientsPage{ - Clients: items, - Page: clients.Page{ - Total: total, - Offset: pm.Offset, - Limit: pm.Limit, - }, - }, nil -} - -const userClientBaseQuery = ` - WITH direct_clients AS ( - SELECT - c.id, - c.name, - c.domain_id, - c.parent_group_id, - c.tags, - c.metadata, - c.identity, - c.secret, - c.created_at, - c.updated_at, - c.updated_by, - c.status, - COALESCE(pg.path, CAST('' AS ltree)) AS parent_group_path, - cr.id AS role_id, - cr."name" AS role_name, - array_agg(cra."action") AS actions, - 'direct' as access_type, - '' AS access_provider_id, - '' AS access_provider_role_id, - '' AS access_provider_role_name, - CAST(array[] AS text[]) AS access_provider_role_actions - FROM - clients_role_members crm - JOIN - clients_role_actions cra ON cra.role_id = crm.role_id - JOIN - clients_roles cr ON cr.id = crm.role_id - JOIN - clients c ON c.id = cr.entity_id - LEFT JOIN - groups pg ON pg.id = c.parent_group_id - WHERE - crm.member_id = :user_id - AND c.domain_id = :domain_id_param - GROUP BY - cr.entity_id, crm.member_id, cr.id, cr."name", c.id, pg.path - ), - direct_groups AS ( - SELECT - g.*, - gr.entity_id AS entity_id, - grm.member_id AS member_id, - gr.id AS role_id, - gr."name" AS role_name, - array_agg(DISTINCT all_actions."action") AS actions - FROM - groups_role_members grm - JOIN - groups_role_actions gra ON gra.role_id = grm.role_id - JOIN - groups_roles gr ON gr.id = grm.role_id - JOIN - "groups" g ON g.id = gr.entity_id - JOIN - groups_role_actions all_actions ON all_actions.role_id = grm.role_id - WHERE - grm.member_id = :user_id - AND g.domain_id = :domain_id_param - AND gra."action" LIKE 'client%' - GROUP BY - gr.entity_id, grm.member_id, gr.id, gr."name", g."path", g.id - ), - direct_groups_with_subgroup AS ( - SELECT - g.*, - gr.entity_id AS entity_id, - grm.member_id AS member_id, - gr.id AS role_id, - gr."name" AS role_name, - array_agg(DISTINCT all_actions."action") AS actions - FROM - groups_role_members grm - JOIN - groups_role_actions gra ON gra.role_id = grm.role_id - JOIN - groups_roles gr ON gr.id = grm.role_id - JOIN - "groups" g ON g.id = gr.entity_id - JOIN - groups_role_actions all_actions ON all_actions.role_id = grm.role_id - WHERE - grm.member_id = :user_id - AND g.domain_id = :domain_id_param - AND gra."action" LIKE 'subgroup_client%' - GROUP BY - gr.entity_id, grm.member_id, gr.id, gr."name", g."path", g.id - ), - direct_leaf_groups_with_subgroup AS ( - SELECT dgws.* - FROM direct_groups_with_subgroup dgws - WHERE NOT EXISTS ( - SELECT 1 - FROM direct_groups_with_subgroup dgws2 - WHERE - dgws2.path @> dgws.path - AND dgws2.id != dgws.id - ) - ), - indirect_child_groups AS ( - SELECT - DISTINCT indirect_child_groups.id as child_id, - indirect_child_groups.*, - dlgws.id as access_provider_id, - dlgws.role_id as access_provider_role_id, - dlgws.role_name as access_provider_role_name, - dlgws.actions as access_provider_role_actions - FROM - direct_leaf_groups_with_subgroup dlgws - JOIN - groups indirect_child_groups ON indirect_child_groups.path <@ dlgws.path - WHERE - indirect_child_groups.domain_id = :domain_id_param - AND NOT EXISTS ( - SELECT 1 - FROM direct_groups_with_subgroup dgws - WHERE dgws.id = indirect_child_groups.id - ) - ), - final_groups AS ( - SELECT - id, - parent_id, - domain_id, - "name", - description, - metadata, - created_at, - updated_at, - updated_by, - status, - "path", - '' AS role_id, - '' AS role_name, - CAST(array[] AS text[]) AS actions, - 'direct_group' AS access_type, - id AS access_provider_id, - role_id AS access_provider_role_id, - role_name AS access_provider_role_name, - actions AS access_provider_role_actions - FROM - direct_groups - UNION - SELECT - id, - parent_id, - domain_id, - "name", - description, - metadata, - created_at, - updated_at, - updated_by, - status, - "path", - '' AS role_id, - '' AS role_name, - CAST(array[] AS text[]) AS actions, - 'indirect_group' AS access_type, - access_provider_id, - access_provider_role_id, - access_provider_role_name, - access_provider_role_actions - FROM - indirect_child_groups - ), - groups_clients AS ( - SELECT - c.id, - c.name, - c.domain_id, - c.parent_group_id, - c.tags, - c.metadata, - c.identity, - c.secret, - c.created_at, - c.updated_at, - c.updated_by, - c.status, - g.path AS parent_group_path, - g.role_id, - g.role_name, - g.actions, - g.access_type, - g.access_provider_id, - g.access_provider_role_id, - g.access_provider_role_name, - g.access_provider_role_actions - FROM - final_groups g - JOIN - clients c ON c.parent_group_id = g.id - WHERE - NOT EXISTS (SELECT 1 FROM direct_clients dc WHERE dc.id = c.id) - UNION - SELECT * FROM direct_clients - ), - final_clients AS ( - SELECT - gc.id, - gc."name", - gc.domain_id, - gc.parent_group_id, - gc.tags, - gc.metadata, - gc.identity, - gc.secret, - gc.created_at, - gc.updated_at, - gc.updated_by, - gc.status, - gc.parent_group_path, - gc.role_id, - gc.role_name, - gc.actions, - gc.access_type, - gc.access_provider_id, - gc.access_provider_role_id, - gc.access_provider_role_name, - gc.access_provider_role_actions - FROM - groups_clients AS gc - UNION - SELECT - dc.id, - dc."name", - dc.domain_id, - dc.parent_group_id, - dc.tags, - dc.metadata, - dc.identity, - dc.secret, - dc.created_at, - dc.updated_at, - dc.updated_by, - dc.status, - g."path" AS parent_group_path, - '' AS role_id, - '' AS role_name, - CAST(array[] AS text[]) AS actions, - 'domain' AS access_type, - d.id AS access_provider_id, - dr.id AS access_provider_role_id, - dr."name" AS access_provider_role_name, - array_agg(dra."action") as access_provider_role_actions - FROM - domains_role_members drm - JOIN - domains_role_actions dra ON dra.role_id = drm.role_id - JOIN - domains_roles dr ON dr.id = drm.role_id - JOIN - domains d ON d.id = dr.entity_id - JOIN - clients dc ON dc.domain_id = d.id - LEFT JOIN - groups g ON dc.parent_group_id = g.id - WHERE - drm.member_id = :user_id - AND d.id = :domain_id_param - AND dra."action" LIKE 'client_%' - AND NOT EXISTS ( - SELECT 1 FROM groups_clients gc - WHERE gc.id = dc.id - ) - GROUP BY - dc.id, d.id, dr.id, g."path" - ) - ` - -func (repo *clientRepo) SearchClients(ctx context.Context, pm clients.Page) (clients.ClientsPage, error) { - query, err := PageQuery(pm) - if err != nil { - return clients.ClientsPage{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - - tq := query - query = applyOrdering(query, pm) - - q := fmt.Sprintf(`SELECT c.id, c.name, c.metadata, c.created_at, c.updated_at FROM clients c %s LIMIT :limit OFFSET :offset;`, query) - - dbPage, err := ToDBClientsPage(pm) - if err != nil { - return clients.ClientsPage{}, errors.Wrap(repoerr.ErrFailedToRetrieveAllGroups, err) - } - - rows, err := repo.DB.NamedQueryContext(ctx, q, dbPage) - if err != nil { - return clients.ClientsPage{}, repo.eh.HandleError(repoerr.ErrFailedToRetrieveAllGroups, err) - } - defer rows.Close() - - var items []clients.Client - for rows.Next() { - dbc := DBClient{} - if err := rows.StructScan(&dbc); err != nil { - return clients.ClientsPage{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - - c, err := ToClient(dbc) - if err != nil { - return clients.ClientsPage{}, err - } - - items = append(items, c) - } - - cq := fmt.Sprintf(`SELECT COUNT(*) FROM clients c %s;`, tq) - total, err := postgres.Total(ctx, repo.DB, cq, dbPage) - if err != nil { - return clients.ClientsPage{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - - page := clients.ClientsPage{ - Clients: items, - Page: clients.Page{ - Total: total, - Offset: pm.Offset, - Limit: pm.Limit, - }, - } - - return page, nil -} - -func (repo *clientRepo) update(ctx context.Context, client clients.Client, query string) (clients.Client, error) { - dbc, err := ToDBClient(client) - if err != nil { - return clients.Client{}, errors.Wrap(repoerr.ErrUpdateEntity, err) - } - - row, err := repo.DB.NamedQueryContext(ctx, query, dbc) - if err != nil { - return clients.Client{}, repo.eh.HandleError(repoerr.ErrUpdateEntity, err) - } - defer row.Close() - - dbc = DBClient{} - if row.Next() { - if err := row.StructScan(&dbc); err != nil { - return clients.Client{}, repo.eh.HandleError(repoerr.ErrUpdateEntity, err) - } - - return ToClient(dbc) - } - - return clients.Client{}, repoerr.ErrNotFound -} - -func (repo *clientRepo) Delete(ctx context.Context, clientIDs ...string) error { - q := "DELETE FROM clients AS c WHERE c.id = ANY(:client_ids) ;" - - params := map[string]any{ - "client_ids": clientIDs, - } - result, err := repo.DB.NamedExecContext(ctx, q, params) - if err != nil { - return repo.eh.HandleError(repoerr.ErrRemoveEntity, err) - } - if rows, _ := result.RowsAffected(); rows == 0 { - return repoerr.ErrNotFound - } - - return nil -} - -type DBClient struct { - ID string `db:"id"` - Name string `db:"name,omitempty"` - Tags pgtype.TextArray `db:"tags,omitempty"` - Identity string `db:"identity"` - Domain string `db:"domain_id"` - ParentGroup sql.NullString `db:"parent_group_id,omitempty"` - Secret string `db:"secret"` - Metadata []byte `db:"metadata,omitempty"` - PrivateMetadata []byte `db:"private_metadata,omitempty"` - CreatedAt time.Time `db:"created_at,omitempty"` - UpdatedAt sql.NullTime `db:"updated_at,omitempty"` - UpdatedBy *string `db:"updated_by,omitempty"` - Status clients.Status `db:"status,omitempty"` - ParentGroupPath sql.NullString `db:"parent_group_path,omitempty"` - RoleID string `db:"role_id,omitempty"` - RoleName string `db:"role_name,omitempty"` - Actions pq.StringArray `db:"actions,omitempty"` - AccessType string `db:"access_type,omitempty"` - AccessProviderId string `db:"access_provider_id,omitempty"` - AccessProviderRoleId string `db:"access_provider_role_id,omitempty"` - AccessProviderRoleName string `db:"access_provider_role_name,omitempty"` - AccessProviderRoleActions pq.StringArray `db:"access_provider_role_actions,omitempty"` - ConnectionTypes pq.Int32Array `db:"connection_types,omitempty"` - MemberID string `db:"member_id,omitempty"` - Roles json.RawMessage `db:"roles,omitempty"` - TotalCount uint64 `db:"total_count"` -} - -func ToDBClient(c clients.Client) (DBClient, error) { - privateMetadata := []byte("{}") - if len(c.PrivateMetadata) > 0 { - b, err := json.Marshal(c.PrivateMetadata) - if err != nil { - return DBClient{}, errors.Wrap(repoerr.ErrMalformedEntity, err) - } - privateMetadata = b - } - metadata := []byte("{}") - if len(c.Metadata) > 0 { - b, err := json.Marshal(c.Metadata) - if err != nil { - return DBClient{}, errors.Wrap(repoerr.ErrMalformedEntity, err) - } - metadata = b - } - var tags pgtype.TextArray - if err := tags.Set(c.Tags); err != nil { - return DBClient{}, err - } - var updatedBy *string - if c.UpdatedBy != "" { - updatedBy = &c.UpdatedBy - } - var updatedAt sql.NullTime - if c.UpdatedAt != (time.Time{}) { - updatedAt = sql.NullTime{Time: c.UpdatedAt, Valid: true} - } - - return DBClient{ - ID: c.ID, - Name: c.Name, - Tags: tags, - Domain: c.Domain, - ParentGroup: toNullString(c.ParentGroup), - Identity: c.Credentials.Identity, - Secret: c.Credentials.Secret, - Metadata: metadata, - PrivateMetadata: privateMetadata, - CreatedAt: c.CreatedAt, - UpdatedAt: updatedAt, - UpdatedBy: updatedBy, - Status: c.Status, - }, nil -} - -func ToClient(t DBClient) (clients.Client, error) { - var privateMetadata, metadata clients.Metadata - if t.PrivateMetadata != nil { - if err := json.Unmarshal([]byte(t.PrivateMetadata), &privateMetadata); err != nil { - return clients.Client{}, errors.Wrap(repoerr.ErrMalformedEntity, err) - } - } - if t.Metadata != nil { - if err := json.Unmarshal([]byte(t.Metadata), &metadata); err != nil { - return clients.Client{}, errors.Wrap(repoerr.ErrMalformedEntity, err) - } - } - - var tags []string - for _, e := range t.Tags.Elements { - tags = append(tags, e.String) - } - - var updatedBy string - if t.UpdatedBy != nil { - updatedBy = *t.UpdatedBy - } - - var updatedAt time.Time - if t.UpdatedAt.Valid { - updatedAt = t.UpdatedAt.Time.UTC() - } - - var connTypes []connections.ConnType - for _, ct := range t.ConnectionTypes { - connType, err := connections.NewType(uint(ct)) - if err != nil { - return clients.Client{}, err - } - connTypes = append(connTypes, connType) - } - - var roles []roles.MemberRoleActions - if t.Roles != nil { - if err := json.Unmarshal(t.Roles, &roles); err != nil { - return clients.Client{}, errors.Wrap(errors.ErrMalformedEntity, err) - } - } - - cli := clients.Client{ - ID: t.ID, - Name: t.Name, - Tags: tags, - Domain: t.Domain, - ParentGroup: toString(t.ParentGroup), - Credentials: clients.Credentials{ - Identity: t.Identity, - Secret: t.Secret, - }, - Metadata: metadata, - PrivateMetadata: privateMetadata, - CreatedAt: t.CreatedAt.UTC(), - UpdatedAt: updatedAt, - UpdatedBy: updatedBy, - Status: t.Status, - ParentGroupPath: toString(t.ParentGroupPath), - RoleID: t.RoleID, - RoleName: t.RoleName, - Actions: t.Actions, - AccessType: t.AccessType, - AccessProviderId: t.AccessProviderId, - AccessProviderRoleId: t.AccessProviderRoleId, - AccessProviderRoleName: t.AccessProviderRoleName, - AccessProviderRoleActions: t.AccessProviderRoleActions, - ConnectionTypes: connTypes, - Roles: roles, - } - return cli, nil -} - -func ToDBClientsPage(pm clients.Page) (dbClientsPage, error) { - _, data, err := postgres.CreateMetadataQuery("", pm.Metadata) - if err != nil { - return dbClientsPage{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - var tags pgtype.TextArray - if err := tags.Set(pm.Tags.Elements); err != nil { - return dbClientsPage{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - return dbClientsPage{ - Offset: pm.Offset, - Limit: pm.Limit, - Name: pm.Name, - Identity: pm.Identity, - Id: pm.ID, - Metadata: data, - Domain: pm.Domain, - Status: pm.Status, - Tags: tags, - GroupID: pm.Group, - ChannelID: pm.Channel, - RoleName: pm.RoleName, - ConnType: pm.ConnectionType, - RoleID: pm.RoleID, - Actions: pm.Actions, - AccessType: pm.AccessType, - IDs: pq.StringArray(pm.IDs), - CreatedFrom: pm.CreatedFrom, - CreatedTo: pm.CreatedTo, - }, nil -} - -type dbClientsPage struct { - Limit uint64 `db:"limit"` - Offset uint64 `db:"offset"` - Name string `db:"name"` - Id string `db:"id"` - Domain string `db:"domain_id"` - Identity string `db:"identity"` - Metadata []byte `db:"metadata"` - Tags pgtype.TextArray `db:"tags"` - Status clients.Status `db:"status"` - GroupID *string `db:"group_id"` - ChannelID string `db:"channel_id"` - ConnType string `db:"type"` - RoleName string `db:"role_name"` - RoleID string `db:"role_id"` - Actions pq.StringArray `db:"actions"` - AccessType string `db:"access_type"` - CreatedFrom time.Time `db:"created_from"` - CreatedTo time.Time `db:"created_to"` - IDs pq.StringArray `db:"ids"` - UserID string `db:"user_id"` - DomainID string `db:"domain_id_param"` -} - -func PageQuery(pm clients.Page) (string, error) { - var query []string - if pm.Name != "" { - query = append(query, "c.name ILIKE '%' || :name || '%'") - } - if pm.Identity != "" { - query = append(query, "c.identity ILIKE '%' || :identity || '%'") - } - if pm.ID != "" { - query = append(query, "c.id = :id") - } - if len(pm.Tags.Elements) > 0 { - switch pm.Tags.Operator { - case clients.AndOp: - query = append(query, "tags @> :tags") - default: // OR - query = append(query, "tags && :tags") - } - } - if len(pm.IDs) != 0 { - query = append(query, "c.id = ANY(:ids)") - } - - if pm.Status != clients.AllStatus { - query = append(query, "c.status = :status") - } - if pm.Domain != "" { - query = append(query, "c.domain_id = :domain_id") - } - - if pm.Group != nil { - switch *pm.Group { - case "": - query = append(query, "c.parent_group_id = '' ") - default: - query = append(query, "c.parent_group_path <@ (SELECT path from groups where id = :group_id) ") - } - } - - if pm.Channel != "" { - query = append(query, "conn.channel_id = :channel_id ") - if pm.ConnectionType != "" { - query = append(query, "conn.type = :conn_type ") - } - } - if pm.AccessType != "" { - query = append(query, "c.access_type = :access_type") - } - if pm.RoleID != "" { - query = append(query, "c.role_id = :role_id") - } - if pm.RoleName != "" { - query = append(query, "c.role_name = :role_name") - } - if len(pm.Actions) != 0 { - query = append(query, "c.actions @> :actions") - } - if len(pm.Metadata) > 0 { - query = append(query, "c.metadata @> :metadata") - } - - if !pm.CreatedFrom.IsZero() { - query = append(query, "c.created_at >= :created_from") - } - if !pm.CreatedTo.IsZero() { - query = append(query, "c.created_at <= :created_to") - } - - var emq string - if len(query) > 0 { - emq = fmt.Sprintf("WHERE %s", strings.Join(query, " AND ")) - } - return emq, nil -} - -func applyOrdering(emq string, pm clients.Page) string { - var orderBy string - switch pm.Order { - case "name": - orderBy = "name" - case "identity": - orderBy = "identity" - case "created_at": - orderBy = "created_at" - case "updated_at": - orderBy = "COALESCE(updated_at, created_at)" - default: - return emq - } - - if pm.Dir == api.AscDir || pm.Dir == api.DescDir { - return fmt.Sprintf("%s ORDER BY %s %s, id %s", emq, orderBy, pm.Dir, pm.Dir) - } - return fmt.Sprintf("%s ORDER BY %s", emq, orderBy) -} - -func applyLimitOffset(query string) string { - return fmt.Sprintf(`%s - LIMIT :limit OFFSET :offset`, query) -} - -func toNullString(s string) sql.NullString { - if s == "" { - return sql.NullString{} - } - - return sql.NullString{ - String: s, - Valid: true, - } -} - -func toString(s sql.NullString) string { - if s.Valid { - return s.String - } - return "" -} - -func (repo *clientRepo) RetrieveByIds(ctx context.Context, ids []string) (clients.ClientsPage, error) { - if len(ids) == 0 { - return clients.ClientsPage{}, nil - } - - pm := clients.Page{IDs: ids} - query, err := PageQuery(pm) - if err != nil { - return clients.ClientsPage{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - - q := fmt.Sprintf(`SELECT c.id, c.name, c.tags, c.identity, c.metadata, COALESCE(c.domain_id, '') AS domain_id, COALESCE(parent_group_id, '') AS parent_group_id, c.status, - c.created_at, c.updated_at, COALESCE(c.updated_by, '') AS updated_by FROM clients c %s ORDER BY c.created_at`, query) - - dbPage, err := ToDBClientsPage(pm) - if err != nil { - return clients.ClientsPage{}, errors.Wrap(repoerr.ErrFailedToRetrieveAllGroups, err) - } - rows, err := repo.DB.NamedQueryContext(ctx, q, dbPage) - if err != nil { - return clients.ClientsPage{}, errors.Wrap(repoerr.ErrFailedToRetrieveAllGroups, err) - } - defer rows.Close() - - var items []clients.Client - for rows.Next() { - dbc := DBClient{} - if err := rows.StructScan(&dbc); err != nil { - return clients.ClientsPage{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - - c, err := ToClient(dbc) - if err != nil { - return clients.ClientsPage{}, err - } - - items = append(items, c) - } - cq := fmt.Sprintf(`SELECT COUNT(*) FROM clients c %s;`, query) - - total, err := postgres.Total(ctx, repo.DB, cq, dbPage) - if err != nil { - return clients.ClientsPage{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - - page := clients.ClientsPage{ - Clients: items, - Page: clients.Page{ - Total: total, - Offset: pm.Offset, - Limit: total, - }, - } - - return page, nil -} - -func (repo *clientRepo) AddConnections(ctx context.Context, conns []clients.Connection) error { - dbConns := toDBConnections(conns) - q := `INSERT INTO connections (channel_id, domain_id, client_id, type) - VALUES (:channel_id, :domain_id, :client_id, :type);` - if _, err := repo.DB.NamedExecContext(ctx, q, dbConns); err != nil { - return postgres.HandleError(repoerr.ErrCreateEntity, err) - } - - return nil -} - -func (repo *clientRepo) RemoveConnections(ctx context.Context, conns []clients.Connection) (retErr error) { - tx, err := repo.DB.BeginTxx(ctx, nil) - if err != nil { - return errors.Wrap(repoerr.ErrRemoveEntity, err) - } - defer func() { - if retErr != nil { - if errRollBack := tx.Rollback(); errRollBack != nil { - retErr = errors.Wrap(retErr, errors.Wrap(apiutil.ErrRollbackTx, errRollBack)) - } - } - }() - - for _, conn := range conns { - query := `DELETE FROM connections WHERE channel_id = :channel_id AND domain_id = :domain_id AND client_id = :client_id` - if uint8(conn.Type) > 0 { - query = query + " AND type = :type " - } - dbConn := toDBConnection(conn) - if _, err := tx.NamedExec(query, dbConn); err != nil { - return errors.Wrap(repoerr.ErrRemoveEntity, errors.Wrap(fmt.Errorf("failed to delete connection for channel_id: %s, domain_id: %s client_id %s", conn.ChannelID, conn.DomainID, conn.ClientID), err)) - } - } - if err := tx.Commit(); err != nil { - return errors.Wrap(repoerr.ErrRemoveEntity, err) - } - return nil -} - -func (repo *clientRepo) SetParentGroup(ctx context.Context, cli clients.Client) error { - q := "UPDATE clients SET parent_group_id = :parent_group_id, updated_at = :updated_at, updated_by = :updated_by WHERE id = :id" - - dbcli, err := ToDBClient(cli) - if err != nil { - return errors.Wrap(repoerr.ErrUpdateEntity, err) - } - result, err := repo.DB.NamedExecContext(ctx, q, dbcli) - if err != nil { - return postgres.HandleError(repoerr.ErrUpdateEntity, err) - } - if rows, _ := result.RowsAffected(); rows == 0 { - return repoerr.ErrNotFound - } - return nil -} - -func (repo *clientRepo) RemoveParentGroup(ctx context.Context, cli clients.Client) error { - q := "UPDATE clients SET parent_group_id = NULL, updated_at = :updated_at, updated_by = :updated_by WHERE id = :id" - dbcli, err := ToDBClient(cli) - if err != nil { - return errors.Wrap(repoerr.ErrUpdateEntity, err) - } - result, err := repo.DB.NamedExecContext(ctx, q, dbcli) - if err != nil { - return postgres.HandleError(repoerr.ErrRemoveEntity, err) - } - if rows, _ := result.RowsAffected(); rows == 0 { - return repoerr.ErrNotFound - } - return nil -} - -func (repo *clientRepo) ClientConnectionsCount(ctx context.Context, id string) (uint64, error) { - query := `SELECT COUNT(*) FROM connections WHERE client_id = :client_id` - dbConn := dbConnection{ClientID: id} - - total, err := postgres.Total(ctx, repo.DB, query, dbConn) - if err != nil { - return 0, postgres.HandleError(repoerr.ErrViewEntity, err) - } - return total, nil -} - -func (repo *clientRepo) DoesClientHaveConnections(ctx context.Context, id string) (bool, error) { - query := `SELECT 1 FROM connections WHERE client_id = :client_id` - dbConn := dbConnection{ClientID: id} - - rows, err := repo.DB.NamedQueryContext(ctx, query, dbConn) - if err != nil { - return false, postgres.HandleError(repoerr.ErrViewEntity, err) - } - defer rows.Close() - - return rows.Next(), nil -} - -func (repo *clientRepo) RemoveChannelConnections(ctx context.Context, channelID string) error { - query := `DELETE FROM connections WHERE channel_id = :channel_id` - - dbConn := dbConnection{ChannelID: channelID} - if _, err := repo.DB.NamedExecContext(ctx, query, dbConn); err != nil { - return errors.Wrap(repoerr.ErrRemoveEntity, err) - } - return nil -} - -func (repo *clientRepo) RemoveClientConnections(ctx context.Context, clientID string) error { - query := `DELETE FROM connections WHERE client_id = :client_id` - - dbConn := dbConnection{ClientID: clientID} - if _, err := repo.DB.NamedExecContext(ctx, query, dbConn); err != nil { - return errors.Wrap(repoerr.ErrRemoveEntity, err) - } - return nil -} - -func (repo *clientRepo) RetrieveParentGroupClients(ctx context.Context, parentGroupID string) ([]clients.Client, error) { - query := `SELECT c.id, c.name, c.tags, c.metadata, COALESCE(c.domain_id, '') AS domain_id, COALESCE(parent_group_id, '') AS parent_group_id, c.status, - c.created_at, c.updated_at, COALESCE(c.updated_by, '') AS updated_by FROM clients c WHERE c.parent_group_id = :parent_group_id ;` - - rows, err := repo.DB.NamedQueryContext(ctx, query, DBClient{ParentGroup: toNullString(parentGroupID)}) - if err != nil { - return []clients.Client{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - defer rows.Close() - - var clis []clients.Client - for rows.Next() { - dbCli := DBClient{} - if err := rows.StructScan(&dbCli); err != nil { - return []clients.Client{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - - cli, err := ToClient(dbCli) - if err != nil { - return []clients.Client{}, err - } - - clis = append(clis, cli) - } - return clis, nil -} - -func (repo *clientRepo) UnsetParentGroupFromClient(ctx context.Context, parentGroupID string) error { - query := "UPDATE clients SET parent_group_id = NULL WHERE parent_group_id = :parent_group_id" - - if _, err := repo.DB.NamedExecContext(ctx, query, DBClient{ParentGroup: toNullString(parentGroupID)}); err != nil { - return errors.Wrap(repoerr.ErrRemoveEntity, err) - } - return nil -} - -type dbConnection struct { - ClientID string `db:"client_id"` - ChannelID string `db:"channel_id"` - DomainID string `db:"domain_id"` - Type connections.ConnType `db:"type"` -} - -func toDBConnections(conns []clients.Connection) []dbConnection { - var dbconns []dbConnection - for _, conn := range conns { - dbconns = append(dbconns, toDBConnection(conn)) - } - return dbconns -} - -func toDBConnection(conn clients.Connection) dbConnection { - return dbConnection{ - ClientID: conn.ClientID, - ChannelID: conn.ChannelID, - DomainID: conn.DomainID, - Type: conn.Type, - } -} diff --git a/clients/postgres/clients_test.go b/clients/postgres/clients_test.go deleted file mode 100644 index 12c864f93..000000000 --- a/clients/postgres/clients_test.go +++ /dev/null @@ -1,4172 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package postgres_test - -import ( - "context" - "fmt" - "strconv" - "strings" - "testing" - "time" - - "github.com/0x6flab/namegenerator" - "github.com/absmach/magistrala/clients" - "github.com/absmach/magistrala/clients/postgres" - "github.com/absmach/magistrala/domains" - dpostgres "github.com/absmach/magistrala/domains/postgres" - "github.com/absmach/magistrala/groups" - gpostgres "github.com/absmach/magistrala/groups/postgres" - "github.com/absmach/magistrala/internal/testsutil" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/connections" - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - "github.com/absmach/magistrala/pkg/roles" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" -) - -const ( - maxNameSize = 1024 - emailSuffix = "@example.com" - defOrder = "created_at" - ascDir = "asc" - descDir = "desc" -) - -var ( - invalidName = strings.Repeat("m", maxNameSize+10) - clientIdentity = "client-identity@example.com" - clientName = "client name" - invalidDomainID = strings.Repeat("m", maxNameSize+10) - namegen = namegenerator.NewGenerator() - validTimestamp = time.Now().UTC().Truncate(time.Millisecond) - validClient = clients.Client{ - ID: testsutil.GenerateUUID(&testing.T{}), - Domain: testsutil.GenerateUUID(&testing.T{}), - Name: namegen.Generate(), - Metadata: map[string]any{"key": "value"}, - PrivateMetadata: map[string]any{"key": "value"}, - CreatedAt: time.Now().UTC().Truncate(time.Microsecond), - Status: clients.EnabledStatus, - } - invalidID = strings.Repeat("a", 37) - directAccess = "direct" - directGroupAccess = "direct_group" - domainAccess = "domain" - availableActions = []string{ - "delete", - "membership", - "read", - "update", - } - domainAvailableActions = []string{ - "client_add_role_users", - "client_connect_to_channel", - "client_create", - "client_delete", - "client_manage_role", - "client_read", - "client_remove_role_users", - "client_set_parent_group", - "client_update", - "client_view_role_users", - } - groupAvailableActions = []string{ - "client_add_role_users", - "client_connect_to_channel", - "client_create", - "client_delete", - "client_manage_role", - "client_read", - "client_remove_role_users", - "client_set_parent_group", - "client_update", - "client_view_role_users", - "subgroup_client_add_role_users", - "subgroup_client_connect_to_channel", - "subgroup_client_create", - "subgroup_client_delete", - "subgroup_client_manage_role", - "subgroup_client_read", - "subgroup_client_remove_role_users", - "subgroup_client_set_parent_group", - "subgroup_client_update", - "subgroup_client_view_role_users", - "subgroup_manage_role", - "subgroup_membership", - "subgroup_read", - "subgroup_remove_role_users", - "subgroup_set_child", - "subgroup_set_parent", - "subgroup_update", - } - errClientSecretNotAvailable = errors.New("client key is not available") -) - -func TestClientsSave(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM clients") - require.Nil(t, err, fmt.Sprintf("clean clients unexpected error: %s", err)) - }) - repo := postgres.NewRepository(database) - - uid := testsutil.GenerateUUID(t) - domainID := testsutil.GenerateUUID(t) - secret := testsutil.GenerateUUID(t) - - duplicateClientID := testsutil.GenerateUUID(t) - - cases := []struct { - desc string - clients []clients.Client - err error - }{ - { - desc: "add new client successfully", - clients: []clients.Client{ - { - ID: uid, - Domain: domainID, - Name: clientName, - Credentials: clients.Credentials{ - Identity: clientIdentity, - Secret: secret, - }, - PrivateMetadata: map[string]any{"key": "value"}, - Metadata: map[string]any{"key": "value"}, - Status: clients.EnabledStatus, - }, - }, - err: nil, - }, - { - desc: "add multiple clients successfully", - clients: []clients.Client{ - { - ID: testsutil.GenerateUUID(t), - Domain: testsutil.GenerateUUID(t), - Name: namegen.Generate(), - Credentials: clients.Credentials{ - Secret: testsutil.GenerateUUID(t), - }, - PrivateMetadata: map[string]any{"key": "value"}, - Metadata: map[string]any{"key": "value"}, - Status: clients.EnabledStatus, - }, - { - ID: testsutil.GenerateUUID(t), - Domain: testsutil.GenerateUUID(t), - Name: namegen.Generate(), - Credentials: clients.Credentials{ - Secret: testsutil.GenerateUUID(t), - }, - PrivateMetadata: map[string]any{"key": "value"}, - Metadata: map[string]any{"key": "value"}, - Status: clients.EnabledStatus, - }, - { - ID: testsutil.GenerateUUID(t), - Domain: testsutil.GenerateUUID(t), - Name: namegen.Generate(), - Credentials: clients.Credentials{ - Secret: testsutil.GenerateUUID(t), - }, - PrivateMetadata: map[string]any{"key": "value"}, - Metadata: map[string]any{"key": "value"}, - Status: clients.EnabledStatus, - }, - }, - err: nil, - }, - { - desc: "add new client with duplicate secret", - clients: []clients.Client{ - { - ID: testsutil.GenerateUUID(t), - Domain: domainID, - Name: namegen.Generate(), - Credentials: clients.Credentials{ - Identity: clientIdentity, - Secret: secret, - }, - PrivateMetadata: map[string]any{"key": "value"}, - Metadata: map[string]any{"key": "value"}, - Status: clients.EnabledStatus, - }, - }, - err: errClientSecretNotAvailable, - }, - { - desc: "add multiple clients with one client having duplicate secret", - clients: []clients.Client{ - { - ID: testsutil.GenerateUUID(t), - Domain: testsutil.GenerateUUID(t), - Name: namegen.Generate(), - Credentials: clients.Credentials{ - Secret: testsutil.GenerateUUID(t), - }, - PrivateMetadata: map[string]any{"key": "value"}, - Metadata: map[string]any{"key": "value"}, - Status: clients.EnabledStatus, - }, - { - ID: testsutil.GenerateUUID(t), - Domain: domainID, - Name: namegen.Generate(), - Credentials: clients.Credentials{ - Identity: clientIdentity, - Secret: secret, - }, - PrivateMetadata: map[string]any{"key": "value"}, - Metadata: map[string]any{"key": "value"}, - Status: clients.EnabledStatus, - }, - }, - err: errClientSecretNotAvailable, - }, - { - desc: "add new client without domain id", - clients: []clients.Client{ - { - ID: testsutil.GenerateUUID(t), - Name: clientName, - Credentials: clients.Credentials{ - Identity: "withoutdomain-client@example.com", - Secret: testsutil.GenerateUUID(t), - }, - PrivateMetadata: map[string]any{"key": "value"}, - Metadata: map[string]any{"key": "value"}, - Status: clients.EnabledStatus, - }, - }, - err: nil, - }, - { - desc: "add client with invalid client id", - clients: []clients.Client{ - { - ID: invalidName, - Domain: domainID, - Name: clientName, - Credentials: clients.Credentials{ - Identity: "invalidid-client@example.com", - Secret: testsutil.GenerateUUID(t), - }, - PrivateMetadata: map[string]any{"key": "value"}, - Metadata: map[string]any{"key": "value"}, - Status: clients.EnabledStatus, - }, - }, - err: repoerr.ErrCreateEntity, - }, - { - desc: "add multiple clients with one client having invalid client id", - clients: []clients.Client{ - { - ID: testsutil.GenerateUUID(t), - Domain: testsutil.GenerateUUID(t), - Name: namegen.Generate(), - Credentials: clients.Credentials{ - Secret: testsutil.GenerateUUID(t), - }, - PrivateMetadata: map[string]any{"key": "value"}, - Metadata: map[string]any{"key": "value"}, - Status: clients.EnabledStatus, - }, - { - ID: invalidName, - Domain: testsutil.GenerateUUID(t), - Name: namegen.Generate(), - Credentials: clients.Credentials{ - Secret: testsutil.GenerateUUID(t), - }, - PrivateMetadata: map[string]any{"key": "value"}, - Metadata: map[string]any{"key": "value"}, - Status: clients.EnabledStatus, - }, - }, - err: repoerr.ErrCreateEntity, - }, - { - desc: "add client with invalid client name", - clients: []clients.Client{ - { - ID: testsutil.GenerateUUID(t), - Name: invalidName, - Domain: domainID, - Credentials: clients.Credentials{ - Identity: "invalidname-client@example.com", - Secret: testsutil.GenerateUUID(t), - }, - PrivateMetadata: map[string]any{"key": "value"}, - Metadata: map[string]any{"key": "value"}, - Status: clients.EnabledStatus, - }, - }, - err: repoerr.ErrCreateEntity, - }, - { - desc: "add client with invalid client domain id", - clients: []clients.Client{ - { - ID: testsutil.GenerateUUID(t), - Domain: invalidDomainID, - Credentials: clients.Credentials{ - Identity: "invaliddomainid-client@example.com", - Secret: testsutil.GenerateUUID(t), - }, - PrivateMetadata: map[string]any{"key": "value"}, - Metadata: map[string]any{"key": "value"}, - Status: clients.EnabledStatus, - }, - }, - err: repoerr.ErrCreateEntity, - }, - { - desc: "add client with invalid client identity", - clients: []clients.Client{ - { - ID: testsutil.GenerateUUID(t), - Name: clientName, - Credentials: clients.Credentials{ - Identity: invalidName, - Secret: testsutil.GenerateUUID(t), - }, - PrivateMetadata: map[string]any{"key": "value"}, - Metadata: map[string]any{"key": "value"}, - Status: clients.EnabledStatus, - }, - }, - err: repoerr.ErrCreateEntity, - }, - { - desc: "add client with a missing client identity", - clients: []clients.Client{ - { - ID: testsutil.GenerateUUID(t), - Domain: testsutil.GenerateUUID(t), - Name: "missing-client-identity", - Credentials: clients.Credentials{ - Identity: "", - Secret: testsutil.GenerateUUID(t), - }, - PrivateMetadata: map[string]any{"key": "value"}, - Metadata: map[string]any{"key": "value"}, - }, - }, - err: nil, - }, - { - desc: "add client with a missing client secret", - clients: []clients.Client{ - { - ID: testsutil.GenerateUUID(t), - Domain: testsutil.GenerateUUID(t), - Credentials: clients.Credentials{ - Identity: "missing-client-secret@example.com", - Secret: "", - }, - PrivateMetadata: map[string]any{"key": "value"}, - Metadata: map[string]any{"key": "value"}, - }, - }, - err: nil, - }, - { - desc: "add a client with invalid private metadata", - clients: []clients.Client{ - { - ID: testsutil.GenerateUUID(t), - Name: namegen.Generate(), - Credentials: clients.Credentials{ - Identity: fmt.Sprintf("%s@example.com", namegen.Generate()), - Secret: testsutil.GenerateUUID(t), - }, - PrivateMetadata: map[string]any{ - "key": make(chan int), - }, - }, - }, - err: errors.ErrMalformedEntity, - }, - { - desc: "add a client with invalid metadata", - clients: []clients.Client{ - { - ID: testsutil.GenerateUUID(t), - Name: namegen.Generate(), - Credentials: clients.Credentials{ - Identity: fmt.Sprintf("%s@example.com", namegen.Generate()), - Secret: testsutil.GenerateUUID(t), - }, - PrivateMetadata: map[string]any{ - "key": make(chan int), - }, - }, - }, - err: errors.ErrMalformedEntity, - }, - { - desc: "add client with duplicate name", - clients: []clients.Client{ - { - ID: duplicateClientID, - Domain: validClient.Domain, - Name: validClient.Name, - PrivateMetadata: map[string]any{"key": "different_value"}, - Metadata: map[string]any{}, - CreatedAt: validTimestamp, - Status: clients.EnabledStatus, - }, - }, - err: nil, - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - rClients, err := repo.Save(context.Background(), tc.clients...) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - if err == nil { - for i := range rClients { - tc.clients[i].Credentials.Secret = rClients[i].Credentials.Secret - } - assert.Equal(t, tc.clients, rClients, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.clients, rClients)) - } - }) - } -} - -func TestClientsRetrieveBySecret(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM clients") - require.Nil(t, err, fmt.Sprintf("clean clients unexpected error: %s", err)) - }) - repo := postgres.NewRepository(database) - - client := clients.Client{ - ID: testsutil.GenerateUUID(t), - Name: clientName, - Credentials: clients.Credentials{ - Identity: clientIdentity, - Secret: testsutil.GenerateUUID(t), - }, - Domain: testsutil.GenerateUUID(t), - Metadata: clients.Metadata{}, - PrivateMetadata: clients.Metadata{}, - Status: clients.EnabledStatus, - } - - _, err := repo.Save(context.Background(), client) - require.Nil(t, err, fmt.Sprintf("unexpected error: %s", err)) - - cases := []struct { - desc string - secret string - id string - response clients.Client - prefix authn.AuthPrefix - err error - }{ - { - desc: "retrieve client by secret with no id", - secret: client.Credentials.Secret, - response: clients.Client{}, - err: repoerr.ErrNotFound, - }, - { - desc: "retrieve client by client ID and secret successfully", - secret: client.Credentials.Secret, - id: client.ID, - prefix: authn.BasicAuth, - response: client, - err: nil, - }, - { - desc: "retrieve client by client ID invalid secret", - secret: "non-existent-secret", - response: clients.Client{}, - err: repoerr.ErrNotFound, - }, - { - desc: "retrieve client by empty secret", - secret: "", - response: clients.Client{}, - err: repoerr.ErrNotFound, - }, - { - desc: "retrieve client by client ID and secret with an invalid ID type", - secret: client.Credentials.Secret, - id: client.ID, - prefix: authn.DomainAuth, - response: clients.Client{}, - err: repoerr.ErrNotFound, - }, - { - desc: "retrieve client by domain ID and secret successfully", - secret: client.Credentials.Secret, - id: client.Domain, - prefix: authn.DomainAuth, - response: client, - err: nil, - }, - { - desc: "retrieve client by domain ID and secret with an invalid ID type", - secret: client.Credentials.Secret, - id: client.Domain, - prefix: authn.BasicAuth, - response: clients.Client{}, - err: repoerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - res, err := repo.RetrieveBySecret(context.Background(), tc.secret, tc.id, tc.prefix) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, res, tc.response, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.response, res)) - }) - } -} - -func TestRetrieveByID(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM clients") - require.Nil(t, err, fmt.Sprintf("clean clients unexpected error: %s", err)) - }) - repo := postgres.NewRepository(database) - - client := clients.Client{ - ID: testsutil.GenerateUUID(t), - Name: clientName, - Credentials: clients.Credentials{ - Identity: clientIdentity, - Secret: testsutil.GenerateUUID(t), - }, - PrivateMetadata: clients.Metadata{ - "key": "value", - }, - Metadata: clients.Metadata{ - "key": "value", - }, - Status: clients.EnabledStatus, - } - - _, err := repo.Save(context.Background(), client) - require.Nil(t, err, fmt.Sprintf("unexpected error: %s", err)) - - cases := []struct { - desc string - id string - response clients.Client - err error - }{ - { - desc: "successfully", - id: client.ID, - response: client, - err: nil, - }, - { - desc: "with invalid id", - id: testsutil.GenerateUUID(t), - response: clients.Client{}, - err: repoerr.ErrNotFound, - }, - { - desc: "with empty id", - id: "", - response: clients.Client{}, - err: repoerr.ErrNotFound, - }, - } - for _, c := range cases { - t.Run(c.desc, func(t *testing.T) { - cli, err := repo.RetrieveByID(context.Background(), c.id) - assert.True(t, errors.Contains(err, c.err), fmt.Sprintf("expected %s got %s\n", c.err, err)) - if err == nil { - assert.Equal(t, client.ID, cli.ID) - assert.Equal(t, client.Name, cli.Name) - assert.Equal(t, client.PrivateMetadata, cli.PrivateMetadata) - assert.Equal(t, client.Metadata, cli.Metadata) - assert.Equal(t, client.Credentials.Identity, cli.Credentials.Identity) - assert.Equal(t, client.Credentials.Secret, cli.Credentials.Secret) - assert.Equal(t, client.Status, cli.Status) - } - }) - } -} - -func TestUpdate(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM clients") - require.Nil(t, err, fmt.Sprintf("clean clients unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - _, err := repo.Save(context.Background(), validClient) - require.Nil(t, err, fmt.Sprintf("save client unexpected error: %s", err)) - - cases := []struct { - desc string - update string - client clients.Client - err error - }{ - { - desc: "update client successfully", - update: "all", - client: clients.Client{ - ID: validClient.ID, - Name: namegen.Generate(), - PrivateMetadata: map[string]any{"key": "value"}, - Metadata: map[string]any{"key": "value"}, - UpdatedAt: validTimestamp, - UpdatedBy: testsutil.GenerateUUID(t), - }, - err: nil, - }, - { - desc: "update client name", - update: "name", - client: clients.Client{ - ID: validClient.ID, - Name: namegen.Generate(), - UpdatedAt: validTimestamp, - UpdatedBy: testsutil.GenerateUUID(t), - }, - err: nil, - }, - { - desc: "update client private metadata", - update: "private_metadata", - client: clients.Client{ - ID: validClient.ID, - PrivateMetadata: map[string]any{"key1": "value1"}, - UpdatedAt: validTimestamp, - UpdatedBy: testsutil.GenerateUUID(t), - }, - err: nil, - }, - { - desc: "update client metadata", - update: "metadata", - client: clients.Client{ - ID: validClient.ID, - Metadata: map[string]any{"key1": "value1"}, - UpdatedAt: validTimestamp, - UpdatedBy: testsutil.GenerateUUID(t), - }, - err: nil, - }, - { - desc: "update client with invalid ID", - update: "all", - client: clients.Client{ - ID: testsutil.GenerateUUID(t), - Name: namegen.Generate(), - PrivateMetadata: map[string]any{"key": "value"}, - Metadata: map[string]any{"key": "value"}, - UpdatedAt: validTimestamp, - UpdatedBy: testsutil.GenerateUUID(t), - }, - err: repoerr.ErrNotFound, - }, - { - desc: "update client with empty ID", - update: "all", - client: clients.Client{ - Name: namegen.Generate(), - PrivateMetadata: map[string]any{"key": "value"}, - Metadata: map[string]any{"key": "value"}, - UpdatedAt: validTimestamp, - UpdatedBy: testsutil.GenerateUUID(t), - }, - err: repoerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - client, err := repo.Update(context.Background(), tc.client) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - if err == nil { - assert.Equal(t, tc.client.ID, client.ID, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.client.ID, client.ID)) - assert.Equal(t, tc.client.UpdatedAt, client.UpdatedAt, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.client.UpdatedAt, client.UpdatedAt)) - assert.Equal(t, tc.client.UpdatedBy, client.UpdatedBy, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.client.UpdatedBy, client.UpdatedBy)) - switch tc.update { - case "all": - assert.Equal(t, tc.client.Name, client.Name, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.client.Name, client.Name)) - assert.Equal(t, tc.client.PrivateMetadata, client.PrivateMetadata, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.client.PrivateMetadata, client.PrivateMetadata)) - assert.Equal(t, tc.client.Metadata, client.Metadata, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.client.Metadata, client.Metadata)) - case "name": - assert.Equal(t, tc.client.Name, client.Name, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.client.Name, client.Name)) - case "private_metadata": - assert.Equal(t, tc.client.PrivateMetadata, client.PrivateMetadata, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.client.PrivateMetadata, client.PrivateMetadata)) - case "metadata": - assert.Equal(t, tc.client.Metadata, client.Metadata, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.client.Metadata, client.Metadata)) - } - } - }) - } -} - -func TestUpdateTags(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM clients") - require.Nil(t, err, fmt.Sprintf("clean clients unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - client1 := generateClient(t, clients.EnabledStatus, repo) - client2 := generateClient(t, clients.DisabledStatus, repo) - - cases := []struct { - desc string - client clients.Client - err error - }{ - { - desc: "for enabled client", - client: clients.Client{ - ID: client1.ID, - Tags: namegen.GenerateMultiple(5), - }, - err: nil, - }, - { - desc: "for disabled client", - client: clients.Client{ - ID: client2.ID, - Tags: namegen.GenerateMultiple(5), - }, - err: repoerr.ErrNotFound, - }, - { - desc: "for invalid client", - client: clients.Client{ - ID: testsutil.GenerateUUID(t), - Tags: namegen.GenerateMultiple(5), - }, - err: repoerr.ErrNotFound, - }, - { - desc: "for empty client", - client: clients.Client{}, - err: repoerr.ErrNotFound, - }, - } - for _, c := range cases { - t.Run(c.desc, func(t *testing.T) { - c.client.UpdatedAt = time.Now().UTC().Truncate(time.Millisecond) - c.client.UpdatedBy = testsutil.GenerateUUID(t) - expected, err := repo.UpdateTags(context.Background(), c.client) - assert.True(t, errors.Contains(err, c.err), fmt.Sprintf("expected %s to contain %s\n", err, c.err)) - if err == nil { - assert.Equal(t, c.client.Tags, expected.Tags) - assert.Equal(t, c.client.UpdatedAt, expected.UpdatedAt) - assert.Equal(t, c.client.UpdatedBy, expected.UpdatedBy) - } - }) - } -} - -func TestUpdateIdentity(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM clients") - require.Nil(t, err, fmt.Sprintf("clean clients unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - client1 := generateClient(t, clients.EnabledStatus, repo) - client2 := generateClient(t, clients.DisabledStatus, repo) - - cases := []struct { - desc string - client clients.Client - err error - }{ - { - desc: "for enabled client", - client: clients.Client{ - ID: client1.ID, - Credentials: clients.Credentials{ - Identity: namegen.Generate() + emailSuffix, - }, - }, - err: nil, - }, - { - desc: "for disabled client", - client: clients.Client{ - ID: client2.ID, - Credentials: clients.Credentials{ - Identity: namegen.Generate() + emailSuffix, - }, - }, - err: repoerr.ErrNotFound, - }, - { - desc: "for invalid client", - client: clients.Client{ - ID: testsutil.GenerateUUID(t), - Credentials: clients.Credentials{ - Identity: namegen.Generate() + emailSuffix, - }, - }, - err: repoerr.ErrNotFound, - }, - { - desc: "for empty client", - client: clients.Client{}, - err: repoerr.ErrNotFound, - }, - } - for _, c := range cases { - t.Run(c.desc, func(t *testing.T) { - c.client.UpdatedAt = time.Now().UTC().Truncate(time.Millisecond) - c.client.UpdatedBy = testsutil.GenerateUUID(t) - expected, err := repo.UpdateIdentity(context.Background(), c.client) - assert.True(t, errors.Contains(err, c.err), fmt.Sprintf("expected %s to contain %s\n", err, c.err)) - if err == nil { - assert.Equal(t, c.client.Credentials.Identity, expected.Credentials.Identity) - assert.Equal(t, c.client.UpdatedAt, expected.UpdatedAt) - assert.Equal(t, c.client.UpdatedBy, expected.UpdatedBy) - } - }) - } -} - -func TestUpdateSecret(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM clients") - require.Nil(t, err, fmt.Sprintf("clean clients unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - client1 := generateClient(t, clients.EnabledStatus, repo) - client2 := generateClient(t, clients.DisabledStatus, repo) - - cases := []struct { - desc string - client clients.Client - err error - }{ - { - desc: "for enabled client", - client: clients.Client{ - ID: client1.ID, - Credentials: clients.Credentials{ - Secret: "newpassword", - }, - }, - err: nil, - }, - { - desc: "for disabled client", - client: clients.Client{ - ID: client2.ID, - Credentials: clients.Credentials{ - Secret: "newpassword", - }, - }, - err: repoerr.ErrNotFound, - }, - { - desc: "for invalid client", - client: clients.Client{ - ID: testsutil.GenerateUUID(t), - Credentials: clients.Credentials{ - Secret: "newpassword", - }, - }, - err: repoerr.ErrNotFound, - }, - { - desc: "for empty client", - client: clients.Client{}, - err: repoerr.ErrNotFound, - }, - } - for _, c := range cases { - t.Run(c.desc, func(t *testing.T) { - c.client.UpdatedAt = time.Now().UTC().Truncate(time.Millisecond) - c.client.UpdatedBy = testsutil.GenerateUUID(t) - _, err := repo.UpdateSecret(context.Background(), c.client) - assert.True(t, errors.Contains(err, c.err), fmt.Sprintf("expected %s to contain %s\n", err, c.err)) - if err == nil { - rc, err := repo.RetrieveByID(context.Background(), c.client.ID) - require.Nil(t, err, fmt.Sprintf("retrieve client by id during update of secret unexpected error: %s", err)) - assert.Equal(t, c.client.Credentials.Secret, rc.Credentials.Secret) - assert.Equal(t, c.client.UpdatedAt, rc.UpdatedAt) - assert.Equal(t, c.client.UpdatedBy, rc.UpdatedBy) - } - }) - } -} - -func TestChangeStatus(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM clients") - require.Nil(t, err, fmt.Sprintf("clean clients unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - client1 := generateClient(t, clients.EnabledStatus, repo) - client2 := generateClient(t, clients.DisabledStatus, repo) - - cases := []struct { - desc string - client clients.Client - err error - }{ - { - desc: "for an enabled client", - client: clients.Client{ - ID: client1.ID, - Status: clients.DisabledStatus, - }, - err: nil, - }, - { - desc: "for a disabled client", - client: clients.Client{ - ID: client2.ID, - Status: clients.EnabledStatus, - }, - err: nil, - }, - { - desc: "for invalid client", - client: clients.Client{ - ID: testsutil.GenerateUUID(t), - Status: clients.DisabledStatus, - }, - err: repoerr.ErrNotFound, - }, - { - desc: "for empty client", - client: clients.Client{}, - err: repoerr.ErrNotFound, - }, - } - - for _, c := range cases { - t.Run(c.desc, func(t *testing.T) { - c.client.UpdatedAt = time.Now().UTC().Truncate(time.Millisecond) - c.client.UpdatedBy = testsutil.GenerateUUID(t) - expected, err := repo.ChangeStatus(context.Background(), c.client) - assert.True(t, errors.Contains(err, c.err), fmt.Sprintf("expected %s to contain %s\n", err, c.err)) - if err == nil { - assert.Equal(t, c.client.Status, expected.Status) - assert.Equal(t, c.client.UpdatedAt, expected.UpdatedAt) - assert.Equal(t, c.client.UpdatedBy, expected.UpdatedBy) - } - }) - } -} - -func TestRetrieveByIDsWithRoles(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM clients") - require.Nil(t, err, fmt.Sprintf("clean clients unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - nClients := uint64(10) - - domainID := testsutil.GenerateUUID(t) - userID := testsutil.GenerateUUID(t) - expectedClients := []clients.Client{} - for range nClients { - client := clients.Client{ - ID: testsutil.GenerateUUID(t), - Domain: domainID, - Name: namegen.Generate(), - Credentials: clients.Credentials{ - Identity: namegen.Generate() + emailSuffix, - Secret: testsutil.GenerateUUID(t), - }, - Tags: namegen.GenerateMultiple(5), - Metadata: clients.Metadata{ - "department": namegen.Generate(), - }, - Status: clients.EnabledStatus, - CreatedAt: time.Now().UTC().Truncate(time.Microsecond), - } - _, err := repo.Save(context.Background(), client) - require.Nil(t, err, fmt.Sprintf("add new client: expected nil got %s\n", err)) - newRolesProvision := []roles.RoleProvision{ - { - Role: roles.Role{ - ID: testsutil.GenerateUUID(t) + "_" + client.ID, - Name: "admin", - EntityID: client.ID, - CreatedAt: validTimestamp, - CreatedBy: userID, - }, - OptionalActions: availableActions, - OptionalMembers: []string{userID}, - }, - } - npr, err := repo.AddRoles(context.Background(), newRolesProvision) - require.Nil(t, err, fmt.Sprintf("add roles unexpected error: %s", err)) - expectedClient := client - expectedClient.Roles = []roles.MemberRoleActions{ - { - RoleID: npr[0].Role.ID, - RoleName: npr[0].Role.Name, - Actions: npr[0].OptionalActions, - AccessType: directAccess, - }, - } - expectedClients = append(expectedClients, expectedClient) - } - - cases := []struct { - desc string - clientID string - userID string - response clients.Client - err error - }{ - { - desc: "retrieve client with role successfully", - clientID: expectedClients[0].ID, - userID: userID, - response: expectedClients[0], - err: nil, - }, - { - desc: "retrieve another client with role successfully", - clientID: expectedClients[1].ID, - userID: userID, - response: expectedClients[1], - err: nil, - }, - { - desc: "retrieve client with invalid client id", - clientID: testsutil.GenerateUUID(t), - userID: userID, - response: clients.Client{}, - err: repoerr.ErrNotFound, - }, - { - desc: "retrieve client with empty client id", - clientID: "", - userID: userID, - response: clients.Client{}, - err: repoerr.ErrNotFound, - }, - { - desc: "retrieve client with invalid user id", - clientID: expectedClients[0].ID, - userID: testsutil.GenerateUUID(t), - response: clients.Client{}, - err: repoerr.ErrNotFound, - }, - { - desc: "retrieve client with empty user id", - clientID: expectedClients[0].ID, - userID: "", - response: clients.Client{}, - err: repoerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - client, err := repo.RetrieveByIDWithRoles(context.Background(), tc.clientID, tc.userID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("expected %s to contain %s\n", err, tc.err)) - if err == nil { - assert.Equal(t, tc.response, client, fmt.Sprintf("expected %v got %v\n", tc.response, client)) - } - }) - } -} - -func TestRetrieveAll(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM clients") - require.Nil(t, err, fmt.Sprintf("clean clients unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - nClients := uint64(200) - connectedClient := clients.Client{} - - channelID := testsutil.GenerateUUID(t) - expectedClients := []clients.Client{} - disabledClients := []clients.Client{} - reversedClients := []clients.Client{} - baseTime := time.Now().UTC().Truncate(time.Millisecond) - for i := uint64(0); i < nClients; i++ { - client := clients.Client{ - ID: testsutil.GenerateUUID(t), - Domain: testsutil.GenerateUUID(t), - Name: namegen.Generate(), - Credentials: clients.Credentials{ - Identity: namegen.Generate() + emailSuffix, - Secret: testsutil.GenerateUUID(t), - }, - Tags: []string{"tag1", "tag2"}, - Metadata: clients.Metadata{ - "department": namegen.Generate(), - }, - Status: clients.EnabledStatus, - CreatedAt: baseTime.Add(time.Duration(i) * time.Millisecond), - UpdatedAt: baseTime.Add(time.Duration(i) * time.Millisecond), - } - if i%50 == 0 { - client.Status = clients.DisabledStatus - } - if i%99 == 0 { - client.Tags = []string{"tag1", "tag3"} - } - _, err := repo.Save(context.Background(), client) - if i == 0 { - conn := clients.Connection{ - ClientID: client.ID, - ChannelID: channelID, - DomainID: client.Domain, - Type: connections.Publish, - } - err = repo.AddConnections(context.Background(), []clients.Connection{conn}) - assert.Nil(t, err, fmt.Sprintf("add connection unexpected error: %s", err)) - connectedClient = client - connectedClient.ConnectionTypes = []connections.ConnType{connections.Publish} - } - require.Nil(t, err, fmt.Sprintf("add new client: expected nil got %s\n", err)) - expectedClients = append(expectedClients, client) - if client.Status == clients.DisabledStatus { - disabledClients = append(disabledClients, client) - } - } - - for i := len(expectedClients) - 1; i >= 0; i-- { - reversedClients = append(reversedClients, expectedClients[i]) - } - - cases := []struct { - desc string - pm clients.Page - response clients.ClientsPage - err error - }{ - { - desc: "with empty page", - pm: clients.Page{}, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 196, - Offset: 0, - Limit: 0, - }, - Clients: []clients.Client(nil), - }, - }, - { - desc: "with offset only", - pm: clients.Page{ - Offset: 50, - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: nClients, - Offset: 50, - Limit: 0, - }, - Clients: []clients.Client(nil), - }, - }, - { - desc: "with limit only", - pm: clients.Page{ - Limit: 10, - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: nClients, - Offset: 0, - Limit: 10, - }, - Clients: expectedClients[:10], - }, - }, - { - desc: "retrieve all clients", - pm: clients.Page{ - Offset: 0, - Limit: nClients, - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: nClients, - Offset: 0, - Limit: nClients, - }, - Clients: expectedClients, - }, - }, - { - desc: "with offset and limit", - pm: clients.Page{ - Offset: 50, - Limit: 50, - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: nClients, - Offset: 50, - Limit: 50, - }, - Clients: expectedClients[50:100], - }, - }, - { - desc: "with offset out of range and limit", - pm: clients.Page{ - Offset: 1000, - Limit: 50, - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: nClients, - Offset: 1000, - Limit: 50, - }, - Clients: []clients.Client(nil), - }, - }, - { - desc: "with offset and limit out of range", - pm: clients.Page{ - Offset: 170, - Limit: 50, - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: nClients, - Offset: 170, - Limit: 50, - }, - Clients: expectedClients[170:200], - }, - }, - { - desc: "with metadata", - pm: clients.Page{ - Offset: 0, - Limit: nClients, - Metadata: expectedClients[0].Metadata, - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 1, - Offset: 0, - Limit: nClients, - }, - Clients: []clients.Client{expectedClients[0]}, - }, - }, - { - desc: "with wrong metadata", - pm: clients.Page{ - Offset: 0, - Limit: nClients, - Metadata: clients.Metadata{ - "faculty": namegen.Generate(), - }, - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 0, - Offset: 0, - Limit: nClients, - }, - Clients: []clients.Client(nil), - }, - }, - { - desc: "with invalid metadata", - pm: clients.Page{ - Offset: 0, - Limit: nClients, - Metadata: clients.Metadata{ - "faculty": make(chan int), - }, - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: uint64(nClients), - Offset: 0, - Limit: nClients, - }, - Clients: []clients.Client(nil), - }, - err: repoerr.ErrViewEntity, - }, - { - desc: "with name", - pm: clients.Page{ - Offset: 0, - Limit: nClients, - Name: expectedClients[0].Name, - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 1, - Offset: 0, - Limit: nClients, - }, - Clients: []clients.Client{expectedClients[0]}, - }, - }, - { - desc: "with wrong name", - pm: clients.Page{ - Offset: 0, - Limit: nClients, - Name: namegen.Generate(), - Status: clients.AllStatus, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 0, - Offset: 0, - Limit: nClients, - }, - Clients: []clients.Client(nil), - }, - }, - { - desc: "with identity", - pm: clients.Page{ - Offset: 0, - Limit: nClients, - Identity: expectedClients[0].Credentials.Identity, - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 1, - Offset: 0, - Limit: nClients, - }, - Clients: []clients.Client{expectedClients[0]}, - }, - }, - { - desc: "with wrong identity", - pm: clients.Page{ - Offset: 0, - Limit: nClients, - Identity: namegen.Generate(), - Status: clients.AllStatus, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 0, - Offset: 0, - Limit: nClients, - }, - Clients: []clients.Client(nil), - }, - }, - { - desc: "with domain", - pm: clients.Page{ - Offset: 0, - Limit: nClients, - Domain: expectedClients[0].Domain, - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 1, - Offset: 0, - Limit: nClients, - }, - Clients: []clients.Client{expectedClients[0]}, - }, - }, - { - desc: "with wrong domain", - pm: clients.Page{ - Offset: 0, - Limit: nClients, - Domain: testsutil.GenerateUUID(t), - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 0, - Offset: 0, - Limit: nClients, - }, - Clients: []clients.Client(nil), - }, - }, - { - desc: "with enabled status", - pm: clients.Page{ - Offset: 0, - Limit: 10, - Status: clients.EnabledStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 196, - Offset: 0, - Limit: 10, - }, - Clients: expectedClients[1:11], - }, - }, - { - desc: "with disabled status", - pm: clients.Page{ - Offset: 0, - Limit: nClients, - Status: clients.DisabledStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 4, - Offset: 0, - Limit: nClients, - }, - Clients: disabledClients, - }, - }, - { - desc: "with combined status", - pm: clients.Page{ - Offset: 0, - Limit: nClients, - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: nClients, - Offset: 0, - Limit: nClients, - }, - Clients: expectedClients, - }, - }, - { - desc: "with the wrong status", - pm: clients.Page{ - Offset: 0, - Limit: nClients, - Status: 10, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 0, - Offset: 0, - Limit: nClients, - }, - Clients: []clients.Client(nil), - }, - }, - { - desc: "with single tag", - pm: clients.Page{ - Offset: 0, - Limit: nClients, - Tags: clients.TagsQuery{Elements: []string{"tag1"}, Operator: clients.OrOp}, - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 200, - Offset: 0, - Limit: uint64(nClients), - }, - Clients: expectedClients, - }, - }, - { - desc: "with multiple tags and OR operator", - pm: clients.Page{ - Offset: 0, - Limit: nClients, - Tags: clients.TagsQuery{Elements: []string{"tag2", "tag3"}, Operator: clients.OrOp}, - Status: clients.AllStatus, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 200, - Offset: 0, - Limit: uint64(nClients), - }, - Clients: expectedClients, - }, - }, - { - desc: "with multiple tags and AND operator", - pm: clients.Page{ - Offset: 0, - Limit: nClients, - Tags: clients.TagsQuery{Elements: []string{"tag1", "tag3"}, Operator: clients.AndOp}, - Status: clients.AllStatus, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 3, - Offset: 0, - Limit: uint64(nClients), - }, - Clients: []clients.Client{expectedClients[0], expectedClients[99], expectedClients[198]}, - }, - }, - { - desc: "with wrong tags", - pm: clients.Page{ - Offset: 0, - Limit: nClients, - Tags: clients.TagsQuery{Elements: []string{namegen.Generate(), namegen.Generate()}, Operator: clients.OrOp}, - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 0, - Offset: 0, - Limit: nClients, - }, - Clients: []clients.Client(nil), - }, - }, - { - desc: "with multiple parameters", - pm: clients.Page{ - Offset: 0, - Limit: nClients, - Metadata: expectedClients[0].Metadata, - Name: expectedClients[0].Name, - Tags: clients.TagsQuery{Elements: []string{expectedClients[0].Tags[0]}, Operator: clients.OrOp}, - Identity: expectedClients[0].Credentials.Identity, - Domain: expectedClients[0].Domain, - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 1, - Offset: 0, - Limit: nClients, - }, - Clients: []clients.Client{expectedClients[0]}, - }, - }, - { - desc: "with id", - pm: clients.Page{ - Offset: 0, - Limit: nClients, - ID: expectedClients[0].ID, - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 1, - Offset: 0, - Limit: nClients, - }, - Clients: []clients.Client{expectedClients[0]}, - }, - }, - { - desc: "with wrong id", - pm: clients.Page{ - Offset: 0, - Limit: nClients, - ID: testsutil.GenerateUUID(t), - Status: clients.AllStatus, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 0, - Offset: 0, - Limit: nClients, - }, - Clients: []clients.Client(nil), - }, - }, - { - desc: "with channel id", - pm: clients.Page{ - Offset: 0, - Limit: nClients, - Channel: channelID, - Status: clients.AllStatus, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 1, - Offset: 0, - Limit: nClients, - }, - Clients: []clients.Client{connectedClient}, - }, - }, - { - desc: "with order by name ascending", - pm: clients.Page{ - Offset: 0, - Limit: 10, - Order: "name", - Dir: ascDir, - Status: clients.AllStatus, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: nClients, - Offset: 0, - Limit: 10, - }, - }, - }, - { - desc: "with order by name descending", - pm: clients.Page{ - Offset: 0, - Limit: 10, - Order: "name", - Dir: descDir, - Status: clients.AllStatus, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: nClients, - Offset: 0, - Limit: 10, - }, - }, - }, - { - desc: "with order by identity ascending", - pm: clients.Page{ - Offset: 0, - Limit: 10, - Order: "identity", - Dir: ascDir, - Status: clients.AllStatus, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: nClients, - Offset: 0, - Limit: 10, - }, - }, - }, - { - desc: "with order by identity descending", - pm: clients.Page{ - Offset: 0, - Limit: 10, - Order: "identity", - Dir: descDir, - Status: clients.AllStatus, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: nClients, - Offset: 0, - Limit: 10, - }, - }, - }, - { - desc: "with order by created_at ascending", - pm: clients.Page{ - Offset: 0, - Limit: 10, - Order: defOrder, - Dir: ascDir, - Status: clients.AllStatus, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: nClients, - Offset: 0, - Limit: 10, - }, - Clients: expectedClients[:10], - }, - }, - { - desc: "with order by created_at descending", - pm: clients.Page{ - Offset: 0, - Limit: 10, - Order: defOrder, - Dir: descDir, - Status: clients.AllStatus, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: nClients, - Offset: 0, - Limit: 10, - }, - Clients: reversedClients[:10], - }, - }, - { - desc: "with order by updated_at ascending", - pm: clients.Page{ - Offset: 0, - Limit: 10, - Order: "updated_at", - Dir: ascDir, - Status: clients.AllStatus, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: nClients, - Offset: 0, - Limit: 10, - }, - }, - }, - { - desc: "with order by updated_at descending", - pm: clients.Page{ - Offset: 0, - Limit: 10, - Order: "updated_at", - Dir: descDir, - Status: clients.AllStatus, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: nClients, - Offset: 0, - Limit: 10, - }, - }, - }, - { - desc: "with created_from", - pm: clients.Page{ - Offset: 0, - Limit: nClients, - Status: clients.AllStatus, - CreatedFrom: baseTime.Add(100 * time.Millisecond), - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 100, - Offset: 0, - Limit: nClients, - }, - Clients: expectedClients[100:], - }, - }, - { - desc: "with created_to", - pm: clients.Page{ - Offset: 0, - Limit: nClients, - Status: clients.AllStatus, - CreatedTo: baseTime.Add(99 * time.Millisecond), - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 100, - Offset: 0, - Limit: nClients, - }, - Clients: expectedClients[:100], - }, - }, - { - desc: "with both created_from and created_to", - pm: clients.Page{ - Offset: 0, - Limit: nClients, - Status: clients.AllStatus, - CreatedFrom: baseTime.Add(50 * time.Millisecond), - CreatedTo: baseTime.Add(149 * time.Millisecond), - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 100, - Offset: 0, - Limit: nClients, - }, - Clients: expectedClients[50:150], - }, - }, - { - desc: "with created_from returning no results", - pm: clients.Page{ - Offset: 0, - Limit: nClients, - Status: clients.AllStatus, - CreatedFrom: baseTime.Add(500 * time.Millisecond), - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 0, - Offset: 0, - Limit: nClients, - }, - Clients: []clients.Client(nil), - }, - }, - { - desc: "with created_to returning no results", - pm: clients.Page{ - Offset: 0, - Limit: nClients, - Status: clients.AllStatus, - CreatedTo: baseTime.Add(-10 * time.Millisecond), - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 0, - Offset: 0, - Limit: nClients, - }, - Clients: []clients.Client(nil), - }, - }, - } - for _, c := range cases { - t.Run(c.desc, func(t *testing.T) { - page, err := repo.RetrieveAll(context.Background(), c.pm) - assert.True(t, errors.Contains(err, c.err), fmt.Sprintf("expected %s to contain %s\n", err, c.err)) - if err == nil { - assert.Equal(t, c.response.Total, page.Total) - assert.Equal(t, c.response.Offset, page.Offset) - assert.Equal(t, c.response.Limit, page.Limit) - if len(c.response.Clients) > 0 { - expected := stripClientDetails(c.response.Clients) - got := stripClientDetails(page.Clients) - assert.ElementsMatch(t, expected, got, fmt.Sprintf("expected %v got %v\n", expected, got)) - } - verifyClientsOrdering(t, page.Clients, c.pm.Order, c.pm.Dir) - } - }) - } -} - -func TestRetrieveUserClients(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM clients") - require.Nil(t, err, fmt.Sprintf("clean clients unexpected error: %s", err)) - _, err = db.Exec("DELETE FROM groups") - require.Nil(t, err, fmt.Sprintf("clean groups unexpected error: %s", err)) - _, err = db.Exec("DELETE FROM domains") - require.Nil(t, err, fmt.Sprintf("clean clients unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - nClients := uint64(10) - - emptyGroupParam := "" - userID := testsutil.GenerateUUID(t) - domainMemberID := testsutil.GenerateUUID(t) - groupMemberID := testsutil.GenerateUUID(t) - channelID := testsutil.GenerateUUID(t) - domain := generateDomain(t, userID, domainMemberID) - group := generateGroup(t, userID, groupMemberID, domain.ID) - groupClient := clients.Client{} - parentGroupClient := clients.Client{} - connectedClient := clients.Client{} - directClients := []clients.Client{} - domainClients := []clients.Client{} - baseTime := time.Now().UTC().Truncate(time.Microsecond) - for i := range nClients { - client := clients.Client{ - ID: testsutil.GenerateUUID(t), - Domain: domain.ID, - Name: namegen.Generate(), - Credentials: clients.Credentials{ - Identity: namegen.Generate() + emailSuffix, - Secret: testsutil.GenerateUUID(t), - }, - Tags: namegen.GenerateMultiple(5), - Metadata: clients.Metadata{ - "department": namegen.Generate(), - }, - Status: clients.EnabledStatus, - CreatedAt: baseTime.Add(time.Duration(i) * time.Microsecond), - UpdatedAt: baseTime.Add(time.Duration(i) * time.Microsecond), - } - if i == 1 { - client.ParentGroup = group.ID - } - _, err := repo.Save(context.Background(), client) - require.Nil(t, err, fmt.Sprintf("add new client: expected nil got %s\n", err)) - newRolesProvision := []roles.RoleProvision{ - { - Role: roles.Role{ - ID: testsutil.GenerateUUID(t) + "_" + client.ID, - Name: "admin", - EntityID: client.ID, - CreatedAt: validTimestamp, - CreatedBy: userID, - }, - OptionalActions: availableActions, - OptionalMembers: []string{userID}, - }, - } - npr, err := repo.AddRoles(context.Background(), newRolesProvision) - require.Nil(t, err, fmt.Sprintf("add roles unexpected error: %s", err)) - directClient := client - directClient.RoleID = npr[0].Role.ID - directClient.RoleName = npr[0].Role.Name - directClient.AccessType = directAccess - directClient.AccessProviderRoleActions = []string{} - if i == 1 { - directClient.ParentGroupPath = group.ID - } - directClients = append(directClients, directClient) - if i == 1 { - parentGroupClient = directClient - parentGroupClient.ParentGroupPath = group.ID - client.ParentGroupPath = group.ID - groupClient = client - groupClient.AccessType = directGroupAccess - groupClient.AccessProviderId = group.ID - groupClient.AccessProviderRoleId = group.Roles[0].RoleID - groupClient.AccessProviderRoleName = group.Roles[0].RoleName - groupClient.AccessProviderRoleActions = groupAvailableActions - } - if i == 2 { - conn := clients.Connection{ - ClientID: client.ID, - ChannelID: channelID, - DomainID: client.Domain, - Type: connections.Publish, - } - err = repo.AddConnections(context.Background(), []clients.Connection{conn}) - assert.Nil(t, err, fmt.Sprintf("add connection unexpected error: %s", err)) - connectedClient = client - connectedClient.RoleID = npr[0].Role.ID - connectedClient.RoleName = npr[0].Role.Name - connectedClient.AccessType = directAccess - connectedClient.AccessProviderRoleActions = []string{} - connectedClient.ConnectionTypes = []connections.ConnType{connections.Publish} - } - domainClient := client - domainClient.AccessType = domainAccess - domainClient.AccessProviderId = domain.ID - domainClient.AccessProviderRoleId = domain.Roles[0].RoleID - domainClient.AccessProviderRoleName = domain.Roles[0].RoleName - domainClient.AccessProviderRoleActions = domainAvailableActions - domainClients = append(domainClients, domainClient) - } - - reversedDirectClients := []clients.Client{} - for i := len(directClients) - 1; i >= 0; i-- { - reversedDirectClients = append(reversedDirectClients, directClients[i]) - } - - cases := []struct { - desc string - domainID string - userID string - pm clients.Page - response clients.ClientsPage - err error - }{ - { - desc: "retrieve clients with empty page", - domainID: domain.ID, - userID: userID, - pm: clients.Page{}, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 10, - Offset: 0, - Limit: 0, - }, - Clients: []clients.Client(nil), - }, - }, - { - desc: "retrieve clients with offset and limit", - domainID: domain.ID, - userID: userID, - pm: clients.Page{ - Offset: 5, - Limit: 10, - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: nClients, - Offset: 5, - Limit: 10, - }, - Clients: directClients[5:10], - }, - }, - { - desc: "retrieve clients with member id of parent group wth direct group access", - domainID: domain.ID, - userID: groupMemberID, - pm: clients.Page{ - Offset: 0, - Limit: 10, - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 1, - Offset: 0, - Limit: 10, - }, - Clients: []clients.Client{groupClient}, - }, - }, - { - desc: "retrieve clients with member id of domain with domain access", - domainID: domain.ID, - userID: domainMemberID, - pm: clients.Page{ - Offset: 0, - Limit: 10, - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 10, - Offset: 0, - Limit: 10, - }, - Clients: domainClients, - }, - }, - { - desc: "retrieve clients connected to a channel", - domainID: domain.ID, - userID: userID, - pm: clients.Page{ - Offset: 0, - Limit: 10, - Channel: channelID, - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 1, - Offset: 0, - Limit: 10, - }, - Clients: []clients.Client{connectedClient}, - }, - }, - { - desc: "retrieve clients with offset out of range and limit", - domainID: domain.ID, - userID: userID, - pm: clients.Page{ - Offset: 1000, - Limit: 50, - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: nClients, - Offset: 1000, - Limit: 50, - }, - Clients: []clients.Client(nil), - }, - }, - { - desc: "retrieve clients with metadata", - domainID: domain.ID, - userID: userID, - pm: clients.Page{ - Offset: 0, - Limit: nClients, - Metadata: directClients[0].Metadata, - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 1, - Offset: 0, - Limit: nClients, - }, - Clients: []clients.Client{directClients[0]}, - }, - }, - { - desc: "retrieve clients with wrong metadata", - domainID: domain.ID, - userID: userID, - pm: clients.Page{ - Offset: 0, - Limit: nClients, - Metadata: clients.Metadata{ - "faculty": namegen.Generate(), - }, - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 0, - Offset: 0, - Limit: nClients, - }, - Clients: []clients.Client(nil), - }, - }, - { - desc: "retrieve clients with invalid metadata", - domainID: domain.ID, - userID: userID, - pm: clients.Page{ - Offset: 0, - Limit: nClients, - Metadata: clients.Metadata{ - "faculty": make(chan int), - }, - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: uint64(nClients), - Offset: 0, - Limit: nClients, - }, - Clients: []clients.Client(nil), - }, - err: repoerr.ErrViewEntity, - }, - { - desc: "retrieve clients with name", - domainID: domain.ID, - userID: userID, - pm: clients.Page{ - Offset: 0, - Limit: nClients, - Name: directClients[0].Name, - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 1, - Offset: 0, - Limit: nClients, - }, - Clients: []clients.Client{directClients[0]}, - }, - }, - { - desc: "retrieve clients with wrong name", - domainID: domain.ID, - userID: userID, - pm: clients.Page{ - Offset: 0, - Limit: nClients, - Name: namegen.Generate(), - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 0, - Offset: 0, - Limit: nClients, - }, - Clients: []clients.Client(nil), - }, - }, - { - desc: "retrieve cliens with identity", - domainID: domain.ID, - userID: userID, - pm: clients.Page{ - Offset: 0, - Limit: nClients, - Identity: directClients[0].Credentials.Identity, - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 1, - Offset: 0, - Limit: nClients, - }, - Clients: []clients.Client{directClients[0]}, - }, - }, - { - desc: "retrieve clients with wrong identity", - domainID: domain.ID, - userID: userID, - pm: clients.Page{ - Offset: 0, - Limit: nClients, - Identity: namegen.Generate(), - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 0, - Offset: 0, - Limit: nClients, - }, - Clients: []clients.Client(nil), - }, - }, - { - desc: "retrieve clients with tag", - domainID: domain.ID, - userID: userID, - pm: clients.Page{ - Offset: 0, - Limit: nClients, - Tags: clients.TagsQuery{Elements: []string{directClients[0].Tags[0]}, Operator: clients.OrOp}, - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 1, - Offset: 0, - Limit: uint64(nClients), - }, - Clients: []clients.Client{directClients[0]}, - }, - }, - { - desc: "retrieve clients with wrong tags", - domainID: domain.ID, - userID: userID, - pm: clients.Page{ - Offset: 0, - Limit: nClients, - Tags: clients.TagsQuery{Elements: []string{namegen.Generate()}, Operator: clients.OrOp}, - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 0, - Offset: 0, - Limit: nClients, - }, - Clients: []clients.Client(nil), - }, - }, - { - desc: "retrieve clients with multiple parameters", - domainID: domain.ID, - userID: userID, - pm: clients.Page{ - Offset: 0, - Limit: nClients, - Metadata: directClients[0].Metadata, - Name: directClients[0].Name, - Tags: clients.TagsQuery{Elements: []string{directClients[0].Tags[0]}, Operator: clients.OrOp}, - Identity: directClients[0].Credentials.Identity, - Status: clients.AllStatus, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 1, - Offset: 0, - Limit: nClients, - }, - Clients: []clients.Client{directClients[0]}, - }, - }, - { - desc: "retrieve clients with id", - domainID: domain.ID, - userID: userID, - pm: clients.Page{ - Offset: 0, - Limit: nClients, - ID: directClients[0].ID, - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 1, - Offset: 0, - Limit: nClients, - }, - Clients: []clients.Client{directClients[0]}, - }, - }, - { - desc: "retrieve clients with wrong id", - domainID: domain.ID, - userID: userID, - pm: clients.Page{ - Offset: 0, - Limit: nClients, - ID: testsutil.GenerateUUID(t), - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 0, - Offset: 0, - Limit: nClients, - }, - Clients: []clients.Client(nil), - }, - }, - { - desc: "retrieve clients with wrong domain id", - domainID: testsutil.GenerateUUID(t), - userID: userID, - pm: clients.Page{ - Offset: 0, - Limit: nClients, - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 0, - Offset: 0, - Limit: nClients, - }, - Clients: []clients.Client(nil), - }, - }, - { - desc: "retrieve clients with wrong user id", - domainID: domain.ID, - userID: testsutil.GenerateUUID(t), - pm: clients.Page{ - Offset: 0, - Limit: nClients, - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 0, - Offset: 0, - Limit: nClients, - }, - Clients: []clients.Client(nil), - }, - }, - { - desc: "retrieve clients with parent group", - domainID: domain.ID, - userID: userID, - pm: clients.Page{ - Offset: 0, - Limit: nClients, - Group: &group.ID, - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 1, - Offset: 0, - Limit: nClients, - }, - Clients: []clients.Client{parentGroupClient}, - }, - err: nil, - }, - { - desc: "retrieve clients with no parent group", - domainID: domain.ID, - userID: userID, - pm: clients.Page{ - Offset: 0, - Limit: nClients, - Group: &emptyGroupParam, - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 0, - Offset: 0, - Limit: nClients, - }, - Clients: []clients.Client{}, - }, - }, - { - desc: "retrieve clients with access type", - domainID: domain.ID, - userID: domainMemberID, - pm: clients.Page{ - Offset: 0, - Limit: 10, - AccessType: domainAccess, - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 10, - Offset: 0, - Limit: 10, - }, - Clients: domainClients, - }, - }, - { - desc: "retrieve clients with wrong access type", - domainID: domain.ID, - userID: userID, - pm: clients.Page{ - Offset: 0, - Limit: nClients, - AccessType: domainAccess, - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 0, - Offset: 0, - Limit: nClients, - }, - Clients: []clients.Client{}, - }, - }, - { - desc: "retrieve clients with role ID", - domainID: domain.ID, - userID: userID, - pm: clients.Page{ - Offset: 0, - Limit: nClients, - RoleID: directClients[0].RoleID, - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 1, - Offset: 0, - Limit: nClients, - }, - Clients: []clients.Client{directClients[0]}, - }, - }, - { - desc: "retrieve clients with wrong role ID", - domainID: domain.ID, - userID: userID, - pm: clients.Page{ - Offset: 0, - Limit: nClients, - RoleID: testsutil.GenerateUUID(t), - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 0, - Offset: 0, - Limit: nClients, - }, - Clients: []clients.Client(nil), - }, - }, - { - desc: "retrieve clients with role name", - domainID: domain.ID, - userID: userID, - pm: clients.Page{ - Offset: 0, - Limit: 1, - RoleName: directClients[0].RoleName, - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 10, - Offset: 0, - Limit: 1, - }, - Clients: directClients[0:1], - }, - }, - { - desc: "retrieve clients with wrong role name", - domainID: domain.ID, - userID: userID, - pm: clients.Page{ - Offset: 0, - Limit: nClients, - RoleName: namegen.Generate(), - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 0, - Offset: 0, - Limit: nClients, - }, - Clients: []clients.Client(nil), - }, - }, - { - desc: "retrieve clients with order by name ascending", - domainID: domain.ID, - userID: userID, - pm: clients.Page{ - Offset: 0, - Limit: 5, - Order: "name", - Dir: ascDir, - Status: clients.AllStatus, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: nClients, - Offset: 0, - Limit: 5, - }, - }, - }, - { - desc: "retrieve clients with order by name descending", - domainID: domain.ID, - userID: userID, - pm: clients.Page{ - Offset: 0, - Limit: 5, - Order: "name", - Dir: descDir, - Status: clients.AllStatus, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: nClients, - Offset: 0, - Limit: 5, - }, - }, - }, - { - desc: "retrieve clients with order by identity ascending", - domainID: domain.ID, - userID: userID, - pm: clients.Page{ - Offset: 0, - Limit: 5, - Order: "identity", - Dir: ascDir, - Status: clients.AllStatus, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: nClients, - Offset: 0, - Limit: 5, - }, - }, - }, - { - desc: "retrieve clients with order by identity descending", - domainID: domain.ID, - userID: userID, - pm: clients.Page{ - Offset: 0, - Limit: 5, - Order: "identity", - Dir: descDir, - Status: clients.AllStatus, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: nClients, - Offset: 0, - Limit: 5, - }, - }, - }, - { - desc: "retrieve clients with order by created_at ascending", - domainID: domain.ID, - userID: userID, - pm: clients.Page{ - Offset: 0, - Limit: 5, - Order: defOrder, - Dir: ascDir, - Status: clients.AllStatus, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: nClients, - Offset: 0, - Limit: 5, - }, - Clients: directClients[:5], - }, - }, - { - desc: "retrieve clients with order by created_at descending", - domainID: domain.ID, - userID: userID, - pm: clients.Page{ - Offset: 0, - Limit: 5, - Order: defOrder, - Dir: descDir, - Status: clients.AllStatus, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: nClients, - Offset: 0, - Limit: 5, - }, - Clients: reversedDirectClients[:5], - }, - }, - { - desc: "retrieve clients with order by updated_at ascending", - domainID: domain.ID, - userID: userID, - pm: clients.Page{ - Offset: 0, - Limit: 5, - Order: "updated_at", - Dir: ascDir, - Status: clients.AllStatus, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: nClients, - Offset: 0, - Limit: 5, - }, - }, - }, - { - desc: "retrieve clients with order by updated_at descending", - domainID: domain.ID, - userID: userID, - pm: clients.Page{ - Offset: 0, - Limit: 5, - Order: "updated_at", - Dir: descDir, - Status: clients.AllStatus, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: nClients, - Offset: 0, - Limit: 5, - }, - }, - }, - { - desc: "retrieve clients connected to a channel with only total", - domainID: domain.ID, - userID: userID, - pm: clients.Page{ - Offset: 0, - Limit: 10, - Channel: channelID, - Status: clients.AllStatus, - OnlyTotal: true, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 1, - Offset: 0, - Limit: 10, - }, - Clients: []clients.Client(nil), - }, - }, - { - desc: "retrieve clients connected to a non-existent channel", - domainID: domain.ID, - userID: userID, - pm: clients.Page{ - Offset: 0, - Limit: 10, - Channel: testsutil.GenerateUUID(t), - Status: clients.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 0, - Offset: 0, - Limit: 10, - }, - Clients: []clients.Client(nil), - }, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - page, err := repo.RetrieveUserClients(context.Background(), tc.domainID, tc.userID, tc.pm) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("expected %s to contain %s\n", err, tc.err)) - if err == nil { - assert.Equal(t, tc.response.Total, page.Total) - assert.Equal(t, tc.response.Offset, page.Offset) - assert.Equal(t, tc.response.Limit, page.Limit) - if len(tc.response.Clients) > 0 { - expected := stripClientDetails(tc.response.Clients) - got := stripClientDetails(page.Clients) - assert.ElementsMatch(t, expected, got, fmt.Sprintf("expected %+v got %+v\n", expected, got)) - } - verifyClientsOrdering(t, page.Clients, tc.pm.Order, tc.pm.Dir) - } - }) - } -} - -func TestSearchClients(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM clients") - require.Nil(t, err, fmt.Sprintf("clean clients unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - name := namegen.Generate() - - nClients := uint64(200) - expectedClients := []clients.Client{} - baseTime := time.Now().UTC().Truncate(time.Microsecond) - for i := 0; i < int(nClients); i++ { - username := name + strconv.Itoa(i) + emailSuffix - client := clients.Client{ - ID: testsutil.GenerateUUID(t), - Name: username, - Credentials: clients.Credentials{ - Identity: username, - Secret: testsutil.GenerateUUID(t), - }, - Metadata: clients.Metadata{ - "department": namegen.Generate(), - }, - PrivateMetadata: clients.Metadata{}, - Status: clients.EnabledStatus, - CreatedAt: baseTime.Add(time.Duration(i) * time.Microsecond), - } - _, err := repo.Save(context.Background(), client) - require.Nil(t, err, fmt.Sprintf("save client unexpected error: %s", err)) - - expectedClients = append(expectedClients, clients.Client{ - ID: client.ID, - Name: client.Name, - Metadata: client.Metadata, - CreatedAt: client.CreatedAt, - }) - } - - page, err := repo.RetrieveAll(context.Background(), clients.Page{Offset: 0, Limit: nClients}) - require.Nil(t, err, fmt.Sprintf("retrieve all clients unexpected error: %s", err)) - assert.Equal(t, nClients, page.Total) - - cases := []struct { - desc string - page clients.Page - response clients.ClientsPage - err error - }{ - { - desc: "with empty page", - page: clients.Page{}, - response: clients.ClientsPage{ - Clients: []clients.Client(nil), - Page: clients.Page{ - Total: nClients, - Offset: 0, - Limit: 0, - }, - }, - err: nil, - }, - { - desc: "with offset only", - page: clients.Page{ - Offset: 50, - }, - response: clients.ClientsPage{ - Clients: []clients.Client(nil), - Page: clients.Page{ - Total: nClients, - Offset: 50, - Limit: 0, - }, - }, - err: nil, - }, - { - desc: "with limit only", - page: clients.Page{ - Limit: 10, - Order: "name", - Dir: ascDir, - }, - response: clients.ClientsPage{ - Clients: expectedClients[0:10], - Page: clients.Page{ - Total: nClients, - Offset: 0, - Limit: 10, - }, - }, - err: nil, - }, - { - desc: "retrieve all clients", - page: clients.Page{ - Offset: 0, - Limit: nClients, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: nClients, - Offset: 0, - Limit: nClients, - }, - Clients: expectedClients, - }, - }, - { - desc: "with offset and limit", - page: clients.Page{ - Offset: 10, - Limit: 10, - Order: "name", - Dir: ascDir, - }, - response: clients.ClientsPage{ - Clients: expectedClients[10:20], - Page: clients.Page{ - Total: nClients, - Offset: 10, - Limit: 10, - }, - }, - err: nil, - }, - { - desc: "with offset out of range and limit", - page: clients.Page{ - Offset: 1000, - Limit: 50, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: nClients, - Offset: 1000, - Limit: 50, - }, - Clients: []clients.Client(nil), - }, - }, - { - desc: "with offset and limit out of range", - page: clients.Page{ - Offset: 190, - Limit: 50, - Order: "name", - Dir: ascDir, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: nClients, - Offset: 190, - Limit: 50, - }, - Clients: expectedClients[190:200], - }, - }, - { - desc: "with shorter name", - page: clients.Page{ - Name: expectedClients[0].Name[:4], - Offset: 0, - Limit: 10, - Order: "name", - Dir: ascDir, - }, - response: clients.ClientsPage{ - Clients: findClients(expectedClients, expectedClients[0].Name[:4], 0, 10), - Page: clients.Page{ - Total: nClients, - Offset: 0, - Limit: 10, - }, - }, - err: nil, - }, - { - desc: "with longer name", - page: clients.Page{ - Name: expectedClients[0].Name, - Offset: 0, - Limit: 10, - }, - response: clients.ClientsPage{ - Clients: []clients.Client{expectedClients[0]}, - Page: clients.Page{ - Total: 1, - Offset: 0, - Limit: 10, - }, - }, - err: nil, - }, - { - desc: "with name SQL injected", - page: clients.Page{ - Name: fmt.Sprintf("%s' OR '1'='1", expectedClients[0].Name[:1]), - Offset: 0, - Limit: 10, - }, - response: clients.ClientsPage{ - Clients: []clients.Client(nil), - Page: clients.Page{ - Total: 0, - Offset: 0, - Limit: 10, - }, - }, - err: nil, - }, - { - desc: "with shorter Identity", - page: clients.Page{ - Identity: expectedClients[0].Name[:4], - Offset: 0, - Limit: 10, - Order: "name", - Dir: ascDir, - }, - response: clients.ClientsPage{ - Clients: findClients(expectedClients, expectedClients[0].Name[:4], 0, 10), - Page: clients.Page{ - Total: nClients, - Offset: 0, - Limit: 10, - }, - }, - err: nil, - }, - { - desc: "with longer Identity", - page: clients.Page{ - Identity: expectedClients[0].Name, - Offset: 0, - Limit: 10, - }, - response: clients.ClientsPage{ - Clients: []clients.Client{expectedClients[0]}, - Page: clients.Page{ - Total: 1, - Offset: 0, - Limit: 10, - }, - }, - err: nil, - }, - { - desc: "with Identity SQL injected", - page: clients.Page{ - Identity: fmt.Sprintf("%s' OR '1'='1", expectedClients[0].Name[:1]), - Offset: 0, - Limit: 10, - }, - response: clients.ClientsPage{ - Clients: []clients.Client(nil), - Page: clients.Page{ - Total: 0, - Offset: 0, - Limit: 10, - }, - }, - err: nil, - }, - { - desc: "with unknown name", - page: clients.Page{ - Name: namegen.Generate(), - Offset: 0, - Limit: 10, - }, - response: clients.ClientsPage{ - Clients: []clients.Client(nil), - Page: clients.Page{ - Total: 0, - Offset: 0, - Limit: 10, - }, - }, - err: nil, - }, - { - desc: "with unknown name SQL injected", - page: clients.Page{ - Name: fmt.Sprintf("%s' OR '1'='1", namegen.Generate()), - Offset: 0, - Limit: 10, - }, - response: clients.ClientsPage{ - Clients: []clients.Client(nil), - Page: clients.Page{ - Total: 0, - Offset: 0, - Limit: 10, - }, - }, - err: nil, - }, - { - desc: "with unknown identity", - page: clients.Page{ - Identity: namegen.Generate(), - Offset: 0, - Limit: 10, - }, - response: clients.ClientsPage{ - Clients: []clients.Client(nil), - Page: clients.Page{ - Total: 0, - Offset: 0, - Limit: 10, - }, - }, - err: nil, - }, - { - desc: "with name in asc order", - page: clients.Page{ - Order: "name", - Dir: ascDir, - Name: expectedClients[0].Name[:1], - Offset: 0, - Limit: 10, - }, - response: clients.ClientsPage{}, - err: nil, - }, - { - desc: "with name in desc order", - page: clients.Page{ - Order: "name", - Dir: descDir, - Name: expectedClients[0].Name[:1], - Offset: 0, - Limit: 10, - }, - response: clients.ClientsPage{}, - err: nil, - }, - { - desc: "with identity in asc order", - page: clients.Page{ - Order: "identity", - Dir: ascDir, - Identity: expectedClients[0].Name[:1], - Offset: 0, - Limit: 10, - }, - response: clients.ClientsPage{}, - err: nil, - }, - { - desc: "with identity in desc order", - page: clients.Page{ - Order: "identity", - Dir: descDir, - Identity: expectedClients[0].Name[:1], - Offset: 0, - Limit: 10, - }, - response: clients.ClientsPage{}, - err: nil, - }, - } - for _, c := range cases { - t.Run(c.desc, func(t *testing.T) { - switch response, err := repo.SearchClients(context.Background(), c.page); { - case err == nil: - if c.page.Order != "" && c.page.Dir != "" { - c.response = response - } - assert.Nil(t, err) - assert.Equal(t, c.response.Total, response.Total) - assert.Equal(t, c.response.Limit, response.Limit) - assert.Equal(t, c.response.Offset, response.Offset) - assert.ElementsMatch(t, response.Clients, c.response.Clients) - default: - assert.True(t, errors.Contains(err, c.err), fmt.Sprintf("expected %s to contain %s\n", err, c.err)) - } - }) - } -} - -func TestRetrieveByIDs(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM clients") - require.Nil(t, err, fmt.Sprintf("clean clients unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - num := 10 - - var items []clients.Client - baseTime := time.Now().UTC().Truncate(time.Microsecond) - for i := 0; i < num; i++ { - name := namegen.Generate() - client := clients.Client{ - ID: testsutil.GenerateUUID(t), - Domain: testsutil.GenerateUUID(t), - Name: name, - Credentials: clients.Credentials{ - Identity: name + emailSuffix, - Secret: testsutil.GenerateUUID(t), - }, - Tags: namegen.GenerateMultiple(5), - Metadata: map[string]any{"name": name}, - CreatedAt: baseTime.Add(time.Duration(i) * time.Microsecond), - Status: clients.EnabledStatus, - } - _, err := repo.Save(context.Background(), client) - require.Nil(t, err, fmt.Sprintf("add new client: expected nil got %s\n", err)) - items = append(items, client) - } - - cases := []struct { - desc string - ids []string - response clients.ClientsPage - err error - }{ - { - desc: "successfully", - ids: getIDs(items[0:3]), - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 3, - }, - Clients: items[0:3], - }, - err: nil, - }, - { - desc: "successfully", - ids: getIDs(items[3:6]), - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 3, - }, - Clients: items[3:6], - }, - err: nil, - }, - { - desc: "with empty ids", - ids: []string{}, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 0, - }, - Clients: []clients.Client(nil), - }, - err: nil, - }, - { - desc: "with valid and invalid ids", - ids: append(getIDs(items[0:3]), testsutil.GenerateUUID(t)), - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 3, - }, - Clients: items[0:3], - }, - err: nil, - }, - { - desc: "with invalid ids", - ids: []string{testsutil.GenerateUUID(t)}, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 0, - }, - Clients: []clients.Client(nil), - }, - err: nil, - }, - } - - for _, c := range cases { - response, err := repo.RetrieveByIds(context.Background(), c.ids) - assert.True(t, errors.Contains(err, c.err), fmt.Sprintf("%s: expected %s got %s\n", c.desc, c.err, err)) - if err == nil { - assert.Equal(t, c.response.Total, response.Total) - expected := stripClientDetails(c.response.Clients) - got := stripClientDetails(response.Clients) - assert.ElementsMatch(t, expected, got) - } - } -} - -func TestAddConnection(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM connections") - require.Nil(t, err, fmt.Sprintf("clean connections unexpected error: %s", err)) - _, err = db.Exec("DELETE FROM clients") - require.Nil(t, err, fmt.Sprintf("clean clients unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - client := generateClient(t, clients.EnabledStatus, repo) - - validConnection := clients.Connection{ - ClientID: client.ID, - ChannelID: testsutil.GenerateUUID(t), - DomainID: client.Domain, - Type: connections.Publish, - } - - cases := []struct { - desc string - connection clients.Connection - err error - }{ - { - desc: "add connection successfully", - connection: validConnection, - err: nil, - }, - { - desc: "add connection with non-existent client", - connection: clients.Connection{ - ClientID: testsutil.GenerateUUID(t), - ChannelID: testsutil.GenerateUUID(t), - DomainID: client.Domain, - Type: connections.Publish, - }, - err: repoerr.ErrCreateEntity, - }, - { - desc: "add connection with non-existent domain", - connection: clients.Connection{ - ClientID: client.ID, - ChannelID: testsutil.GenerateUUID(t), - DomainID: testsutil.GenerateUUID(t), - Type: connections.Publish, - }, - err: repoerr.ErrCreateEntity, - }, - - { - desc: "add connection with invalid client ID", - connection: clients.Connection{ - ClientID: invalidID, - ChannelID: testsutil.GenerateUUID(t), - DomainID: testsutil.GenerateUUID(t), - Type: connections.Publish, - }, - err: repoerr.ErrMalformedEntity, - }, - { - desc: "add connection with invalid channel ID", - connection: clients.Connection{ - ClientID: client.ID, - ChannelID: invalidID, - DomainID: testsutil.GenerateUUID(t), - Type: connections.Publish, - }, - err: repoerr.ErrMalformedEntity, - }, - { - desc: "add connection with invalid domain ID", - connection: clients.Connection{ - ClientID: client.ID, - ChannelID: testsutil.GenerateUUID(t), - DomainID: invalidID, - Type: connections.Publish, - }, - err: repoerr.ErrMalformedEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := repo.AddConnections(context.Background(), []clients.Connection{tc.connection}) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - }) - } -} - -func TestRemoveConnection(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM connections") - require.Nil(t, err, fmt.Sprintf("clean connections unexpected error: %s", err)) - _, err = db.Exec("DELETE FROM clients") - require.Nil(t, err, fmt.Sprintf("clean clients unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - client := generateClient(t, clients.EnabledStatus, repo) - - validConnection := clients.Connection{ - ClientID: client.ID, - ChannelID: testsutil.GenerateUUID(t), - DomainID: client.Domain, - Type: connections.Publish, - } - - err := repo.AddConnections(context.Background(), []clients.Connection{validConnection}) - require.Nil(t, err, fmt.Sprintf("add connection unexpected error: %s", err)) - - cases := []struct { - desc string - connection clients.Connection - err error - }{ - { - desc: "remove connection successfully", - connection: validConnection, - err: nil, - }, - { - desc: "remove connection with non-existent channel", - connection: clients.Connection{ - ClientID: client.ID, - ChannelID: testsutil.GenerateUUID(t), - DomainID: client.Domain, - Type: connections.Publish, - }, - err: nil, - }, - { - desc: "remove connection with non-existent domain", - connection: clients.Connection{ - ClientID: client.ID, - ChannelID: testsutil.GenerateUUID(t), - DomainID: testsutil.GenerateUUID(t), - Type: connections.Publish, - }, - err: nil, - }, - { - desc: "remove connection with non-existent client", - connection: clients.Connection{ - ClientID: testsutil.GenerateUUID(t), - ChannelID: testsutil.GenerateUUID(t), - DomainID: client.Domain, - Type: connections.Publish, - }, - err: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := repo.RemoveConnections(context.Background(), []clients.Connection{tc.connection}) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - }) - } -} - -func TestDelete(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM clients") - require.Nil(t, err, fmt.Sprintf("clean clients unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - client := generateClient(t, clients.EnabledStatus, repo) - - cases := []struct { - desc string - id string - err error - }{ - { - desc: "delete client successfully", - id: client.ID, - err: nil, - }, - { - desc: "delete client with invalid id", - id: testsutil.GenerateUUID(t), - err: repoerr.ErrNotFound, - }, - { - desc: "delete client with empty id", - id: "", - err: repoerr.ErrNotFound, - }, - } - - for _, tc := range cases { - err := repo.Delete(context.Background(), tc.id) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - } -} - -func TestSetParentGroup(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM clients") - require.Nil(t, err, fmt.Sprintf("clean clients unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - validClient := generateClient(t, clients.EnabledStatus, repo) - - cases := []struct { - desc string - id string - parentGroupID string - err error - }{ - { - desc: "set parent group successfully", - id: validClient.ID, - parentGroupID: testsutil.GenerateUUID(t), - err: nil, - }, - { - desc: "set parent group with invalid ID", - id: invalidID, - parentGroupID: testsutil.GenerateUUID(t), - err: repoerr.ErrNotFound, - }, - { - desc: "set parent group with empty ID", - id: "", - parentGroupID: testsutil.GenerateUUID(t), - err: repoerr.ErrNotFound, - }, - { - desc: "set parent group with invalid parent group ID", - id: validClient.ID, - parentGroupID: invalidID, - err: repoerr.ErrMalformedEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := repo.SetParentGroup(context.Background(), clients.Client{ - ID: tc.id, - ParentGroup: tc.parentGroupID, - }) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - if err == nil { - resp, err := repo.RetrieveByID(context.Background(), tc.id) - require.Nil(t, err, fmt.Sprintf("retrieve client unexpected error: %s", err)) - assert.Equal(t, tc.id, resp.ID, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.id, resp.ID)) - assert.Equal(t, tc.parentGroupID, resp.ParentGroup, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.parentGroupID, resp.ParentGroup)) - } - }) - } -} - -func TestRemoveParentGroup(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM clients") - require.Nil(t, err, fmt.Sprintf("clean clients unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - validClient := generateClient(t, clients.EnabledStatus, repo) - - cases := []struct { - desc string - id string - err error - }{ - { - desc: "remove parent group successfully", - id: validClient.ID, - err: nil, - }, - { - desc: "remove parent group with invalid ID", - id: invalidID, - err: repoerr.ErrNotFound, - }, - { - desc: "remove parent group with empty ID", - id: "", - err: repoerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := repo.RemoveParentGroup(context.Background(), clients.Client{ - ID: tc.id, - }) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - if err == nil { - resp, err := repo.RetrieveByID(context.Background(), tc.id) - require.Nil(t, err, fmt.Sprintf("retrieve client unexpected error: %s", err)) - assert.Equal(t, tc.id, resp.ID, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.id, resp.ID)) - assert.Equal(t, "", resp.ParentGroup, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, "", resp.ParentGroup)) - } - }) - } -} - -func TestClientConnectionsCount(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM connections") - require.Nil(t, err, fmt.Sprintf("clean connections unexpected error: %s", err)) - _, err = db.Exec("DELETE FROM clients") - require.Nil(t, err, fmt.Sprintf("clean clients unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - validClient := generateClient(t, clients.EnabledStatus, repo) - - rConnections := []clients.Connection{} - for i := 0; i < 10; i++ { - connection := clients.Connection{ - ClientID: validClient.ID, - ChannelID: testsutil.GenerateUUID(t), - DomainID: validClient.Domain, - Type: connections.Publish, - } - rConnections = append(rConnections, connection) - } - - err := repo.AddConnections(context.Background(), rConnections) - require.Nil(t, err, fmt.Sprintf("add connection unexpected error: %s", err)) - - cases := []struct { - desc string - clientID string - count uint64 - err error - }{ - { - desc: "get client connections count successfully", - clientID: validClient.ID, - count: 10, - err: nil, - }, - { - desc: "get client connections count with non-existent client", - clientID: testsutil.GenerateUUID(t), - count: 0, - err: nil, - }, - { - desc: "get client connections count with empty client ID", - clientID: "", - count: 0, - err: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - count, err := repo.ClientConnectionsCount(context.Background(), tc.clientID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.count, count, fmt.Sprintf("%s: expected %d got %d\n", tc.desc, tc.count, count)) - }) - } -} - -func TestDoesClientHaveConnections(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM connections") - require.Nil(t, err, fmt.Sprintf("clean connections unexpected error: %s", err)) - _, err = db.Exec("DELETE FROM clients") - require.Nil(t, err, fmt.Sprintf("clean clients unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - validClient := generateClient(t, clients.EnabledStatus, repo) - - validConnection := clients.Connection{ - ClientID: validClient.ID, - ChannelID: testsutil.GenerateUUID(t), - DomainID: validClient.Domain, - Type: connections.Publish, - } - - err := repo.AddConnections(context.Background(), []clients.Connection{validConnection}) - require.Nil(t, err, fmt.Sprintf("add connection unexpected error: %s", err)) - - cases := []struct { - desc string - clientID string - has bool - err error - }{ - { - desc: "check if client has connections successfully", - clientID: validClient.ID, - has: true, - err: nil, - }, - { - desc: "check if client has connections with non-existent channel", - clientID: testsutil.GenerateUUID(t), - has: false, - err: nil, - }, - { - desc: "check if client has connections with empty channel ID", - clientID: "", - has: false, - err: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - has, err := repo.DoesClientHaveConnections(context.Background(), tc.clientID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.has, has, fmt.Sprintf("%s: expected %t got %t\n", tc.desc, tc.has, has)) - }) - } -} - -func TestRemoveClientConnections(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM connections") - require.Nil(t, err, fmt.Sprintf("clean connections unexpected error: %s", err)) - _, err = db.Exec("DELETE FROM clients") - require.Nil(t, err, fmt.Sprintf("clean clients unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - validClient := generateClient(t, clients.EnabledStatus, repo) - - validConnection := clients.Connection{ - ClientID: validClient.ID, - ChannelID: testsutil.GenerateUUID(t), - DomainID: validClient.Domain, - Type: connections.Publish, - } - - err := repo.AddConnections(context.Background(), []clients.Connection{validConnection}) - require.Nil(t, err, fmt.Sprintf("add connection unexpected error: %s", err)) - - cases := []struct { - desc string - clientID string - err error - }{ - { - desc: "remove client connections successfully", - clientID: validConnection.ClientID, - err: nil, - }, - { - desc: "remove client connections with non-existent client", - clientID: testsutil.GenerateUUID(t), - err: nil, - }, - { - desc: "remove client connections with empty client ID", - clientID: "", - err: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := repo.RemoveClientConnections(context.Background(), tc.clientID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - }) - } -} - -func TestRemoveChannelConnections(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM connections") - require.Nil(t, err, fmt.Sprintf("clean connections unexpected error: %s", err)) - _, err = db.Exec("DELETE FROM clients") - require.Nil(t, err, fmt.Sprintf("clean clients unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - validClient := generateClient(t, clients.EnabledStatus, repo) - - validConnection := clients.Connection{ - ClientID: validClient.ID, - ChannelID: testsutil.GenerateUUID(t), - DomainID: validClient.Domain, - Type: connections.Publish, - } - - err := repo.AddConnections(context.Background(), []clients.Connection{validConnection}) - require.Nil(t, err, fmt.Sprintf("add connection unexpected error: %s", err)) - - cases := []struct { - desc string - channelID string - err error - }{ - { - desc: "remove channel connections successfully", - channelID: validConnection.ChannelID, - err: nil, - }, - { - desc: "remove channel connections with non-existent channel", - channelID: testsutil.GenerateUUID(t), - err: nil, - }, - { - desc: "remove channel connections with empty channel ID", - channelID: "", - err: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := repo.RemoveChannelConnections(context.Background(), tc.channelID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - }) - } -} - -func TestRetrieveParentGroupClients(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM clients") - require.Nil(t, err, fmt.Sprintf("clean clients unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - var items []clients.Client - parentID := testsutil.GenerateUUID(t) - baseTime := time.Now().UTC().Truncate(time.Microsecond) - for i := 0; i < 10; i++ { - name := namegen.Generate() - client := clients.Client{ - ID: testsutil.GenerateUUID(t), - Domain: testsutil.GenerateUUID(t), - ParentGroup: parentID, - Name: name, - Metadata: map[string]any{"name": name}, - CreatedAt: baseTime.Add(time.Duration(i) * time.Microsecond), - Status: clients.EnabledStatus, - } - items = append(items, client) - } - - _, err := repo.Save(context.Background(), items...) - require.Nil(t, err, fmt.Sprintf("create client unexpected error: %s", err)) - - cases := []struct { - desc string - parentGroupID string - resp []clients.Client - err error - }{ - { - desc: "retrieve parent group clients successfully", - parentGroupID: parentID, - resp: items[:10], - err: nil, - }, - { - desc: "retrieve parent group clients with non-existent client", - parentGroupID: testsutil.GenerateUUID(t), - resp: []clients.Client(nil), - err: nil, - }, - { - desc: "retrieve parent group clients with empty client ID", - parentGroupID: "", - resp: []clients.Client(nil), - err: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - clients, err := repo.RetrieveParentGroupClients(context.Background(), tc.parentGroupID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - if err == nil { - got := stripClientDetails(clients) - expected := stripClientDetails(tc.resp) - assert.Equal(t, len(tc.resp), len(clients), fmt.Sprintf("%s: expected %d got %d\n", tc.desc, len(tc.resp), len(clients))) - assert.ElementsMatch(t, expected, got, fmt.Sprintf("%s: expected %+v got %+v\n", tc.desc, expected, got)) - } - }) - } -} - -func TestUnsetParentGroupFromClients(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM clients") - require.Nil(t, err, fmt.Sprintf("clean clients unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - var items []clients.Client - parentID := testsutil.GenerateUUID(t) - baseTime := time.Now().UTC().Truncate(time.Microsecond) - for i := 0; i < 10; i++ { - name := namegen.Generate() - client := clients.Client{ - ID: testsutil.GenerateUUID(t), - Domain: testsutil.GenerateUUID(t), - ParentGroup: parentID, - Name: name, - Metadata: map[string]any{"name": name}, - CreatedAt: baseTime.Add(time.Duration(i) * time.Microsecond), - Status: clients.EnabledStatus, - } - items = append(items, client) - } - - _, err := repo.Save(context.Background(), items...) - require.Nil(t, err, fmt.Sprintf("create client unexpected error: %s", err)) - - cases := []struct { - desc string - parentGroupID string - err error - }{ - { - desc: "unset parent group from clients successfully", - parentGroupID: parentID, - err: nil, - }, - { - desc: "unset parent group from clients with non-existent id", - parentGroupID: testsutil.GenerateUUID(t), - err: nil, - }, - { - desc: "unset parent group from clients with empty client ID", - parentGroupID: "", - err: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := repo.UnsetParentGroupFromClient(context.Background(), tc.parentGroupID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - }) - } -} - -func generateClient(t *testing.T, status clients.Status, repo clients.Repository) clients.Client { - client := clients.Client{ - ID: testsutil.GenerateUUID(t), - Name: namegen.Generate(), - Credentials: clients.Credentials{ - Identity: namegen.Generate() + emailSuffix, - Secret: testsutil.GenerateUUID(t), - }, - Tags: namegen.GenerateMultiple(5), - PrivateMetadata: clients.Metadata{ - "name": namegen.Generate(), - }, - Metadata: clients.Metadata{ - "name": namegen.Generate(), - }, - Status: status, - CreatedAt: time.Now().UTC().Truncate(time.Millisecond), - } - _, err := repo.Save(context.Background(), client) - require.Nil(t, err, fmt.Sprintf("add new client: expected nil got %s\n", err)) - - return client -} - -func generateDomain(t *testing.T, userID, memberID string) domains.Domain { - domain := domains.Domain{ - ID: testsutil.GenerateUUID(t), - Route: namegen.Generate(), - Status: domains.EnabledStatus, - CreatedAt: time.Now().UTC().Truncate(time.Microsecond), - CreatedBy: userID, - } - - drepo := dpostgres.NewRepository(ddatabase) - _, err := drepo.SaveDomain(context.Background(), domain) - require.Nil(t, err, fmt.Sprintf("add new domain: expected nil got %s\n", err)) - newRolesProvision := []roles.RoleProvision{ - { - Role: roles.Role{ - ID: testsutil.GenerateUUID(t) + "_" + domain.ID, - Name: "admin", - EntityID: domain.ID, - CreatedAt: validTimestamp, - CreatedBy: userID, - }, - OptionalActions: domainAvailableActions, - OptionalMembers: []string{userID, memberID}, - }, - } - _, err = drepo.AddRoles(context.Background(), newRolesProvision) - require.Nil(t, err, fmt.Sprintf("add new role: expected nil got %s\n", err)) - domain.Roles = []roles.MemberRoleActions{ - { - RoleID: newRolesProvision[0].Role.ID, - RoleName: newRolesProvision[0].Role.Name, - Actions: newRolesProvision[0].OptionalActions, - }, - } - - return domain -} - -func generateGroup(t *testing.T, userID, memberID, domainID string) groups.Group { - group := groups.Group{ - ID: testsutil.GenerateUUID(t), - Name: namegen.Generate(), - Domain: domainID, - Status: groups.EnabledStatus, - CreatedAt: time.Now().UTC().Truncate(time.Microsecond), - } - - grepo := gpostgres.New(gdatabase) - _, err := grepo.Save(context.Background(), group) - require.Nil(t, err, fmt.Sprintf("add new domain: expected nil got %s\n", err)) - newRolesProvision := []roles.RoleProvision{ - { - Role: roles.Role{ - ID: testsutil.GenerateUUID(t) + "_" + group.ID, - Name: "admin", - EntityID: group.ID, - CreatedAt: validTimestamp, - CreatedBy: userID, - }, - OptionalActions: groupAvailableActions, - OptionalMembers: []string{userID, memberID}, - }, - } - _, err = grepo.AddRoles(context.Background(), newRolesProvision) - require.Nil(t, err, fmt.Sprintf("add new role: expected nil got %s\n", err)) - group.Roles = []roles.MemberRoleActions{ - { - RoleID: newRolesProvision[0].Role.ID, - RoleName: newRolesProvision[0].Role.Name, - Actions: newRolesProvision[0].OptionalActions, - }, - } - - return group -} - -func getIDs(clis []clients.Client) []string { - var ids []string - for _, client := range clis { - ids = append(ids, client.ID) - } - - return ids -} - -func stripClientDetails(clients []clients.Client) []clients.Client { - for i := range clients { - clients[i].CreatedAt = validTimestamp - clients[i].Credentials.Secret = "" - clients[i].Actions = []string{} - } - - return clients -} - -func findClients(clis []clients.Client, query string, offset, limit uint64) []clients.Client { - rclients := []clients.Client{} - for _, client := range clis { - if strings.Contains(client.Name, query) { - rclients = append(rclients, client) - } - } - - if offset > uint64(len(rclients)) { - return []clients.Client{} - } - - if limit > uint64(len(rclients)) { - return rclients[offset:] - } - - return rclients[offset:limit] -} - -func verifyClientsOrdering(t *testing.T, clients []clients.Client, order, dir string) { - if order == "" || len(clients) <= 1 { - return - } - - switch order { - case "name": - for i := 1; i < len(clients); i++ { - if dir == ascDir { - assert.LessOrEqual(t, clients[i-1].Name, clients[i].Name) - continue - } - assert.GreaterOrEqual(t, clients[i-1].Name, clients[i].Name) - } - case "identity": - for i := 1; i < len(clients); i++ { - if dir == ascDir { - assert.LessOrEqual(t, clients[i-1].Credentials.Identity, clients[i].Credentials.Identity) - continue - } - assert.GreaterOrEqual(t, clients[i-1].Credentials.Identity, clients[i].Credentials.Identity) - } - case "created_at": - for i := 1; i < len(clients); i++ { - if dir == ascDir { - assert.True(t, !clients[i-1].CreatedAt.After(clients[i].CreatedAt)) - continue - } - assert.True(t, !clients[i-1].CreatedAt.Before(clients[i].CreatedAt)) - } - case "updated_at": - for i := 1; i < len(clients); i++ { - if dir == ascDir { - assert.True(t, !clients[i-1].UpdatedAt.After(clients[i].UpdatedAt)) - continue - } - assert.True(t, !clients[i-1].UpdatedAt.Before(clients[i].UpdatedAt)) - } - } -} diff --git a/clients/postgres/doc.go b/clients/postgres/doc.go deleted file mode 100644 index 6e8346350..000000000 --- a/clients/postgres/doc.go +++ /dev/null @@ -1,5 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package postgres contains the database implementation of clients repository layer. -package postgres diff --git a/clients/postgres/errors.go b/clients/postgres/errors.go deleted file mode 100644 index 95d79c237..000000000 --- a/clients/postgres/errors.go +++ /dev/null @@ -1,24 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package postgres - -import "github.com/absmach/magistrala/pkg/errors" - -var _ errors.Mapper = (*duplicateErrors)(nil) - -type duplicateErrors struct{} - -// GetError maps constraint names to known errors. -func (d duplicateErrors) GetError(constraint string) (error, bool) { - switch constraint { - case "clients_domain_id_secret_key": - return errors.NewRequestError("client key is not available"), true - default: - return nil, false - } -} - -func NewDuplicateErrors() errors.Mapper { - return duplicateErrors{} -} diff --git a/clients/postgres/init.go b/clients/postgres/init.go deleted file mode 100644 index 9cd7f98a0..000000000 --- a/clients/postgres/init.go +++ /dev/null @@ -1,128 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package postgres - -import ( - gpostgres "github.com/absmach/magistrala/groups/postgres" - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - rolesPostgres "github.com/absmach/magistrala/pkg/roles/repo/postgres" - _ "github.com/jackc/pgx/v5/stdlib" // required for SQL access - migrate "github.com/rubenv/sql-migrate" -) - -func Migration() (*migrate.MemoryMigrationSource, error) { - clientsRolesMigration, err := rolesPostgres.Migration(rolesTableNamePrefix, entityTableName, entityIDColumnName) - if err != nil { - return &migrate.MemoryMigrationSource{}, errors.Wrap(repoerr.ErrRoleMigration, err) - } - - clientsMigration := &migrate.MemoryMigrationSource{ - Migrations: []*migrate.Migration{ - { - Id: "clients_01", - // VARCHAR(36) for columns with IDs as UUIDS have a maximum of 36 characters - // STATUS 0 to imply enabled and 1 to imply disabled - Up: []string{ - `CREATE TABLE IF NOT EXISTS clients ( - id VARCHAR(36) PRIMARY KEY, - name VARCHAR(1024), - domain_id VARCHAR(36) NOT NULL, - parent_group_id VARCHAR(36) DEFAULT NULL, - identity VARCHAR(254), - secret VARCHAR(4096) NOT NULL, - tags TEXT[], - metadata JSONB, - created_at TIMESTAMP, - updated_at TIMESTAMP, - updated_by VARCHAR(254), - status SMALLINT NOT NULL DEFAULT 0 CHECK (status >= 0), - UNIQUE (domain_id, secret), - UNIQUE (domain_id, name), - UNIQUE (domain_id, id) - )`, - `CREATE TABLE IF NOT EXISTS connections ( - channel_id VARCHAR(36), - domain_id VARCHAR(36), - client_id VARCHAR(36), - type SMALLINT NOT NULL CHECK (type IN (1, 2)), - FOREIGN KEY (client_id, domain_id) REFERENCES clients (id, domain_id) ON DELETE CASCADE ON UPDATE CASCADE, - PRIMARY KEY (channel_id, domain_id, client_id, type) - )`, - }, - Down: []string{ - `DROP TABLE IF EXISTS clients`, - `DROP TABLE IF EXISTS connections`, - }, - }, - { - Id: "clients_02", - Up: []string{ - `ALTER TABLE clients DROP CONSTRAINT IF EXISTS clients_domain_id_name_key`, - }, - Down: []string{ - `ALTER TABLE clients ADD CONSTRAINT clients_domain_id_name_key UNIQUE (domain_id, name)`, - }, - }, - { - Id: "clients_03", - Up: []string{ - `ALTER TABLE clients ALTER COLUMN created_at TYPE TIMESTAMPTZ;`, - `ALTER TABLE clients ALTER COLUMN updated_at TYPE TIMESTAMPTZ;`, - }, - Down: []string{ - `ALTER TABLE clients ALTER COLUMN created_at TYPE TIMESTAMP;`, - `ALTER TABLE clients ALTER COLUMN updated_at TYPE TIMESTAMP;`, - }, - }, - { - Id: "clients_04", - Up: []string{ - `ALTER TABLE clients ADD COLUMN private_metadata JSONB;`, - }, - Down: []string{ - `ALTER TABLE clients DROP COLUMN private_metadata;`, - }, - }, - { - Id: "clients_05", - Up: []string{ - `UPDATE clients - SET metadata = (COALESCE(metadata, '{}'::jsonb) || COALESCE(metadata->'ui', '{}'::jsonb)) - 'ui' - WHERE metadata ? 'ui' AND jsonb_typeof(metadata->'ui') = 'object'`, - `UPDATE clients - SET private_metadata = (COALESCE(private_metadata, '{}'::jsonb) || COALESCE(private_metadata->'ui', '{}'::jsonb)) - 'ui' - WHERE private_metadata ? 'ui' AND jsonb_typeof(private_metadata->'ui') = 'object'`, - }, - Down: []string{ - `SELECT 1`, - }, - }, - { - Id: "clients_06", - Up: []string{ - `CREATE INDEX IF NOT EXISTS idx_clients_domain_id_status ON clients(domain_id, status);`, - `CREATE INDEX IF NOT EXISTS idx_clients_parent_group_id ON clients(parent_group_id);`, - `CREATE INDEX IF NOT EXISTS idx_connections_client_id ON connections(client_id);`, - }, - Down: []string{ - `DROP INDEX IF EXISTS idx_clients_domain_id_status;`, - `DROP INDEX IF EXISTS idx_clients_parent_group_id;`, - `DROP INDEX IF EXISTS idx_connections_client_id;`, - }, - }, - }, - } - - clientsMigration.Migrations = append(clientsMigration.Migrations, clientsRolesMigration.Migrations...) - - groupsMigration, err := gpostgres.Migration() - if err != nil { - return &migrate.MemoryMigrationSource{}, err - } - - clientsMigration.Migrations = append(clientsMigration.Migrations, groupsMigration.Migrations...) - - return clientsMigration, nil -} diff --git a/clients/postgres/setup_test.go b/clients/postgres/setup_test.go deleted file mode 100644 index fad6d7b1d..000000000 --- a/clients/postgres/setup_test.go +++ /dev/null @@ -1,107 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package postgres_test - -import ( - "database/sql" - "fmt" - "log" - "os" - "testing" - "time" - - cpostgres "github.com/absmach/magistrala/clients/postgres" - pgclient "github.com/absmach/magistrala/pkg/postgres" - "github.com/jmoiron/sqlx" - "github.com/ory/dockertest/v3" - "github.com/ory/dockertest/v3/docker" - "go.opentelemetry.io/otel" -) - -var ( - db *sqlx.DB - database pgclient.Database - ddatabase pgclient.Database - gdatabase pgclient.Database - tracer = otel.Tracer("repo_tests") -) - -func TestMain(m *testing.M) { - pool, err := dockertest.NewPool("") - if err != nil { - log.Fatalf("Could not connect to docker: %s", err) - } - - container, err := pool.RunWithOptions(&dockertest.RunOptions{ - Repository: "postgres", - Tag: "16.2-alpine", - Env: []string{ - "POSTGRES_USER=test", - "POSTGRES_PASSWORD=test", - "POSTGRES_DB=test", - "listen_addresses = '*'", - }, - }, func(config *docker.HostConfig) { - config.AutoRemove = true - config.RestartPolicy = docker.RestartPolicy{Name: "no"} - }) - if err != nil { - log.Fatalf("Could not start container: %s", err) - } - - port := container.GetPort("5432/tcp") - - // exponential backoff-retry, because the application in the container might not be ready to accept connections yet - pool.MaxWait = 120 * time.Second - if err := pool.Retry(func() error { - url := fmt.Sprintf("host=localhost port=%s user=test dbname=test password=test sslmode=disable", port) - db, err := sql.Open("pgx", url) - if err != nil { - return err - } - return db.Ping() - }); err != nil { - log.Fatalf("Could not connect to docker: %s", err) - } - - dbConfig := pgclient.Config{ - Host: "localhost", - Port: port, - User: "test", - Pass: "test", - Name: "test", - SSLMode: "disable", - SSLCert: "", - SSLKey: "", - SSLRootCert: "", - } - - mig, err := cpostgres.Migration() - if err != nil { - log.Fatalf("Could not get DB migrations: %s", err) - } - if db, err = pgclient.Setup(dbConfig, *mig); err != nil { - log.Fatalf("Could not setup test DB connection: %s", err) - } - - if db, err = pgclient.Connect(dbConfig); err != nil { - log.Fatalf("Could not setup test DB connection: %s", err) - } - - database = pgclient.NewDatabase(db, dbConfig, tracer) - - ddatabase = pgclient.NewDatabase(db, dbConfig, tracer) - - gdatabase = pgclient.NewDatabase(db, dbConfig, tracer) - - code := m.Run() - - // Defers will not be run when using os.Exit - db.Close() - if err := pool.Purge(container); err != nil { - log.Fatalf("Could not purge container: %s", err) - } - - os.Exit(code) -} diff --git a/clients/private/doc.go b/clients/private/doc.go deleted file mode 100644 index d5e3ff423..000000000 --- a/clients/private/doc.go +++ /dev/null @@ -1,6 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Private package is a service wrapper around the underlying Repository. -// This is used for internal service communication purpose only. -package private diff --git a/clients/private/mocks/service.go b/clients/private/mocks/service.go deleted file mode 100644 index 90cda170c..000000000 --- a/clients/private/mocks/service.go +++ /dev/null @@ -1,469 +0,0 @@ -// Copyright (c) Abstract Machines - -// SPDX-License-Identifier: Apache-2.0 - -// Code generated by mockery; DO NOT EDIT. -// github.com/vektra/mockery -// template: testify - -package mocks - -import ( - "context" - - "github.com/absmach/magistrala/clients" - mock "github.com/stretchr/testify/mock" -) - -// NewService creates a new instance of Service. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. -// The first argument is typically a *testing.T value. -func NewService(t interface { - mock.TestingT - Cleanup(func()) -}) *Service { - mock := &Service{} - mock.Mock.Test(t) - - t.Cleanup(func() { mock.AssertExpectations(t) }) - - return mock -} - -// Service is an autogenerated mock type for the Service type -type Service struct { - mock.Mock -} - -type Service_Expecter struct { - mock *mock.Mock -} - -func (_m *Service) EXPECT() *Service_Expecter { - return &Service_Expecter{mock: &_m.Mock} -} - -// AddConnections provides a mock function for the type Service -func (_mock *Service) AddConnections(ctx context.Context, conns []clients.Connection) error { - ret := _mock.Called(ctx, conns) - - if len(ret) == 0 { - panic("no return value specified for AddConnections") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, []clients.Connection) error); ok { - r0 = returnFunc(ctx, conns) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_AddConnections_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'AddConnections' -type Service_AddConnections_Call struct { - *mock.Call -} - -// AddConnections is a helper method to define mock.On call -// - ctx context.Context -// - conns []clients.Connection -func (_e *Service_Expecter) AddConnections(ctx interface{}, conns interface{}) *Service_AddConnections_Call { - return &Service_AddConnections_Call{Call: _e.mock.On("AddConnections", ctx, conns)} -} - -func (_c *Service_AddConnections_Call) Run(run func(ctx context.Context, conns []clients.Connection)) *Service_AddConnections_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 []clients.Connection - if args[1] != nil { - arg1 = args[1].([]clients.Connection) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Service_AddConnections_Call) Return(err error) *Service_AddConnections_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_AddConnections_Call) RunAndReturn(run func(ctx context.Context, conns []clients.Connection) error) *Service_AddConnections_Call { - _c.Call.Return(run) - return _c -} - -// Authenticate provides a mock function for the type Service -func (_mock *Service) Authenticate(ctx context.Context, key string) (string, error) { - ret := _mock.Called(ctx, key) - - if len(ret) == 0 { - panic("no return value specified for Authenticate") - } - - var r0 string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) (string, error)); ok { - return returnFunc(ctx, key) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string) string); ok { - r0 = returnFunc(ctx, key) - } else { - r0 = ret.Get(0).(string) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = returnFunc(ctx, key) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_Authenticate_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Authenticate' -type Service_Authenticate_Call struct { - *mock.Call -} - -// Authenticate is a helper method to define mock.On call -// - ctx context.Context -// - key string -func (_e *Service_Expecter) Authenticate(ctx interface{}, key interface{}) *Service_Authenticate_Call { - return &Service_Authenticate_Call{Call: _e.mock.On("Authenticate", ctx, key)} -} - -func (_c *Service_Authenticate_Call) Run(run func(ctx context.Context, key string)) *Service_Authenticate_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Service_Authenticate_Call) Return(s string, err error) *Service_Authenticate_Call { - _c.Call.Return(s, err) - return _c -} - -func (_c *Service_Authenticate_Call) RunAndReturn(run func(ctx context.Context, key string) (string, error)) *Service_Authenticate_Call { - _c.Call.Return(run) - return _c -} - -// RemoveChannelConnections provides a mock function for the type Service -func (_mock *Service) RemoveChannelConnections(ctx context.Context, channelID string) error { - ret := _mock.Called(ctx, channelID) - - if len(ret) == 0 { - panic("no return value specified for RemoveChannelConnections") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) error); ok { - r0 = returnFunc(ctx, channelID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RemoveChannelConnections_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveChannelConnections' -type Service_RemoveChannelConnections_Call struct { - *mock.Call -} - -// RemoveChannelConnections is a helper method to define mock.On call -// - ctx context.Context -// - channelID string -func (_e *Service_Expecter) RemoveChannelConnections(ctx interface{}, channelID interface{}) *Service_RemoveChannelConnections_Call { - return &Service_RemoveChannelConnections_Call{Call: _e.mock.On("RemoveChannelConnections", ctx, channelID)} -} - -func (_c *Service_RemoveChannelConnections_Call) Run(run func(ctx context.Context, channelID string)) *Service_RemoveChannelConnections_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Service_RemoveChannelConnections_Call) Return(err error) *Service_RemoveChannelConnections_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RemoveChannelConnections_Call) RunAndReturn(run func(ctx context.Context, channelID string) error) *Service_RemoveChannelConnections_Call { - _c.Call.Return(run) - return _c -} - -// RemoveConnections provides a mock function for the type Service -func (_mock *Service) RemoveConnections(ctx context.Context, conns []clients.Connection) error { - ret := _mock.Called(ctx, conns) - - if len(ret) == 0 { - panic("no return value specified for RemoveConnections") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, []clients.Connection) error); ok { - r0 = returnFunc(ctx, conns) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RemoveConnections_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveConnections' -type Service_RemoveConnections_Call struct { - *mock.Call -} - -// RemoveConnections is a helper method to define mock.On call -// - ctx context.Context -// - conns []clients.Connection -func (_e *Service_Expecter) RemoveConnections(ctx interface{}, conns interface{}) *Service_RemoveConnections_Call { - return &Service_RemoveConnections_Call{Call: _e.mock.On("RemoveConnections", ctx, conns)} -} - -func (_c *Service_RemoveConnections_Call) Run(run func(ctx context.Context, conns []clients.Connection)) *Service_RemoveConnections_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 []clients.Connection - if args[1] != nil { - arg1 = args[1].([]clients.Connection) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Service_RemoveConnections_Call) Return(err error) *Service_RemoveConnections_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RemoveConnections_Call) RunAndReturn(run func(ctx context.Context, conns []clients.Connection) error) *Service_RemoveConnections_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveById provides a mock function for the type Service -func (_mock *Service) RetrieveById(ctx context.Context, id string) (clients.Client, error) { - ret := _mock.Called(ctx, id) - - if len(ret) == 0 { - panic("no return value specified for RetrieveById") - } - - var r0 clients.Client - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) (clients.Client, error)); ok { - return returnFunc(ctx, id) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string) clients.Client); ok { - r0 = returnFunc(ctx, id) - } else { - r0 = ret.Get(0).(clients.Client) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = returnFunc(ctx, id) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RetrieveById_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveById' -type Service_RetrieveById_Call struct { - *mock.Call -} - -// RetrieveById is a helper method to define mock.On call -// - ctx context.Context -// - id string -func (_e *Service_Expecter) RetrieveById(ctx interface{}, id interface{}) *Service_RetrieveById_Call { - return &Service_RetrieveById_Call{Call: _e.mock.On("RetrieveById", ctx, id)} -} - -func (_c *Service_RetrieveById_Call) Run(run func(ctx context.Context, id string)) *Service_RetrieveById_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Service_RetrieveById_Call) Return(client clients.Client, err error) *Service_RetrieveById_Call { - _c.Call.Return(client, err) - return _c -} - -func (_c *Service_RetrieveById_Call) RunAndReturn(run func(ctx context.Context, id string) (clients.Client, error)) *Service_RetrieveById_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveByIds provides a mock function for the type Service -func (_mock *Service) RetrieveByIds(ctx context.Context, ids []string) (clients.ClientsPage, error) { - ret := _mock.Called(ctx, ids) - - if len(ret) == 0 { - panic("no return value specified for RetrieveByIds") - } - - var r0 clients.ClientsPage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, []string) (clients.ClientsPage, error)); ok { - return returnFunc(ctx, ids) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, []string) clients.ClientsPage); ok { - r0 = returnFunc(ctx, ids) - } else { - r0 = ret.Get(0).(clients.ClientsPage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, []string) error); ok { - r1 = returnFunc(ctx, ids) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RetrieveByIds_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveByIds' -type Service_RetrieveByIds_Call struct { - *mock.Call -} - -// RetrieveByIds is a helper method to define mock.On call -// - ctx context.Context -// - ids []string -func (_e *Service_Expecter) RetrieveByIds(ctx interface{}, ids interface{}) *Service_RetrieveByIds_Call { - return &Service_RetrieveByIds_Call{Call: _e.mock.On("RetrieveByIds", ctx, ids)} -} - -func (_c *Service_RetrieveByIds_Call) Run(run func(ctx context.Context, ids []string)) *Service_RetrieveByIds_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 []string - if args[1] != nil { - arg1 = args[1].([]string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Service_RetrieveByIds_Call) Return(clientsPage clients.ClientsPage, err error) *Service_RetrieveByIds_Call { - _c.Call.Return(clientsPage, err) - return _c -} - -func (_c *Service_RetrieveByIds_Call) RunAndReturn(run func(ctx context.Context, ids []string) (clients.ClientsPage, error)) *Service_RetrieveByIds_Call { - _c.Call.Return(run) - return _c -} - -// UnsetParentGroupFromClient provides a mock function for the type Service -func (_mock *Service) UnsetParentGroupFromClient(ctx context.Context, parentGroupID string) error { - ret := _mock.Called(ctx, parentGroupID) - - if len(ret) == 0 { - panic("no return value specified for UnsetParentGroupFromClient") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) error); ok { - r0 = returnFunc(ctx, parentGroupID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_UnsetParentGroupFromClient_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UnsetParentGroupFromClient' -type Service_UnsetParentGroupFromClient_Call struct { - *mock.Call -} - -// UnsetParentGroupFromClient is a helper method to define mock.On call -// - ctx context.Context -// - parentGroupID string -func (_e *Service_Expecter) UnsetParentGroupFromClient(ctx interface{}, parentGroupID interface{}) *Service_UnsetParentGroupFromClient_Call { - return &Service_UnsetParentGroupFromClient_Call{Call: _e.mock.On("UnsetParentGroupFromClient", ctx, parentGroupID)} -} - -func (_c *Service_UnsetParentGroupFromClient_Call) Run(run func(ctx context.Context, parentGroupID string)) *Service_UnsetParentGroupFromClient_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Service_UnsetParentGroupFromClient_Call) Return(err error) *Service_UnsetParentGroupFromClient_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_UnsetParentGroupFromClient_Call) RunAndReturn(run func(ctx context.Context, parentGroupID string) error) *Service_UnsetParentGroupFromClient_Call { - _c.Call.Return(run) - return _c -} diff --git a/clients/private/service.go b/clients/private/service.go deleted file mode 100644 index ee43bbc3e..000000000 --- a/clients/private/service.go +++ /dev/null @@ -1,125 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package private - -import ( - "context" - - "github.com/absmach/magistrala/clients" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/pkg/policies" -) - -type Service interface { - // Authenticate returns client ID for given client key. - Authenticate(ctx context.Context, key string) (string, error) - - RetrieveById(ctx context.Context, id string) (clients.Client, error) - - RetrieveByIds(ctx context.Context, ids []string) (clients.ClientsPage, error) - - AddConnections(ctx context.Context, conns []clients.Connection) error - - RemoveConnections(ctx context.Context, conns []clients.Connection) error - - RemoveChannelConnections(ctx context.Context, channelID string) error - - UnsetParentGroupFromClient(ctx context.Context, parentGroupID string) error -} - -var _ Service = (*service)(nil) - -func New(repo clients.Repository, cache clients.Cache, evaluator policies.Evaluator, policy policies.Service) Service { - return service{ - repo: repo, - cache: cache, - evaluator: evaluator, - policy: policy, - } -} - -type service struct { - repo clients.Repository - cache clients.Cache - evaluator policies.Evaluator - policy policies.Service -} - -func (svc service) Authenticate(ctx context.Context, token string) (string, error) { - id, err := svc.cache.ID(ctx, token) - if err == nil { - return id, nil - } - prefix, id, key, err := authn.AuthUnpack(token) - if err != nil && err != authn.ErrNotEncoded { - return "", err - } - client, err := svc.repo.RetrieveBySecret(ctx, key, id, prefix) - if err != nil { - return "", errors.Wrap(svcerr.ErrAuthorization, err) - } - if err := svc.cache.Save(ctx, token, client.ID); err != nil { - return "", errors.Wrap(svcerr.ErrAuthorization, err) - } - - return client.ID, nil -} - -func (svc service) RetrieveById(ctx context.Context, ids string) (clients.Client, error) { - return svc.repo.RetrieveByID(ctx, ids) -} - -func (svc service) RetrieveByIds(ctx context.Context, ids []string) (clients.ClientsPage, error) { - return svc.repo.RetrieveByIds(ctx, ids) -} - -func (svc service) AddConnections(ctx context.Context, conns []clients.Connection) (err error) { - return svc.repo.AddConnections(ctx, conns) -} - -func (svc service) RemoveConnections(ctx context.Context, conns []clients.Connection) (err error) { - return svc.repo.RemoveConnections(ctx, conns) -} - -func (svc service) RemoveChannelConnections(ctx context.Context, channelID string) error { - return svc.repo.RemoveChannelConnections(ctx, channelID) -} - -func (svc service) UnsetParentGroupFromClient(ctx context.Context, parentGroupID string) (retErr error) { - clients, err := svc.repo.RetrieveParentGroupClients(ctx, parentGroupID) - if err != nil { - return errors.Wrap(svcerr.ErrViewEntity, err) - } - - if len(clients) > 0 { - prs := []policies.Policy{} - for _, client := range clients { - prs = append(prs, policies.Policy{ - SubjectType: policies.GroupType, - Subject: client.ParentGroup, - Relation: policies.ParentGroupRelation, - ObjectType: policies.ClientType, - Object: client.ID, - }) - } - - if err := svc.policy.DeletePolicies(ctx, prs); err != nil { - return errors.Wrap(svcerr.ErrDeletePolicies, err) - } - defer func() { - if retErr != nil { - if errRollback := svc.policy.AddPolicies(ctx, prs); err != nil { - retErr = errors.Wrap(retErr, errors.Wrap(errors.ErrRollbackTx, errRollback)) - } - } - }() - - if err := svc.repo.UnsetParentGroupFromClient(ctx, parentGroupID); err != nil { - return errors.Wrap(svcerr.ErrRemoveEntity, err) - } - } - return nil -} diff --git a/clients/roles.go b/clients/roles.go deleted file mode 100644 index cd5b879cf..000000000 --- a/clients/roles.go +++ /dev/null @@ -1,71 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package clients - -import ( - "encoding/json" - "strings" - - apiutil "github.com/absmach/magistrala/api/http/util" -) - -// Role represents Client role. -type Role uint8 - -// Possible Client role values. -const ( - UserRole Role = iota - AdminRole - - // AllRole is used for querying purposes to list clients irrespective - // of their role - both admin and user. It is never stored in the - // database as the actual Client role and should always be the largest - // value in this enumeration. - AllRole -) - -// String representation of the possible role values. -const ( - Admin = "admin" - User = "user" -) - -// String converts client role to string literal. -func (cs Role) String() string { - switch cs { - case AdminRole: - return Admin - case UserRole: - return User - case AllRole: - return All - default: - return Unknown - } -} - -// ToRole converts string value to a valid Client role. -func ToRole(status string) (Role, error) { - switch status { - case "", User: - return UserRole, nil - case Admin: - return AdminRole, nil - case All: - return AllRole, nil - default: - return Role(0), apiutil.ErrInvalidRole - } -} - -func (r Role) MarshalJSON() ([]byte, error) { - return json.Marshal(r.String()) -} - -func (r *Role) UnmarshalJSON(data []byte) error { - str := strings.Trim(string(data), "\"") - val, err := ToRole(str) - *r = val - return err -} diff --git a/clients/roles_test.go b/clients/roles_test.go deleted file mode 100644 index 9a3e369e6..000000000 --- a/clients/roles_test.go +++ /dev/null @@ -1,175 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package clients_test - -import ( - "testing" - - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/clients" - "github.com/stretchr/testify/assert" -) - -func TestRoleString(t *testing.T) { - cases := []struct { - desc string - role clients.Role - expected string - }{ - { - desc: "User", - role: clients.UserRole, - expected: "user", - }, - { - desc: "Admin", - role: clients.AdminRole, - expected: "admin", - }, - { - desc: "All", - role: clients.AllRole, - expected: "all", - }, - { - desc: "Unknown", - role: clients.Role(100), - expected: "unknown", - }, - } - - for _, c := range cases { - t.Run(c.desc, func(t *testing.T) { - got := c.role.String() - assert.Equal(t, c.expected, got, "String() = %v, expected %v", got, c.expected) - }) - } -} - -func TestToRole(t *testing.T) { - cases := []struct { - desc string - role string - expected clients.Role - err error - }{ - { - desc: "User", - role: "user", - expected: clients.UserRole, - err: nil, - }, - { - desc: "Admin", - role: "admin", - expected: clients.AdminRole, - err: nil, - }, - { - desc: "All", - role: "all", - expected: clients.AllRole, - err: nil, - }, - { - desc: "Unknown", - role: "unknown", - expected: clients.Role(0), - err: apiutil.ErrInvalidRole, - }, - } - - for _, c := range cases { - t.Run(c.desc, func(t *testing.T) { - got, err := clients.ToRole(c.role) - assert.Equal(t, c.err, err, "ToRole() error = %v, expected %v", err, c.err) - assert.Equal(t, c.expected, got, "ToRole() = %v, expected %v", got, c.expected) - }) - } -} - -func TestRoleMarshalJSON(t *testing.T) { - cases := []struct { - desc string - expected []byte - role clients.Role - err error - }{ - { - desc: "User", - expected: []byte(`"user"`), - role: clients.UserRole, - err: nil, - }, - { - desc: "Admin", - expected: []byte(`"admin"`), - role: clients.AdminRole, - err: nil, - }, - { - desc: "All", - expected: []byte(`"all"`), - role: clients.AllRole, - err: nil, - }, - { - desc: "Unknown", - expected: []byte(`"unknown"`), - role: clients.Role(100), - err: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - got, err := tc.role.MarshalJSON() - assert.Equal(t, tc.err, err, "MarshalJSON() error = %v, expected %v", err, tc.err) - assert.Equal(t, tc.expected, got, "MarshalJSON() = %v, expected %v", got, tc.expected) - }) - } -} - -func TestRoleUnmarshalJSON(t *testing.T) { - cases := []struct { - desc string - expected clients.Role - role []byte - err error - }{ - { - desc: "User", - expected: clients.UserRole, - role: []byte(`"user"`), - err: nil, - }, - { - desc: "Admin", - expected: clients.AdminRole, - role: []byte(`"admin"`), - err: nil, - }, - { - desc: "All", - expected: clients.AllRole, - role: []byte(`"all"`), - err: nil, - }, - { - desc: "Unknown", - expected: clients.Role(0), - role: []byte(`"unknown"`), - err: apiutil.ErrInvalidRole, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - var r clients.Role - err := r.UnmarshalJSON(tc.role) - assert.Equal(t, tc.err, err, "UnmarshalJSON() error = %v, expected %v", err, tc.err) - assert.Equal(t, tc.expected, r, "UnmarshalJSON() = %v, expected %v", r, tc.expected) - }) - } -} diff --git a/clients/service.go b/clients/service.go deleted file mode 100644 index 160a1024a..000000000 --- a/clients/service.go +++ /dev/null @@ -1,406 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 -package clients - -import ( - "context" - "time" - - mg "github.com/absmach/magistrala" - grpcChannelsV1 "github.com/absmach/magistrala/api/grpc/channels/v1" - grpcCommonV1 "github.com/absmach/magistrala/api/grpc/common/v1" - grpcGroupsV1 "github.com/absmach/magistrala/api/grpc/groups/v1" - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/pkg/policies" - "github.com/absmach/magistrala/pkg/roles" -) - -var ( - errRollbackRepo = errors.New("failed to rollback repo") - errSetParentGroup = errors.NewRequestError("client already have parent") - errSetSameParentGroup = errors.NewRequestError("client already assigned to the parent group") - errParentGroupDomainID = errors.NewRequestError("parent group has invalid domain id") - errParentGroupDisabled = errors.NewRequestError("parent group is not enabled") -) -var _ Service = (*service)(nil) - -type service struct { - repo Repository - policy policies.Service - channels grpcChannelsV1.ChannelsServiceClient - groups grpcGroupsV1.GroupsServiceClient - cache Cache - idProvider mg.IDProvider - roles.ProvisionManageService -} - -// NewService returns a new Clients service implementation. -func NewService(repo Repository, policy policies.Service, cache Cache, channels grpcChannelsV1.ChannelsServiceClient, groups grpcGroupsV1.GroupsServiceClient, idProvider mg.IDProvider, sIDProvider mg.IDProvider, availableActions []roles.Action, builtInRoles map[roles.BuiltInRoleName][]roles.Action) (Service, error) { - rpms, err := roles.NewProvisionManageService(policies.ClientType, repo, policy, sIDProvider, availableActions, builtInRoles) - if err != nil { - return service{}, err - } - return service{ - repo: repo, - policy: policy, - channels: channels, - groups: groups, - cache: cache, - idProvider: idProvider, - ProvisionManageService: rpms, - }, nil -} - -func (svc service) CreateClients(ctx context.Context, session authn.Session, cls ...Client) (retClients []Client, retRps []roles.RoleProvision, retErr error) { - var clients []Client - for _, c := range cls { - if c.ID == "" { - clientID, err := svc.idProvider.ID() - if err != nil { - return []Client{}, []roles.RoleProvision{}, errors.Wrap(svcerr.ErrIssueProviderID, err) - } - c.ID = clientID - } - if c.Credentials.Secret == "" { - key, err := svc.idProvider.ID() - if err != nil { - return []Client{}, []roles.RoleProvision{}, errors.Wrap(svcerr.ErrIssueProviderID, err) - } - c.Credentials.Secret = key - } - if c.Status != DisabledStatus && c.Status != EnabledStatus { - return []Client{}, []roles.RoleProvision{}, svcerr.ErrInvalidStatus - } - c.Domain = session.DomainID - c.CreatedAt = time.Now().UTC() - clients = append(clients, c) - } - - newClients, err := svc.repo.Save(ctx, clients...) - if err != nil { - return []Client{}, []roles.RoleProvision{}, errors.Wrap(svcerr.ErrCreateEntity, err) - } - newClientIDs := []string{} - for _, newClient := range newClients { - newClientIDs = append(newClientIDs, newClient.ID) - } - - defer func() { - if retErr != nil { - if errRollBack := svc.repo.Delete(ctx, newClientIDs...); errRollBack != nil { - retErr = errors.Wrap(retErr, errors.Wrap(errRollbackRepo, errRollBack)) - } - } - }() - - newBuiltInRoleMembers := map[roles.BuiltInRoleName][]roles.Member{ - BuiltInRoleAdmin: {roles.Member(session.UserID)}, - } - - optionalPolicies := []policies.Policy{} - - for _, newClientID := range newClientIDs { - optionalPolicies = append(optionalPolicies, - policies.Policy{ - Domain: session.DomainID, - SubjectType: policies.DomainType, - Subject: session.DomainID, - Relation: policies.DomainRelation, - ObjectType: policies.ClientType, - Object: newClientID, - }, - ) - } - - rp, err := svc.AddNewEntitiesRoles(ctx, session.DomainID, session.UserID, newClientIDs, optionalPolicies, newBuiltInRoleMembers) - if err != nil { - return []Client{}, []roles.RoleProvision{}, errors.Wrap(svcerr.ErrAddPolicies, err) - } - - return newClients, rp, nil -} - -func (svc service) View(ctx context.Context, session authn.Session, id string, withRoles bool) (Client, error) { - var client Client - var err error - switch withRoles { - case true: - client, err = svc.repo.RetrieveByIDWithRoles(ctx, id, session.UserID) - default: - client, err = svc.repo.RetrieveByID(ctx, id) - } - if err != nil { - return Client{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - return client, nil -} - -func (svc service) ListClients(ctx context.Context, session authn.Session, pm Page) (ClientsPage, error) { - switch session.SuperAdmin { - case true: - pm.Domain = session.DomainID - cp, err := svc.repo.RetrieveAll(ctx, pm) - if err != nil { - return ClientsPage{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - return cp, nil - default: - cp, err := svc.repo.RetrieveUserClients(ctx, session.DomainID, session.UserID, pm) - if err != nil { - return ClientsPage{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - return cp, nil - } -} - -func (svc service) ListUserClients(ctx context.Context, session authn.Session, userID string, pm Page) (ClientsPage, error) { - cp, err := svc.repo.RetrieveUserClients(ctx, session.DomainID, userID, pm) - if err != nil { - return ClientsPage{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - return cp, nil -} - -func (svc service) Update(ctx context.Context, session authn.Session, cli Client) (Client, error) { - client := Client{ - ID: cli.ID, - Name: cli.Name, - Metadata: cli.Metadata, - PrivateMetadata: cli.PrivateMetadata, - UpdatedAt: time.Now().UTC(), - UpdatedBy: session.UserID, - } - client, err := svc.repo.Update(ctx, client) - if err != nil { - return Client{}, errors.Wrap(svcerr.ErrUpdateEntity, err) - } - return client, nil -} - -func (svc service) UpdateTags(ctx context.Context, session authn.Session, cli Client) (Client, error) { - client := Client{ - ID: cli.ID, - Tags: cli.Tags, - UpdatedAt: time.Now().UTC(), - UpdatedBy: session.UserID, - } - client, err := svc.repo.UpdateTags(ctx, client) - if err != nil { - return Client{}, errors.Wrap(svcerr.ErrUpdateEntity, err) - } - return client, nil -} - -func (svc service) UpdateSecret(ctx context.Context, session authn.Session, id, key string) (Client, error) { - client := Client{ - ID: id, - Credentials: Credentials{ - Secret: key, - }, - UpdatedAt: time.Now().UTC(), - UpdatedBy: session.UserID, - Status: EnabledStatus, - } - client, err := svc.repo.UpdateSecret(ctx, client) - if err != nil { - return Client{}, errors.Wrap(svcerr.ErrUpdateEntity, err) - } - if err := svc.cache.Remove(ctx, client.ID); err != nil { - return client, errors.Wrap(svcerr.ErrRemoveEntity, err) - } - return client, nil -} - -func (svc service) Enable(ctx context.Context, session authn.Session, id string) (Client, error) { - client := Client{ - ID: id, - Status: EnabledStatus, - UpdatedAt: time.Now().UTC(), - } - client, err := svc.changeClientStatus(ctx, session, client) - if err != nil { - return Client{}, errors.Wrap(ErrEnableClient, err) - } - - return client, nil -} - -func (svc service) Disable(ctx context.Context, session authn.Session, id string) (Client, error) { - client := Client{ - ID: id, - Status: DisabledStatus, - UpdatedAt: time.Now().UTC(), - } - client, err := svc.changeClientStatus(ctx, session, client) - if err != nil { - return Client{}, errors.Wrap(ErrDisableClient, err) - } - - if err := svc.cache.Remove(ctx, client.ID); err != nil { - return client, errors.Wrap(svcerr.ErrRemoveEntity, err) - } - - return client, nil -} - -func (svc service) SetParentGroup(ctx context.Context, session authn.Session, parentGroupID string, id string) (retErr error) { - cli, err := svc.repo.RetrieveByID(ctx, id) - if err != nil { - return errors.Wrap(svcerr.ErrUpdateEntity, err) - } - switch cli.ParentGroup { - case parentGroupID: - return errors.Wrap(svcerr.ErrConflict, errSetSameParentGroup) - case "": - // No action needed, proceed to next code after switch - default: - return errors.Wrap(svcerr.ErrConflict, errSetParentGroup) - } - - resp, err := svc.groups.RetrieveEntity(ctx, &grpcCommonV1.RetrieveEntityReq{Id: parentGroupID}) - if err != nil { - return errors.Wrap(svcerr.ErrUpdateEntity, err) - } - if resp.GetEntity().GetDomainId() != session.DomainID { - return errors.Wrap(svcerr.ErrUpdateEntity, errParentGroupDomainID) - } - if resp.GetEntity().GetStatus() != uint32(EnabledStatus) { - return errors.Wrap(svcerr.ErrUpdateEntity, errParentGroupDisabled) - } - - var pols []policies.Policy - - pols = append(pols, policies.Policy{ - Domain: session.DomainID, - SubjectType: policies.GroupType, - Subject: parentGroupID, - Relation: policies.ParentGroupRelation, - ObjectType: policies.ClientType, - Object: id, - }) - - if err := svc.policy.AddPolicies(ctx, pols); err != nil { - return errors.Wrap(svcerr.ErrAddPolicies, err) - } - defer func() { - if retErr != nil { - if errRollback := svc.policy.DeletePolicies(ctx, pols); errRollback != nil { - retErr = errors.Wrap(retErr, errors.Wrap(apiutil.ErrRollbackTx, errRollback)) - } - } - }() - cli = Client{ID: id, ParentGroup: parentGroupID, UpdatedBy: session.UserID, UpdatedAt: time.Now().UTC()} - - if err := svc.repo.SetParentGroup(ctx, cli); err != nil { - return errors.Wrap(svcerr.ErrUpdateEntity, err) - } - return nil -} - -func (svc service) RemoveParentGroup(ctx context.Context, session authn.Session, id string) (retErr error) { - cli, err := svc.repo.RetrieveByID(ctx, id) - if err != nil { - return errors.Wrap(svcerr.ErrViewEntity, err) - } - - if cli.ParentGroup != "" { - var pols []policies.Policy - pols = append(pols, policies.Policy{ - Domain: session.DomainID, - SubjectType: policies.GroupType, - Subject: cli.ParentGroup, - Relation: policies.ParentGroupRelation, - ObjectType: policies.ClientType, - Object: id, - }) - - if err := svc.policy.DeletePolicies(ctx, pols); err != nil { - return errors.Wrap(svcerr.ErrDeletePolicies, err) - } - defer func() { - if retErr != nil { - if errRollback := svc.policy.AddPolicies(ctx, pols); errRollback != nil { - retErr = errors.Wrap(retErr, errors.Wrap(apiutil.ErrRollbackTx, errRollback)) - } - } - }() - - cli := Client{ID: id, UpdatedBy: session.UserID, UpdatedAt: time.Now().UTC()} - - if err := svc.repo.RemoveParentGroup(ctx, cli); err != nil { - return errors.Wrap(svcerr.ErrUpdateEntity, err) - } - } - return nil -} - -func (svc service) Delete(ctx context.Context, session authn.Session, id string) error { - ok, err := svc.repo.DoesClientHaveConnections(ctx, id) - if err != nil { - return errors.Wrap(svcerr.ErrRemoveEntity, err) - } - if ok { - if _, err := svc.channels.RemoveClientConnections(ctx, &grpcChannelsV1.RemoveClientConnectionsReq{ClientId: id}); err != nil { - return errors.Wrap(svcerr.ErrRemoveEntity, err) - } - } - - if _, err := svc.repo.ChangeStatus(ctx, Client{ID: id, Status: DeletedStatus}); err != nil { - return errors.Wrap(svcerr.ErrRemoveEntity, err) - } - - if err := svc.cache.Remove(ctx, id); err != nil { - return errors.Wrap(svcerr.ErrRemoveEntity, err) - } - - filterDeletePolicies := []policies.Policy{ - { - SubjectType: policies.ClientType, - Subject: id, - }, - { - ObjectType: policies.ClientType, - Object: id, - }, - } - deletePolicies := []policies.Policy{ - { - SubjectType: policies.DomainType, - Subject: session.DomainID, - Relation: policies.DomainRelation, - ObjectType: policies.ClientType, - Object: id, - }, - } - - if err := svc.RemoveEntitiesRoles(ctx, session.DomainID, session.DomainUserID, []string{id}, filterDeletePolicies, deletePolicies); err != nil { - return errors.Wrap(svcerr.ErrDeletePolicies, err) - } - - if err := svc.repo.Delete(ctx, id); err != nil { - return errors.Wrap(svcerr.ErrRemoveEntity, err) - } - - return nil -} - -func (svc service) changeClientStatus(ctx context.Context, session authn.Session, client Client) (Client, error) { - dbClient, err := svc.repo.RetrieveByID(ctx, client.ID) - if err != nil { - return Client{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - if dbClient.Status == client.Status { - return Client{}, svcerr.ErrStatusAlreadyAssigned - } - - client.UpdatedBy = session.UserID - - client, err = svc.repo.ChangeStatus(ctx, client) - if err != nil { - return Client{}, errors.Wrap(svcerr.ErrUpdateEntity, err) - } - return client, nil -} diff --git a/clients/service_test.go b/clients/service_test.go deleted file mode 100644 index 8dd8c4edf..000000000 --- a/clients/service_test.go +++ /dev/null @@ -1,1285 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package clients_test - -import ( - "context" - "fmt" - "testing" - - grpcChannelsV1 "github.com/absmach/magistrala/api/grpc/channels/v1" - grpcCommonV1 "github.com/absmach/magistrala/api/grpc/common/v1" - apiutil "github.com/absmach/magistrala/api/http/util" - chmocks "github.com/absmach/magistrala/channels/mocks" - "github.com/absmach/magistrala/clients" - climocks "github.com/absmach/magistrala/clients/mocks" - gpmocks "github.com/absmach/magistrala/groups/mocks" - "github.com/absmach/magistrala/internal/testsutil" - smqauthn "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - policysvc "github.com/absmach/magistrala/pkg/policies" - policymocks "github.com/absmach/magistrala/pkg/policies/mocks" - "github.com/absmach/magistrala/pkg/roles" - "github.com/absmach/magistrala/pkg/uuid" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" -) - -var ( - secret = "strongsecret" - validMetadata = clients.Metadata{"role": "client"} - ID = "6e5e10b3-d4df-4758-b426-4929d55ad740" - client = clients.Client{ - ID: ID, - Name: "clientname", - Tags: []string{"tag1", "tag2"}, - Credentials: clients.Credentials{Identity: "clientidentity", Secret: secret}, - PrivateMetadata: validMetadata, - Metadata: validMetadata, - Status: clients.EnabledStatus, - } - clientWithRoles = clients.Client{ - ID: ID, - Name: "clientname", - Tags: []string{"tag1", "tag2"}, - Credentials: clients.Credentials{Identity: "clientidentity", Secret: secret}, - PrivateMetadata: validMetadata, - Metadata: validMetadata, - Status: clients.EnabledStatus, - Roles: []roles.MemberRoleActions{ - { - RoleID: "test_role_id", - RoleName: "test_role_name", - }, - }, - } - validToken = "token" - validID = "d4ebb847-5d0e-4e46-bdd9-b6aceaaa3a22" - wrongID = testsutil.GenerateUUID(&testing.T{}) -) - -var ( - pService *policymocks.Service - cache *climocks.Cache - repo *climocks.Repository - chgRPCClient *chmocks.ChannelsServiceClient - gpgRPCClient *gpmocks.GroupsServiceClient -) - -func newService() clients.Service { - pService = new(policymocks.Service) - cache = new(climocks.Cache) - idProvider := uuid.NewMock() - sidProvider := uuid.NewMock() - repo = new(climocks.Repository) - chgRPCClient = new(chmocks.ChannelsServiceClient) - gpgRPCClient = new(gpmocks.GroupsServiceClient) - availableActions := []roles.Action{} - builtInRoles := map[roles.BuiltInRoleName][]roles.Action{ - clients.BuiltInRoleAdmin: availableActions, - } - tsv, _ := clients.NewService(repo, pService, cache, chgRPCClient, gpgRPCClient, idProvider, sidProvider, availableActions, builtInRoles) - return tsv -} - -func TestCreateClients(t *testing.T) { - svc := newService() - - cases := []struct { - desc string - client clients.Client - token string - addPolicyErr error - deletePolicyErr error - saveErr error - addRoleErr error - deleteErr error - err error - }{ - { - desc: "create a new client successfully", - client: client, - token: validToken, - err: nil, - }, - { - desc: "create an existing client", - client: client, - token: validToken, - saveErr: repoerr.ErrConflict, - err: repoerr.ErrConflict, - }, - { - desc: "create a new client without secret", - client: clients.Client{ - Name: "clientWithoutSecret", - Credentials: clients.Credentials{ - Identity: "newclientwithoutsecret@example.com", - }, - Status: clients.EnabledStatus, - }, - token: validToken, - err: nil, - }, - { - desc: "create a new client without identity", - client: clients.Client{ - Name: "clientWithoutIdentity", - Credentials: clients.Credentials{ - Identity: "newclientwithoutsecret@example.com", - }, - Status: clients.EnabledStatus, - }, - token: validToken, - err: nil, - }, - { - desc: "create a new enabled client with name", - client: clients.Client{ - Name: "clientWithName", - Credentials: clients.Credentials{ - Identity: "newclientwithname@example.com", - Secret: secret, - }, - Status: clients.EnabledStatus, - }, - token: validToken, - err: nil, - }, - - { - desc: "create a new disabled client with name", - client: clients.Client{ - Name: "clientWithName", - Credentials: clients.Credentials{ - Identity: "newclientwithname@example.com", - Secret: secret, - }, - }, - token: validToken, - err: nil, - }, - { - desc: "create a new enabled client with tags", - client: clients.Client{ - Tags: []string{"tag1", "tag2"}, - Credentials: clients.Credentials{ - Identity: "newclientwithtags@example.com", - Secret: secret, - }, - Status: clients.EnabledStatus, - }, - token: validToken, - err: nil, - }, - { - desc: "create a new disabled client with tags", - client: clients.Client{ - Tags: []string{"tag1", "tag2"}, - Credentials: clients.Credentials{ - Identity: "newclientwithtags@example.com", - Secret: secret, - }, - Status: clients.DisabledStatus, - }, - token: validToken, - err: nil, - }, - { - desc: "create a new enabled client with private metadata", - client: clients.Client{ - Credentials: clients.Credentials{ - Identity: "newclientwithmetadata@example.com", - Secret: secret, - }, - PrivateMetadata: validMetadata, - Status: clients.EnabledStatus, - }, - token: validToken, - err: nil, - }, - { - desc: "create a new enabled client with metadata", - client: clients.Client{ - Credentials: clients.Credentials{ - Identity: "newclientwithmetadata@example.com", - Secret: secret, - }, - Metadata: validMetadata, - Status: clients.EnabledStatus, - }, - token: validToken, - err: nil, - }, - { - desc: "create a new disabled client with private metadata", - client: clients.Client{ - Credentials: clients.Credentials{ - Identity: "newclientwithmetadata@example.com", - Secret: secret, - }, - PrivateMetadata: validMetadata, - }, - token: validToken, - err: nil, - }, - { - desc: "create a new disabled client", - client: clients.Client{ - Credentials: clients.Credentials{ - Identity: "newclientwithvalidstatus@example.com", - Secret: secret, - }, - }, - token: validToken, - err: nil, - }, - { - desc: "create a new client with valid disabled status", - client: clients.Client{ - Credentials: clients.Credentials{ - Identity: "newclientwithvalidstatus@example.com", - Secret: secret, - }, - Status: clients.DisabledStatus, - }, - token: validToken, - err: nil, - }, - { - desc: "create a new client with all fields", - client: clients.Client{ - Name: "newclientwithallfields", - Tags: []string{"tag1", "tag2"}, - Credentials: clients.Credentials{ - Identity: "newclientwithallfields@example.com", - Secret: secret, - }, - PrivateMetadata: clients.Metadata{ - "name": "newclientwithallfields", - }, - Metadata: clients.Metadata{ - "name": "newclientwithallfields", - }, - Status: clients.EnabledStatus, - }, - token: validToken, - err: nil, - }, - { - desc: "create a new client with invalid status", - client: clients.Client{ - Credentials: clients.Credentials{ - Identity: "newclientwithinvalidstatus@example.com", - Secret: secret, - }, - Status: clients.AllStatus, - }, - token: validToken, - err: svcerr.ErrInvalidStatus, - }, - { - desc: "create a new client with failed add policies response", - client: clients.Client{ - Credentials: clients.Credentials{ - Identity: "newclientwithfailedpolicy@example.com", - Secret: secret, - }, - Status: clients.EnabledStatus, - }, - token: validToken, - addPolicyErr: svcerr.ErrInvalidPolicy, - err: svcerr.ErrInvalidPolicy, - }, - { - desc: "create a new client with failed delete policies response", - client: clients.Client{ - Credentials: clients.Credentials{ - Identity: "newclientwithfailedpolicy@example.com", - Secret: secret, - }, - Status: clients.EnabledStatus, - }, - token: validToken, - saveErr: repoerr.ErrConflict, - deletePolicyErr: svcerr.ErrInvalidPolicy, - err: repoerr.ErrConflict, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := repo.On("Save", context.Background(), mock.Anything).Return([]clients.Client{tc.client}, tc.saveErr) - policyCall := pService.On("AddPolicies", context.Background(), mock.Anything).Return(tc.addPolicyErr) - policyCall1 := pService.On("DeletePolicies", context.Background(), mock.Anything).Return(tc.deletePolicyErr) - repoCall1 := repo.On("AddRoles", context.Background(), mock.Anything).Return([]roles.RoleProvision{}, tc.addRoleErr) - repoCall2 := repo.On("Delete", context.Background(), mock.Anything).Return(tc.deleteErr) - expected, _, err := svc.CreateClients(context.Background(), smqauthn.Session{}, tc.client) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - if err == nil { - tc.client.ID = expected[0].ID - tc.client.CreatedAt = expected[0].CreatedAt - tc.client.UpdatedAt = expected[0].UpdatedAt - tc.client.Credentials.Secret = expected[0].Credentials.Secret - tc.client.Domain = expected[0].Domain - tc.client.UpdatedBy = expected[0].UpdatedBy - assert.Equal(t, tc.client, expected[0], fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.client, expected[0])) - } - repoCall.Unset() - policyCall.Unset() - policyCall1.Unset() - repoCall1.Unset() - repoCall2.Unset() - }) - } -} - -func TestViewClient(t *testing.T) { - svc := newService() - - cases := []struct { - desc string - clientID string - withRoles bool - response clients.Client - retrieveErr error - err error - }{ - { - desc: "view client successfully", - response: client, - withRoles: false, - clientID: client.ID, - err: nil, - }, - { - desc: "view client successfully with roles", - response: clientWithRoles, - withRoles: true, - clientID: clientWithRoles.ID, - err: nil, - }, - { - desc: "view client with an invalid token", - response: clients.Client{}, - withRoles: false, - clientID: "", - err: svcerr.ErrAuthorization, - }, - { - desc: "view client with valid token and invalid client id", - response: clients.Client{}, - withRoles: false, - clientID: wrongID, - retrieveErr: svcerr.ErrNotFound, - err: svcerr.ErrNotFound, - }, - { - desc: "view client with an invalid token and invalid client id", - response: clients.Client{}, - withRoles: false, - clientID: wrongID, - err: svcerr.ErrAuthorization, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := repo.On("RetrieveByID", context.Background(), tc.clientID).Return(tc.response, tc.err) - repoCall1 := repo.On("RetrieveByIDWithRoles", context.Background(), tc.clientID, mock.Anything).Return(tc.response, tc.err) - rClient, err := svc.View(context.Background(), smqauthn.Session{}, tc.clientID, tc.withRoles) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - - switch tc.withRoles { - case true: - assert.NotEmpty(t, rClient.Roles) - ok := repo.AssertCalled(t, "RetrieveByIDWithRoles", context.Background(), tc.clientID, mock.Anything) - assert.True(t, ok, fmt.Sprintf("RetrieveByIDWithRoles was not called on %s", tc.desc)) - default: - ok := repo.AssertCalled(t, "RetrieveByID", context.Background(), tc.clientID) - assert.True(t, ok, fmt.Sprintf("RetrieveByID was not called on %s", tc.desc)) - } - - assert.Equal(t, tc.response, rClient, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.response, rClient)) - repoCall.Unset() - repoCall1.Unset() - }) - } -} - -func TestListClients(t *testing.T) { - svc := newService() - - adminID := testsutil.GenerateUUID(t) - domainID := testsutil.GenerateUUID(t) - nonAdminID := testsutil.GenerateUUID(t) - - cases := []struct { - desc string - userKind string - session smqauthn.Session - page clients.Page - listObjectsResponse policysvc.PolicyPage - retrieveAllResponse clients.ClientsPage - listPermissionsResponse policysvc.Permissions - response clients.ClientsPage - id string - size uint64 - listObjectsErr error - retrieveAllErr error - listPermissionsErr error - err error - }{ - { - desc: "list all clients successfully as non admin", - userKind: "non-admin", - session: smqauthn.Session{UserID: nonAdminID, DomainID: domainID, SuperAdmin: false}, - id: nonAdminID, - page: clients.Page{ - Offset: 0, - Limit: 100, - }, - listObjectsResponse: policysvc.PolicyPage{Policies: []string{client.ID, client.ID}}, - retrieveAllResponse: clients.ClientsPage{ - Page: clients.Page{ - Total: 2, - Offset: 0, - Limit: 100, - }, - Clients: []clients.Client{client, client}, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 2, - Offset: 0, - Limit: 100, - }, - Clients: []clients.Client{client, client}, - }, - err: nil, - }, - { - desc: "list all clients as non admin with failed to retrieve all", - userKind: "non-admin", - session: smqauthn.Session{UserID: nonAdminID, DomainID: domainID, SuperAdmin: false}, - id: nonAdminID, - page: clients.Page{ - Offset: 0, - Limit: 100, - }, - listObjectsResponse: policysvc.PolicyPage{Policies: []string{client.ID, client.ID}}, - retrieveAllResponse: clients.ClientsPage{}, - response: clients.ClientsPage{}, - retrieveAllErr: repoerr.ErrNotFound, - err: svcerr.ErrNotFound, - }, - { - desc: "list all clients as non admin with failed super admin", - userKind: "non-admin", - session: smqauthn.Session{UserID: nonAdminID, DomainID: domainID, SuperAdmin: false}, - id: nonAdminID, - page: clients.Page{ - Offset: 0, - Limit: 100, - }, - response: clients.ClientsPage{}, - listObjectsResponse: policysvc.PolicyPage{}, - err: nil, - }, - { - desc: "list all clients as non admin with failed to list objects", - userKind: "non-admin", - id: nonAdminID, - page: clients.Page{ - Offset: 0, - Limit: 100, - }, - retrieveAllErr: repoerr.ErrNotFound, - response: clients.ClientsPage{}, - listObjectsResponse: policysvc.PolicyPage{}, - listObjectsErr: svcerr.ErrNotFound, - err: svcerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - retrieveAllCall := repo.On("RetrieveAll", mock.Anything, mock.Anything).Return(tc.retrieveAllResponse, tc.retrieveAllErr) - retrieveUserClientsCall := repo.On("RetrieveUserClients", mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(tc.retrieveAllResponse, tc.retrieveAllErr) - page, err := svc.ListClients(context.Background(), tc.session, tc.page) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.response, page, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.response, page)) - retrieveAllCall.Unset() - retrieveUserClientsCall.Unset() - }) - } - - cases2 := []struct { - desc string - userKind string - session smqauthn.Session - page clients.Page - listObjectsResponse policysvc.PolicyPage - retrieveAllResponse clients.ClientsPage - listPermissionsResponse policysvc.Permissions - response clients.ClientsPage - id string - size uint64 - listObjectsErr error - retrieveAllErr error - listPermissionsErr error - err error - }{ - { - desc: "list all clients as admin successfully", - userKind: "admin", - id: adminID, - session: smqauthn.Session{UserID: adminID, DomainID: domainID, SuperAdmin: true}, - page: clients.Page{ - Offset: 0, - Limit: 100, - Domain: domainID, - }, - listObjectsResponse: policysvc.PolicyPage{Policies: []string{client.ID, client.ID}}, - retrieveAllResponse: clients.ClientsPage{ - Page: clients.Page{ - Total: 2, - Offset: 0, - Limit: 100, - }, - Clients: []clients.Client{client, client}, - }, - response: clients.ClientsPage{ - Page: clients.Page{ - Total: 2, - Offset: 0, - Limit: 100, - }, - Clients: []clients.Client{client, client}, - }, - err: nil, - }, - { - desc: "list all clients as admin with failed to retrieve all", - userKind: "admin", - id: adminID, - session: smqauthn.Session{UserID: adminID, DomainID: domainID, SuperAdmin: true}, - page: clients.Page{ - Offset: 0, - Limit: 100, - Domain: domainID, - }, - listObjectsResponse: policysvc.PolicyPage{}, - retrieveAllResponse: clients.ClientsPage{}, - retrieveAllErr: repoerr.ErrNotFound, - err: svcerr.ErrNotFound, - }, - { - desc: "list all clients as admin with failed to list clients", - userKind: "admin", - id: adminID, - session: smqauthn.Session{UserID: adminID, DomainID: domainID, SuperAdmin: true}, - page: clients.Page{ - Offset: 0, - Limit: 100, - Domain: domainID, - }, - retrieveAllResponse: clients.ClientsPage{}, - retrieveAllErr: repoerr.ErrNotFound, - err: svcerr.ErrNotFound, - }, - } - - for _, tc := range cases2 { - t.Run(tc.desc, func(t *testing.T) { - retrieveAllCall := repo.On("RetrieveAll", mock.Anything, mock.Anything).Return(tc.retrieveAllResponse, tc.retrieveAllErr) - page, err := svc.ListClients(context.Background(), tc.session, tc.page) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.response, page, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.response, page)) - retrieveAllCall.Unset() - }) - } -} - -func TestUpdateClient(t *testing.T) { - svc := newService() - - client1 := client - client2 := client - client1.Name = "Updated client" - client2.PrivateMetadata = clients.Metadata{"role": "test"} - client2.Metadata = clients.Metadata{"role": "test"} - - cases := []struct { - desc string - client clients.Client - session smqauthn.Session - updateResponse clients.Client - updateErr error - err error - }{ - { - desc: "update client name successfully", - client: client1, - session: smqauthn.Session{UserID: validID}, - updateResponse: client1, - err: nil, - }, - { - desc: "update client metadata with valid token", - client: client2, - updateResponse: client2, - session: smqauthn.Session{UserID: validID}, - err: nil, - }, - { - desc: "update client with failed to update repo", - client: client1, - updateResponse: clients.Client{}, - session: smqauthn.Session{UserID: validID}, - updateErr: repoerr.ErrMalformedEntity, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall1 := repo.On("Update", context.Background(), mock.Anything).Return(tc.updateResponse, tc.updateErr) - updatedClient, err := svc.Update(context.Background(), tc.session, tc.client) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.updateResponse, updatedClient, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.updateResponse, updatedClient)) - repoCall1.Unset() - }) - } -} - -func TestUpdateTags(t *testing.T) { - svc := newService() - - client.Tags = []string{"updated"} - - cases := []struct { - desc string - client clients.Client - session smqauthn.Session - updateResponse clients.Client - updateErr error - err error - }{ - { - desc: "update client tags successfully", - client: client, - session: smqauthn.Session{UserID: validID}, - updateResponse: client, - err: nil, - }, - { - desc: "update client tags with failed to update repo", - client: client, - updateResponse: clients.Client{}, - session: smqauthn.Session{UserID: validID}, - updateErr: repoerr.ErrMalformedEntity, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall1 := repo.On("UpdateTags", context.Background(), mock.Anything).Return(tc.updateResponse, tc.updateErr) - updatedClient, err := svc.UpdateTags(context.Background(), tc.session, tc.client) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.updateResponse, updatedClient, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.updateResponse, updatedClient)) - repoCall1.Unset() - }) - } -} - -func TestUpdateSecret(t *testing.T) { - svc := newService() - - cases := []struct { - desc string - client clients.Client - newSecret string - updateSecretResponse clients.Client - session smqauthn.Session - updateErr error - removeErr error - err error - }{ - { - desc: "update client secret successfully", - client: client, - newSecret: "newSecret", - session: smqauthn.Session{UserID: validID}, - updateSecretResponse: clients.Client{ - ID: client.ID, - Credentials: clients.Credentials{ - Identity: client.Credentials.Identity, - Secret: "newSecret", - }, - }, - err: nil, - }, - { - desc: "update client secret with failed to update repo", - client: client, - newSecret: "newSecret", - session: smqauthn.Session{UserID: validID}, - updateSecretResponse: clients.Client{}, - updateErr: repoerr.ErrMalformedEntity, - err: svcerr.ErrUpdateEntity, - }, - { - desc: "update client secret with failed to remove cache", - client: client, - newSecret: "newSecret", - session: smqauthn.Session{UserID: validID}, - updateSecretResponse: clients.Client{ - ID: client.ID, - Credentials: clients.Credentials{ - Identity: client.Credentials.Identity, - Secret: "newSecret", - }, - }, - removeErr: repoerr.ErrRemoveEntity, - err: svcerr.ErrRemoveEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := repo.On("UpdateSecret", context.Background(), mock.Anything).Return(tc.updateSecretResponse, tc.updateErr) - var cacheCall *mock.Call - if tc.updateErr == nil { - cacheCall = cache.On("Remove", context.Background(), tc.updateSecretResponse.ID).Return(tc.removeErr) - } - updatedClient, err := svc.UpdateSecret(context.Background(), tc.session, tc.client.ID, tc.newSecret) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.updateSecretResponse, updatedClient, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.updateSecretResponse, updatedClient)) - repoCall.Unset() - if cacheCall != nil { - cacheCall.Unset() - } - }) - } -} - -func TestEnable(t *testing.T) { - svc := newService() - - enabledClient1 := clients.Client{ID: ID, Credentials: clients.Credentials{Identity: "client1@example.com", Secret: "password"}, Status: clients.EnabledStatus} - disabledClient1 := clients.Client{ID: ID, Credentials: clients.Credentials{Identity: "client3@example.com", Secret: "password"}, Status: clients.DisabledStatus} - endisabledClient1 := disabledClient1 - endisabledClient1.Status = clients.EnabledStatus - - cases := []struct { - desc string - id string - session smqauthn.Session - client clients.Client - changeStatusResponse clients.Client - retrieveByIDResponse clients.Client - changeStatusErr error - retrieveIDErr error - err error - }{ - { - desc: "enable disabled client", - id: disabledClient1.ID, - session: smqauthn.Session{UserID: validID}, - client: disabledClient1, - changeStatusResponse: endisabledClient1, - retrieveByIDResponse: disabledClient1, - err: nil, - }, - { - desc: "enable disabled client with failed to update repo", - id: disabledClient1.ID, - session: smqauthn.Session{UserID: validID}, - client: disabledClient1, - changeStatusResponse: clients.Client{}, - retrieveByIDResponse: disabledClient1, - changeStatusErr: repoerr.ErrMalformedEntity, - err: svcerr.ErrUpdateEntity, - }, - { - desc: "enable enabled client", - id: enabledClient1.ID, - session: smqauthn.Session{UserID: validID}, - client: enabledClient1, - changeStatusResponse: enabledClient1, - retrieveByIDResponse: enabledClient1, - changeStatusErr: svcerr.ErrStatusAlreadyAssigned, - err: svcerr.ErrStatusAlreadyAssigned, - }, - { - desc: "enable non-existing client", - id: wrongID, - session: smqauthn.Session{UserID: validID}, - client: clients.Client{}, - changeStatusResponse: clients.Client{}, - retrieveByIDResponse: clients.Client{}, - retrieveIDErr: repoerr.ErrNotFound, - err: repoerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := repo.On("RetrieveByID", context.Background(), mock.Anything).Return(tc.retrieveByIDResponse, tc.retrieveIDErr) - repoCall1 := repo.On("ChangeStatus", context.Background(), mock.Anything).Return(tc.changeStatusResponse, tc.changeStatusErr) - _, err := svc.Enable(context.Background(), tc.session, tc.id) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - repoCall.Unset() - repoCall1.Unset() - }) - } -} - -func TestDisable(t *testing.T) { - svc := newService() - - enabledClient1 := clients.Client{ID: ID, Credentials: clients.Credentials{Identity: "client1@example.com", Secret: "password"}, Status: clients.EnabledStatus} - disabledClient1 := clients.Client{ID: ID, Credentials: clients.Credentials{Identity: "client3@example.com", Secret: "password"}, Status: clients.DisabledStatus} - disenabledClient1 := enabledClient1 - disenabledClient1.Status = clients.DisabledStatus - - cases := []struct { - desc string - id string - session smqauthn.Session - client clients.Client - changeStatusResponse clients.Client - retrieveByIDResponse clients.Client - changeStatusErr error - retrieveIDErr error - removeErr error - err error - }{ - { - desc: "disable enabled client", - id: enabledClient1.ID, - session: smqauthn.Session{UserID: validID}, - client: enabledClient1, - changeStatusResponse: disenabledClient1, - retrieveByIDResponse: enabledClient1, - err: nil, - }, - { - desc: "disable client with failed to update repo", - id: enabledClient1.ID, - session: smqauthn.Session{UserID: validID}, - client: enabledClient1, - changeStatusResponse: clients.Client{}, - retrieveByIDResponse: enabledClient1, - changeStatusErr: repoerr.ErrMalformedEntity, - err: svcerr.ErrUpdateEntity, - }, - { - desc: "disable disabled client", - id: disabledClient1.ID, - session: smqauthn.Session{UserID: validID}, - client: disabledClient1, - changeStatusResponse: clients.Client{}, - retrieveByIDResponse: disabledClient1, - changeStatusErr: svcerr.ErrStatusAlreadyAssigned, - err: svcerr.ErrStatusAlreadyAssigned, - }, - { - desc: "disable non-existing client", - id: wrongID, - client: clients.Client{}, - session: smqauthn.Session{UserID: validID}, - changeStatusResponse: clients.Client{}, - retrieveByIDResponse: clients.Client{}, - retrieveIDErr: repoerr.ErrNotFound, - err: repoerr.ErrNotFound, - }, - { - desc: "disable client with failed to remove from cache", - id: enabledClient1.ID, - session: smqauthn.Session{UserID: validID}, - client: disabledClient1, - changeStatusResponse: disenabledClient1, - retrieveByIDResponse: enabledClient1, - removeErr: svcerr.ErrRemoveEntity, - err: svcerr.ErrRemoveEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := repo.On("RetrieveByID", context.Background(), mock.Anything).Return(tc.retrieveByIDResponse, tc.retrieveIDErr) - repoCall1 := repo.On("ChangeStatus", context.Background(), mock.Anything).Return(tc.changeStatusResponse, tc.changeStatusErr) - repoCall2 := cache.On("Remove", mock.Anything, mock.Anything).Return(tc.removeErr) - _, err := svc.Disable(context.Background(), tc.session, tc.id) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - repoCall.Unset() - repoCall1.Unset() - repoCall2.Unset() - }) - } -} - -func TestDelete(t *testing.T) { - svc := newService() - - client := clients.Client{ - ID: testsutil.GenerateUUID(t), - } - - cases := []struct { - desc string - clientID string - checkConnectionsRes bool - checkConnectionsErr error - removeConnectionsErr error - changeStatusErr error - deletePoliciesErr error - removeErr error - deleteErr error - err error - }{ - { - desc: "Delete client without connections successfully", - clientID: client.ID, - err: nil, - }, - { - desc: "Delete client with connections", - clientID: client.ID, - checkConnectionsRes: true, - err: nil, - }, - { - desc: "Delete client with failed to check connections", - clientID: client.ID, - checkConnectionsErr: svcerr.ErrRemoveEntity, - err: svcerr.ErrRemoveEntity, - }, - { - desc: "Delete client with failed to remove connections", - clientID: client.ID, - checkConnectionsRes: true, - removeConnectionsErr: svcerr.ErrRemoveEntity, - err: svcerr.ErrRemoveEntity, - }, - { - desc: "Delete cliet with failed to remove from cache", - clientID: client.ID, - removeErr: svcerr.ErrRemoveEntity, - err: svcerr.ErrRemoveEntity, - }, - { - desc: "Delete client with failed to change status", - clientID: client.ID, - changeStatusErr: svcerr.ErrNotFound, - err: svcerr.ErrRemoveEntity, - }, - { - desc: "Delete client with failed to delete policies", - clientID: client.ID, - deletePoliciesErr: svcerr.ErrNotFound, - err: svcerr.ErrDeletePolicies, - }, - { - desc: "Delete client with failed to delete", - clientID: client.ID, - deleteErr: svcerr.ErrNotFound, - err: svcerr.ErrRemoveEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := repo.On("DoesClientHaveConnections", context.Background(), mock.Anything).Return(tc.checkConnectionsRes, tc.checkConnectionsErr) - channelsCall := chgRPCClient.On("RemoveClientConnections", context.Background(), &grpcChannelsV1.RemoveClientConnectionsReq{ClientId: tc.clientID}).Return(&grpcChannelsV1.RemoveClientConnectionsRes{}, tc.removeConnectionsErr) - repoCall1 := cache.On("Remove", mock.Anything, tc.clientID).Return(tc.removeErr) - repoCall2 := repo.On("ChangeStatus", context.Background(), clients.Client{ID: tc.clientID, Status: clients.DeletedStatus}).Return(client, tc.changeStatusErr) - repoCall3 := repo.On("RetrieveEntitiesRolesActionsMembers", context.Background(), []string{tc.clientID}).Return([]roles.EntityActionRole{}, []roles.EntityMemberRole{}, nil) - policyCall1 := pService.On("DeletePolicies", context.Background(), mock.Anything).Return(tc.deletePoliciesErr) - policyCall2 := pService.On("DeletePolicyFilter", context.Background(), mock.Anything).Return(tc.deletePoliciesErr) - repoCall4 := repo.On("Delete", context.Background(), []string{tc.clientID}).Return(tc.deleteErr) - err := svc.Delete(context.Background(), smqauthn.Session{}, tc.clientID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - repoCall.Unset() - repoCall1.Unset() - policyCall1.Unset() - repoCall2.Unset() - channelsCall.Unset() - repoCall3.Unset() - repoCall4.Unset() - policyCall2.Unset() - }) - } -} - -func TestSetParentGroup(t *testing.T) { - svc := newService() - - parentedClient := client - parentedClient.ParentGroup = validID - - cparentedClient := client - cparentedClient.ParentGroup = testsutil.GenerateUUID(t) - - cases := []struct { - desc string - clientID string - parentGroupID string - session smqauthn.Session - retrieveByIDResp clients.Client - retrieveByIDErr error - retrieveEntityResp *grpcCommonV1.RetrieveEntityRes - retrieveEntityErr error - addPoliciesErr error - deletePoliciesErr error - setParentGroupErr error - err error - }{ - { - desc: "set parent group successfully", - clientID: client.ID, - parentGroupID: testsutil.GenerateUUID(t), - session: smqauthn.Session{UserID: validID, DomainID: validID, DomainUserID: validID + "_" + validID}, - retrieveByIDResp: client, - retrieveEntityResp: &grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Id: testsutil.GenerateUUID(t), - DomainId: validID, - Status: uint32(clients.EnabledStatus), - }, - }, - err: nil, - }, - { - desc: "set parent group with failed to retrieve client", - clientID: client.ID, - parentGroupID: testsutil.GenerateUUID(t), - session: smqauthn.Session{UserID: validID, DomainID: validID, DomainUserID: validID + "_" + validID}, - retrieveByIDResp: clients.Client{}, - retrieveByIDErr: svcerr.ErrNotFound, - err: svcerr.ErrUpdateEntity, - }, - { - desc: "set parent group with parent already set", - clientID: parentedClient.ID, - parentGroupID: validID, - session: smqauthn.Session{UserID: validID, DomainID: validID, DomainUserID: validID + "_" + validID}, - retrieveByIDResp: parentedClient, - err: svcerr.ErrConflict, - }, - { - desc: "set parent group of client with existing parent group", - clientID: cparentedClient.ID, - parentGroupID: testsutil.GenerateUUID(t), - session: smqauthn.Session{UserID: validID, DomainID: validID, DomainUserID: validID + "_" + validID}, - retrieveByIDResp: cparentedClient, - err: svcerr.ErrConflict, - }, - { - desc: "set parent group with failed to retrieve entity", - clientID: client.ID, - parentGroupID: testsutil.GenerateUUID(t), - session: smqauthn.Session{UserID: validID, DomainID: validID, DomainUserID: validID + "_" + validID}, - retrieveByIDResp: client, - retrieveEntityErr: svcerr.ErrAuthorization, - err: svcerr.ErrUpdateEntity, - }, - { - desc: "set parent group with parent group from different domain", - clientID: client.ID, - parentGroupID: testsutil.GenerateUUID(t), - session: smqauthn.Session{UserID: validID, DomainID: validID, DomainUserID: validID + "_" + validID}, - retrieveByIDResp: client, - retrieveEntityResp: &grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Id: testsutil.GenerateUUID(t), - DomainId: testsutil.GenerateUUID(t), - Status: uint32(clients.EnabledStatus), - }, - }, - err: svcerr.ErrUpdateEntity, - }, - { - desc: "set parent group with disabled parent group", - clientID: client.ID, - parentGroupID: testsutil.GenerateUUID(t), - session: smqauthn.Session{UserID: validID, DomainID: validID, DomainUserID: validID + "_" + validID}, - retrieveByIDResp: client, - retrieveEntityResp: &grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Id: testsutil.GenerateUUID(t), - DomainId: validID, - Status: uint32(clients.DisabledStatus), - }, - }, - err: svcerr.ErrUpdateEntity, - }, - { - desc: "set parent group with failed to add policies", - clientID: client.ID, - parentGroupID: testsutil.GenerateUUID(t), - session: smqauthn.Session{UserID: validID, DomainID: validID, DomainUserID: validID + "_" + validID}, - retrieveByIDResp: client, - retrieveEntityResp: &grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Id: testsutil.GenerateUUID(t), - DomainId: validID, - Status: uint32(clients.EnabledStatus), - }, - }, - addPoliciesErr: svcerr.ErrUpdateEntity, - err: svcerr.ErrAddPolicies, - }, - { - desc: "set parent group with failed to set parent group", - clientID: client.ID, - parentGroupID: testsutil.GenerateUUID(t), - session: smqauthn.Session{UserID: validID, DomainID: validID, DomainUserID: validID + "_" + validID}, - retrieveByIDResp: client, - retrieveEntityResp: &grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Id: testsutil.GenerateUUID(t), - DomainId: validID, - Status: uint32(clients.EnabledStatus), - }, - }, - setParentGroupErr: svcerr.ErrUpdateEntity, - err: svcerr.ErrUpdateEntity, - }, - { - desc: "set parent group with failed to set parent group and failed rollback", - clientID: client.ID, - parentGroupID: testsutil.GenerateUUID(t), - session: smqauthn.Session{UserID: validID, DomainID: validID, DomainUserID: validID + "_" + validID}, - retrieveByIDResp: client, - retrieveEntityResp: &grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Id: testsutil.GenerateUUID(t), - DomainId: validID, - Status: uint32(clients.EnabledStatus), - }, - }, - setParentGroupErr: svcerr.ErrUpdateEntity, - deletePoliciesErr: svcerr.ErrAuthorization, - err: apiutil.ErrRollbackTx, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - pols := []policysvc.Policy{ - { - Domain: tc.session.DomainID, - SubjectType: policysvc.GroupType, - Subject: tc.parentGroupID, - Relation: policysvc.ParentGroupRelation, - ObjectType: policysvc.ClientType, - Object: tc.clientID, - }, - } - repoCall := repo.On("RetrieveByID", context.Background(), tc.clientID).Return(tc.retrieveByIDResp, tc.retrieveByIDErr) - groupsCall := gpgRPCClient.On("RetrieveEntity", context.Background(), &grpcCommonV1.RetrieveEntityReq{Id: tc.parentGroupID}).Return(tc.retrieveEntityResp, tc.retrieveEntityErr) - policyCall := pService.On("AddPolicies", context.Background(), pols).Return(tc.addPoliciesErr) - policyCall1 := pService.On("DeletePolicies", context.Background(), pols).Return(tc.deletePoliciesErr) - repoCall2 := repo.On("SetParentGroup", context.Background(), mock.Anything).Return(tc.setParentGroupErr) - err := svc.SetParentGroup(context.Background(), tc.session, tc.parentGroupID, tc.clientID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - repoCall.Unset() - groupsCall.Unset() - policyCall.Unset() - repoCall2.Unset() - policyCall1.Unset() - }) - } -} - -func TestRemoveParentGroup(t *testing.T) { - svc := newService() - - parentedGroup := client - parentedGroup.ParentGroup = validID - - cases := []struct { - desc string - clientID string - session smqauthn.Session - retrieveByIDResp clients.Client - retrieveByIDErr error - deletePoliciesErr error - addPoliciesErr error - removeParentGroupErr error - err error - }{ - { - desc: "remove parent group successfully", - clientID: parentedGroup.ID, - session: smqauthn.Session{UserID: validID, DomainID: validID, DomainUserID: validID + "_" + validID}, - retrieveByIDResp: parentedGroup, - err: nil, - }, - { - desc: "remove parent group with failed to retrieve client", - clientID: parentedGroup.ID, - session: smqauthn.Session{UserID: validID, DomainID: validID, DomainUserID: validID + "_" + validID}, - retrieveByIDResp: clients.Client{}, - retrieveByIDErr: svcerr.ErrNotFound, - err: svcerr.ErrViewEntity, - }, - { - desc: "remove parent group with failed to delete policies", - clientID: parentedGroup.ID, - session: smqauthn.Session{UserID: validID, DomainID: validID, DomainUserID: validID + "_" + validID}, - retrieveByIDResp: parentedGroup, - deletePoliciesErr: svcerr.ErrAuthorization, - err: svcerr.ErrDeletePolicies, - }, - { - desc: "remove parent group with failed to remove parent group", - clientID: parentedGroup.ID, - session: smqauthn.Session{UserID: validID, DomainID: validID, DomainUserID: validID + "_" + validID}, - retrieveByIDResp: parentedGroup, - removeParentGroupErr: svcerr.ErrUpdateEntity, - err: svcerr.ErrUpdateEntity, - }, - { - desc: "remove parent group with failed to remove parent group and failed to add policies", - clientID: parentedGroup.ID, - session: smqauthn.Session{UserID: validID, DomainID: validID, DomainUserID: validID + "_" + validID}, - retrieveByIDResp: parentedGroup, - removeParentGroupErr: svcerr.ErrUpdateEntity, - addPoliciesErr: svcerr.ErrUpdateEntity, - err: apiutil.ErrRollbackTx, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - pols := []policysvc.Policy{ - { - Domain: tc.session.DomainID, - SubjectType: policysvc.GroupType, - Subject: tc.retrieveByIDResp.ParentGroup, - Relation: policysvc.ParentGroupRelation, - ObjectType: policysvc.ClientType, - Object: tc.clientID, - }, - } - repoCall := repo.On("RetrieveByID", context.Background(), tc.clientID).Return(tc.retrieveByIDResp, tc.retrieveByIDErr) - policyCall := pService.On("DeletePolicies", context.Background(), pols).Return(tc.deletePoliciesErr) - policyCall1 := pService.On("AddPolicies", context.Background(), pols).Return(tc.addPoliciesErr) - repoCall2 := repo.On("RemoveParentGroup", context.Background(), mock.Anything).Return(tc.removeParentGroupErr) - err := svc.RemoveParentGroup(context.Background(), tc.session, tc.clientID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - repoCall.Unset() - policyCall.Unset() - repoCall2.Unset() - policyCall1.Unset() - }) - } -} diff --git a/clients/standalone/doc.go b/clients/standalone/doc.go deleted file mode 100644 index 9ad956baf..000000000 --- a/clients/standalone/doc.go +++ /dev/null @@ -1,9 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package standalone contains implementation for auth service in -// single-user scenario. Running with a single user provides -// Clients as a standalone service with one admin user who -// manages all the Clients and Channels and does not -// require connection to Auth service. -package standalone diff --git a/clients/standalone/standalone.go b/clients/standalone/standalone.go deleted file mode 100644 index 5d14ffba7..000000000 --- a/clients/standalone/standalone.go +++ /dev/null @@ -1,4 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package standalone diff --git a/clients/status.go b/clients/status.go deleted file mode 100644 index b66bf2e3a..000000000 --- a/clients/status.go +++ /dev/null @@ -1,94 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package clients - -import ( - "encoding/json" - "strings" - - svcerr "github.com/absmach/magistrala/pkg/errors/service" -) - -// Status represents Client status. -type Status uint8 - -// Possible Client status values. -const ( - // EnabledStatus represents enabled Client. - EnabledStatus Status = iota - // DisabledStatus represents disabled Client. - DisabledStatus - // DeletedStatus represents a client that will be deleted. - DeletedStatus - - // AllStatus is used for querying purposes to list clients irrespective - // of their status - both enabled and disabled. It is never stored in the - // database as the actual Client status and should always be the largest - // value in this enumeration. - AllStatus -) - -// String representation of the possible status values. -const ( - Disabled = "disabled" - Enabled = "enabled" - Deleted = "deleted" - All = "all" - Unknown = "unknown" -) - -// String converts client/group status to string literal. -func (s Status) String() string { - switch s { - case DisabledStatus: - return Disabled - case EnabledStatus: - return Enabled - case DeletedStatus: - return Deleted - case AllStatus: - return All - default: - return Unknown - } -} - -// ToStatus converts string value to a valid Client status. -func ToStatus(status string) (Status, error) { - switch status { - case "", Enabled: - return EnabledStatus, nil - case Disabled: - return DisabledStatus, nil - case Deleted: - return DeletedStatus, nil - case All: - return AllStatus, nil - } - return Status(0), svcerr.ErrInvalidStatus -} - -// Custom Marshaller for Client. -func (s Status) MarshalJSON() ([]byte, error) { - return json.Marshal(s.String()) -} - -func (client Client) MarshalJSON() ([]byte, error) { - type Alias Client - return json.Marshal(&struct { - Alias - Status string `json:"status,omitempty"` - }{ - Alias: (Alias)(client), - Status: client.Status.String(), - }) -} - -// Custom Unmarshaler for Client. -func (s *Status) UnmarshalJSON(data []byte) error { - str := strings.Trim(string(data), "\"") - val, err := ToStatus(str) - *s = val - return err -} diff --git a/clients/status_test.go b/clients/status_test.go deleted file mode 100644 index 7c3bb9118..000000000 --- a/clients/status_test.go +++ /dev/null @@ -1,246 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package clients_test - -import ( - "testing" - - "github.com/absmach/magistrala/clients" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/stretchr/testify/assert" -) - -func TestStatusString(t *testing.T) { - cases := []struct { - desc string - status clients.Status - expected string - }{ - { - desc: "Enabled", - status: clients.EnabledStatus, - expected: "enabled", - }, - { - desc: "Disabled", - status: clients.DisabledStatus, - expected: "disabled", - }, - { - desc: "Deleted", - status: clients.DeletedStatus, - expected: "deleted", - }, - { - desc: "All", - status: clients.AllStatus, - expected: "all", - }, - { - desc: "Unknown", - status: clients.Status(100), - expected: "unknown", - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - got := tc.status.String() - assert.Equal(t, tc.expected, got, "String() = %v, expected %v", got, tc.expected) - }) - } -} - -func TestToStatus(t *testing.T) { - cases := []struct { - desc string - status string - expetcted clients.Status - err error - }{ - { - desc: "Enabled", - status: "enabled", - expetcted: clients.EnabledStatus, - err: nil, - }, - { - desc: "Disabled", - status: "disabled", - expetcted: clients.DisabledStatus, - err: nil, - }, - { - desc: "Deleted", - status: "deleted", - expetcted: clients.DeletedStatus, - err: nil, - }, - { - desc: "All", - status: "all", - expetcted: clients.AllStatus, - err: nil, - }, - { - desc: "Unknown", - status: "unknown", - expetcted: clients.Status(0), - err: svcerr.ErrInvalidStatus, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - got, err := clients.ToStatus(tc.status) - assert.Equal(t, tc.err, err, "ToStatus() error = %v, expected %v", err, tc.err) - assert.Equal(t, tc.expetcted, got, "ToStatus() = %v, expected %v", got, tc.expetcted) - }) - } -} - -func TestStatusMarshalJSON(t *testing.T) { - cases := []struct { - desc string - expected []byte - status clients.Status - err error - }{ - { - desc: "Enabled", - expected: []byte(`"enabled"`), - status: clients.EnabledStatus, - err: nil, - }, - { - desc: "Disabled", - expected: []byte(`"disabled"`), - status: clients.DisabledStatus, - err: nil, - }, - { - desc: "Deleted", - expected: []byte(`"deleted"`), - status: clients.DeletedStatus, - err: nil, - }, - { - desc: "All", - expected: []byte(`"all"`), - status: clients.AllStatus, - err: nil, - }, - { - desc: "Unknown", - expected: []byte(`"unknown"`), - status: clients.Status(100), - err: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - got, err := tc.status.MarshalJSON() - assert.Equal(t, tc.err, err, "MarshalJSON() error = %v, expected %v", err, tc.err) - assert.Equal(t, tc.expected, got, "MarshalJSON() = %v, expected %v", got, tc.expected) - }) - } -} - -func TestStatusUnmarshalJSON(t *testing.T) { - cases := []struct { - desc string - expected clients.Status - status []byte - err error - }{ - { - desc: "Enabled", - expected: clients.EnabledStatus, - status: []byte(`"enabled"`), - err: nil, - }, - { - desc: "Disabled", - expected: clients.DisabledStatus, - status: []byte(`"disabled"`), - err: nil, - }, - { - desc: "Deleted", - expected: clients.DeletedStatus, - status: []byte(`"deleted"`), - err: nil, - }, - { - desc: "All", - expected: clients.AllStatus, - status: []byte(`"all"`), - err: nil, - }, - { - desc: "Unknown", - expected: clients.Status(0), - status: []byte(`"unknown"`), - err: svcerr.ErrInvalidStatus, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - var s clients.Status - err := s.UnmarshalJSON(tc.status) - assert.Equal(t, tc.err, err, "UnmarshalJSON() error = %v, expected %v", err, tc.err) - assert.Equal(t, tc.expected, s, "UnmarshalJSON() = %v, expected %v", s, tc.expected) - }) - } -} - -func TestClientMarshalJSON(t *testing.T) { - cases := []struct { - desc string - expected []byte - user clients.Client - err error - }{ - { - desc: "Enabled", - expected: []byte(`{"id":"","credentials":{},"created_at":"0001-01-01T00:00:00Z","updated_at":"0001-01-01T00:00:00Z","status":"enabled"}`), - user: clients.Client{Status: clients.EnabledStatus}, - err: nil, - }, - { - desc: "Disabled", - expected: []byte(`{"id":"","credentials":{},"created_at":"0001-01-01T00:00:00Z","updated_at":"0001-01-01T00:00:00Z","status":"disabled"}`), - user: clients.Client{Status: clients.DisabledStatus}, - err: nil, - }, - { - desc: "Deleted", - expected: []byte(`{"id":"","credentials":{},"created_at":"0001-01-01T00:00:00Z","updated_at":"0001-01-01T00:00:00Z","status":"deleted"}`), - user: clients.Client{Status: clients.DeletedStatus}, - err: nil, - }, - { - desc: "All", - expected: []byte(`{"id":"","credentials":{},"created_at":"0001-01-01T00:00:00Z","updated_at":"0001-01-01T00:00:00Z","status":"all"}`), - user: clients.Client{Status: clients.AllStatus}, - err: nil, - }, - { - desc: "Unknown", - expected: []byte(`{"id":"","credentials":{},"created_at":"0001-01-01T00:00:00Z","updated_at":"0001-01-01T00:00:00Z","status":"unknown"}`), - user: clients.Client{Status: clients.Status(100)}, - err: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - got, err := tc.user.MarshalJSON() - assert.Equal(t, tc.err, err, "MarshalJSON() error = %v, expected %v", err, tc.err) - assert.Equal(t, tc.expected, got, "MarshalJSON() = %v, expected %v", string(got), string(tc.expected)) - }) - } -} diff --git a/cmd/alarms/main.go b/cmd/alarms/main.go index a8aa0b7be..6e87a3a55 100644 --- a/cmd/alarms/main.go +++ b/cmd/alarms/main.go @@ -17,38 +17,30 @@ import ( "github.com/absmach/magistrala/alarms/middleware" "github.com/absmach/magistrala/alarms/operations" alarmsRepo "github.com/absmach/magistrala/alarms/postgres" - dpostgres "github.com/absmach/magistrala/domains/postgres" + "github.com/absmach/magistrala/internal/atom" mglog "github.com/absmach/magistrala/logger" smqauthn "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/authn/authsvc" - authsvcAuthz "github.com/absmach/magistrala/pkg/authz/authsvc" - dconsumer "github.com/absmach/magistrala/pkg/domains/events/consumer" - domainsAuthz "github.com/absmach/magistrala/pkg/domains/grpcclient" - "github.com/absmach/magistrala/pkg/grpcclient" + atomauthn "github.com/absmach/magistrala/pkg/authn/atom" "github.com/absmach/magistrala/pkg/jaeger" "github.com/absmach/magistrala/pkg/messaging" brokerstracing "github.com/absmach/magistrala/pkg/messaging/brokers/tracing" "github.com/absmach/magistrala/pkg/permissions" "github.com/absmach/magistrala/pkg/postgres" "github.com/absmach/magistrala/pkg/prometheus" - rconsumer "github.com/absmach/magistrala/pkg/re/events/consumer" "github.com/absmach/magistrala/pkg/server" httpserver "github.com/absmach/magistrala/pkg/server/http" "github.com/absmach/magistrala/pkg/uuid" - rpostgres "github.com/absmach/magistrala/re/postgres" "github.com/caarlos0/env/v11" "golang.org/x/sync/errgroup" ) const ( - svcName = "alarms" - envPrefixDB = "MG_ALARMS_DB_" - envPrefixHTTP = "MG_ALARMS_HTTP_" - envPrefixAuth = "MG_AUTH_GRPC_" - defDB = "alarms" - defSvcHTTPPort = "8050" - envPrefixDomains = "MG_DOMAINS_GRPC_" - alarmEntity = "alarm" + svcName = "alarms" + envPrefixDB = "MG_ALARMS_DB_" + envPrefixHTTP = "MG_ALARMS_HTTP_" + defDB = "alarms" + defSvcHTTPPort = "8050" + alarmEntity = "alarm" ) type config struct { @@ -57,8 +49,6 @@ type config struct { InstanceID string `env:"MG_ALARMS_INSTANCE_ID" envDefault:""` JaegerURL url.URL `env:"MG_JAEGER_URL" envDefault:"http://localhost:4318/v1/traces"` TraceRatio float64 `env:"MG_JAEGER_TRACE_RATIO" envDefault:"1.0"` - ESURL string `env:"MG_ES_URL" envDefault:"nats://localhost:4222"` - ESConsumerName string `env:"MG_ALARMS_EVENT_CONSUMER" envDefault:"alarms"` PermissionsFile string `env:"MG_PERMISSIONS_FILE" envDefault:"permission.yaml"` } @@ -114,68 +104,20 @@ func main() { repo := alarmsRepo.NewAlarmsRepo(db) - authConfig := grpcclient.Config{} - if err := env.ParseWithOptions(&authConfig, env.Options{Prefix: envPrefixAuth}); err != nil { - logger.Error(fmt.Sprintf("failed to load %s auth configuration : %s", svcName, err)) - exitCode = 1 - return - } - authn, authnClient, err := authsvc.NewAuthentication(ctx, authConfig) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - am := smqauthn.NewAuthNMiddleware(authn) - defer authnClient.Close() - logger.Info("AuthN successfully connected to auth gRPC server " + authnClient.Secure()) - - domsGrpcCfg := grpcclient.Config{} - if err := env.ParseWithOptions(&domsGrpcCfg, env.Options{Prefix: envPrefixDomains}); err != nil { - logger.Error(fmt.Sprintf("failed to load domains gRPC client configuration : %s", err)) - exitCode = 1 - return - } - - domAuthz, _, domainsHandler, err := domainsAuthz.NewAuthorization(ctx, domsGrpcCfg) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer domainsHandler.Close() - - authz, authzHandler, err := authsvcAuthz.NewAuthorization(ctx, authConfig, domAuthz) - if err != nil { - logger.Error("failed to create authz " + err.Error()) - exitCode = 1 - return - } - defer authzHandler.Close() - - logger.Info("AuthZ successfully connected to auth gRPC server " + authzHandler.Secure()) - - ddatabase := postgres.NewDatabase(db, dbConfig, tracer) - drepo := dpostgres.NewRepository(ddatabase) - - if err := dconsumer.DomainsEventsSubscribe(ctx, drepo, cfg.ESURL, cfg.ESConsumerName, logger); err != nil { - logger.Error(fmt.Sprintf("failed to create domains event store : %s", err)) - exitCode = 1 - return - } - - rdatabase := postgres.NewDatabase(db, dbConfig, tracer) - rrepo := rpostgres.NewRepository(rdatabase) - - if err := rconsumer.RulesEventsSubscribe(ctx, rrepo, cfg.ESURL, cfg.ESConsumerName, logger); err != nil { - logger.Error(fmt.Sprintf("failed to subscribe to rules events: %s", err)) + atomCfg := atom.LoadConfig() + if atomCfg.URL == "" { + logger.Error("ATOM_URL is required") exitCode = 1 return } + logger.Info("AuthN configured to use Atom bearer tokens") + logger.Info("AuthZ configured to use Atom PDP") + am := smqauthn.NewAuthNMiddleware(atomauthn.NewAuthentication()) idp := uuid.New() svc := alarms.NewService(idp, repo) + svc = alarms.WithAtom(svc, atom.NewClient(atomCfg)) permConfig, err := permissions.ParsePermissionsFile(cfg.PermissionsFile) if err != nil { @@ -205,7 +147,7 @@ func main() { return } - svc, err = middleware.NewAuthorizationMiddleware(svc, authz, entitiesOps) + svc, err = middleware.NewAtomAuthorizationMiddleware(svc, atom.NewClient(atomCfg), entitiesOps) if err != nil { logger.Error(fmt.Sprintf("failed to create authorization middleware: %s", err)) exitCode = 1 diff --git a/cmd/atom-bootstrap/main.go b/cmd/atom-bootstrap/main.go new file mode 100644 index 000000000..22fc4f27e --- /dev/null +++ b/cmd/atom-bootstrap/main.go @@ -0,0 +1,146 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +// Package main contains the one-shot Magistrala Atom bootstrap command. +package main + +import ( + "context" + "flag" + "fmt" + "log" + "os" + "strconv" + "strings" + "time" + + "github.com/absmach/magistrala/internal/atom" +) + +const ( + defaultRetries = 30 + defaultRetryInterval = 2 * time.Second + defaultTimeout = 30 * time.Second +) + +func main() { + log.SetFlags(log.LstdFlags | log.Lmicroseconds) + + cfg := atom.LoadConfig() + if cfg.URL == "" { + log.Fatal("ATOM_URL is required") + } + + client := atom.NewClient(cfg) + if len(os.Args) > 1 { + switch os.Args[1] { + case "bootstrap-actions": + runBootstrapActions(client) + return + case "provision-tokens": + runProvisionTokens(client, os.Args[2:]) + return + default: + log.Fatalf("unknown command %q", os.Args[1]) + } + } + runBootstrapActions(client) +} + +func runBootstrapActions(client *atom.Client) { + retries := envInt("MG_ATOM_BOOTSTRAP_RETRIES", defaultRetries) + retryInterval := envDuration("MG_ATOM_BOOTSTRAP_RETRY_INTERVAL", defaultRetryInterval) + timeout := envDuration("MG_ATOM_BOOTSTRAP_TIMEOUT", defaultTimeout) + + var lastErr error + for attempt := 1; attempt <= retries; attempt++ { + ctx, cancel := context.WithTimeout(context.Background(), timeout) + err := atom.BootstrapMagistralaActions(ctx, client) + cancel() + if err == nil { + log.Printf("Magistrala Atom action bootstrap completed") + return + } + lastErr = err + if attempt < retries { + log.Printf("Magistrala Atom action bootstrap attempt %d/%d failed: %v; retrying in %s", attempt, retries, err, retryInterval) + time.Sleep(retryInterval) + } + } + + log.Fatalf("Magistrala Atom action bootstrap failed after %d attempts: %v", retries, lastErr) +} + +func runProvisionTokens(client *atom.Client, args []string) { + fs := flag.NewFlagSet("provision-tokens", flag.ExitOnError) + output := fs.String("output", envString("MG_ATOM_TOKENS_OUTPUT", "docker/.env.tokens"), "path to write generated token env file") + rotate := fs.String("rotate", "", "rotate one token by name/env var, or all") + entityID := fs.String("entity-id", envString("ATOM_SERVICE_ENTITY_ID", atom.DefaultServiceEntityID), "Atom service entity ID to receive API keys") + if err := fs.Parse(args); err != nil { + log.Fatal(err) + } + + retries := envInt("MG_ATOM_BOOTSTRAP_RETRIES", defaultRetries) + retryInterval := envDuration("MG_ATOM_BOOTSTRAP_RETRY_INTERVAL", defaultRetryInterval) + timeout := envDuration("MG_ATOM_BOOTSTRAP_TIMEOUT", defaultTimeout) + + var lastErr error + for attempt := 1; attempt <= retries; attempt++ { + ctx, cancel := context.WithTimeout(context.Background(), timeout) + result, err := atom.ProvisionServiceTokens(ctx, client, atom.TokenProvisionOptions{ + OutputPath: *output, + ServiceEntityID: *entityID, + Rotate: *rotate, + }) + cancel() + if err == nil { + log.Printf("Magistrala Atom token provisioning completed: output=%s preserved=%d created=%d rotated=%d", + result.OutputPath, len(result.Preserved), len(result.Created), len(result.Rotated)) + return + } + lastErr = err + if attempt < retries { + log.Printf("Magistrala Atom token provisioning attempt %d/%d failed: %v; retrying in %s", attempt, retries, err, retryInterval) + time.Sleep(retryInterval) + } + } + + log.Fatalf("Magistrala Atom token provisioning failed after %d attempts: %v", retries, lastErr) +} + +func envInt(key string, fallback int) int { + raw := strings.TrimSpace(os.Getenv(key)) + if raw == "" { + return fallback + } + value, err := strconv.Atoi(raw) + if err != nil || value <= 0 { + return fallback + } + return value +} + +func envDuration(key string, fallback time.Duration) time.Duration { + raw := strings.TrimSpace(os.Getenv(key)) + if raw == "" { + return fallback + } + value, err := time.ParseDuration(raw) + if err == nil && value > 0 { + return value + } + seconds, err := strconv.Atoi(raw) + if err == nil && seconds > 0 { + return time.Duration(seconds) * time.Second + } + fmt.Fprintf(os.Stderr, "invalid %s=%q, using %s\n", key, raw, fallback) + return fallback +} + +func envString(key, fallback string) string { + value := strings.TrimSpace(os.Getenv(key)) + if value == "" { + return fallback + } + return value +} diff --git a/cmd/auth/main.go b/cmd/auth/main.go index 79afda012..91900c8ed 100644 --- a/cmd/auth/main.go +++ b/cmd/auth/main.go @@ -26,26 +26,23 @@ import ( apostgres "github.com/absmach/magistrala/auth/postgres" "github.com/absmach/magistrala/auth/tokenizer/asymmetric" "github.com/absmach/magistrala/auth/tokenizer/symmetric" + "github.com/absmach/magistrala/internal/atom" redisclient "github.com/absmach/magistrala/internal/clients/redis" mglog "github.com/absmach/magistrala/logger" "github.com/absmach/magistrala/pkg/jaeger" - "github.com/absmach/magistrala/pkg/policies/spicedb" + "github.com/absmach/magistrala/pkg/policies" pgclient "github.com/absmach/magistrala/pkg/postgres" "github.com/absmach/magistrala/pkg/prometheus" "github.com/absmach/magistrala/pkg/server" grpcserver "github.com/absmach/magistrala/pkg/server/grpc" httpserver "github.com/absmach/magistrala/pkg/server/http" "github.com/absmach/magistrala/pkg/uuid" - v1 "github.com/authzed/authzed-go/proto/authzed/api/v1" - "github.com/authzed/authzed-go/v1" - "github.com/authzed/grpcutil" "github.com/caarlos0/env/v11" "github.com/jmoiron/sqlx" "github.com/redis/go-redis/v9" "go.opentelemetry.io/otel/trace" "golang.org/x/sync/errgroup" "google.golang.org/grpc" - "google.golang.org/grpc/credentials/insecure" "google.golang.org/grpc/reflection" ) @@ -71,10 +68,6 @@ type config struct { ActiveKeyPath string `env:"MG_AUTH_KEYS_ACTIVE_KEY_PATH" envDefault:"./keys/active.key"` RetiringKeyPath string `env:"MG_AUTH_KEYS_RETIRING_KEY_PATH" envDefault:""` InvitationDuration time.Duration `env:"MG_AUTH_INVITATION_DURATION" envDefault:"168h"` - SpicedbHost string `env:"MG_SPICEDB_HOST" envDefault:"localhost"` - SpicedbPort string `env:"MG_SPICEDB_PORT" envDefault:"50051"` - SpicedbSchemaFile string `env:"MG_SPICEDB_SCHEMA_FILE" envDefault:"./docker/spicedb/schema.zed"` - SpicedbPreSharedKey string `env:"MG_SPICEDB_PRE_SHARED_KEY" envDefault:"12345678"` TraceRatio float64 `env:"MG_JAEGER_TRACE_RATIO" envDefault:"1.0"` ESURL string `env:"MG_ES_URL" envDefault:"amqp://guest:guest@localhost:5682/"` CacheURL string `env:"MG_AUTH_CACHE_URL" envDefault:"redis://localhost:6379/0"` @@ -143,12 +136,15 @@ func main() { }() tracer := tp.Tracer(svcName) - spicedbclient, err := initSpiceDB(ctx, cfg) - if err != nil { - logger.Error(fmt.Sprintf("failed to init spicedb grpc client : %s\n", err.Error())) + atomCfg := atom.LoadConfig() + if atomCfg.URL == "" { + logger.Error("ATOM_URL is required for auth authorization") exitCode = 1 return } + atomClient := atom.NewClient(atomCfg) + policyEvaluator := atom.NewPolicyEvaluator(atomClient) + logger.Info("AuthZ configured to use Atom PDP") isSymmetric, err := auth.IsSymmetricAlgorithm(cfg.KeyAlgorithm) if err != nil { @@ -183,7 +179,7 @@ func main() { } } - svc, err := newService(db, tracer, cfg, dbConfig, logger, spicedbclient, cacheclient, cfg.CacheKeyDuration, tokenizer, idProvider) + svc, err := newService(db, tracer, cfg, dbConfig, logger, policyEvaluator, nil, cacheclient, cfg.CacheKeyDuration, tokenizer, idProvider) if err != nil { logger.Error(fmt.Sprintf("failed to create service : %s\n", err.Error())) exitCode = 1 @@ -234,36 +230,6 @@ func main() { } } -func initSpiceDB(ctx context.Context, cfg config) (*authzed.ClientWithExperimental, error) { - client, err := authzed.NewClientWithExperimentalAPIs( - fmt.Sprintf("%s:%s", cfg.SpicedbHost, cfg.SpicedbPort), - grpc.WithTransportCredentials(insecure.NewCredentials()), - grpcutil.WithInsecureBearerToken(cfg.SpicedbPreSharedKey), - ) - if err != nil { - return client, err - } - - if err := initSchema(ctx, client, cfg.SpicedbSchemaFile); err != nil { - return client, err - } - - return client, nil -} - -func initSchema(ctx context.Context, client *authzed.ClientWithExperimental, schemaFilePath string) error { - schemaContent, err := os.ReadFile(schemaFilePath) - if err != nil { - return fmt.Errorf("failed to read spice db schema file : %w", err) - } - - if _, err = client.SchemaServiceClient.WriteSchema(ctx, &v1.WriteSchemaRequest{Schema: string(schemaContent)}); err != nil { - return fmt.Errorf("failed to create schema in spicedb : %w", err) - } - - return nil -} - func validateKeyConfig(isSymmetric bool, cfg config, l *slog.Logger) error { if isSymmetric { if cfg.SecretKey == "secret" { @@ -291,7 +257,7 @@ func validateKeyConfig(isSymmetric bool, cfg config, l *slog.Logger) error { return nil } -func newService(db *sqlx.DB, tracer trace.Tracer, cfg config, dbConfig pgclient.Config, logger *slog.Logger, spicedbClient *authzed.ClientWithExperimental, cacheClient *redis.Client, keyDuration time.Duration, tokenizer auth.Tokenizer, idProvider magistrala.IDProvider) (auth.Service, error) { +func newService(db *sqlx.DB, tracer trace.Tracer, cfg config, dbConfig pgclient.Config, logger *slog.Logger, policyEvaluator policies.Evaluator, policyService policies.Service, cacheClient *redis.Client, keyDuration time.Duration, tokenizer auth.Tokenizer, idProvider magistrala.IDProvider) (auth.Service, error) { patsCache := cache.NewPatsCache(cacheClient, keyDuration) tokensCache, err := cache.NewUserActiveTokensCache(cacheClient) if err != nil { @@ -303,10 +269,7 @@ func newService(db *sqlx.DB, tracer trace.Tracer, cfg config, dbConfig pgclient. patsRepo := apostgres.NewPatRepo(database, patsCache) hasher := hasher.New() - pEvaluator := spicedb.NewPolicyEvaluator(spicedbClient, logger) - pService := spicedb.NewPolicyService(spicedbClient, logger) - - svc := auth.New(keysRepo, patsRepo, nil, tokensCache, hasher, idProvider, tokenizer, pEvaluator, pService, cfg.AccessDuration, cfg.RefreshDuration, cfg.InvitationDuration) + svc := auth.New(keysRepo, patsRepo, nil, tokensCache, hasher, idProvider, tokenizer, policyEvaluator, policyService, cfg.AccessDuration, cfg.RefreshDuration, cfg.InvitationDuration) svc = middleware.NewLogging(svc, logger) counter, latency := prometheus.MakeMetrics("auth", "api") svc = middleware.NewMetrics(svc, counter, latency) diff --git a/cmd/bootstrap/main.go b/cmd/bootstrap/main.go deleted file mode 100644 index 5e815df57..000000000 --- a/cmd/bootstrap/main.go +++ /dev/null @@ -1,238 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package main contains bootstrap main function to start the bootstrap service. -package main - -import ( - "context" - "fmt" - "log" - "log/slog" - "net/url" - "os" - - chclient "github.com/absmach/callhome/pkg/client" - "github.com/absmach/magistrala" - "github.com/absmach/magistrala/bootstrap" - httpapi "github.com/absmach/magistrala/bootstrap/api" - "github.com/absmach/magistrala/bootstrap/events/producer" - bootstraphasher "github.com/absmach/magistrala/bootstrap/hasher" - "github.com/absmach/magistrala/bootstrap/middleware" - bootstrappg "github.com/absmach/magistrala/bootstrap/postgres" - "github.com/absmach/magistrala/bootstrap/tracing" - mglog "github.com/absmach/magistrala/logger" - smqauthn "github.com/absmach/magistrala/pkg/authn" - authsvcAuthn "github.com/absmach/magistrala/pkg/authn/authsvc" - smqauthz "github.com/absmach/magistrala/pkg/authz" - authsvcAuthz "github.com/absmach/magistrala/pkg/authz/authsvc" - domainsAuthz "github.com/absmach/magistrala/pkg/domains/grpcclient" - "github.com/absmach/magistrala/pkg/events/store" - "github.com/absmach/magistrala/pkg/grpcclient" - "github.com/absmach/magistrala/pkg/jaeger" - pgclient "github.com/absmach/magistrala/pkg/postgres" - "github.com/absmach/magistrala/pkg/prometheus" - mgsdk "github.com/absmach/magistrala/pkg/sdk" - "github.com/absmach/magistrala/pkg/server" - httpserver "github.com/absmach/magistrala/pkg/server/http" - "github.com/absmach/magistrala/pkg/uuid" - "github.com/caarlos0/env/v11" - "go.opentelemetry.io/otel/trace" - "golang.org/x/sync/errgroup" -) - -const ( - svcName = "bootstrap" - envPrefixDB = "MG_BOOTSTRAP_DB_" - envPrefixHTTP = "MG_BOOTSTRAP_HTTP_" - envPrefixAuth = "MG_AUTH_GRPC_" - envPrefixDomains = "MG_DOMAINS_GRPC_" - defDB = "bootstrap" - defSvcHTTPPort = "9013" -) - -type config struct { - LogLevel string `env:"MG_BOOTSTRAP_LOG_LEVEL" envDefault:"info"` - EncKey string `env:"MG_BOOTSTRAP_ENCRYPT_KEY" envDefault:"12345678910111213141516171819202"` - ESConsumerName string `env:"MG_BOOTSTRAP_EVENT_CONSUMER" envDefault:"bootstrap"` - ClientsURL string `env:"MG_CLIENTS_URL" envDefault:"http://localhost:9006"` - ChannelsURL string `env:"MG_CHANNELS_URL" envDefault:"http://localhost:9005"` - JaegerURL url.URL `env:"MG_JAEGER_URL" envDefault:"http://localhost:4318/v1/traces"` - SendTelemetry bool `env:"MG_SEND_TELEMETRY" envDefault:"true"` - InstanceID string `env:"MG_BOOTSTRAP_INSTANCE_ID" envDefault:""` - ESURL string `env:"MG_ES_URL" envDefault:"nats://localhost:4222"` - TraceRatio float64 `env:"MG_JAEGER_TRACE_RATIO" envDefault:"1.0"` - SpicedbHost string `env:"MG_SPICEDB_HOST" envDefault:"localhost"` - SpicedbPort string `env:"MG_SPICEDB_PORT" envDefault:"50051"` - SpicedbPreSharedKey string `env:"MG_SPICEDB_PRE_SHARED_KEY" envDefault:"12345678"` -} - -func main() { - ctx, cancel := context.WithCancel(context.Background()) - g, ctx := errgroup.WithContext(ctx) - - cfg := config{} - if err := env.Parse(&cfg); err != nil { - log.Fatalf("failed to load %s configuration : %s", svcName, err) - } - - logger, err := mglog.New(os.Stdout, cfg.LogLevel) - if err != nil { - log.Fatalf("failed to init logger: %s", err.Error()) - } - - var exitCode int - defer mglog.ExitWithError(&exitCode) - - if cfg.InstanceID == "" { - if cfg.InstanceID, err = uuid.New().ID(); err != nil { - logger.Error(fmt.Sprintf("failed to generate instanceID: %s", err)) - exitCode = 1 - return - } - } - - // Create new postgres client - dbConfig := pgclient.Config{Name: defDB} - if err := env.ParseWithOptions(&dbConfig, env.Options{Prefix: envPrefixDB}); err != nil { - logger.Error(err.Error()) - } - migration := bootstrappg.Migration() - - db, err := pgclient.Setup(dbConfig, *migration) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer db.Close() - - tp, err := jaeger.NewProvider(ctx, svcName, cfg.JaegerURL, cfg.InstanceID, cfg.TraceRatio) - if err != nil { - logger.Error(fmt.Sprintf("failed to init Jaeger: %s", err)) - exitCode = 1 - return - } - defer func() { - if err := tp.Shutdown(ctx); err != nil { - logger.Error(fmt.Sprintf("error shutting down tracer provider: %v", err)) - } - }() - tracer := tp.Tracer(svcName) - - grpcCfg := grpcclient.Config{} - if err := env.ParseWithOptions(&grpcCfg, env.Options{Prefix: envPrefixAuth}); err != nil { - logger.Error(fmt.Sprintf("failed to load auth gRPC client configuration : %s", err)) - exitCode = 1 - return - } - authn, authnClient, err := authsvcAuthn.NewAuthentication(ctx, grpcCfg) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - am := smqauthn.NewAuthNMiddleware(authn) - logger.Info("AuthN successfully connected to auth gRPC server " + authnClient.Secure()) - defer authnClient.Close() - - domsGrpcCfg := grpcclient.Config{} - if err := env.ParseWithOptions(&domsGrpcCfg, env.Options{Prefix: envPrefixDomains}); err != nil { - logger.Error(fmt.Sprintf("failed to load domains gRPC client configuration : %s", err)) - exitCode = 1 - return - } - domainsAuthz, _, domainsHandler, err := domainsAuthz.NewAuthorization(ctx, domsGrpcCfg) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer domainsHandler.Close() - - authz, authzClient, err := authsvcAuthz.NewAuthorization(ctx, grpcCfg, domainsAuthz) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer authzClient.Close() - logger.Info("AuthZ successfully connected to auth gRPC server " + authzClient.Secure()) - - database := pgclient.NewDatabase(db, dbConfig, tracer) - - // Create new service - svc, err := newService(ctx, authz, database, tracer, logger, cfg) - if err != nil { - logger.Error(fmt.Sprintf("failed to create %s service: %s", svcName, err)) - exitCode = 1 - return - } - - httpServerConfig := server.Config{Port: defSvcHTTPPort} - if err := env.ParseWithOptions(&httpServerConfig, env.Options{Prefix: envPrefixHTTP}); err != nil { - logger.Error(fmt.Sprintf("failed to load %s HTTP server configuration : %s", svcName, err)) - exitCode = 1 - return - } - hs := httpserver.NewServer(ctx, cancel, svcName, httpServerConfig, httpapi.MakeHandler(svc, am, bootstrap.NewConfigReader([]byte(cfg.EncKey)), logger, cfg.InstanceID), logger) - - if cfg.SendTelemetry { - chc := chclient.New(svcName, magistrala.Version, logger, cancel) - go chc.CallHome(ctx) - } - - // Start servers - g.Go(func() error { - return hs.Start() - }) - g.Go(func() error { - return server.StopSignalHandler(ctx, cancel, logger, svcName, hs) - }) - - if err := g.Wait(); err != nil { - logger.Error(fmt.Sprintf("Bootstrap service terminated: %s", err)) - } -} - -func newService(ctx context.Context, authz smqauthz.Authorization, database pgclient.Database, tracer trace.Tracer, logger *slog.Logger, cfg config) (bootstrap.Service, error) { - repoConfig := bootstrappg.NewConfigRepository(database, logger) - repoProfile := bootstrappg.NewProfileRepository(database, logger) - repoBindings := bootstrappg.NewBindingRepository(database, logger) - - config := mgsdk.Config{ - ClientsURL: cfg.ClientsURL, - ChannelsURL: cfg.ChannelsURL, - } - - sdk := mgsdk.NewSDK(config) - idp := uuid.New() - resolver := bootstrap.NewSDKResolver(sdk) - renderer := bootstrap.NewRenderer() - - svc := bootstrap.New( - repoConfig, - repoProfile, - repoBindings, - resolver, - renderer, - sdk, - bootstraphasher.New(), - []byte(cfg.EncKey), - idp, - ) - - publisher, err := store.NewPublisher(ctx, cfg.ESURL, "bootstrap-es-pub") - if err != nil { - return nil, err - } - - svc = middleware.AuthorizationMiddleware(svc, authz) - svc = producer.NewEventStoreMiddleware(svc, publisher) - svc = middleware.LoggingMiddleware(svc, logger) - counter, latency := prometheus.MakeMetrics(svcName, "api") - svc = middleware.MetricsMiddleware(svc, counter, latency) - svc = tracing.New(svc, tracer) - - return svc, nil -} diff --git a/cmd/certs/main.go b/cmd/certs/main.go index 0c6a0a175..c5fed4919 100644 --- a/cmd/certs/main.go +++ b/cmd/certs/main.go @@ -20,13 +20,11 @@ import ( "github.com/absmach/magistrala/certs/middleware" "github.com/absmach/magistrala/certs/pki" "github.com/absmach/magistrala/certs/postgres" + "github.com/absmach/magistrala/internal/atom" mglog "github.com/absmach/magistrala/logger" smqauthn "github.com/absmach/magistrala/pkg/authn" - authsvcAuthn "github.com/absmach/magistrala/pkg/authn/authsvc" + atomauthn "github.com/absmach/magistrala/pkg/authn/atom" smqauthz "github.com/absmach/magistrala/pkg/authz" - authsvcAuthz "github.com/absmach/magistrala/pkg/authz/authsvc" - domainsAuthz "github.com/absmach/magistrala/pkg/domains/grpcclient" - "github.com/absmach/magistrala/pkg/grpcclient" "github.com/absmach/magistrala/pkg/jaeger" pgclient "github.com/absmach/magistrala/pkg/postgres" "github.com/absmach/magistrala/pkg/prometheus" @@ -43,16 +41,14 @@ import ( ) const ( - svcName = "certs" - envPrefixHTTP = "MG_CERTS_HTTP_" - envPrefixDB = "MG_CERTS_DB_" - envPrefixGRPC = "MG_CERTS_GRPC_" - envPrefixAuth = "MG_AUTH_GRPC_" - envPrefixDomains = "MG_DOMAINS_GRPC_" - defSvcHTTPPort = "9010" - defSvcGRPCPort = "7012" - defDB = "certs" - serviceTokenKey = "SERVICE_TOKEN=" + svcName = "certs" + envPrefixHTTP = "MG_CERTS_HTTP_" + envPrefixDB = "MG_CERTS_DB_" + envPrefixGRPC = "MG_CERTS_GRPC_" + defSvcHTTPPort = "9010" + defSvcGRPCPort = "7012" + defDB = "certs" + serviceTokenKey = "SERVICE_TOKEN=" ) type config struct { @@ -184,44 +180,17 @@ func main() { }() tracer := tp.Tracer(svcName) - domsGrpcCfg := grpcclient.Config{} - if err := env.ParseWithOptions(&domsGrpcCfg, env.Options{Prefix: envPrefixDomains}); err != nil { - logger.Error(fmt.Sprintf("failed to load domains gRPC client configuration : %s", err)) + atomCfg := atom.LoadConfig() + if atomCfg.URL == "" { + logger.Error("ATOM_URL is required") exitCode = 1 return } - domAuthz, _, domainsHandler, err := domainsAuthz.NewAuthorization(ctx, domsGrpcCfg) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer domainsHandler.Close() - - authClientConfig := grpcclient.Config{} - if err := env.ParseWithOptions(&authClientConfig, env.Options{Prefix: envPrefixAuth}); err != nil { - logger.Error(fmt.Sprintf("failed to load %s auth configuration : %s", svcName, err)) - exitCode = 1 - return - } - - authn, authnHandler, err := authsvcAuthn.NewAuthentication(ctx, authClientConfig) - if err != nil { - logger.Error("failed to create authn " + err.Error()) - exitCode = 1 - return - } - defer authnHandler.Close() - logger.Info("Authn successfully connected to auth gRPC server " + authnHandler.Secure()) + atomClient := atom.NewClient(atomCfg) + authn := atomauthn.NewAuthentication() authnMiddleware := smqauthn.NewAuthNMiddleware(authn) - authz, authzHandler, err := authsvcAuthz.NewAuthorization(ctx, authClientConfig, domAuthz) - if err != nil { - logger.Error("failed to create authz " + err.Error()) - exitCode = 1 - return - } - defer authzHandler.Close() - logger.Info("Authz successfully connected to auth gRPC server " + authzHandler.Secure()) + authz := atom.NewAuthorizationCompat(atomClient) + logger.Info("AuthN/AuthZ configured to use Atom") httpServerConfig := mgserver.Config{Port: defSvcHTTPPort} if err := env.ParseWithOptions(&httpServerConfig, env.Options{Prefix: envPrefixHTTP}); err != nil { logger.Error(fmt.Sprintf("failed to load %s gRPC server configuration : %s", svcName, err)) diff --git a/cmd/channels/main.go b/cmd/channels/main.go deleted file mode 100644 index eff6df74d..000000000 --- a/cmd/channels/main.go +++ /dev/null @@ -1,489 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package main contains clients main function to start the clients service. -package main - -import ( - "context" - "fmt" - "log" - "log/slog" - "net/url" - "os" - "time" - - chclient "github.com/absmach/callhome/pkg/client" - "github.com/absmach/magistrala" - grpcChannelsV1 "github.com/absmach/magistrala/api/grpc/channels/v1" - grpcClientsV1 "github.com/absmach/magistrala/api/grpc/clients/v1" - grpcGroupsV1 "github.com/absmach/magistrala/api/grpc/groups/v1" - "github.com/absmach/magistrala/auth" - "github.com/absmach/magistrala/channels" - grpcapi "github.com/absmach/magistrala/channels/api/grpc" - httpapi "github.com/absmach/magistrala/channels/api/http" - "github.com/absmach/magistrala/channels/cache" - "github.com/absmach/magistrala/channels/events" - "github.com/absmach/magistrala/channels/middleware" - channelsOps "github.com/absmach/magistrala/channels/operations" - "github.com/absmach/magistrala/channels/postgres" - pChannels "github.com/absmach/magistrala/channels/private" - clientsOps "github.com/absmach/magistrala/clients/operations" - domainsOps "github.com/absmach/magistrala/domains/operations" - dpostgres "github.com/absmach/magistrala/domains/postgres" - groupsOps "github.com/absmach/magistrala/groups/operations" - gpostgres "github.com/absmach/magistrala/groups/postgres" - redisclient "github.com/absmach/magistrala/internal/clients/redis" - mglog "github.com/absmach/magistrala/logger" - smqauthn "github.com/absmach/magistrala/pkg/authn" - authsvcAuthn "github.com/absmach/magistrala/pkg/authn/authsvc" - jwksAuthn "github.com/absmach/magistrala/pkg/authn/jwks" - smqauthz "github.com/absmach/magistrala/pkg/authz" - authsvcAuthz "github.com/absmach/magistrala/pkg/authz/authsvc" - "github.com/absmach/magistrala/pkg/callout" - pkgDomains "github.com/absmach/magistrala/pkg/domains" - dconsumer "github.com/absmach/magistrala/pkg/domains/events/consumer" - domainsAuthz "github.com/absmach/magistrala/pkg/domains/grpcclient" - gconsumer "github.com/absmach/magistrala/pkg/groups/events/consumer" - "github.com/absmach/magistrala/pkg/grpcclient" - jaegerclient "github.com/absmach/magistrala/pkg/jaeger" - "github.com/absmach/magistrala/pkg/permissions" - "github.com/absmach/magistrala/pkg/policies" - "github.com/absmach/magistrala/pkg/policies/spicedb" - pg "github.com/absmach/magistrala/pkg/postgres" - pgclient "github.com/absmach/magistrala/pkg/postgres" - "github.com/absmach/magistrala/pkg/prometheus" - "github.com/absmach/magistrala/pkg/roles" - "github.com/absmach/magistrala/pkg/server" - grpcserver "github.com/absmach/magistrala/pkg/server/grpc" - httpserver "github.com/absmach/magistrala/pkg/server/http" - "github.com/absmach/magistrala/pkg/sid" - spicedbdecoder "github.com/absmach/magistrala/pkg/spicedb" - "github.com/absmach/magistrala/pkg/uuid" - "github.com/authzed/authzed-go/v1" - "github.com/authzed/grpcutil" - "github.com/caarlos0/env/v11" - "github.com/go-chi/chi/v5" - "github.com/jmoiron/sqlx" - "go.opentelemetry.io/otel/trace" - "golang.org/x/sync/errgroup" - "google.golang.org/grpc" - "google.golang.org/grpc/credentials/insecure" - "google.golang.org/grpc/reflection" -) - -const ( - svcName = "channels" - envPrefixDB = "MG_CHANNELS_DB_" - envPrefixHTTP = "MG_CHANNELS_HTTP_" - envPrefixGRPC = "MG_CHANNELS_GRPC_" - envPrefixAuth = "MG_AUTH_GRPC_" - envPrefixClients = "MG_CLIENTS_GRPC_" - envPrefixGroups = "MG_GROUPS_GRPC_" - envPrefixDomains = "MG_DOMAINS_GRPC_" - envPrefixChannelCallout = "MG_CHANNELS_CALLOUT_" - defDB = "channels" - defSvcHTTPPort = "9005" - defSvcGRPCPort = "7005" -) - -type config struct { - LogLevel string `env:"MG_CHANNELS_LOG_LEVEL" envDefault:"info"` - InstanceID string `env:"MG_CHANNELS_INSTANCE_ID" envDefault:""` - JaegerURL url.URL `env:"MG_JAEGER_URL" envDefault:"http://localhost:4318/v1/traces"` - SendTelemetry bool `env:"MG_SEND_TELEMETRY" envDefault:"true"` - CacheURL string `env:"MG_CHANNELS_CACHE_URL" envDefault:"redis://localhost:6379/0"` - CacheKeyDuration time.Duration `env:"MG_CHANNELS_CACHE_KEY_DURATION" envDefault:"10m"` - ESURL string `env:"MG_ES_URL" envDefault:"amqp://guest:guest@localhost:5682/"` - ESConsumerName string `env:"MG_CHANNELS_EVENT_CONSUMER" envDefault:"channels"` - TraceRatio float64 `env:"MG_JAEGER_TRACE_RATIO" envDefault:"1.0"` - SpicedbHost string `env:"MG_SPICEDB_HOST" envDefault:"localhost"` - SpicedbPort string `env:"MG_SPICEDB_PORT" envDefault:"50051"` - SpicedbPreSharedKey string `env:"MG_SPICEDB_PRE_SHARED_KEY" envDefault:"12345678"` - SpicedbSchemaFile string `env:"MG_SPICEDB_SCHEMA_FILE" envDefault:"schema.zed"` - AuthKeyAlgorithm string `env:"MG_AUTH_KEYS_ALGORITHM" envDefault:"RS256"` - JWKSURL string `env:"MG_AUTH_JWKS_URL" envDefault:"http://auth:9001/keys/.well-known/jwks.json"` - PermissionsFile string `env:"MG_PERMISSIONS_FILE" envDefault:"permission.yaml"` -} - -func main() { - ctx, cancel := context.WithCancel(context.Background()) - g, ctx := errgroup.WithContext(ctx) - - // Create new channels configuration - cfg := config{} - if err := env.Parse(&cfg); err != nil { - log.Fatalf("failed to load %s configuration : %s", svcName, err) - } - - var logger *slog.Logger - logger, err := mglog.New(os.Stdout, cfg.LogLevel) - if err != nil { - log.Fatalf("failed to init logger: %s", err.Error()) - } - - var exitCode int - defer mglog.ExitWithError(&exitCode) - - if cfg.InstanceID == "" { - if cfg.InstanceID, err = uuid.New().ID(); err != nil { - logger.Error(fmt.Sprintf("failed to generate instanceID: %s", err)) - exitCode = 1 - return - } - } - - // Create new database for clients - dbConfig := pgclient.Config{Name: defDB} - if err := env.ParseWithOptions(&dbConfig, env.Options{Prefix: envPrefixDB}); err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - migrations, err := postgres.Migration() - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - db, err := pgclient.Setup(dbConfig, *migrations) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer db.Close() - - tp, err := jaegerclient.NewProvider(ctx, svcName, cfg.JaegerURL, cfg.InstanceID, cfg.TraceRatio) - if err != nil { - logger.Error(fmt.Sprintf("Failed to init Jaeger: %s", err)) - exitCode = 1 - return - } - defer func() { - if err := tp.Shutdown(ctx); err != nil { - logger.Error(fmt.Sprintf("Error shutting down tracer provider: %v", err)) - } - }() - tracer := tp.Tracer(svcName) - - policyEvaluator, policyService, err := newSpiceDBPolicyServiceEvaluator(cfg, logger) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - logger.Info("Policy service are successfully connected to SpiceDB gRPC server") - - grpcCfg := grpcclient.Config{} - if err := env.ParseWithOptions(&grpcCfg, env.Options{Prefix: envPrefixAuth}); err != nil { - logger.Error(fmt.Sprintf("failed to load auth gRPC client configuration : %s", err)) - exitCode = 1 - return - } - - isSymmetric, err := auth.IsSymmetricAlgorithm(cfg.AuthKeyAlgorithm) - if err != nil { - logger.Error(fmt.Sprintf("failed to parse auth key algorithm : %s", err)) - exitCode = 1 - return - } - var authn smqauthn.Authentication - var authnClient grpcclient.Handler - switch { - case !isSymmetric: - authn, authnClient, err = jwksAuthn.NewAuthentication(ctx, cfg.JWKSURL, grpcCfg) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer authnClient.Close() - logger.Info("AuthN successfully set up jwks authentication on " + cfg.JWKSURL) - default: - authn, authnClient, err = authsvcAuthn.NewAuthentication(ctx, grpcCfg) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer authnClient.Close() - logger.Info("AuthN successfully connected to auth gRPC server " + authnClient.Secure()) - } - authnMiddleware := smqauthn.NewAuthNMiddleware(authn) - - domsGrpcCfg := grpcclient.Config{} - if err := env.ParseWithOptions(&domsGrpcCfg, env.Options{Prefix: envPrefixDomains}); err != nil { - logger.Error(fmt.Sprintf("failed to load domains gRPC client configuration : %s", err)) - exitCode = 1 - return - } - domAuthz, _, domainsHandler, err := domainsAuthz.NewAuthorization(ctx, domsGrpcCfg) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer domainsHandler.Close() - - callCfg := callout.Config{} - if err := env.ParseWithOptions(&callCfg, env.Options{Prefix: envPrefixChannelCallout}); err != nil { - logger.Error(fmt.Sprintf("failed to parse callout config : %s", err)) - exitCode = 1 - return - } - - authz, authzClient, err := authsvcAuthz.NewAuthorization(ctx, grpcCfg, domAuthz) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer authzClient.Close() - logger.Info("AuthZ successfully connected to auth gRPC server " + authzClient.Secure()) - - thgrpcCfg := grpcclient.Config{} - if err := env.ParseWithOptions(&thgrpcCfg, env.Options{Prefix: envPrefixClients}); err != nil { - logger.Error(fmt.Sprintf("failed to load clients gRPC client configuration : %s", err)) - exitCode = 1 - return - } - clientsClient, clientsHandler, err := grpcclient.SetupClientsClient(ctx, thgrpcCfg) - if err != nil { - logger.Error(fmt.Sprintf("failed to connect to clients gRPC server: %s", err)) - exitCode = 1 - return - } - defer clientsHandler.Close() - logger.Info("Clients gRPC client successfully connected to clients gRPC server " + clientsHandler.Secure()) - - groupsgRPCCfg := grpcclient.Config{} - if err := env.ParseWithOptions(&groupsgRPCCfg, env.Options{Prefix: envPrefixGroups}); err != nil { - logger.Error(fmt.Sprintf("failed to load groups gRPC client configuration : %s", err)) - exitCode = 1 - return - } - groupsClient, groupsHandler, err := grpcclient.SetupGroupsClient(ctx, groupsgRPCCfg) - if err != nil { - logger.Error(fmt.Sprintf("failed to connect to groups gRPC server: %s", err)) - exitCode = 1 - return - } - defer groupsHandler.Close() - logger.Info("Groups gRPC client successfully connected to groups gRPC server " + groupsHandler.Secure()) - - callout, err := callout.New(callCfg) - if err != nil { - logger.Error(fmt.Sprintf("failed to create new callout: %s", err)) - exitCode = 1 - return - } - - cacheclient, err := redisclient.Connect(cfg.CacheURL) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer cacheclient.Close() - cache := cache.NewChannelsCache(cacheclient, cfg.CacheKeyDuration) - - permConfig, err := permissions.ParsePermissionsFile(cfg.PermissionsFile) - if err != nil { - logger.Error(fmt.Sprintf("failed to parse permissions file: %s", err)) - exitCode = 1 - return - } - - svc, psvc, err := newService(ctx, db, dbConfig, cache, authz, policyEvaluator, policyService, - cfg, tracer, clientsClient, groupsClient, domAuthz, logger, callout, permConfig) - if err != nil { - logger.Error(fmt.Sprintf("failed to create services: %s", err)) - exitCode = 1 - return - } - - ddatabase := pg.NewDatabase(db, dbConfig, tracer) - drepo := dpostgres.NewRepository(ddatabase) - - if err := dconsumer.DomainsEventsSubscribe(ctx, drepo, cfg.ESURL, cfg.ESConsumerName, logger); err != nil { - logger.Error(fmt.Sprintf("failed to create domains event store : %s", err)) - exitCode = 1 - return - } - - gdatabase := pg.NewDatabase(db, dbConfig, tracer) - grepo := gpostgres.New(gdatabase) - - if err := gconsumer.GroupsEventsSubscribe(ctx, grepo, cfg.ESURL, cfg.ESConsumerName, logger); err != nil { - logger.Error(fmt.Sprintf("failed to create groups event store : %s", err)) - exitCode = 1 - return - } - - grpcServerConfig := server.Config{Port: defSvcGRPCPort} - if err := env.ParseWithOptions(&grpcServerConfig, env.Options{Prefix: envPrefixGRPC}); err != nil { - logger.Error(fmt.Sprintf("failed to load %s gRPC server configuration : %s", svcName, err)) - exitCode = 1 - return - } - registerChannelsServer := func(srv *grpc.Server) { - reflection.Register(srv) - grpcChannelsV1.RegisterChannelsServiceServer(srv, grpcapi.NewServer(psvc)) - } - - gs := grpcserver.NewServer(ctx, cancel, svcName, grpcServerConfig, registerChannelsServer, logger) - - httpServerConfig := server.Config{Port: defSvcHTTPPort} - if err := env.ParseWithOptions(&httpServerConfig, env.Options{Prefix: envPrefixHTTP}); err != nil { - logger.Error(fmt.Sprintf("failed to load %s HTTP server configuration : %s", svcName, err)) - exitCode = 1 - return - } - mux := chi.NewRouter() - idp := uuid.New() - httpSvc := httpserver.NewServer(ctx, cancel, svcName, httpServerConfig, httpapi.MakeHandler(svc, authnMiddleware, mux, logger, cfg.InstanceID, idp), logger) - - if cfg.SendTelemetry { - chc := chclient.New(svcName, magistrala.Version, logger, cancel) - go chc.CallHome(ctx) - } - - // Start all servers - g.Go(func() error { - return httpSvc.Start() - }) - - g.Go(func() error { - return gs.Start() - }) - - g.Go(func() error { - return server.StopSignalHandler(ctx, cancel, logger, svcName, httpSvc) - }) - - if err := g.Wait(); err != nil { - logger.Error(fmt.Sprintf("%s service terminated: %s", svcName, err)) - } -} - -func newService(ctx context.Context, db *sqlx.DB, dbConfig pgclient.Config, cache channels.Cache, authz smqauthz.Authorization, - pe policies.Evaluator, ps policies.Service, cfg config, tracer trace.Tracer, clientsClient grpcClientsV1.ClientsServiceClient, - groupsClient grpcGroupsV1.GroupsServiceClient, da pkgDomains.Authorization, logger *slog.Logger, callout callout.Callout, - permConfig *permissions.PermissionConfig, -) (channels.Service, pChannels.Service, error) { - database := pg.NewDatabase(db, dbConfig, tracer) - repo := postgres.NewRepository(database) - - idp := uuid.New() - sidp, err := sid.New() - if err != nil { - return nil, nil, err - } - - availableActions, buildInRoles, err := availableActionsAndBuiltInRoles(cfg.SpicedbSchemaFile) - if err != nil { - return nil, nil, err - } - - svc, err := channels.New(repo, cache, ps, idp, clientsClient, groupsClient, sidp, availableActions, buildInRoles) - if err != nil { - return nil, nil, err - } - - svc, err = events.NewEventStoreMiddleware(ctx, svc, cfg.ESURL) - if err != nil { - return nil, nil, err - } - - svc = middleware.NewTracing(svc, tracer) - - counter, latency := prometheus.MakeMetrics("channels", "api") - svc = middleware.NewMetrics(svc, counter, latency) - - channelOps, channelRoleOps, err := permConfig.GetEntityPermissions("channels") - if err != nil { - return nil, nil, fmt.Errorf("failed to get channel permissions: %w", err) - } - - domainOps, _, err := permConfig.GetEntityPermissions("domains") - if err != nil { - return nil, nil, fmt.Errorf("failed to get domain permissions: %w", err) - } - - groupOps, _, err := permConfig.GetEntityPermissions("groups") - if err != nil { - return nil, nil, fmt.Errorf("failed to get group permissions: %w", err) - } - - clientOps, _, err := permConfig.GetEntityPermissions("clients") - if err != nil { - return nil, nil, fmt.Errorf("failed to get client permissions: %w", err) - } - - entitiesOps, err := permissions.NewEntitiesOperations( - permissions.EntitiesPermission{ - policies.ChannelType: channelOps, - policies.DomainType: domainOps, - policies.GroupType: groupOps, - policies.ClientType: clientOps, - }, - permissions.EntitiesOperationDetails[permissions.Operation]{ - policies.ChannelType: channelsOps.OperationDetails(), - policies.DomainType: domainsOps.OperationDetails(), - policies.GroupType: groupsOps.OperationDetails(), - policies.ClientType: clientsOps.OperationDetails(), - }, - ) - if err != nil { - return nil, nil, fmt.Errorf("failed to create entities operations: %w", err) - } - - roleOps, err := permissions.NewOperations(roles.Operations(), channelRoleOps) - if err != nil { - return nil, nil, fmt.Errorf("failed to create role operations: %w", err) - } - - svc, err = middleware.NewAuthorization(policies.ChannelType, svc, authz, repo, entitiesOps, roleOps) - if err != nil { - return nil, nil, err - } - - svc, err = middleware.NewCallout(svc, repo, entitiesOps, roleOps, callout) - if err != nil { - return nil, nil, err - } - - svc = middleware.NewLogging(svc, logger) - - psvc := pChannels.New(repo, cache, pe, ps, da) - return svc, psvc, err -} - -func newSpiceDBPolicyServiceEvaluator(cfg config, logger *slog.Logger) (policies.Evaluator, policies.Service, error) { - client, err := authzed.NewClientWithExperimentalAPIs( - fmt.Sprintf("%s:%s", cfg.SpicedbHost, cfg.SpicedbPort), - grpc.WithTransportCredentials(insecure.NewCredentials()), - grpcutil.WithInsecureBearerToken(cfg.SpicedbPreSharedKey), - ) - if err != nil { - return nil, nil, err - } - ps := spicedb.NewPolicyService(client, logger) - - pe := spicedb.NewPolicyEvaluator(client, logger) - return pe, ps, nil -} - -func availableActionsAndBuiltInRoles(spicedbSchemaFile string) ([]roles.Action, map[roles.BuiltInRoleName][]roles.Action, error) { - availableActions, err := spicedbdecoder.GetActionsFromSchema(spicedbSchemaFile, policies.ChannelType) - if err != nil { - return []roles.Action{}, map[roles.BuiltInRoleName][]roles.Action{}, err - } - - builtInRoles := map[roles.BuiltInRoleName][]roles.Action{ - channels.BuiltInRoleAdmin: availableActions, - } - - return availableActions, builtInRoles, err -} diff --git a/cmd/cli/main.go b/cmd/cli/main.go deleted file mode 100644 index ee56a8551..000000000 --- a/cmd/cli/main.go +++ /dev/null @@ -1,234 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package main contains cli main function to run the cli. -package main - -import ( - "log" - - "github.com/absmach/magistrala/cli" - sdk "github.com/absmach/magistrala/pkg/sdk" - "github.com/spf13/cobra" -) - -func main() { - msgContentType := string(sdk.CTJSONSenML) - sdkConf := sdk.Config{ - MsgContentType: sdk.ContentType(msgContentType), - } - - // Root - rootCmd := &cobra.Command{ - Use: "magistrala-cli", - PersistentPreRun: func(_ *cobra.Command, _ []string) { - cliConf, err := cli.ParseConfig(sdkConf) - if err != nil { - log.Fatalf("Failed to parse config: %s", err) - } - if cliConf.MsgContentType == "" { - cliConf.MsgContentType = sdk.ContentType(msgContentType) - } - s := sdk.NewSDK(cliConf) - cli.SetSDK(s) - }, - } - // API commands - healthCmd := cli.NewHealthCmd() - usersCmd := cli.NewUsersCmd() - domainsCmd := cli.NewDomainsCmd() - clientsCmd := cli.NewClientsCmd() - groupsCmd := cli.NewGroupsCmd() - channelsCmd := cli.NewChannelsCmd() - messagesCmd := cli.NewMessagesCmd() - configCmd := cli.NewConfigCmd() - invitationsCmd := cli.NewInvitationsCmd() - journalCmd := cli.NewJournalCmd() - certsCmd := cli.NewCertsCmd() - - // Root Commands - rootCmd.AddCommand(healthCmd) - rootCmd.AddCommand(usersCmd) - rootCmd.AddCommand(domainsCmd) - rootCmd.AddCommand(groupsCmd) - rootCmd.AddCommand(clientsCmd) - rootCmd.AddCommand(channelsCmd) - rootCmd.AddCommand(messagesCmd) - rootCmd.AddCommand(configCmd) - rootCmd.AddCommand(invitationsCmd) - rootCmd.AddCommand(journalCmd) - rootCmd.AddCommand(certsCmd) - - // Root Flags - rootCmd.PersistentFlags().StringVarP( - &sdkConf.ClientsURL, - "clients-url", - "t", - sdkConf.ClientsURL, - "Clients service URL", - ) - - rootCmd.PersistentFlags().StringVarP( - &sdkConf.UsersURL, - "users-url", - "u", - sdkConf.UsersURL, - "Users service URL", - ) - - rootCmd.PersistentFlags().StringVarP( - &sdkConf.DomainsURL, - "domains-url", - "d", - sdkConf.DomainsURL, - "Domains service URL", - ) - - rootCmd.PersistentFlags().StringVarP( - &sdkConf.HTTPAdapterURL, - "http-url", - "p", - sdkConf.HTTPAdapterURL, - "HTTP adapter URL", - ) - - rootCmd.PersistentFlags().StringVarP( - &sdkConf.JournalURL, - "journal-url", - "a", - sdkConf.JournalURL, - "Journal Log URL", - ) - - rootCmd.PersistentFlags().StringVarP( - &sdkConf.CertsURL, - "certs-url", - "", - sdkConf.CertsURL, - "Certs service URL", - ) - - rootCmd.PersistentFlags().StringVarP( - &sdkConf.HostURL, - "host-url", - "H", - sdkConf.HostURL, - "Host URL", - ) - - rootCmd.PersistentFlags().StringVarP( - &msgContentType, - "content-type", - "y", - msgContentType, - "Message content type", - ) - - rootCmd.PersistentFlags().BoolVarP( - &sdkConf.TLSVerification, - "insecure", - "i", - sdkConf.TLSVerification, - "Do not check for TLS cert", - ) - - rootCmd.PersistentFlags().StringVarP( - &cli.ConfigPath, - "config", - "c", - cli.ConfigPath, - "Config path", - ) - - rootCmd.PersistentFlags().BoolVarP( - &cli.RawOutput, - "raw", - "r", - cli.RawOutput, - "Enables raw output mode for easier parsing of output", - ) - rootCmd.PersistentFlags().BoolVarP( - &sdkConf.CurlFlag, - "curl", - "x", - false, - "Convert HTTP request to cURL command", - ) - - // Client and Channels Flags - rootCmd.PersistentFlags().Uint64VarP( - &cli.Limit, - "limit", - "l", - 10, - "Limit query parameter", - ) - - rootCmd.PersistentFlags().Uint64VarP( - &cli.Offset, - "offset", - "o", - 0, - "Offset query parameter", - ) - - rootCmd.PersistentFlags().StringVarP( - &cli.Name, - "name", - "n", - "", - "Name query parameter", - ) - - rootCmd.PersistentFlags().StringVarP( - &cli.Identity, - "identity", - "I", - "", - "User identity query parameter", - ) - - rootCmd.PersistentFlags().StringVarP( - &cli.Metadata, - "metadata", - "m", - "", - "Metadata query parameter", - ) - - rootCmd.PersistentFlags().StringVarP( - &cli.Status, - "status", - "S", - "", - "User status query parameter", - ) - - rootCmd.PersistentFlags().StringVarP( - &cli.Topic, - "topic", - "T", - "", - "Subscription topic query parameter", - ) - - rootCmd.PersistentFlags().StringVarP( - &cli.Contact, - "contact", - "C", - "", - "Subscription contact query parameter", - ) - - rootCmd.PersistentFlags().BoolVarP( - &sdkConf.Roles, - "roles", - "R", - false, - "Adds option to display roles for entities", - ) - - if err := rootCmd.Execute(); err != nil { - log.Fatal(err) - } -} diff --git a/cmd/clients/main.go b/cmd/clients/main.go deleted file mode 100644 index bfec90838..000000000 --- a/cmd/clients/main.go +++ /dev/null @@ -1,483 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package main contains clients main function to start the clients service. -package main - -import ( - "context" - "fmt" - "log" - "log/slog" - "net/url" - "os" - "time" - - chclient "github.com/absmach/callhome/pkg/client" - "github.com/absmach/magistrala" - grpcChannelsV1 "github.com/absmach/magistrala/api/grpc/channels/v1" - grpcClientsV1 "github.com/absmach/magistrala/api/grpc/clients/v1" - grpcGroupsV1 "github.com/absmach/magistrala/api/grpc/groups/v1" - "github.com/absmach/magistrala/auth" - "github.com/absmach/magistrala/clients" - grpcapi "github.com/absmach/magistrala/clients/api/grpc" - httpapi "github.com/absmach/magistrala/clients/api/http" - "github.com/absmach/magistrala/clients/cache" - "github.com/absmach/magistrala/clients/events" - "github.com/absmach/magistrala/clients/middleware" - clientsOps "github.com/absmach/magistrala/clients/operations" - "github.com/absmach/magistrala/clients/postgres" - pClients "github.com/absmach/magistrala/clients/private" - doperations "github.com/absmach/magistrala/domains/operations" - dpostgres "github.com/absmach/magistrala/domains/postgres" - goperations "github.com/absmach/magistrala/groups/operations" - gpostgres "github.com/absmach/magistrala/groups/postgres" - redisclient "github.com/absmach/magistrala/internal/clients/redis" - mglog "github.com/absmach/magistrala/logger" - smqauthn "github.com/absmach/magistrala/pkg/authn" - authsvcAuthn "github.com/absmach/magistrala/pkg/authn/authsvc" - jwksAuthn "github.com/absmach/magistrala/pkg/authn/jwks" - smqauthz "github.com/absmach/magistrala/pkg/authz" - authsvcAuthz "github.com/absmach/magistrala/pkg/authz/authsvc" - "github.com/absmach/magistrala/pkg/callout" - dconsumer "github.com/absmach/magistrala/pkg/domains/events/consumer" - domainsAuthz "github.com/absmach/magistrala/pkg/domains/grpcclient" - gconsumer "github.com/absmach/magistrala/pkg/groups/events/consumer" - "github.com/absmach/magistrala/pkg/grpcclient" - jaegerclient "github.com/absmach/magistrala/pkg/jaeger" - "github.com/absmach/magistrala/pkg/permissions" - "github.com/absmach/magistrala/pkg/policies" - "github.com/absmach/magistrala/pkg/policies/spicedb" - pg "github.com/absmach/magistrala/pkg/postgres" - pgclient "github.com/absmach/magistrala/pkg/postgres" - "github.com/absmach/magistrala/pkg/prometheus" - "github.com/absmach/magistrala/pkg/roles" - "github.com/absmach/magistrala/pkg/server" - grpcserver "github.com/absmach/magistrala/pkg/server/grpc" - httpserver "github.com/absmach/magistrala/pkg/server/http" - "github.com/absmach/magistrala/pkg/sid" - spicedbdecoder "github.com/absmach/magistrala/pkg/spicedb" - "github.com/absmach/magistrala/pkg/uuid" - "github.com/authzed/authzed-go/v1" - "github.com/authzed/grpcutil" - "github.com/caarlos0/env/v11" - "github.com/go-chi/chi/v5" - "github.com/jmoiron/sqlx" - "github.com/redis/go-redis/v9" - "go.opentelemetry.io/otel/trace" - "golang.org/x/sync/errgroup" - "google.golang.org/grpc" - "google.golang.org/grpc/credentials/insecure" - "google.golang.org/grpc/reflection" -) - -const ( - svcName = "clients" - envPrefixDB = "MG_CLIENTS_DB_" - envPrefixHTTP = "MG_CLIENTS_HTTP_" - envPrefixGRPC = "MG_CLIENTS_GRPC_" - envPrefixAuth = "MG_AUTH_GRPC_" - envPrefixChannels = "MG_CHANNELS_GRPC_" - envPrefixGroups = "MG_GROUPS_GRPC_" - envPrefixDomains = "MG_DOMAINS_GRPC_" - envPrefixClientCallout = "MG_CLIENTS_CALLOUT_" - defDB = "clients" - defSvcHTTPPort = "9000" - defSvcAuthGRPCPort = "7000" -) - -type config struct { - InstanceID string `env:"MG_CLIENTS_INSTANCE_ID" envDefault:""` - LogLevel string `env:"MG_CLIENTS_LOG_LEVEL" envDefault:"info"` - StandaloneID string `env:"MG_CLIENTS_STANDALONE_ID" envDefault:""` - StandaloneToken string `env:"MG_CLIENTS_STANDALONE_TOKEN" envDefault:""` - CacheURL string `env:"MG_CLIENTS_CACHE_URL" envDefault:"redis://localhost:6379/0"` - CacheKeyDuration time.Duration `env:"MG_CLIENTS_CACHE_KEY_DURATION" envDefault:"1h"` - JaegerURL url.URL `env:"MG_JAEGER_URL" envDefault:"http://localhost:4318/v1/traces"` - SendTelemetry bool `env:"MG_SEND_TELEMETRY" envDefault:"true"` - ESURL string `env:"MG_ES_URL" envDefault:"amqp://guest:guest@localhost:5682/"` - ESConsumerName string `env:"MG_CLIENTS_EVENT_CONSUMER" envDefault:"clients"` - TraceRatio float64 `env:"MG_JAEGER_TRACE_RATIO" envDefault:"1.0"` - SpicedbHost string `env:"MG_SPICEDB_HOST" envDefault:"localhost"` - SpicedbPort string `env:"MG_SPICEDB_PORT" envDefault:"50051"` - SpicedbPreSharedKey string `env:"MG_SPICEDB_PRE_SHARED_KEY" envDefault:"12345678"` - SpicedbSchemaFile string `env:"MG_SPICEDB_SCHEMA_FILE" envDefault:"schema.zed"` - AuthKeyAlgorithm string `env:"MG_AUTH_KEYS_ALGORITHM" envDefault:"RS256"` - JWKSURL string `env:"MG_AUTH_JWKS_URL" envDefault:"http://auth:9001/keys/.well-known/jwks.json"` - PermissionsFile string `env:"MG_PERMISSIONS_FILE" envDefault:"permission.yaml"` -} - -func main() { - ctx, cancel := context.WithCancel(context.Background()) - g, ctx := errgroup.WithContext(ctx) - - // Create new clients configuration - cfg := config{} - if err := env.Parse(&cfg); err != nil { - log.Fatalf("failed to load %s configuration : %s", svcName, err) - } - - var logger *slog.Logger - logger, err := mglog.New(os.Stdout, cfg.LogLevel) - if err != nil { - log.Fatalf("failed to init logger: %s", err.Error()) - } - - var exitCode int - defer mglog.ExitWithError(&exitCode) - - if cfg.InstanceID == "" { - if cfg.InstanceID, err = uuid.New().ID(); err != nil { - logger.Error(fmt.Sprintf("failed to generate instanceID: %s", err)) - exitCode = 1 - return - } - } - - // Create new database for clients - dbConfig := pgclient.Config{Name: defDB} - if err := env.ParseWithOptions(&dbConfig, env.Options{Prefix: envPrefixDB}); err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - tm, err := postgres.Migration() - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - db, err := pgclient.Setup(dbConfig, *tm) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer db.Close() - - tp, err := jaegerclient.NewProvider(ctx, svcName, cfg.JaegerURL, cfg.InstanceID, cfg.TraceRatio) - if err != nil { - logger.Error(fmt.Sprintf("Failed to init Jaeger: %s", err)) - exitCode = 1 - return - } - defer func() { - if err := tp.Shutdown(ctx); err != nil { - logger.Error(fmt.Sprintf("Error shutting down tracer provider: %v", err)) - } - }() - tracer := tp.Tracer(svcName) - - // Setup new redis cache client - cacheclient, err := redisclient.Connect(cfg.CacheURL) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer cacheclient.Close() - - policyEvaluator, policyService, err := newSpiceDBPolicyServiceEvaluator(cfg, logger) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - logger.Info("Policy evaluator and Policy manager are successfully connected to SpiceDB gRPC server") - - grpcCfg := grpcclient.Config{} - if err := env.ParseWithOptions(&grpcCfg, env.Options{Prefix: envPrefixAuth}); err != nil { - logger.Error(fmt.Sprintf("failed to load auth gRPC client configuration : %s", err)) - exitCode = 1 - return - } - - alg, err := auth.IsSymmetricAlgorithm(cfg.AuthKeyAlgorithm) - if err != nil { - logger.Error(fmt.Sprintf("failed to parse auth key algorithm : %s", err)) - exitCode = 1 - return - } - var authn smqauthn.Authentication - var authnClient grpcclient.Handler - switch { - case !alg: - authn, authnClient, err = jwksAuthn.NewAuthentication(ctx, cfg.JWKSURL, grpcCfg) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer authnClient.Close() - logger.Info("AuthN successfully set up jwks authentication on " + cfg.JWKSURL) - default: - authn, authnClient, err = authsvcAuthn.NewAuthentication(ctx, grpcCfg) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer authnClient.Close() - logger.Info("AuthN successfully connected to auth gRPC server " + authnClient.Secure()) - } - authnMiddleware := smqauthn.NewAuthNMiddleware(authn) - - domsGrpcCfg := grpcclient.Config{} - if err := env.ParseWithOptions(&domsGrpcCfg, env.Options{Prefix: envPrefixDomains}); err != nil { - logger.Error(fmt.Sprintf("failed to load domains gRPC client configuration : %s", err)) - exitCode = 1 - return - } - domAuthz, _, domainsHandler, err := domainsAuthz.NewAuthorization(ctx, domsGrpcCfg) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer domainsHandler.Close() - - callCfg := callout.Config{} - if err := env.ParseWithOptions(&callCfg, env.Options{Prefix: envPrefixClientCallout}); err != nil { - logger.Error(fmt.Sprintf("failed to parse callout config : %s", err)) - exitCode = 1 - return - } - - authz, authzClient, err := authsvcAuthz.NewAuthorization(ctx, grpcCfg, domAuthz) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer authzClient.Close() - logger.Info("AuthZ successfully connected to auth gRPC server " + authzClient.Secure()) - - chgrpccfg := grpcclient.Config{} - if err := env.ParseWithOptions(&chgrpccfg, env.Options{Prefix: envPrefixChannels}); err != nil { - logger.Error(fmt.Sprintf("failed to load channels gRPC client configuration : %s", err)) - exitCode = 1 - return - } - channelsgRPC, channelsClient, err := grpcclient.SetupChannelsClient(ctx, chgrpccfg) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - logger.Info("Channels gRPC client successfully connected to channels gRPC server " + channelsClient.Secure()) - defer channelsClient.Close() - - groupsgRPCCfg := grpcclient.Config{} - if err := env.ParseWithOptions(&groupsgRPCCfg, env.Options{Prefix: envPrefixGroups}); err != nil { - logger.Error(fmt.Sprintf("failed to load groups gRPC client configuration : %s", err)) - exitCode = 1 - return - } - groupsClient, groupsHandler, err := grpcclient.SetupGroupsClient(ctx, groupsgRPCCfg) - if err != nil { - logger.Error(fmt.Sprintf("failed to connect to groups gRPC server: %s", err)) - exitCode = 1 - return - } - defer groupsHandler.Close() - logger.Info("Groups gRPC client successfully connected to groups gRPC server " + groupsHandler.Secure()) - - callout, err := callout.New(callCfg) - if err != nil { - logger.Error(fmt.Sprintf("failed to create new callout: %s", err)) - exitCode = 1 - return - } - - permConfig, err := permissions.ParsePermissionsFile(cfg.PermissionsFile) - if err != nil { - logger.Error(fmt.Sprintf("failed to parse permissions file: %s", err)) - exitCode = 1 - return - } - - svc, psvc, err := newService(ctx, db, dbConfig, authz, policyEvaluator, policyService, cacheclient, - cfg, channelsgRPC, groupsClient, tracer, logger, callout, permConfig) - if err != nil { - logger.Error(fmt.Sprintf("failed to create services: %s", err)) - exitCode = 1 - return - } - - ddatabase := pg.NewDatabase(db, dbConfig, tracer) - drepo := dpostgres.NewRepository(ddatabase) - - if err := dconsumer.DomainsEventsSubscribe(ctx, drepo, cfg.ESURL, cfg.ESConsumerName, logger); err != nil { - logger.Error(fmt.Sprintf("failed to create domains event store : %s", err)) - exitCode = 1 - return - } - - gdatabase := pg.NewDatabase(db, dbConfig, tracer) - grepo := gpostgres.New(gdatabase) - - if err := gconsumer.GroupsEventsSubscribe(ctx, grepo, cfg.ESURL, cfg.ESConsumerName, logger); err != nil { - logger.Error(fmt.Sprintf("failed to create groups event store : %s", err)) - exitCode = 1 - return - } - - httpServerConfig := server.Config{Port: defSvcHTTPPort} - if err := env.ParseWithOptions(&httpServerConfig, env.Options{Prefix: envPrefixHTTP}); err != nil { - logger.Error(fmt.Sprintf("failed to load %s HTTP server configuration : %s", svcName, err)) - exitCode = 1 - return - } - mux := chi.NewRouter() - idp := uuid.New() - httpSvc := httpserver.NewServer(ctx, cancel, svcName, httpServerConfig, httpapi.MakeHandler(svc, authnMiddleware, mux, logger, cfg.InstanceID, idp), logger) - - grpcServerConfig := server.Config{Port: defSvcAuthGRPCPort} - if err := env.ParseWithOptions(&grpcServerConfig, env.Options{Prefix: envPrefixGRPC}); err != nil { - logger.Error(fmt.Sprintf("failed to load %s gRPC server configuration : %s", svcName, err)) - exitCode = 1 - return - } - - registerClientsServer := func(srv *grpc.Server) { - reflection.Register(srv) - grpcClientsV1.RegisterClientsServiceServer(srv, grpcapi.NewServer(psvc)) - } - gs := grpcserver.NewServer(ctx, cancel, svcName, grpcServerConfig, registerClientsServer, logger) - - if cfg.SendTelemetry { - chc := chclient.New(svcName, magistrala.Version, logger, cancel) - go chc.CallHome(ctx) - } - - // Start all servers - g.Go(func() error { - return httpSvc.Start() - }) - - g.Go(func() error { - return gs.Start() - }) - - g.Go(func() error { - return server.StopSignalHandler(ctx, cancel, logger, svcName, httpSvc) - }) - - if err := g.Wait(); err != nil { - logger.Error(fmt.Sprintf("%s service terminated: %s", svcName, err)) - } -} - -func newService(ctx context.Context, db *sqlx.DB, dbConfig pgclient.Config, authz smqauthz.Authorization, pe policies.Evaluator, ps policies.Service, cacheClient *redis.Client, cfg config, channels grpcChannelsV1.ChannelsServiceClient, groups grpcGroupsV1.GroupsServiceClient, tracer trace.Tracer, logger *slog.Logger, callout callout.Callout, permConfig *permissions.PermissionConfig) (clients.Service, pClients.Service, error) { - database := pg.NewDatabase(db, dbConfig, tracer) - repo := postgres.NewRepository(database) - - idp := uuid.New() - sidp, err := sid.New() - if err != nil { - return nil, nil, err - } - - // Clients service - cache := cache.NewCache(cacheClient, cfg.CacheKeyDuration) - - availableActions, builtInRoles, err := availableActionsAndBuiltInRoles(cfg.SpicedbSchemaFile) - if err != nil { - return nil, nil, err - } - - csvc, err := clients.NewService(repo, ps, cache, channels, groups, idp, sidp, availableActions, builtInRoles) - if err != nil { - return nil, nil, err - } - - csvc, err = events.NewEventStoreMiddleware(ctx, csvc, cfg.ESURL) - if err != nil { - return nil, nil, err - } - - csvc = middleware.NewTracing(csvc, tracer) - - counter, latency := prometheus.MakeMetrics(svcName, "api") - csvc = middleware.NewMetrics(csvc, counter, latency) - - clientOps, clientRoleOps, err := permConfig.GetEntityPermissions("clients") - if err != nil { - return nil, nil, fmt.Errorf("failed to get client permissions: %w", err) - } - - domainOps, _, err := permConfig.GetEntityPermissions("domains") - if err != nil { - return nil, nil, fmt.Errorf("failed to get domain permissions: %w", err) - } - - groupOps, _, err := permConfig.GetEntityPermissions("groups") - if err != nil { - return nil, nil, fmt.Errorf("failed to get group permissions: %w", err) - } - - entitiesOps, err := permissions.NewEntitiesOperations( - permissions.EntitiesPermission{ - policies.ClientType: clientOps, - policies.DomainType: domainOps, - policies.GroupType: groupOps, - }, - permissions.EntitiesOperationDetails[permissions.Operation]{ - policies.ClientType: clientsOps.OperationDetails(), - policies.DomainType: doperations.OperationDetails(), - policies.GroupType: goperations.OperationDetails(), - }, - ) - if err != nil { - return nil, nil, fmt.Errorf("failed to create entities operations: %w", err) - } - - roleOps, err := permissions.NewOperations(roles.Operations(), clientRoleOps) - if err != nil { - return nil, nil, fmt.Errorf("failed to create role operations: %w", err) - } - - csvc, err = middleware.NewAuthorization(policies.ClientType, csvc, authz, repo, entitiesOps, roleOps) - if err != nil { - return nil, nil, err - } - - csvc, err = middleware.NewCallout(csvc, repo, entitiesOps, roleOps, callout) - if err != nil { - return nil, nil, err - } - - csvc = middleware.NewLogging(csvc, logger) - - isvc := pClients.New(repo, cache, pe, ps) - - return csvc, isvc, err -} - -func newSpiceDBPolicyServiceEvaluator(cfg config, logger *slog.Logger) (policies.Evaluator, policies.Service, error) { - client, err := authzed.NewClientWithExperimentalAPIs( - fmt.Sprintf("%s:%s", cfg.SpicedbHost, cfg.SpicedbPort), - grpc.WithTransportCredentials(insecure.NewCredentials()), - grpcutil.WithInsecureBearerToken(cfg.SpicedbPreSharedKey), - ) - if err != nil { - return nil, nil, err - } - pe := spicedb.NewPolicyEvaluator(client, logger) - ps := spicedb.NewPolicyService(client, logger) - - return pe, ps, nil -} - -func availableActionsAndBuiltInRoles(spicedbSchemaFile string) ([]roles.Action, map[roles.BuiltInRoleName][]roles.Action, error) { - availableActions, err := spicedbdecoder.GetActionsFromSchema(spicedbSchemaFile, policies.ClientType) - if err != nil { - return []roles.Action{}, map[roles.BuiltInRoleName][]roles.Action{}, err - } - - builtInRoles := map[roles.BuiltInRoleName][]roles.Action{ - clients.BuiltInRoleAdmin: availableActions, - } - - return availableActions, builtInRoles, err -} diff --git a/cmd/domains/main.go b/cmd/domains/main.go deleted file mode 100644 index 409e49739..000000000 --- a/cmd/domains/main.go +++ /dev/null @@ -1,378 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package main - -import ( - "context" - "fmt" - "log" - "log/slog" - "net/url" - "os" - "time" - - chclient "github.com/absmach/callhome/pkg/client" - "github.com/absmach/magistrala" - grpcDomainsV1 "github.com/absmach/magistrala/api/grpc/domains/v1" - "github.com/absmach/magistrala/auth" - "github.com/absmach/magistrala/domains" - domainsSvc "github.com/absmach/magistrala/domains" - domainsgrpcapi "github.com/absmach/magistrala/domains/api/grpc" - httpapi "github.com/absmach/magistrala/domains/api/http" - cache "github.com/absmach/magistrala/domains/cache" - "github.com/absmach/magistrala/domains/events" - dmw "github.com/absmach/magistrala/domains/middleware" - doperations "github.com/absmach/magistrala/domains/operations" - dpostgres "github.com/absmach/magistrala/domains/postgres" - "github.com/absmach/magistrala/domains/private" - redisclient "github.com/absmach/magistrala/internal/clients/redis" - mglog "github.com/absmach/magistrala/logger" - smqauthn "github.com/absmach/magistrala/pkg/authn" - authsvcAuthn "github.com/absmach/magistrala/pkg/authn/authsvc" - jwksAuthn "github.com/absmach/magistrala/pkg/authn/jwks" - "github.com/absmach/magistrala/pkg/authz" - authsvcAuthz "github.com/absmach/magistrala/pkg/authz/authsvc" - "github.com/absmach/magistrala/pkg/callout" - domainsAuthz "github.com/absmach/magistrala/pkg/domains/psvc" - "github.com/absmach/magistrala/pkg/grpcclient" - "github.com/absmach/magistrala/pkg/jaeger" - "github.com/absmach/magistrala/pkg/permissions" - "github.com/absmach/magistrala/pkg/policies" - "github.com/absmach/magistrala/pkg/policies/spicedb" - "github.com/absmach/magistrala/pkg/postgres" - pgclient "github.com/absmach/magistrala/pkg/postgres" - "github.com/absmach/magistrala/pkg/prometheus" - "github.com/absmach/magistrala/pkg/roles" - "github.com/absmach/magistrala/pkg/server" - grpcserver "github.com/absmach/magistrala/pkg/server/grpc" - httpserver "github.com/absmach/magistrala/pkg/server/http" - "github.com/absmach/magistrala/pkg/sid" - spicedbdecoder "github.com/absmach/magistrala/pkg/spicedb" - "github.com/absmach/magistrala/pkg/uuid" - "github.com/authzed/authzed-go/v1" - "github.com/authzed/grpcutil" - "github.com/caarlos0/env/v11" - "github.com/go-chi/chi/v5" - "go.opentelemetry.io/otel/trace" - "golang.org/x/sync/errgroup" - "google.golang.org/grpc" - "google.golang.org/grpc/credentials/insecure" - "google.golang.org/grpc/reflection" -) - -const ( - svcName = "domains" - envPrefixHTTP = "MG_DOMAINS_HTTP_" - envPrefixGrpc = "MG_DOMAINS_GRPC_" - envPrefixDB = "MG_DOMAINS_DB_" - envPrefixAuth = "MG_AUTH_GRPC_" - envPrefixDomainCallout = "MG_DOMAINS_CALLOUT_" - defDB = "domains" - defSvcHTTPPort = "9004" - defSvcGRPCPort = "7004" -) - -type config struct { - LogLevel string `env:"MG_DOMAINS_LOG_LEVEL" envDefault:"info"` - JaegerURL url.URL `env:"MG_JAEGER_URL" envDefault:"http://localhost:4318/v1/traces"` - SendTelemetry bool `env:"MG_SEND_TELEMETRY" envDefault:"true"` - CacheURL string `env:"MG_DOMAINS_CACHE_URL" envDefault:"redis://localhost:6379/0"` - CacheKeyDuration time.Duration `env:"MG_DOMAINS_CACHE_KEY_DURATION" envDefault:"10m"` - InstanceID string `env:"MG_DOMAINS_INSTANCE_ID" envDefault:""` - SpicedbHost string `env:"MG_SPICEDB_HOST" envDefault:"localhost"` - SpicedbPort string `env:"MG_SPICEDB_PORT" envDefault:"50051"` - SpicedbSchemaFile string `env:"MG_SPICEDB_SCHEMA_FILE" envDefault:"schema.zed"` - SpicedbPreSharedKey string `env:"MG_SPICEDB_PRE_SHARED_KEY" envDefault:"12345678"` - TraceRatio float64 `env:"MG_JAEGER_TRACE_RATIO" envDefault:"1.0"` - ESURL string `env:"MG_ES_URL" envDefault:"amqp://guest:guest@localhost:5682/"` - AuthKeyAlgorithm string `env:"MG_AUTH_KEYS_ALGORITHM" envDefault:"RS256"` - JWKSURL string `env:"MG_AUTH_JWKS_URL" envDefault:"http://auth:9001/keys/.well-known/jwks.json"` - PermissionsFile string `env:"MG_PERMISSIONS_FILE" envDefault:"permission.yaml"` -} - -func main() { - ctx, cancel := context.WithCancel(context.Background()) - g, ctx := errgroup.WithContext(ctx) - - cfg := config{} - if err := env.Parse(&cfg); err != nil { - log.Fatalf("failed to load %s configuration : %s", svcName, err.Error()) - } - - logger, err := mglog.New(os.Stdout, cfg.LogLevel) - if err != nil { - log.Fatalf("failed to init logger: %s", err.Error()) - } - - var exitCode int - defer mglog.ExitWithError(&exitCode) - - if cfg.InstanceID == "" { - if cfg.InstanceID, err = uuid.New().ID(); err != nil { - logger.Error(fmt.Sprintf("failed to generate instanceID: %s", err)) - exitCode = 1 - return - } - } - - dbConfig := pgclient.Config{Name: defDB} - if err := env.ParseWithOptions(&dbConfig, env.Options{Prefix: envPrefixDB}); err != nil { - logger.Error(err.Error()) - } - - dm, err := dpostgres.Migration() - if err != nil { - logger.Error(fmt.Sprintf("failed create migrations for domain: %s", err.Error())) - exitCode = 1 - return - } - - db, err := pgclient.Setup(dbConfig, *dm) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer db.Close() - - tp, err := jaeger.NewProvider(ctx, svcName, cfg.JaegerURL, cfg.InstanceID, cfg.TraceRatio) - if err != nil { - logger.Error(fmt.Sprintf("failed to init Jaeger: %s", err)) - exitCode = 1 - return - } - defer func() { - if err := tp.Shutdown(ctx); err != nil { - logger.Error(fmt.Sprintf("error shutting down tracer provider: %v", err)) - } - }() - tracer := tp.Tracer(svcName) - - time.Sleep(1 * time.Second) - - clientConfig := grpcclient.Config{} - if err := env.ParseWithOptions(&clientConfig, env.Options{Prefix: envPrefixAuth}); err != nil { - logger.Error(fmt.Sprintf("failed to load auth gRPC server configuration : %s", err)) - exitCode = 1 - return - } - - isSymmetric, err := auth.IsSymmetricAlgorithm(cfg.AuthKeyAlgorithm) - if err != nil { - logger.Error(fmt.Sprintf("failed to parse auth key algorithm : %s", err)) - exitCode = 1 - return - } - var authn smqauthn.Authentication - var authnClient grpcclient.Handler - switch { - case !isSymmetric: - authn, authnClient, err = jwksAuthn.NewAuthentication(ctx, cfg.JWKSURL, clientConfig) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer authnClient.Close() - logger.Info("AuthN successfully set up jwks authentication on " + cfg.JWKSURL) - default: - authn, authnClient, err = authsvcAuthn.NewAuthentication(ctx, clientConfig) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer authnClient.Close() - logger.Info("AuthN successfully connected to auth gRPC server " + authnClient.Secure()) - } - authnMiddleware := smqauthn.NewAuthNMiddleware(authn) - - database := postgres.NewDatabase(db, dbConfig, tracer) - domainsRepo := dpostgres.NewRepository(database) - - cacheclient, err := redisclient.Connect(cfg.CacheURL) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer cacheclient.Close() - cache := cache.NewDomainsCache(cacheclient, cfg.CacheKeyDuration) - - psvc := private.New(domainsRepo, cache) - - domAuthz := domainsAuthz.NewAuthorization(psvc) - - authz, authzHandler, err := authsvcAuthz.NewAuthorization(ctx, clientConfig, domAuthz) - if err != nil { - logger.Error(fmt.Sprintf("authz failed to connect to auth gRPC server : %s", err.Error())) - exitCode = 1 - return - } - defer authzHandler.Close() - logger.Info("Authz successfully connected to auth gRPC server " + authzHandler.Secure()) - - policyService, err := newPolicyService(cfg, logger) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - logger.Info("Policy client successfully connected to spicedb gRPC server") - - callCfg := callout.Config{} - if err := env.ParseWithOptions(&callCfg, env.Options{Prefix: envPrefixDomainCallout}); err != nil { - logger.Error(fmt.Sprintf("failed to parse callout config : %s", err)) - exitCode = 1 - return - } - - call, err := callout.New(callCfg) - if err != nil { - logger.Error(fmt.Sprintf("failed to create new callout: %s", err)) - exitCode = 1 - return - } - - svc, err := newDomainService(ctx, domainsRepo, cache, tracer, cfg, authz, policyService, logger, call) - if err != nil { - logger.Error(fmt.Sprintf("failed to create %s service: %s", svcName, err.Error())) - exitCode = 1 - return - } - - grpcServerConfig := server.Config{Port: defSvcGRPCPort} - if err := env.ParseWithOptions(&grpcServerConfig, env.Options{Prefix: envPrefixGrpc}); err != nil { - logger.Error(fmt.Sprintf("failed to load %s gRPC server configuration : %s", svcName, err.Error())) - exitCode = 1 - return - } - registerDomainsServiceServer := func(srv *grpc.Server) { - reflection.Register(srv) - grpcDomainsV1.RegisterDomainsServiceServer(srv, domainsgrpcapi.NewDomainsServer(psvc)) - } - - gs := grpcserver.NewServer(ctx, cancel, svcName, grpcServerConfig, registerDomainsServiceServer, logger) - - g.Go(func() error { - return gs.Start() - }) - - httpServerConfig := server.Config{Port: defSvcHTTPPort} - if err := env.ParseWithOptions(&httpServerConfig, env.Options{Prefix: envPrefixHTTP}); err != nil { - logger.Error(fmt.Sprintf("failed to load %s HTTP server configuration : %s", svcName, err.Error())) - exitCode = 1 - return - } - mux := chi.NewMux() - idp := uuid.New() - hs := httpserver.NewServer(ctx, cancel, svcName, httpServerConfig, httpapi.MakeHandler(svc, authnMiddleware, mux, logger, cfg.InstanceID, idp), logger) - - g.Go(func() error { - return hs.Start() - }) - - if cfg.SendTelemetry { - chc := chclient.New(svcName, magistrala.Version, logger, cancel) - go chc.CallHome(ctx) - } - - g.Go(func() error { - return server.StopSignalHandler(ctx, cancel, logger, svcName, hs, gs) - }) - - if err := g.Wait(); err != nil { - logger.Error(fmt.Sprintf("domains service terminated: %s", err)) - } -} - -func newDomainService(ctx context.Context, domainsRepo domainsSvc.Repository, cache domainsSvc.Cache, tracer trace.Tracer, cfg config, authz authz.Authorization, policiessvc policies.Service, logger *slog.Logger, callout callout.Callout) (domains.Service, error) { - idProvider := uuid.New() - sidProvider, err := sid.New() - if err != nil { - return nil, fmt.Errorf("failed to init short id provider : %w", err) - } - - availableActions, builtInRoles, err := availableActionsAndBuiltInRoles(cfg.SpicedbSchemaFile) - if err != nil { - return nil, err - } - - permConfig, err := permissions.ParsePermissionsFile(cfg.PermissionsFile) - if err != nil { - return nil, fmt.Errorf("failed to parse permissions file: %w", err) - } - - svc, err := domainsSvc.New(domainsRepo, cache, policiessvc, idProvider, sidProvider, availableActions, builtInRoles) - if err != nil { - return nil, fmt.Errorf("failed to init domain service: %w", err) - } - svc, err = events.NewEventStoreMiddleware(ctx, svc, cfg.ESURL) - if err != nil { - return nil, fmt.Errorf("failed to init domain event store middleware: %w", err) - } - - domainOps, domainRoleOps, err := permConfig.GetEntityPermissions("domains") - if err != nil { - return nil, fmt.Errorf("failed to get domain permissions: %w", err) - } - - entitiesOps, err := permissions.NewEntitiesOperations( - permissions.EntitiesPermission{policies.DomainType: domainOps}, - permissions.EntitiesOperationDetails[permissions.Operation]{policies.DomainType: doperations.OperationDetails()}, - ) - if err != nil { - return nil, fmt.Errorf("failed to create entities operations: %w", err) - } - - roleOps, err := permissions.NewOperations(roles.Operations(), domainRoleOps) - if err != nil { - return nil, fmt.Errorf("failed to create role operations: %w", err) - } - - svc, err = dmw.NewAuthorization(policies.DomainType, svc, authz, entitiesOps, roleOps) - if err != nil { - return nil, err - } - - svc, err = dmw.NewCallout(svc, entitiesOps, roleOps, callout) - if err != nil { - return nil, err - } - - counter, latency := prometheus.MakeMetrics("domains", "api") - svc = dmw.NewMetrics(svc, counter, latency) - - svc = dmw.NewLogging(svc, logger) - - svc = dmw.NewTracing(svc, tracer) - return svc, nil -} - -func newPolicyService(cfg config, logger *slog.Logger) (policies.Service, error) { - client, err := authzed.NewClientWithExperimentalAPIs( - fmt.Sprintf("%s:%s", cfg.SpicedbHost, cfg.SpicedbPort), - grpc.WithTransportCredentials(insecure.NewCredentials()), - grpcutil.WithInsecureBearerToken(cfg.SpicedbPreSharedKey), - ) - if err != nil { - return nil, err - } - policySvc := spicedb.NewPolicyService(client, logger) - - return policySvc, nil -} - -func availableActionsAndBuiltInRoles(spicedbSchemaFile string) ([]roles.Action, map[roles.BuiltInRoleName][]roles.Action, error) { - availableActions, err := spicedbdecoder.GetActionsFromSchema(spicedbSchemaFile, policies.DomainType) - if err != nil { - return []roles.Action{}, map[roles.BuiltInRoleName][]roles.Action{}, err - } - - builtInRoles := map[roles.BuiltInRoleName][]roles.Action{ - domains.BuiltInRoleAdmin: availableActions, - } - - return availableActions, builtInRoles, err -} diff --git a/cmd/fluxmq/main.go b/cmd/fluxmq/main.go index 16a4786b7..f5aa53f20 100644 --- a/cmd/fluxmq/main.go +++ b/cmd/fluxmq/main.go @@ -9,6 +9,7 @@ package main import ( "context" + "errors" "fmt" "log" "net/http" @@ -21,12 +22,15 @@ import ( "connectrpc.com/otelconnect" "github.com/absmach/fluxmq/pkg/proto/auth/v1/authv1connect" fluxmqgrpc "github.com/absmach/magistrala/fluxmq/api/grpc" + fluxmqhttp "github.com/absmach/magistrala/fluxmq/api/http" + "github.com/absmach/magistrala/internal/atom" mglog "github.com/absmach/magistrala/logger" - domainsAuthz "github.com/absmach/magistrala/pkg/domains/grpcclient" - "github.com/absmach/magistrala/pkg/grpcclient" + atomauthn "github.com/absmach/magistrala/pkg/authn/atom" jaegerclient "github.com/absmach/magistrala/pkg/jaeger" "github.com/absmach/magistrala/pkg/messaging" + fluxmqbroker "github.com/absmach/magistrala/pkg/messaging/fluxmq" "github.com/absmach/magistrala/pkg/server" + httpserver "github.com/absmach/magistrala/pkg/server/http" "github.com/absmach/magistrala/pkg/uuid" "github.com/caarlos0/env/v11" "golang.org/x/net/http2" @@ -35,22 +39,55 @@ import ( ) const ( - svcName = "fluxmq-auth" - defSvcGRPCPort = "7016" - envPrefixClients = "MG_CLIENTS_GRPC_" - envPrefixChannels = "MG_CHANNELS_GRPC_" - envPrefixDomains = "MG_DOMAINS_GRPC_" - envPrefixCache = "MG_FLUXMQ_CACHE_" - envPrefixGRPC = "MG_FLUXMQ_GRPC_" + svcName = "fluxmq-auth" + defSvcGRPCPort = "7016" + envPrefixCache = "MG_FLUXMQ_CACHE_" + envPrefixGRPC = "MG_FLUXMQ_GRPC_" + envPrefixHTTP = "MG_FLUXMQ_PUBLISH_HTTP_" ) type config struct { LogLevel string `env:"MG_FLUXMQ_LOG_LEVEL" envDefault:"info"` + BrokerURL string `env:"MG_MESSAGE_BROKER_URL" envDefault:"amqp://guest:guest@localhost:5682/"` JaegerURL url.URL `env:"MG_JAEGER_URL" envDefault:"http://localhost:4318/v1/traces"` TraceRatio float64 `env:"MG_JAEGER_TRACE_RATIO" envDefault:"1.0"` InstanceID string `env:"MG_FLUXMQ_INSTANCE_ID" envDefault:""` } +type fanoutPublisher struct { + publishers []messaging.Publisher +} + +func (fp fanoutPublisher) Publish(ctx context.Context, topic string, msg *messaging.Message) error { + for _, publisher := range fp.publishers { + if err := publisher.Publish(ctx, topic, msg); err != nil { + return err + } + } + return nil +} + +func (fp fanoutPublisher) Close() error { + errs := make([]error, 0, len(fp.publishers)) + for _, publisher := range fp.publishers { + errs = append(errs, publisher.Close()) + } + return errors.Join(errs...) +} + +type writerBridgeHandler struct { + ctx context.Context + publisher messaging.Publisher +} + +func (h writerBridgeHandler) Handle(msg *messaging.Message) error { + return h.publisher.Publish(h.ctx, messaging.EncodeMessageTopic(msg), msg) +} + +func (h writerBridgeHandler) Cancel() error { + return nil +} + func main() { ctx, cancel := context.WithCancel(context.Background()) g, ctx := errgroup.WithContext(ctx) @@ -88,53 +125,18 @@ func main() { } }() - // Connect to Domains gRPC service (needed for topic route resolution). - domsGrpcCfg := grpcclient.Config{} - if err := env.ParseWithOptions(&domsGrpcCfg, env.Options{Prefix: envPrefixDomains}); err != nil { - logger.Error(fmt.Sprintf("failed to load domains gRPC client configuration: %s", err)) + atomCfg := atom.LoadConfig() + if atomCfg.URL == "" { + logger.Error("ATOM_URL is required") exitCode = 1 return } - _, domainsClient, domainsHandler, err := domainsAuthz.NewAuthorization(ctx, domsGrpcCfg) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer domainsHandler.Close() - logger.Info("Domains gRPC client connected " + domainsHandler.Secure()) - - // Connect to Clients gRPC service (authentication). - clientsClientCfg := grpcclient.Config{} - if err := env.ParseWithOptions(&clientsClientCfg, env.Options{Prefix: envPrefixClients}); err != nil { - logger.Error(fmt.Sprintf("failed to load clients gRPC client configuration: %s", err)) - exitCode = 1 - return - } - clientsClient, clientsHandler, err := grpcclient.SetupClientsClient(ctx, clientsClientCfg) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer clientsHandler.Close() - logger.Info("Clients gRPC client connected " + clientsHandler.Secure()) - - // Connect to Channels gRPC service (authorization + route resolution). - channelsClientCfg := grpcclient.Config{} - if err := env.ParseWithOptions(&channelsClientCfg, env.Options{Prefix: envPrefixChannels}); err != nil { - logger.Error(fmt.Sprintf("failed to load channels gRPC client configuration: %s", err)) - exitCode = 1 - return - } - channelsClient, channelsHandler, err := grpcclient.SetupChannelsClient(ctx, channelsClientCfg) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer channelsHandler.Close() - logger.Info("Channels gRPC client connected " + channelsHandler.Secure()) + atomAuthz := atom.NewClient(atomCfg) + authn := atomauthn.NewAuthentication() + clientsClient := atom.NewClientsCompat(authn, atomAuthz) + domainsClient := atom.NewDomainsCompat(atomAuthz) + channelsClient := atom.NewChannelsCompat(atomAuthz) + logger.Info("FluxMQ authentication, authorization, and route resolution configured to use Atom") // Topic parser with cache for route resolution. cacheConfig := messaging.CacheConfig{} @@ -166,7 +168,7 @@ func main() { return } path, handler := authv1connect.NewAuthServiceHandler( - fluxmqgrpc.NewServer(clientsClient, channelsClient, parser), + fluxmqgrpc.NewServer(clientsClient, channelsClient, parser, atomAuthz), connect.WithInterceptors(otelInterceptor), ) mux.Handle(path, handler) @@ -186,6 +188,72 @@ func main() { MaxHeaderBytes: grpcServerConfig.MaxHeaderBytes, } + messagePublisher, err := fluxmqbroker.NewUndeclaredPublisher( + ctx, + cfg.BrokerURL, + fluxmqbroker.ConnectionName("fluxmq-ui-message-publish-proxy"), + ) + if err != nil { + logger.Error(fmt.Sprintf("failed to create publish proxy message publisher: %s", err)) + exitCode = 1 + return + } + defer messagePublisher.Close() + + writerPublisher, err := fluxmqbroker.NewUndeclaredPublisher( + ctx, + cfg.BrokerURL, + fluxmqbroker.Prefix("writers"), + fluxmqbroker.ConnectionName("fluxmq-ui-publish-proxy"), + ) + if err != nil { + logger.Error(fmt.Sprintf("failed to create publish proxy writer publisher: %s", err)) + exitCode = 1 + return + } + defer writerPublisher.Close() + publisher := fanoutPublisher{publishers: []messaging.Publisher{messagePublisher, writerPublisher}} + + writerBridge, err := fluxmqbroker.NewPubSub( + ctx, + cfg.BrokerURL, + logger, + fluxmqbroker.DirectTopicOnly(), + fluxmqbroker.ConnectionName("fluxmq-mqtt-writer-bridge"), + ) + if err != nil { + logger.Error(fmt.Sprintf("failed to create MQTT writer bridge subscriber: %s", err)) + exitCode = 1 + return + } + defer writerBridge.Close() + if err := writerBridge.Subscribe(ctx, messaging.SubscriberConfig{ + ID: cfg.InstanceID + "-mqtt-writer-bridge", + Topic: "m/#", + Handler: writerBridgeHandler{ctx: ctx, publisher: writerPublisher}, + DeliveryPolicy: messaging.DeliverNewPolicy, + }); err != nil { + logger.Error(fmt.Sprintf("failed to subscribe MQTT writer bridge: %s", err)) + exitCode = 1 + return + } + logger.Info("FluxMQ MQTT writer bridge subscribed", "topic", "m/#") + + httpServerConfig := server.Config{Port: "9026"} + if err := env.ParseWithOptions(&httpServerConfig, env.Options{Prefix: envPrefixHTTP}); err != nil { + logger.Error(fmt.Sprintf("failed to load publish proxy HTTP server configuration: %s", err)) + exitCode = 1 + return + } + hs := httpserver.NewServer( + ctx, + cancel, + "fluxmq-publish", + httpServerConfig, + fluxmqhttp.MakePublishHandler(authn, atomAuthz, publisher), + logger, + ) + g.Go(func() error { logger.Info(fmt.Sprintf("%s service h2c server listening at %s", svcName, address)) var err error @@ -202,10 +270,17 @@ func main() { return nil }) + g.Go(func() error { + return hs.Start() + }) + g.Go(func() error { <-ctx.Done() shutdownCtx, shutdownCancel := context.WithTimeout(context.Background(), server.StopWaitTime) //nolint:contextcheck defer shutdownCancel() + if err := hs.Stop(); err != nil { + return fmt.Errorf("failed to shutdown publish proxy server: %w", err) + } if err := httpServer.Shutdown(shutdownCtx); err != nil { //nolint:contextcheck return fmt.Errorf("failed to shutdown %s server: %w", svcName, err) } diff --git a/cmd/groups/main.go b/cmd/groups/main.go deleted file mode 100644 index dcd85b2b8..000000000 --- a/cmd/groups/main.go +++ /dev/null @@ -1,439 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package main contains groups main function to start the groups service. -package main - -import ( - "context" - "fmt" - "log" - "log/slog" - "net/url" - "os" - - chclient "github.com/absmach/callhome/pkg/client" - "github.com/absmach/magistrala" - grpcChannelsV1 "github.com/absmach/magistrala/api/grpc/channels/v1" - grpcClientsV1 "github.com/absmach/magistrala/api/grpc/clients/v1" - grpcGroupsV1 "github.com/absmach/magistrala/api/grpc/groups/v1" - "github.com/absmach/magistrala/auth" - doperations "github.com/absmach/magistrala/domains/operations" - dpostgres "github.com/absmach/magistrala/domains/postgres" - "github.com/absmach/magistrala/groups" - gpsvc "github.com/absmach/magistrala/groups" - grpcapi "github.com/absmach/magistrala/groups/api/grpc" - httpapi "github.com/absmach/magistrala/groups/api/http" - "github.com/absmach/magistrala/groups/events" - "github.com/absmach/magistrala/groups/middleware" - goperations "github.com/absmach/magistrala/groups/operations" - "github.com/absmach/magistrala/groups/postgres" - pgroups "github.com/absmach/magistrala/groups/private" - mglog "github.com/absmach/magistrala/logger" - smqauthn "github.com/absmach/magistrala/pkg/authn" - authsvcAuthn "github.com/absmach/magistrala/pkg/authn/authsvc" - jwksAuthn "github.com/absmach/magistrala/pkg/authn/jwks" - smqauthz "github.com/absmach/magistrala/pkg/authz" - authsvcAuthz "github.com/absmach/magistrala/pkg/authz/authsvc" - "github.com/absmach/magistrala/pkg/callout" - dconsumer "github.com/absmach/magistrala/pkg/domains/events/consumer" - domainsAuthz "github.com/absmach/magistrala/pkg/domains/grpcclient" - "github.com/absmach/magistrala/pkg/grpcclient" - jaegerclient "github.com/absmach/magistrala/pkg/jaeger" - "github.com/absmach/magistrala/pkg/permissions" - "github.com/absmach/magistrala/pkg/policies" - "github.com/absmach/magistrala/pkg/policies/spicedb" - pg "github.com/absmach/magistrala/pkg/postgres" - pgclient "github.com/absmach/magistrala/pkg/postgres" - "github.com/absmach/magistrala/pkg/prometheus" - "github.com/absmach/magistrala/pkg/roles" - "github.com/absmach/magistrala/pkg/server" - grpcserver "github.com/absmach/magistrala/pkg/server/grpc" - httpserver "github.com/absmach/magistrala/pkg/server/http" - "github.com/absmach/magistrala/pkg/sid" - spicedbdecoder "github.com/absmach/magistrala/pkg/spicedb" - "github.com/absmach/magistrala/pkg/uuid" - "github.com/authzed/authzed-go/v1" - "github.com/authzed/grpcutil" - "github.com/caarlos0/env/v11" - "github.com/go-chi/chi/v5" - "github.com/jmoiron/sqlx" - "go.opentelemetry.io/otel/trace" - "golang.org/x/sync/errgroup" - "google.golang.org/grpc" - "google.golang.org/grpc/credentials/insecure" - "google.golang.org/grpc/reflection" -) - -const ( - svcName = "groups" - envPrefixDB = "MG_GROUPS_DB_" - envPrefixHTTP = "MG_GROUPS_HTTP_" - envPrefixgRPC = "MG_GROUPS_GRPC_" - envPrefixAuth = "MG_AUTH_GRPC_" - envPrefixDomains = "MG_DOMAINS_GRPC_" - envPrefixChannels = "MG_CHANNELS_GRPC_" - envPrefixClients = "MG_CLIENTS_GRPC_" - envPrefixGroupCallout = "MG_GROUPS_CALLOUT_" - defDB = "groups" - defSvcHTTPPort = "9004" - defSvcgRPCPort = "7004" -) - -type config struct { - LogLevel string `env:"MG_GROUPS_LOG_LEVEL" envDefault:"info"` - InstanceID string `env:"MG_GROUPS_INSTANCE_ID" envDefault:""` - JaegerURL url.URL `env:"MG_JAEGER_URL" envDefault:"http://localhost:4318/v1/traces"` - SendTelemetry bool `env:"MG_SEND_TELEMETRY" envDefault:"true"` - ESURL string `env:"MG_ES_URL" envDefault:"amqp://guest:guest@localhost:5682/"` - ESConsumerName string `env:"MG_GROUPS_EVENT_CONSUMER" envDefault:"groups"` - TraceRatio float64 `env:"MG_JAEGER_TRACE_RATIO" envDefault:"1.0"` - SpicedbHost string `env:"MG_SPICEDB_HOST" envDefault:"localhost"` - SpicedbPort string `env:"MG_SPICEDB_PORT" envDefault:"50051"` - SpicedbSchemaFile string `env:"MG_SPICEDB_SCHEMA_FILE" envDefault:"schema.zed"` - SpicedbPreSharedKey string `env:"MG_SPICEDB_PRE_SHARED_KEY" envDefault:"12345678"` - AuthKeyAlgorithm string `env:"MG_AUTH_KEYS_ALGORITHM" envDefault:"RS256"` - JWKSURL string `env:"MG_AUTH_JWKS_URL" envDefault:"http://auth:9001/keys/.well-known/jwks.json"` - PermissionsFile string `env:"MG_PERMISSIONS_FILE" envDefault:"permission.yaml"` -} - -func main() { - ctx, cancel := context.WithCancel(context.Background()) - g, ctx := errgroup.WithContext(ctx) - - cfg := config{} - if err := env.Parse(&cfg); err != nil { - log.Fatalf("failed to load %s configuration : %s", svcName, err.Error()) - } - - logger, err := mglog.New(os.Stdout, cfg.LogLevel) - if err != nil { - log.Fatalf("failed to init logger: %s", err.Error()) - } - - var exitCode int - defer mglog.ExitWithError(&exitCode) - - if cfg.InstanceID == "" { - if cfg.InstanceID, err = uuid.New().ID(); err != nil { - logger.Error(fmt.Sprintf("failed to generate instanceID: %s", err)) - exitCode = 1 - return - } - } - - dbConfig := pgclient.Config{Name: defDB} - if err := env.ParseWithOptions(&dbConfig, env.Options{Prefix: envPrefixDB}); err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - gm, err := postgres.Migration() - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - db, err := pgclient.Setup(dbConfig, *gm) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer db.Close() - - tp, err := jaegerclient.NewProvider(ctx, svcName, cfg.JaegerURL, cfg.InstanceID, cfg.TraceRatio) - if err != nil { - logger.Error(fmt.Sprintf("failed to init Jaeger: %s", err)) - exitCode = 1 - return - } - defer func() { - if err := tp.Shutdown(ctx); err != nil { - logger.Error(fmt.Sprintf("error shutting down tracer provider: %v", err)) - } - }() - tracer := tp.Tracer(svcName) - - authClientConfig := grpcclient.Config{} - if err := env.ParseWithOptions(&authClientConfig, env.Options{Prefix: envPrefixAuth}); err != nil { - logger.Error(fmt.Sprintf("failed to load %s auth configuration : %s", svcName, err)) - exitCode = 1 - return - } - - isSymmetric, err := auth.IsSymmetricAlgorithm(cfg.AuthKeyAlgorithm) - if err != nil { - logger.Error(fmt.Sprintf("failed to parse auth key algorithm : %s", err)) - exitCode = 1 - return - } - var authn smqauthn.Authentication - var authnClient grpcclient.Handler - switch { - case !isSymmetric: - authn, authnClient, err = jwksAuthn.NewAuthentication(ctx, cfg.JWKSURL, authClientConfig) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer authnClient.Close() - logger.Info("AuthN successfully set up jwks authentication on " + cfg.JWKSURL) - default: - authn, authnClient, err = authsvcAuthn.NewAuthentication(ctx, authClientConfig) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer authnClient.Close() - logger.Info("AuthN successfully connected to auth gRPC server " + authnClient.Secure()) - } - authnMiddleware := smqauthn.NewAuthNMiddleware(authn) - - domsGrpcCfg := grpcclient.Config{} - if err := env.ParseWithOptions(&domsGrpcCfg, env.Options{Prefix: envPrefixDomains}); err != nil { - logger.Error(fmt.Sprintf("failed to load domains gRPC client configuration : %s", err)) - exitCode = 1 - return - } - domAuthz, _, domainsHandler, err := domainsAuthz.NewAuthorization(ctx, domsGrpcCfg) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer domainsHandler.Close() - - callCfg := callout.Config{} - if err := env.ParseWithOptions(&callCfg, env.Options{Prefix: envPrefixGroupCallout}); err != nil { - logger.Error(fmt.Sprintf("failed to parse callout config : %s", err)) - exitCode = 1 - return - } - - authz, authzHandler, err := authsvcAuthz.NewAuthorization(ctx, authClientConfig, domAuthz) - if err != nil { - logger.Error("failed to create authz " + err.Error()) - exitCode = 1 - return - } - defer authzHandler.Close() - logger.Info("Authz successfully connected to auth gRPC server " + authzHandler.Secure()) - - policyService, err := newPolicyService(cfg, logger) - if err != nil { - logger.Error("failed to create new policies service " + err.Error()) - exitCode = 1 - return - } - logger.Info("Policy client successfully connected to spicedb gRPC server") - - chgrpcCfg := grpcclient.Config{} - if err := env.ParseWithOptions(&chgrpcCfg, env.Options{Prefix: envPrefixChannels}); err != nil { - logger.Error(fmt.Sprintf("failed to load channels gRPC client configuration : %s", err)) - exitCode = 1 - return - } - channelsClient, channelsHandler, err := grpcclient.SetupChannelsClient(ctx, chgrpcCfg) - if err != nil { - logger.Error(fmt.Sprintf("failed to connect to channels gRPC server: %s", err)) - exitCode = 1 - return - } - defer channelsHandler.Close() - logger.Info("Groups gRPC client successfully connected to channels gRPC server " + channelsHandler.Secure()) - - thgrpcCfg := grpcclient.Config{} - if err := env.ParseWithOptions(&thgrpcCfg, env.Options{Prefix: envPrefixClients}); err != nil { - logger.Error(fmt.Sprintf("failed to load clients gRPC client configuration : %s", err)) - exitCode = 1 - return - } - clientsClient, clientsHandler, err := grpcclient.SetupClientsClient(ctx, thgrpcCfg) - if err != nil { - logger.Error(fmt.Sprintf("failed to connect to clients gRPC server: %s", err)) - exitCode = 1 - return - } - defer clientsHandler.Close() - logger.Info("Clients gRPC client successfully connected to clients gRPC server " + clientsHandler.Secure()) - - callout, err := callout.New(callCfg) - if err != nil { - logger.Error(fmt.Sprintf("failed to create new callout: %s", err)) - exitCode = 1 - return - } - - permConfig, err := permissions.ParsePermissionsFile(cfg.PermissionsFile) - if err != nil { - logger.Error(fmt.Sprintf("failed to parse permissions file: %s", err)) - exitCode = 1 - return - } - - svc, psvc, err := newService(ctx, authz, policyService, db, dbConfig, channelsClient, clientsClient, tracer, logger, cfg, callout, permConfig) - if err != nil { - logger.Error(fmt.Sprintf("failed to setup service: %s", err)) - exitCode = 1 - return - } - - ddatabase := pg.NewDatabase(db, dbConfig, tracer) - drepo := dpostgres.NewRepository(ddatabase) - - if err := dconsumer.DomainsEventsSubscribe(ctx, drepo, cfg.ESURL, cfg.ESConsumerName, logger); err != nil { - logger.Error(fmt.Sprintf("failed to create domains event store : %s", err)) - exitCode = 1 - return - } - - httpServerConfig := server.Config{Port: defSvcHTTPPort} - if err := env.ParseWithOptions(&httpServerConfig, env.Options{Prefix: envPrefixHTTP}); err != nil { - logger.Error(fmt.Sprintf("failed to load %s HTTP server configuration : %s", svcName, err.Error())) - exitCode = 1 - return - } - - mux := chi.NewRouter() - idp := uuid.New() - httpSrv := httpserver.NewServer(ctx, cancel, svcName, httpServerConfig, httpapi.MakeHandler(svc, authnMiddleware, mux, logger, cfg.InstanceID, idp), logger) - - grpcServerConfig := server.Config{} - if err := env.ParseWithOptions(&grpcServerConfig, env.Options{Prefix: envPrefixgRPC}); err != nil { - logger.Error(fmt.Sprintf("failed to load %s gRPC server configuration : %s", svcName, err)) - exitCode = 1 - return - } - - registerGroupsServer := func(srv *grpc.Server) { - reflection.Register(srv) - grpcGroupsV1.RegisterGroupsServiceServer(srv, grpcapi.NewServer(psvc)) - } - gs := grpcserver.NewServer(ctx, cancel, svcName, grpcServerConfig, registerGroupsServer, logger) - - if cfg.SendTelemetry { - chc := chclient.New(svcName, magistrala.Version, logger, cancel) - go chc.CallHome(ctx) - } - - g.Go(func() error { - return gs.Start() - }) - - g.Go(func() error { - return httpSrv.Start() - }) - - g.Go(func() error { - return server.StopSignalHandler(ctx, cancel, logger, svcName, httpSrv) - }) - - if err := g.Wait(); err != nil { - logger.Error(fmt.Sprintf("groups service terminated: %s", err)) - } -} - -func newService(ctx context.Context, authz smqauthz.Authorization, policy policies.Service, db *sqlx.DB, - dbConfig pgclient.Config, channels grpcChannelsV1.ChannelsServiceClient, - clients grpcClientsV1.ClientsServiceClient, tracer trace.Tracer, logger *slog.Logger, c config, callout callout.Callout, permConfig *permissions.PermissionConfig, -) (groups.Service, pgroups.Service, error) { - database := pg.NewDatabase(db, dbConfig, tracer) - idp := uuid.New() - sid, err := sid.New() - if err != nil { - return nil, nil, err - } - - availableActions, builtInRoles, err := availableActionsAndBuiltInRoles(c.SpicedbSchemaFile) - if err != nil { - return nil, nil, err - } - - // Creating groups service - repo := postgres.New(database) - svc, err := gpsvc.NewService(repo, policy, idp, channels, clients, sid, availableActions, builtInRoles) - if err != nil { - return nil, nil, err - } - svc, err = events.New(ctx, svc, c.ESURL) - if err != nil { - return nil, nil, err - } - - groupOps, groupRoleOps, err := permConfig.GetEntityPermissions("groups") - if err != nil { - return nil, nil, fmt.Errorf("failed to get group permissions: %w", err) - } - - domainOps, _, err := permConfig.GetEntityPermissions("domains") - if err != nil { - return nil, nil, fmt.Errorf("failed to get domain permissions: %w", err) - } - - entitiesOps, err := permissions.NewEntitiesOperations( - permissions.EntitiesPermission{ - policies.GroupType: groupOps, - policies.DomainType: domainOps, - }, - permissions.EntitiesOperationDetails[permissions.Operation]{ - policies.GroupType: goperations.OperationDetails(), - policies.DomainType: doperations.OperationDetails(), - }, - ) - if err != nil { - return nil, nil, fmt.Errorf("failed to create entities operations: %w", err) - } - - roleOps, err := permissions.NewOperations(roles.Operations(), groupRoleOps) - if err != nil { - return nil, nil, fmt.Errorf("failed to create role operations: %w", err) - } - - svc, err = middleware.NewAuthorization(policies.GroupType, svc, authz, repo, entitiesOps, roleOps) - if err != nil { - return nil, nil, err - } - - svc, err = middleware.NewCallout(svc, repo, entitiesOps, roleOps, callout) - if err != nil { - return nil, nil, err - } - - svc = middleware.NewTracing(svc, tracer) - svc = middleware.NewLogging(svc, logger) - counter, latency := prometheus.MakeMetrics("groups", "api") - svc = middleware.NewMetrics(svc, counter, latency) - - psvc := pgroups.New(repo) - return svc, psvc, err -} - -func newPolicyService(cfg config, logger *slog.Logger) (policies.Service, error) { - client, err := authzed.NewClientWithExperimentalAPIs( - fmt.Sprintf("%s:%s", cfg.SpicedbHost, cfg.SpicedbPort), - grpc.WithTransportCredentials(insecure.NewCredentials()), - grpcutil.WithInsecureBearerToken(cfg.SpicedbPreSharedKey), - ) - if err != nil { - return nil, err - } - policySvc := spicedb.NewPolicyService(client, logger) - - return policySvc, nil -} - -func availableActionsAndBuiltInRoles(spicedbSchemaFile string) ([]roles.Action, map[roles.BuiltInRoleName][]roles.Action, error) { - availableActions, err := spicedbdecoder.GetActionsFromSchema(spicedbSchemaFile, policies.GroupType) - if err != nil { - return []roles.Action{}, map[roles.BuiltInRoleName][]roles.Action{}, err - } - - builtInRoles := map[roles.BuiltInRoleName][]roles.Action{ - groups.BuiltInRoleAdmin: availableActions, - } - - return availableActions, builtInRoles, err -} diff --git a/cmd/journal/main.go b/cmd/journal/main.go index 3f8502721..fad88edea 100644 --- a/cmd/journal/main.go +++ b/cmd/journal/main.go @@ -14,7 +14,7 @@ import ( chclient "github.com/absmach/callhome/pkg/client" "github.com/absmach/magistrala" - "github.com/absmach/magistrala/auth" + "github.com/absmach/magistrala/internal/atom" "github.com/absmach/magistrala/journal" httpapi "github.com/absmach/magistrala/journal/api" "github.com/absmach/magistrala/journal/events" @@ -22,13 +22,9 @@ import ( journalpg "github.com/absmach/magistrala/journal/postgres" mglog "github.com/absmach/magistrala/logger" smqauthn "github.com/absmach/magistrala/pkg/authn" - authsvcAuthn "github.com/absmach/magistrala/pkg/authn/authsvc" - jwksAuthn "github.com/absmach/magistrala/pkg/authn/jwks" + atomauthn "github.com/absmach/magistrala/pkg/authn/atom" smqauthz "github.com/absmach/magistrala/pkg/authz" - authsvcAuthz "github.com/absmach/magistrala/pkg/authz/authsvc" - domainsAuthz "github.com/absmach/magistrala/pkg/domains/grpcclient" "github.com/absmach/magistrala/pkg/events/store" - "github.com/absmach/magistrala/pkg/grpcclient" jaegerclient "github.com/absmach/magistrala/pkg/jaeger" "github.com/absmach/magistrala/pkg/postgres" pgclient "github.com/absmach/magistrala/pkg/postgres" @@ -43,24 +39,20 @@ import ( ) const ( - svcName = "journal" - envPrefixDB = "MG_JOURNAL_DB_" - envPrefixHTTP = "MG_JOURNAL_HTTP_" - envPrefixAuth = "MG_AUTH_GRPC_" - envPrefixDomains = "MG_DOMAINS_GRPC_" - defDB = "journal" - defSvcHTTPPort = "9021" + svcName = "journal" + envPrefixDB = "MG_JOURNAL_DB_" + envPrefixHTTP = "MG_JOURNAL_HTTP_" + defDB = "journal" + defSvcHTTPPort = "9021" ) type config struct { - LogLevel string `env:"MG_JOURNAL_LOG_LEVEL" envDefault:"info"` - ESURL string `env:"MG_ES_URL" envDefault:"amqp://guest:guest@localhost:5682/"` - JaegerURL url.URL `env:"MG_JAEGER_URL" envDefault:"http://localhost:4318/v1/traces"` - SendTelemetry bool `env:"MG_SEND_TELEMETRY" envDefault:"true"` - InstanceID string `env:"MG_JOURNAL_INSTANCE_ID" envDefault:""` - TraceRatio float64 `env:"MG_JAEGER_TRACE_RATIO" envDefault:"1.0"` - AuthKeyAlgorithm string `env:"MG_AUTH_KEYS_ALGORITHM" envDefault:"RS256"` - JWKSURL string `env:"MG_AUTH_JWKS_URL" envDefault:"http://auth:9001/keys/.well-known/jwks.json"` + LogLevel string `env:"MG_JOURNAL_LOG_LEVEL" envDefault:"info"` + ESURL string `env:"MG_ES_URL" envDefault:"amqp://guest:guest@localhost:5682/"` + JaegerURL url.URL `env:"MG_JAEGER_URL" envDefault:"http://localhost:4318/v1/traces"` + SendTelemetry bool `env:"MG_SEND_TELEMETRY" envDefault:"true"` + InstanceID string `env:"MG_JOURNAL_INSTANCE_ID" envDefault:""` + TraceRatio float64 `env:"MG_JAEGER_TRACE_RATIO" envDefault:"1.0"` } func main() { @@ -102,65 +94,17 @@ func main() { } defer db.Close() - authClientCfg := grpcclient.Config{} - if err := env.ParseWithOptions(&authClientCfg, env.Options{Prefix: envPrefixAuth}); err != nil { - logger.Error(fmt.Sprintf("failed to load auth gRPC client configuration : %s", err)) + atomCfg := atom.LoadConfig() + if atomCfg.URL == "" { + logger.Error("ATOM_URL is required") exitCode = 1 return } - - isSymmetric, err := auth.IsSymmetricAlgorithm(cfg.AuthKeyAlgorithm) - if err != nil { - logger.Error(fmt.Sprintf("failed to parse auth key algorithm : %s", err)) - exitCode = 1 - return - } - var authn smqauthn.Authentication - var authnClient grpcclient.Handler - switch { - case !isSymmetric: - authn, authnClient, err = jwksAuthn.NewAuthentication(ctx, cfg.JWKSURL, authClientCfg) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer authnClient.Close() - logger.Info("AuthN successfully set up jwks authentication on " + cfg.JWKSURL) - default: - authn, authnClient, err = authsvcAuthn.NewAuthentication(ctx, authClientCfg) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer authnClient.Close() - logger.Info("AuthN successfully connected to auth gRPC server " + authnClient.Secure()) - } + atomClient := atom.NewClient(atomCfg) + authn := atomauthn.NewAuthentication() authnMiddleware := smqauthn.NewAuthNMiddleware(authn) - - domsGrpcCfg := grpcclient.Config{} - if err := env.ParseWithOptions(&domsGrpcCfg, env.Options{Prefix: envPrefixDomains}); err != nil { - logger.Error(fmt.Sprintf("failed to load domains gRPC client configuration : %s", err)) - exitCode = 1 - return - } - domAuthz, _, domainsHandler, err := domainsAuthz.NewAuthorization(ctx, domsGrpcCfg) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer domainsHandler.Close() - - authz, authzHandler, err := authsvcAuthz.NewAuthorization(ctx, authClientCfg, domAuthz) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer authzHandler.Close() - logger.Info("AuthZ successfully connected to auth gRPC server " + authzHandler.Secure()) + authz := atom.NewAuthorizationCompat(atomClient) + logger.Info("AuthN/AuthZ configured to use Atom") tp, err := jaegerclient.NewProvider(ctx, svcName, cfg.JaegerURL, cfg.InstanceID, cfg.TraceRatio) if err != nil { diff --git a/cmd/notifications/main.go b/cmd/notifications/main.go index 8ab4c16c2..7e3128953 100644 --- a/cmd/notifications/main.go +++ b/cmd/notifications/main.go @@ -13,12 +13,12 @@ import ( chclient "github.com/absmach/callhome/pkg/client" "github.com/absmach/magistrala" + "github.com/absmach/magistrala/internal/atom" mglog "github.com/absmach/magistrala/logger" "github.com/absmach/magistrala/notifications/emailer" "github.com/absmach/magistrala/notifications/events" "github.com/absmach/magistrala/notifications/middleware" "github.com/absmach/magistrala/pkg/events/store" - "github.com/absmach/magistrala/pkg/grpcclient" jaegerclient "github.com/absmach/magistrala/pkg/jaeger" "github.com/absmach/magistrala/pkg/prometheus" "github.com/absmach/magistrala/pkg/server" @@ -28,9 +28,8 @@ import ( ) const ( - svcName = "notifications" - envPrefixUsers = "MG_USERS_GRPC_" - defEmailPort = "25" + svcName = "notifications" + defEmailPort = "25" ) type config struct { @@ -89,21 +88,14 @@ func main() { } }() - usersClientCfg := grpcclient.Config{} - if err := env.ParseWithOptions(&usersClientCfg, env.Options{Prefix: envPrefixUsers}); err != nil { - logger.Error(fmt.Sprintf("failed to load users gRPC client configuration : %s", err)) + atomCfg := atom.LoadConfig() + if atomCfg.URL == "" { + logger.Error("ATOM_URL is required") exitCode = 1 return } - - usersClient, usersHandler, err := grpcclient.SetupUsersClient(ctx, usersClientCfg) - if err != nil { - logger.Error(fmt.Sprintf("failed to setup users gRPC client: %s", err)) - exitCode = 1 - return - } - defer usersHandler.Close() - logger.Info("Successfully connected to users gRPC server " + usersHandler.Secure()) + usersResolver := emailer.NewAtomUserResolver(atom.NewClient(atomCfg)) + logger.Info("Notifications user lookup configured to use Atom") emailerCfg := emailer.Config{ FromAddress: cfg.EmailFromAddress, @@ -118,7 +110,7 @@ func main() { EmailPassword: cfg.EmailPassword, } - notifier, err := emailer.New(usersClient, emailerCfg) + notifier, err := emailer.New(usersResolver, emailerCfg) if err != nil { logger.Error(fmt.Sprintf("failed to create emailer: %s", err)) exitCode = 1 diff --git a/cmd/postgres-reader/main.go b/cmd/postgres-reader/main.go index 327c6e177..0827bf51f 100644 --- a/cmd/postgres-reader/main.go +++ b/cmd/postgres-reader/main.go @@ -14,9 +14,9 @@ import ( chclient "github.com/absmach/callhome/pkg/client" "github.com/absmach/magistrala" grpcReadersV1 "github.com/absmach/magistrala/api/grpc/readers/v1" + "github.com/absmach/magistrala/internal/atom" mglog "github.com/absmach/magistrala/logger" - "github.com/absmach/magistrala/pkg/authn/authsvc" - "github.com/absmach/magistrala/pkg/grpcclient" + atomauthn "github.com/absmach/magistrala/pkg/authn/atom" pgclient "github.com/absmach/magistrala/pkg/postgres" "github.com/absmach/magistrala/pkg/prometheus" "github.com/absmach/magistrala/pkg/server" @@ -36,16 +36,13 @@ import ( ) const ( - svcName = "postgres-reader" - envPrefixDB = "MG_POSTGRES_" - envPrefixHTTP = "MG_POSTGRES_READER_HTTP_" - envPrefixAuth = "MG_AUTH_GRPC_" - envPrefixClients = "MG_CLIENTS_GRPC_" - envPrefixChannels = "MG_CHANNELS_GRPC_" - defDB = "magistrala" - defSvcHTTPPort = "9009" - defSvcGRPCPort = "7009" - envPrefixGrpc = "MG_POSTGRES_READER_GRPC_" + svcName = "postgres-reader" + envPrefixDB = "MG_POSTGRES_" + envPrefixHTTP = "MG_POSTGRES_READER_HTTP_" + defDB = "magistrala" + defSvcHTTPPort = "9009" + defSvcGRPCPort = "7009" + envPrefixGrpc = "MG_POSTGRES_READER_GRPC_" ) type config struct { @@ -106,54 +103,11 @@ func main() { grpcReadersV1.RegisterReadersServiceServer(srv, readersgrpcapi.NewReadersServer(repo)) } - clientsClientCfg := grpcclient.Config{} - if err := env.ParseWithOptions(&clientsClientCfg, env.Options{Prefix: envPrefixClients}); err != nil { - logger.Error(fmt.Sprintf("failed to load clients gRPC client configuration : %s", err)) - exitCode = 1 - return - } - - clientsClient, clientsHandler, err := grpcclient.SetupClientsClient(ctx, clientsClientCfg) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer clientsHandler.Close() - - logger.Info("Clients service gRPC client successfully connected to clients gRPC server " + clientsHandler.Secure()) - - channelsClientCfg := grpcclient.Config{} - if err := env.ParseWithOptions(&channelsClientCfg, env.Options{Prefix: envPrefixChannels}); err != nil { - logger.Error(fmt.Sprintf("failed to load channels gRPC client configuration : %s", err)) - exitCode = 1 - return - } - - channelsClient, channelsHandler, err := grpcclient.SetupChannelsClient(ctx, channelsClientCfg) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer channelsHandler.Close() - logger.Info("Channels service gRPC client successfully connected to channels gRPC server " + channelsHandler.Secure()) - - authnCfg := grpcclient.Config{} - if err := env.ParseWithOptions(&authnCfg, env.Options{Prefix: envPrefixAuth}); err != nil { - logger.Error(fmt.Sprintf("failed to load auth gRPC client configuration : %s", err)) - exitCode = 1 - return - } - - authn, authnHandler, err := authsvc.NewAuthentication(ctx, authnCfg) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer authnHandler.Close() - logger.Info("authn successfully connected to auth gRPC server " + authnHandler.Secure()) + atomCfg := atom.LoadConfig() + authn := atomauthn.NewAuthentication() + clientsClient := atom.NewClientsCompat(authn) + channelsClient := atom.NewChannelsCompat(atom.NewClient(atomCfg)) + logger.Info("AuthN/AuthZ configured to use Atom") httpServerConfig := server.Config{Port: defSvcHTTPPort} if err := env.ParseWithOptions(&httpServerConfig, env.Options{Prefix: envPrefixHTTP}); err != nil { diff --git a/cmd/provision/main.go b/cmd/provision/main.go deleted file mode 100644 index 8aa9711a6..000000000 --- a/cmd/provision/main.go +++ /dev/null @@ -1,214 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package main contains provision main function to start the provision service. -package main - -import ( - "context" - "encoding/json" - "fmt" - "log" - "os" - "reflect" - - chclient "github.com/absmach/callhome/pkg/client" - "github.com/absmach/magistrala" - "github.com/absmach/magistrala/channels" - "github.com/absmach/magistrala/clients" - mglog "github.com/absmach/magistrala/logger" - smqauthn "github.com/absmach/magistrala/pkg/authn" - authnsvc "github.com/absmach/magistrala/pkg/authn/authsvc" - "github.com/absmach/magistrala/pkg/errors" - "github.com/absmach/magistrala/pkg/grpcclient" - mgsdk "github.com/absmach/magistrala/pkg/sdk" - "github.com/absmach/magistrala/pkg/server" - httpserver "github.com/absmach/magistrala/pkg/server/http" - "github.com/absmach/magistrala/pkg/uuid" - "github.com/absmach/magistrala/provision" - httpapi "github.com/absmach/magistrala/provision/api" - "github.com/absmach/magistrala/provision/middleware" - "github.com/caarlos0/env/v11" - "golang.org/x/sync/errgroup" -) - -const ( - svcName = "provision" - contentType = "application/json" - envPrefixAuth = "MG_AUTH_GRPC_" -) - -var ( - errMissingConfigFile = errors.New("missing config file setting") - errFailLoadingConfigFile = errors.New("failed to load config from file") - errFailedToReadBootstrapContent = errors.New("failed to read bootstrap content from envs") -) - -func main() { - ctx, cancel := context.WithCancel(context.Background()) - g, ctx := errgroup.WithContext(ctx) - - cfg, err := loadConfig() - if err != nil { - log.Fatalf("failed to load %s configuration : %s", svcName, err) - } - - logger, err := mglog.New(os.Stdout, cfg.Server.LogLevel) - if err != nil { - log.Fatalf("failed to init logger: %s", err.Error()) - } - - var exitCode int - defer mglog.ExitWithError(&exitCode) - - if cfg.InstanceID == "" { - if cfg.InstanceID, err = uuid.New().ID(); err != nil { - logger.Error(fmt.Sprintf("failed to generate instanceID: %s", err)) - exitCode = 1 - return - } - } - - grpcCfg := grpcclient.Config{} - if err := env.ParseWithOptions(&grpcCfg, env.Options{Prefix: envPrefixAuth}); err != nil { - logger.Error(fmt.Sprintf("failed to load auth gRPC client configuration : %s", err)) - exitCode = 1 - - return - } - authn, authnClient, err := authnsvc.NewAuthentication(ctx, grpcCfg) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - - return - } - defer authnClient.Close() - logger.Info("AuthN successfully connected to auth gRPC server " + authnClient.Secure()) - am := smqauthn.NewAuthNMiddleware(authn) - - if cfgFromFile, err := loadConfigFromFile(cfg.File); err != nil { - logger.Warn(fmt.Sprintf("Continue with settings from env, failed to load from: %s: %s", cfg.File, err)) - } else { - // Merge environment variables and file settings. - mergeConfigs(&cfgFromFile, &cfg) - cfg = cfgFromFile - logger.Info("Continue with settings from file: " + cfg.File) - } - - SDKCfg := mgsdk.Config{ - UsersURL: cfg.Server.UsersURL, - ChannelsURL: cfg.Server.ChannelsURL, - ClientsURL: cfg.Server.ClientsURL, - BootstrapURL: cfg.Server.MgBSURL, - CertsURL: cfg.Server.CertsURL, - MsgContentType: contentType, - TLSVerification: cfg.Server.TLS, - } - mgSdk := mgsdk.NewSDK(SDKCfg) - - svc := provision.New(cfg, mgSdk, logger) - svc = middleware.NewLogging(svc, logger) - - httpServerConfig := server.Config{Host: "", Port: cfg.Server.Port, KeyFile: cfg.Server.ServerKey, CertFile: cfg.Server.ServerCert} - hs := httpserver.NewServer(ctx, cancel, svcName, httpServerConfig, httpapi.MakeHandler(svc, am, logger, cfg.InstanceID), logger) - - if cfg.SendTelemetry { - chc := chclient.New(svcName, magistrala.Version, logger, cancel) - go chc.CallHome(ctx) - } - - g.Go(func() error { - return hs.Start() - }) - - g.Go(func() error { - return server.StopSignalHandler(ctx, cancel, logger, svcName, hs) - }) - - if err := g.Wait(); err != nil { - logger.Error(fmt.Sprintf("Provision service terminated: %s", err)) - } -} - -func loadConfigFromFile(file string) (provision.Config, error) { - _, err := os.Stat(file) - if os.IsNotExist(err) { - return provision.Config{}, errors.Wrap(errMissingConfigFile, err) - } - c, err := provision.Read(file) - if err != nil { - return provision.Config{}, errors.Wrap(errFailLoadingConfigFile, err) - } - return c, nil -} - -func loadConfig() (provision.Config, error) { - cfg := provision.Config{} - if err := env.Parse(&cfg); err != nil { - return provision.Config{}, err - } - - if cfg.Bootstrap.AutoWhiteList && !cfg.Bootstrap.Provision { - return provision.Config{}, errors.New("Can't auto whitelist if auto config save is off") - } - - var content map[string]any - if cfg.BSContent != "" { - if err := json.Unmarshal([]byte(cfg.BSContent), &content); err != nil { - return provision.Config{}, errFailedToReadBootstrapContent - } - } - - cfg.Bootstrap.Content = content - // This is default conf for provision if there is no config file - cfg.Channels = []channels.Channel{ - { - Name: "control-channel", - Metadata: map[string]any{"type": "control"}, - }, { - Name: "data-channel", - Metadata: map[string]any{"type": "data"}, - }, - } - cfg.Clients = []clients.Client{ - { - Name: "client", - Metadata: map[string]any{"external_id": "xxxxxx"}, - }, - } - - return cfg, nil -} - -func mergeConfigs(dst, src any) any { - d := reflect.ValueOf(dst).Elem() - s := reflect.ValueOf(src).Elem() - - for i := 0; i < d.NumField(); i++ { - dField := d.Field(i) - sField := s.Field(i) - switch dField.Kind() { - case reflect.Struct: - dst := dField.Addr().Interface() - src := sField.Addr().Interface() - m := mergeConfigs(dst, src) - val := reflect.ValueOf(m).Elem().Interface() - dField.Set(reflect.ValueOf(val)) - case reflect.Slice: - case reflect.Bool: - if dField.Interface() == false { - dField.Set(reflect.ValueOf(sField.Interface())) - } - case reflect.Int: - if dField.Interface() == 0 { - dField.Set(reflect.ValueOf(sField.Interface())) - } - case reflect.String: - if dField.Interface() == "" { - dField.Set(reflect.ValueOf(sField.Interface())) - } - } - } - return dst -} diff --git a/cmd/re/main.go b/cmd/re/main.go index 587cda66f..09caf3691 100644 --- a/cmd/re/main.go +++ b/cmd/re/main.go @@ -18,16 +18,12 @@ import ( abrokers "github.com/absmach/magistrala/alarms/brokers" grpcReadersV1 "github.com/absmach/magistrala/api/grpc/readers/v1" "github.com/absmach/magistrala/consumers/writers/brokers" - dpostgres "github.com/absmach/magistrala/domains/postgres" + "github.com/absmach/magistrala/internal/atom" "github.com/absmach/magistrala/internal/email" mglog "github.com/absmach/magistrala/logger" smqauthn "github.com/absmach/magistrala/pkg/authn" - authnsvc "github.com/absmach/magistrala/pkg/authn/authsvc" - mgauthz "github.com/absmach/magistrala/pkg/authz" - authzsvc "github.com/absmach/magistrala/pkg/authz/authsvc" + atomauthn "github.com/absmach/magistrala/pkg/authn/atom" "github.com/absmach/magistrala/pkg/callout" - dconsumer "github.com/absmach/magistrala/pkg/domains/events/consumer" - domainsAuthz "github.com/absmach/magistrala/pkg/domains/grpcclient" "github.com/absmach/magistrala/pkg/emailer" "github.com/absmach/magistrala/pkg/grpcclient" jaegerclient "github.com/absmach/magistrala/pkg/jaeger" @@ -36,14 +32,10 @@ import ( smqbrokers "github.com/absmach/magistrala/pkg/messaging/brokers" brokerstracing "github.com/absmach/magistrala/pkg/messaging/brokers/tracing" "github.com/absmach/magistrala/pkg/permissions" - "github.com/absmach/magistrala/pkg/policies" - "github.com/absmach/magistrala/pkg/policies/spicedb" pgclient "github.com/absmach/magistrala/pkg/postgres" "github.com/absmach/magistrala/pkg/prometheus" - "github.com/absmach/magistrala/pkg/roles" "github.com/absmach/magistrala/pkg/server" httpserver "github.com/absmach/magistrala/pkg/server/http" - spicedbdecoder "github.com/absmach/magistrala/pkg/spicedb" "github.com/absmach/magistrala/pkg/ticker" "github.com/absmach/magistrala/pkg/uuid" "github.com/absmach/magistrala/re" @@ -53,14 +45,10 @@ import ( "github.com/absmach/magistrala/re/operations" repg "github.com/absmach/magistrala/re/postgres" grpcClient "github.com/absmach/magistrala/readers/api/grpc" - "github.com/authzed/authzed-go/v1" - "github.com/authzed/grpcutil" "github.com/caarlos0/env/v11" "github.com/go-chi/chi/v5" "go.opentelemetry.io/otel/trace" "golang.org/x/sync/errgroup" - "google.golang.org/grpc" - "google.golang.org/grpc/credentials/insecure" ) const ( @@ -68,11 +56,9 @@ const ( envPrefixDB = "MG_RE_DB_" envPrefixHTTP = "MG_RE_HTTP_" envPrefixCallout = "MG_RE_CALLOUT_" - envPrefixAuth = "MG_AUTH_GRPC_" defDB = "r" defSvcHTTPPort = "9008" envPrefixGrpc = "MG_TIMESCALE_READER_GRPC_" - envPrefixDomains = "MG_DOMAINS_GRPC_" ) // We use a buffered channel to prevent blocking, as logging is an expensive operation. @@ -81,21 +67,17 @@ const ( const channBuffer = 256 type config struct { - LogLevel string `env:"MG_RE_LOG_LEVEL" envDefault:"info"` - InstanceID string `env:"MG_RE_INSTANCE_ID" envDefault:""` - JaegerURL url.URL `env:"MG_JAEGER_URL" envDefault:"http://localhost:4318/v1/traces"` - SendTelemetry bool `env:"MG_SEND_TELEMETRY" envDefault:"true"` - ESURL string `env:"MG_ES_URL" envDefault:"nats://localhost:4222"` - ESConsumerName string `env:"MG_RE_EVENT_CONSUMER" envDefault:"rules_engine"` - CacheURL string `env:"MG_RE_CACHE_URL" envDefault:"redis://localhost:6379/0"` - CacheKeyDuration time.Duration `env:"MG_RE_CACHE_KEY_DURATION" envDefault:"10m"` - TraceRatio float64 `env:"MG_JAEGER_TRACE_RATIO" envDefault:"1.0"` - BrokerURL string `env:"MG_MESSAGE_BROKER_URL" envDefault:"nats://localhost:4222"` - SpicedbHost string `env:"MG_SPICEDB_HOST" envDefault:"localhost"` - SpicedbPort string `env:"MG_SPICEDB_PORT" envDefault:"50051"` - SpicedbPreSharedKey string `env:"MG_SPICEDB_PRE_SHARED_KEY" envDefault:"12345678"` - SpicedbSchemaFile string `env:"MG_SPICEDB_SCHEMA_FILE" envDefault:"schema.zed"` - PermissionsFile string `env:"MG_PERMISSIONS_FILE" envDefault:"permission.yaml"` + LogLevel string `env:"MG_RE_LOG_LEVEL" envDefault:"info"` + InstanceID string `env:"MG_RE_INSTANCE_ID" envDefault:""` + JaegerURL url.URL `env:"MG_JAEGER_URL" envDefault:"http://localhost:4318/v1/traces"` + SendTelemetry bool `env:"MG_SEND_TELEMETRY" envDefault:"true"` + ESURL string `env:"MG_ES_URL" envDefault:"nats://localhost:4222"` + ESConsumerName string `env:"MG_RE_EVENT_CONSUMER" envDefault:"rules_engine"` + CacheURL string `env:"MG_RE_CACHE_URL" envDefault:"redis://localhost:6379/0"` + CacheKeyDuration time.Duration `env:"MG_RE_CACHE_KEY_DURATION" envDefault:"10m"` + TraceRatio float64 `env:"MG_JAEGER_TRACE_RATIO" envDefault:"1.0"` + BrokerURL string `env:"MG_MESSAGE_BROKER_URL" envDefault:"nats://localhost:4222"` + PermissionsFile string `env:"MG_PERMISSIONS_FILE" envDefault:"permission.yaml"` } func main() { @@ -221,60 +203,20 @@ func main() { defer alarmsPub.Close() alarmsPub = brokerstracing.NewPublisher(httpServerConfig, tracer, alarmsPub) - grpcCfg := grpcclient.Config{} - if err := env.ParseWithOptions(&grpcCfg, env.Options{Prefix: envPrefixAuth}); err != nil { - logger.Error(fmt.Sprintf("failed to load auth gRPC client configuration : %s", err)) + atomCfg := atom.LoadConfig() + if atomCfg.URL == "" { + logger.Error("ATOM_URL is required") exitCode = 1 - return } - authn, authnClient, err := authnsvc.NewAuthentication(ctx, grpcCfg) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - - return - } - am := smqauthn.NewAuthNMiddleware(authn) - - defer authnClient.Close() - logger.Info("AuthN successfully connected to auth gRPC server " + authnClient.Secure()) + authnSvc := atomauthn.NewAuthentication() + logger.Info("AuthN configured to use Atom bearer tokens") + am := smqauthn.NewAuthNMiddleware(authnSvc) runInfo := make(chan pkglog.RunInfo, channBuffer) - domsGrpcCfg := grpcclient.Config{} - if err := env.ParseWithOptions(&domsGrpcCfg, env.Options{Prefix: envPrefixDomains}); err != nil { - logger.Error(fmt.Sprintf("failed to load domains gRPC client configuration : %s", err)) - exitCode = 1 - return - } - domAuthz, _, domainsHandler, err := domainsAuthz.NewAuthorization(ctx, domsGrpcCfg) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer domainsHandler.Close() - - authz, authzClient, err := authzsvc.NewAuthorization(ctx, grpcCfg, domAuthz) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer authzClient.Close() - logger.Info("AuthZ successfully connected to auth gRPC server " + authnClient.Secure()) + logger.Info("AuthZ configured to use Atom PDP") database := pgclient.NewDatabase(db, dbConfig, tracer) - - ddatabase := pgclient.NewDatabase(db, dbConfig, tracer) - drepo := dpostgres.NewRepository(ddatabase) - - if err := dconsumer.DomainsEventsSubscribe(ctx, drepo, cfg.ESURL, cfg.ESConsumerName, logger); err != nil { - logger.Error(fmt.Sprintf("failed to create domains event store : %s", err)) - exitCode = 1 - return - } - regrpcCfg := grpcclient.Config{} if err := env.ParseWithOptions(®rpcCfg, env.Options{Prefix: envPrefixGrpc}); err != nil { logger.Error(fmt.Sprintf("failed to load clients gRPC client configuration : %s", err)) @@ -292,7 +234,7 @@ func main() { readersClient := grpcClient.NewReadersClient(client.Connection(), regrpcCfg.Timeout) logger.Info("Readers gRPC client successfully connected to readers gRPC server " + client.Secure()) - svc, err := newService(ctx, cfg, database, runInfo, msgSub, writersPub, alarmsPub, authz, ec, logger, readersClient, callout, tracer) + svc, err := newService(ctx, cfg, database, runInfo, msgSub, writersPub, alarmsPub, ec, logger, readersClient, callout, tracer) if err != nil { logger.Error(fmt.Sprintf("failed to create services: %s", err)) exitCode = 1 @@ -344,7 +286,7 @@ func main() { } } -func newService(ctx context.Context, cfg config, db pgclient.Database, runInfo chan pkglog.RunInfo, rePubSub messaging.PubSub, writersPub, alarmsPub messaging.Publisher, authz mgauthz.Authorization, ec email.Config, logger *slog.Logger, readersClient grpcReadersV1.ReadersServiceClient, callout callout.Callout, tracer trace.Tracer) (re.Service, error) { +func newService(ctx context.Context, cfg config, db pgclient.Database, runInfo chan pkglog.RunInfo, rePubSub messaging.PubSub, writersPub, alarmsPub messaging.Publisher, ec email.Config, logger *slog.Logger, readersClient grpcReadersV1.ReadersServiceClient, callout callout.Callout, tracer trace.Tracer) (re.Service, error) { repo := repg.NewRepository(db) idp := uuid.New() @@ -353,21 +295,14 @@ func newService(ctx context.Context, cfg config, db pgclient.Database, runInfo c logger.Error(fmt.Sprintf("failed to configure e-mailing util: %s", err.Error())) } - policyService, err := newSpiceDBPolicyServiceEvaluator(cfg, logger) - if err != nil { - return nil, err - } - logger.Info("Policy service successfully connected to SpiceDB gRPC server") + atomCfg := atom.LoadConfig() - availableActions, builtInRoles, err := availableActionsAndBuiltInRoles(cfg.SpicedbSchemaFile) - if err != nil { - return nil, fmt.Errorf("failed to get available actions and built-in roles: %w", err) - } - - csvc, err := re.NewService(repo, runInfo, policyService, idp, rePubSub, writersPub, alarmsPub, ticker.NewTicker(time.Second*30), emailerClient, readersClient, availableActions, builtInRoles) + var csvc re.Service + csvc, err = re.NewService(repo, runInfo, idp, rePubSub, writersPub, alarmsPub, ticker.NewTicker(time.Second*30), emailerClient, readersClient) if err != nil { return nil, fmt.Errorf("failed to create RE service: %w", err) } + csvc = re.WithAtom(csvc, atom.NewClient(atomCfg)) csvc, err = events.NewEventStoreMiddleware(ctx, csvc, cfg.ESURL) if err != nil { @@ -379,7 +314,7 @@ func newService(ctx context.Context, cfg config, db pgclient.Database, runInfo c return nil, fmt.Errorf("failed to parse permissions file: %w", err) } - ruleOps, ruleRoleOps, err := permConfig.GetEntityPermissions(operations.EntityType) + ruleOps, _, err := permConfig.GetEntityPermissions(operations.EntityType) if err != nil { return nil, fmt.Errorf("failed to get rule permissions: %w", err) } @@ -396,16 +331,11 @@ func newService(ctx context.Context, cfg config, db pgclient.Database, runInfo c return nil, fmt.Errorf("failed to create entities operations: %w", err) } - roleOps, err := permissions.NewOperations(roles.Operations(), ruleRoleOps) - if err != nil { - return nil, fmt.Errorf("failed to create role operations: %w", err) - } - - csvc, err = middleware.AuthorizationMiddleware(csvc, authz, entitiesOps, roleOps) + csvc, err = middleware.AtomAuthorizationMiddleware(csvc, atom.NewClient(atomCfg), entitiesOps) if err != nil { return nil, err } - csvc, err = middleware.NewCallout(csvc, callout, entitiesOps, roleOps) + csvc, err = middleware.NewCallout(csvc, callout, entitiesOps) if err != nil { return nil, err } @@ -416,30 +346,3 @@ func newService(ctx context.Context, cfg config, db pgclient.Database, runInfo c return csvc, nil } - -func newSpiceDBPolicyServiceEvaluator(cfg config, logger *slog.Logger) (policies.Service, error) { - client, err := authzed.NewClientWithExperimentalAPIs( - fmt.Sprintf("%s:%s", cfg.SpicedbHost, cfg.SpicedbPort), - grpc.WithTransportCredentials(insecure.NewCredentials()), - grpcutil.WithInsecureBearerToken(cfg.SpicedbPreSharedKey), - ) - if err != nil { - return nil, err - } - ps := spicedb.NewPolicyService(client, logger) - - return ps, nil -} - -func availableActionsAndBuiltInRoles(spicedbSchemaFile string) ([]roles.Action, map[roles.BuiltInRoleName][]roles.Action, error) { - availableActions, err := spicedbdecoder.GetActionsFromSchema(spicedbSchemaFile, operations.EntityType) - if err != nil { - return []roles.Action{}, map[roles.BuiltInRoleName][]roles.Action{}, err - } - - builtInRoles := map[roles.BuiltInRoleName][]roles.Action{ - re.BuiltInRoleAdmin: availableActions, - } - - return availableActions, builtInRoles, err -} diff --git a/cmd/reports/main.go b/cmd/reports/main.go index e243bc21b..b5391bfc4 100644 --- a/cmd/reports/main.go +++ b/cmd/reports/main.go @@ -17,29 +17,21 @@ import ( chclient "github.com/absmach/callhome/pkg/client" "github.com/absmach/magistrala" grpcReadersV1 "github.com/absmach/magistrala/api/grpc/readers/v1" - dpostgres "github.com/absmach/magistrala/domains/postgres" + "github.com/absmach/magistrala/internal/atom" "github.com/absmach/magistrala/internal/email" mglog "github.com/absmach/magistrala/logger" smqauthn "github.com/absmach/magistrala/pkg/authn" - authnsvc "github.com/absmach/magistrala/pkg/authn/authsvc" - mgauthz "github.com/absmach/magistrala/pkg/authz" - authzsvc "github.com/absmach/magistrala/pkg/authz/authsvc" + atomauthn "github.com/absmach/magistrala/pkg/authn/atom" "github.com/absmach/magistrala/pkg/callout" - dconsumer "github.com/absmach/magistrala/pkg/domains/events/consumer" - domainsAuthz "github.com/absmach/magistrala/pkg/domains/grpcclient" "github.com/absmach/magistrala/pkg/emailer" "github.com/absmach/magistrala/pkg/grpcclient" jaegerclient "github.com/absmach/magistrala/pkg/jaeger" pkglog "github.com/absmach/magistrala/pkg/logger" "github.com/absmach/magistrala/pkg/permissions" - "github.com/absmach/magistrala/pkg/policies" - "github.com/absmach/magistrala/pkg/policies/spicedb" pgclient "github.com/absmach/magistrala/pkg/postgres" "github.com/absmach/magistrala/pkg/prometheus" - "github.com/absmach/magistrala/pkg/roles" "github.com/absmach/magistrala/pkg/server" httpserver "github.com/absmach/magistrala/pkg/server/http" - spicedbdecoder "github.com/absmach/magistrala/pkg/spicedb" "github.com/absmach/magistrala/pkg/ticker" "github.com/absmach/magistrala/pkg/uuid" grpcClient "github.com/absmach/magistrala/readers/api/grpc" @@ -49,14 +41,10 @@ import ( "github.com/absmach/magistrala/reports/middleware" "github.com/absmach/magistrala/reports/operations" repg "github.com/absmach/magistrala/reports/postgres" - "github.com/authzed/authzed-go/v1" - "github.com/authzed/grpcutil" "github.com/caarlos0/env/v11" "github.com/go-chi/chi/v5" "go.opentelemetry.io/otel/trace" "golang.org/x/sync/errgroup" - "google.golang.org/grpc" - "google.golang.org/grpc/credentials/insecure" ) const ( @@ -64,11 +52,9 @@ const ( envPrefixDB = "MG_REPORTS_DB_" envPrefixHTTP = "MG_REPORTS_HTTP_" envPrefixCallout = "MG_REPORTS_CALLOUT_" - envPrefixAuth = "MG_AUTH_GRPC_" defDB = "repo" defSvcHTTPPort = "9017" envPrefixGrpc = "MG_TIMESCALE_READER_GRPC_" - envPrefixDomains = "MG_DOMAINS_GRPC_" templatePath = "template/reports_default_template.html" reportEntity = "report" ) @@ -90,10 +76,6 @@ type config struct { BrokerURL string `env:"MG_MESSAGE_BROKER_URL" envDefault:"nats://localhost:4222"` DefaultTemplatePath string `env:"MG_REPORTS_DEFAULT_TEMPLATE" envDefault:""` ConverterURL string `env:"MG_PDF_CONVERTER_URL" envDefault:"http://localhost:4000/pdf"` - SpicedbHost string `env:"MG_SPICEDB_HOST" envDefault:"localhost"` - SpicedbPort string `env:"MG_SPICEDB_PORT" envDefault:"50051"` - SpicedbPreSharedKey string `env:"MG_SPICEDB_PRE_SHARED_KEY" envDefault:"12345678"` - SpicedbSchemaFile string `env:"MG_SPICEDB_SCHEMA_FILE" envDefault:"schema.zed"` PermissionsFile string `env:"MG_PERMISSIONS_FILE" envDefault:"permission.yaml"` } @@ -216,56 +198,17 @@ func main() { return } - grpcCfg := grpcclient.Config{} - if err := env.ParseWithOptions(&grpcCfg, env.Options{Prefix: envPrefixAuth}); err != nil { - logger.Error(fmt.Sprintf("failed to load auth gRPC client configuration : %s", err)) - exitCode = 1 - - return - } - authn, authnClient, err := authnsvc.NewAuthentication(ctx, grpcCfg) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - - return - } - am := smqauthn.NewAuthNMiddleware(authn) - defer authnClient.Close() - logger.Info("AuthN successfully connected to auth gRPC server " + authnClient.Secure()) - - domsGrpcCfg := grpcclient.Config{} - if err := env.ParseWithOptions(&domsGrpcCfg, env.Options{Prefix: envPrefixDomains}); err != nil { - logger.Error(fmt.Sprintf("failed to load domains gRPC client configuration : %s", err)) - exitCode = 1 - return - } - domAuthz, _, domainsHandler, err := domainsAuthz.NewAuthorization(ctx, domsGrpcCfg) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer domainsHandler.Close() - - authz, authzClient, err := authzsvc.NewAuthorization(ctx, grpcCfg, domAuthz) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer authzClient.Close() - logger.Info("AuthZ successfully connected to auth gRPC server " + authnClient.Secure()) - - ddatabase := pgclient.NewDatabase(db, dbConfig, tracer) - drepo := dpostgres.NewRepository(ddatabase) - - if err := dconsumer.DomainsEventsSubscribe(ctx, drepo, cfg.ESURL, cfg.ESConsumerName, logger); err != nil { - logger.Error(fmt.Sprintf("failed to create domains event store : %s", err)) + atomCfg := atom.LoadConfig() + if atomCfg.URL == "" { + logger.Error("ATOM_URL is required") exitCode = 1 return } + authnSvc := atomauthn.NewAuthentication() + logger.Info("AuthN configured to use Atom bearer tokens") + am := smqauthn.NewAuthNMiddleware(authnSvc) + logger.Info("AuthZ configured to use Atom PDP") database := pgclient.NewDatabase(db, dbConfig, tracer) regrpcCfg := grpcclient.Config{} if err := env.ParseWithOptions(®rpcCfg, env.Options{Prefix: envPrefixGrpc}); err != nil { @@ -286,7 +229,7 @@ func main() { runInfo := make(chan pkglog.RunInfo, channBuffer) - svc, err := newService(ctx, cfg, database, runInfo, authz, ec, logger, readersClient, template, callout, tracer) + svc, err := newService(ctx, cfg, database, runInfo, ec, logger, readersClient, template, callout, tracer) if err != nil { logger.Error(fmt.Sprintf("failed to create services: %s", err)) exitCode = 1 @@ -326,7 +269,7 @@ func main() { } } -func newService(ctx context.Context, cfg config, db pgclient.Database, runInfo chan pkglog.RunInfo, authz mgauthz.Authorization, ec email.Config, logger *slog.Logger, readersClient grpcReadersV1.ReadersServiceClient, template reports.ReportTemplate, callout callout.Callout, tracer trace.Tracer) (reports.Service, error) { +func newService(ctx context.Context, cfg config, db pgclient.Database, runInfo chan pkglog.RunInfo, ec email.Config, logger *slog.Logger, readersClient grpcReadersV1.ReadersServiceClient, template reports.ReportTemplate, callout callout.Callout, tracer trace.Tracer) (reports.Service, error) { repo := repg.NewRepository(db) idp := uuid.New() @@ -335,21 +278,14 @@ func newService(ctx context.Context, cfg config, db pgclient.Database, runInfo c logger.Error(fmt.Sprintf("failed to configure e-mailing util: %s", err.Error())) } - policyService, err := newSpiceDBPolicyServiceEvaluator(cfg, logger) - if err != nil { - return nil, err - } - logger.Info("Policy service successfully connected to SpiceDB gRPC server") + atomCfg := atom.LoadConfig() - availableActions, builtInRoles, err := availableActionsAndBuiltInRoles(cfg.SpicedbSchemaFile) - if err != nil { - return nil, fmt.Errorf("failed to get available actions and built-in roles: %w", err) - } - - csvc, err := reports.NewService(repo, runInfo, policyService, idp, ticker.NewTicker(time.Second*30), emailClient, readersClient, template, cfg.ConverterURL, availableActions, builtInRoles) + var csvc reports.Service + csvc, err = reports.NewService(repo, runInfo, idp, ticker.NewTicker(time.Second*30), emailClient, readersClient, template, cfg.ConverterURL) if err != nil { return nil, fmt.Errorf("failed to create reports service: %w", err) } + csvc = reports.WithAtom(csvc, atom.NewClient(atomCfg)) csvc, err = reportsevents.NewEventStoreMiddleware(ctx, csvc, cfg.ESURL) if err != nil { @@ -361,7 +297,7 @@ func newService(ctx context.Context, cfg config, db pgclient.Database, runInfo c return nil, fmt.Errorf("failed to parse permissions file: %w", err) } - reportOps, reportRoleOps, err := permConfig.GetEntityPermissions(reportEntity) + reportOps, _, err := permConfig.GetEntityPermissions(reportEntity) if err != nil { return nil, fmt.Errorf("failed to get report permissions: %w", err) } @@ -378,16 +314,11 @@ func newService(ctx context.Context, cfg config, db pgclient.Database, runInfo c return nil, fmt.Errorf("failed to create entities operations: %w", err) } - roleOps, err := permissions.NewOperations(roles.Operations(), reportRoleOps) - if err != nil { - return nil, fmt.Errorf("failed to create role operations: %w", err) - } - - csvc, err = middleware.AuthorizationMiddleware(csvc, authz, entitiesOps, roleOps) + csvc, err = middleware.AtomAuthorizationMiddleware(csvc, atom.NewClient(atomCfg), entitiesOps) if err != nil { return nil, err } - csvc, err = middleware.NewCallout(csvc, callout, entitiesOps, roleOps) + csvc, err = middleware.NewCallout(csvc, callout, entitiesOps) if err != nil { return nil, err } @@ -398,30 +329,3 @@ func newService(ctx context.Context, cfg config, db pgclient.Database, runInfo c return csvc, nil } - -func newSpiceDBPolicyServiceEvaluator(cfg config, logger *slog.Logger) (policies.Service, error) { - client, err := authzed.NewClientWithExperimentalAPIs( - fmt.Sprintf("%s:%s", cfg.SpicedbHost, cfg.SpicedbPort), - grpc.WithTransportCredentials(insecure.NewCredentials()), - grpcutil.WithInsecureBearerToken(cfg.SpicedbPreSharedKey), - ) - if err != nil { - return nil, err - } - ps := spicedb.NewPolicyService(client, logger) - - return ps, nil -} - -func availableActionsAndBuiltInRoles(spicedbSchemaFile string) ([]roles.Action, map[roles.BuiltInRoleName][]roles.Action, error) { - availableActions, err := spicedbdecoder.GetActionsFromSchema(spicedbSchemaFile, reportEntity) - if err != nil { - return []roles.Action{}, map[roles.BuiltInRoleName][]roles.Action{}, err - } - - builtInRoles := map[roles.BuiltInRoleName][]roles.Action{ - reports.BuiltInRoleAdmin: availableActions, - } - - return availableActions, builtInRoles, err -} diff --git a/cmd/timescale-reader/main.go b/cmd/timescale-reader/main.go index 132e5bcc5..a99d4c3e9 100644 --- a/cmd/timescale-reader/main.go +++ b/cmd/timescale-reader/main.go @@ -14,9 +14,9 @@ import ( chclient "github.com/absmach/callhome/pkg/client" "github.com/absmach/magistrala" grpcReadersV1 "github.com/absmach/magistrala/api/grpc/readers/v1" + "github.com/absmach/magistrala/internal/atom" mglog "github.com/absmach/magistrala/logger" - "github.com/absmach/magistrala/pkg/authn/authsvc" - "github.com/absmach/magistrala/pkg/grpcclient" + atomauthn "github.com/absmach/magistrala/pkg/authn/atom" pgclient "github.com/absmach/magistrala/pkg/postgres" "github.com/absmach/magistrala/pkg/prometheus" "github.com/absmach/magistrala/pkg/server" @@ -36,16 +36,13 @@ import ( ) const ( - svcName = "timescaledb-reader" - envPrefixDB = "MG_TIMESCALE_" - envPrefixHTTP = "MG_TIMESCALE_READER_HTTP_" - envPrefixAuth = "MG_AUTH_GRPC_" - envPrefixClients = "MG_CLIENTS_GRPC_" - envPrefixChannels = "MG_CHANNELS_GRPC_" - defDB = "messages" - defSvcHTTPPort = "9011" - defSvcGRPCPort = "7011" - envPrefixGrpc = "MG_TIMESCALE_READER_GRPC_" + svcName = "timescaledb-reader" + envPrefixDB = "MG_TIMESCALE_" + envPrefixHTTP = "MG_TIMESCALE_READER_HTTP_" + defDB = "messages" + defSvcHTTPPort = "9011" + defSvcGRPCPort = "7011" + envPrefixGrpc = "MG_TIMESCALE_READER_GRPC_" ) type config struct { @@ -106,54 +103,11 @@ func main() { grpcReadersV1.RegisterReadersServiceServer(srv, readersgrpcapi.NewReadersServer(repo)) } - clientsClientCfg := grpcclient.Config{} - if err := env.ParseWithOptions(&clientsClientCfg, env.Options{Prefix: envPrefixClients}); err != nil { - logger.Error(fmt.Sprintf("failed to load %s auth configuration : %s", svcName, err)) - exitCode = 1 - return - } - - clientsClient, clientsHandler, err := grpcclient.SetupClientsClient(ctx, clientsClientCfg) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer clientsHandler.Close() - - logger.Info("Clients service gRPC client successfully connected to clients gRPC server " + clientsHandler.Secure()) - - channelsClientCfg := grpcclient.Config{} - if err := env.ParseWithOptions(&channelsClientCfg, env.Options{Prefix: envPrefixChannels}); err != nil { - logger.Error(fmt.Sprintf("failed to load channels gRPC client configuration : %s", err)) - exitCode = 1 - return - } - - channelsClient, channelsHandler, err := grpcclient.SetupChannelsClient(ctx, channelsClientCfg) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer channelsHandler.Close() - logger.Info("Channels service gRPC client successfully connected to channels gRPC server " + channelsHandler.Secure()) - - authnCfg := grpcclient.Config{} - if err := env.ParseWithOptions(&authnCfg, env.Options{Prefix: envPrefixAuth}); err != nil { - logger.Error(fmt.Sprintf("failed to load auth gRPC client configuration : %s", err)) - exitCode = 1 - return - } - - authn, authnHandler, err := authsvc.NewAuthentication(ctx, authnCfg) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer authnHandler.Close() - logger.Info("authn successfully connected to auth gRPC server " + authnHandler.Secure()) + atomCfg := atom.LoadConfig() + authn := atomauthn.NewAuthentication() + clientsClient := atom.NewClientsCompat(authn) + channelsClient := atom.NewChannelsCompat(atom.NewClient(atomCfg)) + logger.Info("AuthN/AuthZ configured to use Atom") httpServerConfig := server.Config{Port: defSvcHTTPPort} if err := env.ParseWithOptions(&httpServerConfig, env.Options{Prefix: envPrefixHTTP}); err != nil { diff --git a/cmd/users/main.go b/cmd/users/main.go deleted file mode 100644 index 223ddcbd7..000000000 --- a/cmd/users/main.go +++ /dev/null @@ -1,434 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package main contains users main function to start the users service. -package main - -import ( - "context" - "fmt" - "log" - "log/slog" - "net/url" - "os" - "regexp" - "time" - - chclient "github.com/absmach/callhome/pkg/client" - "github.com/absmach/magistrala" - grpcDomainsV1 "github.com/absmach/magistrala/api/grpc/domains/v1" - grpcTokenV1 "github.com/absmach/magistrala/api/grpc/token/v1" - grpcUsersV1 "github.com/absmach/magistrala/api/grpc/users/v1" - "github.com/absmach/magistrala/auth" - "github.com/absmach/magistrala/internal/email" - mglog "github.com/absmach/magistrala/logger" - smqauthn "github.com/absmach/magistrala/pkg/authn" - authsvcAuthn "github.com/absmach/magistrala/pkg/authn/authsvc" - jwksAuthn "github.com/absmach/magistrala/pkg/authn/jwks" - smqauthz "github.com/absmach/magistrala/pkg/authz" - authsvcAuthz "github.com/absmach/magistrala/pkg/authz/authsvc" - domainsAuthz "github.com/absmach/magistrala/pkg/domains/grpcclient" - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - "github.com/absmach/magistrala/pkg/grpcclient" - jaegerclient "github.com/absmach/magistrala/pkg/jaeger" - "github.com/absmach/magistrala/pkg/oauth2" - googleoauth "github.com/absmach/magistrala/pkg/oauth2/google" - "github.com/absmach/magistrala/pkg/policies" - "github.com/absmach/magistrala/pkg/policies/spicedb" - pg "github.com/absmach/magistrala/pkg/postgres" - pgclient "github.com/absmach/magistrala/pkg/postgres" - "github.com/absmach/magistrala/pkg/prometheus" - "github.com/absmach/magistrala/pkg/server" - grpcserver "github.com/absmach/magistrala/pkg/server/grpc" - httpserver "github.com/absmach/magistrala/pkg/server/http" - "github.com/absmach/magistrala/pkg/uuid" - "github.com/absmach/magistrala/users" - httpapi "github.com/absmach/magistrala/users/api" - grpcapi "github.com/absmach/magistrala/users/api/grpc" - "github.com/absmach/magistrala/users/emailer" - "github.com/absmach/magistrala/users/events" - "github.com/absmach/magistrala/users/hasher" - "github.com/absmach/magistrala/users/middleware" - "github.com/absmach/magistrala/users/postgres" - pusers "github.com/absmach/magistrala/users/private" - "github.com/authzed/authzed-go/v1" - "github.com/authzed/grpcutil" - "github.com/caarlos0/env/v11" - "github.com/go-chi/chi/v5" - "go.opentelemetry.io/otel/trace" - "golang.org/x/sync/errgroup" - "google.golang.org/grpc" - "google.golang.org/grpc/credentials/insecure" - "google.golang.org/grpc/reflection" -) - -const ( - svcName = "users" - envPrefixDB = "MG_USERS_DB_" - envPrefixHTTP = "MG_USERS_HTTP_" - envPrefixGRPC = "MG_USERS_GRPC_" - envPrefixAuth = "MG_AUTH_GRPC_" - envPrefixDomains = "MG_DOMAINS_GRPC_" - envPrefixGoogle = "MG_GOOGLE_" - defDB = "users" - defSvcHTTPPort = "9002" - defSvcGRPCPort = "7002" -) - -type config struct { - LogLevel string `env:"MG_USERS_LOG_LEVEL" envDefault:"info"` - AdminEmail string `env:"MG_USERS_ADMIN_EMAIL" envDefault:"admin@example.com"` - AdminPassword string `env:"MG_USERS_ADMIN_PASSWORD" envDefault:"12345678"` - AdminUsername string `env:"MG_USERS_ADMIN_USERNAME" envDefault:"admin"` - AdminFirstName string `env:"MG_USERS_ADMIN_FIRST_NAME" envDefault:"super"` - AdminLastName string `env:"MG_USERS_ADMIN_LAST_NAME" envDefault:"admin"` - PassRegexText string `env:"MG_USERS_PASS_REGEX" envDefault:"^.{8,}$"` - JaegerURL url.URL `env:"MG_JAEGER_URL" envDefault:"http://localhost:4318/v1/traces"` - SendTelemetry bool `env:"MG_SEND_TELEMETRY" envDefault:"true"` - InstanceID string `env:"MG_USERS_INSTANCE_ID" envDefault:""` - ESURL string `env:"MG_ES_URL" envDefault:"amqp://guest:guest@localhost:5682/"` - TraceRatio float64 `env:"MG_JAEGER_TRACE_RATIO" envDefault:"1.0"` - SelfRegister bool `env:"MG_USERS_ALLOW_SELF_REGISTER" envDefault:"false"` - OAuthUIRedirectURL string `env:"MG_OAUTH_UI_REDIRECT_URL" envDefault:"http://localhost:9095/domains"` - OAuthUIErrorURL string `env:"MG_OAUTH_UI_ERROR_URL" envDefault:"http://localhost:9095/error"` - DeleteInterval time.Duration `env:"MG_USERS_DELETE_INTERVAL" envDefault:"24h"` - DeleteAfter time.Duration `env:"MG_USERS_DELETE_AFTER" envDefault:"720h"` - SpicedbHost string `env:"MG_SPICEDB_HOST" envDefault:"localhost"` - SpicedbPort string `env:"MG_SPICEDB_PORT" envDefault:"50051"` - SpicedbPreSharedKey string `env:"MG_SPICEDB_PRE_SHARED_KEY" envDefault:"12345678"` - PasswordResetURLPrefix string `env:"MG_PASSWORD_RESET_URL_PREFIX" envDefault:"http://localhost/password/reset"` - PasswordResetEmailTemplate string `env:"MG_PASSWORD_RESET_EMAIL_TEMPLATE" envDefault:"reset-password-email.tmpl"` - VerificationURLPrefix string `env:"MG_VERIFICATION_URL_PREFIX" envDefault:"http://localhost/verify-email"` - VerificationEmailTemplate string `env:"MG_VERIFICATION_EMAIL_TEMPLATE" envDefault:"verification-email.tmpl"` - AuthKeyAlgorithm string `env:"MG_AUTH_KEYS_ALGORITHM" envDefault:"RS256"` - JWKSURL string `env:"MG_AUTH_JWKS_URL" envDefault:"http://auth:9001/keys/.well-known/jwks.json"` - PassRegex *regexp.Regexp -} - -func main() { - ctx, cancel := context.WithCancel(context.Background()) - g, ctx := errgroup.WithContext(ctx) - - cfg := config{} - if err := env.Parse(&cfg); err != nil { - log.Fatalf("failed to load %s configuration : %s", svcName, err.Error()) - } - passRegex, err := regexp.Compile(cfg.PassRegexText) - if err != nil { - log.Fatalf("invalid password validation rules %s\n", cfg.PassRegexText) - } - cfg.PassRegex = passRegex - - logger, err := mglog.New(os.Stdout, cfg.LogLevel) - if err != nil { - log.Fatalf("failed to init logger: %s", err.Error()) - } - - var exitCode int - defer mglog.ExitWithError(&exitCode) - - if cfg.InstanceID == "" { - if cfg.InstanceID, err = uuid.New().ID(); err != nil { - logger.Error(fmt.Sprintf("failed to generate instanceID: %s", err)) - exitCode = 1 - return - } - } - - resetPasswordEmailConfig := email.Config{} - if err := env.Parse(&resetPasswordEmailConfig); err != nil { - logger.Error(fmt.Sprintf("failed to load reset password email configuration : %s", err.Error())) - exitCode = 1 - return - } - resetPasswordEmailConfig.Template = cfg.PasswordResetEmailTemplate - - verificationEmailConfig := email.Config{} - if err := env.Parse(&verificationEmailConfig); err != nil { - logger.Error(fmt.Sprintf("failed to load verification password email configuration : %s", err.Error())) - exitCode = 1 - return - } - verificationEmailConfig.Template = cfg.VerificationEmailTemplate - - dbConfig := pgclient.Config{Name: defDB} - if err := env.ParseWithOptions(&dbConfig, env.Options{Prefix: envPrefixDB}); err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - - migration := postgres.Migration() - db, err := pgclient.Setup(dbConfig, *migration) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer db.Close() - - tp, err := jaegerclient.NewProvider(ctx, svcName, cfg.JaegerURL, cfg.InstanceID, cfg.TraceRatio) - if err != nil { - logger.Error(fmt.Sprintf("failed to init Jaeger: %s", err)) - exitCode = 1 - return - } - defer func() { - if err := tp.Shutdown(ctx); err != nil { - logger.Error(fmt.Sprintf("error shutting down tracer provider: %v", err)) - } - }() - tracer := tp.Tracer(svcName) - - database := pg.NewDatabase(db, dbConfig, tracer) - repo := postgres.NewRepository(database) - - authClientConfig := grpcclient.Config{} - if err := env.ParseWithOptions(&authClientConfig, env.Options{Prefix: envPrefixAuth}); err != nil { - logger.Error(fmt.Sprintf("failed to load %s auth configuration : %s", svcName, err)) - exitCode = 1 - return - } - - tokenClient, tokenHandler, err := grpcclient.SetupTokenClient(ctx, authClientConfig) - if err != nil { - logger.Error("failed to create token gRPC client " + err.Error()) - exitCode = 1 - return - } - defer tokenHandler.Close() - logger.Info("Token service client successfully connected to auth gRPC server " + tokenHandler.Secure()) - - isSymmetric, err := auth.IsSymmetricAlgorithm(cfg.AuthKeyAlgorithm) - if err != nil { - logger.Error(fmt.Sprintf("failed to parse auth key algorithm : %s", err)) - exitCode = 1 - return - } - var authn smqauthn.Authentication - var authnClient grpcclient.Handler - switch { - case !isSymmetric: - authn, authnClient, err = jwksAuthn.NewAuthentication(ctx, cfg.JWKSURL, authClientConfig) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer authnClient.Close() - logger.Info("AuthN successfully set up jwks authentication on " + cfg.JWKSURL) - default: - authn, authnClient, err = authsvcAuthn.NewAuthentication(ctx, authClientConfig) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer authnClient.Close() - logger.Info("AuthN successfully connected to auth gRPC server " + authnClient.Secure()) - } - authnMiddleware := smqauthn.NewAuthNMiddleware(authn) - domsGrpcCfg := grpcclient.Config{} - if err := env.ParseWithOptions(&domsGrpcCfg, env.Options{Prefix: envPrefixDomains}); err != nil { - logger.Error(fmt.Sprintf("failed to load domains gRPC client configuration : %s", err)) - exitCode = 1 - return - } - domAuthz, domainsClient, domainsHandler, err := domainsAuthz.NewAuthorization(ctx, domsGrpcCfg) - if err != nil { - logger.Error(err.Error()) - exitCode = 1 - return - } - defer domainsHandler.Close() - - authz, authzHandler, err := authsvcAuthz.NewAuthorization(ctx, authClientConfig, domAuthz) - if err != nil { - logger.Error("failed to create authz " + err.Error()) - exitCode = 1 - return - } - defer authzHandler.Close() - logger.Info("AuthZ successfully connected to auth gRPC server " + authzHandler.Secure()) - - policyService, err := newPolicyService(cfg, logger) - if err != nil { - logger.Error("failed to create new policies service " + err.Error()) - exitCode = 1 - return - } - logger.Info("Policy client successfully connected to spicedb gRPC server") - - csvc, err := newService(ctx, authz, tokenClient, policyService, domainsClient, repo, tracer, cfg, resetPasswordEmailConfig, verificationEmailConfig, logger) - if err != nil { - logger.Error(fmt.Sprintf("failed to setup service: %s", err)) - exitCode = 1 - return - } - - psvc := pusers.New(repo) - - grpcServerConfig := server.Config{Port: defSvcGRPCPort} - if err := env.ParseWithOptions(&grpcServerConfig, env.Options{Prefix: envPrefixGRPC}); err != nil { - logger.Error(fmt.Sprintf("failed to load %s gRPC server configuration : %s", svcName, err.Error())) - exitCode = 1 - return - } - - registerUsersServer := func(srv *grpc.Server) { - reflection.Register(srv) - grpcUsersV1.RegisterUsersServiceServer(srv, grpcapi.NewServer(psvc)) - } - gs := grpcserver.NewServer(ctx, cancel, svcName, grpcServerConfig, registerUsersServer, logger) - - httpServerConfig := server.Config{Port: defSvcHTTPPort} - if err := env.ParseWithOptions(&httpServerConfig, env.Options{Prefix: envPrefixHTTP}); err != nil { - logger.Error(fmt.Sprintf("failed to load %s HTTP server configuration : %s", svcName, err.Error())) - exitCode = 1 - return - } - - oauthConfig := oauth2.Config{} - if err := env.ParseWithOptions(&oauthConfig, env.Options{Prefix: envPrefixGoogle}); err != nil { - logger.Error(fmt.Sprintf("failed to load %s Google configuration : %s", svcName, err.Error())) - exitCode = 1 - return - } - oauthProvider := googleoauth.NewProvider(oauthConfig, cfg.OAuthUIRedirectURL, cfg.OAuthUIErrorURL) - - mux := chi.NewRouter() - idp := uuid.New() - httpSrv := httpserver.NewServer(ctx, cancel, svcName, httpServerConfig, httpapi.MakeHandler(csvc, authnMiddleware, tokenClient, cfg.SelfRegister, mux, logger, cfg.InstanceID, cfg.PassRegex, idp, oauthProvider), logger) - - if cfg.SendTelemetry { - chc := chclient.New(svcName, magistrala.Version, logger, cancel) - go chc.CallHome(ctx) - } - - g.Go(func() error { - return httpSrv.Start() - }) - - g.Go(func() error { - return gs.Start() - }) - - g.Go(func() error { - return server.StopSignalHandler(ctx, cancel, logger, svcName, httpSrv, gs) - }) - - if err := g.Wait(); err != nil { - logger.Error(fmt.Sprintf("users service terminated: %s", err)) - } -} - -func newService(ctx context.Context, authz smqauthz.Authorization, token grpcTokenV1.TokenServiceClient, policyService policies.Service, domainsClient grpcDomainsV1.DomainsServiceClient, repo users.Repository, tracer trace.Tracer, c config, resetPasswordEmailConfig, verificationEmailConfig email.Config, logger *slog.Logger) (users.Service, error) { - idp := uuid.New() - hsr := hasher.New() - - // Creating users service - emailerClient, err := emailer.New( - c.PasswordResetURLPrefix, - c.VerificationURLPrefix, - &resetPasswordEmailConfig, - &verificationEmailConfig, - ) - if err != nil { - logger.Error(fmt.Sprintf("failed to configure e-mailing util: %s", err.Error())) - return nil, err - } - - svc := users.NewService(token, repo, policyService, emailerClient, hsr, idp) - - svc, err = events.NewEventStoreMiddleware(ctx, svc, c.ESURL) - if err != nil { - return nil, err - } - svc = middleware.NewAuthorization(svc, authz, c.SelfRegister) - - svc = middleware.NewTracing(svc, tracer) - svc = middleware.NewLogging(svc, logger) - counter, latency := prometheus.MakeMetrics(svcName, "api") - svc = middleware.NewMetrics(svc, counter, latency) - - userID, err := createAdmin(ctx, c, repo, hsr, svc) - if err != nil { - logger.Error(fmt.Sprintf("failed to create admin client: %s", err)) - } - if userID != "" { - if err := createAdminPolicy(ctx, userID, policyService); err != nil { - return nil, err - } - } - - users.NewDeleteHandler(ctx, repo, policyService, domainsClient, c.DeleteInterval, c.DeleteAfter, logger) - - return svc, err -} - -func createAdmin(ctx context.Context, c config, repo users.Repository, hsr users.Hasher, svc users.Service) (string, error) { - id, err := uuid.New().ID() - if err != nil { - return "", err - } - hash, err := hsr.Hash(c.AdminPassword) - if err != nil { - return "", err - } - - user := users.User{ - ID: id, - Email: c.AdminEmail, - FirstName: c.AdminFirstName, - LastName: c.AdminLastName, - Credentials: users.Credentials{ - Username: "admin", - Secret: hash, - }, - Metadata: users.Metadata{ - "role": "admin", - }, - CreatedAt: time.Now().UTC(), - UpdatedAt: time.Now().UTC(), - Role: users.AdminRole, - Status: users.EnabledStatus, - } - - if u, err := repo.RetrieveByEmail(ctx, user.Email); err == nil { - return u.ID, nil - } - - if _, err = repo.Save(ctx, user); err != nil { - return "", err - } - return user.ID, nil -} - -func createAdminPolicy(ctx context.Context, userID string, policyService policies.Service) error { - err := policyService.AddPolicy(ctx, policies.Policy{ - SubjectType: policies.UserType, - Subject: userID, - Relation: policies.AdministratorRelation, - Object: policies.MagistralaObject, - ObjectType: policies.PlatformType, - }) - if err != nil && !errors.Contains(err, repoerr.ErrConflict) { - return err - } - return nil -} - -func newPolicyService(cfg config, logger *slog.Logger) (policies.Service, error) { - client, err := authzed.NewClientWithExperimentalAPIs( - fmt.Sprintf("%s:%s", cfg.SpicedbHost, cfg.SpicedbPort), - grpc.WithTransportCredentials(insecure.NewCredentials()), - grpcutil.WithInsecureBearerToken(cfg.SpicedbPreSharedKey), - ) - if err != nil { - return nil, err - } - policySvc := spicedb.NewPolicyService(client, logger) - - return policySvc, nil -} diff --git a/consumers/README.md b/consumers/README.md index 6b608f80b..96eb7f9f2 100644 --- a/consumers/README.md +++ b/consumers/README.md @@ -4,7 +4,7 @@ Consumers provide an abstraction for various “Magistrala consumers”. A consumer is a generic plugin‑style service that handles received messages — for example, writing them to a database, sending notifications, or transforming them. Before consuming, messages from Magistrala can be transformed (e.g. to JSON or SenML) to match what a specific consumer expects. -This service (Notifiers) is optional — to use it, core services must be running (e.g. message broker + clients + channels etc.). +This service (Notifiers) is optional — to use it, services must be running (e.g. message broker + clients + channels etc.). ## Concepts & Consumer Types @@ -25,7 +25,7 @@ When a subscriber receives messages from the message broker: 2. The transformed message is passed to a consumer — either synchronously (BlockingConsumer) or asynchronously (AsyncConsumer). 3. The consumer handles the message (e.g. storing to DB, sending notifications, writing files, etc.). -Consumers are decoupled from core messaging logic, making them flexible and pluggable. +Consumers are decoupled from messaging logic, making them flexible and pluggable. ## Supported Consumers diff --git a/consumers/notifiers/api/logging.go b/consumers/notifiers/api/logging.go index 26fcf1235..1b71ea02f 100644 --- a/consumers/notifiers/api/logging.go +++ b/consumers/notifiers/api/logging.go @@ -20,7 +20,7 @@ type loggingMiddleware struct { svc notifiers.Service } -// LoggingMiddleware adds logging facilities to the core service. +// LoggingMiddleware adds logging facilities to the service. func LoggingMiddleware(svc notifiers.Service, logger *slog.Logger) notifiers.Service { return &loggingMiddleware{logger, svc} } diff --git a/consumers/notifiers/api/metrics.go b/consumers/notifiers/api/metrics.go index 38d0bb4d4..42507f517 100644 --- a/consumers/notifiers/api/metrics.go +++ b/consumers/notifiers/api/metrics.go @@ -21,7 +21,7 @@ type metricsMiddleware struct { svc notifiers.Service } -// MetricsMiddleware instruments core service by tracking request count and latency. +// MetricsMiddleware instruments service by tracking request count and latency. func MetricsMiddleware(svc notifiers.Service, counter metrics.Counter, latency metrics.Histogram) notifiers.Service { return &metricsMiddleware{ counter: counter, diff --git a/docker/.env b/docker/.env index 8c1a94ce9..e00d094e8 100644 --- a/docker/.env +++ b/docker/.env @@ -31,7 +31,9 @@ MG_FLUXMQ_API_PORT_2=9082 MG_FLUXMQ_API_PORT_3=9083 ## Message Broker -MG_MESSAGE_BROKER_URL=amqp://guest:guest@nginx:${MG_FLUXMQ_AMQP091_PORT}/ +MG_MESSAGE_BROKER_URL=amqp://guest:guest@nginx:${MG_NGINX_AMQP_PORT}/ +MG_FLUXMQ_PUBLISH_HTTP_HOST=fluxmq-auth +MG_FLUXMQ_PUBLISH_HTTP_PORT=9026 ## Redis MG_REDIS_TCP_PORT=6379 @@ -39,7 +41,7 @@ MG_REDIS_URL=redis://es-redis:${MG_REDIS_TCP_PORT}/0 ## Event Store MG_ES_TYPE=es_fluxmq -MG_ES_URL=amqp://guest:guest@nginx:5682/ +MG_ES_URL=amqp://guest:guest@nginx:${MG_NGINX_AMQP_PORT}/ ## Jaeger MG_JAEGER_COLLECTOR_OTLP_ENABLED=true @@ -138,6 +140,57 @@ MG_AUTH_GRPC_CLIENT_CERT=${GRPC_MTLS:+./ssl/certs/auth-grpc-client.crt} MG_AUTH_GRPC_CLIENT_KEY=${GRPC_MTLS:+./ssl/certs/auth-grpc-client.key} MG_AUTH_GRPC_CLIENT_CA_CERTS=${GRPC_MTLS:+./ssl/certs/ca.crt} +### Atom Integration +ATOM_URL=http://atom:8080 +ATOM_PUBLIC_URL=http://nginx:80 +ATOM_JWKS_URL=http://atom:8080/.well-known/jwks.json +ATOM_JWT_ISSUER=http://nginx:80 +ATOM_JWT_AUDIENCE=magistrala +ATOM_SIGNUP_ENABLED=true +ATOM_ALLOW_UNVERIFIED_EMAIL_LOGIN=true +ATOM_DEV_ALLOW_UNVERIFIED_EMAIL_LOGIN=true +ATOM_CORS_ALLOWED_ORIGINS=http://localhost:3000,http://localhost +ATOM_UI_HTTP_PORT=3005 +ATOM_SERVICE_TOKEN= +ATOM_SERVICE_USERNAME=admin +ATOM_SERVICE_SECRET=12345678 +ATOM_ADMIN_TOKEN= +ATOM_ADMIN_USERNAME=admin +ATOM_TIMEOUT=5s +ATOM_HTTP_PORT=8080 +ATOM_GRPC_ADDR=0.0.0.0:8081 +ATOM_DB_PORT=6010 +ATOM_DB_USER=atom +ATOM_DB_PASSWORD=atom +ATOM_DB_NAME=atom +ATOM_JWT_SECRET=change-me-in-production +ATOM_JWT_EXPIRY_SECS=3600 +ATOM_KEY_ENCRYPTION_KEY=6uQzr7tCp9cupO8anzo0i6XtfVocixzIK78lB4o6S3E= +ATOM_KEY_ENCRYPTION_KEY_ID=local:v1 +ATOM_ALLOW_PLAINTEXT_SIGNING_KEYS=false +ATOM_ADMIN_SECRET=12345678 +ATOM_MIN_PASSWORD_CHARS=8 +ATOM_CERTS_ENABLED=true +ATOM_CERTS_CA_MODE=file_root_issuer +ATOM_CERTS_ROOT_CA_CERT_PATH=/certs/ca.crt +ATOM_CERTS_ROOT_CA_KEY_PATH=/certs/ca.key +ATOM_CERTS_CA_DIR=./ssl/certs +ATOM_CERTS_LEAF_DEFAULT_TTL_SECS=2592000 +ATOM_CERTS_LEAF_MAX_TTL_SECS=2592000 +# Local compose runs the server-rendered UI and GraphQL traffic through shared +# Docker peer IPs, which can trip Atom's default per-IP GraphQL limits. +ATOM_RATE_LIMIT_ENABLED=false +ATOM_TRUSTED_PROXY_CIDRS= +ATOM_RUST_LOG=info +ATOM_INVITATION_REDIRECT=http://localhost:3000/invitations/accept +ATOM_INVITATION_EXPIRY_SECS=604800 +ATOM_SMTP_HOST=host.docker.internal +ATOM_SMTP_PORT=2525 +ATOM_SMTP_USERNAME= +ATOM_SMTP_PASSWORD= +ATOM_SMTP_FROM=from@example.com +ATOM_SMTP_TLS=none + ### Domains MG_DOMAINS_LOG_LEVEL=debug MG_DOMAINS_HTTP_HOST=domains @@ -163,26 +216,14 @@ MG_DOMAINS_CACHE_URL=redis://domains-redis:${MG_REDIS_TCP_PORT}/0 MG_DOMAINS_CACHE_KEY_DURATION=10m #### Domains Client Config -MG_DOMAINS_URL=http://domains:9003 -MG_DOMAINS_GRPC_URL=domains:7003 +MG_DOMAINS_URL=http://atom:8080/tenants +MG_DOMAINS_GRPC_URL= MG_DOMAINS_GRPC_TIMEOUT=300s MG_DOMAINS_GRPC_CLIENT_CERT=${GRPC_MTLS:+./ssl/certs/domains-grpc-client.crt} MG_DOMAINS_GRPC_CLIENT_KEY=${GRPC_MTLS:+./ssl/certs/domains-grpc-client.key} MG_DOMAINS_GRPC_CLIENT_CA_CERTS=${GRPC_MTLS:+./ssl/certs/ca.crt} -### SpiceDB Datastore config -MG_SPICEDB_DB_USER=magistrala -MG_SPICEDB_DB_PASS=magistrala -MG_SPICEDB_DB_NAME=spicedb -MG_SPICEDB_DB_PORT=5432 - -### SpiceDB config -MG_SPICEDB_PRE_SHARED_KEY="12345678" -MG_SPICEDB_SCHEMA_FILE="/schema.zed" MG_PERMISSIONS_FILE="/permission.yaml" -MG_SPICEDB_HOST=spicedb -MG_SPICEDB_PORT=50051 -MG_SPICEDB_DATASTORE_ENGINE=postgres ### Users MG_USERS_LOG_LEVEL=debug @@ -223,8 +264,8 @@ MG_VERIFICATION_URL_PREFIX=http://localhost/verify-email MG_VERIFICATION_EMAIL_TEMPLATE=verification-email.tmpl #### Users Client Config -MG_USERS_URL=http://users:9002 -MG_USERS_GRPC_URL=users:7002 +MG_USERS_URL=http://atom:8080/entities +MG_USERS_GRPC_URL= MG_USERS_GRPC_TIMEOUT=300s MG_USERS_GRPC_CLIENT_CERT=${GRPC_MTLS:+./ssl/certs/domains-grpc-client.crt} MG_USERS_GRPC_CLIENT_KEY=${GRPC_MTLS:+./ssl/certs/domains-grpc-client.key} @@ -237,6 +278,7 @@ MG_EMAIL_USERNAME=from@example.com MG_EMAIL_PASSWORD=password MG_EMAIL_FROM_ADDRESS=from@example.com MG_EMAIL_FROM_NAME=Example +MG_EMAIL_TEMPLATE= MG_EMAIL_INVITATION_TEMPLATE=invitation-sent-email.tmpl MG_EMAIL_ACCEPTANCE_TEMPLATE=invitation-accepted-email.tmpl MG_EMAIL_REJECTION_TEMPLATE=invitation-rejected-email.tmpl @@ -274,8 +316,8 @@ MG_GROUPS_DB_SSL_ROOT_CERT= MG_GROUPS_INSTANCE_ID= #### Groups Client Config -MG_GROUPS_URL=groups:9004 -MG_GROUPS_GRPC_URL=groups:7004 +MG_GROUPS_URL=http://atom:8080/groups +MG_GROUPS_GRPC_URL= MG_GROUPS_GRPC_TIMEOUT=300s MG_GROUPS_GRPC_CLIENT_CERT=${GRPC_MTLS:+./ssl/certs/groups-grpc-client.crt} MG_GROUPS_GRPC_CLIENT_KEY=${GRPC_MTLS:+./ssl/certs/groups-grpc-client.key} @@ -306,8 +348,8 @@ MG_CLIENTS_DB_SSL_ROOT_CERT= MG_CLIENTS_INSTANCE_ID= #### Clients Client Config -MG_CLIENTS_URL=http://clients:9006 -MG_CLIENTS_GRPC_URL=clients:7006 +MG_CLIENTS_URL=http://atom:8080/entities +MG_CLIENTS_GRPC_URL= MG_CLIENTS_GRPC_TIMEOUT=300s MG_CLIENTS_GRPC_CLIENT_CERT=${GRPC_MTLS:+./ssl/certs/clients-grpc-client.crt} MG_CLIENTS_GRPC_CLIENT_KEY=${GRPC_MTLS:+./ssl/certs/clients-grpc-client.key} @@ -336,8 +378,8 @@ MG_CHANNELS_CACHE_URL=redis://channels-redis:${MG_REDIS_TCP_PORT}/0 MG_CHANNELS_CACHE_KEY_DURATION=10m #### Channels Client Config -MG_CHANNELS_URL=http://channels:9005 -MG_CHANNELS_GRPC_URL=channels:7005 +MG_CHANNELS_URL=http://atom:8080/resources +MG_CHANNELS_GRPC_URL= MG_CHANNELS_GRPC_TIMEOUT=300s MG_CHANNELS_GRPC_CLIENT_CERT=${GRPC_MTLS:+./ssl/certs/channels-grpc-client.crt} MG_CHANNELS_GRPC_CLIENT_KEY=${GRPC_MTLS:+./ssl/certs/channels-grpc-client.key} @@ -357,60 +399,8 @@ MG_FLUXMQ_CACHE_BUFFER_ITEMS=64 MG_COAP_PORT=5683 ## Addons Services -# Certs -MG_CERTS_LOG_LEVEL=debug -MG_CERTS_HTTP_HOST=certs -MG_CERTS_HTTP_PORT=9019 -MG_CERTS_GRPC_HOST=certs -MG_CERTS_GRPC_PORT=7012 -# WARNING: This is a development/testing secret only. -# NEVER use this weak secret in production! Generate a strong random secret for production deployments. -MG_CERTS_SECRET=12345678 -## Certs Database Configuration -MG_CERTS_DB_HOST=certs-db -MG_CERTS_DB_PORT=5432 -MG_CERTS_DB_USER=absmach -MG_CERTS_DB_PASS=absmach -MG_CERTS_DB=certs -MG_CERTS_DB_SSL_MODE=disable -MG_CERTS_DB_MAX_CONNECTIONS=100 - -## OpenBao Configuration for Certs -MG_CERTS_OPENBAO_HOST=http://openbao:8200 -MG_CERTS_OPENBAO_APP_ROLE=absmach -MG_CERTS_OPENBAO_APP_SECRET=absmach -MG_CERTS_OPENBAO_NAMESPACE= -MG_CERTS_OPENBAO_PKI_PATH=pki -MG_CERTS_OPENBAO_ROLE=absmach -MG_CERTS_OPENBAO_SECRET_ID_TTL=720h -MG_CERTS_SERVICE_TOKEN_PATH=/openbao/service_token -MG_CERTS_SECRET_ID_PATH=/openbao/secret_id -MG_CERTS_SECRET_RENEW_THRESHOLD=24h -MG_CERTS_SECRET_CHECK_INTERVAL=1h - -## OpenBao PKI CA Configuration -MG_CERTS_OPENBAO_PKI_CA_CN=Abstract Machines Certificate Authority -MG_CERTS_OPENBAO_PKI_CA_OU=Abstract Machines -MG_CERTS_OPENBAO_PKI_CA_O=AbstractMachines -MG_CERTS_OPENBAO_PKI_CA_C=FRANCE -MG_CERTS_OPENBAO_PKI_CA_L=PARIS -MG_CERTS_OPENBAO_PKI_CA_ST=PARIS -MG_CERTS_OPENBAO_PKI_CA_ADDR=5 Av. Anatole -MG_CERTS_OPENBAO_PKI_CA_PO=75007 -MG_CERTS_OPENBAO_PKI_CA_DNS_NAMES=localhost -MG_CERTS_OPENBAO_PKI_CA_IP_ADDRESSES=127.0.0.1,::1 -MG_CERTS_OPENBAO_PKI_CA_URI_SANS= -MG_CERTS_OPENBAO_PKI_CA_EMAIL_ADDRESSES=info@abstractmachines.rs - -## OpenBao Unseal Keys and Token -MG_CERTS_OPENBAO_UNSEAL_KEY_1= -MG_CERTS_OPENBAO_UNSEAL_KEY_2= -MG_CERTS_OPENBAO_UNSEAL_KEY_3= -MG_CERTS_OPENBAO_ROOT_TOKEN= - - -#### Auth Client Config for Certs Service +#### Auth gRPC client config for addons MG_ADDONS_CERTS_PATH_PREFIX=../../ MG_AUTH_GRPC_URL=auth:7001 MG_AUTH_GRPC_TIMEOUT=300s @@ -418,16 +408,13 @@ MG_AUTH_GRPC_CLIENT_CERT=${GRPC_MTLS:+./ssl/certs/auth-grpc-client.crt} MG_AUTH_GRPC_CLIENT_KEY=${GRPC_MTLS:+./ssl/certs/auth-grpc-client.key} MG_AUTH_GRPC_SERVER_CA_CERTS=${GRPC_MTLS:+./ssl/certs/ca.crt} -#### Domains Client Config for Certs Service -MG_DOMAINS_GRPC_URL=domains:7003 +#### Domains gRPC client config for addons +MG_DOMAINS_GRPC_URL= MG_DOMAINS_GRPC_TIMEOUT=300s MG_DOMAINS_GRPC_CLIENT_CERT=${GRPC_MTLS:+./ssl/certs/domains-grpc-client.crt} MG_DOMAINS_GRPC_CLIENT_KEY=${GRPC_MTLS:+./ssl/certs/domains-grpc-client.key} MG_DOMAINS_GRPC_SERVER_CA_CERTS=${GRPC_MTLS:+./ssl/certs/ca.crt} -MG_CERTS_JAEGER_FRONTEND=16687 -MG_CERTS_JAEGER_OLTP_HTTP=4319 - ### Postgres MG_POSTGRES_HOST=postgres MG_POSTGRES_PORT=5432 @@ -516,10 +503,9 @@ MG_PROVISION_HTTP_PORT=9016 MG_PROVISION_ENV_CLIENTS_TLS=false MG_PROVISION_SERVER_CERT= MG_PROVISION_SERVER_KEY= -MG_PROVISION_USERS_URL=http://users:9002 -MG_PROVISION_CHANNELS_URL=http://channels:9005 -MG_PROVISION_CLIENTS_URL=http://clients:9006 -MG_PROVISION_CERTS_URL=http://certs:9019 +MG_PROVISION_USERS_URL=http://atom:8080/entities +MG_PROVISION_CHANNELS_URL=http://atom:8080/resources +MG_PROVISION_CLIENTS_URL=http://atom:8080/entities MG_PROVISION_USER= MG_PROVISION_USERNAME= MG_PROVISION_PASS= @@ -529,8 +515,6 @@ MG_PROVISION_BS_SVC_URL=http://bootstrap:9013 MG_PROVISION_BS_CONFIG_PROVISIONING=true MG_PROVISION_BS_AUTO_WHITELIST=true MG_PROVISION_BS_CONTENT= -MG_PROVISION_CERTS_HOURS_VALID=2400h -MG_PROVISION_CERTS_RSA_BITS=2048 MG_PROVISION_INSTANCE_ID= ### Postgres Writer @@ -687,7 +671,7 @@ MG_UI_BACKEND_URL=http://ui-backend:9097 MG_UI_VERIFICATION_TLS=false MG_UI_CONTENT_TYPE=application/senml+json # Set to yes to accept the EULA for the UI services. To view the EULA visit: https://github.com/absmach/eula -MG_UI_DOCKER_ACCEPT_EULA=no +MG_UI_DOCKER_ACCEPT_EULA=yes OTEL_SERVICE_NAME=ui-mg OTEL_EXPORTER_OTLP_ENDPOINT=http://jaeger:4318 @@ -715,20 +699,22 @@ MG_UI_BACKEND_DB_SSL_KEY= MG_UI_BACKEND_DB_SSL_ROOT_CERT= ## UI -MG_AUTH_URL=http://auth:9001 -MG_DOMAINS_URL=http://domains:9003 -MG_USERS_URL=http://users:9002 -MG_CLIENTS_URL=http://clients:9006 -MG_CHANNELS_URL=http://channels:9005 -MG_GROUPS_URL=http://groups:9004 +MG_AUTH_URL=http://nginx:80/auth +MG_ATOM_URL=http://nginx:80 +MG_DOMAINS_URL=http://nginx:80/tenants +MG_USERS_URL=http://nginx:80/entities +MG_CLIENTS_URL=http://nginx:80/entities +MG_CHANNELS_URL=http://nginx:80/resources +MG_GROUPS_URL=http://nginx:80/groups MG_BOOTSTRAP_URL=http://bootstrap:9013 -MG_CERTS_URL=http://certs:9019 MG_HTTP_ADAPTER_URL=http://nginx:80/http +MG_PUBLISH_PROXY_URL=http://nginx:80 MG_READER_URL=http://timescale-reader:9011 MG_JOURNAL_URL=http://journal:9021 ### UI Configuration MG_UI_TYPE=mg +MG_UI_CLIENT_TYPE=client MG_UI_BASE_PATH=/ MG_NEXTAUTH_BASE_PATH=/api/auth NEXTAUTH_SECRET=4WdW0Z0tAOyQ/ZAI3YLVV/wNu+yUZXBLDDQ3AGrgfJ4= diff --git a/docker/README.md b/docker/README.md index 45ba2aef6..116497c4c 100644 --- a/docker/README.md +++ b/docker/README.md @@ -13,13 +13,22 @@ Follow the [official Docker Compose installation guide](https://docs.docker.com/ Run the following commands from the project root directory. ```bash -docker compose -f docker/docker-compose.yaml up +make provision_atom_tokens +make run_latest +``` + +`make provision_atom_tokens` starts Atom, creates per-service Atom API keys, and writes them to the generated `docker/.env.tokens` file. That file is local-only and must not be committed. + +If you use `docker compose` directly instead of the Makefile, pass both env files: + +```bash +docker compose -f docker/docker-compose.yaml --env-file docker/.env --env-file docker/.env.tokens up ``` To start additional addon services: ```bash -docker compose -f docker/addons//docker-compose.yaml up +docker compose -f docker/addons//docker-compose.yaml --env-file docker/.env --env-file docker/.env.tokens up ``` To pull images from a specific release in `ghcr.io/absmach/magistrala`, change `MG_RELEASE_TAG` in `.env` before running these commands. @@ -29,7 +38,7 @@ To pull images from a specific release in `ghcr.io/absmach/magistrala`, change ` Magistrala supports configurable MQTT broker and Message broker, which also acts as an events store. Magistrala uses two types of brokers: 1. **MQTT_BROKER**: Handles MQTT communication between MQTT adapters and message broker. This can either be `RabbitMQ` or `NATS`. -2. **MESSAGE_BROKER**: Manages message exchange between Magistrala core, optional, and external services. This can either be `NATS` or `RabbitMQ`. This is used to store messages for distributed processing. +2. **MESSAGE_BROKER**: Manages message exchange between Magistrala services and external services. This can either be `NATS` or `RabbitMQ`. This is used to store messages for distributed processing. Events store: This is used by Magistrala services to store events for distributed processing. Magistrala uses a single service to be the message broker and events store. This can either be `NATS` or `RabbitMQ`. Redis can also be used as an events store, but it requires a message broker to be deployed along with it for message exchange. @@ -198,7 +207,7 @@ The certbot service keeps running and checks renewal twice a day. When a certifi The included `Makefile` defines build and Docker‑build targets for all Magistrala services. Key points: -- `SERVICES`: list of core services (auth, clients, channels, http, coap, mqtt, ws, etc.) +- `SERVICES`: list of services (auth, clients, channels, http, coap, mqtt, ws, etc.) - `DOCKERS`, `DOCKERS_DEV`: build targets for production and development Docker images - `make dockers`, `make dockers_dev`: always tag images as `ghcr.io/absmach/magistrala/` @@ -215,7 +224,8 @@ make dockers # builds all Docker images Start services with Docker compose: ```bash -docker compose -f docker/docker-compose.yaml up +make provision_atom_tokens +make run_latest ``` To clean up: diff --git a/docker/addons/postgres-reader/docker-compose.yaml b/docker/addons/postgres-reader/docker-compose.yaml index 0856a4fd3..b04ae92bb 100644 --- a/docker/addons/postgres-reader/docker-compose.yaml +++ b/docker/addons/postgres-reader/docker-compose.yaml @@ -31,11 +31,10 @@ services: MG_POSTGRES_SSL_CERT: ${MG_POSTGRES_SSL_CERT} MG_POSTGRES_SSL_KEY: ${MG_POSTGRES_SSL_KEY} MG_POSTGRES_SSL_ROOT_CERT: ${MG_POSTGRES_SSL_ROOT_CERT} - MG_CLIENTS_GRPC_URL: ${MG_CLIENTS_GRPC_URL} - MG_CLIENTS_GRPC_TIMEOUT: ${MG_CLIENTS_GRPC_TIMEOUT} - MG_CLIENTS_GRPC_CLIENT_CERT: ${MG_CLIENTS_GRPC_CLIENT_CERT:+/clients-grpc-client.crt} - MG_CLIENTS_GRPC_CLIENT_KEY: ${MG_CLIENTS_GRPC_CLIENT_KEY:+/clients-grpc-client.key} - MG_CLIENTS_GRPC_SERVER_CA_CERTS: ${MG_CLIENTS_GRPC_SERVER_CA_CERTS:+/clients-grpc-server-ca.crt} + ATOM_URL: ${ATOM_URL} + ATOM_SERVICE_TOKEN: ${MG_ATOM_TOKEN_POSTGRES_READER} + ATOM_JWKS_URL: ${ATOM_JWKS_URL} + ATOM_TIMEOUT: ${ATOM_TIMEOUT} MG_POSTGRES_READER_GRPC_URL: ${MG_POSTGRES_READER_GRPC_URL} MG_POSTGRES_READER_GRPC_PORT: ${MG_POSTGRES_READER_GRPC_PORT} MG_POSTGRES_READER_GRPC_HOST: ${MG_POSTGRES_READER_GRPC_HOST} @@ -46,11 +45,6 @@ services: MG_POSTGRES_READER_GRPC_CLIENT_KEY: ${MG_POSTGRES_READER_GRPC_CLIENT_KEY:+/readers-grpc-client.key} MG_POSTGRES_READER_GRPC_SERVER_CERT: ${MG_POSTGRES_READER_GRPC_SERVER_CERT:+./ssl/certs/readers-grpc-server.crt} MG_POSTGRES_READER_GRPC_SERVER_KEY: ${MG_POSTGRES_READER_GRPC_SERVER_KEY:+./ssl/certs/readers-grpc-server.key} - MG_AUTH_GRPC_URL: ${MG_AUTH_GRPC_URL} - MG_AUTH_GRPC_TIMEOUT: ${MG_AUTH_GRPC_TIMEOUT} - MG_AUTH_GRPC_CLIENT_CERT: ${MG_AUTH_GRPC_CLIENT_CERT:+/auth-grpc-client.crt} - MG_AUTH_GRPC_CLIENT_KEY: ${MG_AUTH_GRPC_CLIENT_KEY:+/auth-grpc-client.key} - MG_AUTH_GRPC_SERVER_CA_CERTS: ${MG_AUTH_GRPC_SERVER_CA_CERTS:+/auth-grpc-server-ca.crt} MG_SEND_TELEMETRY: ${MG_SEND_TELEMETRY} MG_POSTGRES_READER_INSTANCE_ID: ${MG_POSTGRES_READER_INSTANCE_ID} ports: @@ -59,37 +53,6 @@ services: networks: - magistrala-base-net volumes: - - type: bind - source: ${MG_ADDONS_CERTS_PATH_PREFIX}${MG_AUTH_GRPC_CLIENT_CERT:-./ssl/placeholder} - target: /auth-grpc-client${MG_AUTH_GRPC_CLIENT_CERT:+.crt} - bind: - create_host_path: true - - type: bind - source: ${MG_ADDONS_CERTS_PATH_PREFIX}${MG_AUTH_GRPC_CLIENT_KEY:-./ssl/placeholder} - target: /auth-grpc-client${MG_AUTH_GRPC_CLIENT_KEY:+.key} - bind: - create_host_path: true - - type: bind - source: ${MG_ADDONS_CERTS_PATH_PREFIX}${MG_AUTH_GRPC_SERVER_CA_CERTS:-./ssl/placeholder} - target: /auth-grpc-server-ca${MG_AUTH_GRPC_SERVER_CA_CERTS:+.crt} - bind: - create_host_path: true - # Clients gRPC mTLS client certificates - - type: bind - source: ${MG_ADDONS_CERTS_PATH_PREFIX}${MG_CLIENTS_GRPC_CLIENT_CERT:-./ssl/placeholder} - target: /clients-grpc-client${MG_CLIENTS_GRPC_CLIENT_CERT:+.crt} - bind: - create_host_path: true - - type: bind - source: ${MG_ADDONS_CERTS_PATH_PREFIX}${MG_CLIENTS_GRPC_CLIENT_KEY:-./ssl/placeholder} - target: /clients-grpc-client${MG_CLIENTS_GRPC_CLIENT_KEY:+.key} - bind: - create_host_path: true - - type: bind - source: ${MG_ADDONS_CERTS_PATH_PREFIX}${MG_CLIENTS_GRPC_SERVER_CA_CERTS:-./ssl/placeholder} - target: /clients-grpc-server-ca${MG_CLIENTS_GRPC_SERVER_CA_CERTS:+.crt} - bind: - create_host_path: true # Reader gRPC mTLS client certificates - type: bind source: ${MG_POSTGRES_READER_GRPC_SERVER_CERT:-./ssl/placeholder} diff --git a/docker/addons/prometheus/metrics/prometheus.yaml b/docker/addons/prometheus/metrics/prometheus.yaml index 3ea4032c4..84be74d9c 100644 --- a/docker/addons/prometheus/metrics/prometheus.yaml +++ b/docker/addons/prometheus/metrics/prometheus.yaml @@ -15,8 +15,7 @@ scrape_configs: enable_http2: true static_configs: - targets: - - magistrala-clients:9000 - - magistrala-users:9002 + - magistrala-atom:8080 - magistrala-http:8008 - magistrala-ws:8186 - magistrala-coap:5683 diff --git a/docker/addons/provision/docker-compose.yaml b/docker/addons/provision/docker-compose.yaml index ba9a5690f..b59ad75be 100644 --- a/docker/addons/provision/docker-compose.yaml +++ b/docker/addons/provision/docker-compose.yaml @@ -33,13 +33,11 @@ services: MG_PROVISION_USERNAME: ${MG_PROVISION_USERNAME} MG_PROVISION_PASS: ${MG_PROVISION_PASS} MG_PROVISION_API_KEY: ${MG_PROVISION_API_KEY} - MG_PROVISION_CERTS_URL: ${MG_PROVISION_CERTS_URL} MG_PROVISION_X509_PROVISIONING: ${MG_PROVISION_X509_PROVISIONING} MG_PROVISION_BS_SVC_URL: ${MG_PROVISION_BS_SVC_URL} MG_PROVISION_BS_CONFIG_PROVISIONING: ${MG_PROVISION_BS_CONFIG_PROVISIONING} MG_PROVISION_BS_AUTO_WHITELIST: ${MG_PROVISION_BS_AUTO_WHITELIST} MG_PROVISION_BS_CONTENT: ${MG_PROVISION_BS_CONTENT} - MG_PROVISION_CERTS_HOURS_VALID: ${MG_PROVISION_CERTS_HOURS_VALID} MG_SEND_TELEMETRY: ${MG_SEND_TELEMETRY} MG_PROVISION_INSTANCE_ID: ${MG_PROVISION_INSTANCE_ID} MG_AUTH_GRPC_URL: ${MG_AUTH_GRPC_URL} diff --git a/docker/docker-compose-ci.yaml b/docker/docker-compose-ci.yaml new file mode 100644 index 000000000..19a6f7757 --- /dev/null +++ b/docker/docker-compose-ci.yaml @@ -0,0 +1,7 @@ +# Copyright (c) Abstract Machines +# SPDX-License-Identifier: Apache-2.0 + +services: + atom-ui: + profiles: + - atom-ui diff --git a/docker/docker-compose.yaml b/docker/docker-compose.yaml index ffc422d3d..39131e9c4 100644 --- a/docker/docker-compose.yaml +++ b/docker/docker-compose.yaml @@ -12,405 +12,121 @@ networks: - subnet: 172.30.0.0/24 volumes: - magistrala-users-db-volume: - magistrala-groups-db-volume: - magistrala-clients-db-volume: - magistrala-channels-db-volume: - magistrala-channels-redis-volume: - magistrala-clients-redis-volume: - magistrala-spicedb-db-volume: - magistrala-auth-db-volume: magistrala-pat-db-volume: - magistrala-domains-db-volume: - magistrala-domains-redis-volume: - magistrala-auth-redis-volume: - magistrala-auth-keys-volume: magistrala-ui-backend-db-volume: magistrala-journal-volume: magistrala-re-db-volume: magistrala-alarms-db-volume: magistrala-reports-db-volume: - magistrala-certs-db-volume: - magistrala-openbao-data: magistrala-timescale-writer-volume: magistrala-fluxmq-node1-volume: magistrala-fluxmq-node2-volume: magistrala-fluxmq-node3-volume: + magistrala-atom-db-volume: services: - spicedb: - image: docker.io/authzed/spicedb:v1.50.0 - container_name: magistrala-spicedb - command: "serve" - restart: "always" + atom-db: + image: postgres:16-alpine + container_name: magistrala-atom-db + restart: on-failure + environment: + POSTGRES_USER: ${ATOM_DB_USER} + POSTGRES_PASSWORD: ${ATOM_DB_PASSWORD} + POSTGRES_DB: ${ATOM_DB_NAME} + ports: + - ${ATOM_DB_PORT}:5432 networks: - magistrala-base-net - ports: - - "8080:8080" - - "9091:9090" - - "50051:50051" - environment: - SPICEDB_GRPC_PRESHARED_KEY: ${MG_SPICEDB_PRE_SHARED_KEY} - SPICEDB_DATASTORE_ENGINE: ${MG_SPICEDB_DATASTORE_ENGINE} - SPICEDB_DATASTORE_CONN_URI: "${MG_SPICEDB_DATASTORE_ENGINE}://${MG_SPICEDB_DB_USER}:${MG_SPICEDB_DB_PASS}@spicedb-db:${MG_SPICEDB_DB_PORT}/${MG_SPICEDB_DB_NAME}?sslmode=disable" + volumes: + - magistrala-atom-db-volume:/var/lib/postgresql/data + healthcheck: + test: ["CMD-SHELL", "pg_isready -U ${ATOM_DB_USER} -d ${ATOM_DB_NAME}"] + interval: 5s + timeout: 5s + retries: 5 + + atom: + image: ghcr.io/absmach/atom:latest + container_name: magistrala-atom + restart: on-failure depends_on: - - spicedb-migrate - - spicedb-migrate: - image: docker.io/authzed/spicedb:v1.50.0 - container_name: magistrala-spicedb-migrate - command: "migrate head" - restart: "on-failure" + atom-db: + condition: service_healthy + environment: + DATABASE_URL: postgres://${ATOM_DB_USER}:${ATOM_DB_PASSWORD}@atom-db:5432/${ATOM_DB_NAME} + LISTEN_ADDR: 0.0.0.0:8080 + GRPC_ADDR: ${ATOM_GRPC_ADDR:-0.0.0.0:8081} + JWT_SECRET: ${ATOM_JWT_SECRET} + JWT_EXPIRY_SECS: ${ATOM_JWT_EXPIRY_SECS} + ATOM_PUBLIC_BASE_URL: ${ATOM_PUBLIC_URL} + ATOM_JWT_ISSUER: ${ATOM_JWT_ISSUER} + ATOM_JWT_AUDIENCE: ${ATOM_JWT_AUDIENCE} + ATOM_SIGNUP_ENABLED: ${ATOM_SIGNUP_ENABLED} + ATOM_ALLOW_UNVERIFIED_EMAIL_LOGIN: ${ATOM_ALLOW_UNVERIFIED_EMAIL_LOGIN} + ATOM_DEV_ALLOW_UNVERIFIED_EMAIL_LOGIN: ${ATOM_DEV_ALLOW_UNVERIFIED_EMAIL_LOGIN} + ATOM_CORS_ALLOWED_ORIGINS: ${ATOM_CORS_ALLOWED_ORIGINS} + ATOM_INVITATION_REDIRECT: ${ATOM_INVITATION_REDIRECT} + ATOM_INVITATION_EXPIRY_SECS: ${ATOM_INVITATION_EXPIRY_SECS} + ATOM_SMTP_HOST: ${ATOM_SMTP_HOST} + ATOM_SMTP_PORT: ${ATOM_SMTP_PORT} + ATOM_SMTP_USERNAME: ${ATOM_SMTP_USERNAME} + ATOM_SMTP_PASSWORD: ${ATOM_SMTP_PASSWORD} + ATOM_SMTP_FROM: ${ATOM_SMTP_FROM} + ATOM_SMTP_TLS: ${ATOM_SMTP_TLS} + ADMIN_SECRET: ${ATOM_ADMIN_SECRET} + ATOM_KEY_ENCRYPTION_KEY: ${ATOM_KEY_ENCRYPTION_KEY} + ATOM_KEY_ENCRYPTION_KEY_ID: ${ATOM_KEY_ENCRYPTION_KEY_ID:-local:v1} + ATOM_ALLOW_PLAINTEXT_SIGNING_KEYS: ${ATOM_ALLOW_PLAINTEXT_SIGNING_KEYS:-false} + ATOM_SERVICE_SECRET: ${ATOM_SERVICE_SECRET} + ATOM_CERTS_ENABLED: ${ATOM_CERTS_ENABLED:-true} + ATOM_CERTS_CA_MODE: ${ATOM_CERTS_CA_MODE:-file_root_issuer} + ATOM_CERTS_ROOT_CA_CERT_PATH: ${ATOM_CERTS_ROOT_CA_CERT_PATH:-/certs/ca.crt} + ATOM_CERTS_ROOT_CA_KEY_PATH: ${ATOM_CERTS_ROOT_CA_KEY_PATH:-/certs/ca.key} + ATOM_CERTS_LEAF_DEFAULT_TTL_SECS: ${ATOM_CERTS_LEAF_DEFAULT_TTL_SECS:-2592000} + ATOM_CERTS_LEAF_MAX_TTL_SECS: ${ATOM_CERTS_LEAF_MAX_TTL_SECS:-2592000} + ATOM_RATE_LIMIT_ENABLED: ${ATOM_RATE_LIMIT_ENABLED:-true} + ATOM_TRUSTED_PROXY_CIDRS: ${ATOM_TRUSTED_PROXY_CIDRS:-} + ATOM_MIN_PASSWORD_CHARS: ${ATOM_MIN_PASSWORD_CHARS} + RUST_LOG: ${ATOM_RUST_LOG} + ports: + - ${ATOM_HTTP_PORT}:8080 + volumes: + - ${ATOM_CERTS_CA_DIR:-./ssl/certs}:/certs:ro networks: - magistrala-base-net - environment: - SPICEDB_DATASTORE_ENGINE: ${MG_SPICEDB_DATASTORE_ENGINE} - SPICEDB_DATASTORE_CONN_URI: "${MG_SPICEDB_DATASTORE_ENGINE}://${MG_SPICEDB_DB_USER}:${MG_SPICEDB_DB_PASS}@spicedb-db:${MG_SPICEDB_DB_PORT}/${MG_SPICEDB_DB_NAME}?sslmode=disable" + + atom-ui: + image: ghcr.io/absmach/atom-ui:latest + container_name: magistrala-atom-ui + restart: on-failure depends_on: - - spicedb-db - - spicedb-db: - image: docker.io/postgres:18.0-alpine3.22 - container_name: magistrala-spicedb-db - networks: - - magistrala-base-net - ports: - - "6010:5432" + - atom environment: - POSTGRES_USER: ${MG_SPICEDB_DB_USER} - POSTGRES_PASSWORD: ${MG_SPICEDB_DB_PASS} - POSTGRES_DB: ${MG_SPICEDB_DB_NAME} - volumes: - - magistrala-spicedb-db-volume:/var/lib/postgresql/data - command: ["postgres", "-c", "track_commit_timestamp=on"] - - auth-db: - image: docker.io/postgres:18.0-alpine3.22 - container_name: magistrala-auth-db - restart: on-failure + ATOM_GRAPHQL_URL: http://atom:8080/graphql ports: - - 6001:5432 - environment: - POSTGRES_USER: ${MG_AUTH_DB_USER} - POSTGRES_PASSWORD: ${MG_AUTH_DB_PASS} - POSTGRES_DB: ${MG_AUTH_DB_NAME} + - ${ATOM_UI_HTTP_PORT:-3005}:3000 networks: - magistrala-base-net - volumes: - - magistrala-auth-db-volume:/var/lib/postgresql/data - auth-redis: - image: docker.io/redis:8.2.2-alpine3.22 - container_name: magistrala-auth-redis + atom-bootstrap: + image: ghcr.io/absmach/magistrala/atom-bootstrap:${MG_RELEASE_TAG} + container_name: magistrala-atom-bootstrap restart: on-failure - networks: - - magistrala-base-net - volumes: - - magistrala-auth-redis-volume:/data - - ./redis/redis.conf:/etc/redis/redis.conf:ro - command: ["redis-server", "/etc/redis/redis.conf"] - - auth: - image: ghcr.io/absmach/magistrala/auth:${MG_RELEASE_TAG} - container_name: magistrala-auth depends_on: - - auth-db - - spicedb - - nginx - expose: - - ${MG_AUTH_GRPC_PORT} - restart: on-failure + - atom environment: - MG_AUTH_LOG_LEVEL: ${MG_AUTH_LOG_LEVEL} - MG_SPICEDB_SCHEMA_FILE: ${MG_SPICEDB_SCHEMA_FILE} - MG_SPICEDB_PRE_SHARED_KEY: ${MG_SPICEDB_PRE_SHARED_KEY} - MG_SPICEDB_HOST: ${MG_SPICEDB_HOST} - MG_SPICEDB_PORT: ${MG_SPICEDB_PORT} - MG_AUTH_INVITATION_DURATION: ${MG_AUTH_INVITATION_DURATION} - MG_AUTH_HTTP_HOST: ${MG_AUTH_HTTP_HOST} - MG_AUTH_HTTP_PORT: ${MG_AUTH_HTTP_PORT} - MG_AUTH_HTTP_SERVER_CERT: ${MG_AUTH_HTTP_SERVER_CERT} - MG_AUTH_HTTP_SERVER_KEY: ${MG_AUTH_HTTP_SERVER_KEY} - MG_AUTH_GRPC_HOST: ${MG_AUTH_GRPC_HOST} - MG_AUTH_GRPC_PORT: ${MG_AUTH_GRPC_PORT} - MG_AUTH_ACCESS_TOKEN_DURATION: ${MG_AUTH_ACCESS_TOKEN_DURATION} - MG_AUTH_REFRESH_TOKEN_DURATION: ${MG_AUTH_REFRESH_TOKEN_DURATION} - MG_AUTH_KEYS_ALGORITHM: ${MG_AUTH_KEYS_ALGORITHM} - MG_AUTH_KEYS_ACTIVE_KEY_PATH: ${MG_AUTH_KEYS_ACTIVE_KEY_PATH:+/keys/active.key} - MG_AUTH_KEYS_RETIRING_KEY_PATH: ${MG_AUTH_KEYS_RETIRING_KEY_PATH:+/keys/retiring.key} - ## Compose supports parameter expansion in environment, - ## Eg: ${VAR:+replacement} or ${VAR+replacement} -> replacement if VAR is set and non-empty, otherwise empty - ## Eg :${VAR:-default} or ${VAR-default} -> value of VAR if set and non-empty, otherwise default - MG_AUTH_GRPC_SERVER_CERT: ${MG_AUTH_GRPC_SERVER_CERT:+/auth-grpc-server.crt} - MG_AUTH_GRPC_SERVER_KEY: ${MG_AUTH_GRPC_SERVER_KEY:+/auth-grpc-server.key} - MG_AUTH_GRPC_SERVER_CA_CERTS: ${MG_AUTH_GRPC_SERVER_CA_CERTS:+/auth-grpc-server-ca.crt} - MG_AUTH_GRPC_CLIENT_CA_CERTS: ${MG_AUTH_GRPC_CLIENT_CA_CERTS:+/auth-grpc-client-ca.crt} - MG_AUTH_DB_HOST: ${MG_AUTH_DB_HOST} - MG_AUTH_DB_PORT: ${MG_AUTH_DB_PORT} - MG_AUTH_DB_USER: ${MG_AUTH_DB_USER} - MG_AUTH_DB_PASS: ${MG_AUTH_DB_PASS} - MG_AUTH_DB_NAME: ${MG_AUTH_DB_NAME} - MG_AUTH_DB_SSL_MODE: ${MG_AUTH_DB_SSL_MODE} - MG_AUTH_DB_SSL_CERT: ${MG_AUTH_DB_SSL_CERT} - MG_AUTH_DB_SSL_KEY: ${MG_AUTH_DB_SSL_KEY} - MG_AUTH_DB_SSL_ROOT_CERT: ${MG_AUTH_DB_SSL_ROOT_CERT} - MG_JAEGER_URL: ${MG_JAEGER_URL} - MG_JAEGER_TRACE_RATIO: ${MG_JAEGER_TRACE_RATIO} - MG_SEND_TELEMETRY: ${MG_SEND_TELEMETRY} - MG_AUTH_ADAPTER_INSTANCE_ID: ${MG_AUTH_ADAPTER_INSTANCE_ID} - MG_ES_URL: ${MG_ES_URL} - MG_AUTH_CACHE_URL: ${MG_AUTH_CACHE_URL} - ports: - - ${MG_AUTH_HTTP_PORT}:${MG_AUTH_HTTP_PORT} - - ${MG_AUTH_GRPC_PORT}:${MG_AUTH_GRPC_PORT} + ATOM_URL: ${ATOM_URL} + ATOM_SERVICE_USERNAME: ${ATOM_SERVICE_USERNAME} + ATOM_SERVICE_SECRET: ${ATOM_SERVICE_SECRET} + ATOM_ADMIN_TOKEN: ${ATOM_ADMIN_TOKEN} + ATOM_ADMIN_USERNAME: ${ATOM_ADMIN_USERNAME} + ATOM_ADMIN_SECRET: ${ATOM_ADMIN_SECRET} + ATOM_TIMEOUT: ${ATOM_TIMEOUT} + MG_ATOM_BOOTSTRAP_RETRIES: ${MG_ATOM_BOOTSTRAP_RETRIES:-30} + MG_ATOM_BOOTSTRAP_RETRY_INTERVAL: ${MG_ATOM_BOOTSTRAP_RETRY_INTERVAL:-2s} + MG_ATOM_BOOTSTRAP_TIMEOUT: ${MG_ATOM_BOOTSTRAP_TIMEOUT:-30s} networks: - magistrala-base-net - volumes: - - ./spicedb/schema.zed:${MG_SPICEDB_SCHEMA_FILE} - - magistrala-pat-db-volume:/magistrala-data - # Auth active private key file - - type: bind - source: ${MG_AUTH_KEYS_ACTIVE_KEY_PATH} - target: /keys/active.key - read_only: true - # Auth retiring private key file (optional, for key rotation) - - type: bind - source: ${MG_AUTH_KEYS_RETIRING_KEY_PATH:-./ssl/placeholder} - target: /keys/retiring.key - read_only: true - bind: - create_host_path: true - # Auth gRPC mTLS server certificates - - type: bind - source: ${MG_AUTH_GRPC_SERVER_CERT:-./ssl/placeholder} - target: /auth-grpc-server.crt - bind: - create_host_path: true - - type: bind - source: ${MG_AUTH_GRPC_SERVER_KEY:-./ssl/placeholder} - target: /auth-grpc-server.key - bind: - create_host_path: true - - type: bind - source: ${MG_AUTH_GRPC_SERVER_CA_CERTS:-./ssl/placeholder} - target: /auth-grpc-server-ca.crt - bind: - create_host_path: true - - type: bind - source: ${MG_AUTH_GRPC_CLIENT_CA_CERTS:-./ssl/placeholder} - target: /auth-grpc-client-ca.crt - bind: - create_host_path: true - # Auth Callout Client Certificates - - type: bind - source: ${MG_AUTH_CALLOUT_CLIENT_CERT:-./ssl/placeholder} - target: /auth-callout-client.crt - bind: - create_host_path: true - - type: bind - source: ${MG_AUTH_CALLOUT_CLIENT_KEY:-./ssl/placeholder} - target: /auth-callout-client.key - bind: - create_host_path: true - - type: bind - source: ${MG_AUTH_CALLOUT_CLIENT_CA_CERTS:-./ssl/placeholder} - target: /auth-callout-client-ca.crt - bind: - create_host_path: true - - domains-db: - image: docker.io/postgres:18.0-alpine3.22 - container_name: magistrala-domains-db - restart: on-failure - ports: - - 6003:5432 - environment: - POSTGRES_USER: ${MG_DOMAINS_DB_USER} - POSTGRES_PASSWORD: ${MG_DOMAINS_DB_PASS} - POSTGRES_DB: ${MG_DOMAINS_DB_NAME} - networks: - - magistrala-base-net - volumes: - - magistrala-domains-db-volume:/var/lib/postgresql/data - - domains-redis: - image: docker.io/redis:8.2.2-alpine3.22 - container_name: magistrala-domains-redis - restart: on-failure - networks: - - magistrala-base-net - volumes: - - magistrala-domains-redis-volume:/data - - domains: - image: ghcr.io/absmach/magistrala/domains:${MG_RELEASE_TAG} - container_name: magistrala-domains - depends_on: - - domains-db - - spicedb - - nginx - expose: - - ${MG_DOMAINS_GRPC_PORT} - restart: on-failure - environment: - MG_DOMAINS_LOG_LEVEL: ${MG_DOMAINS_LOG_LEVEL} - MG_SPICEDB_PRE_SHARED_KEY: ${MG_SPICEDB_PRE_SHARED_KEY} - MG_SPICEDB_HOST: ${MG_SPICEDB_HOST} - MG_SPICEDB_PORT: ${MG_SPICEDB_PORT} - MG_SPICEDB_SCHEMA_FILE: ${MG_SPICEDB_SCHEMA_FILE} - MG_DOMAINS_HTTP_HOST: ${MG_DOMAINS_HTTP_HOST} - MG_DOMAINS_HTTP_PORT: ${MG_DOMAINS_HTTP_PORT} - MG_DOMAINS_HTTP_SERVER_CERT: ${MG_DOMAINS_HTTP_SERVER_CERT} - MG_DOMAINS_HTTP_SERVER_KEY: ${MG_DOMAINS_HTTP_SERVER_KEY} - MG_DOMAINS_GRPC_HOST: ${MG_DOMAINS_GRPC_HOST} - MG_DOMAINS_GRPC_PORT: ${MG_DOMAINS_GRPC_PORT} - ## Compose supports parameter expansion in environment, - ## Eg: ${VAR:+replacement} or ${VAR+replacement} -> replacement if VAR is set and non-empty, otherwise empty - ## Eg :${VAR:-default} or ${VAR-default} -> value of VAR if set and non-empty, otherwise default - MG_DOMAINS_GRPC_SERVER_CERT: ${MG_DOMAINS_GRPC_SERVER_CERT:+/domains-grpc-server.crt} - MG_DOMAINS_GRPC_SERVER_KEY: ${MG_DOMAINS_GRPC_SERVER_KEY:+/domains-grpc-server.key} - MG_DOMAINS_GRPC_SERVER_CA_CERTS: ${MG_DOMAINS_GRPC_SERVER_CA_CERTS:+/domains-grpc-server-ca.crt} - MG_DOMAINS_GRPC_CLIENT_CA_CERTS: ${MG_DOMAINS_GRPC_CLIENT_CA_CERTS:+/domains-grpc-client-ca.crt} - MG_DOMAINS_DB_HOST: ${MG_DOMAINS_DB_HOST} - MG_DOMAINS_DB_PORT: ${MG_DOMAINS_DB_PORT} - MG_DOMAINS_DB_USER: ${MG_DOMAINS_DB_USER} - MG_DOMAINS_DB_PASS: ${MG_DOMAINS_DB_PASS} - MG_DOMAINS_DB_NAME: ${MG_DOMAINS_DB_NAME} - MG_DOMAINS_DB_SSL_MODE: ${MG_DOMAINS_DB_SSL_MODE} - MG_DOMAINS_DB_SSL_CERT: ${MG_DOMAINS_DB_SSL_CERT} - MG_DOMAINS_DB_SSL_KEY: ${MG_DOMAINS_DB_SSL_KEY} - MG_DOMAINS_DB_SSL_ROOT_CERT: ${MG_DOMAINS_DB_SSL_ROOT_CERT} - MG_DOMAINS_INSTANCE_ID: ${MG_DOMAINS_INSTANCE_ID} - MG_ES_URL: ${MG_ES_URL} - MG_DOMAINS_CACHE_URL: ${MG_DOMAINS_CACHE_URL} - MG_DOMAINS_CACHE_KEY_DURATION: ${MG_DOMAINS_CACHE_KEY_DURATION} - MG_AUTH_GRPC_URL: ${MG_AUTH_GRPC_URL} - MG_AUTH_GRPC_TIMEOUT: ${MG_AUTH_GRPC_TIMEOUT} - MG_AUTH_GRPC_CLIENT_CERT: ${MG_AUTH_GRPC_CLIENT_CERT:+/auth-grpc-client.crt} - MG_AUTH_GRPC_CLIENT_KEY: ${MG_AUTH_GRPC_CLIENT_KEY:+/auth-grpc-client.key} - MG_AUTH_GRPC_SERVER_CA_CERTS: ${MG_AUTH_GRPC_SERVER_CA_CERTS:+/auth-grpc-server-ca.crt} - MG_AUTH_KEYS_ALGORITHM: ${MG_AUTH_KEYS_ALGORITHM} - MG_GROUPS_GRPC_URL: ${MG_GROUPS_GRPC_URL} - MG_GROUPS_GRPC_TIMEOUT: ${MG_GROUPS_GRPC_TIMEOUT} - MG_GROUPS_GRPC_CLIENT_CERT: ${MG_GROUPS_GRPC_CLIENT_CERT:+/groups-grpc-client.crt} - MG_GROUPS_GRPC_CLIENT_KEY: ${MG_GROUPS_GRPC_CLIENT_KEY:+/groups-grpc-client.key} - MG_GROUPS_GRPC_SERVER_CA_CERTS: ${MG_GROUPS_GRPC_SERVER_CA_CERTS:+/groups-grpc-server-ca.crt} - MG_CHANNELS_URL: ${MG_CHANNELS_URL} - MG_CHANNELS_GRPC_URL: ${MG_CHANNELS_GRPC_URL} - MG_CHANNELS_GRPC_TIMEOUT: ${MG_CHANNELS_GRPC_TIMEOUT} - MG_CHANNELS_GRPC_CLIENT_CERT: ${MG_CHANNELS_GRPC_CLIENT_CERT:+/channels-grpc-client.crt} - MG_CHANNELS_GRPC_CLIENT_KEY: ${MG_CHANNELS_GRPC_CLIENT_KEY:+/channels-grpc-client.key} - MG_CHANNELS_GRPC_SERVER_CA_CERTS: ${MG_CHANNELS_GRPC_SERVER_CA_CERTS:+/channels-grpc-server-ca.crt} - MG_CLIENTS_GRPC_URL: ${MG_CLIENTS_GRPC_URL} - MG_CLIENTS_GRPC_TIMEOUT: ${MG_CLIENTS_GRPC_TIMEOUT} - MG_CLIENTS_GRPC_CLIENT_CERT: ${MG_CLIENTS_GRPC_CLIENT_CERT:+/clients-grpc-client.crt} - MG_CLIENTS_GRPC_CLIENT_KEY: ${MG_CLIENTS_GRPC_CLIENT_KEY:+/clients-grpc-client.key} - MG_CLIENTS_GRPC_SERVER_CA_CERTS: ${MG_CLIENTS_GRPC_SERVER_CA_CERTS:+/clients-grpc-server-ca.crt} - MG_JAEGER_URL: ${MG_JAEGER_URL} - MG_JAEGER_TRACE_RATIO: ${MG_JAEGER_TRACE_RATIO} - MG_SEND_TELEMETRY: ${MG_SEND_TELEMETRY} - MG_DOMAINS_CALLOUT_URLS: ${MG_DOMAINS_CALLOUT_URLS} - MG_DOMAINS_CALLOUT_METHOD: ${MG_DOMAINS_CALLOUT_METHOD} - MG_DOMAINS_CALLOUT_TLS_VERIFICATION: ${MG_DOMAINS_CALLOUT_TLS_VERIFICATION} - MG_DOMAINS_CALLOUT_TIMEOUT: ${MG_DOMAINS_CALLOUT_TIMEOUT} - MG_DOMAINS_CALLOUT_CA_CERT: ${MG_DOMAINS_CALLOUT_CA_CERT} - MG_DOMAINS_CALLOUT_CERT: ${MG_DOMAINS_CALLOUT_CERT} - MG_DOMAINS_CALLOUT_KEY: ${MG_DOMAINS_CALLOUT_KEY} - MG_DOMAINS_CALLOUT_OPERATIONS: ${MG_DOMAINS_CALLOUT_OPERATIONS} - MG_ALLOW_UNVERIFIED_USER: ${MG_ALLOW_UNVERIFIED_USER} - ports: - - ${MG_DOMAINS_HTTP_PORT}:${MG_DOMAINS_HTTP_PORT} - - ${MG_DOMAINS_GRPC_PORT}:${MG_DOMAINS_GRPC_PORT} - networks: - - magistrala-base-net - volumes: - - ./permission.yaml:/permission.yaml - - ./spicedb/schema.zed:${MG_SPICEDB_SCHEMA_FILE} - # Domains gRPC mTLS server certificates - - type: bind - source: ${MG_DOMAINS_GRPC_SERVER_CERT:-./ssl/placeholder} - target: /domains-grpc-server.crt - bind: - create_host_path: true - - type: bind - source: ${MG_DOMAINS_GRPC_SERVER_KEY:-./ssl/placeholder} - target: /domains-grpc-server.key - bind: - create_host_path: true - - type: bind - source: ${MG_DOMAINS_GRPC_SERVER_CA_CERTS:-./ssl/placeholder} - target: /domains-grpc-server-ca.crt - bind: - create_host_path: true - - type: bind - source: ${MG_DOMAINS_GRPC_CLIENT_CA_CERTS:-./ssl/placeholder} - target: /domains-grpc-client-ca.crt - bind: - create_host_path: true - # Auth gRPC client certificates - - type: bind - source: ${MG_AUTH_GRPC_CLIENT_CERT:-./ssl/placeholder} - target: /auth-grpc-client.crt - bind: - create_host_path: true - - type: bind - source: ${MG_AUTH_GRPC_CLIENT_KEY:-./ssl/placeholder} - target: /auth-grpc-client.key - bind: - create_host_path: true - - type: bind - source: ${MG_AUTH_GRPC_SERVER_CA_CERTS:-./ssl/placeholder} - target: /auth-grpc-server-ca.crt - bind: - create_host_path: true - # Groups gRPC client certificates - - type: bind - source: ${MG_GROUPS_GRPC_CLIENT_CERT:-./ssl/placeholder} - target: /groups-grpc-client.crt - bind: - create_host_path: true - - type: bind - source: ${MG_GROUPS_GRPC_CLIENT_KEY:-./ssl/placeholder} - target: /groups-grpc-client.key - bind: - create_host_path: true - - type: bind - source: ${MG_GROUPS_GRPC_SERVER_CA_CERTS:-./ssl/placeholder} - target: /groups-grpc-server-ca.crt - bind: - create_host_path: true - # Channels gRPC client certificates - - type: bind - source: ${MG_CHANNELS_GRPC_CLIENT_CERT:-./ssl/placeholder} - target: /channels-grpc-client.crt - bind: - create_host_path: true - - type: bind - source: ${MG_CHANNELS_GRPC_CLIENT_KEY:-./ssl/placeholder} - target: /channels-grpc-client.key - bind: - create_host_path: true - - type: bind - source: ${MG_CHANNELS_GRPC_SERVER_CA_CERTS:-./ssl/placeholder} - target: /channels-grpc-server-ca.crt - bind: - create_host_path: true - # Clients gRPC client certificates - - type: bind - source: ${MG_CLIENTS_GRPC_CLIENT_CERT:-./ssl/placeholder} - target: /clients-grpc-client.crt - bind: - create_host_path: true - - type: bind - source: ${MG_CLIENTS_GRPC_CLIENT_KEY:-./ssl/placeholder} - target: /clients-grpc-client.key - bind: - create_host_path: true - - type: bind - source: ${MG_CLIENTS_GRPC_SERVER_CA_CERTS:-./ssl/placeholder} - target: /clients-grpc-server-ca.crt - bind: - create_host_path: true journal-db: image: postgres:16.2-alpine @@ -431,10 +147,12 @@ services: image: ghcr.io/absmach/magistrala/journal:${MG_RELEASE_TAG} container_name: magistrala-journal depends_on: - - journal-db - - auth - - domains - - nginx + journal-db: + condition: service_started + atom-bootstrap: + condition: service_completed_successfully + nginx: + condition: service_started restart: on-failure environment: MG_JOURNAL_LOG_LEVEL: ${MG_JOURNAL_LOG_LEVEL} @@ -451,59 +169,22 @@ services: MG_JOURNAL_DB_SSL_CERT: ${MG_JOURNAL_DB_SSL_CERT} MG_JOURNAL_DB_SSL_KEY: ${MG_JOURNAL_DB_SSL_KEY} MG_JOURNAL_DB_SSL_ROOT_CERT: ${MG_JOURNAL_DB_SSL_ROOT_CERT} - MG_AUTH_GRPC_URL: ${MG_AUTH_GRPC_URL} - MG_AUTH_GRPC_TIMEOUT: ${MG_AUTH_GRPC_TIMEOUT} - MG_AUTH_GRPC_CLIENT_CERT: ${MG_AUTH_GRPC_CLIENT_CERT:+/auth-grpc-client.crt} - MG_AUTH_GRPC_CLIENT_KEY: ${MG_AUTH_GRPC_CLIENT_KEY:+/auth-grpc-client.key} - MG_AUTH_GRPC_SERVER_CA_CERTS: ${MG_AUTH_GRPC_SERVER_CA_CERTS:+/auth-grpc-server-ca.crt} - MG_AUTH_KEYS_ALGORITHM: ${MG_AUTH_KEYS_ALGORITHM} + ATOM_URL: ${ATOM_URL} + ATOM_SERVICE_TOKEN: ${MG_ATOM_TOKEN_JOURNAL} + ATOM_JWKS_URL: ${ATOM_JWKS_URL} + ATOM_JWT_ISSUER: ${ATOM_JWT_ISSUER} + ATOM_JWT_AUDIENCE: ${ATOM_JWT_AUDIENCE} + ATOM_TIMEOUT: ${ATOM_TIMEOUT} MG_ES_URL: ${MG_ES_URL} MG_JAEGER_URL: ${MG_JAEGER_URL} MG_JAEGER_TRACE_RATIO: ${MG_JAEGER_TRACE_RATIO} MG_SEND_TELEMETRY: ${MG_SEND_TELEMETRY} MG_JOURNAL_INSTANCE_ID: ${MG_JOURNAL_INSTANCE_ID} - MG_DOMAINS_GRPC_URL: ${MG_DOMAINS_GRPC_URL} - MG_DOMAINS_GRPC_TIMEOUT: ${MG_DOMAINS_GRPC_TIMEOUT} - MG_DOMAINS_GRPC_CLIENT_CERT: ${MG_DOMAINS_GRPC_CLIENT_CERT:+/domains-grpc-client.crt} - MG_DOMAINS_GRPC_CLIENT_KEY: ${MG_DOMAINS_GRPC_CLIENT_KEY:+/domains-grpc-client.key} - MG_DOMAINS_GRPC_SERVER_CA_CERTS: ${MG_DOMAINS_GRPC_SERVER_CA_CERTS:+/domains-grpc-server-ca.crt} MG_ALLOW_UNVERIFIED_USER: ${MG_ALLOW_UNVERIFIED_USER} ports: - ${MG_JOURNAL_HTTP_PORT}:${MG_JOURNAL_HTTP_PORT} networks: - magistrala-base-net - volumes: - - type: bind - source: ${MG_AUTH_GRPC_CLIENT_CERT:-./ssl/placeholder} - target: /auth-grpc-client.crt - bind: - create_host_path: true - - type: bind - source: ${MG_AUTH_GRPC_CLIENT_KEY:-./ssl/placeholder} - target: /auth-grpc-client.key - bind: - create_host_path: true - - type: bind - source: ${MG_AUTH_GRPC_SERVER_CA_CERTS:-./ssl/placeholder} - target: /auth-grpc-server-ca.crt - bind: - create_host_path: true - - type: bind - source: ${MG_DOMAINS_GRPC_CLIENT_CERT:-./ssl/placeholder} - target: /domains-grpc-client.crt - bind: - create_host_path: true - - type: bind - source: ${MG_DOMAINS_GRPC_CLIENT_KEY:-./ssl/placeholder} - target: /domains-grpc-client.key - bind: - create_host_path: true - - type: bind - source: ${MG_DOMAINS_GRPC_SERVER_CA_CERTS:-./ssl/placeholder} - target: /domains-grpc-server-ca.crt - bind: - create_host_path: true - nginx: image: docker.io/nginx:1.29.2-alpine3.22 container_name: magistrala-nginx @@ -603,555 +284,14 @@ services: sleep 12h & wait $$! done - clients-db: - image: docker.io/postgres:18.0-alpine3.22 - container_name: magistrala-clients-db - restart: on-failure - command: postgres -c "max_connections=${MG_POSTGRES_MAX_CONNECTIONS}" - environment: - POSTGRES_USER: ${MG_CLIENTS_DB_USER} - POSTGRES_PASSWORD: ${MG_CLIENTS_DB_PASS} - POSTGRES_DB: ${MG_CLIENTS_DB_NAME} - MG_POSTGRES_MAX_CONNECTIONS: ${MG_POSTGRES_MAX_CONNECTIONS} - networks: - - magistrala-base-net - ports: - - 6006:5432 - volumes: - - magistrala-clients-db-volume:/var/lib/postgresql/data - - clients-redis: - image: docker.io/redis:8.2.2-alpine3.22 - container_name: magistrala-clients-redis - restart: on-failure - networks: - - magistrala-base-net - volumes: - - magistrala-clients-redis-volume:/data - - clients: - image: ghcr.io/absmach/magistrala/clients:${MG_RELEASE_TAG} - container_name: magistrala-clients - depends_on: - - clients-db - - users - - auth - - nginx - restart: on-failure - environment: - MG_CLIENTS_LOG_LEVEL: ${MG_CLIENTS_LOG_LEVEL} - MG_CLIENTS_STANDALONE_ID: ${MG_CLIENTS_STANDALONE_ID} - MG_CLIENTS_STANDALONE_TOKEN: ${MG_CLIENTS_STANDALONE_TOKEN} - MG_CLIENTS_CACHE_KEY_DURATION: ${MG_CLIENTS_CACHE_KEY_DURATION} - MG_CLIENTS_HTTP_HOST: ${MG_CLIENTS_HTTP_HOST} - MG_CLIENTS_HTTP_PORT: ${MG_CLIENTS_HTTP_PORT} - MG_CLIENTS_GRPC_HOST: ${MG_CLIENTS_GRPC_HOST} - MG_CLIENTS_GRPC_PORT: ${MG_CLIENTS_GRPC_PORT} - ## Compose supports parameter expansion in environment, - ## Eg: ${VAR:+replacement} or ${VAR+replacement} -> replacement if VAR is set and non-empty, otherwise empty - ## Eg :${VAR:-default} or ${VAR-default} -> value of VAR if set and non-empty, otherwise default - MG_CLIENTS_GRPC_SERVER_CERT: ${MG_CLIENTS_GRPC_SERVER_CERT:+/clients-grpc-server.crt} - MG_CLIENTS_GRPC_SERVER_KEY: ${MG_CLIENTS_GRPC_SERVER_KEY:+/clients-grpc-server.key} - MG_CLIENTS_GRPC_SERVER_CA_CERTS: ${MG_CLIENTS_GRPC_SERVER_CA_CERTS:+/clients-grpc-server-ca.crt} - MG_CLIENTS_GRPC_CLIENT_CA_CERTS: ${MG_CLIENTS_GRPC_CLIENT_CA_CERTS:+/clients-grpc-client-ca.crt} - MG_ES_URL: ${MG_ES_URL} - MG_CLIENTS_CACHE_URL: ${MG_CLIENTS_CACHE_URL} - MG_CLIENTS_DB_HOST: ${MG_CLIENTS_DB_HOST} - MG_CLIENTS_DB_PORT: ${MG_CLIENTS_DB_PORT} - MG_CLIENTS_DB_USER: ${MG_CLIENTS_DB_USER} - MG_CLIENTS_DB_PASS: ${MG_CLIENTS_DB_PASS} - MG_CLIENTS_DB_NAME: ${MG_CLIENTS_DB_NAME} - MG_CLIENTS_DB_SSL_MODE: ${MG_CLIENTS_DB_SSL_MODE} - MG_CLIENTS_DB_SSL_CERT: ${MG_CLIENTS_DB_SSL_CERT} - MG_CLIENTS_DB_SSL_KEY: ${MG_CLIENTS_DB_SSL_KEY} - MG_CLIENTS_DB_SSL_ROOT_CERT: ${MG_CLIENTS_DB_SSL_ROOT_CERT} - MG_AUTH_GRPC_URL: ${MG_AUTH_GRPC_URL} - MG_AUTH_GRPC_TIMEOUT: ${MG_AUTH_GRPC_TIMEOUT} - MG_AUTH_GRPC_CLIENT_CERT: ${MG_AUTH_GRPC_CLIENT_CERT:+/auth-grpc-client.crt} - MG_AUTH_GRPC_CLIENT_KEY: ${MG_AUTH_GRPC_CLIENT_KEY:+/auth-grpc-client.key} - MG_AUTH_GRPC_SERVER_CA_CERTS: ${MG_AUTH_GRPC_SERVER_CA_CERTS:+/auth-grpc-server-ca.crt} - MG_AUTH_KEYS_ALGORITHM: ${MG_AUTH_KEYS_ALGORITHM} - MG_CHANNELS_URL: ${MG_CHANNELS_URL} - MG_CHANNELS_GRPC_URL: ${MG_CHANNELS_GRPC_URL} - MG_CHANNELS_GRPC_TIMEOUT: ${MG_CHANNELS_GRPC_TIMEOUT} - MG_CHANNELS_GRPC_CLIENT_CERT: ${MG_CHANNELS_GRPC_CLIENT_CERT:+/channels-grpc-client.crt} - MG_CHANNELS_GRPC_CLIENT_KEY: ${MG_CHANNELS_GRPC_CLIENT_KEY:+/channels-grpc-client.key} - MG_CHANNELS_GRPC_SERVER_CA_CERTS: ${MG_CHANNELS_GRPC_SERVER_CA_CERTS:+/channels-grpc-server-ca.crt} - MG_GROUPS_URL: ${MG_GROUPS_URL} - MG_GROUPS_GRPC_URL: ${MG_GROUPS_GRPC_URL} - MG_GROUPS_GRPC_TIMEOUT: ${MG_GROUPS_GRPC_TIMEOUT} - MG_GROUPS_GRPC_CLIENT_CERT: ${MG_GROUPS_GRPC_CLIENT_CERT:+/groups-grpc-client.crt} - MG_GROUPS_GRPC_CLIENT_KEY: ${MG_GROUPS_GRPC_CLIENT_KEY:+/groups-grpc-client.key} - MG_GROUPS_GRPC_SERVER_CA_CERTS: ${MG_GROUPS_GRPC_SERVER_CA_CERTS:+/groups-grpc-server-ca.crt} - MG_DOMAINS_GRPC_URL: ${MG_DOMAINS_GRPC_URL} - MG_DOMAINS_GRPC_TIMEOUT: ${MG_DOMAINS_GRPC_TIMEOUT} - MG_DOMAINS_GRPC_CLIENT_CERT: ${MG_DOMAINS_GRPC_CLIENT_CERT:+/domains-grpc-client.crt} - MG_DOMAINS_GRPC_CLIENT_KEY: ${MG_DOMAINS_GRPC_CLIENT_KEY:+/domains-grpc-client.key} - MG_DOMAINS_GRPC_SERVER_CA_CERTS: ${MG_DOMAINS_GRPC_SERVER_CA_CERTS:+/domains-grpc-server-ca.crt} - MG_JAEGER_URL: ${MG_JAEGER_URL} - MG_JAEGER_TRACE_RATIO: ${MG_JAEGER_TRACE_RATIO} - MG_SEND_TELEMETRY: ${MG_SEND_TELEMETRY} - MG_SPICEDB_PRE_SHARED_KEY: ${MG_SPICEDB_PRE_SHARED_KEY} - MG_SPICEDB_HOST: ${MG_SPICEDB_HOST} - MG_SPICEDB_PORT: ${MG_SPICEDB_PORT} - MG_SPICEDB_SCHEMA_FILE: ${MG_SPICEDB_SCHEMA_FILE} - MG_CLIENTS_CALLOUT_URLS: ${MG_CLIENTS_CALLOUT_URLS} - MG_CLIENTS_CALLOUT_METHOD: ${MG_CLIENTS_CALLOUT_METHOD} - MG_CLIENTS_CALLOUT_TLS_VERIFICATION: ${MG_CLIENTS_CALLOUT_TLS_VERIFICATION} - MG_CLIENTS_CALLOUT_TIMEOUT: ${MG_CLIENTS_CALLOUT_TIMEOUT} - MG_CLIENTS_CALLOUT_CA_CERT: ${MG_CLIENTS_CALLOUT_CA_CERT} - MG_CLIENTS_CALLOUT_CERT: ${MG_CLIENTS_CALLOUT_CERT} - MG_CLIENTS_CALLOUT_KEY: ${MG_CLIENTS_CALLOUT_KEY} - MG_CLIENTS_CALLOUT_OPERATIONS: ${MG_CLIENTS_CALLOUT_OPERATIONS} - MG_ALLOW_UNVERIFIED_USER: ${MG_ALLOW_UNVERIFIED_USER} - ports: - - ${MG_CLIENTS_HTTP_PORT}:${MG_CLIENTS_HTTP_PORT} - - ${MG_CLIENTS_GRPC_PORT}:${MG_CLIENTS_GRPC_PORT} - networks: - - magistrala-base-net - volumes: - - ./permission.yaml:/permission.yaml - - ./spicedb/schema.zed:${MG_SPICEDB_SCHEMA_FILE} - # Clients gRPC server certificates - - type: bind - source: ${MG_CLIENTS_GRPC_SERVER_CERT:-./ssl/placeholder} - target: /clients-grpc-server.crt - bind: - create_host_path: true - - type: bind - source: ${MG_CLIENTS_GRPC_SERVER_KEY:-./ssl/placeholder} - target: /clients-grpc-server.key - bind: - create_host_path: true - - type: bind - source: ${MG_CLIENTS_GRPC_SERVER_CA_CERTS:-./ssl/placeholder} - target: /clients-grpc-server-ca.crt - bind: - create_host_path: true - - type: bind - source: ${MG_CLIENTS_GRPC_CLIENT_CA_CERTS:-./ssl/placeholder} - target: /clients-grpc-client-ca.crt - bind: - create_host_path: true - # Auth gRPC client certificates - - type: bind - source: ${MG_AUTH_GRPC_CLIENT_CERT:-./ssl/placeholder} - target: /auth-grpc-client.crt - bind: - create_host_path: true - - type: bind - source: ${MG_AUTH_GRPC_CLIENT_KEY:-./ssl/placeholder} - target: /auth-grpc-client.key - bind: - create_host_path: true - - type: bind - source: ${MG_AUTH_GRPC_SERVER_CA_CERTS:-./ssl/placeholder} - target: /auth-grpc-server-ca.crt - bind: - create_host_path: true - # Channel gRPC client certificates - - type: bind - source: ${MG_CHANNELS_GRPC_CLIENT_CERT:-./ssl/placeholder} - target: /channels-grpc-client.crt - bind: - create_host_path: true - - type: bind - source: ${MG_CHANNELS_GRPC_CLIENT_KEY:-./ssl/placeholder} - target: /channels-grpc-client.key - bind: - create_host_path: true - - type: bind - source: ${MG_CHANNELS_GRPC_SERVER_CA_CERTS:-./ssl/placeholder} - target: /channels-grpc-server-ca.crt - bind: - create_host_path: true - # Group gRPC client certificates - - type: bind - source: ${MG_GROUPS_GRPC_CLIENT_CERT:-./ssl/placeholder} - target: /groups-grpc-client.crt - bind: - create_host_path: true - - type: bind - source: ${MG_GROUPS_GRPC_CLIENT_KEY:-./ssl/placeholder} - target: /groups-grpc-client.key - bind: - create_host_path: true - - type: bind - source: ${MG_GROUPS_GRPC_SERVER_CA_CERTS:-./ssl/placeholder} - target: /groups-grpc-server-ca.crt - bind: - create_host_path: true - # Domain gRPC client certificates - - type: bind - source: ${MG_DOMAINS_GRPC_CLIENT_CERT:-./ssl/placeholder} - target: /domains-grpc-client.crt - bind: - create_host_path: true - - type: bind - source: ${MG_DOMAINS_GRPC_CLIENT_KEY:-./ssl/placeholder} - target: /domains-grpc-client.key - bind: - create_host_path: true - - type: bind - source: ${MG_DOMAINS_GRPC_SERVER_CA_CERTS:-./ssl/placeholder} - target: /domains-grpc-server-ca.crt - bind: - create_host_path: true - - channels-db: - image: docker.io/postgres:18.0-alpine3.22 - container_name: magistrala-channels-db - restart: on-failure - command: postgres -c "max_connections=${MG_POSTGRES_MAX_CONNECTIONS}" - environment: - POSTGRES_USER: ${MG_CHANNELS_DB_USER} - POSTGRES_PASSWORD: ${MG_CHANNELS_DB_PASS} - POSTGRES_DB: ${MG_CHANNELS_DB_NAME} - MG_POSTGRES_MAX_CONNECTIONS: ${MG_POSTGRES_MAX_CONNECTIONS} - networks: - - magistrala-base-net - ports: - - 6005:5432 - volumes: - - magistrala-channels-db-volume:/var/lib/postgresql/data - - channels-redis: - image: docker.io/redis:8.2.2-alpine3.22 - container_name: magistrala-channels-redis - restart: on-failure - networks: - - magistrala-base-net - volumes: - - magistrala-channels-redis-volume:/data - - channels: - image: ghcr.io/absmach/magistrala/channels:${MG_RELEASE_TAG} - container_name: magistrala-channels - depends_on: - - channels-db - - channels-redis - - users - - auth - - nginx - restart: on-failure - environment: - MG_CHANNELS_LOG_LEVEL: ${MG_CHANNELS_LOG_LEVEL} - MG_CHANNELS_INSTANCE_ID: ${MG_CHANNELS_INSTANCE_ID} - MG_CHANNELS_HTTP_HOST: ${MG_CHANNELS_HTTP_HOST} - MG_CHANNELS_HTTP_PORT: ${MG_CHANNELS_HTTP_PORT} - MG_CHANNELS_GRPC_HOST: ${MG_CHANNELS_GRPC_HOST} - MG_CHANNELS_GRPC_PORT: ${MG_CHANNELS_GRPC_PORT} - ## Compose supports parameter expansion in environment, - ## Eg: ${VAR:+replacement} or ${VAR+replacement} -> replacement if VAR is set and non-empty, otherwise empty - ## Eg :${VAR:-default} or ${VAR-default} -> value of VAR if set and non-empty, otherwise default - MG_CHANNELS_GRPC_SERVER_CERT: ${MG_CHANNELS_GRPC_SERVER_CERT:+/channels-grpc-server.crt} - MG_CHANNELS_GRPC_SERVER_KEY: ${MG_CHANNELS_GRPC_SERVER_KEY:+/channels-grpc-server.key} - MG_CHANNELS_GRPC_SERVER_CA_CERTS: ${MG_CHANNELS_GRPC_SERVER_CA_CERTS:+/channels-grpc-server-ca.crt} - MG_CHANNELS_GRPC_CLIENT_CA_CERTS: ${MG_CHANNELS_GRPC_CLIENT_CA_CERTS:+/channels-grpc-client-ca.crt} - MG_CHANNELS_DB_HOST: ${MG_CHANNELS_DB_HOST} - MG_CHANNELS_DB_PORT: ${MG_CHANNELS_DB_PORT} - MG_CHANNELS_DB_USER: ${MG_CHANNELS_DB_USER} - MG_CHANNELS_DB_PASS: ${MG_CHANNELS_DB_PASS} - MG_CHANNELS_DB_NAME: ${MG_CHANNELS_DB_NAME} - MG_CHANNELS_DB_SSL_MODE: ${MG_CHANNELS_DB_SSL_MODE} - MG_CHANNELS_DB_SSL_CERT: ${MG_CHANNELS_DB_SSL_CERT} - MG_CHANNELS_DB_SSL_KEY: ${MG_CHANNELS_DB_SSL_KEY} - MG_CHANNELS_DB_SSL_ROOT_CERT: ${MG_CHANNELS_DB_SSL_ROOT_CERT} - MG_CHANNELS_CACHE_URL: ${MG_CHANNELS_CACHE_URL} - MG_CHANNELS_CACHE_KEY_DURATION: ${MG_CHANNELS_CACHE_KEY_DURATION} - MG_AUTH_GRPC_URL: ${MG_AUTH_GRPC_URL} - MG_AUTH_GRPC_TIMEOUT: ${MG_AUTH_GRPC_TIMEOUT} - MG_AUTH_GRPC_CLIENT_CERT: ${MG_AUTH_GRPC_CLIENT_CERT:+/auth-grpc-client.crt} - MG_AUTH_GRPC_CLIENT_KEY: ${MG_AUTH_GRPC_CLIENT_KEY:+/auth-grpc-client.key} - MG_AUTH_GRPC_SERVER_CA_CERTS: ${MG_AUTH_GRPC_SERVER_CA_CERTS:+/auth-grpc-server-ca.crt} - MG_AUTH_KEYS_ALGORITHM: ${MG_AUTH_KEYS_ALGORITHM} - MG_CLIENTS_GRPC_URL: ${MG_CLIENTS_GRPC_URL} - MG_CLIENTS_GRPC_TIMEOUT: ${MG_CLIENTS_GRPC_TIMEOUT} - MG_CLIENTS_GRPC_CLIENT_CERT: ${MG_CLIENTS_GRPC_CLIENT_CERT:+/clients-grpc-client.crt} - MG_CLIENTS_GRPC_CLIENT_KEY: ${MG_CLIENTS_GRPC_CLIENT_KEY:+/clients-grpc-client.key} - MG_CLIENTS_GRPC_SERVER_CA_CERTS: ${MG_CLIENTS_GRPC_SERVER_CA_CERTS:+/clients-grpc-server-ca.crt} - MG_GROUPS_GRPC_URL: ${MG_GROUPS_GRPC_URL} - MG_GROUPS_GRPC_TIMEOUT: ${MG_GROUPS_GRPC_TIMEOUT} - MG_GROUPS_GRPC_CLIENT_CERT: ${MG_GROUPS_GRPC_CLIENT_CERT:+/groups-grpc-client.crt} - MG_GROUPS_GRPC_CLIENT_KEY: ${MG_GROUPS_GRPC_CLIENT_KEY:+/groups-grpc-client.key} - MG_GROUPS_GRPC_SERVER_CA_CERTS: ${MG_GROUPS_GRPC_SERVER_CA_CERTS:+/groups-grpc-server-ca.crt} - MG_DOMAINS_GRPC_URL: ${MG_DOMAINS_GRPC_URL} - MG_DOMAINS_GRPC_TIMEOUT: ${MG_DOMAINS_GRPC_TIMEOUT} - MG_DOMAINS_GRPC_CLIENT_CERT: ${MG_DOMAINS_GRPC_CLIENT_CERT:+/domains-grpc-client.crt} - MG_DOMAINS_GRPC_CLIENT_KEY: ${MG_DOMAINS_GRPC_CLIENT_KEY:+/domains-grpc-client.key} - MG_DOMAINS_GRPC_SERVER_CA_CERTS: ${MG_DOMAINS_GRPC_SERVER_CA_CERTS:+/domains-grpc-server-ca.crt} - MG_ES_URL: ${MG_ES_URL} - MG_JAEGER_URL: ${MG_JAEGER_URL} - MG_JAEGER_TRACE_RATIO: ${MG_JAEGER_TRACE_RATIO} - MG_SEND_TELEMETRY: ${MG_SEND_TELEMETRY} - MG_SPICEDB_PRE_SHARED_KEY: ${MG_SPICEDB_PRE_SHARED_KEY} - MG_SPICEDB_HOST: ${MG_SPICEDB_HOST} - MG_SPICEDB_PORT: ${MG_SPICEDB_PORT} - MG_SPICEDB_SCHEMA_FILE: ${MG_SPICEDB_SCHEMA_FILE} - MG_CHANNELS_CALLOUT_URLS: ${MG_CHANNELS_CALLOUT_URLS} - MG_CHANNELS_CALLOUT_METHOD: ${MG_CHANNELS_CALLOUT_METHOD} - MG_CHANNELS_CALLOUT_TLS_VERIFICATION: ${MG_CHANNELS_CALLOUT_TLS_VERIFICATION} - MG_CHANNELS_CALLOUT_TIMEOUT: ${MG_CHANNELS_CALLOUT_TIMEOUT} - MG_CHANNELS_CALLOUT_CA_CERT: ${MG_CHANNELS_CALLOUT_CA_CERT} - MG_CHANNELS_CALLOUT_CERT: ${MG_CHANNELS_CALLOUT_CERT} - MG_CHANNELS_CALLOUT_KEY: ${MG_CHANNELS_CALLOUT_KEY} - MG_CHANNELS_CALLOUT_OPERATIONS: ${MG_CHANNELS_CALLOUT_OPERATIONS} - MG_ALLOW_UNVERIFIED_USER: ${MG_ALLOW_UNVERIFIED_USER} - ports: - - ${MG_CHANNELS_HTTP_PORT}:${MG_CHANNELS_HTTP_PORT} - - ${MG_CHANNELS_GRPC_PORT}:${MG_CHANNELS_GRPC_PORT} - networks: - - magistrala-base-net - volumes: - - ./permission.yaml:/permission.yaml - - ./spicedb/schema.zed:${MG_SPICEDB_SCHEMA_FILE} - # Channels gRPC server certificates - - type: bind - source: ${MG_CHANNELS_GRPC_SERVER_CERT:-./ssl/placeholder} - target: /channels-grpc-server.crt - bind: - create_host_path: true - - type: bind - source: ${MG_CHANNELS_GRPC_SERVER_KEY:-./ssl/placeholder} - target: /channels-grpc-server.key - bind: - create_host_path: true - - type: bind - source: ${MG_CHANNELS_GRPC_SERVER_CA_CERTS:-./ssl/placeholder} - target: /channels-grpc-server-ca.crt - bind: - create_host_path: true - - type: bind - source: ${MG_CHANNELS_GRPC_CLIENT_CA_CERTS:-./ssl/placeholder} - target: /channels-grpc-client-ca.crt - bind: - create_host_path: true - # Auth gRPC client certificates - - type: bind - source: ${MG_AUTH_GRPC_CLIENT_CERT:-./ssl/placeholder} - target: /auth-grpc-client.crt - bind: - create_host_path: true - - type: bind - source: ${MG_AUTH_GRPC_CLIENT_KEY:-./ssl/placeholder} - target: /auth-grpc-client.key - bind: - create_host_path: true - - type: bind - source: ${MG_AUTH_GRPC_SERVER_CA_CERTS:-./ssl/placeholder} - target: /auth-grpc-server-ca.crt - bind: - create_host_path: true - # Clients gRPC client certificates - - type: bind - source: ${MG_CLIENTS_GRPC_CLIENT_CERT:-./ssl/placeholder} - target: /clients-grpc-client.crt - bind: - create_host_path: true - - type: bind - source: ${MG_CLIENTS_GRPC_CLIENT_KEY:-./ssl/placeholder} - target: /clients-grpc-client.key - bind: - create_host_path: true - - type: bind - source: ${MG_CLIENTS_GRPC_SERVER_CA_CERTS:-./ssl/placeholder} - target: /clients-grpc-server-ca.crt - bind: - create_host_path: true - # Groups gRPC client certificates - - type: bind - source: ${MG_GROUPS_GRPC_CLIENT_CERT:-./ssl/placeholder} - target: /groups-grpc-client.crt - bind: - create_host_path: true - - type: bind - source: ${MG_GROUPS_GRPC_CLIENT_KEY:-./ssl/placeholder} - target: /groups-grpc-client.key - bind: - create_host_path: true - - type: bind - source: ${MG_GROUPS_GRPC_SERVER_CA_CERTS:-./ssl/placeholder} - target: /groups-grpc-server-ca.crt - bind: - create_host_path: true - # Domains gRPC client certificates - - type: bind - source: ${MG_DOMAINS_GRPC_CLIENT_CERT:-./ssl/placeholder} - target: /domains-grpc-client.crt - bind: - create_host_path: true - - type: bind - source: ${MG_DOMAINS_GRPC_CLIENT_KEY:-./ssl/placeholder} - target: /domains-grpc-client.key - bind: - create_host_path: true - - type: bind - source: ${MG_DOMAINS_GRPC_SERVER_CA_CERTS:-./ssl/placeholder} - target: /domains-grpc-server-ca.crt - bind: - create_host_path: true - - users-db: - image: docker.io/postgres:18.0-alpine3.22 - container_name: magistrala-users-db - restart: on-failure - command: postgres -c "max_connections=${MG_POSTGRES_MAX_CONNECTIONS}" - environment: - POSTGRES_USER: ${MG_USERS_DB_USER} - POSTGRES_PASSWORD: ${MG_USERS_DB_PASS} - POSTGRES_DB: ${MG_USERS_DB_NAME} - MG_POSTGRES_MAX_CONNECTIONS: ${MG_POSTGRES_MAX_CONNECTIONS} - ports: - - 6002:5432 - networks: - - magistrala-base-net - volumes: - - magistrala-users-db-volume:/var/lib/postgresql/data - - users: - image: ghcr.io/absmach/magistrala/users:${MG_RELEASE_TAG} - container_name: magistrala-users - depends_on: - - users-db - - auth - - nginx - restart: on-failure - environment: - MG_USERS_LOG_LEVEL: ${MG_USERS_LOG_LEVEL} - MG_USERS_SECRET_KEY: ${MG_USERS_SECRET_KEY} - MG_USERS_ADMIN_EMAIL: ${MG_USERS_ADMIN_EMAIL} - MG_USERS_ADMIN_PASSWORD: ${MG_USERS_ADMIN_PASSWORD} - MG_USERS_ADMIN_USERNAME: ${MG_USERS_ADMIN_USERNAME} - MG_USERS_ADMIN_FIRST_NAME: ${MG_USERS_ADMIN_FIRST_NAME} - MG_USERS_ADMIN_LAST_NAME: ${MG_USERS_ADMIN_LAST_NAME} - MG_USERS_PASS_REGEX: ${MG_USERS_PASS_REGEX} - MG_USERS_HTTP_HOST: ${MG_USERS_HTTP_HOST} - MG_USERS_HTTP_PORT: ${MG_USERS_HTTP_PORT} - MG_USERS_HTTP_SERVER_CERT: ${MG_USERS_HTTP_SERVER_CERT} - MG_USERS_HTTP_SERVER_KEY: ${MG_USERS_HTTP_SERVER_KEY} - MG_USERS_GRPC_HOST: ${MG_USERS_GRPC_HOST} - MG_USERS_GRPC_PORT: ${MG_USERS_GRPC_PORT} - ## Compose supports parameter expansion in environment, - ## Eg: ${VAR:+replacement} or ${VAR+replacement} -> replacement if VAR is set and non-empty, otherwise empty - ## Eg :${VAR:-default} or ${VAR-default} -> value of VAR if set and non-empty, otherwise default - MG_USERS_GRPC_SERVER_CERT: ${MG_USERS_GRPC_SERVER_CERT:+/users-grpc-server.crt} - MG_USERS_GRPC_SERVER_KEY: ${MG_USERS_GRPC_SERVER_KEY:+/users-grpc-server.key} - MG_USERS_GRPC_SERVER_CA_CERTS: ${MG_USERS_GRPC_SERVER_CA_CERTS:+/users-grpc-server-ca.crt} - MG_USERS_GRPC_CLIENT_CA_CERTS: ${MG_USERS_GRPC_CLIENT_CA_CERTS:+/users-grpc-client-ca.crt} - MG_USERS_DB_HOST: ${MG_USERS_DB_HOST} - MG_USERS_DB_PORT: ${MG_USERS_DB_PORT} - MG_USERS_DB_USER: ${MG_USERS_DB_USER} - MG_USERS_DB_PASS: ${MG_USERS_DB_PASS} - MG_USERS_DB_NAME: ${MG_USERS_DB_NAME} - MG_USERS_DB_SSL_MODE: ${MG_USERS_DB_SSL_MODE} - MG_USERS_DB_SSL_CERT: ${MG_USERS_DB_SSL_CERT} - MG_USERS_DB_SSL_KEY: ${MG_USERS_DB_SSL_KEY} - MG_USERS_DB_SSL_ROOT_CERT: ${MG_USERS_DB_SSL_ROOT_CERT} - MG_USERS_ALLOW_SELF_REGISTER: ${MG_USERS_ALLOW_SELF_REGISTER} - MG_EMAIL_HOST: ${MG_EMAIL_HOST} - MG_EMAIL_PORT: ${MG_EMAIL_PORT} - MG_EMAIL_USERNAME: ${MG_EMAIL_USERNAME} - MG_EMAIL_PASSWORD: ${MG_EMAIL_PASSWORD} - MG_EMAIL_FROM_ADDRESS: ${MG_EMAIL_FROM_ADDRESS} - MG_EMAIL_FROM_NAME: ${MG_EMAIL_FROM_NAME} - MG_ES_URL: ${MG_ES_URL} - MG_JAEGER_URL: ${MG_JAEGER_URL} - MG_JAEGER_TRACE_RATIO: ${MG_JAEGER_TRACE_RATIO} - MG_SEND_TELEMETRY: ${MG_SEND_TELEMETRY} - MG_AUTH_GRPC_URL: ${MG_AUTH_GRPC_URL} - MG_AUTH_GRPC_TIMEOUT: ${MG_AUTH_GRPC_TIMEOUT} - MG_AUTH_GRPC_CLIENT_CERT: ${MG_AUTH_GRPC_CLIENT_CERT:+/auth-grpc-client.crt} - MG_AUTH_GRPC_CLIENT_KEY: ${MG_AUTH_GRPC_CLIENT_KEY:+/auth-grpc-client.key} - MG_AUTH_GRPC_SERVER_CA_CERTS: ${MG_AUTH_GRPC_SERVER_CA_CERTS:+/auth-grpc-server-ca.crt} - MG_AUTH_KEYS_ALGORITHM: ${MG_AUTH_KEYS_ALGORITHM} - MG_DOMAINS_GRPC_URL: ${MG_DOMAINS_GRPC_URL} - MG_DOMAINS_GRPC_TIMEOUT: ${MG_DOMAINS_GRPC_TIMEOUT} - MG_DOMAINS_GRPC_CLIENT_CERT: ${MG_DOMAINS_GRPC_CLIENT_CERT:+/domains-grpc-client.crt} - MG_DOMAINS_GRPC_CLIENT_KEY: ${MG_DOMAINS_GRPC_CLIENT_KEY:+/domains-grpc-client.key} - MG_DOMAINS_GRPC_SERVER_CA_CERTS: ${MG_DOMAINS_GRPC_SERVER_CA_CERTS:+/domains-grpc-server-ca.crt} - MG_GOOGLE_CLIENT_ID: ${MG_GOOGLE_CLIENT_ID} - MG_GOOGLE_CLIENT_SECRET: ${MG_GOOGLE_CLIENT_SECRET} - MG_GOOGLE_REDIRECT_URL: ${MG_GOOGLE_REDIRECT_URL} - MG_GOOGLE_STATE: ${MG_GOOGLE_STATE} - MG_OAUTH_UI_REDIRECT_URL: ${MG_OAUTH_UI_REDIRECT_URL} - MG_OAUTH_UI_ERROR_URL: ${MG_OAUTH_UI_ERROR_URL} - MG_USERS_DELETE_INTERVAL: ${MG_USERS_DELETE_INTERVAL} - MG_USERS_DELETE_AFTER: ${MG_USERS_DELETE_AFTER} - MG_SPICEDB_PRE_SHARED_KEY: ${MG_SPICEDB_PRE_SHARED_KEY} - MG_SPICEDB_HOST: ${MG_SPICEDB_HOST} - MG_SPICEDB_PORT: ${MG_SPICEDB_PORT} - MG_PASSWORD_RESET_URL_PREFIX: ${MG_PASSWORD_RESET_URL_PREFIX} - MG_PASSWORD_RESET_EMAIL_TEMPLATE: ${MG_PASSWORD_RESET_EMAIL_TEMPLATE} - MG_VERIFICATION_URL_PREFIX: ${MG_VERIFICATION_URL_PREFIX} - MG_VERIFICATION_EMAIL_TEMPLATE: ${MG_VERIFICATION_EMAIL_TEMPLATE} - MG_ALLOW_UNVERIFIED_USER: ${MG_ALLOW_UNVERIFIED_USER} - ports: - - ${MG_USERS_HTTP_PORT}:${MG_USERS_HTTP_PORT} - - ${MG_USERS_GRPC_PORT}:${MG_USERS_GRPC_PORT} - networks: - - magistrala-base-net - volumes: - - ./templates/${MG_PASSWORD_RESET_EMAIL_TEMPLATE}:/${MG_PASSWORD_RESET_EMAIL_TEMPLATE} - - ./templates/${MG_VERIFICATION_EMAIL_TEMPLATE}:/${MG_VERIFICATION_EMAIL_TEMPLATE} - # Users gRPC server certificates - - type: bind - source: ${MG_USERS_GRPC_SERVER_CERT:-./ssl/placeholder} - target: /users-grpc-server.crt - bind: - create_host_path: true - - type: bind - source: ${MG_USERS_GRPC_SERVER_KEY:-./ssl/placeholder} - target: /users-grpc-server.key - bind: - create_host_path: true - - type: bind - source: ${MG_USERS_GRPC_SERVER_CA_CERTS:-./ssl/placeholder} - target: /users-grpc-server-ca.crt - bind: - create_host_path: true - - type: bind - source: ${MG_USERS_GRPC_CLIENT_CA_CERTS:-./ssl/placeholder} - target: /users-grpc-client-ca.crt - bind: - create_host_path: true - # Auth gRPC client certificates - - type: bind - source: ${MG_AUTH_GRPC_CLIENT_CERT:-./ssl/placeholder} - target: /auth-grpc-client.crt - bind: - create_host_path: true - - type: bind - source: ${MG_AUTH_GRPC_CLIENT_KEY:-./ssl/placeholder} - target: /auth-grpc-client.key - bind: - create_host_path: true - - type: bind - source: ${MG_AUTH_GRPC_SERVER_CA_CERTS:-./ssl/placeholder} - target: /auth-grpc-server-ca.crt - bind: - create_host_path: true - # Domains gRPC client certificates - - type: bind - source: ${MG_DOMAINS_GRPC_CLIENT_CERT:-./ssl/placeholder} - target: /domains-grpc-client.crt - bind: - create_host_path: true - - type: bind - source: ${MG_DOMAINS_GRPC_CLIENT_KEY:-./ssl/placeholder} - target: /domains-grpc-client.key - bind: - create_host_path: true - - type: bind - source: ${MG_DOMAINS_GRPC_SERVER_CA_CERTS:-./ssl/placeholder} - target: /domains-grpc-server-ca.crt - bind: - create_host_path: true - notifications: image: ghcr.io/absmach/magistrala/notifications:${MG_RELEASE_TAG} container_name: magistrala-notifications depends_on: - - nginx + atom-bootstrap: + condition: service_completed_successfully + nginx: + condition: service_started restart: on-failure environment: MG_NOTIFICATIONS_LOG_LEVEL: ${MG_NOTIFICATIONS_LOG_LEVEL} @@ -1169,215 +309,18 @@ services: MG_EMAIL_INVITATION_TEMPLATE: ${MG_EMAIL_INVITATION_TEMPLATE} MG_EMAIL_ACCEPTANCE_TEMPLATE: ${MG_EMAIL_ACCEPTANCE_TEMPLATE} MG_EMAIL_REJECTION_TEMPLATE: ${MG_EMAIL_REJECTION_TEMPLATE} - MG_USERS_GRPC_URL: ${MG_USERS_GRPC_URL} - MG_USERS_GRPC_TIMEOUT: ${MG_USERS_GRPC_TIMEOUT} - MG_USERS_GRPC_CLIENT_CERT: ${MG_USERS_GRPC_CLIENT_CERT:+/users-grpc-client.crt} - MG_USERS_GRPC_CLIENT_KEY: ${MG_USERS_GRPC_CLIENT_KEY:+/users-grpc-client.key} - MG_USERS_GRPC_SERVER_CA_CERTS: ${MG_USERS_GRPC_SERVER_CA_CERTS:+/users-grpc-server-ca.crt} + ATOM_URL: ${ATOM_URL} + ATOM_SERVICE_TOKEN: ${MG_ATOM_TOKEN_NOTIFICATIONS} + ATOM_JWKS_URL: ${ATOM_JWKS_URL} + ATOM_JWT_ISSUER: ${ATOM_JWT_ISSUER} + ATOM_JWT_AUDIENCE: ${ATOM_JWT_AUDIENCE} + ATOM_TIMEOUT: ${ATOM_TIMEOUT} networks: - magistrala-base-net volumes: - ./templates/${MG_EMAIL_INVITATION_TEMPLATE}:/${MG_EMAIL_INVITATION_TEMPLATE} - ./templates/${MG_EMAIL_ACCEPTANCE_TEMPLATE}:/${MG_EMAIL_ACCEPTANCE_TEMPLATE} - ./templates/${MG_EMAIL_REJECTION_TEMPLATE}:/${MG_EMAIL_REJECTION_TEMPLATE} - # Users gRPC client certificates - - type: bind - source: ${MG_USERS_GRPC_CLIENT_CERT:-./ssl/placeholder} - target: /users-grpc-client.crt - bind: - create_host_path: true - - type: bind - source: ${MG_USERS_GRPC_CLIENT_KEY:-./ssl/placeholder} - target: /users-grpc-client.key - bind: - create_host_path: true - - type: bind - source: ${MG_USERS_GRPC_SERVER_CA_CERTS:-./ssl/placeholder} - target: /users-grpc-server-ca.crt - bind: - create_host_path: true - - groups-db: - image: docker.io/postgres:18.0-alpine3.22 - container_name: magistrala-groups-db - restart: on-failure - command: postgres -c "max_connections=${MG_POSTGRES_MAX_CONNECTIONS}" - environment: - POSTGRES_USER: ${MG_GROUPS_DB_USER} - POSTGRES_PASSWORD: ${MG_GROUPS_DB_PASS} - POSTGRES_DB: ${MG_GROUPS_DB_NAME} - MG_POSTGRES_MAX_CONNECTIONS: ${MG_POSTGRES_MAX_CONNECTIONS} - ports: - - 6004:5432 - networks: - - magistrala-base-net - volumes: - - magistrala-groups-db-volume:/var/lib/postgresql/data - - groups: - image: ghcr.io/absmach/magistrala/groups:${MG_RELEASE_TAG} - container_name: magistrala-groups - depends_on: - - groups-db - - auth - - nginx - restart: on-failure - environment: - MG_GROUPS_LOG_LEVEL: ${MG_GROUPS_LOG_LEVEL} - MG_GROUPS_HTTP_HOST: ${MG_GROUPS_HTTP_HOST} - MG_GROUPS_HTTP_PORT: ${MG_GROUPS_HTTP_PORT} - MG_GROUPS_HTTP_SERVER_CERT: ${MG_GROUPS_HTTP_SERVER_CERT} - MG_GROUPS_HTTP_SERVER_KEY: ${MG_GROUPS_HTTP_SERVER_KEY} - MG_GROUPS_GRPC_HOST: ${MG_GROUPS_GRPC_HOST} - MG_GROUPS_GRPC_PORT: ${MG_GROUPS_GRPC_PORT} - ## Compose supports parameter expansion in environment, - ## Eg: ${VAR:+replacement} or ${VAR+replacement} -> replacement if VAR is set and non-empty, otherwise empty - ## Eg :${VAR:-default} or ${VAR-default} -> value of VAR if set and non-empty, otherwise default - MG_GROUPS_GRPC_SERVER_CERT: ${MG_GROUPS_GRPC_SERVER_CERT:+/groups-grpc-server.crt} - MG_GROUPS_GRPC_SERVER_KEY: ${MG_GROUPS_GRPC_SERVER_KEY:+/groups-grpc-server.key} - MG_GROUPS_GRPC_SERVER_CA_CERTS: ${MG_GROUPS_GRPC_SERVER_CA_CERTS:+/groups-grpc-server-ca.crt} - MG_GROUPS_GRPC_CLIENT_CA_CERTS: ${MG_GROUPS_GRPC_CLIENT_CA_CERTS:+/groups-grpc-client-ca.crt} - MG_GROUPS_DB_HOST: ${MG_GROUPS_DB_HOST} - MG_GROUPS_DB_PORT: ${MG_GROUPS_DB_PORT} - MG_GROUPS_DB_USER: ${MG_GROUPS_DB_USER} - MG_GROUPS_DB_PASS: ${MG_GROUPS_DB_PASS} - MG_GROUPS_DB_NAME: ${MG_GROUPS_DB_NAME} - MG_GROUPS_DB_SSL_MODE: ${MG_GROUPS_DB_SSL_MODE} - MG_GROUPS_DB_SSL_CERT: ${MG_GROUPS_DB_SSL_CERT} - MG_GROUPS_DB_SSL_KEY: ${MG_GROUPS_DB_SSL_KEY} - MG_GROUPS_DB_SSL_ROOT_CERT: ${MG_GROUPS_DB_SSL_ROOT_CERT} - MG_CHANNELS_URL: ${MG_CHANNELS_URL} - MG_CHANNELS_GRPC_URL: ${MG_CHANNELS_GRPC_URL} - MG_CHANNELS_GRPC_TIMEOUT: ${MG_CHANNELS_GRPC_TIMEOUT} - MG_CHANNELS_GRPC_CLIENT_CERT: ${MG_CHANNELS_GRPC_CLIENT_CERT:+/channels-grpc-client.crt} - MG_CHANNELS_GRPC_CLIENT_KEY: ${MG_CHANNELS_GRPC_CLIENT_KEY:+/channels-grpc-client.key} - MG_CHANNELS_GRPC_SERVER_CA_CERTS: ${MG_CHANNELS_GRPC_SERVER_CA_CERTS:+/channels-grpc-server-ca.crt} - MG_CLIENTS_GRPC_URL: ${MG_CLIENTS_GRPC_URL} - MG_CLIENTS_GRPC_TIMEOUT: ${MG_CLIENTS_GRPC_TIMEOUT} - MG_CLIENTS_GRPC_CLIENT_CERT: ${MG_CLIENTS_GRPC_CLIENT_CERT:+/clients-grpc-client.crt} - MG_CLIENTS_GRPC_CLIENT_KEY: ${MG_CLIENTS_GRPC_CLIENT_KEY:+/clients-grpc-client.key} - MG_CLIENTS_GRPC_SERVER_CA_CERTS: ${MG_CLIENTS_GRPC_SERVER_CA_CERTS:+/clients-grpc-server-ca.crt} - MG_DOMAINS_GRPC_URL: ${MG_DOMAINS_GRPC_URL} - MG_DOMAINS_GRPC_TIMEOUT: ${MG_DOMAINS_GRPC_TIMEOUT} - MG_DOMAINS_GRPC_CLIENT_CERT: ${MG_DOMAINS_GRPC_CLIENT_CERT:+/domains-grpc-client.crt} - MG_DOMAINS_GRPC_CLIENT_KEY: ${MG_DOMAINS_GRPC_CLIENT_KEY:+/domains-grpc-client.key} - MG_DOMAINS_GRPC_SERVER_CA_CERTS: ${MG_DOMAINS_GRPC_SERVER_CA_CERTS:+/domains-grpc-server-ca.crt} - MG_ES_URL: ${MG_ES_URL} - MG_JAEGER_URL: ${MG_JAEGER_URL} - MG_JAEGER_TRACE_RATIO: ${MG_JAEGER_TRACE_RATIO} - MG_SEND_TELEMETRY: ${MG_SEND_TELEMETRY} - MG_AUTH_GRPC_URL: ${MG_AUTH_GRPC_URL} - MG_AUTH_GRPC_TIMEOUT: ${MG_AUTH_GRPC_TIMEOUT} - MG_AUTH_GRPC_CLIENT_CERT: ${MG_AUTH_GRPC_CLIENT_CERT:+/auth-grpc-client.crt} - MG_AUTH_GRPC_CLIENT_KEY: ${MG_AUTH_GRPC_CLIENT_KEY:+/auth-grpc-client.key} - MG_AUTH_GRPC_SERVER_CA_CERTS: ${MG_AUTH_GRPC_SERVER_CA_CERTS:+/auth-grpc-server-ca.crt} - MG_AUTH_KEYS_ALGORITHM: ${MG_AUTH_KEYS_ALGORITHM} - MG_SPICEDB_PRE_SHARED_KEY: ${MG_SPICEDB_PRE_SHARED_KEY} - MG_SPICEDB_HOST: ${MG_SPICEDB_HOST} - MG_SPICEDB_PORT: ${MG_SPICEDB_PORT} - MG_SPICEDB_SCHEMA_FILE: ${MG_SPICEDB_SCHEMA_FILE} - MG_GROUPS_CALLOUT_URLS: ${MG_GROUPS_CALLOUT_URLS} - MG_GROUPS_CALLOUT_METHOD: ${MG_GROUPS_CALLOUT_METHOD} - MG_GROUPS_CALLOUT_TLS_VERIFICATION: ${MG_GROUPS_CALLOUT_TLS_VERIFICATION} - MG_GROUPS_CALLOUT_TIMEOUT: ${MG_GROUPS_CALLOUT_TIMEOUT} - MG_GROUPS_CALLOUT_CA_CERT: ${MG_GROUPS_CALLOUT_CA_CERT} - MG_GROUPS_CALLOUT_CERT: ${MG_GROUPS_CALLOUT_CERT} - MG_GROUPS_CALLOUT_KEY: ${MG_GROUPS_CALLOUT_KEY} - MG_GROUPS_CALLOUT_OPERATIONS: ${MG_GROUPS_CALLOUT_OPERATIONS} - MG_ALLOW_UNVERIFIED_USER: ${MG_ALLOW_UNVERIFIED_USER} - ports: - - ${MG_GROUPS_HTTP_PORT}:${MG_GROUPS_HTTP_PORT} - - ${MG_GROUPS_GRPC_PORT}:${MG_GROUPS_GRPC_PORT} - networks: - - magistrala-base-net - volumes: - - ./permission.yaml:/permission.yaml - - ./spicedb/schema.zed:${MG_SPICEDB_SCHEMA_FILE} - # Groups gRPC server certificates - - type: bind - source: ${MG_GROUPS_GRPC_SERVER_CERT:-./ssl/placeholder} - target: /groups-grpc-server.crt - bind: - create_host_path: true - - type: bind - source: ${MG_GROUPS_GRPC_SERVER_KEY:-./ssl/placeholder} - target: /groups-grpc-server.key - bind: - create_host_path: true - - type: bind - source: ${MG_GROUPS_GRPC_SERVER_CA_CERTS:-./ssl/placeholder} - target: /groups-grpc-server-ca.crt - bind: - create_host_path: true - - type: bind - source: ${MG_GROUPS_GRPC_CLIENT_CA_CERTS:-./ssl/placeholder} - target: /groups-grpc-client-ca.crt - bind: - create_host_path: true - # Auth gRPC client certificates - - type: bind - source: ${MG_AUTH_GRPC_CLIENT_CERT:-./ssl/placeholder} - target: /auth-grpc-client.crt - bind: - create_host_path: true - - type: bind - source: ${MG_AUTH_GRPC_CLIENT_KEY:-./ssl/placeholder} - target: /auth-grpc-client.key - bind: - create_host_path: true - - type: bind - source: ${MG_AUTH_GRPC_SERVER_CA_CERTS:-./ssl/placeholder} - target: /auth-grpc-server-ca.crt - bind: - create_host_path: true - # Clients gRPC client certificates - - type: bind - source: ${MG_CLIENTS_GRPC_CLIENT_CERT:-./ssl/placeholder} - target: /clients-grpc-client.crt - bind: - create_host_path: true - - type: bind - source: ${MG_CLIENTS_GRPC_CLIENT_KEY:-./ssl/placeholder} - target: /clients-grpc-client.key - bind: - create_host_path: true - - type: bind - source: ${MG_CLIENTS_GRPC_SERVER_CA_CERTS:-./ssl/placeholder} - target: /clients-grpc-server-ca.crt - bind: - create_host_path: true - # Channels gRPC client certificates - - type: bind - source: ${MG_CHANNELS_GRPC_CLIENT_CERT:-./ssl/placeholder} - target: /channels-grpc-client.crt - bind: - create_host_path: true - - type: bind - source: ${MG_CHANNELS_GRPC_CLIENT_KEY:-./ssl/placeholder} - target: /channels-grpc-client.key - bind: - create_host_path: true - - type: bind - source: ${MG_CHANNELS_GRPC_SERVER_CA_CERTS:-./ssl/placeholder} - target: /channels-grpc-server-ca.crt - bind: - create_host_path: true - # Domains gRPC client certificates - - type: bind - source: ${MG_DOMAINS_GRPC_CLIENT_CERT:-./ssl/placeholder} - target: /domains-grpc-client.crt - bind: - create_host_path: true - - type: bind - source: ${MG_DOMAINS_GRPC_CLIENT_KEY:-./ssl/placeholder} - target: /domains-grpc-client.key - bind: - create_host_path: true - - type: bind - source: ${MG_DOMAINS_GRPC_SERVER_CA_CERTS:-./ssl/placeholder} - target: /domains-grpc-server-ca.crt - bind: - create_host_path: true jaeger: image: docker.io/jaegertracing/all-in-one:1.74.0 @@ -1397,7 +340,8 @@ services: user: "0:0" command: ["-config", "/etc/fluxmq/config.yaml"] depends_on: - - fluxmq-auth + fluxmq-auth: + condition: service_started restart: on-failure ports: - ${MG_COAP_PORT}:5683/udp @@ -1415,8 +359,10 @@ services: user: "0:0" command: ["-config", "/etc/fluxmq/config.yaml"] depends_on: - - fluxmq-node1 - - fluxmq-auth + fluxmq-node1: + condition: service_started + fluxmq-auth: + condition: service_started restart: on-failure ports: - ${MG_FLUXMQ_API_PORT_2}:8082 @@ -1433,8 +379,10 @@ services: user: "0:0" command: ["-config", "/etc/fluxmq/config.yaml"] depends_on: - - fluxmq-node1 - - fluxmq-auth + fluxmq-node1: + condition: service_started + fluxmq-auth: + condition: service_started restart: on-failure ports: - ${MG_FLUXMQ_API_PORT_3}:8082 @@ -1448,6 +396,9 @@ services: fluxmq-auth: image: ghcr.io/absmach/magistrala/fluxmq:${MG_RELEASE_TAG} container_name: magistrala-fluxmq-auth + depends_on: + atom-bootstrap: + condition: service_completed_successfully restart: on-failure environment: MG_FLUXMQ_LOG_LEVEL: ${MG_FLUXMQ_LOG_LEVEL} @@ -1457,92 +408,45 @@ services: MG_FLUXMQ_CACHE_NUM_COUNTERS: ${MG_FLUXMQ_CACHE_NUM_COUNTERS} MG_FLUXMQ_CACHE_MAX_COST: ${MG_FLUXMQ_CACHE_MAX_COST} MG_FLUXMQ_CACHE_BUFFER_ITEMS: ${MG_FLUXMQ_CACHE_BUFFER_ITEMS} - MG_CLIENTS_GRPC_URL: ${MG_CLIENTS_GRPC_URL} - MG_CLIENTS_GRPC_TIMEOUT: ${MG_CLIENTS_GRPC_TIMEOUT} - MG_CLIENTS_GRPC_CLIENT_CERT: ${MG_CLIENTS_GRPC_CLIENT_CERT:+/clients-grpc-client.crt} - MG_CLIENTS_GRPC_CLIENT_KEY: ${MG_CLIENTS_GRPC_CLIENT_KEY:+/clients-grpc-client.key} - MG_CLIENTS_GRPC_SERVER_CA_CERTS: ${MG_CLIENTS_GRPC_SERVER_CA_CERTS:+/clients-grpc-server-ca.crt} - MG_CHANNELS_GRPC_URL: ${MG_CHANNELS_GRPC_URL} - MG_CHANNELS_GRPC_TIMEOUT: ${MG_CHANNELS_GRPC_TIMEOUT} - MG_CHANNELS_GRPC_CLIENT_CERT: ${MG_CHANNELS_GRPC_CLIENT_CERT:+/channels-grpc-client.crt} - MG_CHANNELS_GRPC_CLIENT_KEY: ${MG_CHANNELS_GRPC_CLIENT_KEY:+/channels-grpc-client.key} - MG_CHANNELS_GRPC_SERVER_CA_CERTS: ${MG_CHANNELS_GRPC_SERVER_CA_CERTS:+/channels-grpc-server-ca.crt} - MG_DOMAINS_GRPC_URL: ${MG_DOMAINS_GRPC_URL} - MG_DOMAINS_GRPC_TIMEOUT: ${MG_DOMAINS_GRPC_TIMEOUT} - MG_DOMAINS_GRPC_CLIENT_CERT: ${MG_DOMAINS_GRPC_CLIENT_CERT:+/domains-grpc-client.crt} - MG_DOMAINS_GRPC_CLIENT_KEY: ${MG_DOMAINS_GRPC_CLIENT_KEY:+/domains-grpc-client.key} - MG_DOMAINS_GRPC_SERVER_CA_CERTS: ${MG_DOMAINS_GRPC_SERVER_CA_CERTS:+/domains-grpc-server-ca.crt} + MG_MESSAGE_BROKER_URL: ${MG_MESSAGE_BROKER_URL} + MG_FLUXMQ_PUBLISH_HTTP_HOST: ${MG_FLUXMQ_PUBLISH_HTTP_HOST} + MG_FLUXMQ_PUBLISH_HTTP_PORT: ${MG_FLUXMQ_PUBLISH_HTTP_PORT} + ATOM_URL: ${ATOM_URL} + ATOM_SERVICE_TOKEN: ${MG_ATOM_TOKEN_FLUXMQ_AUTH} + ATOM_SERVICE_USERNAME: ${ATOM_SERVICE_USERNAME} + ATOM_SERVICE_SECRET: ${ATOM_SERVICE_SECRET} + ATOM_ADMIN_TOKEN: ${ATOM_ADMIN_TOKEN} + ATOM_ADMIN_USERNAME: ${ATOM_ADMIN_USERNAME} + ATOM_ADMIN_SECRET: ${ATOM_ADMIN_SECRET} + ATOM_JWKS_URL: ${ATOM_JWKS_URL} + ATOM_JWT_ISSUER: ${ATOM_JWT_ISSUER} + ATOM_JWT_AUDIENCE: ${ATOM_JWT_AUDIENCE} + ATOM_TIMEOUT: ${ATOM_TIMEOUT} MG_JAEGER_URL: ${MG_JAEGER_URL} MG_JAEGER_TRACE_RATIO: ${MG_JAEGER_TRACE_RATIO} networks: - magistrala-base-net - volumes: - # Clients gRPC mTLS client certificates - - type: bind - source: ${MG_CLIENTS_GRPC_CLIENT_CERT:-./ssl/placeholder} - target: /clients-grpc-client${MG_CLIENTS_GRPC_CLIENT_CERT:+.crt} - bind: - create_host_path: true - - type: bind - source: ${MG_CLIENTS_GRPC_CLIENT_KEY:-./ssl/placeholder} - target: /clients-grpc-client${MG_CLIENTS_GRPC_CLIENT_KEY:+.key} - bind: - create_host_path: true - - type: bind - source: ${MG_CLIENTS_GRPC_SERVER_CA_CERTS:-./ssl/placeholder} - target: /clients-grpc-server-ca${MG_CLIENTS_GRPC_SERVER_CA_CERTS:+.crt} - bind: - create_host_path: true - # Channels gRPC mTLS client certificates - - type: bind - source: ${MG_CHANNELS_GRPC_CLIENT_CERT:-./ssl/placeholder} - target: /channels-grpc-client${MG_CHANNELS_GRPC_CLIENT_CERT:+.crt} - bind: - create_host_path: true - - type: bind - source: ${MG_CHANNELS_GRPC_CLIENT_KEY:-./ssl/placeholder} - target: /channels-grpc-client${MG_CHANNELS_GRPC_CLIENT_KEY:+.key} - bind: - create_host_path: true - - type: bind - source: ${MG_CHANNELS_GRPC_SERVER_CA_CERTS:-./ssl/placeholder} - target: /channels-grpc-server-ca${MG_CHANNELS_GRPC_SERVER_CA_CERTS:+.crt} - bind: - create_host_path: true - # Domains gRPC mTLS client certificates - - type: bind - source: ${MG_DOMAINS_GRPC_CLIENT_CERT:-./ssl/placeholder} - target: /domains-grpc-client${MG_DOMAINS_GRPC_CLIENT_CERT:+.crt} - bind: - create_host_path: true - - type: bind - source: ${MG_DOMAINS_GRPC_CLIENT_KEY:-./ssl/placeholder} - target: /domains-grpc-client${MG_DOMAINS_GRPC_CLIENT_KEY:+.key} - bind: - create_host_path: true - - type: bind - source: ${MG_DOMAINS_GRPC_SERVER_CA_CERTS:-./ssl/placeholder} - target: /domains-grpc-server-ca${MG_DOMAINS_GRPC_SERVER_CA_CERTS:+.crt} - bind: - create_host_path: true ui: image: ghcr.io/absmach/magistrala/ui-mg:${MG_RELEASE_TAG} container_name: magistrala-ui + restart: unless-stopped ports: - 3000:3000 networks: - magistrala-base-net environment: MG_AUTH_URL: ${MG_AUTH_URL} + MG_ATOM_URL: ${MG_ATOM_URL} + ATOM_URL: ${ATOM_PUBLIC_URL} MG_DOMAINS_URL: ${MG_DOMAINS_URL} MG_USERS_URL: ${MG_USERS_URL} MG_CLIENTS_URL: ${MG_CLIENTS_URL} MG_CHANNELS_URL: ${MG_CHANNELS_URL} MG_GROUPS_URL: ${MG_GROUPS_URL} MG_BOOTSTRAP_URL: ${MG_BOOTSTRAP_URL} - MG_CERTS_URL: ${MG_CERTS_URL} MG_HTTP_ADAPTER_URL: ${MG_HTTP_ADAPTER_URL} + MG_PUBLISH_PROXY_URL: ${MG_PUBLISH_PROXY_URL} MG_READER_URL: ${MG_READER_URL} MG_BACKEND_URL: ${MG_UI_BACKEND_URL} MG_JOURNAL_URL: ${MG_JOURNAL_URL} @@ -1556,6 +460,7 @@ services: MG_UI_BASE_PATH: ${MG_UI_BASE_PATH} MG_NEXTAUTH_BASE_PATH: ${MG_NEXTAUTH_BASE_PATH} MG_UI_TYPE: ${MG_UI_TYPE} + MG_UI_CLIENT_TYPE: ${MG_UI_CLIENT_TYPE} MG_UI_BASEURL: ${MG_UI_BASEURL} NEXTAUTH_URL: ${NEXTAUTH_URL} NEXTAUTH_SECRET: ${NEXTAUTH_SECRET} @@ -1604,6 +509,10 @@ services: MG_BACKEND_DB_SSL_KEY: ${MG_UI_BACKEND_DB_SSL_KEY} MG_BACKEND_DB_SSL_ROOT_CERT: ${MG_UI_BACKEND_DB_SSL_ROOT_CERT} MG_BACKEND_INSTANCE_ID: ${MG_UI_BACKEND_INSTANCE_ID} + ATOM_URL: ${ATOM_URL} + ATOM_SERVICE_USERNAME: ${ATOM_SERVICE_USERNAME} + ATOM_SERVICE_SECRET: ${ATOM_SERVICE_SECRET} + ATOM_TIMEOUT: ${ATOM_TIMEOUT} MG_AUTH_GRPC_URL: ${MG_AUTH_GRPC_URL} MG_AUTH_GRPC_TIMEOUT: ${MG_AUTH_GRPC_TIMEOUT} MG_AUTH_GRPC_CLIENT_CERT: ${MG_AUTH_GRPC_CLIENT_CERT:+/auth-grpc-client.crt} @@ -1635,6 +544,8 @@ services: MG_JAEGER_URL: ${MG_JAEGER_URL} MG_JAEGER_TRACE_RATIO: ${MG_JAEGER_TRACE_RATIO} depends_on: + atom-bootstrap: + condition: service_completed_successfully ui-backend-db: condition: service_healthy seaweedfs-s3: @@ -1775,7 +686,10 @@ services: image: ghcr.io/absmach/magistrala/timescale-reader:${MG_RELEASE_TAG} container_name: magistrala-timescale-reader depends_on: - - timescale + timescale: + condition: service_started + atom-bootstrap: + condition: service_completed_successfully restart: on-failure environment: MG_TIMESCALE_READER_LOG_LEVEL: ${MG_TIMESCALE_READER_LOG_LEVEL} @@ -1792,16 +706,12 @@ services: MG_TIMESCALE_SSL_CERT: ${MG_TIMESCALE_SSL_CERT} MG_TIMESCALE_SSL_KEY: ${MG_TIMESCALE_SSL_KEY} MG_TIMESCALE_SSL_ROOT_CERT: ${MG_TIMESCALE_SSL_ROOT_CERT} - MG_CLIENTS_GRPC_URL: ${MG_CLIENTS_GRPC_URL} - MG_CLIENTS_GRPC_TIMEOUT: ${MG_CLIENTS_GRPC_TIMEOUT} - MG_CLIENTS_GRPC_CLIENT_CERT: ${MG_CLIENTS_GRPC_CLIENT_CERT:+/clients-grpc-client.crt} - MG_CLIENTS_GRPC_CLIENT_KEY: ${MG_CLIENTS_GRPC_CLIENT_KEY:+/clients-grpc-client.key} - MG_CLIENTS_GRPC_SERVER_CA_CERTS: ${MG_CLIENTS_GRPC_SERVER_CA_CERTS:+/clients-grpc-server-ca.crt} - MG_CHANNELS_GRPC_URL: ${MG_CHANNELS_GRPC_URL} - MG_CHANNELS_GRPC_TIMEOUT: ${MG_CHANNELS_GRPC_TIMEOUT} - MG_CHANNELS_GRPC_CLIENT_CERT: ${MG_CHANNELS_GRPC_CLIENT_CERT:+/channels-grpc-client.crt} - MG_CHANNELS_GRPC_CLIENT_KEY: ${MG_CHANNELS_GRPC_CLIENT_KEY:+/channels-grpc-client.key} - MG_CHANNELS_GRPC_SERVER_CA_CERTS: ${MG_CHANNELS_GRPC_SERVER_CA_CERTS:+/channels-grpc-server-ca.crt} + ATOM_URL: ${ATOM_URL} + ATOM_SERVICE_TOKEN: ${MG_ATOM_TOKEN_TIMESCALE_READER} + ATOM_JWKS_URL: ${ATOM_JWKS_URL} + ATOM_JWT_ISSUER: ${ATOM_JWT_ISSUER} + ATOM_JWT_AUDIENCE: ${ATOM_JWT_AUDIENCE} + ATOM_TIMEOUT: ${ATOM_TIMEOUT} MG_TIMESCALE_READER_GRPC_URL: ${MG_TIMESCALE_READER_GRPC_URL} MG_TIMESCALE_READER_GRPC_PORT: ${MG_TIMESCALE_READER_GRPC_PORT} MG_TIMESCALE_READER_GRPC_HOST: ${MG_TIMESCALE_READER_GRPC_HOST} @@ -1812,11 +722,6 @@ services: MG_TIMESCALE_READER_GRPC_CLIENT_KEY: ${MG_TIMESCALE_READER_GRPC_CLIENT_KEY:+/readers-grpc-client.key} MG_TIMESCALE_READER_GRPC_SERVER_CERT: ${MG_TIMESCALE_READER_GRPC_SERVER_CERT:+/readers-grpc-server.crt} MG_TIMESCALE_READER_GRPC_SERVER_KEY: ${MG_TIMESCALE_READER_GRPC_SERVER_KEY:+/readers-grpc-server.key} - MG_AUTH_GRPC_URL: ${MG_AUTH_GRPC_URL} - MG_AUTH_GRPC_TIMEOUT: ${MG_AUTH_GRPC_TIMEOUT} - MG_AUTH_GRPC_CLIENT_CERT: ${MG_AUTH_GRPC_CLIENT_CERT:+/auth-grpc-client.crt} - MG_AUTH_GRPC_CLIENT_KEY: ${MG_AUTH_GRPC_CLIENT_KEY:+/auth-grpc-client.key} - MG_AUTH_GRPC_SERVER_CA_CERTS: ${MG_AUTH_GRPC_SERVER_CA_CERTS:+/auth-grpc-server-ca.crt} MG_SEND_TELEMETRY: ${MG_SEND_TELEMETRY} MG_TIMESCALE_READER_INSTANCE_ID: ${MG_TIMESCALE_READER_INSTANCE_ID} ports: @@ -1825,54 +730,6 @@ services: networks: - magistrala-base-net volumes: - # Auth gRPC client certificates - - type: bind - source: ${MG_AUTH_GRPC_CLIENT_CERT:-./ssl/placeholder} - target: /auth-grpc-client${MG_AUTH_GRPC_CLIENT_CERT:+.crt} - bind: - create_host_path: true - - type: bind - source: ${MG_AUTH_GRPC_CLIENT_KEY:-./ssl/placeholder} - target: /auth-grpc-client${MG_AUTH_GRPC_CLIENT_KEY:+.key} - bind: - create_host_path: true - - type: bind - source: ${MG_AUTH_GRPC_SERVER_CA_CERTS:-./ssl/placeholder} - target: /auth-grpc-server-ca${MG_AUTH_GRPC_SERVER_CA_CERTS:+.crt} - bind: - create_host_path: true - # Clients gRPC client certificates - - type: bind - source: ${MG_CLIENTS_GRPC_CLIENT_CERT:-./ssl/placeholder} - target: /clients-grpc-client${MG_CLIENTS_GRPC_CLIENT_CERT:+.crt} - bind: - create_host_path: true - - type: bind - source: ${MG_CLIENTS_GRPC_CLIENT_KEY:-./ssl/placeholder} - target: /clients-grpc-client${MG_CLIENTS_GRPC_CLIENT_KEY:+.key} - bind: - create_host_path: true - - type: bind - source: ${MG_CLIENTS_GRPC_SERVER_CA_CERTS:-./ssl/placeholder} - target: /clients-grpc-server-ca${MG_CLIENTS_GRPC_SERVER_CA_CERTS:+.crt} - bind: - create_host_path: true - # Channels gRPC client certificates - - type: bind - source: ${MG_CHANNELS_GRPC_CLIENT_CERT:-./ssl/placeholder} - target: /channels-grpc-client${MG_CHANNELS_GRPC_CLIENT_CERT:+.crt} - bind: - create_host_path: true - - type: bind - source: ${MG_CHANNELS_GRPC_CLIENT_KEY:-./ssl/placeholder} - target: /channels-grpc-client${MG_CHANNELS_GRPC_CLIENT_KEY:+.key} - bind: - create_host_path: true - - type: bind - source: ${MG_CHANNELS_GRPC_SERVER_CA_CERTS:-./ssl/placeholder} - target: /channels-grpc-server-ca${MG_CHANNELS_GRPC_SERVER_CA_CERTS:+.crt} - bind: - create_host_path: true # Reader gRPC server and client certificates - type: bind source: ${MG_TIMESCALE_READER_GRPC_SERVER_CERT:-./ssl/placeholder} @@ -1958,12 +815,21 @@ services: image: ghcr.io/absmach/magistrala/re:${MG_RELEASE_TAG} container_name: magistrala-re depends_on: - - re-db - - spicedb-migrate - - nginx + re-db: + condition: service_started + atom-bootstrap: + condition: service_completed_successfully + nginx: + condition: service_started restart: on-failure environment: MG_RE_LOG_LEVEL: ${MG_RE_LOG_LEVEL} + ATOM_URL: ${ATOM_URL} + ATOM_SERVICE_TOKEN: ${MG_ATOM_TOKEN_RE} + ATOM_JWKS_URL: ${ATOM_JWKS_URL} + ATOM_JWT_ISSUER: ${ATOM_JWT_ISSUER} + ATOM_JWT_AUDIENCE: ${ATOM_JWT_AUDIENCE} + ATOM_TIMEOUT: ${ATOM_TIMEOUT} MG_RE_HTTP_PORT: ${MG_RE_HTTP_PORT} MG_RE_HTTP_HOST: ${MG_RE_HTTP_HOST} MG_RE_HTTP_SERVER_CERT: ${MG_RE_HTTP_SERVER_CERT} @@ -1990,15 +856,6 @@ services: MG_JAEGER_URL: ${MG_JAEGER_URL} MG_JAEGER_TRACE_RATIO: ${MG_JAEGER_TRACE_RATIO} MG_SEND_TELEMETRY: ${MG_SEND_TELEMETRY} - MG_AUTH_GRPC_URL: ${MG_AUTH_GRPC_URL} - MG_AUTH_GRPC_TIMEOUT: ${MG_AUTH_GRPC_TIMEOUT} - MG_AUTH_GRPC_CLIENT_CERT: ${MG_AUTH_GRPC_CLIENT_CERT:+/auth-grpc-client.crt} - MG_AUTH_GRPC_CLIENT_KEY: ${MG_AUTH_GRPC_CLIENT_KEY:+/auth-grpc-client.key} - MG_AUTH_GRPC_SERVER_CA_CERTS: ${MG_AUTH_GRPC_SERVER_CA_CERTS:+/auth-grpc-server-ca.crt} - MG_SPICEDB_PRE_SHARED_KEY: ${MG_SPICEDB_PRE_SHARED_KEY} - MG_SPICEDB_HOST: ${MG_SPICEDB_HOST} - MG_SPICEDB_PORT: ${MG_SPICEDB_PORT} - MG_SPICEDB_SCHEMA_FILE: ${MG_SPICEDB_SCHEMA_FILE} MG_PERMISSIONS_FILE: ${MG_PERMISSIONS_FILE} MG_RE_INSTANCE_ID: ${MG_RE_INSTANCE_ID} MG_EMAIL_HOST: ${MG_EMAIL_HOST} @@ -2013,11 +870,6 @@ services: MG_TIMESCALE_READER_GRPC_CLIENT_CERT: ${MG_TIMESCALE_READER_GRPC_CLIENT_CERT} MG_TIMESCALE_READER_GRPC_CLIENT_CA_CERTS: ${MG_TIMESCALE_READER_GRPC_CLIENT_CA_CERTS} MG_TIMESCALE_READER_GRPC_CLIENT_KEY: ${MG_TIMESCALE_READER_GRPC_CLIENT_KEY} - MG_DOMAINS_GRPC_URL: ${MG_DOMAINS_GRPC_URL} - MG_DOMAINS_GRPC_TIMEOUT: ${MG_DOMAINS_GRPC_TIMEOUT} - MG_DOMAINS_GRPC_CLIENT_CERT: ${MG_DOMAINS_GRPC_CLIENT_CERT:+/domains-grpc-client.crt} - MG_DOMAINS_GRPC_CLIENT_KEY: ${MG_DOMAINS_GRPC_CLIENT_KEY:+/domains-grpc-client.key} - MG_DOMAINS_GRPC_SERVER_CA_CERTS: ${MG_DOMAINS_GRPC_SERVER_CA_CERTS:+/domains-grpc-server-ca.crt} MG_ALLOW_UNVERIFIED_USER: ${MG_ALLOW_UNVERIFIED_USER} ports: - ${MG_RE_HTTP_PORT}:${MG_RE_HTTP_PORT} @@ -2025,40 +877,7 @@ services: - magistrala-base-net volumes: - ./permission.yaml:${MG_PERMISSIONS_FILE} - - ./spicedb/schema.zed:${MG_SPICEDB_SCHEMA_FILE} - ./templates/${MG_RE_EMAIL_TEMPLATE}:/email.tmpl - # Auth gRPC client certificates - - type: bind - source: ${MG_AUTH_GRPC_CLIENT_CERT:-./ssl/placeholder} - target: /auth-grpc-client.crt - bind: - create_host_path: true - - type: bind - source: ${MG_AUTH_GRPC_CLIENT_KEY:-./ssl/placeholder} - target: /auth-grpc-client.key - bind: - create_host_path: true - - type: bind - source: ${MG_AUTH_GRPC_SERVER_CA_CERTS:-./ssl/placeholder} - target: /auth-grpc-server-ca.crt - bind: - create_host_path: true - # Domains gRPC client certificates - - type: bind - source: ${MG_DOMAINS_GRPC_CLIENT_CERT:-./ssl/placeholder} - target: /domains-grpc-client.crt - bind: - create_host_path: true - - type: bind - source: ${MG_DOMAINS_GRPC_CLIENT_KEY:-./ssl/placeholder} - target: /domains-grpc-client.key - bind: - create_host_path: true - - type: bind - source: ${MG_DOMAINS_GRPC_SERVER_CA_CERTS:-./ssl/placeholder} - target: /domains-grpc-server-ca.crt - bind: - create_host_path: true alarms-db: image: docker.io/postgres:18.0-alpine3.22 @@ -2080,12 +899,21 @@ services: image: ghcr.io/absmach/magistrala/alarms:${MG_RELEASE_TAG} container_name: magistrala-alarms depends_on: - - alarms-db - - spicedb-migrate - - nginx + alarms-db: + condition: service_started + atom-bootstrap: + condition: service_completed_successfully + nginx: + condition: service_started restart: on-failure environment: MG_ALARMS_LOG_LEVEL: ${MG_ALARMS_LOG_LEVEL} + ATOM_URL: ${ATOM_URL} + ATOM_SERVICE_TOKEN: ${MG_ATOM_TOKEN_ALARMS} + ATOM_JWKS_URL: ${ATOM_JWKS_URL} + ATOM_JWT_ISSUER: ${ATOM_JWT_ISSUER} + ATOM_JWT_AUDIENCE: ${ATOM_JWT_AUDIENCE} + ATOM_TIMEOUT: ${ATOM_TIMEOUT} MG_ALARMS_HTTP_PORT: ${MG_ALARMS_HTTP_PORT} MG_ALARMS_HTTP_HOST: ${MG_ALARMS_HTTP_HOST} MG_ALARMS_HTTP_SERVER_CERT: ${MG_ALARMS_HTTP_SERVER_CERT} @@ -2103,20 +931,6 @@ services: MG_ES_URL: ${MG_ES_URL} MG_JAEGER_URL: ${MG_JAEGER_URL} MG_JAEGER_TRACE_RATIO: ${MG_JAEGER_TRACE_RATIO} - MG_AUTH_GRPC_URL: ${MG_AUTH_GRPC_URL} - MG_AUTH_GRPC_TIMEOUT: ${MG_AUTH_GRPC_TIMEOUT} - MG_AUTH_GRPC_CLIENT_CERT: ${MG_AUTH_GRPC_CLIENT_CERT:+/auth-grpc-client.crt} - MG_AUTH_GRPC_CLIENT_KEY: ${MG_AUTH_GRPC_CLIENT_KEY:+/auth-grpc-client.key} - MG_AUTH_GRPC_SERVER_CA_CERTS: ${MG_AUTH_GRPC_SERVER_CA_CERTS:+/auth-grpc-server-ca.crt} - MG_DOMAINS_GRPC_URL: ${MG_DOMAINS_GRPC_URL} - MG_DOMAINS_GRPC_TIMEOUT: ${MG_DOMAINS_GRPC_TIMEOUT} - MG_DOMAINS_GRPC_CLIENT_CERT: ${MG_DOMAINS_GRPC_CLIENT_CERT:+/domains-grpc-client.crt} - MG_DOMAINS_GRPC_CLIENT_KEY: ${MG_DOMAINS_GRPC_CLIENT_KEY:+/domains-grpc-client.key} - MG_DOMAINS_GRPC_SERVER_CA_CERTS: ${MG_DOMAINS_GRPC_SERVER_CA_CERTS:+/domains-grpc-server-ca.crt} - MG_SPICEDB_PRE_SHARED_KEY: ${MG_SPICEDB_PRE_SHARED_KEY} - MG_SPICEDB_HOST: ${MG_SPICEDB_HOST} - MG_SPICEDB_PORT: ${MG_SPICEDB_PORT} - MG_SPICEDB_SCHEMA_FILE: ${MG_SPICEDB_SCHEMA_FILE} MG_PERMISSIONS_FILE: ${MG_PERMISSIONS_FILE} MG_ALARMS_INSTANCE_ID: ${MG_ALARMS_INSTANCE_ID} MG_ALARMS_EVENT_CONSUMER: ${MG_ALARMS_EVENT_CONSUMER} @@ -2127,39 +941,6 @@ services: - magistrala-base-net volumes: - ./permission.yaml:${MG_PERMISSIONS_FILE} - - ./spicedb/schema.zed:${MG_SPICEDB_SCHEMA_FILE} - # Auth gRPC client certificates - - type: bind - source: ${MG_AUTH_GRPC_CLIENT_CERT:-./ssl/placeholder} - target: /auth-grpc-client.crt - bind: - create_host_path: true - - type: bind - source: ${MG_AUTH_GRPC_CLIENT_KEY:-./ssl/placeholder} - target: /auth-grpc-client.key - bind: - create_host_path: true - - type: bind - source: ${MG_AUTH_GRPC_SERVER_CA_CERTS:-./ssl/placeholder} - target: /auth-grpc-server-ca.crt - bind: - create_host_path: true - # Domains gRPC client certificates - - type: bind - source: ${MG_DOMAINS_GRPC_CLIENT_CERT:-./ssl/placeholder} - target: /domains-grpc-client.crt - bind: - create_host_path: true - - type: bind - source: ${MG_DOMAINS_GRPC_CLIENT_KEY:-./ssl/placeholder} - target: /domains-grpc-client.key - bind: - create_host_path: true - - type: bind - source: ${MG_DOMAINS_GRPC_SERVER_CA_CERTS:-./ssl/placeholder} - target: /domains-grpc-server-ca.crt - bind: - create_host_path: true reports-db: image: docker.io/postgres:18.0-alpine3.22 @@ -2181,12 +962,21 @@ services: image: ghcr.io/absmach/magistrala/reports:${MG_RELEASE_TAG} container_name: magistrala-reports depends_on: - - reports-db - - spicedb-migrate - - nginx + reports-db: + condition: service_started + atom-bootstrap: + condition: service_completed_successfully + nginx: + condition: service_started restart: on-failure environment: MG_REPORTS_LOG_LEVEL: ${MG_REPORTS_LOG_LEVEL} + ATOM_URL: ${ATOM_URL} + ATOM_SERVICE_TOKEN: ${MG_ATOM_TOKEN_REPORTS} + ATOM_JWKS_URL: ${ATOM_JWKS_URL} + ATOM_JWT_ISSUER: ${ATOM_JWT_ISSUER} + ATOM_JWT_AUDIENCE: ${ATOM_JWT_AUDIENCE} + ATOM_TIMEOUT: ${ATOM_TIMEOUT} MG_REPORTS_HTTP_PORT: ${MG_REPORTS_HTTP_PORT} MG_REPORTS_HTTP_HOST: ${MG_REPORTS_HTTP_HOST} MG_REPORTS_HTTP_SERVER_CERT: ${MG_REPORTS_HTTP_SERVER_CERT} @@ -2215,15 +1005,6 @@ services: MG_JAEGER_URL: ${MG_JAEGER_URL} MG_JAEGER_TRACE_RATIO: ${MG_JAEGER_TRACE_RATIO} MG_SEND_TELEMETRY: ${MG_SEND_TELEMETRY} - MG_AUTH_GRPC_URL: ${MG_AUTH_GRPC_URL} - MG_AUTH_GRPC_TIMEOUT: ${MG_AUTH_GRPC_TIMEOUT} - MG_AUTH_GRPC_CLIENT_CERT: ${MG_AUTH_GRPC_CLIENT_CERT:+/auth-grpc-client.crt} - MG_AUTH_GRPC_CLIENT_KEY: ${MG_AUTH_GRPC_CLIENT_KEY:+/auth-grpc-client.key} - MG_AUTH_GRPC_SERVER_CA_CERTS: ${MG_AUTH_GRPC_SERVER_CA_CERTS:+/auth-grpc-server-ca.crt} - MG_SPICEDB_PRE_SHARED_KEY: ${MG_SPICEDB_PRE_SHARED_KEY} - MG_SPICEDB_HOST: ${MG_SPICEDB_HOST} - MG_SPICEDB_PORT: ${MG_SPICEDB_PORT} - MG_SPICEDB_SCHEMA_FILE: ${MG_SPICEDB_SCHEMA_FILE} MG_PERMISSIONS_FILE: ${MG_PERMISSIONS_FILE} MG_REPORTS_INSTANCE_ID: ${MG_RE_INSTANCE_ID} MG_EMAIL_HOST: ${MG_EMAIL_HOST} @@ -2238,11 +1019,6 @@ services: MG_TIMESCALE_READER_GRPC_CLIENT_CERT: ${MG_TIMESCALE_READER_GRPC_CLIENT_CERT} MG_TIMESCALE_READER_GRPC_SERVER_CA_CERTS: ${MG_TIMESCALE_READER_GRPC_SERVER_CA_CERTS} MG_TIMESCALE_READER_GRPC_CLIENT_KEY: ${MG_TIMESCALE_READER_GRPC_CLIENT_KEY} - MG_DOMAINS_GRPC_URL: ${MG_DOMAINS_GRPC_URL} - MG_DOMAINS_GRPC_TIMEOUT: ${MG_DOMAINS_GRPC_TIMEOUT} - MG_DOMAINS_GRPC_CLIENT_CERT: ${MG_DOMAINS_GRPC_CLIENT_CERT:+/domains-grpc-client.crt} - MG_DOMAINS_GRPC_CLIENT_KEY: ${MG_DOMAINS_GRPC_CLIENT_KEY:+/domains-grpc-client.key} - MG_DOMAINS_GRPC_SERVER_CA_CERTS: ${MG_DOMAINS_GRPC_SERVER_CA_CERTS:+/domains-grpc-server-ca.crt} MG_ALLOW_UNVERIFIED_USER: ${MG_ALLOW_UNVERIFIED_USER} ports: - ${MG_REPORTS_HTTP_PORT}:${MG_REPORTS_HTTP_PORT} @@ -2250,40 +1026,7 @@ services: - magistrala-base-net volumes: - ./permission.yaml:${MG_PERMISSIONS_FILE} - - ./spicedb/schema.zed:${MG_SPICEDB_SCHEMA_FILE} - ./templates/${MG_REPORTS_EMAIL_TEMPLATE}:/email.tmpl - # Auth gRPC client certificates - - type: bind - source: ${MG_AUTH_GRPC_CLIENT_CERT:-./ssl/placeholder} - target: /auth-grpc-client.crt - bind: - create_host_path: true - - type: bind - source: ${MG_AUTH_GRPC_CLIENT_KEY:-./ssl/placeholder} - target: /auth-grpc-client.key - bind: - create_host_path: true - - type: bind - source: ${MG_AUTH_GRPC_SERVER_CA_CERTS:-./ssl/placeholder} - target: /auth-grpc-server-ca.crt - bind: - create_host_path: true - # Domains gRPC client certificates - - type: bind - source: ${MG_DOMAINS_GRPC_CLIENT_CERT:-./ssl/placeholder} - target: /domains-grpc-client.crt - bind: - create_host_path: true - - type: bind - source: ${MG_DOMAINS_GRPC_CLIENT_KEY:-./ssl/placeholder} - target: /domains-grpc-client.key - bind: - create_host_path: true - - type: bind - source: ${MG_DOMAINS_GRPC_SERVER_CA_CERTS:-./ssl/placeholder} - target: /domains-grpc-server-ca.crt - bind: - create_host_path: true pdf-generator: image: gotenberg/gotenberg:8.25.1 @@ -2292,153 +1035,3 @@ services: - "4000:3000" networks: - magistrala-base-net - - certs: - image: ghcr.io/absmach/magistrala/certs:${MG_RELEASE_TAG} - container_name: magistrala-certs - depends_on: - openbao: - condition: service_healthy - certs-db: - condition: service_started - restart: on-failure - networks: - - magistrala-base-net - environment: - MG_CERTS_LOG_LEVEL: ${MG_CERTS_LOG_LEVEL} - MG_CERTS_HTTP_HOST: ${MG_CERTS_HTTP_HOST} - MG_CERTS_HTTP_PORT: ${MG_CERTS_HTTP_PORT} - MG_CERTS_GRPC_HOST: ${MG_CERTS_GRPC_HOST} - MG_CERTS_GRPC_PORT: ${MG_CERTS_GRPC_PORT} - MG_JAEGER_URL: ${MG_JAEGER_URL} - MG_JAEGER_TRACE_RATIO: ${MG_JAEGER_TRACE_RATIO} - MG_CERTS_OPENBAO_HOST: ${MG_CERTS_OPENBAO_HOST} - MG_CERTS_OPENBAO_APP_ROLE: ${MG_CERTS_OPENBAO_APP_ROLE} - MG_CERTS_OPENBAO_APP_SECRET: ${MG_CERTS_OPENBAO_APP_SECRET} - MG_CERTS_OPENBAO_NAMESPACE: ${MG_CERTS_OPENBAO_NAMESPACE} - MG_CERTS_OPENBAO_PKI_PATH: ${MG_CERTS_OPENBAO_PKI_PATH} - MG_CERTS_OPENBAO_ROLE: ${MG_CERTS_OPENBAO_ROLE} - MG_CERTS_OPENBAO_SECRET_ID_TTL: ${MG_CERTS_OPENBAO_SECRET_ID_TTL} - MG_CERTS_DB_HOST: ${MG_CERTS_DB_HOST} - MG_CERTS_DB_PORT: ${MG_CERTS_DB_PORT} - MG_CERTS_DB_USER: ${MG_CERTS_DB_USER} - MG_CERTS_DB_PASS: ${MG_CERTS_DB_PASS} - MG_CERTS_DB: ${MG_CERTS_DB} - MG_CERTS_DB_SSL_MODE: ${MG_CERTS_DB_SSL_MODE} - MG_AUTH_GRPC_URL: ${MG_AUTH_GRPC_URL} - MG_AUTH_GRPC_TIMEOUT: ${MG_AUTH_GRPC_TIMEOUT} - MG_AUTH_GRPC_CLIENT_CERT: ${MG_AUTH_GRPC_CLIENT_CERT:+/auth-grpc-client.crt} - MG_AUTH_GRPC_CLIENT_KEY: ${MG_AUTH_GRPC_CLIENT_KEY:+/auth-grpc-client.key} - MG_AUTH_GRPC_SERVER_CA_CERTS: ${MG_AUTH_GRPC_SERVER_CA_CERTS:+/auth-grpc-server-ca.crt} - MG_DOMAINS_GRPC_URL: ${MG_DOMAINS_GRPC_URL} - MG_DOMAINS_GRPC_TIMEOUT: ${MG_DOMAINS_GRPC_TIMEOUT} - MG_DOMAINS_GRPC_CLIENT_CERT: ${MG_DOMAINS_GRPC_CLIENT_CERT:+/domains-grpc-client.crt} - MG_DOMAINS_GRPC_CLIENT_KEY: ${MG_DOMAINS_GRPC_CLIENT_KEY:+/domains-grpc-client.key} - MG_DOMAINS_GRPC_SERVER_CA_CERTS: ${MG_DOMAINS_GRPC_SERVER_CA_CERTS:+/domains-grpc-server-ca.crt} - MG_CERTS_SECRET: ${MG_CERTS_SECRET} - MG_CERTS_SERVICE_TOKEN_PATH: ${MG_CERTS_SERVICE_TOKEN_PATH} - MG_CERTS_SECRET_ID_PATH: ${MG_CERTS_SECRET_ID_PATH} - MG_CERTS_SECRET_RENEW_THRESHOLD: ${MG_CERTS_SECRET_RENEW_THRESHOLD} - MG_CERTS_SECRET_CHECK_INTERVAL: ${MG_CERTS_SECRET_CHECK_INTERVAL} - MG_ALLOW_UNVERIFIED_USER: ${MG_ALLOW_UNVERIFIED_USER} - ports: - - ${MG_CERTS_HTTP_PORT}:${MG_CERTS_HTTP_PORT} - - ${MG_CERTS_GRPC_PORT}:${MG_CERTS_GRPC_PORT} - volumes: - - magistrala-openbao-data:/openbao:ro - # Auth gRPC client certificates - - type: bind - source: ${AM_AUTH_GRPC_CLIENT_CERT:-./ssl/placeholder} - target: /auth-grpc-client.crt - bind: - create_host_path: true - - type: bind - source: ${AM_AUTH_GRPC_CLIENT_KEY:-./ssl/placeholder} - target: /auth-grpc-client.key - bind: - create_host_path: true - - type: bind - source: ${AM_AUTH_GRPC_SERVER_CA_CERTS:-./ssl/placeholder} - target: /auth-grpc-server-ca.crt - bind: - create_host_path: true - # Domains gRPC client certificates - - type: bind - source: ${AM_DOMAINS_GRPC_CLIENT_CERT:-./ssl/placeholder} - target: /domains-grpc-client.crt - bind: - create_host_path: true - - type: bind - source: ${AM_DOMAINS_GRPC_CLIENT_KEY:-./ssl/placeholder} - target: /domains-grpc-client.key - bind: - create_host_path: true - - type: bind - source: ${AM_DOMAINS_GRPC_SERVER_CA_CERTS:-./ssl/placeholder} - target: /domains-grpc-server-ca.crt - bind: - create_host_path: true - - certs-db: - image: docker.io/postgres:16.2-alpine - container_name: magistrala-certs-db - restart: on-failure - networks: - - magistrala-base-net - command: postgres -c "max_connections=${MG_CERTS_DB_MAX_CONNECTIONS}" - environment: - POSTGRES_USER: ${MG_CERTS_DB_USER} - POSTGRES_PASSWORD: ${MG_CERTS_DB_PASS} - POSTGRES_DB: ${MG_CERTS_DB} - ports: - - 5454:5432 - volumes: - - magistrala-certs-db-volume:/var/lib/postgresql/data - - openbao: - image: openbao/openbao:2.4.0 - container_name: magistrala-openbao - restart: on-failure - networks: - - magistrala-base-net - ports: - - 8200:8200 - healthcheck: - test: ["CMD", "sh", "-c", "test -f /opt/openbao/data/service_token"] - interval: 5s - timeout: 3s - retries: 20 - start_period: 30s - environment: - - BAO_ADDR=http://127.0.0.1:8200 - - BAO_LOG_LEVEL=info - - MG_CERTS_OPENBAO_PKI_ROLE=${MG_CERTS_OPENBAO_ROLE} - - MG_CERTS_OPENBAO_APP_ROLE=${MG_CERTS_OPENBAO_APP_ROLE} - - MG_CERTS_OPENBAO_APP_SECRET=${MG_CERTS_OPENBAO_APP_SECRET} - - MG_CERTS_OPENBAO_SECRET_ID_TTL=${MG_CERTS_OPENBAO_SECRET_ID_TTL} - - MG_CERTS_OPENBAO_NAMESPACE=${MG_CERTS_OPENBAO_NAMESPACE} - - MG_CERTS_OPENBAO_PKI_CA_CN=${MG_CERTS_OPENBAO_PKI_CA_CN} - - MG_CERTS_OPENBAO_PKI_CA_OU=${MG_CERTS_OPENBAO_PKI_CA_OU} - - MG_CERTS_OPENBAO_PKI_CA_O=${MG_CERTS_OPENBAO_PKI_CA_O} - - MG_CERTS_OPENBAO_PKI_CA_C=${MG_CERTS_OPENBAO_PKI_CA_C} - - MG_CERTS_OPENBAO_PKI_CA_L=${MG_CERTS_OPENBAO_PKI_CA_L} - - MG_CERTS_OPENBAO_PKI_CA_ST=${MG_CERTS_OPENBAO_PKI_CA_ST} - - MG_CERTS_OPENBAO_PKI_CA_ADDR=${MG_CERTS_OPENBAO_PKI_CA_ADDR} - - MG_CERTS_OPENBAO_PKI_CA_PO=${MG_CERTS_OPENBAO_PKI_CA_PO} - - MG_CERTS_OPENBAO_PKI_CA_DNS_NAMES=${MG_CERTS_OPENBAO_PKI_CA_DNS_NAMES} - - MG_CERTS_OPENBAO_PKI_CA_IP_ADDRESSES=${MG_CERTS_OPENBAO_PKI_CA_IP_ADDRESSES} - - MG_CERTS_OPENBAO_PKI_CA_URI_SANS=${MG_CERTS_OPENBAO_PKI_CA_URI_SANS} - - MG_CERTS_OPENBAO_PKI_CA_EMAIL_ADDRESSES=${MG_CERTS_OPENBAO_PKI_CA_EMAIL_ADDRESSES} - - MG_CERTS_OPENBAO_UNSEAL_KEY_1=${MG_CERTS_OPENBAO_UNSEAL_KEY_1} - - MG_CERTS_OPENBAO_UNSEAL_KEY_2=${MG_CERTS_OPENBAO_UNSEAL_KEY_2} - - MG_CERTS_OPENBAO_UNSEAL_KEY_3=${MG_CERTS_OPENBAO_UNSEAL_KEY_3} - - MG_CERTS_OPENBAO_ROOT_TOKEN=${MG_CERTS_OPENBAO_ROOT_TOKEN} - cap_add: - - IPC_LOCK - mem_swappiness: 0 - volumes: - - magistrala-openbao-data:/opt/openbao/data - - magistrala-openbao-data:/opt/openbao/config - - ./openbao-entrypoint.sh:/entrypoint.sh - entrypoint: /bin/sh - command: /entrypoint.sh diff --git a/docker/fluxmq/node2.yaml b/docker/fluxmq/node2.yaml index 679881026..6bd1e9a68 100644 --- a/docker/fluxmq/node2.yaml +++ b/docker/fluxmq/node2.yaml @@ -34,7 +34,7 @@ server: amqp091: plain: addr: "0.0.0.0:5682" - health_addr: "0.0.0.0:8084" + health_addr: "0.0.0.0:8081" health_enabled: true shutdown_timeout: 30s diff --git a/docker/fluxmq/node3.yaml b/docker/fluxmq/node3.yaml index 286137519..472656692 100644 --- a/docker/fluxmq/node3.yaml +++ b/docker/fluxmq/node3.yaml @@ -34,7 +34,7 @@ server: amqp091: plain: addr: "0.0.0.0:5682" - health_addr: "0.0.0.0:8083" + health_addr: "0.0.0.0:8081" health_enabled: true shutdown_timeout: 30s diff --git a/docker/nginx/entrypoint.sh b/docker/nginx/entrypoint.sh index b04e65a0d..fbe0ea953 100755 --- a/docker/nginx/entrypoint.sh +++ b/docker/nginx/entrypoint.sh @@ -22,6 +22,7 @@ envsubst ' ${MG_RE_HTTP_PORT} ${MG_ALARMS_HTTP_PORT} ${MG_REPORTS_HTTP_PORT} + ${MG_FLUXMQ_PUBLISH_HTTP_PORT} ${MG_NGINX_AMQP_PORT}' < /etc/nginx/nginx.conf.template > /etc/nginx/nginx.conf exec nginx -g "daemon off;" diff --git a/docker/nginx/nginx-key.conf b/docker/nginx/nginx-key.conf index 4747d9358..aaa7d7552 100644 --- a/docker/nginx/nginx-key.conf +++ b/docker/nginx/nginx-key.conf @@ -12,7 +12,7 @@ include /etc/nginx/modules-enabled/*.conf; events { # Explanation: https://serverfault.com/questions/787919/optimal-value-for-nginx-worker-connections - # We'll keep 10k connections per core (assuming one worker per core) + # We'll keep 10k connections for each configured worker. worker_connections 10000; } @@ -53,14 +53,11 @@ http { server_name $dynamic_server_name; set $auth_upstream "auth:${MG_AUTH_HTTP_PORT}"; - set $domains_upstream "domains:${MG_DOMAINS_HTTP_PORT}"; - set $users_upstream "users:${MG_USERS_HTTP_PORT}"; - set $groups_upstream "groups:${MG_GROUPS_HTTP_PORT}"; - set $clients_upstream "clients:${MG_CLIENTS_HTTP_PORT}"; - set $channels_upstream "channels:${MG_CHANNELS_HTTP_PORT}"; + set $atom_upstream "atom:8080"; set $rules_upstream "re:${MG_RE_HTTP_PORT}"; set $alarms_upstream "alarms:${MG_ALARMS_HTTP_PORT}"; set $reports_upstream "reports:${MG_REPORTS_HTTP_PORT}"; + set $publish_upstream "fluxmq-auth:${MG_FLUXMQ_PUBLISH_HTTP_PORT}"; include snippets/ssl.conf; @@ -84,39 +81,38 @@ http { proxy_pass http://$auth_upstream; } - # Proxy pass to domains service - location ~ ^/(domains|invitations) { + # Proxy pass to Atom Auth/OIDC REST endpoints + location ~ ^/auth(/|$) { include snippets/proxy-headers.conf; add_header Access-Control-Expose-Headers Location; - proxy_pass http://$domains_upstream; + proxy_pass http://$atom_upstream; } - # Proxy pass to users service - location ~ ^/(users|password|verify-email|authorize|oauth/callback/[^/]+) { + # Proxy pass to Atom GraphQL and the developer console + location = /graphql { include snippets/proxy-headers.conf; add_header Access-Control-Expose-Headers Location; - proxy_pass http://$users_upstream; + proxy_pass http://$atom_upstream; } - # Proxy pass to groups service - location ~ "^/([a-fA-F0-9]{8}-[a-fA-F0-9]{4}-[a-fA-F0-9]{4}-[a-fA-F0-9]{4}-[a-fA-F0-9]{12})/(groups)" { + # Proxy pass to Atom custom GraphQL-backed endpoints + location ^~ /api/custom { include snippets/proxy-headers.conf; add_header Access-Control-Expose-Headers Location; - proxy_pass http://$groups_upstream; + proxy_pass http://$atom_upstream; } - # Proxy pass to clients service - location ~ "^/([0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12})/(clients)" { + location = /.well-known/jwks.json { include snippets/proxy-headers.conf; add_header Access-Control-Expose-Headers Location; - proxy_pass http://$clients_upstream; + proxy_pass http://$atom_upstream; } - # Proxy pass to channels service - location ~ "^/([0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12})/(channels)" { + # Proxy pass to Atom public certificate endpoints + location ^~ /certs/ { include snippets/proxy-headers.conf; add_header Access-Control-Expose-Headers Location; - proxy_pass http://$channels_upstream; + proxy_pass http://$atom_upstream; } # Proxy pass to rule engine service @@ -142,12 +138,19 @@ http { location /health { include snippets/proxy-headers.conf; - proxy_pass http://$clients_upstream; + proxy_pass http://$atom_upstream; } location /metrics { include snippets/proxy-headers.conf; - proxy_pass http://$clients_upstream; + proxy_pass http://$atom_upstream; + } + + # Proxy user-authenticated UI publishes to FluxMQ through Atom authz. + location ~ "^/([0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12})/channels/([0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12})/messages$" { + include snippets/proxy-headers.conf; + add_header Access-Control-Expose-Headers Location; + proxy_pass http://$publish_upstream; } # Proxy pass to FluxMQ HTTP API diff --git a/docker/nginx/nginx-x509.conf b/docker/nginx/nginx-x509.conf index 8e4f0aafd..34dc52aa2 100644 --- a/docker/nginx/nginx-x509.conf +++ b/docker/nginx/nginx-x509.conf @@ -14,7 +14,7 @@ include /etc/nginx/modules-enabled/*.conf; events { # Explanation: https://serverfault.com/questions/787919/optimal-value-for-nginx-worker-connections - # We'll keep 10k connections per core (assuming one worker per core) + # We'll keep 10k connections for each configured worker. worker_connections 10000; } @@ -60,14 +60,11 @@ http { server_name $dynamic_server_name; set $auth_upstream "auth:${MG_AUTH_HTTP_PORT}"; - set $domains_upstream "domains:${MG_DOMAINS_HTTP_PORT}"; - set $users_upstream "users:${MG_USERS_HTTP_PORT}"; - set $groups_upstream "groups:${MG_GROUPS_HTTP_PORT}"; - set $clients_upstream "clients:${MG_CLIENTS_HTTP_PORT}"; - set $channels_upstream "channels:${MG_CHANNELS_HTTP_PORT}"; + set $atom_upstream "atom:8080"; set $rules_upstream "re:${MG_RE_HTTP_PORT}"; set $alarms_upstream "alarms:${MG_ALARMS_HTTP_PORT}"; set $reports_upstream "reports:${MG_REPORTS_HTTP_PORT}"; + set $publish_upstream "fluxmq-auth:${MG_FLUXMQ_PUBLISH_HTTP_PORT}"; ssl_verify_client optional; include snippets/ssl.conf; @@ -93,39 +90,38 @@ http { proxy_pass http://$auth_upstream; } - # Proxy pass to domains service - location ~ ^/(domains|invitations) { + # Proxy pass to Atom Auth/OIDC REST endpoints + location ~ ^/auth(/|$) { include snippets/proxy-headers.conf; add_header Access-Control-Expose-Headers Location; - proxy_pass http://$domains_upstream; + proxy_pass http://$atom_upstream; } - # Proxy pass to users service - location ~ ^/(users|password|verify-email|authorize|oauth/callback/[^/]+) { + # Proxy pass to Atom GraphQL and the developer console + location = /graphql { include snippets/proxy-headers.conf; add_header Access-Control-Expose-Headers Location; - proxy_pass http://$users_upstream; + proxy_pass http://$atom_upstream; } - # Proxy pass to groups service - location ~ "^/([a-fA-F0-9]{8}-[a-fA-F0-9]{4}-[a-fA-F0-9]{4}-[a-fA-F0-9]{4}-[a-fA-F0-9]{12})/(groups)" { + # Proxy pass to Atom custom GraphQL-backed endpoints + location ^~ /api/custom { include snippets/proxy-headers.conf; add_header Access-Control-Expose-Headers Location; - proxy_pass http://$groups_upstream; + proxy_pass http://$atom_upstream; } - # Proxy pass to clients service - location ~ "^/([0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12})/(clients)" { + location = /.well-known/jwks.json { include snippets/proxy-headers.conf; add_header Access-Control-Expose-Headers Location; - proxy_pass http://$clients_upstream; + proxy_pass http://$atom_upstream; } - # Proxy pass to channels service - location ~ "^/([0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12})/(channels)" { + # Proxy pass to Atom public certificate endpoints + location ^~ /certs/ { include snippets/proxy-headers.conf; add_header Access-Control-Expose-Headers Location; - proxy_pass http://$channels_upstream; + proxy_pass http://$atom_upstream; } # Proxy pass to rule engine service @@ -151,12 +147,20 @@ http { location /health { include snippets/proxy-headers.conf; - proxy_pass http://$clients_upstream; + proxy_pass http://$atom_upstream; } location /metrics { include snippets/proxy-headers.conf; - proxy_pass http://$clients_upstream; + proxy_pass http://$atom_upstream; + } + + # Proxy user-authenticated UI publishes to FluxMQ through Atom authz. + location ~ "^/([0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12})/channels/([0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12})/messages$" { + include snippets/verify-ssl-client.conf; + include snippets/proxy-headers.conf; + add_header Access-Control-Expose-Headers Location; + proxy_pass http://$publish_upstream; } # Proxy pass to FluxMQ HTTP API diff --git a/docker/nginx/snippets/ssl.conf b/docker/nginx/snippets/ssl.conf index 109f6795b..4bf814917 100644 --- a/docker/nginx/snippets/ssl.conf +++ b/docker/nginx/snippets/ssl.conf @@ -12,5 +12,5 @@ ssl_prefer_server_ciphers on; ssl_ciphers "EECDH+AESGCM:EDH+AESGCM:AES256+EECDH:AES256+EDH"; ssl_ecdh_curve secp384r1; ssl_session_tickets off; -resolver 8.8.8.8 8.8.4.4 valid=300s; +resolver 127.0.0.11 ipv6=off valid=10s; resolver_timeout 5s; diff --git a/docker/openbao-entrypoint.sh b/docker/openbao-entrypoint.sh deleted file mode 100755 index c6adff838..000000000 --- a/docker/openbao-entrypoint.sh +++ /dev/null @@ -1,490 +0,0 @@ -#!/bin/sh -# Copyright (c) Abstract Machines -# SPDX-License-Identifier: Apache-2.0 - -set -e - -apk add --no-cache jq - -# Create required directories -mkdir -p /opt/openbao/config /opt/openbao/data /opt/openbao/logs - -cat > /opt/openbao/config/config.hcl << 'EOF' -storage "file" { - path = "/opt/openbao/data" -} -listener "tcp" { - address = "0.0.0.0:8200" - tls_disable = true -} -ui = true -log_level = "Info" -disable_mlock = true -# API timeout settings -default_lease_ttl = "168h" -max_lease_ttl = "720h" -EOF - -export BAO_ADDR=http://127.0.0.1:8200 - -wait_for_bao() { - local retries=30 - while [ "$retries" -gt 0 ]; do - local code=0 - bao status -format=json >/dev/null 2>&1 || code=$? - # exit 0 = unsealed, exit 2 = sealed — both mean the server is up and responding - if [ "$code" -eq 0 ] || [ "$code" -eq 2 ]; then - return 0 - fi - retries=$((retries - 1)) - sleep 2 - done - echo "ERROR: OpenBao server did not become ready in time" >&2 - exit 1 -} - -unseal_bao() { - local key1="$1" key2="$2" key3="$3" - if [ -z "$key1" ] || [ -z "$key2" ] || [ -z "$key3" ]; then - echo "ERROR: One or more unseal keys are empty — cannot unseal" >&2 - exit 1 - fi - bao operator unseal "$key1" - bao operator unseal "$key2" - bao operator unseal "$key3" -} - -create_pki_policy() { - cat > /opt/openbao/config/pki-policy.hcl << EOF -path "pki_int/issue/${MG_CERTS_OPENBAO_PKI_ROLE}" { - capabilities = ["create", "update"] -} -path "pki_int/sign/${MG_CERTS_OPENBAO_PKI_ROLE}" { - capabilities = ["create", "update"] -} -path "pki_int/sign-verbatim/${MG_CERTS_OPENBAO_PKI_ROLE}" { - capabilities = ["create", "update"] -} -path "pki_int/certs" { - capabilities = ["list"] -} -path "pki_int/cert/*" { - capabilities = ["read"] -} -path "pki_int/revoke" { - capabilities = ["create", "update"] -} -path "pki_int/ca" { - capabilities = ["read"] -} -path "pki_int/ca_chain" { - capabilities = ["read"] -} -path "pki_int/crl" { - capabilities = ["read"] -} -path "pki/ca" { - capabilities = ["read"] -} -path "pki/ca_chain" { - capabilities = ["read"] -} -path "pki/crl" { - capabilities = ["read"] -} -path "auth/token/renew-self" { - capabilities = ["update"] -} -path "auth/token/lookup-self" { - capabilities = ["read"] -} -path "sys/renew/*" { - capabilities = ["update"] -} -path "auth/approle/role/${MG_CERTS_OPENBAO_PKI_ROLE}/secret-id" { - capabilities = ["create", "update"] -} -path "auth/approle/role/${MG_CERTS_OPENBAO_PKI_ROLE}/secret-id-accessor/lookup" { - capabilities = ["create", "update"] -} -path "auth/approle/role/${MG_CERTS_OPENBAO_PKI_ROLE}/secret-id-accessor/destroy" { - capabilities = ["create", "update"] -} -EOF - bao policy write pki-policy /opt/openbao/config/pki-policy.hcl > /dev/null -} - -# Check if we have pre-configured unseal keys and root token -if [ -n "$MG_CERTS_OPENBAO_UNSEAL_KEY_1" ] && [ -n "$MG_CERTS_OPENBAO_UNSEAL_KEY_2" ] && [ -n "$MG_CERTS_OPENBAO_UNSEAL_KEY_3" ] && [ -n "$MG_CERTS_OPENBAO_ROOT_TOKEN" ]; then - echo "Using pre-configured unseal keys and root token..." - bao server -config=/opt/openbao/config/config.hcl > /opt/openbao/logs/server.log 2>&1 & - BAO_PID=$! - wait_for_bao - - unseal_bao "$MG_CERTS_OPENBAO_UNSEAL_KEY_1" "$MG_CERTS_OPENBAO_UNSEAL_KEY_2" "$MG_CERTS_OPENBAO_UNSEAL_KEY_3" - - export BAO_TOKEN=$MG_CERTS_OPENBAO_ROOT_TOKEN -else - # Initialize OpenBao if not already done - if [ ! -f /opt/openbao/data/init.json ]; then - echo "Initializing OpenBao for the first time..." - bao server -config=/opt/openbao/config/config.hcl > /opt/openbao/logs/server.log 2>&1 & - BAO_PID=$! - wait_for_bao - - # Initialize with 5 key shares and threshold of 3 - bao operator init -key-shares=5 -key-threshold=3 -format=json > /opt/openbao/data/init.json - - # Extract unseal keys and root token - UNSEAL_KEY_1=$(jq -r '.unseal_keys_b64[0]' /opt/openbao/data/init.json) - UNSEAL_KEY_2=$(jq -r '.unseal_keys_b64[1]' /opt/openbao/data/init.json) - UNSEAL_KEY_3=$(jq -r '.unseal_keys_b64[2]' /opt/openbao/data/init.json) - ROOT_TOKEN=$(jq -r '.root_token' /opt/openbao/data/init.json) - - unseal_bao "$UNSEAL_KEY_1" "$UNSEAL_KEY_2" "$UNSEAL_KEY_3" - - export BAO_TOKEN=$ROOT_TOKEN - echo "OpenBao initialized successfully!" - else - echo "OpenBao already initialized, starting server..." - bao server -config=/opt/openbao/config/config.hcl > /opt/openbao/logs/server.log 2>&1 & - BAO_PID=$! - wait_for_bao - - # Check if OpenBao is sealed and unseal if necessary - if bao status -format=json | jq -e '.sealed == true' >/dev/null 2>&1; then - echo "OpenBao is sealed, unsealing..." - UNSEAL_KEY_1=$(jq -r '.unseal_keys_b64[0]' /opt/openbao/data/init.json) - UNSEAL_KEY_2=$(jq -r '.unseal_keys_b64[1]' /opt/openbao/data/init.json) - UNSEAL_KEY_3=$(jq -r '.unseal_keys_b64[2]' /opt/openbao/data/init.json) - - unseal_bao "$UNSEAL_KEY_1" "$UNSEAL_KEY_2" "$UNSEAL_KEY_3" - echo "OpenBao unsealed successfully!" - else - echo "OpenBao is already unsealed!" - fi - - ROOT_TOKEN=$(jq -r '.root_token' /opt/openbao/data/init.json) - export BAO_TOKEN=$ROOT_TOKEN - fi -fi - -# Configure OpenBao PKI and AppRole if not already configured -if [ ! -f /opt/openbao/data/configured ]; then - echo "Configuring OpenBao PKI and AppRole..." - - # Create namespace if specified - if [ -n "$MG_CERTS_OPENBAO_NAMESPACE" ]; then - if bao namespace create "$MG_CERTS_OPENBAO_NAMESPACE" 2>/tmp/ns_error; then - export BAO_NAMESPACE="$MG_CERTS_OPENBAO_NAMESPACE" - echo "$MG_CERTS_OPENBAO_NAMESPACE" > /opt/openbao/data/namespace - echo "Created namespace: $MG_CERTS_OPENBAO_NAMESPACE" - else - if grep -q "namespace already exists" /tmp/ns_error; then - export BAO_NAMESPACE="$MG_CERTS_OPENBAO_NAMESPACE" - echo "$MG_CERTS_OPENBAO_NAMESPACE" > /opt/openbao/data/namespace - echo "Using existing namespace: $MG_CERTS_OPENBAO_NAMESPACE" - else - echo "ERROR: Failed to create namespace $MG_CERTS_OPENBAO_NAMESPACE:" >&2 - cat /tmp/ns_error >&2 - exit 1 - fi - fi - rm -f /tmp/ns_error - fi - - # Enable authentication methods and secrets engines - if ! bao auth enable approle > /tmp/auth_success 2>/tmp/auth_error; then - if ! grep -q "already in use" /tmp/auth_error; then - echo "ERROR: Failed to enable AppRole auth method:" >&2 - cat /tmp/auth_error >&2 - exit 1 - fi - echo "AppRole already enabled" - fi - rm -f /tmp/auth_error /tmp/auth_success - - # Enable PKI secrets engine - if ! bao secrets enable -path=pki pki > /tmp/pki_success 2>/tmp/pki_error; then - # If the failure wasn’t because the mount already exists, abort - if ! grep -q "already in use" /tmp/pki_error; then - echo "ERROR: Failed to enable PKI secrets engine:" >&2 - cat /tmp/pki_error >&2 - exit 1 - fi - echo "PKI already enabled" - fi - rm -f /tmp/pki_error /tmp/pki_success - - # Configure PKI engine - bao secrets tune -max-lease-ttl=87600h pki > /dev/null - - # Validate required CA environment variables - for var in MG_CERTS_OPENBAO_PKI_CA_CN MG_CERTS_OPENBAO_PKI_CA_O MG_CERTS_OPENBAO_PKI_CA_C; do - eval "value=\$var" - if [ -z "$value" ]; then - echo "ERROR: Required environment variable $var is not set" >&2 - exit 1 - fi - done - - PKI_CMD="bao write -field=certificate pki/root/generate/internal \ - common_name=\"$MG_CERTS_OPENBAO_PKI_CA_CN\" \ - organization=\"$MG_CERTS_OPENBAO_PKI_CA_O\" \ - country=\"$MG_CERTS_OPENBAO_PKI_CA_C\" \ - ttl=87600h \ - key_bits=2048 \ - exclude_cn_from_sans=false" - - [ -n "$MG_CERTS_OPENBAO_PKI_CA_OU" ] && PKI_CMD="$PKI_CMD ou=\"$MG_CERTS_OPENBAO_PKI_CA_OU\"" - [ -n "$MG_CERTS_OPENBAO_PKI_CA_L" ] && PKI_CMD="$PKI_CMD locality=\"$MG_CERTS_OPENBAO_PKI_CA_L\"" - [ -n "$MG_CERTS_OPENBAO_PKI_CA_ST" ] && PKI_CMD="$PKI_CMD province=\"$MG_CERTS_OPENBAO_PKI_CA_ST\"" - [ -n "$MG_CERTS_OPENBAO_PKI_CA_ADDR" ] && PKI_CMD="$PKI_CMD street_address=\"$MG_CERTS_OPENBAO_PKI_CA_ADDR\"" - [ -n "$MG_CERTS_OPENBAO_PKI_CA_PO" ] && PKI_CMD="$PKI_CMD postal_code=\"$MG_CERTS_OPENBAO_PKI_CA_PO\"" - - [ -n "$MG_CERTS_OPENBAO_PKI_CA_DNS_NAMES" ] && PKI_CMD="$PKI_CMD alt_names=\"$MG_CERTS_OPENBAO_PKI_CA_DNS_NAMES\"" - [ -n "$MG_CERTS_OPENBAO_PKI_CA_IP_ADDRESSES" ] && PKI_CMD="$PKI_CMD ip_sans=\"$MG_CERTS_OPENBAO_PKI_CA_IP_ADDRESSES\"" - [ -n "$MG_CERTS_OPENBAO_PKI_CA_URI_SANS" ] && PKI_CMD="$PKI_CMD uri_sans=\"$MG_CERTS_OPENBAO_PKI_CA_URI_SANS\"" - [ -n "$MG_CERTS_OPENBAO_PKI_CA_EMAIL_ADDRESSES" ] && PKI_CMD="$PKI_CMD email_sans=\"$MG_CERTS_OPENBAO_PKI_CA_EMAIL_ADDRESSES\"" - - eval $PKI_CMD > /dev/null - - if [ $? -eq 0 ]; then - echo "OpenBao root CA certificate generated successfully!" - else - echo "ERROR: Failed to generate OpenBao root CA certificate" >&2 - exit 1 - fi - - if ! bao secrets enable -path=pki_int pki > /tmp/pki_int_success 2>/tmp/pki_int_error; then - if ! grep -q "already in use" /tmp/pki_int_error; then - echo "ERROR: Failed to enable intermediate PKI secrets engine:" >&2 - cat /tmp/pki_int_error >&2 - exit 1 - fi - echo "Intermediate PKI already enabled" - fi - rm -f /tmp/pki_int_error /tmp/pki_int_success - - bao secrets tune -max-lease-ttl=8760h pki_int > /dev/null - - INTERMEDIATE_CN="${MG_CERTS_OPENBAO_PKI_CA_CN} Intermediate" - INTERMEDIATE_CSR_CMD="bao write -field=csr pki_int/intermediate/generate/internal \ - common_name=\"$INTERMEDIATE_CN\" \ - organization=\"$MG_CERTS_OPENBAO_PKI_CA_O\" \ - country=\"$MG_CERTS_OPENBAO_PKI_CA_C\" \ - ttl=8760h \ - key_bits=2048" - - [ -n "$MG_CERTS_OPENBAO_PKI_CA_OU" ] && INTERMEDIATE_CSR_CMD="$INTERMEDIATE_CSR_CMD ou=\"$MG_CERTS_OPENBAO_PKI_CA_OU\"" - [ -n "$MG_CERTS_OPENBAO_PKI_CA_L" ] && INTERMEDIATE_CSR_CMD="$INTERMEDIATE_CSR_CMD locality=\"$MG_CERTS_OPENBAO_PKI_CA_L\"" - [ -n "$MG_CERTS_OPENBAO_PKI_CA_ST" ] && INTERMEDIATE_CSR_CMD="$INTERMEDIATE_CSR_CMD province=\"$MG_CERTS_OPENBAO_PKI_CA_ST\"" - [ -n "$MG_CERTS_OPENBAO_PKI_CA_ADDR" ] && INTERMEDIATE_CSR_CMD="$INTERMEDIATE_CSR_CMD street_address=\"$MG_CERTS_OPENBAO_PKI_CA_ADDR\"" - [ -n "$MG_CERTS_OPENBAO_PKI_CA_PO" ] && INTERMEDIATE_CSR_CMD="$INTERMEDIATE_CSR_CMD postal_code=\"$MG_CERTS_OPENBAO_PKI_CA_PO\"" - - [ -n "$MG_CERTS_OPENBAO_PKI_CA_DNS_NAMES" ] && INTERMEDIATE_CSR_CMD="$INTERMEDIATE_CSR_CMD alt_names=\"$MG_CERTS_OPENBAO_PKI_CA_DNS_NAMES\"" - [ -n "$MG_CERTS_OPENBAO_PKI_CA_IP_ADDRESSES" ] && INTERMEDIATE_CSR_CMD="$INTERMEDIATE_CSR_CMD ip_sans=\"$MG_CERTS_OPENBAO_PKI_CA_IP_ADDRESSES\"" - [ -n "$MG_CERTS_OPENBAO_PKI_CA_URI_SANS" ] && INTERMEDIATE_CSR_CMD="$INTERMEDIATE_CSR_CMD uri_sans=\"$MG_CERTS_OPENBAO_PKI_CA_URI_SANS\"" - [ -n "$MG_CERTS_OPENBAO_PKI_CA_EMAIL_ADDRESSES" ] && INTERMEDIATE_CSR_CMD="$INTERMEDIATE_CSR_CMD email_sans=\"$MG_CERTS_OPENBAO_PKI_CA_EMAIL_ADDRESSES\"" - - INTERMEDIATE_CSR=$(eval $INTERMEDIATE_CSR_CMD) - - if [ $? -ne 0 ] || [ -z "$INTERMEDIATE_CSR" ]; then - echo "ERROR: Failed to generate intermediate CA CSR" >&2 - exit 1 - fi - - echo "Intermediate CA CSR generated successfully!" - - INTERMEDIATE_CERT=$(bao write -field=certificate pki/root/sign-intermediate \ - csr="$INTERMEDIATE_CSR" \ - format=pem_bundle \ - ttl=8760h \ - use_csr_values=true) - - if [ $? -ne 0 ] || [ -z "$INTERMEDIATE_CERT" ]; then - echo "ERROR: Failed to sign intermediate CA certificate" >&2 - exit 1 - fi - - echo "Intermediate CA certificate signed successfully!" - - bao write pki/config/urls \ - issuing_certificates='http://127.0.0.1:8200/v1/pki/ca' \ - crl_distribution_points='http://127.0.0.1:8200/v1/pki/crl' \ - ocsp_servers='http://127.0.0.1:8200/v1/pki/ocsp' > /dev/null - - bao write pki_int/config/urls \ - issuing_certificates='http://127.0.0.1:8200/v1/pki_int/ca' \ - crl_distribution_points='http://127.0.0.1:8200/v1/pki_int/crl' \ - ocsp_servers='http://127.0.0.1:8200/v1/pki_int/ocsp' > /dev/null - - bao write pki_int/intermediate/set-signed certificate="$INTERMEDIATE_CERT" > /dev/null - - if [ $? -eq 0 ]; then - echo "Intermediate CA setup completed successfully!" - else - echo "ERROR: Failed to set signed intermediate certificate" >&2 - exit 1 - fi - - echo "$INTERMEDIATE_CERT" > /opt/openbao/data/intermediate_ca.pem - - ROLE_CMD="bao write pki_int/roles/${MG_CERTS_OPENBAO_PKI_ROLE} \ - allow_any_name=true \ - enforce_hostnames=false \ - allow_ip_sans=true \ - allow_localhost=true \ - allow_bare_domains=true \ - allow_subdomains=true \ - allow_glob_domains=true \ - allowed_domains=\"*\" \ - allowed_uri_sans=\"*\" \ - allowed_other_sans=\"*\" \ - server_flag=true \ - client_flag=true \ - code_signing_flag=false \ - email_protection_flag=false \ - key_type=rsa \ - key_bits=2048 \ - key_usage=\"DigitalSignature,KeyEncipherment,KeyAgreement\" \ - ext_key_usage=\"ServerAuth,ClientAuth,OCSPSigning\" \ - use_csr_common_name=true \ - use_csr_sans=true \ - basic_constraints_valid_for_non_ca=true \ - max_ttl=720h \ - ttl=720h" - - eval "$ROLE_CMD" > /dev/null - - create_pki_policy - - # Create AppRole - SECRET_ID_TTL="${MG_CERTS_OPENBAO_SECRET_ID_TTL}" - bao write auth/approle/role/"${MG_CERTS_OPENBAO_PKI_ROLE}" \ - token_policies=pki-policy \ - token_ttl=1h \ - token_max_ttl=4h \ - bind_secret_id=true \ - secret_id_ttl="$SECRET_ID_TTL" > /dev/null - - # Set custom role ID if provided - if [ -n "$MG_CERTS_OPENBAO_APP_ROLE" ]; then - bao write auth/approle/role/"${MG_CERTS_OPENBAO_PKI_ROLE}"/role-id role_id="$MG_CERTS_OPENBAO_APP_ROLE" > /dev/null - fi - - # Set custom secret ID if provided, otherwise generate one - if [ -n "$MG_CERTS_OPENBAO_APP_SECRET" ]; then - bao write auth/approle/role/"${MG_CERTS_OPENBAO_PKI_ROLE}"/custom-secret-id secret_id="$MG_CERTS_OPENBAO_APP_SECRET" > /dev/null - echo "$MG_CERTS_OPENBAO_APP_SECRET" > /opt/openbao/data/secret_id - else - GENERATED_SECRET_ID=$(bao write -field=secret_id -force auth/approle/role/"${MG_CERTS_OPENBAO_PKI_ROLE}"/secret-id) - echo "$GENERATED_SECRET_ID" > /opt/openbao/data/secret_id - fi - - # Generate service token for additional access - SERVICE_TOKEN=$(bao write -field=token auth/token/create \ - policies=pki-policy \ - ttl=24h \ - renewable=true \ - display_name="certs-service" 2>/dev/null) - - echo "SERVICE_TOKEN=$SERVICE_TOKEN" > /opt/openbao/data/service_token - - # Mark configuration as complete - touch /opt/openbao/data/configured - echo "OpenBao configuration completed successfully!" -else - echo "OpenBao already configured, verifying and updating configuration..." - - # Restore namespace if it exists - if [ -f /opt/openbao/data/namespace ] && [ -n "$MG_CERTS_OPENBAO_NAMESPACE" ]; then - SAVED_NAMESPACE=$(cat /opt/openbao/data/namespace) - if [ "$SAVED_NAMESPACE" = "$MG_CERTS_OPENBAO_NAMESPACE" ]; then - export BAO_NAMESPACE="$MG_CERTS_OPENBAO_NAMESPACE" - fi - fi - - # Check if AppRole role exists, create if missing - if ! bao read auth/approle/role/"${MG_CERTS_OPENBAO_PKI_ROLE}" > /dev/null 2>&1; then - if ! bao auth enable approle > /tmp/auth_success 2>/tmp/auth_error; then - if ! grep -q "already in use" /tmp/auth_error; then - echo "ERROR: Failed to enable AppRole auth method:" >&2 - cat /tmp/auth_error >&2 - exit 1 - fi - fi - rm -f /tmp/auth_error /tmp/auth_success - - create_pki_policy - - SECRET_ID_TTL="${MG_CERTS_OPENBAO_SECRET_ID_TTL}" - bao write auth/approle/role/"${MG_CERTS_OPENBAO_PKI_ROLE}" \ - token_policies=pki-policy \ - token_ttl=1h \ - token_max_ttl=4h \ - bind_secret_id=true \ - secret_id_ttl="$SECRET_ID_TTL" > /dev/null - - if [ -n "$MG_CERTS_OPENBAO_APP_ROLE" ]; then - bao write auth/approle/role/"${MG_CERTS_OPENBAO_PKI_ROLE}"/role-id role_id="$MG_CERTS_OPENBAO_APP_ROLE" > /dev/null - fi - fi - - SECRET_ID_VALID=false - if [ -n "$MG_CERTS_OPENBAO_APP_SECRET" ]; then - if bao write -field=client_token auth/approle/login role_id="$MG_CERTS_OPENBAO_APP_ROLE" secret_id="$MG_CERTS_OPENBAO_APP_SECRET" > /dev/null 2>&1; then - SECRET_ID_VALID=true - echo "$MG_CERTS_OPENBAO_APP_SECRET" > /opt/openbao/data/secret_id - fi - elif [ -f /opt/openbao/data/secret_id ]; then - STORED_SECRET_ID=$(cat /opt/openbao/data/secret_id) - if [ -n "$STORED_SECRET_ID" ]; then - ROLE_ID=$(bao read -field=role_id auth/approle/role/"${MG_CERTS_OPENBAO_PKI_ROLE}"/role-id) - if bao write -field=client_token auth/approle/login role_id="$ROLE_ID" secret_id="$STORED_SECRET_ID" > /dev/null 2>&1; then - SECRET_ID_VALID=true - fi - fi - fi - - if [ "$SECRET_ID_VALID" = "false" ]; then - NEW_SECRET_ID=$(bao write -field=secret_id -force auth/approle/role/"${MG_CERTS_OPENBAO_PKI_ROLE}"/secret-id) - - if [ -z "$NEW_SECRET_ID" ]; then - echo "ERROR: Failed to generate new secret ID" >&2 - else - echo "$NEW_SECRET_ID" > /opt/openbao/data/secret_id - echo "Generated new secret ID for certs service" - fi - fi - - # Regenerate service token - SERVICE_TOKEN=$(bao write -field=token auth/token/create \ - policies=pki-policy \ - ttl=24h \ - renewable=true \ - display_name="certs-service" 2>/dev/null) - - if [ -n "$SERVICE_TOKEN" ]; then - echo "SERVICE_TOKEN=$SERVICE_TOKEN" > /opt/openbao/data/service_token - fi -fi - -echo "================================" -echo "OpenBao Production Setup Complete" -echo "================================" -echo "OpenBao Address: http://localhost:8200" -echo "UI Available at: http://localhost:8200/ui" -echo "================================" -echo "IMPORTANT: Store the init.json file securely!" -echo "It contains unseal keys and root token!" -echo "================================" - -echo "OpenBao is ready for certs service on port 8200" - -if [ -n "$BAO_PID" ]; then - wait $BAO_PID -else - echo "ERROR: OpenBao server process ID not available" >&2 - exit 1 -fi diff --git a/domains/README.md b/domains/README.md deleted file mode 100644 index cb114ad3f..000000000 --- a/domains/README.md +++ /dev/null @@ -1,426 +0,0 @@ -# Domains - -The Domains service provides an HTTP API for managing platform domains in Magistrala. Through this API you can create, list, retrieve, update, enable/disable/freeze domains, manage roles & invitations associated with domains, and more. - -For more background on Magistrala concepts, see the [official documentation][doc]. - -## Configuration - -The service is configured through environment variables (unset variables fall back to defaults). - -| Variable | Description | Default | -| ------------------------------------- | ------------------------------------------------------------------------- | ---------------------------------- | -| `MG_DOMAINS_LOG_LEVEL` | Log level for Domains (debug, info, warn, error) | debug | -| `MG_DOMAINS_HTTP_HOST` | Domains service HTTP host | domains | -| `MG_DOMAINS_HTTP_PORT` | Domains service HTTP port | 9003 | -| `MG_DOMAINS_HTTP_SERVER_CERT` | Path to PEM-encoded HTTP server certificate | "" | -| `MG_DOMAINS_HTTP_SERVER_KEY` | Path to PEM-encoded HTTP server key | "" | -| `MG_DOMAINS_GRPC_PORT` | Domains service gRPC port | 7003 | -| `MG_DOMAINS_GRPC_SERVER_CERT` | Path to PEM-encoded gRPC server certificate | "" | -| `MG_DOMAINS_GRPC_SERVER_KEY` | Path to PEM-encoded gRPC server key | "" | -| `MG_DOMAINS_GRPC_SERVER_CA_CERTS` | Path to trusted CA bundle for the gRPC server | "" | -| `MG_DOMAINS_GRPC_CLIENT_CA_CERTS` | Path to client CA bundle to require gRPC mTLS | "" | -| `MG_DOMAINS_DB_HOST` | Database host address | domains-db | -| `MG_DOMAINS_DB_PORT` | Database host port | 5432 | -| `MG_DOMAINS_DB_USER` | Database user | magistrala | -| `MG_DOMAINS_DB_PASS` | Database password | magistrala | -| `MG_DOMAINS_DB_NAME` | Name of the database used by the service | domains | -| `MG_DOMAINS_DB_SSL_MODE` | Database connection SSL mode (disable, require, verify-ca, verify-full) | "" | -| `MG_DOMAINS_DB_SSL_CERT` | Path to the PEM-encoded certificate file | "" | -| `MG_DOMAINS_DB_SSL_KEY` | Path to the PEM-encoded key file | "" | -| `MG_DOMAINS_DB_SSL_ROOT_CERT` | Path to the PEM-encoded root certificate file | "" | -| `MG_DOMAINS_CACHE_URL` | Cache database URL | redis://domains-redis:6379/0 | -| `MG_DOMAINS_CACHE_KEY_DURATION` | Cache key duration for domain status/route lookups | 10m | -| `MG_DOMAINS_INSTANCE_ID` | Domains instance ID (auto-generated when empty) | "" | -| `MG_SPICEDB_HOST` | SpiceDB host for policy checks | magistrala-spicedb | -| `MG_SPICEDB_PORT` | SpiceDB port | 50051 | -| `MG_SPICEDB_SCHEMA_FILE` | Path to SpiceDB schema file used to seed available actions | ./docker/spicedb/schema.schema.zed | -| `MG_SPICEDB_PRE_SHARED_KEY` | SpiceDB preshared key | 12345678 | -| `MG_ES_URL` | Event store URL | nats://localhost:4222 | -| `MG_JAEGER_URL` | Jaeger server URL | | -| `MG_JAEGER_TRACE_RATIO` | Trace sampling ratio | 1.0 | -| `MG_SEND_TELEMETRY` | Send telemetry to the Magistrala call-home server | true | -| `MG_AUTH_GRPC_URL` | Auth service gRPC URL | "" | -| `MG_AUTH_GRPC_TIMEOUT` | Auth service gRPC request timeout | 1s | -| `MG_AUTH_GRPC_CLIENT_CERT` | Path to the PEM-encoded Auth gRPC client certificate | "" | -| `MG_AUTH_GRPC_CLIENT_KEY` | Path to the PEM-encoded Auth gRPC client key | "" | -| `MG_AUTH_GRPC_SERVER_CA_CERTS` | Path to the PEM-encoded Auth gRPC trusted CA bundle | "" | -| `MG_DOMAINS_CALLOUT_URLS` | Comma-separated list of HTTP callout targets invoked on domain operations | "" | -| `MG_DOMAINS_CALLOUT_METHOD` | HTTP method for callouts (POST or GET) | POST | -| `MG_DOMAINS_CALLOUT_TLS_VERIFICATION` | Verify TLS certificates for callouts | true | -| `MG_DOMAINS_CALLOUT_TIMEOUT` | Callout request timeout | 10s | -| `MG_DOMAINS_CALLOUT_KEY` | Client key for mTLS callouts | "" | -| `MG_DOMAINS_CALLOUT_OPERATIONS` | Comma-separated list of operation names that should trigger callouts | "" | - -**Note**: Set `MG_DOMAINS_CALLOUT_OPERATIONS` to a subset of `OpCreateDomain`, `OpRetrieveDomain`, `OpUpdateDomain`, `OpEnableDomain`, `OpDisableDomain`, `OpFreezeDomain`, `OpListDomains`, `OpViewDomainInvitation`, `OpSendInvitation`, `OpAcceptInvitation`, `OpListInvitations`, `OpListDomainInvitations`, `OpRejectInvitation`, or `OpDeleteInvitation` to filter which actions produce callouts. - -## Deployment - -The service is distributed as a Docker container. See the [`domains` section](https://github.com/absmach/magistrala/blob/main/docker/docker-compose.yaml#L215-L310) of the compose file for an example deployment. - -To run the service outside of a container: - -```bash -# download the latest version of the service -git clone https://github.com/absmach/magistrala -cd magistrala - -# compile the domains service -make domains - -# copy binary to $GOBIN -make install - -# set the environment variables and run the service -MG_DOMAINS_LOG_LEVEL=debug \ -MG_DOMAINS_CACHE_URL=redis://domains-redis:6379/0 \ -MG_DOMAINS_CACHE_KEY_DURATION=10m \ -MG_DOMAINS_HTTP_HOST=domains \ -MG_DOMAINS_HTTP_PORT=9003 \ -MG_DOMAINS_HTTP_SERVER_CERT="" \ -MG_DOMAINS_HTTP_SERVER_KEY="" \ -MG_DOMAINS_GRPC_HOST=domains \ -MG_DOMAINS_GRPC_PORT=7003 \ -MG_DOMAINS_GRPC_SERVER_CERT="" \ -MG_DOMAINS_GRPC_SERVER_KEY="" \ -MG_DOMAINS_GRPC_SERVER_CA_CERTS="" \ -MG_DOMAINS_GRPC_CLIENT_CA_CERTS="" \ -MG_DOMAINS_DB_HOST=domains-db \ -MG_DOMAINS_DB_PORT=5432 \ -MG_DOMAINS_DB_USER=magistrala \MG_DOMAINS_DB_PASS=magistrala \MG_DOMAINS_DB_NAME=domains \ -MG_DOMAINS_DB_SSL_MODE="" \ -MG_DOMAINS_DB_SSL_CERT="" \ -MG_DOMAINS_DB_SSL_KEY="" \ -MG_DOMAINS_DB_SSL_ROOT_CERT="" \ -MG_AUTH_GRPC_URL="" \ -MG_AUTH_GRPC_TIMEOUT=1s \ -MG_AUTH_GRPC_CLIENT_CERT="" \ -MG_AUTH_GRPC_CLIENT_KEY="" \ -MG_AUTH_GRPC_SERVER_CA_CERTS="" \ -MG_SPICEDB_HOST=localhost \ -MG_SPICEDB_PORT=50051 \ -MG_SPICEDB_SCHEMA_FILE=./docker/spicedb/schema.schema.zed \ -MG_SPICEDB_PRE_SHARED_KEY=12345678 \ -MG_ES_URL=nats://localhost:4222 \ -MG_JAEGER_URL= \ -MG_JAEGER_TRACE_RATIO=1.0 \ -MG_DOMAINS_CALLOUT_URLS="" \ -MG_DOMAINS_CALLOUT_METHOD=POST \ -MG_DOMAINS_CALLOUT_TLS_VERIFICATION=true \ -MG_DOMAINS_CALLOUT_TIMEOUT=10s \ -MG_DOMAINS_CALLOUT_KEY="" \ -MG_DOMAINS_CALLOUT_OPERATIONS="" \ -MG_SEND_TELEMETRY=true \ -MG_DOMAINS_INSTANCE_ID="" \ -$GOBIN/magistrala-domains -``` - -## Usage - -Domains supports the following operations: - -| Operation | Description | -| ------------------- | ------------------------------------------------------------------------------------- | -| `create` | Create a new domain with a unique route | -| `get` | Retrieve a domain (optionally with role memberships) or list accessible domains | -| `update` | Update a domain’s name, tags, or metadata | -| `enable` | Enable a previously disabled domain | -| `disable` | Disable an active domain | -| `freeze` | Freeze a domain (platform administrators only) | -| `invite` | Send an invitation for a user to join a domain with a specific role | -| `invitations` | List invitations for the current user or for a specific domain | -| `accept/reject` | Accept or reject a pending domain invitation | -| `delete-invitation` | Delete an invitation (inviter, invitee, or admin) | -| `roles` | Create/list/update/delete domain roles; manage role actions and members; list actions | - -### API Examples - -#### Create a Domain - -```bash -curl -X POST http://localhost:9004/domains \ - -H "Authorization: Bearer " \ - -H "Content-Type: application/json" \ - -d '{ - "name": "Edge Tenant", - "route": "edge", - "tags": ["iot", "prod"], - "metadata": { "region": "eu-west-1" } - }' -``` - -Expected response: - -```json -{ - "id": "f2b16e2c-5ad1-4c44-9c1b-0f862ec1c0c8", - "name": "Edge Tenant", - "tags": ["iot", "prod"], - "route": "edge", - "metadata": { "region": "eu-west-1" }, - "status": "enabled", - "created_by": "a5b6c7d8-e901-4fab-9bcd-123456789abc", - "created_at": "2024-10-24T13:31:52Z" -} -``` - -#### List Domains - -```bash -curl -X GET "http://localhost:9004/domains?limit=10&status=enabled" \ - -H "Authorization: Bearer " -``` - -```json -{ - "total": 2, - "offset": 0, - "limit": 10, - "domains": [ - { - "id": "f2b16e2c-5ad1-4c44-9c1b-0f862ec1c0c8", - "name": "Edge Tenant", - "route": "edge", - "status": "enabled", - "created_at": "2024-10-24T13:31:52Z" - }, - { - "id": "7f6a5b4c-3210-4fed-ba98-76543210fedc", - "name": "Sandbox", - "route": "sandbox", - "status": "disabled", - "created_at": "2024-10-10T08:12:04Z" - } - ] -} -``` - -#### Retrieve a Domain (with Roles) - -```bash -curl -X GET "http://localhost:9004/domains/?roles=true" \ - -H "Authorization: Bearer " -``` - -```json -{ - "id": "f2b16e2c-5ad1-4c44-9c1b-0f862ec1c0c8", - "name": "Edge Tenant", - "route": "edge", - "status": "enabled", - "roles": [ - { - "role_id": "b83d25e7-6a49-4c2e-98c5-9323d4d9af7d", - "role_name": "admin", - "actions": ["manage_role_permission", "update_permission", "add_role_users_permission"] - } - ], - "created_at": "2024-10-24T13:31:52Z", - "updated_at": "2024-10-24T13:31:52Z" -} -``` - -#### Update a Domain - -```bash -curl -X PATCH http://localhost:9004/domains/ \ - -H "Authorization: Bearer " \ - -H "Content-Type: application/json" \ - -d '{ - "name": "Edge Operations", - "tags": ["iot", "ops"], - "metadata": { "region": "eu-west-1", "env": "prod" } - }' -``` - -#### Enable, Disable, or Freeze a Domain - -```bash -curl -X POST http://localhost:9004/domains//disable \ - -H "Authorization: Bearer " - -curl -X POST http://localhost:9004/domains//enable \ - -H "Authorization: Bearer " - -curl -X POST http://localhost:9004/domains//freeze \ - -H "Authorization: Bearer " -``` - -#### Send an Invitation - -```bash -curl -X POST http://localhost:9004/domains//invitations \ - -H "Authorization: Bearer " \ - -H "Content-Type: application/json" \ - -d '{ - "invitee_user_id": "", - "role_id": "", - "resend": false - }' -``` - -#### List Domain or User Invitations - -```bash -# For a specific domain -curl -X GET "http://localhost:9004/domains//invitations?limit=10&state=pending" \ - -H "Authorization: Bearer " - -# For the current user -curl -X GET "http://localhost:9004/invitations?limit=10" \ - -H "Authorization: Bearer " -``` - -```json -{ - "total": 1, - "offset": 0, - "limit": 10, - "invitations": [ - { - "invited_by": "a5b6c7d8-e901-4fab-9bcd-123456789abc", - "invitee_user_id": "2c4d6e8f-0a12-4b3c-9d8e-7f6a5b4c3d2e", - "domain_id": "f2b16e2c-5ad1-4c44-9c1b-0f862ec1c0c8", - "domain_name": "Edge Tenant", - "role_id": "b83d25e7-6a49-4c2e-98c5-9323d4d9af7d", - "role_name": "admin", - "actions": ["manage_role_permission", "update_permission"], - "created_at": "2024-10-25T11:03:42Z" - } - ] -} -``` - -#### Accept, Reject, or Delete an Invitation - -```bash -# Accept -curl -X POST http://localhost:9004/invitations/accept \ - -H "Authorization: Bearer " \ - -H "Content-Type: application/json" \ - -d '{ "domain_id": "" }' - -# Reject -curl -X POST http://localhost:9004/invitations/reject \ - -H "Authorization: Bearer " \ - -H "Content-Type: application/json" \ - -d '{ "domain_id": "" }' - -# Delete -curl -X DELETE http://localhost:9004/domains//invitations \ - -H "Authorization: Bearer " \ - -H "Content-Type: application/json" \ - -d '{ "user_id": "" }' -``` - -## Roles Management for Domains - -Domain roles reuse the shared role manager. Supported operations: - -| Operation | Description | -| ------------------------- | -------------------------------------------------------------------- | -| `create-role` | Create a new role for a domain | -| `list-roles` | List all roles assigned to a domain | -| `get-role` | Retrieve details for a specific domain role | -| `update-role` | Update a domain role name | -| `delete-role` | Delete a domain role | -| `add-role-action` | Add one or more actions to a domain role | -| `list-role-actions` | List all actions associated with a domain role | -| `delete-role-action` | Remove a specific action from a domain role | -| `delete-all-role-actions` | Remove all actions from a domain role | -| `add-role-member` | Associate one or more members with a domain role | -| `list-role-members` | List all members of a domain role | -| `delete-role-member` | Remove one or more members from a domain role | -| `delete-all-role-members` | Remove all members from a domain role | -| `list-available-actions` | Retrieve the global list of available domain actions from the schema | - -Example: create a domain role - -```bash -curl -X POST http://localhost:9004/domains//roles \ - -H "Authorization: Bearer " \ - -H "Content-Type: application/json" \ - -d '{ - "role_name": "domain-editor", - "optional_actions": ["update_permission", "read_permission"], - "optional_members": [""] - }' -``` - -To discover allowed actions before creating roles, call: - -```bash -curl -X GET http://localhost:9004/domains/roles/available-actions \ - -H "Authorization: Bearer " -``` - -## Implementation Details - -- Domains and invitations are persisted in PostgreSQL; migrations also create role tables with a `domains_` prefix. -- Redis caches domain status and route-to-ID lookups to speed up authorization. -- Domain lifecycle events are published to the configured event store (`MG_ES_URL`). -- Authorization and role checks are enforced via SpiceDB-backed policy service. -- Optional HTTP callouts can be triggered before operations, using the `MG_DOMAINS_CALLOUT_*` settings. -- Observability: Jaeger tracing, Prometheus metrics at `/metrics`, and a `/health` endpoint. - -### Domains Table - -| Column | Type | Description | -| ------------ | ------------ | --------------------------------------------------- | -| `id` | VARCHAR(36) | UUID of the domain (primary key) | -| `name` | VARCHAR(254) | Human-readable domain name | -| `tags` | TEXT[] | Domain tags | -| `metadata` | JSONB | Arbitrary metadata | -| `route` | VARCHAR(254) | Unique domain route/alias | -| `created_at` | TIMESTAMPTZ | Creation timestamp | -| `updated_at` | TIMESTAMPTZ | Last update timestamp | -| `updated_by` | VARCHAR(254) | Actor who last updated the domain | -| `created_by` | VARCHAR(254) | Actor who created the domain | -| `status` | SMALLINT | 0 = enabled, 1 = disabled, 2 = freezed, 3 = deleted | - -### Invitations Table - -| Column | Type | Description | -| ----------------- | ----------- | ------------------------------------------------ | -| `invited_by` | VARCHAR(36) | User who sent the invitation | -| `invitee_user_id` | VARCHAR(36) | User being invited | -| `domain_id` | VARCHAR(36) | Domain to join (FK to `domains.id`) | -| `role_id` | VARCHAR(36) | Role to grant on acceptance | -| `created_at` | TIMESTAMPTZ | Invitation creation time | -| `updated_at` | TIMESTAMPTZ | Last modification time | -| `confirmed_at` | TIMESTAMPTZ | When the invitation was accepted (if applicable) | -| `rejected_at` | TIMESTAMPTZ | When the invitation was rejected (if applicable) | - -## Best Practices - -- Reserve concise, DNS-friendly `route` values for external-facing domains. -- Use metadata and tags to capture environment, region, and ownership for filtering. -- Prefer `disable` over delete when you need reversible off-boarding; use `freeze` for emergency locks by admins. -- Keep role definitions minimal; grant only the actions needed and audit with `list-role-members`. -- Clean up stale invitations regularly using the domain/user invitation listing endpoints. -- When enabling callouts, narrow `MG_DOMAINS_CALLOUT_OPERATIONS` to the events you must observe. - -## Versioning and Health Check - -The Domains service exposes `/health` with status and build metadata. - -```bash -curl -X GET http://localhost:9004/health \ - -H "accept: application/health+json" -``` - -Example response: - -```json -{ - "status": "pass", - "version": "0.18.0", - "commit": "7d6f4dc4f7f0c1fa3dc24eddfb18bb5073ff4f62", - "description": "domains service", - "build_time": "1970-01-01_00:00:00" -} -``` - -For full API coverage, see the [Domains API documentation](https://docs.api.magistrala.absmach.eu/?urls.primaryName=api%2Fdomains.yaml). - -[doc]: https://magistrala.absmach.eu/docs/ \ No newline at end of file diff --git a/domains/api/grpc/client.go b/domains/api/grpc/client.go deleted file mode 100644 index 16ec231d3..000000000 --- a/domains/api/grpc/client.go +++ /dev/null @@ -1,147 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package grpc - -import ( - "context" - "time" - - grpcCommonV1 "github.com/absmach/magistrala/api/grpc/common/v1" - grpcDomainsV1 "github.com/absmach/magistrala/api/grpc/domains/v1" - grpcapi "github.com/absmach/magistrala/auth/api/grpc" - "github.com/go-kit/kit/endpoint" - kitgrpc "github.com/go-kit/kit/transport/grpc" - "google.golang.org/grpc" -) - -const domainsSvcName = "domains.v1.DomainsService" - -var _ grpcDomainsV1.DomainsServiceClient = (*domainsGrpcClient)(nil) - -type domainsGrpcClient struct { - deleteUserFromDomains endpoint.Endpoint - retrieveStatus endpoint.Endpoint - retrieveIDByRoute endpoint.Endpoint - timeout time.Duration -} - -// NewDomainsClient returns new domains gRPC client instance. -func NewDomainsClient(conn *grpc.ClientConn, timeout time.Duration) grpcDomainsV1.DomainsServiceClient { - return &domainsGrpcClient{ - deleteUserFromDomains: kitgrpc.NewClient( - conn, - domainsSvcName, - "DeleteUserFromDomains", - encodeDeleteUserRequest, - decodeDeleteUserResponse, - grpcDomainsV1.DeleteUserRes{}, - ).Endpoint(), - retrieveStatus: kitgrpc.NewClient( - conn, - domainsSvcName, - "RetrieveStatus", - encodeRetrieveStatusRequest, - decodeRetrieveStatusResponse, - grpcCommonV1.RetrieveEntityRes{}, - ).Endpoint(), - retrieveIDByRoute: kitgrpc.NewClient( - conn, - domainsSvcName, - "RetrieveIDByRoute", - encodeRetrieveIDByRouteRequest, - decodeRetrieveIDByRouteResponse, - grpcCommonV1.RetrieveEntityRes{}, - ).Endpoint(), - timeout: timeout, - } -} - -func (client domainsGrpcClient) DeleteUserFromDomains(ctx context.Context, in *grpcDomainsV1.DeleteUserReq, opts ...grpc.CallOption) (*grpcDomainsV1.DeleteUserRes, error) { - ctx, cancel := context.WithTimeout(ctx, client.timeout) - defer cancel() - - res, err := client.deleteUserFromDomains(ctx, deleteUserPoliciesReq{ - ID: in.GetId(), - }) - if err != nil { - return &grpcDomainsV1.DeleteUserRes{}, grpcapi.DecodeError(err) - } - - dpr := res.(deleteUserRes) - return &grpcDomainsV1.DeleteUserRes{Deleted: dpr.deleted}, nil -} - -func decodeDeleteUserResponse(_ context.Context, grpcRes any) (any, error) { - res := grpcRes.(*grpcDomainsV1.DeleteUserRes) - return deleteUserRes{deleted: res.GetDeleted()}, nil -} - -func encodeDeleteUserRequest(_ context.Context, grpcReq any) (any, error) { - req := grpcReq.(deleteUserPoliciesReq) - return &grpcDomainsV1.DeleteUserReq{ - Id: req.ID, - }, nil -} - -func (client domainsGrpcClient) RetrieveStatus(ctx context.Context, in *grpcCommonV1.RetrieveEntityReq, opts ...grpc.CallOption) (*grpcCommonV1.RetrieveEntityRes, error) { - ctx, cancel := context.WithTimeout(ctx, client.timeout) - defer cancel() - - res, err := client.retrieveStatus(ctx, retrieveStatusReq{ - ID: in.GetId(), - }) - if err != nil { - return &grpcCommonV1.RetrieveEntityRes{}, grpcapi.DecodeError(err) - } - - rdsr := res.(retrieveStatusRes) - return &grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Status: uint32(rdsr.status), - }, - }, nil -} - -func decodeRetrieveStatusResponse(_ context.Context, grpcRes any) (any, error) { - res := grpcRes.(*grpcCommonV1.RetrieveEntityRes) - return retrieveStatusRes{status: uint8(res.Entity.GetStatus())}, nil -} - -func encodeRetrieveStatusRequest(_ context.Context, grpcReq any) (any, error) { - req := grpcReq.(retrieveStatusReq) - return &grpcCommonV1.RetrieveEntityReq{ - Id: req.ID, - }, nil -} - -func (client domainsGrpcClient) RetrieveIDByRoute(ctx context.Context, in *grpcCommonV1.RetrieveIDByRouteReq, opts ...grpc.CallOption) (*grpcCommonV1.RetrieveEntityRes, error) { - ctx, cancel := context.WithTimeout(ctx, client.timeout) - defer cancel() - - res, err := client.retrieveIDByRoute(ctx, retrieveIDByRouteReq{ - Route: in.GetRoute(), - }) - if err != nil { - return &grpcCommonV1.RetrieveEntityRes{}, grpcapi.DecodeError(err) - } - - rbr := res.(retrieveIDByRouteRes) - return &grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Id: rbr.id, - }, - }, nil -} - -func decodeRetrieveIDByRouteResponse(_ context.Context, grpcRes any) (any, error) { - res := grpcRes.(*grpcCommonV1.RetrieveEntityRes) - return retrieveIDByRouteRes{id: res.Entity.GetId()}, nil -} - -func encodeRetrieveIDByRouteRequest(_ context.Context, grpcReq any) (any, error) { - req := grpcReq.(retrieveIDByRouteReq) - return &grpcCommonV1.RetrieveIDByRouteReq{ - Route: req.Route, - }, nil -} diff --git a/domains/api/grpc/doc.go b/domains/api/grpc/doc.go deleted file mode 100644 index ecbe29d55..000000000 --- a/domains/api/grpc/doc.go +++ /dev/null @@ -1,5 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package grpc contains implementation of Domains service gRPC API. -package grpc diff --git a/domains/api/grpc/endpoint.go b/domains/api/grpc/endpoint.go deleted file mode 100644 index dcf7c7717..000000000 --- a/domains/api/grpc/endpoint.go +++ /dev/null @@ -1,62 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package grpc - -import ( - "context" - - domains "github.com/absmach/magistrala/domains/private" - "github.com/go-kit/kit/endpoint" -) - -func deleteUserFromDomainsEndpoint(svc domains.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(deleteUserPoliciesReq) - if err := req.validate(); err != nil { - return deleteUserRes{}, err - } - - if err := svc.DeleteUserFromDomains(ctx, req.ID); err != nil { - return deleteUserRes{}, err - } - - return deleteUserRes{deleted: true}, nil - } -} - -func retrieveStatusEndpoint(svc domains.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(retrieveStatusReq) - if err := req.validate(); err != nil { - return retrieveStatusRes{}, err - } - - status, err := svc.RetrieveStatus(ctx, req.ID) - if err != nil { - return retrieveStatusRes{}, err - } - - return retrieveStatusRes{ - status: uint8(status), - }, nil - } -} - -func retrieveIDByRouteEndpoint(svc domains.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(retrieveIDByRouteReq) - if err := req.validate(); err != nil { - return retrieveIDByRouteRes{}, err - } - - id, err := svc.RetrieveIDByRoute(ctx, req.Route) - if err != nil { - return retrieveIDByRouteRes{}, err - } - - return retrieveIDByRouteRes{ - id: id, - }, nil - } -} diff --git a/domains/api/grpc/endpoint_test.go b/domains/api/grpc/endpoint_test.go deleted file mode 100644 index 2a177fa0a..000000000 --- a/domains/api/grpc/endpoint_test.go +++ /dev/null @@ -1,220 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package grpc_test - -import ( - "context" - "fmt" - "net" - "testing" - "time" - - grpcCommonV1 "github.com/absmach/magistrala/api/grpc/common/v1" - grpcDomainsV1 "github.com/absmach/magistrala/api/grpc/domains/v1" - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/domains" - grpcapi "github.com/absmach/magistrala/domains/api/grpc" - pDomains "github.com/absmach/magistrala/domains/private" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" - "google.golang.org/grpc" - "google.golang.org/grpc/credentials/insecure" -) - -const ( - port = 8081 - secret = "secret" - email = "test@example.com" - id = "testID" - clientsType = "clients" - usersType = "users" - description = "Description" - groupName = "smqx" - adminPermission = "admin" - authoritiesObj = "authorities" - memberRelation = "member" - loginDuration = 30 * time.Minute - refreshDuration = 24 * time.Hour - invalidDuration = 7 * 24 * time.Hour - validToken = "valid" - inValidToken = "invalid" - validPolicy = "valid" -) - -var authAddr = fmt.Sprintf("localhost:%d", port) - -func startGRPCServer(svc pDomains.Service, port int) *grpc.Server { - listener, _ := net.Listen("tcp", fmt.Sprintf(":%d", port)) - server := grpc.NewServer() - grpcDomainsV1.RegisterDomainsServiceServer(server, grpcapi.NewDomainsServer(svc)) - go func() { - err := server.Serve(listener) - assert.Nil(&testing.T{}, err, fmt.Sprintf(`"Unexpected error creating auth server %s"`, err)) - }() - - return server -} - -func TestDeleteUserFromDomains(t *testing.T) { - conn, err := grpc.NewClient(authAddr, grpc.WithTransportCredentials(insecure.NewCredentials())) - assert.Nil(t, err, fmt.Sprintf("Unexpected error creating client connection %s", err)) - grpcClient := grpcapi.NewDomainsClient(conn, time.Second) - - cases := []struct { - desc string - token string - deleteUserReq *grpcDomainsV1.DeleteUserReq - deleteUserRes *grpcDomainsV1.DeleteUserRes - err error - }{ - { - desc: "delete valid req", - token: validToken, - deleteUserReq: &grpcDomainsV1.DeleteUserReq{ - Id: id, - }, - deleteUserRes: &grpcDomainsV1.DeleteUserRes{Deleted: true}, - err: nil, - }, - { - desc: "delete invalid req with invalid token", - token: inValidToken, - deleteUserReq: &grpcDomainsV1.DeleteUserReq{}, - deleteUserRes: &grpcDomainsV1.DeleteUserRes{Deleted: false}, - err: apiutil.ErrMissingID, - }, - { - desc: "delete invalid req with invalid token", - token: inValidToken, - deleteUserReq: &grpcDomainsV1.DeleteUserReq{ - Id: id, - }, - deleteUserRes: &grpcDomainsV1.DeleteUserRes{Deleted: false}, - err: apiutil.ErrMissingPolicyEntityType, - }, - } - for _, tc := range cases { - repoCall := svc.On("DeleteUserFromDomains", mock.Anything, tc.deleteUserReq.Id).Return(tc.err) - dpr, err := grpcClient.DeleteUserFromDomains(context.Background(), tc.deleteUserReq) - assert.Equal(t, tc.deleteUserRes.GetDeleted(), dpr.GetDeleted(), fmt.Sprintf("%s: expected %v got %v", tc.desc, tc.deleteUserRes.GetDeleted(), dpr.GetDeleted())) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - repoCall.Unset() - } -} - -func TestRetrieveStatus(t *testing.T) { - conn, err := grpc.NewClient(authAddr, grpc.WithTransportCredentials(insecure.NewCredentials())) - assert.Nil(t, err, fmt.Sprintf("Unexpected error creating client connection %s", err)) - grpcClient := grpcapi.NewDomainsClient(conn, time.Second) - - cases := []struct { - desc string - token string - retrieveReq *grpcCommonV1.RetrieveEntityReq - svcRes domains.Status - svcErr error - retrieveRes *grpcCommonV1.RetrieveEntityRes - err error - }{ - { - desc: "retrieve status with valid req", - token: validToken, - retrieveReq: &grpcCommonV1.RetrieveEntityReq{ - Id: id, - }, - svcRes: domains.EnabledStatus, - retrieveRes: &grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Status: uint32(domains.EnabledStatus), - }, - }, - err: nil, - }, - { - desc: "retrieve status with empty id", - retrieveReq: &grpcCommonV1.RetrieveEntityReq{ - Id: "", - }, - svcRes: domains.AllStatus, - retrieveRes: &grpcCommonV1.RetrieveEntityRes{}, - err: apiutil.ErrMissingID, - }, - { - desc: "retrieve status with invalid id", - retrieveReq: &grpcCommonV1.RetrieveEntityReq{ - Id: "invalid", - }, - svcRes: domains.AllStatus, - svcErr: svcerr.ErrNotFound, - retrieveRes: &grpcCommonV1.RetrieveEntityRes{}, - err: svcerr.ErrNotFound, - }, - } - for _, tc := range cases { - svcCall := svc.On("RetrieveStatus", mock.Anything, tc.retrieveReq.Id).Return(tc.svcRes, tc.svcErr) - dpr, err := grpcClient.RetrieveStatus(context.Background(), tc.retrieveReq) - assert.Equal(t, tc.retrieveRes.Entity, dpr.Entity, fmt.Sprintf("%s: expected %v got %v", tc.desc, tc.retrieveRes.Entity, dpr.Entity)) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - svcCall.Unset() - } -} - -func TestRetrieveIDByRoute(t *testing.T) { - conn, err := grpc.NewClient(authAddr, grpc.WithTransportCredentials(insecure.NewCredentials())) - assert.Nil(t, err, fmt.Sprintf("Unexpected error creating client connection %s", err)) - grpcClient := grpcapi.NewDomainsClient(conn, time.Second) - - validRoute := "validRoute" - - cases := []struct { - desc string - retrieveReq *grpcCommonV1.RetrieveIDByRouteReq - svcRes string - svcErr error - retrieveRes *grpcCommonV1.RetrieveEntityRes - err error - }{ - { - desc: "retrieve id with valid route", - retrieveReq: &grpcCommonV1.RetrieveIDByRouteReq{ - Route: validRoute, - }, - svcRes: id, - retrieveRes: &grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Id: id, - }, - }, - err: nil, - }, - { - desc: "retrieve id with empty route", - retrieveReq: &grpcCommonV1.RetrieveIDByRouteReq{ - Route: "", - }, - svcRes: "", - retrieveRes: &grpcCommonV1.RetrieveEntityRes{}, - err: apiutil.ErrMissingRoute, - }, - { - desc: "retrieve id with invalid route", - retrieveReq: &grpcCommonV1.RetrieveIDByRouteReq{ - Route: "invalid", - }, - svcRes: "", - svcErr: svcerr.ErrNotFound, - retrieveRes: &grpcCommonV1.RetrieveEntityRes{}, - err: svcerr.ErrNotFound, - }, - } - for _, tc := range cases { - svcCall := svc.On("RetrieveIDByRoute", mock.Anything, tc.retrieveReq.Route).Return(tc.svcRes, tc.svcErr) - dpr, err := grpcClient.RetrieveIDByRoute(context.Background(), tc.retrieveReq) - assert.Equal(t, tc.retrieveRes.Entity, dpr.Entity, fmt.Sprintf("%s: expected %v got %v", tc.desc, tc.retrieveRes.Entity, dpr.Entity)) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - svcCall.Unset() - } -} diff --git a/domains/api/grpc/requests.go b/domains/api/grpc/requests.go deleted file mode 100644 index 536ccd1c0..000000000 --- a/domains/api/grpc/requests.go +++ /dev/null @@ -1,44 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package grpc - -import ( - apiutil "github.com/absmach/magistrala/api/http/util" -) - -type deleteUserPoliciesReq struct { - ID string -} - -func (req deleteUserPoliciesReq) validate() error { - if req.ID == "" { - return apiutil.ErrMissingID - } - - return nil -} - -type retrieveStatusReq struct { - ID string -} - -func (req retrieveStatusReq) validate() error { - if req.ID == "" { - return apiutil.ErrMissingID - } - - return nil -} - -type retrieveIDByRouteReq struct { - Route string -} - -func (req retrieveIDByRouteReq) validate() error { - if req.Route == "" { - return apiutil.ErrMissingRoute - } - - return nil -} diff --git a/domains/api/grpc/responses.go b/domains/api/grpc/responses.go deleted file mode 100644 index 47587c172..000000000 --- a/domains/api/grpc/responses.go +++ /dev/null @@ -1,16 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package grpc - -type deleteUserRes struct { - deleted bool -} - -type retrieveIDByRouteRes struct { - id string -} - -type retrieveStatusRes struct { - status uint8 -} diff --git a/domains/api/grpc/server.go b/domains/api/grpc/server.go deleted file mode 100644 index 7f8d617dc..000000000 --- a/domains/api/grpc/server.go +++ /dev/null @@ -1,117 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package grpc - -import ( - "context" - - grpcCommonV1 "github.com/absmach/magistrala/api/grpc/common/v1" - grpcDomainsV1 "github.com/absmach/magistrala/api/grpc/domains/v1" - grpcapi "github.com/absmach/magistrala/auth/api/grpc" - domains "github.com/absmach/magistrala/domains/private" - kitgrpc "github.com/go-kit/kit/transport/grpc" -) - -var _ grpcDomainsV1.DomainsServiceServer = (*domainsGrpcServer)(nil) - -type domainsGrpcServer struct { - grpcDomainsV1.UnimplementedDomainsServiceServer - deleteUserFromDomains kitgrpc.Handler - retrieveStatus kitgrpc.Handler - retrieveIDByRoute kitgrpc.Handler -} - -func NewDomainsServer(svc domains.Service) grpcDomainsV1.DomainsServiceServer { - return &domainsGrpcServer{ - deleteUserFromDomains: kitgrpc.NewServer( - (deleteUserFromDomainsEndpoint(svc)), - decodeDeleteUserRequest, - encodeDeleteUserResponse, - ), - retrieveStatus: kitgrpc.NewServer( - retrieveStatusEndpoint(svc), - decodeRetrieveStatusRequest, - encodeRetrieveStatusResponse, - ), - retrieveIDByRoute: kitgrpc.NewServer( - retrieveIDByRouteEndpoint(svc), - decodeRetrieveIDByRouteRequest, - encodeRetrieveIDByRouteResponse, - ), - } -} - -func decodeDeleteUserRequest(_ context.Context, grpcReq any) (any, error) { - req := grpcReq.(*grpcDomainsV1.DeleteUserReq) - return deleteUserPoliciesReq{ - ID: req.GetId(), - }, nil -} - -func encodeDeleteUserResponse(_ context.Context, grpcRes any) (any, error) { - res := grpcRes.(deleteUserRes) - return &grpcDomainsV1.DeleteUserRes{Deleted: res.deleted}, nil -} - -func (s *domainsGrpcServer) DeleteUserFromDomains(ctx context.Context, req *grpcDomainsV1.DeleteUserReq) (*grpcDomainsV1.DeleteUserRes, error) { - _, res, err := s.deleteUserFromDomains.ServeGRPC(ctx, req) - if err != nil { - return nil, grpcapi.EncodeError(err) - } - return res.(*grpcDomainsV1.DeleteUserRes), nil -} - -func decodeRetrieveStatusRequest(_ context.Context, grpcReq any) (any, error) { - req := grpcReq.(*grpcCommonV1.RetrieveEntityReq) - - return retrieveStatusReq{ - ID: req.GetId(), - }, nil -} - -func encodeRetrieveStatusResponse(_ context.Context, grpcRes any) (any, error) { - res := grpcRes.(retrieveStatusRes) - - return &grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Status: uint32(res.status), - }, - }, nil -} - -func (s *domainsGrpcServer) RetrieveStatus(ctx context.Context, req *grpcCommonV1.RetrieveEntityReq) (*grpcCommonV1.RetrieveEntityRes, error) { - _, res, err := s.retrieveStatus.ServeGRPC(ctx, req) - if err != nil { - return nil, grpcapi.EncodeError(err) - } - - return res.(*grpcCommonV1.RetrieveEntityRes), nil -} - -func decodeRetrieveIDByRouteRequest(_ context.Context, grpcReq any) (any, error) { - req := grpcReq.(*grpcCommonV1.RetrieveIDByRouteReq) - - return retrieveIDByRouteReq{ - Route: req.GetRoute(), - }, nil -} - -func encodeRetrieveIDByRouteResponse(_ context.Context, grpcRes any) (any, error) { - res := grpcRes.(retrieveIDByRouteRes) - - return &grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Id: res.id, - }, - }, nil -} - -func (s *domainsGrpcServer) RetrieveIDByRoute(ctx context.Context, req *grpcCommonV1.RetrieveIDByRouteReq) (*grpcCommonV1.RetrieveEntityRes, error) { - _, res, err := s.retrieveIDByRoute.ServeGRPC(ctx, req) - if err != nil { - return nil, grpcapi.EncodeError(err) - } - - return res.(*grpcCommonV1.RetrieveEntityRes), nil -} diff --git a/domains/api/grpc/setup_test.go b/domains/api/grpc/setup_test.go deleted file mode 100644 index d21864e58..000000000 --- a/domains/api/grpc/setup_test.go +++ /dev/null @@ -1,24 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package grpc_test - -import ( - "os" - "testing" - - "github.com/absmach/magistrala/domains/private/mocks" -) - -var svc *mocks.Service - -func TestMain(m *testing.M) { - svc = new(mocks.Service) - server := startGRPCServer(svc, port) - - code := m.Run() - - server.GracefulStop() - - os.Exit(code) -} diff --git a/domains/api/http/decode.go b/domains/api/http/decode.go deleted file mode 100644 index b325d561d..000000000 --- a/domains/api/http/decode.go +++ /dev/null @@ -1,303 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package http - -import ( - "context" - "encoding/json" - "net/http" - "strings" - "time" - - api "github.com/absmach/magistrala/api/http" - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/domains" - "github.com/absmach/magistrala/pkg/errors" - "github.com/go-chi/chi/v5" -) - -const ( - inviteeUserIDKey = "invitee_user_id" - domainIDKey = "domain_id" - invitedByKey = "invited_by" - roleIDKey = "role_id" - stateKey = "state" -) - -func decodeCreateDomainRequest(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - req := createDomainReq{} - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - - return req, nil -} - -func decodeRetrieveDomainRequest(_ context.Context, r *http.Request) (any, error) { - roles, err := apiutil.ReadBoolQuery(r, api.RolesKey, false) - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - req := retrieveDomainRequest{ - domainID: chi.URLParam(r, "domainID"), - roles: roles, - } - return req, nil -} - -func decodeUpdateDomainRequest(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - - req := updateDomainReq{ - domainID: chi.URLParam(r, "domainID"), - } - - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - - return req, nil -} - -func decodeListDomainRequest(ctx context.Context, r *http.Request) (any, error) { - page, err := decodePageRequest(ctx, r) - if err != nil { - return nil, err - } - req := listDomainsReq{ - page, - } - - return req, nil -} - -func decodeEnableDomainRequest(_ context.Context, r *http.Request) (any, error) { - req := enableDomainReq{ - domainID: chi.URLParam(r, "domainID"), - } - return req, nil -} - -func decodeDisableDomainRequest(_ context.Context, r *http.Request) (any, error) { - req := disableDomainReq{ - domainID: chi.URLParam(r, "domainID"), - } - return req, nil -} - -func decodeFreezeDomainRequest(_ context.Context, r *http.Request) (any, error) { - req := freezeDomainReq{ - domainID: chi.URLParam(r, "domainID"), - } - return req, nil -} - -func decodePageRequest(_ context.Context, r *http.Request) (domains.Page, error) { - s, err := apiutil.ReadStringQuery(r, api.StatusKey, api.DefClientStatus) - if err != nil { - return domains.Page{}, errors.Wrap(apiutil.ErrValidation, err) - } - st, err := domains.ToStatus(s) - if err != nil { - return domains.Page{}, errors.Wrap(apiutil.ErrValidation, err) - } - o, err := apiutil.ReadNumQuery[uint64](r, api.OffsetKey, api.DefOffset) - if err != nil { - return domains.Page{}, errors.Wrap(apiutil.ErrValidation, err) - } - or, err := apiutil.ReadStringQuery(r, api.OrderKey, api.DefOrder) - if err != nil { - return domains.Page{}, errors.Wrap(apiutil.ErrValidation, err) - } - dir, err := apiutil.ReadStringQuery(r, api.DirKey, api.DefDir) - if err != nil { - return domains.Page{}, errors.Wrap(apiutil.ErrValidation, err) - } - l, err := apiutil.ReadNumQuery[uint64](r, api.LimitKey, api.DefLimit) - if err != nil { - return domains.Page{}, errors.Wrap(apiutil.ErrValidation, err) - } - m, err := apiutil.ReadMetadataQuery(r, api.MetadataKey, nil) - if err != nil { - return domains.Page{}, errors.Wrap(apiutil.ErrValidation, err) - } - n, err := apiutil.ReadStringQuery(r, api.NameKey, "") - if err != nil { - return domains.Page{}, errors.Wrap(apiutil.ErrValidation, err) - } - t, err := apiutil.ReadStringQuery(r, api.TagsKey, "") - if err != nil { - return domains.Page{}, errors.Wrap(apiutil.ErrValidation, err) - } - var tq domains.TagsQuery - if t != "" { - tq = domains.ToTagsQuery(t) - } - - allActions, err := apiutil.ReadStringQuery(r, api.ActionsKey, "") - if err != nil { - return domains.Page{}, errors.Wrap(apiutil.ErrValidation, err) - } - - var actions []string - - allActions = strings.TrimSpace(allActions) - if allActions != "" { - actions = strings.Split(allActions, ",") - } - roleID, err := apiutil.ReadStringQuery(r, api.RoleIDKey, "") - if err != nil { - return domains.Page{}, errors.Wrap(apiutil.ErrValidation, err) - } - - roleName, err := apiutil.ReadStringQuery(r, api.RoleNameKey, "") - if err != nil { - return domains.Page{}, errors.Wrap(apiutil.ErrValidation, err) - } - - id, err := apiutil.ReadStringQuery(r, api.IDOrder, "") - if err != nil { - return domains.Page{}, errors.Wrap(apiutil.ErrValidation, err) - } - - ot, err := apiutil.ReadBoolQuery(r, api.OnlyTotal, false) - if err != nil { - return domains.Page{}, errors.Wrap(apiutil.ErrValidation, err) - } - - cfrom, err := apiutil.ReadStringQuery(r, "created_from", "") - if err != nil { - return domains.Page{}, errors.Wrap(apiutil.ErrValidation, err) - } - cto, err := apiutil.ReadStringQuery(r, "created_to", "") - if err != nil { - return domains.Page{}, errors.Wrap(apiutil.ErrValidation, err) - } - - var createdFrom, createdTo time.Time - if cfrom != "" { - if createdFrom, err = time.Parse(time.RFC3339, cfrom); err != nil { - return domains.Page{}, errors.Wrap(apiutil.ErrInvalidQueryParams, err) - } - } - if cto != "" { - if createdTo, err = time.Parse(time.RFC3339, cto); err != nil { - return domains.Page{}, errors.Wrap(apiutil.ErrInvalidQueryParams, err) - } - } - - return domains.Page{ - Offset: o, - Order: or, - Dir: dir, - Limit: l, - Name: n, - Metadata: m, - Tags: tq, - RoleID: roleID, - RoleName: roleName, - Actions: actions, - Status: st, - ID: id, - OnlyTotal: ot, - CreatedFrom: createdFrom, - CreatedTo: createdTo, - }, nil -} - -func decodeSendInvitationReq(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - - var req sendInvitationReq - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - - return req, nil -} - -func decodeListInvitationsReq(_ context.Context, r *http.Request) (any, error) { - offset, err := apiutil.ReadNumQuery[uint64](r, api.OffsetKey, api.DefOffset) - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - limit, err := apiutil.ReadNumQuery[uint64](r, api.LimitKey, api.DefLimit) - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - inviteeUserID, err := apiutil.ReadStringQuery(r, inviteeUserIDKey, "") - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - invitedBy, err := apiutil.ReadStringQuery(r, invitedByKey, "") - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - roleID, err := apiutil.ReadStringQuery(r, roleIDKey, "") - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - domainID, err := apiutil.ReadStringQuery(r, domainIDKey, "") - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - st, err := apiutil.ReadStringQuery(r, stateKey, domains.AllState.String()) - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - state, err := domains.ToState(st) - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - ot, err := apiutil.ReadBoolQuery(r, api.OnlyTotal, false) - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - req := listInvitationsReq{ - InvitationPageMeta: domains.InvitationPageMeta{ - Offset: offset, - Limit: limit, - InvitedBy: invitedBy, - InviteeUserID: inviteeUserID, - RoleID: roleID, - DomainID: domainID, - State: state, - OnlyTotal: ot, - }, - } - - return req, nil -} - -func decodeAcceptInvitationReq(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - - var req acceptInvitationReq - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - - return req, nil -} - -func decodeDeleteInvitationReq(_ context.Context, r *http.Request) (any, error) { - req := deleteInvitationReq{ - domainID: chi.URLParam(r, "domainID"), - } - - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - - return req, nil -} diff --git a/domains/api/http/endpoint.go b/domains/api/http/endpoint.go deleted file mode 100644 index 7892cd8db..000000000 --- a/domains/api/http/endpoint.go +++ /dev/null @@ -1,304 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package http - -import ( - "context" - - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/domains" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/go-kit/kit/endpoint" -) - -// InvitationSent is the message returned when an invitation is sent. -const InvitationSent = "invitation sent" - -func createDomainEndpoint(svc domains.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(createDomainReq) - if err := req.validate(); err != nil { - return nil, err - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthorization - } - - d := domains.Domain{ - ID: req.ID, - Name: req.Name, - Metadata: req.Metadata, - Tags: req.Tags, - Route: req.Route, - } - domain, _, err := svc.CreateDomain(ctx, session, d) - if err != nil { - return nil, err - } - - return createDomainRes{domain}, nil - } -} - -func retrieveDomainEndpoint(svc domains.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(retrieveDomainRequest) - if err := req.validate(); err != nil { - return nil, err - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthorization - } - - domain, err := svc.RetrieveDomain(ctx, session, req.domainID, req.roles) - if err != nil { - return nil, err - } - return retrieveDomainRes{domain}, nil - } -} - -func updateDomainEndpoint(svc domains.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(updateDomainReq) - if err := req.validate(); err != nil { - return nil, err - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthorization - } - - d := domains.DomainReq{ - Name: req.Name, - Metadata: req.Metadata, - Tags: req.Tags, - } - domain, err := svc.UpdateDomain(ctx, session, req.domainID, d) - if err != nil { - return nil, err - } - - return updateDomainRes{domain}, nil - } -} - -func listDomainsEndpoint(svc domains.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(listDomainsReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthorization - } - - dp, err := svc.ListDomains(ctx, session, req.Page) - if err != nil { - return nil, err - } - return listDomainsRes{dp}, nil - } -} - -func enableDomainEndpoint(svc domains.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(enableDomainReq) - if err := req.validate(); err != nil { - return nil, err - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthorization - } - - if _, err := svc.EnableDomain(ctx, session, req.domainID); err != nil { - return nil, err - } - return enableDomainRes{}, nil - } -} - -func disableDomainEndpoint(svc domains.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(disableDomainReq) - if err := req.validate(); err != nil { - return nil, err - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthorization - } - - if _, err := svc.DisableDomain(ctx, session, req.domainID); err != nil { - return nil, err - } - return disableDomainRes{}, nil - } -} - -func freezeDomainEndpoint(svc domains.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(freezeDomainReq) - if err := req.validate(); err != nil { - return nil, err - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthorization - } - - if _, err := svc.FreezeDomain(ctx, session, req.domainID); err != nil { - return nil, err - } - return freezeDomainRes{}, nil - } -} - -func sendInvitationEndpoint(svc domains.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(sendInvitationReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthorization - } - - invitation := domains.Invitation{ - InviteeUserID: req.InviteeUserID, - DomainID: session.DomainID, - RoleID: req.RoleID, - Resend: req.Resend, - } - - if _, err := svc.SendInvitation(ctx, session, invitation); err != nil { - return nil, err - } - - return sendInvitationRes{ - Message: InvitationSent, - }, nil - } -} - -func listDomainInvitationsEndpoint(svc domains.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(listInvitationsReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthorization - } - - page, err := svc.ListDomainInvitations(ctx, session, req.InvitationPageMeta) - if err != nil { - return nil, err - } - - return listInvitationsRes{ - page, - }, nil - } -} - -func listUserInvitationsEndpoint(svc domains.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(listInvitationsReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthorization - } - - page, err := svc.ListInvitations(ctx, session, req.InvitationPageMeta) - if err != nil { - return nil, err - } - - return listInvitationsRes{ - page, - }, nil - } -} - -func acceptInvitationEndpoint(svc domains.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(acceptInvitationReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthorization - } - - if _, err := svc.AcceptInvitation(ctx, session, req.DomainID); err != nil { - return nil, err - } - - return acceptInvitationRes{}, nil - } -} - -func rejectInvitationEndpoint(svc domains.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(acceptInvitationReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthorization - } - - if _, err := svc.RejectInvitation(ctx, session, req.DomainID); err != nil { - return nil, err - } - - return rejectInvitationRes{}, nil - } -} - -func deleteInvitationEndpoint(svc domains.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(deleteInvitationReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthorization - } - session.DomainID = req.domainID - - if err := svc.DeleteInvitation(ctx, session, req.UserID, req.domainID); err != nil { - return nil, err - } - - return deleteInvitationRes{}, nil - } -} diff --git a/domains/api/http/endpoint_test.go b/domains/api/http/endpoint_test.go deleted file mode 100644 index 296955b61..000000000 --- a/domains/api/http/endpoint_test.go +++ /dev/null @@ -1,1871 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package http_test - -import ( - "encoding/json" - "fmt" - "io" - "net/http" - "net/http/httptest" - "strings" - "testing" - "time" - - api "github.com/absmach/magistrala/api/http" - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/domains" - domainsapi "github.com/absmach/magistrala/domains/api/http" - "github.com/absmach/magistrala/domains/mocks" - "github.com/absmach/magistrala/internal/testsutil" - mglog "github.com/absmach/magistrala/logger" - "github.com/absmach/magistrala/pkg/authn" - smqauthn "github.com/absmach/magistrala/pkg/authn" - authnmock "github.com/absmach/magistrala/pkg/authn/mocks" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/pkg/roles" - "github.com/absmach/magistrala/pkg/uuid" - "github.com/go-chi/chi/v5" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" -) - -const contentType = "application/json" - -var ( - validMetadata = domains.Metadata{"role": "client"} - ID = testsutil.GenerateUUID(&testing.T{}) - domain = domains.Domain{ - ID: ID, - Name: "domainname", - Tags: []string{"tag1", "tag2"}, - Metadata: validMetadata, - Status: domains.EnabledStatus, - Route: "mydomain", - } - validToken = "token" - inValidToken = "invalid" - invalid = "invalid" - userID = testsutil.GenerateUUID(&testing.T{}) - validID = testsutil.GenerateUUID(&testing.T{}) - domainID = testsutil.GenerateUUID(&testing.T{}) -) - -type testRequest struct { - client *http.Client - method string - url string - contentType string - token string - body io.Reader -} - -func (tr testRequest) make() (*http.Response, error) { - req, err := http.NewRequest(tr.method, tr.url, tr.body) - if err != nil { - return nil, err - } - - if tr.token != "" { - req.Header.Set("Authorization", apiutil.BearerPrefix+tr.token) - } - - if tr.contentType != "" { - req.Header.Set("Content-Type", tr.contentType) - } - - req.Header.Set("Referer", "http://localhost") - - return tr.client.Do(req) -} - -func toJSON(data any) string { - jsonData, err := json.Marshal(data) - if err != nil { - return "" - } - return string(jsonData) -} - -func newDomainsServer() (*httptest.Server, *mocks.Service, *authnmock.Authentication) { - logger := mglog.NewMock() - svc := new(mocks.Service) - authn := new(authnmock.Authentication) - mux := chi.NewMux() - idp := uuid.NewMock() - am := smqauthn.NewAuthNMiddleware(authn, smqauthn.WithAllowUnverifiedUser(true)) - domainsapi.MakeHandler(svc, am, mux, logger, "", idp) - return httptest.NewServer(mux), svc, authn -} - -func TestCreateDomain(t *testing.T) { - ds, svc, auth := newDomainsServer() - defer ds.Close() - - cases := []struct { - desc string - domain domains.Domain - token string - session authn.Session - contentType string - svcErr error - status int - authnErr error - err error - }{ - { - desc: "register a new domain successfully", - domain: domains.Domain{ - Name: "test", - Metadata: domains.Metadata{"role": "domain"}, - Tags: []string{"tag1", "tag2"}, - Route: "test", - }, - token: validToken, - contentType: contentType, - status: http.StatusCreated, - err: nil, - }, - { - desc: "register a new domain with empty token", - domain: domains.Domain{ - Name: "test", - Metadata: domains.Metadata{"role": "domain"}, - Tags: []string{"tag1", "tag2"}, - Route: "test", - }, - token: "", - contentType: contentType, - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "register a new domain with invalid token", - domain: domains.Domain{ - Name: "test", - Metadata: domains.Metadata{"role": "domain"}, - Tags: []string{"tag1", "tag2"}, - Route: "test", - }, - token: inValidToken, - contentType: contentType, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "register a new domain with an empty name", - domain: domains.Domain{ - Name: "", - Metadata: domains.Metadata{"role": "domain"}, - Tags: []string{"tag1", "tag2"}, - Route: "test", - }, - token: validToken, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrMissingName, - }, - { - desc: "register a new domain with an empty route", - domain: domains.Domain{ - Name: "test", - Metadata: domains.Metadata{"role": "domain"}, - Tags: []string{"tag1", "tag2"}, - Route: "", - }, - token: validToken, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrMissingRoute, - }, - { - desc: "register a new domain with invalid content type", - domain: domains.Domain{ - Name: "test", - Metadata: domains.Metadata{"role": "domain"}, - Tags: []string{"tag1", "tag2"}, - Route: "test", - }, - token: validToken, - contentType: "application/xml", - status: http.StatusUnsupportedMediaType, - err: apiutil.ErrUnsupportedContentType, - }, - { - desc: "register a new domain that cant be marshalled", - domain: domains.Domain{ - Name: "test", - Metadata: map[string]any{ - "test": make(chan int), - }, - Tags: []string{"tag1", "tag2"}, - Route: "test", - }, - token: validToken, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrMalformedRequestBody, - }, - { - desc: "register domain with service error", - domain: domains.Domain{ - Name: "test", - Metadata: domains.Metadata{"role": "domain"}, - Tags: []string{"tag1", "tag2"}, - Route: "test", - }, - token: validToken, - contentType: contentType, - status: http.StatusUnprocessableEntity, - svcErr: svcerr.ErrCreateEntity, - err: svcerr.ErrCreateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - data := toJSON(tc.domain) - req := testRequest{ - client: ds.Client(), - method: http.MethodPost, - url: fmt.Sprintf("%s/domains", ds.URL), - contentType: tc.contentType, - token: tc.token, - body: strings.NewReader(data), - } - if tc.token == validToken { - tc.session = authn.Session{UserID: userID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - svcCall := svc.On("CreateDomain", mock.Anything, tc.session, tc.domain).Return(tc.domain, []roles.RoleProvision{}, tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - var errRes respBody - err = json.NewDecoder(res.Body).Decode(&errRes) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if errRes.Err != "" || errRes.Message != "" { - err = errors.Wrap(errors.New(errRes.Err), errors.New(errRes.Message)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestListDomains(t *testing.T) { - ds, svc, auth := newDomainsServer() - defer ds.Close() - - cases := []struct { - desc string - token string - session authn.Session - query string - page domains.Page - listDomainsResp domains.DomainsPage - status int - svcErr error - authnErr error - err error - }{ - { - desc: "list domains with valid token", - token: validToken, - page: domains.Page{ - Offset: api.DefOffset, - Limit: api.DefLimit, - Order: api.DefOrder, - Dir: api.DefDir, - }, - status: http.StatusOK, - listDomainsResp: domains.DomainsPage{ - Total: 1, - Domains: []domains.Domain{domain}, - }, - err: nil, - }, - { - desc: "list domains with empty token", - token: "", - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "list domains with invalid token", - token: inValidToken, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "list domains with offset", - token: validToken, - query: "offset=1", - page: domains.Page{ - Offset: 1, - Limit: api.DefLimit, - Order: api.DefOrder, - Dir: api.DefDir, - }, - listDomainsResp: domains.DomainsPage{ - Total: 1, - Domains: []domains.Domain{domain}, - }, - status: http.StatusOK, - err: nil, - }, - { - desc: "list domains with invalid offset", - token: validToken, - query: "offset=invalid", - status: http.StatusBadRequest, - err: apiutil.ErrValidation, - }, - { - desc: "list domains with limit", - token: validToken, - query: "limit=1", - page: domains.Page{ - Offset: api.DefOffset, - Limit: 1, - Order: api.DefOrder, - Dir: api.DefDir, - }, - listDomainsResp: domains.DomainsPage{ - Total: 1, - Domains: []domains.Domain{domain}, - }, - status: http.StatusOK, - err: nil, - }, - { - desc: "list domains with invalid limit", - token: validToken, - query: "limit=invalid", - status: http.StatusBadRequest, - err: apiutil.ErrValidation, - }, - { - desc: "list domains with name", - token: validToken, - listDomainsResp: domains.DomainsPage{ - Total: 1, - Domains: []domains.Domain{domain}, - }, - query: "name=domainname", - page: domains.Page{ - Offset: api.DefOffset, - Limit: api.DefLimit, - Order: api.DefOrder, - Dir: api.DefDir, - Name: "domainname", - }, - status: http.StatusOK, - err: nil, - }, - { - desc: "list domains with empty name", - token: validToken, - query: "name= ", - status: http.StatusBadRequest, - err: apiutil.ErrValidation, - }, - { - desc: "list domains with duplicate name", - token: validToken, - query: "name=1&name=2", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list domains with status", - token: validToken, - listDomainsResp: domains.DomainsPage{ - Total: 1, - Domains: []domains.Domain{domain}, - }, - query: "status=enabled", - page: domains.Page{ - Offset: api.DefOffset, - Limit: api.DefLimit, - Order: api.DefOrder, - Dir: api.DefDir, - Status: domains.EnabledStatus, - }, - status: http.StatusOK, - err: nil, - }, - { - desc: "list domains with invalid status", - token: validToken, - query: "status=invalid", - status: http.StatusBadRequest, - err: apiutil.ErrValidation, - }, - { - desc: "list domains with duplicate status", - token: validToken, - query: "status=enabled&status=disabled", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list domains with single tag", - token: validToken, - page: domains.Page{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Tags: domains.TagsQuery{Elements: []string{"tag1"}, Operator: domains.OrOp}, - }, - listDomainsResp: domains.DomainsPage{ - Total: 1, - Domains: []domains.Domain{domain}, - }, - query: "tags=tag1", - status: http.StatusOK, - err: nil, - }, - { - desc: "list domains with multiple tags and OR operator", - token: validToken, - page: domains.Page{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Tags: domains.TagsQuery{Elements: []string{"tag1", "tag2", "tag3"}, Operator: domains.OrOp}, - }, - listDomainsResp: domains.DomainsPage{ - Total: 1, - Domains: []domains.Domain{domain}, - }, - query: "tags=tag1,tag2,tag3", - status: http.StatusOK, - err: nil, - }, - { - desc: "list domains with multiple tags and AND operator", - token: validToken, - page: domains.Page{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Tags: domains.TagsQuery{Elements: []string{"tag1", "tag2", "tag3"}, Operator: domains.AndOp}, - }, - listDomainsResp: domains.DomainsPage{ - Total: 1, - Domains: []domains.Domain{domain}, - }, - query: "tags=tag1%2Btag2%2Btag3", - status: http.StatusOK, - err: nil, - }, - { - desc: "list domains with duplicate tags", - token: validToken, - query: "tags=tag1&tags=tag2", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list domains with metadata", - token: validToken, - query: "metadata=%7B%22domain%22%3A%20%22example.com%22%7D&", - page: domains.Page{ - Offset: api.DefOffset, - Limit: api.DefLimit, - Order: api.DefOrder, - Dir: api.DefDir, - Metadata: domains.Metadata{ - "domain": "example.com", - }, - }, - listDomainsResp: domains.DomainsPage{ - Total: 1, - Domains: []domains.Domain{domain}, - }, - status: http.StatusOK, - err: nil, - }, - { - desc: "list domains with invalid metadata", - token: validToken, - query: "metadata=invalid", - status: http.StatusBadRequest, - err: apiutil.ErrValidation, - }, - { - desc: "list domains with duplicate metadata", - token: validToken, - query: "metadata=%7B%22domain%22%3A%20%22example.com%22%7D&metadata=%7B%22domain%22%3A%20%22example.com%22%7D", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list domains with role name", - token: validToken, - query: "role_name=view", - page: domains.Page{ - Offset: api.DefOffset, - Limit: api.DefLimit, - Order: api.DefOrder, - Dir: api.DefDir, - RoleName: "view", - }, - listDomainsResp: domains.DomainsPage{ - Total: 1, - Domains: []domains.Domain{domain}, - }, - status: http.StatusOK, - err: nil, - }, - { - desc: "list domains with invalid role name", - token: validToken, - query: "role_name= ", - status: http.StatusBadRequest, - err: apiutil.ErrValidation, - }, - { - desc: "list domains with duplicate role name", - token: validToken, - query: "role_name=view&role_name=view", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list domains with order", - token: validToken, - page: domains.Page{ - Offset: api.DefOffset, - Limit: api.DefLimit, - Order: "name", - Dir: api.DefDir, - }, - query: "order=name", - listDomainsResp: domains.DomainsPage{ - Total: 1, - Domains: []domains.Domain{domain}, - }, - status: http.StatusOK, - }, - { - desc: "list domains with invalid order", - token: validToken, - query: "order= ", - status: http.StatusBadRequest, - err: apiutil.ErrValidation, - }, - { - desc: "list domains with duplicate order", - token: validToken, - query: "order=name&order=name", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list domains with dir", - token: validToken, - page: domains.Page{ - Offset: api.DefOffset, - Limit: api.DefLimit, - Order: api.DefOrder, - Dir: "asc", - }, - query: "dir=asc", - listDomainsResp: domains.DomainsPage{ - Total: 1, - Domains: []domains.Domain{domain}, - }, - status: http.StatusOK, - }, - { - desc: "list domains with invalid dir", - token: validToken, - query: "dir= ", - status: http.StatusBadRequest, - err: apiutil.ErrValidation, - }, - { - desc: "list domains with duplicate dir", - token: validToken, - query: "dir=asc&dir=asc", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list domains with service error", - token: validToken, - page: domains.Page{ - Offset: api.DefOffset, - Limit: api.DefLimit, - Order: api.DefOrder, - Dir: api.DefDir, - }, - status: http.StatusUnprocessableEntity, - listDomainsResp: domains.DomainsPage{}, - svcErr: svcerr.ErrViewEntity, - err: svcerr.ErrViewEntity, - }, - { - desc: "list domains with created_from parameter", - token: validToken, - query: "created_from=2024-01-01T00:00:00Z", - page: domains.Page{ - Offset: api.DefOffset, - Limit: api.DefLimit, - Order: api.DefOrder, - Dir: api.DefDir, - CreatedFrom: time.Date(2024, 1, 1, 0, 0, 0, 0, time.UTC), - }, - listDomainsResp: domains.DomainsPage{ - Total: 1, - Domains: []domains.Domain{domain}, - }, - status: http.StatusOK, - err: nil, - }, - { - desc: "list domains with created_to parameter", - token: validToken, - query: "created_to=2024-12-31T23:59:59Z", - page: domains.Page{ - Offset: api.DefOffset, - Limit: api.DefLimit, - Order: api.DefOrder, - Dir: api.DefDir, - CreatedTo: time.Date(2024, 12, 31, 23, 59, 59, 0, time.UTC), - }, - listDomainsResp: domains.DomainsPage{ - Total: 1, - Domains: []domains.Domain{domain}, - }, - status: http.StatusOK, - err: nil, - }, - { - desc: "list domains with both created_from and created_to parameters", - token: validToken, - query: "created_from=2024-01-01T00:00:00Z&created_to=2024-12-31T23:59:59Z", - page: domains.Page{ - Offset: api.DefOffset, - Limit: api.DefLimit, - Order: api.DefOrder, - Dir: api.DefDir, - CreatedFrom: time.Date(2024, 1, 1, 0, 0, 0, 0, time.UTC), - CreatedTo: time.Date(2024, 12, 31, 23, 59, 59, 0, time.UTC), - }, - listDomainsResp: domains.DomainsPage{ - Total: 1, - Domains: []domains.Domain{domain}, - }, - status: http.StatusOK, - err: nil, - }, - { - desc: "list domains with invalid created_from", - token: validToken, - query: "created_from=invalid-timestamp", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list domains with duplicate created_from", - token: validToken, - query: "created_from=2024-01-01T00:00:00Z&created_from=2024-01-02T00:00:00Z", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list domains with invalid created_to", - token: validToken, - query: "created_to=invalid-timestamp", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list domains with duplicate created_to", - token: validToken, - query: "created_to=2024-12-31T23:59:59Z&created_to=2025-01-01T00:00:00Z", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - client: ds.Client(), - method: http.MethodGet, - url: fmt.Sprintf("%s/domains?", ds.URL) + tc.query, - token: tc.token, - } - if tc.token == validToken { - tc.session = authn.Session{UserID: userID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - svcCall := svc.On("ListDomains", mock.Anything, tc.session, tc.page).Return(tc.listDomainsResp, tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestViewDomain(t *testing.T) { - ds, svc, auth := newDomainsServer() - defer ds.Close() - - cases := []struct { - desc string - token string - session authn.Session - domainID string - status int - svcRes domains.Domain - svcErr error - authnErr error - err error - }{ - { - desc: "view domain successfully", - token: validToken, - domainID: domain.ID, - status: http.StatusOK, - err: nil, - }, - { - desc: "view domain with empty token", - token: "", - domainID: domain.ID, - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "view domain with invalid token", - token: inValidToken, - domainID: domain.ID, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "view domain with service error", - token: validToken, - domainID: invalid, - status: http.StatusUnprocessableEntity, - svcErr: svcerr.ErrViewEntity, - err: svcerr.ErrViewEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - client: ds.Client(), - method: http.MethodGet, - url: fmt.Sprintf("%s/domains/%s", ds.URL, tc.domainID), - token: tc.token, - } - if tc.token == validToken { - tc.session = authn.Session{UserID: userID, DomainID: tc.domainID, DomainUserID: tc.domainID + "_" + userID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - svcCall := svc.On("RetrieveDomain", mock.Anything, tc.session, tc.domainID, false).Return(tc.svcRes, tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - var errRes respBody - err = json.NewDecoder(res.Body).Decode(&errRes) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if errRes.Err != "" || errRes.Message != "" { - err = errors.Wrap(errors.New(errRes.Err), errors.New(errRes.Message)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestUpdateDomain(t *testing.T) { - ds, svc, auth := newDomainsServer() - defer ds.Close() - - updatedName := "test" - updatedMetadata := domains.Metadata{"role": "domain"} - updatedTags := []string{"tag1", "tag2"} - - updatedDomain := domains.Domain{ - ID: ID, - Name: updatedName, - Metadata: updatedMetadata, - Tags: updatedTags, - } - unMetadata := domains.Metadata{ - "test": make(chan int), - } - - cases := []struct { - desc string - token string - session authn.Session - domainID string - updateReq domains.DomainReq - contentType string - status int - svcRes domains.Domain - svcErr error - authnErr error - err error - }{ - { - desc: "update domain successfully", - token: validToken, - domainID: domain.ID, - updateReq: domains.DomainReq{ - Name: &updatedName, - Metadata: &updatedMetadata, - Tags: &updatedTags, - }, - contentType: contentType, - status: http.StatusOK, - svcRes: updatedDomain, - err: nil, - }, - { - desc: "update domain name successfully", - token: validToken, - domainID: domain.ID, - updateReq: domains.DomainReq{ - Name: &updatedName, - }, - contentType: contentType, - status: http.StatusOK, - svcRes: updatedDomain, - err: nil, - }, - { - desc: "update domain tags successfully", - token: validToken, - domainID: domain.ID, - updateReq: domains.DomainReq{ - Tags: &updatedTags, - }, - contentType: contentType, - status: http.StatusOK, - svcRes: updatedDomain, - err: nil, - }, - { - desc: "update domain metadata successfully", - token: validToken, - domainID: domain.ID, - updateReq: domains.DomainReq{ - Metadata: &updatedMetadata, - }, - contentType: contentType, - status: http.StatusOK, - svcRes: updatedDomain, - err: nil, - }, - { - desc: "update domain with empty token", - token: "", - domainID: domain.ID, - updateReq: domains.DomainReq{ - Name: &updatedName, - Metadata: &updatedMetadata, - Tags: &updatedTags, - }, - contentType: contentType, - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "update domain with invalid token", - token: inValidToken, - domainID: domain.ID, - updateReq: domains.DomainReq{ - Name: &updatedName, - Metadata: &updatedMetadata, - Tags: &updatedTags, - }, - contentType: contentType, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "update domain with invalid content type", - token: validToken, - domainID: domain.ID, - updateReq: domains.DomainReq{ - Name: &updatedName, - Metadata: &updatedMetadata, - Tags: &updatedTags, - }, - contentType: "application/xml", - status: http.StatusUnsupportedMediaType, - err: apiutil.ErrUnsupportedContentType, - }, - { - desc: "update domain with data that cant be marshalled", - token: validToken, - domainID: domain.ID, - updateReq: domains.DomainReq{ - Name: &updatedName, - Metadata: &unMetadata, - Tags: &updatedTags, - }, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrMalformedRequestBody, - }, - { - desc: "update domain with invalid id", - token: validToken, - domainID: invalid, - updateReq: domains.DomainReq{ - Name: &updatedName, - Metadata: &updatedMetadata, - Tags: &updatedTags, - }, - contentType: contentType, - status: http.StatusUnprocessableEntity, - svcErr: svcerr.ErrUpdateEntity, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - data := toJSON(tc.updateReq) - req := testRequest{ - client: ds.Client(), - method: http.MethodPatch, - url: fmt.Sprintf("%s/domains/%s", ds.URL, tc.domainID), - body: strings.NewReader(data), - contentType: tc.contentType, - token: tc.token, - } - - if tc.token == validToken { - tc.session = authn.Session{UserID: userID, DomainID: tc.domainID, DomainUserID: tc.domainID + "_" + userID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - svcCall := svc.On("UpdateDomain", mock.Anything, tc.session, tc.domainID, tc.updateReq).Return(tc.svcRes, tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - var errRes respBody - err = json.NewDecoder(res.Body).Decode(&errRes) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if errRes.Err != "" || errRes.Message != "" { - err = errors.Wrap(errors.New(errRes.Err), errors.New(errRes.Message)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestEnableDomain(t *testing.T) { - ds, svc, auth := newDomainsServer() - defer ds.Close() - - cases := []struct { - desc string - token string - session authn.Session - domainID string - status int - svcErr error - svcRes domains.Domain - authnErr error - err error - }{ - { - desc: "enable domain with valid token", - token: validToken, - domainID: domain.ID, - status: http.StatusOK, - svcRes: domain, - err: nil, - }, - { - desc: "enable domain with invalid token", - token: inValidToken, - domainID: domain.ID, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "enable domain with empty token", - token: "", - domainID: domain.ID, - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "enable domain with empty id", - token: validToken, - domainID: "", - status: http.StatusBadRequest, - err: apiutil.ErrMissingID, - }, - { - desc: "enable domain with invalid id", - token: validToken, - domainID: invalid, - status: http.StatusUnprocessableEntity, - svcErr: svcerr.ErrUpdateEntity, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - client: ds.Client(), - method: http.MethodPost, - url: fmt.Sprintf("%s/domains/%s/enable", ds.URL, tc.domainID), - contentType: contentType, - token: tc.token, - } - if tc.token == validToken { - tc.session = authn.Session{UserID: userID, DomainID: tc.domainID, DomainUserID: tc.domainID + "_" + userID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - svcCall := svc.On("EnableDomain", mock.Anything, tc.session, tc.domainID).Return(tc.svcRes, tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestDisableDomain(t *testing.T) { - ds, svc, auth := newDomainsServer() - defer ds.Close() - - cases := []struct { - desc string - token string - session authn.Session - domainID string - status int - svcErr error - svcRes domains.Domain - authnErr error - err error - }{ - { - desc: "disable domain with valid token", - token: validToken, - domainID: domain.ID, - status: http.StatusOK, - svcRes: domain, - err: nil, - }, - { - desc: "disable domain with invalid token", - token: inValidToken, - domainID: domain.ID, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "disable domain with empty token", - token: "", - domainID: domain.ID, - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "disable domain with empty id", - token: validToken, - domainID: "", - status: http.StatusBadRequest, - err: apiutil.ErrMissingID, - }, - { - desc: "disable domain with invalid id", - token: validToken, - domainID: invalid, - status: http.StatusUnprocessableEntity, - svcErr: svcerr.ErrUpdateEntity, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - client: ds.Client(), - method: http.MethodPost, - url: fmt.Sprintf("%s/domains/%s/disable", ds.URL, tc.domainID), - contentType: contentType, - token: tc.token, - } - if tc.token == validToken { - tc.session = authn.Session{UserID: userID, DomainID: tc.domainID, DomainUserID: tc.domainID + "_" + userID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - svcCall := svc.On("DisableDomain", mock.Anything, tc.session, tc.domainID).Return(tc.svcRes, tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestFreezeDomain(t *testing.T) { - ds, svc, auth := newDomainsServer() - defer ds.Close() - - cases := []struct { - desc string - token string - session authn.Session - domainID string - status int - svcErr error - svcRes domains.Domain - authnErr error - err error - }{ - { - desc: "freeze domain with valid token", - token: validToken, - domainID: domain.ID, - status: http.StatusOK, - svcRes: domain, - err: nil, - }, - { - desc: "freeze domain with invalid token", - token: inValidToken, - domainID: domain.ID, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "freeze domain with empty token", - token: "", - domainID: domain.ID, - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "freeze domain with empty id", - token: validToken, - domainID: "", - status: http.StatusBadRequest, - err: apiutil.ErrMissingID, - }, - { - desc: "freeze domain with invalid id", - token: validToken, - domainID: invalid, - status: http.StatusUnprocessableEntity, - svcErr: svcerr.ErrUpdateEntity, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - client: ds.Client(), - method: http.MethodPost, - url: fmt.Sprintf("%s/domains/%s/freeze", ds.URL, tc.domainID), - contentType: contentType, - token: tc.token, - } - if tc.token == validToken { - tc.session = authn.Session{UserID: userID, DomainID: tc.domainID, DomainUserID: tc.domainID + "_" + userID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - svcCall := svc.On("FreezeDomain", mock.Anything, tc.session, tc.domainID).Return(tc.svcRes, tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestSendInvitation(t *testing.T) { - is, svc, auth := newDomainsServer() - - cases := []struct { - desc string - token string - domainID string - data string - session authn.Session - contentType string - status int - authnErr error - svcErr error - }{ - { - desc: "send invitation with valid request", - token: validToken, - domainID: domainID, - data: fmt.Sprintf(`{"invitee_user_id": "%s","role_id": "%s"}`, validID, validID), - status: http.StatusCreated, - contentType: contentType, - svcErr: nil, - }, - { - desc: "send invitation with invalid token", - token: "", - domainID: domainID, - data: fmt.Sprintf(`{"invitee_user_id": "%s","role_id": "%s"}`, validID, validID), - status: http.StatusUnauthorized, - contentType: contentType, - svcErr: nil, - }, - { - desc: "send invitation with empty domain_id", - token: validToken, - domainID: "", - data: fmt.Sprintf(`{"invitee_user_id": "%s","role_id": "%s"}`, validID, validID), - status: http.StatusBadRequest, - contentType: contentType, - svcErr: nil, - }, - { - desc: "send invitation with invalid content type", - token: validToken, - domainID: domainID, - data: fmt.Sprintf(`{"invitee_user_id": "%s","role_id": "%s"}`, validID, validID), - status: http.StatusUnsupportedMediaType, - contentType: "text/plain", - svcErr: nil, - }, - { - desc: "send invitation with invalid data", - token: validToken, - domainID: domainID, - data: `data`, - status: http.StatusBadRequest, - contentType: contentType, - svcErr: nil, - }, - { - desc: "send invitation with service error", - token: validToken, - domainID: domainID, - data: fmt.Sprintf(`{"invitee_user_id": "%s","role_id": "%s"}`, validID, validID), - status: http.StatusForbidden, - contentType: contentType, - svcErr: svcerr.ErrAuthorization, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = authn.Session{UserID: userID, DomainID: domainID, DomainUserID: domainID + "_" + userID} - } - authnCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - repoCall := svc.On("SendInvitation", mock.Anything, tc.session, mock.Anything).Return(domains.Invitation{}, tc.svcErr) - req := testRequest{ - client: is.Client(), - method: http.MethodPost, - url: fmt.Sprintf("%s/domains/%s/invitations", is.URL, tc.domainID), - token: tc.token, - contentType: tc.contentType, - body: strings.NewReader(tc.data), - } - - res, err := req.make() - assert.Nil(t, err, tc.desc) - assert.Equal(t, tc.status, res.StatusCode, tc.desc) - repoCall.Unset() - authnCall.Unset() - }) - } -} - -func TestListDomainInvitations(t *testing.T) { - is, svc, auth := newDomainsServer() - - cases := []struct { - desc string - token string - session authn.Session - domainID string - query string - contentType string - status int - svcErr error - authnErr error - }{ - { - desc: "list invitations with valid request", - token: validToken, - domainID: domainID, - status: http.StatusOK, - contentType: contentType, - svcErr: nil, - }, - { - desc: "list invitations with invalid token", - token: "", - domainID: domainID, - status: http.StatusUnauthorized, - contentType: contentType, - svcErr: nil, - }, - { - desc: "list invitations with offset", - token: validToken, - domainID: domainID, - query: "offset=1", - status: http.StatusOK, - contentType: contentType, - svcErr: nil, - }, - { - desc: "list invitations with invalid offset", - token: validToken, - domainID: domainID, - query: "offset=invalid", - status: http.StatusBadRequest, - contentType: contentType, - svcErr: nil, - }, - { - desc: "list invitations with limit", - token: validToken, - query: "limit=1", - domainID: domainID, - status: http.StatusOK, - contentType: contentType, - svcErr: nil, - }, - { - desc: "list invitations with invalid limit", - token: validToken, - domainID: domainID, - query: "limit=invalid", - status: http.StatusBadRequest, - contentType: contentType, - svcErr: nil, - }, - { - desc: "list invitations with invitee_user_id", - token: validToken, - domainID: domainID, - query: fmt.Sprintf("invitee_user_id=%s", validID), - status: http.StatusOK, - contentType: contentType, - svcErr: nil, - }, - { - desc: "list invitations with duplicate invitee_user_id", - token: validToken, - domainID: domainID, - query: "invitee_user_id=1&invitee_user_id=2", - status: http.StatusBadRequest, - contentType: contentType, - svcErr: nil, - }, - { - desc: "list invitations with invited_by", - token: validToken, - domainID: domainID, - query: fmt.Sprintf("invited_by=%s", validID), - status: http.StatusOK, - contentType: contentType, - svcErr: nil, - }, - { - desc: "list invitations with duplicate invited_by", - token: validToken, - domainID: domainID, - query: "invited_by=1&invited_by=2", - status: http.StatusBadRequest, - contentType: contentType, - svcErr: nil, - }, - { - desc: "list invitations with state", - token: validToken, - domainID: domainID, - query: "state=pending", - status: http.StatusOK, - contentType: contentType, - svcErr: nil, - }, - { - desc: "list invitations with invalid state", - token: validToken, - domainID: domainID, - query: "state=invalid", - status: http.StatusBadRequest, - contentType: contentType, - svcErr: nil, - }, - { - desc: "list invitations with duplicate state", - token: validToken, - domainID: domainID, - query: "state=all&state=all", - status: http.StatusBadRequest, - contentType: contentType, - svcErr: nil, - }, - { - desc: "list invitations with service error", - token: validToken, - domainID: invalid, - status: http.StatusForbidden, - contentType: contentType, - svcErr: svcerr.ErrAuthorization, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = authn.Session{UserID: userID, DomainID: tc.domainID} - } - authnCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - repoCall := svc.On("ListDomainInvitations", mock.Anything, tc.session, mock.Anything).Return(domains.InvitationPage{}, tc.svcErr) - req := testRequest{ - client: is.Client(), - method: http.MethodGet, - url: fmt.Sprintf("%s/domains/%s/invitations?", is.URL, tc.domainID) + tc.query, - token: tc.token, - contentType: tc.contentType, - } - res, err := req.make() - assert.Nil(t, err, tc.desc) - assert.Equal(t, tc.status, res.StatusCode, tc.desc) - repoCall.Unset() - authnCall.Unset() - }) - } -} - -func TestListUserInvitations(t *testing.T) { - is, svc, auth := newDomainsServer() - - cases := []struct { - desc string - token string - session authn.Session - query string - contentType string - status int - svcErr error - authnErr error - }{ - { - desc: "list invitations with valid request", - token: validToken, - status: http.StatusOK, - contentType: contentType, - svcErr: nil, - }, - { - desc: "list invitations with invalid token", - token: "", - status: http.StatusUnauthorized, - contentType: contentType, - svcErr: nil, - }, - { - desc: "list invitations with offset", - token: validToken, - query: "offset=1", - status: http.StatusOK, - contentType: contentType, - svcErr: nil, - }, - { - desc: "list invitations with invalid offset", - token: validToken, - query: "offset=invalid", - status: http.StatusBadRequest, - contentType: contentType, - svcErr: nil, - }, - { - desc: "list invitations with limit", - token: validToken, - query: "limit=1", - status: http.StatusOK, - contentType: contentType, - svcErr: nil, - }, - { - desc: "list invitations with invalid limit", - token: validToken, - query: "limit=invalid", - status: http.StatusBadRequest, - contentType: contentType, - svcErr: nil, - }, - { - desc: "list invitations with invited_by", - token: validToken, - query: fmt.Sprintf("invited_by=%s", validID), - status: http.StatusOK, - contentType: contentType, - svcErr: nil, - }, - { - desc: "list invitations with duplicate invited_by", - token: validToken, - query: "invited_by=1&invited_by=2", - status: http.StatusBadRequest, - contentType: contentType, - svcErr: nil, - }, - { - desc: "list invitations with state", - token: validToken, - query: "state=pending", - status: http.StatusOK, - contentType: contentType, - svcErr: nil, - }, - { - desc: "list invitations with invalid state", - token: validToken, - query: "state=invalid", - status: http.StatusBadRequest, - contentType: contentType, - svcErr: nil, - }, - { - desc: "list invitations with duplicate state", - token: validToken, - query: "state=all&state=all", - status: http.StatusBadRequest, - contentType: contentType, - svcErr: nil, - }, - { - desc: "list invitations with service error", - token: validToken, - status: http.StatusForbidden, - contentType: contentType, - svcErr: svcerr.ErrAuthorization, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = authn.Session{UserID: userID} - } - authnCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - repoCall := svc.On("ListInvitations", mock.Anything, tc.session, mock.Anything).Return(domains.InvitationPage{}, tc.svcErr) - req := testRequest{ - client: is.Client(), - method: http.MethodGet, - url: is.URL + "/invitations?" + tc.query, - token: tc.token, - contentType: tc.contentType, - } - res, err := req.make() - assert.Nil(t, err, tc.desc) - assert.Equal(t, tc.status, res.StatusCode, tc.desc) - repoCall.Unset() - authnCall.Unset() - }) - } -} - -func TestDeleteInvitation(t *testing.T) { - is, svc, auth := newDomainsServer() - - cases := []struct { - desc string - token string - session authn.Session - domainID string - userID string - contentType string - status int - svcErr error - authnErr error - }{ - { - desc: "delete invitation with valid request", - token: validToken, - userID: validID, - domainID: domainID, - status: http.StatusNoContent, - contentType: contentType, - svcErr: nil, - }, - { - desc: "delete invitation with invalid token", - token: "", - userID: validID, - domainID: domainID, - status: http.StatusUnauthorized, - contentType: contentType, - svcErr: nil, - }, - { - desc: "delete invitation with service error", - token: validToken, - userID: validID, - domainID: domainID, - status: http.StatusForbidden, - contentType: contentType, - svcErr: svcerr.ErrAuthorization, - }, - { - desc: "delete invitation with empty invitee_user_id", - token: validToken, - userID: "", - domainID: domainID, - status: http.StatusBadRequest, - contentType: contentType, - svcErr: nil, - }, - { - desc: "delete invitation with empty domain_id", - token: validToken, - userID: validID, - domainID: "", - status: http.StatusBadRequest, - contentType: contentType, - svcErr: nil, - }, - { - desc: "delete invitation with empty invitee_user_id and domain_id", - token: validToken, - userID: "", - domainID: "", - status: http.StatusBadRequest, - contentType: contentType, - svcErr: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = authn.Session{UserID: userID, DomainID: domainID, DomainUserID: domainID + "_" + userID} - } - authnCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - repoCall := svc.On("DeleteInvitation", mock.Anything, tc.session, tc.userID, tc.domainID).Return(tc.svcErr) - - data := fmt.Sprintf(`{"user_id": "%s"}`, tc.userID) - req := testRequest{ - client: is.Client(), - method: http.MethodDelete, - url: fmt.Sprintf("%s/domains/%s/invitations", is.URL, tc.domainID), - token: tc.token, - contentType: tc.contentType, - body: strings.NewReader(data), - } - - res, err := req.make() - assert.Nil(t, err, tc.desc) - assert.Equal(t, tc.status, res.StatusCode, tc.desc) - repoCall.Unset() - authnCall.Unset() - }) - } -} - -func TestAcceptInvitation(t *testing.T) { - is, svc, auth := newDomainsServer() - - cases := []struct { - desc string - token string - session authn.Session - data string - contentType string - status int - svcErr error - authnErr error - }{ - { - desc: "accept invitation with valid request", - data: fmt.Sprintf(`{"domain_id": "%s"}`, validID), - token: validToken, - status: http.StatusNoContent, - contentType: contentType, - svcErr: nil, - }, - { - desc: "accept invitation with invalid token", - token: "", - data: fmt.Sprintf(`{"domain_id": "%s"}`, validID), - status: http.StatusUnauthorized, - contentType: contentType, - svcErr: nil, - }, - { - desc: "accept invitation with service error", - token: validToken, - data: fmt.Sprintf(`{"domain_id": "%s"}`, validID), - status: http.StatusForbidden, - contentType: contentType, - svcErr: svcerr.ErrAuthorization, - }, - { - desc: "accept invitation with invalid content type", - token: validToken, - data: fmt.Sprintf(`{"domain_id": "%s"}`, validID), - status: http.StatusUnsupportedMediaType, - contentType: "text/plain", - svcErr: nil, - }, - { - desc: "accept invitation with invalid data", - token: validToken, - data: `data`, - status: http.StatusBadRequest, - contentType: contentType, - svcErr: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = authn.Session{UserID: userID, DomainID: domainID} - } - authnCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - repoCall := svc.On("AcceptInvitation", mock.Anything, tc.session, mock.Anything).Return(domains.Invitation{}, tc.svcErr) - req := testRequest{ - client: is.Client(), - method: http.MethodPost, - url: is.URL + "/invitations/accept", - token: tc.token, - contentType: tc.contentType, - body: strings.NewReader(tc.data), - } - - res, err := req.make() - assert.Nil(t, err, tc.desc) - assert.Equal(t, tc.status, res.StatusCode, tc.desc) - repoCall.Unset() - authnCall.Unset() - }) - } -} - -func TestRejectInvitation(t *testing.T) { - is, svc, auth := newDomainsServer() - - cases := []struct { - desc string - token string - session authn.Session - data string - contentType string - status int - svcErr error - authnErr error - }{ - { - desc: "reject invitation with valid request", - token: validToken, - data: fmt.Sprintf(`{"domain_id": "%s"}`, validID), - status: http.StatusNoContent, - contentType: contentType, - svcErr: nil, - }, - { - desc: "reject invitation with invalid token", - token: "", - data: fmt.Sprintf(`{"domain_id": "%s"}`, validID), - status: http.StatusUnauthorized, - contentType: contentType, - svcErr: nil, - }, - { - desc: "reject invitation with unauthorized error", - token: validToken, - data: fmt.Sprintf(`{"domain_id": "%s"}`, "invalid"), - status: http.StatusForbidden, - contentType: contentType, - svcErr: svcerr.ErrAuthorization, - }, - { - desc: "reject invitation with invalid content type", - token: validToken, - data: fmt.Sprintf(`{"domain_id": "%s"}`, validID), - status: http.StatusUnsupportedMediaType, - contentType: "text/plain", - svcErr: nil, - }, - { - desc: "reject invitation with invalid data", - token: validToken, - data: `data`, - status: http.StatusBadRequest, - contentType: contentType, - svcErr: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = authn.Session{UserID: userID, DomainID: domainID} - } - authnCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - repoCall := svc.On("RejectInvitation", mock.Anything, tc.session, mock.Anything).Return(domains.Invitation{}, tc.svcErr) - req := testRequest{ - client: is.Client(), - method: http.MethodPost, - url: is.URL + "/invitations/reject", - token: tc.token, - contentType: tc.contentType, - body: strings.NewReader(tc.data), - } - - res, err := req.make() - assert.Nil(t, err, tc.desc) - assert.Equal(t, tc.status, res.StatusCode, tc.desc) - repoCall.Unset() - authnCall.Unset() - }) - } -} - -type respBody struct { - Err string `json:"error"` - Message string `json:"message"` - Total int `json:"total"` - Permissions []string `json:"permissions"` - ID string `json:"id"` - Tags []string `json:"tags"` - Status domains.Status `json:"status"` -} diff --git a/domains/api/http/requests.go b/domains/api/http/requests.go deleted file mode 100644 index 6ecf27a8b..000000000 --- a/domains/api/http/requests.go +++ /dev/null @@ -1,184 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package http - -import ( - api "github.com/absmach/magistrala/api/http" - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/domains" -) - -const maxLimitSize = 100 - -type createDomainReq struct { - ID string `json:"id,omitempty"` - Name string `json:"name"` - Metadata map[string]any `json:"metadata,omitempty"` - Tags []string `json:"tags,omitempty"` - Route string `json:"route"` -} - -func (req createDomainReq) validate() error { - if req.ID != "" { - return api.ValidateUUID(req.ID) - } - if req.Name == "" { - return apiutil.ErrMissingName - } - if req.Route == "" { - return apiutil.ErrMissingRoute - } - if err := validateRoute(req.Route); err != nil { - return err - } - - return nil -} - -type retrieveDomainRequest struct { - domainID string - roles bool -} - -func (req retrieveDomainRequest) validate() error { - if req.domainID == "" { - return apiutil.ErrMissingID - } - - return nil -} - -type updateDomainReq struct { - domainID string - Name *string `json:"name,omitempty"` - Metadata *domains.Metadata `json:"metadata,omitempty"` - Tags *[]string `json:"tags,omitempty"` -} - -func (req updateDomainReq) validate() error { - if req.domainID == "" { - return apiutil.ErrMissingID - } - - return nil -} - -type listDomainsReq struct { - domains.Page -} - -func (req listDomainsReq) validate() error { - switch req.Order { - case "", api.NameOrder, api.CreatedAtOrder, api.UpdatedAtOrder: - default: - return apiutil.ErrInvalidOrder - } - - if req.Dir != "" && (req.Dir != api.DescDir && req.Dir != api.AscDir) { - return apiutil.ErrInvalidDirection - } - - return nil -} - -type enableDomainReq struct { - domainID string -} - -func (req enableDomainReq) validate() error { - if req.domainID == "" { - return apiutil.ErrMissingID - } - - return nil -} - -type disableDomainReq struct { - domainID string -} - -func (req disableDomainReq) validate() error { - if req.domainID == "" { - return apiutil.ErrMissingID - } - - return nil -} - -type freezeDomainReq struct { - domainID string -} - -func (req freezeDomainReq) validate() error { - if req.domainID == "" { - return apiutil.ErrMissingID - } - - return nil -} - -type sendInvitationReq struct { - InviteeUserID string `json:"invitee_user_id,omitempty"` - RoleID string `json:"role_id,omitempty"` - Resend bool `json:"resend,omitempty"` -} - -func (req *sendInvitationReq) validate() error { - if req.InviteeUserID == "" || req.RoleID == "" { - return apiutil.ErrMissingID - } - - return nil -} - -type listInvitationsReq struct { - domains.InvitationPageMeta -} - -func (req *listInvitationsReq) validate() error { - if req.InvitationPageMeta.Limit > maxLimitSize || req.InvitationPageMeta.Limit < 1 { - return apiutil.ErrLimitSize - } - - return nil -} - -type acceptInvitationReq struct { - DomainID string `json:"domain_id,omitempty"` -} - -func (req *acceptInvitationReq) validate() error { - if req.DomainID == "" { - return apiutil.ErrMissingDomainID - } - - return nil -} - -type deleteInvitationReq struct { - UserID string `json:"user_id"` - domainID string -} - -func (req *deleteInvitationReq) validate() error { - if req.UserID == "" { - return apiutil.ErrMissingID - } - if req.domainID == "" { - return apiutil.ErrMissingDomainID - } - - return nil -} - -func validateRoute(route string) error { - if err := api.ValidateUUID(route); err == nil { - return nil - } - if err := api.ValidateRoute(route); err != nil { - return err - } - - return nil -} diff --git a/domains/api/http/responses.go b/domains/api/http/responses.go deleted file mode 100644 index 132ee4158..000000000 --- a/domains/api/http/responses.go +++ /dev/null @@ -1,205 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package http - -import ( - "net/http" - - "github.com/absmach/magistrala" - "github.com/absmach/magistrala/domains" -) - -var ( - _ magistrala.Response = (*createDomainRes)(nil) - _ magistrala.Response = (*retrieveDomainRes)(nil) - _ magistrala.Response = (*listDomainsRes)(nil) - _ magistrala.Response = (*enableDomainRes)(nil) - _ magistrala.Response = (*disableDomainRes)(nil) - _ magistrala.Response = (*freezeDomainRes)(nil) - _ magistrala.Response = (*sendInvitationRes)(nil) - _ magistrala.Response = (*listInvitationsRes)(nil) - _ magistrala.Response = (*acceptInvitationRes)(nil) - _ magistrala.Response = (*rejectInvitationRes)(nil) - _ magistrala.Response = (*deleteInvitationRes)(nil) -) - -type createDomainRes struct { - domains.Domain -} - -func (res createDomainRes) Code() int { - return http.StatusCreated -} - -func (res createDomainRes) Headers() map[string]string { - return map[string]string{} -} - -func (res createDomainRes) Empty() bool { - return false -} - -type retrieveDomainRes struct { - domains.Domain -} - -func (res retrieveDomainRes) Code() int { - return http.StatusOK -} - -func (res retrieveDomainRes) Headers() map[string]string { - return map[string]string{} -} - -func (res retrieveDomainRes) Empty() bool { - return false -} - -type updateDomainRes struct { - domains.Domain -} - -func (res updateDomainRes) Code() int { - return http.StatusOK -} - -func (res updateDomainRes) Headers() map[string]string { - return map[string]string{} -} - -func (res updateDomainRes) Empty() bool { - return false -} - -type listDomainsRes struct { - domains.DomainsPage -} - -func (res listDomainsRes) Code() int { - return http.StatusOK -} - -func (res listDomainsRes) Headers() map[string]string { - return map[string]string{} -} - -func (res listDomainsRes) Empty() bool { - return false -} - -type enableDomainRes struct{} - -func (res enableDomainRes) Code() int { - return http.StatusOK -} - -func (res enableDomainRes) Headers() map[string]string { - return map[string]string{} -} - -func (res enableDomainRes) Empty() bool { - return true -} - -type disableDomainRes struct{} - -func (res disableDomainRes) Code() int { - return http.StatusOK -} - -func (res disableDomainRes) Headers() map[string]string { - return map[string]string{} -} - -func (res disableDomainRes) Empty() bool { - return true -} - -type freezeDomainRes struct{} - -func (res freezeDomainRes) Code() int { - return http.StatusOK -} - -func (res freezeDomainRes) Headers() map[string]string { - return map[string]string{} -} - -func (res freezeDomainRes) Empty() bool { - return true -} - -type sendInvitationRes struct { - Message string `json:"message"` -} - -func (res sendInvitationRes) Code() int { - return http.StatusCreated -} - -func (res sendInvitationRes) Headers() map[string]string { - return map[string]string{} -} - -func (res sendInvitationRes) Empty() bool { - return true -} - -type listInvitationsRes struct { - domains.InvitationPage `json:",inline"` -} - -func (res listInvitationsRes) Code() int { - return http.StatusOK -} - -func (res listInvitationsRes) Headers() map[string]string { - return map[string]string{} -} - -func (res listInvitationsRes) Empty() bool { - return false -} - -type acceptInvitationRes struct{} - -func (res acceptInvitationRes) Code() int { - return http.StatusNoContent -} - -func (res acceptInvitationRes) Headers() map[string]string { - return map[string]string{} -} - -func (res acceptInvitationRes) Empty() bool { - return true -} - -type deleteInvitationRes struct{} - -func (res deleteInvitationRes) Code() int { - return http.StatusNoContent -} - -func (res deleteInvitationRes) Headers() map[string]string { - return map[string]string{} -} - -func (res deleteInvitationRes) Empty() bool { - return true -} - -type rejectInvitationRes struct{} - -func (res rejectInvitationRes) Code() int { - return http.StatusNoContent -} - -func (res rejectInvitationRes) Headers() map[string]string { - return map[string]string{} -} - -func (res rejectInvitationRes) Empty() bool { - return true -} diff --git a/domains/api/http/transport.go b/domains/api/http/transport.go deleted file mode 100644 index 153f54e7d..000000000 --- a/domains/api/http/transport.go +++ /dev/null @@ -1,139 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package http - -import ( - "log/slog" - "net/http" - - "github.com/absmach/magistrala" - api "github.com/absmach/magistrala/api/http" - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/domains" - smqauthn "github.com/absmach/magistrala/pkg/authn" - roleManagerHttp "github.com/absmach/magistrala/pkg/roles/rolemanager/api" - "github.com/go-chi/chi/v5" - kithttp "github.com/go-kit/kit/transport/http" - "github.com/prometheus/client_golang/prometheus/promhttp" - "go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp" -) - -// MakeHandler returns a HTTP handler for Domains and Invitations API endpoints. -func MakeHandler(svc domains.Service, authn smqauthn.AuthNMiddleware, mux *chi.Mux, logger *slog.Logger, instanceID string, idp magistrala.IDProvider) http.Handler { - opts := []kithttp.ServerOption{ - kithttp.ServerErrorEncoder(apiutil.LoggingErrorEncoder(logger, api.EncodeError)), - } - - d := roleManagerHttp.NewDecoder("domainID") - mux.Route("/domains", func(r chi.Router) { - r.Use(api.RequestIDMiddleware(idp)) - - r.Group(func(r chi.Router) { - r.Use(authn.WithOptions(smqauthn.WithDomainCheck(false)).Middleware()) - r.Post("/", otelhttp.NewHandler(kithttp.NewServer( - createDomainEndpoint(svc), - decodeCreateDomainRequest, - api.EncodeResponse, - opts..., - ), "create_domain").ServeHTTP) - - r.Get("/", otelhttp.NewHandler(kithttp.NewServer( - listDomainsEndpoint(svc), - decodeListDomainRequest, - api.EncodeResponse, - opts..., - ), "list_domains").ServeHTTP) - - roleManagerHttp.EntityAvailableActionsRouter(svc, d, r, opts) - }) - - r.Route("/{domainID}", func(r chi.Router) { - r.Use(authn.Middleware()) - r.Get("/", otelhttp.NewHandler(kithttp.NewServer( - retrieveDomainEndpoint(svc), - decodeRetrieveDomainRequest, - api.EncodeResponse, - opts..., - ), "view_domain").ServeHTTP) - - r.Patch("/", otelhttp.NewHandler(kithttp.NewServer( - updateDomainEndpoint(svc), - decodeUpdateDomainRequest, - api.EncodeResponse, - opts..., - ), "update_domain").ServeHTTP) - - r.Post("/enable", otelhttp.NewHandler(kithttp.NewServer( - enableDomainEndpoint(svc), - decodeEnableDomainRequest, - api.EncodeResponse, - opts..., - ), "enable_domain").ServeHTTP) - - r.Post("/disable", otelhttp.NewHandler(kithttp.NewServer( - disableDomainEndpoint(svc), - decodeDisableDomainRequest, - api.EncodeResponse, - opts..., - ), "disable_domain").ServeHTTP) - - r.Post("/freeze", otelhttp.NewHandler(kithttp.NewServer( - freezeDomainEndpoint(svc), - decodeFreezeDomainRequest, - api.EncodeResponse, - opts..., - ), "freeze_domain").ServeHTTP) - roleManagerHttp.EntityRoleMangerRouter(svc, d, r, opts) - }) - - r.Route("/{domainID}/invitations", func(r chi.Router) { - r.Use(authn.Middleware()) - r.Post("/", otelhttp.NewHandler(kithttp.NewServer( - sendInvitationEndpoint(svc), - decodeSendInvitationReq, - api.EncodeResponse, - opts..., - ), "send_invitation").ServeHTTP) - r.Get("/", otelhttp.NewHandler(kithttp.NewServer( - listDomainInvitationsEndpoint(svc), - decodeListInvitationsReq, - api.EncodeResponse, - opts..., - ), "list_domain_invitations").ServeHTTP) - r.Delete("/", otelhttp.NewHandler(kithttp.NewServer( - deleteInvitationEndpoint(svc), - decodeDeleteInvitationReq, - api.EncodeResponse, - opts..., - ), "delete_invitation").ServeHTTP) - }) - }) - - mux.Route("/invitations", func(r chi.Router) { - r.Use(authn.WithOptions(smqauthn.WithDomainCheck(false)).Middleware()) - r.Get("/", otelhttp.NewHandler(kithttp.NewServer( - listUserInvitationsEndpoint(svc), - decodeListInvitationsReq, - api.EncodeResponse, - opts..., - ), "list_user_invitations").ServeHTTP) - r.Post("/accept", otelhttp.NewHandler(kithttp.NewServer( - acceptInvitationEndpoint(svc), - decodeAcceptInvitationReq, - api.EncodeResponse, - opts..., - ), "accept_invitation").ServeHTTP) - r.Post("/reject", otelhttp.NewHandler(kithttp.NewServer( - rejectInvitationEndpoint(svc), - decodeAcceptInvitationReq, - api.EncodeResponse, - opts..., - ), "reject_invitation").ServeHTTP) - }) - - mux.Get("/health", magistrala.Health("domains", instanceID)) - mux.Handle("/metrics", promhttp.Handler()) - - return mux -} diff --git a/domains/builtinroles.go b/domains/builtinroles.go deleted file mode 100644 index d1ea21267..000000000 --- a/domains/builtinroles.go +++ /dev/null @@ -1,8 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package domains - -import "github.com/absmach/magistrala/pkg/roles" - -const BuiltInRoleAdmin roles.BuiltInRoleName = "admin" diff --git a/domains/cache/doc.go b/domains/cache/doc.go deleted file mode 100644 index 55f6d8093..000000000 --- a/domains/cache/doc.go +++ /dev/null @@ -1,6 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package cache contains the domain concept definitions needed to -// support Magistrala Domains cache service functionality. -package cache diff --git a/domains/cache/domains.go b/domains/cache/domains.go deleted file mode 100644 index af7949a72..000000000 --- a/domains/cache/domains.go +++ /dev/null @@ -1,108 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package cache - -import ( - "context" - "time" - - "github.com/absmach/magistrala/domains" - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/redis/go-redis/v9" -) - -var ( - ErrEmptyDomainID = errors.New("domain ID is empty") - ErrEmptyRoute = errors.New("route is empty") -) - -type domainsCache struct { - client *redis.Client - duration time.Duration -} - -func NewDomainsCache(client *redis.Client, duration time.Duration) domains.Cache { - return &domainsCache{ - client: client, - duration: duration, - } -} - -func (dc *domainsCache) SaveStatus(ctx context.Context, domainID string, status domains.Status) error { - if domainID == "" { - return errors.Wrap(repoerr.ErrCreateEntity, ErrEmptyDomainID) - } - statusString := status.String() - if statusString == domains.Unknown { - return errors.Wrap(repoerr.ErrCreateEntity, svcerr.ErrInvalidStatus) - } - if err := dc.client.Set(ctx, domainID, status.String(), dc.duration).Err(); err != nil { - return errors.Wrap(repoerr.ErrCreateEntity, err) - } - - return nil -} - -func (dc *domainsCache) SaveID(ctx context.Context, route, domainID string) error { - if route == "" { - return errors.Wrap(repoerr.ErrCreateEntity, ErrEmptyRoute) - } - if domainID == "" { - return errors.Wrap(repoerr.ErrCreateEntity, ErrEmptyDomainID) - } - if err := dc.client.Set(ctx, route, domainID, dc.duration).Err(); err != nil { - return errors.Wrap(repoerr.ErrCreateEntity, err) - } - - return nil -} - -func (dc *domainsCache) Status(ctx context.Context, domainID string) (domains.Status, error) { - st, err := dc.client.Get(ctx, domainID).Result() - if err != nil { - return domains.AllStatus, errors.Wrap(repoerr.ErrNotFound, err) - } - status, err := domains.ToStatus(st) - if err != nil { - return domains.AllStatus, errors.Wrap(repoerr.ErrNotFound, err) - } - - return status, nil -} - -func (dc *domainsCache) ID(ctx context.Context, route string) (string, error) { - if route == "" { - return "", errors.Wrap(repoerr.ErrNotFound, ErrEmptyRoute) - } - domainID, err := dc.client.Get(ctx, route).Result() - if err != nil { - return "", errors.Wrap(repoerr.ErrNotFound, err) - } - - return domainID, nil -} - -func (dc *domainsCache) RemoveStatus(ctx context.Context, domainID string) error { - if domainID == "" { - return errors.Wrap(repoerr.ErrRemoveEntity, ErrEmptyDomainID) - } - if err := dc.client.Del(ctx, domainID).Err(); err != nil { - return errors.Wrap(repoerr.ErrRemoveEntity, err) - } - - return nil -} - -func (dc *domainsCache) RemoveID(ctx context.Context, route string) error { - if route == "" { - return errors.Wrap(repoerr.ErrRemoveEntity, ErrEmptyRoute) - } - if err := dc.client.Del(ctx, route).Err(); err != nil { - return errors.Wrap(repoerr.ErrRemoveEntity, err) - } - - return nil -} diff --git a/domains/cache/domains_test.go b/domains/cache/domains_test.go deleted file mode 100644 index 716ca9cc9..000000000 --- a/domains/cache/domains_test.go +++ /dev/null @@ -1,340 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package cache_test - -import ( - "context" - "fmt" - "strings" - "testing" - "time" - - "github.com/absmach/magistrala/domains" - "github.com/absmach/magistrala/domains/cache" - "github.com/absmach/magistrala/internal/testsutil" - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - "github.com/redis/go-redis/v9" - "github.com/stretchr/testify/assert" -) - -var ( - testRoute = "test-route" - nonExistent = "non-existing" -) - -func setupDomainsClient(t *testing.T) domains.Cache { - opts, err := redis.ParseURL(redisURL) - assert.Nil(t, err, fmt.Sprintf("got unexpected error on parsing redis URL: %s", err)) - redisClient := redis.NewClient(opts) - - return cache.NewDomainsCache(redisClient, 10*time.Minute) -} - -func TestSaveStatus(t *testing.T) { - dc := setupDomainsClient(t) - - domainID := testsutil.GenerateUUID(t) - - cases := []struct { - desc string - domainID string - status domains.Status - err error - }{ - { - desc: "Save with enabled status", - domainID: domainID, - status: domains.EnabledStatus, - err: nil, - }, - { - desc: "Save with disabled status", - domainID: testsutil.GenerateUUID(t), - status: domains.DisabledStatus, - err: nil, - }, - { - desc: "Save with frozen status", - domainID: testsutil.GenerateUUID(t), - status: domains.FreezeStatus, - err: nil, - }, - { - desc: "Save with empty domain ID", - domainID: "", - status: domains.EnabledStatus, - err: repoerr.ErrCreateEntity, - }, - { - desc: "Save with all status", - domainID: testsutil.GenerateUUID(t), - status: domains.AllStatus, - err: nil, - }, - { - desc: "Save with invalid status", - domainID: testsutil.GenerateUUID(t), - status: domains.Status(6), - err: repoerr.ErrCreateEntity, - }, - { - desc: "Save the same record", - domainID: domainID, - status: domains.EnabledStatus, - err: nil, - }, - { - desc: "Save client with long id ", - domainID: strings.Repeat("a", 513*1024*1024), - status: domains.EnabledStatus, - err: repoerr.ErrCreateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := dc.SaveStatus(context.Background(), tc.domainID, tc.status) - assert.True(t, errors.Contains(err, tc.err)) - }) - } -} - -func TestSaveID(t *testing.T) { - dc := setupDomainsClient(t) - - route := testRoute - domainID := testsutil.GenerateUUID(t) - - cases := []struct { - desc string - route string - domainID string - err error - }{ - { - desc: "Save domain ID with valid route", - route: route, - domainID: domainID, - err: nil, - }, - { - desc: "Save domain ID with empty route", - route: "", - domainID: domainID, - err: repoerr.ErrCreateEntity, - }, - { - desc: "Save domain ID with empty domain ID", - route: route, - domainID: "", - err: repoerr.ErrCreateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := dc.SaveID(context.Background(), tc.route, tc.domainID) - assert.True(t, errors.Contains(err, tc.err)) - }) - } -} - -func TestStatus(t *testing.T) { - dc := setupDomainsClient(t) - - enabledDomainID := testsutil.GenerateUUID(t) - err := dc.SaveStatus(context.Background(), enabledDomainID, domains.EnabledStatus) - assert.Nil(t, err, fmt.Sprintf("Unexpected error while trying to save: %s", err)) - - disabledDomainID := testsutil.GenerateUUID(t) - err = dc.SaveStatus(context.Background(), disabledDomainID, domains.DisabledStatus) - assert.Nil(t, err, fmt.Sprintf("Unexpected error while trying to save: %s", err)) - - frozenDomainID := testsutil.GenerateUUID(t) - err = dc.SaveStatus(context.Background(), frozenDomainID, domains.FreezeStatus) - assert.Nil(t, err, fmt.Sprintf("Unexpected error while trying to save: %s", err)) - - allDomainID := testsutil.GenerateUUID(t) - err = dc.SaveStatus(context.Background(), allDomainID, domains.AllStatus) - assert.Nil(t, err, fmt.Sprintf("Unexpected error while trying to save: %s", err)) - - cases := []struct { - desc string - domainID string - status domains.Status - err error - }{ - { - desc: "Get domain status from cache for enabled domain", - domainID: enabledDomainID, - status: domains.EnabledStatus, - err: nil, - }, - { - desc: "Get domain status from cache for disabled domain", - domainID: disabledDomainID, - status: domains.DisabledStatus, - err: nil, - }, - { - desc: "Get domain status from cache for frozen domain", - domainID: frozenDomainID, - status: domains.FreezeStatus, - err: nil, - }, - { - desc: "Get domain status from cache for all domain", - domainID: allDomainID, - status: domains.AllStatus, - err: nil, - }, - { - desc: "Get domain status from cache for non existing domain", - domainID: testsutil.GenerateUUID(t), - status: domains.AllStatus, - err: repoerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - status, err := dc.Status(context.Background(), tc.domainID) - assert.True(t, errors.Contains(err, tc.err)) - assert.Equal(t, tc.status, status) - }) - } -} - -func TestID(t *testing.T) { - dc := setupDomainsClient(t) - - route := testRoute - domainID := testsutil.GenerateUUID(t) - err := dc.SaveID(context.Background(), route, domainID) - assert.Nil(t, err, fmt.Sprintf("Unexpected error while trying to save: %s", err)) - - cases := []struct { - desc string - route string - domainID string - err error - }{ - { - desc: "Get domain ID from cache for valid route", - route: route, - domainID: domainID, - err: nil, - }, - { - desc: "Get domain ID from cache for non existing route", - route: nonExistent, - domainID: "", - err: repoerr.ErrNotFound, - }, - { - desc: "Get domain ID from cache with empty route", - route: "", - domainID: "", - err: repoerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - id, err := dc.ID(context.Background(), tc.route) - assert.True(t, errors.Contains(err, tc.err)) - assert.Equal(t, tc.domainID, id) - }) - } -} - -func TestRemoveStatus(t *testing.T) { - dc := setupDomainsClient(t) - - domainID := testsutil.GenerateUUID(t) - err := dc.SaveStatus(context.Background(), domainID, domains.EnabledStatus) - assert.Nil(t, err, fmt.Sprintf("Unexpected error while trying to save: %s", err)) - - cases := []struct { - desc string - domainID string - err error - }{ - { - desc: "Remove domain from cache", - domainID: domainID, - err: nil, - }, - { - desc: "Remove domain from cache with empty domain ID", - domainID: "", - err: repoerr.ErrRemoveEntity, - }, - { - desc: "Remove non existing domain from cache", - domainID: testsutil.GenerateUUID(t), - err: nil, - }, - { - desc: "Remove domain from cache with long id", - domainID: strings.Repeat("a", 513*1024*1024), - err: repoerr.ErrRemoveEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := dc.RemoveStatus(context.Background(), tc.domainID) - assert.True(t, errors.Contains(err, tc.err)) - if err == nil { - _, err = dc.Status(context.Background(), tc.domainID) - assert.True(t, errors.Contains(err, repoerr.ErrNotFound)) - } - }) - } -} - -func TestRemoveID(t *testing.T) { - dc := setupDomainsClient(t) - - route := testRoute - domainID := testsutil.GenerateUUID(t) - err := dc.SaveID(context.Background(), route, domainID) - assert.Nil(t, err, fmt.Sprintf("Unexpected error while trying to save: %s", err)) - - cases := []struct { - desc string - route string - err error - }{ - { - desc: "Remove domain ID from cache", - route: route, - err: nil, - }, - { - desc: "Remove domain ID from cache with empty route", - route: "", - err: repoerr.ErrRemoveEntity, - }, - { - desc: "Remove non existing domain ID from cache", - route: nonExistent, - err: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := dc.RemoveID(context.Background(), tc.route) - assert.True(t, errors.Contains(err, tc.err)) - if err == nil { - id, err := dc.ID(context.Background(), tc.route) - assert.True(t, errors.Contains(err, repoerr.ErrNotFound)) - assert.Equal(t, "", id) - } - }) - } -} diff --git a/domains/cache/setup_test.go b/domains/cache/setup_test.go deleted file mode 100644 index 716f0672c..000000000 --- a/domains/cache/setup_test.go +++ /dev/null @@ -1,61 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package cache_test - -import ( - "context" - "fmt" - "log" - "os" - "testing" - - "github.com/ory/dockertest/v3" - "github.com/ory/dockertest/v3/docker" - "github.com/redis/go-redis/v9" -) - -var ( - redisClient *redis.Client - redisURL string -) - -func TestMain(m *testing.M) { - pool, err := dockertest.NewPool("") - if err != nil { - log.Fatalf("Could not connect to docker: %s", err) - } - - container, err := pool.RunWithOptions(&dockertest.RunOptions{ - Repository: "redis", - Tag: "7.2.4-alpine", - }, func(config *docker.HostConfig) { - config.AutoRemove = true - config.RestartPolicy = docker.RestartPolicy{Name: "no"} - }) - if err != nil { - log.Fatalf("Could not start container: %s", err) - } - - redisURL = fmt.Sprintf("redis://localhost:%s/0", container.GetPort("6379/tcp")) - opts, err := redis.ParseURL(redisURL) - if err != nil { - log.Fatalf("Could not parse redis URL: %s", err) - } - - if err := pool.Retry(func() error { - redisClient = redis.NewClient(opts) - - return redisClient.Ping(context.Background()).Err() - }); err != nil { - log.Fatalf("Could not connect to docker: %s", err) - } - - code := m.Run() - - if err := pool.Purge(container); err != nil { - log.Fatalf("Could not purge container: %s", err) - } - - os.Exit(code) -} diff --git a/domains/domains.go b/domains/domains.go deleted file mode 100644 index 24a1bd164..000000000 --- a/domains/domains.go +++ /dev/null @@ -1,325 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package domains - -import ( - "context" - "encoding/json" - "strings" - "time" - - "github.com/absmach/magistrala/pkg/authn" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/pkg/roles" -) - -// Status represents Domain status. -type Status uint8 - -// Possible Domain status values. -const ( - // EnabledStatus represents enabled Domain. - EnabledStatus Status = iota - // DisabledStatus represents disabled Domain. - DisabledStatus - // FreezeStatus represents domain is in freezed state. - FreezeStatus - // DeletedStatus represents domain is in deleted state. - DeletedStatus - - // AllStatus is used for querying purposes to list Domains irrespective - // of their status - enabled, disabled, freezed, deleting. It is never stored in the - // database as the actual domain status and should always be the larger than freeze status - // value in this enumeration. - AllStatus -) - -// String representation of the possible status values. -const ( - Disabled = "disabled" - Enabled = "enabled" - Freezed = "freezed" - Deleted = "deleted" - All = "all" - Unknown = "unknown" -) - -// String converts client/group status to string literal. -func (s Status) String() string { - switch s { - case DisabledStatus: - return Disabled - case EnabledStatus: - return Enabled - case AllStatus: - return All - case FreezeStatus: - return Freezed - case DeletedStatus: - return Deleted - default: - return Unknown - } -} - -// ToStatus converts string value to a valid Domain status. -func ToStatus(status string) (Status, error) { - switch status { - case "", Enabled: - return EnabledStatus, nil - case Disabled: - return DisabledStatus, nil - case Freezed: - return FreezeStatus, nil - case Deleted: - return DeletedStatus, nil - case All: - return AllStatus, nil - } - return Status(0), svcerr.ErrInvalidStatus -} - -// Custom Marshaller for Domains status. -func (s Status) MarshalJSON() ([]byte, error) { - return json.Marshal(s.String()) -} - -// Custom Unmarshaler for Domains status. -func (s *Status) UnmarshalJSON(data []byte) error { - str := strings.Trim(string(data), "\"") - val, err := ToStatus(str) - *s = val - return err -} - -// Metadata represents arbitrary JSON. -type Metadata map[string]any - -type DomainReq struct { - Name *string `json:"name,omitempty"` - Metadata *Metadata `json:"metadata,omitempty"` - Tags *[]string `json:"tags,omitempty"` - Status *Status `json:"status,omitempty"` - UpdatedBy *string `json:"updated_by,omitempty"` - UpdatedAt *time.Time `json:"updated_at,omitempty"` -} - -type Domain struct { - ID string `json:"id"` - Name string `json:"name"` - Metadata Metadata `json:"metadata,omitempty"` - Tags []string `json:"tags,omitempty"` - Route string `json:"route,omitempty"` - Status Status `json:"status"` - RoleID string `json:"role_id,omitempty"` - RoleName string `json:"role_name,omitempty"` - Actions []string `json:"actions,omitempty"` - CreatedBy string `json:"created_by,omitempty"` - CreatedAt time.Time `json:"created_at"` - UpdatedBy string `json:"updated_by,omitempty"` - UpdatedAt time.Time `json:"updated_at,omitempty"` - MemberID string `json:"member_id,omitempty"` - Roles []roles.MemberRoleActions `json:"roles,omitempty"` -} - -type Operator uint8 - -const ( - OrOp Operator = iota - AndOp -) - -type TagsQuery struct { - Elements []string - Operator Operator -} - -func ToTagsQuery(s string) TagsQuery { - switch { - case strings.Contains(s, "+"): - elements := strings.Split(s, "+") - for i := range elements { - elements[i] = strings.TrimSpace(elements[i]) - } - return TagsQuery{Elements: elements, Operator: AndOp} - case strings.Contains(s, ","): - elements := strings.Split(s, ",") - for i := range elements { - elements[i] = strings.TrimSpace(elements[i]) - } - return TagsQuery{Elements: elements, Operator: OrOp} - default: - return TagsQuery{Elements: []string{s}, Operator: OrOp} - } -} - -type Page struct { - Total uint64 `json:"total"` - Offset uint64 `json:"offset"` - Limit uint64 `json:"limit"` - OnlyTotal bool `json:"only_total"` - Name string `json:"name,omitempty"` - Order string `json:"-"` - Dir string `json:"-"` - Metadata Metadata `json:"metadata,omitempty"` - Tags TagsQuery `json:"tags,omitempty"` - RoleName string `json:"role_name,omitempty"` - RoleID string `json:"role_id,omitempty"` - Actions []string `json:"actions,omitempty"` - Status Status `json:"status,omitempty"` - ID string `json:"id,omitempty"` - IDs []string `json:"-"` - Identity string `json:"identity,omitempty"` - UserID string `json:"user_id,omitempty"` - CreatedFrom time.Time `json:"created_from,omitempty"` - CreatedTo time.Time `json:"created_to,omitempty"` -} - -type DomainsPage struct { - Total uint64 `json:"total"` - Offset uint64 `json:"offset,omitempty"` - Limit uint64 `json:"limit,omitempty"` - Domains []Domain `json:"domains,omitempty"` -} - -func (page DomainsPage) MarshalJSON() ([]byte, error) { - type Alias DomainsPage - a := struct { - Alias - }{ - Alias: Alias(page), - } - - if a.Domains == nil { - a.Domains = make([]Domain, 0) - } - - return json.Marshal(a) -} - -type Service interface { - // CreateDomain creates a new domain. - CreateDomain(ctx context.Context, sesssion authn.Session, d Domain) (Domain, []roles.RoleProvision, error) - - // RetrieveDomain retrieves a domain specified by the provided ID. - RetrieveDomain(ctx context.Context, sesssion authn.Session, id string, withRoles bool) (Domain, error) - - // UpdateDomain updates the domain specified by the provided ID. - UpdateDomain(ctx context.Context, sesssion authn.Session, id string, d DomainReq) (Domain, error) - - // EnableDomain enables the domain specified by the provided ID. - EnableDomain(ctx context.Context, sesssion authn.Session, id string) (Domain, error) - - // DisableDomain disables the domain specified by the provided ID. - // Only platform administrators and domain admins can disable domains. - DisableDomain(ctx context.Context, sesssion authn.Session, id string) (Domain, error) - - // FreezeDomain freezes the domain specified by the provided ID. - // Only platform administrators can freeze domains. - FreezeDomain(ctx context.Context, sesssion authn.Session, id string) (Domain, error) - - // ListDomains returns a list of domains. - ListDomains(ctx context.Context, sesssion authn.Session, page Page) (DomainsPage, error) - - // SendInvitation sends an invitation to the given user. - // Only domain administrators and platform administrators can send invitations. - // Returns the enriched invitation with domain and role names populated. - SendInvitation(ctx context.Context, session authn.Session, invitation Invitation) (Invitation, error) - - // ListInvitations returns a list of invitations. - // By default, it will list invitations the current user has received. - ListInvitations(ctx context.Context, session authn.Session, page InvitationPageMeta) (invitations InvitationPage, err error) - - // ListDomainInvitations returns a list of invitations for the domain. - // People who can list invitations are: - // - platform administrators can list all invitations - // - domain administrators can list invitations for their domain - ListDomainInvitations(ctx context.Context, session authn.Session, page InvitationPageMeta) (invitations InvitationPage, err error) - - // AcceptInvitation accepts an invitation by adding the user to the domain. - AcceptInvitation(ctx context.Context, session authn.Session, domainID string) (invitation Invitation, err error) - - // DeleteInvitation deletes an invitation. - // People who can delete invitations are: - // - the invited user: they can delete their own invitations - // - the user who sent the invitation - // - domain administrators - // - platform administrators - DeleteInvitation(ctx context.Context, session authn.Session, inviteeUserID, domainID string) (err error) - - // RejectInvitation rejects an invitation. - // People who can reject invitations are: - // - the invited user: they can reject their own invitations - RejectInvitation(ctx context.Context, session authn.Session, domainID string) (Invitation, error) - - roles.RoleManager -} - -// Repository specifies Domain persistence API. -type Repository interface { - // SaveDomain creates db insert transaction for the given domain. - SaveDomain(ctx context.Context, d Domain) (Domain, error) - - // RetrieveDomainByIDWithRoles retrieves a domain by its unique ID along with member roles. - RetrieveDomainByIDWithRoles(ctx context.Context, id string, memberID string) (Domain, error) - - // RetrieveDomainByID retrieves a domain by its unique ID. - RetrieveDomainByID(ctx context.Context, id string) (Domain, error) - - // RetrieveDomainByRoute retrieves a domain by its unique route. - RetrieveDomainByRoute(ctx context.Context, route string) (Domain, error) - - // RetrieveAllDomainsByIDs retrieves for given Domain IDs. - RetrieveAllDomainsByIDs(ctx context.Context, pm Page) (DomainsPage, error) - - // UpdateDomain updates the domain name and metadata. - UpdateDomain(ctx context.Context, id string, d DomainReq) (Domain, error) - - // DeleteDomain deletes the domain. - DeleteDomain(ctx context.Context, id string) error - - // ListDomains list all the domains - ListDomains(ctx context.Context, pm Page) (DomainsPage, error) - - // CreateInvitation creates an invitation. - SaveInvitation(ctx context.Context, invitation Invitation) (err error) - - // RetrieveInvitation retrieves an invitation. - RetrieveInvitation(ctx context.Context, userID, domainID string) (Invitation, error) - - // RetrieveAllInvitations retrieves all invitations. - RetrieveAllInvitations(ctx context.Context, page InvitationPageMeta) (invitations InvitationPage, err error) - - // UpdateConfirmation updates an invitation by setting the confirmation time. - UpdateConfirmation(ctx context.Context, invitation Invitation) (err error) - - // UpdateRejection updates an invitation by setting the rejection time. - UpdateRejection(ctx context.Context, invitation Invitation) (err error) - - // DeleteUsersInvitations deletes invitation to a provided domain for users with provided user IDs. - DeleteUsersInvitations(ctx context.Context, domainID string, userID ...string) (err error) - - roles.Repository -} - -// Cache contains domains caching interface. -type Cache interface { - // Save stores pair domain status and domain id. - SaveStatus(ctx context.Context, domainID string, status Status) error - - // SaveID stores pair route and domain id. - SaveID(ctx context.Context, route, domainID string) error - - // Status returns domain status for given domain ID. - Status(ctx context.Context, domainID string) (Status, error) - - // ID returns domain ID for given route. - ID(ctx context.Context, route string) (string, error) - - // RemoveStatus removes domain ID and status pair from cache. - RemoveStatus(ctx context.Context, domainID string) error - - // RemoveID removes domain route and ID pair from cache. - RemoveID(ctx context.Context, route string) error -} diff --git a/domains/domains_test.go b/domains/domains_test.go deleted file mode 100644 index 7e55ff3b2..000000000 --- a/domains/domains_test.go +++ /dev/null @@ -1,186 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package domains_test - -import ( - "testing" - - "github.com/absmach/magistrala/domains" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/stretchr/testify/assert" -) - -func TestStatusString(t *testing.T) { - cases := []struct { - desc string - status domains.Status - expected string - }{ - { - desc: "Enabled", - status: domains.EnabledStatus, - expected: "enabled", - }, - { - desc: "Disabled", - status: domains.DisabledStatus, - expected: "disabled", - }, - { - desc: "Freezed", - status: domains.FreezeStatus, - expected: "freezed", - }, - { - desc: "All", - status: domains.AllStatus, - expected: "all", - }, - { - desc: "Unknown", - status: domains.Status(100), - expected: "unknown", - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - got := tc.status.String() - assert.Equal(t, tc.expected, got, "String() = %v, expected %v", got, tc.expected) - }) - } -} - -func TestToStatus(t *testing.T) { - cases := []struct { - desc string - status string - expetcted domains.Status - err error - }{ - { - desc: "Enabled", - status: "enabled", - expetcted: domains.EnabledStatus, - err: nil, - }, - { - desc: "Disabled", - status: "disabled", - expetcted: domains.DisabledStatus, - err: nil, - }, - { - desc: "Freezed", - status: "freezed", - expetcted: domains.FreezeStatus, - err: nil, - }, - { - desc: "All", - status: "all", - expetcted: domains.AllStatus, - err: nil, - }, - { - desc: "Unknown", - status: "unknown", - expetcted: domains.Status(0), - err: svcerr.ErrInvalidStatus, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - got, err := domains.ToStatus(tc.status) - assert.Equal(t, tc.err, err, "ToStatus() error = %v, expected %v", err, tc.err) - assert.Equal(t, tc.expetcted, got, "ToStatus() = %v, expected %v", got, tc.expetcted) - }) - } -} - -func TestStatusMarshalJSON(t *testing.T) { - cases := []struct { - desc string - expected []byte - status domains.Status - err error - }{ - { - desc: "Enabled", - expected: []byte(`"enabled"`), - status: domains.EnabledStatus, - err: nil, - }, - { - desc: "Disabled", - expected: []byte(`"disabled"`), - status: domains.DisabledStatus, - err: nil, - }, - { - desc: "All", - expected: []byte(`"all"`), - status: domains.AllStatus, - err: nil, - }, - { - desc: "Unknown", - expected: []byte(`"unknown"`), - status: domains.Status(100), - err: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - got, err := tc.status.MarshalJSON() - assert.Equal(t, tc.err, err, "MarshalJSON() error = %v, expected %v", err, tc.err) - assert.Equal(t, tc.expected, got, "MarshalJSON() = %v, expected %v", got, tc.expected) - }) - } -} - -func TestStatusUnmarshalJSON(t *testing.T) { - cases := []struct { - desc string - expected domains.Status - status []byte - err error - }{ - { - desc: "Enabled", - expected: domains.EnabledStatus, - status: []byte(`"enabled"`), - err: nil, - }, - { - desc: "Disabled", - expected: domains.DisabledStatus, - status: []byte(`"disabled"`), - err: nil, - }, - { - desc: "All", - expected: domains.AllStatus, - status: []byte(`"all"`), - err: nil, - }, - { - desc: "Unknown", - expected: domains.Status(0), - status: []byte(`"unknown"`), - err: svcerr.ErrInvalidStatus, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - var s domains.Status - err := s.UnmarshalJSON(tc.status) - assert.Equal(t, tc.err, err, "UnmarshalJSON() error = %v, expected %v", err, tc.err) - assert.Equal(t, tc.expected, s, "UnmarshalJSON() = %v, expected %v", s, tc.expected) - }) - } -} diff --git a/domains/events/doc.go b/domains/events/doc.go deleted file mode 100644 index a115b5f92..000000000 --- a/domains/events/doc.go +++ /dev/null @@ -1,6 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package events provides the domain concept definitions needed to -// support Magistrala auth service functionality. -package events diff --git a/domains/events/events.go b/domains/events/events.go deleted file mode 100644 index 0c219e2db..000000000 --- a/domains/events/events.go +++ /dev/null @@ -1,454 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package events - -import ( - "time" - - "github.com/absmach/magistrala/domains" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/events" - "github.com/absmach/magistrala/pkg/roles" -) - -const ( - domainPrefix = "domain." - domainCreate = domainPrefix + "create" - domainRetrieve = domainPrefix + "retrieve" - domainUpdate = domainPrefix + "update" - domainEnable = domainPrefix + "enable" - domainDisable = domainPrefix + "disable" - domainFreeze = domainPrefix + "freeze" - domainList = domainPrefix + "list" - invitationPrefix = "invitation." - invitationSend = invitationPrefix + "send" - invitationAccept = invitationPrefix + "accept" - invitationReject = invitationPrefix + "reject" - invitationList = invitationPrefix + "list" - invitationListDomain = invitationPrefix + "list_domain" - invitationDelete = invitationPrefix + "delete" -) - -var ( - _ events.Event = (*createDomainEvent)(nil) - _ events.Event = (*retrieveDomainEvent)(nil) - _ events.Event = (*updateDomainEvent)(nil) - _ events.Event = (*enableDomainEvent)(nil) - _ events.Event = (*disableDomainEvent)(nil) - _ events.Event = (*freezeDomainEvent)(nil) - _ events.Event = (*listDomainsEvent)(nil) - _ events.Event = (*sendInvitationEvent)(nil) - _ events.Event = (*listInvitationsEvent)(nil) - _ events.Event = (*listDomainInvitationsEvent)(nil) - _ events.Event = (*acceptInvitationEvent)(nil) - _ events.Event = (*rejectInvitationEvent)(nil) - _ events.Event = (*deleteInvitationEvent)(nil) -) - -type createDomainEvent struct { - domains.Domain - rolesProvisioned []roles.RoleProvision - authn.Session - requestID string -} - -func (cde createDomainEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": domainCreate, - "id": cde.ID, - "route": cde.Route, - "status": cde.Status.String(), - "created_at": cde.CreatedAt, - "created_by": cde.CreatedBy, - "roles_provisioned": cde.rolesProvisioned, - "user_id": cde.UserID, - "token_type": cde.Type.String(), - "super_admin": cde.SuperAdmin, - "request_id": cde.requestID, - } - - if cde.Name != "" { - val["name"] = cde.Name - } - if len(cde.Tags) > 0 { - val["tags"] = cde.Tags - } - if cde.Metadata != nil { - val["metadata"] = cde.Metadata - } - - return val, nil -} - -type retrieveDomainEvent struct { - domains.Domain - authn.Session - requestID string -} - -func (rde retrieveDomainEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": domainRetrieve, - "id": rde.ID, - "route": rde.Route, - "status": rde.Status.String(), - "created_at": rde.CreatedAt, - "user_id": rde.UserID, - "token_type": rde.Type.String(), - "super_admin": rde.SuperAdmin, - "request_id": rde.requestID, - } - - if rde.Name != "" { - val["name"] = rde.Name - } - if len(rde.Tags) > 0 { - val["tags"] = rde.Tags - } - if rde.Metadata != nil { - val["metadata"] = rde.Metadata - } - - if !rde.UpdatedAt.IsZero() { - val["updated_at"] = rde.UpdatedAt - } - if rde.UpdatedBy != "" { - val["updated_by"] = rde.UpdatedBy - } - return val, nil -} - -type updateDomainEvent struct { - domain domains.Domain - Session authn.Session - requestID string -} - -func (ude updateDomainEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": domainUpdate, - "id": ude.domain.ID, - "route": ude.domain.Route, - "status": ude.domain.Status.String(), - "created_at": ude.domain.CreatedAt, - "created_by": ude.domain.CreatedBy, - "updated_at": ude.domain.UpdatedAt, - "updated_by": ude.domain.UpdatedBy, - "user_id": ude.Session.UserID, - "token_type": ude.Session.Type.String(), - "super_admin": ude.Session.SuperAdmin, - "request_id": ude.requestID, - } - - if ude.domain.Name != "" { - val["name"] = ude.domain.Name - } - if len(ude.domain.Tags) > 0 { - val["tags"] = ude.domain.Tags - } - if ude.domain.Metadata != nil { - val["metadata"] = ude.domain.Metadata - } - - return val, nil -} - -type enableDomainEvent struct { - domainID string - updatedAt time.Time - updatedBy string - authn.Session - requestID string -} - -func (cdse enableDomainEvent) Encode() (map[string]any, error) { - return map[string]any{ - "operation": domainEnable, - "id": cdse.domainID, - "updated_at": cdse.updatedAt, - "updated_by": cdse.updatedBy, - "user_id": cdse.UserID, - "token_type": cdse.Type.String(), - "super_admin": cdse.SuperAdmin, - "request_id": cdse.requestID, - }, nil -} - -type disableDomainEvent struct { - domainID string - updatedAt time.Time - updatedBy string - authn.Session - requestID string -} - -func (cdse disableDomainEvent) Encode() (map[string]any, error) { - return map[string]any{ - "operation": domainDisable, - "id": cdse.domainID, - "updated_at": cdse.updatedAt, - "updated_by": cdse.updatedBy, - "user_id": cdse.UserID, - "token_type": cdse.Type.String(), - "super_admin": cdse.SuperAdmin, - "request_id": cdse.requestID, - }, nil -} - -type freezeDomainEvent struct { - domainID string - updatedAt time.Time - updatedBy string - authn.Session - requestID string -} - -func (cdse freezeDomainEvent) Encode() (map[string]any, error) { - return map[string]any{ - "operation": domainFreeze, - "id": cdse.domainID, - "updated_at": cdse.updatedAt, - "updated_by": cdse.updatedBy, - "user_id": cdse.UserID, - "token_type": cdse.Type.String(), - "super_admin": cdse.SuperAdmin, - "request_id": cdse.requestID, - }, nil -} - -type listDomainsEvent struct { - domains.Page - total uint64 - userID string - tokenType string - superAdmin bool - requestID string -} - -func (lde listDomainsEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": domainList, - "total": lde.total, - "offset": lde.Offset, - "limit": lde.Limit, - "user_id": lde.userID, - "token_type": lde.tokenType, - "super_admin": lde.superAdmin, - "request_id": lde.requestID, - } - - if lde.Name != "" { - val["name"] = lde.Name - } - if lde.Order != "" { - val["order"] = lde.Order - } - if lde.Dir != "" { - val["dir"] = lde.Dir - } - if lde.Metadata != nil { - val["metadata"] = lde.Metadata - } - if len(lde.Tags.Elements) > 0 { - val["tag"] = lde.Tags.Elements - } - if lde.RoleID != "" { - val["role_id"] = lde.RoleID - } - if lde.RoleName != "" { - val["role_name"] = lde.RoleName - } - if len(lde.Actions) != 0 { - val["actions"] = lde.Actions - } - if lde.Status.String() != "" { - val["status"] = lde.Status.String() - } - if lde.ID != "" { - val["id"] = lde.ID - } - if len(lde.IDs) > 0 { - val["ids"] = lde.IDs - } - if lde.Identity != "" { - val["identity"] = lde.Identity - } - if lde.UserID != "" { - val["user_id"] = lde.UserID - } - - return val, nil -} - -type sendInvitationEvent struct { - invitation domains.Invitation - session authn.Session - requestID string -} - -func (sie sendInvitationEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": invitationSend, - "invitee_user_id": sie.invitation.InviteeUserID, - "domain_id": sie.invitation.DomainID, - "invited_by": sie.session.UserID, - "role_id": sie.invitation.RoleID, - "token_type": sie.session.Type.String(), - "super_admin": sie.session.SuperAdmin, - "request_id": sie.requestID, - } - - if sie.invitation.DomainName != "" { - val["domain_name"] = sie.invitation.DomainName - } - if sie.invitation.RoleName != "" { - val["role_name"] = sie.invitation.RoleName - } - - return val, nil -} - -type listInvitationsEvent struct { - domains.InvitationPageMeta - session authn.Session - requestID string -} - -func (lie listInvitationsEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": invitationList, - "offset": lie.Offset, - "limit": lie.Limit, - "user_id": lie.session.UserID, - "token_type": lie.session.Type.String(), - "request_id": lie.requestID, - } - - if lie.InvitedBy != "" { - val["invited_by"] = lie.InvitedBy - } - if lie.InviteeUserID != "" { - val["invitee_user_id"] = lie.InviteeUserID - } - if lie.DomainID != "" { - val["domain_id"] = lie.DomainID - } - if lie.RoleID != "" { - val["role_id"] = lie.RoleID - } - if lie.State.String() != domains.UnknownState { - val["state"] = lie.State.String() - } - - return val, nil -} - -type listDomainInvitationsEvent struct { - domains.InvitationPageMeta - session authn.Session - requestID string -} - -func (lie listDomainInvitationsEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": invitationListDomain, - "offset": lie.Offset, - "limit": lie.Limit, - "domain_id": lie.session.DomainID, - "token_type": lie.session.Type.String(), - "super_admin": lie.session.SuperAdmin, - "request_id": lie.requestID, - } - - if lie.InvitedBy != "" { - val["invited_by"] = lie.InvitedBy - } - if lie.InviteeUserID != "" { - val["invitee_user_id"] = lie.InviteeUserID - } - if lie.RoleID != "" { - val["role_id"] = lie.RoleID - } - if lie.State.String() != domains.UnknownState { - val["state"] = lie.State.String() - } - - return val, nil -} - -type acceptInvitationEvent struct { - invitation domains.Invitation - session authn.Session - requestID string -} - -func (aie acceptInvitationEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": invitationAccept, - "domain_id": aie.invitation.DomainID, - "invitee_user_id": aie.session.UserID, - "invited_by": aie.invitation.InvitedBy, - "role_id": aie.invitation.RoleID, - "token_type": aie.session.Type.String(), - "super_admin": aie.session.SuperAdmin, - "request_id": aie.requestID, - } - - if aie.invitation.DomainName != "" { - val["domain_name"] = aie.invitation.DomainName - } - if aie.invitation.RoleName != "" { - val["role_name"] = aie.invitation.RoleName - } - - return val, nil -} - -type rejectInvitationEvent struct { - invitation domains.Invitation - session authn.Session - requestID string -} - -func (rie rejectInvitationEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": invitationReject, - "domain_id": rie.invitation.DomainID, - "invitee_user_id": rie.session.UserID, - "invited_by": rie.invitation.InvitedBy, - "role_id": rie.invitation.RoleID, - "token_type": rie.session.Type.String(), - "super_admin": rie.session.SuperAdmin, - "request_id": rie.requestID, - } - - if rie.invitation.DomainName != "" { - val["domain_name"] = rie.invitation.DomainName - } - if rie.invitation.RoleName != "" { - val["role_name"] = rie.invitation.RoleName - } - - return val, nil -} - -type deleteInvitationEvent struct { - inviteeUserID string - domainID string - session authn.Session - requestID string -} - -func (die deleteInvitationEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": invitationDelete, - "invitee_user_id": die.inviteeUserID, - "domain_id": die.domainID, - "token_type": die.session.Type.String(), - "super_admin": die.session.SuperAdmin, - "request_id": die.requestID, - } - - return val, nil -} diff --git a/domains/events/streams.go b/domains/events/streams.go deleted file mode 100644 index 676713c11..000000000 --- a/domains/events/streams.go +++ /dev/null @@ -1,314 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package events - -import ( - "context" - - "github.com/absmach/magistrala/domains" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/events" - "github.com/absmach/magistrala/pkg/events/store" - "github.com/absmach/magistrala/pkg/roles" - rmEvents "github.com/absmach/magistrala/pkg/roles/rolemanager/events" - "github.com/go-chi/chi/v5/middleware" -) - -const ( - magistralaPrefix = "magistrala." - createStream = magistralaPrefix + domainCreate - retrieveStream = magistralaPrefix + domainRetrieve - updateStream = magistralaPrefix + domainUpdate - enableStream = magistralaPrefix + domainEnable - disableStream = magistralaPrefix + domainDisable - freezeStream = magistralaPrefix + domainFreeze - listStream = magistralaPrefix + domainList - sendInvitationStream = magistralaPrefix + invitationSend - acceptInvitationStream = magistralaPrefix + invitationAccept - rejectInvitationStream = magistralaPrefix + invitationReject - listInvitationsStream = magistralaPrefix + invitationList - listDomainInvitationsStream = magistralaPrefix + invitationListDomain - deleteInvitationStream = magistralaPrefix + invitationDelete -) - -var _ domains.Service = (*eventStore)(nil) - -type eventStore struct { - events.Publisher - svc domains.Service - rmEvents.RoleManagerEventStore -} - -// NewEventStoreMiddleware returns wrapper around auth service that sends -// events to event store. -func NewEventStoreMiddleware(ctx context.Context, svc domains.Service, url string) (domains.Service, error) { - publisher, err := store.NewPublisher(ctx, url, "domains-es-pub") - if err != nil { - return nil, err - } - - res := rmEvents.NewRoleManagerEventStore("domains", domainPrefix, svc, publisher) - - return &eventStore{ - svc: svc, - Publisher: publisher, - RoleManagerEventStore: res, - }, nil -} - -func (es *eventStore) CreateDomain(ctx context.Context, session authn.Session, domain domains.Domain) (domains.Domain, []roles.RoleProvision, error) { - domain, rps, err := es.svc.CreateDomain(ctx, session, domain) - if err != nil { - return domain, rps, err - } - - event := createDomainEvent{ - Domain: domain, - rolesProvisioned: rps, - Session: session, - requestID: middleware.GetReqID(ctx), - } - - if err := es.Publish(ctx, createStream, event); err != nil { - return domain, rps, err - } - - return domain, rps, nil -} - -func (es *eventStore) RetrieveDomain(ctx context.Context, session authn.Session, id string, withRoles bool) (domains.Domain, error) { - domain, err := es.svc.RetrieveDomain(ctx, session, id, withRoles) - if err != nil { - return domain, err - } - - event := retrieveDomainEvent{ - domain, - session, - middleware.GetReqID(ctx), - } - - if err := es.Publish(ctx, retrieveStream, event); err != nil { - return domain, err - } - - return domain, nil -} - -func (es *eventStore) UpdateDomain(ctx context.Context, session authn.Session, id string, d domains.DomainReq) (domains.Domain, error) { - domain, err := es.svc.UpdateDomain(ctx, session, id, d) - if err != nil { - return domain, err - } - - event := updateDomainEvent{ - domain: domain, - Session: session, - requestID: middleware.GetReqID(ctx), - } - - if err := es.Publish(ctx, updateStream, event); err != nil { - return domain, err - } - - return domain, nil -} - -func (es *eventStore) EnableDomain(ctx context.Context, session authn.Session, id string) (domains.Domain, error) { - domain, err := es.svc.EnableDomain(ctx, session, id) - if err != nil { - return domain, err - } - - event := enableDomainEvent{ - domainID: id, - updatedAt: domain.UpdatedAt, - updatedBy: domain.UpdatedBy, - Session: session, - requestID: middleware.GetReqID(ctx), - } - - if err := es.Publish(ctx, enableStream, event); err != nil { - return domain, err - } - - return domain, nil -} - -func (es *eventStore) DisableDomain(ctx context.Context, session authn.Session, id string) (domains.Domain, error) { - domain, err := es.svc.DisableDomain(ctx, session, id) - if err != nil { - return domain, err - } - - event := disableDomainEvent{ - domainID: id, - updatedAt: domain.UpdatedAt, - updatedBy: domain.UpdatedBy, - Session: session, - requestID: middleware.GetReqID(ctx), - } - - if err := es.Publish(ctx, disableStream, event); err != nil { - return domain, err - } - - return domain, nil -} - -func (es *eventStore) FreezeDomain(ctx context.Context, session authn.Session, id string) (domains.Domain, error) { - domain, err := es.svc.FreezeDomain(ctx, session, id) - if err != nil { - return domain, err - } - - event := freezeDomainEvent{ - domainID: id, - updatedAt: domain.UpdatedAt, - updatedBy: domain.UpdatedBy, - Session: session, - requestID: middleware.GetReqID(ctx), - } - - if err := es.Publish(ctx, freezeStream, event); err != nil { - return domain, err - } - - return domain, nil -} - -func (es *eventStore) ListDomains(ctx context.Context, session authn.Session, p domains.Page) (domains.DomainsPage, error) { - dp, err := es.svc.ListDomains(ctx, session, p) - if err != nil { - return dp, err - } - - event := listDomainsEvent{ - Page: p, - total: dp.Total, - userID: session.UserID, - tokenType: session.Type.String(), - superAdmin: session.SuperAdmin, - requestID: middleware.GetReqID(ctx), - } - - if err := es.Publish(ctx, listStream, event); err != nil { - return dp, err - } - - return dp, nil -} - -func (es *eventStore) SendInvitation(ctx context.Context, session authn.Session, invitation domains.Invitation) (domains.Invitation, error) { - inv, err := es.svc.SendInvitation(ctx, session, invitation) - if err != nil { - return domains.Invitation{}, err - } - - event := sendInvitationEvent{ - invitation: inv, - session: session, - requestID: middleware.GetReqID(ctx), - } - - if err := es.Publish(ctx, sendInvitationStream, event); err != nil { - return inv, err - } - - return inv, nil -} - -func (es *eventStore) ListInvitations(ctx context.Context, session authn.Session, pm domains.InvitationPageMeta) (domains.InvitationPage, error) { - ip, err := es.svc.ListInvitations(ctx, session, pm) - if err != nil { - return ip, err - } - - event := listInvitationsEvent{ - InvitationPageMeta: pm, - session: session, - requestID: middleware.GetReqID(ctx), - } - - if err := es.Publish(ctx, listInvitationsStream, event); err != nil { - return ip, err - } - - return ip, nil -} - -func (es *eventStore) ListDomainInvitations(ctx context.Context, session authn.Session, pm domains.InvitationPageMeta) (domains.InvitationPage, error) { - ip, err := es.svc.ListDomainInvitations(ctx, session, pm) - if err != nil { - return ip, err - } - - event := listDomainInvitationsEvent{ - InvitationPageMeta: pm, - session: session, - requestID: middleware.GetReqID(ctx), - } - - if err := es.Publish(ctx, listDomainInvitationsStream, event); err != nil { - return ip, err - } - - return ip, nil -} - -func (es *eventStore) AcceptInvitation(ctx context.Context, session authn.Session, domainID string) (domains.Invitation, error) { - inv, err := es.svc.AcceptInvitation(ctx, session, domainID) - if err != nil { - return inv, err - } - - if err := es.RoleManagerEventStore.RoleAddMembersEventPublisher(ctx, inv.DomainID, inv.RoleID, []string{inv.InviteeUserID}); err != nil { - return inv, err - } - - event := acceptInvitationEvent{ - invitation: inv, - session: session, - requestID: middleware.GetReqID(ctx), - } - - if err := es.Publish(ctx, acceptInvitationStream, event); err != nil { - return inv, err - } - return inv, nil -} - -func (es *eventStore) RejectInvitation(ctx context.Context, session authn.Session, domainID string) (domains.Invitation, error) { - inv, err := es.svc.RejectInvitation(ctx, session, domainID) - if err != nil { - return domains.Invitation{}, err - } - - event := rejectInvitationEvent{ - invitation: inv, - session: session, - requestID: middleware.GetReqID(ctx), - } - - if err := es.Publish(ctx, rejectInvitationStream, event); err != nil { - return inv, err - } - - return inv, nil -} - -func (es *eventStore) DeleteInvitation(ctx context.Context, session authn.Session, inviteeUserID, domainID string) error { - if err := es.svc.DeleteInvitation(ctx, session, inviteeUserID, domainID); err != nil { - return err - } - - event := deleteInvitationEvent{ - inviteeUserID: inviteeUserID, - domainID: domainID, - session: session, - requestID: middleware.GetReqID(ctx), - } - - return es.Publish(ctx, deleteInvitationStream, event) -} diff --git a/domains/events/streams_test.go b/domains/events/streams_test.go deleted file mode 100644 index 3b3dff69c..000000000 --- a/domains/events/streams_test.go +++ /dev/null @@ -1,709 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package events_test - -import ( - "context" - "fmt" - "os" - "testing" - "time" - - "github.com/absmach/magistrala/domains" - "github.com/absmach/magistrala/domains/events" - "github.com/absmach/magistrala/domains/mocks" - "github.com/absmach/magistrala/internal/testsutil" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/pkg/roles" - "github.com/go-chi/chi/v5/middleware" - "github.com/redis/go-redis/v9" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" -) - -var ( - storeClient *redis.Client - storeURL string - validSession = authn.Session{ - DomainID: testsutil.GenerateUUID(&testing.T{}), - UserID: testsutil.GenerateUUID(&testing.T{}), - } - validDomain = generateTestDomain(&testing.T{}) - validDomainsPage = domains.DomainsPage{ - Limit: 10, - Offset: 0, - Total: 1, - Domains: []domains.Domain{validDomain}, - } - validInvitation = generateTestInvitation(&testing.T{}) - validInvitationsPage = domains.InvitationPage{ - Total: 1, - Offset: 0, - Limit: 10, - Invitations: []domains.Invitation{validInvitation}, - } -) - -func newEventStoreMiddleware(t *testing.T) (*mocks.Service, domains.Service) { - svc := new(mocks.Service) - nsvc, err := events.NewEventStoreMiddleware(context.Background(), svc, storeURL) - require.Nil(t, err, fmt.Sprintf("create events store middleware failed with unexpected error: %s", err)) - - return svc, nsvc -} - -func TestMain(m *testing.M) { - code := testsutil.RunRedisTest(m, &storeClient, &storeURL) - os.Exit(code) -} - -func TestCreateDomain(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validID := testsutil.GenerateUUID(t) - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, validID) - - cases := []struct { - desc string - session authn.Session - domain domains.Domain - svcRes domains.Domain - svcRoleRes []roles.RoleProvision - svcErr error - resp domains.Domain - respRoleRes []roles.RoleProvision - err error - }{ - { - desc: "publish successfully", - session: validSession, - domain: validDomain, - svcRes: validDomain, - svcRoleRes: []roles.RoleProvision{}, - svcErr: nil, - resp: validDomain, - respRoleRes: []roles.RoleProvision{}, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - domain: validDomain, - svcRes: domains.Domain{}, - svcRoleRes: []roles.RoleProvision{}, - svcErr: svcerr.ErrCreateEntity, - resp: domains.Domain{}, - respRoleRes: []roles.RoleProvision{}, - err: svcerr.ErrCreateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("CreateDomain", validCtx, tc.session, tc.domain).Return(tc.svcRes, tc.svcRoleRes, tc.svcErr) - resp, respRoleRes, err := nsvc.CreateDomain(validCtx, tc.session, tc.domain) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - assert.Equal(t, tc.respRoleRes, respRoleRes, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.respRoleRes, respRoleRes)) - svcCall.Unset() - }) - } -} - -func TestRetrieveDomain(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - domainID string - withRoles bool - svcRes domains.Domain - svcErr error - resp domains.Domain - err error - }{ - { - desc: "publish successfully", - session: validSession, - domainID: validDomain.ID, - withRoles: false, - svcRes: validDomain, - svcErr: nil, - resp: validDomain, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - domainID: validDomain.ID, - withRoles: false, - svcRes: domains.Domain{}, - svcErr: svcerr.ErrViewEntity, - resp: domains.Domain{}, - err: svcerr.ErrViewEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("RetrieveDomain", validCtx, tc.session, tc.domainID, tc.withRoles).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.RetrieveDomain(validCtx, tc.session, tc.domainID, tc.withRoles) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestUpdateDomain(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - updatedDomain := validDomain - updatedDomain.Name = "updatedName" - domainReq := domains.DomainReq{ - Name: &updatedDomain.Name, - } - - cases := []struct { - desc string - session authn.Session - domainID string - domainReq domains.DomainReq - svcRes domains.Domain - svcErr error - resp domains.Domain - err error - }{ - { - desc: "publish successfully", - session: validSession, - domainID: validDomain.ID, - domainReq: domainReq, - svcRes: updatedDomain, - svcErr: nil, - resp: updatedDomain, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - domainID: validDomain.ID, - domainReq: domainReq, - svcRes: domains.Domain{}, - svcErr: svcerr.ErrUpdateEntity, - resp: domains.Domain{}, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("UpdateDomain", validCtx, tc.session, tc.domainID, tc.domainReq).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.UpdateDomain(validCtx, tc.session, tc.domainID, tc.domainReq) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestEnableDomain(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - domainID string - svcRes domains.Domain - svcErr error - resp domains.Domain - err error - }{ - { - desc: "publish successfully", - session: validSession, - domainID: validDomain.ID, - svcRes: validDomain, - svcErr: nil, - resp: validDomain, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - domainID: validDomain.ID, - svcRes: domains.Domain{}, - svcErr: svcerr.ErrUpdateEntity, - resp: domains.Domain{}, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("EnableDomain", validCtx, tc.session, tc.domainID).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.EnableDomain(validCtx, tc.session, tc.domainID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestDisableDomain(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - domainID string - svcRes domains.Domain - svcErr error - resp domains.Domain - err error - }{ - { - desc: "publish successfully", - session: validSession, - domainID: validDomain.ID, - svcRes: validDomain, - svcErr: nil, - resp: validDomain, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - domainID: validDomain.ID, - svcRes: domains.Domain{}, - svcErr: svcerr.ErrUpdateEntity, - resp: domains.Domain{}, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("DisableDomain", validCtx, tc.session, tc.domainID).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.DisableDomain(validCtx, tc.session, tc.domainID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestFreezeDomain(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - domainID string - svcRes domains.Domain - svcErr error - resp domains.Domain - err error - }{ - { - desc: "publish successfully", - session: validSession, - domainID: validDomain.ID, - svcRes: validDomain, - svcErr: nil, - resp: validDomain, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - domainID: validDomain.ID, - svcRes: domains.Domain{}, - svcErr: svcerr.ErrUpdateEntity, - resp: domains.Domain{}, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("FreezeDomain", validCtx, tc.session, tc.domainID).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.FreezeDomain(validCtx, tc.session, tc.domainID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestListDomains(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - pageMeta domains.Page - svcRes domains.DomainsPage - svcErr error - resp domains.DomainsPage - err error - }{ - { - desc: "publish successfully", - session: validSession, - pageMeta: domains.Page{ - Limit: 10, - Offset: 0, - }, - svcRes: validDomainsPage, - svcErr: nil, - resp: validDomainsPage, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - pageMeta: domains.Page{ - Limit: 10, - Offset: 0, - }, - svcRes: domains.DomainsPage{}, - svcErr: svcerr.ErrViewEntity, - resp: domains.DomainsPage{}, - err: svcerr.ErrViewEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("ListDomains", validCtx, tc.session, tc.pageMeta).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.ListDomains(validCtx, tc.session, tc.pageMeta) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestSendInvitation(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - invitation domains.Invitation - svcRes domains.Invitation - svcErr error - resp domains.Invitation - err error - }{ - { - desc: "publish successfully", - session: validSession, - invitation: validInvitation, - svcRes: validInvitation, - svcErr: nil, - resp: validInvitation, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - invitation: validInvitation, - svcRes: domains.Invitation{}, - svcErr: svcerr.ErrCreateEntity, - resp: domains.Invitation{}, - err: svcerr.ErrCreateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("SendInvitation", validCtx, tc.session, tc.invitation).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.SendInvitation(validCtx, tc.session, tc.invitation) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestListInvitations(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - pageMeta domains.InvitationPageMeta - svcRes domains.InvitationPage - svcErr error - resp domains.InvitationPage - err error - }{ - { - desc: "publish successfully", - session: validSession, - pageMeta: domains.InvitationPageMeta{ - Limit: 10, - Offset: 0, - }, - svcRes: validInvitationsPage, - svcErr: nil, - resp: validInvitationsPage, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - pageMeta: domains.InvitationPageMeta{ - Limit: 10, - Offset: 0, - }, - svcRes: domains.InvitationPage{}, - svcErr: svcerr.ErrViewEntity, - resp: domains.InvitationPage{}, - err: svcerr.ErrViewEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("ListInvitations", validCtx, tc.session, tc.pageMeta).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.ListInvitations(validCtx, tc.session, tc.pageMeta) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestListDomainInvitations(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - pageMeta domains.InvitationPageMeta - svcRes domains.InvitationPage - svcErr error - resp domains.InvitationPage - err error - }{ - { - desc: "publish successfully", - session: validSession, - pageMeta: domains.InvitationPageMeta{ - Limit: 10, - Offset: 0, - DomainID: validDomain.ID, - }, - svcRes: validInvitationsPage, - svcErr: nil, - resp: validInvitationsPage, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - pageMeta: domains.InvitationPageMeta{ - Limit: 10, - Offset: 0, - DomainID: validDomain.ID, - }, - svcRes: domains.InvitationPage{}, - svcErr: svcerr.ErrViewEntity, - resp: domains.InvitationPage{}, - err: svcerr.ErrViewEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("ListDomainInvitations", validCtx, tc.session, tc.pageMeta).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.ListDomainInvitations(validCtx, tc.session, tc.pageMeta) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestAcceptInvitation(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - domainID string - svcRes domains.Invitation - svcErr error - resp domains.Invitation - err error - }{ - { - desc: "publish successfully", - session: validSession, - domainID: validDomain.ID, - svcRes: validInvitation, - svcErr: nil, - resp: validInvitation, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - domainID: validDomain.ID, - svcRes: domains.Invitation{}, - svcErr: svcerr.ErrUpdateEntity, - resp: domains.Invitation{}, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("AcceptInvitation", validCtx, tc.session, tc.domainID).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.AcceptInvitation(validCtx, tc.session, tc.domainID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestDeleteInvitation(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - inviteeUserID string - domainID string - svcErr error - err error - }{ - { - desc: "publish successfully", - session: validSession, - inviteeUserID: validInvitation.InvitedBy, - domainID: validDomain.ID, - svcErr: nil, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - inviteeUserID: validInvitation.InvitedBy, - domainID: validDomain.ID, - svcErr: svcerr.ErrRemoveEntity, - err: svcerr.ErrRemoveEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("DeleteInvitation", validCtx, tc.session, tc.inviteeUserID, tc.domainID).Return(tc.svcErr) - err := nsvc.DeleteInvitation(validCtx, tc.session, tc.inviteeUserID, tc.domainID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - svcCall.Unset() - }) - } -} - -func TestRejectInvitation(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - domainID string - svcRes domains.Invitation - svcErr error - resp domains.Invitation - err error - }{ - { - desc: "publish successfully", - session: validSession, - domainID: validDomain.ID, - svcRes: validInvitation, - svcErr: nil, - resp: validInvitation, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - domainID: validDomain.ID, - svcRes: domains.Invitation{}, - svcErr: svcerr.ErrUpdateEntity, - resp: domains.Invitation{}, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("RejectInvitation", validCtx, tc.session, tc.domainID).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.RejectInvitation(validCtx, tc.session, tc.domainID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func generateTestDomain(t *testing.T) domains.Domain { - createdAt, err := time.Parse(time.RFC3339, "2024-01-01T00:00:00Z") - assert.Nil(t, err, fmt.Sprintf("Unexpected error parsing time: %v", err)) - return domains.Domain{ - ID: testsutil.GenerateUUID(t), - Name: "domainname", - Tags: []string{"tag1", "tag2"}, - Metadata: domains.Metadata{"key1": "value1"}, - CreatedAt: createdAt, - UpdatedAt: createdAt, - Status: domains.EnabledStatus, - } -} - -func generateTestInvitation(t *testing.T) domains.Invitation { - createdAt, err := time.Parse(time.RFC3339, "2024-01-01T00:00:00Z") - assert.Nil(t, err, fmt.Sprintf("Unexpected error parsing time: %v", err)) - return domains.Invitation{ - InvitedBy: testsutil.GenerateUUID(t), - InviteeUserID: testsutil.GenerateUUID(t), - DomainID: testsutil.GenerateUUID(t), - RoleID: testsutil.GenerateUUID(t), - CreatedAt: createdAt, - UpdatedAt: createdAt, - } -} diff --git a/domains/invitations.go b/domains/invitations.go deleted file mode 100644 index 9ed80dbea..000000000 --- a/domains/invitations.go +++ /dev/null @@ -1,60 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package domains - -import ( - "encoding/json" - "time" -) - -// Invitation is an invitation to join a domain. -type Invitation struct { - InvitedBy string `json:"invited_by"` - InviteeUserID string `json:"invitee_user_id"` - DomainID string `json:"domain_id"` - DomainName string `json:"domain_name,omitempty"` - RoleID string `json:"role_id,omitempty"` - RoleName string `json:"role_name,omitempty"` - Actions []string `json:"actions,omitempty"` - CreatedAt time.Time `json:"created_at"` - UpdatedAt time.Time `json:"updated_at,omitempty"` - ConfirmedAt time.Time `json:"confirmed_at,omitempty"` - RejectedAt time.Time `json:"rejected_at,omitempty"` - Resend bool `json:"resend,omitempty"` -} - -// InvitationPage is a page of invitations. -type InvitationPage struct { - Total uint64 `json:"total"` - Offset uint64 `json:"offset"` - Limit uint64 `json:"limit"` - Invitations []Invitation `json:"invitations"` -} - -func (page InvitationPage) MarshalJSON() ([]byte, error) { - type Alias InvitationPage - a := struct { - Alias - }{ - Alias: Alias(page), - } - - if a.Invitations == nil { - a.Invitations = make([]Invitation, 0) - } - - return json.Marshal(a) -} - -type InvitationPageMeta struct { - Offset uint64 `json:"offset" db:"offset"` - Limit uint64 `json:"limit" db:"limit"` - OnlyTotal bool `json:"only_total"` - InvitedBy string `json:"invited_by,omitempty" db:"invited_by,omitempty"` - InviteeUserID string `json:"invitee_user_id,omitempty" db:"invitee_user_id,omitempty"` - DomainID string `json:"domain_id,omitempty" db:"domain_id,omitempty"` - RoleID string `json:"role_id,omitempty" db:"role_id,omitempty"` - InvitedByOrUserID string `db:"invited_by_or_user_id,omitempty"` - State State `json:"state,omitempty"` -} diff --git a/domains/invitations_test.go b/domains/invitations_test.go deleted file mode 100644 index fe0e679ff..000000000 --- a/domains/invitations_test.go +++ /dev/null @@ -1,50 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package domains_test - -import ( - "fmt" - "testing" - - "github.com/absmach/magistrala/domains" - "github.com/stretchr/testify/assert" -) - -func TestInvitation_MarshalJSON(t *testing.T) { - cases := []struct { - desc string - page domains.InvitationPage - res string - }{ - { - desc: "empty page", - page: domains.InvitationPage{ - Invitations: []domains.Invitation(nil), - }, - res: `{"total":0,"offset":0,"limit":0,"invitations":[]}`, - }, - { - desc: "page with invitations", - page: domains.InvitationPage{ - Total: 1, - Offset: 0, - Limit: 0, - Invitations: []domains.Invitation{ - { - InvitedBy: "John", - InviteeUserID: "123", - DomainID: "123", - }, - }, - }, - res: `{"total":1,"offset":0,"limit":0,"invitations":[{"invited_by":"John","invitee_user_id":"123","domain_id":"123","created_at":"0001-01-01T00:00:00Z","updated_at":"0001-01-01T00:00:00Z","confirmed_at":"0001-01-01T00:00:00Z","rejected_at":"0001-01-01T00:00:00Z"}]}`, - }, - } - - for _, tc := range cases { - data, err := tc.page.MarshalJSON() - assert.NoError(t, err, "Unexpected error: %v", err) - assert.Equal(t, tc.res, string(data), fmt.Sprintf("%s: expected %s, got %s", tc.desc, tc.res, string(data))) - } -} diff --git a/domains/middleware/authorization.go b/domains/middleware/authorization.go deleted file mode 100644 index 98c4d5184..000000000 --- a/domains/middleware/authorization.go +++ /dev/null @@ -1,285 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package middleware - -import ( - "context" - - "github.com/absmach/magistrala/auth" - "github.com/absmach/magistrala/domains" - "github.com/absmach/magistrala/domains/operations" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/authz" - smqauthz "github.com/absmach/magistrala/pkg/authz" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/pkg/permissions" - "github.com/absmach/magistrala/pkg/policies" - "github.com/absmach/magistrala/pkg/roles" - rolemgr "github.com/absmach/magistrala/pkg/roles/rolemanager/middleware" -) - -var _ domains.Service = (*authorizationMiddleware)(nil) - -// ErrMemberExist indicates that the user is already a member of the domain. -var ErrMemberExist = errors.New("user is already a member of the domain") - -type authorizationMiddleware struct { - svc domains.Service - authz smqauthz.Authorization - entitiesOps permissions.EntitiesOperations[permissions.Operation] - rOps permissions.Operations[permissions.RoleOperation] - rolemgr.RoleManagerAuthorizationMiddleware -} - -// NewAuthorization adds authorization to the domains service. -func NewAuthorization(entityType string, svc domains.Service, authz smqauthz.Authorization, entitiesOps permissions.EntitiesOperations[permissions.Operation], domainRoleOps permissions.Operations[permissions.RoleOperation]) (domains.Service, error) { - if err := entitiesOps.Validate(); err != nil { - return &authorizationMiddleware{}, err - } - - ram, err := rolemgr.NewAuthorization(entityType, svc, authz, domainRoleOps) - if err != nil { - return &authorizationMiddleware{}, err - } - return &authorizationMiddleware{ - svc: svc, - authz: authz, - entitiesOps: entitiesOps, - rOps: domainRoleOps, - RoleManagerAuthorizationMiddleware: ram, - }, nil -} - -func (am *authorizationMiddleware) CreateDomain(ctx context.Context, session authn.Session, d domains.Domain) (domains.Domain, []roles.RoleProvision, error) { - return am.svc.CreateDomain(ctx, session, d) -} - -func (am *authorizationMiddleware) RetrieveDomain(ctx context.Context, session authn.Session, id string, withRoles bool) (domains.Domain, error) { - switch err := am.checkSuperAdmin(ctx, session); { - case err == nil: - session.SuperAdmin = true - return am.svc.RetrieveDomain(ctx, session, id, withRoles) - case errors.Contains(err, svcerr.ErrSuperAdminAction): - default: - return domains.Domain{}, err - } - - if err := am.authorize(ctx, session, policies.DomainType, operations.OpRetrieveDomain, authz.PolicyReq{ - Subject: session.DomainUserID, - SubjectType: policies.UserType, - SubjectKind: policies.UsersKind, - Object: id, - ObjectType: policies.DomainType, - }); err != nil { - return domains.Domain{}, err - } - - return am.svc.RetrieveDomain(ctx, session, id, withRoles) -} - -func (am *authorizationMiddleware) UpdateDomain(ctx context.Context, session authn.Session, id string, d domains.DomainReq) (domains.Domain, error) { - if err := am.authorize(ctx, session, policies.DomainType, operations.OpUpdateDomain, authz.PolicyReq{ - Subject: session.DomainUserID, - SubjectType: policies.UserType, - SubjectKind: policies.UsersKind, - Object: id, - ObjectType: policies.DomainType, - }); err != nil { - return domains.Domain{}, err - } - - return am.svc.UpdateDomain(ctx, session, id, d) -} - -func (am *authorizationMiddleware) EnableDomain(ctx context.Context, session authn.Session, id string) (domains.Domain, error) { - if err := am.authorize(ctx, session, policies.DomainType, operations.OpEnableDomain, authz.PolicyReq{ - Subject: session.DomainUserID, - SubjectType: policies.UserType, - SubjectKind: policies.UsersKind, - Object: id, - ObjectType: policies.DomainType, - }); err != nil { - return domains.Domain{}, err - } - - return am.svc.EnableDomain(ctx, session, id) -} - -func (am *authorizationMiddleware) DisableDomain(ctx context.Context, session authn.Session, id string) (domains.Domain, error) { - if err := am.authorize(ctx, session, policies.DomainType, operations.OpDisableDomain, authz.PolicyReq{ - Subject: session.DomainUserID, - SubjectType: policies.UserType, - SubjectKind: policies.UsersKind, - Object: id, - ObjectType: policies.DomainType, - }); err != nil { - return domains.Domain{}, err - } - - return am.svc.DisableDomain(ctx, session, id) -} - -func (am *authorizationMiddleware) FreezeDomain(ctx context.Context, session authn.Session, id string) (domains.Domain, error) { - // Only SuperAdmin can freeze the domain - if err := am.authz.Authorize(ctx, authz.PolicyReq{ - Subject: session.UserID, - SubjectType: policies.UserType, - SubjectKind: policies.UsersKind, - Permission: policies.AdminPermission, - Object: policies.MagistralaObject, - ObjectType: policies.PlatformType, - }, nil); err != nil { - return domains.Domain{}, err - } - - return am.svc.FreezeDomain(ctx, session, id) -} - -func (am *authorizationMiddleware) ListDomains(ctx context.Context, session authn.Session, page domains.Page) (domains.DomainsPage, error) { - switch err := am.checkSuperAdmin(ctx, session); { - case err == nil: - session.SuperAdmin = true - case errors.Contains(err, svcerr.ErrSuperAdminAction): - default: - return domains.DomainsPage{}, err - } - - return am.svc.ListDomains(ctx, session, page) -} - -func (am *authorizationMiddleware) SendInvitation(ctx context.Context, session authn.Session, invitation domains.Invitation) (domains.Invitation, error) { - if err := am.authorize(ctx, session, policies.DomainType, operations.OpSendDomainInvitation, authz.PolicyReq{ - Subject: session.DomainUserID, - SubjectType: policies.UserType, - SubjectKind: policies.UsersKind, - Object: session.DomainID, - ObjectType: policies.DomainType, - }); err != nil { - return domains.Invitation{}, err - } - - if err := am.checkAdmin(ctx, session); err != nil { - return domains.Invitation{}, err - } - - return am.svc.SendInvitation(ctx, session, invitation) -} - -func (am *authorizationMiddleware) ListInvitations(ctx context.Context, session authn.Session, page domains.InvitationPageMeta) (invs domains.InvitationPage, err error) { - return am.svc.ListInvitations(ctx, session, page) -} - -func (am *authorizationMiddleware) ListDomainInvitations(ctx context.Context, session authn.Session, page domains.InvitationPageMeta) (invs domains.InvitationPage, err error) { - if err := am.authorize(ctx, session, policies.DomainType, operations.OpListDomainInvitations, authz.PolicyReq{ - Subject: session.DomainUserID, - SubjectType: policies.UserType, - SubjectKind: policies.UsersKind, - Object: session.DomainID, - ObjectType: policies.DomainType, - }); err != nil { - return domains.InvitationPage{}, err - } - - return am.svc.ListDomainInvitations(ctx, session, page) -} - -func (am *authorizationMiddleware) AcceptInvitation(ctx context.Context, session authn.Session, domainID string) (inv domains.Invitation, err error) { - return am.svc.AcceptInvitation(ctx, session, domainID) -} - -func (am *authorizationMiddleware) RejectInvitation(ctx context.Context, session authn.Session, domainID string) (domains.Invitation, error) { - return am.svc.RejectInvitation(ctx, session, domainID) -} - -func (am *authorizationMiddleware) DeleteInvitation(ctx context.Context, session authn.Session, inviteeUserID, domainID string) (err error) { - if err := am.authorize(ctx, session, policies.DomainType, operations.OpDeleteDomainInvitation, authz.PolicyReq{ - Subject: session.DomainUserID, - SubjectType: policies.UserType, - SubjectKind: policies.UsersKind, - Object: session.DomainID, - ObjectType: policies.DomainType, - }); err != nil { - return err - } - - return am.svc.DeleteInvitation(ctx, session, inviteeUserID, domainID) -} - -func (am *authorizationMiddleware) authorize(ctx context.Context, session authn.Session, entityType string, op permissions.Operation, authReq authz.PolicyReq) error { - authReq.Domain = session.DomainID - - perm, err := am.entitiesOps.GetPermission(entityType, op) - if err != nil { - return err - } - - authReq.Permission = perm.String() - - var pat *smqauthz.PATReq - if session.PatID != "" { - entityID := authReq.Object - opName := am.entitiesOps.OperationName(entityType, op) - pat = &smqauthz.PATReq{ - UserID: session.UserID, - PatID: session.PatID, - EntityID: entityID, - EntityType: auth.DomainsType.String(), - Operation: opName, - Domain: session.DomainID, - } - } - - if err := am.authz.Authorize(ctx, authReq, pat); err != nil { - return err - } - - return nil -} - -// checkAdmin checks if the given user is a domain or platform administrator. -func (am *authorizationMiddleware) checkAdmin(ctx context.Context, session authn.Session) error { - req := smqauthz.PolicyReq{ - SubjectType: policies.UserType, - SubjectKind: policies.UsersKind, - Subject: session.DomainUserID, - Permission: policies.AdminPermission, - ObjectType: policies.DomainType, - Object: session.DomainID, - } - if err := am.authz.Authorize(ctx, req, nil); err == nil { - return nil - } - - req = smqauthz.PolicyReq{ - SubjectType: policies.UserType, - SubjectKind: policies.UsersKind, - Subject: session.UserID, - Permission: policies.AdminPermission, - ObjectType: policies.PlatformType, - Object: policies.MagistralaObject, - } - - if err := am.authz.Authorize(ctx, req, nil); err == nil { - return nil - } - - return svcerr.ErrAuthorization -} - -func (am *authorizationMiddleware) checkSuperAdmin(ctx context.Context, session authn.Session) error { - if session.Role != authn.SuperAdminRole { - return svcerr.ErrSuperAdminAction - } - if err := am.authz.Authorize(ctx, smqauthz.PolicyReq{ - SubjectType: policies.UserType, - Subject: session.UserID, - Permission: policies.AdminPermission, - ObjectType: policies.PlatformType, - Object: policies.MagistralaObject, - }, nil); err != nil { - return err - } - return nil -} diff --git a/domains/middleware/callout.go b/domains/middleware/callout.go deleted file mode 100644 index 124f778c9..000000000 --- a/domains/middleware/callout.go +++ /dev/null @@ -1,238 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package middleware - -import ( - "context" - "time" - - "github.com/absmach/magistrala/domains" - "github.com/absmach/magistrala/domains/operations" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/callout" - "github.com/absmach/magistrala/pkg/permissions" - "github.com/absmach/magistrala/pkg/policies" - "github.com/absmach/magistrala/pkg/roles" - rolemgr "github.com/absmach/magistrala/pkg/roles/rolemanager/middleware" -) - -var _ domains.Service = (*calloutMiddleware)(nil) - -type calloutMiddleware struct { - svc domains.Service - callout callout.Callout - entitiesOps permissions.EntitiesOperations[permissions.Operation] - rolemgr.RoleManagerCalloutMiddleware -} - -func NewCallout(svc domains.Service, entitiesOps permissions.EntitiesOperations[permissions.Operation], roleOps permissions.Operations[permissions.RoleOperation], callout callout.Callout) (domains.Service, error) { - call, err := rolemgr.NewCallout(policies.DomainType, svc, callout, roleOps) - if err != nil { - return nil, err - } - - if err := entitiesOps.Validate(); err != nil { - return nil, err - } - - return &calloutMiddleware{ - svc: svc, - callout: callout, - entitiesOps: entitiesOps, - RoleManagerCalloutMiddleware: call, - }, nil -} - -func (cm *calloutMiddleware) CreateDomain(ctx context.Context, session authn.Session, d domains.Domain) (domains.Domain, []roles.RoleProvision, error) { - params := map[string]any{ - "entity_id": d.ID, - } - - if err := cm.callOut(ctx, session, policies.DomainType, operations.OpCreateDomain, params); err != nil { - return domains.Domain{}, nil, err - } - - return cm.svc.CreateDomain(ctx, session, d) -} - -func (cm *calloutMiddleware) RetrieveDomain(ctx context.Context, session authn.Session, id string, withRoles bool) (domains.Domain, error) { - params := map[string]any{ - "entity_id": id, - "with_roles": withRoles, - } - - if err := cm.callOut(ctx, session, policies.DomainType, operations.OpRetrieveDomain, params); err != nil { - return domains.Domain{}, err - } - - return cm.svc.RetrieveDomain(ctx, session, id, withRoles) -} - -func (cm *calloutMiddleware) UpdateDomain(ctx context.Context, session authn.Session, id string, d domains.DomainReq) (domains.Domain, error) { - params := map[string]any{ - "entity_id": id, - "domain_req": d, - } - - if err := cm.callOut(ctx, session, policies.DomainType, operations.OpUpdateDomain, params); err != nil { - return domains.Domain{}, err - } - - return cm.svc.UpdateDomain(ctx, session, id, d) -} - -func (cm *calloutMiddleware) EnableDomain(ctx context.Context, session authn.Session, id string) (domains.Domain, error) { - params := map[string]any{ - "entity_id": id, - } - - if err := cm.callOut(ctx, session, policies.DomainType, operations.OpEnableDomain, params); err != nil { - return domains.Domain{}, err - } - - return cm.svc.EnableDomain(ctx, session, id) -} - -func (cm *calloutMiddleware) DisableDomain(ctx context.Context, session authn.Session, id string) (domains.Domain, error) { - params := map[string]any{ - "entity_id": id, - } - - if err := cm.callOut(ctx, session, policies.DomainType, operations.OpDisableDomain, params); err != nil { - return domains.Domain{}, err - } - - return cm.svc.DisableDomain(ctx, session, id) -} - -func (cm *calloutMiddleware) FreezeDomain(ctx context.Context, session authn.Session, id string) (domains.Domain, error) { - params := map[string]any{ - "entity_id": id, - } - - if err := cm.callOut(ctx, session, policies.DomainType, operations.OpFreezeDomain, params); err != nil { - return domains.Domain{}, err - } - - return cm.svc.FreezeDomain(ctx, session, id) -} - -func (cm *calloutMiddleware) ListDomains(ctx context.Context, session authn.Session, page domains.Page) (domains.DomainsPage, error) { - params := map[string]any{ - "page": page, - } - - if err := cm.callOut(ctx, session, policies.DomainType, operations.OpListDomains, params); err != nil { - return domains.DomainsPage{}, err - } - - return cm.svc.ListDomains(ctx, session, page) -} - -func (cm *calloutMiddleware) SendInvitation(ctx context.Context, session authn.Session, invitation domains.Invitation) (domains.Invitation, error) { - params := map[string]any{ - "entity_id": invitation.DomainID, - "invitation": invitation, - } - - // While entity here is technically an invitation, Domain is used as - // the entity in callout since the invitation refers to the domain. - if err := cm.callOut(ctx, session, policies.DomainType, operations.OpSendDomainInvitation, params); err != nil { - return domains.Invitation{}, err - } - - return cm.svc.SendInvitation(ctx, session, invitation) -} - -func (cm *calloutMiddleware) ListInvitations(ctx context.Context, session authn.Session, page domains.InvitationPageMeta) (domains.InvitationPage, error) { - params := map[string]any{ - "page": page, - } - - if err := cm.callOut(ctx, session, policies.DomainType, operations.OpListInvitations, params); err != nil { - return domains.InvitationPage{}, err - } - - return cm.svc.ListInvitations(ctx, session, page) -} - -func (cm *calloutMiddleware) ListDomainInvitations(ctx context.Context, session authn.Session, page domains.InvitationPageMeta) (domains.InvitationPage, error) { - params := map[string]any{ - "entity_id": page.DomainID, - "page": page, - } - - if err := cm.callOut(ctx, session, policies.DomainType, operations.OpListDomainInvitations, params); err != nil { - return domains.InvitationPage{}, err - } - - return cm.svc.ListDomainInvitations(ctx, session, page) -} - -func (cm *calloutMiddleware) AcceptInvitation(ctx context.Context, session authn.Session, domainID string) (domains.Invitation, error) { - params := map[string]any{ - "entity_id": domainID, - } - - // Similar to sending an invitation, Domain is used as the - // entity in callout since the invitation refers to the domain. - if err := cm.callOut(ctx, session, policies.DomainType, operations.OpAcceptInvitation, params); err != nil { - return domains.Invitation{}, err - } - - return cm.svc.AcceptInvitation(ctx, session, domainID) -} - -func (cm *calloutMiddleware) RejectInvitation(ctx context.Context, session authn.Session, domainID string) (domains.Invitation, error) { - params := map[string]any{ - "entity_id": domainID, - } - - // Similar to sending and accepting, Domain is used as - // the entity in callout since the invitation refers to the domain. - if err := cm.callOut(ctx, session, policies.DomainType, operations.OpRejectInvitation, params); err != nil { - return domains.Invitation{}, err - } - - return cm.svc.RejectInvitation(ctx, session, domainID) -} - -func (cm *calloutMiddleware) DeleteInvitation(ctx context.Context, session authn.Session, inviteeUserID, domainID string) error { - params := map[string]any{ - "entity_id": domainID, - "invitee_user_id": inviteeUserID, - } - - if err := cm.callOut(ctx, session, policies.DomainType, operations.OpDeleteDomainInvitation, params); err != nil { - return err - } - - return cm.svc.DeleteInvitation(ctx, session, inviteeUserID, domainID) -} - -func (cm *calloutMiddleware) callOut(ctx context.Context, session authn.Session, entityType string, op permissions.Operation, pld map[string]any) error { - var entityID string - if id, ok := pld["entity_id"].(string); ok { - entityID = id - } - - req := callout.Request{ - BaseRequest: callout.BaseRequest{ - Operation: cm.entitiesOps.OperationName(entityType, op), - EntityType: entityType, - EntityID: entityID, - CallerID: session.UserID, - CallerType: policies.UserType, - DomainID: entityID, - Time: time.Now().UTC(), - }, - Payload: pld, - } - - if err := cm.callout.Callout(ctx, req); err != nil { - return err - } - - return nil -} diff --git a/domains/middleware/doc.go b/domains/middleware/doc.go deleted file mode 100644 index a4c6f9861..000000000 --- a/domains/middleware/doc.go +++ /dev/null @@ -1,9 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package middleware provides authorization, logging, metrics and tracing middleware -// for Magistrala Domains service. -// -// For more details about tracing instrumentation for Magistrala refer to the -// documentation at https://magistrala.absmach.eu/docs/. -package middleware diff --git a/domains/middleware/logging.go b/domains/middleware/logging.go deleted file mode 100644 index c0407f809..000000000 --- a/domains/middleware/logging.go +++ /dev/null @@ -1,294 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -//go:build !test - -package middleware - -import ( - "context" - "log/slog" - "time" - - "github.com/absmach/magistrala/domains" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/roles" - rolemw "github.com/absmach/magistrala/pkg/roles/rolemanager/middleware" - "github.com/go-chi/chi/v5/middleware" -) - -var _ domains.Service = (*loggingMiddleware)(nil) - -type loggingMiddleware struct { - logger *slog.Logger - svc domains.Service - rolemw.RoleManagerLoggingMiddleware -} - -// NewLogging adds logging facilities to the core service. -func NewLogging(svc domains.Service, logger *slog.Logger) domains.Service { - rmlm := rolemw.NewLogging("domains", svc, logger) - return &loggingMiddleware{ - logger: logger, - svc: svc, - RoleManagerLoggingMiddleware: rmlm, - } -} - -func (lm *loggingMiddleware) CreateDomain(ctx context.Context, session authn.Session, d domains.Domain) (do domains.Domain, rps []roles.RoleProvision, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.Group("domain", - slog.String("id", d.ID), - slog.String("name", d.Name), - slog.String("route", d.Route), - ), - } - if err != nil { - args := append(args, slog.String("error", err.Error())) - lm.logger.Warn("Create domain failed", args...) - return - } - lm.logger.Info("Create domain completed successfully", args...) - }(time.Now()) - return lm.svc.CreateDomain(ctx, session, d) -} - -func (lm *loggingMiddleware) RetrieveDomain(ctx context.Context, session authn.Session, id string, withRoles bool) (do domains.Domain, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.String("domain_id", id), - slog.Bool("with_roles", withRoles), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Retrieve domain failed", args...) - return - } - lm.logger.Info("Retrieve domain completed successfully", args...) - }(time.Now()) - return lm.svc.RetrieveDomain(ctx, session, id, withRoles) -} - -func (lm *loggingMiddleware) UpdateDomain(ctx context.Context, session authn.Session, id string, d domains.DomainReq) (do domains.Domain, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.Group("domain", - slog.String("id", id), - slog.Any("name", d.Name), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Update domain failed", args...) - return - } - lm.logger.Info("Update domain completed successfully", args...) - }(time.Now()) - return lm.svc.UpdateDomain(ctx, session, id, d) -} - -func (lm *loggingMiddleware) EnableDomain(ctx context.Context, session authn.Session, id string) (do domains.Domain, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.Group("domain", - slog.String("id", id), - slog.String("name", do.Name), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Enable domain failed", args...) - return - } - lm.logger.Info("Enable domain completed successfully", args...) - }(time.Now()) - return lm.svc.EnableDomain(ctx, session, id) -} - -func (lm *loggingMiddleware) DisableDomain(ctx context.Context, session authn.Session, id string) (do domains.Domain, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.Group("domain", - slog.String("id", id), - slog.String("name", do.Name), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Disable domain failed", args...) - return - } - lm.logger.Info("Disable domain completed successfully", args...) - }(time.Now()) - return lm.svc.DisableDomain(ctx, session, id) -} - -func (lm *loggingMiddleware) FreezeDomain(ctx context.Context, session authn.Session, id string) (do domains.Domain, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.Group("domain", - slog.String("id", id), - slog.String("name", do.Name), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Freeze domain failed", args...) - return - } - lm.logger.Info("Freeze domain completed successfully", args...) - }(time.Now()) - return lm.svc.FreezeDomain(ctx, session, id) -} - -func (lm *loggingMiddleware) ListDomains(ctx context.Context, session authn.Session, page domains.Page) (do domains.DomainsPage, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.Group("page", - slog.Uint64("limit", page.Limit), - slog.Uint64("offset", page.Offset), - slog.Uint64("total", page.Total), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("List domains failed", args...) - return - } - lm.logger.Info("List domains completed successfully", args...) - }(time.Now()) - return lm.svc.ListDomains(ctx, session, page) -} - -func (lm *loggingMiddleware) SendInvitation(ctx context.Context, session authn.Session, invitation domains.Invitation) (inv domains.Invitation, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.String("user_id", session.UserID), - slog.String("invitee_user_id", invitation.InviteeUserID), - slog.String("domain_id", invitation.DomainID), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Send invitation failed", args...) - return - } - lm.logger.Info("Send invitation completed successfully", args...) - }(time.Now()) - return lm.svc.SendInvitation(ctx, session, invitation) -} - -func (lm *loggingMiddleware) ListInvitations(ctx context.Context, session authn.Session, pm domains.InvitationPageMeta) (invs domains.InvitationPage, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.String("user_id", session.UserID), - slog.Group("page", - slog.Uint64("offset", pm.Offset), - slog.Uint64("limit", pm.Limit), - slog.Uint64("total", invs.Total), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("List invitations failed", args...) - return - } - lm.logger.Info("List invitations completed successfully", args...) - }(time.Now()) - return lm.svc.ListInvitations(ctx, session, pm) -} - -func (lm *loggingMiddleware) ListDomainInvitations(ctx context.Context, session authn.Session, pm domains.InvitationPageMeta) (invs domains.InvitationPage, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.String("domain_id", session.DomainID), - slog.Group("page", - slog.Uint64("offset", pm.Offset), - slog.Uint64("limit", pm.Limit), - slog.Uint64("total", invs.Total), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("List domain invitations failed", args...) - return - } - lm.logger.Info("List domain invitations completed successfully", args...) - }(time.Now()) - return lm.svc.ListDomainInvitations(ctx, session, pm) -} - -func (lm *loggingMiddleware) AcceptInvitation(ctx context.Context, session authn.Session, domainID string) (inv domains.Invitation, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.String("user_id", session.UserID), - slog.String("domain_id", domainID), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Accept invitation failed", args...) - return - } - lm.logger.Info("Accept invitation completed successfully", args...) - }(time.Now()) - return lm.svc.AcceptInvitation(ctx, session, domainID) -} - -func (lm *loggingMiddleware) RejectInvitation(ctx context.Context, session authn.Session, domainID string) (inv domains.Invitation, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.String("user_id", session.UserID), - slog.String("domain_id", domainID), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Reject invitation failed", args...) - return - } - lm.logger.Info("Reject invitation completed successfully", args...) - }(time.Now()) - return lm.svc.RejectInvitation(ctx, session, domainID) -} - -func (lm *loggingMiddleware) DeleteInvitation(ctx context.Context, session authn.Session, inviteeUserID, domainID string) (err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.String("user_id", session.UserID), - slog.String("invitee_user_id", inviteeUserID), - slog.String("domain_id", domainID), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Delete invitation failed", args...) - return - } - lm.logger.Info("Delete invitation completed successfully", args...) - }(time.Now()) - return lm.svc.DeleteInvitation(ctx, session, inviteeUserID, domainID) -} diff --git a/domains/middleware/metrics.go b/domains/middleware/metrics.go deleted file mode 100644 index add74a2fd..000000000 --- a/domains/middleware/metrics.go +++ /dev/null @@ -1,142 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -//go:build !test - -package middleware - -import ( - "context" - "time" - - "github.com/absmach/magistrala/domains" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/roles" - rolemw "github.com/absmach/magistrala/pkg/roles/rolemanager/middleware" - "github.com/go-kit/kit/metrics" -) - -var _ domains.Service = (*metricsMiddleware)(nil) - -type metricsMiddleware struct { - counter metrics.Counter - latency metrics.Histogram - svc domains.Service - rolemw.RoleManagerMetricsMiddleware -} - -// NewMetrics instruments core service by tracking request count and latency. -func NewMetrics(svc domains.Service, counter metrics.Counter, latency metrics.Histogram) domains.Service { - rmmw := rolemw.NewMetrics("domains", svc, counter, latency) - - return &metricsMiddleware{ - counter: counter, - latency: latency, - svc: svc, - RoleManagerMetricsMiddleware: rmmw, - } -} - -func (ms *metricsMiddleware) CreateDomain(ctx context.Context, session authn.Session, d domains.Domain) (domains.Domain, []roles.RoleProvision, error) { - defer func(begin time.Time) { - ms.counter.With("method", "create_domain").Add(1) - ms.latency.With("method", "create_domain").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.CreateDomain(ctx, session, d) -} - -func (ms *metricsMiddleware) RetrieveDomain(ctx context.Context, session authn.Session, id string, withRoles bool) (domains.Domain, error) { - defer func(begin time.Time) { - ms.counter.With("method", "retrieve_domain").Add(1) - ms.latency.With("method", "retrieve_domain").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.RetrieveDomain(ctx, session, id, withRoles) -} - -func (ms *metricsMiddleware) UpdateDomain(ctx context.Context, session authn.Session, id string, d domains.DomainReq) (domains.Domain, error) { - defer func(begin time.Time) { - ms.counter.With("method", "update_domain").Add(1) - ms.latency.With("method", "update_domain").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.UpdateDomain(ctx, session, id, d) -} - -func (ms *metricsMiddleware) EnableDomain(ctx context.Context, session authn.Session, id string) (domains.Domain, error) { - defer func(begin time.Time) { - ms.counter.With("method", "enable_domain").Add(1) - ms.latency.With("method", "enable_domain").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.EnableDomain(ctx, session, id) -} - -func (ms *metricsMiddleware) DisableDomain(ctx context.Context, session authn.Session, id string) (domains.Domain, error) { - defer func(begin time.Time) { - ms.counter.With("method", "disable_domain").Add(1) - ms.latency.With("method", "disable_domain").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.DisableDomain(ctx, session, id) -} - -func (ms *metricsMiddleware) FreezeDomain(ctx context.Context, session authn.Session, id string) (domains.Domain, error) { - defer func(begin time.Time) { - ms.counter.With("method", "freeze_domain").Add(1) - ms.latency.With("method", "freeze_domain").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.FreezeDomain(ctx, session, id) -} - -func (ms *metricsMiddleware) ListDomains(ctx context.Context, session authn.Session, page domains.Page) (domains.DomainsPage, error) { - defer func(begin time.Time) { - ms.counter.With("method", "list_domains").Add(1) - ms.latency.With("method", "list_domains").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.ListDomains(ctx, session, page) -} - -func (mm *metricsMiddleware) SendInvitation(ctx context.Context, session authn.Session, invitation domains.Invitation) (domains.Invitation, error) { - defer func(begin time.Time) { - mm.counter.With("method", "send_invitation").Add(1) - mm.latency.With("method", "send_invitation").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return mm.svc.SendInvitation(ctx, session, invitation) -} - -func (mm *metricsMiddleware) ListInvitations(ctx context.Context, session authn.Session, pm domains.InvitationPageMeta) (invs domains.InvitationPage, err error) { - defer func(begin time.Time) { - mm.counter.With("method", "list_invitations").Add(1) - mm.latency.With("method", "list_invitations").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return mm.svc.ListInvitations(ctx, session, pm) -} - -func (mm *metricsMiddleware) ListDomainInvitations(ctx context.Context, session authn.Session, pm domains.InvitationPageMeta) (invs domains.InvitationPage, err error) { - defer func(begin time.Time) { - mm.counter.With("method", "list_invitee_invitations").Add(1) - mm.latency.With("method", "list_invitee_invitations").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return mm.svc.ListDomainInvitations(ctx, session, pm) -} - -func (mm *metricsMiddleware) AcceptInvitation(ctx context.Context, session authn.Session, domainID string) (inv domains.Invitation, err error) { - defer func(begin time.Time) { - mm.counter.With("method", "accept_invitation").Add(1) - mm.latency.With("method", "accept_invitation").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return mm.svc.AcceptInvitation(ctx, session, domainID) -} - -func (mm *metricsMiddleware) RejectInvitation(ctx context.Context, session authn.Session, domainID string) (inv domains.Invitation, err error) { - defer func(begin time.Time) { - mm.counter.With("method", "reject_invitation").Add(1) - mm.latency.With("method", "reject_invitation").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return mm.svc.RejectInvitation(ctx, session, domainID) -} - -func (mm *metricsMiddleware) DeleteInvitation(ctx context.Context, session authn.Session, userID, domainID string) (err error) { - defer func(begin time.Time) { - mm.counter.With("method", "delete_invitation").Add(1) - mm.latency.With("method", "delete_invitation").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return mm.svc.DeleteInvitation(ctx, session, userID, domainID) -} diff --git a/domains/middleware/tracing.go b/domains/middleware/tracing.go deleted file mode 100644 index 6636cdce9..000000000 --- a/domains/middleware/tracing.go +++ /dev/null @@ -1,146 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package middleware - -import ( - "context" - - "github.com/absmach/magistrala/domains" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/roles" - rolemw "github.com/absmach/magistrala/pkg/roles/rolemanager/middleware" - "github.com/absmach/magistrala/pkg/tracing" - "go.opentelemetry.io/otel/attribute" - "go.opentelemetry.io/otel/trace" -) - -var _ domains.Service = (*tracingMiddleware)(nil) - -type tracingMiddleware struct { - tracer trace.Tracer - svc domains.Service - rolemw.RoleManagerTracing -} - -// NewTracing returns a new domains service with tracing capabilities. -func NewTracing(svc domains.Service, tracer trace.Tracer) domains.Service { - return &tracingMiddleware{tracer, svc, rolemw.NewTracing("domain", svc, tracer)} -} - -func (tm *tracingMiddleware) CreateDomain(ctx context.Context, session authn.Session, d domains.Domain) (domains.Domain, []roles.RoleProvision, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "create_domain", trace.WithAttributes( - attribute.String("name", d.Name), - )) - defer span.End() - return tm.svc.CreateDomain(ctx, session, d) -} - -func (tm *tracingMiddleware) RetrieveDomain(ctx context.Context, session authn.Session, id string, withRoles bool) (domains.Domain, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "view_domain", trace.WithAttributes( - attribute.String("id", id), - attribute.Bool("with_roles", withRoles), - )) - defer span.End() - return tm.svc.RetrieveDomain(ctx, session, id, withRoles) -} - -func (tm *tracingMiddleware) UpdateDomain(ctx context.Context, session authn.Session, id string, d domains.DomainReq) (domains.Domain, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "update_domain", trace.WithAttributes( - attribute.String("id", id), - )) - defer span.End() - return tm.svc.UpdateDomain(ctx, session, id, d) -} - -func (tm *tracingMiddleware) EnableDomain(ctx context.Context, session authn.Session, id string) (domains.Domain, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "enable_domain", trace.WithAttributes( - attribute.String("id", id), - )) - defer span.End() - return tm.svc.EnableDomain(ctx, session, id) -} - -func (tm *tracingMiddleware) DisableDomain(ctx context.Context, session authn.Session, id string) (domains.Domain, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "disable_domain", trace.WithAttributes( - attribute.String("id", id), - )) - defer span.End() - return tm.svc.DisableDomain(ctx, session, id) -} - -func (tm *tracingMiddleware) FreezeDomain(ctx context.Context, session authn.Session, id string) (domains.Domain, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "freeze_domain", trace.WithAttributes( - attribute.String("id", id), - )) - defer span.End() - return tm.svc.FreezeDomain(ctx, session, id) -} - -func (tm *tracingMiddleware) ListDomains(ctx context.Context, session authn.Session, p domains.Page) (domains.DomainsPage, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "list_domains") - defer span.End() - return tm.svc.ListDomains(ctx, session, p) -} - -func (tm *tracingMiddleware) SendInvitation(ctx context.Context, session authn.Session, invitation domains.Invitation) (domains.Invitation, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "send_invitation", trace.WithAttributes( - attribute.String("domain_id", invitation.DomainID), - attribute.String("invitee_user_id", invitation.InviteeUserID), - )) - defer span.End() - - return tm.svc.SendInvitation(ctx, session, invitation) -} - -func (tm *tracingMiddleware) ListInvitations(ctx context.Context, session authn.Session, pm domains.InvitationPageMeta) (invs domains.InvitationPage, err error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "list_invitations", trace.WithAttributes( - attribute.Int("limit", int(pm.Limit)), - attribute.Int("offset", int(pm.Offset)), - attribute.String("invitee_user_id", pm.InviteeUserID), - attribute.String("domain_id", pm.DomainID), - attribute.String("invited_by", pm.InvitedBy), - )) - defer span.End() - - return tm.svc.ListInvitations(ctx, session, pm) -} - -func (tm *tracingMiddleware) ListDomainInvitations(ctx context.Context, session authn.Session, pm domains.InvitationPageMeta) (invs domains.InvitationPage, err error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "list_domain_invitations", trace.WithAttributes( - attribute.Int("limit", int(pm.Limit)), - attribute.Int("offset", int(pm.Offset)), - attribute.String("domain_id", session.DomainID), - )) - defer span.End() - - return tm.svc.ListDomainInvitations(ctx, session, pm) -} - -func (tm *tracingMiddleware) AcceptInvitation(ctx context.Context, session authn.Session, domainID string) (inv domains.Invitation, err error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "accept_invitation", trace.WithAttributes( - attribute.String("domain_id", domainID), - )) - defer span.End() - - return tm.svc.AcceptInvitation(ctx, session, domainID) -} - -func (tm *tracingMiddleware) RejectInvitation(ctx context.Context, session authn.Session, domainID string) (inv domains.Invitation, err error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "reject_invitation", trace.WithAttributes( - attribute.String("domain_id", domainID), - )) - defer span.End() - - return tm.svc.RejectInvitation(ctx, session, domainID) -} - -func (tm *tracingMiddleware) DeleteInvitation(ctx context.Context, session authn.Session, inviteeUserID, domainID string) (err error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "delete_invitation", trace.WithAttributes( - attribute.String("invitee_user_id", inviteeUserID), - attribute.String("domain_id", domainID), - )) - defer span.End() - - return tm.svc.DeleteInvitation(ctx, session, inviteeUserID, domainID) -} diff --git a/domains/mocks/cache.go b/domains/mocks/cache.go deleted file mode 100644 index 8a7912143..000000000 --- a/domains/mocks/cache.go +++ /dev/null @@ -1,415 +0,0 @@ -// Copyright (c) Abstract Machines - -// SPDX-License-Identifier: Apache-2.0 - -// Code generated by mockery; DO NOT EDIT. -// github.com/vektra/mockery -// template: testify - -package mocks - -import ( - "context" - - "github.com/absmach/magistrala/domains" - mock "github.com/stretchr/testify/mock" -) - -// NewCache creates a new instance of Cache. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. -// The first argument is typically a *testing.T value. -func NewCache(t interface { - mock.TestingT - Cleanup(func()) -}) *Cache { - mock := &Cache{} - mock.Mock.Test(t) - - t.Cleanup(func() { mock.AssertExpectations(t) }) - - return mock -} - -// Cache is an autogenerated mock type for the Cache type -type Cache struct { - mock.Mock -} - -type Cache_Expecter struct { - mock *mock.Mock -} - -func (_m *Cache) EXPECT() *Cache_Expecter { - return &Cache_Expecter{mock: &_m.Mock} -} - -// ID provides a mock function for the type Cache -func (_mock *Cache) ID(ctx context.Context, route string) (string, error) { - ret := _mock.Called(ctx, route) - - if len(ret) == 0 { - panic("no return value specified for ID") - } - - var r0 string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) (string, error)); ok { - return returnFunc(ctx, route) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string) string); ok { - r0 = returnFunc(ctx, route) - } else { - r0 = ret.Get(0).(string) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = returnFunc(ctx, route) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Cache_ID_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ID' -type Cache_ID_Call struct { - *mock.Call -} - -// ID is a helper method to define mock.On call -// - ctx context.Context -// - route string -func (_e *Cache_Expecter) ID(ctx interface{}, route interface{}) *Cache_ID_Call { - return &Cache_ID_Call{Call: _e.mock.On("ID", ctx, route)} -} - -func (_c *Cache_ID_Call) Run(run func(ctx context.Context, route string)) *Cache_ID_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Cache_ID_Call) Return(s string, err error) *Cache_ID_Call { - _c.Call.Return(s, err) - return _c -} - -func (_c *Cache_ID_Call) RunAndReturn(run func(ctx context.Context, route string) (string, error)) *Cache_ID_Call { - _c.Call.Return(run) - return _c -} - -// RemoveID provides a mock function for the type Cache -func (_mock *Cache) RemoveID(ctx context.Context, route string) error { - ret := _mock.Called(ctx, route) - - if len(ret) == 0 { - panic("no return value specified for RemoveID") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) error); ok { - r0 = returnFunc(ctx, route) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Cache_RemoveID_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveID' -type Cache_RemoveID_Call struct { - *mock.Call -} - -// RemoveID is a helper method to define mock.On call -// - ctx context.Context -// - route string -func (_e *Cache_Expecter) RemoveID(ctx interface{}, route interface{}) *Cache_RemoveID_Call { - return &Cache_RemoveID_Call{Call: _e.mock.On("RemoveID", ctx, route)} -} - -func (_c *Cache_RemoveID_Call) Run(run func(ctx context.Context, route string)) *Cache_RemoveID_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Cache_RemoveID_Call) Return(err error) *Cache_RemoveID_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Cache_RemoveID_Call) RunAndReturn(run func(ctx context.Context, route string) error) *Cache_RemoveID_Call { - _c.Call.Return(run) - return _c -} - -// RemoveStatus provides a mock function for the type Cache -func (_mock *Cache) RemoveStatus(ctx context.Context, domainID string) error { - ret := _mock.Called(ctx, domainID) - - if len(ret) == 0 { - panic("no return value specified for RemoveStatus") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) error); ok { - r0 = returnFunc(ctx, domainID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Cache_RemoveStatus_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveStatus' -type Cache_RemoveStatus_Call struct { - *mock.Call -} - -// RemoveStatus is a helper method to define mock.On call -// - ctx context.Context -// - domainID string -func (_e *Cache_Expecter) RemoveStatus(ctx interface{}, domainID interface{}) *Cache_RemoveStatus_Call { - return &Cache_RemoveStatus_Call{Call: _e.mock.On("RemoveStatus", ctx, domainID)} -} - -func (_c *Cache_RemoveStatus_Call) Run(run func(ctx context.Context, domainID string)) *Cache_RemoveStatus_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Cache_RemoveStatus_Call) Return(err error) *Cache_RemoveStatus_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Cache_RemoveStatus_Call) RunAndReturn(run func(ctx context.Context, domainID string) error) *Cache_RemoveStatus_Call { - _c.Call.Return(run) - return _c -} - -// SaveID provides a mock function for the type Cache -func (_mock *Cache) SaveID(ctx context.Context, route string, domainID string) error { - ret := _mock.Called(ctx, route, domainID) - - if len(ret) == 0 { - panic("no return value specified for SaveID") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) error); ok { - r0 = returnFunc(ctx, route, domainID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Cache_SaveID_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'SaveID' -type Cache_SaveID_Call struct { - *mock.Call -} - -// SaveID is a helper method to define mock.On call -// - ctx context.Context -// - route string -// - domainID string -func (_e *Cache_Expecter) SaveID(ctx interface{}, route interface{}, domainID interface{}) *Cache_SaveID_Call { - return &Cache_SaveID_Call{Call: _e.mock.On("SaveID", ctx, route, domainID)} -} - -func (_c *Cache_SaveID_Call) Run(run func(ctx context.Context, route string, domainID string)) *Cache_SaveID_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Cache_SaveID_Call) Return(err error) *Cache_SaveID_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Cache_SaveID_Call) RunAndReturn(run func(ctx context.Context, route string, domainID string) error) *Cache_SaveID_Call { - _c.Call.Return(run) - return _c -} - -// SaveStatus provides a mock function for the type Cache -func (_mock *Cache) SaveStatus(ctx context.Context, domainID string, status domains.Status) error { - ret := _mock.Called(ctx, domainID, status) - - if len(ret) == 0 { - panic("no return value specified for SaveStatus") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, domains.Status) error); ok { - r0 = returnFunc(ctx, domainID, status) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Cache_SaveStatus_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'SaveStatus' -type Cache_SaveStatus_Call struct { - *mock.Call -} - -// SaveStatus is a helper method to define mock.On call -// - ctx context.Context -// - domainID string -// - status domains.Status -func (_e *Cache_Expecter) SaveStatus(ctx interface{}, domainID interface{}, status interface{}) *Cache_SaveStatus_Call { - return &Cache_SaveStatus_Call{Call: _e.mock.On("SaveStatus", ctx, domainID, status)} -} - -func (_c *Cache_SaveStatus_Call) Run(run func(ctx context.Context, domainID string, status domains.Status)) *Cache_SaveStatus_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 domains.Status - if args[2] != nil { - arg2 = args[2].(domains.Status) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Cache_SaveStatus_Call) Return(err error) *Cache_SaveStatus_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Cache_SaveStatus_Call) RunAndReturn(run func(ctx context.Context, domainID string, status domains.Status) error) *Cache_SaveStatus_Call { - _c.Call.Return(run) - return _c -} - -// Status provides a mock function for the type Cache -func (_mock *Cache) Status(ctx context.Context, domainID string) (domains.Status, error) { - ret := _mock.Called(ctx, domainID) - - if len(ret) == 0 { - panic("no return value specified for Status") - } - - var r0 domains.Status - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) (domains.Status, error)); ok { - return returnFunc(ctx, domainID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string) domains.Status); ok { - r0 = returnFunc(ctx, domainID) - } else { - r0 = ret.Get(0).(domains.Status) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = returnFunc(ctx, domainID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Cache_Status_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Status' -type Cache_Status_Call struct { - *mock.Call -} - -// Status is a helper method to define mock.On call -// - ctx context.Context -// - domainID string -func (_e *Cache_Expecter) Status(ctx interface{}, domainID interface{}) *Cache_Status_Call { - return &Cache_Status_Call{Call: _e.mock.On("Status", ctx, domainID)} -} - -func (_c *Cache_Status_Call) Run(run func(ctx context.Context, domainID string)) *Cache_Status_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Cache_Status_Call) Return(status domains.Status, err error) *Cache_Status_Call { - _c.Call.Return(status, err) - return _c -} - -func (_c *Cache_Status_Call) RunAndReturn(run func(ctx context.Context, domainID string) (domains.Status, error)) *Cache_Status_Call { - _c.Call.Return(run) - return _c -} diff --git a/domains/mocks/domains_client.go b/domains/mocks/domains_client.go deleted file mode 100644 index 111b4c1d9..000000000 --- a/domains/mocks/domains_client.go +++ /dev/null @@ -1,294 +0,0 @@ -// Copyright (c) Abstract Machines - -// SPDX-License-Identifier: Apache-2.0 - -// Code generated by mockery; DO NOT EDIT. -// github.com/vektra/mockery -// template: testify - -package mocks - -import ( - "context" - - v10 "github.com/absmach/magistrala/api/grpc/common/v1" - "github.com/absmach/magistrala/api/grpc/domains/v1" - mock "github.com/stretchr/testify/mock" - "google.golang.org/grpc" -) - -// NewDomainsServiceClient creates a new instance of DomainsServiceClient. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. -// The first argument is typically a *testing.T value. -func NewDomainsServiceClient(t interface { - mock.TestingT - Cleanup(func()) -}) *DomainsServiceClient { - mock := &DomainsServiceClient{} - mock.Mock.Test(t) - - t.Cleanup(func() { mock.AssertExpectations(t) }) - - return mock -} - -// DomainsServiceClient is an autogenerated mock type for the DomainsServiceClient type -type DomainsServiceClient struct { - mock.Mock -} - -type DomainsServiceClient_Expecter struct { - mock *mock.Mock -} - -func (_m *DomainsServiceClient) EXPECT() *DomainsServiceClient_Expecter { - return &DomainsServiceClient_Expecter{mock: &_m.Mock} -} - -// DeleteUserFromDomains provides a mock function for the type DomainsServiceClient -func (_mock *DomainsServiceClient) DeleteUserFromDomains(ctx context.Context, in *v1.DeleteUserReq, opts ...grpc.CallOption) (*v1.DeleteUserRes, error) { - var tmpRet mock.Arguments - if len(opts) > 0 { - tmpRet = _mock.Called(ctx, in, opts) - } else { - tmpRet = _mock.Called(ctx, in) - } - ret := tmpRet - - if len(ret) == 0 { - panic("no return value specified for DeleteUserFromDomains") - } - - var r0 *v1.DeleteUserRes - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, *v1.DeleteUserReq, ...grpc.CallOption) (*v1.DeleteUserRes, error)); ok { - return returnFunc(ctx, in, opts...) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, *v1.DeleteUserReq, ...grpc.CallOption) *v1.DeleteUserRes); ok { - r0 = returnFunc(ctx, in, opts...) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).(*v1.DeleteUserRes) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, *v1.DeleteUserReq, ...grpc.CallOption) error); ok { - r1 = returnFunc(ctx, in, opts...) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// DomainsServiceClient_DeleteUserFromDomains_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'DeleteUserFromDomains' -type DomainsServiceClient_DeleteUserFromDomains_Call struct { - *mock.Call -} - -// DeleteUserFromDomains is a helper method to define mock.On call -// - ctx context.Context -// - in *v1.DeleteUserReq -// - opts ...grpc.CallOption -func (_e *DomainsServiceClient_Expecter) DeleteUserFromDomains(ctx interface{}, in interface{}, opts ...interface{}) *DomainsServiceClient_DeleteUserFromDomains_Call { - return &DomainsServiceClient_DeleteUserFromDomains_Call{Call: _e.mock.On("DeleteUserFromDomains", - append([]interface{}{ctx, in}, opts...)...)} -} - -func (_c *DomainsServiceClient_DeleteUserFromDomains_Call) Run(run func(ctx context.Context, in *v1.DeleteUserReq, opts ...grpc.CallOption)) *DomainsServiceClient_DeleteUserFromDomains_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 *v1.DeleteUserReq - if args[1] != nil { - arg1 = args[1].(*v1.DeleteUserReq) - } - var arg2 []grpc.CallOption - var variadicArgs []grpc.CallOption - if len(args) > 2 { - variadicArgs = args[2].([]grpc.CallOption) - } - arg2 = variadicArgs - run( - arg0, - arg1, - arg2..., - ) - }) - return _c -} - -func (_c *DomainsServiceClient_DeleteUserFromDomains_Call) Return(deleteUserRes *v1.DeleteUserRes, err error) *DomainsServiceClient_DeleteUserFromDomains_Call { - _c.Call.Return(deleteUserRes, err) - return _c -} - -func (_c *DomainsServiceClient_DeleteUserFromDomains_Call) RunAndReturn(run func(ctx context.Context, in *v1.DeleteUserReq, opts ...grpc.CallOption) (*v1.DeleteUserRes, error)) *DomainsServiceClient_DeleteUserFromDomains_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveIDByRoute provides a mock function for the type DomainsServiceClient -func (_mock *DomainsServiceClient) RetrieveIDByRoute(ctx context.Context, in *v10.RetrieveIDByRouteReq, opts ...grpc.CallOption) (*v10.RetrieveEntityRes, error) { - var tmpRet mock.Arguments - if len(opts) > 0 { - tmpRet = _mock.Called(ctx, in, opts) - } else { - tmpRet = _mock.Called(ctx, in) - } - ret := tmpRet - - if len(ret) == 0 { - panic("no return value specified for RetrieveIDByRoute") - } - - var r0 *v10.RetrieveEntityRes - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, *v10.RetrieveIDByRouteReq, ...grpc.CallOption) (*v10.RetrieveEntityRes, error)); ok { - return returnFunc(ctx, in, opts...) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, *v10.RetrieveIDByRouteReq, ...grpc.CallOption) *v10.RetrieveEntityRes); ok { - r0 = returnFunc(ctx, in, opts...) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).(*v10.RetrieveEntityRes) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, *v10.RetrieveIDByRouteReq, ...grpc.CallOption) error); ok { - r1 = returnFunc(ctx, in, opts...) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// DomainsServiceClient_RetrieveIDByRoute_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveIDByRoute' -type DomainsServiceClient_RetrieveIDByRoute_Call struct { - *mock.Call -} - -// RetrieveIDByRoute is a helper method to define mock.On call -// - ctx context.Context -// - in *v10.RetrieveIDByRouteReq -// - opts ...grpc.CallOption -func (_e *DomainsServiceClient_Expecter) RetrieveIDByRoute(ctx interface{}, in interface{}, opts ...interface{}) *DomainsServiceClient_RetrieveIDByRoute_Call { - return &DomainsServiceClient_RetrieveIDByRoute_Call{Call: _e.mock.On("RetrieveIDByRoute", - append([]interface{}{ctx, in}, opts...)...)} -} - -func (_c *DomainsServiceClient_RetrieveIDByRoute_Call) Run(run func(ctx context.Context, in *v10.RetrieveIDByRouteReq, opts ...grpc.CallOption)) *DomainsServiceClient_RetrieveIDByRoute_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 *v10.RetrieveIDByRouteReq - if args[1] != nil { - arg1 = args[1].(*v10.RetrieveIDByRouteReq) - } - var arg2 []grpc.CallOption - var variadicArgs []grpc.CallOption - if len(args) > 2 { - variadicArgs = args[2].([]grpc.CallOption) - } - arg2 = variadicArgs - run( - arg0, - arg1, - arg2..., - ) - }) - return _c -} - -func (_c *DomainsServiceClient_RetrieveIDByRoute_Call) Return(retrieveEntityRes *v10.RetrieveEntityRes, err error) *DomainsServiceClient_RetrieveIDByRoute_Call { - _c.Call.Return(retrieveEntityRes, err) - return _c -} - -func (_c *DomainsServiceClient_RetrieveIDByRoute_Call) RunAndReturn(run func(ctx context.Context, in *v10.RetrieveIDByRouteReq, opts ...grpc.CallOption) (*v10.RetrieveEntityRes, error)) *DomainsServiceClient_RetrieveIDByRoute_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveStatus provides a mock function for the type DomainsServiceClient -func (_mock *DomainsServiceClient) RetrieveStatus(ctx context.Context, in *v10.RetrieveEntityReq, opts ...grpc.CallOption) (*v10.RetrieveEntityRes, error) { - var tmpRet mock.Arguments - if len(opts) > 0 { - tmpRet = _mock.Called(ctx, in, opts) - } else { - tmpRet = _mock.Called(ctx, in) - } - ret := tmpRet - - if len(ret) == 0 { - panic("no return value specified for RetrieveStatus") - } - - var r0 *v10.RetrieveEntityRes - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, *v10.RetrieveEntityReq, ...grpc.CallOption) (*v10.RetrieveEntityRes, error)); ok { - return returnFunc(ctx, in, opts...) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, *v10.RetrieveEntityReq, ...grpc.CallOption) *v10.RetrieveEntityRes); ok { - r0 = returnFunc(ctx, in, opts...) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).(*v10.RetrieveEntityRes) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, *v10.RetrieveEntityReq, ...grpc.CallOption) error); ok { - r1 = returnFunc(ctx, in, opts...) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// DomainsServiceClient_RetrieveStatus_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveStatus' -type DomainsServiceClient_RetrieveStatus_Call struct { - *mock.Call -} - -// RetrieveStatus is a helper method to define mock.On call -// - ctx context.Context -// - in *v10.RetrieveEntityReq -// - opts ...grpc.CallOption -func (_e *DomainsServiceClient_Expecter) RetrieveStatus(ctx interface{}, in interface{}, opts ...interface{}) *DomainsServiceClient_RetrieveStatus_Call { - return &DomainsServiceClient_RetrieveStatus_Call{Call: _e.mock.On("RetrieveStatus", - append([]interface{}{ctx, in}, opts...)...)} -} - -func (_c *DomainsServiceClient_RetrieveStatus_Call) Run(run func(ctx context.Context, in *v10.RetrieveEntityReq, opts ...grpc.CallOption)) *DomainsServiceClient_RetrieveStatus_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 *v10.RetrieveEntityReq - if args[1] != nil { - arg1 = args[1].(*v10.RetrieveEntityReq) - } - var arg2 []grpc.CallOption - var variadicArgs []grpc.CallOption - if len(args) > 2 { - variadicArgs = args[2].([]grpc.CallOption) - } - arg2 = variadicArgs - run( - arg0, - arg1, - arg2..., - ) - }) - return _c -} - -func (_c *DomainsServiceClient_RetrieveStatus_Call) Return(retrieveEntityRes *v10.RetrieveEntityRes, err error) *DomainsServiceClient_RetrieveStatus_Call { - _c.Call.Return(retrieveEntityRes, err) - return _c -} - -func (_c *DomainsServiceClient_RetrieveStatus_Call) RunAndReturn(run func(ctx context.Context, in *v10.RetrieveEntityReq, opts ...grpc.CallOption) (*v10.RetrieveEntityRes, error)) *DomainsServiceClient_RetrieveStatus_Call { - _c.Call.Return(run) - return _c -} diff --git a/domains/mocks/repository.go b/domains/mocks/repository.go deleted file mode 100644 index bf6747ce4..000000000 --- a/domains/mocks/repository.go +++ /dev/null @@ -1,2309 +0,0 @@ -// Copyright (c) Abstract Machines - -// SPDX-License-Identifier: Apache-2.0 - -// Code generated by mockery; DO NOT EDIT. -// github.com/vektra/mockery -// template: testify - -package mocks - -import ( - "context" - - "github.com/absmach/magistrala/domains" - "github.com/absmach/magistrala/pkg/roles" - mock "github.com/stretchr/testify/mock" -) - -// NewRepository creates a new instance of Repository. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. -// The first argument is typically a *testing.T value. -func NewRepository(t interface { - mock.TestingT - Cleanup(func()) -}) *Repository { - mock := &Repository{} - mock.Mock.Test(t) - - t.Cleanup(func() { mock.AssertExpectations(t) }) - - return mock -} - -// Repository is an autogenerated mock type for the Repository type -type Repository struct { - mock.Mock -} - -type Repository_Expecter struct { - mock *mock.Mock -} - -func (_m *Repository) EXPECT() *Repository_Expecter { - return &Repository_Expecter{mock: &_m.Mock} -} - -// AddRoles provides a mock function for the type Repository -func (_mock *Repository) AddRoles(ctx context.Context, rps []roles.RoleProvision) ([]roles.RoleProvision, error) { - ret := _mock.Called(ctx, rps) - - if len(ret) == 0 { - panic("no return value specified for AddRoles") - } - - var r0 []roles.RoleProvision - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, []roles.RoleProvision) ([]roles.RoleProvision, error)); ok { - return returnFunc(ctx, rps) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, []roles.RoleProvision) []roles.RoleProvision); ok { - r0 = returnFunc(ctx, rps) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]roles.RoleProvision) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, []roles.RoleProvision) error); ok { - r1 = returnFunc(ctx, rps) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_AddRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'AddRoles' -type Repository_AddRoles_Call struct { - *mock.Call -} - -// AddRoles is a helper method to define mock.On call -// - ctx context.Context -// - rps []roles.RoleProvision -func (_e *Repository_Expecter) AddRoles(ctx interface{}, rps interface{}) *Repository_AddRoles_Call { - return &Repository_AddRoles_Call{Call: _e.mock.On("AddRoles", ctx, rps)} -} - -func (_c *Repository_AddRoles_Call) Run(run func(ctx context.Context, rps []roles.RoleProvision)) *Repository_AddRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 []roles.RoleProvision - if args[1] != nil { - arg1 = args[1].([]roles.RoleProvision) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_AddRoles_Call) Return(roleProvisions []roles.RoleProvision, err error) *Repository_AddRoles_Call { - _c.Call.Return(roleProvisions, err) - return _c -} - -func (_c *Repository_AddRoles_Call) RunAndReturn(run func(ctx context.Context, rps []roles.RoleProvision) ([]roles.RoleProvision, error)) *Repository_AddRoles_Call { - _c.Call.Return(run) - return _c -} - -// DeleteDomain provides a mock function for the type Repository -func (_mock *Repository) DeleteDomain(ctx context.Context, id string) error { - ret := _mock.Called(ctx, id) - - if len(ret) == 0 { - panic("no return value specified for DeleteDomain") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) error); ok { - r0 = returnFunc(ctx, id) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_DeleteDomain_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'DeleteDomain' -type Repository_DeleteDomain_Call struct { - *mock.Call -} - -// DeleteDomain is a helper method to define mock.On call -// - ctx context.Context -// - id string -func (_e *Repository_Expecter) DeleteDomain(ctx interface{}, id interface{}) *Repository_DeleteDomain_Call { - return &Repository_DeleteDomain_Call{Call: _e.mock.On("DeleteDomain", ctx, id)} -} - -func (_c *Repository_DeleteDomain_Call) Run(run func(ctx context.Context, id string)) *Repository_DeleteDomain_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_DeleteDomain_Call) Return(err error) *Repository_DeleteDomain_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_DeleteDomain_Call) RunAndReturn(run func(ctx context.Context, id string) error) *Repository_DeleteDomain_Call { - _c.Call.Return(run) - return _c -} - -// DeleteUsersInvitations provides a mock function for the type Repository -func (_mock *Repository) DeleteUsersInvitations(ctx context.Context, domainID string, userID ...string) error { - var tmpRet mock.Arguments - if len(userID) > 0 { - tmpRet = _mock.Called(ctx, domainID, userID) - } else { - tmpRet = _mock.Called(ctx, domainID) - } - ret := tmpRet - - if len(ret) == 0 { - panic("no return value specified for DeleteUsersInvitations") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, ...string) error); ok { - r0 = returnFunc(ctx, domainID, userID...) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_DeleteUsersInvitations_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'DeleteUsersInvitations' -type Repository_DeleteUsersInvitations_Call struct { - *mock.Call -} - -// DeleteUsersInvitations is a helper method to define mock.On call -// - ctx context.Context -// - domainID string -// - userID ...string -func (_e *Repository_Expecter) DeleteUsersInvitations(ctx interface{}, domainID interface{}, userID ...interface{}) *Repository_DeleteUsersInvitations_Call { - return &Repository_DeleteUsersInvitations_Call{Call: _e.mock.On("DeleteUsersInvitations", - append([]interface{}{ctx, domainID}, userID...)...)} -} - -func (_c *Repository_DeleteUsersInvitations_Call) Run(run func(ctx context.Context, domainID string, userID ...string)) *Repository_DeleteUsersInvitations_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 []string - var variadicArgs []string - if len(args) > 2 { - variadicArgs = args[2].([]string) - } - arg2 = variadicArgs - run( - arg0, - arg1, - arg2..., - ) - }) - return _c -} - -func (_c *Repository_DeleteUsersInvitations_Call) Return(err error) *Repository_DeleteUsersInvitations_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_DeleteUsersInvitations_Call) RunAndReturn(run func(ctx context.Context, domainID string, userID ...string) error) *Repository_DeleteUsersInvitations_Call { - _c.Call.Return(run) - return _c -} - -// ListDomains provides a mock function for the type Repository -func (_mock *Repository) ListDomains(ctx context.Context, pm domains.Page) (domains.DomainsPage, error) { - ret := _mock.Called(ctx, pm) - - if len(ret) == 0 { - panic("no return value specified for ListDomains") - } - - var r0 domains.DomainsPage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, domains.Page) (domains.DomainsPage, error)); ok { - return returnFunc(ctx, pm) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, domains.Page) domains.DomainsPage); ok { - r0 = returnFunc(ctx, pm) - } else { - r0 = ret.Get(0).(domains.DomainsPage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, domains.Page) error); ok { - r1 = returnFunc(ctx, pm) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_ListDomains_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ListDomains' -type Repository_ListDomains_Call struct { - *mock.Call -} - -// ListDomains is a helper method to define mock.On call -// - ctx context.Context -// - pm domains.Page -func (_e *Repository_Expecter) ListDomains(ctx interface{}, pm interface{}) *Repository_ListDomains_Call { - return &Repository_ListDomains_Call{Call: _e.mock.On("ListDomains", ctx, pm)} -} - -func (_c *Repository_ListDomains_Call) Run(run func(ctx context.Context, pm domains.Page)) *Repository_ListDomains_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 domains.Page - if args[1] != nil { - arg1 = args[1].(domains.Page) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_ListDomains_Call) Return(domainsPage domains.DomainsPage, err error) *Repository_ListDomains_Call { - _c.Call.Return(domainsPage, err) - return _c -} - -func (_c *Repository_ListDomains_Call) RunAndReturn(run func(ctx context.Context, pm domains.Page) (domains.DomainsPage, error)) *Repository_ListDomains_Call { - _c.Call.Return(run) - return _c -} - -// ListEntityMembers provides a mock function for the type Repository -func (_mock *Repository) ListEntityMembers(ctx context.Context, entityID string, pageQuery roles.MembersRolePageQuery) (roles.MembersRolePage, error) { - ret := _mock.Called(ctx, entityID, pageQuery) - - if len(ret) == 0 { - panic("no return value specified for ListEntityMembers") - } - - var r0 roles.MembersRolePage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, roles.MembersRolePageQuery) (roles.MembersRolePage, error)); ok { - return returnFunc(ctx, entityID, pageQuery) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, roles.MembersRolePageQuery) roles.MembersRolePage); ok { - r0 = returnFunc(ctx, entityID, pageQuery) - } else { - r0 = ret.Get(0).(roles.MembersRolePage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, roles.MembersRolePageQuery) error); ok { - r1 = returnFunc(ctx, entityID, pageQuery) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_ListEntityMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ListEntityMembers' -type Repository_ListEntityMembers_Call struct { - *mock.Call -} - -// ListEntityMembers is a helper method to define mock.On call -// - ctx context.Context -// - entityID string -// - pageQuery roles.MembersRolePageQuery -func (_e *Repository_Expecter) ListEntityMembers(ctx interface{}, entityID interface{}, pageQuery interface{}) *Repository_ListEntityMembers_Call { - return &Repository_ListEntityMembers_Call{Call: _e.mock.On("ListEntityMembers", ctx, entityID, pageQuery)} -} - -func (_c *Repository_ListEntityMembers_Call) Run(run func(ctx context.Context, entityID string, pageQuery roles.MembersRolePageQuery)) *Repository_ListEntityMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 roles.MembersRolePageQuery - if args[2] != nil { - arg2 = args[2].(roles.MembersRolePageQuery) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_ListEntityMembers_Call) Return(membersRolePage roles.MembersRolePage, err error) *Repository_ListEntityMembers_Call { - _c.Call.Return(membersRolePage, err) - return _c -} - -func (_c *Repository_ListEntityMembers_Call) RunAndReturn(run func(ctx context.Context, entityID string, pageQuery roles.MembersRolePageQuery) (roles.MembersRolePage, error)) *Repository_ListEntityMembers_Call { - _c.Call.Return(run) - return _c -} - -// RemoveEntityMembers provides a mock function for the type Repository -func (_mock *Repository) RemoveEntityMembers(ctx context.Context, entityID string, members []string) error { - ret := _mock.Called(ctx, entityID, members) - - if len(ret) == 0 { - panic("no return value specified for RemoveEntityMembers") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, []string) error); ok { - r0 = returnFunc(ctx, entityID, members) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RemoveEntityMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveEntityMembers' -type Repository_RemoveEntityMembers_Call struct { - *mock.Call -} - -// RemoveEntityMembers is a helper method to define mock.On call -// - ctx context.Context -// - entityID string -// - members []string -func (_e *Repository_Expecter) RemoveEntityMembers(ctx interface{}, entityID interface{}, members interface{}) *Repository_RemoveEntityMembers_Call { - return &Repository_RemoveEntityMembers_Call{Call: _e.mock.On("RemoveEntityMembers", ctx, entityID, members)} -} - -func (_c *Repository_RemoveEntityMembers_Call) Run(run func(ctx context.Context, entityID string, members []string)) *Repository_RemoveEntityMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RemoveEntityMembers_Call) Return(err error) *Repository_RemoveEntityMembers_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RemoveEntityMembers_Call) RunAndReturn(run func(ctx context.Context, entityID string, members []string) error) *Repository_RemoveEntityMembers_Call { - _c.Call.Return(run) - return _c -} - -// RemoveMemberFromAllRoles provides a mock function for the type Repository -func (_mock *Repository) RemoveMemberFromAllRoles(ctx context.Context, memberID string) error { - ret := _mock.Called(ctx, memberID) - - if len(ret) == 0 { - panic("no return value specified for RemoveMemberFromAllRoles") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) error); ok { - r0 = returnFunc(ctx, memberID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RemoveMemberFromAllRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveMemberFromAllRoles' -type Repository_RemoveMemberFromAllRoles_Call struct { - *mock.Call -} - -// RemoveMemberFromAllRoles is a helper method to define mock.On call -// - ctx context.Context -// - memberID string -func (_e *Repository_Expecter) RemoveMemberFromAllRoles(ctx interface{}, memberID interface{}) *Repository_RemoveMemberFromAllRoles_Call { - return &Repository_RemoveMemberFromAllRoles_Call{Call: _e.mock.On("RemoveMemberFromAllRoles", ctx, memberID)} -} - -func (_c *Repository_RemoveMemberFromAllRoles_Call) Run(run func(ctx context.Context, memberID string)) *Repository_RemoveMemberFromAllRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RemoveMemberFromAllRoles_Call) Return(err error) *Repository_RemoveMemberFromAllRoles_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RemoveMemberFromAllRoles_Call) RunAndReturn(run func(ctx context.Context, memberID string) error) *Repository_RemoveMemberFromAllRoles_Call { - _c.Call.Return(run) - return _c -} - -// RemoveRoles provides a mock function for the type Repository -func (_mock *Repository) RemoveRoles(ctx context.Context, roleIDs []string) error { - ret := _mock.Called(ctx, roleIDs) - - if len(ret) == 0 { - panic("no return value specified for RemoveRoles") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, []string) error); ok { - r0 = returnFunc(ctx, roleIDs) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RemoveRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveRoles' -type Repository_RemoveRoles_Call struct { - *mock.Call -} - -// RemoveRoles is a helper method to define mock.On call -// - ctx context.Context -// - roleIDs []string -func (_e *Repository_Expecter) RemoveRoles(ctx interface{}, roleIDs interface{}) *Repository_RemoveRoles_Call { - return &Repository_RemoveRoles_Call{Call: _e.mock.On("RemoveRoles", ctx, roleIDs)} -} - -func (_c *Repository_RemoveRoles_Call) Run(run func(ctx context.Context, roleIDs []string)) *Repository_RemoveRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 []string - if args[1] != nil { - arg1 = args[1].([]string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RemoveRoles_Call) Return(err error) *Repository_RemoveRoles_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RemoveRoles_Call) RunAndReturn(run func(ctx context.Context, roleIDs []string) error) *Repository_RemoveRoles_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveAllDomainsByIDs provides a mock function for the type Repository -func (_mock *Repository) RetrieveAllDomainsByIDs(ctx context.Context, pm domains.Page) (domains.DomainsPage, error) { - ret := _mock.Called(ctx, pm) - - if len(ret) == 0 { - panic("no return value specified for RetrieveAllDomainsByIDs") - } - - var r0 domains.DomainsPage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, domains.Page) (domains.DomainsPage, error)); ok { - return returnFunc(ctx, pm) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, domains.Page) domains.DomainsPage); ok { - r0 = returnFunc(ctx, pm) - } else { - r0 = ret.Get(0).(domains.DomainsPage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, domains.Page) error); ok { - r1 = returnFunc(ctx, pm) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveAllDomainsByIDs_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveAllDomainsByIDs' -type Repository_RetrieveAllDomainsByIDs_Call struct { - *mock.Call -} - -// RetrieveAllDomainsByIDs is a helper method to define mock.On call -// - ctx context.Context -// - pm domains.Page -func (_e *Repository_Expecter) RetrieveAllDomainsByIDs(ctx interface{}, pm interface{}) *Repository_RetrieveAllDomainsByIDs_Call { - return &Repository_RetrieveAllDomainsByIDs_Call{Call: _e.mock.On("RetrieveAllDomainsByIDs", ctx, pm)} -} - -func (_c *Repository_RetrieveAllDomainsByIDs_Call) Run(run func(ctx context.Context, pm domains.Page)) *Repository_RetrieveAllDomainsByIDs_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 domains.Page - if args[1] != nil { - arg1 = args[1].(domains.Page) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RetrieveAllDomainsByIDs_Call) Return(domainsPage domains.DomainsPage, err error) *Repository_RetrieveAllDomainsByIDs_Call { - _c.Call.Return(domainsPage, err) - return _c -} - -func (_c *Repository_RetrieveAllDomainsByIDs_Call) RunAndReturn(run func(ctx context.Context, pm domains.Page) (domains.DomainsPage, error)) *Repository_RetrieveAllDomainsByIDs_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveAllInvitations provides a mock function for the type Repository -func (_mock *Repository) RetrieveAllInvitations(ctx context.Context, page domains.InvitationPageMeta) (domains.InvitationPage, error) { - ret := _mock.Called(ctx, page) - - if len(ret) == 0 { - panic("no return value specified for RetrieveAllInvitations") - } - - var r0 domains.InvitationPage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, domains.InvitationPageMeta) (domains.InvitationPage, error)); ok { - return returnFunc(ctx, page) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, domains.InvitationPageMeta) domains.InvitationPage); ok { - r0 = returnFunc(ctx, page) - } else { - r0 = ret.Get(0).(domains.InvitationPage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, domains.InvitationPageMeta) error); ok { - r1 = returnFunc(ctx, page) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveAllInvitations_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveAllInvitations' -type Repository_RetrieveAllInvitations_Call struct { - *mock.Call -} - -// RetrieveAllInvitations is a helper method to define mock.On call -// - ctx context.Context -// - page domains.InvitationPageMeta -func (_e *Repository_Expecter) RetrieveAllInvitations(ctx interface{}, page interface{}) *Repository_RetrieveAllInvitations_Call { - return &Repository_RetrieveAllInvitations_Call{Call: _e.mock.On("RetrieveAllInvitations", ctx, page)} -} - -func (_c *Repository_RetrieveAllInvitations_Call) Run(run func(ctx context.Context, page domains.InvitationPageMeta)) *Repository_RetrieveAllInvitations_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 domains.InvitationPageMeta - if args[1] != nil { - arg1 = args[1].(domains.InvitationPageMeta) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RetrieveAllInvitations_Call) Return(invitations domains.InvitationPage, err error) *Repository_RetrieveAllInvitations_Call { - _c.Call.Return(invitations, err) - return _c -} - -func (_c *Repository_RetrieveAllInvitations_Call) RunAndReturn(run func(ctx context.Context, page domains.InvitationPageMeta) (domains.InvitationPage, error)) *Repository_RetrieveAllInvitations_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveAllRoles provides a mock function for the type Repository -func (_mock *Repository) RetrieveAllRoles(ctx context.Context, entityID string, limit uint64, offset uint64) (roles.RolePage, error) { - ret := _mock.Called(ctx, entityID, limit, offset) - - if len(ret) == 0 { - panic("no return value specified for RetrieveAllRoles") - } - - var r0 roles.RolePage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, uint64, uint64) (roles.RolePage, error)); ok { - return returnFunc(ctx, entityID, limit, offset) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, uint64, uint64) roles.RolePage); ok { - r0 = returnFunc(ctx, entityID, limit, offset) - } else { - r0 = ret.Get(0).(roles.RolePage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, uint64, uint64) error); ok { - r1 = returnFunc(ctx, entityID, limit, offset) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveAllRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveAllRoles' -type Repository_RetrieveAllRoles_Call struct { - *mock.Call -} - -// RetrieveAllRoles is a helper method to define mock.On call -// - ctx context.Context -// - entityID string -// - limit uint64 -// - offset uint64 -func (_e *Repository_Expecter) RetrieveAllRoles(ctx interface{}, entityID interface{}, limit interface{}, offset interface{}) *Repository_RetrieveAllRoles_Call { - return &Repository_RetrieveAllRoles_Call{Call: _e.mock.On("RetrieveAllRoles", ctx, entityID, limit, offset)} -} - -func (_c *Repository_RetrieveAllRoles_Call) Run(run func(ctx context.Context, entityID string, limit uint64, offset uint64)) *Repository_RetrieveAllRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 uint64 - if args[2] != nil { - arg2 = args[2].(uint64) - } - var arg3 uint64 - if args[3] != nil { - arg3 = args[3].(uint64) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Repository_RetrieveAllRoles_Call) Return(rolePage roles.RolePage, err error) *Repository_RetrieveAllRoles_Call { - _c.Call.Return(rolePage, err) - return _c -} - -func (_c *Repository_RetrieveAllRoles_Call) RunAndReturn(run func(ctx context.Context, entityID string, limit uint64, offset uint64) (roles.RolePage, error)) *Repository_RetrieveAllRoles_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveDomainByID provides a mock function for the type Repository -func (_mock *Repository) RetrieveDomainByID(ctx context.Context, id string) (domains.Domain, error) { - ret := _mock.Called(ctx, id) - - if len(ret) == 0 { - panic("no return value specified for RetrieveDomainByID") - } - - var r0 domains.Domain - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) (domains.Domain, error)); ok { - return returnFunc(ctx, id) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string) domains.Domain); ok { - r0 = returnFunc(ctx, id) - } else { - r0 = ret.Get(0).(domains.Domain) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = returnFunc(ctx, id) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveDomainByID_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveDomainByID' -type Repository_RetrieveDomainByID_Call struct { - *mock.Call -} - -// RetrieveDomainByID is a helper method to define mock.On call -// - ctx context.Context -// - id string -func (_e *Repository_Expecter) RetrieveDomainByID(ctx interface{}, id interface{}) *Repository_RetrieveDomainByID_Call { - return &Repository_RetrieveDomainByID_Call{Call: _e.mock.On("RetrieveDomainByID", ctx, id)} -} - -func (_c *Repository_RetrieveDomainByID_Call) Run(run func(ctx context.Context, id string)) *Repository_RetrieveDomainByID_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RetrieveDomainByID_Call) Return(domain domains.Domain, err error) *Repository_RetrieveDomainByID_Call { - _c.Call.Return(domain, err) - return _c -} - -func (_c *Repository_RetrieveDomainByID_Call) RunAndReturn(run func(ctx context.Context, id string) (domains.Domain, error)) *Repository_RetrieveDomainByID_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveDomainByIDWithRoles provides a mock function for the type Repository -func (_mock *Repository) RetrieveDomainByIDWithRoles(ctx context.Context, id string, memberID string) (domains.Domain, error) { - ret := _mock.Called(ctx, id, memberID) - - if len(ret) == 0 { - panic("no return value specified for RetrieveDomainByIDWithRoles") - } - - var r0 domains.Domain - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) (domains.Domain, error)); ok { - return returnFunc(ctx, id, memberID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) domains.Domain); ok { - r0 = returnFunc(ctx, id, memberID) - } else { - r0 = ret.Get(0).(domains.Domain) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, string) error); ok { - r1 = returnFunc(ctx, id, memberID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveDomainByIDWithRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveDomainByIDWithRoles' -type Repository_RetrieveDomainByIDWithRoles_Call struct { - *mock.Call -} - -// RetrieveDomainByIDWithRoles is a helper method to define mock.On call -// - ctx context.Context -// - id string -// - memberID string -func (_e *Repository_Expecter) RetrieveDomainByIDWithRoles(ctx interface{}, id interface{}, memberID interface{}) *Repository_RetrieveDomainByIDWithRoles_Call { - return &Repository_RetrieveDomainByIDWithRoles_Call{Call: _e.mock.On("RetrieveDomainByIDWithRoles", ctx, id, memberID)} -} - -func (_c *Repository_RetrieveDomainByIDWithRoles_Call) Run(run func(ctx context.Context, id string, memberID string)) *Repository_RetrieveDomainByIDWithRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RetrieveDomainByIDWithRoles_Call) Return(domain domains.Domain, err error) *Repository_RetrieveDomainByIDWithRoles_Call { - _c.Call.Return(domain, err) - return _c -} - -func (_c *Repository_RetrieveDomainByIDWithRoles_Call) RunAndReturn(run func(ctx context.Context, id string, memberID string) (domains.Domain, error)) *Repository_RetrieveDomainByIDWithRoles_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveDomainByRoute provides a mock function for the type Repository -func (_mock *Repository) RetrieveDomainByRoute(ctx context.Context, route string) (domains.Domain, error) { - ret := _mock.Called(ctx, route) - - if len(ret) == 0 { - panic("no return value specified for RetrieveDomainByRoute") - } - - var r0 domains.Domain - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) (domains.Domain, error)); ok { - return returnFunc(ctx, route) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string) domains.Domain); ok { - r0 = returnFunc(ctx, route) - } else { - r0 = ret.Get(0).(domains.Domain) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = returnFunc(ctx, route) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveDomainByRoute_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveDomainByRoute' -type Repository_RetrieveDomainByRoute_Call struct { - *mock.Call -} - -// RetrieveDomainByRoute is a helper method to define mock.On call -// - ctx context.Context -// - route string -func (_e *Repository_Expecter) RetrieveDomainByRoute(ctx interface{}, route interface{}) *Repository_RetrieveDomainByRoute_Call { - return &Repository_RetrieveDomainByRoute_Call{Call: _e.mock.On("RetrieveDomainByRoute", ctx, route)} -} - -func (_c *Repository_RetrieveDomainByRoute_Call) Run(run func(ctx context.Context, route string)) *Repository_RetrieveDomainByRoute_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RetrieveDomainByRoute_Call) Return(domain domains.Domain, err error) *Repository_RetrieveDomainByRoute_Call { - _c.Call.Return(domain, err) - return _c -} - -func (_c *Repository_RetrieveDomainByRoute_Call) RunAndReturn(run func(ctx context.Context, route string) (domains.Domain, error)) *Repository_RetrieveDomainByRoute_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveEntitiesRolesActionsMembers provides a mock function for the type Repository -func (_mock *Repository) RetrieveEntitiesRolesActionsMembers(ctx context.Context, entityIDs []string) ([]roles.EntityActionRole, []roles.EntityMemberRole, error) { - ret := _mock.Called(ctx, entityIDs) - - if len(ret) == 0 { - panic("no return value specified for RetrieveEntitiesRolesActionsMembers") - } - - var r0 []roles.EntityActionRole - var r1 []roles.EntityMemberRole - var r2 error - if returnFunc, ok := ret.Get(0).(func(context.Context, []string) ([]roles.EntityActionRole, []roles.EntityMemberRole, error)); ok { - return returnFunc(ctx, entityIDs) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, []string) []roles.EntityActionRole); ok { - r0 = returnFunc(ctx, entityIDs) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]roles.EntityActionRole) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, []string) []roles.EntityMemberRole); ok { - r1 = returnFunc(ctx, entityIDs) - } else { - if ret.Get(1) != nil { - r1 = ret.Get(1).([]roles.EntityMemberRole) - } - } - if returnFunc, ok := ret.Get(2).(func(context.Context, []string) error); ok { - r2 = returnFunc(ctx, entityIDs) - } else { - r2 = ret.Error(2) - } - return r0, r1, r2 -} - -// Repository_RetrieveEntitiesRolesActionsMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveEntitiesRolesActionsMembers' -type Repository_RetrieveEntitiesRolesActionsMembers_Call struct { - *mock.Call -} - -// RetrieveEntitiesRolesActionsMembers is a helper method to define mock.On call -// - ctx context.Context -// - entityIDs []string -func (_e *Repository_Expecter) RetrieveEntitiesRolesActionsMembers(ctx interface{}, entityIDs interface{}) *Repository_RetrieveEntitiesRolesActionsMembers_Call { - return &Repository_RetrieveEntitiesRolesActionsMembers_Call{Call: _e.mock.On("RetrieveEntitiesRolesActionsMembers", ctx, entityIDs)} -} - -func (_c *Repository_RetrieveEntitiesRolesActionsMembers_Call) Run(run func(ctx context.Context, entityIDs []string)) *Repository_RetrieveEntitiesRolesActionsMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 []string - if args[1] != nil { - arg1 = args[1].([]string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RetrieveEntitiesRolesActionsMembers_Call) Return(entityActionRoles []roles.EntityActionRole, entityMemberRoles []roles.EntityMemberRole, err error) *Repository_RetrieveEntitiesRolesActionsMembers_Call { - _c.Call.Return(entityActionRoles, entityMemberRoles, err) - return _c -} - -func (_c *Repository_RetrieveEntitiesRolesActionsMembers_Call) RunAndReturn(run func(ctx context.Context, entityIDs []string) ([]roles.EntityActionRole, []roles.EntityMemberRole, error)) *Repository_RetrieveEntitiesRolesActionsMembers_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveEntityRole provides a mock function for the type Repository -func (_mock *Repository) RetrieveEntityRole(ctx context.Context, entityID string, roleID string) (roles.Role, error) { - ret := _mock.Called(ctx, entityID, roleID) - - if len(ret) == 0 { - panic("no return value specified for RetrieveEntityRole") - } - - var r0 roles.Role - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) (roles.Role, error)); ok { - return returnFunc(ctx, entityID, roleID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) roles.Role); ok { - r0 = returnFunc(ctx, entityID, roleID) - } else { - r0 = ret.Get(0).(roles.Role) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, string) error); ok { - r1 = returnFunc(ctx, entityID, roleID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveEntityRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveEntityRole' -type Repository_RetrieveEntityRole_Call struct { - *mock.Call -} - -// RetrieveEntityRole is a helper method to define mock.On call -// - ctx context.Context -// - entityID string -// - roleID string -func (_e *Repository_Expecter) RetrieveEntityRole(ctx interface{}, entityID interface{}, roleID interface{}) *Repository_RetrieveEntityRole_Call { - return &Repository_RetrieveEntityRole_Call{Call: _e.mock.On("RetrieveEntityRole", ctx, entityID, roleID)} -} - -func (_c *Repository_RetrieveEntityRole_Call) Run(run func(ctx context.Context, entityID string, roleID string)) *Repository_RetrieveEntityRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RetrieveEntityRole_Call) Return(role roles.Role, err error) *Repository_RetrieveEntityRole_Call { - _c.Call.Return(role, err) - return _c -} - -func (_c *Repository_RetrieveEntityRole_Call) RunAndReturn(run func(ctx context.Context, entityID string, roleID string) (roles.Role, error)) *Repository_RetrieveEntityRole_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveInvitation provides a mock function for the type Repository -func (_mock *Repository) RetrieveInvitation(ctx context.Context, userID string, domainID string) (domains.Invitation, error) { - ret := _mock.Called(ctx, userID, domainID) - - if len(ret) == 0 { - panic("no return value specified for RetrieveInvitation") - } - - var r0 domains.Invitation - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) (domains.Invitation, error)); ok { - return returnFunc(ctx, userID, domainID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) domains.Invitation); ok { - r0 = returnFunc(ctx, userID, domainID) - } else { - r0 = ret.Get(0).(domains.Invitation) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, string) error); ok { - r1 = returnFunc(ctx, userID, domainID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveInvitation_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveInvitation' -type Repository_RetrieveInvitation_Call struct { - *mock.Call -} - -// RetrieveInvitation is a helper method to define mock.On call -// - ctx context.Context -// - userID string -// - domainID string -func (_e *Repository_Expecter) RetrieveInvitation(ctx interface{}, userID interface{}, domainID interface{}) *Repository_RetrieveInvitation_Call { - return &Repository_RetrieveInvitation_Call{Call: _e.mock.On("RetrieveInvitation", ctx, userID, domainID)} -} - -func (_c *Repository_RetrieveInvitation_Call) Run(run func(ctx context.Context, userID string, domainID string)) *Repository_RetrieveInvitation_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RetrieveInvitation_Call) Return(invitation domains.Invitation, err error) *Repository_RetrieveInvitation_Call { - _c.Call.Return(invitation, err) - return _c -} - -func (_c *Repository_RetrieveInvitation_Call) RunAndReturn(run func(ctx context.Context, userID string, domainID string) (domains.Invitation, error)) *Repository_RetrieveInvitation_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveRole provides a mock function for the type Repository -func (_mock *Repository) RetrieveRole(ctx context.Context, roleID string) (roles.Role, error) { - ret := _mock.Called(ctx, roleID) - - if len(ret) == 0 { - panic("no return value specified for RetrieveRole") - } - - var r0 roles.Role - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) (roles.Role, error)); ok { - return returnFunc(ctx, roleID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string) roles.Role); ok { - r0 = returnFunc(ctx, roleID) - } else { - r0 = ret.Get(0).(roles.Role) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = returnFunc(ctx, roleID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveRole' -type Repository_RetrieveRole_Call struct { - *mock.Call -} - -// RetrieveRole is a helper method to define mock.On call -// - ctx context.Context -// - roleID string -func (_e *Repository_Expecter) RetrieveRole(ctx interface{}, roleID interface{}) *Repository_RetrieveRole_Call { - return &Repository_RetrieveRole_Call{Call: _e.mock.On("RetrieveRole", ctx, roleID)} -} - -func (_c *Repository_RetrieveRole_Call) Run(run func(ctx context.Context, roleID string)) *Repository_RetrieveRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RetrieveRole_Call) Return(role roles.Role, err error) *Repository_RetrieveRole_Call { - _c.Call.Return(role, err) - return _c -} - -func (_c *Repository_RetrieveRole_Call) RunAndReturn(run func(ctx context.Context, roleID string) (roles.Role, error)) *Repository_RetrieveRole_Call { - _c.Call.Return(run) - return _c -} - -// RoleAddActions provides a mock function for the type Repository -func (_mock *Repository) RoleAddActions(ctx context.Context, role roles.Role, actions []string) ([]string, error) { - ret := _mock.Called(ctx, role, actions) - - if len(ret) == 0 { - panic("no return value specified for RoleAddActions") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role, []string) ([]string, error)); ok { - return returnFunc(ctx, role, actions) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role, []string) []string); ok { - r0 = returnFunc(ctx, role, actions) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, roles.Role, []string) error); ok { - r1 = returnFunc(ctx, role, actions) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RoleAddActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleAddActions' -type Repository_RoleAddActions_Call struct { - *mock.Call -} - -// RoleAddActions is a helper method to define mock.On call -// - ctx context.Context -// - role roles.Role -// - actions []string -func (_e *Repository_Expecter) RoleAddActions(ctx interface{}, role interface{}, actions interface{}) *Repository_RoleAddActions_Call { - return &Repository_RoleAddActions_Call{Call: _e.mock.On("RoleAddActions", ctx, role, actions)} -} - -func (_c *Repository_RoleAddActions_Call) Run(run func(ctx context.Context, role roles.Role, actions []string)) *Repository_RoleAddActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RoleAddActions_Call) Return(ops []string, err error) *Repository_RoleAddActions_Call { - _c.Call.Return(ops, err) - return _c -} - -func (_c *Repository_RoleAddActions_Call) RunAndReturn(run func(ctx context.Context, role roles.Role, actions []string) ([]string, error)) *Repository_RoleAddActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleAddMembers provides a mock function for the type Repository -func (_mock *Repository) RoleAddMembers(ctx context.Context, role roles.Role, members []string) ([]string, error) { - ret := _mock.Called(ctx, role, members) - - if len(ret) == 0 { - panic("no return value specified for RoleAddMembers") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role, []string) ([]string, error)); ok { - return returnFunc(ctx, role, members) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role, []string) []string); ok { - r0 = returnFunc(ctx, role, members) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, roles.Role, []string) error); ok { - r1 = returnFunc(ctx, role, members) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RoleAddMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleAddMembers' -type Repository_RoleAddMembers_Call struct { - *mock.Call -} - -// RoleAddMembers is a helper method to define mock.On call -// - ctx context.Context -// - role roles.Role -// - members []string -func (_e *Repository_Expecter) RoleAddMembers(ctx interface{}, role interface{}, members interface{}) *Repository_RoleAddMembers_Call { - return &Repository_RoleAddMembers_Call{Call: _e.mock.On("RoleAddMembers", ctx, role, members)} -} - -func (_c *Repository_RoleAddMembers_Call) Run(run func(ctx context.Context, role roles.Role, members []string)) *Repository_RoleAddMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RoleAddMembers_Call) Return(strings []string, err error) *Repository_RoleAddMembers_Call { - _c.Call.Return(strings, err) - return _c -} - -func (_c *Repository_RoleAddMembers_Call) RunAndReturn(run func(ctx context.Context, role roles.Role, members []string) ([]string, error)) *Repository_RoleAddMembers_Call { - _c.Call.Return(run) - return _c -} - -// RoleCheckActionsExists provides a mock function for the type Repository -func (_mock *Repository) RoleCheckActionsExists(ctx context.Context, roleID string, actions []string) (bool, error) { - ret := _mock.Called(ctx, roleID, actions) - - if len(ret) == 0 { - panic("no return value specified for RoleCheckActionsExists") - } - - var r0 bool - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, []string) (bool, error)); ok { - return returnFunc(ctx, roleID, actions) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, []string) bool); ok { - r0 = returnFunc(ctx, roleID, actions) - } else { - r0 = ret.Get(0).(bool) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, []string) error); ok { - r1 = returnFunc(ctx, roleID, actions) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RoleCheckActionsExists_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleCheckActionsExists' -type Repository_RoleCheckActionsExists_Call struct { - *mock.Call -} - -// RoleCheckActionsExists is a helper method to define mock.On call -// - ctx context.Context -// - roleID string -// - actions []string -func (_e *Repository_Expecter) RoleCheckActionsExists(ctx interface{}, roleID interface{}, actions interface{}) *Repository_RoleCheckActionsExists_Call { - return &Repository_RoleCheckActionsExists_Call{Call: _e.mock.On("RoleCheckActionsExists", ctx, roleID, actions)} -} - -func (_c *Repository_RoleCheckActionsExists_Call) Run(run func(ctx context.Context, roleID string, actions []string)) *Repository_RoleCheckActionsExists_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RoleCheckActionsExists_Call) Return(b bool, err error) *Repository_RoleCheckActionsExists_Call { - _c.Call.Return(b, err) - return _c -} - -func (_c *Repository_RoleCheckActionsExists_Call) RunAndReturn(run func(ctx context.Context, roleID string, actions []string) (bool, error)) *Repository_RoleCheckActionsExists_Call { - _c.Call.Return(run) - return _c -} - -// RoleCheckMembersExists provides a mock function for the type Repository -func (_mock *Repository) RoleCheckMembersExists(ctx context.Context, roleID string, members []string) (bool, error) { - ret := _mock.Called(ctx, roleID, members) - - if len(ret) == 0 { - panic("no return value specified for RoleCheckMembersExists") - } - - var r0 bool - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, []string) (bool, error)); ok { - return returnFunc(ctx, roleID, members) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, []string) bool); ok { - r0 = returnFunc(ctx, roleID, members) - } else { - r0 = ret.Get(0).(bool) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, []string) error); ok { - r1 = returnFunc(ctx, roleID, members) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RoleCheckMembersExists_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleCheckMembersExists' -type Repository_RoleCheckMembersExists_Call struct { - *mock.Call -} - -// RoleCheckMembersExists is a helper method to define mock.On call -// - ctx context.Context -// - roleID string -// - members []string -func (_e *Repository_Expecter) RoleCheckMembersExists(ctx interface{}, roleID interface{}, members interface{}) *Repository_RoleCheckMembersExists_Call { - return &Repository_RoleCheckMembersExists_Call{Call: _e.mock.On("RoleCheckMembersExists", ctx, roleID, members)} -} - -func (_c *Repository_RoleCheckMembersExists_Call) Run(run func(ctx context.Context, roleID string, members []string)) *Repository_RoleCheckMembersExists_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RoleCheckMembersExists_Call) Return(b bool, err error) *Repository_RoleCheckMembersExists_Call { - _c.Call.Return(b, err) - return _c -} - -func (_c *Repository_RoleCheckMembersExists_Call) RunAndReturn(run func(ctx context.Context, roleID string, members []string) (bool, error)) *Repository_RoleCheckMembersExists_Call { - _c.Call.Return(run) - return _c -} - -// RoleListActions provides a mock function for the type Repository -func (_mock *Repository) RoleListActions(ctx context.Context, roleID string) ([]string, error) { - ret := _mock.Called(ctx, roleID) - - if len(ret) == 0 { - panic("no return value specified for RoleListActions") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) ([]string, error)); ok { - return returnFunc(ctx, roleID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string) []string); ok { - r0 = returnFunc(ctx, roleID) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = returnFunc(ctx, roleID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RoleListActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleListActions' -type Repository_RoleListActions_Call struct { - *mock.Call -} - -// RoleListActions is a helper method to define mock.On call -// - ctx context.Context -// - roleID string -func (_e *Repository_Expecter) RoleListActions(ctx interface{}, roleID interface{}) *Repository_RoleListActions_Call { - return &Repository_RoleListActions_Call{Call: _e.mock.On("RoleListActions", ctx, roleID)} -} - -func (_c *Repository_RoleListActions_Call) Run(run func(ctx context.Context, roleID string)) *Repository_RoleListActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RoleListActions_Call) Return(strings []string, err error) *Repository_RoleListActions_Call { - _c.Call.Return(strings, err) - return _c -} - -func (_c *Repository_RoleListActions_Call) RunAndReturn(run func(ctx context.Context, roleID string) ([]string, error)) *Repository_RoleListActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleListMembers provides a mock function for the type Repository -func (_mock *Repository) RoleListMembers(ctx context.Context, roleID string, limit uint64, offset uint64) (roles.MembersPage, error) { - ret := _mock.Called(ctx, roleID, limit, offset) - - if len(ret) == 0 { - panic("no return value specified for RoleListMembers") - } - - var r0 roles.MembersPage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, uint64, uint64) (roles.MembersPage, error)); ok { - return returnFunc(ctx, roleID, limit, offset) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, uint64, uint64) roles.MembersPage); ok { - r0 = returnFunc(ctx, roleID, limit, offset) - } else { - r0 = ret.Get(0).(roles.MembersPage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, uint64, uint64) error); ok { - r1 = returnFunc(ctx, roleID, limit, offset) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RoleListMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleListMembers' -type Repository_RoleListMembers_Call struct { - *mock.Call -} - -// RoleListMembers is a helper method to define mock.On call -// - ctx context.Context -// - roleID string -// - limit uint64 -// - offset uint64 -func (_e *Repository_Expecter) RoleListMembers(ctx interface{}, roleID interface{}, limit interface{}, offset interface{}) *Repository_RoleListMembers_Call { - return &Repository_RoleListMembers_Call{Call: _e.mock.On("RoleListMembers", ctx, roleID, limit, offset)} -} - -func (_c *Repository_RoleListMembers_Call) Run(run func(ctx context.Context, roleID string, limit uint64, offset uint64)) *Repository_RoleListMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 uint64 - if args[2] != nil { - arg2 = args[2].(uint64) - } - var arg3 uint64 - if args[3] != nil { - arg3 = args[3].(uint64) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Repository_RoleListMembers_Call) Return(membersPage roles.MembersPage, err error) *Repository_RoleListMembers_Call { - _c.Call.Return(membersPage, err) - return _c -} - -func (_c *Repository_RoleListMembers_Call) RunAndReturn(run func(ctx context.Context, roleID string, limit uint64, offset uint64) (roles.MembersPage, error)) *Repository_RoleListMembers_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveActions provides a mock function for the type Repository -func (_mock *Repository) RoleRemoveActions(ctx context.Context, role roles.Role, actions []string) error { - ret := _mock.Called(ctx, role, actions) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveActions") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role, []string) error); ok { - r0 = returnFunc(ctx, role, actions) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RoleRemoveActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveActions' -type Repository_RoleRemoveActions_Call struct { - *mock.Call -} - -// RoleRemoveActions is a helper method to define mock.On call -// - ctx context.Context -// - role roles.Role -// - actions []string -func (_e *Repository_Expecter) RoleRemoveActions(ctx interface{}, role interface{}, actions interface{}) *Repository_RoleRemoveActions_Call { - return &Repository_RoleRemoveActions_Call{Call: _e.mock.On("RoleRemoveActions", ctx, role, actions)} -} - -func (_c *Repository_RoleRemoveActions_Call) Run(run func(ctx context.Context, role roles.Role, actions []string)) *Repository_RoleRemoveActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RoleRemoveActions_Call) Return(err error) *Repository_RoleRemoveActions_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RoleRemoveActions_Call) RunAndReturn(run func(ctx context.Context, role roles.Role, actions []string) error) *Repository_RoleRemoveActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveAllActions provides a mock function for the type Repository -func (_mock *Repository) RoleRemoveAllActions(ctx context.Context, role roles.Role) error { - ret := _mock.Called(ctx, role) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveAllActions") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role) error); ok { - r0 = returnFunc(ctx, role) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RoleRemoveAllActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveAllActions' -type Repository_RoleRemoveAllActions_Call struct { - *mock.Call -} - -// RoleRemoveAllActions is a helper method to define mock.On call -// - ctx context.Context -// - role roles.Role -func (_e *Repository_Expecter) RoleRemoveAllActions(ctx interface{}, role interface{}) *Repository_RoleRemoveAllActions_Call { - return &Repository_RoleRemoveAllActions_Call{Call: _e.mock.On("RoleRemoveAllActions", ctx, role)} -} - -func (_c *Repository_RoleRemoveAllActions_Call) Run(run func(ctx context.Context, role roles.Role)) *Repository_RoleRemoveAllActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RoleRemoveAllActions_Call) Return(err error) *Repository_RoleRemoveAllActions_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RoleRemoveAllActions_Call) RunAndReturn(run func(ctx context.Context, role roles.Role) error) *Repository_RoleRemoveAllActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveAllMembers provides a mock function for the type Repository -func (_mock *Repository) RoleRemoveAllMembers(ctx context.Context, role roles.Role) error { - ret := _mock.Called(ctx, role) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveAllMembers") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role) error); ok { - r0 = returnFunc(ctx, role) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RoleRemoveAllMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveAllMembers' -type Repository_RoleRemoveAllMembers_Call struct { - *mock.Call -} - -// RoleRemoveAllMembers is a helper method to define mock.On call -// - ctx context.Context -// - role roles.Role -func (_e *Repository_Expecter) RoleRemoveAllMembers(ctx interface{}, role interface{}) *Repository_RoleRemoveAllMembers_Call { - return &Repository_RoleRemoveAllMembers_Call{Call: _e.mock.On("RoleRemoveAllMembers", ctx, role)} -} - -func (_c *Repository_RoleRemoveAllMembers_Call) Run(run func(ctx context.Context, role roles.Role)) *Repository_RoleRemoveAllMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RoleRemoveAllMembers_Call) Return(err error) *Repository_RoleRemoveAllMembers_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RoleRemoveAllMembers_Call) RunAndReturn(run func(ctx context.Context, role roles.Role) error) *Repository_RoleRemoveAllMembers_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveMembers provides a mock function for the type Repository -func (_mock *Repository) RoleRemoveMembers(ctx context.Context, role roles.Role, members []string) error { - ret := _mock.Called(ctx, role, members) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveMembers") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role, []string) error); ok { - r0 = returnFunc(ctx, role, members) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RoleRemoveMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveMembers' -type Repository_RoleRemoveMembers_Call struct { - *mock.Call -} - -// RoleRemoveMembers is a helper method to define mock.On call -// - ctx context.Context -// - role roles.Role -// - members []string -func (_e *Repository_Expecter) RoleRemoveMembers(ctx interface{}, role interface{}, members interface{}) *Repository_RoleRemoveMembers_Call { - return &Repository_RoleRemoveMembers_Call{Call: _e.mock.On("RoleRemoveMembers", ctx, role, members)} -} - -func (_c *Repository_RoleRemoveMembers_Call) Run(run func(ctx context.Context, role roles.Role, members []string)) *Repository_RoleRemoveMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RoleRemoveMembers_Call) Return(err error) *Repository_RoleRemoveMembers_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RoleRemoveMembers_Call) RunAndReturn(run func(ctx context.Context, role roles.Role, members []string) error) *Repository_RoleRemoveMembers_Call { - _c.Call.Return(run) - return _c -} - -// SaveDomain provides a mock function for the type Repository -func (_mock *Repository) SaveDomain(ctx context.Context, d domains.Domain) (domains.Domain, error) { - ret := _mock.Called(ctx, d) - - if len(ret) == 0 { - panic("no return value specified for SaveDomain") - } - - var r0 domains.Domain - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, domains.Domain) (domains.Domain, error)); ok { - return returnFunc(ctx, d) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, domains.Domain) domains.Domain); ok { - r0 = returnFunc(ctx, d) - } else { - r0 = ret.Get(0).(domains.Domain) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, domains.Domain) error); ok { - r1 = returnFunc(ctx, d) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_SaveDomain_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'SaveDomain' -type Repository_SaveDomain_Call struct { - *mock.Call -} - -// SaveDomain is a helper method to define mock.On call -// - ctx context.Context -// - d domains.Domain -func (_e *Repository_Expecter) SaveDomain(ctx interface{}, d interface{}) *Repository_SaveDomain_Call { - return &Repository_SaveDomain_Call{Call: _e.mock.On("SaveDomain", ctx, d)} -} - -func (_c *Repository_SaveDomain_Call) Run(run func(ctx context.Context, d domains.Domain)) *Repository_SaveDomain_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 domains.Domain - if args[1] != nil { - arg1 = args[1].(domains.Domain) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_SaveDomain_Call) Return(domain domains.Domain, err error) *Repository_SaveDomain_Call { - _c.Call.Return(domain, err) - return _c -} - -func (_c *Repository_SaveDomain_Call) RunAndReturn(run func(ctx context.Context, d domains.Domain) (domains.Domain, error)) *Repository_SaveDomain_Call { - _c.Call.Return(run) - return _c -} - -// SaveInvitation provides a mock function for the type Repository -func (_mock *Repository) SaveInvitation(ctx context.Context, invitation domains.Invitation) error { - ret := _mock.Called(ctx, invitation) - - if len(ret) == 0 { - panic("no return value specified for SaveInvitation") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, domains.Invitation) error); ok { - r0 = returnFunc(ctx, invitation) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_SaveInvitation_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'SaveInvitation' -type Repository_SaveInvitation_Call struct { - *mock.Call -} - -// SaveInvitation is a helper method to define mock.On call -// - ctx context.Context -// - invitation domains.Invitation -func (_e *Repository_Expecter) SaveInvitation(ctx interface{}, invitation interface{}) *Repository_SaveInvitation_Call { - return &Repository_SaveInvitation_Call{Call: _e.mock.On("SaveInvitation", ctx, invitation)} -} - -func (_c *Repository_SaveInvitation_Call) Run(run func(ctx context.Context, invitation domains.Invitation)) *Repository_SaveInvitation_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 domains.Invitation - if args[1] != nil { - arg1 = args[1].(domains.Invitation) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_SaveInvitation_Call) Return(err error) *Repository_SaveInvitation_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_SaveInvitation_Call) RunAndReturn(run func(ctx context.Context, invitation domains.Invitation) error) *Repository_SaveInvitation_Call { - _c.Call.Return(run) - return _c -} - -// UpdateConfirmation provides a mock function for the type Repository -func (_mock *Repository) UpdateConfirmation(ctx context.Context, invitation domains.Invitation) error { - ret := _mock.Called(ctx, invitation) - - if len(ret) == 0 { - panic("no return value specified for UpdateConfirmation") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, domains.Invitation) error); ok { - r0 = returnFunc(ctx, invitation) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_UpdateConfirmation_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateConfirmation' -type Repository_UpdateConfirmation_Call struct { - *mock.Call -} - -// UpdateConfirmation is a helper method to define mock.On call -// - ctx context.Context -// - invitation domains.Invitation -func (_e *Repository_Expecter) UpdateConfirmation(ctx interface{}, invitation interface{}) *Repository_UpdateConfirmation_Call { - return &Repository_UpdateConfirmation_Call{Call: _e.mock.On("UpdateConfirmation", ctx, invitation)} -} - -func (_c *Repository_UpdateConfirmation_Call) Run(run func(ctx context.Context, invitation domains.Invitation)) *Repository_UpdateConfirmation_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 domains.Invitation - if args[1] != nil { - arg1 = args[1].(domains.Invitation) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_UpdateConfirmation_Call) Return(err error) *Repository_UpdateConfirmation_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_UpdateConfirmation_Call) RunAndReturn(run func(ctx context.Context, invitation domains.Invitation) error) *Repository_UpdateConfirmation_Call { - _c.Call.Return(run) - return _c -} - -// UpdateDomain provides a mock function for the type Repository -func (_mock *Repository) UpdateDomain(ctx context.Context, id string, d domains.DomainReq) (domains.Domain, error) { - ret := _mock.Called(ctx, id, d) - - if len(ret) == 0 { - panic("no return value specified for UpdateDomain") - } - - var r0 domains.Domain - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, domains.DomainReq) (domains.Domain, error)); ok { - return returnFunc(ctx, id, d) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, domains.DomainReq) domains.Domain); ok { - r0 = returnFunc(ctx, id, d) - } else { - r0 = ret.Get(0).(domains.Domain) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, domains.DomainReq) error); ok { - r1 = returnFunc(ctx, id, d) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_UpdateDomain_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateDomain' -type Repository_UpdateDomain_Call struct { - *mock.Call -} - -// UpdateDomain is a helper method to define mock.On call -// - ctx context.Context -// - id string -// - d domains.DomainReq -func (_e *Repository_Expecter) UpdateDomain(ctx interface{}, id interface{}, d interface{}) *Repository_UpdateDomain_Call { - return &Repository_UpdateDomain_Call{Call: _e.mock.On("UpdateDomain", ctx, id, d)} -} - -func (_c *Repository_UpdateDomain_Call) Run(run func(ctx context.Context, id string, d domains.DomainReq)) *Repository_UpdateDomain_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 domains.DomainReq - if args[2] != nil { - arg2 = args[2].(domains.DomainReq) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_UpdateDomain_Call) Return(domain domains.Domain, err error) *Repository_UpdateDomain_Call { - _c.Call.Return(domain, err) - return _c -} - -func (_c *Repository_UpdateDomain_Call) RunAndReturn(run func(ctx context.Context, id string, d domains.DomainReq) (domains.Domain, error)) *Repository_UpdateDomain_Call { - _c.Call.Return(run) - return _c -} - -// UpdateRejection provides a mock function for the type Repository -func (_mock *Repository) UpdateRejection(ctx context.Context, invitation domains.Invitation) error { - ret := _mock.Called(ctx, invitation) - - if len(ret) == 0 { - panic("no return value specified for UpdateRejection") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, domains.Invitation) error); ok { - r0 = returnFunc(ctx, invitation) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_UpdateRejection_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateRejection' -type Repository_UpdateRejection_Call struct { - *mock.Call -} - -// UpdateRejection is a helper method to define mock.On call -// - ctx context.Context -// - invitation domains.Invitation -func (_e *Repository_Expecter) UpdateRejection(ctx interface{}, invitation interface{}) *Repository_UpdateRejection_Call { - return &Repository_UpdateRejection_Call{Call: _e.mock.On("UpdateRejection", ctx, invitation)} -} - -func (_c *Repository_UpdateRejection_Call) Run(run func(ctx context.Context, invitation domains.Invitation)) *Repository_UpdateRejection_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 domains.Invitation - if args[1] != nil { - arg1 = args[1].(domains.Invitation) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_UpdateRejection_Call) Return(err error) *Repository_UpdateRejection_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_UpdateRejection_Call) RunAndReturn(run func(ctx context.Context, invitation domains.Invitation) error) *Repository_UpdateRejection_Call { - _c.Call.Return(run) - return _c -} - -// UpdateRole provides a mock function for the type Repository -func (_mock *Repository) UpdateRole(ctx context.Context, ro roles.Role) (roles.Role, error) { - ret := _mock.Called(ctx, ro) - - if len(ret) == 0 { - panic("no return value specified for UpdateRole") - } - - var r0 roles.Role - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role) (roles.Role, error)); ok { - return returnFunc(ctx, ro) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role) roles.Role); ok { - r0 = returnFunc(ctx, ro) - } else { - r0 = ret.Get(0).(roles.Role) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, roles.Role) error); ok { - r1 = returnFunc(ctx, ro) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_UpdateRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateRole' -type Repository_UpdateRole_Call struct { - *mock.Call -} - -// UpdateRole is a helper method to define mock.On call -// - ctx context.Context -// - ro roles.Role -func (_e *Repository_Expecter) UpdateRole(ctx interface{}, ro interface{}) *Repository_UpdateRole_Call { - return &Repository_UpdateRole_Call{Call: _e.mock.On("UpdateRole", ctx, ro)} -} - -func (_c *Repository_UpdateRole_Call) Run(run func(ctx context.Context, ro roles.Role)) *Repository_UpdateRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_UpdateRole_Call) Return(role roles.Role, err error) *Repository_UpdateRole_Call { - _c.Call.Return(role, err) - return _c -} - -func (_c *Repository_UpdateRole_Call) RunAndReturn(run func(ctx context.Context, ro roles.Role) (roles.Role, error)) *Repository_UpdateRole_Call { - _c.Call.Return(run) - return _c -} diff --git a/domains/mocks/service.go b/domains/mocks/service.go deleted file mode 100644 index 998e5c4b7..000000000 --- a/domains/mocks/service.go +++ /dev/null @@ -1,2479 +0,0 @@ -// Copyright (c) Abstract Machines - -// SPDX-License-Identifier: Apache-2.0 - -// Code generated by mockery; DO NOT EDIT. -// github.com/vektra/mockery -// template: testify - -package mocks - -import ( - "context" - - "github.com/absmach/magistrala/domains" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/roles" - mock "github.com/stretchr/testify/mock" -) - -// NewService creates a new instance of Service. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. -// The first argument is typically a *testing.T value. -func NewService(t interface { - mock.TestingT - Cleanup(func()) -}) *Service { - mock := &Service{} - mock.Mock.Test(t) - - t.Cleanup(func() { mock.AssertExpectations(t) }) - - return mock -} - -// Service is an autogenerated mock type for the Service type -type Service struct { - mock.Mock -} - -type Service_Expecter struct { - mock *mock.Mock -} - -func (_m *Service) EXPECT() *Service_Expecter { - return &Service_Expecter{mock: &_m.Mock} -} - -// AcceptInvitation provides a mock function for the type Service -func (_mock *Service) AcceptInvitation(ctx context.Context, session authn.Session, domainID string) (domains.Invitation, error) { - ret := _mock.Called(ctx, session, domainID) - - if len(ret) == 0 { - panic("no return value specified for AcceptInvitation") - } - - var r0 domains.Invitation - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) (domains.Invitation, error)); ok { - return returnFunc(ctx, session, domainID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) domains.Invitation); ok { - r0 = returnFunc(ctx, session, domainID) - } else { - r0 = ret.Get(0).(domains.Invitation) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string) error); ok { - r1 = returnFunc(ctx, session, domainID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_AcceptInvitation_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'AcceptInvitation' -type Service_AcceptInvitation_Call struct { - *mock.Call -} - -// AcceptInvitation is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - domainID string -func (_e *Service_Expecter) AcceptInvitation(ctx interface{}, session interface{}, domainID interface{}) *Service_AcceptInvitation_Call { - return &Service_AcceptInvitation_Call{Call: _e.mock.On("AcceptInvitation", ctx, session, domainID)} -} - -func (_c *Service_AcceptInvitation_Call) Run(run func(ctx context.Context, session authn.Session, domainID string)) *Service_AcceptInvitation_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_AcceptInvitation_Call) Return(invitation domains.Invitation, err error) *Service_AcceptInvitation_Call { - _c.Call.Return(invitation, err) - return _c -} - -func (_c *Service_AcceptInvitation_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, domainID string) (domains.Invitation, error)) *Service_AcceptInvitation_Call { - _c.Call.Return(run) - return _c -} - -// AddRole provides a mock function for the type Service -func (_mock *Service) AddRole(ctx context.Context, session authn.Session, entityID string, roleName string, optionalActions []string, optionalMembers []string) (roles.RoleProvision, error) { - ret := _mock.Called(ctx, session, entityID, roleName, optionalActions, optionalMembers) - - if len(ret) == 0 { - panic("no return value specified for AddRole") - } - - var r0 roles.RoleProvision - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string, []string) (roles.RoleProvision, error)); ok { - return returnFunc(ctx, session, entityID, roleName, optionalActions, optionalMembers) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string, []string) roles.RoleProvision); ok { - r0 = returnFunc(ctx, session, entityID, roleName, optionalActions, optionalMembers) - } else { - r0 = ret.Get(0).(roles.RoleProvision) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, []string, []string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleName, optionalActions, optionalMembers) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_AddRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'AddRole' -type Service_AddRole_Call struct { - *mock.Call -} - -// AddRole is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleName string -// - optionalActions []string -// - optionalMembers []string -func (_e *Service_Expecter) AddRole(ctx interface{}, session interface{}, entityID interface{}, roleName interface{}, optionalActions interface{}, optionalMembers interface{}) *Service_AddRole_Call { - return &Service_AddRole_Call{Call: _e.mock.On("AddRole", ctx, session, entityID, roleName, optionalActions, optionalMembers)} -} - -func (_c *Service_AddRole_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleName string, optionalActions []string, optionalMembers []string)) *Service_AddRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - var arg5 []string - if args[5] != nil { - arg5 = args[5].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - arg5, - ) - }) - return _c -} - -func (_c *Service_AddRole_Call) Return(roleProvision roles.RoleProvision, err error) *Service_AddRole_Call { - _c.Call.Return(roleProvision, err) - return _c -} - -func (_c *Service_AddRole_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleName string, optionalActions []string, optionalMembers []string) (roles.RoleProvision, error)) *Service_AddRole_Call { - _c.Call.Return(run) - return _c -} - -// CreateDomain provides a mock function for the type Service -func (_mock *Service) CreateDomain(ctx context.Context, sesssion authn.Session, d domains.Domain) (domains.Domain, []roles.RoleProvision, error) { - ret := _mock.Called(ctx, sesssion, d) - - if len(ret) == 0 { - panic("no return value specified for CreateDomain") - } - - var r0 domains.Domain - var r1 []roles.RoleProvision - var r2 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, domains.Domain) (domains.Domain, []roles.RoleProvision, error)); ok { - return returnFunc(ctx, sesssion, d) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, domains.Domain) domains.Domain); ok { - r0 = returnFunc(ctx, sesssion, d) - } else { - r0 = ret.Get(0).(domains.Domain) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, domains.Domain) []roles.RoleProvision); ok { - r1 = returnFunc(ctx, sesssion, d) - } else { - if ret.Get(1) != nil { - r1 = ret.Get(1).([]roles.RoleProvision) - } - } - if returnFunc, ok := ret.Get(2).(func(context.Context, authn.Session, domains.Domain) error); ok { - r2 = returnFunc(ctx, sesssion, d) - } else { - r2 = ret.Error(2) - } - return r0, r1, r2 -} - -// Service_CreateDomain_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'CreateDomain' -type Service_CreateDomain_Call struct { - *mock.Call -} - -// CreateDomain is a helper method to define mock.On call -// - ctx context.Context -// - sesssion authn.Session -// - d domains.Domain -func (_e *Service_Expecter) CreateDomain(ctx interface{}, sesssion interface{}, d interface{}) *Service_CreateDomain_Call { - return &Service_CreateDomain_Call{Call: _e.mock.On("CreateDomain", ctx, sesssion, d)} -} - -func (_c *Service_CreateDomain_Call) Run(run func(ctx context.Context, sesssion authn.Session, d domains.Domain)) *Service_CreateDomain_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 domains.Domain - if args[2] != nil { - arg2 = args[2].(domains.Domain) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_CreateDomain_Call) Return(domain domains.Domain, roleProvisions []roles.RoleProvision, err error) *Service_CreateDomain_Call { - _c.Call.Return(domain, roleProvisions, err) - return _c -} - -func (_c *Service_CreateDomain_Call) RunAndReturn(run func(ctx context.Context, sesssion authn.Session, d domains.Domain) (domains.Domain, []roles.RoleProvision, error)) *Service_CreateDomain_Call { - _c.Call.Return(run) - return _c -} - -// DeleteInvitation provides a mock function for the type Service -func (_mock *Service) DeleteInvitation(ctx context.Context, session authn.Session, inviteeUserID string, domainID string) error { - ret := _mock.Called(ctx, session, inviteeUserID, domainID) - - if len(ret) == 0 { - panic("no return value specified for DeleteInvitation") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) error); ok { - r0 = returnFunc(ctx, session, inviteeUserID, domainID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_DeleteInvitation_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'DeleteInvitation' -type Service_DeleteInvitation_Call struct { - *mock.Call -} - -// DeleteInvitation is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - inviteeUserID string -// - domainID string -func (_e *Service_Expecter) DeleteInvitation(ctx interface{}, session interface{}, inviteeUserID interface{}, domainID interface{}) *Service_DeleteInvitation_Call { - return &Service_DeleteInvitation_Call{Call: _e.mock.On("DeleteInvitation", ctx, session, inviteeUserID, domainID)} -} - -func (_c *Service_DeleteInvitation_Call) Run(run func(ctx context.Context, session authn.Session, inviteeUserID string, domainID string)) *Service_DeleteInvitation_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_DeleteInvitation_Call) Return(err error) *Service_DeleteInvitation_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_DeleteInvitation_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, inviteeUserID string, domainID string) error) *Service_DeleteInvitation_Call { - _c.Call.Return(run) - return _c -} - -// DisableDomain provides a mock function for the type Service -func (_mock *Service) DisableDomain(ctx context.Context, sesssion authn.Session, id string) (domains.Domain, error) { - ret := _mock.Called(ctx, sesssion, id) - - if len(ret) == 0 { - panic("no return value specified for DisableDomain") - } - - var r0 domains.Domain - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) (domains.Domain, error)); ok { - return returnFunc(ctx, sesssion, id) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) domains.Domain); ok { - r0 = returnFunc(ctx, sesssion, id) - } else { - r0 = ret.Get(0).(domains.Domain) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string) error); ok { - r1 = returnFunc(ctx, sesssion, id) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_DisableDomain_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'DisableDomain' -type Service_DisableDomain_Call struct { - *mock.Call -} - -// DisableDomain is a helper method to define mock.On call -// - ctx context.Context -// - sesssion authn.Session -// - id string -func (_e *Service_Expecter) DisableDomain(ctx interface{}, sesssion interface{}, id interface{}) *Service_DisableDomain_Call { - return &Service_DisableDomain_Call{Call: _e.mock.On("DisableDomain", ctx, sesssion, id)} -} - -func (_c *Service_DisableDomain_Call) Run(run func(ctx context.Context, sesssion authn.Session, id string)) *Service_DisableDomain_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_DisableDomain_Call) Return(domain domains.Domain, err error) *Service_DisableDomain_Call { - _c.Call.Return(domain, err) - return _c -} - -func (_c *Service_DisableDomain_Call) RunAndReturn(run func(ctx context.Context, sesssion authn.Session, id string) (domains.Domain, error)) *Service_DisableDomain_Call { - _c.Call.Return(run) - return _c -} - -// EnableDomain provides a mock function for the type Service -func (_mock *Service) EnableDomain(ctx context.Context, sesssion authn.Session, id string) (domains.Domain, error) { - ret := _mock.Called(ctx, sesssion, id) - - if len(ret) == 0 { - panic("no return value specified for EnableDomain") - } - - var r0 domains.Domain - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) (domains.Domain, error)); ok { - return returnFunc(ctx, sesssion, id) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) domains.Domain); ok { - r0 = returnFunc(ctx, sesssion, id) - } else { - r0 = ret.Get(0).(domains.Domain) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string) error); ok { - r1 = returnFunc(ctx, sesssion, id) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_EnableDomain_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'EnableDomain' -type Service_EnableDomain_Call struct { - *mock.Call -} - -// EnableDomain is a helper method to define mock.On call -// - ctx context.Context -// - sesssion authn.Session -// - id string -func (_e *Service_Expecter) EnableDomain(ctx interface{}, sesssion interface{}, id interface{}) *Service_EnableDomain_Call { - return &Service_EnableDomain_Call{Call: _e.mock.On("EnableDomain", ctx, sesssion, id)} -} - -func (_c *Service_EnableDomain_Call) Run(run func(ctx context.Context, sesssion authn.Session, id string)) *Service_EnableDomain_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_EnableDomain_Call) Return(domain domains.Domain, err error) *Service_EnableDomain_Call { - _c.Call.Return(domain, err) - return _c -} - -func (_c *Service_EnableDomain_Call) RunAndReturn(run func(ctx context.Context, sesssion authn.Session, id string) (domains.Domain, error)) *Service_EnableDomain_Call { - _c.Call.Return(run) - return _c -} - -// FreezeDomain provides a mock function for the type Service -func (_mock *Service) FreezeDomain(ctx context.Context, sesssion authn.Session, id string) (domains.Domain, error) { - ret := _mock.Called(ctx, sesssion, id) - - if len(ret) == 0 { - panic("no return value specified for FreezeDomain") - } - - var r0 domains.Domain - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) (domains.Domain, error)); ok { - return returnFunc(ctx, sesssion, id) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) domains.Domain); ok { - r0 = returnFunc(ctx, sesssion, id) - } else { - r0 = ret.Get(0).(domains.Domain) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string) error); ok { - r1 = returnFunc(ctx, sesssion, id) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_FreezeDomain_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'FreezeDomain' -type Service_FreezeDomain_Call struct { - *mock.Call -} - -// FreezeDomain is a helper method to define mock.On call -// - ctx context.Context -// - sesssion authn.Session -// - id string -func (_e *Service_Expecter) FreezeDomain(ctx interface{}, sesssion interface{}, id interface{}) *Service_FreezeDomain_Call { - return &Service_FreezeDomain_Call{Call: _e.mock.On("FreezeDomain", ctx, sesssion, id)} -} - -func (_c *Service_FreezeDomain_Call) Run(run func(ctx context.Context, sesssion authn.Session, id string)) *Service_FreezeDomain_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_FreezeDomain_Call) Return(domain domains.Domain, err error) *Service_FreezeDomain_Call { - _c.Call.Return(domain, err) - return _c -} - -func (_c *Service_FreezeDomain_Call) RunAndReturn(run func(ctx context.Context, sesssion authn.Session, id string) (domains.Domain, error)) *Service_FreezeDomain_Call { - _c.Call.Return(run) - return _c -} - -// ListAvailableActions provides a mock function for the type Service -func (_mock *Service) ListAvailableActions(ctx context.Context, session authn.Session) ([]string, error) { - ret := _mock.Called(ctx, session) - - if len(ret) == 0 { - panic("no return value specified for ListAvailableActions") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session) ([]string, error)); ok { - return returnFunc(ctx, session) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session) []string); ok { - r0 = returnFunc(ctx, session) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session) error); ok { - r1 = returnFunc(ctx, session) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_ListAvailableActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ListAvailableActions' -type Service_ListAvailableActions_Call struct { - *mock.Call -} - -// ListAvailableActions is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -func (_e *Service_Expecter) ListAvailableActions(ctx interface{}, session interface{}) *Service_ListAvailableActions_Call { - return &Service_ListAvailableActions_Call{Call: _e.mock.On("ListAvailableActions", ctx, session)} -} - -func (_c *Service_ListAvailableActions_Call) Run(run func(ctx context.Context, session authn.Session)) *Service_ListAvailableActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Service_ListAvailableActions_Call) Return(strings []string, err error) *Service_ListAvailableActions_Call { - _c.Call.Return(strings, err) - return _c -} - -func (_c *Service_ListAvailableActions_Call) RunAndReturn(run func(ctx context.Context, session authn.Session) ([]string, error)) *Service_ListAvailableActions_Call { - _c.Call.Return(run) - return _c -} - -// ListDomainInvitations provides a mock function for the type Service -func (_mock *Service) ListDomainInvitations(ctx context.Context, session authn.Session, page domains.InvitationPageMeta) (domains.InvitationPage, error) { - ret := _mock.Called(ctx, session, page) - - if len(ret) == 0 { - panic("no return value specified for ListDomainInvitations") - } - - var r0 domains.InvitationPage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, domains.InvitationPageMeta) (domains.InvitationPage, error)); ok { - return returnFunc(ctx, session, page) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, domains.InvitationPageMeta) domains.InvitationPage); ok { - r0 = returnFunc(ctx, session, page) - } else { - r0 = ret.Get(0).(domains.InvitationPage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, domains.InvitationPageMeta) error); ok { - r1 = returnFunc(ctx, session, page) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_ListDomainInvitations_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ListDomainInvitations' -type Service_ListDomainInvitations_Call struct { - *mock.Call -} - -// ListDomainInvitations is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - page domains.InvitationPageMeta -func (_e *Service_Expecter) ListDomainInvitations(ctx interface{}, session interface{}, page interface{}) *Service_ListDomainInvitations_Call { - return &Service_ListDomainInvitations_Call{Call: _e.mock.On("ListDomainInvitations", ctx, session, page)} -} - -func (_c *Service_ListDomainInvitations_Call) Run(run func(ctx context.Context, session authn.Session, page domains.InvitationPageMeta)) *Service_ListDomainInvitations_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 domains.InvitationPageMeta - if args[2] != nil { - arg2 = args[2].(domains.InvitationPageMeta) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_ListDomainInvitations_Call) Return(invitations domains.InvitationPage, err error) *Service_ListDomainInvitations_Call { - _c.Call.Return(invitations, err) - return _c -} - -func (_c *Service_ListDomainInvitations_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, page domains.InvitationPageMeta) (domains.InvitationPage, error)) *Service_ListDomainInvitations_Call { - _c.Call.Return(run) - return _c -} - -// ListDomains provides a mock function for the type Service -func (_mock *Service) ListDomains(ctx context.Context, sesssion authn.Session, page domains.Page) (domains.DomainsPage, error) { - ret := _mock.Called(ctx, sesssion, page) - - if len(ret) == 0 { - panic("no return value specified for ListDomains") - } - - var r0 domains.DomainsPage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, domains.Page) (domains.DomainsPage, error)); ok { - return returnFunc(ctx, sesssion, page) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, domains.Page) domains.DomainsPage); ok { - r0 = returnFunc(ctx, sesssion, page) - } else { - r0 = ret.Get(0).(domains.DomainsPage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, domains.Page) error); ok { - r1 = returnFunc(ctx, sesssion, page) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_ListDomains_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ListDomains' -type Service_ListDomains_Call struct { - *mock.Call -} - -// ListDomains is a helper method to define mock.On call -// - ctx context.Context -// - sesssion authn.Session -// - page domains.Page -func (_e *Service_Expecter) ListDomains(ctx interface{}, sesssion interface{}, page interface{}) *Service_ListDomains_Call { - return &Service_ListDomains_Call{Call: _e.mock.On("ListDomains", ctx, sesssion, page)} -} - -func (_c *Service_ListDomains_Call) Run(run func(ctx context.Context, sesssion authn.Session, page domains.Page)) *Service_ListDomains_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 domains.Page - if args[2] != nil { - arg2 = args[2].(domains.Page) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_ListDomains_Call) Return(domainsPage domains.DomainsPage, err error) *Service_ListDomains_Call { - _c.Call.Return(domainsPage, err) - return _c -} - -func (_c *Service_ListDomains_Call) RunAndReturn(run func(ctx context.Context, sesssion authn.Session, page domains.Page) (domains.DomainsPage, error)) *Service_ListDomains_Call { - _c.Call.Return(run) - return _c -} - -// ListEntityMembers provides a mock function for the type Service -func (_mock *Service) ListEntityMembers(ctx context.Context, session authn.Session, entityID string, pq roles.MembersRolePageQuery) (roles.MembersRolePage, error) { - ret := _mock.Called(ctx, session, entityID, pq) - - if len(ret) == 0 { - panic("no return value specified for ListEntityMembers") - } - - var r0 roles.MembersRolePage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, roles.MembersRolePageQuery) (roles.MembersRolePage, error)); ok { - return returnFunc(ctx, session, entityID, pq) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, roles.MembersRolePageQuery) roles.MembersRolePage); ok { - r0 = returnFunc(ctx, session, entityID, pq) - } else { - r0 = ret.Get(0).(roles.MembersRolePage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, roles.MembersRolePageQuery) error); ok { - r1 = returnFunc(ctx, session, entityID, pq) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_ListEntityMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ListEntityMembers' -type Service_ListEntityMembers_Call struct { - *mock.Call -} - -// ListEntityMembers is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - pq roles.MembersRolePageQuery -func (_e *Service_Expecter) ListEntityMembers(ctx interface{}, session interface{}, entityID interface{}, pq interface{}) *Service_ListEntityMembers_Call { - return &Service_ListEntityMembers_Call{Call: _e.mock.On("ListEntityMembers", ctx, session, entityID, pq)} -} - -func (_c *Service_ListEntityMembers_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, pq roles.MembersRolePageQuery)) *Service_ListEntityMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 roles.MembersRolePageQuery - if args[3] != nil { - arg3 = args[3].(roles.MembersRolePageQuery) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_ListEntityMembers_Call) Return(membersRolePage roles.MembersRolePage, err error) *Service_ListEntityMembers_Call { - _c.Call.Return(membersRolePage, err) - return _c -} - -func (_c *Service_ListEntityMembers_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, pq roles.MembersRolePageQuery) (roles.MembersRolePage, error)) *Service_ListEntityMembers_Call { - _c.Call.Return(run) - return _c -} - -// ListInvitations provides a mock function for the type Service -func (_mock *Service) ListInvitations(ctx context.Context, session authn.Session, page domains.InvitationPageMeta) (domains.InvitationPage, error) { - ret := _mock.Called(ctx, session, page) - - if len(ret) == 0 { - panic("no return value specified for ListInvitations") - } - - var r0 domains.InvitationPage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, domains.InvitationPageMeta) (domains.InvitationPage, error)); ok { - return returnFunc(ctx, session, page) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, domains.InvitationPageMeta) domains.InvitationPage); ok { - r0 = returnFunc(ctx, session, page) - } else { - r0 = ret.Get(0).(domains.InvitationPage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, domains.InvitationPageMeta) error); ok { - r1 = returnFunc(ctx, session, page) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_ListInvitations_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ListInvitations' -type Service_ListInvitations_Call struct { - *mock.Call -} - -// ListInvitations is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - page domains.InvitationPageMeta -func (_e *Service_Expecter) ListInvitations(ctx interface{}, session interface{}, page interface{}) *Service_ListInvitations_Call { - return &Service_ListInvitations_Call{Call: _e.mock.On("ListInvitations", ctx, session, page)} -} - -func (_c *Service_ListInvitations_Call) Run(run func(ctx context.Context, session authn.Session, page domains.InvitationPageMeta)) *Service_ListInvitations_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 domains.InvitationPageMeta - if args[2] != nil { - arg2 = args[2].(domains.InvitationPageMeta) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_ListInvitations_Call) Return(invitations domains.InvitationPage, err error) *Service_ListInvitations_Call { - _c.Call.Return(invitations, err) - return _c -} - -func (_c *Service_ListInvitations_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, page domains.InvitationPageMeta) (domains.InvitationPage, error)) *Service_ListInvitations_Call { - _c.Call.Return(run) - return _c -} - -// RejectInvitation provides a mock function for the type Service -func (_mock *Service) RejectInvitation(ctx context.Context, session authn.Session, domainID string) (domains.Invitation, error) { - ret := _mock.Called(ctx, session, domainID) - - if len(ret) == 0 { - panic("no return value specified for RejectInvitation") - } - - var r0 domains.Invitation - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) (domains.Invitation, error)); ok { - return returnFunc(ctx, session, domainID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) domains.Invitation); ok { - r0 = returnFunc(ctx, session, domainID) - } else { - r0 = ret.Get(0).(domains.Invitation) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string) error); ok { - r1 = returnFunc(ctx, session, domainID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RejectInvitation_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RejectInvitation' -type Service_RejectInvitation_Call struct { - *mock.Call -} - -// RejectInvitation is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - domainID string -func (_e *Service_Expecter) RejectInvitation(ctx interface{}, session interface{}, domainID interface{}) *Service_RejectInvitation_Call { - return &Service_RejectInvitation_Call{Call: _e.mock.On("RejectInvitation", ctx, session, domainID)} -} - -func (_c *Service_RejectInvitation_Call) Run(run func(ctx context.Context, session authn.Session, domainID string)) *Service_RejectInvitation_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_RejectInvitation_Call) Return(invitation domains.Invitation, err error) *Service_RejectInvitation_Call { - _c.Call.Return(invitation, err) - return _c -} - -func (_c *Service_RejectInvitation_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, domainID string) (domains.Invitation, error)) *Service_RejectInvitation_Call { - _c.Call.Return(run) - return _c -} - -// RemoveEntityMembers provides a mock function for the type Service -func (_mock *Service) RemoveEntityMembers(ctx context.Context, session authn.Session, entityID string, members []string) error { - ret := _mock.Called(ctx, session, entityID, members) - - if len(ret) == 0 { - panic("no return value specified for RemoveEntityMembers") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, []string) error); ok { - r0 = returnFunc(ctx, session, entityID, members) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RemoveEntityMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveEntityMembers' -type Service_RemoveEntityMembers_Call struct { - *mock.Call -} - -// RemoveEntityMembers is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - members []string -func (_e *Service_Expecter) RemoveEntityMembers(ctx interface{}, session interface{}, entityID interface{}, members interface{}) *Service_RemoveEntityMembers_Call { - return &Service_RemoveEntityMembers_Call{Call: _e.mock.On("RemoveEntityMembers", ctx, session, entityID, members)} -} - -func (_c *Service_RemoveEntityMembers_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, members []string)) *Service_RemoveEntityMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 []string - if args[3] != nil { - arg3 = args[3].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_RemoveEntityMembers_Call) Return(err error) *Service_RemoveEntityMembers_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RemoveEntityMembers_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, members []string) error) *Service_RemoveEntityMembers_Call { - _c.Call.Return(run) - return _c -} - -// RemoveMemberFromAllRoles provides a mock function for the type Service -func (_mock *Service) RemoveMemberFromAllRoles(ctx context.Context, session authn.Session, memberID string) error { - ret := _mock.Called(ctx, session, memberID) - - if len(ret) == 0 { - panic("no return value specified for RemoveMemberFromAllRoles") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) error); ok { - r0 = returnFunc(ctx, session, memberID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RemoveMemberFromAllRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveMemberFromAllRoles' -type Service_RemoveMemberFromAllRoles_Call struct { - *mock.Call -} - -// RemoveMemberFromAllRoles is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - memberID string -func (_e *Service_Expecter) RemoveMemberFromAllRoles(ctx interface{}, session interface{}, memberID interface{}) *Service_RemoveMemberFromAllRoles_Call { - return &Service_RemoveMemberFromAllRoles_Call{Call: _e.mock.On("RemoveMemberFromAllRoles", ctx, session, memberID)} -} - -func (_c *Service_RemoveMemberFromAllRoles_Call) Run(run func(ctx context.Context, session authn.Session, memberID string)) *Service_RemoveMemberFromAllRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_RemoveMemberFromAllRoles_Call) Return(err error) *Service_RemoveMemberFromAllRoles_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RemoveMemberFromAllRoles_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, memberID string) error) *Service_RemoveMemberFromAllRoles_Call { - _c.Call.Return(run) - return _c -} - -// RemoveRole provides a mock function for the type Service -func (_mock *Service) RemoveRole(ctx context.Context, session authn.Session, entityID string, roleID string) error { - ret := _mock.Called(ctx, session, entityID, roleID) - - if len(ret) == 0 { - panic("no return value specified for RemoveRole") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) error); ok { - r0 = returnFunc(ctx, session, entityID, roleID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RemoveRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveRole' -type Service_RemoveRole_Call struct { - *mock.Call -} - -// RemoveRole is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -func (_e *Service_Expecter) RemoveRole(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}) *Service_RemoveRole_Call { - return &Service_RemoveRole_Call{Call: _e.mock.On("RemoveRole", ctx, session, entityID, roleID)} -} - -func (_c *Service_RemoveRole_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string)) *Service_RemoveRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_RemoveRole_Call) Return(err error) *Service_RemoveRole_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RemoveRole_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string) error) *Service_RemoveRole_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveAllRoles provides a mock function for the type Service -func (_mock *Service) RetrieveAllRoles(ctx context.Context, session authn.Session, entityID string, limit uint64, offset uint64) (roles.RolePage, error) { - ret := _mock.Called(ctx, session, entityID, limit, offset) - - if len(ret) == 0 { - panic("no return value specified for RetrieveAllRoles") - } - - var r0 roles.RolePage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, uint64, uint64) (roles.RolePage, error)); ok { - return returnFunc(ctx, session, entityID, limit, offset) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, uint64, uint64) roles.RolePage); ok { - r0 = returnFunc(ctx, session, entityID, limit, offset) - } else { - r0 = ret.Get(0).(roles.RolePage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, uint64, uint64) error); ok { - r1 = returnFunc(ctx, session, entityID, limit, offset) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RetrieveAllRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveAllRoles' -type Service_RetrieveAllRoles_Call struct { - *mock.Call -} - -// RetrieveAllRoles is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - limit uint64 -// - offset uint64 -func (_e *Service_Expecter) RetrieveAllRoles(ctx interface{}, session interface{}, entityID interface{}, limit interface{}, offset interface{}) *Service_RetrieveAllRoles_Call { - return &Service_RetrieveAllRoles_Call{Call: _e.mock.On("RetrieveAllRoles", ctx, session, entityID, limit, offset)} -} - -func (_c *Service_RetrieveAllRoles_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, limit uint64, offset uint64)) *Service_RetrieveAllRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 uint64 - if args[3] != nil { - arg3 = args[3].(uint64) - } - var arg4 uint64 - if args[4] != nil { - arg4 = args[4].(uint64) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_RetrieveAllRoles_Call) Return(rolePage roles.RolePage, err error) *Service_RetrieveAllRoles_Call { - _c.Call.Return(rolePage, err) - return _c -} - -func (_c *Service_RetrieveAllRoles_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, limit uint64, offset uint64) (roles.RolePage, error)) *Service_RetrieveAllRoles_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveDomain provides a mock function for the type Service -func (_mock *Service) RetrieveDomain(ctx context.Context, sesssion authn.Session, id string, withRoles bool) (domains.Domain, error) { - ret := _mock.Called(ctx, sesssion, id, withRoles) - - if len(ret) == 0 { - panic("no return value specified for RetrieveDomain") - } - - var r0 domains.Domain - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, bool) (domains.Domain, error)); ok { - return returnFunc(ctx, sesssion, id, withRoles) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, bool) domains.Domain); ok { - r0 = returnFunc(ctx, sesssion, id, withRoles) - } else { - r0 = ret.Get(0).(domains.Domain) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, bool) error); ok { - r1 = returnFunc(ctx, sesssion, id, withRoles) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RetrieveDomain_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveDomain' -type Service_RetrieveDomain_Call struct { - *mock.Call -} - -// RetrieveDomain is a helper method to define mock.On call -// - ctx context.Context -// - sesssion authn.Session -// - id string -// - withRoles bool -func (_e *Service_Expecter) RetrieveDomain(ctx interface{}, sesssion interface{}, id interface{}, withRoles interface{}) *Service_RetrieveDomain_Call { - return &Service_RetrieveDomain_Call{Call: _e.mock.On("RetrieveDomain", ctx, sesssion, id, withRoles)} -} - -func (_c *Service_RetrieveDomain_Call) Run(run func(ctx context.Context, sesssion authn.Session, id string, withRoles bool)) *Service_RetrieveDomain_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 bool - if args[3] != nil { - arg3 = args[3].(bool) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_RetrieveDomain_Call) Return(domain domains.Domain, err error) *Service_RetrieveDomain_Call { - _c.Call.Return(domain, err) - return _c -} - -func (_c *Service_RetrieveDomain_Call) RunAndReturn(run func(ctx context.Context, sesssion authn.Session, id string, withRoles bool) (domains.Domain, error)) *Service_RetrieveDomain_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveRole provides a mock function for the type Service -func (_mock *Service) RetrieveRole(ctx context.Context, session authn.Session, entityID string, roleID string) (roles.Role, error) { - ret := _mock.Called(ctx, session, entityID, roleID) - - if len(ret) == 0 { - panic("no return value specified for RetrieveRole") - } - - var r0 roles.Role - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) (roles.Role, error)); ok { - return returnFunc(ctx, session, entityID, roleID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) roles.Role); ok { - r0 = returnFunc(ctx, session, entityID, roleID) - } else { - r0 = ret.Get(0).(roles.Role) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RetrieveRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveRole' -type Service_RetrieveRole_Call struct { - *mock.Call -} - -// RetrieveRole is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -func (_e *Service_Expecter) RetrieveRole(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}) *Service_RetrieveRole_Call { - return &Service_RetrieveRole_Call{Call: _e.mock.On("RetrieveRole", ctx, session, entityID, roleID)} -} - -func (_c *Service_RetrieveRole_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string)) *Service_RetrieveRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_RetrieveRole_Call) Return(role roles.Role, err error) *Service_RetrieveRole_Call { - _c.Call.Return(role, err) - return _c -} - -func (_c *Service_RetrieveRole_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string) (roles.Role, error)) *Service_RetrieveRole_Call { - _c.Call.Return(run) - return _c -} - -// RoleAddActions provides a mock function for the type Service -func (_mock *Service) RoleAddActions(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string) ([]string, error) { - ret := _mock.Called(ctx, session, entityID, roleID, actions) - - if len(ret) == 0 { - panic("no return value specified for RoleAddActions") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) ([]string, error)); ok { - return returnFunc(ctx, session, entityID, roleID, actions) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) []string); ok { - r0 = returnFunc(ctx, session, entityID, roleID, actions) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, []string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID, actions) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RoleAddActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleAddActions' -type Service_RoleAddActions_Call struct { - *mock.Call -} - -// RoleAddActions is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - actions []string -func (_e *Service_Expecter) RoleAddActions(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, actions interface{}) *Service_RoleAddActions_Call { - return &Service_RoleAddActions_Call{Call: _e.mock.On("RoleAddActions", ctx, session, entityID, roleID, actions)} -} - -func (_c *Service_RoleAddActions_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string)) *Service_RoleAddActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_RoleAddActions_Call) Return(ops []string, err error) *Service_RoleAddActions_Call { - _c.Call.Return(ops, err) - return _c -} - -func (_c *Service_RoleAddActions_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string) ([]string, error)) *Service_RoleAddActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleAddMembers provides a mock function for the type Service -func (_mock *Service) RoleAddMembers(ctx context.Context, session authn.Session, entityID string, roleID string, members []string) ([]string, error) { - ret := _mock.Called(ctx, session, entityID, roleID, members) - - if len(ret) == 0 { - panic("no return value specified for RoleAddMembers") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) ([]string, error)); ok { - return returnFunc(ctx, session, entityID, roleID, members) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) []string); ok { - r0 = returnFunc(ctx, session, entityID, roleID, members) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, []string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID, members) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RoleAddMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleAddMembers' -type Service_RoleAddMembers_Call struct { - *mock.Call -} - -// RoleAddMembers is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - members []string -func (_e *Service_Expecter) RoleAddMembers(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, members interface{}) *Service_RoleAddMembers_Call { - return &Service_RoleAddMembers_Call{Call: _e.mock.On("RoleAddMembers", ctx, session, entityID, roleID, members)} -} - -func (_c *Service_RoleAddMembers_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, members []string)) *Service_RoleAddMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_RoleAddMembers_Call) Return(strings []string, err error) *Service_RoleAddMembers_Call { - _c.Call.Return(strings, err) - return _c -} - -func (_c *Service_RoleAddMembers_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, members []string) ([]string, error)) *Service_RoleAddMembers_Call { - _c.Call.Return(run) - return _c -} - -// RoleCheckActionsExists provides a mock function for the type Service -func (_mock *Service) RoleCheckActionsExists(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string) (bool, error) { - ret := _mock.Called(ctx, session, entityID, roleID, actions) - - if len(ret) == 0 { - panic("no return value specified for RoleCheckActionsExists") - } - - var r0 bool - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) (bool, error)); ok { - return returnFunc(ctx, session, entityID, roleID, actions) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) bool); ok { - r0 = returnFunc(ctx, session, entityID, roleID, actions) - } else { - r0 = ret.Get(0).(bool) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, []string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID, actions) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RoleCheckActionsExists_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleCheckActionsExists' -type Service_RoleCheckActionsExists_Call struct { - *mock.Call -} - -// RoleCheckActionsExists is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - actions []string -func (_e *Service_Expecter) RoleCheckActionsExists(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, actions interface{}) *Service_RoleCheckActionsExists_Call { - return &Service_RoleCheckActionsExists_Call{Call: _e.mock.On("RoleCheckActionsExists", ctx, session, entityID, roleID, actions)} -} - -func (_c *Service_RoleCheckActionsExists_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string)) *Service_RoleCheckActionsExists_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_RoleCheckActionsExists_Call) Return(b bool, err error) *Service_RoleCheckActionsExists_Call { - _c.Call.Return(b, err) - return _c -} - -func (_c *Service_RoleCheckActionsExists_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string) (bool, error)) *Service_RoleCheckActionsExists_Call { - _c.Call.Return(run) - return _c -} - -// RoleCheckMembersExists provides a mock function for the type Service -func (_mock *Service) RoleCheckMembersExists(ctx context.Context, session authn.Session, entityID string, roleID string, members []string) (bool, error) { - ret := _mock.Called(ctx, session, entityID, roleID, members) - - if len(ret) == 0 { - panic("no return value specified for RoleCheckMembersExists") - } - - var r0 bool - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) (bool, error)); ok { - return returnFunc(ctx, session, entityID, roleID, members) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) bool); ok { - r0 = returnFunc(ctx, session, entityID, roleID, members) - } else { - r0 = ret.Get(0).(bool) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, []string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID, members) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RoleCheckMembersExists_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleCheckMembersExists' -type Service_RoleCheckMembersExists_Call struct { - *mock.Call -} - -// RoleCheckMembersExists is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - members []string -func (_e *Service_Expecter) RoleCheckMembersExists(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, members interface{}) *Service_RoleCheckMembersExists_Call { - return &Service_RoleCheckMembersExists_Call{Call: _e.mock.On("RoleCheckMembersExists", ctx, session, entityID, roleID, members)} -} - -func (_c *Service_RoleCheckMembersExists_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, members []string)) *Service_RoleCheckMembersExists_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_RoleCheckMembersExists_Call) Return(b bool, err error) *Service_RoleCheckMembersExists_Call { - _c.Call.Return(b, err) - return _c -} - -func (_c *Service_RoleCheckMembersExists_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, members []string) (bool, error)) *Service_RoleCheckMembersExists_Call { - _c.Call.Return(run) - return _c -} - -// RoleListActions provides a mock function for the type Service -func (_mock *Service) RoleListActions(ctx context.Context, session authn.Session, entityID string, roleID string) ([]string, error) { - ret := _mock.Called(ctx, session, entityID, roleID) - - if len(ret) == 0 { - panic("no return value specified for RoleListActions") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) ([]string, error)); ok { - return returnFunc(ctx, session, entityID, roleID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) []string); ok { - r0 = returnFunc(ctx, session, entityID, roleID) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RoleListActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleListActions' -type Service_RoleListActions_Call struct { - *mock.Call -} - -// RoleListActions is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -func (_e *Service_Expecter) RoleListActions(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}) *Service_RoleListActions_Call { - return &Service_RoleListActions_Call{Call: _e.mock.On("RoleListActions", ctx, session, entityID, roleID)} -} - -func (_c *Service_RoleListActions_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string)) *Service_RoleListActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_RoleListActions_Call) Return(strings []string, err error) *Service_RoleListActions_Call { - _c.Call.Return(strings, err) - return _c -} - -func (_c *Service_RoleListActions_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string) ([]string, error)) *Service_RoleListActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleListMembers provides a mock function for the type Service -func (_mock *Service) RoleListMembers(ctx context.Context, session authn.Session, entityID string, roleID string, limit uint64, offset uint64) (roles.MembersPage, error) { - ret := _mock.Called(ctx, session, entityID, roleID, limit, offset) - - if len(ret) == 0 { - panic("no return value specified for RoleListMembers") - } - - var r0 roles.MembersPage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, uint64, uint64) (roles.MembersPage, error)); ok { - return returnFunc(ctx, session, entityID, roleID, limit, offset) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, uint64, uint64) roles.MembersPage); ok { - r0 = returnFunc(ctx, session, entityID, roleID, limit, offset) - } else { - r0 = ret.Get(0).(roles.MembersPage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, uint64, uint64) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID, limit, offset) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RoleListMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleListMembers' -type Service_RoleListMembers_Call struct { - *mock.Call -} - -// RoleListMembers is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - limit uint64 -// - offset uint64 -func (_e *Service_Expecter) RoleListMembers(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, limit interface{}, offset interface{}) *Service_RoleListMembers_Call { - return &Service_RoleListMembers_Call{Call: _e.mock.On("RoleListMembers", ctx, session, entityID, roleID, limit, offset)} -} - -func (_c *Service_RoleListMembers_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, limit uint64, offset uint64)) *Service_RoleListMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 uint64 - if args[4] != nil { - arg4 = args[4].(uint64) - } - var arg5 uint64 - if args[5] != nil { - arg5 = args[5].(uint64) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - arg5, - ) - }) - return _c -} - -func (_c *Service_RoleListMembers_Call) Return(membersPage roles.MembersPage, err error) *Service_RoleListMembers_Call { - _c.Call.Return(membersPage, err) - return _c -} - -func (_c *Service_RoleListMembers_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, limit uint64, offset uint64) (roles.MembersPage, error)) *Service_RoleListMembers_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveActions provides a mock function for the type Service -func (_mock *Service) RoleRemoveActions(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string) error { - ret := _mock.Called(ctx, session, entityID, roleID, actions) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveActions") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) error); ok { - r0 = returnFunc(ctx, session, entityID, roleID, actions) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RoleRemoveActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveActions' -type Service_RoleRemoveActions_Call struct { - *mock.Call -} - -// RoleRemoveActions is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - actions []string -func (_e *Service_Expecter) RoleRemoveActions(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, actions interface{}) *Service_RoleRemoveActions_Call { - return &Service_RoleRemoveActions_Call{Call: _e.mock.On("RoleRemoveActions", ctx, session, entityID, roleID, actions)} -} - -func (_c *Service_RoleRemoveActions_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string)) *Service_RoleRemoveActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_RoleRemoveActions_Call) Return(err error) *Service_RoleRemoveActions_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RoleRemoveActions_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string) error) *Service_RoleRemoveActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveAllActions provides a mock function for the type Service -func (_mock *Service) RoleRemoveAllActions(ctx context.Context, session authn.Session, entityID string, roleID string) error { - ret := _mock.Called(ctx, session, entityID, roleID) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveAllActions") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) error); ok { - r0 = returnFunc(ctx, session, entityID, roleID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RoleRemoveAllActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveAllActions' -type Service_RoleRemoveAllActions_Call struct { - *mock.Call -} - -// RoleRemoveAllActions is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -func (_e *Service_Expecter) RoleRemoveAllActions(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}) *Service_RoleRemoveAllActions_Call { - return &Service_RoleRemoveAllActions_Call{Call: _e.mock.On("RoleRemoveAllActions", ctx, session, entityID, roleID)} -} - -func (_c *Service_RoleRemoveAllActions_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string)) *Service_RoleRemoveAllActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_RoleRemoveAllActions_Call) Return(err error) *Service_RoleRemoveAllActions_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RoleRemoveAllActions_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string) error) *Service_RoleRemoveAllActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveAllMembers provides a mock function for the type Service -func (_mock *Service) RoleRemoveAllMembers(ctx context.Context, session authn.Session, entityID string, roleID string) error { - ret := _mock.Called(ctx, session, entityID, roleID) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveAllMembers") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) error); ok { - r0 = returnFunc(ctx, session, entityID, roleID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RoleRemoveAllMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveAllMembers' -type Service_RoleRemoveAllMembers_Call struct { - *mock.Call -} - -// RoleRemoveAllMembers is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -func (_e *Service_Expecter) RoleRemoveAllMembers(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}) *Service_RoleRemoveAllMembers_Call { - return &Service_RoleRemoveAllMembers_Call{Call: _e.mock.On("RoleRemoveAllMembers", ctx, session, entityID, roleID)} -} - -func (_c *Service_RoleRemoveAllMembers_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string)) *Service_RoleRemoveAllMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_RoleRemoveAllMembers_Call) Return(err error) *Service_RoleRemoveAllMembers_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RoleRemoveAllMembers_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string) error) *Service_RoleRemoveAllMembers_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveMembers provides a mock function for the type Service -func (_mock *Service) RoleRemoveMembers(ctx context.Context, session authn.Session, entityID string, roleID string, members []string) error { - ret := _mock.Called(ctx, session, entityID, roleID, members) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveMembers") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) error); ok { - r0 = returnFunc(ctx, session, entityID, roleID, members) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RoleRemoveMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveMembers' -type Service_RoleRemoveMembers_Call struct { - *mock.Call -} - -// RoleRemoveMembers is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - members []string -func (_e *Service_Expecter) RoleRemoveMembers(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, members interface{}) *Service_RoleRemoveMembers_Call { - return &Service_RoleRemoveMembers_Call{Call: _e.mock.On("RoleRemoveMembers", ctx, session, entityID, roleID, members)} -} - -func (_c *Service_RoleRemoveMembers_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, members []string)) *Service_RoleRemoveMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_RoleRemoveMembers_Call) Return(err error) *Service_RoleRemoveMembers_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RoleRemoveMembers_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, members []string) error) *Service_RoleRemoveMembers_Call { - _c.Call.Return(run) - return _c -} - -// SendInvitation provides a mock function for the type Service -func (_mock *Service) SendInvitation(ctx context.Context, session authn.Session, invitation domains.Invitation) (domains.Invitation, error) { - ret := _mock.Called(ctx, session, invitation) - - if len(ret) == 0 { - panic("no return value specified for SendInvitation") - } - - var r0 domains.Invitation - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, domains.Invitation) (domains.Invitation, error)); ok { - return returnFunc(ctx, session, invitation) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, domains.Invitation) domains.Invitation); ok { - r0 = returnFunc(ctx, session, invitation) - } else { - r0 = ret.Get(0).(domains.Invitation) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, domains.Invitation) error); ok { - r1 = returnFunc(ctx, session, invitation) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_SendInvitation_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'SendInvitation' -type Service_SendInvitation_Call struct { - *mock.Call -} - -// SendInvitation is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - invitation domains.Invitation -func (_e *Service_Expecter) SendInvitation(ctx interface{}, session interface{}, invitation interface{}) *Service_SendInvitation_Call { - return &Service_SendInvitation_Call{Call: _e.mock.On("SendInvitation", ctx, session, invitation)} -} - -func (_c *Service_SendInvitation_Call) Run(run func(ctx context.Context, session authn.Session, invitation domains.Invitation)) *Service_SendInvitation_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 domains.Invitation - if args[2] != nil { - arg2 = args[2].(domains.Invitation) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_SendInvitation_Call) Return(invitation1 domains.Invitation, err error) *Service_SendInvitation_Call { - _c.Call.Return(invitation1, err) - return _c -} - -func (_c *Service_SendInvitation_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, invitation domains.Invitation) (domains.Invitation, error)) *Service_SendInvitation_Call { - _c.Call.Return(run) - return _c -} - -// UpdateDomain provides a mock function for the type Service -func (_mock *Service) UpdateDomain(ctx context.Context, sesssion authn.Session, id string, d domains.DomainReq) (domains.Domain, error) { - ret := _mock.Called(ctx, sesssion, id, d) - - if len(ret) == 0 { - panic("no return value specified for UpdateDomain") - } - - var r0 domains.Domain - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, domains.DomainReq) (domains.Domain, error)); ok { - return returnFunc(ctx, sesssion, id, d) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, domains.DomainReq) domains.Domain); ok { - r0 = returnFunc(ctx, sesssion, id, d) - } else { - r0 = ret.Get(0).(domains.Domain) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, domains.DomainReq) error); ok { - r1 = returnFunc(ctx, sesssion, id, d) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_UpdateDomain_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateDomain' -type Service_UpdateDomain_Call struct { - *mock.Call -} - -// UpdateDomain is a helper method to define mock.On call -// - ctx context.Context -// - sesssion authn.Session -// - id string -// - d domains.DomainReq -func (_e *Service_Expecter) UpdateDomain(ctx interface{}, sesssion interface{}, id interface{}, d interface{}) *Service_UpdateDomain_Call { - return &Service_UpdateDomain_Call{Call: _e.mock.On("UpdateDomain", ctx, sesssion, id, d)} -} - -func (_c *Service_UpdateDomain_Call) Run(run func(ctx context.Context, sesssion authn.Session, id string, d domains.DomainReq)) *Service_UpdateDomain_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 domains.DomainReq - if args[3] != nil { - arg3 = args[3].(domains.DomainReq) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_UpdateDomain_Call) Return(domain domains.Domain, err error) *Service_UpdateDomain_Call { - _c.Call.Return(domain, err) - return _c -} - -func (_c *Service_UpdateDomain_Call) RunAndReturn(run func(ctx context.Context, sesssion authn.Session, id string, d domains.DomainReq) (domains.Domain, error)) *Service_UpdateDomain_Call { - _c.Call.Return(run) - return _c -} - -// UpdateRoleName provides a mock function for the type Service -func (_mock *Service) UpdateRoleName(ctx context.Context, session authn.Session, entityID string, roleID string, newRoleName string) (roles.Role, error) { - ret := _mock.Called(ctx, session, entityID, roleID, newRoleName) - - if len(ret) == 0 { - panic("no return value specified for UpdateRoleName") - } - - var r0 roles.Role - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, string) (roles.Role, error)); ok { - return returnFunc(ctx, session, entityID, roleID, newRoleName) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, string) roles.Role); ok { - r0 = returnFunc(ctx, session, entityID, roleID, newRoleName) - } else { - r0 = ret.Get(0).(roles.Role) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID, newRoleName) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_UpdateRoleName_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateRoleName' -type Service_UpdateRoleName_Call struct { - *mock.Call -} - -// UpdateRoleName is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - newRoleName string -func (_e *Service_Expecter) UpdateRoleName(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, newRoleName interface{}) *Service_UpdateRoleName_Call { - return &Service_UpdateRoleName_Call{Call: _e.mock.On("UpdateRoleName", ctx, session, entityID, roleID, newRoleName)} -} - -func (_c *Service_UpdateRoleName_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, newRoleName string)) *Service_UpdateRoleName_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 string - if args[4] != nil { - arg4 = args[4].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_UpdateRoleName_Call) Return(role roles.Role, err error) *Service_UpdateRoleName_Call { - _c.Call.Return(role, err) - return _c -} - -func (_c *Service_UpdateRoleName_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, newRoleName string) (roles.Role, error)) *Service_UpdateRoleName_Call { - _c.Call.Return(run) - return _c -} diff --git a/domains/operations/operations.go b/domains/operations/operations.go deleted file mode 100644 index 41e487683..000000000 --- a/domains/operations/operations.go +++ /dev/null @@ -1,138 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package operations - -import ( - "github.com/absmach/magistrala/pkg/permissions" -) - -const ( - OpCreateDomain permissions.Operation = iota - OpUpdateDomain - OpRetrieveDomain - OpEnableDomain - OpDisableDomain - OpFreezeDomain - OpListDomains - - OpSendDomainInvitation - OpListDomainInvitations - OpDeleteDomainInvitation - - OpViewInvitation - OpListInvitations - OpAcceptInvitation - OpRejectInvitation - - OpCreateDomainClients - OpListDomainClients - OpCreateDomainChannels - OpListDomainChannels - OpCreateDomainGroups - OpListDomainGroups -) - -func OperationDetails() map[permissions.Operation]permissions.OperationDetails { - ops := map[permissions.Operation]permissions.OperationDetails{ - OpUpdateDomain: { - Name: "update", - PermissionRequired: true, - }, - OpRetrieveDomain: { - Name: "read", - PermissionRequired: true, - }, - OpListDomains: { - Name: "list", - PermissionRequired: true, - }, - OpEnableDomain: { - Name: "enable", - PermissionRequired: true, - }, - OpDisableDomain: { - Name: "disable", - PermissionRequired: true, - }, - - // Permission not required, only Super Admin can freeze the domain - OpFreezeDomain: { - Name: "freeze", - PermissionRequired: false, - }, - - OpCreateDomain: { - Name: "create", - PermissionRequired: false, - }, - - // Domain Invitation related permissions - OpSendDomainInvitation: { - Name: "send_invitation", - PermissionRequired: true, - }, - - OpDeleteDomainInvitation: { - Name: "delete_invitation", - PermissionRequired: true, - }, - - OpListDomainInvitations: { - Name: "list_domain_invitation", - PermissionRequired: true, - }, - - // User Invitation related permissions - OpViewInvitation: { - Name: "view_invitation", - PermissionRequired: false, - }, - OpListInvitations: { - Name: "list_invitation", - PermissionRequired: false, - }, - - OpAcceptInvitation: { - Name: "accept_invitation", - PermissionRequired: false, - }, - - OpRejectInvitation: { - Name: "reject_invitation", - PermissionRequired: false, - }, - - // Operations related to entities - OpCreateDomainClients: { - Name: "create_clients", - PermissionRequired: true, - }, - - OpListDomainClients: { - Name: "list_clients", - PermissionRequired: true, - }, - - OpCreateDomainChannels: { - Name: "create_channels", - PermissionRequired: true, - }, - - OpListDomainChannels: { - Name: "list_channels", - PermissionRequired: true, - }, - - OpCreateDomainGroups: { - Name: "create_groups", - PermissionRequired: true, - }, - - OpListDomainGroups: { - Name: "list_groups", - PermissionRequired: true, - }, - } - return ops -} diff --git a/domains/postgres/doc.go b/domains/postgres/doc.go deleted file mode 100644 index ac5c81ae1..000000000 --- a/domains/postgres/doc.go +++ /dev/null @@ -1,6 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package postgres contains Key repository implementations using -// PostgreSQL as the underlying database. -package postgres diff --git a/domains/postgres/domains.go b/domains/postgres/domains.go deleted file mode 100644 index b5ae8f117..000000000 --- a/domains/postgres/domains.go +++ /dev/null @@ -1,791 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package postgres - -import ( - "context" - "database/sql" - "encoding/json" - "fmt" - "strings" - "time" - - api "github.com/absmach/magistrala/api/http" - "github.com/absmach/magistrala/domains" - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - "github.com/absmach/magistrala/pkg/policies" - "github.com/absmach/magistrala/pkg/postgres" - "github.com/absmach/magistrala/pkg/roles" - rolesPostgres "github.com/absmach/magistrala/pkg/roles/repo/postgres" - "github.com/jackc/pgtype" - "github.com/lib/pq" -) - -var _ domains.Repository = (*domainRepo)(nil) - -const ( - rolesTableNamePrefix = "domains" - entityTableName = "domains" - entityIDColumnName = "id" -) - -type domainRepo struct { - db postgres.Database - eh errors.Handler - rolesPostgres.Repository -} - -// NewRepository instantiates a PostgreSQL -// implementation of Domain repository. -func NewRepository(db postgres.Database) domains.Repository { - rmsvcRepo := rolesPostgres.NewRepository(db, policies.DomainType, rolesTableNamePrefix, entityTableName, entityIDColumnName) - errHandlerOptions := []errors.HandlerOption{ - postgres.WithDuplicateErrors(NewDuplicateErrors()), - } - return &domainRepo{ - db: db, - eh: postgres.NewErrorHandler(errHandlerOptions...), - Repository: rmsvcRepo, - } -} - -func (repo domainRepo) SaveDomain(ctx context.Context, d domains.Domain) (dd domains.Domain, err error) { - q := `INSERT INTO domains (id, name, tags, route, metadata, created_at, updated_at, updated_by, created_by, status) - VALUES (:id, :name, :tags, :route, :metadata, :created_at, :updated_at, :updated_by, :created_by, :status) - RETURNING id, name, tags, route, metadata, created_at, updated_at, updated_by, created_by, status;` - - dbd, err := toDBDomain(d) - if err != nil { - return domains.Domain{}, errors.Wrap(repoerr.ErrCreateEntity, errors.ErrRollbackTx) - } - - row, err := repo.db.NamedQueryContext(ctx, q, dbd) - if err != nil { - return domains.Domain{}, repo.eh.HandleError(repoerr.ErrCreateEntity, err) - } - defer row.Close() - - if !row.Next() { - return domains.Domain{}, repoerr.ErrNotFound - } - - dbd = dbDomain{} - if err := row.StructScan(&dbd); err != nil { - return domains.Domain{}, repo.eh.HandleError(repoerr.ErrFailedOpDB, err) - } - - domain, err := toDomain(dbd) - if err != nil { - return domains.Domain{}, errors.Wrap(repoerr.ErrFailedOpDB, err) - } - - return domain, nil -} - -// RetrieveDomainByIDWithRoles retrieves Domain by its unique ID along with member roles. -func (repo domainRepo) RetrieveDomainByIDWithRoles(ctx context.Context, id string, memberID string) (domains.Domain, error) { - q := ` - WITH all_roles AS ( - SELECT - d.id AS domain_id, - drm.member_id AS member_id, - dr.id AS role_id, - dr."name" AS role_name, - jsonb_agg(DISTINCT all_actions."action") AS actions, - 'direct' AS access_type, - '' AS access_provider_path, - '' AS access_provider_id - FROM - domains d - JOIN - domains_roles dr ON - dr.entity_id = d.id - JOIN - domains_role_members drm ON - dr.id = drm.role_id - JOIN - domains_role_actions dra ON - dr.id = dra.role_id - JOIN - domains_role_actions all_actions ON - dr.id = all_actions.role_id - WHERE - d.id = :id - AND - drm.member_id = :member_id - GROUP BY - d.id, - dr.id, - dr."name", - drm.member_id - ), - final_roles AS ( - SELECT - ar.domain_id, - ar.member_id, - jsonb_agg( - jsonb_build_object( - 'role_id', ar.role_id, - 'role_name', ar.role_name, - 'actions', ar.actions, - 'access_type', ar.access_type, - 'access_provider_path', ar.access_provider_path, - 'access_provider_id', ar.access_provider_id - ) - ) AS roles - FROM - all_roles ar - GROUP BY - ar.domain_id, - ar.member_id - ) - SELECT - d.id, - d.name, - d.tags, - d.route, - d.metadata, - d.created_at, - d.updated_at, - d.updated_by, - d.created_by, - d.status, - fr.member_id, - fr.roles - FROM - domains d - JOIN final_roles fr ON - d.id = fr.domain_id - ` - - dbdp := dbDomainsPage{ - ID: id, - UserID: memberID, - } - - rows, err := repo.db.NamedQueryContext(ctx, q, dbdp) - if err != nil { - return domains.Domain{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - defer rows.Close() - - dbd := dbDomain{} - if rows.Next() { - if err = rows.StructScan(&dbd); err != nil { - return domains.Domain{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - - domain, err := toDomain(dbd) - if err != nil { - return domains.Domain{}, errors.Wrap(repoerr.ErrFailedOpDB, err) - } - - return domain, nil - } - return domains.Domain{}, repoerr.ErrNotFound -} - -// RetrieveDomainByID retrieves Domain by its unique ID. -func (repo domainRepo) RetrieveDomainByID(ctx context.Context, id string) (domains.Domain, error) { - q := `SELECT d.id as id, d.name as name, d.tags as tags, d.route as route, d.metadata as metadata, d.created_at as created_at, d.updated_at as updated_at, d.updated_by as updated_by, d.created_by as created_by, d.status as status - FROM domains d WHERE d.id = :id` - - dbdp := dbDomainsPage{ - ID: id, - } - - rows, err := repo.db.NamedQueryContext(ctx, q, dbdp) - if err != nil { - return domains.Domain{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - defer rows.Close() - - dbd := dbDomain{} - if rows.Next() { - if err = rows.StructScan(&dbd); err != nil { - return domains.Domain{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - - domain, err := toDomain(dbd) - if err != nil { - return domains.Domain{}, errors.Wrap(repoerr.ErrFailedOpDB, err) - } - - return domain, nil - } - return domains.Domain{}, repoerr.ErrNotFound -} - -// RetrieveDomainByRoute retrieves Domain by its unique route. -func (repo domainRepo) RetrieveDomainByRoute(ctx context.Context, route string) (domains.Domain, error) { - q := `SELECT d.id as id, d.name as name, d.tags as tags, d.route as route, d.metadata as metadata, d.created_at as created_at, d.updated_at as updated_at, d.updated_by as updated_by, d.created_by as created_by, d.status as status - FROM domains d WHERE d.route = :route` - - dbdom := dbDomain{ - Route: &route, - } - - rows, err := repo.db.NamedQueryContext(ctx, q, dbdom) - if err != nil { - return domains.Domain{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - defer rows.Close() - - dbd := dbDomain{} - if rows.Next() { - if err = rows.StructScan(&dbd); err != nil { - return domains.Domain{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - - domain, err := toDomain(dbd) - if err != nil { - return domains.Domain{}, errors.Wrap(repoerr.ErrFailedOpDB, err) - } - - return domain, nil - } - return domains.Domain{}, repoerr.ErrNotFound -} - -// RetrieveAllByIDs retrieves for given Domain IDs . -func (repo domainRepo) RetrieveAllDomainsByIDs(ctx context.Context, pm domains.Page) (domains.DomainsPage, error) { - if len(pm.IDs) == 0 { - return domains.DomainsPage{}, nil - } - query, err := buildPageQuery(pm) - if err != nil { - return domains.DomainsPage{}, errors.Wrap(repoerr.ErrFailedOpDB, err) - } - - baseQ := `SELECT d.id as id, d.name as name, d.tags as tags, d.route as route, d.metadata as metadata, - d.created_at as created_at, d.updated_at as updated_at, d.updated_by as updated_by, - d.created_by as created_by, d.status as status, COUNT(*) OVER() AS total_count FROM domains d` - - squery := applyOrdering(query, pm) - - q := fmt.Sprintf("%s %s LIMIT %d OFFSET %d;", baseQ, squery, pm.Limit, pm.Offset) - - dbPage, err := toDBDomainsPage(pm) - if err != nil { - return domains.DomainsPage{}, errors.Wrap(repoerr.ErrFailedToRetrieveAllGroups, err) - } - - rows, err := repo.db.NamedQueryContext(ctx, q, dbPage) - if err != nil { - return domains.DomainsPage{}, repo.eh.HandleError(repoerr.ErrFailedToRetrieveAllGroups, err) - } - defer rows.Close() - - var total uint64 - var doms []domains.Domain - for rows.Next() { - dbd := dbDomain{} - if err := rows.StructScan(&dbd); err != nil { - return domains.DomainsPage{}, repo.eh.HandleError(repoerr.ErrFailedToRetrieveAllGroups, err) - } - total = dbd.TotalCount - d, err := toDomain(dbd) - if err != nil { - return domains.DomainsPage{}, repo.eh.HandleError(repoerr.ErrFailedToRetrieveAllGroups, err) - } - doms = append(doms, d) - } - - if len(doms) == 0 { - cq := "SELECT COUNT(*) FROM domains d" - if query != "" { - cq = fmt.Sprintf(" %s %s", cq, query) - } - total, err = postgres.Total(ctx, repo.db, cq, dbPage) - if err != nil { - return domains.DomainsPage{}, repo.eh.HandleError(repoerr.ErrFailedToRetrieveAllGroups, err) - } - } - - return domains.DomainsPage{ - Total: total, - Offset: pm.Offset, - Limit: pm.Limit, - Domains: doms, - }, nil -} - -// ListDomains list domains of user. -func (repo domainRepo) ListDomains(ctx context.Context, pm domains.Page) (domains.DomainsPage, error) { - query, err := buildPageQuery(pm) - if err != nil { - return domains.DomainsPage{}, errors.Wrap(repoerr.ErrFailedOpDB, err) - } - squery := applyOrdering(query, pm) - - q := `SELECT - d.id as id, - d.name as name, - d.tags as tags, - d.route as route, - d.metadata as metadata, - d.created_at as created_at, - d.updated_at as updated_at, - d.updated_by as updated_by, - d.created_by as created_by, - d.status as status, - COUNT(*) OVER() AS total_count - FROM - domains as d - %s - LIMIT :limit OFFSET :offset` - - if pm.UserID != "" { - q = userDomainsBaseQuery + - ` - SELECT - d.id as id, - d.name as name, - d.tags as tags, - d.route as route, - d.metadata as metadata, - d.status as status, - d.role_id AS role_id, - d.role_name AS role_name, - d.actions AS actions, - d.created_at as created_at, - d.updated_at as updated_at, - d.updated_by as updated_by, - d.created_by as created_by, - COUNT(*) OVER() AS total_count - FROM - domains d - %s - LIMIT :limit OFFSET :offset - ` - } - - q = fmt.Sprintf(q, squery) - - dbPage, err := toDBDomainsPage(pm) - if err != nil { - return domains.DomainsPage{}, errors.Wrap(repoerr.ErrFailedToRetrieveAllGroups, err) - } - - if pm.OnlyTotal { - cq := `SELECT COUNT(*) FROM domains as d %s` - if pm.UserID != "" { - cq = userDomainsBaseQuery + cq - } - if query != "" { - cq = fmt.Sprintf(cq, query) - } - total, err := postgres.Total(ctx, repo.db, cq, dbPage) - if err != nil { - return domains.DomainsPage{}, repo.eh.HandleError(repoerr.ErrFailedToRetrieveAllGroups, err) - } - return domains.DomainsPage{Total: total, Offset: pm.Offset, Limit: pm.Limit}, nil - } - - rows, err := repo.db.NamedQueryContext(ctx, q, dbPage) - if err != nil { - return domains.DomainsPage{}, repo.eh.HandleError(repoerr.ErrFailedToRetrieveAllGroups, err) - } - defer rows.Close() - - var total uint64 - var doms []domains.Domain - for rows.Next() { - dbd := dbDomain{} - if err := rows.StructScan(&dbd); err != nil { - return domains.DomainsPage{}, repo.eh.HandleError(repoerr.ErrFailedToRetrieveAllGroups, err) - } - total = dbd.TotalCount - d, err := toDomain(dbd) - if err != nil { - return domains.DomainsPage{}, repo.eh.HandleError(repoerr.ErrFailedToRetrieveAllGroups, err) - } - doms = append(doms, d) - } - - if len(doms) == 0 { - cq := `SELECT COUNT(*) FROM domains as d %s` - if pm.UserID != "" { - cq = userDomainsBaseQuery + cq - } - if query != "" { - cq = fmt.Sprintf(cq, query) - } - total, err = postgres.Total(ctx, repo.db, cq, dbPage) - if err != nil { - return domains.DomainsPage{}, repo.eh.HandleError(repoerr.ErrFailedToRetrieveAllGroups, err) - } - } - - return domains.DomainsPage{ - Total: total, - Offset: pm.Offset, - Limit: pm.Limit, - Domains: doms, - }, nil -} - -// UpdateDomain updates the client name and metadata. -func (repo domainRepo) UpdateDomain(ctx context.Context, id string, dr domains.DomainReq) (domains.Domain, error) { - var query []string - var upq string - d := domains.Domain{ID: id} - - if dr.Name != nil && *dr.Name != "" { - query = append(query, "name = :name") - d.Name = *dr.Name - } - if dr.Metadata != nil { - query = append(query, "metadata = :metadata") - d.Metadata = *dr.Metadata - } - if dr.Tags != nil { - query = append(query, "tags = :tags") - d.Tags = *dr.Tags - } - if dr.Status != nil { - query = append(query, "status = :status") - d.Status = *dr.Status - } - d.UpdatedAt = time.Now().UTC() - if dr.UpdatedAt != nil { - query = append(query, "updated_at = :updated_at") - d.UpdatedAt = *dr.UpdatedAt - } - if dr.UpdatedBy != nil { - query = append(query, "updated_by = :updated_by") - d.UpdatedAt = *dr.UpdatedAt - } - - if len(query) > 0 { - upq = strings.Join(query, ", ") - } - - q := fmt.Sprintf(`UPDATE domains SET %s - WHERE id = :id - RETURNING id, name, tags, route, metadata, created_at, updated_at, updated_by, created_by, status;`, upq) - - dbd, err := toDBDomain(d) - if err != nil { - return domains.Domain{}, errors.Wrap(repoerr.ErrUpdateEntity, err) - } - - row, err := repo.db.NamedQueryContext(ctx, q, dbd) - if err != nil { - return domains.Domain{}, repo.eh.HandleError(repoerr.ErrUpdateEntity, err) - } - defer row.Close() - - if !row.Next() { - return domains.Domain{}, repoerr.ErrNotFound - } - - dbd = dbDomain{} - if err := row.StructScan(&dbd); err != nil { - return domains.Domain{}, repo.eh.HandleError(repoerr.ErrFailedOpDB, err) - } - - domain, err := toDomain(dbd) - if err != nil { - return domains.Domain{}, errors.Wrap(repoerr.ErrFailedOpDB, err) - } - - return domain, nil -} - -// Delete delete domain from database. -func (repo domainRepo) DeleteDomain(ctx context.Context, id string) error { - q := "DELETE FROM domains WHERE id = $1;" - - res, err := repo.db.ExecContext(ctx, q, id) - if err != nil { - return repo.eh.HandleError(repoerr.ErrRemoveEntity, err) - } - if rows, _ := res.RowsAffected(); rows == 0 { - return repoerr.ErrNotFound - } - - return nil -} - -const userDomainsBaseQuery = ` - with domains AS ( - SELECT - d.id as id, - d.name as name, - d.tags as tags, - d.route as route, - d.metadata as metadata, - d.created_at as created_at, - d.updated_at as updated_at, - d.updated_by as updated_by, - d.created_by as created_by, - d.status as status, - dr.entity_id AS entity_id, - drm.member_id AS member_id, - dr.id AS role_id, - dr."name" AS role_name, - array_agg(dra."action") AS actions - FROM - domains_role_members drm - JOIN - domains_role_actions dra ON dra.role_id = drm.role_id - JOIN - domains_roles dr ON dr.id = drm.role_id - JOIN - "domains" d ON d.id = dr.entity_id - WHERE - drm.member_id = :member_id - GROUP BY - dr.entity_id, drm.member_id, dr.id, dr."name", d.id - )` - -func applyOrdering(emq string, pm domains.Page) string { - col := "COALESCE(d.updated_at, d.created_at)" - - switch pm.Order { - case "name": - col = "d.name" - case "created_at": - col = "d.created_at" - case "updated_at", "": - col = "COALESCE(d.updated_at, d.created_at)" - } - - dir := pm.Dir - if dir != api.AscDir && dir != api.DescDir { - dir = api.DescDir - } - - return fmt.Sprintf("%s ORDER BY %s %s, d.id %s", emq, col, dir, dir) -} - -type dbDomain struct { - ID string `db:"id"` - Name string `db:"name"` - Metadata []byte `db:"metadata,omitempty"` - Tags pgtype.TextArray `db:"tags,omitempty"` - Route *string `db:"route,omitempty"` - Status domains.Status `db:"status"` - RoleID string `db:"role_id"` - RoleName string `db:"role_name"` - Actions pq.StringArray `db:"actions"` - CreatedBy string `db:"created_by"` - CreatedAt time.Time `db:"created_at"` - UpdatedBy *string `db:"updated_by,omitempty"` - UpdatedAt sql.NullTime `db:"updated_at,omitempty"` - MemberID string `db:"member_id,omitempty"` - Roles json.RawMessage `db:"roles,omitempty"` - TotalCount uint64 `db:"total_count"` -} - -func toDBDomain(d domains.Domain) (dbDomain, error) { - data := []byte("{}") - if len(d.Metadata) > 0 { - b, err := json.Marshal(d.Metadata) - if err != nil { - return dbDomain{}, errors.Wrap(errors.ErrMalformedEntity, err) - } - data = b - } - var tags pgtype.TextArray - if err := tags.Set(d.Tags); err != nil { - return dbDomain{}, err - } - var route *string - if d.Route != "" { - route = &d.Route - } - - var updatedBy *string - if d.UpdatedBy != "" { - updatedBy = &d.UpdatedBy - } - var updatedAt sql.NullTime - if d.UpdatedAt != (time.Time{}) { - updatedAt = sql.NullTime{Time: d.UpdatedAt, Valid: true} - } - - return dbDomain{ - ID: d.ID, - Name: d.Name, - Metadata: data, - Tags: tags, - Route: route, - Status: d.Status, - RoleID: d.RoleID, - CreatedBy: d.CreatedBy, - CreatedAt: d.CreatedAt, - UpdatedBy: updatedBy, - UpdatedAt: updatedAt, - }, nil -} - -func toDomain(d dbDomain) (domains.Domain, error) { - var metadata domains.Metadata - if d.Metadata != nil { - if err := json.Unmarshal([]byte(d.Metadata), &metadata); err != nil { - return domains.Domain{}, errors.Wrap(errors.ErrMalformedEntity, err) - } - } - var tags []string - for _, e := range d.Tags.Elements { - tags = append(tags, e.String) - } - var route string - if d.Route != nil { - route = *d.Route - } - var updatedBy string - if d.UpdatedBy != nil { - updatedBy = *d.UpdatedBy - } - var updatedAt time.Time - if d.UpdatedAt.Valid { - updatedAt = d.UpdatedAt.Time.UTC() - } - - var mra []roles.MemberRoleActions - if d.Roles != nil { - if err := json.Unmarshal(d.Roles, &mra); err != nil { - return domains.Domain{}, errors.Wrap(errors.ErrMalformedEntity, err) - } - } - - return domains.Domain{ - ID: d.ID, - Name: d.Name, - Metadata: metadata, - Tags: tags, - Route: route, - RoleID: d.RoleID, - RoleName: d.RoleName, - Actions: d.Actions, - Status: d.Status, - CreatedBy: d.CreatedBy, - CreatedAt: d.CreatedAt.UTC(), - UpdatedBy: updatedBy, - UpdatedAt: updatedAt, - MemberID: d.MemberID, - Roles: mra, - }, nil -} - -type dbDomainsPage struct { - Total uint64 `db:"total"` - Limit uint64 `db:"limit"` - Offset uint64 `db:"offset"` - Order string `db:"order"` - Dir string `db:"dir"` - Name string `db:"name"` - RoleID string `db:"role_id"` - RoleName string `db:"role_name"` - Actions pq.StringArray `db:"actions"` - ID string `db:"id"` - IDs pq.StringArray `db:"ids"` - Metadata []byte `db:"metadata"` - Tags pgtype.TextArray `db:"tags"` - Status domains.Status `db:"status"` - UserID string `db:"member_id"` - CreatedFrom time.Time `db:"created_from"` - CreatedTo time.Time `db:"created_to"` -} - -func toDBDomainsPage(pm domains.Page) (dbDomainsPage, error) { - _, data, err := postgres.CreateMetadataQuery("", pm.Metadata) - if err != nil { - return dbDomainsPage{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - var tags pgtype.TextArray - if err := tags.Set(pm.Tags.Elements); err != nil { - return dbDomainsPage{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - return dbDomainsPage{ - Total: pm.Total, - Limit: pm.Limit, - Offset: pm.Offset, - Order: pm.Order, - Dir: pm.Dir, - Name: pm.Name, - RoleID: pm.RoleID, - RoleName: pm.RoleName, - Actions: pm.Actions, - ID: pm.ID, - IDs: pq.StringArray(pm.IDs), - Metadata: data, - Tags: tags, - Status: pm.Status, - UserID: pm.UserID, - CreatedFrom: pm.CreatedFrom, - CreatedTo: pm.CreatedTo, - }, nil -} - -func buildPageQuery(pm domains.Page) (string, error) { - var query []string - var emq string - - if pm.ID != "" { - query = append(query, "d.id = :id") - } - - if len(pm.IDs) != 0 { - query = append(query, "d.id = ANY(:ids)") - } - - if (pm.Status >= domains.EnabledStatus) && (pm.Status < domains.AllStatus) { - query = append(query, "d.status = :status") - } else { - query = append(query, fmt.Sprintf("d.status < %d", domains.AllStatus)) - } - - if pm.Name != "" { - query = append(query, "d.name ILIKE '%' || :name || '%'") - } - - if pm.UserID != "" { - if pm.RoleName != "" { - query = append(query, "d.role_name = :role_name") - } - - if pm.RoleID != "" { - query = append(query, "d.role_id = :role_id") - } - - if len(pm.Actions) != 0 { - query = append(query, "d.actions @> :actions") - } - } - - if len(pm.Tags.Elements) > 0 { - switch pm.Tags.Operator { - case domains.AndOp: - query = append(query, "tags @> :tags") - default: // OR - query = append(query, "tags && :tags") - } - } - - mq, _, err := postgres.CreateMetadataQuery("", pm.Metadata) - if err != nil { - return "", errors.Wrap(repoerr.ErrViewEntity, err) - } - if mq != "" { - query = append(query, mq) - } - - if !pm.CreatedFrom.IsZero() { - query = append(query, "d.created_at >= :created_from") - } - if !pm.CreatedTo.IsZero() { - query = append(query, "d.created_at <= :created_to") - } - - if len(query) > 0 { - emq = fmt.Sprintf("WHERE %s", strings.Join(query, " AND ")) - } - - return emq, nil -} diff --git a/domains/postgres/domains_test.go b/domains/postgres/domains_test.go deleted file mode 100644 index ceba6420f..000000000 --- a/domains/postgres/domains_test.go +++ /dev/null @@ -1,1201 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package postgres_test - -import ( - "context" - "fmt" - "testing" - "time" - - "github.com/absmach/magistrala/domains" - "github.com/absmach/magistrala/domains/postgres" - "github.com/absmach/magistrala/internal/testsutil" - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" -) - -const ( - invalid = "invalid" - ascDir = "asc" - descDir = "desc" - defOrder = "created_at" -) - -var ( - domainID = testsutil.GenerateUUID(&testing.T{}) - userID = testsutil.GenerateUUID(&testing.T{}) - errDomainExists = errors.New("domain already exists") -) - -func TestSaveDomain(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM domains") - require.Nil(t, err, fmt.Sprintf("clean domains unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - cases := []struct { - desc string - domain domains.Domain - err error - }{ - { - desc: "add new domain with all fields successfully", - domain: domains.Domain{ - ID: domainID, - Name: "test", - Route: "test", - Tags: []string{"test"}, - Metadata: map[string]any{ - "test": "test", - }, - CreatedAt: time.Now().UTC().Truncate(time.Microsecond), - UpdatedAt: time.Now().UTC().Truncate(time.Microsecond), - CreatedBy: userID, - UpdatedBy: userID, - Status: domains.EnabledStatus, - }, - err: nil, - }, - { - desc: "add the same domain again", - domain: domains.Domain{ - ID: domainID, - Name: "test", - Route: "test", - Tags: []string{"test"}, - Metadata: map[string]any{ - "test": "test", - }, - CreatedAt: time.Now().UTC().Truncate(time.Microsecond), - UpdatedAt: time.Now().UTC().Truncate(time.Microsecond), - CreatedBy: userID, - UpdatedBy: userID, - Status: domains.EnabledStatus, - }, - err: errDomainExists, - }, - { - desc: "add domain with empty ID", - domain: domains.Domain{ - ID: "", - Name: "test1", - Route: "test1", - Tags: []string{"test"}, - Metadata: map[string]any{ - "test": "test", - }, - CreatedAt: time.Now().UTC().Truncate(time.Microsecond), - UpdatedAt: time.Now().UTC().Truncate(time.Microsecond), - CreatedBy: userID, - UpdatedBy: userID, - Status: domains.EnabledStatus, - }, - err: nil, - }, - { - desc: "add domain with empty route", - domain: domains.Domain{ - ID: testsutil.GenerateUUID(&testing.T{}), - Name: "test1", - Route: "", - Tags: []string{"test"}, - Metadata: map[string]any{ - "test": "test", - }, - CreatedAt: time.Now(), - UpdatedAt: time.Now(), - CreatedBy: userID, - UpdatedBy: userID, - Status: domains.EnabledStatus, - }, - err: repoerr.ErrCreateEntity, - }, - { - desc: "add domain with malformed metadata", - domain: domains.Domain{ - ID: domainID, - Name: "test1", - Route: "test1", - Tags: []string{"test"}, - Metadata: map[string]any{ - "key": make(chan int), - }, - CreatedAt: time.Now(), - UpdatedAt: time.Now(), - CreatedBy: userID, - UpdatedBy: userID, - Status: domains.EnabledStatus, - }, - err: repoerr.ErrCreateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - domain, err := repo.SaveDomain(context.Background(), tc.domain) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - if err == nil { - assert.Equal(t, tc.domain, domain, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.domain, domain)) - } - }) - } -} - -func TestRetrieveByID(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM domains") - require.Nil(t, err, fmt.Sprintf("clean domains unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - domain := domains.Domain{ - ID: domainID, - Name: "test", - Route: "test", - Tags: []string{"test"}, - Metadata: map[string]any{ - "test": "test", - }, - CreatedBy: userID, - UpdatedBy: userID, - CreatedAt: time.Now().UTC().Truncate(time.Microsecond), - UpdatedAt: time.Now().UTC().Truncate(time.Microsecond), - Status: domains.EnabledStatus, - } - - _, err := repo.SaveDomain(context.Background(), domain) - require.Nil(t, err, fmt.Sprintf("failed to save domain %s", domain.ID)) - - cases := []struct { - desc string - domainID string - response domains.Domain - err error - }{ - { - desc: "retrieve existing domain", - domainID: domain.ID, - response: domain, - err: nil, - }, - { - desc: "retrieve non-existing domain", - domainID: invalid, - response: domains.Domain{}, - err: repoerr.ErrNotFound, - }, - { - desc: "retrieve with empty domain id", - domainID: "", - response: domains.Domain{}, - err: repoerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - d, err := repo.RetrieveDomainByID(context.Background(), tc.domainID) - assert.Equal(t, tc.response, d, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.response, d)) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.err, err)) - }) - } -} - -func TestRetrieveByRoute(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM domains") - require.Nil(t, err, fmt.Sprintf("clean domains unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - validRoute := "testRoute" - domain := domains.Domain{ - ID: domainID, - Name: "test", - Route: validRoute, - Tags: []string{"test"}, - Metadata: map[string]any{ - "test": "test", - }, - CreatedBy: userID, - UpdatedBy: userID, - CreatedAt: time.Now().UTC().Truncate(time.Microsecond), - UpdatedAt: time.Now().UTC().Truncate(time.Microsecond), - Status: domains.EnabledStatus, - } - - _, err := repo.SaveDomain(context.Background(), domain) - require.Nil(t, err, fmt.Sprintf("failed to save domain %s", domain.ID)) - - cases := []struct { - desc string - route string - response domains.Domain - err error - }{ - { - desc: "retrieve existing domain", - route: validRoute, - response: domain, - err: nil, - }, - { - desc: "retrieve doamin with invalid route", - route: invalid, - response: domains.Domain{}, - err: repoerr.ErrNotFound, - }, - { - desc: "retrieve with empty domain route", - route: "", - response: domains.Domain{}, - err: repoerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - d, err := repo.RetrieveDomainByRoute(context.Background(), tc.route) - assert.Equal(t, tc.response, d, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.response, d)) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.err, err)) - }) - } -} - -func TestRetrieveAllByIDs(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM domains") - require.Nil(t, err, fmt.Sprintf("clean domains unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - items := []domains.Domain{} - baseTime := time.Now().UTC().Truncate(time.Millisecond) - for i := 0; i < 10; i++ { - domain := domains.Domain{ - ID: testsutil.GenerateUUID(t), - Name: fmt.Sprintf(`"test%d"`, i), - Route: fmt.Sprintf(`"test%d"`, i), - Tags: []string{"test"}, - Metadata: map[string]any{ - "test": "test", - }, - CreatedBy: userID, - UpdatedBy: userID, - Status: domains.EnabledStatus, - CreatedAt: baseTime.Add(time.Duration(i) * time.Millisecond), - } - if i%5 == 0 { - domain.Status = domains.DisabledStatus - domain.Tags = []string{"test", "admin"} - domain.Metadata = map[string]any{ - "test1": "test1", - } - } - _, err := repo.SaveDomain(context.Background(), domain) - require.Nil(t, err, fmt.Sprintf("save domain unexpected error: %s", err)) - items = append(items, domain) - } - - cases := []struct { - desc string - pm domains.Page - response domains.DomainsPage - err error - }{ - { - desc: "retrieve by ids successfully", - pm: domains.Page{ - Offset: 0, - Limit: 10, - IDs: []string{items[1].ID, items[2].ID}, - Order: defOrder, - Dir: ascDir, - }, - response: domains.DomainsPage{ - Total: 2, - Offset: 0, - Limit: 10, - Domains: []domains.Domain{items[1], items[2]}, - }, - err: nil, - }, - { - desc: "retrieve by ids with empty ids", - pm: domains.Page{ - Offset: 0, - Limit: 10, - IDs: []string{}, - }, - response: domains.DomainsPage{ - Total: 0, - Offset: 0, - Limit: 0, - }, - err: nil, - }, - { - desc: "retrieve by ids with invalid ids", - pm: domains.Page{ - Offset: 0, - Limit: 10, - IDs: []string{invalid}, - }, - response: domains.DomainsPage{ - Total: 0, - Offset: 0, - Limit: 10, - }, - err: nil, - }, - { - desc: "retrieve by ids and status", - pm: domains.Page{ - Offset: 0, - Limit: 10, - IDs: []string{items[0].ID, items[1].ID}, - Status: domains.DisabledStatus, - }, - response: domains.DomainsPage{ - Total: 1, - Offset: 0, - Limit: 10, - Domains: []domains.Domain{items[0]}, - }, - }, - { - desc: "retrieve by ids and status with invalid status", - pm: domains.Page{ - Offset: 0, - Limit: 10, - IDs: []string{items[0].ID, items[1].ID}, - Status: 5, - Order: defOrder, - Dir: ascDir, - }, - response: domains.DomainsPage{ - Total: 2, - Offset: 0, - Limit: 10, - Domains: []domains.Domain{items[0], items[1]}, - }, - }, - { - desc: "retrieve by ids and tags", - pm: domains.Page{ - Offset: 0, - Limit: 10, - IDs: []string{items[0].ID, items[1].ID}, - Tags: domains.TagsQuery{Elements: []string{"test"}, Operator: domains.OrOp}, - }, - response: domains.DomainsPage{ - Total: 1, - Offset: 0, - Limit: 10, - Domains: []domains.Domain{items[1]}, - }, - }, - { - desc: "retrieve by ids and metadata", - pm: domains.Page{ - Offset: 0, - Limit: 10, - IDs: []string{items[1].ID, items[2].ID}, - Metadata: map[string]any{ - "test": "test", - }, - Status: domains.EnabledStatus, - Order: defOrder, - Dir: ascDir, - }, - response: domains.DomainsPage{ - Total: 2, - Offset: 0, - Limit: 10, - Domains: items[1:3], - }, - }, - { - desc: "retrieve by ids and metadata with invalid metadata", - pm: domains.Page{ - Offset: 0, - Limit: 10, - IDs: []string{items[1].ID, items[2].ID}, - Metadata: map[string]any{ - "test1": "test1", - }, - Status: domains.EnabledStatus, - }, - response: domains.DomainsPage{ - Total: 0, - Offset: 0, - Limit: 10, - }, - }, - { - desc: "retrieve by ids and malfomed metadata", - pm: domains.Page{ - Offset: 0, - Limit: 10, - IDs: []string{items[1].ID, items[2].ID}, - Metadata: map[string]any{ - "key": make(chan int), - }, - Status: domains.EnabledStatus, - }, - response: domains.DomainsPage{}, - err: repoerr.ErrViewEntity, - }, - { - desc: "retrieve all by ids and id", - pm: domains.Page{ - Offset: 0, - Limit: 10, - ID: items[1].ID, - IDs: []string{items[1].ID, items[2].ID}, - }, - response: domains.DomainsPage{ - Total: 1, - Offset: 0, - Limit: 10, - Domains: []domains.Domain{items[1]}, - }, - }, - { - desc: "retrieve all by ids and id with invalid id", - pm: domains.Page{ - Offset: 0, - Limit: 10, - ID: invalid, - IDs: []string{items[1].ID, items[2].ID}, - }, - response: domains.DomainsPage{ - Total: 0, - Offset: 0, - Limit: 10, - }, - }, - { - desc: "retrieve all by ids and name", - pm: domains.Page{ - Offset: 0, - Limit: 10, - Name: items[1].Name, - IDs: []string{items[1].ID, items[2].ID}, - }, - response: domains.DomainsPage{ - Total: 1, - Offset: 0, - Limit: 10, - Domains: []domains.Domain{items[1]}, - }, - }, - { - desc: "retrieve all by ids with empty page", - pm: domains.Page{}, - response: domains.DomainsPage{}, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - dp, err := repo.RetrieveAllDomainsByIDs(context.Background(), tc.pm) - assert.Equal(t, tc.response, dp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.response, dp)) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.err, err)) - }) - } -} - -func TestUpdate(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM domains") - require.Nil(t, err, fmt.Sprintf("clean domains unexpected error: %s", err)) - }) - - updatedName := "test1" - updatedMetadata := domains.Metadata{ - "test1": "test1", - } - updatedTags := []string{"test1"} - updatedStatus := domains.DisabledStatus - - repo := postgres.NewRepository(database) - - domain := domains.Domain{ - ID: domainID, - Name: "test", - Route: "test", - Tags: []string{"test"}, - Metadata: map[string]any{ - "test": "test", - }, - CreatedBy: userID, - UpdatedBy: userID, - Status: domains.EnabledStatus, - } - - _, err := repo.SaveDomain(context.Background(), domain) - require.Nil(t, err, fmt.Sprintf("failed to save domain %s", domain.ID)) - - cases := []struct { - desc string - domainID string - d domains.DomainReq - response domains.Domain - err error - }{ - { - desc: "update existing domain name and metadata", - domainID: domain.ID, - d: domains.DomainReq{ - Name: &updatedName, - Metadata: &updatedMetadata, - }, - response: domains.Domain{ - ID: domainID, - Name: "test1", - Route: "test", - Tags: []string{"test"}, - Metadata: map[string]any{ - "test1": "test1", - }, - CreatedBy: userID, - UpdatedBy: userID, - Status: domains.EnabledStatus, - UpdatedAt: time.Now(), - }, - err: nil, - }, - { - desc: "update existing domain name, metadata, tags and status", - domainID: domain.ID, - d: domains.DomainReq{ - Name: &updatedName, - Metadata: &updatedMetadata, - Tags: &updatedTags, - Status: &updatedStatus, - }, - response: domains.Domain{ - ID: domainID, - Name: "test1", - Route: "test", - Tags: []string{"test1"}, - Metadata: map[string]any{ - "test1": "test1", - }, - CreatedBy: userID, - UpdatedBy: userID, - Status: domains.DisabledStatus, - UpdatedAt: time.Now(), - }, - err: nil, - }, - { - desc: "update non-existing domain", - domainID: invalid, - d: domains.DomainReq{ - Name: &updatedName, - Metadata: &updatedMetadata, - }, - response: domains.Domain{}, - err: repoerr.ErrNotFound, - }, - { - desc: "update domain with empty ID", - domainID: "", - d: domains.DomainReq{ - Name: &updatedName, - Metadata: &updatedMetadata, - }, - response: domains.Domain{}, - err: repoerr.ErrNotFound, - }, - { - desc: "update domain with malformed metadata", - domainID: domainID, - d: domains.DomainReq{ - Name: &updatedName, - Metadata: &domains.Metadata{"key": make(chan int)}, - }, - response: domains.Domain{}, - err: repoerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - d, err := repo.UpdateDomain(context.Background(), tc.domainID, tc.d) - d.UpdatedAt = tc.response.UpdatedAt - assert.Equal(t, tc.response, d, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.response, d)) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - }) - } -} - -func TestDelete(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM domains") - require.Nil(t, err, fmt.Sprintf("clean domains unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - domain := domains.Domain{ - ID: domainID, - Name: "test", - Route: "test", - Tags: []string{"test"}, - Metadata: map[string]any{ - "test": "test", - }, - CreatedBy: userID, - UpdatedBy: userID, - Status: domains.EnabledStatus, - } - - _, err := repo.SaveDomain(context.Background(), domain) - require.Nil(t, err, fmt.Sprintf("failed to save domain %s", domain.ID)) - - cases := []struct { - desc string - domainID string - err error - }{ - { - desc: "delete existing domain", - domainID: domain.ID, - err: nil, - }, - { - desc: "delete non-existing domain", - domainID: invalid, - err: repoerr.ErrNotFound, - }, - { - desc: "delete domain with empty ID", - domainID: "", - err: repoerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := repo.DeleteDomain(context.Background(), tc.domainID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - }) - } -} - -func TestListDomains(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM domains") - require.Nil(t, err, fmt.Sprintf("clean domains unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - items := []domains.Domain{} - baseTime := time.Now().UTC().Truncate(time.Millisecond) - for i := 0; i < 10; i++ { - domain := domains.Domain{ - ID: testsutil.GenerateUUID(t), - Name: fmt.Sprintf(`"test%d"`, i), - Route: fmt.Sprintf(`"test%d"`, i), - Tags: []string{"tag1", "tag2"}, - Metadata: map[string]any{ - "test": "test", - }, - CreatedBy: userID, - UpdatedBy: userID, - Status: domains.EnabledStatus, - CreatedAt: baseTime.Add(time.Duration(i) * time.Millisecond), - UpdatedAt: baseTime.Add(time.Duration(i) * time.Millisecond), - } - if i%5 == 0 { - domain.Status = domains.DisabledStatus - domain.Metadata = map[string]any{ - "test1": "test1", - } - } - if i%9 == 0 { - domain.Tags = []string{"tag1", "tag3"} - } - _, err := repo.SaveDomain(context.Background(), domain) - require.Nil(t, err, fmt.Sprintf("save domain unexpected error: %s", err)) - items = append(items, domain) - } - - reversedDomains := []domains.Domain{} - for i := len(items) - 1; i >= 0; i-- { - reversedDomains = append(reversedDomains, items[i]) - } - - cases := []struct { - desc string - pm domains.Page - response domains.DomainsPage - err error - }{ - { - desc: "list all domains", - pm: domains.Page{ - Offset: 0, - Limit: 10, - Status: domains.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: domains.DomainsPage{ - Total: 10, - Offset: 0, - Limit: 10, - Domains: items, - }, - err: nil, - }, - { - desc: "list all domains with enabled status", - pm: domains.Page{ - Offset: 0, - Limit: 10, - Status: domains.EnabledStatus, - Order: defOrder, - Dir: ascDir, - }, - response: domains.DomainsPage{ - Total: 8, - Offset: 0, - Limit: 10, - Domains: []domains.Domain{items[1], items[2], items[3], items[4], items[6], items[7], items[8], items[9]}, - }, - err: nil, - }, - { - desc: "list all domains with name", - pm: domains.Page{ - Offset: 0, - Limit: 10, - Name: items[0].Name, - Status: domains.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: domains.DomainsPage{ - Total: 1, - Offset: 0, - Limit: 10, - Domains: []domains.Domain{items[0]}, - }, - err: nil, - }, - { - desc: "list all domains with disabled status", - pm: domains.Page{ - Offset: 0, - Limit: 10, - Status: domains.DisabledStatus, - Order: defOrder, - Dir: ascDir, - }, - response: domains.DomainsPage{ - Total: 2, - Offset: 0, - Limit: 10, - Domains: []domains.Domain{items[0], items[5]}, - }, - err: nil, - }, - { - desc: "list all domains with single tag", - pm: domains.Page{ - Offset: 0, - Limit: 10, - Tags: domains.TagsQuery{Elements: []string{"tag1"}, Operator: domains.OrOp}, - Status: domains.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: domains.DomainsPage{ - Total: 10, - Offset: 0, - Limit: 10, - Domains: items, - }, - err: nil, - }, - { - desc: "list all domain with multiple tags and OR operator", - pm: domains.Page{ - Offset: 0, - Limit: 10, - Tags: domains.TagsQuery{Elements: []string{"tag2", "tag3"}, Operator: domains.OrOp}, - Status: domains.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: domains.DomainsPage{ - Total: 10, - Offset: 0, - Limit: 10, - Domains: items, - }, - err: nil, - }, - { - desc: "retrieve domain with multiple tags and AND operator", - pm: domains.Page{ - Offset: 0, - Limit: 10, - Tags: domains.TagsQuery{Elements: []string{"tag1", "tag3"}, Operator: domains.AndOp}, - Status: domains.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: domains.DomainsPage{ - Total: 2, - Offset: 0, - Limit: 10, - Domains: []domains.Domain{items[0], items[9]}, - }, - }, - { - desc: "retrieve domain with invalid tags", - pm: domains.Page{ - Offset: 0, - Limit: 10, - Tags: domains.TagsQuery{Elements: []string{"invalid-tag"}, Operator: domains.OrOp}, - Status: domains.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: domains.DomainsPage{ - Total: 0, - Offset: 0, - Limit: 10, - Domains: []domains.Domain(nil), - }, - }, - { - desc: "list all domains with metadata", - pm: domains.Page{ - Offset: 0, - Limit: 10, - Metadata: map[string]any{ - "test1": "test1", - }, - Status: domains.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: domains.DomainsPage{ - Total: 2, - Offset: 0, - Limit: 10, - Domains: []domains.Domain{items[0], items[5]}, - }, - err: nil, - }, - { - desc: "list all domains with invalid metadata", - pm: domains.Page{ - Offset: 0, - Limit: 10, - Metadata: map[string]any{ - "key": make(chan int), - }, - Status: domains.AllStatus, - }, - response: domains.DomainsPage{}, - err: repoerr.ErrViewEntity, - }, - { - desc: "list all domains with subject id", - pm: domains.Page{ - Offset: 0, - Limit: 10, - UserID: userID, - Status: domains.AllStatus, - }, - response: domains.DomainsPage{ - Total: 0, - Offset: 0, - Limit: 10, - }, - err: nil, - }, - { - desc: "list domains with id", - pm: domains.Page{ - Offset: 0, - Limit: 10, - ID: items[0].ID, - Status: domains.AllStatus, - }, - response: domains.DomainsPage{ - Total: 1, - Offset: 0, - Limit: 10, - Domains: []domains.Domain{items[0]}, - }, - err: nil, - }, - { - desc: "list domains with invalid id", - pm: domains.Page{ - Offset: 0, - Limit: 10, - ID: invalid, - Status: domains.AllStatus, - }, - response: domains.DomainsPage{ - Total: 0, - Offset: 0, - Limit: 10, - }, - err: nil, - }, - { - desc: "list domains with order by name ascending", - pm: domains.Page{ - Offset: 0, - Limit: 10, - Status: domains.AllStatus, - Order: "name", - Dir: ascDir, - }, - response: domains.DomainsPage{ - Total: 10, - Offset: 0, - Limit: 10, - }, - err: nil, - }, - { - desc: "list domains with order by name descending", - pm: domains.Page{ - Offset: 0, - Limit: 10, - Status: domains.AllStatus, - Order: "name", - Dir: descDir, - }, - response: domains.DomainsPage{ - Total: 10, - Offset: 0, - Limit: 10, - }, - err: nil, - }, - { - desc: "list domains with order by created_at ascending", - pm: domains.Page{ - Offset: 0, - Limit: 10, - Status: domains.AllStatus, - Order: defOrder, - Dir: ascDir, - }, - response: domains.DomainsPage{ - Total: 10, - Offset: 0, - Limit: 10, - Domains: items, - }, - err: nil, - }, - { - desc: "list domains with order by created_at descending", - pm: domains.Page{ - Offset: 0, - Limit: 10, - Status: domains.AllStatus, - Order: defOrder, - Dir: descDir, - }, - response: domains.DomainsPage{ - Total: 10, - Offset: 0, - Limit: 10, - Domains: reversedDomains, - }, - err: nil, - }, - { - desc: "list domains with order by updated_at ascending", - pm: domains.Page{ - Offset: 0, - Limit: 10, - Status: domains.AllStatus, - Order: "updated_at", - Dir: ascDir, - }, - response: domains.DomainsPage{ - Total: 10, - Offset: 0, - Limit: 10, - Domains: items, - }, - err: nil, - }, - { - desc: "list domains with order by updated_at descending", - pm: domains.Page{ - Offset: 0, - Limit: 10, - Status: domains.AllStatus, - Order: "updated_at", - Dir: descDir, - }, - response: domains.DomainsPage{ - Total: 10, - Offset: 0, - Limit: 10, - Domains: reversedDomains, - }, - err: nil, - }, - { - desc: "list domains with created_from filter", - pm: domains.Page{ - Offset: 0, - Limit: 10, - Status: domains.AllStatus, - CreatedFrom: baseTime.Add(5 * time.Millisecond), - Order: "created_at", - Dir: ascDir, - }, - response: domains.DomainsPage{ - Total: 5, - Offset: 0, - Limit: 10, - Domains: []domains.Domain{items[5], items[6], items[7], items[8], items[9]}, - }, - err: nil, - }, - { - desc: "list domains with created_to filter", - pm: domains.Page{ - Offset: 0, - Limit: 10, - Status: domains.AllStatus, - CreatedTo: baseTime.Add(4 * time.Millisecond), - Order: "created_at", - Dir: ascDir, - }, - response: domains.DomainsPage{ - Total: 5, - Offset: 0, - Limit: 10, - Domains: []domains.Domain{items[0], items[1], items[2], items[3], items[4]}, - }, - err: nil, - }, - { - desc: "list domains with both created_from and created_to filters", - pm: domains.Page{ - Offset: 0, - Limit: 10, - Status: domains.AllStatus, - CreatedFrom: baseTime.Add(2 * time.Millisecond), - CreatedTo: baseTime.Add(7 * time.Millisecond), - Order: "created_at", - Dir: ascDir, - }, - response: domains.DomainsPage{ - Total: 6, - Offset: 0, - Limit: 10, - Domains: []domains.Domain{items[2], items[3], items[4], items[5], items[6], items[7]}, - }, - err: nil, - }, - { - desc: "list domains with created_from filter returning no results", - pm: domains.Page{ - Offset: 0, - Limit: 10, - Status: domains.AllStatus, - CreatedFrom: baseTime.Add(20 * time.Millisecond), - Order: "created_at", - Dir: ascDir, - }, - response: domains.DomainsPage{ - Total: 0, - Offset: 0, - Limit: 10, - Domains: []domains.Domain(nil), - }, - err: nil, - }, - { - desc: "list domains with created_to filter returning no results", - pm: domains.Page{ - Offset: 0, - Limit: 10, - Status: domains.AllStatus, - CreatedTo: baseTime.Add(-10 * time.Millisecond), - Order: "created_at", - Dir: ascDir, - }, - response: domains.DomainsPage{ - Total: 0, - Offset: 0, - Limit: 10, - Domains: []domains.Domain(nil), - }, - err: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - dp, err := repo.ListDomains(context.Background(), tc.pm) - assert.Equal(t, tc.response.Total, dp.Total, fmt.Sprintf("%s: expected total %d got %d\n", tc.desc, tc.response.Total, dp.Total)) - assert.Equal(t, tc.response.Offset, dp.Offset, fmt.Sprintf("%s: expected offset %d got %d\n", tc.desc, tc.response.Offset, dp.Offset)) - assert.Equal(t, tc.response.Limit, dp.Limit, fmt.Sprintf("%s: expected limit %d got %d\n", tc.desc, tc.response.Limit, dp.Limit)) - if len(tc.response.Domains) > 0 { - assert.ElementsMatch(t, tc.response.Domains, dp.Domains, fmt.Sprintf("%s: expected domains %v got %v\n", tc.desc, tc.response.Domains, dp.Domains)) - } - verifyDomainsOrdering(t, dp.Domains, tc.pm.Order, tc.pm.Dir) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.err, err)) - }) - } -} - -func verifyDomainsOrdering(t *testing.T, domains []domains.Domain, order, dir string) { - if order == "" || len(domains) <= 1 { - return - } - - for i := 0; i < len(domains)-1; i++ { - switch order { - case "name": - if dir == ascDir { - assert.LessOrEqual(t, domains[i].Name, domains[i+1].Name, fmt.Sprintf("Domains not ordered by name ascending at index %d: %s > %s", i, domains[i].Name, domains[i+1].Name)) - continue - } - assert.GreaterOrEqual(t, domains[i].Name, domains[i+1].Name, fmt.Sprintf("Domains not ordered by name descending at index %d: %s < %s", i, domains[i].Name, domains[i+1].Name)) - case "created_at": - if dir == ascDir { - assert.False(t, domains[i].CreatedAt.After(domains[i+1].CreatedAt), fmt.Sprintf("Domains not ordered by created_at ascending at index %d: %v > %v", i, domains[i].CreatedAt, domains[i+1].CreatedAt)) - continue - } - assert.False(t, domains[i].CreatedAt.Before(domains[i+1].CreatedAt), fmt.Sprintf("Domains not ordered by created_at descending at index %d: %v < %v", i, domains[i].CreatedAt, domains[i+1].CreatedAt)) - case "updated_at": - if dir == ascDir { - assert.False(t, domains[i].UpdatedAt.After(domains[i+1].UpdatedAt), fmt.Sprintf("Domains not ordered by updated_at ascending at index %d: %v > %v", i, domains[i].UpdatedAt, domains[i+1].UpdatedAt)) - continue - } - assert.False(t, domains[i].UpdatedAt.Before(domains[i+1].UpdatedAt), fmt.Sprintf("Domains not ordered by updated_at descending at index %d: %v < %v", i, domains[i].UpdatedAt, domains[i+1].UpdatedAt)) - } - } -} diff --git a/domains/postgres/errors.go b/domains/postgres/errors.go deleted file mode 100644 index 02d94a22f..000000000 --- a/domains/postgres/errors.go +++ /dev/null @@ -1,26 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package postgres - -import "github.com/absmach/magistrala/pkg/errors" - -var _ errors.Mapper = (*duplicateErrors)(nil) - -type duplicateErrors struct{} - -// GetError maps constraint names to known errors. -func (d duplicateErrors) GetError(constraint string) (error, bool) { - switch constraint { - case "domains_route_key": - return errors.ErrRouteNotAvailable, true - case "domains_pkey": - return errors.NewRequestError("domain already exists"), true - default: - return nil, false - } -} - -func NewDuplicateErrors() errors.Mapper { - return duplicateErrors{} -} diff --git a/domains/postgres/init.go b/domains/postgres/init.go deleted file mode 100644 index 0266ba389..000000000 --- a/domains/postgres/init.go +++ /dev/null @@ -1,253 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package postgres - -import ( - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - rolesPostgres "github.com/absmach/magistrala/pkg/roles/repo/postgres" - _ "github.com/jackc/pgx/v5/stdlib" // required for SQL access - migrate "github.com/rubenv/sql-migrate" -) - -// Migration of Domains service. -func Migration() (*migrate.MemoryMigrationSource, error) { - rolesMigration, err := rolesPostgres.Migration(rolesTableNamePrefix, entityTableName, entityIDColumnName) - if err != nil { - return &migrate.MemoryMigrationSource{}, errors.Wrap(repoerr.ErrRoleMigration, err) - } - - domainMigrations := &migrate.MemoryMigrationSource{ - Migrations: []*migrate.Migration{ - { - Id: "domain_1", - Up: []string{ - `CREATE TABLE IF NOT EXISTS domains ( - id VARCHAR(36) PRIMARY KEY, - name VARCHAR(254), - tags TEXT[], - metadata JSONB, - alias VARCHAR(254) NOT NULL UNIQUE, - created_at TIMESTAMP, - updated_at TIMESTAMP, - updated_by VARCHAR(254), - created_by VARCHAR(254), - status SMALLINT NOT NULL DEFAULT 0 CHECK (status >= 0) - );`, - }, - Down: []string{ - `DROP TABLE IF EXISTS domains`, - }, - }, - { - Id: "domain_2", - Up: []string{ - `CREATE TABLE IF NOT EXISTS invitations ( - invited_by VARCHAR(36) NOT NULL, - invitee_user_id VARCHAR(36) NOT NULL, - domain_id VARCHAR(36) NOT NULL, - role_id VARCHAR(36) NOT NULL, - created_at TIMESTAMP NOT NULL, - updated_at TIMESTAMP, - confirmed_at TIMESTAMP, - rejected_at TIMESTAMP, - UNIQUE (invitee_user_id, domain_id), - PRIMARY KEY (invitee_user_id, domain_id), - FOREIGN KEY (domain_id) REFERENCES domains(id) ON DELETE CASCADE - );`, - }, - Down: []string{ - `DROP TABLE IF EXISTS invitations`, - }, - }, - { - Id: "domain_3", - Up: []string{ - `DO $$ - BEGIN - IF EXISTS ( - SELECT 1 - FROM information_schema.columns - WHERE table_schema = current_schema() - AND table_name = 'domains' - AND column_name = 'alias' - ) - AND NOT EXISTS ( - SELECT 1 - FROM information_schema.columns - WHERE table_schema = current_schema() - AND table_name = 'domains' - AND column_name = 'route' - ) THEN - EXECUTE 'ALTER TABLE domains RENAME COLUMN alias TO route;'; - END IF; - END $$;`, - }, - Down: []string{ - `DO $$ - BEGIN - IF EXISTS ( - SELECT 1 - FROM information_schema.columns - WHERE table_schema = current_schema() - AND table_name = 'domains' - AND column_name = 'route' - ) - AND NOT EXISTS ( - SELECT 1 - FROM information_schema.columns - WHERE table_schema = current_schema() - AND table_name = 'domains' - AND column_name = 'alias' - ) THEN - EXECUTE 'ALTER TABLE domains RENAME COLUMN route TO alias;'; - END IF; - END $$;`, - }, - }, - { - Id: "domain_4", - Up: []string{ - `ALTER TABLE domains ALTER COLUMN created_at TYPE TIMESTAMPTZ;`, - `ALTER TABLE domains ALTER COLUMN updated_at TYPE TIMESTAMPTZ;`, - `ALTER TABLE invitations ALTER COLUMN created_at TYPE TIMESTAMPTZ;`, - `ALTER TABLE invitations ALTER COLUMN updated_at TYPE TIMESTAMPTZ;`, - `ALTER TABLE invitations ALTER COLUMN confirmed_at TYPE TIMESTAMPTZ;`, - `ALTER TABLE invitations ALTER COLUMN rejected_at TYPE TIMESTAMPTZ;`, - }, - Down: []string{ - `ALTER TABLE domains ALTER COLUMN created_at TYPE TIMESTAMP;`, - `ALTER TABLE domains ALTER COLUMN updated_at TYPE TIMESTAMP;`, - `ALTER TABLE invitations ALTER COLUMN created_at TYPE TIMESTAMP;`, - `ALTER TABLE invitations ALTER COLUMN updated_at TYPE TIMESTAMP;`, - `ALTER TABLE invitations ALTER COLUMN confirmed_at TYPE TIMESTAMP;`, - `ALTER TABLE invitations ALTER COLUMN rejected_at TYPE TIMESTAMP;`, - }, - }, - { - Id: "domain_5", - Up: []string{ - `DO $$ - BEGIN - IF EXISTS ( - SELECT 1 - FROM pg_constraint c - JOIN pg_class t ON c.conrelid = t.oid - JOIN pg_namespace n ON n.oid = t.relnamespace - WHERE t.relname = 'domains' - AND n.nspname = current_schema() - AND c.conname = 'domains_alias_key' - ) THEN - EXECUTE 'ALTER TABLE domains RENAME CONSTRAINT domains_alias_key TO domains_route_key;'; - END IF; - END $$;`, - }, - Down: []string{ - `DO $$ - BEGIN - IF EXISTS ( - SELECT 1 - FROM pg_constraint c - JOIN pg_class t ON c.conrelid = t.oid - JOIN pg_namespace n ON n.oid = t.relnamespace - WHERE t.relname = 'domains' - AND n.nspname = current_schema() - AND c.conname = 'domains_route_key' - ) THEN - EXECUTE 'ALTER TABLE domains RENAME CONSTRAINT domains_route_key TO domains_alias_key;'; - END IF; - END $$;`, - }, - }, - { - Id: "domain_6", - Up: []string{ - `CREATE INDEX IF NOT EXISTS idx_invitations_invited_by ON invitations(invited_by);`, - `CREATE INDEX IF NOT EXISTS idx_invitations_role_id ON invitations(role_id);`, - }, - Down: []string{ - `DROP INDEX IF EXISTS idx_invitations_invited_by;`, - `DROP INDEX IF EXISTS idx_invitations_role_id;`, - }, - }, - { - Id: "domain_7", - Up: []string{ - `UPDATE domains - SET metadata = (COALESCE(metadata, '{}'::jsonb) || COALESCE(metadata->'ui', '{}'::jsonb)) - 'ui' - WHERE metadata ? 'ui' AND jsonb_typeof(metadata->'ui') = 'object'`, - }, - Down: []string{ - `SELECT 1`, - }, - }, - { - Id: "domains_roles_4", - Up: []string{ - `INSERT INTO domains_role_actions (role_id, action) - SELECT dr.id, a.action - FROM domains_roles dr - CROSS JOIN (VALUES - ('rule_create'), - ('rule_read'), - ('rule_update'), - ('rule_delete'), - ('rule_manage_role'), - ('rule_add_role_users'), - ('rule_remove_role_users'), - ('rule_view_role_users'), - ('alarm_update'), - ('alarm_read'), - ('alarm_delete'), - ('alarm_assign'), - ('alarm_acknowledge'), - ('alarm_resolve'), - ('report_create'), - ('report_read'), - ('report_update'), - ('report_delete'), - ('report_manage_role'), - ('report_add_role_users'), - ('report_remove_role_users'), - ('report_view_role_users') - ) AS a(action) - WHERE dr.name = 'admin' - ON CONFLICT DO NOTHING;`, - }, - Down: []string{ - `DELETE FROM domains_role_actions - WHERE action IN ( - 'rule_create', - 'rule_read', - 'rule_update', - 'rule_delete', - 'rule_manage_role', - 'rule_add_role_users', - 'rule_remove_role_users', - 'rule_view_role_users', - 'alarm_update', - 'alarm_read', - 'alarm_delete', - 'alarm_assign', - 'alarm_acknowledge', - 'alarm_resolve', - 'report_create', - 'report_read', - 'report_update', - 'report_delete', - 'report_manage_role', - 'report_add_role_users', - 'report_remove_role_users', - 'report_view_role_users' - ) - AND role_id IN (SELECT id FROM domains_roles WHERE name = 'admin');`, - }, - }, - }, - } - - domainMigrations.Migrations = append(domainMigrations.Migrations, rolesMigration.Migrations...) - - return domainMigrations, nil -} diff --git a/domains/postgres/invitations.go b/domains/postgres/invitations.go deleted file mode 100644 index 328d9992b..000000000 --- a/domains/postgres/invitations.go +++ /dev/null @@ -1,279 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package postgres - -import ( - "context" - "database/sql" - "fmt" - "strings" - "time" - - "github.com/absmach/magistrala/domains" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - "github.com/absmach/magistrala/pkg/postgres" -) - -func (repo domainRepo) SaveInvitation(ctx context.Context, invitation domains.Invitation) (err error) { - q := `INSERT INTO invitations (invited_by, invitee_user_id, domain_id, role_id, created_at) - VALUES (:invited_by, :invitee_user_id, :domain_id, :role_id, :created_at)` - - dbInv := toDBInvitation(invitation) - if _, err = repo.db.NamedExecContext(ctx, q, dbInv); err != nil { - return postgres.HandleError(repoerr.ErrCreateEntity, err) - } - - return nil -} - -func (repo domainRepo) RetrieveInvitation(ctx context.Context, inviteeUserID, domainID string) (domains.Invitation, error) { - q := `SELECT invited_by, invitee_user_id, domain_id, role_id, created_at, updated_at, confirmed_at, rejected_at FROM invitations WHERE invitee_user_id = :invitee_user_id AND domain_id = :domain_id;` - - dbinv := dbInvitation{ - InviteeUserID: inviteeUserID, - DomainID: domainID, - } - rows, err := repo.db.NamedQueryContext(ctx, q, dbinv) - if err != nil { - return domains.Invitation{}, postgres.HandleError(repoerr.ErrViewEntity, err) - } - defer rows.Close() - - dbinv = dbInvitation{} - if rows.Next() { - if err = rows.StructScan(&dbinv); err != nil { - return domains.Invitation{}, postgres.HandleError(repoerr.ErrViewEntity, err) - } - - return toInvitation(dbinv), nil - } - - return domains.Invitation{}, repoerr.ErrNotFound -} - -func (repo domainRepo) RetrieveAllInvitations(ctx context.Context, pm domains.InvitationPageMeta) (domains.InvitationPage, error) { - query := pageQuery(pm) - - q := fmt.Sprintf(` - SELECT - i.invited_by, - i.invitee_user_id, - i.domain_id, - d."name" AS domain_name, - i.role_id, - dr."name" AS role_name, - i.created_at, - i.updated_at, - i.confirmed_at, - i.rejected_at - FROM - invitations i - LEFT JOIN domains d ON - i.domain_id = d.id - LEFT JOIN domains_roles dr ON - dr.id = i.role_id - %s - LIMIT :limit OFFSET :offset; - `, query) - - var items []domains.Invitation - if !pm.OnlyTotal { - rows, err := repo.db.NamedQueryContext(ctx, q, pm) - if err != nil { - return domains.InvitationPage{}, postgres.HandleError(repoerr.ErrViewEntity, err) - } - defer rows.Close() - - for rows.Next() { - var dbinv dbInvitation - if err = rows.StructScan(&dbinv); err != nil { - return domains.InvitationPage{}, postgres.HandleError(repoerr.ErrViewEntity, err) - } - items = append(items, toInvitation(dbinv)) - } - } - - tq := fmt.Sprintf(` - SELECT - COUNT(*) - FROM - invitations i - LEFT JOIN domains d ON - i.domain_id = d.id - LEFT JOIN domains_roles dr ON - dr.id = i.role_id %s - `, query) - - total, err := postgres.Total(ctx, repo.db, tq, pm) - if err != nil { - return domains.InvitationPage{}, postgres.HandleError(repoerr.ErrViewEntity, err) - } - - invPage := domains.InvitationPage{ - Total: total, - Offset: pm.Offset, - Limit: pm.Limit, - Invitations: items, - } - - return invPage, nil -} - -func (repo domainRepo) UpdateConfirmation(ctx context.Context, invitation domains.Invitation) (err error) { - q := `UPDATE invitations SET confirmed_at = :confirmed_at, updated_at = :updated_at WHERE invitee_user_id = :invitee_user_id AND domain_id = :domain_id` - - dbinv := toDBInvitation(invitation) - result, err := repo.db.NamedExecContext(ctx, q, dbinv) - if err != nil { - return postgres.HandleError(repoerr.ErrUpdateEntity, err) - } - if rows, _ := result.RowsAffected(); rows == 0 { - return repoerr.ErrNotFound - } - - return nil -} - -func (repo domainRepo) UpdateRejection(ctx context.Context, invitation domains.Invitation) (err error) { - q := `UPDATE invitations SET rejected_at = :rejected_at, updated_at = :updated_at WHERE invitee_user_id = :invitee_user_id AND domain_id = :domain_id` - - dbInv := toDBInvitation(invitation) - result, err := repo.db.NamedExecContext(ctx, q, dbInv) - if err != nil { - return postgres.HandleError(repoerr.ErrUpdateEntity, err) - } - if rows, _ := result.RowsAffected(); rows == 0 { - return repoerr.ErrNotFound - } - - return nil -} - -func (repo domainRepo) DeleteUsersInvitations(ctx context.Context, domain string, inviteeUserIDs ...string) (err error) { - if len(inviteeUserIDs) == 0 { - return repoerr.ErrNotFound - } - - q := `DELETE FROM invitations WHERE domain_id = :domain_id AND invitee_user_id = ANY(:invitee_user_ids);` - - params := map[string]any{ - "invitee_user_ids": inviteeUserIDs, - "domain_id": domain, - } - result, err := repo.db.NamedExecContext(ctx, q, params) - if err != nil { - return postgres.HandleError(repoerr.ErrRemoveEntity, err) - } - if rows, _ := result.RowsAffected(); rows == 0 { - return repoerr.ErrNotFound - } - - return nil -} - -func pageQuery(pm domains.InvitationPageMeta) string { - var query []string - var emq string - if pm.DomainID != "" { - query = append(query, "i.domain_id = :domain_id") - } - if pm.InviteeUserID != "" { - query = append(query, "i.invitee_user_id = :invitee_user_id") - } - if pm.InvitedBy != "" { - query = append(query, "i.invited_by = :invited_by") - } - if pm.RoleID != "" { - query = append(query, "i.role_id = :role_id") - } - if pm.InvitedByOrUserID != "" { - query = append(query, "(i.invited_by = :invited_by_or_user_id OR i.invitee_user_id = :invited_by_or_user_id)") - } - if pm.State == domains.Accepted { - query = append(query, "i.confirmed_at IS NOT NULL") - } - if pm.State == domains.Pending { - query = append(query, "i.confirmed_at IS NULL AND rejected_at IS NULL") - } - if pm.State == domains.Rejected { - query = append(query, "i.rejected_at IS NOT NULL") - } - - if len(query) > 0 { - emq = fmt.Sprintf("WHERE %s", strings.Join(query, " AND ")) - } - - return emq -} - -type dbInvitation struct { - InvitedBy string `db:"invited_by"` - InviteeUserID string `db:"invitee_user_id"` - DomainID string `db:"domain_id"` - DomainName sql.NullString `db:"domain_name,omitempty"` - RoleID string `db:"role_id,omitempty"` - RoleName sql.NullString `db:"role_name,omitempty"` - Relation string `db:"relation"` - CreatedAt time.Time `db:"created_at"` - UpdatedAt sql.NullTime `db:"updated_at,omitempty"` - ConfirmedAt sql.NullTime `db:"confirmed_at,omitempty"` - RejectedAt sql.NullTime `db:"rejected_at,omitempty"` -} - -func toDBInvitation(inv domains.Invitation) dbInvitation { - var updatedAt, confirmedAt, rejectedAt sql.NullTime - if inv.UpdatedAt != (time.Time{}) { - updatedAt = sql.NullTime{Time: inv.UpdatedAt, Valid: true} - } - if inv.ConfirmedAt != (time.Time{}) { - confirmedAt = sql.NullTime{Time: inv.ConfirmedAt, Valid: true} - } - if inv.RejectedAt != (time.Time{}) { - rejectedAt = sql.NullTime{Time: inv.RejectedAt, Valid: true} - } - - return dbInvitation{ - InvitedBy: inv.InvitedBy, - InviteeUserID: inv.InviteeUserID, - DomainID: inv.DomainID, - RoleID: inv.RoleID, - CreatedAt: inv.CreatedAt, - UpdatedAt: updatedAt, - ConfirmedAt: confirmedAt, - RejectedAt: rejectedAt, - } -} - -func toInvitation(dbinv dbInvitation) domains.Invitation { - var updatedAt, confirmedAt, rejectedAt time.Time - if dbinv.UpdatedAt.Valid { - updatedAt = dbinv.UpdatedAt.Time - } - if dbinv.ConfirmedAt.Valid { - confirmedAt = dbinv.ConfirmedAt.Time.UTC() - } - if dbinv.RejectedAt.Valid { - rejectedAt = dbinv.RejectedAt.Time.UTC() - } - - return domains.Invitation{ - InvitedBy: dbinv.InvitedBy, - InviteeUserID: dbinv.InviteeUserID, - DomainID: dbinv.DomainID, - DomainName: toString(dbinv.DomainName), - RoleID: dbinv.RoleID, - RoleName: toString(dbinv.RoleName), - CreatedAt: dbinv.CreatedAt.UTC(), - UpdatedAt: updatedAt, - ConfirmedAt: confirmedAt, - RejectedAt: rejectedAt, - } -} - -func toString(s sql.NullString) string { - if s.Valid { - return s.String - } - return "" -} diff --git a/domains/postgres/invitations_test.go b/domains/postgres/invitations_test.go deleted file mode 100644 index 486a3d917..000000000 --- a/domains/postgres/invitations_test.go +++ /dev/null @@ -1,852 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package postgres_test - -import ( - "context" - "fmt" - "strings" - "testing" - "time" - - "github.com/absmach/magistrala/domains" - "github.com/absmach/magistrala/domains/postgres" - "github.com/absmach/magistrala/internal/testsutil" - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" -) - -const roleName = "roleName" - -var invalidUUID = strings.Repeat("a", 37) - -func TestSaveInvitation(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM invitations") - require.Nil(t, err, fmt.Sprintf("clean invitations unexpected error: %s", err)) - _, err = db.Exec("DELETE FROM domains") - require.Nil(t, err, fmt.Sprintf("clean domains unexpected error: %s", err)) - }) - repo := postgres.NewRepository(database) - - dom := saveDomain(t, repo) - userID := testsutil.GenerateUUID(t) - roleID := testsutil.GenerateUUID(t) - - cases := []struct { - desc string - invitation domains.Invitation - err error - }{ - { - desc: "add new invitation successfully", - invitation: domains.Invitation{ - InvitedBy: testsutil.GenerateUUID(t), - InviteeUserID: userID, - DomainID: dom.ID, - RoleID: roleID, - CreatedAt: time.Now(), - }, - err: nil, - }, - { - desc: "add new invitation with an confirmed_at date", - invitation: domains.Invitation{ - InvitedBy: testsutil.GenerateUUID(t), - InviteeUserID: testsutil.GenerateUUID(t), - DomainID: dom.ID, - CreatedAt: time.Now(), - RoleID: roleID, - ConfirmedAt: time.Now(), - }, - err: nil, - }, - { - desc: "add invitation with duplicate invitation", - invitation: domains.Invitation{ - InvitedBy: testsutil.GenerateUUID(t), - InviteeUserID: userID, - DomainID: dom.ID, - RoleID: roleID, - CreatedAt: time.Now(), - }, - err: repoerr.ErrConflict, - }, - { - desc: "add invitation with invalid invitation invited_by", - invitation: domains.Invitation{ - InvitedBy: invalidUUID, - InviteeUserID: testsutil.GenerateUUID(t), - DomainID: dom.ID, - RoleID: roleID, - CreatedAt: time.Now(), - }, - err: repoerr.ErrMalformedEntity, - }, - { - desc: "add invitation with invalid invitation domain", - invitation: domains.Invitation{ - InvitedBy: testsutil.GenerateUUID(t), - InviteeUserID: testsutil.GenerateUUID(t), - DomainID: invalidUUID, - RoleID: roleID, - CreatedAt: time.Now(), - }, - err: repoerr.ErrMalformedEntity, - }, - { - desc: "add invitation with invalid invitation invitee user id", - invitation: domains.Invitation{ - InvitedBy: testsutil.GenerateUUID(t), - InviteeUserID: invalidUUID, - DomainID: testsutil.GenerateUUID(t), - RoleID: roleID, - CreatedAt: time.Now(), - }, - err: repoerr.ErrMalformedEntity, - }, - { - desc: "add invitation with empty invitation domain", - invitation: domains.Invitation{ - InvitedBy: testsutil.GenerateUUID(t), - InviteeUserID: testsutil.GenerateUUID(t), - RoleID: roleID, - CreatedAt: time.Now(), - }, - err: repoerr.ErrCreateEntity, - }, - { - desc: "add invitation with empty invitation invitee user id", - invitation: domains.Invitation{ - InvitedBy: testsutil.GenerateUUID(t), - DomainID: dom.ID, - RoleID: roleID, - CreatedAt: time.Now(), - }, - err: nil, - }, - { - desc: "add invitation with empty invitation invited_by", - invitation: domains.Invitation{ - DomainID: dom.ID, - InviteeUserID: testsutil.GenerateUUID(t), - RoleID: roleID, - CreatedAt: time.Now(), - }, - err: nil, - }, - { - desc: "add invitation with empty invitation role id", - invitation: domains.Invitation{ - InvitedBy: testsutil.GenerateUUID(t), - InviteeUserID: testsutil.GenerateUUID(t), - DomainID: dom.ID, - CreatedAt: time.Now(), - }, - err: nil, - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := repo.SaveInvitation(context.Background(), tc.invitation) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.err, err)) - }) - } -} - -func TestRetrieveInvitation(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM invitations") - require.Nil(t, err, fmt.Sprintf("clean invitations unexpected error: %s", err)) - _, err = db.Exec("DELETE FROM domains") - require.Nil(t, err, fmt.Sprintf("clean domains unexpected error: %s", err)) - }) - repo := postgres.NewRepository(database) - - dom := saveDomain(t, repo) - - invitation := domains.Invitation{ - InvitedBy: testsutil.GenerateUUID(t), - InviteeUserID: testsutil.GenerateUUID(t), - DomainID: dom.ID, - RoleID: testsutil.GenerateUUID(t), - CreatedAt: time.Now().UTC().Truncate(time.Microsecond), - } - - err := repo.SaveInvitation(context.Background(), invitation) - require.Nil(t, err, fmt.Sprintf("create invitation unexpected error: %s", err)) - - cases := []struct { - desc string - userID string - domainID string - response domains.Invitation - err error - }{ - { - desc: "retrieve invitations successfully", - userID: invitation.InviteeUserID, - domainID: invitation.DomainID, - response: invitation, - err: nil, - }, - { - desc: "retrieve invitations with invalid invitee user id", - userID: testsutil.GenerateUUID(t), - domainID: invitation.DomainID, - response: domains.Invitation{}, - err: repoerr.ErrNotFound, - }, - { - desc: "retrieve invitations with invalid invitation domain_id", - userID: invitation.InviteeUserID, - domainID: testsutil.GenerateUUID(t), - response: domains.Invitation{}, - err: repoerr.ErrNotFound, - }, - { - desc: "retrieve invitations with invalid invitee user id and domain_id", - userID: testsutil.GenerateUUID(t), - domainID: testsutil.GenerateUUID(t), - response: domains.Invitation{}, - err: repoerr.ErrNotFound, - }, - { - desc: "retrieve invitations with empty invitee user id", - userID: "", - domainID: invitation.DomainID, - response: domains.Invitation{}, - err: repoerr.ErrNotFound, - }, - { - desc: "retrieve invitations with empty invitation domain_id", - userID: invitation.InviteeUserID, - domainID: "", - response: domains.Invitation{}, - err: repoerr.ErrNotFound, - }, - { - desc: "retrieve invitations with empty invitation user id and domain_id", - userID: "", - domainID: "", - response: domains.Invitation{}, - err: repoerr.ErrNotFound, - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - inv, err := repo.RetrieveInvitation(context.Background(), tc.userID, tc.domainID) - assert.Equal(t, tc.err, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.response, inv, fmt.Sprintf("desc: %s\n", tc.desc)) - }) - } -} - -func TestRetrieveAllInvitations(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM invitations") - require.Nil(t, err, fmt.Sprintf("clean invitations unexpected error: %s", err)) - _, err = db.Exec("DELETE FROM domains") - require.Nil(t, err, fmt.Sprintf("clean domains unexpected error: %s", err)) - }) - repo := postgres.NewRepository(database) - - dom := saveDomain(t, repo) - - num := 200 - - var items []domains.Invitation - for i := 0; i < num; i++ { - invitation := domains.Invitation{ - InvitedBy: testsutil.GenerateUUID(t), - InviteeUserID: testsutil.GenerateUUID(t), - DomainID: dom.ID, - DomainName: dom.Name, - RoleID: testsutil.GenerateUUID(t), - CreatedAt: time.Now().UTC().Truncate(time.Microsecond), - } - err := repo.SaveInvitation(context.Background(), invitation) - require.Nil(t, err, fmt.Sprintf("create invitation unexpected error: %s", err)) - items = append(items, invitation) - } - items[100].ConfirmedAt = time.Now().UTC().Truncate(time.Microsecond) - err := repo.UpdateConfirmation(context.Background(), items[100]) - require.Nil(t, err, fmt.Sprintf("update invitation unexpected error: %s", err)) - - swap := items[100] - items = append(items[:100], items[101:]...) - items = append(items, swap) - - cases := []struct { - desc string - page domains.InvitationPageMeta - response domains.InvitationPage - err error - }{ - { - desc: "retrieve invitations successfully", - page: domains.InvitationPageMeta{ - Offset: 0, - Limit: 10, - }, - response: domains.InvitationPage{ - Total: uint64(num), - Offset: 0, - Limit: 10, - Invitations: items[:10], - }, - err: nil, - }, - { - desc: "retrieve invitations with offset", - page: domains.InvitationPageMeta{ - Offset: 10, - Limit: 10, - }, - response: domains.InvitationPage{ - Total: uint64(num), - Offset: 10, - Limit: 10, - Invitations: items[10:20], - }, - }, - { - desc: "retrieve invitations with limit", - page: domains.InvitationPageMeta{ - Offset: 0, - Limit: 50, - }, - response: domains.InvitationPage{ - Total: uint64(num), - Offset: 0, - Limit: 50, - Invitations: items[:50], - }, - }, - { - desc: "retrieve invitations with offset and limit", - page: domains.InvitationPageMeta{ - Offset: 10, - Limit: 50, - }, - response: domains.InvitationPage{ - Total: uint64(num), - Offset: 10, - Limit: 50, - Invitations: items[10:60], - }, - }, - { - desc: "retrieve invitations with offset out of range", - page: domains.InvitationPageMeta{ - Offset: 1000, - Limit: 50, - }, - response: domains.InvitationPage{ - Total: uint64(num), - Offset: 1000, - Limit: 50, - Invitations: []domains.Invitation(nil), - }, - }, - { - desc: "retrieve invitations with offset and limit out of range", - page: domains.InvitationPageMeta{ - Offset: 170, - Limit: 50, - }, - response: domains.InvitationPage{ - Total: uint64(num), - Offset: 170, - Limit: 50, - Invitations: items[170:200], - }, - }, - { - desc: "retrieve invitations with limit out of range", - page: domains.InvitationPageMeta{ - Offset: 0, - Limit: 1000, - }, - response: domains.InvitationPage{ - Total: uint64(num), - Offset: 0, - Limit: 1000, - Invitations: items, - }, - }, - { - desc: "retrieve invitations with empty page", - page: domains.InvitationPageMeta{}, - response: domains.InvitationPage{ - Total: uint64(num), - Offset: 0, - Limit: 0, - Invitations: []domains.Invitation(nil), - }, - }, - { - desc: "retrieve invitations with domain", - page: domains.InvitationPageMeta{ - DomainID: items[0].DomainID, - Offset: 0, - Limit: 10, - }, - response: domains.InvitationPage{ - Total: uint64(num), - Offset: 0, - Limit: 10, - Invitations: items[:10], - }, - }, - { - desc: "retrieve invitations with invitee user id", - page: domains.InvitationPageMeta{ - InviteeUserID: items[0].InviteeUserID, - Offset: 0, - Limit: 10, - }, - response: domains.InvitationPage{ - Total: 1, - Offset: 0, - Limit: 10, - Invitations: []domains.Invitation{items[0]}, - }, - }, - { - desc: "retrieve invitations with invited_by", - page: domains.InvitationPageMeta{ - InvitedBy: items[0].InvitedBy, - Offset: 0, - Limit: 10, - }, - response: domains.InvitationPage{ - Total: 1, - Offset: 0, - Limit: 10, - Invitations: []domains.Invitation{items[0]}, - }, - }, - { - desc: "retrieve invitations with role_id", - page: domains.InvitationPageMeta{ - RoleID: items[3].RoleID, - Offset: 0, - Limit: 10, - }, - response: domains.InvitationPage{ - Total: 1, - Offset: 0, - Limit: 10, - Invitations: []domains.Invitation{items[3]}, - }, - }, - { - desc: "retrieve invitations with invited_by_or_user_id", - page: domains.InvitationPageMeta{ - InvitedByOrUserID: items[0].InviteeUserID, - Offset: 0, - Limit: 10, - }, - response: domains.InvitationPage{ - Total: 1, - Offset: 0, - Limit: 10, - Invitations: []domains.Invitation{items[0]}, - }, - }, - { - desc: "retrieve invitations with domain_id and invitee user id", - page: domains.InvitationPageMeta{ - DomainID: items[0].DomainID, - InviteeUserID: items[0].InviteeUserID, - Offset: 0, - Limit: 10, - }, - response: domains.InvitationPage{ - Total: 1, - Offset: 0, - Limit: 10, - Invitations: []domains.Invitation{items[0]}, - }, - }, - { - desc: "retrieve invitations with domain_id and invited_by", - page: domains.InvitationPageMeta{ - DomainID: items[0].DomainID, - InvitedBy: items[0].InvitedBy, - Offset: 0, - Limit: 10, - }, - response: domains.InvitationPage{ - Total: 1, - Offset: 0, - Limit: 10, - Invitations: []domains.Invitation{items[0]}, - }, - }, - { - desc: "retrieve invitations with invitee user id and invited_by", - page: domains.InvitationPageMeta{ - InviteeUserID: items[0].InviteeUserID, - InvitedBy: items[0].InvitedBy, - Offset: 0, - Limit: 10, - }, - response: domains.InvitationPage{ - Total: 1, - Offset: 0, - Limit: 10, - Invitations: []domains.Invitation{items[0]}, - }, - }, - { - desc: "retrieve invitations with domain_id, invitee user id and invited_by", - page: domains.InvitationPageMeta{ - DomainID: items[0].DomainID, - InviteeUserID: items[0].InviteeUserID, - InvitedBy: items[0].InvitedBy, - Offset: 0, - Limit: 10, - }, - response: domains.InvitationPage{ - Total: 1, - Offset: 0, - Limit: 10, - Invitations: []domains.Invitation{items[0]}, - }, - }, - { - desc: "retrieve invitations with domain_id, invitee user id, invited_by and role_id", - page: domains.InvitationPageMeta{ - DomainID: items[0].DomainID, - InviteeUserID: items[0].InviteeUserID, - InvitedBy: items[0].InvitedBy, - RoleID: items[0].RoleID, - Offset: 0, - Limit: 10, - }, - response: domains.InvitationPage{ - Total: 1, - Offset: 0, - Limit: 10, - Invitations: []domains.Invitation{items[0]}, - }, - }, - { - desc: "retrieve invitations with invalid domain", - page: domains.InvitationPageMeta{ - DomainID: invalidUUID, - Offset: 0, - Limit: 10, - }, - response: domains.InvitationPage{ - Total: 0, - Offset: 0, - Limit: 10, - Invitations: []domains.Invitation(nil), - }, - }, - { - desc: "retrieve invitations with invalid invitee user id", - page: domains.InvitationPageMeta{ - InviteeUserID: testsutil.GenerateUUID(t), - Offset: 0, - Limit: 10, - }, - response: domains.InvitationPage{ - Total: 0, - Offset: 0, - Limit: 10, - Invitations: []domains.Invitation(nil), - }, - }, - { - desc: "retrieve invitations with invalid invited_by", - page: domains.InvitationPageMeta{ - InvitedBy: invalidUUID, - Offset: 0, - Limit: 10, - }, - response: domains.InvitationPage{ - Total: 0, - Offset: 0, - Limit: 10, - Invitations: []domains.Invitation(nil), - }, - }, - { - desc: "retrieve invitations with invalid role_id", - page: domains.InvitationPageMeta{ - RoleID: invalidUUID, - Offset: 0, - Limit: 10, - }, - response: domains.InvitationPage{ - Total: 0, - Offset: 0, - Limit: 10, - Invitations: []domains.Invitation(nil), - }, - }, - { - desc: "retrieve invitations with accepted state", - page: domains.InvitationPageMeta{ - State: domains.Accepted, - Offset: 0, - Limit: 10, - }, - response: domains.InvitationPage{ - Total: 1, - Offset: 0, - Limit: 10, - Invitations: []domains.Invitation{items[num-1]}, - }, - }, - { - desc: "retrieve invitations with pending state", - page: domains.InvitationPageMeta{ - State: domains.Pending, - Offset: 0, - Limit: 10, - }, - response: domains.InvitationPage{ - Total: uint64(num - 1), - Offset: 0, - Limit: 10, - Invitations: items[0:10], - }, - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - page, err := repo.RetrieveAllInvitations(context.Background(), tc.page) - assert.Equal(t, tc.response.Total, page.Total, fmt.Sprintf("%s: expected %d got %d\n", tc.desc, tc.response.Total, page.Total)) - assert.Equal(t, tc.response.Offset, page.Offset, fmt.Sprintf("%s: expected %d got %d\n", tc.desc, tc.response.Offset, page.Offset)) - assert.Equal(t, tc.response.Limit, page.Limit, fmt.Sprintf("%s: expected %d got %d\n", tc.desc, tc.response.Limit, page.Limit)) - assert.ElementsMatch(t, page.Invitations, tc.response.Invitations, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.response.Invitations, page.Invitations)) - assert.Equal(t, tc.err, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - }) - } -} - -func TestUpdateConfirmation(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM invitations") - require.Nil(t, err, fmt.Sprintf("clean invitations unexpected error: %s", err)) - _, err = db.Exec("DELETE FROM domains") - require.Nil(t, err, fmt.Sprintf("clean domains unexpected error: %s", err)) - }) - repo := postgres.NewRepository(database) - - dom := saveDomain(t, repo) - - invitation := domains.Invitation{ - InvitedBy: testsutil.GenerateUUID(t), - InviteeUserID: testsutil.GenerateUUID(t), - DomainID: dom.ID, - DomainName: dom.Name, - RoleID: testsutil.GenerateUUID(t), - RoleName: roleName, - CreatedAt: time.Now(), - } - err := repo.SaveInvitation(context.Background(), invitation) - require.Nil(t, err, fmt.Sprintf("create invitation unexpected error: %s", err)) - - cases := []struct { - desc string - invitation domains.Invitation - err error - }{ - { - desc: "update invitation successfully", - invitation: domains.Invitation{ - DomainID: invitation.DomainID, - InviteeUserID: invitation.InviteeUserID, - ConfirmedAt: time.Now(), - }, - err: nil, - }, - { - desc: "update invitation with invalid invitee user id", - invitation: domains.Invitation{ - InviteeUserID: testsutil.GenerateUUID(t), - DomainID: invitation.InviteeUserID, - ConfirmedAt: time.Now(), - }, - err: repoerr.ErrNotFound, - }, - { - desc: "update invitation with invalid domain", - invitation: domains.Invitation{ - InviteeUserID: invitation.InviteeUserID, - DomainID: testsutil.GenerateUUID(t), - ConfirmedAt: time.Now(), - }, - err: repoerr.ErrNotFound, - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := repo.UpdateConfirmation(context.Background(), tc.invitation) - assert.Equal(t, tc.err, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - }) - } -} - -func TestUpdateRejection(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM invitations") - require.Nil(t, err, fmt.Sprintf("clean invitations unexpected error: %s", err)) - _, err = db.Exec("DELETE FROM domains") - require.Nil(t, err, fmt.Sprintf("clean domains unexpected error: %s", err)) - }) - repo := postgres.NewRepository(database) - - dom := saveDomain(t, repo) - - invitation := domains.Invitation{ - InvitedBy: testsutil.GenerateUUID(t), - InviteeUserID: testsutil.GenerateUUID(t), - DomainID: dom.ID, - DomainName: dom.Name, - RoleID: testsutil.GenerateUUID(t), - RoleName: roleName, - CreatedAt: time.Now(), - } - err := repo.SaveInvitation(context.Background(), invitation) - require.Nil(t, err, fmt.Sprintf("create invitation unexpected error: %s", err)) - - cases := []struct { - desc string - invitation domains.Invitation - err error - }{ - { - desc: "update invitation successfully", - invitation: domains.Invitation{ - DomainID: invitation.DomainID, - InviteeUserID: invitation.InviteeUserID, - RejectedAt: time.Now(), - }, - err: nil, - }, - { - desc: "update invitation with invalid invitee user id", - invitation: domains.Invitation{ - InviteeUserID: testsutil.GenerateUUID(t), - DomainID: invitation.InviteeUserID, - RejectedAt: time.Now(), - }, - err: repoerr.ErrNotFound, - }, - { - desc: "update invitation with invalid domain", - invitation: domains.Invitation{ - InviteeUserID: invitation.InviteeUserID, - DomainID: testsutil.GenerateUUID(t), - RejectedAt: time.Now(), - }, - err: repoerr.ErrNotFound, - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := repo.UpdateRejection(context.Background(), tc.invitation) - assert.Equal(t, tc.err, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - }) - } -} - -func TestDeleteUsersInvitations(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM invitations") - require.Nil(t, err, fmt.Sprintf("clean invitations unexpected error: %s", err)) - _, err = db.Exec("DELETE FROM domains") - require.Nil(t, err, fmt.Sprintf("clean domains unexpected error: %s", err)) - }) - repo := postgres.NewRepository(database) - - dom := saveDomain(t, repo) - - num := 10 - items := make([]domains.Invitation, 0, num) - - for i := 0; i < num; i++ { - invitation := domains.Invitation{ - InvitedBy: testsutil.GenerateUUID(t), - InviteeUserID: testsutil.GenerateUUID(t), - DomainID: dom.ID, - DomainName: dom.Name, - RoleID: testsutil.GenerateUUID(t), - RoleName: roleName, - CreatedAt: time.Now(), - } - err := repo.SaveInvitation(context.Background(), invitation) - require.Nil(t, err, fmt.Sprintf("create invitation unexpected error: %s", err)) - items = append(items, invitation) - } - - cases := []struct { - desc string - domainID string - userIDs []string - err error - }{ - { - desc: "delete one invitation successfully", - domainID: dom.ID, - userIDs: []string{items[0].InviteeUserID}, - err: nil, - }, - { - desc: "delete multiple invitations successfully", - domainID: dom.ID, - userIDs: []string{items[1].InviteeUserID, items[2].InviteeUserID, items[3].InviteeUserID}, - err: nil, - }, - { - desc: "delete invitation with invalid invitation id", - domainID: dom.ID, - userIDs: []string{testsutil.GenerateUUID(t)}, - err: repoerr.ErrNotFound, - }, - { - desc: "delete invitation with empty user id", - domainID: dom.ID, - userIDs: []string{}, - err: repoerr.ErrNotFound, - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := repo.DeleteUsersInvitations(context.Background(), tc.domainID, tc.userIDs...) - assert.Equal(t, tc.err, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - }) - } -} - -func saveDomain(t *testing.T, repo domains.Repository) domains.Domain { - domain := domains.Domain{ - ID: testsutil.GenerateUUID(t), - Name: "test", - Route: "test", - Tags: []string{"test"}, - Metadata: map[string]any{ - "test": "test", - }, - CreatedBy: userID, - UpdatedBy: userID, - CreatedAt: time.Now().UTC().Truncate(time.Microsecond), - UpdatedAt: time.Now().UTC().Truncate(time.Microsecond), - Status: domains.EnabledStatus, - } - - _, err := repo.SaveDomain(context.Background(), domain) - require.Nil(t, err, fmt.Sprintf("failed to save domain %s", domain.ID)) - - return domain -} diff --git a/domains/postgres/setup_test.go b/domains/postgres/setup_test.go deleted file mode 100644 index c42150724..000000000 --- a/domains/postgres/setup_test.go +++ /dev/null @@ -1,99 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package postgres_test contains tests for PostgreSQL repository -// implementations. -package postgres_test - -import ( - "database/sql" - "fmt" - "log" - "os" - "testing" - "time" - - dpostgres "github.com/absmach/magistrala/domains/postgres" - "github.com/absmach/magistrala/pkg/postgres" - pgclient "github.com/absmach/magistrala/pkg/postgres" - "github.com/jmoiron/sqlx" - dockertest "github.com/ory/dockertest/v3" - "github.com/ory/dockertest/v3/docker" - "go.opentelemetry.io/otel" -) - -var ( - db *sqlx.DB - database postgres.Database - tracer = otel.Tracer("repo_tests") -) - -func TestMain(m *testing.M) { - pool, err := dockertest.NewPool("") - if err != nil { - log.Fatalf("Could not connect to docker: %s", err) - } - - container, err := pool.RunWithOptions(&dockertest.RunOptions{ - Repository: "postgres", - Tag: "16.2-alpine", - Env: []string{ - "POSTGRES_USER=test", - "POSTGRES_PASSWORD=test", - "POSTGRES_DB=test", - "listen_addresses = '*'", - }, - }, func(config *docker.HostConfig) { - config.AutoRemove = true - config.RestartPolicy = docker.RestartPolicy{Name: "no"} - }) - if err != nil { - log.Fatalf("Could not start container: %s", err) - } - - port := container.GetPort("5432/tcp") - - pool.MaxWait = 120 * time.Second - if err := pool.Retry(func() error { - url := fmt.Sprintf("host=localhost port=%s user=test dbname=test password=test sslmode=disable", port) - db, err := sql.Open("pgx", url) - if err != nil { - return err - } - return db.Ping() - }); err != nil { - log.Fatalf("Could not connect to docker: %s", err) - } - - dbConfig := pgclient.Config{ - Host: "localhost", - Port: port, - User: "test", - Pass: "test", - Name: "test", - SSLMode: "disable", - SSLCert: "", - SSLKey: "", - SSLRootCert: "", - } - - dMigration, err := dpostgres.Migration() - if err != nil { - log.Fatalf("Could not apply domains table migration: %v", err) - } - if db, err = pgclient.Setup(dbConfig, *dMigration); err != nil { - log.Fatalf("Could not setup test DB connection: %s", err) - } - - database = postgres.NewDatabase(db, dbConfig, tracer) - - code := m.Run() - - // Defers will not be run when using os.Exit - db.Close() - if err := pool.Purge(container); err != nil { - log.Fatalf("Could not purge container: %s", err) - } - - os.Exit(code) -} diff --git a/domains/private/mocks/service.go b/domains/private/mocks/service.go deleted file mode 100644 index ae2ce340e..000000000 --- a/domains/private/mocks/service.go +++ /dev/null @@ -1,232 +0,0 @@ -// Copyright (c) Abstract Machines - -// SPDX-License-Identifier: Apache-2.0 - -// Code generated by mockery; DO NOT EDIT. -// github.com/vektra/mockery -// template: testify - -package mocks - -import ( - "context" - - "github.com/absmach/magistrala/domains" - mock "github.com/stretchr/testify/mock" -) - -// NewService creates a new instance of Service. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. -// The first argument is typically a *testing.T value. -func NewService(t interface { - mock.TestingT - Cleanup(func()) -}) *Service { - mock := &Service{} - mock.Mock.Test(t) - - t.Cleanup(func() { mock.AssertExpectations(t) }) - - return mock -} - -// Service is an autogenerated mock type for the Service type -type Service struct { - mock.Mock -} - -type Service_Expecter struct { - mock *mock.Mock -} - -func (_m *Service) EXPECT() *Service_Expecter { - return &Service_Expecter{mock: &_m.Mock} -} - -// DeleteUserFromDomains provides a mock function for the type Service -func (_mock *Service) DeleteUserFromDomains(ctx context.Context, id string) error { - ret := _mock.Called(ctx, id) - - if len(ret) == 0 { - panic("no return value specified for DeleteUserFromDomains") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) error); ok { - r0 = returnFunc(ctx, id) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_DeleteUserFromDomains_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'DeleteUserFromDomains' -type Service_DeleteUserFromDomains_Call struct { - *mock.Call -} - -// DeleteUserFromDomains is a helper method to define mock.On call -// - ctx context.Context -// - id string -func (_e *Service_Expecter) DeleteUserFromDomains(ctx interface{}, id interface{}) *Service_DeleteUserFromDomains_Call { - return &Service_DeleteUserFromDomains_Call{Call: _e.mock.On("DeleteUserFromDomains", ctx, id)} -} - -func (_c *Service_DeleteUserFromDomains_Call) Run(run func(ctx context.Context, id string)) *Service_DeleteUserFromDomains_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Service_DeleteUserFromDomains_Call) Return(err error) *Service_DeleteUserFromDomains_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_DeleteUserFromDomains_Call) RunAndReturn(run func(ctx context.Context, id string) error) *Service_DeleteUserFromDomains_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveIDByRoute provides a mock function for the type Service -func (_mock *Service) RetrieveIDByRoute(ctx context.Context, route string) (string, error) { - ret := _mock.Called(ctx, route) - - if len(ret) == 0 { - panic("no return value specified for RetrieveIDByRoute") - } - - var r0 string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) (string, error)); ok { - return returnFunc(ctx, route) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string) string); ok { - r0 = returnFunc(ctx, route) - } else { - r0 = ret.Get(0).(string) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = returnFunc(ctx, route) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RetrieveIDByRoute_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveIDByRoute' -type Service_RetrieveIDByRoute_Call struct { - *mock.Call -} - -// RetrieveIDByRoute is a helper method to define mock.On call -// - ctx context.Context -// - route string -func (_e *Service_Expecter) RetrieveIDByRoute(ctx interface{}, route interface{}) *Service_RetrieveIDByRoute_Call { - return &Service_RetrieveIDByRoute_Call{Call: _e.mock.On("RetrieveIDByRoute", ctx, route)} -} - -func (_c *Service_RetrieveIDByRoute_Call) Run(run func(ctx context.Context, route string)) *Service_RetrieveIDByRoute_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Service_RetrieveIDByRoute_Call) Return(s string, err error) *Service_RetrieveIDByRoute_Call { - _c.Call.Return(s, err) - return _c -} - -func (_c *Service_RetrieveIDByRoute_Call) RunAndReturn(run func(ctx context.Context, route string) (string, error)) *Service_RetrieveIDByRoute_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveStatus provides a mock function for the type Service -func (_mock *Service) RetrieveStatus(ctx context.Context, id string) (domains.Status, error) { - ret := _mock.Called(ctx, id) - - if len(ret) == 0 { - panic("no return value specified for RetrieveStatus") - } - - var r0 domains.Status - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) (domains.Status, error)); ok { - return returnFunc(ctx, id) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string) domains.Status); ok { - r0 = returnFunc(ctx, id) - } else { - r0 = ret.Get(0).(domains.Status) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = returnFunc(ctx, id) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RetrieveStatus_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveStatus' -type Service_RetrieveStatus_Call struct { - *mock.Call -} - -// RetrieveStatus is a helper method to define mock.On call -// - ctx context.Context -// - id string -func (_e *Service_Expecter) RetrieveStatus(ctx interface{}, id interface{}) *Service_RetrieveStatus_Call { - return &Service_RetrieveStatus_Call{Call: _e.mock.On("RetrieveStatus", ctx, id)} -} - -func (_c *Service_RetrieveStatus_Call) Run(run func(ctx context.Context, id string)) *Service_RetrieveStatus_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Service_RetrieveStatus_Call) Return(status domains.Status, err error) *Service_RetrieveStatus_Call { - _c.Call.Return(status, err) - return _c -} - -func (_c *Service_RetrieveStatus_Call) RunAndReturn(run func(ctx context.Context, id string) (domains.Status, error)) *Service_RetrieveStatus_Call { - _c.Call.Return(run) - return _c -} diff --git a/domains/private/service.go b/domains/private/service.go deleted file mode 100644 index 43d5cbd06..000000000 --- a/domains/private/service.go +++ /dev/null @@ -1,87 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package private - -import ( - "context" - - "github.com/absmach/magistrala/domains" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" -) - -const defLimit = 100 - -type Service interface { - RetrieveStatus(ctx context.Context, id string) (domains.Status, error) - DeleteUserFromDomains(ctx context.Context, id string) error - RetrieveIDByRoute(ctx context.Context, route string) (string, error) -} - -var _ Service = (*service)(nil) - -func New(repo domains.Repository, cache domains.Cache) Service { - return service{ - repo: repo, - cache: cache, - } -} - -type service struct { - repo domains.Repository - cache domains.Cache -} - -func (svc service) RetrieveStatus(ctx context.Context, id string) (domains.Status, error) { - status, err := svc.cache.Status(ctx, id) - if err == nil { - return status, nil - } - dom, err := svc.repo.RetrieveDomainByID(ctx, id) - if err != nil { - return domains.AllStatus, errors.Wrap(svcerr.ErrViewEntity, err) - } - status = dom.Status - if err := svc.cache.SaveStatus(ctx, id, status); err != nil { - return domains.AllStatus, errors.Wrap(svcerr.ErrUpdateEntity, err) - } - - return dom.Status, nil -} - -func (svc service) DeleteUserFromDomains(ctx context.Context, id string) (err error) { - domainsPage, err := svc.repo.ListDomains(ctx, domains.Page{UserID: id, Limit: defLimit}) - if err != nil { - return err - } - - if domainsPage.Total > defLimit { - for i := defLimit; i < int(domainsPage.Total); i += defLimit { - page := domains.Page{UserID: id, Offset: uint64(i), Limit: defLimit} - dp, err := svc.repo.ListDomains(ctx, page) - if err != nil { - return err - } - domainsPage.Domains = append(domainsPage.Domains, dp.Domains...) - } - } - - return nil -} - -func (svc service) RetrieveIDByRoute(ctx context.Context, route string) (string, error) { - id, err := svc.cache.ID(ctx, route) - if err == nil { - return id, nil - } - dom, err := svc.repo.RetrieveDomainByRoute(ctx, route) - if err != nil { - return "", errors.Wrap(svcerr.ErrViewEntity, err) - } - if err := svc.cache.SaveID(ctx, route, dom.ID); err != nil { - return "", errors.Wrap(svcerr.ErrUpdateEntity, err) - } - - return dom.ID, nil -} diff --git a/domains/service.go b/domains/service.go deleted file mode 100644 index 2b6cbe854..000000000 --- a/domains/service.go +++ /dev/null @@ -1,424 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package domains - -import ( - "context" - "time" - - "github.com/absmach/magistrala" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/pkg/policies" - "github.com/absmach/magistrala/pkg/roles" -) - -var ( - errCreateDomainPolicy = errors.New("failed to create domain policy") - errRollbackRepo = errors.New("failed to rollback repo") -) - -type service struct { - repo Repository - cache Cache - policy policies.Service - idProvider magistrala.IDProvider - roles.ProvisionManageService -} - -var _ Service = (*service)(nil) - -func New(repo Repository, cache Cache, policy policies.Service, idProvider magistrala.IDProvider, sidProvider magistrala.IDProvider, availableActions []roles.Action, builtInRoles map[roles.BuiltInRoleName][]roles.Action) (Service, error) { - rpms, err := roles.NewProvisionManageService(policies.DomainType, repo, policy, sidProvider, availableActions, builtInRoles) - if err != nil { - return nil, err - } - - return &service{ - repo: repo, - cache: cache, - policy: policy, - idProvider: idProvider, - ProvisionManageService: rpms, - }, nil -} - -func (svc service) CreateDomain(ctx context.Context, session authn.Session, d Domain) (retDo Domain, retRps []roles.RoleProvision, retErr error) { - d.CreatedBy = session.UserID - - if d.ID == "" { - domainID, err := svc.idProvider.ID() - if err != nil { - return Domain{}, []roles.RoleProvision{}, errors.Wrap(svcerr.ErrCreateEntity, err) - } - d.ID = domainID - } - - if d.Status != DisabledStatus && d.Status != EnabledStatus { - return Domain{}, []roles.RoleProvision{}, svcerr.ErrInvalidStatus - } - - d.CreatedAt = time.Now().UTC() - - // Domain is created in repo first, because Roles table have foreign key relation with Domain ID - dom, err := svc.repo.SaveDomain(ctx, d) - if err != nil { - if errors.Contains(err, errors.ErrRouteNotAvailable) { - return Domain{}, []roles.RoleProvision{}, errors.ErrRouteNotAvailable - } - return Domain{}, []roles.RoleProvision{}, errors.Wrap(svcerr.ErrCreateEntity, err) - } - defer func() { - if retErr != nil { - if errRollBack := svc.repo.DeleteDomain(ctx, d.ID); errRollBack != nil { - retErr = errors.Wrap(retErr, errors.Wrap(errRollbackRepo, errRollBack)) - } - } - }() - - newBuiltInRoleMembers := map[roles.BuiltInRoleName][]roles.Member{ - BuiltInRoleAdmin: {roles.Member(session.UserID)}, - } - - optionalPolicies := []policies.Policy{ - { - Subject: policies.MagistralaObject, - SubjectType: policies.PlatformType, - Relation: "organization", - Object: d.ID, - ObjectType: policies.DomainType, - }, - } - - rps, err := svc.AddNewEntitiesRoles(ctx, d.ID, session.UserID, []string{d.ID}, optionalPolicies, newBuiltInRoleMembers) - if err != nil { - return Domain{}, []roles.RoleProvision{}, errors.Wrap(errCreateDomainPolicy, err) - } - - return dom, rps, nil -} - -func (svc service) RetrieveDomain(ctx context.Context, session authn.Session, id string, withRoles bool) (Domain, error) { - var domain Domain - var err error - switch session.SuperAdmin { - case true: - domain, err = svc.repo.RetrieveDomainByID(ctx, id) - default: - switch withRoles { - case true: - domain, err = svc.repo.RetrieveDomainByIDWithRoles(ctx, id, session.UserID) - default: - domain, err = svc.repo.RetrieveDomainByID(ctx, id) - } - } - if err != nil { - return Domain{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - return domain, nil -} - -func (svc service) UpdateDomain(ctx context.Context, session authn.Session, id string, d DomainReq) (Domain, error) { - updatedAt := time.Now().UTC() - d.UpdatedAt = &updatedAt - d.UpdatedBy = &session.UserID - dom, err := svc.repo.UpdateDomain(ctx, id, d) - if err != nil { - return Domain{}, errors.Wrap(svcerr.ErrUpdateEntity, err) - } - return dom, nil -} - -func (svc service) EnableDomain(ctx context.Context, session authn.Session, id string) (Domain, error) { - status := EnabledStatus - updatedAt := time.Now().UTC() - dom, err := svc.repo.UpdateDomain(ctx, id, DomainReq{Status: &status, UpdatedBy: &session.UserID, UpdatedAt: &updatedAt}) - if err != nil { - return Domain{}, errors.Wrap(svcerr.ErrUpdateEntity, err) - } - if err := svc.cache.RemoveStatus(ctx, id); err != nil { - return dom, errors.Wrap(svcerr.ErrRemoveEntity, err) - } - - return dom, nil -} - -func (svc service) DisableDomain(ctx context.Context, session authn.Session, id string) (Domain, error) { - status := DisabledStatus - updatedAt := time.Now().UTC() - dom, err := svc.repo.UpdateDomain(ctx, id, DomainReq{Status: &status, UpdatedBy: &session.UserID, UpdatedAt: &updatedAt}) - if err != nil { - return Domain{}, errors.Wrap(svcerr.ErrUpdateEntity, err) - } - if err := svc.cache.RemoveStatus(ctx, id); err != nil { - return dom, errors.Wrap(svcerr.ErrRemoveEntity, err) - } - - return dom, nil -} - -// Only SuperAdmin can freeze the domain. -func (svc service) FreezeDomain(ctx context.Context, session authn.Session, id string) (Domain, error) { - status := FreezeStatus - updatedAt := time.Now().UTC() - dom, err := svc.repo.UpdateDomain(ctx, id, DomainReq{Status: &status, UpdatedBy: &session.UserID, UpdatedAt: &updatedAt}) - if err != nil { - return Domain{}, errors.Wrap(svcerr.ErrUpdateEntity, err) - } - if err := svc.cache.RemoveStatus(ctx, id); err != nil { - return dom, errors.Wrap(svcerr.ErrRemoveEntity, err) - } - - return dom, nil -} - -func (svc service) ListDomains(ctx context.Context, session authn.Session, p Page) (DomainsPage, error) { - p.UserID = session.UserID - if session.SuperAdmin { - p.UserID = "" - } - - dp, err := svc.repo.ListDomains(ctx, p) - if err != nil { - return DomainsPage{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - return dp, nil -} - -func (svc *service) SendInvitation(ctx context.Context, session authn.Session, invitation Invitation) (Invitation, error) { - role, err := svc.repo.RetrieveRole(ctx, invitation.RoleID) - if err != nil { - return Invitation{}, errors.Wrap(svcerr.ErrInvalidRole, err) - } - invitation.RoleName = role.Name - - // Retrieve domain to get domain name - domain, err := svc.repo.RetrieveDomainByID(ctx, invitation.DomainID) - if err != nil { - return Invitation{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - invitation.DomainName = domain.Name - - invitation.InvitedBy = session.UserID - invitation.CreatedAt = time.Now().UTC() - - if invitation.Resend { - if err := svc.resendInvitation(ctx, invitation); err != nil { - return Invitation{}, err - } - return invitation, nil - } - - if err := svc.repo.SaveInvitation(ctx, invitation); err != nil { - return Invitation{}, errors.Wrap(svcerr.ErrCreateEntity, err) - } - return invitation, nil -} - -func (svc *service) resendInvitation(ctx context.Context, invitation Invitation) error { - inv, err := svc.repo.RetrieveInvitation(ctx, invitation.InviteeUserID, invitation.DomainID) - if err != nil { - return errors.Wrap(svcerr.ErrViewEntity, err) - } - if !inv.ConfirmedAt.IsZero() { - return svcerr.ErrInvitationAlreadyAccepted - } - if !inv.RejectedAt.IsZero() { - invitation.RejectedAt = time.Time{} - invitation.UpdatedAt = time.Now().UTC() - if err := svc.repo.UpdateRejection(ctx, invitation); err != nil { - return errors.Wrap(svcerr.ErrUpdateEntity, err) - } - } - - return nil -} - -func (svc *service) ListInvitations(ctx context.Context, session authn.Session, page InvitationPageMeta) (invitations InvitationPage, err error) { - page.InviteeUserID = session.UserID - ip, err := svc.repo.RetrieveAllInvitations(ctx, page) - if err != nil { - return InvitationPage{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - return ip, nil -} - -func (svc *service) ListDomainInvitations(ctx context.Context, session authn.Session, page InvitationPageMeta) (invitations InvitationPage, err error) { - page.DomainID = session.DomainID - ip, err := svc.repo.RetrieveAllInvitations(ctx, page) - if err != nil { - return InvitationPage{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - return ip, nil -} - -func (svc *service) AcceptInvitation(ctx context.Context, session authn.Session, domainID string) (invitation Invitation, err error) { - inv, err := svc.repo.RetrieveInvitation(ctx, session.UserID, domainID) - if err != nil { - return Invitation{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - - if inv.InviteeUserID != session.UserID { - return Invitation{}, svcerr.ErrAuthorization - } - - if !inv.ConfirmedAt.IsZero() { - return Invitation{}, svcerr.ErrInvitationAlreadyAccepted - } - - if !inv.RejectedAt.IsZero() { - return Invitation{}, svcerr.ErrInvitationAlreadyRejected - } - - inv, err = svc.populateDetails(ctx, inv, domainID) - if err != nil { - return Invitation{}, err - } - - session.DomainID = domainID - - if _, err := svc.RoleAddMembers(ctx, session, domainID, inv.RoleID, []string{session.UserID}); err != nil { - return Invitation{}, errors.Wrap(svcerr.ErrUpdateEntity, err) - } - - inv.ConfirmedAt = time.Now().UTC() - inv.UpdatedAt = inv.ConfirmedAt - - if err := svc.repo.UpdateConfirmation(ctx, inv); err != nil { - return Invitation{}, errors.Wrap(svcerr.ErrUpdateEntity, err) - } - - return inv, nil -} - -func (svc *service) RejectInvitation(ctx context.Context, session authn.Session, domainID string) (Invitation, error) { - inv, err := svc.repo.RetrieveInvitation(ctx, session.UserID, domainID) - if err != nil { - return Invitation{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - - if inv.InviteeUserID != session.UserID { - return Invitation{}, svcerr.ErrAuthorization - } - - if !inv.ConfirmedAt.IsZero() { - return Invitation{}, svcerr.ErrInvitationAlreadyAccepted - } - - if !inv.RejectedAt.IsZero() { - return Invitation{}, svcerr.ErrInvitationAlreadyRejected - } - - inv, err = svc.populateDetails(ctx, inv, domainID) - if err != nil { - return Invitation{}, err - } - - inv.RejectedAt = time.Now().UTC() - inv.UpdatedAt = inv.RejectedAt - - if err := svc.repo.UpdateRejection(ctx, inv); err != nil { - return Invitation{}, errors.Wrap(svcerr.ErrUpdateEntity, err) - } - - return inv, nil -} - -func (svc *service) DeleteInvitation(ctx context.Context, session authn.Session, inviteeUserID, domainID string) error { - if session.UserID == inviteeUserID { - if err := svc.repo.DeleteUsersInvitations(ctx, domainID, inviteeUserID); err != nil { - return errors.Wrap(svcerr.ErrRemoveEntity, err) - } - return nil - } - - inv, err := svc.repo.RetrieveInvitation(ctx, inviteeUserID, domainID) - if err != nil { - return errors.Wrap(svcerr.ErrRemoveEntity, err) - } - - if !inv.ConfirmedAt.IsZero() { - return errors.Wrap(svcerr.ErrRemoveEntity, svcerr.ErrInvitationAlreadyAccepted) - } - - if !inv.RejectedAt.IsZero() { - return errors.Wrap(svcerr.ErrRemoveEntity, svcerr.ErrInvitationAlreadyRejected) - } - - if err := svc.repo.DeleteUsersInvitations(ctx, domainID, inviteeUserID); err != nil { - return errors.Wrap(svcerr.ErrRemoveEntity, err) - } - - return nil -} - -// Add domain and role names for an invitation if they are not already set. -func (svc *service) populateDetails(ctx context.Context, inv Invitation, domainID string) (Invitation, error) { - // Populate domain name if not already set - if inv.DomainName == "" { - domain, err := svc.repo.RetrieveDomainByID(ctx, domainID) - if err != nil { - return Invitation{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - inv.DomainName = domain.Name - } - - // Populate role name if not already set - if inv.RoleName == "" { - role, err := svc.repo.RetrieveRole(ctx, inv.RoleID) - if err != nil { - return Invitation{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - inv.RoleName = role.Name - } - - return inv, nil -} - -// Add addition removal of user from invitations. -func (svc *service) RemoveEntityMembers(ctx context.Context, session authn.Session, entityID string, members []string) error { - if err := svc.repo.DeleteUsersInvitations(ctx, entityID, members...); err != nil && err != repoerr.ErrNotFound { - return err - } - - return svc.ProvisionManageService.RemoveEntityMembers(ctx, session, entityID, members) -} - -func (svc *service) RoleRemoveMembers(ctx context.Context, session authn.Session, entityID, roleID string, members []string) error { - ro, err := svc.repo.RetrieveEntityRole(ctx, entityID, roleID) - if err != nil { - return errors.Wrap(svcerr.ErrViewEntity, err) - } - - if _, err := svc.ProvisionManageService.BuiltInRoleActions(roles.BuiltInRoleName(ro.Name)); err == nil { - membersPage, err := svc.repo.RoleListMembers(ctx, ro.ID, 0, 0) - if err != nil { - return errors.Wrap(svcerr.ErrViewEntity, err) - } - if membersPage.Total <= uint64(len(members)) { - return svcerr.ErrRetainOneMember - } - } - - if err := svc.repo.DeleteUsersInvitations(ctx, entityID, members...); err != nil && err != repoerr.ErrNotFound { - return err - } - - return svc.ProvisionManageService.RoleRemoveMembers(ctx, session, entityID, roleID, members) -} - -func (svc *service) RoleRemoveAllMembers(ctx context.Context, session authn.Session, entityID, roleID string) error { - ro, err := svc.repo.RetrieveEntityRole(ctx, entityID, roleID) - if err != nil { - return errors.Wrap(svcerr.ErrViewEntity, err) - } - - if _, err := svc.ProvisionManageService.BuiltInRoleActions(roles.BuiltInRoleName(ro.Name)); err == nil { - return svcerr.ErrRetainOneMember - } - - return svc.ProvisionManageService.RoleRemoveAllMembers(ctx, session, entityID, roleID) -} diff --git a/domains/service_test.go b/domains/service_test.go deleted file mode 100644 index 56da61f6b..000000000 --- a/domains/service_test.go +++ /dev/null @@ -1,1143 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package domains_test - -import ( - "context" - "fmt" - "testing" - "time" - - "github.com/absmach/magistrala/domains" - "github.com/absmach/magistrala/domains/mocks" - "github.com/absmach/magistrala/groups" - "github.com/absmach/magistrala/internal/testsutil" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - policiesMocks "github.com/absmach/magistrala/pkg/policies/mocks" - "github.com/absmach/magistrala/pkg/roles" - "github.com/absmach/magistrala/pkg/sid" - "github.com/absmach/magistrala/pkg/uuid" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" -) - -const ( - groupName = "smqx" - validID = "d4ebb847-5d0e-4e46-bdd9-b6aceaaa3a22" -) - -var ( - ErrExpiry = errors.New("session is expired") - errAddPolicies = errors.New("failed to add policies") - inValid = "invalid" - valid = "valid" - domain = domains.Domain{ - ID: validID, - Name: groupName, - Tags: []string{"tag1", "tag2"}, - Route: "test", - RoleID: "test_role_id", - CreatedBy: validID, - UpdatedBy: validID, - } - validRoles = []roles.MemberRoleActions{ - { - RoleID: "domain_role_id", - RoleName: "domain_role_name", - Actions: []string{"read", "delete"}, - AccessType: "direct", - }, - } - domainWithRoles = domains.Domain{ - ID: validID, - Name: groupName, - Tags: []string{"tag1", "tag2"}, - Route: "test", - RoleID: "test_role_id", - CreatedBy: validID, - UpdatedBy: validID, - Roles: validRoles, - } - userID = testsutil.GenerateUUID(&testing.T{}) - validSession = authn.Session{UserID: userID} - validInvitation = domains.Invitation{ - InviteeUserID: testsutil.GenerateUUID(&testing.T{}), - DomainID: testsutil.GenerateUUID(&testing.T{}), - RoleID: testsutil.GenerateUUID(&testing.T{}), - } -) - -var ( - drepo *mocks.Repository - dcache *mocks.Cache - policy *policiesMocks.Service -) - -func newService() domains.Service { - drepo = new(mocks.Repository) - dcache = new(mocks.Cache) - idProvider := uuid.NewMock() - sidProvider := sid.NewMock() - policy = new(policiesMocks.Service) - availableActions := []roles.Action{} - builtInRoles := map[roles.BuiltInRoleName][]roles.Action{ - groups.BuiltInRoleAdmin: availableActions, - } - ds, _ := domains.New(drepo, dcache, policy, idProvider, sidProvider, availableActions, builtInRoles) - return ds -} - -func TestCreateDomain(t *testing.T) { - svc := newService() - - cases := []struct { - desc string - d domains.Domain - session authn.Session - userID string - addPoliciesErr error - addRolesErr error - saveDomainErr error - deleteDomainErr error - deletePoliciesErr error - err error - }{ - { - desc: "create domain successfully", - d: domains.Domain{ - Name: groupName, - Status: domains.EnabledStatus, - }, - session: validSession, - err: nil, - }, - { - desc: "create domain with custom id", - d: domains.Domain{ - ID: validID, - Name: groupName, - Status: domains.EnabledStatus, - }, - session: validSession, - err: nil, - }, - { - desc: "create domain with invalid status", - d: domains.Domain{ - Name: groupName, - Status: domains.AllStatus, - }, - session: validSession, - err: svcerr.ErrInvalidStatus, - }, - { - desc: "create domain with failed to save domain", - d: domains.Domain{ - Name: groupName, - Status: domains.EnabledStatus, - }, - session: validSession, - saveDomainErr: svcerr.ErrCreateEntity, - err: svcerr.ErrCreateEntity, - }, - { - desc: "create domain with failed to add policies", - d: domains.Domain{ - Name: groupName, - Status: domains.EnabledStatus, - }, - session: validSession, - addPoliciesErr: errAddPolicies, - err: errAddPolicies, - }, - { - desc: "create domain with failed to add policies and failed rollback", - d: domains.Domain{ - Name: groupName, - Status: domains.EnabledStatus, - }, - session: validSession, - addPoliciesErr: errAddPolicies, - deleteDomainErr: svcerr.ErrRemoveEntity, - err: svcerr.ErrRemoveEntity, - }, - { - desc: "create domain with failed to add roles", - d: domains.Domain{ - Name: groupName, - Status: domains.EnabledStatus, - }, - session: validSession, - addRolesErr: errors.ErrMalformedEntity, - err: errors.ErrMalformedEntity, - }, - { - desc: "create domain with failed to add roles and failed rollback", - d: domains.Domain{ - Name: groupName, - Status: domains.EnabledStatus, - }, - session: validSession, - addRolesErr: errors.ErrMalformedEntity, - deleteDomainErr: errors.ErrMalformedEntity, - err: errors.ErrMalformedEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := drepo.On("SaveDomain", mock.Anything, mock.Anything).Return(tc.d, tc.saveDomainErr) - repoCall1 := drepo.On("DeleteDomain", mock.Anything, mock.Anything).Return(tc.deleteDomainErr) - repoCall2 := drepo.On("AddRoles", mock.Anything, mock.Anything).Return([]roles.RoleProvision{}, tc.addRolesErr) - policyCall := policy.On("AddPolicies", mock.Anything, mock.Anything).Return(tc.addPoliciesErr) - policyCall1 := policy.On("DeletePolicies", mock.Anything, mock.Anything).Return(tc.deletePoliciesErr) - _, _, err := svc.CreateDomain(context.Background(), tc.session, tc.d) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.err, err)) - repoCall.Unset() - repoCall1.Unset() - repoCall2.Unset() - policyCall.Unset() - policyCall1.Unset() - }) - } -} - -func TestRetrieveDomain(t *testing.T) { - svc := newService() - - superAdminSession := validSession - superAdminSession.SuperAdmin = true - - cases := []struct { - desc string - session authn.Session - domainID string - withRoles bool - retrieveDomainRes domains.Domain - retrieveDomainErr error - err error - }{ - { - desc: "retrieve domain successfully as super admin", - session: superAdminSession, - withRoles: false, - domainID: validID, - retrieveDomainRes: domain, - err: nil, - }, - { - desc: "retrieve domain successfully as non super admin", - session: validSession, - withRoles: false, - domainID: validID, - retrieveDomainRes: domain, - err: nil, - }, - { - desc: "retrieve domain successfully as non super admin with roles", - session: validSession, - withRoles: true, - domainID: validID, - retrieveDomainRes: domainWithRoles, - err: nil, - }, - { - desc: "retrieve domain with empty domain id", - session: validSession, - withRoles: false, - domainID: "", - retrieveDomainErr: repoerr.ErrNotFound, - err: svcerr.ErrViewEntity, - }, - { - desc: "retrieve non-existing domain", - session: validSession, - withRoles: false, - domainID: inValid, - retrieveDomainErr: repoerr.ErrNotFound, - err: svcerr.ErrViewEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := drepo.On("RetrieveDomainByID", context.Background(), tc.domainID).Return(tc.retrieveDomainRes, tc.retrieveDomainErr) - repoCall1 := drepo.On("RetrieveDomainByIDWithRoles", context.Background(), tc.domainID, tc.session.UserID).Return(tc.retrieveDomainRes, tc.retrieveDomainErr) - domain, err := svc.RetrieveDomain(context.Background(), tc.session, tc.domainID, tc.withRoles) - assert.True(t, errors.Contains(err, tc.err)) - - switch tc.withRoles { - case true: - assert.Equal(t, domain.Roles, validRoles, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, validRoles, domain.Roles)) - ok := drepo.AssertCalled(t, "RetrieveDomainByIDWithRoles", context.Background(), tc.domainID, tc.session.UserID) - assert.True(t, ok, fmt.Sprintf("RetrieveDomainByIDWithRoles was not called on %s", tc.desc)) - default: - assert.Empty(t, domain.Roles) - ok := drepo.AssertCalled(t, "RetrieveDomainByID", context.Background(), tc.domainID) - assert.True(t, ok, fmt.Sprintf("RetrieveDomainByID was not called on %s", tc.desc)) - } - - assert.Equal(t, tc.retrieveDomainRes, domain) - repoCall.Unset() - repoCall1.Unset() - }) - } -} - -func TestUpdateDomain(t *testing.T) { - svc := newService() - - updatedDomain := domain - updatedDomain.Name = valid - updatedDomain.Route = valid - - cases := []struct { - desc string - session authn.Session - domainID string - updateReq domains.DomainReq - updateRes domains.Domain - updateErr error - err error - }{ - { - desc: "update domain successfully", - session: validSession, - domainID: domain.ID, - updateReq: domains.DomainReq{ - Name: &valid, - }, - updateRes: updatedDomain, - err: nil, - }, - { - desc: "update domain with empty domainID", - session: validSession, - domainID: "", - updateReq: domains.DomainReq{ - Name: &valid, - }, - updateErr: repoerr.ErrNotFound, - err: svcerr.ErrUpdateEntity, - }, - { - desc: "update domain with failed to update", - session: validSession, - domainID: domain.ID, - updateReq: domains.DomainReq{ - Name: &valid, - }, - updateErr: errors.ErrMalformedEntity, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := drepo.On("UpdateDomain", context.Background(), tc.domainID, mock.Anything).Return(tc.updateRes, tc.updateErr) - domain, err := svc.UpdateDomain(context.Background(), tc.session, tc.domainID, tc.updateReq) - assert.True(t, errors.Contains(err, tc.err)) - assert.Equal(t, tc.updateRes, domain) - repoCall.Unset() - }) - } -} - -func TestEnableDomain(t *testing.T) { - svc := newService() - - enabledDomain := domain - enabledDomain.Status = domains.EnabledStatus - - cases := []struct { - desc string - session authn.Session - domainID string - enableRes domains.Domain - enableErr error - cacheErr error - resp domains.Domain - err error - }{ - { - desc: "enable domain successfully", - session: validSession, - domainID: domain.ID, - enableRes: enabledDomain, - resp: enabledDomain, - err: nil, - }, - { - desc: "enable domain with empty domainID", - session: validSession, - domainID: "", - enableErr: repoerr.ErrNotFound, - resp: domains.Domain{}, - err: svcerr.ErrUpdateEntity, - }, - { - desc: "enable domain with failed to enable", - session: validSession, - domainID: domain.ID, - enableErr: errors.ErrMalformedEntity, - resp: domains.Domain{}, - err: svcerr.ErrUpdateEntity, - }, - { - desc: "enable domain with failed to remove cache", - session: validSession, - domainID: domain.ID, - enableRes: enabledDomain, - cacheErr: errors.ErrMalformedEntity, - resp: enabledDomain, - err: svcerr.ErrRemoveEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := drepo.On("UpdateDomain", context.Background(), tc.domainID, mock.Anything).Return(tc.enableRes, tc.enableErr) - cacheCall := dcache.On("RemoveStatus", context.Background(), tc.domainID).Return(tc.cacheErr) - domain, err := svc.EnableDomain(context.Background(), tc.session, tc.domainID) - assert.True(t, errors.Contains(err, tc.err)) - assert.Equal(t, tc.resp, domain) - repoCall.Unset() - cacheCall.Unset() - }) - } -} - -func TestDisableDomain(t *testing.T) { - svc := newService() - - disabledDomain := domain - disabledDomain.Status = domains.DisabledStatus - - cases := []struct { - desc string - session authn.Session - domainID string - disableRes domains.Domain - disableErr error - cacheErr error - resp domains.Domain - err error - }{ - { - desc: "disable domain successfully", - session: validSession, - domainID: domain.ID, - disableRes: disabledDomain, - resp: disabledDomain, - err: nil, - }, - { - desc: "disable domain with empty domainID", - session: validSession, - domainID: "", - disableErr: repoerr.ErrNotFound, - resp: domains.Domain{}, - err: svcerr.ErrUpdateEntity, - }, - { - desc: "disable domain with failed to disable", - session: validSession, - domainID: domain.ID, - disableErr: errors.ErrMalformedEntity, - resp: domains.Domain{}, - err: svcerr.ErrUpdateEntity, - }, - { - desc: "disable domain with failed to remove cache", - session: validSession, - domainID: domain.ID, - disableRes: disabledDomain, - cacheErr: errors.ErrMalformedEntity, - resp: disabledDomain, - err: svcerr.ErrRemoveEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := drepo.On("UpdateDomain", context.Background(), tc.domainID, mock.Anything).Return(tc.disableRes, tc.disableErr) - cacheCall := dcache.On("RemoveStatus", context.Background(), tc.domainID).Return(tc.cacheErr) - domain, err := svc.DisableDomain(context.Background(), tc.session, tc.domainID) - assert.True(t, errors.Contains(err, tc.err)) - assert.Equal(t, tc.disableRes, domain) - repoCall.Unset() - cacheCall.Unset() - }) - } -} - -func TestFreezeDomain(t *testing.T) { - svc := newService() - - freezeDomain := domain - freezeDomain.Status = domains.FreezeStatus - - cases := []struct { - desc string - session authn.Session - domainID string - freezeRes domains.Domain - freezeErr error - cacheErr error - resp domains.Domain - err error - }{ - { - desc: "freeze domain successfully", - session: validSession, - domainID: domain.ID, - freezeRes: freezeDomain, - resp: freezeDomain, - err: nil, - }, - { - desc: "freeze domain with empty domainID", - session: validSession, - domainID: "", - freezeErr: repoerr.ErrNotFound, - resp: domains.Domain{}, - err: svcerr.ErrUpdateEntity, - }, - { - desc: "freeze domain with failed to freeze", - session: validSession, - domainID: domain.ID, - freezeErr: errors.ErrMalformedEntity, - resp: domains.Domain{}, - err: svcerr.ErrUpdateEntity, - }, - { - desc: "freeze domain with failed to remove cache", - session: validSession, - domainID: domain.ID, - freezeRes: freezeDomain, - cacheErr: errors.ErrMalformedEntity, - resp: freezeDomain, - err: svcerr.ErrRemoveEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := drepo.On("UpdateDomain", context.Background(), tc.domainID, mock.Anything).Return(tc.freezeRes, tc.freezeErr) - cacheCall := dcache.On("RemoveStatus", context.Background(), tc.domainID).Return(tc.cacheErr) - domain, err := svc.FreezeDomain(context.Background(), tc.session, tc.domainID) - assert.True(t, errors.Contains(err, tc.err)) - assert.Equal(t, tc.freezeRes, domain) - repoCall.Unset() - cacheCall.Unset() - }) - } -} - -func TestListDomains(t *testing.T) { - svc := newService() - - cases := []struct { - desc string - session authn.Session - domainID string - pageMeta domains.Page - listDomainsRes domains.DomainsPage - listDomainErr error - err error - }{ - { - desc: "list domains successfully", - session: validSession, - domainID: validID, - pageMeta: domains.Page{ - UserID: userID, - Offset: 0, - Limit: 10, - Status: domains.EnabledStatus, - }, - listDomainsRes: domains.DomainsPage{ - Domains: []domains.Domain{domain}, - Offset: 0, - Limit: 10, - Total: 1, - }, - err: nil, - }, - { - desc: "list domains as admin successfully", - session: authn.Session{UserID: validID, SuperAdmin: true}, - domainID: validID, - pageMeta: domains.Page{ - Offset: 0, - Limit: 10, - Status: domains.EnabledStatus, - }, - listDomainsRes: domains.DomainsPage{ - Domains: []domains.Domain{domain}, - Offset: 0, - Limit: 10, - Total: 1, - }, - err: nil, - }, - { - desc: "list domains with repository error on list domains", - session: validSession, - domainID: validID, - pageMeta: domains.Page{ - UserID: userID, - Offset: 0, - Limit: 10, - Status: domains.EnabledStatus, - }, - listDomainErr: errors.ErrMalformedEntity, - err: svcerr.ErrViewEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall1 := drepo.On("ListDomains", context.Background(), tc.pageMeta).Return(tc.listDomainsRes, tc.listDomainErr) - dp, err := svc.ListDomains(context.Background(), tc.session, tc.pageMeta) - assert.True(t, errors.Contains(err, tc.err)) - assert.Equal(t, tc.listDomainsRes, dp) - repoCall1.Unset() - }) - } -} - -func TestSendInvitation(t *testing.T) { - svc := newService() - - rejectedInvitation := validInvitation - rejectedInvitation.RejectedAt = time.Now() - acceptedInvitation := validInvitation - acceptedInvitation.ConfirmedAt = time.Now() - resentInvitation := validInvitation - resentInvitation.Resend = true - - cases := []struct { - desc string - session authn.Session - req domains.Invitation - retrieveRoleErr error - retrieveDomainErr error - createInvitationErr error - retrieveInvRes domains.Invitation - retrieveInvErr error - updateRejectionErr error - err error - }{ - { - desc: "send invitation successful", - session: validSession, - req: validInvitation, - err: nil, - }, - { - desc: "send invitation with invalid role id", - session: validSession, - req: domains.Invitation{ - DomainID: testsutil.GenerateUUID(t), - InviteeUserID: testsutil.GenerateUUID(t), - RoleID: inValid, - }, - retrieveRoleErr: repoerr.ErrNotFound, - err: svcerr.ErrInvalidRole, - }, - { - desc: "send invitation with failed to retrieve domain", - session: validSession, - req: validInvitation, - retrieveDomainErr: repoerr.ErrNotFound, - err: svcerr.ErrViewEntity, - }, - { - desc: "send invitations with failed to save invitation", - session: validSession, - req: validInvitation, - createInvitationErr: repoerr.ErrCreateEntity, - err: svcerr.ErrCreateEntity, - }, - { - desc: "resend invitation successfully", - session: validSession, - req: resentInvitation, - retrieveInvRes: rejectedInvitation, - retrieveInvErr: nil, - err: nil, - }, - { - desc: "resend invitation with failed to retrieve invitation", - session: validSession, - req: resentInvitation, - retrieveInvRes: domains.Invitation{}, - retrieveInvErr: repoerr.ErrNotFound, - err: svcerr.ErrViewEntity, - }, - { - desc: "resend an invitation that is already accepted", - session: validSession, - req: resentInvitation, - retrieveInvRes: acceptedInvitation, - retrieveInvErr: nil, - err: svcerr.ErrInvitationAlreadyAccepted, - }, - { - desc: "resend invitation with failed to update rejection", - session: validSession, - req: resentInvitation, - retrieveInvRes: rejectedInvitation, - retrieveInvErr: nil, - updateRejectionErr: repoerr.ErrUpdateEntity, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := drepo.On("RetrieveRole", context.Background(), tc.req.RoleID).Return(roles.Role{Name: "admin"}, tc.retrieveRoleErr) - repoCall1 := drepo.On("RetrieveDomainByID", context.Background(), tc.req.DomainID).Return(domains.Domain{Name: "test_domain"}, tc.retrieveDomainErr) - repoCall2 := drepo.On("SaveInvitation", context.Background(), mock.Anything).Return(tc.createInvitationErr) - repoCall3 := drepo.On("RetrieveInvitation", context.Background(), tc.req.InviteeUserID, tc.req.DomainID).Return(tc.retrieveInvRes, tc.retrieveInvErr) - repoCall4 := drepo.On("UpdateRejection", context.Background(), mock.Anything).Return(tc.updateRejectionErr) - _, err := svc.SendInvitation(context.Background(), tc.session, tc.req) - assert.True(t, errors.Contains(err, tc.err)) - repoCall.Unset() - repoCall1.Unset() - repoCall2.Unset() - repoCall3.Unset() - repoCall4.Unset() - }) - } -} - -func TestListInvitations(t *testing.T) { - svc := newService() - - validPageMeta := domains.InvitationPageMeta{ - Offset: 0, - Limit: 10, - } - validResp := domains.InvitationPage{ - Total: 1, - Offset: 0, - Limit: 10, - Invitations: []domains.Invitation{ - { - InvitedBy: testsutil.GenerateUUID(t), - InviteeUserID: testsutil.GenerateUUID(t), - DomainID: testsutil.GenerateUUID(t), - RoleID: testsutil.GenerateUUID(t), - RoleName: "admin", - CreatedAt: time.Now().Add(-time.Hour), - UpdatedAt: time.Now().Add(-time.Hour), - ConfirmedAt: time.Now().Add(-time.Hour), - }, - }, - } - - cases := []struct { - desc string - session authn.Session - page domains.InvitationPageMeta - resp domains.InvitationPage - err error - repoErr error - }{ - { - desc: "list invitations successful", - session: validSession, - page: validPageMeta, - resp: validResp, - err: nil, - repoErr: nil, - }, - - { - desc: "list invitations unsuccessful", - session: validSession, - page: validPageMeta, - err: svcerr.ErrViewEntity, - resp: domains.InvitationPage{}, - repoErr: repoerr.ErrViewEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := drepo.On("RetrieveAllInvitations", context.Background(), mock.Anything).Return(tc.resp, tc.repoErr) - resp, err := svc.ListInvitations(context.Background(), tc.session, tc.page) - assert.True(t, errors.Contains(err, tc.err), tc.desc) - assert.Equal(t, tc.resp, resp, tc.desc) - repoCall.Unset() - }) - } -} - -func TestListDomainInvitations(t *testing.T) { - svc := newService() - - validPageMeta := domains.InvitationPageMeta{ - Offset: 0, - Limit: 10, - } - validResp := domains.InvitationPage{ - Total: 1, - Offset: 0, - Limit: 10, - Invitations: []domains.Invitation{ - { - InvitedBy: testsutil.GenerateUUID(t), - InviteeUserID: testsutil.GenerateUUID(t), - DomainID: testsutil.GenerateUUID(t), - RoleID: testsutil.GenerateUUID(t), - RoleName: "admin", - CreatedAt: time.Now().Add(-time.Hour), - UpdatedAt: time.Now().Add(-time.Hour), - ConfirmedAt: time.Now().Add(-time.Hour), - }, - }, - } - - cases := []struct { - desc string - session authn.Session - page domains.InvitationPageMeta - resp domains.InvitationPage - repoErr error - err error - }{ - { - desc: "list domain invitations successful", - session: validSession, - page: validPageMeta, - resp: validResp, - repoErr: nil, - err: nil, - }, - - { - desc: "list domain invitations unsuccessful", - session: validSession, - page: validPageMeta, - resp: domains.InvitationPage{}, - repoErr: repoerr.ErrViewEntity, - err: svcerr.ErrViewEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := drepo.On("RetrieveAllInvitations", context.Background(), mock.Anything).Return(tc.resp, tc.repoErr) - resp, err := svc.ListDomainInvitations(context.Background(), tc.session, tc.page) - assert.True(t, errors.Contains(err, tc.err), tc.desc) - assert.Equal(t, tc.resp, resp, tc.desc) - repoCall.Unset() - }) - } -} - -func TestAcceptInvitation(t *testing.T) { - svc := newService() - - cases := []struct { - desc string - domainID string - session authn.Session - resp domains.Invitation - retrieveInvitationErr error - retrieveDomainErr error - retrieveRoleErr error - updateConfirmationErr error - addRoleMemberErr error - err error - }{ - { - desc: "accept invitation successful", - domainID: validID, - session: validSession, - resp: domains.Invitation{ - InviteeUserID: userID, - DomainID: testsutil.GenerateUUID(t), - RoleID: testsutil.GenerateUUID(t), - }, - err: nil, - }, - { - desc: "accept invitation with failed to retrieve invitation", - session: validSession, - retrieveInvitationErr: repoerr.ErrNotFound, - err: svcerr.ErrNotFound, - }, - { - desc: "accept invitation with of different user", - session: validSession, - resp: domains.Invitation{ - InviteeUserID: testsutil.GenerateUUID(t), - DomainID: testsutil.GenerateUUID(t), - RoleID: testsutil.GenerateUUID(t), - }, - err: svcerr.ErrAuthorization, - }, - { - desc: "accept invitation with failed to add role member", - domainID: validID, - session: validSession, - resp: domains.Invitation{ - InviteeUserID: userID, - DomainID: testsutil.GenerateUUID(t), - RoleID: testsutil.GenerateUUID(t), - }, - addRoleMemberErr: repoerr.ErrMalformedEntity, - err: svcerr.ErrUpdateEntity, - }, - { - desc: "accept invitation with failed update confirmation", - session: validSession, - domainID: validID, - resp: domains.Invitation{ - InviteeUserID: userID, - DomainID: validID, - RoleID: testsutil.GenerateUUID(t), - }, - updateConfirmationErr: repoerr.ErrNotFound, - err: svcerr.ErrUpdateEntity, - }, - { - desc: "accept invitation that is already confirmed", - session: validSession, - domainID: validID, - resp: domains.Invitation{ - InviteeUserID: userID, - DomainID: testsutil.GenerateUUID(t), - RoleID: testsutil.GenerateUUID(t), - ConfirmedAt: time.Now(), - }, - err: svcerr.ErrInvitationAlreadyAccepted, - }, - { - desc: "accept rejected invitation", - session: validSession, - domainID: validID, - resp: domains.Invitation{ - InviteeUserID: userID, - DomainID: testsutil.GenerateUUID(t), - RoleID: testsutil.GenerateUUID(t), - RejectedAt: time.Now(), - }, - err: svcerr.ErrInvitationAlreadyRejected, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := drepo.On("RetrieveInvitation", context.Background(), tc.session.UserID, tc.domainID).Return(tc.resp, tc.retrieveInvitationErr) - repoCall1 := drepo.On("RetrieveDomainByID", context.Background(), tc.domainID).Return(domains.Domain{Name: "test_domain"}, tc.retrieveDomainErr) - repoCall2 := drepo.On("RetrieveRole", context.Background(), tc.resp.RoleID).Return(roles.Role{Name: "admin"}, tc.retrieveRoleErr) - repoCall3 := drepo.On("RetrieveEntityRole", context.Background(), tc.domainID, tc.resp.RoleID).Return(roles.Role{}, tc.addRoleMemberErr) - policyCall := policy.On("AddPolicies", context.Background(), mock.Anything).Return(tc.addRoleMemberErr) - repoCall4 := drepo.On("RoleAddMembers", context.Background(), mock.Anything, []string{tc.resp.InviteeUserID}).Return([]string{}, tc.addRoleMemberErr) - repoCall5 := drepo.On("UpdateConfirmation", context.Background(), mock.Anything).Return(tc.updateConfirmationErr) - _, err := svc.AcceptInvitation(context.Background(), tc.session, tc.domainID) - assert.True(t, errors.Contains(err, tc.err)) - repoCall.Unset() - repoCall1.Unset() - repoCall2.Unset() - repoCall3.Unset() - policyCall.Unset() - repoCall4.Unset() - repoCall5.Unset() - }) - } -} - -func TestRejectInvitation(t *testing.T) { - svc := newService() - - cases := []struct { - desc string - domainID string - session authn.Session - resp domains.Invitation - retrieveInvitationErr error - updateConfirmationErr error - addRoleMemberErr error - err error - }{ - { - desc: "reject invitation successful", - domainID: validID, - session: validSession, - resp: domains.Invitation{ - InviteeUserID: userID, - DomainID: testsutil.GenerateUUID(t), - RoleID: testsutil.GenerateUUID(t), - }, - err: nil, - }, - { - desc: "reject invitation with failed to retrieve invitation", - session: validSession, - retrieveInvitationErr: repoerr.ErrNotFound, - err: svcerr.ErrNotFound, - }, - { - desc: "reject invitation with of different user", - session: validSession, - resp: domains.Invitation{ - InviteeUserID: testsutil.GenerateUUID(t), - DomainID: testsutil.GenerateUUID(t), - RoleID: testsutil.GenerateUUID(t), - }, - err: svcerr.ErrAuthorization, - }, - { - desc: "reject invitation with failed update confirmation", - session: validSession, - domainID: validID, - resp: domains.Invitation{ - InviteeUserID: userID, - DomainID: validID, - RoleID: testsutil.GenerateUUID(t), - }, - updateConfirmationErr: repoerr.ErrNotFound, - err: svcerr.ErrUpdateEntity, - }, - { - desc: "reject invitation that is already confirmed", - session: validSession, - domainID: validID, - resp: domains.Invitation{ - InviteeUserID: userID, - DomainID: testsutil.GenerateUUID(t), - RoleID: testsutil.GenerateUUID(t), - ConfirmedAt: time.Now(), - }, - err: svcerr.ErrInvitationAlreadyAccepted, - }, - { - desc: "reject rejected invitation", - session: validSession, - domainID: validID, - resp: domains.Invitation{ - InviteeUserID: userID, - DomainID: testsutil.GenerateUUID(t), - RoleID: testsutil.GenerateUUID(t), - RejectedAt: time.Now(), - }, - err: svcerr.ErrInvitationAlreadyRejected, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := drepo.On("RetrieveInvitation", context.Background(), tc.session.UserID, tc.domainID).Return(tc.resp, tc.retrieveInvitationErr) - repoCall1 := drepo.On("RetrieveDomainByID", context.Background(), tc.domainID).Return(domains.Domain{Name: "test_domain"}, nil) - repoCall2 := drepo.On("RetrieveRole", context.Background(), tc.resp.RoleID).Return(roles.Role{Name: "admin"}, nil) - repoCall3 := drepo.On("UpdateRejection", context.Background(), mock.Anything).Return(tc.updateConfirmationErr) - _, err := svc.RejectInvitation(context.Background(), tc.session, tc.domainID) - assert.True(t, errors.Contains(err, tc.err)) - repoCall.Unset() - repoCall1.Unset() - repoCall2.Unset() - repoCall3.Unset() - }) - } -} - -func TestDeleteInvitation(t *testing.T) { - svc := newService() - - acceptedInv := validInvitation - acceptedInv.ConfirmedAt = time.Now() - rejectedInv := validInvitation - rejectedInv.RejectedAt = time.Now() - - cases := []struct { - desc string - userID string - domainID string - session authn.Session - resp domains.Invitation - retrieveInvitationErr error - deleteInvitationErr error - err error - }{ - { - desc: "delete invitations successful", - userID: testsutil.GenerateUUID(t), - domainID: testsutil.GenerateUUID(t), - session: validSession, - resp: validInvitation, - err: nil, - }, - { - desc: "delete invitations for the same user", - userID: validInvitation.InviteeUserID, - domainID: validInvitation.DomainID, - resp: validInvitation, - session: authn.Session{UserID: validInvitation.InviteeUserID}, - err: nil, - }, - { - desc: "delete invitations for the invited user", - userID: validInvitation.InviteeUserID, - domainID: validInvitation.DomainID, - session: validSession, - resp: validInvitation, - err: nil, - }, - { - desc: "delete accepted invitation as non invitee user", - userID: validID, - domainID: validInvitation.DomainID, - session: validSession, - resp: acceptedInv, - err: svcerr.ErrInvitationAlreadyAccepted, - }, - { - desc: "delete rejected invitation as non invitee user", - userID: validID, - domainID: validInvitation.DomainID, - session: validSession, - resp: rejectedInv, - err: svcerr.ErrInvitationAlreadyRejected, - }, - { - desc: "delete invitation with error retrieving invitation", - userID: validInvitation.InviteeUserID, - domainID: validInvitation.DomainID, - session: validSession, - resp: domains.Invitation{}, - retrieveInvitationErr: repoerr.ErrNotFound, - err: svcerr.ErrRemoveEntity, - }, - { - desc: "delete invitation with error deleting invitation", - userID: validInvitation.InviteeUserID, - domainID: validInvitation.DomainID, - session: validSession, - resp: domains.Invitation{}, - deleteInvitationErr: repoerr.ErrNotFound, - err: svcerr.ErrRemoveEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := drepo.On("RetrieveInvitation", context.Background(), mock.Anything, mock.Anything).Return(tc.resp, tc.retrieveInvitationErr) - repoCall1 := drepo.On("DeleteUsersInvitations", context.Background(), mock.Anything, mock.Anything).Return(tc.deleteInvitationErr) - err := svc.DeleteInvitation(context.Background(), tc.session, tc.userID, tc.domainID) - assert.True(t, errors.Contains(err, tc.err)) - repoCall.Unset() - repoCall1.Unset() - }) - } -} diff --git a/domains/state.go b/domains/state.go deleted file mode 100644 index 1b9e0b44e..000000000 --- a/domains/state.go +++ /dev/null @@ -1,74 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package domains - -import ( - "encoding/json" - "strings" - - apiutil "github.com/absmach/magistrala/api/http/util" -) - -// State represents invitation state. -type State uint8 - -const ( - AllState State = iota // All is used for querying purposes to list invitations irrespective of their state - both pending and accepted. - Pending // Pending is the state of an invitation that has not been accepted yet. - Accepted // Accepted is the state of an invitation that has been accepted. - Rejected // Rejected is the state of an invitation that has been rejected. -) - -// String representation of the possible state values. -const ( - all = "all" - pending = "pending" - accepted = "accepted" - rejected = "rejected" - UnknownState = "unknown" -) - -// String converts invitation state to string literal. -func (s State) String() string { - switch s { - case AllState: - return all - case Pending: - return pending - case Accepted: - return accepted - case Rejected: - return rejected - default: - return UnknownState - } -} - -// ToState converts string value to a valid invitation state. -func ToState(status string) (State, error) { - switch status { - case all: - return AllState, nil - case pending: - return Pending, nil - case accepted: - return Accepted, nil - case rejected: - return Rejected, nil - } - - return State(0), apiutil.ErrInvitationState -} - -func (s State) MarshalJSON() ([]byte, error) { - return json.Marshal(s.String()) -} - -// Custom Unmarshaler for Client/Groups. -func (s *State) UnmarshalJSON(data []byte) error { - str := strings.Trim(string(data), "\"") - val, err := ToState(str) - *s = val - return err -} diff --git a/domains/state_test.go b/domains/state_test.go deleted file mode 100644 index ac9fe63ef..000000000 --- a/domains/state_test.go +++ /dev/null @@ -1,95 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package domains_test - -import ( - "testing" - - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/domains" - "github.com/stretchr/testify/assert" -) - -func TestState_String(t *testing.T) { - tests := []struct { - name string - state domains.State - expected string - }{ - {"Pending", domains.Pending, "pending"}, - {"Accepted", domains.Accepted, "accepted"}, - {"Rejected", domains.Rejected, "rejected"}, - {"All", domains.AllState, "all"}, - {"Unknown", domains.State(100), "unknown"}, - } - - for _, tt := range tests { - got := tt.state.String() - assert.Equal(t, tt.expected, got, "State.String() = %v, expected %v", got, tt.expected) - } -} - -func TestToState(t *testing.T) { - tests := []struct { - name string - status string - state domains.State - err error - }{ - {"Pending", "pending", domains.Pending, nil}, - {"Accepted", "accepted", domains.Accepted, nil}, - {"Rejected", "rejected", domains.Rejected, nil}, - {"All", "all", domains.AllState, nil}, - {"Unknown", "unknown", domains.State(0), apiutil.ErrInvitationState}, - } - - for _, tt := range tests { - got, err := domains.ToState(tt.status) - assert.Equal(t, tt.err, err, "ToState() error = %v, expected %v", err, tt.err) - assert.Equal(t, tt.state, got, "ToState() = %v, expected %v", got, tt.state) - } -} - -func TestState_MarshalJSON(t *testing.T) { - tests := []struct { - name string - state domains.State - expected []byte - err error - }{ - {"Pending", domains.Pending, []byte(`"pending"`), nil}, - {"Accepted", domains.Accepted, []byte(`"accepted"`), nil}, - {"Rejected", domains.Rejected, []byte(`"rejected"`), nil}, - {"All", domains.AllState, []byte(`"all"`), nil}, - {"Unknown", domains.State(100), []byte(`"unknown"`), nil}, - } - - for _, tt := range tests { - got, err := tt.state.MarshalJSON() - assert.Equal(t, tt.expected, got, "State.MarshalJSON() = %v, expected %v", got, tt.expected) - assert.Equal(t, tt.err, err, "State.MarshalJSON() error = %v, expected %v", err, tt.err) - } -} - -func TestState_UnmarshalJSON(t *testing.T) { - tests := []struct { - name string - data []byte - state domains.State - err error - }{ - {"Pending", []byte(`"pending"`), domains.Pending, nil}, - {"Accepted", []byte(`"accepted"`), domains.Accepted, nil}, - {"Rejected", []byte(`"rejected"`), domains.Rejected, nil}, - {"All", []byte(`"all"`), domains.AllState, nil}, - {"Unknown", []byte(`"unknown"`), domains.State(0), apiutil.ErrInvitationState}, - } - - for _, tt := range tests { - var state domains.State - err := state.UnmarshalJSON(tt.data) - assert.Equal(t, tt.err, err, "State.UnmarshalJSON() error = %v, expected %v", err, tt.err) - assert.Equal(t, tt.state, state, "State.UnmarshalJSON() = %v, expected %v", state, tt.state) - } -} diff --git a/fluxmq/api/grpc/server.go b/fluxmq/api/grpc/server.go index bbce363c9..fc9b95f06 100644 --- a/fluxmq/api/grpc/server.go +++ b/fluxmq/api/grpc/server.go @@ -13,7 +13,7 @@ import ( grpcChannelsV1 "github.com/absmach/magistrala/api/grpc/channels/v1" grpcClientsV1 "github.com/absmach/magistrala/api/grpc/clients/v1" apiutil "github.com/absmach/magistrala/api/http/util" - smqauth "github.com/absmach/magistrala/auth" + "github.com/absmach/magistrala/internal/atom" "github.com/absmach/magistrala/pkg/authn" "github.com/absmach/magistrala/pkg/connections" "github.com/absmach/magistrala/pkg/errors" @@ -28,6 +28,7 @@ type connectServer struct { authv1connect.UnimplementedAuthServiceHandler clients grpcClientsV1.ClientsServiceClient channels grpcChannelsV1.ChannelsServiceClient + atomAuth atom.Authorizer parser messaging.TopicParser } @@ -37,10 +38,16 @@ func NewServer( clients grpcClientsV1.ClientsServiceClient, channels grpcChannelsV1.ChannelsServiceClient, parser messaging.TopicParser, + atomAuth ...atom.Authorizer, ) authv1connect.AuthServiceHandler { + var authz atom.Authorizer + if len(atomAuth) > 0 { + authz = atomAuth[0] + } return &connectServer{ clients: clients, channels: channels, + atomAuth: authz, parser: parser, } } @@ -96,6 +103,33 @@ func (s *connectServer) Authorize(ctx context.Context, req *connect.Request[auth return connect.NewResponse(&authv1.AuthzRes{Authorized: true}), nil } + if s.atomAuth != nil { + action := "subscribe" + if connType == connections.Publish { + action = "publish" + } + res, err := s.atomAuth.CheckAuthz(ctx, atom.AuthzRequest{ + SubjectID: req.Msg.GetExternalId(), + Action: action, + ResourceID: channelID, + ObjectKind: "resource", + ObjectID: channelID, + Context: map[string]any{ + "domain_id": domainID, + "client_type": policies.ClientType, + "connection": connType.String(), + "topic_type": uint32(topicType), + }, + }) + if err != nil { + if shouldDenyAuthorize(err) { + return connect.NewResponse(&authv1.AuthzRes{Authorized: false}), nil + } + return nil, encodeError(err) + } + return connect.NewResponse(&authv1.AuthzRes{Authorized: res.Allowed}), nil + } + ar := &grpcChannelsV1.AuthzReq{ Type: uint32(connType), ClientId: req.Msg.GetExternalId(), @@ -151,7 +185,7 @@ func encodeError(err error) error { err == apiutil.ErrMissingID: return connect.NewError(connect.CodeInvalidArgument, err) case errors.Contains(err, svcerr.ErrAuthentication), - errors.Contains(err, smqauth.ErrKeyExpired): + strings.Contains(err.Error(), "use of expired key"): return connect.NewError(connect.CodeUnauthenticated, err) case errors.Contains(err, svcerr.ErrAuthorization): return connect.NewError(connect.CodePermissionDenied, err) diff --git a/fluxmq/api/http/publish.go b/fluxmq/api/http/publish.go new file mode 100644 index 000000000..ae60c2cd2 --- /dev/null +++ b/fluxmq/api/http/publish.go @@ -0,0 +1,258 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +// Package http exposes user-authenticated message publishing for the MG UI. +package http + +import ( + "context" + "encoding/json" + "fmt" + "net/http" + "strings" + "time" + + "github.com/absmach/magistrala/internal/atom" + smqauthn "github.com/absmach/magistrala/pkg/authn" + "github.com/absmach/magistrala/pkg/messaging" + "github.com/go-chi/chi/v5" +) + +const ( + contentType = "application/json" + httpProto = "http" +) + +type publishRequest struct { + ClientID string `json:"client_id"` + Subtopic string `json:"subtopic"` + Payload json.RawMessage `json:"payload"` +} + +type publishResponse struct { + Status string `json:"status"` +} + +type errorResponse struct { + Error string `json:"error"` +} + +type publishHandler struct { + authn smqauthn.Authentication + atom *atom.Client + publisher messaging.Publisher +} + +// MakePublishHandler returns an HTTP handler that authenticates the user with +// Atom, authorizes publish access in Atom, and writes directly to the message +// broker with the selected client as the publisher identity. +func MakePublishHandler( + authn smqauthn.Authentication, + atomClient *atom.Client, + publisher messaging.Publisher, +) http.Handler { + h := publishHandler{ + authn: authn, + atom: atomClient, + publisher: publisher, + } + r := chi.NewRouter() + r.Post("/{domainID}/channels/{channelID}/messages", h.publish) + r.Get("/health", func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", contentType) + w.WriteHeader(http.StatusOK) + if err := json.NewEncoder(w).Encode(publishResponse{Status: "ok"}); err != nil { + return + } + }) + return r +} + +func (h publishHandler) publish(w http.ResponseWriter, r *http.Request) { + domainID := chi.URLParam(r, "domainID") + channelID := chi.URLParam(r, "channelID") + if domainID == "" || channelID == "" { + writeError(w, http.StatusBadRequest, "domainID and channelID are required") + return + } + + token := bearerToken(r) + if token == "" { + writeError(w, http.StatusUnauthorized, "bearer token is required") + return + } + session, err := h.authn.Authenticate(r.Context(), token) + if err != nil { + writeError(w, http.StatusUnauthorized, "invalid bearer token") + return + } + + var req publishRequest + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + writeError(w, http.StatusBadRequest, "invalid publish request") + return + } + payload, err := payloadBytes(req.Payload) + if err != nil { + writeError(w, http.StatusBadRequest, err.Error()) + return + } + + publisherID := session.UserID + if req.ClientID != "" { + if err := h.ensureClientPublisher(r.Context(), domainID, channelID, session.UserID, req.ClientID); err != nil { + writeError(w, http.StatusForbidden, err.Error()) + return + } + publisherID = req.ClientID + } + if err := h.ensureUserPublish(r.Context(), domainID, channelID, session.UserID, req.ClientID); err != nil { + writeError(w, http.StatusForbidden, err.Error()) + return + } + + subtopic := cleanSubtopic(req.Subtopic) + topic := messaging.EncodeTopicSuffix(domainID, channelID, subtopic) + + msg := &messaging.Message{ + Domain: domainID, + Channel: channelID, + Subtopic: subtopic, + Publisher: publisherID, + ClientId: session.UserID, + Protocol: httpProto, + Payload: payload, + Created: time.Now().UnixNano(), + } + if err := h.publisher.Publish(r.Context(), topic, msg); err != nil { + writeError(w, http.StatusBadGateway, "failed to publish message") + return + } + + w.Header().Set("Content-Type", contentType) + w.WriteHeader(http.StatusAccepted) + if err := json.NewEncoder(w).Encode(publishResponse{Status: "accepted"}); err != nil { + return + } +} + +func (h publishHandler) ensureUserPublish( + ctx context.Context, + domainID string, + channelID string, + userID string, + clientID string, +) error { + res, err := h.atom.CheckAuthz(ctx, atom.AuthzRequest{ + SubjectID: userID, + Action: "publish", + ResourceID: channelID, + ObjectKind: "resource", + ObjectID: channelID, + Context: map[string]any{ + "domain_id": domainID, + "publisher_client_id": clientID, + }, + }) + if err != nil { + return err + } + if !res.Allowed { + return fmt.Errorf("user is not allowed to publish to channel") + } + return nil +} + +func (h publishHandler) ensureClientPublisher( + ctx context.Context, + domainID string, + channelID string, + userID string, + clientID string, +) error { + client, err := h.atom.GetEntity(ctx, clientID) + if err != nil { + return fmt.Errorf("publisher client not found") + } + if client.Kind != "device" && attrString(client.Attributes, "magistrala_kind") != atom.KindClient { + return fmt.Errorf("publisher identity is not a client") + } + if client.TenantID == "" || client.TenantID != domainID { + return fmt.Errorf("publisher client belongs to a different domain") + } + userAccess, err := h.atom.CheckAuthz(ctx, atom.AuthzRequest{ + SubjectID: userID, + Action: "read", + ResourceID: clientID, + ObjectKind: "entity", + ObjectID: clientID, + Context: map[string]any{ + "domain_id": domainID, + }, + }) + if err != nil { + return err + } + if !userAccess.Allowed { + return fmt.Errorf("user is not allowed to use publisher client") + } + res, err := h.atom.CheckAuthz(ctx, atom.AuthzRequest{ + SubjectID: clientID, + Action: "publish", + ResourceID: channelID, + ObjectKind: "resource", + ObjectID: channelID, + Context: map[string]any{ + "domain_id": domainID, + }, + }) + if err != nil { + return err + } + if !res.Allowed { + return fmt.Errorf("publisher client is not connected for publish") + } + return nil +} + +func bearerToken(r *http.Request) string { + token := r.Header.Get("Authorization") + return strings.TrimPrefix(token, "Bearer ") +} + +func payloadBytes(raw json.RawMessage) ([]byte, error) { + if len(raw) == 0 { + return nil, fmt.Errorf("payload is required") + } + if raw[0] == '"' { + var value string + if err := json.Unmarshal(raw, &value); err != nil { + return nil, fmt.Errorf("payload must be a string or JSON value") + } + return []byte(value), nil + } + return raw, nil +} + +func cleanSubtopic(subtopic string) string { + return strings.Trim(strings.ReplaceAll(subtopic, ".", "/"), "/") +} + +func attrString(attrs atom.Attributes, key string) string { + value, ok := attrs[key] + if !ok || value == nil { + return "" + } + if str, ok := value.(string); ok { + return str + } + return fmt.Sprint(value) +} + +func writeError(w http.ResponseWriter, status int, message string) { + w.Header().Set("Content-Type", contentType) + w.WriteHeader(status) + if err := json.NewEncoder(w).Encode(errorResponse{Error: message}); err != nil { + return + } +} diff --git a/go.mod b/go.mod index a827e2b00..d16e839fa 100644 --- a/go.mod +++ b/go.mod @@ -3,14 +3,12 @@ module github.com/absmach/magistrala go 1.26.4 require ( + connectrpc.com/connect v1.20.0 connectrpc.com/otelconnect v0.9.0 github.com/0x6flab/namegenerator v1.4.0 github.com/absmach/callhome v0.18.2 github.com/absmach/fluxmq v0.30.0 github.com/absmach/senml v1.0.8 - github.com/authzed/authzed-go v1.10.0 - github.com/authzed/grpcutil v0.0.0-20260105210157-e237581949c2 - github.com/authzed/spicedb v1.54.0 github.com/caarlos0/env/v10 v10.0.0 github.com/caarlos0/env/v11 v11.4.1 github.com/dgraph-io/ristretto/v2 v2.4.0 @@ -57,10 +55,9 @@ require ( go.opentelemetry.io/otel/sdk v1.44.0 go.opentelemetry.io/otel/trace v1.44.0 golang.org/x/crypto v0.53.0 - golang.org/x/oauth2 v0.36.0 + golang.org/x/net v0.56.0 golang.org/x/sync v0.21.0 gonum.org/v1/gonum v0.17.0 - google.golang.org/genproto/googleapis/rpc v0.0.0-20260610212136-7ab31c22f7ad google.golang.org/grpc v1.81.1 google.golang.org/protobuf v1.36.11 gopkg.in/gomail.v2 v2.0.0-20160411212932-81ebce5c23df @@ -71,26 +68,18 @@ require ( require ( github.com/smarty/assertions v1.16.0 // indirect github.com/smartystreets/goconvey v1.8.1 // indirect + google.golang.org/genproto/googleapis/rpc v0.0.0-20260610212136-7ab31c22f7ad // indirect ) require ( - buf.build/gen/go/bufbuild/protovalidate/protocolbuffers/go v1.36.11-20260415201107-50325440f8f2.1 // indirect - buf.build/go/protovalidate v1.2.0 // indirect - cel.dev/expr v0.25.2 // indirect - cloud.google.com/go/compute/metadata v0.9.0 // indirect - connectrpc.com/connect v1.20.0 dario.cat/mergo v1.0.2 // indirect filippo.io/edwards25519 v1.2.0 // indirect github.com/Azure/go-ansiterm v0.0.0-20250102033503-faa5f7b0171c // indirect github.com/Microsoft/go-winio v0.6.2 // indirect github.com/Nvveen/Gotty v0.0.0-20120604004816-cd527374f1e5 // indirect - github.com/antlr4-go/antlr/v4 v4.13.1 // indirect - github.com/authzed/cel-go v0.20.2 // indirect github.com/beorn7/perks v1.0.1 // indirect - github.com/ccoveille/go-safecast/v2 v2.0.1 // indirect github.com/cenkalti/backoff/v4 v4.3.0 // indirect github.com/cenkalti/backoff/v5 v5.0.3 // indirect - github.com/certifi/gocertifi v0.0.0-20210507211836-431795d63e8d // indirect github.com/cespare/xxhash/v2 v2.3.0 // indirect github.com/containerd/continuity v0.5.0 // indirect github.com/containerd/errdefs v1.0.0 // indirect @@ -103,12 +92,9 @@ require ( github.com/docker/go-units v0.5.0 // indirect github.com/dsnet/golib/memfile v1.0.0 // indirect github.com/dustin/go-humanize v1.0.1 // indirect - github.com/emirpasic/gods v1.18.1 // indirect - github.com/envoyproxy/protoc-gen-validate v1.3.3 // indirect github.com/felixge/httpsnoop v1.1.0 // indirect github.com/fsnotify/fsnotify v1.10.1 // indirect github.com/fxamacker/cbor/v2 v2.9.2 // indirect - github.com/go-errors/errors v1.5.1 // indirect github.com/go-gorp/gorp/v3 v3.1.0 // indirect github.com/go-jose/go-jose/v4 v4.1.4 // indirect github.com/go-kit/log v0.2.1 // indirect @@ -118,9 +104,7 @@ require ( github.com/go-sql-driver/mysql v1.10.0 // indirect github.com/go-viper/mapstructure/v2 v2.5.0 // indirect github.com/goccy/go-json v0.10.6 // indirect - github.com/google/cel-go v0.28.1 // indirect github.com/google/shlex v0.0.0-20191202100458-e7afc7fbc510 // indirect - github.com/grpc-ecosystem/go-grpc-middleware v1.4.0 // indirect github.com/grpc-ecosystem/grpc-gateway/v2 v2.29.0 // indirect github.com/hashicorp/errwrap v1.1.0 // indirect github.com/hashicorp/go-cleanhttp v0.5.2 // indirect @@ -135,7 +119,6 @@ require ( github.com/jackc/pgpassfile v1.0.0 // indirect github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect github.com/jackc/puddle/v2 v2.2.2 // indirect - github.com/jzelinskie/stringz v0.0.3 // indirect github.com/klauspost/compress v1.18.6 // indirect github.com/lestrrat-go/blackmagic v1.0.4 // indirect github.com/lestrrat-go/httpcc v1.0.1 // indirect @@ -156,17 +139,15 @@ require ( github.com/opencontainers/go-digest v1.0.0 // indirect github.com/opencontainers/image-spec v1.1.1 // indirect github.com/opencontainers/runc v1.3.6 // indirect - github.com/pelletier/go-toml/v2 v2.4.1 + github.com/pelletier/go-toml/v2 v2.4.1 // indirect github.com/pion/dtls/v3 v3.1.4 // indirect github.com/pion/logging v0.2.4 // indirect github.com/pion/transport/v4 v4.0.2 // indirect - github.com/planetscale/vtprotobuf v0.6.1-0.20240917153116-6f2963f01587 // indirect github.com/pmezard/go-difflib v1.0.1-0.20181226105442-5d4384ee4fb2 // indirect github.com/prometheus/client_model v0.6.2 // indirect github.com/prometheus/common v0.68.1 // indirect github.com/prometheus/procfs v0.20.1 // indirect - github.com/rabbitmq/amqp091-go v1.12.0 - github.com/rs/zerolog v1.35.1 // indirect + github.com/rabbitmq/amqp091-go v1.11.0 github.com/ryanuber/go-glob v1.0.0 // indirect github.com/sagikazarmark/locafero v0.12.0 // indirect github.com/segmentio/asm v1.2.1 // indirect @@ -174,7 +155,6 @@ require ( github.com/spf13/afero v1.15.0 // indirect github.com/spf13/cast v1.10.0 // indirect github.com/spf13/pflag v1.0.10 // indirect - github.com/stoewer/go-strcase v1.3.1 // indirect github.com/stretchr/objx v0.5.3 // indirect github.com/subosito/gotenv v1.6.0 // indirect github.com/x448/float16 v0.8.4 // indirect @@ -188,7 +168,6 @@ require ( go.uber.org/atomic v1.11.0 // indirect go.yaml.in/yaml/v3 v3.0.4 // indirect golang.org/x/exp v0.0.0-20260611194520-c48552f49976 // indirect - golang.org/x/net v0.56.0 golang.org/x/sys v0.46.0 // indirect golang.org/x/text v0.38.0 // indirect golang.org/x/time v0.15.0 // indirect diff --git a/go.sum b/go.sum index 259b9bd6e..17d30f554 100644 --- a/go.sum +++ b/go.sum @@ -1,14 +1,5 @@ al.essio.dev/pkg/shellescape v1.5.1/go.mod h1:6sIqp7X2P6mThCQ7twERpZTuigpr6KbZWtls1U8I890= -buf.build/gen/go/bufbuild/protovalidate/protocolbuffers/go v1.36.11-20260415201107-50325440f8f2.1 h1:s6hzCXtND/ICdGPTMGk7C+/BFlr2Jg5GyH0NKf4XGXg= -buf.build/gen/go/bufbuild/protovalidate/protocolbuffers/go v1.36.11-20260415201107-50325440f8f2.1/go.mod h1:tvtbpgaVXZX4g6Pn+AnzFycuRK3MOz5HJfEGeEllXYM= -buf.build/go/protovalidate v1.2.0 h1:DQVrUWkmGTBij+kOYv/x2LLxwcLaGKMdzShj1/6/3H0= -buf.build/go/protovalidate v1.2.0/go.mod h1:7rYiQEhqvAipoazpVNBBH2S2f8bjG4huMVy1V2Yofn4= -cel.dev/expr v0.25.2 h1:K6j46C81hXtZQfuX60cVWQFBJahKSE2gfRbNuvr5bFs= -cel.dev/expr v0.25.2/go.mod h1:hrXvqGP6G6gyx8UAHSHJ5RGk//1Oj5nXQ2NI02Nrsg4= -cloud.google.com/go v0.26.0/go.mod h1:aQUYkXzVsufM+DwF1aE+0xfcU+56JwCaLick0ClmMTw= cloud.google.com/go v0.34.0/go.mod h1:aQUYkXzVsufM+DwF1aE+0xfcU+56JwCaLick0ClmMTw= -cloud.google.com/go/compute/metadata v0.9.0 h1:pDUj4QMoPejqq20dK0Pg2N4yG9zIkYGdBtwLoEkH9Zs= -cloud.google.com/go/compute/metadata v0.9.0/go.mod h1:E0bWwX5wTnLPedCKqk3pJmVgCBSM6qQI1yTBdEb3C10= connectrpc.com/connect v1.20.0 h1:6TNDAB+WeNd2uolWNlYczB5E0KNNaVMNUEx8JEUsPmQ= connectrpc.com/connect v1.20.0/go.mod h1:A2ygJrukXwWy32vkCAAHNVguZrqZ+jeZ9rGRnGR4dN4= connectrpc.com/otelconnect v0.9.0 h1:NggB3pzRC3pukQWaYbRHJulxuXvmCKCKkQ9hbrHAWoA= @@ -42,26 +33,13 @@ github.com/alecthomas/template v0.0.0-20190718012654-fb15b899a751/go.mod h1:LOuy github.com/alecthomas/units v0.0.0-20151022065526-2efee857e7cf/go.mod h1:ybxpYRFXyAe+OPACYpWeL0wqObRcbAqCMya13uyzqw0= github.com/alecthomas/units v0.0.0-20190717042225-c3de453c63f4/go.mod h1:ybxpYRFXyAe+OPACYpWeL0wqObRcbAqCMya13uyzqw0= github.com/alecthomas/units v0.0.0-20190924025748-f65c72e2690d/go.mod h1:rBZYJk541a8SKzHPHnH3zbiI+7dagKZ0cgpgrD7Fyho= -github.com/antlr4-go/antlr/v4 v4.13.1 h1:SqQKkuVZ+zWkMMNkjy5FZe5mr5WURWnlpmOuzYWrPrQ= -github.com/antlr4-go/antlr/v4 v4.13.1/go.mod h1:GKmUxMtwp6ZgGwZSva4eWPC5mS6vUAmOABFgjdkM7Nw= -github.com/authzed/authzed-go v1.10.0 h1:GUPzYFnStk1PIBZOkQzqcYA3iHaOlVYVl1UnHIRofKU= -github.com/authzed/authzed-go v1.10.0/go.mod h1:2DL7pg4iqMltwWOSw+wvbEzAK7uRt3545+bkcGYD8D8= -github.com/authzed/cel-go v0.20.2 h1:GlmLecGry7Z8HU0k+hmaHHUV05ZHrsFxduXHtIePvck= -github.com/authzed/cel-go v0.20.2/go.mod h1:pJHVFWbqUHV1J+klQoZubdKswlbxcsbojda3mye9kiU= -github.com/authzed/grpcutil v0.0.0-20260105210157-e237581949c2 h1:ymPD1ugBsXVUpLIG/lnRn1ndgOrsrki/0ZX7uP/S1GI= -github.com/authzed/grpcutil v0.0.0-20260105210157-e237581949c2/go.mod h1:FLssYBs1DrwuItfI411kzqcV8QSqGb/B7PC6snNhjvU= -github.com/authzed/spicedb v1.54.0 h1:XQiFm/G/YpQktYmcLCVrQRWkkLpAjUBZa1SBoEbmd0A= -github.com/authzed/spicedb v1.54.0/go.mod h1:9McK3CNkb831IzWZ1pVsswVoqjZTyDPy3fUF7GcvvOY= github.com/aws/aws-sdk-go v1.34.0/go.mod h1:5zCpMtNQVjRREroY7sYe8lOMRSxkhG6MZveU8YkpAk0= github.com/aws/aws-sdk-go v1.40.45 h1:QN1nsY27ssD/JmW4s83qmSb+uL6DG4GmCDzjmJB4xUI= github.com/aws/aws-sdk-go v1.40.45/go.mod h1:585smgzpB/KqRA+K3y/NL/oYRqQvpNJYvLm+LY1U59Q= -github.com/benbjohnson/clock v1.1.0/go.mod h1:J11/hYXuz8f4ySSvYwY0FKfm+ezbsZBKZxNJlLklBHA= github.com/beorn7/perks v0.0.0-20180321164747-3a771d992973/go.mod h1:Dwedo/Wpr24TaqPxmxbtue+5NUziq4I4S80YR8gNf3Q= github.com/beorn7/perks v1.0.0/go.mod h1:KWe93zE9D1o94FZ5RNwFwVgaQK1VOXiVxmqh+CedLV8= github.com/beorn7/perks v1.0.1 h1:VlbKKnNfV8bJzeqoa4cOKqO6bYr3WgKZxO8Z16+hsOM= github.com/beorn7/perks v1.0.1/go.mod h1:G2ZrVWU2WbWT9wwq4/hrbKbnv/1ERSJQ0ibhJ6rlkpw= -github.com/brianvoe/gofakeit/v6 v6.28.0 h1:Xib46XXuQfmlLS2EXRuJpqcw8St6qSZz75OUo0tgAW4= -github.com/brianvoe/gofakeit/v6 v6.28.0/go.mod h1:Xj58BMSnFqcn/fAQeSK+/PLtC5kSb7FJIq4JyGa8vEs= github.com/bsm/ginkgo/v2 v2.12.0 h1:Ny8MWAHyOepLGlLKYmXG4IEkioBysk6GpaRTLC8zwWs= github.com/bsm/ginkgo/v2 v2.12.0/go.mod h1:SwYbGRRDovPVboqFv0tPTcG1sN61LM1Z4ARdbAV9g4c= github.com/bsm/gomega v1.27.10 h1:yeMWxP2pV2fG3FgAODIY8EiRE3dy0aeFYt4l7wh6yKA= @@ -72,15 +50,10 @@ github.com/caarlos0/env/v11 v11.4.1 h1:fYwH0sWEsBSMPG7t4e/PEfTFzrWrpjyygXyUnWiSw github.com/caarlos0/env/v11 v11.4.1/go.mod h1:qupehSf/Y0TUTsxKywqRt/vJjN5nz6vauiYEUUr8P4U= github.com/cbroglie/mustache v1.0.1 h1:ivMg8MguXq/rrz2eu3tw6g3b16+PQhoTn6EZAhst2mw= github.com/cbroglie/mustache v1.0.1/go.mod h1:R/RUa+SobQ14qkP4jtx5Vke5sDytONDQXNLPY/PO69g= -github.com/ccoveille/go-safecast/v2 v2.0.1 h1:2+mIu3gXtwmWelBia2kkxfB8eP4orTHDH7ClSlWkd6I= -github.com/ccoveille/go-safecast/v2 v2.0.1/go.mod h1:JIYA4CAR33blIDuE6fSwCp2sz1oOBahXnvmdBhOAABs= github.com/cenkalti/backoff/v4 v4.3.0 h1:MyRJ/UdXutAwSAT+s3wNd7MfTIcy71VQueUuFK343L8= github.com/cenkalti/backoff/v4 v4.3.0/go.mod h1:Y3VNntkOUPxTVeUxJ/G5vcM//AlwfmyYozVcomhLiZE= github.com/cenkalti/backoff/v5 v5.0.3 h1:ZN+IMa753KfX5hd8vVaMixjnqRZ3y8CuJKRKj1xcsSM= github.com/cenkalti/backoff/v5 v5.0.3/go.mod h1:rkhZdG3JZukswDf7f0cwqPNk4K0sa+F97BxZthm/crw= -github.com/census-instrumentation/opencensus-proto v0.2.1/go.mod h1:f6KPmirojxKA12rnyqOA5BBL4O983OfeGPqjHWSTneU= -github.com/certifi/gocertifi v0.0.0-20210507211836-431795d63e8d h1:S2NE3iHSwP0XV47EEXL8mWmRdEfGscSJ+7EgePNgt0s= -github.com/certifi/gocertifi v0.0.0-20210507211836-431795d63e8d/go.mod h1:sGbDF6GwGcLpkNXPUTkMRoywsNa/ol15pxFe6ERfguA= github.com/cespare/xxhash/v2 v2.1.1/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs= github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs= github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs= @@ -88,8 +61,6 @@ github.com/cheggaaa/pb/v3 v3.0.5/go.mod h1:X1L61/+36nz9bjIsrDU52qHKOQukUQe2Ge+Yv github.com/chzyer/logex v1.1.10/go.mod h1:+Ywpsq7O8HXn0nuIou7OrIPyXbp3wmkHB+jjWRnGsAI= github.com/chzyer/readline v0.0.0-20180603132655-2972be24d48e/go.mod h1:nSuG5e5PlCu98SY8svDHJxuZscDgtXS6KTTbou5AhLI= github.com/chzyer/test v0.0.0-20180213035817-a1ea475d72b1/go.mod h1:Q3SI9o4m/ZMnBNeIyt5eFwwo7qiLfzFZmjNmxjkiQlU= -github.com/client9/misspell v0.3.4/go.mod h1:qj6jICC3Q7zFZvVWo7KLAzC3yx5G7kyvSDkc90ppPyw= -github.com/cncf/udpa/go v0.0.0-20191209042840-269d4d468f6f/go.mod h1:M8M6+tZqaGXZJjfX53e64911xZQV5JYwmTeXPW+k8Sc= github.com/cockroachdb/apd v1.1.0/go.mod h1:8Sl8LxpKi29FqWXR16WEFZRNSz3SoPzUzeMeY4+DwBQ= github.com/containerd/continuity v0.5.0 h1:7a85HZpCSs+1Zps0Ee3DPSuAWY+0SJM1JNM51nlEVDg= github.com/containerd/continuity v0.5.0/go.mod h1:/lNJvtJKUQStBzpVQ1+rasXO1LAWtUQssk28EZvJ3nE= @@ -130,14 +101,6 @@ github.com/dustin/go-humanize v1.0.1 h1:GzkhY7T5VNhEkwH0PVJgjz+fX1rhBrR7pRT3mDkp github.com/dustin/go-humanize v1.0.1/go.mod h1:Mu1zIs6XwVuF/gI1OepvI0qD18qycQx+mFykh5fBlto= github.com/eclipse/paho.mqtt.golang v1.5.1 h1:/VSOv3oDLlpqR2Epjn1Q7b2bSTplJIeV2ISgCl2W7nE= github.com/eclipse/paho.mqtt.golang v1.5.1/go.mod h1:1/yJCneuyOoCOzKSsOTUc0AJfpsItBGWvYpBLimhArU= -github.com/emirpasic/gods v1.18.1 h1:FXtiHYKDGKCW2KzwZKx0iC0PQmdlorYgdFG9jPXJ1Bc= -github.com/emirpasic/gods v1.18.1/go.mod h1:8tpGGwCnJ5H4r6BWwaV6OrWmMoPhUl5jm/FMNAnJvWQ= -github.com/envoyproxy/go-control-plane v0.9.0/go.mod h1:YTl/9mNaCwkRvm6d1a2C3ymFceY/DCBVvsKhRF0iEA4= -github.com/envoyproxy/go-control-plane v0.9.1-0.20191026205805-5f8ba28d4473/go.mod h1:YTl/9mNaCwkRvm6d1a2C3ymFceY/DCBVvsKhRF0iEA4= -github.com/envoyproxy/go-control-plane v0.9.4/go.mod h1:6rpuAdCZL397s3pYoYcLgu1mIlRU8Am5FuJP05cCM98= -github.com/envoyproxy/protoc-gen-validate v0.1.0/go.mod h1:iSmxcyjqTsJpI2R4NaDN7+kN2VEUnK/pcBlmesArF7c= -github.com/envoyproxy/protoc-gen-validate v1.3.3 h1:MVQghNeW+LZcmXe7SY1V36Z+WFMDjpqGAGacLe2T0ds= -github.com/envoyproxy/protoc-gen-validate v1.3.3/go.mod h1:TsndJ/ngyIdQRhMcVVGDDHINPLWB7C82oDArY51KfB0= github.com/fatih/color v1.7.0/go.mod h1:Zm6kSWBoL9eyXnKyktHP6abPY2pDugNf5KwzbycvMj4= github.com/fatih/color v1.13.0/go.mod h1:kLAiJbzzSOZDVNGyDpeOxJ47H46qBXwg5ILebYFFOfk= github.com/fatih/color v1.19.0 h1:Zp3PiM21/9Ld6FzSKyL5c/BULoe/ONr9KlbYVOfG8+w= @@ -154,8 +117,6 @@ github.com/fxamacker/cbor/v2 v2.9.2 h1:X4Ksno9+x3cz0TZv69ec1hxP/+tymuR8PXQJyDwfh github.com/fxamacker/cbor/v2 v2.9.2/go.mod h1:vM4b+DJCtHn+zz7h3FFp/hDAI9WNWCsZj23V5ytsSxQ= github.com/go-chi/chi/v5 v5.3.0 h1:halUjDxhshgXHMrao5bB8eNBXo/rnzwr8m5m36glehM= github.com/go-chi/chi/v5 v5.3.0/go.mod h1:R+tYY2hNuVUUjxoPtqUdgBqevM9s9njzkTLutVsOCto= -github.com/go-errors/errors v1.5.1 h1:ZwEMSLRCapFLflTpT7NKaAc7ukJ8ZPEjzlxt8rPN8bk= -github.com/go-errors/errors v1.5.1/go.mod h1:sIVyrIiJhuEF+Pj9Ebtd6P/rEYROXFi3BopGUQ5a5Og= github.com/go-gorp/gorp/v3 v3.1.0 h1:ItKF/Vbuj31dmV4jxA1qblpSwkl9g1typ24xoe70IGs= github.com/go-gorp/gorp/v3 v3.1.0/go.mod h1:dLEjIyyRNiXvNZ8PSmzpt1GsWAUK8kjVhEpjH8TixEw= github.com/go-jose/go-jose/v4 v4.1.4 h1:moDMcTHmvE6Groj34emNPLs/qtYXRVcd6S7NHbHz3kA= @@ -192,13 +153,9 @@ github.com/gofrs/uuid v4.0.0+incompatible/go.mod h1:b2aQJv3Z4Fp6yNu3cdSllBxTCLRx github.com/gofrs/uuid/v5 v5.4.0 h1:EfbpCTjqMuGyq5ZJwxqzn3Cbr2d0rUZU7v5ycAk/e/0= github.com/gofrs/uuid/v5 v5.4.0/go.mod h1:CDOjlDMVAtN56jqyRUZh58JT31Tiw7/oQyEXZV+9bD8= github.com/gogo/protobuf v1.1.1/go.mod h1:r8qH/GZQm5c6nD/R0oafs1akxWv10x8SbQlK7atdtwQ= -github.com/gogo/protobuf v1.3.2/go.mod h1:P1XiOD3dCwIKUDQYPy72D8LYyHL2YPYrpS2s69NZV8Q= -github.com/golang/glog v0.0.0-20160126235308-23def4e6c14b/go.mod h1:SBH7ygxi8pfUlaOkMMuAQtPIUF8ecWP5IEl/CR7VP2Q= -github.com/golang/mock v1.1.1/go.mod h1:oTYuIxOrZwtPieC+H1uAHpcLFnEyAGVDL/k47Jfbm0A= github.com/golang/protobuf v1.2.0/go.mod h1:6lQm79b+lXiMfvg/cZm0SGofjICqVBUtrP5yJMmIC1U= github.com/golang/protobuf v1.3.1/go.mod h1:6lQm79b+lXiMfvg/cZm0SGofjICqVBUtrP5yJMmIC1U= github.com/golang/protobuf v1.3.2/go.mod h1:6lQm79b+lXiMfvg/cZm0SGofjICqVBUtrP5yJMmIC1U= -github.com/golang/protobuf v1.3.3/go.mod h1:vzj43D7+SQXF/4pzW/hwtAqwc6iTitCiVSaWz5lYuqw= github.com/golang/protobuf v1.4.0-rc.1/go.mod h1:ceaxUfeHdC40wWswd/P6IGgMaK3YpKi5j83Wpe3EHw8= github.com/golang/protobuf v1.4.0-rc.1.0.20200221234624-67d41d38c208/go.mod h1:xKAWHe0F5eneWXFV3EuXVDTCmh+JuBKY0li0aMyXATA= github.com/golang/protobuf v1.4.0-rc.2/go.mod h1:LlEzMj4AhA7rCAGe4KMBDvJI+AwstrUpVNzEA03Pprs= @@ -208,9 +165,6 @@ github.com/golang/protobuf v1.4.2/go.mod h1:oDoupMAO8OvCJWAcko0GGGIgR6R6ocIYbsSw github.com/golang/protobuf v1.4.3/go.mod h1:oDoupMAO8OvCJWAcko0GGGIgR6R6ocIYbsSw735rRwI= github.com/golang/protobuf v1.5.4 h1:i7eJL8qZTpSEXOPTxNKhASYpMn+8e5Q6AdndVa1dWek= github.com/golang/protobuf v1.5.4/go.mod h1:lnTiLA8Wa4RWRcIUkrtSVa5nRhsEGBg48fD6rSs7xps= -github.com/google/cel-go v0.28.1 h1:YWIwi77J4xIsYUwAF/iIuS6haffzIHS8yWI8glSbLWM= -github.com/google/cel-go v0.28.1/go.mod h1:X0bD6iVNR8pkROSOoHVdgTkzmRcosof7WQqCD6wcMc8= -github.com/google/go-cmp v0.2.0/go.mod h1:oXzfMopK8JAjlY9xF4vHSVASa0yLyX7SntLO5aqRK0M= github.com/google/go-cmp v0.3.0/go.mod h1:8QqcDgzrUqlUb/G2PQTWiueGozuR1884gddMywk6iLU= github.com/google/go-cmp v0.3.1/go.mod h1:8QqcDgzrUqlUb/G2PQTWiueGozuR1884gddMywk6iLU= github.com/google/go-cmp v0.4.0/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/gNBxE= @@ -232,8 +186,6 @@ github.com/gopherjs/gopherjs v1.17.2 h1:fQnZVsXk8uxXIStYb0N4bGk7jeyTalG/wsZjQ25d github.com/gopherjs/gopherjs v1.17.2/go.mod h1:pRRIvn/QzFLrKfvEz3qUuEhtE/zLCWfreZ6J5gM2i+k= github.com/gorilla/websocket v1.5.4-0.20250319132907-e064f32e3674 h1:JeSE6pjso5THxAzdVpqr6/geYxZytqFMBCOtn/ujyeo= github.com/gorilla/websocket v1.5.4-0.20250319132907-e064f32e3674/go.mod h1:r4w70xmWCQKmi1ONH4KIaBptdivuRPyosB9RmPlGEwA= -github.com/grpc-ecosystem/go-grpc-middleware v1.4.0 h1:UH//fgunKIs4JdUbpDl1VZCDaL56wXCB/5+wF6uHfaI= -github.com/grpc-ecosystem/go-grpc-middleware v1.4.0/go.mod h1:g5qyo/la0ALbONm6Vbp88Yd8NsDy6rZz+RcrMPxvld8= github.com/grpc-ecosystem/grpc-gateway/v2 v2.29.0 h1:5VipnvEpbqr2gA2VbM+nYVbkIF28c5ZQfqCBQ5g2xfk= github.com/grpc-ecosystem/grpc-gateway/v2 v2.29.0/go.mod h1:Hyl3n6Twe1hvtd9XUXDec4pTvgMSEixRuQKPTMH2bNs= github.com/hashicorp/errwrap v1.0.0/go.mod h1:YH+1FKiLXxHSkmPseP+kNlulaMuP3n2brvKWEqk/Jc4= @@ -332,9 +284,6 @@ github.com/jtolds/gls v4.20.0+incompatible h1:xdiiI2gbIgH/gLH7ADydsJ1uDOEzR8yvV7 github.com/jtolds/gls v4.20.0+incompatible/go.mod h1:QJZ7F/aHp+rZTRtaJ1ow/lLfFfVYBRgL+9YlvaHOwJU= github.com/julienschmidt/httprouter v1.2.0/go.mod h1:SYymIcj16QtmaHHD7aYtjjsJG7VTCxuUUipMqKk8s4w= github.com/julienschmidt/httprouter v1.3.0/go.mod h1:JR6WtHb+2LUe8TCKY3cZOxFyyO8IZAc4RVcycCCAKdM= -github.com/jzelinskie/stringz v0.0.3 h1:0GhG3lVMYrYtIvRbxvQI6zqRTT1P1xyQlpa0FhfUXas= -github.com/jzelinskie/stringz v0.0.3/go.mod h1:hHYbgxJuNLRw91CmpuFsYEOyQqpDVFg8pvEh23vy4P0= -github.com/kisielk/errcheck v1.5.0/go.mod h1:pFxgyoBC7bSaBwPgfKdkLd5X25qrDl4LWUI2bnpBCr8= github.com/kisielk/gotool v1.0.0/go.mod h1:XhKaO+MFFWcvkIS/tQcRk01m1F5IRFswLeQ+oQHNcck= github.com/klauspost/compress v1.18.6 h1:2jupLlAwFm95+YDR+NwD2MEfFO9d4z4Prjl1XXDjuao= github.com/klauspost/compress v1.18.6/go.mod h1:cwPg85FWrGar70rWktvGQj8/hthj3wpl0PGDogxkrSQ= @@ -430,7 +379,6 @@ github.com/opencontainers/image-spec v1.1.1 h1:y0fUlFfIZhPF1W537XOLg0/fcx6zcHCJw github.com/opencontainers/image-spec v1.1.1/go.mod h1:qpqAh3Dmcf36wStyyWU+kCeDgrGnAve2nCC8+7h8Q0M= github.com/opencontainers/runc v1.3.6 h1:SLGIymCtsk80iNPWgbc8dtjI30r+5mTVV+4dN8/17Sk= github.com/opencontainers/runc v1.3.6/go.mod h1:o1wyv76EDlTkcf0KTFgN8bMWLPvgF/HfX709lDv+rr4= -github.com/opentracing/opentracing-go v1.1.0/go.mod h1:UkNAQd3GIcIGf0SeVgPpRdFStlNbqXla1AfSYxPUl2o= github.com/ory/dockertest/v3 v3.12.0 h1:3oV9d0sDzlSQfHtIaB5k6ghUCVMVLpAY8hwrqoCyRCw= github.com/ory/dockertest/v3 v3.12.0/go.mod h1:aKNDTva3cp8dwOWwb9cWuX84aH5akkxXRvO7KCwWVjE= github.com/pborman/getopt v0.0.0-20170112200414-7148bc3a4c30/go.mod h1:85jBQOZwpVEaDAr341tbn15RS4fCAsIst0qp7i8ex1o= @@ -447,8 +395,6 @@ github.com/pion/transport/v4 v4.0.2/go.mod h1:06hFI+jCFcok2X2MekVufNZ/uzNZXivGBP github.com/pkg/errors v0.8.0/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0= github.com/pkg/errors v0.8.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0= github.com/pkg/errors v0.9.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0= -github.com/planetscale/vtprotobuf v0.6.1-0.20240917153116-6f2963f01587 h1:xzZOeCMQLA/W198ZkdVdt4EKFKJtS26B773zNU377ZY= -github.com/planetscale/vtprotobuf v0.6.1-0.20240917153116-6f2963f01587/go.mod h1:t/avpk3KcrXxUnYOhZhMXJlSEyie6gQbtLq5NM3loB8= github.com/plgd-dev/go-coap/v3 v3.5.3 h1:0MRTXwIasXmTwqUXJjUjHALl8hQxuLBMr/pr4NTBa6U= github.com/plgd-dev/go-coap/v3 v3.5.3/go.mod h1:kgdxil4mi3Bi9s5av/NbQeVwRJ+8N6zGHFEPy7qTRWI= github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= @@ -464,7 +410,6 @@ github.com/prometheus/client_golang v1.23.2 h1:Je96obch5RDVy3FDMndoUsjAhG5Edi49h github.com/prometheus/client_golang v1.23.2/go.mod h1:Tb1a6LWHB3/SPIzCoaDXI4I8UHKeFTEQ1YCr+0Gyqmg= github.com/prometheus/client_model v0.0.0-20180712105110-5c3871d89910/go.mod h1:MbSGuTsp3dbXC40dX6PRTWyKYBIrTGTE9sqQNg2J8bo= github.com/prometheus/client_model v0.0.0-20190129233127-fd36f4220a90/go.mod h1:xMI15A0UPsDsEKsMN9yxemIoYk6Tm2C1GtYGdfGttqA= -github.com/prometheus/client_model v0.0.0-20190812154241-14fe0d1b01d4/go.mod h1:xMI15A0UPsDsEKsMN9yxemIoYk6Tm2C1GtYGdfGttqA= github.com/prometheus/client_model v0.2.0/go.mod h1:xMI15A0UPsDsEKsMN9yxemIoYk6Tm2C1GtYGdfGttqA= github.com/prometheus/client_model v0.6.2 h1:oBsgwpGs7iVziMvrGhE53c/GrLUsZdHnqNwqPLxwZyk= github.com/prometheus/client_model v0.6.2/go.mod h1:y3m2F6Gdpfy6Ut/GBsUqTWZqCUvMVzSfMLjcu6wAwpE= @@ -479,20 +424,16 @@ github.com/prometheus/procfs v0.1.3/go.mod h1:lV6e/gmhEcM9IjHGsFOCxxuZ+z1YqCvr4O github.com/prometheus/procfs v0.6.0/go.mod h1:cz+aTbrPOrUb4q7XlbU9ygM+/jj0fzG6c1xBZuNvfVA= github.com/prometheus/procfs v0.20.1 h1:XwbrGOIplXW/AU3YhIhLODXMJYyC1isLFfYCsTEycfc= github.com/prometheus/procfs v0.20.1/go.mod h1:o9EMBZGRyvDrSPH1RqdxhojkuXstoe4UlK79eF5TGGo= -github.com/rabbitmq/amqp091-go v1.12.0 h1:V0v14Iqfs+MwHWihJt/nGS5Ulu0vw572b2Co3mwunkI= -github.com/rabbitmq/amqp091-go v1.12.0/go.mod h1:Hy4jKW5kQART1u+JkDTF9YYOQUHXqMuhrgxOEeS7G4o= +github.com/rabbitmq/amqp091-go v1.11.0 h1:HxIctVm9Gid/Vtn706necmZ7Wj6pgGI2eqplRbEY8O8= +github.com/rabbitmq/amqp091-go v1.11.0/go.mod h1:Hy4jKW5kQART1u+JkDTF9YYOQUHXqMuhrgxOEeS7G4o= github.com/redis/go-redis/v9 v9.21.0 h1:FPBE4hhbAke+TLmcY3WkpbDffJEomdqPn3HYiqAtL9E= github.com/redis/go-redis/v9 v9.21.0/go.mod h1:v/M13XI1PVCDcm01VtPFOADfZtHf8YW3baQf57KlIkA= -github.com/rodaine/protogofakeit v0.1.1 h1:ZKouljuRM3A+TArppfBqnH8tGZHOwM/pjvtXe9DaXH8= -github.com/rodaine/protogofakeit v0.1.1/go.mod h1:pXn/AstBYMaSfc1/RqH3N82pBuxtWgejz1AlYpY1mI0= github.com/rogpeppe/go-internal v1.3.0/go.mod h1:M8bDsm7K2OlrFYOpmOWEs/qY81heoFRclV5y23lUDJ4= github.com/rogpeppe/go-internal v1.14.1 h1:UQB4HGPB6osV0SQTLymcB4TgvyWu6ZyliaW0tI/otEQ= github.com/rogpeppe/go-internal v1.14.1/go.mod h1:MaRKkUm5W0goXpeCfT7UZI6fk/L7L7so1lCWt35ZSgc= github.com/rs/xid v1.2.1/go.mod h1:+uKXf+4Djp6Md1KODXJxgGQPKngRmWyn10oCKFzNHOQ= github.com/rs/zerolog v1.13.0/go.mod h1:YbFCdg8HfsridGWAh22vktObvhZbQsZXe4/zB0OKkWU= github.com/rs/zerolog v1.15.0/go.mod h1:xYTKnLHcpfU2225ny5qZjxnj9NvkumZYjJHlAThCjNc= -github.com/rs/zerolog v1.35.1 h1:m7xQeoiLIiV0BCEY4Hs+j2NG4Gp2o2KPKmhnnLiazKI= -github.com/rs/zerolog v1.35.1/go.mod h1:EjML9kdfa/RMA7h/6z6pYmq1ykOuA8/mjWaEvGI+jcw= github.com/rubenv/sql-migrate v1.8.1 h1:EPNwCvjAowHI3TnZ+4fQu3a915OpnQoPAjTXCGOy2U0= github.com/rubenv/sql-migrate v1.8.1/go.mod h1:BTIKBORjzyxZDS6dzoiw6eAFYJ1iNlGAtjn4LGeVjS8= github.com/russross/blackfriday/v2 v2.0.1/go.mod h1:+Rmxgy9KzJVeS9/2gXHxylqXiyQDYRxCVz55jmeOWTM= @@ -534,8 +475,6 @@ github.com/spf13/viper v1.21.0 h1:x5S+0EU27Lbphp4UKm1C+1oQO+rKx36vfCoaVebLFSU= github.com/spf13/viper v1.21.0/go.mod h1:P0lhsswPGWD/1lZJ9ny3fYnVqxiegrlNrEmgLjbTCAY= github.com/sqids/sqids-go v0.4.1 h1:eQKYzmAZbLlRwHeHYPF35QhgxwZHLnlmVj9AkIj/rrw= github.com/sqids/sqids-go v0.4.1/go.mod h1:EMwHuPQgSNFS0A49jESTfIQS+066XQTVhukrzEPScl8= -github.com/stoewer/go-strcase v1.3.1 h1:iS0MdW+kVTxgMoE1LAZyMiYJFKlOzLooE4MxjirtkAs= -github.com/stoewer/go-strcase v1.3.1/go.mod h1:fAH5hQ5pehh+j3nZfvwdk2RgEgQjAoM8wodgtPmh1xo= github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= github.com/stretchr/objx v0.1.1/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= github.com/stretchr/objx v0.2.0/go.mod h1:qt09Ya8vawLte6SNmTgCsAVtYtaKzEcn8ATUoHMkEqE= @@ -576,8 +515,6 @@ github.com/xo/terminfo v0.0.0-20220910002029-abceb7e1c41e h1:JVG44RsyaB9T2KIHavM github.com/xo/terminfo v0.0.0-20220910002029-abceb7e1c41e/go.mod h1:RbqR21r5mrJuqunuUZ/Dhy/avygyECGrLceyNeo4LiM= github.com/yuin/gluamapper v0.0.0-20150323120927-d836955830e7 h1:noHsffKZsNfU38DwcXWEPldrTjIZ8FPNKx8mYMGnqjs= github.com/yuin/gluamapper v0.0.0-20150323120927-d836955830e7/go.mod h1:bbMEM6aU1WDF1ErA5YJ0p91652pGv140gGw4Ww3RGp8= -github.com/yuin/goldmark v1.1.27/go.mod h1:3hX8gzYuyVAZsxl0MRgGTJEmQBFcNTphYh9decYSb74= -github.com/yuin/goldmark v1.2.1/go.mod h1:3hX8gzYuyVAZsxl0MRgGTJEmQBFcNTphYh9decYSb74= github.com/yuin/goldmark v1.4.13/go.mod h1:6yULJ656Px+3vBD8DxQVa3kxgyrAnzto9xy5taEt/CY= github.com/yuin/gopher-lua v0.0.0-20200816102855-ee81675732da/go.mod h1:E1AXubJBdNmFERAOucpDIxNzeGfLzg0mYh+UfMWdChA= github.com/yuin/gopher-lua v1.1.2 h1:yF/FjE3hD65tBbt0VXLE13HWS9h34fdzJmrWRXwobGA= @@ -611,21 +548,17 @@ go.uber.org/atomic v1.3.2/go.mod h1:gD2HeocX3+yG+ygLZcrzQJaqmWj9AIm7n08wl/qW/PE= go.uber.org/atomic v1.4.0/go.mod h1:gD2HeocX3+yG+ygLZcrzQJaqmWj9AIm7n08wl/qW/PE= go.uber.org/atomic v1.5.0/go.mod h1:sABNBOSYdrvTF6hTgEIbc7YasKWGhgEQZyfxyTvoXHQ= go.uber.org/atomic v1.6.0/go.mod h1:sABNBOSYdrvTF6hTgEIbc7YasKWGhgEQZyfxyTvoXHQ= -go.uber.org/atomic v1.7.0/go.mod h1:fEN4uk6kAWBTFdckzkM89CLk9XfWZrxpCo0nPH17wJc= go.uber.org/atomic v1.11.0 h1:ZvwS0R+56ePWxUNi+Atn9dWONBPp/AUETXlHW0DxSjE= go.uber.org/atomic v1.11.0/go.mod h1:LUxbIzbOniOlMKjJjyPfpl4v+PKK2cNJn91OQbhoJI0= -go.uber.org/goleak v1.1.10/go.mod h1:8a7PlsEVH3e/a/GLqe5IIrQx6GzcnRmZEufDUTk4A7A= go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto= go.uber.org/goleak v1.3.0/go.mod h1:CoHD4mav9JJNrW/WLlf7HGZPjdw8EucARQHekz1X6bE= go.uber.org/multierr v1.1.0/go.mod h1:wR5kodmAFQ0UK8QlbwjlSNy0Z68gJhDJUG5sjR94q/0= go.uber.org/multierr v1.3.0/go.mod h1:VgVr7evmIr6uPjLBxg28wmKNXyqE9akIJ5XnfpiKl+4= go.uber.org/multierr v1.5.0/go.mod h1:FeouvMocqHpRaaGuG9EjoKcStLC43Zu/fmqdUMPcKYU= -go.uber.org/multierr v1.6.0/go.mod h1:cdWPpRnG4AhwMwsgIHip0KRBQjJy5kYEpYjJxpXp9iU= go.uber.org/tools v0.0.0-20190618225709-2cfd321de3ee/go.mod h1:vJERXedbb3MVM5f9Ejo0C68/HhF8uaILCdgjnY+goOA= go.uber.org/zap v1.9.1/go.mod h1:vwi/ZaCAaUcBkycHslxD9B2zi4UTXhF60s6SWpuDF0Q= go.uber.org/zap v1.10.0/go.mod h1:vwi/ZaCAaUcBkycHslxD9B2zi4UTXhF60s6SWpuDF0Q= go.uber.org/zap v1.13.0/go.mod h1:zwrFLgMcdUuIBviXEYEH1YKNaOBnKXsx2IPda5bBwHM= -go.uber.org/zap v1.18.1/go.mod h1:xg/QME4nWcxGxrpdeYfq7UvYrLh66cuVKdrbD1XF/NI= go.yaml.in/yaml/v2 v2.4.4 h1:tuyd0P+2Ont/d6e2rl3be67goVK4R6deVxCUX5vyPaQ= go.yaml.in/yaml/v2 v2.4.4/go.mod h1:gMZqIpDtDqOfM0uNfy0SkpRhvUryYH0Z6wdMYcacYXQ= go.yaml.in/yaml/v3 v3.0.4 h1:tfq32ie2Jv2UxXFdLJdh3jXuOzWiL1fo0bu/FbuKpbc= @@ -645,33 +578,23 @@ golang.org/x/crypto v0.19.0/go.mod h1:Iy9bg/ha4yyC70EfRS8jz+B6ybOBKMaSxLj6P6oBDf golang.org/x/crypto v0.20.0/go.mod h1:Xwo95rrVNIoSMx9wa1JroENMToLWn3RNVrTBpLHgZPQ= golang.org/x/crypto v0.53.0 h1:QZ4Muo8THX6CizN2vPPd5fBGHyogrdK9fG4wLPFUsto= golang.org/x/crypto v0.53.0/go.mod h1:DNLU434OwVakk9PzuwV8w62mAJpRJL3vsgcfp4Qnsio= -golang.org/x/exp v0.0.0-20190121172915-509febef88a4/go.mod h1:CJ0aWSM057203Lf6IL+f9T1iT9GByDxfZKAQTCR3kQA= golang.org/x/exp v0.0.0-20260611194520-c48552f49976 h1:X8Hz2ImujgbmetVuW+w2YkyZChE3cBpZi2P158rTG9M= golang.org/x/exp v0.0.0-20260611194520-c48552f49976/go.mod h1:vnf4pv9iKZXY58sQE1L86zmNWJ4159e1RkcWiLCkeEY= -golang.org/x/lint v0.0.0-20181026193005-c67002cb31c3/go.mod h1:UVdnD1Gm6xHRNCYTkRU2/jEulfH38KcIWyp/GAMgvoE= -golang.org/x/lint v0.0.0-20190227174305-5b3e6a55c961/go.mod h1:wehouNa3lNwaWXcvxsM5YxQ5yQlVC4a0KAMCusXpPoU= -golang.org/x/lint v0.0.0-20190313153728-d0100b6bd8b3/go.mod h1:6SW0HCj/g11FgYtHlgUYUwCkIfeOF89ocIRzGO/8vkc= golang.org/x/lint v0.0.0-20190930215403-16217165b5de/go.mod h1:6SW0HCj/g11FgYtHlgUYUwCkIfeOF89ocIRzGO/8vkc= golang.org/x/mod v0.0.0-20190513183733-4bf6d317e70e/go.mod h1:mXi4GBBbnImb6dmsKGUJ2LatrhH/nqhxcFungHvyanc= golang.org/x/mod v0.1.1-0.20191105210325-c90efee705ee/go.mod h1:QqPTAvyqsEbceGzBzNggFXnrqF1CaUcvgkdR5Ot7KZg= -golang.org/x/mod v0.2.0/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA= -golang.org/x/mod v0.3.0/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA= golang.org/x/mod v0.6.0-dev.0.20220419223038-86c51ed26bb4/go.mod h1:jJ57K6gSWd91VN4djpZkiMVwK6gcyfeH4XE8wZrZaV4= golang.org/x/mod v0.8.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs= golang.org/x/net v0.0.0-20180724234803-3673e40ba225/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4= -golang.org/x/net v0.0.0-20180826012351-8a410e7b638d/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4= golang.org/x/net v0.0.0-20181114220301-adae6a3d119a/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4= golang.org/x/net v0.0.0-20190108225652-1e06a53dbb7e/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4= -golang.org/x/net v0.0.0-20190213061140-3a22650c66bd/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4= golang.org/x/net v0.0.0-20190311183353-d8887717615a/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg= golang.org/x/net v0.0.0-20190404232315-eb5bcb51f2a3/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg= golang.org/x/net v0.0.0-20190613194153-d28f0bde5980/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s= golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s= golang.org/x/net v0.0.0-20190813141303-74dc4d7220e7/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s= golang.org/x/net v0.0.0-20200202094626-16171245cfb2/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s= -golang.org/x/net v0.0.0-20200226121028-0de0cce0169b/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s= golang.org/x/net v0.0.0-20200625001655-4c5254603344/go.mod h1:/O7V0waA8r7cgGh81Ro3o1hOxt32SMVPicZroKQ2sZA= -golang.org/x/net v0.0.0-20201021035429-f5854403a974/go.mod h1:sp8m0HH+o8qH0wwXwYZr8TS3Oi6o0r6Gce1SSxlDquU= golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg= golang.org/x/net v0.0.0-20220722155237-a158d28d115b/go.mod h1:XRhObCWvk6IyKnWLug+ECip1KBveYUHfp+8e9klMJ9c= golang.org/x/net v0.6.0/go.mod h1:2Tu9+aMcznHK/AK1HMvgo6xiTLG5rD5rZLDS+rp2Bjs= @@ -679,22 +602,16 @@ golang.org/x/net v0.10.0/go.mod h1:0qNGK6F8kojg2nk9dLZ2mShWaEBan6FAoqfSigmmuDg= golang.org/x/net v0.21.0/go.mod h1:bIjVDfnllIU7BJ2DNgfnXvpSvtn8VRwhlsaeUTyUS44= golang.org/x/net v0.56.0 h1:Rw8j/hFzGvJUZwNBXnAtf5sVDVt+65SK2C7IxCxZt5o= golang.org/x/net v0.56.0/go.mod h1:D3Ku6r+V6JROoZK144D2XfMHFcMq/0zSfLelVTCFKec= -golang.org/x/oauth2 v0.0.0-20180821212333-d2e6202438be/go.mod h1:N/0e6XlmueqKjAGxoOufVs8QHGRruUQn6yWY3a++T0U= golang.org/x/oauth2 v0.0.0-20190226205417-e64efc72b421/go.mod h1:gOpvHmFTYa4IltrdGE7lF6nIHvwfUNPOp7c8zoXwtLw= -golang.org/x/oauth2 v0.36.0 h1:peZ/1z27fi9hUOFCAZaHyrpWG5lwe0RJEEEeH0ThlIs= -golang.org/x/oauth2 v0.36.0/go.mod h1:YDBUJMTkDnJS+A4BP4eZBjCqtokkg1hODuPjwiGPO7Q= -golang.org/x/sync v0.0.0-20180314180146-1d60e4601c6f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.0.0-20181108010431-42b317875d0f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.0.0-20181221193216-37e7f081c4d4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.0.0-20190911185100-cd5d95a43a6e/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= -golang.org/x/sync v0.0.0-20201020160332-67f06af15bc9/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.0.0-20201207232520-09787c993a3a/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.0.0-20220722155255-886fb9371eb4/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.1.0/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.21.0 h1:HLII4xRRTtCRkxYp4HNFF0Js/Og6q2i++KXbg0gHCwM= golang.org/x/sync v0.21.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= -golang.org/x/sys v0.0.0-20180830151530-49385e6e1522/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= golang.org/x/sys v0.0.0-20180905080454-ebe1bf3edb33/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= golang.org/x/sys v0.0.0-20181116152217-5ac8a444bdc5/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= golang.org/x/sys v0.0.0-20190204203706-41f3e6584952/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= @@ -711,14 +628,12 @@ golang.org/x/sys v0.0.0-20200223170610-d5e6a3e2c0ae/go.mod h1:h1NjWce9XRLGQEsW7w golang.org/x/sys v0.0.0-20200323222414-85ca7c5b95cd/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/sys v0.0.0-20200615200032-f1bc736245b1/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/sys v0.0.0-20200625212154-ddb9806d33ae/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= -golang.org/x/sys v0.0.0-20200930185726-fdedc70b468f/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/sys v0.0.0-20210124154548-22da62e12c0c/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/sys v0.0.0-20210603081109-ebe580a85c40/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.0.0-20210615035016-665e8c7367d1/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.0.0-20210616094352-59db8d763f22/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.0.0-20210630005230-0f9fa26af87c/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= -golang.org/x/sys v0.0.0-20211025201205-69cdffdb9359/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.0.0-20220328115105-d36c6a25d886/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.0.0-20220520151302-bc2c85ada10a/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.0.0-20220722155257-8c9f86f7a55f/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= @@ -748,20 +663,14 @@ golang.org/x/time v0.0.0-20210220033141-f8bda1e9f3ba/go.mod h1:tRJNPiyCQ0inRvYxb golang.org/x/time v0.15.0 h1:bbrp8t3bGUeFOx08pvsMYRTCVSMk89u4tKbNOZbp88U= golang.org/x/time v0.15.0/go.mod h1:Y4YMaQmXwGQZoFaVFk4YpCt4FLQMYKZe9oeV/f4MSno= golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ= -golang.org/x/tools v0.0.0-20190114222345-bf090417da8b/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ= -golang.org/x/tools v0.0.0-20190226205152-f727befe758c/go.mod h1:9Yl7xja0Znq3iFh3HoIrodX9oNMXvdceNzlUR8zjMvY= golang.org/x/tools v0.0.0-20190311212946-11955173bddd/go.mod h1:LCzVGOaR6xXOjkQ3onu1FJEFr0SW1gC7cKk1uF8kGRs= golang.org/x/tools v0.0.0-20190425163242-31fd60d6bfdc/go.mod h1:RgjU9mgBXZiqYHBnxXauZ1Gv1EHHAz9KjViQ78xBX0Q= -golang.org/x/tools v0.0.0-20190524140312-2c0ae7006135/go.mod h1:RgjU9mgBXZiqYHBnxXauZ1Gv1EHHAz9KjViQ78xBX0Q= golang.org/x/tools v0.0.0-20190621195816-6e04913cbbac/go.mod h1:/rFqwRUd4F7ZHNgwSSTFct+R/Kf4OFW1sUzUTQQTgfc= golang.org/x/tools v0.0.0-20190823170909-c4a336ef6a2f/go.mod h1:b+2E5dAYhXwXZwtnZ6UAqBI28+e2cm9otk0dWdXHAEo= golang.org/x/tools v0.0.0-20191029041327-9cc4af7d6b2c/go.mod h1:b+2E5dAYhXwXZwtnZ6UAqBI28+e2cm9otk0dWdXHAEo= golang.org/x/tools v0.0.0-20191029190741-b9c20aec41a5/go.mod h1:b+2E5dAYhXwXZwtnZ6UAqBI28+e2cm9otk0dWdXHAEo= -golang.org/x/tools v0.0.0-20191108193012-7d206e10da11/go.mod h1:b+2E5dAYhXwXZwtnZ6UAqBI28+e2cm9otk0dWdXHAEo= golang.org/x/tools v0.0.0-20191119224855-298f0cb1881e/go.mod h1:b+2E5dAYhXwXZwtnZ6UAqBI28+e2cm9otk0dWdXHAEo= golang.org/x/tools v0.0.0-20200103221440-774c71fcf114/go.mod h1:TB2adYChydJhpapKDTa4BR/hXlZSLoq2Wpct/0txZ28= -golang.org/x/tools v0.0.0-20200619180055-7c47624df98f/go.mod h1:EkVYQZoAsY45+roYkvgYkIh4xh/qjgUK9TdY2XT94GE= -golang.org/x/tools v0.0.0-20210106214847-113979e3529a/go.mod h1:emZCQorbCU4vsT4fOWvOPXz4eW1wZW4PmDk9uLelYpA= golang.org/x/tools v0.1.12/go.mod h1:hNGJHUnrk76NpqgfD5Aqm5Crs+Hm0VOH/i9J2+nxYbc= golang.org/x/tools v0.6.0/go.mod h1:Xwgl3UAJ/d3gWutnCtw505GrjyAbvKui8lOU390QaIU= golang.org/x/xerrors v0.0.0-20190410155217-1f06c39b4373/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= @@ -772,20 +681,11 @@ golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8T golang.org/x/xerrors v0.0.0-20200804184101-5ec99f83aff1/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= gonum.org/v1/gonum v0.17.0 h1:VbpOemQlsSMrYmn7T2OUvQ4dqxQXU+ouZFQsZOx50z4= gonum.org/v1/gonum v0.17.0/go.mod h1:El3tOrEuMpv2UdMrbNlKEh9vd86bmQ6vqIcDwxEOc1E= -google.golang.org/appengine v1.1.0/go.mod h1:EbEs0AVv82hx2wNQdGPgUI5lhzA/G0D9YwlJXL52JkM= google.golang.org/appengine v1.4.0/go.mod h1:xpcJRLb0r/rnEns0DIKYYv+WjYCduHsrkT7/EB5XEv4= -google.golang.org/genproto v0.0.0-20180817151627-c66870c02cf8/go.mod h1:JiN7NxoALGmiZfu7CAH4rXhgtRTLTxftemlI0sWmxmc= -google.golang.org/genproto v0.0.0-20190819201941-24fa4b261c55/go.mod h1:DMBHOl98Agz4BDEuKkezgsaosCRResVns1a3J2ZsMNc= -google.golang.org/genproto v0.0.0-20200423170343-7949de9c1215/go.mod h1:55QSHmfGQM9UVYDPBsyGGes0y52j32PQ3BqQfXhyH3c= google.golang.org/genproto/googleapis/api v0.0.0-20260610212136-7ab31c22f7ad h1:3iLyITS/sySRwbUKoC7ogfj2Yr1Cjs0pfaRKj5U5HEw= google.golang.org/genproto/googleapis/api v0.0.0-20260610212136-7ab31c22f7ad/go.mod h1:KdNqO+rCIWgFumrNBSEDlDNrkrQnpkax7Tv1WxNY8V4= google.golang.org/genproto/googleapis/rpc v0.0.0-20260610212136-7ab31c22f7ad h1:45WmJvIV6C2+O/jjLkPUH+F3aOj/1miDoU2DD0+NWbg= google.golang.org/genproto/googleapis/rpc v0.0.0-20260610212136-7ab31c22f7ad/go.mod h1:4Hqkh8ycfw05ld/3BWL7rJOSfebL2Q+DVDeRgYgxUU8= -google.golang.org/grpc v1.19.0/go.mod h1:mqu4LbDTu4XGKhr4mRzUsmM4RtVoemTSY81AxZiDr8c= -google.golang.org/grpc v1.23.0/go.mod h1:Y5yQAOtifL1yxbo5wqy6BxZv8vAUGQwXBOALyacEbxg= -google.golang.org/grpc v1.25.1/go.mod h1:c3i+UQWmh7LiEpx4sFZnkU36qjEYZ0imhYfXVyQciAY= -google.golang.org/grpc v1.27.0/go.mod h1:qbnxyOmOxrQa7FizSgH+ReBfzJrCY1pSN7KXBS8abTk= -google.golang.org/grpc v1.29.1/go.mod h1:itym6AZVZYACWQqET3MqgPpjcuV5QH3BxFS3IjizoKk= google.golang.org/grpc v1.81.1 h1:VnnIIZ88UzOOKLukQi+ImGz8O1Wdp8nAGGnvOfEIWQQ= google.golang.org/grpc v1.81.1/go.mod h1:xGH9GfzOyMTGIOXBJmXt+BX/V0kcdQbdcuwQ/zNw42I= google.golang.org/protobuf v0.0.0-20200109180630-ec00e32a8dfd/go.mod h1:DFci5gLYBciE7Vtevhsrf46CRTquxDuWsQurQQe4oz8= @@ -815,18 +715,14 @@ gopkg.in/yaml.v2 v2.2.1/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI= gopkg.in/yaml.v2 v2.2.2/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI= gopkg.in/yaml.v2 v2.2.4/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI= gopkg.in/yaml.v2 v2.2.5/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI= -gopkg.in/yaml.v2 v2.2.8/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI= gopkg.in/yaml.v2 v2.3.0/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI= gopkg.in/yaml.v2 v2.4.0 h1:D8xgwECY7CYvx+Y2n4sBz93Jn9JRvxdiyyo8CTfuKaY= gopkg.in/yaml.v2 v2.4.0/go.mod h1:RDklbk79AGWmwhnvt/jBztapEOGDOx6ZbXqjP6csGnQ= gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= -gopkg.in/yaml.v3 v3.0.0-20210107192922-496545a6307b/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= gotest.tools/v3 v3.5.2 h1:7koQfIKdy+I8UTetycgUqXWSDwpgv193Ka+qRsmBY8Q= gotest.tools/v3 v3.5.2/go.mod h1:LtdLGcnqToBH83WByAAi/wiwSFCArdFIUV/xxN4pcjA= -honnef.co/go/tools v0.0.0-20190102054323-c2f93a96b099/go.mod h1:rf3lG4BRIbNafJWhAfAdb/ePZxsR/4RtNHQocxwk9r4= -honnef.co/go/tools v0.0.0-20190523083050-ea95bdfd59fc/go.mod h1:rf3lG4BRIbNafJWhAfAdb/ePZxsR/4RtNHQocxwk9r4= honnef.co/go/tools v0.0.1-2019.2.3/go.mod h1:a3bituU0lyd329TUQxRnasdCoJDkEUEAqEt0JzvZhAg= moul.io/http2curl v1.0.0 h1:6XwpyZOYsgZJrU8exnG87ncVkU1FVCcTRpwzOkTDUi8= moul.io/http2curl v1.0.0/go.mod h1:f6cULg+e4Md/oW1cYmwW4IWQOVl2lGbmCNGOHvzX2kE= diff --git a/groups/README.md b/groups/README.md deleted file mode 100644 index 001e5e835..000000000 --- a/groups/README.md +++ /dev/null @@ -1,307 +0,0 @@ -# Groups - -The Groups service exposes HTTP and gRPC APIs for organizing entities into hierarchical groups within a domain, managing membership, permissions, and roles. It handles group lifecycle (create/update/enable/disable/delete), parent/child relationships, listings (flat or tree), and role-based access. - -For a deeper overview of Magistrala, see the [official documentation][doc]. - -## Configuration - -The service is configured via environment variables (unset values fall back to defaults). - -| Variable | Description | Default | -| ------------------------------------ | ------------------------------------------------------------------------ | ------------------------------ | -| `MG_GROUPS_LOG_LEVEL` | Log level for Groups (debug, info, warn, error) | debug | -| `MG_GROUPS_HTTP_HOST` | Groups service HTTP host | groups | -| `MG_GROUPS_HTTP_PORT` | Groups service HTTP port | 9004 | -| `MG_GROUPS_HTTP_SERVER_CERT` | Path to PEM-encoded HTTP server certificate | "" | -| `MG_GROUPS_HTTP_SERVER_KEY` | Path to PEM-encoded HTTP server key | "" | -| `MG_GROUPS_HTTP_SERVER_CA_CERTS` | Path to trusted CA bundle for the HTTP server | "" | -| `MG_GROUPS_HTTP_CLIENT_CA_CERTS` | Path to client CA bundle to require HTTP mTLS | "" | -| `MG_GROUPS_GRPC_HOST` | Groups service gRPC host | groups | -| `MG_GROUPS_GRPC_PORT` | Groups service gRPC port | 7004 | -| `MG_GROUPS_GRPC_SERVER_CERT` | Path to PEM-encoded gRPC server certificate | "" | -| `MG_GROUPS_GRPC_SERVER_KEY` | Path to PEM-encoded gRPC server key | "" | -| `MG_GROUPS_GRPC_SERVER_CA_CERTS` | Path to trusted CA bundle for the gRPC server | "" | -| `MG_GROUPS_GRPC_CLIENT_CA_CERTS` | Path to client CA bundle to require gRPC mTLS | "" | -| `MG_GROUPS_DB_HOST` | Database host address | groups-db | -| `MG_GROUPS_DB_PORT` | Database host port | 5432 | -| `MG_GROUPS_DB_USER` | Database user | magistrala | -| `MG_GROUPS_DB_PASS` | Database password | magistrala | -| `MG_GROUPS_DB_NAME` | Name of the database used by the service | groups | -| `MG_GROUPS_DB_SSL_MODE` | Database connection SSL mode (disable, require, verify-ca, verify-full) | disable | -| `MG_GROUPS_DB_SSL_CERT` | Path to the PEM-encoded certificate file | "" | -| `MG_GROUPS_DB_SSL_KEY` | Path to the PEM-encoded key file | "" | -| `MG_GROUPS_DB_SSL_ROOT_CERT` | Path to the PEM-encoded root certificate file | "" | -| `MG_GROUPS_INSTANCE_ID` | Groups instance ID (auto-generated when empty) | "" | -| `MG_GROUPS_EVENT_CONSUMER` | NATS consumer name for domain events | groups | -| `MG_SPICEDB_HOST` | SpiceDB host for policy checks | magistrala-spicedb | -| `MG_SPICEDB_PORT` | SpiceDB port | 50051 | -| `MG_SPICEDB_SCHEMA_FILE` | Path to SpiceDB schema file used to seed available actions | "/schema.zed" | -| `MG_SPICEDB_PRE_SHARED_KEY` | SpiceDB preshared key | 12345678 | -| `MG_ES_URL` | Event store URL | nats://nats:4222 | -| `MG_JAEGER_URL` | Jaeger server URL | | -| `MG_JAEGER_TRACE_RATIO` | Trace sampling ratio | 1.0 | -| `MG_SEND_TELEMETRY` | Send telemetry to the Magistrala call-home server | true | -| `MG_AUTH_GRPC_URL` | Auth service gRPC URL | "" | -| `MG_AUTH_GRPC_TIMEOUT` | Auth service gRPC request timeout | 1s | -| `MG_AUTH_GRPC_CLIENT_CERT` | Path to the PEM-encoded Auth gRPC client certificate | "" | -| `MG_AUTH_GRPC_CLIENT_KEY` | Path to the PEM-encoded Auth gRPC client key | "" | -| `MG_AUTH_GRPC_SERVER_CA_CERTS` | Path to the PEM-encoded Auth gRPC trusted CA bundle | "" | -| `MG_GROUPS_CALLOUT_URLS` | Comma-separated list of HTTP callout targets invoked on group operations | "" | -| `MG_GROUPS_CALLOUT_METHOD` | HTTP method for callouts (POST or GET) | POST | -| `MG_GROUPS_CALLOUT_TLS_VERIFICATION` | Verify TLS certificates for callouts | false | -| `MG_GROUPS_CALLOUT_TIMEOUT` | Callout request timeout | 10s | -| `MG_GROUPS_CALLOUT_CA_CERT` | CA bundle for verifying callout targets | "" | -| `MG_GROUPS_CALLOUT_CERT` | Client certificate for mTLS callouts | "" | -| `MG_GROUPS_CALLOUT_KEY` | Client key for mTLS callouts | "" | -| `MG_GROUPS_CALLOUT_OPERATIONS` | Comma-separated list of operation names that should trigger callouts | "" | - -**Note**: Set `MG_GROUPS_CALLOUT_OPERATIONS` to a subset of `OpCreateGroup`, `OpViewGroup`, `OpUpdateGroup`, `OpEnableGroup`, `OpDisableGroup`, `OpDeleteGroup`, `OpListGroups`, `OpHierarchy`, `OpAddParentGroup`, `OpRemoveParentGroup`, `OpAddChildrenGroups`, `OpRemoveChildrenGroups`, `OpRemoveAllChildrenGroups`, or `OpListChildrenGroups` to filter which actions produce callouts. - -## Deployment - -The service ships as a Docker container. See the [`groups` section](https://github.com/absmach/magistrala/blob/main/docker/docker-compose.yaml#L950-L1035) in `docker-compose.yaml` for deployment configuration. - -To build and run locally: - -```bash -# download the latest version of the service -git clone https://github.com/absmach/magistrala -cd magistrala - -# compile the groups service -make groups - -# copy binary to $GOBIN -make install - -# set the environment variables and run the service -MG_GROUPS_LOG_LEVEL=debug \ -MG_GROUPS_HTTP_HOST=groups \ -MG_GROUPS_HTTP_PORT=9004 \ -MG_GROUPS_HTTP_SERVER_CERT="" \ -MG_GROUPS_HTTP_SERVER_KEY="" \ -MG_GROUPS_GRPC_HOST=groups \ -MG_GROUPS_GRPC_PORT=7004 \ -MG_GROUPS_GRPC_SERVER_CERT="" \ -MG_GROUPS_GRPC_SERVER_KEY="" \ -MG_GROUPS_GRPC_SERVER_CA_CERTS="" \ -MG_GROUPS_GRPC_CLIENT_CA_CERTS="" \ -MG_GROUPS_DB_HOST=groups-db \ -MG_GROUPS_DB_PORT=5432 \ -MG_GROUPS_DB_USER=magistrala \MG_GROUPS_DB_PASS=magistrala \MG_GROUPS_DB_NAME=groups \ -MG_GROUPS_DB_SSL_MODE=disable \ -MG_GROUPS_DB_SSL_CERT="" \ -MG_GROUPS_DB_SSL_KEY="" \ -MG_GROUPS_DB_SSL_ROOT_CERT="" \ -MG_AUTH_GRPC_URL="" \ -MG_AUTH_GRPC_TIMEOUT=1s \ -MG_AUTH_GRPC_CLIENT_CERT="" \ -MG_AUTH_GRPC_CLIENT_KEY="" \ -MG_AUTH_GRPC_SERVER_CA_CERTS="" \ -MG_DOMAINS_GRPC_URL=domains:7003 \ -MG_DOMAINS_GRPC_TIMEOUT=1s \ -MG_DOMAINS_GRPC_CLIENT_CERT="" \ -MG_DOMAINS_GRPC_CLIENT_KEY="" \ -MG_DOMAINS_GRPC_SERVER_CA_CERTS="" \ -MG_CHANNELS_GRPC_URL=channels:7005 \ -MG_CHANNELS_GRPC_TIMEOUT=1s \ -MG_CHANNELS_GRPC_CLIENT_CERT="" \ -MG_CHANNELS_GRPC_CLIENT_KEY="" \ -MG_CHANNELS_GRPC_SERVER_CA_CERTS="" \ -MG_CLIENTS_GRPC_URL=clients:7000 \ -MG_CLIENTS_GRPC_TIMEOUT=1s \ -MG_CLIENTS_GRPC_CLIENT_CERT="" \ -MG_CLIENTS_GRPC_CLIENT_KEY="" \ -MG_CLIENTS_GRPC_SERVER_CA_CERTS="" \ -MG_SPICEDB_HOST=localhost \ -MG_SPICEDB_PORT=50051 \ -MG_SPICEDB_SCHEMA_FILE=schema.zed \ -MG_SPICEDB_PRE_SHARED_KEY=12345678 \ -MG_ES_URL=nats://localhost:4222 \ -MG_JAEGER_URL= \ -MG_JAEGER_TRACE_RATIO=1.0 \ -MG_GROUPS_CALLOUT_URLS="" \ -MG_GROUPS_CALLOUT_METHOD=POST \ -MG_GROUPS_CALLOUT_TLS_VERIFICATION=false \ -MG_GROUPS_CALLOUT_TIMEOUT=10s \ -MG_GROUPS_CALLOUT_CA_CERT="" \ -MG_GROUPS_CALLOUT_CERT="" \ -MG_GROUPS_CALLOUT_KEY="" \ -MG_GROUPS_CALLOUT_OPERATIONS="" \ -MG_SEND_TELEMETRY=true \ -MG_GROUPS_INSTANCE_ID="" \ -$GOBIN/magistrala-groups -``` - -## Usage - -Groups supports the following operations: - -| Operation | Description | -| ---------------------------------- | ----------------------------------------------------------------------- | -| `create` | Create a new group within a domain | -| `list` | List groups (flat list or tree) with filters for metadata, tags, status | -| `get` | Retrieve a single group (optionally with role memberships) | -| `update` | Update a group’s name, description, tags, or metadata | -| `enable` / `disable` | Enable or disable a group | -| `delete` | Permanently delete a group | -| `add-parent` / `remove-parent` | Assign or remove a parent group | -| `add-children` / `remove-children` | Attach or detach child groups (or remove all children) | -| `list-children` | List children at specific depth ranges | -| `hierarchy` | Fetch ancestors/descendants as a tree or list | -| `roles` | Create/list/update/delete group roles; manage role actions and members | - -### API Examples - -#### Create a Group - -```bash -curl -X POST http://localhost:9004//groups \ - -H "Authorization: Bearer " \ - -H "Content-Type: application/json" \ - -d '{ - "name": "edge-devices", - "description": "All edge devices", - "metadata": { "region": "eu-west-1" }, - "tags": ["iot","edge"], - "parent_id": "", - "status": "enabled" - }' -``` - -#### List Groups (flat) - -```bash -curl -X GET "http://localhost:9004//groups?limit=10&status=enabled" \ - -H "Authorization: Bearer " -``` - -#### Retrieve a Group (with Roles) - -```bash -curl -X GET "http://localhost:9004//groups/?roles=true" \ - -H "Authorization: Bearer " -``` - -#### Update a Group - -```bash -curl -X PUT http://localhost:9004//groups/ \ - -H "Authorization: Bearer " \ - -H "Content-Type: application/json" \ - -d '{ - "name": "edge-ops", - "description": "Edge operations", - "metadata": { "region": "eu-west-1", "env": "prod" }, - "tags": ["iot","ops"] - }' -``` - -#### Enable or Disable a Group - -```bash -curl -X POST http://localhost:9004//groups//enable \ - -H "Authorization: Bearer " - -curl -X POST http://localhost:9004//groups//disable \ - -H "Authorization: Bearer " -``` - -#### Manage Hierarchy - -```bash -# Add a parent -curl -X POST http://localhost:9004//groups//parents \ - -H "Authorization: Bearer " \ - -H "Content-Type: application/json" \ - -d '{ "parent_id": "" }' - -# List children between levels 1 and 2 -curl -X GET "http://localhost:9004//groups//children?start_level=1&end_level=2&limit=10" \ - -H "Authorization: Bearer " -``` - -## Roles Management for Groups - -Group roles use the shared role manager. Supported operations mirror domain roles (create, list, view, update, delete roles; add/list/remove actions; add/list/remove members; list available actions). - -Example: create a group role - -```bash -curl -X POST http://localhost:9004//groups//roles \ - -H "Authorization: Bearer " \ - -H "Content-Type: application/json" \ - -d '{ - "role_name": "group-admin", - "optional_actions": ["manage_role_permission", "update_permission"], - "optional_members": [""] - }' -``` - -List available actions for groups: - -```bash -curl -X GET http://localhost:9004//groups/roles/available-actions \ - -H "Authorization: Bearer " -``` - -## Implementation Details - -- Groups are stored in PostgreSQL with `ltree` paths for hierarchy queries; domain migrations are applied alongside group migrations for referential integrity. -- Role tables are provisioned per entity with a `groups_` prefix. -- Event notifications are published to `MG_ES_URL`; domain events are consumed to keep group data aligned. -- Authorization and roles are enforced through SpiceDB and shared policy middleware. -- Optional HTTP callouts (pre-operation hooks) are controlled via `MG_GROUPS_CALLOUT_*`. -- Observability: Jaeger tracing, Prometheus metrics at `/metrics`, and a `/health` endpoint. - -### Groups Table - -| Column | Type | Description | -| ------------- | ------------- | ------------------------------------------------------ | -| `id` | VARCHAR(36) | UUID of the group (primary key) | -| `parent_id` | VARCHAR(36) | Optional parent group (self-referential FK) | -| `domain_id` | VARCHAR(36) | Owning domain | -| `name` | VARCHAR(1024) | Group name | -| `description` | VARCHAR(1024) | Optional description | -| `metadata` | JSONB | Arbitrary metadata | -| `tags` | TEXT[] | Group tags | -| `path` | LTREE | Hierarchical path for fast ancestor/descendant queries | -| `created_at` | TIMESTAMPTZ | Creation timestamp | -| `updated_at` | TIMESTAMPTZ | Last update timestamp | -| `updated_by` | VARCHAR(254) | Actor who last updated the group | -| `status` | SMALLINT | 0 = enabled, 1 = disabled, 2 = deleted | - -## Best Practices - -- Model hierarchy deliberately: keep depth reasonable and avoid cycles by design. -- Use tags/metadata to segment groups by environment, region, or ownership for filtering. -- Prefer `disable` before `delete` when you need reversible off-boarding. -- Use roles sparingly and audit with `list-role-members`; grant only required actions. -- Fetch children with bounded levels to keep queries efficient. -- Limit callouts to necessary operations via `MG_GROUPS_CALLOUT_OPERATIONS`. - -## Versioning and Health Check - -The Groups service exposes `/health` with status and build metadata. - -```bash -curl -X GET http://localhost:9004/health \ - -H "accept: application/health+json" -``` - -Example response: - -```json -{ - "status": "pass", - "version": "0.18.0", - "commit": "7d6f4dc4f7f0c1fa3dc24eddfb18bb5073ff4f62", - "description": "groups service", - "build_time": "1970-01-01_00:00:00" -} -``` - -For full API coverage, see the [Groups API documentation](https://docs.api.magistrala.absmach.eu/?urls.primaryName=api%2Fgroups.yaml). - -[doc]: https://magistrala.absmach.eu/docs/ \ No newline at end of file diff --git a/groups/api/grpc/client.go b/groups/api/grpc/client.go deleted file mode 100644 index 8cb0e83d6..000000000 --- a/groups/api/grpc/client.go +++ /dev/null @@ -1,93 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package grpc - -import ( - "context" - "fmt" - "time" - - grpcCommonV1 "github.com/absmach/magistrala/api/grpc/common/v1" - grpcGroupsV1 "github.com/absmach/magistrala/api/grpc/groups/v1" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/go-kit/kit/endpoint" - kitgrpc "github.com/go-kit/kit/transport/grpc" - "google.golang.org/grpc" - "google.golang.org/grpc/codes" - "google.golang.org/grpc/status" -) - -const svcName = "groups.v1.GroupsService" - -var _ grpcGroupsV1.GroupsServiceClient = (*grpcClient)(nil) - -type grpcClient struct { - timeout time.Duration - retrieveEntity endpoint.Endpoint -} - -// NewClient returns new gRPC client instance. -func NewClient(conn *grpc.ClientConn, timeout time.Duration) grpcGroupsV1.GroupsServiceClient { - return &grpcClient{ - retrieveEntity: kitgrpc.NewClient( - conn, - svcName, - "RetrieveEntity", - encodeRetrieveEntityRequest, - decodeRetrieveEntityResponse, - grpcCommonV1.RetrieveEntityRes{}, - ).Endpoint(), - - timeout: timeout, - } -} - -func (client grpcClient) RetrieveEntity(ctx context.Context, req *grpcCommonV1.RetrieveEntityReq, _ ...grpc.CallOption) (r *grpcCommonV1.RetrieveEntityRes, err error) { - ctx, cancel := context.WithTimeout(ctx, client.timeout) - defer cancel() - - res, err := client.retrieveEntity(ctx, req) - if err != nil { - return &grpcCommonV1.RetrieveEntityRes{}, decodeError(err) - } - typedRes := res.(*grpcCommonV1.RetrieveEntityRes) - - return typedRes, nil -} - -func encodeRetrieveEntityRequest(_ context.Context, grpcReq any) (any, error) { - return grpcReq, nil -} - -func decodeRetrieveEntityResponse(_ context.Context, grpcRes any) (any, error) { - return grpcRes, nil -} - -func decodeError(err error) error { - if st, ok := status.FromError(err); ok { - switch st.Code() { - case codes.Unauthenticated: - return errors.Wrap(svcerr.ErrAuthentication, errors.New(st.Message())) - case codes.PermissionDenied: - return errors.Wrap(svcerr.ErrAuthorization, errors.New(st.Message())) - case codes.InvalidArgument: - return errors.Wrap(errors.ErrMalformedEntity, errors.New(st.Message())) - case codes.FailedPrecondition: - return errors.Wrap(errors.ErrMalformedEntity, errors.New(st.Message())) - case codes.NotFound: - return errors.Wrap(svcerr.ErrNotFound, errors.New(st.Message())) - case codes.AlreadyExists: - return errors.Wrap(svcerr.ErrConflict, errors.New(st.Message())) - case codes.OK: - if msg := st.Message(); msg != "" { - return errors.Wrap(errors.ErrUnidentified, errors.New(msg)) - } - return nil - default: - return errors.Wrap(fmt.Errorf("unexpected gRPC status: %s (status code:%v)", st.Code().String(), st.Code()), errors.New(st.Message())) - } - } - return err -} diff --git a/groups/api/grpc/doc.go b/groups/api/grpc/doc.go deleted file mode 100644 index 20956ee50..000000000 --- a/groups/api/grpc/doc.go +++ /dev/null @@ -1,5 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package grpc contains implementation of Auth service gRPC API. -package grpc diff --git a/groups/api/grpc/endpoint.go b/groups/api/grpc/endpoint.go deleted file mode 100644 index 692be6a92..000000000 --- a/groups/api/grpc/endpoint.go +++ /dev/null @@ -1,23 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package grpc - -import ( - "context" - - groups "github.com/absmach/magistrala/groups/private" - "github.com/go-kit/kit/endpoint" -) - -func retrieveEntityEndpoint(svc groups.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(retrieveEntityReq) - group, err := svc.RetrieveById(ctx, req.Id) - if err != nil { - return retrieveEntityRes{}, err - } - - return retrieveEntityRes{id: group.ID, domain: group.Domain, parentGroup: group.Parent, status: uint8(group.Status)}, nil - } -} diff --git a/groups/api/grpc/endpoint_test.go b/groups/api/grpc/endpoint_test.go deleted file mode 100644 index ddcb5f721..000000000 --- a/groups/api/grpc/endpoint_test.go +++ /dev/null @@ -1,162 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package grpc_test - -import ( - "context" - "fmt" - "net" - "testing" - "time" - - grpcCommonV1 "github.com/absmach/magistrala/api/grpc/common/v1" - grpcGroupsV1 "github.com/absmach/magistrala/api/grpc/groups/v1" - "github.com/absmach/magistrala/groups" - grpcapi "github.com/absmach/magistrala/groups/api/grpc" - prmocks "github.com/absmach/magistrala/groups/private/mocks" - "github.com/absmach/magistrala/internal/nullable" - "github.com/absmach/magistrala/internal/testsutil" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" - "google.golang.org/grpc" - "google.golang.org/grpc/credentials/insecure" -) - -const port = 7004 - -var ( - validID = testsutil.GenerateUUID(&testing.T{}) - valid = "valid" - desc = nullable.New(valid) - validGroupResp = groups.Group{ - ID: testsutil.GenerateUUID(&testing.T{}), - Name: valid, - Description: desc, - Domain: testsutil.GenerateUUID(&testing.T{}), - Parent: testsutil.GenerateUUID(&testing.T{}), - Metadata: groups.Metadata{ - "name": "test", - }, - Children: []*groups.Group{}, - CreatedAt: time.Now().Add(-1 * time.Second), - UpdatedAt: time.Now(), - UpdatedBy: testsutil.GenerateUUID(&testing.T{}), - Status: groups.EnabledStatus, - } -) - -func startGRPCServer(svc *prmocks.Service, port int) { - listener, err := net.Listen("tcp", fmt.Sprintf(":%d", port)) - if err != nil { - panic(fmt.Sprintf("failed to obtain port: %s", err)) - } - server := grpc.NewServer() - grpcGroupsV1.RegisterGroupsServiceServer(server, grpcapi.NewServer(svc)) - go func() { - if err := server.Serve(listener); err != nil { - panic(fmt.Sprintf("failed to serve: %s", err)) - } - }() -} - -func TestRetrieveEntityEndpoint(t *testing.T) { - svc := new(prmocks.Service) - startGRPCServer(svc, port) - grpAddr := fmt.Sprintf("localhost:%d", port) - conn, _ := grpc.NewClient(grpAddr, grpc.WithTransportCredentials(insecure.NewCredentials())) - client := grpcapi.NewClient(conn, time.Second) - - cases := []struct { - desc string - req *grpcCommonV1.RetrieveEntityReq - svcRes groups.Group - svcErr error - res *grpcCommonV1.RetrieveEntityRes - err error - }{ - { - desc: "retrieve group successfully", - req: &grpcCommonV1.RetrieveEntityReq{ - Id: validID, - }, - svcRes: validGroupResp, - svcErr: nil, - res: &grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Id: validGroupResp.ID, - DomainId: validGroupResp.Domain, - ParentGroupId: validGroupResp.Parent, - Status: uint32(validGroupResp.Status), - }, - }, - err: nil, - }, - { - desc: "retrieve group with authentication error", - req: &grpcCommonV1.RetrieveEntityReq{ - Id: validID, - }, - svcErr: svcerr.ErrAuthentication, - res: &grpcCommonV1.RetrieveEntityRes{}, - err: svcerr.ErrAuthentication, - }, - { - desc: "retrieve group with authorization error", - req: &grpcCommonV1.RetrieveEntityReq{ - Id: validID, - }, - svcErr: svcerr.ErrAuthorization, - res: &grpcCommonV1.RetrieveEntityRes{}, - err: svcerr.ErrAuthorization, - }, - { - desc: "retrieve group with not found error", - req: &grpcCommonV1.RetrieveEntityReq{ - Id: validID, - }, - svcErr: svcerr.ErrNotFound, - res: &grpcCommonV1.RetrieveEntityRes{}, - err: svcerr.ErrNotFound, - }, - { - desc: "retrieve group with malformed entity error", - req: &grpcCommonV1.RetrieveEntityReq{ - Id: validID, - }, - svcErr: errors.ErrMalformedEntity, - res: &grpcCommonV1.RetrieveEntityRes{}, - err: errors.ErrMalformedEntity, - }, - { - desc: "retrieve group with conflict error", - req: &grpcCommonV1.RetrieveEntityReq{ - Id: validID, - }, - svcErr: svcerr.ErrConflict, - res: &grpcCommonV1.RetrieveEntityRes{}, - err: svcerr.ErrConflict, - }, - { - desc: "retrieve group with unknown error", - req: &grpcCommonV1.RetrieveEntityReq{ - Id: validID, - }, - svcErr: errors.ErrUnidentified, - res: &grpcCommonV1.RetrieveEntityRes{}, - err: errors.ErrUnidentified, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("RetrieveById", mock.Anything, tc.req.Id).Return(tc.svcRes, tc.svcErr) - res, err := client.RetrieveEntity(context.Background(), tc.req) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s", tc.desc, tc.err, err)) - assert.Equal(t, tc.res, res, fmt.Sprintf("%s: expected %s got %s", tc.desc, tc.res, res)) - svcCall.Unset() - }) - } -} diff --git a/groups/api/grpc/request.go b/groups/api/grpc/request.go deleted file mode 100644 index 4c8286e10..000000000 --- a/groups/api/grpc/request.go +++ /dev/null @@ -1,8 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package grpc - -type retrieveEntityReq struct { - Id string -} diff --git a/groups/api/grpc/responses.go b/groups/api/grpc/responses.go deleted file mode 100644 index 8370e73b8..000000000 --- a/groups/api/grpc/responses.go +++ /dev/null @@ -1,13 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package grpc - -type groupBasic struct { - id string - domain string - parentGroup string - status uint8 -} - -type retrieveEntityRes groupBasic diff --git a/groups/api/grpc/server.go b/groups/api/grpc/server.go deleted file mode 100644 index a07d20404..000000000 --- a/groups/api/grpc/server.go +++ /dev/null @@ -1,89 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package grpc - -import ( - "context" - - grpcCommonV1 "github.com/absmach/magistrala/api/grpc/common/v1" - grpcGroupsV1 "github.com/absmach/magistrala/api/grpc/groups/v1" - apiutil "github.com/absmach/magistrala/api/http/util" - smqauth "github.com/absmach/magistrala/auth" - groups "github.com/absmach/magistrala/groups/private" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - kitgrpc "github.com/go-kit/kit/transport/grpc" - "google.golang.org/grpc/codes" - "google.golang.org/grpc/status" -) - -var _ grpcGroupsV1.GroupsServiceServer = (*grpcServer)(nil) - -type grpcServer struct { - grpcGroupsV1.UnimplementedGroupsServiceServer - retrieveEntity kitgrpc.Handler -} - -// NewServer returns new AuthServiceServer instance. -func NewServer(svc groups.Service) grpcGroupsV1.GroupsServiceServer { - return &grpcServer{ - retrieveEntity: kitgrpc.NewServer( - retrieveEntityEndpoint(svc), - decodeRetrieveEntityRequest, - encodeRetrieveEntityResponse, - ), - } -} - -func (s *grpcServer) RetrieveEntity(ctx context.Context, req *grpcCommonV1.RetrieveEntityReq) (*grpcCommonV1.RetrieveEntityRes, error) { - _, res, err := s.retrieveEntity.ServeGRPC(ctx, req) - if err != nil { - return nil, encodeError(err) - } - return res.(*grpcCommonV1.RetrieveEntityRes), nil -} - -func decodeRetrieveEntityRequest(_ context.Context, grpcReq any) (any, error) { - req := grpcReq.(*grpcCommonV1.RetrieveEntityReq) - return retrieveEntityReq{ - Id: req.GetId(), - }, nil -} - -func encodeRetrieveEntityResponse(_ context.Context, grpcRes any) (any, error) { - res := grpcRes.(retrieveEntityRes) - - return &grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Id: res.id, - DomainId: res.domain, - ParentGroupId: res.parentGroup, - Status: uint32(res.status), - }, - }, nil -} - -func encodeError(err error) error { - switch { - case errors.Contains(err, nil): - return nil - case errors.Contains(err, errors.ErrMalformedEntity), - err == apiutil.ErrInvalidAuthKey, - err == apiutil.ErrMissingID, - err == apiutil.ErrMissingMemberType, - err == apiutil.ErrMissingPolicySub, - err == apiutil.ErrMissingPolicyObj, - err == apiutil.ErrMalformedPolicyAct: - return status.Error(codes.InvalidArgument, err.Error()) - case errors.Contains(err, svcerr.ErrAuthentication), - errors.Contains(err, smqauth.ErrKeyExpired), - err == apiutil.ErrMissingEmail, - err == apiutil.ErrBearerToken: - return status.Error(codes.Unauthenticated, err.Error()) - case errors.Contains(err, svcerr.ErrAuthorization): - return status.Error(codes.PermissionDenied, err.Error()) - default: - return status.Error(codes.Internal, err.Error()) - } -} diff --git a/groups/api/http/decode.go b/groups/api/http/decode.go deleted file mode 100644 index ac4c85ead..000000000 --- a/groups/api/http/decode.go +++ /dev/null @@ -1,346 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package api - -import ( - "context" - "encoding/json" - "net/http" - "strings" - "time" - - api "github.com/absmach/magistrala/api/http" - apiutil "github.com/absmach/magistrala/api/http/util" - groups "github.com/absmach/magistrala/groups" - "github.com/absmach/magistrala/pkg/errors" - "github.com/go-chi/chi/v5" -) - -func DecodeGroupCreate(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - var g groups.Group - if err := json.NewDecoder(r.Body).Decode(&g); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - req := createGroupReq{ - Group: g, - } - - return req, nil -} - -func DecodeListGroupsRequest(_ context.Context, r *http.Request) (any, error) { - pm, err := decodePageMeta(r) - if err != nil { - return nil, err - } - - userID, err := apiutil.ReadStringQuery(r, api.UserKey, "") - if err != nil { - return groups.PageMeta{}, errors.Wrap(apiutil.ErrValidation, err) - } - - groupID, err := apiutil.ReadStringQuery(r, api.GroupKey, "") - if err != nil { - return groups.PageMeta{}, errors.Wrap(apiutil.ErrValidation, err) - } - - req := listGroupsReq{ - PageMeta: pm, - userID: userID, - groupID: groupID, - } - return req, nil -} - -func DecodeGroupUpdate(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - req := updateGroupReq{ - id: chi.URLParam(r, "groupID"), - } - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - return req, nil -} - -func decodeUpdateGroupTags(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - - req := updateGroupTagsReq{ - id: chi.URLParam(r, "groupID"), - } - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - - return req, nil -} - -func DecodeGroupRequest(_ context.Context, r *http.Request) (any, error) { - roles, err := apiutil.ReadBoolQuery(r, api.RolesKey, false) - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - req := groupReq{ - id: chi.URLParam(r, "groupID"), - roles: roles, - } - - return req, nil -} - -func DecodeChangeGroupStatusRequest(_ context.Context, r *http.Request) (any, error) { - req := changeGroupStatusReq{ - id: chi.URLParam(r, "groupID"), - } - return req, nil -} - -func decodeRetrieveGroupHierarchy(_ context.Context, r *http.Request) (any, error) { - hm, err := decodeHierarchyPageMeta(r) - if err != nil { - return nil, err - } - - req := retrieveGroupHierarchyReq{ - id: chi.URLParam(r, "groupID"), - HierarchyPageMeta: hm, - } - return req, nil -} - -func decodeAddParentGroupRequest(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - - req := addParentGroupReq{ - id: chi.URLParam(r, "groupID"), - } - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - return req, nil -} - -func decodeRemoveParentGroupRequest(_ context.Context, r *http.Request) (any, error) { - req := removeParentGroupReq{ - id: chi.URLParam(r, "groupID"), - } - return req, nil -} - -func decodeAddChildrenGroupsRequest(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - req := addChildrenGroupsReq{ - id: chi.URLParam(r, "groupID"), - } - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - return req, nil -} - -func decodeRemoveChildrenGroupsRequest(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - req := removeChildrenGroupsReq{ - id: chi.URLParam(r, "groupID"), - } - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - return req, nil -} - -func decodeRemoveAllChildrenGroupsRequest(_ context.Context, r *http.Request) (any, error) { - req := removeAllChildrenGroupsReq{ - id: chi.URLParam(r, "groupID"), - } - return req, nil -} - -func decodeListChildrenGroupsRequest(_ context.Context, r *http.Request) (any, error) { - pm, err := decodePageMeta(r) - if err != nil { - return nil, err - } - - startLevel, err := apiutil.ReadNumQuery[int64](r, api.StartLevelKey, api.DefStartLevel) - if err != nil { - return groups.PageMeta{}, errors.Wrap(apiutil.ErrValidation, err) - } - - endLevel, err := apiutil.ReadNumQuery[int64](r, api.EndLevelKey, api.DefEndLevel) - if err != nil { - return groups.PageMeta{}, errors.Wrap(apiutil.ErrValidation, err) - } - - req := listChildrenGroupsReq{ - id: chi.URLParam(r, "groupID"), - PageMeta: pm, - startLevel: startLevel, - endLevel: endLevel, - } - return req, nil -} - -func decodeHierarchyPageMeta(r *http.Request) (groups.HierarchyPageMeta, error) { - level, err := apiutil.ReadNumQuery[uint64](r, api.LevelKey, api.DefLevel) - if err != nil { - return groups.HierarchyPageMeta{}, errors.Wrap(apiutil.ErrValidation, err) - } - - tree, err := apiutil.ReadBoolQuery(r, api.TreeKey, false) - if err != nil { - return groups.HierarchyPageMeta{}, errors.Wrap(apiutil.ErrValidation, err) - } - hierarchyDir, err := apiutil.ReadNumQuery[int64](r, api.DirKey, -1) - if err != nil { - return groups.HierarchyPageMeta{}, errors.Wrap(apiutil.ErrValidation, err) - } - - return groups.HierarchyPageMeta{ - Level: level, - Direction: hierarchyDir, - Tree: tree, - }, nil -} - -func decodePageMeta(r *http.Request) (groups.PageMeta, error) { - s, err := apiutil.ReadStringQuery(r, api.StatusKey, api.DefGroupStatus) - if err != nil { - return groups.PageMeta{}, errors.Wrap(apiutil.ErrValidation, err) - } - st, err := groups.ToStatus(s) - if err != nil { - return groups.PageMeta{}, errors.Wrap(apiutil.ErrValidation, err) - } - offset, err := apiutil.ReadNumQuery[uint64](r, api.OffsetKey, api.DefOffset) - if err != nil { - return groups.PageMeta{}, errors.Wrap(apiutil.ErrValidation, err) - } - limit, err := apiutil.ReadNumQuery[uint64](r, api.LimitKey, api.DefLimit) - if err != nil { - return groups.PageMeta{}, errors.Wrap(apiutil.ErrValidation, err) - } - name, err := apiutil.ReadStringQuery(r, api.NameKey, "") - if err != nil { - return groups.PageMeta{}, errors.Wrap(apiutil.ErrValidation, err) - } - id, err := apiutil.ReadStringQuery(r, api.IDOrder, "") - if err != nil { - return groups.PageMeta{}, errors.Wrap(apiutil.ErrValidation, err) - } - meta, err := apiutil.ReadMetadataQuery(r, api.MetadataKey, nil) - if err != nil { - return groups.PageMeta{}, errors.Wrap(apiutil.ErrValidation, err) - } - - allActions, err := apiutil.ReadStringQuery(r, api.ActionsKey, "") - if err != nil { - return groups.PageMeta{}, errors.Wrap(apiutil.ErrValidation, err) - } - - actions := []string{} - - allActions = strings.TrimSpace(allActions) - if allActions != "" { - actions = strings.Split(allActions, ",") - } - roleID, err := apiutil.ReadStringQuery(r, api.RoleIDKey, "") - if err != nil { - return groups.PageMeta{}, errors.Wrap(apiutil.ErrValidation, err) - } - - roleName, err := apiutil.ReadStringQuery(r, api.RoleNameKey, "") - if err != nil { - return groups.PageMeta{}, errors.Wrap(apiutil.ErrValidation, err) - } - - accessType, err := apiutil.ReadStringQuery(r, api.AccessTypeKey, "") - if err != nil { - return groups.PageMeta{}, errors.Wrap(apiutil.ErrValidation, err) - } - - rootGroup, err := apiutil.ReadBoolQuery(r, api.RootGroupKey, false) - if err != nil { - return groups.PageMeta{}, errors.Wrap(apiutil.ErrValidation, err) - } - - ot, err := apiutil.ReadBoolQuery(r, api.OnlyTotal, false) - if err != nil { - return groups.PageMeta{}, errors.Wrap(apiutil.ErrValidation, err) - } - - order, err := apiutil.ReadStringQuery(r, api.OrderKey, api.DefOrder) - if err != nil { - return groups.PageMeta{}, errors.Wrap(apiutil.ErrValidation, err) - } - dir, err := apiutil.ReadStringQuery(r, api.DirKey, api.DefDir) - if err != nil { - return groups.PageMeta{}, errors.Wrap(apiutil.ErrValidation, err) - } - tags, err := apiutil.ReadStringQuery(r, api.TagsKey, "") - if err != nil { - return groups.PageMeta{}, errors.Wrap(apiutil.ErrValidation, err) - } - var tq groups.TagsQuery - if tags != "" { - tq = groups.ToTagsQuery(tags) - } - - cfrom, err := apiutil.ReadStringQuery(r, "created_from", "") - if err != nil { - return groups.PageMeta{}, errors.Wrap(apiutil.ErrValidation, err) - } - cto, err := apiutil.ReadStringQuery(r, "created_to", "") - if err != nil { - return groups.PageMeta{}, errors.Wrap(apiutil.ErrValidation, err) - } - - var createdFrom, createdTo time.Time - if cfrom != "" { - if createdFrom, err = time.Parse(time.RFC3339, cfrom); err != nil { - return groups.PageMeta{}, errors.Wrap(apiutil.ErrInvalidQueryParams, err) - } - } - if cto != "" { - if createdTo, err = time.Parse(time.RFC3339, cto); err != nil { - return groups.PageMeta{}, errors.Wrap(apiutil.ErrInvalidQueryParams, err) - } - } - - ret := groups.PageMeta{ - Offset: offset, - Limit: limit, - Name: name, - ID: id, - Metadata: meta, - Status: st, - RoleName: roleName, - RoleID: roleID, - Actions: actions, - AccessType: accessType, - RootGroup: rootGroup, - OnlyTotal: ot, - Order: order, - Dir: dir, - Tags: tq, - CreatedFrom: createdFrom, - CreatedTo: createdTo, - } - return ret, nil -} diff --git a/groups/api/http/decode_test.go b/groups/api/http/decode_test.go deleted file mode 100644 index 818acf327..000000000 --- a/groups/api/http/decode_test.go +++ /dev/null @@ -1,535 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package api - -import ( - "context" - "fmt" - "net/http" - "net/url" - "strings" - "testing" - "time" - - api "github.com/absmach/magistrala/api/http" - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/groups" - "github.com/absmach/magistrala/pkg/errors" - "github.com/stretchr/testify/assert" -) - -func TestDecodeListGroupsRequest(t *testing.T) { - cases := []struct { - desc string - url string - header map[string][]string - resp any - err error - }{ - { - desc: "valid request with no parameters", - url: "http://localhost:8080", - header: map[string][]string{}, - resp: listGroupsReq{ - PageMeta: groups.PageMeta{ - Limit: 10, - Actions: []string{}, - Dir: "desc", - Order: "updated_at", - }, - }, - err: nil, - }, - { - desc: "valid request with all parameters", - url: "http://localhost:8080?status=enabled&offset=10&limit=10&name=random&metadata={\"test\":\"test\"}&level=2&t&permission=random&list_perms=true", - header: map[string][]string{ - "Authorization": {"Bearer 123"}, - }, - resp: listGroupsReq{ - PageMeta: groups.PageMeta{ - Status: groups.EnabledStatus, - Offset: 10, - Limit: 10, - Name: "random", - Metadata: groups.Metadata{ - "test": "test", - }, - Actions: []string{}, - Dir: "desc", - Order: "updated_at", - }, - }, - err: nil, - }, - { - desc: "valid request with invalid page metadata", - url: "http://localhost:8080?metadata=random", - resp: nil, - err: apiutil.ErrValidation, - }, - { - desc: "valid request with created_from parameter", - url: "http://localhost:8080?created_from=2024-01-01T00:00:00Z", - resp: listGroupsReq{ - PageMeta: groups.PageMeta{ - Limit: 10, - Actions: []string{}, - Dir: "desc", - Order: "updated_at", - CreatedFrom: time.Date(2024, 1, 1, 0, 0, 0, 0, time.UTC), - }, - }, - err: nil, - }, - { - desc: "valid request with created_to parameter", - url: "http://localhost:8080?created_to=2024-12-31T23:59:59Z", - resp: listGroupsReq{ - PageMeta: groups.PageMeta{ - Limit: 10, - Actions: []string{}, - Dir: "desc", - Order: "updated_at", - CreatedTo: time.Date(2024, 12, 31, 23, 59, 59, 0, time.UTC), - }, - }, - err: nil, - }, - { - desc: "valid request with both created_from and created_to parameters", - url: "http://localhost:8080?created_from=2024-01-01T00:00:00Z&created_to=2024-12-31T23:59:59Z", - resp: listGroupsReq{ - PageMeta: groups.PageMeta{ - Limit: 10, - Actions: []string{}, - Dir: "desc", - Order: "updated_at", - CreatedFrom: time.Date(2024, 1, 1, 0, 0, 0, 0, time.UTC), - CreatedTo: time.Date(2024, 12, 31, 23, 59, 59, 0, time.UTC), - }, - }, - err: nil, - }, - { - desc: "invalid request with malformed created_from", - url: "http://localhost:8080?created_from=invalid-timestamp", - resp: nil, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "invalid request with malformed created_to", - url: "http://localhost:8080?created_to=invalid-timestamp", - resp: nil, - err: apiutil.ErrInvalidQueryParams, - }, - } - - for _, tc := range cases { - parsedURL, err := url.Parse(tc.url) - assert.NoError(t, err) - - req := &http.Request{ - URL: parsedURL, - Header: tc.header, - } - resp, err := DecodeListGroupsRequest(context.Background(), req) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.resp, resp)) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("expected error %v to contain %v", err, tc.err)) - } -} - -func TestDecodeRetrieveGroupHierarchy(t *testing.T) { - cases := []struct { - desc string - url string - header map[string][]string - resp any - err error - }{ - { - desc: "valid request with no parameters", - url: "http://localhost:8080", - header: map[string][]string{}, - resp: retrieveGroupHierarchyReq{ - HierarchyPageMeta: groups.HierarchyPageMeta{ - Direction: -1, - }, - }, - err: nil, - }, - { - desc: "valid request with all parameters", - url: "http://localhost:8080?tree=true&level=2&dir=-1", - header: map[string][]string{ - "Authorization": {"Bearer 123"}, - }, - resp: retrieveGroupHierarchyReq{ - HierarchyPageMeta: groups.HierarchyPageMeta{ - Level: 2, - Direction: -1, - Tree: true, - }, - }, - err: nil, - }, - { - desc: "valid request with invalid level", - url: "http://localhost:8080?level=random", - resp: nil, - err: apiutil.ErrValidation, - }, - { - desc: "valid request with invalid tree", - url: "http://localhost:8080?tree=random", - resp: nil, - err: apiutil.ErrValidation, - }, - { - desc: "valid request with invalid direction", - url: "http://localhost:8080?dir=random", - resp: nil, - err: apiutil.ErrValidation, - }, - } - - for _, tc := range cases { - parsedURL, err := url.Parse(tc.url) - assert.NoError(t, err) - - req := &http.Request{ - URL: parsedURL, - Header: tc.header, - } - resp, err := decodeRetrieveGroupHierarchy(context.Background(), req) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.resp, resp)) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("expected error %v to contain %v", err, tc.err)) - } -} - -func TestDecodeListChildrenRequest(t *testing.T) { - cases := []struct { - desc string - url string - header map[string][]string - resp any - err error - }{ - { - desc: "valid request with no parameters", - url: "http://localhost:8080", - header: map[string][]string{}, - resp: listChildrenGroupsReq{ - startLevel: 1, - endLevel: 0, - PageMeta: groups.PageMeta{ - Limit: 10, - Actions: []string{}, - Dir: "desc", - Order: "updated_at", - }, - }, - err: nil, - }, - { - desc: "valid request with all parameters", - url: "http://localhost:8080?status=enabled&offset=10&limit=10&name=random&metadata={\"test\":\"test\"}&level=2&parent_id=random&tree=true&dir=desc&member_kind=random&permission=random&list_perms=true", - header: map[string][]string{ - "Authorization": {"Bearer 123"}, - }, - resp: listChildrenGroupsReq{ - startLevel: 1, - endLevel: 0, - PageMeta: groups.PageMeta{ - Status: groups.EnabledStatus, - Offset: 10, - Limit: 10, - Name: "random", - Metadata: groups.Metadata{ - "test": "test", - }, - Actions: []string{}, - Dir: "desc", - Order: "updated_at", - }, - }, - err: nil, - }, - { - desc: "valid request with invalid page metadata", - url: "http://localhost:8080?metadata=random", - resp: nil, - err: apiutil.ErrValidation, - }, - } - - for _, tc := range cases { - parsedURL, err := url.Parse(tc.url) - assert.NoError(t, err) - - req := &http.Request{ - URL: parsedURL, - Header: tc.header, - } - resp, err := decodeListChildrenGroupsRequest(context.Background(), req) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.resp, resp)) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("expected error %v to contain %v", err, tc.err)) - } -} - -func TestDecodePageMeta(t *testing.T) { - cases := []struct { - desc string - url string - resp groups.PageMeta - err error - }{ - { - desc: "valid request with no parameters", - url: "http://localhost:8080", - resp: groups.PageMeta{ - Limit: 10, - Actions: []string{}, - Dir: "desc", - Order: "updated_at", - }, - err: nil, - }, - { - desc: "valid request with all parameters", - url: "http://localhost:8080?status=enabled&offset=10&limit=10&name=random&metadata={\"test\":\"test\"}", - resp: groups.PageMeta{ - Status: groups.EnabledStatus, - Offset: 10, - Limit: 10, - Name: "random", - Metadata: groups.Metadata{ - "test": "test", - }, - Actions: []string{}, - Dir: "desc", - Order: "updated_at", - }, - err: nil, - }, - { - desc: "valid request with invalid status", - url: "http://localhost:8080?status=random", - resp: groups.PageMeta{}, - err: apiutil.ErrValidation, - }, - { - desc: "valid request with invalid status duplicated", - url: "http://localhost:8080?status=random&status=random", - resp: groups.PageMeta{}, - err: apiutil.ErrValidation, - }, - { - desc: "valid request with invalid offset", - url: "http://localhost:8080?offset=random", - resp: groups.PageMeta{}, - err: apiutil.ErrValidation, - }, - { - desc: "valid request with invalid limit", - url: "http://localhost:8080?limit=random", - resp: groups.PageMeta{}, - err: apiutil.ErrValidation, - }, - { - desc: "valid request with invalid name", - url: "http://localhost:8080?name=random&name=random", - resp: groups.PageMeta{}, - err: apiutil.ErrValidation, - }, - { - desc: "valid request with invalid page metadata", - url: "http://localhost:8080?metadata=random", - resp: groups.PageMeta{}, - err: apiutil.ErrValidation, - }, - } - - for _, tc := range cases { - parsedURL, err := url.Parse(tc.url) - assert.NoError(t, err) - - req := &http.Request{URL: parsedURL} - resp, err := decodePageMeta(req) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("expected error %v to contain %v", err, tc.err)) - } -} - -func TestDecodeGroupCreate(t *testing.T) { - cases := []struct { - desc string - body string - header map[string][]string - resp any - err error - }{ - { - desc: "valid request", - body: `{"name": "random", "description": "valid"}`, - header: map[string][]string{ - "Authorization": {"Bearer 123"}, - "Content-Type": {api.ContentType}, - }, - resp: createGroupReq{ - Group: groups.Group{ - Name: "random", - Description: desc, - }, - }, - err: nil, - }, - { - desc: "invalid content type", - body: `{"name": "random", "description": "random"}`, - header: map[string][]string{ - "Authorization": {"Bearer 123"}, - "Content-Type": {"text/plain"}, - }, - resp: nil, - err: apiutil.ErrUnsupportedContentType, - }, - { - desc: "invalid request body", - body: `data`, - header: map[string][]string{ - "Authorization": {"Bearer 123"}, - "Content-Type": {api.ContentType}, - }, - resp: nil, - err: apiutil.ErrMalformedRequestBody, - }, - } - - for _, tc := range cases { - req, err := http.NewRequest(http.MethodPost, "http://localhost:8080", strings.NewReader(tc.body)) - assert.NoError(t, err) - req.Header = tc.header - resp, err := DecodeGroupCreate(context.Background(), req) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.resp, resp)) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("expected error %v to contain %v", err, tc.err)) - } -} - -func TestDecodeGroupUpdate(t *testing.T) { - cases := []struct { - desc string - body string - header map[string][]string - resp any - err error - }{ - { - desc: "valid request", - body: `{"name": "random", "description": "valid"}`, - header: map[string][]string{ - "Authorization": {"Bearer 123"}, - "Content-Type": {api.ContentType}, - }, - resp: updateGroupReq{ - Name: "random", - Description: desc, - }, - err: nil, - }, - { - desc: "invalid content type", - body: `{"name": "random", "description": "valid"}`, - header: map[string][]string{ - "Authorization": {"Bearer 123"}, - "Content-Type": {"text/plain"}, - }, - resp: nil, - err: apiutil.ErrUnsupportedContentType, - }, - { - desc: "invalid request body", - body: `data`, - header: map[string][]string{ - "Authorization": {"Bearer 123"}, - "Content-Type": {api.ContentType}, - }, - resp: nil, - err: apiutil.ErrMalformedRequestBody, - }, - } - - for _, tc := range cases { - req, err := http.NewRequest(http.MethodPut, "http://localhost:8080", strings.NewReader(tc.body)) - assert.NoError(t, err) - req.Header = tc.header - resp, err := DecodeGroupUpdate(context.Background(), req) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.resp, resp)) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("expected error %v to contain %v", err, tc.err)) - } -} - -func TestDecodeGroupRequest(t *testing.T) { - cases := []struct { - desc string - header map[string][]string - resp any - err error - }{ - { - desc: "valid request", - header: map[string][]string{ - "Authorization": {"Bearer 123"}, - }, - resp: groupReq{}, - err: nil, - }, - { - desc: "empty token", - resp: groupReq{}, - err: nil, - }, - } - - for _, tc := range cases { - req, err := http.NewRequest(http.MethodGet, "http://localhost:8080", http.NoBody) - assert.NoError(t, err) - req.Header = tc.header - resp, err := DecodeGroupRequest(context.Background(), req) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.resp, resp)) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("expected error %v to contain %v", err, tc.err)) - } -} - -func TestDecodeChangeGroupStatus(t *testing.T) { - cases := []struct { - desc string - header map[string][]string - resp any - err error - }{ - { - desc: "valid request", - header: map[string][]string{ - "Authorization": {"Bearer 123"}, - }, - resp: changeGroupStatusReq{}, - err: nil, - }, - { - desc: "empty token", - resp: changeGroupStatusReq{}, - err: nil, - }, - } - - for _, tc := range cases { - req, err := http.NewRequest(http.MethodGet, "http://localhost:8080", http.NoBody) - assert.NoError(t, err) - req.Header = tc.header - resp, err := DecodeChangeGroupStatusRequest(context.Background(), req) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.resp, resp)) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("expected error %v to contain %v", err, tc.err)) - } -} diff --git a/groups/api/http/doc.go b/groups/api/http/doc.go deleted file mode 100644 index 2424852cc..000000000 --- a/groups/api/http/doc.go +++ /dev/null @@ -1,6 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package api contains API-related concerns: endpoint definitions, middlewares -// and all resource representations. -package api diff --git a/groups/api/http/endpoint_test.go b/groups/api/http/endpoint_test.go deleted file mode 100644 index 06796fc39..000000000 --- a/groups/api/http/endpoint_test.go +++ /dev/null @@ -1,2320 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package api - -import ( - "encoding/json" - "fmt" - "io" - "net/http" - "net/http/httptest" - "net/url" - "strings" - "testing" - "time" - - api "github.com/absmach/magistrala/api/http" - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/groups" - "github.com/absmach/magistrala/groups/mocks" - "github.com/absmach/magistrala/internal/nullable" - "github.com/absmach/magistrala/internal/testsutil" - mglog "github.com/absmach/magistrala/logger" - smqauthn "github.com/absmach/magistrala/pkg/authn" - authnmocks "github.com/absmach/magistrala/pkg/authn/mocks" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/pkg/roles" - "github.com/absmach/magistrala/pkg/uuid" - "github.com/go-chi/chi/v5" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" -) - -var ( - desc = nullable.New(valid) - validGroupResp = groups.Group{ - ID: testsutil.GenerateUUID(&testing.T{}), - Name: valid, - Description: desc, - Domain: testsutil.GenerateUUID(&testing.T{}), - Parent: testsutil.GenerateUUID(&testing.T{}), - Metadata: groups.Metadata{ - "name": "test", - }, - Children: []*groups.Group{}, - CreatedAt: time.Now().Add(-1 * time.Second), - UpdatedAt: time.Now(), - UpdatedBy: testsutil.GenerateUUID(&testing.T{}), - Status: groups.EnabledStatus, - } - validID = testsutil.GenerateUUID(&testing.T{}) - validToken = "validToken" - invalidToken = "invalidToken" - contentType = "application/json" - validTimeStamp = time.Date(2024, 1, 1, 0, 0, 0, 0, time.UTC) -) - -func newGroupsServer() (*httptest.Server, *mocks.Service, *authnmocks.Authentication) { - authn := new(authnmocks.Authentication) - svc := new(mocks.Service) - mux := chi.NewRouter() - idp := uuid.NewMock() - logger := mglog.NewMock() - am := smqauthn.NewAuthNMiddleware(authn, smqauthn.WithAllowUnverifiedUser(true)) - mux = MakeHandler(svc, am, mux, logger, "", idp) - - return httptest.NewServer(mux), svc, authn -} - -func TestCreateGroupEndpoint(t *testing.T) { - gs, svc, authn := newGroupsServer() - defer gs.Close() - - reqGroup := groups.Group{ - Name: valid, - Description: desc, - Metadata: map[string]any{ - "name": "test", - }, - } - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - req createGroupReq - contentType string - svcResp groups.Group - svcErr error - authnErr error - status int - err error - }{ - { - desc: "create group successfully", - token: validToken, - domainID: validID, - req: createGroupReq{ - Group: reqGroup, - }, - contentType: contentType, - svcResp: validGroupResp, - status: http.StatusCreated, - err: nil, - }, - { - desc: "create group with invalid token", - token: invalidToken, - session: smqauthn.Session{}, - domainID: validID, - req: createGroupReq{ - Group: reqGroup, - }, - contentType: contentType, - authnErr: svcerr.ErrAuthentication, - status: http.StatusUnauthorized, - err: svcerr.ErrAuthentication, - }, - { - desc: "create group with empty token", - token: "", - session: smqauthn.Session{}, - domainID: validID, - req: createGroupReq{ - Group: reqGroup, - }, - contentType: contentType, - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "create group with empty domainID", - token: validToken, - req: createGroupReq{ - Group: reqGroup, - }, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrMissingDomainID, - }, - { - desc: "create group with missing name", - token: validToken, - domainID: validID, - req: createGroupReq{ - Group: groups.Group{ - Description: desc, - Metadata: map[string]any{ - "name": "test", - }, - }, - }, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrNameSize, - }, - { - desc: "create group with name that is too long", - token: validToken, - domainID: validID, - req: createGroupReq{ - Group: groups.Group{ - Name: strings.Repeat("a", 1025), - Description: desc, - Metadata: map[string]any{ - "name": "test", - }, - }, - }, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrNameSize, - }, - { - desc: "create group with invalid content type", - token: validToken, - domainID: validID, - req: createGroupReq{ - Group: reqGroup, - }, - contentType: "application/xml", - svcResp: validGroupResp, - status: http.StatusUnsupportedMediaType, - err: apiutil.ErrUnsupportedContentType, - }, - { - desc: "create group with service error", - token: validToken, - domainID: validID, - req: createGroupReq{ - Group: reqGroup, - }, - contentType: contentType, - svcResp: groups.Group{}, - svcErr: svcerr.ErrAuthorization, - status: http.StatusForbidden, - err: svcerr.ErrAuthorization, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - data := toJSON(tc.req) - req := testRequest{ - client: gs.Client(), - method: http.MethodPost, - url: fmt.Sprintf("%s/%s/groups/", gs.URL, tc.domainID), - contentType: tc.contentType, - token: tc.token, - body: strings.NewReader(data), - } - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID + "_" + validID, UserID: validID, DomainID: validID} - } - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - svcCall := svc.On("CreateGroup", mock.Anything, tc.session, tc.req.Group).Return(tc.svcResp, []roles.RoleProvision{}, tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - var errRes respBody - err = json.NewDecoder(res.Body).Decode(&errRes) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if errRes.Err != "" || errRes.Message != "" { - err = errors.Wrap(errors.New(errRes.Err), errors.New(errRes.Message)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestViewGroupEndpoint(t *testing.T) { - gs, svc, authn := newGroupsServer() - defer gs.Close() - - cases := []struct { - desc string - token string - id string - domainID string - roles bool - session smqauthn.Session - svcResp groups.Group - svcErr error - resp groups.Group - status int - authnErr error - err error - }{ - { - desc: "view group successfully", - token: validToken, - domainID: validID, - roles: false, - id: validID, - svcResp: validGroupResp, - svcErr: nil, - resp: validGroupResp, - status: http.StatusOK, - err: nil, - }, - { - desc: "view group with invalid token", - token: invalidToken, - session: smqauthn.Session{}, - domainID: validID, - roles: false, - id: validID, - svcResp: validGroupResp, - svcErr: nil, - authnErr: svcerr.ErrAuthentication, - status: http.StatusUnauthorized, - err: svcerr.ErrAuthentication, - }, - { - desc: "view group with empty token", - token: "", - session: smqauthn.Session{}, - domainID: validID, - roles: false, - id: validID, - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "view group with empty domainID", - token: validToken, - id: validID, - status: http.StatusBadRequest, - err: apiutil.ErrMissingDomainID, - }, - { - desc: "view group with service error", - token: validToken, - id: validID, - domainID: validID, - roles: false, - svcResp: validGroupResp, - svcErr: svcerr.ErrAuthorization, - status: http.StatusForbidden, - err: svcerr.ErrAuthorization, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - client: gs.Client(), - method: http.MethodGet, - url: fmt.Sprintf("%s/%s/groups/%s?roles=%v", gs.URL, tc.domainID, tc.id, tc.roles), - token: tc.token, - } - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID + "_" + validID, UserID: validID, DomainID: validID} - } - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - svcCall := svc.On("ViewGroup", mock.Anything, tc.session, tc.id, tc.roles).Return(tc.svcResp, tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - var errRes respBody - err = json.NewDecoder(res.Body).Decode(&errRes) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if errRes.Err != "" || errRes.Message != "" { - err = errors.Wrap(errors.New(errRes.Err), errors.New(errRes.Message)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestUpdateGroupEndpoint(t *testing.T) { - gs, svc, authn := newGroupsServer() - defer gs.Close() - - updateGroupReq := groups.Group{ - ID: validID, - Name: valid, - Description: desc, - Metadata: map[string]any{ - "name": "test", - }, - } - - cases := []struct { - desc string - token string - id string - domainID string - updateReq groups.Group - contentType string - session smqauthn.Session - svcResp groups.Group - svcErr error - resp groups.Group - status int - authnErr error - err error - }{ - { - desc: "update group successfully", - token: validToken, - domainID: validID, - id: validID, - updateReq: updateGroupReq, - contentType: contentType, - svcResp: validGroupResp, - status: http.StatusOK, - err: nil, - }, - { - desc: "update group with invalid token", - token: invalidToken, - session: smqauthn.Session{}, - domainID: validID, - id: validID, - updateReq: updateGroupReq, - contentType: contentType, - authnErr: svcerr.ErrAuthentication, - status: http.StatusUnauthorized, - err: svcerr.ErrAuthentication, - }, - { - desc: "update group with empty token", - token: "", - session: smqauthn.Session{}, - domainID: validID, - id: validID, - updateReq: updateGroupReq, - contentType: contentType, - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "update group with empty domainID", - token: validToken, - id: validID, - updateReq: updateGroupReq, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrMissingDomainID, - }, - { - desc: "update group with name that is too long", - token: validToken, - id: validID, - domainID: validID, - updateReq: groups.Group{ - ID: validID, - Name: strings.Repeat("a", 1025), - Description: desc, - Metadata: map[string]any{ - "name": "test", - }, - }, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrNameSize, - }, - { - desc: "update group with invalid content type", - token: validToken, - id: validID, - domainID: validID, - updateReq: updateGroupReq, - contentType: "application/xml", - svcResp: validGroupResp, - status: http.StatusUnsupportedMediaType, - err: apiutil.ErrUnsupportedContentType, - }, - { - desc: "update group with service error", - token: validToken, - id: validID, - domainID: validID, - updateReq: updateGroupReq, - contentType: contentType, - svcResp: groups.Group{}, - svcErr: svcerr.ErrAuthorization, - status: http.StatusForbidden, - err: svcerr.ErrAuthorization, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - data := toJSON(tc.updateReq) - req := testRequest{ - client: gs.Client(), - method: http.MethodPut, - url: fmt.Sprintf("%s/%s/groups/%s", gs.URL, tc.domainID, tc.id), - contentType: tc.contentType, - token: tc.token, - body: strings.NewReader(data), - } - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID + "_" + validID, UserID: validID, DomainID: validID} - } - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - svcCall := svc.On("UpdateGroup", mock.Anything, tc.session, tc.updateReq).Return(tc.svcResp, tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - var errRes respBody - err = json.NewDecoder(res.Body).Decode(&errRes) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if errRes.Err != "" || errRes.Message != "" { - err = errors.Wrap(errors.New(errRes.Err), errors.New(errRes.Message)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestUpdateGroupTagsEndpoint(t *testing.T) { - gs, svc, authn := newGroupsServer() - defer gs.Close() - - newTag := "newtag" - - cases := []struct { - desc string - token string - id string - domainID string - data string - contentType string - session smqauthn.Session - svcResp groups.Group - svcErr error - resp groups.Group - status int - authnErr error - err error - }{ - { - desc: "update group tags successfully", - token: validToken, - domainID: validID, - id: validID, - data: fmt.Sprintf(`{"tags":["%s"]}`, newTag), - contentType: contentType, - svcResp: validGroupResp, - status: http.StatusOK, - err: nil, - }, - { - desc: "update group tags with invalid token", - token: invalidToken, - session: smqauthn.Session{}, - domainID: validID, - id: validID, - data: fmt.Sprintf(`{"tags":["%s"]}`, newTag), - contentType: contentType, - authnErr: svcerr.ErrAuthentication, - status: http.StatusUnauthorized, - err: svcerr.ErrAuthentication, - }, - { - desc: "update group tags with empty token", - token: "", - session: smqauthn.Session{}, - domainID: validID, - id: validID, - data: fmt.Sprintf(`{"tags":["%s"]}`, newTag), - contentType: contentType, - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "update group tags with empty domainID", - token: validToken, - id: validID, - data: fmt.Sprintf(`{"tags":["%s"]}`, newTag), - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrMissingDomainID, - }, - { - desc: "update group tags with invalid content type", - token: validToken, - id: validID, - domainID: validID, - data: fmt.Sprintf(`{"tags":["%s"]}`, newTag), - contentType: "application/xml", - svcResp: validGroupResp, - status: http.StatusUnsupportedMediaType, - err: apiutil.ErrUnsupportedContentType, - }, - { - desc: "update group tags with service error", - token: validToken, - id: validID, - domainID: validID, - data: fmt.Sprintf(`{"tags":["%s"]}`, newTag), - contentType: contentType, - svcResp: groups.Group{}, - svcErr: svcerr.ErrAuthorization, - status: http.StatusForbidden, - err: svcerr.ErrAuthorization, - }, - { - desc: "update group with malformed request", - token: validToken, - id: validID, - domainID: validID, - contentType: contentType, - data: fmt.Sprintf(`{"tags":["%s"}`, newTag), - status: http.StatusBadRequest, - err: apiutil.ErrMalformedRequestBody, - }, - { - desc: "update group with empty id", - token: validToken, - id: "", - domainID: validID, - contentType: contentType, - data: fmt.Sprintf(`{"tags":["%s"]}`, newTag), - status: http.StatusBadRequest, - err: apiutil.ErrMissingID, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - client: gs.Client(), - method: http.MethodPatch, - url: fmt.Sprintf("%s/%s/groups/%s/tags", gs.URL, tc.domainID, tc.id), - contentType: tc.contentType, - token: tc.token, - body: strings.NewReader(tc.data), - } - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID + "_" + validID, UserID: validID, DomainID: validID} - } - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - svcCall := svc.On("UpdateGroupTags", mock.Anything, tc.session, groups.Group{ID: tc.id, Tags: []string{newTag}}).Return(tc.svcResp, tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - var errRes respBody - err = json.NewDecoder(res.Body).Decode(&errRes) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if errRes.Err != "" || errRes.Message != "" { - err = errors.Wrap(errors.New(errRes.Err), errors.New(errRes.Message)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestEnableGroupEndpoint(t *testing.T) { - gs, svc, authn := newGroupsServer() - defer gs.Close() - - cases := []struct { - desc string - token string - id string - domainID string - session smqauthn.Session - svcResp groups.Group - svcErr error - resp groups.Group - status int - authnErr error - err error - }{ - { - desc: "enable group successfully", - token: validToken, - domainID: validID, - id: validID, - svcResp: validGroupResp, - svcErr: nil, - resp: validGroupResp, - status: http.StatusOK, - err: nil, - }, - { - desc: "enable group with invalid token", - token: invalidToken, - session: smqauthn.Session{}, - domainID: validID, - id: validID, - authnErr: svcerr.ErrAuthentication, - status: http.StatusUnauthorized, - err: svcerr.ErrAuthentication, - }, - { - desc: "enable group with empty token", - token: "", - session: smqauthn.Session{}, - domainID: validID, - id: validID, - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "enable group with empty domainID", - token: validToken, - id: validID, - status: http.StatusBadRequest, - err: apiutil.ErrMissingDomainID, - }, - { - desc: "enable group with service error", - token: validToken, - id: validID, - domainID: validID, - svcResp: groups.Group{}, - svcErr: svcerr.ErrAuthorization, - status: http.StatusForbidden, - err: svcerr.ErrAuthorization, - }, - { - desc: "enable group with empty id", - token: validToken, - id: "", - domainID: validID, - status: http.StatusBadRequest, - err: apiutil.ErrMissingID, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - client: gs.Client(), - method: http.MethodPost, - url: fmt.Sprintf("%s/%s/groups/%s/enable", gs.URL, tc.domainID, tc.id), - token: tc.token, - } - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID + "_" + validID, UserID: validID, DomainID: validID} - } - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - svcCall := svc.On("EnableGroup", mock.Anything, tc.session, tc.id).Return(tc.svcResp, tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - var errRes respBody - err = json.NewDecoder(res.Body).Decode(&errRes) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if errRes.Err != "" || errRes.Message != "" { - err = errors.Wrap(errors.New(errRes.Err), errors.New(errRes.Message)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestDisableGroupEndpoint(t *testing.T) { - gs, svc, authn := newGroupsServer() - defer gs.Close() - - cases := []struct { - desc string - token string - id string - domainID string - session smqauthn.Session - svcResp groups.Group - svcErr error - resp groups.Group - status int - authnErr error - err error - }{ - { - desc: "disable group successfully", - token: validToken, - domainID: validID, - id: validID, - svcResp: validGroupResp, - svcErr: nil, - resp: validGroupResp, - status: http.StatusOK, - err: nil, - }, - { - desc: "disable group with invalid token", - token: invalidToken, - session: smqauthn.Session{}, - domainID: validID, - id: validID, - authnErr: svcerr.ErrAuthentication, - status: http.StatusUnauthorized, - err: svcerr.ErrAuthentication, - }, - { - desc: "disable group with empty token", - token: "", - session: smqauthn.Session{}, - domainID: validID, - id: validID, - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "disable group with empty domainID", - token: validToken, - id: validID, - status: http.StatusBadRequest, - err: apiutil.ErrMissingDomainID, - }, - { - desc: "disable group with service error", - token: validToken, - id: validID, - domainID: validID, - svcResp: groups.Group{}, - svcErr: svcerr.ErrAuthorization, - status: http.StatusForbidden, - err: svcerr.ErrAuthorization, - }, - { - desc: "disable group with empty id", - token: validToken, - id: "", - domainID: validID, - status: http.StatusBadRequest, - err: apiutil.ErrMissingID, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - client: gs.Client(), - method: http.MethodPost, - url: fmt.Sprintf("%s/%s/groups/%s/disable", gs.URL, tc.domainID, tc.id), - token: tc.token, - } - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID + "_" + validID, UserID: validID, DomainID: validID} - } - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - svcCall := svc.On("DisableGroup", mock.Anything, tc.session, tc.id).Return(tc.svcResp, tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - var errRes respBody - err = json.NewDecoder(res.Body).Decode(&errRes) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if errRes.Err != "" || errRes.Message != "" { - err = errors.Wrap(errors.New(errRes.Err), errors.New(errRes.Message)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestListGroups(t *testing.T) { - gs, svc, authn := newGroupsServer() - defer gs.Close() - - cases := []struct { - desc string - query string - domainID string - token string - session smqauthn.Session - pageMeta groups.PageMeta - listGroupsResponse groups.Page - status int - authnErr error - err error - }{ - { - desc: "list groups successfully", - domainID: validID, - token: validToken, - status: http.StatusOK, - pageMeta: groups.PageMeta{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Actions: []string{}, - }, - listGroupsResponse: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 1, - }, - Groups: []groups.Group{validGroupResp}, - }, - err: nil, - }, - { - desc: "list groups with empty token", - domainID: validID, - token: "", - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "list groups with invalid token", - domainID: validID, - token: invalidToken, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "list groups with offset", - domainID: validID, - token: validToken, - pageMeta: groups.PageMeta{ - Offset: 1, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Actions: []string{}, - }, - listGroupsResponse: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 1, - }, - Groups: []groups.Group{validGroupResp}, - }, - query: "offset=1", - status: http.StatusOK, - err: nil, - }, - { - desc: "list groups with invalid offset", - domainID: validID, - token: validToken, - query: "offset=invalid", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list groups with limit", - domainID: validID, - token: validToken, - pageMeta: groups.PageMeta{ - Offset: 0, - Limit: 1, - Order: api.DefOrder, - Dir: api.DefDir, - Actions: []string{}, - }, - listGroupsResponse: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 1, - }, - Groups: []groups.Group{validGroupResp}, - }, - query: "limit=1", - status: http.StatusOK, - err: nil, - }, - { - desc: "list groups with invalid limit", - domainID: validID, - token: validToken, - query: "limit=invalid", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list groups with limit greater than max", - token: validToken, - domainID: validID, - query: fmt.Sprintf("limit=%d", api.MaxLimitSize+1), - status: http.StatusBadRequest, - err: apiutil.ErrLimitSize, - }, - { - desc: "list groups with name", - domainID: validID, - token: validToken, - pageMeta: groups.PageMeta{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Actions: []string{}, - Name: "clientname", - }, - listGroupsResponse: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 1, - }, - Groups: []groups.Group{validGroupResp}, - }, - query: "name=clientname", - status: http.StatusOK, - err: nil, - }, - { - desc: "list groups with duplicate name", - domainID: validID, - token: validToken, - query: "name=1&name=2", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list groups with status", - domainID: validID, - token: validToken, - pageMeta: groups.PageMeta{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Actions: []string{}, - Status: groups.EnabledStatus, - }, - listGroupsResponse: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 1, - }, - Groups: []groups.Group{validGroupResp}, - }, - query: "status=enabled", - status: http.StatusOK, - err: nil, - }, - { - desc: "list groups with invalid status", - domainID: validID, - token: validToken, - query: "status=invalid", - status: http.StatusBadRequest, - err: svcerr.ErrInvalidStatus, - }, - { - desc: "list groups with duplicate status", - domainID: validID, - token: validToken, - query: "status=enabled&status=disabled", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list groups with single tag", - domainID: validID, - token: validToken, - pageMeta: groups.PageMeta{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Actions: []string{}, - Tags: groups.TagsQuery{Elements: []string{"tag1"}, Operator: groups.OrOp}, - }, - listGroupsResponse: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 1, - }, - Groups: []groups.Group{validGroupResp}, - }, - query: "tags=tag1", - status: http.StatusOK, - err: nil, - }, - { - desc: "list groups with multiple tags and OR operator", - domainID: validID, - token: validToken, - pageMeta: groups.PageMeta{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Actions: []string{}, - Tags: groups.TagsQuery{Elements: []string{"tag1", "tag2", "tag3"}, Operator: groups.OrOp}, - }, - listGroupsResponse: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 1, - }, - Groups: []groups.Group{validGroupResp}, - }, - query: "tags=tag1,tag2,tag3", - status: http.StatusOK, - err: nil, - }, - { - desc: "list groups with multiple tags and AND operator", - domainID: validID, - token: validToken, - pageMeta: groups.PageMeta{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Actions: []string{}, - Tags: groups.TagsQuery{Elements: []string{"tag1", "tag2", "tag3"}, Operator: groups.AndOp}, - }, - listGroupsResponse: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 1, - }, - Groups: []groups.Group{validGroupResp}, - }, - query: "tags=tag1%2Btag2%2Btag3", - status: http.StatusOK, - err: nil, - }, - { - desc: "list groups with duplicate tags", - domainID: validID, - token: validToken, - query: "tags=tag1&tags=tag2", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list groups with metadata", - domainID: validID, - token: validToken, - pageMeta: groups.PageMeta{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Actions: []string{}, - Metadata: map[string]any{"domain": "example.com"}, - }, - listGroupsResponse: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 1, - }, - Groups: []groups.Group{validGroupResp}, - }, - query: fmt.Sprintf("metadata=%s", url.PathEscape(`{"domain": "example.com"}`)), - status: http.StatusOK, - err: nil, - }, - { - desc: "list groups with invalid metadata", - domainID: validID, - token: validToken, - query: "metadata=invalid", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list groups with duplicate metadata", - domainID: validID, - token: validToken, - query: fmt.Sprintf("metadata=%s&metadata=%s", url.PathEscape(`{"domain": "example.com"}`), url.PathEscape(`{"domain": "example.com"}`)), - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list groups with created_from", - domainID: validID, - token: validToken, - pageMeta: groups.PageMeta{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Actions: []string{}, - CreatedFrom: validTimeStamp, - }, - listGroupsResponse: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 1, - }, - Groups: []groups.Group{validGroupResp}, - }, - query: "created_from=2024-01-01T00:00:00Z", - status: http.StatusOK, - err: nil, - }, - { - desc: "list groups with created_to", - domainID: validID, - token: validToken, - pageMeta: groups.PageMeta{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Actions: []string{}, - CreatedTo: validTimeStamp, - }, - listGroupsResponse: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 1, - }, - Groups: []groups.Group{validGroupResp}, - }, - query: "created_to=2024-01-01T00:00:00Z", - status: http.StatusOK, - err: nil, - }, - { - desc: "list groups with both created_from and created_to", - domainID: validID, - token: validToken, - pageMeta: groups.PageMeta{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Actions: []string{}, - CreatedFrom: validTimeStamp, - CreatedTo: validTimeStamp, - }, - listGroupsResponse: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 1, - }, - Groups: []groups.Group{validGroupResp}, - }, - query: "created_from=2024-01-01T00:00:00Z&created_to=2024-01-01T00:00:00Z", - status: http.StatusOK, - err: nil, - }, - { - desc: "list groups with invalid created_from", - domainID: validID, - token: validToken, - query: "created_from=invalid-timestamp", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list groups with invalid created_to", - domainID: validID, - token: validToken, - query: "created_to=invalid-timestamp", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - client: gs.Client(), - method: http.MethodGet, - url: gs.URL + "/" + tc.domainID + "/groups?" + tc.query, - contentType: contentType, - token: tc.token, - } - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID + "_" + validID, UserID: validID, DomainID: validID} - } - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - svcCall := svc.On("ListGroups", mock.Anything, tc.session, tc.pageMeta).Return(tc.listGroupsResponse, tc.err) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - var bodyRes respBody - err = json.NewDecoder(res.Body).Decode(&bodyRes) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if bodyRes.Err != "" || bodyRes.Message != "" { - err = errors.Wrap(errors.New(bodyRes.Err), errors.New(bodyRes.Message)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestDeleteGroupEndpoint(t *testing.T) { - gs, svc, authn := newGroupsServer() - defer gs.Close() - - cases := []struct { - desc string - token string - id string - domainID string - session smqauthn.Session - svcErr error - status int - authnErr error - err error - }{ - { - desc: "delete group successfully", - token: validToken, - domainID: validID, - id: validID, - svcErr: nil, - status: http.StatusNoContent, - err: nil, - }, - { - desc: "delete group with invalid token", - token: invalidToken, - session: smqauthn.Session{}, - domainID: validID, - id: validID, - authnErr: svcerr.ErrAuthentication, - status: http.StatusUnauthorized, - err: svcerr.ErrAuthentication, - }, - { - desc: "delete group with empty token", - token: "", - session: smqauthn.Session{}, - domainID: validID, - id: validID, - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "delete group with empty domainID", - token: validToken, - id: validID, - status: http.StatusBadRequest, - err: apiutil.ErrMissingDomainID, - }, - { - desc: "delete group with service error", - token: validToken, - id: validID, - domainID: validID, - svcErr: svcerr.ErrAuthorization, - status: http.StatusForbidden, - err: svcerr.ErrAuthorization, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - client: gs.Client(), - method: http.MethodDelete, - url: fmt.Sprintf("%s/%s/groups/%s", gs.URL, tc.domainID, tc.id), - token: tc.token, - } - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID + "_" + validID, UserID: validID, DomainID: validID} - } - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - svcCall := svc.On("DeleteGroup", mock.Anything, tc.session, tc.id).Return(tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestRetrieveGroupHierarchyEndpoint(t *testing.T) { - gs, svc, authn := newGroupsServer() - defer gs.Close() - - retrieveHierarchyRes := groups.HierarchyPage{ - Groups: []groups.Group{validGroupResp}, - HierarchyPageMeta: groups.HierarchyPageMeta{ - Level: 1, - Direction: -1, - Tree: false, - }, - } - - treeHierarchyRes := groups.HierarchyPage{ - Groups: []groups.Group{validGroupResp}, - HierarchyPageMeta: groups.HierarchyPageMeta{ - Level: 1, - Direction: -1, - Tree: true, - }, - } - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - groupID string - query string - pageMeta groups.HierarchyPageMeta - svcRes groups.HierarchyPage - svcErr error - authnErr error - status int - err error - }{ - { - desc: "retrieve group hierarchy successfully", - token: validToken, - domainID: validID, - groupID: validID, - query: "level=1&dir=-1&tree=false", - pageMeta: groups.HierarchyPageMeta{ - Level: 1, - Direction: -1, - Tree: false, - }, - svcRes: retrieveHierarchyRes, - svcErr: nil, - status: http.StatusOK, - err: nil, - }, - { - desc: "retrieve group hierarchy successfully with tree", - token: validToken, - domainID: validID, - groupID: validID, - query: "level=1&dir=-1&tree=true", - pageMeta: groups.HierarchyPageMeta{ - Level: 1, - Direction: -1, - Tree: true, - }, - svcRes: treeHierarchyRes, - svcErr: nil, - status: http.StatusOK, - err: nil, - }, - { - desc: "retrieve group hierarchy with invalid token", - token: invalidToken, - session: smqauthn.Session{}, - domainID: validID, - groupID: validID, - query: "level=1&dir=-1&tree=false", - authnErr: svcerr.ErrAuthentication, - status: http.StatusUnauthorized, - err: svcerr.ErrAuthentication, - }, - { - desc: "retrieve group hierarchy with empty token", - token: "", - session: smqauthn.Session{}, - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "retrieve group hierarchy with empty domainID", - token: validToken, - groupID: validID, - query: "level=1&dir=-1&tree=false", - status: http.StatusBadRequest, - err: apiutil.ErrMissingDomainID, - }, - { - desc: "retrieve group hierarchy with service error", - token: validToken, - groupID: validID, - domainID: validID, - query: "level=1&dir=-1&tree=false", - pageMeta: groups.HierarchyPageMeta{ - Level: 1, - Direction: -1, - Tree: false, - }, - svcRes: groups.HierarchyPage{}, - svcErr: svcerr.ErrAuthorization, - status: http.StatusForbidden, - err: svcerr.ErrAuthorization, - }, - { - desc: "retrieve group hierarchy with invalid level", - token: validToken, - groupID: validID, - domainID: validID, - query: "level=invalid&dir=-1&tree=false", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "retrieve group hierarchy with invalid direction", - token: validToken, - groupID: validID, - domainID: validID, - query: "level=1&dir=invalid&tree=false", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "retrieve group hierarchy with invalid tree", - token: validToken, - groupID: validID, - domainID: validID, - query: "level=1&dir=-1&tree=invalid", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "retrieve group hierarchy with empty groupID", - token: validToken, - domainID: validID, - query: "level=1&dir=-1&tree=false", - status: http.StatusBadRequest, - err: apiutil.ErrMissingID, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - client: gs.Client(), - method: http.MethodGet, - url: fmt.Sprintf("%s/%s/groups/%s/hierarchy?%s", gs.URL, tc.domainID, tc.groupID, tc.query), - token: tc.token, - } - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID + "_" + validID, UserID: validID, DomainID: validID} - } - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - svcCall := svc.On("RetrieveGroupHierarchy", mock.Anything, tc.session, tc.groupID, tc.pageMeta).Return(tc.svcRes, tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - var bodyRes respBody - err = json.NewDecoder(res.Body).Decode(&bodyRes) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if bodyRes.Err != "" || bodyRes.Message != "" { - err = errors.Wrap(errors.New(bodyRes.Err), errors.New(bodyRes.Message)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestAddParentGroupEndpoint(t *testing.T) { - gs, svc, authn := newGroupsServer() - defer gs.Close() - - cases := []struct { - desc string - token string - id string - domainID string - parentID string - session smqauthn.Session - contentType string - svcErr error - status int - authnErr error - err error - }{ - { - desc: "add parent group successfully", - token: validToken, - domainID: validID, - id: validGroupResp.ID, - parentID: validID, - contentType: contentType, - svcErr: nil, - status: http.StatusOK, - err: nil, - }, - { - desc: "add parent group with invalid token", - token: invalidToken, - session: smqauthn.Session{}, - domainID: validID, - id: validGroupResp.ID, - parentID: validID, - contentType: contentType, - authnErr: svcerr.ErrAuthentication, - status: http.StatusUnauthorized, - err: svcerr.ErrAuthentication, - }, - { - desc: "add parent group with empty token", - token: "", - session: smqauthn.Session{}, - domainID: validID, - id: validGroupResp.ID, - parentID: validID, - contentType: contentType, - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "add parent group with empty domainID", - token: validToken, - id: validGroupResp.ID, - parentID: validID, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrMissingDomainID, - }, - { - desc: "add parent group with service error", - token: validToken, - id: validGroupResp.ID, - domainID: validID, - parentID: validID, - contentType: contentType, - svcErr: svcerr.ErrAuthorization, - status: http.StatusForbidden, - err: svcerr.ErrAuthorization, - }, - { - desc: "add parent group with empty id", - token: validToken, - id: "", - domainID: validID, - parentID: validID, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrMissingID, - }, - { - desc: "add parent group with empty parentID", - token: validToken, - id: validID, - domainID: validID, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrValidation, - }, - { - desc: "add self parenting group", - token: validToken, - id: validID, - domainID: validID, - parentID: validID, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrSelfParentingNotAllowed, - }, - { - desc: "add parent group with invalid content type", - token: validToken, - id: validID, - domainID: validID, - parentID: validID, - contentType: "application/xml", - status: http.StatusUnsupportedMediaType, - err: apiutil.ErrUnsupportedContentType, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - reqData := struct { - ParentID string `json:"parent_id"` - }{ - ParentID: tc.parentID, - } - data := toJSON(reqData) - req := testRequest{ - client: gs.Client(), - method: http.MethodPost, - url: fmt.Sprintf("%s/%s/groups/%s/parent", gs.URL, tc.domainID, tc.id), - token: tc.token, - contentType: tc.contentType, - body: strings.NewReader(data), - } - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID + "_" + validID, UserID: validID, DomainID: validID} - } - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - svcCall := svc.On("AddParentGroup", mock.Anything, tc.session, tc.id, tc.parentID).Return(tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestRemoveParentGroupEndpoint(t *testing.T) { - gs, svc, authn := newGroupsServer() - defer gs.Close() - - cases := []struct { - desc string - token string - id string - domainID string - session smqauthn.Session - svcErr error - status int - authnErr error - err error - }{ - { - desc: "remove parent group successfully", - token: validToken, - domainID: validID, - id: validGroupResp.ID, - svcErr: nil, - status: http.StatusNoContent, - err: nil, - }, - { - desc: "remove parent group with invalid token", - token: invalidToken, - session: smqauthn.Session{}, - domainID: validID, - id: validGroupResp.ID, - authnErr: svcerr.ErrAuthentication, - status: http.StatusUnauthorized, - err: svcerr.ErrAuthentication, - }, - { - desc: "remove parent group with empty token", - token: "", - session: smqauthn.Session{}, - domainID: validID, - id: validGroupResp.ID, - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "remove parent group with empty domainID", - token: validToken, - id: validGroupResp.ID, - status: http.StatusBadRequest, - err: apiutil.ErrMissingDomainID, - }, - { - desc: "remove parent group with service error", - token: validToken, - id: validGroupResp.ID, - domainID: validID, - svcErr: svcerr.ErrAuthorization, - status: http.StatusForbidden, - err: svcerr.ErrAuthorization, - }, - { - desc: "remove parent group with empty id", - token: validToken, - id: "", - domainID: validID, - status: http.StatusBadRequest, - err: apiutil.ErrMissingID, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - client: gs.Client(), - method: http.MethodDelete, - url: fmt.Sprintf("%s/%s/groups/%s/parent", gs.URL, tc.domainID, tc.id), - token: tc.token, - } - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID + "_" + validID, UserID: validID, DomainID: validID} - } - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - svcCall := svc.On("RemoveParentGroup", mock.Anything, tc.session, tc.id).Return(tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestAddChildrenGroupsEndpoint(t *testing.T) { - gs, svc, authn := newGroupsServer() - defer gs.Close() - - cases := []struct { - desc string - token string - id string - domainID string - childrenIDs []string - session smqauthn.Session - contentType string - svcErr error - status int - authnErr error - err error - }{ - { - desc: "add children groups successfully", - token: validToken, - domainID: validID, - id: validGroupResp.ID, - childrenIDs: []string{validID}, - contentType: contentType, - svcErr: nil, - status: http.StatusOK, - err: nil, - }, - { - desc: "add children groups with invalid token", - token: invalidToken, - session: smqauthn.Session{}, - domainID: validID, - id: validGroupResp.ID, - childrenIDs: []string{validID}, - contentType: contentType, - authnErr: svcerr.ErrAuthentication, - status: http.StatusUnauthorized, - err: svcerr.ErrAuthentication, - }, - { - desc: "add children groups with empty token", - token: "", - session: smqauthn.Session{}, - domainID: validID, - id: validGroupResp.ID, - childrenIDs: []string{validID}, - contentType: contentType, - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "add children groups with empty domainID", - token: validToken, - id: validGroupResp.ID, - childrenIDs: []string{validID}, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrMissingDomainID, - }, - { - desc: "add children groups with service error", - token: validToken, - id: validGroupResp.ID, - domainID: validID, - childrenIDs: []string{validID}, - contentType: contentType, - svcErr: svcerr.ErrAuthorization, - status: http.StatusForbidden, - err: svcerr.ErrAuthorization, - }, - { - desc: "add children groups with empty id", - token: validToken, - id: "", - domainID: validID, - childrenIDs: []string{validID}, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrMissingID, - }, - { - desc: "add children groups with empty childrenIDs", - token: validToken, - id: validGroupResp.ID, - domainID: validID, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrValidation, - }, - { - desc: "add children groups with invalid childrenIDs", - token: validToken, - id: validGroupResp.ID, - domainID: validID, - childrenIDs: []string{"invalid"}, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrValidation, - }, - { - desc: "add self children group", - token: validToken, - id: validID, - domainID: validID, - childrenIDs: []string{validID}, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrSelfParentingNotAllowed, - }, - { - desc: "add children groups with invalid content type", - token: validToken, - id: validGroupResp.ID, - domainID: validID, - childrenIDs: []string{validID}, - contentType: "application/xml", - status: http.StatusUnsupportedMediaType, - err: apiutil.ErrUnsupportedContentType, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - reqData := struct { - ChildrenIDs []string `json:"children_ids"` - }{ - ChildrenIDs: tc.childrenIDs, - } - data := toJSON(reqData) - req := testRequest{ - client: gs.Client(), - method: http.MethodPost, - url: fmt.Sprintf("%s/%s/groups/%s/children", gs.URL, tc.domainID, tc.id), - token: tc.token, - contentType: tc.contentType, - body: strings.NewReader(data), - } - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID + "_" + validID, UserID: validID, DomainID: validID} - } - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - svcCall := svc.On("AddChildrenGroups", mock.Anything, tc.session, tc.id, tc.childrenIDs).Return(tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestRemoveChildrenGroupsEndpoint(t *testing.T) { - gs, svc, authn := newGroupsServer() - defer gs.Close() - - cases := []struct { - desc string - token string - id string - domainID string - session smqauthn.Session - childrenIDs []string - contentType string - svcErr error - status int - authnErr error - err error - }{ - { - desc: "remove children groups successfully", - token: validToken, - domainID: validID, - id: validGroupResp.ID, - childrenIDs: []string{validID}, - contentType: contentType, - svcErr: nil, - status: http.StatusNoContent, - err: nil, - }, - { - desc: "remove children groups with invalid token", - token: invalidToken, - session: smqauthn.Session{}, - domainID: validID, - id: validGroupResp.ID, - childrenIDs: []string{validID}, - contentType: contentType, - authnErr: svcerr.ErrAuthentication, - status: http.StatusUnauthorized, - err: svcerr.ErrAuthentication, - }, - { - desc: "remove children groups with empty token", - token: "", - session: smqauthn.Session{}, - domainID: validID, - id: validGroupResp.ID, - childrenIDs: []string{validID}, - contentType: contentType, - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "remove children groups with empty domainID", - token: validToken, - id: validGroupResp.ID, - childrenIDs: []string{validID}, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrMissingDomainID, - }, - { - desc: "remove children groups with service error", - token: validToken, - id: validGroupResp.ID, - domainID: validID, - childrenIDs: []string{validID}, - contentType: contentType, - svcErr: svcerr.ErrAuthorization, - status: http.StatusForbidden, - err: svcerr.ErrAuthorization, - }, - { - desc: "remove children groups with empty id", - token: validToken, - id: "", - domainID: validID, - contentType: contentType, - childrenIDs: []string{validID}, - status: http.StatusBadRequest, - err: apiutil.ErrMissingID, - }, - { - desc: "remove children groups with empty childrenIDs", - token: validToken, - id: validGroupResp.ID, - domainID: validID, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrValidation, - }, - { - desc: "remove children groups with invalid childrenIDs", - token: validToken, - id: validGroupResp.ID, - domainID: validID, - childrenIDs: []string{"invalid"}, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrValidation, - }, - { - desc: "remove children groups with invalid content type", - token: validToken, - id: validGroupResp.ID, - domainID: validID, - childrenIDs: []string{validID}, - contentType: "application/xml", - status: http.StatusUnsupportedMediaType, - err: apiutil.ErrUnsupportedContentType, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - reqData := struct { - ChildrenIDs []string `json:"children_ids"` - }{ - ChildrenIDs: tc.childrenIDs, - } - data := toJSON(reqData) - req := testRequest{ - client: gs.Client(), - method: http.MethodDelete, - url: fmt.Sprintf("%s/%s/groups/%s/children", gs.URL, tc.domainID, tc.id), - token: tc.token, - contentType: tc.contentType, - body: strings.NewReader(data), - } - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID + "_" + validID, UserID: validID, DomainID: validID} - } - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - svcCall := svc.On("RemoveChildrenGroups", mock.Anything, tc.session, tc.id, tc.childrenIDs).Return(tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestRemoveAllChildrenGroupsEndpoint(t *testing.T) { - gs, svc, authn := newGroupsServer() - defer gs.Close() - - cases := []struct { - desc string - token string - id string - domainID string - session smqauthn.Session - svcErr error - status int - authnErr error - err error - }{ - { - desc: "remove all children groups successfully", - token: validToken, - domainID: validID, - id: validGroupResp.ID, - svcErr: nil, - status: http.StatusNoContent, - err: nil, - }, - { - desc: "remove all children groups with invalid token", - token: invalidToken, - session: smqauthn.Session{}, - domainID: validID, - id: validGroupResp.ID, - authnErr: svcerr.ErrAuthentication, - status: http.StatusUnauthorized, - err: svcerr.ErrAuthentication, - }, - { - desc: "remove all children groups with empty token", - token: "", - session: smqauthn.Session{}, - domainID: validID, - id: validGroupResp.ID, - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "remove all children groups with empty domainID", - token: validToken, - id: validGroupResp.ID, - status: http.StatusBadRequest, - err: apiutil.ErrMissingDomainID, - }, - { - desc: "remove all children groups with service error", - token: validToken, - id: validGroupResp.ID, - domainID: validID, - svcErr: svcerr.ErrAuthorization, - status: http.StatusForbidden, - err: svcerr.ErrAuthorization, - }, - { - desc: "remove all children groups with empty id", - token: validToken, - id: "", - domainID: validID, - status: http.StatusBadRequest, - err: apiutil.ErrMissingID, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - client: gs.Client(), - method: http.MethodDelete, - url: fmt.Sprintf("%s/%s/groups/%s/children/all", gs.URL, tc.domainID, tc.id), - token: tc.token, - } - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID + "_" + validID, UserID: validID, DomainID: validID} - } - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - svcCall := svc.On("RemoveAllChildrenGroups", mock.Anything, tc.session, tc.id).Return(tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestListChildrenGroupsEndpoint(t *testing.T) { - gs, svc, authn := newGroupsServer() - defer gs.Close() - - cases := []struct { - desc string - token string - id string - domainID string - session smqauthn.Session - query string - pageMeta groups.PageMeta - svcRes groups.Page - svcErr error - authnErr error - status int - err error - }{ - { - desc: "list children groups successfully", - token: validToken, - domainID: validID, - id: validGroupResp.ID, - query: "limit=1&offset=0", - pageMeta: groups.PageMeta{ - Limit: 1, - Offset: 0, - Actions: []string{}, - Dir: "desc", - Order: "updated_at", - }, - svcRes: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 1, - }, - Groups: []groups.Group{validGroupResp}, - }, - svcErr: nil, - status: http.StatusOK, - err: nil, - }, - { - desc: "list children groups with invalid token", - token: invalidToken, - session: smqauthn.Session{}, - domainID: validID, - id: validGroupResp.ID, - query: "limit=1&offset=0", - pageMeta: groups.PageMeta{ - Limit: 1, - Offset: 0, - Actions: []string{}, - Dir: "desc", - Order: "updated_at", - }, - authnErr: svcerr.ErrAuthentication, - status: http.StatusUnauthorized, - err: svcerr.ErrAuthentication, - }, - { - desc: "list children groups with empty token", - token: "", - session: smqauthn.Session{}, - domainID: validID, - id: validGroupResp.ID, - query: "limit=1&offset=0", - pageMeta: groups.PageMeta{ - Limit: 1, - Offset: 0, - Actions: []string{}, - Dir: "desc", - Order: "updated_at", - }, - status: http.StatusUnauthorized, - err: apiutil.ErrBearerToken, - }, - { - desc: "list children groups with empty domainID", - token: validToken, - id: validGroupResp.ID, - query: "limit=1&offset=0", - status: http.StatusBadRequest, - err: apiutil.ErrMissingDomainID, - }, - { - desc: "list children groups with service error", - token: validToken, - id: validGroupResp.ID, - domainID: validID, - query: "limit=1&offset=0", - pageMeta: groups.PageMeta{ - Limit: 1, - Offset: 0, - Actions: []string{}, - Dir: "desc", - Order: "updated_at", - }, - svcRes: groups.Page{}, - svcErr: svcerr.ErrAuthorization, - status: http.StatusForbidden, - err: svcerr.ErrAuthorization, - }, - { - desc: "list children groups with invalid limit", - token: validToken, - id: validGroupResp.ID, - domainID: validID, - query: "limit=invalid&offset=0", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list children groups with invalid offset", - token: validToken, - id: validGroupResp.ID, - domainID: validID, - query: "limit=1&offset=invalid", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list children groups with empty id", - token: validToken, - domainID: validID, - query: "limit=1&offset=0", - status: http.StatusBadRequest, - err: apiutil.ErrMissingID, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - client: gs.Client(), - method: http.MethodGet, - url: fmt.Sprintf("%s/%s/groups/%s/children?%s", gs.URL, tc.domainID, tc.id, tc.query), - token: tc.token, - } - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID + "_" + validID, UserID: validID, DomainID: validID} - } - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authnErr) - svcCall := svc.On("ListChildrenGroups", mock.Anything, tc.session, tc.id, int64(1), int64(0), tc.pageMeta).Return(tc.svcRes, tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - var bodyRes respBody - err = json.NewDecoder(res.Body).Decode(&bodyRes) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if bodyRes.Err != "" || bodyRes.Message != "" { - err = errors.Wrap(errors.New(bodyRes.Err), errors.New(bodyRes.Message)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -type testRequest struct { - client *http.Client - method string - url string - contentType string - token string - body io.Reader -} - -func (tr testRequest) make() (*http.Response, error) { - req, err := http.NewRequest(tr.method, tr.url, tr.body) - if err != nil { - return nil, err - } - - if tr.token != "" { - req.Header.Set("Authorization", apiutil.BearerPrefix+tr.token) - } - - if tr.contentType != "" { - req.Header.Set("Content-Type", tr.contentType) - } - - req.Header.Set("Referer", "http://localhost") - - return tr.client.Do(req) -} - -func toJSON(data any) string { - jsonData, err := json.Marshal(data) - if err != nil { - return "" - } - return string(jsonData) -} - -type respBody struct { - Err string `json:"error"` - Message string `json:"message"` - Total int `json:"total"` - Permissions []string `json:"permissions"` - ID string `json:"id"` - Tags []string `json:"tags"` - Status groups.Status `json:"status"` -} diff --git a/groups/api/http/endpoints.go b/groups/api/http/endpoints.go deleted file mode 100644 index 3f087a138..000000000 --- a/groups/api/http/endpoints.go +++ /dev/null @@ -1,410 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package api - -import ( - "context" - - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/groups" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/go-kit/kit/endpoint" -) - -func CreateGroupEndpoint(svc groups.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(createGroupReq) - if err := req.validate(); err != nil { - return createGroupRes{created: false}, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return createGroupRes{created: false}, svcerr.ErrAuthentication - } - - group, _, err := svc.CreateGroup(ctx, session, req.Group) - if err != nil { - return createGroupRes{created: false}, err - } - - return createGroupRes{created: true, Group: group}, nil - } -} - -func ViewGroupEndpoint(svc groups.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(groupReq) - if err := req.validate(); err != nil { - return viewGroupRes{}, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return viewGroupRes{}, svcerr.ErrAuthentication - } - - group, err := svc.ViewGroup(ctx, session, req.id, req.roles) - if err != nil { - return viewGroupRes{}, err - } - - return viewGroupRes{Group: group}, nil - } -} - -func UpdateGroupEndpoint(svc groups.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(updateGroupReq) - if err := req.validate(); err != nil { - return updateGroupRes{}, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return updateGroupRes{}, svcerr.ErrAuthentication - } - - group := groups.Group{ - ID: req.id, - Name: req.Name, - Description: req.Description, - Metadata: req.Metadata, - } - - group, err := svc.UpdateGroup(ctx, session, group) - if err != nil { - return updateGroupRes{}, err - } - - return updateGroupRes{Group: group}, nil - } -} - -func updateGroupTagsEndpoint(svc groups.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(updateGroupTagsReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - g := groups.Group{ - ID: req.id, - Tags: req.Tags, - } - g, err := svc.UpdateGroupTags(ctx, session, g) - if err != nil { - return nil, err - } - - return updateGroupRes{Group: g}, nil - } -} - -func EnableGroupEndpoint(svc groups.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(changeGroupStatusReq) - if err := req.validate(); err != nil { - return changeStatusRes{}, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return changeStatusRes{}, svcerr.ErrAuthentication - } - - group, err := svc.EnableGroup(ctx, session, req.id) - if err != nil { - return changeStatusRes{}, err - } - return changeStatusRes{Group: group}, nil - } -} - -func DisableGroupEndpoint(svc groups.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(changeGroupStatusReq) - if err := req.validate(); err != nil { - return changeStatusRes{}, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return changeStatusRes{}, svcerr.ErrAuthentication - } - - group, err := svc.DisableGroup(ctx, session, req.id) - if err != nil { - return changeStatusRes{}, err - } - return changeStatusRes{Group: group}, nil - } -} - -func ListGroupsEndpoint(svc groups.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(listGroupsReq) - - if err := req.validate(); err != nil { - return groupPageRes{}, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return groupPageRes{}, svcerr.ErrAuthentication - } - - var page groups.Page - var err error - switch { - case req.userID != "": - page, err = svc.ListUserGroups(ctx, session, req.userID, req.PageMeta) - default: - page, err = svc.ListGroups(ctx, session, req.PageMeta) - } - if err != nil { - return groupPageRes{}, err - } - - groups := []viewGroupRes{} - for _, g := range page.Groups { - groups = append(groups, toViewGroupRes(g)) - } - - return groupPageRes{ - pageRes: pageRes{ - Limit: page.Limit, - Offset: page.Offset, - Total: page.Total, - }, - Groups: groups, - }, nil - } -} - -func DeleteGroupEndpoint(svc groups.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(groupReq) - if err := req.validate(); err != nil { - return deleteGroupRes{}, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return deleteGroupRes{}, svcerr.ErrAuthentication - } - if err := svc.DeleteGroup(ctx, session, req.id); err != nil { - return deleteGroupRes{}, err - } - return deleteGroupRes{deleted: true}, nil - } -} - -func retrieveGroupHierarchyEndpoint(svc groups.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(retrieveGroupHierarchyReq) - if err := req.validate(); err != nil { - return retrieveGroupHierarchyRes{}, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return changeStatusRes{}, svcerr.ErrAuthentication - } - - hp, err := svc.RetrieveGroupHierarchy(ctx, session, req.id, req.HierarchyPageMeta) - if err != nil { - return retrieveGroupHierarchyRes{}, err - } - if req.HierarchyPageMeta.Tree { - return buildGroupsResponseTree(hp), nil - } - - groups := []viewGroupRes{} - for _, g := range hp.Groups { - groups = append(groups, toViewGroupRes(g)) - } - return retrieveGroupHierarchyRes{Level: hp.Level, Direction: hp.Direction, Groups: groups}, nil - } -} - -func addParentGroupEndpoint(svc groups.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(addParentGroupReq) - if err := req.validate(); err != nil { - return addParentGroupRes{}, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return changeStatusRes{}, svcerr.ErrAuthentication - } - - if err := svc.AddParentGroup(ctx, session, req.id, req.ParentID); err != nil { - return addParentGroupRes{}, err - } - return addParentGroupRes{}, nil - } -} - -func removeParentGroupEndpoint(svc groups.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(removeParentGroupReq) - if err := req.validate(); err != nil { - return removeParentGroupRes{}, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return changeStatusRes{}, svcerr.ErrAuthentication - } - - if err := svc.RemoveParentGroup(ctx, session, req.id); err != nil { - return removeParentGroupRes{}, err - } - return removeParentGroupRes{}, nil - } -} - -func addChildrenGroupsEndpoint(svc groups.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(addChildrenGroupsReq) - if err := req.validate(); err != nil { - return addChildrenGroupsRes{}, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return changeStatusRes{}, svcerr.ErrAuthentication - } - - if err := svc.AddChildrenGroups(ctx, session, req.id, req.ChildrenIDs); err != nil { - return addChildrenGroupsRes{}, err - } - return addChildrenGroupsRes{}, nil - } -} - -func removeChildrenGroupsEndpoint(svc groups.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(removeChildrenGroupsReq) - if err := req.validate(); err != nil { - return removeChildrenGroupsRes{}, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return changeStatusRes{}, svcerr.ErrAuthentication - } - - if err := svc.RemoveChildrenGroups(ctx, session, req.id, req.ChildrenIDs); err != nil { - return removeChildrenGroupsRes{}, err - } - return removeChildrenGroupsRes{}, nil - } -} - -func removeAllChildrenGroupsEndpoint(svc groups.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(removeAllChildrenGroupsReq) - if err := req.validate(); err != nil { - return removeAllChildrenGroupsRes{}, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return changeStatusRes{}, svcerr.ErrAuthentication - } - - if err := svc.RemoveAllChildrenGroups(ctx, session, req.id); err != nil { - return removeAllChildrenGroupsRes{}, err - } - return removeAllChildrenGroupsRes{}, nil - } -} - -func listChildrenGroupsEndpoint(svc groups.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(listChildrenGroupsReq) - if err := req.validate(); err != nil { - return listChildrenGroupsRes{}, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return changeStatusRes{}, svcerr.ErrAuthentication - } - - gp, err := svc.ListChildrenGroups(ctx, session, req.id, req.startLevel, req.endLevel, req.PageMeta) - if err != nil { - return listChildrenGroupsRes{}, err - } - viewGroups := []viewGroupRes{} - - for _, group := range gp.Groups { - viewGroups = append(viewGroups, toViewGroupRes(group)) - } - return listChildrenGroupsRes{ - pageRes: pageRes{ - Limit: gp.Limit, - Offset: gp.Offset, - Total: gp.Total, - }, - Groups: viewGroups, - }, nil - } -} - -func toViewGroupRes(group groups.Group) viewGroupRes { - view := viewGroupRes{ - Group: group, - } - return view -} - -func buildGroupsResponseTree(page groups.HierarchyPage) retrieveGroupHierarchyRes { - groupsMap := map[string]*groups.Group{} - parentsMap := map[string][]*groups.Group{} - for i := range page.Groups { - if _, ok := groupsMap[page.Groups[i].ID]; !ok { - groupsMap[page.Groups[i].ID] = &page.Groups[i] - parentsMap[page.Groups[i].ID] = make([]*groups.Group, 0) - } - } - - for _, group := range groupsMap { - if children, ok := parentsMap[group.Parent]; ok { - children = append(children, group) - parentsMap[group.Parent] = children - } - } - - res := retrieveGroupHierarchyRes{ - Level: page.Level, - Direction: page.Direction, - Groups: []viewGroupRes{}, - } - - for _, group := range groupsMap { - if children, ok := parentsMap[group.ID]; ok { - group.Children = children - } - } - - for _, group := range groupsMap { - view := toViewGroupRes(*group) - if children, ok := parentsMap[group.Parent]; len(children) == 0 || !ok { - res.Groups = append(res.Groups, view) - } - } - - return res -} diff --git a/groups/api/http/requests.go b/groups/api/http/requests.go deleted file mode 100644 index 6cf731ee0..000000000 --- a/groups/api/http/requests.go +++ /dev/null @@ -1,221 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package api - -import ( - api "github.com/absmach/magistrala/api/http" - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/groups" - "github.com/absmach/magistrala/internal/nullable" -) - -type createGroupReq struct { - groups.Group -} - -func (req createGroupReq) validate() error { - if len(req.Name) > api.MaxNameSize || req.Name == "" { - return apiutil.ErrNameSize - } - - return nil -} - -type updateGroupReq struct { - id string - Name string `json:"name,omitempty"` - Description nullable.Value[string] `json:"description,omitempty"` - Metadata map[string]any `json:"metadata,omitempty"` -} - -func (req updateGroupReq) validate() error { - if req.id == "" { - return apiutil.ErrMissingID - } - if len(req.Name) > api.MaxNameSize { - return apiutil.ErrNameSize - } - return nil -} - -type updateGroupTagsReq struct { - id string - Tags []string `json:"tags,omitempty"` -} - -func (req updateGroupTagsReq) validate() error { - if req.id == "" { - return apiutil.ErrMissingID - } - - return nil -} - -type listGroupsReq struct { - groups.PageMeta - userID string - groupID string -} - -func (req listGroupsReq) validate() error { - if req.Limit > api.MaxLimitSize || req.Limit < 1 { - return apiutil.ErrLimitSize - } - - if req.userID != "" && req.groupID != "" { - return apiutil.ErrMultipleEntitiesFilter - } - - switch req.Order { - case "", api.NameOrder, api.CreatedAtOrder, api.UpdatedAtOrder: - default: - return apiutil.ErrInvalidOrder - } - - if req.Dir != "" && (req.Dir != api.DescDir && req.Dir != api.AscDir) { - return apiutil.ErrInvalidDirection - } - - return nil -} - -type groupReq struct { - id string - roles bool -} - -func (req groupReq) validate() error { - if req.id == "" { - return apiutil.ErrMissingID - } - - return nil -} - -type changeGroupStatusReq struct { - id string -} - -func (req changeGroupStatusReq) validate() error { - if req.id == "" { - return apiutil.ErrMissingID - } - return nil -} - -type retrieveGroupHierarchyReq struct { - groups.HierarchyPageMeta - id string -} - -func (req retrieveGroupHierarchyReq) validate() error { - if req.Level > groups.MaxLevel { - return apiutil.ErrLevel - } - if req.id == "" { - return apiutil.ErrMissingID - } - - return nil -} - -type addParentGroupReq struct { - id string - ParentID string `json:"parent_id"` -} - -func (req addParentGroupReq) validate() error { - if req.id == "" { - return apiutil.ErrMissingID - } - if err := api.ValidateUUID(req.ParentID); err != nil { - return err - } - if req.id == req.ParentID { - return apiutil.ErrSelfParentingNotAllowed - } - return nil -} - -type removeParentGroupReq struct { - id string -} - -func (req removeParentGroupReq) validate() error { - if req.id == "" { - return apiutil.ErrMissingID - } - return nil -} - -type addChildrenGroupsReq struct { - id string - ChildrenIDs []string `json:"children_ids"` -} - -func (req addChildrenGroupsReq) validate() error { - if req.id == "" { - return apiutil.ErrMissingID - } - if len(req.ChildrenIDs) == 0 { - return apiutil.ErrMissingChildrenGroupIDs - } - for _, childID := range req.ChildrenIDs { - if err := api.ValidateUUID(childID); err != nil { - return err - } - if req.id == childID { - return apiutil.ErrSelfParentingNotAllowed - } - } - return nil -} - -type removeChildrenGroupsReq struct { - id string - ChildrenIDs []string `json:"children_ids"` -} - -func (req removeChildrenGroupsReq) validate() error { - if req.id == "" { - return apiutil.ErrMissingID - } - if len(req.ChildrenIDs) == 0 { - return apiutil.ErrMissingChildrenGroupIDs - } - for _, childID := range req.ChildrenIDs { - if err := api.ValidateUUID(childID); err != nil { - return err - } - } - return nil -} - -type removeAllChildrenGroupsReq struct { - id string -} - -func (req removeAllChildrenGroupsReq) validate() error { - if req.id == "" { - return apiutil.ErrMissingID - } - return nil -} - -type listChildrenGroupsReq struct { - id string - startLevel int64 - endLevel int64 - groups.PageMeta -} - -func (req listChildrenGroupsReq) validate() error { - if req.id == "" { - return apiutil.ErrMissingID - } - if req.Limit > api.MaxLimitSize || req.Limit < 1 { - return apiutil.ErrLimitSize - } - return nil -} diff --git a/groups/api/http/requests_test.go b/groups/api/http/requests_test.go deleted file mode 100644 index 8c41f8af2..000000000 --- a/groups/api/http/requests_test.go +++ /dev/null @@ -1,495 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package api - -import ( - "fmt" - "strings" - "testing" - - api "github.com/absmach/magistrala/api/http" - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/groups" - "github.com/absmach/magistrala/internal/testsutil" - "github.com/stretchr/testify/assert" -) - -var valid = "valid" - -func TestCreateGroupReqValidation(t *testing.T) { - cases := []struct { - desc string - req createGroupReq - err error - }{ - { - desc: "valid request", - req: createGroupReq{ - Group: groups.Group{ - Name: valid, - }, - }, - err: nil, - }, - { - desc: "long name", - req: createGroupReq{ - Group: groups.Group{ - Name: strings.Repeat("a", api.MaxNameSize+1), - }, - }, - err: apiutil.ErrNameSize, - }, - { - desc: "empty name", - req: createGroupReq{ - Group: groups.Group{}, - }, - err: apiutil.ErrNameSize, - }, - } - - for _, tc := range cases { - err := tc.req.validate() - assert.Equal(t, tc.err, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - } -} - -func TestUpdateGroupReqValidation(t *testing.T) { - cases := []struct { - desc string - req updateGroupReq - err error - }{ - { - desc: "valid request", - req: updateGroupReq{ - id: valid, - Name: valid, - }, - err: nil, - }, - { - desc: "long name", - req: updateGroupReq{ - id: valid, - Name: strings.Repeat("a", api.MaxNameSize+1), - }, - err: apiutil.ErrNameSize, - }, - { - desc: "empty id", - req: updateGroupReq{ - Name: valid, - }, - err: apiutil.ErrMissingID, - }, - } - - for _, tc := range cases { - err := tc.req.validate() - assert.Equal(t, tc.err, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - } -} - -func TestListGroupReqValidation(t *testing.T) { - cases := []struct { - desc string - req listGroupsReq - err error - }{ - { - desc: "valid request", - req: listGroupsReq{ - PageMeta: groups.PageMeta{ - Limit: 10, - }, - }, - err: nil, - }, - { - desc: "invalid lower limit", - req: listGroupsReq{ - PageMeta: groups.PageMeta{ - Limit: 0, - }, - }, - err: apiutil.ErrLimitSize, - }, - { - desc: "invalid upper limit", - req: listGroupsReq{ - PageMeta: groups.PageMeta{ - Limit: api.MaxLimitSize + 1, - }, - }, - err: apiutil.ErrLimitSize, - }, - } - - for _, tc := range cases { - err := tc.req.validate() - assert.Equal(t, tc.err, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - } -} - -func TestGroupReqValidation(t *testing.T) { - cases := []struct { - desc string - req groupReq - err error - }{ - { - desc: "valid request", - req: groupReq{ - id: valid, - }, - err: nil, - }, - - { - desc: "empty id", - req: groupReq{}, - err: apiutil.ErrMissingID, - }, - } - - for _, tc := range cases { - err := tc.req.validate() - assert.Equal(t, tc.err, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - } -} - -func TestChangeGroupStatusReqValidation(t *testing.T) { - cases := []struct { - desc string - req changeGroupStatusReq - err error - }{ - { - desc: "valid request", - req: changeGroupStatusReq{ - id: valid, - }, - err: nil, - }, - { - desc: "empty id", - req: changeGroupStatusReq{}, - err: apiutil.ErrMissingID, - }, - } - - for _, tc := range cases { - err := tc.req.validate() - assert.Equal(t, tc.err, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - } -} - -func TestRetrieveGroupHierarchyReqValidation(t *testing.T) { - cases := []struct { - desc string - req retrieveGroupHierarchyReq - err error - }{ - { - desc: "valid request", - req: retrieveGroupHierarchyReq{ - HierarchyPageMeta: groups.HierarchyPageMeta{ - Tree: true, - Level: 1, - Direction: -1, - }, - id: valid, - }, - }, - { - desc: "invalid level", - req: retrieveGroupHierarchyReq{ - HierarchyPageMeta: groups.HierarchyPageMeta{ - Tree: true, - Level: groups.MaxLevel + 1, - Direction: -1, - }, - id: valid, - }, - err: apiutil.ErrLevel, - }, - { - desc: "empty id", - req: retrieveGroupHierarchyReq{ - HierarchyPageMeta: groups.HierarchyPageMeta{ - Tree: true, - Level: 1, - Direction: -1, - }, - }, - err: apiutil.ErrMissingID, - }, - } - - for _, tc := range cases { - err := tc.req.validate() - assert.Equal(t, tc.err, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - } -} - -func TestAddParentGroupReqValidation(t *testing.T) { - cases := []struct { - desc string - req addParentGroupReq - err error - }{ - { - desc: "valid request", - req: addParentGroupReq{ - id: testsutil.GenerateUUID(t), - ParentID: testsutil.GenerateUUID(t), - }, - err: nil, - }, - { - desc: "empty id", - req: addParentGroupReq{ - ParentID: testsutil.GenerateUUID(t), - }, - err: apiutil.ErrMissingID, - }, - { - desc: "empty parent id", - req: addParentGroupReq{ - id: testsutil.GenerateUUID(t), - }, - err: apiutil.ErrInvalidIDFormat, - }, - { - desc: "invalid parent id", - req: addParentGroupReq{ - id: testsutil.GenerateUUID(t), - ParentID: "invalid", - }, - err: apiutil.ErrInvalidIDFormat, - }, - { - desc: "same id", - req: addParentGroupReq{ - id: validID, - ParentID: validID, - }, - err: apiutil.ErrSelfParentingNotAllowed, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := tc.req.validate() - assert.Equal(t, tc.err, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - }) - } -} - -func TestRemoveParentGroupReqValidation(t *testing.T) { - cases := []struct { - desc string - req removeParentGroupReq - err error - }{ - { - desc: "valid request", - req: removeParentGroupReq{ - id: testsutil.GenerateUUID(t), - }, - err: nil, - }, - { - desc: "empty id", - req: removeParentGroupReq{}, - err: apiutil.ErrMissingID, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := tc.req.validate() - assert.Equal(t, tc.err, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - }) - } -} - -func TestAddChildrenGroupsReqValidation(t *testing.T) { - cases := []struct { - desc string - req addChildrenGroupsReq - err error - }{ - { - desc: "valid request", - req: addChildrenGroupsReq{ - id: testsutil.GenerateUUID(t), - ChildrenIDs: []string{testsutil.GenerateUUID(t)}, - }, - err: nil, - }, - { - desc: "empty id", - req: addChildrenGroupsReq{ - ChildrenIDs: []string{testsutil.GenerateUUID(t)}, - }, - err: apiutil.ErrMissingID, - }, - { - desc: "empty children ids", - req: addChildrenGroupsReq{ - id: testsutil.GenerateUUID(t), - }, - err: apiutil.ErrMissingChildrenGroupIDs, - }, - { - desc: "invalid child id", - req: addChildrenGroupsReq{ - id: testsutil.GenerateUUID(t), - ChildrenIDs: []string{"invalid"}, - }, - err: apiutil.ErrInvalidIDFormat, - }, - { - desc: "self parenting", - req: addChildrenGroupsReq{ - id: validID, - ChildrenIDs: []string{validID, testsutil.GenerateUUID(t), testsutil.GenerateUUID(t)}, - }, - err: apiutil.ErrSelfParentingNotAllowed, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := tc.req.validate() - assert.Equal(t, tc.err, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - }) - } -} - -func TestRemoveChildrenGroupsReqValidation(t *testing.T) { - cases := []struct { - desc string - req removeChildrenGroupsReq - err error - }{ - { - desc: "valid request", - req: removeChildrenGroupsReq{ - id: testsutil.GenerateUUID(t), - ChildrenIDs: []string{testsutil.GenerateUUID(t)}, - }, - err: nil, - }, - { - desc: "empty id", - req: removeChildrenGroupsReq{}, - err: apiutil.ErrMissingID, - }, - { - desc: "empty children ids", - req: removeChildrenGroupsReq{ - id: testsutil.GenerateUUID(t), - }, - err: apiutil.ErrMissingChildrenGroupIDs, - }, - { - desc: "invalid child id", - req: removeChildrenGroupsReq{ - id: testsutil.GenerateUUID(t), - ChildrenIDs: []string{"invalid"}, - }, - err: apiutil.ErrInvalidIDFormat, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := tc.req.validate() - assert.Equal(t, tc.err, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - }) - } -} - -func TestRemoveAllChildrenGroupsReqValidation(t *testing.T) { - cases := []struct { - desc string - req removeAllChildrenGroupsReq - err error - }{ - { - desc: "valid request", - req: removeAllChildrenGroupsReq{ - id: testsutil.GenerateUUID(t), - }, - err: nil, - }, - { - desc: "empty id", - req: removeAllChildrenGroupsReq{}, - err: apiutil.ErrMissingID, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := tc.req.validate() - assert.Equal(t, tc.err, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - }) - } -} - -func TestListChildrenGroupsReqValidation(t *testing.T) { - cases := []struct { - desc string - req listChildrenGroupsReq - err error - }{ - { - desc: "valid request", - req: listChildrenGroupsReq{ - id: validID, - PageMeta: groups.PageMeta{ - Limit: 10, - }, - }, - err: nil, - }, - { - desc: "empty id", - req: listChildrenGroupsReq{}, - err: apiutil.ErrMissingID, - }, - { - desc: "invalid lower limit", - req: listChildrenGroupsReq{ - id: validID, - PageMeta: groups.PageMeta{ - Limit: 0, - }, - }, - err: apiutil.ErrLimitSize, - }, - { - desc: "invalid upper limit", - req: listChildrenGroupsReq{ - id: validID, - PageMeta: groups.PageMeta{ - Limit: api.MaxLimitSize + 1, - }, - }, - err: apiutil.ErrLimitSize, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := tc.req.validate() - assert.Equal(t, tc.err, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - }) - } -} diff --git a/groups/api/http/responses.go b/groups/api/http/responses.go deleted file mode 100644 index 748303994..000000000 --- a/groups/api/http/responses.go +++ /dev/null @@ -1,250 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package api - -import ( - "fmt" - "net/http" - - "github.com/absmach/magistrala" - "github.com/absmach/magistrala/groups" -) - -var ( - _ magistrala.Response = (*createGroupRes)(nil) - _ magistrala.Response = (*groupPageRes)(nil) - _ magistrala.Response = (*changeStatusRes)(nil) - _ magistrala.Response = (*viewGroupRes)(nil) - _ magistrala.Response = (*updateGroupRes)(nil) - _ magistrala.Response = (*retrieveGroupHierarchyRes)(nil) - _ magistrala.Response = (*addParentGroupRes)(nil) - _ magistrala.Response = (*removeParentGroupRes)(nil) - _ magistrala.Response = (*addChildrenGroupsRes)(nil) - _ magistrala.Response = (*removeChildrenGroupsRes)(nil) - _ magistrala.Response = (*removeAllChildrenGroupsRes)(nil) - _ magistrala.Response = (*listChildrenGroupsRes)(nil) -) - -type viewGroupRes struct { - groups.Group `json:",inline"` -} - -func (res viewGroupRes) Code() int { - return http.StatusOK -} - -func (res viewGroupRes) Headers() map[string]string { - return map[string]string{} -} - -func (res viewGroupRes) Empty() bool { - return false -} - -type createGroupRes struct { - groups.Group `json:",inline"` - created bool -} - -func (res createGroupRes) Code() int { - if res.created { - return http.StatusCreated - } - - return http.StatusOK -} - -func (res createGroupRes) Headers() map[string]string { - if res.created { - return map[string]string{ - "Location": fmt.Sprintf("/groups/%s", res.ID), - } - } - - return map[string]string{} -} - -func (res createGroupRes) Empty() bool { - return false -} - -type groupPageRes struct { - pageRes - Groups []viewGroupRes `json:"groups,omitempty"` -} - -type pageRes struct { - Limit uint64 `json:"limit,omitempty"` - Offset uint64 `json:"offset,omitempty"` - Total uint64 `json:"total"` -} - -func (res groupPageRes) Code() int { - return http.StatusOK -} - -func (res groupPageRes) Headers() map[string]string { - return map[string]string{} -} - -func (res groupPageRes) Empty() bool { - return false -} - -type updateGroupRes struct { - groups.Group `json:",inline"` -} - -func (res updateGroupRes) Code() int { - return http.StatusOK -} - -func (res updateGroupRes) Headers() map[string]string { - return map[string]string{} -} - -func (res updateGroupRes) Empty() bool { - return false -} - -type changeStatusRes struct { - groups.Group `json:",inline"` -} - -func (res changeStatusRes) Code() int { - return http.StatusOK -} - -func (res changeStatusRes) Headers() map[string]string { - return map[string]string{} -} - -func (res changeStatusRes) Empty() bool { - return false -} - -type deleteGroupRes struct { - deleted bool -} - -func (res deleteGroupRes) Code() int { - if res.deleted { - return http.StatusNoContent - } - - return http.StatusBadRequest -} - -func (res deleteGroupRes) Headers() map[string]string { - return map[string]string{} -} - -func (res deleteGroupRes) Empty() bool { - return true -} - -type retrieveGroupHierarchyRes struct { - Level uint64 `json:"level"` - Direction int64 `json:"direction"` - Groups []viewGroupRes `json:"groups"` -} - -func (res retrieveGroupHierarchyRes) Code() int { - return http.StatusOK -} - -func (res retrieveGroupHierarchyRes) Headers() map[string]string { - return map[string]string{} -} - -func (res retrieveGroupHierarchyRes) Empty() bool { - return false -} - -type addParentGroupRes struct{} - -func (res addParentGroupRes) Code() int { - return http.StatusOK -} - -func (res addParentGroupRes) Headers() map[string]string { - return map[string]string{} -} - -func (res addParentGroupRes) Empty() bool { - return true -} - -type removeParentGroupRes struct{} - -func (res removeParentGroupRes) Code() int { - return http.StatusNoContent -} - -func (res removeParentGroupRes) Headers() map[string]string { - return map[string]string{} -} - -func (res removeParentGroupRes) Empty() bool { - return true -} - -type addChildrenGroupsRes struct{} - -func (res addChildrenGroupsRes) Code() int { - return http.StatusOK -} - -func (res addChildrenGroupsRes) Headers() map[string]string { - return map[string]string{} -} - -func (res addChildrenGroupsRes) Empty() bool { - return true -} - -type removeChildrenGroupsRes struct{} - -func (res removeChildrenGroupsRes) Code() int { - return http.StatusNoContent -} - -func (res removeChildrenGroupsRes) Headers() map[string]string { - return map[string]string{} -} - -func (res removeChildrenGroupsRes) Empty() bool { - return true -} - -type removeAllChildrenGroupsRes struct{} - -func (res removeAllChildrenGroupsRes) Code() int { - return http.StatusNoContent -} - -func (res removeAllChildrenGroupsRes) Headers() map[string]string { - return map[string]string{} -} - -func (res removeAllChildrenGroupsRes) Empty() bool { - return true -} - -type listChildrenGroupsRes struct { - pageRes - Groups []viewGroupRes `json:"groups"` -} - -func (res listChildrenGroupsRes) Code() int { - return http.StatusOK -} - -func (res listChildrenGroupsRes) Headers() map[string]string { - return map[string]string{} -} - -func (res listChildrenGroupsRes) Empty() bool { - return false -} diff --git a/groups/api/http/transport.go b/groups/api/http/transport.go deleted file mode 100644 index ea09685d4..000000000 --- a/groups/api/http/transport.go +++ /dev/null @@ -1,151 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package api - -import ( - "log/slog" - - "github.com/absmach/magistrala" - api "github.com/absmach/magistrala/api/http" - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/groups" - smqauthn "github.com/absmach/magistrala/pkg/authn" - roleManagerHttp "github.com/absmach/magistrala/pkg/roles/rolemanager/api" - "github.com/go-chi/chi/v5" - kithttp "github.com/go-kit/kit/transport/http" - "github.com/prometheus/client_golang/prometheus/promhttp" - "go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp" -) - -// MakeHandler returns a HTTP handler for Groups API endpoints. -func MakeHandler(svc groups.Service, authn smqauthn.AuthNMiddleware, mux *chi.Mux, logger *slog.Logger, instanceID string, idp magistrala.IDProvider) *chi.Mux { - opts := []kithttp.ServerOption{ - kithttp.ServerErrorEncoder(apiutil.LoggingErrorEncoder(logger, api.EncodeError)), - } - d := roleManagerHttp.NewDecoder("groupID") - - mux.Route("/{domainID}/groups", func(r chi.Router) { - r.Use(authn.Middleware()) - r.Use(api.RequestIDMiddleware(idp)) - - r.Post("/", otelhttp.NewHandler(kithttp.NewServer( - CreateGroupEndpoint(svc), - DecodeGroupCreate, - api.EncodeResponse, - opts..., - ), "create_group").ServeHTTP) - - r.Get("/", otelhttp.NewHandler(kithttp.NewServer( - ListGroupsEndpoint(svc), - DecodeListGroupsRequest, - api.EncodeResponse, - opts..., - ), "list_groups").ServeHTTP) - r = roleManagerHttp.EntityAvailableActionsRouter(svc, d, r, opts) - - r.Route("/{groupID}", func(r chi.Router) { - r.Get("/", otelhttp.NewHandler(kithttp.NewServer( - ViewGroupEndpoint(svc), - DecodeGroupRequest, - api.EncodeResponse, - opts..., - ), "view_group").ServeHTTP) - - r.Put("/", otelhttp.NewHandler(kithttp.NewServer( - UpdateGroupEndpoint(svc), - DecodeGroupUpdate, - api.EncodeResponse, - opts..., - ), "update_group").ServeHTTP) - - r.Patch("/tags", otelhttp.NewHandler(kithttp.NewServer( - updateGroupTagsEndpoint(svc), - decodeUpdateGroupTags, - api.EncodeResponse, - opts..., - ), "update_group_tags").ServeHTTP) - - r.Delete("/", otelhttp.NewHandler(kithttp.NewServer( - DeleteGroupEndpoint(svc), - DecodeGroupRequest, - api.EncodeResponse, - opts..., - ), "delete_group").ServeHTTP) - - r.Post("/enable", otelhttp.NewHandler(kithttp.NewServer( - EnableGroupEndpoint(svc), - DecodeChangeGroupStatusRequest, - api.EncodeResponse, - opts..., - ), "enable_group").ServeHTTP) - - r.Post("/disable", otelhttp.NewHandler(kithttp.NewServer( - DisableGroupEndpoint(svc), - DecodeChangeGroupStatusRequest, - api.EncodeResponse, - opts..., - ), "disable_group").ServeHTTP) - - r = roleManagerHttp.EntityRoleMangerRouter(svc, d, r, opts) - - r.Get("/hierarchy", otelhttp.NewHandler(kithttp.NewServer( - retrieveGroupHierarchyEndpoint(svc), - decodeRetrieveGroupHierarchy, - api.EncodeResponse, - opts..., - ), "retrieve_group_hierarchy").ServeHTTP) - - r.Route("/parent", func(r chi.Router) { - r.Post("/", otelhttp.NewHandler(kithttp.NewServer( - addParentGroupEndpoint(svc), - decodeAddParentGroupRequest, - api.EncodeResponse, - opts..., - ), "add_parent_group").ServeHTTP) - - r.Delete("/", otelhttp.NewHandler(kithttp.NewServer( - removeParentGroupEndpoint(svc), - decodeRemoveParentGroupRequest, - api.EncodeResponse, - opts..., - ), "remove_parent_group").ServeHTTP) - }) - - r.Route("/children", func(r chi.Router) { - r.Post("/", otelhttp.NewHandler(kithttp.NewServer( - addChildrenGroupsEndpoint(svc), - decodeAddChildrenGroupsRequest, - api.EncodeResponse, - opts..., - ), "add_children_groups").ServeHTTP) - - r.Delete("/", otelhttp.NewHandler(kithttp.NewServer( - removeChildrenGroupsEndpoint(svc), - decodeRemoveChildrenGroupsRequest, - api.EncodeResponse, - opts..., - ), "remove_children_groups").ServeHTTP) - - r.Delete("/all", otelhttp.NewHandler(kithttp.NewServer( - removeAllChildrenGroupsEndpoint(svc), - decodeRemoveAllChildrenGroupsRequest, - api.EncodeResponse, - opts..., - ), "remove_all_children_groups").ServeHTTP) - - r.Get("/", otelhttp.NewHandler(kithttp.NewServer( - listChildrenGroupsEndpoint(svc), - decodeListChildrenGroupsRequest, - api.EncodeResponse, - opts..., - ), "list_children_groups").ServeHTTP) - }) - }) - }) - - mux.Get("/health", magistrala.Health("groups", instanceID)) - mux.Handle("/metrics", promhttp.Handler()) - - return mux -} diff --git a/groups/builtinroles.go b/groups/builtinroles.go deleted file mode 100644 index fc647d9eb..000000000 --- a/groups/builtinroles.go +++ /dev/null @@ -1,8 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package groups - -import "github.com/absmach/magistrala/pkg/roles" - -const BuiltInRoleAdmin roles.BuiltInRoleName = "admin" diff --git a/groups/doc.go b/groups/doc.go deleted file mode 100644 index 55e0840d6..000000000 --- a/groups/doc.go +++ /dev/null @@ -1,6 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package groups contains the domain concept definitions needed to support -// Magistrala groups functionality. -package groups diff --git a/groups/errors.go b/groups/errors.go deleted file mode 100644 index b6665fa0b..000000000 --- a/groups/errors.go +++ /dev/null @@ -1,17 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package groups - -import "errors" - -var ( - // ErrInvalidStatus indicates invalid status. - ErrInvalidStatus = errors.New("invalid groups status") - - // ErrEnableGroup indicates error in enabling group. - ErrEnableGroup = errors.New("failed to enable group") - - // ErrDisableGroup indicates error in disabling group. - ErrDisableGroup = errors.New("failed to disable group") -) diff --git a/groups/events/doc.go b/groups/events/doc.go deleted file mode 100644 index f1cd64cb7..000000000 --- a/groups/events/doc.go +++ /dev/null @@ -1,5 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package events contains event source Redis client implementation. -package events diff --git a/groups/events/events.go b/groups/events/events.go deleted file mode 100644 index dcce44e02..000000000 --- a/groups/events/events.go +++ /dev/null @@ -1,460 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package events - -import ( - "time" - - groups "github.com/absmach/magistrala/groups" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/events" - "github.com/absmach/magistrala/pkg/roles" -) - -const ( - groupPrefix = "group." - groupCreate = groupPrefix + "create" - groupUpdate = groupPrefix + "update" - groupUpdateTags = groupPrefix + "update_tags" - groupEnable = groupPrefix + "enable" - groupDisable = groupPrefix + "disable" - groupView = groupPrefix + "view" - groupList = groupPrefix + "list" - groupListUserGroups = groupPrefix + "list_user_groups" - groupRemove = groupPrefix + "remove" - groupRetrieveGroupHierarchy = groupPrefix + "retrieve_group_hierarchy" - groupAddParentGroup = groupPrefix + "add_parent_group" - groupRemoveParentGroup = groupPrefix + "remove_parent_group" - groupAddChildrenGroups = groupPrefix + "add_children_groups" - groupRemoveChildrenGroups = groupPrefix + "remove_children_groups" - groupRemoveAllChildrenGroups = groupPrefix + "remove_all_children_groups" - groupListChildrenGroups = groupPrefix + "list_children_groups" -) - -var ( - _ events.Event = (*createGroupEvent)(nil) - _ events.Event = (*updateGroupEvent)(nil) - _ events.Event = (*changeGroupStatusEvent)(nil) - _ events.Event = (*viewGroupEvent)(nil) - _ events.Event = (*deleteGroupEvent)(nil) - _ events.Event = (*viewGroupEvent)(nil) - _ events.Event = (*listGroupEvent)(nil) - _ events.Event = (*addParentGroupEvent)(nil) - _ events.Event = (*removeParentGroupEvent)(nil) - _ events.Event = (*addChildrenGroupsEvent)(nil) - _ events.Event = (*removeChildrenGroupsEvent)(nil) - _ events.Event = (*removeAllChildrenGroupsEvent)(nil) - _ events.Event = (*listChildrenGroupsEvent)(nil) - _ events.Event = (*retrieveGroupHierarchyEvent)(nil) -) - -type createGroupEvent struct { - groups.Group - rolesProvisioned []roles.RoleProvision - authn.Session - requestID string -} - -func (cge createGroupEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": groupCreate, - "id": cge.ID, - "roles_provisioned": cge.rolesProvisioned, - "status": cge.Status.String(), - "created_at": cge.CreatedAt, - "domain": cge.DomainID, - "user_id": cge.UserID, - "token_type": cge.Type.String(), - "super_admin": cge.SuperAdmin, - "request_id": cge.requestID, - } - - if cge.Parent != "" { - val["parent"] = cge.Parent - } - if cge.Name != "" { - val["name"] = cge.Name - } - if cge.Description.Valid { - val["description"] = cge.Description - } - if cge.Metadata != nil { - val["metadata"] = cge.Metadata - } - if cge.Status.String() != "" { - val["status"] = cge.Status.String() - } - - return val, nil -} - -type updateGroupEvent struct { - groups.Group - authn.Session - operation string - requestID string -} - -func (uge updateGroupEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": uge.operation, - "updated_at": uge.UpdatedAt, - "updated_by": uge.UpdatedBy, - "domain": uge.DomainID, - "user_id": uge.UserID, - "tags": uge.Tags, - "token_type": uge.Type.String(), - "super_admin": uge.SuperAdmin, - "request_id": uge.requestID, - } - - if uge.ID != "" { - val["id"] = uge.ID - } - if uge.Parent != "" { - val["parent"] = uge.Parent - } - if uge.Name != "" { - val["name"] = uge.Name - } - if uge.Description.Valid { - val["description"] = uge.Description - } - if uge.Metadata != nil { - val["metadata"] = uge.Metadata - } - if !uge.CreatedAt.IsZero() { - val["created_at"] = uge.CreatedAt - } - if uge.Status.String() != "" { - val["status"] = uge.Status.String() - } - - return val, nil -} - -type changeGroupStatusEvent struct { - id string - operation string - status string - updatedAt time.Time - updatedBy string - authn.Session - requestID string -} - -func (rge changeGroupStatusEvent) Encode() (map[string]any, error) { - return map[string]any{ - "operation": rge.operation, - "id": rge.id, - "status": rge.status, - "updated_at": rge.updatedAt, - "updated_by": rge.updatedBy, - "domain": rge.DomainID, - "user_id": rge.UserID, - "token_type": rge.Type.String(), - "super_admin": rge.SuperAdmin, - "request_id": rge.requestID, - }, nil -} - -type viewGroupEvent struct { - groups.Group - authn.Session - requestID string -} - -func (vge viewGroupEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": groupView, - "id": vge.ID, - "domain": vge.DomainID, - "user_id": vge.UserID, - "token_type": vge.Type.String(), - "super_admin": vge.SuperAdmin, - "request_id": vge.requestID, - } - - if vge.Parent != "" { - val["parent"] = vge.Parent - } - if vge.Name != "" { - val["name"] = vge.Name - } - if vge.Description.Valid { - val["description"] = vge.Description - } - if vge.Metadata != nil { - val["metadata"] = vge.Metadata - } - if !vge.CreatedAt.IsZero() { - val["created_at"] = vge.CreatedAt - } - if !vge.UpdatedAt.IsZero() { - val["updated_at"] = vge.UpdatedAt - } - if vge.UpdatedBy != "" { - val["updated_by"] = vge.UpdatedBy - } - if vge.Status.String() != "" { - val["status"] = vge.Status.String() - } - - return val, nil -} - -type listGroupEvent struct { - groups.PageMeta - domainID string - userID string - tokenType string - superAdmin bool - requestID string -} - -func (lge listGroupEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": groupList, - "total": lge.Total, - "offset": lge.Offset, - "limit": lge.Limit, - "domain": lge.domainID, - "user_id": lge.userID, - "token_type": lge.tokenType, - "super_admin": lge.superAdmin, - "request_id": lge.requestID, - } - - if lge.Name != "" { - val["name"] = lge.Name - } - if len(lge.Tags.Elements) > 0 { - val["tag"] = lge.Tags.Elements - } - if lge.Metadata != nil { - val["metadata"] = lge.Metadata - } - if lge.Status.String() != "" { - val["status"] = lge.Status.String() - } - - return val, nil -} - -type listUserGroupEvent struct { - userID string - domainID string - groups.PageMeta - tokenType string - superAdmin bool - requestID string -} - -func (luge listUserGroupEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": groupListUserGroups, - "user_id": luge.userID, - "domain": luge.domainID, - "total": luge.Total, - "offset": luge.Offset, - "limit": luge.Limit, - "token_type": luge.tokenType, - "super_admin": luge.superAdmin, - "request_id": luge.requestID, - } - - if luge.Name != "" { - val["name"] = luge.Name - } - if len(luge.Tags.Elements) > 0 { - val["tag"] = luge.Tags.Elements - } - if luge.Metadata != nil { - val["metadata"] = luge.Metadata - } - if luge.Status.String() != "" { - val["status"] = luge.Status.String() - } - - return val, nil -} - -type deleteGroupEvent struct { - id string - authn.Session - requestID string -} - -func (rge deleteGroupEvent) Encode() (map[string]any, error) { - return map[string]any{ - "operation": groupRemove, - "id": rge.id, - "domain": rge.DomainID, - "user_id": rge.UserID, - "token_type": rge.Type.String(), - "super_admin": rge.SuperAdmin, - "request_id": rge.requestID, - }, nil -} - -type retrieveGroupHierarchyEvent struct { - id string - groups.HierarchyPageMeta - authn.Session - requestID string -} - -func (vcge retrieveGroupHierarchyEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": groupRetrieveGroupHierarchy, - "id": vcge.id, - "level": vcge.Level, - "direction": vcge.Direction, - "tree": vcge.Tree, - "domain": vcge.DomainID, - "user_id": vcge.UserID, - "token_type": vcge.Type.String(), - "super_admin": vcge.SuperAdmin, - "request_id": vcge.requestID, - } - return val, nil -} - -type addParentGroupEvent struct { - id string - parentID string - authn.Session - requestID string -} - -func (apge addParentGroupEvent) Encode() (map[string]any, error) { - return map[string]any{ - "operation": groupAddParentGroup, - "id": apge.id, - "parent_id": apge.parentID, - "domain": apge.DomainID, - "user_id": apge.UserID, - "token_type": apge.Type.String(), - "super_admin": apge.SuperAdmin, - "request_id": apge.requestID, - }, nil -} - -type removeParentGroupEvent struct { - id string - authn.Session - requestID string -} - -func (rpge removeParentGroupEvent) Encode() (map[string]any, error) { - return map[string]any{ - "operation": groupRemoveParentGroup, - "id": rpge.id, - "domain": rpge.DomainID, - "user_id": rpge.UserID, - "token_type": rpge.Type.String(), - "super_admin": rpge.SuperAdmin, - "request_id": rpge.requestID, - }, nil -} - -type addChildrenGroupsEvent struct { - id string - childrenIDs []string - authn.Session - requestID string -} - -func (acge addChildrenGroupsEvent) Encode() (map[string]any, error) { - return map[string]any{ - "operation": groupAddChildrenGroups, - "id": acge.id, - "children_ids": acge.childrenIDs, - "domain": acge.DomainID, - "user_id": acge.UserID, - "token_type": acge.Type.String(), - "super_admin": acge.SuperAdmin, - "request_id": acge.requestID, - }, nil -} - -type removeChildrenGroupsEvent struct { - id string - childrenIDs []string - authn.Session - requestID string -} - -func (rcge removeChildrenGroupsEvent) Encode() (map[string]any, error) { - return map[string]any{ - "operation": groupRemoveChildrenGroups, - "id": rcge.id, - "children_ids": rcge.childrenIDs, - "domain": rcge.DomainID, - "user_id": rcge.UserID, - "token_type": rcge.Type.String(), - "super_admin": rcge.SuperAdmin, - "request_id": rcge.requestID, - }, nil -} - -type removeAllChildrenGroupsEvent struct { - id string - authn.Session - requestID string -} - -func (racge removeAllChildrenGroupsEvent) Encode() (map[string]any, error) { - return map[string]any{ - "operation": groupRemoveAllChildrenGroups, - "id": racge.id, - "domain": racge.DomainID, - "user_id": racge.UserID, - "token_type": racge.Type.String(), - "super_admin": racge.SuperAdmin, - "request_id": racge.requestID, - }, nil -} - -type listChildrenGroupsEvent struct { - id string - startLevel int64 - endLevel int64 - groups.PageMeta - domainID string - userID string - tokenType string - superAdmin bool - requestID string -} - -func (vcge listChildrenGroupsEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": groupListChildrenGroups, - "id": vcge.id, - "start_level": vcge.startLevel, - "end_level": vcge.endLevel, - "total": vcge.Total, - "offset": vcge.Offset, - "limit": vcge.Limit, - "domain": vcge.domainID, - "user_id": vcge.userID, - "token_type": vcge.tokenType, - "super_admin": vcge.superAdmin, - "request_id": vcge.requestID, - } - if vcge.Name != "" { - val["name"] = vcge.Name - } - if len(vcge.Tags.Elements) > 0 { - val["tag"] = vcge.Tags.Elements - } - if vcge.Metadata != nil { - val["metadata"] = vcge.Metadata - } - if vcge.Status.String() != "" { - val["status"] = vcge.Status.String() - } - return val, nil -} diff --git a/groups/events/streams.go b/groups/events/streams.go deleted file mode 100644 index 2fdb60ed1..000000000 --- a/groups/events/streams.go +++ /dev/null @@ -1,312 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package events - -import ( - "context" - - "github.com/absmach/magistrala/groups" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/events" - "github.com/absmach/magistrala/pkg/events/store" - "github.com/absmach/magistrala/pkg/roles" - rmEvents "github.com/absmach/magistrala/pkg/roles/rolemanager/events" - "github.com/go-chi/chi/v5/middleware" -) - -const ( - magistralaPrefix = "magistrala." - createStream = magistralaPrefix + groupCreate - updateStream = magistralaPrefix + groupUpdate - updateTagsStream = magistralaPrefix + groupUpdateTags - enableStream = magistralaPrefix + groupEnable - disableStream = magistralaPrefix + groupDisable - viewStream = magistralaPrefix + groupView - listStream = magistralaPrefix + groupList - listUserGroupsStream = magistralaPrefix + groupListUserGroups - removeStream = magistralaPrefix + groupRemove - retrieveHierarchyStream = magistralaPrefix + groupRetrieveGroupHierarchy - addParentStream = magistralaPrefix + groupAddParentGroup - removeParentStream = magistralaPrefix + groupRemoveParentGroup - addChildrenStream = magistralaPrefix + groupAddChildrenGroups - removeChildrenStream = magistralaPrefix + groupRemoveChildrenGroups - removeAllChildrenStream = magistralaPrefix + groupRemoveAllChildrenGroups - listChildrenStream = magistralaPrefix + groupListChildrenGroups -) - -var _ groups.Service = (*eventStore)(nil) - -type eventStore struct { - events.Publisher - svc groups.Service - rmEvents.RoleManagerEventStore -} - -// NewEventStoreMiddleware returns wrapper around clients service that sends -// events to event store. -func New(ctx context.Context, svc groups.Service, url string) (groups.Service, error) { - publisher, err := store.NewPublisher(ctx, url, "groups-es-pub") - if err != nil { - return nil, err - } - rmes := rmEvents.NewRoleManagerEventStore("groups", groupPrefix, svc, publisher) - - return &eventStore{ - svc: svc, - Publisher: publisher, - RoleManagerEventStore: rmes, - }, nil -} - -func (es eventStore) CreateGroup(ctx context.Context, session authn.Session, group groups.Group) (groups.Group, []roles.RoleProvision, error) { - group, rps, err := es.svc.CreateGroup(ctx, session, group) - if err != nil { - return group, rps, err - } - - event := createGroupEvent{ - Group: group, - rolesProvisioned: rps, - Session: session, - requestID: middleware.GetReqID(ctx), - } - - if err := es.Publish(ctx, createStream, event); err != nil { - return group, rps, err - } - - return group, rps, nil -} - -func (es eventStore) UpdateGroup(ctx context.Context, session authn.Session, group groups.Group) (groups.Group, error) { - group, err := es.svc.UpdateGroup(ctx, session, group) - if err != nil { - return group, err - } - - event := updateGroupEvent{ - Group: group, - Session: session, - operation: groupUpdate, - requestID: middleware.GetReqID(ctx), - } - - if err := es.Publish(ctx, updateStream, event); err != nil { - return group, err - } - - return group, nil -} - -func (es *eventStore) UpdateGroupTags(ctx context.Context, session authn.Session, g groups.Group) (groups.Group, error) { - g, err := es.svc.UpdateGroupTags(ctx, session, g) - if err != nil { - return g, err - } - - event := updateGroupEvent{ - Group: g, - Session: session, - operation: groupUpdateTags, - requestID: middleware.GetReqID(ctx), - } - if err := es.Publish(ctx, updateTagsStream, event); err != nil { - return g, err - } - - return g, nil -} - -func (es eventStore) ViewGroup(ctx context.Context, session authn.Session, id string, withRoles bool) (groups.Group, error) { - group, err := es.svc.ViewGroup(ctx, session, id, withRoles) - if err != nil { - return group, err - } - event := viewGroupEvent{ - group, - session, - middleware.GetReqID(ctx), - } - - if err := es.Publish(ctx, viewStream, event); err != nil { - return group, err - } - - return group, nil -} - -func (es eventStore) ListGroups(ctx context.Context, session authn.Session, pm groups.PageMeta) (groups.Page, error) { - gp, err := es.svc.ListGroups(ctx, session, pm) - if err != nil { - return gp, err - } - event := listGroupEvent{ - PageMeta: pm, - domainID: session.DomainID, - userID: session.UserID, - tokenType: session.Type.String(), - superAdmin: session.SuperAdmin, - requestID: middleware.GetReqID(ctx), - } - - if err := es.Publish(ctx, listStream, event); err != nil { - return gp, err - } - - return gp, nil -} - -func (es eventStore) ListUserGroups(ctx context.Context, session authn.Session, userID string, pm groups.PageMeta) (groups.Page, error) { - gp, err := es.svc.ListUserGroups(ctx, session, userID, pm) - if err != nil { - return gp, err - } - event := listUserGroupEvent{ - userID: userID, - PageMeta: pm, - domainID: session.DomainID, - tokenType: session.Type.String(), - superAdmin: session.SuperAdmin, - requestID: middleware.GetReqID(ctx), - } - - if err := es.Publish(ctx, listUserGroupsStream, event); err != nil { - return gp, err - } - - return gp, nil -} - -func (es eventStore) EnableGroup(ctx context.Context, session authn.Session, id string) (groups.Group, error) { - group, err := es.svc.EnableGroup(ctx, session, id) - if err != nil { - return group, err - } - - return es.changeStatus(ctx, session, groupEnable, enableStream, group) -} - -func (es eventStore) DisableGroup(ctx context.Context, session authn.Session, id string) (groups.Group, error) { - group, err := es.svc.DisableGroup(ctx, session, id) - if err != nil { - return group, err - } - - return es.changeStatus(ctx, session, groupDisable, disableStream, group) -} - -func (es eventStore) changeStatus(ctx context.Context, session authn.Session, operation, stream string, group groups.Group) (groups.Group, error) { - event := changeGroupStatusEvent{ - id: group.ID, - operation: operation, - updatedAt: group.UpdatedAt, - updatedBy: group.UpdatedBy, - status: group.Status.String(), - Session: session, - requestID: middleware.GetReqID(ctx), - } - - if err := es.Publish(ctx, stream, event); err != nil { - return group, err - } - - return group, nil -} - -func (es eventStore) DeleteGroup(ctx context.Context, session authn.Session, id string) error { - if err := es.svc.DeleteGroup(ctx, session, id); err != nil { - return err - } - if err := es.Publish(ctx, removeStream, deleteGroupEvent{ - id: id, - Session: session, - requestID: middleware.GetReqID(ctx), - }); err != nil { - return err - } - return nil -} - -func (es eventStore) RetrieveGroupHierarchy(ctx context.Context, session authn.Session, id string, hm groups.HierarchyPageMeta) (groups.HierarchyPage, error) { - g, err := es.svc.RetrieveGroupHierarchy(ctx, session, id, hm) - if err != nil { - return g, err - } - if err := es.Publish(ctx, retrieveHierarchyStream, retrieveGroupHierarchyEvent{id: id, Session: session, HierarchyPageMeta: hm, requestID: middleware.GetReqID(ctx)}); err != nil { - return g, err - } - return g, nil -} - -func (es eventStore) AddParentGroup(ctx context.Context, session authn.Session, id, parentID string) error { - if err := es.svc.AddParentGroup(ctx, session, id, parentID); err != nil { - return err - } - if err := es.Publish(ctx, addParentStream, addParentGroupEvent{id: id, parentID: parentID, Session: session, requestID: middleware.GetReqID(ctx)}); err != nil { - return err - } - return nil -} - -func (es eventStore) RemoveParentGroup(ctx context.Context, session authn.Session, id string) error { - if err := es.svc.RemoveParentGroup(ctx, session, id); err != nil { - return err - } - if err := es.Publish(ctx, removeParentStream, removeParentGroupEvent{id: id, Session: session, requestID: middleware.GetReqID(ctx)}); err != nil { - return err - } - return nil -} - -func (es eventStore) AddChildrenGroups(ctx context.Context, session authn.Session, id string, childrenGroupIDs []string) error { - if err := es.svc.AddChildrenGroups(ctx, session, id, childrenGroupIDs); err != nil { - return err - } - if err := es.Publish(ctx, addChildrenStream, addChildrenGroupsEvent{id: id, Session: session, childrenIDs: childrenGroupIDs, requestID: middleware.GetReqID(ctx)}); err != nil { - return err - } - return nil -} - -func (es eventStore) RemoveChildrenGroups(ctx context.Context, session authn.Session, id string, childrenGroupIDs []string) error { - if err := es.svc.RemoveChildrenGroups(ctx, session, id, childrenGroupIDs); err != nil { - return err - } - if err := es.Publish(ctx, removeChildrenStream, removeChildrenGroupsEvent{id: id, Session: session, childrenIDs: childrenGroupIDs, requestID: middleware.GetReqID(ctx)}); err != nil { - return err - } - - return nil -} - -func (es eventStore) RemoveAllChildrenGroups(ctx context.Context, session authn.Session, id string) error { - if err := es.svc.RemoveAllChildrenGroups(ctx, session, id); err != nil { - return err - } - if err := es.Publish(ctx, removeAllChildrenStream, removeAllChildrenGroupsEvent{id: id, Session: session, requestID: middleware.GetReqID(ctx)}); err != nil { - return err - } - return nil -} - -func (es eventStore) ListChildrenGroups(ctx context.Context, session authn.Session, id string, startLevel, endLevel int64, pm groups.PageMeta) (groups.Page, error) { - g, err := es.svc.ListChildrenGroups(ctx, session, id, startLevel, endLevel, pm) - if err != nil { - return g, err - } - if err := es.Publish(ctx, listChildrenStream, listChildrenGroupsEvent{ - id: id, - domainID: session.DomainID, - startLevel: startLevel, - endLevel: endLevel, - PageMeta: pm, - userID: session.UserID, - tokenType: session.Type.String(), - superAdmin: session.SuperAdmin, - requestID: middleware.GetReqID(ctx), - }); err != nil { - return g, err - } - return g, nil -} diff --git a/groups/events/streams_test.go b/groups/events/streams_test.go deleted file mode 100644 index c2b80ba93..000000000 --- a/groups/events/streams_test.go +++ /dev/null @@ -1,825 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package events_test - -import ( - "context" - "fmt" - "os" - "testing" - "time" - - "github.com/absmach/magistrala/groups" - "github.com/absmach/magistrala/groups/events" - "github.com/absmach/magistrala/groups/mocks" - "github.com/absmach/magistrala/internal/testsutil" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/pkg/roles" - "github.com/go-chi/chi/v5/middleware" - "github.com/redis/go-redis/v9" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" -) - -var ( - storeClient *redis.Client - storeURL string - validSession = authn.Session{ - DomainID: testsutil.GenerateUUID(&testing.T{}), - UserID: testsutil.GenerateUUID(&testing.T{}), - } - validGroup = generateTestGroup(&testing.T{}) - validGroupsPage = groups.Page{ - PageMeta: groups.PageMeta{ - Limit: 10, - Offset: 0, - Total: 1, - }, - Groups: []groups.Group{validGroup}, - } - validHierarchyPage = groups.HierarchyPage{ - HierarchyPageMeta: groups.HierarchyPageMeta{ - Level: 1, - Direction: -1, - Tree: false, - }, - Groups: []groups.Group{validGroup}, - } -) - -func newEventStoreMiddleware(t *testing.T) (*mocks.Service, groups.Service) { - svc := new(mocks.Service) - nsvc, err := events.New(context.Background(), svc, storeURL) - require.Nil(t, err, fmt.Sprintf("create events store middleware failed with unexpected error: %s", err)) - - return svc, nsvc -} - -func TestMain(m *testing.M) { - code := testsutil.RunRedisTest(m, &storeClient, &storeURL) - os.Exit(code) -} - -func TestCreateGroup(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validID := testsutil.GenerateUUID(t) - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, validID) - - cases := []struct { - desc string - session authn.Session - group groups.Group - svcRes groups.Group - svcRoleRes []roles.RoleProvision - svcErr error - resp groups.Group - respRoleRes []roles.RoleProvision - err error - }{ - { - desc: "publish successfully", - session: validSession, - group: validGroup, - svcRes: validGroup, - svcRoleRes: []roles.RoleProvision{}, - svcErr: nil, - resp: validGroup, - respRoleRes: []roles.RoleProvision{}, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - group: validGroup, - svcRes: groups.Group{}, - svcRoleRes: []roles.RoleProvision{}, - svcErr: svcerr.ErrCreateEntity, - resp: groups.Group{}, - respRoleRes: []roles.RoleProvision{}, - err: svcerr.ErrCreateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("CreateGroup", validCtx, tc.session, tc.group).Return(tc.svcRes, tc.svcRoleRes, tc.svcErr) - resp, respRoleRes, err := nsvc.CreateGroup(validCtx, tc.session, tc.group) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - assert.Equal(t, tc.respRoleRes, respRoleRes, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.respRoleRes, respRoleRes)) - svcCall.Unset() - }) - } -} - -func TestViewGroup(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - groupID string - withRoles bool - svcRes groups.Group - svcErr error - resp groups.Group - err error - }{ - { - desc: "publish successfully", - session: validSession, - groupID: validGroup.ID, - withRoles: false, - svcRes: validGroup, - svcErr: nil, - resp: validGroup, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - groupID: validGroup.ID, - withRoles: false, - svcRes: groups.Group{}, - svcErr: svcerr.ErrViewEntity, - resp: groups.Group{}, - err: svcerr.ErrViewEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("ViewGroup", validCtx, tc.session, tc.groupID, tc.withRoles).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.ViewGroup(validCtx, tc.session, tc.groupID, tc.withRoles) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestUpdateGroup(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - updatedGroup := validGroup - updatedGroup.Name = "updatedName" - - cases := []struct { - desc string - session authn.Session - group groups.Group - svcRes groups.Group - svcErr error - resp groups.Group - err error - }{ - { - desc: "publish successfully", - session: validSession, - group: updatedGroup, - svcRes: updatedGroup, - svcErr: nil, - resp: updatedGroup, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - group: updatedGroup, - svcRes: groups.Group{}, - svcErr: svcerr.ErrUpdateEntity, - resp: groups.Group{}, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("UpdateGroup", validCtx, tc.session, tc.group).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.UpdateGroup(validCtx, tc.session, tc.group) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestUpdateGroupTags(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - updatedGroup := validGroup - updatedGroup.Tags = []string{"newTag1", "newTag2"} - - cases := []struct { - desc string - session authn.Session - group groups.Group - svcRes groups.Group - svcErr error - resp groups.Group - err error - }{ - { - desc: "publish successfully", - session: validSession, - group: updatedGroup, - svcRes: updatedGroup, - svcErr: nil, - resp: updatedGroup, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - group: updatedGroup, - svcRes: groups.Group{}, - svcErr: svcerr.ErrUpdateEntity, - resp: groups.Group{}, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("UpdateGroupTags", validCtx, tc.session, tc.group).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.UpdateGroupTags(validCtx, tc.session, tc.group) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestEnableGroup(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - groupID string - svcRes groups.Group - svcErr error - resp groups.Group - err error - }{ - { - desc: "publish successfully", - session: validSession, - groupID: validGroup.ID, - svcRes: validGroup, - svcErr: nil, - resp: validGroup, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - groupID: validGroup.ID, - svcRes: groups.Group{}, - svcErr: svcerr.ErrUpdateEntity, - resp: groups.Group{}, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("EnableGroup", validCtx, tc.session, tc.groupID).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.EnableGroup(validCtx, tc.session, tc.groupID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestDisableGroup(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - groupID string - svcRes groups.Group - svcErr error - resp groups.Group - err error - }{ - { - desc: "publish successfully", - session: validSession, - groupID: validGroup.ID, - svcRes: validGroup, - svcErr: nil, - resp: validGroup, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - groupID: validGroup.ID, - svcRes: groups.Group{}, - svcErr: svcerr.ErrUpdateEntity, - resp: groups.Group{}, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("DisableGroup", validCtx, tc.session, tc.groupID).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.DisableGroup(validCtx, tc.session, tc.groupID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestListGroups(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - pageMeta groups.PageMeta - svcRes groups.Page - svcErr error - resp groups.Page - err error - }{ - { - desc: "publish successfully", - session: validSession, - pageMeta: groups.PageMeta{ - Limit: 10, - Offset: 0, - }, - svcRes: validGroupsPage, - svcErr: nil, - resp: validGroupsPage, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - pageMeta: groups.PageMeta{ - Limit: 10, - Offset: 0, - }, - svcRes: groups.Page{}, - svcErr: svcerr.ErrViewEntity, - resp: groups.Page{}, - err: svcerr.ErrViewEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("ListGroups", validCtx, tc.session, tc.pageMeta).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.ListGroups(validCtx, tc.session, tc.pageMeta) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestListUserGroups(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - userID string - pageMeta groups.PageMeta - svcRes groups.Page - svcErr error - resp groups.Page - err error - }{ - { - desc: "publish successfully", - session: validSession, - userID: validSession.UserID, - pageMeta: groups.PageMeta{ - Limit: 10, - Offset: 0, - }, - svcRes: validGroupsPage, - svcErr: nil, - resp: validGroupsPage, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - userID: validSession.UserID, - pageMeta: groups.PageMeta{ - Limit: 10, - Offset: 0, - }, - svcRes: groups.Page{}, - svcErr: svcerr.ErrViewEntity, - resp: groups.Page{}, - err: svcerr.ErrViewEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("ListUserGroups", validCtx, tc.session, tc.userID, tc.pageMeta).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.ListUserGroups(validCtx, tc.session, tc.userID, tc.pageMeta) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestDeleteGroup(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - groupID string - svcErr error - err error - }{ - { - desc: "publish successfully", - session: validSession, - groupID: validGroup.ID, - svcErr: nil, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - groupID: validGroup.ID, - svcErr: svcerr.ErrRemoveEntity, - err: svcerr.ErrRemoveEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("DeleteGroup", validCtx, tc.session, tc.groupID).Return(tc.svcErr) - err := nsvc.DeleteGroup(validCtx, tc.session, tc.groupID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - svcCall.Unset() - }) - } -} - -func TestRetrieveGroupHierarchy(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - groupID string - pageMeta groups.HierarchyPageMeta - svcRes groups.HierarchyPage - svcErr error - resp groups.HierarchyPage - err error - }{ - { - desc: "publish successfully", - session: validSession, - groupID: validGroup.ID, - pageMeta: groups.HierarchyPageMeta{ - Level: 1, - Direction: -1, - Tree: false, - }, - svcRes: validHierarchyPage, - svcErr: nil, - resp: validHierarchyPage, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - groupID: validGroup.ID, - pageMeta: groups.HierarchyPageMeta{ - Level: 1, - Direction: -1, - Tree: false, - }, - svcRes: groups.HierarchyPage{}, - svcErr: svcerr.ErrViewEntity, - resp: groups.HierarchyPage{}, - err: svcerr.ErrViewEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("RetrieveGroupHierarchy", validCtx, tc.session, tc.groupID, tc.pageMeta).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.RetrieveGroupHierarchy(validCtx, tc.session, tc.groupID, tc.pageMeta) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestAddParentGroup(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - groupID string - parentID string - svcErr error - err error - }{ - { - desc: "publish successfully", - session: validSession, - groupID: validGroup.ID, - parentID: testsutil.GenerateUUID(t), - svcErr: nil, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - groupID: validGroup.ID, - parentID: testsutil.GenerateUUID(t), - svcErr: svcerr.ErrUpdateEntity, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("AddParentGroup", validCtx, tc.session, tc.groupID, tc.parentID).Return(tc.svcErr) - err := nsvc.AddParentGroup(validCtx, tc.session, tc.groupID, tc.parentID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - svcCall.Unset() - }) - } -} - -func TestRemoveParentGroup(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - groupID string - svcErr error - err error - }{ - { - desc: "publish successfully", - session: validSession, - groupID: validGroup.ID, - svcErr: nil, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - groupID: validGroup.ID, - svcErr: svcerr.ErrUpdateEntity, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("RemoveParentGroup", validCtx, tc.session, tc.groupID).Return(tc.svcErr) - err := nsvc.RemoveParentGroup(validCtx, tc.session, tc.groupID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - svcCall.Unset() - }) - } -} - -func TestAddChildrenGroups(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - groupID string - childrenGroupIDs []string - svcErr error - err error - }{ - { - desc: "publish successfully", - session: validSession, - groupID: validGroup.ID, - childrenGroupIDs: []string{testsutil.GenerateUUID(t)}, - svcErr: nil, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - groupID: validGroup.ID, - childrenGroupIDs: []string{testsutil.GenerateUUID(t)}, - svcErr: svcerr.ErrUpdateEntity, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("AddChildrenGroups", validCtx, tc.session, tc.groupID, tc.childrenGroupIDs).Return(tc.svcErr) - err := nsvc.AddChildrenGroups(validCtx, tc.session, tc.groupID, tc.childrenGroupIDs) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - svcCall.Unset() - }) - } -} - -func TestRemoveChildrenGroups(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - groupID string - childrenGroupIDs []string - svcErr error - err error - }{ - { - desc: "publish successfully", - session: validSession, - groupID: validGroup.ID, - childrenGroupIDs: []string{testsutil.GenerateUUID(t)}, - svcErr: nil, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - groupID: validGroup.ID, - childrenGroupIDs: []string{testsutil.GenerateUUID(t)}, - svcErr: svcerr.ErrUpdateEntity, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("RemoveChildrenGroups", validCtx, tc.session, tc.groupID, tc.childrenGroupIDs).Return(tc.svcErr) - err := nsvc.RemoveChildrenGroups(validCtx, tc.session, tc.groupID, tc.childrenGroupIDs) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - svcCall.Unset() - }) - } -} - -func TestRemoveAllChildrenGroups(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - groupID string - svcErr error - err error - }{ - { - desc: "publish successfully", - session: validSession, - groupID: validGroup.ID, - svcErr: nil, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - groupID: validGroup.ID, - svcErr: svcerr.ErrUpdateEntity, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("RemoveAllChildrenGroups", validCtx, tc.session, tc.groupID).Return(tc.svcErr) - err := nsvc.RemoveAllChildrenGroups(validCtx, tc.session, tc.groupID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - svcCall.Unset() - }) - } -} - -func TestListChildrenGroups(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - groupID string - startLevel int64 - endLevel int64 - pageMeta groups.PageMeta - svcRes groups.Page - svcErr error - resp groups.Page - err error - }{ - { - desc: "publish successfully", - session: validSession, - groupID: validGroup.ID, - startLevel: 1, - endLevel: 5, - pageMeta: groups.PageMeta{ - Limit: 10, - Offset: 0, - }, - svcRes: validGroupsPage, - svcErr: nil, - resp: validGroupsPage, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - groupID: validGroup.ID, - startLevel: 1, - endLevel: 5, - pageMeta: groups.PageMeta{ - Limit: 10, - Offset: 0, - }, - svcRes: groups.Page{}, - svcErr: svcerr.ErrViewEntity, - resp: groups.Page{}, - err: svcerr.ErrViewEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("ListChildrenGroups", validCtx, tc.session, tc.groupID, tc.startLevel, tc.endLevel, tc.pageMeta).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.ListChildrenGroups(validCtx, tc.session, tc.groupID, tc.startLevel, tc.endLevel, tc.pageMeta) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func generateTestGroup(t *testing.T) groups.Group { - createdAt, err := time.Parse(time.RFC3339, "2024-01-01T00:00:00Z") - assert.Nil(t, err, fmt.Sprintf("Unexpected error parsing time: %v", err)) - return groups.Group{ - ID: testsutil.GenerateUUID(t), - Name: "groupname", - Domain: testsutil.GenerateUUID(t), - Tags: []string{"tag1", "tag2"}, - Metadata: groups.Metadata{"key1": "value1"}, - CreatedAt: createdAt, - UpdatedAt: createdAt, - Status: groups.EnabledStatus, - Level: 1, - } -} diff --git a/groups/groups.go b/groups/groups.go deleted file mode 100644 index 8313495a4..000000000 --- a/groups/groups.go +++ /dev/null @@ -1,184 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package groups - -import ( - "context" - "time" - - "github.com/absmach/magistrala/internal/nullable" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/roles" -) - -// MaxLevel represents the maximum group hierarchy level. -const ( - MaxLevel = uint64(20) - MaxPathLength = 20 -) - -// Metadata represents arbitrary JSON. -type Metadata map[string]any - -// Group represents the group of Clients. -// Indicates a level in tree hierarchy. Root node is level 1. -// Path in a tree consisting of group IDs -// Paths are unique per domain. -type Group struct { - ID string `json:"id"` - Domain string `json:"domain_id,omitempty"` - Parent string `json:"parent_id,omitempty"` - Name string `json:"name"` - Description nullable.Value[string] `json:"description,omitempty"` - Tags []string `json:"tags,omitempty"` - Metadata Metadata `json:"metadata,omitempty"` - Level int `json:"level,omitempty"` - Path string `json:"path,omitempty"` - Children []*Group `json:"children,omitempty"` - CreatedAt time.Time `json:"created_at"` - UpdatedAt time.Time `json:"updated_at,omitempty"` - UpdatedBy string `json:"updated_by,omitempty"` - Status Status `json:"status"` - RoleID string `json:"role_id,omitempty"` - RoleName string `json:"role_name,omitempty"` - Actions []string `json:"actions,omitempty"` - AccessType string `json:"access_type,omitempty"` - AccessProviderId string `json:"access_provider_id,omitempty"` - AccessProviderRoleId string `json:"access_provider_role_id,omitempty"` - AccessProviderRoleName string `json:"access_provider_role_name,omitempty"` - AccessProviderRoleActions []string `json:"access_provider_role_actions,omitempty"` - MemberId string `json:"member_id,omitempty"` - Roles []roles.MemberRoleActions `json:"roles,omitempty"` -} - -type Member struct { - ID string `json:"id"` - Type string `json:"type"` -} - -// Memberships contains page related metadata as well as list of memberships that -// belong to this page. -type MembersPage struct { - Total uint64 `json:"total"` - Offset uint64 `json:"offset"` - Limit uint64 `json:"limit"` - Members []Member `json:"members"` -} - -// Page contains page related metadata as well as list -// of Groups that belong to the page. -type Page struct { - PageMeta - Groups []Group -} - -type HierarchyPageMeta struct { - Level uint64 `json:"level"` - Direction int64 `json:"direction"` // ancestors (+1) or descendants (-1) - // - `true` - result is JSON tree representing groups hierarchy, - // - `false` - result is JSON array of groups. - Tree bool `json:"tree"` -} -type HierarchyPage struct { - HierarchyPageMeta - Groups []Group -} - -// Repository specifies a group persistence API. -type Repository interface { - // Save group. - Save(ctx context.Context, g Group) (Group, error) - - // Update a group. - Update(ctx context.Context, g Group) (Group, error) - - // Update a group's tags. - UpdateTags(ctx context.Context, g Group) (Group, error) - - // RetrieveByID retrieves group by its id. - RetrieveByID(ctx context.Context, id string) (Group, error) - - RetrieveByIDAndUser(ctx context.Context, domainID, userID, groupID string) (Group, error) - - RetrieveByIDWithRoles(ctx context.Context, groupID, memberID string) (Group, error) - - // RetrieveAll retrieves all groups. - RetrieveAll(ctx context.Context, pm PageMeta) (Page, error) - - // RetrieveByIDs retrieves group by ids and query. - RetrieveByIDs(ctx context.Context, pm PageMeta, ids ...string) (Page, error) - - RetrieveHierarchy(ctx context.Context, domainID, userID, groupID string, hm HierarchyPageMeta) (HierarchyPage, error) - - // ChangeStatus changes groups status to active or inactive - ChangeStatus(ctx context.Context, group Group) (Group, error) - - // AssignParentGroup assigns parent group id to a given group id - AssignParentGroup(ctx context.Context, parentGroupID string, groupIDs ...string) error - - // UnassignParentGroup unassign parent group id fr given group id - UnassignParentGroup(ctx context.Context, parentGroupID string, groupIDs ...string) error - - UnassignAllChildrenGroups(ctx context.Context, id string) error - - RetrieveUserGroups(ctx context.Context, domainID, userID string, pm PageMeta) (Page, error) - - // RetrieveChildrenGroups at given level in ltree - // Condition: startLevel == 0 and endLevel < 0, Retrieve all children groups from parent group level, Example: If we pass startLevel 0 and endLevel -1, then function will return all children of parent group - // Condition: startLevel > 0 and endLevel == 0, Retrieve specific level of children groups from parent group level, Example: If we pass startLevel 1 and endLevel 0, then function will return children of parent group from level 1 - // Condition: startLevel > 0 and endLevel < 0, Retrieve all children groups from specific level from parent group level, Example: If we pass startLevel 2 and endLevel -1, then function will return all children of parent group from level 2 - // Condition: startLevel > 0 and endLevel > 0, Retrieve children groups between specific level from parent group level, Example: If we pass startLevel 3 and endLevel 5, then function will return all children of parent group between level 3 and 5 - RetrieveChildrenGroups(ctx context.Context, domainID, userID, groupID string, startLevel, endLevel int64, pm PageMeta) (Page, error) - - RetrieveAllParentGroups(ctx context.Context, domainID, userID, groupID string, pm PageMeta) (Page, error) - // Delete a group - Delete(ctx context.Context, groupID string) error - - roles.Repository -} - -type Service interface { - // CreateGroup creates new group. - CreateGroup(ctx context.Context, session authn.Session, g Group) (Group, []roles.RoleProvision, error) - - // UpdateGroup updates the group identified by the provided ID. - UpdateGroup(ctx context.Context, session authn.Session, g Group) (Group, error) - - // UpdateGroupTags updates the groups's tags. - UpdateGroupTags(ctx context.Context, session authn.Session, group Group) (Group, error) - - // ViewGroup retrieves data about the group identified by ID. - ViewGroup(ctx context.Context, session authn.Session, id string, withRoles bool) (Group, error) - - // ListGroups retrieves groups for given filters. - ListGroups(ctx context.Context, session authn.Session, pm PageMeta) (Page, error) - - // ListGroups retrieves user accessible groups for given filters. - ListUserGroups(ctx context.Context, session authn.Session, userID string, pm PageMeta) (Page, error) - - // EnableGroup logically enables the group identified with the provided ID. - EnableGroup(ctx context.Context, session authn.Session, id string) (Group, error) - - // DisableGroup logically disables the group identified with the provided ID. - DisableGroup(ctx context.Context, session authn.Session, id string) (Group, error) - - // DeleteGroup delete the given group id - DeleteGroup(ctx context.Context, session authn.Session, id string) error - - RetrieveGroupHierarchy(ctx context.Context, session authn.Session, id string, hm HierarchyPageMeta) (HierarchyPage, error) - - AddParentGroup(ctx context.Context, session authn.Session, id, parentID string) error - - RemoveParentGroup(ctx context.Context, session authn.Session, id string) error - - AddChildrenGroups(ctx context.Context, session authn.Session, id string, childrenGroupIDs []string) error - - RemoveChildrenGroups(ctx context.Context, session authn.Session, id string, childrenGroupIDs []string) error - - RemoveAllChildrenGroups(ctx context.Context, session authn.Session, id string) error - - ListChildrenGroups(ctx context.Context, session authn.Session, id string, startLevel, endLevel int64, pm PageMeta) (Page, error) - - roles.RoleManager -} diff --git a/groups/middleware/authorization.go b/groups/middleware/authorization.go deleted file mode 100644 index d935a28d1..000000000 --- a/groups/middleware/authorization.go +++ /dev/null @@ -1,415 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package middleware - -import ( - "context" - "fmt" - - "github.com/absmach/magistrala/auth" - dOperations "github.com/absmach/magistrala/domains/operations" - "github.com/absmach/magistrala/groups" - "github.com/absmach/magistrala/groups/operations" - "github.com/absmach/magistrala/pkg/authn" - smqauthz "github.com/absmach/magistrala/pkg/authz" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/pkg/permissions" - "github.com/absmach/magistrala/pkg/policies" - "github.com/absmach/magistrala/pkg/roles" - rolemgr "github.com/absmach/magistrala/pkg/roles/rolemanager/middleware" -) - -var ( - errView = errors.New("not authorized to view group") - errUpdate = errors.New("not authorized to update group") - errUpdateTags = errors.New("not authorized to update group tags") - errEnable = errors.New("not authorized to enable group") - errDisable = errors.New("not authorized to disable group") - errDelete = errors.New("not authorized to delete group") - errViewHierarchy = errors.New("not authorized to view group parent/children hierarchy") - errListChildrenGroups = errors.New("not authorized to view chidden groups of group") - errSetParentGroup = errors.New("not authorized to set parent group to group") - errRemoveParentGroup = errors.New("not authorized to remove parent group from group") - errSetChildrenGroups = errors.New("not authorized to set children groups to group") - errRemoveChildrenGroups = errors.New("not authorized to remove children groups from group") - errParentGroupSetChildGroup = errors.New("not authorized to set child group in parent group") - errParentGroupRemoveChildGroup = errors.New("not authorized to remove child group from parent group") - errChildGroupSetParentGroup = errors.New("not authorized to set parent group to child group") - errDomainCreateGroups = errors.New("not authorized to create groups in domain") - errDomainListGroups = errors.New("not authorized to list groups in domain") -) - -var _ groups.Service = (*authorizationMiddleware)(nil) - -type authorizationMiddleware struct { - svc groups.Service - repo groups.Repository - authz smqauthz.Authorization - entitiesOps permissions.EntitiesOperations[permissions.Operation] - rolemgr.RoleManagerAuthorizationMiddleware -} - -// NewAuthorization adds authorization to the groups service. -func NewAuthorization( - entityType string, - svc groups.Service, - authz smqauthz.Authorization, - repo groups.Repository, - entitiesOps permissions.EntitiesOperations[permissions.Operation], - roleOps permissions.Operations[permissions.RoleOperation], -) (groups.Service, error) { - if err := entitiesOps.Validate(); err != nil { - return nil, err - } - ram, err := rolemgr.NewAuthorization(policies.GroupType, svc, authz, roleOps) - if err != nil { - return nil, err - } - - return &authorizationMiddleware{ - svc: svc, - authz: authz, - repo: repo, - entitiesOps: entitiesOps, - RoleManagerAuthorizationMiddleware: ram, - }, nil -} - -func (am *authorizationMiddleware) CreateGroup(ctx context.Context, session authn.Session, g groups.Group) (groups.Group, []roles.RoleProvision, error) { - if err := am.authorize(ctx, session, policies.DomainType, dOperations.OpCreateDomainGroups, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - SubjectKind: policies.UsersKind, - Subject: session.DomainUserID, - Object: session.DomainID, - ObjectType: policies.DomainType, - }); err != nil { - return groups.Group{}, []roles.RoleProvision{}, errors.Wrap(errDomainCreateGroups, err) - } - - if g.Parent != "" { - if err := am.authorize(ctx, session, policies.GroupType, operations.OpAddChildrenGroups, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - SubjectKind: policies.UsersKind, - Subject: session.DomainUserID, - Object: g.Parent, - ObjectType: policies.GroupType, - }); err != nil { - return groups.Group{}, []roles.RoleProvision{}, errors.Wrap(errParentGroupSetChildGroup, err) - } - } - - return am.svc.CreateGroup(ctx, session, g) -} - -func (am *authorizationMiddleware) UpdateGroup(ctx context.Context, session authn.Session, g groups.Group) (groups.Group, error) { - if err := am.authorize(ctx, session, policies.GroupType, operations.OpUpdateGroup, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - SubjectKind: policies.UsersKind, - Subject: session.DomainUserID, - Object: g.ID, - ObjectType: policies.GroupType, - }); err != nil { - return groups.Group{}, errors.Wrap(errUpdate, err) - } - - return am.svc.UpdateGroup(ctx, session, g) -} - -func (am *authorizationMiddleware) UpdateGroupTags(ctx context.Context, session authn.Session, group groups.Group) (groups.Group, error) { - if err := am.authorize(ctx, session, policies.GroupType, operations.OpUpdateGroupTags, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - Subject: session.DomainUserID, - ObjectType: policies.GroupType, - Object: group.ID, - }); err != nil { - return groups.Group{}, errors.Wrap(errUpdateTags, err) - } - - return am.svc.UpdateGroupTags(ctx, session, group) -} - -func (am *authorizationMiddleware) ViewGroup(ctx context.Context, session authn.Session, id string, withRoles bool) (groups.Group, error) { - if err := am.authorize(ctx, session, policies.GroupType, operations.OpViewGroup, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - SubjectKind: policies.UsersKind, - Subject: session.DomainUserID, - Object: id, - ObjectType: policies.GroupType, - }); err != nil { - return groups.Group{}, errors.Wrap(errView, err) - } - - return am.svc.ViewGroup(ctx, session, id, withRoles) -} - -func (am *authorizationMiddleware) ListGroups(ctx context.Context, session authn.Session, gm groups.PageMeta) (groups.Page, error) { - if err := am.checkSuperAdmin(ctx, session); err == nil { - session.SuperAdmin = true - return am.svc.ListGroups(ctx, session, gm) - } - if err := am.authorize(ctx, session, policies.DomainType, dOperations.OpListDomainGroups, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - SubjectKind: policies.UsersKind, - Subject: session.DomainUserID, - Object: session.DomainID, - ObjectType: policies.DomainType, - }); err != nil { - return groups.Page{}, errors.Wrap(errDomainListGroups, err) - } - - return am.svc.ListGroups(ctx, session, gm) -} - -func (am *authorizationMiddleware) ListUserGroups(ctx context.Context, session authn.Session, userID string, pm groups.PageMeta) (groups.Page, error) { - if err := am.checkSuperAdmin(ctx, session); err == nil { - session.SuperAdmin = true - return am.svc.ListGroups(ctx, session, pm) - } - if err := am.authorize(ctx, session, policies.DomainType, dOperations.OpListDomainGroups, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - SubjectKind: policies.UsersKind, - Subject: session.DomainUserID, - Object: session.DomainID, - ObjectType: policies.DomainType, - }); err != nil { - return groups.Page{}, errors.Wrap(errDomainListGroups, err) - } - - return am.svc.ListUserGroups(ctx, session, userID, pm) -} - -func (am *authorizationMiddleware) EnableGroup(ctx context.Context, session authn.Session, id string) (groups.Group, error) { - if err := am.authorize(ctx, session, policies.GroupType, operations.OpEnableGroup, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - Subject: session.DomainUserID, - Object: id, - ObjectType: policies.GroupType, - }); err != nil { - return groups.Group{}, errors.Wrap(errEnable, err) - } - - return am.svc.EnableGroup(ctx, session, id) -} - -func (am *authorizationMiddleware) DisableGroup(ctx context.Context, session authn.Session, id string) (groups.Group, error) { - if err := am.authorize(ctx, session, policies.GroupType, operations.OpDisableGroup, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - Subject: session.DomainUserID, - Object: id, - ObjectType: policies.GroupType, - }); err != nil { - return groups.Group{}, errors.Wrap(errDisable, err) - } - - return am.svc.DisableGroup(ctx, session, id) -} - -func (am *authorizationMiddleware) DeleteGroup(ctx context.Context, session authn.Session, id string) error { - if err := am.authorize(ctx, session, policies.GroupType, operations.OpDeleteGroup, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - Subject: session.DomainUserID, - Object: id, - ObjectType: policies.GroupType, - }); err != nil { - return errors.Wrap(errDelete, err) - } - - return am.svc.DeleteGroup(ctx, session, id) -} - -func (am *authorizationMiddleware) RetrieveGroupHierarchy(ctx context.Context, session authn.Session, id string, hm groups.HierarchyPageMeta) (groups.HierarchyPage, error) { - if err := am.authorize(ctx, session, policies.GroupType, operations.OpRetrieveGroupHierarchy, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - Subject: session.DomainUserID, - Object: id, - ObjectType: policies.GroupType, - }); err != nil { - return groups.HierarchyPage{}, errors.Wrap(errViewHierarchy, err) - } - - return am.svc.RetrieveGroupHierarchy(ctx, session, id, hm) -} - -func (am *authorizationMiddleware) AddParentGroup(ctx context.Context, session authn.Session, id, parentID string) error { - if err := am.authorize(ctx, session, policies.GroupType, operations.OpAddParentGroup, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - Subject: session.DomainUserID, - Object: id, - ObjectType: policies.GroupType, - }); err != nil { - return errors.Wrap(errSetParentGroup, err) - } - - if err := am.authorize(ctx, session, policies.GroupType, operations.OpAddChildrenGroups, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - Subject: session.DomainUserID, - Object: parentID, - ObjectType: policies.GroupType, - }); err != nil { - return errors.Wrap(errParentGroupSetChildGroup, err) - } - - return am.svc.AddParentGroup(ctx, session, id, parentID) -} - -func (am *authorizationMiddleware) RemoveParentGroup(ctx context.Context, session authn.Session, id string) error { - if err := am.authorize(ctx, session, policies.GroupType, operations.OpRemoveParentGroup, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - Subject: session.DomainUserID, - Object: id, - ObjectType: policies.GroupType, - }); err != nil { - return errors.Wrap(errRemoveParentGroup, err) - } - - group, err := am.repo.RetrieveByID(ctx, id) - if err != nil { - return errors.Wrap(svcerr.ErrViewEntity, err) - } - - if group.Parent != "" { - if err := am.authorize(ctx, session, policies.GroupType, operations.OpRemoveParentGroup, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - Subject: session.DomainUserID, - Object: group.Parent, - ObjectType: policies.GroupType, - }); err != nil { - return errors.Wrap(errParentGroupRemoveChildGroup, err) - } - } - - return am.svc.RemoveParentGroup(ctx, session, id) -} - -func (am *authorizationMiddleware) AddChildrenGroups(ctx context.Context, session authn.Session, id string, childrenGroupIDs []string) error { - if err := am.authorize(ctx, session, policies.GroupType, operations.OpAddChildrenGroups, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - Subject: session.DomainUserID, - Object: id, - ObjectType: policies.GroupType, - }); err != nil { - return errors.Wrap(errSetChildrenGroups, err) - } - - for _, childID := range childrenGroupIDs { - if err := am.authorize(ctx, session, policies.GroupType, operations.OpAddParentGroup, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - Subject: session.DomainUserID, - Object: childID, - ObjectType: policies.GroupType, - }); err != nil { - return errors.Wrap(errChildGroupSetParentGroup, errors.Wrap(fmt.Errorf("child group id: %s", childID), err)) - } - } - - return am.svc.AddChildrenGroups(ctx, session, id, childrenGroupIDs) -} - -func (am *authorizationMiddleware) RemoveChildrenGroups(ctx context.Context, session authn.Session, id string, childrenGroupIDs []string) error { - if err := am.authorize(ctx, session, policies.GroupType, operations.OpRemoveChildrenGroups, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - Subject: session.DomainUserID, - Object: id, - ObjectType: policies.GroupType, - }); err != nil { - return errors.Wrap(errRemoveChildrenGroups, err) - } - - return am.svc.RemoveChildrenGroups(ctx, session, id, childrenGroupIDs) -} - -func (am *authorizationMiddleware) RemoveAllChildrenGroups(ctx context.Context, session authn.Session, id string) error { - if err := am.authorize(ctx, session, policies.GroupType, operations.OpRemoveAllChildrenGroups, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - Subject: session.DomainUserID, - Object: id, - ObjectType: policies.GroupType, - }); err != nil { - return err - } - - return am.svc.RemoveAllChildrenGroups(ctx, session, id) -} - -func (am *authorizationMiddleware) ListChildrenGroups(ctx context.Context, session authn.Session, id string, startLevel, endLevel int64, pm groups.PageMeta) (groups.Page, error) { - if err := am.authorize(ctx, session, policies.GroupType, operations.OpListChildrenGroups, smqauthz.PolicyReq{ - Domain: session.DomainID, - SubjectType: policies.UserType, - Subject: session.DomainUserID, - Object: id, - ObjectType: policies.GroupType, - }); err != nil { - return groups.Page{}, errors.Wrap(errListChildrenGroups, err) - } - - return am.svc.ListChildrenGroups(ctx, session, id, startLevel, endLevel, pm) -} - -func (am *authorizationMiddleware) checkSuperAdmin(ctx context.Context, session authn.Session) error { - if session.Role != authn.SuperAdminRole { - return svcerr.ErrSuperAdminAction - } - if err := am.authz.Authorize(ctx, smqauthz.PolicyReq{ - SubjectType: policies.UserType, - Subject: session.UserID, - Permission: policies.AdminPermission, - ObjectType: policies.PlatformType, - Object: policies.MagistralaObject, - }, nil); err != nil { - return err - } - return nil -} - -func (am *authorizationMiddleware) authorize(ctx context.Context, session authn.Session, entityType string, op permissions.Operation, pr smqauthz.PolicyReq) error { - pr.Domain = session.DomainID - - perm, err := am.entitiesOps.GetPermission(entityType, op) - if err != nil { - return err - } - pr.Permission = perm.String() - - var pat *smqauthz.PATReq - if session.PatID != "" { - entityID := pr.Object - opName := am.entitiesOps.OperationName(entityType, op) - if op == dOperations.OpListDomainGroups || op == dOperations.OpCreateDomainGroups { - entityID = auth.AnyIDs - } - pat = &smqauthz.PATReq{ - UserID: session.UserID, - PatID: session.PatID, - EntityID: entityID, - EntityType: auth.GroupsType.String(), - Operation: opName, - Domain: session.DomainID, - } - } - - if err := am.authz.Authorize(ctx, pr, pat); err != nil { - return err - } - return nil -} diff --git a/groups/middleware/callout.go b/groups/middleware/callout.go deleted file mode 100644 index 32b0c8273..000000000 --- a/groups/middleware/callout.go +++ /dev/null @@ -1,285 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package middleware - -import ( - "context" - "time" - - dOperations "github.com/absmach/magistrala/domains/operations" - "github.com/absmach/magistrala/groups" - "github.com/absmach/magistrala/groups/operations" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/callout" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/pkg/permissions" - "github.com/absmach/magistrala/pkg/policies" - "github.com/absmach/magistrala/pkg/roles" - rolemgr "github.com/absmach/magistrala/pkg/roles/rolemanager/middleware" -) - -var _ groups.Service = (*calloutMiddleware)(nil) - -type calloutMiddleware struct { - svc groups.Service - repo groups.Repository - callout callout.Callout - entitiesOps permissions.EntitiesOperations[permissions.Operation] - rolemgr.RoleManagerCalloutMiddleware -} - -func NewCallout(svc groups.Service, repo groups.Repository, entitiesOps permissions.EntitiesOperations[permissions.Operation], roleOps permissions.Operations[permissions.RoleOperation], callout callout.Callout) (groups.Service, error) { - call, err := rolemgr.NewCallout(policies.GroupType, svc, callout, roleOps) - if err != nil { - return nil, err - } - - if err := entitiesOps.Validate(); err != nil { - return nil, err - } - - return &calloutMiddleware{ - svc: svc, - repo: repo, - callout: callout, - entitiesOps: entitiesOps, - RoleManagerCalloutMiddleware: call, - }, nil -} - -func (cm *calloutMiddleware) CreateGroup(ctx context.Context, session authn.Session, g groups.Group) (groups.Group, []roles.RoleProvision, error) { - params := map[string]any{ - "entities": []groups.Group{g}, - "count": 1, - } - - if err := cm.callOut(ctx, session, policies.DomainType, dOperations.OpCreateDomainGroups, params); err != nil { - return groups.Group{}, nil, err - } - - return cm.svc.CreateGroup(ctx, session, g) -} - -func (cm *calloutMiddleware) UpdateGroup(ctx context.Context, session authn.Session, group groups.Group) (groups.Group, error) { - params := map[string]any{ - "entity_id": group.ID, - "group": group, - } - - if err := cm.callOut(ctx, session, policies.GroupType, operations.OpUpdateGroup, params); err != nil { - return groups.Group{}, err - } - - return cm.svc.UpdateGroup(ctx, session, group) -} - -func (cm *calloutMiddleware) UpdateGroupTags(ctx context.Context, session authn.Session, group groups.Group) (groups.Group, error) { - params := map[string]any{ - "entity_id": group.ID, - "tags": group.Tags, - } - - if err := cm.callOut(ctx, session, policies.GroupType, operations.OpUpdateGroupTags, params); err != nil { - return groups.Group{}, err - } - - return cm.svc.UpdateGroupTags(ctx, session, group) -} - -func (cm *calloutMiddleware) ViewGroup(ctx context.Context, session authn.Session, id string, withRoles bool) (groups.Group, error) { - params := map[string]any{ - "entity_id": id, - } - - if err := cm.callOut(ctx, session, policies.GroupType, operations.OpViewGroup, params); err != nil { - return groups.Group{}, err - } - - return cm.svc.ViewGroup(ctx, session, id, withRoles) -} - -func (cm *calloutMiddleware) ListGroups(ctx context.Context, session authn.Session, gm groups.PageMeta) (groups.Page, error) { - params := map[string]any{ - "pagemeta": gm, - } - - if err := cm.callOut(ctx, session, policies.DomainType, dOperations.OpListDomainGroups, params); err != nil { - return groups.Page{}, err - } - - return cm.svc.ListGroups(ctx, session, gm) -} - -func (cm *calloutMiddleware) ListUserGroups(ctx context.Context, session authn.Session, userID string, gm groups.PageMeta) (groups.Page, error) { - params := map[string]any{ - "user_id": userID, - "pagemeta": gm, - } - - if err := cm.callOut(ctx, session, policies.GroupType, operations.OpListUserGroups, params); err != nil { - return groups.Page{}, err - } - - return cm.svc.ListUserGroups(ctx, session, userID, gm) -} - -func (cm *calloutMiddleware) EnableGroup(ctx context.Context, session authn.Session, id string) (groups.Group, error) { - params := map[string]any{ - "entity_id": id, - } - - if err := cm.callOut(ctx, session, policies.GroupType, operations.OpEnableGroup, params); err != nil { - return groups.Group{}, err - } - - return cm.svc.EnableGroup(ctx, session, id) -} - -func (cm *calloutMiddleware) DisableGroup(ctx context.Context, session authn.Session, id string) (groups.Group, error) { - params := map[string]any{ - "entity_id": id, - } - - if err := cm.callOut(ctx, session, policies.GroupType, operations.OpDisableGroup, params); err != nil { - return groups.Group{}, err - } - - return cm.svc.DisableGroup(ctx, session, id) -} - -func (cm *calloutMiddleware) DeleteGroup(ctx context.Context, session authn.Session, id string) error { - params := map[string]any{ - "entity_id": id, - } - - if err := cm.callOut(ctx, session, policies.GroupType, operations.OpDeleteGroup, params); err != nil { - return err - } - - return cm.svc.DeleteGroup(ctx, session, id) -} - -func (cm *calloutMiddleware) RetrieveGroupHierarchy(ctx context.Context, session authn.Session, id string, hm groups.HierarchyPageMeta) (groups.HierarchyPage, error) { - params := map[string]any{ - "entity_id": id, - "hierarchy_pagemeta": hm, - } - - if err := cm.callOut(ctx, session, policies.GroupType, operations.OpRetrieveGroupHierarchy, params); err != nil { - return groups.HierarchyPage{}, err - } - - return cm.svc.RetrieveGroupHierarchy(ctx, session, id, hm) -} - -func (cm *calloutMiddleware) AddParentGroup(ctx context.Context, session authn.Session, id, parentID string) error { - params := map[string]any{ - "entity_id": id, - "parent_id": parentID, - } - - if err := cm.callOut(ctx, session, policies.GroupType, operations.OpAddParentGroup, params); err != nil { - return err - } - - return cm.svc.AddParentGroup(ctx, session, id, parentID) -} - -func (cm *calloutMiddleware) RemoveParentGroup(ctx context.Context, session authn.Session, id string) error { - group, err := cm.repo.RetrieveByID(ctx, id) - if err != nil { - return errors.Wrap(svcerr.ErrViewEntity, err) - } - - params := map[string]any{ - "entity_id": id, - "parent_id": group.Parent, - } - - if err := cm.callOut(ctx, session, policies.GroupType, operations.OpRemoveParentGroup, params); err != nil { - return err - } - - return cm.svc.RemoveParentGroup(ctx, session, id) -} - -func (cm *calloutMiddleware) AddChildrenGroups(ctx context.Context, session authn.Session, id string, childrenGroupIDs []string) error { - params := map[string]any{ - "entity_id": id, - "children_group_ids": childrenGroupIDs, - } - - if err := cm.callOut(ctx, session, policies.GroupType, operations.OpAddChildrenGroups, params); err != nil { - return err - } - - return cm.svc.AddChildrenGroups(ctx, session, id, childrenGroupIDs) -} - -func (cm *calloutMiddleware) RemoveChildrenGroups(ctx context.Context, session authn.Session, id string, childrenGroupIDs []string) error { - params := map[string]any{ - "entity_id": id, - "children_group_ids": childrenGroupIDs, - } - - if err := cm.callOut(ctx, session, policies.GroupType, operations.OpRemoveChildrenGroups, params); err != nil { - return err - } - - return cm.svc.RemoveChildrenGroups(ctx, session, id, childrenGroupIDs) -} - -func (cm *calloutMiddleware) RemoveAllChildrenGroups(ctx context.Context, session authn.Session, id string) error { - params := map[string]any{ - "entity_id": id, - } - - if err := cm.callOut(ctx, session, policies.GroupType, operations.OpRemoveAllChildrenGroups, params); err != nil { - return err - } - - return cm.svc.RemoveAllChildrenGroups(ctx, session, id) -} - -func (cm *calloutMiddleware) ListChildrenGroups(ctx context.Context, session authn.Session, id string, startLevel, endLevel int64, pm groups.PageMeta) (groups.Page, error) { - params := map[string]any{ - "entity_id": id, - "start_level": startLevel, - "end_level": endLevel, - "pagemeta": pm, - } - - if err := cm.callOut(ctx, session, policies.GroupType, operations.OpListChildrenGroups, params); err != nil { - return groups.Page{}, err - } - - return cm.svc.ListChildrenGroups(ctx, session, id, startLevel, endLevel, pm) -} - -func (cm *calloutMiddleware) callOut(ctx context.Context, session authn.Session, entityType string, op permissions.Operation, pld map[string]any) error { - var entityID string - if id, ok := pld["entity_id"].(string); ok { - entityID = id - } - - req := callout.Request{ - BaseRequest: callout.BaseRequest{ - Operation: cm.entitiesOps.OperationName(entityType, op), - EntityType: entityType, - EntityID: entityID, - CallerID: session.UserID, - CallerType: policies.UserType, - DomainID: session.DomainID, - Time: time.Now().UTC(), - }, - Payload: pld, - } - - if err := cm.callout.Callout(ctx, req); err != nil { - return err - } - - return nil -} diff --git a/groups/middleware/doc.go b/groups/middleware/doc.go deleted file mode 100644 index a4c6f9861..000000000 --- a/groups/middleware/doc.go +++ /dev/null @@ -1,9 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package middleware provides authorization, logging, metrics and tracing middleware -// for Magistrala Domains service. -// -// For more details about tracing instrumentation for Magistrala refer to the -// documentation at https://magistrala.absmach.eu/docs/. -package middleware diff --git a/groups/middleware/logging.go b/groups/middleware/logging.go deleted file mode 100644 index c9849194f..000000000 --- a/groups/middleware/logging.go +++ /dev/null @@ -1,372 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package middleware - -import ( - "context" - "log/slog" - "time" - - "github.com/absmach/magistrala/groups" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/roles" - rolemw "github.com/absmach/magistrala/pkg/roles/rolemanager/middleware" - "github.com/go-chi/chi/v5/middleware" -) - -var _ groups.Service = (*loggingMiddleware)(nil) - -type loggingMiddleware struct { - logger *slog.Logger - svc groups.Service - rolemw.RoleManagerLoggingMiddleware -} - -// NewLogging adds logging facilities to the groups service. -func NewLogging(svc groups.Service, logger *slog.Logger) groups.Service { - return &loggingMiddleware{logger, svc, rolemw.NewLogging("groups", svc, logger)} -} - -// CreateGroup logs the create_group request. It logs the group name, id and token and the time it took to complete the request. -// If the request fails, it logs the error. -func (lm *loggingMiddleware) CreateGroup(ctx context.Context, session authn.Session, group groups.Group) (g groups.Group, rps []roles.RoleProvision, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("domain_id", session.DomainID), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.Group("group", - slog.String("id", g.ID), - slog.String("name", g.Name), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Create group failed", args...) - return - } - lm.logger.Info("Create group completed successfully", args...) - }(time.Now()) - return lm.svc.CreateGroup(ctx, session, group) -} - -// UpdateGroup logs the update_group request. It logs the group name, id and the time it took to complete the request. -// If the request fails, it logs the error. -func (lm *loggingMiddleware) UpdateGroup(ctx context.Context, session authn.Session, group groups.Group) (g groups.Group, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("domain_id", session.DomainID), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.Group("group", - slog.String("id", group.ID), - slog.String("name", group.Name), - slog.Any("metadata", group.Metadata), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Update group failed", args...) - return - } - lm.logger.Info("Update group completed successfully", args...) - }(time.Now()) - return lm.svc.UpdateGroup(ctx, session, group) -} - -func (lm *loggingMiddleware) UpdateGroupTags(ctx context.Context, session authn.Session, group groups.Group) (g groups.Group, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("domain_id", session.DomainID), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.Group("group", - slog.String("id", g.ID), - slog.String("name", g.Name), - slog.Any("tags", g.Tags), - ), - } - if err != nil { - args := append(args, slog.String("error", err.Error())) - lm.logger.Warn("Update group tags failed", args...) - return - } - lm.logger.Info("Update group tags completed successfully", args...) - }(time.Now()) - return lm.svc.UpdateGroupTags(ctx, session, group) -} - -// ViewGroup logs the view_group request. It logs the group name, id and the time it took to complete the request. -// If the request fails, it logs the error. -func (lm *loggingMiddleware) ViewGroup(ctx context.Context, session authn.Session, id string, withRoles bool) (g groups.Group, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("domain_id", session.DomainID), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.Group("group", - slog.String("id", g.ID), - slog.String("name", g.Name), - slog.Bool("with_roles", withRoles), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("View group failed", args...) - return - } - lm.logger.Info("View group completed successfully", args...) - }(time.Now()) - return lm.svc.ViewGroup(ctx, session, id, withRoles) -} - -// ListGroups logs the list_groups request. It logs the page metadata and the time it took to complete the request. -// If the request fails, it logs the error. -func (lm *loggingMiddleware) ListGroups(ctx context.Context, session authn.Session, pm groups.PageMeta) (cg groups.Page, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("domain_id", session.DomainID), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.Group("page", - slog.Uint64("limit", pm.Limit), - slog.Uint64("offset", pm.Offset), - slog.Uint64("total", cg.Total), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("List groups failed", args...) - return - } - lm.logger.Info("List groups completed successfully", args...) - }(time.Now()) - return lm.svc.ListGroups(ctx, session, pm) -} - -func (lm *loggingMiddleware) ListUserGroups(ctx context.Context, session authn.Session, userID string, pm groups.PageMeta) (cg groups.Page, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.String("user_id", userID), - slog.String("domain_id", session.DomainID), - slog.Group("page", - slog.Uint64("limit", pm.Limit), - slog.Uint64("offset", pm.Offset), - slog.Uint64("total", cg.Total), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("List user groups failed", args...) - return - } - lm.logger.Info("List user groups completed successfully", args...) - }(time.Now()) - return lm.svc.ListUserGroups(ctx, session, userID, pm) -} - -// EnableGroup logs the enable_group request. It logs the group name, id and the time it took to complete the request. -// If the request fails, it logs the error. -func (lm *loggingMiddleware) EnableGroup(ctx context.Context, session authn.Session, id string) (g groups.Group, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("domain_id", session.DomainID), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.Group("group", - slog.String("id", id), - slog.String("name", g.Name), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Enable group failed", args...) - return - } - lm.logger.Info("Enable group completed successfully", args...) - }(time.Now()) - return lm.svc.EnableGroup(ctx, session, id) -} - -// DisableGroup logs the disable_group request. It logs the group id and the time it took to complete the request. -// If the request fails, it logs the error. -func (lm *loggingMiddleware) DisableGroup(ctx context.Context, session authn.Session, id string) (g groups.Group, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("domain_id", session.DomainID), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.Group("group", - slog.String("id", id), - slog.String("name", g.Name), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Disable group failed", args...) - return - } - lm.logger.Info("Disable group completed successfully", args...) - }(time.Now()) - return lm.svc.DisableGroup(ctx, session, id) -} - -func (lm *loggingMiddleware) DeleteGroup(ctx context.Context, session authn.Session, id string) (err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("domain_id", session.DomainID), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.String("group_id", id), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Delete group failed", args...) - return - } - lm.logger.Info("Delete group completed successfully", args...) - }(time.Now()) - return lm.svc.DeleteGroup(ctx, session, id) -} - -func (lm *loggingMiddleware) RetrieveGroupHierarchy(ctx context.Context, session authn.Session, id string, hm groups.HierarchyPageMeta) (gp groups.HierarchyPage, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.String("group_id", id), - slog.String("domain_id", session.DomainID), - slog.Group("page", - slog.Uint64("level", hm.Level), - slog.Int64("direction", hm.Direction), - slog.Bool("tree", hm.Tree), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Retrieve group hierarchy failed", args...) - return - } - lm.logger.Info("Retrieve group hierarchy completed successfully", args...) - }(time.Now()) - return lm.svc.RetrieveGroupHierarchy(ctx, session, id, hm) -} - -func (lm *loggingMiddleware) AddParentGroup(ctx context.Context, session authn.Session, id, parentID string) (err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("domain_id", session.DomainID), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.String("group_id", id), - slog.String("parent_group_id", parentID), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Add parent group failed", args...) - return - } - lm.logger.Info("Add parent group completed successfully", args...) - }(time.Now()) - return lm.svc.AddParentGroup(ctx, session, id, parentID) -} - -func (lm *loggingMiddleware) RemoveParentGroup(ctx context.Context, session authn.Session, id string) (err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("domain_id", session.DomainID), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.String("group_id", id), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Remove parent group failed", args...) - return - } - lm.logger.Info("Remove parent group completed successfully", args...) - }(time.Now()) - return lm.svc.RemoveParentGroup(ctx, session, id) -} - -func (lm *loggingMiddleware) AddChildrenGroups(ctx context.Context, session authn.Session, id string, childrenGroupIDs []string) (err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("domain_id", session.DomainID), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.String("group_id", id), - slog.Any("children_group_ids", childrenGroupIDs), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Add children groups failed", args...) - return - } - lm.logger.Info("Add parent group completed successfully", args...) - }(time.Now()) - return lm.svc.AddChildrenGroups(ctx, session, id, childrenGroupIDs) -} - -func (lm *loggingMiddleware) RemoveChildrenGroups(ctx context.Context, session authn.Session, id string, childrenGroupIDs []string) (err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("domain_id", session.DomainID), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.String("group_id", id), - slog.Any("children_group_ids", childrenGroupIDs), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Remove children groups failed", args...) - return - } - lm.logger.Info("Remove parent group completed successfully", args...) - }(time.Now()) - return lm.svc.RemoveChildrenGroups(ctx, session, id, childrenGroupIDs) -} - -func (lm *loggingMiddleware) RemoveAllChildrenGroups(ctx context.Context, session authn.Session, id string) (err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("domain_id", session.DomainID), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.String("group_id", id), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Remove all children groups failed", args...) - return - } - lm.logger.Info("Remove all parent group completed successfully", args...) - }(time.Now()) - return lm.svc.RemoveAllChildrenGroups(ctx, session, id) -} - -func (lm *loggingMiddleware) ListChildrenGroups(ctx context.Context, session authn.Session, id string, startLevel, endLevel int64, pm groups.PageMeta) (gp groups.Page, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("domain_id", session.DomainID), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.String("group_id", id), - slog.Group("page", - slog.Uint64("limit", pm.Limit), - slog.Uint64("offset", pm.Offset), - slog.Uint64("total", gp.Total), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("List children groups failed", args...) - return - } - lm.logger.Info("List children groups completed successfully", args...) - }(time.Now()) - return lm.svc.ListChildrenGroups(ctx, session, id, startLevel, endLevel, pm) -} diff --git a/groups/middleware/metrics.go b/groups/middleware/metrics.go deleted file mode 100644 index 0601aa8d8..000000000 --- a/groups/middleware/metrics.go +++ /dev/null @@ -1,170 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package middleware - -import ( - "context" - "time" - - "github.com/absmach/magistrala/groups" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/roles" - rolemw "github.com/absmach/magistrala/pkg/roles/rolemanager/middleware" - "github.com/go-kit/kit/metrics" -) - -var _ groups.Service = (*metricsMiddleware)(nil) - -type metricsMiddleware struct { - counter metrics.Counter - latency metrics.Histogram - svc groups.Service - rolemw.RoleManagerMetricsMiddleware -} - -// NewMetrics instruments policies service by tracking request count and latency. -func NewMetrics(svc groups.Service, counter metrics.Counter, latency metrics.Histogram) groups.Service { - rmm := rolemw.NewMetrics("group", svc, counter, latency) - return &metricsMiddleware{ - counter: counter, - latency: latency, - svc: svc, - RoleManagerMetricsMiddleware: rmm, - } -} - -// CreateGroup instruments CreateGroup method with metrics. -func (ms *metricsMiddleware) CreateGroup(ctx context.Context, session authn.Session, g groups.Group) (groups.Group, []roles.RoleProvision, error) { - defer func(begin time.Time) { - ms.counter.With("method", "create_group").Add(1) - ms.latency.With("method", "create_group").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.CreateGroup(ctx, session, g) -} - -// UpdateGroup instruments UpdateGroup method with metrics. -func (ms *metricsMiddleware) UpdateGroup(ctx context.Context, session authn.Session, group groups.Group) (rGroup groups.Group, err error) { - defer func(begin time.Time) { - ms.counter.With("method", "update_group").Add(1) - ms.latency.With("method", "update_group").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.UpdateGroup(ctx, session, group) -} - -// UpdateGroupTags instruments UpdateGroupTags method with metrics. -func (ms *metricsMiddleware) UpdateGroupTags(ctx context.Context, session authn.Session, group groups.Group) (groups.Group, error) { - defer func(begin time.Time) { - ms.counter.With("method", "update_group_tags").Add(1) - ms.latency.With("method", "update_group_tags").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.UpdateGroupTags(ctx, session, group) -} - -// ViewGroup instruments ViewGroup method with metrics. -func (ms *metricsMiddleware) ViewGroup(ctx context.Context, session authn.Session, id string, withRoles bool) (g groups.Group, err error) { - defer func(begin time.Time) { - ms.counter.With("method", "view_group").Add(1) - ms.latency.With("method", "view_group").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.ViewGroup(ctx, session, id, withRoles) -} - -// ListGroups instruments ListGroups method with metrics. -func (ms *metricsMiddleware) ListGroups(ctx context.Context, session authn.Session, pm groups.PageMeta) (cg groups.Page, err error) { - defer func(begin time.Time) { - ms.counter.With("method", "list_groups").Add(1) - ms.latency.With("method", "list_groups").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.ListGroups(ctx, session, pm) -} - -func (ms *metricsMiddleware) ListUserGroups(ctx context.Context, session authn.Session, userID string, pm groups.PageMeta) (cg groups.Page, err error) { - defer func(begin time.Time) { - ms.counter.With("method", "list_user_groups").Add(1) - ms.latency.With("method", "list_user_groups").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.ListUserGroups(ctx, session, userID, pm) -} - -// EnableGroup instruments EnableGroup method with metrics. -func (ms *metricsMiddleware) EnableGroup(ctx context.Context, session authn.Session, id string) (g groups.Group, err error) { - defer func(begin time.Time) { - ms.counter.With("method", "enable_group").Add(1) - ms.latency.With("method", "enable_group").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.EnableGroup(ctx, session, id) -} - -// DisableGroup instruments DisableGroup method with metrics. -func (ms *metricsMiddleware) DisableGroup(ctx context.Context, session authn.Session, id string) (g groups.Group, err error) { - defer func(begin time.Time) { - ms.counter.With("method", "disable_group").Add(1) - ms.latency.With("method", "disable_group").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.DisableGroup(ctx, session, id) -} - -func (ms *metricsMiddleware) DeleteGroup(ctx context.Context, session authn.Session, id string) (err error) { - defer func(begin time.Time) { - ms.counter.With("method", "delete_group").Add(1) - ms.latency.With("method", "delete_group").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.DeleteGroup(ctx, session, id) -} - -func (ms *metricsMiddleware) RetrieveGroupHierarchy(ctx context.Context, session authn.Session, id string, hm groups.HierarchyPageMeta) (groups.HierarchyPage, error) { - defer func(begin time.Time) { - ms.counter.With("method", "list_parent_groups").Add(1) - ms.latency.With("method", "list_parent_groups").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.RetrieveGroupHierarchy(ctx, session, id, hm) -} - -func (ms *metricsMiddleware) AddParentGroup(ctx context.Context, session authn.Session, id, parentID string) error { - defer func(begin time.Time) { - ms.counter.With("method", "add_parent_group").Add(1) - ms.latency.With("method", "add_parent_group").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.AddParentGroup(ctx, session, id, parentID) -} - -func (ms *metricsMiddleware) RemoveParentGroup(ctx context.Context, session authn.Session, id string) error { - defer func(begin time.Time) { - ms.counter.With("method", "remove_parent_group").Add(1) - ms.latency.With("method", "remove_parent_group").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.RemoveParentGroup(ctx, session, id) -} - -func (ms *metricsMiddleware) AddChildrenGroups(ctx context.Context, session authn.Session, id string, childrenGroupIDs []string) error { - defer func(begin time.Time) { - ms.counter.With("method", "add_children_groups").Add(1) - ms.latency.With("method", "add_children_groups").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.AddChildrenGroups(ctx, session, id, childrenGroupIDs) -} - -func (ms *metricsMiddleware) RemoveChildrenGroups(ctx context.Context, session authn.Session, id string, childrenGroupIDs []string) error { - defer func(begin time.Time) { - ms.counter.With("method", "remove_children_groups").Add(1) - ms.latency.With("method", "remove_children_groups").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.RemoveChildrenGroups(ctx, session, id, childrenGroupIDs) -} - -func (ms *metricsMiddleware) RemoveAllChildrenGroups(ctx context.Context, session authn.Session, id string) error { - defer func(begin time.Time) { - ms.counter.With("method", "remove_all_children_groups").Add(1) - ms.latency.With("method", "remove_all_children_groups").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.RemoveAllChildrenGroups(ctx, session, id) -} - -func (ms *metricsMiddleware) ListChildrenGroups(ctx context.Context, session authn.Session, id string, startLevel, endLevel int64, pm groups.PageMeta) (groups.Page, error) { - defer func(begin time.Time) { - ms.counter.With("method", "list_children_groups").Add(1) - ms.latency.With("method", "list_children_groups").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.ListChildrenGroups(ctx, session, id, startLevel, endLevel, pm) -} diff --git a/groups/middleware/tracing.go b/groups/middleware/tracing.go deleted file mode 100644 index 4a972303e..000000000 --- a/groups/middleware/tracing.go +++ /dev/null @@ -1,200 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package middleware - -import ( - "context" - "fmt" - - "github.com/absmach/magistrala/groups" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/roles" - rolemw "github.com/absmach/magistrala/pkg/roles/rolemanager/middleware" - "github.com/absmach/magistrala/pkg/tracing" - "go.opentelemetry.io/otel/attribute" - "go.opentelemetry.io/otel/trace" -) - -var _ groups.Service = (*tracingMiddleware)(nil) - -type tracingMiddleware struct { - tracer trace.Tracer - svc groups.Service - rolemw.RoleManagerTracing -} - -// NewTracing returns a new groups service with tracing capabilities. -func NewTracing(svc groups.Service, tracer trace.Tracer) groups.Service { - return &tracingMiddleware{tracer, svc, rolemw.NewTracing("group", svc, tracer)} -} - -// CreateGroup traces the "CreateGroup" operation of the wrapped groups.Service. -func (tm *tracingMiddleware) CreateGroup(ctx context.Context, session authn.Session, g groups.Group) (groups.Group, []roles.RoleProvision, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_create_group") - defer span.End() - - return tm.svc.CreateGroup(ctx, session, g) -} - -// ViewGroup traces the "ViewGroup" operation of the wrapped groups.Service. -func (tm *tracingMiddleware) ViewGroup(ctx context.Context, session authn.Session, id string, withRoles bool) (groups.Group, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_view_group", trace.WithAttributes(attribute.String("id", id), attribute.Bool("with_roles", withRoles))) - defer span.End() - - return tm.svc.ViewGroup(ctx, session, id, withRoles) -} - -// ListGroups traces the "ListGroups" operation of the wrapped groups.Service. -func (tm *tracingMiddleware) ListGroups(ctx context.Context, session authn.Session, pm groups.PageMeta) (groups.Page, error) { - attr := []attribute.KeyValue{ - attribute.String("name", pm.Name), - attribute.StringSlice("tags", pm.Tags.Elements), - attribute.String("status", pm.Status.String()), - attribute.Int64("offset", int64(pm.Offset)), - attribute.Int64("limit", int64(pm.Limit)), - } - for k, v := range pm.Metadata { - attr = append(attr, attribute.String(k, fmt.Sprintf("%v", v))) - } - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_list_groups", trace.WithAttributes(attr...)) - defer span.End() - - return tm.svc.ListGroups(ctx, session, pm) -} - -func (tm *tracingMiddleware) ListUserGroups(ctx context.Context, session authn.Session, userID string, pm groups.PageMeta) (groups.Page, error) { - attr := []attribute.KeyValue{ - attribute.String("user_id", userID), - attribute.String("name", pm.Name), - attribute.StringSlice("tag", pm.Tags.Elements), - attribute.String("status", pm.Status.String()), - attribute.Int64("offset", int64(pm.Offset)), - attribute.Int64("limit", int64(pm.Limit)), - } - for k, v := range pm.Metadata { - attr = append(attr, attribute.String(k, fmt.Sprintf("%v", v))) - } - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_list_user_groups", trace.WithAttributes(attr...)) - defer span.End() - - return tm.svc.ListUserGroups(ctx, session, userID, pm) -} - -// UpdateGroup traces the "UpdateGroup" operation of the wrapped groups.Service. -func (tm *tracingMiddleware) UpdateGroup(ctx context.Context, session authn.Session, g groups.Group) (groups.Group, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_update_group") - defer span.End() - - return tm.svc.UpdateGroup(ctx, session, g) -} - -// UpdateGroupTags traces the "UpdateGroupTags" operation of the wrapped groups.Service. -func (tm *tracingMiddleware) UpdateGroupTags(ctx context.Context, session authn.Session, group groups.Group) (groups.Group, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_update_group_tags", trace.WithAttributes( - attribute.String("id", group.ID), - attribute.StringSlice("tags", group.Tags), - )) - defer span.End() - - return tm.svc.UpdateGroupTags(ctx, session, group) -} - -// EnableGroup traces the "EnableGroup" operation of the wrapped groups.Service. -func (tm *tracingMiddleware) EnableGroup(ctx context.Context, session authn.Session, id string) (groups.Group, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_enable_group", trace.WithAttributes(attribute.String("id", id))) - defer span.End() - - return tm.svc.EnableGroup(ctx, session, id) -} - -// DisableGroup traces the "DisableGroup" operation of the wrapped groups.Service. -func (tm *tracingMiddleware) DisableGroup(ctx context.Context, session authn.Session, id string) (groups.Group, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_disable_group", trace.WithAttributes(attribute.String("id", id))) - defer span.End() - - return tm.svc.DisableGroup(ctx, session, id) -} - -func (tm *tracingMiddleware) RetrieveGroupHierarchy(ctx context.Context, session authn.Session, id string, hm groups.HierarchyPageMeta) (groups.HierarchyPage, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_list_group_hierarchy", - trace.WithAttributes( - attribute.String("id", id), - attribute.Int64("level", int64(hm.Level)), - attribute.Int64("direction", hm.Direction), - attribute.Bool("tree", hm.Tree), - )) - defer span.End() - - return tm.svc.RetrieveGroupHierarchy(ctx, session, id, hm) -} - -func (tm *tracingMiddleware) AddParentGroup(ctx context.Context, session authn.Session, id, parentID string) error { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_add_parent_group", - trace.WithAttributes( - attribute.String("id", id), - attribute.String("parent_id", parentID), - )) - defer span.End() - return tm.svc.AddParentGroup(ctx, session, id, parentID) -} - -func (tm *tracingMiddleware) RemoveParentGroup(ctx context.Context, session authn.Session, id string) error { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_remove_parent_group", trace.WithAttributes(attribute.String("id", id))) - defer span.End() - return tm.svc.RemoveParentGroup(ctx, session, id) -} - -func (tm *tracingMiddleware) AddChildrenGroups(ctx context.Context, session authn.Session, id string, childrenGroupIDs []string) error { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_add_children_groups", - trace.WithAttributes( - attribute.String("id", id), - attribute.StringSlice("children_group_ids", childrenGroupIDs), - )) - - defer span.End() - return tm.svc.AddChildrenGroups(ctx, session, id, childrenGroupIDs) -} - -func (tm *tracingMiddleware) RemoveChildrenGroups(ctx context.Context, session authn.Session, id string, childrenGroupIDs []string) error { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_remove_children_groups", - trace.WithAttributes( - attribute.String("id", id), - attribute.StringSlice("children_group_ids", childrenGroupIDs), - )) - defer span.End() - return tm.svc.RemoveChildrenGroups(ctx, session, id, childrenGroupIDs) -} - -func (tm *tracingMiddleware) RemoveAllChildrenGroups(ctx context.Context, session authn.Session, id string) error { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_remove_all_children_groups", trace.WithAttributes(attribute.String("id", id))) - defer span.End() - return tm.svc.RemoveAllChildrenGroups(ctx, session, id) -} - -func (tm *tracingMiddleware) ListChildrenGroups(ctx context.Context, session authn.Session, id string, startLevel, endLevel int64, pm groups.PageMeta) (groups.Page, error) { - attr := []attribute.KeyValue{ - attribute.String("id", id), - attribute.String("name", pm.Name), - attribute.StringSlice("tags", pm.Tags.Elements), - attribute.String("status", pm.Status.String()), - attribute.Int64("start_level", startLevel), - attribute.Int64("end_level", endLevel), - attribute.Int64("offset", int64(pm.Offset)), - attribute.Int64("limit", int64(pm.Limit)), - } - for k, v := range pm.Metadata { - attr = append(attr, attribute.String(k, fmt.Sprintf("%v", v))) - } - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_list_children_groups", trace.WithAttributes(attr...)) - defer span.End() - return tm.svc.ListChildrenGroups(ctx, session, id, startLevel, endLevel, pm) -} - -// DeleteGroup traces the "DeleteGroup" operation of the wrapped groups.Service. -func (tm *tracingMiddleware) DeleteGroup(ctx context.Context, session authn.Session, id string) error { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_delete_group", trace.WithAttributes(attribute.String("id", id))) - defer span.End() - - return tm.svc.DeleteGroup(ctx, session, id) -} diff --git a/groups/mocks/doc.go b/groups/mocks/doc.go deleted file mode 100644 index 16ed198af..000000000 --- a/groups/mocks/doc.go +++ /dev/null @@ -1,5 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package mocks contains mocks for testing purposes. -package mocks diff --git a/groups/mocks/groups_client.go b/groups/mocks/groups_client.go deleted file mode 100644 index dce62077f..000000000 --- a/groups/mocks/groups_client.go +++ /dev/null @@ -1,127 +0,0 @@ -// Copyright (c) Abstract Machines - -// SPDX-License-Identifier: Apache-2.0 - -// Code generated by mockery; DO NOT EDIT. -// github.com/vektra/mockery -// template: testify - -package mocks - -import ( - "context" - - "github.com/absmach/magistrala/api/grpc/common/v1" - mock "github.com/stretchr/testify/mock" - "google.golang.org/grpc" -) - -// NewGroupsServiceClient creates a new instance of GroupsServiceClient. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. -// The first argument is typically a *testing.T value. -func NewGroupsServiceClient(t interface { - mock.TestingT - Cleanup(func()) -}) *GroupsServiceClient { - mock := &GroupsServiceClient{} - mock.Mock.Test(t) - - t.Cleanup(func() { mock.AssertExpectations(t) }) - - return mock -} - -// GroupsServiceClient is an autogenerated mock type for the GroupsServiceClient type -type GroupsServiceClient struct { - mock.Mock -} - -type GroupsServiceClient_Expecter struct { - mock *mock.Mock -} - -func (_m *GroupsServiceClient) EXPECT() *GroupsServiceClient_Expecter { - return &GroupsServiceClient_Expecter{mock: &_m.Mock} -} - -// RetrieveEntity provides a mock function for the type GroupsServiceClient -func (_mock *GroupsServiceClient) RetrieveEntity(ctx context.Context, in *v1.RetrieveEntityReq, opts ...grpc.CallOption) (*v1.RetrieveEntityRes, error) { - var tmpRet mock.Arguments - if len(opts) > 0 { - tmpRet = _mock.Called(ctx, in, opts) - } else { - tmpRet = _mock.Called(ctx, in) - } - ret := tmpRet - - if len(ret) == 0 { - panic("no return value specified for RetrieveEntity") - } - - var r0 *v1.RetrieveEntityRes - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, *v1.RetrieveEntityReq, ...grpc.CallOption) (*v1.RetrieveEntityRes, error)); ok { - return returnFunc(ctx, in, opts...) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, *v1.RetrieveEntityReq, ...grpc.CallOption) *v1.RetrieveEntityRes); ok { - r0 = returnFunc(ctx, in, opts...) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).(*v1.RetrieveEntityRes) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, *v1.RetrieveEntityReq, ...grpc.CallOption) error); ok { - r1 = returnFunc(ctx, in, opts...) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// GroupsServiceClient_RetrieveEntity_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveEntity' -type GroupsServiceClient_RetrieveEntity_Call struct { - *mock.Call -} - -// RetrieveEntity is a helper method to define mock.On call -// - ctx context.Context -// - in *v1.RetrieveEntityReq -// - opts ...grpc.CallOption -func (_e *GroupsServiceClient_Expecter) RetrieveEntity(ctx interface{}, in interface{}, opts ...interface{}) *GroupsServiceClient_RetrieveEntity_Call { - return &GroupsServiceClient_RetrieveEntity_Call{Call: _e.mock.On("RetrieveEntity", - append([]interface{}{ctx, in}, opts...)...)} -} - -func (_c *GroupsServiceClient_RetrieveEntity_Call) Run(run func(ctx context.Context, in *v1.RetrieveEntityReq, opts ...grpc.CallOption)) *GroupsServiceClient_RetrieveEntity_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 *v1.RetrieveEntityReq - if args[1] != nil { - arg1 = args[1].(*v1.RetrieveEntityReq) - } - var arg2 []grpc.CallOption - var variadicArgs []grpc.CallOption - if len(args) > 2 { - variadicArgs = args[2].([]grpc.CallOption) - } - arg2 = variadicArgs - run( - arg0, - arg1, - arg2..., - ) - }) - return _c -} - -func (_c *GroupsServiceClient_RetrieveEntity_Call) Return(retrieveEntityRes *v1.RetrieveEntityRes, err error) *GroupsServiceClient_RetrieveEntity_Call { - _c.Call.Return(retrieveEntityRes, err) - return _c -} - -func (_c *GroupsServiceClient_RetrieveEntity_Call) RunAndReturn(run func(ctx context.Context, in *v1.RetrieveEntityReq, opts ...grpc.CallOption) (*v1.RetrieveEntityRes, error)) *GroupsServiceClient_RetrieveEntity_Call { - _c.Call.Return(run) - return _c -} diff --git a/groups/mocks/repository.go b/groups/mocks/repository.go deleted file mode 100644 index 37c2e2426..000000000 --- a/groups/mocks/repository.go +++ /dev/null @@ -1,2624 +0,0 @@ -// Copyright (c) Abstract Machines - -// SPDX-License-Identifier: Apache-2.0 - -// Code generated by mockery; DO NOT EDIT. -// github.com/vektra/mockery -// template: testify - -package mocks - -import ( - "context" - - "github.com/absmach/magistrala/groups" - "github.com/absmach/magistrala/pkg/roles" - mock "github.com/stretchr/testify/mock" -) - -// NewRepository creates a new instance of Repository. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. -// The first argument is typically a *testing.T value. -func NewRepository(t interface { - mock.TestingT - Cleanup(func()) -}) *Repository { - mock := &Repository{} - mock.Mock.Test(t) - - t.Cleanup(func() { mock.AssertExpectations(t) }) - - return mock -} - -// Repository is an autogenerated mock type for the Repository type -type Repository struct { - mock.Mock -} - -type Repository_Expecter struct { - mock *mock.Mock -} - -func (_m *Repository) EXPECT() *Repository_Expecter { - return &Repository_Expecter{mock: &_m.Mock} -} - -// AddRoles provides a mock function for the type Repository -func (_mock *Repository) AddRoles(ctx context.Context, rps []roles.RoleProvision) ([]roles.RoleProvision, error) { - ret := _mock.Called(ctx, rps) - - if len(ret) == 0 { - panic("no return value specified for AddRoles") - } - - var r0 []roles.RoleProvision - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, []roles.RoleProvision) ([]roles.RoleProvision, error)); ok { - return returnFunc(ctx, rps) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, []roles.RoleProvision) []roles.RoleProvision); ok { - r0 = returnFunc(ctx, rps) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]roles.RoleProvision) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, []roles.RoleProvision) error); ok { - r1 = returnFunc(ctx, rps) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_AddRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'AddRoles' -type Repository_AddRoles_Call struct { - *mock.Call -} - -// AddRoles is a helper method to define mock.On call -// - ctx context.Context -// - rps []roles.RoleProvision -func (_e *Repository_Expecter) AddRoles(ctx interface{}, rps interface{}) *Repository_AddRoles_Call { - return &Repository_AddRoles_Call{Call: _e.mock.On("AddRoles", ctx, rps)} -} - -func (_c *Repository_AddRoles_Call) Run(run func(ctx context.Context, rps []roles.RoleProvision)) *Repository_AddRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 []roles.RoleProvision - if args[1] != nil { - arg1 = args[1].([]roles.RoleProvision) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_AddRoles_Call) Return(roleProvisions []roles.RoleProvision, err error) *Repository_AddRoles_Call { - _c.Call.Return(roleProvisions, err) - return _c -} - -func (_c *Repository_AddRoles_Call) RunAndReturn(run func(ctx context.Context, rps []roles.RoleProvision) ([]roles.RoleProvision, error)) *Repository_AddRoles_Call { - _c.Call.Return(run) - return _c -} - -// AssignParentGroup provides a mock function for the type Repository -func (_mock *Repository) AssignParentGroup(ctx context.Context, parentGroupID string, groupIDs ...string) error { - var tmpRet mock.Arguments - if len(groupIDs) > 0 { - tmpRet = _mock.Called(ctx, parentGroupID, groupIDs) - } else { - tmpRet = _mock.Called(ctx, parentGroupID) - } - ret := tmpRet - - if len(ret) == 0 { - panic("no return value specified for AssignParentGroup") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, ...string) error); ok { - r0 = returnFunc(ctx, parentGroupID, groupIDs...) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_AssignParentGroup_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'AssignParentGroup' -type Repository_AssignParentGroup_Call struct { - *mock.Call -} - -// AssignParentGroup is a helper method to define mock.On call -// - ctx context.Context -// - parentGroupID string -// - groupIDs ...string -func (_e *Repository_Expecter) AssignParentGroup(ctx interface{}, parentGroupID interface{}, groupIDs ...interface{}) *Repository_AssignParentGroup_Call { - return &Repository_AssignParentGroup_Call{Call: _e.mock.On("AssignParentGroup", - append([]interface{}{ctx, parentGroupID}, groupIDs...)...)} -} - -func (_c *Repository_AssignParentGroup_Call) Run(run func(ctx context.Context, parentGroupID string, groupIDs ...string)) *Repository_AssignParentGroup_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 []string - var variadicArgs []string - if len(args) > 2 { - variadicArgs = args[2].([]string) - } - arg2 = variadicArgs - run( - arg0, - arg1, - arg2..., - ) - }) - return _c -} - -func (_c *Repository_AssignParentGroup_Call) Return(err error) *Repository_AssignParentGroup_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_AssignParentGroup_Call) RunAndReturn(run func(ctx context.Context, parentGroupID string, groupIDs ...string) error) *Repository_AssignParentGroup_Call { - _c.Call.Return(run) - return _c -} - -// ChangeStatus provides a mock function for the type Repository -func (_mock *Repository) ChangeStatus(ctx context.Context, group groups.Group) (groups.Group, error) { - ret := _mock.Called(ctx, group) - - if len(ret) == 0 { - panic("no return value specified for ChangeStatus") - } - - var r0 groups.Group - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, groups.Group) (groups.Group, error)); ok { - return returnFunc(ctx, group) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, groups.Group) groups.Group); ok { - r0 = returnFunc(ctx, group) - } else { - r0 = ret.Get(0).(groups.Group) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, groups.Group) error); ok { - r1 = returnFunc(ctx, group) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_ChangeStatus_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ChangeStatus' -type Repository_ChangeStatus_Call struct { - *mock.Call -} - -// ChangeStatus is a helper method to define mock.On call -// - ctx context.Context -// - group groups.Group -func (_e *Repository_Expecter) ChangeStatus(ctx interface{}, group interface{}) *Repository_ChangeStatus_Call { - return &Repository_ChangeStatus_Call{Call: _e.mock.On("ChangeStatus", ctx, group)} -} - -func (_c *Repository_ChangeStatus_Call) Run(run func(ctx context.Context, group groups.Group)) *Repository_ChangeStatus_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 groups.Group - if args[1] != nil { - arg1 = args[1].(groups.Group) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_ChangeStatus_Call) Return(group1 groups.Group, err error) *Repository_ChangeStatus_Call { - _c.Call.Return(group1, err) - return _c -} - -func (_c *Repository_ChangeStatus_Call) RunAndReturn(run func(ctx context.Context, group groups.Group) (groups.Group, error)) *Repository_ChangeStatus_Call { - _c.Call.Return(run) - return _c -} - -// Delete provides a mock function for the type Repository -func (_mock *Repository) Delete(ctx context.Context, groupID string) error { - ret := _mock.Called(ctx, groupID) - - if len(ret) == 0 { - panic("no return value specified for Delete") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) error); ok { - r0 = returnFunc(ctx, groupID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_Delete_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Delete' -type Repository_Delete_Call struct { - *mock.Call -} - -// Delete is a helper method to define mock.On call -// - ctx context.Context -// - groupID string -func (_e *Repository_Expecter) Delete(ctx interface{}, groupID interface{}) *Repository_Delete_Call { - return &Repository_Delete_Call{Call: _e.mock.On("Delete", ctx, groupID)} -} - -func (_c *Repository_Delete_Call) Run(run func(ctx context.Context, groupID string)) *Repository_Delete_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_Delete_Call) Return(err error) *Repository_Delete_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_Delete_Call) RunAndReturn(run func(ctx context.Context, groupID string) error) *Repository_Delete_Call { - _c.Call.Return(run) - return _c -} - -// ListEntityMembers provides a mock function for the type Repository -func (_mock *Repository) ListEntityMembers(ctx context.Context, entityID string, pageQuery roles.MembersRolePageQuery) (roles.MembersRolePage, error) { - ret := _mock.Called(ctx, entityID, pageQuery) - - if len(ret) == 0 { - panic("no return value specified for ListEntityMembers") - } - - var r0 roles.MembersRolePage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, roles.MembersRolePageQuery) (roles.MembersRolePage, error)); ok { - return returnFunc(ctx, entityID, pageQuery) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, roles.MembersRolePageQuery) roles.MembersRolePage); ok { - r0 = returnFunc(ctx, entityID, pageQuery) - } else { - r0 = ret.Get(0).(roles.MembersRolePage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, roles.MembersRolePageQuery) error); ok { - r1 = returnFunc(ctx, entityID, pageQuery) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_ListEntityMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ListEntityMembers' -type Repository_ListEntityMembers_Call struct { - *mock.Call -} - -// ListEntityMembers is a helper method to define mock.On call -// - ctx context.Context -// - entityID string -// - pageQuery roles.MembersRolePageQuery -func (_e *Repository_Expecter) ListEntityMembers(ctx interface{}, entityID interface{}, pageQuery interface{}) *Repository_ListEntityMembers_Call { - return &Repository_ListEntityMembers_Call{Call: _e.mock.On("ListEntityMembers", ctx, entityID, pageQuery)} -} - -func (_c *Repository_ListEntityMembers_Call) Run(run func(ctx context.Context, entityID string, pageQuery roles.MembersRolePageQuery)) *Repository_ListEntityMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 roles.MembersRolePageQuery - if args[2] != nil { - arg2 = args[2].(roles.MembersRolePageQuery) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_ListEntityMembers_Call) Return(membersRolePage roles.MembersRolePage, err error) *Repository_ListEntityMembers_Call { - _c.Call.Return(membersRolePage, err) - return _c -} - -func (_c *Repository_ListEntityMembers_Call) RunAndReturn(run func(ctx context.Context, entityID string, pageQuery roles.MembersRolePageQuery) (roles.MembersRolePage, error)) *Repository_ListEntityMembers_Call { - _c.Call.Return(run) - return _c -} - -// RemoveEntityMembers provides a mock function for the type Repository -func (_mock *Repository) RemoveEntityMembers(ctx context.Context, entityID string, members []string) error { - ret := _mock.Called(ctx, entityID, members) - - if len(ret) == 0 { - panic("no return value specified for RemoveEntityMembers") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, []string) error); ok { - r0 = returnFunc(ctx, entityID, members) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RemoveEntityMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveEntityMembers' -type Repository_RemoveEntityMembers_Call struct { - *mock.Call -} - -// RemoveEntityMembers is a helper method to define mock.On call -// - ctx context.Context -// - entityID string -// - members []string -func (_e *Repository_Expecter) RemoveEntityMembers(ctx interface{}, entityID interface{}, members interface{}) *Repository_RemoveEntityMembers_Call { - return &Repository_RemoveEntityMembers_Call{Call: _e.mock.On("RemoveEntityMembers", ctx, entityID, members)} -} - -func (_c *Repository_RemoveEntityMembers_Call) Run(run func(ctx context.Context, entityID string, members []string)) *Repository_RemoveEntityMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RemoveEntityMembers_Call) Return(err error) *Repository_RemoveEntityMembers_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RemoveEntityMembers_Call) RunAndReturn(run func(ctx context.Context, entityID string, members []string) error) *Repository_RemoveEntityMembers_Call { - _c.Call.Return(run) - return _c -} - -// RemoveMemberFromAllRoles provides a mock function for the type Repository -func (_mock *Repository) RemoveMemberFromAllRoles(ctx context.Context, memberID string) error { - ret := _mock.Called(ctx, memberID) - - if len(ret) == 0 { - panic("no return value specified for RemoveMemberFromAllRoles") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) error); ok { - r0 = returnFunc(ctx, memberID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RemoveMemberFromAllRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveMemberFromAllRoles' -type Repository_RemoveMemberFromAllRoles_Call struct { - *mock.Call -} - -// RemoveMemberFromAllRoles is a helper method to define mock.On call -// - ctx context.Context -// - memberID string -func (_e *Repository_Expecter) RemoveMemberFromAllRoles(ctx interface{}, memberID interface{}) *Repository_RemoveMemberFromAllRoles_Call { - return &Repository_RemoveMemberFromAllRoles_Call{Call: _e.mock.On("RemoveMemberFromAllRoles", ctx, memberID)} -} - -func (_c *Repository_RemoveMemberFromAllRoles_Call) Run(run func(ctx context.Context, memberID string)) *Repository_RemoveMemberFromAllRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RemoveMemberFromAllRoles_Call) Return(err error) *Repository_RemoveMemberFromAllRoles_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RemoveMemberFromAllRoles_Call) RunAndReturn(run func(ctx context.Context, memberID string) error) *Repository_RemoveMemberFromAllRoles_Call { - _c.Call.Return(run) - return _c -} - -// RemoveRoles provides a mock function for the type Repository -func (_mock *Repository) RemoveRoles(ctx context.Context, roleIDs []string) error { - ret := _mock.Called(ctx, roleIDs) - - if len(ret) == 0 { - panic("no return value specified for RemoveRoles") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, []string) error); ok { - r0 = returnFunc(ctx, roleIDs) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RemoveRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveRoles' -type Repository_RemoveRoles_Call struct { - *mock.Call -} - -// RemoveRoles is a helper method to define mock.On call -// - ctx context.Context -// - roleIDs []string -func (_e *Repository_Expecter) RemoveRoles(ctx interface{}, roleIDs interface{}) *Repository_RemoveRoles_Call { - return &Repository_RemoveRoles_Call{Call: _e.mock.On("RemoveRoles", ctx, roleIDs)} -} - -func (_c *Repository_RemoveRoles_Call) Run(run func(ctx context.Context, roleIDs []string)) *Repository_RemoveRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 []string - if args[1] != nil { - arg1 = args[1].([]string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RemoveRoles_Call) Return(err error) *Repository_RemoveRoles_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RemoveRoles_Call) RunAndReturn(run func(ctx context.Context, roleIDs []string) error) *Repository_RemoveRoles_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveAll provides a mock function for the type Repository -func (_mock *Repository) RetrieveAll(ctx context.Context, pm groups.PageMeta) (groups.Page, error) { - ret := _mock.Called(ctx, pm) - - if len(ret) == 0 { - panic("no return value specified for RetrieveAll") - } - - var r0 groups.Page - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, groups.PageMeta) (groups.Page, error)); ok { - return returnFunc(ctx, pm) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, groups.PageMeta) groups.Page); ok { - r0 = returnFunc(ctx, pm) - } else { - r0 = ret.Get(0).(groups.Page) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, groups.PageMeta) error); ok { - r1 = returnFunc(ctx, pm) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveAll_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveAll' -type Repository_RetrieveAll_Call struct { - *mock.Call -} - -// RetrieveAll is a helper method to define mock.On call -// - ctx context.Context -// - pm groups.PageMeta -func (_e *Repository_Expecter) RetrieveAll(ctx interface{}, pm interface{}) *Repository_RetrieveAll_Call { - return &Repository_RetrieveAll_Call{Call: _e.mock.On("RetrieveAll", ctx, pm)} -} - -func (_c *Repository_RetrieveAll_Call) Run(run func(ctx context.Context, pm groups.PageMeta)) *Repository_RetrieveAll_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 groups.PageMeta - if args[1] != nil { - arg1 = args[1].(groups.PageMeta) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RetrieveAll_Call) Return(page groups.Page, err error) *Repository_RetrieveAll_Call { - _c.Call.Return(page, err) - return _c -} - -func (_c *Repository_RetrieveAll_Call) RunAndReturn(run func(ctx context.Context, pm groups.PageMeta) (groups.Page, error)) *Repository_RetrieveAll_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveAllParentGroups provides a mock function for the type Repository -func (_mock *Repository) RetrieveAllParentGroups(ctx context.Context, domainID string, userID string, groupID string, pm groups.PageMeta) (groups.Page, error) { - ret := _mock.Called(ctx, domainID, userID, groupID, pm) - - if len(ret) == 0 { - panic("no return value specified for RetrieveAllParentGroups") - } - - var r0 groups.Page - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string, string, groups.PageMeta) (groups.Page, error)); ok { - return returnFunc(ctx, domainID, userID, groupID, pm) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string, string, groups.PageMeta) groups.Page); ok { - r0 = returnFunc(ctx, domainID, userID, groupID, pm) - } else { - r0 = ret.Get(0).(groups.Page) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, string, string, groups.PageMeta) error); ok { - r1 = returnFunc(ctx, domainID, userID, groupID, pm) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveAllParentGroups_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveAllParentGroups' -type Repository_RetrieveAllParentGroups_Call struct { - *mock.Call -} - -// RetrieveAllParentGroups is a helper method to define mock.On call -// - ctx context.Context -// - domainID string -// - userID string -// - groupID string -// - pm groups.PageMeta -func (_e *Repository_Expecter) RetrieveAllParentGroups(ctx interface{}, domainID interface{}, userID interface{}, groupID interface{}, pm interface{}) *Repository_RetrieveAllParentGroups_Call { - return &Repository_RetrieveAllParentGroups_Call{Call: _e.mock.On("RetrieveAllParentGroups", ctx, domainID, userID, groupID, pm)} -} - -func (_c *Repository_RetrieveAllParentGroups_Call) Run(run func(ctx context.Context, domainID string, userID string, groupID string, pm groups.PageMeta)) *Repository_RetrieveAllParentGroups_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 groups.PageMeta - if args[4] != nil { - arg4 = args[4].(groups.PageMeta) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Repository_RetrieveAllParentGroups_Call) Return(page groups.Page, err error) *Repository_RetrieveAllParentGroups_Call { - _c.Call.Return(page, err) - return _c -} - -func (_c *Repository_RetrieveAllParentGroups_Call) RunAndReturn(run func(ctx context.Context, domainID string, userID string, groupID string, pm groups.PageMeta) (groups.Page, error)) *Repository_RetrieveAllParentGroups_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveAllRoles provides a mock function for the type Repository -func (_mock *Repository) RetrieveAllRoles(ctx context.Context, entityID string, limit uint64, offset uint64) (roles.RolePage, error) { - ret := _mock.Called(ctx, entityID, limit, offset) - - if len(ret) == 0 { - panic("no return value specified for RetrieveAllRoles") - } - - var r0 roles.RolePage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, uint64, uint64) (roles.RolePage, error)); ok { - return returnFunc(ctx, entityID, limit, offset) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, uint64, uint64) roles.RolePage); ok { - r0 = returnFunc(ctx, entityID, limit, offset) - } else { - r0 = ret.Get(0).(roles.RolePage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, uint64, uint64) error); ok { - r1 = returnFunc(ctx, entityID, limit, offset) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveAllRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveAllRoles' -type Repository_RetrieveAllRoles_Call struct { - *mock.Call -} - -// RetrieveAllRoles is a helper method to define mock.On call -// - ctx context.Context -// - entityID string -// - limit uint64 -// - offset uint64 -func (_e *Repository_Expecter) RetrieveAllRoles(ctx interface{}, entityID interface{}, limit interface{}, offset interface{}) *Repository_RetrieveAllRoles_Call { - return &Repository_RetrieveAllRoles_Call{Call: _e.mock.On("RetrieveAllRoles", ctx, entityID, limit, offset)} -} - -func (_c *Repository_RetrieveAllRoles_Call) Run(run func(ctx context.Context, entityID string, limit uint64, offset uint64)) *Repository_RetrieveAllRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 uint64 - if args[2] != nil { - arg2 = args[2].(uint64) - } - var arg3 uint64 - if args[3] != nil { - arg3 = args[3].(uint64) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Repository_RetrieveAllRoles_Call) Return(rolePage roles.RolePage, err error) *Repository_RetrieveAllRoles_Call { - _c.Call.Return(rolePage, err) - return _c -} - -func (_c *Repository_RetrieveAllRoles_Call) RunAndReturn(run func(ctx context.Context, entityID string, limit uint64, offset uint64) (roles.RolePage, error)) *Repository_RetrieveAllRoles_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveByID provides a mock function for the type Repository -func (_mock *Repository) RetrieveByID(ctx context.Context, id string) (groups.Group, error) { - ret := _mock.Called(ctx, id) - - if len(ret) == 0 { - panic("no return value specified for RetrieveByID") - } - - var r0 groups.Group - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) (groups.Group, error)); ok { - return returnFunc(ctx, id) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string) groups.Group); ok { - r0 = returnFunc(ctx, id) - } else { - r0 = ret.Get(0).(groups.Group) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = returnFunc(ctx, id) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveByID_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveByID' -type Repository_RetrieveByID_Call struct { - *mock.Call -} - -// RetrieveByID is a helper method to define mock.On call -// - ctx context.Context -// - id string -func (_e *Repository_Expecter) RetrieveByID(ctx interface{}, id interface{}) *Repository_RetrieveByID_Call { - return &Repository_RetrieveByID_Call{Call: _e.mock.On("RetrieveByID", ctx, id)} -} - -func (_c *Repository_RetrieveByID_Call) Run(run func(ctx context.Context, id string)) *Repository_RetrieveByID_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RetrieveByID_Call) Return(group groups.Group, err error) *Repository_RetrieveByID_Call { - _c.Call.Return(group, err) - return _c -} - -func (_c *Repository_RetrieveByID_Call) RunAndReturn(run func(ctx context.Context, id string) (groups.Group, error)) *Repository_RetrieveByID_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveByIDAndUser provides a mock function for the type Repository -func (_mock *Repository) RetrieveByIDAndUser(ctx context.Context, domainID string, userID string, groupID string) (groups.Group, error) { - ret := _mock.Called(ctx, domainID, userID, groupID) - - if len(ret) == 0 { - panic("no return value specified for RetrieveByIDAndUser") - } - - var r0 groups.Group - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string, string) (groups.Group, error)); ok { - return returnFunc(ctx, domainID, userID, groupID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string, string) groups.Group); ok { - r0 = returnFunc(ctx, domainID, userID, groupID) - } else { - r0 = ret.Get(0).(groups.Group) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, string, string) error); ok { - r1 = returnFunc(ctx, domainID, userID, groupID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveByIDAndUser_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveByIDAndUser' -type Repository_RetrieveByIDAndUser_Call struct { - *mock.Call -} - -// RetrieveByIDAndUser is a helper method to define mock.On call -// - ctx context.Context -// - domainID string -// - userID string -// - groupID string -func (_e *Repository_Expecter) RetrieveByIDAndUser(ctx interface{}, domainID interface{}, userID interface{}, groupID interface{}) *Repository_RetrieveByIDAndUser_Call { - return &Repository_RetrieveByIDAndUser_Call{Call: _e.mock.On("RetrieveByIDAndUser", ctx, domainID, userID, groupID)} -} - -func (_c *Repository_RetrieveByIDAndUser_Call) Run(run func(ctx context.Context, domainID string, userID string, groupID string)) *Repository_RetrieveByIDAndUser_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Repository_RetrieveByIDAndUser_Call) Return(group groups.Group, err error) *Repository_RetrieveByIDAndUser_Call { - _c.Call.Return(group, err) - return _c -} - -func (_c *Repository_RetrieveByIDAndUser_Call) RunAndReturn(run func(ctx context.Context, domainID string, userID string, groupID string) (groups.Group, error)) *Repository_RetrieveByIDAndUser_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveByIDWithRoles provides a mock function for the type Repository -func (_mock *Repository) RetrieveByIDWithRoles(ctx context.Context, groupID string, memberID string) (groups.Group, error) { - ret := _mock.Called(ctx, groupID, memberID) - - if len(ret) == 0 { - panic("no return value specified for RetrieveByIDWithRoles") - } - - var r0 groups.Group - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) (groups.Group, error)); ok { - return returnFunc(ctx, groupID, memberID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) groups.Group); ok { - r0 = returnFunc(ctx, groupID, memberID) - } else { - r0 = ret.Get(0).(groups.Group) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, string) error); ok { - r1 = returnFunc(ctx, groupID, memberID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveByIDWithRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveByIDWithRoles' -type Repository_RetrieveByIDWithRoles_Call struct { - *mock.Call -} - -// RetrieveByIDWithRoles is a helper method to define mock.On call -// - ctx context.Context -// - groupID string -// - memberID string -func (_e *Repository_Expecter) RetrieveByIDWithRoles(ctx interface{}, groupID interface{}, memberID interface{}) *Repository_RetrieveByIDWithRoles_Call { - return &Repository_RetrieveByIDWithRoles_Call{Call: _e.mock.On("RetrieveByIDWithRoles", ctx, groupID, memberID)} -} - -func (_c *Repository_RetrieveByIDWithRoles_Call) Run(run func(ctx context.Context, groupID string, memberID string)) *Repository_RetrieveByIDWithRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RetrieveByIDWithRoles_Call) Return(group groups.Group, err error) *Repository_RetrieveByIDWithRoles_Call { - _c.Call.Return(group, err) - return _c -} - -func (_c *Repository_RetrieveByIDWithRoles_Call) RunAndReturn(run func(ctx context.Context, groupID string, memberID string) (groups.Group, error)) *Repository_RetrieveByIDWithRoles_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveByIDs provides a mock function for the type Repository -func (_mock *Repository) RetrieveByIDs(ctx context.Context, pm groups.PageMeta, ids ...string) (groups.Page, error) { - var tmpRet mock.Arguments - if len(ids) > 0 { - tmpRet = _mock.Called(ctx, pm, ids) - } else { - tmpRet = _mock.Called(ctx, pm) - } - ret := tmpRet - - if len(ret) == 0 { - panic("no return value specified for RetrieveByIDs") - } - - var r0 groups.Page - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, groups.PageMeta, ...string) (groups.Page, error)); ok { - return returnFunc(ctx, pm, ids...) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, groups.PageMeta, ...string) groups.Page); ok { - r0 = returnFunc(ctx, pm, ids...) - } else { - r0 = ret.Get(0).(groups.Page) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, groups.PageMeta, ...string) error); ok { - r1 = returnFunc(ctx, pm, ids...) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveByIDs_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveByIDs' -type Repository_RetrieveByIDs_Call struct { - *mock.Call -} - -// RetrieveByIDs is a helper method to define mock.On call -// - ctx context.Context -// - pm groups.PageMeta -// - ids ...string -func (_e *Repository_Expecter) RetrieveByIDs(ctx interface{}, pm interface{}, ids ...interface{}) *Repository_RetrieveByIDs_Call { - return &Repository_RetrieveByIDs_Call{Call: _e.mock.On("RetrieveByIDs", - append([]interface{}{ctx, pm}, ids...)...)} -} - -func (_c *Repository_RetrieveByIDs_Call) Run(run func(ctx context.Context, pm groups.PageMeta, ids ...string)) *Repository_RetrieveByIDs_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 groups.PageMeta - if args[1] != nil { - arg1 = args[1].(groups.PageMeta) - } - var arg2 []string - var variadicArgs []string - if len(args) > 2 { - variadicArgs = args[2].([]string) - } - arg2 = variadicArgs - run( - arg0, - arg1, - arg2..., - ) - }) - return _c -} - -func (_c *Repository_RetrieveByIDs_Call) Return(page groups.Page, err error) *Repository_RetrieveByIDs_Call { - _c.Call.Return(page, err) - return _c -} - -func (_c *Repository_RetrieveByIDs_Call) RunAndReturn(run func(ctx context.Context, pm groups.PageMeta, ids ...string) (groups.Page, error)) *Repository_RetrieveByIDs_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveChildrenGroups provides a mock function for the type Repository -func (_mock *Repository) RetrieveChildrenGroups(ctx context.Context, domainID string, userID string, groupID string, startLevel int64, endLevel int64, pm groups.PageMeta) (groups.Page, error) { - ret := _mock.Called(ctx, domainID, userID, groupID, startLevel, endLevel, pm) - - if len(ret) == 0 { - panic("no return value specified for RetrieveChildrenGroups") - } - - var r0 groups.Page - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string, string, int64, int64, groups.PageMeta) (groups.Page, error)); ok { - return returnFunc(ctx, domainID, userID, groupID, startLevel, endLevel, pm) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string, string, int64, int64, groups.PageMeta) groups.Page); ok { - r0 = returnFunc(ctx, domainID, userID, groupID, startLevel, endLevel, pm) - } else { - r0 = ret.Get(0).(groups.Page) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, string, string, int64, int64, groups.PageMeta) error); ok { - r1 = returnFunc(ctx, domainID, userID, groupID, startLevel, endLevel, pm) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveChildrenGroups_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveChildrenGroups' -type Repository_RetrieveChildrenGroups_Call struct { - *mock.Call -} - -// RetrieveChildrenGroups is a helper method to define mock.On call -// - ctx context.Context -// - domainID string -// - userID string -// - groupID string -// - startLevel int64 -// - endLevel int64 -// - pm groups.PageMeta -func (_e *Repository_Expecter) RetrieveChildrenGroups(ctx interface{}, domainID interface{}, userID interface{}, groupID interface{}, startLevel interface{}, endLevel interface{}, pm interface{}) *Repository_RetrieveChildrenGroups_Call { - return &Repository_RetrieveChildrenGroups_Call{Call: _e.mock.On("RetrieveChildrenGroups", ctx, domainID, userID, groupID, startLevel, endLevel, pm)} -} - -func (_c *Repository_RetrieveChildrenGroups_Call) Run(run func(ctx context.Context, domainID string, userID string, groupID string, startLevel int64, endLevel int64, pm groups.PageMeta)) *Repository_RetrieveChildrenGroups_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 int64 - if args[4] != nil { - arg4 = args[4].(int64) - } - var arg5 int64 - if args[5] != nil { - arg5 = args[5].(int64) - } - var arg6 groups.PageMeta - if args[6] != nil { - arg6 = args[6].(groups.PageMeta) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - arg5, - arg6, - ) - }) - return _c -} - -func (_c *Repository_RetrieveChildrenGroups_Call) Return(page groups.Page, err error) *Repository_RetrieveChildrenGroups_Call { - _c.Call.Return(page, err) - return _c -} - -func (_c *Repository_RetrieveChildrenGroups_Call) RunAndReturn(run func(ctx context.Context, domainID string, userID string, groupID string, startLevel int64, endLevel int64, pm groups.PageMeta) (groups.Page, error)) *Repository_RetrieveChildrenGroups_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveEntitiesRolesActionsMembers provides a mock function for the type Repository -func (_mock *Repository) RetrieveEntitiesRolesActionsMembers(ctx context.Context, entityIDs []string) ([]roles.EntityActionRole, []roles.EntityMemberRole, error) { - ret := _mock.Called(ctx, entityIDs) - - if len(ret) == 0 { - panic("no return value specified for RetrieveEntitiesRolesActionsMembers") - } - - var r0 []roles.EntityActionRole - var r1 []roles.EntityMemberRole - var r2 error - if returnFunc, ok := ret.Get(0).(func(context.Context, []string) ([]roles.EntityActionRole, []roles.EntityMemberRole, error)); ok { - return returnFunc(ctx, entityIDs) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, []string) []roles.EntityActionRole); ok { - r0 = returnFunc(ctx, entityIDs) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]roles.EntityActionRole) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, []string) []roles.EntityMemberRole); ok { - r1 = returnFunc(ctx, entityIDs) - } else { - if ret.Get(1) != nil { - r1 = ret.Get(1).([]roles.EntityMemberRole) - } - } - if returnFunc, ok := ret.Get(2).(func(context.Context, []string) error); ok { - r2 = returnFunc(ctx, entityIDs) - } else { - r2 = ret.Error(2) - } - return r0, r1, r2 -} - -// Repository_RetrieveEntitiesRolesActionsMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveEntitiesRolesActionsMembers' -type Repository_RetrieveEntitiesRolesActionsMembers_Call struct { - *mock.Call -} - -// RetrieveEntitiesRolesActionsMembers is a helper method to define mock.On call -// - ctx context.Context -// - entityIDs []string -func (_e *Repository_Expecter) RetrieveEntitiesRolesActionsMembers(ctx interface{}, entityIDs interface{}) *Repository_RetrieveEntitiesRolesActionsMembers_Call { - return &Repository_RetrieveEntitiesRolesActionsMembers_Call{Call: _e.mock.On("RetrieveEntitiesRolesActionsMembers", ctx, entityIDs)} -} - -func (_c *Repository_RetrieveEntitiesRolesActionsMembers_Call) Run(run func(ctx context.Context, entityIDs []string)) *Repository_RetrieveEntitiesRolesActionsMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 []string - if args[1] != nil { - arg1 = args[1].([]string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RetrieveEntitiesRolesActionsMembers_Call) Return(entityActionRoles []roles.EntityActionRole, entityMemberRoles []roles.EntityMemberRole, err error) *Repository_RetrieveEntitiesRolesActionsMembers_Call { - _c.Call.Return(entityActionRoles, entityMemberRoles, err) - return _c -} - -func (_c *Repository_RetrieveEntitiesRolesActionsMembers_Call) RunAndReturn(run func(ctx context.Context, entityIDs []string) ([]roles.EntityActionRole, []roles.EntityMemberRole, error)) *Repository_RetrieveEntitiesRolesActionsMembers_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveEntityRole provides a mock function for the type Repository -func (_mock *Repository) RetrieveEntityRole(ctx context.Context, entityID string, roleID string) (roles.Role, error) { - ret := _mock.Called(ctx, entityID, roleID) - - if len(ret) == 0 { - panic("no return value specified for RetrieveEntityRole") - } - - var r0 roles.Role - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) (roles.Role, error)); ok { - return returnFunc(ctx, entityID, roleID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) roles.Role); ok { - r0 = returnFunc(ctx, entityID, roleID) - } else { - r0 = ret.Get(0).(roles.Role) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, string) error); ok { - r1 = returnFunc(ctx, entityID, roleID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveEntityRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveEntityRole' -type Repository_RetrieveEntityRole_Call struct { - *mock.Call -} - -// RetrieveEntityRole is a helper method to define mock.On call -// - ctx context.Context -// - entityID string -// - roleID string -func (_e *Repository_Expecter) RetrieveEntityRole(ctx interface{}, entityID interface{}, roleID interface{}) *Repository_RetrieveEntityRole_Call { - return &Repository_RetrieveEntityRole_Call{Call: _e.mock.On("RetrieveEntityRole", ctx, entityID, roleID)} -} - -func (_c *Repository_RetrieveEntityRole_Call) Run(run func(ctx context.Context, entityID string, roleID string)) *Repository_RetrieveEntityRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RetrieveEntityRole_Call) Return(role roles.Role, err error) *Repository_RetrieveEntityRole_Call { - _c.Call.Return(role, err) - return _c -} - -func (_c *Repository_RetrieveEntityRole_Call) RunAndReturn(run func(ctx context.Context, entityID string, roleID string) (roles.Role, error)) *Repository_RetrieveEntityRole_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveHierarchy provides a mock function for the type Repository -func (_mock *Repository) RetrieveHierarchy(ctx context.Context, domainID string, userID string, groupID string, hm groups.HierarchyPageMeta) (groups.HierarchyPage, error) { - ret := _mock.Called(ctx, domainID, userID, groupID, hm) - - if len(ret) == 0 { - panic("no return value specified for RetrieveHierarchy") - } - - var r0 groups.HierarchyPage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string, string, groups.HierarchyPageMeta) (groups.HierarchyPage, error)); ok { - return returnFunc(ctx, domainID, userID, groupID, hm) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string, string, groups.HierarchyPageMeta) groups.HierarchyPage); ok { - r0 = returnFunc(ctx, domainID, userID, groupID, hm) - } else { - r0 = ret.Get(0).(groups.HierarchyPage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, string, string, groups.HierarchyPageMeta) error); ok { - r1 = returnFunc(ctx, domainID, userID, groupID, hm) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveHierarchy_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveHierarchy' -type Repository_RetrieveHierarchy_Call struct { - *mock.Call -} - -// RetrieveHierarchy is a helper method to define mock.On call -// - ctx context.Context -// - domainID string -// - userID string -// - groupID string -// - hm groups.HierarchyPageMeta -func (_e *Repository_Expecter) RetrieveHierarchy(ctx interface{}, domainID interface{}, userID interface{}, groupID interface{}, hm interface{}) *Repository_RetrieveHierarchy_Call { - return &Repository_RetrieveHierarchy_Call{Call: _e.mock.On("RetrieveHierarchy", ctx, domainID, userID, groupID, hm)} -} - -func (_c *Repository_RetrieveHierarchy_Call) Run(run func(ctx context.Context, domainID string, userID string, groupID string, hm groups.HierarchyPageMeta)) *Repository_RetrieveHierarchy_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 groups.HierarchyPageMeta - if args[4] != nil { - arg4 = args[4].(groups.HierarchyPageMeta) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Repository_RetrieveHierarchy_Call) Return(hierarchyPage groups.HierarchyPage, err error) *Repository_RetrieveHierarchy_Call { - _c.Call.Return(hierarchyPage, err) - return _c -} - -func (_c *Repository_RetrieveHierarchy_Call) RunAndReturn(run func(ctx context.Context, domainID string, userID string, groupID string, hm groups.HierarchyPageMeta) (groups.HierarchyPage, error)) *Repository_RetrieveHierarchy_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveRole provides a mock function for the type Repository -func (_mock *Repository) RetrieveRole(ctx context.Context, roleID string) (roles.Role, error) { - ret := _mock.Called(ctx, roleID) - - if len(ret) == 0 { - panic("no return value specified for RetrieveRole") - } - - var r0 roles.Role - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) (roles.Role, error)); ok { - return returnFunc(ctx, roleID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string) roles.Role); ok { - r0 = returnFunc(ctx, roleID) - } else { - r0 = ret.Get(0).(roles.Role) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = returnFunc(ctx, roleID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveRole' -type Repository_RetrieveRole_Call struct { - *mock.Call -} - -// RetrieveRole is a helper method to define mock.On call -// - ctx context.Context -// - roleID string -func (_e *Repository_Expecter) RetrieveRole(ctx interface{}, roleID interface{}) *Repository_RetrieveRole_Call { - return &Repository_RetrieveRole_Call{Call: _e.mock.On("RetrieveRole", ctx, roleID)} -} - -func (_c *Repository_RetrieveRole_Call) Run(run func(ctx context.Context, roleID string)) *Repository_RetrieveRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RetrieveRole_Call) Return(role roles.Role, err error) *Repository_RetrieveRole_Call { - _c.Call.Return(role, err) - return _c -} - -func (_c *Repository_RetrieveRole_Call) RunAndReturn(run func(ctx context.Context, roleID string) (roles.Role, error)) *Repository_RetrieveRole_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveUserGroups provides a mock function for the type Repository -func (_mock *Repository) RetrieveUserGroups(ctx context.Context, domainID string, userID string, pm groups.PageMeta) (groups.Page, error) { - ret := _mock.Called(ctx, domainID, userID, pm) - - if len(ret) == 0 { - panic("no return value specified for RetrieveUserGroups") - } - - var r0 groups.Page - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string, groups.PageMeta) (groups.Page, error)); ok { - return returnFunc(ctx, domainID, userID, pm) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string, groups.PageMeta) groups.Page); ok { - r0 = returnFunc(ctx, domainID, userID, pm) - } else { - r0 = ret.Get(0).(groups.Page) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, string, groups.PageMeta) error); ok { - r1 = returnFunc(ctx, domainID, userID, pm) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveUserGroups_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveUserGroups' -type Repository_RetrieveUserGroups_Call struct { - *mock.Call -} - -// RetrieveUserGroups is a helper method to define mock.On call -// - ctx context.Context -// - domainID string -// - userID string -// - pm groups.PageMeta -func (_e *Repository_Expecter) RetrieveUserGroups(ctx interface{}, domainID interface{}, userID interface{}, pm interface{}) *Repository_RetrieveUserGroups_Call { - return &Repository_RetrieveUserGroups_Call{Call: _e.mock.On("RetrieveUserGroups", ctx, domainID, userID, pm)} -} - -func (_c *Repository_RetrieveUserGroups_Call) Run(run func(ctx context.Context, domainID string, userID string, pm groups.PageMeta)) *Repository_RetrieveUserGroups_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 groups.PageMeta - if args[3] != nil { - arg3 = args[3].(groups.PageMeta) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Repository_RetrieveUserGroups_Call) Return(page groups.Page, err error) *Repository_RetrieveUserGroups_Call { - _c.Call.Return(page, err) - return _c -} - -func (_c *Repository_RetrieveUserGroups_Call) RunAndReturn(run func(ctx context.Context, domainID string, userID string, pm groups.PageMeta) (groups.Page, error)) *Repository_RetrieveUserGroups_Call { - _c.Call.Return(run) - return _c -} - -// RoleAddActions provides a mock function for the type Repository -func (_mock *Repository) RoleAddActions(ctx context.Context, role roles.Role, actions []string) ([]string, error) { - ret := _mock.Called(ctx, role, actions) - - if len(ret) == 0 { - panic("no return value specified for RoleAddActions") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role, []string) ([]string, error)); ok { - return returnFunc(ctx, role, actions) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role, []string) []string); ok { - r0 = returnFunc(ctx, role, actions) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, roles.Role, []string) error); ok { - r1 = returnFunc(ctx, role, actions) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RoleAddActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleAddActions' -type Repository_RoleAddActions_Call struct { - *mock.Call -} - -// RoleAddActions is a helper method to define mock.On call -// - ctx context.Context -// - role roles.Role -// - actions []string -func (_e *Repository_Expecter) RoleAddActions(ctx interface{}, role interface{}, actions interface{}) *Repository_RoleAddActions_Call { - return &Repository_RoleAddActions_Call{Call: _e.mock.On("RoleAddActions", ctx, role, actions)} -} - -func (_c *Repository_RoleAddActions_Call) Run(run func(ctx context.Context, role roles.Role, actions []string)) *Repository_RoleAddActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RoleAddActions_Call) Return(ops []string, err error) *Repository_RoleAddActions_Call { - _c.Call.Return(ops, err) - return _c -} - -func (_c *Repository_RoleAddActions_Call) RunAndReturn(run func(ctx context.Context, role roles.Role, actions []string) ([]string, error)) *Repository_RoleAddActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleAddMembers provides a mock function for the type Repository -func (_mock *Repository) RoleAddMembers(ctx context.Context, role roles.Role, members []string) ([]string, error) { - ret := _mock.Called(ctx, role, members) - - if len(ret) == 0 { - panic("no return value specified for RoleAddMembers") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role, []string) ([]string, error)); ok { - return returnFunc(ctx, role, members) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role, []string) []string); ok { - r0 = returnFunc(ctx, role, members) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, roles.Role, []string) error); ok { - r1 = returnFunc(ctx, role, members) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RoleAddMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleAddMembers' -type Repository_RoleAddMembers_Call struct { - *mock.Call -} - -// RoleAddMembers is a helper method to define mock.On call -// - ctx context.Context -// - role roles.Role -// - members []string -func (_e *Repository_Expecter) RoleAddMembers(ctx interface{}, role interface{}, members interface{}) *Repository_RoleAddMembers_Call { - return &Repository_RoleAddMembers_Call{Call: _e.mock.On("RoleAddMembers", ctx, role, members)} -} - -func (_c *Repository_RoleAddMembers_Call) Run(run func(ctx context.Context, role roles.Role, members []string)) *Repository_RoleAddMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RoleAddMembers_Call) Return(strings []string, err error) *Repository_RoleAddMembers_Call { - _c.Call.Return(strings, err) - return _c -} - -func (_c *Repository_RoleAddMembers_Call) RunAndReturn(run func(ctx context.Context, role roles.Role, members []string) ([]string, error)) *Repository_RoleAddMembers_Call { - _c.Call.Return(run) - return _c -} - -// RoleCheckActionsExists provides a mock function for the type Repository -func (_mock *Repository) RoleCheckActionsExists(ctx context.Context, roleID string, actions []string) (bool, error) { - ret := _mock.Called(ctx, roleID, actions) - - if len(ret) == 0 { - panic("no return value specified for RoleCheckActionsExists") - } - - var r0 bool - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, []string) (bool, error)); ok { - return returnFunc(ctx, roleID, actions) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, []string) bool); ok { - r0 = returnFunc(ctx, roleID, actions) - } else { - r0 = ret.Get(0).(bool) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, []string) error); ok { - r1 = returnFunc(ctx, roleID, actions) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RoleCheckActionsExists_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleCheckActionsExists' -type Repository_RoleCheckActionsExists_Call struct { - *mock.Call -} - -// RoleCheckActionsExists is a helper method to define mock.On call -// - ctx context.Context -// - roleID string -// - actions []string -func (_e *Repository_Expecter) RoleCheckActionsExists(ctx interface{}, roleID interface{}, actions interface{}) *Repository_RoleCheckActionsExists_Call { - return &Repository_RoleCheckActionsExists_Call{Call: _e.mock.On("RoleCheckActionsExists", ctx, roleID, actions)} -} - -func (_c *Repository_RoleCheckActionsExists_Call) Run(run func(ctx context.Context, roleID string, actions []string)) *Repository_RoleCheckActionsExists_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RoleCheckActionsExists_Call) Return(b bool, err error) *Repository_RoleCheckActionsExists_Call { - _c.Call.Return(b, err) - return _c -} - -func (_c *Repository_RoleCheckActionsExists_Call) RunAndReturn(run func(ctx context.Context, roleID string, actions []string) (bool, error)) *Repository_RoleCheckActionsExists_Call { - _c.Call.Return(run) - return _c -} - -// RoleCheckMembersExists provides a mock function for the type Repository -func (_mock *Repository) RoleCheckMembersExists(ctx context.Context, roleID string, members []string) (bool, error) { - ret := _mock.Called(ctx, roleID, members) - - if len(ret) == 0 { - panic("no return value specified for RoleCheckMembersExists") - } - - var r0 bool - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, []string) (bool, error)); ok { - return returnFunc(ctx, roleID, members) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, []string) bool); ok { - r0 = returnFunc(ctx, roleID, members) - } else { - r0 = ret.Get(0).(bool) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, []string) error); ok { - r1 = returnFunc(ctx, roleID, members) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RoleCheckMembersExists_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleCheckMembersExists' -type Repository_RoleCheckMembersExists_Call struct { - *mock.Call -} - -// RoleCheckMembersExists is a helper method to define mock.On call -// - ctx context.Context -// - roleID string -// - members []string -func (_e *Repository_Expecter) RoleCheckMembersExists(ctx interface{}, roleID interface{}, members interface{}) *Repository_RoleCheckMembersExists_Call { - return &Repository_RoleCheckMembersExists_Call{Call: _e.mock.On("RoleCheckMembersExists", ctx, roleID, members)} -} - -func (_c *Repository_RoleCheckMembersExists_Call) Run(run func(ctx context.Context, roleID string, members []string)) *Repository_RoleCheckMembersExists_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RoleCheckMembersExists_Call) Return(b bool, err error) *Repository_RoleCheckMembersExists_Call { - _c.Call.Return(b, err) - return _c -} - -func (_c *Repository_RoleCheckMembersExists_Call) RunAndReturn(run func(ctx context.Context, roleID string, members []string) (bool, error)) *Repository_RoleCheckMembersExists_Call { - _c.Call.Return(run) - return _c -} - -// RoleListActions provides a mock function for the type Repository -func (_mock *Repository) RoleListActions(ctx context.Context, roleID string) ([]string, error) { - ret := _mock.Called(ctx, roleID) - - if len(ret) == 0 { - panic("no return value specified for RoleListActions") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) ([]string, error)); ok { - return returnFunc(ctx, roleID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string) []string); ok { - r0 = returnFunc(ctx, roleID) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = returnFunc(ctx, roleID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RoleListActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleListActions' -type Repository_RoleListActions_Call struct { - *mock.Call -} - -// RoleListActions is a helper method to define mock.On call -// - ctx context.Context -// - roleID string -func (_e *Repository_Expecter) RoleListActions(ctx interface{}, roleID interface{}) *Repository_RoleListActions_Call { - return &Repository_RoleListActions_Call{Call: _e.mock.On("RoleListActions", ctx, roleID)} -} - -func (_c *Repository_RoleListActions_Call) Run(run func(ctx context.Context, roleID string)) *Repository_RoleListActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RoleListActions_Call) Return(strings []string, err error) *Repository_RoleListActions_Call { - _c.Call.Return(strings, err) - return _c -} - -func (_c *Repository_RoleListActions_Call) RunAndReturn(run func(ctx context.Context, roleID string) ([]string, error)) *Repository_RoleListActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleListMembers provides a mock function for the type Repository -func (_mock *Repository) RoleListMembers(ctx context.Context, roleID string, limit uint64, offset uint64) (roles.MembersPage, error) { - ret := _mock.Called(ctx, roleID, limit, offset) - - if len(ret) == 0 { - panic("no return value specified for RoleListMembers") - } - - var r0 roles.MembersPage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, uint64, uint64) (roles.MembersPage, error)); ok { - return returnFunc(ctx, roleID, limit, offset) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, uint64, uint64) roles.MembersPage); ok { - r0 = returnFunc(ctx, roleID, limit, offset) - } else { - r0 = ret.Get(0).(roles.MembersPage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, uint64, uint64) error); ok { - r1 = returnFunc(ctx, roleID, limit, offset) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RoleListMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleListMembers' -type Repository_RoleListMembers_Call struct { - *mock.Call -} - -// RoleListMembers is a helper method to define mock.On call -// - ctx context.Context -// - roleID string -// - limit uint64 -// - offset uint64 -func (_e *Repository_Expecter) RoleListMembers(ctx interface{}, roleID interface{}, limit interface{}, offset interface{}) *Repository_RoleListMembers_Call { - return &Repository_RoleListMembers_Call{Call: _e.mock.On("RoleListMembers", ctx, roleID, limit, offset)} -} - -func (_c *Repository_RoleListMembers_Call) Run(run func(ctx context.Context, roleID string, limit uint64, offset uint64)) *Repository_RoleListMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 uint64 - if args[2] != nil { - arg2 = args[2].(uint64) - } - var arg3 uint64 - if args[3] != nil { - arg3 = args[3].(uint64) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Repository_RoleListMembers_Call) Return(membersPage roles.MembersPage, err error) *Repository_RoleListMembers_Call { - _c.Call.Return(membersPage, err) - return _c -} - -func (_c *Repository_RoleListMembers_Call) RunAndReturn(run func(ctx context.Context, roleID string, limit uint64, offset uint64) (roles.MembersPage, error)) *Repository_RoleListMembers_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveActions provides a mock function for the type Repository -func (_mock *Repository) RoleRemoveActions(ctx context.Context, role roles.Role, actions []string) error { - ret := _mock.Called(ctx, role, actions) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveActions") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role, []string) error); ok { - r0 = returnFunc(ctx, role, actions) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RoleRemoveActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveActions' -type Repository_RoleRemoveActions_Call struct { - *mock.Call -} - -// RoleRemoveActions is a helper method to define mock.On call -// - ctx context.Context -// - role roles.Role -// - actions []string -func (_e *Repository_Expecter) RoleRemoveActions(ctx interface{}, role interface{}, actions interface{}) *Repository_RoleRemoveActions_Call { - return &Repository_RoleRemoveActions_Call{Call: _e.mock.On("RoleRemoveActions", ctx, role, actions)} -} - -func (_c *Repository_RoleRemoveActions_Call) Run(run func(ctx context.Context, role roles.Role, actions []string)) *Repository_RoleRemoveActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RoleRemoveActions_Call) Return(err error) *Repository_RoleRemoveActions_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RoleRemoveActions_Call) RunAndReturn(run func(ctx context.Context, role roles.Role, actions []string) error) *Repository_RoleRemoveActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveAllActions provides a mock function for the type Repository -func (_mock *Repository) RoleRemoveAllActions(ctx context.Context, role roles.Role) error { - ret := _mock.Called(ctx, role) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveAllActions") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role) error); ok { - r0 = returnFunc(ctx, role) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RoleRemoveAllActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveAllActions' -type Repository_RoleRemoveAllActions_Call struct { - *mock.Call -} - -// RoleRemoveAllActions is a helper method to define mock.On call -// - ctx context.Context -// - role roles.Role -func (_e *Repository_Expecter) RoleRemoveAllActions(ctx interface{}, role interface{}) *Repository_RoleRemoveAllActions_Call { - return &Repository_RoleRemoveAllActions_Call{Call: _e.mock.On("RoleRemoveAllActions", ctx, role)} -} - -func (_c *Repository_RoleRemoveAllActions_Call) Run(run func(ctx context.Context, role roles.Role)) *Repository_RoleRemoveAllActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RoleRemoveAllActions_Call) Return(err error) *Repository_RoleRemoveAllActions_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RoleRemoveAllActions_Call) RunAndReturn(run func(ctx context.Context, role roles.Role) error) *Repository_RoleRemoveAllActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveAllMembers provides a mock function for the type Repository -func (_mock *Repository) RoleRemoveAllMembers(ctx context.Context, role roles.Role) error { - ret := _mock.Called(ctx, role) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveAllMembers") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role) error); ok { - r0 = returnFunc(ctx, role) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RoleRemoveAllMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveAllMembers' -type Repository_RoleRemoveAllMembers_Call struct { - *mock.Call -} - -// RoleRemoveAllMembers is a helper method to define mock.On call -// - ctx context.Context -// - role roles.Role -func (_e *Repository_Expecter) RoleRemoveAllMembers(ctx interface{}, role interface{}) *Repository_RoleRemoveAllMembers_Call { - return &Repository_RoleRemoveAllMembers_Call{Call: _e.mock.On("RoleRemoveAllMembers", ctx, role)} -} - -func (_c *Repository_RoleRemoveAllMembers_Call) Run(run func(ctx context.Context, role roles.Role)) *Repository_RoleRemoveAllMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RoleRemoveAllMembers_Call) Return(err error) *Repository_RoleRemoveAllMembers_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RoleRemoveAllMembers_Call) RunAndReturn(run func(ctx context.Context, role roles.Role) error) *Repository_RoleRemoveAllMembers_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveMembers provides a mock function for the type Repository -func (_mock *Repository) RoleRemoveMembers(ctx context.Context, role roles.Role, members []string) error { - ret := _mock.Called(ctx, role, members) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveMembers") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role, []string) error); ok { - r0 = returnFunc(ctx, role, members) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RoleRemoveMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveMembers' -type Repository_RoleRemoveMembers_Call struct { - *mock.Call -} - -// RoleRemoveMembers is a helper method to define mock.On call -// - ctx context.Context -// - role roles.Role -// - members []string -func (_e *Repository_Expecter) RoleRemoveMembers(ctx interface{}, role interface{}, members interface{}) *Repository_RoleRemoveMembers_Call { - return &Repository_RoleRemoveMembers_Call{Call: _e.mock.On("RoleRemoveMembers", ctx, role, members)} -} - -func (_c *Repository_RoleRemoveMembers_Call) Run(run func(ctx context.Context, role roles.Role, members []string)) *Repository_RoleRemoveMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RoleRemoveMembers_Call) Return(err error) *Repository_RoleRemoveMembers_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RoleRemoveMembers_Call) RunAndReturn(run func(ctx context.Context, role roles.Role, members []string) error) *Repository_RoleRemoveMembers_Call { - _c.Call.Return(run) - return _c -} - -// Save provides a mock function for the type Repository -func (_mock *Repository) Save(ctx context.Context, g groups.Group) (groups.Group, error) { - ret := _mock.Called(ctx, g) - - if len(ret) == 0 { - panic("no return value specified for Save") - } - - var r0 groups.Group - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, groups.Group) (groups.Group, error)); ok { - return returnFunc(ctx, g) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, groups.Group) groups.Group); ok { - r0 = returnFunc(ctx, g) - } else { - r0 = ret.Get(0).(groups.Group) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, groups.Group) error); ok { - r1 = returnFunc(ctx, g) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_Save_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Save' -type Repository_Save_Call struct { - *mock.Call -} - -// Save is a helper method to define mock.On call -// - ctx context.Context -// - g groups.Group -func (_e *Repository_Expecter) Save(ctx interface{}, g interface{}) *Repository_Save_Call { - return &Repository_Save_Call{Call: _e.mock.On("Save", ctx, g)} -} - -func (_c *Repository_Save_Call) Run(run func(ctx context.Context, g groups.Group)) *Repository_Save_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 groups.Group - if args[1] != nil { - arg1 = args[1].(groups.Group) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_Save_Call) Return(group groups.Group, err error) *Repository_Save_Call { - _c.Call.Return(group, err) - return _c -} - -func (_c *Repository_Save_Call) RunAndReturn(run func(ctx context.Context, g groups.Group) (groups.Group, error)) *Repository_Save_Call { - _c.Call.Return(run) - return _c -} - -// UnassignAllChildrenGroups provides a mock function for the type Repository -func (_mock *Repository) UnassignAllChildrenGroups(ctx context.Context, id string) error { - ret := _mock.Called(ctx, id) - - if len(ret) == 0 { - panic("no return value specified for UnassignAllChildrenGroups") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) error); ok { - r0 = returnFunc(ctx, id) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_UnassignAllChildrenGroups_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UnassignAllChildrenGroups' -type Repository_UnassignAllChildrenGroups_Call struct { - *mock.Call -} - -// UnassignAllChildrenGroups is a helper method to define mock.On call -// - ctx context.Context -// - id string -func (_e *Repository_Expecter) UnassignAllChildrenGroups(ctx interface{}, id interface{}) *Repository_UnassignAllChildrenGroups_Call { - return &Repository_UnassignAllChildrenGroups_Call{Call: _e.mock.On("UnassignAllChildrenGroups", ctx, id)} -} - -func (_c *Repository_UnassignAllChildrenGroups_Call) Run(run func(ctx context.Context, id string)) *Repository_UnassignAllChildrenGroups_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_UnassignAllChildrenGroups_Call) Return(err error) *Repository_UnassignAllChildrenGroups_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_UnassignAllChildrenGroups_Call) RunAndReturn(run func(ctx context.Context, id string) error) *Repository_UnassignAllChildrenGroups_Call { - _c.Call.Return(run) - return _c -} - -// UnassignParentGroup provides a mock function for the type Repository -func (_mock *Repository) UnassignParentGroup(ctx context.Context, parentGroupID string, groupIDs ...string) error { - var tmpRet mock.Arguments - if len(groupIDs) > 0 { - tmpRet = _mock.Called(ctx, parentGroupID, groupIDs) - } else { - tmpRet = _mock.Called(ctx, parentGroupID) - } - ret := tmpRet - - if len(ret) == 0 { - panic("no return value specified for UnassignParentGroup") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, ...string) error); ok { - r0 = returnFunc(ctx, parentGroupID, groupIDs...) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_UnassignParentGroup_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UnassignParentGroup' -type Repository_UnassignParentGroup_Call struct { - *mock.Call -} - -// UnassignParentGroup is a helper method to define mock.On call -// - ctx context.Context -// - parentGroupID string -// - groupIDs ...string -func (_e *Repository_Expecter) UnassignParentGroup(ctx interface{}, parentGroupID interface{}, groupIDs ...interface{}) *Repository_UnassignParentGroup_Call { - return &Repository_UnassignParentGroup_Call{Call: _e.mock.On("UnassignParentGroup", - append([]interface{}{ctx, parentGroupID}, groupIDs...)...)} -} - -func (_c *Repository_UnassignParentGroup_Call) Run(run func(ctx context.Context, parentGroupID string, groupIDs ...string)) *Repository_UnassignParentGroup_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 []string - var variadicArgs []string - if len(args) > 2 { - variadicArgs = args[2].([]string) - } - arg2 = variadicArgs - run( - arg0, - arg1, - arg2..., - ) - }) - return _c -} - -func (_c *Repository_UnassignParentGroup_Call) Return(err error) *Repository_UnassignParentGroup_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_UnassignParentGroup_Call) RunAndReturn(run func(ctx context.Context, parentGroupID string, groupIDs ...string) error) *Repository_UnassignParentGroup_Call { - _c.Call.Return(run) - return _c -} - -// Update provides a mock function for the type Repository -func (_mock *Repository) Update(ctx context.Context, g groups.Group) (groups.Group, error) { - ret := _mock.Called(ctx, g) - - if len(ret) == 0 { - panic("no return value specified for Update") - } - - var r0 groups.Group - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, groups.Group) (groups.Group, error)); ok { - return returnFunc(ctx, g) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, groups.Group) groups.Group); ok { - r0 = returnFunc(ctx, g) - } else { - r0 = ret.Get(0).(groups.Group) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, groups.Group) error); ok { - r1 = returnFunc(ctx, g) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_Update_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Update' -type Repository_Update_Call struct { - *mock.Call -} - -// Update is a helper method to define mock.On call -// - ctx context.Context -// - g groups.Group -func (_e *Repository_Expecter) Update(ctx interface{}, g interface{}) *Repository_Update_Call { - return &Repository_Update_Call{Call: _e.mock.On("Update", ctx, g)} -} - -func (_c *Repository_Update_Call) Run(run func(ctx context.Context, g groups.Group)) *Repository_Update_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 groups.Group - if args[1] != nil { - arg1 = args[1].(groups.Group) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_Update_Call) Return(group groups.Group, err error) *Repository_Update_Call { - _c.Call.Return(group, err) - return _c -} - -func (_c *Repository_Update_Call) RunAndReturn(run func(ctx context.Context, g groups.Group) (groups.Group, error)) *Repository_Update_Call { - _c.Call.Return(run) - return _c -} - -// UpdateRole provides a mock function for the type Repository -func (_mock *Repository) UpdateRole(ctx context.Context, ro roles.Role) (roles.Role, error) { - ret := _mock.Called(ctx, ro) - - if len(ret) == 0 { - panic("no return value specified for UpdateRole") - } - - var r0 roles.Role - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role) (roles.Role, error)); ok { - return returnFunc(ctx, ro) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role) roles.Role); ok { - r0 = returnFunc(ctx, ro) - } else { - r0 = ret.Get(0).(roles.Role) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, roles.Role) error); ok { - r1 = returnFunc(ctx, ro) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_UpdateRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateRole' -type Repository_UpdateRole_Call struct { - *mock.Call -} - -// UpdateRole is a helper method to define mock.On call -// - ctx context.Context -// - ro roles.Role -func (_e *Repository_Expecter) UpdateRole(ctx interface{}, ro interface{}) *Repository_UpdateRole_Call { - return &Repository_UpdateRole_Call{Call: _e.mock.On("UpdateRole", ctx, ro)} -} - -func (_c *Repository_UpdateRole_Call) Run(run func(ctx context.Context, ro roles.Role)) *Repository_UpdateRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_UpdateRole_Call) Return(role roles.Role, err error) *Repository_UpdateRole_Call { - _c.Call.Return(role, err) - return _c -} - -func (_c *Repository_UpdateRole_Call) RunAndReturn(run func(ctx context.Context, ro roles.Role) (roles.Role, error)) *Repository_UpdateRole_Call { - _c.Call.Return(run) - return _c -} - -// UpdateTags provides a mock function for the type Repository -func (_mock *Repository) UpdateTags(ctx context.Context, g groups.Group) (groups.Group, error) { - ret := _mock.Called(ctx, g) - - if len(ret) == 0 { - panic("no return value specified for UpdateTags") - } - - var r0 groups.Group - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, groups.Group) (groups.Group, error)); ok { - return returnFunc(ctx, g) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, groups.Group) groups.Group); ok { - r0 = returnFunc(ctx, g) - } else { - r0 = ret.Get(0).(groups.Group) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, groups.Group) error); ok { - r1 = returnFunc(ctx, g) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_UpdateTags_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateTags' -type Repository_UpdateTags_Call struct { - *mock.Call -} - -// UpdateTags is a helper method to define mock.On call -// - ctx context.Context -// - g groups.Group -func (_e *Repository_Expecter) UpdateTags(ctx interface{}, g interface{}) *Repository_UpdateTags_Call { - return &Repository_UpdateTags_Call{Call: _e.mock.On("UpdateTags", ctx, g)} -} - -func (_c *Repository_UpdateTags_Call) Run(run func(ctx context.Context, g groups.Group)) *Repository_UpdateTags_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 groups.Group - if args[1] != nil { - arg1 = args[1].(groups.Group) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_UpdateTags_Call) Return(group groups.Group, err error) *Repository_UpdateTags_Call { - _c.Call.Return(group, err) - return _c -} - -func (_c *Repository_UpdateTags_Call) RunAndReturn(run func(ctx context.Context, g groups.Group) (groups.Group, error)) *Repository_UpdateTags_Call { - _c.Call.Return(run) - return _c -} diff --git a/groups/mocks/service.go b/groups/mocks/service.go deleted file mode 100644 index f093bf7cc..000000000 --- a/groups/mocks/service.go +++ /dev/null @@ -1,2686 +0,0 @@ -// Copyright (c) Abstract Machines - -// SPDX-License-Identifier: Apache-2.0 - -// Code generated by mockery; DO NOT EDIT. -// github.com/vektra/mockery -// template: testify - -package mocks - -import ( - "context" - - "github.com/absmach/magistrala/groups" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/roles" - mock "github.com/stretchr/testify/mock" -) - -// NewService creates a new instance of Service. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. -// The first argument is typically a *testing.T value. -func NewService(t interface { - mock.TestingT - Cleanup(func()) -}) *Service { - mock := &Service{} - mock.Mock.Test(t) - - t.Cleanup(func() { mock.AssertExpectations(t) }) - - return mock -} - -// Service is an autogenerated mock type for the Service type -type Service struct { - mock.Mock -} - -type Service_Expecter struct { - mock *mock.Mock -} - -func (_m *Service) EXPECT() *Service_Expecter { - return &Service_Expecter{mock: &_m.Mock} -} - -// AddChildrenGroups provides a mock function for the type Service -func (_mock *Service) AddChildrenGroups(ctx context.Context, session authn.Session, id string, childrenGroupIDs []string) error { - ret := _mock.Called(ctx, session, id, childrenGroupIDs) - - if len(ret) == 0 { - panic("no return value specified for AddChildrenGroups") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, []string) error); ok { - r0 = returnFunc(ctx, session, id, childrenGroupIDs) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_AddChildrenGroups_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'AddChildrenGroups' -type Service_AddChildrenGroups_Call struct { - *mock.Call -} - -// AddChildrenGroups is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - id string -// - childrenGroupIDs []string -func (_e *Service_Expecter) AddChildrenGroups(ctx interface{}, session interface{}, id interface{}, childrenGroupIDs interface{}) *Service_AddChildrenGroups_Call { - return &Service_AddChildrenGroups_Call{Call: _e.mock.On("AddChildrenGroups", ctx, session, id, childrenGroupIDs)} -} - -func (_c *Service_AddChildrenGroups_Call) Run(run func(ctx context.Context, session authn.Session, id string, childrenGroupIDs []string)) *Service_AddChildrenGroups_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 []string - if args[3] != nil { - arg3 = args[3].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_AddChildrenGroups_Call) Return(err error) *Service_AddChildrenGroups_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_AddChildrenGroups_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, id string, childrenGroupIDs []string) error) *Service_AddChildrenGroups_Call { - _c.Call.Return(run) - return _c -} - -// AddParentGroup provides a mock function for the type Service -func (_mock *Service) AddParentGroup(ctx context.Context, session authn.Session, id string, parentID string) error { - ret := _mock.Called(ctx, session, id, parentID) - - if len(ret) == 0 { - panic("no return value specified for AddParentGroup") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) error); ok { - r0 = returnFunc(ctx, session, id, parentID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_AddParentGroup_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'AddParentGroup' -type Service_AddParentGroup_Call struct { - *mock.Call -} - -// AddParentGroup is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - id string -// - parentID string -func (_e *Service_Expecter) AddParentGroup(ctx interface{}, session interface{}, id interface{}, parentID interface{}) *Service_AddParentGroup_Call { - return &Service_AddParentGroup_Call{Call: _e.mock.On("AddParentGroup", ctx, session, id, parentID)} -} - -func (_c *Service_AddParentGroup_Call) Run(run func(ctx context.Context, session authn.Session, id string, parentID string)) *Service_AddParentGroup_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_AddParentGroup_Call) Return(err error) *Service_AddParentGroup_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_AddParentGroup_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, id string, parentID string) error) *Service_AddParentGroup_Call { - _c.Call.Return(run) - return _c -} - -// AddRole provides a mock function for the type Service -func (_mock *Service) AddRole(ctx context.Context, session authn.Session, entityID string, roleName string, optionalActions []string, optionalMembers []string) (roles.RoleProvision, error) { - ret := _mock.Called(ctx, session, entityID, roleName, optionalActions, optionalMembers) - - if len(ret) == 0 { - panic("no return value specified for AddRole") - } - - var r0 roles.RoleProvision - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string, []string) (roles.RoleProvision, error)); ok { - return returnFunc(ctx, session, entityID, roleName, optionalActions, optionalMembers) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string, []string) roles.RoleProvision); ok { - r0 = returnFunc(ctx, session, entityID, roleName, optionalActions, optionalMembers) - } else { - r0 = ret.Get(0).(roles.RoleProvision) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, []string, []string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleName, optionalActions, optionalMembers) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_AddRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'AddRole' -type Service_AddRole_Call struct { - *mock.Call -} - -// AddRole is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleName string -// - optionalActions []string -// - optionalMembers []string -func (_e *Service_Expecter) AddRole(ctx interface{}, session interface{}, entityID interface{}, roleName interface{}, optionalActions interface{}, optionalMembers interface{}) *Service_AddRole_Call { - return &Service_AddRole_Call{Call: _e.mock.On("AddRole", ctx, session, entityID, roleName, optionalActions, optionalMembers)} -} - -func (_c *Service_AddRole_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleName string, optionalActions []string, optionalMembers []string)) *Service_AddRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - var arg5 []string - if args[5] != nil { - arg5 = args[5].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - arg5, - ) - }) - return _c -} - -func (_c *Service_AddRole_Call) Return(roleProvision roles.RoleProvision, err error) *Service_AddRole_Call { - _c.Call.Return(roleProvision, err) - return _c -} - -func (_c *Service_AddRole_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleName string, optionalActions []string, optionalMembers []string) (roles.RoleProvision, error)) *Service_AddRole_Call { - _c.Call.Return(run) - return _c -} - -// CreateGroup provides a mock function for the type Service -func (_mock *Service) CreateGroup(ctx context.Context, session authn.Session, g groups.Group) (groups.Group, []roles.RoleProvision, error) { - ret := _mock.Called(ctx, session, g) - - if len(ret) == 0 { - panic("no return value specified for CreateGroup") - } - - var r0 groups.Group - var r1 []roles.RoleProvision - var r2 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, groups.Group) (groups.Group, []roles.RoleProvision, error)); ok { - return returnFunc(ctx, session, g) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, groups.Group) groups.Group); ok { - r0 = returnFunc(ctx, session, g) - } else { - r0 = ret.Get(0).(groups.Group) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, groups.Group) []roles.RoleProvision); ok { - r1 = returnFunc(ctx, session, g) - } else { - if ret.Get(1) != nil { - r1 = ret.Get(1).([]roles.RoleProvision) - } - } - if returnFunc, ok := ret.Get(2).(func(context.Context, authn.Session, groups.Group) error); ok { - r2 = returnFunc(ctx, session, g) - } else { - r2 = ret.Error(2) - } - return r0, r1, r2 -} - -// Service_CreateGroup_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'CreateGroup' -type Service_CreateGroup_Call struct { - *mock.Call -} - -// CreateGroup is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - g groups.Group -func (_e *Service_Expecter) CreateGroup(ctx interface{}, session interface{}, g interface{}) *Service_CreateGroup_Call { - return &Service_CreateGroup_Call{Call: _e.mock.On("CreateGroup", ctx, session, g)} -} - -func (_c *Service_CreateGroup_Call) Run(run func(ctx context.Context, session authn.Session, g groups.Group)) *Service_CreateGroup_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 groups.Group - if args[2] != nil { - arg2 = args[2].(groups.Group) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_CreateGroup_Call) Return(group groups.Group, roleProvisions []roles.RoleProvision, err error) *Service_CreateGroup_Call { - _c.Call.Return(group, roleProvisions, err) - return _c -} - -func (_c *Service_CreateGroup_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, g groups.Group) (groups.Group, []roles.RoleProvision, error)) *Service_CreateGroup_Call { - _c.Call.Return(run) - return _c -} - -// DeleteGroup provides a mock function for the type Service -func (_mock *Service) DeleteGroup(ctx context.Context, session authn.Session, id string) error { - ret := _mock.Called(ctx, session, id) - - if len(ret) == 0 { - panic("no return value specified for DeleteGroup") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) error); ok { - r0 = returnFunc(ctx, session, id) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_DeleteGroup_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'DeleteGroup' -type Service_DeleteGroup_Call struct { - *mock.Call -} - -// DeleteGroup is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - id string -func (_e *Service_Expecter) DeleteGroup(ctx interface{}, session interface{}, id interface{}) *Service_DeleteGroup_Call { - return &Service_DeleteGroup_Call{Call: _e.mock.On("DeleteGroup", ctx, session, id)} -} - -func (_c *Service_DeleteGroup_Call) Run(run func(ctx context.Context, session authn.Session, id string)) *Service_DeleteGroup_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_DeleteGroup_Call) Return(err error) *Service_DeleteGroup_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_DeleteGroup_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, id string) error) *Service_DeleteGroup_Call { - _c.Call.Return(run) - return _c -} - -// DisableGroup provides a mock function for the type Service -func (_mock *Service) DisableGroup(ctx context.Context, session authn.Session, id string) (groups.Group, error) { - ret := _mock.Called(ctx, session, id) - - if len(ret) == 0 { - panic("no return value specified for DisableGroup") - } - - var r0 groups.Group - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) (groups.Group, error)); ok { - return returnFunc(ctx, session, id) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) groups.Group); ok { - r0 = returnFunc(ctx, session, id) - } else { - r0 = ret.Get(0).(groups.Group) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string) error); ok { - r1 = returnFunc(ctx, session, id) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_DisableGroup_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'DisableGroup' -type Service_DisableGroup_Call struct { - *mock.Call -} - -// DisableGroup is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - id string -func (_e *Service_Expecter) DisableGroup(ctx interface{}, session interface{}, id interface{}) *Service_DisableGroup_Call { - return &Service_DisableGroup_Call{Call: _e.mock.On("DisableGroup", ctx, session, id)} -} - -func (_c *Service_DisableGroup_Call) Run(run func(ctx context.Context, session authn.Session, id string)) *Service_DisableGroup_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_DisableGroup_Call) Return(group groups.Group, err error) *Service_DisableGroup_Call { - _c.Call.Return(group, err) - return _c -} - -func (_c *Service_DisableGroup_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, id string) (groups.Group, error)) *Service_DisableGroup_Call { - _c.Call.Return(run) - return _c -} - -// EnableGroup provides a mock function for the type Service -func (_mock *Service) EnableGroup(ctx context.Context, session authn.Session, id string) (groups.Group, error) { - ret := _mock.Called(ctx, session, id) - - if len(ret) == 0 { - panic("no return value specified for EnableGroup") - } - - var r0 groups.Group - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) (groups.Group, error)); ok { - return returnFunc(ctx, session, id) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) groups.Group); ok { - r0 = returnFunc(ctx, session, id) - } else { - r0 = ret.Get(0).(groups.Group) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string) error); ok { - r1 = returnFunc(ctx, session, id) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_EnableGroup_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'EnableGroup' -type Service_EnableGroup_Call struct { - *mock.Call -} - -// EnableGroup is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - id string -func (_e *Service_Expecter) EnableGroup(ctx interface{}, session interface{}, id interface{}) *Service_EnableGroup_Call { - return &Service_EnableGroup_Call{Call: _e.mock.On("EnableGroup", ctx, session, id)} -} - -func (_c *Service_EnableGroup_Call) Run(run func(ctx context.Context, session authn.Session, id string)) *Service_EnableGroup_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_EnableGroup_Call) Return(group groups.Group, err error) *Service_EnableGroup_Call { - _c.Call.Return(group, err) - return _c -} - -func (_c *Service_EnableGroup_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, id string) (groups.Group, error)) *Service_EnableGroup_Call { - _c.Call.Return(run) - return _c -} - -// ListAvailableActions provides a mock function for the type Service -func (_mock *Service) ListAvailableActions(ctx context.Context, session authn.Session) ([]string, error) { - ret := _mock.Called(ctx, session) - - if len(ret) == 0 { - panic("no return value specified for ListAvailableActions") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session) ([]string, error)); ok { - return returnFunc(ctx, session) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session) []string); ok { - r0 = returnFunc(ctx, session) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session) error); ok { - r1 = returnFunc(ctx, session) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_ListAvailableActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ListAvailableActions' -type Service_ListAvailableActions_Call struct { - *mock.Call -} - -// ListAvailableActions is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -func (_e *Service_Expecter) ListAvailableActions(ctx interface{}, session interface{}) *Service_ListAvailableActions_Call { - return &Service_ListAvailableActions_Call{Call: _e.mock.On("ListAvailableActions", ctx, session)} -} - -func (_c *Service_ListAvailableActions_Call) Run(run func(ctx context.Context, session authn.Session)) *Service_ListAvailableActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Service_ListAvailableActions_Call) Return(strings []string, err error) *Service_ListAvailableActions_Call { - _c.Call.Return(strings, err) - return _c -} - -func (_c *Service_ListAvailableActions_Call) RunAndReturn(run func(ctx context.Context, session authn.Session) ([]string, error)) *Service_ListAvailableActions_Call { - _c.Call.Return(run) - return _c -} - -// ListChildrenGroups provides a mock function for the type Service -func (_mock *Service) ListChildrenGroups(ctx context.Context, session authn.Session, id string, startLevel int64, endLevel int64, pm groups.PageMeta) (groups.Page, error) { - ret := _mock.Called(ctx, session, id, startLevel, endLevel, pm) - - if len(ret) == 0 { - panic("no return value specified for ListChildrenGroups") - } - - var r0 groups.Page - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, int64, int64, groups.PageMeta) (groups.Page, error)); ok { - return returnFunc(ctx, session, id, startLevel, endLevel, pm) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, int64, int64, groups.PageMeta) groups.Page); ok { - r0 = returnFunc(ctx, session, id, startLevel, endLevel, pm) - } else { - r0 = ret.Get(0).(groups.Page) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, int64, int64, groups.PageMeta) error); ok { - r1 = returnFunc(ctx, session, id, startLevel, endLevel, pm) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_ListChildrenGroups_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ListChildrenGroups' -type Service_ListChildrenGroups_Call struct { - *mock.Call -} - -// ListChildrenGroups is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - id string -// - startLevel int64 -// - endLevel int64 -// - pm groups.PageMeta -func (_e *Service_Expecter) ListChildrenGroups(ctx interface{}, session interface{}, id interface{}, startLevel interface{}, endLevel interface{}, pm interface{}) *Service_ListChildrenGroups_Call { - return &Service_ListChildrenGroups_Call{Call: _e.mock.On("ListChildrenGroups", ctx, session, id, startLevel, endLevel, pm)} -} - -func (_c *Service_ListChildrenGroups_Call) Run(run func(ctx context.Context, session authn.Session, id string, startLevel int64, endLevel int64, pm groups.PageMeta)) *Service_ListChildrenGroups_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 int64 - if args[3] != nil { - arg3 = args[3].(int64) - } - var arg4 int64 - if args[4] != nil { - arg4 = args[4].(int64) - } - var arg5 groups.PageMeta - if args[5] != nil { - arg5 = args[5].(groups.PageMeta) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - arg5, - ) - }) - return _c -} - -func (_c *Service_ListChildrenGroups_Call) Return(page groups.Page, err error) *Service_ListChildrenGroups_Call { - _c.Call.Return(page, err) - return _c -} - -func (_c *Service_ListChildrenGroups_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, id string, startLevel int64, endLevel int64, pm groups.PageMeta) (groups.Page, error)) *Service_ListChildrenGroups_Call { - _c.Call.Return(run) - return _c -} - -// ListEntityMembers provides a mock function for the type Service -func (_mock *Service) ListEntityMembers(ctx context.Context, session authn.Session, entityID string, pq roles.MembersRolePageQuery) (roles.MembersRolePage, error) { - ret := _mock.Called(ctx, session, entityID, pq) - - if len(ret) == 0 { - panic("no return value specified for ListEntityMembers") - } - - var r0 roles.MembersRolePage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, roles.MembersRolePageQuery) (roles.MembersRolePage, error)); ok { - return returnFunc(ctx, session, entityID, pq) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, roles.MembersRolePageQuery) roles.MembersRolePage); ok { - r0 = returnFunc(ctx, session, entityID, pq) - } else { - r0 = ret.Get(0).(roles.MembersRolePage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, roles.MembersRolePageQuery) error); ok { - r1 = returnFunc(ctx, session, entityID, pq) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_ListEntityMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ListEntityMembers' -type Service_ListEntityMembers_Call struct { - *mock.Call -} - -// ListEntityMembers is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - pq roles.MembersRolePageQuery -func (_e *Service_Expecter) ListEntityMembers(ctx interface{}, session interface{}, entityID interface{}, pq interface{}) *Service_ListEntityMembers_Call { - return &Service_ListEntityMembers_Call{Call: _e.mock.On("ListEntityMembers", ctx, session, entityID, pq)} -} - -func (_c *Service_ListEntityMembers_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, pq roles.MembersRolePageQuery)) *Service_ListEntityMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 roles.MembersRolePageQuery - if args[3] != nil { - arg3 = args[3].(roles.MembersRolePageQuery) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_ListEntityMembers_Call) Return(membersRolePage roles.MembersRolePage, err error) *Service_ListEntityMembers_Call { - _c.Call.Return(membersRolePage, err) - return _c -} - -func (_c *Service_ListEntityMembers_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, pq roles.MembersRolePageQuery) (roles.MembersRolePage, error)) *Service_ListEntityMembers_Call { - _c.Call.Return(run) - return _c -} - -// ListGroups provides a mock function for the type Service -func (_mock *Service) ListGroups(ctx context.Context, session authn.Session, pm groups.PageMeta) (groups.Page, error) { - ret := _mock.Called(ctx, session, pm) - - if len(ret) == 0 { - panic("no return value specified for ListGroups") - } - - var r0 groups.Page - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, groups.PageMeta) (groups.Page, error)); ok { - return returnFunc(ctx, session, pm) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, groups.PageMeta) groups.Page); ok { - r0 = returnFunc(ctx, session, pm) - } else { - r0 = ret.Get(0).(groups.Page) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, groups.PageMeta) error); ok { - r1 = returnFunc(ctx, session, pm) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_ListGroups_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ListGroups' -type Service_ListGroups_Call struct { - *mock.Call -} - -// ListGroups is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - pm groups.PageMeta -func (_e *Service_Expecter) ListGroups(ctx interface{}, session interface{}, pm interface{}) *Service_ListGroups_Call { - return &Service_ListGroups_Call{Call: _e.mock.On("ListGroups", ctx, session, pm)} -} - -func (_c *Service_ListGroups_Call) Run(run func(ctx context.Context, session authn.Session, pm groups.PageMeta)) *Service_ListGroups_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 groups.PageMeta - if args[2] != nil { - arg2 = args[2].(groups.PageMeta) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_ListGroups_Call) Return(page groups.Page, err error) *Service_ListGroups_Call { - _c.Call.Return(page, err) - return _c -} - -func (_c *Service_ListGroups_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, pm groups.PageMeta) (groups.Page, error)) *Service_ListGroups_Call { - _c.Call.Return(run) - return _c -} - -// ListUserGroups provides a mock function for the type Service -func (_mock *Service) ListUserGroups(ctx context.Context, session authn.Session, userID string, pm groups.PageMeta) (groups.Page, error) { - ret := _mock.Called(ctx, session, userID, pm) - - if len(ret) == 0 { - panic("no return value specified for ListUserGroups") - } - - var r0 groups.Page - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, groups.PageMeta) (groups.Page, error)); ok { - return returnFunc(ctx, session, userID, pm) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, groups.PageMeta) groups.Page); ok { - r0 = returnFunc(ctx, session, userID, pm) - } else { - r0 = ret.Get(0).(groups.Page) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, groups.PageMeta) error); ok { - r1 = returnFunc(ctx, session, userID, pm) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_ListUserGroups_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ListUserGroups' -type Service_ListUserGroups_Call struct { - *mock.Call -} - -// ListUserGroups is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - userID string -// - pm groups.PageMeta -func (_e *Service_Expecter) ListUserGroups(ctx interface{}, session interface{}, userID interface{}, pm interface{}) *Service_ListUserGroups_Call { - return &Service_ListUserGroups_Call{Call: _e.mock.On("ListUserGroups", ctx, session, userID, pm)} -} - -func (_c *Service_ListUserGroups_Call) Run(run func(ctx context.Context, session authn.Session, userID string, pm groups.PageMeta)) *Service_ListUserGroups_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 groups.PageMeta - if args[3] != nil { - arg3 = args[3].(groups.PageMeta) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_ListUserGroups_Call) Return(page groups.Page, err error) *Service_ListUserGroups_Call { - _c.Call.Return(page, err) - return _c -} - -func (_c *Service_ListUserGroups_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, userID string, pm groups.PageMeta) (groups.Page, error)) *Service_ListUserGroups_Call { - _c.Call.Return(run) - return _c -} - -// RemoveAllChildrenGroups provides a mock function for the type Service -func (_mock *Service) RemoveAllChildrenGroups(ctx context.Context, session authn.Session, id string) error { - ret := _mock.Called(ctx, session, id) - - if len(ret) == 0 { - panic("no return value specified for RemoveAllChildrenGroups") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) error); ok { - r0 = returnFunc(ctx, session, id) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RemoveAllChildrenGroups_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveAllChildrenGroups' -type Service_RemoveAllChildrenGroups_Call struct { - *mock.Call -} - -// RemoveAllChildrenGroups is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - id string -func (_e *Service_Expecter) RemoveAllChildrenGroups(ctx interface{}, session interface{}, id interface{}) *Service_RemoveAllChildrenGroups_Call { - return &Service_RemoveAllChildrenGroups_Call{Call: _e.mock.On("RemoveAllChildrenGroups", ctx, session, id)} -} - -func (_c *Service_RemoveAllChildrenGroups_Call) Run(run func(ctx context.Context, session authn.Session, id string)) *Service_RemoveAllChildrenGroups_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_RemoveAllChildrenGroups_Call) Return(err error) *Service_RemoveAllChildrenGroups_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RemoveAllChildrenGroups_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, id string) error) *Service_RemoveAllChildrenGroups_Call { - _c.Call.Return(run) - return _c -} - -// RemoveChildrenGroups provides a mock function for the type Service -func (_mock *Service) RemoveChildrenGroups(ctx context.Context, session authn.Session, id string, childrenGroupIDs []string) error { - ret := _mock.Called(ctx, session, id, childrenGroupIDs) - - if len(ret) == 0 { - panic("no return value specified for RemoveChildrenGroups") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, []string) error); ok { - r0 = returnFunc(ctx, session, id, childrenGroupIDs) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RemoveChildrenGroups_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveChildrenGroups' -type Service_RemoveChildrenGroups_Call struct { - *mock.Call -} - -// RemoveChildrenGroups is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - id string -// - childrenGroupIDs []string -func (_e *Service_Expecter) RemoveChildrenGroups(ctx interface{}, session interface{}, id interface{}, childrenGroupIDs interface{}) *Service_RemoveChildrenGroups_Call { - return &Service_RemoveChildrenGroups_Call{Call: _e.mock.On("RemoveChildrenGroups", ctx, session, id, childrenGroupIDs)} -} - -func (_c *Service_RemoveChildrenGroups_Call) Run(run func(ctx context.Context, session authn.Session, id string, childrenGroupIDs []string)) *Service_RemoveChildrenGroups_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 []string - if args[3] != nil { - arg3 = args[3].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_RemoveChildrenGroups_Call) Return(err error) *Service_RemoveChildrenGroups_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RemoveChildrenGroups_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, id string, childrenGroupIDs []string) error) *Service_RemoveChildrenGroups_Call { - _c.Call.Return(run) - return _c -} - -// RemoveEntityMembers provides a mock function for the type Service -func (_mock *Service) RemoveEntityMembers(ctx context.Context, session authn.Session, entityID string, members []string) error { - ret := _mock.Called(ctx, session, entityID, members) - - if len(ret) == 0 { - panic("no return value specified for RemoveEntityMembers") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, []string) error); ok { - r0 = returnFunc(ctx, session, entityID, members) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RemoveEntityMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveEntityMembers' -type Service_RemoveEntityMembers_Call struct { - *mock.Call -} - -// RemoveEntityMembers is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - members []string -func (_e *Service_Expecter) RemoveEntityMembers(ctx interface{}, session interface{}, entityID interface{}, members interface{}) *Service_RemoveEntityMembers_Call { - return &Service_RemoveEntityMembers_Call{Call: _e.mock.On("RemoveEntityMembers", ctx, session, entityID, members)} -} - -func (_c *Service_RemoveEntityMembers_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, members []string)) *Service_RemoveEntityMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 []string - if args[3] != nil { - arg3 = args[3].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_RemoveEntityMembers_Call) Return(err error) *Service_RemoveEntityMembers_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RemoveEntityMembers_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, members []string) error) *Service_RemoveEntityMembers_Call { - _c.Call.Return(run) - return _c -} - -// RemoveMemberFromAllRoles provides a mock function for the type Service -func (_mock *Service) RemoveMemberFromAllRoles(ctx context.Context, session authn.Session, memberID string) error { - ret := _mock.Called(ctx, session, memberID) - - if len(ret) == 0 { - panic("no return value specified for RemoveMemberFromAllRoles") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) error); ok { - r0 = returnFunc(ctx, session, memberID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RemoveMemberFromAllRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveMemberFromAllRoles' -type Service_RemoveMemberFromAllRoles_Call struct { - *mock.Call -} - -// RemoveMemberFromAllRoles is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - memberID string -func (_e *Service_Expecter) RemoveMemberFromAllRoles(ctx interface{}, session interface{}, memberID interface{}) *Service_RemoveMemberFromAllRoles_Call { - return &Service_RemoveMemberFromAllRoles_Call{Call: _e.mock.On("RemoveMemberFromAllRoles", ctx, session, memberID)} -} - -func (_c *Service_RemoveMemberFromAllRoles_Call) Run(run func(ctx context.Context, session authn.Session, memberID string)) *Service_RemoveMemberFromAllRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_RemoveMemberFromAllRoles_Call) Return(err error) *Service_RemoveMemberFromAllRoles_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RemoveMemberFromAllRoles_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, memberID string) error) *Service_RemoveMemberFromAllRoles_Call { - _c.Call.Return(run) - return _c -} - -// RemoveParentGroup provides a mock function for the type Service -func (_mock *Service) RemoveParentGroup(ctx context.Context, session authn.Session, id string) error { - ret := _mock.Called(ctx, session, id) - - if len(ret) == 0 { - panic("no return value specified for RemoveParentGroup") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) error); ok { - r0 = returnFunc(ctx, session, id) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RemoveParentGroup_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveParentGroup' -type Service_RemoveParentGroup_Call struct { - *mock.Call -} - -// RemoveParentGroup is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - id string -func (_e *Service_Expecter) RemoveParentGroup(ctx interface{}, session interface{}, id interface{}) *Service_RemoveParentGroup_Call { - return &Service_RemoveParentGroup_Call{Call: _e.mock.On("RemoveParentGroup", ctx, session, id)} -} - -func (_c *Service_RemoveParentGroup_Call) Run(run func(ctx context.Context, session authn.Session, id string)) *Service_RemoveParentGroup_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_RemoveParentGroup_Call) Return(err error) *Service_RemoveParentGroup_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RemoveParentGroup_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, id string) error) *Service_RemoveParentGroup_Call { - _c.Call.Return(run) - return _c -} - -// RemoveRole provides a mock function for the type Service -func (_mock *Service) RemoveRole(ctx context.Context, session authn.Session, entityID string, roleID string) error { - ret := _mock.Called(ctx, session, entityID, roleID) - - if len(ret) == 0 { - panic("no return value specified for RemoveRole") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) error); ok { - r0 = returnFunc(ctx, session, entityID, roleID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RemoveRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveRole' -type Service_RemoveRole_Call struct { - *mock.Call -} - -// RemoveRole is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -func (_e *Service_Expecter) RemoveRole(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}) *Service_RemoveRole_Call { - return &Service_RemoveRole_Call{Call: _e.mock.On("RemoveRole", ctx, session, entityID, roleID)} -} - -func (_c *Service_RemoveRole_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string)) *Service_RemoveRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_RemoveRole_Call) Return(err error) *Service_RemoveRole_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RemoveRole_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string) error) *Service_RemoveRole_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveAllRoles provides a mock function for the type Service -func (_mock *Service) RetrieveAllRoles(ctx context.Context, session authn.Session, entityID string, limit uint64, offset uint64) (roles.RolePage, error) { - ret := _mock.Called(ctx, session, entityID, limit, offset) - - if len(ret) == 0 { - panic("no return value specified for RetrieveAllRoles") - } - - var r0 roles.RolePage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, uint64, uint64) (roles.RolePage, error)); ok { - return returnFunc(ctx, session, entityID, limit, offset) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, uint64, uint64) roles.RolePage); ok { - r0 = returnFunc(ctx, session, entityID, limit, offset) - } else { - r0 = ret.Get(0).(roles.RolePage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, uint64, uint64) error); ok { - r1 = returnFunc(ctx, session, entityID, limit, offset) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RetrieveAllRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveAllRoles' -type Service_RetrieveAllRoles_Call struct { - *mock.Call -} - -// RetrieveAllRoles is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - limit uint64 -// - offset uint64 -func (_e *Service_Expecter) RetrieveAllRoles(ctx interface{}, session interface{}, entityID interface{}, limit interface{}, offset interface{}) *Service_RetrieveAllRoles_Call { - return &Service_RetrieveAllRoles_Call{Call: _e.mock.On("RetrieveAllRoles", ctx, session, entityID, limit, offset)} -} - -func (_c *Service_RetrieveAllRoles_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, limit uint64, offset uint64)) *Service_RetrieveAllRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 uint64 - if args[3] != nil { - arg3 = args[3].(uint64) - } - var arg4 uint64 - if args[4] != nil { - arg4 = args[4].(uint64) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_RetrieveAllRoles_Call) Return(rolePage roles.RolePage, err error) *Service_RetrieveAllRoles_Call { - _c.Call.Return(rolePage, err) - return _c -} - -func (_c *Service_RetrieveAllRoles_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, limit uint64, offset uint64) (roles.RolePage, error)) *Service_RetrieveAllRoles_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveGroupHierarchy provides a mock function for the type Service -func (_mock *Service) RetrieveGroupHierarchy(ctx context.Context, session authn.Session, id string, hm groups.HierarchyPageMeta) (groups.HierarchyPage, error) { - ret := _mock.Called(ctx, session, id, hm) - - if len(ret) == 0 { - panic("no return value specified for RetrieveGroupHierarchy") - } - - var r0 groups.HierarchyPage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, groups.HierarchyPageMeta) (groups.HierarchyPage, error)); ok { - return returnFunc(ctx, session, id, hm) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, groups.HierarchyPageMeta) groups.HierarchyPage); ok { - r0 = returnFunc(ctx, session, id, hm) - } else { - r0 = ret.Get(0).(groups.HierarchyPage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, groups.HierarchyPageMeta) error); ok { - r1 = returnFunc(ctx, session, id, hm) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RetrieveGroupHierarchy_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveGroupHierarchy' -type Service_RetrieveGroupHierarchy_Call struct { - *mock.Call -} - -// RetrieveGroupHierarchy is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - id string -// - hm groups.HierarchyPageMeta -func (_e *Service_Expecter) RetrieveGroupHierarchy(ctx interface{}, session interface{}, id interface{}, hm interface{}) *Service_RetrieveGroupHierarchy_Call { - return &Service_RetrieveGroupHierarchy_Call{Call: _e.mock.On("RetrieveGroupHierarchy", ctx, session, id, hm)} -} - -func (_c *Service_RetrieveGroupHierarchy_Call) Run(run func(ctx context.Context, session authn.Session, id string, hm groups.HierarchyPageMeta)) *Service_RetrieveGroupHierarchy_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 groups.HierarchyPageMeta - if args[3] != nil { - arg3 = args[3].(groups.HierarchyPageMeta) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_RetrieveGroupHierarchy_Call) Return(hierarchyPage groups.HierarchyPage, err error) *Service_RetrieveGroupHierarchy_Call { - _c.Call.Return(hierarchyPage, err) - return _c -} - -func (_c *Service_RetrieveGroupHierarchy_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, id string, hm groups.HierarchyPageMeta) (groups.HierarchyPage, error)) *Service_RetrieveGroupHierarchy_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveRole provides a mock function for the type Service -func (_mock *Service) RetrieveRole(ctx context.Context, session authn.Session, entityID string, roleID string) (roles.Role, error) { - ret := _mock.Called(ctx, session, entityID, roleID) - - if len(ret) == 0 { - panic("no return value specified for RetrieveRole") - } - - var r0 roles.Role - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) (roles.Role, error)); ok { - return returnFunc(ctx, session, entityID, roleID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) roles.Role); ok { - r0 = returnFunc(ctx, session, entityID, roleID) - } else { - r0 = ret.Get(0).(roles.Role) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RetrieveRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveRole' -type Service_RetrieveRole_Call struct { - *mock.Call -} - -// RetrieveRole is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -func (_e *Service_Expecter) RetrieveRole(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}) *Service_RetrieveRole_Call { - return &Service_RetrieveRole_Call{Call: _e.mock.On("RetrieveRole", ctx, session, entityID, roleID)} -} - -func (_c *Service_RetrieveRole_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string)) *Service_RetrieveRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_RetrieveRole_Call) Return(role roles.Role, err error) *Service_RetrieveRole_Call { - _c.Call.Return(role, err) - return _c -} - -func (_c *Service_RetrieveRole_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string) (roles.Role, error)) *Service_RetrieveRole_Call { - _c.Call.Return(run) - return _c -} - -// RoleAddActions provides a mock function for the type Service -func (_mock *Service) RoleAddActions(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string) ([]string, error) { - ret := _mock.Called(ctx, session, entityID, roleID, actions) - - if len(ret) == 0 { - panic("no return value specified for RoleAddActions") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) ([]string, error)); ok { - return returnFunc(ctx, session, entityID, roleID, actions) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) []string); ok { - r0 = returnFunc(ctx, session, entityID, roleID, actions) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, []string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID, actions) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RoleAddActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleAddActions' -type Service_RoleAddActions_Call struct { - *mock.Call -} - -// RoleAddActions is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - actions []string -func (_e *Service_Expecter) RoleAddActions(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, actions interface{}) *Service_RoleAddActions_Call { - return &Service_RoleAddActions_Call{Call: _e.mock.On("RoleAddActions", ctx, session, entityID, roleID, actions)} -} - -func (_c *Service_RoleAddActions_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string)) *Service_RoleAddActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_RoleAddActions_Call) Return(ops []string, err error) *Service_RoleAddActions_Call { - _c.Call.Return(ops, err) - return _c -} - -func (_c *Service_RoleAddActions_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string) ([]string, error)) *Service_RoleAddActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleAddMembers provides a mock function for the type Service -func (_mock *Service) RoleAddMembers(ctx context.Context, session authn.Session, entityID string, roleID string, members []string) ([]string, error) { - ret := _mock.Called(ctx, session, entityID, roleID, members) - - if len(ret) == 0 { - panic("no return value specified for RoleAddMembers") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) ([]string, error)); ok { - return returnFunc(ctx, session, entityID, roleID, members) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) []string); ok { - r0 = returnFunc(ctx, session, entityID, roleID, members) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, []string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID, members) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RoleAddMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleAddMembers' -type Service_RoleAddMembers_Call struct { - *mock.Call -} - -// RoleAddMembers is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - members []string -func (_e *Service_Expecter) RoleAddMembers(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, members interface{}) *Service_RoleAddMembers_Call { - return &Service_RoleAddMembers_Call{Call: _e.mock.On("RoleAddMembers", ctx, session, entityID, roleID, members)} -} - -func (_c *Service_RoleAddMembers_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, members []string)) *Service_RoleAddMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_RoleAddMembers_Call) Return(strings []string, err error) *Service_RoleAddMembers_Call { - _c.Call.Return(strings, err) - return _c -} - -func (_c *Service_RoleAddMembers_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, members []string) ([]string, error)) *Service_RoleAddMembers_Call { - _c.Call.Return(run) - return _c -} - -// RoleCheckActionsExists provides a mock function for the type Service -func (_mock *Service) RoleCheckActionsExists(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string) (bool, error) { - ret := _mock.Called(ctx, session, entityID, roleID, actions) - - if len(ret) == 0 { - panic("no return value specified for RoleCheckActionsExists") - } - - var r0 bool - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) (bool, error)); ok { - return returnFunc(ctx, session, entityID, roleID, actions) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) bool); ok { - r0 = returnFunc(ctx, session, entityID, roleID, actions) - } else { - r0 = ret.Get(0).(bool) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, []string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID, actions) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RoleCheckActionsExists_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleCheckActionsExists' -type Service_RoleCheckActionsExists_Call struct { - *mock.Call -} - -// RoleCheckActionsExists is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - actions []string -func (_e *Service_Expecter) RoleCheckActionsExists(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, actions interface{}) *Service_RoleCheckActionsExists_Call { - return &Service_RoleCheckActionsExists_Call{Call: _e.mock.On("RoleCheckActionsExists", ctx, session, entityID, roleID, actions)} -} - -func (_c *Service_RoleCheckActionsExists_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string)) *Service_RoleCheckActionsExists_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_RoleCheckActionsExists_Call) Return(b bool, err error) *Service_RoleCheckActionsExists_Call { - _c.Call.Return(b, err) - return _c -} - -func (_c *Service_RoleCheckActionsExists_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string) (bool, error)) *Service_RoleCheckActionsExists_Call { - _c.Call.Return(run) - return _c -} - -// RoleCheckMembersExists provides a mock function for the type Service -func (_mock *Service) RoleCheckMembersExists(ctx context.Context, session authn.Session, entityID string, roleID string, members []string) (bool, error) { - ret := _mock.Called(ctx, session, entityID, roleID, members) - - if len(ret) == 0 { - panic("no return value specified for RoleCheckMembersExists") - } - - var r0 bool - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) (bool, error)); ok { - return returnFunc(ctx, session, entityID, roleID, members) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) bool); ok { - r0 = returnFunc(ctx, session, entityID, roleID, members) - } else { - r0 = ret.Get(0).(bool) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, []string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID, members) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RoleCheckMembersExists_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleCheckMembersExists' -type Service_RoleCheckMembersExists_Call struct { - *mock.Call -} - -// RoleCheckMembersExists is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - members []string -func (_e *Service_Expecter) RoleCheckMembersExists(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, members interface{}) *Service_RoleCheckMembersExists_Call { - return &Service_RoleCheckMembersExists_Call{Call: _e.mock.On("RoleCheckMembersExists", ctx, session, entityID, roleID, members)} -} - -func (_c *Service_RoleCheckMembersExists_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, members []string)) *Service_RoleCheckMembersExists_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_RoleCheckMembersExists_Call) Return(b bool, err error) *Service_RoleCheckMembersExists_Call { - _c.Call.Return(b, err) - return _c -} - -func (_c *Service_RoleCheckMembersExists_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, members []string) (bool, error)) *Service_RoleCheckMembersExists_Call { - _c.Call.Return(run) - return _c -} - -// RoleListActions provides a mock function for the type Service -func (_mock *Service) RoleListActions(ctx context.Context, session authn.Session, entityID string, roleID string) ([]string, error) { - ret := _mock.Called(ctx, session, entityID, roleID) - - if len(ret) == 0 { - panic("no return value specified for RoleListActions") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) ([]string, error)); ok { - return returnFunc(ctx, session, entityID, roleID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) []string); ok { - r0 = returnFunc(ctx, session, entityID, roleID) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RoleListActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleListActions' -type Service_RoleListActions_Call struct { - *mock.Call -} - -// RoleListActions is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -func (_e *Service_Expecter) RoleListActions(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}) *Service_RoleListActions_Call { - return &Service_RoleListActions_Call{Call: _e.mock.On("RoleListActions", ctx, session, entityID, roleID)} -} - -func (_c *Service_RoleListActions_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string)) *Service_RoleListActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_RoleListActions_Call) Return(strings []string, err error) *Service_RoleListActions_Call { - _c.Call.Return(strings, err) - return _c -} - -func (_c *Service_RoleListActions_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string) ([]string, error)) *Service_RoleListActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleListMembers provides a mock function for the type Service -func (_mock *Service) RoleListMembers(ctx context.Context, session authn.Session, entityID string, roleID string, limit uint64, offset uint64) (roles.MembersPage, error) { - ret := _mock.Called(ctx, session, entityID, roleID, limit, offset) - - if len(ret) == 0 { - panic("no return value specified for RoleListMembers") - } - - var r0 roles.MembersPage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, uint64, uint64) (roles.MembersPage, error)); ok { - return returnFunc(ctx, session, entityID, roleID, limit, offset) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, uint64, uint64) roles.MembersPage); ok { - r0 = returnFunc(ctx, session, entityID, roleID, limit, offset) - } else { - r0 = ret.Get(0).(roles.MembersPage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, uint64, uint64) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID, limit, offset) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RoleListMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleListMembers' -type Service_RoleListMembers_Call struct { - *mock.Call -} - -// RoleListMembers is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - limit uint64 -// - offset uint64 -func (_e *Service_Expecter) RoleListMembers(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, limit interface{}, offset interface{}) *Service_RoleListMembers_Call { - return &Service_RoleListMembers_Call{Call: _e.mock.On("RoleListMembers", ctx, session, entityID, roleID, limit, offset)} -} - -func (_c *Service_RoleListMembers_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, limit uint64, offset uint64)) *Service_RoleListMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 uint64 - if args[4] != nil { - arg4 = args[4].(uint64) - } - var arg5 uint64 - if args[5] != nil { - arg5 = args[5].(uint64) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - arg5, - ) - }) - return _c -} - -func (_c *Service_RoleListMembers_Call) Return(membersPage roles.MembersPage, err error) *Service_RoleListMembers_Call { - _c.Call.Return(membersPage, err) - return _c -} - -func (_c *Service_RoleListMembers_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, limit uint64, offset uint64) (roles.MembersPage, error)) *Service_RoleListMembers_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveActions provides a mock function for the type Service -func (_mock *Service) RoleRemoveActions(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string) error { - ret := _mock.Called(ctx, session, entityID, roleID, actions) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveActions") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) error); ok { - r0 = returnFunc(ctx, session, entityID, roleID, actions) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RoleRemoveActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveActions' -type Service_RoleRemoveActions_Call struct { - *mock.Call -} - -// RoleRemoveActions is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - actions []string -func (_e *Service_Expecter) RoleRemoveActions(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, actions interface{}) *Service_RoleRemoveActions_Call { - return &Service_RoleRemoveActions_Call{Call: _e.mock.On("RoleRemoveActions", ctx, session, entityID, roleID, actions)} -} - -func (_c *Service_RoleRemoveActions_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string)) *Service_RoleRemoveActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_RoleRemoveActions_Call) Return(err error) *Service_RoleRemoveActions_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RoleRemoveActions_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string) error) *Service_RoleRemoveActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveAllActions provides a mock function for the type Service -func (_mock *Service) RoleRemoveAllActions(ctx context.Context, session authn.Session, entityID string, roleID string) error { - ret := _mock.Called(ctx, session, entityID, roleID) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveAllActions") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) error); ok { - r0 = returnFunc(ctx, session, entityID, roleID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RoleRemoveAllActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveAllActions' -type Service_RoleRemoveAllActions_Call struct { - *mock.Call -} - -// RoleRemoveAllActions is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -func (_e *Service_Expecter) RoleRemoveAllActions(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}) *Service_RoleRemoveAllActions_Call { - return &Service_RoleRemoveAllActions_Call{Call: _e.mock.On("RoleRemoveAllActions", ctx, session, entityID, roleID)} -} - -func (_c *Service_RoleRemoveAllActions_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string)) *Service_RoleRemoveAllActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_RoleRemoveAllActions_Call) Return(err error) *Service_RoleRemoveAllActions_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RoleRemoveAllActions_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string) error) *Service_RoleRemoveAllActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveAllMembers provides a mock function for the type Service -func (_mock *Service) RoleRemoveAllMembers(ctx context.Context, session authn.Session, entityID string, roleID string) error { - ret := _mock.Called(ctx, session, entityID, roleID) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveAllMembers") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) error); ok { - r0 = returnFunc(ctx, session, entityID, roleID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RoleRemoveAllMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveAllMembers' -type Service_RoleRemoveAllMembers_Call struct { - *mock.Call -} - -// RoleRemoveAllMembers is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -func (_e *Service_Expecter) RoleRemoveAllMembers(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}) *Service_RoleRemoveAllMembers_Call { - return &Service_RoleRemoveAllMembers_Call{Call: _e.mock.On("RoleRemoveAllMembers", ctx, session, entityID, roleID)} -} - -func (_c *Service_RoleRemoveAllMembers_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string)) *Service_RoleRemoveAllMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_RoleRemoveAllMembers_Call) Return(err error) *Service_RoleRemoveAllMembers_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RoleRemoveAllMembers_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string) error) *Service_RoleRemoveAllMembers_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveMembers provides a mock function for the type Service -func (_mock *Service) RoleRemoveMembers(ctx context.Context, session authn.Session, entityID string, roleID string, members []string) error { - ret := _mock.Called(ctx, session, entityID, roleID, members) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveMembers") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) error); ok { - r0 = returnFunc(ctx, session, entityID, roleID, members) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RoleRemoveMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveMembers' -type Service_RoleRemoveMembers_Call struct { - *mock.Call -} - -// RoleRemoveMembers is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - members []string -func (_e *Service_Expecter) RoleRemoveMembers(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, members interface{}) *Service_RoleRemoveMembers_Call { - return &Service_RoleRemoveMembers_Call{Call: _e.mock.On("RoleRemoveMembers", ctx, session, entityID, roleID, members)} -} - -func (_c *Service_RoleRemoveMembers_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, members []string)) *Service_RoleRemoveMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_RoleRemoveMembers_Call) Return(err error) *Service_RoleRemoveMembers_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RoleRemoveMembers_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, members []string) error) *Service_RoleRemoveMembers_Call { - _c.Call.Return(run) - return _c -} - -// UpdateGroup provides a mock function for the type Service -func (_mock *Service) UpdateGroup(ctx context.Context, session authn.Session, g groups.Group) (groups.Group, error) { - ret := _mock.Called(ctx, session, g) - - if len(ret) == 0 { - panic("no return value specified for UpdateGroup") - } - - var r0 groups.Group - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, groups.Group) (groups.Group, error)); ok { - return returnFunc(ctx, session, g) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, groups.Group) groups.Group); ok { - r0 = returnFunc(ctx, session, g) - } else { - r0 = ret.Get(0).(groups.Group) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, groups.Group) error); ok { - r1 = returnFunc(ctx, session, g) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_UpdateGroup_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateGroup' -type Service_UpdateGroup_Call struct { - *mock.Call -} - -// UpdateGroup is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - g groups.Group -func (_e *Service_Expecter) UpdateGroup(ctx interface{}, session interface{}, g interface{}) *Service_UpdateGroup_Call { - return &Service_UpdateGroup_Call{Call: _e.mock.On("UpdateGroup", ctx, session, g)} -} - -func (_c *Service_UpdateGroup_Call) Run(run func(ctx context.Context, session authn.Session, g groups.Group)) *Service_UpdateGroup_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 groups.Group - if args[2] != nil { - arg2 = args[2].(groups.Group) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_UpdateGroup_Call) Return(group groups.Group, err error) *Service_UpdateGroup_Call { - _c.Call.Return(group, err) - return _c -} - -func (_c *Service_UpdateGroup_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, g groups.Group) (groups.Group, error)) *Service_UpdateGroup_Call { - _c.Call.Return(run) - return _c -} - -// UpdateGroupTags provides a mock function for the type Service -func (_mock *Service) UpdateGroupTags(ctx context.Context, session authn.Session, group groups.Group) (groups.Group, error) { - ret := _mock.Called(ctx, session, group) - - if len(ret) == 0 { - panic("no return value specified for UpdateGroupTags") - } - - var r0 groups.Group - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, groups.Group) (groups.Group, error)); ok { - return returnFunc(ctx, session, group) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, groups.Group) groups.Group); ok { - r0 = returnFunc(ctx, session, group) - } else { - r0 = ret.Get(0).(groups.Group) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, groups.Group) error); ok { - r1 = returnFunc(ctx, session, group) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_UpdateGroupTags_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateGroupTags' -type Service_UpdateGroupTags_Call struct { - *mock.Call -} - -// UpdateGroupTags is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - group groups.Group -func (_e *Service_Expecter) UpdateGroupTags(ctx interface{}, session interface{}, group interface{}) *Service_UpdateGroupTags_Call { - return &Service_UpdateGroupTags_Call{Call: _e.mock.On("UpdateGroupTags", ctx, session, group)} -} - -func (_c *Service_UpdateGroupTags_Call) Run(run func(ctx context.Context, session authn.Session, group groups.Group)) *Service_UpdateGroupTags_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 groups.Group - if args[2] != nil { - arg2 = args[2].(groups.Group) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_UpdateGroupTags_Call) Return(group1 groups.Group, err error) *Service_UpdateGroupTags_Call { - _c.Call.Return(group1, err) - return _c -} - -func (_c *Service_UpdateGroupTags_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, group groups.Group) (groups.Group, error)) *Service_UpdateGroupTags_Call { - _c.Call.Return(run) - return _c -} - -// UpdateRoleName provides a mock function for the type Service -func (_mock *Service) UpdateRoleName(ctx context.Context, session authn.Session, entityID string, roleID string, newRoleName string) (roles.Role, error) { - ret := _mock.Called(ctx, session, entityID, roleID, newRoleName) - - if len(ret) == 0 { - panic("no return value specified for UpdateRoleName") - } - - var r0 roles.Role - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, string) (roles.Role, error)); ok { - return returnFunc(ctx, session, entityID, roleID, newRoleName) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, string) roles.Role); ok { - r0 = returnFunc(ctx, session, entityID, roleID, newRoleName) - } else { - r0 = ret.Get(0).(roles.Role) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID, newRoleName) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_UpdateRoleName_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateRoleName' -type Service_UpdateRoleName_Call struct { - *mock.Call -} - -// UpdateRoleName is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - newRoleName string -func (_e *Service_Expecter) UpdateRoleName(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, newRoleName interface{}) *Service_UpdateRoleName_Call { - return &Service_UpdateRoleName_Call{Call: _e.mock.On("UpdateRoleName", ctx, session, entityID, roleID, newRoleName)} -} - -func (_c *Service_UpdateRoleName_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, newRoleName string)) *Service_UpdateRoleName_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 string - if args[4] != nil { - arg4 = args[4].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_UpdateRoleName_Call) Return(role roles.Role, err error) *Service_UpdateRoleName_Call { - _c.Call.Return(role, err) - return _c -} - -func (_c *Service_UpdateRoleName_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, newRoleName string) (roles.Role, error)) *Service_UpdateRoleName_Call { - _c.Call.Return(run) - return _c -} - -// ViewGroup provides a mock function for the type Service -func (_mock *Service) ViewGroup(ctx context.Context, session authn.Session, id string, withRoles bool) (groups.Group, error) { - ret := _mock.Called(ctx, session, id, withRoles) - - if len(ret) == 0 { - panic("no return value specified for ViewGroup") - } - - var r0 groups.Group - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, bool) (groups.Group, error)); ok { - return returnFunc(ctx, session, id, withRoles) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, bool) groups.Group); ok { - r0 = returnFunc(ctx, session, id, withRoles) - } else { - r0 = ret.Get(0).(groups.Group) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, bool) error); ok { - r1 = returnFunc(ctx, session, id, withRoles) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_ViewGroup_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ViewGroup' -type Service_ViewGroup_Call struct { - *mock.Call -} - -// ViewGroup is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - id string -// - withRoles bool -func (_e *Service_Expecter) ViewGroup(ctx interface{}, session interface{}, id interface{}, withRoles interface{}) *Service_ViewGroup_Call { - return &Service_ViewGroup_Call{Call: _e.mock.On("ViewGroup", ctx, session, id, withRoles)} -} - -func (_c *Service_ViewGroup_Call) Run(run func(ctx context.Context, session authn.Session, id string, withRoles bool)) *Service_ViewGroup_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 bool - if args[3] != nil { - arg3 = args[3].(bool) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_ViewGroup_Call) Return(group groups.Group, err error) *Service_ViewGroup_Call { - _c.Call.Return(group, err) - return _c -} - -func (_c *Service_ViewGroup_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, id string, withRoles bool) (groups.Group, error)) *Service_ViewGroup_Call { - _c.Call.Return(run) - return _c -} diff --git a/groups/operations/operations.go b/groups/operations/operations.go deleted file mode 100644 index f28922582..000000000 --- a/groups/operations/operations.go +++ /dev/null @@ -1,105 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package operations - -import "github.com/absmach/magistrala/pkg/permissions" - -// Group Operations. -const ( - OpViewGroup permissions.Operation = iota - OpUpdateGroup - OpUpdateGroupTags - OpEnableGroup - OpDisableGroup - OpRetrieveGroupHierarchy - OpAddParentGroup - OpRemoveParentGroup - OpAddChildrenGroups - OpRemoveChildrenGroups - OpRemoveAllChildrenGroups - OpListChildrenGroups - OpDeleteGroup - OpGroupSetChildClient - OpGroupRemoveChildClient - OpGroupSetChildChannel - OpGroupRemoveChildChannel - OpListUserGroups -) - -func OperationDetails() map[permissions.Operation]permissions.OperationDetails { - return map[permissions.Operation]permissions.OperationDetails{ - OpViewGroup: { - Name: "view", - PermissionRequired: true, - }, - OpUpdateGroup: { - Name: "update", - PermissionRequired: true, - }, - OpUpdateGroupTags: { - Name: "update_tags", - PermissionRequired: true, - }, - OpEnableGroup: { - Name: "enable", - PermissionRequired: true, - }, - OpDisableGroup: { - Name: "disable", - PermissionRequired: true, - }, - OpRetrieveGroupHierarchy: { - Name: "retrieve_group_hierarchy", - PermissionRequired: true, - }, - OpAddParentGroup: { - Name: "add_parent_group", - PermissionRequired: true, - }, - OpRemoveParentGroup: { - Name: "remove_parent_group", - PermissionRequired: true, - }, - OpAddChildrenGroups: { - Name: "add_children_groups", - PermissionRequired: true, - }, - OpRemoveChildrenGroups: { - Name: "remove_children_groups", - PermissionRequired: true, - }, - OpRemoveAllChildrenGroups: { - Name: "remove_all_children_groups", - PermissionRequired: true, - }, - OpListChildrenGroups: { - Name: "list_children_groups", - PermissionRequired: true, - }, - OpDeleteGroup: { - Name: "delete", - PermissionRequired: true, - }, - OpGroupSetChildClient: { - Name: "set_child_client", - PermissionRequired: true, - }, - OpGroupRemoveChildClient: { - Name: "remove_child_client", - PermissionRequired: true, - }, - OpGroupSetChildChannel: { - Name: "set_child_channel", - PermissionRequired: true, - }, - OpGroupRemoveChildChannel: { - Name: "remove_child_channel", - PermissionRequired: true, - }, - OpListUserGroups: { - Name: "list_user_groups", - PermissionRequired: false, // hardcoded to superadmin - }, - } -} diff --git a/groups/page.go b/groups/page.go deleted file mode 100644 index 14dfb48c9..000000000 --- a/groups/page.go +++ /dev/null @@ -1,64 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package groups - -import ( - "strings" - "time" -) - -type Operator uint8 - -const ( - OrOp Operator = iota - AndOp -) - -type TagsQuery struct { - Elements []string - Operator Operator -} - -func ToTagsQuery(s string) TagsQuery { - switch { - case strings.Contains(s, "+"): - elements := strings.Split(s, "+") - for i := range elements { - elements[i] = strings.TrimSpace(elements[i]) - } - return TagsQuery{Elements: elements, Operator: AndOp} - case strings.Contains(s, ","): - elements := strings.Split(s, ",") - for i := range elements { - elements[i] = strings.TrimSpace(elements[i]) - } - return TagsQuery{Elements: elements, Operator: OrOp} - default: - return TagsQuery{Elements: []string{s}, Operator: OrOp} - } -} - -// PageMeta contains page metadata that helps navigation. -type PageMeta struct { - Total uint64 `json:"total"` - Offset uint64 `json:"offset"` - Limit uint64 `json:"limit"` - OnlyTotal bool `json:"only_total"` - Name string `json:"name,omitempty"` - ID string `json:"id,omitempty"` - Dir string `json:"dir,omitempty"` - Order string `json:"order,omitempty"` - Path string `json:"path,omitempty"` - DomainID string `json:"domain_id,omitempty"` - Tags TagsQuery `json:"tags,omitempty"` - Metadata Metadata `json:"metadata,omitempty"` - Status Status `json:"status,omitempty"` - RoleName string `json:"role_name,omitempty"` - RoleID string `json:"role_id,omitempty"` - Actions []string `json:"actions,omitempty"` - AccessType string `json:"access_type,omitempty"` - RootGroup bool `json:"root_group,omitempty"` - CreatedFrom time.Time `json:"created_from,omitempty"` - CreatedTo time.Time `json:"created_to,omitempty"` -} diff --git a/groups/postgres/doc.go b/groups/postgres/doc.go deleted file mode 100644 index 96fe21175..000000000 --- a/groups/postgres/doc.go +++ /dev/null @@ -1,5 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package postgres contains the database implementation of groups repository layer. -package postgres diff --git a/groups/postgres/errors.go b/groups/postgres/errors.go deleted file mode 100644 index 8a2be3d28..000000000 --- a/groups/postgres/errors.go +++ /dev/null @@ -1,26 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package postgres - -import "github.com/absmach/magistrala/pkg/errors" - -var _ errors.Mapper = (*duplicateErrors)(nil) - -var errCyclicParentGroup = errors.NewRequestError("cyclic parent, group is parent of requested group") - -type duplicateErrors struct{} - -// GetError maps constraint names to known errors. -func (d duplicateErrors) GetError(constraint string) (error, bool) { - switch constraint { - case "groups_pkey": - return errors.NewRequestError("group id already exists"), true - default: - return nil, false - } -} - -func NewDuplicateErrors() errors.Mapper { - return duplicateErrors{} -} diff --git a/groups/postgres/groups.go b/groups/postgres/groups.go deleted file mode 100644 index 960a7ea50..000000000 --- a/groups/postgres/groups.go +++ /dev/null @@ -1,1488 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package postgres - -import ( - "context" - "database/sql" - "encoding/json" - "fmt" - "strings" - "time" - - api "github.com/absmach/magistrala/api/http" - groups "github.com/absmach/magistrala/groups" - "github.com/absmach/magistrala/internal/nullable" - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - "github.com/absmach/magistrala/pkg/policies" - "github.com/absmach/magistrala/pkg/postgres" - "github.com/absmach/magistrala/pkg/roles" - rolesPostgres "github.com/absmach/magistrala/pkg/roles/repo/postgres" - "github.com/jackc/pgtype" - "github.com/jmoiron/sqlx" - "github.com/lib/pq" -) - -var _ groups.Repository = (*groupRepository)(nil) - -const ( - rolesTableNamePrefix = "groups" - entityTableName = "groups" - entityIDColumnName = "id" -) - -var ( - errParentGroupID = errors.New("parent group id is empty") - errParentGroupPath = errors.New("parent group path is empty") - errParentSuffix = errors.New("parent group path doesn't have parent id suffix") -) - -type groupRepository struct { - db postgres.Database - eh errors.Handler - rolesPostgres.Repository -} - -// New instantiates a PostgreSQL implementation of group -// repository. -func New(db postgres.Database) groups.Repository { - roleRepo := rolesPostgres.NewRepository(db, policies.GroupType, rolesTableNamePrefix, entityTableName, entityIDColumnName) - errHandlerOptions := []errors.HandlerOption{ - postgres.WithDuplicateErrors(NewDuplicateErrors()), - } - return &groupRepository{ - db: db, - eh: postgres.NewErrorHandler(errHandlerOptions...), - Repository: roleRepo, - } -} - -func (repo groupRepository) Save(ctx context.Context, g groups.Group) (groups.Group, error) { - q, computedPath, err := repo.getInsertQuery(ctx, g) - if err != nil { - return groups.Group{}, errors.Wrap(repoerr.ErrCreateEntity, err) - } - dbg, err := toDBGroup(g) - if err != nil { - return groups.Group{}, repo.eh.HandleError(repoerr.ErrCreateEntity, err) - } - if computedPath != "" { - dbg.Path = computedPath - } - - row, err := repo.db.NamedQueryContext(ctx, q, dbg) - if err != nil { - return groups.Group{}, repo.eh.HandleError(repoerr.ErrCreateEntity, err) - } - - defer row.Close() - row.Next() - dbg = dbGroup{} - if err := row.StructScan(&dbg); err != nil { - return groups.Group{}, repo.eh.HandleError(repoerr.ErrCreateEntity, err) - } - - return toGroup(dbg) -} - -func (repo groupRepository) Update(ctx context.Context, g groups.Group) (groups.Group, error) { - var query []string - var upq string - if g.Name != "" { - query = append(query, "name = :name,") - } - if g.Description.Valid { - query = append(query, "description = :description,") - } - if g.Metadata != nil { - query = append(query, "metadata = :metadata,") - } - if len(query) > 0 { - upq = strings.Join(query, " ") - } - g.Status = groups.EnabledStatus - q := fmt.Sprintf(`UPDATE groups SET %s updated_at = :updated_at, updated_by = :updated_by - WHERE id = :id AND status = :status - RETURNING id, name, tags, description, domain_id, COALESCE(parent_id, '') AS parent_id, metadata, created_at, updated_at, updated_by, status`, upq) - - dbu, err := toDBGroup(g) - if err != nil { - return groups.Group{}, errors.Wrap(repoerr.ErrUpdateEntity, err) - } - - row, err := repo.db.NamedQueryContext(ctx, q, dbu) - if err != nil { - return groups.Group{}, repo.eh.HandleError(repoerr.ErrUpdateEntity, err) - } - - defer row.Close() - if ok := row.Next(); !ok { - return groups.Group{}, errors.Wrap(repoerr.ErrNotFound, row.Err()) - } - dbu = dbGroup{} - if err := row.StructScan(&dbu); err != nil { - return groups.Group{}, repo.eh.HandleError(repoerr.ErrUpdateEntity, err) - } - return toGroup(dbu) -} - -func (repo groupRepository) UpdateTags(ctx context.Context, group groups.Group) (groups.Group, error) { - q := `UPDATE groups SET tags = :tags, updated_at = :updated_at, updated_by = :updated_by - WHERE id = :id AND status = :status - RETURNING id, name, tags, metadata, COALESCE(domain_id, '') AS domain_id, COALESCE(parent_id, '') AS parent_id, status, created_at, updated_at, updated_by` - group.Status = groups.EnabledStatus - - dbg, err := toDBGroup(group) - if err != nil { - return groups.Group{}, errors.Wrap(repoerr.ErrUpdateEntity, err) - } - - row, err := repo.db.NamedQueryContext(ctx, q, dbg) - if err != nil { - return groups.Group{}, repo.eh.HandleError(repoerr.ErrUpdateEntity, err) - } - defer row.Close() - - dbg = dbGroup{} - if row.Next() { - if err := row.StructScan(&dbg); err != nil { - return groups.Group{}, repo.eh.HandleError(repoerr.ErrUpdateEntity, err) - } - - return toGroup(dbg) - } - - return groups.Group{}, repoerr.ErrNotFound -} - -func (repo groupRepository) ChangeStatus(ctx context.Context, group groups.Group) (groups.Group, error) { - qc := `UPDATE groups SET status = :status, updated_at = :updated_at, updated_by = :updated_by WHERE id = :id - RETURNING id, name, tags, description, domain_id, COALESCE(parent_id, '') AS parent_id, metadata, created_at, updated_at, updated_by, status` - - dbg, err := toDBGroup(group) - if err != nil { - return groups.Group{}, errors.Wrap(repoerr.ErrUpdateEntity, err) - } - row, err := repo.db.NamedQueryContext(ctx, qc, dbg) - if err != nil { - return groups.Group{}, repo.eh.HandleError(repoerr.ErrUpdateEntity, err) - } - defer row.Close() - if ok := row.Next(); !ok { - return groups.Group{}, errors.Wrap(repoerr.ErrNotFound, row.Err()) - } - dbg = dbGroup{} - if err := row.StructScan(&dbg); err != nil { - return groups.Group{}, repo.eh.HandleError(repoerr.ErrUpdateEntity, err) - } - - return toGroup(dbg) -} - -func (repo groupRepository) RetrieveByID(ctx context.Context, id string) (groups.Group, error) { - q := `SELECT id, name, tags, domain_id, COALESCE(parent_id, '') AS parent_id, description, metadata, created_at, updated_at, updated_by, status, path FROM groups - WHERE id = :id` - - dbg := dbGroup{ - ID: id, - } - - row, err := repo.db.NamedQueryContext(ctx, q, dbg) - if err != nil { - return groups.Group{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - defer row.Close() - - dbg = dbGroup{} - if ok := row.Next(); !ok { - return groups.Group{}, repoerr.ErrNotFound - } - if err := row.StructScan(&dbg); err != nil { - return groups.Group{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - return toGroup(dbg) -} - -func (repo groupRepository) RetrieveByIDWithRoles(ctx context.Context, id, memberID string) (groups.Group, error) { - query := ` - WITH selected_group AS ( - SELECT - g.id, - g.parent_id, - g.domain_id, - g.path AS parent_group_path - FROM - groups g - WHERE - g.id = :id - LIMIT 1 - ), - selected_group_roles AS ( - SELECT - sg.id AS group_id, - grm.member_id AS member_id, - gr.id AS role_id, - gr.name AS role_name, - jsonb_agg(DISTINCT gra.action) AS actions, - g.path AS access_provider_path, - gr.entity_id AS access_provider_id, - CASE - WHEN gr.entity_id = sg.id THEN 'direct' - WHEN gr.entity_id = sg.parent_id THEN 'direct_group' - ELSE 'indirect_group' - END AS access_type - FROM - groups g - JOIN - groups_roles gr ON gr.entity_id = g.id - JOIN - groups_role_members grm ON gr.id = grm.role_id - JOIN - groups_role_actions gra ON gr.id = gra.role_id - JOIN - selected_group sg ON g.path @> sg.parent_group_path - WHERE - grm.member_id = :member_id - AND ( - (gr.entity_id = sg.id) - OR (gr.entity_id <> sg.id AND gra.action LIKE 'subgroup%%') - ) - GROUP BY - sg.id, gr.entity_id, gr.id, gr.name, g.path, grm.member_id, sg.parent_id - ), - selected_domain_roles AS ( - SELECT - sg.id AS group_id, - drm.member_id AS member_id, - dr.id AS role_id, - dr.name AS role_name, - jsonb_agg(DISTINCT all_actions.action) AS actions, - CAST('' AS ltree) access_provider_path, - 'domain' AS access_type, - dr.entity_id AS access_provider_id - FROM - domains d - JOIN - selected_group sg ON sg.domain_id = d.id - JOIN - domains_roles dr ON dr.entity_id = d.id - JOIN - domains_role_members drm ON dr.id = drm.role_id - JOIN - domains_role_actions dra ON dr.id = dra.role_id - JOIN - domains_role_actions all_actions ON dr.id = all_actions.role_id - WHERE - drm.member_id = :member_id - AND dra.action LIKE 'group%%' - GROUP BY - sg.id, dr.entity_id, dr.id, dr.name, drm.member_id - ), - all_roles AS ( - SELECT - sgr.group_id, - sgr.member_id, - sgr.role_id AS role_id, - sgr.role_name AS role_name, - sgr.actions AS actions, - sgr.access_type AS access_type, - sgr.access_provider_path AS access_provider_path, - sgr.access_provider_id AS access_provider_id - FROM - selected_group_roles sgr - UNION - SELECT - sdr.group_id, - sdr.member_id, - sdr.role_id AS role_id, - sdr.role_name AS role_name, - sdr.actions AS actions, - sdr.access_type AS access_type, - sdr.access_provider_path AS access_provider_path, - sdr.access_provider_id AS access_provider_id - FROM - selected_domain_roles sdr - ), - final_roles AS ( - SELECT - ar.group_id, - ar.member_id, - jsonb_agg( - jsonb_build_object( - 'role_id', ar.role_id, - 'role_name', ar.role_name, - 'actions', ar.actions, - 'access_type', ar.access_type, - 'access_provider_path', ar.access_provider_path, - 'access_provider_id', ar.access_provider_id - ) - ) AS roles - FROM all_roles ar - GROUP BY - ar.group_id, ar.member_id - ) - SELECT - g.id, - g.parent_id, - g.domain_id, - g.name, - g.tags, - g.description, - g.path, - g.metadata, - g.created_at, - g.updated_at, - g.updated_by, - g.status, - fr.member_id, - fr.roles - FROM groups g - JOIN final_roles fr ON fr.group_id = g.id - ` - - parameters := map[string]any{ - "id": id, - "member_id": memberID, - } - row, err := repo.db.NamedQueryContext(ctx, query, parameters) - if err != nil { - return groups.Group{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - defer row.Close() - - dbg := dbGroup{} - if !row.Next() { - return groups.Group{}, repoerr.ErrNotFound - } - - if err := row.StructScan(&dbg); err != nil { - return groups.Group{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - - return toGroup(dbg) -} - -func (repo groupRepository) RetrieveByIDAndUser(ctx context.Context, domainID, userID, groupID string) (groups.Group, error) { - baseQuery := userGroupsBaseQuery - - dbg := dbGroup{ID: groupID, UserID: userID, DomainIDParam: domainID} - q := fmt.Sprintf(`%s - SELECT - g.id, - g.name, - g.domain_id, - COALESCE(g.parent_id, '') AS parent_id, - g.tags, - g.description, - g.metadata, - g.created_at, - g.updated_at, - g.updated_by, - g.status, - g.path as path, - g.role_id, - g.role_name, - g.actions, - g.access_type, - g.access_provider_id, - g.access_provider_role_id, - g.access_provider_role_name, - g.access_provider_role_actions - FROM - final_groups g - WHERE - g.id = :id - LIMIT 1 - ; - `, - baseQuery) - - row, err := repo.db.NamedQueryContext(ctx, q, dbg) - if err != nil { - return groups.Group{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - defer row.Close() - - dbg = dbGroup{} - if ok := row.Next(); !ok { - return groups.Group{}, repoerr.ErrNotFound - } - if err := row.StructScan(&dbg); err != nil { - return groups.Group{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - return toGroup(dbg) -} - -func (repo groupRepository) RetrieveAll(ctx context.Context, pm groups.PageMeta) (groups.Page, error) { - query := buildQuery(pm) - - if pm.RootGroup { - query += " AND nlevel(g.path) = 1 " - } - - orderClause := "" - var orderBy string - switch pm.Order { - case "name": - orderBy = "g.name" - case "created_at": - orderBy = "g.created_at" - case "updated_at": - orderBy = "COALESCE(g.updated_at, g.created_at)" - } - - if orderBy != "" { - dir := pm.Dir - if dir != api.AscDir && dir != api.DescDir { - dir = api.DescDir - } - orderClause = fmt.Sprintf("ORDER BY %s %s, g.id %s", orderBy, dir, dir) - } - - dbPageMeta, err := toDBGroupPageMeta(pm) - if err != nil { - return groups.Page{}, errors.Wrap(repoerr.ErrFailedToRetrieveAllGroups, err) - } - - if pm.OnlyTotal { - cq := fmt.Sprintf(`SELECT COUNT(*) FROM groups g %s;`, query) - total, err := postgres.Total(ctx, repo.db, cq, dbPageMeta) - if err != nil { - return groups.Page{}, repo.eh.HandleError(repoerr.ErrFailedToRetrieveAllGroups, err) - } - page := groups.Page{PageMeta: pm} - page.Total = total - return page, nil - } - - q := fmt.Sprintf(`SELECT g.id, g.domain_id, tags, COALESCE(g.parent_id, '') AS parent_id, g.name, g.description, - g.metadata, g.created_at, g.updated_at, g.updated_by, g.status, - COUNT(*) OVER() AS total_count FROM groups g %s %s LIMIT :limit OFFSET :offset;`, query, orderClause) - - rows, err := repo.db.NamedQueryContext(ctx, q, dbPageMeta) - if err != nil { - return groups.Page{}, repo.eh.HandleError(repoerr.ErrFailedToRetrieveAllGroups, err) - } - defer rows.Close() - - var total uint64 - var items []groups.Group - for rows.Next() { - dbg := dbGroup{} - if err := rows.StructScan(&dbg); err != nil { - return groups.Page{}, repo.eh.HandleError(repoerr.ErrFailedToRetrieveAllGroups, err) - } - total = dbg.TotalCount - g, err := toGroup(dbg) - if err != nil { - return groups.Page{}, err - } - items = append(items, g) - } - - if len(items) == 0 { - cq := fmt.Sprintf(`SELECT COUNT(*) FROM groups g %s;`, query) - total, err = postgres.Total(ctx, repo.db, cq, dbPageMeta) - if err != nil { - return groups.Page{}, repo.eh.HandleError(repoerr.ErrFailedToRetrieveAllGroups, err) - } - } - - page := groups.Page{PageMeta: pm} - page.Total = total - page.Groups = items - return page, nil -} - -func (repo groupRepository) RetrieveByIDs(ctx context.Context, pm groups.PageMeta, ids ...string) (groups.Page, error) { - if (len(ids) == 0) && (pm.DomainID == "") { - return groups.Page{PageMeta: groups.PageMeta{Offset: pm.Offset, Limit: pm.Limit}}, nil - } - query := buildQuery(pm, ids...) - - q := fmt.Sprintf(`SELECT DISTINCT g.id, g.domain_id, tags, COALESCE(g.parent_id, '') AS parent_id, g.name, g.tags, g.description, - g.metadata, g.created_at, g.updated_at, g.updated_by, g.status, - COUNT(*) OVER() AS total_count FROM groups g %s ORDER BY g.created_at LIMIT :limit OFFSET :offset;`, query) - - dbPageMeta, err := toDBGroupPageMeta(pm) - if err != nil { - return groups.Page{}, errors.Wrap(repoerr.ErrFailedToRetrieveAllGroups, err) - } - dbPageMeta.IDs = pq.StringArray(ids) - rows, err := repo.db.NamedQueryContext(ctx, q, dbPageMeta) - if err != nil { - return groups.Page{}, repo.eh.HandleError(repoerr.ErrFailedToRetrieveAllGroups, err) - } - defer rows.Close() - - var total uint64 - var items []groups.Group - for rows.Next() { - dbg := dbGroup{} - if err := rows.StructScan(&dbg); err != nil { - return groups.Page{}, repo.eh.HandleError(repoerr.ErrFailedToRetrieveAllGroups, err) - } - total = dbg.TotalCount - g, err := toGroup(dbg) - if err != nil { - return groups.Page{}, err - } - items = append(items, g) - } - - if len(items) == 0 { - cq := fmt.Sprintf(`SELECT COUNT(*) FROM ( - SELECT DISTINCT g.id FROM groups g %s - ) AS subquery;`, query) - total, err = postgres.Total(ctx, repo.db, cq, dbPageMeta) - if err != nil { - return groups.Page{}, repo.eh.HandleError(repoerr.ErrFailedToRetrieveAllGroups, err) - } - } - - page := groups.Page{PageMeta: pm} - page.Total = total - page.Groups = items - return page, nil -} - -func (repo groupRepository) RetrieveHierarchy(ctx context.Context, domainID, userID, groupID string, hm groups.HierarchyPageMeta) (groups.HierarchyPage, error) { - var dirQuery string - switch { - case hm.Direction >= 0: - dirQuery = "g.path @> (SELECT path FROM groups WHERE id = :id)" - default: - dirQuery = "g.path <@ (SELECT path FROM groups WHERE id = :id)" - } - - baseQuery := userGroupsBaseQuery - query := fmt.Sprintf(`%s, - target_hierarchy AS ( - SELECT - g.id, - g.parent_id, - g.domain_id, - g.name, - g.tags, - g.description, - g.metadata, - g.created_at, - g.updated_at, - g.updated_by, - g.status, - g.path, - nlevel(g.path) AS level - FROM - groups g - WHERE - %s - ), - filtered_hierarchy AS ( - SELECT - th.* - FROM - target_hierarchy th - JOIN - final_groups fg ON th.id = fg.id - ) - SELECT - * - FROM - filtered_hierarchy - ORDER BY path; - `, baseQuery, dirQuery) - - parameters := map[string]any{ - "id": groupID, - "level": hm.Level, - "user_id": userID, - "domain_id_param": domainID, - } - - rows, err := repo.db.NamedQueryContext(ctx, query, parameters) - if err != nil { - return groups.HierarchyPage{}, repo.eh.HandleError(repoerr.ErrFailedToRetrieveAllGroups, err) - } - defer rows.Close() - - items, err := repo.processRows(rows) - if err != nil { - return groups.HierarchyPage{}, repo.eh.HandleError(repoerr.ErrFailedToRetrieveAllGroups, err) - } - - return groups.HierarchyPage{HierarchyPageMeta: hm, Groups: items}, nil -} - -func (repo groupRepository) AssignParentGroup(ctx context.Context, parentGroupID string, groupIDs ...string) (err error) { - if len(groupIDs) == 0 { - return nil - } - - tx, err := repo.db.BeginTxx(ctx, nil) - if err != nil { - return repo.eh.HandleError(repoerr.ErrUpdateEntity, err) - } - defer func() { - if err != nil { - if errRollback := tx.Rollback(); errRollback != nil { - err = errors.Wrap(err, errRollback) - } - } - }() - - pq := `SELECT id, path FROM groups WHERE id = $1 LIMIT 1;` - rows, err := tx.Queryx(pq, parentGroupID) - if err != nil { - return repo.eh.HandleError(repoerr.ErrUpdateEntity, err) - } - defer rows.Close() - - pGroups, err := repo.processRows(rows) - if err != nil { - return repo.eh.HandleError(repoerr.ErrUpdateEntity, err) - } - if len(pGroups) == 0 { - return repoerr.ErrUpdateEntity - } - pGroup := pGroups[0] - - if pGroup.ID == "" { - return errors.Wrap(repoerr.ErrViewEntity, errParentGroupID) - } - if pGroup.Path == "" { - return errors.Wrap(repoerr.ErrViewEntity, errParentGroupPath) - } - if !strings.HasSuffix(pGroup.Path, pGroup.ID) { - return errors.Wrap(repoerr.ErrViewEntity, errParentSuffix) - } - sPaths := strings.Split(pGroup.Path, ".") // 021b9f24-5337-469b-abfa-586f5813dd41.bd4a1fea-6303-4dca-9628-301cd1165a8c.c7e8f389-11e9-4849-a474-e186012ddf38 - for _, sPath := range sPaths { - for _, cgid := range groupIDs { - if sPath == cgid { - return errors.Wrap(repoerr.ErrUpdateEntity, errCyclicParentGroup) - } - } - } - - query := ` UPDATE groups - SET parent_id = :parent_id - WHERE id = ANY(:children_group_ids) - RETURNING id, path;` - - params := map[string]any{ - "parent_id": pGroup.ID, - "children_group_ids": groupIDs, - } - - crows, err := tx.NamedQuery(query, params) - if err != nil { - return repo.eh.HandleError(repoerr.ErrUpdateEntity, err) - } - defer crows.Close() - cgroups, err := repo.processRows(crows) - if err != nil { - return repo.eh.HandleError(repoerr.ErrUpdateEntity, err) - } - - childrenPaths := []string{} - for _, cg := range cgroups { - spath := strings.Split(cg.Path, ".") - if len(spath) > 0 { - childrenPaths = append(childrenPaths, cg.Path) - } - } - - query = `UPDATE groups - SET path = text2ltree(COALESCE($1, '') || '.' || ltree2text(path)) - WHERE path <@ ANY($2::ltree[]);` - - if _, err := tx.Exec(query, pGroup.Path, childrenPaths); err != nil { - return repo.eh.HandleError(repoerr.ErrUpdateEntity, err) - } - - if err := tx.Commit(); err != nil { - return repo.eh.HandleError(repoerr.ErrUpdateEntity, err) - } - return nil -} - -func (repo groupRepository) UnassignParentGroup(ctx context.Context, parentGroupID string, groupIDs ...string) (err error) { - if len(groupIDs) == 0 { - return nil - } - - tx, err := repo.db.BeginTxx(ctx, nil) - if err != nil { - return repo.eh.HandleError(repoerr.ErrUpdateEntity, err) - } - defer func() { - if err != nil { - if errRollback := tx.Rollback(); errRollback != nil { - err = errors.Wrap(err, errRollback) - } - } - }() - pq := `SELECT id, path FROM groups WHERE id = $1 LIMIT 1;` - rows, err := tx.Queryx(pq, parentGroupID) - if err != nil { - return repo.eh.HandleError(repoerr.ErrUpdateEntity, err) - } - defer rows.Close() - - pGroups, err := repo.processRows(rows) - if err != nil { - return repo.eh.HandleError(repoerr.ErrUpdateEntity, err) - } - if len(pGroups) == 0 { - return repoerr.ErrUpdateEntity - } - pGroup := pGroups[0] - - if pGroup.ID == "" { - return errors.Wrap(repoerr.ErrViewEntity, errParentGroupID) - } - if pGroup.Path == "" { - return errors.Wrap(repoerr.ErrViewEntity, errParentGroupPath) - } - - query := `UPDATE groups - SET parent_id = NULL - WHERE id = ANY(:children_group_ids) AND parent_id = :parent_id - RETURNING id, path;` - - parameters := map[string]any{ - "parent_id": pGroup.ID, - "children_group_ids": groupIDs, - } - crows, err := tx.NamedQuery(query, parameters) - if err != nil { - return repo.eh.HandleError(repoerr.ErrUpdateEntity, err) - } - defer crows.Close() - cgroups, err := repo.processRows(crows) - if err != nil { - return repo.eh.HandleError(repoerr.ErrUpdateEntity, err) - } - - childrenPaths := []string{} - for _, cg := range cgroups { - spath := strings.Split(cg.Path, ".") - if len(spath) > 0 { - childrenPaths = append(childrenPaths, cg.Path) - } - } - - query = `UPDATE groups - SET path = text2ltree(replace(ltree2text(path), $1 || '.', '')) - WHERE path <@ ANY($2::ltree[]);` - - if _, err := tx.Exec(query, pGroup.Path, childrenPaths); err != nil { - return repo.eh.HandleError(repoerr.ErrUpdateEntity, err) - } - - if err := tx.Commit(); err != nil { - return repo.eh.HandleError(repoerr.ErrUpdateEntity, err) - } - return nil -} - -func (repo groupRepository) UnassignAllChildrenGroups(ctx context.Context, id string) error { - query := ` - UPDATE groups AS g SET - parent_id = NULL - WHERE g.parent_id = :parent_id ; - ` - - result, err := repo.db.NamedExecContext(ctx, query, dbGroup{ParentID: &id}) - if err != nil { - return repo.eh.HandleError(repoerr.ErrUpdateEntity, err) - } - if rows, _ := result.RowsAffected(); rows == 0 { - return repoerr.ErrNotFound - } - - return nil -} - -func (repo groupRepository) Delete(ctx context.Context, groupID string) error { - q := "DELETE FROM groups AS g WHERE g.id = $1;" - - result, err := repo.db.ExecContext(ctx, q, groupID) - if err != nil { - return repo.eh.HandleError(repoerr.ErrRemoveEntity, err) - } - if rows, _ := result.RowsAffected(); rows == 0 { - return repoerr.ErrNotFound - } - return nil -} - -func (repo groupRepository) RetrieveAllParentGroups(ctx context.Context, domainID, userID, groupID string, pm groups.PageMeta) (groups.Page, error) { - cGroup, err := repo.RetrieveByID(ctx, groupID) - if err != nil { - return groups.Page{}, err - } - - query := buildQuery(pm) - - levelCondition := "g.path @> CAST(:path AS ltree) " - - switch { - case query == "": - query = " WHERE " + levelCondition - default: - query = query + " AND " + levelCondition - } - - pm.Path = cGroup.Path - return repo.retrieveGroups(ctx, domainID, userID, query, pm) -} - -func (repo groupRepository) RetrieveChildrenGroups(ctx context.Context, domainID, userID, groupID string, startLevel, endLevel int64, pm groups.PageMeta) (groups.Page, error) { - pGroup, err := repo.RetrieveByID(ctx, groupID) - if err != nil { - return groups.Page{}, err - } - - query := buildQuery(pm) - - levelCondition := "" - switch { - // Retrieve all children groups from parent group level - case startLevel == 0 && endLevel < 0: - levelCondition = " path ~ CAST(:path || '.*' AS lquery) " - - // Retrieve specific level of children groups from parent group level - case (startLevel > 0) && (startLevel == endLevel || endLevel == 0): - levelCondition = fmt.Sprintf(" path ~ CAST(:path || '.*{%d}' AS lquery) ", startLevel) - - // Retrieve all children groups from specific level from parent group level - case startLevel > 0 && endLevel < 0: - levelCondition = fmt.Sprintf(" path ~ CAST(:path || '.*{%d,}' AS lquery) ", startLevel) - - // Retrieve children groups between specific level from parent group level - case startLevel > 0 && endLevel > 0 && startLevel < endLevel: - levelCondition = fmt.Sprintf(" path ~ CAST(:path || '.*{%d,%d}' AS lquery) ", startLevel, endLevel) - default: - return groups.Page{}, errors.Wrap(repoerr.ErrViewEntity, fmt.Errorf("invalid level range: start level: %d end level: %d", startLevel, endLevel)) - } - - switch { - case query == "": - query = " WHERE " + levelCondition - default: - query = query + " AND " + levelCondition - } - - pm.Path = pGroup.Path - return repo.retrieveGroups(ctx, domainID, userID, query, pm) -} - -func (repo groupRepository) RetrieveUserGroups(ctx context.Context, domainID, userID string, pm groups.PageMeta) (groups.Page, error) { - query := buildQuery(pm) - if pm.RootGroup { - query += (` AND - NOT EXISTS ( - SELECT 1 - FROM groups anc - JOIN final_groups fg - ON fg.id = anc.id - WHERE anc.domain_id = g.domain_id - AND anc.path @> g.path - AND anc.id <> g.id - )`) - } - - return repo.retrieveGroups(ctx, domainID, userID, query, pm) -} - -func (repo groupRepository) retrieveGroups(ctx context.Context, domainID, userID, query string, pm groups.PageMeta) (groups.Page, error) { - baseQuery := userGroupsBaseQuery - - orderClause := "" - var orderBy string - switch pm.Order { - case "name": - orderBy = "g.name" - case "created_at": - orderBy = "g.created_at" - case "updated_at", "": - orderBy = "COALESCE(g.updated_at, g.created_at)" - } - - if orderBy != "" { - dir := pm.Dir - if dir != api.AscDir && dir != api.DescDir { - dir = api.DescDir - } - orderClause = fmt.Sprintf("ORDER BY %s %s, g.id %s", orderBy, dir, dir) - } - - dbPageMeta, err := toDBGroupPageMeta(pm) - if err != nil { - return groups.Page{}, errors.Wrap(repoerr.ErrFailedToRetrieveAllGroups, err) - } - dbPageMeta.UserID = userID - dbPageMeta.DomainIDParam = domainID - - if pm.OnlyTotal { - cq := fmt.Sprintf(`%s - SELECT COUNT(*) AS total_count - FROM final_groups g - %s; - `, baseQuery, query) - - total, err := postgres.Total(ctx, repo.db, cq, dbPageMeta) - if err != nil { - return groups.Page{}, repo.eh.HandleError(repoerr.ErrFailedToRetrieveAllGroups, err) - } - - page := groups.Page{PageMeta: pm} - page.Total = total - return page, nil - } - - q := fmt.Sprintf(`%s - SELECT - g.id, - g.name, - g.domain_id, - COALESCE(g.parent_id, '') AS parent_id, - g.description, - g.tags, - g.metadata, - g.created_at, - g.updated_at, - g.updated_by, - g.status, - g.path as path, - g.role_id, - g.role_name, - g.actions, - g.access_type, - g.access_provider_id, - g.access_provider_role_id, - g.access_provider_role_name, - g.access_provider_role_actions, - COUNT(*) OVER() AS total_count - FROM final_groups g - %s - %s - LIMIT :limit OFFSET :offset;`, - baseQuery, query, orderClause) - - rows, err := repo.db.NamedQueryContext(ctx, q, dbPageMeta) - if err != nil { - return groups.Page{}, repo.eh.HandleError(repoerr.ErrFailedToRetrieveAllGroups, err) - } - defer rows.Close() - - var total uint64 - var items []groups.Group - for rows.Next() { - dbg := dbGroup{} - if err := rows.StructScan(&dbg); err != nil { - return groups.Page{}, repo.eh.HandleError(repoerr.ErrFailedToRetrieveAllGroups, err) - } - - total = dbg.TotalCount - - group, err := toGroup(dbg) - if err != nil { - return groups.Page{}, repo.eh.HandleError(repoerr.ErrFailedToRetrieveAllGroups, err) - } - items = append(items, group) - } - - if len(items) == 0 { - cq := fmt.Sprintf(`%s - SELECT COUNT(*) AS total_count - FROM final_groups g - %s; - `, baseQuery, query) - - total, err = postgres.Total(ctx, repo.db, cq, dbPageMeta) - if err != nil { - return groups.Page{}, repo.eh.HandleError(repoerr.ErrFailedToRetrieveAllGroups, err) - } - } - - page := groups.Page{PageMeta: pm} - page.Total = total - page.Groups = items - return page, nil -} - -const userGroupsBaseQuery = ` -WITH direct_groups AS ( -SELECT - g.*, - gr.entity_id AS entity_id, - grm.member_id AS member_id, - gr.id AS role_id, - gr."name" AS role_name, - array_agg(gra."action") AS actions -FROM - groups_role_members grm -JOIN - groups_role_actions gra ON gra.role_id = grm.role_id -JOIN - groups_roles gr ON gr.id = grm.role_id -JOIN - "groups" g ON g.id = gr.entity_id -WHERE - grm.member_id = :user_id - AND g.domain_id = :domain_id_param -GROUP BY - gr.entity_id, grm.member_id, gr.id, gr."name", g."path", g.id -), -direct_groups_with_subgroup AS ( - SELECT - g.*, - gr.entity_id AS entity_id, - grm.member_id AS member_id, - gr.id AS role_id, - gr."name" AS role_name, - array_agg(DISTINCT gra."action") AS actions - FROM - groups_role_members grm - JOIN - groups_role_actions gra ON gra.role_id = grm.role_id - JOIN - groups_roles gr ON gr.id = grm.role_id - JOIN - "groups" g ON g.id = gr.entity_id - WHERE - grm.member_id = :user_id - AND g.domain_id = :domain_id_param - GROUP BY - gr.entity_id, grm.member_id, gr.id, gr."name", g."path", g.id - HAVING - bool_or(gra."action" LIKE 'subgroup_%') -), -direct_leaf_groups_with_subgroup AS ( - SELECT dgws.* - FROM direct_groups_with_subgroup dgws - WHERE NOT EXISTS ( - SELECT 1 - FROM direct_groups_with_subgroup dgws2 - WHERE - dgws2.path @> dgws.path - AND dgws2.id != dgws.id - ) -), -indirect_child_groups AS ( - SELECT - DISTINCT indirect_child_groups.id as child_id, - indirect_child_groups.*, - dlgws.id as access_provider_id, - dlgws.role_id as access_provider_role_id, - dlgws.role_name as access_provider_role_name, - dlgws.actions as access_provider_role_actions - FROM - direct_leaf_groups_with_subgroup dlgws - JOIN - groups indirect_child_groups ON indirect_child_groups.path <@ dlgws.path - WHERE - indirect_child_groups.domain_id = :domain_id_param - AND NOT EXISTS ( - SELECT 1 - FROM direct_groups_with_subgroup dgws - WHERE dgws.id = indirect_child_groups.id - ) -), -direct_indirect_groups as ( - SELECT - id, - parent_id, - domain_id, - "name", - tags, - description, - metadata, - created_at, - updated_at, - updated_by, - status, - "path", - role_id, - role_name, - actions, - 'direct' AS access_type, - '' AS access_provider_id, - '' AS access_provider_role_id, - '' AS access_provider_role_name, - CAST(array[] AS text[]) AS access_provider_role_actions - FROM - direct_groups - UNION - SELECT - id, - parent_id, - domain_id, - "name", - tags, - description, - metadata, - created_at, - updated_at, - updated_by, - status, - "path", - '' AS role_id, - '' AS role_name, - CAST(array[] AS text[]) AS actions, - 'indirect' AS access_type, - access_provider_id, - access_provider_role_id, - access_provider_role_name, - access_provider_role_actions - FROM - indirect_child_groups -), -final_groups AS ( - SELECT - dig.id, - dig.parent_id, - dig.domain_id, - dig."name", - dig.tags, - dig.description, - dig.metadata, - dig.created_at, - dig.updated_at, - dig.updated_by, - dig.status, - dig."path", - dig.role_id, - dig.role_name, - dig.actions, - dig.access_type, - dig.access_provider_id, - dig.access_provider_role_id, - dig.access_provider_role_name, - dig.access_provider_role_actions - FROM - direct_indirect_groups as dig - UNION - SELECT - dg.id, - dg.parent_id, - dg.domain_id, - dg."name", - dg.tags, - dg.description, - dg.metadata, - dg.created_at, - dg.updated_at, - dg.updated_by, - dg.status, - dg."path", - '' AS role_id, - '' AS role_name, - CAST(array[] AS text[]) AS actions, - 'domain' AS access_type, - d.id AS access_provider_id, - dr.id AS access_provider_role_id, - dr."name" AS access_provider_role_name, - array_agg(dra."action") as actions - FROM - domains_role_members drm - JOIN - domains_role_actions dra ON dra.role_id = drm.role_id - JOIN - domains_roles dr ON dr.id = drm.role_id - JOIN - domains d ON d.id = dr.entity_id - JOIN - "groups" dg ON dg.domain_id = d.id - WHERE - drm.member_id = :user_id - AND d.id = :domain_id_param - AND dra."action" LIKE 'group_%' - AND NOT EXISTS ( - SELECT 1 FROM direct_indirect_groups dig - WHERE dig.id = dg.id - ) - GROUP BY - dg.id, d.id, dr.id -) - ` - -func buildQuery(gm groups.PageMeta, ids ...string) string { - queries := []string{} - - if len(ids) > 0 { - queries = append(queries, "id = ANY(:ids)") - } - if gm.Name != "" { - queries = append(queries, "g.name ILIKE '%' || :name || '%'") - } - if gm.ID != "" { - queries = append(queries, "g.id = :id") - } - if gm.Status != groups.AllStatus { - queries = append(queries, "g.status = :status") - } - if len(gm.Tags.Elements) > 0 { - switch gm.Tags.Operator { - case groups.AndOp: - queries = append(queries, "tags @> :tags") - default: // OR - queries = append(queries, "tags && :tags") - } - } - if gm.DomainID != "" { - queries = append(queries, "g.domain_id = :domain_id") - } - if gm.AccessType != "" { - queries = append(queries, "g.access_type = :access_type") - } - if gm.RoleID != "" { - queries = append(queries, "g.role_id = :role_id") - } - if gm.RoleName != "" { - queries = append(queries, "g.role_name = :role_name") - } - if len(gm.Actions) != 0 { - queries = append(queries, "g.actions @> :actions") - } - if len(gm.Metadata) > 0 { - queries = append(queries, "g.metadata @> :metadata") - } - if !gm.CreatedFrom.IsZero() { - queries = append(queries, "g.created_at >= :created_from") - } - if !gm.CreatedTo.IsZero() { - queries = append(queries, "g.created_at <= :created_to") - } - if len(queries) > 0 { - return fmt.Sprintf("WHERE %s", strings.Join(queries, " AND ")) - } - - return "" -} - -type dbGroup struct { - ID string `db:"id"` - ParentID *string `db:"parent_id,omitempty"` - DomainID string `db:"domain_id,omitempty"` - Name string `db:"name"` - Description sql.NullString `db:"description,omitempty"` - Tags pgtype.TextArray `db:"tags,omitempty"` - Level int `db:"level"` - Path string `db:"path,omitempty"` - Metadata []byte `db:"metadata,omitempty"` - CreatedAt time.Time `db:"created_at"` - UpdatedAt sql.NullTime `db:"updated_at,omitempty"` - UpdatedBy *string `db:"updated_by,omitempty"` - Status groups.Status `db:"status"` - RoleID string `db:"role_id"` - RoleName string `db:"role_name"` - Actions pq.StringArray `db:"actions"` - AccessType string `db:"access_type"` - AccessProviderId string `db:"access_provider_id"` - AccessProviderRoleId string `db:"access_provider_role_id"` - AccessProviderRoleName string `db:"access_provider_role_name"` - AccessProviderRoleActions pq.StringArray `db:"access_provider_role_actions"` - MemberID string `db:"member_id,omitempty"` - Roles json.RawMessage `db:"roles,omitempty"` - TotalCount uint64 `db:"total_count"` - UserID string `db:"user_id,omitempty"` - DomainIDParam string `db:"domain_id_param,omitempty"` -} - -func toDBGroup(g groups.Group) (dbGroup, error) { - data := []byte("{}") - if len(g.Metadata) > 0 { - b, err := json.Marshal(g.Metadata) - if err != nil { - return dbGroup{}, errors.Wrap(errors.ErrMalformedEntity, err) - } - data = b - } - var tags pgtype.TextArray - if err := tags.Set(g.Tags); err != nil { - return dbGroup{}, err - } - var parentID *string - if g.Parent != "" { - parentID = &g.Parent - } - var updatedAt sql.NullTime - if !g.UpdatedAt.IsZero() { - updatedAt = sql.NullTime{Time: g.UpdatedAt, Valid: true} - } - var updatedBy *string - if g.UpdatedBy != "" { - updatedBy = &g.UpdatedBy - } - return dbGroup{ - ID: g.ID, - Name: g.Name, - ParentID: parentID, - DomainID: g.Domain, - Description: sql.NullString{String: g.Description.Value, Valid: g.Description.Valid}, - Tags: tags, - Metadata: data, - Path: g.Path, - CreatedAt: g.CreatedAt, - UpdatedAt: updatedAt, - UpdatedBy: updatedBy, - Status: g.Status, - }, nil -} - -func toGroup(g dbGroup) (groups.Group, error) { - var metadata groups.Metadata - if g.Metadata != nil { - if err := json.Unmarshal(g.Metadata, &metadata); err != nil { - return groups.Group{}, errors.Wrap(repoerr.ErrMalformedEntity, err) - } - } - var tags []string - for _, e := range g.Tags.Elements { - tags = append(tags, e.String) - } - var parentID string - if g.ParentID != nil { - parentID = *g.ParentID - } - var updatedAt time.Time - if g.UpdatedAt.Valid { - updatedAt = g.UpdatedAt.Time.UTC() - } - var updatedBy string - if g.UpdatedBy != nil { - updatedBy = *g.UpdatedBy - } - - var roles []roles.MemberRoleActions - if g.Roles != nil { - if err := json.Unmarshal(g.Roles, &roles); err != nil { - return groups.Group{}, errors.Wrap(errors.ErrMalformedEntity, err) - } - } - - return groups.Group{ - ID: g.ID, - Name: g.Name, - Parent: parentID, - Domain: g.DomainID, - Description: nullable.Value[string]{Value: g.Description.String, Valid: g.Description.Valid}, - Tags: tags, - Metadata: metadata, - Level: g.Level, - Path: g.Path, - UpdatedAt: updatedAt, - UpdatedBy: updatedBy, - CreatedAt: g.CreatedAt.UTC(), - Status: g.Status, - RoleID: g.RoleID, - RoleName: g.RoleName, - Actions: g.Actions, - AccessType: g.AccessType, - AccessProviderId: g.AccessProviderId, - AccessProviderRoleId: g.AccessProviderRoleId, - AccessProviderRoleName: g.AccessProviderRoleName, - AccessProviderRoleActions: g.AccessProviderRoleActions, - Roles: roles, - }, nil -} - -func toDBGroupPageMeta(pm groups.PageMeta) (dbGroupPageMeta, error) { - data := []byte("{}") - if len(pm.Metadata) > 0 { - b, err := json.Marshal(pm.Metadata) - if err != nil { - return dbGroupPageMeta{}, errors.Wrap(errors.ErrMalformedEntity, err) - } - data = b - } - var tags pgtype.TextArray - if err := tags.Set(pm.Tags.Elements); err != nil { - return dbGroupPageMeta{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - return dbGroupPageMeta{ - ID: pm.ID, - Name: pm.Name, - Metadata: data, - Tags: tags, - Total: pm.Total, - Offset: pm.Offset, - Limit: pm.Limit, - DomainID: pm.DomainID, - Status: pm.Status, - RoleName: pm.RoleName, - RoleID: pm.RoleID, - Actions: pm.Actions, - AccessType: pm.AccessType, - Path: pm.Path, - CreatedFrom: pm.CreatedFrom, - CreatedTo: pm.CreatedTo, - }, nil -} - -type dbGroupPageMeta struct { - ID string `db:"id"` - Name string `db:"name"` - ParentID string `db:"parent_id"` - DomainID string `db:"domain_id"` - Metadata []byte `db:"metadata"` - Path string `db:"path"` - Level uint64 `db:"level"` - Total uint64 `db:"total"` - Limit uint64 `db:"limit"` - Offset uint64 `db:"offset"` - Subject string `db:"subject"` - RoleName string `db:"role_name"` - RoleID string `db:"role_id"` - Actions pq.StringArray `db:"actions"` - AccessType string `db:"access_type"` - Status groups.Status `db:"status"` - Tags pgtype.TextArray `db:"tags"` - IDs pq.StringArray `db:"ids"` - CreatedFrom time.Time `db:"created_from"` - CreatedTo time.Time `db:"created_to"` - UserID string `db:"user_id"` - DomainIDParam string `db:"domain_id_param"` -} - -func (repo groupRepository) processRows(rows *sqlx.Rows) ([]groups.Group, error) { - var items []groups.Group - for rows.Next() { - dbg := dbGroup{} - if err := rows.StructScan(&dbg); err != nil { - return items, err - } - group, err := toGroup(dbg) - if err != nil { - return items, err - } - items = append(items, group) - } - return items, nil -} - -func (repo groupRepository) getInsertQuery(c context.Context, g groups.Group) (string, string, error) { - switch { - case g.Parent != "": - parent, err := repo.RetrieveByID(c, g.Parent) - if err != nil { - return "", "", err - } - path := parent.Path + "." + g.ID - if len(strings.Split(path, ".")) > groups.MaxPathLength { - return "", "", fmt.Errorf("reached max nested depth") - } - return `INSERT INTO groups (name, description, tags, id, domain_id, parent_id, metadata, created_at, status, path) - VALUES (:name, :description, :tags, :id, :domain_id, :parent_id, :metadata, :created_at, :status, CAST(:path AS ltree)) - RETURNING id, name, description, tags, domain_id, COALESCE(parent_id, '') AS parent_id, metadata, created_at, status, path, nlevel(path) as level;`, path, nil - default: - return `INSERT INTO groups (name, description, tags, id, domain_id, metadata, created_at, status, path) - VALUES (:name, :description, :tags, :id, :domain_id, :metadata, :created_at, :status, :id) - RETURNING id, name, description, tags, domain_id, COALESCE(parent_id, '') AS parent_id, metadata, created_at, status, path, nlevel(path) as level;`, "", nil - } -} diff --git a/groups/postgres/groups_test.go b/groups/postgres/groups_test.go deleted file mode 100644 index 00c6545b2..000000000 --- a/groups/postgres/groups_test.go +++ /dev/null @@ -1,2477 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package postgres_test - -import ( - "context" - "fmt" - "strings" - "testing" - "time" - - "github.com/0x6flab/namegenerator" - "github.com/absmach/magistrala/groups" - "github.com/absmach/magistrala/groups/postgres" - "github.com/absmach/magistrala/internal/nullable" - "github.com/absmach/magistrala/internal/testsutil" - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - "github.com/absmach/magistrala/pkg/roles" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" -) - -var ( - namegen = namegenerator.NewGenerator() - invalidID = strings.Repeat("a", 37) - validTimestamp = time.Now().UTC().Truncate(time.Millisecond) - description = strings.Repeat("a", 64) - desc = nullable.New(description) - invalidDescription = strings.Repeat("a", 1025) - invalidDesc = nullable.New(invalidDescription) - - validGroup = groups.Group{ - ID: testsutil.GenerateUUID(&testing.T{}), - Domain: testsutil.GenerateUUID(&testing.T{}), - Name: namegen.Generate(), - Tags: []string{"tag1", "tag2"}, - Description: desc, - Metadata: map[string]any{"key": "value"}, - CreatedAt: time.Now().UTC().Truncate(time.Microsecond), - Status: groups.EnabledStatus, - } - directAccess = "direct" - ascDir = "asc" - descDir = "desc" - availableActions = []string{ - "update", - "read", - "membership", - "delete", - "subgroup_create", - "subgroup_client_create", - "subgroup_channel_create", - "subgroup_update", - "subgroup_read", - "subgroup_membership", - "subgroup_delete", - "subgroup_set_child", - "subgroup_set_parent", - "subgroup_manage_role", - "subgroup_add_role_users", - "subgroup_remove_role_users", - "subgroup_view_role_users", - } - errGroupExists = errors.NewRequestError("group id already exists") -) - -func TestSave(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM groups") - require.Nil(t, err, fmt.Sprintf("clean groups unexpected error: %s", err)) - }) - - validGroupRes := validGroup - validGroupRes.Path = validGroup.ID - validGroupRes.Level = 1 - - repo := postgres.New(database) - - parentGroup := validGroup - parentGroup.ID = testsutil.GenerateUUID(t) - parentGroup.Name = namegen.Generate() - - pgroup, err := repo.Save(context.Background(), parentGroup) - require.Nil(t, err, fmt.Sprintf("save group unexpected error: %s", err)) - - validChildGroup := validGroup - validChildGroup.ID = testsutil.GenerateUUID(t) - validChildGroup.Name = namegen.Generate() - validChildGroup.Parent = pgroup.ID - validChildGroupRes := validChildGroup - validChildGroupRes.Path = fmt.Sprintf("%s.%s", pgroup.Path, validChildGroupRes.ID) - validChildGroupRes.Level = 2 - duplicateGroupID := testsutil.GenerateUUID(t) - - cases := []struct { - desc string - group groups.Group - resp groups.Group - err error - }{ - { - desc: "add new group successfully", - group: validGroup, - resp: validGroupRes, - err: nil, - }, - { - desc: "add duplicate group", - group: validGroup, - err: errGroupExists, - }, - { - desc: "add group with parent", - group: validChildGroup, - resp: validChildGroupRes, - err: nil, - }, - { - desc: "add group with invalid ID", - group: groups.Group{ - ID: invalidID, - Domain: testsutil.GenerateUUID(t), - Name: namegen.Generate(), - Description: desc, - Metadata: map[string]any{"key": "value"}, - CreatedAt: validTimestamp, - Status: groups.EnabledStatus, - }, - err: repoerr.ErrCreateEntity, - }, - { - desc: "add group with invalid domain", - group: groups.Group{ - ID: testsutil.GenerateUUID(t), - Domain: invalidID, - Name: namegen.Generate(), - Description: desc, - Metadata: map[string]any{"key": "value"}, - CreatedAt: validTimestamp, - Status: groups.EnabledStatus, - }, - err: repoerr.ErrCreateEntity, - }, - { - desc: "add group with invalid parent", - group: groups.Group{ - ID: testsutil.GenerateUUID(t), - Parent: testsutil.GenerateUUID(t), - Name: namegen.Generate(), - Description: desc, - Metadata: map[string]any{"key": "value"}, - CreatedAt: validTimestamp, - Status: groups.EnabledStatus, - }, - err: repoerr.ErrNotFound, - }, - { - desc: "add group with invalid name", - group: groups.Group{ - ID: testsutil.GenerateUUID(t), - Domain: testsutil.GenerateUUID(t), - Name: strings.Repeat("a", 1025), - Description: desc, - Metadata: map[string]any{"key": "value"}, - CreatedAt: validTimestamp, - Status: groups.EnabledStatus, - }, - err: repoerr.ErrCreateEntity, - }, - { - desc: "add group with invalid description", - group: groups.Group{ - ID: testsutil.GenerateUUID(t), - Domain: testsutil.GenerateUUID(t), - Name: namegen.Generate(), - Description: invalidDesc, - Metadata: map[string]any{"key": "value"}, - CreatedAt: validTimestamp, - Status: groups.EnabledStatus, - }, - err: repoerr.ErrCreateEntity, - }, - { - desc: "add group with invalid metadata", - group: groups.Group{ - ID: testsutil.GenerateUUID(t), - Domain: testsutil.GenerateUUID(t), - Name: namegen.Generate(), - Description: desc, - Metadata: map[string]any{ - "key": make(chan int), - }, - CreatedAt: validTimestamp, - Status: groups.EnabledStatus, - }, - err: repoerr.ErrMalformedEntity, - }, - { - desc: "add group with invalid domain", - group: groups.Group{ - ID: testsutil.GenerateUUID(t), - Name: namegen.Generate(), - Domain: invalidID, - Description: desc, - Metadata: map[string]any{"key": "value"}, - CreatedAt: validTimestamp, - Status: groups.EnabledStatus, - }, - err: repoerr.ErrCreateEntity, - }, - { - desc: "add group with duplicate name", - group: groups.Group{ - ID: duplicateGroupID, - Domain: validGroup.Domain, - Name: validGroup.Name, - Description: desc, - Metadata: map[string]any{"key": "different_value"}, - CreatedAt: validTimestamp, - Status: groups.EnabledStatus, - Path: duplicateGroupID, - Level: 1, - }, - resp: groups.Group{ - ID: duplicateGroupID, - Domain: validGroup.Domain, - Name: validGroup.Name, - Description: desc, - Metadata: map[string]any{"key": "different_value"}, - CreatedAt: validTimestamp, - Status: groups.EnabledStatus, - Path: duplicateGroupID, - Level: 1, - }, - err: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - group, err := repo.Save(context.Background(), tc.group) - assert.Equal(t, tc.resp, group, fmt.Sprintf("%s: expected %v got %+v\n", tc.desc, tc.resp, group)) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - }) - } -} - -func TestUpdate(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM groups") - require.Nil(t, err, fmt.Sprintf("clean groups unexpected error: %s", err)) - }) - - repo := postgres.New(database) - - group, err := repo.Save(context.Background(), validGroup) - require.Nil(t, err, fmt.Sprintf("save group unexpected error: %s", err)) - - cases := []struct { - desc string - update string - group groups.Group - err error - }{ - { - desc: "update group successfully", - update: "all", - group: groups.Group{ - ID: group.ID, - Name: namegen.Generate(), - Description: desc, - Metadata: map[string]any{"key": "value"}, - UpdatedAt: validTimestamp, - UpdatedBy: testsutil.GenerateUUID(t), - }, - err: nil, - }, - { - desc: "update group name", - update: "name", - group: groups.Group{ - ID: group.ID, - Name: namegen.Generate(), - UpdatedAt: validTimestamp, - UpdatedBy: testsutil.GenerateUUID(t), - }, - err: nil, - }, - { - desc: "update group description", - update: "description", - group: groups.Group{ - ID: group.ID, - Description: desc, - UpdatedAt: validTimestamp, - UpdatedBy: testsutil.GenerateUUID(t), - }, - err: nil, - }, - { - desc: "update group metadata", - update: "metadata", - group: groups.Group{ - ID: group.ID, - Metadata: map[string]any{"key1": "value1"}, - UpdatedAt: validTimestamp, - UpdatedBy: testsutil.GenerateUUID(t), - }, - err: nil, - }, - { - desc: "update group with invalid ID", - update: "all", - group: groups.Group{ - ID: testsutil.GenerateUUID(t), - Name: namegen.Generate(), - Description: desc, - Metadata: map[string]any{"key": "value"}, - UpdatedAt: validTimestamp, - UpdatedBy: testsutil.GenerateUUID(t), - }, - err: repoerr.ErrNotFound, - }, - { - desc: "update group with empty ID", - update: "all", - group: groups.Group{ - Name: namegen.Generate(), - Description: desc, - Metadata: map[string]any{"key": "value"}, - UpdatedAt: validTimestamp, - UpdatedBy: testsutil.GenerateUUID(t), - }, - err: repoerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - group, err := repo.Update(context.Background(), tc.group) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - if err == nil { - assert.Equal(t, tc.group.ID, group.ID, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.group.ID, group.ID)) - assert.Equal(t, tc.group.UpdatedAt, group.UpdatedAt, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.group.UpdatedAt, group.UpdatedAt)) - assert.Equal(t, tc.group.UpdatedBy, group.UpdatedBy, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.group.UpdatedBy, group.UpdatedBy)) - switch tc.update { - case "all": - assert.Equal(t, tc.group.Name, group.Name, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.group.Name, group.Name)) - assert.Equal(t, tc.group.Description, group.Description, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.group.Description, group.Description)) - assert.Equal(t, tc.group.Metadata, group.Metadata, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.group.Metadata, group.Metadata)) - case "name": - assert.Equal(t, tc.group.Name, group.Name, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.group.Name, group.Name)) - case "description": - assert.Equal(t, tc.group.Description, group.Description, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.group.Description, group.Description)) - case "metadata": - assert.Equal(t, tc.group.Metadata, group.Metadata, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.group.Metadata, group.Metadata)) - } - } - }) - } -} - -func TestUpdateTags(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM groups") - require.Nil(t, err, fmt.Sprintf("clean groups unexpected error: %s", err)) - }) - - repo := postgres.New(database) - - _, err := repo.Save(context.Background(), validGroup) - require.Nil(t, err, fmt.Sprintf("save group unexpected error: %s", err)) - - cases := []struct { - desc string - group groups.Group - err error - }{ - { - desc: "update group tags", - group: groups.Group{ - ID: validGroup.ID, - Tags: []string{"tag3", "tag4"}, - UpdatedAt: validTimestamp, - UpdatedBy: testsutil.GenerateUUID(t), - }, - err: nil, - }, - { - desc: "update group with invalid ID", - group: groups.Group{ - ID: testsutil.GenerateUUID(t), - Tags: []string{"tag3", "tag4"}, - UpdatedAt: validTimestamp, - UpdatedBy: testsutil.GenerateUUID(t), - }, - err: repoerr.ErrNotFound, - }, - { - desc: "update group with empty ID", - group: groups.Group{ - Tags: []string{"tag3", "tag4"}, - UpdatedAt: validTimestamp, - UpdatedBy: testsutil.GenerateUUID(t), - }, - err: repoerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - group, err := repo.UpdateTags(context.Background(), tc.group) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - if err == nil { - assert.Equal(t, tc.group.ID, group.ID, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.group.ID, group.ID)) - assert.Equal(t, tc.group.UpdatedAt, group.UpdatedAt, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.group.UpdatedAt, group.UpdatedAt)) - assert.Equal(t, tc.group.UpdatedBy, group.UpdatedBy, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.group.UpdatedBy, group.UpdatedBy)) - assert.Equal(t, tc.group.Tags, group.Tags, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.group.Tags, group.Tags)) - } - }) - } -} - -func TestChangeStatus(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM groups") - require.Nil(t, err, fmt.Sprintf("clean groups unexpected error: %s", err)) - }) - - repo := postgres.New(database) - - group, err := repo.Save(context.Background(), validGroup) - require.Nil(t, err, fmt.Sprintf("save group unexpected error: %s", err)) - - cases := []struct { - desc string - group groups.Group - err error - }{ - { - desc: "change status group successfully", - group: groups.Group{ - ID: group.ID, - Status: groups.DisabledStatus, - UpdatedAt: validTimestamp, - UpdatedBy: testsutil.GenerateUUID(t), - }, - err: nil, - }, - { - desc: "change status group with invalid ID", - group: groups.Group{ - ID: testsutil.GenerateUUID(t), - Status: groups.DisabledStatus, - UpdatedAt: validTimestamp, - UpdatedBy: testsutil.GenerateUUID(t), - }, - err: repoerr.ErrNotFound, - }, - { - desc: "change status group with empty ID", - group: groups.Group{ - Status: groups.DisabledStatus, - UpdatedAt: validTimestamp, - UpdatedBy: testsutil.GenerateUUID(t), - }, - err: repoerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - group, err := repo.ChangeStatus(context.Background(), tc.group) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - if err == nil { - assert.Equal(t, tc.group.ID, group.ID, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.group.ID, group.ID)) - assert.Equal(t, tc.group.UpdatedAt, group.UpdatedAt, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.group.UpdatedAt, group.UpdatedAt)) - assert.Equal(t, tc.group.UpdatedBy, group.UpdatedBy, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.group.UpdatedBy, group.UpdatedBy)) - assert.Equal(t, tc.group.Status, group.Status, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.group.Status, group.Status)) - } - }) - } -} - -func TestRetrieveByID(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM groups") - require.Nil(t, err, fmt.Sprintf("clean groups unexpected error: %s", err)) - }) - - repo := postgres.New(database) - - validGroupRes := validGroup - validGroupRes.Path = validGroup.ID - - group, err := repo.Save(context.Background(), validGroup) - require.Nil(t, err, fmt.Sprintf("save group unexpected error: %s", err)) - - cases := []struct { - desc string - id string - group groups.Group - resp groups.Group - err error - }{ - { - desc: "retrieve group by id successfully", - id: group.ID, - group: validGroup, - resp: validGroupRes, - err: nil, - }, - { - desc: "retrieve group by id with invalid ID", - id: invalidID, - group: groups.Group{}, - err: repoerr.ErrNotFound, - }, - { - desc: "retrieve group by id with empty ID", - id: "", - group: groups.Group{}, - err: repoerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - group, err := repo.RetrieveByID(context.Background(), tc.id) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - if err == nil { - assert.Nil(t, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, group, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.group, group)) - } - }) - } -} - -func TestRetrieveByIDAndUser(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM groups") - require.Nil(t, err, fmt.Sprintf("clean groups unexpected error: %s", err)) - }) - - repo := postgres.New(database) - - domainID := testsutil.GenerateUUID(t) - userID := testsutil.GenerateUUID(t) - num := 10 - items := []groups.Group{} - for i := 0; i < num; i++ { - name := namegen.Generate() - group := groups.Group{ - ID: testsutil.GenerateUUID(t), - Domain: domainID, - Name: name, - Description: desc, - Metadata: map[string]any{"name": name}, - CreatedAt: validTimestamp, - Status: groups.EnabledStatus, - } - grp, err := repo.Save(context.Background(), group) - require.Nil(t, err, fmt.Sprintf("create group unexpected error: %s", err)) - newRolesProvision := []roles.RoleProvision{ - { - Role: roles.Role{ - ID: testsutil.GenerateUUID(t) + "_" + grp.ID, - Name: "admin", - EntityID: grp.ID, - CreatedAt: validTimestamp, - CreatedBy: userID, - }, - OptionalActions: availableActions, - OptionalMembers: []string{userID}, - }, - } - _, err = repo.AddRoles(context.Background(), newRolesProvision) - require.Nil(t, err, fmt.Sprintf("add roles unexpected error: %s", err)) - ngrp := grp - ngrp.RoleID = newRolesProvision[0].Role.ID - ngrp.RoleName = newRolesProvision[0].Role.Name - ngrp.AccessType = directAccess - items = append(items, ngrp) - } - - cases := []struct { - desc string - groupID string - userID string - domainID string - resp groups.Group - err error - }{ - { - desc: "retrieve group by id and user successfully", - groupID: items[0].ID, - userID: userID, - domainID: domainID, - resp: items[0], - err: nil, - }, - { - desc: "retrieve group by id and user successfully", - groupID: items[5].ID, - userID: userID, - domainID: domainID, - resp: items[5], - err: nil, - }, - { - desc: "retrieve group by id and user with invalid group ID", - groupID: invalidID, - userID: userID, - domainID: domainID, - err: repoerr.ErrNotFound, - }, - { - desc: "retrieve group by id and user with empty group ID", - groupID: "", - userID: userID, - domainID: domainID, - err: repoerr.ErrNotFound, - }, - { - desc: "retrieve group by id and user with invalid user ID", - groupID: items[0].ID, - userID: testsutil.GenerateUUID(t), - domainID: domainID, - err: repoerr.ErrNotFound, - }, - { - desc: "retrieve group by id and user with empty user ID", - groupID: items[0].ID, - userID: "", - domainID: domainID, - err: repoerr.ErrNotFound, - }, - { - desc: "retrieve group by id and user with invalid domain ID", - groupID: items[0].ID, - userID: userID, - domainID: invalidID, - err: repoerr.ErrNotFound, - }, - { - desc: "retrieve group by id and user with empty domain ID", - groupID: items[0].ID, - userID: userID, - domainID: "", - err: repoerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - group, err := repo.RetrieveByIDAndUser(context.Background(), tc.domainID, tc.userID, tc.groupID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - if err == nil { - group.Actions = nil - group.Level = 1 - group.AccessProviderRoleActions = nil - assert.Equal(t, tc.resp, group, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, group)) - } - }) - } -} - -func TestRetrieveAll(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM groups") - require.Nil(t, err, fmt.Sprintf("clean groups unexpected error: %s", err)) - }) - - repo := postgres.New(database) - num := 200 - baseTime := time.Now().UTC().Truncate(time.Millisecond) - - var items []groups.Group - parentID := "" - for i := 0; i < num; i++ { - name := namegen.Generate() - group := groups.Group{ - ID: testsutil.GenerateUUID(t), - Domain: testsutil.GenerateUUID(t), - Parent: parentID, - Name: name, - Description: desc, - Metadata: map[string]any{"name": name}, - CreatedAt: baseTime.Add(time.Duration(i) * time.Millisecond), - UpdatedAt: baseTime.Add(time.Duration(i) * time.Millisecond), - Status: groups.EnabledStatus, - Tags: []string{"tag1", "tag2"}, - } - if i%99 == 0 { - group.Tags = []string{"tag1", "tag3"} - } - _, err := repo.Save(context.Background(), group) - require.Nil(t, err, fmt.Sprintf("create group unexpected error: %s", err)) - items = append(items, group) - if i%20 == 0 { - parentID = group.ID - } - } - - reversedGroups := []groups.Group{} - for i := len(items) - 1; i >= 0; i-- { - reversedGroups = append(reversedGroups, items[i]) - } - - cases := []struct { - desc string - page groups.Page - response groups.Page - err error - }{ - { - desc: "retrieve groups successfully", - page: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: 0, - Limit: 10, - Order: "created_at", - Dir: ascDir, - }, - }, - response: groups.Page{ - PageMeta: groups.PageMeta{ - Total: uint64(num), - Offset: 0, - Limit: 10, - }, - Groups: items[:10], - }, - err: nil, - }, - { - desc: "retrieve groups with offset", - page: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: 10, - Limit: 10, - Order: "created_at", - Dir: ascDir, - }, - }, - response: groups.Page{ - PageMeta: groups.PageMeta{ - Total: uint64(num), - Offset: 10, - Limit: 10, - }, - Groups: items[10:20], - }, - err: nil, - }, - { - desc: "retrieve groups with limit", - page: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: 0, - Limit: 50, - - Order: "created_at", - Dir: ascDir, - }, - }, - response: groups.Page{ - PageMeta: groups.PageMeta{ - Total: uint64(num), - Offset: 0, - Limit: 50, - }, - Groups: items[:50], - }, - err: nil, - }, - { - desc: "retrieve groups with offset and limit", - page: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: 50, - Limit: 50, - Order: "created_at", - Dir: ascDir, - }, - }, - response: groups.Page{ - PageMeta: groups.PageMeta{ - Total: uint64(num), - Offset: 50, - Limit: 50, - }, - Groups: items[50:100], - }, - err: nil, - }, - { - desc: "retrieve groups with offset out of range", - page: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: 1000, - Limit: 50, - Order: "created_at", - Dir: ascDir, - }, - }, - response: groups.Page{ - PageMeta: groups.PageMeta{ - Total: uint64(num), - Offset: 1000, - Limit: 50, - }, - Groups: []groups.Group(nil), - }, - err: nil, - }, - { - desc: "retrieve groups with offset and limit out of range", - page: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: 170, - Limit: 50, - Order: "created_at", - Dir: ascDir, - }, - }, - response: groups.Page{ - PageMeta: groups.PageMeta{ - Total: uint64(num), - Offset: 170, - Limit: 50, - }, - Groups: items[170:200], - }, - err: nil, - }, - { - desc: "retrieve groups with limit out of range", - page: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: 0, - Limit: 1000, - }, - }, - response: groups.Page{ - PageMeta: groups.PageMeta{ - Total: uint64(num), - Offset: 0, - Limit: 1000, - }, - Groups: items, - }, - err: nil, - }, - { - desc: "retrieve groups with empty page", - page: groups.Page{}, - response: groups.Page{ - PageMeta: groups.PageMeta{ - Total: uint64(num), - Offset: 0, - Limit: 0, - }, - Groups: []groups.Group(nil), - }, - err: nil, - }, - { - desc: "retrieve groups with name", - page: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: 0, - Limit: 10, - Name: items[0].Name, - }, - }, - response: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 1, - Offset: 0, - Limit: 10, - }, - Groups: []groups.Group{items[0]}, - }, - err: nil, - }, - { - desc: "retrieve groups with domain", - page: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: 0, - Limit: 10, - DomainID: items[0].Domain, - }, - }, - response: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 1, - Offset: 0, - Limit: 10, - }, - Groups: []groups.Group{items[0]}, - }, - err: nil, - }, - { - desc: "retrieve groups with metadata", - page: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: 0, - Limit: 10, - Metadata: items[0].Metadata, - }, - }, - response: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 1, - Offset: 0, - Limit: 10, - }, - Groups: []groups.Group{items[0]}, - }, - err: nil, - }, - { - desc: "retrieve groups with invalid metadata", - page: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: 0, - Limit: 10, - Metadata: map[string]any{ - "key": make(chan int), - }, - }, - }, - response: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 0, - Offset: 0, - Limit: 10, - }, - Groups: []groups.Group(nil), - }, - err: errors.ErrMalformedEntity, - }, - { - desc: "retrieve groups with id", - page: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: 0, - Limit: 10, - ID: items[0].ID, - }, - }, - response: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 1, - Offset: 0, - Limit: 10, - }, - Groups: []groups.Group{items[0]}, - }, - err: nil, - }, - { - desc: "retrieve groups with wrong id", - page: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: 0, - Limit: 10, - ID: "wrong", - }, - }, - response: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 0, - Offset: 0, - Limit: 10, - }, - Groups: []groups.Group(nil), - }, - err: nil, - }, - { - desc: "retrieve groups with order by name ascending", - page: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: 0, - Limit: 10, - Order: "name", - Dir: ascDir, - }, - }, - response: groups.Page{ - PageMeta: groups.PageMeta{ - Total: uint64(num), - Offset: 0, - Limit: 10, - }, - }, - err: nil, - }, - { - desc: "retrieve groups with order by name descending", - page: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: 0, - Limit: 10, - Order: "name", - Dir: descDir, - }, - }, - response: groups.Page{ - PageMeta: groups.PageMeta{ - Total: uint64(num), - Offset: 0, - Limit: 10, - }, - }, - err: nil, - }, - { - desc: "retrieve groups with order by created_at ascending", - page: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: 0, - Limit: 10, - Order: "created_at", - Dir: ascDir, - }, - }, - response: groups.Page{ - PageMeta: groups.PageMeta{ - Total: uint64(num), - Offset: 0, - Limit: 10, - }, - Groups: items[:10], - }, - err: nil, - }, - { - desc: "retrieve groups with order by created_at descending", - page: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: 0, - Limit: 10, - Order: "created_at", - Dir: descDir, - }, - }, - response: groups.Page{ - PageMeta: groups.PageMeta{ - Total: uint64(num), - Offset: 0, - Limit: 10, - }, - Groups: reversedGroups[:10], - }, - err: nil, - }, - { - desc: "retrieve groups with order by updated_at ascending", - page: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: 0, - Limit: 10, - Order: "updated_at", - Dir: ascDir, - }, - }, - response: groups.Page{ - PageMeta: groups.PageMeta{ - Total: uint64(num), - Offset: 0, - Limit: 10, - }, - Groups: items[:10], - }, - err: nil, - }, - { - desc: "retrieve groups with order by updated_at descending", - page: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: 0, - Limit: 10, - Order: "updated_at", - Dir: descDir, - }, - }, - response: groups.Page{ - PageMeta: groups.PageMeta{ - Total: uint64(num), - Offset: 0, - Limit: 10, - }, - Groups: reversedGroups[:10], - }, - err: nil, - }, - { - desc: "retrieve groups with single tag", - page: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: 0, - Limit: uint64(num), - Tags: groups.TagsQuery{Elements: []string{"tag1"}, Operator: groups.OrOp}, - Status: groups.AllStatus, - }, - }, - response: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 200, - Offset: 0, - Limit: uint64(num), - }, - Groups: items, - }, - err: nil, - }, - { - desc: "retrieve group with multiple tags and OR operator", - page: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: 0, - Limit: uint64(num), - Tags: groups.TagsQuery{Elements: []string{"tag2", "tag3"}, Operator: groups.OrOp}, - Status: groups.AllStatus, - }, - }, - response: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 200, - Offset: 0, - Limit: uint64(num), - }, - Groups: items, - }, - }, - { - desc: "retrieve group with multiple tags and AND operator", - page: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: 0, - Limit: uint64(num), - Tags: groups.TagsQuery{Elements: []string{"tag1", "tag3"}, Operator: groups.AndOp}, - Status: groups.AllStatus, - }, - }, - response: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 3, - Offset: 0, - Limit: uint64(num), - }, - Groups: []groups.Group{items[0], items[99], items[198]}, - }, - }, - { - desc: "retrieve group with invalid tags", - page: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: 0, - Limit: uint64(num), - Tags: groups.TagsQuery{Elements: []string{namegen.Generate(), namegen.Generate()}, Operator: groups.OrOp}, - Status: groups.AllStatus, - }, - }, - response: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 0, - Offset: 0, - Limit: uint64(num), - }, - Groups: []groups.Group(nil), - }, - }, - { - desc: "retrieve groups with created_from", - page: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: 0, - Limit: 200, - Order: "created_at", - Dir: ascDir, - CreatedFrom: baseTime.Add(100 * time.Millisecond), - }, - }, - response: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 100, - Offset: 0, - Limit: 200, - }, - Groups: items[100:], - }, - err: nil, - }, - { - desc: "retrieve groups with created_to", - page: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: 0, - Limit: 200, - Order: "created_at", - Dir: ascDir, - CreatedTo: baseTime.Add(99 * time.Millisecond), - }, - }, - response: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 100, - Offset: 0, - Limit: 200, - }, - Groups: items[:100], - }, - err: nil, - }, - { - desc: "retrieve groups with both created_from and created_to", - page: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: 0, - Limit: 200, - Order: "created_at", - Dir: ascDir, - CreatedFrom: baseTime.Add(50 * time.Millisecond), - CreatedTo: baseTime.Add(149 * time.Millisecond), - }, - }, - response: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 100, - Offset: 0, - Limit: 200, - }, - Groups: items[50:150], - }, - err: nil, - }, - { - desc: "retrieve groups with created_from returning no results", - page: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: 0, - Limit: 10, - Order: "created_at", - Dir: ascDir, - CreatedFrom: baseTime.Add(1000 * time.Millisecond), - }, - }, - response: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 0, - Offset: 0, - Limit: 10, - }, - Groups: []groups.Group(nil), - }, - err: nil, - }, - { - desc: "retrieve groups with created_to returning no results", - page: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: 0, - Limit: 10, - Order: "created_at", - Dir: ascDir, - CreatedTo: baseTime.Add(-1 * time.Millisecond), - }, - }, - response: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 0, - Offset: 0, - Limit: 10, - }, - Groups: []groups.Group(nil), - }, - err: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - switch groups, err := repo.RetrieveAll(context.Background(), tc.page.PageMeta); { - case err == nil: - assert.Nil(t, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.response.Total, groups.Total, fmt.Sprintf("%s: expected %d got %d\n", tc.desc, tc.response.Total, groups.Total)) - assert.Equal(t, tc.response.Limit, groups.Limit, fmt.Sprintf("%s: expected %d got %d\n", tc.desc, tc.response.Limit, groups.Limit)) - assert.Equal(t, tc.response.Offset, groups.Offset, fmt.Sprintf("%s: expected %d got %d\n", tc.desc, tc.response.Offset, groups.Offset)) - got := stripGroupDetails(groups.Groups) - if len(tc.response.Groups) > 0 { - resp := stripGroupDetails(tc.response.Groups) - assert.ElementsMatch(t, resp, got, fmt.Sprintf("%s: expected %+v got %+v\n", tc.desc, resp, got)) - } - verifyGroupsOrdering(t, groups.Groups, tc.page.PageMeta.Order, tc.page.PageMeta.Dir) - default: - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - } - }) - } -} - -func TestRetrieveByIDs(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM groups") - require.Nil(t, err, fmt.Sprintf("clean groups unexpected error: %s", err)) - }) - - repo := postgres.New(database) - num := 200 - - var items []groups.Group - parentID := "" - for i := 0; i < num; i++ { - name := namegen.Generate() - group := groups.Group{ - ID: testsutil.GenerateUUID(t), - Domain: testsutil.GenerateUUID(t), - Parent: parentID, - Name: name, - Description: desc, - Metadata: map[string]any{"name": name}, - CreatedAt: time.Now().UTC().Truncate(time.Microsecond), - Status: groups.EnabledStatus, - } - _, err := repo.Save(context.Background(), group) - require.Nil(t, err, fmt.Sprintf("create invitation unexpected error: %s", err)) - items = append(items, group) - if i%20 == 0 { - parentID = group.ID - } - } - - cases := []struct { - desc string - page groups.Page - ids []string - response groups.Page - err error - }{ - { - desc: "retrieve groups successfully", - page: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: 0, - Limit: 10, - }, - }, - ids: getIDs(items[0:3]), - response: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 3, - Offset: 0, - Limit: 10, - }, - Groups: items[0:3], - }, - err: nil, - }, - { - desc: "retrieve groups with empty ids", - page: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: 0, - Limit: 10, - }, - }, - ids: []string{}, - response: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: 0, - Limit: 10, - }, - Groups: []groups.Group(nil), - }, - err: nil, - }, - { - desc: "retrieve groups with empty ids but with domain", - page: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: 0, - Limit: 10, - DomainID: items[0].Domain, - }, - }, - ids: []string{}, - response: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 1, - Offset: 0, - Limit: 10, - }, - Groups: []groups.Group{items[0]}, - }, - err: nil, - }, - { - desc: "retrieve groups with offset", - page: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: 10, - Limit: 10, - }, - }, - ids: getIDs(items[0:20]), - response: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 20, - Offset: 10, - Limit: 10, - }, - Groups: items[10:20], - }, - err: nil, - }, - { - desc: "retrieve groups with offset out of range", - page: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: 1000, - Limit: 50, - }, - }, - ids: getIDs(items[0:20]), - response: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 20, - Offset: 1000, - Limit: 50, - }, - Groups: []groups.Group(nil), - }, - err: nil, - }, - { - desc: "retrieve groups with offset and limit out of range", - page: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: 15, - Limit: 10, - }, - }, - ids: getIDs(items[0:20]), - response: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 20, - Offset: 15, - Limit: 10, - }, - Groups: items[15:20], - }, - err: nil, - }, - { - desc: "retrieve groups with limit out of range", - page: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: 0, - Limit: 1000, - }, - }, - ids: getIDs(items[0:20]), - response: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 20, - Offset: 0, - Limit: 1000, - }, - Groups: items[:20], - }, - err: nil, - }, - { - desc: "retrieve groups with name", - page: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: 0, - Limit: 10, - Name: items[0].Name, - }, - }, - ids: getIDs(items[0:20]), - response: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 1, - Offset: 0, - Limit: 10, - }, - Groups: []groups.Group{items[0]}, - }, - err: nil, - }, - { - desc: "retrieve groups with domain", - page: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: 0, - Limit: 10, - DomainID: items[0].Domain, - }, - }, - ids: getIDs(items[0:20]), - response: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 1, - Offset: 0, - Limit: 10, - }, - Groups: []groups.Group{items[0]}, - }, - err: nil, - }, - { - desc: "retrieve groups with metadata", - page: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: 0, - Limit: 10, - Metadata: items[0].Metadata, - }, - }, - ids: getIDs(items[0:20]), - response: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 1, - Offset: 0, - Limit: 10, - }, - Groups: []groups.Group{items[0]}, - }, - err: nil, - }, - { - desc: "retrieve groups with invalid metadata", - page: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: 0, - Limit: 10, - Metadata: map[string]any{ - "key": make(chan int), - }, - }, - }, - ids: getIDs(items[0:20]), - response: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 0, - Offset: 0, - Limit: 10, - }, - Groups: []groups.Group(nil), - }, - err: errors.ErrMalformedEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - groups, err := repo.RetrieveByIDs(context.Background(), tc.page.PageMeta, tc.ids...) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - if err == nil { - assert.Equal(t, tc.response.Total, groups.Total, fmt.Sprintf("%s: expected %d got %d\n", tc.desc, tc.response.Total, groups.Total)) - assert.Equal(t, tc.response.Limit, groups.Limit, fmt.Sprintf("%s: expected %d got %d\n", tc.desc, tc.response.Limit, groups.Limit)) - assert.Equal(t, tc.response.Offset, groups.Offset, fmt.Sprintf("%s: expected %d got %d\n", tc.desc, tc.response.Offset, groups.Offset)) - got := stripGroupDetails(groups.Groups) - resp := stripGroupDetails(tc.response.Groups) - assert.ElementsMatch(t, resp, got, fmt.Sprintf("%s: expected %+v got %+v\n", tc.desc, resp, got)) - } - }) - } -} - -func TestDelete(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM groups") - require.Nil(t, err, fmt.Sprintf("clean groups unexpected error: %s", err)) - }) - - repo := postgres.New(database) - - group, err := repo.Save(context.Background(), validGroup) - require.Nil(t, err, fmt.Sprintf("save group unexpected error: %s", err)) - - cases := []struct { - desc string - id string - err error - }{ - { - desc: "delete group successfully", - id: group.ID, - err: nil, - }, - { - desc: "delete group with invalid ID", - id: invalidID, - err: repoerr.ErrNotFound, - }, - { - desc: "delete group with empty ID", - id: "", - err: repoerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := repo.Delete(context.Background(), tc.id) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - }) - } -} - -func TestAssignParentGroup(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM groups") - require.Nil(t, err, fmt.Sprintf("clean groups unexpected error: %s", err)) - }) - - repo := postgres.New(database) - - num := 10 - - var items []groups.Group - for i := 0; i < num; i++ { - name := namegen.Generate() - group := groups.Group{ - ID: testsutil.GenerateUUID(t), - Domain: testsutil.GenerateUUID(t), - Name: name, - Description: desc, - Metadata: map[string]any{"name": name}, - CreatedAt: validTimestamp, - Status: groups.EnabledStatus, - } - _, err := repo.Save(context.Background(), group) - require.Nil(t, err, fmt.Sprintf("create invitation unexpected error: %s", err)) - items = append(items, group) - } - - cases := []struct { - desc string - id string - ids []string - err error - }{ - { - desc: "assign parent group successfully", - id: items[0].ID, - ids: []string{items[1].ID, items[2].ID, items[3].ID, items[4].ID, items[5].ID}, - err: nil, - }, - { - desc: "assign parent group with invalid ID", - id: testsutil.GenerateUUID(t), - ids: []string{items[1].ID, items[2].ID, items[3].ID, items[4].ID, items[5].ID}, - err: repoerr.ErrUpdateEntity, - }, - { - desc: "assign parent group with empty ID", - id: "", - ids: []string{items[1].ID, items[2].ID, items[3].ID, items[4].ID, items[5].ID}, - err: repoerr.ErrUpdateEntity, - }, - { - desc: "assign parent group with invalid group IDs", - id: items[0].ID, - ids: []string{testsutil.GenerateUUID(t), testsutil.GenerateUUID(t), testsutil.GenerateUUID(t), testsutil.GenerateUUID(t), testsutil.GenerateUUID(t)}, - err: nil, - }, - { - desc: "assign parent group with empty group IDs", - id: items[0].ID, - ids: []string{}, - err: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := repo.AssignParentGroup(context.Background(), tc.id, tc.ids...) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - }) - } -} - -func TestUnassignParentGroup(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM groups") - require.Nil(t, err, fmt.Sprintf("clean groups unexpected error: %s", err)) - }) - - repo := postgres.New(database) - - num := 10 - - var items []groups.Group - parentID := "" - for i := 0; i < num; i++ { - name := namegen.Generate() - group := groups.Group{ - ID: testsutil.GenerateUUID(t), - Domain: testsutil.GenerateUUID(t), - Parent: parentID, - Name: name, - Description: desc, - Metadata: map[string]any{"name": name}, - CreatedAt: time.Now().UTC().Truncate(time.Microsecond), - Status: groups.EnabledStatus, - } - _, err := repo.Save(context.Background(), group) - require.Nil(t, err, fmt.Sprintf("create invitation unexpected error: %s", err)) - items = append(items, group) - if i == 0 { - parentID = group.ID - } - } - - cases := []struct { - desc string - id string - ids []string - err error - }{ - { - desc: "un-assign parent group successfully", - id: items[0].ID, - ids: []string{items[1].ID, items[2].ID, items[3].ID, items[4].ID, items[5].ID}, - err: nil, - }, - { - desc: "un-assign parent group with invalid ID", - id: testsutil.GenerateUUID(t), - ids: []string{items[1].ID, items[2].ID, items[3].ID, items[4].ID, items[5].ID}, - err: repoerr.ErrUpdateEntity, - }, - { - desc: "un-assign parent group with empty ID", - id: "", - ids: []string{items[1].ID, items[2].ID, items[3].ID, items[4].ID, items[5].ID}, - err: repoerr.ErrUpdateEntity, - }, - { - desc: "un-assign parent group with invalid group IDs", - id: items[0].ID, - ids: []string{testsutil.GenerateUUID(t), testsutil.GenerateUUID(t), testsutil.GenerateUUID(t), testsutil.GenerateUUID(t), testsutil.GenerateUUID(t)}, - err: nil, - }, - { - desc: "un-assign parent group with empty group IDs", - id: items[0].ID, - ids: []string{}, - err: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := repo.UnassignParentGroup(context.Background(), tc.id, tc.ids...) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - }) - } -} - -func TestUnassignAllChildrenGroups(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM groups") - require.Nil(t, err, fmt.Sprintf("clean groups unexpected error: %s", err)) - }) - - repo := postgres.New(database) - - num := 10 - - var items []groups.Group - parentID := "" - for i := 0; i < num; i++ { - name := namegen.Generate() - group := groups.Group{ - ID: testsutil.GenerateUUID(t), - Domain: testsutil.GenerateUUID(t), - Parent: parentID, - Name: name, - Description: desc, - Metadata: map[string]any{"name": name}, - CreatedAt: time.Now().UTC().Truncate(time.Microsecond), - Status: groups.EnabledStatus, - } - _, err := repo.Save(context.Background(), group) - require.Nil(t, err, fmt.Sprintf("create invitation unexpected error: %s", err)) - items = append(items, group) - if i == 0 { - parentID = group.ID - } - } - - cases := []struct { - desc string - id string - err error - }{ - { - desc: "un-assign all children groups successfully", - id: items[0].ID, - err: nil, - }, - { - desc: "un-assign all children groups with invalid ID", - id: testsutil.GenerateUUID(t), - err: repoerr.ErrNotFound, - }, - { - desc: "un-assign all children groups with empty ID", - id: "", - err: repoerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := repo.UnassignAllChildrenGroups(context.Background(), tc.id) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - }) - } -} - -func TestRetrieveHierarchy(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM groups") - require.Nil(t, err, fmt.Sprintf("clean groups unexpected error: %s", err)) - }) - - repo := postgres.New(database) - - userID := testsutil.GenerateUUID(t) - domainID := testsutil.GenerateUUID(t) - num := 10 - - var items []groups.Group - parentID := "" - for i := 0; i < num; i++ { - name := namegen.Generate() - group := groups.Group{ - ID: testsutil.GenerateUUID(t), - Domain: domainID, - Parent: parentID, - Name: name, - Description: desc, - Metadata: map[string]any{"name": name}, - CreatedAt: time.Now().UTC().Truncate(time.Microsecond), - Status: groups.EnabledStatus, - } - _, err := repo.Save(context.Background(), group) - require.Nil(t, err, fmt.Sprintf("create group unexpected error: %s", err)) - newRolesProvision := []roles.RoleProvision{ - { - Role: roles.Role{ - ID: testsutil.GenerateUUID(t) + "_" + group.ID, - Name: "admin", - EntityID: group.ID, - CreatedAt: validTimestamp, - CreatedBy: userID, - }, - OptionalActions: availableActions, - OptionalMembers: []string{userID}, - }, - } - _, err = repo.AddRoles(context.Background(), newRolesProvision) - require.Nil(t, err, fmt.Sprintf("add roles unexpected error: %s", err)) - items = append(items, group) - if i == 0 { - parentID = group.ID - } - } - - cases := []struct { - desc string - groupID string - userID string - domainID string - hm groups.HierarchyPageMeta - resp groups.HierarchyPage - err error - }{ - { - desc: "retrieve ancestors successfully", - groupID: items[1].ID, - userID: userID, - domainID: domainID, - hm: groups.HierarchyPageMeta{ - Level: 1, - Direction: +1, - Tree: false, - }, - resp: groups.HierarchyPage{ - Groups: []groups.Group{items[0], items[1]}, - HierarchyPageMeta: groups.HierarchyPageMeta{ - Level: 1, - Direction: +1, - Tree: false, - }, - }, - err: nil, - }, - { - desc: "retrieve descendants successfully", - groupID: items[0].ID, - userID: userID, - domainID: domainID, - hm: groups.HierarchyPageMeta{ - Level: 1, - Direction: -1, - Tree: false, - }, - resp: groups.HierarchyPage{ - Groups: items, - HierarchyPageMeta: groups.HierarchyPageMeta{ - Level: 1, - Direction: -1, - Tree: false, - }, - }, - err: nil, - }, - { - desc: "retrieve hierarchy with invalid ID", - groupID: testsutil.GenerateUUID(t), - userID: userID, - domainID: domainID, - err: nil, - }, - { - desc: "retrieve hierarchy with empty ID", - groupID: "", - userID: userID, - domainID: domainID, - err: nil, - }, - { - desc: "retrieve hierarchy with invalid domain ID", - groupID: items[0].ID, - userID: userID, - domainID: testsutil.GenerateUUID(t), - hm: groups.HierarchyPageMeta{ - Level: 1, - Direction: -1, - Tree: false, - }, - resp: groups.HierarchyPage{ - Groups: []groups.Group(nil), - HierarchyPageMeta: groups.HierarchyPageMeta{ - Level: 1, - Direction: -1, - Tree: false, - }, - }, - err: nil, - }, - { - desc: "retrieve hierarchy with invalid user ID", - groupID: items[0].ID, - userID: testsutil.GenerateUUID(t), - domainID: domainID, - hm: groups.HierarchyPageMeta{ - Level: 1, - Direction: -1, - Tree: false, - }, - resp: groups.HierarchyPage{ - Groups: []groups.Group(nil), - HierarchyPageMeta: groups.HierarchyPageMeta{ - Level: 1, - Direction: -1, - Tree: false, - }, - }, - err: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - gpPage, err := repo.RetrieveHierarchy(context.Background(), tc.domainID, tc.userID, tc.groupID, tc.hm) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - if err == nil { - got := stripGroupDetails(gpPage.Groups) - resp := stripGroupDetails(tc.resp.Groups) - assert.ElementsMatch(t, resp, got, fmt.Sprintf("%s: expected %+v got %+v\n", tc.desc, resp, got)) - } - }) - } -} - -func TestRetrieveAllParentGroups(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM groups") - require.Nil(t, err, fmt.Sprintf("clean groups unexpected error: %s", err)) - }) - - repo := postgres.New(database) - - parentID := "" - domainID := testsutil.GenerateUUID(t) - userID := testsutil.GenerateUUID(t) - num := 10 - halfindex := num/2 + 1 - items := []groups.Group{} - for i := 0; i < num; i++ { - name := namegen.Generate() - group := groups.Group{ - ID: testsutil.GenerateUUID(t), - Domain: domainID, - Name: name, - Parent: parentID, - Description: desc, - Metadata: map[string]any{"name": name}, - CreatedAt: validTimestamp, - Status: groups.EnabledStatus, - } - grp, err := repo.Save(context.Background(), group) - require.Nil(t, err, fmt.Sprintf("create group unexpected error: %s", err)) - parentID = grp.ID - newRolesProvision := []roles.RoleProvision{ - { - Role: roles.Role{ - ID: testsutil.GenerateUUID(t) + "_" + grp.ID, - Name: "admin", - EntityID: grp.ID, - CreatedAt: validTimestamp, - CreatedBy: userID, - }, - OptionalActions: availableActions, - OptionalMembers: []string{userID}, - }, - } - _, err = repo.AddRoles(context.Background(), newRolesProvision) - require.Nil(t, err, fmt.Sprintf("add roles unexpected error: %s", err)) - ngrp := grp - ngrp.RoleID = newRolesProvision[0].Role.ID - ngrp.RoleName = newRolesProvision[0].Role.Name - ngrp.AccessType = directAccess - items = append(items, ngrp) - } - - cases := []struct { - desc string - id string - domainID string - userID string - pageMeta groups.PageMeta - resp groups.Page - err error - }{ - { - desc: "retrieve all parent groups successfully", - id: items[num-1].ID, - domainID: domainID, - userID: userID, - pageMeta: groups.PageMeta{ - Offset: 0, - Limit: 20, - }, - resp: groups.Page{ - PageMeta: groups.PageMeta{ - Total: uint64(num), - }, - Groups: items, - }, - err: nil, - }, - { - desc: "retrieve half of all parent groups successfully", - id: items[num/2].ID, - domainID: domainID, - userID: userID, - pageMeta: groups.PageMeta{ - Offset: 0, - Limit: 20, - }, - resp: groups.Page{ - PageMeta: groups.PageMeta{ - Total: uint64(halfindex), - }, - Groups: items[:halfindex], - }, - err: nil, - }, - { - desc: "retrieve all parent groups with invalid group ID", - id: testsutil.GenerateUUID(t), - domainID: domainID, - userID: userID, - pageMeta: groups.PageMeta{ - Offset: 0, - Limit: 20, - }, - resp: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 0, - }, - Groups: []groups.Group(nil), - }, - err: repoerr.ErrNotFound, - }, - { - desc: "retrieve all parent groups with empty group ID", - id: "", - domainID: domainID, - userID: userID, - pageMeta: groups.PageMeta{ - Offset: 0, - Limit: 20, - }, - resp: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 0, - }, - Groups: []groups.Group(nil), - }, - err: repoerr.ErrNotFound, - }, - { - desc: "retrieve all parent groups with invalid domain ID", - id: items[num-1].ID, - domainID: testsutil.GenerateUUID(t), - userID: userID, - pageMeta: groups.PageMeta{ - Offset: 0, - Limit: 20, - }, - resp: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 0, - }, - Groups: []groups.Group(nil), - }, - err: nil, - }, - { - desc: "retrieve all parent groups with invalid user ID", - id: items[num-1].ID, - domainID: domainID, - userID: testsutil.GenerateUUID(t), - pageMeta: groups.PageMeta{ - Offset: 0, - Limit: 20, - }, - resp: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 0, - }, - Groups: []groups.Group(nil), - }, - err: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - groups, err := repo.RetrieveAllParentGroups(context.Background(), tc.domainID, tc.userID, tc.id, tc.pageMeta) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - if err == nil { - assert.Equal(t, tc.resp.Total, groups.Total, fmt.Sprintf("%s: expected %d got %d\n", tc.desc, tc.resp.Total, groups.Total)) - got := stripGroupDetails(groups.Groups) - resp := stripGroupDetails(tc.resp.Groups) - assert.ElementsMatch(t, resp, got, fmt.Sprintf("%s: expected %+v got %+v\n", tc.desc, resp, got)) - } - }) - } -} - -func TestRetrieveChildrenGroups(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM groups") - require.Nil(t, err, fmt.Sprintf("clean groups unexpected error: %s", err)) - }) - - repo := postgres.New(database) - - parentID := "" - domainID := testsutil.GenerateUUID(t) - userID := testsutil.GenerateUUID(t) - num := 10 - items := []groups.Group{} - for i := 0; i < num; i++ { - name := namegen.Generate() - group := groups.Group{ - ID: testsutil.GenerateUUID(t), - Domain: domainID, - Name: name, - Parent: parentID, - Description: desc, - Metadata: map[string]any{"name": name}, - CreatedAt: validTimestamp, - Status: groups.EnabledStatus, - } - grp, err := repo.Save(context.Background(), group) - require.Nil(t, err, fmt.Sprintf("create group unexpected error: %s", err)) - parentID = grp.ID - newRolesProvision := []roles.RoleProvision{ - { - Role: roles.Role{ - ID: testsutil.GenerateUUID(t) + "_" + grp.ID, - Name: "admin", - EntityID: grp.ID, - CreatedAt: validTimestamp, - CreatedBy: userID, - }, - OptionalActions: availableActions, - OptionalMembers: []string{userID}, - }, - } - _, err = repo.AddRoles(context.Background(), newRolesProvision) - require.Nil(t, err, fmt.Sprintf("add roles unexpected error: %s", err)) - ngrp := grp - ngrp.RoleID = newRolesProvision[0].Role.ID - ngrp.RoleName = newRolesProvision[0].Role.Name - ngrp.AccessType = directAccess - items = append(items, ngrp) - } - - cases := []struct { - desc string - id string - domainID string - userID string - startLevel int64 - endLevel int64 - pageMeta groups.PageMeta - resp groups.Page - err error - }{ - { - desc: "retrieve children groups from parent group level successfully", - id: items[0].ID, - domainID: domainID, - userID: userID, - startLevel: 0, - endLevel: -1, - pageMeta: groups.PageMeta{ - Offset: 0, - Limit: 20, - }, - resp: groups.Page{ - PageMeta: groups.PageMeta{ - Total: uint64(num), - }, - Groups: items, - }, - err: nil, - }, - { - desc: "Retrieve specific level of children groups from parent group level", - id: items[0].ID, - domainID: domainID, - userID: userID, - startLevel: 1, - endLevel: 1, - pageMeta: groups.PageMeta{ - Offset: 0, - Limit: 20, - }, - resp: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 1, - }, - Groups: []groups.Group{items[1]}, - }, - err: nil, - }, - { - desc: "Retrieve all children groups from specific level from parent group level", - id: items[0].ID, - pageMeta: groups.PageMeta{ - Offset: 0, - Limit: 20, - }, - domainID: domainID, - userID: userID, - startLevel: 2, - endLevel: -1, - resp: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 8, - }, - Groups: items[2:], - }, - err: nil, - }, - { - desc: "Retrieve all children groups from specific level to specific level from parent group level", - id: items[0].ID, - pageMeta: groups.PageMeta{ - Offset: 0, - Limit: 20, - }, - domainID: domainID, - userID: userID, - startLevel: 1, - endLevel: 2, - resp: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 2, - }, - Groups: items[1:3], - }, - err: nil, - }, - { - desc: "Retrieve all children groups with invalid group ID", - id: testsutil.GenerateUUID(t), - domainID: domainID, - userID: userID, - startLevel: 0, - endLevel: -1, - pageMeta: groups.PageMeta{ - Offset: 0, - Limit: 20, - }, - resp: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 0, - }, - Groups: []groups.Group(nil), - }, - err: repoerr.ErrNotFound, - }, - { - desc: "Retrieve all children groups with empty group ID", - id: "", - domainID: domainID, - userID: userID, - startLevel: 0, - endLevel: -1, - pageMeta: groups.PageMeta{ - Offset: 0, - Limit: 20, - }, - resp: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 0, - }, - Groups: []groups.Group(nil), - }, - err: repoerr.ErrNotFound, - }, - { - desc: "Retrieve all children groups with invalid domain ID", - id: items[0].ID, - domainID: testsutil.GenerateUUID(t), - userID: userID, - startLevel: 0, - endLevel: -1, - pageMeta: groups.PageMeta{ - Offset: 0, - Limit: 20, - }, - resp: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 0, - }, - Groups: []groups.Group(nil), - }, - err: nil, - }, - { - desc: "Retrieve all children groups with invalid user ID", - id: items[0].ID, - domainID: domainID, - userID: testsutil.GenerateUUID(t), - startLevel: 0, - endLevel: -1, - pageMeta: groups.PageMeta{ - Offset: 0, - Limit: 20, - }, - resp: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 0, - }, - Groups: []groups.Group(nil), - }, - err: nil, - }, - { - desc: "Retrieve all children groups with invalid start level", - id: items[0].ID, - domainID: domainID, - userID: userID, - startLevel: -1, - endLevel: -1, - pageMeta: groups.PageMeta{ - Offset: 0, - Limit: 20, - }, - resp: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 0, - }, - Groups: []groups.Group(nil), - }, - err: repoerr.ErrViewEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - groups, err := repo.RetrieveChildrenGroups(context.Background(), tc.domainID, tc.userID, tc.id, tc.startLevel, tc.endLevel, tc.pageMeta) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - if err == nil { - assert.Equal(t, tc.resp.Total, groups.Total, fmt.Sprintf("%s: expected %d got %d\n", tc.desc, tc.resp.Total, groups.Total)) - got := stripGroupDetails(groups.Groups) - resp := stripGroupDetails(tc.resp.Groups) - assert.ElementsMatch(t, resp, got, fmt.Sprintf("%s: expected %+v got %+v\n", tc.desc, resp, got)) - } - }) - } -} - -func getIDs(groups []groups.Group) []string { - var ids []string - for _, group := range groups { - ids = append(ids, group.ID) - } - - return ids -} - -func stripGroupDetails(groups []groups.Group) []groups.Group { - for i := range groups { - groups[i].Level = 0 - groups[i].Path = "" - groups[i].CreatedAt = validTimestamp - groups[i].UpdatedAt = validTimestamp - groups[i].Actions = nil - groups[i].AccessProviderRoleActions = nil - } - - return groups -} - -func verifyGroupsOrdering(t *testing.T, groups []groups.Group, order, dir string) { - if order == "" || len(groups) <= 1 { - return - } - - for i := 0; i < len(groups)-1; i++ { - switch order { - case "name": - if dir == ascDir { - assert.LessOrEqual(t, groups[i].Name, groups[i+1].Name, fmt.Sprintf("Groups not ordered by name ascending at index %d: %s > %s", i, groups[i].Name, groups[i+1].Name)) - continue - } - assert.GreaterOrEqual(t, groups[i].Name, groups[i+1].Name, fmt.Sprintf("Groups not ordered by name descending at index %d: %s < %s", i, groups[i].Name, groups[i+1].Name)) - case "created_at": - if dir == ascDir { - assert.False(t, groups[i].CreatedAt.After(groups[i+1].CreatedAt), fmt.Sprintf("Groups not ordered by created_at ascending at index %d: %v > %v", i, groups[i].CreatedAt, groups[i+1].CreatedAt)) - continue - } - assert.False(t, groups[i].CreatedAt.Before(groups[i+1].CreatedAt), fmt.Sprintf("Groups not ordered by created_at descending at index %d: %v < %v", i, groups[i].CreatedAt, groups[i+1].CreatedAt)) - case "updated_at": - if dir == ascDir { - assert.False(t, groups[i].UpdatedAt.After(groups[i+1].UpdatedAt), fmt.Sprintf("Groups not ordered by updated_at ascending at index %d: %v > %v", i, groups[i].UpdatedAt, groups[i+1].UpdatedAt)) - continue - } - assert.False(t, groups[i].UpdatedAt.Before(groups[i+1].UpdatedAt), fmt.Sprintf("Groups not ordered by updated_at descending at index %d: %v < %v", i, groups[i].UpdatedAt, groups[i+1].UpdatedAt)) - } - } -} diff --git a/groups/postgres/init.go b/groups/postgres/init.go deleted file mode 100644 index f275c5bd0..000000000 --- a/groups/postgres/init.go +++ /dev/null @@ -1,119 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package postgres - -import ( - dpostgres "github.com/absmach/magistrala/domains/postgres" - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - rolesPostgres "github.com/absmach/magistrala/pkg/roles/repo/postgres" - _ "github.com/jackc/pgx/v5/stdlib" // required for SQL access - migrate "github.com/rubenv/sql-migrate" -) - -func Migration() (*migrate.MemoryMigrationSource, error) { - rolesMigration, err := rolesPostgres.Migration(rolesTableNamePrefix, entityTableName, entityIDColumnName) - if err != nil { - return &migrate.MemoryMigrationSource{}, errors.Wrap(repoerr.ErrRoleMigration, err) - } - - groupsMigration := &migrate.MemoryMigrationSource{ - Migrations: []*migrate.Migration{ - { - Id: "groups_01", - Up: []string{ - `CREATE TABLE IF NOT EXISTS groups ( - id VARCHAR(36) PRIMARY KEY, - parent_id VARCHAR(36), - domain_id VARCHAR(36) NOT NULL, - name VARCHAR(1024) NOT NULL, - description VARCHAR(1024), - metadata JSONB, - created_at TIMESTAMP, - updated_at TIMESTAMP, - updated_by VARCHAR(254), - status SMALLINT NOT NULL DEFAULT 0 CHECK (status >= 0), - UNIQUE (domain_id, name), - FOREIGN KEY (parent_id) REFERENCES groups (id) ON DELETE SET NULL, - CHECK (id != parent_id) - )`, - }, - Down: []string{ - `DROP TABLE IF EXISTS groups`, - }, - }, - { - Id: "groups_02", - Up: []string{ - `CREATE EXTENSION IF NOT EXISTS LTREE`, - `ALTER TABLE groups ADD COLUMN IF NOT EXISTS path LTREE`, - `CREATE INDEX IF NOT EXISTS path_gist_idx ON groups USING GIST (path);`, - }, - Down: []string{ - `DROP TABLE IF EXISTS groups`, - `DROP EXTENSION IF EXISTS LTREE`, - }, - }, - { - Id: "groups_03", - Up: []string{ - `ALTER TABLE groups DROP CONSTRAINT IF EXISTS groups_domain_id_name_key`, - }, - Down: []string{ - `ALTER TABLE groups ADD CONSTRAINT groups_domain_id_name_key UNIQUE (domain_id, name)`, - }, - }, - { - Id: "groups_04", - Up: []string{ - `ALTER TABLE groups ADD COLUMN IF NOT EXISTS tags TEXT[]`, - }, - Down: []string{ - `ALTER TABLE groups DROP COLUMN tags`, - }, - }, - { - Id: "groups_05", - Up: []string{ - `ALTER TABLE groups ALTER COLUMN created_at TYPE TIMESTAMPTZ;`, - `ALTER TABLE groups ALTER COLUMN updated_at TYPE TIMESTAMPTZ;`, - }, - Down: []string{ - `ALTER TABLE groups ALTER COLUMN created_at TYPE TIMESTAMP;`, - `ALTER TABLE groups ALTER COLUMN updated_at TYPE TIMESTAMP;`, - }, - }, - { - Id: "groups_06", - Up: []string{ - `UPDATE groups - SET metadata = (COALESCE(metadata, '{}'::jsonb) || COALESCE(metadata->'ui', '{}'::jsonb)) - 'ui' - WHERE metadata ? 'ui' AND jsonb_typeof(metadata->'ui') = 'object'`, - }, - Down: []string{ - `SELECT 1`, - }, - }, - { - Id: "groups_07", - Up: []string{ - `CREATE INDEX IF NOT EXISTS idx_groups_domain_id_status ON groups(domain_id, status);`, - }, - Down: []string{ - `DROP INDEX IF EXISTS idx_groups_domain_id_status;`, - }, - }, - }, - } - - groupsMigration.Migrations = append(groupsMigration.Migrations, rolesMigration.Migrations...) - - domainsMigrations, err := dpostgres.Migration() - if err != nil { - return nil, err - } - groupsMigration.Migrations = append(groupsMigration.Migrations, domainsMigrations.Migrations...) - - return groupsMigration, nil -} diff --git a/groups/postgres/setup_test.go b/groups/postgres/setup_test.go deleted file mode 100644 index e243bb4e4..000000000 --- a/groups/postgres/setup_test.go +++ /dev/null @@ -1,98 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package postgres_test - -import ( - "database/sql" - "fmt" - "log" - "os" - "testing" - "time" - - gpostgres "github.com/absmach/magistrala/groups/postgres" - "github.com/absmach/magistrala/pkg/postgres" - pgclient "github.com/absmach/magistrala/pkg/postgres" - "github.com/jmoiron/sqlx" - "github.com/ory/dockertest/v3" - "github.com/ory/dockertest/v3/docker" - "go.opentelemetry.io/otel" -) - -var ( - db *sqlx.DB - database postgres.Database - tracer = otel.Tracer("repo_tests") -) - -func TestMain(m *testing.M) { - pool, err := dockertest.NewPool("") - if err != nil { - log.Fatalf("Could not connect to docker: %s", err) - } - - container, err := pool.RunWithOptions(&dockertest.RunOptions{ - Repository: "postgres", - Tag: "16.2-alpine", - Env: []string{ - "POSTGRES_USER=test", - "POSTGRES_PASSWORD=test", - "POSTGRES_DB=test", - "listen_addresses = '*'", - }, - }, func(config *docker.HostConfig) { - config.AutoRemove = true - config.RestartPolicy = docker.RestartPolicy{Name: "no"} - }) - if err != nil { - log.Fatalf("Could not start container: %s", err) - } - - port := container.GetPort("5432/tcp") - - // exponential backoff-retry, because the application in the container might not be ready to accept connections yet - pool.MaxWait = 120 * time.Second - if err := pool.Retry(func() error { - url := fmt.Sprintf("host=localhost port=%s user=test dbname=test password=test sslmode=disable", port) - db, err := sql.Open("pgx", url) - if err != nil { - return err - } - return db.Ping() - }); err != nil { - log.Fatalf("Could not connect to docker: %s", err) - } - - dbConfig := pgclient.Config{ - Host: "localhost", - Port: port, - User: "test", - Pass: "test", - Name: "test", - SSLMode: "disable", - SSLCert: "", - SSLKey: "", - SSLRootCert: "", - } - - gmig, err := gpostgres.Migration() - if err != nil { - log.Fatalf("Could not get groups migration : %s", err) - } - if db, err = pgclient.Setup(dbConfig, *gmig); err != nil { - log.Fatalf("Could not setup test DB connection: %s", err) - } - - database = postgres.NewDatabase(db, dbConfig, tracer) - - code := m.Run() - - // Defers will not be run when using os.Exit - db.Close() - if err := pool.Purge(container); err != nil { - log.Fatalf("Could not purge container: %s", err) - } - - os.Exit(code) -} diff --git a/groups/private/mocks/service.go b/groups/private/mocks/service.go deleted file mode 100644 index ab6fa9315..000000000 --- a/groups/private/mocks/service.go +++ /dev/null @@ -1,109 +0,0 @@ -// Copyright (c) Abstract Machines - -// SPDX-License-Identifier: Apache-2.0 - -// Code generated by mockery; DO NOT EDIT. -// github.com/vektra/mockery -// template: testify - -package mocks - -import ( - "context" - - "github.com/absmach/magistrala/groups" - mock "github.com/stretchr/testify/mock" -) - -// NewService creates a new instance of Service. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. -// The first argument is typically a *testing.T value. -func NewService(t interface { - mock.TestingT - Cleanup(func()) -}) *Service { - mock := &Service{} - mock.Mock.Test(t) - - t.Cleanup(func() { mock.AssertExpectations(t) }) - - return mock -} - -// Service is an autogenerated mock type for the Service type -type Service struct { - mock.Mock -} - -type Service_Expecter struct { - mock *mock.Mock -} - -func (_m *Service) EXPECT() *Service_Expecter { - return &Service_Expecter{mock: &_m.Mock} -} - -// RetrieveById provides a mock function for the type Service -func (_mock *Service) RetrieveById(ctx context.Context, id string) (groups.Group, error) { - ret := _mock.Called(ctx, id) - - if len(ret) == 0 { - panic("no return value specified for RetrieveById") - } - - var r0 groups.Group - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) (groups.Group, error)); ok { - return returnFunc(ctx, id) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string) groups.Group); ok { - r0 = returnFunc(ctx, id) - } else { - r0 = ret.Get(0).(groups.Group) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = returnFunc(ctx, id) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RetrieveById_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveById' -type Service_RetrieveById_Call struct { - *mock.Call -} - -// RetrieveById is a helper method to define mock.On call -// - ctx context.Context -// - id string -func (_e *Service_Expecter) RetrieveById(ctx interface{}, id interface{}) *Service_RetrieveById_Call { - return &Service_RetrieveById_Call{Call: _e.mock.On("RetrieveById", ctx, id)} -} - -func (_c *Service_RetrieveById_Call) Run(run func(ctx context.Context, id string)) *Service_RetrieveById_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Service_RetrieveById_Call) Return(group groups.Group, err error) *Service_RetrieveById_Call { - _c.Call.Return(group, err) - return _c -} - -func (_c *Service_RetrieveById_Call) RunAndReturn(run func(ctx context.Context, id string) (groups.Group, error)) *Service_RetrieveById_Call { - _c.Call.Return(run) - return _c -} diff --git a/groups/private/service.go b/groups/private/service.go deleted file mode 100644 index 3c190faf8..000000000 --- a/groups/private/service.go +++ /dev/null @@ -1,28 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package private - -import ( - "context" - - "github.com/absmach/magistrala/groups" -) - -type Service interface { - RetrieveById(ctx context.Context, id string) (groups.Group, error) -} - -var _ Service = (*service)(nil) - -func New(repo groups.Repository) Service { - return service{repo} -} - -type service struct { - repo groups.Repository -} - -func (svc service) RetrieveById(ctx context.Context, ids string) (groups.Group, error) { - return svc.repo.RetrieveByID(ctx, ids) -} diff --git a/groups/service.go b/groups/service.go deleted file mode 100644 index 9c05fdb6a..000000000 --- a/groups/service.go +++ /dev/null @@ -1,465 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package groups - -import ( - "context" - "fmt" - "time" - - "github.com/absmach/magistrala" - grpcChannelsV1 "github.com/absmach/magistrala/api/grpc/channels/v1" - grpcClientsV1 "github.com/absmach/magistrala/api/grpc/clients/v1" - apiutil "github.com/absmach/magistrala/api/http/util" - smqauthn "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/pkg/policies" - "github.com/absmach/magistrala/pkg/roles" -) - -var ( - ErrGroupIDs = errors.New("invalid group ids") - errChangeGroupStatus = errors.NewServiceError("failed to change group status") - errGroupHaveParent = errors.NewRequestError("group already have parent") - errDifferentParent = errors.NewRequestError("groups have different parent") -) - -type service struct { - repo Repository - policy policies.Service - idProvider magistrala.IDProvider - channels grpcChannelsV1.ChannelsServiceClient - clients grpcClientsV1.ClientsServiceClient - - roles.ProvisionManageService -} - -// NewService returns a new groups service implementation. -func NewService(repo Repository, policy policies.Service, idp magistrala.IDProvider, channels grpcChannelsV1.ChannelsServiceClient, clients grpcClientsV1.ClientsServiceClient, sidProvider magistrala.IDProvider, availableActions []roles.Action, builtInRoles map[roles.BuiltInRoleName][]roles.Action) (Service, error) { - rpms, err := roles.NewProvisionManageService(policies.GroupType, repo, policy, sidProvider, availableActions, builtInRoles) - if err != nil { - return service{}, err - } - return service{ - repo: repo, - policy: policy, - idProvider: idp, - channels: channels, - clients: clients, - ProvisionManageService: rpms, - }, nil -} - -func (svc service) CreateGroup(ctx context.Context, session smqauthn.Session, g Group) (retGr Group, retRps []roles.RoleProvision, retErr error) { - groupID, err := svc.idProvider.ID() - if err != nil { - return Group{}, []roles.RoleProvision{}, errors.Wrap(svcerr.ErrCreateEntity, err) - } - if g.Status != EnabledStatus && g.Status != DisabledStatus { - return Group{}, []roles.RoleProvision{}, svcerr.ErrInvalidStatus - } - - g.ID = groupID - g.CreatedAt = time.Now().UTC() - g.Domain = session.DomainID - - saved, err := svc.repo.Save(ctx, g) - if err != nil { - return Group{}, []roles.RoleProvision{}, errors.Wrap(svcerr.ErrCreateEntity, err) - } - - defer func() { - if retErr != nil { - if errRollback := svc.repo.Delete(ctx, saved.ID); errRollback != nil { - retErr = errors.Wrap(retErr, errors.Wrap(apiutil.ErrRollbackTx, errRollback)) - } - } - }() - - oprs := []policies.Policy{} - - oprs = append(oprs, policies.Policy{ - Domain: session.DomainID, - SubjectType: policies.DomainType, - Subject: session.DomainID, - Relation: policies.DomainRelation, - ObjectType: policies.GroupType, - Object: saved.ID, - }) - if saved.Parent != "" { - oprs = append(oprs, policies.Policy{ - Domain: session.DomainID, - SubjectType: policies.GroupType, - Subject: saved.Parent, - Relation: policies.ParentGroupRelation, - ObjectType: policies.GroupType, - ObjectKind: policies.NewGroupKind, - Object: saved.ID, - }) - } - newBuiltInRoleMembers := map[roles.BuiltInRoleName][]roles.Member{ - BuiltInRoleAdmin: {roles.Member(session.UserID)}, - } - rp, err := svc.AddNewEntitiesRoles(ctx, session.DomainID, session.UserID, []string{saved.ID}, oprs, newBuiltInRoleMembers) - if err != nil { - return Group{}, []roles.RoleProvision{}, errors.Wrap(svcerr.ErrAddPolicies, err) - } - - return saved, rp, nil -} - -func (svc service) ViewGroup(ctx context.Context, session smqauthn.Session, id string, withRoles bool) (Group, error) { - var group Group - var err error - switch withRoles { - case true: - group, err = svc.repo.RetrieveByIDWithRoles(ctx, id, session.UserID) - default: - group, err = svc.repo.RetrieveByID(ctx, id) - } - if err != nil { - return Group{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - - return group, nil -} - -func (svc service) ListGroups(ctx context.Context, session smqauthn.Session, gm PageMeta) (Page, error) { - switch session.SuperAdmin { - case true: - gm.DomainID = session.DomainID - page, err := svc.repo.RetrieveAll(ctx, gm) - if err != nil { - return Page{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - return page, nil - default: - page, err := svc.repo.RetrieveUserGroups(ctx, session.DomainID, session.UserID, gm) - if err != nil { - return Page{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - return page, nil - } -} - -func (svc service) ListUserGroups(ctx context.Context, session smqauthn.Session, userID string, pm PageMeta) (Page, error) { - page, err := svc.repo.RetrieveUserGroups(ctx, session.DomainID, userID, pm) - if err != nil { - return Page{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - return page, nil -} - -func (svc service) UpdateGroup(ctx context.Context, session smqauthn.Session, g Group) (Group, error) { - g.UpdatedAt = time.Now().UTC() - g.UpdatedBy = session.UserID - - group, err := svc.repo.Update(ctx, g) - if err != nil { - return Group{}, errors.Wrap(svcerr.ErrUpdateEntity, err) - } - return group, nil -} - -func (svc service) UpdateGroupTags(ctx context.Context, session smqauthn.Session, g Group) (Group, error) { - group := Group{ - ID: g.ID, - Tags: g.Tags, - UpdatedAt: time.Now(), - UpdatedBy: session.UserID, - } - group, err := svc.repo.UpdateTags(ctx, group) - if err != nil { - return Group{}, errors.Wrap(svcerr.ErrUpdateEntity, err) - } - return group, nil -} - -func (svc service) EnableGroup(ctx context.Context, session smqauthn.Session, id string) (Group, error) { - group := Group{ - ID: id, - Status: EnabledStatus, - UpdatedAt: time.Now().UTC(), - } - group, err := svc.changeGroupStatus(ctx, session, group) - if err != nil { - return Group{}, errors.Wrap(errChangeGroupStatus, err) - } - return group, nil -} - -func (svc service) DisableGroup(ctx context.Context, session smqauthn.Session, id string) (Group, error) { - group := Group{ - ID: id, - Status: DisabledStatus, - UpdatedAt: time.Now().UTC(), - } - group, err := svc.changeGroupStatus(ctx, session, group) - if err != nil { - return Group{}, errors.Wrap(errChangeGroupStatus, err) - } - return group, nil -} - -func (svc service) RetrieveGroupHierarchy(ctx context.Context, session smqauthn.Session, id string, hm HierarchyPageMeta) (HierarchyPage, error) { - hp, err := svc.repo.RetrieveHierarchy(ctx, session.DomainID, session.UserID, id, hm) - if err != nil { - return HierarchyPage{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - return hp, nil -} - -func (svc service) AddParentGroup(ctx context.Context, session smqauthn.Session, id, parentID string) (retErr error) { - group, err := svc.repo.RetrieveByID(ctx, id) - if err != nil { - return errors.Wrap(svcerr.ErrViewEntity, err) - } - - var pols []policies.Policy - if group.Parent != "" { - return errors.Wrap(svcerr.ErrConflict, fmt.Errorf("%s group already have parent", group.ID)) - } - pols = append(pols, policies.Policy{ - Domain: session.DomainID, - SubjectType: policies.GroupType, - Subject: parentID, - Relation: policies.ParentGroupRelation, - ObjectType: policies.GroupType, - Object: group.ID, - }) - - if err := svc.policy.AddPolicies(ctx, pols); err != nil { - return errors.Wrap(svcerr.ErrAddPolicies, err) - } - defer func() { - if retErr != nil { - if errRollback := svc.policy.DeletePolicies(ctx, pols); errRollback != nil { - retErr = errors.Wrap(retErr, errors.Wrap(apiutil.ErrRollbackTx, errRollback)) - } - } - }() - - if err := svc.repo.AssignParentGroup(ctx, parentID, group.ID); err != nil { - return err - } - return nil -} - -func (svc service) RemoveParentGroup(ctx context.Context, session smqauthn.Session, id string) (retErr error) { - group, err := svc.repo.RetrieveByID(ctx, id) - if err != nil { - return errors.Wrap(svcerr.ErrViewEntity, err) - } - - if group.Parent != "" { - var pols []policies.Policy - pols = append(pols, policies.Policy{ - Domain: session.DomainID, - SubjectType: policies.GroupType, - Subject: group.Parent, - Relation: policies.ParentGroupRelation, - ObjectType: policies.GroupType, - Object: group.ID, - }) - - if err := svc.policy.DeletePolicies(ctx, pols); err != nil { - return errors.Wrap(svcerr.ErrDeletePolicies, err) - } - defer func() { - if retErr != nil { - if errRollback := svc.policy.AddPolicies(ctx, pols); errRollback != nil { - retErr = errors.Wrap(retErr, errors.Wrap(apiutil.ErrRollbackTx, errRollback)) - } - } - }() - if err := svc.repo.UnassignParentGroup(ctx, group.Parent, group.ID); err != nil { - return errors.Wrap(svcerr.ErrRemoveEntity, err) - } - - return nil - } - - return nil -} - -func (svc service) AddChildrenGroups(ctx context.Context, session smqauthn.Session, parentGroupID string, childrenGroupIDs []string) (retErr error) { - childrenGroupsPage, err := svc.repo.RetrieveByIDs(ctx, PageMeta{Limit: 1<<63 - 1}, childrenGroupIDs...) - if err != nil { - return errors.Wrap(svcerr.ErrViewEntity, err) - } - if len(childrenGroupsPage.Groups) == 0 { - return ErrGroupIDs - } - - for _, childGroup := range childrenGroupsPage.Groups { - if childGroup.Parent != "" { - return errors.Wrap(svcerr.ErrConflict, errGroupHaveParent) - } - } - - var pols []policies.Policy - for _, childGroup := range childrenGroupsPage.Groups { - pols = append(pols, policies.Policy{ - Domain: session.DomainID, - SubjectType: policies.GroupType, - Subject: parentGroupID, - Relation: policies.ParentGroupRelation, - ObjectType: policies.GroupType, - Object: childGroup.ID, - }) - } - - if err := svc.policy.AddPolicies(ctx, pols); err != nil { - return errors.Wrap(svcerr.ErrAddPolicies, err) - } - defer func() { - if retErr != nil { - if errRollback := svc.policy.DeletePolicies(ctx, pols); errRollback != nil { - retErr = errors.Wrap(retErr, errors.Wrap(apiutil.ErrRollbackTx, errRollback)) - } - } - }() - if err = svc.repo.AssignParentGroup(ctx, parentGroupID, childrenGroupIDs...); err != nil { - return errors.Wrap(svcerr.ErrUpdateEntity, err) - } - - return nil -} - -func (svc service) RemoveChildrenGroups(ctx context.Context, session smqauthn.Session, parentGroupID string, childrenGroupIDs []string) (retErr error) { - childrenGroupsPage, err := svc.repo.RetrieveByIDs(ctx, PageMeta{Limit: 1<<63 - 1}, childrenGroupIDs...) - if err != nil { - return errors.Wrap(svcerr.ErrViewEntity, err) - } - if len(childrenGroupsPage.Groups) == 0 { - return ErrGroupIDs - } - - var pols []policies.Policy - - for _, group := range childrenGroupsPage.Groups { - if group.Parent != "" && group.Parent != parentGroupID { - return errors.Wrap(svcerr.ErrConflict, errDifferentParent) - } - pols = append(pols, policies.Policy{ - Domain: session.DomainID, - SubjectType: policies.GroupType, - Subject: parentGroupID, - Relation: policies.ParentGroupRelation, - ObjectType: policies.GroupType, - Object: group.ID, - }) - } - - if err := svc.policy.DeletePolicies(ctx, pols); err != nil { - return errors.Wrap(svcerr.ErrDeletePolicies, err) - } - defer func() { - if retErr != nil { - if errRollback := svc.policy.AddPolicies(ctx, pols); errRollback != nil { - retErr = errors.Wrap(retErr, errors.Wrap(apiutil.ErrRollbackTx, errRollback)) - } - } - }() - if err := svc.repo.UnassignParentGroup(ctx, parentGroupID, childrenGroupIDs...); err != nil { - return errors.Wrap(svcerr.ErrUpdateEntity, err) - } - - return nil -} - -func (svc service) RemoveAllChildrenGroups(ctx context.Context, session smqauthn.Session, id string) error { - pol := policies.Policy{ - Domain: session.DomainID, - SubjectType: policies.GroupType, - Subject: id, - Relation: policies.ParentGroupRelation, - ObjectType: policies.GroupType, - } - - if err := svc.policy.DeletePolicyFilter(ctx, pol); err != nil { - return errors.Wrap(svcerr.ErrDeletePolicies, err) - } - if err := svc.repo.UnassignAllChildrenGroups(ctx, id); err != nil { - return errors.Wrap(svcerr.ErrRemoveEntity, err) - } - - return nil -} - -func (svc service) ListChildrenGroups(ctx context.Context, session smqauthn.Session, id string, startLevel, endLevel int64, pm PageMeta) (Page, error) { - page, err := svc.repo.RetrieveChildrenGroups(ctx, session.DomainID, session.UserID, id, startLevel, endLevel, pm) - if err != nil { - return Page{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - return page, nil -} - -func (svc service) DeleteGroup(ctx context.Context, session smqauthn.Session, id string) error { - if _, err := svc.channels.UnsetParentGroupFromChannels(ctx, &grpcChannelsV1.UnsetParentGroupFromChannelsReq{ParentGroupId: id}); err != nil { - return errors.Wrap(svcerr.ErrRemoveEntity, err) - } - - if _, err := svc.clients.UnsetParentGroupFromClient(ctx, &grpcClientsV1.UnsetParentGroupFromClientReq{ParentGroupId: id}); err != nil { - return errors.Wrap(svcerr.ErrRemoveEntity, err) - } - - g, err := svc.repo.ChangeStatus(ctx, Group{ID: id, Status: DeletedStatus}) - if err != nil { - return errors.Wrap(svcerr.ErrRemoveEntity, err) - } - - filterDeletePolicies := []policies.Policy{ - { - SubjectType: policies.GroupType, - Subject: id, - }, - { - ObjectType: policies.GroupType, - Object: id, - }, - } - deletePolicies := []policies.Policy{ - { - SubjectType: policies.DomainType, - Subject: session.DomainID, - Relation: policies.DomainRelation, - ObjectType: policies.GroupType, - Object: id, - }, - } - if g.Parent != "" { - deletePolicies = append(deletePolicies, policies.Policy{ - Domain: session.DomainID, - SubjectType: policies.GroupType, - Subject: g.Parent, - Relation: policies.ParentGroupRelation, - ObjectType: policies.GroupType, - Object: id, - }) - } - if err := svc.RemoveEntitiesRoles(ctx, session.DomainID, session.DomainUserID, []string{id}, filterDeletePolicies, deletePolicies); err != nil { - return errors.Wrap(svcerr.ErrDeletePolicies, err) - } - - if err := svc.repo.Delete(ctx, id); err != nil { - return errors.Wrap(svcerr.ErrRemoveEntity, err) - } - - return nil -} - -func (svc service) changeGroupStatus(ctx context.Context, session smqauthn.Session, group Group) (Group, error) { - dbGroup, err := svc.repo.RetrieveByID(ctx, group.ID) - if err != nil { - return Group{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - if dbGroup.Status == group.Status { - return Group{}, svcerr.ErrStatusAlreadyAssigned - } - - group.UpdatedBy = session.UserID - return svc.repo.ChangeStatus(ctx, group) -} diff --git a/groups/service_test.go b/groups/service_test.go deleted file mode 100644 index 4a9d73ad7..000000000 --- a/groups/service_test.go +++ /dev/null @@ -1,1324 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package groups_test - -import ( - "context" - "fmt" - "testing" - "time" - - "github.com/0x6flab/namegenerator" - grpcChannelsV1 "github.com/absmach/magistrala/api/grpc/channels/v1" - grpcClientsV1 "github.com/absmach/magistrala/api/grpc/clients/v1" - apiutil "github.com/absmach/magistrala/api/http/util" - chmocks "github.com/absmach/magistrala/channels/mocks" - climocks "github.com/absmach/magistrala/clients/mocks" - "github.com/absmach/magistrala/groups" - "github.com/absmach/magistrala/groups/mocks" - "github.com/absmach/magistrala/internal/nullable" - "github.com/absmach/magistrala/internal/testsutil" - "github.com/absmach/magistrala/pkg/authn" - smqauthn "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - policysvc "github.com/absmach/magistrala/pkg/policies" - policymocks "github.com/absmach/magistrala/pkg/policies/mocks" - "github.com/absmach/magistrala/pkg/roles" - "github.com/absmach/magistrala/pkg/uuid" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" -) - -var ( - idProvider = uuid.New() - namegen = namegenerator.NewGenerator() - description = namegen.Generate() - desc = nullable.Value[string]{Valid: true, Value: description} - validGroup = groups.Group{ - ID: testsutil.GenerateUUID(&testing.T{}), - Name: namegen.Generate(), - Description: desc, - Metadata: map[string]any{ - "key": "value", - }, - Status: groups.EnabledStatus, - } - validGroupWithRoles = groups.Group{ - ID: testsutil.GenerateUUID(&testing.T{}), - Name: namegen.Generate(), - Description: desc, - Metadata: map[string]any{ - "key": "value", - }, - Status: groups.EnabledStatus, - Roles: []roles.MemberRoleActions{ - { - RoleID: "test-id", - RoleName: "test-name", - AccessType: "direct", - }, - }, - } - parentGroupID = testsutil.GenerateUUID(&testing.T{}) - childGroupID = testsutil.GenerateUUID(&testing.T{}) - childGroup = groups.Group{ - ID: childGroupID, - Name: namegen.Generate(), - Description: desc, - Metadata: map[string]any{ - "key": "value", - }, - Status: groups.EnabledStatus, - Parent: parentGroupID, - } - children = []*groups.Group{&childGroup} - parentGroup = groups.Group{ - ID: parentGroupID, - Name: namegen.Generate(), - Description: desc, - Metadata: map[string]any{ - "key": "value", - }, - Status: groups.EnabledStatus, - Children: children, - } - validID = testsutil.GenerateUUID(&testing.T{}) - validSession = authn.Session{UserID: validID, DomainID: validID, DomainUserID: validID} -) - -var ( - repo *mocks.Repository - policies *policymocks.Service - channels *chmocks.ChannelsServiceClient - clients *climocks.ClientsServiceClient -) - -func newService(t *testing.T) groups.Service { - repo = new(mocks.Repository) - policies = new(policymocks.Service) - channels = new(chmocks.ChannelsServiceClient) - clients = new(climocks.ClientsServiceClient) - availableActions := []roles.Action{} - builtInRoles := map[roles.BuiltInRoleName][]roles.Action{ - groups.BuiltInRoleAdmin: availableActions, - } - svc, err := groups.NewService(repo, policies, idProvider, channels, clients, idProvider, availableActions, builtInRoles) - assert.Nil(t, err, fmt.Sprintf(" Unexpected error while creating service %v", err)) - return svc -} - -func TestCreateGroup(t *testing.T) { - svc := newService(t) - - cases := []struct { - desc string - group groups.Group - saveResp groups.Group - saveErr error - deleteErr error - addPoliciesErr error - deletePoliciesErr error - addRoleErr error - err error - }{ - { - desc: "create group successfully", - group: validGroup, - saveResp: groups.Group{ - ID: testsutil.GenerateUUID(t), - CreatedAt: time.Now(), - Domain: validID, - }, - err: nil, - }, - { - desc: "create group with invalid status", - group: groups.Group{ - Name: namegen.Generate(), - Description: desc, - Status: groups.Status(100), - }, - err: svcerr.ErrInvalidStatus, - }, - { - desc: "create group successfully with parent", - group: groups.Group{ - Name: namegen.Generate(), - Description: desc, - Status: groups.EnabledStatus, - Parent: testsutil.GenerateUUID(t), - }, - saveResp: groups.Group{ - ID: testsutil.GenerateUUID(t), - CreatedAt: time.Now(), - Domain: testsutil.GenerateUUID(t), - Parent: testsutil.GenerateUUID(t), - }, - err: nil, - }, - { - desc: "create group with failed to save", - group: validGroup, - saveResp: groups.Group{}, - saveErr: errors.ErrMalformedEntity, - err: errors.ErrMalformedEntity, - }, - { - desc: " create group with failed to add policies", - group: validGroup, - saveResp: groups.Group{ - ID: testsutil.GenerateUUID(t), - CreatedAt: time.Now(), - Domain: validID, - }, - addPoliciesErr: svcerr.ErrAuthorization, - err: svcerr.ErrAddPolicies, - }, - { - desc: " create group with failed to add policies and failed rollback", - group: validGroup, - saveResp: groups.Group{ - ID: testsutil.GenerateUUID(t), - CreatedAt: time.Now(), - Domain: validID, - }, - addPoliciesErr: svcerr.ErrAuthorization, - deleteErr: svcerr.ErrRemoveEntity, - err: svcerr.ErrRemoveEntity, - }, - { - desc: "create group with failed to add roles", - group: validGroup, - saveResp: groups.Group{ - ID: testsutil.GenerateUUID(t), - CreatedAt: time.Now(), - Domain: validID, - }, - addRoleErr: svcerr.ErrCreateEntity, - err: svcerr.ErrAddPolicies, - }, - { - desc: "create groups with failed to add roles and failed to delete policies", - group: validGroup, - saveResp: groups.Group{ - ID: testsutil.GenerateUUID(t), - CreatedAt: time.Now(), - Domain: validID, - }, - addRoleErr: svcerr.ErrCreateEntity, - deletePoliciesErr: svcerr.ErrRemoveEntity, - err: svcerr.ErrAddPolicies, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := repo.On("Save", context.Background(), mock.Anything).Return(tc.saveResp, tc.saveErr) - policyCall := policies.On("AddPolicies", context.Background(), mock.Anything).Return(tc.addPoliciesErr) - policyCall1 := policies.On("DeletePolicies", context.Background(), mock.Anything).Return(tc.deletePoliciesErr) - repoCall1 := repo.On("AddRoles", context.Background(), mock.Anything).Return([]roles.RoleProvision{}, tc.addRoleErr) - repoCall2 := repo.On("Delete", context.Background(), mock.Anything).Return(tc.deleteErr) - got, _, err := svc.CreateGroup(context.Background(), validSession, tc.group) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("expected error %v but got %v", tc.err, err)) - if err == nil { - assert.NotEmpty(t, got.ID) - assert.NotEmpty(t, got.CreatedAt) - assert.NotEmpty(t, got.Domain) - assert.WithinDuration(t, time.Now(), got.CreatedAt, 2*time.Second) - ok := repoCall.Parent.AssertCalled(t, "Save", context.Background(), mock.Anything) - assert.True(t, ok, fmt.Sprintf("Save was not called on %s", tc.desc)) - } - repoCall.Unset() - policyCall.Unset() - policyCall1.Unset() - repoCall1.Unset() - repoCall2.Unset() - }) - } -} - -func TestViewGroup(t *testing.T) { - svc := newService(t) - - cases := []struct { - desc string - session smqauthn.Session - id string - withRoles bool - repoResp groups.Group - repoErr error - err error - }{ - { - desc: "view group successfully", - id: validGroup.ID, - session: validSession, - withRoles: false, - repoResp: validGroup, - }, - { - desc: "view group successfully with roles", - id: validGroupWithRoles.ID, - session: validSession, - withRoles: true, - repoResp: validGroupWithRoles, - }, - { - desc: "view group with failed to retrieve", - id: testsutil.GenerateUUID(t), - session: validSession, - withRoles: false, - repoErr: repoerr.ErrNotFound, - err: svcerr.ErrViewEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := repo.On("RetrieveByID", context.Background(), tc.id).Return(tc.repoResp, tc.repoErr) - repoCall1 := repo.On("RetrieveByIDWithRoles", context.Background(), tc.id, tc.session.UserID).Return(tc.repoResp, tc.repoErr) - got, err := svc.ViewGroup(context.Background(), validSession, tc.id, tc.withRoles) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("expected error %v to contain %v", err, tc.err)) - if err == nil { - switch tc.withRoles { - case true: - assert.Equal(t, tc.repoResp, got) - ok := repo.AssertCalled(t, "RetrieveByIDWithRoles", context.Background(), tc.id, tc.session.UserID) - assert.True(t, ok, fmt.Sprintf("RetrieveByIDWithRoles was not called on %s", tc.desc)) - default: - assert.Equal(t, tc.repoResp, got) - ok := repo.AssertCalled(t, "RetrieveByID", context.Background(), tc.id) - assert.True(t, ok, fmt.Sprintf("RetrieveByID was not called on %s", tc.desc)) - } - } - repoCall.Unset() - repoCall1.Unset() - }) - } -} - -func TestUpdateGroup(t *testing.T) { - svc := newService(t) - - cases := []struct { - desc string - group groups.Group - repoResp groups.Group - repoErr error - err error - }{ - { - desc: "update group successfully", - group: groups.Group{ - ID: testsutil.GenerateUUID(t), - Name: namegen.Generate(), - }, - repoResp: validGroup, - }, - { - desc: "update group with repo error", - group: groups.Group{ - ID: testsutil.GenerateUUID(t), - Name: namegen.Generate(), - }, - repoErr: repoerr.ErrNotFound, - err: svcerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := repo.On("Update", context.Background(), mock.Anything).Return(tc.repoResp, tc.repoErr) - got, err := svc.UpdateGroup(context.Background(), validSession, tc.group) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("expected error %v to contain %v", err, tc.err)) - if err == nil { - assert.Equal(t, tc.repoResp, got) - ok := repo.AssertCalled(t, "Update", context.Background(), mock.Anything) - assert.True(t, ok, fmt.Sprintf("Update was not called on %s", tc.desc)) - } - repoCall.Unset() - }) - } -} - -func TestUpdateGroupTags(t *testing.T) { - svc := newService(t) - - cases := []struct { - desc string - updateReq groups.Group - repoResp groups.Group - repoErr error - err error - }{ - { - desc: "update group tags successfully", - updateReq: groups.Group{ - ID: testsutil.GenerateUUID(t), - Tags: []string{"tag1", "tag2"}, - }, - repoResp: groups.Group{ - ID: testsutil.GenerateUUID(t), - Tags: []string{"tag1", "tag2"}, - }, - }, - { - desc: "update group tags with repo error", - updateReq: groups.Group{ - ID: testsutil.GenerateUUID(t), - Tags: []string{"tag1", "tag2"}, - }, - repoErr: repoerr.ErrNotFound, - err: svcerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := repo.On("UpdateTags", context.Background(), mock.Anything).Return(tc.repoResp, tc.repoErr) - got, err := svc.UpdateGroupTags(context.Background(), validSession, tc.updateReq) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("expected error %v to contain %v", err, tc.err)) - if err == nil { - assert.Equal(t, tc.repoResp, got) - ok := repo.AssertCalled(t, "UpdateTags", context.Background(), mock.Anything) - assert.True(t, ok, fmt.Sprintf("UpdateTags was not called on %s", tc.desc)) - } - repoCall.Unset() - }) - } -} - -func TestEnableGroup(t *testing.T) { - svc := newService(t) - - cases := []struct { - desc string - id string - retrieveResp groups.Group - retrieveErr error - changeResp groups.Group - changeErr error - err error - }{ - { - desc: "enable group successfully", - id: testsutil.GenerateUUID(t), - retrieveResp: groups.Group{ - Status: groups.DisabledStatus, - }, - changeResp: validGroup, - }, - { - desc: "enable group with enabled group", - id: testsutil.GenerateUUID(t), - retrieveResp: groups.Group{ - Status: groups.EnabledStatus, - }, - err: svcerr.ErrStatusAlreadyAssigned, - }, - { - desc: "enable group with retrieve error", - id: testsutil.GenerateUUID(t), - retrieveResp: groups.Group{}, - retrieveErr: repoerr.ErrNotFound, - err: repoerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := repo.On("RetrieveByID", context.Background(), tc.id).Return(tc.retrieveResp, tc.retrieveErr) - repoCall1 := repo.On("ChangeStatus", context.Background(), mock.Anything).Return(tc.changeResp, tc.changeErr) - got, err := svc.EnableGroup(context.Background(), validSession, tc.id) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("expected error %v to contain %v", err, tc.err)) - if err == nil { - assert.Equal(t, tc.changeResp, got) - ok := repo.AssertCalled(t, "RetrieveByID", context.Background(), tc.id) - assert.True(t, ok, fmt.Sprintf("RetrieveByID was not called on %s", tc.desc)) - } - repoCall.Unset() - repoCall1.Unset() - }) - } -} - -func TestDisableGroup(t *testing.T) { - svc := newService(t) - - cases := []struct { - desc string - id string - retrieveResp groups.Group - retrieveErr error - changeResp groups.Group - changeErr error - err error - }{ - { - desc: "disable group successfully", - id: testsutil.GenerateUUID(t), - retrieveResp: groups.Group{ - Status: groups.EnabledStatus, - }, - changeResp: validGroup, - }, - { - desc: "disable group with disabled group", - id: testsutil.GenerateUUID(t), - retrieveResp: groups.Group{ - Status: groups.DisabledStatus, - }, - err: svcerr.ErrStatusAlreadyAssigned, - }, - { - desc: "disable group with retrieve error", - id: testsutil.GenerateUUID(t), - retrieveResp: groups.Group{}, - retrieveErr: repoerr.ErrNotFound, - err: repoerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := repo.On("RetrieveByID", context.Background(), tc.id).Return(tc.retrieveResp, tc.retrieveErr) - repoCall1 := repo.On("ChangeStatus", context.Background(), mock.Anything).Return(tc.changeResp, tc.changeErr) - got, err := svc.DisableGroup(context.Background(), validSession, tc.id) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("expected error %v to contain %v", err, tc.err)) - if err == nil { - assert.Equal(t, tc.changeResp, got) - ok := repo.AssertCalled(t, "RetrieveByID", context.Background(), tc.id) - assert.True(t, ok, fmt.Sprintf("RetrieveByID was not called on %s", tc.desc)) - } - repoCall.Unset() - repoCall1.Unset() - }) - } -} - -func TestListGroups(t *testing.T) { - svc := newService(t) - - cases := []struct { - desc string - session smqauthn.Session - pageMeta groups.PageMeta - retrieveAllRes groups.Page - retrieveAllErr error - retrieveUserGroupRes groups.Page - retrieveUserGroupErr error - resp groups.Page - err error - }{ - { - desc: "list groups as super admin successfully", - session: smqauthn.Session{UserID: validID, DomainID: validID, DomainUserID: validID, SuperAdmin: true}, - pageMeta: groups.PageMeta{ - Limit: 10, - Offset: 0, - DomainID: validID, - }, - retrieveAllRes: groups.Page{ - Groups: []groups.Group{validGroup}, - PageMeta: groups.PageMeta{ - Total: 1, - }, - }, - resp: groups.Page{ - Groups: []groups.Group{validGroup}, - PageMeta: groups.PageMeta{ - Total: 1, - }, - }, - err: nil, - }, - { - desc: "list groups as super admin with failed to retrieve", - session: smqauthn.Session{UserID: validID, DomainID: validID, DomainUserID: validID, SuperAdmin: true}, - pageMeta: groups.PageMeta{ - Limit: 10, - Offset: 0, - DomainID: validID, - }, - retrieveAllErr: repoerr.ErrNotFound, - resp: groups.Page{}, - err: repoerr.ErrNotFound, - }, - { - desc: "list groups as non admin successfully", - session: validSession, - pageMeta: groups.PageMeta{ - Limit: 10, - Offset: 0, - }, - retrieveUserGroupRes: groups.Page{ - Groups: []groups.Group{validGroup}, - PageMeta: groups.PageMeta{ - Total: 1, - }, - }, - resp: groups.Page{ - Groups: []groups.Group{validGroup}, - PageMeta: groups.PageMeta{ - Total: 1, - }, - }, - err: nil, - }, - { - desc: "list groups as non admin with failed to retrieve user groups", - session: validSession, - pageMeta: groups.PageMeta{ - Limit: 10, - Offset: 0, - }, - retrieveUserGroupErr: repoerr.ErrNotFound, - resp: groups.Page{}, - err: svcerr.ErrViewEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := repo.On("RetrieveAll", context.Background(), tc.pageMeta).Return(tc.retrieveAllRes, tc.retrieveAllErr) - repoCall1 := repo.On("RetrieveUserGroups", context.Background(), tc.session.DomainID, tc.session.UserID, tc.pageMeta).Return(tc.retrieveUserGroupRes, tc.retrieveUserGroupErr) - got, err := svc.ListGroups(context.Background(), tc.session, tc.pageMeta) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("expected error %v to contain %v", err, tc.err)) - assert.Equal(t, tc.resp, got) - repoCall.Unset() - repoCall1.Unset() - }) - } -} - -func TestListUserGroups(t *testing.T) { - svc := newService(t) - - cases := []struct { - desc string - session smqauthn.Session - userID string - pageMeta groups.PageMeta - retrieveUserGroupRes groups.Page - retrieveUserGroupErr error - resp groups.Page - err error - }{ - { - desc: "list user groups successfully", - session: validSession, - userID: validID, - pageMeta: groups.PageMeta{ - Limit: 10, - Offset: 0, - }, - retrieveUserGroupRes: groups.Page{ - Groups: []groups.Group{validGroup}, - PageMeta: groups.PageMeta{ - Total: 1, - }, - }, - resp: groups.Page{ - Groups: []groups.Group{validGroup}, - PageMeta: groups.PageMeta{ - Total: 1, - }, - }, - err: nil, - }, - { - desc: "list user groups with failed to retrieve", - session: validSession, - userID: validID, - pageMeta: groups.PageMeta{ - Limit: 10, - Offset: 0, - }, - retrieveUserGroupErr: repoerr.ErrNotFound, - resp: groups.Page{}, - err: svcerr.ErrViewEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := repo.On("RetrieveUserGroups", context.Background(), tc.session.DomainID, tc.userID, tc.pageMeta).Return(tc.retrieveUserGroupRes, tc.retrieveUserGroupErr) - got, err := svc.ListUserGroups(context.Background(), tc.session, tc.userID, tc.pageMeta) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("expected error %v to contain %v", err, tc.err)) - assert.Equal(t, tc.resp, got) - repoCall.Unset() - }) - } -} - -func TestRetrieveGroupHierarchy(t *testing.T) { - svc := newService(t) - - cases := []struct { - desc string - id string - pageMeta groups.HierarchyPageMeta - retrieveHierarchyRes groups.HierarchyPage - retrieveHierarchyErr error - err error - }{ - { - desc: "retrieve group hierarchy successfully", - id: parentGroup.ID, - pageMeta: groups.HierarchyPageMeta{ - Level: 1, - Direction: -1, - Tree: false, - }, - retrieveHierarchyRes: groups.HierarchyPage{ - HierarchyPageMeta: groups.HierarchyPageMeta{ - Level: 1, - Direction: -1, - Tree: false, - }, - Groups: []groups.Group{parentGroup}, - }, - err: nil, - }, - { - desc: "retrieve group hierarchy with failed to retrieve hierarchy", - id: parentGroup.ID, - pageMeta: groups.HierarchyPageMeta{ - Level: 1, - Direction: -1, - Tree: false, - }, - retrieveHierarchyErr: repoerr.ErrNotFound, - err: repoerr.ErrNotFound, - }, - { - desc: "retrieve group hierarchy with invalid group ID", - id: testsutil.GenerateUUID(t), - pageMeta: groups.HierarchyPageMeta{ - Level: 1, - Direction: -1, - Tree: false, - }, - retrieveHierarchyErr: nil, - err: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := repo.On("RetrieveHierarchy", context.Background(), validSession.DomainID, validSession.UserID, tc.id, tc.pageMeta).Return(tc.retrieveHierarchyRes, tc.retrieveHierarchyErr) - _, err := svc.RetrieveGroupHierarchy(context.Background(), validSession, tc.id, tc.pageMeta) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("expected error %v to contain %v", err, tc.err)) - if tc.err == nil { - ok := repo.AssertCalled(t, "RetrieveHierarchy", context.Background(), validSession.DomainID, validSession.UserID, tc.id, tc.pageMeta) - assert.True(t, ok, fmt.Sprintf("RetrieveHierarchy was not called on %s", tc.desc)) - } - repoCall.Unset() - }) - } -} - -func TestAddParentGroup(t *testing.T) { - svc := newService(t) - - cases := []struct { - desc string - id string - parentID string - retrieveResp groups.Group - retrieveErr error - addPoliciesErr error - deletePoliciesErr error - assignParentErr error - err error - }{ - { - desc: "add parent group successfully", - id: validGroup.ID, - parentID: parentGroupID, - retrieveResp: validGroup, - err: nil, - }, - { - desc: "add parent group with failed to retrieve", - id: validGroup.ID, - parentID: parentGroupID, - retrieveErr: repoerr.ErrNotFound, - err: repoerr.ErrNotFound, - }, - { - desc: "add parent group to group with parent", - id: childGroupID, - parentID: parentGroupID, - retrieveResp: childGroup, - err: svcerr.ErrConflict, - }, - { - desc: "add parent group with failed to add policies", - id: validGroup.ID, - parentID: parentGroupID, - retrieveResp: validGroup, - addPoliciesErr: svcerr.ErrAuthorization, - err: svcerr.ErrAddPolicies, - }, - { - desc: "add parent group with repo error in assign parent group", - id: validGroup.ID, - parentID: parentGroupID, - retrieveResp: validGroup, - assignParentErr: repoerr.ErrNotFound, - err: repoerr.ErrNotFound, - }, - { - desc: "add parent group with repo error in assign parent group and failed to delete policies", - id: validGroup.ID, - parentID: parentGroupID, - retrieveResp: validGroup, - assignParentErr: repoerr.ErrNotFound, - deletePoliciesErr: svcerr.ErrAuthorization, - err: apiutil.ErrRollbackTx, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - pol := policysvc.Policy{ - Domain: validID, - SubjectType: policysvc.GroupType, - Subject: tc.parentID, - Relation: policysvc.ParentGroupRelation, - ObjectType: policysvc.GroupType, - Object: tc.id, - } - repoCall := repo.On("RetrieveByID", context.Background(), tc.id).Return(tc.retrieveResp, tc.retrieveErr) - policyCall := policies.On("AddPolicies", context.Background(), []policysvc.Policy{pol}).Return(tc.addPoliciesErr) - policyCall1 := policies.On("DeletePolicies", context.Background(), []policysvc.Policy{pol}).Return(tc.deletePoliciesErr) - repoCall1 := repo.On("AssignParentGroup", context.Background(), tc.parentID, []string{tc.id}).Return(tc.assignParentErr) - err := svc.AddParentGroup(context.Background(), validSession, tc.id, tc.parentID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("expected error %v to contain %v", err, tc.err)) - ok := repo.AssertCalled(t, "RetrieveByID", context.Background(), tc.id) - assert.True(t, ok, fmt.Sprintf("RetrieveByID was not called on %s", tc.desc)) - repoCall.Unset() - policyCall.Unset() - policyCall1.Unset() - repoCall1.Unset() - }) - } -} - -func TestRemoveParentGroup(t *testing.T) { - svc := newService(t) - - cases := []struct { - desc string - id string - retrieveResp groups.Group - retrieveErr error - deletePoliciesErr error - addPoliciesErr error - unassignParentErr error - err error - }{ - { - desc: "remove parent group successfully", - id: childGroupID, - retrieveResp: childGroup, - err: nil, - }, - { - desc: "remove parent group with failed to retrieve", - id: childGroupID, - retrieveErr: repoerr.ErrNotFound, - err: repoerr.ErrNotFound, - }, - { - desc: "remove parent group with no parent", - id: validGroup.ID, - retrieveResp: validGroup, - err: nil, - }, - { - desc: "remove parent group with failed to delete policies", - id: childGroupID, - retrieveResp: childGroup, - deletePoliciesErr: svcerr.ErrAuthorization, - err: svcerr.ErrDeletePolicies, - }, - { - desc: "remove parent group with repo error in unassign parent group", - id: childGroupID, - retrieveResp: childGroup, - unassignParentErr: repoerr.ErrNotFound, - err: repoerr.ErrNotFound, - }, - { - desc: "remove parent group with repo error in unassign parent group and failed to add policies", - id: childGroupID, - retrieveResp: childGroup, - unassignParentErr: repoerr.ErrNotFound, - addPoliciesErr: svcerr.ErrAuthorization, - err: apiutil.ErrRollbackTx, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - pol := policysvc.Policy{ - Domain: validID, - SubjectType: policysvc.GroupType, - Subject: tc.retrieveResp.Parent, - Relation: policysvc.ParentGroupRelation, - ObjectType: policysvc.GroupType, - Object: tc.id, - } - repoCall := repo.On("RetrieveByID", context.Background(), tc.id).Return(tc.retrieveResp, tc.retrieveErr) - policyCall := policies.On("DeletePolicies", context.Background(), []policysvc.Policy{pol}).Return(tc.deletePoliciesErr) - policyCall1 := policies.On("AddPolicies", context.Background(), []policysvc.Policy{pol}).Return(tc.addPoliciesErr) - repoCall1 := repo.On("UnassignParentGroup", context.Background(), tc.retrieveResp.Parent, []string{tc.id}).Return(tc.unassignParentErr) - err := svc.RemoveParentGroup(context.Background(), validSession, tc.id) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("expected error %v to contain %v", err, tc.err)) - ok := repo.AssertCalled(t, "RetrieveByID", context.Background(), tc.id) - assert.True(t, ok, fmt.Sprintf("RetrieveByID was not called on %s", tc.desc)) - repoCall.Unset() - policyCall.Unset() - policyCall1.Unset() - repoCall1.Unset() - }) - } -} - -func TestAddChildrenGroups(t *testing.T) { - svc := newService(t) - - cases := []struct { - desc string - parentID string - childrenIDs []string - retrieveResp groups.Page - retrieveErr error - addPoliciesErr error - deletePoliciesErr error - assignParentErr error - err error - }{ - { - desc: "add children groups successfully", - parentID: parentGroupID, - childrenIDs: []string{validGroup.ID}, - retrieveResp: groups.Page{ - Groups: []groups.Group{validGroup}, - PageMeta: groups.PageMeta{ - Total: 1, - }, - }, - err: nil, - }, - { - desc: "add children groups with failed to retrieve", - parentID: parentGroupID, - childrenIDs: []string{validGroup.ID}, - retrieveErr: repoerr.ErrNotFound, - err: repoerr.ErrNotFound, - }, - { - desc: "add non existent child group", - parentID: parentGroupID, - childrenIDs: []string{testsutil.GenerateUUID(&testing.T{})}, - retrieveResp: groups.Page{}, - err: groups.ErrGroupIDs, - }, - { - desc: "add child group with parent", - parentID: parentGroupID, - childrenIDs: []string{childGroupID}, - retrieveResp: groups.Page{ - Groups: []groups.Group{childGroup}, - PageMeta: groups.PageMeta{ - Total: 1, - }, - }, - err: svcerr.ErrConflict, - }, - { - desc: "add children groups with failed to add policies", - parentID: parentGroupID, - childrenIDs: []string{validGroup.ID}, - retrieveResp: groups.Page{ - Groups: []groups.Group{validGroup}, - PageMeta: groups.PageMeta{ - Total: 1, - }, - }, - addPoliciesErr: svcerr.ErrAuthorization, - err: svcerr.ErrAddPolicies, - }, - { - desc: "add children groups with repo error in assign children groups", - parentID: parentGroupID, - childrenIDs: []string{validGroup.ID}, - retrieveResp: groups.Page{ - Groups: []groups.Group{validGroup}, - PageMeta: groups.PageMeta{ - Total: 1, - }, - }, - assignParentErr: repoerr.ErrNotFound, - err: repoerr.ErrNotFound, - }, - { - desc: "add children groups with repo error in assign children groups and failed to delete policies", - parentID: parentGroupID, - childrenIDs: []string{validGroup.ID}, - retrieveResp: groups.Page{ - Groups: []groups.Group{validGroup}, - PageMeta: groups.PageMeta{ - Total: 1, - }, - }, - assignParentErr: repoerr.ErrNotFound, - deletePoliciesErr: svcerr.ErrAuthorization, - err: apiutil.ErrRollbackTx, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - pol := policysvc.Policy{ - Domain: validID, - SubjectType: policysvc.GroupType, - Subject: tc.parentID, - Relation: policysvc.ParentGroupRelation, - ObjectType: policysvc.GroupType, - Object: validGroup.ID, - } - repoCall := repo.On("RetrieveByIDs", context.Background(), groups.PageMeta{Limit: 1<<63 - 1}, tc.childrenIDs).Return(tc.retrieveResp, tc.retrieveErr) - policyCall := policies.On("AddPolicies", context.Background(), []policysvc.Policy{pol}).Return(tc.addPoliciesErr) - policyCall1 := policies.On("DeletePolicies", context.Background(), []policysvc.Policy{pol}).Return(tc.deletePoliciesErr) - repoCall1 := repo.On("AssignParentGroup", context.Background(), tc.parentID, tc.childrenIDs).Return(tc.assignParentErr) - err := svc.AddChildrenGroups(context.Background(), validSession, tc.parentID, tc.childrenIDs) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("expected error %v to contain %v", err, tc.err)) - repoCall.Unset() - policyCall.Unset() - policyCall1.Unset() - repoCall1.Unset() - }) - } -} - -func TestRemoveChildrenGroups(t *testing.T) { - svc := newService(t) - - cases := []struct { - desc string - parentID string - childrenIDs []string - retrieveResp groups.Page - retrieveErr error - deletePoliciesErr error - addPoliciesErr error - unassignParentErr error - err error - }{ - { - desc: "remove children groups successfully", - parentID: parentGroupID, - childrenIDs: []string{childGroupID}, - retrieveResp: groups.Page{ - Groups: []groups.Group{childGroup}, - PageMeta: groups.PageMeta{ - Total: 1, - }, - }, - err: nil, - }, - { - desc: "remove children groups with failed to retrieve", - parentID: parentGroupID, - childrenIDs: []string{childGroupID}, - retrieveErr: repoerr.ErrNotFound, - err: repoerr.ErrNotFound, - }, - { - desc: "remove non existent child group", - parentID: parentGroupID, - childrenIDs: []string{testsutil.GenerateUUID(&testing.T{})}, - retrieveResp: groups.Page{}, - err: groups.ErrGroupIDs, - }, - { - desc: "remove children groups from different parent", - parentID: validGroup.ID, - childrenIDs: []string{childGroupID}, - retrieveResp: groups.Page{ - Groups: []groups.Group{childGroup}, - PageMeta: groups.PageMeta{ - Total: 1, - }, - }, - err: svcerr.ErrConflict, - }, - { - desc: "remove children groups with failed to delete policies", - parentID: parentGroupID, - childrenIDs: []string{childGroupID}, - retrieveResp: groups.Page{ - Groups: []groups.Group{childGroup}, - PageMeta: groups.PageMeta{ - Total: 1, - }, - }, - deletePoliciesErr: svcerr.ErrAuthorization, - err: svcerr.ErrDeletePolicies, - }, - { - desc: "remove children groups with repo error in unassign children groups", - parentID: parentGroupID, - childrenIDs: []string{childGroupID}, - retrieveResp: groups.Page{ - Groups: []groups.Group{childGroup}, - PageMeta: groups.PageMeta{ - Total: 1, - }, - }, - unassignParentErr: repoerr.ErrNotFound, - err: repoerr.ErrNotFound, - }, - { - desc: "remove children groups with repo error in unassign children groups and failed to add policies", - parentID: parentGroupID, - childrenIDs: []string{childGroupID}, - retrieveResp: groups.Page{ - Groups: []groups.Group{childGroup}, - PageMeta: groups.PageMeta{ - Total: 1, - }, - }, - unassignParentErr: repoerr.ErrNotFound, - addPoliciesErr: svcerr.ErrAuthorization, - err: apiutil.ErrRollbackTx, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - pol := policysvc.Policy{ - Domain: validID, - SubjectType: policysvc.GroupType, - Subject: tc.parentID, - Relation: policysvc.ParentGroupRelation, - ObjectType: policysvc.GroupType, - Object: childGroupID, - } - repoCall := repo.On("RetrieveByIDs", context.Background(), groups.PageMeta{Limit: 1<<63 - 1}, tc.childrenIDs).Return(tc.retrieveResp, tc.retrieveErr) - policyCall := policies.On("DeletePolicies", context.Background(), []policysvc.Policy{pol}).Return(tc.deletePoliciesErr) - policyCall1 := policies.On("AddPolicies", context.Background(), []policysvc.Policy{pol}).Return(tc.addPoliciesErr) - repoCall1 := repo.On("UnassignParentGroup", context.Background(), tc.parentID, tc.childrenIDs).Return(tc.unassignParentErr) - err := svc.RemoveChildrenGroups(context.Background(), validSession, tc.parentID, tc.childrenIDs) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("expected error %v to contain %v", err, tc.err)) - repoCall.Unset() - policyCall.Unset() - policyCall1.Unset() - repoCall1.Unset() - }) - } -} - -func TestRemoveAllChildrenGroups(t *testing.T) { - svc := newService(t) - - cases := []struct { - desc string - parentID string - deletePolicyErr error - unassignAllChildrenErr error - err error - }{ - { - desc: "remove all children groups successfully", - parentID: parentGroupID, - err: nil, - }, - { - desc: "remove all children groups with failed to delete policy", - parentID: parentGroupID, - deletePolicyErr: svcerr.ErrAuthorization, - err: svcerr.ErrDeletePolicies, - }, - { - desc: "remove all children groups with failed to unassign all children", - parentID: parentGroupID, - deletePolicyErr: nil, - unassignAllChildrenErr: repoerr.ErrNotFound, - err: repoerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - policyCall := policies.On("DeletePolicyFilter", context.Background(), policysvc.Policy{ - Domain: validID, - SubjectType: policysvc.GroupType, - Subject: tc.parentID, - Relation: policysvc.ParentGroupRelation, - ObjectType: policysvc.GroupType, - }).Return(tc.deletePolicyErr) - repoCall := repo.On("UnassignAllChildrenGroups", context.Background(), tc.parentID).Return(tc.unassignAllChildrenErr) - err := svc.RemoveAllChildrenGroups(context.Background(), validSession, tc.parentID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("expected error %v to contain %v", err, tc.err)) - policyCall.Unset() - repoCall.Unset() - }) - } -} - -func TestListAllChildrenGroups(t *testing.T) { - svc := newService(t) - - cases := []struct { - desc string - session smqauthn.Session - pageMeta groups.PageMeta - parentID string - startLevel int64 - endLevel int64 - retrieveRes groups.Page - retrieveErr error - resp groups.Page - err error - }{ - { - desc: "list all children groups successfully", - session: validSession, - parentID: parentGroupID, - pageMeta: groups.PageMeta{ - Limit: 10, - Offset: 0, - }, - startLevel: 0, - endLevel: -1, - retrieveRes: groups.Page{ - Groups: []groups.Group{childGroup}, - PageMeta: groups.PageMeta{ - Total: 1, - }, - }, - resp: groups.Page{ - Groups: []groups.Group{childGroup}, - PageMeta: groups.PageMeta{ - Total: 1, - }, - }, - err: nil, - }, - { - desc: "list all children groups with failed to retrieve", - session: validSession, - parentID: parentGroupID, - pageMeta: groups.PageMeta{ - Limit: 10, - Offset: 0, - }, - startLevel: 0, - endLevel: -1, - retrieveErr: repoerr.ErrNotFound, - resp: groups.Page{}, - err: svcerr.ErrViewEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := repo.On("RetrieveChildrenGroups", context.Background(), tc.session.DomainID, tc.session.UserID, tc.parentID, tc.startLevel, tc.endLevel, tc.pageMeta).Return(tc.retrieveRes, tc.retrieveErr) - page, err := svc.ListChildrenGroups(context.Background(), tc.session, tc.parentID, tc.startLevel, tc.endLevel, tc.pageMeta) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("expected error %v to contain %v", err, tc.err)) - assert.Equal(t, tc.resp, page) - repoCall.Unset() - }) - } -} - -func TestDeleteGroup(t *testing.T) { - svc := newService(t) - - cases := []struct { - desc string - id string - changeStatusRes groups.Group - changeStatusErr error - deletePoliciesErr error - deleteErr error - unsetFromChannels error - unsetFromClients error - err error - }{ - { - desc: "delete group successfully", - id: validGroup.ID, - err: nil, - }, - { - desc: "delete group with parent successfully", - id: childGroupID, - changeStatusRes: childGroup, - err: nil, - }, - { - desc: "delete group with failed to remove parent group from channels", - id: validGroup.ID, - unsetFromChannels: svcerr.ErrRemoveEntity, - err: svcerr.ErrRemoveEntity, - }, - { - desc: "delete group with failed to remove parent group from clients", - id: validGroup.ID, - unsetFromChannels: nil, - unsetFromClients: svcerr.ErrRemoveEntity, - err: svcerr.ErrRemoveEntity, - }, - { - desc: "delete group with failed to change status", - id: validGroup.ID, - changeStatusErr: repoerr.ErrNotFound, - err: repoerr.ErrNotFound, - }, - { - desc: "delete group with failed to delete", - id: validGroup.ID, - changeStatusRes: validGroup, - deleteErr: repoerr.ErrNotFound, - err: repoerr.ErrNotFound, - }, - { - desc: "delete group with failed to delete policies", - id: validGroup.ID, - changeStatusRes: validGroup, - deleteErr: nil, - deletePoliciesErr: svcerr.ErrAuthorization, - err: svcerr.ErrDeletePolicies, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := repo.On("ChangeStatus", context.Background(), groups.Group{ID: tc.id, Status: groups.DeletedStatus}).Return(tc.changeStatusRes, tc.changeStatusErr) - repoCall1 := repo.On("Delete", context.Background(), tc.id).Return(tc.deleteErr) - svcCall := channels.On("UnsetParentGroupFromChannels", context.Background(), &grpcChannelsV1.UnsetParentGroupFromChannelsReq{ParentGroupId: tc.id}).Return(&grpcChannelsV1.UnsetParentGroupFromChannelsRes{}, tc.unsetFromChannels) - svcCall1 := clients.On("UnsetParentGroupFromClient", context.Background(), &grpcClientsV1.UnsetParentGroupFromClientReq{ParentGroupId: tc.id}).Return(&grpcClientsV1.UnsetParentGroupFromClientRes{}, tc.unsetFromClients) - repoCall2 := repo.On("RetrieveEntitiesRolesActionsMembers", context.Background(), []string{tc.id}).Return([]roles.EntityActionRole{}, []roles.EntityMemberRole{}, nil) - policyCall := policies.On("DeletePolicyFilter", context.Background(), mock.Anything).Return(tc.deletePoliciesErr) - policyCall1 := policies.On("DeletePolicies", context.Background(), mock.Anything).Return(nil) - err := svc.DeleteGroup(context.Background(), validSession, tc.id) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("expected error %v to contain %v", err, tc.err)) - policyCall.Unset() - repoCall.Unset() - repoCall1.Unset() - svcCall.Unset() - svcCall1.Unset() - repoCall2.Unset() - policyCall1.Unset() - }) - } -} diff --git a/groups/status.go b/groups/status.go deleted file mode 100644 index 3a45357ea..000000000 --- a/groups/status.go +++ /dev/null @@ -1,83 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package groups - -import ( - "encoding/json" - "strings" - - svcerr "github.com/absmach/magistrala/pkg/errors/service" -) - -// Status represents Group status. -type Status uint8 - -// Possible Group status values. -const ( - // EnabledStatus represents enabled Group. - EnabledStatus Status = iota - // DisabledStatus represents disabled Group. - DisabledStatus - // DeletedStatus represents deleted Group. - DeletedStatus - - // AllStatus is used for querying purposes to list groups irrespective - // of their status - both active and inactive. It is never stored in the - // database as the actual Group status and should always be the largest - // value in this enumeration. - AllStatus -) - -// String representation of the possible status values. -const ( - Disabled = "disabled" - Enabled = "enabled" - Deleted = "deleted" - All = "all" - Unknown = "unknown" -) - -// String converts group status to string literal. -func (s Status) String() string { - switch s { - case DisabledStatus: - return Disabled - case EnabledStatus: - return Enabled - case DeletedStatus: - return Deleted - case AllStatus: - return All - default: - return Unknown - } -} - -// ToStatus converts string value to a valid Group status. -func ToStatus(status string) (Status, error) { - switch status { - case Disabled: - return DisabledStatus, nil - case Enabled: - return EnabledStatus, nil - case Deleted: - return DeletedStatus, nil - case All: - return AllStatus, nil - } - return Status(0), svcerr.ErrInvalidStatus -} - -// Custom Marshaller for Status. -func (s Status) MarshalJSON() ([]byte, error) { - return json.Marshal(s.String()) -} - -// Custom Unmarshaler for Status. -func (s *Status) UnmarshalJSON(data []byte) error { - str := strings.Trim(string(data), "\"") - val, err := ToStatus(str) - *s = val - return err -} diff --git a/groups/status_test.go b/groups/status_test.go deleted file mode 100644 index 1e3a627e9..000000000 --- a/groups/status_test.go +++ /dev/null @@ -1,52 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package groups_test - -import ( - "testing" - - "github.com/absmach/magistrala/groups" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/stretchr/testify/assert" -) - -func TestStatus_String(t *testing.T) { - cases := []struct { - name string - status groups.Status - expected string - }{ - {"Enabled", groups.EnabledStatus, "enabled"}, - {"Disabled", groups.DisabledStatus, "disabled"}, - {"Deleted", groups.DeletedStatus, "deleted"}, - {"All", groups.AllStatus, "all"}, - {"Unknown", groups.Status(100), "unknown"}, - } - - for _, tc := range cases { - got := tc.status.String() - assert.Equal(t, tc.expected, got, "Status.String() = %v, expected %v", got, tc.expected) - } -} - -func TestToStatus(t *testing.T) { - cases := []struct { - name string - status string - gstatus groups.Status - err error - }{ - {"Enabled", "enabled", groups.EnabledStatus, nil}, - {"Disabled", "disabled", groups.DisabledStatus, nil}, - {"Deleted", "deleted", groups.DeletedStatus, nil}, - {"All", "all", groups.AllStatus, nil}, - {"Unknown", "unknown", groups.Status(0), svcerr.ErrInvalidStatus}, - } - - for _, tc := range cases { - got, err := groups.ToStatus(tc.status) - assert.Equal(t, tc.err, err, "ToStatus() error = %v, expected %v", err, tc.err) - assert.Equal(t, tc.gstatus, got, "ToStatus() = %v, expected %v", got, tc.gstatus) - } -} diff --git a/internal/atom/authz.go b/internal/atom/authz.go new file mode 100644 index 000000000..d799bc727 --- /dev/null +++ b/internal/atom/authz.go @@ -0,0 +1,77 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package atom + +import ( + "context" + + "github.com/absmach/magistrala/pkg/authn" + "github.com/absmach/magistrala/pkg/errors" + "github.com/absmach/magistrala/pkg/policies" +) + +type Authorizer interface { + CheckAuthz(ctx context.Context, req AuthzRequest) (AuthzResponse, error) +} + +func Authorize(ctx context.Context, client Authorizer, session authn.Session, action, legacyObjectType, objectID, resourceKind string) error { + req := AuthzRequest{ + SubjectID: SubjectID(session), + Action: CapabilityName(action), + ResourceID: resourceID(legacyObjectType, objectID), + ObjectKind: ObjectKind(legacyObjectType, resourceKind), + ObjectID: objectID, + Context: map[string]any{ + "domain_id": session.DomainID, + "legacy_object_type": legacyObjectType, + }, + } + res, err := client.CheckAuthz(ctx, req) + if err != nil { + return errors.Wrap(errors.ErrAuthorization, err) + } + if !res.Allowed { + return errors.ErrAuthorization + } + return nil +} + +func SubjectID(session authn.Session) string { + if session.UserID != "" { + return session.UserID + } + return session.DomainUserID +} + +func ObjectKind(legacyObjectType, resourceKind string) string { + switch legacyObjectType { + case policies.DomainType: + return atomObjectKindTenant + case policies.PlatformType: + return policies.PlatformType + case policies.ClientType: + return atomObjectKindEntity + case policies.GroupType: + return atomObjectKindGroup + case policies.ChannelType, policies.RulesType, policies.ReportsType, policies.AlarmsType: + return atomObjectKindResource + } + switch resourceKind { + case KindClient, atomKindDevice: + return atomObjectKindEntity + case atomKindGroup: + return atomObjectKindGroup + case KindChannel, KindRule, KindReport, KindAlarm: + return atomObjectKindResource + default: + return resourceKind + } +} + +func resourceID(legacyObjectType, objectID string) string { + if legacyObjectType == policies.DomainType || legacyObjectType == policies.PlatformType { + return "" + } + return objectID +} diff --git a/internal/atom/authz_compat.go b/internal/atom/authz_compat.go new file mode 100644 index 000000000..5fd8b9d17 --- /dev/null +++ b/internal/atom/authz_compat.go @@ -0,0 +1,124 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package atom + +import ( + "context" + "strings" + + smqauthz "github.com/absmach/magistrala/pkg/authz" + "github.com/absmach/magistrala/pkg/errors" + "github.com/absmach/magistrala/pkg/policies" +) + +type AuthorizationCompat struct { + Client Authorizer +} + +func NewAuthorizationCompat(client Authorizer) AuthorizationCompat { + return AuthorizationCompat{Client: client} +} + +func (a AuthorizationCompat) Authorize(ctx context.Context, pr smqauthz.PolicyReq, _ *smqauthz.PATReq) error { + subjectID := pr.Subject + if subjectID == "" { + return errors.ErrAuthentication + } + objectKind := ObjectKind(pr.ObjectType, legacyResourceKind(pr.ObjectKind, pr.ObjectType)) + res, err := a.Client.CheckAuthz(ctx, AuthzRequest{ + SubjectID: subjectID, + Action: CapabilityName(pr.Permission), + ResourceID: resourceID(pr.ObjectType, pr.Object), + ObjectKind: objectKind, + ObjectID: pr.Object, + Context: map[string]any{ + "domain_id": pr.Domain, + "legacy_object_kind": pr.ObjectKind, + "legacy_object_type": pr.ObjectType, + "legacy_permission": pr.Permission, + "legacy_relation": pr.Relation, + "legacy_subject_kind": pr.SubjectKind, + "legacy_subject_type": pr.SubjectType, + }, + }) + if err != nil { + return errors.Wrap(errors.ErrAuthorization, err) + } + if !res.Allowed { + return errors.ErrAuthorization + } + return nil +} + +func CapabilityName(action string) string { + normalized := strings.ToLower(strings.TrimSpace(action)) + switch { + case normalized == policies.AdminPermission, + normalized == "admin_permission", + strings.Contains(normalized, "manage_role"): + return atomActionManage + case normalized == policies.ViewPermission, + normalized == atomActionRead, + strings.Contains(normalized, "read"), + strings.Contains(normalized, "view"): + return atomActionRead + case normalized == policies.CreatePermission, + normalized == "write", + strings.Contains(normalized, "create"), + strings.Contains(normalized, "update"), + strings.Contains(normalized, "edit"), + strings.Contains(normalized, "enable"), + strings.Contains(normalized, "disable"), + strings.Contains(normalized, "assign"), + strings.Contains(normalized, "acknowledge"), + strings.Contains(normalized, "resolve"): + return atomActionWrite + case normalized == policies.DeletePermission, + strings.Contains(normalized, "delete"), + strings.Contains(normalized, "remove"): + return atomActionDelete + case normalized == policies.PublishPermission: + return atomActionPublish + case normalized == policies.SubscribePermission: + return atomActionSubscribe + case normalized == "generate", normalized == "execute": + return atomActionExecute + case normalized == atomActionList: + return atomActionList + default: + return normalized + } +} + +func legacyResourceKind(objectKind, objectType string) string { + switch objectKind { + case policies.ChannelsKind, policies.NewChannelKind: + return KindChannel + case policies.ClientsKind, policies.NewClientKind: + return "client" + case policies.GroupsKind, policies.NewGroupKind: + return atomObjectKindGroup + case policies.DomainsKind: + return atomObjectKindTenant + default: + switch objectType { + case policies.ChannelType: + return KindChannel + case policies.ClientType: + return "client" + case policies.GroupType: + return atomObjectKindGroup + case policies.DomainType: + return atomObjectKindTenant + case policies.RulesType: + return KindRule + case policies.ReportsType: + return KindReport + case policies.AlarmsType: + return KindAlarm + default: + return objectType + } + } +} diff --git a/internal/atom/authz_test.go b/internal/atom/authz_test.go new file mode 100644 index 000000000..718e601a3 --- /dev/null +++ b/internal/atom/authz_test.go @@ -0,0 +1,93 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package atom_test + +import ( + "context" + "testing" + + channelsv1 "github.com/absmach/magistrala/api/grpc/channels/v1" + "github.com/absmach/magistrala/internal/atom" + "github.com/absmach/magistrala/pkg/authn" + "github.com/absmach/magistrala/pkg/connections" + "github.com/absmach/magistrala/pkg/errors" + "github.com/absmach/magistrala/pkg/policies" + "github.com/stretchr/testify/assert" +) + +type authzClient struct { + req atom.AuthzRequest + res atom.AuthzResponse + err error +} + +func (c *authzClient) CheckAuthz(_ context.Context, req atom.AuthzRequest) (atom.AuthzResponse, error) { + c.req = req + return c.res, c.err +} + +func TestAuthorizeBuildsResourceRequest(t *testing.T) { + client := &authzClient{res: atom.AuthzResponse{Allowed: true}} + session := authn.Session{UserID: "user-1", DomainID: "domain-1"} + + err := atom.Authorize(context.Background(), client, session, "view", policies.RulesType, "rule-1", atom.KindRule) + + assert.NoError(t, err) + assert.Equal(t, atom.AuthzRequest{ + SubjectID: "user-1", + Action: "read", + ResourceID: "rule-1", + ObjectKind: "resource", + ObjectID: "rule-1", + Context: map[string]any{ + "domain_id": "domain-1", + "legacy_object_type": policies.RulesType, + }, + }, client.req) +} + +func TestAuthorizeBuildsTenantRequest(t *testing.T) { + client := &authzClient{res: atom.AuthzResponse{Allowed: true}} + session := authn.Session{UserID: "user-1", DomainID: "domain-1"} + + err := atom.Authorize(context.Background(), client, session, "create", policies.DomainType, "domain-1", atom.KindRule) + + assert.NoError(t, err) + assert.Equal(t, "tenant", client.req.ObjectKind) + assert.Equal(t, "domain-1", client.req.ObjectID) + assert.Empty(t, client.req.ResourceID) +} + +func TestAuthorizeDenied(t *testing.T) { + client := &authzClient{res: atom.AuthzResponse{Allowed: false}} + + err := atom.Authorize(context.Background(), client, authn.Session{UserID: "user-1"}, "view", policies.RulesType, "rule-1", atom.KindRule) + + assert.True(t, errors.Contains(err, errors.ErrAuthorization)) +} + +func TestChannelsCompatAuthorizeBuildsResourceRequest(t *testing.T) { + client := &authzClient{res: atom.AuthzResponse{Allowed: true}} + compat := atom.NewChannelsCompat(client) + + res, err := compat.Authorize(context.Background(), &channelsv1.AuthzReq{ + ClientId: "domain-1_user-1", + DomainId: "domain-1", + Type: uint32(connections.Subscribe), + ChannelId: "channel-1", + }) + + assert.NoError(t, err) + assert.True(t, res.GetAuthorized()) + assert.Equal(t, atom.AuthzRequest{ + SubjectID: "user-1", + Action: "subscribe", + ResourceID: "channel-1", + ObjectKind: "resource", + ObjectID: "channel-1", + Context: map[string]any{ + "domain_id": "domain-1", + }, + }, client.req) +} diff --git a/internal/atom/bootstrap.go b/internal/atom/bootstrap.go new file mode 100644 index 000000000..ca546d291 --- /dev/null +++ b/internal/atom/bootstrap.go @@ -0,0 +1,161 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package atom + +import ( + "context" + "fmt" +) + +var magistralaActionDescriptions = map[string]string{ + atomActionRead: "Read / view an object", + atomActionWrite: "Create or update an object", + atomActionDelete: "Delete an object", + atomActionManage: "Full administrative control", + atomActionPublish: "Publish messages to a channel", + atomActionSubscribe: "Subscribe to channel messages", + atomActionExecute: "Execute a command or action", + atomActionList: "List objects", +} + +var magistralaActionApplicability = []CapabilityApplicabilitySpec{ + {ActionName: atomActionWrite, ObjectKind: atomObjectKindTenant}, + + {ActionName: atomActionRead, ObjectKind: atomObjectKindGroup}, + {ActionName: atomActionWrite, ObjectKind: atomObjectKindGroup}, + {ActionName: atomActionDelete, ObjectKind: atomObjectKindGroup}, + {ActionName: atomActionManage, ObjectKind: atomObjectKindGroup}, + {ActionName: atomActionList, ObjectKind: atomObjectKindGroup}, + + {ActionName: atomActionRead, ObjectKind: atomObjectKindResource, ObjectType: "resource:channel"}, + {ActionName: atomActionWrite, ObjectKind: atomObjectKindResource, ObjectType: "resource:channel"}, + {ActionName: atomActionDelete, ObjectKind: atomObjectKindResource, ObjectType: "resource:channel"}, + {ActionName: atomActionManage, ObjectKind: atomObjectKindResource, ObjectType: "resource:channel"}, + {ActionName: atomActionPublish, ObjectKind: atomObjectKindResource, ObjectType: "resource:channel"}, + {ActionName: atomActionSubscribe, ObjectKind: atomObjectKindResource, ObjectType: "resource:channel"}, + + {ActionName: atomActionRead, ObjectKind: atomObjectKindResource, ObjectType: "resource:rule"}, + {ActionName: atomActionWrite, ObjectKind: atomObjectKindResource, ObjectType: "resource:rule"}, + {ActionName: atomActionDelete, ObjectKind: atomObjectKindResource, ObjectType: "resource:rule"}, + {ActionName: atomActionManage, ObjectKind: atomObjectKindResource, ObjectType: "resource:rule"}, + {ActionName: atomActionExecute, ObjectKind: atomObjectKindResource, ObjectType: "resource:rule"}, + {ActionName: atomActionList, ObjectKind: atomObjectKindResource, ObjectType: "resource:rule"}, + + {ActionName: atomActionRead, ObjectKind: atomObjectKindResource, ObjectType: "resource:report"}, + {ActionName: atomActionWrite, ObjectKind: atomObjectKindResource, ObjectType: "resource:report"}, + {ActionName: atomActionDelete, ObjectKind: atomObjectKindResource, ObjectType: "resource:report"}, + {ActionName: atomActionManage, ObjectKind: atomObjectKindResource, ObjectType: "resource:report"}, + {ActionName: atomActionExecute, ObjectKind: atomObjectKindResource, ObjectType: "resource:report"}, + {ActionName: atomActionList, ObjectKind: atomObjectKindResource, ObjectType: "resource:report"}, + + {ActionName: atomActionRead, ObjectKind: atomObjectKindResource, ObjectType: "resource:alarm"}, + {ActionName: atomActionWrite, ObjectKind: atomObjectKindResource, ObjectType: "resource:alarm"}, + {ActionName: atomActionDelete, ObjectKind: atomObjectKindResource, ObjectType: "resource:alarm"}, + {ActionName: atomActionManage, ObjectKind: atomObjectKindResource, ObjectType: "resource:alarm"}, + {ActionName: atomActionList, ObjectKind: atomObjectKindResource, ObjectType: "resource:alarm"}, +} + +var magistralaActionAssignmentRules = []ActionAssignmentRuleSpec{ + { + EntityKind: atomKindDevice, + ActionName: atomActionPublish, + ObjectKind: atomObjectKindResource, + ObjectType: "resource:channel", + Decision: "allow", + }, + { + EntityKind: atomKindDevice, + ActionName: atomActionSubscribe, + ObjectKind: atomObjectKindResource, + ObjectType: "resource:channel", + Decision: "allow", + }, +} + +// BootstrapMagistralaActions installs Magistrala-specific action applicability in Atom. +// It is safe to call repeatedly during startup. +func BootstrapMagistralaActions(ctx context.Context, client *Client) error { + if client == nil { + return fmt.Errorf("atom client is nil") + } + capabilities, err := client.ListCapabilities(ctx) + if err != nil { + return fmt.Errorf("list atom actions: %w", err) + } + byName := map[string]Capability{} + for _, capability := range capabilities.Items { + byName[capability.Name] = capability + } + + for _, spec := range magistralaActionApplicability { + capability, ok := byName[spec.ActionName] + if !ok { + description := spec.Description + if description == "" { + description = magistralaActionDescriptions[spec.ActionName] + } + capability, err = client.CreateCapability(ctx, spec.ActionName, description) + if err != nil { + if !IsConflict(err) { + return fmt.Errorf("create atom action %q: %w", spec.ActionName, err) + } + id, lookupErr := client.CapabilityID(ctx, spec.ActionName) + if lookupErr != nil { + return fmt.Errorf("lookup existing atom action %q after conflict: %w", spec.ActionName, lookupErr) + } + capability = Capability{ID: id, Name: spec.ActionName, Description: description} + } + byName[spec.ActionName] = capability + } + if _, err := client.AddCapabilityApplicability(ctx, capability.ID, spec.ObjectKind, spec.ObjectType); err != nil { + return fmt.Errorf("add atom applicability %s -> %s:%s: %w", spec.ActionName, spec.ObjectKind, spec.ObjectType, err) + } + } + + for _, spec := range magistralaActionAssignmentRules { + if err := ensureActionAssignmentRule(ctx, client, spec); err != nil { + return fmt.Errorf("ensure atom assignment guardrail %s %s %s:%s: %w", spec.EntityKind, spec.ActionName, spec.ObjectKind, spec.ObjectType, err) + } + } + return nil +} + +func ensureActionAssignmentRule(ctx context.Context, client *Client, spec ActionAssignmentRuleSpec) error { + rules, err := client.ListActionAssignmentRules(ctx, spec) + if err != nil { + return err + } + if actionAssignmentRuleExists(rules.Items, spec) { + return nil + } + if _, err := client.CreateActionAssignmentRule(ctx, spec); err != nil { + if !IsConflict(err) { + return err + } + rules, lookupErr := client.ListActionAssignmentRules(ctx, spec) + if lookupErr != nil { + return fmt.Errorf("lookup existing rule after conflict: %w", lookupErr) + } + if actionAssignmentRuleExists(rules.Items, spec) { + return nil + } + return err + } + return nil +} + +func actionAssignmentRuleExists(rules []ActionAssignmentRule, spec ActionAssignmentRuleSpec) bool { + for _, rule := range rules { + if rule.TenantID == spec.TenantID && + rule.EntityKind == spec.EntityKind && + rule.ActionName == spec.ActionName && + rule.ObjectKind == spec.ObjectKind && + rule.ObjectType == spec.ObjectType && + rule.Decision == spec.Decision && + rule.IsAbsolute == spec.IsAbsolute { + return true + } + } + return false +} diff --git a/internal/atom/bootstrap_test.go b/internal/atom/bootstrap_test.go new file mode 100644 index 000000000..aa9ae3c8a --- /dev/null +++ b/internal/atom/bootstrap_test.go @@ -0,0 +1,158 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package atom + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" +) + +func TestBootstrapMagistralaActionsCreatesMissingActionsAndApplicability(t *testing.T) { + actions := map[string]Capability{ + atomActionRead: {ID: "read-id", Name: atomActionRead}, + atomActionWrite: {ID: "write-id", Name: atomActionWrite}, + atomActionDelete: {ID: "delete-id", Name: atomActionDelete}, + atomActionManage: {ID: "manage-id", Name: atomActionManage}, + } + var applicability []map[string]any + var assignmentRules []map[string]any + + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost || r.URL.Path != atomGraphQLPath { + t.Fatalf("unexpected request: %s %s", r.Method, r.URL.Path) + } + var payload struct { + Query string `json:"query"` + Variables map[string]any `json:"variables"` + } + if err := json.NewDecoder(r.Body).Decode(&payload); err != nil { + t.Fatalf("decode request: %v", err) + } + + switch { + case strings.Contains(payload.Query, "query Actions"): + items := make([]Capability, 0, len(actions)) + for _, action := range actions { + items = append(items, action) + } + _ = json.NewEncoder(w).Encode(map[string]any{ + "data": map[string]any{ + "actions": map[string]any{"items": items, "total": len(items)}, + }, + }) + case strings.Contains(payload.Query, "query ActionAssignmentRules"): + _ = json.NewEncoder(w).Encode(map[string]any{ + "data": map[string]any{ + "actionAssignmentRules": map[string]any{"items": []map[string]any{}, "total": 0}, + }, + }) + case strings.Contains(payload.Query, "createActionAssignmentRule"): + input := payload.Variables["input"].(map[string]any) + assignmentRules = append(assignmentRules, input) + _ = json.NewEncoder(w).Encode(map[string]any{ + "data": map[string]any{ + "createActionAssignmentRule": map[string]any{ + "id": input["actionName"].(string) + "-rule-id", + "tenant_id": "", + "entity_kind": input["entityKind"], + "action_name": input["actionName"], + "object_kind": input["objectKind"], + "object_type": input["objectType"], + "decision": input["decision"], + "is_absolute": input["isAbsolute"], + "created_at": "2026-06-18T00:00:00Z", + }, + }, + }) + case strings.Contains(payload.Query, "createAction"): + input := payload.Variables["input"].(map[string]any) + name := input["name"].(string) + action := Capability{ID: name + "-id", Name: name} + actions[name] = action + _ = json.NewEncoder(w).Encode(map[string]any{ + "data": map[string]any{"createAction": action}, + }) + case strings.Contains(payload.Query, "addActionApplicability"): + input := payload.Variables["input"].(map[string]any) + applicability = append(applicability, input) + _ = json.NewEncoder(w).Encode(map[string]any{ + "data": map[string]any{ + "addActionApplicability": map[string]any{ + "action_id": input["actionId"], + "action_name": "action", + "object_kind": input["objectKind"], + "object_type": input["objectType"], + "description": "", + }, + }, + }) + default: + t.Fatalf("unexpected GraphQL payload: %s", payload.Query) + } + })) + defer srv.Close() + + client := NewClient(Config{URL: srv.URL, Timeout: time.Second}) + if err := BootstrapMagistralaActions(context.Background(), client); err != nil { + t.Fatalf("bootstrap failed: %v", err) + } + + for _, name := range []string{atomActionRead, atomActionWrite, atomActionDelete, atomActionManage, atomActionPublish, atomActionSubscribe, atomActionExecute, atomActionList} { + if _, ok := actions[name]; !ok { + t.Fatalf("action %q was not ensured", name) + } + } + if len(applicability) != len(magistralaActionApplicability) { + t.Fatalf("unexpected applicability count: got %d want %d", len(applicability), len(magistralaActionApplicability)) + } + assertApplicability(t, applicability, "write-id", atomObjectKindTenant, "") + assertApplicability(t, applicability, "read-id", atomObjectKindGroup, "") + assertApplicability(t, applicability, "write-id", atomObjectKindGroup, "") + assertApplicability(t, applicability, "delete-id", atomObjectKindGroup, "") + assertApplicability(t, applicability, "manage-id", atomObjectKindGroup, "") + assertApplicability(t, applicability, "list-id", atomObjectKindGroup, "") + assertApplicability(t, applicability, "publish-id", atomObjectKindResource, "resource:channel") + assertApplicability(t, applicability, "execute-id", atomObjectKindResource, "resource:rule") + assertApplicability(t, applicability, "list-id", atomObjectKindResource, "resource:rule") + assertApplicability(t, applicability, "execute-id", atomObjectKindResource, "resource:report") + assertApplicability(t, applicability, "list-id", atomObjectKindResource, "resource:report") + assertApplicability(t, applicability, "manage-id", atomObjectKindResource, "resource:alarm") + assertApplicability(t, applicability, "list-id", atomObjectKindResource, "resource:alarm") + if len(assignmentRules) != len(magistralaActionAssignmentRules) { + t.Fatalf("unexpected assignment guardrail count: got %d want %d", len(assignmentRules), len(magistralaActionAssignmentRules)) + } + assertAssignmentRule(t, assignmentRules, atomKindDevice, atomActionPublish, atomObjectKindResource, "resource:channel", "allow") + assertAssignmentRule(t, assignmentRules, atomKindDevice, atomActionSubscribe, atomObjectKindResource, "resource:channel", "allow") +} + +func assertApplicability(t *testing.T, entries []map[string]any, actionID, objectKind, objectType string) { + t.Helper() + for _, entry := range entries { + entryObjectType, _ := entry["objectType"].(string) + if entry["actionId"] == actionID && entry["objectKind"] == objectKind && entryObjectType == objectType { + return + } + } + t.Fatalf("missing applicability action=%s object=%s:%s", actionID, objectKind, objectType) +} + +func assertAssignmentRule(t *testing.T, entries []map[string]any, entityKind, actionName, objectKind, objectType, decision string) { + t.Helper() + for _, entry := range entries { + if entry["entityKind"] == entityKind && + entry["actionName"] == actionName && + entry["objectKind"] == objectKind && + entry["objectType"] == objectType && + entry["decision"] == decision && + entry["isAbsolute"] == false { + return + } + } + t.Fatalf("missing assignment guardrail entity=%s action=%s object=%s:%s decision=%s", entityKind, actionName, objectKind, objectType, decision) +} diff --git a/internal/atom/client.go b/internal/atom/client.go new file mode 100644 index 000000000..7078db5b5 --- /dev/null +++ b/internal/atom/client.go @@ -0,0 +1,924 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package atom + +import ( + "bytes" + "context" + "encoding/json" + "fmt" + "io" + "net/http" + "strings" +) + +type Client struct { + baseURL string + token string + adminUsername string + adminSecret string + userAgent string + httpClient *http.Client +} + +func NewClient(cfg Config) *Client { + timeout := cfg.Timeout + if timeout == 0 { + timeout = defaultTimeout + } + return &Client{ + baseURL: strings.TrimRight(cfg.URL, "/"), + token: cfg.Token, + adminUsername: cfg.AdminUsername, + adminSecret: cfg.AdminSecret, + userAgent: cfg.UserAgent, + httpClient: &http.Client{ + Timeout: timeout, + }, + } +} + +func (c *Client) UpsertTenant(ctx context.Context, tenant Tenant) error { + if tenant.ID == "" { + _, err := c.CreateTenant(ctx, tenant) + return err + } + if _, err := c.CreateTenant(ctx, tenant); err == nil || !IsConflict(err) { + return err + } + _, err := c.UpdateTenant(ctx, tenant.ID, tenant) + return err +} + +func (c *Client) CreateTenant(ctx context.Context, tenant Tenant) (Tenant, error) { + var out struct { + CreateTenant Tenant `json:"createTenant"` + } + err := c.graphQL(ctx, `mutation CreateTenant($input: CreateTenantInput!) { + createTenant(input: $input) { id name route status tags attributes created_by: createdBy updated_by: updatedBy created_at: createdAt updated_at: updatedAt } + }`, map[string]any{"input": tenantCreateInput(tenant)}, &out) + return out.CreateTenant, err +} + +func (c *Client) GetTenant(ctx context.Context, id string) (Tenant, error) { + var out struct { + Tenant Tenant `json:"tenant"` + } + err := c.graphQL(ctx, `query Tenant($id: ID!) { + tenant(id: $id) { id name route status tags attributes created_by: createdBy updated_by: updatedBy created_at: createdAt updated_at: updatedAt } + }`, map[string]any{"id": id}, &out) + return out.Tenant, err +} + +func (c *Client) UpdateTenant(ctx context.Context, id string, tenant Tenant) (Tenant, error) { + var out struct { + UpdateTenant Tenant `json:"updateTenant"` + } + err := c.graphQL(ctx, `mutation UpdateTenant($id: ID!, $input: UpdateTenantInput!) { + updateTenant(id: $id, input: $input) { id name route status tags attributes created_by: createdBy updated_by: updatedBy created_at: createdAt updated_at: updatedAt } + }`, map[string]any{"id": id, "input": tenantUpdateInput(tenant)}, &out) + return out.UpdateTenant, err +} + +func (c *Client) ChangeTenantStatus(ctx context.Context, id, action string) (Tenant, error) { + field := map[string]string{ + "enable": "enableTenant", + "disable": "disableTenant", + "freeze": "freezeTenant", + }[action] + if field == "" { + return Tenant{}, Error{StatusCode: http.StatusBadRequest, Message: "unsupported tenant status action: " + action} + } + var out map[string]Tenant + err := c.graphQL(ctx, fmt.Sprintf(`mutation ChangeTenantStatus($id: ID!) { + %s(id: $id) { id name route status tags attributes created_by: createdBy updated_by: updatedBy created_at: createdAt updated_at: updatedAt } + }`, field), map[string]any{"id": id}, &out) + if err != nil { + return Tenant{}, err + } + return out[field], nil +} + +func (c *Client) UpsertEntity(ctx context.Context, entity Entity) error { + if entity.ID == "" { + _, err := c.CreateEntity(ctx, entity) + return err + } + if _, err := c.CreateEntity(ctx, entity); err == nil || !IsConflict(err) { + return err + } + _, err := c.UpdateEntity(ctx, entity.ID, entity) + return err +} + +func (c *Client) CreateEntity(ctx context.Context, entity Entity) (Entity, error) { + var out struct { + CreateEntity Entity `json:"createEntity"` + } + err := c.graphQL(ctx, `mutation CreateEntity($input: CreateEntityInput!) { + createEntity(input: $input) { id kind name tenant_id: tenantId status attributes created_at: createdAt updated_at: updatedAt } + }`, map[string]any{"input": entityCreateInput(entity)}, &out) + return out.CreateEntity, err +} + +func (c *Client) GetEntity(ctx context.Context, id string) (Entity, error) { + var out struct { + Entity Entity `json:"entity"` + } + err := c.graphQL(ctx, `query Entity($id: ID!) { + entity(id: $id) { id kind name tenant_id: tenantId status attributes created_at: createdAt updated_at: updatedAt } + }`, map[string]any{"id": id}, &out) + return out.Entity, err +} + +func (c *Client) UpdateEntity(ctx context.Context, id string, entity Entity) (Entity, error) { + var out struct { + UpdateEntity Entity `json:"updateEntity"` + } + err := c.graphQL(ctx, `mutation UpdateEntity($id: ID!, $input: UpdateEntityInput!) { + updateEntity(id: $id, input: $input) { id kind name tenant_id: tenantId status attributes created_at: createdAt updated_at: updatedAt } + }`, map[string]any{"id": id, "input": entityUpdateInput(entity)}, &out) + return out.UpdateEntity, err +} + +func (c *Client) UpsertGroup(ctx context.Context, group Group) error { + if group.ID == "" { + _, err := c.CreateGroup(ctx, group) + return err + } + if _, err := c.CreateGroup(ctx, group); err == nil || !IsConflict(err) { + return err + } + _, err := c.UpdateGroup(ctx, group.ID, group) + return err +} + +func (c *Client) CreateGroup(ctx context.Context, group Group) (Group, error) { + var out struct { + CreateGroup Group `json:"createGroup"` + } + err := c.graphQL(ctx, `mutation CreateGroup($input: CreateGroupInput!) { + createGroup(input: $input) { id name tenant_id: tenantId description parent_id: parentId status attributes created_at: createdAt updated_at: updatedAt } + }`, map[string]any{"input": groupCreateInput(group)}, &out) + return out.CreateGroup, err +} + +func (c *Client) GetGroup(ctx context.Context, id string) (Group, error) { + var out struct { + Group Group `json:"group"` + } + err := c.graphQL(ctx, `query Group($id: ID!) { + group(id: $id) { id name tenant_id: tenantId description parent_id: parentId status attributes created_at: createdAt updated_at: updatedAt } + }`, map[string]any{"id": id}, &out) + return out.Group, err +} + +func (c *Client) UpdateGroup(ctx context.Context, id string, group Group) (Group, error) { + var out struct { + UpdateGroup Group `json:"updateGroup"` + } + err := c.graphQL(ctx, `mutation UpdateGroup($id: ID!, $input: UpdateGroupInput!) { + updateGroup(id: $id, input: $input) { id name tenant_id: tenantId description parent_id: parentId status attributes created_at: createdAt updated_at: updatedAt } + }`, map[string]any{"id": id, "input": groupUpdateInput(group)}, &out) + return out.UpdateGroup, err +} + +func (c *Client) UpsertResource(ctx context.Context, resource Resource) error { + if resource.ID == "" { + _, err := c.CreateResource(ctx, resource) + return err + } + if _, err := c.CreateResource(ctx, resource); err == nil || !IsConflict(err) { + return err + } + _, err := c.UpdateResource(ctx, resource.ID, resource) + return err +} + +func (c *Client) CreateResource(ctx context.Context, resource Resource) (Resource, error) { + var out struct { + CreateResource Resource `json:"createResource"` + } + err := c.graphQL(ctx, `mutation CreateResource($input: CreateResourceInput!) { + createResource(input: $input) { id kind name tenant_id: tenantId owner_id: ownerId attributes created_at: createdAt updated_at: updatedAt } + }`, map[string]any{"input": resourceCreateInput(resource)}, &out) + return out.CreateResource, err +} + +func (c *Client) GetResource(ctx context.Context, id string) (Resource, error) { + var out struct { + Resource Resource `json:"resource"` + } + err := c.graphQL(ctx, `query Resource($id: ID!) { + resource(id: $id) { id kind name tenant_id: tenantId owner_id: ownerId attributes created_at: createdAt updated_at: updatedAt } + }`, map[string]any{"id": id}, &out) + return out.Resource, err +} + +func (c *Client) UpdateResource(ctx context.Context, id string, resource Resource) (Resource, error) { + var out struct { + UpdateResource Resource `json:"updateResource"` + } + err := c.graphQL(ctx, `mutation UpdateResource($id: ID!, $input: UpdateResourceInput!) { + updateResource(id: $id, input: $input) { id kind name tenant_id: tenantId owner_id: ownerId attributes created_at: createdAt updated_at: updatedAt } + }`, map[string]any{"id": id, "input": resourceUpdateInput(resource)}, &out) + return out.UpdateResource, err +} + +func (c *Client) DeleteTenant(ctx context.Context, id string) error { + return c.graphQL(ctx, `mutation DeleteTenant($id: ID!) { deleteTenant(id: $id) }`, map[string]any{"id": id}, nil) +} + +func (c *Client) ListTenants(ctx context.Context, q Query) (TenantList, error) { + var out struct { + Tenants TenantList `json:"tenants"` + } + err := c.graphQL(ctx, `query Tenants($q: String, $name: String, $route: String, $status: TenantStatus, $limit: Int, $offset: Int) { + tenants(q: $q, name: $name, route: $route, status: $status, limit: $limit, offset: $offset) { + total + items { id name route status tags attributes created_by: createdBy updated_by: updatedBy created_at: createdAt updated_at: updatedAt } + } + }`, queryVariables(q), &out) + return out.Tenants, err +} + +func (c *Client) CheckAuthz(ctx context.Context, req AuthzRequest) (AuthzResponse, error) { + var out struct { + AuthzCheck AuthzResponse `json:"authzCheck"` + } + err := c.graphQL(ctx, `mutation AuthzCheck($input: AuthzCheckInput!) { + authzCheck(input: $input) { allowed reason } + }`, map[string]any{"input": authzInput(req)}, &out) + return out.AuthzCheck, err +} + +func (c *Client) CheckAuthzWithToken(ctx context.Context, token string, req AuthzRequest) (AuthzResponse, error) { + var out struct { + AuthzCheck AuthzResponse `json:"authzCheck"` + } + err := c.graphQLWithToken(ctx, `mutation AuthzCheck($input: AuthzCheckInput!) { + authzCheck(input: $input) { allowed reason } + }`, map[string]any{"input": authzInput(req)}, &out, token) + return out.AuthzCheck, err +} + +func (c *Client) ListCapabilities(ctx context.Context) (CapabilityList, error) { + var out struct { + Actions CapabilityList `json:"actions"` + } + err := c.graphQL(ctx, `query Actions($limit: Int!) { + actions(limit: $limit) { total items { id name description } } + }`, map[string]any{"limit": 100}, &out) + return out.Actions, err +} + +func (c *Client) CapabilityID(ctx context.Context, name string) (string, error) { + list, err := c.ListCapabilities(ctx) + if err != nil { + return "", err + } + for _, capability := range list.Items { + if capability.Name == name { + return capability.ID, nil + } + } + return "", Error{StatusCode: http.StatusNotFound, Message: "capability " + name + " not found"} +} + +func (c *Client) CreateCapability(ctx context.Context, name, description string) (Capability, error) { + var out struct { + CreateAction Capability `json:"createAction"` + } + input := map[string]any{"name": name} + setIfNotEmpty(input, "description", description) + err := c.graphQL(ctx, `mutation CreateAction($input: CreateActionInput!) { + createAction(input: $input) { id name description } + }`, map[string]any{"input": input}, &out) + return out.CreateAction, err +} + +func (c *Client) AddCapabilityApplicability(ctx context.Context, actionID, objectKind, objectType string) (CapabilityApplicability, error) { + var out struct { + AddActionApplicability CapabilityApplicability `json:"addActionApplicability"` + } + input := map[string]any{ + "actionId": actionID, + "objectKind": objectKind, + } + setIfNotEmpty(input, "objectType", objectType) + err := c.graphQL(ctx, `mutation AddActionApplicability($input: AddActionApplicabilityInput!) { + addActionApplicability(input: $input) { + action_id: actionId + action_name: actionName + description + object_kind: objectKind + object_type: objectType + } + }`, map[string]any{"input": input}, &out) + return out.AddActionApplicability, err +} + +func (c *Client) ListActionAssignmentRules(ctx context.Context, spec ActionAssignmentRuleSpec) (ActionAssignmentRuleList, error) { + var out struct { + ActionAssignmentRules ActionAssignmentRuleList `json:"actionAssignmentRules"` + } + vars := map[string]any{"limit": 100, "offset": 0} + setIfNotEmpty(vars, "tenantId", spec.TenantID) + setIfNotEmpty(vars, "entityKind", spec.EntityKind) + setIfNotEmpty(vars, "actionName", spec.ActionName) + setIfNotEmpty(vars, "objectKind", spec.ObjectKind) + setIfNotEmpty(vars, "objectType", spec.ObjectType) + setIfNotEmpty(vars, "decision", spec.Decision) + err := c.graphQL(ctx, `query ActionAssignmentRules( + $tenantId: ID, + $entityKind: EntityKind, + $actionName: String, + $objectKind: String, + $objectType: String, + $decision: ActionAssignmentRuleDecision, + $limit: Int!, + $offset: Int! + ) { + actionAssignmentRules( + tenantId: $tenantId, + entityKind: $entityKind, + actionName: $actionName, + objectKind: $objectKind, + objectType: $objectType, + decision: $decision, + limit: $limit, + offset: $offset + ) { + total + items { + id + tenant_id: tenantId + entity_kind: entityKind + action_name: actionName + object_kind: objectKind + object_type: objectType + decision + is_absolute: isAbsolute + created_at: createdAt + } + } + }`, vars, &out) + return out.ActionAssignmentRules, err +} + +func (c *Client) CreateActionAssignmentRule(ctx context.Context, spec ActionAssignmentRuleSpec) (ActionAssignmentRule, error) { + var out struct { + CreateActionAssignmentRule ActionAssignmentRule `json:"createActionAssignmentRule"` + } + input := map[string]any{ + "entityKind": spec.EntityKind, + "actionName": spec.ActionName, + "objectKind": spec.ObjectKind, + "decision": spec.Decision, + "isAbsolute": spec.IsAbsolute, + } + setIfNotEmpty(input, "tenantId", spec.TenantID) + setIfNotEmpty(input, "objectType", spec.ObjectType) + err := c.graphQL(ctx, `mutation CreateActionAssignmentRule($input: CreateActionAssignmentRuleInput!) { + createActionAssignmentRule(input: $input) { + id + tenant_id: tenantId + entity_kind: entityKind + action_name: actionName + object_kind: objectKind + object_type: objectType + decision + is_absolute: isAbsolute + created_at: createdAt + } + }`, map[string]any{"input": input}, &out) + return out.CreateActionAssignmentRule, err +} + +func (c *Client) CreatePermissionBlock(ctx context.Context, block CreatePermissionBlock) (PermissionBlock, error) { + var out struct { + CreatePermissionBlock PermissionBlock `json:"createPermissionBlock"` + } + err := c.graphQL(ctx, `mutation CreatePermissionBlock($input: CreatePermissionBlockInput!) { + createPermissionBlock(input: $input) { + id tenant_id: tenantId scope_mode: scopeMode object_kind: objectKind object_type: objectType object_id: objectId group_id: groupId effect conditions + actions { id name description } + } + }`, map[string]any{"input": permissionBlockInput(block)}, &out) + return out.CreatePermissionBlock, err +} + +func (c *Client) CreateDirectPolicy(ctx context.Context, policy CreateDirectPolicy) (DirectPolicy, error) { + var out struct { + CreateDirectPolicy DirectPolicy `json:"createDirectPolicy"` + } + err := c.graphQL(ctx, `mutation CreateDirectPolicy($input: CreateDirectPolicyInput!) { + createDirectPolicy(input: $input) { + id tenant_id: tenantId subject_kind: subjectKind subject_id: subjectId permission_block_id: permissionBlockId created_at: createdAt + permission_block: permissionBlock { + id tenant_id: tenantId scope_mode: scopeMode object_kind: objectKind object_type: objectType object_id: objectId group_id: groupId effect conditions + actions { id name description } + } + } + }`, map[string]any{"input": directPolicyInput(policy)}, &out) + return out.CreateDirectPolicy, err +} + +func (c *Client) ListDirectPolicies(ctx context.Context, q DirectPolicyQuery) (DirectPolicyList, error) { + var out struct { + DirectPolicies DirectPolicyList `json:"directPolicies"` + } + err := c.graphQL(ctx, `query DirectPolicies($tenantId: ID, $subjectKind: SubjectKind, $subjectId: ID, $limit: Int, $offset: Int) { + directPolicies(tenantId: $tenantId, subjectKind: $subjectKind, subjectId: $subjectId, limit: $limit, offset: $offset) { + total + items { + id tenant_id: tenantId subject_kind: subjectKind subject_id: subjectId permission_block_id: permissionBlockId created_at: createdAt + permission_block: permissionBlock { + id tenant_id: tenantId scope_mode: scopeMode object_kind: objectKind object_type: objectType object_id: objectId group_id: groupId effect conditions + actions { id name description } + } + } + } + }`, directPolicyQueryVariables(q), &out) + return out.DirectPolicies, err +} + +func (c *Client) DeleteDirectPolicy(ctx context.Context, id string) error { + return c.graphQL(ctx, `mutation DeleteDirectPolicy($id: ID!) { deleteDirectPolicy(id: $id) }`, map[string]any{"id": id}, nil) +} + +func (c *Client) AuthorizedObjectIDs(ctx context.Context, q AuthorizedObjectIDsQuery) (AuthorizedObjectIDs, error) { + var out struct { + AuthorizedObjectIDs AuthorizedObjectIDs `json:"authorizedObjectIds"` + } + err := c.graphQL(ctx, `query AuthorizedObjectIDs($input: AuthorizedObjectIdsInput!) { + authorizedObjectIds(input: $input) { + ids + total + } + }`, authorizedObjectIDVariables(q), &out) + return out.AuthorizedObjectIDs, err +} + +func (c *Client) LoginPassword(ctx context.Context, identifier, secret string) (LoginResponse, error) { + var out LoginResponse + err := c.doWithToken(ctx, http.MethodPost, "/auth/login", LoginRequest{ + Identifier: identifier, + Secret: secret, + Kind: "password", + }, &out, "") + return out, err +} + +func (c *Client) Introspect(ctx context.Context, token string) (IntrospectionResponse, error) { + var out IntrospectionResponse + err := c.doWithToken(ctx, http.MethodGet, "/auth/introspect", nil, &out, token) + return out, err +} + +func (c *Client) DeleteEntity(ctx context.Context, id string) error { + return c.graphQL(ctx, `mutation DeleteEntity($id: ID!) { deleteEntity(id: $id) }`, map[string]any{"id": id}, nil) +} + +func (c *Client) CreatePassword(ctx context.Context, entityID, password string) error { + return c.graphQL(ctx, `mutation CreatePassword($entityId: ID!, $password: String!) { + createPassword(entityId: $entityId, password: $password) + }`, map[string]any{"entityId": entityID, "password": password}, nil) +} + +func (c *Client) CreateAPIKey(ctx context.Context, entityID, description string) (APIKeyResponse, error) { + var out struct { + CreateAPIKey APIKeyResponse `json:"createApiKey"` + } + err := c.graphQL(ctx, `mutation CreateAPIKey($entityId: ID!, $input: CreateApiKeyInput!) { + createApiKey(entityId: $entityId, input: $input) { + credentialId + key + expiresAt + } + }`, map[string]any{ + "entityId": entityID, + "input": map[string]any{ + "description": description, + }, + }, &out) + return out.CreateAPIKey, err +} + +func (c *Client) RevokeCredential(ctx context.Context, entityID, credentialID string) error { + return c.graphQL(ctx, `mutation RevokeCredential($entityId: ID!, $credentialId: ID!) { + revokeCredential(entityId: $entityId, credentialId: $credentialId) + }`, map[string]any{"entityId": entityID, "credentialId": credentialID}, nil) +} + +func (c *Client) ListEntities(ctx context.Context, q Query) (EntityList, error) { + var out struct { + Entities EntityList `json:"entities"` + } + err := c.graphQL(ctx, `query Entities($q: String, $kind: EntityKind, $tenantId: ID, $status: EntityStatus, $limit: Int, $offset: Int) { + entities(q: $q, kind: $kind, tenantId: $tenantId, status: $status, limit: $limit, offset: $offset) { + total + items { id kind name tenant_id: tenantId status attributes created_at: createdAt updated_at: updatedAt } + } + }`, objectQueryVariables(q), &out) + return out.Entities, err +} + +func (c *Client) DeleteGroup(ctx context.Context, id string) error { + return c.graphQL(ctx, `mutation DeleteGroup($id: ID!) { deleteGroup(id: $id) }`, map[string]any{"id": id}, nil) +} + +func (c *Client) ListGroups(ctx context.Context, q Query) (GroupList, error) { + var out struct { + Groups GroupList `json:"groups"` + } + err := c.graphQL(ctx, `query Groups($q: String, $tenantId: ID, $status: EntityStatus, $limit: Int, $offset: Int) { + groups(q: $q, tenantId: $tenantId, status: $status, limit: $limit, offset: $offset) { + total + items { id name tenant_id: tenantId description parent_id: parentId status attributes created_at: createdAt updated_at: updatedAt } + } + }`, objectQueryVariables(q), &out) + return out.Groups, err +} + +func (c *Client) DeleteResource(ctx context.Context, id string) error { + return c.graphQL(ctx, `mutation DeleteResource($id: ID!) { deleteResource(id: $id) }`, map[string]any{"id": id}, nil) +} + +func (c *Client) ListResources(ctx context.Context, q Query) (ResourceList, error) { + var out struct { + Resources ResourceList `json:"resources"` + } + vars := objectQueryVariables(q) + if q.Name != "" && q.Q == "" { + vars["q"] = q.Name + } + err := c.graphQL(ctx, `query Resources($q: String, $kind: String, $tenantId: ID, $limit: Int, $offset: Int) { + resources(q: $q, kind: $kind, tenantId: $tenantId, limit: $limit, offset: $offset) { + total + items { id kind name tenant_id: tenantId owner_id: ownerId attributes created_at: createdAt updated_at: updatedAt } + } + }`, vars, &out) + return out.Resources, err +} + +type graphQLRequest struct { + Query string `json:"query"` + Variables map[string]any `json:"variables,omitempty"` +} + +type graphQLErrorItem struct { + Message string `json:"message"` +} + +type graphQLResponse struct { + Data json.RawMessage `json:"data"` + Errors []graphQLErrorItem `json:"errors,omitempty"` +} + +func (c *Client) graphQL(ctx context.Context, query string, variables map[string]any, out any) error { + var response graphQLResponse + if err := c.do(ctx, http.MethodPost, "/graphql", graphQLRequest{Query: query, Variables: variables}, &response); err != nil { + return err + } + return decodeGraphQLResponse(response, out) +} + +func (c *Client) graphQLWithToken(ctx context.Context, query string, variables map[string]any, out any, token string) error { + var response graphQLResponse + if err := c.doWithToken(ctx, http.MethodPost, "/graphql", graphQLRequest{Query: query, Variables: variables}, &response, token); err != nil { + return err + } + return decodeGraphQLResponse(response, out) +} + +func decodeGraphQLResponse(response graphQLResponse, out any) error { + if len(response.Errors) > 0 { + return graphQLErr(response.Errors) + } + if out == nil { + return nil + } + if len(response.Data) == 0 { + return Error{StatusCode: http.StatusInternalServerError, Message: "atom GraphQL response did not include data"} + } + return json.Unmarshal(response.Data, out) +} + +func graphQLErr(errors []graphQLErrorItem) error { + messages := make([]string, 0, len(errors)) + for _, err := range errors { + if err.Message != "" { + messages = append(messages, err.Message) + } + } + message := strings.Join(messages, "; ") + lower := strings.ToLower(message) + status := http.StatusBadRequest + switch { + case strings.Contains(lower, "duplicate") || strings.Contains(lower, "already exists") || strings.Contains(lower, "unique"): + status = http.StatusConflict + case strings.Contains(lower, "not found"): + status = http.StatusNotFound + case strings.Contains(lower, "unauthenticated") || strings.Contains(lower, "authentication"): + status = http.StatusUnauthorized + case strings.Contains(lower, "forbidden") || strings.Contains(lower, "authorization") || strings.Contains(lower, "access denied"): + status = http.StatusForbidden + } + return Error{StatusCode: status, Message: message} +} + +func tenantCreateInput(tenant Tenant) map[string]any { + input := map[string]any{"name": tenant.Name} + setIfNotEmpty(input, "id", tenant.ID) + setIfNotEmpty(input, "route", tenant.Route) + if tenant.Tags != nil { + input["tags"] = tenant.Tags + } + if tenant.Attributes != nil { + input["attributes"] = tenant.Attributes + } + return input +} + +func tenantUpdateInput(tenant Tenant) map[string]any { + input := map[string]any{} + setIfNotEmpty(input, "name", tenant.Name) + setIfNotEmpty(input, "route", tenant.Route) + if tenant.Tags != nil { + input["tags"] = tenant.Tags + } + if tenant.Attributes != nil { + input["attributes"] = tenant.Attributes + } + return input +} + +func entityCreateInput(entity Entity) map[string]any { + input := map[string]any{"name": entity.Name} + setIfNotEmpty(input, "id", entity.ID) + setIfNotEmpty(input, "kind", entity.Kind) + setIfNotEmpty(input, "tenantId", entity.TenantID) + if entity.Attributes != nil { + input["attributes"] = entity.Attributes + } else { + input["attributes"] = map[string]any{} + } + return input +} + +func entityUpdateInput(entity Entity) map[string]any { + input := map[string]any{} + setIfNotEmpty(input, "name", entity.Name) + setIfNotEmpty(input, "status", entity.Status) + if entity.Attributes != nil { + input["attributes"] = entity.Attributes + } + return input +} + +func groupCreateInput(group Group) map[string]any { + input := map[string]any{"name": group.Name} + setIfNotEmpty(input, "id", group.ID) + setIfNotEmpty(input, "tenantId", group.TenantID) + setIfNotEmpty(input, "description", group.Description) + if group.Attributes != nil { + input["attributes"] = group.Attributes + } + return input +} + +func groupUpdateInput(group Group) map[string]any { + input := map[string]any{} + setIfNotEmpty(input, "name", group.Name) + setIfNotEmpty(input, "description", group.Description) + setIfNotEmpty(input, "status", group.Status) + if group.Attributes != nil { + input["attributes"] = group.Attributes + } + return input +} + +func resourceCreateInput(resource Resource) map[string]any { + input := map[string]any{"kind": resource.Kind} + setIfNotEmpty(input, "id", resource.ID) + setIfNotEmpty(input, "name", resource.Name) + setIfNotEmpty(input, "tenantId", resource.TenantID) + setIfNotEmpty(input, "ownerId", resource.OwnerID) + if resource.Attributes != nil { + input["attributes"] = resource.Attributes + } + return input +} + +func resourceUpdateInput(resource Resource) map[string]any { + input := map[string]any{} + setIfNotEmpty(input, "name", resource.Name) + if resource.Attributes != nil { + input["attributes"] = resource.Attributes + } + return input +} + +func authzInput(req AuthzRequest) map[string]any { + input := map[string]any{ + "subjectId": req.SubjectID, + "action": req.Action, + } + setIfNotEmpty(input, "resourceId", req.ResourceID) + setIfNotEmpty(input, "objectKind", req.ObjectKind) + setIfNotEmpty(input, "objectId", req.ObjectID) + if req.Context != nil { + input["context"] = req.Context + } + return input +} + +func permissionBlockInput(block CreatePermissionBlock) map[string]any { + input := map[string]any{ + "scopeMode": block.ScopeMode, + "actionIds": block.ActionIDs, + } + setIfNotEmpty(input, "tenantId", block.TenantID) + setIfNotEmpty(input, "objectKind", block.ObjectKind) + setIfNotEmpty(input, "objectType", block.ObjectType) + setIfNotEmpty(input, "objectId", block.ObjectID) + setIfNotEmpty(input, "groupId", block.GroupID) + setIfNotEmpty(input, "effect", block.Effect) + if block.Conditions != nil { + input["conditions"] = block.Conditions + } + return input +} + +func directPolicyInput(policy CreateDirectPolicy) map[string]any { + input := map[string]any{ + "subjectKind": policy.SubjectKind, + "subjectId": policy.SubjectID, + "permissionBlockId": policy.PermissionBlockID, + } + setIfNotEmpty(input, "tenantId", policy.TenantID) + return input +} + +func directPolicyQueryVariables(q DirectPolicyQuery) map[string]any { + vars := map[string]any{} + setIfNotEmpty(vars, "tenantId", q.TenantID) + setIfNotEmpty(vars, "subjectKind", q.SubjectKind) + setIfNotEmpty(vars, "subjectId", q.SubjectID) + if q.Limit > 0 { + vars["limit"] = int(q.Limit) + } + if q.Offset > 0 { + vars["offset"] = int(q.Offset) + } + return vars +} + +func authorizedObjectIDVariables(q AuthorizedObjectIDsQuery) map[string]any { + input := map[string]any{ + "subjectId": q.SubjectID, + "action": q.Action, + "objectKind": q.ObjectKind, + } + setIfNotEmpty(input, "objectType", q.ObjectType) + setIfNotEmpty(input, "tenantId", q.TenantID) + setIfNotEmpty(input, "q", q.Q) + if q.Limit > 0 { + input["limit"] = int(q.Limit) + } + if q.Offset > 0 { + input["offset"] = int(q.Offset) + } + return map[string]any{"input": input} +} + +func queryVariables(q Query) map[string]any { + vars := map[string]any{} + setIfNotEmpty(vars, "q", q.Q) + setIfNotEmpty(vars, "name", q.Name) + setIfNotEmpty(vars, "route", q.Route) + setIfNotEmpty(vars, "kind", q.Kind) + setIfNotEmpty(vars, "tenantId", q.TenantID) + setIfNotEmpty(vars, "status", q.Status) + if q.Limit > 0 { + vars["limit"] = int(q.Limit) + } + if q.Offset > 0 { + vars["offset"] = int(q.Offset) + } + return vars +} + +func objectQueryVariables(q Query) map[string]any { + vars := map[string]any{} + setIfNotEmpty(vars, "q", q.Q) + setIfNotEmpty(vars, "kind", q.Kind) + setIfNotEmpty(vars, "tenantId", q.TenantID) + setIfNotEmpty(vars, "status", q.Status) + if q.Limit > 0 { + vars["limit"] = int(q.Limit) + } + if q.Offset > 0 { + vars["offset"] = int(q.Offset) + } + return vars +} + +func setIfNotEmpty(values map[string]any, key, value string) { + if value != "" { + values[key] = value + } +} + +func (c *Client) do(ctx context.Context, method, path string, in, out any) error { + token := c.token + if token == "" && c.adminSecret != "" { + adminToken, err := c.loginAdmin(ctx) + if err != nil { + return err + } + token = adminToken + } + return c.doWithToken(ctx, method, path, in, out, token) +} + +func (c *Client) loginAdmin(ctx context.Context) (string, error) { + username := c.adminUsername + if username == "" { + username = defaultAdminUsername + } + resp, err := c.LoginPassword(ctx, username, c.adminSecret) + if err != nil { + return "", err + } + return resp.Token, nil +} + +func (c *Client) doWithToken(ctx context.Context, method, path string, in, out any, token string) error { + if c.baseURL == "" { + return Error{StatusCode: 0, Message: "atom URL is empty"} + } + + var body io.Reader + if in != nil { + data, err := json.Marshal(in) + if err != nil { + return err + } + body = bytes.NewReader(data) + } + + req, err := http.NewRequestWithContext(ctx, method, c.baseURL+path, body) + if err != nil { + return err + } + if in != nil { + req.Header.Set("Content-Type", "application/json") + } + req.Header.Set("Accept", "application/json") + if c.userAgent != "" { + req.Header.Set("User-Agent", c.userAgent) + } + if token != "" { + req.Header.Set("Authorization", "Bearer "+token) + } + + resp, err := c.httpClient.Do(req) + if err != nil { + return err + } + defer resp.Body.Close() + + if resp.StatusCode < http.StatusOK || resp.StatusCode >= http.StatusMultipleChoices { + msg, _ := io.ReadAll(io.LimitReader(resp.Body, 4096)) + return Error{StatusCode: resp.StatusCode, Message: strings.TrimSpace(string(msg))} + } + if out == nil || resp.StatusCode == http.StatusNoContent { + return nil + } + return json.NewDecoder(resp.Body).Decode(out) +} + +type Error struct { + StatusCode int + Message string +} + +func (e Error) Error() string { + if e.Message == "" { + return fmt.Sprintf("atom request failed with status %d", e.StatusCode) + } + return fmt.Sprintf("atom request failed with status %d: %s", e.StatusCode, e.Message) +} + +func IsConflict(err error) bool { + ae, ok := err.(Error) + return ok && ae.StatusCode == http.StatusConflict +} + +func IsNotFound(err error) bool { + ae, ok := err.(Error) + return ok && ae.StatusCode == http.StatusNotFound +} diff --git a/internal/atom/client_test.go b/internal/atom/client_test.go new file mode 100644 index 000000000..67fd50f08 --- /dev/null +++ b/internal/atom/client_test.go @@ -0,0 +1,111 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package atom + +import ( + "context" + "encoding/json" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" +) + +func TestUpsertResourceCreatesThenUpdatesOnConflict(t *testing.T) { + var operations []string + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost || r.URL.Path != atomGraphQLPath { + t.Fatalf("unexpected request: %s %s", r.Method, r.URL.Path) + } + body, _ := io.ReadAll(r.Body) + payload := string(body) + switch { + case strings.Contains(payload, "createResource"): + operations = append(operations, "createResource") + _ = json.NewEncoder(w).Encode(map[string]any{ + "errors": []map[string]string{{"message": "duplicate key value violates unique constraint"}}, + }) + return + case strings.Contains(payload, "updateResource"): + operations = append(operations, "updateResource") + _ = json.NewEncoder(w).Encode(map[string]any{ + "data": map[string]any{ + "updateResource": map[string]any{"id": "res-1", "kind": KindChannel, "name": "ch"}, + }, + }) + return + default: + t.Fatalf("unexpected GraphQL payload: %s", payload) + } + })) + defer srv.Close() + + client := NewClient(Config{URL: srv.URL, Timeout: time.Second}) + if err := client.UpsertResource(context.Background(), Resource{ID: "res-1", Kind: KindChannel, Name: "ch"}); err != nil { + t.Fatalf("upsert failed: %v", err) + } + + want := []string{"createResource", "updateResource"} + if len(operations) != len(want) { + t.Fatalf("unexpected operation count: got %v want %v", operations, want) + } + for i := range want { + if operations[i] != want[i] { + t.Fatalf("unexpected operations: got %v want %v", operations, want) + } + } +} + +func TestListResources(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost || r.URL.Path != atomGraphQLPath { + t.Fatalf("unexpected request: %s %s", r.Method, r.URL.Path) + } + var payload struct { + Variables map[string]any `json:"variables"` + } + if err := json.NewDecoder(r.Body).Decode(&payload); err != nil { + t.Fatalf("decode request: %v", err) + } + if payload.Variables["kind"] != KindRule || payload.Variables["tenantId"] != testDomainID { + t.Fatalf("unexpected variables: %+v", payload.Variables) + } + _ = json.NewEncoder(w).Encode(map[string]any{ + "data": map[string]any{ + "resources": map[string]any{ + "items": []Resource{{ID: "rule-1", Kind: KindRule, Name: "high-temp"}}, + "total": 1, + }, + }, + }) + })) + defer srv.Close() + + client := NewClient(Config{URL: srv.URL, Timeout: time.Second}) + got, err := client.ListResources(context.Background(), Query{Kind: KindRule, TenantID: testDomainID}) + if err != nil { + t.Fatalf("list failed: %v", err) + } + if got.Total != 1 || got.Items[0].ID != "rule-1" { + t.Fatalf("unexpected list: %+v", got) + } +} + +func TestLoadConfig(t *testing.T) { + t.Setenv("ATOM_URL", "http://atom:8080/") + t.Setenv("ATOM_ADMIN_TOKEN", "token") + t.Setenv("ATOM_ADMIN_USERNAME", "admin") + t.Setenv("ATOM_ADMIN_SECRET", "secret") + t.Setenv("ATOM_TIMEOUT", "3s") + + cfg := LoadConfig() + if cfg.URL != "http://atom:8080" || cfg.JWKSURL != "http://atom:8080/.well-known/jwks.json" || cfg.Token != "token" || cfg.AdminUsername != "admin" || cfg.AdminSecret != "secret" { + t.Fatalf("unexpected config: %+v", cfg) + } + if cfg.Timeout != 3*time.Second { + t.Fatalf("unexpected timeout: %s", cfg.Timeout) + } +} diff --git a/internal/atom/config.go b/internal/atom/config.go new file mode 100644 index 000000000..1fbec108b --- /dev/null +++ b/internal/atom/config.go @@ -0,0 +1,69 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package atom + +import ( + "os" + "strconv" + "strings" + "time" +) + +const ( + defaultTimeout = 5 * time.Second + defaultAdminUsername = "admin" +) + +// Config controls Magistrala's optional Atom integration. +type Config struct { + URL string + JWKSURL string + JWTIssuer string + JWTAudience string + Token string + AdminUsername string + AdminSecret string + Timeout time.Duration + UserAgent string +} + +// LoadConfig reads Atom integration settings from environment variables. +func LoadConfig() Config { + atomURL := strings.TrimRight(os.Getenv("ATOM_URL"), "/") + return Config{ + URL: atomURL, + JWKSURL: envString("ATOM_JWKS_URL", atomURL+"/.well-known/jwks.json"), + JWTIssuer: envString("ATOM_JWT_ISSUER", envString("ATOM_PUBLIC_URL", atomURL)), + JWTAudience: envString("ATOM_JWT_AUDIENCE", "magistrala"), + Token: envString("ATOM_SERVICE_TOKEN", os.Getenv("ATOM_ADMIN_TOKEN")), + AdminUsername: envString("ATOM_SERVICE_USERNAME", envString("ATOM_ADMIN_USERNAME", defaultAdminUsername)), + AdminSecret: envString("ATOM_SERVICE_SECRET", os.Getenv("ATOM_ADMIN_SECRET")), + Timeout: envDuration("ATOM_TIMEOUT", defaultTimeout), + UserAgent: "magistrala-atom-integration", + } +} + +func envString(key, fallback string) string { + v := strings.TrimSpace(os.Getenv(key)) + if v == "" { + return fallback + } + return v +} + +func envDuration(key string, fallback time.Duration) time.Duration { + v := os.Getenv(key) + if v == "" { + return fallback + } + d, err := time.ParseDuration(v) + if err == nil { + return d + } + seconds, err := strconv.Atoi(v) + if err == nil && seconds > 0 { + return time.Duration(seconds) * time.Second + } + return fallback +} diff --git a/internal/atom/constants.go b/internal/atom/constants.go new file mode 100644 index 000000000..290803f38 --- /dev/null +++ b/internal/atom/constants.go @@ -0,0 +1,42 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package atom + +const ( + atomActionRead = "read" + atomActionWrite = "write" + atomActionDelete = "delete" + atomActionManage = "manage" + atomActionPublish = "publish" + atomActionSubscribe = "subscribe" + atomActionExecute = "execute" + atomActionList = "list" +) + +const ( + atomStatusActive = "active" + atomStatusInactive = "inactive" + atomStatusEnabled = "enabled" + atomStatusDisabled = "disabled" + atomStatusFrozen = "frozen" + atomStatusSuspended = "suspended" + atomStatusDeleted = "deleted" +) + +const ( + atomKindDevice = "device" + atomKindGroup = "group" + atomKindHuman = "human" +) + +const ( + atomObjectKindEntity = "entity" + atomObjectKindGroup = "group" + atomObjectKindResource = "resource" + atomObjectKindTenant = "tenant" +) + +const atomScopeModeObject = "object" + +const atomGraphQLPath = "/graphql" diff --git a/internal/atom/grpc_compat.go b/internal/atom/grpc_compat.go new file mode 100644 index 000000000..53485f3fa --- /dev/null +++ b/internal/atom/grpc_compat.go @@ -0,0 +1,220 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package atom + +import ( + "context" + "net/http" + "strings" + + channelsv1 "github.com/absmach/magistrala/api/grpc/channels/v1" + clientsv1 "github.com/absmach/magistrala/api/grpc/clients/v1" + commonv1 "github.com/absmach/magistrala/api/grpc/common/v1" + domainsv1 "github.com/absmach/magistrala/api/grpc/domains/v1" + smqauthn "github.com/absmach/magistrala/pkg/authn" + "github.com/absmach/magistrala/pkg/connections" + "google.golang.org/grpc" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" +) + +type AtomClientsCompat struct { + Authn smqauthn.Authentication + Client *Client +} + +func NewClientsCompat(authn smqauthn.Authentication, client ...*Client) clientsv1.ClientsServiceClient { + atomClient := NewClient(LoadConfig()) + if len(client) > 0 && client[0] != nil { + atomClient = client[0] + } + return AtomClientsCompat{Authn: authn, Client: atomClient} +} + +func (c AtomClientsCompat) Authenticate(ctx context.Context, in *clientsv1.AuthnReq, _ ...grpc.CallOption) (*clientsv1.AuthnRes, error) { + token := in.GetToken() + if prefix, id, key, err := smqauthn.AuthUnpack(token); err == nil { + switch prefix { + case smqauthn.BasicAuth: + res, loginErr := c.Client.LoginPassword(ctx, id, key) + if loginErr == nil { + return &clientsv1.AuthnRes{Authenticated: true, Id: res.EntityID}, nil + } + if !isAtomUnauthorized(loginErr) { + return nil, loginErr + } + token = key + case smqauthn.DomainAuth: + token = key + case smqauthn.Unknown: + token = key + } + } + session, err := c.Authn.Authenticate(ctx, token) + if err != nil { + return nil, err + } + return &clientsv1.AuthnRes{Authenticated: true, Id: session.UserID}, nil +} + +func isAtomUnauthorized(err error) bool { + atomErr, ok := err.(Error) + return ok && atomErr.StatusCode == http.StatusUnauthorized +} + +func (c AtomClientsCompat) RetrieveEntity(context.Context, *commonv1.RetrieveEntityReq, ...grpc.CallOption) (*commonv1.RetrieveEntityRes, error) { + return nil, status.Error(codes.Unimplemented, "atom clients compatibility only supports Authenticate") +} + +func (c AtomClientsCompat) RetrieveEntities(context.Context, *commonv1.RetrieveEntitiesReq, ...grpc.CallOption) (*commonv1.RetrieveEntitiesRes, error) { + return nil, status.Error(codes.Unimplemented, "atom clients compatibility only supports Authenticate") +} + +func (c AtomClientsCompat) AddConnections(context.Context, *commonv1.AddConnectionsReq, ...grpc.CallOption) (*commonv1.AddConnectionsRes, error) { + return nil, status.Error(codes.Unimplemented, "atom clients compatibility only supports Authenticate") +} + +func (c AtomClientsCompat) RemoveConnections(context.Context, *commonv1.RemoveConnectionsReq, ...grpc.CallOption) (*commonv1.RemoveConnectionsRes, error) { + return nil, status.Error(codes.Unimplemented, "atom clients compatibility only supports Authenticate") +} + +func (c AtomClientsCompat) RemoveChannelConnections(context.Context, *clientsv1.RemoveChannelConnectionsReq, ...grpc.CallOption) (*clientsv1.RemoveChannelConnectionsRes, error) { + return nil, status.Error(codes.Unimplemented, "atom clients compatibility only supports Authenticate") +} + +func (c AtomClientsCompat) UnsetParentGroupFromClient(context.Context, *clientsv1.UnsetParentGroupFromClientReq, ...grpc.CallOption) (*clientsv1.UnsetParentGroupFromClientRes, error) { + return nil, status.Error(codes.Unimplemented, "atom clients compatibility only supports Authenticate") +} + +type AtomDomainsCompat struct { + Client *Client +} + +func NewDomainsCompat(client *Client) domainsv1.DomainsServiceClient { + return AtomDomainsCompat{Client: client} +} + +func (c AtomDomainsCompat) DeleteUserFromDomains(context.Context, *domainsv1.DeleteUserReq, ...grpc.CallOption) (*domainsv1.DeleteUserRes, error) { + return nil, status.Error(codes.Unimplemented, "atom domains compatibility does not delete user memberships") +} + +func (c AtomDomainsCompat) RetrieveStatus(ctx context.Context, in *commonv1.RetrieveEntityReq, _ ...grpc.CallOption) (*commonv1.RetrieveEntityRes, error) { + tenant, err := c.Client.GetTenant(ctx, in.GetId()) + if err != nil { + return nil, err + } + return &commonv1.RetrieveEntityRes{Entity: &commonv1.EntityBasic{ + Id: tenant.ID, + Status: atomStatusCode(tenant.Status), + }}, nil +} + +func (c AtomDomainsCompat) RetrieveIDByRoute(ctx context.Context, in *commonv1.RetrieveIDByRouteReq, _ ...grpc.CallOption) (*commonv1.RetrieveEntityRes, error) { + tenants, err := c.Client.ListTenants(ctx, Query{Route: in.GetRoute(), Limit: 1}) + if err != nil { + return nil, err + } + if len(tenants.Items) == 0 { + return nil, status.Errorf(codes.NotFound, "tenant route %q not found", in.GetRoute()) + } + tenant := tenants.Items[0] + return &commonv1.RetrieveEntityRes{Entity: &commonv1.EntityBasic{ + Id: tenant.ID, + Status: atomStatusCode(tenant.Status), + }}, nil +} + +type AtomChannelsCompat struct { + Client Authorizer + Atom *Client +} + +func NewChannelsCompat(client Authorizer) channelsv1.ChannelsServiceClient { + atomClient, _ := client.(*Client) + return AtomChannelsCompat{Client: client, Atom: atomClient} +} + +func (c AtomChannelsCompat) Authorize(ctx context.Context, in *channelsv1.AuthzReq, _ ...grpc.CallOption) (*channelsv1.AuthzRes, error) { + action := "subscribe" + if connections.ConnType(in.GetType()) == connections.Publish { + action = "publish" + } + subjectID := strings.TrimPrefix(in.GetClientId(), in.GetDomainId()+"_") + resp, err := c.Client.CheckAuthz(ctx, AuthzRequest{ + SubjectID: subjectID, + Action: action, + ResourceID: in.GetChannelId(), + ObjectKind: atomObjectKindResource, + ObjectID: in.GetChannelId(), + Context: map[string]any{ + "domain_id": in.GetDomainId(), + }, + }) + if err != nil { + return nil, err + } + return &channelsv1.AuthzRes{Authorized: resp.Allowed}, nil +} + +func (c AtomChannelsCompat) RemoveClientConnections(context.Context, *channelsv1.RemoveClientConnectionsReq, ...grpc.CallOption) (*channelsv1.RemoveClientConnectionsRes, error) { + return nil, status.Error(codes.Unimplemented, "atom channels compatibility only supports Authorize") +} + +func (c AtomChannelsCompat) UnsetParentGroupFromChannels(context.Context, *channelsv1.UnsetParentGroupFromChannelsReq, ...grpc.CallOption) (*channelsv1.UnsetParentGroupFromChannelsRes, error) { + return nil, status.Error(codes.Unimplemented, "atom channels compatibility only supports Authorize") +} + +func (c AtomChannelsCompat) RetrieveEntity(context.Context, *commonv1.RetrieveEntityReq, ...grpc.CallOption) (*commonv1.RetrieveEntityRes, error) { + return nil, status.Error(codes.Unimplemented, "atom channels compatibility requires a concrete Atom client") +} + +func (c AtomChannelsCompat) RetrieveIDByRoute(ctx context.Context, in *commonv1.RetrieveIDByRouteReq, _ ...grpc.CallOption) (*commonv1.RetrieveEntityRes, error) { + if c.Atom == nil { + return nil, status.Error(codes.Unimplemented, "atom channels compatibility requires a concrete Atom client") + } + resources, err := c.Atom.ListResources(ctx, Query{ + Kind: KindChannel, + TenantID: in.GetDomainId(), + Q: in.GetRoute(), + Limit: 20, + }) + if err != nil { + return nil, err + } + for _, resource := range resources.Items { + if resource.Name == in.GetRoute() || attrString(resource.Attributes, "route") == in.GetRoute() { + return &commonv1.RetrieveEntityRes{Entity: &commonv1.EntityBasic{ + Id: resource.ID, + DomainId: resource.TenantID, + Status: atomStatusCode(attrString(resource.Attributes, "status")), + }}, nil + } + } + return nil, status.Errorf(codes.NotFound, "channel route %q not found", in.GetRoute()) +} + +func atomStatusCode(value string) uint32 { + switch strings.ToLower(value) { + case "", atomStatusActive, atomStatusEnabled: + return 0 + case atomStatusInactive, atomStatusDisabled, atomStatusFrozen, atomStatusSuspended: + return 1 + case atomStatusDeleted: + return 2 + default: + return 0 + } +} + +func attrString(attrs Attributes, key string) string { + if attrs == nil { + return "" + } + value, ok := attrs[key] + if !ok || value == nil { + return "" + } + str, _ := value.(string) + return str +} diff --git a/internal/atom/grpc_compat_test.go b/internal/atom/grpc_compat_test.go new file mode 100644 index 000000000..d9d6015a7 --- /dev/null +++ b/internal/atom/grpc_compat_test.go @@ -0,0 +1,107 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package atom + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + "time" + + clientsv1 "github.com/absmach/magistrala/api/grpc/clients/v1" + smqauthn "github.com/absmach/magistrala/pkg/authn" +) + +type recordingAuthn struct { + called bool + token string + session smqauthn.Session + err error +} + +func (r *recordingAuthn) Authenticate(_ context.Context, token string) (smqauthn.Session, error) { + r.called = true + r.token = token + return r.session, r.err +} + +func TestAtomClientsCompatAuthenticatesBasicPasswordWithAtomLogin(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost || r.URL.Path != "/auth/login" { + t.Fatalf("unexpected request: %s %s", r.Method, r.URL.Path) + } + var got LoginRequest + if err := json.NewDecoder(r.Body).Decode(&got); err != nil { + t.Fatalf("decode login request: %v", err) + } + if got.Identifier != testEntityID || got.Secret != testDeviceSecret || got.Kind != "password" { + t.Fatalf("unexpected login request: %+v", got) + } + _ = json.NewEncoder(w).Encode(LoginResponse{ + Token: "jwt", + EntityID: testEntityID, + SessionID: "session-1", + ExpiresAt: time.Now().Add(time.Hour), + }) + })) + defer srv.Close() + + fallback := &recordingAuthn{} + compat := NewClientsCompat(fallback, NewClient(Config{URL: srv.URL, Timeout: time.Second})) + token := smqauthn.AuthPack(smqauthn.BasicAuth, testEntityID, testDeviceSecret) + + res, err := compat.Authenticate(context.Background(), &clientsv1.AuthnReq{Token: token}) + if err != nil { + t.Fatalf("authenticate basic password: %v", err) + } + if !res.GetAuthenticated() || res.GetId() != testEntityID { + t.Fatalf("unexpected response: %+v", res) + } + if fallback.called { + t.Fatal("token fallback should not be called after successful Atom password login") + } +} + +func TestAtomClientsCompatFallsBackToBearerTokenWhenBasicPasswordRejected(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + http.Error(w, "invalid credentials", http.StatusUnauthorized) + })) + defer srv.Close() + + fallback := &recordingAuthn{session: smqauthn.Session{UserID: "entity-2"}} + compat := NewClientsCompat(fallback, NewClient(Config{URL: srv.URL, Timeout: time.Second})) + token := smqauthn.AuthPack(smqauthn.BasicAuth, testEntityID, "atom_token") + + res, err := compat.Authenticate(context.Background(), &clientsv1.AuthnReq{Token: token}) + if err != nil { + t.Fatalf("authenticate fallback token: %v", err) + } + if !fallback.called || fallback.token != "atom_token" { + t.Fatalf("unexpected fallback call: called=%v token=%q", fallback.called, fallback.token) + } + if !res.GetAuthenticated() || res.GetId() != "entity-2" { + t.Fatalf("unexpected response: %+v", res) + } +} + +func TestAtomClientsCompatDoesNotHideAtomPasswordLoginFailures(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + http.Error(w, "atom unavailable", http.StatusInternalServerError) + })) + defer srv.Close() + + fallback := &recordingAuthn{session: smqauthn.Session{UserID: "entity-2"}} + compat := NewClientsCompat(fallback, NewClient(Config{URL: srv.URL, Timeout: time.Second})) + token := smqauthn.AuthPack(smqauthn.BasicAuth, testEntityID, testDeviceSecret) + + _, err := compat.Authenticate(context.Background(), &clientsv1.AuthnReq{Token: token}) + if err == nil { + t.Fatal("expected Atom login failure") + } + if fallback.called { + t.Fatal("token fallback should not be called for non-authentication Atom failures") + } +} diff --git a/internal/atom/mapping.go b/internal/atom/mapping.go new file mode 100644 index 000000000..8c1b7c8ce --- /dev/null +++ b/internal/atom/mapping.go @@ -0,0 +1,184 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package atom + +const ( + KindUser = "user" + KindClient = "client" + KindChannel = "channel" + KindRule = "rule" + KindAlarm = "alarm" + KindReport = "report" +) + +func TenantFromFields(f ObjectFields) Tenant { + return Tenant{ + ID: f.ID, + Name: f.Name, + Route: f.Route, + Tags: cloneStrings(f.Tags), + Status: tenantStatus(f.Status), + CreatedBy: f.CreatedBy, + UpdatedBy: f.UpdatedBy, + Attributes: compact(Attributes{ + "source": "magistrala", + "metadata": cloneMap(f.Metadata), + "created_at": timeString(f.CreatedAt), + "updated_at": timeString(f.UpdatedAt), + }), + } +} + +func EntityFromFields(f ObjectFields) Entity { + return Entity{ + ID: f.ID, + Kind: entityKind(f.Kind), + Name: f.Name, + TenantID: f.TenantID, + Status: entityStatus(f.Status), + Attributes: compact(Attributes{ + "source": "magistrala", + "magistrala_kind": f.Kind, + "tags": cloneStrings(f.Tags), + "metadata": cloneMap(f.Metadata), + "private_metadata": cloneMap(f.Private), + "parent_group_id": f.ParentID, + "created_at": timeString(f.CreatedAt), + "updated_at": timeString(f.UpdatedAt), + "updated_by": f.UpdatedBy, + }), + } +} + +func tenantStatus(status string) string { + switch status { + case atomStatusEnabled: + return atomStatusActive + case atomStatusDisabled: + return atomStatusInactive + case "freezed": + return atomStatusFrozen + case atomStatusDeleted: + return atomStatusDeleted + default: + return status + } +} + +func entityStatus(status string) string { + switch status { + case atomStatusEnabled: + return atomStatusActive + case atomStatusDisabled, atomStatusDeleted: + return atomStatusInactive + default: + return status + } +} + +func entityKind(kind string) string { + switch kind { + case KindUser: + return atomKindHuman + case KindClient: + return atomKindDevice + default: + return kind + } +} + +func GroupFromFields(f ObjectFields) Group { + return Group{ + ID: f.ID, + Name: f.Name, + TenantID: f.TenantID, + Description: f.Description, + ParentID: f.ParentID, + Status: entityStatus(f.Status), + Attributes: compact(Attributes{ + "source": "magistrala", + "parent_id": f.ParentID, + "tags": cloneStrings(f.Tags), + "metadata": cloneMap(f.Metadata), + "status": f.Status, + "created_at": timeString(f.CreatedAt), + "updated_at": timeString(f.UpdatedAt), + "updated_by": f.UpdatedBy, + }), + } +} + +func ResourceFromFields(f ObjectFields) Resource { + return Resource{ + ID: f.ID, + Kind: f.Kind, + Name: f.Name, + TenantID: f.TenantID, + OwnerID: f.OwnerID, + Attributes: compact(Attributes{ + "source": "magistrala", + "status": f.Status, + "route": f.Route, + "parent_group_id": f.ParentID, + "tags": cloneStrings(f.Tags), + "metadata": cloneMap(f.Metadata), + "created_at": timeString(f.CreatedAt), + "updated_at": timeString(f.UpdatedAt), + "updated_by": f.UpdatedBy, + }), + } +} + +func compact(attrs Attributes) Attributes { + for k, v := range attrs { + switch val := v.(type) { + case string: + if val == "" { + delete(attrs, k) + } + case []string: + if len(val) == 0 { + delete(attrs, k) + } + case map[string]any: + if len(val) == 0 { + delete(attrs, k) + } + case nil: + delete(attrs, k) + } + } + return attrs +} + +func cloneStrings(in []string) []string { + if len(in) == 0 { + return nil + } + out := make([]string, len(in)) + copy(out, in) + return out +} + +func cloneMap(in map[string]any) map[string]any { + if len(in) == 0 { + return nil + } + out := make(map[string]any, len(in)) + for k, v := range in { + out[k] = v + } + return out +} + +func timeString(t interface { + IsZero() bool + Format(string) string +}, +) string { + if t.IsZero() { + return "" + } + return t.Format("2006-01-02T15:04:05.999999999Z07:00") +} diff --git a/internal/atom/mapping_test.go b/internal/atom/mapping_test.go new file mode 100644 index 000000000..c5214f7b5 --- /dev/null +++ b/internal/atom/mapping_test.go @@ -0,0 +1,71 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package atom + +import ( + "testing" + "time" +) + +func TestTenantFromFields(t *testing.T) { + created := time.Date(2026, 4, 30, 10, 11, 12, 0, time.UTC) + got := TenantFromFields(ObjectFields{ + ID: "domain-1", + Name: "Acme", + Route: "acme", + Status: "enabled", + Tags: []string{"prod"}, + Metadata: map[string]any{"tier": "gold"}, + CreatedBy: "user-1", + CreatedAt: created, + }) + + if got.ID != "domain-1" || got.Name != "Acme" || got.Route != "acme" { + t.Fatalf("unexpected tenant: %+v", got) + } + if got.Attributes["source"] != "magistrala" { + t.Fatalf("missing source attribute: %+v", got.Attributes) + } + if got.Attributes["created_at"] != "2026-04-30T10:11:12Z" { + t.Fatalf("unexpected created_at: %v", got.Attributes["created_at"]) + } +} + +func TestEntityFromFields(t *testing.T) { + got := EntityFromFields(ObjectFields{ + ID: "client-1", + Kind: KindClient, + Name: "pump", + TenantID: "domain-1", + Status: "enabled", + ParentID: "group-1", + Tags: []string{"field"}, + }) + + if got.Kind != "device" || got.TenantID != "domain-1" { + t.Fatalf("unexpected entity: %+v", got) + } + if got.Attributes["magistrala_kind"] != KindClient { + t.Fatalf("missing magistrala kind: %+v", got.Attributes) + } + if got.Attributes["parent_group_id"] != "group-1" { + t.Fatalf("missing parent group: %+v", got.Attributes) + } +} + +func TestResourceFromFieldsOmitsEmptyValues(t *testing.T) { + got := ResourceFromFields(ObjectFields{ + ID: "channel-1", + Kind: KindChannel, + Name: "telemetry", + TenantID: "domain-1", + }) + + if _, ok := got.Attributes["tags"]; ok { + t.Fatalf("empty tags should be omitted: %+v", got.Attributes) + } + if got.Attributes["source"] != "magistrala" { + t.Fatalf("missing source attribute: %+v", got.Attributes) + } +} diff --git a/internal/atom/policy.go b/internal/atom/policy.go new file mode 100644 index 000000000..d4cf831ed --- /dev/null +++ b/internal/atom/policy.go @@ -0,0 +1,86 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package atom + +import ( + "context" + "strings" + + "github.com/absmach/magistrala/pkg/errors" + "github.com/absmach/magistrala/pkg/policies" +) + +type PolicyEvaluator struct { + client Authorizer +} + +func NewPolicyEvaluator(client Authorizer) PolicyEvaluator { + return PolicyEvaluator{client: client} +} + +func (pe PolicyEvaluator) CheckPolicy(ctx context.Context, pr policies.Policy) error { + res, err := pe.client.CheckAuthz(ctx, AuthzRequest{ + SubjectID: policySubjectID(pr), + Action: policyAction(pr), + ResourceID: policyResourceID(pr), + ObjectKind: policyObjectKind(pr), + ObjectID: pr.Object, + Context: map[string]any{ + "domain_id": pr.Domain, + "legacy_object_type": pr.ObjectType, + "legacy_relation": pr.Relation, + }, + }) + if err != nil { + return errors.Wrap(errors.ErrAuthorization, err) + } + if !res.Allowed { + return errors.ErrAuthorization + } + return nil +} + +func policySubjectID(pr policies.Policy) string { + if pr.Domain != "" { + return strings.TrimPrefix(pr.Subject, pr.Domain+"_") + } + return pr.Subject +} + +func policyAction(pr policies.Policy) string { + if pr.Permission != "" { + return CapabilityName(pr.Permission) + } + return CapabilityName(pr.Relation) +} + +func policyObjectKind(pr policies.Policy) string { + switch pr.ObjectType { + case policies.DomainType: + return atomObjectKindTenant + case policies.PlatformType: + return policies.PlatformType + case policies.ClientType: + return atomObjectKindEntity + case policies.GroupType: + return atomObjectKindGroup + case policies.ChannelType: + return atomObjectKindResource + case policies.RulesType: + return atomObjectKindResource + case policies.ReportsType: + return atomObjectKindResource + case policies.AlarmsType: + return atomObjectKindResource + default: + return pr.ObjectType + } +} + +func policyResourceID(pr policies.Policy) string { + if pr.ObjectType == policies.DomainType || pr.ObjectType == policies.PlatformType { + return "" + } + return pr.Object +} diff --git a/internal/atom/policy_service.go b/internal/atom/policy_service.go new file mode 100644 index 000000000..d95ec5e6c --- /dev/null +++ b/internal/atom/policy_service.go @@ -0,0 +1,266 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package atom + +import ( + "context" + stderrors "errors" + + "github.com/absmach/magistrala/pkg/policies" +) + +const policyPageLimit uint64 = 100 + +var errUnsupportedPolicyOperation = stderrors.New("atom policy service: unsupported policy operation") + +type policyClient interface { + Authorizer + AuthorizedObjectIDs(ctx context.Context, q AuthorizedObjectIDsQuery) (AuthorizedObjectIDs, error) +} + +type policyWriter interface { + CapabilityID(ctx context.Context, name string) (string, error) + CreatePermissionBlock(ctx context.Context, block CreatePermissionBlock) (PermissionBlock, error) + CreateDirectPolicy(ctx context.Context, policy CreateDirectPolicy) (DirectPolicy, error) + ListDirectPolicies(ctx context.Context, q DirectPolicyQuery) (DirectPolicyList, error) + DeleteDirectPolicy(ctx context.Context, id string) error +} + +type PolicyService struct { + client policyClient +} + +func NewPolicyService(client policyClient) PolicyService { + return PolicyService{client: client} +} + +func (ps PolicyService) AddPolicy(ctx context.Context, pr policies.Policy) error { + writer, ok := ps.client.(policyWriter) + if !ok { + return errUnsupportedPolicyOperation + } + capID, err := writer.CapabilityID(ctx, CapabilityName(pr.Permission)) + if err != nil { + return err + } + block, err := writer.CreatePermissionBlock(ctx, CreatePermissionBlock{ + TenantID: pr.Domain, + ScopeMode: policyGrantScopeMode(pr), + ObjectKind: policyGrantObjectKind(pr), + ObjectType: policyGrantObjectType(pr), + ObjectID: policyGrantObjectID(pr), + Effect: "allow", + Conditions: map[string]any{}, + ActionIDs: []string{capID}, + }) + if err != nil { + return err + } + _, err = writer.CreateDirectPolicy(ctx, CreateDirectPolicy{ + TenantID: pr.Domain, + SubjectKind: policyGrantSubjectKind(pr), + SubjectID: policySubjectID(pr), + PermissionBlockID: block.ID, + }) + return err +} + +func (ps PolicyService) AddPolicies(ctx context.Context, prs []policies.Policy) error { + for _, pr := range prs { + if err := ps.AddPolicy(ctx, pr); err != nil { + return err + } + } + return nil +} + +func (ps PolicyService) DeletePolicyFilter(ctx context.Context, pr policies.Policy) error { + writer, ok := ps.client.(policyWriter) + if !ok { + return errUnsupportedPolicyOperation + } + capID, err := writer.CapabilityID(ctx, CapabilityName(pr.Permission)) + if err != nil { + return err + } + page, err := writer.ListDirectPolicies(ctx, DirectPolicyQuery{ + TenantID: pr.Domain, + SubjectKind: policyGrantSubjectKind(pr), + SubjectID: policySubjectID(pr), + Limit: policyPageLimit, + }) + if err != nil { + return err + } + for _, policy := range page.Items { + if !directPolicyMatches(policy, capID, pr) { + continue + } + if err := writer.DeleteDirectPolicy(ctx, policy.ID); err != nil { + return err + } + } + return nil +} + +func (ps PolicyService) DeletePolicies(ctx context.Context, prs []policies.Policy) error { + for _, pr := range prs { + if err := ps.DeletePolicyFilter(ctx, pr); err != nil { + return err + } + } + return nil +} + +func (ps PolicyService) ListObjects(ctx context.Context, pr policies.Policy, _ string, limit uint64) (policies.PolicyPage, error) { + page, err := ps.ListAllObjects(ctx, pr) + if err != nil { + return policies.PolicyPage{}, err + } + if limit == 0 || uint64(len(page.Policies)) <= limit { + return page, nil + } + page.Policies = page.Policies[:limit] + return page, nil +} + +func (ps PolicyService) ListAllObjects(ctx context.Context, pr policies.Policy) (policies.PolicyPage, error) { + if !isSupportedObjectList(pr) { + return policies.PolicyPage{}, errUnsupportedPolicyOperation + } + + var ids []string + for offset := uint64(0); ; offset += policyPageLimit { + page, err := ps.client.AuthorizedObjectIDs(ctx, AuthorizedObjectIDsQuery{ + SubjectID: policySubjectID(pr), + Action: CapabilityName(pr.Permission), + ObjectKind: policyObjectKind(pr), + ObjectType: entityKind(KindClient), + TenantID: pr.Domain, + Limit: policyPageLimit, + Offset: offset, + }) + if err != nil { + return policies.PolicyPage{}, err + } + + ids = append(ids, page.IDs...) + + if uint64(len(page.IDs)) < policyPageLimit || offset+uint64(len(page.IDs)) >= page.Total { + break + } + } + + return policies.PolicyPage{Policies: ids}, nil +} + +func (ps PolicyService) CountObjects(ctx context.Context, pr policies.Policy) (uint64, error) { + page, err := ps.ListAllObjects(ctx, pr) + if err != nil { + return 0, err + } + return uint64(len(page.Policies)), nil +} + +func (ps PolicyService) ListSubjects(context.Context, policies.Policy, string, uint64) (policies.PolicyPage, error) { + return policies.PolicyPage{}, errUnsupportedPolicyOperation +} + +func (ps PolicyService) ListAllSubjects(context.Context, policies.Policy) (policies.PolicyPage, error) { + return policies.PolicyPage{}, errUnsupportedPolicyOperation +} + +func (ps PolicyService) CountSubjects(context.Context, policies.Policy) (uint64, error) { + return 0, errUnsupportedPolicyOperation +} + +func (ps PolicyService) ListPermissions(context.Context, policies.Policy, []string) (policies.Permissions, error) { + return nil, errUnsupportedPolicyOperation +} + +func isSupportedObjectList(pr policies.Policy) bool { + return pr.SubjectType == policies.UserType && + pr.Subject != "" && + pr.ObjectType == policies.ClientType && + pr.Permission == policies.ViewPermission +} + +func policyGrantSubjectKind(pr policies.Policy) string { + if pr.SubjectType == policies.GroupType || pr.SubjectKind == policies.GroupsKind { + return atomObjectKindGroup + } + return atomObjectKindEntity +} + +func policyGrantScopeMode(pr policies.Policy) string { + switch pr.ObjectType { + case policies.PlatformType: + return "platform" + case policies.DomainType: + return atomObjectKindTenant + default: + return atomScopeModeObject + } +} + +func policyGrantObjectKind(pr policies.Policy) string { + if policyGrantScopeMode(pr) != atomScopeModeObject { + return "" + } + switch pr.ObjectType { + case policies.ClientType: + return atomObjectKindEntity + case policies.GroupType: + return atomObjectKindGroup + } + return atomObjectKindResource +} + +func policyGrantObjectType(pr policies.Policy) string { + if policyGrantScopeMode(pr) != atomScopeModeObject { + return "" + } + if pr.ObjectType == policies.ClientType { + return atomObjectKindEntity + ":" + entityKind(KindClient) + } + switch pr.ObjectType { + case policies.ChannelType: + return "resource:" + KindChannel + case policies.RulesType: + return "resource:" + KindRule + case policies.ReportsType: + return "resource:" + KindReport + case policies.AlarmsType: + return "resource:" + KindAlarm + case policies.GroupType: + return "" + default: + return "" + } +} + +func policyGrantObjectID(pr policies.Policy) string { + if policyGrantScopeMode(pr) != "object" { + return "" + } + return policyResourceID(pr) +} + +func directPolicyMatches(policy DirectPolicy, actionID string, pr policies.Policy) bool { + block := policy.PermissionBlock + if block.ID == "" || block.ScopeMode != policyGrantScopeMode(pr) { + return false + } + if block.ObjectKind != policyGrantObjectKind(pr) || + block.ObjectType != policyGrantObjectType(pr) || + block.ObjectID != policyGrantObjectID(pr) { + return false + } + for _, action := range block.Actions { + if action.ID == actionID { + return true + } + } + return false +} diff --git a/internal/atom/policy_service_test.go b/internal/atom/policy_service_test.go new file mode 100644 index 000000000..ddcbcf587 --- /dev/null +++ b/internal/atom/policy_service_test.go @@ -0,0 +1,234 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package atom + +import ( + "context" + "testing" + + "github.com/absmach/magistrala/pkg/policies" +) + +type fakePolicyClient struct { + authorized AuthorizedObjectIDs + queries []AuthorizedObjectIDsQuery + capID string + blocks []CreatePermissionBlock + created []CreateDirectPolicy + policies []DirectPolicy + deleted []string +} + +func (f *fakePolicyClient) AuthorizedObjectIDs(_ context.Context, q AuthorizedObjectIDsQuery) (AuthorizedObjectIDs, error) { + f.queries = append(f.queries, q) + return f.authorized, nil +} + +func (f *fakePolicyClient) CheckAuthz(context.Context, AuthzRequest) (AuthzResponse, error) { + return AuthzResponse{Allowed: true}, nil +} + +func (f *fakePolicyClient) CapabilityID(context.Context, string) (string, error) { + if f.capID == "" { + return "cap-publish", nil + } + return f.capID, nil +} + +func (f *fakePolicyClient) CreatePermissionBlock(_ context.Context, block CreatePermissionBlock) (PermissionBlock, error) { + f.blocks = append(f.blocks, block) + return PermissionBlock{ + ID: "block-1", + TenantID: block.TenantID, + ScopeMode: block.ScopeMode, + ObjectKind: block.ObjectKind, + ObjectType: block.ObjectType, + ObjectID: block.ObjectID, + Effect: block.Effect, + Conditions: block.Conditions, + Actions: []Capability{{ID: block.ActionIDs[0]}}, + }, nil +} + +func (f *fakePolicyClient) CreateDirectPolicy(_ context.Context, policy CreateDirectPolicy) (DirectPolicy, error) { + f.created = append(f.created, policy) + return DirectPolicy{ID: "policy-1", PermissionBlockID: policy.PermissionBlockID}, nil +} + +func (f *fakePolicyClient) ListDirectPolicies(context.Context, DirectPolicyQuery) (DirectPolicyList, error) { + return DirectPolicyList{Items: f.policies, Total: uint64(len(f.policies))}, nil +} + +func (f *fakePolicyClient) DeleteDirectPolicy(_ context.Context, id string) error { + f.deleted = append(f.deleted, id) + return nil +} + +func TestPolicyServiceListAllObjectsUsesAtomAuthorizedObjectIds(t *testing.T) { + client := &fakePolicyClient{ + authorized: AuthorizedObjectIDs{IDs: []string{"client-2"}, Total: 1}, + } + svc := NewPolicyService(client) + + page, err := svc.ListAllObjects(context.Background(), policies.Policy{ + SubjectType: policies.UserType, + Subject: testDomainID + "_user-1", + Domain: testDomainID, + ObjectType: policies.ClientType, + Permission: policies.ViewPermission, + }) + if err != nil { + t.Fatalf("list objects failed: %v", err) + } + if len(page.Policies) != 1 || page.Policies[0] != "client-2" { + t.Fatalf("unexpected policies: %+v", page.Policies) + } + if len(client.queries) != 1 { + t.Fatalf("unexpected authorized object queries: %d", len(client.queries)) + } + query := client.queries[0] + if query.SubjectID != "user-1" || + query.Action != atomActionRead || + query.ObjectKind != atomObjectKindEntity || + query.ObjectType != entityKind(KindClient) || + query.TenantID != testDomainID { + t.Fatalf("unexpected authorized object query: %+v", query) + } +} + +func TestPolicyServiceAddPolicyCreatesInternalCapabilityPolicy(t *testing.T) { + client := &fakePolicyClient{capID: "cap-publish"} + svc := NewPolicyService(client) + + err := svc.AddPolicy(context.Background(), policies.Policy{ + Domain: testDomainID, + Subject: testDomainID + "_client-1", + SubjectType: policies.ClientType, + Object: "channel-1", + ObjectType: policies.ChannelType, + Permission: policies.PublishPermission, + }) + if err != nil { + t.Fatalf("add policy failed: %v", err) + } + if len(client.blocks) != 1 || len(client.created) != 1 { + t.Fatalf("expected one permission block and direct policy, got %d/%d", len(client.blocks), len(client.created)) + } + block := client.blocks[0] + if block.TenantID != testDomainID || + block.ScopeMode != atomScopeModeObject || + block.ObjectKind != atomObjectKindResource || + block.ObjectType != "resource:channel" || + block.ObjectID != "channel-1" || + block.Effect != "allow" || + len(block.ActionIDs) != 1 || + block.ActionIDs[0] != "cap-publish" { + t.Fatalf("unexpected permission block: %+v", block) + } + created := client.created[0] + if created.TenantID != testDomainID || + created.SubjectKind != atomObjectKindEntity || + created.SubjectID != "client-1" || + created.PermissionBlockID != "block-1" { + t.Fatalf("unexpected direct policy: %+v", created) + } +} + +func TestPolicyServiceAddPolicyCreatesGroupCapabilityPolicy(t *testing.T) { + client := &fakePolicyClient{capID: "cap-read"} + svc := NewPolicyService(client) + + err := svc.AddPolicy(context.Background(), policies.Policy{ + Domain: testDomainID, + Subject: testDomainID + "_user-1", + SubjectType: policies.UserType, + Object: "group-1", + ObjectType: policies.GroupType, + Permission: policies.ViewPermission, + }) + if err != nil { + t.Fatalf("add policy failed: %v", err) + } + if len(client.blocks) != 1 || len(client.created) != 1 { + t.Fatalf("expected one permission block and direct policy, got %d/%d", len(client.blocks), len(client.created)) + } + block := client.blocks[0] + if block.TenantID != testDomainID || + block.ScopeMode != atomScopeModeObject || + block.ObjectKind != atomObjectKindGroup || + block.ObjectType != "" || + block.ObjectID != "group-1" || + block.Effect != "allow" || + len(block.ActionIDs) != 1 || + block.ActionIDs[0] != "cap-read" { + t.Fatalf("unexpected permission block: %+v", block) + } + created := client.created[0] + if created.TenantID != testDomainID || + created.SubjectKind != atomObjectKindEntity || + created.SubjectID != "user-1" || + created.PermissionBlockID != "block-1" { + t.Fatalf("unexpected direct policy: %+v", created) + } +} + +func TestPolicyServiceDeletePolicyFilterRemovesMatchingCapabilityPolicy(t *testing.T) { + client := &fakePolicyClient{ + capID: "cap-subscribe", + policies: []DirectPolicy{ + { + ID: "keep", + PermissionBlock: PermissionBlock{ + ID: "keep-block", + ScopeMode: "object", + ObjectKind: atomObjectKindResource, + ObjectType: "resource:channel", + ObjectID: "channel-1", + Actions: []Capability{{ID: "cap-other"}}, + }, + }, + { + ID: "delete", + PermissionBlock: PermissionBlock{ + ID: "delete-block", + ScopeMode: "object", + ObjectKind: atomObjectKindResource, + ObjectType: "resource:channel", + ObjectID: "channel-1", + Actions: []Capability{{ID: "cap-subscribe"}}, + }, + }, + }, + } + svc := NewPolicyService(client) + + err := svc.DeletePolicyFilter(context.Background(), policies.Policy{ + Domain: testDomainID, + Subject: "domain-1_client-1", + SubjectType: policies.ClientType, + Object: "channel-1", + ObjectType: policies.ChannelType, + Permission: policies.SubscribePermission, + }) + if err != nil { + t.Fatalf("delete policy failed: %v", err) + } + if len(client.deleted) != 1 || client.deleted[0] != "delete" { + t.Fatalf("unexpected deleted policies: %+v", client.deleted) + } +} + +func TestPolicyServiceUnsupportedOperation(t *testing.T) { + svc := NewPolicyService(&fakePolicyClient{}) + + _, err := svc.ListAllObjects(context.Background(), policies.Policy{ + SubjectType: policies.UserType, + Subject: "user-1", + ObjectType: policies.ChannelType, + Permission: policies.ViewPermission, + }) + if err == nil { + t.Fatal("expected unsupported operation error") + } +} diff --git a/internal/atom/policy_test.go b/internal/atom/policy_test.go new file mode 100644 index 000000000..f47f2614d --- /dev/null +++ b/internal/atom/policy_test.go @@ -0,0 +1,55 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package atom_test + +import ( + "context" + "testing" + + "github.com/absmach/magistrala/internal/atom" + "github.com/absmach/magistrala/pkg/errors" + "github.com/absmach/magistrala/pkg/policies" + "github.com/stretchr/testify/assert" +) + +func TestPolicyEvaluatorCheckPolicy(t *testing.T) { + client := &authzClient{res: atom.AuthzResponse{Allowed: true}} + evaluator := atom.NewPolicyEvaluator(client) + + err := evaluator.CheckPolicy(context.Background(), policies.Policy{ + Domain: "domain-1", + Subject: "domain-1_user-1", + Permission: policies.ViewPermission, + ObjectType: policies.RulesType, + Object: "rule-1", + }) + + assert.NoError(t, err) + assert.Equal(t, atom.AuthzRequest{ + SubjectID: "user-1", + Action: "read", + ResourceID: "rule-1", + ObjectKind: "resource", + ObjectID: "rule-1", + Context: map[string]any{ + "domain_id": "domain-1", + "legacy_object_type": policies.RulesType, + "legacy_relation": "", + }, + }, client.req) +} + +func TestPolicyEvaluatorDenied(t *testing.T) { + client := &authzClient{res: atom.AuthzResponse{Allowed: false}} + evaluator := atom.NewPolicyEvaluator(client) + + err := evaluator.CheckPolicy(context.Background(), policies.Policy{ + Subject: "user-1", + Permission: policies.AdminPermission, + ObjectType: policies.PlatformType, + Object: policies.MagistralaObject, + }) + + assert.True(t, errors.Contains(err, errors.ErrAuthorization)) +} diff --git a/internal/atom/projector.go b/internal/atom/projector.go new file mode 100644 index 000000000..c86485ce6 --- /dev/null +++ b/internal/atom/projector.go @@ -0,0 +1,17 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package atom + +import "context" + +type Projector interface { + UpsertTenant(ctx context.Context, tenant Tenant) error + UpsertEntity(ctx context.Context, entity Entity) error + UpsertGroup(ctx context.Context, group Group) error + UpsertResource(ctx context.Context, resource Resource) error + DeleteTenant(ctx context.Context, id string) error + DeleteEntity(ctx context.Context, id string) error + DeleteGroup(ctx context.Context, id string) error + DeleteResource(ctx context.Context, id string) error +} diff --git a/internal/atom/service_tokens.go b/internal/atom/service_tokens.go new file mode 100644 index 000000000..7d5ce0764 --- /dev/null +++ b/internal/atom/service_tokens.go @@ -0,0 +1,245 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package atom + +import ( + "bufio" + "context" + "fmt" + "os" + "path/filepath" + "strings" +) + +const DefaultServiceEntityID = "00000000-0000-0000-0000-000000000003" + +type ServiceTokenSpec struct { + Name string + Env string + Description string +} + +type TokenProvisionOptions struct { + OutputPath string + ServiceEntityID string + Rotate string + Specs []ServiceTokenSpec +} + +type TokenProvisionResult struct { + OutputPath string + Preserved []string + Created []string + Rotated []string +} + +func DefaultServiceTokenSpecs() []ServiceTokenSpec { + return []ServiceTokenSpec{ + {Name: "fluxmq-auth", Env: "MG_ATOM_TOKEN_FLUXMQ_AUTH", Description: "Magistrala Docker Compose token for fluxmq-auth"}, + {Name: "fluxmq-node1", Env: "MG_ATOM_TOKEN_FLUXMQ_NODE1", Description: "Magistrala Docker Compose token for fluxmq-node1"}, + {Name: "fluxmq-node2", Env: "MG_ATOM_TOKEN_FLUXMQ_NODE2", Description: "Magistrala Docker Compose token for fluxmq-node2"}, + {Name: "fluxmq-node3", Env: "MG_ATOM_TOKEN_FLUXMQ_NODE3", Description: "Magistrala Docker Compose token for fluxmq-node3"}, + {Name: "journal", Env: "MG_ATOM_TOKEN_JOURNAL", Description: "Magistrala Docker Compose token for journal"}, + {Name: "notifications", Env: "MG_ATOM_TOKEN_NOTIFICATIONS", Description: "Magistrala Docker Compose token for notifications"}, + {Name: "timescale-reader", Env: "MG_ATOM_TOKEN_TIMESCALE_READER", Description: "Magistrala Docker Compose token for timescale-reader"}, + {Name: "re", Env: "MG_ATOM_TOKEN_RE", Description: "Magistrala Docker Compose token for rule engine"}, + {Name: "alarms", Env: "MG_ATOM_TOKEN_ALARMS", Description: "Magistrala Docker Compose token for alarms"}, + {Name: "reports", Env: "MG_ATOM_TOKEN_REPORTS", Description: "Magistrala Docker Compose token for reports"}, + {Name: "postgres-reader", Env: "MG_ATOM_TOKEN_POSTGRES_READER", Description: "Magistrala Docker Compose token for postgres-reader"}, + } +} + +func ProvisionServiceTokens(ctx context.Context, client *Client, opts TokenProvisionOptions) (TokenProvisionResult, error) { + if client == nil { + return TokenProvisionResult{}, fmt.Errorf("atom client is nil") + } + if strings.TrimSpace(opts.OutputPath) == "" { + return TokenProvisionResult{}, fmt.Errorf("token output path is required") + } + entityID := strings.TrimSpace(opts.ServiceEntityID) + if entityID == "" { + entityID = DefaultServiceEntityID + } + specs := opts.Specs + if len(specs) == 0 { + specs = DefaultServiceTokenSpecs() + } + rotate, err := normalizeRotation(opts.Rotate, specs) + if err != nil { + return TokenProvisionResult{}, err + } + + existing, err := readTokenEnvFile(opts.OutputPath) + if err != nil { + return TokenProvisionResult{}, err + } + + values := make(map[string]string, len(specs)) + result := TokenProvisionResult{OutputPath: opts.OutputPath} + for _, spec := range specs { + token := strings.TrimSpace(existing[spec.Env]) + shouldRotate := rotate["all"] || rotate[spec.Env] + if token != "" && !shouldRotate { + active, err := client.TokenActive(ctx, token) + if err == nil && active { + values[spec.Env] = token + result.Preserved = append(result.Preserved, spec.Env) + continue + } + } + if token != "" && shouldRotate { + credentialID, ok := CredentialIDFromAPIKey(token) + if ok { + if err := client.RevokeCredential(ctx, entityID, credentialID); err != nil && !IsNotFound(err) { + return TokenProvisionResult{}, fmt.Errorf("revoke %s credential %s: %w", spec.Env, credentialID, err) + } + } + } + created, err := client.CreateAPIKey(ctx, entityID, spec.Description) + if err != nil { + return TokenProvisionResult{}, fmt.Errorf("create %s token: %w", spec.Env, err) + } + if strings.TrimSpace(created.Key) == "" { + return TokenProvisionResult{}, fmt.Errorf("create %s token: atom returned an empty key", spec.Env) + } + values[spec.Env] = created.Key + if shouldRotate { + result.Rotated = append(result.Rotated, spec.Env) + } else { + result.Created = append(result.Created, spec.Env) + } + } + + if err := writeTokenEnvFile(opts.OutputPath, specs, values); err != nil { + return TokenProvisionResult{}, err + } + return result, nil +} + +func (c *Client) TokenActive(ctx context.Context, token string) (bool, error) { + res, err := c.Introspect(ctx, token) + if err != nil { + return false, err + } + return res.Active, nil +} + +func CredentialIDFromAPIKey(token string) (string, bool) { + rest, ok := strings.CutPrefix(strings.TrimSpace(token), "atom_") + if !ok { + return "", false + } + idHex, secretHex, ok := strings.Cut(rest, "_") + if !ok || len(idHex) != 32 || len(secretHex) != 64 || !isLowerHex(idHex) || !isLowerHex(secretHex) { + return "", false + } + return fmt.Sprintf("%s-%s-%s-%s-%s", idHex[0:8], idHex[8:12], idHex[12:16], idHex[16:20], idHex[20:32]), true +} + +func normalizeRotation(raw string, specs []ServiceTokenSpec) (map[string]bool, error) { + rotation := map[string]bool{} + raw = strings.TrimSpace(raw) + if raw == "" { + return rotation, nil + } + if strings.EqualFold(raw, "all") { + rotation["all"] = true + return rotation, nil + } + lookup := map[string]string{} + for _, spec := range specs { + lookup[strings.ToLower(spec.Env)] = spec.Env + lookup[strings.ToLower(strings.TrimPrefix(spec.Env, "MG_ATOM_TOKEN_"))] = spec.Env + lookup[strings.ToLower(strings.ReplaceAll(spec.Name, "-", "_"))] = spec.Env + } + key := strings.ToLower(strings.ReplaceAll(raw, "-", "_")) + env, ok := lookup[key] + if !ok { + return nil, fmt.Errorf("unknown token rotation target %q", raw) + } + rotation[env] = true + return rotation, nil +} + +func readTokenEnvFile(path string) (map[string]string, error) { + values := map[string]string{} + file, err := os.Open(path) + if err != nil { + if os.IsNotExist(err) { + return values, nil + } + return nil, fmt.Errorf("read token env file: %w", err) + } + defer file.Close() + + scanner := bufio.NewScanner(file) + for scanner.Scan() { + line := strings.TrimSpace(scanner.Text()) + if line == "" || strings.HasPrefix(line, "#") { + continue + } + key, value, ok := strings.Cut(line, "=") + if !ok { + continue + } + key = strings.TrimSpace(key) + value = strings.TrimSpace(value) + if key != "" { + values[key] = value + } + } + if err := scanner.Err(); err != nil { + return nil, fmt.Errorf("scan token env file: %w", err) + } + return values, nil +} + +func writeTokenEnvFile(path string, specs []ServiceTokenSpec, values map[string]string) error { + dir := filepath.Dir(path) + if err := os.MkdirAll(dir, 0o700); err != nil { + return fmt.Errorf("create token env directory: %w", err) + } + tmp, err := os.CreateTemp(dir, ".env.tokens-*") + if err != nil { + return fmt.Errorf("create token env temp file: %w", err) + } + tmpPath := tmp.Name() + defer func() { _ = os.Remove(tmpPath) }() + + if err := tmp.Chmod(0o600); err != nil { + _ = tmp.Close() + return fmt.Errorf("secure token env temp file: %w", err) + } + if _, err := fmt.Fprintln(tmp, "# Generated by atom-bootstrap provision-tokens. Do not commit."); err != nil { + _ = tmp.Close() + return err + } + for _, spec := range specs { + value := strings.TrimSpace(values[spec.Env]) + if value == "" { + _ = tmp.Close() + return fmt.Errorf("missing generated token for %s", spec.Env) + } + if _, err := fmt.Fprintf(tmp, "%s=%s\n", spec.Env, value); err != nil { + _ = tmp.Close() + return fmt.Errorf("write token env file: %w", err) + } + } + if err := tmp.Close(); err != nil { + return fmt.Errorf("close token env temp file: %w", err) + } + if err := os.Rename(tmpPath, path); err != nil { + return fmt.Errorf("replace token env file: %w", err) + } + return nil +} + +func isLowerHex(value string) bool { + for _, r := range value { + if (r < '0' || r > '9') && (r < 'a' || r > 'f') { + return false + } + } + return true +} diff --git a/internal/atom/service_tokens_test.go b/internal/atom/service_tokens_test.go new file mode 100644 index 000000000..b010ebc82 --- /dev/null +++ b/internal/atom/service_tokens_test.go @@ -0,0 +1,254 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package atom + +import ( + "context" + "encoding/json" + "fmt" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "strings" + "testing" + "time" +) + +func TestProvisionServiceTokensCreatesMissingToken(t *testing.T) { + fake := newFakeAtomTokenServer(t, nil) + defer fake.Close() + + output := filepath.Join(t.TempDir(), ".env.tokens") + result, err := ProvisionServiceTokens(context.Background(), fake.Client(), TokenProvisionOptions{ + OutputPath: output, + Specs: []ServiceTokenSpec{testTokenSpec()}, + }) + if err != nil { + t.Fatalf("provision tokens failed: %v", err) + } + if !containsString(result.Created, testTokenSpec().Env) { + t.Fatalf("expected token to be created, got result %+v", result) + } + if len(fake.created) != 1 { + t.Fatalf("unexpected create count: %d", len(fake.created)) + } + values, err := readTokenEnvFile(output) + if err != nil { + t.Fatalf("read token env file: %v", err) + } + if values[testTokenSpec().Env] == "" { + t.Fatalf("expected generated token in env file") + } + info, err := os.Stat(output) + if err != nil { + t.Fatalf("stat token env file: %v", err) + } + if got := info.Mode().Perm(); got != 0o600 { + t.Fatalf("unexpected token env permissions: got %s want -rw-------", got) + } + assertNoTempTokenFiles(t, filepath.Dir(output)) +} + +func TestProvisionServiceTokensPreservesExistingActiveToken(t *testing.T) { + token := apiKeyForCredentialID("11111111-1111-1111-1111-111111111111") + fake := newFakeAtomTokenServer(t, map[string]bool{token: true}) + defer fake.Close() + + output := filepath.Join(t.TempDir(), ".env.tokens") + if err := os.WriteFile(output, []byte(testTokenSpec().Env+"="+token+"\n"), 0o600); err != nil { + t.Fatalf("write existing token file: %v", err) + } + + result, err := ProvisionServiceTokens(context.Background(), fake.Client(), TokenProvisionOptions{ + OutputPath: output, + Specs: []ServiceTokenSpec{testTokenSpec()}, + }) + if err != nil { + t.Fatalf("provision tokens failed: %v", err) + } + if !containsString(result.Preserved, testTokenSpec().Env) { + t.Fatalf("expected token to be preserved, got result %+v", result) + } + if len(fake.created) != 0 { + t.Fatalf("expected no new API key, got %d", len(fake.created)) + } + values, err := readTokenEnvFile(output) + if err != nil { + t.Fatalf("read token env file: %v", err) + } + if got := values[testTokenSpec().Env]; got != token { + t.Fatalf("expected preserved token, got %q", got) + } +} + +func TestProvisionServiceTokensRotatesToken(t *testing.T) { + oldCredentialID := "11111111-1111-1111-1111-111111111111" + token := apiKeyForCredentialID(oldCredentialID) + fake := newFakeAtomTokenServer(t, map[string]bool{token: true}) + defer fake.Close() + + output := filepath.Join(t.TempDir(), ".env.tokens") + if err := os.WriteFile(output, []byte(testTokenSpec().Env+"="+token+"\n"), 0o600); err != nil { + t.Fatalf("write existing token file: %v", err) + } + + result, err := ProvisionServiceTokens(context.Background(), fake.Client(), TokenProvisionOptions{ + OutputPath: output, + Rotate: "journal", + Specs: []ServiceTokenSpec{testTokenSpec()}, + }) + if err != nil { + t.Fatalf("provision tokens failed: %v", err) + } + if !containsString(result.Rotated, testTokenSpec().Env) { + t.Fatalf("expected token to be rotated, got result %+v", result) + } + if !containsString(fake.revoked, oldCredentialID) { + t.Fatalf("expected old credential to be revoked, got %v", fake.revoked) + } + values, err := readTokenEnvFile(output) + if err != nil { + t.Fatalf("read token env file: %v", err) + } + if got := values[testTokenSpec().Env]; got == "" || got == token { + t.Fatalf("expected rotated token, got %q", got) + } +} + +func TestCredentialIDFromAPIKey(t *testing.T) { + want := "11111111-2222-3333-4444-555555555555" + got, ok := CredentialIDFromAPIKey(apiKeyForCredentialID(want)) + if !ok { + t.Fatalf("expected credential id to parse") + } + if got != want { + t.Fatalf("unexpected credential id: got %s want %s", got, want) + } + if _, ok := CredentialIDFromAPIKey("not-an-api-key"); ok { + t.Fatalf("expected invalid token to be rejected") + } +} + +type fakeAtomTokenServer struct { + t *testing.T + server *httptest.Server + active map[string]bool + + created []map[string]any + revoked []string + nextID int +} + +func newFakeAtomTokenServer(t *testing.T, active map[string]bool) *fakeAtomTokenServer { + t.Helper() + fake := &fakeAtomTokenServer{ + t: t, + active: active, + } + if fake.active == nil { + fake.active = map[string]bool{} + } + fake.server = httptest.NewServer(http.HandlerFunc(fake.handle)) + return fake +} + +func (f *fakeAtomTokenServer) Close() { + f.server.Close() +} + +func (f *fakeAtomTokenServer) Client() *Client { + return NewClient(Config{URL: f.server.URL, Token: "admin-token", Timeout: time.Second}) +} + +func (f *fakeAtomTokenServer) handle(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/auth/introspect": + token := strings.TrimPrefix(r.Header.Get("Authorization"), "Bearer ") + if err := json.NewEncoder(w).Encode(IntrospectionResponse{Active: f.active[token], EntityID: "entity-1"}); err != nil { + f.t.Fatalf("encode introspection response: %v", err) + } + case atomGraphQLPath: + f.handleGraphQL(w, r) + default: + f.t.Fatalf("unexpected request: %s %s", r.Method, r.URL.Path) + } +} + +func (f *fakeAtomTokenServer) handleGraphQL(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + f.t.Fatalf("unexpected GraphQL method: %s", r.Method) + } + var payload struct { + Query string `json:"query"` + Variables map[string]any `json:"variables"` + } + if err := json.NewDecoder(r.Body).Decode(&payload); err != nil { + f.t.Fatalf("decode GraphQL request: %v", err) + } + + switch { + case strings.Contains(payload.Query, "createApiKey"): + input := payload.Variables["input"].(map[string]any) + f.created = append(f.created, input) + f.nextID++ + credentialID := credentialIDForIndex(f.nextID) + key := apiKeyForCredentialID(credentialID) + f.active[key] = true + if err := json.NewEncoder(w).Encode(map[string]any{ + "data": map[string]any{ + "createApiKey": APIKeyResponse{ + CredentialID: credentialID, + Key: key, + }, + }, + }); err != nil { + f.t.Fatalf("encode create API key response: %v", err) + } + case strings.Contains(payload.Query, "revokeCredential"): + credentialID := payload.Variables["credentialId"].(string) + f.revoked = append(f.revoked, credentialID) + if err := json.NewEncoder(w).Encode(map[string]any{ + "data": map[string]any{"revokeCredential": true}, + }); err != nil { + f.t.Fatalf("encode revoke credential response: %v", err) + } + default: + f.t.Fatalf("unexpected GraphQL payload: %s", payload.Query) + } +} + +func testTokenSpec() ServiceTokenSpec { + return ServiceTokenSpec{Name: "journal", Env: "MG_ATOM_TOKEN_JOURNAL", Description: "test journal token"} +} + +func apiKeyForCredentialID(id string) string { + return "atom_" + strings.ReplaceAll(id, "-", "") + "_" + strings.Repeat("a", 64) +} + +func credentialIDForIndex(index int) string { + return fmt.Sprintf("aaaaaaaa-aaaa-aaaa-aaaa-%012d", index) +} + +func containsString(values []string, want string) bool { + for _, value := range values { + if value == want { + return true + } + } + return false +} + +func assertNoTempTokenFiles(t *testing.T, dir string) { + t.Helper() + entries, err := os.ReadDir(dir) + if err != nil { + t.Fatalf("read output directory: %v", err) + } + for _, entry := range entries { + if strings.HasPrefix(entry.Name(), ".env.tokens-") { + t.Fatalf("temporary token file was not removed: %s", entry.Name()) + } + } +} diff --git a/internal/atom/test_constants_test.go b/internal/atom/test_constants_test.go new file mode 100644 index 000000000..23f937d1c --- /dev/null +++ b/internal/atom/test_constants_test.go @@ -0,0 +1,10 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package atom + +const ( + testDeviceSecret = "device-secret" + testDomainID = "domain-1" + testEntityID = "entity-1" +) diff --git a/internal/atom/token.go b/internal/atom/token.go new file mode 100644 index 000000000..2cc3c1de5 --- /dev/null +++ b/internal/atom/token.go @@ -0,0 +1,159 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package atom + +import ( + "context" + "errors" + "fmt" + "io" + "net/http" + "strings" + "sync" + "time" + + "github.com/lestrrat-go/jwx/v2/jwk" + "github.com/lestrrat-go/jwx/v2/jws" + "github.com/lestrrat-go/jwx/v2/jwt" +) + +var ErrInvalidBearerToken = errors.New("invalid bearer token") + +const jwksCacheDuration = 5 * time.Minute + +type TokenVerifier struct { + jwksURL string + issuer string + audience string + httpClient *http.Client + client *Client + cache jwk.Set + cachedAt time.Time + mu sync.RWMutex +} + +func NewTokenVerifier(cfg Config) *TokenVerifier { + timeout := cfg.Timeout + if timeout == 0 { + timeout = defaultTimeout + } + return &TokenVerifier{ + jwksURL: cfg.JWKSURL, + issuer: cfg.JWTIssuer, + audience: cfg.JWTAudience, + client: NewClient(cfg), + httpClient: &http.Client{ + Timeout: timeout, + }, + } +} + +func (v *TokenVerifier) VerifyTokenClaims(ctx context.Context, token string) (TokenClaims, error) { + if strings.HasPrefix(token, "atom_") { + res, err := v.client.Introspect(ctx, token) + if err != nil || !res.Active || res.EntityID == "" { + return TokenClaims{}, ErrInvalidBearerToken + } + return TokenClaims{ + SubjectID: res.EntityID, + SessionID: res.SessionID, + TenantID: res.TenantID, + }, nil + } + + set, err := v.fetchJWKS(ctx, false) + if err != nil { + return TokenClaims{}, err + } + tkn, err := v.parseVerifiedToken(token, set) + if err != nil { + set, refreshErr := v.fetchJWKS(ctx, true) + if refreshErr == nil { + tkn, err = v.parseVerifiedToken(token, set) + } + } + if err != nil { + return TokenClaims{}, ErrInvalidBearerToken + } + claims := TokenClaims{ + SubjectID: tkn.Subject(), + ExpiresAt: tkn.Expiration().Unix(), + IssuedAt: tkn.IssuedAt().Unix(), + } + if claims.SubjectID == "" { + return TokenClaims{}, ErrInvalidBearerToken + } + if sid, ok := stringClaim(tkn, "sid"); ok { + claims.SessionID = sid + } + if tid, ok := stringClaim(tkn, "tid"); ok { + claims.TenantID = tid + } + return claims, nil +} + +func (v *TokenVerifier) fetchJWKS(ctx context.Context, force bool) (jwk.Set, error) { + if !force { + v.mu.RLock() + if v.cache != nil && time.Since(v.cachedAt) < jwksCacheDuration { + set := v.cache + v.mu.RUnlock() + return set, nil + } + v.mu.RUnlock() + } + req, err := http.NewRequestWithContext(ctx, http.MethodGet, v.jwksURL, nil) + if err != nil { + return nil, err + } + req.Header.Set("Accept", "application/json") + res, err := v.httpClient.Do(req) + if err != nil { + return nil, err + } + defer res.Body.Close() + if res.StatusCode != http.StatusOK { + body, _ := io.ReadAll(io.LimitReader(res.Body, 1024)) + return nil, fmt.Errorf("fetch atom jwks: status=%d body=%s", res.StatusCode, string(body)) + } + body, err := io.ReadAll(res.Body) + if err != nil { + return nil, err + } + set, err := jwk.Parse(body) + if err != nil { + return nil, err + } + v.mu.Lock() + v.cache = set + v.cachedAt = time.Now() + v.mu.Unlock() + return set, nil +} + +func (v *TokenVerifier) parseVerifiedToken(token string, set jwk.Set) (jwt.Token, error) { + options := []jwt.ParseOption{ + jwt.WithValidate(true), + jwt.WithKeySet(set, jws.WithInferAlgorithmFromKey(true)), + } + if v.issuer != "" { + options = append(options, jwt.WithIssuer(v.issuer)) + } + if v.audience != "" { + options = append(options, jwt.WithAudience(v.audience)) + } + return jwt.Parse( + []byte(token), + options..., + ) +} + +func stringClaim(tkn jwt.Token, name string) (string, bool) { + value, ok := tkn.Get(name) + if !ok { + return "", false + } + str, ok := value.(string) + return str, ok +} diff --git a/internal/atom/token_test.go b/internal/atom/token_test.go new file mode 100644 index 000000000..d60a86e54 --- /dev/null +++ b/internal/atom/token_test.go @@ -0,0 +1,134 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package atom + +import ( + "context" + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rand" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/lestrrat-go/jwx/v2/jwa" + "github.com/lestrrat-go/jwx/v2/jwk" + "github.com/lestrrat-go/jwx/v2/jwt" +) + +func TestTokenVerifierVerifiesAtomJWT(t *testing.T) { + token, jwksURL := signedAtomTokenServer(t, time.Now().Add(time.Hour)) + + claims, err := NewTokenVerifier(Config{ + JWKSURL: jwksURL, + JWTIssuer: "http://atom:8080", + JWTAudience: "magistrala", + Timeout: time.Second, + }).VerifyTokenClaims(context.Background(), token) + if err != nil { + t.Fatalf("verify token: %v", err) + } + if claims.SubjectID != "entity-1" || claims.SessionID != "session-1" || claims.TenantID != "tenant-1" { + t.Fatalf("unexpected claims: %+v", claims) + } +} + +func TestTokenVerifierRejectsUnsignedPayload(t *testing.T) { + _, jwksURL := signedAtomTokenServer(t, time.Now().Add(time.Hour)) + + _, err := NewTokenVerifier(Config{JWKSURL: jwksURL, Timeout: time.Second}).VerifyTokenClaims(context.Background(), "eyJhbGciOiJub25lIn0.eyJzdWIiOiJlbnRpdHktMSJ9.") + if err == nil { + t.Fatal("expected unsigned token to fail") + } +} + +func TestTokenVerifierRejectsExpiredToken(t *testing.T) { + token, jwksURL := signedAtomTokenServer(t, time.Now().Add(-time.Hour)) + + _, err := NewTokenVerifier(Config{JWKSURL: jwksURL, Timeout: time.Second}).VerifyTokenClaims(context.Background(), token) + if err == nil { + t.Fatal("expected expired token to fail") + } +} + +func TestTokenVerifierIntrospectsAtomAPIKey(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/auth/introspect" || r.Header.Get("Authorization") != "Bearer atom_test" { + t.Fatalf("unexpected request: %s %s", r.Method, r.URL.Path) + } + _ = json.NewEncoder(w).Encode(IntrospectionResponse{ + Active: true, + EntityID: "entity-2", + TenantID: "tenant-2", + }) + })) + defer srv.Close() + + claims, err := NewTokenVerifier(Config{URL: srv.URL, JWKSURL: srv.URL + "/jwks", Timeout: time.Second}).VerifyTokenClaims(context.Background(), "atom_test") + if err != nil { + t.Fatalf("verify api key: %v", err) + } + if claims.SubjectID != "entity-2" || claims.TenantID != "tenant-2" { + t.Fatalf("unexpected claims: %+v", claims) + } +} + +func signedAtomTokenServer(t *testing.T, expiry time.Time) (string, string) { + t.Helper() + privateKey, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + if err != nil { + t.Fatalf("generate key: %v", err) + } + privateJWK, err := jwk.FromRaw(privateKey) + if err != nil { + t.Fatalf("private jwk: %v", err) + } + if err := privateJWK.Set(jwk.AlgorithmKey, jwa.ES256); err != nil { + t.Fatalf("set private alg: %v", err) + } + if err := privateJWK.Set(jwk.KeyIDKey, "kid-1"); err != nil { + t.Fatalf("set private kid: %v", err) + } + publicJWK, err := jwk.FromRaw(privateKey.PublicKey) + if err != nil { + t.Fatalf("public jwk: %v", err) + } + if err := publicJWK.Set(jwk.AlgorithmKey, jwa.ES256); err != nil { + t.Fatalf("set public alg: %v", err) + } + if err := publicJWK.Set(jwk.KeyIDKey, "kid-1"); err != nil { + t.Fatalf("set public kid: %v", err) + } + set := jwk.NewSet() + if err := set.AddKey(publicJWK); err != nil { + t.Fatalf("add key: %v", err) + } + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "application/json") + if err := json.NewEncoder(w).Encode(set); err != nil { + t.Fatalf("write jwks: %v", err) + } + })) + t.Cleanup(srv.Close) + + tkn, err := jwt.NewBuilder(). + Issuer("http://atom:8080"). + Audience([]string{"magistrala"}). + Subject("entity-1"). + Claim("sid", "session-1"). + Claim("tid", "tenant-1"). + IssuedAt(time.Now()). + Expiration(expiry). + Build() + if err != nil { + t.Fatalf("build token: %v", err) + } + signed, err := jwt.Sign(tkn, jwt.WithKey(jwa.ES256, privateJWK)) + if err != nil { + t.Fatalf("sign token: %v", err) + } + return string(signed), srv.URL +} diff --git a/internal/atom/types.go b/internal/atom/types.go new file mode 100644 index 000000000..179f16f37 --- /dev/null +++ b/internal/atom/types.go @@ -0,0 +1,278 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package atom + +import "time" + +type Attributes map[string]any + +type Tenant struct { + ID string `json:"id,omitempty"` + Name string `json:"name"` + Route string `json:"route,omitempty"` + Tags []string `json:"tags,omitempty"` + Status string `json:"status,omitempty"` + Attributes Attributes `json:"attributes,omitempty"` + CreatedBy string `json:"created_by,omitempty"` + UpdatedBy string `json:"updated_by,omitempty"` + CreatedAt time.Time `json:"created_at,omitempty"` + UpdatedAt time.Time `json:"updated_at,omitempty"` +} + +type Entity struct { + ID string `json:"id,omitempty"` + Kind string `json:"kind"` + Name string `json:"name"` + TenantID string `json:"tenant_id,omitempty"` + Status string `json:"status,omitempty"` + Attributes Attributes `json:"attributes,omitempty"` + CreatedAt time.Time `json:"created_at,omitempty"` + UpdatedAt time.Time `json:"updated_at,omitempty"` +} + +type Group struct { + ID string `json:"id,omitempty"` + Name string `json:"name"` + TenantID string `json:"tenant_id,omitempty"` + Description string `json:"description,omitempty"` + ParentID string `json:"parent_id,omitempty"` + Status string `json:"status,omitempty"` + Attributes Attributes `json:"attributes,omitempty"` + CreatedAt time.Time `json:"created_at,omitempty"` + UpdatedAt time.Time `json:"updated_at,omitempty"` +} + +type Resource struct { + ID string `json:"id,omitempty"` + Kind string `json:"kind"` + Name string `json:"name"` + TenantID string `json:"tenant_id,omitempty"` + OwnerID string `json:"owner_id,omitempty"` + Attributes Attributes `json:"attributes,omitempty"` + CreatedAt time.Time `json:"created_at,omitempty"` + UpdatedAt time.Time `json:"updated_at,omitempty"` +} + +type Query struct { + IDs []string + Q string + Kind string + TenantID string + Name string + Route string + Status string + Limit uint64 + Offset uint64 +} + +type AuthzRequest struct { + SubjectID string `json:"subject_id"` + Action string `json:"action"` + ResourceID string `json:"resource_id,omitempty"` + ObjectKind string `json:"object_kind,omitempty"` + ObjectID string `json:"object_id,omitempty"` + Context map[string]any `json:"context,omitempty"` +} + +type AuthzResponse struct { + Allowed bool `json:"allowed"` + Reason string `json:"reason,omitempty"` +} + +type Capability struct { + ID string `json:"id"` + Name string `json:"name"` + Description string `json:"description,omitempty"` +} + +type CapabilityList struct { + Items []Capability `json:"items"` + Total int64 `json:"total,omitempty"` +} + +type CapabilityApplicability struct { + ActionID string `json:"action_id"` + ActionName string `json:"action_name"` + Description string `json:"description,omitempty"` + ObjectKind string `json:"object_kind"` + ObjectType string `json:"object_type,omitempty"` +} + +type CapabilityApplicabilitySpec struct { + ActionName string + Description string + ObjectKind string + ObjectType string +} + +type ActionAssignmentRule struct { + ID string `json:"id"` + TenantID string `json:"tenant_id,omitempty"` + EntityKind string `json:"entity_kind"` + ActionName string `json:"action_name"` + ObjectKind string `json:"object_kind"` + ObjectType string `json:"object_type,omitempty"` + Decision string `json:"decision"` + IsAbsolute bool `json:"is_absolute"` + CreatedAt string `json:"created_at,omitempty"` +} + +type ActionAssignmentRuleList struct { + Items []ActionAssignmentRule `json:"items"` + Total int64 `json:"total,omitempty"` +} + +type ActionAssignmentRuleSpec struct { + TenantID string + EntityKind string + ActionName string + ObjectKind string + ObjectType string + Decision string + IsAbsolute bool +} + +type PermissionBlock struct { + ID string `json:"id"` + TenantID string `json:"tenant_id,omitempty"` + ScopeMode string `json:"scope_mode"` + ObjectKind string `json:"object_kind,omitempty"` + ObjectType string `json:"object_type,omitempty"` + ObjectID string `json:"object_id,omitempty"` + GroupID string `json:"group_id,omitempty"` + Effect string `json:"effect"` + Conditions map[string]any `json:"conditions,omitempty"` + Actions []Capability `json:"actions,omitempty"` +} + +type CreatePermissionBlock struct { + TenantID string `json:"tenant_id,omitempty"` + ScopeMode string `json:"scope_mode"` + ObjectKind string `json:"object_kind,omitempty"` + ObjectType string `json:"object_type,omitempty"` + ObjectID string `json:"object_id,omitempty"` + GroupID string `json:"group_id,omitempty"` + Effect string `json:"effect,omitempty"` + Conditions map[string]any `json:"conditions,omitempty"` + ActionIDs []string `json:"action_ids"` +} + +type DirectPolicy struct { + ID string `json:"id"` + TenantID string `json:"tenant_id,omitempty"` + SubjectKind string `json:"subject_kind"` + SubjectID string `json:"subject_id"` + PermissionBlockID string `json:"permission_block_id"` + PermissionBlock PermissionBlock `json:"permission_block,omitempty"` + CreatedAt time.Time `json:"created_at,omitempty"` +} + +type CreateDirectPolicy struct { + TenantID string `json:"tenant_id,omitempty"` + SubjectKind string `json:"subject_kind"` + SubjectID string `json:"subject_id"` + PermissionBlockID string `json:"permission_block_id"` +} + +type DirectPolicyQuery struct { + TenantID string + SubjectKind string + SubjectID string + Limit uint64 + Offset uint64 +} + +type DirectPolicyList struct { + Items []DirectPolicy `json:"items"` + Total uint64 `json:"total"` +} + +type AuthorizedObjectIDsQuery struct { + SubjectID string + Action string + ObjectKind string + ObjectType string + TenantID string + Q string + Limit uint64 + Offset uint64 +} + +type AuthorizedObjectIDs struct { + IDs []string `json:"ids"` + Total uint64 `json:"total"` +} + +type TokenClaims struct { + SubjectID string `json:"sub"` + SessionID string `json:"sid,omitempty"` + TenantID string `json:"tid,omitempty"` + ExpiresAt int64 `json:"exp,omitempty"` + IssuedAt int64 `json:"iat,omitempty"` +} + +type IntrospectionResponse struct { + Active bool `json:"active"` + EntityID string `json:"entity_id"` + TenantID string `json:"tenant_id,omitempty"` + SessionID string `json:"session_id,omitempty"` +} + +type LoginRequest struct { + Identifier string `json:"identifier"` + Secret string `json:"secret"` + Kind string `json:"kind,omitempty"` +} + +type LoginResponse struct { + Token string `json:"token"` + EntityID string `json:"entity_id"` + SessionID string `json:"session_id"` + ExpiresAt time.Time `json:"expires_at"` +} + +type APIKeyResponse struct { + CredentialID string `json:"credentialId"` + Key string `json:"key"` + ExpiresAt *time.Time `json:"expiresAt,omitempty"` +} + +type ResourceList struct { + Items []Resource `json:"items"` + Total uint64 `json:"total"` +} + +type TenantList struct { + Items []Tenant `json:"items"` + Total uint64 `json:"total"` +} + +type EntityList struct { + Items []Entity `json:"items"` + Total uint64 `json:"total"` +} + +type GroupList struct { + Items []Group `json:"items"` + Total uint64 `json:"total"` +} + +type ObjectFields struct { + ID string + Kind string + Name string + TenantID string + OwnerID string + Status string + Route string + ParentID string + Tags []string + Metadata map[string]any + Private map[string]any + CreatedBy string + UpdatedBy string + CreatedAt time.Time + UpdatedAt time.Time + Description string +} diff --git a/journal/middleware/authorization.go b/journal/middleware/authorization.go index 9845edbdc..065b01fb2 100644 --- a/journal/middleware/authorization.go +++ b/journal/middleware/authorization.go @@ -39,7 +39,7 @@ func (am *authorizationMiddleware) RetrieveAll(ctx context.Context, session smqa permission := readPermission objectType := page.EntityType.String() object := page.EntityID - subject := session.DomainUserID + subject := subjectID(session) // If the entity is a user, we need to check if the user is an admin if page.EntityType.String() == policies.UserType { @@ -70,7 +70,7 @@ func (am *authorizationMiddleware) RetrieveClientTelemetry(ctx context.Context, Domain: session.DomainID, SubjectType: policies.UserType, SubjectKind: policies.UsersKind, - Subject: session.DomainUserID, + Subject: subjectID(session), Permission: readPermission, ObjectType: policies.ClientType, Object: clientID, @@ -82,3 +82,10 @@ func (am *authorizationMiddleware) RetrieveClientTelemetry(ctx context.Context, return am.svc.RetrieveClientTelemetry(ctx, session, clientID) } + +func subjectID(session smqauthn.Session) string { + if session.UserID != "" { + return session.UserID + } + return session.DomainUserID +} diff --git a/notifications/README.md b/notifications/README.md index 09ec39280..3ce05eb06 100644 --- a/notifications/README.md +++ b/notifications/README.md @@ -9,19 +9,19 @@ This service listens to invitation events from the domains service and sends ema - Someone accepts their domain invitation (`invitation.accept`) - Someone rejects their domain invitation (`invitation.reject`) -The service uses gRPC to fetch user information from the users service and sends styled email notifications using SMTP. +The service fetches user information from Atom entities and sends styled email notifications using SMTP. ## Features - **Event-Driven**: Listens to invitation events from the event store (NATS/RabbitMQ) -- **gRPC Integration**: Fetches user details (name, email) from the users service +- **Atom Integration**: Fetches user details (name, email) from Atom - **Beautiful Email Templates**: Styled HTML email templates with Magistrala branding (#083662) - **Configurable**: Email server settings and templates are fully configurable ## Architecture ``` -domains service → event store → notifications service → users service (gRPC) +domains service → event store → notifications service → Atom ↓ SMTP Server → Email Recipients ``` @@ -49,12 +49,12 @@ The service is configured using environment variables: - `MG_EMAIL_ACCEPTANCE_TEMPLATE` - Path to acceptance email template - `MG_EMAIL_REJECTION_TEMPLATE` - Path to rejection email template -### gRPC Configuration (Users Service) -- `MG_USERS_GRPC_URL` - Users service gRPC URL -- `MG_USERS_GRPC_TIMEOUT` - gRPC request timeout -- `MG_USERS_GRPC_CLIENT_CERT` - Client certificate path -- `MG_USERS_GRPC_CLIENT_KEY` - Client key path -- `MG_USERS_GRPC_SERVER_CA_CERTS` - Server CA certificates path +### Atom Configuration +- `ATOM_URL` - Atom HTTP URL +- `ATOM_SERVICE_TOKEN` - Service bearer token, if provisioned +- `ATOM_SERVICE_USERNAME` / `ATOM_SERVICE_SECRET` - Service credential fallback +- `ATOM_ADMIN_USERNAME` / `ATOM_ADMIN_SECRET` - Admin credential fallback +- `ATOM_TIMEOUT` - Atom HTTP timeout ## Running the Service @@ -108,6 +108,6 @@ The service consists of: ## Dependencies -- Users service (gRPC) - for fetching user information +- Atom - for fetching user information - Event store (NATS/RabbitMQ) - for receiving invitation events - SMTP server - for sending emails diff --git a/notifications/emailer/atom_users.go b/notifications/emailer/atom_users.go new file mode 100644 index 000000000..bffc61a59 --- /dev/null +++ b/notifications/emailer/atom_users.go @@ -0,0 +1,70 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package emailer + +import ( + "context" + "fmt" + + "github.com/absmach/magistrala/internal/atom" +) + +// AtomUserResolver resolves notification users from Atom entities. +type AtomUserResolver struct { + client *atom.Client +} + +// NewAtomUserResolver creates an Atom-backed notification user resolver. +func NewAtomUserResolver(client *atom.Client) AtomUserResolver { + return AtomUserResolver{client: client} +} + +// FetchUsers loads users by ID from Atom. +func (r AtomUserResolver) FetchUsers(ctx context.Context, userIDs []string) (map[string]User, error) { + users := make(map[string]User, len(userIDs)) + for _, userID := range userIDs { + if userID == "" { + continue + } + entity, err := r.client.GetEntity(ctx, userID) + if err != nil { + return nil, fmt.Errorf("fetch atom entity %s: %w", userID, err) + } + users[userID] = atomEntityUser(entity) + } + return users, nil +} + +func atomEntityUser(entity atom.Entity) User { + email := attrString(entity.Attributes, "email") + if email == "" { + email = attrString(entity.Attributes, "primary_email") + } + username := attrString(entity.Attributes, "username") + if username == "" { + username = entity.Name + } + return User{ + ID: entity.ID, + Email: email, + Username: username, + FirstName: attrString(entity.Attributes, "first_name"), + LastName: attrString(entity.Attributes, "last_name"), + } +} + +func attrString(attrs atom.Attributes, key string) string { + value, ok := attrs[key] + if !ok || value == nil { + return "" + } + switch typed := value.(type) { + case string: + return typed + case fmt.Stringer: + return typed.String() + default: + return fmt.Sprint(typed) + } +} diff --git a/notifications/emailer/emailer.go b/notifications/emailer/emailer.go index 5cbad941f..56affef5c 100644 --- a/notifications/emailer/emailer.go +++ b/notifications/emailer/emailer.go @@ -9,7 +9,6 @@ import ( "unicode" "unicode/utf8" - grpcUsersV1 "github.com/absmach/magistrala/api/grpc/users/v1" "github.com/absmach/magistrala/internal/email" "github.com/absmach/magistrala/notifications" "github.com/absmach/magistrala/pkg/errors" @@ -27,8 +26,22 @@ const ( var _ notifications.Notifier = (*notifier)(nil) +// User contains the user fields needed by notification templates. +type User struct { + ID string + Email string + Username string + FirstName string + LastName string +} + +// UserResolver resolves notification principals from the active identity store. +type UserResolver interface { + FetchUsers(ctx context.Context, userIDs []string) (map[string]User, error) +} + type notifier struct { - usersClient grpcUsersV1.UsersServiceClient + users UserResolver agents map[notifications.NotificationType]*email.Agent fromName string domainAltName string @@ -49,7 +62,7 @@ type Config struct { } // New creates a new email notifier. -func New(usersClient grpcUsersV1.UsersServiceClient, cfg Config) (notifications.Notifier, error) { +func New(users UserResolver, cfg Config) (notifications.Notifier, error) { templates := map[notifications.NotificationType]string{ notifications.Invitation: cfg.InvitationTemplate, notifications.Acceptance: cfg.AcceptanceTemplate, @@ -75,7 +88,7 @@ func New(usersClient grpcUsersV1.UsersServiceClient, cfg Config) (notifications. } return ¬ifier{ - usersClient: usersClient, + users: users, agents: agents, fromName: cfg.FromName, domainAltName: cfg.DomainAltName, @@ -156,27 +169,11 @@ func (n *notifier) buildEmailContent(notifType notifications.NotificationType, i } } -func (n *notifier) fetchUsers(ctx context.Context, userIDs []string) (map[string]*grpcUsersV1.User, error) { - req := &grpcUsersV1.RetrieveUsersReq{ - Ids: userIDs, - Limit: uint64(len(userIDs)), - Offset: 0, - } - - res, err := n.usersClient.RetrieveUsers(ctx, req) - if err != nil { - return nil, err - } - - users := make(map[string]*grpcUsersV1.User) - for _, user := range res.Users { - users[user.Id] = user - } - - return users, nil +func (n *notifier) fetchUsers(ctx context.Context, userIDs []string) (map[string]User, error) { + return n.users.FetchUsers(ctx, userIDs) } -func (n *notifier) userDisplayName(user *grpcUsersV1.User) string { +func (n *notifier) userDisplayName(user User) string { if user.FirstName != "" && user.LastName != "" { return fmt.Sprintf("%s %s", user.FirstName, user.LastName) } @@ -189,7 +186,7 @@ func (n *notifier) userDisplayName(user *grpcUsersV1.User) string { if user.Email != "" { return user.Email } - return user.Id + return user.ID } func titleFirst(s string) string { diff --git a/notifications/emailer/emailer_test.go b/notifications/emailer/emailer_test.go index d0791abc7..84833489e 100644 --- a/notifications/emailer/emailer_test.go +++ b/notifications/emailer/emailer_test.go @@ -9,12 +9,10 @@ import ( "os" "testing" - grpcUsersV1 "github.com/absmach/magistrala/api/grpc/users/v1" "github.com/absmach/magistrala/notifications" "github.com/absmach/magistrala/notifications/emailer" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/mock" - "google.golang.org/grpc" ) const ( @@ -37,15 +35,14 @@ const ( type mockUsersClient struct { mock.Mock - grpcUsersV1.UsersServiceClient } -func (m *mockUsersClient) RetrieveUsers(ctx context.Context, req *grpcUsersV1.RetrieveUsersReq, opts ...grpc.CallOption) (*grpcUsersV1.RetrieveUsersRes, error) { - args := m.Called(ctx, req, opts) +func (m *mockUsersClient) FetchUsers(ctx context.Context, userIDs []string) (map[string]emailer.User, error) { + args := m.Called(ctx, userIDs) if args.Get(0) == nil { return nil, args.Error(1) } - return args.Get(0).(*grpcUsersV1.RetrieveUsersRes), args.Error(1) + return args.Get(0).(map[string]emailer.User), args.Error(1) } func TestNotify(t *testing.T) { @@ -90,24 +87,22 @@ func TestNotify(t *testing.T) { RoleName: roleName, }, setupMock: func() { - usersClient.On("RetrieveUsers", mock.Anything, mock.MatchedBy(func(req *grpcUsersV1.RetrieveUsersReq) bool { - return len(req.Ids) == 2 && - ((req.Ids[0] == inviterID && req.Ids[1] == inviteeID) || - (req.Ids[0] == inviteeID && req.Ids[1] == inviterID)) - }), mock.Anything).Return(&grpcUsersV1.RetrieveUsersRes{ - Users: []*grpcUsersV1.User{ - { - Id: inviterID, - Email: inviterEmail, - FirstName: inviterFirst, - LastName: inviterLast, - }, - { - Id: inviteeID, - Email: inviteeEmail, - FirstName: inviteeFirst, - LastName: inviteeLast, - }, + usersClient.On("FetchUsers", mock.Anything, mock.MatchedBy(func(userIDs []string) bool { + return len(userIDs) == 2 && + ((userIDs[0] == inviterID && userIDs[1] == inviteeID) || + (userIDs[0] == inviteeID && userIDs[1] == inviterID)) + })).Return(map[string]emailer.User{ + inviterID: { + ID: inviterID, + Email: inviterEmail, + FirstName: inviterFirst, + LastName: inviterLast, + }, + inviteeID: { + ID: inviteeID, + Email: inviteeEmail, + FirstName: inviteeFirst, + LastName: inviteeLast, }, }, nil).Once() }, @@ -125,22 +120,20 @@ func TestNotify(t *testing.T) { RoleName: roleName, }, setupMock: func() { - usersClient.On("RetrieveUsers", mock.Anything, mock.MatchedBy(func(req *grpcUsersV1.RetrieveUsersReq) bool { - return len(req.Ids) == 2 - }), mock.Anything).Return(&grpcUsersV1.RetrieveUsersRes{ - Users: []*grpcUsersV1.User{ - { - Id: inviterID, - Email: inviterEmail, - FirstName: inviterFirst, - LastName: inviterLast, - }, - { - Id: inviteeID, - Email: inviteeEmail, - FirstName: inviteeFirst, - LastName: inviteeLast, - }, + usersClient.On("FetchUsers", mock.Anything, mock.MatchedBy(func(userIDs []string) bool { + return len(userIDs) == 2 + })).Return(map[string]emailer.User{ + inviterID: { + ID: inviterID, + Email: inviterEmail, + FirstName: inviterFirst, + LastName: inviterLast, + }, + inviteeID: { + ID: inviteeID, + Email: inviteeEmail, + FirstName: inviteeFirst, + LastName: inviteeLast, }, }, nil).Once() }, @@ -158,22 +151,20 @@ func TestNotify(t *testing.T) { RoleName: roleName, }, setupMock: func() { - usersClient.On("RetrieveUsers", mock.Anything, mock.MatchedBy(func(req *grpcUsersV1.RetrieveUsersReq) bool { - return len(req.Ids) == 2 - }), mock.Anything).Return(&grpcUsersV1.RetrieveUsersRes{ - Users: []*grpcUsersV1.User{ - { - Id: inviterID, - Email: inviterEmail, - FirstName: inviterFirst, - LastName: inviterLast, - }, - { - Id: inviteeID, - Email: inviteeEmail, - FirstName: inviteeFirst, - LastName: inviteeLast, - }, + usersClient.On("FetchUsers", mock.Anything, mock.MatchedBy(func(userIDs []string) bool { + return len(userIDs) == 2 + })).Return(map[string]emailer.User{ + inviterID: { + ID: inviterID, + Email: inviterEmail, + FirstName: inviterFirst, + LastName: inviterLast, + }, + inviteeID: { + ID: inviteeID, + Email: inviteeEmail, + FirstName: inviteeFirst, + LastName: inviteeLast, }, }, nil).Once() }, @@ -191,9 +182,9 @@ func TestNotify(t *testing.T) { RoleName: roleName, }, setupMock: func() { - usersClient.On("RetrieveUsers", mock.Anything, mock.Anything, mock.Anything).Return(nil, fmt.Errorf("grpc error")).Once() + usersClient.On("FetchUsers", mock.Anything, mock.Anything).Return(nil, fmt.Errorf("atom error")).Once() }, - expectedError: fmt.Errorf("grpc error"), + expectedError: fmt.Errorf("atom error"), }, } diff --git a/pkg/README.md b/pkg/README.md index e143d173a..962992866 100644 --- a/pkg/README.md +++ b/pkg/README.md @@ -27,7 +27,7 @@ import "github.com/absmach/magistrala/pkg/authn" | `events` | Event store client abstractions and subscriber utilities. | | `prometheus` | Metrics collectors for request counts/latency. | | `jaeger`, `tracing` | OpenTelemetry tracing configuration and instrumentation helpers. | -| `channels`, `clients`, `groups`, `domains`, `roles` | Shared types and helpers for core Magistrala domain services. | +| `channels`, `clients`, `groups`, `domains`, `roles` | Shared types and helpers for Magistrala domain services. | | `messaging`, `connections`, `callout` | Messaging DTOs, connection types, and outbound callout helpers. | | `sdk` | Go SDK for interacting with Magistrala services. | | `errors` | Error wrappers with consistent error typing. | diff --git a/pkg/authn/atom/authn.go b/pkg/authn/atom/authn.go new file mode 100644 index 000000000..14822a95c --- /dev/null +++ b/pkg/authn/atom/authn.go @@ -0,0 +1,37 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package atom + +import ( + "context" + + atomcore "github.com/absmach/magistrala/internal/atom" + "github.com/absmach/magistrala/pkg/authn" + "github.com/absmach/magistrala/pkg/errors" + svcerr "github.com/absmach/magistrala/pkg/errors/service" +) + +type authentication struct { + verifier *atomcore.TokenVerifier +} + +var _ authn.Authentication = (*authentication)(nil) + +func NewAuthentication() authn.Authentication { + return authentication{verifier: atomcore.NewTokenVerifier(atomcore.LoadConfig())} +} + +func (a authentication) Authenticate(ctx context.Context, token string) (authn.Session, error) { + claims, err := a.verifier.VerifyTokenClaims(ctx, token) + if err != nil { + return authn.Session{}, errors.Wrap(svcerr.ErrAuthentication, err) + } + return authn.Session{ + Type: authn.AccessToken, + UserID: claims.SubjectID, + DomainID: claims.TenantID, + Role: authn.UserRole, + Verified: true, + }, nil +} diff --git a/pkg/authz/authsvc/authz.go b/pkg/authz/authsvc/authz.go index d8c6a1dee..fad5c0953 100644 --- a/pkg/authz/authsvc/authz.go +++ b/pkg/authz/authsvc/authz.go @@ -8,7 +8,6 @@ import ( grpcAuthV1 "github.com/absmach/magistrala/api/grpc/auth/v1" "github.com/absmach/magistrala/auth/api/grpc/auth" - "github.com/absmach/magistrala/domains" "github.com/absmach/magistrala/pkg/authz" pkgDomians "github.com/absmach/magistrala/pkg/domains" "github.com/absmach/magistrala/pkg/errors" @@ -102,7 +101,7 @@ func (a authorization) checkDomain(ctx context.Context, subjectType, subject, do } switch status { - case domains.FreezeStatus: + case pkgDomians.FreezeStatus: _, err := a.authSvcClient.Authorize(ctx, &grpcAuthV1.AuthZReq{ PolicyReq: &grpcAuthV1.PolicyReq{ Subject: subject, @@ -114,7 +113,7 @@ func (a authorization) checkDomain(ctx context.Context, subjectType, subject, do }) return err - case domains.DisabledStatus: + case pkgDomians.DisabledStatus: _, err := a.authSvcClient.Authorize(ctx, &grpcAuthV1.AuthZReq{ PolicyReq: &grpcAuthV1.PolicyReq{ Subject: subject, @@ -126,7 +125,7 @@ func (a authorization) checkDomain(ctx context.Context, subjectType, subject, do }) return err - case domains.EnabledStatus: + case pkgDomians.EnabledStatus: _, err := a.authSvcClient.Authorize(ctx, &grpcAuthV1.AuthZReq{ PolicyReq: &grpcAuthV1.PolicyReq{ Subject: subject, diff --git a/pkg/bootstrap/events/consumer/doc.go b/pkg/bootstrap/events/consumer/doc.go deleted file mode 100644 index ac76e55de..000000000 --- a/pkg/bootstrap/events/consumer/doc.go +++ /dev/null @@ -1,6 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package consumer contains events consumer for client and channel events -// consumed by the Bootstrap service. -package consumer diff --git a/pkg/bootstrap/events/consumer/streams.go b/pkg/bootstrap/events/consumer/streams.go deleted file mode 100644 index 96d673923..000000000 --- a/pkg/bootstrap/events/consumer/streams.go +++ /dev/null @@ -1,46 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package consumer - -import ( - "context" - "log/slog" - - "github.com/absmach/magistrala/bootstrap" - "github.com/absmach/magistrala/pkg/events" - "github.com/absmach/magistrala/pkg/events/store" -) - -const stream = "events.magistrala.*.*" - -type eventHandler struct { - svc bootstrap.Service -} - -// BootstrapEventsSubscribe subscribes bootstrap config-state handlers to the event store. -func BootstrapEventsSubscribe(ctx context.Context, svc bootstrap.Service, esURL, esConsumerName string, logger *slog.Logger) error { - subscriber, err := store.NewSubscriber(ctx, esURL, "bootstrap-es-sub", logger) - if err != nil { - return err - } - - subConfig := events.SubscriberConfig{ - Stream: stream, - Consumer: esConsumerName, - Handler: NewEventHandler(svc), - Ordered: true, - } - return subscriber.Subscribe(ctx, subConfig) -} - -// NewEventHandler returns bootstrap events handler. -func NewEventHandler(svc bootstrap.Service) events.EventHandler { - return &eventHandler{ - svc: svc, - } -} - -func (es *eventHandler) Handle(_ context.Context, _ events.Event) error { - return nil -} diff --git a/pkg/bootstrap/events/doc.go b/pkg/bootstrap/events/doc.go deleted file mode 100644 index 61cd808b0..000000000 --- a/pkg/bootstrap/events/doc.go +++ /dev/null @@ -1,6 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package events provides the events sourcing of bootstrap -// provide replication in other service and definitions needed to support it -package events diff --git a/pkg/channels/events/consumer/decode.go b/pkg/channels/events/consumer/decode.go deleted file mode 100644 index 0b9b6c5dc..000000000 --- a/pkg/channels/events/consumer/decode.go +++ /dev/null @@ -1,297 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package consumer - -import ( - "time" - - "github.com/absmach/magistrala/channels" - "github.com/absmach/magistrala/pkg/connections" - "github.com/absmach/magistrala/pkg/errors" - "github.com/absmach/magistrala/pkg/roles" - rconsumer "github.com/absmach/magistrala/pkg/roles/rolemanager/events/consumer" -) - -const layout = "2006-01-02T15:04:05.999999Z" - -var ( - errDecodeCreateChannelEvent = errors.New("failed to decode channel create event") - errDecodeUpdateChannelEvent = errors.New("failed to decode channel update event") - errDecodeChangeStatusChannelEvent = errors.New("failed to decode channel change status event") - errDecodeRemoveChannelEvent = errors.New("failed to decode channel remove event") - errDecodeSetParentGroupEvent = errors.New("failed to decode channel set parent event") - errDecodeRemoveParentGroupEvent = errors.New("failed to decode channel remove parent event") - errDecodeConnectEvent = errors.New("failed to decode channel connect event") - errDeocodeDisconnectEvent = errors.New("failed to decode channel disconnect event") - - errID = errors.New("missing or invalid 'id'") - errDomain = errors.New("missing or invalid 'domain'") - errStatus = errors.New("missing or invalid 'status'") - errTags = errors.New("invalid 'tags'") - errConvertStatus = errors.New("failed to convert status") - errChannelIDs = errors.New("missing or invalid 'channel_ids' in connection") - errClientIDs = errors.New("missing or invalid 'client_ids' in connection") - errConnType = errors.New("missing or invalid 'type' in connection") - errCreatedAt = errors.New("failed to parse 'created_at' time") - errUpdatedAt = errors.New("failed to parse 'updated_at' time") -) - -func ToChannel(data map[string]any) (channels.Channel, error) { - var c channels.Channel - id, ok := data["id"].(string) - if !ok { - return channels.Channel{}, errID - } - c.ID = id - - dom, ok := data["domain"].(string) - if !ok { - return channels.Channel{}, errDomain - } - c.Domain = dom - - st, ok := data["status"].(string) - if !ok { - return channels.Channel{}, errStatus - } - status, err := channels.ToStatus(st) - if err != nil { - return channels.Channel{}, errConvertStatus - } - c.Status = status - - cat, ok := data["created_at"].(string) - if !ok { - return channels.Channel{}, errCreatedAt - } - ct, err := time.Parse(layout, cat) - if err != nil { - return channels.Channel{}, errors.Wrap(errCreatedAt, err) - } - c.CreatedAt = ct - - // Following fields of channels are allowed to be empty. - name, ok := data["name"].(string) - if ok { - c.Name = name - } - - parent, ok := data["parent_group_id"].(string) - if ok { - c.ParentGroup = parent - } - - itags, ok := data["tags"].([]any) - if ok { - tags, err := rconsumer.ToStrings(itags) - if err != nil { - return channels.Channel{}, errors.Wrap(errTags, err) - } - c.Tags = tags - } - - meta, ok := data["metadata"].(map[string]any) - if ok { - c.Metadata = meta - } - - uby, ok := data["updated_by"].(string) - if ok { - c.UpdatedBy = uby - } - - uat, ok := data["updated_at"].(string) - if ok { - ut, err := time.Parse(layout, uat) - if err != nil { - return channels.Channel{}, errors.Wrap(errUpdatedAt, err) - } - c.UpdatedAt = ut - } - - return c, nil -} - -func ToConnections(data map[string]any) ([]channels.Connection, error) { - var connTypes []connections.ConnType - domain, ok := data["domain"].(string) - if !ok { - return nil, errDomain - } - - ityp, ok := data["types"].([]any) - if !ok { - return nil, errConnType - } - typs, err := rconsumer.ToStrings(ityp) - if err != nil { - return nil, errors.Wrap(errConnType, err) - } - for _, typ := range typs { - connType, err := connections.ParseConnType(typ) - if err != nil { - return nil, errors.Wrap(errConnType, err) - } - connTypes = append(connTypes, connType) - } - - ichanIDs, ok := data["channel_ids"].([]any) - if !ok { - return []channels.Connection{}, errChannelIDs - } - channelIDs, err := rconsumer.ToStrings(ichanIDs) - if err != nil { - return []channels.Connection{}, errors.Wrap(errChannelIDs, err) - } - - iclIDs, ok := data["client_ids"].([]any) - if !ok { - return []channels.Connection{}, errClientIDs - } - clientIDs, err := rconsumer.ToStrings(iclIDs) - if err != nil { - return []channels.Connection{}, errors.Wrap(errClientIDs, err) - } - - var conns []channels.Connection - for _, chanID := range channelIDs { - for _, clientID := range clientIDs { - for _, connType := range connTypes { - conns = append(conns, channels.Connection{ - ChannelID: chanID, - ClientID: clientID, - Type: connType, - DomainID: domain, - }) - } - } - } - - return conns, nil -} - -func decodeCreateChannelEvent(data map[string]any) (channels.Channel, []roles.RoleProvision, error) { - c, err := ToChannel(data) - if err != nil { - return channels.Channel{}, []roles.RoleProvision{}, errors.Wrap(errDecodeCreateChannelEvent, err) - } - irps, ok := data["roles_provisioned"].([]any) - if !ok { - return channels.Channel{}, []roles.RoleProvision{}, errors.Wrap(errDecodeCreateChannelEvent, errors.New("missing or invalid 'roles_provisioned'")) - } - rps, err := rconsumer.ToRoleProvisions(irps) - if err != nil { - return channels.Channel{}, []roles.RoleProvision{}, errors.Wrap(errDecodeCreateChannelEvent, err) - } - - return c, rps, nil -} - -func decodeUpdateChannelEvent(data map[string]any) (channels.Channel, error) { - c, err := ToChannel(data) - if err != nil { - return channels.Channel{}, errors.Wrap(errDecodeUpdateChannelEvent, err) - } - return c, nil -} - -func decodeChangeStatusChannelEvent(data map[string]any) (channels.Channel, error) { - c, err := ToChannelStatus(data) - if err != nil { - return channels.Channel{}, errors.Wrap(errDecodeChangeStatusChannelEvent, err) - } - return c, nil -} - -func ToChannelStatus(data map[string]any) (channels.Channel, error) { - var c channels.Channel - id, ok := data["id"].(string) - if !ok { - return channels.Channel{}, errID - } - c.ID = id - - stat, ok := data["status"].(string) - if !ok { - return channels.Channel{}, errStatus - } - st, err := channels.ToStatus(stat) - if err != nil { - return channels.Channel{}, errors.Wrap(errConvertStatus, err) - } - c.Status = st - - uat, ok := data["updated_at"].(string) - if ok { - ut, err := time.Parse(layout, uat) - if err != nil { - return channels.Channel{}, errors.Wrap(errUpdatedAt, err) - } - c.UpdatedAt = ut - } - - uby, ok := data["updated_by"].(string) - if ok { - c.UpdatedBy = uby - } - - return c, nil -} - -func decodeRemoveChannelEvent(data map[string]any) (channels.Channel, error) { - var c channels.Channel - id, ok := data["id"].(string) - if !ok { - return channels.Channel{}, errors.Wrap(errDecodeRemoveChannelEvent, errID) - } - c.ID = id - - return c, nil -} - -func decodeConnectEvent(data map[string]any) ([]channels.Connection, error) { - conns, err := ToConnections(data) - if err != nil { - return []channels.Connection{}, errors.Wrap(errDecodeConnectEvent, err) - } - - return conns, nil -} - -func decodeDisconnectEvent(data map[string]any) ([]channels.Connection, error) { - conns, err := ToConnections(data) - if err != nil { - return []channels.Connection{}, errors.Wrap(errDeocodeDisconnectEvent, err) - } - - return conns, nil -} - -func decodeSetParentGroupEvent(data map[string]any) (channels.Channel, error) { - id, ok := data["id"].(string) - if !ok { - return channels.Channel{}, errors.Wrap(errDecodeSetParentGroupEvent, errID) - } - - parent, ok := data["parent_group_id"].(string) - if !ok { - return channels.Channel{}, errors.Wrap(errDecodeSetParentGroupEvent, errID) - } - - return channels.Channel{ - ID: id, - ParentGroup: parent, - }, nil -} - -func decodeRemoveParentGroupEvent(data map[string]any) (channels.Channel, error) { - id, ok := data["id"].(string) - if !ok { - return channels.Channel{}, errors.Wrap(errDecodeRemoveParentGroupEvent, errID) - } - - return channels.Channel{ - ID: id, - }, nil -} diff --git a/pkg/channels/events/consumer/doc.go b/pkg/channels/events/consumer/doc.go deleted file mode 100644 index 21aaf047a..000000000 --- a/pkg/channels/events/consumer/doc.go +++ /dev/null @@ -1,6 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package consumer contains events consumer for events -// published by channels service. -package consumer diff --git a/pkg/channels/events/consumer/streams.go b/pkg/channels/events/consumer/streams.go deleted file mode 100644 index 72cb554b0..000000000 --- a/pkg/channels/events/consumer/streams.go +++ /dev/null @@ -1,217 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package consumer - -import ( - "context" - "log/slog" - - "github.com/absmach/magistrala/channels" - "github.com/absmach/magistrala/pkg/errors" - "github.com/absmach/magistrala/pkg/events" - "github.com/absmach/magistrala/pkg/events/store" - rconsumer "github.com/absmach/magistrala/pkg/roles/rolemanager/events/consumer" -) - -const ( - stream = "events.magistrala.channel.*" - - create = "channel.create" - update = "channel.update" - updateTags = "channel.update_tags" - enable = "channel.enable" - disable = "channel.disable" - remove = "channel.remove" - connect = "channel.connect" - disconnect = "channel.disconnect" - setParentGroup = "channel.set_parent" - removeParentGroup = "channel.remove_parent" -) - -var ( - errNoOperationKey = errors.New("operation key is not found in event message") - errCreateChannelEvent = errors.New("failed to consume channel create event") - errUpdateChannelEvent = errors.New("failed to consume channel update event") - errChangeStatusChannelEvent = errors.New("failed to consume channel change status event") - errRemoveChannelEvent = errors.New("failed to consume channel remove event") - errConnectEvent = errors.New("failed to consume channel connect event") - errDisconnectEvent = errors.New("failed to consume channel disconnect event") - errSetParentGroupEvent = errors.New("failed to consume channel add parent group event") - errRemoveParentGroupEvent = errors.New("failed to consume channel remove parent group event") -) - -type eventHandler struct { - repo channels.Repository - rolesEventHandler rconsumer.EventHandler -} - -func ChannelsEventsSubscribe(ctx context.Context, repo channels.Repository, esURL, esConsumerName string, logger *slog.Logger) error { - subscriber, err := store.NewSubscriber(ctx, esURL, "channels-es-sub", logger) - if err != nil { - return err - } - - subConfig := events.SubscriberConfig{ - Stream: stream, - Consumer: esConsumerName, - Handler: NewEventHandler(repo), - Ordered: true, - } - return subscriber.Subscribe(ctx, subConfig) -} - -// NewEventHandler returns new event store handler. -func NewEventHandler(repo channels.Repository) events.EventHandler { - reh := rconsumer.NewEventHandler("channel", repo) - return &eventHandler{ - repo: repo, - rolesEventHandler: reh, - } -} - -func (es *eventHandler) Handle(ctx context.Context, event events.Event) error { - msg, err := event.Encode() - if err != nil { - return err - } - - op, ok := msg["operation"] - - if !ok { - return errNoOperationKey - } - switch op { - case create: - return es.createChannelHandler(ctx, msg) - case update: - return es.updateChannelHandler(ctx, msg) - case updateTags: - return es.updateChannelTagsHandler(ctx, msg) - case enable, disable: - return es.changeStatusChannelHandler(ctx, msg) - case remove: - return es.removeChannelHandler(ctx, msg) - case connect: - return es.connectChannelHandler(ctx, msg) - case disconnect: - return es.disconnectChannelHandler(ctx, msg) - case setParentGroup: - return es.setParentGroupHandler(ctx, msg) - case removeParentGroup: - return es.removeParentGroupHandler(ctx, msg) - } - - return es.rolesEventHandler.Handle(ctx, op, msg) -} - -func (es *eventHandler) createChannelHandler(ctx context.Context, data map[string]any) error { - c, rps, err := decodeCreateChannelEvent(data) - if err != nil { - return errors.Wrap(errCreateChannelEvent, err) - } - - if _, err := es.repo.Save(ctx, c); err != nil { - return errors.Wrap(errCreateChannelEvent, err) - } - if _, err := es.repo.AddRoles(ctx, rps); err != nil { - return errors.Wrap(errCreateChannelEvent, err) - } - - return nil -} - -func (es *eventHandler) updateChannelHandler(ctx context.Context, data map[string]any) error { - c, err := decodeUpdateChannelEvent(data) - if err != nil { - return errors.Wrap(errUpdateChannelEvent, err) - } - - if _, err := es.repo.Update(ctx, c); err != nil { - return errors.Wrap(errUpdateChannelEvent, err) - } - - return nil -} - -func (es *eventHandler) updateChannelTagsHandler(ctx context.Context, data map[string]any) error { - c, err := decodeUpdateChannelEvent(data) - if err != nil { - return errors.Wrap(errUpdateChannelEvent, err) - } - - if _, err := es.repo.UpdateTags(ctx, c); err != nil { - return errors.Wrap(errUpdateChannelEvent, err) - } - - return nil -} - -func (es *eventHandler) changeStatusChannelHandler(ctx context.Context, data map[string]any) error { - c, err := decodeChangeStatusChannelEvent(data) - if err != nil { - return errors.Wrap(errChangeStatusChannelEvent, err) - } - - if _, err := es.repo.ChangeStatus(ctx, c); err != nil { - return errors.Wrap(errChangeStatusChannelEvent, err) - } - - return nil -} - -func (es *eventHandler) removeChannelHandler(ctx context.Context, data map[string]any) error { - c, err := decodeRemoveChannelEvent(data) - if err != nil { - return errors.Wrap(errRemoveChannelEvent, err) - } - - if err := es.repo.Remove(ctx, c.ID); err != nil { - return errors.Wrap(errRemoveChannelEvent, err) - } - return nil -} - -func (es *eventHandler) connectChannelHandler(ctx context.Context, data map[string]any) error { - c, err := decodeConnectEvent(data) - if err != nil { - return errors.Wrap(errConnectEvent, err) - } - if err := es.repo.AddConnections(ctx, c); err != nil { - return errors.Wrap(errConnectEvent, err) - } - return nil -} - -func (es *eventHandler) disconnectChannelHandler(ctx context.Context, data map[string]any) error { - c, err := decodeDisconnectEvent(data) - if err != nil { - return errors.Wrap(errDisconnectEvent, err) - } - if err := es.repo.RemoveConnections(ctx, c); err != nil { - return errors.Wrap(errDisconnectEvent, err) - } - return nil -} - -func (es *eventHandler) setParentGroupHandler(ctx context.Context, data map[string]any) error { - c, err := decodeSetParentGroupEvent(data) - if err != nil { - return errors.Wrap(errSetParentGroupEvent, err) - } - if err := es.repo.SetParentGroup(ctx, c); err != nil { - return errors.Wrap(errSetParentGroupEvent, err) - } - return nil -} - -func (es *eventHandler) removeParentGroupHandler(ctx context.Context, data map[string]any) error { - c, err := decodeRemoveParentGroupEvent(data) - if err != nil { - return errors.Wrap(errRemoveParentGroupEvent, err) - } - if err := es.repo.RemoveParentGroup(ctx, c); err != nil { - return errors.Wrap(errRemoveParentGroupEvent, err) - } - return nil -} diff --git a/pkg/clients/events/consumer/decode.go b/pkg/clients/events/consumer/decode.go deleted file mode 100644 index 090d504fa..000000000 --- a/pkg/clients/events/consumer/decode.go +++ /dev/null @@ -1,225 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package consumer - -import ( - "time" - - "github.com/absmach/magistrala/clients" - "github.com/absmach/magistrala/pkg/errors" - "github.com/absmach/magistrala/pkg/roles" - rconsumer "github.com/absmach/magistrala/pkg/roles/rolemanager/events/consumer" -) - -const layout = "2006-01-02T15:04:05.999999Z" - -var ( - errDecodeCreateClientEvent = errors.New("failed to decode client create event") - errDecodeUpdateClientEvent = errors.New("failed to decode client update event") - errDecodeChangeStatusClientEvent = errors.New("failed to decode client change status event") - errDecodeRemoveClientEvent = errors.New("failed to decode client remove event") - errDecodeSetParentGroupEvent = errors.New("failed to decode client set parent event") - errDecodeRemoveParentGroupEvent = errors.New("failed to decode client remove parent event") - - errID = errors.New("missing or invalid 'id'") - errDomain = errors.New("missing or invalid 'domain'") - errStatus = errors.New("missing or invalid 'status'") - errTags = errors.New("invalid 'tags'") - errConvertStatus = errors.New("failed to convert status") - errCreatedAt = errors.New("failed to parse 'created_at' time") - errUpdatedAt = errors.New("failed to parse 'updated_at' time") -) - -func ToClient(data map[string]any) (clients.Client, error) { - var c clients.Client - id, ok := data["id"].(string) - if !ok { - return clients.Client{}, errID - } - c.ID = id - - dom, ok := data["domain"].(string) - if !ok { - return clients.Client{}, errDomain - } - c.Domain = dom - - st, ok := data["status"].(string) - if !ok { - return clients.Client{}, errStatus - } - status, err := clients.ToStatus(st) - if err != nil { - return clients.Client{}, errConvertStatus - } - c.Status = status - - cat, ok := data["created_at"].(string) - if !ok { - return clients.Client{}, errCreatedAt - } - ct, err := time.Parse(layout, cat) - if err != nil { - return clients.Client{}, errors.Wrap(errCreatedAt, err) - } - c.CreatedAt = ct - - // Following fields of clients are allowed to be empty. - name, ok := data["name"].(string) - if ok { - c.Name = name - } - - identity, ok := data["identity"].(string) - if ok { - c.Identity = identity - } - - parent, ok := data["parent_group_id"].(string) - if ok { - c.ParentGroup = parent - } - - itags, ok := data["tags"].([]any) - if ok { - tags, err := rconsumer.ToStrings(itags) - if err != nil { - return clients.Client{}, errors.Wrap(errTags, err) - } - c.Tags = tags - } - - meta, ok := data["metadata"].(map[string]any) - if ok { - c.Metadata = meta - } - - pmeta, ok := data["private_metadata"].(map[string]any) - if ok { - c.PrivateMetadata = pmeta - } - - uby, ok := data["updated_by"].(string) - if ok { - c.UpdatedBy = uby - } - - uat, ok := data["updated_at"].(string) - if ok { - ut, err := time.Parse(layout, uat) - if err != nil { - return clients.Client{}, errors.Wrap(errUpdatedAt, err) - } - c.UpdatedAt = ut - } - - return c, nil -} - -func decodeCreateClientEvent(data map[string]any) (clients.Client, []roles.RoleProvision, error) { - c, err := ToClient(data) - if err != nil { - return clients.Client{}, []roles.RoleProvision{}, errors.Wrap(errDecodeCreateClientEvent, err) - } - irps, ok := data["roles_provisioned"].([]any) - if !ok { - return clients.Client{}, []roles.RoleProvision{}, errors.Wrap(errDecodeCreateClientEvent, errors.New("missing or invalid 'roles_provisioned'")) - } - rps, err := rconsumer.ToRoleProvisions(irps) - if err != nil { - return clients.Client{}, []roles.RoleProvision{}, errors.Wrap(errDecodeCreateClientEvent, err) - } - - return c, rps, nil -} - -func decodeUpdateClientEvent(data map[string]any) (clients.Client, error) { - c, err := ToClient(data) - if err != nil { - return clients.Client{}, errors.Wrap(errDecodeUpdateClientEvent, err) - } - return c, nil -} - -func decodeChangeStatusClientEvent(data map[string]any) (clients.Client, error) { - c, err := ToClientStatus(data) - if err != nil { - return clients.Client{}, errors.Wrap(errDecodeChangeStatusClientEvent, err) - } - return c, nil -} - -func ToClientStatus(data map[string]any) (clients.Client, error) { - var c clients.Client - id, ok := data["id"].(string) - if !ok { - return clients.Client{}, errID - } - c.ID = id - - stat, ok := data["status"].(string) - if !ok { - return clients.Client{}, errStatus - } - st, err := clients.ToStatus(stat) - if err != nil { - return clients.Client{}, errors.Wrap(errConvertStatus, err) - } - c.Status = st - - uat, ok := data["updated_at"].(string) - if ok { - ut, err := time.Parse(layout, uat) - if err != nil { - return clients.Client{}, errors.Wrap(errUpdatedAt, err) - } - c.UpdatedAt = ut - } - - uby, ok := data["updated_by"].(string) - if ok { - c.UpdatedBy = uby - } - - return c, nil -} - -func decodeRemoveClientEvent(data map[string]any) (clients.Client, error) { - var c clients.Client - id, ok := data["id"].(string) - if !ok { - return clients.Client{}, errors.Wrap(errDecodeRemoveClientEvent, errID) - } - c.ID = id - - return c, nil -} - -func decodeSetParentGroupEvent(data map[string]any) (clients.Client, error) { - id, ok := data["id"].(string) - if !ok { - return clients.Client{}, errors.Wrap(errDecodeSetParentGroupEvent, errID) - } - - parent, ok := data["parent_group_id"].(string) - if !ok { - return clients.Client{}, errors.Wrap(errDecodeSetParentGroupEvent, errID) - } - - return clients.Client{ - ID: id, - ParentGroup: parent, - }, nil -} - -func decodeRemoveParentGroupEvent(data map[string]any) (clients.Client, error) { - id, ok := data["id"].(string) - if !ok { - return clients.Client{}, errors.Wrap(errDecodeRemoveParentGroupEvent, errID) - } - - return clients.Client{ - ID: id, - }, nil -} diff --git a/pkg/clients/events/consumer/doc.go b/pkg/clients/events/consumer/doc.go deleted file mode 100644 index 4621b5011..000000000 --- a/pkg/clients/events/consumer/doc.go +++ /dev/null @@ -1,6 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package consumer contains events consumer for events -// published by clients service. -package consumer diff --git a/pkg/clients/events/consumer/streams.go b/pkg/clients/events/consumer/streams.go deleted file mode 100644 index 8fc486129..000000000 --- a/pkg/clients/events/consumer/streams.go +++ /dev/null @@ -1,187 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package consumer - -import ( - "context" - "log/slog" - - "github.com/absmach/magistrala/clients" - "github.com/absmach/magistrala/pkg/errors" - "github.com/absmach/magistrala/pkg/events" - "github.com/absmach/magistrala/pkg/events/store" - rconsumer "github.com/absmach/magistrala/pkg/roles/rolemanager/events/consumer" -) - -const ( - stream = "events.magistrala.client.*" - - create = "client.create" - update = "client.update" - updateTags = "client.update_tags" - enable = "client.enable" - disable = "client.disable" - remove = "client.remove" - setParentGroup = "client.set_parent" - removeParentGroup = "client.remove_parent" -) - -var ( - errNoOperationKey = errors.New("operation key is not found in event message") - errCreateClientEvent = errors.New("failed to consume client create event") - errUpdateClientEvent = errors.New("failed to consume client update event") - errChangeStatusClientEvent = errors.New("failed to consume client change status event") - errRemoveClientEvent = errors.New("failed to consume client remove event") - errSetParentGroupEvent = errors.New("failed to consume client add parent group event") - errRemoveParentGroupEvent = errors.New("failed to consume client remove parent group event") -) - -type eventHandler struct { - repo clients.Repository - rolesEventHandler rconsumer.EventHandler -} - -func ClientsEventsSubscribe(ctx context.Context, repo clients.Repository, esURL, esConsumerName string, logger *slog.Logger) error { - subscriber, err := store.NewSubscriber(ctx, esURL, "clients-es-sub", logger) - if err != nil { - return err - } - - subConfig := events.SubscriberConfig{ - Stream: stream, - Consumer: esConsumerName, - Handler: NewEventHandler(repo), - Ordered: true, - } - return subscriber.Subscribe(ctx, subConfig) -} - -// NewEventHandler returns new event store handler. -func NewEventHandler(repo clients.Repository) events.EventHandler { - reh := rconsumer.NewEventHandler("client", repo) - return &eventHandler{ - repo: repo, - rolesEventHandler: reh, - } -} - -func (es *eventHandler) Handle(ctx context.Context, event events.Event) error { - msg, err := event.Encode() - if err != nil { - return err - } - - op, ok := msg["operation"] - - if !ok { - return errNoOperationKey - } - switch op { - case create: - return es.createClientHandler(ctx, msg) - case update: - return es.updateClientHandler(ctx, msg) - case updateTags: - return es.updateClientTagsHandler(ctx, msg) - case enable, disable: - return es.changeStatusClientHandler(ctx, msg) - case remove: - return es.removeClientHandler(ctx, msg) - case setParentGroup: - return es.setParentGroupHandler(ctx, msg) - case removeParentGroup: - return es.removeParentGroupHandler(ctx, msg) - } - - return es.rolesEventHandler.Handle(ctx, op, msg) -} - -func (es *eventHandler) createClientHandler(ctx context.Context, data map[string]any) error { - c, rps, err := decodeCreateClientEvent(data) - if err != nil { - return errors.Wrap(errCreateClientEvent, err) - } - - if _, err := es.repo.Save(ctx, c); err != nil { - return errors.Wrap(errCreateClientEvent, err) - } - if _, err := es.repo.AddRoles(ctx, rps); err != nil { - return errors.Wrap(errCreateClientEvent, err) - } - - return nil -} - -func (es *eventHandler) updateClientHandler(ctx context.Context, data map[string]any) error { - c, err := decodeUpdateClientEvent(data) - if err != nil { - return errors.Wrap(errUpdateClientEvent, err) - } - - if _, err := es.repo.Update(ctx, c); err != nil { - return errors.Wrap(errUpdateClientEvent, err) - } - - return nil -} - -func (es *eventHandler) updateClientTagsHandler(ctx context.Context, data map[string]any) error { - c, err := decodeUpdateClientEvent(data) - if err != nil { - return errors.Wrap(errUpdateClientEvent, err) - } - - if _, err := es.repo.UpdateTags(ctx, c); err != nil { - return errors.Wrap(errUpdateClientEvent, err) - } - - return nil -} - -func (es *eventHandler) changeStatusClientHandler(ctx context.Context, data map[string]any) error { - c, err := decodeChangeStatusClientEvent(data) - if err != nil { - return errors.Wrap(errChangeStatusClientEvent, err) - } - - if _, err := es.repo.ChangeStatus(ctx, c); err != nil { - return errors.Wrap(errChangeStatusClientEvent, err) - } - - return nil -} - -func (es *eventHandler) removeClientHandler(ctx context.Context, data map[string]any) error { - c, err := decodeRemoveClientEvent(data) - if err != nil { - return errors.Wrap(errRemoveClientEvent, err) - } - - if err := es.repo.Delete(ctx, c.ID); err != nil { - return errors.Wrap(errRemoveClientEvent, err) - } - return nil -} - -func (es *eventHandler) setParentGroupHandler(ctx context.Context, data map[string]any) error { - c, err := decodeSetParentGroupEvent(data) - if err != nil { - return errors.Wrap(errSetParentGroupEvent, err) - } - if err := es.repo.SetParentGroup(ctx, c); err != nil { - return errors.Wrap(errSetParentGroupEvent, err) - } - return nil -} - -func (es *eventHandler) removeParentGroupHandler(ctx context.Context, data map[string]any) error { - c, err := decodeRemoveParentGroupEvent(data) - if err != nil { - return errors.Wrap(errRemoveParentGroupEvent, err) - } - if err := es.repo.RemoveParentGroup(ctx, c); err != nil { - return errors.Wrap(errRemoveParentGroupEvent, err) - } - return nil -} diff --git a/pkg/domains/authz.go b/pkg/domains/authz.go index 9e5fa4380..f590100e6 100644 --- a/pkg/domains/authz.go +++ b/pkg/domains/authz.go @@ -5,10 +5,17 @@ package domains import ( "context" +) - "github.com/absmach/magistrala/domains" +type Status uint8 + +const ( + EnabledStatus Status = iota + DisabledStatus + FreezeStatus + AllStatus ) type Authorization interface { - RetrieveStatus(ctx context.Context, id string) (domains.Status, error) + RetrieveStatus(ctx context.Context, id string) (Status, error) } diff --git a/pkg/domains/events/consumer/decode.go b/pkg/domains/events/consumer/decode.go deleted file mode 100644 index 3979a1758..000000000 --- a/pkg/domains/events/consumer/decode.go +++ /dev/null @@ -1,271 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package consumer - -import ( - "fmt" - "time" - - "github.com/absmach/magistrala/domains" - "github.com/absmach/magistrala/pkg/errors" - "github.com/absmach/magistrala/pkg/roles" - rconsumer "github.com/absmach/magistrala/pkg/roles/rolemanager/events/consumer" -) - -const ( - layout = "2006-01-02T15:04:05.999999Z" -) - -var ( - errDecodeCreateDomainEvent = errors.New("failed to decode domain create event") - errDecodeUpdateDomainEvent = errors.New("failed to decode domain update event") - errDecodeEnableDomainEvent = errors.New("failed to decode domain enable event") - errDecodeDisableDomainEvent = errors.New("failed to decode domain disable event") - errDecodeFreezeDomainEvent = errors.New("failed to decode domain freeze event") - errDecodeRemoveDomainsEvent = errors.New("failed to decode domain remove event") - - errID = errors.New("missing or invalid 'id'") - errName = errors.New("missing or invalid 'name'") - errRoute = errors.New("missing or invalid 'route'") - errTags = errors.New("invalid 'tags'") - errStatus = errors.New("missing or invalid 'status'") - errConvertStatus = errors.New("failed to convert status") - errCreatedBy = errors.New("missing or invalid 'created_by'") - errCreatedAt = errors.New("failed to parse 'created_at' time") - errUpdatedAt = errors.New("failed to parse 'updated_at' time") -) - -func ToDomains(data map[string]any) (domains.Domain, error) { - var d domains.Domain - id, ok := data["id"].(string) - if !ok { - return domains.Domain{}, errID - } - d.ID = id - - name, ok := data["name"].(string) - if !ok { - return domains.Domain{}, errName - } - d.Name = name - - stat, ok := data["status"].(string) - if !ok { - return domains.Domain{}, errStatus - } - st, err := domains.ToStatus(stat) - if err != nil { - return domains.Domain{}, errors.Wrap(errConvertStatus, err) - } - d.Status = st - - route, ok := data["route"].(string) - if !ok { - return domains.Domain{}, errRoute - } - d.Route = route - - cby, ok := data["created_by"].(string) - if !ok { - return domains.Domain{}, errCreatedBy - } - d.CreatedBy = cby - - cat, ok := data["created_at"].(string) - if !ok { - return domains.Domain{}, errCreatedAt - } - ct, err := time.Parse(layout, cat) - if err != nil { - return domains.Domain{}, errors.Wrap(errCreatedAt, err) - } - d.CreatedAt = ct - - // Following fields of groups are allowed to be empty. - itags, ok := data["tags"].([]any) - if ok { - tags, err := rconsumer.ToStrings(itags) - if err != nil { - return domains.Domain{}, errors.Wrap(errTags, err) - } - d.Tags = tags - } - - meta, ok := data["metadata"].(map[string]any) - if ok { - d.Metadata = meta - } - - uby, ok := data["updated_by"].(string) - if ok { - d.UpdatedBy = uby - } - - uat, ok := data["updated_at"].(string) - if ok { - ut, err := time.Parse(layout, uat) - if err != nil { - return domains.Domain{}, errors.Wrap(errUpdatedAt, err) - } - d.UpdatedAt = ut - } - - return d, nil -} - -func decodeCreateDomainEvent(data map[string]any) (domains.Domain, []roles.RoleProvision, error) { - d, err := ToDomains(data) - if err != nil { - return domains.Domain{}, []roles.RoleProvision{}, errors.Wrap(errDecodeCreateDomainEvent, err) - } - irps, ok := data["roles_provisioned"].([]any) - if !ok { - return domains.Domain{}, []roles.RoleProvision{}, errors.Wrap(errDecodeCreateDomainEvent, errors.New("missing or invalid 'roles_provisioned'")) - } - rps, err := rconsumer.ToRoleProvisions(irps) - if err != nil { - return domains.Domain{}, []roles.RoleProvision{}, errors.Wrap(errDecodeCreateDomainEvent, err) - } - - return d, rps, nil -} - -func decodeUpdateDomainEvent(data map[string]any) (domains.Domain, error) { - var d domains.Domain - - id, ok := data["id"].(string) - if !ok { - return domains.Domain{}, errors.Wrap(errDecodeUpdateDomainEvent, errID) - } - d.ID = id - - name, ok := data["name"].(string) - if ok { - d.Name = name - } - - route, ok := data["route"].(string) - if ok { - d.Route = route - } - - itags, ok := data["tags"].([]any) - if ok { - tags, err := rconsumer.ToStrings(itags) - if err != nil { - return domains.Domain{}, errors.Wrap(errDecodeUpdateDomainEvent, err) - } - d.Tags = tags - } - - meta, ok := data["metadata"].(map[string]any) - if ok { - d.Metadata = meta - } - - uby, ok := data["updated_by"].(string) - if ok { - d.UpdatedBy = uby - } - - uat, ok := data["updated_at"].(string) - if ok { - ut, err := time.Parse(layout, uat) - if err != nil { - return domains.Domain{}, errors.Wrap(errDecodeUpdateDomainEvent, errors.Wrap(errUpdatedAt, err)) - } - d.UpdatedAt = ut - } - - return d, nil -} - -func decodeEnableDomainEvent(data map[string]any) (domains.Domain, error) { - var d domains.Domain - id, ok := data["id"].(string) - if !ok { - return domains.Domain{}, errors.Wrap(errDecodeEnableDomainEvent, errID) - } - d.ID = id - - uby, ok := data["updated_by"].(string) - if ok { - d.UpdatedBy = uby - } - - uat, ok := data["updated_at"].(string) - if ok { - ut, err := time.Parse(layout, uat) - if err != nil { - return domains.Domain{}, errors.Wrap(errDecodeEnableDomainEvent, errors.Wrap(errUpdatedAt, err)) - } - d.UpdatedAt = ut - } - - return d, nil -} - -func decodeDisableDomainEvent(data map[string]any) (domains.Domain, error) { - var d domains.Domain - id, ok := data["id"].(string) - if !ok { - return domains.Domain{}, errors.Wrap(errDecodeDisableDomainEvent, errID) - } - d.ID = id - - uby, ok := data["updated_by"].(string) - if ok { - d.UpdatedBy = uby - } - - uat, ok := data["updated_at"].(string) - if ok { - ut, err := time.Parse(layout, uat) - if err != nil { - return domains.Domain{}, errors.Wrap(errDecodeDisableDomainEvent, errors.Wrap(errUpdatedAt, err)) - } - d.UpdatedAt = ut - } - - return d, nil -} - -func decodeFreezeDomainEvent(data map[string]any) (domains.Domain, error) { - var d domains.Domain - id, ok := data["id"].(string) - if !ok { - return domains.Domain{}, errors.Wrap(errDecodeFreezeDomainEvent, errID) - } - d.ID = id - - uby, ok := data["updated_by"].(string) - if ok { - d.UpdatedBy = uby - } - - uat, ok := data["updated_at"].(string) - if ok { - ut, err := time.Parse(layout, uat) - if err != nil { - return domains.Domain{}, errors.Wrap(errDecodeFreezeDomainEvent, errors.Wrap(errUpdatedAt, err)) - } - d.UpdatedAt = ut - } - - return d, nil -} - -func decodeUserDeleteDomainEvent(_ map[string]any) (domains.Domain, error) { - return domains.Domain{}, fmt.Errorf("not implemented decode domain user delete event ") -} - -func decodeDeleteDomainEvent(data map[string]any) (domains.Domain, error) { - var d domains.Domain - id, ok := data["id"].(string) - if !ok { - return domains.Domain{}, errors.Wrap(errDecodeRemoveDomainsEvent, errID) - } - d.ID = id - return d, nil -} diff --git a/pkg/domains/events/consumer/doc.go b/pkg/domains/events/consumer/doc.go deleted file mode 100644 index a99efc936..000000000 --- a/pkg/domains/events/consumer/doc.go +++ /dev/null @@ -1,4 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package consumer diff --git a/pkg/domains/events/consumer/stream.go b/pkg/domains/events/consumer/stream.go deleted file mode 100644 index 152089c2a..000000000 --- a/pkg/domains/events/consumer/stream.go +++ /dev/null @@ -1,204 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package consumer - -import ( - "context" - "fmt" - "log/slog" - - "github.com/absmach/magistrala/domains" - "github.com/absmach/magistrala/pkg/errors" - "github.com/absmach/magistrala/pkg/events" - "github.com/absmach/magistrala/pkg/events/store" - "github.com/absmach/magistrala/pkg/messaging" - rconsumer "github.com/absmach/magistrala/pkg/roles/rolemanager/events/consumer" -) - -const ( - stream = "events.magistrala.domain.*" - - create = "domain.create" - update = "domain.update" - enable = "domain.enable" - disable = "domain.disable" - freeze = "domain.freeze" - delete = "domain.delete" - userDelete = "domain.user_delete" -) - -var ( - errNoOperationKey = errors.New("operation key is not found in event message") - errCreateDomainEvent = errors.New("failed to consume domain create event") - errUpdateDomainEvent = errors.New("failed to consume domain update event") - errEnableDomainGroupEvent = errors.New("failed to consume domain enable event") - errDisableDomainGroupEvent = errors.New("failed to consume domain disable event") - errFreezeDomainGroupEvent = errors.New("failed to consume domain freeze event") - errUserDeleteDomainEvent = errors.New("failed to consume domain user delete event") - errDeleteDomainEvent = errors.New("failed to consume domain delete event") -) - -type eventHandler struct { - repo domains.Repository - rolesEventHandler rconsumer.EventHandler -} - -func DomainsEventsSubscribe(ctx context.Context, repo domains.Repository, esURL, esConsumerName string, logger *slog.Logger) error { - subscriber, err := store.NewSubscriber(ctx, esURL, "domains-es-sub", logger) - if err != nil { - return err - } - - subConfig := events.SubscriberConfig{ - Stream: stream, - Consumer: esConsumerName, - Handler: NewEventHandler(repo), - DeliveryPolicy: messaging.DeliverNewPolicy, - Ordered: true, - } - return subscriber.Subscribe(ctx, subConfig) -} - -// NewEventHandler returns new event store handler. -func NewEventHandler(repo domains.Repository) events.EventHandler { - reh := rconsumer.NewEventHandler("domain", repo) - return &eventHandler{ - repo: repo, - rolesEventHandler: reh, - } -} - -func (es *eventHandler) Handle(ctx context.Context, event events.Event) error { - msg, err := event.Encode() - if err != nil { - return err - } - - op, ok := msg["operation"] - - if !ok { - return errNoOperationKey - } - switch op { - case create: - return es.createDomainHandler(ctx, msg) - case update: - return es.updateDomainHandler(ctx, msg) - case enable: - return es.enableDomainHandler(ctx, msg) - case disable: - return es.disableDomainHandler(ctx, msg) - case freeze: - return es.freezeDomainHandler(ctx, msg) - case userDelete: - return es.userDeleteDomainHandler(ctx, msg) - case delete: - return es.deleteDomainHandler(ctx, msg) - } - - return es.rolesEventHandler.Handle(ctx, op, msg) -} - -func (es *eventHandler) createDomainHandler(ctx context.Context, data map[string]any) error { - d, rps, err := decodeCreateDomainEvent(data) - if err != nil { - return errors.Wrap(errCreateDomainEvent, err) - } - - if _, err := es.repo.SaveDomain(ctx, d); err != nil { - return errors.Wrap(errCreateDomainEvent, err) - } - if _, err := es.repo.AddRoles(ctx, rps); err != nil { - return errors.Wrap(errCreateDomainEvent, err) - } - - return nil -} - -func (es *eventHandler) updateDomainHandler(ctx context.Context, data map[string]any) error { - d, err := decodeUpdateDomainEvent(data) - if err != nil { - return errors.Wrap(errUpdateDomainEvent, err) - } - - if _, err := es.repo.UpdateDomain( - ctx, - d.ID, - domains.DomainReq{ - Name: &d.Name, - Metadata: &d.Metadata, - Tags: &d.Tags, - UpdatedBy: &d.UpdatedBy, - UpdatedAt: &d.UpdatedAt, - }, - ); err != nil { - return errors.Wrap(errUpdateDomainEvent, err) - } - - return nil -} - -func (es *eventHandler) enableDomainHandler(ctx context.Context, data map[string]any) error { - d, err := decodeEnableDomainEvent(data) - if err != nil { - return errors.Wrap(errEnableDomainGroupEvent, err) - } - - enabled := domains.EnabledStatus - if _, err := es.repo.UpdateDomain(ctx, d.ID, domains.DomainReq{Status: &enabled, UpdatedBy: &d.UpdatedBy, UpdatedAt: &d.UpdatedAt}); err != nil { - return errors.Wrap(errEnableDomainGroupEvent, err) - } - - return nil -} - -func (es *eventHandler) disableDomainHandler(ctx context.Context, data map[string]any) error { - d, err := decodeDisableDomainEvent(data) - if err != nil { - return errors.Wrap(errDisableDomainGroupEvent, err) - } - - disabled := domains.DisabledStatus - if _, err := es.repo.UpdateDomain(ctx, d.ID, domains.DomainReq{Status: &disabled, UpdatedBy: &d.UpdatedBy, UpdatedAt: &d.UpdatedAt}); err != nil { - return errors.Wrap(errDisableDomainGroupEvent, err) - } - - return nil -} - -func (es *eventHandler) freezeDomainHandler(ctx context.Context, data map[string]any) error { - d, err := decodeFreezeDomainEvent(data) - if err != nil { - return errors.Wrap(errFreezeDomainGroupEvent, err) - } - - freeze := domains.FreezeStatus - if _, err := es.repo.UpdateDomain(ctx, d.ID, domains.DomainReq{Status: &freeze, UpdatedBy: &d.UpdatedBy, UpdatedAt: &d.UpdatedAt}); err != nil { - return errors.Wrap(errFreezeDomainGroupEvent, err) - } - - return nil -} - -func (es *eventHandler) userDeleteDomainHandler(_ context.Context, data map[string]any) error { - _, err := decodeUserDeleteDomainEvent(data) - if err != nil { - return errors.Wrap(errUserDeleteDomainEvent, err) - } - - return fmt.Errorf("not implemented user delete domain handler") -} - -func (es *eventHandler) deleteDomainHandler(ctx context.Context, data map[string]any) error { - d, err := decodeDeleteDomainEvent(data) - if err != nil { - return errors.Wrap(errDeleteDomainEvent, err) - } - - if err := es.repo.DeleteDomain(ctx, d.ID); err != nil { - return errors.Wrap(errDeleteDomainEvent, err) - } - - return nil -} diff --git a/pkg/domains/grpcclient/authz.go b/pkg/domains/grpcclient/authz.go index 87553ce6d..699b805d4 100644 --- a/pkg/domains/grpcclient/authz.go +++ b/pkg/domains/grpcclient/authz.go @@ -8,7 +8,6 @@ import ( grpcCommonV1 "github.com/absmach/magistrala/api/grpc/common/v1" grpcDomainsV1 "github.com/absmach/magistrala/api/grpc/domains/v1" - "github.com/absmach/magistrala/domains" pkgDomains "github.com/absmach/magistrala/pkg/domains" "github.com/absmach/magistrala/pkg/grpcclient" ) @@ -28,14 +27,14 @@ func NewAuthorization(ctx context.Context, cfg grpcclient.Config) (pkgDomains.Au return authorization{domainsSvcClient: domainsClient}, domainsClient, domainsHandler, nil } -func (a authorization) RetrieveStatus(ctx context.Context, id string) (domains.Status, error) { +func (a authorization) RetrieveStatus(ctx context.Context, id string) (pkgDomains.Status, error) { req := grpcCommonV1.RetrieveEntityReq{ Id: id, } res, err := a.domainsSvcClient.RetrieveStatus(ctx, &req) if err != nil { - return domains.AllStatus, err + return pkgDomains.AllStatus, err } - return domains.Status(res.Entity.GetStatus()), nil + return pkgDomains.Status(res.Entity.GetStatus()), nil } diff --git a/pkg/domains/psvc/authz.go b/pkg/domains/psvc/authz.go deleted file mode 100644 index 41d3b4689..000000000 --- a/pkg/domains/psvc/authz.go +++ /dev/null @@ -1,35 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package domainscache - -import ( - "context" - - "github.com/absmach/magistrala/domains" - "github.com/absmach/magistrala/domains/private" - pkgDomains "github.com/absmach/magistrala/pkg/domains" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" -) - -type authorization struct { - psvc private.Service -} - -var _ pkgDomains.Authorization = (*authorization)(nil) - -func NewAuthorization(psvc private.Service) pkgDomains.Authorization { - return authorization{ - psvc: psvc, - } -} - -func (a authorization) RetrieveStatus(ctx context.Context, id string) (domains.Status, error) { - status, err := a.psvc.RetrieveStatus(ctx, id) - if err != nil { - return domains.AllStatus, errors.Wrap(svcerr.ErrViewEntity, err) - } - - return status, nil -} diff --git a/pkg/groups/events/consumer/decode.go b/pkg/groups/events/consumer/decode.go deleted file mode 100644 index 3eedd0cb8..000000000 --- a/pkg/groups/events/consumer/decode.go +++ /dev/null @@ -1,254 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package consumer - -import ( - "time" - - "github.com/absmach/magistrala/groups" - "github.com/absmach/magistrala/internal/nullable" - "github.com/absmach/magistrala/pkg/errors" - "github.com/absmach/magistrala/pkg/roles" - rconsumer "github.com/absmach/magistrala/pkg/roles/rolemanager/events/consumer" -) - -var ( - errDecodeCreateGroupEvent = errors.New("failed to decode group create event") - errDecodeUpdateGroupEvent = errors.New("failed to decode group update event") - errDecodeChangeStatusGroupEvent = errors.New("failed to decode group change status event") - errDecodeRemoveGroupEvent = errors.New("failed to decode group remove event") - errDecodeAddParentGroupEvent = errors.New("failed to decode group add parent event") - errDecodeRemoveParentGroupEvent = errors.New("failed to decode group remove parent event") - errDecodeAddChildrenGroupsEvent = errors.New("failed to decode group add children groups event") - errDecodeRemoveChildrenGroupsEvent = errors.New("failed to decode group remove children groups event") - - errID = errors.New("missing or invalid 'id'") - errName = errors.New("missing or invalid 'name'") - errDomain = errors.New("missing or invalid 'domain'") - errParent = errors.New("missing or invalid 'parent'") - errChildrenIDs = errors.New("missing or invalid 'children_ids'") - errStatus = errors.New("missing or invalid 'status'") - errConvertStatus = errors.New("failed to convert status") - errCreatedAt = errors.New("failed to parse 'created_at' time") - errUpdatedAt = errors.New("failed to parse 'updated_at' time") -) - -const layout = "2006-01-02T15:04:05.999999Z" - -func ToGroups(data map[string]any) (groups.Group, error) { - var g groups.Group - id, ok := data["id"].(string) - if !ok { - return groups.Group{}, errID - } - g.ID = id - - name, ok := data["name"].(string) - if !ok { - return groups.Group{}, errName - } - g.Name = name - - dom, ok := data["domain"].(string) - if !ok { - return groups.Group{}, errDomain - } - g.Domain = dom - - stat, ok := data["status"].(string) - if !ok { - return groups.Group{}, errStatus - } - st, err := groups.ToStatus(stat) - if err != nil { - return groups.Group{}, errors.Wrap(errConvertStatus, err) - } - g.Status = st - - cat, ok := data["created_at"].(string) - if !ok { - return groups.Group{}, errCreatedAt - } - ct, err := time.Parse(layout, cat) - if err != nil { - return groups.Group{}, errors.Wrap(errCreatedAt, err) - } - g.CreatedAt = ct - - // Following fields of groups are allowed to be empty. - - desc, ok := data["description"].(string) - if ok { - g.Description = nullable.New(desc) - } - - parent, ok := data["parent"].(string) - if ok { - g.Parent = parent - } - - meta, ok := data["metadata"].(map[string]any) - if ok { - g.Metadata = meta - } - - uby, ok := data["updated_by"].(string) - if ok { - g.UpdatedBy = uby - } - - uat, ok := data["updated_at"].(string) - if ok { - ut, err := time.Parse(layout, uat) - if err != nil { - return groups.Group{}, errors.Wrap(errUpdatedAt, err) - } - g.UpdatedAt = ut - } - - return g, nil -} - -func decodeCreateGroupEvent(data map[string]any) (groups.Group, []roles.RoleProvision, error) { - g, err := ToGroups(data) - if err != nil { - return groups.Group{}, []roles.RoleProvision{}, errors.Wrap(errDecodeCreateGroupEvent, err) - } - irps, ok := data["roles_provisioned"].([]any) - if !ok { - return groups.Group{}, []roles.RoleProvision{}, errors.Wrap(errDecodeCreateGroupEvent, errors.New("missing or invalid 'roles_provisioned'")) - } - rps, err := rconsumer.ToRoleProvisions(irps) - if err != nil { - return groups.Group{}, []roles.RoleProvision{}, errors.Wrap(errDecodeCreateGroupEvent, err) - } - - return g, rps, nil -} - -func decodeUpdateGroupEvent(data map[string]any) (groups.Group, error) { - g, err := ToGroups(data) - if err != nil { - return groups.Group{}, errors.Wrap(errDecodeUpdateGroupEvent, err) - } - return g, nil -} - -func ToGroupStatus(data map[string]any) (groups.Group, error) { - var g groups.Group - id, ok := data["id"].(string) - if !ok { - return groups.Group{}, errID - } - g.ID = id - - stat, ok := data["status"].(string) - if !ok { - return groups.Group{}, errStatus - } - st, err := groups.ToStatus(stat) - if err != nil { - return groups.Group{}, errors.Wrap(errConvertStatus, err) - } - g.Status = st - - uat, ok := data["updated_at"].(string) - if ok { - ut, err := time.Parse(layout, uat) - if err != nil { - return groups.Group{}, errors.Wrap(errUpdatedAt, err) - } - g.UpdatedAt = ut - } - - uby, ok := data["updated_by"].(string) - if ok { - g.UpdatedBy = uby - } - - return g, nil -} - -func decodeChangeStatusGroupEvent(data map[string]any) (groups.Group, error) { - g, err := ToGroupStatus(data) - if err != nil { - return groups.Group{}, errors.Wrap(errDecodeChangeStatusGroupEvent, err) - } - return g, nil -} - -func decodeRemoveGroupEvent(data map[string]any) (groups.Group, error) { - var g groups.Group - id, ok := data["id"].(string) - if !ok { - return groups.Group{}, errors.Wrap(errDecodeRemoveGroupEvent, errID) - } - g.ID = id - - return g, nil -} - -func decodeAddParentGroupEvent(data map[string]any) (id string, parent string, err error) { - id, ok := data["id"].(string) - if !ok { - return "", "", errors.Wrap(errAddParentGroupEvent, errID) - } - - parent, ok = data["parent_id"].(string) - if !ok { - return "", "", errors.Wrap(errDecodeAddParentGroupEvent, errParent) - } - - return id, parent, nil -} - -func decodeRemoveParentGroupEvent(data map[string]any) (id string, err error) { - id, ok := data["id"].(string) - if !ok { - return "", errors.Wrap(errDecodeRemoveParentGroupEvent, errID) - } - - return id, nil -} - -func decodeAddChildrenGroupEvent(data map[string]any) (id string, childrenIDs []string, err error) { - id, ok := data["id"].(string) - if !ok { - return "", []string{}, errors.Wrap(errDecodeAddChildrenGroupsEvent, errID) - } - chIDs, ok := data["children_ids"].([]any) - if !ok { - return "", []string{}, errors.Wrap(errDecodeAddChildrenGroupsEvent, errChildrenIDs) - } - cids, err := rconsumer.ToStrings(chIDs) - if err != nil { - return "", []string{}, errors.Wrap(errDecodeAddChildrenGroupsEvent, errors.Wrap(errChildrenIDs, err)) - } - return id, cids, nil -} - -func decodeRemoveChildrenGroupEvent(data map[string]any) (id string, childrenIDs []string, err error) { - id, ok := data["id"].(string) - if !ok { - return "", []string{}, errors.Wrap(errDecodeRemoveChildrenGroupsEvent, errID) - } - chIDs, ok := data["children_ids"].([]any) - if !ok { - return "", []string{}, errors.Wrap(errDecodeRemoveChildrenGroupsEvent, errChildrenIDs) - } - cids, err := rconsumer.ToStrings(chIDs) - if err != nil { - return "", []string{}, errors.Wrap(errDecodeRemoveChildrenGroupsEvent, errors.Wrap(errChildrenIDs, err)) - } - return id, cids, nil -} - -func decodeRemoveAllChildrenGroupEvent(data map[string]any) (id string, err error) { - id, ok := data["id"].(string) - if !ok { - return "", errors.Wrap(errDecodeRemoveChildrenGroupsEvent, errID) - } - - return id, nil -} diff --git a/pkg/groups/events/consumer/doc.go b/pkg/groups/events/consumer/doc.go deleted file mode 100644 index f3fea76f1..000000000 --- a/pkg/groups/events/consumer/doc.go +++ /dev/null @@ -1,6 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package consumer contains events consumer for events -// published by Bootstrap service. -package consumer diff --git a/pkg/groups/events/consumer/streams.go b/pkg/groups/events/consumer/streams.go deleted file mode 100644 index 14992ef75..000000000 --- a/pkg/groups/events/consumer/streams.go +++ /dev/null @@ -1,223 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package consumer - -import ( - "context" - "log/slog" - - "github.com/absmach/magistrala/groups" - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - "github.com/absmach/magistrala/pkg/events" - "github.com/absmach/magistrala/pkg/events/store" - rconsumer "github.com/absmach/magistrala/pkg/roles/rolemanager/events/consumer" -) - -const ( - stream = "events.magistrala.group.*" - - create = "group.create" - update = "group.update" - enable = "group.enable" - disable = "group.disable" - remove = "group.remove" - addParentGroup = "group.add_parent_group" - removeParentGroup = "group.remove_parent_group" - addChildrenGroups = "group.add_children_groups" - removeChildrenGroups = "group.remove_children_groups" - removeAllChildrenGroups = "group.remove_all_children_groups" -) - -var ( - errNoOperationKey = errors.New("operation key is not found in event message") - errCreateGroupEvent = errors.New("failed to consume group create event") - errUpdateGroupEvent = errors.New("failed to consume group update event") - errChangeStatusGroupEvent = errors.New("failed to consume group change status event") - errRemoveGroupEvent = errors.New("failed to consume group remove event") - errAddParentGroupEvent = errors.New("failed to consume group add parent group event") - errRemoveParentGroupEvent = errors.New("failed to consume group remove parent group event") - errAddChildrenGroupEvent = errors.New("failed to consume group add children groups event") - errRemoveChildrenGroupEvent = errors.New("failed to consume group remove children groups event") - errRemoveAllChildrenGroupEvent = errors.New("failed to consume group remove all children groups event") -) - -type eventHandler struct { - repo groups.Repository - rolesEventHandler rconsumer.EventHandler -} - -func GroupsEventsSubscribe(ctx context.Context, repo groups.Repository, esURL, esConsumerName string, logger *slog.Logger) error { - subscriber, err := store.NewSubscriber(ctx, esURL, "groups-es-sub", logger) - if err != nil { - return err - } - - subConfig := events.SubscriberConfig{ - Stream: stream, - Consumer: esConsumerName, - Handler: NewEventHandler(repo), - Ordered: true, - } - return subscriber.Subscribe(ctx, subConfig) -} - -// NewEventHandler returns new event store handler. -func NewEventHandler(repo groups.Repository) events.EventHandler { - reh := rconsumer.NewEventHandler("group", repo) - return &eventHandler{ - repo: repo, - rolesEventHandler: reh, - } -} - -func (es *eventHandler) Handle(ctx context.Context, event events.Event) error { - msg, err := event.Encode() - if err != nil { - return err - } - - op, ok := msg["operation"] - - if !ok { - return errNoOperationKey - } - switch op { - case create: - return es.createGroupHandler(ctx, msg) - case update: - return es.updateGroupHandler(ctx, msg) - case enable, disable: - return es.changeStatusGroupHandler(ctx, msg) - case remove: - return es.removeGroupHandler(ctx, msg) - case addParentGroup: - return es.addParentGroupHandler(ctx, msg) - case removeParentGroup: - return es.removeParentGroupHandler(ctx, msg) - case addChildrenGroups: - return es.addChildrenGroupsHandler(ctx, msg) - case removeChildrenGroups: - return es.removeChildrenGroupsHandler(ctx, msg) - case removeAllChildrenGroups: - return es.removeAllChildrenGroupsHandler(ctx, msg) - } - - return es.rolesEventHandler.Handle(ctx, op, msg) -} - -func (es *eventHandler) createGroupHandler(ctx context.Context, data map[string]any) error { - g, rps, err := decodeCreateGroupEvent(data) - if err != nil { - return errors.Wrap(errCreateGroupEvent, err) - } - - if _, err := es.repo.Save(ctx, g); err != nil { - return errors.Wrap(errCreateGroupEvent, err) - } - if _, err := es.repo.AddRoles(ctx, rps); err != nil { - return errors.Wrap(errCreateGroupEvent, err) - } - - return nil -} - -func (es *eventHandler) updateGroupHandler(ctx context.Context, data map[string]any) error { - g, err := decodeUpdateGroupEvent(data) - if err != nil { - return errors.Wrap(errUpdateGroupEvent, err) - } - - if _, err := es.repo.Update(ctx, g); err != nil { - return errors.Wrap(errUpdateGroupEvent, err) - } - - return nil -} - -func (es *eventHandler) changeStatusGroupHandler(ctx context.Context, data map[string]any) error { - g, err := decodeChangeStatusGroupEvent(data) - if err != nil { - return errors.Wrap(errChangeStatusGroupEvent, err) - } - - if _, err := es.repo.ChangeStatus(ctx, g); err != nil { - return errors.Wrap(errChangeStatusGroupEvent, err) - } - - return nil -} - -func (es *eventHandler) removeGroupHandler(ctx context.Context, data map[string]any) error { - g, err := decodeRemoveGroupEvent(data) - if err != nil { - return errors.Wrap(errRemoveGroupEvent, err) - } - - if err := es.repo.Delete(ctx, g.ID); err != nil { - return errors.Wrap(errRemoveGroupEvent, err) - } - return nil -} - -func (es *eventHandler) addParentGroupHandler(ctx context.Context, data map[string]any) error { - id, parent, err := decodeAddParentGroupEvent(data) - if err != nil { - return errors.Wrap(errAddParentGroupEvent, err) - } - if err := es.repo.AssignParentGroup(ctx, parent, id); err != nil { - return errors.Wrap(errAddParentGroupEvent, err) - } - return nil -} - -func (es *eventHandler) removeParentGroupHandler(ctx context.Context, data map[string]any) error { - id, err := decodeRemoveParentGroupEvent(data) - if err != nil { - return errors.Wrap(errRemoveParentGroupEvent, err) - } - g, err := es.repo.RetrieveByID(ctx, id) - if err != nil { - return errors.Wrap(errRemoveParentGroupEvent, err) - } - if err := es.repo.UnassignParentGroup(ctx, g.Parent, id); err != nil { - return errors.Wrap(errRemoveParentGroupEvent, err) - } - return nil -} - -func (es *eventHandler) addChildrenGroupsHandler(ctx context.Context, data map[string]any) error { - id, cids, err := decodeAddChildrenGroupEvent(data) - if err != nil { - return errors.Wrap(errAddChildrenGroupEvent, err) - } - - if err := es.repo.AssignParentGroup(ctx, id, cids...); err != nil { - return errors.Wrap(errAddChildrenGroupEvent, err) - } - return nil -} - -func (es *eventHandler) removeChildrenGroupsHandler(ctx context.Context, data map[string]any) error { - id, cids, err := decodeRemoveChildrenGroupEvent(data) - if err != nil { - return errors.Wrap(errRemoveChildrenGroupEvent, err) - } - - if err := es.repo.UnassignParentGroup(ctx, id, cids...); err != nil { - return errors.Wrap(errRemoveChildrenGroupEvent, err) - } - return nil -} - -func (es *eventHandler) removeAllChildrenGroupsHandler(ctx context.Context, data map[string]any) error { - id, err := decodeRemoveAllChildrenGroupEvent(data) - if err != nil { - return errors.Wrap(errRemoveAllChildrenGroupEvent, err) - } - if err := es.repo.UnassignAllChildrenGroups(ctx, id); err != nil && err != repoerr.ErrNotFound { - return errors.Wrap(errRemoveAllChildrenGroupEvent, err) - } - return nil -} diff --git a/pkg/grpcclient/client.go b/pkg/grpcclient/client.go index c0d89b79e..0603bd4da 100644 --- a/pkg/grpcclient/client.go +++ b/pkg/grpcclient/client.go @@ -13,11 +13,6 @@ import ( grpcTokenV1 "github.com/absmach/magistrala/api/grpc/token/v1" grpcUsersV1 "github.com/absmach/magistrala/api/grpc/users/v1" tokengrpc "github.com/absmach/magistrala/auth/api/grpc/token" - channelsgrpc "github.com/absmach/magistrala/channels/api/grpc" - clientsauth "github.com/absmach/magistrala/clients/api/grpc" - domainsgrpc "github.com/absmach/magistrala/domains/api/grpc" - groupsgrpc "github.com/absmach/magistrala/groups/api/grpc" - usersgrpc "github.com/absmach/magistrala/users/api/grpc" grpchealth "google.golang.org/grpc/health/grpc_health_v1" ) @@ -55,7 +50,7 @@ func SetupDomainsClient(ctx context.Context, cfg Config) (grpcDomainsV1.DomainsS return nil, nil, err } - return domainsgrpc.NewDomainsClient(client.Connection(), cfg.Timeout), client, nil + return grpcDomainsV1.NewDomainsServiceClient(client.Connection()), client, nil } // SetupClientsClient loads clients gRPC configuration and creates new clients gRPC client. @@ -69,7 +64,7 @@ func SetupClientsClient(ctx context.Context, cfg Config) (grpcClientsV1.ClientsS return nil, nil, err } - return clientsauth.NewClient(client.Connection(), cfg.Timeout), client, nil + return grpcClientsV1.NewClientsServiceClient(client.Connection()), client, nil } // SetupChannelsClient loads channels gRPC configuration and creates new channels gRPC client. @@ -83,7 +78,7 @@ func SetupChannelsClient(ctx context.Context, cfg Config) (grpcChannelsV1.Channe return nil, nil, err } - return channelsgrpc.NewClient(client.Connection(), cfg.Timeout), client, nil + return grpcChannelsV1.NewChannelsServiceClient(client.Connection()), client, nil } // SetupGroupsClient loads groups gRPC configuration and creates new groups gRPC client. @@ -97,7 +92,7 @@ func SetupGroupsClient(ctx context.Context, cfg Config) (grpcGroupsV1.GroupsServ return nil, nil, err } - return groupsgrpc.NewClient(client.Connection(), cfg.Timeout), client, nil + return grpcGroupsV1.NewGroupsServiceClient(client.Connection()), client, nil } // SetupUsersClient loads users gRPC configuration and creates new users gRPC client. @@ -111,5 +106,5 @@ func SetupUsersClient(ctx context.Context, cfg Config) (grpcUsersV1.UsersService return nil, nil, err } - return usersgrpc.NewClient(client.Connection(), cfg.Timeout), client, nil + return grpcUsersV1.NewUsersServiceClient(client.Connection()), client, nil } diff --git a/pkg/grpcclient/client_test.go b/pkg/grpcclient/client_test.go deleted file mode 100644 index f27e4d216..000000000 --- a/pkg/grpcclient/client_test.go +++ /dev/null @@ -1,168 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package grpcclient_test - -import ( - "context" - "fmt" - "testing" - "time" - - grpcClientsV1 "github.com/absmach/magistrala/api/grpc/clients/v1" - grpcDomainsV1 "github.com/absmach/magistrala/api/grpc/domains/v1" - grpcTokenV1 "github.com/absmach/magistrala/api/grpc/token/v1" - tokengrpcapi "github.com/absmach/magistrala/auth/api/grpc/token" - "github.com/absmach/magistrala/auth/mocks" - clientsgrpcapi "github.com/absmach/magistrala/clients/api/grpc" - climocks "github.com/absmach/magistrala/clients/private/mocks" - domainsgrpcapi "github.com/absmach/magistrala/domains/api/grpc" - domainsMocks "github.com/absmach/magistrala/domains/private/mocks" - mglog "github.com/absmach/magistrala/logger" - "github.com/absmach/magistrala/pkg/errors" - "github.com/absmach/magistrala/pkg/grpcclient" - "github.com/absmach/magistrala/pkg/server" - grpcserver "github.com/absmach/magistrala/pkg/server/grpc" - "github.com/stretchr/testify/assert" - "google.golang.org/grpc" -) - -func TestSetupToken(t *testing.T) { - ctx, cancel := context.WithCancel(context.Background()) - defer cancel() - registerAuthServiceServer := func(srv *grpc.Server) { - grpcTokenV1.RegisterTokenServiceServer(srv, tokengrpcapi.NewTokenServer(new(mocks.Service))) - } - gs := grpcserver.NewServer(ctx, cancel, "auth", server.Config{Port: "12345"}, registerAuthServiceServer, mglog.NewMock()) - go func() { - err := gs.Start() - assert.Nil(t, err, fmt.Sprintf(`"Unexpected error creating server %s"`, err)) - }() - defer func() { - err := gs.Stop() - assert.Nil(t, err, fmt.Sprintf(`"Unexpected error stopping server %s"`, err)) - }() - - cases := []struct { - desc string - config grpcclient.Config - err error - }{ - { - desc: "successful", - config: grpcclient.Config{ - URL: "localhost:12345", - Timeout: time.Second, - }, - err: nil, - }, - { - desc: "failed with empty URL", - config: grpcclient.Config{ - URL: "", - Timeout: time.Second, - }, - err: errors.New("service is not serving"), - }, - } - - for _, c := range cases { - t.Run(c.desc, func(t *testing.T) { - client, handler, err := grpcclient.SetupTokenClient(context.Background(), c.config) - assert.True(t, errors.Contains(err, c.err), fmt.Sprintf("expected %s to contain %s", err, c.err)) - if err == nil { - assert.NotNil(t, client) - assert.NotNil(t, handler) - } - }) - } -} - -func TestSetupClientsClient(t *testing.T) { - ctx, cancel := context.WithCancel(context.Background()) - defer cancel() - - registerClientsServiceServer := func(srv *grpc.Server) { - grpcClientsV1.RegisterClientsServiceServer(srv, clientsgrpcapi.NewServer(new(climocks.Service))) - } - gs := grpcserver.NewServer(ctx, cancel, "clients", server.Config{Port: "12345"}, registerClientsServiceServer, mglog.NewMock()) - go func() { - err := gs.Start() - assert.Nil(t, err, fmt.Sprintf(`"Unexpected error creating server %s"`, err)) - }() - time.Sleep(time.Second) - defer func() { - err := gs.Stop() - assert.Nil(t, err, fmt.Sprintf(`"Unexpected error stopping server %s"`, err)) - }() - - cases := []struct { - desc string - config grpcclient.Config - err error - }{ - { - desc: "successful", - config: grpcclient.Config{ - URL: "localhost:12345", - Timeout: time.Second, - }, - err: nil, - }, - } - - for _, c := range cases { - t.Run(c.desc, func(t *testing.T) { - client, handler, err := grpcclient.SetupClientsClient(context.Background(), c.config) - assert.True(t, errors.Contains(err, c.err), fmt.Sprintf("expected %s to contain %s", err, c.err)) - if err == nil { - assert.NotNil(t, client) - assert.NotNil(t, handler) - } - }) - } -} - -func TestSetupDomainsClient(t *testing.T) { - ctx, cancel := context.WithCancel(context.Background()) - defer cancel() - registerDomainsServiceServer := func(srv *grpc.Server) { - grpcDomainsV1.RegisterDomainsServiceServer(srv, domainsgrpcapi.NewDomainsServer(new(domainsMocks.Service))) - } - gs := grpcserver.NewServer(ctx, cancel, "domains", server.Config{Port: "12345"}, registerDomainsServiceServer, mglog.NewMock()) - go func() { - err := gs.Start() - assert.Nil(t, err, fmt.Sprintf("Unexpected error creating server %s", err)) - }() - time.Sleep(time.Second) - defer func() { - err := gs.Stop() - assert.Nil(t, err, fmt.Sprintf("Unexpected error stopping server %s", err)) - }() - - cases := []struct { - desc string - config grpcclient.Config - err error - }{ - { - desc: "successfully", - config: grpcclient.Config{ - URL: "localhost:12345", - Timeout: time.Second, - }, - err: nil, - }, - } - - for _, c := range cases { - t.Run(c.desc, func(t *testing.T) { - client, handler, err := grpcclient.SetupDomainsClient(context.Background(), c.config) - assert.True(t, errors.Contains(err, c.err), fmt.Sprintf("expected %s to contain %s", err, c.err)) - if err == nil { - assert.NotNil(t, client) - assert.NotNil(t, handler) - } - }) - } -} diff --git a/pkg/messaging/fluxmq/options.go b/pkg/messaging/fluxmq/options.go index 842642235..992bdc57d 100644 --- a/pkg/messaging/fluxmq/options.go +++ b/pkg/messaging/fluxmq/options.go @@ -20,6 +20,7 @@ type options struct { prefix string connectionName string directTopicIngress bool + directTopicOnly bool } func defaultOptions() options { @@ -85,6 +86,25 @@ func DirectTopicIngress() messaging.Option { } } +// DirectTopicOnly subscribes only to regular MQTT topics and skips stream queue +// consumption. This is intended for bridge services that observe broker-native +// topics without also consuming queued messages. +func DirectTopicOnly() messaging.Option { + return func(val any) error { + switch v := val.(type) { + case *publisher: + return nil + case *pubsub: + v.directTopicIngress = true + v.directTopicOnly = true + default: + return ErrInvalidType + } + + return nil + } +} + // JSStreamConfig is a no-op for FluxMQ AMQP backend and exists only to keep // option-compatibility with legacy NATS broker wrappers. func JSStreamConfig(_ jetstream.StreamConfig) messaging.Option { diff --git a/pkg/messaging/fluxmq/publisher.go b/pkg/messaging/fluxmq/publisher.go index cfc268af1..7113ddc95 100644 --- a/pkg/messaging/fluxmq/publisher.go +++ b/pkg/messaging/fluxmq/publisher.go @@ -21,7 +21,17 @@ type publisher struct { } // NewPublisher creates a FluxMQ-backed message publisher. -func NewPublisher(_ context.Context, url string, opts ...messaging.Option) (messaging.Publisher, error) { +func NewPublisher(ctx context.Context, url string, opts ...messaging.Option) (messaging.Publisher, error) { + return newPublisher(ctx, url, true, opts...) +} + +// NewUndeclaredPublisher creates a FluxMQ-backed publisher without declaring +// the stream queue. Use it only when another service owns queue declaration. +func NewUndeclaredPublisher(ctx context.Context, url string, opts ...messaging.Option) (messaging.Publisher, error) { + return newPublisher(ctx, url, false, opts...) +} + +func newPublisher(_ context.Context, url string, declare bool, opts ...messaging.Option) (messaging.Publisher, error) { pub := &publisher{ options: defaultOptions(), } @@ -52,9 +62,11 @@ func NewPublisher(_ context.Context, url string, opts ...messaging.Option) (mess if err := client.Connect(); err != nil { return nil, err } - if err := declareStream(client, pub.prefix); err != nil { - _ = client.Close() - return nil, err + if declare { + if err := declareStream(client, pub.prefix); err != nil { + _ = client.Close() + return nil, err + } } pub.client = client diff --git a/pkg/messaging/fluxmq/pubsub.go b/pkg/messaging/fluxmq/pubsub.go index 80f4931d1..0c74147d4 100644 --- a/pkg/messaging/fluxmq/pubsub.go +++ b/pkg/messaging/fluxmq/pubsub.go @@ -8,7 +8,6 @@ import ( "errors" "fmt" "log/slog" - "strconv" "strings" "sync" "time" @@ -94,29 +93,31 @@ func (ps *pubsub) Subscribe(_ context.Context, cfg messaging.SubscriberConfig) e } group := formatConsumerName(cfg.Topic, cfg.ID) - opts := &fluxamqp.StreamConsumeOptions{ - QueueName: ps.prefix, - Filter: streamFilter(ps.prefix, cfg.Topic), - ConsumerGroup: group, - } + sub := subscription{} - switch cfg.DeliveryPolicy { - case messaging.DeliverNewPolicy: - opts.Offset = "last" - case messaging.DeliverAllPolicy: - opts.Offset = "first" - } - - if err := ps.client.SubscribeToStream(opts, func(msg *fluxamqp.QueueMessage) { - if err := ps.handle(cfg.Handler, msg); err != nil { - ps.logWarn("failed to process FluxMQ stream message", "error", err, "topic", cfg.Topic, "consumer_group", group) + if !ps.directTopicOnly { + opts := &fluxamqp.StreamConsumeOptions{ + QueueName: ps.prefix, + Filter: streamFilter(ps.prefix, cfg.Topic), + ConsumerGroup: group, } - }); err != nil { - return err - } - sub := subscription{ - streamTopic: queueFilter(ps.prefix, cfg.Topic), + switch cfg.DeliveryPolicy { + case messaging.DeliverNewPolicy: + opts.Offset = "last" + case messaging.DeliverAllPolicy: + opts.Offset = "first" + } + + if err := ps.client.SubscribeToStream(opts, func(msg *fluxamqp.QueueMessage) { + if err := ps.handle(cfg.Handler, msg); err != nil { + ps.logWarn("failed to process FluxMQ stream message", "error", err, "topic", cfg.Topic, "consumer_group", group) + } + }); err != nil { + return err + } + + sub.streamTopic = queueFilter(ps.prefix, cfg.Topic) } if ps.directTopicIngress { // Subscribe to regular MQTT topics so that messages published directly @@ -127,7 +128,9 @@ func (ps *pubsub) Subscribe(_ context.Context, cfg messaging.SubscriberConfig) e ps.logWarn("failed to process FluxMQ topic message", "error", err, "topic", sub.mqttTopic) } }); err != nil { - _ = ps.client.UnsubscribeFromStream(sub.streamTopic) + if sub.streamTopic != "" { + _ = ps.client.UnsubscribeFromStream(sub.streamTopic) + } return err } @@ -156,7 +159,10 @@ func (ps *pubsub) Unsubscribe(_ context.Context, id, topic string) error { return ErrNotSubscribed } - streamErr := ps.client.UnsubscribeFromStream(sub.streamTopic) + var streamErr error + if sub.streamTopic != "" { + streamErr = ps.client.UnsubscribeFromStream(sub.streamTopic) + } var topicErr error if sub.mqttTopic != "" { topicErr = ps.client.Unsubscribe(sub.mqttTopic) @@ -226,11 +232,12 @@ func messageFromDelivery(body []byte, headers map[string]any, ts time.Time, pref protocol = "mqtt" } - created := ts.UnixNano() - if s := stringHeader(headers, "created"); s != "" { - if v, err := strconv.ParseInt(s, 10, 64); err == nil { - created = v - } + created := time.Now().UnixNano() + if !ts.IsZero() { + created = ts.UnixNano() + } + if v, ok := int64Header(headers, "created"); ok { + created = v } return &messaging.Message{ diff --git a/pkg/messaging/fluxmq/pubsub_test.go b/pkg/messaging/fluxmq/pubsub_test.go index dfbef3f96..65e3d2b06 100644 --- a/pkg/messaging/fluxmq/pubsub_test.go +++ b/pkg/messaging/fluxmq/pubsub_test.go @@ -139,7 +139,7 @@ func TestMessageFromDelivery(t *testing.T) { { name: "use explicit publisher header when present", body: []byte("raw"), - headers: map[string]any{"external_id": "tenant-user", "client_id": "client-22"}, + headers: map[string]any{"external_id": "tenant-user", "client_id": "client-22", "created": int64(1710000000000000250)}, ts: time.Unix(1710000000, 250), prefix: "m", mqttTopic: "m/dom/c/ch", @@ -151,7 +151,7 @@ func TestMessageFromDelivery(t *testing.T) { Publisher: "tenant-user", ClientId: "client-22", Protocol: "mqtt", - Created: time.Unix(1710000000, 250).UnixNano(), + Created: 1710000000000000250, }, }, { @@ -214,3 +214,28 @@ func TestMessageFromDelivery(t *testing.T) { }) } } + +func TestMessageFromDeliveryZeroTimestampFallsBackToNow(t *testing.T) { + before := time.Now().UnixNano() + got, err := messageFromDelivery([]byte("raw"), nil, time.Time{}, "m", "m/dom/c/ch") + after := time.Now().UnixNano() + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if got.Created < before || got.Created > after { + t.Fatalf("expected created timestamp between %d and %d, got %d", before, after, got.Created) + } +} + +func TestDirectTopicOnlyEnablesDirectIngressAndSkipsStream(t *testing.T) { + ps := &pubsub{} + if err := DirectTopicOnly()(ps); err != nil { + t.Fatalf("unexpected error: %v", err) + } + if !ps.directTopicIngress { + t.Fatal("expected direct topic ingress to be enabled") + } + if !ps.directTopicOnly { + t.Fatal("expected stream consumption to be skipped") + } +} diff --git a/pkg/messaging/fluxmq/topic.go b/pkg/messaging/fluxmq/topic.go index ecb80f567..c1ee6af63 100644 --- a/pkg/messaging/fluxmq/topic.go +++ b/pkg/messaging/fluxmq/topic.go @@ -5,6 +5,7 @@ package fluxmq import ( "fmt" + "strconv" "strings" fluxamqp "github.com/absmach/fluxmq/client/amqp" @@ -113,6 +114,32 @@ func stringHeader(headers map[string]any, key string) string { } } +func int64Header(headers map[string]any, key string) (int64, bool) { + if headers == nil { + return 0, false + } + v, ok := headers[key] + if !ok { + return 0, false + } + switch val := v.(type) { + case int64: + return val, true + case int: + return int64(val), true + case int32: + return int64(val), true + case string: + parsed, err := strconv.ParseInt(val, 10, 64) + return parsed, err == nil + case []byte: + parsed, err := strconv.ParseInt(string(val), 10, 64) + return parsed, err == nil + default: + return 0, false + } +} + func declareStream(client *fluxamqp.Client, prefix string) error { _, err := client.DeclareStreamQueue(&fluxamqp.StreamQueueOptions{ Name: prefix, diff --git a/pkg/messaging/message_identity.go b/pkg/messaging/message_identity.go index b66e99b4d..78e493079 100644 --- a/pkg/messaging/message_identity.go +++ b/pkg/messaging/message_identity.go @@ -3,14 +3,16 @@ package messaging -// ClientIdentity returns the transport client identifier carried by the message. -// It falls back to Publisher for backward compatibility with older messages. +// ClientIdentity returns the authenticated application identity carried by the +// message. FluxMQ stores the protocol connection identifier in client_id and the +// Atom/Magistrala entity identifier in publisher/external_id, so publisher wins +// when both are present. func (m *Message) ClientIdentity() string { if m == nil { return "" } - if clientID := m.GetClientId(); clientID != "" { - return clientID + if publisher := m.GetPublisher(); publisher != "" { + return publisher } - return m.GetPublisher() + return m.GetClientId() } diff --git a/pkg/messaging/message_identity_test.go b/pkg/messaging/message_identity_test.go new file mode 100644 index 000000000..d136420a8 --- /dev/null +++ b/pkg/messaging/message_identity_test.go @@ -0,0 +1,43 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package messaging + +import "testing" + +func TestClientIdentity(t *testing.T) { + cases := []struct { + name string + msg *Message + want string + }{ + { + name: "nil message", + msg: nil, + want: "", + }, + { + name: "publisher wins over transport client id", + msg: &Message{ + Publisher: "entity-1", + ClientId: "amqp091:connection", + }, + want: "entity-1", + }, + { + name: "fallback to transport client id for legacy messages", + msg: &Message{ + ClientId: "legacy-client", + }, + want: "legacy-client", + }, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + if got := tc.msg.ClientIdentity(); got != tc.want { + t.Fatalf("expected %q, got %q", tc.want, got) + } + }) + } +} diff --git a/pkg/messaging/topics_test.go b/pkg/messaging/topics_test.go deleted file mode 100644 index 7ce73bc8c..000000000 --- a/pkg/messaging/topics_test.go +++ /dev/null @@ -1,1022 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package messaging_test - -import ( - "context" - "fmt" - "testing" - "time" - - grpcCommonV1 "github.com/absmach/magistrala/api/grpc/common/v1" - chmocks "github.com/absmach/magistrala/channels/mocks" - dmocks "github.com/absmach/magistrala/domains/mocks" - "github.com/absmach/magistrala/internal/testsutil" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/pkg/messaging" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" -) - -var ( - validRoute = "valid-route" - invalidRoute = "invalid-route" - channelID = testsutil.GenerateUUID(&testing.T{}) - domainID = testsutil.GenerateUUID(&testing.T{}) - topicFmt = "m/%s/c/%s" - healthTopicFmt = "hc/%s" - subtopic = "subtopic" - topicSubtopicFmt = "m/%s/c/%s/%s" - cachedTopic = fmt.Sprintf(topicSubtopicFmt, domainID, channelID, subtopic) -) - -func setupResolver() (messaging.TopicResolver, *dmocks.DomainsServiceClient, *chmocks.ChannelsServiceClient) { - channels := new(chmocks.ChannelsServiceClient) - domains := new(dmocks.DomainsServiceClient) - resolver := messaging.NewTopicResolver(channels, domains) - - return resolver, domains, channels -} - -func setupParser() (messaging.TopicParser, *dmocks.DomainsServiceClient, *chmocks.ChannelsServiceClient, error) { - channels := new(chmocks.ChannelsServiceClient) - domains := new(dmocks.DomainsServiceClient) - parser, err := messaging.NewTopicParser(messaging.DefaultCacheConfig, channels, domains) - if err != nil { - return nil, nil, nil, err - } - - return parser, domains, channels, nil -} - -var ParsePublisherTopicTestCases = []struct { - desc string - topic string - domainID string - channelID string - subtopic string - topicType messaging.TopicType - err error -}{ - { - desc: "valid topic with subtopic /m/domain123/c/channel456/devices/temp", - topic: "/m/domain123/c/channel456/devices/temp", - domainID: "domain123", - channelID: "channel456", - subtopic: "devices/temp", - topicType: messaging.MessageType, - err: nil, - }, - { - desc: "valid topic with URL encoded subtopic /m/domain123/c/channel456/devices%2Ftemp%2Fdata", - topic: "/m/domain123/c/channel456/devices%2Ftemp%2Fdata", - domainID: "domain123", - channelID: "channel456", - subtopic: "devices/temp/data", - topicType: messaging.MessageType, - }, - { - desc: "valid topic with subtopic /m/domain/c/channel/extra/extra2", - topic: "/m/domain/c/channel/extra/extra2", - domainID: "domain", - channelID: "channel", - subtopic: "extra/extra2", - topicType: messaging.MessageType, - }, - { - desc: "valid topic without subtopic /m/domain123/c/channel456", - topic: "/m/domain123/c/channel456", - domainID: "domain123", - channelID: "channel456", - subtopic: "", - topicType: messaging.MessageType, - }, - { - desc: "valid topic with trailing slash /m/domain123/c/channel456/devices/temp/", - topic: "/m/domain123/c/channel456/devices/temp/", - domainID: "domain123", - channelID: "channel456", - subtopic: "devices/temp", - topicType: messaging.MessageType, - }, - { - desc: "valid health check topic", - topic: fmt.Sprintf(healthTopicFmt, domainID), - domainID: domainID, - channelID: "", - subtopic: "", - topicType: messaging.HealthType, - err: nil, - }, - { - desc: "invalid health check topic with empty domain", - topic: "hc/", - domainID: "", - channelID: "", - subtopic: "", - topicType: messaging.InvalidType, - err: messaging.ErrMalformedTopic, - }, - { - desc: "invalid topic format (missing parts) /m/domain123/c/", - topic: "/m/domain123/c/", - domainID: "domain123", - channelID: "", - subtopic: "", - err: messaging.ErrMalformedTopic, - }, - { - desc: "invalid topic format (missing domain) /m//c/channel123", - topic: "/m//c/channel123", - domainID: "", - channelID: "", - subtopic: "", - err: messaging.ErrMalformedTopic, - }, - { - desc: "invalid topic format (missing channel) /m/domain123/c/", - topic: "/m/domain123/c//subtopic", - domainID: "domain123", - channelID: "", - subtopic: "", - err: messaging.ErrMalformedTopic, - }, - { - desc: "topic with wildcards + and # /m/domain123/c/channel456/devices/+/temp/#", - topic: "/m/domain123/c/channel456/devices/+/temp/#", - domainID: "domain123", - channelID: "channel456", - subtopic: "", - err: messaging.ErrMalformedTopic, - }, - { - desc: "invalid domain name m/domain*123/c/channel456/devices/+/temp/#", - topic: "m/domain*123/c/channel456/devices/+/temp/#", - domainID: "", - channelID: "channel456", - subtopic: "devices.*.temp.>", - err: messaging.ErrMalformedTopic, - }, - { - desc: "invalid subtopic /m/domain123/c/channel456/sub/a*b/topic", - topic: "/m/domain123/c/channel456/sub/a*b/topic", - domainID: "domain123", - channelID: "channel456", - subtopic: "", - err: messaging.ErrMalformedTopic, - }, - { - desc: "invalid subtopic /m/domain123/c/channel456/sub/a>b/topic", - topic: "/m/domain123/c/channel456/sub/a>b/topic", - domainID: "domain123", - channelID: "channel456", - subtopic: "", - err: messaging.ErrMalformedTopic, - }, - { - desc: "invalid subtopic /m/domain123/c/channel456/sub/a#b/topic", - topic: "/m/domain123/c/channel456/sub/a#b/topic", - domainID: "domain123", - channelID: "channel456", - subtopic: "", - err: messaging.ErrMalformedTopic, - }, - { - desc: "invalid subtopic /m/domain123/c/channel456/sub/a+b/topic", - topic: "/m/domain123/c/channel456/sub/a+b/topic", - domainID: "domain123", - channelID: "channel456", - subtopic: "", - err: messaging.ErrMalformedTopic, - }, - { - desc: "invalid subtopic /m/domain123/c/channel456/sub/a//b/topic", - topic: "/m/domain123/c/channel456/sub/a//b/topic", - domainID: "domain123", - channelID: "channel456", - subtopic: "", - err: messaging.ErrMalformedTopic, - }, - { - desc: "invalid topic regex \"not-a-topic\"", - topic: "not-a-topic", - domainID: "", - channelID: "", - subtopic: "", - err: messaging.ErrMalformedTopic, - }, - { - desc: "extra segment before prefix /extra/m/domain/c/channel", - topic: "/extra/m/domain/c/channel", - err: messaging.ErrMalformedTopic, - }, -} - -func TestParsePublishTopic(t *testing.T) { - for _, tc := range ParsePublisherTopicTestCases { - t.Run(tc.desc, func(t *testing.T) { - domainID, channelID, subtopic, topicType, err := messaging.ParsePublishTopic(tc.topic) - assert.True(t, errors.Contains(err, tc.err), "expected error %v, got %v", tc.err, err) - if err == nil { - assert.Equal(t, tc.domainID, domainID) - assert.Equal(t, tc.channelID, channelID) - assert.Equal(t, tc.subtopic, subtopic) - assert.Equal(t, tc.topicType, topicType) - } - }) - } -} - -func BenchmarkParsePublisherTopic(b *testing.B) { - for _, tc := range ParsePublisherTopicTestCases { - b.Run(tc.desc, func(b *testing.B) { - for b.Loop() { - _, _, _, _, _ = messaging.ParsePublishTopic(tc.topic) - } - }) - } -} - -var ParseSubscribeTestCases = []struct { - desc string - topic string - domainID string - channelID string - subtopic string - topicType messaging.TopicType - err error -}{ - { - desc: "valid topic with subtopic /m/domain123/c/channel456/devices/temp", - topic: "/m/domain123/c/channel456/devices/temp", - domainID: "domain123", - channelID: "channel456", - subtopic: "devices/temp", - topicType: messaging.MessageType, - }, - { - desc: "topic with wildcards + and # /m/domain123/c/channel456/devices/+/temp/#", - topic: "/m/domain123/c/channel456/devices/+/temp/#", - domainID: "domain123", - channelID: "channel456", - subtopic: "devices/+/temp/#", - topicType: messaging.MessageType, - }, - { - desc: "valid topic without subtopic /m/domain123/c/channel456", - topic: "/m/domain123/c/channel456", - domainID: "domain123", - channelID: "channel456", - subtopic: "", - topicType: messaging.MessageType, - }, - { - desc: "valid topic with trailing slash /m/domain123/c/channel456/devices/temp/", - topic: "/m/domain123/c/channel456/devices/temp/", - domainID: "domain123", - channelID: "channel456", - subtopic: "devices/temp", - topicType: messaging.MessageType, - }, - { - desc: "valid health check topic", - topic: fmt.Sprintf(healthTopicFmt, domainID), - domainID: domainID, - channelID: "", - subtopic: "", - topicType: messaging.HealthType, - err: nil, - }, - { - desc: "invalid health check topic with empty domain", - topic: "hc/", - domainID: "", - channelID: "", - subtopic: "", - topicType: messaging.InvalidType, - err: messaging.ErrMalformedTopic, - }, - { - desc: "invalid topic format (missing channel) /m/domain123/c/", - topic: "/m/domain123/c/", - domainID: "domain123", - channelID: "", - subtopic: "", - err: messaging.ErrMalformedTopic, - }, - { - desc: "invalid topic format (missing domain) /m//c/channel123", - topic: "/m//c/channel123", - domainID: "", - channelID: "", - subtopic: "", - err: messaging.ErrMalformedTopic, - }, - { - desc: "invalid topic format (missing channel) /m/domain123/c/", - topic: "/m/domain123/c//subtopic", - domainID: "domain123", - channelID: "", - subtopic: "", - err: messaging.ErrMalformedTopic, - }, - { - desc: "valid domain with wildcards m/domain*123/c/channel456/devices/+/temp/#", - topic: "m/domain*123/c/channel456/devices/+/temp/#", - domainID: "domain*123", - channelID: "channel456", - subtopic: "devices/+/temp/#", - topicType: messaging.MessageType, - }, - { - desc: "invalid subtopic /m/domain123/c/channel456/sub/a*b/topic", - topic: "/m/domain123/c/channel456/sub/a*b/topic", - domainID: "domain123", - channelID: "channel456", - subtopic: "", - err: messaging.ErrMalformedTopic, - }, - { - desc: "invalid subtopic /m/domain123/c/channel456/sub/a>b/topic", - topic: "/m/domain123/c/channel456/sub/a>b/topic", - domainID: "domain123", - channelID: "channel456", - subtopic: "", - err: messaging.ErrMalformedTopic, - }, - { - desc: "invalid subtopic /m/domain123/c/channel456/sub/a#b/topic", - topic: "/m/domain123/c/channel456/sub/a#b/topic", - domainID: "domain123", - channelID: "channel456", - subtopic: "", - err: messaging.ErrMalformedTopic, - }, - { - desc: "invalid subtopic /m/domain123/c/channel456/sub/a+b/topic", - topic: "/m/domain123/c/channel456/sub/a+b/topic", - domainID: "domain123", - channelID: "channel456", - subtopic: "", - err: messaging.ErrMalformedTopic, - }, - { - desc: "invalid subtopic /m/domain123/c/channel456/sub/a//b/topic", - topic: "/m/domain123/c/channel456/sub/a//b/topic", - domainID: "domain123", - channelID: "channel456", - subtopic: "", - err: messaging.ErrMalformedTopic, - }, - { - desc: "invalid subtopic /m/domain123/c/channel456/sub/a/ /b/topic", - topic: "/m/domain123/c/channel456/sub/a/ /b/topic", - domainID: "domain123", - channelID: "channel456", - subtopic: "", - err: messaging.ErrMalformedTopic, - }, - { - desc: "completely invalid topic \"invalid-topic\"", - topic: "invalid-topic", - domainID: "", - channelID: "", - subtopic: "", - err: messaging.ErrMalformedTopic, - }, - { - desc: "extra segment before prefix /extra/m/domain/c/channel", - topic: "/extra/m/domain/c/channel", - err: messaging.ErrMalformedTopic, - }, -} - -func TestParseSubscribeTopic(t *testing.T) { - for _, tc := range ParseSubscribeTestCases { - t.Run(tc.desc, func(t *testing.T) { - domainID, channelID, subtopic, topicType, err := messaging.ParseSubscribeTopic(tc.topic) - assert.True(t, errors.Contains(err, tc.err), "expected error %v, got %v", tc.err, err) - if err == nil { - assert.Equal(t, tc.domainID, domainID) - assert.Equal(t, tc.channelID, channelID) - assert.Equal(t, tc.subtopic, subtopic) - assert.Equal(t, tc.topicType, topicType) - } - }) - } -} - -func BenchmarkParseSubscribeTopic(b *testing.B) { - for _, tc := range ParseSubscribeTestCases { - b.Run(tc.desc, func(b *testing.B) { - for b.Loop() { - _, _, _, _, _ = messaging.ParseSubscribeTopic(tc.topic) - } - }) - } -} - -func TestEncodeTopic(t *testing.T) { - cases := []struct { - desc string - domainID string - channelID string - subtopic string - expected string - }{ - { - desc: "with subtopic", - domainID: "domain1", - channelID: "chan1", - subtopic: "dev/sensor/temp", - expected: "m/domain1/c/chan1/dev/sensor/temp", - }, - { - desc: "without subtopic", - domainID: "domain1", - channelID: "chan1", - subtopic: "", - expected: "m/domain1/c/chan1", - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - actual := messaging.EncodeTopic(tc.domainID, tc.channelID, tc.subtopic) - assert.Equal(t, tc.expected, actual) - }) - } -} - -func TestEncodeTopicSuffix(t *testing.T) { - cases := []struct { - desc string - domainID string - channelID string - subtopic string - expected string - }{ - { - desc: "with subtopic", - domainID: "domain1", - channelID: "chan1", - subtopic: "dev/sensor/temp", - expected: "domain1/c/chan1/dev/sensor/temp", - }, - { - desc: "without subtopic", - domainID: "domain1", - channelID: "chan1", - subtopic: "", - expected: "domain1/c/chan1", - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - actual := messaging.EncodeTopicSuffix(tc.domainID, tc.channelID, tc.subtopic) - assert.Equal(t, tc.expected, actual) - }) - } -} - -func TestMessage_EncodeTopicSuffix(t *testing.T) { - cases := []struct { - desc string - message *messaging.Message - expected string - }{ - { - desc: "with subtopic", - message: &messaging.Message{ - Domain: "domainX", - Channel: "chanX", - Subtopic: "device/123/status", - }, - expected: "domainX/c/chanX/device/123/status", - }, - { - desc: "without subtopic", - message: &messaging.Message{ - Domain: "domainY", - Channel: "chanY", - }, - expected: "domainY/c/chanY", - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - actual := messaging.EncodeMessageTopic(tc.message) - assert.Equal(t, tc.expected, actual) - }) - } -} - -func TestMessage_EncodeToMQTTTopic(t *testing.T) { - cases := []struct { - desc string - message *messaging.Message - expected string - }{ - { - desc: "with subtopic", - message: &messaging.Message{ - Domain: "domainA", - Channel: "chanA", - Subtopic: "dev/1/temp", - }, - expected: "m/domainA/c/chanA/dev/1/temp", - }, - { - desc: "without subtopic", - message: &messaging.Message{ - Domain: "domainB", - Channel: "chanB", - }, - expected: "m/domainB/c/chanB", - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - actual := messaging.EncodeMessageMQTTTopic(tc.message) - assert.Equal(t, tc.expected, actual) - }) - } -} - -func TestResolve(t *testing.T) { - resolver, domains, channels := setupResolver() - - cases := []struct { - desc string - domain string - channel string - domainID string - channelID string - isRoute bool - domainsErr error - channelsErr error - err error - }{ - { - desc: "valid domainID and channelID", - domain: domainID, - channel: channelID, - domainID: domainID, - channelID: channelID, - isRoute: false, - err: nil, - }, - { - desc: "valid domain route and channel ID", - domain: validRoute, - channel: channelID, - domainID: domainID, - channelID: channelID, - isRoute: true, - err: nil, - }, - { - desc: "valid domain ID and channel route", - domain: domainID, - channel: validRoute, - domainID: domainID, - channelID: channelID, - isRoute: true, - err: nil, - }, - { - desc: "valid domain route and channel route", - domain: validRoute, - channel: validRoute, - domainID: domainID, - channelID: channelID, - isRoute: true, - err: nil, - }, - { - desc: "invalid domain route and valid channel", - domain: invalidRoute, - channel: channelID, - domainID: "", - channelID: "", - domainsErr: svcerr.ErrNotFound, - err: messaging.ErrFailedResolveDomain, - }, - { - desc: "valid domain and invalid channel", - domain: domainID, - channel: invalidRoute, - domainID: domainID, - channelID: "", - channelsErr: svcerr.ErrNotFound, - err: messaging.ErrFailedResolveChannel, - }, - { - desc: "empty domain", - domain: "", - channel: channelID, - domainID: "", - channelID: "", - err: messaging.ErrEmptyRouteID, - }, - { - desc: "empty channel", - domain: domainID, - channel: "", - domainID: domainID, - channelID: "", - err: nil, - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - domainsCall := domains.On("RetrieveIDByRoute", mock.Anything, &grpcCommonV1.RetrieveIDByRouteReq{Route: tc.domain}).Return(&grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Id: tc.domainID, - }, - }, tc.domainsErr) - channelsCall := channels.On("RetrieveIDByRoute", mock.Anything, &grpcCommonV1.RetrieveIDByRouteReq{Route: tc.channel, DomainId: tc.domainID}).Return(&grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Id: tc.channelID, - }, - }, tc.channelsErr) - domainID, channelID, isRoute, err := resolver.Resolve(context.Background(), tc.domain, tc.channel) - assert.True(t, errors.Contains(err, tc.err), "expected error %v, got %v", tc.err, err) - if err == nil { - assert.Equal(t, tc.domainID, domainID, "expected domain ID %s, got %s", tc.domainID, domainID) - assert.Equal(t, tc.channelID, channelID, "expected channel ID %s, got %s", tc.channelID, channelID) - assert.Equal(t, tc.isRoute, isRoute, "expected isRoute %t, got %t", tc.isRoute, isRoute) - } - domainsCall.Unset() - channelsCall.Unset() - }) - } -} - -func TestResolveTopic(t *testing.T) { - resolver, domains, channels := setupResolver() - - cases := []struct { - desc string - topic string - domain string - channel string - domainID string - channelID string - domainsErr error - channelsErr error - response string - err error - }{ - { - desc: "valid topic with domainID and channelID", - topic: fmt.Sprintf(topicFmt, domainID, channelID), - domain: domainID, - channel: channelID, - domainID: domainID, - channelID: channelID, - response: fmt.Sprintf(topicFmt, domainID, channelID), - err: nil, - }, - { - desc: "valid topic with domain route and channel ID", - topic: fmt.Sprintf(topicFmt, validRoute, channelID), - domain: validRoute, - channel: channelID, - domainID: domainID, - channelID: channelID, - response: fmt.Sprintf(topicFmt, domainID, channelID), - err: nil, - }, - { - desc: "valid topic with domain ID and channel route", - topic: fmt.Sprintf(topicFmt, domainID, validRoute), - domain: domainID, - channel: validRoute, - domainID: domainID, - channelID: channelID, - response: fmt.Sprintf(topicFmt, domainID, channelID), - err: nil, - }, - { - desc: "valid topic with domain route and channel route", - topic: fmt.Sprintf(topicFmt, validRoute, validRoute), - domain: validRoute, - channel: validRoute, - domainID: domainID, - channelID: channelID, - response: fmt.Sprintf(topicFmt, domainID, channelID), - err: nil, - }, - { - desc: "invalid topic with invalid domain route and valid channel", - topic: fmt.Sprintf(topicFmt, invalidRoute, channelID), - domain: invalidRoute, - channel: channelID, - domainID: "", - channelID: "", - domainsErr: svcerr.ErrNotFound, - err: messaging.ErrFailedResolveDomain, - }, - { - desc: "valid topic with valid topic with domainID and channelID and subtopic", - topic: fmt.Sprintf(topicFmt, domainID, channelID) + "/subtopic", - domain: domainID, - channel: channelID, - domainID: domainID, - channelID: channelID, - response: fmt.Sprintf(topicFmt, domainID, channelID) + "/subtopic", - err: nil, - }, - { - desc: "invalid topic with empty domain", - topic: fmt.Sprintf(topicFmt, "", channelID), - domain: "", - channel: channelID, - domainID: "", - channelID: "", - err: messaging.ErrMalformedTopic, - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - domainsCall := domains.On("RetrieveIDByRoute", mock.Anything, &grpcCommonV1.RetrieveIDByRouteReq{Route: tc.domain}).Return(&grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Id: tc.domainID, - }, - }, tc.domainsErr) - channelsCall := channels.On("RetrieveIDByRoute", mock.Anything, &grpcCommonV1.RetrieveIDByRouteReq{Route: tc.channel, DomainId: tc.domainID}).Return(&grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Id: tc.channelID, - }, - }, tc.channelsErr) - rtopic, err := resolver.ResolveTopic(context.Background(), tc.topic) - assert.True(t, errors.Contains(err, tc.err), "expected error %v, got %v", tc.err, err) - if err == nil { - assert.Equal(t, tc.response, rtopic, "expected topic %s, got %s", tc.response, rtopic) - } - domainsCall.Unset() - channelsCall.Unset() - }) - } -} - -func TestParserPublishTopic(t *testing.T) { - parser, domains, channels, err := setupParser() - assert.Nil(t, err, fmt.Sprintf("unexpected error while setting up parser: %v", err)) - - udomainID := testsutil.GenerateUUID(t) - uchannelID := testsutil.GenerateUUID(t) - - cachedInvalidTopic := "m/invalid-domain/c" - - dom, ch, st, tt, err := parser.ParsePublishTopic(context.Background(), cachedTopic, false) - assert.Nil(t, err, fmt.Sprintf("unexpected error while publishing topic: %v", err)) - assert.Equal(t, domainID, dom, "expected domainID %s, got %s", domainID, dom) - assert.Equal(t, channelID, ch, "expected channelID %s, got %s", channelID, ch) - assert.Equal(t, subtopic, st, "expected subtopic %s, got %s", subtopic, st) - assert.Equal(t, messaging.MessageType, tt, "expected topic type %v, got %v", messaging.MessageType, tt) - - dom, ch, st, tt, err = parser.ParsePublishTopic(context.Background(), cachedInvalidTopic, false) - assert.NotNil(t, err, "expected error for invalid cached topic") - assert.Equal(t, "", dom, "expected empty domainID for invalid topic") - assert.Equal(t, "", ch, "expected empty channelID for invalid topic") - assert.Equal(t, "", st, "expected empty subtopic for invalid topic") - assert.Equal(t, messaging.InvalidType, tt, "expected unknown topic type for invalid topic") - time.Sleep(10 * time.Millisecond) // Ensure cache is populated - - cases := []struct { - desc string - topic string - resolve bool - domain string - channel string - domainID string - channelID string - subtopic string - topicType messaging.TopicType - domainsErr error - channelsErr error - err error - }{ - { - desc: "valid uncached topic with domainID and channelID", - topic: fmt.Sprintf(topicFmt, udomainID, uchannelID) + "/subtopic", - resolve: true, - domain: udomainID, - channel: uchannelID, - domainID: udomainID, - channelID: uchannelID, - subtopic: subtopic, - topicType: messaging.MessageType, - err: nil, - }, - { - desc: "valid cached topic with domainID and channelID", - topic: cachedTopic, - domain: domainID, - channel: channelID, - domainID: domainID, - channelID: channelID, - subtopic: subtopic, - topicType: messaging.MessageType, - err: nil, - }, - { - desc: "invalid uncached topic with invalid format", - topic: "invalid-topic", - domain: "", - channel: "", - domainID: "", - channelID: "", - err: messaging.ErrMalformedTopic, - }, - { - desc: "invalid cached topic with invalid format", - topic: cachedInvalidTopic, - domain: "", - channel: "", - domainID: "", - channelID: "", - err: messaging.ErrMalformedTopic, - }, - { - desc: "valid uncached topic with domain and channel routes", - topic: fmt.Sprintf(topicFmt, validRoute, validRoute) + "/subtopic", - resolve: true, - domain: validRoute, - channel: validRoute, - domainID: domainID, - channelID: channelID, - subtopic: subtopic, - topicType: messaging.MessageType, - err: nil, - }, - { - desc: "valid uncached topic with failed domain resolution", - topic: fmt.Sprintf(topicFmt, invalidRoute, uchannelID) + "/subtopic", - resolve: true, - domain: invalidRoute, - channel: uchannelID, - domainID: "", - channelID: "", - domainsErr: svcerr.ErrNotFound, - err: messaging.ErrFailedResolveDomain, - }, - { - desc: "valid uncached healthcheck topic", - topic: fmt.Sprintf(healthTopicFmt, domainID), - domain: domainID, - channel: "", - domainID: domainID, - channelID: "", - subtopic: "", - topicType: messaging.HealthType, - err: nil, - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - domainsCall := domains.On("RetrieveIDByRoute", mock.Anything, &grpcCommonV1.RetrieveIDByRouteReq{Route: tc.domain}).Return(&grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Id: tc.domainID, - }, - }, tc.domainsErr) - channelsCall := channels.On("RetrieveIDByRoute", mock.Anything, &grpcCommonV1.RetrieveIDByRouteReq{Route: tc.channel, DomainId: tc.domainID}).Return(&grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Id: tc.channelID, - }, - }, tc.channelsErr) - domainID, channelID, subtopic, topicType, err := parser.ParsePublishTopic(context.Background(), tc.topic, tc.resolve) - assert.True(t, errors.Contains(err, tc.err), "expected error %v, got %v", tc.err, err) - if err == nil { - assert.Equal(t, tc.domainID, domainID, "expected domainID %s, got %s", tc.domainID, domainID) - assert.Equal(t, tc.channelID, channelID, "expected channelID %s, got %s", tc.channelID, channelID) - assert.Equal(t, tc.subtopic, subtopic, "expected subtopic %s, got %s", tc.subtopic, subtopic) - assert.Equal(t, tc.topicType, topicType, "expected topic type %v, got %v", tc.topicType, topicType) - } - domainsCall.Unset() - channelsCall.Unset() - }) - } -} - -func BenchmarkParserPublishTopic(b *testing.B) { - parser, _, _, err := setupParser() - if err != nil { - b.Fatalf("unexpected error while setting up parser: %v", err) - } - - for _, tc := range ParsePublisherTopicTestCases { - b.Run(tc.desc, func(b *testing.B) { - for b.Loop() { - _, _, _, _, _ = parser.ParsePublishTopic(context.Background(), tc.topic, false) - } - }) - } -} - -func TestParserSubscribeTopic(t *testing.T) { - parser, domains, channels, err := setupParser() - assert.Nil(t, err, fmt.Sprintf("unexpected error while setting up parser: %v", err)) - - cases := []struct { - desc string - topic string - resolve bool - domain string - channel string - domainID string - channelID string - subtopic string - topicType messaging.TopicType - domainsErr error - channelsErr error - err error - }{ - { - desc: "valid topic with domainID and channelID", - topic: fmt.Sprintf(topicFmt, domainID, channelID), - resolve: true, - domain: domainID, - channel: channelID, - domainID: domainID, - channelID: channelID, - topicType: messaging.MessageType, - err: nil, - }, - { - desc: "valid topic with domainID and channelID and subtopic", - topic: fmt.Sprintf(topicSubtopicFmt, domainID, channelID, subtopic), - resolve: true, - domain: domainID, - channel: channelID, - domainID: domainID, - channelID: channelID, - subtopic: subtopic, - topicType: messaging.MessageType, - err: nil, - }, - { - desc: "valid topic with domain and channel routes", - topic: fmt.Sprintf(topicFmt, validRoute, validRoute), - resolve: true, - domain: validRoute, - channel: validRoute, - domainID: domainID, - channelID: channelID, - topicType: messaging.MessageType, - err: nil, - }, - { - desc: "invalid topic with invalid format", - topic: "invalid-topic", - resolve: false, - domain: "", - channel: "", - domainID: "", - channelID: "", - err: messaging.ErrMalformedTopic, - }, - { - desc: "valid topic with invalid domain route", - topic: fmt.Sprintf(topicFmt, invalidRoute, validRoute), - resolve: true, - domain: invalidRoute, - channel: validRoute, - domainID: "", - channelID: "", - domainsErr: svcerr.ErrNotFound, - err: messaging.ErrFailedResolveDomain, - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - domainsCall := domains.On("RetrieveIDByRoute", mock.Anything, &grpcCommonV1.RetrieveIDByRouteReq{Route: tc.domain}).Return(&grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Id: tc.domainID, - }, - }, tc.domainsErr) - channelsCall := channels.On("RetrieveIDByRoute", mock.Anything, &grpcCommonV1.RetrieveIDByRouteReq{Route: tc.channel, DomainId: tc.domainID}).Return(&grpcCommonV1.RetrieveEntityRes{ - Entity: &grpcCommonV1.EntityBasic{ - Id: tc.channelID, - }, - }, tc.channelsErr) - dom, ch, st, tt, err := parser.ParseSubscribeTopic(context.Background(), tc.topic, tc.resolve) - assert.True(t, errors.Contains(err, tc.err), "expected error %v, got %v", tc.err, err) - if err == nil { - assert.Equal(t, tc.domainID, dom, "expected domainID %s, got %s", tc.domainID, dom) - assert.Equal(t, tc.channelID, ch, "expected channelID %s, got %s", tc.channelID, ch) - assert.Equal(t, tc.subtopic, st, "expected subtopic %s, got %s", tc.subtopic, st) - assert.Equal(t, tc.topicType, tt, "expected topic type %v, got %v", tc.topicType, tt) - } - domainsCall.Unset() - channelsCall.Unset() - }) - } -} diff --git a/pkg/oauth2/doc.go b/pkg/oauth2/doc.go deleted file mode 100644 index 2d7e006f5..000000000 --- a/pkg/oauth2/doc.go +++ /dev/null @@ -1,6 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package oauth2 contains the domain concept definitions needed to support -// Magistrala ui service OAuth2 functionality. -package oauth2 diff --git a/pkg/oauth2/google/doc.go b/pkg/oauth2/google/doc.go deleted file mode 100644 index 74f7ada5e..000000000 --- a/pkg/oauth2/google/doc.go +++ /dev/null @@ -1,6 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package google contains the domain concept definitions needed to support -// Magistrala services for Google OAuth2 functionality. -package google diff --git a/pkg/oauth2/google/provider.go b/pkg/oauth2/google/provider.go deleted file mode 100644 index c44634d5c..000000000 --- a/pkg/oauth2/google/provider.go +++ /dev/null @@ -1,113 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package google - -import ( - "context" - "io" - "net/http" - "net/url" - "time" - - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - mgoauth2 "github.com/absmach/magistrala/pkg/oauth2" - uclient "github.com/absmach/magistrala/users" - "golang.org/x/oauth2" - googleoauth2 "golang.org/x/oauth2/google" -) - -const ( - providerName = "google" - defTimeout = 1 * time.Minute - userInfoURL = "https://www.googleapis.com/oauth2/v2/userinfo?access_token=" - tokenInfoURL = "https://oauth2.googleapis.com/tokeninfo?access_token=" -) - -var scopes = []string{ - "https://www.googleapis.com/auth/userinfo.email", - "https://www.googleapis.com/auth/userinfo.profile", -} - -var httpClient = &http.Client{ - Timeout: defTimeout, -} - -var _ mgoauth2.Provider = (*config)(nil) - -type config struct { - config *oauth2.Config - state string - uiRedirectURL string - errorURL string -} - -// NewProvider returns a new Google OAuth provider. -func NewProvider(cfg mgoauth2.Config, uiRedirectURL, errorURL string) mgoauth2.Provider { - return &config{ - config: &oauth2.Config{ - ClientID: cfg.ClientID, - ClientSecret: cfg.ClientSecret, - Endpoint: googleoauth2.Endpoint, - RedirectURL: cfg.RedirectURL, - Scopes: scopes, - }, - state: cfg.State, - uiRedirectURL: uiRedirectURL, - errorURL: errorURL, - } -} - -func (cfg *config) Name() string { - return providerName -} - -func (cfg *config) State() string { - return cfg.state -} - -func (cfg *config) RedirectURL() string { - return cfg.uiRedirectURL -} - -func (cfg *config) ErrorURL() string { - return cfg.errorURL -} - -func (cfg *config) IsEnabled() bool { - return cfg.config.ClientID != "" && cfg.config.ClientSecret != "" -} - -func (cfg *config) Exchange(ctx context.Context, code string) (oauth2.Token, error) { - token, err := cfg.config.Exchange(ctx, code) - if err != nil { - return oauth2.Token{}, err - } - - return *token, nil -} - -func (cfg *config) UserInfo(accessToken string) (uclient.User, error) { - resp, err := httpClient.Get(userInfoURL + url.QueryEscape(accessToken)) - if err != nil { - return uclient.User{}, err - } - defer resp.Body.Close() - - if resp.StatusCode != http.StatusOK { - return uclient.User{}, svcerr.ErrAuthentication - } - - data, err := io.ReadAll(resp.Body) - if err != nil { - return uclient.User{}, err - } - - user, err := mgoauth2.NormalizeUser(data, providerName) - if err != nil { - return uclient.User{}, errors.Wrap(err, svcerr.ErrAuthentication) - } - - return user, nil -} diff --git a/pkg/oauth2/mocks/provider.go b/pkg/oauth2/mocks/provider.go deleted file mode 100644 index 8ee04e0b5..000000000 --- a/pkg/oauth2/mocks/provider.go +++ /dev/null @@ -1,390 +0,0 @@ -// Copyright (c) Abstract Machines - -// SPDX-License-Identifier: Apache-2.0 - -// Code generated by mockery; DO NOT EDIT. -// github.com/vektra/mockery -// template: testify - -package mocks - -import ( - "context" - - "github.com/absmach/magistrala/users" - mock "github.com/stretchr/testify/mock" - "golang.org/x/oauth2" -) - -// NewProvider creates a new instance of Provider. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. -// The first argument is typically a *testing.T value. -func NewProvider(t interface { - mock.TestingT - Cleanup(func()) -}) *Provider { - mock := &Provider{} - mock.Mock.Test(t) - - t.Cleanup(func() { mock.AssertExpectations(t) }) - - return mock -} - -// Provider is an autogenerated mock type for the Provider type -type Provider struct { - mock.Mock -} - -type Provider_Expecter struct { - mock *mock.Mock -} - -func (_m *Provider) EXPECT() *Provider_Expecter { - return &Provider_Expecter{mock: &_m.Mock} -} - -// ErrorURL provides a mock function for the type Provider -func (_mock *Provider) ErrorURL() string { - ret := _mock.Called() - - if len(ret) == 0 { - panic("no return value specified for ErrorURL") - } - - var r0 string - if returnFunc, ok := ret.Get(0).(func() string); ok { - r0 = returnFunc() - } else { - r0 = ret.Get(0).(string) - } - return r0 -} - -// Provider_ErrorURL_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ErrorURL' -type Provider_ErrorURL_Call struct { - *mock.Call -} - -// ErrorURL is a helper method to define mock.On call -func (_e *Provider_Expecter) ErrorURL() *Provider_ErrorURL_Call { - return &Provider_ErrorURL_Call{Call: _e.mock.On("ErrorURL")} -} - -func (_c *Provider_ErrorURL_Call) Run(run func()) *Provider_ErrorURL_Call { - _c.Call.Run(func(args mock.Arguments) { - run() - }) - return _c -} - -func (_c *Provider_ErrorURL_Call) Return(s string) *Provider_ErrorURL_Call { - _c.Call.Return(s) - return _c -} - -func (_c *Provider_ErrorURL_Call) RunAndReturn(run func() string) *Provider_ErrorURL_Call { - _c.Call.Return(run) - return _c -} - -// Exchange provides a mock function for the type Provider -func (_mock *Provider) Exchange(ctx context.Context, code string) (oauth2.Token, error) { - ret := _mock.Called(ctx, code) - - if len(ret) == 0 { - panic("no return value specified for Exchange") - } - - var r0 oauth2.Token - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) (oauth2.Token, error)); ok { - return returnFunc(ctx, code) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string) oauth2.Token); ok { - r0 = returnFunc(ctx, code) - } else { - r0 = ret.Get(0).(oauth2.Token) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = returnFunc(ctx, code) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Provider_Exchange_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Exchange' -type Provider_Exchange_Call struct { - *mock.Call -} - -// Exchange is a helper method to define mock.On call -// - ctx context.Context -// - code string -func (_e *Provider_Expecter) Exchange(ctx interface{}, code interface{}) *Provider_Exchange_Call { - return &Provider_Exchange_Call{Call: _e.mock.On("Exchange", ctx, code)} -} - -func (_c *Provider_Exchange_Call) Run(run func(ctx context.Context, code string)) *Provider_Exchange_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Provider_Exchange_Call) Return(token oauth2.Token, err error) *Provider_Exchange_Call { - _c.Call.Return(token, err) - return _c -} - -func (_c *Provider_Exchange_Call) RunAndReturn(run func(ctx context.Context, code string) (oauth2.Token, error)) *Provider_Exchange_Call { - _c.Call.Return(run) - return _c -} - -// IsEnabled provides a mock function for the type Provider -func (_mock *Provider) IsEnabled() bool { - ret := _mock.Called() - - if len(ret) == 0 { - panic("no return value specified for IsEnabled") - } - - var r0 bool - if returnFunc, ok := ret.Get(0).(func() bool); ok { - r0 = returnFunc() - } else { - r0 = ret.Get(0).(bool) - } - return r0 -} - -// Provider_IsEnabled_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'IsEnabled' -type Provider_IsEnabled_Call struct { - *mock.Call -} - -// IsEnabled is a helper method to define mock.On call -func (_e *Provider_Expecter) IsEnabled() *Provider_IsEnabled_Call { - return &Provider_IsEnabled_Call{Call: _e.mock.On("IsEnabled")} -} - -func (_c *Provider_IsEnabled_Call) Run(run func()) *Provider_IsEnabled_Call { - _c.Call.Run(func(args mock.Arguments) { - run() - }) - return _c -} - -func (_c *Provider_IsEnabled_Call) Return(b bool) *Provider_IsEnabled_Call { - _c.Call.Return(b) - return _c -} - -func (_c *Provider_IsEnabled_Call) RunAndReturn(run func() bool) *Provider_IsEnabled_Call { - _c.Call.Return(run) - return _c -} - -// Name provides a mock function for the type Provider -func (_mock *Provider) Name() string { - ret := _mock.Called() - - if len(ret) == 0 { - panic("no return value specified for Name") - } - - var r0 string - if returnFunc, ok := ret.Get(0).(func() string); ok { - r0 = returnFunc() - } else { - r0 = ret.Get(0).(string) - } - return r0 -} - -// Provider_Name_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Name' -type Provider_Name_Call struct { - *mock.Call -} - -// Name is a helper method to define mock.On call -func (_e *Provider_Expecter) Name() *Provider_Name_Call { - return &Provider_Name_Call{Call: _e.mock.On("Name")} -} - -func (_c *Provider_Name_Call) Run(run func()) *Provider_Name_Call { - _c.Call.Run(func(args mock.Arguments) { - run() - }) - return _c -} - -func (_c *Provider_Name_Call) Return(s string) *Provider_Name_Call { - _c.Call.Return(s) - return _c -} - -func (_c *Provider_Name_Call) RunAndReturn(run func() string) *Provider_Name_Call { - _c.Call.Return(run) - return _c -} - -// RedirectURL provides a mock function for the type Provider -func (_mock *Provider) RedirectURL() string { - ret := _mock.Called() - - if len(ret) == 0 { - panic("no return value specified for RedirectURL") - } - - var r0 string - if returnFunc, ok := ret.Get(0).(func() string); ok { - r0 = returnFunc() - } else { - r0 = ret.Get(0).(string) - } - return r0 -} - -// Provider_RedirectURL_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RedirectURL' -type Provider_RedirectURL_Call struct { - *mock.Call -} - -// RedirectURL is a helper method to define mock.On call -func (_e *Provider_Expecter) RedirectURL() *Provider_RedirectURL_Call { - return &Provider_RedirectURL_Call{Call: _e.mock.On("RedirectURL")} -} - -func (_c *Provider_RedirectURL_Call) Run(run func()) *Provider_RedirectURL_Call { - _c.Call.Run(func(args mock.Arguments) { - run() - }) - return _c -} - -func (_c *Provider_RedirectURL_Call) Return(s string) *Provider_RedirectURL_Call { - _c.Call.Return(s) - return _c -} - -func (_c *Provider_RedirectURL_Call) RunAndReturn(run func() string) *Provider_RedirectURL_Call { - _c.Call.Return(run) - return _c -} - -// State provides a mock function for the type Provider -func (_mock *Provider) State() string { - ret := _mock.Called() - - if len(ret) == 0 { - panic("no return value specified for State") - } - - var r0 string - if returnFunc, ok := ret.Get(0).(func() string); ok { - r0 = returnFunc() - } else { - r0 = ret.Get(0).(string) - } - return r0 -} - -// Provider_State_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'State' -type Provider_State_Call struct { - *mock.Call -} - -// State is a helper method to define mock.On call -func (_e *Provider_Expecter) State() *Provider_State_Call { - return &Provider_State_Call{Call: _e.mock.On("State")} -} - -func (_c *Provider_State_Call) Run(run func()) *Provider_State_Call { - _c.Call.Run(func(args mock.Arguments) { - run() - }) - return _c -} - -func (_c *Provider_State_Call) Return(s string) *Provider_State_Call { - _c.Call.Return(s) - return _c -} - -func (_c *Provider_State_Call) RunAndReturn(run func() string) *Provider_State_Call { - _c.Call.Return(run) - return _c -} - -// UserInfo provides a mock function for the type Provider -func (_mock *Provider) UserInfo(accessToken string) (users.User, error) { - ret := _mock.Called(accessToken) - - if len(ret) == 0 { - panic("no return value specified for UserInfo") - } - - var r0 users.User - var r1 error - if returnFunc, ok := ret.Get(0).(func(string) (users.User, error)); ok { - return returnFunc(accessToken) - } - if returnFunc, ok := ret.Get(0).(func(string) users.User); ok { - r0 = returnFunc(accessToken) - } else { - r0 = ret.Get(0).(users.User) - } - if returnFunc, ok := ret.Get(1).(func(string) error); ok { - r1 = returnFunc(accessToken) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Provider_UserInfo_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UserInfo' -type Provider_UserInfo_Call struct { - *mock.Call -} - -// UserInfo is a helper method to define mock.On call -// - accessToken string -func (_e *Provider_Expecter) UserInfo(accessToken interface{}) *Provider_UserInfo_Call { - return &Provider_UserInfo_Call{Call: _e.mock.On("UserInfo", accessToken)} -} - -func (_c *Provider_UserInfo_Call) Run(run func(accessToken string)) *Provider_UserInfo_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 string - if args[0] != nil { - arg0 = args[0].(string) - } - run( - arg0, - ) - }) - return _c -} - -func (_c *Provider_UserInfo_Call) Return(user users.User, err error) *Provider_UserInfo_Call { - _c.Call.Return(user, err) - return _c -} - -func (_c *Provider_UserInfo_Call) RunAndReturn(run func(accessToken string) (users.User, error)) *Provider_UserInfo_Call { - _c.Call.Return(run) - return _c -} diff --git a/pkg/oauth2/normalize.go b/pkg/oauth2/normalize.go deleted file mode 100644 index 63d5518b5..000000000 --- a/pkg/oauth2/normalize.go +++ /dev/null @@ -1,97 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package oauth2 - -import ( - "encoding/json" - "fmt" - "strings" - - "github.com/absmach/magistrala/users" -) - -type normalizedUser struct { - ID string `json:"id"` - FirstName string `json:"first_name"` - LastName string `json:"last_name"` - Username string `json:"username"` - Email string `json:"email"` - Picture string `json:"picture"` -} - -func NormalizeUser(data []byte, provider string) (users.User, error) { - var raw map[string]any - if err := json.Unmarshal(data, &raw); err != nil { - return users.User{}, err - } - - normalized := normalizeProfile(raw) - - userBytes, err := json.Marshal(normalized) - if err != nil { - return users.User{}, err - } - - var user normalizedUser - if err := json.Unmarshal(userBytes, &user); err != nil { - return users.User{}, err - } - - if err := validateUser(user); err != nil { - return users.User{}, err - } - - return users.User{ - ID: user.ID, - FirstName: user.FirstName, - LastName: user.LastName, - Email: user.Email, - ProfilePicture: user.Picture, - Metadata: users.Metadata{"oauth_provider": provider}, - }, nil -} - -func normalizeProfile(raw map[string]any) map[string]any { - normalized := make(map[string]any) - - keyMap := map[string][]string{ - "id": {"id"}, - "first_name": {"given_name", "first_name", "givenName", "firstname"}, - "last_name": {"family_name", "last_name", "familyName", "lastname"}, - "username": {"username", "user_name", "userName"}, - "email": {"email", "email_address", "emailAddress"}, - "picture": {"picture", "profile_picture", "profilePicture", "avatar"}, - } - - for stdKey, variants := range keyMap { - for _, variant := range variants { - if val, ok := raw[variant]; ok { - normalized[stdKey] = val - break - } - } - } - - return normalized -} - -func validateUser(user normalizedUser) error { - var missing []string - if user.ID == "" { - missing = append(missing, "id") - } - if user.FirstName == "" { - missing = append(missing, "first_name") - } - if user.LastName == "" { - missing = append(missing, "last_name") - } - if user.Email == "" { - missing = append(missing, "email") - } - if len(missing) > 0 { - return fmt.Errorf("missing required fields: %s", strings.Join(missing, ", ")) - } - return nil -} diff --git a/pkg/oauth2/normalize_test.go b/pkg/oauth2/normalize_test.go deleted file mode 100644 index 15614086f..000000000 --- a/pkg/oauth2/normalize_test.go +++ /dev/null @@ -1,157 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package oauth2 - -import ( - "testing" - - "github.com/absmach/magistrala/users" - "github.com/stretchr/testify/assert" -) - -func TestNormalizeUser(t *testing.T) { - cases := []struct { - desc string - inputJSON string - provider string - wantUser users.User - wantErrStr string - }{ - { - desc: "valid user with standard keys", - inputJSON: `{ - "id": "123", - "given_name": "Jane", - "family_name": "Doe", - "email": "jane@example.com", - "picture": "pic.jpg" - }`, - provider: "google", - wantUser: users.User{ - ID: "123", - FirstName: "Jane", - LastName: "Doe", - Email: "jane@example.com", - ProfilePicture: "pic.jpg", - Metadata: users.Metadata{"oauth_provider": "google"}, - }, - wantErrStr: "", - }, - { - desc: "missing required fields", - inputJSON: `{ - "given_name": "Jane" - }`, - provider: "google", - wantUser: users.User{}, - wantErrStr: "missing required fields: id, last_name, email", - }, - { - desc: "invalid JSON", - inputJSON: `{invalid json`, - provider: "google", - wantUser: users.User{}, - wantErrStr: "invalid character", - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - user, err := NormalizeUser([]byte(tc.inputJSON), tc.provider) - if tc.wantErrStr != "" { - assert.Error(t, err) - assert.Contains(t, err.Error(), tc.wantErrStr) - assert.Equal(t, tc.wantUser, user) - } else { - assert.NoError(t, err) - assert.Equal(t, tc.wantUser, user) - } - }) - } -} - -func TestNormalizeProfile(t *testing.T) { - cases := []struct { - desc string - raw map[string]any - expected map[string]any - }{ - { - desc: "maps all variants to normalized keys", - raw: map[string]any{ - "id": "id123", - "givenName": "John", - "familyName": "Smith", - "user_name": "jsmith", - "emailAddress": "john@smith.com", - "profilePicture": "pic.png", - }, - expected: map[string]any{ - "id": "id123", - "first_name": "John", - "last_name": "Smith", - "username": "jsmith", - "email": "john@smith.com", - "picture": "pic.png", - }, - }, - { - desc: "missing keys returns empty map", - raw: map[string]any{"foo": "bar"}, - expected: map[string]any{}, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - got := normalizeProfile(tc.raw) - assert.Equal(t, tc.expected, got) - }) - } -} - -func TestValidateUser(t *testing.T) { - cases := []struct { - desc string - user normalizedUser - wantErr string - }{ - { - desc: "valid user returns nil error", - user: normalizedUser{ - ID: "1", - FirstName: "F", - LastName: "L", - Email: "e@example.com", - }, - wantErr: "", - }, - { - desc: "missing id returns error", - user: normalizedUser{ - FirstName: "F", - LastName: "L", - Email: "e@example.com", - }, - wantErr: "missing required fields: id", - }, - { - desc: "multiple missing fields returns all in error", - user: normalizedUser{}, - wantErr: "missing required fields: id, first_name, last_name, email", - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := validateUser(tc.user) - if tc.wantErr == "" { - assert.NoError(t, err) - } else { - assert.Error(t, err) - assert.Equal(t, tc.wantErr, err.Error()) - } - }) - } -} diff --git a/pkg/oauth2/oauth2.go b/pkg/oauth2/oauth2.go deleted file mode 100644 index d8ec5745b..000000000 --- a/pkg/oauth2/oauth2.go +++ /dev/null @@ -1,44 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package oauth2 - -import ( - "context" - - "github.com/absmach/magistrala/users" - "golang.org/x/oauth2" -) - -// Config is the configuration for the OAuth2 provider. -type Config struct { - ClientID string `env:"CLIENT_ID" envDefault:""` - ClientSecret string `env:"CLIENT_SECRET" envDefault:""` - State string `env:"STATE" envDefault:""` - RedirectURL string `env:"REDIRECT_URL" envDefault:""` -} - -// Provider is an interface that provides the OAuth2 flow for a specific provider -// (e.g. Google, GitHub, etc.) -type Provider interface { - // Name returns the name of the OAuth2 provider. - Name() string - - // State returns the current state for the OAuth2 flow. - State() string - - // RedirectURL returns the URL to redirect the user to after completing the OAuth2 flow. - RedirectURL() string - - // ErrorURL returns the URL to redirect the user to in case of an error during the OAuth2 flow. - ErrorURL() string - - // IsEnabled checks if the OAuth2 provider is enabled. - IsEnabled() bool - - // Exchange converts an authorization code into a token. - Exchange(ctx context.Context, code string) (oauth2.Token, error) - - // UserInfo retrieves the user's information using the access token. - UserInfo(accessToken string) (users.User, error) -} diff --git a/pkg/policies/service.go b/pkg/policies/service.go index cb7aab772..996c6956a 100644 --- a/pkg/policies/service.go +++ b/pkg/policies/service.go @@ -66,8 +66,7 @@ type PolicyPage struct { type Permissions []string -// PolicyService facilitates the communication to authorization -// services and implements Authz functionalities for spicedb. +// Service facilitates communication with an authorization backend. type Service interface { // AddPolicy creates a policy for the given subject, so that, after // AddPolicy, `subject` has a `relation` on `object`. Returns a non-nil diff --git a/pkg/policies/spicedb/doc.go b/pkg/policies/spicedb/doc.go deleted file mode 100644 index beac26947..000000000 --- a/pkg/policies/spicedb/doc.go +++ /dev/null @@ -1,5 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package server contains the HTTP, gRPC and CoAP server implementation. -package spicedb diff --git a/pkg/policies/spicedb/evaluator.go b/pkg/policies/spicedb/evaluator.go deleted file mode 100644 index e40b7207b..000000000 --- a/pkg/policies/spicedb/evaluator.go +++ /dev/null @@ -1,64 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package spicedb - -import ( - "context" - "log/slog" - - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/pkg/policies" - v1 "github.com/authzed/authzed-go/proto/authzed/api/v1" - "github.com/authzed/authzed-go/v1" -) - -type policyEvaluator struct { - client *authzed.ClientWithExperimental - permissionClient v1.PermissionsServiceClient - logger *slog.Logger -} - -func NewPolicyEvaluator(client *authzed.ClientWithExperimental, logger *slog.Logger) policies.Evaluator { - return &policyEvaluator{ - client: client, - permissionClient: client.PermissionsServiceClient, - logger: logger, - } -} - -func (pe *policyEvaluator) CheckPolicy(ctx context.Context, pr policies.Policy) error { - checkReq := v1.CheckPermissionRequest{ - // FullyConsistent means little caching will be available, which means performance will suffer. - // Only use if a ZedToken is not available or absolutely latest information is required. - // If we want to avoid FullyConsistent and to improve the performance of spicedb, then we need to cache the ZEDTOKEN whenever RELATIONS is created or updated. - // Instead of using FullyConsistent we need to use Consistency_AtLeastAsFresh, code looks like below one. - // Consistency: &v1.Consistency{ - // Requirement: &v1.Consistency_AtLeastAsFresh{ - // AtLeastAsFresh: getRelationTupleZedTokenFromCache() , - // } - // }, - // Reference: https://authzed.com/docs/reference/api-consistency - Consistency: &v1.Consistency{ - Requirement: &v1.Consistency_FullyConsistent{ - FullyConsistent: true, - }, - }, - Resource: &v1.ObjectReference{ObjectType: pr.ObjectType, ObjectId: pr.Object}, - Permission: pr.Permission, - Subject: &v1.SubjectReference{Object: &v1.ObjectReference{ObjectType: pr.SubjectType, ObjectId: pr.Subject}, OptionalRelation: pr.SubjectRelation}, - } - - resp, err := pe.permissionClient.CheckPermission(ctx, &checkReq) - if err != nil { - return handleSpicedbError(err) - } - if resp.Permissionship == v1.CheckPermissionResponse_PERMISSIONSHIP_HAS_PERMISSION { - return nil - } - if reason, ok := v1.CheckPermissionResponse_Permissionship_name[int32(resp.Permissionship)]; ok { - return errors.Wrap(svcerr.ErrAuthorization, errors.New(reason)) - } - return svcerr.ErrAuthorization -} diff --git a/pkg/policies/spicedb/service.go b/pkg/policies/spicedb/service.go deleted file mode 100644 index 1ebce14e2..000000000 --- a/pkg/policies/spicedb/service.go +++ /dev/null @@ -1,906 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package spicedb - -import ( - "context" - "fmt" - "io" - "log/slog" - - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/pkg/policies" - v1 "github.com/authzed/authzed-go/proto/authzed/api/v1" - "github.com/authzed/authzed-go/v1" - gstatus "google.golang.org/genproto/googleapis/rpc/status" - "google.golang.org/grpc/codes" - "google.golang.org/grpc/status" -) - -const defRetrieveAllLimit = 1000 - -var ( - errAddPolicies = errors.New("failed to add policies") - errRetrievePolicies = errors.New("failed to retrieve policies") - errRemovePolicies = errors.New("failed to remove the policies") - errNoPolicies = errors.New("no policies provided") - errInternal = errors.New("spicedb internal error") - errPlatform = errors.New("invalid platform id") -) - -var ( - defClientsFilterPermissions = []string{ - policies.AdminPermission, - policies.DeletePermission, - policies.EditPermission, - policies.ViewPermission, - policies.SharePermission, - policies.PublishPermission, - policies.SubscribePermission, - } - - defGroupsFilterPermissions = []string{ - policies.AdminPermission, - policies.DeletePermission, - policies.EditPermission, - policies.ViewPermission, - policies.MembershipPermission, - policies.SharePermission, - } - - defDomainsFilterPermissions = []string{ - policies.AdminPermission, - policies.EditPermission, - policies.ViewPermission, - policies.MembershipPermission, - policies.SharePermission, - } - - defPlatformFilterPermissions = []string{ - policies.AdminPermission, - policies.MembershipPermission, - } -) - -type policyService struct { - client *authzed.ClientWithExperimental - permissionClient v1.PermissionsServiceClient - logger *slog.Logger -} - -func NewPolicyService(client *authzed.ClientWithExperimental, logger *slog.Logger) policies.Service { - return &policyService{ - client: client, - permissionClient: client.PermissionsServiceClient, - logger: logger, - } -} - -func (ps *policyService) AddPolicy(ctx context.Context, pr policies.Policy) error { - if err := ps.policyValidation(pr); err != nil { - return errors.Wrap(svcerr.ErrInvalidPolicy, err) - } - precond, err := ps.addPolicyPreCondition(ctx, pr) - if err != nil { - return err - } - - updates := []*v1.RelationshipUpdate{ - { - Operation: v1.RelationshipUpdate_OPERATION_CREATE, - Relationship: &v1.Relationship{ - Resource: &v1.ObjectReference{ObjectType: pr.ObjectType, ObjectId: pr.Object}, - Relation: pr.Relation, - Subject: &v1.SubjectReference{Object: &v1.ObjectReference{ObjectType: pr.SubjectType, ObjectId: pr.Subject}, OptionalRelation: pr.SubjectRelation}, - }, - }, - } - _, err = ps.permissionClient.WriteRelationships(ctx, &v1.WriteRelationshipsRequest{Updates: updates, OptionalPreconditions: precond}) - if err != nil { - return errors.Wrap(errAddPolicies, handleSpicedbError(err)) - } - - return nil -} - -func (ps *policyService) AddPolicies(ctx context.Context, prs []policies.Policy) error { - updates := []*v1.RelationshipUpdate{} - var preconds []*v1.Precondition - for _, pr := range prs { - if err := ps.policyValidation(pr); err != nil { - return errors.Wrap(svcerr.ErrInvalidPolicy, err) - } - precond, err := ps.addPolicyPreCondition(ctx, pr) - if err != nil { - return err - } - preconds = append(preconds, precond...) - updates = append(updates, &v1.RelationshipUpdate{ - Operation: v1.RelationshipUpdate_OPERATION_CREATE, - Relationship: &v1.Relationship{ - Resource: &v1.ObjectReference{ObjectType: pr.ObjectType, ObjectId: pr.Object}, - Relation: pr.Relation, - Subject: &v1.SubjectReference{Object: &v1.ObjectReference{ObjectType: pr.SubjectType, ObjectId: pr.Subject}, OptionalRelation: pr.SubjectRelation}, - }, - }) - } - if len(updates) == 0 { - return errors.Wrap(errors.ErrMalformedEntity, errNoPolicies) - } - _, err := ps.permissionClient.WriteRelationships(ctx, &v1.WriteRelationshipsRequest{Updates: updates, OptionalPreconditions: preconds}) - if err != nil { - return errors.Wrap(errAddPolicies, handleSpicedbError(err)) - } - - return nil -} - -func (ps *policyService) DeletePolicyFilter(ctx context.Context, pr policies.Policy) error { - req := &v1.DeleteRelationshipsRequest{ - RelationshipFilter: &v1.RelationshipFilter{ - ResourceType: pr.ObjectType, - OptionalResourceId: pr.Object, - OptionalResourceIdPrefix: pr.ObjectPrefix, - }, - } - - if pr.Relation != "" { - req.RelationshipFilter.OptionalRelation = pr.Relation - } - - if pr.SubjectType != "" { - req.RelationshipFilter.OptionalSubjectFilter = &v1.SubjectFilter{ - SubjectType: pr.SubjectType, - } - if pr.Subject != "" { - req.RelationshipFilter.OptionalSubjectFilter.OptionalSubjectId = pr.Subject - } - if pr.SubjectRelation != "" { - req.RelationshipFilter.OptionalSubjectFilter.OptionalRelation = &v1.SubjectFilter_RelationFilter{ - Relation: pr.SubjectRelation, - } - } - } - - if _, err := ps.permissionClient.DeleteRelationships(ctx, req); err != nil { - return errors.Wrap(errRemovePolicies, handleSpicedbError(err)) - } - - return nil -} - -func (ps *policyService) DeletePolicies(ctx context.Context, prs []policies.Policy) error { - updates := []*v1.RelationshipUpdate{} - for _, pr := range prs { - if err := ps.policyValidation(pr); err != nil { - return errors.Wrap(svcerr.ErrInvalidPolicy, err) - } - updates = append(updates, &v1.RelationshipUpdate{ - Operation: v1.RelationshipUpdate_OPERATION_DELETE, - Relationship: &v1.Relationship{ - Resource: &v1.ObjectReference{ObjectType: pr.ObjectType, ObjectId: pr.Object}, - Relation: pr.Relation, - Subject: &v1.SubjectReference{Object: &v1.ObjectReference{ObjectType: pr.SubjectType, ObjectId: pr.Subject}, OptionalRelation: pr.SubjectRelation}, - }, - }) - } - if len(updates) == 0 { - return errors.Wrap(errors.ErrMalformedEntity, errNoPolicies) - } - _, err := ps.permissionClient.WriteRelationships(ctx, &v1.WriteRelationshipsRequest{Updates: updates}) - if err != nil { - return errors.Wrap(errRemovePolicies, handleSpicedbError(err)) - } - - return nil -} - -func (ps *policyService) ListObjects(ctx context.Context, pr policies.Policy, nextPageToken string, limit uint64) (policies.PolicyPage, error) { - if limit <= 0 { - limit = 100 - } - res, npt, err := ps.retrieveObjects(ctx, pr, nextPageToken, limit) - if err != nil { - return policies.PolicyPage{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - var page policies.PolicyPage - for _, tuple := range res { - page.Policies = append(page.Policies, tuple.Object) - } - page.NextPageToken = npt - - return page, nil -} - -func (ps *policyService) ListAllObjects(ctx context.Context, pr policies.Policy) (policies.PolicyPage, error) { - res, err := ps.retrieveAllObjects(ctx, pr) - if err != nil { - return policies.PolicyPage{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - var page policies.PolicyPage - for _, tuple := range res { - page.Policies = append(page.Policies, tuple.Object) - } - - return page, nil -} - -func (ps *policyService) CountObjects(ctx context.Context, pr policies.Policy) (uint64, error) { - var count uint64 - nextPageToken := "" - for { - relationTuples, npt, err := ps.retrieveObjects(ctx, pr, nextPageToken, defRetrieveAllLimit) - if err != nil { - return count, err - } - count = count + uint64(len(relationTuples)) - if npt == "" { - break - } - nextPageToken = npt - } - - return count, nil -} - -func (ps *policyService) ListSubjects(ctx context.Context, pr policies.Policy, nextPageToken string, limit uint64) (policies.PolicyPage, error) { - if limit <= 0 { - limit = 100 - } - res, npt, err := ps.retrieveSubjects(ctx, pr, nextPageToken, limit) - if err != nil { - return policies.PolicyPage{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - var page policies.PolicyPage - for _, tuple := range res { - page.Policies = append(page.Policies, tuple.Subject) - } - page.NextPageToken = npt - - return page, nil -} - -func (ps *policyService) ListAllSubjects(ctx context.Context, pr policies.Policy) (policies.PolicyPage, error) { - res, err := ps.retrieveAllSubjects(ctx, pr) - if err != nil { - return policies.PolicyPage{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - var page policies.PolicyPage - for _, tuple := range res { - page.Policies = append(page.Policies, tuple.Subject) - } - - return page, nil -} - -func (ps *policyService) CountSubjects(ctx context.Context, pr policies.Policy) (uint64, error) { - var count uint64 - nextPageToken := "" - for { - relationTuples, npt, err := ps.retrieveSubjects(ctx, pr, nextPageToken, defRetrieveAllLimit) - if err != nil { - return count, err - } - count = count + uint64(len(relationTuples)) - if npt == "" { - break - } - nextPageToken = npt - } - - return count, nil -} - -func (ps *policyService) ListPermissions(ctx context.Context, pr policies.Policy, permissionsFilter []string) (policies.Permissions, error) { - if len(permissionsFilter) == 0 { - switch pr.ObjectType { - case policies.ClientType: - permissionsFilter = defClientsFilterPermissions - case policies.GroupType: - permissionsFilter = defGroupsFilterPermissions - case policies.PlatformType: - permissionsFilter = defPlatformFilterPermissions - case policies.DomainType: - permissionsFilter = defDomainsFilterPermissions - default: - return nil, svcerr.ErrMalformedEntity - } - } - pers, err := ps.retrievePermissions(ctx, pr, permissionsFilter) - if err != nil { - return []string{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - - return pers, nil -} - -func (ps *policyService) policyValidation(pr policies.Policy) error { - if pr.ObjectType == policies.PlatformType && pr.Object != policies.MagistralaObject { - return errPlatform - } - - return nil -} - -func (ps *policyService) addPolicyPreCondition(ctx context.Context, pr policies.Policy) ([]*v1.Precondition, error) { - // Checks are required for following ( -> means adding) - // 1.) user -> group (both user groups and channels) - // 2.) user -> client - // 3.) group -> group (both for adding parent_group and channels) - // 4.) group (channel) -> client - // 5.) user -> domain - - switch { - // 1.) user -> group (both user groups and channels) - // Checks : - // - USER with ANY RELATION to DOMAIN - // - GROUP with DOMAIN RELATION to DOMAIN - case pr.SubjectType == policies.UserType && pr.ObjectType == policies.GroupType: - return ps.userGroupPreConditions(ctx, pr) - - // 2.) user -> client - // Checks : - // - USER with ANY RELATION to DOMAIN - // - CLIENT with DOMAIN RELATION to DOMAIN - case pr.SubjectType == policies.UserType && pr.ObjectType == policies.ClientType: - return ps.userClientPreConditions(ctx, pr) - - // 3.) group -> group (both for adding parent_group and channels) - // Checks : - // - CHILD_GROUP with out PARENT_GROUP RELATION with any GROUP - case pr.SubjectType == policies.GroupType && pr.ObjectType == policies.GroupType: - return groupPreConditions(pr) - - // 4.) group (channel) -> client - // Checks : - // - GROUP (channel) with DOMAIN RELATION to DOMAIN - // - NO GROUP should not have PARENT_GROUP RELATION with GROUP (channel) - // - CLIENT with DOMAIN RELATION to DOMAIN - // case pr.SubjectType == policies.GroupType && pr.ObjectType == policies.ClientType: - // return channelClientPreCondition(pr) - - // 5.) user -> domain - // Checks : - // - User doesn't have any relation with domain - case pr.SubjectType == policies.UserType && pr.ObjectType == policies.DomainType: - return ps.userDomainPreConditions(ctx, pr) - - // Check client and group not belongs to other domain before adding to domain - case pr.SubjectType == policies.DomainType && pr.Relation == policies.DomainRelation && (pr.ObjectType == policies.ClientType || pr.ObjectType == policies.GroupType): - preconds := []*v1.Precondition{ - { - Operation: v1.Precondition_OPERATION_MUST_NOT_MATCH, - Filter: &v1.RelationshipFilter{ - ResourceType: pr.ObjectType, - OptionalResourceId: pr.Object, - OptionalRelation: policies.DomainRelation, - OptionalSubjectFilter: &v1.SubjectFilter{ - SubjectType: policies.DomainType, - }, - }, - }, - } - return preconds, nil - } - - return nil, nil -} - -func (ps *policyService) userGroupPreConditions(ctx context.Context, pr policies.Policy) ([]*v1.Precondition, error) { - var preconds []*v1.Precondition - - // user should not have any relation with group - preconds = append(preconds, &v1.Precondition{ - Operation: v1.Precondition_OPERATION_MUST_NOT_MATCH, - Filter: &v1.RelationshipFilter{ - ResourceType: policies.GroupType, - OptionalResourceId: pr.Object, - OptionalSubjectFilter: &v1.SubjectFilter{ - SubjectType: policies.UserType, - OptionalSubjectId: pr.Subject, - }, - }, - }) - isSuperAdmin := false - if err := ps.checkPolicy(ctx, policies.Policy{ - Subject: pr.Subject, - SubjectType: pr.SubjectType, - Permission: policies.AdminPermission, - Object: policies.MagistralaObject, - ObjectType: policies.PlatformType, - }); err == nil { - isSuperAdmin = true - } - - if !isSuperAdmin { - preconds = append(preconds, &v1.Precondition{ - Operation: v1.Precondition_OPERATION_MUST_MATCH, - Filter: &v1.RelationshipFilter{ - ResourceType: policies.DomainType, - OptionalResourceId: pr.Domain, - OptionalSubjectFilter: &v1.SubjectFilter{ - SubjectType: policies.UserType, - OptionalSubjectId: pr.Subject, - }, - }, - }) - } - switch { - case pr.ObjectKind == policies.NewGroupKind || pr.ObjectKind == policies.NewChannelKind: - preconds = append(preconds, - &v1.Precondition{ - Operation: v1.Precondition_OPERATION_MUST_NOT_MATCH, - Filter: &v1.RelationshipFilter{ - ResourceType: policies.GroupType, - OptionalResourceId: pr.Object, - OptionalRelation: policies.DomainRelation, - OptionalSubjectFilter: &v1.SubjectFilter{ - SubjectType: policies.DomainType, - }, - }, - }, - ) - default: - preconds = append(preconds, - &v1.Precondition{ - Operation: v1.Precondition_OPERATION_MUST_MATCH, - Filter: &v1.RelationshipFilter{ - ResourceType: policies.GroupType, - OptionalResourceId: pr.Object, - OptionalRelation: policies.DomainRelation, - OptionalSubjectFilter: &v1.SubjectFilter{ - SubjectType: policies.DomainType, - OptionalSubjectId: pr.Domain, - }, - }, - }, - ) - } - - return preconds, nil -} - -func (ps *policyService) userClientPreConditions(ctx context.Context, pr policies.Policy) ([]*v1.Precondition, error) { - var preconds []*v1.Precondition - - // user should not have any relation with client - preconds = append(preconds, &v1.Precondition{ - Operation: v1.Precondition_OPERATION_MUST_NOT_MATCH, - Filter: &v1.RelationshipFilter{ - ResourceType: policies.ClientType, - OptionalResourceId: pr.Object, - OptionalSubjectFilter: &v1.SubjectFilter{ - SubjectType: policies.UserType, - OptionalSubjectId: pr.Subject, - }, - }, - }) - - isSuperAdmin := false - if err := ps.checkPolicy(ctx, policies.Policy{ - Subject: pr.Subject, - SubjectType: pr.SubjectType, - Permission: policies.AdminPermission, - Object: policies.MagistralaObject, - ObjectType: policies.PlatformType, - }); err == nil { - isSuperAdmin = true - } - - if !isSuperAdmin { - preconds = append(preconds, &v1.Precondition{ - Operation: v1.Precondition_OPERATION_MUST_MATCH, - Filter: &v1.RelationshipFilter{ - ResourceType: policies.DomainType, - OptionalResourceId: pr.Domain, - OptionalSubjectFilter: &v1.SubjectFilter{ - SubjectType: policies.UserType, - OptionalSubjectId: pr.Subject, - }, - }, - }) - } - switch { - // For New client - // - CLIENT without DOMAIN RELATION to ANY DOMAIN - case pr.ObjectKind == policies.NewClientKind: - preconds = append(preconds, - &v1.Precondition{ - Operation: v1.Precondition_OPERATION_MUST_NOT_MATCH, - Filter: &v1.RelationshipFilter{ - ResourceType: policies.ClientType, - OptionalResourceId: pr.Object, - OptionalRelation: policies.DomainRelation, - OptionalSubjectFilter: &v1.SubjectFilter{ - SubjectType: policies.DomainType, - }, - }, - }, - ) - default: - // For existing client - // - CLIENT without DOMAIN RELATION to ANY DOMAIN - preconds = append(preconds, - &v1.Precondition{ - Operation: v1.Precondition_OPERATION_MUST_MATCH, - Filter: &v1.RelationshipFilter{ - ResourceType: policies.ClientType, - OptionalResourceId: pr.Object, - OptionalRelation: policies.DomainRelation, - OptionalSubjectFilter: &v1.SubjectFilter{ - SubjectType: policies.DomainType, - OptionalSubjectId: pr.Domain, - }, - }, - }, - ) - } - - return preconds, nil -} - -func (ps *policyService) userDomainPreConditions(ctx context.Context, pr policies.Policy) ([]*v1.Precondition, error) { - var preconds []*v1.Precondition - - if err := ps.checkPolicy(ctx, policies.Policy{ - Subject: pr.Subject, - SubjectType: pr.SubjectType, - Permission: policies.AdminPermission, - Object: policies.MagistralaObject, - ObjectType: policies.PlatformType, - }); err == nil { - return preconds, fmt.Errorf("use already exists in domain") - } - - // user should not have any relation with domain. - preconds = append(preconds, &v1.Precondition{ - Operation: v1.Precondition_OPERATION_MUST_NOT_MATCH, - Filter: &v1.RelationshipFilter{ - ResourceType: policies.DomainType, - OptionalResourceId: pr.Object, - OptionalSubjectFilter: &v1.SubjectFilter{ - SubjectType: policies.UserType, - OptionalSubjectId: pr.Subject, - }, - }, - }) - - return preconds, nil -} - -func (ps *policyService) checkPolicy(ctx context.Context, pr policies.Policy) error { - checkReq := v1.CheckPermissionRequest{ - // FullyConsistent means little caching will be available, which means performance will suffer. - // Only use if a ZedToken is not available or absolutely latest information is required. - // If we want to avoid FullyConsistent and to improve the performance of spicedb, then we need to cache the ZEDTOKEN whenever RELATIONS is created or updated. - // Instead of using FullyConsistent we need to use Consistency_AtLeastAsFresh, code looks like below one. - // Consistency: &v1.Consistency{ - // Requirement: &v1.Consistency_AtLeastAsFresh{ - // AtLeastAsFresh: getRelationTupleZedTokenFromCache() , - // } - // }, - // Reference: https://authzed.com/docs/reference/api-consistency - Consistency: &v1.Consistency{ - Requirement: &v1.Consistency_FullyConsistent{ - FullyConsistent: true, - }, - }, - Resource: &v1.ObjectReference{ObjectType: pr.ObjectType, ObjectId: pr.Object}, - Permission: pr.Permission, - Subject: &v1.SubjectReference{Object: &v1.ObjectReference{ObjectType: pr.SubjectType, ObjectId: pr.Subject}, OptionalRelation: pr.SubjectRelation}, - } - - resp, err := ps.permissionClient.CheckPermission(ctx, &checkReq) - if err != nil { - return handleSpicedbError(err) - } - if resp.Permissionship == v1.CheckPermissionResponse_PERMISSIONSHIP_HAS_PERMISSION { - return nil - } - if reason, ok := v1.CheckPermissionResponse_Permissionship_name[int32(resp.Permissionship)]; ok { - return errors.Wrap(svcerr.ErrAuthorization, errors.New(reason)) - } - return svcerr.ErrAuthorization -} - -func (ps *policyService) retrieveObjects(ctx context.Context, pr policies.Policy, nextPageToken string, limit uint64) ([]policies.Policy, string, error) { - resourceReq := &v1.LookupResourcesRequest{ - Consistency: &v1.Consistency{ - Requirement: &v1.Consistency_FullyConsistent{ - FullyConsistent: true, - }, - }, - ResourceObjectType: pr.ObjectType, - Permission: pr.Permission, - Subject: &v1.SubjectReference{Object: &v1.ObjectReference{ObjectType: pr.SubjectType, ObjectId: pr.Subject}, OptionalRelation: pr.SubjectRelation}, - OptionalLimit: uint32(limit), - } - if nextPageToken != "" { - resourceReq.OptionalCursor = &v1.Cursor{Token: nextPageToken} - } - stream, err := ps.permissionClient.LookupResources(ctx, resourceReq) - if err != nil { - return nil, "", errors.Wrap(errRetrievePolicies, handleSpicedbError(err)) - } - resources := []*v1.LookupResourcesResponse{} - var token string - for { - resp, err := stream.Recv() - switch err { - case nil: - resources = append(resources, resp) - case io.EOF: - if len(resources) > 0 && resources[len(resources)-1].AfterResultCursor != nil { - token = resources[len(resources)-1].AfterResultCursor.Token - } - return objectsToAuthPolicies(resources), token, nil - default: - if len(resources) > 0 && resources[len(resources)-1].AfterResultCursor != nil { - token = resources[len(resources)-1].AfterResultCursor.Token - } - return []policies.Policy{}, token, errors.Wrap(errRetrievePolicies, handleSpicedbError(err)) - } - } -} - -func (ps *policyService) retrieveAllObjects(ctx context.Context, pr policies.Policy) ([]policies.Policy, error) { - resourceReq := &v1.LookupResourcesRequest{ - Consistency: &v1.Consistency{ - Requirement: &v1.Consistency_FullyConsistent{ - FullyConsistent: true, - }, - }, - ResourceObjectType: pr.ObjectType, - Permission: pr.Permission, - Subject: &v1.SubjectReference{Object: &v1.ObjectReference{ObjectType: pr.SubjectType, ObjectId: pr.Subject}, OptionalRelation: pr.SubjectRelation}, - } - stream, err := ps.permissionClient.LookupResources(ctx, resourceReq) - if err != nil { - return nil, errors.Wrap(errRetrievePolicies, handleSpicedbError(err)) - } - tuples := []policies.Policy{} - for { - resp, err := stream.Recv() - switch { - case errors.Contains(err, io.EOF): - return tuples, nil - case err != nil: - return tuples, errors.Wrap(errRetrievePolicies, handleSpicedbError(err)) - default: - tuples = append(tuples, policies.Policy{Object: resp.ResourceObjectId}) - } - } -} - -func (ps *policyService) retrieveSubjects(ctx context.Context, pr policies.Policy, nextPageToken string, limit uint64) ([]policies.Policy, string, error) { - subjectsReq := v1.LookupSubjectsRequest{ - Consistency: &v1.Consistency{ - Requirement: &v1.Consistency_FullyConsistent{ - FullyConsistent: true, - }, - }, - Resource: &v1.ObjectReference{ObjectType: pr.ObjectType, ObjectId: pr.Object}, - Permission: pr.Permission, - SubjectObjectType: pr.SubjectType, - OptionalSubjectRelation: pr.SubjectRelation, - OptionalConcreteLimit: uint32(limit), - WildcardOption: v1.LookupSubjectsRequest_WILDCARD_OPTION_INCLUDE_WILDCARDS, - } - if nextPageToken != "" { - subjectsReq.OptionalCursor = &v1.Cursor{Token: nextPageToken} - } - stream, err := ps.permissionClient.LookupSubjects(ctx, &subjectsReq) - if err != nil { - return nil, "", errors.Wrap(errRetrievePolicies, handleSpicedbError(err)) - } - subjects := []*v1.LookupSubjectsResponse{} - var token string - for { - resp, err := stream.Recv() - - switch err { - case nil: - subjects = append(subjects, resp) - case io.EOF: - if len(subjects) > 0 && subjects[len(subjects)-1].AfterResultCursor != nil { - token = subjects[len(subjects)-1].AfterResultCursor.Token - } - return subjectsToAuthPolicies(subjects), token, nil - default: - if len(subjects) > 0 && subjects[len(subjects)-1].AfterResultCursor != nil { - token = subjects[len(subjects)-1].AfterResultCursor.Token - } - return []policies.Policy{}, token, errors.Wrap(errRetrievePolicies, handleSpicedbError(err)) - } - } -} - -func (ps *policyService) retrieveAllSubjects(ctx context.Context, pr policies.Policy) ([]policies.Policy, error) { - var tuples []policies.Policy - nextPageToken := "" - for i := 0; ; i++ { - relationTuples, npt, err := ps.retrieveSubjects(ctx, pr, nextPageToken, defRetrieveAllLimit) - if err != nil { - return tuples, err - } - tuples = append(tuples, relationTuples...) - if npt == "" || (len(tuples) < defRetrieveAllLimit) { - break - } - nextPageToken = npt - } - return tuples, nil -} - -func (ps *policyService) retrievePermissions(ctx context.Context, pr policies.Policy, filterPermission []string) (policies.Permissions, error) { - var permissionChecks []*v1.CheckBulkPermissionsRequestItem - for _, fp := range filterPermission { - permissionChecks = append(permissionChecks, &v1.CheckBulkPermissionsRequestItem{ - Resource: &v1.ObjectReference{ - ObjectType: pr.ObjectType, - ObjectId: pr.Object, - }, - Permission: fp, - Subject: &v1.SubjectReference{ - Object: &v1.ObjectReference{ - ObjectType: pr.SubjectType, - ObjectId: pr.Subject, - }, - OptionalRelation: pr.SubjectRelation, - }, - }) - } - resp, err := ps.client.PermissionsServiceClient.CheckBulkPermissions(ctx, &v1.CheckBulkPermissionsRequest{ - Consistency: &v1.Consistency{ - Requirement: &v1.Consistency_FullyConsistent{ - FullyConsistent: true, - }, - }, - Items: permissionChecks, - }) - if err != nil { - return policies.Permissions{}, errors.Wrap(errRetrievePolicies, handleSpicedbError(err)) - } - - permissions := []string{} - for _, pair := range resp.Pairs { - if pair.GetError() != nil { - s := pair.GetError() - return policies.Permissions{}, errors.Wrap(errRetrievePolicies, convertGRPCStatusToError(convertToGrpcStatus(s))) - } - item := pair.GetItem() - req := pair.GetRequest() - if item != nil && req != nil && item.Permissionship == v1.CheckPermissionResponse_PERMISSIONSHIP_HAS_PERMISSION { - permissions = append(permissions, req.GetPermission()) - } - } - return permissions, nil -} - -func groupPreConditions(pr policies.Policy) ([]*v1.Precondition, error) { - // - PARENT_GROUP (subject) with DOMAIN RELATION to DOMAIN - precond := []*v1.Precondition{ - { - Operation: v1.Precondition_OPERATION_MUST_MATCH, - Filter: &v1.RelationshipFilter{ - ResourceType: policies.GroupType, - OptionalResourceId: pr.Subject, - OptionalRelation: policies.DomainRelation, - OptionalSubjectFilter: &v1.SubjectFilter{ - SubjectType: policies.DomainType, - OptionalSubjectId: pr.Domain, - }, - }, - }, - } - if pr.ObjectKind != policies.ChannelsKind { - precond = append(precond, - &v1.Precondition{ - Operation: v1.Precondition_OPERATION_MUST_NOT_MATCH, - Filter: &v1.RelationshipFilter{ - ResourceType: policies.GroupType, - OptionalResourceId: pr.Object, - OptionalRelation: policies.ParentGroupRelation, - OptionalSubjectFilter: &v1.SubjectFilter{ - SubjectType: policies.GroupType, - }, - }, - }, - ) - } - switch { - // - NEW CHILD_GROUP (object) with out DOMAIN RELATION to ANY DOMAIN - case pr.ObjectType == policies.GroupType && pr.ObjectKind == policies.NewGroupKind: - precond = append(precond, - &v1.Precondition{ - Operation: v1.Precondition_OPERATION_MUST_NOT_MATCH, - Filter: &v1.RelationshipFilter{ - ResourceType: policies.GroupType, - OptionalResourceId: pr.Object, - OptionalRelation: policies.DomainRelation, - OptionalSubjectFilter: &v1.SubjectFilter{ - SubjectType: policies.DomainType, - }, - }, - }, - ) - default: - // - CHILD_GROUP (object) with DOMAIN RELATION to DOMAIN - precond = append(precond, - &v1.Precondition{ - Operation: v1.Precondition_OPERATION_MUST_MATCH, - Filter: &v1.RelationshipFilter{ - ResourceType: policies.GroupType, - OptionalResourceId: pr.Object, - OptionalRelation: policies.DomainRelation, - OptionalSubjectFilter: &v1.SubjectFilter{ - SubjectType: policies.DomainType, - OptionalSubjectId: pr.Domain, - }, - }, - }, - ) - } - return precond, nil -} - -func objectsToAuthPolicies(objects []*v1.LookupResourcesResponse) []policies.Policy { - var policyList []policies.Policy - for _, obj := range objects { - policyList = append(policyList, policies.Policy{ - Object: obj.GetResourceObjectId(), - }) - } - return policyList -} - -func subjectsToAuthPolicies(subjects []*v1.LookupSubjectsResponse) []policies.Policy { - var policyList []policies.Policy - for _, sub := range subjects { - policyList = append(policyList, policies.Policy{ - Subject: sub.Subject.GetSubjectObjectId(), - }) - } - return policyList -} - -func handleSpicedbError(err error) error { - if st, ok := status.FromError(err); ok { - return convertGRPCStatusToError(st) - } - return err -} - -func convertToGrpcStatus(gst *gstatus.Status) *status.Status { - st := status.New(codes.Code(gst.Code), gst.GetMessage()) - return st -} - -func convertGRPCStatusToError(st *status.Status) error { - switch st.Code() { - case codes.NotFound: - return errors.Wrap(repoerr.ErrNotFound, errors.New(st.Message())) - case codes.InvalidArgument: - return errors.Wrap(errors.ErrMalformedEntity, errors.New(st.Message())) - case codes.AlreadyExists: - return errors.Wrap(repoerr.ErrConflict, errors.New(st.Message())) - case codes.Unauthenticated: - return errors.Wrap(svcerr.ErrAuthentication, errors.New(st.Message())) - case codes.Internal: - return errors.Wrap(errInternal, errors.New(st.Message())) - case codes.OK: - if msg := st.Message(); msg != "" { - return errors.Wrap(errors.ErrUnidentified, errors.New(msg)) - } - return nil - case codes.FailedPrecondition: - return errors.Wrap(errors.ErrMalformedEntity, errors.New(st.Message())) - case codes.PermissionDenied: - return errors.Wrap(svcerr.ErrAuthorization, errors.New(st.Message())) - default: - return errors.Wrap(fmt.Errorf("unexpected gRPC status: %s (status code:%v)", st.Code().String(), st.Code()), errors.New(st.Message())) - } -} diff --git a/pkg/re/events/consumer/decode.go b/pkg/re/events/consumer/decode.go index 0f4f15937..41066aa77 100644 --- a/pkg/re/events/consumer/decode.go +++ b/pkg/re/events/consumer/decode.go @@ -8,8 +8,6 @@ import ( "time" "github.com/absmach/magistrala/pkg/errors" - "github.com/absmach/magistrala/pkg/roles" - rconsumer "github.com/absmach/magistrala/pkg/roles/rolemanager/events/consumer" "github.com/absmach/magistrala/pkg/schedule" "github.com/absmach/magistrala/re" ) @@ -33,6 +31,7 @@ var ( errUpdatedAt = errors.New("failed to parse 'updated_at' time") errDecodeLogic = errors.New("failed to decode 'logic'") errDecodeSchedule = errors.New("failed to decode 'schedule'") + errStringValue = errors.New("invalid string value") ) // ToRule decodes a map[string]any event payload into a re.Rule. @@ -82,7 +81,7 @@ func ToRule(data map[string]any) (re.Rule, error) { } if itags, ok := data["tags"].([]any); ok { - tags, err := rconsumer.ToStrings(itags) + tags, err := toStrings(itags) if err != nil { return re.Rule{}, errors.Wrap(errTags, err) } @@ -138,21 +137,25 @@ func ToRule(data map[string]any) (re.Rule, error) { return r, nil } -func decodeAddRuleEvent(data map[string]any) (re.Rule, []roles.RoleProvision, error) { +func toStrings(values []any) ([]string, error) { + strings := make([]string, 0, len(values)) + for _, value := range values { + str, ok := value.(string) + if !ok { + return nil, errStringValue + } + strings = append(strings, str) + } + return strings, nil +} + +func decodeAddRuleEvent(data map[string]any) (re.Rule, error) { r, err := ToRule(data) if err != nil { - return re.Rule{}, nil, errors.Wrap(errDecodeAddRuleEvent, err) + return re.Rule{}, errors.Wrap(errDecodeAddRuleEvent, err) } - var rps []roles.RoleProvision - if irps, ok := data["roles_provisioned"].([]any); ok { - rps, err = rconsumer.ToRoleProvisions(irps) - if err != nil { - return re.Rule{}, nil, errors.Wrap(errDecodeAddRuleEvent, err) - } - } - - return r, rps, nil + return r, nil } func decodeUpdateRuleEvent(data map[string]any) (re.Rule, error) { diff --git a/pkg/re/events/consumer/stream.go b/pkg/re/events/consumer/stream.go index c0b011fa9..864320b89 100644 --- a/pkg/re/events/consumer/stream.go +++ b/pkg/re/events/consumer/stream.go @@ -10,7 +10,6 @@ import ( "github.com/absmach/magistrala/pkg/errors" "github.com/absmach/magistrala/pkg/events" "github.com/absmach/magistrala/pkg/events/store" - rconsumer "github.com/absmach/magistrala/pkg/roles/rolemanager/events/consumer" "github.com/absmach/magistrala/re" ) @@ -38,8 +37,7 @@ var ( ) type eventHandler struct { - repo re.Repository - rolesEventHandler rconsumer.EventHandler + repo re.Repository } func RulesEventsSubscribe(ctx context.Context, repo re.Repository, esURL, esConsumerName string, logger *slog.Logger) error { @@ -59,10 +57,8 @@ func RulesEventsSubscribe(ctx context.Context, repo re.Repository, esURL, esCons // NewEventHandler returns new event store handler. func NewEventHandler(repo re.Repository) events.EventHandler { - reh := rconsumer.NewEventHandler("rule", repo) return &eventHandler{ - repo: repo, - rolesEventHandler: reh, + repo: repo, } } @@ -94,11 +90,11 @@ func (es *eventHandler) Handle(ctx context.Context, event events.Event) error { return es.removeRuleHandler(ctx, msg) } - return es.rolesEventHandler.Handle(ctx, op, msg) + return nil } func (es *eventHandler) addRuleHandler(ctx context.Context, data map[string]any) error { - r, rps, err := decodeAddRuleEvent(data) + r, err := decodeAddRuleEvent(data) if err != nil { return errors.Wrap(errAddRuleEvent, err) } @@ -107,10 +103,6 @@ func (es *eventHandler) addRuleHandler(ctx context.Context, data map[string]any) return errors.Wrap(errAddRuleEvent, err) } - if _, err := es.repo.AddRoles(ctx, rps); err != nil { - return errors.Wrap(errAddRuleEvent, err) - } - return nil } diff --git a/pkg/roles/mocks/provisioner.go b/pkg/roles/mocks/provisioner.go deleted file mode 100644 index 6e9c989a2..000000000 --- a/pkg/roles/mocks/provisioner.go +++ /dev/null @@ -1,217 +0,0 @@ -// Copyright (c) Abstract Machines - -// SPDX-License-Identifier: Apache-2.0 - -// Code generated by mockery; DO NOT EDIT. -// github.com/vektra/mockery -// template: testify - -package mocks - -import ( - "context" - - "github.com/absmach/magistrala/pkg/policies" - "github.com/absmach/magistrala/pkg/roles" - mock "github.com/stretchr/testify/mock" -) - -// NewProvisioner creates a new instance of Provisioner. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. -// The first argument is typically a *testing.T value. -func NewProvisioner(t interface { - mock.TestingT - Cleanup(func()) -}) *Provisioner { - mock := &Provisioner{} - mock.Mock.Test(t) - - t.Cleanup(func() { mock.AssertExpectations(t) }) - - return mock -} - -// Provisioner is an autogenerated mock type for the Provisioner type -type Provisioner struct { - mock.Mock -} - -type Provisioner_Expecter struct { - mock *mock.Mock -} - -func (_m *Provisioner) EXPECT() *Provisioner_Expecter { - return &Provisioner_Expecter{mock: &_m.Mock} -} - -// AddNewEntitiesRoles provides a mock function for the type Provisioner -func (_mock *Provisioner) AddNewEntitiesRoles(ctx context.Context, domainID string, userID string, entityIDs []string, optionalEntityPolicies []policies.Policy, newBuiltInRoleMembers map[roles.BuiltInRoleName][]roles.Member) ([]roles.RoleProvision, error) { - ret := _mock.Called(ctx, domainID, userID, entityIDs, optionalEntityPolicies, newBuiltInRoleMembers) - - if len(ret) == 0 { - panic("no return value specified for AddNewEntitiesRoles") - } - - var r0 []roles.RoleProvision - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string, []string, []policies.Policy, map[roles.BuiltInRoleName][]roles.Member) ([]roles.RoleProvision, error)); ok { - return returnFunc(ctx, domainID, userID, entityIDs, optionalEntityPolicies, newBuiltInRoleMembers) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string, []string, []policies.Policy, map[roles.BuiltInRoleName][]roles.Member) []roles.RoleProvision); ok { - r0 = returnFunc(ctx, domainID, userID, entityIDs, optionalEntityPolicies, newBuiltInRoleMembers) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]roles.RoleProvision) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, string, []string, []policies.Policy, map[roles.BuiltInRoleName][]roles.Member) error); ok { - r1 = returnFunc(ctx, domainID, userID, entityIDs, optionalEntityPolicies, newBuiltInRoleMembers) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Provisioner_AddNewEntitiesRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'AddNewEntitiesRoles' -type Provisioner_AddNewEntitiesRoles_Call struct { - *mock.Call -} - -// AddNewEntitiesRoles is a helper method to define mock.On call -// - ctx context.Context -// - domainID string -// - userID string -// - entityIDs []string -// - optionalEntityPolicies []policies.Policy -// - newBuiltInRoleMembers map[roles.BuiltInRoleName][]roles.Member -func (_e *Provisioner_Expecter) AddNewEntitiesRoles(ctx interface{}, domainID interface{}, userID interface{}, entityIDs interface{}, optionalEntityPolicies interface{}, newBuiltInRoleMembers interface{}) *Provisioner_AddNewEntitiesRoles_Call { - return &Provisioner_AddNewEntitiesRoles_Call{Call: _e.mock.On("AddNewEntitiesRoles", ctx, domainID, userID, entityIDs, optionalEntityPolicies, newBuiltInRoleMembers)} -} - -func (_c *Provisioner_AddNewEntitiesRoles_Call) Run(run func(ctx context.Context, domainID string, userID string, entityIDs []string, optionalEntityPolicies []policies.Policy, newBuiltInRoleMembers map[roles.BuiltInRoleName][]roles.Member)) *Provisioner_AddNewEntitiesRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 []string - if args[3] != nil { - arg3 = args[3].([]string) - } - var arg4 []policies.Policy - if args[4] != nil { - arg4 = args[4].([]policies.Policy) - } - var arg5 map[roles.BuiltInRoleName][]roles.Member - if args[5] != nil { - arg5 = args[5].(map[roles.BuiltInRoleName][]roles.Member) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - arg5, - ) - }) - return _c -} - -func (_c *Provisioner_AddNewEntitiesRoles_Call) Return(roleProvisions []roles.RoleProvision, err error) *Provisioner_AddNewEntitiesRoles_Call { - _c.Call.Return(roleProvisions, err) - return _c -} - -func (_c *Provisioner_AddNewEntitiesRoles_Call) RunAndReturn(run func(ctx context.Context, domainID string, userID string, entityIDs []string, optionalEntityPolicies []policies.Policy, newBuiltInRoleMembers map[roles.BuiltInRoleName][]roles.Member) ([]roles.RoleProvision, error)) *Provisioner_AddNewEntitiesRoles_Call { - _c.Call.Return(run) - return _c -} - -// RemoveEntitiesRoles provides a mock function for the type Provisioner -func (_mock *Provisioner) RemoveEntitiesRoles(ctx context.Context, domainID string, userID string, entityIDs []string, optionalFilterDeletePolicies []policies.Policy, optionalDeletePolicies []policies.Policy) error { - ret := _mock.Called(ctx, domainID, userID, entityIDs, optionalFilterDeletePolicies, optionalDeletePolicies) - - if len(ret) == 0 { - panic("no return value specified for RemoveEntitiesRoles") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string, []string, []policies.Policy, []policies.Policy) error); ok { - r0 = returnFunc(ctx, domainID, userID, entityIDs, optionalFilterDeletePolicies, optionalDeletePolicies) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Provisioner_RemoveEntitiesRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveEntitiesRoles' -type Provisioner_RemoveEntitiesRoles_Call struct { - *mock.Call -} - -// RemoveEntitiesRoles is a helper method to define mock.On call -// - ctx context.Context -// - domainID string -// - userID string -// - entityIDs []string -// - optionalFilterDeletePolicies []policies.Policy -// - optionalDeletePolicies []policies.Policy -func (_e *Provisioner_Expecter) RemoveEntitiesRoles(ctx interface{}, domainID interface{}, userID interface{}, entityIDs interface{}, optionalFilterDeletePolicies interface{}, optionalDeletePolicies interface{}) *Provisioner_RemoveEntitiesRoles_Call { - return &Provisioner_RemoveEntitiesRoles_Call{Call: _e.mock.On("RemoveEntitiesRoles", ctx, domainID, userID, entityIDs, optionalFilterDeletePolicies, optionalDeletePolicies)} -} - -func (_c *Provisioner_RemoveEntitiesRoles_Call) Run(run func(ctx context.Context, domainID string, userID string, entityIDs []string, optionalFilterDeletePolicies []policies.Policy, optionalDeletePolicies []policies.Policy)) *Provisioner_RemoveEntitiesRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 []string - if args[3] != nil { - arg3 = args[3].([]string) - } - var arg4 []policies.Policy - if args[4] != nil { - arg4 = args[4].([]policies.Policy) - } - var arg5 []policies.Policy - if args[5] != nil { - arg5 = args[5].([]policies.Policy) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - arg5, - ) - }) - return _c -} - -func (_c *Provisioner_RemoveEntitiesRoles_Call) Return(err error) *Provisioner_RemoveEntitiesRoles_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Provisioner_RemoveEntitiesRoles_Call) RunAndReturn(run func(ctx context.Context, domainID string, userID string, entityIDs []string, optionalFilterDeletePolicies []policies.Policy, optionalDeletePolicies []policies.Policy) error) *Provisioner_RemoveEntitiesRoles_Call { - _c.Call.Return(run) - return _c -} diff --git a/pkg/roles/mocks/repository.go b/pkg/roles/mocks/repository.go deleted file mode 100644 index a76e48a3f..000000000 --- a/pkg/roles/mocks/repository.go +++ /dev/null @@ -1,1396 +0,0 @@ -// Copyright (c) Abstract Machines - -// SPDX-License-Identifier: Apache-2.0 - -// Code generated by mockery; DO NOT EDIT. -// github.com/vektra/mockery -// template: testify - -package mocks - -import ( - "context" - - "github.com/absmach/magistrala/pkg/roles" - mock "github.com/stretchr/testify/mock" -) - -// NewRepository creates a new instance of Repository. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. -// The first argument is typically a *testing.T value. -func NewRepository(t interface { - mock.TestingT - Cleanup(func()) -}) *Repository { - mock := &Repository{} - mock.Mock.Test(t) - - t.Cleanup(func() { mock.AssertExpectations(t) }) - - return mock -} - -// Repository is an autogenerated mock type for the Repository type -type Repository struct { - mock.Mock -} - -type Repository_Expecter struct { - mock *mock.Mock -} - -func (_m *Repository) EXPECT() *Repository_Expecter { - return &Repository_Expecter{mock: &_m.Mock} -} - -// AddRoles provides a mock function for the type Repository -func (_mock *Repository) AddRoles(ctx context.Context, rps []roles.RoleProvision) ([]roles.RoleProvision, error) { - ret := _mock.Called(ctx, rps) - - if len(ret) == 0 { - panic("no return value specified for AddRoles") - } - - var r0 []roles.RoleProvision - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, []roles.RoleProvision) ([]roles.RoleProvision, error)); ok { - return returnFunc(ctx, rps) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, []roles.RoleProvision) []roles.RoleProvision); ok { - r0 = returnFunc(ctx, rps) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]roles.RoleProvision) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, []roles.RoleProvision) error); ok { - r1 = returnFunc(ctx, rps) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_AddRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'AddRoles' -type Repository_AddRoles_Call struct { - *mock.Call -} - -// AddRoles is a helper method to define mock.On call -// - ctx context.Context -// - rps []roles.RoleProvision -func (_e *Repository_Expecter) AddRoles(ctx interface{}, rps interface{}) *Repository_AddRoles_Call { - return &Repository_AddRoles_Call{Call: _e.mock.On("AddRoles", ctx, rps)} -} - -func (_c *Repository_AddRoles_Call) Run(run func(ctx context.Context, rps []roles.RoleProvision)) *Repository_AddRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 []roles.RoleProvision - if args[1] != nil { - arg1 = args[1].([]roles.RoleProvision) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_AddRoles_Call) Return(roleProvisions []roles.RoleProvision, err error) *Repository_AddRoles_Call { - _c.Call.Return(roleProvisions, err) - return _c -} - -func (_c *Repository_AddRoles_Call) RunAndReturn(run func(ctx context.Context, rps []roles.RoleProvision) ([]roles.RoleProvision, error)) *Repository_AddRoles_Call { - _c.Call.Return(run) - return _c -} - -// ListEntityMembers provides a mock function for the type Repository -func (_mock *Repository) ListEntityMembers(ctx context.Context, entityID string, pageQuery roles.MembersRolePageQuery) (roles.MembersRolePage, error) { - ret := _mock.Called(ctx, entityID, pageQuery) - - if len(ret) == 0 { - panic("no return value specified for ListEntityMembers") - } - - var r0 roles.MembersRolePage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, roles.MembersRolePageQuery) (roles.MembersRolePage, error)); ok { - return returnFunc(ctx, entityID, pageQuery) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, roles.MembersRolePageQuery) roles.MembersRolePage); ok { - r0 = returnFunc(ctx, entityID, pageQuery) - } else { - r0 = ret.Get(0).(roles.MembersRolePage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, roles.MembersRolePageQuery) error); ok { - r1 = returnFunc(ctx, entityID, pageQuery) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_ListEntityMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ListEntityMembers' -type Repository_ListEntityMembers_Call struct { - *mock.Call -} - -// ListEntityMembers is a helper method to define mock.On call -// - ctx context.Context -// - entityID string -// - pageQuery roles.MembersRolePageQuery -func (_e *Repository_Expecter) ListEntityMembers(ctx interface{}, entityID interface{}, pageQuery interface{}) *Repository_ListEntityMembers_Call { - return &Repository_ListEntityMembers_Call{Call: _e.mock.On("ListEntityMembers", ctx, entityID, pageQuery)} -} - -func (_c *Repository_ListEntityMembers_Call) Run(run func(ctx context.Context, entityID string, pageQuery roles.MembersRolePageQuery)) *Repository_ListEntityMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 roles.MembersRolePageQuery - if args[2] != nil { - arg2 = args[2].(roles.MembersRolePageQuery) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_ListEntityMembers_Call) Return(membersRolePage roles.MembersRolePage, err error) *Repository_ListEntityMembers_Call { - _c.Call.Return(membersRolePage, err) - return _c -} - -func (_c *Repository_ListEntityMembers_Call) RunAndReturn(run func(ctx context.Context, entityID string, pageQuery roles.MembersRolePageQuery) (roles.MembersRolePage, error)) *Repository_ListEntityMembers_Call { - _c.Call.Return(run) - return _c -} - -// RemoveEntityMembers provides a mock function for the type Repository -func (_mock *Repository) RemoveEntityMembers(ctx context.Context, entityID string, members []string) error { - ret := _mock.Called(ctx, entityID, members) - - if len(ret) == 0 { - panic("no return value specified for RemoveEntityMembers") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, []string) error); ok { - r0 = returnFunc(ctx, entityID, members) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RemoveEntityMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveEntityMembers' -type Repository_RemoveEntityMembers_Call struct { - *mock.Call -} - -// RemoveEntityMembers is a helper method to define mock.On call -// - ctx context.Context -// - entityID string -// - members []string -func (_e *Repository_Expecter) RemoveEntityMembers(ctx interface{}, entityID interface{}, members interface{}) *Repository_RemoveEntityMembers_Call { - return &Repository_RemoveEntityMembers_Call{Call: _e.mock.On("RemoveEntityMembers", ctx, entityID, members)} -} - -func (_c *Repository_RemoveEntityMembers_Call) Run(run func(ctx context.Context, entityID string, members []string)) *Repository_RemoveEntityMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RemoveEntityMembers_Call) Return(err error) *Repository_RemoveEntityMembers_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RemoveEntityMembers_Call) RunAndReturn(run func(ctx context.Context, entityID string, members []string) error) *Repository_RemoveEntityMembers_Call { - _c.Call.Return(run) - return _c -} - -// RemoveMemberFromAllRoles provides a mock function for the type Repository -func (_mock *Repository) RemoveMemberFromAllRoles(ctx context.Context, memberID string) error { - ret := _mock.Called(ctx, memberID) - - if len(ret) == 0 { - panic("no return value specified for RemoveMemberFromAllRoles") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) error); ok { - r0 = returnFunc(ctx, memberID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RemoveMemberFromAllRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveMemberFromAllRoles' -type Repository_RemoveMemberFromAllRoles_Call struct { - *mock.Call -} - -// RemoveMemberFromAllRoles is a helper method to define mock.On call -// - ctx context.Context -// - memberID string -func (_e *Repository_Expecter) RemoveMemberFromAllRoles(ctx interface{}, memberID interface{}) *Repository_RemoveMemberFromAllRoles_Call { - return &Repository_RemoveMemberFromAllRoles_Call{Call: _e.mock.On("RemoveMemberFromAllRoles", ctx, memberID)} -} - -func (_c *Repository_RemoveMemberFromAllRoles_Call) Run(run func(ctx context.Context, memberID string)) *Repository_RemoveMemberFromAllRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RemoveMemberFromAllRoles_Call) Return(err error) *Repository_RemoveMemberFromAllRoles_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RemoveMemberFromAllRoles_Call) RunAndReturn(run func(ctx context.Context, memberID string) error) *Repository_RemoveMemberFromAllRoles_Call { - _c.Call.Return(run) - return _c -} - -// RemoveRoles provides a mock function for the type Repository -func (_mock *Repository) RemoveRoles(ctx context.Context, roleIDs []string) error { - ret := _mock.Called(ctx, roleIDs) - - if len(ret) == 0 { - panic("no return value specified for RemoveRoles") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, []string) error); ok { - r0 = returnFunc(ctx, roleIDs) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RemoveRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveRoles' -type Repository_RemoveRoles_Call struct { - *mock.Call -} - -// RemoveRoles is a helper method to define mock.On call -// - ctx context.Context -// - roleIDs []string -func (_e *Repository_Expecter) RemoveRoles(ctx interface{}, roleIDs interface{}) *Repository_RemoveRoles_Call { - return &Repository_RemoveRoles_Call{Call: _e.mock.On("RemoveRoles", ctx, roleIDs)} -} - -func (_c *Repository_RemoveRoles_Call) Run(run func(ctx context.Context, roleIDs []string)) *Repository_RemoveRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 []string - if args[1] != nil { - arg1 = args[1].([]string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RemoveRoles_Call) Return(err error) *Repository_RemoveRoles_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RemoveRoles_Call) RunAndReturn(run func(ctx context.Context, roleIDs []string) error) *Repository_RemoveRoles_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveAllRoles provides a mock function for the type Repository -func (_mock *Repository) RetrieveAllRoles(ctx context.Context, entityID string, limit uint64, offset uint64) (roles.RolePage, error) { - ret := _mock.Called(ctx, entityID, limit, offset) - - if len(ret) == 0 { - panic("no return value specified for RetrieveAllRoles") - } - - var r0 roles.RolePage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, uint64, uint64) (roles.RolePage, error)); ok { - return returnFunc(ctx, entityID, limit, offset) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, uint64, uint64) roles.RolePage); ok { - r0 = returnFunc(ctx, entityID, limit, offset) - } else { - r0 = ret.Get(0).(roles.RolePage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, uint64, uint64) error); ok { - r1 = returnFunc(ctx, entityID, limit, offset) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveAllRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveAllRoles' -type Repository_RetrieveAllRoles_Call struct { - *mock.Call -} - -// RetrieveAllRoles is a helper method to define mock.On call -// - ctx context.Context -// - entityID string -// - limit uint64 -// - offset uint64 -func (_e *Repository_Expecter) RetrieveAllRoles(ctx interface{}, entityID interface{}, limit interface{}, offset interface{}) *Repository_RetrieveAllRoles_Call { - return &Repository_RetrieveAllRoles_Call{Call: _e.mock.On("RetrieveAllRoles", ctx, entityID, limit, offset)} -} - -func (_c *Repository_RetrieveAllRoles_Call) Run(run func(ctx context.Context, entityID string, limit uint64, offset uint64)) *Repository_RetrieveAllRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 uint64 - if args[2] != nil { - arg2 = args[2].(uint64) - } - var arg3 uint64 - if args[3] != nil { - arg3 = args[3].(uint64) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Repository_RetrieveAllRoles_Call) Return(rolePage roles.RolePage, err error) *Repository_RetrieveAllRoles_Call { - _c.Call.Return(rolePage, err) - return _c -} - -func (_c *Repository_RetrieveAllRoles_Call) RunAndReturn(run func(ctx context.Context, entityID string, limit uint64, offset uint64) (roles.RolePage, error)) *Repository_RetrieveAllRoles_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveEntitiesRolesActionsMembers provides a mock function for the type Repository -func (_mock *Repository) RetrieveEntitiesRolesActionsMembers(ctx context.Context, entityIDs []string) ([]roles.EntityActionRole, []roles.EntityMemberRole, error) { - ret := _mock.Called(ctx, entityIDs) - - if len(ret) == 0 { - panic("no return value specified for RetrieveEntitiesRolesActionsMembers") - } - - var r0 []roles.EntityActionRole - var r1 []roles.EntityMemberRole - var r2 error - if returnFunc, ok := ret.Get(0).(func(context.Context, []string) ([]roles.EntityActionRole, []roles.EntityMemberRole, error)); ok { - return returnFunc(ctx, entityIDs) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, []string) []roles.EntityActionRole); ok { - r0 = returnFunc(ctx, entityIDs) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]roles.EntityActionRole) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, []string) []roles.EntityMemberRole); ok { - r1 = returnFunc(ctx, entityIDs) - } else { - if ret.Get(1) != nil { - r1 = ret.Get(1).([]roles.EntityMemberRole) - } - } - if returnFunc, ok := ret.Get(2).(func(context.Context, []string) error); ok { - r2 = returnFunc(ctx, entityIDs) - } else { - r2 = ret.Error(2) - } - return r0, r1, r2 -} - -// Repository_RetrieveEntitiesRolesActionsMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveEntitiesRolesActionsMembers' -type Repository_RetrieveEntitiesRolesActionsMembers_Call struct { - *mock.Call -} - -// RetrieveEntitiesRolesActionsMembers is a helper method to define mock.On call -// - ctx context.Context -// - entityIDs []string -func (_e *Repository_Expecter) RetrieveEntitiesRolesActionsMembers(ctx interface{}, entityIDs interface{}) *Repository_RetrieveEntitiesRolesActionsMembers_Call { - return &Repository_RetrieveEntitiesRolesActionsMembers_Call{Call: _e.mock.On("RetrieveEntitiesRolesActionsMembers", ctx, entityIDs)} -} - -func (_c *Repository_RetrieveEntitiesRolesActionsMembers_Call) Run(run func(ctx context.Context, entityIDs []string)) *Repository_RetrieveEntitiesRolesActionsMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 []string - if args[1] != nil { - arg1 = args[1].([]string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RetrieveEntitiesRolesActionsMembers_Call) Return(entityActionRoles []roles.EntityActionRole, entityMemberRoles []roles.EntityMemberRole, err error) *Repository_RetrieveEntitiesRolesActionsMembers_Call { - _c.Call.Return(entityActionRoles, entityMemberRoles, err) - return _c -} - -func (_c *Repository_RetrieveEntitiesRolesActionsMembers_Call) RunAndReturn(run func(ctx context.Context, entityIDs []string) ([]roles.EntityActionRole, []roles.EntityMemberRole, error)) *Repository_RetrieveEntitiesRolesActionsMembers_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveEntityRole provides a mock function for the type Repository -func (_mock *Repository) RetrieveEntityRole(ctx context.Context, entityID string, roleID string) (roles.Role, error) { - ret := _mock.Called(ctx, entityID, roleID) - - if len(ret) == 0 { - panic("no return value specified for RetrieveEntityRole") - } - - var r0 roles.Role - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) (roles.Role, error)); ok { - return returnFunc(ctx, entityID, roleID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) roles.Role); ok { - r0 = returnFunc(ctx, entityID, roleID) - } else { - r0 = ret.Get(0).(roles.Role) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, string) error); ok { - r1 = returnFunc(ctx, entityID, roleID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveEntityRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveEntityRole' -type Repository_RetrieveEntityRole_Call struct { - *mock.Call -} - -// RetrieveEntityRole is a helper method to define mock.On call -// - ctx context.Context -// - entityID string -// - roleID string -func (_e *Repository_Expecter) RetrieveEntityRole(ctx interface{}, entityID interface{}, roleID interface{}) *Repository_RetrieveEntityRole_Call { - return &Repository_RetrieveEntityRole_Call{Call: _e.mock.On("RetrieveEntityRole", ctx, entityID, roleID)} -} - -func (_c *Repository_RetrieveEntityRole_Call) Run(run func(ctx context.Context, entityID string, roleID string)) *Repository_RetrieveEntityRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RetrieveEntityRole_Call) Return(role roles.Role, err error) *Repository_RetrieveEntityRole_Call { - _c.Call.Return(role, err) - return _c -} - -func (_c *Repository_RetrieveEntityRole_Call) RunAndReturn(run func(ctx context.Context, entityID string, roleID string) (roles.Role, error)) *Repository_RetrieveEntityRole_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveRole provides a mock function for the type Repository -func (_mock *Repository) RetrieveRole(ctx context.Context, roleID string) (roles.Role, error) { - ret := _mock.Called(ctx, roleID) - - if len(ret) == 0 { - panic("no return value specified for RetrieveRole") - } - - var r0 roles.Role - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) (roles.Role, error)); ok { - return returnFunc(ctx, roleID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string) roles.Role); ok { - r0 = returnFunc(ctx, roleID) - } else { - r0 = ret.Get(0).(roles.Role) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = returnFunc(ctx, roleID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveRole' -type Repository_RetrieveRole_Call struct { - *mock.Call -} - -// RetrieveRole is a helper method to define mock.On call -// - ctx context.Context -// - roleID string -func (_e *Repository_Expecter) RetrieveRole(ctx interface{}, roleID interface{}) *Repository_RetrieveRole_Call { - return &Repository_RetrieveRole_Call{Call: _e.mock.On("RetrieveRole", ctx, roleID)} -} - -func (_c *Repository_RetrieveRole_Call) Run(run func(ctx context.Context, roleID string)) *Repository_RetrieveRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RetrieveRole_Call) Return(role roles.Role, err error) *Repository_RetrieveRole_Call { - _c.Call.Return(role, err) - return _c -} - -func (_c *Repository_RetrieveRole_Call) RunAndReturn(run func(ctx context.Context, roleID string) (roles.Role, error)) *Repository_RetrieveRole_Call { - _c.Call.Return(run) - return _c -} - -// RoleAddActions provides a mock function for the type Repository -func (_mock *Repository) RoleAddActions(ctx context.Context, role roles.Role, actions []string) ([]string, error) { - ret := _mock.Called(ctx, role, actions) - - if len(ret) == 0 { - panic("no return value specified for RoleAddActions") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role, []string) ([]string, error)); ok { - return returnFunc(ctx, role, actions) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role, []string) []string); ok { - r0 = returnFunc(ctx, role, actions) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, roles.Role, []string) error); ok { - r1 = returnFunc(ctx, role, actions) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RoleAddActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleAddActions' -type Repository_RoleAddActions_Call struct { - *mock.Call -} - -// RoleAddActions is a helper method to define mock.On call -// - ctx context.Context -// - role roles.Role -// - actions []string -func (_e *Repository_Expecter) RoleAddActions(ctx interface{}, role interface{}, actions interface{}) *Repository_RoleAddActions_Call { - return &Repository_RoleAddActions_Call{Call: _e.mock.On("RoleAddActions", ctx, role, actions)} -} - -func (_c *Repository_RoleAddActions_Call) Run(run func(ctx context.Context, role roles.Role, actions []string)) *Repository_RoleAddActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RoleAddActions_Call) Return(ops []string, err error) *Repository_RoleAddActions_Call { - _c.Call.Return(ops, err) - return _c -} - -func (_c *Repository_RoleAddActions_Call) RunAndReturn(run func(ctx context.Context, role roles.Role, actions []string) ([]string, error)) *Repository_RoleAddActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleAddMembers provides a mock function for the type Repository -func (_mock *Repository) RoleAddMembers(ctx context.Context, role roles.Role, members []string) ([]string, error) { - ret := _mock.Called(ctx, role, members) - - if len(ret) == 0 { - panic("no return value specified for RoleAddMembers") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role, []string) ([]string, error)); ok { - return returnFunc(ctx, role, members) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role, []string) []string); ok { - r0 = returnFunc(ctx, role, members) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, roles.Role, []string) error); ok { - r1 = returnFunc(ctx, role, members) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RoleAddMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleAddMembers' -type Repository_RoleAddMembers_Call struct { - *mock.Call -} - -// RoleAddMembers is a helper method to define mock.On call -// - ctx context.Context -// - role roles.Role -// - members []string -func (_e *Repository_Expecter) RoleAddMembers(ctx interface{}, role interface{}, members interface{}) *Repository_RoleAddMembers_Call { - return &Repository_RoleAddMembers_Call{Call: _e.mock.On("RoleAddMembers", ctx, role, members)} -} - -func (_c *Repository_RoleAddMembers_Call) Run(run func(ctx context.Context, role roles.Role, members []string)) *Repository_RoleAddMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RoleAddMembers_Call) Return(strings []string, err error) *Repository_RoleAddMembers_Call { - _c.Call.Return(strings, err) - return _c -} - -func (_c *Repository_RoleAddMembers_Call) RunAndReturn(run func(ctx context.Context, role roles.Role, members []string) ([]string, error)) *Repository_RoleAddMembers_Call { - _c.Call.Return(run) - return _c -} - -// RoleCheckActionsExists provides a mock function for the type Repository -func (_mock *Repository) RoleCheckActionsExists(ctx context.Context, roleID string, actions []string) (bool, error) { - ret := _mock.Called(ctx, roleID, actions) - - if len(ret) == 0 { - panic("no return value specified for RoleCheckActionsExists") - } - - var r0 bool - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, []string) (bool, error)); ok { - return returnFunc(ctx, roleID, actions) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, []string) bool); ok { - r0 = returnFunc(ctx, roleID, actions) - } else { - r0 = ret.Get(0).(bool) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, []string) error); ok { - r1 = returnFunc(ctx, roleID, actions) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RoleCheckActionsExists_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleCheckActionsExists' -type Repository_RoleCheckActionsExists_Call struct { - *mock.Call -} - -// RoleCheckActionsExists is a helper method to define mock.On call -// - ctx context.Context -// - roleID string -// - actions []string -func (_e *Repository_Expecter) RoleCheckActionsExists(ctx interface{}, roleID interface{}, actions interface{}) *Repository_RoleCheckActionsExists_Call { - return &Repository_RoleCheckActionsExists_Call{Call: _e.mock.On("RoleCheckActionsExists", ctx, roleID, actions)} -} - -func (_c *Repository_RoleCheckActionsExists_Call) Run(run func(ctx context.Context, roleID string, actions []string)) *Repository_RoleCheckActionsExists_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RoleCheckActionsExists_Call) Return(b bool, err error) *Repository_RoleCheckActionsExists_Call { - _c.Call.Return(b, err) - return _c -} - -func (_c *Repository_RoleCheckActionsExists_Call) RunAndReturn(run func(ctx context.Context, roleID string, actions []string) (bool, error)) *Repository_RoleCheckActionsExists_Call { - _c.Call.Return(run) - return _c -} - -// RoleCheckMembersExists provides a mock function for the type Repository -func (_mock *Repository) RoleCheckMembersExists(ctx context.Context, roleID string, members []string) (bool, error) { - ret := _mock.Called(ctx, roleID, members) - - if len(ret) == 0 { - panic("no return value specified for RoleCheckMembersExists") - } - - var r0 bool - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, []string) (bool, error)); ok { - return returnFunc(ctx, roleID, members) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, []string) bool); ok { - r0 = returnFunc(ctx, roleID, members) - } else { - r0 = ret.Get(0).(bool) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, []string) error); ok { - r1 = returnFunc(ctx, roleID, members) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RoleCheckMembersExists_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleCheckMembersExists' -type Repository_RoleCheckMembersExists_Call struct { - *mock.Call -} - -// RoleCheckMembersExists is a helper method to define mock.On call -// - ctx context.Context -// - roleID string -// - members []string -func (_e *Repository_Expecter) RoleCheckMembersExists(ctx interface{}, roleID interface{}, members interface{}) *Repository_RoleCheckMembersExists_Call { - return &Repository_RoleCheckMembersExists_Call{Call: _e.mock.On("RoleCheckMembersExists", ctx, roleID, members)} -} - -func (_c *Repository_RoleCheckMembersExists_Call) Run(run func(ctx context.Context, roleID string, members []string)) *Repository_RoleCheckMembersExists_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RoleCheckMembersExists_Call) Return(b bool, err error) *Repository_RoleCheckMembersExists_Call { - _c.Call.Return(b, err) - return _c -} - -func (_c *Repository_RoleCheckMembersExists_Call) RunAndReturn(run func(ctx context.Context, roleID string, members []string) (bool, error)) *Repository_RoleCheckMembersExists_Call { - _c.Call.Return(run) - return _c -} - -// RoleListActions provides a mock function for the type Repository -func (_mock *Repository) RoleListActions(ctx context.Context, roleID string) ([]string, error) { - ret := _mock.Called(ctx, roleID) - - if len(ret) == 0 { - panic("no return value specified for RoleListActions") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) ([]string, error)); ok { - return returnFunc(ctx, roleID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string) []string); ok { - r0 = returnFunc(ctx, roleID) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = returnFunc(ctx, roleID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RoleListActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleListActions' -type Repository_RoleListActions_Call struct { - *mock.Call -} - -// RoleListActions is a helper method to define mock.On call -// - ctx context.Context -// - roleID string -func (_e *Repository_Expecter) RoleListActions(ctx interface{}, roleID interface{}) *Repository_RoleListActions_Call { - return &Repository_RoleListActions_Call{Call: _e.mock.On("RoleListActions", ctx, roleID)} -} - -func (_c *Repository_RoleListActions_Call) Run(run func(ctx context.Context, roleID string)) *Repository_RoleListActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RoleListActions_Call) Return(strings []string, err error) *Repository_RoleListActions_Call { - _c.Call.Return(strings, err) - return _c -} - -func (_c *Repository_RoleListActions_Call) RunAndReturn(run func(ctx context.Context, roleID string) ([]string, error)) *Repository_RoleListActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleListMembers provides a mock function for the type Repository -func (_mock *Repository) RoleListMembers(ctx context.Context, roleID string, limit uint64, offset uint64) (roles.MembersPage, error) { - ret := _mock.Called(ctx, roleID, limit, offset) - - if len(ret) == 0 { - panic("no return value specified for RoleListMembers") - } - - var r0 roles.MembersPage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, uint64, uint64) (roles.MembersPage, error)); ok { - return returnFunc(ctx, roleID, limit, offset) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, uint64, uint64) roles.MembersPage); ok { - r0 = returnFunc(ctx, roleID, limit, offset) - } else { - r0 = ret.Get(0).(roles.MembersPage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, uint64, uint64) error); ok { - r1 = returnFunc(ctx, roleID, limit, offset) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RoleListMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleListMembers' -type Repository_RoleListMembers_Call struct { - *mock.Call -} - -// RoleListMembers is a helper method to define mock.On call -// - ctx context.Context -// - roleID string -// - limit uint64 -// - offset uint64 -func (_e *Repository_Expecter) RoleListMembers(ctx interface{}, roleID interface{}, limit interface{}, offset interface{}) *Repository_RoleListMembers_Call { - return &Repository_RoleListMembers_Call{Call: _e.mock.On("RoleListMembers", ctx, roleID, limit, offset)} -} - -func (_c *Repository_RoleListMembers_Call) Run(run func(ctx context.Context, roleID string, limit uint64, offset uint64)) *Repository_RoleListMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 uint64 - if args[2] != nil { - arg2 = args[2].(uint64) - } - var arg3 uint64 - if args[3] != nil { - arg3 = args[3].(uint64) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Repository_RoleListMembers_Call) Return(membersPage roles.MembersPage, err error) *Repository_RoleListMembers_Call { - _c.Call.Return(membersPage, err) - return _c -} - -func (_c *Repository_RoleListMembers_Call) RunAndReturn(run func(ctx context.Context, roleID string, limit uint64, offset uint64) (roles.MembersPage, error)) *Repository_RoleListMembers_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveActions provides a mock function for the type Repository -func (_mock *Repository) RoleRemoveActions(ctx context.Context, role roles.Role, actions []string) error { - ret := _mock.Called(ctx, role, actions) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveActions") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role, []string) error); ok { - r0 = returnFunc(ctx, role, actions) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RoleRemoveActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveActions' -type Repository_RoleRemoveActions_Call struct { - *mock.Call -} - -// RoleRemoveActions is a helper method to define mock.On call -// - ctx context.Context -// - role roles.Role -// - actions []string -func (_e *Repository_Expecter) RoleRemoveActions(ctx interface{}, role interface{}, actions interface{}) *Repository_RoleRemoveActions_Call { - return &Repository_RoleRemoveActions_Call{Call: _e.mock.On("RoleRemoveActions", ctx, role, actions)} -} - -func (_c *Repository_RoleRemoveActions_Call) Run(run func(ctx context.Context, role roles.Role, actions []string)) *Repository_RoleRemoveActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RoleRemoveActions_Call) Return(err error) *Repository_RoleRemoveActions_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RoleRemoveActions_Call) RunAndReturn(run func(ctx context.Context, role roles.Role, actions []string) error) *Repository_RoleRemoveActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveAllActions provides a mock function for the type Repository -func (_mock *Repository) RoleRemoveAllActions(ctx context.Context, role roles.Role) error { - ret := _mock.Called(ctx, role) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveAllActions") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role) error); ok { - r0 = returnFunc(ctx, role) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RoleRemoveAllActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveAllActions' -type Repository_RoleRemoveAllActions_Call struct { - *mock.Call -} - -// RoleRemoveAllActions is a helper method to define mock.On call -// - ctx context.Context -// - role roles.Role -func (_e *Repository_Expecter) RoleRemoveAllActions(ctx interface{}, role interface{}) *Repository_RoleRemoveAllActions_Call { - return &Repository_RoleRemoveAllActions_Call{Call: _e.mock.On("RoleRemoveAllActions", ctx, role)} -} - -func (_c *Repository_RoleRemoveAllActions_Call) Run(run func(ctx context.Context, role roles.Role)) *Repository_RoleRemoveAllActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RoleRemoveAllActions_Call) Return(err error) *Repository_RoleRemoveAllActions_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RoleRemoveAllActions_Call) RunAndReturn(run func(ctx context.Context, role roles.Role) error) *Repository_RoleRemoveAllActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveAllMembers provides a mock function for the type Repository -func (_mock *Repository) RoleRemoveAllMembers(ctx context.Context, role roles.Role) error { - ret := _mock.Called(ctx, role) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveAllMembers") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role) error); ok { - r0 = returnFunc(ctx, role) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RoleRemoveAllMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveAllMembers' -type Repository_RoleRemoveAllMembers_Call struct { - *mock.Call -} - -// RoleRemoveAllMembers is a helper method to define mock.On call -// - ctx context.Context -// - role roles.Role -func (_e *Repository_Expecter) RoleRemoveAllMembers(ctx interface{}, role interface{}) *Repository_RoleRemoveAllMembers_Call { - return &Repository_RoleRemoveAllMembers_Call{Call: _e.mock.On("RoleRemoveAllMembers", ctx, role)} -} - -func (_c *Repository_RoleRemoveAllMembers_Call) Run(run func(ctx context.Context, role roles.Role)) *Repository_RoleRemoveAllMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RoleRemoveAllMembers_Call) Return(err error) *Repository_RoleRemoveAllMembers_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RoleRemoveAllMembers_Call) RunAndReturn(run func(ctx context.Context, role roles.Role) error) *Repository_RoleRemoveAllMembers_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveMembers provides a mock function for the type Repository -func (_mock *Repository) RoleRemoveMembers(ctx context.Context, role roles.Role, members []string) error { - ret := _mock.Called(ctx, role, members) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveMembers") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role, []string) error); ok { - r0 = returnFunc(ctx, role, members) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RoleRemoveMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveMembers' -type Repository_RoleRemoveMembers_Call struct { - *mock.Call -} - -// RoleRemoveMembers is a helper method to define mock.On call -// - ctx context.Context -// - role roles.Role -// - members []string -func (_e *Repository_Expecter) RoleRemoveMembers(ctx interface{}, role interface{}, members interface{}) *Repository_RoleRemoveMembers_Call { - return &Repository_RoleRemoveMembers_Call{Call: _e.mock.On("RoleRemoveMembers", ctx, role, members)} -} - -func (_c *Repository_RoleRemoveMembers_Call) Run(run func(ctx context.Context, role roles.Role, members []string)) *Repository_RoleRemoveMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RoleRemoveMembers_Call) Return(err error) *Repository_RoleRemoveMembers_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RoleRemoveMembers_Call) RunAndReturn(run func(ctx context.Context, role roles.Role, members []string) error) *Repository_RoleRemoveMembers_Call { - _c.Call.Return(run) - return _c -} - -// UpdateRole provides a mock function for the type Repository -func (_mock *Repository) UpdateRole(ctx context.Context, ro roles.Role) (roles.Role, error) { - ret := _mock.Called(ctx, ro) - - if len(ret) == 0 { - panic("no return value specified for UpdateRole") - } - - var r0 roles.Role - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role) (roles.Role, error)); ok { - return returnFunc(ctx, ro) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role) roles.Role); ok { - r0 = returnFunc(ctx, ro) - } else { - r0 = ret.Get(0).(roles.Role) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, roles.Role) error); ok { - r1 = returnFunc(ctx, ro) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_UpdateRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateRole' -type Repository_UpdateRole_Call struct { - *mock.Call -} - -// UpdateRole is a helper method to define mock.On call -// - ctx context.Context -// - ro roles.Role -func (_e *Repository_Expecter) UpdateRole(ctx interface{}, ro interface{}) *Repository_UpdateRole_Call { - return &Repository_UpdateRole_Call{Call: _e.mock.On("UpdateRole", ctx, ro)} -} - -func (_c *Repository_UpdateRole_Call) Run(run func(ctx context.Context, ro roles.Role)) *Repository_UpdateRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_UpdateRole_Call) Return(role roles.Role, err error) *Repository_UpdateRole_Call { - _c.Call.Return(role, err) - return _c -} - -func (_c *Repository_UpdateRole_Call) RunAndReturn(run func(ctx context.Context, ro roles.Role) (roles.Role, error)) *Repository_UpdateRole_Call { - _c.Call.Return(run) - return _c -} diff --git a/pkg/roles/mocks/role_manager.go b/pkg/roles/mocks/role_manager.go deleted file mode 100644 index 49f8f9a10..000000000 --- a/pkg/roles/mocks/role_manager.go +++ /dev/null @@ -1,1525 +0,0 @@ -// Copyright (c) Abstract Machines - -// SPDX-License-Identifier: Apache-2.0 - -// Code generated by mockery; DO NOT EDIT. -// github.com/vektra/mockery -// template: testify - -package mocks - -import ( - "context" - - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/roles" - mock "github.com/stretchr/testify/mock" -) - -// NewRoleManager creates a new instance of RoleManager. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. -// The first argument is typically a *testing.T value. -func NewRoleManager(t interface { - mock.TestingT - Cleanup(func()) -}) *RoleManager { - mock := &RoleManager{} - mock.Mock.Test(t) - - t.Cleanup(func() { mock.AssertExpectations(t) }) - - return mock -} - -// RoleManager is an autogenerated mock type for the RoleManager type -type RoleManager struct { - mock.Mock -} - -type RoleManager_Expecter struct { - mock *mock.Mock -} - -func (_m *RoleManager) EXPECT() *RoleManager_Expecter { - return &RoleManager_Expecter{mock: &_m.Mock} -} - -// AddRole provides a mock function for the type RoleManager -func (_mock *RoleManager) AddRole(ctx context.Context, session authn.Session, entityID string, roleName string, optionalActions []string, optionalMembers []string) (roles.RoleProvision, error) { - ret := _mock.Called(ctx, session, entityID, roleName, optionalActions, optionalMembers) - - if len(ret) == 0 { - panic("no return value specified for AddRole") - } - - var r0 roles.RoleProvision - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string, []string) (roles.RoleProvision, error)); ok { - return returnFunc(ctx, session, entityID, roleName, optionalActions, optionalMembers) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string, []string) roles.RoleProvision); ok { - r0 = returnFunc(ctx, session, entityID, roleName, optionalActions, optionalMembers) - } else { - r0 = ret.Get(0).(roles.RoleProvision) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, []string, []string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleName, optionalActions, optionalMembers) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// RoleManager_AddRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'AddRole' -type RoleManager_AddRole_Call struct { - *mock.Call -} - -// AddRole is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleName string -// - optionalActions []string -// - optionalMembers []string -func (_e *RoleManager_Expecter) AddRole(ctx interface{}, session interface{}, entityID interface{}, roleName interface{}, optionalActions interface{}, optionalMembers interface{}) *RoleManager_AddRole_Call { - return &RoleManager_AddRole_Call{Call: _e.mock.On("AddRole", ctx, session, entityID, roleName, optionalActions, optionalMembers)} -} - -func (_c *RoleManager_AddRole_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleName string, optionalActions []string, optionalMembers []string)) *RoleManager_AddRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - var arg5 []string - if args[5] != nil { - arg5 = args[5].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - arg5, - ) - }) - return _c -} - -func (_c *RoleManager_AddRole_Call) Return(roleProvision roles.RoleProvision, err error) *RoleManager_AddRole_Call { - _c.Call.Return(roleProvision, err) - return _c -} - -func (_c *RoleManager_AddRole_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleName string, optionalActions []string, optionalMembers []string) (roles.RoleProvision, error)) *RoleManager_AddRole_Call { - _c.Call.Return(run) - return _c -} - -// ListAvailableActions provides a mock function for the type RoleManager -func (_mock *RoleManager) ListAvailableActions(ctx context.Context, session authn.Session) ([]string, error) { - ret := _mock.Called(ctx, session) - - if len(ret) == 0 { - panic("no return value specified for ListAvailableActions") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session) ([]string, error)); ok { - return returnFunc(ctx, session) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session) []string); ok { - r0 = returnFunc(ctx, session) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session) error); ok { - r1 = returnFunc(ctx, session) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// RoleManager_ListAvailableActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ListAvailableActions' -type RoleManager_ListAvailableActions_Call struct { - *mock.Call -} - -// ListAvailableActions is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -func (_e *RoleManager_Expecter) ListAvailableActions(ctx interface{}, session interface{}) *RoleManager_ListAvailableActions_Call { - return &RoleManager_ListAvailableActions_Call{Call: _e.mock.On("ListAvailableActions", ctx, session)} -} - -func (_c *RoleManager_ListAvailableActions_Call) Run(run func(ctx context.Context, session authn.Session)) *RoleManager_ListAvailableActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *RoleManager_ListAvailableActions_Call) Return(strings []string, err error) *RoleManager_ListAvailableActions_Call { - _c.Call.Return(strings, err) - return _c -} - -func (_c *RoleManager_ListAvailableActions_Call) RunAndReturn(run func(ctx context.Context, session authn.Session) ([]string, error)) *RoleManager_ListAvailableActions_Call { - _c.Call.Return(run) - return _c -} - -// ListEntityMembers provides a mock function for the type RoleManager -func (_mock *RoleManager) ListEntityMembers(ctx context.Context, session authn.Session, entityID string, pq roles.MembersRolePageQuery) (roles.MembersRolePage, error) { - ret := _mock.Called(ctx, session, entityID, pq) - - if len(ret) == 0 { - panic("no return value specified for ListEntityMembers") - } - - var r0 roles.MembersRolePage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, roles.MembersRolePageQuery) (roles.MembersRolePage, error)); ok { - return returnFunc(ctx, session, entityID, pq) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, roles.MembersRolePageQuery) roles.MembersRolePage); ok { - r0 = returnFunc(ctx, session, entityID, pq) - } else { - r0 = ret.Get(0).(roles.MembersRolePage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, roles.MembersRolePageQuery) error); ok { - r1 = returnFunc(ctx, session, entityID, pq) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// RoleManager_ListEntityMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ListEntityMembers' -type RoleManager_ListEntityMembers_Call struct { - *mock.Call -} - -// ListEntityMembers is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - pq roles.MembersRolePageQuery -func (_e *RoleManager_Expecter) ListEntityMembers(ctx interface{}, session interface{}, entityID interface{}, pq interface{}) *RoleManager_ListEntityMembers_Call { - return &RoleManager_ListEntityMembers_Call{Call: _e.mock.On("ListEntityMembers", ctx, session, entityID, pq)} -} - -func (_c *RoleManager_ListEntityMembers_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, pq roles.MembersRolePageQuery)) *RoleManager_ListEntityMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 roles.MembersRolePageQuery - if args[3] != nil { - arg3 = args[3].(roles.MembersRolePageQuery) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *RoleManager_ListEntityMembers_Call) Return(membersRolePage roles.MembersRolePage, err error) *RoleManager_ListEntityMembers_Call { - _c.Call.Return(membersRolePage, err) - return _c -} - -func (_c *RoleManager_ListEntityMembers_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, pq roles.MembersRolePageQuery) (roles.MembersRolePage, error)) *RoleManager_ListEntityMembers_Call { - _c.Call.Return(run) - return _c -} - -// RemoveEntityMembers provides a mock function for the type RoleManager -func (_mock *RoleManager) RemoveEntityMembers(ctx context.Context, session authn.Session, entityID string, members []string) error { - ret := _mock.Called(ctx, session, entityID, members) - - if len(ret) == 0 { - panic("no return value specified for RemoveEntityMembers") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, []string) error); ok { - r0 = returnFunc(ctx, session, entityID, members) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// RoleManager_RemoveEntityMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveEntityMembers' -type RoleManager_RemoveEntityMembers_Call struct { - *mock.Call -} - -// RemoveEntityMembers is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - members []string -func (_e *RoleManager_Expecter) RemoveEntityMembers(ctx interface{}, session interface{}, entityID interface{}, members interface{}) *RoleManager_RemoveEntityMembers_Call { - return &RoleManager_RemoveEntityMembers_Call{Call: _e.mock.On("RemoveEntityMembers", ctx, session, entityID, members)} -} - -func (_c *RoleManager_RemoveEntityMembers_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, members []string)) *RoleManager_RemoveEntityMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 []string - if args[3] != nil { - arg3 = args[3].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *RoleManager_RemoveEntityMembers_Call) Return(err error) *RoleManager_RemoveEntityMembers_Call { - _c.Call.Return(err) - return _c -} - -func (_c *RoleManager_RemoveEntityMembers_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, members []string) error) *RoleManager_RemoveEntityMembers_Call { - _c.Call.Return(run) - return _c -} - -// RemoveMemberFromAllRoles provides a mock function for the type RoleManager -func (_mock *RoleManager) RemoveMemberFromAllRoles(ctx context.Context, session authn.Session, memberID string) error { - ret := _mock.Called(ctx, session, memberID) - - if len(ret) == 0 { - panic("no return value specified for RemoveMemberFromAllRoles") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) error); ok { - r0 = returnFunc(ctx, session, memberID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// RoleManager_RemoveMemberFromAllRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveMemberFromAllRoles' -type RoleManager_RemoveMemberFromAllRoles_Call struct { - *mock.Call -} - -// RemoveMemberFromAllRoles is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - memberID string -func (_e *RoleManager_Expecter) RemoveMemberFromAllRoles(ctx interface{}, session interface{}, memberID interface{}) *RoleManager_RemoveMemberFromAllRoles_Call { - return &RoleManager_RemoveMemberFromAllRoles_Call{Call: _e.mock.On("RemoveMemberFromAllRoles", ctx, session, memberID)} -} - -func (_c *RoleManager_RemoveMemberFromAllRoles_Call) Run(run func(ctx context.Context, session authn.Session, memberID string)) *RoleManager_RemoveMemberFromAllRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *RoleManager_RemoveMemberFromAllRoles_Call) Return(err error) *RoleManager_RemoveMemberFromAllRoles_Call { - _c.Call.Return(err) - return _c -} - -func (_c *RoleManager_RemoveMemberFromAllRoles_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, memberID string) error) *RoleManager_RemoveMemberFromAllRoles_Call { - _c.Call.Return(run) - return _c -} - -// RemoveRole provides a mock function for the type RoleManager -func (_mock *RoleManager) RemoveRole(ctx context.Context, session authn.Session, entityID string, roleID string) error { - ret := _mock.Called(ctx, session, entityID, roleID) - - if len(ret) == 0 { - panic("no return value specified for RemoveRole") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) error); ok { - r0 = returnFunc(ctx, session, entityID, roleID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// RoleManager_RemoveRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveRole' -type RoleManager_RemoveRole_Call struct { - *mock.Call -} - -// RemoveRole is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -func (_e *RoleManager_Expecter) RemoveRole(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}) *RoleManager_RemoveRole_Call { - return &RoleManager_RemoveRole_Call{Call: _e.mock.On("RemoveRole", ctx, session, entityID, roleID)} -} - -func (_c *RoleManager_RemoveRole_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string)) *RoleManager_RemoveRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *RoleManager_RemoveRole_Call) Return(err error) *RoleManager_RemoveRole_Call { - _c.Call.Return(err) - return _c -} - -func (_c *RoleManager_RemoveRole_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string) error) *RoleManager_RemoveRole_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveAllRoles provides a mock function for the type RoleManager -func (_mock *RoleManager) RetrieveAllRoles(ctx context.Context, session authn.Session, entityID string, limit uint64, offset uint64) (roles.RolePage, error) { - ret := _mock.Called(ctx, session, entityID, limit, offset) - - if len(ret) == 0 { - panic("no return value specified for RetrieveAllRoles") - } - - var r0 roles.RolePage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, uint64, uint64) (roles.RolePage, error)); ok { - return returnFunc(ctx, session, entityID, limit, offset) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, uint64, uint64) roles.RolePage); ok { - r0 = returnFunc(ctx, session, entityID, limit, offset) - } else { - r0 = ret.Get(0).(roles.RolePage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, uint64, uint64) error); ok { - r1 = returnFunc(ctx, session, entityID, limit, offset) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// RoleManager_RetrieveAllRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveAllRoles' -type RoleManager_RetrieveAllRoles_Call struct { - *mock.Call -} - -// RetrieveAllRoles is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - limit uint64 -// - offset uint64 -func (_e *RoleManager_Expecter) RetrieveAllRoles(ctx interface{}, session interface{}, entityID interface{}, limit interface{}, offset interface{}) *RoleManager_RetrieveAllRoles_Call { - return &RoleManager_RetrieveAllRoles_Call{Call: _e.mock.On("RetrieveAllRoles", ctx, session, entityID, limit, offset)} -} - -func (_c *RoleManager_RetrieveAllRoles_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, limit uint64, offset uint64)) *RoleManager_RetrieveAllRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 uint64 - if args[3] != nil { - arg3 = args[3].(uint64) - } - var arg4 uint64 - if args[4] != nil { - arg4 = args[4].(uint64) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *RoleManager_RetrieveAllRoles_Call) Return(rolePage roles.RolePage, err error) *RoleManager_RetrieveAllRoles_Call { - _c.Call.Return(rolePage, err) - return _c -} - -func (_c *RoleManager_RetrieveAllRoles_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, limit uint64, offset uint64) (roles.RolePage, error)) *RoleManager_RetrieveAllRoles_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveRole provides a mock function for the type RoleManager -func (_mock *RoleManager) RetrieveRole(ctx context.Context, session authn.Session, entityID string, roleID string) (roles.Role, error) { - ret := _mock.Called(ctx, session, entityID, roleID) - - if len(ret) == 0 { - panic("no return value specified for RetrieveRole") - } - - var r0 roles.Role - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) (roles.Role, error)); ok { - return returnFunc(ctx, session, entityID, roleID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) roles.Role); ok { - r0 = returnFunc(ctx, session, entityID, roleID) - } else { - r0 = ret.Get(0).(roles.Role) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// RoleManager_RetrieveRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveRole' -type RoleManager_RetrieveRole_Call struct { - *mock.Call -} - -// RetrieveRole is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -func (_e *RoleManager_Expecter) RetrieveRole(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}) *RoleManager_RetrieveRole_Call { - return &RoleManager_RetrieveRole_Call{Call: _e.mock.On("RetrieveRole", ctx, session, entityID, roleID)} -} - -func (_c *RoleManager_RetrieveRole_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string)) *RoleManager_RetrieveRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *RoleManager_RetrieveRole_Call) Return(role roles.Role, err error) *RoleManager_RetrieveRole_Call { - _c.Call.Return(role, err) - return _c -} - -func (_c *RoleManager_RetrieveRole_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string) (roles.Role, error)) *RoleManager_RetrieveRole_Call { - _c.Call.Return(run) - return _c -} - -// RoleAddActions provides a mock function for the type RoleManager -func (_mock *RoleManager) RoleAddActions(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string) ([]string, error) { - ret := _mock.Called(ctx, session, entityID, roleID, actions) - - if len(ret) == 0 { - panic("no return value specified for RoleAddActions") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) ([]string, error)); ok { - return returnFunc(ctx, session, entityID, roleID, actions) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) []string); ok { - r0 = returnFunc(ctx, session, entityID, roleID, actions) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, []string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID, actions) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// RoleManager_RoleAddActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleAddActions' -type RoleManager_RoleAddActions_Call struct { - *mock.Call -} - -// RoleAddActions is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - actions []string -func (_e *RoleManager_Expecter) RoleAddActions(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, actions interface{}) *RoleManager_RoleAddActions_Call { - return &RoleManager_RoleAddActions_Call{Call: _e.mock.On("RoleAddActions", ctx, session, entityID, roleID, actions)} -} - -func (_c *RoleManager_RoleAddActions_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string)) *RoleManager_RoleAddActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *RoleManager_RoleAddActions_Call) Return(ops []string, err error) *RoleManager_RoleAddActions_Call { - _c.Call.Return(ops, err) - return _c -} - -func (_c *RoleManager_RoleAddActions_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string) ([]string, error)) *RoleManager_RoleAddActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleAddMembers provides a mock function for the type RoleManager -func (_mock *RoleManager) RoleAddMembers(ctx context.Context, session authn.Session, entityID string, roleID string, members []string) ([]string, error) { - ret := _mock.Called(ctx, session, entityID, roleID, members) - - if len(ret) == 0 { - panic("no return value specified for RoleAddMembers") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) ([]string, error)); ok { - return returnFunc(ctx, session, entityID, roleID, members) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) []string); ok { - r0 = returnFunc(ctx, session, entityID, roleID, members) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, []string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID, members) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// RoleManager_RoleAddMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleAddMembers' -type RoleManager_RoleAddMembers_Call struct { - *mock.Call -} - -// RoleAddMembers is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - members []string -func (_e *RoleManager_Expecter) RoleAddMembers(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, members interface{}) *RoleManager_RoleAddMembers_Call { - return &RoleManager_RoleAddMembers_Call{Call: _e.mock.On("RoleAddMembers", ctx, session, entityID, roleID, members)} -} - -func (_c *RoleManager_RoleAddMembers_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, members []string)) *RoleManager_RoleAddMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *RoleManager_RoleAddMembers_Call) Return(strings []string, err error) *RoleManager_RoleAddMembers_Call { - _c.Call.Return(strings, err) - return _c -} - -func (_c *RoleManager_RoleAddMembers_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, members []string) ([]string, error)) *RoleManager_RoleAddMembers_Call { - _c.Call.Return(run) - return _c -} - -// RoleCheckActionsExists provides a mock function for the type RoleManager -func (_mock *RoleManager) RoleCheckActionsExists(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string) (bool, error) { - ret := _mock.Called(ctx, session, entityID, roleID, actions) - - if len(ret) == 0 { - panic("no return value specified for RoleCheckActionsExists") - } - - var r0 bool - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) (bool, error)); ok { - return returnFunc(ctx, session, entityID, roleID, actions) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) bool); ok { - r0 = returnFunc(ctx, session, entityID, roleID, actions) - } else { - r0 = ret.Get(0).(bool) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, []string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID, actions) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// RoleManager_RoleCheckActionsExists_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleCheckActionsExists' -type RoleManager_RoleCheckActionsExists_Call struct { - *mock.Call -} - -// RoleCheckActionsExists is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - actions []string -func (_e *RoleManager_Expecter) RoleCheckActionsExists(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, actions interface{}) *RoleManager_RoleCheckActionsExists_Call { - return &RoleManager_RoleCheckActionsExists_Call{Call: _e.mock.On("RoleCheckActionsExists", ctx, session, entityID, roleID, actions)} -} - -func (_c *RoleManager_RoleCheckActionsExists_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string)) *RoleManager_RoleCheckActionsExists_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *RoleManager_RoleCheckActionsExists_Call) Return(b bool, err error) *RoleManager_RoleCheckActionsExists_Call { - _c.Call.Return(b, err) - return _c -} - -func (_c *RoleManager_RoleCheckActionsExists_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string) (bool, error)) *RoleManager_RoleCheckActionsExists_Call { - _c.Call.Return(run) - return _c -} - -// RoleCheckMembersExists provides a mock function for the type RoleManager -func (_mock *RoleManager) RoleCheckMembersExists(ctx context.Context, session authn.Session, entityID string, roleID string, members []string) (bool, error) { - ret := _mock.Called(ctx, session, entityID, roleID, members) - - if len(ret) == 0 { - panic("no return value specified for RoleCheckMembersExists") - } - - var r0 bool - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) (bool, error)); ok { - return returnFunc(ctx, session, entityID, roleID, members) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) bool); ok { - r0 = returnFunc(ctx, session, entityID, roleID, members) - } else { - r0 = ret.Get(0).(bool) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, []string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID, members) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// RoleManager_RoleCheckMembersExists_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleCheckMembersExists' -type RoleManager_RoleCheckMembersExists_Call struct { - *mock.Call -} - -// RoleCheckMembersExists is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - members []string -func (_e *RoleManager_Expecter) RoleCheckMembersExists(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, members interface{}) *RoleManager_RoleCheckMembersExists_Call { - return &RoleManager_RoleCheckMembersExists_Call{Call: _e.mock.On("RoleCheckMembersExists", ctx, session, entityID, roleID, members)} -} - -func (_c *RoleManager_RoleCheckMembersExists_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, members []string)) *RoleManager_RoleCheckMembersExists_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *RoleManager_RoleCheckMembersExists_Call) Return(b bool, err error) *RoleManager_RoleCheckMembersExists_Call { - _c.Call.Return(b, err) - return _c -} - -func (_c *RoleManager_RoleCheckMembersExists_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, members []string) (bool, error)) *RoleManager_RoleCheckMembersExists_Call { - _c.Call.Return(run) - return _c -} - -// RoleListActions provides a mock function for the type RoleManager -func (_mock *RoleManager) RoleListActions(ctx context.Context, session authn.Session, entityID string, roleID string) ([]string, error) { - ret := _mock.Called(ctx, session, entityID, roleID) - - if len(ret) == 0 { - panic("no return value specified for RoleListActions") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) ([]string, error)); ok { - return returnFunc(ctx, session, entityID, roleID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) []string); ok { - r0 = returnFunc(ctx, session, entityID, roleID) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// RoleManager_RoleListActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleListActions' -type RoleManager_RoleListActions_Call struct { - *mock.Call -} - -// RoleListActions is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -func (_e *RoleManager_Expecter) RoleListActions(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}) *RoleManager_RoleListActions_Call { - return &RoleManager_RoleListActions_Call{Call: _e.mock.On("RoleListActions", ctx, session, entityID, roleID)} -} - -func (_c *RoleManager_RoleListActions_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string)) *RoleManager_RoleListActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *RoleManager_RoleListActions_Call) Return(strings []string, err error) *RoleManager_RoleListActions_Call { - _c.Call.Return(strings, err) - return _c -} - -func (_c *RoleManager_RoleListActions_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string) ([]string, error)) *RoleManager_RoleListActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleListMembers provides a mock function for the type RoleManager -func (_mock *RoleManager) RoleListMembers(ctx context.Context, session authn.Session, entityID string, roleID string, limit uint64, offset uint64) (roles.MembersPage, error) { - ret := _mock.Called(ctx, session, entityID, roleID, limit, offset) - - if len(ret) == 0 { - panic("no return value specified for RoleListMembers") - } - - var r0 roles.MembersPage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, uint64, uint64) (roles.MembersPage, error)); ok { - return returnFunc(ctx, session, entityID, roleID, limit, offset) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, uint64, uint64) roles.MembersPage); ok { - r0 = returnFunc(ctx, session, entityID, roleID, limit, offset) - } else { - r0 = ret.Get(0).(roles.MembersPage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, uint64, uint64) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID, limit, offset) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// RoleManager_RoleListMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleListMembers' -type RoleManager_RoleListMembers_Call struct { - *mock.Call -} - -// RoleListMembers is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - limit uint64 -// - offset uint64 -func (_e *RoleManager_Expecter) RoleListMembers(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, limit interface{}, offset interface{}) *RoleManager_RoleListMembers_Call { - return &RoleManager_RoleListMembers_Call{Call: _e.mock.On("RoleListMembers", ctx, session, entityID, roleID, limit, offset)} -} - -func (_c *RoleManager_RoleListMembers_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, limit uint64, offset uint64)) *RoleManager_RoleListMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 uint64 - if args[4] != nil { - arg4 = args[4].(uint64) - } - var arg5 uint64 - if args[5] != nil { - arg5 = args[5].(uint64) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - arg5, - ) - }) - return _c -} - -func (_c *RoleManager_RoleListMembers_Call) Return(membersPage roles.MembersPage, err error) *RoleManager_RoleListMembers_Call { - _c.Call.Return(membersPage, err) - return _c -} - -func (_c *RoleManager_RoleListMembers_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, limit uint64, offset uint64) (roles.MembersPage, error)) *RoleManager_RoleListMembers_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveActions provides a mock function for the type RoleManager -func (_mock *RoleManager) RoleRemoveActions(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string) error { - ret := _mock.Called(ctx, session, entityID, roleID, actions) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveActions") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) error); ok { - r0 = returnFunc(ctx, session, entityID, roleID, actions) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// RoleManager_RoleRemoveActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveActions' -type RoleManager_RoleRemoveActions_Call struct { - *mock.Call -} - -// RoleRemoveActions is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - actions []string -func (_e *RoleManager_Expecter) RoleRemoveActions(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, actions interface{}) *RoleManager_RoleRemoveActions_Call { - return &RoleManager_RoleRemoveActions_Call{Call: _e.mock.On("RoleRemoveActions", ctx, session, entityID, roleID, actions)} -} - -func (_c *RoleManager_RoleRemoveActions_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string)) *RoleManager_RoleRemoveActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *RoleManager_RoleRemoveActions_Call) Return(err error) *RoleManager_RoleRemoveActions_Call { - _c.Call.Return(err) - return _c -} - -func (_c *RoleManager_RoleRemoveActions_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string) error) *RoleManager_RoleRemoveActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveAllActions provides a mock function for the type RoleManager -func (_mock *RoleManager) RoleRemoveAllActions(ctx context.Context, session authn.Session, entityID string, roleID string) error { - ret := _mock.Called(ctx, session, entityID, roleID) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveAllActions") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) error); ok { - r0 = returnFunc(ctx, session, entityID, roleID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// RoleManager_RoleRemoveAllActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveAllActions' -type RoleManager_RoleRemoveAllActions_Call struct { - *mock.Call -} - -// RoleRemoveAllActions is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -func (_e *RoleManager_Expecter) RoleRemoveAllActions(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}) *RoleManager_RoleRemoveAllActions_Call { - return &RoleManager_RoleRemoveAllActions_Call{Call: _e.mock.On("RoleRemoveAllActions", ctx, session, entityID, roleID)} -} - -func (_c *RoleManager_RoleRemoveAllActions_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string)) *RoleManager_RoleRemoveAllActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *RoleManager_RoleRemoveAllActions_Call) Return(err error) *RoleManager_RoleRemoveAllActions_Call { - _c.Call.Return(err) - return _c -} - -func (_c *RoleManager_RoleRemoveAllActions_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string) error) *RoleManager_RoleRemoveAllActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveAllMembers provides a mock function for the type RoleManager -func (_mock *RoleManager) RoleRemoveAllMembers(ctx context.Context, session authn.Session, entityID string, roleID string) error { - ret := _mock.Called(ctx, session, entityID, roleID) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveAllMembers") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) error); ok { - r0 = returnFunc(ctx, session, entityID, roleID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// RoleManager_RoleRemoveAllMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveAllMembers' -type RoleManager_RoleRemoveAllMembers_Call struct { - *mock.Call -} - -// RoleRemoveAllMembers is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -func (_e *RoleManager_Expecter) RoleRemoveAllMembers(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}) *RoleManager_RoleRemoveAllMembers_Call { - return &RoleManager_RoleRemoveAllMembers_Call{Call: _e.mock.On("RoleRemoveAllMembers", ctx, session, entityID, roleID)} -} - -func (_c *RoleManager_RoleRemoveAllMembers_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string)) *RoleManager_RoleRemoveAllMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *RoleManager_RoleRemoveAllMembers_Call) Return(err error) *RoleManager_RoleRemoveAllMembers_Call { - _c.Call.Return(err) - return _c -} - -func (_c *RoleManager_RoleRemoveAllMembers_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string) error) *RoleManager_RoleRemoveAllMembers_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveMembers provides a mock function for the type RoleManager -func (_mock *RoleManager) RoleRemoveMembers(ctx context.Context, session authn.Session, entityID string, roleID string, members []string) error { - ret := _mock.Called(ctx, session, entityID, roleID, members) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveMembers") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) error); ok { - r0 = returnFunc(ctx, session, entityID, roleID, members) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// RoleManager_RoleRemoveMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveMembers' -type RoleManager_RoleRemoveMembers_Call struct { - *mock.Call -} - -// RoleRemoveMembers is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - members []string -func (_e *RoleManager_Expecter) RoleRemoveMembers(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, members interface{}) *RoleManager_RoleRemoveMembers_Call { - return &RoleManager_RoleRemoveMembers_Call{Call: _e.mock.On("RoleRemoveMembers", ctx, session, entityID, roleID, members)} -} - -func (_c *RoleManager_RoleRemoveMembers_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, members []string)) *RoleManager_RoleRemoveMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *RoleManager_RoleRemoveMembers_Call) Return(err error) *RoleManager_RoleRemoveMembers_Call { - _c.Call.Return(err) - return _c -} - -func (_c *RoleManager_RoleRemoveMembers_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, members []string) error) *RoleManager_RoleRemoveMembers_Call { - _c.Call.Return(run) - return _c -} - -// UpdateRoleName provides a mock function for the type RoleManager -func (_mock *RoleManager) UpdateRoleName(ctx context.Context, session authn.Session, entityID string, roleID string, newRoleName string) (roles.Role, error) { - ret := _mock.Called(ctx, session, entityID, roleID, newRoleName) - - if len(ret) == 0 { - panic("no return value specified for UpdateRoleName") - } - - var r0 roles.Role - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, string) (roles.Role, error)); ok { - return returnFunc(ctx, session, entityID, roleID, newRoleName) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, string) roles.Role); ok { - r0 = returnFunc(ctx, session, entityID, roleID, newRoleName) - } else { - r0 = ret.Get(0).(roles.Role) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID, newRoleName) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// RoleManager_UpdateRoleName_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateRoleName' -type RoleManager_UpdateRoleName_Call struct { - *mock.Call -} - -// UpdateRoleName is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - newRoleName string -func (_e *RoleManager_Expecter) UpdateRoleName(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, newRoleName interface{}) *RoleManager_UpdateRoleName_Call { - return &RoleManager_UpdateRoleName_Call{Call: _e.mock.On("UpdateRoleName", ctx, session, entityID, roleID, newRoleName)} -} - -func (_c *RoleManager_UpdateRoleName_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, newRoleName string)) *RoleManager_UpdateRoleName_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 string - if args[4] != nil { - arg4 = args[4].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *RoleManager_UpdateRoleName_Call) Return(role roles.Role, err error) *RoleManager_UpdateRoleName_Call { - _c.Call.Return(role, err) - return _c -} - -func (_c *RoleManager_UpdateRoleName_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, newRoleName string) (roles.Role, error)) *RoleManager_UpdateRoleName_Call { - _c.Call.Return(run) - return _c -} diff --git a/pkg/roles/provisionmanage.go b/pkg/roles/provisionmanage.go deleted file mode 100644 index 6efcc4d64..000000000 --- a/pkg/roles/provisionmanage.go +++ /dev/null @@ -1,653 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package roles - -import ( - "context" - "fmt" - "time" - - "github.com/absmach/magistrala" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/pkg/policies" -) - -var ( - errRemoveOptionalDeletePolicies = errors.New("failed to delete the additional requested policies") - errRemoveOptionalFilterDeletePolicies = errors.New("failed to filter delete the additional requested policies") - errRollbackRoles = errors.New("failed to rollback roles") -) - -type roleProvisionerManger interface { - RoleManager - Provisioner -} - -var _ roleProvisionerManger = (*ProvisionManageService)(nil) - -type ProvisionManageService struct { - entityType string - repo Repository - sidProvider magistrala.IDProvider - policy policies.Service - actions []Action - builtInRoles map[BuiltInRoleName][]Action -} - -func NewProvisionManageService(entityType string, repo Repository, policy policies.Service, sidProvider magistrala.IDProvider, actions []Action, builtInRoles map[BuiltInRoleName][]Action) (ProvisionManageService, error) { - rm := ProvisionManageService{ - entityType: entityType, - repo: repo, - sidProvider: sidProvider, - policy: policy, - actions: actions, - builtInRoles: builtInRoles, - } - return rm, nil -} - -func (pms ProvisionManageService) BuiltInRoleActions(name BuiltInRoleName) ([]Action, error) { - actions, ok := pms.builtInRoles[name] - if !ok { - return nil, errors.Wrap(svcerr.ErrNotFound, fmt.Errorf("role %s not found", name)) - } - return actions, nil -} - -func toRolesActions(actions []string) []Action { - roActions := []Action{} - for _, action := range actions { - roActions = append(roActions, Action(action)) - } - return roActions -} - -func roleActionsToString(roActions []Action) []string { - actions := []string{} - for _, roAction := range roActions { - actions = append(actions, roAction.String()) - } - return actions -} - -func roleMembersToString(roMems []Member) []string { - mems := []string{} - for _, roMem := range roMems { - mems = append(mems, roMem.String()) - } - return mems -} - -func (r ProvisionManageService) isActionAllowed(action Action) bool { - for _, cap := range r.actions { - if cap == action { - return true - } - } - return false -} - -func (r ProvisionManageService) validateActions(actions []Action) error { - for _, ac := range actions { - action := Action(ac) - if !r.isActionAllowed(action) { - return errors.Wrap(svcerr.ErrMalformedEntity, fmt.Errorf("invalid action %s ", action)) - } - } - return nil -} - -func (r ProvisionManageService) RemoveEntitiesRoles(ctx context.Context, domainID, userID string, entityIDs []string, optionalFilterDeletePolicies []policies.Policy, optionalDeletePolicies []policies.Policy) error { - ears, emrs, err := r.repo.RetrieveEntitiesRolesActionsMembers(ctx, entityIDs) - if err != nil { - return err - } - - deletePolicies := []policies.Policy{} - for _, ear := range ears { - deletePolicies = append(deletePolicies, policies.Policy{ - Subject: ear.RoleID, - SubjectRelation: policies.MemberRelation, - SubjectType: policies.RoleType, - Relation: ear.Action, - ObjectType: r.entityType, - Object: ear.EntityID, - }) - } - for _, emr := range emrs { - deletePolicies = append(deletePolicies, policies.Policy{ - Subject: policies.EncodeDomainUserID(domainID, emr.MemberID), - SubjectType: policies.UserType, - Relation: policies.MemberRelation, - ObjectType: policies.RoleType, - Object: emr.RoleID, - }) - } - - if err := r.policy.DeletePolicies(ctx, deletePolicies); err != nil { - return errors.Wrap(svcerr.ErrDeletePolicies, err) - } - - if len(optionalDeletePolicies) > 1 { - if err := r.policy.DeletePolicies(ctx, optionalDeletePolicies); err != nil { - return errors.Wrap(errRemoveOptionalDeletePolicies, err) - } - } - - for _, optionalFilterDeletePolicy := range optionalFilterDeletePolicies { - if err := r.policy.DeletePolicyFilter(ctx, optionalFilterDeletePolicy); err != nil { - return errors.Wrap(errRemoveOptionalFilterDeletePolicies, err) - } - } - return nil -} - -func (r ProvisionManageService) AddNewEntitiesRoles(ctx context.Context, domainID, userID string, entityIDs []string, optionalEntityPolicies []policies.Policy, newBuiltInRoleMembers map[BuiltInRoleName][]Member) (retRolesProvision []RoleProvision, retErr error) { - var newRolesProvision []RoleProvision - p := []policies.Policy{} - - for _, entityID := range entityIDs { - for defaultRole, defaultRoleMembers := range newBuiltInRoleMembers { - actions, ok := r.builtInRoles[defaultRole] - if !ok { - return []RoleProvision{}, fmt.Errorf("default role %s not found in in-built roles", defaultRole) - } - - sid, err := r.sidProvider.ID() - if err != nil { - return []RoleProvision{}, errors.Wrap(svcerr.ErrCreateEntity, err) - } - - id := r.entityType + "_" + sid - if err := r.validateActions(actions); err != nil { - return []RoleProvision{}, errors.Wrap(svcerr.ErrMalformedEntity, err) - } - - members := roleMembersToString(defaultRoleMembers) - caps := roleActionsToString(actions) - - newRolesProvision = append(newRolesProvision, RoleProvision{ - Role: Role{ - ID: id, - Name: defaultRole.String(), - EntityID: entityID, - CreatedAt: time.Now().UTC(), - CreatedBy: userID, - }, - OptionalActions: caps, - OptionalMembers: members, - }) - - for _, cap := range caps { - p = append(p, policies.Policy{ - SubjectType: policies.RoleType, - SubjectRelation: policies.MemberRelation, - Subject: id, - Relation: cap, - Object: entityID, - ObjectType: r.entityType, - }) - } - - for _, member := range members { - p = append(p, policies.Policy{ - SubjectType: policies.UserType, - Subject: policies.EncodeDomainUserID(domainID, member), - Relation: policies.MemberRelation, - Object: id, - ObjectType: policies.RoleType, - }) - } - } - } - p = append(p, optionalEntityPolicies...) - - if len(p) > 0 { - if err := r.policy.AddPolicies(ctx, p); err != nil { - return []RoleProvision{}, errors.Wrap(svcerr.ErrCreateEntity, err) - } - defer func() { - if retErr != nil { - if errRollBack := r.policy.DeletePolicies(ctx, p); errRollBack != nil { - retErr = errors.Wrap(retErr, errors.Wrap(errRollbackRoles, errRollBack)) - } - } - }() - } - - rp, err := r.repo.AddRoles(ctx, newRolesProvision) - if err != nil { - return []RoleProvision{}, errors.Wrap(svcerr.ErrCreateEntity, err) - } - - return rp, nil -} - -func (r ProvisionManageService) AddRole(ctx context.Context, session authn.Session, entityID string, roleName string, optionalActions []string, optionalMembers []string) (retRoleProvision RoleProvision, retErr error) { - sid, err := r.sidProvider.ID() - if err != nil { - return RoleProvision{}, errors.Wrap(svcerr.ErrCreateEntity, err) - } - - id := r.entityType + "_" + sid - - if err := r.validateActions(toRolesActions(optionalActions)); err != nil { - return RoleProvision{}, errors.Wrap(svcerr.ErrMalformedEntity, err) - } - - newRoleProvisions := []RoleProvision{ - { - Role: Role{ - ID: id, - Name: roleName, - EntityID: entityID, - CreatedAt: time.Now().UTC(), - CreatedBy: session.UserID, - }, - OptionalActions: optionalActions, - OptionalMembers: optionalMembers, - }, - } - prs := []policies.Policy{} - - for _, cap := range optionalActions { - prs = append(prs, policies.Policy{ - SubjectType: policies.RoleType, - SubjectRelation: policies.MemberRelation, - Subject: id, - Relation: cap, - Object: entityID, - ObjectType: r.entityType, - }) - } - - for _, member := range optionalMembers { - prs = append(prs, policies.Policy{ - SubjectType: policies.UserType, - Subject: policies.EncodeDomainUserID(session.DomainID, member), - Relation: policies.MemberRelation, - Object: id, - ObjectType: policies.RoleType, - }) - } - - if len(prs) > 0 { - if err := r.policy.AddPolicies(ctx, prs); err != nil { - return RoleProvision{}, errors.Wrap(svcerr.ErrCreateEntity, err) - } - - defer func() { - if retErr != nil { - if errRollBack := r.policy.DeletePolicies(ctx, prs); errRollBack != nil { - retErr = errors.Wrap(retErr, errors.Wrap(errRollbackRoles, errRollBack)) - } - } - }() - } - - rp, err := r.repo.AddRoles(ctx, newRoleProvisions) - if err != nil { - return RoleProvision{}, errors.Wrap(svcerr.ErrCreateEntity, err) - } - - if len(rp) == 0 { - return RoleProvision{}, svcerr.ErrCreateEntity - } - - return rp[0], nil -} - -func (r ProvisionManageService) RemoveRole(ctx context.Context, session authn.Session, entityID, roleID string) error { - ro, err := r.repo.RetrieveEntityRole(ctx, entityID, roleID) - if err != nil { - return errors.Wrap(svcerr.ErrRemoveEntity, err) - } - req := policies.Policy{ - SubjectType: policies.RoleType, - Subject: ro.ID, - } - if err := r.policy.DeletePolicyFilter(ctx, req); err != nil { - return errors.Wrap(svcerr.ErrRemoveEntity, err) - } - - if err := r.repo.RemoveRoles(ctx, []string{ro.ID}); err != nil { - return errors.Wrap(svcerr.ErrRemoveEntity, err) - } - return nil -} - -func (r ProvisionManageService) UpdateRoleName(ctx context.Context, session authn.Session, entityID, roleID, newRoleName string) (Role, error) { - ro, err := r.repo.RetrieveEntityRole(ctx, entityID, roleID) - if err != nil { - return Role{}, errors.Wrap(svcerr.ErrUpdateEntity, err) - } - ro, err = r.repo.UpdateRole(ctx, Role{ - ID: ro.ID, - EntityID: entityID, - Name: newRoleName, - UpdatedBy: session.UserID, - UpdatedAt: time.Now().UTC(), - }) - if err != nil { - return Role{}, errors.Wrap(svcerr.ErrUpdateEntity, err) - } - return ro, nil -} - -func (r ProvisionManageService) RetrieveRole(ctx context.Context, session authn.Session, entityID, roleID string) (Role, error) { - ro, err := r.repo.RetrieveEntityRole(ctx, entityID, roleID) - if err != nil { - return Role{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - return ro, nil -} - -func (r ProvisionManageService) RetrieveAllRoles(ctx context.Context, session authn.Session, entityID string, limit, offset uint64) (RolePage, error) { - ros, err := r.repo.RetrieveAllRoles(ctx, entityID, limit, offset) - if err != nil { - return RolePage{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - return ros, nil -} - -func (r ProvisionManageService) ListAvailableActions(ctx context.Context, session authn.Session) ([]string, error) { - acts := []string{} - for _, a := range r.actions { - acts = append(acts, string(a)) - } - return acts, nil -} - -func (r ProvisionManageService) RoleAddActions(ctx context.Context, session authn.Session, entityID, roleID string, actions []string) (retActs []string, retErr error) { - ro, err := r.repo.RetrieveEntityRole(ctx, entityID, roleID) - if err != nil { - return []string{}, errors.Wrap(svcerr.ErrUpdateEntity, err) - } - - if len(actions) == 0 { - return []string{}, svcerr.ErrMalformedEntity - } - - if err := r.validateActions(toRolesActions(actions)); err != nil { - return []string{}, errors.Wrap(svcerr.ErrMalformedEntity, err) - } - - prs := []policies.Policy{} - for _, cap := range actions { - prs = append(prs, policies.Policy{ - SubjectType: policies.RoleType, - SubjectRelation: policies.MemberRelation, - Subject: ro.ID, - Relation: cap, - Object: entityID, - ObjectType: r.entityType, - }) - } - - if err := r.policy.AddPolicies(ctx, prs); err != nil { - return []string{}, errors.Wrap(svcerr.ErrAddPolicies, err) - } - - defer func() { - if retErr != nil { - if errRollBack := r.policy.DeletePolicies(ctx, prs); errRollBack != nil { - retErr = errors.Wrap(retErr, errors.Wrap(errRollbackRoles, errRollBack)) - } - } - }() - - ro.UpdatedAt = time.Now().UTC() - ro.UpdatedBy = session.UserID - - resActs, err := r.repo.RoleAddActions(ctx, ro, actions) - if err != nil { - return []string{}, errors.Wrap(svcerr.ErrCreateEntity, err) - } - return resActs, nil -} - -func (r ProvisionManageService) RoleListActions(ctx context.Context, session authn.Session, entityID, roleID string) ([]string, error) { - ro, err := r.repo.RetrieveEntityRole(ctx, entityID, roleID) - if err != nil { - return []string{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - - acts, err := r.repo.RoleListActions(ctx, ro.ID) - if err != nil { - return []string{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - return acts, nil -} - -func (r ProvisionManageService) RoleCheckActionsExists(ctx context.Context, session authn.Session, entityID, roleID string, actions []string) (bool, error) { - ro, err := r.repo.RetrieveEntityRole(ctx, entityID, roleID) - if err != nil { - return false, errors.Wrap(svcerr.ErrViewEntity, err) - } - - result, err := r.repo.RoleCheckActionsExists(ctx, ro.ID, actions) - if err != nil { - return true, errors.Wrap(svcerr.ErrViewEntity, err) - } - return result, nil -} - -func (r ProvisionManageService) RoleRemoveActions(ctx context.Context, session authn.Session, entityID, roleID string, actions []string) (err error) { - ro, err := r.repo.RetrieveEntityRole(ctx, entityID, roleID) - if err != nil { - return errors.Wrap(svcerr.ErrRemoveEntity, err) - } - - if len(actions) == 0 { - return svcerr.ErrMalformedEntity - } - - prs := []policies.Policy{} - for _, op := range actions { - prs = append(prs, policies.Policy{ - SubjectType: policies.RoleType, - SubjectRelation: policies.MemberRelation, - Subject: ro.ID, - Relation: op, - Object: entityID, - ObjectType: r.entityType, - }) - } - - if err := r.policy.DeletePolicies(ctx, prs); err != nil { - return errors.Wrap(svcerr.ErrDeletePolicies, err) - } - ro.UpdatedAt = time.Now().UTC() - ro.UpdatedBy = session.UserID - if err := r.repo.RoleRemoveActions(ctx, ro, actions); err != nil { - return errors.Wrap(svcerr.ErrRemoveEntity, err) - } - return nil -} - -func (r ProvisionManageService) RoleRemoveAllActions(ctx context.Context, session authn.Session, entityID, roleID string) error { - ro, err := r.repo.RetrieveEntityRole(ctx, entityID, roleID) - if err != nil { - return errors.Wrap(svcerr.ErrRemoveEntity, err) - } - - prs := policies.Policy{ - SubjectType: policies.RoleType, - Subject: ro.ID, - } - - if err := r.policy.DeletePolicyFilter(ctx, prs); err != nil { - return errors.Wrap(svcerr.ErrDeletePolicies, err) - } - - ro.UpdatedAt = time.Now().UTC() - ro.UpdatedBy = session.UserID - - if err := r.repo.RoleRemoveAllActions(ctx, ro); err != nil { - return errors.Wrap(svcerr.ErrRemoveEntity, err) - } - return nil -} - -func (r ProvisionManageService) RoleAddMembers(ctx context.Context, session authn.Session, entityID, roleID string, members []string) (retMems []string, retErr error) { - ro, err := r.repo.RetrieveEntityRole(ctx, entityID, roleID) - if err != nil { - return []string{}, errors.Wrap(svcerr.ErrUpdateEntity, err) - } - - if len(members) == 0 { - return []string{}, svcerr.ErrMalformedEntity - } - - prs := []policies.Policy{} - for _, mem := range members { - prs = append(prs, policies.Policy{ - SubjectType: policies.UserType, - Subject: policies.EncodeDomainUserID(session.DomainID, mem), - Relation: policies.MemberRelation, - Object: ro.ID, - ObjectType: policies.RoleType, - }) - } - - if err := r.policy.AddPolicies(ctx, prs); err != nil { - return []string{}, errors.Wrap(svcerr.ErrAddPolicies, err) - } - - defer func() { - if retErr != nil { - if errRollBack := r.policy.DeletePolicies(ctx, prs); errRollBack != nil { - retErr = errors.Wrap(retErr, errors.Wrap(errRollbackRoles, errRollBack)) - } - } - }() - - ro.UpdatedAt = time.Now().UTC() - ro.UpdatedBy = session.UserID - - mems, err := r.repo.RoleAddMembers(ctx, ro, members) - if err != nil { - return []string{}, errors.Wrap(svcerr.ErrCreateEntity, err) - } - return mems, nil -} - -func (r ProvisionManageService) RoleListMembers(ctx context.Context, session authn.Session, entityID, roleID string, limit, offset uint64) (MembersPage, error) { - ro, err := r.repo.RetrieveEntityRole(ctx, entityID, roleID) - if err != nil { - return MembersPage{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - - mp, err := r.repo.RoleListMembers(ctx, ro.ID, limit, offset) - if err != nil { - return MembersPage{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - return mp, nil -} - -func (r ProvisionManageService) RoleCheckMembersExists(ctx context.Context, session authn.Session, entityID, roleID string, members []string) (bool, error) { - ro, err := r.repo.RetrieveEntityRole(ctx, entityID, roleID) - if err != nil { - return false, errors.Wrap(svcerr.ErrViewEntity, err) - } - - result, err := r.repo.RoleCheckMembersExists(ctx, ro.ID, members) - if err != nil { - return true, errors.Wrap(svcerr.ErrViewEntity, err) - } - return result, nil -} - -func (r ProvisionManageService) RoleRemoveMembers(ctx context.Context, session authn.Session, entityID, roleID string, members []string) (err error) { - ro, err := r.repo.RetrieveEntityRole(ctx, entityID, roleID) - if err != nil { - return errors.Wrap(svcerr.ErrRemoveEntity, err) - } - - if len(members) == 0 { - return svcerr.ErrMalformedEntity - } - - prs := []policies.Policy{} - for _, mem := range members { - prs = append(prs, policies.Policy{ - SubjectType: policies.UserType, - Subject: policies.EncodeDomainUserID(session.DomainID, mem), - Relation: policies.MemberRelation, - Object: ro.ID, - ObjectType: policies.RoleType, - }) - } - - if err := r.policy.DeletePolicies(ctx, prs); err != nil { - return errors.Wrap(svcerr.ErrDeletePolicies, err) - } - - ro.UpdatedAt = time.Now().UTC() - ro.UpdatedBy = session.UserID - if err := r.repo.RoleRemoveMembers(ctx, ro, members); err != nil { - return errors.Wrap(svcerr.ErrRemoveEntity, err) - } - return nil -} - -func (r ProvisionManageService) RoleRemoveAllMembers(ctx context.Context, session authn.Session, entityID, roleID string) (err error) { - ro, err := r.repo.RetrieveEntityRole(ctx, entityID, roleID) - if err != nil { - return errors.Wrap(svcerr.ErrRemoveEntity, err) - } - - prs := policies.Policy{ - ObjectType: policies.RoleType, - Object: ro.ID, - SubjectType: policies.UserType, - } - - if err := r.policy.DeletePolicyFilter(ctx, prs); err != nil { - return errors.Wrap(svcerr.ErrDeletePolicies, err) - } - - ro.UpdatedAt = time.Now().UTC() - ro.UpdatedBy = session.UserID - - if err := r.repo.RoleRemoveAllMembers(ctx, ro); err != nil { - return errors.Wrap(svcerr.ErrRemoveEntity, err) - } - return nil -} - -func (r ProvisionManageService) ListEntityMembers(ctx context.Context, session authn.Session, entityID string, pageQuery MembersRolePageQuery) (MembersRolePage, error) { - mp, err := r.repo.ListEntityMembers(ctx, entityID, pageQuery) - if err != nil { - return MembersRolePage{}, err - } - return mp, nil -} - -func (r ProvisionManageService) RemoveEntityMembers(ctx context.Context, session authn.Session, entityID string, members []string) error { - if err := r.repo.RemoveEntityMembers(ctx, entityID, members); err != nil { - return err - } - return nil -} - -func (r ProvisionManageService) RemoveMemberFromAllRoles(ctx context.Context, session authn.Session, member string) (err error) { - if err := r.repo.RemoveMemberFromAllRoles(ctx, member); err != nil { - return errors.Wrap(svcerr.ErrRemoveEntity, err) - } - - prs := policies.Policy{ - ObjectType: policies.RoleType, - ObjectPrefix: r.entityType + "_", - SubjectType: policies.UserType, - } - - if err := r.policy.DeletePolicyFilter(ctx, prs); err != nil { - return errors.Wrap(svcerr.ErrDeletePolicies, err) - } - - return fmt.Errorf("not implemented") -} diff --git a/pkg/roles/repo/doc.go b/pkg/roles/repo/doc.go deleted file mode 100644 index 13812d96c..000000000 --- a/pkg/roles/repo/doc.go +++ /dev/null @@ -1,4 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package repo diff --git a/pkg/roles/repo/postgres/doc.go b/pkg/roles/repo/postgres/doc.go deleted file mode 100644 index 211292207..000000000 --- a/pkg/roles/repo/postgres/doc.go +++ /dev/null @@ -1,4 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package postgres diff --git a/pkg/roles/repo/postgres/init.go b/pkg/roles/repo/postgres/init.go deleted file mode 100644 index e2ebcab3e..000000000 --- a/pkg/roles/repo/postgres/init.go +++ /dev/null @@ -1,82 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package postgres - -import ( - "fmt" - - _ "github.com/jackc/pgx/v5/stdlib" // required for SQL access - migrate "github.com/rubenv/sql-migrate" -) - -// Migration of Auth service. -func Migration(rolesTableNamePrefix, entityTableName, entityIDColumnName string) (*migrate.MemoryMigrationSource, error) { - if entityTableName == "" || entityIDColumnName == "" { - return nil, fmt.Errorf("invalid entity Table Name or column name") - } - - return &migrate.MemoryMigrationSource{ - Migrations: []*migrate.Migration{ - { - Id: fmt.Sprintf("%s_roles_1", rolesTableNamePrefix), - Up: []string{ - fmt.Sprintf(`CREATE TABLE IF NOT EXISTS %s_roles ( - id VARCHAR(254) NOT NULL PRIMARY KEY, - name varchar(200) NOT NULL, - entity_id VARCHAR(36) NOT NULL, - created_at TIMESTAMP, - updated_at TIMESTAMP, - updated_by VARCHAR(254), - created_by VARCHAR(254), - CONSTRAINT %s_roles_unique_role_name_entity_id_constraint UNIQUE (name, entity_id), - CONSTRAINT %s_roles_fk_entity_id FOREIGN KEY(entity_id) REFERENCES %s(%s) ON DELETE CASCADE - );`, rolesTableNamePrefix, rolesTableNamePrefix, rolesTableNamePrefix, entityTableName, entityIDColumnName), - - fmt.Sprintf(`CREATE TABLE IF NOT EXISTS %s_role_actions ( - role_id VARCHAR(254) NOT NULL, - action VARCHAR(254) NOT NULL, - CONSTRAINT %s_role_actions_unique_role_action_constraint UNIQUE (role_id, action), - CONSTRAINT %s_role_actions_fk_roles_id FOREIGN KEY(role_id) REFERENCES %s_roles(id) ON DELETE CASCADE - );`, rolesTableNamePrefix, rolesTableNamePrefix, rolesTableNamePrefix, rolesTableNamePrefix), - - fmt.Sprintf(`CREATE TABLE IF NOT EXISTS %s_role_members ( - role_id VARCHAR(254) NOT NULL, - member_id VARCHAR(254) NOT NULL, - entity_id VARCHAR(36) NOT NULL, - CONSTRAINT %s_role_members_unique_role_member_constraint UNIQUE (role_id, member_id), - CONSTRAINT %s_role_members_unique_entity_member_constraint UNIQUE (member_id, entity_id), - CONSTRAINT %s_role_members_fk_roles_id FOREIGN KEY(role_id) REFERENCES %s_roles(id) ON DELETE CASCADE - );`, rolesTableNamePrefix, rolesTableNamePrefix, rolesTableNamePrefix, rolesTableNamePrefix, rolesTableNamePrefix), - }, - Down: []string{ - fmt.Sprintf(`DROP TABLE IF EXISTS %s_roles`, rolesTableNamePrefix), - fmt.Sprintf(`DROP TABLE IF EXISTS %s_role_actions`, rolesTableNamePrefix), - fmt.Sprintf(`DROP TABLE IF EXISTS %s_role_members`, rolesTableNamePrefix), - }, - }, - { - Id: fmt.Sprintf("%s_roles_2", rolesTableNamePrefix), - Up: []string{ - fmt.Sprintf(`ALTER TABLE %s_roles ALTER COLUMN created_at TYPE TIMESTAMPTZ;`, rolesTableNamePrefix), - fmt.Sprintf(`ALTER TABLE %s_roles ALTER COLUMN updated_at TYPE TIMESTAMPTZ;`, rolesTableNamePrefix), - }, - Down: []string{ - fmt.Sprintf(`ALTER TABLE %s_roles ALTER COLUMN created_at TYPE TIMESTAMP;`, rolesTableNamePrefix), - fmt.Sprintf(`ALTER TABLE %s_roles ALTER COLUMN updated_at TYPE TIMESTAMP;`, rolesTableNamePrefix), - }, - }, - { - Id: fmt.Sprintf("%s_roles_3", rolesTableNamePrefix), - Up: []string{ - fmt.Sprintf(`CREATE INDEX IF NOT EXISTS idx_%s_role_members_member_id ON %s_role_members(member_id);`, rolesTableNamePrefix, rolesTableNamePrefix), - fmt.Sprintf(`CREATE INDEX IF NOT EXISTS idx_%s_role_actions_action ON %s_role_actions(action text_pattern_ops);`, rolesTableNamePrefix, rolesTableNamePrefix), - }, - Down: []string{ - fmt.Sprintf(`DROP INDEX IF EXISTS idx_%s_role_members_member_id;`, rolesTableNamePrefix), - fmt.Sprintf(`DROP INDEX IF EXISTS idx_%s_role_actions_action;`, rolesTableNamePrefix), - }, - }, - }, - }, nil -} diff --git a/pkg/roles/repo/postgres/roles.go b/pkg/roles/repo/postgres/roles.go deleted file mode 100644 index 76e54bf10..000000000 --- a/pkg/roles/repo/postgres/roles.go +++ /dev/null @@ -1,1493 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package postgres - -import ( - "context" - "database/sql" - "encoding/json" - "fmt" - "strings" - "time" - - api "github.com/absmach/magistrala/api/http" - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - "github.com/absmach/magistrala/pkg/policies" - "github.com/absmach/magistrala/pkg/postgres" - "github.com/absmach/magistrala/pkg/roles" - "github.com/lib/pq" -) - -var _ roles.Repository = (*Repository)(nil) - -type Repository struct { - db postgres.Database - tableNamePrefix string - entityTableName string - entityIDColumnName string - membersListBaseQuery string -} - -// NewRepository instantiates a PostgreSQL -// implementation of Roles repository. -func NewRepository(db postgres.Database, entityType, tableNamePrefix, entityTableName, entityIDColumnName string) Repository { - var membersListBaseQuery string - - switch entityType { - case policies.ChannelType: - membersListBaseQuery = channelMembersListBaseQuery() - case policies.ClientType: - membersListBaseQuery = clientMembersListBaseQuery() - case policies.GroupType: - membersListBaseQuery = groupMembersListBaseQuery() - case policies.DomainType: - membersListBaseQuery = domainMembersListBaseQuery() - case policies.RulesType: - membersListBaseQuery = rulesMembersListBaseQuery() - case policies.ReportsType: - membersListBaseQuery = reportsMembersListBaseQuery() - } - - return Repository{ - db: db, - tableNamePrefix: tableNamePrefix, - entityTableName: entityTableName, - entityIDColumnName: entityIDColumnName, - membersListBaseQuery: membersListBaseQuery, - } -} - -type dbPage struct { - ID string `db:"id"` - Name string `db:"name"` - EntityID string `db:"entity_id"` - RoleID string `db:"role_id"` - Limit uint64 `db:"limit"` - Offset uint64 `db:"offset"` -} -type dbRole struct { - ID string `db:"id"` - Name string `db:"name"` - EntityID string `db:"entity_id"` - CreatedBy *string `db:"created_by"` - CreatedAt sql.NullTime `db:"created_at"` - UpdatedBy *string `db:"updated_by"` - UpdatedAt sql.NullTime `db:"updated_at"` -} - -type dbMemberRoles struct { - MemberID string `db:"member_id,omitempty"` - Roles json.RawMessage `db:"roles,omitempty"` -} - -type dbEntityActionRole struct { - EntityID string `db:"entity_id"` - Action string `db:"action"` - RoleID string `db:"role_id"` -} -type dbEntityMemberRole struct { - EntityID string `db:"entity_id"` - MemberID string `db:"member_id"` - RoleID string `db:"role_id"` -} - -func dbToEntityActionRole(dbs []dbEntityActionRole) []roles.EntityActionRole { - var r []roles.EntityActionRole - for _, d := range dbs { - r = append(r, roles.EntityActionRole{ - EntityID: d.EntityID, - Action: d.Action, - RoleID: d.RoleID, - }) - } - return r -} - -func dbToEntityMemberRole(dbs []dbEntityMemberRole) []roles.EntityMemberRole { - var r []roles.EntityMemberRole - for _, d := range dbs { - r = append(r, roles.EntityMemberRole{ - EntityID: d.EntityID, - MemberID: d.MemberID, - RoleID: d.RoleID, - }) - } - return r -} - -type dbRoleAction struct { - RoleID string `db:"role_id"` - Action string `db:"action"` -} - -type dbRoleMember struct { - RoleID string `db:"role_id"` - EntityID string `db:"entity_id"` - MemberID string `db:"member_id"` -} - -func toDBRoles(role roles.Role) dbRole { - var createdBy *string - if role.CreatedBy != "" { - createdBy = &role.CreatedBy - } - var createdAt sql.NullTime - if role.CreatedAt != (time.Time{}) && !role.CreatedAt.IsZero() { - createdAt = sql.NullTime{Time: role.CreatedAt, Valid: true} - } - - var updatedBy *string - if role.UpdatedBy != "" { - updatedBy = &role.UpdatedBy - } - var updatedAt sql.NullTime - if role.UpdatedAt != (time.Time{}) && !role.UpdatedAt.IsZero() { - updatedAt = sql.NullTime{Time: role.UpdatedAt, Valid: true} - } - - return dbRole{ - ID: role.ID, - Name: role.Name, - EntityID: role.EntityID, - CreatedBy: createdBy, - CreatedAt: createdAt, - UpdatedBy: updatedBy, - UpdatedAt: updatedAt, - } -} - -func toRole(r dbRole) roles.Role { - var createdBy string - if r.CreatedBy != nil { - createdBy = *r.CreatedBy - } - var createdAt time.Time - if r.CreatedAt.Valid { - createdAt = r.CreatedAt.Time.UTC() - } - - var updatedBy string - if r.UpdatedBy != nil { - updatedBy = *r.UpdatedBy - } - var updatedAt time.Time - if r.UpdatedAt.Valid { - updatedAt = r.UpdatedAt.Time.UTC() - } - - return roles.Role{ - ID: r.ID, - Name: r.Name, - EntityID: r.EntityID, - CreatedBy: createdBy, - CreatedAt: createdAt, - UpdatedBy: updatedBy, - UpdatedAt: updatedAt, - } -} - -func (repo *Repository) AddRoles(ctx context.Context, rps []roles.RoleProvision) ([]roles.RoleProvision, error) { - tx, err := repo.db.BeginTxx(ctx, nil) - if err != nil { - return []roles.RoleProvision{}, errors.Wrap(repoerr.ErrCreateEntity, err) - } - defer func() { - if err != nil { - if errRollback := tx.Rollback(); errRollback != nil { - err = errors.Wrap(errors.Wrap(apiutil.ErrRollbackTx, errRollback), err) - } - } - }() - - for _, rp := range rps { - q := fmt.Sprintf(`INSERT INTO %s_roles (id, name, entity_id, created_by, created_at, updated_by, updated_at) - VALUES (:id, :name, :entity_id, :created_by, :created_at, :updated_by, :updated_at);`, repo.tableNamePrefix) - - if _, err := tx.NamedExec(q, toDBRoles(rp.Role)); err != nil { - return []roles.RoleProvision{}, postgres.HandleError(repoerr.ErrCreateEntity, err) - } - - if len(rp.OptionalActions) > 0 { - capq := fmt.Sprintf(`INSERT INTO %s_role_actions (role_id, action) - VALUES (:role_id, :action) - RETURNING role_id, action`, repo.tableNamePrefix) - - rCaps := []dbRoleAction{} - for _, cap := range rp.OptionalActions { - rCaps = append(rCaps, dbRoleAction{ - RoleID: rp.ID, - Action: string(cap), - }) - } - if _, err := tx.NamedExec(capq, rCaps); err != nil { - return []roles.RoleProvision{}, postgres.HandleError(repoerr.ErrCreateEntity, err) - } - } - - if len(rp.OptionalMembers) > 0 { - mq := fmt.Sprintf(`INSERT INTO %s_role_members (role_id, entity_id, member_id) - VALUES (:role_id, :entity_id, :member_id) - RETURNING role_id, entity_id, member_id`, repo.tableNamePrefix) - - rMems := []dbRoleMember{} - for _, m := range rp.OptionalMembers { - rMems = append(rMems, dbRoleMember{ - RoleID: rp.ID, - MemberID: m, - EntityID: rp.EntityID, - }) - } - if _, err := tx.NamedExec(mq, rMems); err != nil { - return []roles.RoleProvision{}, postgres.HandleError(repoerr.ErrCreateEntity, err) - } - } - } - - if err := tx.Commit(); err != nil { - return []roles.RoleProvision{}, postgres.HandleError(repoerr.ErrCreateEntity, err) - } - - return rps, nil -} - -func (repo *Repository) RemoveRoles(ctx context.Context, roleIDs []string) error { - q := fmt.Sprintf("DELETE FROM %s_roles WHERE id = ANY(:role_ids) ;", repo.tableNamePrefix) - - params := map[string]any{ - "role_ids": roleIDs, - } - result, err := repo.db.NamedExecContext(ctx, q, params) - if err != nil { - return postgres.HandleError(repoerr.ErrRemoveEntity, err) - } - if rows, _ := result.RowsAffected(); rows == 0 { - return repoerr.ErrNotFound - } - - return nil -} - -// Update only role name, don't update ID. -func (repo *Repository) UpdateRole(ctx context.Context, role roles.Role) (roles.Role, error) { - var query []string - var upq string - if role.Name != "" { - query = append(query, "name = :name,") - } - - if len(query) > 0 { - upq = strings.Join(query, " ") - } - - q := fmt.Sprintf(`UPDATE %s_roles SET %s updated_at = :updated_at, updated_by = :updated_by - WHERE id = :id - RETURNING id, name, entity_id, created_by, created_at, updated_by, updated_at`, - repo.tableNamePrefix, upq) - - row, err := repo.db.NamedQueryContext(ctx, q, toDBRoles(role)) - if err != nil { - return roles.Role{}, postgres.HandleError(repoerr.ErrUpdateEntity, err) - } - defer row.Close() - - dbr := dbRole{} - if row.Next() { - if err := row.StructScan(&dbr); err != nil { - return roles.Role{}, errors.Wrap(repoerr.ErrUpdateEntity, err) - } - return toRole(dbr), nil - } - - return roles.Role{}, repoerr.ErrNotFound -} - -func (repo *Repository) RetrieveRole(ctx context.Context, roleID string) (roles.Role, error) { - q := fmt.Sprintf(`SELECT id, name, entity_id, created_by, created_at, updated_by, updated_at - FROM %s_roles WHERE id = :id`, repo.tableNamePrefix) - - dbr := dbRole{ - ID: roleID, - } - - rows, err := repo.db.NamedQueryContext(ctx, q, dbr) - if err != nil { - return roles.Role{}, postgres.HandleError(repoerr.ErrViewEntity, err) - } - defer rows.Close() - - dbr = dbRole{} - if rows.Next() { - if err = rows.StructScan(&dbr); err != nil { - return roles.Role{}, postgres.HandleError(repoerr.ErrViewEntity, err) - } - - return toRole(dbr), nil - } - - return roles.Role{}, repoerr.ErrNotFound -} - -func (repo *Repository) RetrieveEntityRole(ctx context.Context, entityID, roleID string) (roles.Role, error) { - q := fmt.Sprintf(`SELECT id, name, entity_id, created_by, created_at, updated_by, updated_at - FROM %s_roles WHERE entity_id = :entity_id and id = :id`, repo.tableNamePrefix) - - dbr := dbRole{ - EntityID: entityID, - ID: roleID, - } - - rows, err := repo.db.NamedQueryContext(ctx, q, dbr) - if err != nil { - return roles.Role{}, postgres.HandleError(repoerr.ErrViewEntity, err) - } - defer rows.Close() - - dbr = dbRole{} - if rows.Next() { - if err = rows.StructScan(&dbr); err != nil { - return roles.Role{}, postgres.HandleError(repoerr.ErrViewEntity, err) - } - - return toRole(dbr), nil - } - - return roles.Role{}, repoerr.ErrNotFound -} - -func (repo *Repository) RetrieveAllRoles(ctx context.Context, entityID string, limit, offset uint64) (roles.RolePage, error) { - q := fmt.Sprintf(`SELECT id, name, entity_id, created_by, created_at, updated_by, updated_at - FROM %s_roles WHERE entity_id = :entity_id ORDER BY created_at LIMIT :limit OFFSET :offset;`, repo.tableNamePrefix) - - dbp := dbPage{ - EntityID: entityID, - Limit: limit, - Offset: offset, - } - - rows, err := repo.db.NamedQueryContext(ctx, q, dbp) - if err != nil { - return roles.RolePage{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - defer rows.Close() - - items := []roles.Role{} - for rows.Next() { - dbr := dbRole{} - if err := rows.StructScan(&dbr); err != nil { - return roles.RolePage{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - - items = append(items, toRole(dbr)) - } - cq := fmt.Sprintf(`SELECT COUNT(*) FROM %s_roles WHERE entity_id = :entity_id`, repo.tableNamePrefix) - - total, err := postgres.Total(ctx, repo.db, cq, dbp) - if err != nil { - return roles.RolePage{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - - page := roles.RolePage{ - Roles: items, - Total: total, - Offset: offset, - Limit: limit, - } - - return page, nil -} - -func (repo *Repository) RoleAddActions(ctx context.Context, role roles.Role, actions []string) (caps []string, err error) { - tx, err := repo.db.BeginTxx(ctx, nil) - if err != nil { - return []string{}, errors.Wrap(repoerr.ErrCreateEntity, err) - } - defer func() { - if err != nil { - if errRollback := tx.Rollback(); errRollback != nil { - err = errors.Wrap(errors.Wrap(apiutil.ErrRollbackTx, errRollback), err) - } - } - }() - - capq := fmt.Sprintf(`INSERT INTO %s_role_actions (role_id, action) - VALUES (:role_id, :action) - RETURNING role_id, action`, repo.tableNamePrefix) - - rCaps := []dbRoleAction{} - for _, cap := range actions { - rCaps = append(rCaps, dbRoleAction{ - RoleID: role.ID, - Action: string(cap), - }) - } - if _, err := tx.NamedExecContext(ctx, capq, rCaps); err != nil { - return []string{}, postgres.HandleError(repoerr.ErrCreateEntity, err) - } - - upq := fmt.Sprintf(`UPDATE %s_roles SET updated_at = :updated_at, updated_by = :updated_by WHERE id = :id;`, repo.tableNamePrefix) - if _, err := tx.NamedExecContext(ctx, upq, toDBRoles(role)); err != nil { - return []string{}, postgres.HandleError(repoerr.ErrCreateEntity, err) - } - - if err := tx.Commit(); err != nil { - return []string{}, postgres.HandleError(repoerr.ErrCreateEntity, err) - } - - return actions, nil -} - -func (repo *Repository) RoleListActions(ctx context.Context, roleID string) ([]string, error) { - q := fmt.Sprintf(`SELECT role_id, action FROM %s_role_actions WHERE role_id = :role_id ;`, repo.tableNamePrefix) - - dbrcap := dbRoleAction{ - RoleID: roleID, - } - - rows, err := repo.db.NamedQueryContext(ctx, q, dbrcap) - if err != nil { - return []string{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - defer rows.Close() - - items := []string{} - for rows.Next() { - dbrcap = dbRoleAction{} - if err := rows.StructScan(&dbrcap); err != nil { - return []string{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - - items = append(items, dbrcap.Action) - } - return items, nil -} - -func (repo *Repository) RoleCheckActionsExists(ctx context.Context, roleID string, actions []string) (bool, error) { - q := fmt.Sprintf(`SELECT COUNT(*) FROM %s_role_actions WHERE role_id = $1 AND action = ANY($2)`, repo.tableNamePrefix) - - var count int - err := repo.db.QueryRowxContext(ctx, q, roleID, pq.Array(actions)).Scan(&count) - if err != nil { - return false, errors.Wrap(repoerr.ErrViewEntity, err) - } - - // Check if the count matches the number of actions provided - if count != len(actions) { - return false, nil - } - - return true, nil -} - -func (repo *Repository) RoleRemoveActions(ctx context.Context, role roles.Role, actions []string) (err error) { - tx, err := repo.db.BeginTxx(ctx, nil) - if err != nil { - return errors.Wrap(repoerr.ErrRemoveEntity, err) - } - defer func() { - if err != nil { - if errRollback := tx.Rollback(); errRollback != nil { - err = errors.Wrap(errors.Wrap(apiutil.ErrRollbackTx, errRollback), err) - } - } - }() - - q := fmt.Sprintf(`DELETE FROM %s_role_actions WHERE role_id = :role_id AND action = ANY(:actions)`, repo.tableNamePrefix) - - params := map[string]any{ - "role_id": role.ID, - "actions": actions, - } - - if _, err := tx.NamedExec(q, params); err != nil { - return errors.Wrap(repoerr.ErrRemoveEntity, err) - } - - upq := fmt.Sprintf(`UPDATE %s_roles SET updated_at = :updated_at, updated_by = :updated_by WHERE id = :id;`, repo.tableNamePrefix) - if _, err := tx.NamedExec(upq, toDBRoles(role)); err != nil { - return postgres.HandleError(repoerr.ErrRemoveEntity, err) - } - - if err := tx.Commit(); err != nil { - return errors.Wrap(repoerr.ErrRemoveEntity, err) - } - - return nil -} - -func (repo *Repository) RoleRemoveAllActions(ctx context.Context, role roles.Role) error { - tx, err := repo.db.BeginTxx(ctx, nil) - if err != nil { - return errors.Wrap(repoerr.ErrRemoveEntity, err) - } - defer func() { - if err != nil { - if errRollback := tx.Rollback(); errRollback != nil { - err = errors.Wrap(errors.Wrap(apiutil.ErrRollbackTx, errRollback), err) - } - } - }() - - q := fmt.Sprintf(`DELETE FROM %s_role_actions WHERE role_id = :role_id `, repo.tableNamePrefix) - - dbrcap := dbRoleAction{RoleID: role.ID} - - if _, err := tx.NamedExec(q, dbrcap); err != nil { - return errors.Wrap(repoerr.ErrRemoveEntity, err) - } - - upq := fmt.Sprintf(`UPDATE %s_roles SET updated_at = :updated_at, updated_by = :updated_by WHERE id = :id;`, repo.tableNamePrefix) - if _, err := tx.NamedExec(upq, toDBRoles(role)); err != nil { - return postgres.HandleError(repoerr.ErrRemoveEntity, err) - } - - if err := tx.Commit(); err != nil { - return errors.Wrap(repoerr.ErrRemoveEntity, err) - } - - return nil -} - -func (repo *Repository) RoleAddMembers(ctx context.Context, role roles.Role, members []string) ([]string, error) { - mq := fmt.Sprintf(`INSERT INTO %s_role_members (role_id, entity_id, member_id) - VALUES (:role_id, :entity_id, :member_id) - RETURNING role_id, :entity_id, member_id`, repo.tableNamePrefix) - - tx, err := repo.db.BeginTxx(ctx, nil) - if err != nil { - return []string{}, errors.Wrap(repoerr.ErrCreateEntity, err) - } - defer func() { - if err != nil { - if errRollback := tx.Rollback(); errRollback != nil { - err = errors.Wrap(errors.Wrap(apiutil.ErrRollbackTx, errRollback), err) - } - } - }() - - rMems := []dbRoleMember{} - for _, m := range members { - rMems = append(rMems, dbRoleMember{ - RoleID: role.ID, - EntityID: role.EntityID, - MemberID: m, - }) - } - if _, err := tx.NamedExec(mq, rMems); err != nil { - return []string{}, postgres.HandleError(repoerr.ErrCreateEntity, err) - } - - upq := fmt.Sprintf(`UPDATE %s_roles SET updated_at = :updated_at, updated_by = :updated_by WHERE id = :id;`, repo.tableNamePrefix) - if _, err := tx.NamedExec(upq, toDBRoles(role)); err != nil { - return []string{}, postgres.HandleError(repoerr.ErrCreateEntity, err) - } - - if err := tx.Commit(); err != nil { - return []string{}, postgres.HandleError(repoerr.ErrCreateEntity, err) - } - - return members, nil -} - -func (repo *Repository) RoleListMembers(ctx context.Context, roleID string, limit, offset uint64) (roles.MembersPage, error) { - q := fmt.Sprintf(`SELECT role_id, member_id FROM %s_role_members WHERE role_id = :role_id LIMIT :limit OFFSET :offset;`, repo.tableNamePrefix) - - dbp := dbPage{ - RoleID: roleID, - Limit: limit, - Offset: offset, - } - - rows, err := repo.db.NamedQueryContext(ctx, q, dbp) - if err != nil { - return roles.MembersPage{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - defer rows.Close() - - items := []string{} - for rows.Next() { - dbrmems := dbRoleMember{} - if err := rows.StructScan(&dbrmems); err != nil { - return roles.MembersPage{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - - items = append(items, dbrmems.MemberID) - } - - cq := fmt.Sprintf(`SELECT COUNT(*) FROM %s_role_members WHERE role_id = :role_id`, repo.tableNamePrefix) - - total, err := postgres.Total(ctx, repo.db, cq, dbp) - if err != nil { - return roles.MembersPage{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - - return roles.MembersPage{ - Members: items, - Total: total, - Offset: offset, - Limit: limit, - }, nil -} - -func (repo *Repository) RoleCheckMembersExists(ctx context.Context, roleID string, members []string) (bool, error) { - q := fmt.Sprintf(`SELECT COUNT(*) FROM %s_role_members WHERE role_id = $1 AND member_id = ANY($2)`, repo.tableNamePrefix) - - var count int - err := repo.db.QueryRowxContext(ctx, q, roleID, pq.Array(members)).Scan(&count) - if err != nil { - return false, errors.Wrap(repoerr.ErrViewEntity, err) - } - - if count != len(members) { - return false, nil - } - - return true, nil -} - -func (repo *Repository) RoleRemoveMembers(ctx context.Context, role roles.Role, members []string) (err error) { - tx, err := repo.db.BeginTxx(ctx, nil) - if err != nil { - return errors.Wrap(repoerr.ErrRemoveEntity, err) - } - defer func() { - if err != nil { - if errRollback := tx.Rollback(); errRollback != nil { - err = errors.Wrap(errors.Wrap(apiutil.ErrRollbackTx, errRollback), err) - } - } - }() - - q := fmt.Sprintf(`DELETE FROM %s_role_members WHERE role_id = :role_id AND member_id = ANY(:member_ids)`, repo.tableNamePrefix) - - params := map[string]any{ - "role_id": role.ID, - "member_ids": members, - } - - if _, err := tx.NamedExec(q, params); err != nil { - return errors.Wrap(repoerr.ErrRemoveEntity, err) - } - - upq := fmt.Sprintf(`UPDATE %s_roles SET updated_at = :updated_at, updated_by = :updated_by WHERE id = :id;`, repo.tableNamePrefix) - if _, err := tx.NamedExec(upq, toDBRoles(role)); err != nil { - return postgres.HandleError(repoerr.ErrRemoveEntity, err) - } - - if err := tx.Commit(); err != nil { - return errors.Wrap(repoerr.ErrRemoveEntity, err) - } - return nil -} - -func (repo *Repository) RoleRemoveAllMembers(ctx context.Context, role roles.Role) (err error) { - tx, err := repo.db.BeginTxx(ctx, nil) - if err != nil { - return errors.Wrap(repoerr.ErrRemoveEntity, err) - } - defer func() { - if err != nil { - if errRollback := tx.Rollback(); errRollback != nil { - err = errors.Wrap(errors.Wrap(apiutil.ErrRollbackTx, errRollback), err) - } - } - }() - q := fmt.Sprintf(`DELETE FROM %s_role_members WHERE role_id = :role_id `, repo.tableNamePrefix) - - dbrcap := dbRoleAction{RoleID: role.ID} - - if _, err := tx.NamedExec(q, dbrcap); err != nil { - return errors.Wrap(repoerr.ErrRemoveEntity, err) - } - - upq := fmt.Sprintf(`UPDATE %s_roles SET updated_at = :updated_at, updated_by = :updated_by WHERE id = :id;`, repo.tableNamePrefix) - if _, err := tx.NamedExec(upq, toDBRoles(role)); err != nil { - return postgres.HandleError(repoerr.ErrRemoveEntity, err) - } - - if err := tx.Commit(); err != nil { - return errors.Wrap(repoerr.ErrRemoveEntity, err) - } - return nil -} - -func (repo *Repository) RetrieveEntitiesRolesActionsMembers(ctx context.Context, entityIDs []string) ([]roles.EntityActionRole, []roles.EntityMemberRole, error) { - params := map[string]any{ - "entity_ids": entityIDs, - } - - clientsActionsRolesQuery := fmt.Sprintf(`SELECT e.%s AS entity_id , era."action" AS "action", er.id AS role_id - FROM %s e - JOIN %s_roles er ON er.entity_id = e.%s - JOIN %s_role_actions era ON era.role_id = er.id - WHERE e.%s = ANY(:entity_ids); - `, repo.entityIDColumnName, repo.entityTableName, repo.tableNamePrefix, repo.entityIDColumnName, repo.tableNamePrefix, repo.entityIDColumnName) - rows, err := repo.db.NamedQueryContext(ctx, clientsActionsRolesQuery, params) - if err != nil { - return []roles.EntityActionRole{}, []roles.EntityMemberRole{}, postgres.HandleError(repoerr.ErrViewEntity, err) - } - - defer rows.Close() - dbears := []dbEntityActionRole{} - for rows.Next() { - dbear := dbEntityActionRole{} - if err = rows.StructScan(&dbear); err != nil { - return []roles.EntityActionRole{}, []roles.EntityMemberRole{}, postgres.HandleError(repoerr.ErrViewEntity, err) - } - - dbears = append(dbears, dbear) - } - clientsMembersRolesQuery := fmt.Sprintf(`SELECT e.%s AS entity_id , erm.member_id AS member_id, er.id AS role_id - FROM %s e - JOIN %s_roles er ON er.entity_id = e.%s - JOIN %s_role_members erm ON erm.role_id = er.id - WHERE e.%s = ANY(:entity_ids); - `, repo.entityIDColumnName, repo.entityTableName, repo.tableNamePrefix, repo.entityIDColumnName, repo.tableNamePrefix, repo.entityIDColumnName) - - rows, err = repo.db.NamedQueryContext(ctx, clientsMembersRolesQuery, params) - if err != nil { - return []roles.EntityActionRole{}, []roles.EntityMemberRole{}, postgres.HandleError(repoerr.ErrViewEntity, err) - } - - defer rows.Close() - dbemrs := []dbEntityMemberRole{} - for rows.Next() { - dbemr := dbEntityMemberRole{} - if err = rows.StructScan(&dbemr); err != nil { - return []roles.EntityActionRole{}, []roles.EntityMemberRole{}, postgres.HandleError(repoerr.ErrViewEntity, err) - } - - dbemrs = append(dbemrs, dbemr) - } - return dbToEntityActionRole(dbears), dbToEntityMemberRole(dbemrs), nil -} - -func (repo *Repository) ListEntityMembers(ctx context.Context, entityID string, pageQuery roles.MembersRolePageQuery) (roles.MembersRolePage, error) { - dbPageQuery, err := toDBMembersRolePageQuery(pageQuery) - if err != nil { - return roles.MembersRolePage{}, err - } - dbPageQuery.EntityID = entityID - - entityMembersQuery := fmt.Sprintf(` - %s - SELECT - member_id, - roles - FROM - members - `, repo.membersListBaseQuery) - - entityMembersQuery = applyConditions(entityMembersQuery, pageQuery) - entityMembersQuery = applyOrdering(entityMembersQuery, pageQuery) - entityMembersQuery = applyLimitOffset(entityMembersQuery) - - rows, err := repo.db.NamedQueryContext(ctx, entityMembersQuery, dbPageQuery) - if err != nil { - return roles.MembersRolePage{}, postgres.HandleError(repoerr.ErrViewEntity, err) - } - - defer rows.Close() - mems := []roles.MemberRoles{} - for rows.Next() { - var dbmr dbMemberRoles - if err = rows.StructScan(&dbmr); err != nil { - return roles.MembersRolePage{}, postgres.HandleError(repoerr.ErrViewEntity, err) - } - - var roleActions []roles.MemberRoleActions - if err := json.Unmarshal(dbmr.Roles, &roleActions); err != nil { - return roles.MembersRolePage{}, fmt.Errorf("failed to unmarshal roles JSON: %w", err) - } - mems = append(mems, roles.MemberRoles{MemberID: dbmr.MemberID, Roles: roleActions}) - } - - entityMembersCountQuery := fmt.Sprintf(` - %s - SELECT - COUNT(*) - FROM - members - `, repo.membersListBaseQuery) - - entityMembersCountQuery = applyConditions(entityMembersCountQuery, pageQuery) - - total, err := postgres.Total(ctx, repo.db, entityMembersCountQuery, dbPageQuery) - if err != nil { - return roles.MembersRolePage{}, err - } - - return roles.MembersRolePage{ - Total: total, - Limit: pageQuery.Limit, - Offset: pageQuery.Offset, - Members: mems, - }, nil -} - -func (repo *Repository) RemoveEntityMembers(ctx context.Context, entityID string, memberIDs []string) error { - return nil -} - -func (repo *Repository) RemoveMemberFromAllRoles(ctx context.Context, memberID string) (err error) { - return nil -} - -func (repo *Repository) SetMemberListBaseQuery(query string) { - repo.membersListBaseQuery = query -} - -func applyConditions(query string, pageQuery roles.MembersRolePageQuery) string { - var whereClause []string - - if pageQuery.RoleID != "" { - whereClause = append(whereClause, " roles @> :role_id ") - } - if pageQuery.RoleName != "" { - whereClause = append(whereClause, " roles @> :role_name ") - } - if len(pageQuery.Actions) != 0 { - whereClause = append(whereClause, " roles @> :actions ") - } - if pageQuery.AccessType != "" { - whereClause = append(whereClause, " roles @> :access_type ") - } - if pageQuery.AccessProviderID != "" { - whereClause = append(whereClause, " roles @> :access_provider_id ") - } - - var whereCondition string - if len(whereClause) != 0 { - whereCondition = "WHERE " + strings.Join(whereClause, " AND ") - } - - return fmt.Sprintf(`%s - %s`, query, whereCondition) -} - -func applyOrdering(query string, pageQuery roles.MembersRolePageQuery) string { - switch pageQuery.Order { - case "access_provider_id", "role_name", "role_id", "access_type": - query = fmt.Sprintf("%s ORDER BY %s", query, pageQuery.Order) - if pageQuery.Dir == api.AscDir || pageQuery.Dir == api.DescDir { - query = fmt.Sprintf("%s %s", query, pageQuery.Dir) - } - } - return query -} - -func applyLimitOffset(query string) string { - return fmt.Sprintf(`%s - LIMIT :limit OFFSET :offset`, query) -} - -type dbMembersRolePageQuery struct { - Offset uint64 `db:"offset"` - Limit uint64 `db:"limit"` - OrderBy string `db:"order_by"` - Direction string `db:"dir"` - AccessProviderID json.RawMessage `db:"access_provider_id"` - RoleId json.RawMessage `db:"role_id"` - RoleName json.RawMessage `db:"role_name"` - Actions json.RawMessage `db:"actions"` - AccessType json.RawMessage `db:"access_type"` - EntityID string `db:"entity_id"` -} - -func toDBMembersRolePageQuery(pageQuery roles.MembersRolePageQuery) (dbMembersRolePageQuery, error) { - actions := []byte("{}") - if len(pageQuery.Actions) != 0 { - var err error - jactions := []struct { - Actions []string `json:"actions"` - }{ - { - Actions: pageQuery.Actions, - }, - } - actions, err = json.Marshal(jactions) - if err != nil { - return dbMembersRolePageQuery{}, err - } - } - - accessProviderID := []byte("{}") - if pageQuery.AccessProviderID != "" { - accessProviderID = []byte(fmt.Sprintf("[{\"access_provider_id\" : \"%s\"}]", pageQuery.AccessProviderID)) - } - - roleID := []byte("{}") - if pageQuery.RoleID != "" { - roleID = []byte(fmt.Sprintf("[{\"role_id\" : \"%s\"}]", pageQuery.RoleID)) - } - - roleName := []byte("{}") - if pageQuery.RoleName != "" { - roleName = []byte(fmt.Sprintf("[{\"role_name\" : \"%s\"}]", pageQuery.RoleName)) - } - - accessType := []byte("{}") - if pageQuery.AccessType != "" { - accessType = []byte(fmt.Sprintf("[{\"access_type\" : \"%s\"}]", pageQuery.AccessType)) - } - - return dbMembersRolePageQuery{ - Offset: pageQuery.Offset, - Limit: pageQuery.Limit, - OrderBy: pageQuery.Order, - Direction: pageQuery.Dir, - AccessProviderID: accessProviderID, - RoleId: roleID, - RoleName: roleName, - Actions: actions, - AccessType: accessType, - }, nil -} - -func rulesMembersListBaseQuery() string { - return ` -WITH ungrouped_members AS ( - SELECT - rr.id, - rr.name, - rrm.member_id, - ARRAY_AGG(DISTINCT rra."action") AS actions, - 'direct' AS access_type, - '' AS access_provider_id - FROM - rules_roles rr - JOIN rules_role_members rrm ON rrm.role_id = rr.id - JOIN rules_role_actions rra ON rra.role_id = rr.id - WHERE - rr.entity_id = :entity_id - GROUP BY - rr.id, - rrm.member_id -UNION - SELECT - dr.id, - dr.name, - drm.member_id, - ARRAY_AGG(DISTINCT agg_dra."action") AS actions, - 'domain' AS access_type, - d.id AS access_provider_id - FROM - rules r - JOIN domains d ON d.id = r.domain_id - JOIN domains_roles dr ON dr.entity_id = d.id - JOIN domains_role_members drm ON dr.id = drm.role_id - JOIN domains_role_actions dra ON dr.id = dra.role_id - JOIN domains_role_actions agg_dra ON agg_dra.role_id = dr.id - WHERE - r.id = :entity_id - AND dra."action" LIKE 'rule%' - GROUP BY - dr.id, - drm.member_id, - d.id -), -members AS ( - SELECT - um.member_id, - JSONB_AGG( - JSON_BUILD_OBJECT( - 'role_id', um.id, - 'role_name', um.name, - 'actions', um.actions, - 'access_type', um.access_type, - 'access_provider_id', um.access_provider_id - ) - ) AS roles - FROM - ungrouped_members um - GROUP BY - um.member_id -) - ` -} - -func reportsMembersListBaseQuery() string { - return ` -WITH ungrouped_members AS ( - SELECT - rr.id, - rr.name, - rrm.member_id, - ARRAY_AGG(DISTINCT rra."action") AS actions, - 'direct' AS access_type, - '' AS access_provider_id - FROM - reports_roles rr - JOIN reports_role_members rrm ON rrm.role_id = rr.id - JOIN reports_role_actions rra ON rra.role_id = rr.id - WHERE - rr.entity_id = :entity_id - GROUP BY - rr.id, - rrm.member_id -UNION - SELECT - dr.id, - dr.name, - drm.member_id, - ARRAY_AGG(DISTINCT agg_dra."action") AS actions, - 'domain' AS access_type, - d.id AS access_provider_id - FROM - report_config rc - JOIN domains d ON d.id = rc.domain_id - JOIN domains_roles dr ON dr.entity_id = d.id - JOIN domains_role_members drm ON dr.id = drm.role_id - JOIN domains_role_actions dra ON dr.id = dra.role_id - JOIN domains_role_actions agg_dra ON agg_dra.role_id = dr.id - WHERE - rc.id = :entity_id - AND dra."action" LIKE 'report%' - GROUP BY - dr.id, - drm.member_id, - d.id -), -members AS ( - SELECT - um.member_id, - JSONB_AGG( - JSON_BUILD_OBJECT( - 'role_id', um.id, - 'role_name', um.name, - 'actions', um.actions, - 'access_type', um.access_type, - 'access_provider_id', um.access_provider_id - ) - ) AS roles - FROM - ungrouped_members um - GROUP BY - um.member_id -) - ` -} - -func domainMembersListBaseQuery() string { - return ` -WITH ungrouped_members AS ( - SELECT - dr.id, - dr.name, - drm.member_id, - ARRAY_AGG(DISTINCT all_actions.action) AS actions, - 'direct' AS access_type, - '' AS access_provider_id - FROM - domains_role_members drm - JOIN domains_roles dr ON - dr.id = drm.role_id - JOIN domains_role_actions dra ON - dra.role_id = dr.id - JOIN domains_role_actions all_actions ON - all_actions.role_id = drm.role_id - WHERE - dr.entity_id = :entity_id - GROUP BY - dr.id, - drm.member_id -), -members AS ( - SELECT - um.member_id, - JSONB_AGG( - JSON_BUILD_OBJECT( - 'role_id', um.id, - 'role_name', um.name, - 'actions', um.actions, - 'access_type', um.access_type, - 'access_provider_id', um.access_provider_id - ) - ) AS roles - FROM - ungrouped_members um - GROUP BY - um.member_id -) - ` -} - -func groupMembersListBaseQuery() string { - return ` -WITH ungrouped_members AS ( - SELECT - gr."name", - gr.id, - grm.member_id, - ARRAY_AGG(DISTINCT agg_gra."action") AS actions, - CASE - WHEN g.id = :entity_id THEN 'direct' - ELSE 'indirect_group' - END AS access_type, - CASE - WHEN g.id = :entity_id THEN '' - ELSE g.id - END AS access_provider_id - FROM - "groups" g - JOIN - groups_roles gr ON - gr.entity_id = g.id - JOIN - groups_role_members grm ON - grm.role_id = gr.id - JOIN - groups_role_actions gra ON - gra.role_id = gr.id - JOIN - groups_role_actions agg_gra ON - agg_gra.role_id = gr.id - WHERE - g.path @> ( - SELECT - "path" - FROM - "groups" - WHERE - id = :entity_id - LIMIT 1 - ) - AND ( - g.id = :entity_id - OR gra."action" LIKE 'subgroup%' - ) -- -- If g.id = , it allows all actions. If g.id <> , it only allows actions matching 'subgroup%'. - GROUP BY - gr.id, - grm.member_id, - g.id -UNION - SELECT - dr."name", - dr.id, - drm.member_id, - ARRAY_AGG(DISTINCT agg_dra."action") AS actions, - 'domain' AS access_type, - d.id AS access_provider_id - FROM - "groups" g - JOIN - domains d ON - d.id = g.domain_id - JOIN - domains_roles dr ON - dr.entity_id = d.id - JOIN - domains_role_members drm ON - dr.id = drm.role_id - JOIN - domains_role_actions dra ON - dr.id = dra.role_id - JOIN - domains_role_actions agg_dra ON - agg_dra.role_id = dr.id - WHERE - g.id = :entity_id - AND - dra."action" LIKE 'group%' - GROUP BY - dr.id, - drm.member_id, - d.id -), -members AS ( - SELECT - um.member_id, - JSONB_AGG( - JSON_BUILD_OBJECT( - 'role_id', um.id, - 'role_name', um.name, - 'actions', um.actions, - 'access_type', um.access_type, - 'access_provider_id', um.access_provider_id - ) - ) AS roles - FROM - ungrouped_members um - GROUP BY - um.member_id -) - ` -} - -func clientMembersListBaseQuery() string { - return ` -WITH client_group AS ( - SELECT - c.id, - c.parent_group_id, - c.domain_id, - g."path" AS parent_group_path - FROM - clients c - LEFT JOIN - "groups" g ON - g.id = c.parent_group_id - WHERE - c.id = :entity_id - LIMIT 1 -), -ungrouped_members AS ( - SELECT - cr."name", - cr.id, - crm.member_id, - ARRAY_AGG(DISTINCT cra."action") AS actions, - 'direct' AS access_type, - '' AS access_provider_id, - CAST('' AS LTREE) AS access_provider_path - FROM - client_group cg - JOIN - clients_roles cr ON - cr.entity_id = cg.id - JOIN - clients_role_members crm ON - crm.role_id = cr.id - JOIN - clients_role_actions cra ON - cra.role_id = cr.id - GROUP BY - cr.id, - crm.member_id - UNION - SELECT - gr."name", - gr.id, - grm.member_id, - ARRAY_AGG(DISTINCT agg_gra."action") AS actions, - CASE - WHEN g.id = cg.parent_group_id THEN 'direct_group' - ELSE 'indirect_group' - END AS access_type, - g.id AS access_provider_id, - g.path AS access_provider_path - FROM - client_group cg - JOIN - "groups" g ON - g.PATH @> cg.parent_group_path - JOIN - groups_roles gr ON - g.id = gr.entity_id - JOIN - groups_role_members grm ON - grm.role_id = gr.id - JOIN - groups_role_actions gra ON - gra.role_id = gr.id - JOIN - groups_role_actions agg_gra ON - agg_gra.role_id = gr.id - WHERE - ( - gra."action" LIKE 'client%%' - AND g.id = cg.parent_group_id - ) - OR - ( - gra."action" LIKE 'subgroup_client%%' - AND g.id <> cg.parent_group_id - ) - GROUP BY - gr.id, - grm.member_id, - g.id, - cg.parent_group_id - UNION - SELECT - dr."name", - dr.id, - drm.member_id, - ARRAY_AGG(DISTINCT agg_dra."action") AS actions, - 'domain' AS access_type, - d.id AS access_provider_id, - CAST('' AS LTREE) AS access_provider_path - FROM - client_group cg - JOIN - domains d ON - d.id = cg.domain_id - JOIN - domains_roles dr ON - dr.entity_id = d.id - JOIN - domains_role_members drm ON - dr.id = drm.role_id - JOIN - domains_role_actions dra ON - dr.id = dra.role_id - JOIN - domains_role_actions agg_dra ON - agg_dra.role_id = dr.id - WHERE - dra."action" LIKE 'client%' - GROUP BY - dr.id, - drm.member_id, - d.id -), -members AS ( - SELECT - um.member_id, - JSONB_AGG( - JSON_BUILD_OBJECT( - 'role_id', um.id, - 'role_name', um.name, - 'actions', um.actions, - 'access_type', um.access_type, - 'access_provider_id', um.access_provider_id, - 'access_provider_path', um.access_provider_path - ) - ) AS roles - FROM - ungrouped_members um - GROUP BY - um.member_id -) - ` -} - -func channelMembersListBaseQuery() string { - return ` -WITH channel_group AS ( - SELECT - c.id, - c.parent_group_id, - c.domain_id, - g."path" AS parent_group_path - FROM - channels c - LEFT JOIN - "groups" g ON - g.id = c.parent_group_id - WHERE - c.id = :entity_id - LIMIT 1 -), -ungrouped_members AS ( - SELECT - cr."name", - cr.id, - crm.member_id, - ARRAY_AGG(DISTINCT cra."action") AS actions, - 'direct' AS access_type, - '' AS access_provider_id, - CAST('' AS LTREE) AS access_provider_path - FROM - channel_group cg - JOIN - channels_roles cr ON - cr.entity_id = cg.id - JOIN - channels_role_members crm ON - crm.role_id = cr.id - JOIN - channels_role_actions cra ON - cra.role_id = cr.id - GROUP BY - cr.id, - crm.member_id - UNION - SELECT - gr."name", - gr.id, - grm.member_id, - ARRAY_AGG(DISTINCT agg_gra."action") AS actions, - CASE - WHEN g.id = cg.parent_group_id THEN 'direct_group' - ELSE 'indirect_group' - END AS access_type, - g.id AS access_provider_id, - g.path AS access_provider_path - FROM - channel_group cg - JOIN - "groups" g ON - g.PATH @> cg.parent_group_path - JOIN - groups_roles gr ON - g.id = gr.entity_id - JOIN - groups_role_members grm ON - grm.role_id = gr.id - JOIN - groups_role_actions gra ON - gra.role_id = gr.id - JOIN - groups_role_actions agg_gra ON - agg_gra.role_id = gr.id - WHERE - ( - gra."action" LIKE 'channel%%' - AND g.id = cg.parent_group_id - ) - OR - ( - gra."action" LIKE 'subgroup_channel%%' - AND g.id <> cg.parent_group_id - ) - GROUP BY - gr.id, - grm.member_id, - g.id, - cg.parent_group_id - UNION - SELECT - dr."name", - dr.id, - drm.member_id, - ARRAY_AGG(DISTINCT agg_dra."action") AS actions, - 'domain' AS access_type, - d.id AS access_provider_id, - CAST('' AS LTREE) AS access_provider_path - FROM - channel_group cg - JOIN - domains d ON - d.id = cg.domain_id - JOIN - domains_roles dr ON - dr.entity_id = d.id - JOIN - domains_role_members drm ON - dr.id = drm.role_id - JOIN - domains_role_actions dra ON - dr.id = dra.role_id - JOIN - domains_role_actions agg_dra ON - agg_dra.role_id = dr.id - WHERE - dra."action" LIKE 'channel%' - GROUP BY - dr.id, - drm.member_id, - d.id -), -members AS ( - SELECT - um.member_id, - JSONB_AGG( - JSON_BUILD_OBJECT( - 'role_id', um.id, - 'role_name', um.name, - 'actions', um.actions, - 'access_type', um.access_type, - 'access_provider_id', um.access_provider_id, - 'access_provider_path', um.access_provider_path - ) - ) AS roles - FROM - ungrouped_members um - GROUP BY - um.member_id -) - ` -} diff --git a/pkg/roles/rolemanager/api/decoders.go b/pkg/roles/rolemanager/api/decoders.go deleted file mode 100644 index f0963fb9b..000000000 --- a/pkg/roles/rolemanager/api/decoders.go +++ /dev/null @@ -1,287 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package http - -import ( - "context" - "encoding/json" - "net/http" - "strings" - - api "github.com/absmach/magistrala/api/http" - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/pkg/errors" - "github.com/go-chi/chi/v5" -) - -type Decoder struct { - entityIDTemplate string -} - -func NewDecoder(entityIDTemplate string) Decoder { - return Decoder{entityIDTemplate} -} - -func (d Decoder) DecodeCreateRole(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - req := createRoleReq{ - token: apiutil.ExtractBearerToken(r), - entityID: chi.URLParam(r, d.entityIDTemplate), - } - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - return req, nil -} - -func (d Decoder) DecodeListRoles(_ context.Context, r *http.Request) (any, error) { - o, err := apiutil.ReadNumQuery[uint64](r, api.OffsetKey, api.DefOffset) - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - l, err := apiutil.ReadNumQuery[uint64](r, api.LimitKey, api.DefLimit) - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - req := listRolesReq{ - token: apiutil.ExtractBearerToken(r), - entityID: chi.URLParam(r, d.entityIDTemplate), - limit: l, - offset: o, - } - return req, nil -} - -func (d Decoder) DecodeListEntityMembers(_ context.Context, r *http.Request) (any, error) { - o, err := apiutil.ReadNumQuery[uint64](r, api.OffsetKey, api.DefOffset) - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - l, err := apiutil.ReadNumQuery[uint64](r, api.LimitKey, api.DefLimit) - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - order, err := apiutil.ReadStringQuery(r, api.OrderKey, "") - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - dir, err := apiutil.ReadStringQuery(r, api.LimitKey, "") - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - accessProviderID, err := apiutil.ReadStringQuery(r, api.AccessProviderIDKey, "") - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - accessType, err := apiutil.ReadStringQuery(r, api.AccessTypeKey, "") - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - roleId, err := apiutil.ReadStringQuery(r, api.RoleIDKey, "") - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - roleName, err := apiutil.ReadStringQuery(r, api.RoleNameKey, "") - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - allActions, err := apiutil.ReadStringQuery(r, api.ActionsKey, "") - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - actions := []string{} - - allActions = strings.TrimSpace(allActions) - if allActions != "" { - actions = strings.Split(allActions, ",") - } - - req := listEntityMembersReq{ - token: apiutil.ExtractBearerToken(r), - entityID: chi.URLParam(r, d.entityIDTemplate), - limit: l, - offset: o, - order: order, - dir: dir, - accessProviderID: accessProviderID, - roleId: roleId, - roleName: roleName, - actions: actions, - accessType: accessType, - } - return req, nil -} - -func (d Decoder) DecodeRemoveEntityMembers(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - - req := removeEntityMembersReq{ - token: apiutil.ExtractBearerToken(r), - entityID: chi.URLParam(r, d.entityIDTemplate), - } - - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - return req, nil -} - -func (d Decoder) DecodeViewRole(_ context.Context, r *http.Request) (any, error) { - req := viewRoleReq{ - token: apiutil.ExtractBearerToken(r), - entityID: chi.URLParam(r, d.entityIDTemplate), - roleID: chi.URLParam(r, "roleID"), - } - return req, nil -} - -func (d Decoder) DecodeUpdateRole(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - req := updateRoleReq{ - token: apiutil.ExtractBearerToken(r), - entityID: chi.URLParam(r, d.entityIDTemplate), - roleID: chi.URLParam(r, "roleID"), - } - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - return req, nil -} - -func (d Decoder) DecodeDeleteRole(_ context.Context, r *http.Request) (any, error) { - req := deleteRoleReq{ - token: apiutil.ExtractBearerToken(r), - entityID: chi.URLParam(r, d.entityIDTemplate), - roleID: chi.URLParam(r, "roleID"), - } - return req, nil -} - -func (d Decoder) DecodeListAvailableActions(_ context.Context, r *http.Request) (any, error) { - req := listAvailableActionsReq{ - token: apiutil.ExtractBearerToken(r), - } - return req, nil -} - -func (d Decoder) DecodeAddRoleActions(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - req := addRoleActionsReq{ - token: apiutil.ExtractBearerToken(r), - entityID: chi.URLParam(r, d.entityIDTemplate), - roleID: chi.URLParam(r, "roleID"), - } - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - return req, nil -} - -func (d Decoder) DecodeListRoleActions(_ context.Context, r *http.Request) (any, error) { - req := listRoleActionsReq{ - token: apiutil.ExtractBearerToken(r), - entityID: chi.URLParam(r, d.entityIDTemplate), - roleID: chi.URLParam(r, "roleID"), - } - return req, nil -} - -func (d Decoder) DecodeDeleteRoleActions(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - req := deleteRoleActionsReq{ - token: apiutil.ExtractBearerToken(r), - entityID: chi.URLParam(r, d.entityIDTemplate), - roleID: chi.URLParam(r, "roleID"), - } - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - return req, nil -} - -func (d Decoder) DecodeDeleteAllRoleActions(_ context.Context, r *http.Request) (any, error) { - req := deleteAllRoleActionsReq{ - token: apiutil.ExtractBearerToken(r), - entityID: chi.URLParam(r, d.entityIDTemplate), - roleID: chi.URLParam(r, "roleID"), - } - return req, nil -} - -func (d Decoder) DecodeAddRoleMembers(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - req := addRoleMembersReq{ - token: apiutil.ExtractBearerToken(r), - entityID: chi.URLParam(r, d.entityIDTemplate), - roleID: chi.URLParam(r, "roleID"), - } - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - return req, nil -} - -func (d Decoder) DecodeListRoleMembers(_ context.Context, r *http.Request) (any, error) { - o, err := apiutil.ReadNumQuery[uint64](r, api.OffsetKey, api.DefOffset) - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - l, err := apiutil.ReadNumQuery[uint64](r, api.LimitKey, api.DefLimit) - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - req := listRoleMembersReq{ - token: apiutil.ExtractBearerToken(r), - entityID: chi.URLParam(r, d.entityIDTemplate), - roleID: chi.URLParam(r, "roleID"), - limit: l, - offset: o, - } - return req, nil -} - -func (d Decoder) DecodeDeleteRoleMembers(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - req := deleteRoleMembersReq{ - token: apiutil.ExtractBearerToken(r), - entityID: chi.URLParam(r, d.entityIDTemplate), - roleID: chi.URLParam(r, "roleID"), - } - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - return req, nil -} - -func (d Decoder) DecodeDeleteAllRoleMembers(_ context.Context, r *http.Request) (any, error) { - req := deleteAllRoleMembersReq{ - token: apiutil.ExtractBearerToken(r), - entityID: chi.URLParam(r, d.entityIDTemplate), - roleID: chi.URLParam(r, "roleID"), - } - return req, nil -} diff --git a/pkg/roles/rolemanager/api/endpoints.go b/pkg/roles/rolemanager/api/endpoints.go deleted file mode 100644 index dbba08967..000000000 --- a/pkg/roles/rolemanager/api/endpoints.go +++ /dev/null @@ -1,341 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package http - -import ( - "context" - - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/pkg/roles" - "github.com/go-kit/kit/endpoint" -) - -func CreateRoleEndpoint(svc roles.RoleManager) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(createRoleReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - rp, err := svc.AddRole(ctx, session, req.entityID, req.RoleName, req.OptionalActions, req.OptionalMembers) - if err != nil { - return nil, err - } - return createRoleRes{RoleProvision: rp}, nil - } -} - -func ListRolesEndpoint(svc roles.RoleManager) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(listRolesReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - ros, err := svc.RetrieveAllRoles(ctx, session, req.entityID, req.limit, req.offset) - if err != nil { - return nil, err - } - return listRolesRes{RolePage: ros}, nil - } -} - -func ListEntityMembersEndpoint(svc roles.RoleManager) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(listEntityMembersReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - pageQuery := roles.MembersRolePageQuery{ - Offset: req.offset, - Limit: req.limit, - AccessProviderID: req.accessProviderID, - Order: req.order, - Dir: req.dir, - RoleID: req.roleId, - RoleName: req.roleName, - Actions: req.actions, - AccessType: req.accessType, - } - - mems, err := svc.ListEntityMembers(ctx, session, req.entityID, pageQuery) - if err != nil { - return nil, err - } - return listEntityMembersRes{mems}, nil - } -} - -func RemoveEntityMembersEndpoint(svc roles.RoleManager) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(removeEntityMembersReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - if err := svc.RemoveEntityMembers(ctx, session, req.entityID, req.MemberIDs); err != nil { - return nil, err - } - return deleteEntityMembersRes{}, nil - } -} - -func ViewRoleEndpoint(svc roles.RoleManager) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(viewRoleReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - ro, err := svc.RetrieveRole(ctx, session, req.entityID, req.roleID) - if err != nil { - return nil, err - } - return viewRoleRes{Role: ro}, nil - } -} - -func UpdateRoleEndpoint(svc roles.RoleManager) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(updateRoleReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - ro, err := svc.UpdateRoleName(ctx, session, req.entityID, req.roleID, req.Name) - if err != nil { - return nil, err - } - return updateRoleRes{Role: ro}, nil - } -} - -func DeleteRoleEndpoint(svc roles.RoleManager) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(deleteRoleReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - if err := svc.RemoveRole(ctx, session, req.entityID, req.roleID); err != nil { - return nil, err - } - return deleteRoleRes{}, nil - } -} - -func ListAvailableActionsEndpoint(svc roles.RoleManager) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(listAvailableActionsReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - acts, err := svc.ListAvailableActions(ctx, session) - if err != nil { - return nil, err - } - return listAvailableActionsRes{acts}, nil - } -} - -func AddRoleActionsEndpoint(svc roles.RoleManager) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(addRoleActionsReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - caps, err := svc.RoleAddActions(ctx, session, req.entityID, req.roleID, req.Actions) - if err != nil { - return nil, err - } - return addRoleActionsRes{Actions: caps}, nil - } -} - -func ListRoleActionsEndpoint(svc roles.RoleManager) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(listRoleActionsReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - caps, err := svc.RoleListActions(ctx, session, req.entityID, req.roleID) - if err != nil { - return nil, err - } - return listRoleActionsRes{Actions: caps}, nil - } -} - -func DeleteRoleActionsEndpoint(svc roles.RoleManager) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(deleteRoleActionsReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - if err := svc.RoleRemoveActions(ctx, session, req.entityID, req.roleID, req.Actions); err != nil { - return nil, err - } - return deleteRoleActionsRes{}, nil - } -} - -func DeleteAllRoleActionsEndpoint(svc roles.RoleManager) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(deleteAllRoleActionsReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - if err := svc.RoleRemoveAllActions(ctx, session, req.entityID, req.roleID); err != nil { - return nil, err - } - return deleteAllRoleActionsRes{}, nil - } -} - -func AddRoleMembersEndpoint(svc roles.RoleManager) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(addRoleMembersReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - members, err := svc.RoleAddMembers(ctx, session, req.entityID, req.roleID, req.Members) - if err != nil { - return nil, err - } - return addRoleMembersRes{members}, nil - } -} - -func ListRoleMembersEndpoint(svc roles.RoleManager) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(listRoleMembersReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - mp, err := svc.RoleListMembers(ctx, session, req.entityID, req.roleID, req.limit, req.offset) - if err != nil { - return nil, err - } - return listRoleMembersRes{mp}, nil - } -} - -func DeleteRoleMembersEndpoint(svc roles.RoleManager) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(deleteRoleMembersReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - if err := svc.RoleRemoveMembers(ctx, session, req.entityID, req.roleID, req.Members); err != nil { - return nil, err - } - return deleteRoleMembersRes{}, nil - } -} - -func DeleteAllRoleMembersEndpoint(svc roles.RoleManager) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(deleteAllRoleMembersReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - if err := svc.RoleRemoveAllMembers(ctx, session, req.entityID, req.roleID); err != nil { - return nil, err - } - return deleteAllRoleMemberRes{}, nil - } -} diff --git a/pkg/roles/rolemanager/api/requests.go b/pkg/roles/rolemanager/api/requests.go deleted file mode 100644 index 0a3a91514..000000000 --- a/pkg/roles/rolemanager/api/requests.go +++ /dev/null @@ -1,347 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package http - -import ( - api "github.com/absmach/magistrala/api/http" - apiutil "github.com/absmach/magistrala/api/http/util" -) - -type createRoleReq struct { - token string - entityID string - RoleName string `json:"role_name"` - OptionalActions []string `json:"optional_actions"` - OptionalMembers []string `json:"optional_members"` -} - -func (req createRoleReq) validate() error { - if req.token == "" { - return apiutil.ErrBearerToken - } - if err := api.ValidateUUID(req.entityID); err != nil { - return err - } - if len(req.RoleName) == 0 { - return apiutil.ErrMissingRoleName - } - if len(req.RoleName) > 200 { - return apiutil.ErrNameSize - } - - return nil -} - -type listRolesReq struct { - token string - entityID string - limit uint64 - offset uint64 -} - -func (req listRolesReq) validate() error { - if req.token == "" { - return apiutil.ErrBearerToken - } - if req.entityID == "" { - return apiutil.ErrMissingID - } - if req.limit > api.MaxLimitSize || req.limit < 1 { - return apiutil.ErrLimitSize - } - return nil -} - -type listEntityMembersReq struct { - token string - entityID string - limit uint64 - offset uint64 - dir string - order string - accessProviderID string - roleId string - roleName string - actions []string - accessType string -} - -func (req listEntityMembersReq) validate() error { - if req.token == "" { - return apiutil.ErrBearerToken - } - if req.entityID == "" { - return apiutil.ErrMissingID - } - if req.limit > api.MaxLimitSize || req.limit < 1 { - return apiutil.ErrLimitSize - } - return nil -} - -type removeEntityMembersReq struct { - token string - entityID string - MemberIDs []string `json:"member_ids"` -} - -func (req removeEntityMembersReq) validate() error { - if req.token == "" { - return apiutil.ErrBearerToken - } - if req.entityID == "" { - return apiutil.ErrMissingID - } - if len(req.MemberIDs) == 0 { - return apiutil.ErrMissingMemberIDs - } - return nil -} - -type viewRoleReq struct { - token string - entityID string - roleID string -} - -func (req viewRoleReq) validate() error { - if req.token == "" { - return apiutil.ErrBearerToken - } - if req.entityID == "" { - return apiutil.ErrMissingID - } - if req.roleID == "" { - return apiutil.ErrMissingRoleID - } - return nil -} - -type updateRoleReq struct { - token string - entityID string - roleID string - Name string `json:"name"` -} - -func (req updateRoleReq) validate() error { - if req.token == "" { - return apiutil.ErrBearerToken - } - if req.entityID == "" { - return apiutil.ErrMissingID - } - if req.roleID == "" { - return apiutil.ErrMissingRoleID - } - if req.Name == "" { - return apiutil.ErrMissingRoleName - } - return nil -} - -type deleteRoleReq struct { - token string - entityID string - roleID string -} - -func (req deleteRoleReq) validate() error { - if req.token == "" { - return apiutil.ErrBearerToken - } - if req.entityID == "" { - return apiutil.ErrMissingID - } - if req.roleID == "" { - return apiutil.ErrMissingRoleID - } - return nil -} - -type listAvailableActionsReq struct { - token string -} - -func (req listAvailableActionsReq) validate() error { - if req.token == "" { - return apiutil.ErrBearerToken - } - return nil -} - -type addRoleActionsReq struct { - token string - entityID string - roleID string - Actions []string `json:"actions"` -} - -func (req addRoleActionsReq) validate() error { - if req.token == "" { - return apiutil.ErrBearerToken - } - if req.entityID == "" { - return apiutil.ErrMissingID - } - if req.roleID == "" { - return apiutil.ErrMissingRoleID - } - - if len(req.Actions) == 0 { - return apiutil.ErrMissingPolicyEntityType - } - return nil -} - -type listRoleActionsReq struct { - token string - entityID string - roleID string -} - -func (req listRoleActionsReq) validate() error { - if req.token == "" { - return apiutil.ErrBearerToken - } - if req.entityID == "" { - return apiutil.ErrMissingID - } - if req.roleID == "" { - return apiutil.ErrMissingRoleID - } - return nil -} - -type deleteRoleActionsReq struct { - token string - entityID string - roleID string - Actions []string `json:"actions"` -} - -func (req deleteRoleActionsReq) validate() error { - if req.token == "" { - return apiutil.ErrBearerToken - } - if req.entityID == "" { - return apiutil.ErrMissingID - } - if req.roleID == "" { - return apiutil.ErrMissingRoleID - } - - if len(req.Actions) == 0 { - return apiutil.ErrMissingPolicyEntityType - } - return nil -} - -type deleteAllRoleActionsReq struct { - token string - entityID string - roleID string -} - -func (req deleteAllRoleActionsReq) validate() error { - if req.token == "" { - return apiutil.ErrBearerToken - } - if req.entityID == "" { - return apiutil.ErrMissingID - } - if req.roleID == "" { - return apiutil.ErrMissingRoleID - } - return nil -} - -type addRoleMembersReq struct { - token string - entityID string - roleID string - Members []string `json:"members"` -} - -func (req addRoleMembersReq) validate() error { - if req.token == "" { - return apiutil.ErrBearerToken - } - if req.entityID == "" { - return apiutil.ErrMissingID - } - if req.roleID == "" { - return apiutil.ErrMissingRoleID - } - if len(req.Members) == 0 { - return apiutil.ErrMissingRoleMembers - } - return nil -} - -type listRoleMembersReq struct { - token string - entityID string - roleID string - limit uint64 - offset uint64 -} - -func (req listRoleMembersReq) validate() error { - if req.token == "" { - return apiutil.ErrBearerToken - } - if req.entityID == "" { - return apiutil.ErrMissingID - } - if req.roleID == "" { - return apiutil.ErrMissingRoleID - } - if req.limit > api.MaxLimitSize || req.limit < 1 { - return apiutil.ErrLimitSize - } - return nil -} - -type deleteRoleMembersReq struct { - token string - entityID string - roleID string - Members []string `json:"members"` -} - -func (req deleteRoleMembersReq) validate() error { - if req.token == "" { - return apiutil.ErrBearerToken - } - if req.entityID == "" { - return apiutil.ErrMissingID - } - if req.roleID == "" { - return apiutil.ErrMissingRoleID - } - if len(req.Members) == 0 { - return apiutil.ErrMissingRoleMembers - } - return nil -} - -type deleteAllRoleMembersReq struct { - token string - entityID string - roleID string -} - -func (req deleteAllRoleMembersReq) validate() error { - if req.token == "" { - return apiutil.ErrBearerToken - } - if req.entityID == "" { - return apiutil.ErrMissingID - } - if req.roleID == "" { - return apiutil.ErrMissingRoleID - } - return nil -} diff --git a/pkg/roles/rolemanager/api/responses.go b/pkg/roles/rolemanager/api/responses.go deleted file mode 100644 index 72b07aa76..000000000 --- a/pkg/roles/rolemanager/api/responses.go +++ /dev/null @@ -1,272 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package http - -import ( - "net/http" - - "github.com/absmach/magistrala" - "github.com/absmach/magistrala/pkg/roles" -) - -var ( - _ magistrala.Response = (*createRoleRes)(nil) - _ magistrala.Response = (*listRolesRes)(nil) - _ magistrala.Response = (*viewRoleRes)(nil) - _ magistrala.Response = (*updateRoleRes)(nil) - _ magistrala.Response = (*deleteRoleRes)(nil) - _ magistrala.Response = (*listAvailableActionsRes)(nil) - _ magistrala.Response = (*addRoleActionsRes)(nil) - _ magistrala.Response = (*listRoleActionsRes)(nil) - _ magistrala.Response = (*deleteRoleActionsRes)(nil) - _ magistrala.Response = (*deleteAllRoleActionsRes)(nil) - _ magistrala.Response = (*addRoleMembersRes)(nil) - _ magistrala.Response = (*listRoleMembersRes)(nil) - _ magistrala.Response = (*deleteRoleMembersRes)(nil) - _ magistrala.Response = (*deleteAllRoleMemberRes)(nil) -) - -type createRoleRes struct { - roles.RoleProvision -} - -func (res createRoleRes) Code() int { - return http.StatusCreated -} - -func (res createRoleRes) Headers() map[string]string { - return map[string]string{} -} - -func (res createRoleRes) Empty() bool { - return false -} - -type listRolesRes struct { - roles.RolePage -} - -func (res listRolesRes) Code() int { - return http.StatusOK -} - -func (res listRolesRes) Headers() map[string]string { - return map[string]string{} -} - -func (res listRolesRes) Empty() bool { - return false -} - -type listEntityMembersRes struct { - roles.MembersRolePage -} - -func (res listEntityMembersRes) Code() int { - return http.StatusOK -} - -func (res listEntityMembersRes) Headers() map[string]string { - return map[string]string{} -} - -func (res listEntityMembersRes) Empty() bool { - return false -} - -type deleteEntityMembersRes struct{} - -func (res deleteEntityMembersRes) Code() int { - return http.StatusNoContent -} - -func (res deleteEntityMembersRes) Headers() map[string]string { - return map[string]string{} -} - -func (res deleteEntityMembersRes) Empty() bool { - return true -} - -type viewRoleRes struct { - roles.Role -} - -func (res viewRoleRes) Code() int { - return http.StatusOK -} - -func (res viewRoleRes) Headers() map[string]string { - return map[string]string{} -} - -func (res viewRoleRes) Empty() bool { - return false -} - -type updateRoleRes struct { - roles.Role -} - -func (res updateRoleRes) Code() int { - return http.StatusOK -} - -func (res updateRoleRes) Headers() map[string]string { - return map[string]string{} -} - -func (res updateRoleRes) Empty() bool { - return false -} - -type deleteRoleRes struct{} - -func (res deleteRoleRes) Code() int { - return http.StatusNoContent -} - -func (res deleteRoleRes) Headers() map[string]string { - return map[string]string{} -} - -func (res deleteRoleRes) Empty() bool { - return true -} - -type listAvailableActionsRes struct { - AvailableActions []string `json:"available_actions"` -} - -func (res listAvailableActionsRes) Code() int { - return http.StatusOK -} - -func (res listAvailableActionsRes) Headers() map[string]string { - return map[string]string{} -} - -func (res listAvailableActionsRes) Empty() bool { - return false -} - -type addRoleActionsRes struct { - Actions []string `json:"actions"` -} - -func (res addRoleActionsRes) Code() int { - return http.StatusOK -} - -func (res addRoleActionsRes) Headers() map[string]string { - return map[string]string{} -} - -func (res addRoleActionsRes) Empty() bool { - return false -} - -type listRoleActionsRes struct { - Actions []string `json:"actions"` -} - -func (res listRoleActionsRes) Code() int { - return http.StatusOK -} - -func (res listRoleActionsRes) Headers() map[string]string { - return map[string]string{} -} - -func (res listRoleActionsRes) Empty() bool { - return false -} - -type deleteRoleActionsRes struct{} - -func (res deleteRoleActionsRes) Code() int { - return http.StatusNoContent -} - -func (res deleteRoleActionsRes) Headers() map[string]string { - return map[string]string{} -} - -func (res deleteRoleActionsRes) Empty() bool { - return true -} - -type deleteAllRoleActionsRes struct{} - -func (res deleteAllRoleActionsRes) Code() int { - return http.StatusNoContent -} - -func (res deleteAllRoleActionsRes) Headers() map[string]string { - return map[string]string{} -} - -func (res deleteAllRoleActionsRes) Empty() bool { - return true -} - -type addRoleMembersRes struct { - Members []string `json:"members"` -} - -func (res addRoleMembersRes) Code() int { - return http.StatusOK -} - -func (res addRoleMembersRes) Headers() map[string]string { - return map[string]string{} -} - -func (res addRoleMembersRes) Empty() bool { - return false -} - -type listRoleMembersRes struct { - roles.MembersPage -} - -func (res listRoleMembersRes) Code() int { - return http.StatusOK -} - -func (res listRoleMembersRes) Headers() map[string]string { - return map[string]string{} -} - -func (res listRoleMembersRes) Empty() bool { - return false -} - -type deleteRoleMembersRes struct{} - -func (res deleteRoleMembersRes) Code() int { - return http.StatusNoContent -} - -func (res deleteRoleMembersRes) Headers() map[string]string { - return map[string]string{} -} - -func (res deleteRoleMembersRes) Empty() bool { - return true -} - -type deleteAllRoleMemberRes struct{} - -func (res deleteAllRoleMemberRes) Code() int { - return http.StatusNoContent -} - -func (res deleteAllRoleMemberRes) Headers() map[string]string { - return map[string]string{} -} - -func (res deleteAllRoleMemberRes) Empty() bool { - return true -} diff --git a/pkg/roles/rolemanager/api/router.go b/pkg/roles/rolemanager/api/router.go deleted file mode 100644 index 11aeebbfe..000000000 --- a/pkg/roles/rolemanager/api/router.go +++ /dev/null @@ -1,140 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package http - -import ( - api "github.com/absmach/magistrala/api/http" - "github.com/absmach/magistrala/pkg/roles" - "github.com/go-chi/chi/v5" - kithttp "github.com/go-kit/kit/transport/http" - "go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp" -) - -func EntityRoleMangerRouter(svc roles.RoleManager, d Decoder, r chi.Router, opts []kithttp.ServerOption) chi.Router { - r.Route("/roles", func(r chi.Router) { - r.Post("/", otelhttp.NewHandler(kithttp.NewServer( - CreateRoleEndpoint(svc), - d.DecodeCreateRole, - api.EncodeResponse, - opts..., - ), "create_role").ServeHTTP) - - r.Get("/", otelhttp.NewHandler(kithttp.NewServer( - ListRolesEndpoint(svc), - d.DecodeListRoles, - api.EncodeResponse, - opts..., - ), "list_roles").ServeHTTP) - - r.Get("/members", otelhttp.NewHandler(kithttp.NewServer( - ListEntityMembersEndpoint(svc), - d.DecodeListEntityMembers, - api.EncodeResponse, - opts..., - ), "list_entity_members").ServeHTTP) - - r.Delete("/", otelhttp.NewHandler(kithttp.NewServer( - RemoveEntityMembersEndpoint(svc), - d.DecodeListEntityMembers, - api.EncodeResponse, - opts..., - ), "delete_entity_members").ServeHTTP) - - r.Route("/{roleID}", func(r chi.Router) { - r.Get("/", otelhttp.NewHandler(kithttp.NewServer( - ViewRoleEndpoint(svc), - d.DecodeViewRole, - api.EncodeResponse, - opts..., - ), "view_role").ServeHTTP) - - r.Put("/", otelhttp.NewHandler(kithttp.NewServer( - UpdateRoleEndpoint(svc), - d.DecodeUpdateRole, - api.EncodeResponse, - opts..., - ), "update_role").ServeHTTP) - - r.Delete("/", otelhttp.NewHandler(kithttp.NewServer( - DeleteRoleEndpoint(svc), - d.DecodeDeleteRole, - api.EncodeResponse, - opts..., - ), "delete_role").ServeHTTP) - - r.Route("/actions", func(r chi.Router) { - r.Post("/", otelhttp.NewHandler(kithttp.NewServer( - AddRoleActionsEndpoint(svc), - d.DecodeAddRoleActions, - api.EncodeResponse, - opts..., - ), "add_role_actions").ServeHTTP) - - r.Get("/", otelhttp.NewHandler(kithttp.NewServer( - ListRoleActionsEndpoint(svc), - d.DecodeListRoleActions, - api.EncodeResponse, - opts..., - ), "list_role_actions").ServeHTTP) - - r.Post("/delete", otelhttp.NewHandler(kithttp.NewServer( - DeleteRoleActionsEndpoint(svc), - d.DecodeDeleteRoleActions, - api.EncodeResponse, - opts..., - ), "delete_role_actions").ServeHTTP) - - r.Post("/delete-all", otelhttp.NewHandler(kithttp.NewServer( - DeleteAllRoleActionsEndpoint(svc), - d.DecodeDeleteAllRoleActions, - api.EncodeResponse, - opts..., - ), "delete_all_role_actions").ServeHTTP) - }) - - r.Route("/members", func(r chi.Router) { - r.Post("/", otelhttp.NewHandler(kithttp.NewServer( - AddRoleMembersEndpoint(svc), - d.DecodeAddRoleMembers, - api.EncodeResponse, - opts..., - ), "add_role_members").ServeHTTP) - - r.Get("/", otelhttp.NewHandler(kithttp.NewServer( - ListRoleMembersEndpoint(svc), - d.DecodeListRoleMembers, - api.EncodeResponse, - opts..., - ), "list_role_members").ServeHTTP) - - r.Post("/delete", otelhttp.NewHandler(kithttp.NewServer( - DeleteRoleMembersEndpoint(svc), - d.DecodeDeleteRoleMembers, - api.EncodeResponse, - opts..., - ), "delete_role_members").ServeHTTP) - - r.Post("/delete-all", otelhttp.NewHandler(kithttp.NewServer( - DeleteAllRoleMembersEndpoint(svc), - d.DecodeDeleteAllRoleMembers, - api.EncodeResponse, - opts..., - ), "delete_all_role_members").ServeHTTP) - }) - }) - }) - - return r -} - -func EntityAvailableActionsRouter(svc roles.RoleManager, d Decoder, r chi.Router, opts []kithttp.ServerOption) chi.Router { - r.Get("/roles/available-actions", otelhttp.NewHandler(kithttp.NewServer( - ListAvailableActionsEndpoint(svc), - d.DecodeListAvailableActions, - api.EncodeResponse, - opts..., - ), "list_available_actions").ServeHTTP) - - return r -} diff --git a/pkg/roles/rolemanager/events/consumer/decode.go b/pkg/roles/rolemanager/events/consumer/decode.go deleted file mode 100644 index 3e918098b..000000000 --- a/pkg/roles/rolemanager/events/consumer/decode.go +++ /dev/null @@ -1,146 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package consumer - -import ( - "time" - - "github.com/absmach/magistrala/pkg/errors" - "github.com/absmach/magistrala/pkg/roles" -) - -var ( - errID = errors.New("missing or invalid 'id'") - errRoleID = errors.New("missing or invalid 'role_id'") - errName = errors.New("missing or invalid 'name'") - errEntityID = errors.New("missing or invalid 'entity_id'") - errActions = errors.New("missing or invalid 'actions'") - errMembers = errors.New("missing or invalid 'members'") - errCreatedAt = errors.New("failed to parse 'created_at' time") - errUpdatedAt = errors.New("failed to parse 'updated_at' time") - errNotString = errors.New("not string type") - - errInvalidRoleProvision = errors.New("invalid 'role_provisions'") - errRoleProvision = errors.New("failed to convert role_provisions interface'") - errRoleProvisionMembers = errors.New("failed to convert role_provisions member interface'") - errRoleProvisionActions = errors.New("failed to convert role_provisions action interface'") -) - -const ( - layout = "2006-01-02T15:04:05.999999Z" -) - -func ToRole(data map[string]any) (roles.Role, error) { - var r roles.Role - - id, ok := data["id"].(string) - if !ok { - return roles.Role{}, errID - } - r.ID = id - - name, ok := data["name"].(string) - if !ok { - return roles.Role{}, errName - } - r.Name = name - - eid, ok := data["entity_id"].(string) - if !ok { - return roles.Role{}, errEntityID - } - r.EntityID = eid - - // Following fields of groups are allowed to be empty. - - cat, ok := data["created_at"].(string) - if ok { - ct, err := time.Parse(layout, cat) - if err != nil { - return roles.Role{}, errors.Wrap(errCreatedAt, err) - } - r.CreatedAt = ct - } - - cby, ok := data["created_by"].(string) - if ok { - r.CreatedBy = cby - } - - uat, ok := data["updated_at"].(string) - if ok { - ut, err := time.Parse(layout, uat) - if err != nil { - return roles.Role{}, errors.Wrap(errUpdatedAt, err) - } - r.UpdatedAt = ut - } - - uby, ok := data["updated_by"].(string) - if ok { - r.UpdatedBy = uby - } - - return r, nil -} - -func ToStrings(data []any) ([]string, error) { - var strs []string - for _, i := range data { - str, ok := i.(string) - if !ok { - return []string{}, errNotString - } - strs = append(strs, str) - } - return strs, nil -} - -func ToRoleProvision(data map[string]any) (roles.RoleProvision, error) { - var rp roles.RoleProvision - - r, err := ToRole(data) - if err != nil { - return roles.RoleProvision{}, err - } - rp.Role = r - - // Following fields of groups are allowed to be empty. - - opActs, ok := data["optional_actions"].([]any) - if ok { - a, err := ToStrings(opActs) - if err != nil { - return roles.RoleProvision{}, errors.Wrap(errRoleProvisionActions, err) - } - rp.OptionalActions = a - } - - opMems, ok := data["optional_members"].([]any) - if ok { - m, err := ToStrings(opMems) - if err != nil { - return roles.RoleProvision{}, errors.Wrap(errRoleProvisionMembers, err) - } - rp.OptionalMembers = m - } - - return rp, nil -} - -func ToRoleProvisions(data []any) ([]roles.RoleProvision, error) { - var rps []roles.RoleProvision - for _, d := range data { - irp, ok := d.(map[string]any) - if !ok { - return []roles.RoleProvision{}, errInvalidRoleProvision - } - rp, err := ToRoleProvision(irp) - if err != nil { - return []roles.RoleProvision{}, errors.Wrap(errRoleProvision, err) - } - rps = append(rps, rp) - } - return rps, nil -} diff --git a/pkg/roles/rolemanager/events/consumer/handler.go b/pkg/roles/rolemanager/events/consumer/handler.go deleted file mode 100644 index 85dbc701c..000000000 --- a/pkg/roles/rolemanager/events/consumer/handler.go +++ /dev/null @@ -1,263 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package consumer - -import ( - "context" - "fmt" - - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - "github.com/absmach/magistrala/pkg/roles" - "github.com/absmach/magistrala/pkg/roles/rolemanager/events" -) - -const ( - errAddEntityRoleEvent = "failed to consume %s add role event : %w" - errUpdateEntityRoleEvent = "failed to consume %s update role event : %w" - errRemoveEntityRoleEvent = "failed to consume %s remove role event : %w" - errAddEntityRoleActionsEvent = "failed to consume %s add role actions event : %w" - errRemoveEntityRoleActionsEvent = "failed to consume %s remove role actions event : %w" - errRemoveEntityRoleAllActionsEvent = "failed to consume %s remove role all actions event : %w" - errAddEntityRoleMembersEvent = "failed to consume %s add role members event : %w" - errRemoveEntityRoleMembersEvent = "failed to consume %s remove role members event : %w" - errRemoveEntityRoleAllMembersEvent = "failed to consume %s remove role all members event : %w" -) - -type EventHandler struct { - entityType string - repo roles.Repository - addRole string - removeRole string - updateRole string - addRoleActions string - removeRoleActions string - removeAllRoleActions string - addRoleMembers string - removeRoleMembers string - removeRoleAllMembers string - removeMemberFromAllRoles string - removeEntityMembers string -} - -func NewEventHandler(entityType string, repo roles.Repository) EventHandler { - return EventHandler{ - entityType: entityType, - repo: repo, - addRole: entityType + "." + events.AddRole, - removeRole: entityType + "." + events.RemoveRole, - updateRole: entityType + "." + events.UpdateRole, - addRoleActions: entityType + "." + events.AddRoleActions, - removeRoleActions: entityType + "." + events.RemoveRoleActions, - removeAllRoleActions: entityType + "." + events.RemoveAllRoleActions, - addRoleMembers: entityType + "." + events.AddRoleMembers, - removeRoleMembers: entityType + "." + events.RemoveRoleMembers, - removeRoleAllMembers: entityType + "." + events.RemoveRoleAllMembers, - removeMemberFromAllRoles: entityType + "." + events.RemoveMemberFromAllRoles, - removeEntityMembers: entityType + "." + events.RemoveEntityMembers, - } -} - -func (es *EventHandler) Handle(ctx context.Context, op any, msg map[string]any) error { - switch op { - case es.addRole: - return es.AddEntityRoleHandler(ctx, msg) - case es.removeRole: - return es.RemoveEntityRoleHandler(ctx, msg) - case es.updateRole: - return es.UpdateEntityRoleHandler(ctx, msg) - case es.addRoleActions: - return es.AddEntityRoleActionsHandler(ctx, msg) - case es.removeRoleActions: - return es.RemoveEntityRoleActionsHandler(ctx, msg) - case es.removeAllRoleActions: - return es.RemoveAllEntityRoleActionsHandler(ctx, msg) - case es.addRoleMembers: - return es.AddEntityRoleMembersHandler(ctx, msg) - case es.removeRoleMembers: - return es.RemoveEntityRoleMembersHandler(ctx, msg) - case es.removeRoleAllMembers: - return es.RemoveAllMembersFromEntityRoleHandler(ctx, msg) - case es.removeEntityMembers: - return es.RemoveEntityMembersHandler(ctx, msg) - case es.removeMemberFromAllRoles: - return es.RemoveMemberFromAllEntityHandler(ctx, msg) - } - return nil -} - -func (es *EventHandler) AddEntityRoleHandler(ctx context.Context, data map[string]any) error { - rps, err := ToRoleProvision(data) - if err != nil { - return fmt.Errorf(errAddEntityRoleEvent, es.entityType, err) - } - if _, err := es.repo.AddRoles(ctx, []roles.RoleProvision{rps}); err != nil { - if !errors.Contains(err, repoerr.ErrConflict) { - return fmt.Errorf(errAddEntityRoleEvent, es.entityType, err) - } - } - - return nil -} - -func (es *EventHandler) UpdateEntityRoleHandler(ctx context.Context, data map[string]any) error { - ro, err := ToRole(data) - if err != nil { - return fmt.Errorf(errUpdateEntityRoleEvent, es.entityType, err) - } - - if _, err = es.repo.UpdateRole(ctx, ro); err != nil { - return fmt.Errorf(errUpdateEntityRoleEvent, es.entityType, err) - } - - return nil -} - -func (es *EventHandler) RemoveEntityRoleHandler(ctx context.Context, data map[string]any) error { - id, ok := data["role_id"].(string) - if !ok { - return fmt.Errorf(errRemoveEntityRoleEvent, es.entityType, errRoleID) - } - - if err := es.repo.RemoveRoles(ctx, []string{id}); err != nil { - return fmt.Errorf(errRemoveEntityRoleEvent, es.entityType, err) - } - - return nil -} - -func (es *EventHandler) AddEntityRoleActionsHandler(ctx context.Context, data map[string]any) error { - id, ok := data["role_id"].(string) - if !ok { - return fmt.Errorf(errAddEntityRoleActionsEvent, es.entityType, errRoleID) - } - iacts, ok := data["actions"].([]any) - if !ok { - return fmt.Errorf(errAddEntityRoleActionsEvent, es.entityType, errActions) - } - acts, err := ToStrings(iacts) - if err != nil { - return fmt.Errorf(errAddEntityRoleActionsEvent, es.entityType, err) - } - - if _, err := es.repo.RoleAddActions(ctx, roles.Role{ID: id}, acts); err != nil { - return fmt.Errorf(errAddEntityRoleActionsEvent, es.entityType, err) - } - - return nil -} - -func (es *EventHandler) RemoveEntityRoleActionsHandler(ctx context.Context, data map[string]any) error { - id, ok := data["role_id"].(string) - if !ok { - return fmt.Errorf(errAddEntityRoleActionsEvent, es.entityType, errRoleID) - } - iacts, ok := data["actions"].([]any) - if !ok { - return fmt.Errorf(errAddEntityRoleActionsEvent, es.entityType, errActions) - } - acts, err := ToStrings(iacts) - if err != nil { - return fmt.Errorf(errAddEntityRoleActionsEvent, es.entityType, err) - } - - if err := es.repo.RoleRemoveActions(ctx, roles.Role{ID: id}, acts); err != nil { - return fmt.Errorf(errAddEntityRoleActionsEvent, es.entityType, err) - } - return nil -} - -func (es *EventHandler) RemoveAllEntityRoleActionsHandler(ctx context.Context, data map[string]any) error { - id, ok := data["role_id"].(string) - if !ok { - return fmt.Errorf(errRemoveEntityRoleAllActionsEvent, es.entityType, errRoleID) - } - - if err := es.repo.RoleRemoveAllActions(ctx, roles.Role{ID: id}); err != nil { - return fmt.Errorf(errRemoveEntityRoleAllActionsEvent, es.entityType, err) - } - return nil -} - -func (es *EventHandler) AddEntityRoleMembersHandler(ctx context.Context, data map[string]any) error { - id, ok := data["role_id"].(string) - if !ok { - return fmt.Errorf(errAddEntityRoleMembersEvent, es.entityType, errRoleID) - } - entityID, ok := data["entity_id"].(string) - if !ok { - return fmt.Errorf(errRemoveEntityRoleAllMembersEvent, es.entityType, errEntityID) - } - imems, ok := data["members"].([]any) - if !ok { - return fmt.Errorf(errAddEntityRoleMembersEvent, es.entityType, errMembers) - } - mems, err := ToStrings(imems) - if err != nil { - return fmt.Errorf(errAddEntityRoleMembersEvent, es.entityType, err) - } - - if _, err := es.repo.RoleAddMembers(ctx, roles.Role{ID: id, EntityID: entityID}, mems); err != nil { - return fmt.Errorf(errAddEntityRoleMembersEvent, es.entityType, err) - } - - return nil -} - -func (es *EventHandler) RemoveEntityRoleMembersHandler(ctx context.Context, data map[string]any) error { - id, ok := data["role_id"].(string) - if !ok { - return fmt.Errorf(errRemoveEntityRoleMembersEvent, es.entityType, errRoleID) - } - imems, ok := data["members"].([]any) - if !ok { - return fmt.Errorf(errRemoveEntityRoleMembersEvent, es.entityType, errMembers) - } - mems, err := ToStrings(imems) - if err != nil { - return fmt.Errorf(errRemoveEntityRoleMembersEvent, es.entityType, err) - } - - if err := es.repo.RoleRemoveMembers(ctx, roles.Role{ID: id}, mems); err != nil { - return fmt.Errorf(errRemoveEntityRoleMembersEvent, es.entityType, err) - } - - return nil -} - -func (es *EventHandler) RemoveAllMembersFromEntityRoleHandler(ctx context.Context, data map[string]any) error { - id, ok := data["role_id"].(string) - if !ok { - return fmt.Errorf(errRemoveEntityRoleAllMembersEvent, es.entityType, errRoleID) - } - - if err := es.repo.RoleRemoveAllMembers(ctx, roles.Role{ID: id}); err != nil { - return fmt.Errorf(errRemoveEntityRoleAllMembersEvent, es.entityType, err) - } - return nil -} - -func (es *EventHandler) RemoveEntityMembersHandler(ctx context.Context, data map[string]any) error { - entityID, ok := data["entity_id"].(string) - if !ok { - return fmt.Errorf(errRemoveEntityRoleAllMembersEvent, es.entityType, errEntityID) - } - imems, ok := data["members"].([]any) - if !ok { - return fmt.Errorf(errRemoveEntityRoleMembersEvent, es.entityType, errMembers) - } - mems, err := ToStrings(imems) - if err != nil { - return fmt.Errorf(errRemoveEntityRoleMembersEvent, es.entityType, err) - } - - // added when repo is implemented. - _ = entityID - _ = mems - return nil -} - -func (es *EventHandler) RemoveMemberFromAllEntityHandler(ctx context.Context, data map[string]any) error { - return nil -} diff --git a/pkg/roles/rolemanager/events/doc.go b/pkg/roles/rolemanager/events/doc.go deleted file mode 100644 index a115b5f92..000000000 --- a/pkg/roles/rolemanager/events/doc.go +++ /dev/null @@ -1,6 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package events provides the domain concept definitions needed to -// support Magistrala auth service functionality. -package events diff --git a/pkg/roles/rolemanager/events/events.go b/pkg/roles/rolemanager/events/events.go deleted file mode 100644 index 9092a5a69..000000000 --- a/pkg/roles/rolemanager/events/events.go +++ /dev/null @@ -1,406 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package events - -import ( - "github.com/absmach/magistrala/pkg/events" - "github.com/absmach/magistrala/pkg/roles" -) - -const ( - AddRole = "role.add" - RemoveRole = "role.remove" - UpdateRole = "role.update" - ViewRole = "role.view" - ViewAllRole = "role.view_all" - ListAvailableActions = "role.list_available_actions" - AddRoleActions = "role.actions.add" - ListRoleActions = "role.actions.ist" - CheckRoleActions = "role.actions.check" - RemoveRoleActions = "role.actions.remove" - RemoveAllRoleActions = "role.actions.remove_all" - AddRoleMembers = "role.members.add" - ListRoleMembers = "role.members.list" - CheckRoleMembers = "role.members.check" - RemoveRoleMembers = "role.members.remove" - RemoveRoleAllMembers = "role.members.remove_all" - ListEntityMembers = "members.list" - RemoveEntityMembers = "members.remove" - RemoveMemberFromAllRoles = "role.members.remove_from_all_roles" -) - -var ( - _ events.Event = (*addRoleEvent)(nil) - _ events.Event = (*removeRoleEvent)(nil) - _ events.Event = (*updateRoleEvent)(nil) - _ events.Event = (*retrieveRoleEvent)(nil) - _ events.Event = (*retrieveAllRolesEvent)(nil) - _ events.Event = (*listAvailableActionsEvent)(nil) - _ events.Event = (*roleAddActionsEvent)(nil) - _ events.Event = (*roleListActionsEvent)(nil) - _ events.Event = (*roleCheckActionsExistsEvent)(nil) - _ events.Event = (*roleRemoveActionsEvent)(nil) - _ events.Event = (*roleRemoveAllActionsEvent)(nil) - _ events.Event = (*roleAddMembersEvent)(nil) - _ events.Event = (*roleListMembersEvent)(nil) - _ events.Event = (*roleCheckMembersExistsEvent)(nil) - _ events.Event = (*roleRemoveMembersEvent)(nil) - _ events.Event = (*roleRemoveAllMembersEvent)(nil) - _ events.Event = (*listEntityMembersEvent)(nil) - _ events.Event = (*removeEntityMembersEvent)(nil) - _ events.Event = (*removeMemberFromAllRolesEvent)(nil) -) - -type addRoleEvent struct { - operationPrefix string - roles.RoleProvision - requestID string -} - -func (are addRoleEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": are.operationPrefix + AddRole, - "id": are.ID, - "name": are.Name, - "entity_id": are.EntityID, - "created_by": are.CreatedBy, - "created_at": are.CreatedAt, - "updated_by": are.UpdatedBy, - "updated_at": are.UpdatedAt, - "optional_actions": are.OptionalActions, - "optional_members": are.OptionalMembers, - "request_id": are.requestID, - } - return val, nil -} - -type removeRoleEvent struct { - operationPrefix string - entityID string - roleID string - requestID string -} - -func (rre removeRoleEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": rre.operationPrefix + RemoveRole, - "entity_id": rre.entityID, - "role_id": rre.roleID, - "request_id": rre.requestID, - } - return val, nil -} - -type updateRoleEvent struct { - operationPrefix string - roles.Role - requestID string -} - -func (ure updateRoleEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": ure.operationPrefix + UpdateRole, - "id": ure.ID, - "name": ure.Name, - "entity_id": ure.EntityID, - "created_by": ure.CreatedBy, - "created_at": ure.CreatedAt, - "updated_by": ure.UpdatedBy, - "updated_at": ure.UpdatedAt, - "request_id": ure.requestID, - } - return val, nil -} - -type retrieveRoleEvent struct { - operationPrefix string - roles.Role - requestID string -} - -func (rre retrieveRoleEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": rre.operationPrefix + ViewRole, - "id": rre.ID, - "name": rre.Name, - "entity_id": rre.EntityID, - "created_by": rre.CreatedBy, - "created_at": rre.CreatedAt, - "updated_by": rre.UpdatedBy, - "updated_at": rre.UpdatedAt, - "request_id": rre.requestID, - } - return val, nil -} - -type retrieveAllRolesEvent struct { - operationPrefix string - entityID string - limit uint64 - offset uint64 - requestID string -} - -func (rare retrieveAllRolesEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": rare.operationPrefix + ViewAllRole, - "entity_id": rare.entityID, - "limit": rare.limit, - "offset": rare.offset, - "request_id": rare.requestID, - } - return val, nil -} - -type listAvailableActionsEvent struct { - operationPrefix string - requestID string -} - -func (laae listAvailableActionsEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": laae.operationPrefix + ListAvailableActions, - "request_id": laae.requestID, - } - return val, nil -} - -type roleAddActionsEvent struct { - operationPrefix string - entityID string - roleID string - actions []string - requestID string -} - -func (raae roleAddActionsEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": raae.operationPrefix + AddRoleActions, - "entity_id": raae.entityID, - "role_id": raae.roleID, - "actions": raae.actions, - "request_id": raae.requestID, - } - return val, nil -} - -type roleListActionsEvent struct { - operationPrefix string - entityID string - roleID string - requestID string -} - -func (rlae roleListActionsEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": rlae.operationPrefix + ListRoleActions, - "entity_id": rlae.entityID, - "role_id": rlae.roleID, - "request_id": rlae.requestID, - } - return val, nil -} - -type roleCheckActionsExistsEvent struct { - operationPrefix string - entityID string - roleID string - actions []string - isAllExists bool - requestID string -} - -func (rcaee roleCheckActionsExistsEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": rcaee.operationPrefix + CheckRoleActions, - "entity_id": rcaee.entityID, - "role_id": rcaee.roleID, - "actions": rcaee.actions, - "is_all_exists": rcaee.isAllExists, - "request_id": rcaee.requestID, - } - return val, nil -} - -type roleRemoveActionsEvent struct { - operationPrefix string - entityID string - roleID string - actions []string - requestID string -} - -func (rrae roleRemoveActionsEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": rrae.operationPrefix + RemoveRoleActions, - "entity_id": rrae.entityID, - "role_id": rrae.roleID, - "actions": rrae.actions, - "request_id": rrae.requestID, - } - return val, nil -} - -type roleRemoveAllActionsEvent struct { - operationPrefix string - entityID string - roleID string - requestID string -} - -func (rraae roleRemoveAllActionsEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": rraae.operationPrefix + RemoveAllRoleActions, - "entity_id": rraae.entityID, - "role_id": rraae.roleID, - "request_id": rraae.requestID, - } - return val, nil -} - -type roleAddMembersEvent struct { - operationPrefix string - entityID string - roleID string - members []string - requestID string -} - -func (rame roleAddMembersEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": rame.operationPrefix + AddRoleMembers, - "entity_id": rame.entityID, - "role_id": rame.roleID, - "members": rame.members, - "request_id": rame.requestID, - } - return val, nil -} - -type roleListMembersEvent struct { - operationPrefix string - entityID string - roleID string - limit uint64 - offset uint64 - requestID string -} - -func (rlme roleListMembersEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": rlme.operationPrefix + ListRoleMembers, - "entity_id": rlme.entityID, - "role_id": rlme.roleID, - "limit": rlme.limit, - "offset": rlme.offset, - "request_id": rlme.requestID, - } - return val, nil -} - -type roleCheckMembersExistsEvent struct { - operationPrefix string - entityID string - roleID string - members []string - requestID string -} - -func (rcmee roleCheckMembersExistsEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": rcmee.operationPrefix + CheckRoleMembers, - "entity_id": rcmee.entityID, - "role_id": rcmee.roleID, - "members": rcmee.members, - "request_id": rcmee.requestID, - } - return val, nil -} - -type roleRemoveMembersEvent struct { - operationPrefix string - entityID string - roleID string - members []string - requestID string -} - -func (rrme roleRemoveMembersEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": rrme.operationPrefix + RemoveRoleMembers, - "entity_id": rrme.entityID, - "role_id": rrme.roleID, - "members": rrme.members, - "request_id": rrme.requestID, - } - return val, nil -} - -type roleRemoveAllMembersEvent struct { - operationPrefix string - entityID string - roleID string - requestID string -} - -func (rrame roleRemoveAllMembersEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": rrame.operationPrefix + RemoveRoleAllMembers, - "entity_id": rrame.entityID, - "role_id": rrame.roleID, - "request_id": rrame.requestID, - } - return val, nil -} - -type listEntityMembersEvent struct { - operationPrefix string - entityID string - limit uint64 - offset uint64 - requestID string -} - -func (leme listEntityMembersEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": leme.operationPrefix + ListEntityMembers, - "entity_id": leme.entityID, - "limit": leme.limit, - "offset": leme.offset, - "request_id": leme.requestID, - } - return val, nil -} - -type removeEntityMembersEvent struct { - operationPrefix string - entityID string - members []string - requestID string -} - -func (reme removeEntityMembersEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": reme.operationPrefix + RemoveEntityMembers, - "entity_id": reme.entityID, - "members": reme.members, - "request_id": reme.requestID, - } - return val, nil -} - -type removeMemberFromAllRolesEvent struct { - operationPrefix string - memberID string - requestID string -} - -func (rmare removeMemberFromAllRolesEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": rmare.operationPrefix + RemoveMemberFromAllRoles, - "member_id": rmare.memberID, - "request_id": rmare.requestID, - } - return val, nil -} diff --git a/pkg/roles/rolemanager/events/streams.go b/pkg/roles/rolemanager/events/streams.go deleted file mode 100644 index 73d15bdcb..000000000 --- a/pkg/roles/rolemanager/events/streams.go +++ /dev/null @@ -1,381 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package events - -import ( - "context" - - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/events" - "github.com/absmach/magistrala/pkg/roles" - "github.com/go-chi/chi/v5/middleware" -) - -const ( - magistralaPrefix = "magistrala." - rolesPrefix = "roles" -) - -var _ roles.RoleManager = (*RoleManagerEventStore)(nil) - -type RoleManagerEventStore struct { - events.Publisher - svc roles.RoleManager - operationPrefix string - svcName string - streamID string -} - -// NewEventStoreMiddleware returns wrapper around auth service that sends -// events to event store. -func NewRoleManagerEventStore(svcName, operationPrefix string, svc roles.RoleManager, publisher events.Publisher) RoleManagerEventStore { - return RoleManagerEventStore{ - svcName: svcName, - operationPrefix: operationPrefix, - svc: svc, - streamID: magistralaPrefix + operationPrefix + rolesPrefix, - Publisher: publisher, - } -} - -func (rmes *RoleManagerEventStore) AddRole(ctx context.Context, session authn.Session, entityID, roleName string, optionalActions []string, optionalMembers []string) (roles.RoleProvision, error) { - nrp, err := rmes.svc.AddRole(ctx, session, entityID, roleName, optionalActions, optionalMembers) - if err != nil { - return nrp, err - } - - e := addRoleEvent{ - operationPrefix: rmes.operationPrefix, - RoleProvision: nrp, - requestID: middleware.GetReqID(ctx), - } - if err := rmes.Publish(ctx, rmes.streamID, e); err != nil { - return nrp, err - } - return nrp, nil -} - -func (rmes *RoleManagerEventStore) RemoveRole(ctx context.Context, session authn.Session, entityID, roleID string) error { - if err := rmes.svc.RemoveRole(ctx, session, entityID, roleID); err != nil { - return err - } - e := removeRoleEvent{ - operationPrefix: rmes.operationPrefix, - roleID: roleID, - entityID: entityID, - requestID: middleware.GetReqID(ctx), - } - if err := rmes.Publish(ctx, rmes.streamID, e); err != nil { - return err - } - return nil -} - -func (rmes *RoleManagerEventStore) UpdateRoleName(ctx context.Context, session authn.Session, entityID, roleID, newRoleName string) (roles.Role, error) { - ro, err := rmes.svc.UpdateRoleName(ctx, session, entityID, roleID, newRoleName) - if err != nil { - return ro, err - } - - e := updateRoleEvent{ - operationPrefix: rmes.operationPrefix, - Role: ro, - requestID: middleware.GetReqID(ctx), - } - if err := rmes.Publish(ctx, rmes.streamID, e); err != nil { - return ro, err - } - return ro, nil -} - -func (rmes *RoleManagerEventStore) RetrieveRole(ctx context.Context, session authn.Session, entityID, roleID string) (roles.Role, error) { - ro, err := rmes.svc.RetrieveRole(ctx, session, entityID, roleID) - if err != nil { - return ro, err - } - e := retrieveRoleEvent{ - operationPrefix: rmes.operationPrefix, - Role: ro, - requestID: middleware.GetReqID(ctx), - } - if err := rmes.Publish(ctx, rmes.streamID, e); err != nil { - return ro, err - } - return ro, nil -} - -func (rmes *RoleManagerEventStore) RetrieveAllRoles(ctx context.Context, session authn.Session, entityID string, limit, offset uint64) (roles.RolePage, error) { - rp, err := rmes.svc.RetrieveAllRoles(ctx, session, entityID, limit, offset) - if err != nil { - return rp, err - } - - e := retrieveAllRolesEvent{ - operationPrefix: rmes.operationPrefix, - entityID: entityID, - limit: limit, - offset: offset, - requestID: middleware.GetReqID(ctx), - } - if err := rmes.Publish(ctx, rmes.streamID, e); err != nil { - return rp, err - } - return rp, nil -} - -func (rmes *RoleManagerEventStore) ListAvailableActions(ctx context.Context, session authn.Session) ([]string, error) { - actions, err := rmes.svc.ListAvailableActions(ctx, session) - if err != nil { - return actions, err - } - e := listAvailableActionsEvent{ - operationPrefix: rmes.operationPrefix, - requestID: middleware.GetReqID(ctx), - } - if err := rmes.Publish(ctx, rmes.streamID, e); err != nil { - return actions, err - } - return actions, nil -} - -func (rmes *RoleManagerEventStore) RoleAddActions(ctx context.Context, session authn.Session, entityID, roleID string, actions []string) ([]string, error) { - actions, err := rmes.svc.RoleAddActions(ctx, session, entityID, roleID, actions) - if err != nil { - return actions, err - } - e := roleAddActionsEvent{ - operationPrefix: rmes.operationPrefix, - entityID: entityID, - roleID: roleID, - actions: actions, - requestID: middleware.GetReqID(ctx), - } - if err := rmes.Publish(ctx, rmes.streamID, e); err != nil { - return actions, err - } - return actions, nil -} - -func (rmes *RoleManagerEventStore) RoleListActions(ctx context.Context, session authn.Session, entityID, roleID string) ([]string, error) { - actions, err := rmes.svc.RoleListActions(ctx, session, entityID, roleID) - if err != nil { - return actions, err - } - - e := roleListActionsEvent{ - operationPrefix: rmes.operationPrefix, - entityID: entityID, - roleID: roleID, - requestID: middleware.GetReqID(ctx), - } - if err := rmes.Publish(ctx, rmes.streamID, e); err != nil { - return actions, err - } - return actions, nil -} - -func (rmes *RoleManagerEventStore) RoleCheckActionsExists(ctx context.Context, session authn.Session, entityID, roleID string, actions []string) (bool, error) { - isAllExists, err := rmes.svc.RoleCheckActionsExists(ctx, session, entityID, roleID, actions) - if err != nil { - return isAllExists, err - } - - e := roleCheckActionsExistsEvent{ - operationPrefix: rmes.operationPrefix, - entityID: entityID, - roleID: roleID, - actions: actions, - isAllExists: isAllExists, - requestID: middleware.GetReqID(ctx), - } - if err := rmes.Publish(ctx, rmes.streamID, e); err != nil { - return isAllExists, err - } - return isAllExists, nil -} - -func (rmes *RoleManagerEventStore) RoleRemoveActions(ctx context.Context, session authn.Session, entityID, roleID string, actions []string) (err error) { - if err := rmes.svc.RoleRemoveActions(ctx, session, entityID, roleID, actions); err != nil { - return err - } - - e := roleRemoveActionsEvent{ - operationPrefix: rmes.operationPrefix, - entityID: entityID, - roleID: roleID, - actions: actions, - requestID: middleware.GetReqID(ctx), - } - if err := rmes.Publish(ctx, rmes.streamID, e); err != nil { - return err - } - return nil -} - -func (rmes *RoleManagerEventStore) RoleRemoveAllActions(ctx context.Context, session authn.Session, entityID, roleID string) error { - if err := rmes.svc.RoleRemoveAllActions(ctx, session, entityID, roleID); err != nil { - return err - } - - e := roleRemoveAllActionsEvent{ - operationPrefix: rmes.operationPrefix, - entityID: entityID, - roleID: roleID, - requestID: middleware.GetReqID(ctx), - } - if err := rmes.Publish(ctx, rmes.streamID, e); err != nil { - return err - } - return nil -} - -func (rmes *RoleManagerEventStore) RoleAddMembers(ctx context.Context, session authn.Session, entityID, roleID string, members []string) ([]string, error) { - mems, err := rmes.svc.RoleAddMembers(ctx, session, entityID, roleID, members) - if err != nil { - return mems, err - } - - err = rmes.RoleAddMembersEventPublisher(ctx, entityID, roleID, mems) - return mems, err -} - -func (rmes *RoleManagerEventStore) RoleAddMembersEventPublisher(ctx context.Context, entityID, roleID string, members []string) error { - e := roleAddMembersEvent{ - operationPrefix: rmes.operationPrefix, - entityID: entityID, - roleID: roleID, - members: members, - requestID: middleware.GetReqID(ctx), - } - if err := rmes.Publish(ctx, rmes.streamID, e); err != nil { - return err - } - return nil -} - -func (rmes *RoleManagerEventStore) RoleListMembers(ctx context.Context, session authn.Session, entityID, roleID string, limit, offset uint64) (roles.MembersPage, error) { - mp, err := rmes.svc.RoleListMembers(ctx, session, entityID, roleID, limit, offset) - if err != nil { - return mp, err - } - - e := roleListMembersEvent{ - operationPrefix: rmes.operationPrefix, - entityID: entityID, - roleID: roleID, - limit: limit, - offset: offset, - requestID: middleware.GetReqID(ctx), - } - if err := rmes.Publish(ctx, rmes.streamID, e); err != nil { - return mp, err - } - return mp, nil -} - -func (rmes *RoleManagerEventStore) RoleCheckMembersExists(ctx context.Context, session authn.Session, entityID, roleID string, members []string) (bool, error) { - isAllExists, err := rmes.svc.RoleCheckMembersExists(ctx, session, entityID, roleID, members) - if err != nil { - return isAllExists, err - } - - e := roleCheckMembersExistsEvent{ - operationPrefix: rmes.operationPrefix, - entityID: entityID, - roleID: roleID, - members: members, - requestID: middleware.GetReqID(ctx), - } - if err := rmes.Publish(ctx, rmes.streamID, e); err != nil { - return isAllExists, err - } - return isAllExists, nil -} - -func (rmes *RoleManagerEventStore) RoleRemoveMembers(ctx context.Context, session authn.Session, entityID, roleID string, members []string) (err error) { - if err := rmes.svc.RoleRemoveMembers(ctx, session, entityID, roleID, members); err != nil { - return err - } - - e := roleRemoveMembersEvent{ - operationPrefix: rmes.operationPrefix, - entityID: entityID, - roleID: roleID, - members: members, - requestID: middleware.GetReqID(ctx), - } - if err := rmes.Publish(ctx, rmes.streamID, e); err != nil { - return err - } - return nil -} - -func (rmes *RoleManagerEventStore) RoleRemoveAllMembers(ctx context.Context, session authn.Session, entityID, roleID string) (err error) { - if err := rmes.svc.RoleRemoveAllMembers(ctx, session, entityID, roleID); err != nil { - return err - } - - e := roleRemoveAllMembersEvent{ - operationPrefix: rmes.operationPrefix, - entityID: entityID, - roleID: roleID, - requestID: middleware.GetReqID(ctx), - } - if err := rmes.Publish(ctx, rmes.streamID, e); err != nil { - return err - } - return nil -} - -func (rmes *RoleManagerEventStore) ListEntityMembers(ctx context.Context, session authn.Session, entityID string, pageQuery roles.MembersRolePageQuery) (roles.MembersRolePage, error) { - mems, err := rmes.svc.ListEntityMembers(ctx, session, entityID, pageQuery) - if err != nil { - return mems, err - } - - e := listEntityMembersEvent{ - operationPrefix: rmes.operationPrefix, - entityID: entityID, - limit: pageQuery.Limit, - offset: pageQuery.Offset, - requestID: middleware.GetReqID(ctx), - } - if err := rmes.Publish(ctx, rmes.streamID, e); err != nil { - return mems, err - } - return mems, nil -} - -func (rmes *RoleManagerEventStore) RemoveEntityMembers(ctx context.Context, session authn.Session, entityID string, members []string) error { - if err := rmes.svc.RemoveEntityMembers(ctx, session, entityID, members); err != nil { - return err - } - - e := removeEntityMembersEvent{ - operationPrefix: rmes.operationPrefix, - entityID: entityID, - members: members, - requestID: middleware.GetReqID(ctx), - } - if err := rmes.Publish(ctx, rmes.streamID, e); err != nil { - return err - } - return nil -} - -func (rmes *RoleManagerEventStore) RemoveMemberFromAllRoles(ctx context.Context, session authn.Session, memberID string) (err error) { - if err := rmes.svc.RemoveMemberFromAllRoles(ctx, session, memberID); err != nil { - return err - } - - e := removeMemberFromAllRolesEvent{ - operationPrefix: rmes.operationPrefix, - memberID: memberID, - requestID: middleware.GetReqID(ctx), - } - if err := rmes.Publish(ctx, rmes.streamID, e); err != nil { - return err - } - return nil -} diff --git a/pkg/roles/rolemanager/middleware/authorization.go b/pkg/roles/rolemanager/middleware/authorization.go deleted file mode 100644 index 155c03341..000000000 --- a/pkg/roles/rolemanager/middleware/authorization.go +++ /dev/null @@ -1,373 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package middleware - -import ( - "context" - "fmt" - - "github.com/absmach/magistrala/auth" - "github.com/absmach/magistrala/pkg/authn" - smqauthz "github.com/absmach/magistrala/pkg/authz" - "github.com/absmach/magistrala/pkg/errors" - "github.com/absmach/magistrala/pkg/permissions" - "github.com/absmach/magistrala/pkg/policies" - "github.com/absmach/magistrala/pkg/roles" -) - -var _ roles.RoleManager = (*RoleManagerAuthorizationMiddleware)(nil) - -type RoleManagerAuthorizationMiddleware struct { - entityType string - svc roles.RoleManager - authz smqauthz.Authorization - ops permissions.Operations[permissions.RoleOperation] -} - -// NewAuthorization adds authorization for role related methods to the core service. -func NewAuthorization(entityType string, svc roles.RoleManager, authz smqauthz.Authorization, roleOps permissions.Operations[permissions.RoleOperation]) (RoleManagerAuthorizationMiddleware, error) { - if err := roleOps.Validate(); err != nil { - return RoleManagerAuthorizationMiddleware{}, err - } - - ram := RoleManagerAuthorizationMiddleware{ - entityType: entityType, - svc: svc, - authz: authz, - ops: roleOps, - } - - return ram, nil -} - -func (ram RoleManagerAuthorizationMiddleware) AddRole(ctx context.Context, session authn.Session, entityID, roleName string, optionalActions []string, optionalMembers []string) (roles.RoleProvision, error) { - if err := ram.authorize(ctx, session, roles.OpAddRole, smqauthz.PolicyReq{ - Domain: session.DomainID, - Subject: session.DomainUserID, - SubjectType: policies.UserType, - SubjectKind: policies.UsersKind, - Object: entityID, - ObjectType: ram.entityType, - }); err != nil { - return roles.RoleProvision{}, err - } - if err := ram.validateMembers(ctx, session, optionalMembers); err != nil { - return roles.RoleProvision{}, err - } - return ram.svc.AddRole(ctx, session, entityID, roleName, optionalActions, optionalMembers) -} - -func (ram RoleManagerAuthorizationMiddleware) RemoveRole(ctx context.Context, session authn.Session, entityID, roleID string) error { - if err := ram.authorize(ctx, session, roles.OpRemoveRole, smqauthz.PolicyReq{ - Domain: session.DomainID, - Subject: session.DomainUserID, - SubjectType: policies.UserType, - SubjectKind: policies.UsersKind, - Object: entityID, - ObjectType: ram.entityType, - }); err != nil { - return err - } - return ram.svc.RemoveRole(ctx, session, entityID, roleID) -} - -func (ram RoleManagerAuthorizationMiddleware) UpdateRoleName(ctx context.Context, session authn.Session, entityID, roleID, newRoleName string) (roles.Role, error) { - if err := ram.authorize(ctx, session, roles.OpUpdateRoleName, smqauthz.PolicyReq{ - Domain: session.DomainID, - Subject: session.DomainUserID, - SubjectType: policies.UserType, - SubjectKind: policies.UsersKind, - Object: entityID, - ObjectType: ram.entityType, - }); err != nil { - return roles.Role{}, err - } - return ram.svc.UpdateRoleName(ctx, session, entityID, roleID, newRoleName) -} - -func (ram RoleManagerAuthorizationMiddleware) RetrieveRole(ctx context.Context, session authn.Session, entityID, roleID string) (roles.Role, error) { - if err := ram.authorize(ctx, session, roles.OpRetrieveRole, smqauthz.PolicyReq{ - Domain: session.DomainID, - Subject: session.DomainUserID, - SubjectType: policies.UserType, - SubjectKind: policies.UsersKind, - Object: entityID, - ObjectType: ram.entityType, - }); err != nil { - return roles.Role{}, err - } - return ram.svc.RetrieveRole(ctx, session, entityID, roleID) -} - -func (ram RoleManagerAuthorizationMiddleware) RetrieveAllRoles(ctx context.Context, session authn.Session, entityID string, limit, offset uint64) (roles.RolePage, error) { - if err := ram.authorize(ctx, session, roles.OpRetrieveAllRoles, smqauthz.PolicyReq{ - Domain: session.DomainID, - Subject: session.DomainUserID, - SubjectType: policies.UserType, - SubjectKind: policies.UsersKind, - Object: entityID, - ObjectType: ram.entityType, - }); err != nil { - return roles.RolePage{}, err - } - return ram.svc.RetrieveAllRoles(ctx, session, entityID, limit, offset) -} - -func (ram RoleManagerAuthorizationMiddleware) ListAvailableActions(ctx context.Context, session authn.Session) ([]string, error) { - return ram.svc.ListAvailableActions(ctx, session) -} - -func (ram RoleManagerAuthorizationMiddleware) RoleAddActions(ctx context.Context, session authn.Session, entityID, roleID string, actions []string) (ops []string, err error) { - if err := ram.authorize(ctx, session, roles.OpRoleAddActions, smqauthz.PolicyReq{ - Domain: session.DomainID, - Subject: session.DomainUserID, - SubjectType: policies.UserType, - SubjectKind: policies.UsersKind, - Object: entityID, - ObjectType: ram.entityType, - }); err != nil { - return []string{}, err - } - - return ram.svc.RoleAddActions(ctx, session, entityID, roleID, actions) -} - -func (ram RoleManagerAuthorizationMiddleware) RoleListActions(ctx context.Context, session authn.Session, entityID, roleID string) ([]string, error) { - if err := ram.authorize(ctx, session, roles.OpRoleListActions, smqauthz.PolicyReq{ - Domain: session.DomainID, - Subject: session.DomainUserID, - SubjectType: policies.UserType, - SubjectKind: policies.UsersKind, - Object: entityID, - ObjectType: ram.entityType, - }); err != nil { - return []string{}, err - } - - return ram.svc.RoleListActions(ctx, session, entityID, roleID) -} - -func (ram RoleManagerAuthorizationMiddleware) RoleCheckActionsExists(ctx context.Context, session authn.Session, entityID, roleID string, actions []string) (bool, error) { - if err := ram.authorize(ctx, session, roles.OpRoleCheckActionsExists, smqauthz.PolicyReq{ - Domain: session.DomainID, - Subject: session.DomainUserID, - SubjectType: policies.UserType, - SubjectKind: policies.UsersKind, - Object: entityID, - ObjectType: ram.entityType, - }); err != nil { - return false, err - } - return ram.svc.RoleCheckActionsExists(ctx, session, entityID, roleID, actions) -} - -func (ram RoleManagerAuthorizationMiddleware) RoleRemoveActions(ctx context.Context, session authn.Session, entityID, roleID string, actions []string) (err error) { - if err := ram.authorize(ctx, session, roles.OpRoleRemoveActions, smqauthz.PolicyReq{ - Domain: session.DomainID, - Subject: session.DomainUserID, - SubjectType: policies.UserType, - SubjectKind: policies.UsersKind, - Object: entityID, - ObjectType: ram.entityType, - }); err != nil { - return err - } - - return ram.svc.RoleRemoveActions(ctx, session, entityID, roleID, actions) -} - -func (ram RoleManagerAuthorizationMiddleware) RoleRemoveAllActions(ctx context.Context, session authn.Session, entityID, roleID string) error { - if err := ram.authorize(ctx, session, roles.OpRoleRemoveAllActions, smqauthz.PolicyReq{ - Domain: session.DomainID, - Subject: session.DomainUserID, - SubjectType: policies.UserType, - SubjectKind: policies.UsersKind, - Object: entityID, - ObjectType: ram.entityType, - }); err != nil { - return err - } - return ram.svc.RoleRemoveAllActions(ctx, session, entityID, roleID) -} - -func (ram RoleManagerAuthorizationMiddleware) RoleAddMembers(ctx context.Context, session authn.Session, entityID, roleID string, members []string) ([]string, error) { - if err := ram.authorize(ctx, session, roles.OpRoleAddMembers, smqauthz.PolicyReq{ - Domain: session.DomainID, - Subject: session.DomainUserID, - SubjectType: policies.UserType, - SubjectKind: policies.UsersKind, - Object: entityID, - ObjectType: ram.entityType, - }); err != nil { - return []string{}, err - } - - if err := ram.validateMembers(ctx, session, members); err != nil { - return []string{}, err - } - return ram.svc.RoleAddMembers(ctx, session, entityID, roleID, members) -} - -func (ram RoleManagerAuthorizationMiddleware) RoleListMembers(ctx context.Context, session authn.Session, entityID, roleID string, limit, offset uint64) (roles.MembersPage, error) { - if err := ram.authorize(ctx, session, roles.OpRoleListMembers, smqauthz.PolicyReq{ - Domain: session.DomainID, - Subject: session.DomainUserID, - SubjectType: policies.UserType, - SubjectKind: policies.UsersKind, - Object: entityID, - ObjectType: ram.entityType, - }); err != nil { - return roles.MembersPage{}, err - } - return ram.svc.RoleListMembers(ctx, session, entityID, roleID, limit, offset) -} - -func (ram RoleManagerAuthorizationMiddleware) RoleCheckMembersExists(ctx context.Context, session authn.Session, entityID, roleID string, members []string) (bool, error) { - if err := ram.authorize(ctx, session, roles.OpRoleCheckMembersExists, smqauthz.PolicyReq{ - Domain: session.DomainID, - Subject: session.DomainUserID, - SubjectType: policies.UserType, - SubjectKind: policies.UsersKind, - Object: entityID, - ObjectType: ram.entityType, - }); err != nil { - return false, err - } - return ram.svc.RoleCheckMembersExists(ctx, session, entityID, roleID, members) -} - -func (ram RoleManagerAuthorizationMiddleware) RoleRemoveAllMembers(ctx context.Context, session authn.Session, entityID, roleID string) (err error) { - if err := ram.authorize(ctx, session, roles.OpRoleRemoveAllMembers, smqauthz.PolicyReq{ - Domain: session.DomainID, - Subject: session.DomainUserID, - SubjectType: policies.UserType, - SubjectKind: policies.UsersKind, - Object: entityID, - ObjectType: ram.entityType, - }); err != nil { - return err - } - return ram.svc.RoleRemoveAllMembers(ctx, session, entityID, roleID) -} - -func (ram RoleManagerAuthorizationMiddleware) ListEntityMembers(ctx context.Context, session authn.Session, entityID string, pageQuery roles.MembersRolePageQuery) (roles.MembersRolePage, error) { - if err := ram.authorize(ctx, session, roles.OpRoleListMembers, smqauthz.PolicyReq{ - Domain: session.DomainID, - Subject: session.DomainUserID, - SubjectType: policies.UserType, - SubjectKind: policies.UsersKind, - Object: entityID, - ObjectType: ram.entityType, - }); err != nil { - return roles.MembersRolePage{}, err - } - return ram.svc.ListEntityMembers(ctx, session, entityID, pageQuery) -} - -func (ram RoleManagerAuthorizationMiddleware) RemoveEntityMembers(ctx context.Context, session authn.Session, entityID string, members []string) error { - if err := ram.authorize(ctx, session, roles.OpRoleRemoveAllMembers, smqauthz.PolicyReq{ - Domain: session.DomainID, - Subject: session.DomainUserID, - SubjectType: policies.UserType, - SubjectKind: policies.UsersKind, - Object: entityID, - ObjectType: ram.entityType, - }); err != nil { - return err - } - return ram.svc.RemoveEntityMembers(ctx, session, entityID, members) -} - -func (ram RoleManagerAuthorizationMiddleware) RoleRemoveMembers(ctx context.Context, session authn.Session, entityID, roleID string, members []string) (err error) { - if err := ram.authorize(ctx, session, roles.OpRoleRemoveMembers, smqauthz.PolicyReq{ - Domain: session.DomainID, - Subject: session.DomainUserID, - SubjectType: policies.UserType, - SubjectKind: policies.UsersKind, - Object: entityID, - ObjectType: ram.entityType, - }); err != nil { - return err - } - return ram.svc.RoleRemoveMembers(ctx, session, entityID, roleID, members) -} - -func (ram RoleManagerAuthorizationMiddleware) authorize(ctx context.Context, session authn.Session, op permissions.RoleOperation, pr smqauthz.PolicyReq) error { - pr.Domain = session.DomainID - - perm, err := ram.ops.GetPermission(op) - if err != nil { - return err - } - - pr.Permission = perm.String() - - var pat *smqauthz.PATReq - if session.PatID != "" { - opName := ram.ops.OperationName(op) - var patEntityType string - switch pr.ObjectType { - case policies.GroupType: - patEntityType = auth.GroupsType.String() - case policies.ClientType: - patEntityType = auth.ClientsType.String() - case policies.ChannelType: - patEntityType = auth.ChannelsType.String() - default: - return errors.Wrap(errors.ErrAuthorization, fmt.Errorf("unsupported entity type for PAT: %s", pr.ObjectType)) - } - pat = &smqauthz.PATReq{ - UserID: session.UserID, - PatID: session.PatID, - EntityID: pr.Object, - EntityType: patEntityType, - Operation: auth.RoleOperationPrefix + opName, - Domain: session.DomainID, - } - } - - if err := ram.authz.Authorize(ctx, pr, pat); err != nil { - return errors.Wrap(errors.ErrAuthorization, err) - } - - return nil -} - -func (ram RoleManagerAuthorizationMiddleware) RemoveMemberFromAllRoles(ctx context.Context, session authn.Session, memberID string) (err error) { - return ram.svc.RemoveMemberFromAllRoles(ctx, session, memberID) -} - -func (ram RoleManagerAuthorizationMiddleware) validateMembers(ctx context.Context, session authn.Session, members []string) error { - switch ram.entityType { - case policies.DomainType: - for _, member := range members { - if err := ram.authz.Authorize(ctx, smqauthz.PolicyReq{ - Permission: policies.MembershipPermission, - Subject: member, - SubjectType: policies.UserType, - SubjectKind: policies.UsersKind, - Object: policies.MagistralaObject, - ObjectType: policies.PlatformType, - }, nil); err != nil { - return errors.Wrap(errors.ErrMissingMember, err) - } - } - return nil - - default: - for _, member := range members { - if err := ram.authz.Authorize(ctx, smqauthz.PolicyReq{ - Permission: policies.MembershipPermission, - Subject: policies.EncodeDomainUserID(session.DomainID, member), - SubjectType: policies.UserType, - SubjectKind: policies.UsersKind, - Object: session.DomainID, - ObjectType: policies.DomainType, - }, nil); err != nil { - return errors.Wrap(errors.ErrMissingDomainMember, err) - } - } - return nil - } -} diff --git a/pkg/roles/rolemanager/middleware/callout.go b/pkg/roles/rolemanager/middleware/callout.go deleted file mode 100644 index 734a34017..000000000 --- a/pkg/roles/rolemanager/middleware/callout.go +++ /dev/null @@ -1,275 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package middleware - -import ( - "context" - "time" - - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/callout" - "github.com/absmach/magistrala/pkg/permissions" - "github.com/absmach/magistrala/pkg/policies" - "github.com/absmach/magistrala/pkg/roles" -) - -var _ roles.RoleManager = (*RoleManagerCalloutMiddleware)(nil) - -type RoleManagerCalloutMiddleware struct { - entityType string - svc roles.RoleManager - callout callout.Callout - roleOps permissions.Operations[permissions.RoleOperation] -} - -func NewCallout(entityType string, svc roles.RoleManager, callout callout.Callout, roleOps permissions.Operations[permissions.RoleOperation]) (RoleManagerCalloutMiddleware, error) { - if err := roleOps.Validate(); err != nil { - return RoleManagerCalloutMiddleware{}, err - } - - return RoleManagerCalloutMiddleware{ - svc: svc, - callout: callout, - entityType: entityType, - roleOps: roleOps, - }, nil -} - -func (rcm *RoleManagerCalloutMiddleware) AddRole(ctx context.Context, session authn.Session, entityID, roleName string, optionalActions []string, optionalMembers []string) (roles.RoleProvision, error) { - params := map[string]any{ - "entity_id": entityID, - "role_name": roleName, - "optional_actions": optionalActions, - "optional_members": optionalMembers, - "count": 1, - } - if err := rcm.callOut(ctx, session, roles.OpAddRole, params); err != nil { - return roles.RoleProvision{}, err - } - return rcm.svc.AddRole(ctx, session, entityID, roleName, optionalActions, optionalMembers) -} - -func (rcm *RoleManagerCalloutMiddleware) RemoveRole(ctx context.Context, session authn.Session, entityID, roleID string) error { - params := map[string]any{ - "entity_id": entityID, - "role_id": roleID, - } - if err := rcm.callOut(ctx, session, roles.OpRemoveRole, params); err != nil { - return err - } - return rcm.svc.RemoveRole(ctx, session, entityID, roleID) -} - -func (rcm *RoleManagerCalloutMiddleware) UpdateRoleName(ctx context.Context, session authn.Session, entityID, roleID, newRoleName string) (roles.Role, error) { - params := map[string]any{ - "entity_id": entityID, - "role_id": roleID, - "new_role_name": newRoleName, - } - if err := rcm.callOut(ctx, session, roles.OpUpdateRoleName, params); err != nil { - return roles.Role{}, err - } - return rcm.svc.UpdateRoleName(ctx, session, entityID, roleID, newRoleName) -} - -func (rcm *RoleManagerCalloutMiddleware) RetrieveRole(ctx context.Context, session authn.Session, entityID, roleID string) (roles.Role, error) { - params := map[string]any{ - "entity_id": entityID, - "role_id": roleID, - } - if err := rcm.callOut(ctx, session, roles.OpRetrieveRole, params); err != nil { - return roles.Role{}, err - } - return rcm.svc.RetrieveRole(ctx, session, entityID, roleID) -} - -func (rcm *RoleManagerCalloutMiddleware) RetrieveAllRoles(ctx context.Context, session authn.Session, entityID string, limit, offset uint64) (roles.RolePage, error) { - params := map[string]any{ - "entity_id": entityID, - "limit": limit, - "offset": offset, - } - if err := rcm.callOut(ctx, session, roles.OpRetrieveAllRoles, params); err != nil { - return roles.RolePage{}, err - } - return rcm.svc.RetrieveAllRoles(ctx, session, entityID, limit, offset) -} - -func (rcm *RoleManagerCalloutMiddleware) ListAvailableActions(ctx context.Context, session authn.Session) ([]string, error) { - params := map[string]any{} - if err := rcm.callOut(ctx, session, roles.OpListAvailableActions, params); err != nil { - return []string{}, err - } - return rcm.svc.ListAvailableActions(ctx, session) -} - -func (rcm *RoleManagerCalloutMiddleware) RoleAddActions(ctx context.Context, session authn.Session, entityID, roleID string, actions []string) ([]string, error) { - params := map[string]any{ - "entity_id": entityID, - "role_id": roleID, - "actions": actions, - } - if err := rcm.callOut(ctx, session, roles.OpRoleAddActions, params); err != nil { - return []string{}, err - } - return rcm.svc.RoleAddActions(ctx, session, entityID, roleID, actions) -} - -func (rcm *RoleManagerCalloutMiddleware) RoleListActions(ctx context.Context, session authn.Session, entityID, roleID string) ([]string, error) { - params := map[string]any{ - "entity_id": entityID, - "role_id": roleID, - } - if err := rcm.callOut(ctx, session, roles.OpRoleListActions, params); err != nil { - return []string{}, err - } - return rcm.svc.RoleListActions(ctx, session, entityID, roleID) -} - -func (rcm *RoleManagerCalloutMiddleware) RoleCheckActionsExists(ctx context.Context, session authn.Session, entityID, roleID string, actions []string) (bool, error) { - params := map[string]any{ - "entity_id": entityID, - "role_id": roleID, - "actions": actions, - } - if err := rcm.callOut(ctx, session, roles.OpRoleCheckActionsExists, params); err != nil { - return false, err - } - return rcm.svc.RoleCheckActionsExists(ctx, session, entityID, roleID, actions) -} - -func (rcm *RoleManagerCalloutMiddleware) RoleRemoveActions(ctx context.Context, session authn.Session, entityID, roleID string, actions []string) error { - params := map[string]any{ - "entity_id": entityID, - "role_id": roleID, - "actions": actions, - } - if err := rcm.callOut(ctx, session, roles.OpRoleRemoveActions, params); err != nil { - return err - } - return rcm.svc.RoleRemoveActions(ctx, session, entityID, roleID, actions) -} - -func (rcm *RoleManagerCalloutMiddleware) RoleRemoveAllActions(ctx context.Context, session authn.Session, entityID, roleID string) error { - params := map[string]any{ - "entity_id": entityID, - "role_id": roleID, - } - if err := rcm.callOut(ctx, session, roles.OpRoleRemoveAllActions, params); err != nil { - return err - } - return rcm.svc.RoleRemoveAllActions(ctx, session, entityID, roleID) -} - -func (rcm *RoleManagerCalloutMiddleware) RoleAddMembers(ctx context.Context, session authn.Session, entityID, roleID string, members []string) ([]string, error) { - params := map[string]any{ - "entity_id": entityID, - "role_id": roleID, - "members": members, - } - if err := rcm.callOut(ctx, session, roles.OpRoleAddMembers, params); err != nil { - return []string{}, err - } - return rcm.svc.RoleAddMembers(ctx, session, entityID, roleID, members) -} - -func (rcm *RoleManagerCalloutMiddleware) RoleListMembers(ctx context.Context, session authn.Session, entityID, roleID string, limit, offset uint64) (roles.MembersPage, error) { - params := map[string]any{ - "entity_id": entityID, - "role_id": roleID, - "limit": limit, - "offset": offset, - } - if err := rcm.callOut(ctx, session, roles.OpRoleListMembers, params); err != nil { - return roles.MembersPage{}, err - } - return rcm.svc.RoleListMembers(ctx, session, entityID, roleID, limit, offset) -} - -func (rcm *RoleManagerCalloutMiddleware) RoleCheckMembersExists(ctx context.Context, session authn.Session, entityID, roleID string, members []string) (bool, error) { - params := map[string]any{ - "entity_id": entityID, - "role_id": roleID, - "members": members, - } - if err := rcm.callOut(ctx, session, roles.OpRoleCheckMembersExists, params); err != nil { - return false, err - } - return rcm.svc.RoleCheckMembersExists(ctx, session, entityID, roleID, members) -} - -func (rcm *RoleManagerCalloutMiddleware) RoleRemoveAllMembers(ctx context.Context, session authn.Session, entityID, roleID string) error { - params := map[string]any{ - "entity_id": entityID, - "role_id": roleID, - } - if err := rcm.callOut(ctx, session, roles.OpRoleRemoveAllMembers, params); err != nil { - return err - } - return rcm.svc.RoleRemoveAllMembers(ctx, session, entityID, roleID) -} - -func (rcm *RoleManagerCalloutMiddleware) ListEntityMembers(ctx context.Context, session authn.Session, entityID string, pageQuery roles.MembersRolePageQuery) (roles.MembersRolePage, error) { - params := map[string]any{ - "entity_id": entityID, - "page_query": pageQuery, - } - if err := rcm.callOut(ctx, session, roles.OpRoleListMembers, params); err != nil { - return roles.MembersRolePage{}, err - } - return rcm.svc.ListEntityMembers(ctx, session, entityID, pageQuery) -} - -func (rcm *RoleManagerCalloutMiddleware) RemoveEntityMembers(ctx context.Context, session authn.Session, entityID string, members []string) error { - params := map[string]any{ - "entity_id": entityID, - "members": members, - } - if err := rcm.callOut(ctx, session, roles.OpRoleRemoveAllMembers, params); err != nil { - return err - } - return rcm.svc.RemoveEntityMembers(ctx, session, entityID, members) -} - -func (rcm *RoleManagerCalloutMiddleware) RoleRemoveMembers(ctx context.Context, session authn.Session, entityID, roleID string, members []string) error { - params := map[string]any{ - "entity_id": entityID, - "role_id": roleID, - "members": members, - } - if err := rcm.callOut(ctx, session, roles.OpRoleRemoveMembers, params); err != nil { - return err - } - return rcm.svc.RoleRemoveMembers(ctx, session, entityID, roleID, members) -} - -func (rcm *RoleManagerCalloutMiddleware) RemoveMemberFromAllRoles(ctx context.Context, session authn.Session, memberID string) error { - return rcm.svc.RemoveMemberFromAllRoles(ctx, session, memberID) -} - -func (rcm *RoleManagerCalloutMiddleware) callOut(ctx context.Context, session authn.Session, op permissions.RoleOperation, pld map[string]any) error { - var entityID string - if id, ok := pld["entity_id"].(string); ok { - entityID = id - } - - req := callout.Request{ - BaseRequest: callout.BaseRequest{ - Operation: rcm.roleOps.OperationName(op), - EntityType: rcm.entityType, - EntityID: entityID, - CallerID: session.UserID, - CallerType: policies.UserType, - DomainID: session.DomainID, - Time: time.Now().UTC(), - }, - Payload: pld, - } - - if err := rcm.callout.Callout(ctx, req); err != nil { - return err - } - - return nil -} diff --git a/pkg/roles/rolemanager/middleware/doc.go b/pkg/roles/rolemanager/middleware/doc.go deleted file mode 100644 index d8a741174..000000000 --- a/pkg/roles/rolemanager/middleware/doc.go +++ /dev/null @@ -1,9 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package middleware provides authorization, logging, metrics and tracing middleware -// for Magistrala RoleManager service. -// -// For more details about tracing instrumentation for Magistrala refer to the -// documentation at https://magistrala.absmach.eu/docs/. -package middleware diff --git a/pkg/roles/rolemanager/middleware/logging.go b/pkg/roles/rolemanager/middleware/logging.go deleted file mode 100644 index e0ff4e4ce..000000000 --- a/pkg/roles/rolemanager/middleware/logging.go +++ /dev/null @@ -1,388 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -//go:build !test - -package middleware - -import ( - "context" - "fmt" - "log/slog" - "time" - - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/roles" -) - -var _ roles.RoleManager = (*RoleManagerLoggingMiddleware)(nil) - -type RoleManagerLoggingMiddleware struct { - svcName string - svc roles.RoleManager - logger *slog.Logger -} - -// NewLogging adds logging facilities to the core service. -func NewLogging(svcName string, svc roles.RoleManager, logger *slog.Logger) RoleManagerLoggingMiddleware { - return RoleManagerLoggingMiddleware{ - svcName: svcName, - svc: svc, - logger: logger, - } -} - -func (lm *RoleManagerLoggingMiddleware) AddRole(ctx context.Context, session authn.Session, entityID, roleName string, optionalActions []string, optionalMembers []string) (ro roles.RoleProvision, err error) { - prefix := fmt.Sprintf("Add %s roles", lm.svcName) - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.Group(lm.svcName+"_add_role", - slog.String("entity_id", entityID), - slog.String("role_name", roleName), - slog.Any("optional_actions", optionalActions), - slog.Any("optional_members", optionalMembers), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn(prefix+" failed", args...) - return - } - lm.logger.Info(prefix+" completed successfully", args...) - }(time.Now()) - return lm.svc.AddRole(ctx, session, entityID, roleName, optionalActions, optionalMembers) -} - -func (lm *RoleManagerLoggingMiddleware) RemoveRole(ctx context.Context, session authn.Session, entityID, roleID string) (err error) { - prefix := fmt.Sprintf("Delete %s role", lm.svcName) - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.Group(lm.svcName+"_delete_role", - slog.String("entity_id", entityID), - slog.String("role_id", roleID), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn(prefix+" failed", args...) - return - } - lm.logger.Info(prefix+" completed successfully", args...) - }(time.Now()) - return lm.svc.RemoveRole(ctx, session, entityID, roleID) -} - -func (lm *RoleManagerLoggingMiddleware) UpdateRoleName(ctx context.Context, session authn.Session, entityID, roleID, newRoleName string) (ro roles.Role, err error) { - prefix := fmt.Sprintf("Update %s role name", lm.svcName) - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.Group(lm.svcName+"_update_role_name", - slog.String("entity_id", entityID), - slog.String("role_id", roleID), - slog.String("new_role_name", newRoleName), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn(prefix+" failed", args...) - return - } - lm.logger.Info(prefix+" completed successfully", args...) - }(time.Now()) - return lm.svc.UpdateRoleName(ctx, session, entityID, roleID, newRoleName) -} - -func (lm *RoleManagerLoggingMiddleware) RetrieveRole(ctx context.Context, session authn.Session, entityID, roleID string) (ro roles.Role, err error) { - prefix := fmt.Sprintf("Retrieve %s role", lm.svcName) - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.Group(lm.svcName+"_retrieve_role", - slog.String("entity_id", entityID), - slog.String("role_id", roleID), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn(prefix+" failed", args...) - return - } - lm.logger.Info(prefix+" completed successfully", args...) - }(time.Now()) - return lm.svc.RetrieveRole(ctx, session, entityID, roleID) -} - -func (lm *RoleManagerLoggingMiddleware) RetrieveAllRoles(ctx context.Context, session authn.Session, entityID string, limit, offset uint64) (rp roles.RolePage, err error) { - prefix := fmt.Sprintf("List %s roles", lm.svcName) - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.Group(lm.svcName+"_roles_retrieve_all", - slog.String("entity_id", entityID), - slog.Uint64("limit", limit), - slog.Uint64("offset", offset), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn(prefix+" failed", args...) - return - } - lm.logger.Info(prefix+" completed successfully", args...) - }(time.Now()) - return lm.svc.RetrieveAllRoles(ctx, session, entityID, limit, offset) -} - -func (lm *RoleManagerLoggingMiddleware) ListAvailableActions(ctx context.Context, session authn.Session) (acts []string, err error) { - prefix := fmt.Sprintf("List %s available actions", lm.svcName) - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.Group(lm.svcName + "_list_available_actions"), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn(prefix+" failed", args...) - return - } - lm.logger.Info(prefix+" completed successfully", args...) - }(time.Now()) - return lm.svc.ListAvailableActions(ctx, session) -} - -func (lm *RoleManagerLoggingMiddleware) RoleAddActions(ctx context.Context, session authn.Session, entityID, roleID string, actions []string) (caps []string, err error) { - prefix := fmt.Sprintf("%s role add actions", lm.svcName) - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.Group(lm.svcName+"_role_add_actions", - slog.String("entity_id", entityID), - slog.String("role_id", roleID), - slog.Any("actions", actions), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn(prefix+" failed", args...) - return - } - lm.logger.Info(prefix+" completed successfully", args...) - }(time.Now()) - return lm.svc.RoleAddActions(ctx, session, entityID, roleID, actions) -} - -func (lm *RoleManagerLoggingMiddleware) RoleListActions(ctx context.Context, session authn.Session, entityID, roleID string) (roOps []string, err error) { - prefix := fmt.Sprintf("%s role list actions", lm.svcName) - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.Group(lm.svcName+"_list_role_actions", - slog.String("entity_id", entityID), - slog.String("role_id", roleID), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn(prefix+" failed", args...) - return - } - lm.logger.Info(prefix+" completed successfully", args...) - }(time.Now()) - return lm.svc.RoleListActions(ctx, session, entityID, roleID) -} - -func (lm *RoleManagerLoggingMiddleware) RoleCheckActionsExists(ctx context.Context, session authn.Session, entityID, roleID string, actions []string) (bool, error) { - return lm.svc.RoleCheckActionsExists(ctx, session, entityID, roleID, actions) -} - -func (lm *RoleManagerLoggingMiddleware) RoleRemoveActions(ctx context.Context, session authn.Session, entityID, roleID string, actions []string) (err error) { - prefix := fmt.Sprintf("%s role remove actions", lm.svcName) - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.Group(lm.svcName+"_role_remove_actions", - slog.String("entity_id", entityID), - slog.String("role_id", roleID), - slog.Any("actions", actions), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn(prefix+" failed", args...) - return - } - lm.logger.Info(prefix+" completed successfully", args...) - }(time.Now()) - return lm.svc.RoleRemoveActions(ctx, session, entityID, roleID, actions) -} - -func (lm *RoleManagerLoggingMiddleware) RoleRemoveAllActions(ctx context.Context, session authn.Session, entityID, roleID string) (err error) { - prefix := fmt.Sprintf("%s role remove all actions", lm.svcName) - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.Group(lm.svcName+"_role_remove_all_actions", - slog.String("entity_id", entityID), - slog.String("role_id", roleID), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn(prefix+" failed", args...) - return - } - lm.logger.Info(prefix+" completed successfully", args...) - }(time.Now()) - return lm.svc.RoleRemoveAllActions(ctx, session, entityID, roleID) -} - -func (lm *RoleManagerLoggingMiddleware) RoleAddMembers(ctx context.Context, session authn.Session, entityID, roleID string, members []string) (mems []string, err error) { - prefix := fmt.Sprintf("%s role add members", lm.svcName) - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.Group(lm.svcName+"_role_add_members", - slog.String("entity_id", entityID), - slog.String("role_id", roleID), - slog.Any("members", members), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn(prefix+" failed", args...) - return - } - lm.logger.Info(prefix+" completed successfully", args...) - }(time.Now()) - return lm.svc.RoleAddMembers(ctx, session, entityID, roleID, members) -} - -func (lm *RoleManagerLoggingMiddleware) RoleListMembers(ctx context.Context, session authn.Session, entityID, roleID string, limit, offset uint64) (mp roles.MembersPage, err error) { - prefix := fmt.Sprintf("%s role list members", lm.svcName) - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.Group(lm.svcName+"_role_add_members", - slog.String("entity_id", entityID), - slog.String("role_id", roleID), - slog.Uint64("limit", limit), - slog.Uint64("offset", offset), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn(prefix+" failed", args...) - return - } - lm.logger.Info(prefix+" completed successfully", args...) - }(time.Now()) - return lm.svc.RoleListMembers(ctx, session, entityID, roleID, limit, offset) -} - -func (lm *RoleManagerLoggingMiddleware) RoleCheckMembersExists(ctx context.Context, session authn.Session, entityID, roleID string, members []string) (bool, error) { - return lm.svc.RoleCheckMembersExists(ctx, session, entityID, roleID, members) -} - -func (lm *RoleManagerLoggingMiddleware) RoleRemoveMembers(ctx context.Context, session authn.Session, entityID, roleID string, members []string) (err error) { - prefix := fmt.Sprintf("%s role remove members", lm.svcName) - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.Group(lm.svcName+"_role_remove_members", - slog.String("entity_id", entityID), - slog.String("role_id", roleID), - slog.Any("members", members), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn(prefix+" failed", args...) - return - } - lm.logger.Info(prefix+" completed successfully", args...) - }(time.Now()) - return lm.svc.RoleRemoveMembers(ctx, session, entityID, roleID, members) -} - -func (lm *RoleManagerLoggingMiddleware) RoleRemoveAllMembers(ctx context.Context, session authn.Session, entityID, roleID string) (err error) { - prefix := fmt.Sprintf("%s role remove all members", lm.svcName) - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.Group(lm.svcName+"_role_remove_all_members", - slog.String("entity_id", entityID), - slog.String("role_id", roleID), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn(prefix+" failed", args...) - return - } - lm.logger.Info(prefix+" completed successfully", args...) - }(time.Now()) - return lm.svc.RoleRemoveAllMembers(ctx, session, entityID, roleID) -} - -func (lm *RoleManagerLoggingMiddleware) ListEntityMembers(ctx context.Context, session authn.Session, entityID string, pageQuery roles.MembersRolePageQuery) (mems roles.MembersRolePage, err error) { - prefix := fmt.Sprintf("%s list entity members", lm.svcName) - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.Group(lm.svcName+"_remove_entity_members", - slog.String("entity_id", entityID), - slog.Uint64("limit", pageQuery.Limit), - slog.Uint64("offset", pageQuery.Offset), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn(prefix+" failed", args...) - return - } - lm.logger.Info(prefix+" completed successfully", args...) - }(time.Now()) - return lm.svc.ListEntityMembers(ctx, session, entityID, pageQuery) -} - -func (lm *RoleManagerLoggingMiddleware) RemoveEntityMembers(ctx context.Context, session authn.Session, entityID string, members []string) (err error) { - prefix := fmt.Sprintf("%s remove entity members", lm.svcName) - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.Group(lm.svcName+"_remove_entity_members", - slog.String("entity_id", entityID), - slog.Any("member_ids", members), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn(prefix+" failed", args...) - return - } - lm.logger.Info(prefix+" completed successfully", args...) - }(time.Now()) - return lm.svc.RemoveEntityMembers(ctx, session, entityID, members) -} - -func (lm *RoleManagerLoggingMiddleware) RemoveMemberFromAllRoles(ctx context.Context, session authn.Session, memberID string) (err error) { - prefix := fmt.Sprintf("%s remove members from all roles", lm.svcName) - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.Group(lm.svcName+"_remove_members_from_all_roles", - slog.Any("member_id", memberID), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn(prefix+" failed", args...) - return - } - lm.logger.Info(prefix+" completed successfully", args...) - }(time.Now()) - return lm.svc.RemoveMemberFromAllRoles(ctx, session, memberID) -} diff --git a/pkg/roles/rolemanager/middleware/meterics.go b/pkg/roles/rolemanager/middleware/meterics.go deleted file mode 100644 index 13ba45b70..000000000 --- a/pkg/roles/rolemanager/middleware/meterics.go +++ /dev/null @@ -1,109 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -//go:build !test - -package middleware - -import ( - "context" - - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/roles" - "github.com/go-kit/kit/metrics" -) - -var _ roles.RoleManager = (*RoleManagerMetricsMiddleware)(nil) - -type RoleManagerMetricsMiddleware struct { - svcName string - svc roles.RoleManager - counter metrics.Counter - latency metrics.Histogram -} - -// NewMetrics instruments core service by tracking request count and latency. -func NewMetrics(svcName string, svc roles.RoleManager, counter metrics.Counter, latency metrics.Histogram) RoleManagerMetricsMiddleware { - return RoleManagerMetricsMiddleware{ - svcName: svcName, - svc: svc, - counter: counter, - latency: latency, - } -} - -func (rmm *RoleManagerMetricsMiddleware) AddRole(ctx context.Context, session authn.Session, entityID, roleName string, optionalActions []string, optionalMembers []string) (roles.RoleProvision, error) { - return rmm.svc.AddRole(ctx, session, entityID, roleName, optionalActions, optionalMembers) -} - -func (rmm *RoleManagerMetricsMiddleware) RemoveRole(ctx context.Context, session authn.Session, entityID, roleID string) error { - return rmm.svc.RemoveRole(ctx, session, entityID, roleID) -} - -func (rmm *RoleManagerMetricsMiddleware) UpdateRoleName(ctx context.Context, session authn.Session, entityID, roleID, newRoleName string) (roles.Role, error) { - return rmm.svc.UpdateRoleName(ctx, session, entityID, roleID, newRoleName) -} - -func (rmm *RoleManagerMetricsMiddleware) RetrieveRole(ctx context.Context, session authn.Session, entityID, roleID string) (roles.Role, error) { - return rmm.svc.RetrieveRole(ctx, session, entityID, roleID) -} - -func (rmm *RoleManagerMetricsMiddleware) RetrieveAllRoles(ctx context.Context, session authn.Session, entityID string, limit, offset uint64) (roles.RolePage, error) { - return rmm.svc.RetrieveAllRoles(ctx, session, entityID, limit, offset) -} - -func (rmm *RoleManagerMetricsMiddleware) ListAvailableActions(ctx context.Context, session authn.Session) ([]string, error) { - return rmm.svc.ListAvailableActions(ctx, session) -} - -func (rmm *RoleManagerMetricsMiddleware) RoleAddActions(ctx context.Context, session authn.Session, entityID, roleID string, actions []string) (caps []string, err error) { - return rmm.svc.RoleAddActions(ctx, session, entityID, roleID, actions) -} - -func (rmm *RoleManagerMetricsMiddleware) RoleListActions(ctx context.Context, session authn.Session, entityID, roleID string) ([]string, error) { - return rmm.svc.RoleListActions(ctx, session, entityID, roleID) -} - -func (rmm *RoleManagerMetricsMiddleware) RoleCheckActionsExists(ctx context.Context, session authn.Session, entityID, roleID string, actions []string) (bool, error) { - return rmm.svc.RoleCheckActionsExists(ctx, session, entityID, roleID, actions) -} - -func (rmm *RoleManagerMetricsMiddleware) RoleRemoveActions(ctx context.Context, session authn.Session, entityID, roleID string, actions []string) (err error) { - return rmm.svc.RoleRemoveActions(ctx, session, entityID, roleID, actions) -} - -func (rmm *RoleManagerMetricsMiddleware) RoleRemoveAllActions(ctx context.Context, session authn.Session, entityID, roleID string) error { - return rmm.svc.RoleRemoveAllActions(ctx, session, entityID, roleID) -} - -func (rmm *RoleManagerMetricsMiddleware) RoleAddMembers(ctx context.Context, session authn.Session, entityID, roleID string, members []string) ([]string, error) { - return rmm.svc.RoleAddMembers(ctx, session, entityID, roleID, members) -} - -func (rmm *RoleManagerMetricsMiddleware) RoleListMembers(ctx context.Context, session authn.Session, entityID, roleID string, limit, offset uint64) (roles.MembersPage, error) { - return rmm.svc.RoleListMembers(ctx, session, entityID, roleID, limit, offset) -} - -func (rmm *RoleManagerMetricsMiddleware) RoleCheckMembersExists(ctx context.Context, session authn.Session, entityID, roleID string, members []string) (bool, error) { - return rmm.svc.RoleCheckMembersExists(ctx, session, entityID, roleID, members) -} - -func (rmm *RoleManagerMetricsMiddleware) RoleRemoveMembers(ctx context.Context, session authn.Session, entityID, roleID string, members []string) (err error) { - return rmm.svc.RoleRemoveMembers(ctx, session, entityID, roleID, members) -} - -func (rmm *RoleManagerMetricsMiddleware) RoleRemoveAllMembers(ctx context.Context, session authn.Session, entityID, roleID string) (err error) { - return rmm.svc.RoleRemoveAllMembers(ctx, session, entityID, roleID) -} - -func (rmm *RoleManagerMetricsMiddleware) ListEntityMembers(ctx context.Context, session authn.Session, entityID string, pageQuery roles.MembersRolePageQuery) (roles.MembersRolePage, error) { - return rmm.svc.ListEntityMembers(ctx, session, entityID, pageQuery) -} - -func (rmm *RoleManagerMetricsMiddleware) RemoveEntityMembers(ctx context.Context, session authn.Session, entityID string, members []string) error { - return rmm.svc.RemoveEntityMembers(ctx, session, entityID, members) -} - -func (rmm *RoleManagerMetricsMiddleware) RemoveMemberFromAllRoles(ctx context.Context, session authn.Session, memberID string) (err error) { - return rmm.svc.RemoveMemberFromAllRoles(ctx, session, memberID) -} diff --git a/pkg/roles/rolemanager/middleware/tracing.go b/pkg/roles/rolemanager/middleware/tracing.go deleted file mode 100644 index 2fd9af0b6..000000000 --- a/pkg/roles/rolemanager/middleware/tracing.go +++ /dev/null @@ -1,101 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package middleware - -import ( - "context" - - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/roles" - "go.opentelemetry.io/otel/trace" -) - -var _ roles.RoleManager = (*RoleManagerTracing)(nil) - -type RoleManagerTracing struct { - svcName string - roles roles.RoleManager - tracer trace.Tracer -} - -// NewTracing adds tracing facilities to the core service. -func NewTracing(svcName string, svc roles.RoleManager, tracer trace.Tracer) RoleManagerTracing { - return RoleManagerTracing{svcName, svc, tracer} -} - -func (rtm *RoleManagerTracing) AddRole(ctx context.Context, session authn.Session, entityID, roleName string, optionalActions []string, optionalMembers []string) (roles.RoleProvision, error) { - return rtm.roles.AddRole(ctx, session, entityID, roleName, optionalActions, optionalMembers) -} - -func (rtm *RoleManagerTracing) RemoveRole(ctx context.Context, session authn.Session, entityID, roleID string) error { - return rtm.roles.RemoveRole(ctx, session, entityID, roleID) -} - -func (rtm *RoleManagerTracing) UpdateRoleName(ctx context.Context, session authn.Session, entityID, roleID, newRoleName string) (roles.Role, error) { - return rtm.roles.UpdateRoleName(ctx, session, entityID, roleID, newRoleName) -} - -func (rtm *RoleManagerTracing) RetrieveRole(ctx context.Context, session authn.Session, entityID, roleID string) (roles.Role, error) { - return rtm.roles.RetrieveRole(ctx, session, entityID, roleID) -} - -func (rtm *RoleManagerTracing) RetrieveAllRoles(ctx context.Context, session authn.Session, entityID string, limit, offset uint64) (roles.RolePage, error) { - return rtm.roles.RetrieveAllRoles(ctx, session, entityID, limit, offset) -} - -func (rtm *RoleManagerTracing) ListAvailableActions(ctx context.Context, session authn.Session) ([]string, error) { - return rtm.roles.ListAvailableActions(ctx, session) -} - -func (rtm *RoleManagerTracing) RoleAddActions(ctx context.Context, session authn.Session, entityID, roleID string, actions []string) (ops []string, err error) { - return rtm.roles.RoleAddActions(ctx, session, entityID, roleID, actions) -} - -func (rtm *RoleManagerTracing) RoleListActions(ctx context.Context, session authn.Session, entityID, roleID string) ([]string, error) { - return rtm.roles.RoleListActions(ctx, session, entityID, roleID) -} - -func (rtm *RoleManagerTracing) RoleCheckActionsExists(ctx context.Context, session authn.Session, entityID, roleID string, actions []string) (bool, error) { - return rtm.roles.RoleCheckActionsExists(ctx, session, entityID, roleID, actions) -} - -func (rtm *RoleManagerTracing) RoleRemoveActions(ctx context.Context, session authn.Session, entityID, roleID string, actions []string) (err error) { - return rtm.roles.RoleRemoveActions(ctx, session, entityID, roleID, actions) -} - -func (rtm *RoleManagerTracing) RoleRemoveAllActions(ctx context.Context, session authn.Session, entityID, roleID string) error { - return rtm.roles.RoleRemoveAllActions(ctx, session, entityID, roleID) -} - -func (rtm *RoleManagerTracing) RoleAddMembers(ctx context.Context, session authn.Session, entityID, roleID string, members []string) ([]string, error) { - return rtm.roles.RoleAddMembers(ctx, session, entityID, roleID, members) -} - -func (rtm *RoleManagerTracing) RoleListMembers(ctx context.Context, session authn.Session, entityID, roleID string, limit, offset uint64) (roles.MembersPage, error) { - return rtm.roles.RoleListMembers(ctx, session, entityID, roleID, limit, offset) -} - -func (rtm *RoleManagerTracing) RoleCheckMembersExists(ctx context.Context, session authn.Session, entityID, roleID string, members []string) (bool, error) { - return rtm.roles.RoleCheckMembersExists(ctx, session, entityID, roleID, members) -} - -func (rtm *RoleManagerTracing) RoleRemoveMembers(ctx context.Context, session authn.Session, entityID, roleID string, members []string) (err error) { - return rtm.roles.RoleRemoveMembers(ctx, session, entityID, roleID, members) -} - -func (rtm *RoleManagerTracing) RoleRemoveAllMembers(ctx context.Context, session authn.Session, entityID, roleID string) (err error) { - return rtm.roles.RoleRemoveAllMembers(ctx, session, entityID, roleID) -} - -func (rtm *RoleManagerTracing) ListEntityMembers(ctx context.Context, session authn.Session, entityID string, pageQuery roles.MembersRolePageQuery) (roles.MembersRolePage, error) { - return rtm.roles.ListEntityMembers(ctx, session, entityID, pageQuery) -} - -func (rtm *RoleManagerTracing) RemoveEntityMembers(ctx context.Context, session authn.Session, entityID string, members []string) error { - return rtm.roles.RemoveEntityMembers(ctx, session, entityID, members) -} - -func (rtm *RoleManagerTracing) RemoveMemberFromAllRoles(ctx context.Context, session authn.Session, memberID string) (err error) { - return rtm.roles.RemoveMemberFromAllRoles(ctx, session, memberID) -} diff --git a/pkg/roles/roles.go b/pkg/roles/roles.go deleted file mode 100644 index a96be3945..000000000 --- a/pkg/roles/roles.go +++ /dev/null @@ -1,275 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package roles - -import ( - "context" - "time" - - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/permissions" - "github.com/absmach/magistrala/pkg/policies" -) - -type Action string - -func (ac Action) String() string { - return string(ac) -} - -type Member string - -func (mem Member) String() string { - return string(mem) -} - -type RoleName string - -func (r RoleName) String() string { - return string(r) -} - -type BuiltInRoleName RoleName - -func (b BuiltInRoleName) ToRoleName() RoleName { - return RoleName(b) -} - -func (b BuiltInRoleName) String() string { - return string(b) -} - -type Role struct { - ID string `json:"id"` - Name string `json:"name"` - EntityID string `json:"entity_id"` - CreatedBy string `json:"created_by"` - CreatedAt time.Time `json:"created_at"` - UpdatedBy string `json:"updated_by"` - UpdatedAt time.Time `json:"updated_at"` -} - -type RoleProvision struct { - Role - OptionalActions []string `json:"optional_actions"` - OptionalMembers []string `json:"optional_members"` -} - -type RolePage struct { - Total uint64 `json:"total"` - Offset uint64 `json:"offset"` - Limit uint64 `json:"limit"` - Roles []Role `json:"roles"` -} - -type MemberRoleActions struct { - RoleID string `json:"role_id"` - RoleName string `json:"role_name"` - Actions []string `json:"actions,omitempty"` - AccessProviderID string `json:"access_provider_id,omitempty"` - AccessProviderPath string `json:"access_provider_path,omitempty"` - AccessType string `json:"access_type,omitempty"` -} -type MemberRoles struct { - MemberID string `json:"member_id,omitempty"` - Roles []MemberRoleActions `json:"roles,omitempty"` -} - -type MembersRolePage struct { - Total uint64 `json:"total"` - Offset uint64 `json:"offset"` - Limit uint64 `json:"limit"` - Members []MemberRoles `json:"members"` -} - -type MembersRolePageQuery struct { - Total uint64 `json:"total"` - Offset uint64 `json:"offset"` - Limit uint64 `json:"limit"` - Order string `json:"order_by"` - Dir string `json:"dir"` - AccessProviderID string `json:"access_provider_id"` - RoleID string `json:"role_id"` - RoleName string `json:"role_name"` - Actions []string `json:"actions"` - AccessType string `json:"access_type"` -} - -type MembersPage struct { - Total uint64 `json:"total"` - Offset uint64 `json:"offset"` - Limit uint64 `json:"limit"` - Members []string `json:"members"` -} - -type EntityActionRole struct { - EntityID string `json:"entity_id"` - Action string `json:"action"` - RoleID string `json:"role_id"` -} -type EntityMemberRole struct { - EntityID string `json:"entity_id"` - MemberID string `json:"member_id"` - RoleID string `json:"role_id"` -} - -type Provisioner interface { - AddNewEntitiesRoles(ctx context.Context, domainID, userID string, entityIDs []string, optionalEntityPolicies []policies.Policy, newBuiltInRoleMembers map[BuiltInRoleName][]Member) ([]RoleProvision, error) - RemoveEntitiesRoles(ctx context.Context, domainID, userID string, entityIDs []string, optionalFilterDeletePolicies []policies.Policy, optionalDeletePolicies []policies.Policy) error -} - -type RoleManager interface { - // Add New role to entity - AddRole(ctx context.Context, session authn.Session, entityID, roleName string, optionalActions []string, optionalMembers []string) (RoleProvision, error) - - // Remove removes the roles of entity. - RemoveRole(ctx context.Context, session authn.Session, entityID, roleID string) error - - // UpdateName update the name of the entity role. - UpdateRoleName(ctx context.Context, session authn.Session, entityID, roleID, newRoleName string) (Role, error) - - RetrieveRole(ctx context.Context, session authn.Session, entityID, roleID string) (Role, error) - - RetrieveAllRoles(ctx context.Context, session authn.Session, entityID string, limit, offset uint64) (RolePage, error) - - ListAvailableActions(ctx context.Context, session authn.Session) ([]string, error) - - RoleAddActions(ctx context.Context, session authn.Session, entityID, roleID string, actions []string) (ops []string, err error) - - RoleListActions(ctx context.Context, session authn.Session, entityID, roleID string) ([]string, error) - - RoleCheckActionsExists(ctx context.Context, session authn.Session, entityID, roleID string, actions []string) (bool, error) - - RoleRemoveActions(ctx context.Context, session authn.Session, entityID, roleID string, actions []string) (err error) - - RoleRemoveAllActions(ctx context.Context, session authn.Session, entityID, roleID string) error - - RoleAddMembers(ctx context.Context, session authn.Session, entityID, roleID string, members []string) ([]string, error) - - RoleListMembers(ctx context.Context, session authn.Session, entityID, roleID string, limit, offset uint64) (MembersPage, error) - - RoleCheckMembersExists(ctx context.Context, session authn.Session, entityID, roleID string, members []string) (bool, error) - - RoleRemoveMembers(ctx context.Context, session authn.Session, entityID, roleID string, members []string) (err error) - - RoleRemoveAllMembers(ctx context.Context, session authn.Session, entityID, roleID string) (err error) - - ListEntityMembers(ctx context.Context, session authn.Session, entityID string, pq MembersRolePageQuery) (MembersRolePage, error) - - RemoveEntityMembers(ctx context.Context, session authn.Session, entityID string, members []string) (err error) - - RemoveMemberFromAllRoles(ctx context.Context, session authn.Session, memberID string) (err error) -} - -type Repository interface { - AddRoles(ctx context.Context, rps []RoleProvision) ([]RoleProvision, error) - RemoveRoles(ctx context.Context, roleIDs []string) error - UpdateRole(ctx context.Context, ro Role) (Role, error) - RetrieveRole(ctx context.Context, roleID string) (Role, error) - RetrieveEntityRole(ctx context.Context, entityID, roleID string) (Role, error) - RetrieveAllRoles(ctx context.Context, entityID string, limit, offset uint64) (RolePage, error) - RoleAddActions(ctx context.Context, role Role, actions []string) (ops []string, err error) - RoleListActions(ctx context.Context, roleID string) ([]string, error) - RoleCheckActionsExists(ctx context.Context, roleID string, actions []string) (bool, error) - RoleRemoveActions(ctx context.Context, role Role, actions []string) (err error) - RoleRemoveAllActions(ctx context.Context, role Role) error - RoleAddMembers(ctx context.Context, role Role, members []string) ([]string, error) - RoleListMembers(ctx context.Context, roleID string, limit, offset uint64) (MembersPage, error) - RoleCheckMembersExists(ctx context.Context, roleID string, members []string) (bool, error) - RoleRemoveMembers(ctx context.Context, role Role, members []string) (err error) - RoleRemoveAllMembers(ctx context.Context, role Role) (err error) - RetrieveEntitiesRolesActionsMembers(ctx context.Context, entityIDs []string) ([]EntityActionRole, []EntityMemberRole, error) - ListEntityMembers(ctx context.Context, entityID string, pageQuery MembersRolePageQuery) (MembersRolePage, error) - RemoveEntityMembers(ctx context.Context, entityID string, members []string) error - RemoveMemberFromAllRoles(ctx context.Context, memberID string) (err error) -} - -const ( - OpAddRole permissions.RoleOperation = iota - OpRemoveRole - OpUpdateRoleName - OpRetrieveRole - OpRetrieveAllRoles - OpRoleAddActions - OpRoleListActions - OpRoleCheckActionsExists - OpRoleRemoveActions - OpRoleRemoveAllActions - OpRoleAddMembers - OpRoleListMembers - OpRoleCheckMembersExists - OpRoleRemoveMembers - OpRoleRemoveAllMembers - OpListAvailableActions -) - -func Operations() map[permissions.RoleOperation]permissions.OperationDetails { - ops := map[permissions.RoleOperation]permissions.OperationDetails{ - OpAddRole: { - Name: "add", - PermissionRequired: true, - }, - OpRemoveRole: { - Name: "remove", - PermissionRequired: true, - }, - OpUpdateRoleName: { - Name: "update", - PermissionRequired: true, - }, - OpRetrieveRole: { - Name: "retrieve", - PermissionRequired: true, - }, - OpRetrieveAllRoles: { - Name: "retrieve_all", - PermissionRequired: true, - }, - OpRoleAddActions: { - Name: "add_actions", - PermissionRequired: true, - }, - OpRoleListActions: { - Name: "list_actions", - PermissionRequired: true, - }, - OpRoleCheckActionsExists: { - Name: "check_actions_exists", - PermissionRequired: true, - }, - OpRoleRemoveActions: { - Name: "remove_actions", - PermissionRequired: true, - }, - OpRoleRemoveAllActions: { - Name: "remove_all_actions", - PermissionRequired: true, - }, - OpRoleAddMembers: { - Name: "add_members", - PermissionRequired: true, - }, - OpRoleListMembers: { - Name: "list_members", - PermissionRequired: true, - }, - OpRoleCheckMembersExists: { - Name: "check_members_exists", - PermissionRequired: true, - }, - OpRoleRemoveMembers: { - Name: "remove_members", - PermissionRequired: true, - }, - OpRoleRemoveAllMembers: { - Name: "remove_all_members", - PermissionRequired: true, - }, - OpListAvailableActions: { - Name: "list_available_actions", - PermissionRequired: false, - }, - } - return ops -} diff --git a/pkg/sdk/alarms_test.go b/pkg/sdk/alarms_test.go deleted file mode 100644 index 81e4ac1a6..000000000 --- a/pkg/sdk/alarms_test.go +++ /dev/null @@ -1,390 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package sdk_test - -import ( - "context" - "net/http/httptest" - "testing" - "time" - - "github.com/absmach/magistrala/alarms" - "github.com/absmach/magistrala/alarms/api" - amocks "github.com/absmach/magistrala/alarms/mocks" - mglog "github.com/absmach/magistrala/logger" - smqauthn "github.com/absmach/magistrala/pkg/authn" - authnmocks "github.com/absmach/magistrala/pkg/authn/mocks" - "github.com/absmach/magistrala/pkg/errors" - "github.com/absmach/magistrala/pkg/sdk" - "github.com/absmach/magistrala/pkg/uuid" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" -) - -const alarmID = "alarm-1" - -var testAlarm = sdk.Alarm{ - ID: alarmID, - RuleID: "rule-1", - DomainID: domainID, - ChannelID: "chan-1", - ClientID: "client-1", - Subtopic: "subtopic", - Status: "active", - Measurement: "temperature", - Value: "30.5", - Unit: "C", - Threshold: "25", - Cause: "threshold_exceeded", - Severity: 80, - AssigneeID: "user-1", - Metadata: sdk.Metadata{"key": "value"}, -} - -func setupAlarms() (*httptest.Server, *amocks.Service, *authnmocks.Authentication) { - asvc := new(amocks.Service) - logger := mglog.NewMock() - authn := new(authnmocks.Authentication) - am := smqauthn.NewAuthNMiddleware(authn, smqauthn.WithAllowUnverifiedUser(true)) - idp := uuid.NewMock() - mux := api.MakeHandler(asvc, logger, idp, "", am) - return httptest.NewServer(mux), asvc, authn -} - -func TestUpdateAlarm(t *testing.T) { - as, asvc, auth := setupAlarms() - defer as.Close() - - conf := sdk.Config{ - AlarmsURL: as.URL, - } - mgsdk := sdk.NewSDK(conf) - - updated := testAlarm - updated.Status = "cleared" - - svcAlarm := alarms.Alarm{ - ID: alarmID, - RuleID: "rule-1", - DomainID: domainID, - ChannelID: "chan-1", - ClientID: "client-1", - Subtopic: "subtopic", - Status: alarms.ClearedStatus, - Measurement: "temperature", - Value: "30.5", - Unit: "C", - Threshold: "25", - Cause: "threshold_exceeded", - Severity: 80, - AssigneeID: "user-1", - Metadata: alarms.Metadata{"key": "value"}, - } - - cases := []struct { - desc string - alarm sdk.Alarm - token string - session smqauthn.Session - svcRes alarms.Alarm - svcErr error - authenticateErr error - wantErr bool - resp sdk.Alarm - }{ - { - desc: "update alarm successfully", - alarm: updated, - token: validToken, - svcRes: svcAlarm, - resp: testAlarm, - }, - { - desc: "update alarm with empty token", - alarm: updated, - token: "", - wantErr: true, - }, - { - desc: "update non-existent alarm", - alarm: sdk.Alarm{ID: "non-existent"}, - token: validToken, - svcErr: errors.New("not found"), - wantErr: true, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := asvc.On("UpdateAlarm", mock.Anything, tc.session, mock.Anything).Return(tc.svcRes, tc.svcErr) - result, err := mgsdk.UpdateAlarm(context.Background(), tc.alarm, domainID, tc.token) - assert.Equal(t, tc.wantErr, err != nil) - if !tc.wantErr { - assert.NotEmpty(t, result.ID) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestViewAlarm(t *testing.T) { - as, asvc, auth := setupAlarms() - defer as.Close() - - conf := sdk.Config{ - AlarmsURL: as.URL, - } - mgsdk := sdk.NewSDK(conf) - - svcAlarm := alarms.Alarm{ - ID: alarmID, - RuleID: "rule-1", - DomainID: domainID, - ChannelID: "chan-1", - ClientID: "client-1", - Subtopic: "subtopic", - Status: alarms.ActiveStatus, - Measurement: "temperature", - Value: "30.5", - Unit: "C", - Threshold: "25", - Cause: "threshold_exceeded", - Severity: 80, - AssigneeID: "user-1", - Metadata: alarms.Metadata{"key": "value"}, - } - - cases := []struct { - desc string - id string - token string - session smqauthn.Session - svcRes alarms.Alarm - svcErr error - authenticateErr error - wantErr bool - }{ - { - desc: "view alarm successfully", - id: alarmID, - token: validToken, - svcRes: svcAlarm, - }, - { - desc: "view alarm with empty token", - id: alarmID, - token: "", - wantErr: true, - }, - { - desc: "view non-existent alarm", - id: "non-existent", - token: validToken, - svcErr: errors.New("not found"), - wantErr: true, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := asvc.On("ViewAlarm", mock.Anything, tc.session, tc.id).Return(tc.svcRes, tc.svcErr) - result, err := mgsdk.ViewAlarm(context.Background(), tc.id, domainID, tc.token) - assert.Equal(t, tc.wantErr, err != nil) - if !tc.wantErr { - assert.NotEmpty(t, result.ID) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestListAlarms(t *testing.T) { - as, asvc, auth := setupAlarms() - defer as.Close() - - conf := sdk.Config{ - AlarmsURL: as.URL, - } - mgsdk := sdk.NewSDK(conf) - - svcAlarm := alarms.Alarm{ - ID: alarmID, - RuleID: "rule-1", - DomainID: domainID, - ChannelID: "chan-1", - ClientID: "client-1", - Subtopic: "subtopic", - Status: alarms.ActiveStatus, - Measurement: "temperature", - Value: "30.5", - Unit: "C", - Threshold: "25", - Cause: "threshold_exceeded", - Severity: 80, - AssigneeID: "user-1", - Metadata: alarms.Metadata{"key": "value"}, - } - - svcAlarmsPage := alarms.AlarmsPage{ - Total: 2, - Offset: 0, - Limit: 10, - Alarms: []alarms.Alarm{svcAlarm}, - } - - cases := []struct { - desc string - pm sdk.PageMetadata - token string - session smqauthn.Session - svcRes alarms.AlarmsPage - svcErr error - authenticateErr error - wantErr bool - }{ - { - desc: "list alarms successfully", - pm: sdk.PageMetadata{Offset: 0, Limit: 10}, - token: validToken, - svcRes: svcAlarmsPage, - }, - { - desc: "list alarms with status and entity filters", - pm: sdk.PageMetadata{ - Limit: 5, - Status: "active", - ChannelID: "chan-1", - ClientID: "client-1", - RuleID: "rule-1", - AssigneeID: "user-1", - Severity: 80, - }, - token: validToken, - svcRes: svcAlarmsPage, - }, - { - desc: "list alarms with time range and sorting", - pm: sdk.PageMetadata{ - Limit: 10, - CreatedFrom: time.Date(2024, 1, 1, 0, 0, 0, 0, time.UTC), - CreatedTo: time.Date(2024, 12, 31, 0, 0, 0, 0, time.UTC), - Order: "created_at", - Dir: "asc", - }, - token: validToken, - svcRes: svcAlarmsPage, - }, - { - desc: "list alarms with actor filters", - pm: sdk.PageMetadata{ - Limit: 10, - UpdatedBy: "user-2", - AssignedBy: "user-3", - AcknowledgedBy: "user-4", - ResolvedBy: "user-5", - Subtopic: "subtopic-1", - }, - token: validToken, - svcRes: svcAlarmsPage, - }, - { - desc: "list alarms with empty metadata excludes severity", - pm: sdk.PageMetadata{}, - token: validToken, - svcRes: alarms.AlarmsPage{}, - }, - { - desc: "list alarms with zero severity excluded", - pm: sdk.PageMetadata{Status: "active", Severity: 0}, - token: validToken, - svcRes: alarms.AlarmsPage{}, - }, - { - desc: "list alarms with empty token", - pm: sdk.PageMetadata{Limit: 10}, - token: "", - wantErr: true, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := asvc.On("ListAlarms", mock.Anything, tc.session, mock.Anything).Return(tc.svcRes, tc.svcErr) - result, err := mgsdk.ListAlarms(context.Background(), tc.pm, domainID, tc.token) - assert.Equal(t, tc.wantErr, err != nil) - if !tc.wantErr { - assert.Equal(t, tc.svcRes.Total, result.Total) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestDeleteAlarm(t *testing.T) { - as, asvc, auth := setupAlarms() - defer as.Close() - - conf := sdk.Config{ - AlarmsURL: as.URL, - } - mgsdk := sdk.NewSDK(conf) - - cases := []struct { - desc string - id string - token string - session smqauthn.Session - svcErr error - authenticateErr error - wantErr bool - }{ - { - desc: "delete alarm successfully", - id: alarmID, - token: validToken, - }, - { - desc: "delete alarm with empty token", - id: alarmID, - token: "", - wantErr: true, - }, - { - desc: "delete non-existent alarm", - id: "non-existent", - token: validToken, - svcErr: errors.New("not found"), - wantErr: true, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := asvc.On("DeleteAlarm", mock.Anything, tc.session, tc.id).Return(tc.svcErr) - err := mgsdk.DeleteAlarm(context.Background(), tc.id, domainID, tc.token) - assert.Equal(t, tc.wantErr, err != nil) - svcCall.Unset() - authCall.Unset() - }) - } -} diff --git a/pkg/sdk/bootstrap_test.go b/pkg/sdk/bootstrap_test.go deleted file mode 100644 index bb85eca18..000000000 --- a/pkg/sdk/bootstrap_test.go +++ /dev/null @@ -1,1230 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package sdk_test - -import ( - "context" - "crypto/aes" - "crypto/cipher" - "crypto/rand" - "encoding/json" - "fmt" - "io" - "net/http" - "net/http/httptest" - "testing" - - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/bootstrap" - "github.com/absmach/magistrala/bootstrap/api" - bmocks "github.com/absmach/magistrala/bootstrap/mocks" - "github.com/absmach/magistrala/internal/testsutil" - mglog "github.com/absmach/magistrala/logger" - smqauthn "github.com/absmach/magistrala/pkg/authn" - authnmocks "github.com/absmach/magistrala/pkg/authn/mocks" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - sdk "github.com/absmach/magistrala/pkg/sdk" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" -) - -var ( - externalId = testsutil.GenerateUUID(&testing.T{}) - externalKey = testsutil.GenerateUUID(&testing.T{}) - clientId = testsutil.GenerateUUID(&testing.T{}) - channel1Id = testsutil.GenerateUUID(&testing.T{}) - channel2Id = testsutil.GenerateUUID(&testing.T{}) - clientCert = "newcert" - clientKey = "newkey" - caCert = "newca" - content = "newcontent" - bsName = "test" - encKey = []byte("1234567891011121") - - bootstrapConfig = bootstrap.Config{ - ID: clientId, - Name: bsName, - ClientCert: clientCert, - ClientKey: clientKey, - CACert: caCert, - ExternalID: externalId, - ExternalKey: externalKey, - Content: content, - Status: bootstrap.Inactive, - } - - sdkBootstrapConfig = sdk.BootstrapConfig{ - ExternalID: externalId, - ExternalKey: externalKey, - ID: clientId, - Name: bsName, - ClientCert: clientCert, - ClientKey: clientKey, - CACert: caCert, - Content: content, - Status: sdk.BootstrapDisabledStatus, - } - - sdkBootstrapListRes = sdk.BootstrapConfig{ - ID: clientId, - ExternalID: externalId, - Name: bsName, - Content: content, - Status: sdk.BootstrapDisabledStatus, - } - - sdkBootstrapCertRes = sdk.BootstrapConfig{ - ID: clientId, - ClientCert: clientCert, - ClientKey: clientKey, - CACert: caCert, - } - - sdkBootstrapReadRes = sdk.BootstrapConfig{ - ID: clientId, - Content: content, - ClientCert: clientCert, - ClientKey: clientKey, - CACert: caCert, - } - - readConfigResponse = struct { - ID string `json:"id"` - Content string `json:"content,omitempty"` - ClientCert string `json:"client_cert,omitempty"` - ClientKey string `json:"client_key,omitempty"` - CACert string `json:"ca_cert,omitempty"` - }{ - ID: clientId, - Content: content, - ClientCert: clientCert, - ClientKey: clientKey, - CACert: caCert, - } -) - -var ( - errMarshalChan = errors.New("json: unsupported type: chan int") - errJSONEOF = errors.New("unexpected end of JSON input") -) - -func setupBootstrap() (*httptest.Server, *bmocks.Service, *bmocks.ConfigReader, *authnmocks.Authentication) { - bsvc := new(bmocks.Service) - reader := new(bmocks.ConfigReader) - logger := mglog.NewMock() - authn := new(authnmocks.Authentication) - am := smqauthn.NewAuthNMiddleware(authn, smqauthn.WithAllowUnverifiedUser(true)) - - mux := api.MakeHandler(bsvc, am, reader, logger, "") - - return httptest.NewServer(mux), bsvc, reader, authn -} - -func bootstrapSession() smqauthn.Session { - return smqauthn.Session{ - DomainUserID: domainID + "_" + validID, - UserID: validID, - DomainID: domainID, - } -} - -func TestAddBootstrap(t *testing.T) { - bs, bsvc, _, auth := setupBootstrap() - defer bs.Close() - - mgsdk := sdk.NewSDK(sdk.Config{BootstrapURL: bs.URL}) - - createCfg := sdkBootstrapConfig - createCfg.ID = "" - - svcReq := bootstrap.Config{ - ExternalID: externalId, - ExternalKey: externalKey, - Name: bsName, - ClientCert: clientCert, - ClientKey: clientKey, - CACert: caCert, - Content: content, - } - - cases := []struct { - desc string - token string - cfg sdk.BootstrapConfig - svcReq bootstrap.Config - svcRes bootstrap.Config - svcErr error - authErr error - expectSvcCall bool - expectedID string - expectedSDKErr errors.SDKError - }{ - { - desc: "add successfully", - token: validToken, - cfg: createCfg, - svcReq: svcReq, - svcRes: bootstrapConfig, - expectSvcCall: true, - expectedID: clientId, - }, - { - desc: "add with invalid token", - token: invalidToken, - cfg: createCfg, - authErr: svcerr.ErrAuthentication, - expectedSDKErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "add with config that cannot be marshalled", - token: validToken, - cfg: sdk.BootstrapConfig{ - RenderContext: map[string]any{ - "broken": make(chan int), - }, - ExternalID: externalId, - ExternalKey: externalKey, - }, - expectedSDKErr: errors.NewSDKError(errMarshalChan), - }, - { - desc: "add with missing required fields", - token: validToken, - cfg: sdk.BootstrapConfig{}, - expectedSDKErr: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "add with service failure", - token: validToken, - cfg: createCfg, - svcReq: svcReq, - svcRes: bootstrap.Config{}, - svcErr: svcerr.ErrNotFound, - expectSvcCall: true, - expectedSDKErr: errors.NewSDKErrorWithStatus(svcerr.ErrNotFound, http.StatusNotFound), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - session := smqauthn.Session{} - if tc.token == validToken { - session = bootstrapSession() - } - - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(session, tc.authErr) - - var svcCall *mock.Call - if tc.expectSvcCall { - svcCall = bsvc.On("Add", mock.Anything, session, tc.token, tc.svcReq).Return(tc.svcRes, tc.svcErr) - } - - resp, err := mgsdk.AddBootstrap(context.Background(), tc.cfg, domainID, tc.token) - - assert.Equal(t, tc.expectedSDKErr, err) - assert.Equal(t, tc.expectedID, resp) - if tc.expectSvcCall { - svcCall.Unset() - } - authCall.Unset() - }) - } -} - -func TestListBootstraps(t *testing.T) { - bs, bsvc, _, auth := setupBootstrap() - defer bs.Close() - - mgsdk := sdk.NewSDK(sdk.Config{BootstrapURL: bs.URL}) - - cases := []struct { - desc string - token string - pm sdk.PageMetadata - svcResp bootstrap.ConfigsPage - svcErr error - authErr error - expectSvcCall bool - expectedResp sdk.BootstrapPage - expectedSDKErr errors.SDKError - }{ - { - desc: "list successfully", - token: validToken, - pm: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - }, - svcResp: bootstrap.ConfigsPage{ - Total: 1, - Offset: 0, - Limit: 10, - Configs: []bootstrap.Config{bootstrapConfig}, - }, - expectSvcCall: true, - expectedResp: sdk.BootstrapPage{ - PageRes: sdk.PageRes{ - Total: 1, - Offset: 0, - Limit: 10, - }, - Configs: []sdk.BootstrapConfig{sdkBootstrapListRes}, - }, - }, - { - desc: "list with invalid token", - token: invalidToken, - pm: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - }, - authErr: svcerr.ErrAuthentication, - expectedSDKErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "list with invalid query params", - token: validToken, - pm: sdk.PageMetadata{ - Metadata: map[string]any{ - "test": make(chan int), - }, - }, - expectedSDKErr: errors.NewSDKError(errMarshalChan), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - var authCall *mock.Call - session := smqauthn.Session{} - if tc.token == validToken { - session = bootstrapSession() - } - if tc.expectedSDKErr == nil || tc.authErr != nil { - authCall = auth.On("Authenticate", mock.Anything, tc.token).Return(session, tc.authErr) - } - - var svcCall *mock.Call - if tc.expectSvcCall { - svcCall = bsvc.On("List", mock.Anything, session, mock.Anything, tc.pm.Offset, tc.pm.Limit).Return(tc.svcResp, tc.svcErr) - } - - resp, err := mgsdk.Bootstraps(context.Background(), tc.pm, domainID, tc.token) - - assert.Equal(t, tc.expectedSDKErr, err) - assert.Equal(t, tc.expectedResp, resp) - if svcCall != nil { - svcCall.Unset() - } - if authCall != nil { - authCall.Unset() - } - }) - } -} - -func TestWhitelist(t *testing.T) { - bs, bsvc, _, auth := setupBootstrap() - defer bs.Close() - - mgsdk := sdk.NewSDK(sdk.Config{BootstrapURL: bs.URL}) - - cases := []struct { - desc string - token string - clientID string - status sdk.BootstrapStatus - method string - svcResp bootstrap.Config - svcErr error - authErr error - expectSvcCall bool - expectedSDKErr errors.SDKError - }{ - { - desc: "enable bootstrap successfully", - token: validToken, - clientID: clientId, - status: sdk.BootstrapEnabledStatus, - method: "EnableConfig", - svcResp: bootstrap.Config{ID: clientId, Status: bootstrap.Active}, - expectSvcCall: true, - }, - { - desc: "disable bootstrap successfully", - token: validToken, - clientID: clientId, - status: sdk.BootstrapDisabledStatus, - method: "DisableConfig", - svcResp: bootstrap.Config{ID: clientId, Status: bootstrap.Inactive}, - expectSvcCall: true, - }, - { - desc: "whitelist with invalid token", - token: invalidToken, - clientID: clientId, - status: sdk.BootstrapEnabledStatus, - authErr: svcerr.ErrAuthentication, - expectedSDKErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "whitelist with invalid status", - token: validToken, - clientID: clientId, - status: sdk.BootstrapStatus("invalid"), - expectedSDKErr: errors.NewSDKErrorWithStatus(errors.New("invalid bootstrap status"), http.StatusBadRequest), - }, - { - desc: "whitelist with empty client id", - token: validToken, - clientID: "", - status: sdk.BootstrapEnabledStatus, - expectedSDKErr: errors.NewSDKError(apiutil.ErrMissingID), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - var authCall *mock.Call - session := smqauthn.Session{} - if tc.token == validToken { - session = bootstrapSession() - } - if tc.clientID != "" && (tc.status == sdk.BootstrapDisabledStatus || tc.status == sdk.BootstrapEnabledStatus) { - authCall = auth.On("Authenticate", mock.Anything, tc.token).Return(session, tc.authErr) - } - - var svcCall *mock.Call - if tc.expectSvcCall { - svcCall = bsvc.On(tc.method, mock.Anything, session, tc.clientID).Return(tc.svcResp, tc.svcErr) - } - - err := mgsdk.Whitelist(context.Background(), tc.clientID, tc.status, domainID, tc.token) - - assert.Equal(t, tc.expectedSDKErr, err) - if svcCall != nil { - svcCall.Unset() - } - if authCall != nil { - authCall.Unset() - } - }) - } -} - -func TestViewBootstrap(t *testing.T) { - bs, bsvc, _, auth := setupBootstrap() - defer bs.Close() - - mgsdk := sdk.NewSDK(sdk.Config{BootstrapURL: bs.URL}) - - cases := []struct { - desc string - token string - id string - svcResp bootstrap.Config - svcErr error - authErr error - expectSvcCall bool - expectedResp sdk.BootstrapConfig - expectedSDKErr errors.SDKError - }{ - { - desc: "view successfully", - token: validToken, - id: clientId, - svcResp: bootstrapConfig, - expectSvcCall: true, - expectedResp: sdkBootstrapListRes, - }, - { - desc: "view with invalid token", - token: invalidToken, - id: clientId, - authErr: svcerr.ErrAuthentication, - expectedSDKErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "view with non-existent client id", - token: validToken, - id: invalid, - svcResp: bootstrap.Config{}, - svcErr: svcerr.ErrNotFound, - expectSvcCall: true, - expectedSDKErr: errors.NewSDKErrorWithStatus(svcerr.ErrNotFound, http.StatusNotFound), - }, - { - desc: "view with empty client id", - token: validToken, - id: "", - expectedSDKErr: errors.NewSDKError(apiutil.ErrMissingID), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - var authCall *mock.Call - session := smqauthn.Session{} - if tc.token == validToken { - session = bootstrapSession() - } - if tc.id != "" { - authCall = auth.On("Authenticate", mock.Anything, tc.token).Return(session, tc.authErr) - } - - var svcCall *mock.Call - if tc.expectSvcCall { - svcCall = bsvc.On("View", mock.Anything, session, tc.id).Return(tc.svcResp, tc.svcErr) - } - - resp, err := mgsdk.ViewBootstrap(context.Background(), tc.id, domainID, tc.token) - - assert.Equal(t, tc.expectedSDKErr, err) - assert.Equal(t, tc.expectedResp, resp) - if svcCall != nil { - svcCall.Unset() - } - if authCall != nil { - authCall.Unset() - } - }) - } -} - -func TestUpdateBootstrap(t *testing.T) { - bs, bsvc, _, auth := setupBootstrap() - defer bs.Close() - - mgsdk := sdk.NewSDK(sdk.Config{BootstrapURL: bs.URL}) - - cases := []struct { - desc string - token string - cfg sdk.BootstrapConfig - svcReq bootstrap.Config - svcErr error - authErr error - expectSvcCall bool - expectedSDKErr errors.SDKError - }{ - { - desc: "update successfully", - token: validToken, - cfg: sdkBootstrapConfig, - svcReq: bootstrap.Config{ - ID: clientId, - Name: bsName, - Content: content, - }, - expectSvcCall: true, - }, - { - desc: "update with invalid token", - token: invalidToken, - cfg: sdkBootstrapConfig, - authErr: svcerr.ErrAuthentication, - expectedSDKErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "update with empty client id", - token: validToken, - cfg: sdk.BootstrapConfig{}, - expectedSDKErr: errors.NewSDKError(apiutil.ErrMissingID), - }, - { - desc: "update with config that cannot be marshalled", - token: validToken, - cfg: sdk.BootstrapConfig{ - ID: clientId, - RenderContext: map[string]any{ - "broken": make(chan int), - }, - }, - expectedSDKErr: errors.NewSDKError(errMarshalChan), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - var authCall *mock.Call - session := smqauthn.Session{} - if tc.token == validToken { - session = bootstrapSession() - } - if tc.cfg.ID != "" { - authCall = auth.On("Authenticate", mock.Anything, tc.token).Return(session, tc.authErr) - } - - var svcCall *mock.Call - if tc.expectSvcCall { - svcCall = bsvc.On("Update", mock.Anything, session, tc.svcReq).Return(tc.svcErr) - } - - err := mgsdk.UpdateBootstrap(context.Background(), tc.cfg, domainID, tc.token) - - assert.Equal(t, tc.expectedSDKErr, err) - if svcCall != nil { - svcCall.Unset() - } - if authCall != nil { - authCall.Unset() - } - }) - } -} - -func TestUpdateBootstrapCerts(t *testing.T) { - bs, bsvc, _, auth := setupBootstrap() - defer bs.Close() - - mgsdk := sdk.NewSDK(sdk.Config{BootstrapURL: bs.URL}) - - cases := []struct { - desc string - token string - id string - cert string - key string - ca string - svcResp bootstrap.Config - svcErr error - authErr error - expectSvcCall bool - expectedResp sdk.BootstrapConfig - expectedSDKErr errors.SDKError - }{ - { - desc: "update certs successfully", - token: validToken, - id: clientId, - cert: clientCert, - key: clientKey, - ca: caCert, - svcResp: bootstrapConfig, - expectSvcCall: true, - expectedResp: sdkBootstrapCertRes, - }, - { - desc: "update certs with invalid token", - token: invalidToken, - id: clientId, - cert: clientCert, - key: clientKey, - ca: caCert, - authErr: svcerr.ErrAuthentication, - expectedSDKErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "update certs with empty id", - token: validToken, - id: "", - cert: clientCert, - key: clientKey, - ca: caCert, - expectedSDKErr: errors.NewSDKError(apiutil.ErrMissingID), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - var authCall *mock.Call - session := smqauthn.Session{} - if tc.token == validToken { - session = bootstrapSession() - } - if tc.id != "" { - authCall = auth.On("Authenticate", mock.Anything, tc.token).Return(session, tc.authErr) - } - - var svcCall *mock.Call - if tc.expectSvcCall { - svcCall = bsvc.On("UpdateCert", mock.Anything, session, tc.id, tc.cert, tc.key, tc.ca).Return(tc.svcResp, tc.svcErr) - } - - resp, err := mgsdk.UpdateBootstrapCerts(context.Background(), tc.id, tc.cert, tc.key, tc.ca, domainID, tc.token) - - assert.Equal(t, tc.expectedSDKErr, err) - assert.Equal(t, tc.expectedResp, resp) - if svcCall != nil { - svcCall.Unset() - } - if authCall != nil { - authCall.Unset() - } - }) - } -} - -func TestUpdateBootstrapConnection(t *testing.T) { - mgsdk := sdk.NewSDK(sdk.Config{}) - - err := mgsdk.UpdateBootstrapConnection(context.Background(), clientId, []string{channel1Id, channel2Id}, domainID, validToken) - assert.Equal(t, errors.NewSDKError(errors.New("bootstrap connection updates are no longer supported")), err) - - err = mgsdk.UpdateBootstrapConnection(context.Background(), "", []string{channel1Id}, domainID, validToken) - assert.Equal(t, errors.NewSDKError(apiutil.ErrMissingID), err) -} - -func TestRemoveBootstrap(t *testing.T) { - bs, bsvc, _, auth := setupBootstrap() - defer bs.Close() - - mgsdk := sdk.NewSDK(sdk.Config{BootstrapURL: bs.URL}) - - cases := []struct { - desc string - token string - id string - svcErr error - authErr error - expectSvcCall bool - expectedSDKErr errors.SDKError - }{ - { - desc: "remove successfully", - token: validToken, - id: clientId, - expectSvcCall: true, - }, - { - desc: "remove with invalid token", - token: invalidToken, - id: clientId, - authErr: svcerr.ErrAuthentication, - expectedSDKErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "remove with empty id", - token: validToken, - id: "", - expectedSDKErr: errors.NewSDKError(apiutil.ErrMissingID), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - var authCall *mock.Call - session := smqauthn.Session{} - if tc.token == validToken { - session = bootstrapSession() - } - if tc.id != "" { - authCall = auth.On("Authenticate", mock.Anything, tc.token).Return(session, tc.authErr) - } - - var svcCall *mock.Call - if tc.expectSvcCall { - svcCall = bsvc.On("Remove", mock.Anything, session, tc.id).Return(tc.svcErr) - } - - err := mgsdk.RemoveBootstrap(context.Background(), tc.id, domainID, tc.token) - - assert.Equal(t, tc.expectedSDKErr, err) - if svcCall != nil { - svcCall.Unset() - } - if authCall != nil { - authCall.Unset() - } - }) - } -} - -func TestBootstrap(t *testing.T) { - bs, bsvc, reader, _ := setupBootstrap() - defer bs.Close() - - mgsdk := sdk.NewSDK(sdk.Config{BootstrapURL: bs.URL}) - - cases := []struct { - desc string - externalID string - externalKey string - svcResp bootstrap.Config - svcErr error - readerResp any - readerErr error - expectSvcCall bool - expectedResp sdk.BootstrapConfig - expectedSDKErr errors.SDKError - }{ - { - desc: "bootstrap successfully", - externalID: externalId, - externalKey: externalKey, - svcResp: bootstrapConfig, - readerResp: readConfigResponse, - expectSvcCall: true, - expectedResp: sdkBootstrapReadRes, - }, - { - desc: "bootstrap with reader error", - externalID: externalId, - externalKey: externalKey, - svcResp: bootstrapConfig, - readerErr: errJSONEOF, - expectSvcCall: true, - expectedSDKErr: errors.NewSDKErrorWithStatus(errJSONEOF, http.StatusInternalServerError), - }, - { - desc: "bootstrap with malformed response", - externalID: externalId, - externalKey: externalKey, - svcResp: bootstrapConfig, - readerResp: []byte{0}, - expectSvcCall: true, - expectedSDKErr: errors.NewSDKError(errors.New("json: cannot unmarshal string into Go value of type sdk.BootstrapConfig")), - }, - { - desc: "bootstrap with empty id", - externalID: "", - externalKey: externalKey, - expectedSDKErr: errors.NewSDKError(apiutil.ErrMissingID), - }, - { - desc: "bootstrap with empty key", - externalID: externalId, - externalKey: "", - expectedSDKErr: errors.NewSDKErrorWithStatus(apiutil.ErrBearerKey, http.StatusUnauthorized), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - var svcCall *mock.Call - var readerCall *mock.Call - if tc.expectSvcCall { - svcCall = bsvc.On("Bootstrap", mock.Anything, tc.externalKey, tc.externalID, false).Return(tc.svcResp, tc.svcErr) - readerCall = reader.On("ReadConfig", tc.svcResp, false).Return(tc.readerResp, tc.readerErr) - } - - resp, err := mgsdk.Bootstrap(context.Background(), tc.externalID, tc.externalKey) - - assert.Equal(t, tc.expectedSDKErr, err) - assert.Equal(t, tc.expectedResp, resp) - if svcCall != nil { - svcCall.Unset() - } - if readerCall != nil { - readerCall.Unset() - } - }) - } -} - -func TestBootstrapSecure(t *testing.T) { - bs, bsvc, reader, _ := setupBootstrap() - defer bs.Close() - - mgsdk := sdk.NewSDK(sdk.Config{BootstrapURL: bs.URL}) - - body, err := json.Marshal(readConfigResponse) - assert.Nil(t, err, fmt.Sprintf("Marshalling bootstrap response expected to succeed: %s.\n", err)) - - encResponse, err := encrypt(body, encKey) - assert.Nil(t, err, fmt.Sprintf("Encrypting bootstrap response expected to succeed: %s.\n", err)) - - cases := []struct { - desc string - externalID string - externalKey string - cryptoKey string - svcResp bootstrap.Config - svcErr error - readerResp []byte - readerErr error - expectSvcCall bool - expectedResp sdk.BootstrapConfig - expectedSDKErr errors.SDKError - }{ - { - desc: "secure bootstrap successfully", - externalID: externalId, - externalKey: externalKey, - cryptoKey: string(encKey), - svcResp: bootstrapConfig, - readerResp: encResponse, - expectSvcCall: true, - expectedResp: sdkBootstrapReadRes, - }, - { - desc: "secure bootstrap with invalid crypto key", - externalID: externalId, - externalKey: externalKey, - cryptoKey: invalid, - expectedSDKErr: errors.NewSDKError(errors.New("crypto/aes: invalid key size 7")), - }, - { - desc: "secure bootstrap with reader error", - externalID: externalId, - externalKey: externalKey, - cryptoKey: string(encKey), - svcResp: bootstrapConfig, - readerErr: errJSONEOF, - expectSvcCall: true, - expectedSDKErr: errors.NewSDKErrorWithStatus(errJSONEOF, http.StatusInternalServerError), - }, - { - desc: "secure bootstrap with malformed response", - externalID: externalId, - externalKey: externalKey, - cryptoKey: string(encKey), - svcResp: bootstrapConfig, - readerResp: []byte{0}, - expectSvcCall: true, - expectedSDKErr: errors.NewSDKError(errJSONEOF), - }, - { - desc: "secure bootstrap with empty id", - externalID: "", - externalKey: externalKey, - cryptoKey: string(encKey), - expectedSDKErr: errors.NewSDKError(apiutil.ErrMissingID), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - var svcCall *mock.Call - var readerCall *mock.Call - if tc.expectSvcCall { - svcCall = bsvc.On("Bootstrap", mock.Anything, mock.Anything, tc.externalID, true).Return(tc.svcResp, tc.svcErr) - readerCall = reader.On("ReadConfig", tc.svcResp, true).Return(tc.readerResp, tc.readerErr) - } - - resp, err := mgsdk.BootstrapSecure(context.Background(), tc.externalID, tc.externalKey, tc.cryptoKey) - - assert.Equal(t, tc.expectedSDKErr, err) - assert.Equal(t, tc.expectedResp, resp) - if svcCall != nil { - svcCall.Unset() - } - if readerCall != nil { - readerCall.Unset() - } - }) - } -} - -func TestCreateBootstrapProfile(t *testing.T) { - bs, bsvc, _, auth := setupBootstrap() - defer bs.Close() - - mgsdk := sdk.NewSDK(sdk.Config{BootstrapURL: bs.URL}) - - profile := sdk.BootstrapProfile{ - Name: "gateway-profile", - ContentFormat: "go-template", - } - saved := bootstrap.Profile{ - ID: testsutil.GenerateUUID(t), - DomainID: domainID, - Name: "gateway-profile", - ContentFormat: bootstrap.ContentFormatGoTemplate, - } - - cases := []struct { - desc string - token string - profile sdk.BootstrapProfile - svcResp bootstrap.Profile - svcErr error - authErr error - expectedSDKErr errors.SDKError - }{ - { - desc: "create profile successfully", - token: validToken, - profile: profile, - svcResp: saved, - }, - { - desc: "create profile with invalid token", - token: invalidToken, - profile: profile, - authErr: svcerr.ErrAuthentication, - expectedSDKErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "create profile with empty name", - token: validToken, - profile: sdk.BootstrapProfile{}, - expectedSDKErr: errors.NewSDKErrorWithStatus(apiutil.ErrMissingName, http.StatusBadRequest), - }, - { - desc: "create profile with service error", - token: validToken, - profile: profile, - svcErr: svcerr.ErrCreateEntity, - expectedSDKErr: errors.NewSDKErrorWithStatus(svcerr.ErrCreateEntity, http.StatusUnprocessableEntity), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - session := smqauthn.Session{} - if tc.token == validToken { - session = bootstrapSession() - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(session, tc.authErr) - svcCall := bsvc.On("CreateProfile", mock.Anything, session, mock.Anything).Return(tc.svcResp, tc.svcErr) - - _, err := mgsdk.CreateBootstrapProfile(context.Background(), tc.profile, domainID, tc.token) - assert.Equal(t, tc.expectedSDKErr, err) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestViewBootstrapProfile(t *testing.T) { - bs, bsvc, _, auth := setupBootstrap() - defer bs.Close() - - mgsdk := sdk.NewSDK(sdk.Config{BootstrapURL: bs.URL}) - - profileID := testsutil.GenerateUUID(t) - saved := bootstrap.Profile{ - ID: profileID, - DomainID: domainID, - Name: "gateway-profile", - ContentFormat: bootstrap.ContentFormatGoTemplate, - } - expected := sdk.BootstrapProfile{ - ID: profileID, - DomainID: domainID, - Name: "gateway-profile", - ContentFormat: "go-template", - } - - cases := []struct { - desc string - token string - profileID string - svcResp bootstrap.Profile - svcErr error - authErr error - expectedResp sdk.BootstrapProfile - expectedSDKErr errors.SDKError - }{ - { - desc: "view profile successfully", - token: validToken, - profileID: profileID, - svcResp: saved, - expectedResp: expected, - }, - { - desc: "view profile with invalid token", - token: invalidToken, - profileID: profileID, - authErr: svcerr.ErrAuthentication, - expectedSDKErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "view profile with empty id", - token: validToken, - profileID: "", - expectedSDKErr: errors.NewSDKError(apiutil.ErrMissingID), - }, - { - desc: "view profile not found", - token: validToken, - profileID: profileID, - svcErr: svcerr.ErrNotFound, - expectedSDKErr: errors.NewSDKErrorWithStatus(svcerr.ErrNotFound, http.StatusNotFound), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - session := smqauthn.Session{} - if tc.token == validToken { - session = bootstrapSession() - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(session, tc.authErr) - svcCall := bsvc.On("ViewProfile", mock.Anything, session, tc.profileID).Return(tc.svcResp, tc.svcErr) - - resp, err := mgsdk.ViewBootstrapProfile(context.Background(), tc.profileID, domainID, tc.token) - assert.Equal(t, tc.expectedSDKErr, err) - if tc.expectedSDKErr == nil { - assert.Equal(t, tc.expectedResp, resp) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestUpdateBootstrapProfile(t *testing.T) { - bs, bsvc, _, auth := setupBootstrap() - defer bs.Close() - - mgsdk := sdk.NewSDK(sdk.Config{BootstrapURL: bs.URL}) - - profileID := testsutil.GenerateUUID(t) - updatedProfile := bootstrap.Profile{ - ID: profileID, - DomainID: domainID, - Name: "updated-name", - ContentFormat: bootstrap.ContentFormatYAML, - } - expectedResp := sdk.BootstrapProfile{ - ID: profileID, - DomainID: domainID, - Name: "updated-name", - ContentFormat: "yaml", - } - - cases := []struct { - desc string - token string - profile sdk.BootstrapProfile - svcResp bootstrap.Profile - svcErr error - authErr error - expectedResp sdk.BootstrapProfile - expectedSDKErr errors.SDKError - }{ - { - desc: "update profile successfully", - token: validToken, - profile: sdk.BootstrapProfile{ - ID: profileID, - Name: "updated-name", - ContentFormat: "yaml", - }, - svcResp: updatedProfile, - expectedResp: expectedResp, - }, - { - desc: "update profile with invalid token", - token: invalidToken, - profile: sdk.BootstrapProfile{ - ID: profileID, - }, - authErr: svcerr.ErrAuthentication, - expectedSDKErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "update profile with empty id", - token: validToken, - profile: sdk.BootstrapProfile{}, - expectedSDKErr: errors.NewSDKError(apiutil.ErrMissingID), - }, - { - desc: "update profile not found", - token: validToken, - profile: sdk.BootstrapProfile{ - ID: profileID, - }, - svcErr: svcerr.ErrNotFound, - expectedSDKErr: errors.NewSDKErrorWithStatus(svcerr.ErrNotFound, http.StatusNotFound), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - session := smqauthn.Session{} - if tc.token == validToken { - session = bootstrapSession() - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(session, tc.authErr) - svcCall := bsvc.On("UpdateProfile", mock.Anything, session, mock.Anything).Return(tc.svcResp, tc.svcErr) - - resp, err := mgsdk.UpdateBootstrapProfile(context.Background(), tc.profile, domainID, tc.token) - assert.Equal(t, tc.expectedSDKErr, err) - if tc.expectedSDKErr == nil { - assert.Equal(t, tc.expectedResp, resp) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestBootstrapProfiles(t *testing.T) { - bs, bsvc, _, auth := setupBootstrap() - defer bs.Close() - - mgsdk := sdk.NewSDK(sdk.Config{BootstrapURL: bs.URL}) - - profiles := bootstrap.ProfilesPage{ - Total: 2, - Offset: 0, - Limit: 10, - Profiles: []bootstrap.Profile{ - {ID: testsutil.GenerateUUID(t), DomainID: domainID, Name: "p1", ContentFormat: bootstrap.ContentFormatGoTemplate}, - {ID: testsutil.GenerateUUID(t), DomainID: domainID, Name: "p2", ContentFormat: bootstrap.ContentFormatYAML}, - }, - } - - cases := []struct { - desc string - token string - pageMeta sdk.PageMetadata - svcResp bootstrap.ProfilesPage - svcErr error - authErr error - expectedCount int - expectedSDKErr errors.SDKError - }{ - { - desc: "list profiles successfully", - token: validToken, - pageMeta: sdk.PageMetadata{Offset: 0, Limit: 10}, - svcResp: profiles, - expectedCount: 2, - }, - { - desc: "list profiles filtered by name", - token: validToken, - pageMeta: sdk.PageMetadata{Offset: 0, Limit: 10, Name: "p1"}, - svcResp: bootstrap.ProfilesPage{Total: 1, Profiles: profiles.Profiles[:1]}, - expectedCount: 1, - }, - { - desc: "list profiles with invalid token", - token: invalidToken, - pageMeta: sdk.PageMetadata{Offset: 0, Limit: 10}, - authErr: svcerr.ErrAuthentication, - expectedSDKErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - session := smqauthn.Session{} - if tc.token == validToken { - session = bootstrapSession() - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(session, tc.authErr) - svcCall := bsvc.On("ListProfiles", mock.Anything, session, mock.Anything, mock.Anything, mock.Anything).Return(tc.svcResp, tc.svcErr) - - resp, err := mgsdk.BootstrapProfiles(context.Background(), tc.pageMeta, domainID, tc.token) - assert.Equal(t, tc.expectedSDKErr, err) - if tc.expectedSDKErr == nil { - assert.Equal(t, tc.expectedCount, len(resp.Profiles)) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func encrypt(in, encKey []byte) ([]byte, error) { - block, err := aes.NewCipher(encKey) - if err != nil { - return nil, err - } - ciphertext := make([]byte, aes.BlockSize+len(in)) - iv := ciphertext[:aes.BlockSize] - if _, err := io.ReadFull(rand.Reader, iv); err != nil { - return nil, err - } - stream := cipher.NewCFBEncrypter(block, iv) - stream.XORKeyStream(ciphertext[aes.BlockSize:], in) - return ciphertext, nil -} diff --git a/pkg/sdk/certs_metadata_test.go b/pkg/sdk/certs_metadata_test.go deleted file mode 100644 index 2e44f2e36..000000000 --- a/pkg/sdk/certs_metadata_test.go +++ /dev/null @@ -1,74 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package sdk - -import ( - "net/url" - "testing" - - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" -) - -func TestPageMetadataQueryWithCertFilters(t *testing.T) { - pm := PageMetadata{ - EntityID: "entity-id", - CommonName: "device-cn", - Organization: []string{"Acme", "QA"}, - OrganizationalUnit: []string{"Platform"}, - Country: []string{"RS"}, - Province: []string{"Belgrade"}, - Locality: []string{"Belgrade"}, - StreetAddress: []string{"Nemanjina 4"}, - PostalCode: []string{"11000"}, - DNSNames: []string{"device.local"}, - IPAddresses: []string{"127.0.0.1"}, - EmailAddresses: []string{"device@example.com"}, - TTL: "24h", - } - - encoded, err := pm.query() - require.NoError(t, err) - - values, err := url.ParseQuery(encoded) - require.NoError(t, err) - - assert.Equal(t, "entity-id", values.Get("entity_id")) - assert.Equal(t, "device-cn", values.Get("common_name")) - assert.Equal(t, "24h", values.Get("ttl")) - assert.ElementsMatch(t, []string{"Acme", "QA"}, values["organization"]) - assert.Equal(t, []string{"Platform"}, values["organizational_unit"]) - assert.Equal(t, []string{"RS"}, values["country"]) - assert.Equal(t, []string{"Belgrade"}, values["province"]) - assert.Equal(t, []string{"Belgrade"}, values["locality"]) - assert.Equal(t, []string{"Nemanjina 4"}, values["street_address"]) - assert.Equal(t, []string{"11000"}, values["postal_code"]) - assert.Equal(t, []string{"device.local"}, values["dns_names"]) - assert.Equal(t, []string{"127.0.0.1"}, values["ip_addresses"]) - assert.Equal(t, []string{"device@example.com"}, values["email_addresses"]) -} - -func TestCertStatusAliases(t *testing.T) { - assert.Equal(t, CertValid, Valid) - assert.Equal(t, CertRevoked, Revoked) - assert.Equal(t, CertUnknown, Unknown) -} - -func TestCertTypeString(t *testing.T) { - tests := []struct { - desc string - typ CertType - expected string - }{ - {desc: "root", typ: RootCA, expected: "root"}, - {desc: "intermediate", typ: IntermediateCA, expected: "intermediate"}, - {desc: "unknown", typ: CertType(99), expected: "unknown"}, - } - - for _, tc := range tests { - t.Run(tc.desc, func(t *testing.T) { - assert.Equal(t, tc.expected, tc.typ.String()) - }) - } -} diff --git a/pkg/sdk/certs_test.go b/pkg/sdk/certs_test.go deleted file mode 100644 index 25cb5297c..000000000 --- a/pkg/sdk/certs_test.go +++ /dev/null @@ -1,1032 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package sdk_test - -import ( - "context" - "fmt" - "net/http" - "net/http/httptest" - "testing" - - "github.com/absmach/magistrala/certs" - httpapi "github.com/absmach/magistrala/certs/api/http" - "github.com/absmach/magistrala/certs/mocks" - mglog "github.com/absmach/magistrala/logger" - smqauthn "github.com/absmach/magistrala/pkg/authn" - authnmocks "github.com/absmach/magistrala/pkg/authn/mocks" - "github.com/absmach/magistrala/pkg/errors" - "github.com/absmach/magistrala/pkg/sdk" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" -) - -const ( - certsInstanceID = "5de9b29a-feb9-11ed-be56-0242ac120002" - certsContentType = "application/senml+json" - serialNum = "8e7a30c-bc9f-22de-ae67-1342bc139507" - certsID = "c333e6f-59bb-4c39-9e13-3a2766af8ba5" - ttl = "10h" - commonName = "test" - token = "token" - agentToken = "agent-token" - certsDomainID = "domain-certsID" -) - -func setupCerts() (*httptest.Server, *mocks.Service, *authnmocks.Authentication) { - svc := new(mocks.Service) - logger := mglog.NewMock() - authn := new(authnmocks.Authentication) - am := smqauthn.NewAuthNMiddleware(authn, smqauthn.WithAllowUnverifiedUser(true)) - handler := httpapi.MakeHandler(svc, am, logger, certsInstanceID, agentToken) - - return httptest.NewServer(handler), svc, authn -} - -func TestIssueCert(t *testing.T) { - ts, svc, auth := setupCerts() - defer ts.Close() - - sdkConfig := sdk.Config{ - CertsURL: ts.URL, - MsgContentType: certsContentType, - TLSVerification: false, - } - - ctsdk := sdk.NewSDK(sdkConfig) - - ipAddr := []string{"192.128.101.82"} - cases := []struct { - desc string - entityID string - ttl string - ipAddrs []string - commonName string - svcresp certs.Certificate - svcerr error - authenticateErr error - err errors.SDKError - sdkCert sdk.Certificate - domain string - token string - session smqauthn.Session - }{ - { - desc: "IssueCert success", - entityID: certsID, - ttl: ttl, - ipAddrs: ipAddr, - commonName: commonName, - svcresp: certs.Certificate{ - SerialNumber: serialNum, - }, - sdkCert: sdk.Certificate{ - SerialNumber: serialNum, - }, - svcerr: nil, - err: nil, - domain: certsDomainID, - token: token, - }, - { - desc: "IssueCert failure", - entityID: certsID, - ttl: ttl, - ipAddrs: ipAddr, - commonName: commonName, - svcresp: certs.Certificate{}, - svcerr: certs.ErrCreateEntity, - err: errors.NewSDKErrorWithStatus(certs.ErrCreateEntity, http.StatusUnprocessableEntity), - domain: certsDomainID, - token: token, - }, - { - desc: "IssueCert with empty entityID", - entityID: `""`, - ttl: ttl, - ipAddrs: ipAddr, - commonName: commonName, - svcresp: certs.Certificate{}, - svcerr: certs.ErrMalformedEntity, - err: errors.NewSDKErrorWithStatus(certs.ErrMalformedEntity, http.StatusBadRequest), - domain: certsDomainID, - token: token, - }, - { - desc: "IssueCert with empty ipAddrs", - entityID: certsID, - ttl: ttl, - commonName: commonName, - svcresp: certs.Certificate{SerialNumber: serialNum}, - sdkCert: sdk.Certificate{ - SerialNumber: serialNum, - }, - svcerr: nil, - err: nil, - domain: certsDomainID, - token: token, - }, - { - desc: "IssueCert with empty ttl", - entityID: certsID, - ttl: "", - ipAddrs: ipAddr, - commonName: commonName, - svcresp: certs.Certificate{SerialNumber: serialNum}, - sdkCert: sdk.Certificate{ - SerialNumber: serialNum, - }, - svcerr: nil, - err: nil, - domain: certsDomainID, - token: token, - }, - { - desc: "IssueCert with empty commonName", - entityID: certsID, - ttl: ttl, - ipAddrs: ipAddr, - commonName: "", - svcresp: certs.Certificate{}, - svcerr: certs.ErrMalformedEntity, - err: errors.NewSDKErrorWithStatus(certs.ErrMalformedEntity, http.StatusBadRequest), - domain: certsDomainID, - token: token, - }, - { - desc: "IssueCert with empty token", - entityID: certsID, - ttl: ttl, - ipAddrs: ipAddr, - commonName: commonName, - svcresp: certs.Certificate{}, - svcerr: nil, - err: errors.NewSDKErrorWithStatus(errors.New("missing or invalid bearer user token"), http.StatusUnauthorized), - domain: certsDomainID, - token: "", - }, - { - desc: "IssueCert with empty domain", - entityID: certsID, - ttl: ttl, - ipAddrs: ipAddr, - commonName: commonName, - svcresp: certs.Certificate{}, - svcerr: nil, - err: errors.NewSDKErrorWithStatus(errors.New("missing domainID"), http.StatusBadRequest), - domain: "", - token: token, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == token { - tc.session = smqauthn.Session{DomainUserID: certsID, UserID: certsID, DomainID: certsDomainID} - } - - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := svc.On("IssueCert", mock.Anything, tc.session, tc.entityID, tc.ttl, tc.ipAddrs, certs.SubjectOptions{CommonName: tc.commonName}).Return(tc.svcresp, tc.svcerr) - resp, err := ctsdk.IssueCert(context.Background(), tc.entityID, tc.ttl, tc.ipAddrs, sdk.Options{CommonName: tc.commonName}, tc.domain, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - assert.Equal(t, tc.sdkCert.SerialNumber, resp.SerialNumber) - ok := svcCall.Parent.AssertCalled(t, "IssueCert", mock.Anything, tc.session, tc.entityID, tc.ttl, tc.ipAddrs, certs.SubjectOptions{CommonName: tc.commonName}) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestRevokeCert(t *testing.T) { - ts, svc, auth := setupCerts() - defer ts.Close() - - sdkConfig := sdk.Config{ - CertsURL: ts.URL, - MsgContentType: certsContentType, - TLSVerification: false, - } - - ctsdk := sdk.NewSDK(sdkConfig) - - cases := []struct { - desc string - serial string - svcresp string - svcerr error - authenticateErr error - err errors.SDKError - domain string - token string - session smqauthn.Session - }{ - { - desc: "RevokeCert success", - serial: serialNum, - svcerr: nil, - err: nil, - domain: certsDomainID, - token: token, - }, - { - desc: "RevokeCert failure", - serial: serialNum, - svcerr: certs.ErrUpdateEntity, - err: errors.NewSDKErrorWithStatus(certs.ErrUpdateEntity, http.StatusUnprocessableEntity), - domain: certsDomainID, - token: token, - }, - { - desc: "RevokeCert with empty serial", - serial: "", - svcerr: certs.ErrMalformedEntity, - err: errors.NewSDKErrorWithStatus(certs.ErrMalformedEntity, http.StatusBadRequest), - domain: certsDomainID, - token: token, - }, - { - desc: "RevokeCert with empty token", - serial: serialNum, - svcerr: nil, - err: errors.NewSDKErrorWithStatus(errors.New("missing or invalid bearer user token"), http.StatusUnauthorized), - domain: certsDomainID, - token: "", - }, - { - desc: "RevokeCert with empty domain", - serial: serialNum, - svcerr: nil, - err: errors.NewSDKErrorWithStatus(errors.New("missing domainID"), http.StatusBadRequest), - domain: "", - token: token, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == token { - tc.session = smqauthn.Session{DomainUserID: certsID, UserID: certsID, DomainID: certsDomainID} - } - - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := svc.On("RevokeBySerial", mock.Anything, tc.session, tc.serial).Return(tc.svcerr) - - err := ctsdk.RevokeCert(context.Background(), tc.serial, tc.domain, tc.token) - assert.Equal(t, tc.err, err) - if tc.desc != "RevokeCert with empty serial" && tc.desc != "RevokeCert with empty token" && tc.desc != "RevokeCert with empty domain" { - ok := svcCall.Parent.AssertCalled(t, "RevokeBySerial", mock.Anything, tc.session, tc.serial) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestDeleteCert(t *testing.T) { - ts, svc, auth := setupCerts() - defer ts.Close() - - sdkConfig := sdk.Config{ - CertsURL: ts.URL, - MsgContentType: certsContentType, - TLSVerification: false, - } - - ctsdk := sdk.NewSDK(sdkConfig) - - cases := []struct { - desc string - entityID string - svcresp string - svcerr error - authenticateErr error - err errors.SDKError - domain string - token string - session smqauthn.Session - }{ - { - desc: "DeleteCert success", - entityID: certsID, - svcerr: nil, - err: nil, - domain: certsDomainID, - token: token, - }, - { - desc: "DeleteCert failure", - entityID: certsID, - svcerr: certs.ErrUpdateEntity, - err: errors.NewSDKErrorWithStatus(certs.ErrUpdateEntity, http.StatusUnprocessableEntity), - domain: certsDomainID, - token: token, - }, - { - desc: "DeleteCert with empty entity certsID", - entityID: "", - svcerr: certs.ErrMalformedEntity, - err: errors.NewSDKErrorWithStatus(certs.ErrMalformedEntity, http.StatusBadRequest), - domain: certsDomainID, - token: token, - }, - { - desc: "DeleteCert with empty token", - entityID: certsID, - svcerr: nil, - err: errors.NewSDKErrorWithStatus(errors.New("missing or invalid bearer user token"), http.StatusUnauthorized), - domain: certsDomainID, - token: "", - }, - { - desc: "DeleteCert with empty domain", - entityID: certsID, - svcerr: nil, - err: errors.NewSDKErrorWithStatus(errors.New("missing domainID"), http.StatusBadRequest), - domain: "", - token: token, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == token { - tc.session = smqauthn.Session{DomainUserID: certsID, UserID: certsID, DomainID: certsDomainID} - } - - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := svc.On("RevokeAll", mock.Anything, tc.session, tc.entityID).Return(tc.svcerr) - - err := ctsdk.DeleteCert(context.Background(), tc.entityID, tc.domain, tc.token) - assert.Equal(t, tc.err, err) - if tc.desc != "DeleteCert with empty entity certsID" && tc.desc != "DeleteCert with empty token" && tc.desc != "DeleteCert with empty domain" { - ok := svcCall.Parent.AssertCalled(t, "RevokeAll", mock.Anything, tc.session, tc.entityID) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestRenewCert(t *testing.T) { - ts, svc, auth := setupCerts() - defer ts.Close() - - sdkConfig := sdk.Config{ - CertsURL: ts.URL, - MsgContentType: certsContentType, - TLSVerification: false, - } - - ctsdk := sdk.NewSDK(sdkConfig) - - cases := []struct { - desc string - serial string - svcresp certs.Certificate - svcerr error - authenticateErr error - err errors.SDKError - expected sdk.Certificate - domain string - token string - session smqauthn.Session - }{ - { - desc: "RenewCert success", - serial: serialNum, - svcresp: certs.Certificate{ - SerialNumber: "new-serial-123", - EntityID: "test-entity", - }, - svcerr: nil, - err: nil, - expected: sdk.Certificate{ - SerialNumber: "new-serial-123", - EntityID: "test-entity", - }, - domain: certsDomainID, - token: token, - }, - { - desc: "RenewCert failure", - serial: serialNum, - svcresp: certs.Certificate{}, - svcerr: certs.ErrUpdateEntity, - err: errors.NewSDKErrorWithStatus(certs.ErrUpdateEntity, http.StatusUnprocessableEntity), - expected: sdk.Certificate{}, - domain: certsDomainID, - token: token, - }, - { - desc: "RenewCert with empty serial", - serial: "", - svcresp: certs.Certificate{}, - svcerr: certs.ErrMalformedEntity, - err: errors.NewSDKErrorWithStatus(certs.ErrMalformedEntity, http.StatusBadRequest), - expected: sdk.Certificate{}, - domain: certsDomainID, - token: token, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == token { - tc.session = smqauthn.Session{DomainUserID: certsID, UserID: certsID, DomainID: certsDomainID} - } - - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := svc.On("RenewCert", mock.Anything, tc.session, tc.serial).Return(tc.svcresp, tc.svcerr) - - cert, err := ctsdk.RenewCert(context.Background(), tc.serial, tc.domain, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - assert.Equal(t, tc.expected, cert) - } else { - assert.Equal(t, sdk.Certificate{}, cert) - } - if tc.desc != "RenewCert with empty serial" { - ok := svcCall.Parent.AssertCalled(t, "RenewCert", mock.Anything, tc.session, tc.serial) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestListCerts(t *testing.T) { - ts, svc, auth := setupCerts() - defer ts.Close() - - sdkConfig := sdk.Config{ - CertsURL: ts.URL, - MsgContentType: certsContentType, - TLSVerification: false, - } - - ctsdk := sdk.NewSDK(sdkConfig) - - cases := []struct { - desc string - svcResp certs.CertificatePage - sdkPm sdk.PageMetadata - svcerr error - authenticateErr error - err errors.SDKError - domain string - token string - session smqauthn.Session - }{ - { - desc: "ListCerts success", - sdkPm: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - }, - svcResp: certs.CertificatePage{ - PageMetadata: certs.PageMetadata{ - Total: 1, - Offset: 0, - Limit: 10, - }, - Certificates: []certs.Certificate{ - { - SerialNumber: serialNum, - }, - }, - }, - domain: certsDomainID, - token: token, - }, - { - desc: "ListCerts success with entity certsID", - sdkPm: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - EntityID: certsID, - }, - svcResp: certs.CertificatePage{ - PageMetadata: certs.PageMetadata{ - Total: 1, - Offset: 0, - Limit: 10, - }, - Certificates: []certs.Certificate{ - { - SerialNumber: serialNum, - EntityID: certsID, - }, - }, - }, - domain: certsDomainID, - token: token, - }, - { - desc: "ListCerts failure", - sdkPm: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - }, - svcerr: certs.ErrViewEntity, - err: errors.NewSDKErrorWithStatus(certs.ErrViewEntity, http.StatusUnprocessableEntity), - domain: certsDomainID, - token: token, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == token { - tc.session = smqauthn.Session{DomainUserID: certsID, UserID: certsID, DomainID: certsDomainID} - } - - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := svc.On("ListCerts", mock.Anything, tc.session, mock.Anything).Return(tc.svcResp, tc.svcerr) - - resp, err := ctsdk.ListCerts(context.Background(), tc.sdkPm, tc.domain, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - assert.Equal(t, tc.svcResp.Total, resp.Total) - assert.Equal(t, tc.svcResp.Certificates[0].SerialNumber, resp.Certificates[0].SerialNumber) - if tc.desc == "ListCerts success with entity certsID" { - assert.Equal(t, tc.svcResp.Certificates[0].EntityID, resp.Certificates[0].EntityID) - } - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestViewCert(t *testing.T) { - ts, svc, auth := setupCerts() - defer ts.Close() - - sdkConfig := sdk.Config{ - CertsURL: ts.URL, - MsgContentType: certsContentType, - TLSVerification: false, - } - - ctsdk := sdk.NewSDK(sdkConfig) - - cert := sdk.Certificate{ - SerialNumber: serialNum, - } - - cases := []struct { - desc string - serial string - svcresp certs.Certificate - svcerr error - authenticateErr error - err errors.SDKError - sdkCert sdk.Certificate - domain string - token string - session smqauthn.Session - }{ - { - desc: "ViewCert success", - serial: serialNum, - svcresp: certs.Certificate{ - SerialNumber: serialNum, - }, - sdkCert: cert, - svcerr: nil, - err: nil, - domain: certsDomainID, - token: token, - }, - { - desc: "ViewCert failure", - serial: serialNum, - svcresp: certs.Certificate{}, - svcerr: certs.ErrViewEntity, - err: errors.NewSDKErrorWithStatus(certs.ErrViewEntity, http.StatusUnprocessableEntity), - domain: certsDomainID, - token: token, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == token { - tc.session = smqauthn.Session{DomainUserID: certsID, UserID: certsID, DomainID: certsDomainID} - } - - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := svc.On("ViewCert", mock.Anything, tc.session, tc.serial).Return(tc.svcresp, tc.svcerr) - - c, err := ctsdk.ViewCert(context.Background(), tc.serial, tc.domain, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "ViewCert", mock.Anything, tc.session, tc.serial) - assert.True(t, ok) - } - assert.Equal(t, tc.sdkCert.SerialNumber, c.SerialNumber, fmt.Sprintf("expected: %v, got: %v", tc.sdkCert.SerialNumber, c.SerialNumber)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestDownloadCACert(t *testing.T) { - ts, svc, _ := setupCerts() - defer ts.Close() - - sdkConfig := sdk.Config{ - CertsURL: ts.URL, - MsgContentType: certsContentType, - TLSVerification: false, - } - - ctsdk := sdk.NewSDK(sdkConfig) - - cert := sdk.Certificate{ - SerialNumber: serialNum, - } - - cases := []struct { - desc string - svcresp certs.Certificate - svcerr error - err errors.SDKError - sdkCert sdk.Certificate - }{ - { - desc: "Download CA successfully", - svcresp: certs.Certificate{ - SerialNumber: serialNum, - Certificate: []byte("cert"), - Key: []byte("key"), - }, - sdkCert: cert, - svcerr: nil, - err: nil, - }, - { - desc: "Download CA failure", - svcresp: certs.Certificate{}, - svcerr: certs.ErrViewEntity, - err: errors.NewSDKErrorWithStatus(certs.ErrViewEntity, http.StatusUnprocessableEntity), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("RetrieveCAChain", mock.Anything).Return(tc.svcresp, tc.svcerr) - - _, err := ctsdk.DownloadCA(context.Background()) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RetrieveCAChain", mock.Anything) - assert.True(t, ok) - } - svcCall.Unset() - }) - } -} - -func TestViewCA(t *testing.T) { - ts, svc, _ := setupCerts() - defer ts.Close() - - sdkConfig := sdk.Config{ - CertsURL: ts.URL, - MsgContentType: certsContentType, - TLSVerification: false, - } - - ctsdk := sdk.NewSDK(sdkConfig) - - cert := sdk.Certificate{ - SerialNumber: serialNum, - Certificate: "cert", - Key: "Key", - } - - cases := []struct { - desc string - svcresp certs.Certificate - svcerr error - err errors.SDKError - sdkCert sdk.Certificate - }{ - { - desc: "ViewCA success", - svcresp: certs.Certificate{ - Certificate: []byte("cert"), - }, - sdkCert: cert, - svcerr: nil, - err: nil, - }, - { - desc: "ViewCA failure", - svcresp: certs.Certificate{}, - svcerr: certs.ErrViewEntity, - err: errors.NewSDKErrorWithStatus(certs.ErrViewEntity, http.StatusUnprocessableEntity), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("RetrieveCAChain", mock.Anything).Return(tc.svcresp, tc.svcerr) - - c, err := ctsdk.ViewCA(context.Background()) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RetrieveCAChain", mock.Anything) - assert.True(t, ok) - } - assert.Equal(t, tc.sdkCert.Certificate, c.Certificate, fmt.Sprintf("expected: %v, got: %v", tc.sdkCert.Certificate, c.Certificate)) - svcCall.Unset() - }) - } -} - -func TestGenerateCRL(t *testing.T) { - ts, svc, _ := setupCerts() - defer ts.Close() - - sdkConfig := sdk.Config{ - CertsURL: ts.URL, - MsgContentType: certsContentType, - TLSVerification: false, - } - - ctsdk := sdk.NewSDK(sdkConfig) - - crlData := []byte("mock-crl-data") - - cases := []struct { - desc string - svcresp []byte - svcerr error - err errors.SDKError - }{ - { - desc: "GenerateCRL success", - svcresp: crlData, - svcerr: nil, - err: nil, - }, - { - desc: "GenerateCRL failure", - svcresp: nil, - svcerr: certs.ErrFailedCertCreation, - err: errors.NewSDKErrorWithStatus(certs.ErrFailedCertCreation, http.StatusUnprocessableEntity), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("GenerateCRL", mock.Anything).Return(tc.svcresp, tc.svcerr) - - resp, err := ctsdk.GenerateCRL(context.Background()) - assert.Equal(t, tc.err, err) - if tc.err == nil { - assert.Equal(t, tc.svcresp, resp) - ok := svcCall.Parent.AssertCalled(t, "GenerateCRL", mock.Anything) - assert.True(t, ok) - } - svcCall.Unset() - }) - } -} - -func TestRevokeAll(t *testing.T) { - ts, svc, auth := setupCerts() - defer ts.Close() - - sdkConfig := sdk.Config{ - CertsURL: ts.URL, - MsgContentType: certsContentType, - TLSVerification: false, - } - - ctsdk := sdk.NewSDK(sdkConfig) - - cases := []struct { - desc string - entityID string - svcerr error - authenticateErr error - err errors.SDKError - domain string - token string - session smqauthn.Session - }{ - { - desc: "RevokeAll success", - entityID: certsID, - svcerr: nil, - err: nil, - domain: certsDomainID, - token: token, - }, - { - desc: "RevokeAll failure", - entityID: certsID, - svcerr: certs.ErrUpdateEntity, - err: errors.NewSDKErrorWithStatus(certs.ErrUpdateEntity, http.StatusUnprocessableEntity), - domain: certsDomainID, - token: token, - }, - { - desc: "RevokeAll with empty entityID", - entityID: "", - svcerr: certs.ErrMalformedEntity, - err: errors.NewSDKErrorWithStatus(certs.ErrMalformedEntity, http.StatusBadRequest), - domain: certsDomainID, - token: token, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == token { - tc.session = smqauthn.Session{DomainUserID: certsID, UserID: certsID, DomainID: certsDomainID} - } - - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := svc.On("RevokeAll", mock.Anything, tc.session, tc.entityID).Return(tc.svcerr) - - err := ctsdk.RevokeAll(context.Background(), tc.entityID, tc.domain, tc.token) - assert.Equal(t, tc.err, err) - if tc.desc != "RevokeAll with empty entityID" { - ok := svcCall.Parent.AssertCalled(t, "RevokeAll", mock.Anything, tc.session, tc.entityID) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestGetEntityID(t *testing.T) { - ts, svc, auth := setupCerts() - defer ts.Close() - - sdkConfig := sdk.Config{ - CertsURL: ts.URL, - MsgContentType: certsContentType, - TLSVerification: false, - } - - ctsdk := sdk.NewSDK(sdkConfig) - - entityID := "test-entity-certsID" - - cases := []struct { - desc string - serial string - svcresp certs.Certificate - svcerr error - authenticateErr error - err errors.SDKError - expected string - domain string - token string - session smqauthn.Session - }{ - { - desc: "GetEntityID success", - serial: serialNum, - svcresp: certs.Certificate{ - SerialNumber: serialNum, - EntityID: entityID, - }, - svcerr: nil, - err: nil, - expected: entityID, - domain: certsDomainID, - token: token, - }, - { - desc: "GetEntityID failure", - serial: serialNum, - svcresp: certs.Certificate{}, - svcerr: certs.ErrViewEntity, - err: errors.NewSDKErrorWithStatus(certs.ErrViewEntity, http.StatusUnprocessableEntity), - expected: "", - domain: certsDomainID, - token: token, - }, - { - desc: "GetEntityID with empty serial", - serial: "", - svcresp: certs.Certificate{}, - svcerr: certs.ErrMalformedEntity, - err: errors.NewSDKErrorWithStatus(certs.ErrMalformedEntity, http.StatusBadRequest), - expected: "", - domain: certsDomainID, - token: token, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == token { - tc.session = smqauthn.Session{DomainUserID: certsID, UserID: certsID, DomainID: certsDomainID} - } - - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - var svcCall *mock.Call - if tc.desc == "GetEntityID with empty serial" { - // Empty serial routes to ListCerts endpoint instead of ViewCert - svcCall = svc.On("ListCerts", mock.Anything, tc.session, mock.Anything).Return(certs.CertificatePage{}, tc.svcerr) - } else { - svcCall = svc.On("ViewCert", mock.Anything, tc.session, tc.serial).Return(tc.svcresp, tc.svcerr) - } - - resp, err := ctsdk.EntityID(context.Background(), tc.serial, tc.domain, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.expected, resp) - if tc.desc != "GetEntityID with empty serial" { - ok := svcCall.Parent.AssertCalled(t, "ViewCert", mock.Anything, tc.session, tc.serial) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestIssueFromCSRInternal(t *testing.T) { - ts, svc, auth := setupCerts() - defer ts.Close() - - sdkConfig := sdk.Config{ - CertsURL: ts.URL, - MsgContentType: certsContentType, - TLSVerification: false, - } - - ctsdk := sdk.NewSDK(sdkConfig) - - cert := sdk.Certificate{ - SerialNumber: serialNum, - } - - cases := []struct { - desc string - entityID string - ttl string - csr string - svcresp certs.Certificate - svcerr error - err errors.SDKError - sdkCert sdk.Certificate - }{ - { - desc: "IssueFromCSRInternal success", - entityID: certsID, - ttl: ttl, - csr: "valid-csr-content", - svcresp: certs.Certificate{ - SerialNumber: serialNum, - Certificate: []byte("cert"), - Key: []byte("key"), - }, - sdkCert: cert, - svcerr: nil, - err: nil, - }, - { - desc: "IssueFromCSRInternal failure", - entityID: certsID, - ttl: ttl, - csr: "invalid-csr-content", - svcresp: certs.Certificate{}, - svcerr: certs.ErrFailedCertCreation, - err: errors.NewSDKErrorWithStatus(certs.ErrFailedCertCreation, http.StatusUnprocessableEntity), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - agentSession := smqauthn.Session{DomainUserID: certsID, UserID: certsID, DomainID: certsDomainID} - authCall := auth.On("Authenticate", mock.Anything, agentToken).Return(agentSession, nil) - svcCall := svc.On("IssueFromCSRInternal", mock.Anything, tc.entityID, tc.ttl, mock.Anything).Return(tc.svcresp, tc.svcerr) - - c, err := ctsdk.IssueFromCSRInternal(context.Background(), tc.entityID, tc.ttl, tc.csr, agentToken) - assert.Equal(t, tc.err, err) - if tc.err == nil { - assert.Equal(t, tc.sdkCert.SerialNumber, c.SerialNumber) - ok := svcCall.Parent.AssertCalled(t, "IssueFromCSRInternal", mock.Anything, tc.entityID, tc.ttl, mock.Anything) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} diff --git a/pkg/sdk/channels.go b/pkg/sdk/channels.go index d321a27a6..196d8fbfd 100644 --- a/pkg/sdk/channels.go +++ b/pkg/sdk/channels.go @@ -12,7 +12,6 @@ import ( apiutil "github.com/absmach/magistrala/api/http/util" "github.com/absmach/magistrala/pkg/errors" - "github.com/absmach/magistrala/pkg/roles" ) const ( @@ -22,19 +21,18 @@ const ( // Channel represents magistrala channel. type Channel struct { - ID string `json:"id,omitempty"` - Name string `json:"name,omitempty"` - Tags []string `json:"tags,omitempty"` - Route string `json:"route,omitempty"` - ParentGroup string `json:"parent_group_id,omitempty"` - DomainID string `json:"domain_id,omitempty"` - Metadata Metadata `json:"metadata,omitempty"` - CreatedAt time.Time `json:"created_at,omitempty"` - UpdatedAt time.Time `json:"updated_at,omitempty"` - UpdatedBy string `json:"updated_by,omitempty"` - Status string `json:"status,omitempty"` - Permissions []string `json:"permissions,omitempty"` - Roles []roles.MemberRoleActions `json:"roles,omitempty"` + ID string `json:"id,omitempty"` + Name string `json:"name,omitempty"` + Tags []string `json:"tags,omitempty"` + Route string `json:"route,omitempty"` + ParentGroup string `json:"parent_group_id,omitempty"` + DomainID string `json:"domain_id,omitempty"` + Metadata Metadata `json:"metadata,omitempty"` + CreatedAt time.Time `json:"created_at,omitempty"` + UpdatedAt time.Time `json:"updated_at,omitempty"` + UpdatedBy string `json:"updated_by,omitempty"` + Status string `json:"status,omitempty"` + Permissions []string `json:"permissions,omitempty"` } func (sdk mgSDK) CreateChannel(ctx context.Context, c Channel, domainID, token string) (Channel, errors.SDKError) { diff --git a/pkg/sdk/channels_test.go b/pkg/sdk/channels_test.go deleted file mode 100644 index 8647c943d..000000000 --- a/pkg/sdk/channels_test.go +++ /dev/null @@ -1,2139 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package sdk_test - -import ( - "context" - "fmt" - "net/http" - "net/http/httptest" - "strings" - "testing" - "time" - - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/channels" - chapi "github.com/absmach/magistrala/channels/api/http" - chmocks "github.com/absmach/magistrala/channels/mocks" - "github.com/absmach/magistrala/internal/testsutil" - mglog "github.com/absmach/magistrala/logger" - smqauthn "github.com/absmach/magistrala/pkg/authn" - authnmocks "github.com/absmach/magistrala/pkg/authn/mocks" - "github.com/absmach/magistrala/pkg/connections" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/pkg/roles" - sdk "github.com/absmach/magistrala/pkg/sdk" - "github.com/absmach/magistrala/pkg/uuid" - "github.com/go-chi/chi/v5" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" -) - -var ( - channelName = "channelName" - newName = "newName" - valid = "valid" - channel = generateTestChannel(&testing.T{}) -) - -func setupChannels() (*httptest.Server, *chmocks.Service, *authnmocks.Authentication) { - svc := new(chmocks.Service) - logger := mglog.NewMock() - authn := new(authnmocks.Authentication) - mux := chi.NewRouter() - idp := uuid.NewMock() - am := smqauthn.NewAuthNMiddleware(authn, smqauthn.WithAllowUnverifiedUser(true)) - chapi.MakeHandler(svc, am, mux, logger, "", idp) - - return httptest.NewServer(mux), svc, authn -} - -func TestCreateChannel(t *testing.T) { - ts, gsvc, auth := setupChannels() - defer ts.Close() - - createChannelReq := channels.Channel{ - Name: channel.Name, - Route: channel.Route, - Metadata: channels.Metadata{"role": "client"}, - Status: channels.EnabledStatus, - } - - channelReq := sdk.Channel{ - Name: channel.Name, - Route: channel.Route, - Metadata: validMetadata, - Status: channels.EnabledStatus.String(), - } - - parentID := testsutil.GenerateUUID(&testing.T{}) - pChannel := channel - pChannel.ParentGroup = parentID - - iChannel := convertChannel(channel) - iChannel.Metadata = channels.Metadata{ - "test": make(chan int), - } - - conf := sdk.Config{ - ChannelsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - cases := []struct { - desc string - channelReq sdk.Channel - domainID string - token string - session smqauthn.Session - createChannelReq channels.Channel - svcRes []channels.Channel - svcErr error - authenticateRes smqauthn.Session - authenticateErr error - response sdk.Channel - err errors.SDKError - }{ - { - desc: "create channel successfully", - channelReq: channelReq, - domainID: domainID, - token: validToken, - createChannelReq: createChannelReq, - svcRes: []channels.Channel{convertChannel(channel)}, - svcErr: nil, - response: channel, - err: nil, - }, - { - desc: "create channel with existing name", - channelReq: channelReq, - domainID: domainID, - token: validToken, - createChannelReq: createChannelReq, - svcRes: []channels.Channel{}, - svcErr: svcerr.ErrCreateEntity, - response: sdk.Channel{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrCreateEntity, http.StatusUnprocessableEntity), - }, - { - desc: "create channel that can't be marshalled", - channelReq: sdk.Channel{ - Name: "test", - Metadata: map[string]any{ - "test": make(chan int), - }, - }, - domainID: domainID, - token: validToken, - createChannelReq: channels.Channel{}, - svcRes: []channels.Channel{}, - svcErr: nil, - response: sdk.Channel{}, - err: errors.NewSDKError(fmt.Errorf("json: unsupported type: chan int")), - }, - { - desc: "create channel with parent group", - channelReq: sdk.Channel{ - Name: channel.Name, - Route: channel.Route, - ParentGroup: parentID, - Status: channels.EnabledStatus.String(), - }, - domainID: domainID, - token: validToken, - createChannelReq: channels.Channel{ - Name: channel.Name, - ParentGroup: parentID, - Route: channel.Route, - Status: channels.EnabledStatus, - }, - svcRes: []channels.Channel{convertChannel(pChannel)}, - svcErr: nil, - response: pChannel, - err: nil, - }, - { - desc: "create channel with invalid parent", - channelReq: sdk.Channel{ - Name: channel.Name, - Route: channel.Route, - ParentGroup: wrongID, - Status: channels.EnabledStatus.String(), - }, - domainID: domainID, - token: validToken, - createChannelReq: channels.Channel{ - Name: channel.Name, - ParentGroup: wrongID, - Route: channel.Route, - Status: channels.EnabledStatus, - }, - svcRes: []channels.Channel{}, - svcErr: svcerr.ErrCreateEntity, - response: sdk.Channel{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrCreateEntity, http.StatusUnprocessableEntity), - }, - { - desc: "create a channel with every field defined", - channelReq: sdk.Channel{ - ID: channel.ID, - ParentGroup: parentID, - Route: channel.Route, - Name: channel.Name, - Metadata: validMetadata, - CreatedAt: channel.CreatedAt, - UpdatedAt: channel.UpdatedAt, - Status: channels.EnabledStatus.String(), - }, - domainID: domainID, - token: validToken, - createChannelReq: channels.Channel{ - ID: channel.ID, - ParentGroup: parentID, - Route: channel.Route, - Name: channel.Name, - Metadata: channels.Metadata{"role": "client"}, - CreatedAt: channel.CreatedAt, - UpdatedAt: channel.UpdatedAt, - Status: channels.EnabledStatus, - }, - svcRes: []channels.Channel{convertChannel(pChannel)}, - svcErr: nil, - response: pChannel, - err: nil, - }, - { - desc: "create channel with response that can't be unmarshalled", - channelReq: channelReq, - domainID: domainID, - token: validToken, - createChannelReq: createChannelReq, - svcRes: []channels.Channel{iChannel}, - svcErr: nil, - response: sdk.Channel{}, - err: errors.NewSDKError(fmt.Errorf("unexpected end of JSON input")), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: fmt.Sprintf("%s_%s", domainID, validID), UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := gsvc.On("CreateChannels", mock.Anything, tc.session, []channels.Channel{tc.createChannelReq}).Return(tc.svcRes, []roles.RoleProvision{}, tc.svcErr) - resp, err := mgsdk.CreateChannel(context.Background(), tc.channelReq, tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "CreateChannels", mock.Anything, tc.session, []channels.Channel{tc.createChannelReq}) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestCreateChannels(t *testing.T) { - ts, gsvc, auth := setupChannels() - defer ts.Close() - - var chs []sdk.Channel - conf := sdk.Config{ - ChannelsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - for i := 0; i < 3; i++ { - gr := generateTestChannel(t) - chs = append(chs, gr) - } - - cases := []struct { - desc string - domainID string - token string - session smqauthn.Session - channelsReq []sdk.Channel - createChannelsReq []channels.Channel - svcRes []channels.Channel - svcErr error - authenticateErr error - response []sdk.Channel - err errors.SDKError - }{ - { - desc: "create channels successfully", - domainID: domainID, - token: validToken, - channelsReq: chs, - createChannelsReq: convertChannels(chs), - svcRes: convertChannels(chs), - svcErr: nil, - response: chs, - err: nil, - }, - { - desc: "create channels with invalid token", - domainID: domainID, - token: invalidToken, - channelsReq: chs, - createChannelsReq: convertChannels(chs), - svcRes: []channels.Channel{}, - authenticateErr: svcerr.ErrAuthentication, - response: []sdk.Channel{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "create channels with empty token", - domainID: validID, - token: "", - channelsReq: chs, - createChannelsReq: convertChannels(chs), - svcRes: []channels.Channel{}, - svcErr: nil, - response: []sdk.Channel{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "create channels with service response that can,t be marshalled", - domainID: domainID, - token: validToken, - channelsReq: []sdk.Channel{ - { - ID: generateUUID(t), - Name: "channel_1", - Route: valid, - Metadata: map[string]any{ - "test": make(chan int), - }, - }, - }, - createChannelsReq: convertChannels(chs), - svcRes: []channels.Channel{}, - svcErr: nil, - response: []sdk.Channel{}, - err: errors.NewSDKError(fmt.Errorf("json: unsupported type: chan int")), - }, - { - desc: "create channels with service response that can't be unmarshalled", - domainID: domainID, - token: validToken, - channelsReq: chs, - createChannelsReq: convertChannels(chs), - svcRes: []channels.Channel{ - { - ID: generateUUID(t), - Metadata: channels.Metadata{ - "test": make(chan int), - }, - }, - }, - svcErr: nil, - response: []sdk.Channel{}, - err: errors.NewSDKError(fmt.Errorf("unexpected end of JSON input")), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: fmt.Sprintf("%s_%s", domainID, validID), UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := gsvc.On("CreateChannels", mock.Anything, tc.session, tc.createChannelsReq).Return(tc.svcRes, []roles.RoleProvision{}, tc.svcErr) - resp, err := mgsdk.CreateChannels(context.Background(), tc.channelsReq, tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestListChannels(t *testing.T) { - ts, gsvc, auth := setupChannels() - defer ts.Close() - - var chs []sdk.Channel - conf := sdk.Config{ - ChannelsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - for i := 10; i < 100; i++ { - gr := generateTestChannel(t) - chs = append(chs, gr) - } - - cases := []struct { - desc string - domainID string - token string - session smqauthn.Session - status channels.Status - total uint64 - offset uint64 - limit uint64 - level int - name string - metadata sdk.Metadata - channelsPageMeta channels.Page - svcRes channels.ChannelsPage - svcErr error - authenticateRes smqauthn.Session - authenticateErr error - response sdk.ChannelsPage - err errors.SDKError - }{ - { - desc: "list channels successfully", - token: validToken, - domainID: domainID, - limit: limit, - offset: offset, - total: total, - channelsPageMeta: channels.Page{ - Actions: []string{}, - Order: "updated_at", - Dir: "desc", - Offset: offset, - Limit: limit, - }, - svcRes: channels.ChannelsPage{ - Page: channels.Page{ - Total: uint64(len(chs[offset:limit])), - }, - Channels: convertChannels(chs[offset:limit]), - }, - response: sdk.ChannelsPage{ - PageRes: sdk.PageRes{ - Total: uint64(len(chs[offset:limit])), - }, - Channels: chs[offset:limit], - }, - err: nil, - }, - { - desc: "list channels with invalid token", - token: invalidToken, - domainID: domainID, - offset: offset, - limit: limit, - channelsPageMeta: channels.Page{ - Actions: []string{}, - Order: "updated_at", - Dir: "desc", - Offset: offset, - Limit: limit, - }, - svcRes: channels.ChannelsPage{}, - authenticateErr: svcerr.ErrAuthentication, - response: sdk.ChannelsPage{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "list channels with empty token", - token: "", - domainID: validID, - offset: offset, - limit: limit, - channelsPageMeta: channels.Page{ - Actions: []string{}, - Order: "updated_at", - Dir: "desc", - }, - svcRes: channels.ChannelsPage{}, - svcErr: nil, - response: sdk.ChannelsPage{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "list channels with zero limit", - token: validToken, - domainID: domainID, - offset: offset, - limit: 0, - channelsPageMeta: channels.Page{ - Actions: []string{}, - Order: "updated_at", - Dir: "desc", - Offset: offset, - Limit: 10, - }, - svcRes: channels.ChannelsPage{ - Page: channels.Page{ - Total: uint64(len(chs[offset:])), - }, - Channels: convertChannels(chs[offset:limit]), - }, - svcErr: nil, - response: sdk.ChannelsPage{ - PageRes: sdk.PageRes{ - Total: uint64(len(chs[offset:])), - }, - Channels: chs[offset:limit], - }, - err: nil, - }, - { - desc: "list channels with limit greater than max", - token: validToken, - domainID: domainID, - offset: offset, - limit: 110, - channelsPageMeta: channels.Page{ - Actions: []string{}, - Order: "updated_at", - Dir: "desc", - }, - svcRes: channels.ChannelsPage{}, - svcErr: nil, - response: sdk.ChannelsPage{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrLimitSize, http.StatusBadRequest), - }, - { - desc: "list channels with level", - token: validToken, - domainID: domainID, - offset: 0, - limit: 1, - level: 1, - channelsPageMeta: channels.Page{ - Actions: []string{}, - Order: "updated_at", - Dir: "desc", - Offset: offset, - Limit: 1, - }, - svcRes: channels.ChannelsPage{ - Page: channels.Page{ - Total: 1, - }, - Channels: convertChannels(chs[0:1]), - }, - svcErr: nil, - response: sdk.ChannelsPage{ - PageRes: sdk.PageRes{ - Total: 1, - }, - Channels: chs[0:1], - }, - err: nil, - }, - { - desc: "list channels with metadata", - token: validToken, - domainID: domainID, - offset: 0, - limit: 10, - metadata: sdk.Metadata{"name": "client_89"}, - channelsPageMeta: channels.Page{ - Actions: []string{}, - Order: "updated_at", - Dir: "desc", - Offset: offset, - Limit: 10, - Metadata: channels.Metadata{"name": "client_89"}, - }, - svcRes: channels.ChannelsPage{ - Page: channels.Page{ - Total: 1, - }, - Channels: convertChannels([]sdk.Channel{chs[89]}), - }, - svcErr: nil, - response: sdk.ChannelsPage{ - PageRes: sdk.PageRes{ - Total: 1, - }, - Channels: []sdk.Channel{chs[89]}, - }, - err: nil, - }, - { - desc: "list channels with invalid metadata", - token: validToken, - domainID: domainID, - offset: 0, - limit: 10, - metadata: sdk.Metadata{ - "test": make(chan int), - }, - channelsPageMeta: channels.Page{ - Actions: []string{}, - Order: "updated_at", - Dir: "desc", - }, - svcRes: channels.ChannelsPage{}, - svcErr: nil, - response: sdk.ChannelsPage{}, - err: errors.NewSDKError(fmt.Errorf("json: unsupported type: chan int")), - }, - { - desc: "list channels with service response that can't be unmarshalled", - token: validToken, - domainID: domainID, - offset: 0, - limit: 10, - channelsPageMeta: channels.Page{ - Actions: []string{}, - Order: "updated_at", - Dir: "desc", - Offset: 0, - Limit: 10, - }, - svcRes: channels.ChannelsPage{ - Page: channels.Page{ - Total: 1, - }, - Channels: []channels.Channel{{ - ID: generateUUID(t), - Metadata: channels.Metadata{ - "test": make(chan int), - }, - }}, - }, - svcErr: nil, - response: sdk.ChannelsPage{}, - err: errors.NewSDKError(fmt.Errorf("unexpected end of JSON input")), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - pm := sdk.PageMetadata{ - Offset: tc.offset, - Limit: tc.limit, - Level: uint64(tc.level), - Metadata: tc.metadata, - } - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := gsvc.On("ListChannels", mock.Anything, tc.session, tc.channelsPageMeta).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.Channels(context.Background(), pm, tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "ListChannels", mock.Anything, tc.session, tc.channelsPageMeta) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestViewChannel(t *testing.T) { - ts, gsvc, auth := setupChannels() - defer ts.Close() - - channelRes := convertChannel(channel) - conf := sdk.Config{ - ChannelsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - channelResRoles := sdk.Config{ - ChannelsURL: ts.URL, - Roles: true, - } - mgsdkRoles := sdk.NewSDK(channelResRoles) - - cases := []struct { - desc string - domainID string - token string - session smqauthn.Session - withRoles bool - channelID string - svcRes channels.Channel - svcErr error - authenticateErr error - response sdk.Channel - err errors.SDKError - }{ - { - desc: "view channel successfully", - domainID: domainID, - token: validToken, - withRoles: false, - channelID: channelRes.ID, - svcRes: channelRes, - svcErr: nil, - response: channel, - err: nil, - }, - { - desc: "view channel successfully with roles", - domainID: domainID, - token: validToken, - withRoles: true, - channelID: channelRes.ID, - svcRes: channelRes, - svcErr: nil, - response: channel, - err: nil, - }, - { - desc: "view channel with invalid token", - domainID: domainID, - token: invalidToken, - withRoles: false, - channelID: channelRes.ID, - svcRes: channels.Channel{}, - authenticateErr: svcerr.ErrAuthentication, - response: sdk.Channel{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "view channel with empty token", - domainID: domainID, - token: "", - withRoles: false, - channelID: channelRes.ID, - svcRes: channels.Channel{}, - svcErr: nil, - response: sdk.Channel{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "view channel for wrong id", - domainID: domainID, - token: validToken, - withRoles: false, - channelID: wrongID, - svcRes: channels.Channel{}, - svcErr: svcerr.ErrNotFound, - response: sdk.Channel{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrNotFound, http.StatusNotFound), - }, - { - desc: "view channel with empty channel id", - domainID: domainID, - token: validToken, - withRoles: false, - channelID: "", - svcRes: channels.Channel{}, - svcErr: nil, - response: sdk.Channel{}, - err: errors.NewSDKError(apiutil.ErrMissingID), - }, - { - desc: "view channel with service response that can't be unmarshalled", - domainID: domainID, - token: validToken, - withRoles: false, - channelID: channelRes.ID, - svcRes: channels.Channel{ - ID: generateUUID(t), - Metadata: channels.Metadata{ - "test": make(chan int), - }, - }, - svcErr: nil, - response: sdk.Channel{}, - err: errors.NewSDKError(fmt.Errorf("unexpected end of JSON input")), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := gsvc.On("ViewChannel", mock.Anything, tc.session, tc.channelID, tc.withRoles).Return(tc.svcRes, tc.svcErr) - - var resp sdk.Channel - var err error - - switch tc.withRoles { - case true: - resp, err = mgsdkRoles.Channel(context.Background(), tc.channelID, tc.domainID, tc.token) - default: - resp, err = mgsdk.Channel(context.Background(), tc.channelID, tc.domainID, tc.token) - } - - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.withRoles { - assert.Equal(t, resp.Roles, validRoles, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, validRoles, resp.Roles)) - } - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "ViewChannel", mock.Anything, tc.session, tc.channelID, tc.withRoles) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestUpdateChannel(t *testing.T) { - ts, gsvc, auth := setupChannels() - defer ts.Close() - - conf := sdk.Config{ - ChannelsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - mChannel := convertChannel(channel) - mChannel.Metadata = channels.Metadata{ - "field": "value2", - } - msdkChannel := channel - msdkChannel.Metadata = sdk.Metadata{ - "field": "value2", - } - - nChannel := convertChannel(channel) - nChannel.Name = newName - nsdkChannel := channel - nsdkChannel.Name = newName - - aChannel := convertChannel(channel) - aChannel.Name = newName - aChannel.Metadata = channels.Metadata{"field": "value2"} - asdkChannel := channel - asdkChannel.Name = newName - asdkChannel.Metadata = sdk.Metadata{"field": "value2"} - - cases := []struct { - desc string - domainID string - token string - session smqauthn.Session - channelReq sdk.Channel - updateChannelReq channels.Channel - svcRes channels.Channel - svcErr error - authenticateErr error - response sdk.Channel - err errors.SDKError - }{ - { - desc: "update channel name", - domainID: domainID, - token: validToken, - channelReq: sdk.Channel{ - ID: channel.ID, - Name: newName, - }, - updateChannelReq: channels.Channel{ - ID: channel.ID, - Name: newName, - }, - svcRes: nChannel, - svcErr: nil, - response: nsdkChannel, - err: nil, - }, - { - desc: "update channel metadata", - domainID: domainID, - token: validToken, - channelReq: sdk.Channel{ - ID: channel.ID, - Metadata: sdk.Metadata{ - "field": "value2", - }, - }, - updateChannelReq: channels.Channel{ - ID: channel.ID, - Metadata: channels.Metadata{"field": "value2"}, - }, - svcRes: mChannel, - svcErr: nil, - response: msdkChannel, - err: nil, - }, - { - desc: "update channel with every field defined", - domainID: domainID, - token: validToken, - channelReq: sdk.Channel{ - ID: channel.ID, - Name: newName, - Metadata: sdk.Metadata{"field": "value2"}, - }, - updateChannelReq: channels.Channel{ - ID: channel.ID, - Name: newName, - - Metadata: channels.Metadata{"field": "value2"}, - }, - svcRes: aChannel, - svcErr: nil, - response: asdkChannel, - err: nil, - }, - { - desc: "update channel name with invalid channel id", - domainID: domainID, - token: validToken, - channelReq: sdk.Channel{ - ID: wrongID, - Name: newName, - }, - updateChannelReq: channels.Channel{ - ID: wrongID, - Name: newName, - }, - svcRes: channels.Channel{}, - svcErr: svcerr.ErrNotFound, - response: sdk.Channel{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrNotFound, http.StatusNotFound), - }, - { - desc: "update channel description with invalid channel id", - domainID: domainID, - token: validToken, - channelReq: sdk.Channel{ - ID: wrongID, - }, - updateChannelReq: channels.Channel{ - ID: wrongID, - }, - svcRes: channels.Channel{}, - svcErr: svcerr.ErrNotFound, - response: sdk.Channel{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrNotFound, http.StatusNotFound), - }, - { - desc: "update channel metadata with invalid channel id", - domainID: domainID, - token: validToken, - channelReq: sdk.Channel{ - ID: wrongID, - Metadata: sdk.Metadata{ - "field": "value2", - }, - }, - updateChannelReq: channels.Channel{ - ID: wrongID, - Metadata: channels.Metadata{"field": "value2"}, - }, - svcRes: channels.Channel{}, - svcErr: svcerr.ErrNotFound, - response: sdk.Channel{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrNotFound, http.StatusNotFound), - }, - { - desc: "update channel with invalid token", - domainID: domainID, - token: invalidToken, - channelReq: sdk.Channel{ - ID: channel.ID, - Name: newName, - }, - updateChannelReq: channels.Channel{ - ID: channel.ID, - Name: newName, - }, - svcRes: channels.Channel{}, - authenticateErr: svcerr.ErrAuthentication, - response: sdk.Channel{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "update channel with empty token", - domainID: domainID, - token: "", - channelReq: sdk.Channel{ - ID: channel.ID, - Name: newName, - }, - updateChannelReq: channels.Channel{ - ID: channel.ID, - Name: newName, - }, - svcRes: channels.Channel{}, - svcErr: nil, - response: sdk.Channel{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "update channel with name that is too long", - domainID: domainID, - token: validToken, - channelReq: sdk.Channel{ - ID: channel.ID, - Name: strings.Repeat("a", 1025), - }, - updateChannelReq: channels.Channel{}, - svcRes: channels.Channel{}, - svcErr: nil, - response: sdk.Channel{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrNameSize, http.StatusBadRequest), - }, - { - desc: "update channel that can't be marshalled", - domainID: domainID, - token: validToken, - channelReq: sdk.Channel{ - ID: channel.ID, - Name: "test", - Metadata: map[string]any{ - "test": make(chan int), - }, - }, - updateChannelReq: channels.Channel{}, - svcRes: channels.Channel{}, - svcErr: nil, - response: sdk.Channel{}, - err: errors.NewSDKError(fmt.Errorf("json: unsupported type: chan int")), - }, - { - desc: "update channel with service response that can't be unmarshalled", - domainID: domainID, - token: validToken, - channelReq: sdk.Channel{ - ID: channel.ID, - Name: newName, - }, - updateChannelReq: channels.Channel{ - ID: channel.ID, - Name: newName, - }, - svcRes: channels.Channel{ - ID: generateUUID(t), - Metadata: channels.Metadata{ - "test": make(chan int), - }, - }, - svcErr: nil, - response: sdk.Channel{}, - err: errors.NewSDKError(fmt.Errorf("unexpected end of JSON input")), - }, - { - desc: "update channel with empty channel id", - domainID: domainID, - token: validToken, - channelReq: sdk.Channel{ - Name: newName, - }, - updateChannelReq: channels.Channel{}, - svcRes: channels.Channel{}, - svcErr: nil, - response: sdk.Channel{}, - err: errors.NewSDKError(apiutil.ErrMissingID), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := gsvc.On("UpdateChannel", mock.Anything, tc.session, tc.updateChannelReq).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.UpdateChannel(context.Background(), tc.channelReq, tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "UpdateChannel", mock.Anything, tc.session, tc.updateChannelReq) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestUpdateChannelTags(t *testing.T) { - ts, tsvc, auth := setupChannels() - defer ts.Close() - - sdkChannel := generateTestChannel(t) - updatedChannel := sdkChannel - updatedChannel.Tags = []string{"newTag1", "newTag2"} - updateChannelReq := sdk.Channel{ - ID: sdkChannel.ID, - Tags: updatedChannel.Tags, - } - - conf := sdk.Config{ - ChannelsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - cases := []struct { - desc string - domainID string - token string - session smqauthn.Session - updateChannelReq sdk.Channel - svcReq channels.Channel - svcRes channels.Channel - svcErr error - authenticateErr error - response sdk.Channel - err errors.SDKError - }{ - { - desc: "update channel tags successfully", - domainID: domainID, - token: validToken, - updateChannelReq: updateChannelReq, - svcReq: convertChannel(updateChannelReq), - svcRes: convertChannel(updatedChannel), - svcErr: nil, - response: updatedChannel, - err: nil, - }, - { - desc: "update channel tags with an invalid token", - domainID: domainID, - token: invalidToken, - updateChannelReq: updateChannelReq, - svcReq: convertChannel(updateChannelReq), - svcRes: channels.Channel{}, - authenticateErr: svcerr.ErrAuthorization, - response: sdk.Channel{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusUnauthorized), - }, - { - desc: "update channel tags with empty token", - domainID: domainID, - token: "", - updateChannelReq: updateChannelReq, - svcReq: convertChannel(updateChannelReq), - svcRes: channels.Channel{}, - svcErr: nil, - response: sdk.Channel{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "update channel tags with an invalid channel id", - domainID: domainID, - token: validToken, - updateChannelReq: sdk.Channel{ - ID: wrongID, - Tags: updatedChannel.Tags, - }, - svcReq: convertChannel(sdk.Channel{ - ID: wrongID, - Tags: updatedChannel.Tags, - }), - svcRes: channels.Channel{}, - svcErr: svcerr.ErrUpdateEntity, - response: sdk.Channel{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrUpdateEntity, http.StatusUnprocessableEntity), - }, - { - desc: "update channel tags with empty channel id", - domainID: domainID, - token: validToken, - updateChannelReq: sdk.Channel{ - ID: "", - Tags: updatedChannel.Tags, - }, - svcReq: convertChannel(sdk.Channel{ - ID: "", - Tags: updatedChannel.Tags, - }), - svcRes: channels.Channel{}, - svcErr: nil, - response: sdk.Channel{}, - err: errors.NewSDKError(apiutil.ErrMissingID), - }, - { - desc: "update channel tags with a request that can't be marshalled", - domainID: domainID, - token: validToken, - updateChannelReq: sdk.Channel{ - ID: "test", - Metadata: map[string]any{ - "test": make(chan int), - }, - }, - svcReq: channels.Channel{}, - svcRes: channels.Channel{}, - svcErr: nil, - response: sdk.Channel{}, - err: errors.NewSDKError(fmt.Errorf("json: unsupported type: chan int")), - }, - { - desc: "update channel tags with a response that can't be unmarshalled", - domainID: domainID, - token: validToken, - updateChannelReq: updateChannelReq, - svcReq: convertChannel(updateChannelReq), - svcRes: channels.Channel{ - Name: updatedChannel.Name, - Tags: updatedChannel.Tags, - Metadata: channels.Metadata{ - "test": make(chan int), - }, - }, - svcErr: nil, - response: sdk.Channel{}, - err: errors.NewSDKError(fmt.Errorf("unexpected end of JSON input")), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, mock.Anything).Return(tc.session, tc.authenticateErr) - svcCall := tsvc.On("UpdateChannelTags", mock.Anything, tc.session, tc.svcReq).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.UpdateChannelTags(context.Background(), tc.updateChannelReq, tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "UpdateChannelTags", mock.Anything, tc.session, tc.svcReq) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestEnableChannel(t *testing.T) { - ts, gsvc, auth := setupChannels() - defer ts.Close() - - conf := sdk.Config{ - ChannelsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - cases := []struct { - desc string - domainID string - token string - session smqauthn.Session - channelID string - svcRes channels.Channel - svcErr error - authenticateErr error - response sdk.Channel - err errors.SDKError - }{ - { - desc: "enable channel successfully", - domainID: domainID, - token: validToken, - channelID: channel.ID, - svcRes: convertChannel(channel), - svcErr: nil, - response: channel, - err: nil, - }, - { - desc: "enable channel with invalid token", - domainID: domainID, - token: invalidToken, - channelID: channel.ID, - svcRes: channels.Channel{}, - authenticateErr: svcerr.ErrAuthentication, - response: sdk.Channel{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "enable channel with empty token", - domainID: domainID, - token: "", - channelID: channel.ID, - svcRes: channels.Channel{}, - svcErr: nil, - response: sdk.Channel{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "enable channel with invalid channel id", - domainID: domainID, - token: validToken, - channelID: wrongID, - svcRes: channels.Channel{}, - svcErr: svcerr.ErrNotFound, - response: sdk.Channel{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrNotFound, http.StatusNotFound), - }, - { - desc: "enable channel with empty channel id", - domainID: domainID, - token: validToken, - channelID: "", - svcRes: channels.Channel{}, - svcErr: nil, - response: sdk.Channel{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "enable channel with service response that can't be unmarshalled", - domainID: domainID, - token: validToken, - channelID: channel.ID, - svcRes: channels.Channel{ - ID: generateUUID(t), - Metadata: channels.Metadata{ - "test": make(chan int), - }, - }, - svcErr: nil, - response: sdk.Channel{}, - err: errors.NewSDKError(fmt.Errorf("unexpected end of JSON input")), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := gsvc.On("EnableChannel", mock.Anything, tc.session, tc.channelID).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.EnableChannel(context.Background(), tc.channelID, tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "EnableChannel", mock.Anything, tc.session, tc.channelID) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestDisableChannel(t *testing.T) { - ts, gsvc, auth := setupChannels() - defer ts.Close() - - conf := sdk.Config{ - ChannelsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - dChannel := channel - dChannel.Status = channels.DisabledStatus.String() - - cases := []struct { - desc string - domainID string - token string - session smqauthn.Session - channelID string - svcRes channels.Channel - svcErr error - authenticateErr error - response sdk.Channel - err errors.SDKError - }{ - { - desc: "disable channel successfully", - domainID: domainID, - token: validToken, - channelID: channel.ID, - svcRes: convertChannel(dChannel), - svcErr: nil, - response: dChannel, - err: nil, - }, - { - desc: "disable channel with invalid token", - domainID: domainID, - token: invalidToken, - channelID: channel.ID, - svcRes: channels.Channel{}, - authenticateErr: svcerr.ErrAuthentication, - response: sdk.Channel{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "disable channel with empty token", - domainID: domainID, - token: "", - channelID: channel.ID, - svcRes: channels.Channel{}, - svcErr: nil, - response: sdk.Channel{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "disable channel with invalid channel id", - domainID: domainID, - token: validToken, - channelID: wrongID, - svcRes: channels.Channel{}, - svcErr: svcerr.ErrNotFound, - response: sdk.Channel{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrNotFound, http.StatusNotFound), - }, - { - desc: "disable channel with empty channel id", - domainID: domainID, - token: validToken, - channelID: "", - svcRes: channels.Channel{}, - svcErr: nil, - response: sdk.Channel{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "disable channel with service response that can't be unmarshalled", - domainID: domainID, - token: validToken, - channelID: channel.ID, - svcRes: channels.Channel{ - ID: generateUUID(t), - Metadata: channels.Metadata{ - "test": make(chan int), - }, - }, - svcErr: nil, - response: sdk.Channel{}, - err: errors.NewSDKError(fmt.Errorf("unexpected end of JSON input")), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := gsvc.On("DisableChannel", mock.Anything, tc.session, tc.channelID).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.DisableChannel(context.Background(), tc.channelID, tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "DisableChannel", mock.Anything, tc.session, tc.channelID) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestDeleteChannel(t *testing.T) { - ts, gsvc, auth := setupChannels() - defer ts.Close() - - conf := sdk.Config{ - ChannelsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - cases := []struct { - desc string - domainID string - token string - session smqauthn.Session - channelID string - svcErr error - authenticateErr error - err errors.SDKError - }{ - { - desc: "delete channel successfully", - domainID: domainID, - token: validToken, - channelID: channel.ID, - svcErr: nil, - err: nil, - }, - { - desc: "delete channel with invalid token", - domainID: domainID, - token: invalidToken, - channelID: channel.ID, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "delete channel with empty token", - domainID: domainID, - token: "", - channelID: channel.ID, - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "delete channel with invalid channel id", - domainID: domainID, - token: validToken, - channelID: wrongID, - svcErr: svcerr.ErrRemoveEntity, - err: errors.NewSDKErrorWithStatus(svcerr.ErrRemoveEntity, http.StatusUnprocessableEntity), - }, - { - desc: "delete channel with empty channel id", - domainID: domainID, - token: validToken, - channelID: "", - svcErr: svcerr.ErrRemoveEntity, - err: errors.NewSDKError(apiutil.ErrMissingID), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := gsvc.On("RemoveChannel", mock.Anything, tc.session, tc.channelID).Return(tc.svcErr) - err := mgsdk.DeleteChannel(context.Background(), tc.channelID, tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RemoveChannel", mock.Anything, tc.session, tc.channelID) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestConnect(t *testing.T) { - ts, gsvc, auth := setupChannels() - defer ts.Close() - - conf := sdk.Config{ - ChannelsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - clientID := generateUUID(t) - - cases := []struct { - desc string - domainID string - token string - session smqauthn.Session - connection sdk.Connection - svcErr error - authenticateRes smqauthn.Session - authenticateErr error - err errors.SDKError - }{ - { - desc: "connect successfully", - domainID: domainID, - token: validToken, - connection: sdk.Connection{ - ChannelIDs: []string{channel.ID}, - ClientIDs: []string{clientID}, - Types: []string{"Publish", "Subscribe"}, - }, - svcErr: nil, - err: nil, - }, - { - desc: "connect with invalid token", - domainID: domainID, - token: invalidToken, - connection: sdk.Connection{ - ChannelIDs: []string{channel.ID}, - ClientIDs: []string{clientID}, - Types: []string{"Publish", "Subscribe"}, - }, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "connect with empty token", - domainID: domainID, - token: "", - connection: sdk.Connection{ - ChannelIDs: []string{channel.ID}, - ClientIDs: []string{clientID}, - Types: []string{"Publish", "Subscribe"}, - }, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "connect with invalid channel id", - domainID: domainID, - token: validToken, - connection: sdk.Connection{ - ChannelIDs: []string{wrongID}, - ClientIDs: []string{clientID}, - Types: []string{"Publish", "Subscribe"}, - }, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "connect with empty channel id", - domainID: domainID, - token: validToken, - connection: sdk.Connection{ - ChannelIDs: []string{}, - ClientIDs: []string{clientID}, - Types: []string{"Publish", "Subscribe"}, - }, - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "connect with empty client id", - domainID: domainID, - token: validToken, - connection: sdk.Connection{ - ChannelIDs: []string{channel.ID}, - ClientIDs: []string{}, - Types: []string{"Publish", "Subscribe"}, - }, - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - connTypes := []connections.ConnType{} - for _, ct := range tc.connection.Types { - connType, err := connections.ParseConnType(ct) - assert.Nil(t, err, fmt.Sprintf("error parsing connection type %s", ct)) - connTypes = append(connTypes, connType) - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := gsvc.On("Connect", mock.Anything, tc.session, tc.connection.ChannelIDs, tc.connection.ClientIDs, connTypes).Return(tc.svcErr) - err := mgsdk.Connect(context.Background(), tc.connection, tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "Connect", mock.Anything, tc.session, tc.connection.ChannelIDs, tc.connection.ClientIDs, connTypes) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestDisconnect(t *testing.T) { - ts, gsvc, auth := setupChannels() - defer ts.Close() - - conf := sdk.Config{ - ChannelsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - clientID := generateUUID(t) - - cases := []struct { - desc string - domainID string - token string - session smqauthn.Session - disconnect sdk.Connection - svcErr error - authenticateRes smqauthn.Session - authenticateErr error - err errors.SDKError - }{ - { - desc: "disconnect successfully", - domainID: domainID, - token: validToken, - disconnect: sdk.Connection{ - ChannelIDs: []string{channel.ID}, - ClientIDs: []string{clientID}, - Types: []string{"Publish", "Subscribe"}, - }, - svcErr: nil, - err: nil, - }, - { - desc: "disconnect with invalid token", - domainID: domainID, - token: invalidToken, - disconnect: sdk.Connection{ - ChannelIDs: []string{channel.ID}, - ClientIDs: []string{clientID}, - Types: []string{"Publish", "Subscribe"}, - }, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "disconnect with empty token", - domainID: domainID, - token: "", - disconnect: sdk.Connection{ - ChannelIDs: []string{channel.ID}, - ClientIDs: []string{clientID}, - Types: []string{"Publish", "Subscribe"}, - }, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "disconnect with invalid channel id", - domainID: domainID, - token: validToken, - disconnect: sdk.Connection{ - ChannelIDs: []string{wrongID}, - ClientIDs: []string{clientID}, - Types: []string{"Publish", "Subscribe"}, - }, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(apiutil.ErrInvalidIDFormat, http.StatusBadRequest), - }, - { - desc: "disconnect with empty channel id", - domainID: domainID, - token: validToken, - disconnect: sdk.Connection{ - ChannelIDs: []string{}, - ClientIDs: []string{clientID}, - Types: []string{"Publish", "Subscribe"}, - }, - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "disconnect with empty client id", - domainID: domainID, - token: validToken, - disconnect: sdk.Connection{ - ChannelIDs: []string{channel.ID}, - ClientIDs: []string{}, - Types: []string{"Publish", "Subscribe"}, - }, - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - connTypes := []connections.ConnType{} - for _, ct := range tc.disconnect.Types { - connType, err := connections.ParseConnType(ct) - assert.Nil(t, err, fmt.Sprintf("error parsing connection type %s", ct)) - connTypes = append(connTypes, connType) - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := gsvc.On("Disconnect", mock.Anything, tc.session, tc.disconnect.ChannelIDs, tc.disconnect.ClientIDs, connTypes).Return(tc.svcErr) - err := mgsdk.Disconnect(context.Background(), tc.disconnect, tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "Disconnect", mock.Anything, tc.session, tc.disconnect.ChannelIDs, tc.disconnect.ClientIDs, connTypes) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestConnectClients(t *testing.T) { - ts, gsvc, auth := setupChannels() - defer ts.Close() - - conf := sdk.Config{ - ChannelsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - clientID := generateUUID(t) - - cases := []struct { - desc string - domainID string - token string - session smqauthn.Session - channelID string - clientID string - connType string - svcErr error - authenticateRes smqauthn.Session - authenticateErr error - err errors.SDKError - }{ - { - desc: "connect successfully", - domainID: domainID, - token: validToken, - channelID: channel.ID, - clientID: clientID, - connType: "Publish", - svcErr: nil, - err: nil, - }, - { - desc: "connect with invalid token", - domainID: domainID, - token: invalidToken, - channelID: channel.ID, - clientID: clientID, - connType: "Publish", - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "connect with empty token", - domainID: domainID, - token: "", - channelID: channel.ID, - clientID: clientID, - connType: "Publish", - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "connect with invalid channel id", - domainID: domainID, - token: validToken, - channelID: wrongID, - clientID: clientID, - connType: "Publish", - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "connect with empty channel id", - domainID: domainID, - token: validToken, - channelID: "", - clientID: clientID, - connType: "Publish", - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "connect with empty client id", - domainID: domainID, - token: validToken, - channelID: channel.ID, - clientID: "", - connType: "Publish", - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrInvalidIDFormat, http.StatusBadRequest), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - connType, err := connections.ParseConnType(tc.connType) - assert.Nil(t, err, fmt.Sprintf("error parsing connection type %s", tc.connType)) - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := gsvc.On("Connect", mock.Anything, tc.session, []string{tc.channelID}, []string{tc.clientID}, []connections.ConnType{connType}).Return(tc.svcErr) - err = mgsdk.ConnectClients(context.Background(), tc.channelID, []string{tc.clientID}, []string{tc.connType}, tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "Connect", mock.Anything, tc.session, []string{tc.channelID}, []string{tc.clientID}, []connections.ConnType{connType}) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestDisconnectClients(t *testing.T) { - ts, gsvc, auth := setupChannels() - defer ts.Close() - - conf := sdk.Config{ - ChannelsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - clientID := generateUUID(t) - - cases := []struct { - desc string - domainID string - token string - session smqauthn.Session - channelID string - clientID string - connType string - svcErr error - authenticateErr error - err errors.SDKError - }{ - { - desc: "disconnect successfully", - domainID: domainID, - token: validToken, - channelID: channel.ID, - clientID: clientID, - connType: "Publish", - svcErr: nil, - err: nil, - }, - { - desc: "disconnect with invalid token", - domainID: domainID, - token: invalidToken, - channelID: channel.ID, - clientID: clientID, - connType: "Publish", - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "disconnect with empty token", - domainID: domainID, - token: "", - channelID: channel.ID, - clientID: clientID, - connType: "Publish", - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "disconnect with invalid channel id", - domainID: domainID, - token: validToken, - channelID: wrongID, - clientID: clientID, - connType: "Publish", - err: errors.NewSDKErrorWithStatus(apiutil.ErrInvalidIDFormat, http.StatusBadRequest), - }, - { - desc: "disconnect with empty channel id", - domainID: domainID, - token: validToken, - channelID: "", - clientID: clientID, - connType: "Publish", - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "disconnect with empty client id", - domainID: domainID, - token: validToken, - channelID: channel.ID, - clientID: "", - connType: "Publish", - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrInvalidIDFormat, http.StatusBadRequest), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - connType, err := connections.ParseConnType(tc.connType) - assert.Nil(t, err, fmt.Sprintf("error parsing connection type %s", tc.connType)) - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := gsvc.On("Disconnect", mock.Anything, tc.session, []string{tc.channelID}, []string{tc.clientID}, []connections.ConnType{connType}).Return(tc.svcErr) - err = mgsdk.DisconnectClients(context.Background(), tc.channelID, []string{tc.clientID}, []string{tc.connType}, tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "Disconnect", mock.Anything, tc.session, []string{tc.channelID}, []string{tc.clientID}, []connections.ConnType{connType}) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestSetChannelParent(t *testing.T) { - ts, gsvc, auth := setupChannels() - defer ts.Close() - - conf := sdk.Config{ - ChannelsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - parentID := testsutil.GenerateUUID(t) - - cases := []struct { - desc string - domainID string - token string - session smqauthn.Session - channelID string - parentID string - svcErr error - authenticateErr error - err errors.SDKError - }{ - { - desc: "set channel parent successfully", - domainID: domainID, - token: validToken, - channelID: channel.ID, - parentID: parentID, - svcErr: nil, - err: nil, - }, - { - desc: "set channel parent with invalid token", - domainID: domainID, - token: invalidToken, - channelID: channel.ID, - parentID: parentID, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "set channel parent with empty token", - domainID: domainID, - token: "", - channelID: channel.ID, - parentID: parentID, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "set channel parent with invalid channel id", - domainID: domainID, - token: validToken, - channelID: wrongID, - parentID: parentID, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "set channel parent with empty channel id", - domainID: domainID, - token: validToken, - channelID: "", - parentID: parentID, - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "set channel parent with empty parent id", - domainID: domainID, - token: validToken, - channelID: channel.ID, - parentID: "", - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingParentGroupID, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := gsvc.On("SetParentGroup", mock.Anything, tc.session, tc.parentID, tc.channelID).Return(tc.svcErr) - err := mgsdk.SetChannelParent(context.Background(), tc.channelID, tc.domainID, tc.parentID, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "SetParentGroup", mock.Anything, tc.session, tc.parentID, tc.channelID) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestRemoveChannelParent(t *testing.T) { - ts, gsvc, auth := setupChannels() - defer ts.Close() - - conf := sdk.Config{ - ChannelsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - parentID := testsutil.GenerateUUID(t) - - cases := []struct { - desc string - domainID string - token string - session smqauthn.Session - channelID string - parentID string - svcErr error - authenticateErr error - err errors.SDKError - }{ - { - desc: "remove channel parent successfully", - domainID: domainID, - token: validToken, - channelID: channel.ID, - parentID: parentID, - svcErr: nil, - err: nil, - }, - { - desc: "remove channel parent with invalid token", - domainID: domainID, - token: invalidToken, - channelID: channel.ID, - parentID: parentID, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "remove channel parent with empty token", - domainID: domainID, - token: "", - channelID: channel.ID, - parentID: parentID, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "remove channel parent with invalid channel id", - domainID: domainID, - token: validToken, - channelID: wrongID, - parentID: parentID, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "remove channel parent with empty channel id", - domainID: domainID, - token: validToken, - channelID: "", - parentID: parentID, - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := gsvc.On("RemoveParentGroup", mock.Anything, tc.session, tc.channelID).Return(tc.svcErr) - err := mgsdk.RemoveChannelParent(context.Background(), tc.channelID, tc.domainID, tc.parentID, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RemoveParentGroup", mock.Anything, tc.session, tc.channelID) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func generateTestChannel(t *testing.T) sdk.Channel { - createdAt, err := time.Parse(time.RFC3339, "2023-03-03T00:00:00Z") - assert.Nil(t, err, fmt.Sprintf("unexpected error %s", err)) - updatedAt := createdAt - ch := sdk.Channel{ - ID: testsutil.GenerateUUID(&testing.T{}), - DomainID: testsutil.GenerateUUID(&testing.T{}), - Name: channelName, - Route: valid, - Metadata: sdk.Metadata{"role": "client"}, - CreatedAt: createdAt, - UpdatedAt: updatedAt, - Status: channels.EnabledStatus.String(), - Roles: validRoles, - } - return ch -} diff --git a/pkg/sdk/clients.go b/pkg/sdk/clients.go index c0d4bcb81..3269451fb 100644 --- a/pkg/sdk/clients.go +++ b/pkg/sdk/clients.go @@ -12,7 +12,6 @@ import ( apiutil "github.com/absmach/magistrala/api/http/util" "github.com/absmach/magistrala/pkg/errors" - "github.com/absmach/magistrala/pkg/roles" ) const ( @@ -27,20 +26,19 @@ const ( // Client represents magistrala client. type Client struct { - ID string `json:"id,omitempty"` - Name string `json:"name,omitempty"` - Tags []string `json:"tags,omitempty"` - DomainID string `json:"domain_id,omitempty"` - ParentGroup string `json:"parent_group_id,omitempty"` - Credentials ClientCredentials `json:"credentials"` - Metadata map[string]any `json:"metadata,omitempty"` - PrivateMetadata map[string]any `json:"private_metadata,omitempty"` - CreatedAt time.Time `json:"created_at,omitempty"` - UpdatedAt time.Time `json:"updated_at,omitempty"` - UpdatedBy string `json:"updated_by,omitempty"` - Status string `json:"status,omitempty"` - Permissions []string `json:"permissions,omitempty"` - Roles []roles.MemberRoleActions `json:"roles,omitempty"` + ID string `json:"id,omitempty"` + Name string `json:"name,omitempty"` + Tags []string `json:"tags,omitempty"` + DomainID string `json:"domain_id,omitempty"` + ParentGroup string `json:"parent_group_id,omitempty"` + Credentials ClientCredentials `json:"credentials"` + Metadata map[string]any `json:"metadata,omitempty"` + PrivateMetadata map[string]any `json:"private_metadata,omitempty"` + CreatedAt time.Time `json:"created_at,omitempty"` + UpdatedAt time.Time `json:"updated_at,omitempty"` + UpdatedBy string `json:"updated_by,omitempty"` + Status string `json:"status,omitempty"` + Permissions []string `json:"permissions,omitempty"` } type ClientCredentials struct { diff --git a/pkg/sdk/clients_test.go b/pkg/sdk/clients_test.go deleted file mode 100644 index a35109314..000000000 --- a/pkg/sdk/clients_test.go +++ /dev/null @@ -1,3286 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package sdk_test - -import ( - "context" - "fmt" - "net/http" - "net/http/httptest" - "strings" - "testing" - "time" - - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/clients" - api "github.com/absmach/magistrala/clients/api/http" - "github.com/absmach/magistrala/clients/mocks" - "github.com/absmach/magistrala/internal/testsutil" - mglog "github.com/absmach/magistrala/logger" - smqauthn "github.com/absmach/magistrala/pkg/authn" - authnmocks "github.com/absmach/magistrala/pkg/authn/mocks" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/pkg/roles" - sdk "github.com/absmach/magistrala/pkg/sdk" - "github.com/absmach/magistrala/pkg/uuid" - "github.com/go-chi/chi/v5" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" -) - -var clientID = "fe6b4e92-cc98-425e-b0aa-000000000001" - -func setupClients() (*httptest.Server, *mocks.Service, *authnmocks.Authentication) { - tsvc := new(mocks.Service) - - logger := mglog.NewMock() - mux := chi.NewRouter() - idp := uuid.NewMock() - authn := new(authnmocks.Authentication) - am := smqauthn.NewAuthNMiddleware(authn, smqauthn.WithAllowUnverifiedUser(true)) - api.MakeHandler(tsvc, am, mux, logger, "", idp) - - return httptest.NewServer(mux), tsvc, authn -} - -func TestCreateClient(t *testing.T) { - ts, tsvc, auth := setupClients() - defer ts.Close() - - client := generateTestClient(t, false) - createClientReq := sdk.Client{ - Name: client.Name, - Tags: client.Tags, - Credentials: client.Credentials, - Metadata: client.Metadata, - PrivateMetadata: client.PrivateMetadata, - Status: client.Status, - } - - conf := sdk.Config{ - ClientsURL: ts.URL, - } - - mgsdk := sdk.NewSDK(conf) - - cases := []struct { - desc string - domainID string - token string - session smqauthn.Session - createClientReq sdk.Client - svcReq clients.Client - svcRes []clients.Client - svcErr error - authenticateErr error - response sdk.Client - err errors.SDKError - }{ - { - desc: "create new client successfully", - domainID: domainID, - token: validToken, - createClientReq: createClientReq, - svcReq: convertClient(createClientReq), - svcRes: []clients.Client{convertClient(client)}, - svcErr: nil, - response: client, - err: nil, - }, - { - desc: "create new client with invalid token", - domainID: domainID, - token: invalidToken, - createClientReq: createClientReq, - svcReq: convertClient(createClientReq), - svcRes: []clients.Client{}, - authenticateErr: svcerr.ErrAuthentication, - response: sdk.Client{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "create new client with empty token", - domainID: domainID, - token: "", - createClientReq: createClientReq, - svcReq: convertClient(createClientReq), - svcRes: []clients.Client{}, - svcErr: nil, - response: sdk.Client{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "create an existing client", - domainID: domainID, - token: validToken, - createClientReq: createClientReq, - svcReq: convertClient(createClientReq), - svcRes: []clients.Client{}, - svcErr: svcerr.ErrCreateEntity, - response: sdk.Client{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrCreateEntity, http.StatusUnprocessableEntity), - }, - { - desc: "create a client with name too long", - domainID: domainID, - token: validToken, - createClientReq: sdk.Client{ - Name: strings.Repeat("a", 1025), - Tags: client.Tags, - Credentials: client.Credentials, - PrivateMetadata: client.PrivateMetadata, - Metadata: client.Metadata, - Status: client.Status, - }, - svcReq: clients.Client{}, - svcRes: []clients.Client{}, - svcErr: nil, - response: sdk.Client{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrNameSize, http.StatusBadRequest), - }, - { - desc: "create a client with invalid id", - domainID: domainID, - token: validToken, - createClientReq: sdk.Client{ - ID: "123456789", - Name: client.Name, - Tags: client.Tags, - Credentials: client.Credentials, - PrivateMetadata: client.PrivateMetadata, - Metadata: client.Metadata, - Status: client.Status, - }, - svcReq: clients.Client{}, - svcRes: []clients.Client{}, - svcErr: nil, - response: sdk.Client{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrInvalidIDFormat, http.StatusBadRequest), - }, - { - desc: "create a client with a request that can't be marshalled", - domainID: domainID, - token: validToken, - createClientReq: sdk.Client{ - Name: valid, - PrivateMetadata: map[string]any{ - valid: make(chan int), - }, - }, - svcReq: clients.Client{}, - svcRes: []clients.Client{}, - svcErr: nil, - response: sdk.Client{}, - err: errors.NewSDKError(fmt.Errorf("json: unsupported type: chan int")), - }, - { - desc: "create a client with a response that can't be unmarshalled", - domainID: domainID, - token: validToken, - createClientReq: createClientReq, - svcReq: convertClient(createClientReq), - svcRes: []clients.Client{{ - Name: client.Name, - Tags: client.Tags, - Credentials: clients.Credentials(client.Credentials), - PrivateMetadata: clients.Metadata{ - "test": make(chan int), - }, - }}, - svcErr: nil, - response: sdk.Client{}, - err: errors.NewSDKError(fmt.Errorf("unexpected end of JSON input")), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, mock.Anything).Return(tc.session, tc.authenticateErr) - svcCall := tsvc.On("CreateClients", mock.Anything, tc.session, []clients.Client{tc.svcReq}).Return(tc.svcRes, []roles.RoleProvision{}, tc.svcErr) - resp, err := mgsdk.CreateClient(context.Background(), tc.createClientReq, tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "CreateClients", mock.Anything, tc.session, []clients.Client{tc.svcReq}) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestCreateClients(t *testing.T) { - ts, tsvc, auth := setupClients() - defer ts.Close() - - sdkClients := []sdk.Client{} - for i := 0; i < 3; i++ { - client := generateTestClient(t, false) - sdkClients = append(sdkClients, client) - } - - conf := sdk.Config{ - ClientsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - cases := []struct { - desc string - domainID string - token string - session smqauthn.Session - createClientsRequest []sdk.Client - svcReq []clients.Client - svcRes []clients.Client - svcErr error - authenticateErr error - response []sdk.Client - err errors.SDKError - }{ - { - desc: "create new clients successfully", - domainID: domainID, - token: validToken, - createClientsRequest: sdkClients, - svcReq: convertClients(sdkClients...), - svcRes: convertClients(sdkClients...), - svcErr: nil, - response: sdkClients, - err: nil, - }, - { - desc: "create new clients with invalid token", - domainID: domainID, - token: invalidToken, - createClientsRequest: sdkClients, - svcReq: convertClients(sdkClients...), - svcRes: []clients.Client{}, - authenticateErr: svcerr.ErrAuthentication, - response: []sdk.Client{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "create new clients with empty token", - domainID: domainID, - token: "", - createClientsRequest: sdkClients, - svcReq: convertClients(sdkClients...), - svcRes: []clients.Client{}, - svcErr: nil, - response: []sdk.Client{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "create new clients with a request that can't be marshalled", - domainID: domainID, - token: validToken, - createClientsRequest: []sdk.Client{{Name: "test", PrivateMetadata: map[string]any{"test": make(chan int)}}}, - svcReq: convertClients(sdkClients...), - svcRes: []clients.Client{}, - svcErr: nil, - response: []sdk.Client{}, - err: errors.NewSDKError(fmt.Errorf("json: unsupported type: chan int")), - }, - { - desc: "create new clients with a response that can't be unmarshalled", - domainID: domainID, - token: validToken, - createClientsRequest: sdkClients, - svcReq: convertClients(sdkClients...), - svcRes: []clients.Client{{ - Name: sdkClients[0].Name, - Tags: sdkClients[0].Tags, - Credentials: clients.Credentials(sdkClients[0].Credentials), - PrivateMetadata: clients.Metadata{ - "test": make(chan int), - }, - }}, - svcErr: nil, - response: []sdk.Client{}, - err: errors.NewSDKError(fmt.Errorf("unexpected end of JSON input")), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, mock.Anything).Return(tc.session, tc.authenticateErr) - svcCall := tsvc.On("CreateClients", mock.Anything, tc.session, tc.svcReq).Return(tc.svcRes, []roles.RoleProvision{}, tc.svcErr) - resp, err := mgsdk.CreateClients(context.Background(), tc.createClientsRequest, tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "CreateClients", mock.Anything, tc.session, tc.svcReq) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestListClients(t *testing.T) { - ts, tsvc, auth := setupClients() - defer ts.Close() - - var sdkClients []sdk.Client - for i := 10; i < 100; i++ { - c := generateTestClient(t, false) - if i == 50 { - c.Status = clients.DisabledStatus.String() - c.Tags = []string{"tag1", "tag2"} - } - sdkClients = append(sdkClients, c) - } - - conf := sdk.Config{ - ClientsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - cases := []struct { - desc string - token string - domainID string - session smqauthn.Session - pageMeta sdk.PageMetadata - svcReq clients.Page - svcRes clients.ClientsPage - svcErr error - authenticateErr error - response sdk.ClientsPage - err errors.SDKError - }{ - { - desc: "list all clients successfully", - domainID: domainID, - token: validToken, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 100, - }, - svcReq: clients.Page{ - Actions: []string{}, - Order: "updated_at", - Dir: "desc", - Offset: 0, - Limit: 100, - }, - svcRes: clients.ClientsPage{ - Page: clients.Page{ - Offset: 0, - Limit: 100, - Total: uint64(len(sdkClients)), - }, - Clients: convertClients(sdkClients...), - }, - svcErr: nil, - response: sdk.ClientsPage{ - PageRes: sdk.PageRes{ - Limit: 100, - Total: uint64(len(sdkClients)), - }, - Clients: sdkClients, - }, - }, - { - desc: "list all clients with an invalid token", - domainID: domainID, - token: invalidToken, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 100, - }, - svcReq: clients.Page{ - Actions: []string{}, - Order: "updated_at", - Dir: "desc", - Offset: 0, - Limit: 100, - }, - svcRes: clients.ClientsPage{}, - authenticateErr: svcerr.ErrAuthentication, - response: sdk.ClientsPage{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "list all clients with limit greater than max", - domainID: domainID, - token: validToken, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 1000, - }, - svcReq: clients.Page{ - Actions: []string{}, - Order: "updated_at", - Dir: "desc", - }, - svcRes: clients.ClientsPage{}, - svcErr: nil, - response: sdk.ClientsPage{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrLimitSize, http.StatusBadRequest), - }, - { - desc: "list all clients with name size greater than max", - domainID: domainID, - token: validToken, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 100, - Name: strings.Repeat("a", 1025), - }, - svcReq: clients.Page{ - Actions: []string{}, - Order: "updated_at", - Dir: "desc", - }, - svcRes: clients.ClientsPage{}, - svcErr: nil, - response: sdk.ClientsPage{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrNameSize, http.StatusBadRequest), - }, - { - desc: "list all clients with status", - domainID: domainID, - token: validToken, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 100, - Status: clients.DisabledStatus.String(), - }, - svcReq: clients.Page{ - Actions: []string{}, - Order: "updated_at", - Dir: "desc", - Offset: 0, - Limit: 100, - Status: clients.DisabledStatus, - }, - svcRes: clients.ClientsPage{ - Page: clients.Page{ - Offset: 0, - Limit: 100, - Total: 1, - }, - Clients: convertClients(sdkClients[50]), - }, - svcErr: nil, - response: sdk.ClientsPage{ - PageRes: sdk.PageRes{ - Limit: 100, - Total: 1, - }, - Clients: []sdk.Client{sdkClients[50]}, - }, - err: nil, - }, - { - desc: "list all clients with tags", - domainID: domainID, - token: validToken, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 100, - Tags: sdk.TagsQuery{Elements: []string{"tag1"}, Operator: sdk.OrOp}, - }, - svcReq: clients.Page{ - Actions: []string{}, - Order: "updated_at", - Dir: "desc", - Offset: 0, - Limit: 100, - Tags: clients.TagsQuery{Elements: []string{"tag1"}, Operator: clients.OrOp}, - }, - svcRes: clients.ClientsPage{ - Page: clients.Page{ - Offset: 0, - Limit: 100, - Total: 1, - }, - Clients: convertClients(sdkClients[50]), - }, - svcErr: nil, - response: sdk.ClientsPage{ - PageRes: sdk.PageRes{ - Limit: 100, - Total: 1, - }, - Clients: []sdk.Client{sdkClients[50]}, - }, - err: nil, - }, - { - desc: "list all clients with invalid metadata", - domainID: domainID, - token: validToken, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 100, - Metadata: map[string]any{ - "test": make(chan int), - }, - }, - svcReq: clients.Page{ - Actions: []string{}, - Order: "updated_at", - Dir: "desc", - }, - svcRes: clients.ClientsPage{}, - svcErr: nil, - response: sdk.ClientsPage{}, - err: errors.NewSDKError(fmt.Errorf("json: unsupported type: chan int")), - }, - { - desc: "list all clients with response that can't be unmarshalled", - domainID: domainID, - token: validToken, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 100, - }, - svcReq: clients.Page{ - Actions: []string{}, - Order: "updated_at", - Dir: "desc", - Offset: 0, - Limit: 100, - }, - svcRes: clients.ClientsPage{ - Page: clients.Page{ - Offset: 0, - Limit: 100, - Total: 1, - }, - Clients: []clients.Client{{ - Name: sdkClients[0].Name, - Tags: sdkClients[0].Tags, - Credentials: clients.Credentials(sdkClients[0].Credentials), - PrivateMetadata: clients.Metadata{ - "test": make(chan int), - }, - }}, - }, - svcErr: nil, - response: sdk.ClientsPage{}, - err: errors.NewSDKError(fmt.Errorf("unexpected end of JSON input")), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, mock.Anything).Return(tc.session, tc.authenticateErr) - svcCall := tsvc.On("ListClients", mock.Anything, tc.session, tc.svcReq).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.Clients(context.Background(), tc.pageMeta, tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "ListClients", mock.Anything, tc.session, tc.svcReq) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestViewClient(t *testing.T) { - ts, tsvc, auth := setupClients() - defer ts.Close() - - sdkClient := generateTestClient(t, false) - sdkClientWithRoles := generateTestClient(t, true) - conf := sdk.Config{ - ClientsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - confRoles := sdk.Config{ - ClientsURL: ts.URL, - Roles: true, - } - mgsdkRoles := sdk.NewSDK(confRoles) - - cases := []struct { - desc string - domainID string - token string - session smqauthn.Session - withRoles bool - clientID string - svcRes clients.Client - svcErr error - authenticateErr error - response sdk.Client - err errors.SDKError - }{ - { - desc: "view client successfully", - domainID: domainID, - token: validToken, - withRoles: false, - clientID: sdkClient.ID, - svcRes: convertClient(sdkClient), - svcErr: nil, - response: sdkClient, - err: nil, - }, - { - desc: "view client successfully with roles", - domainID: domainID, - token: validToken, - withRoles: true, - clientID: sdkClientWithRoles.ID, - svcRes: convertClient(sdkClientWithRoles), - svcErr: nil, - response: sdkClientWithRoles, - err: nil, - }, - { - desc: "view client with an invalid token", - domainID: domainID, - token: invalidToken, - withRoles: false, - clientID: sdkClient.ID, - svcRes: clients.Client{}, - authenticateErr: svcerr.ErrAuthorization, - response: sdk.Client{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusUnauthorized), - }, - { - desc: "view client with empty token", - domainID: domainID, - token: "", - withRoles: false, - clientID: sdkClient.ID, - svcRes: clients.Client{}, - svcErr: nil, - response: sdk.Client{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "view client with an invalid client id", - domainID: domainID, - token: validToken, - withRoles: false, - clientID: wrongID, - svcRes: clients.Client{}, - svcErr: svcerr.ErrNotFound, - response: sdk.Client{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrNotFound, http.StatusNotFound), - }, - { - desc: "view client with empty client id", - domainID: domainID, - token: validToken, - withRoles: false, - clientID: "", - svcRes: clients.Client{}, - svcErr: nil, - response: sdk.Client{}, - err: errors.NewSDKError(apiutil.ErrMissingID), - }, - { - desc: "view client with response that can't be unmarshalled", - domainID: domainID, - token: validToken, - withRoles: false, - clientID: sdkClient.ID, - svcRes: clients.Client{ - Name: sdkClient.Name, - Tags: sdkClient.Tags, - Credentials: clients.Credentials(sdkClient.Credentials), - PrivateMetadata: clients.Metadata{ - "test": make(chan int), - }, - }, - svcErr: nil, - response: sdk.Client{}, - err: errors.NewSDKError(fmt.Errorf("unexpected end of JSON input")), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, mock.Anything).Return(tc.session, tc.authenticateErr) - svcCall := tsvc.On("View", mock.Anything, tc.session, tc.clientID, tc.withRoles).Return(tc.svcRes, tc.svcErr) - - var resp sdk.Client - var err error - switch tc.withRoles { - case true: - resp, err = mgsdkRoles.Client(context.Background(), tc.clientID, tc.domainID, tc.token) - default: - resp, err = mgsdk.Client(context.Background(), tc.clientID, tc.domainID, tc.token) - } - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.withRoles { - assert.Equal(t, resp.Roles, validRoles, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, validRoles, resp.Roles)) - } - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "View", mock.Anything, tc.session, tc.clientID, tc.withRoles) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestUpdateClient(t *testing.T) { - ts, tsvc, auth := setupClients() - defer ts.Close() - - sdkClient := generateTestClient(t, false) - updatedClient := sdkClient - updatedClient.Name = "newName" - updatedClient.Metadata = map[string]any{ - "newKey": "newValue", - } - updatedClient.PrivateMetadata = map[string]any{ - "privateKey": "privateValue", - } - updateClientReq := sdk.Client{ - ID: sdkClient.ID, - Name: updatedClient.Name, - Metadata: updatedClient.Metadata, - PrivateMetadata: updatedClient.PrivateMetadata, - } - - conf := sdk.Config{ - ClientsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - cases := []struct { - desc string - domainID string - token string - session smqauthn.Session - updateClientReq sdk.Client - svcReq clients.Client - svcRes clients.Client - svcErr error - authenticateErr error - response sdk.Client - err errors.SDKError - }{ - { - desc: "update client successfully", - domainID: domainID, - token: validToken, - updateClientReq: updateClientReq, - svcReq: convertClient(updateClientReq), - svcRes: convertClient(updatedClient), - svcErr: nil, - response: updatedClient, - err: nil, - }, - { - desc: "update client with an invalid token", - domainID: domainID, - token: invalidToken, - updateClientReq: updateClientReq, - svcReq: convertClient(updateClientReq), - svcRes: clients.Client{}, - authenticateErr: svcerr.ErrAuthorization, - response: sdk.Client{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusUnauthorized), - }, - { - desc: "update client with empty token", - domainID: domainID, - token: "", - updateClientReq: updateClientReq, - svcReq: convertClient(updateClientReq), - svcRes: clients.Client{}, - svcErr: nil, - response: sdk.Client{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "update client with an invalid client id", - domainID: domainID, - token: validToken, - updateClientReq: sdk.Client{ - ID: wrongID, - Name: updatedClient.Name, - }, - svcReq: convertClient(sdk.Client{ - ID: wrongID, - Name: updatedClient.Name, - }), - svcRes: clients.Client{}, - svcErr: svcerr.ErrUpdateEntity, - response: sdk.Client{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrUpdateEntity, http.StatusUnprocessableEntity), - }, - { - desc: "update client with empty client id", - domainID: domainID, - token: validToken, - - updateClientReq: sdk.Client{ - ID: "", - Name: updatedClient.Name, - }, - svcReq: convertClient(sdk.Client{ - ID: "", - Name: updatedClient.Name, - }), - svcRes: clients.Client{}, - svcErr: nil, - response: sdk.Client{}, - err: errors.NewSDKError(apiutil.ErrMissingID), - }, - { - desc: "update client with a request that can't be marshalled", - domainID: domainID, - token: validToken, - - updateClientReq: sdk.Client{ - ID: valid, - Metadata: map[string]any{ - "test": make(chan int), - }, - }, - svcReq: clients.Client{}, - svcRes: clients.Client{}, - svcErr: nil, - response: sdk.Client{}, - err: errors.NewSDKError(fmt.Errorf("json: unsupported type: chan int")), - }, - { - desc: "update client with a response that can't be unmarshalled", - domainID: domainID, - token: validToken, - updateClientReq: updateClientReq, - svcReq: convertClient(updateClientReq), - svcRes: clients.Client{ - Name: updatedClient.Name, - Tags: updatedClient.Tags, - Credentials: clients.Credentials(updatedClient.Credentials), - Metadata: clients.Metadata{ - "test": make(chan int), - }, - }, - svcErr: nil, - response: sdk.Client{}, - err: errors.NewSDKError(fmt.Errorf("unexpected end of JSON input")), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, mock.Anything).Return(tc.session, tc.authenticateErr) - svcCall := tsvc.On("Update", mock.Anything, tc.session, tc.svcReq).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.UpdateClient(context.Background(), tc.updateClientReq, tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "Update", mock.Anything, tc.session, tc.svcReq) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestUpdateClientTags(t *testing.T) { - ts, tsvc, auth := setupClients() - defer ts.Close() - - sdkClient := generateTestClient(t, false) - updatedClient := sdkClient - updatedClient.Tags = []string{"newTag1", "newTag2"} - updateClientReq := sdk.Client{ - ID: sdkClient.ID, - Tags: updatedClient.Tags, - } - - conf := sdk.Config{ - ClientsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - cases := []struct { - desc string - domainID string - token string - session smqauthn.Session - updateClientReq sdk.Client - svcReq clients.Client - svcRes clients.Client - svcErr error - authenticateErr error - response sdk.Client - err errors.SDKError - }{ - { - desc: "update client tags successfully", - domainID: domainID, - token: validToken, - updateClientReq: updateClientReq, - svcReq: convertClient(updateClientReq), - svcRes: convertClient(updatedClient), - svcErr: nil, - response: updatedClient, - err: nil, - }, - { - desc: "update client tags with an invalid token", - domainID: domainID, - token: invalidToken, - updateClientReq: updateClientReq, - svcReq: convertClient(updateClientReq), - svcRes: clients.Client{}, - authenticateErr: svcerr.ErrAuthorization, - response: sdk.Client{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusUnauthorized), - }, - { - desc: "update client tags with empty token", - domainID: domainID, - token: "", - updateClientReq: updateClientReq, - svcReq: convertClient(updateClientReq), - svcRes: clients.Client{}, - svcErr: nil, - response: sdk.Client{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "update client tags with an invalid client id", - domainID: domainID, - token: validToken, - updateClientReq: sdk.Client{ - ID: wrongID, - Tags: updatedClient.Tags, - }, - svcReq: convertClient(sdk.Client{ - ID: wrongID, - Tags: updatedClient.Tags, - }), - svcRes: clients.Client{}, - svcErr: svcerr.ErrUpdateEntity, - response: sdk.Client{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrUpdateEntity, http.StatusUnprocessableEntity), - }, - { - desc: "update client tags with empty client id", - domainID: domainID, - token: validToken, - updateClientReq: sdk.Client{ - ID: "", - Tags: updatedClient.Tags, - }, - svcReq: convertClient(sdk.Client{ - ID: "", - Tags: updatedClient.Tags, - }), - svcRes: clients.Client{}, - svcErr: nil, - response: sdk.Client{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "update client tags with a request that can't be marshalled", - domainID: domainID, - token: validToken, - updateClientReq: sdk.Client{ - ID: valid, - Metadata: map[string]any{ - "test": make(chan int), - }, - }, - svcReq: clients.Client{}, - svcRes: clients.Client{}, - svcErr: nil, - response: sdk.Client{}, - err: errors.NewSDKError(fmt.Errorf("json: unsupported type: chan int")), - }, - { - desc: "update client tags with a response that can't be unmarshalled", - domainID: domainID, - token: validToken, - updateClientReq: updateClientReq, - svcReq: convertClient(updateClientReq), - svcRes: clients.Client{ - Name: updatedClient.Name, - Tags: updatedClient.Tags, - Credentials: clients.Credentials(updatedClient.Credentials), - Metadata: clients.Metadata{ - "test": make(chan int), - }, - }, - svcErr: nil, - response: sdk.Client{}, - err: errors.NewSDKError(fmt.Errorf("unexpected end of JSON input")), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, mock.Anything).Return(tc.session, tc.authenticateErr) - svcCall := tsvc.On("UpdateTags", mock.Anything, tc.session, tc.svcReq).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.UpdateClientTags(context.Background(), tc.updateClientReq, tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "UpdateTags", mock.Anything, tc.session, tc.svcReq) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestUpdateClientSecret(t *testing.T) { - ts, tsvc, auth := setupClients() - defer ts.Close() - - sdkClient := generateTestClient(t, false) - newSecret := generateUUID(t) - updatedClient := sdkClient - updatedClient.Credentials.Secret = newSecret - - conf := sdk.Config{ - ClientsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - cases := []struct { - desc string - domainID string - token string - session smqauthn.Session - clientID string - newSecret string - svcRes clients.Client - svcErr error - authenticateErr error - response sdk.Client - err errors.SDKError - }{ - { - desc: "update client secret successfully", - domainID: domainID, - token: validToken, - clientID: sdkClient.ID, - newSecret: newSecret, - svcRes: convertClient(updatedClient), - svcErr: nil, - response: updatedClient, - err: nil, - }, - { - desc: "update client secret with an invalid token", - domainID: domainID, - token: invalidToken, - clientID: sdkClient.ID, - newSecret: newSecret, - svcRes: clients.Client{}, - authenticateErr: svcerr.ErrAuthorization, - response: sdk.Client{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusUnauthorized), - }, - { - desc: "update client secret with empty token", - domainID: domainID, - token: "", - clientID: sdkClient.ID, - newSecret: newSecret, - svcRes: clients.Client{}, - svcErr: nil, - response: sdk.Client{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "update client secret with an invalid client id", - domainID: domainID, - token: validToken, - clientID: wrongID, - newSecret: newSecret, - svcRes: clients.Client{}, - svcErr: svcerr.ErrUpdateEntity, - response: sdk.Client{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrUpdateEntity, http.StatusUnprocessableEntity), - }, - { - desc: "update client secret with empty client id", - domainID: domainID, - token: validToken, - clientID: "", - newSecret: newSecret, - svcRes: clients.Client{}, - svcErr: nil, - response: sdk.Client{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "update client with empty new secret", - domainID: domainID, - token: validToken, - clientID: sdkClient.ID, - newSecret: "", - svcRes: clients.Client{}, - svcErr: nil, - response: sdk.Client{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingSecret, http.StatusBadRequest), - }, - { - desc: "update client secret with a response that can't be unmarshalled", - domainID: domainID, - token: validToken, - clientID: sdkClient.ID, - newSecret: newSecret, - svcRes: clients.Client{ - Name: updatedClient.Name, - Tags: updatedClient.Tags, - Credentials: clients.Credentials(updatedClient.Credentials), - Metadata: clients.Metadata{ - "test": make(chan int), - }, - }, - svcErr: nil, - response: sdk.Client{}, - err: errors.NewSDKError(fmt.Errorf("unexpected end of JSON input")), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, mock.Anything).Return(tc.session, tc.authenticateErr) - svcCall := tsvc.On("UpdateSecret", mock.Anything, tc.session, tc.clientID, tc.newSecret).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.UpdateClientSecret(context.Background(), tc.clientID, tc.newSecret, tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "UpdateSecret", mock.Anything, tc.session, tc.clientID, tc.newSecret) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestEnableClient(t *testing.T) { - ts, tsvc, auth := setupClients() - defer ts.Close() - - client := generateTestClient(t, false) - enabledClient := client - enabledClient.Status = clients.EnabledStatus.String() - - conf := sdk.Config{ - ClientsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - cases := []struct { - desc string - domainID string - token string - session smqauthn.Session - clientID string - svcRes clients.Client - svcErr error - authenticateErr error - response sdk.Client - err errors.SDKError - }{ - { - desc: "enable client successfully", - domainID: domainID, - token: validToken, - clientID: client.ID, - svcRes: convertClient(enabledClient), - svcErr: nil, - response: enabledClient, - err: nil, - }, - { - desc: "enable client with an invalid token", - domainID: domainID, - token: invalidToken, - clientID: client.ID, - svcRes: clients.Client{}, - authenticateErr: svcerr.ErrAuthorization, - response: sdk.Client{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusUnauthorized), - }, - { - desc: "enable client with an invalid client id", - domainID: domainID, - token: validToken, - clientID: wrongID, - svcRes: clients.Client{}, - svcErr: svcerr.ErrEnableClient, - response: sdk.Client{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrEnableClient, http.StatusUnprocessableEntity), - }, - { - desc: "enable client with empty client id", - domainID: domainID, - token: validToken, - clientID: "", - svcRes: clients.Client{}, - svcErr: nil, - response: sdk.Client{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "enable client with a response that can't be unmarshalled", - domainID: domainID, - token: validToken, - clientID: client.ID, - svcRes: clients.Client{ - Name: enabledClient.Name, - Tags: enabledClient.Tags, - Credentials: clients.Credentials(enabledClient.Credentials), - Metadata: clients.Metadata{ - "test": make(chan int), - }, - }, - svcErr: nil, - response: sdk.Client{}, - err: errors.NewSDKError(fmt.Errorf("unexpected end of JSON input")), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, mock.Anything).Return(tc.session, tc.authenticateErr) - svcCall := tsvc.On("Enable", mock.Anything, tc.session, tc.clientID).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.EnableClient(context.Background(), tc.clientID, tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "Enable", mock.Anything, tc.session, tc.clientID) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestDisableClient(t *testing.T) { - ts, tsvc, auth := setupClients() - defer ts.Close() - - client := generateTestClient(t, false) - disabledClient := client - disabledClient.Status = clients.DisabledStatus.String() - - conf := sdk.Config{ - ClientsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - cases := []struct { - desc string - domainID string - token string - session smqauthn.Session - clientID string - svcRes clients.Client - svcErr error - authenticateErr error - response sdk.Client - err errors.SDKError - }{ - { - desc: "disable client successfully", - domainID: domainID, - token: validToken, - clientID: client.ID, - svcRes: convertClient(disabledClient), - svcErr: nil, - response: disabledClient, - err: nil, - }, - { - desc: "disable client with an invalid token", - domainID: domainID, - token: invalidToken, - clientID: client.ID, - svcRes: clients.Client{}, - authenticateErr: svcerr.ErrAuthorization, - response: sdk.Client{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusUnauthorized), - }, - { - desc: "disable client with an invalid client id", - domainID: domainID, - token: validToken, - clientID: wrongID, - svcRes: clients.Client{}, - svcErr: svcerr.ErrDisableClient, - response: sdk.Client{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrDisableClient, http.StatusUnprocessableEntity), - }, - { - desc: "disable client with empty client id", - domainID: domainID, - token: validToken, - clientID: "", - svcRes: clients.Client{}, - svcErr: nil, - response: sdk.Client{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "disable client with a response that can't be unmarshalled", - domainID: domainID, - token: validToken, - clientID: client.ID, - svcRes: clients.Client{ - Name: disabledClient.Name, - Tags: disabledClient.Tags, - Credentials: clients.Credentials(disabledClient.Credentials), - Metadata: clients.Metadata{ - "test": make(chan int), - }, - }, - svcErr: nil, - response: sdk.Client{}, - err: errors.NewSDKError(fmt.Errorf("unexpected end of JSON input")), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, mock.Anything).Return(tc.session, tc.authenticateErr) - svcCall := tsvc.On("Disable", mock.Anything, tc.session, tc.clientID).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.DisableClient(context.Background(), tc.clientID, tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "Disable", mock.Anything, tc.session, tc.clientID) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestDeleteClient(t *testing.T) { - ts, tsvc, auth := setupClients() - defer ts.Close() - - client := generateTestClient(t, false) - - conf := sdk.Config{ - ClientsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - cases := []struct { - desc string - domainID string - token string - session smqauthn.Session - clientID string - svcErr error - authenticateErr error - err errors.SDKError - }{ - { - desc: "delete client successfully", - domainID: domainID, - token: validToken, - clientID: client.ID, - svcErr: nil, - err: nil, - }, - { - desc: "delete client with an invalid token", - domainID: domainID, - token: invalidToken, - clientID: client.ID, - authenticateErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusUnauthorized), - }, - { - desc: "delete client with empty token", - domainID: domainID, - token: "", - clientID: client.ID, - svcErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "delete client with an invalid client id", - domainID: domainID, - token: validToken, - clientID: wrongID, - svcErr: svcerr.ErrRemoveEntity, - err: errors.NewSDKErrorWithStatus(svcerr.ErrRemoveEntity, http.StatusUnprocessableEntity), - }, - { - desc: "delete client with empty client id", - domainID: domainID, - token: validToken, - clientID: "", - svcErr: nil, - err: errors.NewSDKError(apiutil.ErrMissingID), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, mock.Anything).Return(tc.session, tc.authenticateErr) - svcCall := tsvc.On("Delete", mock.Anything, tc.session, tc.clientID).Return(tc.svcErr) - err := mgsdk.DeleteClient(context.Background(), tc.clientID, tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "Delete", mock.Anything, tc.session, tc.clientID) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestSetClientParent(t *testing.T) { - ts, csvc, auth := setupClients() - defer ts.Close() - - conf := sdk.Config{ - ClientsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - parentID := testsutil.GenerateUUID(t) - - cases := []struct { - desc string - domainID string - token string - session smqauthn.Session - clientID string - parentID string - svcErr error - authenticateErr error - err errors.SDKError - }{ - { - desc: "set client parent successfully", - domainID: domainID, - token: validToken, - clientID: clientID, - parentID: parentID, - svcErr: nil, - err: nil, - }, - { - desc: "set client parent with invalid token", - domainID: domainID, - token: invalidToken, - clientID: clientID, - parentID: parentID, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "set client parent with empty token", - domainID: domainID, - token: "", - clientID: clientID, - parentID: parentID, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "set client parent with invalid client id", - domainID: domainID, - token: validToken, - clientID: wrongID, - parentID: parentID, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "set client parent with empty client id", - domainID: domainID, - token: validToken, - clientID: "", - parentID: parentID, - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "set client parent with empty parent id", - domainID: domainID, - token: validToken, - clientID: clientID, - parentID: "", - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingParentGroupID, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("SetParentGroup", mock.Anything, tc.session, tc.parentID, tc.clientID).Return(tc.svcErr) - err := mgsdk.SetClientParent(context.Background(), tc.clientID, tc.domainID, tc.parentID, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "SetParentGroup", mock.Anything, tc.session, tc.parentID, tc.clientID) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestRemoveClientParent(t *testing.T) { - ts, csvc, auth := setupClients() - defer ts.Close() - - conf := sdk.Config{ - ClientsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - parentID := testsutil.GenerateUUID(t) - - cases := []struct { - desc string - domainID string - token string - session smqauthn.Session - clientID string - parentID string - svcErr error - authenticateErr error - err errors.SDKError - }{ - { - desc: "remove client parent successfully", - domainID: domainID, - token: validToken, - clientID: clientID, - parentID: parentID, - svcErr: nil, - err: nil, - }, - { - desc: "remove client parent with invalid token", - domainID: domainID, - token: invalidToken, - clientID: clientID, - parentID: parentID, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "remove client parent with empty token", - domainID: domainID, - token: "", - clientID: clientID, - parentID: parentID, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "remove client parent with invalid client id", - domainID: domainID, - token: validToken, - clientID: wrongID, - parentID: parentID, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "remove client parent with empty client id", - domainID: domainID, - token: validToken, - clientID: "", - parentID: parentID, - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("RemoveParentGroup", mock.Anything, tc.session, tc.clientID).Return(tc.svcErr) - err := mgsdk.RemoveClientParent(context.Background(), tc.clientID, tc.domainID, tc.parentID, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RemoveParentGroup", mock.Anything, tc.session, tc.clientID) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestCreateClientRole(t *testing.T) { - ts, csvc, auth := setupClients() - defer ts.Close() - - conf := sdk.Config{ - ClientsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - optionalActions := []string{"create", "update"} - optionalMembers := []string{testsutil.GenerateUUID(t), testsutil.GenerateUUID(t)} - rReq := sdk.RoleReq{ - RoleName: roleName, - OptionalActions: optionalActions, - OptionalMembers: optionalMembers, - } - userID := testsutil.GenerateUUID(t) - now := time.Now().UTC() - role := roles.Role{ - ID: testsutil.GenerateUUID(t), - Name: rReq.RoleName, - EntityID: clientID, - CreatedBy: userID, - CreatedAt: now, - } - roleProvision := roles.RoleProvision{ - Role: role, - OptionalActions: optionalActions, - OptionalMembers: optionalMembers, - } - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - clientID string - roleReq sdk.RoleReq - svcRes roles.RoleProvision - svcErr error - authenticateErr error - response sdk.Role - err errors.SDKError - }{ - { - desc: "create client role successfully", - token: validToken, - domainID: domainID, - clientID: clientID, - roleReq: rReq, - svcRes: roleProvision, - svcErr: nil, - response: convertRoleProvision(roleProvision), - err: nil, - }, - { - desc: "create client role with invalid token", - token: invalidToken, - domainID: domainID, - clientID: clientID, - roleReq: rReq, - svcRes: roles.RoleProvision{}, - authenticateErr: svcerr.ErrAuthentication, - response: sdk.Role{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "create client role with empty token", - token: "", - domainID: domainID, - clientID: clientID, - roleReq: rReq, - svcRes: roles.RoleProvision{}, - response: sdk.Role{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "create client role with invalid client id", - token: validToken, - domainID: domainID, - clientID: testsutil.GenerateUUID(t), - roleReq: rReq, - svcRes: roles.RoleProvision{}, - svcErr: svcerr.ErrAuthorization, - response: sdk.Role{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "create client role with empty client id", - token: validToken, - domainID: domainID, - clientID: "", - roleReq: rReq, - svcRes: roles.RoleProvision{}, - svcErr: nil, - response: sdk.Role{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrInvalidIDFormat, http.StatusBadRequest), - }, - { - desc: "create client role with empty role name", - token: validToken, - domainID: domainID, - clientID: clientID, - roleReq: sdk.RoleReq{ - RoleName: "", - OptionalActions: []string{"create", "update"}, - OptionalMembers: []string{testsutil.GenerateUUID(t), testsutil.GenerateUUID(t)}, - }, - svcRes: roles.RoleProvision{}, - svcErr: nil, - response: sdk.Role{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingRoleName, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("AddRole", mock.Anything, tc.session, tc.clientID, tc.roleReq.RoleName, tc.roleReq.OptionalActions, tc.roleReq.OptionalMembers).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.CreateClientRole(context.Background(), tc.clientID, tc.domainID, tc.roleReq, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "AddRole", mock.Anything, tc.session, tc.clientID, tc.roleReq.RoleName, tc.roleReq.OptionalActions, tc.roleReq.OptionalMembers) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestListClientRoles(t *testing.T) { - ts, csvc, auth := setupClients() - defer ts.Close() - - conf := sdk.Config{ - ClientsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - role := roles.Role{ - ID: testsutil.GenerateUUID(t), - Name: roleName, - EntityID: clientID, - CreatedBy: testsutil.GenerateUUID(t), - CreatedAt: time.Now().UTC(), - } - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - clientID string - pageMeta sdk.PageMetadata - svcRes roles.RolePage - svcErr error - authenticateErr error - response sdk.RolesPage - err errors.SDKError - }{ - { - desc: "list client roles successfully", - token: validToken, - domainID: domainID, - clientID: clientID, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - }, - svcRes: roles.RolePage{ - Total: 1, - Offset: 0, - Limit: 10, - Roles: []roles.Role{role}, - }, - svcErr: nil, - response: sdk.RolesPage{ - Total: 1, - Offset: 0, - Limit: 10, - Roles: []sdk.Role{convertRole(role)}, - }, - err: nil, - }, - { - desc: "list client roles with invalid token", - token: invalidToken, - domainID: domainID, - clientID: clientID, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - }, - svcRes: roles.RolePage{}, - authenticateErr: svcerr.ErrAuthentication, - response: sdk.RolesPage{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "list client roles with empty token", - token: "", - domainID: domainID, - clientID: clientID, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - }, - svcRes: roles.RolePage{}, - response: sdk.RolesPage{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "list client roles with invalid client id", - token: validToken, - domainID: domainID, - clientID: testsutil.GenerateUUID(t), - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - }, - svcRes: roles.RolePage{}, - svcErr: svcerr.ErrAuthorization, - response: sdk.RolesPage{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "list client roles with empty client id", - token: validToken, - domainID: domainID, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - }, - clientID: "", - svcRes: roles.RolePage{}, - svcErr: nil, - response: sdk.RolesPage{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("RetrieveAllRoles", mock.Anything, tc.session, tc.clientID, tc.pageMeta.Limit, tc.pageMeta.Offset).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.ClientRoles(context.Background(), tc.clientID, tc.domainID, tc.pageMeta, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RetrieveAllRoles", mock.Anything, tc.session, tc.clientID, tc.pageMeta.Limit, tc.pageMeta.Offset) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestViewClientRole(t *testing.T) { - ts, csvc, auth := setupClients() - defer ts.Close() - - conf := sdk.Config{ - ClientsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - role := roles.Role{ - ID: testsutil.GenerateUUID(t), - Name: roleName, - EntityID: clientID, - CreatedBy: testsutil.GenerateUUID(t), - CreatedAt: time.Now().UTC(), - } - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - clientID string - roleID string - svcRes roles.Role - svcErr error - authenticateErr error - response sdk.Role - err errors.SDKError - }{ - { - desc: "view client role successfully", - token: validToken, - domainID: domainID, - clientID: clientID, - roleID: role.ID, - svcRes: role, - svcErr: nil, - response: convertRole(role), - err: nil, - }, - { - desc: "view client role with invalid token", - token: invalidToken, - domainID: domainID, - clientID: clientID, - roleID: role.ID, - svcRes: roles.Role{}, - authenticateErr: svcerr.ErrAuthentication, - response: sdk.Role{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "view client role with empty token", - token: "", - domainID: domainID, - clientID: clientID, - roleID: role.ID, - svcRes: roles.Role{}, - response: sdk.Role{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "view client role with invalid client id", - token: validToken, - domainID: domainID, - clientID: testsutil.GenerateUUID(t), - roleID: role.ID, - svcRes: roles.Role{}, - svcErr: svcerr.ErrAuthorization, - response: sdk.Role{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "view client role with empty client id", - token: validToken, - domainID: domainID, - clientID: "", - roleID: role.ID, - svcRes: roles.Role{}, - svcErr: nil, - response: sdk.Role{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "view client role with invalid role id", - token: validToken, - domainID: domainID, - clientID: clientID, - roleID: invalid, - svcRes: roles.Role{}, - svcErr: svcerr.ErrAuthorization, - response: sdk.Role{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("RetrieveRole", mock.Anything, tc.session, tc.clientID, tc.roleID).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.ClientRole(context.Background(), tc.clientID, tc.roleID, tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RetrieveRole", mock.Anything, tc.session, tc.clientID, tc.roleID) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestUpdateClientRole(t *testing.T) { - ts, csvc, auth := setupClients() - defer ts.Close() - - conf := sdk.Config{ - ClientsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - roleID := testsutil.GenerateUUID(t) - newRoleName := valid - userID := testsutil.GenerateUUID(t) - createdAt := time.Now().UTC().Add(-time.Hour) - role := roles.Role{ - ID: testsutil.GenerateUUID(t), - Name: newRoleName, - EntityID: clientID, - CreatedBy: userID, - CreatedAt: createdAt, - UpdatedBy: userID, - UpdatedAt: time.Now().UTC(), - } - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - clientID string - roleID string - newRoleName string - svcRes roles.Role - svcErr error - authenticateErr error - response sdk.Role - err errors.SDKError - }{ - { - desc: "update client role successfully", - token: validToken, - domainID: domainID, - clientID: clientID, - roleID: roleID, - newRoleName: newRoleName, - svcRes: role, - svcErr: nil, - response: convertRole(role), - err: nil, - }, - { - desc: "update client role with invalid token", - token: invalidToken, - domainID: domainID, - clientID: clientID, - roleID: roleID, - newRoleName: newRoleName, - svcRes: roles.Role{}, - authenticateErr: svcerr.ErrAuthentication, - response: sdk.Role{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "update client role with empty token", - token: "", - domainID: domainID, - clientID: clientID, - roleID: roleID, - newRoleName: newRoleName, - svcRes: roles.Role{}, - response: sdk.Role{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "update client role with invalid client id", - token: validToken, - domainID: domainID, - clientID: testsutil.GenerateUUID(t), - roleID: roleID, - newRoleName: newRoleName, - svcRes: roles.Role{}, - svcErr: svcerr.ErrAuthorization, - response: sdk.Role{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "update client role with empty client id", - token: validToken, - domainID: domainID, - clientID: "", - roleID: roleID, - newRoleName: newRoleName, - svcRes: roles.Role{}, - svcErr: nil, - response: sdk.Role{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("UpdateRoleName", mock.Anything, tc.session, tc.clientID, tc.roleID, tc.newRoleName).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.UpdateClientRole(context.Background(), tc.clientID, tc.roleID, tc.newRoleName, tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "UpdateRoleName", mock.Anything, tc.session, tc.clientID, tc.roleID, tc.newRoleName) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestDeleteClientRole(t *testing.T) { - ts, csvc, auth := setupClients() - defer ts.Close() - - conf := sdk.Config{ - ClientsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - roleID := testsutil.GenerateUUID(t) - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - clientID string - roleID string - svcErr error - authenticateErr error - err errors.SDKError - }{ - { - desc: "delete client role successfully", - token: validToken, - domainID: domainID, - clientID: clientID, - roleID: roleID, - svcErr: nil, - err: nil, - }, - { - desc: "delete client role with invalid token", - token: invalidToken, - domainID: domainID, - clientID: clientID, - roleID: roleID, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "delete client role with empty token", - token: "", - domainID: domainID, - clientID: clientID, - roleID: roleID, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "delete client role with invalid client id", - token: validToken, - domainID: domainID, - clientID: testsutil.GenerateUUID(t), - roleID: roleID, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "delete client role with empty client id", - token: validToken, - domainID: domainID, - clientID: "", - roleID: roleID, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "delete client role with invalid role id", - token: validToken, - domainID: domainID, - clientID: clientID, - roleID: invalid, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("RemoveRole", mock.Anything, tc.session, tc.clientID, tc.roleID).Return(tc.svcErr) - err := mgsdk.DeleteClientRole(context.Background(), tc.clientID, tc.roleID, tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RemoveRole", mock.Anything, tc.session, tc.clientID, tc.roleID) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestAddClientRoleActions(t *testing.T) { - ts, csvc, auth := setupClients() - defer ts.Close() - - conf := sdk.Config{ - ClientsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - roleID := testsutil.GenerateUUID(t) - actions := []string{"create", "update"} - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - clientID string - roleID string - actions []string - svcRes []string - svcErr error - authenticateErr error - response []string - err errors.SDKError - }{ - { - desc: "add client role actions successfully", - token: validToken, - domainID: domainID, - clientID: clientID, - roleID: roleID, - actions: actions, - svcRes: actions, - svcErr: nil, - response: actions, - err: nil, - }, - { - desc: "add client role actions with invalid token", - token: invalidToken, - domainID: domainID, - clientID: clientID, - roleID: roleID, - actions: actions, - authenticateErr: svcerr.ErrAuthentication, - response: []string{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "add client role actions with empty token", - token: "", - domainID: domainID, - clientID: clientID, - roleID: roleID, - actions: actions, - response: []string{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "add client role actions with invalid client id", - token: validToken, - domainID: domainID, - clientID: testsutil.GenerateUUID(t), - roleID: roleID, - actions: actions, - svcErr: svcerr.ErrAuthorization, - response: []string{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "add client role actions with empty client id", - token: validToken, - domainID: domainID, - clientID: "", - roleID: roleID, - actions: actions, - response: []string{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "add client role actions with invalid role id", - token: validToken, - domainID: domainID, - clientID: clientID, - roleID: invalid, - actions: actions, - svcErr: svcerr.ErrAuthorization, - response: []string{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "add client role actions with empty actions", - token: validToken, - domainID: domainID, - clientID: clientID, - roleID: roleID, - actions: []string{}, - svcErr: nil, - response: []string{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingPolicyEntityType, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("RoleAddActions", mock.Anything, tc.session, tc.clientID, tc.roleID, tc.actions).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.AddClientRoleActions(context.Background(), tc.clientID, tc.roleID, tc.domainID, tc.actions, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RoleAddActions", mock.Anything, tc.session, tc.clientID, tc.roleID, tc.actions) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestListClientRoleActions(t *testing.T) { - ts, csvc, auth := setupClients() - defer ts.Close() - - conf := sdk.Config{ - ClientsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - roleID := testsutil.GenerateUUID(t) - actions := []string{"create", "update"} - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - clientID string - roleID string - svcRes []string - svcErr error - authenticateErr error - response []string - err errors.SDKError - }{ - { - desc: "list client role actions successfully", - token: validToken, - domainID: domainID, - clientID: clientID, - roleID: roleID, - svcRes: actions, - svcErr: nil, - response: actions, - err: nil, - }, - { - desc: "list client role actions with invalid token", - token: invalidToken, - domainID: domainID, - clientID: clientID, - roleID: roleID, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "list client role actions with empty token", - token: "", - domainID: domainID, - clientID: clientID, - roleID: roleID, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "list client role actions with invalid client id", - token: validToken, - domainID: domainID, - clientID: testsutil.GenerateUUID(t), - roleID: roleID, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "list client role actions with empty client id", - token: validToken, - domainID: domainID, - clientID: "", - roleID: roleID, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "list client role actions with invalid role id", - token: validToken, - domainID: domainID, - clientID: clientID, - roleID: invalid, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "list client role actions with empty role id", - token: validToken, - domainID: domainID, - clientID: clientID, - roleID: "", - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingRoleID, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("RoleListActions", mock.Anything, tc.session, tc.clientID, tc.roleID).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.ClientRoleActions(context.Background(), tc.clientID, tc.roleID, tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RoleListActions", mock.Anything, tc.session, tc.clientID, tc.roleID) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestRemoveClientRoleActions(t *testing.T) { - ts, csvc, auth := setupClients() - defer ts.Close() - - conf := sdk.Config{ - ClientsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - roleID := testsutil.GenerateUUID(t) - actions := []string{"create", "update"} - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - clientID string - roleID string - actions []string - svcErr error - authenticateErr error - err errors.SDKError - }{ - { - desc: "remove client role actions successfully", - token: validToken, - domainID: domainID, - clientID: clientID, - roleID: roleID, - actions: actions, - svcErr: nil, - err: nil, - }, - { - desc: "remove client role actions with invalid token", - token: invalidToken, - domainID: domainID, - clientID: clientID, - roleID: roleID, - actions: actions, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "remove client role actions with empty token", - token: "", - domainID: domainID, - clientID: clientID, - roleID: roleID, - actions: actions, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "remove client role actions with invalid client id", - token: validToken, - domainID: domainID, - clientID: testsutil.GenerateUUID(t), - roleID: roleID, - actions: actions, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "remove client role actions with empty client id", - token: validToken, - domainID: domainID, - clientID: "", - roleID: roleID, - actions: actions, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "remove client role actions with invalid role id", - token: validToken, - domainID: domainID, - clientID: clientID, - roleID: invalid, - actions: actions, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "remove client role actions with empty actions", - token: validToken, - domainID: domainID, - clientID: clientID, - roleID: roleID, - actions: []string{}, - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingPolicyEntityType, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("RoleRemoveActions", mock.Anything, tc.session, tc.clientID, tc.roleID, tc.actions).Return(tc.svcErr) - err := mgsdk.RemoveClientRoleActions(context.Background(), tc.clientID, tc.roleID, tc.domainID, tc.actions, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RoleRemoveActions", mock.Anything, tc.session, tc.clientID, tc.roleID, tc.actions) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestRemoveAllClientRoleActions(t *testing.T) { - ts, csvc, auth := setupClients() - defer ts.Close() - - conf := sdk.Config{ - ClientsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - roleID := testsutil.GenerateUUID(t) - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - clientID string - roleID string - svcErr error - authenticateErr error - err errors.SDKError - }{ - { - desc: "remove all client role actions successfully", - token: validToken, - domainID: domainID, - clientID: clientID, - roleID: roleID, - svcErr: nil, - err: nil, - }, - { - desc: "remove all client role actions with invalid token", - token: invalidToken, - domainID: domainID, - clientID: clientID, - roleID: roleID, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "remove all client role actions with empty token", - token: "", - domainID: domainID, - clientID: clientID, - roleID: roleID, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "remove all client role actions with invalid client id", - token: validToken, - domainID: domainID, - clientID: testsutil.GenerateUUID(t), - roleID: roleID, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "remove all client role actions with empty client id", - token: validToken, - domainID: domainID, - clientID: "", - roleID: roleID, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "remove all client role actions with invalid role id", - token: validToken, - domainID: domainID, - clientID: clientID, - roleID: invalid, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "remove all client role actions with empty role id", - token: validToken, - domainID: domainID, - clientID: clientID, - roleID: "", - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingRoleID, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("RoleRemoveAllActions", mock.Anything, tc.session, tc.clientID, tc.roleID).Return(tc.svcErr) - err := mgsdk.RemoveAllClientRoleActions(context.Background(), tc.clientID, tc.roleID, tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RoleRemoveAllActions", mock.Anything, tc.session, tc.clientID, tc.roleID) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestAddClientRoleMembers(t *testing.T) { - ts, csvc, auth := setupClients() - defer ts.Close() - - conf := sdk.Config{ - ClientsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - roleID := testsutil.GenerateUUID(t) - members := []string{"user1", "user2"} - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - clientID string - roleID string - members []string - svcRes []string - svcErr error - authenticateErr error - response []string - err errors.SDKError - }{ - { - desc: "add client role members successfully", - token: validToken, - domainID: domainID, - clientID: clientID, - roleID: roleID, - members: members, - svcRes: members, - svcErr: nil, - response: members, - err: nil, - }, - { - desc: "add client role members with invalid token", - token: invalidToken, - domainID: domainID, - clientID: clientID, - roleID: roleID, - members: members, - authenticateErr: svcerr.ErrAuthentication, - response: []string{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "add client role members with empty token", - token: "", - domainID: domainID, - clientID: clientID, - roleID: roleID, - members: members, - response: []string{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "add client role members with invalid client id", - token: validToken, - domainID: domainID, - clientID: testsutil.GenerateUUID(t), - roleID: roleID, - members: members, - svcErr: svcerr.ErrAuthorization, - response: []string{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "add client role members with empty client id", - token: validToken, - domainID: domainID, - clientID: "", - roleID: roleID, - members: members, - response: []string{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "add client role members with invalid role id", - token: validToken, - domainID: domainID, - clientID: clientID, - roleID: invalid, - members: members, - svcErr: svcerr.ErrAuthorization, - response: []string{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "add client role members with empty members", - token: validToken, - domainID: domainID, - clientID: clientID, - roleID: roleID, - members: []string{}, - svcErr: nil, - response: []string{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingRoleMembers, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("RoleAddMembers", mock.Anything, tc.session, tc.clientID, tc.roleID, tc.members).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.AddClientRoleMembers(context.Background(), tc.clientID, tc.roleID, tc.domainID, tc.members, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RoleAddMembers", mock.Anything, tc.session, tc.clientID, tc.roleID, tc.members) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestListClientRoleMembers(t *testing.T) { - ts, csvc, auth := setupClients() - defer ts.Close() - - conf := sdk.Config{ - ClientsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - roleID := testsutil.GenerateUUID(t) - members := []string{"user1", "user2"} - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - clientID string - roleID string - pageMeta sdk.PageMetadata - svcRes roles.MembersPage - svcErr error - authenticateErr error - response sdk.RoleMembersPage - err errors.SDKError - }{ - { - desc: "list client role members successfully", - token: validToken, - domainID: domainID, - clientID: clientID, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 5, - }, - roleID: roleID, - svcRes: roles.MembersPage{ - Total: 2, - Offset: 0, - Limit: 5, - Members: members, - }, - svcErr: nil, - response: sdk.RoleMembersPage{ - Total: 2, - Offset: 0, - Limit: 5, - Members: members, - }, - err: nil, - }, - { - desc: "list client role members with invalid token", - token: invalidToken, - domainID: domainID, - clientID: clientID, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 5, - }, - roleID: roleID, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "list client role members with empty token", - token: "", - domainID: domainID, - clientID: clientID, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 5, - }, - roleID: roleID, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "list client role members with invalid client id", - token: validToken, - domainID: domainID, - clientID: testsutil.GenerateUUID(t), - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 5, - }, - roleID: roleID, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "list client role members with empty client id", - token: validToken, - domainID: domainID, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 5, - }, - clientID: "", - roleID: roleID, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "list client role members with invalid role id", - token: validToken, - domainID: domainID, - clientID: clientID, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 5, - }, - roleID: invalid, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "list client role members with empty role id", - token: validToken, - domainID: domainID, - clientID: clientID, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 5, - }, - roleID: "", - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingRoleID, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("RoleListMembers", mock.Anything, tc.session, tc.clientID, tc.roleID, tc.pageMeta.Limit, tc.pageMeta.Offset).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.ClientRoleMembers(context.Background(), tc.clientID, tc.roleID, tc.domainID, tc.pageMeta, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RoleListMembers", mock.Anything, tc.session, tc.clientID, tc.roleID, tc.pageMeta.Limit, tc.pageMeta.Offset) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestRemoveClientRoleMembers(t *testing.T) { - ts, csvc, auth := setupClients() - defer ts.Close() - - conf := sdk.Config{ - ClientsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - roleID := testsutil.GenerateUUID(t) - members := []string{"user1", "user2"} - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - clientID string - roleID string - members []string - svcErr error - authenticateErr error - err errors.SDKError - }{ - { - desc: "remove client role members successfully", - token: validToken, - domainID: domainID, - clientID: clientID, - roleID: roleID, - members: members, - svcErr: nil, - err: nil, - }, - { - desc: "remove client role members with invalid token", - token: invalidToken, - domainID: domainID, - clientID: clientID, - roleID: roleID, - members: members, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "remove client role members with empty token", - token: "", - domainID: domainID, - clientID: clientID, - roleID: roleID, - members: members, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "remove client role members with invalid client id", - token: validToken, - domainID: domainID, - clientID: testsutil.GenerateUUID(t), - roleID: roleID, - members: members, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "remove client role members with empty client id", - token: validToken, - domainID: domainID, - clientID: "", - roleID: roleID, - members: members, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "remove client role members with invalid role id", - token: validToken, - domainID: domainID, - clientID: clientID, - roleID: invalid, - members: members, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "remove client role members with empty members", - token: validToken, - domainID: domainID, - clientID: clientID, - roleID: roleID, - members: []string{}, - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingRoleMembers, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("RoleRemoveMembers", mock.Anything, tc.session, tc.clientID, tc.roleID, tc.members).Return(tc.svcErr) - err := mgsdk.RemoveClientRoleMembers(context.Background(), tc.clientID, tc.roleID, tc.domainID, tc.members, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RoleRemoveMembers", mock.Anything, tc.session, tc.clientID, tc.roleID, tc.members) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestRemoveAllClientRoleMembers(t *testing.T) { - ts, csvc, auth := setupClients() - defer ts.Close() - - conf := sdk.Config{ - ClientsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - roleID := testsutil.GenerateUUID(t) - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - clientID string - roleID string - svcErr error - authenticateErr error - err errors.SDKError - }{ - { - desc: "remove all client role members successfully", - token: validToken, - domainID: domainID, - clientID: clientID, - roleID: roleID, - svcErr: nil, - err: nil, - }, - { - desc: "remove all client role members with invalid token", - token: invalidToken, - domainID: domainID, - clientID: clientID, - roleID: roleID, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "remove all client role members with empty token", - token: "", - domainID: domainID, - clientID: clientID, - roleID: roleID, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "remove all client role members with invalid client id", - token: validToken, - domainID: domainID, - clientID: testsutil.GenerateUUID(t), - roleID: roleID, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "remove all client role members with empty client id", - token: validToken, - domainID: domainID, - clientID: "", - roleID: roleID, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "remove all client role members with invalid role id", - token: validToken, - domainID: domainID, - clientID: clientID, - roleID: invalid, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "remove all client role members with empty role id", - token: validToken, - domainID: domainID, - clientID: clientID, - roleID: "", - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingRoleID, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("RoleRemoveAllMembers", mock.Anything, tc.session, tc.clientID, tc.roleID).Return(tc.svcErr) - err := mgsdk.RemoveAllClientRoleMembers(context.Background(), tc.clientID, tc.roleID, tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RoleRemoveAllMembers", mock.Anything, tc.session, tc.clientID, tc.roleID) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestListAvailableClientRoleActions(t *testing.T) { - ts, csvc, auth := setupClients() - defer ts.Close() - - conf := sdk.Config{ - ClientsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - actions := []string{"create", "update"} - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - svcRes []string - svcErr error - authenticateErr error - response []string - err errors.SDKError - }{ - { - desc: "list available role actions successfully", - token: validToken, - domainID: domainID, - svcRes: actions, - svcErr: nil, - response: actions, - err: nil, - }, - { - desc: "list available role actions with invalid token", - token: invalidToken, - domainID: domainID, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "list available role actions with empty token", - token: "", - domainID: domainID, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "list available role actions with empty domain id", - token: validToken, - domainID: "", - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingDomainID, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("ListAvailableActions", mock.Anything, tc.session).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.AvailableClientRoleActions(context.Background(), tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "ListAvailableActions", mock.Anything, tc.session) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func generateTestClient(t *testing.T, withRoles bool) sdk.Client { - createdAt, err := time.Parse(time.RFC3339, "2023-03-03T00:00:00Z") - assert.Nil(t, err, fmt.Sprintf("unexpected error %s", err)) - updatedAt := createdAt - var rl []roles.MemberRoleActions - if withRoles { - rl = validRoles - } - return sdk.Client{ - ID: testsutil.GenerateUUID(t), - Name: "clientname", - Credentials: sdk.ClientCredentials{ - Identity: "client@example.com", - Secret: generateUUID(t), - }, - Tags: []string{"tag1", "tag2"}, - Metadata: validMetadata, - PrivateMetadata: validMetadata, - Status: clients.EnabledStatus.String(), - CreatedAt: createdAt, - UpdatedAt: updatedAt, - Roles: rl, - } -} diff --git a/pkg/sdk/consumers_test.go b/pkg/sdk/consumers_test.go deleted file mode 100644 index 8644450c2..000000000 --- a/pkg/sdk/consumers_test.go +++ /dev/null @@ -1,454 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package sdk_test - -import ( - "context" - "fmt" - "net/http" - "net/http/httptest" - "testing" - - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/consumers/notifiers" - httpapi "github.com/absmach/magistrala/consumers/notifiers/api" - notmocks "github.com/absmach/magistrala/consumers/notifiers/mocks" - "github.com/absmach/magistrala/internal/testsutil" - mglog "github.com/absmach/magistrala/logger" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/pkg/sdk" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" -) - -var ( - ownerID = testsutil.GenerateUUID(&testing.T{}) - subID = testsutil.GenerateUUID(&testing.T{}) - sdkSubReq = sdk.Subscription{ - Topic: "topic", - Contact: "contact", - } - sdkSubRes = sdk.Subscription{ - Topic: "topic", - Contact: "contact", - OwnerID: ownerID, - ID: subID, - } - notSubReq = notifiers.Subscription{ - Contact: "contact", - Topic: "topic", - } - notSubRes = notifiers.Subscription{ - Contact: "contact", - Topic: "topic", - OwnerID: ownerID, - ID: subID, - } - instanceID = "instanceID" -) - -func setupSubscriptions() (*httptest.Server, *notmocks.Service) { - nsvc := new(notmocks.Service) - logger := mglog.NewMock() - mux := httpapi.MakeHandler(nsvc, logger, instanceID) - - return httptest.NewServer(mux), nsvc -} - -func TestCreateSubscription(t *testing.T) { - ts, nsvc := setupSubscriptions() - defer ts.Close() - - sdkConf := sdk.Config{ - UsersURL: ts.URL, - MsgContentType: contentType, - TLSVerification: false, - } - - mgsdk := sdk.NewSDK(sdkConf) - - cases := []struct { - desc string - subscription sdk.Subscription - token string - empty bool - id string - svcReq notifiers.Subscription - svcErr error - svcRes string - err errors.SDKError - }{ - { - desc: "create new subscription", - subscription: sdkSubReq, - token: validToken, - empty: false, - svcReq: notSubReq, - svcRes: subID, - svcErr: nil, - err: nil, - }, - { - desc: "create new subscription with empty token", - subscription: sdkSubReq, - token: "", - empty: true, - svcReq: notifiers.Subscription{}, - svcRes: "", - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "create new subscription with invalid token", - subscription: sdkSubReq, - token: invalidToken, - empty: true, - svcReq: notSubReq, - svcRes: "", - svcErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "create new subscription with empty topic", - subscription: sdk.Subscription{ - Topic: "", - Contact: "contact", - }, - token: validToken, - empty: true, - svcReq: notifiers.Subscription{}, - svcErr: nil, - svcRes: "", - err: errors.NewSDKErrorWithStatus(apiutil.ErrInvalidTopic, http.StatusBadRequest), - }, - { - desc: "create new subscription with empty contact", - subscription: sdk.Subscription{ - Topic: "topic", - Contact: "", - }, - token: validToken, - empty: true, - svcReq: notifiers.Subscription{}, - svcErr: nil, - svcRes: "", - err: errors.NewSDKErrorWithStatus(apiutil.ErrInvalidContact, http.StatusBadRequest), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := nsvc.On("CreateSubscription", mock.Anything, tc.token, tc.svcReq).Return(tc.svcRes, tc.svcErr) - loc, err := mgsdk.CreateSubscription(context.Background(), tc.subscription.Topic, tc.subscription.Contact, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.empty, loc == "") - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "CreateSubscription", mock.Anything, tc.token, tc.svcReq) - assert.True(t, ok) - } - svcCall.Unset() - }) - } -} - -func TestViewSubscription(t *testing.T) { - ts, nsvc := setupSubscriptions() - defer ts.Close() - sdkConf := sdk.Config{ - UsersURL: ts.URL, - MsgContentType: contentType, - TLSVerification: false, - } - - mgsdk := sdk.NewSDK(sdkConf) - - cases := []struct { - desc string - subID string - token string - svcRes notifiers.Subscription - svcErr error - response sdk.Subscription - err errors.SDKError - }{ - { - desc: "view existing subscription", - subID: subID, - token: validToken, - svcRes: notSubRes, - svcErr: nil, - response: sdkSubRes, - err: nil, - }, - { - desc: "view non-existent subscription", - subID: wrongID, - token: validToken, - svcRes: notifiers.Subscription{}, - svcErr: svcerr.ErrNotFound, - response: sdk.Subscription{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrNotFound, http.StatusNotFound), - }, - { - desc: "view subscription with invalid token", - subID: subID, - token: invalidToken, - svcRes: notifiers.Subscription{}, - svcErr: svcerr.ErrAuthentication, - response: sdk.Subscription{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "view subscription with empty token", - subID: subID, - token: "", - svcRes: notifiers.Subscription{}, - svcErr: nil, - response: sdk.Subscription{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := nsvc.On("ViewSubscription", mock.Anything, tc.token, tc.subID).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.ViewSubscription(context.Background(), tc.subID, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "ViewSubscription", mock.Anything, tc.token, tc.subID) - assert.True(t, ok) - } - svcCall.Unset() - }) - } -} - -func TestListSubscription(t *testing.T) { - ts, nsvc := setupSubscriptions() - defer ts.Close() - sdkConf := sdk.Config{ - UsersURL: ts.URL, - MsgContentType: contentType, - TLSVerification: false, - } - - mgsdk := sdk.NewSDK(sdkConf) - nSubs := 10 - noSubs := []notifiers.Subscription{} - sdSubs := []sdk.Subscription{} - for i := 0; i < nSubs; i++ { - nosub := notifiers.Subscription{ - OwnerID: ownerID, - Topic: fmt.Sprintf("topic_%d", i), - Contact: fmt.Sprintf("contact_%d", i), - } - noSubs = append(noSubs, nosub) - sdsub := sdk.Subscription{ - OwnerID: ownerID, - Topic: fmt.Sprintf("topic_%d", i), - Contact: fmt.Sprintf("contact_%d", i), - } - sdSubs = append(sdSubs, sdsub) - } - - cases := []struct { - desc string - token string - pageMeta sdk.PageMetadata - svcReq notifiers.PageMetadata - svcRes notifiers.Page - svcErr error - response sdk.SubscriptionPage - err errors.SDKError - }{ - { - desc: "list all subscription", - token: validToken, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - }, - svcReq: notifiers.PageMetadata{ - Offset: 0, - Limit: 10, - }, - svcRes: notifiers.Page{ - Total: 10, - Subscriptions: noSubs, - }, - svcErr: nil, - response: sdk.SubscriptionPage{ - PageRes: sdk.PageRes{ - Total: 10, - }, - Subscriptions: sdSubs, - }, - err: nil, - }, - { - desc: "list subscription with specific topic", - token: validToken, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - Topic: "topic_1", - }, - svcReq: notifiers.PageMetadata{ - Offset: 0, - Limit: 10, - Topic: "topic_1", - }, - svcRes: notifiers.Page{ - Total: uint(len(noSubs[1:2])), - Subscriptions: noSubs[1:2], - }, - svcErr: nil, - response: sdk.SubscriptionPage{ - PageRes: sdk.PageRes{ - Total: uint64(len(sdSubs[1:2])), - }, - Subscriptions: sdSubs[1:2], - }, - err: nil, - }, - { - desc: "list subscription with specific contact", - token: validToken, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - Contact: "contact_1", - }, - svcReq: notifiers.PageMetadata{ - Offset: 0, - Limit: 10, - Contact: "contact_1", - }, - svcRes: notifiers.Page{ - Total: uint(len(noSubs[1:2])), - Subscriptions: noSubs[1:2], - }, - svcErr: nil, - response: sdk.SubscriptionPage{ - PageRes: sdk.PageRes{ - Total: uint64(len(sdSubs[1:2])), - }, - Subscriptions: sdSubs[1:2], - }, - err: nil, - }, - { - desc: "list subscription with invalid token", - token: invalidToken, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - }, - svcReq: notifiers.PageMetadata{ - Offset: 0, - Limit: 10, - }, - svcRes: notifiers.Page{}, - svcErr: svcerr.ErrAuthentication, - response: sdk.SubscriptionPage{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "list subscription with empty token", - token: "", - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - }, - svcReq: notifiers.PageMetadata{}, - svcRes: notifiers.Page{}, - svcErr: nil, - response: sdk.SubscriptionPage{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := nsvc.On("ListSubscriptions", mock.Anything, tc.token, tc.svcReq).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.ListSubscriptions(context.Background(), tc.pageMeta, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "ListSubscriptions", mock.Anything, tc.token, tc.svcReq) - assert.True(t, ok) - } - svcCall.Unset() - }) - } -} - -func TestDeleteSubscription(t *testing.T) { - ts, nsvc := setupSubscriptions() - defer ts.Close() - sdkConf := sdk.Config{ - UsersURL: ts.URL, - MsgContentType: contentType, - TLSVerification: false, - } - - mgsdk := sdk.NewSDK(sdkConf) - - cases := []struct { - desc string - subID string - token string - svcErr error - err errors.SDKError - }{ - { - desc: "delete existing subscription", - subID: subID, - token: validToken, - svcErr: nil, - err: nil, - }, - { - desc: "delete non-existent subscription", - subID: wrongID, - token: validToken, - svcErr: svcerr.ErrRemoveEntity, - err: errors.NewSDKErrorWithStatus(svcerr.ErrRemoveEntity, http.StatusUnprocessableEntity), - }, - { - desc: "delete subscription with invalid token", - subID: subID, - token: invalidToken, - svcErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "delete subscription with empty token", - subID: subID, - token: "", - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "delete subscription with empty subID", - subID: "", - token: validToken, - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := nsvc.On("RemoveSubscription", mock.Anything, tc.token, tc.subID).Return(tc.svcErr) - err := mgsdk.DeleteSubscription(context.Background(), tc.subID, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RemoveSubscription", mock.Anything, tc.token, tc.subID) - assert.True(t, ok) - } - svcCall.Unset() - }) - } -} diff --git a/pkg/sdk/domains.go b/pkg/sdk/domains.go index f3a68b33f..b3a96dc50 100644 --- a/pkg/sdk/domains.go +++ b/pkg/sdk/domains.go @@ -12,7 +12,6 @@ import ( apiutil "github.com/absmach/magistrala/api/http/util" "github.com/absmach/magistrala/pkg/errors" - "github.com/absmach/magistrala/pkg/roles" ) const ( @@ -22,19 +21,18 @@ const ( // Domain represents magistrala domain. type Domain struct { - ID string `json:"id,omitempty"` - Name string `json:"name,omitempty"` - Metadata Metadata `json:"metadata,omitempty"` - Tags []string `json:"tags,omitempty"` - Route string `json:"route,omitempty"` - Status string `json:"status,omitempty"` - Permission string `json:"permission,omitempty"` - CreatedBy string `json:"created_by,omitempty"` - CreatedAt time.Time `json:"created_at,omitempty"` - UpdatedBy string `json:"updated_by,omitempty"` - UpdatedAt time.Time `json:"updated_at,omitempty"` - Permissions []string `json:"permissions,omitempty"` - Roles []roles.MemberRoleActions `json:"roles,omitempty"` + ID string `json:"id,omitempty"` + Name string `json:"name,omitempty"` + Metadata Metadata `json:"metadata,omitempty"` + Tags []string `json:"tags,omitempty"` + Route string `json:"route,omitempty"` + Status string `json:"status,omitempty"` + Permission string `json:"permission,omitempty"` + CreatedBy string `json:"created_by,omitempty"` + CreatedAt time.Time `json:"created_at,omitempty"` + UpdatedBy string `json:"updated_by,omitempty"` + UpdatedAt time.Time `json:"updated_at,omitempty"` + Permissions []string `json:"permissions,omitempty"` } func (sdk mgSDK) CreateDomain(ctx context.Context, domain Domain, token string) (Domain, errors.SDKError) { diff --git a/pkg/sdk/domains_test.go b/pkg/sdk/domains_test.go deleted file mode 100644 index 78715a091..000000000 --- a/pkg/sdk/domains_test.go +++ /dev/null @@ -1,2369 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package sdk_test - -import ( - "context" - "fmt" - "net/http" - "net/http/httptest" - "testing" - "time" - - api "github.com/absmach/magistrala/api/http" - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/domains" - domainapi "github.com/absmach/magistrala/domains/api/http" - "github.com/absmach/magistrala/domains/mocks" - "github.com/absmach/magistrala/internal/testsutil" - mglog "github.com/absmach/magistrala/logger" - smqauthn "github.com/absmach/magistrala/pkg/authn" - authnmocks "github.com/absmach/magistrala/pkg/authn/mocks" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/pkg/roles" - sdk "github.com/absmach/magistrala/pkg/sdk" - "github.com/absmach/magistrala/pkg/uuid" - "github.com/go-chi/chi/v5" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" -) - -var ( - authDomain, sdkDomain = generateTestDomain(&testing.T{}) - authDomainReq = domains.Domain{ - Name: authDomain.Name, - Metadata: authDomain.Metadata, - Tags: authDomain.Tags, - Route: authDomain.Route, - } - validRoles = []roles.MemberRoleActions{ - { - RoleID: "domain_role_id", - RoleName: "domain_role_name", - Actions: []string{"read", "delete"}, - AccessType: "direct", - }, - } - sdkDomainReq = sdk.Domain{ - Name: sdkDomain.Name, - Metadata: sdkDomain.Metadata, - Tags: sdkDomain.Tags, - Route: sdkDomain.Route, - Roles: validRoles, - } - updatedDomianName = "updated-domain" -) - -func setupDomains() (*httptest.Server, *mocks.Service, *authnmocks.Authentication) { - svc := new(mocks.Service) - logger := mglog.NewMock() - mux := chi.NewRouter() - idp := uuid.NewMock() - authn := new(authnmocks.Authentication) - am := smqauthn.NewAuthNMiddleware(authn, smqauthn.WithAllowUnverifiedUser(true)) - - handler := domainapi.MakeHandler(svc, am, mux, logger, "", idp) - return httptest.NewServer(handler), svc, authn -} - -func TestCreateDomain(t *testing.T) { - ds, svc, auth := setupDomains() - defer ds.Close() - - sdkConf := sdk.Config{ - DomainsURL: ds.URL, - MsgContentType: contentType, - } - - mgsdk := sdk.NewSDK(sdkConf) - - cases := []struct { - desc string - token string - session smqauthn.Session - domain sdk.Domain - svcReq domains.Domain - svcRes domains.Domain - svcErr error - authnErr error - response sdk.Domain - err error - }{ - { - desc: "create domain successfully", - token: validToken, - domain: sdkDomainReq, - svcReq: authDomainReq, - svcRes: authDomain, - svcErr: nil, - response: sdkDomain, - err: nil, - }, - { - desc: "create domain with invalid token", - token: invalidToken, - domain: sdkDomainReq, - svcReq: authDomainReq, - svcRes: domains.Domain{}, - authnErr: svcerr.ErrAuthentication, - response: sdk.Domain{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "create domain with empty token", - token: "", - domain: sdkDomainReq, - svcReq: authDomainReq, - svcRes: domains.Domain{}, - svcErr: nil, - response: sdk.Domain{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "create domain with empty name", - token: validToken, - domain: sdk.Domain{ - Name: "", - Metadata: sdkDomain.Metadata, - Tags: sdkDomain.Tags, - Route: sdkDomain.Route, - }, - svcReq: domains.Domain{}, - svcRes: domains.Domain{}, - svcErr: nil, - response: sdk.Domain{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingName, http.StatusBadRequest), - }, - { - desc: "create domain with request that cannot be marshalled", - token: validToken, - domain: sdk.Domain{ - Name: sdkDomain.Name, - Metadata: sdk.Metadata{ - "key": make(chan int), - }, - }, - svcReq: domains.Domain{}, - svcRes: domains.Domain{}, - svcErr: nil, - response: sdk.Domain{}, - err: errors.NewSDKError(fmt.Errorf("json: unsupported type: chan int")), - }, - { - desc: "create domain with response that cannot be unmarshalled", - token: validToken, - domain: sdkDomainReq, - svcReq: authDomainReq, - svcRes: domains.Domain{ - ID: authDomain.ID, - Name: authDomain.Name, - Metadata: domains.Metadata{ - "key": make(chan int), - }, - }, - svcErr: nil, - response: sdk.Domain{}, - err: errors.NewSDKError(fmt.Errorf("unexpected end of JSON input")), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, mock.Anything).Return(tc.session, tc.authnErr) - svcCall := svc.On("CreateDomain", mock.Anything, tc.session, tc.svcReq).Return(tc.svcRes, []roles.RoleProvision{}, tc.svcErr) - resp, err := mgsdk.CreateDomain(context.Background(), tc.domain, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "CreateDomain", mock.Anything, tc.session, tc.svcReq) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestUpdateDomain(t *testing.T) { - ds, svc, authn := setupDomains() - defer ds.Close() - - sdkConf := sdk.Config{ - DomainsURL: ds.URL, - MsgContentType: contentType, - } - - mgsdk := sdk.NewSDK(sdkConf) - - upDomainSDK := sdkDomain - upDomainSDK.Name = updatedDomianName - upDomainAuth := authDomain - upDomainAuth.Name = updatedDomianName - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - domain sdk.Domain - svcRes domains.Domain - svcErr error - authnErr error - response sdk.Domain - err error - }{ - { - desc: "update domain successfully", - token: validToken, - domainID: sdkDomain.ID, - domain: sdk.Domain{ - ID: sdkDomain.ID, - Name: updatedDomianName, - }, - svcRes: upDomainAuth, - svcErr: nil, - response: upDomainSDK, - err: nil, - }, - { - desc: "update domain with invalid token", - token: invalidToken, - domainID: sdkDomain.ID, - domain: sdk.Domain{ - ID: sdkDomain.ID, - Name: updatedDomianName, - }, - svcRes: domains.Domain{}, - authnErr: svcerr.ErrAuthentication, - response: sdk.Domain{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "update domain with empty token", - token: "", - domainID: sdkDomain.ID, - domain: sdk.Domain{ - ID: sdkDomain.ID, - Name: updatedDomianName, - }, - svcRes: domains.Domain{}, - svcErr: nil, - response: sdk.Domain{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "update domain with invalid domain ID", - token: validToken, - domainID: wrongID, - domain: sdk.Domain{ - ID: wrongID, - Name: updatedDomianName, - }, - svcRes: domains.Domain{}, - svcErr: svcerr.ErrAuthorization, - response: sdk.Domain{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "update domain with empty id", - token: validToken, - domainID: "", - domain: sdk.Domain{ - Name: sdkDomain.Name, - }, - svcRes: domains.Domain{}, - svcErr: nil, - response: sdk.Domain{}, - err: errors.NewSDKError(apiutil.ErrMissingID), - }, - { - desc: "update domain with request that cannot be marshalled", - token: validToken, - domainID: sdkDomain.ID, - domain: sdk.Domain{ - ID: sdkDomain.ID, - Name: sdkDomain.Name, - Metadata: sdk.Metadata{ - "key": make(chan int), - }, - }, - svcRes: domains.Domain{}, - svcErr: nil, - response: sdk.Domain{}, - err: errors.NewSDKError(fmt.Errorf("json: unsupported type: chan int")), - }, - { - desc: "update domain with response that cannot be unmarshalled", - token: validToken, - domainID: sdkDomain.ID, - domain: sdk.Domain{ - ID: sdkDomain.ID, - Name: sdkDomain.Name, - }, - svcRes: domains.Domain{ - ID: authDomain.ID, - Name: authDomain.Name, - Metadata: domains.Metadata{ - "key": make(chan int), - }, - }, - svcErr: nil, - response: sdk.Domain{}, - err: errors.NewSDKError(fmt.Errorf("unexpected end of JSON input")), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: tc.domainID + "_" + validID, UserID: validID, DomainID: tc.domainID} - } - authCall := authn.On("Authenticate", mock.Anything, mock.Anything).Return(tc.session, tc.authnErr) - svcCall := svc.On("UpdateDomain", mock.Anything, tc.session, tc.domainID, mock.Anything).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.UpdateDomain(context.Background(), tc.domain, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "UpdateDomain", mock.Anything, tc.session, tc.domainID, mock.Anything) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestViewDomain(t *testing.T) { - ds, svc, authn := setupDomains() - defer ds.Close() - - sdkConf := sdk.Config{ - DomainsURL: ds.URL, - MsgContentType: contentType, - } - - mgsdk := sdk.NewSDK(sdkConf) - - sdkConfRoles := sdk.Config{ - DomainsURL: ds.URL, - Roles: true, - } - mgsdkRoles := sdk.NewSDK(sdkConfRoles) - - cases := []struct { - desc string - token string - session smqauthn.Session - withRoles bool - domainID string - svcRes domains.Domain - svcErr error - authnErr error - response sdk.Domain - err error - }{ - { - desc: "view domain successfully", - token: validToken, - domainID: sdkDomain.ID, - withRoles: false, - svcRes: authDomain, - svcErr: nil, - response: sdkDomain, - err: nil, - }, - { - desc: "view domain successfully with roles", - token: validToken, - domainID: sdkDomain.ID, - withRoles: true, - svcRes: authDomain, - svcErr: nil, - response: sdkDomain, - err: nil, - }, - { - desc: "view domain with invalid token", - token: invalidToken, - domainID: sdkDomain.ID, - withRoles: false, - svcRes: domains.Domain{}, - authnErr: svcerr.ErrAuthentication, - response: sdk.Domain{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "view domain with empty token", - token: "", - domainID: sdkDomain.ID, - withRoles: false, - svcRes: domains.Domain{}, - svcErr: nil, - response: sdk.Domain{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "view domain with invalid domain ID", - token: validToken, - domainID: wrongID, - withRoles: false, - svcRes: domains.Domain{}, - svcErr: svcerr.ErrAuthorization, - response: sdk.Domain{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "view domain with empty id", - token: validToken, - domainID: "", - withRoles: false, - svcRes: domains.Domain{}, - svcErr: nil, - response: sdk.Domain{}, - err: errors.NewSDKError(apiutil.ErrMissingID), - }, - { - desc: "view domain with response that cannot be unmarshalled", - token: validToken, - domainID: sdkDomain.ID, - withRoles: false, - svcRes: domains.Domain{ - ID: authDomain.ID, - Name: authDomain.Name, - Metadata: domains.Metadata{ - "key": make(chan int), - }, - }, - svcErr: nil, - response: sdk.Domain{}, - err: errors.NewSDKError(fmt.Errorf("unexpected end of JSON input")), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: tc.domainID + "_" + validID, UserID: validID, DomainID: tc.domainID} - } - authCall := authn.On("Authenticate", mock.Anything, mock.Anything).Return(tc.session, tc.authnErr) - svcCall := svc.On("RetrieveDomain", mock.Anything, tc.session, tc.domainID, tc.withRoles).Return(tc.svcRes, tc.svcErr) - - var resp sdk.Domain - var err error - - switch tc.withRoles { - case true: - resp, err = mgsdkRoles.Domain(context.Background(), tc.domainID, tc.token) - default: - resp, err = mgsdk.Domain(context.Background(), tc.domainID, tc.token) - } - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.withRoles { - assert.Equal(t, resp.Roles, validRoles, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, validRoles, resp.Roles)) - } - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RetrieveDomain", mock.Anything, tc.session, tc.domainID, false) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestListDomians(t *testing.T) { - ds, svc, authn := setupDomains() - defer ds.Close() - - sdkConf := sdk.Config{ - DomainsURL: ds.URL, - MsgContentType: contentType, - } - - mgsdk := sdk.NewSDK(sdkConf) - - cases := []struct { - desc string - token string - session smqauthn.Session - pageMeta sdk.PageMetadata - svcReq domains.Page - svcRes domains.DomainsPage - svcErr error - authnErr error - response sdk.DomainsPage - err error - }{ - { - desc: "list domains successfully", - token: validToken, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - }, - svcReq: domains.Page{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - }, - svcRes: domains.DomainsPage{ - Total: 1, - Domains: []domains.Domain{authDomain}, - }, - svcErr: nil, - response: sdk.DomainsPage{ - PageRes: sdk.PageRes{ - Total: 1, - }, - Domains: []sdk.Domain{sdkDomain}, - }, - err: nil, - }, - { - desc: "list domains with invalid token", - token: invalidToken, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - }, - svcReq: domains.Page{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - }, - svcRes: domains.DomainsPage{}, - authnErr: svcerr.ErrAuthentication, - response: sdk.DomainsPage{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "list domains with empty token", - token: "", - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - }, - svcReq: domains.Page{}, - svcRes: domains.DomainsPage{}, - svcErr: nil, - response: sdk.DomainsPage{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "list domains with invalid page metadata", - token: validToken, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - Metadata: sdk.Metadata{ - "key": make(chan int), - }, - }, - svcReq: domains.Page{}, - svcRes: domains.DomainsPage{}, - svcErr: nil, - response: sdk.DomainsPage{}, - err: errors.NewSDKError(fmt.Errorf("json: unsupported type: chan int")), - }, - { - desc: "list domains with request that cannot be marshalled", - token: validToken, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - }, - svcReq: domains.Page{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - }, - svcRes: domains.DomainsPage{ - Total: 1, - Domains: []domains.Domain{{ - Name: authDomain.Name, - Metadata: domains.Metadata{"key": make(chan int)}, - }}, - }, - svcErr: nil, - response: sdk.DomainsPage{}, - err: errors.NewSDKError(fmt.Errorf("unexpected end of JSON input")), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := authn.On("Authenticate", mock.Anything, mock.Anything).Return(tc.session, tc.authnErr) - svcCall := svc.On("ListDomains", mock.Anything, tc.session, tc.svcReq).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.Domains(context.Background(), tc.pageMeta, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "ListDomains", mock.Anything, tc.session, mock.Anything) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestEnableDomain(t *testing.T) { - ds, svc, authn := setupDomains() - defer ds.Close() - - sdkConf := sdk.Config{ - DomainsURL: ds.URL, - MsgContentType: contentType, - } - - mgsdk := sdk.NewSDK(sdkConf) - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - svcRes domains.Domain - svcErr error - authnErr error - err error - }{ - { - desc: "enable domain successfully", - token: validToken, - domainID: sdkDomain.ID, - svcRes: authDomain, - svcErr: nil, - err: nil, - }, - { - desc: "enable domain with invalid token", - token: invalidToken, - domainID: sdkDomain.ID, - svcRes: domains.Domain{}, - authnErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "enable domain with empty token", - token: "", - domainID: sdkDomain.ID, - svcRes: domains.Domain{}, - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "enable domain with empty domain id", - token: validToken, - domainID: "", - svcRes: domains.Domain{}, - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingDomainID, http.StatusBadRequest), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: tc.domainID + "_" + validID, UserID: validID, DomainID: tc.domainID} - } - authCall := authn.On("Authenticate", mock.Anything, mock.Anything).Return(tc.session, tc.authnErr) - svcCall := svc.On("EnableDomain", mock.Anything, tc.session, tc.domainID).Return(tc.svcRes, tc.svcErr) - err := mgsdk.EnableDomain(context.Background(), tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "EnableDomain", mock.Anything, tc.session, tc.domainID) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestDisableDomain(t *testing.T) { - ds, svc, authn := setupDomains() - defer ds.Close() - - sdkConf := sdk.Config{ - DomainsURL: ds.URL, - MsgContentType: contentType, - } - - mgsdk := sdk.NewSDK(sdkConf) - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - svcRes domains.Domain - svcErr error - authnErr error - err error - }{ - { - desc: "disable domain successfully", - token: validToken, - domainID: sdkDomain.ID, - svcRes: authDomain, - svcErr: nil, - err: nil, - }, - { - desc: "disable domain with invalid token", - token: invalidToken, - domainID: sdkDomain.ID, - svcRes: domains.Domain{}, - authnErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "disable domain with empty token", - token: "", - domainID: sdkDomain.ID, - svcRes: domains.Domain{}, - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "disable domain with empty domain id", - token: validToken, - domainID: "", - svcRes: domains.Domain{}, - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingDomainID, http.StatusBadRequest), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: tc.domainID + "_" + validID, UserID: validID, DomainID: tc.domainID} - } - authCall := authn.On("Authenticate", mock.Anything, mock.Anything).Return(tc.session, tc.authnErr) - svcCall := svc.On("DisableDomain", mock.Anything, tc.session, tc.domainID).Return(tc.svcRes, tc.svcErr) - err := mgsdk.DisableDomain(context.Background(), tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "DisableDomain", mock.Anything, tc.session, tc.domainID) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestFreezeDomain(t *testing.T) { - ds, svc, authn := setupDomains() - defer ds.Close() - - sdkConf := sdk.Config{ - DomainsURL: ds.URL, - MsgContentType: contentType, - } - - mgsdk := sdk.NewSDK(sdkConf) - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - svcRes domains.Domain - svcErr error - authnErr error - err error - }{ - { - desc: "freeze domain successfully", - token: validToken, - domainID: sdkDomain.ID, - svcRes: authDomain, - svcErr: nil, - err: nil, - }, - { - desc: "freeze domain with invalid token", - token: invalidToken, - domainID: sdkDomain.ID, - svcRes: domains.Domain{}, - authnErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "freeze domain with empty token", - token: "", - domainID: sdkDomain.ID, - svcRes: domains.Domain{}, - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "freeze domain with empty domain id", - token: validToken, - domainID: "", - svcRes: domains.Domain{}, - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingDomainID, http.StatusBadRequest), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: tc.domainID + "_" + validID, UserID: validID, DomainID: tc.domainID} - } - authCall := authn.On("Authenticate", mock.Anything, mock.Anything).Return(tc.session, tc.authnErr) - svcCall := svc.On("FreezeDomain", mock.Anything, tc.session, tc.domainID).Return(tc.svcRes, tc.svcErr) - err := mgsdk.FreezeDomain(context.Background(), tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "FreezeDomain", mock.Anything, tc.session, tc.domainID) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestCreateDomainRole(t *testing.T) { - ts, csvc, auth := setupDomains() - defer ts.Close() - - conf := sdk.Config{ - DomainsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - optionalActions := []string{"create", "update"} - optionalMembers := []string{testsutil.GenerateUUID(t), testsutil.GenerateUUID(t)} - rReq := sdk.RoleReq{ - RoleName: roleName, - OptionalActions: optionalActions, - OptionalMembers: optionalMembers, - } - userID := testsutil.GenerateUUID(t) - now := time.Now().UTC() - role := roles.Role{ - ID: testsutil.GenerateUUID(t), - Name: rReq.RoleName, - EntityID: domainID, - CreatedBy: userID, - CreatedAt: now, - } - roleProvision := roles.RoleProvision{ - Role: role, - OptionalActions: optionalActions, - OptionalMembers: optionalMembers, - } - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - roleReq sdk.RoleReq - svcRes roles.RoleProvision - svcErr error - authenticateErr error - response sdk.Role - err errors.SDKError - }{ - { - desc: "create domain role successfully", - token: validToken, - domainID: domainID, - roleReq: rReq, - svcRes: roleProvision, - svcErr: nil, - response: convertRoleProvision(roleProvision), - err: nil, - }, - { - desc: "create domain role with invalid token", - token: invalidToken, - domainID: domainID, - roleReq: rReq, - svcRes: roles.RoleProvision{}, - authenticateErr: svcerr.ErrAuthentication, - response: sdk.Role{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "create domain role with empty token", - token: "", - domainID: domainID, - roleReq: rReq, - svcRes: roles.RoleProvision{}, - response: sdk.Role{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "create domain role with invalid domain id", - token: validToken, - domainID: testsutil.GenerateUUID(t), - roleReq: rReq, - svcRes: roles.RoleProvision{}, - svcErr: svcerr.ErrAuthorization, - response: sdk.Role{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "create domain role with empty domain id", - token: validToken, - domainID: "", - roleReq: rReq, - svcRes: roles.RoleProvision{}, - svcErr: nil, - response: sdk.Role{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingDomainID, http.StatusBadRequest), - }, - { - desc: "create domain role with empty role name", - token: validToken, - domainID: domainID, - roleReq: sdk.RoleReq{ - RoleName: "", - OptionalActions: []string{"create", "update"}, - OptionalMembers: []string{testsutil.GenerateUUID(t), testsutil.GenerateUUID(t)}, - }, - svcRes: roles.RoleProvision{}, - svcErr: nil, - response: sdk.Role{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingRoleName, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: tc.domainID + "_" + validID, UserID: validID, DomainID: tc.domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("AddRole", mock.Anything, tc.session, tc.domainID, tc.roleReq.RoleName, tc.roleReq.OptionalActions, tc.roleReq.OptionalMembers).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.CreateDomainRole(context.Background(), tc.domainID, tc.roleReq, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "AddRole", mock.Anything, tc.session, tc.domainID, tc.roleReq.RoleName, tc.roleReq.OptionalActions, tc.roleReq.OptionalMembers) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestListDomainRoles(t *testing.T) { - ts, csvc, auth := setupDomains() - defer ts.Close() - - conf := sdk.Config{ - DomainsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - role := roles.Role{ - ID: testsutil.GenerateUUID(t), - Name: roleName, - EntityID: domainID, - CreatedBy: testsutil.GenerateUUID(t), - CreatedAt: time.Now().UTC(), - } - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - pageMeta sdk.PageMetadata - svcRes roles.RolePage - svcErr error - authenticateErr error - response sdk.RolesPage - err errors.SDKError - }{ - { - desc: "list domain roles successfully", - token: validToken, - domainID: domainID, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - }, - svcRes: roles.RolePage{ - Total: 1, - Offset: 0, - Limit: 10, - Roles: []roles.Role{role}, - }, - svcErr: nil, - response: sdk.RolesPage{ - Total: 1, - Offset: 0, - Limit: 10, - Roles: []sdk.Role{convertRole(role)}, - }, - err: nil, - }, - { - desc: "list domain roles with invalid token", - token: invalidToken, - domainID: domainID, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - }, - svcRes: roles.RolePage{}, - authenticateErr: svcerr.ErrAuthentication, - response: sdk.RolesPage{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "list domain roles with empty token", - token: "", - domainID: domainID, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - }, - svcRes: roles.RolePage{}, - response: sdk.RolesPage{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "list domain roles with invalid domain id", - token: validToken, - domainID: testsutil.GenerateUUID(t), - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - }, - svcRes: roles.RolePage{}, - svcErr: svcerr.ErrAuthorization, - response: sdk.RolesPage{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "list domain roles with empty domain id", - token: validToken, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - }, - domainID: "", - svcRes: roles.RolePage{}, - svcErr: nil, - response: sdk.RolesPage{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingDomainID, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: tc.domainID + "_" + validID, UserID: validID, DomainID: tc.domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("RetrieveAllRoles", mock.Anything, tc.session, tc.domainID, tc.pageMeta.Limit, tc.pageMeta.Offset).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.DomainRoles(context.Background(), tc.domainID, tc.pageMeta, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RetrieveAllRoles", mock.Anything, tc.session, tc.domainID, tc.pageMeta.Limit, tc.pageMeta.Offset) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestViewClietRole(t *testing.T) { - ts, csvc, auth := setupDomains() - defer ts.Close() - - conf := sdk.Config{ - DomainsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - role := roles.Role{ - ID: testsutil.GenerateUUID(t), - Name: roleName, - EntityID: domainID, - CreatedBy: testsutil.GenerateUUID(t), - CreatedAt: time.Now().UTC(), - } - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - roleID string - svcRes roles.Role - svcErr error - authenticateErr error - response sdk.Role - err errors.SDKError - }{ - { - desc: "view domain role successfully", - token: validToken, - domainID: domainID, - roleID: role.ID, - svcRes: role, - svcErr: nil, - response: convertRole(role), - err: nil, - }, - { - desc: "view domain role with invalid token", - token: invalidToken, - domainID: domainID, - roleID: role.ID, - svcRes: roles.Role{}, - authenticateErr: svcerr.ErrAuthentication, - response: sdk.Role{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "view domain role with empty token", - token: "", - domainID: domainID, - roleID: role.ID, - svcRes: roles.Role{}, - response: sdk.Role{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "view domain role with invalid domain id", - token: validToken, - domainID: testsutil.GenerateUUID(t), - roleID: role.ID, - svcRes: roles.Role{}, - svcErr: svcerr.ErrAuthorization, - response: sdk.Role{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "view domain role with empty domain id", - token: validToken, - domainID: "", - roleID: role.ID, - svcRes: roles.Role{}, - svcErr: nil, - response: sdk.Role{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingDomainID, http.StatusBadRequest), - }, - { - desc: "view domain role with invalid role id", - token: validToken, - domainID: domainID, - roleID: invalid, - svcRes: roles.Role{}, - svcErr: svcerr.ErrAuthorization, - response: sdk.Role{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: tc.domainID + "_" + validID, UserID: validID, DomainID: tc.domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("RetrieveRole", mock.Anything, tc.session, tc.domainID, tc.roleID).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.DomainRole(context.Background(), tc.domainID, tc.roleID, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RetrieveRole", mock.Anything, tc.session, tc.domainID, tc.roleID) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestUpdateDomainRole(t *testing.T) { - ts, csvc, auth := setupDomains() - defer ts.Close() - - conf := sdk.Config{ - DomainsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - roleID := testsutil.GenerateUUID(t) - newRoleName := "newTest" - userID := testsutil.GenerateUUID(t) - createdAt := time.Now().UTC().Add(-time.Hour) - role := roles.Role{ - ID: testsutil.GenerateUUID(t), - Name: newRoleName, - EntityID: domainID, - CreatedBy: userID, - CreatedAt: createdAt, - UpdatedBy: userID, - UpdatedAt: time.Now().UTC(), - } - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - roleID string - newRoleName string - svcRes roles.Role - svcErr error - authenticateErr error - response sdk.Role - err errors.SDKError - }{ - { - desc: "update domain role successfully", - token: validToken, - domainID: domainID, - roleID: roleID, - newRoleName: newRoleName, - svcRes: role, - svcErr: nil, - response: convertRole(role), - err: nil, - }, - { - desc: "update domain role with invalid token", - token: invalidToken, - domainID: domainID, - roleID: roleID, - newRoleName: newRoleName, - svcRes: roles.Role{}, - authenticateErr: svcerr.ErrAuthentication, - response: sdk.Role{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "update domain role with empty token", - token: "", - domainID: domainID, - roleID: roleID, - newRoleName: newRoleName, - svcRes: roles.Role{}, - response: sdk.Role{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "update domain role with invalid domain id", - token: validToken, - domainID: testsutil.GenerateUUID(t), - roleID: roleID, - newRoleName: newRoleName, - svcRes: roles.Role{}, - svcErr: svcerr.ErrAuthorization, - response: sdk.Role{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "update domain role with empty domain id", - token: validToken, - domainID: "", - roleID: roleID, - newRoleName: newRoleName, - svcRes: roles.Role{}, - svcErr: nil, - response: sdk.Role{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingDomainID, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: tc.domainID + "_" + validID, UserID: validID, DomainID: tc.domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("UpdateRoleName", mock.Anything, tc.session, tc.domainID, tc.roleID, tc.newRoleName).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.UpdateDomainRole(context.Background(), tc.domainID, tc.roleID, tc.newRoleName, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "UpdateRoleName", mock.Anything, tc.session, tc.domainID, tc.roleID, tc.newRoleName) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestDeleteDomainRole(t *testing.T) { - ts, csvc, auth := setupDomains() - defer ts.Close() - - conf := sdk.Config{ - DomainsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - roleID := testsutil.GenerateUUID(t) - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - roleID string - svcErr error - authenticateErr error - err errors.SDKError - }{ - { - desc: "delete domain role successfully", - token: validToken, - domainID: domainID, - roleID: roleID, - svcErr: nil, - err: nil, - }, - { - desc: "delete domain role with invalid token", - token: invalidToken, - domainID: domainID, - roleID: roleID, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "delete domain role with empty token", - token: "", - domainID: domainID, - roleID: roleID, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "delete domain role with invalid domain id", - token: validToken, - domainID: testsutil.GenerateUUID(t), - roleID: roleID, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "delete domain role with empty domain id", - token: validToken, - domainID: "", - roleID: roleID, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingDomainID, http.StatusBadRequest), - }, - { - desc: "delete domain role with invalid role id", - token: validToken, - domainID: domainID, - roleID: invalid, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: tc.domainID + "_" + validID, UserID: validID, DomainID: tc.domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("RemoveRole", mock.Anything, tc.session, tc.domainID, tc.roleID).Return(tc.svcErr) - err := mgsdk.DeleteDomainRole(context.Background(), tc.domainID, tc.roleID, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RemoveRole", mock.Anything, tc.session, tc.domainID, tc.roleID) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestAddDomainRoleActions(t *testing.T) { - ts, csvc, auth := setupDomains() - defer ts.Close() - - conf := sdk.Config{ - DomainsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - roleID := testsutil.GenerateUUID(t) - actions := []string{"create", "update"} - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - roleID string - actions []string - svcRes []string - svcErr error - authenticateErr error - response []string - err errors.SDKError - }{ - { - desc: "add domain role actions successfully", - token: validToken, - domainID: domainID, - roleID: roleID, - actions: actions, - svcRes: actions, - svcErr: nil, - response: actions, - err: nil, - }, - { - desc: "add domain role actions with invalid token", - token: invalidToken, - domainID: domainID, - roleID: roleID, - actions: actions, - authenticateErr: svcerr.ErrAuthentication, - response: []string{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "add domain role actions with empty token", - token: "", - domainID: domainID, - roleID: roleID, - actions: actions, - response: []string{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "add domain role actions with invalid domain id", - token: validToken, - domainID: testsutil.GenerateUUID(t), - roleID: roleID, - actions: actions, - svcErr: svcerr.ErrAuthorization, - response: []string{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "add domain role actions with empty domain id", - token: validToken, - domainID: "", - roleID: roleID, - actions: actions, - response: []string{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingDomainID, http.StatusBadRequest), - }, - { - desc: "add domain role actions with invalid role id", - token: validToken, - domainID: domainID, - roleID: invalid, - actions: actions, - svcErr: svcerr.ErrAuthorization, - response: []string{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "add domain role actions with empty actions", - token: validToken, - domainID: domainID, - roleID: roleID, - actions: []string{}, - svcErr: nil, - response: []string{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingPolicyEntityType, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: tc.domainID + "_" + validID, UserID: validID, DomainID: tc.domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("RoleAddActions", mock.Anything, tc.session, tc.domainID, tc.roleID, tc.actions).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.AddDomainRoleActions(context.Background(), tc.domainID, tc.roleID, tc.actions, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RoleAddActions", mock.Anything, tc.session, tc.domainID, tc.roleID, tc.actions) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestListDomainRoleActions(t *testing.T) { - ts, csvc, auth := setupDomains() - defer ts.Close() - - conf := sdk.Config{ - DomainsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - roleID := testsutil.GenerateUUID(t) - actions := []string{"create", "update"} - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - roleID string - svcRes []string - svcErr error - authenticateErr error - response []string - err errors.SDKError - }{ - { - desc: "list domain role actions successfully", - token: validToken, - domainID: domainID, - roleID: roleID, - svcRes: actions, - svcErr: nil, - response: actions, - err: nil, - }, - { - desc: "list domain role actions with invalid token", - token: invalidToken, - domainID: domainID, - roleID: roleID, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "list domain role actions with empty token", - token: "", - domainID: domainID, - roleID: roleID, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "list domain role actions with invalid domain id", - token: validToken, - domainID: testsutil.GenerateUUID(t), - roleID: roleID, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "list domain role actions with empty domain id", - token: validToken, - domainID: "", - roleID: roleID, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingDomainID, http.StatusBadRequest), - }, - { - desc: "list domain role actions with invalid role id", - token: validToken, - domainID: domainID, - roleID: invalid, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "list domain role actions with empty role id", - token: validToken, - domainID: domainID, - roleID: "", - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingRoleID, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: tc.domainID + "_" + validID, UserID: validID, DomainID: tc.domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("RoleListActions", mock.Anything, tc.session, tc.domainID, tc.roleID).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.DomainRoleActions(context.Background(), tc.domainID, tc.roleID, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RoleListActions", mock.Anything, tc.session, tc.domainID, tc.roleID) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestRemoveDomainRoleActions(t *testing.T) { - ts, csvc, auth := setupDomains() - defer ts.Close() - - conf := sdk.Config{ - DomainsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - roleID := testsutil.GenerateUUID(t) - actions := []string{"create", "update"} - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - roleID string - actions []string - svcErr error - authenticateErr error - err errors.SDKError - }{ - { - desc: "remove domain role actions successfully", - token: validToken, - domainID: domainID, - roleID: roleID, - actions: actions, - svcErr: nil, - err: nil, - }, - { - desc: "remove domain role actions with invalid token", - token: invalidToken, - domainID: domainID, - roleID: roleID, - actions: actions, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "remove domain role actions with empty token", - token: "", - domainID: domainID, - roleID: roleID, - actions: actions, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "remove domain role actions with invalid domain id", - token: validToken, - domainID: testsutil.GenerateUUID(t), - roleID: roleID, - actions: actions, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "remove domain role actions with empty domain id", - token: validToken, - domainID: "", - roleID: roleID, - actions: actions, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingDomainID, http.StatusBadRequest), - }, - { - desc: "remove domain role actions with invalid role id", - token: validToken, - domainID: domainID, - roleID: invalid, - actions: actions, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "remove domain role actions with empty actions", - token: validToken, - domainID: domainID, - roleID: roleID, - actions: []string{}, - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingPolicyEntityType, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: tc.domainID + "_" + validID, UserID: validID, DomainID: tc.domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("RoleRemoveActions", mock.Anything, tc.session, tc.domainID, tc.roleID, tc.actions).Return(tc.svcErr) - err := mgsdk.RemoveDomainRoleActions(context.Background(), tc.domainID, tc.roleID, tc.actions, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RoleRemoveActions", mock.Anything, tc.session, tc.domainID, tc.roleID, tc.actions) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestRemoveAllDomainRoleActions(t *testing.T) { - ts, csvc, auth := setupDomains() - defer ts.Close() - - conf := sdk.Config{ - DomainsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - roleID := testsutil.GenerateUUID(t) - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - roleID string - svcErr error - authenticateErr error - err errors.SDKError - }{ - { - desc: "remove all domain role actions successfully", - token: validToken, - domainID: domainID, - roleID: roleID, - svcErr: nil, - err: nil, - }, - { - desc: "remove all domain role actions with invalid token", - token: invalidToken, - domainID: domainID, - roleID: roleID, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "remove all domain role actions with empty token", - token: "", - domainID: domainID, - roleID: roleID, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "remove all domain role actions with invalid domain id", - token: validToken, - domainID: testsutil.GenerateUUID(t), - roleID: roleID, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "remove all domain role actions with empty domain id", - token: validToken, - domainID: "", - roleID: roleID, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingDomainID, http.StatusBadRequest), - }, - { - desc: "remove all domain role actions with invalid role id", - token: validToken, - domainID: domainID, - roleID: invalid, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "remove all domain role actions with empty role id", - token: validToken, - domainID: domainID, - roleID: "", - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingRoleID, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: tc.domainID + "_" + validID, UserID: validID, DomainID: tc.domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("RoleRemoveAllActions", mock.Anything, tc.session, tc.domainID, tc.roleID).Return(tc.svcErr) - err := mgsdk.RemoveAllDomainRoleActions(context.Background(), tc.domainID, tc.roleID, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RoleRemoveAllActions", mock.Anything, tc.session, tc.domainID, tc.roleID) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestAddDomainRoleMembers(t *testing.T) { - ts, csvc, auth := setupDomains() - defer ts.Close() - - conf := sdk.Config{ - DomainsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - roleID := testsutil.GenerateUUID(t) - members := []string{"user1", "user2"} - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - roleID string - members []string - svcRes []string - svcErr error - authenticateErr error - response []string - err errors.SDKError - }{ - { - desc: "add domain role members successfully", - token: validToken, - domainID: domainID, - roleID: roleID, - members: members, - svcRes: members, - svcErr: nil, - response: members, - err: nil, - }, - { - desc: "add domain role members with invalid token", - token: invalidToken, - domainID: domainID, - roleID: roleID, - members: members, - authenticateErr: svcerr.ErrAuthentication, - response: []string{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "add domain role members with empty token", - token: "", - domainID: domainID, - roleID: roleID, - members: members, - response: []string{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "add domain role members with invalid domain id", - token: validToken, - domainID: testsutil.GenerateUUID(t), - roleID: roleID, - members: members, - svcErr: svcerr.ErrAuthorization, - response: []string{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "add domain role members with empty domain id", - token: validToken, - domainID: "", - roleID: roleID, - members: members, - response: []string{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingDomainID, http.StatusBadRequest), - }, - { - desc: "add domain role members with invalid role id", - token: validToken, - domainID: domainID, - roleID: invalid, - members: members, - svcErr: svcerr.ErrAuthorization, - response: []string{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "add domain role members with empty members", - token: validToken, - domainID: domainID, - roleID: roleID, - members: []string{}, - svcErr: nil, - response: []string{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingRoleMembers, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: tc.domainID + "_" + validID, UserID: validID, DomainID: tc.domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("RoleAddMembers", mock.Anything, tc.session, tc.domainID, tc.roleID, tc.members).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.AddDomainRoleMembers(context.Background(), tc.domainID, tc.roleID, tc.members, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RoleAddMembers", mock.Anything, tc.session, tc.domainID, tc.roleID, tc.members) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestListDomainRoleMembers(t *testing.T) { - ts, csvc, auth := setupDomains() - defer ts.Close() - - conf := sdk.Config{ - DomainsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - roleID := testsutil.GenerateUUID(t) - members := []string{"user1", "user2"} - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - roleID string - pageMeta sdk.PageMetadata - svcRes roles.MembersPage - svcErr error - authenticateErr error - response sdk.RoleMembersPage - err errors.SDKError - }{ - { - desc: "list domain role members successfully", - token: validToken, - domainID: domainID, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 5, - }, - roleID: roleID, - svcRes: roles.MembersPage{ - Total: 2, - Offset: 0, - Limit: 5, - Members: members, - }, - svcErr: nil, - response: sdk.RoleMembersPage{ - Total: 2, - Offset: 0, - Limit: 5, - Members: members, - }, - err: nil, - }, - { - desc: "list domain role members with invalid token", - token: invalidToken, - domainID: domainID, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 5, - }, - roleID: roleID, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "list domain role members with empty token", - token: "", - domainID: domainID, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 5, - }, - roleID: roleID, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "list domain role members with invalid domain id", - token: validToken, - domainID: testsutil.GenerateUUID(t), - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 5, - }, - roleID: roleID, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "list domain role members with empty domain id", - token: validToken, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 5, - }, - domainID: "", - roleID: roleID, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingDomainID, http.StatusBadRequest), - }, - { - desc: "list domain role members with invalid role id", - token: validToken, - domainID: domainID, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 5, - }, - roleID: invalid, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "list domain role members with empty role id", - token: validToken, - domainID: domainID, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 5, - }, - roleID: "", - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingRoleID, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: tc.domainID + "_" + validID, UserID: validID, DomainID: tc.domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("RoleListMembers", mock.Anything, tc.session, tc.domainID, tc.roleID, tc.pageMeta.Limit, tc.pageMeta.Offset).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.DomainRoleMembers(context.Background(), tc.domainID, tc.roleID, tc.pageMeta, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RoleListMembers", mock.Anything, tc.session, tc.domainID, tc.roleID, tc.pageMeta.Limit, tc.pageMeta.Offset) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestRemoveDomainRoleMembers(t *testing.T) { - ts, csvc, auth := setupDomains() - defer ts.Close() - - conf := sdk.Config{ - DomainsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - roleID := testsutil.GenerateUUID(t) - members := []string{"user1", "user2"} - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - roleID string - members []string - svcErr error - authenticateErr error - err errors.SDKError - }{ - { - desc: "remove domain role members successfully", - token: validToken, - domainID: domainID, - roleID: roleID, - members: members, - svcErr: nil, - err: nil, - }, - { - desc: "remove domain role members with invalid token", - token: invalidToken, - domainID: domainID, - roleID: roleID, - members: members, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "remove domain role members with empty token", - token: "", - domainID: domainID, - roleID: roleID, - members: members, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "remove domain role members with invalid domain id", - token: validToken, - domainID: testsutil.GenerateUUID(t), - roleID: roleID, - members: members, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "remove domain role members with empty domain id", - token: validToken, - domainID: "", - roleID: roleID, - members: members, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingDomainID, http.StatusBadRequest), - }, - { - desc: "remove domain role members with invalid role id", - token: validToken, - domainID: domainID, - roleID: invalid, - members: members, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "remove domain role members with empty members", - token: validToken, - domainID: domainID, - roleID: roleID, - members: []string{}, - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingRoleMembers, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: tc.domainID + "_" + validID, UserID: validID, DomainID: tc.domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("RoleRemoveMembers", mock.Anything, tc.session, tc.domainID, tc.roleID, tc.members).Return(tc.svcErr) - err := mgsdk.RemoveDomainRoleMembers(context.Background(), tc.domainID, tc.roleID, tc.members, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RoleRemoveMembers", mock.Anything, tc.session, tc.domainID, tc.roleID, tc.members) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestRemoveAllDomainRoleMembers(t *testing.T) { - ts, csvc, auth := setupDomains() - defer ts.Close() - - conf := sdk.Config{ - DomainsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - roleID := testsutil.GenerateUUID(t) - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - roleID string - svcErr error - authenticateErr error - err errors.SDKError - }{ - { - desc: "remove all domain role members successfully", - token: validToken, - domainID: domainID, - roleID: roleID, - svcErr: nil, - err: nil, - }, - { - desc: "remove all domain role members with invalid token", - token: invalidToken, - domainID: domainID, - roleID: roleID, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "remove all domain role members with empty token", - token: "", - domainID: domainID, - roleID: roleID, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "remove all domain role members with invalid domain id", - token: validToken, - domainID: testsutil.GenerateUUID(t), - roleID: roleID, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "remove all domain role members with empty domain id", - token: validToken, - domainID: "", - roleID: roleID, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingDomainID, http.StatusBadRequest), - }, - { - desc: "remove all domain role members with invalid role id", - token: validToken, - domainID: domainID, - roleID: invalid, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "remove all domain role members with empty role id", - token: validToken, - domainID: domainID, - roleID: "", - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingRoleID, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: tc.domainID + "_" + validID, UserID: validID, DomainID: tc.domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("RoleRemoveAllMembers", mock.Anything, tc.session, tc.domainID, tc.roleID).Return(tc.svcErr) - err := mgsdk.RemoveAllDomainRoleMembers(context.Background(), tc.domainID, tc.roleID, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RoleRemoveAllMembers", mock.Anything, tc.session, tc.domainID, tc.roleID) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestListAvailableDomainRoleActions(t *testing.T) { - ts, csvc, auth := setupDomains() - defer ts.Close() - - conf := sdk.Config{ - DomainsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - actions := []string{"create", "update"} - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - svcRes []string - svcErr error - authenticateErr error - response []string - err errors.SDKError - }{ - { - desc: "list available role actions successfully", - token: validToken, - svcRes: actions, - svcErr: nil, - response: actions, - err: nil, - }, - { - desc: "list available role actions with invalid token", - token: invalidToken, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "list available role actions with empty token", - token: "", - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("ListAvailableActions", mock.Anything, tc.session).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.AvailableDomainRoleActions(context.Background(), tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "ListAvailableActions", mock.Anything, tc.session) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func generateTestDomain(t *testing.T) (domains.Domain, sdk.Domain) { - createdAt, err := time.Parse(time.RFC3339, "2024-04-01T00:00:00Z") - assert.Nil(t, err, fmt.Sprintf("Unexpected error parsing time: %s", err)) - ownerID := testsutil.GenerateUUID(t) - ad := domains.Domain{ - ID: testsutil.GenerateUUID(t), - Name: "test-domain", - Metadata: domains.Metadata(validMetadata), - Tags: []string{"tag1", "tag2"}, - Route: "test-route", - Status: domains.EnabledStatus, - CreatedBy: ownerID, - CreatedAt: createdAt, - UpdatedBy: ownerID, - UpdatedAt: createdAt, - Roles: validRoles, - } - - sd := sdk.Domain{ - ID: ad.ID, - Name: ad.Name, - Metadata: validMetadata, - Tags: ad.Tags, - Route: ad.Route, - Status: ad.Status.String(), - CreatedBy: ad.CreatedBy, - CreatedAt: ad.CreatedAt, - UpdatedBy: ad.UpdatedBy, - UpdatedAt: ad.UpdatedAt, - Roles: ad.Roles, - } - return ad, sd -} diff --git a/pkg/sdk/groups.go b/pkg/sdk/groups.go index 325e3af91..99a4f063f 100644 --- a/pkg/sdk/groups.go +++ b/pkg/sdk/groups.go @@ -12,7 +12,6 @@ import ( apiutil "github.com/absmach/magistrala/api/http/util" "github.com/absmach/magistrala/pkg/errors" - "github.com/absmach/magistrala/pkg/roles" ) const ( @@ -27,29 +26,28 @@ const ( // Path in a tree consisting of group IDs // Paths are unique per owner. type Group struct { - ID string `json:"id,omitempty"` - DomainID string `json:"domain_id,omitempty"` - ParentID string `json:"parent_id,omitempty"` - Name string `json:"name,omitempty"` - Description string `json:"description,omitempty"` - Tags []string `json:"tags,omitempty"` - Metadata Metadata `json:"metadata,omitempty"` - Level int `json:"level,omitempty"` - Path string `json:"path,omitempty"` - Children []*Group `json:"children,omitempty"` - CreatedAt time.Time `json:"created_at,omitempty"` - UpdatedAt time.Time `json:"updated_at,omitempty"` - UpdatedBy string `json:"updated_by,omitempty"` - Status string `json:"status,omitempty"` - RoleID string `json:"role_id,omitempty"` - RoleName string `json:"role_name,omitempty"` - Actions []string `json:"actions,omitempty"` - AccessType string `json:"access_type,omitempty"` - AccessProviderId string `json:"access_provider_id,omitempty"` - AccessProviderRoleId string `json:"access_provider_role_id,omitempty"` - AccessProviderRoleName string `json:"access_provider_role_name,omitempty"` - AccessProviderRoleActions []string `json:"access_provider_role_actions,omitempty"` - Roles []roles.MemberRoleActions `json:"roles,omitempty"` + ID string `json:"id,omitempty"` + DomainID string `json:"domain_id,omitempty"` + ParentID string `json:"parent_id,omitempty"` + Name string `json:"name,omitempty"` + Description string `json:"description,omitempty"` + Tags []string `json:"tags,omitempty"` + Metadata Metadata `json:"metadata,omitempty"` + Level int `json:"level,omitempty"` + Path string `json:"path,omitempty"` + Children []*Group `json:"children,omitempty"` + CreatedAt time.Time `json:"created_at,omitempty"` + UpdatedAt time.Time `json:"updated_at,omitempty"` + UpdatedBy string `json:"updated_by,omitempty"` + Status string `json:"status,omitempty"` + RoleID string `json:"-"` + RoleName string `json:"-"` + Actions []string `json:"-"` + AccessType string `json:"-"` + AccessProviderId string `json:"-"` + AccessProviderRoleId string `json:"-"` + AccessProviderRoleName string `json:"-"` + AccessProviderRoleActions []string `json:"-"` } func (sdk mgSDK) CreateGroup(ctx context.Context, g Group, domainID, token string) (Group, errors.SDKError) { diff --git a/pkg/sdk/groups_test.go b/pkg/sdk/groups_test.go deleted file mode 100644 index 2be27801f..000000000 --- a/pkg/sdk/groups_test.go +++ /dev/null @@ -1,3849 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package sdk_test - -import ( - "context" - "fmt" - "net/http" - "net/http/httptest" - "strings" - "testing" - "time" - - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/groups" - httpapi "github.com/absmach/magistrala/groups/api/http" - "github.com/absmach/magistrala/groups/mocks" - "github.com/absmach/magistrala/internal/nullable" - "github.com/absmach/magistrala/internal/testsutil" - mglog "github.com/absmach/magistrala/logger" - smqauthn "github.com/absmach/magistrala/pkg/authn" - authnmocks "github.com/absmach/magistrala/pkg/authn/mocks" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - oauth2mocks "github.com/absmach/magistrala/pkg/oauth2/mocks" - "github.com/absmach/magistrala/pkg/roles" - sdk "github.com/absmach/magistrala/pkg/sdk" - "github.com/absmach/magistrala/pkg/uuid" - "github.com/go-chi/chi/v5" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" -) - -var ( - sdkGroup = generateTestGroup(&testing.T{}) - group = convertGroup(sdkGroup) - updatedName = "updated_name" - updatedDescription = "updated_description" -) - -func setupGroups() (*httptest.Server, *mocks.Service, *authnmocks.Authentication) { - svc := new(mocks.Service) - - logger := mglog.NewMock() - mux := chi.NewRouter() - idp := uuid.NewMock() - provider := new(oauth2mocks.Provider) - provider.On("Name").Return(roleName) - authn := new(authnmocks.Authentication) - am := smqauthn.NewAuthNMiddleware(authn, smqauthn.WithAllowUnverifiedUser(true)) - httpapi.MakeHandler(svc, am, mux, logger, "", idp) - - return httptest.NewServer(mux), svc, authn -} - -func TestCreateGroup(t *testing.T) { - ts, gsvc, auth := setupGroups() - defer ts.Close() - - conf := sdk.Config{ - GroupsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - createGroupReq := sdk.Group{ - Name: gName, - Description: description, - Metadata: validMetadata, - } - pGroup := group - pGroup.Parent = testsutil.GenerateUUID(t) - psdkGroup := sdkGroup - psdkGroup.ParentID = pGroup.Parent - - uGroup := group - uGroup.Metadata = groups.Metadata{ - "key": make(chan int), - } - - desc := nullable.New(description) - - cases := []struct { - desc string - domainID string - token string - session smqauthn.Session - groupReq sdk.Group - svcReq groups.Group - svcRes groups.Group - svcErr error - authenticateErr error - response sdk.Group - err errors.SDKError - }{ - { - desc: "create group successfully", - domainID: domainID, - token: validToken, - groupReq: createGroupReq, - svcReq: groups.Group{ - Name: gName, - Description: desc, - Metadata: groups.Metadata{"role": "client"}, - }, - svcRes: group, - svcErr: nil, - response: sdkGroup, - err: nil, - }, - { - desc: "create group with existing name", - domainID: domainID, - token: validToken, - groupReq: createGroupReq, - svcReq: groups.Group{ - Name: gName, - Description: desc, - Metadata: groups.Metadata{"role": "client"}, - }, - svcRes: group, - svcErr: nil, - response: sdkGroup, - err: nil, - }, - { - desc: "create group with parent", - domainID: domainID, - token: validToken, - groupReq: sdk.Group{ - Name: gName, - Description: description, - Metadata: validMetadata, - ParentID: pGroup.Parent, - }, - svcReq: groups.Group{ - Name: gName, - Description: desc, - Metadata: groups.Metadata{"role": "client"}, - Parent: pGroup.Parent, - }, - svcRes: pGroup, - svcErr: nil, - response: psdkGroup, - err: nil, - }, - { - desc: "create group with invalid parent", - domainID: domainID, - token: validToken, - groupReq: sdk.Group{ - Name: gName, - Description: description, - Metadata: validMetadata, - ParentID: wrongID, - }, - svcReq: groups.Group{ - Name: gName, - Description: desc, - Metadata: groups.Metadata{"role": "client"}, - Parent: wrongID, - }, - svcRes: groups.Group{}, - svcErr: svcerr.ErrAuthorization, - response: sdk.Group{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "create group with invalid token", - domainID: domainID, - token: invalidToken, - groupReq: sdk.Group{ - Name: gName, - Description: description, - Metadata: validMetadata, - }, - svcReq: groups.Group{ - Name: gName, - Description: desc, - Metadata: groups.Metadata{"role": "client"}, - }, - svcRes: groups.Group{}, - authenticateErr: svcerr.ErrAuthentication, - response: sdk.Group{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "create group with empty token", - domainID: domainID, - token: "", - groupReq: sdk.Group{ - Name: gName, - Description: description, - Metadata: validMetadata, - }, - svcReq: groups.Group{}, - svcRes: groups.Group{}, - svcErr: nil, - response: sdk.Group{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "create group with missing name", - domainID: domainID, - token: validToken, - groupReq: sdk.Group{ - Description: description, - Metadata: validMetadata, - }, - svcReq: groups.Group{}, - svcRes: groups.Group{}, - svcErr: nil, - response: sdk.Group{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrNameSize, http.StatusBadRequest), - }, - { - desc: "create group with name that is too long", - domainID: domainID, - token: validToken, - groupReq: sdk.Group{ - Name: strings.Repeat("a", 1025), - Description: description, - Metadata: validMetadata, - }, - svcReq: groups.Group{}, - svcRes: groups.Group{}, - svcErr: nil, - response: sdk.Group{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrNameSize, http.StatusBadRequest), - }, - { - desc: "create group with request that cannot be marshalled", - domainID: domainID, - token: validToken, - groupReq: sdk.Group{ - Name: gName, - Description: description, - Metadata: sdk.Metadata{ - "key": make(chan int), - }, - }, - svcReq: groups.Group{}, - svcRes: groups.Group{}, - svcErr: nil, - response: sdk.Group{}, - err: errors.NewSDKError(fmt.Errorf("json: unsupported type: chan int")), - }, - { - desc: "create group with service response that cannot be unmarshalled", - domainID: domainID, - token: validToken, - groupReq: sdk.Group{ - Name: gName, - Description: description, - Metadata: validMetadata, - }, - svcReq: groups.Group{ - Name: gName, - Description: desc, - Metadata: groups.Metadata{"role": "client"}, - }, - svcRes: uGroup, - svcErr: nil, - response: sdk.Group{}, - err: errors.NewSDKError(fmt.Errorf("unexpected end of JSON input")), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := gsvc.On("CreateGroup", mock.Anything, tc.session, tc.svcReq).Return(tc.svcRes, []roles.RoleProvision{}, tc.svcErr) - resp, err := mgsdk.CreateGroup(context.Background(), tc.groupReq, tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "CreateGroup", mock.Anything, tc.session, tc.svcReq) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestListGroups(t *testing.T) { - ts, gsvc, auth := setupGroups() - defer ts.Close() - - var grps []sdk.Group - conf := sdk.Config{ - GroupsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - for i := 10; i < 100; i++ { - gr := sdk.Group{ - ID: generateUUID(t), - Name: fmt.Sprintf("group_%d", i), - Metadata: sdk.Metadata{"name": fmt.Sprintf("user_%d", i)}, - Status: groups.EnabledStatus.String(), - } - grps = append(grps, gr) - } - - cases := []struct { - desc string - token string - domainID string - session smqauthn.Session - pageMeta sdk.PageMetadata - svcReq groups.PageMeta - svcRes groups.Page - svcErr error - authenticateErr error - response sdk.GroupsPage - err errors.SDKError - }{ - { - desc: "list groups successfully", - domainID: domainID, - token: validToken, - pageMeta: sdk.PageMetadata{ - Offset: offset, - Limit: 100, - Order: "created_at", - Direction: "asc", - }, - svcReq: groups.PageMeta{ - Offset: offset, - Limit: 100, - Order: "created_at", - Dir: "asc", - Actions: []string{}, - }, - svcRes: groups.Page{ - PageMeta: groups.PageMeta{ - Total: uint64(len(grps)), - }, - Groups: convertGroups(grps), - }, - response: sdk.GroupsPage{ - PageRes: sdk.PageRes{ - Total: uint64(len(grps)), - }, - Groups: grps, - }, - err: nil, - }, - { - desc: "list groups with invalid token", - token: invalidToken, - domainID: domainID, - pageMeta: sdk.PageMetadata{ - Offset: offset, - Limit: 100, - }, - svcReq: groups.PageMeta{ - Offset: offset, - Limit: 100, - Actions: []string{}, - }, - svcRes: groups.Page{}, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "list groups with empty token", - domainID: domainID, - token: "", - pageMeta: sdk.PageMetadata{ - Offset: offset, - Limit: 100, - }, - svcReq: groups.PageMeta{}, - svcRes: groups.Page{}, - svcErr: nil, - response: sdk.GroupsPage{}, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "list groups with zero limit", - domainID: domainID, - token: validToken, - pageMeta: sdk.PageMetadata{ - Offset: offset, - Limit: 0, - Order: "created_at", - Direction: "asc", - }, - svcReq: groups.PageMeta{ - Offset: offset, - Limit: 10, - Order: "created_at", - Dir: "asc", - Actions: []string{}, - }, - svcRes: groups.Page{ - PageMeta: groups.PageMeta{ - Total: uint64(len(grps[0:10])), - }, - Groups: convertGroups(grps[0:10]), - }, - svcErr: nil, - response: sdk.GroupsPage{ - PageRes: sdk.PageRes{ - Total: uint64(len(grps[0:10])), - }, - Groups: grps[0:10], - }, - err: nil, - }, - { - desc: "list groups with limit greater than max", - domainID: domainID, - token: validToken, - pageMeta: sdk.PageMetadata{ - Offset: offset, - Limit: 110, - }, - svcReq: groups.PageMeta{}, - svcRes: groups.Page{}, - svcErr: nil, - response: sdk.GroupsPage{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrLimitSize, http.StatusBadRequest), - }, - { - desc: "list groups with given name", - domainID: domainID, - token: validToken, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - Order: "created_at", - Direction: "asc", - Metadata: sdk.Metadata{ - "name": "user_89", - }, - }, - svcReq: groups.PageMeta{ - Offset: 0, - Limit: 10, - Order: "created_at", - Dir: "asc", - Metadata: groups.Metadata{ - "name": "user_89", - }, - Actions: []string{}, - }, - svcRes: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 1, - }, - Groups: convertGroups([]sdk.Group{grps[89]}), - }, - svcErr: nil, - response: sdk.GroupsPage{ - PageRes: sdk.PageRes{ - Total: 1, - }, - Groups: []sdk.Group{grps[89]}, - }, - err: nil, - }, - { - desc: "list groups with invalid page metadata", - domainID: domainID, - token: validToken, - pageMeta: sdk.PageMetadata{ - Offset: offset, - Limit: limit, - Metadata: sdk.Metadata{ - "key": make(chan int), - }, - }, - svcReq: groups.PageMeta{}, - svcRes: groups.Page{}, - svcErr: nil, - response: sdk.GroupsPage{}, - err: errors.NewSDKError(fmt.Errorf("json: unsupported type: chan int")), - }, - { - desc: "list groups with service response that cannot be unmarshalled", - domainID: domainID, - token: validToken, - pageMeta: sdk.PageMetadata{ - Offset: offset, - Limit: limit, - Order: "created_at", - Direction: "asc", - }, - svcReq: groups.PageMeta{ - Offset: offset, - Limit: limit, - Order: "created_at", - Dir: "asc", - Actions: []string{}, - }, - svcRes: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 1, - }, - Groups: []groups.Group{{ - ID: generateUUID(t), - Name: "group_1", - Metadata: groups.Metadata{ - "key": make(chan int), - }, - }}, - }, - svcErr: nil, - response: sdk.GroupsPage{}, - err: errors.NewSDKError(fmt.Errorf("unexpected end of JSON input")), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := gsvc.On("ListGroups", mock.Anything, tc.session, tc.svcReq).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.Groups(context.Background(), tc.pageMeta, tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "ListGroups", mock.Anything, tc.session, tc.svcReq) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestViewGroup(t *testing.T) { - ts, gsvc, auth := setupGroups() - defer ts.Close() - - conf := sdk.Config{ - GroupsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - confRoles := sdk.Config{ - GroupsURL: ts.URL, - Roles: true, - } - mgsdkRoles := sdk.NewSDK(confRoles) - - cases := []struct { - desc string - domainID string - token string - session smqauthn.Session - withRoles bool - groupID string - svcRes groups.Group - svcErr error - authenticateErr error - response sdk.Group - err errors.SDKError - }{ - { - desc: "view group successfully", - domainID: domainID, - token: validToken, - withRoles: false, - groupID: group.ID, - svcRes: group, - svcErr: nil, - response: sdkGroup, - err: nil, - }, - { - desc: "view group successfully with roles", - domainID: domainID, - token: validToken, - withRoles: true, - groupID: group.ID, - svcRes: group, - svcErr: nil, - response: sdkGroup, - err: nil, - }, - { - desc: "view group with invalid token", - domainID: domainID, - token: invalidToken, - withRoles: false, - groupID: group.ID, - svcRes: groups.Group{}, - authenticateErr: svcerr.ErrAuthentication, - response: sdk.Group{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "view group with empty token", - domainID: domainID, - token: "", - withRoles: false, - groupID: group.ID, - svcRes: groups.Group{}, - svcErr: nil, - response: sdk.Group{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "view group with invalid group id", - domainID: domainID, - token: validToken, - withRoles: false, - groupID: wrongID, - svcRes: groups.Group{}, - svcErr: svcerr.ErrNotFound, - response: sdk.Group{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrNotFound, http.StatusNotFound), - }, - { - desc: "view group with service response that cannot be unmarshalled", - domainID: domainID, - token: validToken, - withRoles: false, - groupID: group.ID, - svcRes: groups.Group{ - ID: group.ID, - Name: "group_1", - Metadata: groups.Metadata{ - "key": make(chan int), - }, - }, - svcErr: nil, - response: sdk.Group{}, - err: errors.NewSDKError(fmt.Errorf("unexpected end of JSON input")), - }, - { - desc: "view group with empty id", - domainID: domainID, - token: validToken, - withRoles: false, - groupID: "", - svcRes: groups.Group{}, - svcErr: nil, - response: sdk.Group{}, - err: errors.NewSDKError(apiutil.ErrMissingID), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := gsvc.On("ViewGroup", mock.Anything, tc.session, tc.groupID, tc.withRoles).Return(tc.svcRes, tc.svcErr) - - var resp sdk.Group - var err error - - switch tc.withRoles { - case true: - resp, err = mgsdkRoles.Group(context.Background(), tc.groupID, tc.domainID, tc.token) - default: - resp, err = mgsdk.Group(context.Background(), tc.groupID, tc.domainID, tc.token) - } - - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.withRoles { - assert.Equal(t, resp.Roles, validRoles, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, validRoles, resp.Roles)) - } - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "ViewGroup", mock.Anything, tc.session, tc.groupID, tc.withRoles) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestUpdateGroup(t *testing.T) { - ts, gsvc, auth := setupGroups() - defer ts.Close() - - upGroup := sdkGroup - upGroup.Name = updatedName - upGroup.Description = updatedDescription - upGroup.Metadata = sdk.Metadata{"key": "value"} - - conf := sdk.Config{ - GroupsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - group.ID = generateUUID(t) - - updatedDesc := nullable.New(updatedDescription) - - cases := []struct { - desc string - domainID string - token string - session smqauthn.Session - groupReq sdk.Group - svcReq groups.Group - svcRes groups.Group - svcErr error - authenticateErr error - response sdk.Group - err errors.SDKError - }{ - { - desc: "update group successfully", - domainID: domainID, - token: validToken, - groupReq: sdk.Group{ - ID: group.ID, - Name: updatedName, - Description: updatedDescription, - Metadata: sdk.Metadata{"key": "value"}, - }, - svcReq: groups.Group{ - ID: group.ID, - Name: updatedName, - Description: updatedDesc, - Metadata: groups.Metadata{"key": "value"}, - }, - svcRes: convertGroup(upGroup), - svcErr: nil, - response: upGroup, - err: nil, - }, - { - desc: "update group name with invalid group id", - domainID: domainID, - token: validToken, - groupReq: sdk.Group{ - ID: wrongID, - Name: updatedName, - Description: updatedDescription, - Metadata: sdk.Metadata{"key": "value"}, - }, - svcReq: groups.Group{ - ID: wrongID, - Name: updatedName, - Description: updatedDesc, - Metadata: groups.Metadata{"key": "value"}, - }, - svcRes: groups.Group{}, - svcErr: svcerr.ErrNotFound, - response: sdk.Group{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrNotFound, http.StatusNotFound), - }, - { - desc: "update group name with invalid token", - domainID: domainID, - token: invalidToken, - groupReq: sdk.Group{ - ID: group.ID, - Name: updatedName, - Description: updatedDescription, - Metadata: sdk.Metadata{"key": "value"}, - }, - svcReq: groups.Group{ - ID: group.ID, - Name: updatedName, - Description: updatedDesc, - Metadata: groups.Metadata{"key": "value"}, - }, - svcRes: groups.Group{}, - authenticateErr: svcerr.ErrAuthentication, - response: sdk.Group{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "update group name with empty token", - domainID: domainID, - token: "", - groupReq: sdk.Group{ - ID: group.ID, - Name: updatedName, - Description: updatedDescription, - Metadata: sdk.Metadata{"key": "value"}, - }, - svcReq: groups.Group{}, - svcRes: groups.Group{}, - svcErr: nil, - response: sdk.Group{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "update group with empty id", - domainID: domainID, - token: validToken, - groupReq: sdk.Group{ - ID: "", - Name: updatedName, - Description: updatedDescription, - Metadata: sdk.Metadata{"key": "value"}, - }, - svcReq: groups.Group{}, - svcRes: groups.Group{}, - svcErr: nil, - response: sdk.Group{}, - err: errors.NewSDKError(apiutil.ErrMissingID), - }, - { - desc: "update group with request that can't be marshalled", - domainID: domainID, - token: validToken, - groupReq: sdk.Group{ - ID: group.ID, - Name: updatedName, - Description: updatedDescription, - Metadata: sdk.Metadata{"key": make(chan int)}, - }, - svcReq: groups.Group{}, - svcRes: groups.Group{}, - svcErr: nil, - response: sdk.Group{}, - err: errors.NewSDKError(fmt.Errorf("json: unsupported type: chan int")), - }, - { - desc: "update group with service response that cannot be unmarshalled", - domainID: domainID, - token: validToken, - groupReq: sdk.Group{ - ID: group.ID, - Name: updatedName, - Description: updatedDescription, - Metadata: sdk.Metadata{"key": "value"}, - }, - svcReq: groups.Group{ - ID: group.ID, - Name: updatedName, - Description: updatedDesc, - Metadata: groups.Metadata{"key": "value"}, - }, - svcRes: groups.Group{ - ID: group.ID, - Name: updatedName, - Metadata: groups.Metadata{ - "key": make(chan int), - }, - }, - svcErr: nil, - response: sdk.Group{}, - err: errors.NewSDKError(fmt.Errorf("unexpected end of JSON input")), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := gsvc.On("UpdateGroup", mock.Anything, tc.session, tc.svcReq).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.UpdateGroup(context.Background(), tc.groupReq, tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "UpdateGroup", mock.Anything, tc.session, tc.svcReq) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestUpdateGroupTags(t *testing.T) { - ts, tsvc, auth := setupGroups() - defer ts.Close() - - sdkGroup := generateTestGroup(t) - updatedGroup := sdkGroup - updatedGroup.Tags = []string{"newTag1", "newTag2"} - updateGroupReq := sdk.Group{ - ID: sdkGroup.ID, - Tags: updatedGroup.Tags, - } - - conf := sdk.Config{ - GroupsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - cases := []struct { - desc string - domainID string - token string - session smqauthn.Session - updateGroupReq sdk.Group - svcReq groups.Group - svcRes groups.Group - svcErr error - authenticateErr error - response sdk.Group - err errors.SDKError - }{ - { - desc: "update group tags successfully", - domainID: domainID, - token: validToken, - updateGroupReq: updateGroupReq, - svcReq: convertGroup(updateGroupReq), - svcRes: convertGroup(updatedGroup), - svcErr: nil, - response: updatedGroup, - err: nil, - }, - { - desc: "update group tags with an invalid token", - domainID: domainID, - token: invalidToken, - updateGroupReq: updateGroupReq, - svcReq: convertGroup(updateGroupReq), - svcRes: groups.Group{}, - authenticateErr: svcerr.ErrAuthorization, - response: sdk.Group{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusUnauthorized), - }, - { - desc: "update group tags with empty token", - domainID: domainID, - token: "", - updateGroupReq: updateGroupReq, - svcReq: convertGroup(updateGroupReq), - svcRes: groups.Group{}, - svcErr: nil, - response: sdk.Group{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "update group tags with an invalid group id", - domainID: domainID, - token: validToken, - updateGroupReq: sdk.Group{ - ID: wrongID, - Tags: updatedGroup.Tags, - }, - svcReq: convertGroup(sdk.Group{ - ID: wrongID, - Tags: updatedGroup.Tags, - }), - svcRes: groups.Group{}, - svcErr: svcerr.ErrUpdateEntity, - response: sdk.Group{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrUpdateEntity, http.StatusUnprocessableEntity), - }, - { - desc: "update group tags with empty group id", - domainID: domainID, - token: validToken, - updateGroupReq: sdk.Group{ - ID: "", - Tags: updatedGroup.Tags, - }, - svcReq: convertGroup(sdk.Group{ - ID: "", - Tags: updatedGroup.Tags, - }), - svcRes: groups.Group{}, - svcErr: nil, - response: sdk.Group{}, - err: errors.NewSDKError(apiutil.ErrMissingID), - }, - { - desc: "update group tags with a request that can't be marshalled", - domainID: domainID, - token: validToken, - updateGroupReq: sdk.Group{ - ID: "test", - Metadata: map[string]any{ - "test": make(chan int), - }, - }, - svcReq: groups.Group{}, - svcRes: groups.Group{}, - svcErr: nil, - response: sdk.Group{}, - err: errors.NewSDKError(fmt.Errorf("json: unsupported type: chan int")), - }, - { - desc: "update group tags with a response that can't be unmarshalled", - domainID: domainID, - token: validToken, - updateGroupReq: updateGroupReq, - svcReq: convertGroup(updateGroupReq), - svcRes: groups.Group{ - Name: updatedGroup.Name, - Tags: updatedGroup.Tags, - Metadata: groups.Metadata{ - "test": make(chan int), - }, - }, - svcErr: nil, - response: sdk.Group{}, - err: errors.NewSDKError(fmt.Errorf("unexpected end of JSON input")), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, mock.Anything).Return(tc.session, tc.authenticateErr) - svcCall := tsvc.On("UpdateGroupTags", mock.Anything, tc.session, tc.svcReq).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.UpdateGroupTags(context.Background(), tc.updateGroupReq, tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "UpdateGroupTags", mock.Anything, tc.session, tc.svcReq) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestEnableGroup(t *testing.T) { - ts, gsvc, auth := setupGroups() - defer ts.Close() - - conf := sdk.Config{ - GroupsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - enGroup := sdkGroup - enGroup.Status = groups.EnabledStatus.String() - - cases := []struct { - desc string - domainID string - token string - session smqauthn.Session - groupID string - svcRes groups.Group - svcErr error - authenticateErr error - response sdk.Group - err errors.SDKError - }{ - { - desc: "enable group successfully", - domainID: domainID, - token: validToken, - groupID: group.ID, - svcRes: convertGroup(enGroup), - svcErr: nil, - response: enGroup, - err: nil, - }, - { - desc: "enable group with invalid group id", - domainID: domainID, - token: validToken, - groupID: wrongID, - svcRes: groups.Group{}, - svcErr: svcerr.ErrNotFound, - response: sdk.Group{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrNotFound, http.StatusNotFound), - }, - { - desc: "enable group with invalid token", - domainID: domainID, - token: invalidToken, - groupID: group.ID, - svcRes: groups.Group{}, - authenticateErr: svcerr.ErrAuthentication, - response: sdk.Group{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "enable group with empty token", - domainID: domainID, - token: "", - groupID: group.ID, - svcRes: groups.Group{}, - svcErr: nil, - response: sdk.Group{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "enable group with empty id", - domainID: domainID, - token: validToken, - groupID: "", - svcRes: groups.Group{}, - svcErr: nil, - response: sdk.Group{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "enable group with service response that cannot be unmarshalled", - domainID: domainID, - token: validToken, - groupID: group.ID, - svcRes: groups.Group{ - ID: group.ID, - Name: "group_1", - Metadata: groups.Metadata{ - "key": make(chan int), - }, - }, - svcErr: nil, - response: sdk.Group{}, - err: errors.NewSDKError(fmt.Errorf("unexpected end of JSON input")), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := gsvc.On("EnableGroup", mock.Anything, tc.session, tc.groupID).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.EnableGroup(context.Background(), tc.groupID, tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "EnableGroup", mock.Anything, tc.session, tc.groupID) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestDisableGroup(t *testing.T) { - ts, gsvc, auth := setupGroups() - defer ts.Close() - - conf := sdk.Config{ - GroupsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - disGroup := sdkGroup - disGroup.Status = groups.DisabledStatus.String() - - cases := []struct { - desc string - domainID string - token string - session smqauthn.Session - groupID string - svcRes groups.Group - svcErr error - authenticateErr error - response sdk.Group - err errors.SDKError - }{ - { - desc: "disable group successfully", - domainID: domainID, - token: validToken, - groupID: group.ID, - svcRes: convertGroup(disGroup), - svcErr: nil, - response: disGroup, - err: nil, - }, - { - desc: "disable group with invalid group id", - domainID: domainID, - token: validToken, - groupID: wrongID, - svcRes: groups.Group{}, - svcErr: svcerr.ErrNotFound, - response: sdk.Group{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrNotFound, http.StatusNotFound), - }, - { - desc: "disable group with invalid token", - domainID: domainID, - token: invalidToken, - groupID: group.ID, - svcRes: groups.Group{}, - authenticateErr: svcerr.ErrAuthentication, - response: sdk.Group{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "disable group with empty token", - domainID: domainID, - token: "", - groupID: group.ID, - svcRes: groups.Group{}, - svcErr: nil, - response: sdk.Group{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "disable group with empty id", - domainID: domainID, - token: validToken, - groupID: "", - svcRes: groups.Group{}, - svcErr: nil, - response: sdk.Group{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "disable group with service response that cannot be unmarshalled", - domainID: domainID, - token: validToken, - groupID: group.ID, - svcRes: groups.Group{ - ID: group.ID, - Name: "group_1", - Metadata: groups.Metadata{ - "key": make(chan int), - }, - }, - svcErr: nil, - response: sdk.Group{}, - err: errors.NewSDKError(fmt.Errorf("unexpected end of JSON input")), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := gsvc.On("DisableGroup", mock.Anything, tc.session, tc.groupID).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.DisableGroup(context.Background(), tc.groupID, tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "DisableGroup", mock.Anything, tc.session, tc.groupID) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestDeleteGroup(t *testing.T) { - ts, gsvc, auth := setupGroups() - defer ts.Close() - - conf := sdk.Config{ - GroupsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - cases := []struct { - desc string - domainID string - token string - session smqauthn.Session - groupID string - svcErr error - authenticateErr error - err errors.SDKError - }{ - { - desc: "delete group successfully", - domainID: domainID, - token: validToken, - groupID: group.ID, - svcErr: nil, - err: nil, - }, - { - desc: "delete group with invalid group id", - domainID: domainID, - token: validToken, - groupID: wrongID, - svcErr: svcerr.ErrRemoveEntity, - err: errors.NewSDKErrorWithStatus(svcerr.ErrRemoveEntity, http.StatusUnprocessableEntity), - }, - { - desc: "delete group with invalid token", - domainID: domainID, - token: invalidToken, - groupID: group.ID, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "delete group with empty token", - domainID: domainID, - token: "", - groupID: group.ID, - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "delete group with empty id", - domainID: domainID, - token: validToken, - groupID: "", - svcErr: nil, - err: errors.NewSDKError(apiutil.ErrMissingID), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := gsvc.On("DeleteGroup", mock.Anything, tc.session, tc.groupID).Return(tc.svcErr) - err := mgsdk.DeleteGroup(context.Background(), tc.groupID, tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "DeleteGroup", mock.Anything, tc.session, tc.groupID) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestSetGroupParent(t *testing.T) { - ts, csvc, auth := setupGroups() - defer ts.Close() - - conf := sdk.Config{ - GroupsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - groupID := testsutil.GenerateUUID(t) - parentID := testsutil.GenerateUUID(t) - - cases := []struct { - desc string - domainID string - token string - session smqauthn.Session - groupID string - parentID string - svcErr error - authenticateErr error - err errors.SDKError - }{ - { - desc: "set group parent successfully", - domainID: domainID, - token: validToken, - groupID: groupID, - parentID: parentID, - svcErr: nil, - err: nil, - }, - { - desc: "set group parent with invalid token", - domainID: domainID, - token: invalidToken, - groupID: groupID, - parentID: parentID, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "set group parent with empty token", - domainID: domainID, - token: "", - groupID: groupID, - parentID: parentID, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "set group parent with invalid group id", - domainID: domainID, - token: validToken, - groupID: wrongID, - parentID: parentID, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "set group parent with empty group id", - domainID: domainID, - token: validToken, - groupID: "", - parentID: parentID, - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "set group parent with empty parent id", - domainID: domainID, - token: validToken, - groupID: groupID, - parentID: "", - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrInvalidIDFormat, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("AddParentGroup", mock.Anything, tc.session, tc.groupID, tc.parentID).Return(tc.svcErr) - err := mgsdk.SetGroupParent(context.Background(), tc.groupID, tc.domainID, tc.parentID, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "AddParentGroup", mock.Anything, tc.session, tc.groupID, tc.parentID) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestRemoveGroupParent(t *testing.T) { - ts, csvc, auth := setupGroups() - defer ts.Close() - - conf := sdk.Config{ - GroupsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - groupID := testsutil.GenerateUUID(t) - parentID := testsutil.GenerateUUID(t) - - cases := []struct { - desc string - domainID string - token string - session smqauthn.Session - groupID string - parentID string - svcErr error - authenticateErr error - err errors.SDKError - }{ - { - desc: "remove group parent successfully", - domainID: domainID, - token: validToken, - groupID: groupID, - parentID: parentID, - svcErr: nil, - err: nil, - }, - { - desc: "remove group parent with invalid token", - domainID: domainID, - token: invalidToken, - groupID: groupID, - parentID: parentID, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "remove group parent with empty token", - domainID: domainID, - token: "", - groupID: groupID, - parentID: parentID, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "remove group parent with invalid group id", - domainID: domainID, - token: validToken, - groupID: wrongID, - parentID: parentID, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "remove group parent with empty group id", - domainID: domainID, - token: validToken, - groupID: "", - parentID: parentID, - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("RemoveParentGroup", mock.Anything, tc.session, tc.groupID).Return(tc.svcErr) - err := mgsdk.RemoveGroupParent(context.Background(), tc.groupID, tc.domainID, tc.parentID, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RemoveParentGroup", mock.Anything, tc.session, tc.groupID) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestAddChildrenGroups(t *testing.T) { - ts, csvc, auth := setupGroups() - defer ts.Close() - - conf := sdk.Config{ - GroupsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - groupID := testsutil.GenerateUUID(t) - childID := testsutil.GenerateUUID(t) - - cases := []struct { - desc string - domainID string - token string - session smqauthn.Session - groupID string - childrenIDs []string - svcErr error - authenticateErr error - err errors.SDKError - }{ - { - desc: "add children group successfully", - domainID: domainID, - token: validToken, - groupID: groupID, - childrenIDs: []string{childID}, - svcErr: nil, - err: nil, - }, - { - desc: "add children group with invalid token", - domainID: domainID, - token: invalidToken, - groupID: groupID, - childrenIDs: []string{childID}, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "add children group with empty token", - domainID: domainID, - token: "", - groupID: groupID, - childrenIDs: []string{childID}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "add children group with invalid group id", - domainID: domainID, - token: validToken, - groupID: wrongID, - childrenIDs: []string{childID}, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "add children group with empty group id", - domainID: domainID, - token: validToken, - groupID: "", - childrenIDs: []string{childID}, - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "add children group with empty children ids", - domainID: domainID, - token: validToken, - groupID: groupID, - childrenIDs: []string{}, - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingChildrenGroupIDs, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("AddChildrenGroups", mock.Anything, tc.session, tc.groupID, tc.childrenIDs).Return(tc.svcErr) - err := mgsdk.AddChildren(context.Background(), tc.groupID, tc.domainID, tc.childrenIDs, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "AddChildrenGroups", mock.Anything, tc.session, tc.groupID, tc.childrenIDs) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestRemoveChildrenGroups(t *testing.T) { - ts, csvc, auth := setupGroups() - defer ts.Close() - - conf := sdk.Config{ - GroupsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - groupID := testsutil.GenerateUUID(t) - childID := testsutil.GenerateUUID(t) - - cases := []struct { - desc string - domainID string - token string - session smqauthn.Session - groupID string - childrenIDs []string - svcErr error - authenticateErr error - err errors.SDKError - }{ - { - desc: "remove children group successfully", - domainID: domainID, - token: validToken, - groupID: groupID, - childrenIDs: []string{childID}, - svcErr: nil, - err: nil, - }, - { - desc: "remove children group with invalid token", - domainID: domainID, - token: invalidToken, - groupID: groupID, - childrenIDs: []string{childID}, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "remove children group with empty token", - domainID: domainID, - token: "", - groupID: groupID, - childrenIDs: []string{childID}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "remove children group with invalid group id", - domainID: domainID, - token: validToken, - groupID: wrongID, - childrenIDs: []string{childID}, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "remove children group with empty group id", - domainID: domainID, - token: validToken, - groupID: "", - childrenIDs: []string{childID}, - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "remove children group with empty children ids", - domainID: domainID, - token: validToken, - groupID: groupID, - childrenIDs: []string{}, - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingChildrenGroupIDs, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("RemoveChildrenGroups", mock.Anything, tc.session, tc.groupID, tc.childrenIDs).Return(tc.svcErr) - err := mgsdk.RemoveChildren(context.Background(), tc.groupID, tc.domainID, tc.childrenIDs, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RemoveChildrenGroups", mock.Anything, tc.session, tc.groupID, tc.childrenIDs) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestRemoveAllChildrenGroups(t *testing.T) { - ts, csvc, auth := setupGroups() - defer ts.Close() - - conf := sdk.Config{ - GroupsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - groupID := testsutil.GenerateUUID(t) - - cases := []struct { - desc string - domainID string - token string - session smqauthn.Session - groupID string - svcErr error - authenticateErr error - err errors.SDKError - }{ - { - desc: "remove all children group successfully", - domainID: domainID, - token: validToken, - groupID: groupID, - svcErr: nil, - err: nil, - }, - { - desc: "remove all children group with invalid token", - domainID: domainID, - token: invalidToken, - groupID: groupID, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "remove all children group with empty token", - domainID: domainID, - token: "", - groupID: groupID, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "remove all children group with invalid group id", - domainID: domainID, - token: validToken, - groupID: wrongID, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "remove all children group with empty group id", - domainID: domainID, - token: validToken, - groupID: "", - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("RemoveAllChildrenGroups", mock.Anything, tc.session, tc.groupID).Return(tc.svcErr) - err := mgsdk.RemoveAllChildren(context.Background(), tc.groupID, tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RemoveAllChildrenGroups", mock.Anything, tc.session, tc.groupID) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestListChildrenGroups(t *testing.T) { - ts, gsvc, auth := setupGroups() - defer ts.Close() - - var grps []sdk.Group - conf := sdk.Config{ - GroupsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - parentID := "" - for i := 10; i < 100; i++ { - gr := sdk.Group{ - ID: generateUUID(t), - Name: fmt.Sprintf("group_%d", i), - Metadata: sdk.Metadata{"name": fmt.Sprintf("user_%d", i)}, - Status: groups.EnabledStatus.String(), - ParentID: parentID, - Level: -1, - } - parentID = gr.ID - grps = append(grps, gr) - } - childID := grps[0].ID - - cases := []struct { - desc string - token string - domainID string - session smqauthn.Session - childID string - pageMeta sdk.PageMetadata - svcReq groups.Page - svcRes groups.Page - svcErr error - authenticateErr error - response sdk.GroupsPage - err errors.SDKError - }{ - { - desc: "list children groups successfully", - domainID: domainID, - token: validToken, - childID: childID, - pageMeta: sdk.PageMetadata{ - Offset: offset, - Limit: limit, - }, - svcReq: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: offset, - Limit: limit, - }, - }, - svcRes: groups.Page{ - PageMeta: groups.PageMeta{ - Total: uint64(len(grps[offset:limit])), - }, - Groups: convertGroups(grps[offset:limit]), - }, - response: sdk.GroupsPage{ - PageRes: sdk.PageRes{ - Total: uint64(len(grps[offset:limit])), - }, - Groups: grps[offset:limit], - }, - err: nil, - }, - { - desc: "list children groups with invalid token", - domainID: domainID, - token: invalidToken, - childID: childID, - pageMeta: sdk.PageMetadata{ - Offset: offset, - Limit: limit, - }, - svcReq: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: offset, - Limit: limit, - }, - }, - svcRes: groups.Page{}, - authenticateErr: svcerr.ErrAuthentication, - response: sdk.GroupsPage{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "list children groups with empty token", - domainID: domainID, - token: "", - childID: childID, - pageMeta: sdk.PageMetadata{ - Offset: offset, - Limit: limit, - }, - svcReq: groups.Page{}, - svcRes: groups.Page{}, - svcErr: nil, - response: sdk.GroupsPage{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "list children groups with zero limit", - domainID: domainID, - token: validToken, - childID: childID, - pageMeta: sdk.PageMetadata{ - Offset: offset, - Limit: 0, - }, - svcReq: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: offset, - Limit: 10, - }, - }, - svcRes: groups.Page{ - PageMeta: groups.PageMeta{ - Total: uint64(len(grps[offset:10])), - }, - Groups: convertGroups(grps[offset:10]), - }, - response: sdk.GroupsPage{ - PageRes: sdk.PageRes{ - Total: uint64(len(grps[offset:10])), - }, - Groups: grps[offset:10], - }, - err: nil, - }, - { - desc: "list children groups with limit greater than max", - domainID: domainID, - token: validToken, - childID: childID, - pageMeta: sdk.PageMetadata{ - Offset: offset, - Limit: 110, - }, - svcReq: groups.Page{}, - svcRes: groups.Page{}, - svcErr: nil, - response: sdk.GroupsPage{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrLimitSize, http.StatusBadRequest), - }, - { - desc: "list children groups with given metadata", - domainID: domainID, - token: validToken, - childID: childID, - pageMeta: sdk.PageMetadata{ - Offset: offset, - Limit: limit, - Metadata: sdk.Metadata{ - "name": "user_89", - }, - }, - svcReq: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: offset, - Limit: limit, - Metadata: groups.Metadata{ - "name": "user_89", - }, - }, - }, - svcRes: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 1, - }, - Groups: convertGroups([]sdk.Group{grps[89]}), - }, - response: sdk.GroupsPage{ - PageRes: sdk.PageRes{ - Total: 1, - }, - Groups: []sdk.Group{grps[89]}, - }, - err: nil, - }, - { - desc: "list children groups with invalid page metadata", - domainID: domainID, - token: validToken, - childID: childID, - pageMeta: sdk.PageMetadata{ - Offset: offset, - Limit: limit, - Metadata: sdk.Metadata{ - "key": make(chan int), - }, - }, - svcReq: groups.Page{}, - svcRes: groups.Page{}, - svcErr: nil, - response: sdk.GroupsPage{}, - err: errors.NewSDKError(fmt.Errorf("json: unsupported type: chan int")), - }, - { - desc: "list children groups with service response that cannot be unmarshalled", - domainID: domainID, - token: validToken, - childID: childID, - pageMeta: sdk.PageMetadata{ - Offset: offset, - Limit: limit, - }, - svcReq: groups.Page{ - PageMeta: groups.PageMeta{ - Offset: offset, - Limit: limit, - }, - }, - svcRes: groups.Page{ - PageMeta: groups.PageMeta{ - Total: 1, - }, - Groups: []groups.Group{{ - ID: generateUUID(t), - Name: "group_1", - Metadata: groups.Metadata{ - "key": make(chan int), - }, - Level: -1, - }}, - }, - svcErr: nil, - response: sdk.GroupsPage{}, - err: errors.NewSDKError(fmt.Errorf("unexpected end of JSON input")), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := gsvc.On("ListChildrenGroups", mock.Anything, tc.session, tc.childID, int64(1), int64(0), mock.Anything).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.Children(context.Background(), tc.childID, tc.domainID, tc.pageMeta, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "ListChildrenGroups", mock.Anything, tc.session, tc.childID, int64(1), int64(0), mock.Anything) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestHierarchy(t *testing.T) { - ts, gsvc, auth := setupGroups() - defer ts.Close() - - var grps []sdk.Group - conf := sdk.Config{ - GroupsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - parentID := "" - for i := 10; i < 100; i++ { - gr := sdk.Group{ - ID: generateUUID(t), - Name: fmt.Sprintf("group_%d", i), - Metadata: sdk.Metadata{"name": fmt.Sprintf("user_%d", i)}, - Status: groups.EnabledStatus.String(), - ParentID: parentID, - Level: -1, - } - parentID = gr.ID - grps = append(grps, gr) - } - childID := grps[0].ID - - cases := []struct { - desc string - token string - domainID string - session smqauthn.Session - groupID string - pageMeta sdk.PageMetadata - svcReq groups.HierarchyPageMeta - svcRes groups.HierarchyPage - svcErr error - authenticateErr error - response sdk.GroupsHierarchyPage - err errors.SDKError - }{ - { - desc: "list hierarchy successfully", - domainID: domainID, - token: validToken, - groupID: childID, - pageMeta: sdk.PageMetadata{ - Level: 2, - Tree: false, - }, - svcReq: groups.HierarchyPageMeta{ - Level: 2, - Direction: -1, - Tree: false, - }, - svcRes: groups.HierarchyPage{ - HierarchyPageMeta: groups.HierarchyPageMeta{ - Level: 2, - Direction: +1, - Tree: false, - }, - Groups: convertGroups(grps[1:]), - }, - response: sdk.GroupsHierarchyPage{ - Level: 2, - Direction: +1, - Groups: grps[1:], - }, - err: nil, - }, - { - desc: "list hierarchy with invalid token", - domainID: domainID, - token: validToken, - groupID: childID, - pageMeta: sdk.PageMetadata{ - Level: 2, - Tree: false, - }, - svcReq: groups.HierarchyPageMeta{ - Level: 2, - Direction: -1, - Tree: false, - }, - authenticateErr: svcerr.ErrAuthentication, - response: sdk.GroupsHierarchyPage{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "list hierarchy with empty token", - domainID: domainID, - token: "", - groupID: childID, - pageMeta: sdk.PageMetadata{ - Level: 2, - Tree: false, - }, - svcReq: groups.HierarchyPageMeta{ - Level: 2, - Direction: -1, - Tree: false, - }, - svcErr: nil, - response: sdk.GroupsHierarchyPage{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "list hierarchy with invalid group id", - domainID: domainID, - token: validToken, - groupID: wrongID, - pageMeta: sdk.PageMetadata{ - Level: 2, - Tree: false, - }, - svcReq: groups.HierarchyPageMeta{ - Level: 2, - Direction: -1, - Tree: false, - }, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "list hierarchy with response that cannot be unmarshalled", - domainID: domainID, - token: validToken, - groupID: childID, - pageMeta: sdk.PageMetadata{ - Level: 2, - Tree: false, - }, - svcReq: groups.HierarchyPageMeta{ - Level: 2, - Direction: -1, - Tree: false, - }, - svcRes: groups.HierarchyPage{ - HierarchyPageMeta: groups.HierarchyPageMeta{ - Level: 2, - Direction: +1, - Tree: false, - }, - Groups: []groups.Group{{ - ID: generateUUID(t), - Name: "group_1", - Metadata: groups.Metadata{ - "key": make(chan int), - }, - Level: -1, - }}, - }, - svcErr: nil, - response: sdk.GroupsHierarchyPage{}, - err: errors.NewSDKError(fmt.Errorf("unexpected end of JSON input")), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := gsvc.On("RetrieveGroupHierarchy", mock.Anything, tc.session, tc.groupID, tc.svcReq).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.Hierarchy(context.Background(), tc.groupID, tc.domainID, tc.pageMeta, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RetrieveGroupHierarchy", mock.Anything, tc.session, tc.groupID, tc.svcReq) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestCreateGroupRole(t *testing.T) { - ts, csvc, auth := setupGroups() - defer ts.Close() - - conf := sdk.Config{ - GroupsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - optionalActions := []string{"create", "update"} - optionalMembers := []string{testsutil.GenerateUUID(t), testsutil.GenerateUUID(t)} - rReq := sdk.RoleReq{ - RoleName: roleName, - OptionalActions: optionalActions, - OptionalMembers: optionalMembers, - } - userID := testsutil.GenerateUUID(t) - groupID := testsutil.GenerateUUID(t) - now := time.Now().UTC() - role := roles.Role{ - ID: testsutil.GenerateUUID(t), - Name: rReq.RoleName, - EntityID: groupID, - CreatedBy: userID, - CreatedAt: now, - } - roleProvision := roles.RoleProvision{ - Role: role, - OptionalActions: optionalActions, - OptionalMembers: optionalMembers, - } - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - groupID string - roleReq sdk.RoleReq - svcRes roles.RoleProvision - svcErr error - authenticateErr error - response sdk.Role - err errors.SDKError - }{ - { - desc: "create group role successfully", - token: validToken, - domainID: domainID, - groupID: groupID, - roleReq: rReq, - svcRes: roleProvision, - svcErr: nil, - response: convertRoleProvision(roleProvision), - err: nil, - }, - { - desc: "create group role with invalid token", - token: invalidToken, - domainID: domainID, - groupID: groupID, - roleReq: rReq, - svcRes: roles.RoleProvision{}, - authenticateErr: svcerr.ErrAuthentication, - response: sdk.Role{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "create group role with empty token", - token: "", - domainID: domainID, - groupID: groupID, - roleReq: rReq, - svcRes: roles.RoleProvision{}, - response: sdk.Role{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "create group role with invalid group id", - token: validToken, - domainID: domainID, - groupID: testsutil.GenerateUUID(t), - roleReq: rReq, - svcRes: roles.RoleProvision{}, - svcErr: svcerr.ErrAuthorization, - response: sdk.Role{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "create group role with empty group id", - token: validToken, - domainID: domainID, - groupID: "", - roleReq: rReq, - svcRes: roles.RoleProvision{}, - svcErr: nil, - response: sdk.Role{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrInvalidIDFormat, http.StatusBadRequest), - }, - { - desc: "create group role with empty role name", - token: validToken, - domainID: domainID, - groupID: groupID, - roleReq: sdk.RoleReq{ - RoleName: "", - OptionalActions: []string{"create", "update"}, - OptionalMembers: []string{testsutil.GenerateUUID(t), testsutil.GenerateUUID(t)}, - }, - svcRes: roles.RoleProvision{}, - svcErr: nil, - response: sdk.Role{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingRoleName, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("AddRole", mock.Anything, tc.session, tc.groupID, tc.roleReq.RoleName, tc.roleReq.OptionalActions, tc.roleReq.OptionalMembers).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.CreateGroupRole(context.Background(), tc.groupID, tc.domainID, tc.roleReq, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "AddRole", mock.Anything, tc.session, tc.groupID, tc.roleReq.RoleName, tc.roleReq.OptionalActions, tc.roleReq.OptionalMembers) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestListGroupRoles(t *testing.T) { - ts, csvc, auth := setupGroups() - defer ts.Close() - - conf := sdk.Config{ - GroupsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - groupID := testsutil.GenerateUUID(t) - role := roles.Role{ - ID: testsutil.GenerateUUID(t), - Name: roleName, - EntityID: groupID, - CreatedBy: testsutil.GenerateUUID(t), - CreatedAt: time.Now().UTC(), - } - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - groupID string - pageMeta sdk.PageMetadata - svcRes roles.RolePage - svcErr error - authenticateErr error - response sdk.RolesPage - err errors.SDKError - }{ - { - desc: "list group roles successfully", - token: validToken, - domainID: domainID, - groupID: groupID, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - }, - svcRes: roles.RolePage{ - Total: 1, - Offset: 0, - Limit: 10, - Roles: []roles.Role{role}, - }, - svcErr: nil, - response: sdk.RolesPage{ - Total: 1, - Offset: 0, - Limit: 10, - Roles: []sdk.Role{convertRole(role)}, - }, - err: nil, - }, - { - desc: "list group roles with invalid token", - token: invalidToken, - domainID: domainID, - groupID: groupID, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - }, - svcRes: roles.RolePage{}, - authenticateErr: svcerr.ErrAuthentication, - response: sdk.RolesPage{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "list group roles with empty token", - token: "", - domainID: domainID, - groupID: groupID, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - }, - svcRes: roles.RolePage{}, - response: sdk.RolesPage{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "list group roles with invalid group id", - token: validToken, - domainID: domainID, - groupID: testsutil.GenerateUUID(t), - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - }, - svcRes: roles.RolePage{}, - svcErr: svcerr.ErrAuthorization, - response: sdk.RolesPage{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "list group roles with empty group id", - token: validToken, - domainID: domainID, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - }, - groupID: "", - svcRes: roles.RolePage{}, - svcErr: nil, - response: sdk.RolesPage{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("RetrieveAllRoles", mock.Anything, tc.session, tc.groupID, tc.pageMeta.Limit, tc.pageMeta.Offset).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.GroupRoles(context.Background(), tc.groupID, tc.domainID, tc.pageMeta, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RetrieveAllRoles", mock.Anything, tc.session, tc.groupID, tc.pageMeta.Limit, tc.pageMeta.Offset) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestViewGroupRole(t *testing.T) { - ts, csvc, auth := setupGroups() - defer ts.Close() - - conf := sdk.Config{ - GroupsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - groupID := testsutil.GenerateUUID(t) - role := roles.Role{ - ID: testsutil.GenerateUUID(t), - Name: roleName, - EntityID: groupID, - CreatedBy: testsutil.GenerateUUID(t), - CreatedAt: time.Now().UTC(), - } - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - groupID string - roleID string - svcRes roles.Role - svcErr error - authenticateErr error - response sdk.Role - err errors.SDKError - }{ - { - desc: "view group role successfully", - token: validToken, - domainID: domainID, - groupID: groupID, - roleID: role.ID, - svcRes: role, - svcErr: nil, - response: convertRole(role), - err: nil, - }, - { - desc: "view group role with invalid token", - token: invalidToken, - domainID: domainID, - groupID: groupID, - roleID: role.ID, - svcRes: roles.Role{}, - authenticateErr: svcerr.ErrAuthentication, - response: sdk.Role{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "view group role with empty token", - token: "", - domainID: domainID, - groupID: groupID, - roleID: role.ID, - svcRes: roles.Role{}, - response: sdk.Role{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "view group role with invalid group id", - token: validToken, - domainID: domainID, - groupID: testsutil.GenerateUUID(t), - roleID: role.ID, - svcRes: roles.Role{}, - svcErr: svcerr.ErrAuthorization, - response: sdk.Role{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "view group role with empty group id", - token: validToken, - domainID: domainID, - groupID: "", - roleID: role.ID, - svcRes: roles.Role{}, - svcErr: nil, - response: sdk.Role{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "view group role with invalid role id", - token: validToken, - domainID: domainID, - groupID: groupID, - roleID: invalid, - svcRes: roles.Role{}, - svcErr: svcerr.ErrAuthorization, - response: sdk.Role{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("RetrieveRole", mock.Anything, tc.session, tc.groupID, tc.roleID).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.GroupRole(context.Background(), tc.groupID, tc.roleID, tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RetrieveRole", mock.Anything, tc.session, tc.groupID, tc.roleID) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestUpdateGroupRole(t *testing.T) { - ts, csvc, auth := setupGroups() - defer ts.Close() - - conf := sdk.Config{ - GroupsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - groupID := testsutil.GenerateUUID(t) - roleID := testsutil.GenerateUUID(t) - newRoleName := "newTest" - userID := testsutil.GenerateUUID(t) - createdAt := time.Now().UTC().Add(-time.Hour) - role := roles.Role{ - ID: testsutil.GenerateUUID(t), - Name: newRoleName, - EntityID: groupID, - CreatedBy: userID, - CreatedAt: createdAt, - UpdatedBy: userID, - UpdatedAt: time.Now().UTC(), - } - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - groupID string - roleID string - newRoleName string - svcRes roles.Role - svcErr error - authenticateErr error - response sdk.Role - err errors.SDKError - }{ - { - desc: "update group role successfully", - token: validToken, - domainID: domainID, - groupID: groupID, - roleID: roleID, - newRoleName: newRoleName, - svcRes: role, - svcErr: nil, - response: convertRole(role), - err: nil, - }, - { - desc: "update group role with invalid token", - token: invalidToken, - domainID: domainID, - groupID: groupID, - roleID: roleID, - newRoleName: newRoleName, - svcRes: roles.Role{}, - authenticateErr: svcerr.ErrAuthentication, - response: sdk.Role{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "update group role with empty token", - token: "", - domainID: domainID, - groupID: groupID, - roleID: roleID, - newRoleName: newRoleName, - svcRes: roles.Role{}, - response: sdk.Role{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "update group role with invalid group id", - token: validToken, - domainID: domainID, - groupID: testsutil.GenerateUUID(t), - roleID: roleID, - newRoleName: newRoleName, - svcRes: roles.Role{}, - svcErr: svcerr.ErrAuthorization, - response: sdk.Role{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "update group role with empty group id", - token: validToken, - domainID: domainID, - groupID: "", - roleID: roleID, - newRoleName: newRoleName, - svcRes: roles.Role{}, - svcErr: nil, - response: sdk.Role{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("UpdateRoleName", mock.Anything, tc.session, tc.groupID, tc.roleID, tc.newRoleName).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.UpdateGroupRole(context.Background(), tc.groupID, tc.roleID, tc.newRoleName, tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "UpdateRoleName", mock.Anything, tc.session, tc.groupID, tc.roleID, tc.newRoleName) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestDeleteGroupRole(t *testing.T) { - ts, csvc, auth := setupGroups() - defer ts.Close() - - conf := sdk.Config{ - GroupsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - roleID := testsutil.GenerateUUID(t) - groupID := testsutil.GenerateUUID(t) - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - groupID string - roleID string - svcErr error - authenticateErr error - err errors.SDKError - }{ - { - desc: "delete group role successfully", - token: validToken, - domainID: domainID, - groupID: groupID, - roleID: roleID, - svcErr: nil, - err: nil, - }, - { - desc: "delete group role with invalid token", - token: invalidToken, - domainID: domainID, - groupID: groupID, - roleID: roleID, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "delete group role with empty token", - token: "", - domainID: domainID, - groupID: groupID, - roleID: roleID, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "delete group role with invalid group id", - token: validToken, - domainID: domainID, - groupID: testsutil.GenerateUUID(t), - roleID: roleID, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "delete group role with empty group id", - token: validToken, - domainID: domainID, - groupID: "", - roleID: roleID, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "delete group role with invalid role id", - token: validToken, - domainID: domainID, - groupID: groupID, - roleID: invalid, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("RemoveRole", mock.Anything, tc.session, tc.groupID, tc.roleID).Return(tc.svcErr) - err := mgsdk.DeleteGroupRole(context.Background(), tc.groupID, tc.roleID, tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RemoveRole", mock.Anything, tc.session, tc.groupID, tc.roleID) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestAddGroupRoleActions(t *testing.T) { - ts, csvc, auth := setupGroups() - defer ts.Close() - - conf := sdk.Config{ - GroupsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - roleID := testsutil.GenerateUUID(t) - actions := []string{"create", "update"} - groupID := testsutil.GenerateUUID(t) - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - groupID string - roleID string - actions []string - svcRes []string - svcErr error - authenticateErr error - response []string - err errors.SDKError - }{ - { - desc: "add group role actions successfully", - token: validToken, - domainID: domainID, - groupID: groupID, - roleID: roleID, - actions: actions, - svcRes: actions, - svcErr: nil, - response: actions, - err: nil, - }, - { - desc: "add group role actions with invalid token", - token: invalidToken, - domainID: domainID, - groupID: groupID, - roleID: roleID, - actions: actions, - authenticateErr: svcerr.ErrAuthentication, - response: []string{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "add group role actions with empty token", - token: "", - domainID: domainID, - groupID: groupID, - roleID: roleID, - actions: actions, - response: []string{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "add group role actions with invalid group id", - token: validToken, - domainID: domainID, - groupID: testsutil.GenerateUUID(t), - roleID: roleID, - actions: actions, - svcErr: svcerr.ErrAuthorization, - response: []string{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "add group role actions with empty group id", - token: validToken, - domainID: domainID, - groupID: "", - roleID: roleID, - actions: actions, - response: []string{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "add group role actions with invalid role id", - token: validToken, - domainID: domainID, - groupID: groupID, - roleID: invalid, - actions: actions, - svcErr: svcerr.ErrAuthorization, - response: []string{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "add group role actions with empty actions", - token: validToken, - domainID: domainID, - groupID: groupID, - roleID: roleID, - actions: []string{}, - svcErr: nil, - response: []string{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingPolicyEntityType, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("RoleAddActions", mock.Anything, tc.session, tc.groupID, tc.roleID, tc.actions).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.AddGroupRoleActions(context.Background(), tc.groupID, tc.roleID, tc.domainID, tc.actions, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RoleAddActions", mock.Anything, tc.session, tc.groupID, tc.roleID, tc.actions) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestListGroupRoleActions(t *testing.T) { - ts, csvc, auth := setupGroups() - defer ts.Close() - - conf := sdk.Config{ - GroupsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - roleID := testsutil.GenerateUUID(t) - actions := []string{"create", "update"} - groupID := testsutil.GenerateUUID(t) - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - groupID string - roleID string - svcRes []string - svcErr error - authenticateErr error - response []string - err errors.SDKError - }{ - { - desc: "list group role actions successfully", - token: validToken, - domainID: domainID, - groupID: groupID, - roleID: roleID, - svcRes: actions, - svcErr: nil, - response: actions, - err: nil, - }, - { - desc: "list group role actions with invalid token", - token: invalidToken, - domainID: domainID, - groupID: groupID, - roleID: roleID, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "list group role actions with empty token", - token: "", - domainID: domainID, - groupID: groupID, - roleID: roleID, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "list group role actions with invalid group id", - token: validToken, - domainID: domainID, - groupID: testsutil.GenerateUUID(t), - roleID: roleID, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "list group role actions with empty group id", - token: validToken, - domainID: domainID, - groupID: "", - roleID: roleID, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "list group role actions with invalid role id", - token: validToken, - domainID: domainID, - groupID: groupID, - roleID: invalid, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "list group role actions with empty role id", - token: validToken, - domainID: domainID, - groupID: groupID, - roleID: "", - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingRoleID, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("RoleListActions", mock.Anything, tc.session, tc.groupID, tc.roleID).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.GroupRoleActions(context.Background(), tc.groupID, tc.roleID, tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RoleListActions", mock.Anything, tc.session, tc.groupID, tc.roleID) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestRemoveGroupRoleActions(t *testing.T) { - ts, csvc, auth := setupGroups() - defer ts.Close() - - conf := sdk.Config{ - GroupsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - roleID := testsutil.GenerateUUID(t) - actions := []string{"create", "update"} - groupID := testsutil.GenerateUUID(t) - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - groupID string - roleID string - actions []string - svcErr error - authenticateErr error - err errors.SDKError - }{ - { - desc: "remove group role actions successfully", - token: validToken, - domainID: domainID, - groupID: groupID, - roleID: roleID, - actions: actions, - svcErr: nil, - err: nil, - }, - { - desc: "remove group role actions with invalid token", - token: invalidToken, - domainID: domainID, - groupID: groupID, - roleID: roleID, - actions: actions, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "remove group role actions with empty token", - token: "", - domainID: domainID, - groupID: groupID, - roleID: roleID, - actions: actions, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "remove group role actions with invalid group id", - token: validToken, - domainID: domainID, - groupID: testsutil.GenerateUUID(t), - roleID: roleID, - actions: actions, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "remove group role actions with empty group id", - token: validToken, - domainID: domainID, - groupID: "", - roleID: roleID, - actions: actions, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "remove group role actions with invalid role id", - token: validToken, - domainID: domainID, - groupID: groupID, - roleID: invalid, - actions: actions, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "remove group role actions with empty actions", - token: validToken, - domainID: domainID, - groupID: groupID, - roleID: roleID, - actions: []string{}, - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingPolicyEntityType, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("RoleRemoveActions", mock.Anything, tc.session, tc.groupID, tc.roleID, tc.actions).Return(tc.svcErr) - err := mgsdk.RemoveGroupRoleActions(context.Background(), tc.groupID, tc.roleID, tc.domainID, tc.actions, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RoleRemoveActions", mock.Anything, tc.session, tc.groupID, tc.roleID, tc.actions) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestRemoveAllGroupRoleActions(t *testing.T) { - ts, csvc, auth := setupGroups() - defer ts.Close() - - conf := sdk.Config{ - GroupsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - roleID := testsutil.GenerateUUID(t) - groupID := testsutil.GenerateUUID(t) - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - groupID string - roleID string - svcErr error - authenticateErr error - err errors.SDKError - }{ - { - desc: "remove all group role actions successfully", - token: validToken, - domainID: domainID, - groupID: groupID, - roleID: roleID, - svcErr: nil, - err: nil, - }, - { - desc: "remove all group role actions with invalid token", - token: invalidToken, - domainID: domainID, - groupID: groupID, - roleID: roleID, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "remove all group role actions with empty token", - token: "", - domainID: domainID, - groupID: groupID, - roleID: roleID, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "remove all group role actions with invalid group id", - token: validToken, - domainID: domainID, - groupID: testsutil.GenerateUUID(t), - roleID: roleID, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "remove all group role actions with empty group id", - token: validToken, - domainID: domainID, - groupID: "", - roleID: roleID, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "remove all group role actions with invalid role id", - token: validToken, - domainID: domainID, - groupID: groupID, - roleID: invalid, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "remove all group role actions with empty role id", - token: validToken, - domainID: domainID, - groupID: groupID, - roleID: "", - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingRoleID, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("RoleRemoveAllActions", mock.Anything, tc.session, tc.groupID, tc.roleID).Return(tc.svcErr) - err := mgsdk.RemoveAllGroupRoleActions(context.Background(), tc.groupID, tc.roleID, tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RoleRemoveAllActions", mock.Anything, tc.session, tc.groupID, tc.roleID) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestAddGroupRoleMembers(t *testing.T) { - ts, csvc, auth := setupGroups() - defer ts.Close() - - conf := sdk.Config{ - GroupsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - roleID := testsutil.GenerateUUID(t) - members := []string{"user1", "user2"} - groupID := testsutil.GenerateUUID(t) - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - groupID string - roleID string - members []string - svcRes []string - svcErr error - authenticateErr error - response []string - err errors.SDKError - }{ - { - desc: "add group role members successfully", - token: validToken, - domainID: domainID, - groupID: groupID, - roleID: roleID, - members: members, - svcRes: members, - svcErr: nil, - response: members, - err: nil, - }, - { - desc: "add group role members with invalid token", - token: invalidToken, - domainID: domainID, - groupID: groupID, - roleID: roleID, - members: members, - authenticateErr: svcerr.ErrAuthentication, - response: []string{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "add group role members with empty token", - token: "", - domainID: domainID, - groupID: groupID, - roleID: roleID, - members: members, - response: []string{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "add group role members with invalid group id", - token: validToken, - domainID: domainID, - groupID: testsutil.GenerateUUID(t), - roleID: roleID, - members: members, - svcErr: svcerr.ErrAuthorization, - response: []string{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "add group role members with empty group id", - token: validToken, - domainID: domainID, - groupID: "", - roleID: roleID, - members: members, - response: []string{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "add group role members with invalid role id", - token: validToken, - domainID: domainID, - groupID: groupID, - roleID: invalid, - members: members, - svcErr: svcerr.ErrAuthorization, - response: []string{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "add group role members with empty members", - token: validToken, - domainID: domainID, - groupID: groupID, - roleID: roleID, - members: []string{}, - svcErr: nil, - response: []string{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingRoleMembers, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("RoleAddMembers", mock.Anything, tc.session, tc.groupID, tc.roleID, tc.members).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.AddGroupRoleMembers(context.Background(), tc.groupID, tc.roleID, tc.domainID, tc.members, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RoleAddMembers", mock.Anything, tc.session, tc.groupID, tc.roleID, tc.members) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestListGroupRoleMembers(t *testing.T) { - ts, csvc, auth := setupGroups() - defer ts.Close() - - conf := sdk.Config{ - GroupsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - roleID := testsutil.GenerateUUID(t) - members := []string{"user1", "user2"} - groupID := testsutil.GenerateUUID(t) - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - groupID string - roleID string - pageMeta sdk.PageMetadata - svcRes roles.MembersPage - svcErr error - authenticateErr error - response sdk.RoleMembersPage - err errors.SDKError - }{ - { - desc: "list group role members successfully", - token: validToken, - domainID: domainID, - groupID: groupID, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 5, - }, - roleID: roleID, - svcRes: roles.MembersPage{ - Total: 2, - Offset: 0, - Limit: 5, - Members: members, - }, - svcErr: nil, - response: sdk.RoleMembersPage{ - Total: 2, - Offset: 0, - Limit: 5, - Members: members, - }, - err: nil, - }, - { - desc: "list group role members with invalid token", - token: invalidToken, - domainID: domainID, - groupID: groupID, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 5, - }, - roleID: roleID, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "list group role members with empty token", - token: "", - domainID: domainID, - groupID: groupID, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 5, - }, - roleID: roleID, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "list group role members with invalid group id", - token: validToken, - domainID: domainID, - groupID: testsutil.GenerateUUID(t), - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 5, - }, - roleID: roleID, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "list group role members with empty group id", - token: validToken, - domainID: domainID, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 5, - }, - groupID: "", - roleID: roleID, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "list group role members with invalid role id", - token: validToken, - domainID: domainID, - groupID: groupID, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 5, - }, - roleID: invalid, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "list group role members with empty role id", - token: validToken, - domainID: domainID, - groupID: groupID, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 5, - }, - roleID: "", - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingRoleID, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("RoleListMembers", mock.Anything, tc.session, tc.groupID, tc.roleID, tc.pageMeta.Limit, tc.pageMeta.Offset).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.GroupRoleMembers(context.Background(), tc.groupID, tc.roleID, tc.domainID, tc.pageMeta, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RoleListMembers", mock.Anything, tc.session, tc.groupID, tc.roleID, tc.pageMeta.Limit, tc.pageMeta.Offset) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestRemoveGroupRoleMembers(t *testing.T) { - ts, csvc, auth := setupGroups() - defer ts.Close() - - conf := sdk.Config{ - GroupsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - roleID := testsutil.GenerateUUID(t) - members := []string{"user1", "user2"} - groupID := testsutil.GenerateUUID(t) - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - groupID string - roleID string - members []string - svcErr error - authenticateErr error - err errors.SDKError - }{ - { - desc: "remove group role members successfully", - token: validToken, - domainID: domainID, - groupID: groupID, - roleID: roleID, - members: members, - svcErr: nil, - err: nil, - }, - { - desc: "remove group role members with invalid token", - token: invalidToken, - domainID: domainID, - groupID: groupID, - roleID: roleID, - members: members, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "remove group role members with empty token", - token: "", - domainID: domainID, - groupID: groupID, - roleID: roleID, - members: members, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "remove group role members with invalid group id", - token: validToken, - domainID: domainID, - groupID: testsutil.GenerateUUID(t), - roleID: roleID, - members: members, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "remove group role members with empty group id", - token: validToken, - domainID: domainID, - groupID: "", - roleID: roleID, - members: members, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "remove group role members with invalid role id", - token: validToken, - domainID: domainID, - groupID: groupID, - roleID: invalid, - members: members, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "remove group role members with empty members", - token: validToken, - domainID: domainID, - groupID: groupID, - roleID: roleID, - members: []string{}, - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingRoleMembers, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("RoleRemoveMembers", mock.Anything, tc.session, tc.groupID, tc.roleID, tc.members).Return(tc.svcErr) - err := mgsdk.RemoveGroupRoleMembers(context.Background(), tc.groupID, tc.roleID, tc.domainID, tc.members, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RoleRemoveMembers", mock.Anything, tc.session, tc.groupID, tc.roleID, tc.members) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestRemoveAllGroupRoleMembers(t *testing.T) { - ts, csvc, auth := setupGroups() - defer ts.Close() - - conf := sdk.Config{ - GroupsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - roleID := testsutil.GenerateUUID(t) - groupID := testsutil.GenerateUUID(t) - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - groupID string - roleID string - svcErr error - authenticateErr error - err errors.SDKError - }{ - { - desc: "remove all group role members successfully", - token: validToken, - domainID: domainID, - groupID: groupID, - roleID: roleID, - svcErr: nil, - err: nil, - }, - { - desc: "remove all group role members with invalid token", - token: invalidToken, - domainID: domainID, - groupID: groupID, - roleID: roleID, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "remove all group role members with empty token", - token: "", - domainID: domainID, - groupID: groupID, - roleID: roleID, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "remove all group role members with invalid group id", - token: validToken, - domainID: domainID, - groupID: testsutil.GenerateUUID(t), - roleID: roleID, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "remove all group role members with empty group id", - token: validToken, - domainID: domainID, - groupID: "", - roleID: roleID, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "remove all group role members with invalid role id", - token: validToken, - domainID: domainID, - groupID: groupID, - roleID: invalid, - svcErr: svcerr.ErrAuthorization, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "remove all group role members with empty role id", - token: validToken, - domainID: domainID, - groupID: groupID, - roleID: "", - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingRoleID, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("RoleRemoveAllMembers", mock.Anything, tc.session, tc.groupID, tc.roleID).Return(tc.svcErr) - err := mgsdk.RemoveAllGroupRoleMembers(context.Background(), tc.groupID, tc.roleID, tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RoleRemoveAllMembers", mock.Anything, tc.session, tc.groupID, tc.roleID) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestListAvailableGroupRoleActions(t *testing.T) { - ts, csvc, auth := setupGroups() - defer ts.Close() - - conf := sdk.Config{ - GroupsURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - actions := []string{"create", "update"} - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - svcRes []string - svcErr error - authenticateErr error - response []string - err errors.SDKError - }{ - { - desc: "list available role actions successfully", - token: validToken, - domainID: domainID, - svcRes: actions, - svcErr: nil, - response: actions, - err: nil, - }, - { - desc: "list available role actions with invalid token", - token: invalidToken, - domainID: domainID, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "list available role actions with empty token", - token: "", - domainID: domainID, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "list available role actions with empty domain id", - token: validToken, - domainID: "", - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingDomainID, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := csvc.On("ListAvailableActions", mock.Anything, tc.session).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.AvailableGroupRoleActions(context.Background(), tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "ListAvailableActions", mock.Anything, tc.session) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func generateTestGroup(t *testing.T) sdk.Group { - createdAt, err := time.Parse(time.RFC3339, "2023-03-03T00:00:00Z") - assert.Nil(t, err, fmt.Sprintf("unexpected error %s", err)) - updatedAt := createdAt - gr := sdk.Group{ - ID: testsutil.GenerateUUID(t), - DomainID: testsutil.GenerateUUID(t), - Name: gName, - Description: description, - Metadata: sdk.Metadata{"role": "client"}, - CreatedAt: createdAt, - UpdatedAt: updatedAt, - Status: groups.EnabledStatus.String(), - Roles: validRoles, - } - return gr -} diff --git a/pkg/sdk/health_test.go b/pkg/sdk/health_test.go deleted file mode 100644 index f1a854eea..000000000 --- a/pkg/sdk/health_test.go +++ /dev/null @@ -1,126 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package sdk_test - -import ( - "fmt" - "testing" - - "github.com/absmach/magistrala" - "github.com/absmach/magistrala/pkg/errors" - sdk "github.com/absmach/magistrala/pkg/sdk" - "github.com/stretchr/testify/assert" -) - -func TestHealth(t *testing.T) { - clientsTs, _, _ := setupClients() - defer clientsTs.Close() - - usersTs, _, _ := setupUsers() - defer usersTs.Close() - - groupsTs, _, _ := setupGroups() - defer groupsTs.Close() - - channelsTs, _, _ := setupChannels() - defer channelsTs.Close() - - domainsTs, _, _ := setupDomains() - defer domainsTs.Close() - - journalTs, _, _ := setupJournal() - defer journalTs.Close() - - fluxmqTs := setupFluxMQ("any") - defer fluxmqTs.Close() - - sdkConf := sdk.Config{ - ClientsURL: clientsTs.URL, - UsersURL: usersTs.URL, - HTTPAdapterURL: fluxmqTs.URL, - GroupsURL: groupsTs.URL, - ChannelsURL: channelsTs.URL, - DomainsURL: domainsTs.URL, - JournalURL: journalTs.URL, - MsgContentType: contentType, - TLSVerification: false, - } - - mgsdk := sdk.NewSDK(sdkConf) - cases := []struct { - desc string - service string - empty bool - description string - status string - err errors.SDKError - }{ - { - desc: "get clients service health check", - service: "clients", - empty: false, - err: nil, - description: "clients service", - status: "pass", - }, - { - desc: "get users service health check", - service: "users", - empty: false, - err: nil, - description: "users service", - status: "pass", - }, - { - desc: "get groups service health check", - service: "groups", - empty: false, - err: nil, - description: "groups service", - status: "pass", - }, - { - desc: "get channels service health check", - service: "channels", - empty: false, - err: nil, - description: "channels service", - status: "pass", - }, - { - desc: "get domains service health check", - service: "domains", - empty: false, - err: nil, - description: "domains service", - status: "pass", - }, - { - desc: "get journal service health check", - service: "journal", - empty: false, - err: nil, - description: "journal-log service", - status: "pass", - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - h, err := mgsdk.Health(tc.service) - assert.Equal(t, tc.err, err, fmt.Sprintf("%s: expected error %s, got %s", tc.desc, tc.err, err)) - assert.Equal(t, tc.status, h.Status, fmt.Sprintf("%s: expected %s status, got %s", tc.desc, tc.status, h.Status)) - assert.Equal(t, tc.empty, h.Version == "", fmt.Sprintf("%s: expected non-empty version", tc.desc)) - assert.Equal(t, magistrala.Commit, h.Commit, fmt.Sprintf("%s: expected non-empty commit", tc.desc)) - assert.Equal(t, tc.description, h.Description, fmt.Sprintf("%s: expected proper description, got %s", tc.desc, h.Description)) - assert.Equal(t, magistrala.BuildTime, h.BuildTime, fmt.Sprintf("%s: expected default epoch date, got %s", tc.desc, h.BuildTime)) - }) - } - - // FluxMQ returns a simpler health response without version/commit/description. - t.Run("get fluxmq service health check", func(t *testing.T) { - h, err := mgsdk.Health("fluxmq") - assert.Nil(t, err) - assert.Equal(t, "healthy", h.Status) - }) -} diff --git a/pkg/sdk/invitations_test.go b/pkg/sdk/invitations_test.go deleted file mode 100644 index bca45034e..000000000 --- a/pkg/sdk/invitations_test.go +++ /dev/null @@ -1,468 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package sdk_test - -import ( - "context" - "fmt" - "net/http" - "testing" - "time" - - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/domains" - "github.com/absmach/magistrala/internal/testsutil" - smqauthn "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - sdk "github.com/absmach/magistrala/pkg/sdk" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" -) - -var ( - sdkInvitation = generateTestInvitation(&testing.T{}) - invitation = convertInvitation(sdkInvitation) -) - -func TestSendInvitation(t *testing.T) { - is, svc, auth := setupDomains() - defer is.Close() - - conf := sdk.Config{ - DomainsURL: is.URL, - } - mgsdk := sdk.NewSDK(conf) - - sendInvitationReq := sdk.Invitation{ - InviteeUserID: invitation.InviteeUserID, - DomainID: invitation.DomainID, - RoleID: invitation.RoleID, - } - - cases := []struct { - desc string - token string - session smqauthn.Session - sendInvitationReq sdk.Invitation - svcReq domains.Invitation - authenticateErr error - svcErr error - err error - }{ - { - desc: "send invitation successfully", - token: validToken, - sendInvitationReq: sendInvitationReq, - svcReq: convertInvitation(sendInvitationReq), - svcErr: nil, - err: nil, - }, - { - desc: "send invitation with invalid token", - token: invalidToken, - sendInvitationReq: sendInvitationReq, - svcReq: convertInvitation(sendInvitationReq), - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "send invitation with empty token", - token: "", - sendInvitationReq: sendInvitationReq, - svcReq: domains.Invitation{}, - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "send invitation with empty userID", - token: validToken, - sendInvitationReq: sdk.Invitation{ - InviteeUserID: "", - DomainID: invitation.DomainID, - RoleID: invitation.RoleID, - }, - svcReq: domains.Invitation{}, - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "send invitation with empty role ID", - token: validToken, - sendInvitationReq: sdk.Invitation{ - InviteeUserID: invitation.InviteeUserID, - DomainID: invitation.DomainID, - RoleID: "", - }, - svcReq: domains.Invitation{}, - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "send inviation with invalid domainID", - token: validToken, - sendInvitationReq: sdk.Invitation{ - InviteeUserID: invitation.InviteeUserID, - DomainID: wrongID, - RoleID: invitation.RoleID, - }, - svcReq: domains.Invitation{ - InviteeUserID: invitation.InviteeUserID, - DomainID: wrongID, - RoleID: invitation.RoleID, - }, - svcErr: svcerr.ErrCreateEntity, - err: errors.NewSDKErrorWithStatus(svcerr.ErrCreateEntity, http.StatusUnprocessableEntity), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == valid { - tc.session = smqauthn.Session{ - UserID: tc.sendInvitationReq.InviteeUserID, - DomainID: tc.sendInvitationReq.DomainID, - DomainUserID: fmt.Sprintf("%s_%s", tc.sendInvitationReq.DomainID, tc.sendInvitationReq.InviteeUserID), - } - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := svc.On("SendInvitation", mock.Anything, tc.session, tc.svcReq).Return(domains.Invitation{}, tc.svcErr) - err := mgsdk.SendInvitation(context.Background(), tc.sendInvitationReq, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "SendInvitation", mock.Anything, tc.session, tc.svcReq) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestListInvitation(t *testing.T) { - is, svc, auth := setupDomains() - defer is.Close() - - conf := sdk.Config{ - DomainsURL: is.URL, - } - mgsdk := sdk.NewSDK(conf) - - cases := []struct { - desc string - token string - session smqauthn.Session - pageMeta sdk.PageMetadata - svcReq domains.InvitationPageMeta - svcRes domains.InvitationPage - svcErr error - authenticateErr error - response sdk.InvitationPage - err error - }{ - { - desc: "list invitations successfully", - token: validToken, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - }, - svcReq: domains.InvitationPageMeta{ - Offset: 0, - Limit: 10, - }, - svcRes: domains.InvitationPage{ - Total: 1, - Invitations: []domains.Invitation{invitation}, - }, - svcErr: nil, - response: sdk.InvitationPage{ - Total: 1, - Invitations: []sdk.Invitation{sdkInvitation}, - }, - err: nil, - }, - { - desc: "list invitations with invalid token", - token: invalidToken, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - }, - svcReq: domains.InvitationPageMeta{ - Offset: 0, - Limit: 10, - }, - svcRes: domains.InvitationPage{}, - authenticateErr: svcerr.ErrAuthentication, - response: sdk.InvitationPage{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "list invitations with empty token", - token: "", - pageMeta: sdk.PageMetadata{}, - svcRes: domains.InvitationPage{}, - svcErr: nil, - response: sdk.InvitationPage{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "list invitations with limit greater than max limit", - token: validToken, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 101, - }, - svcReq: domains.InvitationPageMeta{}, - svcRes: domains.InvitationPage{}, - svcErr: nil, - response: sdk.InvitationPage{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrLimitSize, http.StatusBadRequest), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == valid { - tc.session = smqauthn.Session{DomainUserID: validID, UserID: validID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := svc.On("ListInvitations", mock.Anything, tc.session, tc.svcReq).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.Invitations(context.Background(), tc.pageMeta, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "ListInvitations", mock.Anything, tc.session, tc.svcReq) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestAcceptInvitation(t *testing.T) { - is, svc, auth := setupDomains() - defer is.Close() - - conf := sdk.Config{ - DomainsURL: is.URL, - } - mgsdk := sdk.NewSDK(conf) - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - authenticateErr error - svcErr error - err error - }{ - { - desc: "accept invitation successfully", - token: validToken, - domainID: invitation.DomainID, - svcErr: nil, - err: nil, - }, - { - desc: "accept invitation with invalid token", - token: invalidToken, - domainID: invitation.DomainID, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "accept invitation with empty token", - token: "", - domainID: invitation.DomainID, - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "accept invitation with invalid domainID", - token: validToken, - domainID: wrongID, - svcErr: svcerr.ErrNotFound, - err: errors.NewSDKErrorWithStatus(svcerr.ErrNotFound, http.StatusNotFound), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == valid { - tc.session = smqauthn.Session{DomainUserID: validID, UserID: validID, DomainID: validID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := svc.On("AcceptInvitation", mock.Anything, tc.session, tc.domainID).Return(domains.Invitation{}, tc.svcErr) - err := mgsdk.AcceptInvitation(context.Background(), tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "AcceptInvitation", mock.Anything, tc.session, tc.domainID) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestRejectInvitation(t *testing.T) { - is, svc, auth := setupDomains() - defer is.Close() - - conf := sdk.Config{ - DomainsURL: is.URL, - } - mgsdk := sdk.NewSDK(conf) - - cases := []struct { - desc string - token string - session smqauthn.Session - domainID string - authenticateErr error - svcErr error - err error - }{ - { - desc: "reject invitation successfully", - token: validToken, - domainID: invitation.DomainID, - svcErr: nil, - err: nil, - }, - { - desc: "reject invitation with invalid token", - token: invalidToken, - domainID: invitation.DomainID, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "reject invitation with empty token", - token: "", - domainID: invitation.DomainID, - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "reject invitation with invalid domainID", - token: validToken, - domainID: wrongID, - svcErr: svcerr.ErrNotFound, - err: errors.NewSDKErrorWithStatus(svcerr.ErrNotFound, http.StatusNotFound), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == valid { - tc.session = smqauthn.Session{DomainUserID: validID, UserID: validID, DomainID: validID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := svc.On("RejectInvitation", mock.Anything, tc.session, tc.domainID).Return(domains.Invitation{}, tc.svcErr) - err := mgsdk.RejectInvitation(context.Background(), tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RejectInvitation", mock.Anything, tc.session, tc.domainID) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestDeleteInvitation(t *testing.T) { - is, svc, auth := setupDomains() - defer is.Close() - - conf := sdk.Config{ - DomainsURL: is.URL, - } - mgsdk := sdk.NewSDK(conf) - - cases := []struct { - desc string - token string - session smqauthn.Session - inviteeUserID string - domainID string - authenticateErr error - svcErr error - err error - }{ - { - desc: "delete invitation successfully", - token: validToken, - inviteeUserID: invitation.InviteeUserID, - domainID: invitation.DomainID, - svcErr: nil, - err: nil, - }, - { - desc: "delete invitation with invalid token", - token: invalidToken, - inviteeUserID: invitation.InviteeUserID, - domainID: invitation.DomainID, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "delete invitation with empty token", - token: "", - inviteeUserID: invitation.InviteeUserID, - domainID: invitation.DomainID, - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "delete invitation with empty domainID", - token: validToken, - inviteeUserID: invitation.InviteeUserID, - domainID: "", - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingDomainID, http.StatusBadRequest), - }, - { - desc: "delete invitation with invalid domainID", - token: validToken, - inviteeUserID: invitation.InviteeUserID, - domainID: wrongID, - svcErr: svcerr.ErrNotFound, - err: errors.NewSDKErrorWithStatus(svcerr.ErrNotFound, http.StatusNotFound), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == valid { - tc.session = smqauthn.Session{UserID: tc.inviteeUserID, DomainID: tc.domainID, DomainUserID: fmt.Sprintf("%s_%s", tc.domainID, tc.inviteeUserID)} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := svc.On("DeleteInvitation", mock.Anything, tc.session, tc.inviteeUserID, tc.domainID).Return(tc.svcErr) - err := mgsdk.DeleteInvitation(context.Background(), tc.inviteeUserID, tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "DeleteInvitation", mock.Anything, tc.session, tc.inviteeUserID, tc.domainID) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func generateTestInvitation(t *testing.T) sdk.Invitation { - createdAt, err := time.Parse(time.RFC3339, "2024-01-01T00:00:00Z") - assert.Nil(t, err, fmt.Sprintf("Unexpected error parsing time: %v", err)) - return sdk.Invitation{ - InvitedBy: testsutil.GenerateUUID(t), - InviteeUserID: testsutil.GenerateUUID(t), - DomainID: testsutil.GenerateUUID(t), - RoleID: testsutil.GenerateUUID(t), - RoleName: "admin", - Actions: []string{"read", "update"}, - CreatedAt: createdAt, - UpdatedAt: createdAt, - } -} diff --git a/pkg/sdk/journal_test.go b/pkg/sdk/journal_test.go deleted file mode 100644 index eb58082cf..000000000 --- a/pkg/sdk/journal_test.go +++ /dev/null @@ -1,361 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package sdk_test - -import ( - "context" - "fmt" - "net/http" - "net/http/httptest" - "testing" - "time" - - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/journal" - "github.com/absmach/magistrala/journal/api" - "github.com/absmach/magistrala/journal/mocks" - mglog "github.com/absmach/magistrala/logger" - smqauthn "github.com/absmach/magistrala/pkg/authn" - authnmocks "github.com/absmach/magistrala/pkg/authn/mocks" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - sdk "github.com/absmach/magistrala/pkg/sdk" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" -) - -func setupJournal() (*httptest.Server, *mocks.Service, *authnmocks.Authentication) { - svc := new(mocks.Service) - authn := new(authnmocks.Authentication) - logger := mglog.NewMock() - am := smqauthn.NewAuthNMiddleware(authn, smqauthn.WithAllowUnverifiedUser(true)) - mux := api.MakeHandler(svc, am, logger, "journal-log", "test") - - return httptest.NewServer(mux), svc, authn -} - -func TestRetrieveJournal(t *testing.T) { - js, svc, authn := setupJournal() - defer js.Close() - - testJournal := generateTestJournal(t) - validEntityType := "group" - - sdkConf := sdk.Config{ - JournalURL: js.URL, - } - - mgsdk := sdk.NewSDK(sdkConf) - - cases := []struct { - desc string - token string - session smqauthn.Session - entityType string - entityID string - domainID string - pageMeta sdk.PageMetadata - svcReq journal.Page - svcRes journal.JournalsPage - svcErr error - authnErr error - response sdk.JournalsPage - err error - }{ - { - desc: "retrieve user journal successfully", - token: validToken, - entityType: "user", - entityID: validID, - domainID: domainID, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - }, - svcReq: journal.Page{ - Offset: 0, - Limit: 10, - EntityID: validID, - EntityType: journal.UserEntity, - Direction: "desc", - }, - svcRes: journal.JournalsPage{ - Total: 1, - Journals: []journal.Journal{convertJournal(testJournal)}, - }, - svcErr: nil, - response: sdk.JournalsPage{ - Total: 1, - Journals: []sdk.Journal{testJournal}, - }, - err: nil, - }, - { - desc: "retrieve channel journal successfully", - token: validToken, - entityType: "channel", - entityID: validID, - domainID: domainID, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - }, - svcReq: journal.Page{ - Offset: 0, - Limit: 10, - EntityID: validID, - EntityType: journal.ChannelEntity, - Direction: "desc", - }, - svcRes: journal.JournalsPage{ - Total: 1, - Journals: []journal.Journal{convertJournal(testJournal)}, - }, - svcErr: nil, - response: sdk.JournalsPage{ - Total: 1, - Journals: []sdk.Journal{testJournal}, - }, - err: nil, - }, - { - desc: "retrieve group journal successfully", - token: validToken, - entityType: "group", - entityID: validID, - domainID: domainID, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - }, - svcReq: journal.Page{ - Offset: 0, - Limit: 10, - EntityID: validID, - EntityType: journal.GroupEntity, - Direction: "desc", - }, - svcRes: journal.JournalsPage{ - Total: 1, - Journals: []journal.Journal{convertJournal(testJournal)}, - }, - svcErr: nil, - response: sdk.JournalsPage{ - Total: 1, - Journals: []sdk.Journal{testJournal}, - }, - err: nil, - }, - { - desc: "retrieve client journal successfully", - token: validToken, - entityType: "client", - entityID: validID, - domainID: domainID, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - }, - svcReq: journal.Page{ - Offset: 0, - Limit: 10, - EntityID: validID, - EntityType: journal.ClientEntity, - Direction: "desc", - }, - svcRes: journal.JournalsPage{ - Total: 1, - Journals: []journal.Journal{convertJournal(testJournal)}, - }, - svcErr: nil, - response: sdk.JournalsPage{ - Total: 1, - Journals: []sdk.Journal{testJournal}, - }, - err: nil, - }, - { - desc: "retrieve journal with invalid token", - token: invalidToken, - entityType: validEntityType, - entityID: validID, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - }, - svcReq: journal.Page{ - Offset: 0, - Limit: 10, - EntityID: validID, - EntityType: journal.GroupEntity, - Direction: "desc", - }, - svcRes: journal.JournalsPage{}, - authnErr: svcerr.ErrAuthentication, - response: sdk.JournalsPage{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "retrieve journal with empty token", - token: "", - entityType: validEntityType, - entityID: validID, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - }, - svcReq: journal.Page{}, - svcRes: journal.JournalsPage{}, - svcErr: nil, - response: sdk.JournalsPage{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "retrieve journal with invalid entity type", - token: validToken, - entityType: "invalid", - entityID: validID, - domainID: domainID, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - }, - svcReq: journal.Page{}, - svcRes: journal.JournalsPage{}, - svcErr: nil, - response: sdk.JournalsPage{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrInvalidEntityType, http.StatusBadRequest), - }, - { - desc: "retrieve journal with empty entity ID", - token: validToken, - entityType: validEntityType, - entityID: "", - domainID: domainID, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - }, - svcReq: journal.Page{}, - svcRes: journal.JournalsPage{}, - svcErr: nil, - response: sdk.JournalsPage{}, - err: errors.NewSDKError(apiutil.ErrMissingID), - }, - { - desc: "retrieve journal with empty entity type", - token: validToken, - entityType: "", - entityID: validID, - domainID: domainID, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - }, - svcReq: journal.Page{}, - svcRes: journal.JournalsPage{}, - svcErr: nil, - response: sdk.JournalsPage{}, - err: errors.NewSDKError(apiutil.ErrMissingEntityType), - }, - { - desc: "retrieve journal with limit greater than default", - token: validToken, - entityType: validEntityType, - entityID: validID, - domainID: domainID, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 1000, - }, - svcReq: journal.Page{}, - svcRes: journal.JournalsPage{}, - svcErr: nil, - response: sdk.JournalsPage{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrLimitSize, http.StatusBadRequest), - }, - { - desc: "retrieve journal with invalid page metadata", - token: validToken, - entityType: validEntityType, - entityID: validID, - domainID: domainID, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - Metadata: map[string]any{ - "key": make(chan int), - }, - }, - svcReq: journal.Page{}, - svcRes: journal.JournalsPage{}, - svcErr: nil, - response: sdk.JournalsPage{}, - err: errors.NewSDKError(fmt.Errorf("json: unsupported type: chan int")), - }, - { - desc: "retrieve journal with response that cannot be unmarshalled", - token: validToken, - entityType: validEntityType, - entityID: validID, - domainID: domainID, - pageMeta: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - }, - svcReq: journal.Page{ - Offset: 0, - Limit: 10, - EntityID: validID, - EntityType: journal.GroupEntity, - Direction: "desc", - }, - svcRes: journal.JournalsPage{ - Total: 1, - Journals: []journal.Journal{{ - ID: validID, - Operation: "create", - OccurredAt: time.Now(), - Attributes: validMetadata, - Metadata: map[string]any{ - "key": make(chan int), - }, - }}, - }, - svcErr: nil, - response: sdk.JournalsPage{}, - err: errors.NewSDKError(fmt.Errorf("unexpected end of JSON input")), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: fmt.Sprintf("%s_%s", domainID, validID), UserID: validID, DomainID: domainID} - } - authCall := authn.On("Authenticate", mock.Anything, mock.Anything).Return(tc.session, tc.authnErr) - svcCall := svc.On("RetrieveAll", mock.Anything, tc.session, tc.svcReq).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.Journal(context.Background(), tc.entityType, tc.entityID, tc.domainID, tc.pageMeta, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RetrieveAll", mock.Anything, tc.session, tc.svcReq) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func generateTestJournal(t *testing.T) sdk.Journal { - occuredAt, err := time.Parse(time.RFC3339, "2024-01-01T00:00:00Z") - assert.Nil(t, err, fmt.Sprintf("Unexpected error parsing time: %v", err)) - return sdk.Journal{ - ID: validID, - Operation: "create", - OccurredAt: occuredAt, - Attributes: validMetadata, - Metadata: validMetadata, - } -} diff --git a/pkg/sdk/message_test.go b/pkg/sdk/message_test.go deleted file mode 100644 index fd5d449bd..000000000 --- a/pkg/sdk/message_test.go +++ /dev/null @@ -1,180 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package sdk_test - -import ( - "context" - "encoding/json" - "fmt" - "io" - "net/http" - "net/http/httptest" - "testing" - - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/pkg/errors" - sdk "github.com/absmach/magistrala/pkg/sdk" - "github.com/stretchr/testify/assert" -) - -type publishReq struct { - Topic string `json:"topic"` - Payload []byte `json:"payload"` - QoS byte `json:"qos"` - Retain bool `json:"retain"` -} - -func setupFluxMQ(secret string, expectedTopic ...string) *httptest.Server { - mux := http.NewServeMux() - - mux.HandleFunc("POST /publish", func(w http.ResponseWriter, r *http.Request) { - username := r.Header.Get("X-FluxMQ-Username") - auth := r.Header.Get("Authorization") - if username == "" || auth == "" || auth != "Bearer "+secret { - http.Error(w, "unauthorized", http.StatusUnauthorized) - return - } - - body, err := io.ReadAll(r.Body) - if err != nil { - http.Error(w, "bad request", http.StatusBadRequest) - return - } - defer r.Body.Close() - - var req publishReq - if err := json.Unmarshal(body, &req); err != nil { - http.Error(w, "invalid json", http.StatusBadRequest) - return - } - - if req.Topic == "" { - http.Error(w, "empty topic", http.StatusBadRequest) - return - } - if len(expectedTopic) > 0 && req.Topic != expectedTopic[0] { - http.Error(w, fmt.Sprintf("unexpected topic: %s", req.Topic), http.StatusBadRequest) - return - } - - w.Header().Set("Content-Type", "application/json") - w.WriteHeader(http.StatusOK) - fmt.Fprint(w, `{"status":"ok"}`) - }) - - mux.HandleFunc("GET /health", func(w http.ResponseWriter, _ *http.Request) { - w.Header().Set("Content-Type", "application/json") - w.WriteHeader(http.StatusOK) - fmt.Fprint(w, `{"status":"healthy"}`) - }) - - return httptest.NewServer(mux) -} - -func TestSendMessage(t *testing.T) { - clientSecret := "validSecret" - - cases := []struct { - desc string - topic string - domainID string - wantTopic string - msg string - secret string - err errors.SDKError - }{ - { - desc: "publish message successfully", - topic: "channelID", - domainID: "domainID", - wantTopic: "m/domainID/c/channelID", - msg: `[{"n":"current","t":-1,"v":1.6}]`, - secret: clientSecret, - err: nil, - }, - { - desc: "publish message with subtopic", - topic: "channelID/sub/topic", - domainID: "domainID", - wantTopic: "m/domainID/c/channelID/sub/topic", - msg: `[{"n":"current","t":-1,"v":1.6}]`, - secret: clientSecret, - err: nil, - }, - { - desc: "publish message with invalid secret", - topic: "channelID", - domainID: "domainID", - wantTopic: "m/domainID/c/channelID", - msg: `[{"n":"current","t":-1,"v":1.6}]`, - secret: "invalid", - err: errors.NewSDKErrorWithStatus(errors.Wrap(errors.New(""), errors.New("")), http.StatusUnauthorized), - }, - { - desc: "publish message with empty secret", - topic: "channelID", - domainID: "domainID", - wantTopic: "m/domainID/c/channelID", - msg: `[{"n":"current","t":-1,"v":1.6}]`, - secret: "", - err: errors.NewSDKErrorWithStatus(errors.Wrap(errors.New(""), errors.New("")), http.StatusUnauthorized), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - ts := setupFluxMQ(clientSecret, tc.wantTopic) - defer ts.Close() - - sdkConf := sdk.Config{ - HTTPAdapterURL: ts.URL, - MsgContentType: "application/senml+json", - TLSVerification: false, - } - mgsdk := sdk.NewSDK(sdkConf) - - err := mgsdk.SendMessage(context.Background(), tc.domainID, tc.topic, tc.msg, tc.secret) - if tc.err != nil { - assert.NotNil(t, err, fmt.Sprintf("%s: expected error, got nil", tc.desc)) - } else { - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error: %v", tc.desc, err)) - } - }) - } -} - -func TestSetContentType(t *testing.T) { - sdkConf := sdk.Config{ - MsgContentType: "application/senml+json", - TLSVerification: false, - } - mgsdk := sdk.NewSDK(sdkConf) - - cases := []struct { - desc string - cType sdk.ContentType - err errors.SDKError - }{ - { - desc: "set senml+json content type", - cType: "application/senml+json", - err: nil, - }, - { - desc: "set json content type", - cType: "application/json", - err: nil, - }, - { - desc: "set invalid content type", - cType: "invalid", - err: errors.NewSDKError(apiutil.ErrUnsupportedContentType), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - err := mgsdk.SetContentType(tc.cType) - assert.Equal(t, tc.err, err, fmt.Sprintf("%s: expected error %s, got %s", tc.desc, tc.err, err)) - }) - } -} diff --git a/pkg/sdk/messages_test.go b/pkg/sdk/messages_test.go deleted file mode 100644 index cbaf0f1b7..000000000 --- a/pkg/sdk/messages_test.go +++ /dev/null @@ -1,240 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package sdk_test - -import ( - "context" - "net/http" - "net/http/httptest" - "testing" - - grpcChannelsV1 "github.com/absmach/magistrala/api/grpc/channels/v1" - apiutil "github.com/absmach/magistrala/api/http/util" - chmocks "github.com/absmach/magistrala/channels/mocks" - climocks "github.com/absmach/magistrala/clients/mocks" - smqauthn "github.com/absmach/magistrala/pkg/authn" - authnmocks "github.com/absmach/magistrala/pkg/authn/mocks" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/pkg/sdk" - "github.com/absmach/magistrala/pkg/transformers/senml" - "github.com/absmach/magistrala/readers" - readersapi "github.com/absmach/magistrala/readers/api/http" - readersmocks "github.com/absmach/magistrala/readers/mocks" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" -) - -func setupReaders() (*httptest.Server, *authnmocks.Authentication, *readersmocks.MessageRepository) { - repo := new(readersmocks.MessageRepository) - authn := new(authnmocks.Authentication) - clientsGRPCClient = new(climocks.ClientsServiceClient) - channelsGRPCClient = new(chmocks.ChannelsServiceClient) - - mux := readersapi.MakeHandler(repo, authn, clientsGRPCClient, channelsGRPCClient, "test", "") - return httptest.NewServer(mux), authn, repo -} - -func TestReadMessages(t *testing.T) { - ts, authn, repo := setupReaders() - defer ts.Close() - - channelID := "channelID" - msgValue := 1.6 - boolVal := true - msg := senml.Message{ - Name: "current", - Time: 1720000000, - Value: &msgValue, - Publisher: validID, - } - invalidMsg := "[{\"n\":\"current\",\"t\":-1,\"v\":1.6}]" - - sdkConf := sdk.Config{ - ReaderURL: ts.URL, - } - - mgsdk := sdk.NewSDK(sdkConf) - - cases := []struct { - desc string - token string - chanName string - domainID string - messagePageMeta sdk.MessagePageMetadata - authzErr error - authnErr error - repoRes readers.MessagesPage - repoErr error - response sdk.MessagesPage - err errors.SDKError - }{ - { - desc: "read messages successfully", - token: validToken, - chanName: channelID, - domainID: validID, - messagePageMeta: sdk.MessagePageMetadata{ - PageMetadata: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - Level: 0, - }, - Publisher: validID, - BoolValue: &boolVal, - }, - repoRes: readers.MessagesPage{ - Total: 1, - Messages: []readers.Message{msg}, - }, - repoErr: nil, - response: sdk.MessagesPage{ - PageRes: sdk.PageRes{ - Total: 1, - }, - Messages: []senml.Message{msg}, - }, - err: nil, - }, - { - desc: "read messages successfully with subtopic", - token: validToken, - chanName: channelID + "/subtopic", - domainID: validID, - messagePageMeta: sdk.MessagePageMetadata{ - PageMetadata: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - }, - Publisher: validID, - }, - repoRes: readers.MessagesPage{ - Total: 1, - Messages: []readers.Message{msg}, - }, - repoErr: nil, - response: sdk.MessagesPage{ - PageRes: sdk.PageRes{ - Total: 1, - }, - Messages: []senml.Message{msg}, - }, - err: nil, - }, - { - desc: "read messages with invalid token", - token: invalidToken, - chanName: channelID, - domainID: validID, - messagePageMeta: sdk.MessagePageMetadata{ - PageMetadata: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - }, - Subtopic: "subtopic", - Publisher: validID, - }, - authzErr: svcerr.ErrAuthorization, - repoRes: readers.MessagesPage{}, - response: sdk.MessagesPage{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthorization, http.StatusForbidden), - }, - { - desc: "read messages with empty token", - token: "", - chanName: channelID, - domainID: validID, - messagePageMeta: sdk.MessagePageMetadata{ - PageMetadata: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - }, - Subtopic: "subtopic", - Publisher: validID, - }, - authnErr: svcerr.ErrAuthentication, - repoRes: readers.MessagesPage{}, - response: sdk.MessagesPage{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "read messages with empty channel ID", - token: validToken, - chanName: "", - domainID: validID, - messagePageMeta: sdk.MessagePageMetadata{ - PageMetadata: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - }, - Subtopic: "subtopic", - Publisher: validID, - }, - repoRes: readers.MessagesPage{}, - repoErr: nil, - response: sdk.MessagesPage{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "read messages with invalid message page metadata", - token: validToken, - chanName: channelID, - domainID: validID, - messagePageMeta: sdk.MessagePageMetadata{ - PageMetadata: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - Metadata: map[string]any{ - "key": make(chan int), - }, - }, - Subtopic: "subtopic", - Publisher: validID, - }, - repoRes: readers.MessagesPage{}, - repoErr: nil, - response: sdk.MessagesPage{}, - err: errors.NewSDKError(errors.New("json: unsupported type: chan int")), - }, - { - desc: "read messages with response that cannot be unmarshalled", - token: validToken, - chanName: channelID, - domainID: validID, - messagePageMeta: sdk.MessagePageMetadata{ - PageMetadata: sdk.PageMetadata{ - Offset: 0, - Limit: 10, - }, - Subtopic: "subtopic", - Publisher: validID, - }, - repoRes: readers.MessagesPage{ - Total: 1, - Messages: []readers.Message{invalidMsg}, - }, - repoErr: nil, - response: sdk.MessagesPage{}, - err: errors.NewSDKError(errors.New("json: cannot unmarshal string into Go struct field MessagesPage.messages of type senml.Message")), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - authCall1 := authn.On("Authenticate", mock.Anything, tc.token).Return(smqauthn.Session{UserID: validID}, tc.authnErr) - authzCall := channelsGRPCClient.On("Authorize", mock.Anything, mock.Anything).Return(&grpcChannelsV1.AuthzRes{Authorized: true}, tc.authzErr) - repoCall := repo.On("ReadAll", channelID, mock.Anything).Return(tc.repoRes, tc.repoErr) - response, err := mgsdk.ReadMessages(context.Background(), tc.messagePageMeta, tc.chanName, tc.domainID, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, response) - if tc.err == nil { - ok := repoCall.Parent.AssertCalled(t, "ReadAll", channelID, mock.Anything) - assert.True(t, ok) - } - authCall1.Unset() - authzCall.Unset() - repoCall.Unset() - }) - } -} diff --git a/pkg/sdk/reports.go b/pkg/sdk/reports.go index 4f7b262a3..5b52dbdb4 100644 --- a/pkg/sdk/reports.go +++ b/pkg/sdk/reports.go @@ -12,7 +12,6 @@ import ( "time" "github.com/absmach/magistrala/pkg/errors" - "github.com/absmach/magistrala/pkg/roles" ) const ( @@ -22,21 +21,20 @@ const ( // ReportConfig represents a report configuration. type ReportConfig struct { - ID string `json:"id,omitempty"` - Name string `json:"name,omitempty"` - Description string `json:"description,omitempty"` - DomainID string `json:"domain_id,omitempty"` - Schedule any `json:"schedule,omitempty"` - Config any `json:"config,omitempty"` - Email any `json:"email,omitempty"` - Metrics any `json:"metrics,omitempty"` - ReportTemplate ReportTemplate `json:"report_template,omitempty"` - Status string `json:"status,omitempty"` - CreatedAt time.Time `json:"created_at,omitempty"` - CreatedBy string `json:"created_by,omitempty"` - UpdatedAt time.Time `json:"updated_at,omitempty"` - UpdatedBy string `json:"updated_by,omitempty"` - Roles []roles.MemberRoleActions `json:"roles,omitempty"` + ID string `json:"id,omitempty"` + Name string `json:"name,omitempty"` + Description string `json:"description,omitempty"` + DomainID string `json:"domain_id,omitempty"` + Schedule any `json:"schedule,omitempty"` + Config any `json:"config,omitempty"` + Email any `json:"email,omitempty"` + Metrics any `json:"metrics,omitempty"` + ReportTemplate ReportTemplate `json:"report_template,omitempty"` + Status string `json:"status,omitempty"` + CreatedAt time.Time `json:"created_at,omitempty"` + CreatedBy string `json:"created_by,omitempty"` + UpdatedAt time.Time `json:"updated_at,omitempty"` + UpdatedBy string `json:"updated_by,omitempty"` } type ReportTemplate any diff --git a/pkg/sdk/reports_test.go b/pkg/sdk/reports_test.go deleted file mode 100644 index d0e9462fd..000000000 --- a/pkg/sdk/reports_test.go +++ /dev/null @@ -1,867 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package sdk_test - -import ( - "context" - "errors" - "net/http/httptest" - "testing" - "time" - - mglog "github.com/absmach/magistrala/logger" - smqauthn "github.com/absmach/magistrala/pkg/authn" - authnmocks "github.com/absmach/magistrala/pkg/authn/mocks" - pkgSch "github.com/absmach/magistrala/pkg/schedule" - "github.com/absmach/magistrala/pkg/sdk" - "github.com/absmach/magistrala/reports" - "github.com/absmach/magistrala/reports/api" - rmocks "github.com/absmach/magistrala/reports/mocks" - "github.com/go-chi/chi/v5" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" -) - -const ( - reportConfigID = "report-config-1" - reportName = "daily-report" - reportUpdatedName = "updated daily-report" - reportDescription = "Daily temperature report" - reportUpdatedDesc = "updated Daily temperature report" - validTemplate = ` - - - {{$.Title}} - - - -
-

{{$.Title}}

-

Generated on: {{$.GeneratedDate}}

-
-
-

Messages

- {{range .Messages}} -
-

Time: {{formatTime .Time}}

-

Value: {{formatValue .}}

-
- {{end}} -
- -` -) - -var ( - now = time.Now().UTC().Truncate(time.Minute) - future = now.Add(1 * time.Hour) - schedule = pkgSch.Schedule{ - StartDateTime: future, - Recurring: pkgSch.Daily, - RecurringPeriod: 1, - Time: future, - } - metrics = []reports.ReqMetric{ - { - ChannelID: "channel1", - ClientIDs: []string{"client1"}, - Name: "metric_name", - }, - } - config = reports.MetricConfig{ - From: "now()-1h", - To: "now()", - Title: "test_title", - Aggregation: reports.AggConfig{AggType: reports.AggregationAVG, Interval: "1h"}, - } - email = reports.EmailSetting{ - To: []string{"test@example.com"}, - Subject: "Test Report", - } - - testReportConfig = sdk.ReportConfig{ - ID: reportConfigID, - Name: reportName, - Description: reportDescription, - DomainID: domainID, - Status: "enabled", - Schedule: schedule, - Metrics: metrics, - Config: &config, - Email: &email, - } -) - -func setupReports() (*httptest.Server, *rmocks.Service, *authnmocks.Authentication) { - rsvc := new(rmocks.Service) - log := mglog.NewMock() - authn := new(authnmocks.Authentication) - am := smqauthn.NewAuthNMiddleware(authn, smqauthn.WithAllowUnverifiedUser(true)) - mux := chi.NewRouter() - _ = api.MakeHandler(rsvc, am, mux, log, "") - return httptest.NewServer(mux), rsvc, authn -} - -func TestAddReportConfig(t *testing.T) { - rs, rsvc, auth := setupReports() - defer rs.Close() - - conf := sdk.Config{ - ReportsURL: rs.URL, - } - mgsdk := sdk.NewSDK(conf) - - svcCfg := reports.ReportConfig{ - ID: reportConfigID, - Name: reportName, - Description: reportDescription, - DomainID: domainID, - Status: reports.EnabledStatus, - Schedule: schedule, - Metrics: []reports.ReqMetric{ - { - ChannelID: "channel1", - ClientIDs: []string{"client1"}, - Name: "metric_name", - }, - }, - Config: &reports.MetricConfig{ - From: "now()-1h", - To: "now()", - Title: "test_title", - Aggregation: reports.AggConfig{AggType: reports.AggregationAVG, Interval: "1h"}, - }, - Email: &reports.EmailSetting{ - To: []string{"test@example.com"}, - Subject: "Test Report", - }, - } - - cases := []struct { - desc string - cfg sdk.ReportConfig - token string - session smqauthn.Session - svcRes reports.ReportConfig - svcErr error - authenticateErr error - wantErr bool - }{ - { - desc: "add report config successfully", - cfg: testReportConfig, - token: validToken, - svcRes: svcCfg, - }, - { - desc: "add report config with empty token", - cfg: sdk.ReportConfig{Name: "daily-report"}, - token: "", - wantErr: true, - svcErr: errors.New("missing or invalid bearer user token"), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := rsvc.On("AddReportConfig", mock.Anything, tc.session, mock.Anything).Return(tc.svcRes, tc.svcErr) - result, err := mgsdk.AddReportConfig(context.Background(), tc.cfg, domainID, tc.token) - assert.Equal(t, tc.wantErr, err != nil) - if !tc.wantErr { - assert.NotEmpty(t, result.ID) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestViewReportConfig(t *testing.T) { - rs, rsvc, auth := setupReports() - defer rs.Close() - - conf := sdk.Config{ - ReportsURL: rs.URL, - } - mgsdk := sdk.NewSDK(conf) - - svcCfg := reports.ReportConfig{ - ID: reportConfigID, - Name: reportName, - Description: reportDescription, - DomainID: domainID, - Status: reports.EnabledStatus, - Metrics: metrics, - Config: &config, - Email: &email, - } - - cases := []struct { - desc string - id string - token string - session smqauthn.Session - svcRes reports.ReportConfig - svcErr error - authenticateErr error - wantErr bool - }{ - { - desc: "view report config successfully", - id: reportConfigID, - token: validToken, - svcRes: svcCfg, - }, - { - desc: "view report config with empty token", - id: reportConfigID, - token: "", - wantErr: true, - }, - { - desc: "view non-existent report config", - id: "non-existent", - token: validToken, - svcErr: errors.New("not found"), - wantErr: true, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := rsvc.On("ViewReportConfig", mock.Anything, tc.session, tc.id, mock.Anything).Return(tc.svcRes, tc.svcErr) - result, err := mgsdk.ViewReportConfig(context.Background(), tc.id, domainID, tc.token) - assert.Equal(t, tc.wantErr, err != nil) - if !tc.wantErr { - assert.NotEmpty(t, result.ID) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestUpdateReportConfig(t *testing.T) { - rs, rsvc, auth := setupReports() - defer rs.Close() - - conf := sdk.Config{ - ReportsURL: rs.URL, - } - mgsdk := sdk.NewSDK(conf) - - updatedConfig := testReportConfig - updatedConfig.Name = reportUpdatedName - updatedConfig.Description = reportUpdatedDesc - - svcCfg := reports.ReportConfig{ - ID: reportConfigID, - Name: reportUpdatedName, - Description: reportUpdatedDesc, - DomainID: domainID, - Status: reports.EnabledStatus, - Metrics: metrics, - Config: &config, - Email: &email, - } - - cases := []struct { - desc string - cfg sdk.ReportConfig - token string - session smqauthn.Session - svcRes reports.ReportConfig - svcErr error - authenticateErr error - wantErr bool - }{ - { - desc: "update report config successfully", - cfg: updatedConfig, - token: validToken, - svcRes: svcCfg, - }, - { - desc: "update report config with empty token", - cfg: sdk.ReportConfig{ID: reportConfigID, Name: "updated-report"}, - token: "", - wantErr: true, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := rsvc.On("UpdateReportConfig", mock.Anything, tc.session, mock.Anything).Return(tc.svcRes, tc.svcErr) - result, err := mgsdk.UpdateReportConfig(context.Background(), tc.cfg, domainID, tc.token) - assert.Equal(t, tc.wantErr, err != nil) - if !tc.wantErr { - assert.NotEmpty(t, result.ID) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestUpdateReportSchedule(t *testing.T) { - rs, rsvc, auth := setupReports() - defer rs.Close() - - conf := sdk.Config{ - ReportsURL: rs.URL, - } - mgsdk := sdk.NewSDK(conf) - - svcCfg := reports.ReportConfig{ - ID: reportConfigID, - Name: reportName, - Status: reports.EnabledStatus, - } - - cases := []struct { - desc string - cfg sdk.ReportConfig - token string - session smqauthn.Session - svcRes reports.ReportConfig - svcErr error - authenticateErr error - wantErr bool - }{ - { - desc: "update report schedule successfully", - cfg: sdk.ReportConfig{ID: reportConfigID, Schedule: map[string]any{"cron": "0 9 * * *"}}, - token: validToken, - svcRes: svcCfg, - }, - { - desc: "update report schedule with empty token", - cfg: sdk.ReportConfig{ID: reportConfigID}, - token: "", - wantErr: true, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := rsvc.On("UpdateReportSchedule", mock.Anything, tc.session, mock.Anything).Return(tc.svcRes, tc.svcErr) - result, err := mgsdk.UpdateReportSchedule(context.Background(), tc.cfg, domainID, tc.token) - assert.Equal(t, tc.wantErr, err != nil) - if !tc.wantErr { - assert.NotEmpty(t, result.ID) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestRemoveReportConfig(t *testing.T) { - rs, rsvc, auth := setupReports() - defer rs.Close() - - conf := sdk.Config{ - ReportsURL: rs.URL, - } - mgsdk := sdk.NewSDK(conf) - - cases := []struct { - desc string - id string - token string - session smqauthn.Session - svcErr error - authenticateErr error - wantErr bool - }{ - { - desc: "remove report config successfully", - id: reportConfigID, - token: validToken, - }, - { - desc: "remove report config with empty token", - id: reportConfigID, - token: "", - wantErr: true, - }, - { - desc: "remove non-existent report config", - id: "non-existent", - token: validToken, - svcErr: errors.New("not found"), - wantErr: true, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := rsvc.On("RemoveReportConfig", mock.Anything, tc.session, tc.id).Return(tc.svcErr) - err := mgsdk.RemoveReportConfig(context.Background(), tc.id, domainID, tc.token) - assert.Equal(t, tc.wantErr, err != nil) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestListReportsConfig(t *testing.T) { - rs, rsvc, auth := setupReports() - defer rs.Close() - - conf := sdk.Config{ - ReportsURL: rs.URL, - } - mgsdk := sdk.NewSDK(conf) - - svcPage := reports.ReportConfigPage{} - - cases := []struct { - desc string - pm sdk.PageMetadata - token string - session smqauthn.Session - svcRes reports.ReportConfigPage - svcErr error - authenticateErr error - wantErr bool - }{ - { - desc: "list reports config successfully", - pm: sdk.PageMetadata{Offset: 0, Limit: 10}, - token: validToken, - svcRes: svcPage, - }, - { - desc: "list reports config with filters", - pm: sdk.PageMetadata{ - Limit: 10, - Name: "daily", - Status: "enabled", - Dir: "desc", - Order: "created_at", - }, - token: validToken, - svcRes: svcPage, - }, - { - desc: "list reports config with empty metadata excludes filter params", - pm: sdk.PageMetadata{}, - token: validToken, - svcRes: reports.ReportConfigPage{}, - }, - { - desc: "list reports config with empty token", - pm: sdk.PageMetadata{Limit: 10}, - token: "", - wantErr: true, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := rsvc.On("ListReportsConfig", mock.Anything, tc.session, mock.Anything).Return(tc.svcRes, tc.svcErr) - result, err := mgsdk.ListReportsConfig(context.Background(), tc.pm, domainID, tc.token) - assert.Equal(t, tc.wantErr, err != nil) - if !tc.wantErr { - assert.NotNil(t, result) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestEnableReportConfig(t *testing.T) { - rs, rsvc, auth := setupReports() - defer rs.Close() - - conf := sdk.Config{ - ReportsURL: rs.URL, - } - mgsdk := sdk.NewSDK(conf) - - svcCfg := reports.ReportConfig{ - ID: reportConfigID, - Status: reports.EnabledStatus, - } - - cases := []struct { - desc string - id string - token string - session smqauthn.Session - svcRes reports.ReportConfig - svcErr error - authenticateErr error - wantErr bool - }{ - { - desc: "enable report config successfully", - id: reportConfigID, - token: validToken, - svcRes: svcCfg, - }, - { - desc: "enable report config with empty token", - id: reportConfigID, - token: "", - wantErr: true, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := rsvc.On("EnableReportConfig", mock.Anything, tc.session, tc.id).Return(tc.svcRes, tc.svcErr) - result, err := mgsdk.EnableReportConfig(context.Background(), tc.id, domainID, tc.token) - assert.Equal(t, tc.wantErr, err != nil) - if !tc.wantErr { - assert.NotEmpty(t, result.ID) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestDisableReportConfig(t *testing.T) { - rs, rsvc, auth := setupReports() - defer rs.Close() - - conf := sdk.Config{ - ReportsURL: rs.URL, - } - mgsdk := sdk.NewSDK(conf) - - svcCfg := reports.ReportConfig{ - ID: reportConfigID, - Status: reports.DisabledStatus, - } - - cases := []struct { - desc string - id string - token string - session smqauthn.Session - svcRes reports.ReportConfig - svcErr error - authenticateErr error - wantErr bool - }{ - { - desc: "disable report config successfully", - id: reportConfigID, - token: validToken, - svcRes: svcCfg, - }, - { - desc: "disable report config with empty token", - id: reportConfigID, - token: "", - wantErr: true, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := rsvc.On("DisableReportConfig", mock.Anything, tc.session, tc.id).Return(tc.svcRes, tc.svcErr) - result, err := mgsdk.DisableReportConfig(context.Background(), tc.id, domainID, tc.token) - assert.Equal(t, tc.wantErr, err != nil) - if !tc.wantErr { - assert.NotEmpty(t, result.ID) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestUpdateReportTemplate(t *testing.T) { - rs, rsvc, auth := setupReports() - defer rs.Close() - - conf := sdk.Config{ - ReportsURL: rs.URL, - } - mgsdk := sdk.NewSDK(conf) - - cases := []struct { - desc string - cfg sdk.ReportConfig - token string - session smqauthn.Session - svcErr error - authenticateErr error - wantErr bool - }{ - { - desc: "update report template successfully", - cfg: sdk.ReportConfig{ - ID: reportConfigID, - ReportTemplate: validTemplate, - }, - token: validToken, - }, - { - desc: "update report template with empty token", - cfg: sdk.ReportConfig{ID: reportConfigID}, - token: "", - wantErr: true, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := rsvc.On("UpdateReportTemplate", mock.Anything, tc.session, mock.Anything).Return(tc.svcErr) - err := mgsdk.UpdateReportTemplate(context.Background(), tc.cfg, domainID, tc.token) - assert.Equal(t, tc.wantErr, err != nil) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestViewReportTemplate(t *testing.T) { - rs, rsvc, auth := setupReports() - defer rs.Close() - - conf := sdk.Config{ - ReportsURL: rs.URL, - } - mgsdk := sdk.NewSDK(conf) - - svcTmpl := reports.ReportTemplate(validTemplate) - - cases := []struct { - desc string - id string - token string - session smqauthn.Session - svcRes reports.ReportTemplate - svcErr error - authenticateErr error - wantErr bool - }{ - { - desc: "view report template successfully", - id: reportConfigID, - token: validToken, - svcRes: svcTmpl, - }, - { - desc: "view report template with empty token", - id: reportConfigID, - token: "", - wantErr: true, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := rsvc.On("ViewReportTemplate", mock.Anything, tc.session, tc.id).Return(tc.svcRes, tc.svcErr) - result, err := mgsdk.ViewReportTemplate(context.Background(), tc.id, domainID, tc.token) - assert.Equal(t, tc.wantErr, err != nil) - if !tc.wantErr { - assert.NotEmpty(t, result) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestDeleteReportTemplate(t *testing.T) { - rs, rsvc, auth := setupReports() - defer rs.Close() - - conf := sdk.Config{ - ReportsURL: rs.URL, - } - mgsdk := sdk.NewSDK(conf) - - cases := []struct { - desc string - id string - token string - session smqauthn.Session - svcErr error - authenticateErr error - wantErr bool - }{ - { - desc: "delete report template successfully", - id: reportConfigID, - token: validToken, - }, - { - desc: "delete report template with empty token", - id: reportConfigID, - token: "", - wantErr: true, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := rsvc.On("DeleteReportTemplate", mock.Anything, tc.session, tc.id).Return(tc.svcErr) - err := mgsdk.DeleteReportTemplate(context.Background(), tc.id, domainID, tc.token) - assert.Equal(t, tc.wantErr, err != nil) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestGenerateReport(t *testing.T) { - rs, rsvc, auth := setupReports() - defer rs.Close() - - conf := sdk.Config{ - ReportsURL: rs.URL, - } - mgsdk := sdk.NewSDK(conf) - - svcPage := reports.ReportPage{} - - config := sdk.ReportConfig{ - ID: reportConfigID, - Name: reportName, - Description: reportDescription, - DomainID: domainID, - Metrics: metrics, - Config: &config, - ReportTemplate: reports.ReportTemplate(validTemplate), - } - - cases := []struct { - desc string - cfg sdk.ReportConfig - action sdk.ReportAction - token string - session smqauthn.Session - svcRes reports.ReportPage - svcErr error - authenticateErr error - wantErr bool - }{ - { - desc: "generate report successfully", - cfg: config, - action: sdk.ViewReportAction, - token: validToken, - svcRes: svcPage, - }, - { - desc: "generate report with download action", - cfg: config, - action: sdk.DownloadReportAction, - token: validToken, - svcRes: svcPage, - }, - { - desc: "generate report with empty token", - cfg: config, - action: sdk.ViewReportAction, - token: "", - wantErr: true, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{ - DomainUserID: domainID + "_" + validID, - UserID: validID, - DomainID: domainID, - } - } - - authCall := auth.On( - "Authenticate", - mock.Anything, - tc.token, - ).Return(tc.session, tc.authenticateErr) - - svcCall := rsvc.On( - "GenerateReport", - mock.Anything, - tc.session, - mock.Anything, - mock.Anything, - ).Return(tc.svcRes, tc.svcErr) - - page, file, err := mgsdk.GenerateReport( - context.Background(), - tc.cfg, - tc.action, - domainID, - tc.token, - ) - - assert.Equal(t, tc.wantErr, err != nil) - - if !tc.wantErr { - if tc.action == sdk.DownloadReportAction { - // download should return file - assert.NotNil(t, file) - } else { - // view/email should return page - assert.Equal(t, tc.svcRes.Total, page.Total) - } - } - - svcCall.Unset() - authCall.Unset() - }) - } -} diff --git a/pkg/sdk/rules.go b/pkg/sdk/rules.go index 288f15368..89630b500 100644 --- a/pkg/sdk/rules.go +++ b/pkg/sdk/rules.go @@ -10,29 +10,27 @@ import ( "net/http" "github.com/absmach/magistrala/pkg/errors" - "github.com/absmach/magistrala/pkg/roles" ) const rulesEndpoint = "rules" // Rule represents a rule configuration. type Rule struct { - ID string `json:"id,omitempty"` - Name string `json:"name,omitempty"` - DomainID string `json:"domain,omitempty"` - Metadata Metadata `json:"metadata,omitempty"` - Tags []string `json:"tags,omitempty"` - InputChannel string `json:"input_channel,omitempty"` - InputTopic string `json:"input_topic,omitempty"` - Logic any `json:"logic,omitempty"` - Outputs any `json:"outputs,omitempty"` - Schedule any `json:"schedule,omitempty"` - Status string `json:"status,omitempty"` - CreatedAt string `json:"created_at,omitempty"` - CreatedBy string `json:"created_by,omitempty"` - UpdatedAt string `json:"updated_at,omitempty"` - UpdatedBy string `json:"updated_by,omitempty"` - Roles []roles.MemberRoleActions `json:"roles,omitempty"` + ID string `json:"id,omitempty"` + Name string `json:"name,omitempty"` + DomainID string `json:"domain,omitempty"` + Metadata Metadata `json:"metadata,omitempty"` + Tags []string `json:"tags,omitempty"` + InputChannel string `json:"input_channel,omitempty"` + InputTopic string `json:"input_topic,omitempty"` + Logic any `json:"logic,omitempty"` + Outputs any `json:"outputs,omitempty"` + Schedule any `json:"schedule,omitempty"` + Status string `json:"status,omitempty"` + CreatedAt string `json:"created_at,omitempty"` + CreatedBy string `json:"created_by,omitempty"` + UpdatedAt string `json:"updated_at,omitempty"` + UpdatedBy string `json:"updated_by,omitempty"` } type Page struct { diff --git a/pkg/sdk/rules_test.go b/pkg/sdk/rules_test.go deleted file mode 100644 index 415e78ef4..000000000 --- a/pkg/sdk/rules_test.go +++ /dev/null @@ -1,586 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package sdk_test - -import ( - "context" - "errors" - "net/http/httptest" - "testing" - - mglog "github.com/absmach/magistrala/logger" - smqauthn "github.com/absmach/magistrala/pkg/authn" - authnmocks "github.com/absmach/magistrala/pkg/authn/mocks" - "github.com/absmach/magistrala/pkg/roles" - "github.com/absmach/magistrala/pkg/sdk" - "github.com/absmach/magistrala/re" - "github.com/absmach/magistrala/re/api" - remocks "github.com/absmach/magistrala/re/mocks" - "github.com/go-chi/chi/v5" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" -) - -const ruleID = "rule-1" - -var testRule = sdk.Rule{ - ID: ruleID, - Name: "temperature-rule", - InputChannel: "chan-1", - InputTopic: "sensors/temperature", - Status: "enabled", - Tags: []string{"temperature", "alerts"}, -} - -func setupRules() (*httptest.Server, *remocks.Service, *authnmocks.Authentication) { - rsvc := new(remocks.Service) - log := mglog.NewMock() - authn := new(authnmocks.Authentication) - am := smqauthn.NewAuthNMiddleware(authn, smqauthn.WithAllowUnverifiedUser(true)) - mux := chi.NewRouter() - _ = api.MakeHandler(rsvc, am, mux, log, "") - return httptest.NewServer(mux), rsvc, authn -} - -func TestAddRule(t *testing.T) { - rs, rsvc, auth := setupRules() - defer rs.Close() - - conf := sdk.Config{ - RulesEngineURL: rs.URL, - } - mgsdk := sdk.NewSDK(conf) - - svcRule := re.Rule{ - ID: ruleID, - Name: "temperature-rule", - InputChannel: "chan-1", - Status: re.EnabledStatus, - } - - cases := []struct { - desc string - rule sdk.Rule - token string - session smqauthn.Session - svcRes re.Rule - svcErr error - authenticateErr error - wantErr bool - }{ - { - desc: "add rule successfully", - rule: sdk.Rule{Name: "temp-rule", InputChannel: "chan-1"}, - token: validToken, - svcRes: svcRule, - }, - { - desc: "add rule with empty token", - rule: sdk.Rule{Name: "temp-rule"}, - token: "", - wantErr: true, - }, - { - desc: "add rule with bad request", - rule: sdk.Rule{}, - token: validToken, - svcErr: errors.New("bad request"), - wantErr: true, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := rsvc.On("AddRule", mock.Anything, tc.session, mock.Anything).Return(tc.svcRes, []roles.RoleProvision(nil), tc.svcErr) - result, err := mgsdk.AddRule(context.Background(), tc.rule, domainID, tc.token) - assert.Equal(t, tc.wantErr, err != nil) - if !tc.wantErr { - assert.NotEmpty(t, result.ID) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestViewRule(t *testing.T) { - rs, rsvc, auth := setupRules() - defer rs.Close() - - conf := sdk.Config{ - RulesEngineURL: rs.URL, - } - mgsdk := sdk.NewSDK(conf) - - svcRule := re.Rule{ - ID: ruleID, - Name: "temperature-rule", - InputChannel: "chan-1", - Status: re.EnabledStatus, - } - - cases := []struct { - desc string - id string - token string - session smqauthn.Session - svcRes re.Rule - svcErr error - authenticateErr error - wantErr bool - }{ - { - desc: "view rule successfully", - id: ruleID, - token: validToken, - svcRes: svcRule, - }, - { - desc: "view rule with empty token", - id: ruleID, - token: "", - wantErr: true, - }, - { - desc: "view non-existent rule", - id: "non-existent", - token: validToken, - svcErr: errors.New("not found"), - wantErr: true, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := rsvc.On("ViewRule", mock.Anything, tc.session, tc.id, mock.Anything).Return(tc.svcRes, tc.svcErr) - result, err := mgsdk.ViewRule(context.Background(), tc.id, domainID, tc.token) - assert.Equal(t, tc.wantErr, err != nil) - if !tc.wantErr { - assert.NotEmpty(t, result.ID) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestUpdateRule(t *testing.T) { - rs, rsvc, auth := setupRules() - defer rs.Close() - - conf := sdk.Config{ - RulesEngineURL: rs.URL, - } - mgsdk := sdk.NewSDK(conf) - - updatedRule := testRule - updatedRule.Name = "updated-rule" - - svcRule := re.Rule{ - ID: ruleID, - Name: "updated-rule", - InputChannel: "chan-1", - InputTopic: "sensors/temperature", - Status: re.EnabledStatus, - Tags: []string{"temperature", "alerts"}, - } - - cases := []struct { - desc string - rule sdk.Rule - token string - session smqauthn.Session - svcRes re.Rule - svcErr error - authenticateErr error - wantErr bool - }{ - { - desc: "update rule successfully", - rule: updatedRule, - token: validToken, - svcRes: svcRule, - }, - { - desc: "update rule with empty token", - rule: updatedRule, - token: "", - wantErr: true, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := rsvc.On("UpdateRule", mock.Anything, tc.session, mock.Anything).Return(tc.svcRes, tc.svcErr) - result, err := mgsdk.UpdateRule(context.Background(), tc.rule, domainID, tc.token) - assert.Equal(t, tc.wantErr, err != nil) - if !tc.wantErr { - assert.NotEmpty(t, result.ID) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestUpdateRuleTags(t *testing.T) { - rs, rsvc, auth := setupRules() - defer rs.Close() - - conf := sdk.Config{ - RulesEngineURL: rs.URL, - } - mgsdk := sdk.NewSDK(conf) - - svcRule := re.Rule{ - ID: ruleID, - Tags: []string{"new-tag"}, - } - - cases := []struct { - desc string - rule sdk.Rule - token string - session smqauthn.Session - svcRes re.Rule - svcErr error - authenticateErr error - wantErr bool - }{ - { - desc: "update rule tags successfully", - rule: sdk.Rule{ID: ruleID, Tags: []string{"new-tag"}}, - token: validToken, - svcRes: svcRule, - }, - { - desc: "update rule tags with empty token", - rule: sdk.Rule{ID: ruleID}, - token: "", - wantErr: true, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := rsvc.On("UpdateRuleTags", mock.Anything, tc.session, mock.Anything).Return(tc.svcRes, tc.svcErr) - result, err := mgsdk.UpdateRuleTags(context.Background(), tc.rule, domainID, tc.token) - assert.Equal(t, tc.wantErr, err != nil) - if !tc.wantErr { - assert.NotEmpty(t, result.ID) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestUpdateRuleSchedule(t *testing.T) { - rs, rsvc, auth := setupRules() - defer rs.Close() - - conf := sdk.Config{ - RulesEngineURL: rs.URL, - } - mgsdk := sdk.NewSDK(conf) - - svcRule := re.Rule{ - ID: ruleID, - } - - cases := []struct { - desc string - rule sdk.Rule - token string - session smqauthn.Session - svcRes re.Rule - svcErr error - authenticateErr error - wantErr bool - }{ - { - desc: "update rule schedule successfully", - rule: sdk.Rule{ID: ruleID, Schedule: map[string]any{"cron": "0 * * * *"}}, - token: validToken, - svcRes: svcRule, - }, - { - desc: "update rule schedule with empty token", - rule: sdk.Rule{ID: ruleID}, - token: "", - wantErr: true, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := rsvc.On("UpdateRuleSchedule", mock.Anything, tc.session, mock.Anything).Return(tc.svcRes, tc.svcErr) - result, err := mgsdk.UpdateRuleSchedule(context.Background(), tc.rule, domainID, tc.token) - assert.Equal(t, tc.wantErr, err != nil) - if !tc.wantErr { - assert.NotEmpty(t, result.ID) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestListRules(t *testing.T) { - rs, rsvc, auth := setupRules() - defer rs.Close() - - conf := sdk.Config{ - RulesEngineURL: rs.URL, - } - mgsdk := sdk.NewSDK(conf) - - svcPage := re.Page{} - - cases := []struct { - desc string - pm sdk.PageMetadata - token string - session smqauthn.Session - svcRes re.Page - svcErr error - authenticateErr error - wantErr bool - }{ - { - desc: "list rules successfully", - pm: sdk.PageMetadata{Offset: 0, Limit: 10}, - token: validToken, - svcRes: svcPage, - }, - { - desc: "list rules with filters", - pm: sdk.PageMetadata{ - Limit: 5, - Name: "temp", - Status: "enabled", - InputChannel: "chan-1", - Tag: "temperature", - Dir: "desc", - Order: "created_at", - }, - token: validToken, - svcRes: svcPage, - }, - { - desc: "list rules with empty metadata excludes filter params", - pm: sdk.PageMetadata{}, - token: validToken, - svcRes: re.Page{}, - }, - { - desc: "list rules with empty token", - pm: sdk.PageMetadata{Limit: 10}, - token: "", - wantErr: true, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := rsvc.On("ListRules", mock.Anything, tc.session, mock.Anything).Return(tc.svcRes, tc.svcErr) - result, err := mgsdk.ListRules(context.Background(), tc.pm, domainID, tc.token) - assert.Equal(t, tc.wantErr, err != nil) - if !tc.wantErr { - assert.NotNil(t, result) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestEnableRule(t *testing.T) { - rs, rsvc, auth := setupRules() - defer rs.Close() - - conf := sdk.Config{ - RulesEngineURL: rs.URL, - } - mgsdk := sdk.NewSDK(conf) - - svcRule := re.Rule{ - ID: ruleID, - Status: re.EnabledStatus, - } - - cases := []struct { - desc string - id string - token string - session smqauthn.Session - svcRes re.Rule - svcErr error - authenticateErr error - wantErr bool - }{ - { - desc: "enable rule successfully", - id: ruleID, - token: validToken, - svcRes: svcRule, - }, - { - desc: "enable rule with empty token", - id: ruleID, - token: "", - wantErr: true, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := rsvc.On("EnableRule", mock.Anything, tc.session, tc.id).Return(tc.svcRes, tc.svcErr) - result, err := mgsdk.EnableRule(context.Background(), tc.id, domainID, tc.token) - assert.Equal(t, tc.wantErr, err != nil) - if !tc.wantErr { - assert.NotEmpty(t, result.ID) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestDisableRule(t *testing.T) { - rs, rsvc, auth := setupRules() - defer rs.Close() - - conf := sdk.Config{ - RulesEngineURL: rs.URL, - } - mgsdk := sdk.NewSDK(conf) - - svcRule := re.Rule{ - ID: ruleID, - Status: re.DisabledStatus, - } - - cases := []struct { - desc string - id string - token string - session smqauthn.Session - svcRes re.Rule - svcErr error - authenticateErr error - wantErr bool - }{ - { - desc: "disable rule successfully", - id: ruleID, - token: validToken, - svcRes: svcRule, - }, - { - desc: "disable rule with empty token", - id: ruleID, - token: "", - wantErr: true, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := rsvc.On("DisableRule", mock.Anything, tc.session, tc.id).Return(tc.svcRes, tc.svcErr) - result, err := mgsdk.DisableRule(context.Background(), tc.id, domainID, tc.token) - assert.Equal(t, tc.wantErr, err != nil) - if !tc.wantErr { - assert.NotEmpty(t, result.ID) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestRemoveRule(t *testing.T) { - rs, rsvc, auth := setupRules() - defer rs.Close() - - conf := sdk.Config{ - RulesEngineURL: rs.URL, - } - mgsdk := sdk.NewSDK(conf) - - cases := []struct { - desc string - id string - token string - session smqauthn.Session - svcErr error - authenticateErr error - wantErr bool - }{ - { - desc: "remove rule successfully", - id: ruleID, - token: validToken, - }, - { - desc: "remove rule with empty token", - id: ruleID, - token: "", - wantErr: true, - }, - { - desc: "remove non-existent rule", - id: "non-existent", - token: validToken, - svcErr: errors.New("not found"), - wantErr: true, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: domainID + "_" + validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := rsvc.On("RemoveRule", mock.Anything, tc.session, tc.id).Return(tc.svcErr) - err := mgsdk.RemoveRule(context.Background(), tc.id, domainID, tc.token) - assert.Equal(t, tc.wantErr, err != nil) - svcCall.Unset() - authCall.Unset() - }) - } -} diff --git a/pkg/sdk/setup_test.go b/pkg/sdk/setup_test.go deleted file mode 100644 index 143cb0901..000000000 --- a/pkg/sdk/setup_test.go +++ /dev/null @@ -1,314 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package sdk_test - -import ( - "fmt" - "os" - "regexp" - "testing" - "time" - - "github.com/absmach/magistrala/channels" - chmocks "github.com/absmach/magistrala/channels/mocks" - "github.com/absmach/magistrala/clients" - climocks "github.com/absmach/magistrala/clients/mocks" - "github.com/absmach/magistrala/domains" - groups "github.com/absmach/magistrala/groups" - "github.com/absmach/magistrala/internal/nullable" - "github.com/absmach/magistrala/internal/testsutil" - "github.com/absmach/magistrala/journal" - "github.com/absmach/magistrala/pkg/roles" - sdk "github.com/absmach/magistrala/pkg/sdk" - "github.com/absmach/magistrala/pkg/uuid" - "github.com/absmach/magistrala/users" - "github.com/stretchr/testify/assert" -) - -const ( - invalidIdentity = "invalididentity" - Identity = "identity" - Email = "email" - InvalidEmail = "invalidemail" - secret = "strongsecret" - invalidToken = "invalid" - contentType = sdk.CTJSON - invalid = "invalid" - wrongID = "wrongID" - roleName = "roleName" -) - -var ( - idProvider = uuid.New() - validMetadata = sdk.Metadata{"role": "client"} - user = generateTestUser(&testing.T{}) - description = "shortdescription" - gName = "groupname" - validToken = "valid" - limit uint64 = 5 - offset uint64 = 0 - total uint64 = 200 - passRegex = regexp.MustCompile("^.{8,}$") - validID = testsutil.GenerateUUID(&testing.T{}) - - clientsGRPCClient *climocks.ClientsServiceClient - channelsGRPCClient *chmocks.ChannelsServiceClient -) - -func generateUUID(t *testing.T) string { - ulid, err := idProvider.ID() - assert.Nil(t, err, fmt.Sprintf("unexpected error: %s", err)) - - return ulid -} - -func convertUsers(cs []sdk.User) []users.User { - ccs := []users.User{} - - for _, c := range cs { - ccs = append(ccs, convertUser(c)) - } - - return ccs -} - -func convertClients(cs ...sdk.Client) []clients.Client { - ccs := []clients.Client{} - - for _, c := range cs { - ccs = append(ccs, convertClient(c)) - } - - return ccs -} - -func convertGroups(cs []sdk.Group) []groups.Group { - cgs := []groups.Group{} - - for _, c := range cs { - cgs = append(cgs, convertGroup(c)) - } - - return cgs -} - -func convertChannels(cs []sdk.Channel) []channels.Channel { - chs := []channels.Channel{} - - for _, c := range cs { - chs = append(chs, convertChannel(c)) - } - - return chs -} - -func convertGroup(g sdk.Group) groups.Group { - if g.Status == "" { - g.Status = groups.EnabledStatus.String() - } - status, err := groups.ToStatus(g.Status) - if err != nil { - return groups.Group{} - } - var desc nullable.Value[string] - if g.Description != "" { - desc = nullable.New(g.Description) - } - - return groups.Group{ - ID: g.ID, - Domain: g.DomainID, - Parent: g.ParentID, - Name: g.Name, - Description: desc, - Tags: g.Tags, - Metadata: groups.Metadata(g.Metadata), - Level: g.Level, - Path: g.Path, - Children: convertChildren(g.Children), - CreatedAt: g.CreatedAt, - UpdatedAt: g.UpdatedAt, - Status: status, - RoleID: g.RoleID, - RoleName: g.RoleName, - Actions: g.Actions, - AccessType: g.AccessType, - AccessProviderId: g.AccessProviderId, - AccessProviderRoleId: g.AccessProviderRoleId, - AccessProviderRoleName: g.AccessProviderRoleName, - AccessProviderRoleActions: g.AccessProviderRoleActions, - Roles: g.Roles, - } -} - -func convertChildren(gs []*sdk.Group) []*groups.Group { - var cg []*groups.Group - - if len(gs) == 0 { - return cg - } - - for _, g := range gs { - insert := convertGroup(*g) - cg = append(cg, &insert) - } - - return cg -} - -func convertUser(c sdk.User) users.User { - if c.Status == "" { - c.Status = users.EnabledStatus.String() - } - status, err := users.ToStatus(c.Status) - if err != nil { - return users.User{} - } - role, err := users.ToRole(c.Role) - if err != nil { - return users.User{} - } - return users.User{ - ID: c.ID, - FirstName: c.FirstName, - LastName: c.LastName, - Tags: c.Tags, - Email: c.Email, - Credentials: users.Credentials(c.Credentials), - Metadata: users.Metadata(c.Metadata), - PrivateMetadata: users.Metadata(c.PrivateMetadata), - CreatedAt: c.CreatedAt, - UpdatedAt: c.UpdatedAt, - Status: status, - Role: role, - ProfilePicture: c.ProfilePicture, - } -} - -func convertClient(c sdk.Client) clients.Client { - if c.Status == "" { - c.Status = clients.EnabledStatus.String() - } - status, err := clients.ToStatus(c.Status) - if err != nil { - return clients.Client{} - } - return clients.Client{ - ID: c.ID, - Name: c.Name, - Tags: c.Tags, - Domain: c.DomainID, - ParentGroup: c.ParentGroup, - Credentials: clients.Credentials(c.Credentials), - Metadata: clients.Metadata(c.Metadata), - PrivateMetadata: clients.Metadata(c.PrivateMetadata), - CreatedAt: c.CreatedAt, - UpdatedAt: c.UpdatedAt, - UpdatedBy: c.UpdatedBy, - Status: status, - Roles: c.Roles, - } -} - -func convertChannel(g sdk.Channel) channels.Channel { - if g.Status == "" { - g.Status = channels.EnabledStatus.String() - } - status, err := channels.ToStatus(g.Status) - if err != nil { - return channels.Channel{} - } - return channels.Channel{ - ID: g.ID, - Name: g.Name, - Tags: g.Tags, - ParentGroup: g.ParentGroup, - Route: g.Route, - Domain: g.DomainID, - Metadata: channels.Metadata(g.Metadata), - CreatedAt: g.CreatedAt, - UpdatedAt: g.UpdatedAt, - UpdatedBy: g.UpdatedBy, - Status: status, - Roles: g.Roles, - } -} - -func convertInvitation(i sdk.Invitation) domains.Invitation { - return domains.Invitation{ - InvitedBy: i.InvitedBy, - InviteeUserID: i.InviteeUserID, - DomainID: i.DomainID, - RoleID: i.RoleID, - RoleName: i.RoleName, - Actions: i.Actions, - CreatedAt: i.CreatedAt, - UpdatedAt: i.UpdatedAt, - ConfirmedAt: i.ConfirmedAt, - RejectedAt: i.RejectedAt, - } -} - -func convertJournal(j sdk.Journal) journal.Journal { - return journal.Journal{ - ID: j.ID, - Operation: j.Operation, - OccurredAt: j.OccurredAt, - Attributes: j.Attributes, - Metadata: j.Metadata, - } -} - -func generateTestUser(t *testing.T) sdk.User { - createdAt, err := time.Parse(time.RFC3339, "2024-01-01T00:00:00Z") - assert.Nil(t, err, fmt.Sprintf("Unexpected error parsing time: %v", err)) - return sdk.User{ - ID: generateUUID(t), - FirstName: "userfirstname", - LastName: "userlastname", - Email: "useremail@example.com", - Credentials: sdk.Credentials{ - Username: "username", - Secret: secret, - }, - Tags: []string{"tag1", "tag2"}, - Metadata: validMetadata, - PrivateMetadata: validMetadata, - CreatedAt: createdAt, - UpdatedAt: createdAt, - Status: users.EnabledStatus.String(), - Role: users.UserRole.String(), - } -} - -func convertRole(r roles.Role) sdk.Role { - return sdk.Role{ - ID: r.ID, - Name: r.Name, - EntityID: r.EntityID, - CreatedBy: r.CreatedBy, - CreatedAt: r.CreatedAt, - UpdatedBy: r.UpdatedBy, - UpdatedAt: r.UpdatedAt, - } -} - -func convertRoleProvision(r roles.RoleProvision) sdk.Role { - return sdk.Role{ - ID: r.ID, - Name: r.Name, - EntityID: r.EntityID, - CreatedBy: r.CreatedBy, - CreatedAt: r.CreatedAt, - UpdatedBy: r.UpdatedBy, - UpdatedAt: r.UpdatedAt, - OptionalActions: r.OptionalActions, - OptionalMembers: r.OptionalMembers, - } -} - -func TestMain(m *testing.M) { - exitCode := m.Run() - os.Exit(exitCode) -} diff --git a/pkg/sdk/tokens_test.go b/pkg/sdk/tokens_test.go deleted file mode 100644 index f2220ed5a..000000000 --- a/pkg/sdk/tokens_test.go +++ /dev/null @@ -1,186 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package sdk_test - -import ( - "context" - "net/http" - "testing" - - grpcTokenV1 "github.com/absmach/magistrala/api/grpc/token/v1" - apiutil "github.com/absmach/magistrala/api/http/util" - smqauth "github.com/absmach/magistrala/auth" - smqauthn "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - sdk "github.com/absmach/magistrala/pkg/sdk" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" -) - -func TestIssueToken(t *testing.T) { - ts, svc, _ := setupUsers() - defer ts.Close() - - client := generateTestUser(t) - token := generateTestToken() - - conf := sdk.Config{ - UsersURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - cases := []struct { - desc string - login sdk.Login - svcRes *grpcTokenV1.Token - svcErr error - response sdk.Token - err errors.SDKError - }{ - { - desc: "issue token successfully", - login: sdk.Login{ - Username: client.Credentials.Username, - Password: client.Credentials.Secret, - }, - svcRes: &grpcTokenV1.Token{ - AccessToken: token.AccessToken, - RefreshToken: &token.RefreshToken, - AccessType: smqauth.AccessKey.String(), - }, - svcErr: nil, - response: token, - err: nil, - }, - { - desc: "issue token with invalid identity", - login: sdk.Login{ - Username: invalidIdentity, - Password: client.Credentials.Secret, - }, - svcRes: &grpcTokenV1.Token{}, - svcErr: svcerr.ErrAuthentication, - response: sdk.Token{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "issue token with invalid secret", - login: sdk.Login{ - Username: client.Credentials.Username, - Password: "invalid", - }, - svcRes: &grpcTokenV1.Token{}, - svcErr: svcerr.ErrLogin, - response: sdk.Token{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrLogin, http.StatusUnauthorized), - }, - { - desc: "issue token with empty identity", - login: sdk.Login{ - Username: "", - Password: client.Credentials.Secret, - }, - svcRes: &grpcTokenV1.Token{}, - svcErr: nil, - response: sdk.Token{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingUsernameEmail, http.StatusBadRequest), - }, - { - desc: "issue token with empty secret", - login: sdk.Login{ - Username: client.Credentials.Username, - Password: "", - }, - svcRes: &grpcTokenV1.Token{}, - svcErr: nil, - response: sdk.Token{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingPass, http.StatusBadRequest), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("IssueToken", mock.Anything, tc.login.Username, tc.login.Password, tc.login.Description).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.CreateToken(context.Background(), tc.login) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "IssueToken", mock.Anything, tc.login.Username, tc.login.Password, tc.login.Description) - assert.True(t, ok) - } - svcCall.Unset() - }) - } -} - -func TestRefreshToken(t *testing.T) { - ts, svc, auth := setupUsers() - defer ts.Close() - - token := generateTestToken() - - conf := sdk.Config{ - UsersURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - cases := []struct { - desc string - token string - svcRes *grpcTokenV1.Token - svcErr error - identifyErr error - response sdk.Token - err errors.SDKError - }{ - { - desc: "refresh token successfully", - token: token.RefreshToken, - svcRes: &grpcTokenV1.Token{ - AccessToken: token.AccessToken, - RefreshToken: &token.RefreshToken, - AccessType: token.AccessType, - }, - response: token, - err: nil, - }, - { - desc: "refresh token with invalid token", - token: invalidToken, - svcRes: nil, - identifyErr: svcerr.ErrAuthentication, - response: sdk.Token{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "refresh token with empty token", - token: "", - response: sdk.Token{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - authCall := auth.On("Authenticate", mock.Anything, mock.Anything).Return(smqauthn.Session{DomainUserID: validID, UserID: validID, DomainID: validID}, tc.identifyErr) - svcCall := svc.On("RefreshToken", mock.Anything, smqauthn.Session{DomainUserID: validID, UserID: validID, DomainID: validID}, tc.token).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.RefreshToken(context.Background(), tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "RefreshToken", mock.Anything, smqauthn.Session{DomainUserID: validID, UserID: validID, DomainID: validID}, tc.token) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func generateTestToken() sdk.Token { - return sdk.Token{ - AccessToken: "access_token", - RefreshToken: "refresh_token", - AccessType: smqauth.AccessKey.String(), - } -} diff --git a/pkg/sdk/transport_test.go b/pkg/sdk/transport_test.go deleted file mode 100644 index e1b34cd6f..000000000 --- a/pkg/sdk/transport_test.go +++ /dev/null @@ -1,157 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package sdk_test - -import ( - "context" - "net" - "net/http" - "net/http/httptest" - "strings" - "sync/atomic" - "testing" - - "github.com/absmach/magistrala/pkg/sdk" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" -) - -// TestTransport verifies that IdleConnTimeout=90s keeps connections pooled for -// healthy servers, and that network errors (EOF, reset) surface as descriptive errors. -func TestTransport(t *testing.T) { - cases := []struct { - desc string - serverFunc func(t *testing.T) (url string, cleanup func()) - ctxFunc func() context.Context - wantErr bool - errContains string - }{ - { - desc: "make request successfully with connection reuse", - serverFunc: func(t *testing.T) (string, func()) { - t.Helper() - var connCount atomic.Int32 - srv := httptest.NewUnstartedServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - w.Header().Set("Content-Type", "application/json") - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte(`{"clients":[{"id":"1","name":"test-client"}]}`)) - })) - srv.Config.ConnState = func(_ net.Conn, state http.ConnState) { - if state == http.StateNew { - connCount.Add(1) - } - } - srv.Start() - return srv.URL, func() { - srv.Close() - assert.Equal(t, int32(1), connCount.Load(), "expected connections to be reused (keep-alives enabled)") - } - }, - wantErr: false, - }, - { - desc: "make request with server closing connection", - serverFunc: func(t *testing.T) (string, func()) { - t.Helper() - ln, err := net.Listen("tcp", "127.0.0.1:0") - require.NoError(t, err) - go func() { - for { - conn, err := ln.Accept() - if err != nil { - return - } - conn.Close() - } - }() - return "http://" + ln.Addr().String(), func() { ln.Close() } - }, - wantErr: true, - errContains: "request failed", - }, - { - desc: "make request with connection reset by peer", - serverFunc: func(t *testing.T) (string, func()) { - t.Helper() - ln, err := net.Listen("tcp", "127.0.0.1:0") - require.NoError(t, err) - go func() { - for { - conn, err := ln.Accept() - if err != nil { - return - } - - tcpConn, ok := conn.(*net.TCPConn) - if ok { - _ = tcpConn.SetLinger(0) - } - conn.Close() - } - }() - return "http://" + ln.Addr().String(), func() { ln.Close() } - }, - wantErr: true, - errContains: "request failed", - }, - { - desc: "make request with unreachable server", - serverFunc: func(t *testing.T) (string, func()) { - t.Helper() - ln, err := net.Listen("tcp", "127.0.0.1:0") - require.NoError(t, err) - addr := ln.Addr().String() - ln.Close() - return "http://" + addr, func() {} - }, - wantErr: true, - errContains: "request failed", - }, - { - desc: "make request with cancelled context", - serverFunc: func(t *testing.T) (string, func()) { - t.Helper() - srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - w.WriteHeader(http.StatusOK) - })) - return srv.URL, srv.Close - }, - ctxFunc: func() context.Context { - ctx, cancel := context.WithCancel(context.Background()) - cancel() - return ctx - }, - wantErr: true, - errContains: "request failed", - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - url, cleanup := tc.serverFunc(t) - defer cleanup() - - smqsdk := sdk.NewSDK(sdk.Config{ClientsURL: url}) - - ctx := context.Background() - if tc.ctxFunc != nil { - ctx = tc.ctxFunc() - } - - client := sdk.Client{Name: "test-client"} - for i := 0; i < 2; i++ { - _, err := smqsdk.CreateClients(ctx, []sdk.Client{client}, domainID, validToken) - if tc.wantErr { - require.Error(t, err) - if tc.errContains != "" { - assert.True(t, strings.Contains(err.Error(), tc.errContains), - "expected error %q to contain %q", err.Error(), tc.errContains) - } - break - } - require.NoError(t, err) - } - }) - } -} diff --git a/pkg/sdk/users_test.go b/pkg/sdk/users_test.go deleted file mode 100644 index fa2f82a5a..000000000 --- a/pkg/sdk/users_test.go +++ /dev/null @@ -1,2494 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package sdk_test - -import ( - "context" - "fmt" - "net/http" - "net/http/httptest" - "strings" - "testing" - "time" - - grpcTokenV1 "github.com/absmach/magistrala/api/grpc/token/v1" - api "github.com/absmach/magistrala/api/http" - apiutil "github.com/absmach/magistrala/api/http/util" - authmocks "github.com/absmach/magistrala/auth/mocks" - mglog "github.com/absmach/magistrala/logger" - smqauthn "github.com/absmach/magistrala/pkg/authn" - authnmocks "github.com/absmach/magistrala/pkg/authn/mocks" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - oauth2mocks "github.com/absmach/magistrala/pkg/oauth2/mocks" - sdk "github.com/absmach/magistrala/pkg/sdk" - "github.com/absmach/magistrala/pkg/uuid" - "github.com/absmach/magistrala/users" - httpapi "github.com/absmach/magistrala/users/api" - umocks "github.com/absmach/magistrala/users/mocks" - "github.com/go-chi/chi/v5" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" -) - -var ( - id = generateUUID(&testing.T{}) - domainID = "c717fa97-ffd9-40cb-8cf9-7c2859059395" -) - -func setupUsers() (*httptest.Server, *umocks.Service, *authnmocks.Authentication) { - usvc := new(umocks.Service) - logger := mglog.NewMock() - mux := chi.NewRouter() - idp := uuid.NewMock() - provider := new(oauth2mocks.Provider) - provider.On("Name").Return("test") - authn := new(authnmocks.Authentication) - am := smqauthn.NewAuthNMiddleware(authn, smqauthn.WithDomainCheck(false), smqauthn.WithAllowUnverifiedUser(true)) - token := new(authmocks.TokenServiceClient) - httpapi.MakeHandler(usvc, am, token, true, mux, logger, "", passRegex, idp, provider) - - return httptest.NewServer(mux), usvc, authn -} - -func TestCreateUser(t *testing.T) { - ts, svc, _ := setupUsers() - defer ts.Close() - - createSdkUserReq := sdk.User{ - FirstName: user.FirstName, - LastName: user.LastName, - Email: user.Email, - Tags: user.Tags, - Credentials: user.Credentials, - Metadata: user.Metadata, - PrivateMetadata: user.PrivateMetadata, - Status: user.Status, - } - - conf := sdk.Config{ - UsersURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - cases := []struct { - desc string - token string - createSdkUserReq sdk.User - svcReq users.User - svcRes users.User - svcErr error - response sdk.User - err errors.SDKError - }{ - { - desc: "register new user successfully", - token: validToken, - createSdkUserReq: createSdkUserReq, - svcReq: convertUser(createSdkUserReq), - svcRes: convertUser(user), - svcErr: nil, - response: user, - err: nil, - }, - { - desc: "register existing user", - token: validToken, - createSdkUserReq: createSdkUserReq, - svcReq: convertUser(createSdkUserReq), - svcRes: users.User{}, - svcErr: svcerr.ErrCreateEntity, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrCreateEntity, http.StatusUnprocessableEntity), - }, - { - desc: "register user with invalid token", - token: invalidToken, - createSdkUserReq: createSdkUserReq, - svcReq: convertUser(createSdkUserReq), - svcRes: users.User{}, - svcErr: svcerr.ErrAuthentication, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "register user with empty token", - token: "", - createSdkUserReq: createSdkUserReq, - svcReq: convertUser(createSdkUserReq), - svcRes: users.User{}, - svcErr: svcerr.ErrAuthentication, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "register empty credentials user", - token: validToken, - createSdkUserReq: sdk.User{ - FirstName: createSdkUserReq.FirstName, - LastName: createSdkUserReq.LastName, - Email: createSdkUserReq.Email, - Credentials: sdk.Credentials{ - Username: "", - Secret: "", - }, - }, - svcReq: users.User{}, - svcRes: users.User{}, - svcErr: nil, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingUsername, http.StatusBadRequest), - }, - { - desc: "register user with first name too long", - token: validToken, - createSdkUserReq: sdk.User{ - FirstName: strings.Repeat("a", 1025), - Credentials: createSdkUserReq.Credentials, - PrivateMetadata: createSdkUserReq.PrivateMetadata, - Metadata: createSdkUserReq.Metadata, - Tags: createSdkUserReq.Tags, - }, - svcReq: users.User{}, - svcRes: users.User{}, - svcErr: nil, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrNameSize, http.StatusBadRequest), - }, - { - desc: "register user with empty userName", - token: validToken, - createSdkUserReq: sdk.User{ - FirstName: createSdkUserReq.FirstName, - LastName: createSdkUserReq.LastName, - Email: createSdkUserReq.Email, - Credentials: sdk.Credentials{ - Username: "", - Secret: createSdkUserReq.Credentials.Secret, - }, - PrivateMetadata: createSdkUserReq.PrivateMetadata, - Metadata: createSdkUserReq.Metadata, - Tags: createSdkUserReq.Tags, - }, - svcReq: users.User{}, - svcRes: users.User{}, - svcErr: nil, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingUsername, http.StatusBadRequest), - }, - { - desc: "register user with empty secret", - token: validToken, - createSdkUserReq: sdk.User{ - FirstName: createSdkUserReq.FirstName, - LastName: createSdkUserReq.LastName, - Email: createSdkUserReq.Email, - Credentials: sdk.Credentials{ - Username: createSdkUserReq.Credentials.Username, - Secret: "", - }, - PrivateMetadata: createSdkUserReq.PrivateMetadata, - Metadata: createSdkUserReq.Metadata, - Tags: createSdkUserReq.Tags, - }, - svcReq: users.User{}, - svcRes: users.User{}, - svcErr: nil, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingPass, http.StatusBadRequest), - }, - { - desc: "register user with secret that is too short", - token: validToken, - createSdkUserReq: sdk.User{ - FirstName: createSdkUserReq.FirstName, - LastName: createSdkUserReq.LastName, - Email: createSdkUserReq.Email, - Credentials: sdk.Credentials{ - Username: createSdkUserReq.Credentials.Username, - Secret: "weak", - }, - PrivateMetadata: createSdkUserReq.PrivateMetadata, - Metadata: createSdkUserReq.Metadata, - Tags: createSdkUserReq.Tags, - }, - svcReq: users.User{}, - svcRes: users.User{}, - svcErr: nil, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrPasswordFormat, http.StatusBadRequest), - }, - { - desc: "register a user with request that can't be marshalled", - token: validToken, - createSdkUserReq: sdk.User{ - Credentials: sdk.Credentials{ - Username: "user", - Secret: "12345678", - }, - FirstName: createSdkUserReq.FirstName, - LastName: createSdkUserReq.LastName, - Email: createSdkUserReq.Email, - Metadata: map[string]any{ - "test": make(chan int), - }, - }, - svcReq: users.User{}, - svcRes: users.User{}, - svcErr: nil, - response: sdk.User{}, - err: errors.NewSDKError(fmt.Errorf("json: unsupported type: chan int")), - }, - { - desc: "register a user with response that can't be unmarshalled", - token: validToken, - createSdkUserReq: createSdkUserReq, - svcReq: convertUser(createSdkUserReq), - svcRes: users.User{ - ID: id, - FirstName: createSdkUserReq.FirstName, - LastName: createSdkUserReq.LastName, - Email: createSdkUserReq.Email, - Credentials: users.Credentials{ - Username: createSdkUserReq.Credentials.Username, - Secret: createSdkUserReq.Credentials.Secret, - }, - Metadata: users.Metadata{ - "key": make(chan int), - }, - }, - svcErr: nil, - response: sdk.User{}, - err: errors.NewSDKError(fmt.Errorf("unexpected end of JSON input")), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("Register", mock.Anything, smqauthn.Session{}, tc.svcReq, true).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.CreateUser(context.Background(), tc.createSdkUserReq, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "Register", mock.Anything, smqauthn.Session{}, tc.svcReq, true) - assert.True(t, ok) - } - svcCall.Unset() - }) - } -} - -func TestListUsers(t *testing.T) { - ts, svc, auth := setupUsers() - defer ts.Close() - - var cls []sdk.User - conf := sdk.Config{ - UsersURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - for i := 10; i < 100; i++ { - cl := sdk.User{ - ID: generateUUID(t), - FirstName: fmt.Sprintf("user_%d", i), - Credentials: sdk.Credentials{ - Username: fmt.Sprintf("Username_%d", i), - Secret: fmt.Sprintf("password_%d", i), - }, - Metadata: sdk.Metadata{"name": fmt.Sprintf("user_%d", i)}, - Status: users.EnabledStatus.String(), - Role: users.UserRole.String(), - } - if i == 50 { - cl.Status = users.DisabledStatus.String() - cl.Tags = []string{"tag1", "tag2"} - } - cls = append(cls, cl) - } - - cases := []struct { - desc string - token string - session smqauthn.Session - pageMeta sdk.PageMetadata - svcReq users.Page - svcRes users.UsersPage - svcErr error - authenticateErr error - response sdk.UsersPage - err errors.SDKError - }{ - { - desc: "list users successfully", - token: validToken, - pageMeta: sdk.PageMetadata{ - Offset: offset, - Limit: limit, - }, - svcReq: users.Page{ - Offset: offset, - Limit: limit, - Order: api.DefOrder, - Dir: api.DefDir, - }, - svcRes: users.UsersPage{ - Page: users.Page{ - Total: uint64(len(cls[offset:limit])), - }, - Users: convertUsers(cls[offset:limit]), - }, - response: sdk.UsersPage{ - PageRes: sdk.PageRes{ - Total: uint64(len(cls[offset:limit])), - }, - Users: cls[offset:limit], - }, - err: nil, - }, - { - desc: "list users with invalid token", - token: invalidToken, - session: smqauthn.Session{}, - pageMeta: sdk.PageMetadata{ - Offset: offset, - Limit: limit, - }, - svcReq: users.Page{ - Offset: offset, - Limit: limit, - Order: api.DefOrder, - Dir: api.DefDir, - }, - svcRes: users.UsersPage{}, - svcErr: svcerr.ErrAuthentication, - authenticateErr: svcerr.ErrAuthentication, - response: sdk.UsersPage{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "list users with empty token", - token: "", - pageMeta: sdk.PageMetadata{ - Offset: offset, - Limit: limit, - }, - svcReq: users.Page{}, - svcRes: users.UsersPage{}, - svcErr: nil, - authenticateErr: apiutil.ErrBearerToken, - response: sdk.UsersPage{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "list users with zero limit", - token: validToken, - pageMeta: sdk.PageMetadata{ - Offset: offset, - Limit: 0, - }, - svcReq: users.Page{ - Offset: offset, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - }, - svcRes: users.UsersPage{ - Page: users.Page{ - Total: uint64(len(cls[offset:10])), - }, - Users: convertUsers(cls[offset:10]), - }, - response: sdk.UsersPage{ - PageRes: sdk.PageRes{ - Total: uint64(len(cls[offset:10])), - }, - Users: cls[offset:10], - }, - err: nil, - }, - { - desc: "list users with limit greater than max", - token: validToken, - pageMeta: sdk.PageMetadata{ - Offset: offset, - Limit: 101, - }, - svcReq: users.Page{}, - svcRes: users.UsersPage{}, - svcErr: nil, - response: sdk.UsersPage{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrLimitSize, http.StatusBadRequest), - }, - { - desc: "list users with given metadata", - token: validToken, - pageMeta: sdk.PageMetadata{ - Offset: offset, - Limit: limit, - Metadata: sdk.Metadata{"name": "user_99"}, - }, - svcReq: users.Page{ - Offset: offset, - Limit: limit, - Metadata: users.Metadata{"name": "user_99"}, - Order: api.DefOrder, - Dir: api.DefDir, - }, - svcRes: users.UsersPage{ - Page: users.Page{ - Total: 1, - }, - Users: []users.User{convertUser(cls[89])}, - }, - svcErr: nil, - response: sdk.UsersPage{ - PageRes: sdk.PageRes{ - Total: 1, - }, - Users: []sdk.User{cls[89]}, - }, - err: nil, - }, - { - desc: "list users with given status", - token: validToken, - pageMeta: sdk.PageMetadata{ - Offset: offset, - Limit: limit, - Status: users.DisabledStatus.String(), - }, - svcReq: users.Page{ - Offset: offset, - Limit: limit, - Status: users.DisabledStatus, - Order: api.DefOrder, - Dir: api.DefDir, - }, - svcRes: users.UsersPage{ - Page: users.Page{ - Total: 1, - }, - Users: []users.User{convertUser(cls[50])}, - }, - svcErr: nil, - response: sdk.UsersPage{ - PageRes: sdk.PageRes{ - Total: 1, - }, - Users: []sdk.User{cls[50]}, - }, - err: nil, - }, - { - desc: "list users with given tag", - token: validToken, - pageMeta: sdk.PageMetadata{ - Offset: offset, - Limit: limit, - Tags: sdk.TagsQuery{Elements: []string{"tag1"}, Operator: sdk.OrOp}, - }, - svcReq: users.Page{ - Offset: offset, - Limit: limit, - Tags: users.TagsQuery{Elements: []string{"tag1"}, Operator: users.OrOp}, - Order: api.DefOrder, - Dir: api.DefDir, - }, - svcRes: users.UsersPage{ - Page: users.Page{ - Total: 1, - }, - Users: []users.User{convertUser(cls[50])}, - }, - svcErr: nil, - response: sdk.UsersPage{ - PageRes: sdk.PageRes{ - Total: 1, - }, - Users: []sdk.User{cls[50]}, - }, - err: nil, - }, - { - desc: "list users with CreatedFrom", - token: validToken, - pageMeta: sdk.PageMetadata{ - Offset: offset, - Limit: limit, - CreatedFrom: time.Date(2024, 1, 1, 0, 0, 0, 0, time.UTC), - }, - svcReq: users.Page{ - Offset: offset, - Limit: limit, - CreatedFrom: time.Date(2024, 1, 1, 0, 0, 0, 0, time.UTC), - Order: api.DefOrder, - Dir: api.DefDir, - }, - svcRes: users.UsersPage{ - Page: users.Page{ - Total: uint64(len(cls[offset:limit])), - }, - Users: convertUsers(cls[offset:limit]), - }, - svcErr: nil, - response: sdk.UsersPage{ - PageRes: sdk.PageRes{ - Total: uint64(len(cls[offset:limit])), - }, - Users: cls[offset:limit], - }, - err: nil, - }, - { - desc: "list users with CreatedTo", - token: validToken, - pageMeta: sdk.PageMetadata{ - Offset: offset, - Limit: limit, - CreatedTo: time.Date(2025, 12, 31, 23, 59, 59, 0, time.UTC), - }, - svcReq: users.Page{ - Offset: offset, - Limit: limit, - CreatedTo: time.Date(2025, 12, 31, 23, 59, 59, 0, time.UTC), - Order: api.DefOrder, - Dir: api.DefDir, - }, - svcRes: users.UsersPage{ - Page: users.Page{ - Total: uint64(len(cls[offset:limit])), - }, - Users: convertUsers(cls[offset:limit]), - }, - svcErr: nil, - response: sdk.UsersPage{ - PageRes: sdk.PageRes{ - Total: uint64(len(cls[offset:limit])), - }, - Users: cls[offset:limit], - }, - err: nil, - }, - { - desc: "list users with both CreatedFrom and CreatedTo", - token: validToken, - pageMeta: sdk.PageMetadata{ - Offset: offset, - Limit: limit, - CreatedFrom: time.Date(2024, 1, 1, 0, 0, 0, 0, time.UTC), - CreatedTo: time.Date(2025, 12, 31, 23, 59, 59, 0, time.UTC), - }, - svcReq: users.Page{ - Offset: offset, - Limit: limit, - CreatedFrom: time.Date(2024, 1, 1, 0, 0, 0, 0, time.UTC), - CreatedTo: time.Date(2025, 12, 31, 23, 59, 59, 0, time.UTC), - Order: api.DefOrder, - Dir: api.DefDir, - }, - svcRes: users.UsersPage{ - Page: users.Page{ - Total: 2, - }, - Users: []users.User{convertUser(cls[10]), convertUser(cls[20])}, - }, - svcErr: nil, - response: sdk.UsersPage{ - PageRes: sdk.PageRes{ - Total: 2, - }, - Users: []sdk.User{cls[10], cls[20]}, - }, - err: nil, - }, - { - desc: "list users with request that can't be marshalled", - token: validToken, - pageMeta: sdk.PageMetadata{ - Offset: offset, - Limit: limit, - Metadata: sdk.Metadata{ - "test": make(chan int), - }, - }, - svcReq: users.Page{ - Offset: offset, - Limit: limit, - Order: api.DefOrder, - Dir: api.DefDir, - }, - svcRes: users.UsersPage{}, - svcErr: nil, - response: sdk.UsersPage{}, - err: errors.NewSDKError(fmt.Errorf("json: unsupported type: chan int")), - }, - { - desc: "list users with response that can't be unmarshalled", - token: validToken, - pageMeta: sdk.PageMetadata{ - Offset: offset, - Limit: limit, - }, - svcReq: users.Page{ - Offset: offset, - Limit: limit, - Order: api.DefOrder, - Dir: api.DefDir, - }, - svcRes: users.UsersPage{ - Page: users.Page{ - Total: uint64(len(cls[offset:limit])), - }, - Users: []users.User{ - { - ID: id, - FirstName: "user_99", - Metadata: users.Metadata{ - "key": make(chan int), - }, - }, - }, - }, - response: sdk.UsersPage{}, - err: errors.NewSDKError(fmt.Errorf("unexpected end of JSON input")), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := svc.On("ListUsers", mock.Anything, tc.session, tc.svcReq).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.Users(context.Background(), tc.pageMeta, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.token != "" { - ok := authCall.Parent.AssertCalled(t, "Authenticate", mock.Anything, tc.token) - assert.True(t, ok) - } - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "ListUsers", mock.Anything, tc.session, tc.svcReq) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestSearchUsers(t *testing.T) { - ts, svc, auth := setupUsers() - defer ts.Close() - - var cls []sdk.User - conf := sdk.Config{ - UsersURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - for i := 10; i < 100; i++ { - cl := sdk.User{ - ID: generateUUID(t), - FirstName: fmt.Sprintf("user_%d", i), - Email: fmt.Sprintf("email_%d", i), - Credentials: sdk.Credentials{ - Username: fmt.Sprintf("Username_%d", i), - Secret: fmt.Sprintf("password_%d", i), - }, - Metadata: sdk.Metadata{"name": fmt.Sprintf("user_%d", i)}, - Status: users.EnabledStatus.String(), - Role: users.UserRole.String(), - } - if i == 50 { - cl.Status = users.DisabledStatus.String() - cl.Tags = []string{"tag1", "tag2"} - } - cls = append(cls, cl) - } - - cases := []struct { - desc string - token string - page sdk.PageMetadata - response []sdk.User - searchreturn users.UsersPage - err errors.SDKError - authenticateErr error - }{ - { - desc: "search for users", - token: validToken, - err: nil, - page: sdk.PageMetadata{ - Offset: offset, - Limit: limit, - Username: "user_20", - }, - response: []sdk.User{cls[10]}, - searchreturn: users.UsersPage{ - Users: []users.User{convertUser(cls[10])}, - Page: users.Page{ - Total: 1, - Offset: offset, - Limit: limit, - }, - }, - }, - { - desc: "search for users with invalid token", - token: invalidToken, - page: sdk.PageMetadata{ - Offset: offset, - Limit: limit, - Username: "user_10", - }, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - response: nil, - authenticateErr: svcerr.ErrAuthentication, - }, - { - desc: "search for users with empty token", - token: "", - page: sdk.PageMetadata{ - Offset: offset, - Limit: limit, - Username: "user_10", - }, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - response: nil, - authenticateErr: svcerr.ErrAuthentication, - }, - { - desc: "search for users with empty query", - token: validToken, - page: sdk.PageMetadata{ - Offset: offset, - Limit: limit, - FirstName: "", - }, - err: errors.NewSDKErrorWithStatus(apiutil.ErrEmptySearchQuery, http.StatusBadRequest), - }, - { - desc: "search for users with invalid length of query", - token: validToken, - page: sdk.PageMetadata{ - Offset: offset, - Limit: limit, - Username: "a", - }, - err: errors.NewSDKErrorWithStatus(apiutil.ErrValidation, http.StatusBadRequest), - }, - { - desc: "search for users with invalid limit", - token: validToken, - page: sdk.PageMetadata{ - Offset: offset, - Limit: 0, - Username: "user_10", - }, - err: errors.NewSDKErrorWithStatus(apiutil.ErrLimitSize, http.StatusBadRequest), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(smqauthn.Session{DomainUserID: validID, UserID: validID, DomainID: domainID}, tc.authenticateErr) - svcCall := svc.On("SearchUsers", mock.Anything, mock.Anything).Return(tc.searchreturn, tc.err) - page, err := mgsdk.SearchUsers(context.Background(), tc.page, tc.token) - assert.Equal(t, tc.err, err, fmt.Sprintf("%s: expected error %v, got %v", tc.desc, tc.err, err)) - assert.Equal(t, tc.response, page.Users, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.response, page.Users)) - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestViewUser(t *testing.T) { - ts, svc, auth := setupUsers() - defer ts.Close() - - conf := sdk.Config{ - UsersURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - cases := []struct { - desc string - token string - session smqauthn.Session - userID string - svcRes users.User - svcErr error - authenticateErr error - response sdk.User - err errors.SDKError - }{ - { - desc: "view user successfully", - token: validToken, - userID: user.ID, - svcRes: convertUser(user), - svcErr: nil, - response: user, - err: nil, - }, - { - desc: "view user with invalid token", - token: invalidToken, - userID: user.ID, - svcRes: users.User{}, - svcErr: svcerr.ErrAuthentication, - authenticateErr: svcerr.ErrAuthentication, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "view user with empty token", - token: "", - userID: user.ID, - svcRes: users.User{}, - svcErr: svcerr.ErrAuthentication, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "view user with invalid id", - token: validToken, - userID: wrongID, - svcRes: users.User{}, - svcErr: svcerr.ErrNotFound, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrNotFound, http.StatusNotFound), - }, - { - desc: "view user with empty id", - token: validToken, - userID: "", - svcRes: users.User{}, - svcErr: nil, - response: sdk.User{}, - err: errors.NewSDKError(apiutil.ErrMissingID), - }, - { - desc: "view user with response that can't be unmarshalled", - token: validToken, - userID: user.ID, - svcRes: users.User{ - ID: id, - FirstName: user.FirstName, - LastName: user.LastName, - Metadata: users.Metadata{ - "key": make(chan int), - }, - }, - svcErr: nil, - response: sdk.User{}, - err: errors.NewSDKError(fmt.Errorf("unexpected end of JSON input")), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := svc.On("View", mock.Anything, tc.session, tc.userID).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.User(context.Background(), tc.userID, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "View", mock.Anything, tc.session, tc.userID) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestUserProfile(t *testing.T) { - ts, svc, auth := setupUsers() - defer ts.Close() - - conf := sdk.Config{ - UsersURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - cases := []struct { - desc string - token string - session smqauthn.Session - svcRes users.User - svcErr error - authenticateErr error - response sdk.User - err errors.SDKError - }{ - { - desc: "view user profile successfully", - token: validToken, - svcRes: convertUser(user), - svcErr: nil, - response: user, - err: nil, - }, - { - desc: "view user profile with invalid token", - token: invalidToken, - svcRes: users.User{}, - svcErr: nil, - authenticateErr: svcerr.ErrAuthentication, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "view user profile with empty token", - token: "", - svcRes: users.User{}, - svcErr: nil, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "view user profile with response that can't be unmarshalled", - token: validToken, - svcRes: users.User{ - ID: id, - FirstName: user.FirstName, - Metadata: users.Metadata{ - "key": make(chan int), - }, - }, - svcErr: nil, - response: sdk.User{}, - err: errors.NewSDKError(fmt.Errorf("unexpected end of JSON input")), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := svc.On("ViewProfile", mock.Anything, tc.session).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.UserProfile(context.Background(), tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "ViewProfile", mock.Anything, tc.session) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestUpdateUser(t *testing.T) { - ts, svc, auth := setupUsers() - defer ts.Close() - - conf := sdk.Config{ - UsersURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - updatedName := "updatedName" - updatedUser := user - updatedUser.FirstName = updatedName - - cases := []struct { - desc string - token string - session smqauthn.Session - updateUserReq sdk.User - userID string - svcReq users.UserReq - svcRes users.User - svcErr error - authenticateErr error - response sdk.User - err errors.SDKError - }{ - { - desc: "update user name with valid token", - token: validToken, - updateUserReq: sdk.User{ - ID: user.ID, - FirstName: updatedName, - }, - userID: user.ID, - svcReq: users.UserReq{ - FirstName: &updatedName, - }, - svcRes: convertUser(updatedUser), - svcErr: nil, - response: updatedUser, - err: nil, - }, - { - desc: "update user name with invalid token", - token: invalidToken, - updateUserReq: sdk.User{ - ID: user.ID, - FirstName: updatedName, - }, - userID: user.ID, - svcReq: users.UserReq{ - FirstName: &updatedName, - }, - svcRes: users.User{}, - authenticateErr: svcerr.ErrAuthentication, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "update user name with invalid id", - token: validToken, - updateUserReq: sdk.User{ - ID: wrongID, - FirstName: updatedName, - }, - userID: wrongID, - svcReq: users.UserReq{ - FirstName: &updatedName, - }, - svcRes: users.User{}, - svcErr: svcerr.ErrUpdateEntity, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrUpdateEntity, http.StatusUnprocessableEntity), - }, - { - desc: "update user name with empty token", - token: "", - updateUserReq: sdk.User{ - ID: user.ID, - FirstName: updatedName, - }, - userID: user.ID, - svcReq: users.UserReq{ - FirstName: &updatedName, - }, - svcRes: users.User{}, - svcErr: nil, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "update user name with empty id", - token: validToken, - updateUserReq: sdk.User{ - ID: "", - FirstName: updatedName, - }, - userID: "", - svcReq: users.UserReq{ - FirstName: &updatedName, - }, - svcRes: users.User{}, - svcErr: nil, - response: sdk.User{}, - err: errors.NewSDKError(apiutil.ErrMissingID), - }, - { - desc: "update user with request that can't be marshalled", - token: validToken, - updateUserReq: sdk.User{ - ID: generateUUID(t), - Metadata: map[string]any{ - "test": make(chan int), - }, - }, - svcReq: users.UserReq{}, - svcRes: users.User{}, - svcErr: nil, - response: sdk.User{}, - err: errors.NewSDKError(fmt.Errorf("json: unsupported type: chan int")), - }, - { - desc: "update user with response that can't be unmarshalled", - token: validToken, - updateUserReq: sdk.User{ - ID: user.ID, - FirstName: updatedName, - }, - userID: user.ID, - svcReq: users.UserReq{ - FirstName: &updatedName, - }, - svcRes: users.User{ - ID: id, - FirstName: updatedName, - Metadata: users.Metadata{ - "key": make(chan int), - }, - }, - svcErr: nil, - response: sdk.User{}, - err: errors.NewSDKError(fmt.Errorf("unexpected end of JSON input")), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := svc.On("Update", mock.Anything, tc.session, tc.userID, tc.svcReq).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.UpdateUser(context.Background(), tc.updateUserReq, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "Update", mock.Anything, tc.session, tc.userID, tc.svcReq) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestUpdateUserTags(t *testing.T) { - ts, svc, auth := setupUsers() - defer ts.Close() - - conf := sdk.Config{ - UsersURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - updatedTags := []string{"updatedTag1", "updatedTag2"} - - updatedUser := user - updatedUser.Tags = updatedTags - - cases := []struct { - desc string - token string - session smqauthn.Session - updateUserReq sdk.User - userID string - svcReq users.UserReq - svcRes users.User - svcErr error - authenticateErr error - response sdk.User - err errors.SDKError - }{ - { - desc: "update user tags with valid token", - token: validToken, - updateUserReq: sdk.User{ - ID: user.ID, - Tags: updatedTags, - }, - userID: user.ID, - svcReq: users.UserReq{ - Tags: &updatedTags, - }, - svcRes: convertUser(updatedUser), - svcErr: nil, - response: updatedUser, - err: nil, - }, - { - desc: "update user tags with invalid token", - token: invalidToken, - updateUserReq: sdk.User{ - ID: user.ID, - Tags: updatedTags, - }, - userID: user.ID, - svcReq: users.UserReq{ - Tags: &updatedTags, - }, - svcRes: users.User{}, - authenticateErr: svcerr.ErrAuthentication, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "update user tags with empty token", - token: "", - updateUserReq: sdk.User{ - ID: user.ID, - Tags: updatedTags, - }, - svcReq: users.UserReq{}, - svcRes: users.User{}, - svcErr: nil, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "update user tags with invalid id", - token: validToken, - updateUserReq: sdk.User{ - ID: wrongID, - Tags: updatedTags, - }, - userID: wrongID, - svcReq: users.UserReq{ - Tags: &updatedTags, - }, - svcRes: users.User{}, - svcErr: svcerr.ErrUpdateEntity, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrUpdateEntity, http.StatusUnprocessableEntity), - }, - { - desc: "update user tags with empty id", - token: validToken, - updateUserReq: sdk.User{ - ID: "", - Tags: updatedTags, - }, - userID: "", - svcReq: users.UserReq{}, - svcRes: users.User{}, - svcErr: nil, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "update user tags with request that can't be marshalled", - token: validToken, - updateUserReq: sdk.User{ - ID: generateUUID(t), - Metadata: map[string]any{ - "test": make(chan int), - }, - }, - svcReq: users.UserReq{}, - svcRes: users.User{}, - svcErr: nil, - response: sdk.User{}, - err: errors.NewSDKError(fmt.Errorf("json: unsupported type: chan int")), - }, - { - desc: "update user tags with response that can't be unmarshalled", - token: validToken, - updateUserReq: sdk.User{ - ID: user.ID, - Tags: updatedTags, - }, - userID: user.ID, - svcReq: users.UserReq{ - Tags: &updatedTags, - }, - svcRes: users.User{ - ID: id, - Tags: updatedTags, - Metadata: users.Metadata{ - "key": make(chan int), - }, - }, - svcErr: nil, - response: sdk.User{}, - err: errors.NewSDKError(fmt.Errorf("unexpected end of JSON input")), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := svc.On("UpdateTags", mock.Anything, tc.session, tc.userID, tc.svcReq).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.UpdateUserTags(context.Background(), tc.updateUserReq, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "UpdateTags", mock.Anything, tc.session, tc.userID, tc.svcReq) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestUpdateUserEmail(t *testing.T) { - ts, svc, auth := setupUsers() - defer ts.Close() - - conf := sdk.Config{ - UsersURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - updatedEmail := "updatedEmail@email.com" - updatedUser := user - updatedUser.Email = updatedEmail - - cases := []struct { - desc string - token string - session smqauthn.Session - updateUserReq sdk.User - svcReq string - svcRes users.User - svcErr error - authenticateErr error - response sdk.User - err errors.SDKError - }{ - { - desc: "update email with valid token", - token: validToken, - updateUserReq: sdk.User{ - ID: user.ID, - Email: updatedEmail, - Credentials: sdk.Credentials{ - Secret: user.Credentials.Secret, - }, - }, - svcReq: updatedEmail, - svcRes: convertUser(updatedUser), - svcErr: nil, - response: updatedUser, - err: nil, - }, - { - desc: "update email with invalid token", - token: invalidToken, - updateUserReq: sdk.User{ - ID: user.ID, - Email: updatedEmail, - Credentials: sdk.Credentials{ - Secret: user.Credentials.Secret, - }, - }, - svcReq: updatedEmail, - svcRes: users.User{}, - authenticateErr: svcerr.ErrAuthentication, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "update email with empty token", - token: "", - updateUserReq: sdk.User{ - ID: user.ID, - Email: updatedEmail, - Credentials: sdk.Credentials{ - Secret: user.Credentials.Secret, - }, - }, - svcReq: updatedEmail, - svcRes: users.User{}, - svcErr: nil, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "update email with invalid id", - token: validToken, - updateUserReq: sdk.User{ - ID: wrongID, - Email: updatedEmail, - Credentials: sdk.Credentials{ - Secret: user.Credentials.Secret, - }, - }, - svcReq: updatedEmail, - svcRes: users.User{}, - svcErr: svcerr.ErrUpdateEntity, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrUpdateEntity, http.StatusUnprocessableEntity), - }, - { - desc: "update email with empty id", - token: validToken, - updateUserReq: sdk.User{ - ID: "", - Email: updatedEmail, - Credentials: sdk.Credentials{ - Secret: user.Credentials.Secret, - }, - }, - svcReq: updatedEmail, - svcRes: users.User{}, - svcErr: nil, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "update email with response that can't be unmarshalled", - token: validToken, - updateUserReq: sdk.User{ - ID: user.ID, - Email: updatedEmail, - Credentials: sdk.Credentials{ - Secret: user.Credentials.Secret, - }, - }, - svcReq: updatedEmail, - svcRes: users.User{ - ID: id, - FirstName: updatedEmail, - Metadata: users.Metadata{ - "key": make(chan int), - }, - }, - svcErr: nil, - response: sdk.User{}, - err: errors.NewSDKError(fmt.Errorf("unexpected end of JSON input")), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := svc.On("UpdateEmail", mock.Anything, tc.session, tc.updateUserReq.ID, tc.svcReq).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.UpdateUserEmail(context.Background(), tc.updateUserReq, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "UpdateEmail", mock.Anything, tc.session, tc.updateUserReq.ID, tc.svcReq) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestResetPasswordRequest(t *testing.T) { - ts, svc, _ := setupUsers() - defer ts.Close() - - defHost := "http://localhost" - - conf := sdk.Config{ - UsersURL: ts.URL, - HostURL: defHost, - } - mgsdk := sdk.NewSDK(conf) - - validEmail := "test@email.com" - - cases := []struct { - desc string - email string - svcRes users.User - svcErr error - issueRes *grpcTokenV1.Token - issueErr error - err errors.SDKError - }{ - { - desc: "reset password request with valid email", - email: validEmail, - svcRes: convertUser(user), - svcErr: nil, - issueRes: &grpcTokenV1.Token{AccessToken: validToken, RefreshToken: &validToken}, - err: nil, - }, - { - desc: "reset password request with invalid email", - email: "invalidemail", - svcRes: users.User{}, - svcErr: svcerr.ErrNotFound, - issueRes: &grpcTokenV1.Token{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrNotFound, http.StatusNotFound), - }, - { - desc: "reset password request with empty email", - email: "", - svcRes: users.User{}, - svcErr: nil, - issueRes: &grpcTokenV1.Token{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingEmail, http.StatusBadRequest), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("SendPasswordReset", mock.Anything, tc.email).Return(tc.svcErr) - err := mgsdk.ResetPasswordRequest(context.Background(), tc.email) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "SendPasswordReset", mock.Anything, tc.email) - assert.True(t, ok) - } - svcCall.Unset() - }) - } -} - -func TestResetPassword(t *testing.T) { - ts, svc, auth := setupUsers() - defer ts.Close() - - conf := sdk.Config{ - UsersURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - newPassword := "newPassword" - - cases := []struct { - desc string - token string - session smqauthn.Session - newPassword string - confPassword string - svcErr error - authenticateErr error - err errors.SDKError - }{ - { - desc: "reset password successfully", - token: validToken, - session: smqauthn.Session{UserID: validID, DomainID: domainID}, - newPassword: newPassword, - confPassword: newPassword, - svcErr: nil, - err: nil, - }, - { - desc: "reset password with invalid token", - token: invalidToken, - newPassword: newPassword, - confPassword: newPassword, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "reset password with empty token", - token: "", - newPassword: newPassword, - confPassword: newPassword, - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "reset password with empty new password", - token: validToken, - session: smqauthn.Session{UserID: validID, DomainID: domainID}, - newPassword: "", - confPassword: newPassword, - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingPass, http.StatusBadRequest), - }, - { - desc: "reset password with empty confirm password", - token: validToken, - session: smqauthn.Session{UserID: validID, DomainID: domainID}, - newPassword: newPassword, - confPassword: "", - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingConfPass, http.StatusBadRequest), - }, - { - desc: "reset password with new password not matching confirm password", - token: validToken, - session: smqauthn.Session{UserID: validID, DomainID: domainID}, - newPassword: newPassword, - confPassword: "wrongPassword", - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrInvalidResetPass, http.StatusBadRequest), - }, - { - desc: "reset password with weak password", - token: validToken, - session: smqauthn.Session{UserID: validID, DomainID: domainID}, - newPassword: "weak", - confPassword: "weak", - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrPasswordFormat, http.StatusBadRequest), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := svc.On("ResetSecret", mock.Anything, tc.session, tc.newPassword).Return(tc.svcErr) - err := mgsdk.ResetPassword(context.Background(), tc.newPassword, tc.confPassword, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "ResetSecret", mock.Anything, tc.session, tc.newPassword) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestUpdatePassword(t *testing.T) { - ts, svc, auth := setupUsers() - defer ts.Close() - - conf := sdk.Config{ - UsersURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - newPassword := "newPassword" - updatedUser := user - updatedUser.Credentials.Secret = newPassword - - cases := []struct { - desc string - token string - session smqauthn.Session - oldPassword string - newPassword string - svcRes users.User - svcErr error - authenticateErr error - response sdk.User - err errors.SDKError - }{ - { - desc: "update password successfully", - token: validToken, - oldPassword: secret, - newPassword: newPassword, - svcRes: convertUser(updatedUser), - svcErr: nil, - response: updatedUser, - err: nil, - }, - { - desc: "update password with invalid token", - token: invalidToken, - oldPassword: secret, - newPassword: newPassword, - svcRes: users.User{}, - authenticateErr: svcerr.ErrAuthentication, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "update password with empty token", - token: "", - oldPassword: secret, - newPassword: newPassword, - svcRes: users.User{}, - svcErr: nil, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "update password with empty old password", - token: validToken, - oldPassword: "", - newPassword: newPassword, - svcRes: users.User{}, - svcErr: nil, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingPass, http.StatusBadRequest), - }, - { - desc: "update password with empty new password", - token: validToken, - oldPassword: secret, - newPassword: "", - svcRes: users.User{}, - svcErr: nil, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingPass, http.StatusBadRequest), - }, - { - desc: "update password with invalid new password", - token: validToken, - oldPassword: secret, - newPassword: "weak", - svcRes: users.User{}, - svcErr: nil, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrPasswordFormat, http.StatusBadRequest), - }, - { - desc: "update password with invalid old password", - token: validToken, - oldPassword: "wrongPassword", - newPassword: newPassword, - svcRes: users.User{}, - svcErr: svcerr.ErrLogin, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrLogin, http.StatusUnauthorized), - }, - { - desc: "update password with response that can't be unmarshalled", - token: validToken, - oldPassword: secret, - newPassword: newPassword, - svcRes: users.User{ - ID: id, - FirstName: user.FirstName, - Metadata: users.Metadata{ - "key": make(chan int), - }, - }, - svcErr: nil, - response: sdk.User{}, - err: errors.NewSDKError(fmt.Errorf("unexpected end of JSON input")), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := svc.On("UpdateSecret", mock.Anything, tc.session, tc.oldPassword, tc.newPassword).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.UpdatePassword(context.Background(), tc.oldPassword, tc.newPassword, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "UpdateSecret", mock.Anything, tc.session, tc.oldPassword, tc.newPassword) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestUpdateUserRole(t *testing.T) { - ts, svc, auth := setupUsers() - defer ts.Close() - - conf := sdk.Config{ - UsersURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - updatedUser := user - updatedRole := users.AdminRole.String() - updatedUser.Role = updatedRole - - cases := []struct { - desc string - token string - session smqauthn.Session - updateUserReq sdk.User - svcReq users.User - svcRes users.User - svcErr error - authenticateErr error - response sdk.User - err errors.SDKError - }{ - { - desc: "update user role with valid token", - token: validToken, - updateUserReq: sdk.User{ - ID: user.ID, - Role: updatedRole, - Email: user.Email, - }, - svcReq: users.User{ - ID: user.ID, - Role: users.AdminRole, - }, - svcRes: convertUser(updatedUser), - svcErr: nil, - response: updatedUser, - err: nil, - }, - { - desc: "update user role with invalid token", - token: invalidToken, - updateUserReq: sdk.User{ - ID: user.ID, - Role: updatedRole, - }, - svcReq: users.User{ - ID: user.ID, - Role: users.AdminRole, - }, - svcRes: users.User{}, - authenticateErr: svcerr.ErrAuthentication, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "update user role with empty token", - token: "", - updateUserReq: sdk.User{ - ID: user.ID, - Role: updatedRole, - }, - svcReq: users.User{}, - svcRes: users.User{}, - svcErr: nil, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "update user role with invalid id", - token: validToken, - updateUserReq: sdk.User{ - ID: wrongID, - Role: updatedRole, - }, - svcReq: users.User{ - ID: wrongID, - Role: users.AdminRole, - }, - svcRes: users.User{}, - svcErr: svcerr.ErrUpdateEntity, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrUpdateEntity, http.StatusUnprocessableEntity), - }, - { - desc: "update user role with empty id", - token: validToken, - updateUserReq: sdk.User{ - ID: "", - Role: updatedRole, - }, - svcReq: users.User{}, - svcRes: users.User{}, - svcErr: nil, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "update user role with request that can't be marshalled", - token: validToken, - updateUserReq: sdk.User{ - ID: generateUUID(t), - Metadata: map[string]any{ - "test": make(chan int), - }, - }, - svcReq: users.User{}, - svcRes: users.User{}, - svcErr: nil, - response: sdk.User{}, - err: errors.NewSDKError(fmt.Errorf("json: unsupported type: chan int")), - }, - { - desc: "update user role with response that can't be unmarshalled", - token: validToken, - updateUserReq: sdk.User{ - ID: user.ID, - Role: updatedRole, - }, - svcReq: users.User{ - ID: user.ID, - Role: users.AdminRole, - }, - svcRes: users.User{ - ID: id, - Role: users.AdminRole, - Metadata: users.Metadata{ - "key": make(chan int), - }, - }, - svcErr: nil, - response: sdk.User{}, - err: errors.NewSDKError(fmt.Errorf("unexpected end of JSON input")), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := svc.On("UpdateRole", mock.Anything, tc.session, tc.svcReq).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.UpdateUserRole(context.Background(), tc.updateUserReq, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "UpdateRole", mock.Anything, tc.session, tc.svcReq) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestUpdateUsername(t *testing.T) { - ts, svc, auth := setupUsers() - defer ts.Close() - - conf := sdk.Config{ - UsersURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - updatedUser := user - updatedUsername := "updatedUsername" - updatedUser.Credentials.Username = updatedUsername - - cases := []struct { - desc string - token string - session smqauthn.Session - updateUserReq sdk.User - svcReq users.User - svcRes users.User - svcErr error - authenticateErr error - response sdk.User - err errors.SDKError - }{ - { - desc: "update username with valid token", - token: validToken, - updateUserReq: sdk.User{ - ID: user.ID, - Credentials: sdk.Credentials{ - Username: updatedUsername, - }, - }, - svcReq: users.User{ - ID: user.ID, - Credentials: users.Credentials{ - Username: updatedUsername, - }, - }, - svcRes: convertUser(updatedUser), - svcErr: nil, - response: updatedUser, - err: nil, - }, - { - desc: "update username with invalid token", - token: invalidToken, - updateUserReq: sdk.User{ - ID: user.ID, - Credentials: sdk.Credentials{ - Username: updatedUsername, - }, - }, - svcReq: users.User{ - ID: user.ID, - Credentials: users.Credentials{ - Username: updatedUsername, - }, - }, - svcRes: users.User{}, - authenticateErr: svcerr.ErrAuthentication, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "update username with empty token", - token: "", - updateUserReq: sdk.User{ - ID: user.ID, - Credentials: sdk.Credentials{ - Username: updatedUsername, - }, - }, - svcReq: users.User{}, - svcRes: users.User{}, - svcErr: nil, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "update username with invalid id", - token: validToken, - updateUserReq: sdk.User{ - ID: wrongID, - Credentials: sdk.Credentials{ - Username: updatedUsername, - }, - }, - svcReq: users.User{ - ID: wrongID, - Credentials: users.Credentials{ - Username: updatedUsername, - }, - }, - svcRes: users.User{}, - svcErr: svcerr.ErrUpdateEntity, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrUpdateEntity, http.StatusUnprocessableEntity), - }, - { - desc: "update username with empty id", - token: validToken, - updateUserReq: sdk.User{ - ID: "", - Credentials: sdk.Credentials{ - Username: updatedUsername, - }, - }, - svcReq: users.User{}, - svcRes: users.User{}, - svcErr: nil, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "update username with response that can't be unmarshalled", - token: validToken, - updateUserReq: sdk.User{ - ID: user.ID, - Credentials: sdk.Credentials{ - Username: updatedUsername, - }, - }, - svcReq: users.User{ - ID: user.ID, - Credentials: users.Credentials{ - Username: updatedUsername, - }, - }, - svcRes: users.User{ - ID: id, - Credentials: users.Credentials{ - Username: updatedUsername, - }, - Metadata: users.Metadata{ - "key": make(chan int), - }, - }, - svcErr: nil, - response: sdk.User{}, - err: errors.NewSDKError(fmt.Errorf("unexpected end of JSON input")), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := svc.On("UpdateUsername", mock.Anything, tc.session, tc.svcReq.ID, tc.svcReq.Credentials.Username).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.UpdateUsername(context.Background(), tc.updateUserReq, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "UpdateUsername", mock.Anything, tc.session, tc.svcReq.ID, tc.svcReq.Credentials.Username) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestUpdateProfilePicture(t *testing.T) { - ts, svc, auth := setupUsers() - defer ts.Close() - - conf := sdk.Config{ - UsersURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - updatedProfilePicture := "http://updated.com/profile.jpg" - updatedUser := user - updatedUser.Email = updatedProfilePicture - - cases := []struct { - desc string - token string - session smqauthn.Session - updateUserReq sdk.User - userID string - svcReq users.UserReq - svcRes users.User - svcErr error - authenticateErr error - response sdk.User - err errors.SDKError - }{ - { - desc: "update profile picture with valid token", - token: validToken, - updateUserReq: sdk.User{ - ID: user.ID, - ProfilePicture: updatedProfilePicture, - }, - userID: user.ID, - svcReq: users.UserReq{ - ProfilePicture: &updatedProfilePicture, - }, - svcRes: convertUser(updatedUser), - svcErr: nil, - response: updatedUser, - err: nil, - }, - { - desc: "update profile picture with invalid token", - token: invalidToken, - updateUserReq: sdk.User{ - ID: user.ID, - ProfilePicture: updatedProfilePicture, - }, - userID: user.ID, - svcReq: users.UserReq{ - ProfilePicture: &updatedProfilePicture, - }, - svcRes: users.User{}, - authenticateErr: svcerr.ErrAuthentication, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "update profile picture with empty token", - token: "", - updateUserReq: sdk.User{ - ID: user.ID, - ProfilePicture: updatedProfilePicture, - }, - userID: user.ID, - svcReq: users.UserReq{ - ProfilePicture: &updatedProfilePicture, - }, - svcRes: users.User{}, - svcErr: nil, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "update profile picture with invalid id", - token: validToken, - updateUserReq: sdk.User{ - ID: wrongID, - ProfilePicture: updatedProfilePicture, - }, - userID: wrongID, - svcReq: users.UserReq{ - ProfilePicture: &updatedProfilePicture, - }, - svcRes: users.User{}, - svcErr: svcerr.ErrUpdateEntity, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrUpdateEntity, http.StatusUnprocessableEntity), - }, - { - desc: "update profile picture with empty id", - token: validToken, - updateUserReq: sdk.User{ - ID: "", - ProfilePicture: updatedProfilePicture, - }, - userID: "", - svcReq: users.UserReq{ - ProfilePicture: &updatedProfilePicture, - }, - svcRes: users.User{}, - svcErr: nil, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "update profile picture with request that can't be marshalled", - token: validToken, - updateUserReq: sdk.User{ - ID: generateUUID(t), - Metadata: map[string]any{ - "test": make(chan int), - }, - }, - svcReq: users.UserReq{}, - svcRes: users.User{}, - svcErr: nil, - response: sdk.User{}, - err: errors.NewSDKError(fmt.Errorf("json: unsupported type: chan int")), - }, - { - desc: "update profile picture with response that can't be unmarshalled", - token: validToken, - updateUserReq: sdk.User{ - ID: user.ID, - ProfilePicture: updatedProfilePicture, - }, - userID: user.ID, - svcReq: users.UserReq{ - ProfilePicture: &updatedProfilePicture, - }, - svcRes: users.User{ - ID: id, - Metadata: users.Metadata{ - "key": make(chan int), - }, - }, - svcErr: nil, - response: sdk.User{}, - err: errors.NewSDKError(fmt.Errorf("unexpected end of JSON input")), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := svc.On("UpdateProfilePicture", mock.Anything, tc.session, tc.userID, tc.svcReq).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.UpdateProfilePicture(context.Background(), tc.updateUserReq, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "UpdateProfilePicture", mock.Anything, tc.session, tc.userID, tc.svcReq) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestEnableUser(t *testing.T) { - ts, svc, auth := setupUsers() - defer ts.Close() - - conf := sdk.Config{ - UsersURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - enabledUser := user - enabledUser.Status = users.EnabledStatus.String() - - cases := []struct { - desc string - token string - session smqauthn.Session - userID string - svcRes users.User - svcErr error - authenticateErr error - response sdk.User - err errors.SDKError - }{ - { - desc: "enable user with valid token", - token: validToken, - userID: user.ID, - svcRes: convertUser(enabledUser), - svcErr: nil, - response: enabledUser, - err: nil, - }, - { - desc: "enable user with invalid token", - token: invalidToken, - userID: user.ID, - svcRes: users.User{}, - authenticateErr: svcerr.ErrAuthentication, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "enable user with empty token", - token: "", - userID: user.ID, - svcRes: users.User{}, - svcErr: nil, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := svc.On("Enable", mock.Anything, tc.session, tc.userID).Return(tc.svcRes, tc.svcErr) - - resp, err := mgsdk.EnableUser(context.Background(), tc.userID, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "Enable", mock.Anything, tc.session, tc.userID) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestDisableUser(t *testing.T) { - ts, svc, auth := setupUsers() - defer ts.Close() - - conf := sdk.Config{ - UsersURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - disabledUser := user - disabledUser.Status = users.DisabledStatus.String() - - cases := []struct { - desc string - token string - session smqauthn.Session - userID string - svcRes users.User - svcErr error - authenticateErr error - response sdk.User - err errors.SDKError - }{ - { - desc: "disable user with valid token", - token: validToken, - userID: user.ID, - svcRes: convertUser(disabledUser), - svcErr: nil, - - response: disabledUser, - err: nil, - }, - { - desc: "disable user with invalid token", - token: invalidToken, - userID: user.ID, - svcRes: users.User{}, - authenticateErr: svcerr.ErrAuthentication, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "disable user with empty token", - token: "", - userID: user.ID, - svcRes: users.User{}, - svcErr: nil, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "disable user with invalid id", - token: validToken, - userID: wrongID, - svcRes: users.User{}, - svcErr: svcerr.ErrUpdateEntity, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(svcerr.ErrUpdateEntity, http.StatusUnprocessableEntity), - }, - { - desc: "disable user with empty id", - token: validToken, - userID: "", - svcRes: users.User{}, - svcErr: nil, - response: sdk.User{}, - err: errors.NewSDKErrorWithStatus(apiutil.ErrMissingID, http.StatusBadRequest), - }, - { - desc: "disable user with response that can't be unmarshalled", - token: validToken, - userID: user.ID, - svcRes: users.User{ - ID: id, - Status: users.DisabledStatus, - Metadata: users.Metadata{ - "key": make(chan int), - }, - }, - svcErr: nil, - response: sdk.User{}, - err: errors.NewSDKError(fmt.Errorf("unexpected end of JSON input")), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := svc.On("Disable", mock.Anything, tc.session, tc.userID).Return(tc.svcRes, tc.svcErr) - resp, err := mgsdk.DisableUser(context.Background(), tc.userID, tc.token) - assert.Equal(t, tc.err, err) - assert.Equal(t, tc.response, resp) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "Disable", mock.Anything, tc.session, tc.userID) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} - -func TestDeleteUser(t *testing.T) { - ts, svc, auth := setupUsers() - defer ts.Close() - - conf := sdk.Config{ - UsersURL: ts.URL, - } - mgsdk := sdk.NewSDK(conf) - - cases := []struct { - desc string - token string - session smqauthn.Session - userID string - svcErr error - authenticateErr error - err errors.SDKError - }{ - { - desc: "delete user successfully", - token: validToken, - userID: validID, - svcErr: nil, - err: nil, - }, - { - desc: "delete user with invalid token", - token: invalidToken, - userID: validID, - authenticateErr: svcerr.ErrAuthentication, - err: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, http.StatusUnauthorized), - }, - { - desc: "delete user with empty token", - token: "", - userID: validID, - svcErr: nil, - err: errors.NewSDKErrorWithStatus(apiutil.ErrBearerToken, http.StatusUnauthorized), - }, - { - desc: "delete user with invalid id", - token: validToken, - userID: wrongID, - svcErr: svcerr.ErrRemoveEntity, - err: errors.NewSDKErrorWithStatus(svcerr.ErrRemoveEntity, http.StatusUnprocessableEntity), - }, - { - desc: "delete user with empty id", - token: validToken, - userID: "", - svcErr: nil, - err: errors.NewSDKError(apiutil.ErrMissingID), - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - if tc.token == validToken { - tc.session = smqauthn.Session{DomainUserID: validID, UserID: validID, DomainID: domainID} - } - authCall := auth.On("Authenticate", mock.Anything, tc.token).Return(tc.session, tc.authenticateErr) - svcCall := svc.On("Delete", mock.Anything, tc.session, tc.userID).Return(tc.svcErr) - err := mgsdk.DeleteUser(context.Background(), tc.userID, tc.token) - assert.Equal(t, tc.err, err) - if tc.err == nil { - ok := svcCall.Parent.AssertCalled(t, "Delete", mock.Anything, tc.session, tc.userID) - assert.True(t, ok) - } - svcCall.Unset() - authCall.Unset() - }) - } -} diff --git a/pkg/spicedb/schemadecoder.go b/pkg/spicedb/schemadecoder.go deleted file mode 100644 index 42d901878..000000000 --- a/pkg/spicedb/schemadecoder.go +++ /dev/null @@ -1,69 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package spicedb - -import ( - "fmt" - "io" - "os" - "strings" - - "github.com/absmach/magistrala/pkg/roles" - corev1 "github.com/authzed/spicedb/pkg/proto/core/v1" - "github.com/authzed/spicedb/pkg/schemadsl/compiler" - "github.com/authzed/spicedb/pkg/schemadsl/input" -) - -func GetActionsFromSchema(schemaPath string, objectType string) ([]roles.Action, error) { - objectType = strings.TrimSpace(objectType) - if objectType == "" { - return []roles.Action{}, fmt.Errorf("object type is empty string") - } - - file, err := os.Open(schemaPath) - if err != nil { - return []roles.Action{}, err - } - data, err := io.ReadAll(file) - if err != nil { - return []roles.Action{}, err - } - - compiledSchema, err := compiler.Compile(compiler.InputSchema{ - Source: input.Source("schema"), - SchemaString: string(data), - }, compiler.AllowUnprefixedObjectType()) - if err != nil { - return []roles.Action{}, err - } - - actions := []roles.Action{} - for _, od := range compiledSchema.ObjectDefinitions { - if objectType == od.Name { - for _, relation := range od.Relation { - if relation.UsersetRewrite == nil && relation.TypeInformation != nil && isAction(relation.TypeInformation) { - relName := strings.TrimSpace(relation.GetName()) - if relName == "" { - return []roles.Action{}, fmt.Errorf("got empty relation name") - } - actions = append(actions, roles.Action(relName)) - } - } - } - } - - if len(actions) == 0 { - return []roles.Action{}, fmt.Errorf("no actions found for type %s", objectType) - } - return actions, nil -} - -func isAction(ti *corev1.TypeInformation) bool { - for _, ar := range ti.AllowedDirectRelations { - if ar.GetNamespace() == "role" && ar.GetRelation() == "member" { - return true - } - } - return false -} diff --git a/provision/README.md b/provision/README.md deleted file mode 100644 index 80d539d3d..000000000 --- a/provision/README.md +++ /dev/null @@ -1,232 +0,0 @@ -# Provision service - -Provision service provides an HTTP API to create initial Magistrala resources for gateways or edge deployments. It can create clients and channels based on a configurable layout, optionally create bootstrap configurations, whitelist clients, and issue X.509 certificates for mTLS. - -For gateways to communicate with [Magistrala][magistrala], configuration is required (MQTT host, client, channels, certificates). A gateway can fetch bootstrap configuration from the [Bootstrap][bootstrap] service using its `` and ``. The [Agent][agent] service is typically used on gateways to retrieve that configuration. - -You can create bootstrap configuration directly via [Bootstrap][bootstrap] or through Provision. [Magistrala UI][mgxui] uses the Bootstrap service; Provision is intended to automate gateway setups where one physical gateway may require multiple clients and channels (for example, [Agent][agent] and [Export][export]). This setup is defined as a **provision layout**. - -## Configuration - -The service is configured using environment variables and/or a TOML config file. Defaults below are from `provision/config.go`. Docker add-on examples are in `docker/addons/provision/docker-compose.yaml` and [docker/.env](https://github.com/absmach/magistrala/blob/main/docker/.env). The binary reads `MG_PROVISION_*` variables; the add-on compose file uses `MG_PROVISION_*`, so ensure the container receives the expected names. - -### Core service - -| Variable | Description | Default | -| --- | --- | --- | -| `MG_PROVISION_HTTP_PORT` | Provision service listening port | `9016` | -| `MG_PROVISION_LOG_LEVEL` | Service log level | `info` | -| `MG_PROVISION_ENV_CLIENTS_TLS` | SDK TLS verification | `false` | -| `MG_PROVISION_SERVER_CERT` | HTTPS server certificate | "" | -| `MG_PROVISION_SERVER_KEY` | HTTPS server key | "" | -| `MG_SEND_TELEMETRY` | Send telemetry to Magistrala call-home server | `true` | -| `MG_MQTT_ADAPTER_INSTANCE_ID` | Instance ID used in health output | "" | - -### Magistrala endpoints and credentials - -| Variable | Description | Default | -| --- | --- | --- | -| `MG_PROVISION_USERS_LOCATION` | Users service URL | `http://localhost` | -| `MG_PROVISION_CLIENTS_LOCATION` | Clients service URL | `http://localhost` | -| `MG_PROVISION_CERTS_LOCATION` | Certs service URL (certs SDK) | `http://localhost` | -| `MG_PROVISION_BS_SVC_URL` | Bootstrap service URL | `http://localhost:9000` | -| `MG_PROVISION_CERTS_SVC_URL` | Certs service URL (Magistrala SDK) | `http://localhost:9019` | -| `MG_PROVISION_USERNAME` | Magistrala username | `user` | -| `MG_PROVISION_PASS` | Magistrala password | `test` | -| `MG_PROVISION_API_KEY` | Magistrala authentication token | "" | -| `MG_PROVISION_EMAIL` | Magistrala user email | `test@example.com` | -| `MG_PROVISION_DOMAIN_ID` | Default domain ID (unused by HTTP API) | "" | - -### Provisioning behavior - -| Variable | Description | Default | -| --- | --- | --- | -| `MG_PROVISION_CONFIG_FILE` | Provision config file | `config.toml` | -| `MG_PROVISION_X509_PROVISIONING` | Issue client certificates during provisioning | `false` | -| `MG_PROVISION_BS_CONFIG_PROVISIONING` | Save client config in Bootstrap | `true` | -| `MG_PROVISION_BS_AUTO_WHITELIST` | Auto-whitelist client | `true` | -| `MG_PROVISION_BS_CONTENT` | Bootstrap config content (JSON string) | "" | -| `MG_PROVISION_CERTS_HOURS_VALID` | Client cert validity period | `2400h` | - -## Features - -- **Layout-driven provisioning**: Create clients and channels from a predefined layout. -- **Bootstrap integration**: Create bootstrap configs and optionally whitelist clients. -- **X.509 certificates**: Issue client certificates during provisioning when enabled. -- **Gateway metadata**: Enrich gateway clients with control/data/export channel IDs. -- **Observability**: `/metrics` and `/health` endpoints. - -## Provision layout - -Provision layout is configured in a TOML file (see `provision/configs/config.toml` or `docker/addons/provision/configs/config.toml`). If the file exists, it is loaded and any missing fields are filled with env values. The layout defines which clients and channels will be created when calling `/mapping`. - -Default behavior (when no config file is loaded) creates one client and two channels: `control` and `data`. - -Notes: - -- At least one client must include `external_id` in metadata. This value is replaced with the `external_id` from the provisioning request and is used for bootstrap creation. -- Channel metadata `type` is reserved for `control`, `data`, and `export` and is used to enrich gateway metadata. -- Bootstrap content can be provided via `bootstrap.content` in the TOML file or as JSON through `MG_PROVISION_BS_CONTENT`. - -Example layout: - -```toml -[[clients]] - name = "client" - - [clients.metadata] - external_id = "xxxxxx" - -[[channels]] - name = "control-channel" - - [channels.metadata] - type = "control" - -[[channels]] - name = "data-channel" - - [channels.metadata] - type = "data" - -[[channels]] - name = "export-channel" - - [channels.metadata] - type = "data" -``` - -## Authentication - -Provision uses Magistrala APIs and requires a valid token. There are three ways to provide it: - -- `Authorization: Bearer ` on each request. -- `MG_PROVISION_API_KEY` in env or TOML (used when no header token is provided). -- `MG_PROVISION_USERNAME` and `MG_PROVISION_PASS` in env or TOML (used to create an access token when no header token is provided). - -`POST /{domainID}/mapping` can create its own token using API key or username/password if no `Authorization` header is provided. The `Authorization` header takes precedence when present. `GET /{domainID}/mapping` always requires a bearer token. - -## Architecture - -### Runtime flow - -1. The service loads configuration from env and optionally merges a config file. -2. `POST /{domainID}/mapping` validates the request and ensures a token exists. -3. Clients are created from the configured layout (external ID is injected into metadata). -4. Channels are created with names prefixed by the request `name`. -5. If enabled, bootstrap configs are created and clients are whitelisted (connected to channels). -6. If X.509 provisioning is enabled, certificates are issued and returned in the response. - -## Running - -Provision service can be run standalone or via Docker Compose. - -Standalone: - -```bash -make provision - -MG_PROVISION_BS_SVC_URL=http://localhost:9013 \ -MG_PROVISION_CLIENTS_LOCATION=http://localhost:9006 \ -MG_PROVISION_USERS_LOCATION=http://localhost:9002 \ -MG_PROVISION_CONFIG_FILE=provision/configs/config.toml \ -./build/provision -``` - -Docker Compose (add-on): - -```bash -docker compose -f docker/docker-compose.yaml -f docker/addons/provision/docker-compose.yaml up provision -``` - -## Usage - -The Provision service exposes the following endpoints: - -| Operation | Method & Path | Description | -| --- | --- | --- | -| `provision` | `POST /{domainID}/mapping` | Create clients, channels, bootstrap config, and optional certs | -| `mapping` | `GET /{domainID}/mapping` | Return bootstrap content from config | -| `health` | `GET /health` | Service health check | - -### Example: Provision a gateway - -When credentials are available via env/config, you can omit the `Authorization` header. `Content-Type` must be exactly `application/json`. - -```bash -curl -s -S -X POST http://localhost://mapping \ - -H 'Content-Type: application/json' \ - -d '{"name": "gateway-a", "external_id": "33:52:77:99:43", "external_key": "223334fw2"}' -``` - -If you want to supply a token explicitly: - -```bash -curl -s -S -X POST http://localhost://mapping \ - -H "Authorization: Bearer " \ - -H 'Content-Type: application/json' \ - -d '{"name": "gateway-a", "external_id": "", "external_key": ""}' -``` - -Response contains created clients, channels, and optional certificate data: - -```json -{ - "clients": [ - { - "id": "c22b0c0f-8c03-40da-a06b-37ed3a72c8d1", - "name": "client", - "key": "007cce56-e0eb-40d6-b2b9-ed348a97d1eb", - "metadata": { - "external_id": "33:52:79:C3:43" - } - } - ], - "channels": [ - { - "id": "064c680e-181b-4b58-975e-6983313a5170", - "name": "control-channel", - "metadata": { - "type": "control" - } - }, - { - "id": "579da92d-6078-4801-a18a-dd1cfa2aa44f", - "name": "data-channel", - "metadata": { - "type": "data" - } - } - ], - "whitelisted": { - "c22b0c0f-8c03-40da-a06b-37ed3a72c8d1": true - } -} -``` - -### Example: Read bootstrap mapping - -```bash -curl -s -S -X GET http://localhost://mapping \ - -H "Authorization: Bearer " \ - -H 'Content-Type: application/json' -``` - -## Certificates - -When `MG_PROVISION_X509_PROVISIONING=true`, the provisioning flow issues certificates for each client and returns them in the response as `client_cert`, `client_key`, and `ca_cert`. The certificate TTL is controlled by `MG_PROVISION_CERTS_HOURS_VALID`. - -## Testing - -```bash -go test ./provision/... -``` - -For an in-depth explanation of our Provision Service, see the [official documentation][doc]. - -[doc]: https://magistrala.absmach.eu/docs/dev-guide/services/provision/ -[magistrala]: https://github.com/absmach/magistrala -[bootstrap]: https://github.com/absmach/magistrala/tree/main/bootstrap -[export]: https://github.com/absmach/export -[agent]: https://github.com/absmach/agent -[mgxui]: https://github.com/absmach/magistrala/ui diff --git a/provision/api/doc.go b/provision/api/doc.go deleted file mode 100644 index 2424852cc..000000000 --- a/provision/api/doc.go +++ /dev/null @@ -1,6 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package api contains API-related concerns: endpoint definitions, middlewares -// and all resource representations. -package api diff --git a/provision/api/endpoint.go b/provision/api/endpoint.go deleted file mode 100644 index 2a02b49ad..000000000 --- a/provision/api/endpoint.go +++ /dev/null @@ -1,75 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package api - -import ( - "context" - - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/provision" - "github.com/go-kit/kit/endpoint" -) - -func doProvision(svc provision.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthorization - } - req := request.(provisionReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - res, err := svc.Provision(ctx, session.DomainID, req.token, req.Name, req.ExternalID, req.ExternalKey) - if err != nil { - return nil, err - } - - provisionResponse := provisionRes{ - Clients: res.Clients, - Channels: res.Channels, - ClientCert: res.ClientCert, - ClientKey: res.ClientKey, - CACert: res.CACert, - Whitelisted: res.Whitelisted, - } - - return provisionResponse, nil - } -} - -func getMapping(svc provision.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - res := svc.Mapping() - - return mappingRes{Data: res}, nil - } -} - -func issueCert(svc provision.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthorization - } - req := request.(certReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - cert, key, err := svc.Cert(ctx, session.DomainID, req.token, req.ClientID, req.TTL) - if err != nil { - return nil, err - } - - return certRes{ - Certificate: cert, - Key: key, - }, nil - } -} diff --git a/provision/api/endpoint_test.go b/provision/api/endpoint_test.go deleted file mode 100644 index ba40b6c15..000000000 --- a/provision/api/endpoint_test.go +++ /dev/null @@ -1,329 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package api_test - -import ( - "fmt" - "io" - "net/http" - "net/http/httptest" - "strings" - "testing" - - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/auth" - "github.com/absmach/magistrala/internal/testsutil" - mglog "github.com/absmach/magistrala/logger" - smqauthn "github.com/absmach/magistrala/pkg/authn" - authnmocks "github.com/absmach/magistrala/pkg/authn/mocks" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/provision" - "github.com/absmach/magistrala/provision/api" - mocks "github.com/absmach/magistrala/provision/mocks" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" -) - -var ( - validToken = "valid" - validContenType = "application/json" - validID = testsutil.GenerateUUID(&testing.T{}) - userID = testsutil.GenerateUUID(&testing.T{}) - domainID = testsutil.GenerateUUID(&testing.T{}) - validSession = smqauthn.Session{ - DomainUserID: auth.EncodeDomainUserID(domainID, userID), - UserID: userID, - DomainID: domainID, - } -) - -type testRequest struct { - client *http.Client - method string - url string - token string - contentType string - body io.Reader -} - -func (tr testRequest) make() (*http.Response, error) { - req, err := http.NewRequest(tr.method, tr.url, tr.body) - if err != nil { - return nil, err - } - - if tr.token != "" { - req.Header.Set("Authorization", apiutil.BearerPrefix+tr.token) - } - - if tr.contentType != "" { - req.Header.Set("Content-Type", tr.contentType) - } - - return tr.client.Do(req) -} - -func newProvisionServer() (*httptest.Server, *mocks.Service, *authnmocks.Authentication) { - svc := new(mocks.Service) - - logger := mglog.NewMock() - authn := new(authnmocks.Authentication) - am := smqauthn.NewAuthNMiddleware(authn, smqauthn.WithAllowUnverifiedUser(true)) - mux := api.MakeHandler(svc, am, logger, "test") - return httptest.NewServer(mux), svc, authn -} - -func TestProvision(t *testing.T) { - is, svc, authn := newProvisionServer() - - cases := []struct { - desc string - token string - domainID string - data string - contentType string - status int - authnRes smqauthn.Session - authnErr error - svcErr error - }{ - { - desc: "valid request", - token: validToken, - domainID: validID, - data: fmt.Sprintf(`{"name": "test", "external_id": "%s", "external_key": "%s"}`, validID, validID), - status: http.StatusCreated, - contentType: validContenType, - authnRes: validSession, - svcErr: nil, - }, - { - desc: "request with empty external id", - token: validToken, - domainID: validID, - data: fmt.Sprintf(`{"name": "test", "external_key": "%s"}`, validID), - status: http.StatusBadRequest, - contentType: validContenType, - authnRes: validSession, - }, - { - desc: "request with empty external key", - token: validToken, - domainID: validID, - data: fmt.Sprintf(`{"name": "test", "external_id": "%s"}`, validID), - status: http.StatusUnauthorized, - contentType: validContenType, - authnRes: validSession, - svcErr: nil, - }, - { - desc: "empty token", - token: "", - domainID: validID, - data: fmt.Sprintf(`{"name": "test", "external_id": "%s", "external_key": "%s"}`, validID, validID), - status: http.StatusUnauthorized, - contentType: validContenType, - authnRes: smqauthn.Session{}, - authnErr: errors.ErrAuthentication, - svcErr: nil, - }, - { - desc: "invalid content type", - token: validToken, - domainID: validID, - data: fmt.Sprintf(`{"name": "test", "external_id": "%s", "external_key": "%s"}`, validID, validID), - status: http.StatusUnsupportedMediaType, - contentType: "text/plain", - authnRes: validSession, - svcErr: nil, - }, - { - desc: "invalid request", - token: validToken, - domainID: validID, - data: `data`, - status: http.StatusBadRequest, - contentType: validContenType, - authnRes: validSession, - svcErr: nil, - }, - { - desc: "service error", - token: validToken, - domainID: validID, - data: fmt.Sprintf(`{"name": "test", "external_id": "%s", "external_key": "%s"}`, validID, validID), - status: http.StatusForbidden, - contentType: validContenType, - authnRes: validSession, - svcErr: svcerr.ErrAuthorization, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.authnRes, tc.authnErr) - repocall := svc.On("Provision", mock.Anything, validID, tc.token, "test", validID, validID).Return(provision.Result{}, tc.svcErr) - req := testRequest{ - client: is.Client(), - method: http.MethodPost, - url: is.URL + fmt.Sprintf("/%s/mapping", tc.domainID), - token: tc.token, - contentType: tc.contentType, - body: strings.NewReader(tc.data), - } - - resp, err := req.make() - assert.Nil(t, err, tc.desc) - assert.Equal(t, tc.status, resp.StatusCode, tc.desc) - authCall.Unset() - repocall.Unset() - }) - } -} - -func TestMapping(t *testing.T) { - is, svc, authn := newProvisionServer() - - cases := []struct { - desc string - token string - domainID string - contentType string - status int - authnRes smqauthn.Session - authnErr error - svcErr error - }{ - { - desc: "valid request", - token: validToken, - domainID: validID, - status: http.StatusOK, - contentType: validContenType, - svcErr: nil, - authnRes: validSession, - authnErr: nil, - }, - { - desc: "empty token", - token: "", - domainID: validID, - status: http.StatusUnauthorized, - contentType: validContenType, - svcErr: nil, - authnRes: smqauthn.Session{}, - authnErr: errors.ErrAuthentication, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.authnRes, tc.authnErr) - repocall := svc.On("Mapping").Return(map[string]any{}, tc.svcErr) - req := testRequest{ - client: is.Client(), - method: http.MethodGet, - url: is.URL + fmt.Sprintf("/%s/mapping", tc.domainID), - token: tc.token, - contentType: tc.contentType, - } - - resp, err := req.make() - assert.Nil(t, err, tc.desc) - assert.Equal(t, tc.status, resp.StatusCode, tc.desc) - authCall.Unset() - repocall.Unset() - }) - } -} - -func TestCert(t *testing.T) { - is, svc, authn := newProvisionServer() - - cases := []struct { - desc string - token string - domainID string - data string - contentType string - status int - authnRes smqauthn.Session - authnErr error - svcErr error - }{ - { - desc: "valid request", - token: validToken, - domainID: validID, - data: fmt.Sprintf(`{"client_id": "%s", "ttl": "1h"}`, validID), - status: http.StatusCreated, - contentType: validContenType, - authnRes: validSession, - svcErr: nil, - }, - { - desc: "empty token", - token: "", - domainID: validID, - data: fmt.Sprintf(`{"client_id": "%s", "ttl": "1h"}`, validID), - status: http.StatusUnauthorized, - contentType: validContenType, - authnRes: smqauthn.Session{}, - authnErr: errors.ErrAuthentication, - svcErr: nil, - }, - { - desc: "invalid content type", - token: validToken, - domainID: validID, - data: fmt.Sprintf(`{"client_id": "%s", "ttl": "1h"}`, validID), - status: http.StatusUnsupportedMediaType, - contentType: "text/plain", - authnRes: validSession, - svcErr: nil, - }, - { - desc: "invalid request", - token: validToken, - domainID: validID, - data: `data`, - status: http.StatusBadRequest, - contentType: validContenType, - authnRes: validSession, - svcErr: nil, - }, - { - desc: "service error", - token: validToken, - domainID: validID, - data: fmt.Sprintf(`{"client_id": "%s", "ttl": "1h"}`, validID), - status: http.StatusForbidden, - contentType: validContenType, - authnRes: validSession, - svcErr: svcerr.ErrAuthorization, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.authnRes, tc.authnErr) - repocall := svc.On("Cert", mock.Anything, validID, tc.token, validID, "1h").Return("cert", "key", tc.svcErr) - req := testRequest{ - client: is.Client(), - method: http.MethodPost, - url: is.URL + fmt.Sprintf("/%s/cert", tc.domainID), - token: tc.token, - contentType: tc.contentType, - body: strings.NewReader(tc.data), - } - - resp, err := req.make() - assert.Nil(t, err, tc.desc) - assert.Equal(t, tc.status, resp.StatusCode, tc.desc) - authCall.Unset() - repocall.Unset() - }) - } -} diff --git a/provision/api/requests.go b/provision/api/requests.go deleted file mode 100644 index f93677e61..000000000 --- a/provision/api/requests.go +++ /dev/null @@ -1,43 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package api - -import apiutil "github.com/absmach/magistrala/api/http/util" - -type provisionReq struct { - token string - Name string `json:"name"` - ExternalID string `json:"external_id"` - ExternalKey string `json:"external_key"` -} - -func (req provisionReq) validate() error { - if req.ExternalID == "" { - return apiutil.ErrMissingID - } - - if req.ExternalKey == "" { - return apiutil.ErrBearerKey - } - - if req.Name == "" { - return apiutil.ErrMissingName - } - - return nil -} - -type certReq struct { - token string - ClientID string `json:"client_id"` - TTL string `json:"ttl,omitempty"` -} - -func (req certReq) validate() error { - if req.ClientID == "" { - return apiutil.ErrMissingID - } - - return nil -} diff --git a/provision/api/requests_test.go b/provision/api/requests_test.go deleted file mode 100644 index 0158489dd..000000000 --- a/provision/api/requests_test.go +++ /dev/null @@ -1,58 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package api - -import ( - "fmt" - "testing" - - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/internal/testsutil" - "github.com/absmach/magistrala/pkg/errors" - "github.com/stretchr/testify/assert" -) - -func TestProvisioReq(t *testing.T) { - cases := []struct { - desc string - req provisionReq - err error - }{ - { - desc: "valid request", - req: provisionReq{ - token: "token", - Name: "name", - ExternalID: testsutil.GenerateUUID(t), - ExternalKey: testsutil.GenerateUUID(t), - }, - err: nil, - }, - { - desc: "empty external id", - req: provisionReq{ - token: "token", - Name: "name", - ExternalID: "", - ExternalKey: testsutil.GenerateUUID(t), - }, - err: apiutil.ErrMissingID, - }, - { - desc: "empty external key", - req: provisionReq{ - token: "token", - Name: "name", - ExternalID: testsutil.GenerateUUID(t), - ExternalKey: "", - }, - err: apiutil.ErrBearerKey, - }, - } - - for _, tc := range cases { - err := tc.req.validate() - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected `%v` got `%v`", tc.desc, tc.err, err)) - } -} diff --git a/provision/api/responses.go b/provision/api/responses.go deleted file mode 100644 index ccac4deae..000000000 --- a/provision/api/responses.go +++ /dev/null @@ -1,72 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package api - -import ( - "encoding/json" - "net/http" - - "github.com/absmach/magistrala" - "github.com/absmach/magistrala/pkg/sdk" -) - -var _ magistrala.Response = (*provisionRes)(nil) - -type provisionRes struct { - Clients []sdk.Client `json:"clients"` - Channels []sdk.Channel `json:"channels"` - ClientCert map[string]string `json:"client_cert,omitempty"` - ClientKey map[string]string `json:"client_key,omitempty"` - CACert string `json:"ca_cert,omitempty"` - Whitelisted map[string]bool `json:"whitelisted,omitempty"` -} - -func (res provisionRes) Code() int { - return http.StatusCreated -} - -func (res provisionRes) Headers() map[string]string { - return map[string]string{} -} - -func (res provisionRes) Empty() bool { - return false -} - -type mappingRes struct { - Data any -} - -func (res mappingRes) Code() int { - return http.StatusOK -} - -func (res mappingRes) Headers() map[string]string { - return map[string]string{} -} - -func (res mappingRes) Empty() bool { - return false -} - -type certRes struct { - Certificate string `json:"certificate"` - Key string `json:"key"` -} - -func (res certRes) Code() int { - return http.StatusCreated -} - -func (res certRes) Headers() map[string]string { - return map[string]string{} -} - -func (res certRes) Empty() bool { - return false -} - -func (res mappingRes) MarshalJSON() ([]byte, error) { - return json.Marshal(res.Data) -} diff --git a/provision/api/transport.go b/provision/api/transport.go deleted file mode 100644 index 973a8ab02..000000000 --- a/provision/api/transport.go +++ /dev/null @@ -1,96 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package api - -import ( - "context" - "encoding/json" - "log/slog" - "net/http" - - "github.com/absmach/magistrala" - api "github.com/absmach/magistrala/api/http" - apiutil "github.com/absmach/magistrala/api/http/util" - smqauthn "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/errors" - "github.com/absmach/magistrala/provision" - "github.com/go-chi/chi/v5" - kithttp "github.com/go-kit/kit/transport/http" - "github.com/prometheus/client_golang/prometheus/promhttp" -) - -const ( - contentType = "application/json" -) - -// MakeHandler returns a HTTP handler for API endpoints. -func MakeHandler(svc provision.Service, authn smqauthn.AuthNMiddleware, logger *slog.Logger, instanceID string) http.Handler { - opts := []kithttp.ServerOption{ - kithttp.ServerErrorEncoder(apiutil.LoggingErrorEncoder(logger, api.EncodeError)), - } - - r := chi.NewRouter() - - r.Route("/{domainID}", func(r chi.Router) { - r.Use(authn.WithOptions(smqauthn.WithDomainCheck(true)).Middleware()) - r.Route("/mapping", func(r chi.Router) { - r.Post("/", kithttp.NewServer( - doProvision(svc), - decodeProvisionRequest, - api.EncodeResponse, - opts..., - ).ServeHTTP) - r.Get("/", kithttp.NewServer( - getMapping(svc), - decodeMappingRequest, - api.EncodeResponse, - opts..., - ).ServeHTTP) - }) - r.Post("/cert", kithttp.NewServer( - issueCert(svc), - decodeCertRequest, - api.EncodeResponse, - opts..., - ).ServeHTTP) - }) - r.Handle("/metrics", promhttp.Handler()) - r.Get("/health", magistrala.Health("provision", instanceID)) - - return r -} - -func decodeProvisionRequest(_ context.Context, r *http.Request) (any, error) { - if r.Header.Get("Content-Type") != contentType { - return nil, apiutil.ErrUnsupportedContentType - } - - req := provisionReq{ - token: apiutil.ExtractBearerToken(r), - } - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - - return req, nil -} - -func decodeMappingRequest(_ context.Context, r *http.Request) (any, error) { - return nil, nil -} - -func decodeCertRequest(_ context.Context, r *http.Request) (any, error) { - if r.Header.Get("Content-Type") != contentType { - return nil, apiutil.ErrUnsupportedContentType - } - - req := certReq{ - token: apiutil.ExtractBearerToken(r), - } - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - - return req, nil -} diff --git a/provision/config.go b/provision/config.go deleted file mode 100644 index 32801554f..000000000 --- a/provision/config.go +++ /dev/null @@ -1,120 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package provision - -import ( - "fmt" - "os" - - "github.com/absmach/magistrala/channels" - "github.com/absmach/magistrala/clients" - "github.com/absmach/magistrala/pkg/errors" - "github.com/pelletier/go-toml" -) - -var errFailedToReadConfig = errors.New("failed to read config file") - -// ServiceConf represents service config. -type ServiceConf struct { - Port string `toml:"port" env:"MG_PROVISION_HTTP_PORT" envDefault:"9016"` - LogLevel string `toml:"log_level" env:"MG_PROVISION_LOG_LEVEL" envDefault:"info"` - TLS bool `toml:"tls" env:"MG_PROVISION_ENV_CLIENTS_TLS" envDefault:"false"` - ServerCert string `toml:"server_cert" env:"MG_PROVISION_SERVER_CERT" envDefault:""` - ServerKey string `toml:"server_key" env:"MG_PROVISION_SERVER_KEY" envDefault:""` - ClientsURL string `toml:"clients_url" env:"MG_PROVISION_CLIENTS_URL" envDefault:"http://localhost"` - ChannelsURL string `toml:"channels_url" env:"MG_PROVISION_CHANNELS_URL" envDefault:"http://localhost"` - UsersURL string `toml:"users_url" env:"MG_PROVISION_USERS_URL" envDefault:"http://localhost"` - CertsURL string `toml:"certs_url" env:"MG_PROVISION_CERTS_URL" envDefault:"http://localhost"` - MgEmail string `toml:"mg_email" env:"MG_PROVISION_EMAIL" envDefault:"test@example.com"` - MgUsername string `toml:"mg_username" env:"MG_PROVISION_USERNAME" envDefault:"user"` - MgPass string `toml:"mg_pass" env:"MG_PROVISION_PASS" envDefault:"test"` - MgDomainID string `toml:"mg_domain_id" env:"MG_PROVISION_DOMAIN_ID" envDefault:""` - MgAPIKey string `toml:"mg_api_key" env:"MG_PROVISION_API_KEY" envDefault:""` - MgBSURL string `toml:"mg_bs_url" env:"MG_PROVISION_BS_SVC_URL" envDefault:"http://localhost:9000"` -} - -// Bootstrap represetns the Bootstrap config. -type Bootstrap struct { - X509Provision bool `toml:"x509_provision" env:"MG_PROVISION_X509_PROVISIONING" envDefault:"false"` - Provision bool `toml:"provision" env:"MG_PROVISION_BS_CONFIG_PROVISIONING" envDefault:"true"` - AutoWhiteList bool `toml:"autowhite_list" env:"MG_PROVISION_BS_AUTO_WHITELIST" envDefault:"true"` - ProfileID string `toml:"profile_id" env:"MG_PROVISION_BS_PROFILE_ID" envDefault:""` - RenderContext map[string]any `toml:"render_context,omitempty"` - Bindings []BootstrapBinding `toml:"bindings,omitempty"` - Content map[string]any `toml:"content"` -} - -// BootstrapBinding maps a bootstrap profile slot to one of the resources -// created by provision. -type BootstrapBinding struct { - Slot string `toml:"slot" json:"slot"` - Type string `toml:"type" json:"type"` - Name string `toml:"name" json:"name,omitempty"` - MetadataKey string `toml:"metadata_key" json:"metadata_key,omitempty"` - MetadataValue string `toml:"metadata_value" json:"metadata_value,omitempty"` -} - -// Gateway represetns the Gateway config. -type Gateway struct { - Type string `toml:"type" json:"type"` - ExternalID string `toml:"external_id" json:"external_id"` - ExternalKey string `toml:"external_key" json:"external_key"` - CtrlChannelID string `toml:"ctrl_channel_id" json:"ctrl_channel_id"` - DataChannelID string `toml:"data_channel_id" json:"data_channel_id"` - ExportChannelID string `toml:"export_channel_id" json:"export_channel_id"` - CfgID string `toml:"cfg_id" json:"cfg_id"` -} - -// Cert represetns the certificate config. -type Cert struct { - TTL string `json:"ttl" toml:"ttl" env:"MG_PROVISION_CERTS_HOURS_VALID" envDefault:"2400h"` -} - -// Config struct of Provision. -type Config struct { - File string `toml:"file" env:"MG_PROVISION_CONFIG_FILE" envDefault:"config.toml"` - Server ServiceConf `toml:"server" mapstructure:"server"` - Bootstrap Bootstrap `toml:"bootstrap" mapstructure:"bootstrap"` - Clients []clients.Client `toml:"clients" mapstructure:"clients"` - Channels []channels.Channel `toml:"channels" mapstructure:"channels"` - Cert Cert `toml:"cert" mapstructure:"cert"` - BSContent string `env:"MG_PROVISION_BS_CONTENT" envDefault:""` - SendTelemetry bool `env:"MG_SEND_TELEMETRY" envDefault:"true"` - InstanceID string `env:"MG_MQTT_ADAPTER_INSTANCE_ID" envDefault:""` -} - -// Save - store config in a file. -func Save(c Config, file string) error { - if file == "" { - return errors.ErrEmptyPath - } - - b, err := toml.Marshal(c) - if err != nil { - return errors.Wrap(errFailedToReadConfig, err) - } - if err := os.WriteFile(file, b, 0o644); err != nil { - return fmt.Errorf("Error writing toml: %w", err) - } - - return nil -} - -// Read - retrieve config from a file. -func Read(file string) (Config, error) { - data, err := os.ReadFile(file) - if err != nil { - return Config{}, errors.Wrap(errFailedToReadConfig, err) - } - - var c Config - if err := toml.Unmarshal(data, &c); err != nil { - return Config{}, fmt.Errorf("Error unmarshaling toml: %w", err) - } - if len(c.Bootstrap.RenderContext) == 0 { - c.Bootstrap.RenderContext = nil - } - - return c, nil -} diff --git a/provision/config_test.go b/provision/config_test.go deleted file mode 100644 index a00bfc5dc..000000000 --- a/provision/config_test.go +++ /dev/null @@ -1,229 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package provision_test - -import ( - "fmt" - "os" - "testing" - - "github.com/absmach/magistrala/channels" - "github.com/absmach/magistrala/clients" - "github.com/absmach/magistrala/pkg/connections" - "github.com/absmach/magistrala/pkg/errors" - "github.com/absmach/magistrala/provision" - "github.com/pelletier/go-toml" - "github.com/stretchr/testify/assert" -) - -var ( - validConfig = provision.Config{ - Server: provision.ServiceConf{ - Port: "9016", - LogLevel: "info", - TLS: false, - }, - Bootstrap: provision.Bootstrap{ - X509Provision: true, - Provision: true, - AutoWhiteList: true, - Content: map[string]any{ - "test": "test", - }, - }, - Clients: []clients.Client{ - { - ID: "1234567890", - Name: "test", - Tags: []string{"test"}, - Metadata: map[string]any{ - "test": "test", - }, - PrivateMetadata: clients.Metadata{}, - Actions: []string{}, - AccessProviderRoleActions: []string{}, - ConnectionTypes: []connections.ConnType{}, - }, - }, - Channels: []channels.Channel{ - { - ID: "1234567890", - Name: "test", - Tags: []string{"test"}, - Metadata: map[string]any{ - "test": "test", - }, - Actions: []string{}, - AccessProviderRoleActions: []string{}, - ConnectionTypes: []connections.ConnType{}, - }, - }, - Cert: provision.Cert{}, - SendTelemetry: true, - InstanceID: "1234567890", - } - validConfigFile = "./config.toml" - invalidConfig = provision.Config{ - Bootstrap: provision.Bootstrap{ - Content: map[string]any{ - "invalid": make(chan int), - }, - }, - } - invalidConfigFile = "./invalid.toml" -) - -func createInvalidConfigFile() error { - config := map[string]any{ - "invalid": "invalid", - } - b, err := toml.Marshal(config) - if err != nil { - return err - } - - f, err := os.Create(invalidConfigFile) - if err != nil { - return err - } - - if _, err = f.Write(b); err != nil { - return err - } - - return nil -} - -func createValidConfigFile() error { - b, err := toml.Marshal(validConfig) - if err != nil { - return err - } - - f, err := os.Create(validConfigFile) - if err != nil { - return err - } - - if _, err = f.Write(b); err != nil { - return err - } - - return nil -} - -func TestSave(t *testing.T) { - cases := []struct { - desc string - cfg provision.Config - file string - err error - }{ - { - desc: "save valid config", - cfg: validConfig, - file: validConfigFile, - err: nil, - }, - { - desc: "save valid config with empty file name", - cfg: validConfig, - file: "", - err: errors.ErrEmptyPath, - }, - { - desc: "save empty config with valid config file", - cfg: provision.Config{}, - file: validConfigFile, - err: nil, - }, - { - desc: "save empty config with empty file name", - cfg: provision.Config{}, - file: "", - err: errors.ErrEmptyPath, - }, - { - desc: "save invalid config", - cfg: invalidConfig, - file: invalidConfigFile, - err: errors.New("failed to read config file"), - }, - } - - for _, c := range cases { - t.Run(c.desc, func(t *testing.T) { - err := provision.Save(c.cfg, c.file) - assert.True(t, errors.Contains(err, c.err), fmt.Sprintf("expected: %v, got: %v", c.err, err)) - - if err == nil { - defer func() { - if c.file != "" { - err := os.Remove(c.file) - assert.NoError(t, err) - } - }() - - cfg, err := provision.Read(c.file) - if c.cfg.Bootstrap.Content == nil { - c.cfg.Bootstrap.Content = map[string]any{} - } - assert.Equal(t, c.err, err) - assert.Equal(t, c.cfg, cfg) - } - }) - } -} - -func TestRead(t *testing.T) { - err := createInvalidConfigFile() - assert.NoError(t, err) - - err = createValidConfigFile() - assert.NoError(t, err) - - t.Cleanup(func() { - err := os.Remove(invalidConfigFile) - assert.NoError(t, err) - err = os.Remove(validConfigFile) - assert.NoError(t, err) - }) - - cases := []struct { - desc string - file string - cfg provision.Config - err error - }{ - { - desc: "read valid config", - file: validConfigFile, - cfg: validConfig, - err: nil, - }, - { - desc: "read invalid config", - file: invalidConfigFile, - cfg: invalidConfig, - err: nil, - }, - { - desc: "read empty config", - file: "", - cfg: provision.Config{}, - err: errors.New("failed to read config file"), - }, - } - - for _, c := range cases { - t.Run(c.desc, func(t *testing.T) { - cfg, err := provision.Read(c.file) - if c.desc == "read invalid config" { - c.cfg.Bootstrap.Content = nil - } - assert.True(t, errors.Contains(err, c.err), fmt.Sprintf("expected: %v, got: %v", c.err, err)) - assert.Equal(t, c.cfg, cfg) - }) - } -} diff --git a/provision/configs/config.toml b/provision/configs/config.toml deleted file mode 100644 index 650ed3518..000000000 --- a/provision/configs/config.toml +++ /dev/null @@ -1,47 +0,0 @@ -# Copyright (c) Abstract Machines -# SPDX-License-Identifier: Apache-2.0 - -file = "config.toml" - -[bootstrap] - autowhite_list = true - content = "" - provision = true - x509_provision = false - - -[server] - LogLevel = "info" - ca_certs = "" - http_port = "8190" - mg_api_key = "" - mg_bs_url = "http://localhost:9013" - mg_certs_url = "http://localhost:9019" - mg_pass = "" - mg_user = "" - mqtt_url = "" - port = "" - server_cert = "" - server_key = "" - clients_location = "http://localhost:9006" - tls = true - users_location = "" - -[[clients]] - name = "client" - - [client.metadata] - external_id = "xxxxxx" - - -[[channels]] - name = "control-channel" - - [channels.metadata] - type = "control" - -[[channels]] - name = "data-channel" - - [channels.metadata] - type = "data" diff --git a/provision/doc.go b/provision/doc.go deleted file mode 100644 index e9b855294..000000000 --- a/provision/doc.go +++ /dev/null @@ -1,6 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package provision contains domain concept definitions needed to support -// Provision service feature, i.e. automate provision process. -package provision diff --git a/provision/middleware/logging.go b/provision/middleware/logging.go deleted file mode 100644 index 136aa99c1..000000000 --- a/provision/middleware/logging.go +++ /dev/null @@ -1,71 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package middleware - -import ( - "context" - "log/slog" - "time" - - "github.com/absmach/magistrala/provision" -) - -var _ provision.Service = (*loggingMiddleware)(nil) - -type loggingMiddleware struct { - logger *slog.Logger - svc provision.Service -} - -// NewLogging adds logging facilities to the core service. -func NewLogging(svc provision.Service, logger *slog.Logger) provision.Service { - return &loggingMiddleware{logger, svc} -} - -func (lm *loggingMiddleware) Provision(ctx context.Context, domainID, token, name, externalID, externalKey string) (res provision.Result, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("name", name), - slog.String("external_id", externalID), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Provision failed", args...) - return - } - lm.logger.Info("Provision completed successfully", args...) - }(time.Now()) - - return lm.svc.Provision(ctx, domainID, token, name, externalID, externalKey) -} - -func (lm *loggingMiddleware) Cert(ctx context.Context, domainID, token, clientID, duration string) (cert, key string, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("client_id", clientID), - slog.String("ttl", duration), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Client certificate creation failed", args...) - return - } - lm.logger.Info("Client certificate created successfully", args...) - }(time.Now()) - - return lm.svc.Cert(ctx, domainID, token, clientID, duration) -} - -func (lm *loggingMiddleware) Mapping() (res map[string]any) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - } - lm.logger.Info("Mapping completed successfully", args...) - }(time.Now()) - - return lm.svc.Mapping() -} diff --git a/provision/mocks/service.go b/provision/mocks/service.go deleted file mode 100644 index 0b6cde9f9..000000000 --- a/provision/mocks/service.go +++ /dev/null @@ -1,269 +0,0 @@ -// Copyright (c) Abstract Machines - -// SPDX-License-Identifier: Apache-2.0 - -// Code generated by mockery; DO NOT EDIT. -// github.com/vektra/mockery -// template: testify - -package mocks - -import ( - "context" - - "github.com/absmach/magistrala/provision" - mock "github.com/stretchr/testify/mock" -) - -// NewService creates a new instance of Service. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. -// The first argument is typically a *testing.T value. -func NewService(t interface { - mock.TestingT - Cleanup(func()) -}) *Service { - mock := &Service{} - mock.Mock.Test(t) - - t.Cleanup(func() { mock.AssertExpectations(t) }) - - return mock -} - -// Service is an autogenerated mock type for the Service type -type Service struct { - mock.Mock -} - -type Service_Expecter struct { - mock *mock.Mock -} - -func (_m *Service) EXPECT() *Service_Expecter { - return &Service_Expecter{mock: &_m.Mock} -} - -// Cert provides a mock function for the type Service -func (_mock *Service) Cert(ctx context.Context, domainID string, token string, clientID string, duration string) (string, string, error) { - ret := _mock.Called(ctx, domainID, token, clientID, duration) - - if len(ret) == 0 { - panic("no return value specified for Cert") - } - - var r0 string - var r1 string - var r2 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string, string, string) (string, string, error)); ok { - return returnFunc(ctx, domainID, token, clientID, duration) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string, string, string) string); ok { - r0 = returnFunc(ctx, domainID, token, clientID, duration) - } else { - r0 = ret.Get(0).(string) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, string, string, string) string); ok { - r1 = returnFunc(ctx, domainID, token, clientID, duration) - } else { - r1 = ret.Get(1).(string) - } - if returnFunc, ok := ret.Get(2).(func(context.Context, string, string, string, string) error); ok { - r2 = returnFunc(ctx, domainID, token, clientID, duration) - } else { - r2 = ret.Error(2) - } - return r0, r1, r2 -} - -// Service_Cert_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Cert' -type Service_Cert_Call struct { - *mock.Call -} - -// Cert is a helper method to define mock.On call -// - ctx context.Context -// - domainID string -// - token string -// - clientID string -// - duration string -func (_e *Service_Expecter) Cert(ctx interface{}, domainID interface{}, token interface{}, clientID interface{}, duration interface{}) *Service_Cert_Call { - return &Service_Cert_Call{Call: _e.mock.On("Cert", ctx, domainID, token, clientID, duration)} -} - -func (_c *Service_Cert_Call) Run(run func(ctx context.Context, domainID string, token string, clientID string, duration string)) *Service_Cert_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 string - if args[4] != nil { - arg4 = args[4].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_Cert_Call) Return(s string, s1 string, err error) *Service_Cert_Call { - _c.Call.Return(s, s1, err) - return _c -} - -func (_c *Service_Cert_Call) RunAndReturn(run func(ctx context.Context, domainID string, token string, clientID string, duration string) (string, string, error)) *Service_Cert_Call { - _c.Call.Return(run) - return _c -} - -// Mapping provides a mock function for the type Service -func (_mock *Service) Mapping() map[string]any { - ret := _mock.Called() - - if len(ret) == 0 { - panic("no return value specified for Mapping") - } - - var r0 map[string]any - if returnFunc, ok := ret.Get(0).(func() map[string]any); ok { - r0 = returnFunc() - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).(map[string]any) - } - } - return r0 -} - -// Service_Mapping_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Mapping' -type Service_Mapping_Call struct { - *mock.Call -} - -// Mapping is a helper method to define mock.On call -func (_e *Service_Expecter) Mapping() *Service_Mapping_Call { - return &Service_Mapping_Call{Call: _e.mock.On("Mapping")} -} - -func (_c *Service_Mapping_Call) Run(run func()) *Service_Mapping_Call { - _c.Call.Run(func(args mock.Arguments) { - run() - }) - return _c -} - -func (_c *Service_Mapping_Call) Return(stringToV map[string]any) *Service_Mapping_Call { - _c.Call.Return(stringToV) - return _c -} - -func (_c *Service_Mapping_Call) RunAndReturn(run func() map[string]any) *Service_Mapping_Call { - _c.Call.Return(run) - return _c -} - -// Provision provides a mock function for the type Service -func (_mock *Service) Provision(ctx context.Context, domainID string, token string, name string, externalID string, externalKey string) (provision.Result, error) { - ret := _mock.Called(ctx, domainID, token, name, externalID, externalKey) - - if len(ret) == 0 { - panic("no return value specified for Provision") - } - - var r0 provision.Result - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string, string, string, string) (provision.Result, error)); ok { - return returnFunc(ctx, domainID, token, name, externalID, externalKey) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string, string, string, string) provision.Result); ok { - r0 = returnFunc(ctx, domainID, token, name, externalID, externalKey) - } else { - r0 = ret.Get(0).(provision.Result) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, string, string, string, string) error); ok { - r1 = returnFunc(ctx, domainID, token, name, externalID, externalKey) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_Provision_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Provision' -type Service_Provision_Call struct { - *mock.Call -} - -// Provision is a helper method to define mock.On call -// - ctx context.Context -// - domainID string -// - token string -// - name string -// - externalID string -// - externalKey string -func (_e *Service_Expecter) Provision(ctx interface{}, domainID interface{}, token interface{}, name interface{}, externalID interface{}, externalKey interface{}) *Service_Provision_Call { - return &Service_Provision_Call{Call: _e.mock.On("Provision", ctx, domainID, token, name, externalID, externalKey)} -} - -func (_c *Service_Provision_Call) Run(run func(ctx context.Context, domainID string, token string, name string, externalID string, externalKey string)) *Service_Provision_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 string - if args[4] != nil { - arg4 = args[4].(string) - } - var arg5 string - if args[5] != nil { - arg5 = args[5].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - arg5, - ) - }) - return _c -} - -func (_c *Service_Provision_Call) Return(result provision.Result, err error) *Service_Provision_Call { - _c.Call.Return(result, err) - return _c -} - -func (_c *Service_Provision_Call) RunAndReturn(run func(ctx context.Context, domainID string, token string, name string, externalID string, externalKey string) (provision.Result, error)) *Service_Provision_Call { - _c.Call.Return(run) - return _c -} diff --git a/provision/service.go b/provision/service.go deleted file mode 100644 index e1afad771..000000000 --- a/provision/service.go +++ /dev/null @@ -1,504 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package provision - -import ( - "context" - "encoding/json" - "fmt" - "log/slog" - - "github.com/absmach/magistrala/pkg/errors" - "github.com/absmach/magistrala/pkg/sdk" - smqSDK "github.com/absmach/magistrala/pkg/sdk" -) - -const ( - externalIDKey = "external_id" - gateway = "gateway" - - control = "control" - data = "data" - export = "export" -) - -var ( - ErrUnauthorized = errors.NewAuthNError("unauthorized access") - ErrFailedToCreateToken = errors.NewAuthNError("failed to create access token") - ErrEmptyClientsList = errors.NewRequestError("clients list in configuration empty") - ErrClientUpdate = errors.NewRequestError("failed to update client") - ErrEmptyChannelsList = errors.NewRequestError("channels list in configuration is empty") - ErrFailedChannelCreation = errors.NewRequestError("failed to create channel") - ErrFailedChannelRetrieval = errors.NewRequestError("failed to retrieve channel") - ErrFailedClientCreation = errors.NewRequestError("failed to create client") - ErrFailedClientRetrieval = errors.NewRequestError("failed to retrieve client") - ErrMissingCredentials = errors.NewRequestError("missing credentials") - ErrFailedBootstrapRetrieval = errors.NewServiceError("failed to retrieve bootstrap") - ErrFailedCertCreation = errors.NewServiceError("failed to create certificates") - ErrFailedCertView = errors.NewServiceError("failed to view certificate") - ErrFailedBootstrap = errors.NewServiceError("failed to create bootstrap config") - ErrFailedBootstrapValidate = errors.NewServiceError("failed to validate bootstrap config creation") - ErrFailedBootstrapBinding = errors.NewServiceError("failed to bind bootstrap resources") - ErrGatewayUpdate = errors.NewServiceError("failed to update gateway metadata") -) - -var _ Service = (*provisionService)(nil) - -// Service specifies Provision service API. -type Service interface { - // Provision is the only method this API specifies. Depending on the configuration, - // the following actions will can be executed: - // - create a Client based on external_id (eg. MAC address) - // - create multiple Channels - // - create Bootstrap configuration - // - enable created Bootstrap enrollments - Provision(ctx context.Context, domainID, token, name, externalID, externalKey string) (Result, error) - - // Mapping returns current configuration used for provision - // useful for using in ui to create configuration that matches - // one created with Provision method. - Mapping() map[string]any - - // Certs creates certificate for clients that communicate over mTLS - // A duration string is a possibly signed sequence of decimal numbers, - // each with optional fraction and a unit suffix, such as "300ms", "-1.5h" or "2h45m". - // Valid time units are "ns", "us" (or "µs"), "ms", "s", "m", "h". - Cert(ctx context.Context, domainID, token, clientID, duration string) (string, string, error) -} - -type provisionService struct { - logger *slog.Logger - sdk sdk.SDK - conf Config -} - -// Result represent what is created with additional info. -type Result struct { - Clients []smqSDK.Client `json:"clients,omitempty"` - Channels []smqSDK.Channel `json:"channels,omitempty"` - ClientCert map[string]string `json:"client_cert,omitempty"` - ClientKey map[string]string `json:"client_key,omitempty"` - CACert string `json:"ca_cert,omitempty"` - Whitelisted map[string]bool `json:"whitelisted,omitempty"` - Error string `json:"error,omitempty"` -} - -// New returns new provision service. -func New(cfg Config, mgsdk sdk.SDK, logger *slog.Logger) Service { - return &provisionService{ - logger: logger, - conf: cfg, - sdk: mgsdk, - } -} - -// Mapping retrieves current configuration. -func (ps *provisionService) Mapping() map[string]any { - return ps.conf.Bootstrap.Content -} - -// Provision is provision method for creating setup according to -// provision layout specified in config.toml. -func (ps *provisionService) Provision(ctx context.Context, domainID, token, name, externalID, externalKey string) (res Result, err error) { - var channels []smqSDK.Channel - var clients []smqSDK.Client - var bootstrapIDs []string - defer ps.recover(ctx, &err, &clients, &channels, &bootstrapIDs, domainID, token) - - token, err = ps.createTokenIfEmpty(ctx, token) - if err != nil { - return res, errors.Wrap(ErrFailedToCreateToken, err) - } - - if len(ps.conf.Clients) == 0 { - return res, ErrEmptyClientsList - } - if len(ps.conf.Channels) == 0 { - return res, ErrEmptyChannelsList - } - for _, c := range ps.conf.Clients { - // If client in configs contains metadata with external_id - // set value for it from the provision request - if _, ok := c.Metadata[externalIDKey]; ok { - c.Metadata[externalIDKey] = externalID - } - - cli := smqSDK.Client{ - Metadata: c.Metadata, - } - if name == "" { - name = c.Name - } - cli.Name = name - cli, err := ps.sdk.CreateClient(ctx, cli, domainID, token) - if err != nil { - res.Error = err.Error() - return res, errors.Wrap(ErrFailedClientCreation, err) - } - - // Get newly created client (in order to get the key). - cli, err = ps.sdk.Client(ctx, cli.ID, domainID, token) - if err != nil { - e := errors.Wrap(err, fmt.Errorf("client id: %s", cli.ID)) - return res, errors.Wrap(ErrFailedClientRetrieval, e) - } - clients = append(clients, cli) - } - - for _, channel := range ps.conf.Channels { - ch := smqSDK.Channel{ - Name: name + "_" + channel.Name, - Metadata: smqSDK.Metadata(channel.Metadata), - } - ch, err := ps.sdk.CreateChannel(ctx, ch, domainID, token) - if err != nil { - return res, errors.Wrap(ErrFailedChannelCreation, err) - } - ch, err = ps.sdk.Channel(ctx, ch.ID, domainID, token) - if err != nil { - e := errors.Wrap(err, fmt.Errorf("channel id: %s", ch.ID)) - return res, errors.Wrap(ErrFailedChannelRetrieval, e) - } - channels = append(channels, ch) - } - - res = Result{ - Clients: clients, - Channels: channels, - Whitelisted: map[string]bool{}, - ClientCert: map[string]string{}, - ClientKey: map[string]string{}, - } - - content, err := json.Marshal(ps.conf.Bootstrap.Content) - if err != nil { - return Result{}, errors.Wrap(ErrFailedBootstrap, err) - } - - bootstrapConfigs := make(map[string]sdk.BootstrapConfig) - var gatewayConfig sdk.BootstrapConfig - var gatewayClientID string - for _, c := range clients { - if ps.conf.Bootstrap.Provision && needsBootstrap(c) { - bsReq := sdk.BootstrapConfig{ - ExternalID: externalID, - ExternalKey: externalKey, - Name: name, - CACert: res.CACert, - ClientCert: "", - ClientKey: "", - Content: string(content), - ProfileID: ps.conf.Bootstrap.ProfileID, - RenderContext: ps.bootstrapRenderContext(externalID, name), - } - bsid, err := ps.sdk.AddBootstrap(ctx, bsReq, domainID, token) - if err != nil { - return Result{}, errors.Wrap(ErrFailedBootstrap, err) - } - bootstrapIDs = append(bootstrapIDs, bsid) - - bsConfig, err := ps.sdk.ViewBootstrap(ctx, bsid, domainID, token) - if err != nil { - return Result{}, errors.Wrap(ErrFailedBootstrapValidate, err) - } - bootstrapConfigs[c.ID] = bsConfig - gatewayConfig = bsConfig - gatewayClientID = c.ID - - if err := ps.bindBootstrapResources(ctx, bsConfig.ID, c, clients, channels, domainID, token); err != nil { - return Result{}, errors.Wrap(ErrFailedBootstrapBinding, err) - } - } - - if ps.conf.Bootstrap.X509Provision { - var cert smqSDK.Certificate - - cert, err = ps.sdk.IssueCert(ctx, c.ID, ps.conf.Cert.TTL, nil, smqSDK.Options{}, domainID, token) - if err != nil { - e := errors.Wrap(err, fmt.Errorf("client id: %s", c.ID)) - return res, errors.Wrap(ErrFailedCertCreation, e) - } - cert, err := ps.sdk.ViewCert(ctx, cert.SerialNumber, domainID, token) - if err != nil { - return res, errors.Wrap(ErrFailedCertView, err) - } - - res.ClientCert[c.ID] = cert.Certificate - res.ClientKey[c.ID] = cert.Key - res.CACert = "" - - if bsConfig, ok := bootstrapConfigs[c.ID]; ok { - updated, err := ps.sdk.UpdateBootstrapCerts(ctx, bsConfig.ID, cert.Certificate, cert.Key, "", domainID, token) - if err != nil { - return Result{}, errors.Wrap(ErrFailedCertCreation, err) - } - bootstrapConfigs[c.ID] = updated - if gatewayClientID == c.ID { - gatewayConfig = updated - } - } - } - - if ps.conf.Bootstrap.AutoWhiteList { - if bsConfig, ok := bootstrapConfigs[c.ID]; ok { - if err := ps.sdk.Whitelist(ctx, bsConfig.ID, smqSDK.BootstrapEnabledStatus, domainID, token); err != nil { - res.Error = err.Error() - return res, ErrClientUpdate - } - res.Whitelisted[bsConfig.ID] = true - } - } - } - - if gatewayClientID != "" && gatewayConfig.ID != "" { - if err = ps.updateGateway(ctx, domainID, token, gatewayClientID, gatewayConfig, externalKey, channels); err != nil { - return res, err - } - } - return res, nil -} - -func (ps *provisionService) Cert(ctx context.Context, domainID, token, clientID, ttl string) (string, string, error) { - token, err := ps.createTokenIfEmpty(ctx, token) - if err != nil { - return "", "", errors.Wrap(ErrFailedToCreateToken, err) - } - - c, err := ps.sdk.Client(ctx, clientID, domainID, token) - if err != nil { - return "", "", errors.Wrap(ErrUnauthorized, err) - } - cert, err := ps.sdk.IssueCert(ctx, c.ID, ps.conf.Cert.TTL, []string{}, smqSDK.Options{}, domainID, token) - if err != nil { - return "", "", errors.Wrap(ErrFailedCertCreation, err) - } - cert, err = ps.sdk.ViewCert(ctx, cert.SerialNumber, domainID, token) - if err != nil { - return "", "", errors.Wrap(ErrFailedCertView, err) - } - return cert.Certificate, cert.Key, err -} - -func (ps *provisionService) createTokenIfEmpty(ctx context.Context, token string) (string, error) { - if token != "" { - return token, nil - } - - // If no token in request is provided - // use API key provided in config file or env - if ps.conf.Server.MgAPIKey != "" { - return ps.conf.Server.MgAPIKey, nil - } - - // If no API key use username and password provided to create access token. - if ps.conf.Server.MgUsername == "" || ps.conf.Server.MgPass == "" { - return token, ErrMissingCredentials - } - - u := smqSDK.Login{ - Username: ps.conf.Server.MgUsername, - Password: ps.conf.Server.MgPass, - } - tkn, err := ps.sdk.CreateToken(ctx, u) - if err != nil { - return token, errors.Wrap(ErrFailedToCreateToken, err) - } - - return tkn.AccessToken, nil -} - -func (ps *provisionService) updateGateway(ctx context.Context, domainID, token, gatewayClientID string, bs sdk.BootstrapConfig, externalKey string, channels []smqSDK.Channel) error { - var gw Gateway - for _, ch := range channels { - switch ch.Metadata["type"] { - case control: - gw.CtrlChannelID = ch.ID - case data: - gw.DataChannelID = ch.ID - case export: - gw.ExportChannelID = ch.ID - } - } - gw.ExternalID = bs.ExternalID - gw.ExternalKey = externalKey - gw.CfgID = bs.ID - gw.Type = gateway - - c, sdkerr := ps.sdk.Client(ctx, gatewayClientID, domainID, token) - if sdkerr != nil { - return errors.Wrap(ErrGatewayUpdate, sdkerr) - } - b, err := json.Marshal(gw) - if err != nil { - return errors.Wrap(ErrGatewayUpdate, err) - } - if err := json.Unmarshal(b, &c.Metadata); err != nil { - return errors.Wrap(ErrGatewayUpdate, err) - } - if _, err := ps.sdk.UpdateClient(ctx, c, domainID, token); err != nil { - return errors.Wrap(ErrGatewayUpdate, err) - } - return nil -} - -func (ps *provisionService) bootstrapRenderContext(externalID, name string) map[string]any { - renderContext := make(map[string]any, len(ps.conf.Bootstrap.RenderContext)+2) - for k, v := range ps.conf.Bootstrap.RenderContext { - renderContext[k] = v - } - if externalID != "" { - renderContext[externalIDKey] = externalID - } - if name != "" { - renderContext["name"] = name - } - if len(renderContext) == 0 { - return nil - } - return renderContext -} - -func (ps *provisionService) bindBootstrapResources(ctx context.Context, configID string, bootstrapClient smqSDK.Client, clients []smqSDK.Client, channels []smqSDK.Channel, domainID, token string) error { - if len(ps.conf.Bootstrap.Bindings) == 0 { - return nil - } - - requests := make([]smqSDK.BootstrapBindingRequest, 0, len(ps.conf.Bootstrap.Bindings)) - for _, binding := range ps.conf.Bootstrap.Bindings { - resourceID := ps.bindingResourceID(binding, bootstrapClient, clients, channels) - if resourceID == "" { - return fmt.Errorf("resource for bootstrap binding slot %q not found", binding.Slot) - } - requests = append(requests, smqSDK.BootstrapBindingRequest{ - Slot: binding.Slot, - Type: binding.Type, - ResourceID: resourceID, - }) - } - - return ps.sdk.BindBootstrapResources(ctx, configID, requests, domainID, token) -} - -func (ps *provisionService) bindingResourceID(binding BootstrapBinding, bootstrapClient smqSDK.Client, clients []smqSDK.Client, channels []smqSDK.Channel) string { - switch binding.Type { - case "client": - if matchesClientBinding(binding, bootstrapClient) { - return bootstrapClient.ID - } - for _, client := range clients { - if matchesClientBinding(binding, client) { - return client.ID - } - } - case "channel": - for _, channel := range channels { - if matchesChannelBinding(binding, channel) { - return channel.ID - } - } - } - return "" -} - -func matchesClientBinding(binding BootstrapBinding, client smqSDK.Client) bool { - if binding.Name != "" && client.Name != binding.Name { - return false - } - if binding.MetadataKey != "" { - return metadataValue(client.Metadata, binding.MetadataKey) == binding.MetadataValue - } - return binding.Name != "" || client.ID != "" -} - -func matchesChannelBinding(binding BootstrapBinding, channel smqSDK.Channel) bool { - if binding.Name != "" && channel.Name != binding.Name { - return false - } - if binding.MetadataKey != "" { - return metadataValue(channel.Metadata, binding.MetadataKey) == binding.MetadataValue - } - return binding.Name != "" || channel.ID != "" -} - -func metadataValue(metadata map[string]any, key string) string { - if metadata == nil { - return "" - } - if value, ok := metadata[key]; ok { - return fmt.Sprint(value) - } - return "" -} - -func (ps *provisionService) errLog(err error) { - if err != nil { - ps.logger.Error(fmt.Sprintf("Error recovering: %s", err)) - } -} - -func clean(ctx context.Context, ps *provisionService, clients []smqSDK.Client, channels []smqSDK.Channel, domainID, token string) { - for _, t := range clients { - err := ps.sdk.DeleteClient(ctx, t.ID, domainID, token) - ps.errLog(err) - } - for _, c := range channels { - err := ps.sdk.DeleteChannel(ctx, c.ID, domainID, token) - ps.errLog(err) - } -} - -func (ps *provisionService) removeBootstraps(ctx context.Context, ids []string, domainID, token string) { - for _, id := range ids { - ps.errLog(ps.sdk.RemoveBootstrap(ctx, id, domainID, token)) - } -} - -func (ps *provisionService) recover(ctx context.Context, e *error, ths *[]smqSDK.Client, chs *[]smqSDK.Channel, bootstrapIDs *[]string, domainID, token string) { - if e == nil { - return - } - clients, channels, bootstraps, err := *ths, *chs, *bootstrapIDs, *e - - if errors.Contains(err, ErrFailedClientRetrieval) || errors.Contains(err, ErrFailedChannelCreation) { - for _, c := range clients { - err := ps.sdk.DeleteClient(ctx, c.ID, domainID, token) - ps.errLog(err) - } - return - } - - if errors.Contains(err, ErrFailedBootstrap) || errors.Contains(err, ErrFailedChannelRetrieval) { - clean(ctx, ps, clients, channels, domainID, token) - return - } - - if errors.Contains(err, ErrFailedBootstrapValidate) || errors.Contains(err, ErrFailedCertCreation) || errors.Contains(err, ErrFailedBootstrapBinding) { - clean(ctx, ps, clients, channels, domainID, token) - ps.removeBootstraps(ctx, bootstraps, domainID, token) - return - } - - if errors.Contains(err, ErrClientUpdate) || errors.Contains(err, ErrGatewayUpdate) { - clean(ctx, ps, clients, channels, domainID, token) - for _, c := range clients { - if ps.conf.Bootstrap.X509Provision && needsBootstrap(c) { - err := ps.sdk.RevokeCert(ctx, c.ID, domainID, token) - ps.errLog(err) - } - } - ps.removeBootstraps(ctx, bootstraps, domainID, token) - return - } -} - -func needsBootstrap(c smqSDK.Client) bool { - if c.Metadata == nil { - return false - } - - if _, ok := c.Metadata[externalIDKey]; ok { - return true - } - return false -} diff --git a/provision/service_test.go b/provision/service_test.go deleted file mode 100644 index 14d9fb9e3..000000000 --- a/provision/service_test.go +++ /dev/null @@ -1,449 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package provision_test - -import ( - "context" - "fmt" - "testing" - - "github.com/absmach/magistrala/internal/testsutil" - mglog "github.com/absmach/magistrala/logger" - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - smqSDK "github.com/absmach/magistrala/pkg/sdk" - sdkmocks "github.com/absmach/magistrala/pkg/sdk/mocks" - "github.com/absmach/magistrala/provision" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" -) - -var validToken = "valid" - -func TestMapping(t *testing.T) { - mgsdk := new(sdkmocks.SDK) - svc := provision.New(validConfig, mgsdk, mglog.NewMock()) - - cases := []struct { - desc string - content map[string]any - sdkerr error - err error - }{ - { - desc: "valid request", - content: validConfig.Bootstrap.Content, - sdkerr: nil, - err: nil, - }, - } - - for _, c := range cases { - t.Run(c.desc, func(t *testing.T) { - content := svc.Mapping() - assert.Equal(t, c.content, content) - }) - } -} - -func TestCert(t *testing.T) { - cases := []struct { - desc string - config provision.Config - domainID string - token string - returnedToken string - clientID string - ttl string - serial string - cert string - key string - sdkClientErr error - sdkCertErr error - sdkTokenErr error - err error - }{ - { - desc: "valid", - config: validConfig, - domainID: testsutil.GenerateUUID(t), - token: validToken, - clientID: testsutil.GenerateUUID(t), - ttl: "1h", - cert: "cert", - key: "key", - sdkClientErr: nil, - sdkCertErr: nil, - sdkTokenErr: nil, - err: nil, - }, - { - desc: "empty token with config API key", - config: provision.Config{ - Server: provision.ServiceConf{MgAPIKey: "key"}, - Cert: provision.Cert{TTL: "1h"}, - }, - domainID: testsutil.GenerateUUID(t), - token: "", - returnedToken: "key", - clientID: testsutil.GenerateUUID(t), - ttl: "1h", - cert: "cert", - key: "key", - sdkClientErr: nil, - sdkCertErr: nil, - sdkTokenErr: nil, - err: nil, - }, - { - desc: "empty token with username and password", - config: provision.Config{ - Server: provision.ServiceConf{ - MgUsername: "testUsername", - MgPass: "12345678", - MgDomainID: testsutil.GenerateUUID(t), - }, - Cert: provision.Cert{TTL: "1h"}, - }, - domainID: testsutil.GenerateUUID(t), - token: "", - returnedToken: validToken, - clientID: testsutil.GenerateUUID(t), - ttl: "1h", - cert: "cert", - key: "key", - sdkClientErr: nil, - sdkCertErr: nil, - sdkTokenErr: nil, - err: nil, - }, - { - desc: "empty token with username and invalid password", - config: provision.Config{ - Server: provision.ServiceConf{ - MgUsername: "testUsername", - MgPass: "12345678", - MgDomainID: testsutil.GenerateUUID(t), - }, - Cert: provision.Cert{TTL: "1h"}, - }, - domainID: testsutil.GenerateUUID(t), - token: "", - clientID: testsutil.GenerateUUID(t), - ttl: "1h", - cert: "", - key: "", - sdkClientErr: nil, - sdkCertErr: nil, - sdkTokenErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, 401), - err: provision.ErrFailedToCreateToken, - }, - { - desc: "empty token with empty username and password", - config: provision.Config{ - Server: provision.ServiceConf{}, - Cert: provision.Cert{TTL: "1h"}, - }, - domainID: testsutil.GenerateUUID(t), - token: "", - clientID: testsutil.GenerateUUID(t), - ttl: "1h", - cert: "", - key: "", - sdkClientErr: nil, - sdkCertErr: nil, - sdkTokenErr: nil, - err: provision.ErrMissingCredentials, - }, - { - desc: "invalid clientID", - config: validConfig, - domainID: testsutil.GenerateUUID(t), - token: "invalid", - clientID: testsutil.GenerateUUID(t), - ttl: "1h", - cert: "", - key: "", - sdkClientErr: errors.NewSDKErrorWithStatus(svcerr.ErrAuthentication, 401), - sdkCertErr: nil, - sdkTokenErr: nil, - err: provision.ErrUnauthorized, - }, - { - desc: "invalid clientID", - config: validConfig, - domainID: testsutil.GenerateUUID(t), - token: validToken, - clientID: "invalid", - ttl: "1h", - cert: "", - key: "", - sdkClientErr: errors.NewSDKErrorWithStatus(repoerr.ErrNotFound, 404), - sdkCertErr: nil, - sdkTokenErr: nil, - err: provision.ErrUnauthorized, - }, - { - desc: "failed to issue cert", - config: validConfig, - domainID: testsutil.GenerateUUID(t), - token: validToken, - clientID: testsutil.GenerateUUID(t), - ttl: "1h", - cert: "", - key: "", - sdkClientErr: nil, - sdkTokenErr: nil, - sdkCertErr: errors.NewSDKError(repoerr.ErrCreateEntity), - err: repoerr.ErrCreateEntity, - }, - } - mgsdk := new(sdkmocks.SDK) - for _, c := range cases { - t.Run(c.desc, func(t *testing.T) { - svc := provision.New(c.config, mgsdk, mglog.NewMock()) - - call1 := mgsdk.On("Client", mock.Anything, c.clientID, c.domainID, mock.Anything).Return(smqSDK.Client{ID: c.clientID}, c.sdkClientErr) - var call2 *mock.Call - switch c.token { - case "": - call2 = mgsdk.On("IssueCert", context.Background(), c.clientID, c.config.Cert.TTL, []string{}, smqSDK.Options{}, c.domainID, c.returnedToken).Return(smqSDK.Certificate{SerialNumber: c.serial}, c.sdkCertErr) - default: - call2 = mgsdk.On("IssueCert", context.Background(), c.clientID, c.config.Cert.TTL, []string{}, smqSDK.Options{}, c.domainID, c.token).Return(smqSDK.Certificate{SerialNumber: c.serial}, c.sdkCertErr) - } - call3 := mgsdk.On("ViewCert", mock.Anything, c.serial, mock.Anything, mock.Anything).Return(smqSDK.Certificate{Certificate: c.cert, Key: c.key}, c.sdkCertErr) - - login := smqSDK.Login{ - Username: c.config.Server.MgUsername, - Password: c.config.Server.MgPass, - } - call4 := mgsdk.On("CreateToken", mock.Anything, login).Return(smqSDK.Token{AccessToken: validToken}, c.sdkTokenErr) - cert, key, err := svc.Cert(context.Background(), c.domainID, c.token, c.clientID, c.ttl) - assert.Equal(t, c.cert, cert) - assert.Equal(t, c.key, key) - assert.True(t, errors.Contains(err, c.err), fmt.Sprintf("expected error %v, got %v", c.err, err)) - call1.Unset() - call2.Unset() - call3.Unset() - call4.Unset() - }) - } -} - -func TestProvisionUsesBootstrapEnrollmentID(t *testing.T) { - cfg := validConfig - cfg.Bootstrap = provision.Bootstrap{ - X509Provision: true, - Provision: true, - AutoWhiteList: true, - ProfileID: "gateway-profile", - RenderContext: map[string]any{ - "site": "warehouse-1", - }, - Bindings: []provision.BootstrapBinding{ - { - Slot: "mqtt_client", - Type: "client", - }, - { - Slot: "control", - Type: "channel", - MetadataKey: "type", - MetadataValue: "control", - }, - }, - Content: map[string]any{ - "broker": "mqtt://localhost:1883", - }, - } - cfg.Clients[0].Metadata = map[string]any{ - "external_id": "placeholder", - } - cfg.Channels[0].Name = "control-channel" - cfg.Channels[0].Metadata = map[string]any{ - "type": "control", - } - cfg.Cert.TTL = "1h" - - const ( - name = "gateway-1" - externalID = "AA:BB:CC:DD" - externalKey = "secret" - certPEM = "cert-pem" - keyPEM = "key-pem" - serial = "serial-1" - ) - - clientID := testsutil.GenerateUUID(t) - channelID := testsutil.GenerateUUID(t) - bootstrapID := testsutil.GenerateUUID(t) - domainID := testsutil.GenerateUUID(t) - - clientMetadata := map[string]any{ - "external_id": externalID, - } - var updatedClient smqSDK.Client - - mgsdk := new(sdkmocks.SDK) - svc := provision.New(cfg, mgsdk, mglog.NewMock()) - - createClientCall := mgsdk.On( - "CreateClient", - mock.Anything, - mock.Anything, - domainID, - validToken, - ).Return(smqSDK.Client{ID: clientID}, nil) - - clientCall := mgsdk.On( - "Client", - mock.Anything, - clientID, - domainID, - validToken, - ).Return(smqSDK.Client{ID: clientID, Name: name, Metadata: clientMetadata}, nil).Twice() - - createChannelCall := mgsdk.On( - "CreateChannel", - mock.Anything, - mock.Anything, - domainID, - validToken, - ).Return(smqSDK.Channel{ID: channelID}, nil) - - channelCall := mgsdk.On( - "Channel", - mock.Anything, - channelID, - domainID, - validToken, - ).Return(smqSDK.Channel{ID: channelID, Metadata: smqSDK.Metadata{"type": "control"}}, nil) - - addBootstrapCall := mgsdk.On( - "AddBootstrap", - mock.Anything, - mock.MatchedBy(func(cfg smqSDK.BootstrapConfig) bool { - return cfg.ProfileID == "gateway-profile" && - cfg.RenderContext["site"] == "warehouse-1" && - cfg.RenderContext["external_id"] == externalID && - cfg.RenderContext["name"] == name - }), - domainID, - validToken, - ).Return(bootstrapID, nil) - - viewBootstrapCall := mgsdk.On( - "ViewBootstrap", - mock.Anything, - bootstrapID, - domainID, - validToken, - ).Return(smqSDK.BootstrapConfig{ - ID: bootstrapID, - ExternalID: externalID, - ExternalKey: externalKey, - }, nil) - - bindBootstrapResourcesCall := mgsdk.On( - "BindBootstrapResources", - mock.Anything, - bootstrapID, - []smqSDK.BootstrapBindingRequest{ - { - Slot: "mqtt_client", - Type: "client", - ResourceID: clientID, - }, - { - Slot: "control", - Type: "channel", - ResourceID: channelID, - }, - }, - domainID, - validToken, - ).Return(nil) - - issueCertCall := mgsdk.On( - "IssueCert", - mock.Anything, - clientID, - cfg.Cert.TTL, - mock.Anything, - mock.Anything, - domainID, - validToken, - ).Return(smqSDK.Certificate{SerialNumber: serial}, nil) - - viewCertCall := mgsdk.On( - "ViewCert", - mock.Anything, - serial, - domainID, - validToken, - ).Return(smqSDK.Certificate{Certificate: certPEM, Key: keyPEM}, nil) - - updateBootstrapCertsCall := mgsdk.On( - "UpdateBootstrapCerts", - mock.Anything, - bootstrapID, - certPEM, - keyPEM, - "", - domainID, - validToken, - ).Return(smqSDK.BootstrapConfig{ - ID: bootstrapID, - ExternalID: externalID, - ExternalKey: externalKey, - ClientCert: certPEM, - ClientKey: keyPEM, - }, nil) - - whitelistCall := mgsdk.On( - "Whitelist", - mock.Anything, - bootstrapID, - smqSDK.BootstrapEnabledStatus, - domainID, - validToken, - ).Return(nil) - - updateClientCall := mgsdk.On( - "UpdateClient", - mock.Anything, - mock.Anything, - domainID, - validToken, - ).Run(func(args mock.Arguments) { - updatedClient = args.Get(1).(smqSDK.Client) - }).Return(smqSDK.Client{ID: clientID}, nil) - - res, err := svc.Provision(context.Background(), domainID, validToken, name, externalID, externalKey) - assert.NoError(t, err) - assert.Len(t, res.Clients, 1) - assert.Len(t, res.Channels, 1) - assert.True(t, res.Whitelisted[bootstrapID]) - assert.Equal(t, certPEM, res.ClientCert[clientID]) - assert.Equal(t, keyPEM, res.ClientKey[clientID]) - assert.Equal(t, clientID, updatedClient.ID) - assert.Equal(t, bootstrapID, updatedClient.Metadata["cfg_id"]) - assert.Equal(t, externalID, updatedClient.Metadata["external_id"]) - assert.Equal(t, channelID, updatedClient.Metadata["ctrl_channel_id"]) - assert.Equal(t, "gateway", updatedClient.Metadata["type"]) - - createClientCall.Unset() - clientCall.Unset() - createChannelCall.Unset() - channelCall.Unset() - _ = addBootstrapCall - viewBootstrapCall.Unset() - bindBootstrapResourcesCall.Unset() - issueCertCall.Unset() - viewCertCall.Unset() - updateBootstrapCertsCall.Unset() - whitelistCall.Unset() - updateClientCall.Unset() -} diff --git a/re/README.md b/re/README.md index ca93be223..4fff510e0 100644 --- a/re/README.md +++ b/re/README.md @@ -36,20 +36,15 @@ The service is configured using the following environment variables (values show | `MG_RE_DB_SSL_KEY` | PostgreSQL SSL client key | "" | | `MG_RE_DB_SSL_ROOT_CERT` | PostgreSQL SSL root cert | "" | -### Auth and domains gRPC +### Atom | Variable | Description | Default | | --- | --- | --- | -| `MG_AUTH_GRPC_URL` | Auth gRPC endpoint | `auth:7001` | -| `MG_AUTH_GRPC_TIMEOUT` | Auth gRPC timeout | `300s` | -| `MG_AUTH_GRPC_CLIENT_CERT` | Auth gRPC client cert path | `${GRPC_MTLS:+./ssl/certs/auth-grpc-client.crt}` | -| `MG_AUTH_GRPC_CLIENT_KEY` | Auth gRPC client key path | `${GRPC_MTLS:+./ssl/certs/auth-grpc-client.key}` | -| `MG_AUTH_GRPC_SERVER_CA_CERTS` | Auth gRPC server CA path | `${GRPC_MTLS:+./ssl/certs/ca.crt}` | -| `MG_DOMAINS_GRPC_URL` | Domains gRPC endpoint | `domains:7003` | -| `MG_DOMAINS_GRPC_TIMEOUT` | Domains gRPC timeout | `300s` | -| `MG_DOMAINS_GRPC_CLIENT_CERT` | Domains gRPC client cert path | `${GRPC_MTLS:+./ssl/certs/domains-grpc-client.crt}` | -| `MG_DOMAINS_GRPC_CLIENT_KEY` | Domains gRPC client key path | `${GRPC_MTLS:+./ssl/certs/domains-grpc-client.key}` | -| `MG_DOMAINS_GRPC_SERVER_CA_CERTS` | Domains gRPC server CA path | `${GRPC_MTLS:+./ssl/certs/ca.crt}` | +| `ATOM_URL` | Atom HTTP endpoint | `http://atom:8080` | +| `ATOM_JWKS_URL` | Atom JWKS endpoint for JWT verification | `http://atom:8080/.well-known/jwks.json` | +| `ATOM_ADMIN_USERNAME` | Atom admin login for service projections | `atom-admin` | +| `ATOM_ADMIN_SECRET` | Atom admin secret for service projections | `change-me` | +| `ATOM_TIMEOUT` | Atom request timeout | `5s` | | `MG_ALLOW_UNVERIFIED_USER` | Allow unverified users to access | `true` | ### Readers gRPC diff --git a/re/api/endpoints.go b/re/api/endpoints.go index fcc57f97b..04f2442f0 100644 --- a/re/api/endpoints.go +++ b/re/api/endpoints.go @@ -25,7 +25,7 @@ func addRuleEndpoint(s re.Service) endpoint.Endpoint { if err := req.validate(); err != nil { return addRuleRes{}, err } - rule, _, err := s.AddRule(ctx, session, req.Rule) + rule, err := s.AddRule(ctx, session, req.Rule) if err != nil { return addRuleRes{}, err } diff --git a/re/api/endpoints_test.go b/re/api/endpoints_test.go index 50a5acb4b..5bfbfd9c2 100644 --- a/re/api/endpoints_test.go +++ b/re/api/endpoints_test.go @@ -22,7 +22,6 @@ import ( authnmocks "github.com/absmach/magistrala/pkg/authn/mocks" "github.com/absmach/magistrala/pkg/errors" svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/pkg/roles" pkgSch "github.com/absmach/magistrala/pkg/schedule" "github.com/absmach/magistrala/re" "github.com/absmach/magistrala/re/api" @@ -237,7 +236,7 @@ func TestAddRuleEndpoint(t *testing.T) { } authCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.authnRes, tc.authnErr) - svcCall := svc.On("AddRule", mock.Anything, tc.authnRes, tc.rule).Return(tc.svcRes, []roles.RoleProvision{}, tc.svcErr) + svcCall := svc.On("AddRule", mock.Anything, tc.authnRes, tc.rule).Return(tc.svcRes, tc.svcErr) res, err := req.make() assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) diff --git a/re/api/transport.go b/re/api/transport.go index 3d9404268..44d09bbbb 100644 --- a/re/api/transport.go +++ b/re/api/transport.go @@ -15,7 +15,6 @@ import ( apiutil "github.com/absmach/magistrala/api/http/util" smqauthn "github.com/absmach/magistrala/pkg/authn" "github.com/absmach/magistrala/pkg/errors" - roleManagerHttp "github.com/absmach/magistrala/pkg/roles/rolemanager/api" "github.com/absmach/magistrala/re" "github.com/go-chi/chi/v5" kithttp "github.com/go-kit/kit/transport/http" @@ -37,8 +36,6 @@ func MakeHandler(svc re.Service, authn smqauthn.AuthNMiddleware, mux *chi.Mux, l r.Use(authn.WithOptions(smqauthn.WithDomainCheck(true)).Middleware()) r.Route("/{domainID}", func(r chi.Router) { r.Route("/rules", func(r chi.Router) { - d := roleManagerHttp.NewDecoder("ruleID") - r.Post("/", otelhttp.NewHandler(kithttp.NewServer( addRuleEndpoint(svc), decodeAddRuleRequest, @@ -53,8 +50,6 @@ func MakeHandler(svc re.Service, authn smqauthn.AuthNMiddleware, mux *chi.Mux, l opts..., ), "list_rules").ServeHTTP) - r = roleManagerHttp.EntityAvailableActionsRouter(svc, d, r, opts) - r.Route("/{ruleID}", func(r chi.Router) { r.Get("/", otelhttp.NewHandler(kithttp.NewServer( viewRuleEndpoint(svc), @@ -104,8 +99,6 @@ func MakeHandler(svc re.Service, authn smqauthn.AuthNMiddleware, mux *chi.Mux, l api.EncodeResponse, opts..., ), "disable_rule").ServeHTTP) - - roleManagerHttp.EntityRoleMangerRouter(svc, d, r, opts) }) }) }) diff --git a/re/atom.go b/re/atom.go new file mode 100644 index 000000000..2bba99dc1 --- /dev/null +++ b/re/atom.go @@ -0,0 +1,96 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package re + +import ( + "context" + + "github.com/absmach/magistrala/internal/atom" + "github.com/absmach/magistrala/pkg/authn" +) + +type atomService struct { + Service + projector atom.Projector +} + +func WithAtom(svc Service, projector atom.Projector) Service { + if projector == nil { + return svc + } + return atomService{Service: svc, projector: projector} +} + +func (svc atomService) AddRule(ctx context.Context, session authn.Session, r Rule) (Rule, error) { + rule, err := svc.Service.AddRule(ctx, session, r) + if err != nil { + return rule, err + } + if err := svc.projector.UpsertResource(ctx, ruleProjection(rule)); err != nil { + return rule, nil + } + return rule, nil +} + +func (svc atomService) UpdateRule(ctx context.Context, session authn.Session, r Rule) (Rule, error) { + rule, err := svc.Service.UpdateRule(ctx, session, r) + return svc.upsertAfterRuleChange(ctx, rule, err) +} + +func (svc atomService) UpdateRuleTags(ctx context.Context, session authn.Session, r Rule) (Rule, error) { + rule, err := svc.Service.UpdateRuleTags(ctx, session, r) + return svc.upsertAfterRuleChange(ctx, rule, err) +} + +func (svc atomService) UpdateRuleSchedule(ctx context.Context, session authn.Session, r Rule) (Rule, error) { + rule, err := svc.Service.UpdateRuleSchedule(ctx, session, r) + return svc.upsertAfterRuleChange(ctx, rule, err) +} + +func (svc atomService) EnableRule(ctx context.Context, session authn.Session, id string) (Rule, error) { + rule, err := svc.Service.EnableRule(ctx, session, id) + return svc.upsertAfterRuleChange(ctx, rule, err) +} + +func (svc atomService) DisableRule(ctx context.Context, session authn.Session, id string) (Rule, error) { + rule, err := svc.Service.DisableRule(ctx, session, id) + return svc.upsertAfterRuleChange(ctx, rule, err) +} + +func (svc atomService) RemoveRule(ctx context.Context, session authn.Session, id string) error { + if err := svc.Service.RemoveRule(ctx, session, id); err != nil { + return err + } + _ = svc.projector.DeleteResource(ctx, id) + return nil +} + +func (svc atomService) upsertAfterRuleChange(ctx context.Context, rule Rule, err error) (Rule, error) { + if err != nil { + return rule, err + } + if err := svc.projector.UpsertResource(ctx, ruleProjection(rule)); err != nil { + return rule, nil + } + return rule, nil +} + +func ruleProjection(r Rule) atom.Resource { + res := atom.ResourceFromFields(atom.ObjectFields{ + ID: r.ID, + Kind: atom.KindRule, + Name: r.Name, + TenantID: r.DomainID, + OwnerID: r.CreatedBy, + Status: r.Status.String(), + Tags: r.Tags, + CreatedBy: r.CreatedBy, + UpdatedBy: r.UpdatedBy, + CreatedAt: r.CreatedAt, + UpdatedAt: r.UpdatedAt, + }) + res.Attributes["input_channel"] = r.InputChannel + res.Attributes["input_topic"] = r.InputTopic + return res +} diff --git a/re/atom_test.go b/re/atom_test.go new file mode 100644 index 000000000..73961cfc3 --- /dev/null +++ b/re/atom_test.go @@ -0,0 +1,41 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package re + +import ( + "testing" + "time" + + "github.com/absmach/magistrala/pkg/schedule" +) + +func TestRuleProjectionOmitsMetadataAndSchedule(t *testing.T) { + got := ruleProjection(Rule{ + ID: "rule-1", + Name: "high-temp", + DomainID: "domain-1", + CreatedBy: "user-1", + Status: EnabledStatus, + Tags: []string{"smoke"}, + Metadata: Metadata{"flow": "encoded-flow", "other": "value"}, + InputChannel: "channel-1", + InputTopic: "messages", + Schedule: schedule.Schedule{ + Time: time.Date(2026, 6, 26, 17, 0, 0, 0, time.UTC), + }, + }) + + if _, ok := got.Attributes["metadata"]; ok { + t.Fatalf("rule metadata should not be projected to Atom attributes: %+v", got.Attributes) + } + if _, ok := got.Attributes["scheduled_at"]; ok { + t.Fatalf("rule schedule should not be projected to Atom attributes: %+v", got.Attributes) + } + if got.Attributes["input_channel"] != "channel-1" { + t.Fatalf("unexpected input_channel: %+v", got.Attributes) + } + if got.Attributes["input_topic"] != "messages" { + t.Fatalf("unexpected input_topic: %+v", got.Attributes) + } +} diff --git a/re/builtinroles.go b/re/builtinroles.go deleted file mode 100644 index 6e5f05051..000000000 --- a/re/builtinroles.go +++ /dev/null @@ -1,8 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package re - -import "github.com/absmach/magistrala/pkg/roles" - -const BuiltInRoleAdmin roles.BuiltInRoleName = "admin" diff --git a/re/events/events.go b/re/events/events.go index a704fde1f..7e807c286 100644 --- a/re/events/events.go +++ b/re/events/events.go @@ -8,7 +8,6 @@ import ( "github.com/absmach/magistrala/pkg/authn" "github.com/absmach/magistrala/pkg/events" - "github.com/absmach/magistrala/pkg/roles" "github.com/absmach/magistrala/re" ) @@ -60,8 +59,7 @@ func (bre baseRuleEvent) Encode() map[string]any { } type createRuleEvent struct { - rule re.Rule - rolesProvisioned []roles.RoleProvision + rule re.Rule baseRuleEvent } @@ -72,7 +70,6 @@ func (cre createRuleEvent) Encode() (map[string]any, error) { } maps.Copy(val, cre.baseRuleEvent.Encode()) val["operation"] = ruleCreate - val["roles_provisioned"] = cre.rolesProvisioned return val, nil } diff --git a/re/events/streams.go b/re/events/streams.go index c15033ae0..cff166702 100644 --- a/re/events/streams.go +++ b/re/events/streams.go @@ -10,8 +10,6 @@ import ( "github.com/absmach/magistrala/pkg/events" "github.com/absmach/magistrala/pkg/events/store" "github.com/absmach/magistrala/pkg/messaging" - "github.com/absmach/magistrala/pkg/roles" - rmEvents "github.com/absmach/magistrala/pkg/roles/rolemanager/events" "github.com/absmach/magistrala/re" "github.com/go-chi/chi/v5/middleware" ) @@ -34,7 +32,6 @@ var _ re.Service = (*eventStore)(nil) type eventStore struct { events.Publisher svc re.Service - rmEvents.RoleManagerEventStore } // NewEventStoreMiddleware returns wrapper around rules service that sends @@ -45,29 +42,25 @@ func NewEventStoreMiddleware(ctx context.Context, svc re.Service, url string) (r return nil, err } - res := rmEvents.NewRoleManagerEventStore("rules", rulePrefix, svc, publisher) - return &eventStore{ - svc: svc, - Publisher: publisher, - RoleManagerEventStore: res, + svc: svc, + Publisher: publisher, }, nil } -func (es *eventStore) AddRule(ctx context.Context, session authn.Session, r re.Rule) (re.Rule, []roles.RoleProvision, error) { - rule, rps, err := es.svc.AddRule(ctx, session, r) +func (es *eventStore) AddRule(ctx context.Context, session authn.Session, r re.Rule) (re.Rule, error) { + rule, err := es.svc.AddRule(ctx, session, r) if err != nil { - return rule, rps, err + return rule, err } event := createRuleEvent{ - rule: rule, - rolesProvisioned: rps, - baseRuleEvent: newBaseRuleEvent(session, middleware.GetReqID(ctx)), + rule: rule, + baseRuleEvent: newBaseRuleEvent(session, middleware.GetReqID(ctx)), } if err := es.Publish(ctx, CreateStream, event); err != nil { - return rule, rps, err + return rule, err } - return rule, rps, nil + return rule, nil } func (es *eventStore) ListRules(ctx context.Context, session authn.Session, pm re.PageMeta) (re.Page, error) { diff --git a/re/events/streams_test.go b/re/events/streams_test.go index b7f3954ad..543f2e5c1 100644 --- a/re/events/streams_test.go +++ b/re/events/streams_test.go @@ -15,7 +15,6 @@ import ( "github.com/absmach/magistrala/pkg/errors" svcerr "github.com/absmach/magistrala/pkg/errors/service" "github.com/absmach/magistrala/pkg/messaging" - "github.com/absmach/magistrala/pkg/roles" "github.com/absmach/magistrala/re" "github.com/absmach/magistrala/re/events" "github.com/absmach/magistrala/re/mocks" @@ -60,41 +59,38 @@ func TestAddRule(t *testing.T) { validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) cases := []struct { - desc string - session authn.Session - rule re.Rule - svcRes re.Rule - svcRoleRes []roles.RoleProvision - svcErr error - resp re.Rule - err error + desc string + session authn.Session + rule re.Rule + svcRes re.Rule + svcErr error + resp re.Rule + err error }{ { - desc: "publish successfully", - session: validSession, - rule: validRule, - svcRes: validRule, - svcRoleRes: []roles.RoleProvision{}, - svcErr: nil, - resp: validRule, - err: nil, + desc: "publish successfully", + session: validSession, + rule: validRule, + svcRes: validRule, + svcErr: nil, + resp: validRule, + err: nil, }, { - desc: "failed to publish with service error", - session: validSession, - rule: validRule, - svcRes: re.Rule{}, - svcRoleRes: []roles.RoleProvision{}, - svcErr: svcerr.ErrCreateEntity, - resp: re.Rule{}, - err: svcerr.ErrCreateEntity, + desc: "failed to publish with service error", + session: validSession, + rule: validRule, + svcRes: re.Rule{}, + svcErr: svcerr.ErrCreateEntity, + resp: re.Rule{}, + err: svcerr.ErrCreateEntity, }, } for _, tc := range cases { t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("AddRule", validCtx, tc.session, tc.rule).Return(tc.svcRes, tc.svcRoleRes, tc.svcErr) - resp, _, err := nsvc.AddRule(validCtx, tc.session, tc.rule) + svcCall := svc.On("AddRule", validCtx, tc.session, tc.rule).Return(tc.svcRes, tc.svcErr) + resp, err := nsvc.AddRule(validCtx, tc.session, tc.rule) assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) svcCall.Unset() diff --git a/re/middleware/authorization.go b/re/middleware/authorization.go index d02d38e00..f34c9c1dd 100644 --- a/re/middleware/authorization.go +++ b/re/middleware/authorization.go @@ -7,6 +7,7 @@ import ( "context" "github.com/absmach/magistrala/auth" + "github.com/absmach/magistrala/internal/atom" "github.com/absmach/magistrala/pkg/authn" smqauthz "github.com/absmach/magistrala/pkg/authz" "github.com/absmach/magistrala/pkg/errors" @@ -14,8 +15,6 @@ import ( "github.com/absmach/magistrala/pkg/messaging" "github.com/absmach/magistrala/pkg/permissions" "github.com/absmach/magistrala/pkg/policies" - "github.com/absmach/magistrala/pkg/roles" - rolemgr "github.com/absmach/magistrala/pkg/roles/rolemanager/middleware" "github.com/absmach/magistrala/re" "github.com/absmach/magistrala/re/operations" ) @@ -30,30 +29,36 @@ var ( type authorizationMiddleware struct { svc re.Service authz smqauthz.Authorization + atomAuthz atom.Authorizer entitiesOps permissions.EntitiesOperations[permissions.Operation] - rolemgr.RoleManagerAuthorizationMiddleware } // AuthorizationMiddleware adds authorization to the re service. -func AuthorizationMiddleware(svc re.Service, authz smqauthz.Authorization, entitiesOps permissions.EntitiesOperations[permissions.Operation], roleOps permissions.Operations[permissions.RoleOperation]) (re.Service, error) { +func AuthorizationMiddleware(svc re.Service, authz smqauthz.Authorization, entitiesOps permissions.EntitiesOperations[permissions.Operation]) (re.Service, error) { if err := entitiesOps.Validate(); err != nil { return nil, err } - ram, err := rolemgr.NewAuthorization(operations.EntityType, svc, authz, roleOps) - if err != nil { - return nil, err - } return &authorizationMiddleware{ - svc: svc, - authz: authz, - entitiesOps: entitiesOps, - RoleManagerAuthorizationMiddleware: ram, + svc: svc, + authz: authz, + entitiesOps: entitiesOps, }, nil } -func (am *authorizationMiddleware) AddRule(ctx context.Context, session authn.Session, r re.Rule) (re.Rule, []roles.RoleProvision, error) { +func AtomAuthorizationMiddleware(svc re.Service, authz atom.Authorizer, entitiesOps permissions.EntitiesOperations[permissions.Operation]) (re.Service, error) { + if err := entitiesOps.Validate(); err != nil { + return nil, err + } + return &authorizationMiddleware{ + svc: svc, + atomAuthz: authz, + entitiesOps: entitiesOps, + }, nil +} + +func (am *authorizationMiddleware) AddRule(ctx context.Context, session authn.Session, r re.Rule) (re.Rule, error) { if err := am.authorize(ctx, operations.OpAddRule, session, policies.DomainType, session.DomainID); err != nil { - return re.Rule{}, nil, errors.Wrap(errDomainCreateRules, err) + return re.Rule{}, errors.Wrap(errDomainCreateRules, err) } return am.svc.AddRule(ctx, session, r) @@ -96,6 +101,9 @@ func (am *authorizationMiddleware) ListRules(ctx context.Context, session authn. case err == nil: session.SuperAdmin = true case errors.Contains(err, svcerr.ErrSuperAdminAction): + if err := am.authorize(ctx, operations.OpListRules, session, operations.EntityType, auth.AnyIDs); err != nil { + return re.Page{}, errors.Wrap(errDomainViewRules, err) + } default: return re.Page{}, err } @@ -144,6 +152,9 @@ func (am *authorizationMiddleware) authorize(ctx context.Context, op permissions if err != nil { return err } + if am.atomAuthz != nil { + return atom.Authorize(ctx, am.atomAuthz, session, perm.String(), objType, obj, atom.KindRule) + } pr := smqauthz.PolicyReq{ Domain: session.DomainID, @@ -183,6 +194,9 @@ func (am *authorizationMiddleware) checkSuperAdmin(ctx context.Context, session if session.Role != authn.SuperAdminRole { return svcerr.ErrSuperAdminAction } + if am.atomAuthz != nil { + return atom.Authorize(ctx, am.atomAuthz, session, policies.AdminPermission, policies.PlatformType, policies.MagistralaObject, policies.PlatformType) + } if err := am.authz.Authorize(ctx, smqauthz.PolicyReq{ SubjectType: policies.UserType, Subject: session.UserID, diff --git a/re/middleware/authorization_test.go b/re/middleware/authorization_test.go new file mode 100644 index 000000000..155043e1f --- /dev/null +++ b/re/middleware/authorization_test.go @@ -0,0 +1,105 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package middleware + +import ( + "context" + "testing" + + "github.com/absmach/magistrala/auth" + "github.com/absmach/magistrala/internal/atom" + "github.com/absmach/magistrala/pkg/authn" + pkgerrors "github.com/absmach/magistrala/pkg/errors" + "github.com/absmach/magistrala/pkg/permissions" + "github.com/absmach/magistrala/re" + "github.com/absmach/magistrala/re/mocks" + "github.com/absmach/magistrala/re/operations" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/mock" + "github.com/stretchr/testify/require" +) + +type recordingAtomAuthorizer struct { + allowed bool + reqs []atom.AuthzRequest +} + +func (a *recordingAtomAuthorizer) CheckAuthz(_ context.Context, req atom.AuthzRequest) (atom.AuthzResponse, error) { + a.reqs = append(a.reqs, req) + return atom.AuthzResponse{Allowed: a.allowed}, nil +} + +func TestListRulesAuthorizesRegularUser(t *testing.T) { + svc := mocks.NewService(t) + pm := re.PageMeta{Limit: 10} + session := authn.Session{UserID: "user-1", DomainID: "domain-1", DomainUserID: "domain-1_user-1"} + authz := &recordingAtomAuthorizer{allowed: true} + wrapped, err := AtomAuthorizationMiddleware(svc, authz, testEntitiesOps(t)) + require.NoError(t, err) + + svc.On("ListRules", mock.Anything, session, pm).Return(re.Page{Limit: 10}, nil).Once() + page, err := wrapped.ListRules(context.Background(), session, pm) + + require.NoError(t, err) + assert.Equal(t, uint64(10), page.Limit) + require.Len(t, authz.reqs, 1) + assert.Equal(t, atom.AuthzRequest{ + SubjectID: "user-1", + Action: "list", + ResourceID: auth.AnyIDs, + ObjectKind: "resource", + ObjectID: auth.AnyIDs, + Context: map[string]any{ + "domain_id": "domain-1", + "legacy_object_type": operations.EntityType, + }, + }, authz.reqs[0]) +} + +func TestListRulesDeniedRegularUserDoesNotDelegate(t *testing.T) { + svc := mocks.NewService(t) + authz := &recordingAtomAuthorizer{allowed: false} + wrapped, err := AtomAuthorizationMiddleware(svc, authz, testEntitiesOps(t)) + require.NoError(t, err) + + _, err = wrapped.ListRules(context.Background(), authn.Session{UserID: "user-1", DomainID: "domain-1"}, re.PageMeta{}) + + assert.True(t, pkgerrors.Contains(err, pkgerrors.ErrAuthorization)) + require.Len(t, authz.reqs, 1) +} + +func TestListRulesSuperAdminSkipsListAuthorization(t *testing.T) { + svc := mocks.NewService(t) + pm := re.PageMeta{Limit: 10} + session := authn.Session{UserID: "admin-1", DomainID: "domain-1", Role: authn.SuperAdminRole} + authz := &recordingAtomAuthorizer{allowed: true} + wrapped, err := AtomAuthorizationMiddleware(svc, authz, testEntitiesOps(t)) + require.NoError(t, err) + + svc.On("ListRules", mock.Anything, mock.MatchedBy(func(s authn.Session) bool { + return s.SuperAdmin + }), pm).Return(re.Page{Limit: 10}, nil).Once() + _, err = wrapped.ListRules(context.Background(), session, pm) + + require.NoError(t, err) + require.Len(t, authz.reqs, 1) + assert.Equal(t, "manage", authz.reqs[0].Action) +} + +func testEntitiesOps(t *testing.T) permissions.EntitiesOperations[permissions.Operation] { + t.Helper() + details := operations.OperationDetails() + perms := make(map[string]permissions.Permission, len(details)) + for _, detail := range details { + if detail.PermissionRequired { + perms[detail.Name] = permissions.Permission(detail.Name) + } + } + entitiesOps, err := permissions.NewEntitiesOperations( + permissions.EntitiesPermission{operations.EntityType: perms}, + permissions.EntitiesOperationDetails[permissions.Operation]{operations.EntityType: details}, + ) + require.NoError(t, err) + return entitiesOps +} diff --git a/re/middleware/callout.go b/re/middleware/callout.go index b1cabf497..73c26f141 100644 --- a/re/middleware/callout.go +++ b/re/middleware/callout.go @@ -12,9 +12,6 @@ import ( "github.com/absmach/magistrala/pkg/messaging" "github.com/absmach/magistrala/pkg/permissions" "github.com/absmach/magistrala/pkg/policies" - mgPolicies "github.com/absmach/magistrala/pkg/policies" - "github.com/absmach/magistrala/pkg/roles" - rolemw "github.com/absmach/magistrala/pkg/roles/rolemanager/middleware" "github.com/absmach/magistrala/re" "github.com/absmach/magistrala/re/operations" ) @@ -25,37 +22,30 @@ type calloutMiddleware struct { svc re.Service callout callout.Callout entitiesOps permissions.EntitiesOperations[permissions.Operation] - rolemw.RoleManagerCalloutMiddleware } const entityType = "rule" -func NewCallout(svc re.Service, callout callout.Callout, entitiesOps permissions.EntitiesOperations[permissions.Operation], roleOps permissions.Operations[permissions.RoleOperation]) (re.Service, error) { - call, err := rolemw.NewCallout(mgPolicies.RulesType, svc, callout, roleOps) - if err != nil { - return nil, err - } - +func NewCallout(svc re.Service, callout callout.Callout, entitiesOps permissions.EntitiesOperations[permissions.Operation]) (re.Service, error) { if err := entitiesOps.Validate(); err != nil { return nil, err } return &calloutMiddleware{ - svc: svc, - callout: callout, - entitiesOps: entitiesOps, - RoleManagerCalloutMiddleware: call, + svc: svc, + callout: callout, + entitiesOps: entitiesOps, }, nil } -func (cm *calloutMiddleware) AddRule(ctx context.Context, session authn.Session, r re.Rule) (re.Rule, []roles.RoleProvision, error) { +func (cm *calloutMiddleware) AddRule(ctx context.Context, session authn.Session, r re.Rule) (re.Rule, error) { params := map[string]any{ "entities": r, "count": 1, } if err := cm.callOut(ctx, session, operations.OpAddRule, params); err != nil { - return re.Rule{}, nil, err + return re.Rule{}, err } return cm.svc.AddRule(ctx, session, r) diff --git a/re/middleware/logging.go b/re/middleware/logging.go index 1313afd06..41e924594 100644 --- a/re/middleware/logging.go +++ b/re/middleware/logging.go @@ -11,8 +11,6 @@ import ( "github.com/absmach/magistrala/pkg/authn" "github.com/absmach/magistrala/pkg/messaging" - "github.com/absmach/magistrala/pkg/roles" - rolemw "github.com/absmach/magistrala/pkg/roles/rolemanager/middleware" "github.com/absmach/magistrala/re" ) @@ -21,18 +19,16 @@ var _ re.Service = (*loggingMiddleware)(nil) type loggingMiddleware struct { logger *slog.Logger svc re.Service - rolemw.RoleManagerLoggingMiddleware } func LoggingMiddleware(svc re.Service, logger *slog.Logger) re.Service { return &loggingMiddleware{ - logger: logger, - svc: svc, - RoleManagerLoggingMiddleware: rolemw.NewLogging("re", svc, logger), + logger: logger, + svc: svc, } } -func (lm *loggingMiddleware) AddRule(ctx context.Context, session authn.Session, r re.Rule) (res re.Rule, rps []roles.RoleProvision, err error) { +func (lm *loggingMiddleware) AddRule(ctx context.Context, session authn.Session, r re.Rule) (res re.Rule, err error) { defer func(begin time.Time) { args := []any{ slog.String("duration", time.Since(begin).String()), @@ -46,7 +42,7 @@ func (lm *loggingMiddleware) AddRule(ctx context.Context, session authn.Session, } lm.logger.Info("Add rule completed successfully", args...) }(time.Now()) - res, rps, err = lm.svc.AddRule(ctx, session, r) + res, err = lm.svc.AddRule(ctx, session, r) return } diff --git a/re/middleware/metrics.go b/re/middleware/metrics.go index a93158bb8..e776292f3 100644 --- a/re/middleware/metrics.go +++ b/re/middleware/metrics.go @@ -9,8 +9,6 @@ import ( "github.com/absmach/magistrala/pkg/authn" "github.com/absmach/magistrala/pkg/messaging" - "github.com/absmach/magistrala/pkg/roles" - rolemw "github.com/absmach/magistrala/pkg/roles/rolemanager/middleware" "github.com/absmach/magistrala/re" "github.com/go-kit/kit/metrics" ) @@ -19,21 +17,19 @@ type metricsMiddleware struct { counter metrics.Counter latency metrics.Histogram service re.Service - rolemw.RoleManagerMetricsMiddleware } var _ re.Service = (*metricsMiddleware)(nil) func NewMetricsMiddleware(counter metrics.Counter, latency metrics.Histogram, service re.Service) re.Service { return &metricsMiddleware{ - counter: counter, - latency: latency, - service: service, - RoleManagerMetricsMiddleware: rolemw.NewMetrics("re", service, counter, latency), + counter: counter, + latency: latency, + service: service, } } -func (mm *metricsMiddleware) AddRule(ctx context.Context, session authn.Session, r re.Rule) (re.Rule, []roles.RoleProvision, error) { +func (mm *metricsMiddleware) AddRule(ctx context.Context, session authn.Session, r re.Rule) (re.Rule, error) { defer func(begin time.Time) { mm.counter.With("method", "add_rule").Add(1) mm.latency.With("method", "add_rule").Observe(time.Since(begin).Seconds()) diff --git a/re/middleware/tracing.go b/re/middleware/tracing.go index 146f2977f..4ff6d1881 100644 --- a/re/middleware/tracing.go +++ b/re/middleware/tracing.go @@ -8,8 +8,6 @@ import ( "github.com/absmach/magistrala/pkg/authn" "github.com/absmach/magistrala/pkg/messaging" - "github.com/absmach/magistrala/pkg/roles" - rolemw "github.com/absmach/magistrala/pkg/roles/rolemanager/middleware" smqTracing "github.com/absmach/magistrala/pkg/tracing" "github.com/absmach/magistrala/re" "go.opentelemetry.io/otel/attribute" @@ -19,20 +17,18 @@ import ( type tracingMiddleware struct { tracer trace.Tracer svc re.Service - rolemw.RoleManagerTracing } var _ re.Service = (*tracingMiddleware)(nil) func NewTracingMiddleware(tracer trace.Tracer, svc re.Service) re.Service { return &tracingMiddleware{ - tracer: tracer, - svc: svc, - RoleManagerTracing: rolemw.NewTracing("re", svc, tracer), + tracer: tracer, + svc: svc, } } -func (tm *tracingMiddleware) AddRule(ctx context.Context, session authn.Session, r re.Rule) (re.Rule, []roles.RoleProvision, error) { +func (tm *tracingMiddleware) AddRule(ctx context.Context, session authn.Session, r re.Rule) (re.Rule, error) { ctx, span := smqTracing.StartSpan(ctx, tm.tracer, "add_rule", trace.WithAttributes( attribute.String("name", r.Name), attribute.String("domain_id", r.DomainID), diff --git a/re/mocks/repository.go b/re/mocks/repository.go index 4e5c2b1d6..4bfe2844a 100644 --- a/re/mocks/repository.go +++ b/re/mocks/repository.go @@ -12,7 +12,6 @@ import ( "context" "time" - "github.com/absmach/magistrala/pkg/roles" "github.com/absmach/magistrala/re" mock "github.com/stretchr/testify/mock" ) @@ -44,74 +43,6 @@ func (_m *Repository) EXPECT() *Repository_Expecter { return &Repository_Expecter{mock: &_m.Mock} } -// AddRoles provides a mock function for the type Repository -func (_mock *Repository) AddRoles(ctx context.Context, rps []roles.RoleProvision) ([]roles.RoleProvision, error) { - ret := _mock.Called(ctx, rps) - - if len(ret) == 0 { - panic("no return value specified for AddRoles") - } - - var r0 []roles.RoleProvision - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, []roles.RoleProvision) ([]roles.RoleProvision, error)); ok { - return returnFunc(ctx, rps) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, []roles.RoleProvision) []roles.RoleProvision); ok { - r0 = returnFunc(ctx, rps) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]roles.RoleProvision) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, []roles.RoleProvision) error); ok { - r1 = returnFunc(ctx, rps) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_AddRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'AddRoles' -type Repository_AddRoles_Call struct { - *mock.Call -} - -// AddRoles is a helper method to define mock.On call -// - ctx context.Context -// - rps []roles.RoleProvision -func (_e *Repository_Expecter) AddRoles(ctx interface{}, rps interface{}) *Repository_AddRoles_Call { - return &Repository_AddRoles_Call{Call: _e.mock.On("AddRoles", ctx, rps)} -} - -func (_c *Repository_AddRoles_Call) Run(run func(ctx context.Context, rps []roles.RoleProvision)) *Repository_AddRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 []roles.RoleProvision - if args[1] != nil { - arg1 = args[1].([]roles.RoleProvision) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_AddRoles_Call) Return(roleProvisions []roles.RoleProvision, err error) *Repository_AddRoles_Call { - _c.Call.Return(roleProvisions, err) - return _c -} - -func (_c *Repository_AddRoles_Call) RunAndReturn(run func(ctx context.Context, rps []roles.RoleProvision) ([]roles.RoleProvision, error)) *Repository_AddRoles_Call { - _c.Call.Return(run) - return _c -} - // AddRule provides a mock function for the type Repository func (_mock *Repository) AddRule(ctx context.Context, r re.Rule) (re.Rule, error) { ret := _mock.Called(ctx, r) @@ -244,327 +175,6 @@ func (_c *Repository_ListAllRules_Call) RunAndReturn(run func(ctx context.Contex return _c } -// ListEntityMembers provides a mock function for the type Repository -func (_mock *Repository) ListEntityMembers(ctx context.Context, entityID string, pageQuery roles.MembersRolePageQuery) (roles.MembersRolePage, error) { - ret := _mock.Called(ctx, entityID, pageQuery) - - if len(ret) == 0 { - panic("no return value specified for ListEntityMembers") - } - - var r0 roles.MembersRolePage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, roles.MembersRolePageQuery) (roles.MembersRolePage, error)); ok { - return returnFunc(ctx, entityID, pageQuery) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, roles.MembersRolePageQuery) roles.MembersRolePage); ok { - r0 = returnFunc(ctx, entityID, pageQuery) - } else { - r0 = ret.Get(0).(roles.MembersRolePage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, roles.MembersRolePageQuery) error); ok { - r1 = returnFunc(ctx, entityID, pageQuery) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_ListEntityMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ListEntityMembers' -type Repository_ListEntityMembers_Call struct { - *mock.Call -} - -// ListEntityMembers is a helper method to define mock.On call -// - ctx context.Context -// - entityID string -// - pageQuery roles.MembersRolePageQuery -func (_e *Repository_Expecter) ListEntityMembers(ctx interface{}, entityID interface{}, pageQuery interface{}) *Repository_ListEntityMembers_Call { - return &Repository_ListEntityMembers_Call{Call: _e.mock.On("ListEntityMembers", ctx, entityID, pageQuery)} -} - -func (_c *Repository_ListEntityMembers_Call) Run(run func(ctx context.Context, entityID string, pageQuery roles.MembersRolePageQuery)) *Repository_ListEntityMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 roles.MembersRolePageQuery - if args[2] != nil { - arg2 = args[2].(roles.MembersRolePageQuery) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_ListEntityMembers_Call) Return(membersRolePage roles.MembersRolePage, err error) *Repository_ListEntityMembers_Call { - _c.Call.Return(membersRolePage, err) - return _c -} - -func (_c *Repository_ListEntityMembers_Call) RunAndReturn(run func(ctx context.Context, entityID string, pageQuery roles.MembersRolePageQuery) (roles.MembersRolePage, error)) *Repository_ListEntityMembers_Call { - _c.Call.Return(run) - return _c -} - -// ListUserRules provides a mock function for the type Repository -func (_mock *Repository) ListUserRules(ctx context.Context, userID string, pm re.PageMeta) (re.Page, error) { - ret := _mock.Called(ctx, userID, pm) - - if len(ret) == 0 { - panic("no return value specified for ListUserRules") - } - - var r0 re.Page - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, re.PageMeta) (re.Page, error)); ok { - return returnFunc(ctx, userID, pm) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, re.PageMeta) re.Page); ok { - r0 = returnFunc(ctx, userID, pm) - } else { - r0 = ret.Get(0).(re.Page) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, re.PageMeta) error); ok { - r1 = returnFunc(ctx, userID, pm) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_ListUserRules_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ListUserRules' -type Repository_ListUserRules_Call struct { - *mock.Call -} - -// ListUserRules is a helper method to define mock.On call -// - ctx context.Context -// - userID string -// - pm re.PageMeta -func (_e *Repository_Expecter) ListUserRules(ctx interface{}, userID interface{}, pm interface{}) *Repository_ListUserRules_Call { - return &Repository_ListUserRules_Call{Call: _e.mock.On("ListUserRules", ctx, userID, pm)} -} - -func (_c *Repository_ListUserRules_Call) Run(run func(ctx context.Context, userID string, pm re.PageMeta)) *Repository_ListUserRules_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 re.PageMeta - if args[2] != nil { - arg2 = args[2].(re.PageMeta) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_ListUserRules_Call) Return(page re.Page, err error) *Repository_ListUserRules_Call { - _c.Call.Return(page, err) - return _c -} - -func (_c *Repository_ListUserRules_Call) RunAndReturn(run func(ctx context.Context, userID string, pm re.PageMeta) (re.Page, error)) *Repository_ListUserRules_Call { - _c.Call.Return(run) - return _c -} - -// RemoveEntityMembers provides a mock function for the type Repository -func (_mock *Repository) RemoveEntityMembers(ctx context.Context, entityID string, members []string) error { - ret := _mock.Called(ctx, entityID, members) - - if len(ret) == 0 { - panic("no return value specified for RemoveEntityMembers") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, []string) error); ok { - r0 = returnFunc(ctx, entityID, members) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RemoveEntityMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveEntityMembers' -type Repository_RemoveEntityMembers_Call struct { - *mock.Call -} - -// RemoveEntityMembers is a helper method to define mock.On call -// - ctx context.Context -// - entityID string -// - members []string -func (_e *Repository_Expecter) RemoveEntityMembers(ctx interface{}, entityID interface{}, members interface{}) *Repository_RemoveEntityMembers_Call { - return &Repository_RemoveEntityMembers_Call{Call: _e.mock.On("RemoveEntityMembers", ctx, entityID, members)} -} - -func (_c *Repository_RemoveEntityMembers_Call) Run(run func(ctx context.Context, entityID string, members []string)) *Repository_RemoveEntityMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RemoveEntityMembers_Call) Return(err error) *Repository_RemoveEntityMembers_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RemoveEntityMembers_Call) RunAndReturn(run func(ctx context.Context, entityID string, members []string) error) *Repository_RemoveEntityMembers_Call { - _c.Call.Return(run) - return _c -} - -// RemoveMemberFromAllRoles provides a mock function for the type Repository -func (_mock *Repository) RemoveMemberFromAllRoles(ctx context.Context, memberID string) error { - ret := _mock.Called(ctx, memberID) - - if len(ret) == 0 { - panic("no return value specified for RemoveMemberFromAllRoles") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) error); ok { - r0 = returnFunc(ctx, memberID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RemoveMemberFromAllRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveMemberFromAllRoles' -type Repository_RemoveMemberFromAllRoles_Call struct { - *mock.Call -} - -// RemoveMemberFromAllRoles is a helper method to define mock.On call -// - ctx context.Context -// - memberID string -func (_e *Repository_Expecter) RemoveMemberFromAllRoles(ctx interface{}, memberID interface{}) *Repository_RemoveMemberFromAllRoles_Call { - return &Repository_RemoveMemberFromAllRoles_Call{Call: _e.mock.On("RemoveMemberFromAllRoles", ctx, memberID)} -} - -func (_c *Repository_RemoveMemberFromAllRoles_Call) Run(run func(ctx context.Context, memberID string)) *Repository_RemoveMemberFromAllRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RemoveMemberFromAllRoles_Call) Return(err error) *Repository_RemoveMemberFromAllRoles_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RemoveMemberFromAllRoles_Call) RunAndReturn(run func(ctx context.Context, memberID string) error) *Repository_RemoveMemberFromAllRoles_Call { - _c.Call.Return(run) - return _c -} - -// RemoveRoles provides a mock function for the type Repository -func (_mock *Repository) RemoveRoles(ctx context.Context, roleIDs []string) error { - ret := _mock.Called(ctx, roleIDs) - - if len(ret) == 0 { - panic("no return value specified for RemoveRoles") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, []string) error); ok { - r0 = returnFunc(ctx, roleIDs) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RemoveRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveRoles' -type Repository_RemoveRoles_Call struct { - *mock.Call -} - -// RemoveRoles is a helper method to define mock.On call -// - ctx context.Context -// - roleIDs []string -func (_e *Repository_Expecter) RemoveRoles(ctx interface{}, roleIDs interface{}) *Repository_RemoveRoles_Call { - return &Repository_RemoveRoles_Call{Call: _e.mock.On("RemoveRoles", ctx, roleIDs)} -} - -func (_c *Repository_RemoveRoles_Call) Run(run func(ctx context.Context, roleIDs []string)) *Repository_RemoveRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 []string - if args[1] != nil { - arg1 = args[1].([]string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RemoveRoles_Call) Return(err error) *Repository_RemoveRoles_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RemoveRoles_Call) RunAndReturn(run func(ctx context.Context, roleIDs []string) error) *Repository_RemoveRoles_Call { - _c.Call.Return(run) - return _c -} - // RemoveRule provides a mock function for the type Repository func (_mock *Repository) RemoveRule(ctx context.Context, id string) error { ret := _mock.Called(ctx, id) @@ -622,1114 +232,6 @@ func (_c *Repository_RemoveRule_Call) RunAndReturn(run func(ctx context.Context, return _c } -// RetrieveAllRoles provides a mock function for the type Repository -func (_mock *Repository) RetrieveAllRoles(ctx context.Context, entityID string, limit uint64, offset uint64) (roles.RolePage, error) { - ret := _mock.Called(ctx, entityID, limit, offset) - - if len(ret) == 0 { - panic("no return value specified for RetrieveAllRoles") - } - - var r0 roles.RolePage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, uint64, uint64) (roles.RolePage, error)); ok { - return returnFunc(ctx, entityID, limit, offset) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, uint64, uint64) roles.RolePage); ok { - r0 = returnFunc(ctx, entityID, limit, offset) - } else { - r0 = ret.Get(0).(roles.RolePage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, uint64, uint64) error); ok { - r1 = returnFunc(ctx, entityID, limit, offset) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveAllRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveAllRoles' -type Repository_RetrieveAllRoles_Call struct { - *mock.Call -} - -// RetrieveAllRoles is a helper method to define mock.On call -// - ctx context.Context -// - entityID string -// - limit uint64 -// - offset uint64 -func (_e *Repository_Expecter) RetrieveAllRoles(ctx interface{}, entityID interface{}, limit interface{}, offset interface{}) *Repository_RetrieveAllRoles_Call { - return &Repository_RetrieveAllRoles_Call{Call: _e.mock.On("RetrieveAllRoles", ctx, entityID, limit, offset)} -} - -func (_c *Repository_RetrieveAllRoles_Call) Run(run func(ctx context.Context, entityID string, limit uint64, offset uint64)) *Repository_RetrieveAllRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 uint64 - if args[2] != nil { - arg2 = args[2].(uint64) - } - var arg3 uint64 - if args[3] != nil { - arg3 = args[3].(uint64) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Repository_RetrieveAllRoles_Call) Return(rolePage roles.RolePage, err error) *Repository_RetrieveAllRoles_Call { - _c.Call.Return(rolePage, err) - return _c -} - -func (_c *Repository_RetrieveAllRoles_Call) RunAndReturn(run func(ctx context.Context, entityID string, limit uint64, offset uint64) (roles.RolePage, error)) *Repository_RetrieveAllRoles_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveByIDWithRoles provides a mock function for the type Repository -func (_mock *Repository) RetrieveByIDWithRoles(ctx context.Context, id string, memberID string) (re.Rule, error) { - ret := _mock.Called(ctx, id, memberID) - - if len(ret) == 0 { - panic("no return value specified for RetrieveByIDWithRoles") - } - - var r0 re.Rule - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) (re.Rule, error)); ok { - return returnFunc(ctx, id, memberID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) re.Rule); ok { - r0 = returnFunc(ctx, id, memberID) - } else { - r0 = ret.Get(0).(re.Rule) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, string) error); ok { - r1 = returnFunc(ctx, id, memberID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveByIDWithRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveByIDWithRoles' -type Repository_RetrieveByIDWithRoles_Call struct { - *mock.Call -} - -// RetrieveByIDWithRoles is a helper method to define mock.On call -// - ctx context.Context -// - id string -// - memberID string -func (_e *Repository_Expecter) RetrieveByIDWithRoles(ctx interface{}, id interface{}, memberID interface{}) *Repository_RetrieveByIDWithRoles_Call { - return &Repository_RetrieveByIDWithRoles_Call{Call: _e.mock.On("RetrieveByIDWithRoles", ctx, id, memberID)} -} - -func (_c *Repository_RetrieveByIDWithRoles_Call) Run(run func(ctx context.Context, id string, memberID string)) *Repository_RetrieveByIDWithRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RetrieveByIDWithRoles_Call) Return(rule re.Rule, err error) *Repository_RetrieveByIDWithRoles_Call { - _c.Call.Return(rule, err) - return _c -} - -func (_c *Repository_RetrieveByIDWithRoles_Call) RunAndReturn(run func(ctx context.Context, id string, memberID string) (re.Rule, error)) *Repository_RetrieveByIDWithRoles_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveEntitiesRolesActionsMembers provides a mock function for the type Repository -func (_mock *Repository) RetrieveEntitiesRolesActionsMembers(ctx context.Context, entityIDs []string) ([]roles.EntityActionRole, []roles.EntityMemberRole, error) { - ret := _mock.Called(ctx, entityIDs) - - if len(ret) == 0 { - panic("no return value specified for RetrieveEntitiesRolesActionsMembers") - } - - var r0 []roles.EntityActionRole - var r1 []roles.EntityMemberRole - var r2 error - if returnFunc, ok := ret.Get(0).(func(context.Context, []string) ([]roles.EntityActionRole, []roles.EntityMemberRole, error)); ok { - return returnFunc(ctx, entityIDs) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, []string) []roles.EntityActionRole); ok { - r0 = returnFunc(ctx, entityIDs) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]roles.EntityActionRole) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, []string) []roles.EntityMemberRole); ok { - r1 = returnFunc(ctx, entityIDs) - } else { - if ret.Get(1) != nil { - r1 = ret.Get(1).([]roles.EntityMemberRole) - } - } - if returnFunc, ok := ret.Get(2).(func(context.Context, []string) error); ok { - r2 = returnFunc(ctx, entityIDs) - } else { - r2 = ret.Error(2) - } - return r0, r1, r2 -} - -// Repository_RetrieveEntitiesRolesActionsMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveEntitiesRolesActionsMembers' -type Repository_RetrieveEntitiesRolesActionsMembers_Call struct { - *mock.Call -} - -// RetrieveEntitiesRolesActionsMembers is a helper method to define mock.On call -// - ctx context.Context -// - entityIDs []string -func (_e *Repository_Expecter) RetrieveEntitiesRolesActionsMembers(ctx interface{}, entityIDs interface{}) *Repository_RetrieveEntitiesRolesActionsMembers_Call { - return &Repository_RetrieveEntitiesRolesActionsMembers_Call{Call: _e.mock.On("RetrieveEntitiesRolesActionsMembers", ctx, entityIDs)} -} - -func (_c *Repository_RetrieveEntitiesRolesActionsMembers_Call) Run(run func(ctx context.Context, entityIDs []string)) *Repository_RetrieveEntitiesRolesActionsMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 []string - if args[1] != nil { - arg1 = args[1].([]string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RetrieveEntitiesRolesActionsMembers_Call) Return(entityActionRoles []roles.EntityActionRole, entityMemberRoles []roles.EntityMemberRole, err error) *Repository_RetrieveEntitiesRolesActionsMembers_Call { - _c.Call.Return(entityActionRoles, entityMemberRoles, err) - return _c -} - -func (_c *Repository_RetrieveEntitiesRolesActionsMembers_Call) RunAndReturn(run func(ctx context.Context, entityIDs []string) ([]roles.EntityActionRole, []roles.EntityMemberRole, error)) *Repository_RetrieveEntitiesRolesActionsMembers_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveEntityRole provides a mock function for the type Repository -func (_mock *Repository) RetrieveEntityRole(ctx context.Context, entityID string, roleID string) (roles.Role, error) { - ret := _mock.Called(ctx, entityID, roleID) - - if len(ret) == 0 { - panic("no return value specified for RetrieveEntityRole") - } - - var r0 roles.Role - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) (roles.Role, error)); ok { - return returnFunc(ctx, entityID, roleID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) roles.Role); ok { - r0 = returnFunc(ctx, entityID, roleID) - } else { - r0 = ret.Get(0).(roles.Role) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, string) error); ok { - r1 = returnFunc(ctx, entityID, roleID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveEntityRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveEntityRole' -type Repository_RetrieveEntityRole_Call struct { - *mock.Call -} - -// RetrieveEntityRole is a helper method to define mock.On call -// - ctx context.Context -// - entityID string -// - roleID string -func (_e *Repository_Expecter) RetrieveEntityRole(ctx interface{}, entityID interface{}, roleID interface{}) *Repository_RetrieveEntityRole_Call { - return &Repository_RetrieveEntityRole_Call{Call: _e.mock.On("RetrieveEntityRole", ctx, entityID, roleID)} -} - -func (_c *Repository_RetrieveEntityRole_Call) Run(run func(ctx context.Context, entityID string, roleID string)) *Repository_RetrieveEntityRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RetrieveEntityRole_Call) Return(role roles.Role, err error) *Repository_RetrieveEntityRole_Call { - _c.Call.Return(role, err) - return _c -} - -func (_c *Repository_RetrieveEntityRole_Call) RunAndReturn(run func(ctx context.Context, entityID string, roleID string) (roles.Role, error)) *Repository_RetrieveEntityRole_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveRole provides a mock function for the type Repository -func (_mock *Repository) RetrieveRole(ctx context.Context, roleID string) (roles.Role, error) { - ret := _mock.Called(ctx, roleID) - - if len(ret) == 0 { - panic("no return value specified for RetrieveRole") - } - - var r0 roles.Role - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) (roles.Role, error)); ok { - return returnFunc(ctx, roleID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string) roles.Role); ok { - r0 = returnFunc(ctx, roleID) - } else { - r0 = ret.Get(0).(roles.Role) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = returnFunc(ctx, roleID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveRole' -type Repository_RetrieveRole_Call struct { - *mock.Call -} - -// RetrieveRole is a helper method to define mock.On call -// - ctx context.Context -// - roleID string -func (_e *Repository_Expecter) RetrieveRole(ctx interface{}, roleID interface{}) *Repository_RetrieveRole_Call { - return &Repository_RetrieveRole_Call{Call: _e.mock.On("RetrieveRole", ctx, roleID)} -} - -func (_c *Repository_RetrieveRole_Call) Run(run func(ctx context.Context, roleID string)) *Repository_RetrieveRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RetrieveRole_Call) Return(role roles.Role, err error) *Repository_RetrieveRole_Call { - _c.Call.Return(role, err) - return _c -} - -func (_c *Repository_RetrieveRole_Call) RunAndReturn(run func(ctx context.Context, roleID string) (roles.Role, error)) *Repository_RetrieveRole_Call { - _c.Call.Return(run) - return _c -} - -// RoleAddActions provides a mock function for the type Repository -func (_mock *Repository) RoleAddActions(ctx context.Context, role roles.Role, actions []string) ([]string, error) { - ret := _mock.Called(ctx, role, actions) - - if len(ret) == 0 { - panic("no return value specified for RoleAddActions") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role, []string) ([]string, error)); ok { - return returnFunc(ctx, role, actions) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role, []string) []string); ok { - r0 = returnFunc(ctx, role, actions) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, roles.Role, []string) error); ok { - r1 = returnFunc(ctx, role, actions) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RoleAddActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleAddActions' -type Repository_RoleAddActions_Call struct { - *mock.Call -} - -// RoleAddActions is a helper method to define mock.On call -// - ctx context.Context -// - role roles.Role -// - actions []string -func (_e *Repository_Expecter) RoleAddActions(ctx interface{}, role interface{}, actions interface{}) *Repository_RoleAddActions_Call { - return &Repository_RoleAddActions_Call{Call: _e.mock.On("RoleAddActions", ctx, role, actions)} -} - -func (_c *Repository_RoleAddActions_Call) Run(run func(ctx context.Context, role roles.Role, actions []string)) *Repository_RoleAddActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RoleAddActions_Call) Return(ops []string, err error) *Repository_RoleAddActions_Call { - _c.Call.Return(ops, err) - return _c -} - -func (_c *Repository_RoleAddActions_Call) RunAndReturn(run func(ctx context.Context, role roles.Role, actions []string) ([]string, error)) *Repository_RoleAddActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleAddMembers provides a mock function for the type Repository -func (_mock *Repository) RoleAddMembers(ctx context.Context, role roles.Role, members []string) ([]string, error) { - ret := _mock.Called(ctx, role, members) - - if len(ret) == 0 { - panic("no return value specified for RoleAddMembers") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role, []string) ([]string, error)); ok { - return returnFunc(ctx, role, members) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role, []string) []string); ok { - r0 = returnFunc(ctx, role, members) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, roles.Role, []string) error); ok { - r1 = returnFunc(ctx, role, members) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RoleAddMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleAddMembers' -type Repository_RoleAddMembers_Call struct { - *mock.Call -} - -// RoleAddMembers is a helper method to define mock.On call -// - ctx context.Context -// - role roles.Role -// - members []string -func (_e *Repository_Expecter) RoleAddMembers(ctx interface{}, role interface{}, members interface{}) *Repository_RoleAddMembers_Call { - return &Repository_RoleAddMembers_Call{Call: _e.mock.On("RoleAddMembers", ctx, role, members)} -} - -func (_c *Repository_RoleAddMembers_Call) Run(run func(ctx context.Context, role roles.Role, members []string)) *Repository_RoleAddMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RoleAddMembers_Call) Return(strings []string, err error) *Repository_RoleAddMembers_Call { - _c.Call.Return(strings, err) - return _c -} - -func (_c *Repository_RoleAddMembers_Call) RunAndReturn(run func(ctx context.Context, role roles.Role, members []string) ([]string, error)) *Repository_RoleAddMembers_Call { - _c.Call.Return(run) - return _c -} - -// RoleCheckActionsExists provides a mock function for the type Repository -func (_mock *Repository) RoleCheckActionsExists(ctx context.Context, roleID string, actions []string) (bool, error) { - ret := _mock.Called(ctx, roleID, actions) - - if len(ret) == 0 { - panic("no return value specified for RoleCheckActionsExists") - } - - var r0 bool - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, []string) (bool, error)); ok { - return returnFunc(ctx, roleID, actions) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, []string) bool); ok { - r0 = returnFunc(ctx, roleID, actions) - } else { - r0 = ret.Get(0).(bool) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, []string) error); ok { - r1 = returnFunc(ctx, roleID, actions) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RoleCheckActionsExists_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleCheckActionsExists' -type Repository_RoleCheckActionsExists_Call struct { - *mock.Call -} - -// RoleCheckActionsExists is a helper method to define mock.On call -// - ctx context.Context -// - roleID string -// - actions []string -func (_e *Repository_Expecter) RoleCheckActionsExists(ctx interface{}, roleID interface{}, actions interface{}) *Repository_RoleCheckActionsExists_Call { - return &Repository_RoleCheckActionsExists_Call{Call: _e.mock.On("RoleCheckActionsExists", ctx, roleID, actions)} -} - -func (_c *Repository_RoleCheckActionsExists_Call) Run(run func(ctx context.Context, roleID string, actions []string)) *Repository_RoleCheckActionsExists_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RoleCheckActionsExists_Call) Return(b bool, err error) *Repository_RoleCheckActionsExists_Call { - _c.Call.Return(b, err) - return _c -} - -func (_c *Repository_RoleCheckActionsExists_Call) RunAndReturn(run func(ctx context.Context, roleID string, actions []string) (bool, error)) *Repository_RoleCheckActionsExists_Call { - _c.Call.Return(run) - return _c -} - -// RoleCheckMembersExists provides a mock function for the type Repository -func (_mock *Repository) RoleCheckMembersExists(ctx context.Context, roleID string, members []string) (bool, error) { - ret := _mock.Called(ctx, roleID, members) - - if len(ret) == 0 { - panic("no return value specified for RoleCheckMembersExists") - } - - var r0 bool - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, []string) (bool, error)); ok { - return returnFunc(ctx, roleID, members) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, []string) bool); ok { - r0 = returnFunc(ctx, roleID, members) - } else { - r0 = ret.Get(0).(bool) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, []string) error); ok { - r1 = returnFunc(ctx, roleID, members) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RoleCheckMembersExists_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleCheckMembersExists' -type Repository_RoleCheckMembersExists_Call struct { - *mock.Call -} - -// RoleCheckMembersExists is a helper method to define mock.On call -// - ctx context.Context -// - roleID string -// - members []string -func (_e *Repository_Expecter) RoleCheckMembersExists(ctx interface{}, roleID interface{}, members interface{}) *Repository_RoleCheckMembersExists_Call { - return &Repository_RoleCheckMembersExists_Call{Call: _e.mock.On("RoleCheckMembersExists", ctx, roleID, members)} -} - -func (_c *Repository_RoleCheckMembersExists_Call) Run(run func(ctx context.Context, roleID string, members []string)) *Repository_RoleCheckMembersExists_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RoleCheckMembersExists_Call) Return(b bool, err error) *Repository_RoleCheckMembersExists_Call { - _c.Call.Return(b, err) - return _c -} - -func (_c *Repository_RoleCheckMembersExists_Call) RunAndReturn(run func(ctx context.Context, roleID string, members []string) (bool, error)) *Repository_RoleCheckMembersExists_Call { - _c.Call.Return(run) - return _c -} - -// RoleListActions provides a mock function for the type Repository -func (_mock *Repository) RoleListActions(ctx context.Context, roleID string) ([]string, error) { - ret := _mock.Called(ctx, roleID) - - if len(ret) == 0 { - panic("no return value specified for RoleListActions") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) ([]string, error)); ok { - return returnFunc(ctx, roleID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string) []string); ok { - r0 = returnFunc(ctx, roleID) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = returnFunc(ctx, roleID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RoleListActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleListActions' -type Repository_RoleListActions_Call struct { - *mock.Call -} - -// RoleListActions is a helper method to define mock.On call -// - ctx context.Context -// - roleID string -func (_e *Repository_Expecter) RoleListActions(ctx interface{}, roleID interface{}) *Repository_RoleListActions_Call { - return &Repository_RoleListActions_Call{Call: _e.mock.On("RoleListActions", ctx, roleID)} -} - -func (_c *Repository_RoleListActions_Call) Run(run func(ctx context.Context, roleID string)) *Repository_RoleListActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RoleListActions_Call) Return(strings []string, err error) *Repository_RoleListActions_Call { - _c.Call.Return(strings, err) - return _c -} - -func (_c *Repository_RoleListActions_Call) RunAndReturn(run func(ctx context.Context, roleID string) ([]string, error)) *Repository_RoleListActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleListMembers provides a mock function for the type Repository -func (_mock *Repository) RoleListMembers(ctx context.Context, roleID string, limit uint64, offset uint64) (roles.MembersPage, error) { - ret := _mock.Called(ctx, roleID, limit, offset) - - if len(ret) == 0 { - panic("no return value specified for RoleListMembers") - } - - var r0 roles.MembersPage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, uint64, uint64) (roles.MembersPage, error)); ok { - return returnFunc(ctx, roleID, limit, offset) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, uint64, uint64) roles.MembersPage); ok { - r0 = returnFunc(ctx, roleID, limit, offset) - } else { - r0 = ret.Get(0).(roles.MembersPage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, uint64, uint64) error); ok { - r1 = returnFunc(ctx, roleID, limit, offset) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RoleListMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleListMembers' -type Repository_RoleListMembers_Call struct { - *mock.Call -} - -// RoleListMembers is a helper method to define mock.On call -// - ctx context.Context -// - roleID string -// - limit uint64 -// - offset uint64 -func (_e *Repository_Expecter) RoleListMembers(ctx interface{}, roleID interface{}, limit interface{}, offset interface{}) *Repository_RoleListMembers_Call { - return &Repository_RoleListMembers_Call{Call: _e.mock.On("RoleListMembers", ctx, roleID, limit, offset)} -} - -func (_c *Repository_RoleListMembers_Call) Run(run func(ctx context.Context, roleID string, limit uint64, offset uint64)) *Repository_RoleListMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 uint64 - if args[2] != nil { - arg2 = args[2].(uint64) - } - var arg3 uint64 - if args[3] != nil { - arg3 = args[3].(uint64) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Repository_RoleListMembers_Call) Return(membersPage roles.MembersPage, err error) *Repository_RoleListMembers_Call { - _c.Call.Return(membersPage, err) - return _c -} - -func (_c *Repository_RoleListMembers_Call) RunAndReturn(run func(ctx context.Context, roleID string, limit uint64, offset uint64) (roles.MembersPage, error)) *Repository_RoleListMembers_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveActions provides a mock function for the type Repository -func (_mock *Repository) RoleRemoveActions(ctx context.Context, role roles.Role, actions []string) error { - ret := _mock.Called(ctx, role, actions) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveActions") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role, []string) error); ok { - r0 = returnFunc(ctx, role, actions) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RoleRemoveActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveActions' -type Repository_RoleRemoveActions_Call struct { - *mock.Call -} - -// RoleRemoveActions is a helper method to define mock.On call -// - ctx context.Context -// - role roles.Role -// - actions []string -func (_e *Repository_Expecter) RoleRemoveActions(ctx interface{}, role interface{}, actions interface{}) *Repository_RoleRemoveActions_Call { - return &Repository_RoleRemoveActions_Call{Call: _e.mock.On("RoleRemoveActions", ctx, role, actions)} -} - -func (_c *Repository_RoleRemoveActions_Call) Run(run func(ctx context.Context, role roles.Role, actions []string)) *Repository_RoleRemoveActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RoleRemoveActions_Call) Return(err error) *Repository_RoleRemoveActions_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RoleRemoveActions_Call) RunAndReturn(run func(ctx context.Context, role roles.Role, actions []string) error) *Repository_RoleRemoveActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveAllActions provides a mock function for the type Repository -func (_mock *Repository) RoleRemoveAllActions(ctx context.Context, role roles.Role) error { - ret := _mock.Called(ctx, role) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveAllActions") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role) error); ok { - r0 = returnFunc(ctx, role) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RoleRemoveAllActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveAllActions' -type Repository_RoleRemoveAllActions_Call struct { - *mock.Call -} - -// RoleRemoveAllActions is a helper method to define mock.On call -// - ctx context.Context -// - role roles.Role -func (_e *Repository_Expecter) RoleRemoveAllActions(ctx interface{}, role interface{}) *Repository_RoleRemoveAllActions_Call { - return &Repository_RoleRemoveAllActions_Call{Call: _e.mock.On("RoleRemoveAllActions", ctx, role)} -} - -func (_c *Repository_RoleRemoveAllActions_Call) Run(run func(ctx context.Context, role roles.Role)) *Repository_RoleRemoveAllActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RoleRemoveAllActions_Call) Return(err error) *Repository_RoleRemoveAllActions_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RoleRemoveAllActions_Call) RunAndReturn(run func(ctx context.Context, role roles.Role) error) *Repository_RoleRemoveAllActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveAllMembers provides a mock function for the type Repository -func (_mock *Repository) RoleRemoveAllMembers(ctx context.Context, role roles.Role) error { - ret := _mock.Called(ctx, role) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveAllMembers") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role) error); ok { - r0 = returnFunc(ctx, role) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RoleRemoveAllMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveAllMembers' -type Repository_RoleRemoveAllMembers_Call struct { - *mock.Call -} - -// RoleRemoveAllMembers is a helper method to define mock.On call -// - ctx context.Context -// - role roles.Role -func (_e *Repository_Expecter) RoleRemoveAllMembers(ctx interface{}, role interface{}) *Repository_RoleRemoveAllMembers_Call { - return &Repository_RoleRemoveAllMembers_Call{Call: _e.mock.On("RoleRemoveAllMembers", ctx, role)} -} - -func (_c *Repository_RoleRemoveAllMembers_Call) Run(run func(ctx context.Context, role roles.Role)) *Repository_RoleRemoveAllMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RoleRemoveAllMembers_Call) Return(err error) *Repository_RoleRemoveAllMembers_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RoleRemoveAllMembers_Call) RunAndReturn(run func(ctx context.Context, role roles.Role) error) *Repository_RoleRemoveAllMembers_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveMembers provides a mock function for the type Repository -func (_mock *Repository) RoleRemoveMembers(ctx context.Context, role roles.Role, members []string) error { - ret := _mock.Called(ctx, role, members) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveMembers") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role, []string) error); ok { - r0 = returnFunc(ctx, role, members) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RoleRemoveMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveMembers' -type Repository_RoleRemoveMembers_Call struct { - *mock.Call -} - -// RoleRemoveMembers is a helper method to define mock.On call -// - ctx context.Context -// - role roles.Role -// - members []string -func (_e *Repository_Expecter) RoleRemoveMembers(ctx interface{}, role interface{}, members interface{}) *Repository_RoleRemoveMembers_Call { - return &Repository_RoleRemoveMembers_Call{Call: _e.mock.On("RoleRemoveMembers", ctx, role, members)} -} - -func (_c *Repository_RoleRemoveMembers_Call) Run(run func(ctx context.Context, role roles.Role, members []string)) *Repository_RoleRemoveMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RoleRemoveMembers_Call) Return(err error) *Repository_RoleRemoveMembers_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RoleRemoveMembers_Call) RunAndReturn(run func(ctx context.Context, role roles.Role, members []string) error) *Repository_RoleRemoveMembers_Call { - _c.Call.Return(run) - return _c -} - -// UpdateRole provides a mock function for the type Repository -func (_mock *Repository) UpdateRole(ctx context.Context, ro roles.Role) (roles.Role, error) { - ret := _mock.Called(ctx, ro) - - if len(ret) == 0 { - panic("no return value specified for UpdateRole") - } - - var r0 roles.Role - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role) (roles.Role, error)); ok { - return returnFunc(ctx, ro) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role) roles.Role); ok { - r0 = returnFunc(ctx, ro) - } else { - r0 = ret.Get(0).(roles.Role) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, roles.Role) error); ok { - r1 = returnFunc(ctx, ro) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_UpdateRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateRole' -type Repository_UpdateRole_Call struct { - *mock.Call -} - -// UpdateRole is a helper method to define mock.On call -// - ctx context.Context -// - ro roles.Role -func (_e *Repository_Expecter) UpdateRole(ctx interface{}, ro interface{}) *Repository_UpdateRole_Call { - return &Repository_UpdateRole_Call{Call: _e.mock.On("UpdateRole", ctx, ro)} -} - -func (_c *Repository_UpdateRole_Call) Run(run func(ctx context.Context, ro roles.Role)) *Repository_UpdateRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_UpdateRole_Call) Return(role roles.Role, err error) *Repository_UpdateRole_Call { - _c.Call.Return(role, err) - return _c -} - -func (_c *Repository_UpdateRole_Call) RunAndReturn(run func(ctx context.Context, ro roles.Role) (roles.Role, error)) *Repository_UpdateRole_Call { - _c.Call.Return(run) - return _c -} - // UpdateRule provides a mock function for the type Repository func (_mock *Repository) UpdateRule(ctx context.Context, r re.Rule) (re.Rule, error) { ret := _mock.Called(ctx, r) diff --git a/re/mocks/service.go b/re/mocks/service.go index 44fb36309..b53ff53d4 100644 --- a/re/mocks/service.go +++ b/re/mocks/service.go @@ -13,7 +13,6 @@ import ( "github.com/absmach/magistrala/pkg/authn" "github.com/absmach/magistrala/pkg/messaging" - "github.com/absmach/magistrala/pkg/roles" "github.com/absmach/magistrala/re" mock "github.com/stretchr/testify/mock" ) @@ -45,98 +44,8 @@ func (_m *Service) EXPECT() *Service_Expecter { return &Service_Expecter{mock: &_m.Mock} } -// AddRole provides a mock function for the type Service -func (_mock *Service) AddRole(ctx context.Context, session authn.Session, entityID string, roleName string, optionalActions []string, optionalMembers []string) (roles.RoleProvision, error) { - ret := _mock.Called(ctx, session, entityID, roleName, optionalActions, optionalMembers) - - if len(ret) == 0 { - panic("no return value specified for AddRole") - } - - var r0 roles.RoleProvision - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string, []string) (roles.RoleProvision, error)); ok { - return returnFunc(ctx, session, entityID, roleName, optionalActions, optionalMembers) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string, []string) roles.RoleProvision); ok { - r0 = returnFunc(ctx, session, entityID, roleName, optionalActions, optionalMembers) - } else { - r0 = ret.Get(0).(roles.RoleProvision) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, []string, []string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleName, optionalActions, optionalMembers) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_AddRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'AddRole' -type Service_AddRole_Call struct { - *mock.Call -} - -// AddRole is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleName string -// - optionalActions []string -// - optionalMembers []string -func (_e *Service_Expecter) AddRole(ctx interface{}, session interface{}, entityID interface{}, roleName interface{}, optionalActions interface{}, optionalMembers interface{}) *Service_AddRole_Call { - return &Service_AddRole_Call{Call: _e.mock.On("AddRole", ctx, session, entityID, roleName, optionalActions, optionalMembers)} -} - -func (_c *Service_AddRole_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleName string, optionalActions []string, optionalMembers []string)) *Service_AddRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - var arg5 []string - if args[5] != nil { - arg5 = args[5].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - arg5, - ) - }) - return _c -} - -func (_c *Service_AddRole_Call) Return(roleProvision roles.RoleProvision, err error) *Service_AddRole_Call { - _c.Call.Return(roleProvision, err) - return _c -} - -func (_c *Service_AddRole_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleName string, optionalActions []string, optionalMembers []string) (roles.RoleProvision, error)) *Service_AddRole_Call { - _c.Call.Return(run) - return _c -} - // AddRule provides a mock function for the type Service -func (_mock *Service) AddRule(ctx context.Context, session authn.Session, r re.Rule) (re.Rule, []roles.RoleProvision, error) { +func (_mock *Service) AddRule(ctx context.Context, session authn.Session, r re.Rule) (re.Rule, error) { ret := _mock.Called(ctx, session, r) if len(ret) == 0 { @@ -144,9 +53,8 @@ func (_mock *Service) AddRule(ctx context.Context, session authn.Session, r re.R } var r0 re.Rule - var r1 []roles.RoleProvision - var r2 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, re.Rule) (re.Rule, []roles.RoleProvision, error)); ok { + var r1 error + if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, re.Rule) (re.Rule, error)); ok { return returnFunc(ctx, session, r) } if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, re.Rule) re.Rule); ok { @@ -154,19 +62,12 @@ func (_mock *Service) AddRule(ctx context.Context, session authn.Session, r re.R } else { r0 = ret.Get(0).(re.Rule) } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, re.Rule) []roles.RoleProvision); ok { + if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, re.Rule) error); ok { r1 = returnFunc(ctx, session, r) } else { - if ret.Get(1) != nil { - r1 = ret.Get(1).([]roles.RoleProvision) - } + r1 = ret.Error(1) } - if returnFunc, ok := ret.Get(2).(func(context.Context, authn.Session, re.Rule) error); ok { - r2 = returnFunc(ctx, session, r) - } else { - r2 = ret.Error(2) - } - return r0, r1, r2 + return r0, r1 } // Service_AddRule_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'AddRule' @@ -205,12 +106,12 @@ func (_c *Service_AddRule_Call) Run(run func(ctx context.Context, session authn. return _c } -func (_c *Service_AddRule_Call) Return(rule re.Rule, roleProvisions []roles.RoleProvision, err error) *Service_AddRule_Call { - _c.Call.Return(rule, roleProvisions, err) +func (_c *Service_AddRule_Call) Return(rule re.Rule, err error) *Service_AddRule_Call { + _c.Call.Return(rule, err) return _c } -func (_c *Service_AddRule_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, r re.Rule) (re.Rule, []roles.RoleProvision, error)) *Service_AddRule_Call { +func (_c *Service_AddRule_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, r re.Rule) (re.Rule, error)) *Service_AddRule_Call { _c.Call.Return(run) return _c } @@ -454,152 +355,6 @@ func (_c *Service_Handle_Call) RunAndReturn(run func(msg *messaging.Message) err return _c } -// ListAvailableActions provides a mock function for the type Service -func (_mock *Service) ListAvailableActions(ctx context.Context, session authn.Session) ([]string, error) { - ret := _mock.Called(ctx, session) - - if len(ret) == 0 { - panic("no return value specified for ListAvailableActions") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session) ([]string, error)); ok { - return returnFunc(ctx, session) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session) []string); ok { - r0 = returnFunc(ctx, session) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session) error); ok { - r1 = returnFunc(ctx, session) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_ListAvailableActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ListAvailableActions' -type Service_ListAvailableActions_Call struct { - *mock.Call -} - -// ListAvailableActions is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -func (_e *Service_Expecter) ListAvailableActions(ctx interface{}, session interface{}) *Service_ListAvailableActions_Call { - return &Service_ListAvailableActions_Call{Call: _e.mock.On("ListAvailableActions", ctx, session)} -} - -func (_c *Service_ListAvailableActions_Call) Run(run func(ctx context.Context, session authn.Session)) *Service_ListAvailableActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Service_ListAvailableActions_Call) Return(strings []string, err error) *Service_ListAvailableActions_Call { - _c.Call.Return(strings, err) - return _c -} - -func (_c *Service_ListAvailableActions_Call) RunAndReturn(run func(ctx context.Context, session authn.Session) ([]string, error)) *Service_ListAvailableActions_Call { - _c.Call.Return(run) - return _c -} - -// ListEntityMembers provides a mock function for the type Service -func (_mock *Service) ListEntityMembers(ctx context.Context, session authn.Session, entityID string, pq roles.MembersRolePageQuery) (roles.MembersRolePage, error) { - ret := _mock.Called(ctx, session, entityID, pq) - - if len(ret) == 0 { - panic("no return value specified for ListEntityMembers") - } - - var r0 roles.MembersRolePage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, roles.MembersRolePageQuery) (roles.MembersRolePage, error)); ok { - return returnFunc(ctx, session, entityID, pq) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, roles.MembersRolePageQuery) roles.MembersRolePage); ok { - r0 = returnFunc(ctx, session, entityID, pq) - } else { - r0 = ret.Get(0).(roles.MembersRolePage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, roles.MembersRolePageQuery) error); ok { - r1 = returnFunc(ctx, session, entityID, pq) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_ListEntityMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ListEntityMembers' -type Service_ListEntityMembers_Call struct { - *mock.Call -} - -// ListEntityMembers is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - pq roles.MembersRolePageQuery -func (_e *Service_Expecter) ListEntityMembers(ctx interface{}, session interface{}, entityID interface{}, pq interface{}) *Service_ListEntityMembers_Call { - return &Service_ListEntityMembers_Call{Call: _e.mock.On("ListEntityMembers", ctx, session, entityID, pq)} -} - -func (_c *Service_ListEntityMembers_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, pq roles.MembersRolePageQuery)) *Service_ListEntityMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 roles.MembersRolePageQuery - if args[3] != nil { - arg3 = args[3].(roles.MembersRolePageQuery) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_ListEntityMembers_Call) Return(membersRolePage roles.MembersRolePage, err error) *Service_ListEntityMembers_Call { - _c.Call.Return(membersRolePage, err) - return _c -} - -func (_c *Service_ListEntityMembers_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, pq roles.MembersRolePageQuery) (roles.MembersRolePage, error)) *Service_ListEntityMembers_Call { - _c.Call.Return(run) - return _c -} - // ListRules provides a mock function for the type Service func (_mock *Service) ListRules(ctx context.Context, session authn.Session, pm re.PageMeta) (re.Page, error) { ret := _mock.Called(ctx, session, pm) @@ -672,207 +427,6 @@ func (_c *Service_ListRules_Call) RunAndReturn(run func(ctx context.Context, ses return _c } -// RemoveEntityMembers provides a mock function for the type Service -func (_mock *Service) RemoveEntityMembers(ctx context.Context, session authn.Session, entityID string, members []string) error { - ret := _mock.Called(ctx, session, entityID, members) - - if len(ret) == 0 { - panic("no return value specified for RemoveEntityMembers") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, []string) error); ok { - r0 = returnFunc(ctx, session, entityID, members) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RemoveEntityMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveEntityMembers' -type Service_RemoveEntityMembers_Call struct { - *mock.Call -} - -// RemoveEntityMembers is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - members []string -func (_e *Service_Expecter) RemoveEntityMembers(ctx interface{}, session interface{}, entityID interface{}, members interface{}) *Service_RemoveEntityMembers_Call { - return &Service_RemoveEntityMembers_Call{Call: _e.mock.On("RemoveEntityMembers", ctx, session, entityID, members)} -} - -func (_c *Service_RemoveEntityMembers_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, members []string)) *Service_RemoveEntityMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 []string - if args[3] != nil { - arg3 = args[3].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_RemoveEntityMembers_Call) Return(err error) *Service_RemoveEntityMembers_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RemoveEntityMembers_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, members []string) error) *Service_RemoveEntityMembers_Call { - _c.Call.Return(run) - return _c -} - -// RemoveMemberFromAllRoles provides a mock function for the type Service -func (_mock *Service) RemoveMemberFromAllRoles(ctx context.Context, session authn.Session, memberID string) error { - ret := _mock.Called(ctx, session, memberID) - - if len(ret) == 0 { - panic("no return value specified for RemoveMemberFromAllRoles") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) error); ok { - r0 = returnFunc(ctx, session, memberID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RemoveMemberFromAllRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveMemberFromAllRoles' -type Service_RemoveMemberFromAllRoles_Call struct { - *mock.Call -} - -// RemoveMemberFromAllRoles is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - memberID string -func (_e *Service_Expecter) RemoveMemberFromAllRoles(ctx interface{}, session interface{}, memberID interface{}) *Service_RemoveMemberFromAllRoles_Call { - return &Service_RemoveMemberFromAllRoles_Call{Call: _e.mock.On("RemoveMemberFromAllRoles", ctx, session, memberID)} -} - -func (_c *Service_RemoveMemberFromAllRoles_Call) Run(run func(ctx context.Context, session authn.Session, memberID string)) *Service_RemoveMemberFromAllRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_RemoveMemberFromAllRoles_Call) Return(err error) *Service_RemoveMemberFromAllRoles_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RemoveMemberFromAllRoles_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, memberID string) error) *Service_RemoveMemberFromAllRoles_Call { - _c.Call.Return(run) - return _c -} - -// RemoveRole provides a mock function for the type Service -func (_mock *Service) RemoveRole(ctx context.Context, session authn.Session, entityID string, roleID string) error { - ret := _mock.Called(ctx, session, entityID, roleID) - - if len(ret) == 0 { - panic("no return value specified for RemoveRole") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) error); ok { - r0 = returnFunc(ctx, session, entityID, roleID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RemoveRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveRole' -type Service_RemoveRole_Call struct { - *mock.Call -} - -// RemoveRole is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -func (_e *Service_Expecter) RemoveRole(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}) *Service_RemoveRole_Call { - return &Service_RemoveRole_Call{Call: _e.mock.On("RemoveRole", ctx, session, entityID, roleID)} -} - -func (_c *Service_RemoveRole_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string)) *Service_RemoveRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_RemoveRole_Call) Return(err error) *Service_RemoveRole_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RemoveRole_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string) error) *Service_RemoveRole_Call { - _c.Call.Return(run) - return _c -} - // RemoveRule provides a mock function for the type Service func (_mock *Service) RemoveRule(ctx context.Context, session authn.Session, id string) error { ret := _mock.Called(ctx, session, id) @@ -936,966 +490,6 @@ func (_c *Service_RemoveRule_Call) RunAndReturn(run func(ctx context.Context, se return _c } -// RetrieveAllRoles provides a mock function for the type Service -func (_mock *Service) RetrieveAllRoles(ctx context.Context, session authn.Session, entityID string, limit uint64, offset uint64) (roles.RolePage, error) { - ret := _mock.Called(ctx, session, entityID, limit, offset) - - if len(ret) == 0 { - panic("no return value specified for RetrieveAllRoles") - } - - var r0 roles.RolePage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, uint64, uint64) (roles.RolePage, error)); ok { - return returnFunc(ctx, session, entityID, limit, offset) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, uint64, uint64) roles.RolePage); ok { - r0 = returnFunc(ctx, session, entityID, limit, offset) - } else { - r0 = ret.Get(0).(roles.RolePage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, uint64, uint64) error); ok { - r1 = returnFunc(ctx, session, entityID, limit, offset) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RetrieveAllRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveAllRoles' -type Service_RetrieveAllRoles_Call struct { - *mock.Call -} - -// RetrieveAllRoles is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - limit uint64 -// - offset uint64 -func (_e *Service_Expecter) RetrieveAllRoles(ctx interface{}, session interface{}, entityID interface{}, limit interface{}, offset interface{}) *Service_RetrieveAllRoles_Call { - return &Service_RetrieveAllRoles_Call{Call: _e.mock.On("RetrieveAllRoles", ctx, session, entityID, limit, offset)} -} - -func (_c *Service_RetrieveAllRoles_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, limit uint64, offset uint64)) *Service_RetrieveAllRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 uint64 - if args[3] != nil { - arg3 = args[3].(uint64) - } - var arg4 uint64 - if args[4] != nil { - arg4 = args[4].(uint64) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_RetrieveAllRoles_Call) Return(rolePage roles.RolePage, err error) *Service_RetrieveAllRoles_Call { - _c.Call.Return(rolePage, err) - return _c -} - -func (_c *Service_RetrieveAllRoles_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, limit uint64, offset uint64) (roles.RolePage, error)) *Service_RetrieveAllRoles_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveRole provides a mock function for the type Service -func (_mock *Service) RetrieveRole(ctx context.Context, session authn.Session, entityID string, roleID string) (roles.Role, error) { - ret := _mock.Called(ctx, session, entityID, roleID) - - if len(ret) == 0 { - panic("no return value specified for RetrieveRole") - } - - var r0 roles.Role - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) (roles.Role, error)); ok { - return returnFunc(ctx, session, entityID, roleID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) roles.Role); ok { - r0 = returnFunc(ctx, session, entityID, roleID) - } else { - r0 = ret.Get(0).(roles.Role) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RetrieveRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveRole' -type Service_RetrieveRole_Call struct { - *mock.Call -} - -// RetrieveRole is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -func (_e *Service_Expecter) RetrieveRole(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}) *Service_RetrieveRole_Call { - return &Service_RetrieveRole_Call{Call: _e.mock.On("RetrieveRole", ctx, session, entityID, roleID)} -} - -func (_c *Service_RetrieveRole_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string)) *Service_RetrieveRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_RetrieveRole_Call) Return(role roles.Role, err error) *Service_RetrieveRole_Call { - _c.Call.Return(role, err) - return _c -} - -func (_c *Service_RetrieveRole_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string) (roles.Role, error)) *Service_RetrieveRole_Call { - _c.Call.Return(run) - return _c -} - -// RoleAddActions provides a mock function for the type Service -func (_mock *Service) RoleAddActions(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string) ([]string, error) { - ret := _mock.Called(ctx, session, entityID, roleID, actions) - - if len(ret) == 0 { - panic("no return value specified for RoleAddActions") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) ([]string, error)); ok { - return returnFunc(ctx, session, entityID, roleID, actions) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) []string); ok { - r0 = returnFunc(ctx, session, entityID, roleID, actions) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, []string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID, actions) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RoleAddActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleAddActions' -type Service_RoleAddActions_Call struct { - *mock.Call -} - -// RoleAddActions is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - actions []string -func (_e *Service_Expecter) RoleAddActions(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, actions interface{}) *Service_RoleAddActions_Call { - return &Service_RoleAddActions_Call{Call: _e.mock.On("RoleAddActions", ctx, session, entityID, roleID, actions)} -} - -func (_c *Service_RoleAddActions_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string)) *Service_RoleAddActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_RoleAddActions_Call) Return(ops []string, err error) *Service_RoleAddActions_Call { - _c.Call.Return(ops, err) - return _c -} - -func (_c *Service_RoleAddActions_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string) ([]string, error)) *Service_RoleAddActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleAddMembers provides a mock function for the type Service -func (_mock *Service) RoleAddMembers(ctx context.Context, session authn.Session, entityID string, roleID string, members []string) ([]string, error) { - ret := _mock.Called(ctx, session, entityID, roleID, members) - - if len(ret) == 0 { - panic("no return value specified for RoleAddMembers") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) ([]string, error)); ok { - return returnFunc(ctx, session, entityID, roleID, members) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) []string); ok { - r0 = returnFunc(ctx, session, entityID, roleID, members) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, []string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID, members) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RoleAddMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleAddMembers' -type Service_RoleAddMembers_Call struct { - *mock.Call -} - -// RoleAddMembers is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - members []string -func (_e *Service_Expecter) RoleAddMembers(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, members interface{}) *Service_RoleAddMembers_Call { - return &Service_RoleAddMembers_Call{Call: _e.mock.On("RoleAddMembers", ctx, session, entityID, roleID, members)} -} - -func (_c *Service_RoleAddMembers_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, members []string)) *Service_RoleAddMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_RoleAddMembers_Call) Return(strings []string, err error) *Service_RoleAddMembers_Call { - _c.Call.Return(strings, err) - return _c -} - -func (_c *Service_RoleAddMembers_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, members []string) ([]string, error)) *Service_RoleAddMembers_Call { - _c.Call.Return(run) - return _c -} - -// RoleCheckActionsExists provides a mock function for the type Service -func (_mock *Service) RoleCheckActionsExists(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string) (bool, error) { - ret := _mock.Called(ctx, session, entityID, roleID, actions) - - if len(ret) == 0 { - panic("no return value specified for RoleCheckActionsExists") - } - - var r0 bool - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) (bool, error)); ok { - return returnFunc(ctx, session, entityID, roleID, actions) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) bool); ok { - r0 = returnFunc(ctx, session, entityID, roleID, actions) - } else { - r0 = ret.Get(0).(bool) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, []string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID, actions) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RoleCheckActionsExists_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleCheckActionsExists' -type Service_RoleCheckActionsExists_Call struct { - *mock.Call -} - -// RoleCheckActionsExists is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - actions []string -func (_e *Service_Expecter) RoleCheckActionsExists(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, actions interface{}) *Service_RoleCheckActionsExists_Call { - return &Service_RoleCheckActionsExists_Call{Call: _e.mock.On("RoleCheckActionsExists", ctx, session, entityID, roleID, actions)} -} - -func (_c *Service_RoleCheckActionsExists_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string)) *Service_RoleCheckActionsExists_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_RoleCheckActionsExists_Call) Return(b bool, err error) *Service_RoleCheckActionsExists_Call { - _c.Call.Return(b, err) - return _c -} - -func (_c *Service_RoleCheckActionsExists_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string) (bool, error)) *Service_RoleCheckActionsExists_Call { - _c.Call.Return(run) - return _c -} - -// RoleCheckMembersExists provides a mock function for the type Service -func (_mock *Service) RoleCheckMembersExists(ctx context.Context, session authn.Session, entityID string, roleID string, members []string) (bool, error) { - ret := _mock.Called(ctx, session, entityID, roleID, members) - - if len(ret) == 0 { - panic("no return value specified for RoleCheckMembersExists") - } - - var r0 bool - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) (bool, error)); ok { - return returnFunc(ctx, session, entityID, roleID, members) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) bool); ok { - r0 = returnFunc(ctx, session, entityID, roleID, members) - } else { - r0 = ret.Get(0).(bool) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, []string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID, members) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RoleCheckMembersExists_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleCheckMembersExists' -type Service_RoleCheckMembersExists_Call struct { - *mock.Call -} - -// RoleCheckMembersExists is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - members []string -func (_e *Service_Expecter) RoleCheckMembersExists(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, members interface{}) *Service_RoleCheckMembersExists_Call { - return &Service_RoleCheckMembersExists_Call{Call: _e.mock.On("RoleCheckMembersExists", ctx, session, entityID, roleID, members)} -} - -func (_c *Service_RoleCheckMembersExists_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, members []string)) *Service_RoleCheckMembersExists_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_RoleCheckMembersExists_Call) Return(b bool, err error) *Service_RoleCheckMembersExists_Call { - _c.Call.Return(b, err) - return _c -} - -func (_c *Service_RoleCheckMembersExists_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, members []string) (bool, error)) *Service_RoleCheckMembersExists_Call { - _c.Call.Return(run) - return _c -} - -// RoleListActions provides a mock function for the type Service -func (_mock *Service) RoleListActions(ctx context.Context, session authn.Session, entityID string, roleID string) ([]string, error) { - ret := _mock.Called(ctx, session, entityID, roleID) - - if len(ret) == 0 { - panic("no return value specified for RoleListActions") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) ([]string, error)); ok { - return returnFunc(ctx, session, entityID, roleID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) []string); ok { - r0 = returnFunc(ctx, session, entityID, roleID) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RoleListActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleListActions' -type Service_RoleListActions_Call struct { - *mock.Call -} - -// RoleListActions is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -func (_e *Service_Expecter) RoleListActions(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}) *Service_RoleListActions_Call { - return &Service_RoleListActions_Call{Call: _e.mock.On("RoleListActions", ctx, session, entityID, roleID)} -} - -func (_c *Service_RoleListActions_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string)) *Service_RoleListActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_RoleListActions_Call) Return(strings []string, err error) *Service_RoleListActions_Call { - _c.Call.Return(strings, err) - return _c -} - -func (_c *Service_RoleListActions_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string) ([]string, error)) *Service_RoleListActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleListMembers provides a mock function for the type Service -func (_mock *Service) RoleListMembers(ctx context.Context, session authn.Session, entityID string, roleID string, limit uint64, offset uint64) (roles.MembersPage, error) { - ret := _mock.Called(ctx, session, entityID, roleID, limit, offset) - - if len(ret) == 0 { - panic("no return value specified for RoleListMembers") - } - - var r0 roles.MembersPage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, uint64, uint64) (roles.MembersPage, error)); ok { - return returnFunc(ctx, session, entityID, roleID, limit, offset) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, uint64, uint64) roles.MembersPage); ok { - r0 = returnFunc(ctx, session, entityID, roleID, limit, offset) - } else { - r0 = ret.Get(0).(roles.MembersPage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, uint64, uint64) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID, limit, offset) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RoleListMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleListMembers' -type Service_RoleListMembers_Call struct { - *mock.Call -} - -// RoleListMembers is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - limit uint64 -// - offset uint64 -func (_e *Service_Expecter) RoleListMembers(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, limit interface{}, offset interface{}) *Service_RoleListMembers_Call { - return &Service_RoleListMembers_Call{Call: _e.mock.On("RoleListMembers", ctx, session, entityID, roleID, limit, offset)} -} - -func (_c *Service_RoleListMembers_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, limit uint64, offset uint64)) *Service_RoleListMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 uint64 - if args[4] != nil { - arg4 = args[4].(uint64) - } - var arg5 uint64 - if args[5] != nil { - arg5 = args[5].(uint64) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - arg5, - ) - }) - return _c -} - -func (_c *Service_RoleListMembers_Call) Return(membersPage roles.MembersPage, err error) *Service_RoleListMembers_Call { - _c.Call.Return(membersPage, err) - return _c -} - -func (_c *Service_RoleListMembers_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, limit uint64, offset uint64) (roles.MembersPage, error)) *Service_RoleListMembers_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveActions provides a mock function for the type Service -func (_mock *Service) RoleRemoveActions(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string) error { - ret := _mock.Called(ctx, session, entityID, roleID, actions) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveActions") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) error); ok { - r0 = returnFunc(ctx, session, entityID, roleID, actions) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RoleRemoveActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveActions' -type Service_RoleRemoveActions_Call struct { - *mock.Call -} - -// RoleRemoveActions is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - actions []string -func (_e *Service_Expecter) RoleRemoveActions(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, actions interface{}) *Service_RoleRemoveActions_Call { - return &Service_RoleRemoveActions_Call{Call: _e.mock.On("RoleRemoveActions", ctx, session, entityID, roleID, actions)} -} - -func (_c *Service_RoleRemoveActions_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string)) *Service_RoleRemoveActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_RoleRemoveActions_Call) Return(err error) *Service_RoleRemoveActions_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RoleRemoveActions_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string) error) *Service_RoleRemoveActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveAllActions provides a mock function for the type Service -func (_mock *Service) RoleRemoveAllActions(ctx context.Context, session authn.Session, entityID string, roleID string) error { - ret := _mock.Called(ctx, session, entityID, roleID) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveAllActions") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) error); ok { - r0 = returnFunc(ctx, session, entityID, roleID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RoleRemoveAllActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveAllActions' -type Service_RoleRemoveAllActions_Call struct { - *mock.Call -} - -// RoleRemoveAllActions is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -func (_e *Service_Expecter) RoleRemoveAllActions(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}) *Service_RoleRemoveAllActions_Call { - return &Service_RoleRemoveAllActions_Call{Call: _e.mock.On("RoleRemoveAllActions", ctx, session, entityID, roleID)} -} - -func (_c *Service_RoleRemoveAllActions_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string)) *Service_RoleRemoveAllActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_RoleRemoveAllActions_Call) Return(err error) *Service_RoleRemoveAllActions_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RoleRemoveAllActions_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string) error) *Service_RoleRemoveAllActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveAllMembers provides a mock function for the type Service -func (_mock *Service) RoleRemoveAllMembers(ctx context.Context, session authn.Session, entityID string, roleID string) error { - ret := _mock.Called(ctx, session, entityID, roleID) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveAllMembers") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) error); ok { - r0 = returnFunc(ctx, session, entityID, roleID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RoleRemoveAllMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveAllMembers' -type Service_RoleRemoveAllMembers_Call struct { - *mock.Call -} - -// RoleRemoveAllMembers is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -func (_e *Service_Expecter) RoleRemoveAllMembers(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}) *Service_RoleRemoveAllMembers_Call { - return &Service_RoleRemoveAllMembers_Call{Call: _e.mock.On("RoleRemoveAllMembers", ctx, session, entityID, roleID)} -} - -func (_c *Service_RoleRemoveAllMembers_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string)) *Service_RoleRemoveAllMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_RoleRemoveAllMembers_Call) Return(err error) *Service_RoleRemoveAllMembers_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RoleRemoveAllMembers_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string) error) *Service_RoleRemoveAllMembers_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveMembers provides a mock function for the type Service -func (_mock *Service) RoleRemoveMembers(ctx context.Context, session authn.Session, entityID string, roleID string, members []string) error { - ret := _mock.Called(ctx, session, entityID, roleID, members) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveMembers") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) error); ok { - r0 = returnFunc(ctx, session, entityID, roleID, members) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RoleRemoveMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveMembers' -type Service_RoleRemoveMembers_Call struct { - *mock.Call -} - -// RoleRemoveMembers is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - members []string -func (_e *Service_Expecter) RoleRemoveMembers(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, members interface{}) *Service_RoleRemoveMembers_Call { - return &Service_RoleRemoveMembers_Call{Call: _e.mock.On("RoleRemoveMembers", ctx, session, entityID, roleID, members)} -} - -func (_c *Service_RoleRemoveMembers_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, members []string)) *Service_RoleRemoveMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_RoleRemoveMembers_Call) Return(err error) *Service_RoleRemoveMembers_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RoleRemoveMembers_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, members []string) error) *Service_RoleRemoveMembers_Call { - _c.Call.Return(run) - return _c -} - // StartScheduler provides a mock function for the type Service func (_mock *Service) StartScheduler(ctx context.Context) error { ret := _mock.Called(ctx) @@ -1947,90 +541,6 @@ func (_c *Service_StartScheduler_Call) RunAndReturn(run func(ctx context.Context return _c } -// UpdateRoleName provides a mock function for the type Service -func (_mock *Service) UpdateRoleName(ctx context.Context, session authn.Session, entityID string, roleID string, newRoleName string) (roles.Role, error) { - ret := _mock.Called(ctx, session, entityID, roleID, newRoleName) - - if len(ret) == 0 { - panic("no return value specified for UpdateRoleName") - } - - var r0 roles.Role - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, string) (roles.Role, error)); ok { - return returnFunc(ctx, session, entityID, roleID, newRoleName) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, string) roles.Role); ok { - r0 = returnFunc(ctx, session, entityID, roleID, newRoleName) - } else { - r0 = ret.Get(0).(roles.Role) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID, newRoleName) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_UpdateRoleName_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateRoleName' -type Service_UpdateRoleName_Call struct { - *mock.Call -} - -// UpdateRoleName is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - newRoleName string -func (_e *Service_Expecter) UpdateRoleName(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, newRoleName interface{}) *Service_UpdateRoleName_Call { - return &Service_UpdateRoleName_Call{Call: _e.mock.On("UpdateRoleName", ctx, session, entityID, roleID, newRoleName)} -} - -func (_c *Service_UpdateRoleName_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, newRoleName string)) *Service_UpdateRoleName_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 string - if args[4] != nil { - arg4 = args[4].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_UpdateRoleName_Call) Return(role roles.Role, err error) *Service_UpdateRoleName_Call { - _c.Call.Return(role, err) - return _c -} - -func (_c *Service_UpdateRoleName_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, newRoleName string) (roles.Role, error)) *Service_UpdateRoleName_Call { - _c.Call.Return(run) - return _c -} - // UpdateRule provides a mock function for the type Service func (_mock *Service) UpdateRule(ctx context.Context, session authn.Session, r re.Rule) (re.Rule, error) { ret := _mock.Called(ctx, session, r) diff --git a/re/postgres/init.go b/re/postgres/init.go index d710fc3e0..099695b59 100644 --- a/re/postgres/init.go +++ b/re/postgres/init.go @@ -4,19 +4,11 @@ package postgres import ( - dpostgres "github.com/absmach/magistrala/domains/postgres" - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - rolesPostgres "github.com/absmach/magistrala/pkg/roles/repo/postgres" _ "github.com/jackc/pgx/v5/stdlib" // required for SQL access migrate "github.com/rubenv/sql-migrate" ) func Migration() (*migrate.MemoryMigrationSource, error) { - rolesMigration, err := rolesPostgres.Migration(rolesTableNamePrefix, entityTableName, entityIDColumnName) - if err != nil { - return &migrate.MemoryMigrationSource{}, errors.Wrap(repoerr.ErrRoleMigration, err) - } rulesMigration := &migrate.MemoryMigrationSource{ Migrations: []*migrate.Migration{ { @@ -138,13 +130,5 @@ func Migration() (*migrate.MemoryMigrationSource, error) { }, } - rulesMigration.Migrations = append(rulesMigration.Migrations, rolesMigration.Migrations...) - - domainsMigration, err := dpostgres.Migration() - if err != nil { - return &migrate.MemoryMigrationSource{}, errors.Wrap(repoerr.ErrRoleMigration, err) - } - rulesMigration.Migrations = append(rulesMigration.Migrations, domainsMigration.Migrations...) - return rulesMigration, nil } diff --git a/re/postgres/repository.go b/re/postgres/repository.go index 35545ab9b..c05e09623 100644 --- a/re/postgres/repository.go +++ b/re/postgres/repository.go @@ -13,28 +13,17 @@ import ( api "github.com/absmach/magistrala/api/http" "github.com/absmach/magistrala/pkg/errors" repoerr "github.com/absmach/magistrala/pkg/errors/repository" - mgPolicies "github.com/absmach/magistrala/pkg/policies" "github.com/absmach/magistrala/pkg/postgres" - rolesPostgres "github.com/absmach/magistrala/pkg/roles/repo/postgres" "github.com/absmach/magistrala/re" ) -const ( - rolesTableNamePrefix = "rules" - entityTableName = "rules" - entityIDColumnName = "id" -) - type PostgresRepository struct { DB postgres.Database - rolesPostgres.Repository } func NewRepository(db postgres.Database) re.Repository { - rolesRepo := rolesPostgres.NewRepository(db, mgPolicies.RulesType, rolesTableNamePrefix, entityTableName, entityIDColumnName) return &PostgresRepository{ - DB: db, - Repository: rolesRepo, + DB: db, } } @@ -95,157 +84,6 @@ func (repo *PostgresRepository) ViewRule(ctx context.Context, id string) (re.Rul return ret, nil } -func (repo *PostgresRepository) RetrieveByIDWithRoles(ctx context.Context, id, memberID string) (re.Rule, error) { - query := ` - WITH selected_rule AS ( - SELECT - r.id, - r.domain_id - FROM - rules r - WHERE - r.id = :id - LIMIT 1 - ), - selected_rule_roles AS ( - SELECT - rr.entity_id AS rule_id, - rrm.member_id AS member_id, - rr.id AS role_id, - rr."name" AS role_name, - jsonb_agg(DISTINCT rra."action") AS actions, - 'direct' AS access_type, - '' AS access_provider_id - FROM - rules_roles rr - JOIN - rules_role_members rrm ON rr.id = rrm.role_id - JOIN - rules_role_actions rra ON rr.id = rra.role_id - JOIN - selected_rule sr ON sr.id = rr.entity_id - AND rrm.member_id = :member_id - GROUP BY - rr.entity_id, rr.id, rr.name, rrm.member_id - ), - selected_domain_roles AS ( - SELECT - sr.id AS rule_id, - drm.member_id AS member_id, - dr.id AS role_id, - dr."name" AS role_name, - jsonb_agg(DISTINCT all_actions."action") AS actions, - 'domain' AS access_type, - dr.entity_id AS access_provider_id - FROM - domains d - JOIN - selected_rule sr ON sr.domain_id = d.id - JOIN - domains_roles dr ON dr.entity_id = d.id - JOIN - domains_role_members drm ON dr.id = drm.role_id - JOIN - domains_role_actions dra ON dr.id = dra.role_id - JOIN - domains_role_actions all_actions ON dr.id = all_actions.role_id - WHERE - drm.member_id = :member_id - AND dra."action" LIKE 'rule%' - GROUP BY - sr.id, dr.entity_id, dr.id, dr."name", drm.member_id - ), - all_roles AS ( - SELECT - srr.rule_id, - srr.member_id, - srr.role_id, - srr.role_name, - srr.actions, - srr.access_type, - srr.access_provider_id - FROM - selected_rule_roles srr - UNION - SELECT - sdr.rule_id, - sdr.member_id, - sdr.role_id, - sdr.role_name, - sdr.actions, - sdr.access_type, - sdr.access_provider_id - FROM - selected_domain_roles sdr - ), - final_roles AS ( - SELECT - ar.rule_id, - ar.member_id, - jsonb_agg( - jsonb_build_object( - 'role_id', ar.role_id, - 'role_name', ar.role_name, - 'actions', ar.actions, - 'access_type', ar.access_type, - 'access_provider_id', ar.access_provider_id - ) - ) AS roles - FROM all_roles ar - GROUP BY - ar.rule_id, ar.member_id - ) - SELECT - r2.id, - r2."name", - r2.domain_id, - r2.tags, - r2.metadata, - r2.input_channel, - r2.input_topic, - r2.outputs, - r2.status, - r2.logic_type, - r2.logic_value, - r2.time, - r2.recurring, - r2.recurring_period, - r2.start_datetime, - r2.created_at, - r2.created_by, - r2.updated_at, - r2.updated_by, - fr.member_id, - fr.roles - FROM rules r2 - JOIN final_roles fr ON fr.rule_id = r2.id - ` - parameters := map[string]any{ - "id": id, - "member_id": memberID, - } - row, err := repo.DB.NamedQueryContext(ctx, query, parameters) - if err != nil { - return re.Rule{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - defer row.Close() - - dbrule := dbRule{} - if !row.Next() { - return re.Rule{}, repoerr.ErrNotFound - } - - if err := row.StructScan(&dbrule); err != nil { - return re.Rule{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - - r, err := dbToRule(dbrule) - if err != nil { - return re.Rule{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - return r, nil -} - func (repo *PostgresRepository) UpdateRuleStatus(ctx context.Context, r re.Rule) (re.Rule, error) { q := `UPDATE rules SET status = :status, updated_at = :updated_at, updated_by = :updated_by @@ -402,104 +240,6 @@ func (repo *PostgresRepository) ListAllRules(ctx context.Context, pm re.PageMeta return ret, nil } -func (repo *PostgresRepository) ListUserRules(ctx context.Context, userID string, pm re.PageMeta) (re.Page, error) { - pm.UserID = userID - - additionalConditions := pageRulesQueryConditions(pm) - additionalWhereClause := "" - if len(additionalConditions) > 0 { - additionalWhereClause = "AND " + strings.Join(additionalConditions, " AND ") - } - - orderClause := rulesOrderClause(pm) - pgData := rulesPageData(pm) - - innerQ := fmt.Sprintf(` - WITH direct_rules AS ( - SELECT r.id, r.name, r.domain_id, r.tags, r.metadata, r.input_channel, r.input_topic, - r.logic_type, r.logic_value, r.outputs, r.start_datetime, r.time, - r.recurring, r.recurring_period, r.created_at, r.created_by, r.updated_at, r.updated_by, r.status, - rr.id AS role_id, - rr."name" AS role_name, - array_remove(array_agg(DISTINCT rra."action"), NULL) AS actions, - 'direct' AS access_type, - '' AS access_provider_id, - '' AS access_provider_role_id, - '' AS access_provider_role_name, - CAST(array[] AS text[]) AS access_provider_role_actions - FROM rules_role_members rrm - JOIN rules_roles rr ON rr.id = rrm.role_id - JOIN rules r ON r.id = rr.entity_id - LEFT JOIN rules_role_actions rra ON rra.role_id = rrm.role_id - WHERE rrm.member_id = :user_id - %s - GROUP BY r.id, rr.id, rr."name" - ), - domain_rules AS ( - SELECT r.id, r.name, r.domain_id, r.tags, r.metadata, r.input_channel, r.input_topic, - r.logic_type, r.logic_value, r.outputs, r.start_datetime, r.time, - r.recurring, r.recurring_period, r.created_at, r.created_by, r.updated_at, r.updated_by, r.status, - '' AS role_id, - '' AS role_name, - CAST(array[] AS text[]) AS actions, - 'domain' AS access_type, - d.id AS access_provider_id, - dr.id AS access_provider_role_id, - dr."name" AS access_provider_role_name, - array_agg(DISTINCT dra."action") AS access_provider_role_actions - FROM domains_role_members drm - JOIN domains_role_actions dra ON dra.role_id = drm.role_id - JOIN domains_roles dr ON dr.id = drm.role_id - JOIN domains d ON d.id = dr.entity_id - JOIN rules r ON r.domain_id = d.id - WHERE drm.member_id = :user_id - AND dra.action LIKE 'rule%%' - AND NOT EXISTS (SELECT 1 FROM direct_rules tmp WHERE tmp.id = r.id) - %s - GROUP BY r.id, d.id, dr.id, dr."name" - ) - SELECT * FROM direct_rules - UNION ALL - SELECT * FROM domain_rules - `, additionalWhereClause, additionalWhereClause) - - q := fmt.Sprintf(` - SELECT * FROM (%s) AS sub %s %s; - `, innerQ, orderClause, pgData) - - rows, err := repo.DB.NamedQueryContext(ctx, q, pm) - if err != nil { - return re.Page{}, err - } - defer rows.Close() - - var rules []re.Rule - for rows.Next() { - var r dbRule - if err := rows.StructScan(&r); err != nil { - return re.Page{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - ret, err := dbToRule(r) - if err != nil { - return re.Page{}, err - } - rules = append(rules, ret) - } - - cq := fmt.Sprintf(`SELECT COUNT(*) FROM (%s) AS count_sub;`, innerQ) - total, err := postgres.Total(ctx, repo.DB, cq, pm) - if err != nil { - return re.Page{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - - return re.Page{ - Total: total, - Offset: pm.Offset, - Limit: pm.Limit, - Rules: rules, - }, nil -} - func (repo *PostgresRepository) UpdateRuleDue(ctx context.Context, id string, due time.Time) (re.Rule, error) { q := ` UPDATE rules diff --git a/re/postgres/repository_test.go b/re/postgres/repository_test.go index 4dda7f3f4..635f16fb2 100644 --- a/re/postgres/repository_test.go +++ b/re/postgres/repository_test.go @@ -935,224 +935,6 @@ func TestListRules(t *testing.T) { } } -func TestListUserRules(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM domains_role_actions") - assert.Nil(t, err, fmt.Sprintf("clean domains_role_actions unexpected error: %s", err)) - _, err = db.Exec("DELETE FROM domains_role_members") - assert.Nil(t, err, fmt.Sprintf("clean domains_role_members unexpected error: %s", err)) - _, err = db.Exec("DELETE FROM domains_roles") - assert.Nil(t, err, fmt.Sprintf("clean domains_roles unexpected error: %s", err)) - _, err = db.Exec("DELETE FROM domains") - assert.Nil(t, err, fmt.Sprintf("clean domains unexpected error: %s", err)) - _, err = db.Exec("DELETE FROM rules") - assert.Nil(t, err, fmt.Sprintf("clean rules unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - domainID := generateUUID(t) - domainRoute := generateUUID(t) - userID := generateUUID(t) - domainUserID := generateUUID(t) - otherUserID := generateUUID(t) - channelID := generateUUID(t) - - _, err := db.Exec(`INSERT INTO domains (id, name, route, status) VALUES ($1, $2, $3, $4)`, domainID, namegen.Generate(), domainRoute, 0) - assert.Nil(t, err, fmt.Sprintf("insert domains unexpected error: %s", err)) - - // Create 10 rules; assign the first 4 to userID via a role. - var allRules []re.Rule - for i := range 10 { - r := re.Rule{ - ID: generateUUID(t), - Name: namegen.Generate(), - DomainID: domainID, - InputChannel: channelID, - Logic: re.Script{Type: re.LuaType, Value: "return true"}, - Status: re.EnabledStatus, - CreatedAt: time.Now().UTC().Add(time.Duration(i) * time.Minute).Truncate(time.Microsecond), - CreatedBy: generateUUID(t), - UpdatedAt: time.Now().UTC().Add(time.Duration(i) * time.Minute).Truncate(time.Microsecond), - UpdatedBy: generateUUID(t), - } - rule, err := repo.AddRule(context.Background(), r) - assert.Nil(t, err, fmt.Sprintf("unexpected error: %s", err)) - allRules = append(allRules, rule) - } - - // Assign userID to the first 4 rules via direct role INSERT. - for i := range 4 { - roleID := generateUUID(t) - _, err := db.Exec(`INSERT INTO rules_roles (id, name, entity_id) VALUES ($1, $2, $3)`, roleID, "admin", allRules[i].ID) - assert.Nil(t, err, fmt.Sprintf("insert rules_roles unexpected error: %s", err)) - _, err = db.Exec(`INSERT INTO rules_role_members (role_id, member_id, entity_id) VALUES ($1, $2, $3)`, roleID, userID, allRules[i].ID) - assert.Nil(t, err, fmt.Sprintf("insert rules_role_members unexpected error: %s", err)) - } - - domainRoleID := generateUUID(t) - _, err = db.Exec(`INSERT INTO domains_roles (id, name, entity_id) VALUES ($1, $2, $3)`, domainRoleID, "admin", domainID) - assert.Nil(t, err, fmt.Sprintf("insert domains_roles unexpected error: %s", err)) - _, err = db.Exec(`INSERT INTO domains_role_members (role_id, member_id, entity_id) VALUES ($1, $2, $3)`, domainRoleID, domainUserID, domainID) - assert.Nil(t, err, fmt.Sprintf("insert domains_role_members unexpected error: %s", err)) - _, err = db.Exec(`INSERT INTO domains_role_actions (role_id, action) VALUES ($1, $2)`, domainRoleID, "rule_read") - assert.Nil(t, err, fmt.Sprintf("insert domains_role_actions unexpected error: %s", err)) - - cases := []struct { - desc string - userID string - pm re.PageMeta - count int - err error - }{ - { - desc: "list user rules returns only accessible rules", - userID: userID, - pm: re.PageMeta{ - Offset: 0, - Limit: 100, - Status: re.AllStatus, - }, - count: 4, - err: nil, - }, - { - desc: "list user rules with offset", - userID: userID, - pm: re.PageMeta{ - Offset: 2, - Limit: 100, - Status: re.AllStatus, - }, - count: 2, - err: nil, - }, - { - desc: "list user rules with limit", - userID: userID, - pm: re.PageMeta{ - Offset: 0, - Limit: 2, - Status: re.AllStatus, - }, - count: 2, - err: nil, - }, - { - desc: "list user rules with domain filter", - userID: userID, - pm: re.PageMeta{ - Domain: domainID, - Offset: 0, - Limit: 100, - Status: re.AllStatus, - }, - count: 4, - err: nil, - }, - { - desc: "list user rules with channel filter", - userID: userID, - pm: re.PageMeta{ - InputChannel: channelID, - Offset: 0, - Limit: 100, - Status: re.AllStatus, - }, - count: 4, - err: nil, - }, - { - desc: "list user rules with non-existing domain returns 0", - userID: userID, - pm: re.PageMeta{ - Domain: generateUUID(t), - Offset: 0, - Limit: 100, - Status: re.AllStatus, - }, - count: 0, - err: nil, - }, - { - desc: "list user rules via domain role returns all domain rules", - userID: domainUserID, - pm: re.PageMeta{ - Offset: 0, - Limit: 100, - Status: re.AllStatus, - }, - count: 10, - err: nil, - }, - { - desc: "list rules for user with no role assignments returns 0", - userID: otherUserID, - pm: re.PageMeta{ - Offset: 0, - Limit: 100, - Status: re.AllStatus, - }, - count: 0, - err: nil, - }, - { - desc: "list user rules ordered by name ascending", - userID: userID, - pm: re.PageMeta{ - Offset: 0, - Limit: 100, - Status: re.AllStatus, - Order: nameOrder, - Dir: ascDir, - }, - count: 4, - err: nil, - }, - { - desc: "list user rules ordered by created_at descending", - userID: userID, - pm: re.PageMeta{ - Offset: 0, - Limit: 100, - Status: re.AllStatus, - Order: createdAtOrder, - Dir: descDir, - }, - count: 4, - err: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - page, err := repo.ListUserRules(context.Background(), tc.userID, tc.pm) - if tc.err != nil { - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - return - } - assert.Nil(t, err, fmt.Sprintf("unexpected error: %s", err)) - assert.Equal(t, tc.count, len(page.Rules), fmt.Sprintf("%s: expected %d rules, got %d", tc.desc, tc.count, len(page.Rules))) - if len(page.Rules) > 1 { - switch tc.pm.Order { - case nameOrder: - if tc.pm.Dir == ascDir { - assert.True(t, sort.SliceIsSorted(page.Rules, func(i, j int) bool { - return page.Rules[i].Name <= page.Rules[j].Name - }), "Expected names to be sorted ascending") - } - case createdAtOrder: - if tc.pm.Dir == descDir { - assert.True(t, sort.SliceIsSorted(page.Rules, func(i, j int) bool { - return page.Rules[i].CreatedAt.After(page.Rules[j].CreatedAt) - }), "Expected created_at to be sorted descending") - } - } - } - }) - } -} - func TestRemoveRule(t *testing.T) { t.Cleanup(func() { _, err := db.Exec("DELETE FROM rules") diff --git a/re/postgres/rule.go b/re/postgres/rule.go index 7b67ff650..0dfc94230 100644 --- a/re/postgres/rule.go +++ b/re/postgres/rule.go @@ -9,44 +9,32 @@ import ( "time" "github.com/absmach/magistrala/pkg/errors" - "github.com/absmach/magistrala/pkg/roles" "github.com/absmach/magistrala/pkg/schedule" "github.com/absmach/magistrala/re" "github.com/jackc/pgtype" - "github.com/lib/pq" ) // dbRule represents the database structure for a Rule. type dbRule struct { - ID string `db:"id"` - Name string `db:"name"` - DomainID string `db:"domain_id"` - Tags pgtype.TextArray `db:"tags,omitempty"` - Metadata []byte `db:"metadata,omitempty"` - InputChannel string `db:"input_channel"` - InputTopic sql.NullString `db:"input_topic"` - LogicType re.ScriptType `db:"logic_type"` - LogicValue string `db:"logic_value"` - Outputs []byte `db:"outputs"` - StartDateTime sql.NullTime `db:"start_datetime"` - Time sql.NullTime `db:"time"` - Recurring schedule.Recurring `db:"recurring"` - RecurringPeriod uint `db:"recurring_period"` - Status re.Status `db:"status"` - CreatedAt time.Time `db:"created_at"` - CreatedBy string `db:"created_by"` - UpdatedAt time.Time `db:"updated_at"` - UpdatedBy string `db:"updated_by"` - MemberID string `db:"member_id,omitempty"` - RoleID string `db:"role_id,omitempty"` - RoleName string `db:"role_name,omitempty"` - Actions pq.StringArray `db:"actions,omitempty"` - AccessType string `db:"access_type,omitempty"` - AccessProviderId string `db:"access_provider_id,omitempty"` - AccessProviderRoleId string `db:"access_provider_role_id,omitempty"` - AccessProviderRoleName string `db:"access_provider_role_name,omitempty"` - AccessProviderRoleActions pq.StringArray `db:"access_provider_role_actions,omitempty"` - Roles json.RawMessage `db:"roles,omitempty"` + ID string `db:"id"` + Name string `db:"name"` + DomainID string `db:"domain_id"` + Tags pgtype.TextArray `db:"tags,omitempty"` + Metadata []byte `db:"metadata,omitempty"` + InputChannel string `db:"input_channel"` + InputTopic sql.NullString `db:"input_topic"` + LogicType re.ScriptType `db:"logic_type"` + LogicValue string `db:"logic_value"` + Outputs []byte `db:"outputs"` + StartDateTime sql.NullTime `db:"start_datetime"` + Time sql.NullTime `db:"time"` + Recurring schedule.Recurring `db:"recurring"` + RecurringPeriod uint `db:"recurring_period"` + Status re.Status `db:"status"` + CreatedAt time.Time `db:"created_at"` + CreatedBy string `db:"created_by"` + UpdatedAt time.Time `db:"updated_at"` + UpdatedBy string `db:"updated_by"` } func ruleToDb(r re.Rule) (dbRule, error) { @@ -120,13 +108,6 @@ func dbToRule(dto dbRule) (re.Rule, error) { } } - var roles []roles.MemberRoleActions - if dto.Roles != nil { - if err := json.Unmarshal(dto.Roles, &roles); err != nil { - return re.Rule{}, errors.Wrap(errors.ErrMalformedEntity, err) - } - } - return re.Rule{ ID: dto.ID, Name: dto.Name, @@ -146,20 +127,11 @@ func dbToRule(dto dbRule) (re.Rule, error) { Recurring: dto.Recurring, RecurringPeriod: dto.RecurringPeriod, }, - Status: dto.Status, - CreatedAt: dto.CreatedAt, - CreatedBy: dto.CreatedBy, - UpdatedAt: dto.UpdatedAt, - UpdatedBy: dto.UpdatedBy, - RoleID: dto.RoleID, - RoleName: dto.RoleName, - Actions: []string(dto.Actions), - AccessType: dto.AccessType, - AccessProviderId: dto.AccessProviderId, - AccessProviderRoleId: dto.AccessProviderRoleId, - AccessProviderRoleName: dto.AccessProviderRoleName, - AccessProviderRoleActions: []string(dto.AccessProviderRoleActions), - Roles: roles, + Status: dto.Status, + CreatedAt: dto.CreatedAt, + CreatedBy: dto.CreatedBy, + UpdatedAt: dto.UpdatedAt, + UpdatedBy: dto.UpdatedBy, }, nil } diff --git a/re/rule.go b/re/rule.go index bf1e49735..8a5e3af5e 100644 --- a/re/rule.go +++ b/re/rule.go @@ -11,7 +11,6 @@ import ( "github.com/absmach/magistrala/pkg/authn" "github.com/absmach/magistrala/pkg/errors" "github.com/absmach/magistrala/pkg/messaging" - "github.com/absmach/magistrala/pkg/roles" "github.com/absmach/magistrala/pkg/schedule" "github.com/absmach/magistrala/re/outputs" ) @@ -60,16 +59,6 @@ type Rule struct { CreatedBy string `json:"created_by"` UpdatedAt time.Time `json:"updated_at"` UpdatedBy string `json:"updated_by"` - // Extended - RoleID string `json:"role_id,omitempty"` - RoleName string `json:"role_name,omitempty"` - Actions []string `json:"actions,omitempty"` - AccessType string `json:"access_type,omitempty"` - AccessProviderId string `json:"access_provider_id,omitempty"` - AccessProviderRoleId string `json:"access_provider_role_id,omitempty"` - AccessProviderRoleName string `json:"access_provider_role_name,omitempty"` - AccessProviderRoleActions []string `json:"access_provider_role_actions,omitempty"` - Roles []roles.MemberRoleActions `json:"roles,omitempty"` } // EventEncode converts a Rule struct to map[string]any at event producer. @@ -237,7 +226,7 @@ type Page struct { type Service interface { messaging.MessageHandler - AddRule(ctx context.Context, session authn.Session, r Rule) (Rule, []roles.RoleProvision, error) + AddRule(ctx context.Context, session authn.Session, r Rule) (Rule, error) ViewRule(ctx context.Context, session authn.Session, id string, withRoles bool) (Rule, error) UpdateRule(ctx context.Context, session authn.Session, r Rule) (Rule, error) UpdateRuleTags(ctx context.Context, session authn.Session, r Rule) (Rule, error) @@ -248,20 +237,16 @@ type Service interface { DisableRule(ctx context.Context, session authn.Session, id string) (Rule, error) StartScheduler(ctx context.Context) error - roles.RoleManager } type Repository interface { AddRule(ctx context.Context, r Rule) (Rule, error) ViewRule(ctx context.Context, id string) (Rule, error) - RetrieveByIDWithRoles(ctx context.Context, id, memberID string) (Rule, error) UpdateRule(ctx context.Context, r Rule) (Rule, error) UpdateRuleTags(ctx context.Context, r Rule) (Rule, error) UpdateRuleSchedule(ctx context.Context, r Rule) (Rule, error) RemoveRule(ctx context.Context, id string) error UpdateRuleStatus(ctx context.Context, r Rule) (Rule, error) ListAllRules(ctx context.Context, pm PageMeta) (Page, error) - ListUserRules(ctx context.Context, userID string, pm PageMeta) (Page, error) UpdateRuleDue(ctx context.Context, id string, due time.Time) (Rule, error) - roles.Repository } diff --git a/re/service.go b/re/service.go index d92233038..a7544e610 100644 --- a/re/service.go +++ b/re/service.go @@ -15,10 +15,7 @@ import ( svcerr "github.com/absmach/magistrala/pkg/errors/service" pkglog "github.com/absmach/magistrala/pkg/logger" "github.com/absmach/magistrala/pkg/messaging" - "github.com/absmach/magistrala/pkg/policies" - "github.com/absmach/magistrala/pkg/roles" "github.com/absmach/magistrala/pkg/ticker" - "github.com/absmach/magistrala/re/operations" ) var ( @@ -36,39 +33,33 @@ type re struct { ticker ticker.Ticker email emailer.Emailer readers grpcReadersV1.ReadersServiceClient - roles.ProvisionManageService } -func NewService(repo Repository, runInfo chan pkglog.RunInfo, policy policies.Service, idp magistrala.IDProvider, rePubSub messaging.PubSub, writersPub, alarmsPub messaging.Publisher, tck ticker.Ticker, emailer emailer.Emailer, readers grpcReadersV1.ReadersServiceClient, availableActions []roles.Action, builtInRoles map[roles.BuiltInRoleName][]roles.Action) (Service, error) { - rpms, err := roles.NewProvisionManageService(operations.EntityType, repo, policy, idp, availableActions, builtInRoles) - if err != nil { - return nil, err - } +func NewService(repo Repository, runInfo chan pkglog.RunInfo, idp magistrala.IDProvider, rePubSub messaging.PubSub, writersPub, alarmsPub messaging.Publisher, tck ticker.Ticker, emailer emailer.Emailer, readers grpcReadersV1.ReadersServiceClient) (Service, error) { return &re{ - repo: repo, - idp: idp, - runInfo: runInfo, - rePubSub: rePubSub, - writersPub: writersPub, - alarmsPub: alarmsPub, - ticker: tck, - email: emailer, - readers: readers, - ProvisionManageService: rpms, + repo: repo, + idp: idp, + runInfo: runInfo, + rePubSub: rePubSub, + writersPub: writersPub, + alarmsPub: alarmsPub, + ticker: tck, + email: emailer, + readers: readers, }, nil } -func (re *re) AddRule(ctx context.Context, session authn.Session, r Rule) (retRule Rule, retRps []roles.RoleProvision, retErr error) { +func (re *re) AddRule(ctx context.Context, session authn.Session, r Rule) (retRule Rule, retErr error) { if r.Logic.Type == GoType && goKeywordRegex.MatchString(r.Logic.Value) { - return Rule{}, nil, errors.Wrap(svcerr.ErrMalformedEntity, ErrGoroutinesNotAllowed) + return Rule{}, errors.Wrap(svcerr.ErrMalformedEntity, ErrGoroutinesNotAllowed) } if r.Logic.Type == GoType && panicRegex.MatchString(r.Logic.Value) { - return Rule{}, nil, errors.Wrap(svcerr.ErrMalformedEntity, ErrPanicNotAllowed) + return Rule{}, errors.Wrap(svcerr.ErrMalformedEntity, ErrPanicNotAllowed) } id, err := re.idp.ID() if err != nil { - return Rule{}, nil, err + return Rule{}, err } now := time.Now().UTC() r.CreatedAt = now @@ -84,7 +75,7 @@ func (re *re) AddRule(ctx context.Context, session authn.Session, r Rule) (retRu rule, err := re.repo.AddRule(ctx, r) if err != nil { - return Rule{}, nil, errors.Wrap(svcerr.ErrCreateEntity, err) + return Rule{}, errors.Wrap(svcerr.ErrCreateEntity, err) } defer func() { @@ -95,37 +86,11 @@ func (re *re) AddRule(ctx context.Context, session authn.Session, r Rule) (retRu } }() - newBuiltInRoleMembers := map[roles.BuiltInRoleName][]roles.Member{ - BuiltInRoleAdmin: {roles.Member(session.UserID)}, - } - - optionalPolicies := []policies.Policy{ - { - SubjectType: policies.DomainType, - Subject: session.DomainID, - Relation: policies.DomainRelation, - ObjectType: operations.EntityType, - Object: rule.ID, - }, - } - - rps, err := re.AddNewEntitiesRoles(ctx, session.DomainID, session.UserID, []string{rule.ID}, optionalPolicies, newBuiltInRoleMembers) - if err != nil { - return Rule{}, nil, errors.Wrap(svcerr.ErrAddPolicies, err) - } - - return rule, rps, nil + return rule, nil } func (re *re) ViewRule(ctx context.Context, session authn.Session, id string, withRoles bool) (Rule, error) { - var rule Rule - var err error - switch withRoles { - case true: - rule, err = re.repo.RetrieveByIDWithRoles(ctx, id, session.UserID) - default: - rule, err = re.repo.ViewRule(ctx, id) - } + rule, err := re.repo.ViewRule(ctx, id) if err != nil { return Rule{}, errors.Wrap(svcerr.ErrViewEntity, err) } @@ -175,14 +140,7 @@ func (re *re) UpdateRuleSchedule(ctx context.Context, session authn.Session, r R func (re *re) ListRules(ctx context.Context, session authn.Session, pm PageMeta) (Page, error) { pm.Domain = session.DomainID - if session.SuperAdmin { - page, err := re.repo.ListAllRules(ctx, pm) - if err != nil { - return Page{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - return page, nil - } - page, err := re.repo.ListUserRules(ctx, session.UserID, pm) + page, err := re.repo.ListAllRules(ctx, pm) if err != nil { return Page{}, errors.Wrap(svcerr.ErrViewEntity, err) } diff --git a/re/service_test.go b/re/service_test.go index a0a067db4..0cf77b607 100644 --- a/re/service_test.go +++ b/re/service_test.go @@ -21,7 +21,6 @@ import ( "github.com/absmach/magistrala/pkg/messaging" pubsubmocks "github.com/absmach/magistrala/pkg/messaging/mocks" policymocks "github.com/absmach/magistrala/pkg/policies/mocks" - "github.com/absmach/magistrala/pkg/roles" pkgSch "github.com/absmach/magistrala/pkg/schedule" tmocks "github.com/absmach/magistrala/pkg/ticker/mocks" "github.com/absmach/magistrala/pkg/uuid" @@ -69,11 +68,7 @@ func newService(t *testing.T, runInfo chan pkglog.RunInfo) (re.Service, *mocks.R readersSvc := new(readmocks.ReadersServiceClient) e := new(emocks.Emailer) policy := new(policymocks.Service) - availableActions := []roles.Action{} - builtInRoles := map[roles.BuiltInRoleName][]roles.Action{ - "admin": availableActions, - } - svc, err := re.NewService(repo, runInfo, policy, idProvider, pubsub, pubsub, pubsub, mockTicker, e, readersSvc, availableActions, builtInRoles) + svc, err := re.NewService(repo, runInfo, idProvider, pubsub, pubsub, pubsub, mockTicker, e, readersSvc) if err != nil { t.Fatalf("Failed to create service: %v", err) } @@ -82,19 +77,15 @@ func newService(t *testing.T, runInfo chan pkglog.RunInfo) (re.Service, *mocks.R func TestAddRule(t *testing.T) { // nolint:dogsled - svc, repo, _, _, _, policies := newService(t, make(chan pkglog.RunInfo)) + svc, repo, _, _, _, _ := newService(t, make(chan pkglog.RunInfo)) ruleName := namegen.Generate() now := time.Now().Add(time.Hour) cases := []struct { - desc string - session authn.Session - rule re.Rule - res re.Rule - err error - addPoliciesErr error - deletePolicies error - addRoleErr error - deleteErr error + desc string + session authn.Session + rule re.Rule + res re.Rule + err error }{ { desc: "Add rule successfully", @@ -124,10 +115,7 @@ func TestAddRule(t *testing.T) { CreatedBy: userID, DomainID: domainID, }, - err: nil, - addPoliciesErr: nil, - addRoleErr: nil, - deleteErr: nil, + err: nil, }, { desc: "Add rule with failed repo", @@ -144,11 +132,7 @@ func TestAddRule(t *testing.T) { Time: now, }, }, - err: repoerr.ErrCreateEntity, - addPoliciesErr: nil, - deletePolicies: nil, - addRoleErr: nil, - deleteErr: nil, + err: repoerr.ErrCreateEntity, }, { desc: "Add rule with non-zero StartDateTime", @@ -180,136 +164,7 @@ func TestAddRule(t *testing.T) { CreatedBy: userID, DomainID: domainID, }, - err: nil, - addPoliciesErr: nil, - addRoleErr: nil, - deleteErr: nil, - }, - { - desc: "Add rule with failed to add roles and failed to delete policies", - session: authn.Session{ - UserID: userID, - DomainID: domainID, - }, - rule: re.Rule{ - Name: ruleName, - InputChannel: inputChannel, - Schedule: pkgSch.Schedule{ - Recurring: pkgSch.Daily, - RecurringPeriod: 1, - Time: now, - }, - }, - res: re.Rule{ - Name: ruleName, - ID: ruleID, - InputChannel: inputChannel, - Schedule: pkgSch.Schedule{ - Recurring: pkgSch.Daily, - RecurringPeriod: 1, - Time: now, - }, - Status: re.EnabledStatus, - CreatedBy: userID, - DomainID: domainID, - }, - addRoleErr: svcerr.ErrCreateEntity, - deletePolicies: svcerr.ErrRemoveEntity, - err: svcerr.ErrRemoveEntity, - }, - { - desc: "Add rule with failed to add policies", - session: authn.Session{ - UserID: userID, - DomainID: domainID, - }, - rule: re.Rule{ - Name: ruleName, - InputChannel: inputChannel, - Schedule: pkgSch.Schedule{ - Recurring: pkgSch.Daily, - RecurringPeriod: 1, - Time: now, - }, - }, - res: re.Rule{ - Name: ruleName, - ID: ruleID, - InputChannel: inputChannel, - Schedule: pkgSch.Schedule{ - Recurring: pkgSch.Daily, - RecurringPeriod: 1, - Time: now, - }, - Status: re.EnabledStatus, - CreatedBy: userID, - DomainID: domainID, - }, - addPoliciesErr: svcerr.ErrAuthorization, - err: svcerr.ErrAddPolicies, - }, - { - desc: "Add rule with failed to add policies and failed rollback", - session: authn.Session{ - UserID: userID, - DomainID: domainID, - }, - rule: re.Rule{ - Name: ruleName, - InputChannel: inputChannel, - Schedule: pkgSch.Schedule{ - Recurring: pkgSch.Daily, - RecurringPeriod: 1, - Time: now, - }, - }, - res: re.Rule{ - Name: ruleName, - ID: ruleID, - InputChannel: inputChannel, - Schedule: pkgSch.Schedule{ - Recurring: pkgSch.Daily, - RecurringPeriod: 1, - Time: now, - }, - Status: re.EnabledStatus, - CreatedBy: userID, - DomainID: domainID, - }, - addPoliciesErr: svcerr.ErrAuthorization, - deleteErr: svcerr.ErrRemoveEntity, - err: svcerr.ErrRollbackRepo, - }, - { - desc: "Add rule with failed to add roles", - session: authn.Session{ - UserID: userID, - DomainID: domainID, - }, - rule: re.Rule{ - Name: ruleName, - InputChannel: inputChannel, - Schedule: pkgSch.Schedule{ - Recurring: pkgSch.Daily, - RecurringPeriod: 1, - Time: now, - }, - }, - res: re.Rule{ - Name: ruleName, - ID: ruleID, - InputChannel: inputChannel, - Schedule: pkgSch.Schedule{ - Recurring: pkgSch.Daily, - RecurringPeriod: 1, - Time: now, - }, - Status: re.EnabledStatus, - CreatedBy: userID, - DomainID: domainID, - }, - addRoleErr: svcerr.ErrCreateEntity, - err: svcerr.ErrAddPolicies, + err: nil, }, { desc: "Add rule with Go script containing goroutines", @@ -330,10 +185,7 @@ func TestAddRule(t *testing.T) { Time: now, }, }, - err: re.ErrGoroutinesNotAllowed, - addPoliciesErr: nil, - addRoleErr: nil, - deleteErr: nil, + err: re.ErrGoroutinesNotAllowed, }, { desc: "Add rule with Go script containing panic", @@ -354,162 +206,66 @@ func TestAddRule(t *testing.T) { Time: now, }, }, - err: re.ErrPanicNotAllowed, - addPoliciesErr: nil, - addRoleErr: nil, - deleteErr: nil, - }, - { - desc: "Add rule with failed to add roles and failed to delete policies", - session: authn.Session{ - UserID: userID, - DomainID: domainID, - }, - rule: re.Rule{ - Name: ruleName, - InputChannel: inputChannel, - Schedule: pkgSch.Schedule{ - Recurring: pkgSch.Daily, - RecurringPeriod: 1, - Time: now, - }, - }, - res: re.Rule{ - Name: ruleName, - ID: ruleID, - InputChannel: inputChannel, - Schedule: pkgSch.Schedule{ - Recurring: pkgSch.Daily, - RecurringPeriod: 1, - Time: now, - }, - Status: re.EnabledStatus, - CreatedBy: userID, - DomainID: domainID, - }, - addRoleErr: svcerr.ErrCreateEntity, - deletePolicies: svcerr.ErrRemoveEntity, - err: svcerr.ErrRemoveEntity, - }, - { - desc: "Add rule with failed to add policies", - session: authn.Session{ - UserID: userID, - DomainID: domainID, - }, - rule: re.Rule{ - Name: ruleName, - InputChannel: inputChannel, - Schedule: pkgSch.Schedule{ - Recurring: pkgSch.Daily, - RecurringPeriod: 1, - Time: now, - }, - }, - res: re.Rule{ - Name: ruleName, - ID: ruleID, - InputChannel: inputChannel, - Schedule: pkgSch.Schedule{ - Recurring: pkgSch.Daily, - RecurringPeriod: 1, - Time: now, - }, - Status: re.EnabledStatus, - CreatedBy: userID, - DomainID: domainID, - }, - addPoliciesErr: svcerr.ErrAuthorization, - err: svcerr.ErrAddPolicies, - }, - { - desc: "Add rule with failed to add policies and failed rollback", - session: authn.Session{ - UserID: userID, - DomainID: domainID, - }, - rule: re.Rule{ - Name: ruleName, - InputChannel: inputChannel, - Schedule: pkgSch.Schedule{ - Recurring: pkgSch.Daily, - RecurringPeriod: 1, - Time: now, - }, - }, - res: re.Rule{ - Name: ruleName, - ID: ruleID, - InputChannel: inputChannel, - Schedule: pkgSch.Schedule{ - Recurring: pkgSch.Daily, - RecurringPeriod: 1, - Time: now, - }, - Status: re.EnabledStatus, - CreatedBy: userID, - DomainID: domainID, - }, - addPoliciesErr: svcerr.ErrAuthorization, - deleteErr: svcerr.ErrRemoveEntity, - err: svcerr.ErrRollbackRepo, - }, - { - desc: "Add rule with failed to add roles", - session: authn.Session{ - UserID: userID, - DomainID: domainID, - }, - rule: re.Rule{ - Name: ruleName, - InputChannel: inputChannel, - Schedule: pkgSch.Schedule{ - Recurring: pkgSch.Daily, - RecurringPeriod: 1, - Time: now, - }, - }, - res: re.Rule{ - Name: ruleName, - ID: ruleID, - InputChannel: inputChannel, - Schedule: pkgSch.Schedule{ - Recurring: pkgSch.Daily, - RecurringPeriod: 1, - Time: now, - }, - Status: re.EnabledStatus, - CreatedBy: userID, - DomainID: domainID, - }, - addRoleErr: svcerr.ErrCreateEntity, - err: svcerr.ErrAddPolicies, + err: re.ErrPanicNotAllowed, }, } for _, tc := range cases { t.Run(tc.desc, func(t *testing.T) { repoCall := repo.On("AddRule", mock.Anything, mock.Anything).Return(tc.res, tc.err) - policyCall := policies.On("AddPolicies", context.Background(), mock.Anything).Return(tc.addPoliciesErr) - policyCall2 := policies.On("DeletePolicies", context.Background(), mock.Anything).Return(tc.deletePolicies) - repoCall1 := repo.On("AddRoles", context.Background(), mock.Anything).Return([]roles.RoleProvision{}, tc.addRoleErr) - repoCall2 := repo.On("Remove", context.Background(), mock.Anything).Return(tc.deleteErr) - res, _, err := svc.AddRule(context.Background(), tc.session, tc.rule) + res, err := svc.AddRule(context.Background(), tc.session, tc.rule) assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) if err == nil { assert.NotEmpty(t, res.ID, "expected non-empty result in ID") assert.Equal(t, tc.rule.Name, res.Name) assert.Equal(t, tc.rule.Schedule, res.Schedule) } - policyCall.Unset() - policyCall2.Unset() repoCall.Unset() - repoCall1.Unset() - repoCall2.Unset() }) } } +func TestAddRuleWithoutRoleProvisioning(t *testing.T) { + repo := new(mocks.Repository) + mockTicker := new(tmocks.Ticker) + idProvider := uuid.NewMock() + pubsub := pubsubmocks.NewPubSub(t) + readersSvc := new(readmocks.ReadersServiceClient) + e := new(emocks.Emailer) + + svc, err := re.NewService(repo, make(chan pkglog.RunInfo), idProvider, pubsub, pubsub, pubsub, mockTicker, e, readersSvc) + if err != nil { + t.Fatalf("Failed to create service: %v", err) + } + + session := authn.Session{ + UserID: userID, + DomainID: domainID, + } + rule := re.Rule{ + Name: ruleName, + InputChannel: inputChannel, + Schedule: pkgSch.Schedule{ + Recurring: pkgSch.Daily, + RecurringPeriod: 1, + Time: time.Now().Add(time.Hour), + }, + } + saved := rule + saved.ID = ruleID + saved.Status = re.EnabledStatus + saved.CreatedBy = userID + saved.DomainID = domainID + + repo.On("AddRule", mock.Anything, mock.Anything).Return(saved, nil).Once() + + res, err := svc.AddRule(context.Background(), session, rule) + assert.NoError(t, err) + assert.Equal(t, saved.ID, res.ID) + repo.AssertNotCalled(t, "AddRoles", mock.Anything, mock.Anything) + repo.AssertExpectations(t) +} + func TestViewRule(t *testing.T) { // nolint:dogsled svc, repo, _, _, _, _ := newService(t, make(chan pkglog.RunInfo)) @@ -982,11 +738,7 @@ func TestListRules(t *testing.T) { for _, tc := range cases { t.Run(tc.desc, func(t *testing.T) { var repoCall *mock.Call - if tc.superAdmin { - repoCall = repo.On("ListAllRules", mock.Anything, mock.Anything).Return(tc.res, tc.err) - } else { - repoCall = repo.On("ListUserRules", mock.Anything, mock.Anything, mock.Anything).Return(tc.res, tc.err) - } + repoCall = repo.On("ListAllRules", mock.Anything, mock.Anything).Return(tc.res, tc.err) res, err := svc.ListRules(context.Background(), tc.session, tc.pageMeta) assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) @@ -1175,7 +927,6 @@ func TestHandle(t *testing.T) { page re.Page listErr error publishErr error - emailErr error expectErr bool }{ { @@ -1421,37 +1172,6 @@ func TestHandle(t *testing.T) { }, listErr: nil, }, - { - desc: "consume message with Lua script and failed Email output", - message: &messaging.Message{ - Channel: inputChannel, - Created: now.Unix(), - Payload: []byte(`{"temperature": 25.5}`), - }, - page: re.Page{ - Rules: []re.Rule{ - { - ID: testsutil.GenerateUUID(t), - Name: namegen.Generate(), - InputChannel: inputChannel, - Status: re.EnabledStatus, - Logic: re.Script{ - Type: re.LuaType, - Value: `return message.payload`, - }, - Outputs: re.Outputs{ - &outputs.Email{ - To: []string{"test@example.com"}, - Subject: "Temperature Alert", - Content: "Temperature: {{.Result}}", - }, - }, - Schedule: schedule, - }, - }, - }, - emailErr: errors.New("failed to send email"), - }, { desc: "consume message with rules using GoType", message: &messaging.Message{ @@ -1849,8 +1569,8 @@ func TestHandle(t *testing.T) { err = tc.listErr } }) - repoCall1 := pubmocks.On("Publish", mock.Anything, mock.Anything, mock.Anything).Return(tc.publishErr) - repoCall2 := emailer.On("SendEmailNotification", mock.Anything, mock.Anything, mock.Anything, mock.Anything, mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(tc.emailErr) + repoCall1 := pubmocks.On("Publish", mock.Anything, mock.Anything, mock.Anything).Return(tc.publishErr).Maybe() + repoCall2 := emailer.On("SendEmailNotification", mock.Anything, mock.Anything, mock.Anything, mock.Anything, mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(nil).Maybe() err = svc.Handle(tc.message) assert.Nil(t, err) diff --git a/readers/api/http/endpoint_test.go b/readers/api/http/endpoint_test.go deleted file mode 100644 index f3688f081..000000000 --- a/readers/api/http/endpoint_test.go +++ /dev/null @@ -1,967 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package http_test - -import ( - "encoding/json" - "fmt" - "net/http" - "net/http/httptest" - "testing" - "time" - - grpcChannelsV1 "github.com/absmach/magistrala/api/grpc/channels/v1" - grpcClientsV1 "github.com/absmach/magistrala/api/grpc/clients/v1" - apiutil "github.com/absmach/magistrala/api/http/util" - chmocks "github.com/absmach/magistrala/channels/mocks" - climocks "github.com/absmach/magistrala/clients/mocks" - "github.com/absmach/magistrala/internal/testsutil" - smqauthn "github.com/absmach/magistrala/pkg/authn" - authnmocks "github.com/absmach/magistrala/pkg/authn/mocks" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/pkg/transformers/senml" - "github.com/absmach/magistrala/readers" - customhttp "github.com/absmach/magistrala/readers/api/http" - "github.com/absmach/magistrala/readers/mocks" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" -) - -const ( - svcName = "test-service" - clientToken = "1" - userToken = "token" - invalidToken = "invalid" - numOfMessages = 100 - valueFields = 5 - subtopic = "topic" - mqttProt = "mqtt" - httpProt = "http" - msgName = "temperature" - instanceID = "5de9b29a-feb9-11ed-be56-0242ac120002" - domainID = "b4d7d79e-fd99-4c2b-ac09-524e43df6888" -) - -var ( - v float64 = 5 - vs = "value" - vb = true - vd = "dataValue" - sum float64 = 42 - validSession = smqauthn.Session{UserID: testsutil.GenerateUUID(&testing.T{})} -) - -func newServer(repo *mocks.MessageRepository, authn *authnmocks.Authentication, clients *climocks.ClientsServiceClient, channels *chmocks.ChannelsServiceClient) *httptest.Server { - mux := customhttp.MakeHandler(repo, authn, clients, channels, svcName, instanceID) - return httptest.NewServer(mux) -} - -type testRequest struct { - client *http.Client - method string - url string - token string - key string -} - -func (tr testRequest) make() (*http.Response, error) { - req, err := http.NewRequest(tr.method, tr.url, http.NoBody) - if err != nil { - return nil, err - } - if tr.token != "" { - req.Header.Set("Authorization", apiutil.BearerPrefix+tr.token) - } - if tr.key != "" { - req.Header.Set("Authorization", apiutil.ClientPrefix+tr.key) - } - - return tr.client.Do(req) -} - -func TestReadAll(t *testing.T) { - chanID := testsutil.GenerateUUID(t) - pubID := testsutil.GenerateUUID(t) - pubID2 := testsutil.GenerateUUID(t) - - now := time.Now().Unix() - - var messages []senml.Message - var queryMsgs []senml.Message - var valueMsgs []senml.Message - var boolMsgs []senml.Message - var stringMsgs []senml.Message - var dataMsgs []senml.Message - - for i := 0; i < numOfMessages; i++ { - // Mix possible values as well as value sum. - msg := senml.Message{ - Channel: chanID, - Publisher: pubID, - Protocol: mqttProt, - Time: float64(now - int64(i)), - Name: "name", - } - - count := i % valueFields - switch count { - case 0: - msg.Value = &v - valueMsgs = append(valueMsgs, msg) - case 1: - msg.BoolValue = &vb - boolMsgs = append(boolMsgs, msg) - case 2: - msg.StringValue = &vs - stringMsgs = append(stringMsgs, msg) - case 3: - msg.DataValue = &vd - dataMsgs = append(dataMsgs, msg) - case 4: - msg.Sum = &sum - msg.Subtopic = subtopic - msg.Protocol = httpProt - msg.Publisher = pubID2 - msg.Name = msgName - queryMsgs = append(queryMsgs, msg) - } - - messages = append(messages, msg) - } - - repo := new(mocks.MessageRepository) - authn := new(authnmocks.Authentication) - clients := new(climocks.ClientsServiceClient) - channels := new(chmocks.ChannelsServiceClient) - ts := newServer(repo, authn, clients, channels) - defer ts.Close() - - cases := []struct { - desc string - req string - url string - token string - key string - status int - res pageRes - authnErr error - authzRes *grpcChannelsV1.AuthzRes - authzErr error - err error - }{ - { - desc: "read page with valid offset and limit", - url: fmt.Sprintf("%s/%s/channels/%s/messages?offset=0&limit=10", ts.URL, domainID, chanID), - token: userToken, - status: http.StatusOK, - res: pageRes{ - PageMetadata: readers.PageMetadata{Limit: 10, Format: "messages", Order: "time", Dir: "desc"}, - Total: uint64(len(messages)), - Messages: messages[0:10], - }, - }, - { - desc: "read page with valid offset and limit as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?offset=0&limit=10", ts.URL, domainID, chanID), - token: userToken, - status: http.StatusOK, - res: pageRes{ - PageMetadata: readers.PageMetadata{Limit: 10, Format: "messages", Order: "time", Dir: "desc"}, - Total: uint64(len(messages)), - Messages: messages[0:10], - }, - }, - { - desc: "read page with negative offset as client", - url: fmt.Sprintf("%s/%s/channels/%s/messages?offset=-1&limit=10", ts.URL, domainID, chanID), - key: clientToken, - status: http.StatusBadRequest, - }, - { - desc: "read page with negative limit as client", - url: fmt.Sprintf("%s/%s/channels/%s/messages?offset=0&limit=-10", ts.URL, domainID, chanID), - key: clientToken, - status: http.StatusBadRequest, - }, - { - desc: "read page with zero limit as client", - url: fmt.Sprintf("%s/%s/channels/%s/messages?offset=0&limit=0", ts.URL, domainID, chanID), - key: clientToken, - status: http.StatusBadRequest, - }, - { - desc: "read page with non-integer offset as client", - url: fmt.Sprintf("%s/%s/channels/%s/messages?offset=abc&limit=10", ts.URL, domainID, chanID), - key: clientToken, - status: http.StatusBadRequest, - }, - { - desc: "read page with non-integer limit as client", - url: fmt.Sprintf("%s/%s/channels/%s/messages?offset=0&limit=abc", ts.URL, domainID, chanID), - key: clientToken, - status: http.StatusBadRequest, - }, - { - desc: "read page with invalid channel id as client", - url: fmt.Sprintf("%s/%s/channels/%s/messages?offset=0&limit=10", ts.URL, domainID, ""), - key: clientToken, - status: http.StatusBadRequest, - }, - { - desc: "read page with multiple offset as client", - url: fmt.Sprintf("%s/%s/channels/%s/messages?offset=0&offset=1&limit=10", ts.URL, domainID, chanID), - key: clientToken, - status: http.StatusBadRequest, - }, - { - desc: "read page with multiple limit as client", - url: fmt.Sprintf("%s/%s/channels/%s/messages?offset=0&limit=20&limit=10", ts.URL, domainID, chanID), - key: clientToken, - status: http.StatusBadRequest, - }, - { - desc: "read page with empty token as client", - url: fmt.Sprintf("%s/%s/channels/%s/messages?offset=0&limit=10", ts.URL, domainID, chanID), - token: "", - authnErr: svcerr.ErrAuthentication, - status: http.StatusUnauthorized, - err: svcerr.ErrAuthentication, - }, - { - desc: "read page with default offset as client", - url: fmt.Sprintf("%s/%s/channels/%s/messages?limit=10", ts.URL, domainID, chanID), - key: clientToken, - status: http.StatusOK, - res: pageRes{ - PageMetadata: readers.PageMetadata{Limit: 10, Format: "messages", Order: "time", Dir: "desc"}, - Total: uint64(len(messages)), - Messages: messages[0:10], - }, - }, - { - desc: "read page with default limit as client", - url: fmt.Sprintf("%s/%s/channels/%s/messages?offset=0", ts.URL, domainID, chanID), - key: clientToken, - status: http.StatusOK, - res: pageRes{ - PageMetadata: readers.PageMetadata{Limit: 10, Format: "messages", Order: "time", Dir: "desc"}, - Total: uint64(len(messages)), - Messages: messages[0:10], - }, - }, - { - desc: "read page with senml format as client", - url: fmt.Sprintf("%s/%s/channels/%s/messages?format=messages", ts.URL, domainID, chanID), - key: clientToken, - status: http.StatusOK, - res: pageRes{ - PageMetadata: readers.PageMetadata{Limit: 10, Format: "messages", Order: "time", Dir: "desc"}, - Total: uint64(len(messages)), - Messages: messages[0:10], - }, - }, - { - desc: "read page with subtopic as client", - url: fmt.Sprintf("%s/%s/channels/%s/messages?subtopic=%s&protocol=%s", ts.URL, domainID, chanID, subtopic, httpProt), - key: clientToken, - status: http.StatusOK, - res: pageRes{ - PageMetadata: readers.PageMetadata{Limit: 10, Subtopic: subtopic, Format: "messages", Protocol: httpProt, Order: "time", Dir: "desc"}, - Total: uint64(len(queryMsgs)), - Messages: queryMsgs[0:10], - }, - }, - { - desc: "read page with subtopic and protocol as client", - url: fmt.Sprintf("%s/%s/channels/%s/messages?subtopic=%s&protocol=%s", ts.URL, domainID, chanID, subtopic, httpProt), - key: clientToken, - status: http.StatusOK, - res: pageRes{ - PageMetadata: readers.PageMetadata{Limit: 10, Subtopic: subtopic, Format: "messages", Protocol: httpProt, Order: "time", Dir: "desc"}, - Total: uint64(len(queryMsgs)), - Messages: queryMsgs[0:10], - }, - }, - { - desc: "read page with publisher as client", - url: fmt.Sprintf("%s/%s/channels/%s/messages?publisher=%s", ts.URL, domainID, chanID, pubID2), - key: clientToken, - status: http.StatusOK, - res: pageRes{ - PageMetadata: readers.PageMetadata{Limit: 10, Format: "messages", Publisher: pubID2, Order: "time", Dir: "desc"}, - Total: uint64(len(queryMsgs)), - Messages: queryMsgs[0:10], - }, - }, - { - desc: "read page with protocol as client", - url: fmt.Sprintf("%s/%s/channels/%s/messages?protocol=http", ts.URL, domainID, chanID), - key: clientToken, - status: http.StatusOK, - res: pageRes{ - PageMetadata: readers.PageMetadata{Limit: 10, Format: "messages", Protocol: httpProt, Order: "time", Dir: "desc"}, - Total: uint64(len(queryMsgs)), - Messages: queryMsgs[0:10], - }, - }, - { - desc: "read page with name as client", - url: fmt.Sprintf("%s/%s/channels/%s/messages?name=%s", ts.URL, domainID, chanID, msgName), - key: clientToken, - status: http.StatusOK, - res: pageRes{ - PageMetadata: readers.PageMetadata{Limit: 10, Format: "messages", Name: msgName, Order: "time", Dir: "desc"}, - Total: uint64(len(queryMsgs)), - Messages: queryMsgs[0:10], - }, - }, - { - desc: "read page with value as client", - url: fmt.Sprintf("%s/%s/channels/%s/messages?v=%f", ts.URL, domainID, chanID, v), - key: clientToken, - status: http.StatusOK, - res: pageRes{ - PageMetadata: readers.PageMetadata{Limit: 10, Format: "messages", Value: v, Order: "time", Dir: "desc"}, - Total: uint64(len(valueMsgs)), - Messages: valueMsgs[0:10], - }, - }, - { - desc: "read page with value and equal comparator as client", - url: fmt.Sprintf("%s/%s/channels/%s/messages?v=%f&comparator=%s", ts.URL, domainID, chanID, v, readers.EqualKey), - key: clientToken, - status: http.StatusOK, - res: pageRes{ - PageMetadata: readers.PageMetadata{Limit: 10, Format: "messages", Value: v, Comparator: readers.EqualKey, Order: "time", Dir: "desc"}, - Total: uint64(len(valueMsgs)), - Messages: valueMsgs[0:10], - }, - }, - { - desc: "read page with value and lower-than comparator as client", - url: fmt.Sprintf("%s/%s/channels/%s/messages?v=%f&comparator=%s", ts.URL, domainID, chanID, v+1, readers.LowerThanKey), - key: clientToken, - status: http.StatusOK, - res: pageRes{ - PageMetadata: readers.PageMetadata{Limit: 10, Format: "messages", Value: v + 1, Comparator: readers.LowerThanKey, Order: "time", Dir: "desc"}, - Total: uint64(len(valueMsgs)), - Messages: valueMsgs[0:10], - }, - }, - { - desc: "read page with value and lower-than-or-equal comparator as client", - url: fmt.Sprintf("%s/%s/channels/%s/messages?v=%f&comparator=%s", ts.URL, domainID, chanID, v+1, readers.LowerThanEqualKey), - key: clientToken, - - status: http.StatusOK, - res: pageRes{ - PageMetadata: readers.PageMetadata{Limit: 10, Format: "messages", Value: v + 1, Comparator: readers.LowerThanEqualKey, Order: "time", Dir: "desc"}, - Total: uint64(len(valueMsgs)), - Messages: valueMsgs[0:10], - }, - }, - { - desc: "read page with value and greater-than comparator as client", - url: fmt.Sprintf("%s/%s/channels/%s/messages?v=%f&comparator=%s", ts.URL, domainID, chanID, v-1, readers.GreaterThanKey), - key: clientToken, - status: http.StatusOK, - res: pageRes{ - PageMetadata: readers.PageMetadata{Limit: 10, Format: "messages", Value: v - 1, Comparator: readers.GreaterThanKey, Order: "time", Dir: "desc"}, - Total: uint64(len(valueMsgs)), - Messages: valueMsgs[0:10], - }, - }, - { - desc: "read page with value and greater-than-or-equal comparator as client", - url: fmt.Sprintf("%s/%s/channels/%s/messages?v=%f&comparator=%s", ts.URL, domainID, chanID, v-1, readers.GreaterThanEqualKey), - key: clientToken, - status: http.StatusOK, - res: pageRes{ - PageMetadata: readers.PageMetadata{Limit: 10, Format: "messages", Value: v - 1, Comparator: readers.GreaterThanEqualKey, Order: "time", Dir: "desc"}, - Total: uint64(len(valueMsgs)), - Messages: valueMsgs[0:10], - }, - }, - { - desc: "read page with non-float value as client", - url: fmt.Sprintf("%s/%s/channels/%s/messages?v=ab01", ts.URL, domainID, chanID), - key: clientToken, - status: http.StatusBadRequest, - }, - { - desc: "read page with value and wrong comparator as client", - url: fmt.Sprintf("%s/%s/channels/%s/messages?v=%f&comparator=wrong", ts.URL, domainID, chanID, v-1), - key: clientToken, - status: http.StatusBadRequest, - }, - { - desc: "read page with boolean value as client", - url: fmt.Sprintf("%s/%s/channels/%s/messages?vb=true", ts.URL, domainID, chanID), - key: clientToken, - status: http.StatusOK, - res: pageRes{ - PageMetadata: readers.PageMetadata{Limit: 10, Format: "messages", BoolValue: true, Order: "time", Dir: "desc"}, - Total: uint64(len(boolMsgs)), - Messages: boolMsgs[0:10], - }, - }, - { - desc: "read page with non-boolean value as client", - url: fmt.Sprintf("%s/%s/channels/%s/messages?vb=yes", ts.URL, domainID, chanID), - key: clientToken, - status: http.StatusBadRequest, - }, - { - desc: "read page with string value as client", - url: fmt.Sprintf("%s/%s/channels/%s/messages?vs=%s", ts.URL, domainID, chanID, vs), - key: clientToken, - status: http.StatusOK, - res: pageRes{ - PageMetadata: readers.PageMetadata{Limit: 10, Format: "messages", StringValue: vs, Order: "time", Dir: "desc"}, - Total: uint64(len(stringMsgs)), - Messages: stringMsgs[0:10], - }, - }, - { - desc: "read page with data value as client", - url: fmt.Sprintf("%s/%s/channels/%s/messages?vd=%s", ts.URL, domainID, chanID, vd), - key: clientToken, - status: http.StatusOK, - res: pageRes{ - PageMetadata: readers.PageMetadata{Limit: 10, Format: "messages", DataValue: vd, Order: "time", Dir: "desc"}, - Total: uint64(len(dataMsgs)), - Messages: dataMsgs[0:10], - }, - }, - { - desc: "read page with non-float from as client", - url: fmt.Sprintf("%s/%s/channels/%s/messages?from=ABCD", ts.URL, domainID, chanID), - key: clientToken, - status: http.StatusBadRequest, - }, - { - desc: "read page with non-float to as client", - url: fmt.Sprintf("%s/%s/channels/%s/messages?to=ABCD", ts.URL, domainID, chanID), - key: clientToken, - status: http.StatusBadRequest, - }, - { - desc: "read page with from/to as client", - url: fmt.Sprintf("%s/%s/channels/%s/messages?from=%f&to=%f", ts.URL, domainID, chanID, messages[19].Time, messages[4].Time), - key: clientToken, - status: http.StatusOK, - res: pageRes{ - PageMetadata: readers.PageMetadata{Limit: 10, Format: "messages", From: messages[19].Time, To: messages[4].Time, Order: "time", Dir: "desc"}, - Total: uint64(len(messages[5:20])), - Messages: messages[5:15], - }, - }, - { - desc: "read page with aggregation as client", - url: fmt.Sprintf("%s/%s/channels/%s/messages?aggregation=MAX", ts.URL, domainID, chanID), - key: clientToken, - status: http.StatusBadRequest, - }, - { - desc: "read page with interval as client", - url: fmt.Sprintf("%s/%s/channels/%s/messages?interval=10h", ts.URL, domainID, chanID), - key: clientToken, - status: http.StatusOK, - res: pageRes{ - PageMetadata: readers.PageMetadata{Limit: 10, Format: "messages", Order: "time", Dir: "desc"}, - Total: uint64(len(messages)), - Messages: messages[0:10], - }, - }, - { - desc: "read page with aggregation and interval as client", - url: fmt.Sprintf("%s/%s/channels/%s/messages?aggregation=MAX&interval=10h", ts.URL, domainID, chanID), - key: clientToken, - status: http.StatusBadRequest, - }, - { - desc: "read page with aggregation, interval, to and from as client", - url: fmt.Sprintf("%s/%s/channels/%s/messages?aggregation=MAX&interval=10h&from=%f&to=%f", ts.URL, domainID, chanID, messages[19].Time, messages[4].Time), - key: clientToken, - status: http.StatusOK, - res: pageRes{ - PageMetadata: readers.PageMetadata{Limit: 10, Format: "messages", Aggregation: "MAX", Interval: "10h", From: messages[19].Time, To: messages[4].Time, Order: "time", Dir: "desc"}, - Total: uint64(len(messages[5:20])), - Messages: messages[5:15], - }, - }, - { - desc: "read page with invalid aggregation and valid interval, to and from as client", - url: fmt.Sprintf("%s/%s/channels/%s/messages?aggregation=invalid&interval=10h&from=%f&to=%f", ts.URL, domainID, chanID, messages[19].Time, messages[4].Time), - key: clientToken, - status: http.StatusBadRequest, - }, - { - desc: "read page with invalid interval and valid aggregation, to and from as client", - url: fmt.Sprintf("%s/%s/channels/%s/messages?aggregation=MAX&interval=10hrs&from=%f&to=%f", ts.URL, domainID, chanID, messages[19].Time, messages[4].Time), - key: clientToken, - status: http.StatusBadRequest, - }, - { - desc: "read page with aggregation, interval and to with missing from as client", - url: fmt.Sprintf("%s/%s/channels/%s/messages?aggregation=MAX&interval=10h&to=%f", ts.URL, domainID, chanID, messages[4].Time), - key: clientToken, - status: http.StatusBadRequest, - }, - { - desc: "read page with aggregation, interval and to with invalid from as client", - url: fmt.Sprintf("%s/%s/channels/%s/messages?aggregation=MAX&interval=10h&to=ABCD&from=%f", ts.URL, domainID, chanID, messages[4].Time), - key: clientToken, - status: http.StatusBadRequest, - }, - { - desc: "read page with aggregation, interval and to with invalid to as client", - url: fmt.Sprintf("%s/%s/channels/%s/messages?aggregation=MAX&interval=10h&from=%f&to=ABCD", ts.URL, domainID, chanID, messages[4].Time), - key: clientToken, - status: http.StatusBadRequest, - }, - { - desc: "read page with valid offset and limit as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?offset=0&limit=10", ts.URL, domainID, chanID), - token: userToken, - status: http.StatusOK, - res: pageRes{ - PageMetadata: readers.PageMetadata{Limit: 10, Format: "messages", Order: "time", Dir: "desc"}, - Total: uint64(len(messages)), - Messages: messages[0:10], - }, - }, - { - desc: "read page with invalid client key", - url: fmt.Sprintf("%s/%s/channels/%s/messages?offset=0&limit=10", ts.URL, domainID, chanID), - key: "invalid", - authnErr: svcerr.ErrAuthentication, - status: http.StatusUnauthorized, - err: svcerr.ErrAuthentication, - }, - { - desc: "read page with unauthorized client key", - url: fmt.Sprintf("%s/%s/channels/%s/messages?offset=0&limit=10", ts.URL, domainID, chanID), - key: clientToken, - authnErr: nil, - authzRes: &grpcChannelsV1.AuthzRes{Authorized: false}, - status: http.StatusForbidden, - err: svcerr.ErrAuthorization, - }, - { - desc: "read page with negative offset as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?offset=-1&limit=10", ts.URL, domainID, chanID), - token: userToken, - status: http.StatusBadRequest, - }, - { - desc: "read page with negative limit as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?offset=0&limit=-10", ts.URL, domainID, chanID), - token: userToken, - status: http.StatusBadRequest, - }, - { - desc: "read page with zero limit as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?offset=0&limit=0", ts.URL, domainID, chanID), - token: userToken, - status: http.StatusBadRequest, - }, - { - desc: "read page with non-integer offset as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?offset=abc&limit=10", ts.URL, domainID, chanID), - token: userToken, - status: http.StatusBadRequest, - }, - { - desc: "read page with non-integer limit as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?offset=0&limit=abc", ts.URL, domainID, chanID), - token: userToken, - status: http.StatusBadRequest, - }, - { - desc: "read page with invalid channel id as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?offset=0&limit=10", ts.URL, domainID, ""), - token: userToken, - status: http.StatusBadRequest, - }, - { - desc: "read page with invalid token as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?offset=0&limit=10", ts.URL, domainID, chanID), - token: invalidToken, - authnErr: svcerr.ErrAuthentication, - status: http.StatusUnauthorized, - err: svcerr.ErrAuthorization, - }, - { - desc: "read page with unauthorized as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?offset=0&limit=10", ts.URL, domainID, chanID), - token: userToken, - authzRes: &grpcChannelsV1.AuthzRes{Authorized: false}, - status: http.StatusForbidden, - err: svcerr.ErrAuthorization, - }, - { - desc: "read page with multiple offset as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?offset=0&offset=1&limit=10", ts.URL, domainID, chanID), - token: userToken, - status: http.StatusBadRequest, - }, - { - desc: "read page with multiple limit as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?offset=0&limit=20&limit=10", ts.URL, domainID, chanID), - token: userToken, - status: http.StatusBadRequest, - }, - { - desc: "read page with empty token as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?offset=0&limit=10", ts.URL, domainID, chanID), - token: "", - authnErr: svcerr.ErrAuthentication, - status: http.StatusUnauthorized, - err: svcerr.ErrAuthorization, - }, - { - desc: "read page with default offset as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?limit=10", ts.URL, domainID, chanID), - token: userToken, - status: http.StatusOK, - res: pageRes{ - PageMetadata: readers.PageMetadata{Limit: 10, Format: "messages", Order: "time", Dir: "desc"}, - Total: uint64(len(messages)), - Messages: messages[0:10], - }, - }, - { - desc: "read page with default limit as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?offset=0", ts.URL, domainID, chanID), - token: userToken, - status: http.StatusOK, - res: pageRes{ - PageMetadata: readers.PageMetadata{Limit: 10, Format: "messages", Order: "time", Dir: "desc"}, - Total: uint64(len(messages)), - Messages: messages[0:10], - }, - }, - { - desc: "read page with senml format as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?format=messages", ts.URL, domainID, chanID), - token: userToken, - status: http.StatusOK, - res: pageRes{ - PageMetadata: readers.PageMetadata{Limit: 10, Format: "messages", Order: "time", Dir: "desc"}, - Total: uint64(len(messages)), - Messages: messages[0:10], - }, - }, - { - desc: "read page with subtopic as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?subtopic=%s&protocol=%s", ts.URL, domainID, chanID, subtopic, httpProt), - token: userToken, - status: http.StatusOK, - res: pageRes{ - PageMetadata: readers.PageMetadata{Limit: 10, Format: "messages", Subtopic: subtopic, Protocol: httpProt, Order: "time", Dir: "desc"}, - Total: uint64(len(queryMsgs)), - Messages: queryMsgs[0:10], - }, - }, - { - desc: "read page with subtopic and protocol as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?subtopic=%s&protocol=%s", ts.URL, domainID, chanID, subtopic, httpProt), - token: userToken, - status: http.StatusOK, - res: pageRes{ - PageMetadata: readers.PageMetadata{Limit: 10, Format: "messages", Subtopic: subtopic, Protocol: httpProt, Order: "time", Dir: "desc"}, - Total: uint64(len(queryMsgs)), - Messages: queryMsgs[0:10], - }, - }, - { - desc: "read page with publisher as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?publisher=%s", ts.URL, domainID, chanID, pubID2), - token: userToken, - status: http.StatusOK, - res: pageRes{ - PageMetadata: readers.PageMetadata{Limit: 10, Format: "messages", Publisher: pubID2, Order: "time", Dir: "desc"}, - Total: uint64(len(queryMsgs)), - Messages: queryMsgs[0:10], - }, - }, - { - desc: "read page with protocol as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?protocol=http", ts.URL, domainID, chanID), - token: userToken, - status: http.StatusOK, - res: pageRes{ - PageMetadata: readers.PageMetadata{Limit: 10, Format: "messages", Protocol: httpProt, Order: "time", Dir: "desc"}, - Total: uint64(len(queryMsgs)), - Messages: queryMsgs[0:10], - }, - }, - { - desc: "read page with name as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?name=%s", ts.URL, domainID, chanID, msgName), - token: userToken, - status: http.StatusOK, - res: pageRes{ - PageMetadata: readers.PageMetadata{Limit: 10, Format: "messages", Name: msgName, Order: "time", Dir: "desc"}, - Total: uint64(len(queryMsgs)), - Messages: queryMsgs[0:10], - }, - }, - { - desc: "read page with value as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?v=%f", ts.URL, domainID, chanID, v), - token: userToken, - status: http.StatusOK, - res: pageRes{ - PageMetadata: readers.PageMetadata{Limit: 10, Format: "messages", Value: v, Order: "time", Dir: "desc"}, - Total: uint64(len(valueMsgs)), - Messages: valueMsgs[0:10], - }, - }, - { - desc: "read page with value and equal comparator as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?v=%f&comparator=%s", ts.URL, domainID, chanID, v, readers.EqualKey), - token: userToken, - status: http.StatusOK, - res: pageRes{ - PageMetadata: readers.PageMetadata{Limit: 10, Format: "messages", Value: v, Comparator: readers.EqualKey, Order: "time", Dir: "desc"}, - Total: uint64(len(valueMsgs)), - Messages: valueMsgs[0:10], - }, - }, - { - desc: "read page with value and lower-than comparator as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?v=%f&comparator=%s", ts.URL, domainID, chanID, v+1, readers.LowerThanKey), - token: userToken, - status: http.StatusOK, - res: pageRes{ - PageMetadata: readers.PageMetadata{Limit: 10, Format: "messages", Value: v + 1, Comparator: readers.LowerThanKey, Order: "time", Dir: "desc"}, - Total: uint64(len(valueMsgs)), - Messages: valueMsgs[0:10], - }, - }, - { - desc: "read page with value and lower-than-or-equal comparator as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?v=%f&comparator=%s", ts.URL, domainID, chanID, v+1, readers.LowerThanEqualKey), - token: userToken, - status: http.StatusOK, - res: pageRes{ - PageMetadata: readers.PageMetadata{Limit: 10, Format: "messages", Value: v + 1, Comparator: readers.LowerThanEqualKey, Order: "time", Dir: "desc"}, - Total: uint64(len(valueMsgs)), - Messages: valueMsgs[0:10], - }, - }, - { - desc: "read page with value and greater-than comparator as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?v=%f&comparator=%s", ts.URL, domainID, chanID, v-1, readers.GreaterThanKey), - token: userToken, - status: http.StatusOK, - res: pageRes{ - PageMetadata: readers.PageMetadata{Limit: 10, Order: "time", Dir: "desc", Format: "messages", Value: v - 1, Comparator: readers.GreaterThanKey}, - Total: uint64(len(valueMsgs)), - Messages: valueMsgs[0:10], - }, - }, - { - desc: "read page with value and greater-than-or-equal comparator as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?v=%f&comparator=%s", ts.URL, domainID, chanID, v-1, readers.GreaterThanEqualKey), - token: userToken, - status: http.StatusOK, - res: pageRes{ - PageMetadata: readers.PageMetadata{Order: "time", Dir: "desc", Limit: 10, Format: "messages", Value: v - 1, Comparator: readers.GreaterThanEqualKey}, - Total: uint64(len(valueMsgs)), - Messages: valueMsgs[0:10], - }, - }, - { - desc: "read page with non-float value as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?v=ab01", ts.URL, domainID, chanID), - token: userToken, - status: http.StatusBadRequest, - }, - { - desc: "read page with value and wrong comparator as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?v=%f&comparator=wrong", ts.URL, domainID, chanID, v-1), - token: userToken, - status: http.StatusBadRequest, - }, - { - desc: "read page with boolean value as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?vb=true", ts.URL, domainID, chanID), - token: userToken, - status: http.StatusOK, - res: pageRes{ - PageMetadata: readers.PageMetadata{Limit: 10, Format: "messages", BoolValue: true, Order: "time", Dir: "desc"}, - Total: uint64(len(boolMsgs)), - Messages: boolMsgs[0:10], - }, - }, - { - desc: "read page with non-boolean value as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?vb=yes", ts.URL, domainID, chanID), - token: userToken, - status: http.StatusBadRequest, - }, - { - desc: "read page with string value as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?vs=%s", ts.URL, domainID, chanID, vs), - token: userToken, - status: http.StatusOK, - res: pageRes{ - PageMetadata: readers.PageMetadata{Limit: 10, Format: "messages", StringValue: vs, Order: "time", Dir: "desc"}, - Total: uint64(len(stringMsgs)), - Messages: stringMsgs[0:10], - }, - }, - { - desc: "read page with data value as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?vd=%s", ts.URL, domainID, chanID, vd), - token: userToken, - status: http.StatusOK, - res: pageRes{ - PageMetadata: readers.PageMetadata{Limit: 10, Format: "messages", DataValue: vd, Order: "time", Dir: "desc"}, - Total: uint64(len(dataMsgs)), - Messages: dataMsgs[0:10], - }, - }, - { - desc: "read page with non-float from as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?from=ABCD", ts.URL, domainID, chanID), - token: userToken, - status: http.StatusBadRequest, - }, - { - desc: "read page with non-float to as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?to=ABCD", ts.URL, domainID, chanID), - token: userToken, - status: http.StatusBadRequest, - }, - { - desc: "read page with from/to as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?from=%f&to=%f", ts.URL, domainID, chanID, messages[19].Time, messages[4].Time), - token: userToken, - status: http.StatusOK, - res: pageRes{ - PageMetadata: readers.PageMetadata{Limit: 10, Format: "messages", From: messages[19].Time, To: messages[4].Time, Order: "time", Dir: "desc"}, - Total: uint64(len(messages[5:20])), - Messages: messages[5:15], - }, - }, - { - desc: "read page with aggregation as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?aggregation=MAX", ts.URL, domainID, chanID), - key: userToken, - status: http.StatusBadRequest, - }, - { - desc: "read page with interval as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?interval=10h", ts.URL, domainID, chanID), - key: userToken, - status: http.StatusOK, - res: pageRes{ - PageMetadata: readers.PageMetadata{Limit: 10, Format: "messages", Order: "time", Dir: "desc"}, - Total: uint64(len(messages)), - Messages: messages[0:10], - }, - }, - { - desc: "read page with aggregation and interval as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?aggregation=MAX&interval=10h", ts.URL, domainID, chanID), - key: userToken, - status: http.StatusBadRequest, - }, - { - desc: "read page with aggregation, interval, to and from as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?aggregation=MAX&interval=10h&from=%f&to=%f", ts.URL, domainID, chanID, messages[19].Time, messages[4].Time), - key: userToken, - status: http.StatusOK, - res: pageRes{ - PageMetadata: readers.PageMetadata{Limit: 10, Format: "messages", Aggregation: "MAX", Interval: "10h", From: messages[19].Time, To: messages[4].Time, Order: "time", Dir: "desc"}, - Total: uint64(len(messages[5:20])), - Messages: messages[5:15], - }, - }, - { - desc: "read page with invalid aggregation and valid interval, to and from as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?aggregation=invalid&interval=10h&from=%f&to=%f", ts.URL, domainID, chanID, messages[19].Time, messages[4].Time), - key: userToken, - status: http.StatusBadRequest, - }, - { - desc: "read page with invalid interval and valid aggregation, to and from as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?aggregation=MAX&interval=10hrs&from=%f&to=%f", ts.URL, domainID, chanID, messages[19].Time, messages[4].Time), - key: userToken, - status: http.StatusBadRequest, - }, - { - desc: "read page with aggregation, interval and to with missing from as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?aggregation=MAX&interval=10h&to=%f", ts.URL, domainID, chanID, messages[4].Time), - key: userToken, - status: http.StatusBadRequest, - }, - { - desc: "read page with aggregation, interval and to with invalid from as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?aggregation=MAX&interval=10h&to=ABCD&from=%f", ts.URL, domainID, chanID, messages[4].Time), - key: userToken, - status: http.StatusBadRequest, - }, - { - desc: "read page with aggregation, interval and to with invalid to as user", - url: fmt.Sprintf("%s/%s/channels/%s/messages?aggregation=MAX&interval=10h&from=%f&to=ABCD", ts.URL, domainID, chanID, messages[4].Time), - key: userToken, - status: http.StatusBadRequest, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - authnCall := authn.On("Authenticate", mock.Anything, tc.token).Return(validSession, tc.authnErr) - if tc.key != "" { - authnCall = clients.On("Authenticate", mock.Anything, &grpcClientsV1.AuthnReq{ - Token: smqauthn.AuthPack(smqauthn.DomainAuth, domainID, tc.key), - }).Return(&grpcClientsV1.AuthnRes{Id: testsutil.GenerateUUID(t), Authenticated: true}, tc.authnErr) - } - if tc.authzRes == nil { - tc.authzRes = &grpcChannelsV1.AuthzRes{Authorized: true} - } - authzCall := channels.On("Authorize", mock.Anything, mock.Anything).Return(tc.authzRes, tc.authzErr) - repoCall := repo.On("ReadAll", chanID, tc.res.PageMetadata).Return(readers.MessagesPage{Total: tc.res.Total, Messages: fromSenml(tc.res.Messages)}, nil) - req := testRequest{ - client: ts.Client(), - method: http.MethodGet, - url: tc.url, - token: tc.token, - key: tc.key, - } - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - - var page pageRes - err = json.NewDecoder(res.Body).Decode(&page) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected %d got %d", tc.desc, tc.status, res.StatusCode)) - assert.Equal(t, tc.res.Total, page.Total, fmt.Sprintf("%s: expected %d got %d", tc.desc, tc.res.Total, page.Total)) - assert.ElementsMatch(t, tc.res.Messages, page.Messages, fmt.Sprintf("%s: got incorrect body from response", tc.desc)) - authzCall.Unset() - authnCall.Unset() - repoCall.Unset() - }) - } -} - -type pageRes struct { - readers.PageMetadata - Total uint64 `json:"total"` - Messages []senml.Message `json:"messages"` -} - -func fromSenml(in []senml.Message) []readers.Message { - var ret []readers.Message - for _, m := range in { - ret = append(ret, m) - } - return ret -} diff --git a/readers/middleware/logging.go b/readers/middleware/logging.go index ffce2d5a3..4e7e8de15 100644 --- a/readers/middleware/logging.go +++ b/readers/middleware/logging.go @@ -19,7 +19,7 @@ type loggingMiddleware struct { svc readers.MessageRepository } -// LoggingMiddleware adds logging facilities to the core service. +// LoggingMiddleware adds logging facilities to the service. func LoggingMiddleware(svc readers.MessageRepository, logger *slog.Logger) readers.MessageRepository { return &loggingMiddleware{ logger: logger, diff --git a/readers/middleware/metrics.go b/readers/middleware/metrics.go index c88525897..86a208609 100644 --- a/readers/middleware/metrics.go +++ b/readers/middleware/metrics.go @@ -20,7 +20,7 @@ type metricsMiddleware struct { svc readers.MessageRepository } -// MetricsMiddleware instruments core service by tracking request count and latency. +// MetricsMiddleware instruments service by tracking request count and latency. func MetricsMiddleware(svc readers.MessageRepository, counter metrics.Counter, latency metrics.Histogram) readers.MessageRepository { return &metricsMiddleware{ counter: counter, diff --git a/reports/README.md b/reports/README.md index a0820c9cc..57a055fd7 100644 --- a/reports/README.md +++ b/reports/README.md @@ -36,20 +36,15 @@ The service is configured using the following environment variables (values show | `MG_REPORTS_DB_SSL_KEY` | PostgreSQL SSL client key | "" | | `MG_REPORTS_DB_SSL_ROOT_CERT` | PostgreSQL SSL root cert | "" | -### Auth and domains gRPC +### Atom | Variable | Description | Default | | --- | --- | --- | -| `MG_AUTH_GRPC_URL` | Auth gRPC endpoint | `auth:7001` | -| `MG_AUTH_GRPC_TIMEOUT` | Auth gRPC timeout | `300s` | -| `MG_AUTH_GRPC_CLIENT_CERT` | Auth gRPC client cert path | `${GRPC_MTLS:+./ssl/certs/auth-grpc-client.crt}` | -| `MG_AUTH_GRPC_CLIENT_KEY` | Auth gRPC client key path | `${GRPC_MTLS:+./ssl/certs/auth-grpc-client.key}` | -| `MG_AUTH_GRPC_SERVER_CA_CERTS` | Auth gRPC server CA path | `${GRPC_MTLS:+./ssl/certs/ca.crt}` | -| `MG_DOMAINS_GRPC_URL` | Domains gRPC endpoint | `domains:7003` | -| `MG_DOMAINS_GRPC_TIMEOUT` | Domains gRPC timeout | `300s` | -| `MG_DOMAINS_GRPC_CLIENT_CERT` | Domains gRPC client cert path | `${GRPC_MTLS:+./ssl/certs/domains-grpc-client.crt}` | -| `MG_DOMAINS_GRPC_CLIENT_KEY` | Domains gRPC client key path | `${GRPC_MTLS:+./ssl/certs/domains-grpc-client.key}` | -| `MG_DOMAINS_GRPC_SERVER_CA_CERTS` | Domains gRPC server CA path | `${GRPC_MTLS:+./ssl/certs/ca.crt}` | +| `ATOM_URL` | Atom HTTP endpoint | `http://atom:8080` | +| `ATOM_JWKS_URL` | Atom JWKS endpoint for JWT verification | `http://atom:8080/.well-known/jwks.json` | +| `ATOM_ADMIN_USERNAME` | Atom admin login for service projections | `atom-admin` | +| `ATOM_ADMIN_SECRET` | Atom admin secret for service projections | `change-me` | +| `ATOM_TIMEOUT` | Atom request timeout | `5s` | | `MG_ALLOW_UNVERIFIED_USER` | Allow unverified users to access | `true` | | `MG_SPICEDB_PRE_SHARED_KEY` | SpiceDB pre-shared key | `12345678` | | `MG_SPICEDB_HOST` | SpiceDB host | `magistrala-spicedb` | diff --git a/reports/api/transport.go b/reports/api/transport.go index 808622911..dffd2bc0a 100644 --- a/reports/api/transport.go +++ b/reports/api/transport.go @@ -16,7 +16,6 @@ import ( apiutil "github.com/absmach/magistrala/api/http/util" smqauthn "github.com/absmach/magistrala/pkg/authn" "github.com/absmach/magistrala/pkg/errors" - roleManagerHttp "github.com/absmach/magistrala/pkg/roles/rolemanager/api" "github.com/absmach/magistrala/reports" "github.com/go-chi/chi/v5" kithttp "github.com/go-kit/kit/transport/http" @@ -39,8 +38,6 @@ func MakeHandler(svc reports.Service, authn smqauthn.AuthNMiddleware, mux *chi.M r.Use(authn.WithOptions(smqauthn.WithDomainCheck(true)).Middleware()) r.Route("/{domainID}", func(r chi.Router) { r.Route("/reports", func(r chi.Router) { - d := roleManagerHttp.NewDecoder("reportID") - r.Post("/", otelhttp.NewHandler(kithttp.NewServer( generateReportEndpoint(svc), decodeGenerateReportRequest, @@ -48,8 +45,6 @@ func MakeHandler(svc reports.Service, authn smqauthn.AuthNMiddleware, mux *chi.M opts..., ), "generate_report").ServeHTTP) - r = roleManagerHttp.EntityAvailableActionsRouter(svc, d, r, opts) - r.Route("/configs", func(r chi.Router) { r.Post("/", otelhttp.NewHandler(kithttp.NewServer( addReportConfigEndpoint(svc), @@ -128,8 +123,6 @@ func MakeHandler(svc reports.Service, authn smqauthn.AuthNMiddleware, mux *chi.M api.EncodeResponse, opts..., ), "delete_report_template").ServeHTTP) - - roleManagerHttp.EntityRoleMangerRouter(svc, d, r, opts) }) }) }) diff --git a/reports/atom.go b/reports/atom.go new file mode 100644 index 000000000..37cefbcea --- /dev/null +++ b/reports/atom.go @@ -0,0 +1,91 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package reports + +import ( + "context" + + "github.com/absmach/magistrala/internal/atom" + "github.com/absmach/magistrala/pkg/authn" +) + +type atomService struct { + Service + projector atom.Projector +} + +func WithAtom(svc Service, projector atom.Projector) Service { + if projector == nil { + return svc + } + return atomService{Service: svc, projector: projector} +} + +func (svc atomService) AddReportConfig(ctx context.Context, session authn.Session, cfg ReportConfig) (ReportConfig, error) { + report, err := svc.Service.AddReportConfig(ctx, session, cfg) + if err != nil { + return report, err + } + if err := svc.projector.UpsertResource(ctx, reportProjection(report)); err != nil { + return report, nil + } + return report, nil +} + +func (svc atomService) UpdateReportConfig(ctx context.Context, session authn.Session, cfg ReportConfig) (ReportConfig, error) { + report, err := svc.Service.UpdateReportConfig(ctx, session, cfg) + return svc.upsertAfterReportChange(ctx, report, err) +} + +func (svc atomService) UpdateReportSchedule(ctx context.Context, session authn.Session, cfg ReportConfig) (ReportConfig, error) { + report, err := svc.Service.UpdateReportSchedule(ctx, session, cfg) + return svc.upsertAfterReportChange(ctx, report, err) +} + +func (svc atomService) EnableReportConfig(ctx context.Context, session authn.Session, id string) (ReportConfig, error) { + report, err := svc.Service.EnableReportConfig(ctx, session, id) + return svc.upsertAfterReportChange(ctx, report, err) +} + +func (svc atomService) DisableReportConfig(ctx context.Context, session authn.Session, id string) (ReportConfig, error) { + report, err := svc.Service.DisableReportConfig(ctx, session, id) + return svc.upsertAfterReportChange(ctx, report, err) +} + +func (svc atomService) RemoveReportConfig(ctx context.Context, session authn.Session, id string) error { + if err := svc.Service.RemoveReportConfig(ctx, session, id); err != nil { + return err + } + _ = svc.projector.DeleteResource(ctx, id) + return nil +} + +func (svc atomService) upsertAfterReportChange(ctx context.Context, report ReportConfig, err error) (ReportConfig, error) { + if err != nil { + return report, err + } + if err := svc.projector.UpsertResource(ctx, reportProjection(report)); err != nil { + return report, nil + } + return report, nil +} + +func reportProjection(r ReportConfig) atom.Resource { + res := atom.ResourceFromFields(atom.ObjectFields{ + ID: r.ID, + Kind: atom.KindReport, + Name: r.Name, + TenantID: r.DomainID, + OwnerID: r.CreatedBy, + Status: r.Status.String(), + Metadata: map[string]any{"description": r.Description}, + CreatedBy: r.CreatedBy, + UpdatedBy: r.UpdatedBy, + CreatedAt: r.CreatedAt, + UpdatedAt: r.UpdatedAt, + Description: r.Description, + }) + res.Attributes["scheduled_at"] = r.Schedule.Time + return res +} diff --git a/reports/builtinroles.go b/reports/builtinroles.go deleted file mode 100644 index 5121ee5e1..000000000 --- a/reports/builtinroles.go +++ /dev/null @@ -1,8 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package reports - -import "github.com/absmach/magistrala/pkg/roles" - -const BuiltInRoleAdmin roles.BuiltInRoleName = "admin" diff --git a/reports/events/streams.go b/reports/events/streams.go index 174264341..706387319 100644 --- a/reports/events/streams.go +++ b/reports/events/streams.go @@ -9,7 +9,6 @@ import ( "github.com/absmach/magistrala/pkg/authn" "github.com/absmach/magistrala/pkg/events" "github.com/absmach/magistrala/pkg/events/store" - rmEvents "github.com/absmach/magistrala/pkg/roles/rolemanager/events" "github.com/absmach/magistrala/reports" "github.com/go-chi/chi/v5/middleware" ) @@ -25,7 +24,6 @@ var _ reports.Service = (*eventStore)(nil) type eventStore struct { events.Publisher svc reports.Service - rmEvents.RoleManagerEventStore } func NewEventStoreMiddleware(ctx context.Context, svc reports.Service, url string) (reports.Service, error) { @@ -34,12 +32,9 @@ func NewEventStoreMiddleware(ctx context.Context, svc reports.Service, url strin return nil, err } - res := rmEvents.NewRoleManagerEventStore("reports", reportPrefix, svc, publisher) - return &eventStore{ - svc: svc, - Publisher: publisher, - RoleManagerEventStore: res, + svc: svc, + Publisher: publisher, }, nil } diff --git a/reports/middleware/authorization.go b/reports/middleware/authorization.go index 409b05187..620a08ef4 100644 --- a/reports/middleware/authorization.go +++ b/reports/middleware/authorization.go @@ -7,13 +7,13 @@ import ( "context" "github.com/absmach/magistrala/auth" + "github.com/absmach/magistrala/internal/atom" "github.com/absmach/magistrala/pkg/authn" smqauthz "github.com/absmach/magistrala/pkg/authz" "github.com/absmach/magistrala/pkg/errors" svcerr "github.com/absmach/magistrala/pkg/errors/service" "github.com/absmach/magistrala/pkg/permissions" "github.com/absmach/magistrala/pkg/policies" - rolemgr "github.com/absmach/magistrala/pkg/roles/rolemanager/middleware" "github.com/absmach/magistrala/reports" "github.com/absmach/magistrala/reports/operations" ) @@ -33,24 +33,30 @@ var ( type authorizationMiddleware struct { svc reports.Service authz smqauthz.Authorization + atomAuthz atom.Authorizer entitiesOps permissions.EntitiesOperations[permissions.Operation] - rolemgr.RoleManagerAuthorizationMiddleware } // AuthorizationMiddleware adds authorization to the reports service. -func AuthorizationMiddleware(svc reports.Service, authz smqauthz.Authorization, entitiesOps permissions.EntitiesOperations[permissions.Operation], roleOps permissions.Operations[permissions.RoleOperation]) (reports.Service, error) { +func AuthorizationMiddleware(svc reports.Service, authz smqauthz.Authorization, entitiesOps permissions.EntitiesOperations[permissions.Operation]) (reports.Service, error) { if err := entitiesOps.Validate(); err != nil { return nil, err } - ram, err := rolemgr.NewAuthorization(operations.EntityType, svc, authz, roleOps) - if err != nil { + return &authorizationMiddleware{ + svc: svc, + authz: authz, + entitiesOps: entitiesOps, + }, nil +} + +func AtomAuthorizationMiddleware(svc reports.Service, authz atom.Authorizer, entitiesOps permissions.EntitiesOperations[permissions.Operation]) (reports.Service, error) { + if err := entitiesOps.Validate(); err != nil { return nil, err } return &authorizationMiddleware{ - svc: svc, - authz: authz, - entitiesOps: entitiesOps, - RoleManagerAuthorizationMiddleware: ram, + svc: svc, + atomAuthz: authz, + entitiesOps: entitiesOps, }, nil } @@ -99,6 +105,9 @@ func (am *authorizationMiddleware) ListReportsConfig(ctx context.Context, sessio case err == nil: session.SuperAdmin = true case errors.Contains(err, svcerr.ErrSuperAdminAction): + if err := am.authorize(ctx, operations.OpListReportsConfig, session, operations.EntityType, auth.AnyIDs); err != nil { + return reports.ReportConfigPage{}, errors.Wrap(errDomainViewConfigs, err) + } default: return reports.ReportConfigPage{}, err } @@ -163,6 +172,9 @@ func (am *authorizationMiddleware) authorize(ctx context.Context, op permissions if err != nil { return err } + if am.atomAuthz != nil { + return atom.Authorize(ctx, am.atomAuthz, session, perm.String(), objType, obj, atom.KindReport) + } pr := smqauthz.PolicyReq{ Domain: session.DomainID, @@ -202,6 +214,9 @@ func (am *authorizationMiddleware) checkSuperAdmin(ctx context.Context, session if session.Role != authn.SuperAdminRole { return svcerr.ErrSuperAdminAction } + if am.atomAuthz != nil { + return atom.Authorize(ctx, am.atomAuthz, session, policies.AdminPermission, policies.PlatformType, policies.MagistralaObject, policies.PlatformType) + } if err := am.authz.Authorize(ctx, smqauthz.PolicyReq{ SubjectType: policies.UserType, Subject: session.UserID, diff --git a/reports/middleware/authorization_test.go b/reports/middleware/authorization_test.go new file mode 100644 index 000000000..74b7e08bd --- /dev/null +++ b/reports/middleware/authorization_test.go @@ -0,0 +1,105 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package middleware + +import ( + "context" + "testing" + + "github.com/absmach/magistrala/auth" + "github.com/absmach/magistrala/internal/atom" + "github.com/absmach/magistrala/pkg/authn" + pkgerrors "github.com/absmach/magistrala/pkg/errors" + "github.com/absmach/magistrala/pkg/permissions" + "github.com/absmach/magistrala/reports" + "github.com/absmach/magistrala/reports/mocks" + "github.com/absmach/magistrala/reports/operations" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/mock" + "github.com/stretchr/testify/require" +) + +type recordingAtomAuthorizer struct { + allowed bool + reqs []atom.AuthzRequest +} + +func (a *recordingAtomAuthorizer) CheckAuthz(_ context.Context, req atom.AuthzRequest) (atom.AuthzResponse, error) { + a.reqs = append(a.reqs, req) + return atom.AuthzResponse{Allowed: a.allowed}, nil +} + +func TestListReportsConfigAuthorizesRegularUser(t *testing.T) { + svc := mocks.NewService(t) + pm := reports.PageMeta{Limit: 10} + session := authn.Session{UserID: "user-1", DomainID: "domain-1", DomainUserID: "domain-1_user-1"} + authz := &recordingAtomAuthorizer{allowed: true} + wrapped, err := AtomAuthorizationMiddleware(svc, authz, testEntitiesOps(t)) + require.NoError(t, err) + + svc.On("ListReportsConfig", mock.Anything, session, pm).Return(reports.ReportConfigPage{PageMeta: reports.PageMeta{Limit: 10}}, nil).Once() + page, err := wrapped.ListReportsConfig(context.Background(), session, pm) + + require.NoError(t, err) + assert.Equal(t, uint64(10), page.Limit) + require.Len(t, authz.reqs, 1) + assert.Equal(t, atom.AuthzRequest{ + SubjectID: "user-1", + Action: "list", + ResourceID: auth.AnyIDs, + ObjectKind: "resource", + ObjectID: auth.AnyIDs, + Context: map[string]any{ + "domain_id": "domain-1", + "legacy_object_type": operations.EntityType, + }, + }, authz.reqs[0]) +} + +func TestListReportsConfigDeniedRegularUserDoesNotDelegate(t *testing.T) { + svc := mocks.NewService(t) + authz := &recordingAtomAuthorizer{allowed: false} + wrapped, err := AtomAuthorizationMiddleware(svc, authz, testEntitiesOps(t)) + require.NoError(t, err) + + _, err = wrapped.ListReportsConfig(context.Background(), authn.Session{UserID: "user-1", DomainID: "domain-1"}, reports.PageMeta{}) + + assert.True(t, pkgerrors.Contains(err, pkgerrors.ErrAuthorization)) + require.Len(t, authz.reqs, 1) +} + +func TestListReportsConfigSuperAdminSkipsListAuthorization(t *testing.T) { + svc := mocks.NewService(t) + pm := reports.PageMeta{Limit: 10} + session := authn.Session{UserID: "admin-1", DomainID: "domain-1", Role: authn.SuperAdminRole} + authz := &recordingAtomAuthorizer{allowed: true} + wrapped, err := AtomAuthorizationMiddleware(svc, authz, testEntitiesOps(t)) + require.NoError(t, err) + + svc.On("ListReportsConfig", mock.Anything, mock.MatchedBy(func(s authn.Session) bool { + return s.SuperAdmin + }), pm).Return(reports.ReportConfigPage{PageMeta: reports.PageMeta{Limit: 10}}, nil).Once() + _, err = wrapped.ListReportsConfig(context.Background(), session, pm) + + require.NoError(t, err) + require.Len(t, authz.reqs, 1) + assert.Equal(t, "manage", authz.reqs[0].Action) +} + +func testEntitiesOps(t *testing.T) permissions.EntitiesOperations[permissions.Operation] { + t.Helper() + details := operations.OperationDetails() + perms := make(map[string]permissions.Permission, len(details)) + for _, detail := range details { + if detail.PermissionRequired { + perms[detail.Name] = permissions.Permission(detail.Name) + } + } + entitiesOps, err := permissions.NewEntitiesOperations( + permissions.EntitiesPermission{operations.EntityType: perms}, + permissions.EntitiesOperationDetails[permissions.Operation]{operations.EntityType: details}, + ) + require.NoError(t, err) + return entitiesOps +} diff --git a/reports/middleware/callout.go b/reports/middleware/callout.go index 2b0ae3e8f..e5ac79940 100644 --- a/reports/middleware/callout.go +++ b/reports/middleware/callout.go @@ -11,8 +11,6 @@ import ( "github.com/absmach/magistrala/pkg/callout" "github.com/absmach/magistrala/pkg/permissions" "github.com/absmach/magistrala/pkg/policies" - mgPolicies "github.com/absmach/magistrala/pkg/policies" - rolemw "github.com/absmach/magistrala/pkg/roles/rolemanager/middleware" "github.com/absmach/magistrala/reports" "github.com/absmach/magistrala/reports/operations" ) @@ -23,26 +21,19 @@ type calloutMiddleware struct { svc reports.Service callout callout.Callout entitiesOps permissions.EntitiesOperations[permissions.Operation] - rolemw.RoleManagerCalloutMiddleware } const entityType = "report" -func NewCallout(svc reports.Service, callout callout.Callout, entitiesOps permissions.EntitiesOperations[permissions.Operation], roleOps permissions.Operations[permissions.RoleOperation]) (reports.Service, error) { - call, err := rolemw.NewCallout(mgPolicies.ReportsType, svc, callout, roleOps) - if err != nil { - return nil, err - } - +func NewCallout(svc reports.Service, callout callout.Callout, entitiesOps permissions.EntitiesOperations[permissions.Operation]) (reports.Service, error) { if err := entitiesOps.Validate(); err != nil { return nil, err } return &calloutMiddleware{ - svc: svc, - callout: callout, - entitiesOps: entitiesOps, - RoleManagerCalloutMiddleware: call, + svc: svc, + callout: callout, + entitiesOps: entitiesOps, }, nil } diff --git a/reports/middleware/logging.go b/reports/middleware/logging.go index 94a1bf74a..389f7463a 100644 --- a/reports/middleware/logging.go +++ b/reports/middleware/logging.go @@ -9,7 +9,6 @@ import ( "time" "github.com/absmach/magistrala/pkg/authn" - rolemw "github.com/absmach/magistrala/pkg/roles/rolemanager/middleware" "github.com/absmach/magistrala/reports" ) @@ -18,14 +17,12 @@ var _ reports.Service = (*loggingMiddleware)(nil) type loggingMiddleware struct { logger *slog.Logger svc reports.Service - rolemw.RoleManagerLoggingMiddleware } func LoggingMiddleware(svc reports.Service, logger *slog.Logger) reports.Service { return &loggingMiddleware{ - logger: logger, - svc: svc, - RoleManagerLoggingMiddleware: rolemw.NewLogging("reports", svc, logger), + logger: logger, + svc: svc, } } diff --git a/reports/middleware/metrics.go b/reports/middleware/metrics.go index 68a2d6e2a..9e5fc23de 100644 --- a/reports/middleware/metrics.go +++ b/reports/middleware/metrics.go @@ -8,7 +8,6 @@ import ( "time" "github.com/absmach/magistrala/pkg/authn" - rolemw "github.com/absmach/magistrala/pkg/roles/rolemanager/middleware" "github.com/absmach/magistrala/reports" "github.com/go-kit/kit/metrics" ) @@ -17,17 +16,15 @@ type metricsMiddleware struct { counter metrics.Counter latency metrics.Histogram service reports.Service - rolemw.RoleManagerMetricsMiddleware } var _ reports.Service = (*metricsMiddleware)(nil) func NewMetricsMiddleware(counter metrics.Counter, latency metrics.Histogram, service reports.Service) reports.Service { return &metricsMiddleware{ - counter: counter, - latency: latency, - service: service, - RoleManagerMetricsMiddleware: rolemw.NewMetrics("reports", service, counter, latency), + counter: counter, + latency: latency, + service: service, } } diff --git a/reports/middleware/tracing.go b/reports/middleware/tracing.go index c2bd714c8..c289a5ee1 100644 --- a/reports/middleware/tracing.go +++ b/reports/middleware/tracing.go @@ -7,7 +7,6 @@ import ( "context" "github.com/absmach/magistrala/pkg/authn" - rolemw "github.com/absmach/magistrala/pkg/roles/rolemanager/middleware" smqTracing "github.com/absmach/magistrala/pkg/tracing" "github.com/absmach/magistrala/reports" "go.opentelemetry.io/otel/attribute" @@ -17,16 +16,14 @@ import ( type tracingMiddleware struct { tracer trace.Tracer svc reports.Service - rolemw.RoleManagerTracing } var _ reports.Service = (*tracingMiddleware)(nil) func NewTracingMiddleware(tracer trace.Tracer, svc reports.Service) reports.Service { return &tracingMiddleware{ - tracer: tracer, - svc: svc, - RoleManagerTracing: rolemw.NewTracing("reports", svc, tracer), + tracer: tracer, + svc: svc, } } diff --git a/reports/mocks/repository.go b/reports/mocks/repository.go index 61e6df1a0..4cd1e142d 100644 --- a/reports/mocks/repository.go +++ b/reports/mocks/repository.go @@ -12,7 +12,6 @@ import ( "context" "time" - "github.com/absmach/magistrala/pkg/roles" "github.com/absmach/magistrala/reports" mock "github.com/stretchr/testify/mock" ) @@ -110,74 +109,6 @@ func (_c *Repository_AddReportConfig_Call) RunAndReturn(run func(ctx context.Con return _c } -// AddRoles provides a mock function for the type Repository -func (_mock *Repository) AddRoles(ctx context.Context, rps []roles.RoleProvision) ([]roles.RoleProvision, error) { - ret := _mock.Called(ctx, rps) - - if len(ret) == 0 { - panic("no return value specified for AddRoles") - } - - var r0 []roles.RoleProvision - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, []roles.RoleProvision) ([]roles.RoleProvision, error)); ok { - return returnFunc(ctx, rps) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, []roles.RoleProvision) []roles.RoleProvision); ok { - r0 = returnFunc(ctx, rps) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]roles.RoleProvision) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, []roles.RoleProvision) error); ok { - r1 = returnFunc(ctx, rps) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_AddRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'AddRoles' -type Repository_AddRoles_Call struct { - *mock.Call -} - -// AddRoles is a helper method to define mock.On call -// - ctx context.Context -// - rps []roles.RoleProvision -func (_e *Repository_Expecter) AddRoles(ctx interface{}, rps interface{}) *Repository_AddRoles_Call { - return &Repository_AddRoles_Call{Call: _e.mock.On("AddRoles", ctx, rps)} -} - -func (_c *Repository_AddRoles_Call) Run(run func(ctx context.Context, rps []roles.RoleProvision)) *Repository_AddRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 []roles.RoleProvision - if args[1] != nil { - arg1 = args[1].([]roles.RoleProvision) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_AddRoles_Call) Return(roleProvisions []roles.RoleProvision, err error) *Repository_AddRoles_Call { - _c.Call.Return(roleProvisions, err) - return _c -} - -func (_c *Repository_AddRoles_Call) RunAndReturn(run func(ctx context.Context, rps []roles.RoleProvision) ([]roles.RoleProvision, error)) *Repository_AddRoles_Call { - _c.Call.Return(run) - return _c -} - // DeleteReportTemplate provides a mock function for the type Repository func (_mock *Repository) DeleteReportTemplate(ctx context.Context, domainID string, reportID string) error { ret := _mock.Called(ctx, domainID, reportID) @@ -307,270 +238,6 @@ func (_c *Repository_ListAllReportsConfig_Call) RunAndReturn(run func(ctx contex return _c } -// ListEntityMembers provides a mock function for the type Repository -func (_mock *Repository) ListEntityMembers(ctx context.Context, entityID string, pageQuery roles.MembersRolePageQuery) (roles.MembersRolePage, error) { - ret := _mock.Called(ctx, entityID, pageQuery) - - if len(ret) == 0 { - panic("no return value specified for ListEntityMembers") - } - - var r0 roles.MembersRolePage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, roles.MembersRolePageQuery) (roles.MembersRolePage, error)); ok { - return returnFunc(ctx, entityID, pageQuery) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, roles.MembersRolePageQuery) roles.MembersRolePage); ok { - r0 = returnFunc(ctx, entityID, pageQuery) - } else { - r0 = ret.Get(0).(roles.MembersRolePage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, roles.MembersRolePageQuery) error); ok { - r1 = returnFunc(ctx, entityID, pageQuery) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_ListEntityMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ListEntityMembers' -type Repository_ListEntityMembers_Call struct { - *mock.Call -} - -// ListEntityMembers is a helper method to define mock.On call -// - ctx context.Context -// - entityID string -// - pageQuery roles.MembersRolePageQuery -func (_e *Repository_Expecter) ListEntityMembers(ctx interface{}, entityID interface{}, pageQuery interface{}) *Repository_ListEntityMembers_Call { - return &Repository_ListEntityMembers_Call{Call: _e.mock.On("ListEntityMembers", ctx, entityID, pageQuery)} -} - -func (_c *Repository_ListEntityMembers_Call) Run(run func(ctx context.Context, entityID string, pageQuery roles.MembersRolePageQuery)) *Repository_ListEntityMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 roles.MembersRolePageQuery - if args[2] != nil { - arg2 = args[2].(roles.MembersRolePageQuery) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_ListEntityMembers_Call) Return(membersRolePage roles.MembersRolePage, err error) *Repository_ListEntityMembers_Call { - _c.Call.Return(membersRolePage, err) - return _c -} - -func (_c *Repository_ListEntityMembers_Call) RunAndReturn(run func(ctx context.Context, entityID string, pageQuery roles.MembersRolePageQuery) (roles.MembersRolePage, error)) *Repository_ListEntityMembers_Call { - _c.Call.Return(run) - return _c -} - -// ListUserReportsConfig provides a mock function for the type Repository -func (_mock *Repository) ListUserReportsConfig(ctx context.Context, userID string, pm reports.PageMeta) (reports.ReportConfigPage, error) { - ret := _mock.Called(ctx, userID, pm) - - if len(ret) == 0 { - panic("no return value specified for ListUserReportsConfig") - } - - var r0 reports.ReportConfigPage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, reports.PageMeta) (reports.ReportConfigPage, error)); ok { - return returnFunc(ctx, userID, pm) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, reports.PageMeta) reports.ReportConfigPage); ok { - r0 = returnFunc(ctx, userID, pm) - } else { - r0 = ret.Get(0).(reports.ReportConfigPage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, reports.PageMeta) error); ok { - r1 = returnFunc(ctx, userID, pm) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_ListUserReportsConfig_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ListUserReportsConfig' -type Repository_ListUserReportsConfig_Call struct { - *mock.Call -} - -// ListUserReportsConfig is a helper method to define mock.On call -// - ctx context.Context -// - userID string -// - pm reports.PageMeta -func (_e *Repository_Expecter) ListUserReportsConfig(ctx interface{}, userID interface{}, pm interface{}) *Repository_ListUserReportsConfig_Call { - return &Repository_ListUserReportsConfig_Call{Call: _e.mock.On("ListUserReportsConfig", ctx, userID, pm)} -} - -func (_c *Repository_ListUserReportsConfig_Call) Run(run func(ctx context.Context, userID string, pm reports.PageMeta)) *Repository_ListUserReportsConfig_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 reports.PageMeta - if args[2] != nil { - arg2 = args[2].(reports.PageMeta) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_ListUserReportsConfig_Call) Return(reportConfigPage reports.ReportConfigPage, err error) *Repository_ListUserReportsConfig_Call { - _c.Call.Return(reportConfigPage, err) - return _c -} - -func (_c *Repository_ListUserReportsConfig_Call) RunAndReturn(run func(ctx context.Context, userID string, pm reports.PageMeta) (reports.ReportConfigPage, error)) *Repository_ListUserReportsConfig_Call { - _c.Call.Return(run) - return _c -} - -// RemoveEntityMembers provides a mock function for the type Repository -func (_mock *Repository) RemoveEntityMembers(ctx context.Context, entityID string, members []string) error { - ret := _mock.Called(ctx, entityID, members) - - if len(ret) == 0 { - panic("no return value specified for RemoveEntityMembers") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, []string) error); ok { - r0 = returnFunc(ctx, entityID, members) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RemoveEntityMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveEntityMembers' -type Repository_RemoveEntityMembers_Call struct { - *mock.Call -} - -// RemoveEntityMembers is a helper method to define mock.On call -// - ctx context.Context -// - entityID string -// - members []string -func (_e *Repository_Expecter) RemoveEntityMembers(ctx interface{}, entityID interface{}, members interface{}) *Repository_RemoveEntityMembers_Call { - return &Repository_RemoveEntityMembers_Call{Call: _e.mock.On("RemoveEntityMembers", ctx, entityID, members)} -} - -func (_c *Repository_RemoveEntityMembers_Call) Run(run func(ctx context.Context, entityID string, members []string)) *Repository_RemoveEntityMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RemoveEntityMembers_Call) Return(err error) *Repository_RemoveEntityMembers_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RemoveEntityMembers_Call) RunAndReturn(run func(ctx context.Context, entityID string, members []string) error) *Repository_RemoveEntityMembers_Call { - _c.Call.Return(run) - return _c -} - -// RemoveMemberFromAllRoles provides a mock function for the type Repository -func (_mock *Repository) RemoveMemberFromAllRoles(ctx context.Context, memberID string) error { - ret := _mock.Called(ctx, memberID) - - if len(ret) == 0 { - panic("no return value specified for RemoveMemberFromAllRoles") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) error); ok { - r0 = returnFunc(ctx, memberID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RemoveMemberFromAllRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveMemberFromAllRoles' -type Repository_RemoveMemberFromAllRoles_Call struct { - *mock.Call -} - -// RemoveMemberFromAllRoles is a helper method to define mock.On call -// - ctx context.Context -// - memberID string -func (_e *Repository_Expecter) RemoveMemberFromAllRoles(ctx interface{}, memberID interface{}) *Repository_RemoveMemberFromAllRoles_Call { - return &Repository_RemoveMemberFromAllRoles_Call{Call: _e.mock.On("RemoveMemberFromAllRoles", ctx, memberID)} -} - -func (_c *Repository_RemoveMemberFromAllRoles_Call) Run(run func(ctx context.Context, memberID string)) *Repository_RemoveMemberFromAllRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RemoveMemberFromAllRoles_Call) Return(err error) *Repository_RemoveMemberFromAllRoles_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RemoveMemberFromAllRoles_Call) RunAndReturn(run func(ctx context.Context, memberID string) error) *Repository_RemoveMemberFromAllRoles_Call { - _c.Call.Return(run) - return _c -} - // RemoveReportConfig provides a mock function for the type Repository func (_mock *Repository) RemoveReportConfig(ctx context.Context, id string) error { ret := _mock.Called(ctx, id) @@ -628,1105 +295,6 @@ func (_c *Repository_RemoveReportConfig_Call) RunAndReturn(run func(ctx context. return _c } -// RemoveRoles provides a mock function for the type Repository -func (_mock *Repository) RemoveRoles(ctx context.Context, roleIDs []string) error { - ret := _mock.Called(ctx, roleIDs) - - if len(ret) == 0 { - panic("no return value specified for RemoveRoles") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, []string) error); ok { - r0 = returnFunc(ctx, roleIDs) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RemoveRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveRoles' -type Repository_RemoveRoles_Call struct { - *mock.Call -} - -// RemoveRoles is a helper method to define mock.On call -// - ctx context.Context -// - roleIDs []string -func (_e *Repository_Expecter) RemoveRoles(ctx interface{}, roleIDs interface{}) *Repository_RemoveRoles_Call { - return &Repository_RemoveRoles_Call{Call: _e.mock.On("RemoveRoles", ctx, roleIDs)} -} - -func (_c *Repository_RemoveRoles_Call) Run(run func(ctx context.Context, roleIDs []string)) *Repository_RemoveRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 []string - if args[1] != nil { - arg1 = args[1].([]string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RemoveRoles_Call) Return(err error) *Repository_RemoveRoles_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RemoveRoles_Call) RunAndReturn(run func(ctx context.Context, roleIDs []string) error) *Repository_RemoveRoles_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveAllRoles provides a mock function for the type Repository -func (_mock *Repository) RetrieveAllRoles(ctx context.Context, entityID string, limit uint64, offset uint64) (roles.RolePage, error) { - ret := _mock.Called(ctx, entityID, limit, offset) - - if len(ret) == 0 { - panic("no return value specified for RetrieveAllRoles") - } - - var r0 roles.RolePage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, uint64, uint64) (roles.RolePage, error)); ok { - return returnFunc(ctx, entityID, limit, offset) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, uint64, uint64) roles.RolePage); ok { - r0 = returnFunc(ctx, entityID, limit, offset) - } else { - r0 = ret.Get(0).(roles.RolePage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, uint64, uint64) error); ok { - r1 = returnFunc(ctx, entityID, limit, offset) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveAllRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveAllRoles' -type Repository_RetrieveAllRoles_Call struct { - *mock.Call -} - -// RetrieveAllRoles is a helper method to define mock.On call -// - ctx context.Context -// - entityID string -// - limit uint64 -// - offset uint64 -func (_e *Repository_Expecter) RetrieveAllRoles(ctx interface{}, entityID interface{}, limit interface{}, offset interface{}) *Repository_RetrieveAllRoles_Call { - return &Repository_RetrieveAllRoles_Call{Call: _e.mock.On("RetrieveAllRoles", ctx, entityID, limit, offset)} -} - -func (_c *Repository_RetrieveAllRoles_Call) Run(run func(ctx context.Context, entityID string, limit uint64, offset uint64)) *Repository_RetrieveAllRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 uint64 - if args[2] != nil { - arg2 = args[2].(uint64) - } - var arg3 uint64 - if args[3] != nil { - arg3 = args[3].(uint64) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Repository_RetrieveAllRoles_Call) Return(rolePage roles.RolePage, err error) *Repository_RetrieveAllRoles_Call { - _c.Call.Return(rolePage, err) - return _c -} - -func (_c *Repository_RetrieveAllRoles_Call) RunAndReturn(run func(ctx context.Context, entityID string, limit uint64, offset uint64) (roles.RolePage, error)) *Repository_RetrieveAllRoles_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveByIDWithRoles provides a mock function for the type Repository -func (_mock *Repository) RetrieveByIDWithRoles(ctx context.Context, id string, memberID string) (reports.ReportConfig, error) { - ret := _mock.Called(ctx, id, memberID) - - if len(ret) == 0 { - panic("no return value specified for RetrieveByIDWithRoles") - } - - var r0 reports.ReportConfig - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) (reports.ReportConfig, error)); ok { - return returnFunc(ctx, id, memberID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) reports.ReportConfig); ok { - r0 = returnFunc(ctx, id, memberID) - } else { - r0 = ret.Get(0).(reports.ReportConfig) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, string) error); ok { - r1 = returnFunc(ctx, id, memberID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveByIDWithRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveByIDWithRoles' -type Repository_RetrieveByIDWithRoles_Call struct { - *mock.Call -} - -// RetrieveByIDWithRoles is a helper method to define mock.On call -// - ctx context.Context -// - id string -// - memberID string -func (_e *Repository_Expecter) RetrieveByIDWithRoles(ctx interface{}, id interface{}, memberID interface{}) *Repository_RetrieveByIDWithRoles_Call { - return &Repository_RetrieveByIDWithRoles_Call{Call: _e.mock.On("RetrieveByIDWithRoles", ctx, id, memberID)} -} - -func (_c *Repository_RetrieveByIDWithRoles_Call) Run(run func(ctx context.Context, id string, memberID string)) *Repository_RetrieveByIDWithRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RetrieveByIDWithRoles_Call) Return(reportConfig reports.ReportConfig, err error) *Repository_RetrieveByIDWithRoles_Call { - _c.Call.Return(reportConfig, err) - return _c -} - -func (_c *Repository_RetrieveByIDWithRoles_Call) RunAndReturn(run func(ctx context.Context, id string, memberID string) (reports.ReportConfig, error)) *Repository_RetrieveByIDWithRoles_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveEntitiesRolesActionsMembers provides a mock function for the type Repository -func (_mock *Repository) RetrieveEntitiesRolesActionsMembers(ctx context.Context, entityIDs []string) ([]roles.EntityActionRole, []roles.EntityMemberRole, error) { - ret := _mock.Called(ctx, entityIDs) - - if len(ret) == 0 { - panic("no return value specified for RetrieveEntitiesRolesActionsMembers") - } - - var r0 []roles.EntityActionRole - var r1 []roles.EntityMemberRole - var r2 error - if returnFunc, ok := ret.Get(0).(func(context.Context, []string) ([]roles.EntityActionRole, []roles.EntityMemberRole, error)); ok { - return returnFunc(ctx, entityIDs) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, []string) []roles.EntityActionRole); ok { - r0 = returnFunc(ctx, entityIDs) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]roles.EntityActionRole) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, []string) []roles.EntityMemberRole); ok { - r1 = returnFunc(ctx, entityIDs) - } else { - if ret.Get(1) != nil { - r1 = ret.Get(1).([]roles.EntityMemberRole) - } - } - if returnFunc, ok := ret.Get(2).(func(context.Context, []string) error); ok { - r2 = returnFunc(ctx, entityIDs) - } else { - r2 = ret.Error(2) - } - return r0, r1, r2 -} - -// Repository_RetrieveEntitiesRolesActionsMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveEntitiesRolesActionsMembers' -type Repository_RetrieveEntitiesRolesActionsMembers_Call struct { - *mock.Call -} - -// RetrieveEntitiesRolesActionsMembers is a helper method to define mock.On call -// - ctx context.Context -// - entityIDs []string -func (_e *Repository_Expecter) RetrieveEntitiesRolesActionsMembers(ctx interface{}, entityIDs interface{}) *Repository_RetrieveEntitiesRolesActionsMembers_Call { - return &Repository_RetrieveEntitiesRolesActionsMembers_Call{Call: _e.mock.On("RetrieveEntitiesRolesActionsMembers", ctx, entityIDs)} -} - -func (_c *Repository_RetrieveEntitiesRolesActionsMembers_Call) Run(run func(ctx context.Context, entityIDs []string)) *Repository_RetrieveEntitiesRolesActionsMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 []string - if args[1] != nil { - arg1 = args[1].([]string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RetrieveEntitiesRolesActionsMembers_Call) Return(entityActionRoles []roles.EntityActionRole, entityMemberRoles []roles.EntityMemberRole, err error) *Repository_RetrieveEntitiesRolesActionsMembers_Call { - _c.Call.Return(entityActionRoles, entityMemberRoles, err) - return _c -} - -func (_c *Repository_RetrieveEntitiesRolesActionsMembers_Call) RunAndReturn(run func(ctx context.Context, entityIDs []string) ([]roles.EntityActionRole, []roles.EntityMemberRole, error)) *Repository_RetrieveEntitiesRolesActionsMembers_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveEntityRole provides a mock function for the type Repository -func (_mock *Repository) RetrieveEntityRole(ctx context.Context, entityID string, roleID string) (roles.Role, error) { - ret := _mock.Called(ctx, entityID, roleID) - - if len(ret) == 0 { - panic("no return value specified for RetrieveEntityRole") - } - - var r0 roles.Role - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) (roles.Role, error)); ok { - return returnFunc(ctx, entityID, roleID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) roles.Role); ok { - r0 = returnFunc(ctx, entityID, roleID) - } else { - r0 = ret.Get(0).(roles.Role) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, string) error); ok { - r1 = returnFunc(ctx, entityID, roleID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveEntityRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveEntityRole' -type Repository_RetrieveEntityRole_Call struct { - *mock.Call -} - -// RetrieveEntityRole is a helper method to define mock.On call -// - ctx context.Context -// - entityID string -// - roleID string -func (_e *Repository_Expecter) RetrieveEntityRole(ctx interface{}, entityID interface{}, roleID interface{}) *Repository_RetrieveEntityRole_Call { - return &Repository_RetrieveEntityRole_Call{Call: _e.mock.On("RetrieveEntityRole", ctx, entityID, roleID)} -} - -func (_c *Repository_RetrieveEntityRole_Call) Run(run func(ctx context.Context, entityID string, roleID string)) *Repository_RetrieveEntityRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RetrieveEntityRole_Call) Return(role roles.Role, err error) *Repository_RetrieveEntityRole_Call { - _c.Call.Return(role, err) - return _c -} - -func (_c *Repository_RetrieveEntityRole_Call) RunAndReturn(run func(ctx context.Context, entityID string, roleID string) (roles.Role, error)) *Repository_RetrieveEntityRole_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveRole provides a mock function for the type Repository -func (_mock *Repository) RetrieveRole(ctx context.Context, roleID string) (roles.Role, error) { - ret := _mock.Called(ctx, roleID) - - if len(ret) == 0 { - panic("no return value specified for RetrieveRole") - } - - var r0 roles.Role - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) (roles.Role, error)); ok { - return returnFunc(ctx, roleID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string) roles.Role); ok { - r0 = returnFunc(ctx, roleID) - } else { - r0 = ret.Get(0).(roles.Role) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = returnFunc(ctx, roleID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveRole' -type Repository_RetrieveRole_Call struct { - *mock.Call -} - -// RetrieveRole is a helper method to define mock.On call -// - ctx context.Context -// - roleID string -func (_e *Repository_Expecter) RetrieveRole(ctx interface{}, roleID interface{}) *Repository_RetrieveRole_Call { - return &Repository_RetrieveRole_Call{Call: _e.mock.On("RetrieveRole", ctx, roleID)} -} - -func (_c *Repository_RetrieveRole_Call) Run(run func(ctx context.Context, roleID string)) *Repository_RetrieveRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RetrieveRole_Call) Return(role roles.Role, err error) *Repository_RetrieveRole_Call { - _c.Call.Return(role, err) - return _c -} - -func (_c *Repository_RetrieveRole_Call) RunAndReturn(run func(ctx context.Context, roleID string) (roles.Role, error)) *Repository_RetrieveRole_Call { - _c.Call.Return(run) - return _c -} - -// RoleAddActions provides a mock function for the type Repository -func (_mock *Repository) RoleAddActions(ctx context.Context, role roles.Role, actions []string) ([]string, error) { - ret := _mock.Called(ctx, role, actions) - - if len(ret) == 0 { - panic("no return value specified for RoleAddActions") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role, []string) ([]string, error)); ok { - return returnFunc(ctx, role, actions) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role, []string) []string); ok { - r0 = returnFunc(ctx, role, actions) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, roles.Role, []string) error); ok { - r1 = returnFunc(ctx, role, actions) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RoleAddActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleAddActions' -type Repository_RoleAddActions_Call struct { - *mock.Call -} - -// RoleAddActions is a helper method to define mock.On call -// - ctx context.Context -// - role roles.Role -// - actions []string -func (_e *Repository_Expecter) RoleAddActions(ctx interface{}, role interface{}, actions interface{}) *Repository_RoleAddActions_Call { - return &Repository_RoleAddActions_Call{Call: _e.mock.On("RoleAddActions", ctx, role, actions)} -} - -func (_c *Repository_RoleAddActions_Call) Run(run func(ctx context.Context, role roles.Role, actions []string)) *Repository_RoleAddActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RoleAddActions_Call) Return(ops []string, err error) *Repository_RoleAddActions_Call { - _c.Call.Return(ops, err) - return _c -} - -func (_c *Repository_RoleAddActions_Call) RunAndReturn(run func(ctx context.Context, role roles.Role, actions []string) ([]string, error)) *Repository_RoleAddActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleAddMembers provides a mock function for the type Repository -func (_mock *Repository) RoleAddMembers(ctx context.Context, role roles.Role, members []string) ([]string, error) { - ret := _mock.Called(ctx, role, members) - - if len(ret) == 0 { - panic("no return value specified for RoleAddMembers") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role, []string) ([]string, error)); ok { - return returnFunc(ctx, role, members) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role, []string) []string); ok { - r0 = returnFunc(ctx, role, members) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, roles.Role, []string) error); ok { - r1 = returnFunc(ctx, role, members) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RoleAddMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleAddMembers' -type Repository_RoleAddMembers_Call struct { - *mock.Call -} - -// RoleAddMembers is a helper method to define mock.On call -// - ctx context.Context -// - role roles.Role -// - members []string -func (_e *Repository_Expecter) RoleAddMembers(ctx interface{}, role interface{}, members interface{}) *Repository_RoleAddMembers_Call { - return &Repository_RoleAddMembers_Call{Call: _e.mock.On("RoleAddMembers", ctx, role, members)} -} - -func (_c *Repository_RoleAddMembers_Call) Run(run func(ctx context.Context, role roles.Role, members []string)) *Repository_RoleAddMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RoleAddMembers_Call) Return(strings []string, err error) *Repository_RoleAddMembers_Call { - _c.Call.Return(strings, err) - return _c -} - -func (_c *Repository_RoleAddMembers_Call) RunAndReturn(run func(ctx context.Context, role roles.Role, members []string) ([]string, error)) *Repository_RoleAddMembers_Call { - _c.Call.Return(run) - return _c -} - -// RoleCheckActionsExists provides a mock function for the type Repository -func (_mock *Repository) RoleCheckActionsExists(ctx context.Context, roleID string, actions []string) (bool, error) { - ret := _mock.Called(ctx, roleID, actions) - - if len(ret) == 0 { - panic("no return value specified for RoleCheckActionsExists") - } - - var r0 bool - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, []string) (bool, error)); ok { - return returnFunc(ctx, roleID, actions) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, []string) bool); ok { - r0 = returnFunc(ctx, roleID, actions) - } else { - r0 = ret.Get(0).(bool) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, []string) error); ok { - r1 = returnFunc(ctx, roleID, actions) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RoleCheckActionsExists_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleCheckActionsExists' -type Repository_RoleCheckActionsExists_Call struct { - *mock.Call -} - -// RoleCheckActionsExists is a helper method to define mock.On call -// - ctx context.Context -// - roleID string -// - actions []string -func (_e *Repository_Expecter) RoleCheckActionsExists(ctx interface{}, roleID interface{}, actions interface{}) *Repository_RoleCheckActionsExists_Call { - return &Repository_RoleCheckActionsExists_Call{Call: _e.mock.On("RoleCheckActionsExists", ctx, roleID, actions)} -} - -func (_c *Repository_RoleCheckActionsExists_Call) Run(run func(ctx context.Context, roleID string, actions []string)) *Repository_RoleCheckActionsExists_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RoleCheckActionsExists_Call) Return(b bool, err error) *Repository_RoleCheckActionsExists_Call { - _c.Call.Return(b, err) - return _c -} - -func (_c *Repository_RoleCheckActionsExists_Call) RunAndReturn(run func(ctx context.Context, roleID string, actions []string) (bool, error)) *Repository_RoleCheckActionsExists_Call { - _c.Call.Return(run) - return _c -} - -// RoleCheckMembersExists provides a mock function for the type Repository -func (_mock *Repository) RoleCheckMembersExists(ctx context.Context, roleID string, members []string) (bool, error) { - ret := _mock.Called(ctx, roleID, members) - - if len(ret) == 0 { - panic("no return value specified for RoleCheckMembersExists") - } - - var r0 bool - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, []string) (bool, error)); ok { - return returnFunc(ctx, roleID, members) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, []string) bool); ok { - r0 = returnFunc(ctx, roleID, members) - } else { - r0 = ret.Get(0).(bool) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, []string) error); ok { - r1 = returnFunc(ctx, roleID, members) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RoleCheckMembersExists_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleCheckMembersExists' -type Repository_RoleCheckMembersExists_Call struct { - *mock.Call -} - -// RoleCheckMembersExists is a helper method to define mock.On call -// - ctx context.Context -// - roleID string -// - members []string -func (_e *Repository_Expecter) RoleCheckMembersExists(ctx interface{}, roleID interface{}, members interface{}) *Repository_RoleCheckMembersExists_Call { - return &Repository_RoleCheckMembersExists_Call{Call: _e.mock.On("RoleCheckMembersExists", ctx, roleID, members)} -} - -func (_c *Repository_RoleCheckMembersExists_Call) Run(run func(ctx context.Context, roleID string, members []string)) *Repository_RoleCheckMembersExists_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RoleCheckMembersExists_Call) Return(b bool, err error) *Repository_RoleCheckMembersExists_Call { - _c.Call.Return(b, err) - return _c -} - -func (_c *Repository_RoleCheckMembersExists_Call) RunAndReturn(run func(ctx context.Context, roleID string, members []string) (bool, error)) *Repository_RoleCheckMembersExists_Call { - _c.Call.Return(run) - return _c -} - -// RoleListActions provides a mock function for the type Repository -func (_mock *Repository) RoleListActions(ctx context.Context, roleID string) ([]string, error) { - ret := _mock.Called(ctx, roleID) - - if len(ret) == 0 { - panic("no return value specified for RoleListActions") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) ([]string, error)); ok { - return returnFunc(ctx, roleID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string) []string); ok { - r0 = returnFunc(ctx, roleID) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = returnFunc(ctx, roleID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RoleListActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleListActions' -type Repository_RoleListActions_Call struct { - *mock.Call -} - -// RoleListActions is a helper method to define mock.On call -// - ctx context.Context -// - roleID string -func (_e *Repository_Expecter) RoleListActions(ctx interface{}, roleID interface{}) *Repository_RoleListActions_Call { - return &Repository_RoleListActions_Call{Call: _e.mock.On("RoleListActions", ctx, roleID)} -} - -func (_c *Repository_RoleListActions_Call) Run(run func(ctx context.Context, roleID string)) *Repository_RoleListActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RoleListActions_Call) Return(strings []string, err error) *Repository_RoleListActions_Call { - _c.Call.Return(strings, err) - return _c -} - -func (_c *Repository_RoleListActions_Call) RunAndReturn(run func(ctx context.Context, roleID string) ([]string, error)) *Repository_RoleListActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleListMembers provides a mock function for the type Repository -func (_mock *Repository) RoleListMembers(ctx context.Context, roleID string, limit uint64, offset uint64) (roles.MembersPage, error) { - ret := _mock.Called(ctx, roleID, limit, offset) - - if len(ret) == 0 { - panic("no return value specified for RoleListMembers") - } - - var r0 roles.MembersPage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, uint64, uint64) (roles.MembersPage, error)); ok { - return returnFunc(ctx, roleID, limit, offset) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, uint64, uint64) roles.MembersPage); ok { - r0 = returnFunc(ctx, roleID, limit, offset) - } else { - r0 = ret.Get(0).(roles.MembersPage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, uint64, uint64) error); ok { - r1 = returnFunc(ctx, roleID, limit, offset) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RoleListMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleListMembers' -type Repository_RoleListMembers_Call struct { - *mock.Call -} - -// RoleListMembers is a helper method to define mock.On call -// - ctx context.Context -// - roleID string -// - limit uint64 -// - offset uint64 -func (_e *Repository_Expecter) RoleListMembers(ctx interface{}, roleID interface{}, limit interface{}, offset interface{}) *Repository_RoleListMembers_Call { - return &Repository_RoleListMembers_Call{Call: _e.mock.On("RoleListMembers", ctx, roleID, limit, offset)} -} - -func (_c *Repository_RoleListMembers_Call) Run(run func(ctx context.Context, roleID string, limit uint64, offset uint64)) *Repository_RoleListMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 uint64 - if args[2] != nil { - arg2 = args[2].(uint64) - } - var arg3 uint64 - if args[3] != nil { - arg3 = args[3].(uint64) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Repository_RoleListMembers_Call) Return(membersPage roles.MembersPage, err error) *Repository_RoleListMembers_Call { - _c.Call.Return(membersPage, err) - return _c -} - -func (_c *Repository_RoleListMembers_Call) RunAndReturn(run func(ctx context.Context, roleID string, limit uint64, offset uint64) (roles.MembersPage, error)) *Repository_RoleListMembers_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveActions provides a mock function for the type Repository -func (_mock *Repository) RoleRemoveActions(ctx context.Context, role roles.Role, actions []string) error { - ret := _mock.Called(ctx, role, actions) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveActions") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role, []string) error); ok { - r0 = returnFunc(ctx, role, actions) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RoleRemoveActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveActions' -type Repository_RoleRemoveActions_Call struct { - *mock.Call -} - -// RoleRemoveActions is a helper method to define mock.On call -// - ctx context.Context -// - role roles.Role -// - actions []string -func (_e *Repository_Expecter) RoleRemoveActions(ctx interface{}, role interface{}, actions interface{}) *Repository_RoleRemoveActions_Call { - return &Repository_RoleRemoveActions_Call{Call: _e.mock.On("RoleRemoveActions", ctx, role, actions)} -} - -func (_c *Repository_RoleRemoveActions_Call) Run(run func(ctx context.Context, role roles.Role, actions []string)) *Repository_RoleRemoveActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RoleRemoveActions_Call) Return(err error) *Repository_RoleRemoveActions_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RoleRemoveActions_Call) RunAndReturn(run func(ctx context.Context, role roles.Role, actions []string) error) *Repository_RoleRemoveActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveAllActions provides a mock function for the type Repository -func (_mock *Repository) RoleRemoveAllActions(ctx context.Context, role roles.Role) error { - ret := _mock.Called(ctx, role) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveAllActions") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role) error); ok { - r0 = returnFunc(ctx, role) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RoleRemoveAllActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveAllActions' -type Repository_RoleRemoveAllActions_Call struct { - *mock.Call -} - -// RoleRemoveAllActions is a helper method to define mock.On call -// - ctx context.Context -// - role roles.Role -func (_e *Repository_Expecter) RoleRemoveAllActions(ctx interface{}, role interface{}) *Repository_RoleRemoveAllActions_Call { - return &Repository_RoleRemoveAllActions_Call{Call: _e.mock.On("RoleRemoveAllActions", ctx, role)} -} - -func (_c *Repository_RoleRemoveAllActions_Call) Run(run func(ctx context.Context, role roles.Role)) *Repository_RoleRemoveAllActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RoleRemoveAllActions_Call) Return(err error) *Repository_RoleRemoveAllActions_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RoleRemoveAllActions_Call) RunAndReturn(run func(ctx context.Context, role roles.Role) error) *Repository_RoleRemoveAllActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveAllMembers provides a mock function for the type Repository -func (_mock *Repository) RoleRemoveAllMembers(ctx context.Context, role roles.Role) error { - ret := _mock.Called(ctx, role) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveAllMembers") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role) error); ok { - r0 = returnFunc(ctx, role) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RoleRemoveAllMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveAllMembers' -type Repository_RoleRemoveAllMembers_Call struct { - *mock.Call -} - -// RoleRemoveAllMembers is a helper method to define mock.On call -// - ctx context.Context -// - role roles.Role -func (_e *Repository_Expecter) RoleRemoveAllMembers(ctx interface{}, role interface{}) *Repository_RoleRemoveAllMembers_Call { - return &Repository_RoleRemoveAllMembers_Call{Call: _e.mock.On("RoleRemoveAllMembers", ctx, role)} -} - -func (_c *Repository_RoleRemoveAllMembers_Call) Run(run func(ctx context.Context, role roles.Role)) *Repository_RoleRemoveAllMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RoleRemoveAllMembers_Call) Return(err error) *Repository_RoleRemoveAllMembers_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RoleRemoveAllMembers_Call) RunAndReturn(run func(ctx context.Context, role roles.Role) error) *Repository_RoleRemoveAllMembers_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveMembers provides a mock function for the type Repository -func (_mock *Repository) RoleRemoveMembers(ctx context.Context, role roles.Role, members []string) error { - ret := _mock.Called(ctx, role, members) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveMembers") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role, []string) error); ok { - r0 = returnFunc(ctx, role, members) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_RoleRemoveMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveMembers' -type Repository_RoleRemoveMembers_Call struct { - *mock.Call -} - -// RoleRemoveMembers is a helper method to define mock.On call -// - ctx context.Context -// - role roles.Role -// - members []string -func (_e *Repository_Expecter) RoleRemoveMembers(ctx interface{}, role interface{}, members interface{}) *Repository_RoleRemoveMembers_Call { - return &Repository_RoleRemoveMembers_Call{Call: _e.mock.On("RoleRemoveMembers", ctx, role, members)} -} - -func (_c *Repository_RoleRemoveMembers_Call) Run(run func(ctx context.Context, role roles.Role, members []string)) *Repository_RoleRemoveMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - var arg2 []string - if args[2] != nil { - arg2 = args[2].([]string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RoleRemoveMembers_Call) Return(err error) *Repository_RoleRemoveMembers_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_RoleRemoveMembers_Call) RunAndReturn(run func(ctx context.Context, role roles.Role, members []string) error) *Repository_RoleRemoveMembers_Call { - _c.Call.Return(run) - return _c -} - // UpdateReportConfig provides a mock function for the type Repository func (_mock *Repository) UpdateReportConfig(ctx context.Context, cfg reports.ReportConfig) (reports.ReportConfig, error) { ret := _mock.Called(ctx, cfg) @@ -2066,72 +634,6 @@ func (_c *Repository_UpdateReportTemplate_Call) RunAndReturn(run func(ctx contex return _c } -// UpdateRole provides a mock function for the type Repository -func (_mock *Repository) UpdateRole(ctx context.Context, ro roles.Role) (roles.Role, error) { - ret := _mock.Called(ctx, ro) - - if len(ret) == 0 { - panic("no return value specified for UpdateRole") - } - - var r0 roles.Role - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role) (roles.Role, error)); ok { - return returnFunc(ctx, ro) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, roles.Role) roles.Role); ok { - r0 = returnFunc(ctx, ro) - } else { - r0 = ret.Get(0).(roles.Role) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, roles.Role) error); ok { - r1 = returnFunc(ctx, ro) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_UpdateRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateRole' -type Repository_UpdateRole_Call struct { - *mock.Call -} - -// UpdateRole is a helper method to define mock.On call -// - ctx context.Context -// - ro roles.Role -func (_e *Repository_Expecter) UpdateRole(ctx interface{}, ro interface{}) *Repository_UpdateRole_Call { - return &Repository_UpdateRole_Call{Call: _e.mock.On("UpdateRole", ctx, ro)} -} - -func (_c *Repository_UpdateRole_Call) Run(run func(ctx context.Context, ro roles.Role)) *Repository_UpdateRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 roles.Role - if args[1] != nil { - arg1 = args[1].(roles.Role) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_UpdateRole_Call) Return(role roles.Role, err error) *Repository_UpdateRole_Call { - _c.Call.Return(role, err) - return _c -} - -func (_c *Repository_UpdateRole_Call) RunAndReturn(run func(ctx context.Context, ro roles.Role) (roles.Role, error)) *Repository_UpdateRole_Call { - _c.Call.Return(run) - return _c -} - // ViewReportConfig provides a mock function for the type Repository func (_mock *Repository) ViewReportConfig(ctx context.Context, id string) (reports.ReportConfig, error) { ret := _mock.Called(ctx, id) diff --git a/reports/mocks/service.go b/reports/mocks/service.go index dbf86039a..96fe7bdd1 100644 --- a/reports/mocks/service.go +++ b/reports/mocks/service.go @@ -12,7 +12,6 @@ import ( "context" "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/roles" "github.com/absmach/magistrala/reports" mock "github.com/stretchr/testify/mock" ) @@ -116,96 +115,6 @@ func (_c *Service_AddReportConfig_Call) RunAndReturn(run func(ctx context.Contex return _c } -// AddRole provides a mock function for the type Service -func (_mock *Service) AddRole(ctx context.Context, session authn.Session, entityID string, roleName string, optionalActions []string, optionalMembers []string) (roles.RoleProvision, error) { - ret := _mock.Called(ctx, session, entityID, roleName, optionalActions, optionalMembers) - - if len(ret) == 0 { - panic("no return value specified for AddRole") - } - - var r0 roles.RoleProvision - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string, []string) (roles.RoleProvision, error)); ok { - return returnFunc(ctx, session, entityID, roleName, optionalActions, optionalMembers) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string, []string) roles.RoleProvision); ok { - r0 = returnFunc(ctx, session, entityID, roleName, optionalActions, optionalMembers) - } else { - r0 = ret.Get(0).(roles.RoleProvision) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, []string, []string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleName, optionalActions, optionalMembers) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_AddRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'AddRole' -type Service_AddRole_Call struct { - *mock.Call -} - -// AddRole is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleName string -// - optionalActions []string -// - optionalMembers []string -func (_e *Service_Expecter) AddRole(ctx interface{}, session interface{}, entityID interface{}, roleName interface{}, optionalActions interface{}, optionalMembers interface{}) *Service_AddRole_Call { - return &Service_AddRole_Call{Call: _e.mock.On("AddRole", ctx, session, entityID, roleName, optionalActions, optionalMembers)} -} - -func (_c *Service_AddRole_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleName string, optionalActions []string, optionalMembers []string)) *Service_AddRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - var arg5 []string - if args[5] != nil { - arg5 = args[5].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - arg5, - ) - }) - return _c -} - -func (_c *Service_AddRole_Call) Return(roleProvision roles.RoleProvision, err error) *Service_AddRole_Call { - _c.Call.Return(roleProvision, err) - return _c -} - -func (_c *Service_AddRole_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleName string, optionalActions []string, optionalMembers []string) (roles.RoleProvision, error)) *Service_AddRole_Call { - _c.Call.Return(run) - return _c -} - // DeleteReportTemplate provides a mock function for the type Service func (_mock *Service) DeleteReportTemplate(ctx context.Context, session authn.Session, id string) error { ret := _mock.Called(ctx, session, id) @@ -491,152 +400,6 @@ func (_c *Service_GenerateReport_Call) RunAndReturn(run func(ctx context.Context return _c } -// ListAvailableActions provides a mock function for the type Service -func (_mock *Service) ListAvailableActions(ctx context.Context, session authn.Session) ([]string, error) { - ret := _mock.Called(ctx, session) - - if len(ret) == 0 { - panic("no return value specified for ListAvailableActions") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session) ([]string, error)); ok { - return returnFunc(ctx, session) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session) []string); ok { - r0 = returnFunc(ctx, session) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session) error); ok { - r1 = returnFunc(ctx, session) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_ListAvailableActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ListAvailableActions' -type Service_ListAvailableActions_Call struct { - *mock.Call -} - -// ListAvailableActions is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -func (_e *Service_Expecter) ListAvailableActions(ctx interface{}, session interface{}) *Service_ListAvailableActions_Call { - return &Service_ListAvailableActions_Call{Call: _e.mock.On("ListAvailableActions", ctx, session)} -} - -func (_c *Service_ListAvailableActions_Call) Run(run func(ctx context.Context, session authn.Session)) *Service_ListAvailableActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Service_ListAvailableActions_Call) Return(strings []string, err error) *Service_ListAvailableActions_Call { - _c.Call.Return(strings, err) - return _c -} - -func (_c *Service_ListAvailableActions_Call) RunAndReturn(run func(ctx context.Context, session authn.Session) ([]string, error)) *Service_ListAvailableActions_Call { - _c.Call.Return(run) - return _c -} - -// ListEntityMembers provides a mock function for the type Service -func (_mock *Service) ListEntityMembers(ctx context.Context, session authn.Session, entityID string, pq roles.MembersRolePageQuery) (roles.MembersRolePage, error) { - ret := _mock.Called(ctx, session, entityID, pq) - - if len(ret) == 0 { - panic("no return value specified for ListEntityMembers") - } - - var r0 roles.MembersRolePage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, roles.MembersRolePageQuery) (roles.MembersRolePage, error)); ok { - return returnFunc(ctx, session, entityID, pq) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, roles.MembersRolePageQuery) roles.MembersRolePage); ok { - r0 = returnFunc(ctx, session, entityID, pq) - } else { - r0 = ret.Get(0).(roles.MembersRolePage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, roles.MembersRolePageQuery) error); ok { - r1 = returnFunc(ctx, session, entityID, pq) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_ListEntityMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ListEntityMembers' -type Service_ListEntityMembers_Call struct { - *mock.Call -} - -// ListEntityMembers is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - pq roles.MembersRolePageQuery -func (_e *Service_Expecter) ListEntityMembers(ctx interface{}, session interface{}, entityID interface{}, pq interface{}) *Service_ListEntityMembers_Call { - return &Service_ListEntityMembers_Call{Call: _e.mock.On("ListEntityMembers", ctx, session, entityID, pq)} -} - -func (_c *Service_ListEntityMembers_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, pq roles.MembersRolePageQuery)) *Service_ListEntityMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 roles.MembersRolePageQuery - if args[3] != nil { - arg3 = args[3].(roles.MembersRolePageQuery) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_ListEntityMembers_Call) Return(membersRolePage roles.MembersRolePage, err error) *Service_ListEntityMembers_Call { - _c.Call.Return(membersRolePage, err) - return _c -} - -func (_c *Service_ListEntityMembers_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, pq roles.MembersRolePageQuery) (roles.MembersRolePage, error)) *Service_ListEntityMembers_Call { - _c.Call.Return(run) - return _c -} - // ListReportsConfig provides a mock function for the type Service func (_mock *Service) ListReportsConfig(ctx context.Context, session authn.Session, pm reports.PageMeta) (reports.ReportConfigPage, error) { ret := _mock.Called(ctx, session, pm) @@ -709,138 +472,6 @@ func (_c *Service_ListReportsConfig_Call) RunAndReturn(run func(ctx context.Cont return _c } -// RemoveEntityMembers provides a mock function for the type Service -func (_mock *Service) RemoveEntityMembers(ctx context.Context, session authn.Session, entityID string, members []string) error { - ret := _mock.Called(ctx, session, entityID, members) - - if len(ret) == 0 { - panic("no return value specified for RemoveEntityMembers") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, []string) error); ok { - r0 = returnFunc(ctx, session, entityID, members) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RemoveEntityMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveEntityMembers' -type Service_RemoveEntityMembers_Call struct { - *mock.Call -} - -// RemoveEntityMembers is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - members []string -func (_e *Service_Expecter) RemoveEntityMembers(ctx interface{}, session interface{}, entityID interface{}, members interface{}) *Service_RemoveEntityMembers_Call { - return &Service_RemoveEntityMembers_Call{Call: _e.mock.On("RemoveEntityMembers", ctx, session, entityID, members)} -} - -func (_c *Service_RemoveEntityMembers_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, members []string)) *Service_RemoveEntityMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 []string - if args[3] != nil { - arg3 = args[3].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_RemoveEntityMembers_Call) Return(err error) *Service_RemoveEntityMembers_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RemoveEntityMembers_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, members []string) error) *Service_RemoveEntityMembers_Call { - _c.Call.Return(run) - return _c -} - -// RemoveMemberFromAllRoles provides a mock function for the type Service -func (_mock *Service) RemoveMemberFromAllRoles(ctx context.Context, session authn.Session, memberID string) error { - ret := _mock.Called(ctx, session, memberID) - - if len(ret) == 0 { - panic("no return value specified for RemoveMemberFromAllRoles") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) error); ok { - r0 = returnFunc(ctx, session, memberID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RemoveMemberFromAllRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveMemberFromAllRoles' -type Service_RemoveMemberFromAllRoles_Call struct { - *mock.Call -} - -// RemoveMemberFromAllRoles is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - memberID string -func (_e *Service_Expecter) RemoveMemberFromAllRoles(ctx interface{}, session interface{}, memberID interface{}) *Service_RemoveMemberFromAllRoles_Call { - return &Service_RemoveMemberFromAllRoles_Call{Call: _e.mock.On("RemoveMemberFromAllRoles", ctx, session, memberID)} -} - -func (_c *Service_RemoveMemberFromAllRoles_Call) Run(run func(ctx context.Context, session authn.Session, memberID string)) *Service_RemoveMemberFromAllRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_RemoveMemberFromAllRoles_Call) Return(err error) *Service_RemoveMemberFromAllRoles_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RemoveMemberFromAllRoles_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, memberID string) error) *Service_RemoveMemberFromAllRoles_Call { - _c.Call.Return(run) - return _c -} - // RemoveReportConfig provides a mock function for the type Service func (_mock *Service) RemoveReportConfig(ctx context.Context, session authn.Session, id string) error { ret := _mock.Called(ctx, session, id) @@ -904,1035 +535,6 @@ func (_c *Service_RemoveReportConfig_Call) RunAndReturn(run func(ctx context.Con return _c } -// RemoveRole provides a mock function for the type Service -func (_mock *Service) RemoveRole(ctx context.Context, session authn.Session, entityID string, roleID string) error { - ret := _mock.Called(ctx, session, entityID, roleID) - - if len(ret) == 0 { - panic("no return value specified for RemoveRole") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) error); ok { - r0 = returnFunc(ctx, session, entityID, roleID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RemoveRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RemoveRole' -type Service_RemoveRole_Call struct { - *mock.Call -} - -// RemoveRole is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -func (_e *Service_Expecter) RemoveRole(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}) *Service_RemoveRole_Call { - return &Service_RemoveRole_Call{Call: _e.mock.On("RemoveRole", ctx, session, entityID, roleID)} -} - -func (_c *Service_RemoveRole_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string)) *Service_RemoveRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_RemoveRole_Call) Return(err error) *Service_RemoveRole_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RemoveRole_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string) error) *Service_RemoveRole_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveAllRoles provides a mock function for the type Service -func (_mock *Service) RetrieveAllRoles(ctx context.Context, session authn.Session, entityID string, limit uint64, offset uint64) (roles.RolePage, error) { - ret := _mock.Called(ctx, session, entityID, limit, offset) - - if len(ret) == 0 { - panic("no return value specified for RetrieveAllRoles") - } - - var r0 roles.RolePage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, uint64, uint64) (roles.RolePage, error)); ok { - return returnFunc(ctx, session, entityID, limit, offset) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, uint64, uint64) roles.RolePage); ok { - r0 = returnFunc(ctx, session, entityID, limit, offset) - } else { - r0 = ret.Get(0).(roles.RolePage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, uint64, uint64) error); ok { - r1 = returnFunc(ctx, session, entityID, limit, offset) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RetrieveAllRoles_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveAllRoles' -type Service_RetrieveAllRoles_Call struct { - *mock.Call -} - -// RetrieveAllRoles is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - limit uint64 -// - offset uint64 -func (_e *Service_Expecter) RetrieveAllRoles(ctx interface{}, session interface{}, entityID interface{}, limit interface{}, offset interface{}) *Service_RetrieveAllRoles_Call { - return &Service_RetrieveAllRoles_Call{Call: _e.mock.On("RetrieveAllRoles", ctx, session, entityID, limit, offset)} -} - -func (_c *Service_RetrieveAllRoles_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, limit uint64, offset uint64)) *Service_RetrieveAllRoles_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 uint64 - if args[3] != nil { - arg3 = args[3].(uint64) - } - var arg4 uint64 - if args[4] != nil { - arg4 = args[4].(uint64) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_RetrieveAllRoles_Call) Return(rolePage roles.RolePage, err error) *Service_RetrieveAllRoles_Call { - _c.Call.Return(rolePage, err) - return _c -} - -func (_c *Service_RetrieveAllRoles_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, limit uint64, offset uint64) (roles.RolePage, error)) *Service_RetrieveAllRoles_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveRole provides a mock function for the type Service -func (_mock *Service) RetrieveRole(ctx context.Context, session authn.Session, entityID string, roleID string) (roles.Role, error) { - ret := _mock.Called(ctx, session, entityID, roleID) - - if len(ret) == 0 { - panic("no return value specified for RetrieveRole") - } - - var r0 roles.Role - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) (roles.Role, error)); ok { - return returnFunc(ctx, session, entityID, roleID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) roles.Role); ok { - r0 = returnFunc(ctx, session, entityID, roleID) - } else { - r0 = ret.Get(0).(roles.Role) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RetrieveRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveRole' -type Service_RetrieveRole_Call struct { - *mock.Call -} - -// RetrieveRole is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -func (_e *Service_Expecter) RetrieveRole(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}) *Service_RetrieveRole_Call { - return &Service_RetrieveRole_Call{Call: _e.mock.On("RetrieveRole", ctx, session, entityID, roleID)} -} - -func (_c *Service_RetrieveRole_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string)) *Service_RetrieveRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_RetrieveRole_Call) Return(role roles.Role, err error) *Service_RetrieveRole_Call { - _c.Call.Return(role, err) - return _c -} - -func (_c *Service_RetrieveRole_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string) (roles.Role, error)) *Service_RetrieveRole_Call { - _c.Call.Return(run) - return _c -} - -// RoleAddActions provides a mock function for the type Service -func (_mock *Service) RoleAddActions(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string) ([]string, error) { - ret := _mock.Called(ctx, session, entityID, roleID, actions) - - if len(ret) == 0 { - panic("no return value specified for RoleAddActions") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) ([]string, error)); ok { - return returnFunc(ctx, session, entityID, roleID, actions) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) []string); ok { - r0 = returnFunc(ctx, session, entityID, roleID, actions) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, []string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID, actions) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RoleAddActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleAddActions' -type Service_RoleAddActions_Call struct { - *mock.Call -} - -// RoleAddActions is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - actions []string -func (_e *Service_Expecter) RoleAddActions(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, actions interface{}) *Service_RoleAddActions_Call { - return &Service_RoleAddActions_Call{Call: _e.mock.On("RoleAddActions", ctx, session, entityID, roleID, actions)} -} - -func (_c *Service_RoleAddActions_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string)) *Service_RoleAddActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_RoleAddActions_Call) Return(ops []string, err error) *Service_RoleAddActions_Call { - _c.Call.Return(ops, err) - return _c -} - -func (_c *Service_RoleAddActions_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string) ([]string, error)) *Service_RoleAddActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleAddMembers provides a mock function for the type Service -func (_mock *Service) RoleAddMembers(ctx context.Context, session authn.Session, entityID string, roleID string, members []string) ([]string, error) { - ret := _mock.Called(ctx, session, entityID, roleID, members) - - if len(ret) == 0 { - panic("no return value specified for RoleAddMembers") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) ([]string, error)); ok { - return returnFunc(ctx, session, entityID, roleID, members) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) []string); ok { - r0 = returnFunc(ctx, session, entityID, roleID, members) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, []string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID, members) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RoleAddMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleAddMembers' -type Service_RoleAddMembers_Call struct { - *mock.Call -} - -// RoleAddMembers is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - members []string -func (_e *Service_Expecter) RoleAddMembers(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, members interface{}) *Service_RoleAddMembers_Call { - return &Service_RoleAddMembers_Call{Call: _e.mock.On("RoleAddMembers", ctx, session, entityID, roleID, members)} -} - -func (_c *Service_RoleAddMembers_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, members []string)) *Service_RoleAddMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_RoleAddMembers_Call) Return(strings []string, err error) *Service_RoleAddMembers_Call { - _c.Call.Return(strings, err) - return _c -} - -func (_c *Service_RoleAddMembers_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, members []string) ([]string, error)) *Service_RoleAddMembers_Call { - _c.Call.Return(run) - return _c -} - -// RoleCheckActionsExists provides a mock function for the type Service -func (_mock *Service) RoleCheckActionsExists(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string) (bool, error) { - ret := _mock.Called(ctx, session, entityID, roleID, actions) - - if len(ret) == 0 { - panic("no return value specified for RoleCheckActionsExists") - } - - var r0 bool - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) (bool, error)); ok { - return returnFunc(ctx, session, entityID, roleID, actions) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) bool); ok { - r0 = returnFunc(ctx, session, entityID, roleID, actions) - } else { - r0 = ret.Get(0).(bool) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, []string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID, actions) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RoleCheckActionsExists_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleCheckActionsExists' -type Service_RoleCheckActionsExists_Call struct { - *mock.Call -} - -// RoleCheckActionsExists is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - actions []string -func (_e *Service_Expecter) RoleCheckActionsExists(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, actions interface{}) *Service_RoleCheckActionsExists_Call { - return &Service_RoleCheckActionsExists_Call{Call: _e.mock.On("RoleCheckActionsExists", ctx, session, entityID, roleID, actions)} -} - -func (_c *Service_RoleCheckActionsExists_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string)) *Service_RoleCheckActionsExists_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_RoleCheckActionsExists_Call) Return(b bool, err error) *Service_RoleCheckActionsExists_Call { - _c.Call.Return(b, err) - return _c -} - -func (_c *Service_RoleCheckActionsExists_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string) (bool, error)) *Service_RoleCheckActionsExists_Call { - _c.Call.Return(run) - return _c -} - -// RoleCheckMembersExists provides a mock function for the type Service -func (_mock *Service) RoleCheckMembersExists(ctx context.Context, session authn.Session, entityID string, roleID string, members []string) (bool, error) { - ret := _mock.Called(ctx, session, entityID, roleID, members) - - if len(ret) == 0 { - panic("no return value specified for RoleCheckMembersExists") - } - - var r0 bool - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) (bool, error)); ok { - return returnFunc(ctx, session, entityID, roleID, members) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) bool); ok { - r0 = returnFunc(ctx, session, entityID, roleID, members) - } else { - r0 = ret.Get(0).(bool) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, []string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID, members) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RoleCheckMembersExists_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleCheckMembersExists' -type Service_RoleCheckMembersExists_Call struct { - *mock.Call -} - -// RoleCheckMembersExists is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - members []string -func (_e *Service_Expecter) RoleCheckMembersExists(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, members interface{}) *Service_RoleCheckMembersExists_Call { - return &Service_RoleCheckMembersExists_Call{Call: _e.mock.On("RoleCheckMembersExists", ctx, session, entityID, roleID, members)} -} - -func (_c *Service_RoleCheckMembersExists_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, members []string)) *Service_RoleCheckMembersExists_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_RoleCheckMembersExists_Call) Return(b bool, err error) *Service_RoleCheckMembersExists_Call { - _c.Call.Return(b, err) - return _c -} - -func (_c *Service_RoleCheckMembersExists_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, members []string) (bool, error)) *Service_RoleCheckMembersExists_Call { - _c.Call.Return(run) - return _c -} - -// RoleListActions provides a mock function for the type Service -func (_mock *Service) RoleListActions(ctx context.Context, session authn.Session, entityID string, roleID string) ([]string, error) { - ret := _mock.Called(ctx, session, entityID, roleID) - - if len(ret) == 0 { - panic("no return value specified for RoleListActions") - } - - var r0 []string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) ([]string, error)); ok { - return returnFunc(ctx, session, entityID, roleID) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) []string); ok { - r0 = returnFunc(ctx, session, entityID, roleID) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).([]string) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RoleListActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleListActions' -type Service_RoleListActions_Call struct { - *mock.Call -} - -// RoleListActions is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -func (_e *Service_Expecter) RoleListActions(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}) *Service_RoleListActions_Call { - return &Service_RoleListActions_Call{Call: _e.mock.On("RoleListActions", ctx, session, entityID, roleID)} -} - -func (_c *Service_RoleListActions_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string)) *Service_RoleListActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_RoleListActions_Call) Return(strings []string, err error) *Service_RoleListActions_Call { - _c.Call.Return(strings, err) - return _c -} - -func (_c *Service_RoleListActions_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string) ([]string, error)) *Service_RoleListActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleListMembers provides a mock function for the type Service -func (_mock *Service) RoleListMembers(ctx context.Context, session authn.Session, entityID string, roleID string, limit uint64, offset uint64) (roles.MembersPage, error) { - ret := _mock.Called(ctx, session, entityID, roleID, limit, offset) - - if len(ret) == 0 { - panic("no return value specified for RoleListMembers") - } - - var r0 roles.MembersPage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, uint64, uint64) (roles.MembersPage, error)); ok { - return returnFunc(ctx, session, entityID, roleID, limit, offset) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, uint64, uint64) roles.MembersPage); ok { - r0 = returnFunc(ctx, session, entityID, roleID, limit, offset) - } else { - r0 = ret.Get(0).(roles.MembersPage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, uint64, uint64) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID, limit, offset) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RoleListMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleListMembers' -type Service_RoleListMembers_Call struct { - *mock.Call -} - -// RoleListMembers is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - limit uint64 -// - offset uint64 -func (_e *Service_Expecter) RoleListMembers(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, limit interface{}, offset interface{}) *Service_RoleListMembers_Call { - return &Service_RoleListMembers_Call{Call: _e.mock.On("RoleListMembers", ctx, session, entityID, roleID, limit, offset)} -} - -func (_c *Service_RoleListMembers_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, limit uint64, offset uint64)) *Service_RoleListMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 uint64 - if args[4] != nil { - arg4 = args[4].(uint64) - } - var arg5 uint64 - if args[5] != nil { - arg5 = args[5].(uint64) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - arg5, - ) - }) - return _c -} - -func (_c *Service_RoleListMembers_Call) Return(membersPage roles.MembersPage, err error) *Service_RoleListMembers_Call { - _c.Call.Return(membersPage, err) - return _c -} - -func (_c *Service_RoleListMembers_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, limit uint64, offset uint64) (roles.MembersPage, error)) *Service_RoleListMembers_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveActions provides a mock function for the type Service -func (_mock *Service) RoleRemoveActions(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string) error { - ret := _mock.Called(ctx, session, entityID, roleID, actions) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveActions") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) error); ok { - r0 = returnFunc(ctx, session, entityID, roleID, actions) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RoleRemoveActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveActions' -type Service_RoleRemoveActions_Call struct { - *mock.Call -} - -// RoleRemoveActions is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - actions []string -func (_e *Service_Expecter) RoleRemoveActions(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, actions interface{}) *Service_RoleRemoveActions_Call { - return &Service_RoleRemoveActions_Call{Call: _e.mock.On("RoleRemoveActions", ctx, session, entityID, roleID, actions)} -} - -func (_c *Service_RoleRemoveActions_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string)) *Service_RoleRemoveActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_RoleRemoveActions_Call) Return(err error) *Service_RoleRemoveActions_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RoleRemoveActions_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, actions []string) error) *Service_RoleRemoveActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveAllActions provides a mock function for the type Service -func (_mock *Service) RoleRemoveAllActions(ctx context.Context, session authn.Session, entityID string, roleID string) error { - ret := _mock.Called(ctx, session, entityID, roleID) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveAllActions") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) error); ok { - r0 = returnFunc(ctx, session, entityID, roleID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RoleRemoveAllActions_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveAllActions' -type Service_RoleRemoveAllActions_Call struct { - *mock.Call -} - -// RoleRemoveAllActions is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -func (_e *Service_Expecter) RoleRemoveAllActions(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}) *Service_RoleRemoveAllActions_Call { - return &Service_RoleRemoveAllActions_Call{Call: _e.mock.On("RoleRemoveAllActions", ctx, session, entityID, roleID)} -} - -func (_c *Service_RoleRemoveAllActions_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string)) *Service_RoleRemoveAllActions_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_RoleRemoveAllActions_Call) Return(err error) *Service_RoleRemoveAllActions_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RoleRemoveAllActions_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string) error) *Service_RoleRemoveAllActions_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveAllMembers provides a mock function for the type Service -func (_mock *Service) RoleRemoveAllMembers(ctx context.Context, session authn.Session, entityID string, roleID string) error { - ret := _mock.Called(ctx, session, entityID, roleID) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveAllMembers") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) error); ok { - r0 = returnFunc(ctx, session, entityID, roleID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RoleRemoveAllMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveAllMembers' -type Service_RoleRemoveAllMembers_Call struct { - *mock.Call -} - -// RoleRemoveAllMembers is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -func (_e *Service_Expecter) RoleRemoveAllMembers(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}) *Service_RoleRemoveAllMembers_Call { - return &Service_RoleRemoveAllMembers_Call{Call: _e.mock.On("RoleRemoveAllMembers", ctx, session, entityID, roleID)} -} - -func (_c *Service_RoleRemoveAllMembers_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string)) *Service_RoleRemoveAllMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_RoleRemoveAllMembers_Call) Return(err error) *Service_RoleRemoveAllMembers_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RoleRemoveAllMembers_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string) error) *Service_RoleRemoveAllMembers_Call { - _c.Call.Return(run) - return _c -} - -// RoleRemoveMembers provides a mock function for the type Service -func (_mock *Service) RoleRemoveMembers(ctx context.Context, session authn.Session, entityID string, roleID string, members []string) error { - ret := _mock.Called(ctx, session, entityID, roleID, members) - - if len(ret) == 0 { - panic("no return value specified for RoleRemoveMembers") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, []string) error); ok { - r0 = returnFunc(ctx, session, entityID, roleID, members) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RoleRemoveMembers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RoleRemoveMembers' -type Service_RoleRemoveMembers_Call struct { - *mock.Call -} - -// RoleRemoveMembers is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - members []string -func (_e *Service_Expecter) RoleRemoveMembers(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, members interface{}) *Service_RoleRemoveMembers_Call { - return &Service_RoleRemoveMembers_Call{Call: _e.mock.On("RoleRemoveMembers", ctx, session, entityID, roleID, members)} -} - -func (_c *Service_RoleRemoveMembers_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, members []string)) *Service_RoleRemoveMembers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 []string - if args[4] != nil { - arg4 = args[4].([]string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_RoleRemoveMembers_Call) Return(err error) *Service_RoleRemoveMembers_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RoleRemoveMembers_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, members []string) error) *Service_RoleRemoveMembers_Call { - _c.Call.Return(run) - return _c -} - // StartScheduler provides a mock function for the type Service func (_mock *Service) StartScheduler(ctx context.Context) error { ret := _mock.Called(ctx) @@ -2191,90 +793,6 @@ func (_c *Service_UpdateReportTemplate_Call) RunAndReturn(run func(ctx context.C return _c } -// UpdateRoleName provides a mock function for the type Service -func (_mock *Service) UpdateRoleName(ctx context.Context, session authn.Session, entityID string, roleID string, newRoleName string) (roles.Role, error) { - ret := _mock.Called(ctx, session, entityID, roleID, newRoleName) - - if len(ret) == 0 { - panic("no return value specified for UpdateRoleName") - } - - var r0 roles.Role - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, string) (roles.Role, error)); ok { - return returnFunc(ctx, session, entityID, roleID, newRoleName) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string, string) roles.Role); ok { - r0 = returnFunc(ctx, session, entityID, roleID, newRoleName) - } else { - r0 = ret.Get(0).(roles.Role) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string, string) error); ok { - r1 = returnFunc(ctx, session, entityID, roleID, newRoleName) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_UpdateRoleName_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateRoleName' -type Service_UpdateRoleName_Call struct { - *mock.Call -} - -// UpdateRoleName is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - entityID string -// - roleID string -// - newRoleName string -func (_e *Service_Expecter) UpdateRoleName(ctx interface{}, session interface{}, entityID interface{}, roleID interface{}, newRoleName interface{}) *Service_UpdateRoleName_Call { - return &Service_UpdateRoleName_Call{Call: _e.mock.On("UpdateRoleName", ctx, session, entityID, roleID, newRoleName)} -} - -func (_c *Service_UpdateRoleName_Call) Run(run func(ctx context.Context, session authn.Session, entityID string, roleID string, newRoleName string)) *Service_UpdateRoleName_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - var arg4 string - if args[4] != nil { - arg4 = args[4].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - arg4, - ) - }) - return _c -} - -func (_c *Service_UpdateRoleName_Call) Return(role roles.Role, err error) *Service_UpdateRoleName_Call { - _c.Call.Return(role, err) - return _c -} - -func (_c *Service_UpdateRoleName_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, entityID string, roleID string, newRoleName string) (roles.Role, error)) *Service_UpdateRoleName_Call { - _c.Call.Return(run) - return _c -} - // ViewReportConfig provides a mock function for the type Service func (_mock *Service) ViewReportConfig(ctx context.Context, session authn.Session, id string, withRoles bool) (reports.ReportConfig, error) { ret := _mock.Called(ctx, session, id, withRoles) diff --git a/reports/postgres/init.go b/reports/postgres/init.go index 14c3c815e..03fdb8bec 100644 --- a/reports/postgres/init.go +++ b/reports/postgres/init.go @@ -4,19 +4,11 @@ package postgres import ( - dpostgres "github.com/absmach/magistrala/domains/postgres" - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - rolesPostgres "github.com/absmach/magistrala/pkg/roles/repo/postgres" _ "github.com/jackc/pgx/v5/stdlib" // required for SQL access migrate "github.com/rubenv/sql-migrate" ) func Migration() (*migrate.MemoryMigrationSource, error) { - rolesMigration, err := rolesPostgres.Migration(rolesTableNamePrefix, entityTableName, entityIDColumnName) - if err != nil { - return &migrate.MemoryMigrationSource{}, errors.Wrap(repoerr.ErrRoleMigration, err) - } reportsMigration := &migrate.MemoryMigrationSource{ Migrations: []*migrate.Migration{ { @@ -104,13 +96,5 @@ func Migration() (*migrate.MemoryMigrationSource, error) { }, } - reportsMigration.Migrations = append(reportsMigration.Migrations, rolesMigration.Migrations...) - - domainsMigration, err := dpostgres.Migration() - if err != nil { - return &migrate.MemoryMigrationSource{}, errors.Wrap(repoerr.ErrRoleMigration, err) - } - reportsMigration.Migrations = append(reportsMigration.Migrations, domainsMigration.Migrations...) - return reportsMigration, nil } diff --git a/reports/postgres/reports.go b/reports/postgres/reports.go index 32ac857bb..86a2da1f7 100644 --- a/reports/postgres/reports.go +++ b/reports/postgres/reports.go @@ -9,41 +9,29 @@ import ( "time" "github.com/absmach/magistrala/pkg/errors" - "github.com/absmach/magistrala/pkg/roles" "github.com/absmach/magistrala/pkg/schedule" "github.com/absmach/magistrala/reports" - "github.com/lib/pq" ) // dbReport represents the database structure for a Report. type dbReport struct { - ID string `db:"id"` - Name string `db:"name"` - Description string `db:"description"` - DomainID string `db:"domain_id"` - StartDateTime sql.NullTime `db:"start_datetime"` - Due sql.NullTime `db:"due"` - Recurring schedule.Recurring `db:"recurring"` - RecurringPeriod uint `db:"recurring_period"` - Status reports.Status `db:"status"` - CreatedAt time.Time `db:"created_at"` - CreatedBy string `db:"created_by"` - UpdatedAt time.Time `db:"updated_at"` - UpdatedBy string `db:"updated_by"` - Config []byte `db:"config,omitempty"` - Metrics []byte `db:"metrics"` - Email []byte `db:"email"` - ReportTemplate reports.ReportTemplate `db:"report_template"` - MemberID string `db:"member_id,omitempty"` - RoleID string `db:"role_id,omitempty"` - RoleName string `db:"role_name,omitempty"` - Actions pq.StringArray `db:"actions,omitempty"` - AccessType string `db:"access_type,omitempty"` - AccessProviderId string `db:"access_provider_id,omitempty"` - AccessProviderRoleId string `db:"access_provider_role_id,omitempty"` - AccessProviderRoleName string `db:"access_provider_role_name,omitempty"` - AccessProviderRoleActions pq.StringArray `db:"access_provider_role_actions,omitempty"` - Roles json.RawMessage `db:"roles,omitempty"` + ID string `db:"id"` + Name string `db:"name"` + Description string `db:"description"` + DomainID string `db:"domain_id"` + StartDateTime sql.NullTime `db:"start_datetime"` + Due sql.NullTime `db:"due"` + Recurring schedule.Recurring `db:"recurring"` + RecurringPeriod uint `db:"recurring_period"` + Status reports.Status `db:"status"` + CreatedAt time.Time `db:"created_at"` + CreatedBy string `db:"created_by"` + UpdatedAt time.Time `db:"updated_at"` + UpdatedBy string `db:"updated_by"` + Config []byte `db:"config,omitempty"` + Metrics []byte `db:"metrics"` + Email []byte `db:"email"` + ReportTemplate reports.ReportTemplate `db:"report_template"` } func reportToDb(r reports.ReportConfig) (dbReport, error) { @@ -125,13 +113,6 @@ func dbToReport(dto dbReport) (reports.ReportConfig, error) { } } - var roles []roles.MemberRoleActions - if dto.Roles != nil { - if err := json.Unmarshal(dto.Roles, &roles); err != nil { - return reports.ReportConfig{}, errors.Wrap(errors.ErrMalformedEntity, err) - } - } - rpt := reports.ReportConfig{ ID: dto.ID, Name: dto.Name, @@ -145,22 +126,13 @@ func dbToReport(dto dbReport) (reports.ReportConfig, error) { Recurring: dto.Recurring, RecurringPeriod: dto.RecurringPeriod, }, - Email: &email, - Status: dto.Status, - CreatedAt: dto.CreatedAt, - CreatedBy: dto.CreatedBy, - UpdatedAt: dto.UpdatedAt, - UpdatedBy: dto.UpdatedBy, - ReportTemplate: dto.ReportTemplate, - RoleID: dto.RoleID, - RoleName: dto.RoleName, - Actions: []string(dto.Actions), - AccessType: dto.AccessType, - AccessProviderId: dto.AccessProviderId, - AccessProviderRoleId: dto.AccessProviderRoleId, - AccessProviderRoleName: dto.AccessProviderRoleName, - AccessProviderRoleActions: []string(dto.AccessProviderRoleActions), - Roles: roles, + Email: &email, + Status: dto.Status, + CreatedAt: dto.CreatedAt, + CreatedBy: dto.CreatedBy, + UpdatedAt: dto.UpdatedAt, + UpdatedBy: dto.UpdatedBy, + ReportTemplate: dto.ReportTemplate, } return rpt, nil diff --git a/reports/postgres/repository.go b/reports/postgres/repository.go index 09a5d8c95..4098ef457 100644 --- a/reports/postgres/repository.go +++ b/reports/postgres/repository.go @@ -13,33 +13,22 @@ import ( api "github.com/absmach/magistrala/api/http" "github.com/absmach/magistrala/pkg/errors" repoerr "github.com/absmach/magistrala/pkg/errors/repository" - mgPolicies "github.com/absmach/magistrala/pkg/policies" "github.com/absmach/magistrala/pkg/postgres" - rolesPostgres "github.com/absmach/magistrala/pkg/roles/repo/postgres" "github.com/absmach/magistrala/reports" ) -const ( - rolesTableNamePrefix = "reports" - entityTableName = "report_config" - entityIDColumnName = "id" -) - type PostgresRepository struct { DB postgres.Database eh errors.Handler - rolesPostgres.Repository } func NewRepository(db postgres.Database) reports.Repository { - rolesRepo := rolesPostgres.NewRepository(db, mgPolicies.ReportsType, rolesTableNamePrefix, entityTableName, entityIDColumnName) errHandlerOptions := []errors.HandlerOption{ postgres.WithDuplicateErrors(NewDuplicateErrors()), } return &PostgresRepository{ - DB: db, - eh: postgres.NewErrorHandler(errHandlerOptions...), - Repository: rolesRepo, + DB: db, + eh: postgres.NewErrorHandler(errHandlerOptions...), } } @@ -103,155 +92,6 @@ func (repo *PostgresRepository) ViewReportConfig(ctx context.Context, id string) return rpt, nil } -func (repo *PostgresRepository) RetrieveByIDWithRoles(ctx context.Context, id, memberID string) (reports.ReportConfig, error) { - query := ` - WITH selected_report AS ( - SELECT - r.id, - r.domain_id - FROM - report_config r - WHERE - r.id = :id - LIMIT 1 - ), - selected_report_roles AS ( - SELECT - rr.entity_id AS report_id, - rrm.member_id AS member_id, - rr.id AS role_id, - rr."name" AS role_name, - jsonb_agg(DISTINCT rra."action") AS actions, - 'direct' AS access_type, - '' AS access_provider_id - FROM - reports_roles rr - JOIN - reports_role_members rrm ON rr.id = rrm.role_id - JOIN - reports_role_actions rra ON rr.id = rra.role_id - JOIN - selected_report sr ON sr.id = rr.entity_id - AND rrm.member_id = :member_id - GROUP BY - rr.entity_id, rr.id, rr.name, rrm.member_id - ), - selected_domain_roles AS ( - SELECT - sr.id AS report_id, - drm.member_id AS member_id, - dr.id AS role_id, - dr."name" AS role_name, - jsonb_agg(DISTINCT all_actions."action") AS actions, - 'domain' AS access_type, - dr.entity_id AS access_provider_id - FROM - domains d - JOIN - selected_report sr ON sr.domain_id = d.id - JOIN - domains_roles dr ON dr.entity_id = d.id - JOIN - domains_role_members drm ON dr.id = drm.role_id - JOIN - domains_role_actions dra ON dr.id = dra.role_id - JOIN - domains_role_actions all_actions ON dr.id = all_actions.role_id - WHERE - drm.member_id = :member_id - AND dra."action" LIKE 'report%' - GROUP BY - sr.id, dr.entity_id, dr.id, dr."name", drm.member_id - ), - all_roles AS ( - SELECT - srr.report_id, - srr.member_id, - srr.role_id, - srr.role_name, - srr.actions, - srr.access_type, - srr.access_provider_id - FROM - selected_report_roles srr - UNION - SELECT - sdr.report_id, - sdr.member_id, - sdr.role_id, - sdr.role_name, - sdr.actions, - sdr.access_type, - sdr.access_provider_id - FROM - selected_domain_roles sdr - ), - final_roles AS ( - SELECT - ar.report_id, - ar.member_id, - jsonb_agg( - jsonb_build_object( - 'role_id', ar.role_id, - 'role_name', ar.role_name, - 'actions', ar.actions, - 'access_type', ar.access_type, - 'access_provider_id', ar.access_provider_id - ) - ) AS roles - FROM all_roles ar - GROUP BY - ar.report_id, ar.member_id - ) - SELECT - r2.id, - r2."name", - r2.description, - r2.domain_id, - r2.status, - r2.created_at, - r2.created_by, - r2.updated_at, - r2.updated_by, - r2.due, - r2.recurring, - r2.recurring_period, - r2.start_datetime, - r2.config, - r2.email, - r2.metrics, - r2.report_template, - fr.member_id, - fr.roles - FROM report_config r2 - JOIN final_roles fr ON fr.report_id = r2.id - ` - parameters := map[string]any{ - "id": id, - "member_id": memberID, - } - row, err := repo.DB.NamedQueryContext(ctx, query, parameters) - if err != nil { - return reports.ReportConfig{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - defer row.Close() - - dbreport := dbReport{} - if !row.Next() { - return reports.ReportConfig{}, repoerr.ErrNotFound - } - - if err := row.StructScan(&dbreport); err != nil { - return reports.ReportConfig{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - - cfg, err := dbToReport(dbreport) - if err != nil { - return reports.ReportConfig{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - return cfg, nil -} - func (repo *PostgresRepository) UpdateReportConfigStatus(ctx context.Context, cfg reports.ReportConfig) (reports.ReportConfig, error) { q := `UPDATE report_config SET status = :status, updated_at = :updated_at, updated_by = :updated_by WHERE id = :id @@ -451,103 +291,6 @@ func (repo *PostgresRepository) ListAllReportsConfig(ctx context.Context, pm rep return ret, nil } -func (repo *PostgresRepository) ListUserReportsConfig(ctx context.Context, userID string, pm reports.PageMeta) (reports.ReportConfigPage, error) { - pm.UserID = userID - - additionalConditions := pageReportQueryConditions(pm) - additionalWhereClause := "" - if len(additionalConditions) > 0 { - additionalWhereClause = "AND " + strings.Join(additionalConditions, " AND ") - } - - orderClause := reportsOrderClause(pm) - pgData := reportsPageData(pm) - - innerQ := fmt.Sprintf(` - WITH direct_reports AS ( - SELECT rc.id, rc.name, rc.description, rc.domain_id, rc.metrics, rc.email, rc.config, - rc.start_datetime, rc.due, rc.recurring, rc.recurring_period, - rc.created_at, rc.created_by, rc.updated_at, rc.updated_by, rc.status, - rr.id AS role_id, - rr."name" AS role_name, - array_remove(array_agg(DISTINCT rra."action"), NULL) AS actions, - 'direct' AS access_type, - '' AS access_provider_id, - '' AS access_provider_role_id, - '' AS access_provider_role_name, - CAST(array[] AS text[]) AS access_provider_role_actions - FROM reports_role_members rrm - JOIN reports_roles rr ON rr.id = rrm.role_id - JOIN report_config rc ON rc.id = rr.entity_id - LEFT JOIN reports_role_actions rra ON rra.role_id = rrm.role_id - WHERE rrm.member_id = :user_id - %s - GROUP BY rc.id, rr.id, rr."name" - ), - domain_reports AS ( - SELECT rc.id, rc.name, rc.description, rc.domain_id, rc.metrics, rc.email, rc.config, - rc.start_datetime, rc.due, rc.recurring, rc.recurring_period, - rc.created_at, rc.created_by, rc.updated_at, rc.updated_by, rc.status, - '' AS role_id, - '' AS role_name, - CAST(array[] AS text[]) AS actions, - 'domain' AS access_type, - d.id AS access_provider_id, - dr.id AS access_provider_role_id, - dr."name" AS access_provider_role_name, - array_agg(DISTINCT dra."action") AS access_provider_role_actions - FROM domains_role_members drm - JOIN domains_role_actions dra ON dra.role_id = drm.role_id - JOIN domains_roles dr ON dr.id = drm.role_id - JOIN domains d ON d.id = dr.entity_id - JOIN report_config rc ON rc.domain_id = d.id - WHERE drm.member_id = :user_id - AND dra.action LIKE 'report%%' - AND NOT EXISTS (SELECT 1 FROM direct_reports tmp WHERE tmp.id = rc.id) - %s - GROUP BY rc.id, d.id, dr.id, dr."name" - ) - SELECT * FROM direct_reports - UNION ALL - SELECT * FROM domain_reports - `, additionalWhereClause, additionalWhereClause) - - q := fmt.Sprintf(` - SELECT * FROM (%s) AS sub %s %s; - `, innerQ, orderClause, pgData) - - rows, err := repo.DB.NamedQueryContext(ctx, q, pm) - if err != nil { - return reports.ReportConfigPage{}, err - } - defer rows.Close() - - cfgs := []reports.ReportConfig{} - for rows.Next() { - var r dbReport - if err := rows.StructScan(&r); err != nil { - return reports.ReportConfigPage{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - rpt, err := dbToReport(r) - if err != nil { - return reports.ReportConfigPage{}, err - } - cfgs = append(cfgs, rpt) - } - - cq := fmt.Sprintf(`SELECT COUNT(*) FROM (%s) AS count_sub;`, innerQ) - total, err := postgres.Total(ctx, repo.DB, cq, pm) - if err != nil { - return reports.ReportConfigPage{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - pm.Total = total - - return reports.ReportConfigPage{ - PageMeta: pm, - ReportConfigs: cfgs, - }, nil -} - func (repo *PostgresRepository) UpdateReportDue(ctx context.Context, id string, due time.Time) (reports.ReportConfig, error) { q := ` UPDATE report_config diff --git a/reports/postgres/repository_test.go b/reports/postgres/repository_test.go index 1a50e732b..14ae8408e 100644 --- a/reports/postgres/repository_test.go +++ b/reports/postgres/repository_test.go @@ -483,164 +483,6 @@ func TestListReportsConfig(t *testing.T) { } } -func TestListUserReportsConfig(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM domains_role_actions") - require.Nil(t, err, fmt.Sprintf("clean domains_role_actions unexpected error: %s", err)) - _, err = db.Exec("DELETE FROM domains_role_members") - require.Nil(t, err, fmt.Sprintf("clean domains_role_members unexpected error: %s", err)) - _, err = db.Exec("DELETE FROM domains_roles") - require.Nil(t, err, fmt.Sprintf("clean domains_roles unexpected error: %s", err)) - _, err = db.Exec("DELETE FROM domains") - require.Nil(t, err, fmt.Sprintf("clean domains unexpected error: %s", err)) - _, err = db.Exec("DELETE FROM report_config") - require.Nil(t, err, fmt.Sprintf("clean report_config unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - domainID := generateUUID(t) - domainRoute := generateUUID(t) - userID := generateUUID(t) - domainUserID := generateUUID(t) - otherUserID := generateUUID(t) - - _, err := db.Exec(`INSERT INTO domains (id, name, route, status) VALUES ($1, $2, $3, $4)`, domainID, namegen.Generate(), domainRoute, 0) - require.Nil(t, err, fmt.Sprintf("insert domains unexpected error: %s", err)) - - num := 10 - var allCfgs []reports.ReportConfig - for i := range num { - cfg := reports.ReportConfig{ - ID: generateUUID(t), - Name: fmt.Sprintf("Report-%d", i), - DomainID: domainID, - Status: reports.EnabledStatus, - CreatedAt: time.Now().UTC().Add(time.Duration(i) * time.Minute), - UpdatedAt: time.Now().UTC().Add(time.Duration(i) * time.Minute), - Metrics: []reports.ReqMetric{}, - } - cfg, err := repo.AddReportConfig(context.Background(), cfg) - require.Nil(t, err, fmt.Sprintf("unexpected error: %s", err)) - allCfgs = append(allCfgs, cfg) - } - - // Assign userID to the first 5 report configs via direct role INSERT. - for i := range 5 { - roleID := generateUUID(t) - _, err := db.Exec(`INSERT INTO reports_roles (id, name, entity_id) VALUES ($1, $2, $3)`, roleID, "admin", allCfgs[i].ID) - require.Nil(t, err, fmt.Sprintf("insert reports_roles unexpected error: %s", err)) - _, err = db.Exec(`INSERT INTO reports_role_members (role_id, member_id, entity_id) VALUES ($1, $2, $3)`, roleID, userID, allCfgs[i].ID) - require.Nil(t, err, fmt.Sprintf("insert reports_role_members unexpected error: %s", err)) - } - - domainRoleID := generateUUID(t) - _, err = db.Exec(`INSERT INTO domains_roles (id, name, entity_id) VALUES ($1, $2, $3)`, domainRoleID, "admin", domainID) - require.Nil(t, err, fmt.Sprintf("insert domains_roles unexpected error: %s", err)) - _, err = db.Exec(`INSERT INTO domains_role_members (role_id, member_id, entity_id) VALUES ($1, $2, $3)`, domainRoleID, domainUserID, domainID) - require.Nil(t, err, fmt.Sprintf("insert domains_role_members unexpected error: %s", err)) - _, err = db.Exec(`INSERT INTO domains_role_actions (role_id, action) VALUES ($1, $2)`, domainRoleID, "report_read") - require.Nil(t, err, fmt.Sprintf("insert domains_role_actions unexpected error: %s", err)) - - cases := []struct { - desc string - userID string - pageMeta reports.PageMeta - size int - err error - }{ - { - desc: "list user reports returns only accessible reports", - userID: userID, - pageMeta: reports.PageMeta{ - Domain: domainID, - Limit: 100, - Offset: 0, - }, - size: 5, - err: nil, - }, - { - desc: "list user reports with limit", - userID: userID, - pageMeta: reports.PageMeta{ - Domain: domainID, - Limit: 3, - Offset: 0, - }, - size: 3, - err: nil, - }, - { - desc: "list user reports with offset", - userID: userID, - pageMeta: reports.PageMeta{ - Domain: domainID, - Limit: 100, - Offset: 3, - }, - size: 2, - err: nil, - }, - { - desc: "list user reports with enabled status filter", - userID: userID, - pageMeta: reports.PageMeta{ - Domain: domainID, - Limit: 100, - Status: reports.EnabledStatus, - }, - size: 5, - err: nil, - }, - { - desc: "list reports for user with no role assignments returns 0", - userID: otherUserID, - pageMeta: reports.PageMeta{ - Domain: domainID, - Limit: 100, - Offset: 0, - }, - size: 0, - err: nil, - }, - { - desc: "list user reports via domain role returns all domain reports", - userID: domainUserID, - pageMeta: reports.PageMeta{ - Domain: domainID, - Limit: 100, - Offset: 0, - }, - size: 10, - err: nil, - }, - { - desc: "list user reports with non-existing domain returns 0", - userID: userID, - pageMeta: reports.PageMeta{ - Domain: generateUUID(t), - Limit: 100, - Offset: 0, - }, - size: 0, - err: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - page, err := repo.ListUserReportsConfig(context.Background(), tc.userID, tc.pageMeta) - if tc.err != nil { - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - return - } - require.Nil(t, err, fmt.Sprintf("unexpected error: %s", err)) - require.Equal(t, tc.size, len(page.ReportConfigs), fmt.Sprintf("%s: expected %d reports, got %d", tc.desc, tc.size, len(page.ReportConfigs))) - }) - } -} - func TestUpdateReportSchedule(t *testing.T) { t.Cleanup(func() { _, err := db.Exec("DELETE FROM report_config") diff --git a/reports/reports.go b/reports/reports.go index 3e4dca2b8..4e19e151c 100644 --- a/reports/reports.go +++ b/reports/reports.go @@ -14,7 +14,6 @@ import ( "github.com/absmach/magistrala/pkg/authn" "github.com/absmach/magistrala/pkg/errors" "github.com/absmach/magistrala/pkg/reltime" - "github.com/absmach/magistrala/pkg/roles" "github.com/absmach/magistrala/pkg/schedule" "github.com/absmach/magistrala/pkg/transformers/senml" ) @@ -167,16 +166,6 @@ type ReportConfig struct { CreatedBy string `json:"created_by,omitempty"` UpdatedAt time.Time `json:"updated_at"` UpdatedBy string `json:"updated_by,omitempty"` - // Extended - RoleID string `json:"role_id,omitempty"` - RoleName string `json:"role_name,omitempty"` - Actions []string `json:"actions,omitempty"` - AccessType string `json:"access_type,omitempty"` - AccessProviderId string `json:"access_provider_id,omitempty"` - AccessProviderRoleId string `json:"access_provider_role_id,omitempty"` - AccessProviderRoleName string `json:"access_provider_role_name,omitempty"` - AccessProviderRoleActions []string `json:"access_provider_role_actions,omitempty"` - Roles []roles.MemberRoleActions `json:"roles,omitempty"` } type ReportConfigPage struct { @@ -410,19 +399,16 @@ type PageMeta struct { type Repository interface { AddReportConfig(ctx context.Context, cfg ReportConfig) (ReportConfig, error) ViewReportConfig(ctx context.Context, id string) (ReportConfig, error) - RetrieveByIDWithRoles(ctx context.Context, id, memberID string) (ReportConfig, error) UpdateReportConfig(ctx context.Context, cfg ReportConfig) (ReportConfig, error) UpdateReportSchedule(ctx context.Context, cfg ReportConfig) (ReportConfig, error) RemoveReportConfig(ctx context.Context, id string) error UpdateReportConfigStatus(ctx context.Context, cfg ReportConfig) (ReportConfig, error) ListAllReportsConfig(ctx context.Context, pm PageMeta) (ReportConfigPage, error) - ListUserReportsConfig(ctx context.Context, userID string, pm PageMeta) (ReportConfigPage, error) UpdateReportDue(ctx context.Context, id string, due time.Time) (ReportConfig, error) UpdateReportTemplate(ctx context.Context, domainID, reportID string, template ReportTemplate) error ViewReportTemplate(ctx context.Context, domainID, reportID string) (ReportTemplate, error) DeleteReportTemplate(ctx context.Context, domainID, reportID string) error - roles.Repository } type Service interface { @@ -441,5 +427,4 @@ type Service interface { GenerateReport(ctx context.Context, session authn.Session, config ReportConfig, action ReportAction) (ReportPage, error) StartScheduler(ctx context.Context) error - roles.RoleManager } diff --git a/reports/service.go b/reports/service.go index 5c208bc7c..af900b64d 100644 --- a/reports/service.go +++ b/reports/service.go @@ -17,12 +17,9 @@ import ( "github.com/absmach/magistrala/pkg/errors" svcerr "github.com/absmach/magistrala/pkg/errors/service" pkglog "github.com/absmach/magistrala/pkg/logger" - "github.com/absmach/magistrala/pkg/policies" "github.com/absmach/magistrala/pkg/reltime" - "github.com/absmach/magistrala/pkg/roles" "github.com/absmach/magistrala/pkg/ticker" "github.com/absmach/magistrala/pkg/transformers/senml" - "github.com/absmach/magistrala/reports/operations" ) const limit = 1000 @@ -36,24 +33,18 @@ type report struct { readers grpcReadersV1.ReadersServiceClient defaultTemplate ReportTemplate converterURL string - roles.ProvisionManageService } -func NewService(repo Repository, runInfo chan pkglog.RunInfo, policy policies.Service, idp magistrala.IDProvider, tck ticker.Ticker, emailer emailer.Emailer, readers grpcReadersV1.ReadersServiceClient, template ReportTemplate, converterURL string, availableActions []roles.Action, builtInRoles map[roles.BuiltInRoleName][]roles.Action) (Service, error) { - rpms, err := roles.NewProvisionManageService(operations.EntityType, repo, policy, idp, availableActions, builtInRoles) - if err != nil { - return nil, err - } +func NewService(repo Repository, runInfo chan pkglog.RunInfo, idp magistrala.IDProvider, tck ticker.Ticker, emailer emailer.Emailer, readers grpcReadersV1.ReadersServiceClient, template ReportTemplate, converterURL string) (Service, error) { return &report{ - repo: repo, - idp: idp, - runInfo: runInfo, - email: emailer, - ticker: tck, - readers: readers, - defaultTemplate: template, - converterURL: converterURL, - ProvisionManageService: rpms, + repo: repo, + idp: idp, + runInfo: runInfo, + email: emailer, + ticker: tck, + readers: readers, + defaultTemplate: template, + converterURL: converterURL, }, nil } @@ -88,37 +79,11 @@ func (r *report) AddReportConfig(ctx context.Context, session authn.Session, cfg } }() - newBuiltInRoleMembers := map[roles.BuiltInRoleName][]roles.Member{ - BuiltInRoleAdmin: {roles.Member(session.UserID)}, - } - - optionalPolicies := []policies.Policy{ - { - SubjectType: policies.DomainType, - Subject: session.DomainID, - Relation: policies.DomainRelation, - ObjectType: operations.EntityType, - Object: reportConfig.ID, - }, - } - - _, err = r.AddNewEntitiesRoles(ctx, session.DomainID, session.UserID, []string{reportConfig.ID}, optionalPolicies, newBuiltInRoleMembers) - if err != nil { - return ReportConfig{}, errors.Wrap(svcerr.ErrAddPolicies, err) - } - return reportConfig, nil } func (r *report) ViewReportConfig(ctx context.Context, session authn.Session, id string, withRoles bool) (ReportConfig, error) { - var cfg ReportConfig - var err error - switch withRoles { - case true: - cfg, err = r.repo.RetrieveByIDWithRoles(ctx, id, session.UserID) - default: - cfg, err = r.repo.ViewReportConfig(ctx, id) - } + cfg, err := r.repo.ViewReportConfig(ctx, id) if err != nil { return ReportConfig{}, errors.Wrap(svcerr.ErrViewEntity, err) } @@ -159,14 +124,7 @@ func (r *report) RemoveReportConfig(ctx context.Context, session authn.Session, func (r *report) ListReportsConfig(ctx context.Context, session authn.Session, pm PageMeta) (ReportConfigPage, error) { pm.Domain = session.DomainID - if session.SuperAdmin { - page, err := r.repo.ListAllReportsConfig(ctx, pm) - if err != nil { - return ReportConfigPage{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - return page, nil - } - page, err := r.repo.ListUserReportsConfig(ctx, session.UserID, pm) + page, err := r.repo.ListAllReportsConfig(ctx, pm) if err != nil { return ReportConfigPage{}, errors.Wrap(svcerr.ErrViewEntity, err) } diff --git a/reports/service_test.go b/reports/service_test.go index 2cce103e2..175d220b5 100644 --- a/reports/service_test.go +++ b/reports/service_test.go @@ -18,7 +18,6 @@ import ( svcerr "github.com/absmach/magistrala/pkg/errors/service" pkglog "github.com/absmach/magistrala/pkg/logger" policymocks "github.com/absmach/magistrala/pkg/policies/mocks" - "github.com/absmach/magistrala/pkg/roles" pkgSch "github.com/absmach/magistrala/pkg/schedule" tmocks "github.com/absmach/magistrala/pkg/ticker/mocks" "github.com/absmach/magistrala/pkg/uuid" @@ -62,12 +61,7 @@ func newService(t *testing.T, runInfo chan pkglog.RunInfo) (reports.Service, *mo e := new(emocks.Emailer) policy := new(policymocks.Service) - availableActions := []roles.Action{} - builtInRoles := map[roles.BuiltInRoleName][]roles.Action{ - "admin": availableActions, - } - - svc, err := reports.NewService(repo, runInfo, policy, idProvider, mockTicker, e, readersSvc, template, "", availableActions, builtInRoles) + svc, err := reports.NewService(repo, runInfo, idProvider, mockTicker, e, readersSvc, template, "") if err != nil { t.Fatalf("Failed to create service: %v", err) } @@ -75,18 +69,14 @@ func newService(t *testing.T, runInfo chan pkglog.RunInfo) (reports.Service, *mo } func TestAddReportConfig(t *testing.T) { - svc, repo, _, policies := newService(t, make(chan pkglog.RunInfo)) + svc, repo, _, _ := newService(t, make(chan pkglog.RunInfo)) cases := []struct { - desc string - session authn.Session - cfg reports.ReportConfig - res reports.ReportConfig - err error - addPoliciesErr error - deletePolicies error - addRoleErr error - deleteErr error + desc string + session authn.Session + cfg reports.ReportConfig + res reports.ReportConfig + err error }{ { desc: "Add report config successfully", @@ -98,11 +88,8 @@ func TestAddReportConfig(t *testing.T) { Name: reportName, Schedule: schedule, }, - res: rptConfig, - err: nil, - addPoliciesErr: nil, - addRoleErr: nil, - deleteErr: nil, + res: rptConfig, + err: nil, }, { desc: "Add report config with failed repo", @@ -114,79 +101,13 @@ func TestAddReportConfig(t *testing.T) { Name: reportName, Schedule: schedule, }, - err: repoerr.ErrCreateEntity, - addPoliciesErr: nil, - deletePolicies: nil, - addRoleErr: nil, - deleteErr: nil, - }, - { - desc: "Add report config with failed to add policies", - session: authn.Session{ - UserID: userID, - DomainID: domainID, - }, - cfg: reports.ReportConfig{ - Name: reportName, - Schedule: schedule, - }, - res: rptConfig, - addPoliciesErr: svcerr.ErrAuthorization, - err: svcerr.ErrAddPolicies, - }, - { - desc: "Add report config with failed to add policies and failed rollback", - session: authn.Session{ - UserID: userID, - DomainID: domainID, - }, - cfg: reports.ReportConfig{ - Name: reportName, - Schedule: schedule, - }, - res: rptConfig, - addPoliciesErr: svcerr.ErrAuthorization, - deleteErr: svcerr.ErrRemoveEntity, - err: svcerr.ErrRollbackRepo, - }, - { - desc: "Add report config with failed to add roles", - session: authn.Session{ - UserID: userID, - DomainID: domainID, - }, - cfg: reports.ReportConfig{ - Name: reportName, - Schedule: schedule, - }, - res: rptConfig, - addRoleErr: svcerr.ErrCreateEntity, - err: svcerr.ErrAddPolicies, - }, - { - desc: "Add report config with failed to add roles and failed to delete policies", - session: authn.Session{ - UserID: userID, - DomainID: domainID, - }, - cfg: reports.ReportConfig{ - Name: reportName, - Schedule: schedule, - }, - res: rptConfig, - addRoleErr: svcerr.ErrCreateEntity, - deletePolicies: svcerr.ErrRemoveEntity, - err: svcerr.ErrRemoveEntity, + err: repoerr.ErrCreateEntity, }, } for _, tc := range cases { t.Run(tc.desc, func(t *testing.T) { repoCall := repo.On("AddReportConfig", mock.Anything, mock.Anything).Return(tc.res, tc.err) - policyCall := policies.On("AddPolicies", context.Background(), mock.Anything).Return(tc.addPoliciesErr) - policyCall2 := policies.On("DeletePolicies", context.Background(), mock.Anything).Return(tc.deletePolicies) - repoCall1 := repo.On("AddRoles", context.Background(), mock.Anything).Return([]roles.RoleProvision{}, tc.addRoleErr) - repoCall2 := repo.On("Remove", context.Background(), mock.Anything).Return(tc.deleteErr) res, err := svc.AddReportConfig(context.Background(), tc.session, tc.cfg) assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) if err == nil { @@ -194,15 +115,46 @@ func TestAddReportConfig(t *testing.T) { assert.Equal(t, tc.cfg.Name, res.Name) assert.Equal(t, tc.cfg.Schedule, res.Schedule) } - policyCall.Unset() - policyCall2.Unset() repoCall.Unset() - repoCall1.Unset() - repoCall2.Unset() }) } } +func TestAddReportConfigWithoutRoleProvisioning(t *testing.T) { + repo := new(mocks.Repository) + mockTicker := new(tmocks.Ticker) + idProvider := uuid.NewMock() + readersSvc := new(readmocks.ReadersServiceClient) + e := new(emocks.Emailer) + + svc, err := reports.NewService(repo, make(chan pkglog.RunInfo), idProvider, mockTicker, e, readersSvc, template, "") + if err != nil { + t.Fatalf("Failed to create service: %v", err) + } + + session := authn.Session{ + UserID: userID, + DomainID: domainID, + } + cfg := reports.ReportConfig{ + Name: reportName, + Schedule: schedule, + } + saved := cfg + saved.ID = rptConfig.ID + saved.Status = reports.EnabledStatus + saved.CreatedBy = userID + saved.DomainID = domainID + + repo.On("AddReportConfig", mock.Anything, mock.Anything).Return(saved, nil).Once() + + res, err := svc.AddReportConfig(context.Background(), session, cfg) + assert.NoError(t, err) + assert.Equal(t, saved.ID, res.ID) + repo.AssertNotCalled(t, "AddRoles", mock.Anything, mock.Anything) + repo.AssertExpectations(t) +} + func TestViewReportConfig(t *testing.T) { svc, repo, _, _ := newService(t, make(chan pkglog.RunInfo)) @@ -435,11 +387,7 @@ func TestListReportsConfig(t *testing.T) { for _, tc := range cases { t.Run(tc.desc, func(t *testing.T) { var repoCall *mock.Call - if tc.superAdmin { - repoCall = repo.On("ListAllReportsConfig", mock.Anything, mock.Anything).Return(tc.res, tc.err) - } else { - repoCall = repo.On("ListUserReportsConfig", mock.Anything, mock.Anything, mock.Anything).Return(tc.res, tc.err) - } + repoCall = repo.On("ListAllReportsConfig", mock.Anything, mock.Anything).Return(tc.res, tc.err) res, err := svc.ListReportsConfig(context.Background(), tc.session, tc.pageMeta) assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) if err == nil { diff --git a/scripts/re-backfill-roles/main.go b/scripts/re-backfill-roles/main.go deleted file mode 100644 index 565f4f447..000000000 --- a/scripts/re-backfill-roles/main.go +++ /dev/null @@ -1,536 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package main backfills missing built-in roles for rules. -package main - -import ( - "context" - "database/sql" - "fmt" - "io" - "log" - "os" - "strings" - - mglog "github.com/absmach/magistrala/logger" - "github.com/absmach/magistrala/pkg/errors" - "github.com/absmach/magistrala/pkg/policies" - "github.com/absmach/magistrala/pkg/policies/spicedb" - pgclient "github.com/absmach/magistrala/pkg/postgres" - "github.com/absmach/magistrala/pkg/roles" - spicedbdecoder "github.com/absmach/magistrala/pkg/spicedb" - "github.com/absmach/magistrala/pkg/uuid" - "github.com/absmach/magistrala/re" - "github.com/absmach/magistrala/re/operations" - repg "github.com/absmach/magistrala/re/postgres" - v1 "github.com/authzed/authzed-go/proto/authzed/api/v1" - "github.com/authzed/authzed-go/v1" - "github.com/authzed/grpcutil" - "go.opentelemetry.io/otel/trace/noop" - "google.golang.org/grpc" - "google.golang.org/grpc/credentials/insecure" -) - -const cmdName = "re_backfill_roles" - -var ( - logLevel = "info" - dryRun = false - limit = 0 - defaultMemberID = "" - spicedbHost = "localhost" - spicedbPort = "50051" - spicedbPreSharedKey = "12345678" - spicedbSchemaFile = "docker/spicedb/schema.zed" - dbConfig = pgclient.Config{ - Host: "localhost", - Port: "6009", - User: "magistrala", - Pass: "magistrala", - Name: "rules_engine", - SSLMode: "disable", - } -) - -type missingRule struct { - ID string `db:"id"` - Name string `db:"name"` - DomainID string `db:"domain_id"` - CreatedBy sql.NullString `db:"created_by"` -} - -func main() { - ctx := context.Background() - - if limit < 0 { - log.Fatalf("invalid limit %d: limit must be >= 0", limit) - } - - logger, err := mglog.New(os.Stdout, logLevel) - if err != nil { - log.Fatalf("failed to init logger: %s", err) - } - - var exitCode int - defer mglog.ExitWithError(&exitCode) - - sqlDB, err := pgclient.Connect(dbConfig) - if err != nil { - logger.Error("failed to connect to postgres", "error", err) - exitCode = 1 - return - } - defer sqlDB.Close() - - database := pgclient.NewDatabase(sqlDB, dbConfig, noop.NewTracerProvider().Tracer(cmdName)) - rulesRepo := repg.NewRepository(database) - - rulesWithoutRoles, err := listRulesWithoutRoles(ctx, database, limit) - if err != nil { - logger.Error("failed to list rules without roles", "error", err) - exitCode = 1 - return - } - - logger.Info("loaded rules without roles", "count", len(rulesWithoutRoles), "dry_run", dryRun) - if len(rulesWithoutRoles) == 0 { - return - } - - availableActions, builtInRoles, err := availableActionsAndBuiltInRoles(spicedbSchemaFile) - if err != nil { - logger.Error("failed to load built-in role actions", "error", err) - exitCode = 1 - return - } - - adminRoleActions, err := builtInRoleActionStrings(builtInRoles, re.BuiltInRoleAdmin) - if err != nil { - logger.Error("failed to resolve built-in admin role actions", "error", err) - exitCode = 1 - return - } - - authzedClient, err := newAuthzedClient(spicedbHost, spicedbPort, spicedbPreSharedKey) - if err != nil { - logger.Error("failed to connect to spicedb", "error", err) - exitCode = 1 - return - } - - if dryRun { - var processed, skipped int - - for _, rule := range rulesWithoutRoles { - memberID := strings.TrimSpace(rule.CreatedBy.String) - if memberID == "" { - memberID = strings.TrimSpace(defaultMemberID) - } - if rule.DomainID == "" { - skipped++ - logger.Warn("skipping rule without domain_id", "rule_id", rule.ID, "name", rule.Name) - continue - } - if memberID == "" { - skipped++ - logger.Warn("skipping rule without created_by and no default member override", "rule_id", rule.ID, "name", rule.Name) - continue - } - - isDomainMember, err := isDomainRoleMember(ctx, database, rule.DomainID, memberID) - if err != nil { - skipped++ - logger.Warn( - "skipping rule after failed domain membership check", - "rule_id", rule.ID, - "name", rule.Name, - "domain_id", rule.DomainID, - "member_id", memberID, - "error", err, - ) - continue - } - - candidatePolicies := []policies.Policy{ - { - SubjectType: policies.DomainType, - Subject: rule.DomainID, - Relation: policies.DomainRelation, - ObjectType: operations.EntityType, - Object: rule.ID, - }, - } - policiesToAdd, existingPolicies, err := filterMissingPolicies(ctx, authzedClient.PermissionsServiceClient, candidatePolicies) - if err != nil { - skipped++ - logger.Warn( - "skipping rule after failed spicedb policy lookup", - "rule_id", rule.ID, - "name", rule.Name, - "domain_id", rule.DomainID, - "error", err, - ) - continue - } - for _, existing := range existingPolicies { - logger.Info( - "dry run: spicedb policy already exists, will not be re-added", - "rule_id", rule.ID, - "subject_type", existing.SubjectType, - "subject", existing.Subject, - "relation", existing.Relation, - "object_type", existing.ObjectType, - "object", existing.Object, - ) - } - - if !isDomainMember { - logger.Warn( - "created_by user is not a member of the domain; role will be provisioned without member", - "rule_id", rule.ID, - "name", rule.Name, - "domain_id", rule.DomainID, - "member_id", memberID, - "created_by_exists_in_domain", false, - "role_actions", adminRoleActions, - ) - processed++ - logger.Info( - "dry run: would provision missing built-in role without member", - "rule_id", rule.ID, - "name", rule.Name, - "domain_id", rule.DomainID, - "member_id", memberID, - "created_by_exists_in_domain", false, - "role_actions", adminRoleActions, - "role_name", re.BuiltInRoleAdmin.String(), - "new_optional_policies", len(policiesToAdd), - "existing_optional_policies", len(existingPolicies), - ) - continue - } - - processed++ - logger.Info( - "dry run: would provision missing built-in role", - "rule_id", rule.ID, - "name", rule.Name, - "domain_id", rule.DomainID, - "member_id", memberID, - "created_by_exists_in_domain", true, - "role_actions", adminRoleActions, - "role_name", re.BuiltInRoleAdmin.String(), - "new_optional_policies", len(policiesToAdd), - "existing_optional_policies", len(existingPolicies), - ) - } - - logger.Info( - "backfill finished", - "processed", processed, - "skipped", skipped, - "failed", 0, - "dry_run", true, - ) - return - } - - policyService := spicedb.NewPolicyService(authzedClient, logger) - - provisioner, err := roles.NewProvisionManageService( - operations.EntityType, - rulesRepo, - policyService, - uuid.New(), - availableActions, - builtInRoles, - ) - if err != nil { - logger.Error("failed to create roles provisioner", "error", err) - exitCode = 1 - return - } - - var processed, skipped, failed int - - for _, rule := range rulesWithoutRoles { - memberID := strings.TrimSpace(rule.CreatedBy.String) - if memberID == "" { - memberID = strings.TrimSpace(defaultMemberID) - } - if rule.DomainID == "" { - skipped++ - logger.Warn("skipping rule without domain_id", "rule_id", rule.ID, "name", rule.Name) - continue - } - if memberID == "" { - skipped++ - logger.Warn("skipping rule without created_by and no default member override", "rule_id", rule.ID, "name", rule.Name) - continue - } - - assignMembers := []roles.Member{} - isDomainMember, err := isDomainRoleMember(ctx, database, rule.DomainID, memberID) - if err != nil { - failed++ - logger.Error( - "failed to check domain membership before provisioning role", - "rule_id", rule.ID, - "name", rule.Name, - "domain_id", rule.DomainID, - "member_id", memberID, - "error", err, - ) - continue - } - if isDomainMember { - assignMembers = []roles.Member{roles.Member(memberID)} - } else { - logger.Warn( - "created_by user is not a member of the domain; provisioning role without member", - "rule_id", rule.ID, - "name", rule.Name, - "domain_id", rule.DomainID, - "member_id", memberID, - "created_by_exists_in_domain", false, - "role_actions", adminRoleActions, - ) - } - - candidatePolicies := []policies.Policy{ - { - SubjectType: policies.DomainType, - Subject: rule.DomainID, - Relation: policies.DomainRelation, - ObjectType: operations.EntityType, - Object: rule.ID, - }, - } - optionalPolicies, existingPolicies, err := filterMissingPolicies(ctx, authzedClient.PermissionsServiceClient, candidatePolicies) - if err != nil { - failed++ - logger.Error( - "failed to check existing spicedb policies", - "rule_id", rule.ID, - "name", rule.Name, - "domain_id", rule.DomainID, - "error", err, - ) - continue - } - for _, existing := range existingPolicies { - logger.Info( - "spicedb policy already exists, skipping re-add", - "rule_id", rule.ID, - "subject_type", existing.SubjectType, - "subject", existing.Subject, - "relation", existing.Relation, - "object_type", existing.ObjectType, - "object", existing.Object, - ) - } - - newBuiltInRoleMembers := map[roles.BuiltInRoleName][]roles.Member{ - re.BuiltInRoleAdmin: assignMembers, - } - - if _, err := provisioner.AddNewEntitiesRoles( - ctx, - rule.DomainID, - memberID, - []string{rule.ID}, - optionalPolicies, - newBuiltInRoleMembers, - ); err != nil { - failed++ - logger.Error( - "failed to provision missing built-in role", - "rule_id", rule.ID, - "name", rule.Name, - "domain_id", rule.DomainID, - "member_id", memberID, - "error", err, - ) - continue - } - - processed++ - logger.Info( - "provisioned missing built-in role", - "rule_id", rule.ID, - "name", rule.Name, - "domain_id", rule.DomainID, - "member_id", memberID, - "created_by_exists_in_domain", isDomainMember, - "member_added", len(assignMembers) > 0, - "role_actions", adminRoleActions, - "role_name", re.BuiltInRoleAdmin.String(), - "new_optional_policies", len(optionalPolicies), - "existing_optional_policies", len(existingPolicies), - ) - } - - logger.Info( - "backfill finished", - "processed", processed, - "skipped", skipped, - "failed", failed, - "dry_run", dryRun, - ) - - if failed > 0 { - exitCode = 1 - } -} - -func listRulesWithoutRoles(ctx context.Context, db pgclient.Database, limit int) ([]missingRule, error) { - params := map[string]any{} - - query := ` - SELECT r.id, r.name, r.domain_id, r.created_by - FROM rules r - WHERE NOT EXISTS ( - SELECT 1 - FROM rules_roles rr - WHERE rr.entity_id = r.id - ) - ` - - query += " ORDER BY r.created_at ASC NULLS LAST, r.id ASC" - - if limit > 0 { - query += " LIMIT :limit" - params["limit"] = limit - } - - rows, err := db.NamedQueryContext(ctx, query, params) - if err != nil { - return nil, errors.Wrap(fmt.Errorf("failed to query rules without roles"), err) - } - defer rows.Close() - - var rules []missingRule - for rows.Next() { - var rule missingRule - if err := rows.StructScan(&rule); err != nil { - return nil, errors.Wrap(fmt.Errorf("failed to scan rule without role"), err) - } - rules = append(rules, rule) - } - if err := rows.Err(); err != nil { - return nil, errors.Wrap(fmt.Errorf("failed to iterate rules without roles"), err) - } - - return rules, nil -} - -func isDomainRoleMember(ctx context.Context, db pgclient.Database, domainID, memberID string) (bool, error) { - const query = ` - SELECT EXISTS ( - SELECT 1 - FROM domains_role_members drm - WHERE drm.entity_id = $1 AND drm.member_id = $2 - ) - ` - - var exists bool - if err := db.QueryRowxContext(ctx, query, domainID, memberID).Scan(&exists); err != nil { - return false, errors.Wrap(fmt.Errorf("failed to check domain role membership"), err) - } - - return exists, nil -} - -func newAuthzedClient(spicedbHost, spicedbPort, spicedbPreSharedKey string) (*authzed.ClientWithExperimental, error) { - return authzed.NewClientWithExperimentalAPIs( - fmt.Sprintf("%s:%s", spicedbHost, spicedbPort), - grpc.WithTransportCredentials(insecure.NewCredentials()), - grpcutil.WithInsecureBearerToken(spicedbPreSharedKey), - ) -} - -// filterMissingPolicies splits the given policies into those that do not yet -// exist in SpiceDB (returned first) and those that already exist (returned -// second). Any error from SpiceDB short-circuits with an empty result. -func filterMissingPolicies(ctx context.Context, permClient v1.PermissionsServiceClient, ps []policies.Policy) ([]policies.Policy, []policies.Policy, error) { - missing := make([]policies.Policy, 0, len(ps)) - existing := make([]policies.Policy, 0) - for _, p := range ps { - ok, err := policyExists(ctx, permClient, p) - if err != nil { - return nil, nil, err - } - if ok { - existing = append(existing, p) - continue - } - missing = append(missing, p) - } - return missing, existing, nil -} - -// policyExists returns true when SpiceDB already contains a relationship -// matching the supplied policy on (object_type, object, relation, subject_type, -// subject). The lookup is fully consistent and capped at one row. -func policyExists(ctx context.Context, permClient v1.PermissionsServiceClient, p policies.Policy) (bool, error) { - req := &v1.ReadRelationshipsRequest{ - Consistency: &v1.Consistency{ - Requirement: &v1.Consistency_FullyConsistent{FullyConsistent: true}, - }, - RelationshipFilter: &v1.RelationshipFilter{ - ResourceType: p.ObjectType, - OptionalResourceId: p.Object, - OptionalRelation: p.Relation, - OptionalSubjectFilter: &v1.SubjectFilter{ - SubjectType: p.SubjectType, - OptionalSubjectId: p.Subject, - }, - }, - OptionalLimit: 1, - } - - stream, err := permClient.ReadRelationships(ctx, req) - if err != nil { - return false, errors.Wrap(fmt.Errorf("failed to read spicedb relationships"), err) - } - - for { - _, err := stream.Recv() - switch { - case err == nil: - return true, nil - case errors.Contains(err, io.EOF): - return false, nil - default: - return false, errors.Wrap(fmt.Errorf("failed to receive spicedb relationship"), err) - } - } -} - -func availableActionsAndBuiltInRoles(spicedbSchemaFile string) ([]roles.Action, map[roles.BuiltInRoleName][]roles.Action, error) { - availableActions, err := spicedbdecoder.GetActionsFromSchema(spicedbSchemaFile, operations.EntityType) - if err != nil { - return []roles.Action{}, map[roles.BuiltInRoleName][]roles.Action{}, err - } - - builtInRoles := map[roles.BuiltInRoleName][]roles.Action{ - re.BuiltInRoleAdmin: availableActions, - } - - return availableActions, builtInRoles, nil -} - -func builtInRoleActionStrings(builtInRoles map[roles.BuiltInRoleName][]roles.Action, roleName roles.BuiltInRoleName) ([]string, error) { - actions, ok := builtInRoles[roleName] - if !ok { - return nil, fmt.Errorf("built-in role %q not found", roleName) - } - - ret := make([]string, 0, len(actions)) - for _, action := range actions { - ret = append(ret, action.String()) - } - - return ret, nil -} diff --git a/scripts/reports-backfill-roles/main.go b/scripts/reports-backfill-roles/main.go deleted file mode 100644 index d4e977495..000000000 --- a/scripts/reports-backfill-roles/main.go +++ /dev/null @@ -1,532 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package main backfills missing built-in roles for reports. -package main - -import ( - "context" - "database/sql" - "fmt" - "io" - "log" - "os" - "strings" - - mglog "github.com/absmach/magistrala/logger" - "github.com/absmach/magistrala/pkg/errors" - "github.com/absmach/magistrala/pkg/policies" - "github.com/absmach/magistrala/pkg/policies/spicedb" - pgclient "github.com/absmach/magistrala/pkg/postgres" - "github.com/absmach/magistrala/pkg/roles" - spicedbdecoder "github.com/absmach/magistrala/pkg/spicedb" - "github.com/absmach/magistrala/pkg/uuid" - "github.com/absmach/magistrala/reports" - "github.com/absmach/magistrala/reports/operations" - repg "github.com/absmach/magistrala/reports/postgres" - v1 "github.com/authzed/authzed-go/proto/authzed/api/v1" - "github.com/authzed/authzed-go/v1" - "github.com/authzed/grpcutil" - "go.opentelemetry.io/otel/trace/noop" - "google.golang.org/grpc" - "google.golang.org/grpc/credentials/insecure" -) - -const ( - cmdName = "reports_backfill_roles" -) - -var ( - logLevel = "info" - dryRun = false - limit = 0 - defaultMemberID = "" - spicedbHost = "localhost" - spicedbPort = "50051" - spicedbPreSharedKey = "12345678" - spicedbSchemaFile = "docker/spicedb/schema.zed" - dbConfig = pgclient.Config{ - Host: "localhost", - Port: "6020", - User: "magistrala", - Pass: "magistrala", - Name: "reports", - SSLMode: "disable", - } -) - -type missingReport struct { - ID string `db:"id"` - Name string `db:"name"` - DomainID string `db:"domain_id"` - CreatedBy sql.NullString `db:"created_by"` -} - -func main() { - ctx := context.Background() - - if limit < 0 { - log.Fatalf("invalid limit %d: limit must be >= 0", limit) - } - - logger, err := mglog.New(os.Stdout, logLevel) - if err != nil { - log.Fatalf("failed to init logger: %s", err) - } - - var exitCode int - defer mglog.ExitWithError(&exitCode) - - sqlDB, err := pgclient.Connect(dbConfig) - if err != nil { - logger.Error("failed to connect to postgres", "error", err) - exitCode = 1 - return - } - defer sqlDB.Close() - - database := pgclient.NewDatabase(sqlDB, dbConfig, noop.NewTracerProvider().Tracer(cmdName)) - reportsRepo := repg.NewRepository(database) - - reportsWithoutRoles, err := listReportsWithoutRoles(ctx, database, limit) - if err != nil { - logger.Error("failed to list reports without roles", "error", err) - exitCode = 1 - return - } - - logger.Info("loaded reports without roles", "count", len(reportsWithoutRoles), "dry_run", dryRun) - if len(reportsWithoutRoles) == 0 { - return - } - - availableActions, builtInRoles, err := availableActionsAndBuiltInRoles(spicedbSchemaFile) - if err != nil { - logger.Error("failed to load built-in role actions", "error", err) - exitCode = 1 - return - } - - adminRoleActions, err := builtInRoleActionStrings(builtInRoles, reports.BuiltInRoleAdmin) - if err != nil { - logger.Error("failed to resolve built-in admin role actions", "error", err) - exitCode = 1 - return - } - - authzedClient, err := newAuthzedClient(spicedbHost, spicedbPort, spicedbPreSharedKey) - if err != nil { - logger.Error("failed to connect to spicedb", "error", err) - exitCode = 1 - return - } - - if dryRun { - var processed, skipped int - - for _, report := range reportsWithoutRoles { - memberID := strings.TrimSpace(report.CreatedBy.String) - if memberID == "" { - memberID = strings.TrimSpace(defaultMemberID) - } - if report.DomainID == "" { - skipped++ - logger.Warn("skipping report without domain_id", "report_id", report.ID, "name", report.Name) - continue - } - if memberID == "" { - skipped++ - logger.Warn("skipping report without created_by and no default member override", "report_id", report.ID, "name", report.Name) - continue - } - - isDomainMember, err := isDomainRoleMember(ctx, database, report.DomainID, memberID) - if err != nil { - skipped++ - logger.Warn( - "skipping report after failed domain membership check", - "report_id", report.ID, - "name", report.Name, - "domain_id", report.DomainID, - "member_id", memberID, - "error", err, - ) - continue - } - - candidatePolicies := []policies.Policy{ - { - SubjectType: policies.DomainType, - Subject: report.DomainID, - Relation: policies.DomainRelation, - ObjectType: operations.EntityType, - Object: report.ID, - }, - } - policiesToAdd, existingPolicies, err := filterMissingPolicies(ctx, authzedClient.PermissionsServiceClient, candidatePolicies) - if err != nil { - skipped++ - logger.Warn( - "skipping report after failed spicedb policy lookup", - "report_id", report.ID, - "name", report.Name, - "domain_id", report.DomainID, - "error", err, - ) - continue - } - for _, existing := range existingPolicies { - logger.Info( - "dry run: spicedb policy already exists, will not be re-added", - "report_id", report.ID, - "subject_type", existing.SubjectType, - "subject", existing.Subject, - "relation", existing.Relation, - "object_type", existing.ObjectType, - "object", existing.Object, - ) - } - - if !isDomainMember { - logger.Warn( - "created_by user is not a member of the domain; role will be provisioned without member", - "report_id", report.ID, - "name", report.Name, - "domain_id", report.DomainID, - "member_id", memberID, - "created_by_exists_in_domain", false, - "role_actions", adminRoleActions, - ) - processed++ - logger.Info( - "dry run: would provision missing built-in role without member", - "report_id", report.ID, - "name", report.Name, - "domain_id", report.DomainID, - "member_id", memberID, - "created_by_exists_in_domain", false, - "role_actions", adminRoleActions, - "role_name", reports.BuiltInRoleAdmin.String(), - "new_optional_policies", len(policiesToAdd), - "existing_optional_policies", len(existingPolicies), - ) - continue - } - - processed++ - logger.Info( - "dry run: would provision missing built-in role", - "report_id", report.ID, - "name", report.Name, - "domain_id", report.DomainID, - "member_id", memberID, - "created_by_exists_in_domain", true, - "role_actions", adminRoleActions, - "role_name", reports.BuiltInRoleAdmin.String(), - "new_optional_policies", len(policiesToAdd), - "existing_optional_policies", len(existingPolicies), - ) - } - - logger.Info( - "backfill finished", - "processed", processed, - "skipped", skipped, - "failed", 0, - "dry_run", true, - ) - return - } - - policyService := spicedb.NewPolicyService(authzedClient, logger) - - provisioner, err := roles.NewProvisionManageService( - operations.EntityType, - reportsRepo, - policyService, - uuid.New(), - availableActions, - builtInRoles, - ) - if err != nil { - logger.Error("failed to create roles provisioner", "error", err) - exitCode = 1 - return - } - - var processed, skipped, failed int - - for _, report := range reportsWithoutRoles { - memberID := strings.TrimSpace(report.CreatedBy.String) - if memberID == "" { - memberID = strings.TrimSpace(defaultMemberID) - } - if report.DomainID == "" { - skipped++ - logger.Warn("skipping report without domain_id", "report_id", report.ID, "name", report.Name) - continue - } - if memberID == "" { - skipped++ - logger.Warn("skipping report without created_by and no default member override", "report_id", report.ID, "name", report.Name) - continue - } - - assignMembers := []roles.Member{} - isDomainMember, err := isDomainRoleMember(ctx, database, report.DomainID, memberID) - if err != nil { - failed++ - logger.Error( - "failed to check domain membership before provisioning role", - "report_id", report.ID, - "name", report.Name, - "domain_id", report.DomainID, - "member_id", memberID, - "error", err, - ) - continue - } - if isDomainMember { - assignMembers = []roles.Member{roles.Member(memberID)} - } else { - logger.Warn( - "created_by user is not a member of the domain; provisioning role without member", - "report_id", report.ID, - "name", report.Name, - "domain_id", report.DomainID, - "member_id", memberID, - "created_by_exists_in_domain", false, - "role_actions", adminRoleActions, - ) - } - - candidatePolicies := []policies.Policy{ - { - SubjectType: policies.DomainType, - Subject: report.DomainID, - Relation: policies.DomainRelation, - ObjectType: operations.EntityType, - Object: report.ID, - }, - } - optionalPolicies, existingPolicies, err := filterMissingPolicies(ctx, authzedClient.PermissionsServiceClient, candidatePolicies) - if err != nil { - failed++ - logger.Error( - "failed to check existing spicedb policies", - "report_id", report.ID, - "name", report.Name, - "domain_id", report.DomainID, - "error", err, - ) - continue - } - for _, existing := range existingPolicies { - logger.Info( - "spicedb policy already exists, skipping re-add", - "report_id", report.ID, - "subject_type", existing.SubjectType, - "subject", existing.Subject, - "relation", existing.Relation, - "object_type", existing.ObjectType, - "object", existing.Object, - ) - } - - newBuiltInRoleMembers := map[roles.BuiltInRoleName][]roles.Member{ - reports.BuiltInRoleAdmin: assignMembers, - } - - if _, err := provisioner.AddNewEntitiesRoles( - ctx, - report.DomainID, - memberID, - []string{report.ID}, - optionalPolicies, - newBuiltInRoleMembers, - ); err != nil { - failed++ - logger.Error( - "failed to provision missing built-in role", - "report_id", report.ID, - "name", report.Name, - "domain_id", report.DomainID, - "member_id", memberID, - "error", err, - ) - continue - } - - processed++ - logger.Info( - "provisioned missing built-in role", - "report_id", report.ID, - "name", report.Name, - "domain_id", report.DomainID, - "member_id", memberID, - "created_by_exists_in_domain", isDomainMember, - "member_added", len(assignMembers) > 0, - "role_actions", adminRoleActions, - "role_name", reports.BuiltInRoleAdmin.String(), - "new_optional_policies", len(optionalPolicies), - "existing_optional_policies", len(existingPolicies), - ) - } - - logger.Info( - "backfill finished", - "processed", processed, - "skipped", skipped, - "failed", failed, - "dry_run", dryRun, - ) - - if failed > 0 { - exitCode = 1 - } -} - -func listReportsWithoutRoles(ctx context.Context, db pgclient.Database, limit int) ([]missingReport, error) { - params := map[string]any{} - - query := ` - SELECT rc.id, rc.name, rc.domain_id, rc.created_by - FROM report_config rc - WHERE NOT EXISTS ( - SELECT 1 - FROM reports_roles rr - WHERE rr.entity_id = rc.id - ) - ` - - query += " ORDER BY rc.created_at ASC NULLS LAST, rc.id ASC" - - if limit > 0 { - query += " LIMIT :limit" - params["limit"] = limit - } - - rows, err := db.NamedQueryContext(ctx, query, params) - if err != nil { - return nil, errors.Wrap(fmt.Errorf("failed to query reports without roles"), err) - } - defer rows.Close() - - var reps []missingReport - for rows.Next() { - var rep missingReport - if err := rows.StructScan(&rep); err != nil { - return nil, errors.Wrap(fmt.Errorf("failed to scan report without role"), err) - } - reps = append(reps, rep) - } - if err := rows.Err(); err != nil { - return nil, errors.Wrap(fmt.Errorf("failed to iterate reports without roles"), err) - } - - return reps, nil -} - -func isDomainRoleMember(ctx context.Context, db pgclient.Database, domainID, memberID string) (bool, error) { - const query = ` - SELECT EXISTS ( - SELECT 1 - FROM domains_role_members drm - WHERE drm.entity_id = $1 AND drm.member_id = $2 - ) - ` - - var exists bool - if err := db.QueryRowxContext(ctx, query, domainID, memberID).Scan(&exists); err != nil { - return false, errors.Wrap(fmt.Errorf("failed to check domain role membership"), err) - } - - return exists, nil -} - -func newAuthzedClient(spicedbHost, spicedbPort, spicedbPreSharedKey string) (*authzed.ClientWithExperimental, error) { - return authzed.NewClientWithExperimentalAPIs( - fmt.Sprintf("%s:%s", spicedbHost, spicedbPort), - grpc.WithTransportCredentials(insecure.NewCredentials()), - grpcutil.WithInsecureBearerToken(spicedbPreSharedKey), - ) -} - -func filterMissingPolicies(ctx context.Context, permClient v1.PermissionsServiceClient, ps []policies.Policy) ([]policies.Policy, []policies.Policy, error) { - missing := make([]policies.Policy, 0, len(ps)) - existing := make([]policies.Policy, 0) - for _, p := range ps { - ok, err := policyExists(ctx, permClient, p) - if err != nil { - return nil, nil, err - } - if ok { - existing = append(existing, p) - continue - } - missing = append(missing, p) - } - return missing, existing, nil -} - -func policyExists(ctx context.Context, permClient v1.PermissionsServiceClient, p policies.Policy) (bool, error) { - req := &v1.ReadRelationshipsRequest{ - Consistency: &v1.Consistency{ - Requirement: &v1.Consistency_FullyConsistent{FullyConsistent: true}, - }, - RelationshipFilter: &v1.RelationshipFilter{ - ResourceType: p.ObjectType, - OptionalResourceId: p.Object, - OptionalRelation: p.Relation, - OptionalSubjectFilter: &v1.SubjectFilter{ - SubjectType: p.SubjectType, - OptionalSubjectId: p.Subject, - }, - }, - OptionalLimit: 1, - } - - stream, err := permClient.ReadRelationships(ctx, req) - if err != nil { - return false, errors.Wrap(fmt.Errorf("failed to read spicedb relationships"), err) - } - - for { - _, err := stream.Recv() - switch { - case err == nil: - return true, nil - case errors.Contains(err, io.EOF): - return false, nil - default: - return false, errors.Wrap(fmt.Errorf("failed to receive spicedb relationship"), err) - } - } -} - -func availableActionsAndBuiltInRoles(spicedbSchemaFile string) ([]roles.Action, map[roles.BuiltInRoleName][]roles.Action, error) { - availableActions, err := spicedbdecoder.GetActionsFromSchema(spicedbSchemaFile, operations.EntityType) - if err != nil { - return []roles.Action{}, map[roles.BuiltInRoleName][]roles.Action{}, err - } - - builtInRoles := map[roles.BuiltInRoleName][]roles.Action{ - reports.BuiltInRoleAdmin: availableActions, - } - - return availableActions, builtInRoles, nil -} - -func builtInRoleActionStrings(builtInRoles map[roles.BuiltInRoleName][]roles.Action, roleName roles.BuiltInRoleName) ([]string, error) { - actions, ok := builtInRoles[roleName] - if !ok { - return nil, fmt.Errorf("built-in role %q not found", roleName) - } - - ret := make([]string, 0, len(actions)) - for _, action := range actions { - ret = append(ret, action.String()) - } - - return ret, nil -} diff --git a/scripts/seed-test-data/TESTING.md b/scripts/seed-test-data/TESTING.md deleted file mode 100644 index bc610f090..000000000 --- a/scripts/seed-test-data/TESTING.md +++ /dev/null @@ -1,424 +0,0 @@ -# Backfill Roles — Testing Guide - -This document covers end-to-end testing of the migration scripts on the `migrations` branch: - -| Script | Purpose | -|--------|---------| -| `scripts/re-backfill-roles/` | Backfills missing built-in admin roles for **rules** (RE service) | -| `scripts/reports-backfill-roles/` | Backfills missing built-in admin roles for **reports** | -| `scripts/seed-test-data/` | Seeds all required test data across databases and SpiceDB | -| `domains/postgres/init.go` | Migration adding `alarm_*` and `report_*` actions to the domain admin role | - ---- - -## Prerequisites - -### Infrastructure - -Start the required containers: - -```bash -cd docker -docker compose up -d \ - spicedb-db spicedb-migrate spicedb \ - auth-db auth \ - domains-db domains \ - re-db re \ - reports-db reports \ - alarms-db alarms -``` - -Wait until all services are healthy. Each service applies its own Postgres migrations on startup, creating the required schemas. The `auth` service is required because it writes the SpiceDB schema on startup — without it, the seed script and backfill scripts fail with `object definition not found`. - -### Connection Details (docker-compose defaults) - -| Service | Host | Port | User | Password | Database | -|---------|------|------|------|----------|----------| -| Domains DB | localhost | 6003 | magistrala | magistrala | domains | -| RE DB | localhost | 6009 | magistrala | magistrala | rules_engine | -| Reports DB | localhost | 6020 | magistrala | magistrala | reports | -| Alarms DB | localhost | 6019 | magistrala | magistrala | alarms | -| SpiceDB gRPC | localhost | 50051 | — | 12345678 (pre-shared key) | — | - -### Fix Hard-Coded Configs (if needed) - -The backfill scripts have hard-coded database configs. Before running, verify they match your environment: - -**`scripts/re-backfill-roles/main.go` (lines 48–55):** - -```go -dbConfig = pgclient.Config{ - Host: "localhost", - Port: "6009", // docker-compose: 6009 (NOT 15432) - User: "magistrala", // docker-compose: magistrala (NOT postgres) - Pass: "magistrala", // docker-compose: magistrala (NOT supermq) - Name: "rules_engine", - SSLMode: "disable", -} -``` - -**`scripts/reports-backfill-roles/main.go` (lines 49–56):** - -```go -dbConfig = pgclient.Config{ - Host: "localhost", - Port: "6020", // docker-compose: 6020 (NOT 15432) - User: "magistrala", // docker-compose: magistrala (NOT postgres) - Pass: "magistrala", // docker-compose: magistrala (NOT supermq) - Name: "reports", - SSLMode: "disable", -} -``` - -**Both scripts — SpiceDB schema file (line 47):** - -```go -spicedbSchemaFile = "docker/spicedb/schema.zed" // NOT combined-schema.zed -``` - ---- - -## Step 1 — Seed Test Data - -```bash -go run ./scripts/seed-test-data/ -``` - -This inserts deterministic test data across all four databases and SpiceDB. It is idempotent (uses `ON CONFLICT DO NOTHING`), so re-running is safe. - -### What Gets Created - -**Domain:** - -| ID | Name | -|----|------| -| `d0000000-0000-0000-0000-000000000001` | seed-test-domain | - -**Users:** - -| ID | Domain Membership | -|----|-------------------| -| `u0000000-0000-0000-0000-000000000001` (user1) | Member of domain (in `domains_role_members`) | -| `u0000000-0000-0000-0000-000000000002` (user2) | NOT a domain member | - -**Rules (RE DB) — 6 rules, 4 orphans:** - -| Rule ID | Name | Scenario | -|---------|------|----------| -| `r0000000-...-000000000001` | rule-1-member-creator | Orphan. `created_by=user1` (domain member). Backfill should create role **with** member. | -| `r0000000-...-000000000002` | rule-2-nonmember-creator | Orphan. `created_by=user2` (NOT member). Backfill should create role **without** member. | -| `r0000000-...-000000000003` | rule-3-spicedb-exists | Orphan. `created_by=user1`. SpiceDB parent relation **pre-seeded**. Tests `policyExists` check. | -| `r0000000-...-000000000004` | rule-4-null-creator | Orphan. `created_by=NULL`. Should be **skipped**. | -| `r0000000-...-000000000005` | rule-5-no-domain | Orphan. `domain_id=""`. Should be **skipped**. | -| `r0000000-...-000000000006` | rule-6-has-role-already | Has `rules_roles` entry. Should **NOT appear** in orphan list. | - -**Reports (Reports DB) — 5 reports, 3 orphans:** - -| Report ID | Name | Scenario | -|-----------|------|----------| -| `rp000000-...-000000000001` | report-1-member-creator | Orphan. `created_by=user1`. Backfill should create role **with** member. | -| `rp000000-...-000000000002` | report-2-nonmember-creator | Orphan. `created_by=user2`. Backfill should create role **without** member. | -| `rp000000-...-000000000003` | report-3-spicedb-exists | Orphan. `created_by=user1`. SpiceDB parent **pre-seeded**. Tests `policyExists`. | -| `rp000000-...-000000000004` | report-4-null-creator | Orphan. `created_by=NULL`. Should be **skipped**. | -| `rp000000-...-000000000005` | report-5-has-role-already | Has `reports_roles` entry. Should **NOT appear**. | - -**Alarms (Alarms DB) — 2 alarms:** - -| Alarm ID | Linked Rule | -|----------|-------------| -| `a0000000-...-000000000001` | rule-1 | -| `a0000000-...-000000000002` | rule-2 | - -**SpiceDB (pre-seeded parent relations):** - -``` -rule:r0000000-...-000000000003#domain@domain:d0000000-...-000000000001 -report:rp000000-...-000000000003#domain@domain:d0000000-...-000000000001 -``` - ---- - -## Step 2 — Verify Seed Data - -### Check orphan rules in RE DB - -```bash -psql -h localhost -p 6009 -U magistrala -d rules_engine -c " - SELECT r.id, r.name, r.domain_id, r.created_by - FROM rules r - WHERE NOT EXISTS (SELECT 1 FROM rules_roles rr WHERE rr.entity_id = r.id) - ORDER BY r.name;" -``` - -**Expected:** 5 rows (rule-1 through rule-5). **rule-6 should NOT appear** (it has a role). - -### Check orphan reports in Reports DB - -```bash -psql -h localhost -p 6020 -U magistrala -d reports -c " - SELECT rc.id, rc.name, rc.domain_id, rc.created_by - FROM report_config rc - WHERE NOT EXISTS (SELECT 1 FROM reports_roles rr WHERE rr.entity_id = rc.id) - ORDER BY rc.name;" -``` - -**Expected:** 4 rows (report-1 through report-4). **report-5 should NOT appear**. - -### Check domain membership - -```bash -psql -h localhost -p 6009 -U magistrala -d rules_engine -c " - SELECT * FROM domains_role_members - WHERE entity_id = 'd0000000-0000-0000-0000-000000000001';" -``` - -**Expected:** 1 row for `user1`. `user2` should NOT be present. - -### Check SpiceDB pre-seeded relationships - -```bash -zed relationship read rule \ - --insecure --endpoint localhost:50051 --token 12345678 -``` - -**Expected:** At least one relationship for `rule:r0000000-...-000000000003#domain@domain:d0000000-...-000000000001`. - ---- - -## Step 3 — Test RE Backfill (Dry Run) - -Set `dryRun = true` in `scripts/re-backfill-roles/main.go`, then: - -```bash -go run ./scripts/re-backfill-roles/ -``` - -### Expected Log Output - -| Rule | Expected Log | -|------|-------------| -| rule-1 | `"dry run: would provision missing built-in role"` with `created_by_exists_in_domain=true` | -| rule-2 | `"created_by user is not a member of the domain"` + `"dry run: would provision missing built-in role without member"` | -| rule-3 | `"dry run: spicedb policy already exists, will not be re-added"` + `"dry run: would provision"` with `new_optional_policies=0, existing_optional_policies=1` | -| rule-4 | `"skipping rule without created_by and no default member override"` | -| rule-5 | `"skipping rule without domain_id"` | -| rule-6 | Does NOT appear at all | - -### Verify No Side Effects - -```bash -# Postgres: no new roles created -psql -h localhost -p 6009 -U magistrala -d rules_engine -c " - SELECT COUNT(*) FROM rules_roles - WHERE entity_id IN ( - 'r0000000-0000-0000-0000-000000000001', - 'r0000000-0000-0000-0000-000000000002', - 'r0000000-0000-0000-0000-000000000003' - );" -``` - -**Expected:** `0` (dry run should not write anything). - ---- - -## Step 4 — Test RE Backfill (Real Run) - -Set `dryRun = false` in `scripts/re-backfill-roles/main.go`, then: - -```bash -go run ./scripts/re-backfill-roles/ -``` - -### Expected Log Output - -| Rule | Expected Log | -|------|-------------| -| rule-1 | `"provisioned missing built-in role"` with `member_added=true` | -| rule-2 | `"provisioned missing built-in role"` with `member_added=false` | -| rule-3 | `"spicedb policy already exists, skipping re-add"` + `"provisioned missing built-in role"` with `new_optional_policies=0` | -| rule-4 | `"skipping rule without created_by"` | -| rule-5 | `"skipping rule without domain_id"` | - -Final summary should show: `processed=3, skipped=2, failed=0`. - -### Verify in Postgres - -```bash -# New roles exist for rules 1, 2, 3 -psql -h localhost -p 6009 -U magistrala -d rules_engine -c " - SELECT rr.id, rr.entity_id, rr.name, rr.created_by - FROM rules_roles rr - ORDER BY rr.entity_id;" - -# Role members: rule-1 should have user1; rule-2 and rule-3 check based on domain membership -psql -h localhost -p 6009 -U magistrala -d rules_engine -c " - SELECT rrm.role_id, rrm.member_id, rrm.entity_id - FROM rules_role_members rrm - ORDER BY rrm.entity_id;" -``` - -### Verify in SpiceDB - -```bash -zed relationship read rule \ - --insecure --endpoint localhost:50051 --token 12345678 -``` - -**Expected:** Parent relations for rule-1 and rule-2 are newly created. Rule-3 already had one (no duplicate). - ---- - -## Step 5 — Test Idempotency (Re-Run) - -Run the same backfill again without any changes: - -```bash -go run ./scripts/re-backfill-roles/ -``` - -**Expected:** `"loaded rules without roles" count=2` and `"backfill finished" processed=0, skipped=2, failed=0`. The two remaining rows are rule-4 (`created_by=NULL`) and rule-5 (no `domain_id`), which always re-appear in the orphan query and are skipped each run. The idempotency signal is `processed=0` — no roles or SpiceDB writes are duplicated. - ---- - -## Step 6 — Test Partial State (policyExists Path) - -This specifically validates the SpiceDB pre-check. Simulate a scenario where Postgres lost the role but SpiceDB still has the parent relation: - -```bash -# Delete just the Postgres role for rule-1 -psql -h localhost -p 6009 -U magistrala -d rules_engine -c " - DELETE FROM rules_roles - WHERE entity_id = 'r0000000-0000-0000-0000-000000000001';" - -# Re-run backfill -go run ./scripts/re-backfill-roles/ -``` - -### Expected - -- `"loaded rules without roles" count=1` (only rule-1 reappears) -- `"spicedb policy already exists, skipping re-add"` — the parent relation is detected and filtered out -- `"provisioned missing built-in role"` with `new_optional_policies=0, existing_optional_policies=1` -- The role row is re-created in Postgres **without** a duplicate SpiceDB write - -### Verify - -```bash -# Postgres: role restored -psql -h localhost -p 6009 -U magistrala -d rules_engine -c " - SELECT * FROM rules_roles - WHERE entity_id = 'r0000000-0000-0000-0000-000000000001';" - -# SpiceDB: still exactly one parent relation (no duplicate) -zed relationship read rule:r0000000-0000-0000-0000-000000000001 \ - --insecure --endpoint localhost:50051 --token 12345678 -``` - ---- - -## Step 7 — Test Reports Backfill - -Repeat Steps 3–6 for the reports backfill script: - -```bash -# Dry run (set dryRun = true first) -go run ./scripts/reports-backfill-roles/ - -# Real run (set dryRun = false) -go run ./scripts/reports-backfill-roles/ -``` - -### Expected behavior - -| Report | Expected | -|--------|----------| -| report-1 | Role provisioned **with** member (user1 is domain member) | -| report-2 | Role provisioned **without** member (user2 not in domain) | -| report-3 | `"spicedb policy already exists"` + role provisioned with `new_optional_policies=0` | -| report-4 | Skipped (NULL `created_by`) | -| report-5 | Does not appear (already has role) | - -### Verify - -```bash -psql -h localhost -p 6020 -U magistrala -d reports -c " - SELECT rr.id, rr.entity_id, rr.name - FROM reports_roles rr - ORDER BY rr.entity_id;" - -zed relationship read report \ - --insecure --endpoint localhost:50051 --token 12345678 -``` - ---- - -## Step 8 — Verify Domains Migration - -The `domains/postgres/init.go` change adds `alarm_*` and `report_*` actions to the domain admin role. This is applied by the domains service on startup. - -```bash -psql -h localhost -p 6003 -U magistrala -d domains -c " - SELECT action FROM domains_role_actions - WHERE role_id IN (SELECT id FROM domains_roles WHERE name = 'admin') - ORDER BY action;" -``` - -**Expected:** The result should include all of these new actions: - -``` -alarm_acknowledge -alarm_assign -alarm_delete -alarm_read -alarm_resolve -alarm_update -report_add_role_users -report_create -report_delete -report_manage_role -report_read -report_remove_role_users -report_update -report_view_role_users -``` - ---- - -## Step 9 — Cleanup (Optional) - -To reset and re-test from scratch: - -```bash -# Remove all seeded data from RE DB -psql -h localhost -p 6009 -U magistrala -d rules_engine -c " - DELETE FROM rules WHERE id LIKE 'r0000000-%'; - DELETE FROM domains WHERE id = 'd0000000-0000-0000-0000-000000000001';" - -# Remove all seeded data from Reports DB -psql -h localhost -p 6020 -U magistrala -d reports -c " - DELETE FROM report_config WHERE id LIKE 'rp000000-%'; - DELETE FROM domains WHERE id = 'd0000000-0000-0000-0000-000000000001';" - -# Remove all seeded data from Alarms DB -psql -h localhost -p 6019 -U magistrala -d alarms -c " - DELETE FROM alarms WHERE id LIKE 'a0000000-%'; - DELETE FROM domains WHERE id = 'd0000000-0000-0000-0000-000000000001';" - -# Remove SpiceDB relationships -zed relationship delete rule --insecure --endpoint localhost:50051 --token 12345678 -zed relationship delete report --insecure --endpoint localhost:50051 --token 12345678 -``` - -Then re-run `go run ./scripts/seed-test-data/` to start fresh. - ---- - -## Troubleshooting - -| Problem | Solution | -|---------|----------| -| `failed to connect to postgres` | Verify containers are running: `docker compose ps`. Check ports with `docker compose port re-db 5432`. | -| `failed to read spicedb relationships` | Ensure SpiceDB is running and schema is loaded. Check: `zed schema read --insecure --endpoint localhost:50051 --token 12345678` | -| `failed to load built-in role actions` | Verify `spicedbSchemaFile` points to `docker/spicedb/schema.zed` (not `combined-schema.zed`). | -| `no such table` errors during seed | Services haven't run yet to apply migrations. Start the full service (`re`, `reports`, `alarms`) at least once. | -| Script exits with `count=0` unexpectedly | All rules/reports already have roles. Check with the orphan queries from Step 2. | diff --git a/scripts/seed-test-data/main.go b/scripts/seed-test-data/main.go deleted file mode 100644 index d8b535f39..000000000 --- a/scripts/seed-test-data/main.go +++ /dev/null @@ -1,451 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package main seeds test data across the domains, RE, reports, and alarms -// databases so that the backfill-roles scripts can be tested end-to-end. -// -// It creates one domain, two users (one domain member, one not), several -// rules/reports with and without pre-existing roles, a couple of alarms, and -// optionally a SpiceDB parent relation for one rule to exercise the -// "policy already exists" path. -// -// All IDs are deterministic so re-running is idempotent (INSERT … ON CONFLICT DO NOTHING). -package main - -import ( - "context" - "database/sql" - "fmt" - "log" - "time" - - v1 "github.com/authzed/authzed-go/proto/authzed/api/v1" - "github.com/authzed/authzed-go/v1" - "github.com/authzed/grpcutil" - _ "github.com/jackc/pgx/v5/stdlib" - "google.golang.org/grpc" - "google.golang.org/grpc/credentials/insecure" -) - -// --------------------------------------------------------------------------- -// Configuration — edit these to match your environment. -// Default values target the standard docker-compose setup. -// --------------------------------------------------------------------------- - -var ( - // Domains database. - domainsDB = dbConfig{host: "localhost", port: "6003", user: "magistrala", pass: "magistrala", name: "domains"} - - // RE (rules engine) database. - reDB = dbConfig{host: "localhost", port: "6009", user: "magistrala", pass: "magistrala", name: "rules_engine"} - - // Reports database. - reportsDB = dbConfig{host: "localhost", port: "6020", user: "magistrala", pass: "magistrala", name: "reports"} - - // Alarms database. - alarmsDB = dbConfig{host: "localhost", port: "6019", user: "magistrala", pass: "magistrala", name: "alarms"} - - // SpiceDB. - spicedbHost = "localhost" - spicedbPort = "50051" - spicedbPreSharedKey = "12345678" - - // Whether to write a SpiceDB parent relation for rule-3 to test the - // "policy already exists" code path. - seedSpiceDB = true -) - -// --------------------------------------------------------------------------- -// Deterministic test IDs -// --------------------------------------------------------------------------- - -const ( - domainID = "d0000000-0000-0000-0000-000000000001" - domainName = "seed-test-domain" - - // user1 will be a domain member; user2 will NOT. - user1ID = "u0000000-0000-0000-0000-000000000001" - user2ID = "u0000000-0000-0000-0000-000000000002" - - // Domain role (admin). - domainRoleID = "dr000000-0000-0000-0000-000000000001" - domainRoleName = "admin" - - // Rules — orphans (no rules_roles entry). - rule1ID = "r0000000-0000-0000-0000-000000000001" // created_by=user1 (domain member) → backfill should assign member - rule2ID = "r0000000-0000-0000-0000-000000000002" // created_by=user2 (NOT member) → backfill should provision role without member - rule3ID = "r0000000-0000-0000-0000-000000000003" // created_by=user1, SpiceDB parent already exists → test policyExists - rule4ID = "r0000000-0000-0000-0000-000000000004" // created_by=NULL → should be skipped - rule5ID = "r0000000-0000-0000-0000-000000000005" // empty domain_id → should be skipped - - // Rule with pre-existing role — should NOT appear in orphan list. - rule6ID = "r0000000-0000-0000-0000-000000000006" - rule6RoleID = "rr000000-0000-0000-0000-000000000006" - - // Reports — orphans (no reports_roles entry). - report1ID = "rp000000-0000-0000-0000-000000000001" // created_by=user1 (domain member) - report2ID = "rp000000-0000-0000-0000-000000000002" // created_by=user2 (NOT member) - report3ID = "rp000000-0000-0000-0000-000000000003" // created_by=user1, SpiceDB parent already exists - report4ID = "rp000000-0000-0000-0000-000000000004" // created_by=NULL → skipped - - // Report with pre-existing role. - report5ID = "rp000000-0000-0000-0000-000000000005" - report5RoleID = "rpr00000-0000-0000-0000-000000000005" - - // Alarms (live in alarms DB). - alarm1ID = "a0000000-0000-0000-0000-000000000001" - alarm2ID = "a0000000-0000-0000-0000-000000000002" -) - -type dbConfig struct { - host, port, user, pass, name string -} - -func (c dbConfig) dsn() string { - return fmt.Sprintf("host=%s port=%s user=%s password=%s dbname=%s sslmode=disable", c.host, c.port, c.user, c.pass, c.name) -} - -func main() { - ctx := context.Background() - now := time.Now().UTC() - - // ----------------------------------------------------------------------- - // 1. Seed Domains DB - // ----------------------------------------------------------------------- - log.Println("connecting to domains DB ...") - ddb := mustConnect(domainsDB) - defer ddb.Close() - - seedDomainTables(ctx, ddb, now) - log.Println("domains DB seeded") - - // ----------------------------------------------------------------------- - // 2. Seed RE DB (includes domain tables + rules) - // ----------------------------------------------------------------------- - log.Println("connecting to RE DB ...") - rdb := mustConnect(reDB) - defer rdb.Close() - - seedDomainTables(ctx, rdb, now) - seedRules(ctx, rdb, now) - log.Println("RE DB seeded") - - // ----------------------------------------------------------------------- - // 3. Seed Reports DB (includes domain tables + report_config) - // ----------------------------------------------------------------------- - log.Println("connecting to reports DB ...") - rpdb := mustConnect(reportsDB) - defer rpdb.Close() - - seedDomainTables(ctx, rpdb, now) - seedReports(ctx, rpdb, now) - log.Println("reports DB seeded") - - // ----------------------------------------------------------------------- - // 4. Seed Alarms DB (includes domain + RE tables + alarms) - // ----------------------------------------------------------------------- - log.Println("connecting to alarms DB ...") - adb := mustConnect(alarmsDB) - defer adb.Close() - - seedDomainTables(ctx, adb, now) - seedRulesMinimal(ctx, adb, now) // alarms DB has rules tables via RE migration - seedAlarms(ctx, adb, now) - log.Println("alarms DB seeded") - - // ----------------------------------------------------------------------- - // 5. Optionally seed SpiceDB (parent relations for rule3 and report3) - // ----------------------------------------------------------------------- - if seedSpiceDB { - log.Println("connecting to SpiceDB ...") - seedSpiceDBRelationships(ctx) - log.Println("SpiceDB seeded") - } - - log.Println("all seed data inserted successfully") - printSummary() -} - -// --------------------------------------------------------------------------- -// Domain tables (identical across all DBs that include domain migrations) -// --------------------------------------------------------------------------- - -func seedDomainTables(ctx context.Context, db *sql.DB, now time.Time) { - mustExec(ctx, db, ` - INSERT INTO domains (id, name, tags, metadata, route, created_at, updated_at, created_by, status) - VALUES ($1, $2, '{}', '{}', $3, $4, $4, $5, 0) - ON CONFLICT (id) DO NOTHING`, - domainID, domainName, "seed-test-domain", now, user1ID) - - // Domain admin role - mustExec(ctx, db, ` - INSERT INTO domains_roles (id, name, entity_id, created_at, updated_at, created_by) - VALUES ($1, $2, $3, $4, $4, $5) - ON CONFLICT (id) DO NOTHING`, - domainRoleID, domainRoleName, domainID, now, user1ID) - - // Domain role actions (a representative subset) - actions := []string{ - "domain_update", "domain_read", "domain_membership", - "domain_manage_role", "domain_add_role_users", "domain_remove_role_users", "domain_view_role_users", - "rule_create", "rule_read", "rule_update", "rule_delete", - "rule_manage_role", "rule_add_role_users", "rule_remove_role_users", "rule_view_role_users", - "report_create", "report_read", "report_update", "report_delete", - "report_manage_role", "report_add_role_users", "report_remove_role_users", "report_view_role_users", - "alarm_update", "alarm_read", "alarm_delete", "alarm_assign", "alarm_acknowledge", "alarm_resolve", - } - for _, action := range actions { - mustExec(ctx, db, ` - INSERT INTO domains_role_actions (role_id, action) - VALUES ($1, $2) - ON CONFLICT DO NOTHING`, - domainRoleID, action) - } - - // user1 is a domain member; user2 is NOT. - mustExec(ctx, db, ` - INSERT INTO domains_role_members (role_id, member_id, entity_id) - VALUES ($1, $2, $3) - ON CONFLICT DO NOTHING`, - domainRoleID, user1ID, domainID) -} - -// --------------------------------------------------------------------------- -// Rules (RE DB) -// --------------------------------------------------------------------------- - -func seedRules(ctx context.Context, db *sql.DB, now time.Time) { - type rule struct { - id, name, domainID string - createdBy *string - } - - u1 := strPtr(user1ID) - u2 := strPtr(user2ID) - - rules := []rule{ - {rule1ID, "rule-1-member-creator", domainID, u1}, - {rule2ID, "rule-2-nonmember-creator", domainID, u2}, - {rule3ID, "rule-3-spicedb-exists", domainID, u1}, - {rule4ID, "rule-4-null-creator", domainID, nil}, - {rule5ID, "rule-5-no-domain", "", u1}, - {rule6ID, "rule-6-has-role-already", domainID, u1}, - } - - for _, r := range rules { - mustExec(ctx, db, ` - INSERT INTO rules (id, name, domain_id, created_by, created_at, status, logic_type) - VALUES ($1, $2, $3, $4, $5, 0, 0) - ON CONFLICT (id) DO NOTHING`, - r.id, r.name, r.domainID, r.createdBy, now) - } - - // rule6 already has a role → should NOT appear in orphan list. - mustExec(ctx, db, ` - INSERT INTO rules_roles (id, name, entity_id, created_at, updated_at, created_by) - VALUES ($1, 'admin', $2, $3, $3, $4) - ON CONFLICT (id) DO NOTHING`, - rule6RoleID, rule6ID, now, user1ID) - - mustExec(ctx, db, ` - INSERT INTO rules_role_actions (role_id, action) - VALUES ($1, 'rule_read') - ON CONFLICT DO NOTHING`, - rule6RoleID) - - mustExec(ctx, db, ` - INSERT INTO rules_role_members (role_id, member_id, entity_id) - VALUES ($1, $2, $3) - ON CONFLICT DO NOTHING`, - rule6RoleID, user1ID, rule6ID) -} - -// seedRulesMinimal inserts the same rules into the alarms DB (which has RE -// tables) so that foreign key constraints on rule_id can be satisfied. -func seedRulesMinimal(ctx context.Context, db *sql.DB, now time.Time) { - for _, r := range []struct{ id, name string }{ - {rule1ID, "rule-1-member-creator"}, - {rule2ID, "rule-2-nonmember-creator"}, - } { - mustExec(ctx, db, ` - INSERT INTO rules (id, name, domain_id, created_by, created_at, status, logic_type) - VALUES ($1, $2, $3, $4, $5, 0, 0) - ON CONFLICT (id) DO NOTHING`, - r.id, r.name, domainID, user1ID, now) - } -} - -// --------------------------------------------------------------------------- -// Reports -// --------------------------------------------------------------------------- - -func seedReports(ctx context.Context, db *sql.DB, now time.Time) { - type report struct { - id, name, domainID string - createdBy *string - } - - u1 := strPtr(user1ID) - u2 := strPtr(user2ID) - - reports := []report{ - {report1ID, "report-1-member-creator", domainID, u1}, - {report2ID, "report-2-nonmember-creator", domainID, u2}, - {report3ID, "report-3-spicedb-exists", domainID, u1}, - {report4ID, "report-4-null-creator", domainID, nil}, - {report5ID, "report-5-has-role-already", domainID, u1}, - } - - for _, r := range reports { - mustExec(ctx, db, ` - INSERT INTO report_config (id, name, domain_id, created_by, created_at, status) - VALUES ($1, $2, $3, $4, $5, 0) - ON CONFLICT (id) DO NOTHING`, - r.id, r.name, r.domainID, r.createdBy, now) - } - - // report5 already has a role. - mustExec(ctx, db, ` - INSERT INTO reports_roles (id, name, entity_id, created_at, updated_at, created_by) - VALUES ($1, 'admin', $2, $3, $3, $4) - ON CONFLICT (id) DO NOTHING`, - report5RoleID, report5ID, now, user1ID) - - mustExec(ctx, db, ` - INSERT INTO reports_role_actions (role_id, action) - VALUES ($1, 'report_read') - ON CONFLICT DO NOTHING`, - report5RoleID) - - mustExec(ctx, db, ` - INSERT INTO reports_role_members (role_id, member_id, entity_id) - VALUES ($1, $2, $3) - ON CONFLICT DO NOTHING`, - report5RoleID, user1ID, report5ID) -} - -// --------------------------------------------------------------------------- -// Alarms -// --------------------------------------------------------------------------- - -func seedAlarms(ctx context.Context, db *sql.DB, now time.Time) { - for _, a := range []struct{ id, ruleID string }{ - {alarm1ID, rule1ID}, - {alarm2ID, rule2ID}, - } { - mustExec(ctx, db, ` - INSERT INTO alarms (id, rule_id, domain_id, channel_id, subtopic, client_id, - measurement, value, unit, threshold, cause, status, severity, created_at) - VALUES ($1, $2, $3, 'ch000000-0000-0000-0000-000000000001', 'test/topic', - 'cl000000-0000-0000-0000-000000000001', 'temperature', '42.5', 'C', '40.0', - 'exceeded threshold', 0, 1, $4) - ON CONFLICT (id) DO NOTHING`, - a.id, a.ruleID, domainID, now) - } -} - -// --------------------------------------------------------------------------- -// SpiceDB — write a parent relation for rule3 and report3 so the -// "policy already exists" code path is exercised. -// --------------------------------------------------------------------------- - -func seedSpiceDBRelationships(ctx context.Context) { - addr := fmt.Sprintf("%s:%s", spicedbHost, spicedbPort) - client, err := authzed.NewClientWithExperimentalAPIs( - addr, - grpc.WithTransportCredentials(insecure.NewCredentials()), - grpcutil.WithInsecureBearerToken(spicedbPreSharedKey), - ) - if err != nil { - log.Printf("WARNING: failed to connect to SpiceDB at %s: %v (skipping SpiceDB seed)", addr, err) - return - } - - // Use TOUCH so re-running is idempotent. - updates := []*v1.RelationshipUpdate{ - { - Operation: v1.RelationshipUpdate_OPERATION_TOUCH, - Relationship: &v1.Relationship{ - Resource: &v1.ObjectReference{ObjectType: "rule", ObjectId: rule3ID}, - Relation: "domain", - Subject: &v1.SubjectReference{Object: &v1.ObjectReference{ObjectType: "domain", ObjectId: domainID}}, - }, - }, - { - Operation: v1.RelationshipUpdate_OPERATION_TOUCH, - Relationship: &v1.Relationship{ - Resource: &v1.ObjectReference{ObjectType: "report", ObjectId: report3ID}, - Relation: "domain", - Subject: &v1.SubjectReference{Object: &v1.ObjectReference{ObjectType: "domain", ObjectId: domainID}}, - }, - }, - } - - _, err = client.WriteRelationships(ctx, &v1.WriteRelationshipsRequest{Updates: updates}) - if err != nil { - log.Printf("WARNING: failed to write SpiceDB relationships: %v", err) - return - } - log.Printf("wrote %d SpiceDB relationships (TOUCH)", len(updates)) -} - -// --------------------------------------------------------------------------- -// Helpers -// --------------------------------------------------------------------------- - -func mustConnect(cfg dbConfig) *sql.DB { - db, err := sql.Open("pgx", cfg.dsn()) - if err != nil { - log.Fatalf("failed to open %s: %v", cfg.name, err) - } - ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) - if err := db.PingContext(ctx); err != nil { - cancel() - log.Fatalf("failed to ping %s at %s:%s: %v", cfg.name, cfg.host, cfg.port, err) - } - cancel() - return db -} - -func mustExec(ctx context.Context, db *sql.DB, query string, args ...any) { - if _, err := db.ExecContext(ctx, query, args...); err != nil { - log.Fatalf("exec failed: %v\nquery: %s\nargs: %v", err, query, args) - } -} - -func strPtr(s string) *string { return &s } - -func printSummary() { - fmt.Print(` -=== SEED DATA SUMMARY === - -Domain: ` + domainID + ` ("` + domainName + `") - -Users: - user1 (domain member): ` + user1ID + ` - user2 (NOT a member): ` + user2ID + ` - -Rules (RE DB): - ` + rule1ID + ` rule-1-member-creator orphan, created_by=user1 → expect role WITH member - ` + rule2ID + ` rule-2-nonmember-creator orphan, created_by=user2 → expect role WITHOUT member - ` + rule3ID + ` rule-3-spicedb-exists orphan, created_by=user1, SpiceDB parent pre-seeded → test policyExists - ` + rule4ID + ` rule-4-null-creator orphan, created_by=NULL → expect SKIPPED - ` + rule5ID + ` rule-5-no-domain orphan, domain_id="" → expect SKIPPED - ` + rule6ID + ` rule-6-has-role-already HAS role entry → should NOT appear in orphan list - -Reports (Reports DB): - ` + report1ID + ` report-1-member-creator orphan, created_by=user1 → expect role WITH member - ` + report2ID + ` report-2-nonmember-creator orphan, created_by=user2 → expect role WITHOUT member - ` + report3ID + ` report-3-spicedb-exists orphan, created_by=user1, SpiceDB parent pre-seeded → test policyExists - ` + report4ID + ` report-4-null-creator orphan, created_by=NULL → expect SKIPPED - ` + report5ID + ` report-5-has-role-already HAS role entry → should NOT appear in orphan list - -Alarms (Alarms DB): - ` + alarm1ID + ` alarm-1 (rule1) - ` + alarm2ID + ` alarm-2 (rule2) - -SpiceDB: - rule:` + rule3ID + `#domain@domain:` + domainID + ` (pre-seeded) - report:` + report3ID + `#domain@domain:` + domainID + ` (pre-seeded) -`) -} diff --git a/tools/atom-migration/Dockerfile b/tools/atom-migration/Dockerfile new file mode 100644 index 000000000..a39570ce9 --- /dev/null +++ b/tools/atom-migration/Dockerfile @@ -0,0 +1,18 @@ +# Copyright (c) Abstract Machines +# SPDX-License-Identifier: Apache-2.0 + +# Build a standalone atom-migration binary so the tool can be run without +# `go run` and without the Go toolchain on the host. +FROM golang:1.26-alpine AS build +WORKDIR /src +COPY go.mod go.sum ./ +RUN go mod download +COPY . . +RUN CGO_ENABLED=0 go build -o /atom-migration ./tools/atom-migration + +FROM alpine:3.22 +RUN adduser -D -u 10001 migrator +WORKDIR /work +COPY --from=build /atom-migration /usr/local/bin/atom-migration +USER migrator +ENTRYPOINT ["atom-migration"] diff --git a/tools/atom-migration/PLAN.md b/tools/atom-migration/PLAN.md new file mode 100644 index 000000000..90d430a9f --- /dev/null +++ b/tools/atom-migration/PLAN.md @@ -0,0 +1,363 @@ +# Magistrala v0.30.0 → Atom IAM Migration Plan + +Offline, one-shot migration that reads the per-service Magistrala Postgres +databases (default Docker Compose deployment) and writes a single Atom Postgres +database. Implemented as Go scripts. + +--- + +## 1. Decisions (locked) + +| Topic | Decision | +|-------|----------| +| Passwords | **Force reset.** Users migrate with no `password` credential; they reset via Atom's email flow on first login. (bcrypt → argon2 is not convertible without plaintext.) | +| Scope | Core IAM + roles & policies + connections + PATs + rules/reports/alarms as Atom resources. | +| Execution | **Offline one-shot.** Stop Magistrala app services, snapshot, transform, load, start Atom. | +| IDs | **Preserve Magistrala UUIDs** as Atom UUIDs (PKs and FKs). Magistrala IDs are 36-char UUID strings — directly usable as Atom `UUID` PKs. Keeps audit trails, message payloads, external references, and SpiceDB-derived links intact. | + +--- + +## 2. Source vs target topology + +### Source (Magistrala, default compose — separate DB container per service) + +All `magistrala/magistrala`, port 5432, on network `magistrala-base-net`: + +| Container | DB name | Tables we read | +|-----------|---------|----------------| +| `domains-db` | `domains` | `domains`, `invitations`, `domains_roles`, `domains_role_actions`, `domains_role_members` | +| `users-db` | `users` | `users`, `users_verifications` | +| `clients-db` | `clients` | `clients`, `connections`, `clients_roles*`, plus its embedded `groups` copy | +| `channels-db` | `channels` | `channels`, `connections`, `channels_roles*` | +| `groups-db` | `groups` | `groups`, `groups_roles*` | +| `auth-db` | `auth` | `pats`, `pat_scopes` (skip `keys` — short-lived JWTs; skip legacy `policies`/`domains` mirror) | +| `re-db` | `rules_engine` | `rules`, `rules_roles*` | +| `reports-db` | `reports` | `report_config`, `reports_roles*` | +| `alarms-db` | `alarms` | `alarms` | + +> Note: `groups` migrations are embedded into clients **and** channels **and** the +> standalone groups service. The authoritative groups data for default compose is +> the `groups-db`/`groups` database — read groups from there. + +### Target (Atom — single Postgres) + +`atom` DB, schema from the squashed `migrations/001_initial.sql`. Relevant tables: +`tenants, entities, credentials, entity_emails, resources, principal_groups, +object_groups, *_group_hierarchy, object_group_entities, object_group_resources, +tenant_memberships, roles, permission_blocks, permission_block_actions, +role_permission_blocks, role_assignments, direct_policies, profiles, +profile_versions`. + +Seeded fixed UUIDs already present in Atom (do **not** collide): +`...0001` admin entity, `...0002` atom-admin role, `...0003` mg-service entity, +`...0004` mg-service role, `...0005` authenticated-users group, +`...0006` domain-creator role, `...0007/8/9` permission blocks. + +--- + +## 3. Entity mapping + +### 3.1 domains → tenants +| Magistrala `domains` | Atom `tenants` | +|---|---| +| `id` | `id` (preserved) | +| `name` | `name` | +| `route` | `alias` (must be slug: lowercase, `^[a-z0-9]([a-z0-9-]{0,61}[a-z0-9])?$`, **not** UUID-shaped — see §6) | +| `tags` | `tags` | +| `metadata` | `attributes` | +| `status` 0/1 | `status` `active`/`inactive` | +| `created_at/by`, `updated_at/by` | same (created_by/updated_by only if the user id also migrates) | + +### 3.2 users → entities (kind=`human`) +| Magistrala `users` | Atom | +|---|---| +| `id` | `entities.id` (preserved), `kind='human'`, `tenant_id=NULL` (users are global; domain membership handled in §4), `profile_id` / `profile_version_id` = seeded active `user` profile | +| `first_name,last_name,username,profile_picture,auth_provider` | `entities.attributes` JSON | +| `metadata` | merged into `attributes` | +| `email` | `entity_emails.email`; `verified_at` → `entity_emails.verified_at` | +| `secret` (bcrypt) | **dropped** — no credential created (force-reset) | +| `status` 0/1 | `entities.status` `active`/`inactive` | +| `role` (0 user / 1 admin) | if admin → also assign Atom `atom-admin` role (`role_assignments`) | + +Every migrated human is also inserted into Atom's seeded `authenticated-users` +principal group, matching Atom's normal `create_entity` side effect. + +`users_verifications` → not migrated (transient OTPs). + +### 3.3 clients (things/devices) → entities (kind=`device`) +| Magistrala `clients` | Atom | +|---|---| +| `id` | `entities.id` (preserved), `kind='device'`, `tenant_id=domain_id`, `profile_id` / `profile_version_id` = seeded active `client` profile | +| `name` | `entities.name` (per-tenant unique — handle collisions, §6) | +| `tags,metadata,private_metadata` | `attributes` | +| `identity` | `attributes.identity` (and/or `entities.alias` if slug-valid) | +| `secret` (plaintext) | `credentials` row: `kind='api_key'`, `entity_id=client.id`, `secret_hash`=argon2(secret), `identifier`=client.id, `status` per client status. **See §5 — key-format caveat.** | +| `status` 0/1 | `entities.status` | +| `parent_group_id` | `object_group_entities` membership (§3.5) | + +### 3.4 channels → resources (kind=`channel`) +| Magistrala `channels` | Atom `resources` | +|---|---| +| `id` | `id` (preserved), `kind='channel'`, `tenant_id=domain_id` | +| `name` | `name` | +| `route` | `alias` (slug rules, §6) | +| `tags,metadata` | `attributes` | +| `created_by` | `owner_id` (if the user migrated) | +| `parent_group_id` | `object_group_resources` membership | + +### 3.4b rules / reports / alarms → resources + +Rules-engine rules, report configs, and alarms are domain-scoped objects with no +Atom-native table, so they become **resources** alongside channels, distinguished +by `kind`. Service-specific columns Atom resources lack are folded into +`attributes` (JSONB). `tenant_id=domain_id`; rows whose domain has no surviving +tenant are skipped (§6.4). `owner_id`=`created_by` when that user migrated +(alarms have no `created_by` → NULL). Names are deduped per tenant (§6.3) since +these tables carry no `(domain_id, name)` constraint; alarms (no name) use +`measurement` with an `id` fallback. + +| Source | Atom `resources` | +|---|---| +| `rules_engine.rules` (id) | `kind='rule'`; `input_channel/topic, outputs, logic_type/value, recurring*, time, start_datetime, tags, status` → `attributes` | +| `reports.report_config` (id) | `kind='report'`; `description, config, email, metrics, report_template, due, recurring*, start_datetime, status` → `attributes` | +| `alarms.alarms` (id) | `kind='alarm'`; `rule_id, channel_id, client_id, subtopic, measurement, value, unit, threshold, cause, severity, alarm_status, assignee/assigned/acknowledged/resolved*` → `attributes` | + +> Atom `resources.kind` must permit `rule`, `report`, `alarm` (in addition to +> `channel`). Rules and reports have object-specific role families and those are +> migrated as resource-scoped roles. Alarms have no role family or +> `parent_group_id`, so only the resource rows are migrated for alarms. + +### 3.5 groups → object_groups +Magistrala groups organize clients/channels within a domain (hierarchical, +`parent_id` + `path` ltree). Map to **object_groups**: +| Magistrala `groups` | Atom `object_groups` | +|---|---| +| `id,name,description,metadata,status` | same fields (`metadata`→`attributes`, status mapped) | +| `domain_id` | `tenant_id` | +| `parent_id` | `object_group_hierarchy(parent_id, child_id, tenant_id)` | +Client/channel `parent_group_id` → `object_group_entities` / `object_group_resources`. + +> Confirmed: Magistrala groups have no user-membership table — only +> `parent_group_id` (clients/channels) and group-scoped roles. They are always +> object groupings → `object_groups`. There is no principal-group case to handle. + +--- + +## 4. Roles, policies, memberships + +Magistrala authz = per-service role tables (`_roles`, `_role_actions`, +`_role_members`) enforced via SpiceDB. Atom replaces SpiceDB; we reconstruct authz +**from the SQL role tables** (no need to read SpiceDB directly). Migrated role +families: `domains`, `clients`, `channels`, `groups`, `rules`, and `reports`. + +Magistrala role shape: each role row is bound to a specific object instance +(`entity_id` = the domain/client/channel/group id), has a set of action strings, +and a set of member ids (users). + +Mapping per Magistrala role row → Atom: +1. **`roles`** row (preserve `id`, `name`, `tenant_id` = owning domain when known). +2. **`permission_blocks`** row scoped to the object: + - domain roles → `scope_mode='tenant'`, `tenant_id=entity_id` + - client/channel/rule/report roles → `scope_mode='object'`, `object_id=entity_id` + - group roles → supported Atom group scopes: + direct `client*` / `channel*` actions use `group_direct_objects`; + `subgroup_client*` / `subgroup_channel*` actions use + `group_descendant_objects`; `subgroup*` group actions use + `group_descendant_groups`; direct group-management actions use object scope + on the object group itself. + - `effect='allow'`, `conditions='{}'` +3. **`permission_block_actions`** ← map each Magistrala action string to Atom actions + via a translation table (below). +4. **`role_permission_blocks`** links role → block. +5. **`role_assignments`** one per `_role_member` whose user migrated + (`subject_kind='entity'`, `subject_id=member_id`, `role_id`, `tenant_id`). +6. For **domain** role members, also insert **`tenant_memberships`** + (`tenant_id`, `entity_id`, `status='active'`). + +### Action translation table (Magistrala → Atom action names) +Atom actions: `read, create, write, delete, revoke, rotate, publish, subscribe, +execute, manage, policy.manage, role.manage, authz.check`. + +| Magistrala action (examples) | Atom action | +|---|---| +| `read`, `view`, `*_read`, `*_view_*` | `read` | +| `create`, `*_create*` | `create` | +| `update`, `*_update` | `write` | +| `delete`, `*_delete*` | `delete` | +| `publish` | `publish` | +| `subscribe` | `subscribe` | +| `admin`, `manage`, `*_manage_role` | `manage` | +| `*_add_role_users`, `*_remove_role_users`, `membership*` | `policy.manage` | +| `*_view_role_users` | `read` | +| (unmapped) | log + default to `manage`, or skip — configurable | + +> Magistrala's full action vocabulary is generated from the SpiceDB schema +> (`docker/spicedb/schema.zed`). The migrator ships a complete map derived from +> that schema; anything missing is reported in the dry-run, never silently dropped. + +### Invitations +`domains.invitations` → `tenant_invitations` (preserve domain_id, invitee_user_id, +invited_by, role_id when the role migrated; map confirmed/rejected timestamps). +Pending only; accepted ones are already reflected as memberships. If the source +invitation references a stale role, the Atom `role_id` is written as `NULL` to +avoid a foreign-key failure while preserving the invitation record. + +--- + +## 5. Credentials (PATs + device secrets) + +### Device secrets (clients.secret) — RESOLVED: re-issue keys + +Magistrala stores the device secret in plaintext (looked up `WHERE secret = ...`). +Atom **cannot reuse it.** Verified against Atom source (`src/auth.rs`): +- `auth_from_api_key` calls `parse_api_key` then looks up `WHERE c.id = ` — lookup is by the credential UUID **embedded in the key**, not by an + identifier. +- `parse_api_key` requires exactly `atom_<32 hex cred-id>_<64 hex secret>` (secret + must be 32 raw bytes); anything else is rejected as malformed. + +So a raw Magistrala secret neither fits the format nor is reachable by lookup. +**Resolution: re-issue.** `phaseDeviceCreds` (`newAtomAPIKey`) mints a fresh +`atom__` per device, stores `argon2(raw 32-byte secret)` with +`credentials.id = credId`, and exports `device-keys-.csv` +(`client_id, domain_id, identity, api_key`, mode 0600) for re-provisioning +(bootstrap configs / device reflash). Credential id is derived (uuidv5 of client +id) so re-runs are idempotent; the plaintext key is only emitted by the apply run +that generated it. Validated: emitted key parses and argon2-verifies exactly as +Atom's auth path does. + +### PATs (auth.pats + pat_scopes) — RESOLVED: re-issue + +`pats.secret` is hashed (Magistrala PAT format), so plaintext isn't recoverable; +even if it were, it would not fit Atom's `atom__` format (same +constraint as device keys above). So PATs are **re-issue, no exception.** +`pat_scopes` are preserved in the credential `metadata.scopes` array +(`domain_id, entity_type, operation, entity_id`) for reference / future policy +reconstruction. +- Migrate metadata as `credentials(kind='api_key', entity_id=user_id, + identifier=pat.id, metadata={name,description,scopes,expires_at,...}, + status=revoked?‘revoked’:‘active’, expires_at)`. +- Because the secret can't be verified by Atom argon2, **mark migrated PATs as + needing re-issue** (same class of problem as passwords). Report them; do not + fabricate a usable secret. `pat_scopes` preserved in credential `metadata` for + reference / future policy reconstruction. + +--- + +## 6. Data-quality guardrails (pre-flight, fail loud) — IMPLEMENTED + +`preflight.go` runs all checks read-only before any write. Blocking issues abort +an `--apply` run (`preflightGate`); warnings are advisory. Dry-run reports both. +Checks below; the email check matters mainly for dumps merged across instances +(a single Magistrala enforces email/username uniqueness in its own tables). +1. **Tenant alias** (domain.route): lowercase-fold; must match slug regex and **not** + be UUID-shaped; globally unique case-insensitively. Offending rows → report, + require operator fix or null the alias. +2. **Entity/resource alias** (client.identity / channel.route): same slug rule, + unique per tenant. On violation, drop the alias (keep UUID) rather than abort. +3. **Per-tenant name uniqueness**: `entities(name, tenant_id)` and + `resources` names — Magistrala already enforces `(domain_id, name)` so this is + usually safe; still verify (users have no tenant, so global human-name dupes are + fine because humans are keyed by id/email). +4. **FK integrity**: skip rows whose `domain_id` has no surviving tenant; skip + role members whose user did not migrate; set stale invitation roles to `NULL`; + report all cases. +5. **Email uniqueness**: `entity_emails.email` is globally UNIQUE — dedupe/report + conflicting user emails before load. + +--- + +## 7. Load ordering (FK-safe) + +1. tenants +2. entities (human, device) — without created_by/updated_by FKs first… +3. …then backfill tenants.created_by/updated_by and resources.owner_id +4. entity_emails +5. credentials (device api_key; PAT metadata) +6. resources (channels, rules, reports, alarms) +7. object_groups → object_group_hierarchy → object_group_entities/resources +8. roles → permission_blocks → permission_block_actions → role_permission_blocks +9. role_assignments, direct_policies +10. tenant_memberships +11. tenant_invitations + +All inserts use deterministic PKs (preserved IDs) + `ON CONFLICT DO NOTHING`/upsert +so the migration is **idempotent / re-runnable**. + +--- + +## 8. Go program design + +``` +tools/atom-migration/ + main.go // flags: --dry-run (default), --apply, --report-dir + config.go // reads docker/.env for DB hosts/ports/creds/names + source/ // one reader per source DB (sqlx/pgx), returns typed structs + domains.go users.go clients.go channels.go groups.go auth.go + transform/ // pure funcs: MG structs -> Atom rows; status/alias/action maps + target/ // Atom writer: ordered, transactional, ON CONFLICT upserts + preflight/ // §6 guardrails -> report + report/ // JSON+markdown summary: counts, skips, conflicts, todo lists +``` + +- **Connectivity:** run the migrator as a one-shot container on + `magistrala-base-net` (compose `run --rm`) so it resolves `domains-db`, + `users-db`, … and the Atom `postgres` by service name. Alternatively expose host + ports and run from host. DB creds/names pulled from `docker/.env` keys + (`MG_*_DB_*`) and Atom's `.env` (`POSTGRES_*`). +- **Idempotent & resumable:** every write upserts on preserved PK; safe to re-run. +- **Dry-run first:** default mode reads + transforms + validates + writes a report, + touches nothing. `--apply` runs the load in one transaction per phase. +- **argon2** for device secrets via `golang.org/x/crypto/argon2` matching Atom's + params (confirm Atom's argon2 config: variant/m/t/p) so hashes verify. + +--- + +## 9. Cutover runbook (offline) + +1. `docker compose stop` Magistrala **app** services (users, clients, channels, + groups, domains, auth, adapters) — keep the `*-db` containers running. +2. Backup: `pg_dump` each source DB. +3. Start Atom Postgres; let Atom run once to apply `migrations/001_initial.sql` + (or run migrations standalone), then stop Atom app. +4. Run migrator `--dry-run`; review report; fix guardrail violations (§6). +5. Run migrator `--apply`. +6. Start Atom app; verify (§10). +7. Point remaining Magistrala services (messaging/bootstrap/certs/readers) at Atom + for authn/authz; decommission SpiceDB + users/clients/channels/groups/domains/auth. + +--- + +## 10. Verification + +`--verify` (`verify.go`, read-only) reconciles a completed migration: +- every source id that should have migrated exists in Atom (tenants, human + + device entities, resources, object_groups); missing rows → blocking. +- every device→channel connection has a matching authz edge (direct_policy + + object-scope block + publish/subscribe action); missing → blocking. + +Still recommended manually post-cutover: +- Spot `POST /authz/check` for a sample of (user, domain, action) and + (device, channel, publish) allowed pre-migration. +- Admin login (seeded atom-admin) works; a migrated user completes password reset. +- A re-issued device key authenticates (§5). + +--- + +## 11. Open items — status + +1. **Device key format** (§5) — RESOLVED. Atom looks up by embedded cred UUID and + requires `atom_<32hex>_<64hex>`; MG secrets can't be carried → re-issue + CSV + export. Implemented + validated. +2. **argon2 params** — RESOLVED. Atom uses `Argon2::default()` (argon2id, v=19, + m=19456, t=2, p=1, 32-byte tag); migrator emits the matching PHC string and + re-issued keys verify against Atom's path. +3. **Groups semantics** (§3.5) — RESOLVED. Magistrala groups have no user-member + table; only `parent_group_id` (clients/channels) + group-scoped roles. They are + object groupings → `object_groups`. No principal-group case. +4. **PAT re-issue** (§5) — RESOLVED. Re-issue, no exception (hashed + format). +5. `created_by`/`owner_id` to non-migrated/system users — current behaviour: + `tenants.created_by/updated_by` and `resources.owner_id` are set only when the + referenced user migrated, else left NULL. Operator may prefer pointing these at + the seeded admin entity (`...0001`); a one-line change in the backfill/owner + logic if desired. **Only remaining choice.** diff --git a/tools/atom-migration/README.md b/tools/atom-migration/README.md new file mode 100644 index 000000000..9c9f1d66f --- /dev/null +++ b/tools/atom-migration/README.md @@ -0,0 +1,141 @@ +# atom-migration + +Offline, idempotent migrator: Magistrala v0.30.0 (per-service Postgres) → Atom IAM +(single Postgres). See [PLAN.md](./PLAN.md) for the full mapping and runbook. + +## Build + +Plain binary: + +```bash +go build ./tools/atom-migration/ +``` + +Or a Docker image, so you can run the tool repeatedly without `go run` and +without a Go toolchain on the host (build context is the repo root): + +```bash +docker build -f tools/atom-migration/Dockerfile -t magistrala/atom-migration:dev . +``` + +## Start only the source databases + +The migrator reads Postgres directly — it does **not** need the Magistrala app +services running. To migrate from restored volumes, start just the nine source DB +containers (`--no-deps` keeps compose from pulling in the app services they +depend on): + +```bash +docker compose -f docker/docker-compose.yaml up -d --no-deps \ + auth-db users-db domains-db clients-db channels-db groups-db \ + re-db reports-db alarms-db +``` + +They mount the `magistrala_magistrala--db-volume` volumes and attach to +`magistrala-base-net`, where the migrator resolves them by service name. + +## Run (dry-run is default — writes nothing) + +The migrator needs to reach every source DB **and** the Atom DB. Atom must already +have its schema applied (run Atom once, or apply `migrations/001_initial.sql`). + +Run the image on the compose network. Mount the repo at `/work` so the default +`--env docker/.env` and `--report-dir` resolve, and reach source DBs by their +compose service names: + +If Atom runs as its own compose project (its containers live on `atom_default`), +that network is isolated from `magistrala-base-net`, so the migrator cannot reach +the Atom DB yet. Bridge the Atom Postgres onto the migrator's network once, then +address it by its container name `atom-postgres-1`: + +```bash +docker network connect magistrala-base-net atom-postgres-1 +``` + +```bash +docker run --rm --network magistrala-base-net \ + --user "$(id -u):$(id -g)" \ + -v "$PWD":/work -w /work \ + magistrala/atom-migration:dev \ + --env docker/.env \ + --atom-dsn 'host=atom-postgres-1 port=5432 user=atom password=atom dbname=atom sslmode=disable' +``` + +`--user` makes the container write the report as your host user; without it the +image's non-root user cannot create `report/` in the bind-mounted repo. + +The `network connect` is **not persistent** — re-run it after any +`docker compose down`/recreate of the Atom stack (the bridge silently drops and +`postgres` stops resolving). For a permanent setup, declare `magistrala-base-net` +as an external network in Atom's compose instead. + +Add `--apply` / `--verify` as described below. + +### As a one-shot container on the compose network (recommended) + +```bash +go run ./tools/atom-migration \ + --env docker/.env --from-host \ + --atom-dsn 'host=127.0.0.1 port=5432 user=atom password=atom dbname=atom sslmode=disable' +``` + +## Apply + +```bash +go run ./tools/atom-migration ... --apply +``` + +A pre-flight runs first (read-only). Same checks run in dry-run for the report. + +Magistrala dropped several uniqueness constraints that Atom still enforces +(`tenants.name`; device and group `name` per tenant; tenant/entity/resource +`alias`). The migrator resolves these automatically: in each collision the +oldest row keeps its value and the rest get a deterministic, id-derived suffix +(`name (a1b2c3d4)` / `alias-a1b2c3d4`). Renames are reported as `renamed.*` +counts plus a warning per row, and are stable across re-runs. So these no longer +block the apply — only issues the tool cannot safely auto-fix do (e.g. duplicate +user emails, which are identities). + +## Verify (after apply) + +```bash +go run ./tools/atom-migration ... --verify +``` + +Read-only reconciliation: every source row that should have migrated must exist in +Atom (tenants, entities, resources, object_groups) and every device→channel +connection must have a matching authz edge. Missing rows are reported as blocking. + +## Flags + +| flag | default | meaning | +| ------------------- | ----------------------------- | ---------------------------------------------------- | +| `--env` | `docker/.env` | Magistrala env file (reads `MG_*_DB_*`) | +| `--atom-dsn` | compose default | Atom Postgres DSN (or `ATOM_DATABASE_URL`) | +| `--from-host` | false | rewrite source hosts to 127.0.0.1 (use mapped ports) | +| `--apply` | false | perform the load (omit = dry-run) | +| `--report-dir` | `tools/atom-migration/report` | JSON+markdown report output | +| `--unmapped-action` | `manage` | fallback for unmapped MG actions: `manage` or `skip` | + +## Credentials are re-issued, not carried + +Atom authenticates API keys by a credential UUID embedded in the key +(`atom_<32hex>_<64hex>`, argon2 over the raw 32-byte secret — see Atom +`src/auth.rs`). Magistrala secrets fit neither the format nor the lookup, so: + +- **Device keys** are re-issued. On `--apply` the migrator writes + `report/device-keys-.csv` (`client_id,domain_id,identity,api_key`, mode + 0600). Re-provision devices/bootstrap configs from it, then delete it — the + plaintext secret is shown only once. +- **User passwords** (bcrypt → argon2 unconvertible): users land with no password + credential. Report's `password_reset` TODO lists every user for the email reset. +- **PAT secrets** (hashed + format): metadata migrates, secret must be re-issued. + Report's `pat_reissue` TODO. +- Transient data not migrated: OTP verifications, short-lived auth `keys`, login + attempts. + +## Idempotency + +Every write upserts on a preserved/derived primary key (`ON CONFLICT DO NOTHING`), +so the tool is safe to re-run. Derived UUIDs (roles, permission blocks, policies) +use uuidv5 so they are stable across runs. diff --git a/tools/atom-migration/config.go b/tools/atom-migration/config.go new file mode 100644 index 000000000..08d640160 --- /dev/null +++ b/tools/atom-migration/config.go @@ -0,0 +1,135 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package main + +import ( + "bufio" + "fmt" + "os" + "strings" +) + +// dbConn is one source/target Postgres connection target. +type dbConn struct { + Host string + Port string + User string + Pass string + Name string + SSL string +} + +func (d dbConn) DSN() string { + ssl := d.SSL + if ssl == "" { + ssl = "disable" + } + return fmt.Sprintf("host=%s port=%s user=%s password=%s dbname=%s sslmode=%s", + d.Host, d.Port, d.User, d.Pass, d.Name, ssl) +} + +// config holds every source DB plus the Atom target DSN. +type config struct { + Domains dbConn + Users dbConn + Clients dbConn + Channels dbConn + Groups dbConn + Auth dbConn + RE dbConn // rules engine + Reports dbConn + Alarms dbConn + + AtomDSN string + UnmappedAction string +} + +// loadConfig reads docker/.env for MG_*_DB_* keys. When fromHost is true the +// service-name hosts are rewritten to 127.0.0.1 with the mapped host port (the +// caller is then responsible for exposing those ports in compose). +func loadConfig(envPath, atomDSN string, fromHost bool) (config, error) { + env, err := parseEnvFile(envPath) + if err != nil { + return config{}, err + } + + mk := func(prefix string) dbConn { + c := dbConn{ + Host: env[prefix+"_DB_HOST"], + Port: orDef(env[prefix+"_DB_PORT"], "5432"), + User: orDef(env[prefix+"_DB_USER"], "magistrala"), + Pass: orDef(env[prefix+"_DB_PASS"], "magistrala"), + Name: env[prefix+"_DB_NAME"], + SSL: orDef(env[prefix+"_DB_SSL_MODE"], "disable"), + } + if fromHost { + c.Host = "127.0.0.1" + } + return c + } + + cfg := config{ + Domains: mk("MG_DOMAINS"), + Users: mk("MG_USERS"), + Clients: mk("MG_CLIENTS"), + Channels: mk("MG_CHANNELS"), + Groups: mk("MG_GROUPS"), + Auth: mk("MG_AUTH"), + RE: mk("MG_RE"), + Reports: mk("MG_REPORTS"), + Alarms: mk("MG_ALARMS"), + AtomDSN: atomDSN, + } + + // Default names if .env omitted them. + defName := map[*string]string{ + &cfg.Domains.Name: "domains", &cfg.Users.Name: "users", + &cfg.Clients.Name: "clients", &cfg.Channels.Name: "channels", + &cfg.Groups.Name: "groups", &cfg.Auth.Name: "auth", + &cfg.RE.Name: "rules_engine", &cfg.Reports.Name: "reports", + &cfg.Alarms.Name: "alarms", + } + for p, n := range defName { + if *p == "" { + *p = n + } + } + + if cfg.AtomDSN == "" { + // Fall back to Atom compose defaults; override with --atom-dsn or + // ATOM_DATABASE_URL for real runs. + cfg.AtomDSN = "host=127.0.0.1 port=5432 user=atom password=atom dbname=atom sslmode=disable" + } + return cfg, nil +} + +func parseEnvFile(path string) (map[string]string, error) { + f, err := os.Open(path) + if err != nil { + return nil, fmt.Errorf("open %s: %w", path, err) + } + defer f.Close() + + out := map[string]string{} + sc := bufio.NewScanner(f) + for sc.Scan() { + line := strings.TrimSpace(sc.Text()) + if line == "" || strings.HasPrefix(line, "#") { + continue + } + k, v, ok := strings.Cut(line, "=") + if !ok { + continue + } + out[strings.TrimSpace(k)] = strings.Trim(strings.TrimSpace(v), `"'`) + } + return out, sc.Err() +} + +func orDef(v, def string) string { + if v == "" { + return def + } + return v +} diff --git a/tools/atom-migration/crypto.go b/tools/atom-migration/crypto.go new file mode 100644 index 000000000..9224c6512 --- /dev/null +++ b/tools/atom-migration/crypto.go @@ -0,0 +1,67 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package main + +import ( + "crypto/rand" + "encoding/base64" + "encoding/hex" + "fmt" + "strings" + + "github.com/google/uuid" + "golang.org/x/crypto/argon2" +) + +// Atom uses argon2 crate 0.5 `Argon2::default()` => argon2id, v=19, +// m=19456 KiB, t=2, p=1, 32-byte tag. We emit the matching PHC string so the +// hash verifies in Atom. +const ( + argonMemory = 19456 + argonTime = 2 + argonThreads = 1 + argonKeyLen = 32 + argonSaltLen = 16 +) + +// hashArgon2id produces a PHC-encoded argon2id hash compatible with Atom. +func hashArgon2id(secret []byte) (string, error) { + salt := make([]byte, argonSaltLen) + if _, err := rand.Read(salt); err != nil { + return "", err + } + key := argon2.IDKey(secret, salt, argonTime, argonMemory, argonThreads, argonKeyLen) + + b64 := base64.RawStdEncoding // PHC uses unpadded standard base64 + return fmt.Sprintf( + "$argon2id$v=%d$m=%d,t=%d,p=%d$%s$%s", + argon2.Version, argonMemory, argonTime, argonThreads, + b64.EncodeToString(salt), b64.EncodeToString(key), + ), nil +} + +// newAtomAPIKey mints a fresh Atom-format API key for a device. Atom expects +// `atom_<32hex-credId>_<64hex-secret>` and verifies argon2 over the raw 32 secret +// bytes (see atom src/auth.rs parse_api_key / auth_from_api_key). Magistrala's +// own device secret cannot be reused (arbitrary format, looked up differently), +// so the key is re-issued and must be re-provisioned to the device. +// +// credID is derived deterministically from the client id so re-runs are +// idempotent (same credential row id), but the returned plaintext key is only +// usable from the run that generated it. +func newAtomAPIKey(clientID string) (credID, plaintextKey, secretHash string, err error) { + cu := uuid.NewSHA1(uuidNS, []byte("devcred|"+clientID)) + credIDHex := strings.ReplaceAll(cu.String(), "-", "") + + secret := make([]byte, 32) + if _, err = rand.Read(secret); err != nil { + return "", "", "", err + } + hash, err := hashArgon2id(secret) + if err != nil { + return "", "", "", err + } + key := "atom_" + credIDHex + "_" + hex.EncodeToString(secret) + return cu.String(), key, hash, nil +} diff --git a/tools/atom-migration/db.go b/tools/atom-migration/db.go new file mode 100644 index 000000000..0cb689df2 --- /dev/null +++ b/tools/atom-migration/db.go @@ -0,0 +1,23 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package main + +import ( + "context" + "fmt" + + _ "github.com/jackc/pgx/v5/stdlib" // pgx stdlib driver + "github.com/jmoiron/sqlx" +) + +func openDB(ctx context.Context, name, dsn string) (*sqlx.DB, error) { + db, err := sqlx.Open("pgx", dsn) + if err != nil { + return nil, fmt.Errorf("open %s: %w", name, err) + } + if err := db.PingContext(ctx); err != nil { + return nil, fmt.Errorf("ping %s: %w", name, err) + } + return db, nil +} diff --git a/tools/atom-migration/dedup.go b/tools/atom-migration/dedup.go new file mode 100644 index 000000000..e0d496fa0 --- /dev/null +++ b/tools/atom-migration/dedup.go @@ -0,0 +1,304 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package main + +import ( + "context" + "database/sql" + "sort" + "strings" +) + +// Magistrala dropped several uniqueness constraints that Atom still enforces +// (tenants.name, object device/group name per tenant, tenant/entity/resource +// alias). buildDedup precomputes collision-free names and aliases so the load +// never trips those indexes. Renaming is deterministic: rows are processed in a +// stable order (created_at, id) and the loser of a collision keeps its original +// value plus an id-derived suffix, so re-runs produce identical output. +func (m *migrator) buildDedup(ctx context.Context, rep *report) error { + doms, err := readDomains(ctx, m.domainsDB) + if err != nil { + return err + } + sort.SliceStable(doms, func(i, j int) bool { + return earlier(doms[i].CreatedAt, doms[i].ID, doms[j].CreatedAt, doms[j].ID) + }) + domSet := map[string]bool{} + for _, d := range doms { + domSet[d.ID] = true + } + + tNames := newAllocator(false) + tAlias := newAllocator(true) + for _, d := range doms { + base := strings.TrimSpace(nsToStr(d.Name)) + if base == "" { + base = d.ID + } + final := tNames.take("", base, d.ID) + if final != base { + rep.warnf("tenant %s name %q -> %q (tenants.name is UNIQUE)", d.ID, base, final) + rep.count("renamed.tenants", 1) + } + m.tenantName[d.ID] = final + m.tenantAlias[d.ID] = m.dedupAlias(tAlias, "", d.Route, "tenant "+d.ID, rep) + } + + clients, err := readClients(ctx, m.clientsDB) + if err != nil { + return err + } + sort.SliceStable(clients, func(i, j int) bool { + return earlier(clients[i].CreatedAt, clients[i].ID, clients[j].CreatedAt, clients[j].ID) + }) + dNames := newAllocator(false) + cAlias := newAllocator(true) + for _, c := range clients { + if !domSet[c.DomainID] { + continue + } + base := firstNonEmpty(strings.TrimSpace(nsToStr(c.Name)), c.ID) + final := dNames.take(c.DomainID, base, c.ID) + if final != base { + rep.warnf("device %s name %q -> %q (entities(name, tenant_id) is UNIQUE)", c.ID, base, final) + rep.count("renamed.devices", 1) + } + m.deviceName[c.ID] = final + m.clientAlias[c.ID] = m.dedupAlias(cAlias, c.DomainID, c.Identity, "client "+c.ID, rep) + } + + chans, err := readChannels(ctx, m.channelsDB) + if err != nil { + return err + } + sort.SliceStable(chans, func(i, j int) bool { + return earlier(chans[i].CreatedAt, chans[i].ID, chans[j].CreatedAt, chans[j].ID) + }) + chAlias := newAllocator(true) + for _, ch := range chans { + if !domSet[ch.DomainID] { + continue + } + m.channelAlias[ch.ID] = m.dedupAlias(chAlias, ch.DomainID, ch.Route, "channel "+ch.ID, rep) + } + + grps, err := readGroups(ctx, m.groupsDB) + if err != nil { + return err + } + sort.SliceStable(grps, func(i, j int) bool { + return earlier(grps[i].CreatedAt, grps[i].ID, grps[j].CreatedAt, grps[j].ID) + }) + gNames := newAllocator(false) + for _, g := range grps { + if !domSet[g.DomainID] { + continue + } + base := firstNonEmpty(strings.TrimSpace(g.Name), g.ID) + final := gNames.take(g.DomainID, base, g.ID) + if final != base { + rep.warnf("group %s name %q -> %q (object_groups(name, tenant_id) is UNIQUE)", g.ID, base, final) + rep.count("renamed.groups", 1) + } + m.groupName[g.ID] = final + } + + if err := m.dedupUsers(ctx, rep); err != nil { + return err + } + return nil +} + +// dedupUsers computes collision-free entity names for users. Magistrala users +// are platform-global, so they land in Atom's tenant-less namespace +// (entities.tenant_id IS NULL). Migration 006 collapsed NULL tenant_id to a +// sentinel UUID, making that whole namespace a single unique scope shared with +// the bootstrap system entities (admin, mg-service). Pre-seed those existing +// tenant-less names so a colliding migrated user is renamed rather than +// tripping idx_entities_name_tenant. Names already owned by a user we are about +// to migrate (same id) are not reserved: those re-run idempotently via +// ON CONFLICT (id) and must keep their original name. +func (m *migrator) dedupUsers(ctx context.Context, rep *report) error { + users, err := readUsers(ctx, m.usersDB) + if err != nil { + return err + } + srcID := map[string]bool{} + for _, u := range users { + srcID[u.ID] = true + } + + uNames := newAllocator(false) + rows, err := m.atom.QueryxContext(ctx, + `SELECT id, name FROM entities WHERE tenant_id IS NULL`) + if err != nil { + return err + } + for rows.Next() { + var id, name string + if err := rows.Scan(&id, &name); err != nil { + rows.Close() + return err + } + if srcID[id] { + continue // re-run of this user; ON CONFLICT (id) keeps its name + } + uNames.take("", name, id) // reserve the system/foreign name + } + rows.Close() + + sort.SliceStable(users, func(i, j int) bool { + return earlier(users[i].CreatedAt, users[i].ID, users[j].CreatedAt, users[j].ID) + }) + for _, u := range users { + base := firstNonEmpty(u.Username.String, u.Email.String, u.ID) + final := uNames.take("", base, u.ID) + if final != base { + rep.warnf("user %s name %q -> %q (entities(name, tenant_id) is UNIQUE)", u.ID, base, final) + rep.count("renamed.users", 1) + } + m.userName[u.ID] = final + } + return nil +} + +// dedupAlias normalizes a candidate alias and makes it unique within scope. +// Invalid slugs are dropped (as before); a valid alias that collides keeps its +// value plus an id-derived suffix instead of blocking the load. Returns "" when +// no alias should be set (caller writes NULL). +func (m *migrator) dedupAlias(a *allocator, scope string, raw sql.NullString, label string, rep *report) string { + if !raw.Valid || raw.String == "" { + return "" + } + norm, ok := normalizeAlias(raw.String) + if !ok { + rep.warnf("alias dropped for %s: %q not a valid slug", label, raw.String) + rep.skip("alias_dropped") + return "" + } + final := a.take(scope, norm, idSuffix(raw)) + if final != norm { + rep.warnf("alias for %s %q -> %q (alias is UNIQUE within tenant)", label, norm, final) + rep.count("renamed.aliases", 1) + } + return final +} + +// idSuffix derives a stable, slug-safe suffix source. raw aliases carry no id, +// so dedupAlias passes the raw alias itself; collisions then disambiguate on a +// short hash of it, which is deterministic per source value. +func idSuffix(raw sql.NullString) string { + return shortHash(raw.String) +} + +// allocator hands out unique strings within a namespace. caseFold true compares +// case-insensitively (aliases); names compare exactly. +type allocator struct { + used map[string]map[string]bool + caseFold bool +} + +func newAllocator(caseFold bool) *allocator { + return &allocator{used: map[string]map[string]bool{}, caseFold: caseFold} +} + +func (a *allocator) key(s string) string { + if a.caseFold { + return strings.ToLower(s) + } + return s +} + +// take returns base if free in scope, else base with an id-derived suffix. Long +// alias values are trimmed so the result stays within the 63-char slug limit. +func (a *allocator) take(scope, base, id string) string { + bucket := a.used[scope] + if bucket == nil { + bucket = map[string]bool{} + a.used[scope] = bucket + } + if !bucket[a.key(base)] { + bucket[a.key(base)] = true + return base + } + for _, suf := range []string{short(id), id} { + cand := withSuffix(base, suf, a.caseFold) + if !bucket[a.key(cand)] { + bucket[a.key(cand)] = true + return cand + } + } + // Pathological fallback: keep extending until unique. + cand := withSuffix(base, id, a.caseFold) + for bucket[a.key(cand)] { + cand += "x" + } + bucket[a.key(cand)] = true + return cand +} + +// withSuffix appends "-suf"; for aliases it caps the total at the 63-char slug +// limit by trimming the base. +func withSuffix(base, suf string, alias bool) string { + sep := "-" + if !alias { + sep = " (" + } + if alias { + if max := 63 - len(suf) - 1; len(base) > max && max > 0 { + base = strings.TrimRight(base[:max], "-") + } + return base + sep + suf + } + return base + sep + suf + ")" +} + +func short(id string) string { + if len(id) > 8 { + return id[:8] + } + return id +} + +// shortHash is a small deterministic hex tag (FNV-1a, 8 hex chars). +func shortHash(s string) string { + const ( + offset = 2166136261 + prime = 16777619 + ) + h := uint32(offset) + for i := 0; i < len(s); i++ { + h ^= uint32(s[i]) + h *= prime + } + const hexd = "0123456789abcdef" + out := make([]byte, 8) + for i := 7; i >= 0; i-- { + out[i] = hexd[h&0xf] + h >>= 4 + } + return string(out) +} + +// earlier orders rows by (created_at, id): the oldest row wins a collision and +// keeps its original name/alias. Rows with an unknown created_at sort last; id +// breaks ties so the order is total and stable across runs. +func earlier(ta sql.NullTime, ida string, tb sql.NullTime, idb string) bool { + switch { + case ta.Valid && tb.Valid && !ta.Time.Equal(tb.Time): + return ta.Time.Before(tb.Time) + case ta.Valid != tb.Valid: + return ta.Valid + default: + return ida < idb + } +} + +// aliasOrNil converts a computed alias ("" = none) into a SQL argument. +func aliasOrNil(s string) any { + if s == "" { + return nil + } + return s +} diff --git a/tools/atom-migration/main.go b/tools/atom-migration/main.go new file mode 100644 index 000000000..fc39b0e0f --- /dev/null +++ b/tools/atom-migration/main.go @@ -0,0 +1,96 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +// Command atom-migration migrates a Magistrala v0.30.0 deployment (per-service +// Postgres databases) into a single Atom IAM Postgres database. +// +// It is offline and idempotent: every write upserts on a preserved/derived +// primary key, so the tool is safe to re-run. Default mode is --dry-run, which +// reads, transforms, validates and reports without writing anything. +// +// See PLAN.md in this directory for the full mapping and runbook. +package main + +import ( + "context" + "flag" + "fmt" + "log" + "os" + "time" +) + +func main() { + os.Exit(runMain()) +} + +func runMain() int { + var ( + envPath = flag.String("env", "docker/.env", "path to Magistrala docker/.env (DB hosts/creds)") + atomDSN = flag.String("atom-dsn", envOr("ATOM_DATABASE_URL", ""), "Atom Postgres DSN (overrides --env atom block)") + apply = flag.Bool("apply", false, "perform the load (default is dry-run: read+validate+report only)") + reportDir = flag.String("report-dir", "tools/atom-migration/report", "directory for the JSON+markdown report") + fromHost = flag.Bool("from-host", false, "connect to source DBs via localhost mapped ports instead of compose service names") + unmappedOK = flag.String("unmapped-action", "manage", "fallback Atom action for unmapped Magistrala actions: manage|skip") + verify = flag.Bool("verify", false, "verify a completed migration (reconcile source vs Atom); read-only, no load") + ) + flag.Parse() + + cfg, err := loadConfig(*envPath, *atomDSN, *fromHost) + if err != nil { + log.Printf("config: %v", err) + return 1 + } + cfg.UnmappedAction = *unmappedOK + + mode := "DRY-RUN" + switch { + case *verify: + mode = "VERIFY" + case *apply: + mode = "APPLY" + } + log.Printf("atom-migration starting (%s)", mode) + + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Hour) + defer cancel() + + m, err := newMigrator(ctx, cfg, *apply) + if err != nil { + log.Printf("init: %v", err) + return 1 + } + m.reportDir = *reportDir + defer m.Close() + + rep := newReport(mode) + run := m.Run + if *verify { + run = m.Verify + } + if err := run(ctx, rep); err != nil { + rep.Errorf("fatal: %v", err) + if writeErr := rep.Write(*reportDir); writeErr != nil { + log.Printf("report: %v", writeErr) + } + log.Printf("%s: %v", mode, err) + return 1 + } + + if err := rep.Write(*reportDir); err != nil { + log.Printf("report: %v", err) + return 1 + } + fmt.Print(rep.Summary()) + if rep.HasBlocking() && *apply { + return 2 + } + return 0 +} + +func envOr(k, def string) string { + if v := os.Getenv(k); v != "" { + return v + } + return def +} diff --git a/tools/atom-migration/maps.go b/tools/atom-migration/maps.go new file mode 100644 index 000000000..75ffb796a --- /dev/null +++ b/tools/atom-migration/maps.go @@ -0,0 +1,145 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package main + +import ( + "regexp" + "strings" + + "github.com/google/uuid" +) + +// uuidNS is the deterministic namespace for derived Atom UUIDs (roles, +// permission blocks) whose Magistrala source id is not a UUID. Using uuidv5 +// keeps the migration idempotent across re-runs. +var uuidNS = uuid.MustParse("a70a0000-0000-5000-a000-000000000000") + +const ( + statusActive = "active" + + actionCreate = "create" + actionDelete = "delete" + actionManage = "manage" + actionPublish = "publish" + actionRead = "read" + actionSubscribe = "subscribe" +) + +func derivedUUID(parts ...string) string { + return uuid.NewSHA1(uuidNS, []byte(strings.Join(parts, "|"))).String() +} + +// --- status mapping (Magistrala smallint -> Atom enum text) --- + +// entityStatus maps to entities.status / object_groups.status (active/inactive/suspended). +func entityStatus(s int16) string { + switch s { + case 0: + return statusActive + case 1: + return "inactive" + default: + return "suspended" + } +} + +// tenantStatus maps to tenants.status (active/inactive/frozen/deleted). +func tenantStatus(s int16) string { + switch s { + case 0: + return statusActive + case 1: + return "inactive" + case 2: + return "frozen" + default: + return statusActive + } +} + +// --- alias normalization (Atom 004/005 slug rules) --- + +var ( + slugRe = regexp.MustCompile(`^[a-z0-9]([a-z0-9-]{0,61}[a-z0-9])?$`) + uuidRe = regexp.MustCompile(`^([0-9a-f]{32}|[0-9a-f]{8}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12})$`) +) + +// normalizeAlias lowercases and validates a candidate alias. Returns ("", false) +// when the value cannot be a valid Atom alias (caller keeps the UUID identity and +// drops the alias, logging it). +func normalizeAlias(raw string) (string, bool) { + a := strings.ToLower(strings.TrimSpace(raw)) + if a == "" { + return "", false + } + if uuidRe.MatchString(a) { + return "", false // UUID-shaped aliases collide with id-addressing + } + if !slugRe.MatchString(a) { + return "", false + } + return a, true +} + +// --- action vocabulary mapping (Magistrala -> Atom) --- + +// Atom actions: read create write delete revoke rotate publish subscribe execute +// manage policy.manage role.manage authz.check. +// +// mapAction translates a Magistrala role action string. Returns ("", false) when +// unmapped; caller decides fallback per --unmapped-action. +func mapAction(a string) (string, bool) { + a = strings.ToLower(strings.TrimSpace(a)) + switch a { + case actionPublish: + return actionPublish, true + case actionSubscribe: + return actionSubscribe, true + case actionCreate: + return actionCreate, true + case "update": + return "write", true + case actionDelete: + return actionDelete, true + case "read", "view": + return actionRead, true + case "admin", actionManage: + return actionManage, true + } + switch { + case strings.HasSuffix(a, "_publish"): + return actionPublish, true + case strings.HasSuffix(a, "_subscribe"): + return actionSubscribe, true + case strings.HasSuffix(a, "_view_role_users"), strings.HasSuffix(a, "_read"), + strings.HasSuffix(a, "_view"), strings.Contains(a, "_read_"): + return actionRead, true + case strings.Contains(a, actionCreate): + return actionCreate, true + case strings.HasSuffix(a, "_update"): + return "write", true + case strings.Contains(a, "_delete"): + return actionDelete, true + case strings.HasSuffix(a, "_manage_role"): + return actionManage, true + case strings.HasSuffix(a, "_add_role_users"), strings.HasSuffix(a, "_remove_role_users"), + strings.Contains(a, "membership"), strings.Contains(a, "_connect"): + return "policy.manage", true + case strings.HasSuffix(a, "_share"), strings.HasSuffix(a, "_unshare"): + return "policy.manage", true + } + return "", false +} + +// connectionAction maps Magistrala connection type (1=publish, 2=subscribe). +func connectionAction(t int16) (string, bool) { + switch t { + case 1: + return actionPublish, true + case 2: + return actionSubscribe, true + default: + return "", false + } +} diff --git a/tools/atom-migration/migrator.go b/tools/atom-migration/migrator.go new file mode 100644 index 000000000..482e2a8d0 --- /dev/null +++ b/tools/atom-migration/migrator.go @@ -0,0 +1,1199 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package main + +import ( + "context" + "database/sql" + "encoding/csv" + "encoding/json" + "fmt" + "os" + "path/filepath" + "strings" + "time" + + "github.com/jmoiron/sqlx" +) + +type migrator struct { + cfg config + apply bool + + domainsDB *sqlx.DB + usersDB *sqlx.DB + clientsDB *sqlx.DB + channelsDB *sqlx.DB + groupsDB *sqlx.DB + authDB *sqlx.DB + reDB *sqlx.DB + reportsDB *sqlx.DB + alarmsDB *sqlx.DB + atom *sqlx.DB + + profileID map[string]string // profile key (e.g. "user","client") -> uuid + profileVersionID map[string]string // profile key -> latest active profile_versions.id + actionID map[string]string // action name -> uuid + + migratedUsers map[string]bool + migratedRoles map[string]bool + tenants map[string]bool // domain ids that became tenants + clientDomain map[string]string // client id -> domain id + channelDomain map[string]string + groupDomain map[string]string + resourceDomain map[string]string // resource id -> tenant id + + // Collision-free names/aliases computed by buildDedup (Atom enforces unique + // constraints Magistrala dropped). Keyed by source id; "" alias = NULL. + tenantName map[string]string + userName map[string]string + deviceName map[string]string + groupName map[string]string + tenantAlias map[string]string + clientAlias map[string]string + channelAlias map[string]string + + reportDir string + deviceKeys [][]string // client_id, domain_id, identity, plaintext key (apply only) +} + +func newMigrator(ctx context.Context, cfg config, apply bool) (*migrator, error) { + m := &migrator{ + cfg: cfg, + apply: apply, + profileID: map[string]string{}, + profileVersionID: map[string]string{}, + actionID: map[string]string{}, + migratedUsers: map[string]bool{}, + migratedRoles: map[string]bool{}, + tenants: map[string]bool{}, + clientDomain: map[string]string{}, + channelDomain: map[string]string{}, + groupDomain: map[string]string{}, + resourceDomain: map[string]string{}, + tenantName: map[string]string{}, + userName: map[string]string{}, + deviceName: map[string]string{}, + groupName: map[string]string{}, + tenantAlias: map[string]string{}, + clientAlias: map[string]string{}, + channelAlias: map[string]string{}, + } + var err error + open := func(name, dsn string) *sqlx.DB { + if err != nil { + return nil + } + var db *sqlx.DB + db, err = openDB(ctx, name, dsn) + return db + } + m.domainsDB = open("domains", cfg.Domains.DSN()) + m.usersDB = open("users", cfg.Users.DSN()) + m.clientsDB = open("clients", cfg.Clients.DSN()) + m.channelsDB = open("channels", cfg.Channels.DSN()) + m.groupsDB = open("groups", cfg.Groups.DSN()) + m.authDB = open("auth", cfg.Auth.DSN()) + m.reDB = open("rules_engine", cfg.RE.DSN()) + m.reportsDB = open("reports", cfg.Reports.DSN()) + m.alarmsDB = open("alarms", cfg.Alarms.DSN()) + m.atom = open("atom", cfg.AtomDSN) + if err != nil { + return nil, err + } + return m, nil +} + +const authenticatedUsersGroupID = "00000000-0000-0000-0000-000000000005" + +func (m *migrator) Close() { + for _, db := range []*sqlx.DB{m.domainsDB, m.usersDB, m.clientsDB, m.channelsDB, m.groupsDB, m.authDB, m.reDB, m.reportsDB, m.alarmsDB, m.atom} { + if db != nil { + _ = db.Close() + } + } +} + +func (m *migrator) Run(ctx context.Context, rep *report) error { + if err := m.loadLookups(ctx); err != nil { + return fmt.Errorf("load atom lookups: %w", err) + } + if err := m.preflight(ctx, rep); err != nil { + return fmt.Errorf("preflight: %w", err) + } + if err := m.preflightGate(rep); err != nil { + return err + } + if err := m.buildDedup(ctx, rep); err != nil { + return fmt.Errorf("dedup: %w", err) + } + // Order is FK-safe; see PLAN.md §7. + phases := []struct { + name string + fn func(context.Context, *report) error + }{ + {"tenants", m.phaseTenants}, + {"entities.users", m.phaseUsers}, + {"entities.clients", m.phaseClients}, + {"credentials.devices", m.phaseDeviceCreds}, + {"resources.channels", m.phaseChannels}, + {"resources.rules", m.phaseRules}, + {"resources.reports", m.phaseReports}, + {"resources.alarms", m.phaseAlarms}, + {"object_groups", m.phaseGroups}, + {"group_membership", m.phaseGroupMembership}, + {"roles", m.phaseRoles}, + {"connections", m.phaseConnections}, + {"credentials.pats", m.phasePATs}, + {"tenant_invitations", m.phaseInvitations}, + {"backfill.tenant_actors", m.phaseBackfill}, + } + for _, p := range phases { + if err := p.fn(ctx, rep); err != nil { + return fmt.Errorf("phase %s: %w", p.name, err) + } + } + if m.apply && len(m.deviceKeys) > 0 { + if err := m.writeDeviceKeys(); err != nil { + return fmt.Errorf("write device keys: %w", err) + } + } + return nil +} + +// writeDeviceKeys exports the re-issued device API keys (secret shown once) for +// re-provisioning. Treat the file as a secret and delete after use. +func (m *migrator) writeDeviceKeys() error { + if err := os.MkdirAll(m.reportDir, 0o755); err != nil { + return err + } + path := filepath.Join(m.reportDir, "device-keys-"+time.Now().UTC().Format("20060102-150405")+".csv") + f, err := os.OpenFile(path, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, 0o600) + if err != nil { + return err + } + defer f.Close() + w := csv.NewWriter(f) + defer w.Flush() + if err := w.Write([]string{"client_id", "domain_id", "identity", "api_key"}); err != nil { + return err + } + return w.WriteAll(m.deviceKeys) +} + +func (m *migrator) loadLookups(ctx context.Context) error { + rows, err := m.atom.QueryxContext(ctx, + `SELECT p.key, p.id, pv.id + FROM profiles p + LEFT JOIN LATERAL ( + SELECT id + FROM profile_versions + WHERE profile_id = p.id + AND status = 'active' + ORDER BY version DESC + LIMIT 1 + ) pv ON TRUE + WHERE p.object_kind = 'entity'`) + if err != nil { + return err + } + for rows.Next() { + var k, id string + var versionID sql.NullString + if err := rows.Scan(&k, &id, &versionID); err != nil { + return err + } + m.profileID[k] = id + if versionID.Valid { + m.profileVersionID[k] = versionID.String + } + } + rows.Close() + + ar, err := m.atom.QueryxContext(ctx, `SELECT name, id FROM actions`) + if err != nil { + return err + } + defer ar.Close() + for ar.Next() { + var n, id string + if err := ar.Scan(&n, &id); err != nil { + return err + } + m.actionID[n] = id + } + return nil +} + +// exec runs a write in apply mode; in dry-run it is a no-op. Returns the row +// count target for reporting (always 1 on success path). +func (m *migrator) exec(ctx context.Context, query string, args ...any) error { + if !m.apply { + return nil + } + _, err := m.atom.ExecContext(ctx, query, args...) + return err +} + +// --- phases --- + +func (m *migrator) phaseTenants(ctx context.Context, rep *report) error { + rows, err := readDomains(ctx, m.domainsDB) + if err != nil { + return err + } + for _, d := range rows { + alias := aliasOrNil(m.tenantAlias[d.ID]) + if err := m.exec(ctx, + `INSERT INTO tenants (id, name, alias, status, tags, attributes, created_at, updated_at) + VALUES ($1,$2,$3,$4,$5,$6,$7,$8) ON CONFLICT (id) DO NOTHING`, + d.ID, m.tenantName[d.ID], alias, tenantStatus(d.Status), + pqArr(d.Tags), attrs(d.Metadata, nil), ntToTime(d.CreatedAt), ntPtr(d.UpdatedAt), + ); err != nil { + return err + } + m.tenants[d.ID] = true + rep.count("tenants", 1) + } + return nil +} + +func (m *migrator) phaseUsers(ctx context.Context, rep *report) error { + rows, err := readUsers(ctx, m.usersDB) + if err != nil { + return err + } + prof := nullStr(m.profileID["user"]) + profVer := nullStr(m.profileVersionID["user"]) + for _, u := range rows { + extra := map[string]any{} + putStr(extra, "first_name", u.FirstName) + putStr(extra, "last_name", u.LastName) + putStr(extra, "username", u.Username) + putStr(extra, "profile_picture", u.ProfilePicture) + putStr(extra, "auth_provider", u.AuthProvider) + name := m.userName[u.ID] + + if err := m.exec(ctx, + `INSERT INTO entities (id, kind, name, tenant_id, status, attributes, profile_id, profile_version_id, created_at, updated_at) + VALUES ($1,'human',$2,NULL,$3,$4,$5,$6,$7,$8) ON CONFLICT (id) DO NOTHING`, + u.ID, name, entityStatus(u.Status), attrs(u.Metadata, extra), prof, + profVer, ntToTime(u.CreatedAt), ntPtr(u.UpdatedAt), + ); err != nil { + return err + } + m.migratedUsers[u.ID] = true + rep.count("entities.users", 1) + + if err := m.exec(ctx, + `INSERT INTO principal_group_members (group_id, entity_id) + VALUES ($1,$2) ON CONFLICT DO NOTHING`, + authenticatedUsersGroupID, u.ID, + ); err != nil { + return err + } + rep.count("principal_group_members.authenticated_users", 1) + + // email (force-reset: no password credential migrated) + if u.Email.Valid && u.Email.String != "" { + if err := m.exec(ctx, + `INSERT INTO entity_emails (entity_id, email, verified_at) + VALUES ($1,$2,$3) ON CONFLICT (entity_id) DO NOTHING`, + u.ID, u.Email.String, ntPtr(u.VerifiedAt), + ); err != nil { + // likely global email-unique conflict + rep.warnf("user %s email %q not inserted (conflict?)", u.ID, u.Email.String) + rep.skip("email_conflict") + } else { + rep.count("entity_emails", 1) + } + } + rep.todo("password_reset", fmt.Sprintf("%s (%s)", u.ID, u.Email.String)) + + // platform admin role assignment + if u.Role.Valid && u.Role.Int16 == 1 { + if err := m.exec(ctx, + `INSERT INTO role_assignments (id, tenant_id, subject_kind, subject_id, role_id) + VALUES ($1,NULL,'entity',$2,'00000000-0000-0000-0000-000000000002') + ON CONFLICT DO NOTHING`, + derivedUUID("ra", "admin", u.ID), u.ID, + ); err != nil { + return err + } + rep.count("role_assignments", 1) + } + } + return nil +} + +func (m *migrator) phaseClients(ctx context.Context, rep *report) error { + rows, err := readClients(ctx, m.clientsDB) + if err != nil { + return err + } + prof := nullStr(m.profileID["client"]) + profVer := nullStr(m.profileVersionID["client"]) + for _, c := range rows { + if !m.tenants[c.DomainID] { + rep.skip("client_orphan_domain") + rep.warnf("client %s skipped: domain %s not migrated", c.ID, c.DomainID) + continue + } + m.clientDomain[c.ID] = c.DomainID + extra := map[string]any{} + putStr(extra, "identity", c.Identity) + putTags(extra, c.Tags) + putJSON(extra, "private_metadata", c.PrivateMeta) + alias := aliasOrNil(m.clientAlias[c.ID]) + name := m.deviceName[c.ID] + if err := m.exec(ctx, + `INSERT INTO entities (id, kind, name, tenant_id, status, attributes, profile_id, profile_version_id, alias, created_at, updated_at) + VALUES ($1,'device',$2,$3,$4,$5,$6,$7,$8,$9,$10) ON CONFLICT (id) DO NOTHING`, + c.ID, name, c.DomainID, entityStatus(c.Status), attrs(c.Metadata, extra), + prof, profVer, alias, ntToTime(c.CreatedAt), ntPtr(c.UpdatedAt), + ); err != nil { + return err + } + rep.count("entities.clients", 1) + } + return nil +} + +func (m *migrator) phaseDeviceCreds(ctx context.Context, rep *report) error { + rows, err := readClients(ctx, m.clientsDB) + if err != nil { + return err + } + for _, c := range rows { + if _, ok := m.clientDomain[c.ID]; !ok || !c.Secret.Valid || c.Secret.String == "" { + continue + } + // Atom cannot reuse the Magistrala secret (format + lookup differ); the + // device key is re-issued in atom__ form and exported for + // re-provisioning. See newAtomAPIKey / PLAN §5. + credID, plaintext, hash, err := newAtomAPIKey(c.ID) + if err != nil { + return err + } + if err := m.exec(ctx, + `INSERT INTO credentials (id, entity_id, kind, identifier, secret_hash, metadata, status) + VALUES ($1,$2,'api_key',$3,$4,'{"source":"magistrala-client-reissued"}',$5) + ON CONFLICT (id) DO NOTHING`, + credID, c.ID, c.ID, hash, statusCred(c.Status), + ); err != nil { + return err + } + if m.apply { + m.deviceKeys = append(m.deviceKeys, []string{c.ID, c.DomainID, c.Identity.String, plaintext}) + } + rep.count("credentials.devices", 1) + } + if len(m.deviceKeys) > 0 { + rep.todo("device_reprovision", fmt.Sprintf("%d device keys re-issued -> see device-keys CSV in report dir", len(m.deviceKeys))) + } + return nil +} + +func (m *migrator) phaseChannels(ctx context.Context, rep *report) error { + rows, err := readChannels(ctx, m.channelsDB) + if err != nil { + return err + } + for _, ch := range rows { + if !m.tenants[ch.DomainID] { + rep.skip("channel_orphan_domain") + continue + } + m.channelDomain[ch.ID] = ch.DomainID + m.resourceDomain[ch.ID] = ch.DomainID + alias := aliasOrNil(m.channelAlias[ch.ID]) + owner := sql.NullString{} + if ch.CreatedBy.Valid && m.migratedUsers[ch.CreatedBy.String] { + owner = ch.CreatedBy + } + extra := map[string]any{"status": entityStatus(ch.Status)} + putTags(extra, ch.Tags) + if err := m.exec(ctx, + `INSERT INTO resources (id, kind, name, tenant_id, owner_id, attributes, alias, created_at, updated_at) + VALUES ($1,'channel',$2,$3,$4,$5,$6,$7,$8) ON CONFLICT (id) DO NOTHING`, + ch.ID, nsToStr(ch.Name), ch.DomainID, owner, attrs(ch.Metadata, extra), + alias, ntToTime(ch.CreatedAt), ntPtr(ch.UpdatedAt), + ); err != nil { + return err + } + rep.count("resources.channels", 1) + } + return nil +} + +// insertResource writes one row into Atom resources (kind = channel/rule/report/ +// alarm). Entity-specific columns Magistrala has but Atom resources lack are +// folded into the attributes JSONB. ON CONFLICT (id) keeps it idempotent. +func (m *migrator) insertResource(ctx context.Context, id, kind, name, tenant string, owner sql.NullString, attributes string, createdAt time.Time, updatedAt any) error { + return m.exec(ctx, + `INSERT INTO resources (id, kind, name, tenant_id, owner_id, attributes, created_at, updated_at) + VALUES ($1,$2,$3,$4,$5,$6,$7,$8) ON CONFLICT (id) DO NOTHING`, + id, kind, name, tenant, owner, attributes, createdAt, updatedAt) +} + +// ownerOf returns created_by only when it maps to a migrated human (resources. +// owner_id is an FK to entities); otherwise NULL. +func (m *migrator) ownerOf(createdBy sql.NullString) sql.NullString { + if createdBy.Valid && m.migratedUsers[createdBy.String] { + return createdBy + } + return sql.NullString{} +} + +// uniqueResName makes a resource name unique within a tenant (Atom enforces +// resources(name, tenant_id); rules/reports/alarms carry no such Magistrala +// constraint, so same-tenant dups are possible). Suffixes -2, -3, … on collision. +func uniqueResName(seen map[string]int, tenant, name string) string { + key := tenant + "|" + strings.ToLower(name) + seen[key]++ + if seen[key] == 1 { + return name + } + return fmt.Sprintf("%s-%d", name, seen[key]) +} + +// phaseRules: rules_engine.rules -> resources (kind=rule). +func (m *migrator) phaseRules(ctx context.Context, rep *report) error { + rows, err := readRules(ctx, m.reDB) + if err != nil { + return err + } + seen := map[string]int{} + for _, r := range rows { + if !m.tenants[r.DomainID] { + rep.skip("rule_orphan_domain") + continue + } + m.resourceDomain[r.ID] = r.DomainID + name := uniqueResName(seen, r.DomainID, firstNonEmpty(nsToStr(r.Name), r.ID)) + extra := map[string]any{ + "status": entityStatus(r.Status), + "logic_type": r.LogicType, + } + putStr(extra, "input_channel", r.InputChannel) + putStr(extra, "input_topic", r.InputTopic) + putStr(extra, "updated_by", r.UpdatedBy) + if len(r.Outputs) > 0 { + extra["outputs"] = r.Outputs + } + if len(r.LogicValue) > 0 { + extra["logic_value"] = r.LogicValue // base64-encoded in JSON + } + if r.Recurring.Valid { + extra["recurring"] = r.Recurring.Int16 + } + if r.RecurringPeriod.Valid { + extra["recurring_period"] = r.RecurringPeriod.Int16 + } + if r.Time.Valid { + extra["time"] = r.Time.Time + } + if r.StartDatetime.Valid { + extra["start_datetime"] = r.StartDatetime.Time + } + if len(r.Tags) > 0 { + extra["tags"] = []string(r.Tags) + } + if err := m.insertResource(ctx, r.ID, "rule", name, r.DomainID, m.ownerOf(r.CreatedBy), + attrs(r.Metadata, extra), ntToTime(r.CreatedAt), ntPtr(r.UpdatedAt)); err != nil { + return err + } + rep.count("resources.rules", 1) + } + return nil +} + +// phaseReports: reports.report_config -> resources (kind=report). +func (m *migrator) phaseReports(ctx context.Context, rep *report) error { + rows, err := readReports(ctx, m.reportsDB) + if err != nil { + return err + } + seen := map[string]int{} + for _, rp := range rows { + if !m.tenants[rp.DomainID] { + rep.skip("report_orphan_domain") + continue + } + m.resourceDomain[rp.ID] = rp.DomainID + name := uniqueResName(seen, rp.DomainID, firstNonEmpty(nsToStr(rp.Name), rp.ID)) + extra := map[string]any{"status": entityStatus(rp.Status)} + putStr(extra, "description", rp.Description) + putStr(extra, "report_template", rp.ReportTemplate) + putStr(extra, "updated_by", rp.UpdatedBy) + if len(rp.Config) > 0 { + extra["config"] = rp.Config + } + if len(rp.Email) > 0 { + extra["email"] = rp.Email + } + if len(rp.Metrics) > 0 { + extra["metrics"] = rp.Metrics + } + if rp.Due.Valid { + extra["due"] = rp.Due.Time + } + if rp.Recurring.Valid { + extra["recurring"] = rp.Recurring.Int16 + } + if rp.RecurringPeriod.Valid { + extra["recurring_period"] = rp.RecurringPeriod.Int16 + } + if rp.StartDatetime.Valid { + extra["start_datetime"] = rp.StartDatetime.Time + } + if err := m.insertResource(ctx, rp.ID, "report", name, rp.DomainID, m.ownerOf(rp.CreatedBy), + attrs(nil, extra), ntToTime(rp.CreatedAt), ntPtr(rp.UpdatedAt)); err != nil { + return err + } + rep.count("resources.reports", 1) + } + return nil +} + +// phaseAlarms: alarms.alarms -> resources (kind=alarm). Alarms have no name in +// Magistrala; the measurement is used (id fallback), deduped per tenant. +func (m *migrator) phaseAlarms(ctx context.Context, rep *report) error { + rows, err := readAlarms(ctx, m.alarmsDB) + if err != nil { + return err + } + seen := map[string]int{} + for _, a := range rows { + if !m.tenants[a.DomainID] { + rep.skip("alarm_orphan_domain") + continue + } + m.resourceDomain[a.ID] = a.DomainID + name := uniqueResName(seen, a.DomainID, firstNonEmpty(a.Measurement, a.ID)) + extra := map[string]any{ + "rule_id": a.RuleID, + "channel_id": a.ChannelID, + "client_id": a.ClientID, + "subtopic": a.Subtopic, + "measurement": a.Measurement, + "value": a.Value, + "unit": a.Unit, + "threshold": a.Threshold, + "cause": a.Cause, + "alarm_status": a.Status, + "severity": a.Severity, + } + putStr(extra, "assignee_id", a.AssigneeID) + putStr(extra, "updated_by", a.UpdatedBy) + putStr(extra, "assigned_by", a.AssignedBy) + putStr(extra, "acknowledged_by", a.AcknowledgedBy) + putStr(extra, "resolved_by", a.ResolvedBy) + if a.AssignedAt.Valid { + extra["assigned_at"] = a.AssignedAt.Time + } + if a.AcknowledgedAt.Valid { + extra["acknowledged_at"] = a.AcknowledgedAt.Time + } + if a.ResolvedAt.Valid { + extra["resolved_at"] = a.ResolvedAt.Time + } + // Alarms carry no created_by; owner_id stays NULL. + if err := m.insertResource(ctx, a.ID, "alarm", name, a.DomainID, sql.NullString{}, + attrs(a.Metadata, extra), ntToTime(a.CreatedAt), ntPtr(a.UpdatedAt)); err != nil { + return err + } + rep.count("resources.alarms", 1) + } + return nil +} + +func (m *migrator) phaseGroups(ctx context.Context, rep *report) error { + rows, err := readGroups(ctx, m.groupsDB) + if err != nil { + return err + } + // First pass: groups. Second pass: hierarchy (parent must exist). + for _, g := range rows { + if !m.tenants[g.DomainID] { + rep.skip("group_orphan_domain") + continue + } + m.groupDomain[g.ID] = g.DomainID + if err := m.exec(ctx, + `INSERT INTO object_groups (id, name, tenant_id, description, status, attributes, created_at, updated_at) + VALUES ($1,$2,$3,$4,$5,$6,$7,$8) ON CONFLICT (id) DO NOTHING`, + g.ID, m.groupName[g.ID], g.DomainID, nsToStr(g.Description), entityStatus(g.Status), + attrs(g.Metadata, map[string]any{"tags": []string(g.Tags)}), + ntToTime(g.CreatedAt), ntToTime(g.UpdatedAt), // object_groups.updated_at is NOT NULL + ); err != nil { + return err + } + rep.count("object_groups", 1) + } + for _, g := range rows { + _, haveChild := m.groupDomain[g.ID] + _, haveParent := m.groupDomain[g.ParentID.String] + if !g.ParentID.Valid || !haveChild || !haveParent { + continue + } + if err := m.exec(ctx, + `INSERT INTO object_group_hierarchy (parent_id, child_id, tenant_id) + VALUES ($1,$2,$3) ON CONFLICT (child_id) DO NOTHING`, + g.ParentID.String, g.ID, g.DomainID, + ); err != nil { + return err + } + rep.count("object_group_hierarchy", 1) + } + return nil +} + +func (m *migrator) phaseGroupMembership(ctx context.Context, rep *report) error { + // clients -> object_group_entities + clients, err := readClients(ctx, m.clientsDB) + if err != nil { + return err + } + for _, c := range clients { + _, haveClient := m.clientDomain[c.ID] + _, haveGroup := m.groupDomain[c.ParentGroupID.String] + if !c.ParentGroupID.Valid || !haveClient || !haveGroup { + continue + } + if err := m.exec(ctx, + `INSERT INTO object_group_entities (group_id, entity_id, tenant_id) + VALUES ($1,$2,$3) ON CONFLICT (entity_id) DO NOTHING`, + c.ParentGroupID.String, c.ID, c.DomainID, + ); err != nil { + return err + } + rep.count("object_group_entities", 1) + } + // channels -> object_group_resources + chans, err := readChannels(ctx, m.channelsDB) + if err != nil { + return err + } + for _, ch := range chans { + _, haveChan := m.channelDomain[ch.ID] + _, haveGroup := m.groupDomain[ch.ParentGroupID.String] + if !ch.ParentGroupID.Valid || !haveChan || !haveGroup { + continue + } + if err := m.exec(ctx, + `INSERT INTO object_group_resources (group_id, resource_id, tenant_id) + VALUES ($1,$2,$3) ON CONFLICT (resource_id) DO NOTHING`, + ch.ParentGroupID.String, ch.ID, ch.DomainID, + ); err != nil { + return err + } + rep.count("object_group_resources", 1) + } + return nil +} + +// roleScope describes how a Magistrala role family maps to a permission_block. +type roleScope struct { + prefix string + db *sqlx.DB + domainOf func(entityID string) (string, bool) // object id -> tenant id + scopeMode string + objKind string // for object scope +} + +func (m *migrator) phaseRoles(ctx context.Context, rep *report) error { + families := []roleScope{ + {"domains", m.domainsDB, func(id string) (string, bool) { return id, m.tenants[id] }, "tenant", ""}, + {"clients", m.clientsDB, func(id string) (string, bool) { d, ok := m.clientDomain[id]; return d, ok }, "object", "entity"}, + {"channels", m.channelsDB, func(id string) (string, bool) { d, ok := m.channelDomain[id]; return d, ok }, "object", "resource"}, + {"rules", m.reDB, func(id string) (string, bool) { d, ok := m.resourceDomain[id]; return d, ok }, "object", "resource"}, + {"reports", m.reportsDB, func(id string) (string, bool) { d, ok := m.resourceDomain[id]; return d, ok }, "object", "resource"}, + {"groups", m.groupsDB, func(id string) (string, bool) { d, ok := m.groupDomain[id]; return d, ok }, "group", ""}, + } + for _, f := range families { + if err := m.migrateRoleFamily(ctx, rep, f); err != nil { + return fmt.Errorf("%s roles: %w", f.prefix, err) + } + } + return nil +} + +func (m *migrator) migrateRoleFamily(ctx context.Context, rep *report, f roleScope) error { + roles, acts, mems, err := readRoleFamily(ctx, f.db, f.prefix) + if err != nil { + return err + } + actsByRole := map[string][]string{} + for _, a := range acts { + actsByRole[a.RoleID] = append(actsByRole[a.RoleID], a.Action) + } + for _, r := range roles { + tenant, ok := f.domainOf(r.EntityID) + if !ok { + rep.skip("role_orphan_object") + continue + } + roleID := derivedUUID("role", f.prefix, r.ID) + + // Atom roles are unique on (name, tenant_id). Magistrala object roles are + // per-instance, so many objects in one tenant can share a role name + // (e.g. every client has an "admin" role). Embed the object id to keep the + // Atom role name unique within the tenant. + roleName := f.prefix + ":" + r.EntityID + ":" + r.Name + if f.scopeMode == "tenant" { + roleName = f.prefix + ":" + r.Name // domain roles: one set per tenant + } + if err := m.exec(ctx, + `INSERT INTO roles (id, name, tenant_id, description, created_at, updated_at) + VALUES ($1,$2,$3,$4,$5,$6) ON CONFLICT (id) DO NOTHING`, + roleID, roleName, tenant, "migrated from "+f.prefix+"_roles "+r.ID, + ntToTime(r.CreatedAt), ntPtr(r.UpdatedAt), + ); err != nil { + return err + } + m.migratedRoles[roleID] = true + rep.count("roles", 1) + + // actions -> permission_block_actions + seen := map[string]bool{} // block id + atom action + insertedBlocks := map[string]bool{} + for _, raw := range actsByRole[r.ID] { + atom, ok := mapAction(raw) + if !ok { + if m.cfg.UnmappedAction == "skip" { + rep.skip("action_unmapped_skipped") + rep.warnf("unmapped action %q (%s) skipped", raw, f.prefix) + continue + } + atom = "manage" + rep.warnf("unmapped action %q (%s) -> manage", raw, f.prefix) + } + aid, ok := m.actionID[atom] + if !ok { + continue + } + for _, block := range f.blockPlans(r.ID, r.EntityID, tenant, raw) { + if !insertedBlocks[block.ID] { + if err := m.insertBlock(ctx, block); err != nil { + return err + } + if err := m.exec(ctx, + `INSERT INTO role_permission_blocks (role_id, permission_block_id) + VALUES ($1,$2) ON CONFLICT DO NOTHING`, roleID, block.ID); err != nil { + return err + } + insertedBlocks[block.ID] = true + } + seenKey := block.ID + "|" + atom + if seen[seenKey] { + continue + } + seen[seenKey] = true + if err := m.exec(ctx, + `INSERT INTO permission_block_actions (permission_block_id, action_id) + VALUES ($1,$2) ON CONFLICT DO NOTHING`, block.ID, aid); err != nil { + return err + } + rep.count("permission_block_actions", 1) + } + } + + // members -> role_assignments (+ tenant_memberships for domain roles) + for _, mem := range mems { + if mem.RoleID != r.ID { + continue + } + if !m.migratedUsers[mem.MemberID] { + rep.skip("role_orphan_member") + rep.warnf("%s role %s member %s skipped: user not migrated", f.prefix, r.ID, mem.MemberID) + continue + } + if err := m.exec(ctx, + `INSERT INTO role_assignments (id, tenant_id, subject_kind, subject_id, role_id) + VALUES ($1,$2,'entity',$3,$4) ON CONFLICT DO NOTHING`, + derivedUUID("ra", f.prefix, r.ID, mem.MemberID), tenant, mem.MemberID, roleID, + ); err != nil { + return err + } + rep.count("role_assignments", 1) + if f.prefix == "domains" { + if err := m.exec(ctx, + `INSERT INTO tenant_memberships (tenant_id, entity_id, status) + VALUES ($1,$2,'active') ON CONFLICT DO NOTHING`, + tenant, mem.MemberID); err != nil { + return err + } + rep.count("tenant_memberships", 1) + } + } + } + return nil +} + +type permissionBlockPlan struct { + ID string + TenantID string + ScopeMode string + ObjectKind any + ObjectType any + ObjectID any + GroupID any +} + +func (f roleScope) blockPlans(roleID, objectID, tenant, rawAction string) []permissionBlockPlan { + if f.scopeMode != "group" { + blockID := derivedUUID("block", f.prefix, roleID) + switch f.scopeMode { + case "tenant": + return []permissionBlockPlan{{ + ID: blockID, TenantID: tenant, ScopeMode: "tenant", + }} + default: + return []permissionBlockPlan{{ + ID: blockID, TenantID: tenant, ScopeMode: "object", + ObjectKind: f.objKind, ObjectID: objectID, + }} + } + } + + action := strings.ToLower(strings.TrimSpace(rawAction)) + switch { + case strings.HasPrefix(action, "subgroup_client"): + return []permissionBlockPlan{groupObjectBlock(roleID, tenant, objectID, "descendant", "entity", "entity:device")} + case strings.HasPrefix(action, "client"): + return []permissionBlockPlan{groupObjectBlock(roleID, tenant, objectID, "direct", "entity", "entity:device")} + case strings.HasPrefix(action, "subgroup_channel"): + return []permissionBlockPlan{groupObjectBlock(roleID, tenant, objectID, "descendant", "resource", "resource:channel")} + case strings.HasPrefix(action, "channel"): + return []permissionBlockPlan{groupObjectBlock(roleID, tenant, objectID, "direct", "resource", "resource:channel")} + case strings.HasPrefix(action, "subgroup"): + return []permissionBlockPlan{groupKindBlock(roleID, tenant, objectID, "descendant")} + default: + return []permissionBlockPlan{{ + ID: derivedUUID("block", "groups", roleID, "self"), + TenantID: tenant, ScopeMode: "object", ObjectKind: "group", ObjectID: objectID, + }} + } +} + +func groupObjectBlock(roleID, tenant, groupID, depth, objectKind, objectType string) permissionBlockPlan { + scopeMode := "group_direct_objects" + if depth == "descendant" { + scopeMode = "group_descendant_objects" + } + return permissionBlockPlan{ + ID: derivedUUID("block", "groups", roleID, depth, objectType), + TenantID: tenant, + ScopeMode: scopeMode, + ObjectKind: objectKind, + ObjectType: objectType, + GroupID: groupID, + } +} + +func groupKindBlock(roleID, tenant, groupID, depth string) permissionBlockPlan { + scopeMode := "group_child_groups" + if depth == "descendant" { + scopeMode = "group_descendant_groups" + } + return permissionBlockPlan{ + ID: derivedUUID("block", "groups", roleID, depth, "group"), + TenantID: tenant, + ScopeMode: scopeMode, + GroupID: groupID, + } +} + +func (m *migrator) insertBlock(ctx context.Context, b permissionBlockPlan) error { + return m.exec(ctx, + `INSERT INTO permission_blocks + (id, tenant_id, scope_mode, object_kind, object_type, object_id, group_id, effect, conditions) + VALUES ($1,$2,$3,$4,$5,$6,$7,'allow','{}') ON CONFLICT (id) DO NOTHING`, + b.ID, b.TenantID, b.ScopeMode, b.ObjectKind, b.ObjectType, b.ObjectID, b.GroupID) +} + +func (m *migrator) phaseConnections(ctx context.Context, rep *report) error { + // Both clients-db and channels-db keep their own connections copy; union and + // dedup so neither side's view is missed. + cliConns, err := readConnections(ctx, m.clientsDB) + if err != nil { + return err + } + chConns, err := readConnections(ctx, m.channelsDB) + if err != nil { + return err + } + seenConn := map[string]bool{} + conns := make([]srcConnection, 0, len(cliConns)+len(chConns)) + for _, c := range append(cliConns, chConns...) { + k := c.ChannelID + "|" + c.ClientID + "|" + c.DomainID + "|" + fmt.Sprint(c.Type) + if seenConn[k] { + continue + } + seenConn[k] = true + conns = append(conns, c) + } + for _, c := range conns { + dom, ok := m.channelDomain[c.ChannelID] + if !ok { + rep.skip("conn_orphan_channel") + continue + } + clientDom, ok := m.clientDomain[c.ClientID] + if !ok { + rep.skip("conn_orphan_client") + continue + } + if clientDom != dom || c.DomainID != dom { + rep.skip("conn_domain_mismatch") + rep.warnf("connection client=%s channel=%s skipped: connection domain=%s client domain=%s channel domain=%s", c.ClientID, c.ChannelID, c.DomainID, clientDom, dom) + continue + } + act, ok := connectionAction(c.Type) + if !ok { + rep.skip("conn_bad_type") + continue + } + aid, ok := m.actionID[act] + if !ok { + continue + } + blockID := derivedUUID("connblock", c.ChannelID, fmt.Sprint(c.Type)) + if err := m.exec(ctx, + `INSERT INTO permission_blocks (id, tenant_id, scope_mode, object_kind, object_id, effect, conditions) + VALUES ($1,$2,'object','resource',$3,'allow','{}') ON CONFLICT (id) DO NOTHING`, + blockID, dom, c.ChannelID); err != nil { + return err + } + if err := m.exec(ctx, + `INSERT INTO permission_block_actions (permission_block_id, action_id) + VALUES ($1,$2) ON CONFLICT DO NOTHING`, blockID, aid); err != nil { + return err + } + if err := m.exec(ctx, + `INSERT INTO direct_policies (id, tenant_id, subject_kind, subject_id, permission_block_id) + VALUES ($1,$2,'entity',$3,$4) ON CONFLICT (id) DO NOTHING`, + derivedUUID("dp", c.ChannelID, c.ClientID, fmt.Sprint(c.Type)), dom, c.ClientID, blockID, + ); err != nil { + return err + } + rep.count("connections", 1) + } + return nil +} + +func (m *migrator) phasePATs(ctx context.Context, rep *report) error { + pats, err := readPATs(ctx, m.authDB) + if err != nil { + return err + } + scopeRows, err := readPATScopes(ctx, m.authDB) + if err != nil { + return err + } + scopesByPAT := map[string][]map[string]string{} + for _, s := range scopeRows { + scopesByPAT[s.PatID] = append(scopesByPAT[s.PatID], map[string]string{ + "domain_id": s.DomainID.String, "entity_type": s.EntityType, + "operation": s.Operation, "entity_id": s.EntityID, + }) + } + for _, p := range pats { + if !p.UserID.Valid || !m.migratedUsers[p.UserID.String] { + rep.skip("pat_orphan_user") + continue + } + status := "active" + if p.Revoked.Valid && p.Revoked.Bool { + status = "revoked" + } + meta, err := json.Marshal(map[string]any{ + "source": "magistrala-pat", "name": p.Name, "description": p.Desc.String, + "needs_reissue": true, "scopes": scopesByPAT[p.ID], + }) + if err != nil { + return err + } + // secret_hash NULL: Magistrala PAT secret is not convertible (see PLAN §5). + if err := m.exec(ctx, + `INSERT INTO credentials (id, entity_id, kind, identifier, metadata, status, expires_at) + VALUES ($1,$2,'api_key',$3,$4,$5,$6) ON CONFLICT (id) DO NOTHING`, + p.ID, p.UserID.String, p.ID, string(meta), status, ntPtr(p.ExpiresAt), + ); err != nil { + return err + } + rep.count("credentials.pats", 1) + rep.todo("pat_reissue", fmt.Sprintf("%s (user %s)", p.ID, p.UserID.String)) + } + return nil +} + +func (m *migrator) phaseInvitations(ctx context.Context, rep *report) error { + invs, err := readInvitations(ctx, m.domainsDB) + if err != nil { + return err + } + for _, iv := range invs { + // Only pending invitations matter; accepted ones are already memberships. + if iv.ConfirmedAt.Valid { + rep.skip("invitation_already_accepted") + continue + } + if !m.tenants[iv.DomainID] || !m.migratedUsers[iv.InvitedBy] { + rep.skip("invitation_orphan") + continue + } + invitee := sql.NullString{} + if m.migratedUsers[iv.InviteeID] { + invitee = sql.NullString{String: iv.InviteeID, Valid: true} + } + roleID := derivedUUID("role", "domains", iv.RoleID) + roleArg := any(roleID) + if !m.migratedRoles[roleID] { + roleArg = nil + rep.skip("invitation_orphan_role") + rep.warnf("invitation for tenant %s invitee %s has missing role %s; role_id set NULL", iv.DomainID, iv.InviteeID, iv.RoleID) + } + if err := m.exec(ctx, + `INSERT INTO tenant_invitations (id, tenant_id, invitee_user_id, invited_by, role_id, created_at, rejected_at) + VALUES ($1,$2,$3,$4,$5,$6,$7) ON CONFLICT (id) DO NOTHING`, + derivedUUID("inv", iv.DomainID, iv.InviteeID), iv.DomainID, invitee, iv.InvitedBy, + roleArg, ntToTime(iv.CreatedAt), ntPtr(iv.RejectedAt), + ); err != nil { + return err + } + rep.count("tenant_invitations", 1) + } + return nil +} + +// phaseBackfill sets tenants.created_by/updated_by now that entities exist (the +// FK could not be satisfied during phaseTenants). Only points at migrated users. +func (m *migrator) phaseBackfill(ctx context.Context, rep *report) error { + doms, err := readDomains(ctx, m.domainsDB) + if err != nil { + return err + } + for _, d := range doms { + if !m.tenants[d.ID] { + continue + } + cb := actorOrNil(d.CreatedBy, m.migratedUsers) + ub := actorOrNil(d.UpdatedBy, m.migratedUsers) + if cb == nil && ub == nil { + continue + } + if err := m.exec(ctx, + `UPDATE tenants SET created_by = COALESCE($2, created_by), + updated_by = COALESCE($3, updated_by) WHERE id = $1`, + d.ID, cb, ub); err != nil { + return err + } + rep.count("backfill.tenant_actors", 1) + } + return nil +} + +func actorOrNil(ns sql.NullString, migrated map[string]bool) any { + if ns.Valid && migrated[ns.String] { + return ns.String + } + return nil +} + +// --- small helpers --- + +func nsToStr(ns sql.NullString) string { + if ns.Valid { + return ns.String + } + return "" +} + +func nullStr(s string) any { + if s == "" { + return nil + } + return s +} + +func ntToTime(nt sql.NullTime) time.Time { + if nt.Valid { + return nt.Time + } + return time.Now().UTC() +} + +func ntPtr(nt sql.NullTime) any { + if nt.Valid { + return nt.Time + } + return nil +} + +func pqArr(a []string) any { + if len(a) == 0 { + return "{}" + } + out := "{" + for i, s := range a { + if i > 0 { + out += "," + } + out += `"` + s + `"` + } + return out + "}" +} + +func putStr(m map[string]any, k string, v sql.NullString) { + if v.Valid && v.String != "" { + m[k] = v.String + } +} + +func putTags(m map[string]any, tags []string) { + if len(tags) > 0 { + m["tags"] = []string(tags) + } +} + +func putJSON(m map[string]any, k string, b []byte) { + if len(b) > 0 { + m[k] = json.RawMessage(b) + } +} + +func firstNonEmpty(vals ...string) string { + for _, v := range vals { + if v != "" { + return v + } + } + return "" +} + +func statusCred(s int16) string { + if s == 0 { + return "active" + } + return "revoked" +} + +// attrs merges Magistrala metadata jsonb with extra keys into an Atom attributes +// JSON string. +func attrs(meta []byte, extra map[string]any) string { + out := map[string]any{} + if len(meta) > 0 { + _ = json.Unmarshal(meta, &out) + } + for k, v := range extra { + out[k] = v + } + b, err := json.Marshal(out) + if err != nil { + return "{}" + } + return string(b) +} diff --git a/tools/atom-migration/preflight.go b/tools/atom-migration/preflight.go new file mode 100644 index 000000000..c7815aa04 --- /dev/null +++ b/tools/atom-migration/preflight.go @@ -0,0 +1,318 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package main + +import ( + "context" + "fmt" + "strings" +) + +// preflight runs read-only data-quality checks before any write. Blocking issues +// (rep.block) abort an --apply run; warnings (rep.warn) are advisory. See PLAN §6. +func (m *migrator) preflight(ctx context.Context, rep *report) error { + checks := []func(context.Context, *report) error{ + m.pfEmails, + m.pfHumanNames, + m.pfTenantNames, + m.pfTenantAlias, + m.pfEntityResourceAlias, + m.pfClientNames, + m.pfGroupNames, + m.pfOrphans, + } + for _, c := range checks { + if err := c(ctx, rep); err != nil { + return err + } + } + return nil +} + +// dupGroups returns the keys that appear more than once (with their count). +func dupGroups(keyOf func() []string) map[string]int { + seen := map[string]int{} + for _, k := range keyOf() { + if k != "" { + seen[k]++ + } + } + for k, n := range seen { + if n < 2 { + delete(seen, k) + } + } + return seen +} + +// pfEmails: Atom entity_emails.email is globally UNIQUE. Magistrala enforces this +// in the users table, but a dump merged across instances can break it. +func (m *migrator) pfEmails(ctx context.Context, rep *report) error { + users, err := readUsers(ctx, m.usersDB) + if err != nil { + return err + } + dups := dupGroups(func() []string { + out := make([]string, 0, len(users)) + for _, u := range users { + if u.Email.Valid { + out = append(out, strings.ToLower(strings.TrimSpace(u.Email.String))) + } + } + return out + }) + for email, n := range dups { + rep.blockf("email %q used by %d users (entity_emails.email is UNIQUE)", email, n) + } + return nil +} + +// pfHumanNames: entities(name, tenant_id) is UNIQUE; humans have tenant_id NULL so +// they share one global name namespace. name = first of username/email/id. +func (m *migrator) pfHumanNames(ctx context.Context, rep *report) error { + users, err := readUsers(ctx, m.usersDB) + if err != nil { + return err + } + dups := dupGroups(func() []string { + out := make([]string, 0, len(users)) + for _, u := range users { + out = append(out, firstNonEmpty(u.Username.String, u.Email.String, u.ID)) + } + return out + }) + for name, n := range dups { + // entities(name, tenant_id) is NULLS DISTINCT and humans have tenant_id + // NULL, so duplicate human names do NOT violate the index — advisory only. + rep.warnf("human entity name %q used by %d users (allowed; tenant NULL is NULLS DISTINCT)", name, n) + } + return nil +} + +// pfTenantNames: tenants.name is UNIQUE and NOT NULL. Magistrala domain.name is +// nullable and non-unique, so empty or duplicate names break the load. +func (m *migrator) pfTenantNames(ctx context.Context, rep *report) error { + doms, err := readDomains(ctx, m.domainsDB) + if err != nil { + return err + } + names := make([]string, 0, len(doms)) + for _, d := range doms { + if !d.Name.Valid || strings.TrimSpace(d.Name.String) == "" { + rep.warnf("domain %s has empty name -> will use its id (tenants.name is NOT NULL UNIQUE)", d.ID) + continue + } + names = append(names, d.Name.String) + } + for name, n := range dupGroups(func() []string { return names }) { + rep.warnf("domain name %q used by %d domains -> duplicates auto-renamed (tenants.name is UNIQUE)", name, n) + } + return nil +} + +// pfGroupNames: object_groups(name, tenant_id) is UNIQUE. Magistrala dropped the +// groups (domain_id, name) constraint, so same-domain dups are possible. +func (m *migrator) pfGroupNames(ctx context.Context, rep *report) error { + grps, err := readGroups(ctx, m.groupsDB) + if err != nil { + return err + } + keys := make([]string, 0, len(grps)) + for _, g := range grps { + keys = append(keys, g.DomainID+"|"+g.Name) + } + for k, n := range dupGroups(func() []string { return keys }) { + rep.warnf("group name collision (%s) across %d groups in one tenant -> duplicates auto-renamed", k, n) + } + return nil +} + +// pfTenantAlias: domain.route -> tenants.alias. Globally unique, case-folded, +// slug-shaped, not UUID-shaped. Invalid shape => alias dropped (warn). Case-fold +// collision among otherwise-valid aliases => block. +func (m *migrator) pfTenantAlias(ctx context.Context, rep *report) error { + doms, err := readDomains(ctx, m.domainsDB) + if err != nil { + return err + } + valid := []string{} + for _, d := range doms { + if !d.Route.Valid || d.Route.String == "" { + continue + } + a, ok := normalizeAlias(d.Route.String) + if !ok { + rep.warnf("tenant %s alias %q invalid slug -> dropped", d.ID, d.Route.String) + continue + } + valid = append(valid, a) + } + for a, n := range dupGroups(func() []string { return valid }) { + rep.warnf("tenant alias %q collides case-insensitively across %d domains -> duplicates auto-suffixed", a, n) + } + return nil +} + +// pfEntityResourceAlias: client.identity / channel.route are unique per tenant +// (case-folded). Collision within a domain => block (would violate Atom's unique +// index mid-apply). Invalid shape => warn (dropped). +func (m *migrator) pfEntityResourceAlias(ctx context.Context, rep *report) error { + clients, err := readClients(ctx, m.clientsDB) + if err != nil { + return err + } + chans, err := readChannels(ctx, m.channelsDB) + if err != nil { + return err + } + // key = domain|alias + cKeys := []string{} + for _, c := range clients { + if !c.Identity.Valid || c.Identity.String == "" { + continue + } + a, ok := normalizeAlias(c.Identity.String) + if !ok { + rep.warnf("client %s alias %q invalid slug -> dropped", c.ID, c.Identity.String) + continue + } + cKeys = append(cKeys, c.DomainID+"|"+a) + } + for k, n := range dupGroups(func() []string { return cKeys }) { + rep.warnf("device alias collision (%s) across %d clients in one tenant -> duplicates auto-suffixed", k, n) + } + rKeys := []string{} + for _, ch := range chans { + if !ch.Route.Valid || ch.Route.String == "" { + continue + } + a, ok := normalizeAlias(ch.Route.String) + if !ok { + rep.warnf("channel %s alias %q invalid slug -> dropped", ch.ID, ch.Route.String) + continue + } + rKeys = append(rKeys, ch.DomainID+"|"+a) + } + for k, n := range dupGroups(func() []string { return rKeys }) { + rep.warnf("channel alias collision (%s) across %d channels in one tenant -> duplicates auto-suffixed", k, n) + } + return nil +} + +// pfClientNames: device entities are unique on (name, tenant_id). Magistrala +// dropped the (domain_id, name) unique constraint, so same-domain name dups are +// possible and would break the Atom insert. +func (m *migrator) pfClientNames(ctx context.Context, rep *report) error { + clients, err := readClients(ctx, m.clientsDB) + if err != nil { + return err + } + keys := []string{} + for _, c := range clients { + keys = append(keys, c.DomainID+"|"+firstNonEmpty(c.Name.String, c.ID)) + } + for k, n := range dupGroups(func() []string { return keys }) { + rep.warnf("device name collision (%s) across %d clients in one tenant -> duplicates auto-renamed", k, n) + } + return nil +} + +// pfOrphans: clients/channels/groups whose domain_id has no surviving domain are +// skipped during load. Advisory only. +func (m *migrator) pfOrphans(ctx context.Context, rep *report) error { + doms, err := readDomains(ctx, m.domainsDB) + if err != nil { + return err + } + domSet := map[string]bool{} + for _, d := range doms { + domSet[d.ID] = true + } + count := func(get func() []string, label string) { + n := 0 + for _, id := range get() { + if !domSet[id] { + n++ + } + } + if n > 0 { + rep.warnf("%d %s reference a missing domain -> will be skipped", n, label) + } + } + clients, err := readClients(ctx, m.clientsDB) + if err != nil { + return err + } + count(func() []string { + out := make([]string, len(clients)) + for i, c := range clients { + out[i] = c.DomainID + } + return out + }, "clients") + chans, err := readChannels(ctx, m.channelsDB) + if err != nil { + return err + } + count(func() []string { + out := make([]string, len(chans)) + for i, c := range chans { + out[i] = c.DomainID + } + return out + }, "channels") + grps, err := readGroups(ctx, m.groupsDB) + if err != nil { + return err + } + count(func() []string { + out := make([]string, len(grps)) + for i, g := range grps { + out[i] = g.DomainID + } + return out + }, "groups") + rules, err := readRules(ctx, m.reDB) + if err != nil { + return err + } + count(func() []string { + out := make([]string, len(rules)) + for i, r := range rules { + out[i] = r.DomainID + } + return out + }, "rules") + reports, err := readReports(ctx, m.reportsDB) + if err != nil { + return err + } + count(func() []string { + out := make([]string, len(reports)) + for i, r := range reports { + out[i] = r.DomainID + } + return out + }, "reports") + alarms, err := readAlarms(ctx, m.alarmsDB) + if err != nil { + return err + } + count(func() []string { + out := make([]string, len(alarms)) + for i, a := range alarms { + out[i] = a.DomainID + } + return out + }, "alarms") + return nil +} + +// preflightGate aborts an --apply run when blocking issues exist. +func (m *migrator) preflightGate(rep *report) error { + if m.apply && rep.HasBlocking() { + return fmt.Errorf("preflight found %d blocking issue(s); aborting apply (see report)", len(rep.Blocking)) + } + return nil +} diff --git a/tools/atom-migration/report.go b/tools/atom-migration/report.go new file mode 100644 index 000000000..81b85a5a3 --- /dev/null +++ b/tools/atom-migration/report.go @@ -0,0 +1,174 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package main + +import ( + "encoding/json" + "fmt" + "os" + "path/filepath" + "sort" + "strings" + "sync" + "time" +) + +// report accumulates per-phase counts, skips, and blocking issues. +type report struct { + mu sync.Mutex + + Mode string `json:"mode"` + StartedAt time.Time `json:"started_at"` + Counts map[string]int `json:"counts"` // phase -> rows written/planned + Skipped map[string]int `json:"skipped"` // reason -> count + Warnings []string `json:"warnings"` // non-blocking + Blocking []string `json:"blocking"` // must fix before --apply + Errors []string `json:"errors"` + + // Follow-up lists surfaced for operators (force-reset users, re-issue PATs, etc.) + Todo map[string][]string `json:"todo"` +} + +func newReport(mode string) *report { + return &report{ + Mode: mode, + StartedAt: time.Now().UTC(), + Counts: map[string]int{}, + Skipped: map[string]int{}, + Todo: map[string][]string{}, + } +} + +func (r *report) count(phase string, n int) { + r.mu.Lock() + r.Counts[phase] += n + r.mu.Unlock() +} + +func (r *report) skip(reason string) { + r.mu.Lock() + r.Skipped[reason]++ + r.mu.Unlock() +} + +func (r *report) warnf(format string, a ...any) { + r.mu.Lock() + r.Warnings = append(r.Warnings, fmt.Sprintf(format, a...)) + r.mu.Unlock() +} + +func (r *report) blockf(format string, a ...any) { + r.mu.Lock() + r.Blocking = append(r.Blocking, fmt.Sprintf(format, a...)) + r.mu.Unlock() +} + +func (r *report) Errorf(format string, a ...any) { + r.mu.Lock() + r.Errors = append(r.Errors, fmt.Sprintf(format, a...)) + r.mu.Unlock() +} + +func (r *report) todo(bucket, item string) { + r.mu.Lock() + r.Todo[bucket] = append(r.Todo[bucket], item) + r.mu.Unlock() +} + +func (r *report) HasBlocking() bool { + r.mu.Lock() + defer r.mu.Unlock() + return len(r.Blocking) > 0 || len(r.Errors) > 0 +} + +func (r *report) Write(dir string) error { + if err := os.MkdirAll(dir, 0o755); err != nil { + return err + } + stamp := r.StartedAt.Format("20060102-150405") + b, err := json.MarshalIndent(r, "", " ") + if err != nil { + return err + } + if err := os.WriteFile(filepath.Join(dir, "report-"+stamp+".json"), b, 0o644); err != nil { + return err + } + return os.WriteFile(filepath.Join(dir, "report-"+stamp+".md"), []byte(r.markdown()), 0o644) +} + +func (r *report) Summary() string { + r.mu.Lock() + defer r.mu.Unlock() + var b strings.Builder + fmt.Fprintf(&b, "\n=== atom-migration %s ===\n", r.Mode) + for _, k := range sortedKeys(r.Counts) { + fmt.Fprintf(&b, " %-28s %d\n", k, r.Counts[k]) + } + if len(r.Skipped) > 0 { + b.WriteString(" skipped:\n") + for _, k := range sortedKeys(r.Skipped) { + fmt.Fprintf(&b, " %-26s %d\n", k, r.Skipped[k]) + } + } + fmt.Fprintf(&b, " warnings=%d blocking=%d errors=%d\n", + len(r.Warnings), len(r.Blocking), len(r.Errors)) + for _, x := range r.Blocking { + fmt.Fprintf(&b, " BLOCK: %s\n", x) + } + for _, x := range r.Errors { + fmt.Fprintf(&b, " ERROR: %s\n", x) + } + return b.String() +} + +func (r *report) markdown() string { + r.mu.Lock() + defer r.mu.Unlock() + var b strings.Builder + fmt.Fprintf(&b, "# atom-migration report (%s)\n\n%s\n\n", r.Mode, r.StartedAt.Format(time.RFC3339)) + b.WriteString("## Counts\n\n| phase | rows |\n|---|---|\n") + for _, k := range sortedKeys(r.Counts) { + fmt.Fprintf(&b, "| %s | %d |\n", k, r.Counts[k]) + } + if len(r.Skipped) > 0 { + b.WriteString("\n## Skipped\n\n| reason | count |\n|---|---|\n") + for _, k := range sortedKeys(r.Skipped) { + fmt.Fprintf(&b, "| %s | %d |\n", k, r.Skipped[k]) + } + } + section := func(title string, items []string) { + if len(items) == 0 { + return + } + fmt.Fprintf(&b, "\n## %s\n\n", title) + for _, x := range items { + fmt.Fprintf(&b, "- %s\n", x) + } + } + section("Blocking", r.Blocking) + section("Errors", r.Errors) + section("Warnings", r.Warnings) + for _, bucket := range sortedKeys2(r.Todo) { + section("TODO: "+bucket, r.Todo[bucket]) + } + return b.String() +} + +func sortedKeys(m map[string]int) []string { + ks := make([]string, 0, len(m)) + for k := range m { + ks = append(ks, k) + } + sort.Strings(ks) + return ks +} + +func sortedKeys2(m map[string][]string) []string { + ks := make([]string, 0, len(m)) + for k := range m { + ks = append(ks, k) + } + sort.Strings(ks) + return ks +} diff --git a/tools/atom-migration/source.go b/tools/atom-migration/source.go new file mode 100644 index 000000000..7bd960d6c --- /dev/null +++ b/tools/atom-migration/source.go @@ -0,0 +1,323 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package main + +import ( + "context" + "database/sql" + "encoding/json" + + "github.com/jmoiron/sqlx" + "github.com/lib/pq" +) + +// --- source row structs (Magistrala) --- + +type srcDomain struct { + ID string `db:"id"` + Name sql.NullString `db:"name"` + Route sql.NullString `db:"route"` + Tags pq.StringArray `db:"tags"` + Metadata []byte `db:"metadata"` + CreatedAt sql.NullTime `db:"created_at"` + UpdatedAt sql.NullTime `db:"updated_at"` + CreatedBy sql.NullString `db:"created_by"` + UpdatedBy sql.NullString `db:"updated_by"` + Status int16 `db:"status"` +} + +type srcUser struct { + ID string `db:"id"` + FirstName sql.NullString `db:"first_name"` + LastName sql.NullString `db:"last_name"` + Username sql.NullString `db:"username"` + Email sql.NullString `db:"email"` + Metadata []byte `db:"metadata"` + ProfilePicture sql.NullString `db:"profile_picture"` + AuthProvider sql.NullString `db:"auth_provider"` + Status int16 `db:"status"` + Role sql.NullInt16 `db:"role"` + VerifiedAt sql.NullTime `db:"verified_at"` + CreatedAt sql.NullTime `db:"created_at"` + UpdatedAt sql.NullTime `db:"updated_at"` +} + +type srcClient struct { + ID string `db:"id"` + Name sql.NullString `db:"name"` + DomainID string `db:"domain_id"` + ParentGroupID sql.NullString `db:"parent_group_id"` + Identity sql.NullString `db:"identity"` + Secret sql.NullString `db:"secret"` + Tags pq.StringArray `db:"tags"` + Metadata []byte `db:"metadata"` + PrivateMeta []byte `db:"private_metadata"` + Status int16 `db:"status"` + CreatedAt sql.NullTime `db:"created_at"` + UpdatedAt sql.NullTime `db:"updated_at"` +} + +type srcChannel struct { + ID string `db:"id"` + Name sql.NullString `db:"name"` + DomainID string `db:"domain_id"` + ParentGroupID sql.NullString `db:"parent_group_id"` + Route sql.NullString `db:"route"` + Tags pq.StringArray `db:"tags"` + Metadata []byte `db:"metadata"` + CreatedBy sql.NullString `db:"created_by"` + Status int16 `db:"status"` + CreatedAt sql.NullTime `db:"created_at"` + UpdatedAt sql.NullTime `db:"updated_at"` +} + +type srcConnection struct { + ChannelID string `db:"channel_id"` + DomainID string `db:"domain_id"` + ClientID string `db:"client_id"` + Type int16 `db:"type"` +} + +type srcGroup struct { + ID string `db:"id"` + ParentID sql.NullString `db:"parent_id"` + DomainID string `db:"domain_id"` + Name string `db:"name"` + Description sql.NullString `db:"description"` + Metadata []byte `db:"metadata"` + Tags pq.StringArray `db:"tags"` + Status int16 `db:"status"` + CreatedAt sql.NullTime `db:"created_at"` + UpdatedAt sql.NullTime `db:"updated_at"` +} + +// srcRole / action / member are generic across the *_roles families. +type srcRole struct { + ID string `db:"id"` + Name string `db:"name"` + EntityID string `db:"entity_id"` + CreatedAt sql.NullTime `db:"created_at"` + UpdatedAt sql.NullTime `db:"updated_at"` +} + +type srcRoleAction struct { + RoleID string `db:"role_id"` + Action string `db:"action"` +} + +type srcRoleMember struct { + RoleID string `db:"role_id"` + MemberID string `db:"member_id"` +} + +type srcPAT struct { + ID string `db:"id"` + Name string `db:"name"` + UserID sql.NullString `db:"user_id"` + Desc sql.NullString `db:"description"` + ExpiresAt sql.NullTime `db:"expires_at"` + Revoked sql.NullBool `db:"revoked"` + IssuedAt sql.NullTime `db:"issued_at"` +} + +// srcRule is a rules-engine rule (rules_engine.rules) -> Atom resource kind=rule. +type srcRule struct { + ID string `db:"id"` + Name sql.NullString `db:"name"` + DomainID string `db:"domain_id"` + Metadata []byte `db:"metadata"` + CreatedBy sql.NullString `db:"created_by"` + CreatedAt sql.NullTime `db:"created_at"` + UpdatedAt sql.NullTime `db:"updated_at"` + UpdatedBy sql.NullString `db:"updated_by"` + InputChannel sql.NullString `db:"input_channel"` + InputTopic sql.NullString `db:"input_topic"` + Outputs json.RawMessage `db:"outputs"` + Status int16 `db:"status"` + LogicType int16 `db:"logic_type"` + LogicValue []byte `db:"logic_value"` + Time sql.NullTime `db:"time"` + Recurring sql.NullInt16 `db:"recurring"` + RecurringPeriod sql.NullInt16 `db:"recurring_period"` + StartDatetime sql.NullTime `db:"start_datetime"` + Tags pq.StringArray `db:"tags"` +} + +// srcReport is a report config (reports.report_config) -> Atom resource kind=report. +type srcReport struct { + ID string `db:"id"` + Name sql.NullString `db:"name"` + Description sql.NullString `db:"description"` + DomainID string `db:"domain_id"` + Status int16 `db:"status"` + CreatedAt sql.NullTime `db:"created_at"` + CreatedBy sql.NullString `db:"created_by"` + UpdatedAt sql.NullTime `db:"updated_at"` + UpdatedBy sql.NullString `db:"updated_by"` + Due sql.NullTime `db:"due"` + Recurring sql.NullInt16 `db:"recurring"` + RecurringPeriod sql.NullInt16 `db:"recurring_period"` + StartDatetime sql.NullTime `db:"start_datetime"` + Config json.RawMessage `db:"config"` + Email json.RawMessage `db:"email"` + Metrics json.RawMessage `db:"metrics"` + ReportTemplate sql.NullString `db:"report_template"` +} + +// srcAlarm is an alarm (alarms.alarms) -> Atom resource kind=alarm. +type srcAlarm struct { + ID string `db:"id"` + RuleID string `db:"rule_id"` + DomainID string `db:"domain_id"` + ChannelID string `db:"channel_id"` + Subtopic string `db:"subtopic"` + ClientID string `db:"client_id"` + Measurement string `db:"measurement"` + Value string `db:"value"` + Unit string `db:"unit"` + Threshold string `db:"threshold"` + Cause string `db:"cause"` + Status int16 `db:"status"` + Severity int16 `db:"severity"` + AssigneeID sql.NullString `db:"assignee_id"` + CreatedAt sql.NullTime `db:"created_at"` + UpdatedAt sql.NullTime `db:"updated_at"` + UpdatedBy sql.NullString `db:"updated_by"` + AssignedAt sql.NullTime `db:"assigned_at"` + AssignedBy sql.NullString `db:"assigned_by"` + AcknowledgedAt sql.NullTime `db:"acknowledged_at"` + AcknowledgedBy sql.NullString `db:"acknowledged_by"` + ResolvedAt sql.NullTime `db:"resolved_at"` + ResolvedBy sql.NullString `db:"resolved_by"` + Metadata []byte `db:"metadata"` +} + +// --- readers --- + +func readDomains(ctx context.Context, db *sqlx.DB) ([]srcDomain, error) { + var out []srcDomain + q := `SELECT id, name, route, tags, metadata, created_at, updated_at, created_by, updated_by, status FROM domains` + return out, db.SelectContext(ctx, &out, q) +} + +func readUsers(ctx context.Context, db *sqlx.DB) ([]srcUser, error) { + var out []srcUser + q := `SELECT id, first_name, last_name, username, email, metadata, profile_picture, + auth_provider, status, role, verified_at, created_at, updated_at + FROM users` + return out, db.SelectContext(ctx, &out, q) +} + +func readClients(ctx context.Context, db *sqlx.DB) ([]srcClient, error) { + var out []srcClient + q := `SELECT id, name, domain_id, parent_group_id, identity, secret, tags, metadata, private_metadata, + status, created_at, updated_at + FROM clients` + return out, db.SelectContext(ctx, &out, q) +} + +func readConnections(ctx context.Context, db *sqlx.DB) ([]srcConnection, error) { + var out []srcConnection + return out, db.SelectContext(ctx, &out, `SELECT channel_id, domain_id, client_id, type FROM connections`) +} + +func readChannels(ctx context.Context, db *sqlx.DB) ([]srcChannel, error) { + var out []srcChannel + q := `SELECT id, name, domain_id, parent_group_id, route, tags, metadata, created_by, + status, created_at, updated_at + FROM channels` + return out, db.SelectContext(ctx, &out, q) +} + +func readGroups(ctx context.Context, db *sqlx.DB) ([]srcGroup, error) { + var out []srcGroup + q := `SELECT id, parent_id, domain_id, name, description, metadata, tags, status, + created_at, updated_at + FROM groups` + return out, db.SelectContext(ctx, &out, q) +} + +func readRules(ctx context.Context, db *sqlx.DB) ([]srcRule, error) { + var out []srcRule + q := `SELECT id, name, domain_id, metadata, created_by, created_at, updated_at, updated_by, + input_channel, input_topic, outputs, status, logic_type, logic_value, + "time", recurring, recurring_period, start_datetime, tags + FROM rules` + return out, db.SelectContext(ctx, &out, q) +} + +func readReports(ctx context.Context, db *sqlx.DB) ([]srcReport, error) { + var out []srcReport + q := `SELECT id, name, description, domain_id, status, created_at, created_by, updated_at, + updated_by, due, recurring, recurring_period, start_datetime, + config, email, metrics, report_template + FROM report_config` + return out, db.SelectContext(ctx, &out, q) +} + +func readAlarms(ctx context.Context, db *sqlx.DB) ([]srcAlarm, error) { + var out []srcAlarm + q := `SELECT id, rule_id, domain_id, channel_id, subtopic, client_id, measurement, value, + unit, threshold, cause, status, severity, assignee_id, created_at, updated_at, + updated_by, assigned_at, assigned_by, acknowledged_at, acknowledged_by, + resolved_at, resolved_by, metadata + FROM alarms` + return out, db.SelectContext(ctx, &out, q) +} + +// readRoleFamily reads _roles, _role_actions, _role_members for one service. +func readRoleFamily(ctx context.Context, db *sqlx.DB, prefix string) ([]srcRole, []srcRoleAction, []srcRoleMember, error) { + var roles []srcRole + if err := db.SelectContext(ctx, &roles, + `SELECT id, name, entity_id, created_at, updated_at FROM `+prefix+`_roles`); err != nil { + return nil, nil, nil, err + } + var acts []srcRoleAction + if err := db.SelectContext(ctx, &acts, + `SELECT role_id, action FROM `+prefix+`_role_actions`); err != nil { + return nil, nil, nil, err + } + var mems []srcRoleMember + if err := db.SelectContext(ctx, &mems, + `SELECT role_id, member_id FROM `+prefix+`_role_members`); err != nil { + return nil, nil, nil, err + } + return roles, acts, mems, nil +} + +func readPATs(ctx context.Context, db *sqlx.DB) ([]srcPAT, error) { + var out []srcPAT + q := `SELECT id, name, user_id, description, expires_at, revoked, issued_at FROM pats` + return out, db.SelectContext(ctx, &out, q) +} + +type srcPATScope struct { + PatID string `db:"pat_id"` + DomainID sql.NullString `db:"domain_id"` + EntityType string `db:"entity_type"` + Operation string `db:"operation"` + EntityID string `db:"entity_id"` +} + +func readPATScopes(ctx context.Context, db *sqlx.DB) ([]srcPATScope, error) { + var out []srcPATScope + q := `SELECT pat_id, domain_id, entity_type, operation, entity_id FROM pat_scopes` + return out, db.SelectContext(ctx, &out, q) +} + +type srcInvitation struct { + InvitedBy string `db:"invited_by"` + InviteeID string `db:"invitee_user_id"` + DomainID string `db:"domain_id"` + RoleID string `db:"role_id"` + CreatedAt sql.NullTime `db:"created_at"` + ConfirmedAt sql.NullTime `db:"confirmed_at"` + RejectedAt sql.NullTime `db:"rejected_at"` +} + +func readInvitations(ctx context.Context, db *sqlx.DB) ([]srcInvitation, error) { + var out []srcInvitation + q := `SELECT invited_by, invitee_user_id, domain_id, role_id, created_at, confirmed_at, rejected_at FROM invitations` + return out, db.SelectContext(ctx, &out, q) +} diff --git a/tools/atom-migration/verify.go b/tools/atom-migration/verify.go new file mode 100644 index 000000000..fd2cc6c80 --- /dev/null +++ b/tools/atom-migration/verify.go @@ -0,0 +1,257 @@ +// Copyright (c) Abstract Machines +// SPDX-License-Identifier: Apache-2.0 + +package main + +import ( + "context" + + "github.com/jmoiron/sqlx" +) + +// Verify reconciles a completed migration: every source row that should have +// migrated must exist in Atom, and a sample of reconstructed authz edges +// (device→channel publish/subscribe) must be present. Read-only. +func (m *migrator) Verify(ctx context.Context, rep *report) error { + domSet, err := m.domainSet(ctx) + if err != nil { + return err + } + + atomTenants, err := idSet(ctx, m.atom, `SELECT id::text FROM tenants`) + if err != nil { + return err + } + atomEntities, err := idSet(ctx, m.atom, `SELECT id::text FROM entities`) + if err != nil { + return err + } + atomAuthenticatedUsers, err := idSet(ctx, m.atom, + `SELECT entity_id::text FROM principal_group_members WHERE group_id = $1`, authenticatedUsersGroupID) + if err != nil { + return err + } + atomResources, err := idSet(ctx, m.atom, `SELECT id::text FROM resources`) + if err != nil { + return err + } + atomGroups, err := idSet(ctx, m.atom, `SELECT id::text FROM object_groups`) + if err != nil { + return err + } + + // 1. tenants + doms, err := readDomains(ctx, m.domainsDB) + if err != nil { + return err + } + m.reconcile(rep, "tenants", idsOf(len(doms), func(i int) (string, bool) { return doms[i].ID, true }), atomTenants) + + // 2. human entities + users, err := readUsers(ctx, m.usersDB) + if err != nil { + return err + } + userIDs := idsOf(len(users), func(i int) (string, bool) { return users[i].ID, true }) + m.reconcile(rep, "entities.users", userIDs, atomEntities) + m.reconcile(rep, "principal_group_members.authenticated_users", userIDs, atomAuthenticatedUsers) + + // 3. device entities (only those with a valid domain were migrated) + clients, err := readClients(ctx, m.clientsDB) + if err != nil { + return err + } + m.reconcile(rep, "entities.clients", idsOf(len(clients), func(i int) (string, bool) { + return clients[i].ID, domSet[clients[i].DomainID] + }), atomEntities) + + // 4. resources + chans, err := readChannels(ctx, m.channelsDB) + if err != nil { + return err + } + m.reconcile(rep, "resources.channels", idsOf(len(chans), func(i int) (string, bool) { + return chans[i].ID, domSet[chans[i].DomainID] + }), atomResources) + + // 4b. resources: rules, reports, alarms + rules, err := readRules(ctx, m.reDB) + if err != nil { + return err + } + m.reconcile(rep, "resources.rules", idsOf(len(rules), func(i int) (string, bool) { + return rules[i].ID, domSet[rules[i].DomainID] + }), atomResources) + reports, err := readReports(ctx, m.reportsDB) + if err != nil { + return err + } + m.reconcile(rep, "resources.reports", idsOf(len(reports), func(i int) (string, bool) { + return reports[i].ID, domSet[reports[i].DomainID] + }), atomResources) + alarms, err := readAlarms(ctx, m.alarmsDB) + if err != nil { + return err + } + m.reconcile(rep, "resources.alarms", idsOf(len(alarms), func(i int) (string, bool) { + return alarms[i].ID, domSet[alarms[i].DomainID] + }), atomResources) + + // 5. object_groups + grps, err := readGroups(ctx, m.groupsDB) + if err != nil { + return err + } + m.reconcile(rep, "object_groups", idsOf(len(grps), func(i int) (string, bool) { + return grps[i].ID, domSet[grps[i].DomainID] + }), atomGroups) + + // 6. authz spot-check: every connection must have a device->channel policy. + if err := m.verifyConnections(ctx, rep, domSet); err != nil { + return err + } + return nil +} + +func (m *migrator) domainSet(ctx context.Context) (map[string]bool, error) { + doms, err := readDomains(ctx, m.domainsDB) + if err != nil { + return nil, err + } + s := map[string]bool{} + for _, d := range doms { + s[d.ID] = true + } + return s, nil +} + +// reconcile counts how many expected ids are missing from the atom set. +func (m *migrator) reconcile(rep *report, label string, expected []string, atom map[string]bool) { + missing := 0 + for _, id := range expected { + if !atom[id] { + missing++ + } + } + rep.count("verify."+label+".expected", len(expected)) + if missing > 0 { + rep.blockf("verify %s: %d of %d expected rows missing from Atom", label, missing, len(expected)) + } else { + rep.count("verify."+label+".ok", len(expected)) + } +} + +func (m *migrator) verifyConnections(ctx context.Context, rep *report, domSet map[string]bool) error { + cli, err := readConnections(ctx, m.clientsDB) + if err != nil { + return err + } + ch, err := readConnections(ctx, m.channelsDB) + if err != nil { + return err + } + clients, err := readClients(ctx, m.clientsDB) + if err != nil { + return err + } + clientDomain := map[string]string{} + for _, c := range clients { + if domSet[c.DomainID] { + clientDomain[c.ID] = c.DomainID + } + } + channels, err := readChannels(ctx, m.channelsDB) + if err != nil { + return err + } + channelDomain := map[string]string{} + for _, c := range channels { + if domSet[c.DomainID] { + channelDomain[c.ID] = c.DomainID + } + } + // Build atom edge set: subject_id | channel(object_id) | action. + edges := map[string]bool{} + rows, err := m.atom.QueryxContext(ctx, ` + SELECT dp.subject_id::text, pb.object_id::text, a.name + FROM direct_policies dp + JOIN permission_blocks pb ON pb.id = dp.permission_block_id + JOIN permission_block_actions pba ON pba.permission_block_id = pb.id + JOIN actions a ON a.id = pba.action_id + WHERE pb.scope_mode = 'object' AND pb.object_kind = 'resource'`) + if err != nil { + return err + } + for rows.Next() { + var s, o, act string + if err := rows.Scan(&s, &o, &act); err != nil { + rows.Close() + return err + } + edges[s+"|"+o+"|"+act] = true + } + rows.Close() + + seen := map[string]bool{} + expected, missing := 0, 0 + for _, c := range append(cli, ch...) { + act, ok := connectionAction(c.Type) + if !ok { + continue + } + chDom, ok := channelDomain[c.ChannelID] + if !ok { + continue + } + clDom, ok := clientDomain[c.ClientID] + if !ok || clDom != chDom || c.DomainID != chDom { + continue + } + k := c.ClientID + "|" + c.ChannelID + "|" + act + if seen[k] { + continue + } + seen[k] = true + expected++ + if !edges[k] { + missing++ + } + } + rep.count("verify.connections.expected", expected) + if missing > 0 { + rep.blockf("verify connections: %d of %d device->channel edges missing", missing, expected) + } else { + rep.count("verify.connections.ok", expected) + } + return nil +} + +// --- helpers --- + +func idSet(ctx context.Context, db *sqlx.DB, query string, args ...any) (map[string]bool, error) { + rows, err := db.QueryxContext(ctx, query, args...) + if err != nil { + return nil, err + } + defer rows.Close() + s := map[string]bool{} + for rows.Next() { + var id string + if err := rows.Scan(&id); err != nil { + return nil, err + } + s[id] = true + } + return s, rows.Err() +} + +// idsOf collects ids for indices 0..n-1 where the picker's second return is true. +func idsOf(n int, pick func(i int) (string, bool)) []string { + out := make([]string, 0, n) + for i := range n { + if id, ok := pick(i); ok { + out = append(out, id) + } + } + return out +} diff --git a/tools/config/.mockery.yaml b/tools/config/.mockery.yaml index 4339ac18e..44d7af2b3 100644 --- a/tools/config/.mockery.yaml +++ b/tools/config/.mockery.yaml @@ -12,20 +12,6 @@ force-file-write: true include-auto-generated: true packages: - github.com/absmach/magistrala/api/grpc/clients/v1: - interfaces: - ClientsServiceClient: - config: - dir: "./clients/mocks" - structname: "ClientsServiceClient" - filename: "clients_client.go" - github.com/absmach/magistrala/api/grpc/domains/v1: - interfaces: - DomainsServiceClient: - config: - dir: "./domains/mocks" - structname: "DomainsServiceClient" - filename: "domains_client.go" github.com/absmach/magistrala/api/grpc/token/v1: interfaces: TokenServiceClient: @@ -33,20 +19,6 @@ packages: dir: "./auth/mocks" structname: "TokenServiceClient" filename: "token_client.go" - github.com/absmach/magistrala/api/grpc/channels/v1: - interfaces: - ChannelsServiceClient: - config: - dir: "./channels/mocks" - structname: "ChannelsServiceClient" - filename: "channels_client.go" - github.com/absmach/magistrala/api/grpc/groups/v1: - interfaces: - GroupsServiceClient: - config: - dir: "./groups/mocks" - structname: "GroupsServiceClient" - filename: "groups_client.go" github.com/absmach/magistrala/api/grpc/certs/v1: interfaces: CertsServiceClient: @@ -79,22 +51,6 @@ packages: PATS: PATSRepository: Service: - github.com/absmach/magistrala/channels: - interfaces: - Cache: - Repository: - Service: - github.com/absmach/magistrala/channels/private: - interfaces: - Service: - github.com/absmach/magistrala/clients: - interfaces: - Repository: - Cache: - Service: - github.com/absmach/magistrala/clients/private: - interfaces: - Service: github.com/absmach/magistrala/certs: interfaces: Agent: @@ -107,21 +63,6 @@ packages: interfaces: Service: SubscriptionsRepository: - github.com/absmach/magistrala/domains: - interfaces: - Repository: - Cache: - Service: - github.com/absmach/magistrala/domains/private: - interfaces: - Service: - github.com/absmach/magistrala/groups: - interfaces: - Repository: - Service: - github.com/absmach/magistrala/groups/private: - interfaces: - Service: github.com/absmach/magistrala/journal: interfaces: Repository: @@ -143,18 +84,10 @@ packages: github.com/absmach/magistrala/pkg/messaging: interfaces: PubSub: - github.com/absmach/magistrala/pkg/oauth2: - interfaces: - Provider: github.com/absmach/magistrala/pkg/policies: interfaces: Evaluator: Service: - github.com/absmach/magistrala/pkg/roles: - interfaces: - Provisioner: - RoleManager: - Repository: github.com/absmach/magistrala/pkg/callout: interfaces: Callout: @@ -168,18 +101,6 @@ packages: interfaces: Repository: Service: - github.com/absmach/magistrala/bootstrap: - interfaces: - ConfigRepository: - ConfigReader: - Service: - ProfileRepository: - BindingStore: - BindingResolver: - Renderer: - github.com/absmach/magistrala/provision: - interfaces: - Service: github.com/absmach/magistrala/alarms: interfaces: Service: @@ -188,12 +109,6 @@ packages: interfaces: Service: Repository: - github.com/absmach/magistrala/users: - interfaces: - Emailer: - Hasher: - Repository: - Service: github.com/absmach/magistrala/notifications: interfaces: Notifier: diff --git a/users/README.md b/users/README.md deleted file mode 100644 index 70b0eb87b..000000000 --- a/users/README.md +++ /dev/null @@ -1,329 +0,0 @@ -# Users - -Users service provides an HTTP API for managing users. Through this API clients are able to do the following actions: - -- register new accounts -- login -- manage account(s) (list, update, delete) - -For in-depth explanation of the aforementioned scenarios, as well as thorough understanding of Magistrala, please check out the [official documentation][doc]. - -## Configuration - -The service is configured using the environment variables presented in the following table. Note that any unset variables will be replaced with their default values. - -| Variable | Description | Default | -| ---------------------------------- | ----------------------------------------------------------------------- | --------------------------------- | -| `MG_USERS_LOG_LEVEL` | Log level for users service (debug, info, warn, error) | info | -| `MG_USERS_ADMIN_EMAIL` | Default user, created on startup | | -| `MG_USERS_ADMIN_PASSWORD` | Default user password, created on startup | 12345678 | -| `MG_USERS_PASS_REGEX` | Password regex | ^.{8,}$ | -| `MG_USERS_HTTP_HOST` | Users service HTTP host | localhost | -| `MG_USERS_HTTP_PORT` | Users service HTTP port | 9002 | -| `MG_USERS_HTTP_SERVER_CERT` | Path to the PEM encoded server certificate file | "" | -| `MG_USERS_HTTP_SERVER_KEY` | Path to the PEM encoded server key file | "" | -| `MG_USERS_HTTP_SERVER_CA_CERTS` | Path to the PEM encoded server CA certificate file | "" | -| `MG_USERS_HTTP_CLIENT_CA_CERTS` | Path to the PEM encoded client CA certificate file | "" | -| `MG_AUTH_GRPC_URL` | Auth service GRPC URL | localhost:8181 | -| `MG_AUTH_GRPC_TIMEOUT` | Auth service GRPC timeout | 1s | -| `MG_AUTH_GRPC_CLIENT_CERT` | Path to the PEM encoded client certificate file | "" | -| `MG_AUTH_GRPC_CLIENT_KEY` | Path to the PEM encoded client key file | "" | -| `MG_AUTH_GRPC_SERVER_CA_CERTS` | Path to the PEM encoded server CA certificate file | "" | -| `MG_USERS_DB_HOST` | Database host address | localhost | -| `MG_USERS_DB_PORT` | Database host port | 5432 | -| `MG_USERS_DB_USER` | Database user | magistrala | -| `MG_USERS_DB_PASS` | Database password | magistrala | -| `MG_USERS_DB_NAME` | Name of the database used by the service | users | -| `MG_USERS_DB_SSL_MODE` | Database connection SSL mode (disable, require, verify-ca, verify-full) | disable | -| `MG_USERS_DB_SSL_CERT` | Path to the PEM encoded certificate file | "" | -| `MG_USERS_DB_SSL_KEY` | Path to the PEM encoded key file | "" | -| `MG_USERS_DB_SSL_ROOT_CERT` | Path to the PEM encoded root certificate file | "" | -| `MG_EMAIL_HOST` | Mail server host | localhost | -| `MG_EMAIL_PORT` | Mail server port | 25 | -| `MG_EMAIL_USERNAME` | Mail server username | "" | -| `MG_EMAIL_PASSWORD` | Mail server password | "" | -| `MG_EMAIL_FROM_ADDRESS` | Email "from" address | "" | -| `MG_EMAIL_FROM_NAME` | Email "from" name | "" | -| `MG_PASSWORD_RESET_URL_PREFIX` | Password reset URL prefix | | -| `MG_PASSWORD_RESET_EMAIL_TEMPLATE` | Password reset email template | reset-password-email.tmpl | -| `MG_VERIFICATION_URL_PREFIX` | Verification URL prefix | | -| `MG_VERIFICATION_EMAIL_TEMPLATE` | Verification email template | verification-email.tmpl | -| `MG_USERS_ES_URL` | Event store URL | | -| `MG_JAEGER_URL` | Jaeger server URL | | -| `MG_OAUTH_UI_REDIRECT_URL` | OAuth UI redirect URL | | -| `MG_OAUTH_UI_ERROR_URL` | OAuth UI error URL | | -| `MG_USERS_DELETE_INTERVAL` | Interval for deleting users | 24h | -| `MG_USERS_DELETE_AFTER` | Time after which users are deleted | 720h | -| `MG_JAEGER_TRACE_RATIO` | Jaeger sampling ratio | 1.0 | -| `MG_SEND_TELEMETRY` | Send telemetry to magistrala call home server. | true | -| `MG_USERS_INSTANCE_ID` | Magistrala instance ID | "" | - -## Deployment - -The service itself is distributed as Docker container. Check the [`users`](https://github.com/absmach/magistrala/blob/main/docker/docker-compose.yaml) service section in docker-compose file to see how service is deployed. - -To start the service outside of the container, execute the following shell script: - -```bash -# download the latest version of the service -git clone https://github.com/absmach/magistrala - -cd magistrala - -# compile the service -make users - -# copy binary to bin -make install - -# set the environment variables and run the service -MG_USERS_LOG_LEVEL=info \ -MG_USERS_ADMIN_EMAIL=admin@example.com \ -MG_USERS_ADMIN_PASSWORD=12345678 \ -MG_USERS_PASS_REGEX="^.{8,}$" \ -MG_USERS_HTTP_HOST=localhost \ -MG_USERS_HTTP_PORT=9002 \ -MG_USERS_HTTP_SERVER_CERT="" \ -MG_USERS_HTTP_SERVER_KEY="" \ -MG_USERS_HTTP_SERVER_CA_CERTS="" \ -MG_USERS_HTTP_CLIENT_CA_CERTS="" \ -MG_AUTH_GRPC_URL=localhost:8181 \ -MG_AUTH_GRPC_TIMEOUT=1s \ -MG_AUTH_GRPC_CLIENT_CERT="" \ -MG_AUTH_GRPC_CLIENT_KEY="" \ -MG_AUTH_GRPC_SERVER_CA_CERTS="" \ -MG_USERS_DB_HOST=localhost \ -MG_USERS_DB_PORT=5432 \ -MG_USERS_DB_USER=magistrala \MG_USERS_DB_PASS=magistrala \MG_USERS_DB_NAME=users \ -MG_USERS_DB_SSL_MODE=disable \ -MG_USERS_DB_SSL_CERT="" \ -MG_USERS_DB_SSL_KEY="" \ -MG_USERS_DB_SSL_ROOT_CERT="" \ -MG_EMAIL_HOST=smtp.mailtrap.io \ -MG_EMAIL_PORT=2525 \ -MG_EMAIL_USERNAME="18bf7f7070513" \ -MG_EMAIL_PASSWORD="2b0d302e775b1e" \ -MG_EMAIL_FROM_ADDRESS=from@example.com \ -MG_EMAIL_FROM_NAME=Example \ -MG_PASSWORD_RESET_URL_PREFIX=http://localhost:9002/password/reset \ -MG_PASSWORD_RESET_EMAIL_TEMPLATE=docker/templates/reset-password-email.tmpl \ -MG_VERIFICATION_URL_PREFIX=http://localhost:9002/users/verify-email \ -MG_VERIFICATION_EMAIL_TEMPLATE=docker/templates/verification-email.tmpl \ -MG_USERS_ES_URL=nats://localhost:4222 \ -MG_JAEGER_URL=http://localhost:14268/api/traces \ -MG_JAEGER_TRACE_RATIO=1.0 \ -MG_SEND_TELEMETRY=true \ -MG_OAUTH_UI_REDIRECT_URL=http://localhost:9095/domains \ -MG_OAUTH_UI_ERROR_URL=http://localhost:9095/error \ -MG_USERS_DELETE_INTERVAL=24h \ -MG_USERS_DELETE_AFTER=720h \ -MG_USERS_INSTANCE_ID="" \ -$GOBIN/magistrala-users -``` - -If `MG_EMAIL_TEMPLATE` doesn't point to any file service will function but password reset functionality will not work. The email environment variables are used to send emails with password reset link. The service expects a file in Go template format. The template should be something like [this](https://github.com/absmach/magistrala/blob/main/docker/templates/users.tmpl). - -Setting `MG_USERS_HTTP_SERVER_CERT` and `MG_USERS_HTTP_SERVER_KEY` will enable TLS against the service. The service expects a file in PEM format for both the certificate and the key. Setting `MG_USERS_HTTP_SERVER_CA_CERTS` will enable TLS against the service trusting only those CAs that are provided. The service expects a file in PEM format of trusted CAs. Setting `MG_USERS_HTTP_CLIENT_CA_CERTS` will enable TLS against the service trusting only those CAs that are provided. The service expects a file in PEM format of trusted CAs. - -Setting `MG_AUTH_GRPC_CLIENT_CERT` and `MG_AUTH_GRPC_CLIENT_KEY` will enable TLS against the auth service. The service expects a file in PEM format for both the certificate and the key. Setting `MG_AUTH_GRPC_SERVER_CA_CERTS` will enable TLS against the auth service trusting only those CAs that are provided. The service expects a file in PEM format of trusted CAs. - -## HTTP API - -Base URL defaults to `http://localhost:9002`. Unless otherwise noted, endpoints require `Authorization: Bearer `. - -### Usage - -| Operation | Description | -| ----------------- | ------------------------------------------------------------------------------------------------------------ | -| Register | Create a user; optionally protected if self-registration is disabled. | -| Issue token | Exchange identity (email/username) and secret for access/refresh tokens. | -| Refresh token | Exchange a refresh token for a new access token. | -| Profile | Fetch the authenticated user profile. | -| List/search users | Page and filter users. | -| View user | Retrieve a user by ID . | -| Update user | Patch names/metadata/tags/profile picture; update email/username/role/tags/password via dedicated endpoints. | -| Status | Enable/disable a user or delete a user. | -| Verification | Send verification email; verify via emailed link. | -| Password reset | Request a reset link and set a new password. | - -### Best practices - -- Disable self-registration in production; onboard users via admin tokens or your IdP. -- Keep `allow_unverified_user` false and require email verification before granting domain roles. -- Enforce TLS for HTTP and mTLS for gRPC by setting server/client cert env vars. -- Harden passwords with `MG_USERS_PASS_REGEX` and rotate credentials; purge stale accounts via `MG_USERS_DELETE_AFTER`. -- Rate-limit token issuance and password reset endpoints at your API gateway; export Prometheus metrics to watch for abuse. -- Store SMTP credentials and certificates in a secrets manager; avoid embedding secrets in images or repos. - -### API examples - -#### Register a user - -```bash -curl -X POST "http://localhost:9002/users" \ - -H "Content-Type: application/json" \ - -d '{ - "first_name": "Ada", - "last_name": "Lovelace", - "credentials": { "username": "ada", "secret": "changeMe123" }, - "email": "ada@example.com", - "role": 0, - "status": 0, - "tags": ["iot", "beta"], - "metadata": { "team": "core" } - }' -``` - -Expected response (201 Created): - -```json -{ - "id": "c0b0c68c-5b93-4a93-8f1a-5d63a3f5c3c7", - "first_name": "Ada", - "last_name": "Lovelace", - "email": "ada@example.com", - "role": 0, - "status": 0, - "tags": ["iot", "beta"], - "metadata": { "team": "core" }, - "created_at": "2024-10-24T13:31:52Z" -} -``` - -#### Issue access/refresh tokens (login) - -```bash -curl -X POST "http://localhost:9002/users/tokens/issue" \ - -H "Content-Type: application/json" \ - -d '{ "identity": "ada@example.com", "secret": "changeMe123" }' -``` - -Expected response (201 Created): - -```json -{ - "access_token": "eyJhbGciOi...", - "refresh_token": "eyJhbGciOi...", - "access_type": "Bearer" -} -``` - -#### View authenticated profile - -```bash -curl -X GET "http://localhost:9002/users/profile" \ - -H "Authorization: Bearer $ACCESS_TOKEN" -``` - -Expected response: - -```json -{ - "id": "c0b0c68c-5b93-4a93-8f1a-5d63a3f5c3c7", - "first_name": "Ada", - "last_name": "Lovelace", - "email": "ada@example.com", - "role": 0, - "status": 0, - "tags": ["iot", "beta"], - "metadata": { "team": "core" }, - "verified_at": "2024-10-24T14:02:00Z", - "created_at": "2024-10-24T13:31:52Z", - "updated_at": "2024-10-24T14:02:00Z" -} -``` - -#### List users - -```bash -curl -X GET "http://localhost:9002/users?limit=5&status=enabled&dir=desc" \ - -H "Authorization: Bearer $ACCESS_TOKEN" -``` - -Expected response: - -```json -{ - "total": 2, - "offset": 0, - "limit": 5, - "users": [ - { - "id": "c0b0c68c-5b93-4a93-8f1a-5d63a3f5c3c7", - "first_name": "Ada", - "last_name": "Lovelace", - "email": "ada@example.com", - "role": 0, - "status": 0, - "tags": ["iot", "beta"], - "created_at": "2024-10-24T13:31:52Z" - } - ] -} -``` - -#### Update user metadata and name - -```bash -curl -X PATCH "http://localhost:9002/users/${USER_ID}" \ - -H "Authorization: Bearer $ACCESS_TOKEN" \ - -H "Content-Type: application/json" \ - -d '{ - "first_name": "Ada", - "last_name": "Byron", - "metadata": { "team": "edge" }, - "tags": ["edge", "beta"] - }' -``` - -Expected response: - -```json -{ - "id": "c0b0c68c-5b93-4a93-8f1a-5d63a3f5c3c7", - "first_name": "Ada", - "last_name": "Byron", - "email": "ada@example.com", - "tags": ["edge", "beta"], - "metadata": { "team": "edge" }, - "status": 0, - "role": 0, - "updated_at": "2024-10-24T14:45:10Z", - "updated_by": "a5b6c7d8-e901-4fab-9bcd-123456789abc" -} -``` - -#### Request password reset - -```bash -curl -X POST "http://localhost:9002/password/reset-request" \ - -H "Content-Type: application/json" \ - -d '{ "email": "ada@example.com" }' -``` - -Expected response (201 Created): - -```json -{ "msg": "Email with reset link is sent" } -``` - -#### Health check - -```bash -curl -X GET "http://localhost:9002/health" -``` - -Expected response: - -```json -{ - "status": "pass", - "version": "0.18.0", - "commit": "ffffffff", - "description": "users service", - "build_time": "1970-01-01_00:00:00", - "instance_id": "b4f1d5d2-4f24-4c2a-9a40-123456789abc" -} -``` - -[doc]: https://magistrala.absmach.eu/docs/ \ No newline at end of file diff --git a/users/api/doc.go b/users/api/doc.go deleted file mode 100644 index 2424852cc..000000000 --- a/users/api/doc.go +++ /dev/null @@ -1,6 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package api contains API-related concerns: endpoint definitions, middlewares -// and all resource representations. -package api diff --git a/users/api/endpoint_test.go b/users/api/endpoint_test.go deleted file mode 100644 index b5f1860b3..000000000 --- a/users/api/endpoint_test.go +++ /dev/null @@ -1,3173 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package api_test - -import ( - "encoding/json" - "fmt" - "io" - "net/http" - "net/http/httptest" - "net/url" - "regexp" - "strings" - "testing" - "time" - - grpcTokenV1 "github.com/absmach/magistrala/api/grpc/token/v1" - api "github.com/absmach/magistrala/api/http" - apiutil "github.com/absmach/magistrala/api/http/util" - authmocks "github.com/absmach/magistrala/auth/mocks" - "github.com/absmach/magistrala/internal/testsutil" - mglog "github.com/absmach/magistrala/logger" - smqauthn "github.com/absmach/magistrala/pkg/authn" - authnmocks "github.com/absmach/magistrala/pkg/authn/mocks" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - oauth2mocks "github.com/absmach/magistrala/pkg/oauth2/mocks" - "github.com/absmach/magistrala/pkg/uuid" - "github.com/absmach/magistrala/users" - usersapi "github.com/absmach/magistrala/users/api" - "github.com/absmach/magistrala/users/mocks" - "github.com/go-chi/chi/v5" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" -) - -var ( - secret = "strongsecret" - validCMetadata = users.Metadata{"role": "user"} - user = users.User{ - ID: testsutil.GenerateUUID(&testing.T{}), - LastName: "doe", - FirstName: "jane", - Tags: []string{"foo", "bar"}, - Email: "useremail@example.com", - Credentials: users.Credentials{Username: "username", Secret: secret}, - Metadata: validCMetadata, - PrivateMetadata: validCMetadata, - Status: users.EnabledStatus, - } - validToken = "valid" - inValidToken = "invalid" - inValid = "invalid" - validID = "d4ebb847-5d0e-4e46-bdd9-b6aceaaa3a22" - passRegex = regexp.MustCompile("^.{8,}$") - testReferer = "http://localhost" - domainID = testsutil.GenerateUUID(&testing.T{}) - verifiedSession = smqauthn.Session{UserID: validID, DomainID: domainID, Verified: true} - validTimeStamp = time.Date(2024, 1, 1, 0, 0, 0, 0, time.UTC) -) - -const contentType = "application/json" - -type testRequest struct { - user *http.Client - method string - url string - contentType string - referer string - token string - body io.Reader -} - -func (tr testRequest) make() (*http.Response, error) { - req, err := http.NewRequest(tr.method, tr.url, tr.body) - if err != nil { - return nil, err - } - - if tr.token != "" { - req.Header.Set("Authorization", apiutil.BearerPrefix+tr.token) - } - - if tr.contentType != "" { - req.Header.Set("Content-Type", tr.contentType) - } - - req.Header.Set("Referer", tr.referer) - - return tr.user.Do(req) -} - -func newUsersServer() (*httptest.Server, *mocks.Service, *authnmocks.Authentication) { - svc := new(mocks.Service) - logger := mglog.NewMock() - mux := chi.NewRouter() - idp := uuid.NewMock() - provider := new(oauth2mocks.Provider) - provider.On("Name").Return("test") - authn := new(authnmocks.Authentication) - am := smqauthn.NewAuthNMiddleware(authn) - token := new(authmocks.TokenServiceClient) - usersapi.MakeHandler(svc, am, token, true, mux, logger, "", passRegex, idp, provider) - - return httptest.NewServer(mux), svc, authn -} - -func toJSON(data any) string { - jsonData, err := json.Marshal(data) - if err != nil { - return "" - } - return string(jsonData) -} - -func TestRegister(t *testing.T) { - us, svc, _ := newUsersServer() - defer us.Close() - - cases := []struct { - desc string - user users.User - token string - contentType string - status int - err error - }{ - { - desc: "register a new user with a valid token", - user: user, - token: validToken, - contentType: contentType, - status: http.StatusCreated, - err: nil, - }, - { - desc: "register an existing user", - user: user, - token: validToken, - contentType: contentType, - status: http.StatusBadRequest, - err: svcerr.ErrConflict, - }, - { - desc: "register a user that can't be marshalled", - user: users.User{ - Email: "user@example.com", - Credentials: users.Credentials{ - Secret: "12345678", - }, - PrivateMetadata: map[string]any{ - "test": make(chan int), - }, - }, - token: validToken, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrMalformedRequestBody, - }, - { - desc: "register user with invalid status", - user: users.User{ - Email: "newclientwithinvalidstatus@example.com", - FirstName: "newclientwithinvalidstatus", - LastName: "newclientwithinvalidstatus", - Credentials: users.Credentials{ - Username: "username", - Secret: secret, - }, - Status: users.AllStatus, - }, - token: validToken, - contentType: contentType, - status: http.StatusBadRequest, - err: svcerr.ErrInvalidStatus, - }, - { - desc: "register a user with name too long", - user: users.User{ - FirstName: strings.Repeat("a", 1025), - LastName: "newuserwithnametoolong", - Email: "newuserwithinvalidname@example.com", - Credentials: users.Credentials{ - Secret: secret, - }, - }, - token: validToken, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrNameSize, - }, - { - desc: "register user with invalid content type", - user: user, - token: validToken, - contentType: "application/xml", - status: http.StatusUnsupportedMediaType, - err: apiutil.ErrUnsupportedContentType, - }, - { - desc: "register user with empty request body", - user: users.User{}, - token: validToken, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrMissingFirstName, - }, - { - desc: "register user with invalid username", - user: users.User{ - FirstName: "newuserwithinvalidusername", - LastName: "newuserwithinvalidusername", - Email: "user@example.com", - Credentials: users.Credentials{ - Username: "invalid username", - Secret: secret, - }, - Status: users.EnabledStatus, - }, - token: validToken, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrInvalidUsername, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - data := toJSON(tc.user) - req := testRequest{ - user: us.Client(), - method: http.MethodPost, - url: fmt.Sprintf("%s/users/", us.URL), - contentType: tc.contentType, - token: tc.token, - body: strings.NewReader(data), - } - - svcCall := svc.On("Register", mock.Anything, smqauthn.Session{}, tc.user, true).Return(tc.user, tc.err) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - var errRes respBody - err = json.NewDecoder(res.Body).Decode(&errRes) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if errRes.Err != "" || errRes.Message != "" { - err = errors.Wrap(errors.New(errRes.Err), errors.New(errRes.Message)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - }) - } -} - -func TestView(t *testing.T) { - us, svc, authn := newUsersServer() - defer us.Close() - - cases := []struct { - desc string - token string - id string - status int - authnRes smqauthn.Session - authnErr error - svcErr error - err error - }{ - { - desc: "view user as admin with valid token", - token: validToken, - id: user.ID, - status: http.StatusOK, - authnRes: verifiedSession, - err: nil, - }, - { - desc: "view user with invalid token", - token: inValidToken, - id: user.ID, - status: http.StatusUnauthorized, - authnRes: smqauthn.Session{}, - authnErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "view user with empty token", - token: "", - id: user.ID, - status: http.StatusUnauthorized, - authnRes: smqauthn.Session{}, - authnErr: svcerr.ErrAuthentication, - err: apiutil.ErrBearerToken, - }, - { - desc: "view user as normal user successfully", - token: validToken, - id: user.ID, - status: http.StatusOK, - authnRes: verifiedSession, - err: nil, - }, - { - desc: "view user with invalid ID", - token: validToken, - id: inValid, - status: http.StatusUnprocessableEntity, - authnRes: verifiedSession, - svcErr: svcerr.ErrViewEntity, - err: svcerr.ErrViewEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - user: us.Client(), - method: http.MethodGet, - url: fmt.Sprintf("%s/users/%s", us.URL, tc.id), - token: tc.token, - } - - authnCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.authnRes, tc.authnErr) - svcCall := svc.On("View", mock.Anything, tc.authnRes, tc.id).Return(users.User{}, tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - var errRes respBody - err = json.NewDecoder(res.Body).Decode(&errRes) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if errRes.Err != "" || errRes.Message != "" { - err = errors.Wrap(errors.New(errRes.Err), errors.New(errRes.Message)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authnCall.Unset() - }) - } -} - -func TestViewProfile(t *testing.T) { - us, svc, authn := newUsersServer() - defer us.Close() - - cases := []struct { - desc string - token string - id string - status int - authnRes smqauthn.Session - authnErr error - svcErr error - err error - }{ - { - desc: "view profile with valid token", - token: validToken, - id: user.ID, - status: http.StatusOK, - authnRes: verifiedSession, - err: nil, - }, - { - desc: "view profile with invalid token", - token: inValidToken, - id: user.ID, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - authnRes: smqauthn.Session{}, - err: svcerr.ErrAuthentication, - }, - { - desc: "view profile with empty token", - token: "", - id: user.ID, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - authnRes: smqauthn.Session{}, - err: apiutil.ErrBearerToken, - }, - { - desc: "view profile with service error", - token: validToken, - id: user.ID, - status: http.StatusUnprocessableEntity, - authnRes: verifiedSession, - svcErr: svcerr.ErrViewEntity, - err: svcerr.ErrViewEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - user: us.Client(), - method: http.MethodGet, - url: fmt.Sprintf("%s/users/profile", us.URL), - token: tc.token, - } - - authnCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.authnRes, tc.authnErr) - svcCall := svc.On("ViewProfile", mock.Anything, tc.authnRes).Return(users.User{}, tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - var errRes respBody - err = json.NewDecoder(res.Body).Decode(&errRes) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if errRes.Err != "" || errRes.Message != "" { - err = errors.Wrap(errors.New(errRes.Err), errors.New(errRes.Message)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authnCall.Unset() - }) - } -} - -func TestListUsers(t *testing.T) { - us, svc, authn := newUsersServer() - defer us.Close() - - cases := []struct { - desc string - query string - token string - pageMeta users.Page - listUsersResponse users.UsersPage - status int - authnRes smqauthn.Session - authnErr error - err error - }{ - { - desc: "list users as admin with valid token", - token: validToken, - status: http.StatusOK, - pageMeta: users.Page{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - }, - listUsersResponse: users.UsersPage{ - Page: users.Page{ - Total: 1, - }, - Users: []users.User{user}, - }, - authnRes: verifiedSession, - err: nil, - }, - { - desc: "list users with empty token", - token: "", - status: http.StatusUnauthorized, - authnRes: smqauthn.Session{}, - authnErr: svcerr.ErrAuthentication, - err: apiutil.ErrBearerToken, - }, - { - desc: "list users with invalid token", - token: inValidToken, - status: http.StatusUnauthorized, - authnRes: smqauthn.Session{}, - authnErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "list users with offset", - token: validToken, - pageMeta: users.Page{ - Offset: 1, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - }, - listUsersResponse: users.UsersPage{ - Page: users.Page{ - Offset: 1, - Total: 1, - }, - Users: []users.User{user}, - }, - query: "offset=1", - status: http.StatusOK, - authnRes: verifiedSession, - err: nil, - }, - { - desc: "list users with invalid offset", - token: validToken, - query: "offset=invalid", - status: http.StatusBadRequest, - authnRes: verifiedSession, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list users with limit", - token: validToken, - pageMeta: users.Page{ - Offset: 0, - Limit: 1, - Order: api.DefOrder, - Dir: api.DefDir, - }, - listUsersResponse: users.UsersPage{ - Page: users.Page{ - Limit: 1, - Total: 1, - }, - Users: []users.User{user}, - }, - query: "limit=1", - status: http.StatusOK, - authnRes: verifiedSession, - err: nil, - }, - { - desc: "list users with invalid limit", - token: validToken, - query: "limit=invalid", - status: http.StatusBadRequest, - authnRes: verifiedSession, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list users with limit greater than max", - token: validToken, - query: fmt.Sprintf("limit=%d", api.MaxLimitSize+1), - status: http.StatusBadRequest, - authnRes: verifiedSession, - err: apiutil.ErrLimitSize, - }, - { - desc: "list users with username", - token: validToken, - pageMeta: users.Page{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Username: "username", - }, - listUsersResponse: users.UsersPage{ - Page: users.Page{ - Total: 1, - }, - Users: []users.User{user}, - }, - query: "username=username", - status: http.StatusOK, - authnRes: verifiedSession, - err: nil, - }, - { - desc: "list users with duplicate username", - token: validToken, - query: "username=1&username=2", - status: http.StatusBadRequest, - authnRes: verifiedSession, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list users with first name", - token: validToken, - pageMeta: users.Page{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - FirstName: "firstname", - }, - listUsersResponse: users.UsersPage{ - Page: users.Page{ - Total: 1, - }, - Users: []users.User{user}, - }, - query: "first_name=firstname", - status: http.StatusOK, - authnRes: verifiedSession, - err: nil, - }, - { - desc: "list users with duplicate firstname", - token: validToken, - query: "status=invalid", - status: http.StatusBadRequest, - authnRes: verifiedSession, - err: svcerr.ErrInvalidStatus, - }, - { - desc: "list users with duplicate status", - token: validToken, - query: "status=enabled&status=disabled", - status: http.StatusBadRequest, - authnRes: verifiedSession, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list users with lastname", - token: validToken, - pageMeta: users.Page{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - LastName: "lastname", - }, - listUsersResponse: users.UsersPage{ - Page: users.Page{ - Total: 1, - }, - Users: []users.User{user}, - }, - query: "last_name=lastname", - status: http.StatusOK, - authnRes: verifiedSession, - err: nil, - }, - { - desc: "list users with duplicate lastname", - token: validToken, - query: "last_name=lastname1&last_name=lastname2", - status: http.StatusBadRequest, - authnRes: verifiedSession, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list users with status", - token: validToken, - pageMeta: users.Page{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Status: users.EnabledStatus, - }, - listUsersResponse: users.UsersPage{ - Page: users.Page{ - Total: 1, - }, - Users: []users.User{user}, - }, - query: "status=enabled", - status: http.StatusOK, - authnRes: verifiedSession, - err: nil, - }, - { - desc: "list users with invalid status", - token: validToken, - query: "status=invalid", - status: http.StatusBadRequest, - authnRes: verifiedSession, - err: svcerr.ErrInvalidStatus, - }, - { - desc: "list users with duplicate status", - token: validToken, - query: "status=enabled&status=disabled", - status: http.StatusBadRequest, - authnRes: verifiedSession, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list users with single tag", - token: validToken, - pageMeta: users.Page{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Tags: users.TagsQuery{Elements: []string{"tag1"}, Operator: users.OrOp}, - }, - listUsersResponse: users.UsersPage{ - Page: users.Page{ - Total: 1, - }, - Users: []users.User{user}, - }, - query: "tags=tag1", - status: http.StatusOK, - authnRes: verifiedSession, - err: nil, - }, - { - desc: "list users with multiple tags and OR operator", - token: validToken, - pageMeta: users.Page{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Tags: users.TagsQuery{Elements: []string{"tag1", "tag2", "tag3"}, Operator: users.OrOp}, - }, - listUsersResponse: users.UsersPage{ - Page: users.Page{ - Total: 1, - }, - Users: []users.User{user}, - }, - query: "tags=tag1,tag2,tag3", - status: http.StatusOK, - authnRes: verifiedSession, - err: nil, - }, - { - desc: "list users with multiple tags and AND operator", - token: validToken, - pageMeta: users.Page{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Tags: users.TagsQuery{Elements: []string{"tag1", "tag2", "tag3"}, Operator: users.AndOp}, - }, - listUsersResponse: users.UsersPage{ - Page: users.Page{ - Total: 1, - }, - Users: []users.User{user}, - }, - query: "tags=tag1%2Btag2%2Btag3", - status: http.StatusOK, - authnRes: verifiedSession, - err: nil, - }, - { - desc: "list users with duplicate tags", - token: validToken, - query: "tags=tag1&tags=tag2", - status: http.StatusBadRequest, - authnRes: verifiedSession, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list users with metadata", - token: validToken, - pageMeta: users.Page{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Metadata: users.Metadata{"domain": "example.com"}, - }, - listUsersResponse: users.UsersPage{ - Page: users.Page{ - Total: 1, - }, - Users: []users.User{user}, - }, - query: "metadata=" + url.PathEscape(`{"domain": "example.com"}`), - status: http.StatusOK, - authnRes: verifiedSession, - err: nil, - }, - { - desc: "list users with invalid metadata", - token: validToken, - query: "metadata=invalid", - status: http.StatusBadRequest, - authnRes: verifiedSession, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list users with duplicate metadata", - token: validToken, - query: fmt.Sprintf("metadata=%s&metadata=%s", url.PathEscape(`{"domain": "example.com"}`), url.PathEscape(`{"domain": "example.com"}`)), - status: http.StatusBadRequest, - authnRes: verifiedSession, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list users with email", - token: validToken, - query: fmt.Sprintf("email=%s", user.Email), - pageMeta: users.Page{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Email: user.Email, - }, - listUsersResponse: users.UsersPage{ - Page: users.Page{ - Total: 1, - }, - Users: []users.User{user}, - }, - status: http.StatusOK, - authnRes: verifiedSession, - err: nil, - }, - { - desc: "list users with duplicate email", - token: validToken, - query: "email=1&email=2", - status: http.StatusBadRequest, - authnRes: verifiedSession, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list users with order", - pageMeta: users.Page{ - Offset: 0, - Limit: 10, - Order: "username", - Dir: api.DefDir, - }, - listUsersResponse: users.UsersPage{ - Page: users.Page{ - Total: 1, - }, - Users: []users.User{ - user, - }, - }, - token: validToken, - query: "order=username", - status: http.StatusOK, - authnRes: verifiedSession, - err: nil, - }, - { - desc: "list users with duplicate order", - token: validToken, - query: "order=name&order=name", - status: http.StatusBadRequest, - authnRes: verifiedSession, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list users with invalid order direction", - token: validToken, - query: "dir=invalid", - status: http.StatusBadRequest, - authnRes: verifiedSession, - err: apiutil.ErrInvalidDirection, - }, - { - desc: "list users with duplicate order direction", - token: validToken, - query: "dir=asc&dir=asc", - status: http.StatusBadRequest, - authnRes: verifiedSession, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list users with created_from", - token: validToken, - query: "created_from=2024-01-01T00:00:00Z", - pageMeta: users.Page{ - Offset: 0, - Limit: 10, - Dir: api.DefDir, - Order: api.DefOrder, - CreatedFrom: validTimeStamp, - }, - listUsersResponse: users.UsersPage{ - Page: users.Page{ - Total: 1, - }, - Users: []users.User{user}, - }, - status: http.StatusOK, - authnRes: verifiedSession, - err: nil, - }, - { - desc: "list users with created_to", - token: validToken, - query: "created_to=2024-01-01T00:00:00Z", - pageMeta: users.Page{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - CreatedTo: validTimeStamp, - }, - listUsersResponse: users.UsersPage{ - Page: users.Page{ - Total: 1, - }, - Users: []users.User{user}, - }, - status: http.StatusOK, - authnRes: verifiedSession, - err: nil, - }, - { - desc: "list users with both created_from and created_to", - token: validToken, - query: "created_from=2024-01-01T00:00:00Z&created_to=2024-01-01T00:00:00Z", - pageMeta: users.Page{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - CreatedFrom: validTimeStamp, - CreatedTo: validTimeStamp, - }, - listUsersResponse: users.UsersPage{ - Page: users.Page{ - Total: 1, - }, - Users: []users.User{user}, - }, - status: http.StatusOK, - authnRes: verifiedSession, - err: nil, - }, - { - desc: "list users with invalid created_from format", - token: validToken, - query: "created_from=invalid-date", - status: http.StatusBadRequest, - authnRes: verifiedSession, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list users with invalid created_to format", - token: validToken, - query: "created_to=invalid-date", - status: http.StatusBadRequest, - authnRes: verifiedSession, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list users with duplicate created_from", - token: validToken, - query: "created_from=2024-01-01T00:00:00Z&created_from=2024-01-02T00:00:00Z", - status: http.StatusBadRequest, - authnRes: verifiedSession, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list users with duplicate created_to", - token: validToken, - query: "created_to=2024-12-31T23:59:59Z&created_to=2024-12-30T23:59:59Z", - status: http.StatusBadRequest, - authnRes: verifiedSession, - err: apiutil.ErrInvalidQueryParams, - }, - { - desc: "list users with created_from and others", - token: validToken, - query: "created_from=2024-01-01T00:00:00Z&status=enabled&limit=10", - pageMeta: users.Page{ - Offset: 0, - Limit: 10, - Order: api.DefOrder, - Dir: api.DefDir, - Status: users.EnabledStatus, - CreatedFrom: validTimeStamp, - }, - listUsersResponse: users.UsersPage{ - Page: users.Page{ - Total: 1, - Limit: 10, - }, - Users: []users.User{user}, - }, - status: http.StatusOK, - authnRes: verifiedSession, - err: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - user: us.Client(), - method: http.MethodGet, - url: us.URL + "/users?" + tc.query, - contentType: contentType, - token: tc.token, - } - - authnCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.authnRes, tc.authnErr) - svcCall := svc.On("ListUsers", mock.Anything, tc.authnRes, tc.pageMeta).Return(tc.listUsersResponse, tc.err) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - var bodyRes respBody - err = json.NewDecoder(res.Body).Decode(&bodyRes) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if bodyRes.Err != "" || bodyRes.Message != "" { - err = errors.Wrap(errors.New(bodyRes.Err), errors.New(bodyRes.Message)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authnCall.Unset() - }) - } -} - -func TestSearchUsers(t *testing.T) { - us, svc, authn := newUsersServer() - defer us.Close() - - cases := []struct { - desc string - token string - page users.Page - status int - query string - listUsersResponse users.UsersPage - authnErr error - svcErr error - err error - }{ - { - desc: "search users with valid token", - token: validToken, - status: http.StatusOK, - query: "username=username", - listUsersResponse: users.UsersPage{ - Page: users.Page{ - Total: 1, - }, - Users: []users.User{user}, - }, - err: nil, - }, - { - desc: "search users with empty token", - token: "", - query: "username=username", - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: apiutil.ErrBearerToken, - }, - { - desc: "search users with invalid token", - token: inValidToken, - query: "username=username", - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "search users with offset", - token: validToken, - listUsersResponse: users.UsersPage{ - Page: users.Page{ - Offset: 1, - Total: 1, - }, - Users: []users.User{user}, - }, - query: "username=username&offset=1", - status: http.StatusOK, - err: nil, - }, - { - desc: "search users with invalid offset", - token: validToken, - query: "username=username&offset=invalid", - status: http.StatusBadRequest, - err: apiutil.ErrValidation, - }, - { - desc: "search users with limit", - token: validToken, - listUsersResponse: users.UsersPage{ - Page: users.Page{ - Limit: 1, - Total: 1, - }, - Users: []users.User{user}, - }, - query: "username=username&limit=1", - status: http.StatusOK, - err: nil, - }, - { - desc: "search users with invalid limit", - token: validToken, - query: "username=username&limit=invalid", - status: http.StatusBadRequest, - err: apiutil.ErrValidation, - }, - { - desc: "search users with empty query", - token: validToken, - query: "", - status: http.StatusBadRequest, - err: apiutil.ErrEmptySearchQuery, - }, - { - desc: "search users with invalid length of query", - token: validToken, - query: "username=a", - status: http.StatusBadRequest, - err: apiutil.ErrLenSearchQuery, - }, - { - desc: "serach users with service error", - token: validToken, - query: "username=username", - status: http.StatusUnprocessableEntity, - svcErr: svcerr.ErrViewEntity, - err: svcerr.ErrViewEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - user: us.Client(), - method: http.MethodGet, - url: fmt.Sprintf("%s/users/search?", us.URL) + tc.query, - token: tc.token, - } - - authnCall := authn.On("Authenticate", mock.Anything, tc.token).Return(verifiedSession, tc.authnErr) - svcCall := svc.On("SearchUsers", mock.Anything, mock.Anything).Return( - users.UsersPage{ - Page: tc.listUsersResponse.Page, - Users: tc.listUsersResponse.Users, - }, - tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authnCall.Unset() - }) - } -} - -func TestUpdate(t *testing.T) { - us, svc, authn := newUsersServer() - defer us.Close() - - newName := "newname" - newMetadata := users.Metadata{"newkey": "newvalue"} - - cases := []struct { - desc string - id string - data string - userResponse users.User - token string - authnRes smqauthn.Session - authnErr error - contentType string - status int - err error - }{ - { - desc: "update as admin user with valid token", - id: user.ID, - data: fmt.Sprintf(`{"name":"%s","metadata":%s, "private_metadata":%s}`, newName, toJSON(newMetadata), toJSON(newMetadata)), - token: validToken, - authnRes: verifiedSession, - contentType: contentType, - userResponse: users.User{ - ID: user.ID, - FirstName: newName, - Metadata: newMetadata, - PrivateMetadata: newMetadata, - }, - status: http.StatusOK, - err: nil, - }, - { - desc: "update as normal user with valid token", - id: user.ID, - data: fmt.Sprintf(`{"name":"%s","metadata":%s}`, newName, toJSON(newMetadata)), - token: validToken, - authnRes: verifiedSession, - contentType: contentType, - userResponse: users.User{ - ID: user.ID, - FirstName: newName, - Metadata: newMetadata, - }, - status: http.StatusOK, - err: nil, - }, - { - desc: "update user with invalid token", - id: user.ID, - data: fmt.Sprintf(`{"name":"%s","metadata":%s}`, newName, toJSON(newMetadata)), - token: inValidToken, - authnRes: smqauthn.Session{UserID: validID, DomainID: validID, Verified: true}, - contentType: contentType, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "update user with empty token", - id: user.ID, - data: fmt.Sprintf(`{"name":"%s","metadata":%s}`, newName, toJSON(newMetadata)), - token: "", - authnRes: smqauthn.Session{UserID: validID, DomainID: validID, Verified: true}, - contentType: contentType, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: apiutil.ErrBearerToken, - }, - { - desc: "update user with invalid id", - id: inValid, - data: fmt.Sprintf(`{"name":"%s","metadata":%s}`, newName, toJSON(newMetadata)), - token: validToken, - authnRes: verifiedSession, - contentType: contentType, - status: http.StatusForbidden, - err: svcerr.ErrAuthorization, - }, - { - desc: "update user with invalid contentype", - id: user.ID, - data: fmt.Sprintf(`{"name":"%s","metadata":%s}`, newName, toJSON(newMetadata)), - token: validToken, - authnRes: verifiedSession, - contentType: "application/xml", - status: http.StatusUnsupportedMediaType, - err: apiutil.ErrUnsupportedContentType, - }, - { - desc: "update user with malformed data", - id: user.ID, - data: fmt.Sprintf(`{"name":%s}`, "invalid"), - token: validToken, - authnRes: verifiedSession, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrMalformedRequestBody, - }, - { - desc: "update user with empty id", - id: " ", - data: fmt.Sprintf(`{"name":"%s","metadata":%s}`, newName, toJSON(newMetadata)), - token: validToken, - authnRes: verifiedSession, - contentType: contentType, - status: http.StatusUnprocessableEntity, - err: svcerr.ErrViewEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - user: us.Client(), - method: http.MethodPatch, - url: fmt.Sprintf("%s/users/%s", us.URL, tc.id), - contentType: tc.contentType, - token: tc.token, - body: strings.NewReader(tc.data), - } - authnCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.authnRes, tc.authnErr) - svcCall := svc.On("Update", mock.Anything, tc.authnRes, tc.id, mock.Anything).Return(tc.userResponse, tc.err) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - var resBody respBody - err = json.NewDecoder(res.Body).Decode(&resBody) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if resBody.Err != "" || resBody.Message != "" { - err = errors.Wrap(errors.New(resBody.Err), errors.New(resBody.Message)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authnCall.Unset() - }) - } -} - -func TestUpdateTags(t *testing.T) { - us, svc, authn := newUsersServer() - defer us.Close() - - defer us.Close() - newTag := "newtag" - - cases := []struct { - desc string - id string - data string - contentType string - userResponse users.User - token string - authnRes smqauthn.Session - authnErr error - status int - err error - }{ - { - desc: "updateuser tags as admin with valid token", - id: user.ID, - data: fmt.Sprintf(`{"tags":["%s"]}`, newTag), - contentType: contentType, - userResponse: users.User{ - ID: user.ID, - Tags: []string{newTag}, - }, - token: validToken, - authnRes: verifiedSession, - status: http.StatusOK, - err: nil, - }, - { - desc: "updateuser tags as normal user with valid token", - id: user.ID, - data: fmt.Sprintf(`{"tags":["%s"]}`, newTag), - contentType: contentType, - userResponse: users.User{ - ID: user.ID, - Tags: []string{newTag}, - }, - token: validToken, - authnRes: verifiedSession, - status: http.StatusOK, - err: nil, - }, - { - desc: "update user tags with empty token", - id: user.ID, - data: fmt.Sprintf(`{"tags":["%s"]}`, newTag), - contentType: contentType, - token: "", - authnRes: smqauthn.Session{UserID: validID, DomainID: validID, Verified: true}, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: apiutil.ErrBearerToken, - }, - { - desc: "update user tags with invalid token", - id: user.ID, - data: fmt.Sprintf(`{"tags":["%s"]}`, newTag), - contentType: contentType, - token: inValidToken, - authnRes: smqauthn.Session{UserID: validID, DomainID: validID, Verified: true}, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "update user tags with invalid id", - id: user.ID, - data: fmt.Sprintf(`{"tags":["%s"]}`, newTag), - contentType: contentType, - token: validToken, - authnRes: verifiedSession, - status: http.StatusForbidden, - err: svcerr.ErrAuthorization, - }, - { - desc: "update user tags with invalid contentype", - id: user.ID, - data: fmt.Sprintf(`{"tags":["%s"]}`, newTag), - contentType: "application/xml", - token: validToken, - authnRes: verifiedSession, - status: http.StatusUnsupportedMediaType, - err: apiutil.ErrUnsupportedContentType, - }, - { - desc: "update user tags with empty id", - id: "", - data: fmt.Sprintf(`{"tags":["%s"]}`, newTag), - contentType: contentType, - token: validToken, - authnRes: verifiedSession, - status: http.StatusBadRequest, - err: apiutil.ErrMissingID, - }, - { - desc: "update user with malfomed data", - id: user.ID, - data: fmt.Sprintf(`{"tags":%s}`, newTag), - contentType: contentType, - token: validToken, - authnRes: verifiedSession, - status: http.StatusBadRequest, - err: apiutil.ErrMalformedRequestBody, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - user: us.Client(), - method: http.MethodPatch, - url: fmt.Sprintf("%s/users/%s/tags", us.URL, tc.id), - contentType: tc.contentType, - token: tc.token, - body: strings.NewReader(tc.data), - } - - authnCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.authnRes, tc.authnErr) - svcCall := svc.On("UpdateTags", mock.Anything, tc.authnRes, tc.id, mock.Anything).Return(tc.userResponse, tc.err) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - var resBody respBody - err = json.NewDecoder(res.Body).Decode(&resBody) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if resBody.Err != "" || resBody.Message != "" { - err = errors.Wrap(errors.New(resBody.Err), errors.New(resBody.Message)) - } - if err == nil { - assert.Equal(t, tc.userResponse.Tags, resBody.Tags, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.userResponse.Tags, resBody.Tags)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authnCall.Unset() - }) - } -} - -func TestUpdateEmail(t *testing.T) { - us, svc, authn := newUsersServer() - defer us.Close() - - newuseremail := "newuseremail@example.com" - - cases := []struct { - desc string - data string - user users.User - contentType string - token string - authnRes smqauthn.Session - authnErr error - status int - svcErr error - err error - }{ - { - desc: "update user email as admin with valid token", - data: fmt.Sprintf(`{"email": "%s"}`, newuseremail), - user: users.User{ - ID: user.ID, - Email: newuseremail, - Credentials: users.Credentials{ - Secret: "secret", - }, - }, - contentType: contentType, - token: validToken, - authnRes: verifiedSession, - status: http.StatusOK, - err: nil, - }, - { - desc: "update user email as normal user with valid token", - data: fmt.Sprintf(`{"email": "%s"}`, newuseremail), - user: users.User{ - ID: user.ID, - Email: newuseremail, - Credentials: users.Credentials{ - Secret: "secret", - }, - }, - contentType: contentType, - token: validToken, - authnRes: smqauthn.Session{UserID: validID, DomainID: validID}, - status: http.StatusOK, - err: nil, - }, - { - desc: "update user email with empty token", - data: fmt.Sprintf(`{"email": "%s"}`, newuseremail), - user: users.User{ - ID: user.ID, - Email: newuseremail, - Credentials: users.Credentials{ - Secret: "secret", - }, - }, - contentType: contentType, - token: "", - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: apiutil.ErrBearerToken, - }, - { - desc: "update user email with invalid token", - data: fmt.Sprintf(`{"email": "%s"}`, newuseremail), - user: users.User{ - ID: user.ID, - Email: newuseremail, - Credentials: users.Credentials{ - Secret: "secret", - }, - }, - contentType: contentType, - token: inValid, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "update user email with empty id", - data: fmt.Sprintf(`{"email": "%s"}`, newuseremail), - user: users.User{ - ID: "", - Email: newuseremail, - Credentials: users.Credentials{ - Secret: "secret", - }, - }, - contentType: contentType, - token: validToken, - authnRes: smqauthn.Session{UserID: validID, DomainID: validID}, - status: http.StatusBadRequest, - err: apiutil.ErrMissingID, - }, - { - desc: "update user email with invalid contentype", - data: fmt.Sprintf(`{"email": "%s"}`, ""), - user: users.User{ - ID: user.ID, - Email: newuseremail, - Credentials: users.Credentials{ - Secret: "secret", - }, - }, - contentType: "application/xml", - token: validToken, - status: http.StatusUnsupportedMediaType, - err: apiutil.ErrUnsupportedContentType, - }, - { - desc: "update user email with malformed data", - data: fmt.Sprintf(`{"email": %s}`, "invalid"), - user: users.User{ - ID: user.ID, - Email: "", - Credentials: users.Credentials{ - Secret: "secret", - }, - }, - token: validToken, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrMalformedRequestBody, - }, - { - desc: "update user email with service error", - data: fmt.Sprintf(`{"email": "%s"}`, newuseremail), - user: users.User{ - ID: user.ID, - Email: newuseremail, - }, - contentType: contentType, - token: validToken, - authnRes: verifiedSession, - status: http.StatusUnprocessableEntity, - svcErr: svcerr.ErrUpdateEntity, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - user: us.Client(), - method: http.MethodPatch, - url: fmt.Sprintf("%s/users/%s/email", us.URL, tc.user.ID), - contentType: tc.contentType, - token: tc.token, - body: strings.NewReader(tc.data), - } - - authnCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.authnRes, tc.authnErr) - svcCall := svc.On("UpdateEmail", mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(tc.user, tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - var resBody respBody - err = json.NewDecoder(res.Body).Decode(&resBody) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if resBody.Err != "" || resBody.Message != "" { - err = errors.Wrap(errors.New(resBody.Err), errors.New(resBody.Message)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authnCall.Unset() - }) - } -} - -func TestUpdateUsername(t *testing.T) { - us, svc, authn := newUsersServer() - defer us.Close() - - newusername := "newusername" - - cases := []struct { - desc string - data string - user users.User - contentType string - token string - authnRes smqauthn.Session - authnErr error - status int - err error - }{ - { - desc: "update username as admin with valid token", - data: fmt.Sprintf(`{"username": "%s"}`, newusername), - user: users.User{ - ID: user.ID, - Credentials: users.Credentials{ - Username: newusername, - }, - }, - contentType: contentType, - token: validToken, - authnRes: verifiedSession, - status: http.StatusOK, - err: nil, - }, - { - desc: "update username with empty token", - data: fmt.Sprintf(`{"username": "%s"}`, newusername), - user: users.User{ - ID: user.ID, - Credentials: users.Credentials{ - Username: newusername, - }, - }, - authnRes: smqauthn.Session{Type: smqauthn.AccessToken, Verified: true}, - contentType: contentType, - token: "", - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: apiutil.ErrBearerToken, - }, - { - desc: "update username with invalid token", - data: fmt.Sprintf(`{"username": "%s"}`, newusername), - user: users.User{ - ID: user.ID, - Credentials: users.Credentials{ - Username: newusername, - }, - }, - authnRes: smqauthn.Session{Type: smqauthn.AccessToken, Verified: true}, - contentType: contentType, - token: inValid, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "update username with empty id", - data: fmt.Sprintf(`{"username": "%s"}`, newusername), - user: users.User{ - ID: "", - Credentials: users.Credentials{ - Username: newusername, - }, - }, - authnRes: smqauthn.Session{UserID: validID, DomainID: validID, Verified: true}, - contentType: contentType, - token: validToken, - status: http.StatusBadRequest, - err: apiutil.ErrMissingID, - }, - { - desc: "update username with invalid contentype", - data: fmt.Sprintf(`{"username": "%s"}`, ""), - user: users.User{ - ID: user.ID, - Credentials: users.Credentials{ - Username: newusername, - }, - }, - authnRes: smqauthn.Session{UserID: validID, DomainID: validID, Verified: true}, - contentType: "application/xml", - token: validToken, - status: http.StatusUnsupportedMediaType, - err: apiutil.ErrUnsupportedContentType, - }, - { - desc: "update user email with malformed data", - data: fmt.Sprintf(`{"email": %s}`, "invalid"), - user: users.User{ - ID: user.ID, - Credentials: users.Credentials{ - Username: newusername, - }, - }, - authnRes: smqauthn.Session{UserID: validID, DomainID: validID, Verified: true}, - token: validToken, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrMalformedRequestBody, - }, - { - desc: "update username with invalid username", - data: fmt.Sprintf(`{"username": "%s"}`, "invalid"), - user: users.User{ - ID: user.ID, - Credentials: users.Credentials{ - Username: newusername, - }, - }, - authnRes: smqauthn.Session{UserID: validID, DomainID: validID, Verified: true}, - contentType: contentType, - token: validToken, - status: http.StatusUnprocessableEntity, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - user: us.Client(), - method: http.MethodPatch, - url: fmt.Sprintf("%s/users/%s/username", us.URL, tc.user.ID), - contentType: tc.contentType, - token: tc.token, - body: strings.NewReader(tc.data), - } - - authnCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.authnRes, tc.authnErr) - svcCall := svc.On("UpdateUsername", mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(tc.user, tc.err) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - var resBody respBody - err = json.NewDecoder(res.Body).Decode(&resBody) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if resBody.Err != "" || resBody.Message != "" { - err = errors.Wrap(errors.New(resBody.Err), errors.New(resBody.Message)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authnCall.Unset() - }) - } -} - -func TestUpdateProfilePicture(t *testing.T) { - us, svc, authn := newUsersServer() - defer us.Close() - - newprofilepicture := "https://example.com/newprofilepicture" - - cases := []struct { - desc string - data string - user users.User - contentType string - token string - authnRes smqauthn.Session - authnErr error - status int - svcErr error - err error - }{ - { - desc: "update profile picture as admin with valid token", - data: fmt.Sprintf(`{"profile_picture": "%s"}`, newprofilepicture), - user: users.User{ - ID: user.ID, - ProfilePicture: newprofilepicture, - }, - contentType: contentType, - token: validToken, - authnRes: smqauthn.Session{UserID: validID, DomainID: domainID, Role: smqauthn.SuperAdminRole}, - status: http.StatusOK, - err: nil, - }, - { - desc: "update profile picture with empty token", - data: fmt.Sprintf(`{"profile_picture": "%s"}`, newprofilepicture), - user: users.User{}, - authnRes: smqauthn.Session{Type: smqauthn.AccessToken, Verified: true}, - contentType: contentType, - token: "", - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: apiutil.ErrBearerToken, - }, - { - desc: "update profile_picture with invalid token", - data: fmt.Sprintf(`{"profile_picture": "%s"}`, newprofilepicture), - user: users.User{}, - contentType: contentType, - token: inValid, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "update profile_picture with empty id", - data: fmt.Sprintf(`{"profile_picture": "%s"}`, newprofilepicture), - user: users.User{ - ID: "", - ProfilePicture: newprofilepicture, - }, - - contentType: contentType, - token: validToken, - authnRes: smqauthn.Session{UserID: validID, DomainID: validID, Verified: true}, - status: http.StatusBadRequest, - err: apiutil.ErrMissingID, - }, - { - desc: "update profile_picture with invalid contentype", - data: fmt.Sprintf(`{"profile_picture": "%s"}`, ""), - user: users.User{ - ID: user.ID, - ProfilePicture: newprofilepicture, - }, - authnRes: smqauthn.Session{Type: smqauthn.AccessToken, Verified: true}, - contentType: "application/xml", - token: validToken, - status: http.StatusUnsupportedMediaType, - err: apiutil.ErrUnsupportedContentType, - }, - { - desc: "update profile picture with malformed data", - data: fmt.Sprintf(`{"profile_picture": %s}`, "invalid"), - user: users.User{}, - authnRes: smqauthn.Session{Type: smqauthn.AccessToken, Verified: true}, - token: validToken, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrMalformedRequestBody, - }, - { - desc: "update profile picture with failed to update", - data: fmt.Sprintf(`{"profile_picture": "%s"}`, "invalid"), - user: users.User{ - ID: user.ID, - }, - authnRes: smqauthn.Session{Type: smqauthn.AccessToken, Verified: true}, - contentType: contentType, - token: validToken, - status: http.StatusUnprocessableEntity, - svcErr: svcerr.ErrUpdateEntity, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - user: us.Client(), - method: http.MethodPatch, - url: fmt.Sprintf("%s/users/%s/picture", us.URL, tc.user.ID), - contentType: tc.contentType, - token: tc.token, - body: strings.NewReader(tc.data), - } - - authnCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.authnRes, tc.authnErr) - svcCall := svc.On("UpdateProfilePicture", mock.Anything, tc.authnRes, tc.user.ID, mock.Anything).Return(tc.user, tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - var resBody respBody - err = json.NewDecoder(res.Body).Decode(&resBody) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if resBody.Err != "" || resBody.Message != "" { - err = errors.Wrap(errors.New(resBody.Err), errors.New(resBody.Message)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authnCall.Unset() - }) - } -} - -func TestPasswordResetRequest(t *testing.T) { - us, svc, _ := newUsersServer() - defer us.Close() - - testemail := "test@example.com" - testhost := "example.com" - - cases := []struct { - desc string - data string - contentType string - referer string - status int - generateErr error - sendErr error - err error - }{ - { - desc: "password reset request with valid email", - data: fmt.Sprintf(`{"email": "%s", "host": "%s"}`, testemail, testhost), - contentType: contentType, - referer: testReferer, - status: http.StatusCreated, - err: nil, - }, - { - desc: "password reset request with empty email", - data: fmt.Sprintf(`{"email": "%s", "host": "%s"}`, "", testhost), - contentType: contentType, - referer: testReferer, - status: http.StatusBadRequest, - err: apiutil.ErrValidation, - }, - { - desc: "password reset request with invalid email", - data: fmt.Sprintf(`{"email": "%s", "host": "%s"}`, "invalid", testhost), - contentType: contentType, - referer: testReferer, - status: http.StatusNotFound, - generateErr: svcerr.ErrNotFound, - err: svcerr.ErrNotFound, - }, - { - desc: "password reset with malformed data", - data: fmt.Sprintf(`{"email": %s, "host": %s}`, testemail, testhost), - contentType: contentType, - referer: testReferer, - status: http.StatusBadRequest, - err: apiutil.ErrValidation, - }, - { - desc: "password reset with invalid content type", - data: fmt.Sprintf(`{"email": "%s", "host": "%s"}`, testemail, testhost), - contentType: "application/xml", - referer: testReferer, - status: http.StatusUnsupportedMediaType, - err: apiutil.ErrValidation, - }, - { - desc: "password reset with failed to issue token", - data: fmt.Sprintf(`{"email": "%s", "host": "%s"}`, testemail, testhost), - contentType: contentType, - referer: testReferer, - status: http.StatusUnauthorized, - generateErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - user: us.Client(), - method: http.MethodPost, - url: fmt.Sprintf("%s/password/reset-request", us.URL), - contentType: tc.contentType, - referer: tc.referer, - body: strings.NewReader(tc.data), - } - svcCall := svc.On("SendPasswordReset", mock.Anything, mock.Anything).Return(tc.generateErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - }) - } -} - -func TestSendVerification(t *testing.T) { - us, svc, authn := newUsersServer() - defer us.Close() - - cases := []struct { - desc string - token string - status int - authnRes smqauthn.Session - authnErr error - svcErr error - err error - }{ - { - desc: "send verification with valid token", - token: validToken, - status: http.StatusOK, - authnRes: verifiedSession, - err: nil, - }, - { - desc: "send verification with invalid token", - token: inValidToken, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - authnRes: smqauthn.Session{}, - err: svcerr.ErrAuthentication, - }, - { - desc: "send verification with empty token", - token: "", - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - authnRes: smqauthn.Session{}, - err: apiutil.ErrBearerToken, - }, - { - desc: "send verification with service error", - token: validToken, - status: http.StatusUnprocessableEntity, - authnRes: verifiedSession, - svcErr: svcerr.ErrCreateEntity, - err: svcerr.ErrCreateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - user: us.Client(), - method: http.MethodPost, - url: fmt.Sprintf("%s/users/send-verification", us.URL), - token: tc.token, - } - - authnCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.authnRes, tc.authnErr) - svcCall := svc.On("SendVerification", mock.Anything, tc.authnRes).Return(tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - body, err := io.ReadAll(res.Body) - if err != nil { - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while reading response body: %s", tc.desc, err)) - } - defer res.Body.Close() - var errRes respBody - if len(body) > 0 { - if err := json.Unmarshal(body, &errRes); err != nil { - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while unmarshal response body: %s", tc.desc, err)) - } - } - if errRes.Err != "" || errRes.Message != "" { - err = errors.Wrap(errors.New(errRes.Err), errors.New(errRes.Message)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %v got %v", tc.desc, tc.err, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authnCall.Unset() - }) - } -} - -func TestVerifyEmail(t *testing.T) { - us, svc, _ := newUsersServer() - defer us.Close() - - cases := []struct { - desc string - token string - status int - svcErr error - err error - }{ - { - desc: "verify email with valid token", - token: validToken, - status: http.StatusOK, - err: nil, - }, - { - desc: "verify email with empty token", - token: "", - status: http.StatusBadRequest, - err: apiutil.ErrInvalidVerification, - }, - { - desc: "verify email with service error", - token: validToken, - status: http.StatusUnprocessableEntity, - svcErr: svcerr.ErrUpdateEntity, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - user: us.Client(), - method: http.MethodGet, - url: fmt.Sprintf("%s/verify-email?token=%s", us.URL, tc.token), - } - - svcCall := svc.On("VerifyEmail", mock.Anything, mock.Anything).Return(users.User{}, tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - body, err := io.ReadAll(res.Body) - if err != nil { - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while reading response body: %s", tc.desc, err)) - } - defer res.Body.Close() - var errRes respBody - if len(body) > 0 { - if err := json.Unmarshal(body, &errRes); err != nil { - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while unmarshal response body: %s", tc.desc, err)) - } - } - if errRes.Err != "" || errRes.Message != "" { - err = errors.Wrap(errors.New(errRes.Err), errors.New(errRes.Message)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %v got %v", tc.desc, tc.err, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - }) - } -} - -func TestPasswordReset(t *testing.T) { - us, svc, authn := newUsersServer() - defer us.Close() - - strongPass := "StrongPassword" - - cases := []struct { - desc string - data string - token string - contentType string - status int - authnRes smqauthn.Session - authnErr error - svcErr error - err error - }{ - { - desc: "password reset with valid token", - data: fmt.Sprintf(`{"token": "%s", "password": "%s", "confirm_password": "%s"}`, validToken, strongPass, strongPass), - token: validToken, - authnRes: smqauthn.Session{Type: smqauthn.AccessToken, Verified: true}, - contentType: contentType, - status: http.StatusCreated, - err: nil, - }, - { - desc: "password reset with forgotten password", - data: fmt.Sprintf(`{"token": "%s", "password": "%s", "confirm_password": "%s"}`, validToken, strongPass, strongPass), - token: validToken, - authnRes: smqauthn.Session{Type: smqauthn.AccessToken, Verified: false}, - contentType: contentType, - status: http.StatusCreated, - err: nil, - }, - { - desc: "password reset with invalid token", - data: fmt.Sprintf(`{"token": "%s", "password": "%s", "confirm_password": "%s"}`, inValidToken, strongPass, strongPass), - token: inValidToken, - contentType: contentType, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "password reset to weak password", - data: fmt.Sprintf(`{"token": "%s", "password": "%s", "confirm_password": "%s"}`, validToken, "weak", "weak"), - token: validToken, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrPasswordFormat, - }, - { - desc: "password reset with empty token", - data: fmt.Sprintf(`{"token": "%s", "password": "%s", "confirm_password": "%s"}`, "", strongPass, strongPass), - token: "", - contentType: contentType, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: apiutil.ErrBearerToken, - }, - { - desc: "password reset with empty password", - data: fmt.Sprintf(`{"token": "%s", "password": "%s", "confirm_password": "%s"}`, validToken, "", ""), - token: validToken, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrValidation, - }, - { - desc: "password reset with malformed data", - data: fmt.Sprintf(`{"token": "%s", "password": %s, "confirm_password": %s}`, validToken, strongPass, strongPass), - token: validToken, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrValidation, - }, - { - desc: "password reset with invalid contentype", - data: fmt.Sprintf(`{"token": "%s", "password": "%s", "confirm_password": "%s"}`, validToken, strongPass, strongPass), - token: validToken, - status: http.StatusUnsupportedMediaType, - err: apiutil.ErrValidation, - }, - { - desc: "password reset with service error", - data: fmt.Sprintf(`{"token": "%s", "password": "%s", "confirm_password": "%s"}`, validToken, strongPass, strongPass), - token: validToken, - contentType: contentType, - status: http.StatusUnprocessableEntity, - svcErr: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - user: us.Client(), - method: http.MethodPut, - url: fmt.Sprintf("%s/password/reset", us.URL), - contentType: tc.contentType, - referer: testReferer, - token: tc.token, - body: strings.NewReader(tc.data), - } - authnCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.authnRes, tc.authnErr) - svcCall := svc.On("ResetSecret", mock.Anything, tc.authnRes, mock.Anything).Return(tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authnCall.Unset() - }) - } -} - -func TestUpdateRole(t *testing.T) { - us, svc, authn := newUsersServer() - defer us.Close() - - cases := []struct { - desc string - data string - userID string - token string - contentType string - authnRes smqauthn.Session - authnErr error - status int - svcErr error - err error - }{ - { - desc: "update user role as admin with valid token", - data: fmt.Sprintf(`{"role": "%s"}`, "admin"), - userID: user.ID, - token: validToken, - authnRes: verifiedSession, - contentType: contentType, - status: http.StatusOK, - err: nil, - }, - { - desc: "update user role as normal user with valid token", - data: fmt.Sprintf(`{"role": "%s"}`, "admin"), - userID: user.ID, - token: validToken, - authnRes: verifiedSession, - contentType: contentType, - status: http.StatusOK, - err: nil, - }, - { - desc: "update user role with invalid token", - data: fmt.Sprintf(`{"role": "%s"}`, "admin"), - userID: user.ID, - token: inValidToken, - contentType: contentType, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "update user role with empty token", - data: fmt.Sprintf(`{"role": "%s"}`, "admin"), - userID: user.ID, - token: "", - contentType: contentType, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: apiutil.ErrBearerToken, - }, - { - desc: "update user with invalid role", - data: fmt.Sprintf(`{"role": "%s"}`, "invalid"), - userID: user.ID, - token: validToken, - authnRes: verifiedSession, - contentType: contentType, - status: http.StatusBadRequest, - err: svcerr.ErrInvalidRole, - }, - { - desc: "update user with invalid contentype", - data: fmt.Sprintf(`{"role": "%s"}`, "admin"), - userID: user.ID, - token: validToken, - authnRes: verifiedSession, - contentType: "application/xml", - status: http.StatusUnsupportedMediaType, - err: apiutil.ErrUnsupportedContentType, - }, - { - desc: "update user with malformed data", - data: fmt.Sprintf(`{"role": %s}`, "admin"), - userID: user.ID, - token: validToken, - authnRes: verifiedSession, - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrMalformedRequestBody, - }, - { - desc: "update user with service error", - data: fmt.Sprintf(`{"role": "%s"}`, "admin"), - userID: user.ID, - token: validToken, - authnRes: verifiedSession, - contentType: contentType, - status: http.StatusUnprocessableEntity, - svcErr: svcerr.ErrUpdateEntity, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - user: us.Client(), - method: http.MethodPatch, - url: fmt.Sprintf("%s/users/%s/role", us.URL, tc.userID), - contentType: tc.contentType, - token: tc.token, - body: strings.NewReader(tc.data), - } - - authnCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.authnRes, tc.authnErr) - svcCall := svc.On("UpdateRole", mock.Anything, tc.authnRes, mock.Anything).Return(users.User{}, tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - var resBody respBody - err = json.NewDecoder(res.Body).Decode(&resBody) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if resBody.Err != "" || resBody.Message != "" { - err = errors.Wrap(errors.New(resBody.Err), errors.New(resBody.Message)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authnCall.Unset() - }) - } -} - -func TestUpdateSecret(t *testing.T) { - us, svc, authn := newUsersServer() - defer us.Close() - - cases := []struct { - desc string - data string - user users.User - contentType string - token string - status int - authnRes smqauthn.Session - authnErr error - err error - }{ - { - desc: "update user secret with valid token", - data: `{"old_secret": "strongersecret", "new_secret": "strongersecret"}`, - user: users.User{ - ID: user.ID, - Email: "username", - Credentials: users.Credentials{ - Secret: "strongersecret", - }, - }, - contentType: contentType, - token: validToken, - authnRes: verifiedSession, - status: http.StatusOK, - err: nil, - }, - { - desc: "update user secret with empty token", - data: `{"old_secret": "strongersecret", "new_secret": "strongersecret"}`, - user: users.User{ - ID: user.ID, - Email: "username", - Credentials: users.Credentials{ - Secret: "strongersecret", - }, - }, - token: "", - authnRes: verifiedSession, - contentType: contentType, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: apiutil.ErrBearerToken, - }, - { - desc: "update user secret with invalid token", - data: `{"old_secret": "strongersecret", "new_secret": "strongersecret"}`, - user: users.User{ - ID: user.ID, - Email: "username", - Credentials: users.Credentials{ - Secret: "strongersecret", - }, - }, - contentType: contentType, - token: inValid, - authnRes: verifiedSession, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - - { - desc: "update user secret with empty secret", - data: `{"old_secret": "", "new_secret": "strongersecret"}`, - user: users.User{ - ID: user.ID, - Email: "username", - Credentials: users.Credentials{ - Secret: "", - }, - }, - contentType: contentType, - token: validToken, - authnRes: verifiedSession, - status: http.StatusBadRequest, - err: apiutil.ErrMissingPass, - }, - { - desc: "update user secret with invalid contentype", - data: `{"old_secret": "strongersecret", "new_secret": "strongersecret"}`, - user: users.User{ - ID: user.ID, - Email: "username", - Credentials: users.Credentials{ - Secret: "", - }, - }, - contentType: "application/xml", - token: validToken, - authnRes: verifiedSession, - status: http.StatusUnsupportedMediaType, - err: apiutil.ErrUnsupportedContentType, - }, - { - desc: "update user secret with malformed data", - data: fmt.Sprintf(`{"secret": %s}`, "invalid"), - user: users.User{ - ID: user.ID, - Email: "username", - Credentials: users.Credentials{ - Secret: "", - }, - }, - contentType: contentType, - token: validToken, - authnRes: verifiedSession, - status: http.StatusBadRequest, - err: apiutil.ErrMalformedRequestBody, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - user: us.Client(), - method: http.MethodPatch, - url: fmt.Sprintf("%s/users/secret", us.URL), - contentType: tc.contentType, - token: tc.token, - body: strings.NewReader(tc.data), - } - - authnCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.authnRes, tc.authnErr) - svcCall := svc.On("UpdateSecret", mock.Anything, tc.authnRes, mock.Anything, mock.Anything).Return(tc.user, tc.err) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - var resBody respBody - err = json.NewDecoder(res.Body).Decode(&resBody) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if resBody.Err != "" || resBody.Message != "" { - err = errors.Wrap(errors.New(resBody.Err), errors.New(resBody.Message)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authnCall.Unset() - }) - } -} - -func TestIssueToken(t *testing.T) { - us, svc, _ := newUsersServer() - defer us.Close() - - validUsername := "valid" - validDescription := "test token" - dataFormat := `{"username": "%s", "password": "%s"}` - dataFormatWithDesc := `{"username": "%s", "password": "%s", "description": "%s"}` - - cases := []struct { - desc string - data string - contentType string - status int - err error - }{ - { - desc: "issue token with valid identity and secret", - data: fmt.Sprintf(dataFormat, validUsername, secret), - contentType: contentType, - status: http.StatusCreated, - err: nil, - }, - { - desc: "issue token with valid identity, secret and description", - data: fmt.Sprintf(dataFormatWithDesc, validUsername, secret, validDescription), - contentType: contentType, - status: http.StatusCreated, - err: nil, - }, - { - desc: "issue token with empty identity", - data: fmt.Sprintf(dataFormat, "", secret), - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrMissingUsernameEmail, - }, - { - desc: "issue token with empty secret", - data: fmt.Sprintf(dataFormat, validUsername, ""), - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrMissingPass, - }, - { - desc: "issue token with invalid email", - data: fmt.Sprintf(dataFormat, "invalid", secret), - contentType: contentType, - status: http.StatusUnauthorized, - err: svcerr.ErrAuthentication, - }, - { - desc: "issues token with malformed data", - data: fmt.Sprintf(dataFormat, validUsername, secret), - contentType: contentType, - status: http.StatusBadRequest, - err: apiutil.ErrMalformedRequestBody, - }, - { - desc: "issue token with invalid contentype", - data: fmt.Sprintf(dataFormat, "invalid", secret), - contentType: "application/xml", - status: http.StatusUnsupportedMediaType, - err: apiutil.ErrUnsupportedContentType, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - user: us.Client(), - method: http.MethodPost, - url: fmt.Sprintf("%s/users/tokens/issue", us.URL), - contentType: tc.contentType, - body: strings.NewReader(tc.data), - } - - svcCall := svc.On("IssueToken", mock.Anything, mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(&grpcTokenV1.Token{AccessToken: validToken}, tc.err) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - if tc.err != nil { - var resBody respBody - err = json.NewDecoder(res.Body).Decode(&resBody) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if resBody.Err != "" || resBody.Message != "" { - err = errors.Wrap(errors.New(resBody.Err), errors.New(resBody.Message)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - } - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - }) - } -} - -func TestRefreshToken(t *testing.T) { - us, svc, authn := newUsersServer() - defer us.Close() - - cases := []struct { - desc string - data string - contentType string - token string - authnRes smqauthn.Session - authnErr error - status int - refreshErr error - err error - }{ - { - desc: "refresh token with valid token", - data: fmt.Sprintf(`{"refresh_token": "%s", "domain_id": "%s"}`, validToken, validID), - contentType: contentType, - token: validToken, - authnRes: verifiedSession, - status: http.StatusCreated, - err: nil, - }, - { - desc: "refresh token with invalid token", - data: fmt.Sprintf(`{"refresh_token": "%s", "domain_id": "%s"}`, inValidToken, validID), - contentType: contentType, - token: inValidToken, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "refresh token with empty token", - data: fmt.Sprintf(`{"refresh_token": "%s", "domain_id": "%s"}`, "", validID), - contentType: contentType, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: apiutil.ErrBearerToken, - }, - { - desc: "refresh token with invalid domain", - data: fmt.Sprintf(`{"refresh_token": "%s", "domain_id": "%s"}`, validToken, "invalid"), - contentType: contentType, - token: validToken, - authnRes: verifiedSession, - status: http.StatusUnauthorized, - err: svcerr.ErrAuthentication, - }, - { - desc: "refresh token with malformed data", - data: fmt.Sprintf(`{"refresh_token": %s, "domain_id": %s}`, validToken, validID), - contentType: contentType, - token: validToken, - authnRes: verifiedSession, - status: http.StatusBadRequest, - err: apiutil.ErrMalformedRequestBody, - }, - { - desc: "refresh token with invalid contentype", - data: fmt.Sprintf(`{"refresh_token": "%s", "domain_id": "%s"}`, validToken, validID), - contentType: "application/xml", - token: validToken, - authnRes: verifiedSession, - status: http.StatusUnsupportedMediaType, - err: apiutil.ErrUnsupportedContentType, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - user: us.Client(), - method: http.MethodPost, - url: fmt.Sprintf("%s/users/tokens/refresh", us.URL), - contentType: tc.contentType, - body: strings.NewReader(tc.data), - token: tc.token, - } - authnCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.authnRes, tc.authnErr) - svcCall := svc.On("RefreshToken", mock.Anything, tc.authnRes, tc.token, mock.Anything).Return(&grpcTokenV1.Token{AccessToken: validToken}, tc.err) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - if tc.err != nil { - var resBody respBody - err = json.NewDecoder(res.Body).Decode(&resBody) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if resBody.Err != "" || resBody.Message != "" { - err = errors.Wrap(errors.New(resBody.Err), errors.New(resBody.Message)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - } - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authnCall.Unset() - }) - } -} - -func TestRevokeRefreshToken(t *testing.T) { - us, svc, authn := newUsersServer() - defer us.Close() - - cases := []struct { - desc string - data string - contentType string - token string - authnRes smqauthn.Session - authnErr error - status int - svcErr error - err error - }{ - { - desc: "revoke refresh token with valid token", - data: fmt.Sprintf(`{"token_id": "%s"}`, validToken), - contentType: contentType, - token: validToken, - authnRes: verifiedSession, - status: http.StatusNoContent, - err: nil, - }, - { - desc: "revoke refresh token with invalid token", - data: fmt.Sprintf(`{"token_id": "%s"}`, validToken), - contentType: contentType, - token: inValidToken, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "revoke refresh token with empty token", - data: fmt.Sprintf(`{"token_id": "%s"}`, validToken), - contentType: contentType, - token: "", - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: apiutil.ErrBearerToken, - }, - { - desc: "revoke refresh token with empty token id", - data: `{"token_id": ""}`, - contentType: contentType, - token: validToken, - authnRes: verifiedSession, - status: http.StatusBadRequest, - err: apiutil.ErrMissingID, - }, - { - desc: "revoke refresh token with malformed data", - data: fmt.Sprintf(`{"token_id": %s}`, validToken), - contentType: contentType, - token: validToken, - authnRes: verifiedSession, - status: http.StatusBadRequest, - err: apiutil.ErrMalformedRequestBody, - }, - { - desc: "revoke refresh token with invalid content type", - data: fmt.Sprintf(`{"token_id": "%s"}`, validToken), - contentType: "application/xml", - token: validToken, - authnRes: verifiedSession, - status: http.StatusUnsupportedMediaType, - err: apiutil.ErrUnsupportedContentType, - }, - { - desc: "revoke refresh token with service error", - data: fmt.Sprintf(`{"token_id": "%s"}`, validToken), - contentType: contentType, - token: validToken, - authnRes: verifiedSession, - status: http.StatusUnprocessableEntity, - svcErr: svcerr.ErrViewEntity, - err: svcerr.ErrViewEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - user: us.Client(), - method: http.MethodPost, - url: fmt.Sprintf("%s/users/tokens/revoke", us.URL), - contentType: tc.contentType, - body: strings.NewReader(tc.data), - token: tc.token, - } - authnCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.authnRes, tc.authnErr) - svcCall := svc.On("RevokeRefreshToken", mock.Anything, tc.authnRes, mock.Anything).Return(tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - if tc.err != nil { - var resBody respBody - err = json.NewDecoder(res.Body).Decode(&resBody) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if resBody.Err != "" || resBody.Message != "" { - err = errors.Wrap(errors.New(resBody.Err), errors.New(resBody.Message)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - } - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authnCall.Unset() - }) - } -} - -func TestListActiveRefreshTokens(t *testing.T) { - us, svc, authn := newUsersServer() - defer us.Close() - - cases := []struct { - desc string - token string - authnRes smqauthn.Session - authnErr error - status int - svcRes *grpcTokenV1.ListUserRefreshTokensRes - svcErr error - err error - }{ - { - desc: "list active refresh tokens with valid token", - token: validToken, - authnRes: verifiedSession, - status: http.StatusOK, - svcRes: &grpcTokenV1.ListUserRefreshTokensRes{ - RefreshTokens: []*grpcTokenV1.RefreshToken{ - {Id: "token1", Description: "token-1"}, - {Id: "token2", Description: "token-2"}, - }, - }, - err: nil, - }, - { - desc: "list active refresh tokens with invalid token", - token: inValidToken, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "list active refresh tokens with empty token", - token: "", - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: apiutil.ErrBearerToken, - }, - { - desc: "list active refresh tokens with service error", - token: validToken, - authnRes: verifiedSession, - status: http.StatusUnprocessableEntity, - svcErr: svcerr.ErrViewEntity, - err: svcerr.ErrViewEntity, - }, - { - desc: "list active refresh tokens with empty list", - token: validToken, - authnRes: verifiedSession, - status: http.StatusOK, - svcRes: &grpcTokenV1.ListUserRefreshTokensRes{ - RefreshTokens: []*grpcTokenV1.RefreshToken{}, - }, - err: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - req := testRequest{ - user: us.Client(), - method: http.MethodGet, - url: fmt.Sprintf("%s/users/tokens/refresh-tokens", us.URL), - token: tc.token, - } - authnCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.authnRes, tc.authnErr) - svcCall := svc.On("ListActiveRefreshTokens", mock.Anything, tc.authnRes).Return(tc.svcRes, tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - if tc.err != nil { - var resBody respBody - err = json.NewDecoder(res.Body).Decode(&resBody) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if resBody.Err != "" || resBody.Message != "" { - err = errors.Wrap(errors.New(resBody.Err), errors.New(resBody.Message)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - } - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authnCall.Unset() - }) - } -} - -func TestEnable(t *testing.T) { - us, svc, authn := newUsersServer() - defer us.Close() - cases := []struct { - desc string - user users.User - response users.User - token string - authnRes smqauthn.Session - authnErr error - status int - svcErr error - err error - }{ - { - desc: "enable user as admin with valid token", - user: user, - response: users.User{ - ID: user.ID, - Status: users.EnabledStatus, - }, - token: validToken, - authnRes: verifiedSession, - status: http.StatusOK, - err: nil, - }, - { - desc: "enable user as normal user with valid token", - user: user, - response: users.User{ - ID: user.ID, - Status: users.EnabledStatus, - }, - token: validToken, - authnRes: verifiedSession, - status: http.StatusOK, - err: nil, - }, - { - desc: "enable user with invalid token", - user: user, - token: inValidToken, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "enable user with empty id", - user: users.User{ - ID: "", - }, - token: validToken, - authnRes: verifiedSession, - status: http.StatusBadRequest, - err: apiutil.ErrMissingID, - }, - { - desc: "enable user with service error", - user: user, - token: validToken, - authnRes: verifiedSession, - status: http.StatusUnprocessableEntity, - svcErr: svcerr.ErrEnableUser, - err: svcerr.ErrEnableUser, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - data := toJSON(tc.user) - req := testRequest{ - user: us.Client(), - method: http.MethodPost, - url: fmt.Sprintf("%s/users/%s/enable", us.URL, tc.user.ID), - contentType: contentType, - token: tc.token, - body: strings.NewReader(data), - } - - authnCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.authnRes, tc.authnErr) - svcCall := svc.On("Enable", mock.Anything, tc.authnRes, mock.Anything).Return(tc.user, tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - if tc.err != nil { - var resBody respBody - err = json.NewDecoder(res.Body).Decode(&resBody) - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error while decoding response body: %s", tc.desc, err)) - if resBody.Err != "" || resBody.Message != "" { - err = errors.Wrap(errors.New(resBody.Err), errors.New(resBody.Message)) - } - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - } - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authnCall.Unset() - }) - } -} - -func TestDisable(t *testing.T) { - us, svc, authn := newUsersServer() - defer us.Close() - - cases := []struct { - desc string - user users.User - response users.User - token string - authnRes smqauthn.Session - authnErr error - status int - svcErr error - err error - }{ - { - desc: "disable user as admin with valid token", - user: user, - response: users.User{ - ID: user.ID, - Status: users.DisabledStatus, - }, - token: validToken, - authnRes: smqauthn.Session{UserID: validID, DomainID: domainID, SuperAdmin: true, Verified: true}, - status: http.StatusOK, - err: nil, - }, - { - desc: "disable user as normal user with valid token", - user: user, - response: users.User{ - ID: user.ID, - Status: users.DisabledStatus, - }, - token: validToken, - authnRes: verifiedSession, - status: http.StatusOK, - err: nil, - }, - { - desc: "disable user with invalid token", - user: user, - token: inValidToken, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "disable user with empty id", - user: users.User{ - ID: "", - }, - token: validToken, - authnRes: verifiedSession, - status: http.StatusBadRequest, - err: apiutil.ErrMissingID, - }, - { - desc: "disable user with service error", - user: user, - token: validToken, - authnRes: verifiedSession, - status: http.StatusUnprocessableEntity, - svcErr: svcerr.ErrDisableUser, - err: svcerr.ErrDisableUser, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - data := toJSON(tc.user) - req := testRequest{ - user: us.Client(), - method: http.MethodPost, - url: fmt.Sprintf("%s/users/%s/disable", us.URL, tc.user.ID), - contentType: contentType, - token: tc.token, - body: strings.NewReader(data), - } - - authnCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.authnRes, tc.authnErr) - svcCall := svc.On("Disable", mock.Anything, mock.Anything, mock.Anything).Return(tc.user, tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - svcCall.Unset() - authnCall.Unset() - }) - } -} - -func TestDelete(t *testing.T) { - us, svc, authn := newUsersServer() - defer us.Close() - - cases := []struct { - desc string - user users.User - response users.User - token string - authnRes smqauthn.Session - authnErr error - status int - svcErr error - err error - }{ - { - desc: "delete user as admin with valid token", - user: user, - response: users.User{ - ID: user.ID, - }, - token: validToken, - authnRes: verifiedSession, - status: http.StatusNoContent, - err: nil, - }, - { - desc: "delete user with invalid token", - user: user, - token: inValidToken, - status: http.StatusUnauthorized, - authnErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "delete user with empty id", - user: users.User{ - ID: "", - }, - token: validToken, - authnRes: verifiedSession, - status: http.StatusMethodNotAllowed, - err: apiutil.ErrMissingID, - }, - { - desc: "delete user with service error", - user: user, - token: validToken, - authnRes: verifiedSession, - status: http.StatusUnprocessableEntity, - svcErr: svcerr.ErrRemoveEntity, - err: svcerr.ErrRemoveEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - data := toJSON(tc.user) - req := testRequest{ - user: us.Client(), - method: http.MethodDelete, - url: fmt.Sprintf("%s/users/%s", us.URL, tc.user.ID), - contentType: contentType, - token: tc.token, - body: strings.NewReader(data), - } - authnCall := authn.On("Authenticate", mock.Anything, tc.token).Return(tc.authnRes, tc.authnErr) - repoCall := svc.On("Delete", mock.Anything, tc.authnRes, tc.user.ID).Return(tc.svcErr) - res, err := req.make() - assert.Nil(t, err, fmt.Sprintf("%s: unexpected error %s", tc.desc, err)) - assert.Equal(t, tc.status, res.StatusCode, fmt.Sprintf("%s: expected status code %d got %d", tc.desc, tc.status, res.StatusCode)) - repoCall.Unset() - authnCall.Unset() - }) - } -} - -type respBody struct { - Err string `json:"error"` - Message string `json:"message"` - Total int `json:"total"` - ID string `json:"id"` - Tags []string `json:"tags"` - Role users.Role `json:"role"` - Status users.Status `json:"status"` -} diff --git a/users/api/endpoints.go b/users/api/endpoints.go deleted file mode 100644 index 125752736..000000000 --- a/users/api/endpoints.go +++ /dev/null @@ -1,556 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package api - -import ( - "context" - - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/users" - "github.com/go-kit/kit/endpoint" -) - -func registrationEndpoint(svc users.Service, selfRegister bool) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(createUserReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - session := authn.Session{} - - var ok bool - if !selfRegister { - session, ok = ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - } - - user, err := svc.Register(ctx, session, req.User, selfRegister) - if err != nil { - return nil, err - } - - return createUserRes{ - User: user, - created: true, - }, nil - } -} - -func sendVerificationEndpoint(svc users.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - _ = request.(sendVerificationReq) - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - if err := svc.SendVerification(ctx, session); err != nil { - return sendVerificationRes{}, err - } - - return sendVerificationRes{}, nil - } -} - -func verifyEmailEndpoint(svc users.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(verifyEmailReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - if _, err := svc.VerifyEmail(ctx, req.token); err != nil { - return verifyEmailRes{}, err - } - - return verifyEmailRes{}, nil - } -} - -func viewEndpoint(svc users.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(viewUserReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - user, err := svc.View(ctx, session, req.id) - if err != nil { - return nil, err - } - - return viewUserRes{User: user}, nil - } -} - -func viewProfileEndpoint(svc users.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - client, err := svc.ViewProfile(ctx, session) - if err != nil { - return nil, err - } - - return viewUserRes{User: client}, nil - } -} - -func listUsersEndpoint(svc users.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(listUsersReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - pm := users.Page{ - Status: req.status, - Offset: req.offset, - Limit: req.limit, - OnlyTotal: req.onlyTotal, - Username: req.userName, - Tags: req.tags, - Metadata: req.metadata, - FirstName: req.firstName, - LastName: req.lastName, - Email: req.email, - Order: req.order, - Dir: req.dir, - Id: req.id, - CreatedFrom: req.createdFrom, - CreatedTo: req.createdTo, - } - - page, err := svc.ListUsers(ctx, session, pm) - if err != nil { - return nil, err - } - - res := usersPageRes{ - pageRes: pageRes{ - Total: page.Total, - Offset: page.Offset, - Limit: page.Limit, - }, - Users: []viewUserRes{}, - } - for _, user := range page.Users { - res.Users = append(res.Users, viewUserRes{User: user}) - } - - return res, nil - } -} - -func searchUsersEndpoint(svc users.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(searchUsersReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - pm := users.Page{ - Offset: req.Offset, - Limit: req.Limit, - Username: req.Username, - FirstName: req.FirstName, - LastName: req.LastName, - Id: req.Id, - Order: req.Order, - Dir: req.Dir, - } - page, err := svc.SearchUsers(ctx, pm) - if err != nil { - return nil, err - } - - res := usersPageRes{ - pageRes: pageRes{ - Total: page.Total, - Offset: page.Offset, - Limit: page.Limit, - }, - Users: []viewUserRes{}, - } - for _, user := range page.Users { - res.Users = append(res.Users, viewUserRes{User: user}) - } - - return res, nil - } -} - -func updateEndpoint(svc users.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(updateUserReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - usr := users.UserReq{ - FirstName: req.FirstName, - LastName: req.LastName, - Metadata: req.Metadata, - PrivateMetadata: req.PrivateMetadata, - } - - user, err := svc.Update(ctx, session, req.id, usr) - if err != nil { - return nil, err - } - - return updateUserRes{User: user}, nil - } -} - -func updateTagsEndpoint(svc users.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(updateUserTagsReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - usr := users.UserReq{ - Tags: req.Tags, - } - - user, err := svc.UpdateTags(ctx, session, req.id, usr) - if err != nil { - return nil, err - } - - return updateUserRes{User: user}, nil - } -} - -func updateEmailEndpoint(svc users.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(updateEmailReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - user, err := svc.UpdateEmail(ctx, session, req.id, req.Email) - if err != nil { - return nil, err - } - - return updateUserRes{User: user}, nil - } -} - -// Password reset request endpoint. -// When successful password reset link is generated. -// Link is generated using MG_TOKEN_RESET_ENDPOINT env. -// and value from Referer header for host. -// {Referer}+{MG_TOKEN_RESET_ENDPOINT}+{token=TOKEN} -// http://magistrala.com/reset-request?token=xxxxxxxxxxx. -// Email with a link is being sent to the user. -// When user clicks on a link it should get the ui with form to -// enter new password, when form is submitted token and new password -// must be sent as PUT request to 'password/reset' passwordResetEndpoint. -func passwordResetRequestEndpoint(svc users.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(passResetReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - if err := svc.SendPasswordReset(ctx, req.Email); err != nil { - return nil, err - } - - return passResetReqRes{Msg: MailSent}, nil - } -} - -// This is endpoint that actually sets new password in password reset flow. -// When user clicks on a link in email finally ends on this endpoint as explained in -// the comment above. -func passwordResetEndpoint(svc users.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(resetTokenReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - if err := svc.ResetSecret(ctx, session, req.Password); err != nil { - return nil, err - } - - return passChangeRes{}, nil - } -} - -func updateSecretEndpoint(svc users.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(updateUserSecretReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - user, err := svc.UpdateSecret(ctx, session, req.OldSecret, req.NewSecret) - if err != nil { - return nil, err - } - - return updateUserRes{User: user}, nil - } -} - -func updateUsernameEndpoint(svc users.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(updateUsernameReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthorization - } - - user, err := svc.UpdateUsername(ctx, session, req.id, req.Username) - if err != nil { - return nil, err - } - - return updateUserRes{User: user}, nil - } -} - -func updateProfilePictureEndpoint(svc users.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(updateProfilePictureReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - usr := users.UserReq{ - ProfilePicture: req.ProfilePicture, - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthorization - } - - user, err := svc.UpdateProfilePicture(ctx, session, req.id, usr) - if err != nil { - return nil, err - } - - return updateUserRes{User: user}, nil - } -} - -func updateRoleEndpoint(svc users.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(updateUserRoleReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - user := users.User{ - ID: req.id, - Role: req.role, - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - user, err := svc.UpdateRole(ctx, session, user) - if err != nil { - return nil, err - } - - return updateUserRes{User: user}, nil - } -} - -func issueTokenEndpoint(svc users.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(loginUserReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - token, err := svc.IssueToken(ctx, req.Username, req.Password, req.Description) - if err != nil { - return nil, err - } - - return tokenRes{ - AccessToken: token.GetAccessToken(), - RefreshToken: token.GetRefreshToken(), - AccessType: token.GetAccessType(), - }, nil - } -} - -func refreshTokenEndpoint(svc users.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(tokenReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - token, err := svc.RefreshToken(ctx, session, req.RefreshToken) - if err != nil { - return nil, err - } - - return tokenRes{ - AccessToken: token.GetAccessToken(), - RefreshToken: token.GetRefreshToken(), - AccessType: token.GetAccessType(), - }, nil - } -} - -func revokeRefreshTokenEndpoint(svc users.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(revokeTokenReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - err := svc.RevokeRefreshToken(ctx, session, req.TokenID) - if err != nil { - return nil, err - } - - return revokeRes{}, nil - } -} - -func listActiveRefreshTokensEndpoint(svc users.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - refreshTokens, err := svc.ListActiveRefreshTokens(ctx, session) - if err != nil { - return nil, err - } - - return listRefreshTokensRes{RefreshTokens: refreshTokens.GetRefreshTokens()}, nil - } -} - -func enableEndpoint(svc users.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(changeUserStatusReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - user, err := svc.Enable(ctx, session, req.id) - if err != nil { - return nil, err - } - - return changeUserStatusRes{User: user}, nil - } -} - -func disableEndpoint(svc users.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(changeUserStatusReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - user, err := svc.Disable(ctx, session, req.id) - if err != nil { - return nil, err - } - - return changeUserStatusRes{User: user}, nil - } -} - -func deleteEndpoint(svc users.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(changeUserStatusReq) - if err := req.validate(); err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - session, ok := ctx.Value(authn.SessionKey).(authn.Session) - if !ok { - return nil, svcerr.ErrAuthentication - } - - if err := svc.Delete(ctx, session, req.id); err != nil { - return nil, err - } - - return deleteUserRes{true}, nil - } -} diff --git a/users/api/grpc/client.go b/users/api/grpc/client.go deleted file mode 100644 index 759430575..000000000 --- a/users/api/grpc/client.go +++ /dev/null @@ -1,148 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package grpc - -import ( - "context" - "time" - - grpcUsersV1 "github.com/absmach/magistrala/api/grpc/users/v1" - grpcapi "github.com/absmach/magistrala/auth/api/grpc" - "github.com/absmach/magistrala/users" - "github.com/go-kit/kit/endpoint" - kitgrpc "github.com/go-kit/kit/transport/grpc" - "google.golang.org/grpc" -) - -const usersSvcName = "users.v1.UsersService" - -var _ grpcUsersV1.UsersServiceClient = (*usersGrpcClient)(nil) - -type usersGrpcClient struct { - retrieveUsers endpoint.Endpoint - timeout time.Duration -} - -// NewClient returns new users gRPC client instance. -func NewClient(conn *grpc.ClientConn, timeout time.Duration) grpcUsersV1.UsersServiceClient { - return &usersGrpcClient{ - retrieveUsers: kitgrpc.NewClient( - conn, - usersSvcName, - "RetrieveUsers", - encodeRetrieveUsersRequest, - decodeRetrieveUsersResponse, - grpcUsersV1.RetrieveUsersRes{}, - ).Endpoint(), - timeout: timeout, - } -} - -func (client usersGrpcClient) RetrieveUsers(ctx context.Context, in *grpcUsersV1.RetrieveUsersReq, opts ...grpc.CallOption) (*grpcUsersV1.RetrieveUsersRes, error) { - ctx, cancel := context.WithTimeout(ctx, client.timeout) - defer cancel() - - res, err := client.retrieveUsers(ctx, retrieveUsersReq{ - ids: in.GetIds(), - offset: in.GetOffset(), - limit: in.GetLimit(), - }) - if err != nil { - return &grpcUsersV1.RetrieveUsersRes{}, grpcapi.DecodeError(err) - } - - rur := res.(retrieveUsersRes) - - usersPB, err := toProtoUsers(rur.users) - if err != nil { - return &grpcUsersV1.RetrieveUsersRes{}, err - } - - return &grpcUsersV1.RetrieveUsersRes{ - Total: rur.total, - Limit: rur.limit, - Offset: rur.offset, - Users: usersPB, - }, nil -} - -func decodeRetrieveUsersResponse(_ context.Context, grpcRes any) (any, error) { - res := grpcRes.(*grpcUsersV1.RetrieveUsersRes) - - usersDomain, err := usersFromProto(res.GetUsers()) - if err != nil { - return nil, err - } - - return retrieveUsersRes{ - users: usersDomain, - total: res.GetTotal(), - limit: res.GetLimit(), - offset: res.GetOffset(), - }, nil -} - -func encodeRetrieveUsersRequest(_ context.Context, grpcReq any) (any, error) { - req := grpcReq.(retrieveUsersReq) - return &grpcUsersV1.RetrieveUsersReq{ - Ids: req.ids, - Offset: req.offset, - Limit: req.limit, - }, nil -} - -func usersFromProto(us []*grpcUsersV1.User) ([]users.User, error) { - var res []users.User - for _, u := range us { - du, err := userFromProto(u) - if err != nil { - return nil, err - } - res = append(res, du) - } - - return res, nil -} - -func userFromProto(u *grpcUsersV1.User) (users.User, error) { - metadata := users.Metadata(nil) - if u.GetMetadata() != nil { - metadata = users.Metadata(u.GetMetadata().AsMap()) - } - privateMetadata := users.Metadata(nil) - if u.GetPrivateMetadata() != nil { - privateMetadata = users.Metadata(u.GetPrivateMetadata().AsMap()) - } - - user := users.User{ - ID: u.GetId(), - FirstName: u.GetFirstName(), - LastName: u.GetLastName(), - Tags: u.GetTags(), - Metadata: metadata, - PrivateMetadata: privateMetadata, - Status: users.Status(u.GetStatus()), - Role: users.Role(u.GetRole()), - ProfilePicture: u.GetProfilePicture(), - Credentials: users.Credentials{ - Username: u.GetUsername(), - }, - Email: u.GetEmail(), - UpdatedBy: u.GetUpdatedBy(), - Permissions: u.GetPermissions(), - AuthProvider: u.GetAuthProvider(), - } - - if u.GetCreatedAt() != nil { - user.CreatedAt = u.GetCreatedAt().AsTime() - } - if u.GetUpdatedAt() != nil { - user.UpdatedAt = u.GetUpdatedAt().AsTime() - } - if u.GetVerifiedAt() != nil { - user.VerifiedAt = u.GetVerifiedAt().AsTime() - } - - return user, nil -} diff --git a/users/api/grpc/doc.go b/users/api/grpc/doc.go deleted file mode 100644 index ce2c0fe9d..000000000 --- a/users/api/grpc/doc.go +++ /dev/null @@ -1,5 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package grpc contains implementation of Users service gRPC API. -package grpc diff --git a/users/api/grpc/endpoint.go b/users/api/grpc/endpoint.go deleted file mode 100644 index 3e32e70b0..000000000 --- a/users/api/grpc/endpoint.go +++ /dev/null @@ -1,33 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package grpc - -import ( - "context" - - pusers "github.com/absmach/magistrala/users/private" - "github.com/go-kit/kit/endpoint" -) - -func retrieveUsersEndpoint(svc pusers.Service) endpoint.Endpoint { - return func(ctx context.Context, request any) (any, error) { - req := request.(retrieveUsersReq) - - if err := req.validate(); err != nil { - return retrieveUsersRes{}, err - } - - page, err := svc.RetrieveByIDs(ctx, req.ids, req.offset, req.limit) - if err != nil { - return retrieveUsersRes{}, err - } - - return retrieveUsersRes{ - users: page.Users, - total: page.Total, - limit: page.Limit, - offset: page.Offset, - }, nil - } -} diff --git a/users/api/grpc/requests.go b/users/api/grpc/requests.go deleted file mode 100644 index 2840d9b89..000000000 --- a/users/api/grpc/requests.go +++ /dev/null @@ -1,22 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package grpc - -import ( - apiutil "github.com/absmach/magistrala/api/http/util" -) - -type retrieveUsersReq struct { - ids []string - offset uint64 - limit uint64 -} - -func (req retrieveUsersReq) validate() error { - if len(req.ids) == 0 { - return apiutil.ErrMissingUserID - } - - return nil -} diff --git a/users/api/grpc/responses.go b/users/api/grpc/responses.go deleted file mode 100644 index e754da1e4..000000000 --- a/users/api/grpc/responses.go +++ /dev/null @@ -1,13 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package grpc - -import "github.com/absmach/magistrala/users" - -type retrieveUsersRes struct { - users []users.User - total uint64 - limit uint64 - offset uint64 -} diff --git a/users/api/grpc/server.go b/users/api/grpc/server.go deleted file mode 100644 index 0c2cba330..000000000 --- a/users/api/grpc/server.go +++ /dev/null @@ -1,128 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package grpc - -import ( - "context" - - grpcUsersV1 "github.com/absmach/magistrala/api/grpc/users/v1" - grpcapi "github.com/absmach/magistrala/auth/api/grpc" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/users" - pusers "github.com/absmach/magistrala/users/private" - kitgrpc "github.com/go-kit/kit/transport/grpc" - structpb "google.golang.org/protobuf/types/known/structpb" - timestamppb "google.golang.org/protobuf/types/known/timestamppb" -) - -var _ grpcUsersV1.UsersServiceServer = (*usersGrpcServer)(nil) - -type usersGrpcServer struct { - grpcUsersV1.UnimplementedUsersServiceServer - retrieveUsers kitgrpc.Handler -} - -func NewServer(svc pusers.Service) grpcUsersV1.UsersServiceServer { - return &usersGrpcServer{ - retrieveUsers: kitgrpc.NewServer( - retrieveUsersEndpoint(svc), - decodeRetrieveUsersRequest, - encodeRetrieveUsersResponse, - ), - } -} - -func decodeRetrieveUsersRequest(_ context.Context, grpcReq any) (any, error) { - req := grpcReq.(*grpcUsersV1.RetrieveUsersReq) - return retrieveUsersReq{ - ids: req.GetIds(), - offset: req.GetOffset(), - limit: req.GetLimit(), - }, nil -} - -func encodeRetrieveUsersResponse(_ context.Context, grpcRes any) (any, error) { - res := grpcRes.(retrieveUsersRes) - - usersPB, err := toProtoUsers(res.users) - if err != nil { - return nil, err - } - - return &grpcUsersV1.RetrieveUsersRes{ - Total: res.total, - Limit: res.limit, - Offset: res.offset, - Users: usersPB, - }, nil -} - -func (s *usersGrpcServer) RetrieveUsers(ctx context.Context, req *grpcUsersV1.RetrieveUsersReq) (*grpcUsersV1.RetrieveUsersRes, error) { - _, res, err := s.retrieveUsers.ServeGRPC(ctx, req) - if err != nil { - return nil, grpcapi.EncodeError(err) - } - - return res.(*grpcUsersV1.RetrieveUsersRes), nil -} - -func toProtoUsers(us []users.User) ([]*grpcUsersV1.User, error) { - var res []*grpcUsersV1.User - for _, u := range us { - pu, err := toProtoUser(u) - if err != nil { - return nil, err - } - res = append(res, pu) - } - - return res, nil -} - -func toProtoUser(u users.User) (*grpcUsersV1.User, error) { - var metadata, privateMetadata *structpb.Struct - var err error - if u.Metadata != nil { - metadata, err = structpb.NewStruct(u.Metadata) - if err != nil { - return nil, errors.Wrap(svcerr.ErrViewEntity, err) - } - } - if u.PrivateMetadata != nil { - privateMetadata, err = structpb.NewStruct(u.PrivateMetadata) - if err != nil { - return nil, errors.Wrap(svcerr.ErrViewEntity, err) - } - } - - pu := &grpcUsersV1.User{ - Id: u.ID, - FirstName: u.FirstName, - LastName: u.LastName, - Tags: u.Tags, - Metadata: metadata, - PrivateMetadata: privateMetadata, - Status: uint32(u.Status), - Role: uint32(u.Role), - ProfilePicture: u.ProfilePicture, - Username: u.Credentials.Username, - Email: u.Email, - UpdatedBy: u.UpdatedBy, - AuthProvider: u.AuthProvider, - Permissions: u.Permissions, - } - - if !u.CreatedAt.IsZero() { - pu.CreatedAt = timestamppb.New(u.CreatedAt) - } - if !u.UpdatedAt.IsZero() { - pu.UpdatedAt = timestamppb.New(u.UpdatedAt) - } - if !u.VerifiedAt.IsZero() { - pu.VerifiedAt = timestamppb.New(u.VerifiedAt) - } - - return pu, nil -} diff --git a/users/api/requests.go b/users/api/requests.go deleted file mode 100644 index f90af3db9..000000000 --- a/users/api/requests.go +++ /dev/null @@ -1,348 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package api - -import ( - "net/url" - "time" - - api "github.com/absmach/magistrala/api/http" - apiutil "github.com/absmach/magistrala/api/http/util" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/users" -) - -const maxLimitSize = 100 - -type createUserReq struct { - users.User -} - -func (req createUserReq) validate() error { - if len(req.User.FirstName) > api.MaxNameSize { - return apiutil.ErrNameSize - } - if len(req.User.LastName) > api.MaxNameSize { - return apiutil.ErrNameSize - } - if req.User.FirstName == "" { - return apiutil.ErrMissingFirstName - } - if req.User.LastName == "" { - return apiutil.ErrMissingLastName - } - if req.User.Credentials.Username == "" { - return apiutil.ErrMissingUsername - } - if err := api.ValidateUserName(req.User.Credentials.Username); err != nil { - return err - } - // Username must not be a valid email format due to username/email login. - if err := api.ValidateEmail(req.User.Credentials.Username); err == nil { - return apiutil.ErrInvalidUsername - } - - if req.User.Email == "" { - return apiutil.ErrMissingEmail - } - // Email must be in a valid format. - if err := api.ValidateEmail(req.User.Email); err != nil { - return err - } - if req.User.Credentials.Secret == "" { - return apiutil.ErrMissingPass - } - if !passRegex.MatchString(req.User.Credentials.Secret) { - return apiutil.ErrPasswordFormat - } - if req.User.Status == users.AllStatus { - return svcerr.ErrInvalidStatus - } - if req.User.ProfilePicture != "" { - if _, err := url.Parse(req.User.ProfilePicture); err != nil { - return apiutil.ErrInvalidProfilePictureURL - } - } - - return req.User.Validate() -} - -type sendVerificationReq struct{} - -type verifyEmailReq struct { - token string -} - -func (req verifyEmailReq) validate() error { - if req.token == "" { - return apiutil.ErrInvalidVerification - } - - return nil -} - -type viewUserReq struct { - id string -} - -func (req viewUserReq) validate() error { - if req.id == "" { - return apiutil.ErrMissingID - } - - return nil -} - -type listUsersReq struct { - status users.Status - offset uint64 - limit uint64 - onlyTotal bool - userName string - tags users.TagsQuery - firstName string - lastName string - email string - metadata users.Metadata - order string - dir string - id string - createdFrom time.Time - createdTo time.Time -} - -func (req listUsersReq) validate() error { - if req.limit > maxLimitSize || req.limit < 1 { - return apiutil.ErrLimitSize - } - - switch req.order { - case "", api.CreatedAtOrder, api.UpdatedAtOrder, api.FirstNameKey, api.LastNameKey, api.UsernameKey, api.EmailKey: - default: - return apiutil.ErrInvalidOrder - } - - if req.dir != "" && (req.dir != api.AscDir && req.dir != api.DescDir) { - return apiutil.ErrInvalidDirection - } - - return nil -} - -type searchUsersReq struct { - Offset uint64 - Limit uint64 - Username string - FirstName string - LastName string - Id string - Order string - Dir string -} - -func (req searchUsersReq) validate() error { - if req.Username == "" && req.Id == "" && req.FirstName == "" && req.LastName == "" { - return apiutil.ErrEmptySearchQuery - } - - return nil -} - -type updateUserReq struct { - id string - FirstName *string `json:"first_name,omitempty"` - LastName *string `json:"last_name,omitempty"` - Metadata *users.Metadata `json:"metadata,omitempty"` - PrivateMetadata *users.Metadata `json:"private_metadata,omitempty"` -} - -func (req updateUserReq) validate() error { - if req.id == "" { - return apiutil.ErrMissingID - } - - return nil -} - -type updateUserTagsReq struct { - id string - Tags *[]string `json:"tags,omitempty"` -} - -func (req updateUserTagsReq) validate() error { - if req.id == "" { - return apiutil.ErrMissingID - } - - return nil -} - -type updateUserRoleReq struct { - id string - role users.Role - Role string `json:"role,omitempty"` -} - -func (req updateUserRoleReq) validate() error { - if req.id == "" { - return apiutil.ErrMissingID - } - - return nil -} - -type updateEmailReq struct { - id string - Email string `json:"email,omitempty"` -} - -func (req updateEmailReq) validate() error { - if req.id == "" { - return apiutil.ErrMissingID - } - if err := api.ValidateEmail(req.Email); err != nil { - return err - } - - return nil -} - -type updateUserSecretReq struct { - OldSecret string `json:"old_secret,omitempty"` - NewSecret string `json:"new_secret,omitempty"` -} - -func (req updateUserSecretReq) validate() error { - if req.OldSecret == "" || req.NewSecret == "" { - return apiutil.ErrMissingPass - } - if !passRegex.MatchString(req.NewSecret) { - return apiutil.ErrPasswordFormat - } - - return nil -} - -type updateUsernameReq struct { - id string - Username string `json:"username,omitempty"` -} - -func (req updateUsernameReq) validate() error { - if req.id == "" { - return apiutil.ErrMissingID - } - if len(req.Username) > api.MaxNameSize { - return apiutil.ErrNameSize - } - if req.Username == "" { - return apiutil.ErrMissingUsername - } - - return nil -} - -type updateProfilePictureReq struct { - id string - ProfilePicture *string `json:"profile_picture,omitempty"` -} - -func (req updateProfilePictureReq) validate() error { - if req.id == "" { - return apiutil.ErrMissingID - } - if req.ProfilePicture != nil { - if _, err := url.Parse(*req.ProfilePicture); err != nil { - return apiutil.ErrInvalidProfilePictureURL - } - } - return nil -} - -type changeUserStatusReq struct { - id string -} - -func (req changeUserStatusReq) validate() error { - if req.id == "" { - return apiutil.ErrMissingID - } - - return nil -} - -type loginUserReq struct { - Username string `json:"username,omitempty"` - Password string `json:"password,omitempty"` - Description string `json:"description,omitempty"` -} - -func (req loginUserReq) validate() error { - if req.Username == "" { - return apiutil.ErrMissingUsernameEmail - } - if req.Password == "" { - return apiutil.ErrMissingPass - } - - return nil -} - -type tokenReq struct { - RefreshToken string `json:"refresh_token,omitempty"` -} - -func (req tokenReq) validate() error { - if req.RefreshToken == "" { - return apiutil.ErrBearerToken - } - - return nil -} - -type revokeTokenReq struct { - TokenID string `json:"token_id,omitempty"` -} - -func (req revokeTokenReq) validate() error { - if req.TokenID == "" { - return apiutil.ErrMissingID - } - - return nil -} - -type passResetReq struct { - Email string `json:"email"` -} - -func (req passResetReq) validate() error { - if req.Email == "" { - return apiutil.ErrMissingEmail - } - - return nil -} - -type resetTokenReq struct { - Password string `json:"password"` - ConfPass string `json:"confirm_password"` -} - -func (req resetTokenReq) validate() error { - if req.Password == "" { - return apiutil.ErrMissingPass - } - if req.ConfPass == "" { - return apiutil.ErrMissingConfPass - } - if req.Password != req.ConfPass { - return apiutil.ErrInvalidResetPass - } - if !passRegex.MatchString(req.ConfPass) { - return apiutil.ErrPasswordFormat - } - - return nil -} diff --git a/users/api/requests_test.go b/users/api/requests_test.go deleted file mode 100644 index 26ead54ad..000000000 --- a/users/api/requests_test.go +++ /dev/null @@ -1,649 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package api - -import ( - "net/url" - "strings" - "testing" - - api "github.com/absmach/magistrala/api/http" - apiutil "github.com/absmach/magistrala/api/http/util" - "github.com/absmach/magistrala/internal/testsutil" - "github.com/absmach/magistrala/users" - "github.com/stretchr/testify/assert" -) - -const ( - valid = "valid" - secret = "QJg58*aMan7j" - name = "user" - validEmail = "example@domain.com" -) - -var validID = testsutil.GenerateUUID(&testing.T{}) - -func TestCreateUserReqValidate(t *testing.T) { - cases := []struct { - desc string - req createUserReq - err error - }{ - { - desc: "valid request", - req: createUserReq{ - User: users.User{ - ID: validID, - FirstName: valid, - LastName: valid, - Email: validEmail, - Credentials: users.Credentials{ - Username: valid, - Secret: secret, - }, - }, - }, - err: nil, - }, - { - desc: "name too long", - req: createUserReq{ - User: users.User{ - ID: validID, - FirstName: strings.Repeat("a", api.MaxNameSize+1), - LastName: valid, - }, - }, - err: apiutil.ErrNameSize, - }, - { - desc: "missing email in request", - req: createUserReq{ - User: users.User{ - ID: validID, - FirstName: valid, - LastName: valid, - Credentials: users.Credentials{ - Username: valid, - Secret: secret, - }, - }, - }, - err: apiutil.ErrMissingEmail, - }, - { - desc: "missing secret in request", - req: createUserReq{ - User: users.User{ - ID: validID, - FirstName: valid, - LastName: valid, - Email: validEmail, - Credentials: users.Credentials{ - Username: valid, - }, - }, - }, - err: apiutil.ErrMissingPass, - }, - { - desc: "invalid secret in request", - req: createUserReq{ - User: users.User{ - ID: validID, - FirstName: valid, - LastName: valid, - Email: validEmail, - Credentials: users.Credentials{ - Username: valid, - Secret: "invalid", - }, - }, - }, - err: apiutil.ErrPasswordFormat, - }, - { - desc: "missing username in request", - req: createUserReq{ - User: users.User{ - ID: validID, - FirstName: valid, - LastName: valid, - Email: validEmail, - Credentials: users.Credentials{ - Username: "", - Secret: secret, - }, - }, - }, - err: apiutil.ErrMissingUsername, - }, - { - desc: "username that is too long in request", - req: createUserReq{ - User: users.User{ - ID: validID, - FirstName: valid, - LastName: valid, - Email: validEmail, - Credentials: users.Credentials{ - Username: strings.Repeat("a", 33), - Secret: secret, - }, - }, - }, - err: apiutil.ErrInvalidUsername, - }, - { - desc: "invalid username format in request", - req: createUserReq{ - User: users.User{ - ID: validID, - FirstName: valid, - LastName: valid, - Email: validEmail, - Credentials: users.Credentials{ - Username: "_invalid@username", - Secret: secret, - }, - }, - }, - err: apiutil.ErrInvalidUsername, - }, - } - for _, tc := range cases { - err := tc.req.validate() - assert.Equal(t, tc.err, err, "%s: expected %s got %s\n", tc.desc, tc.err, err) - } -} - -func TestViewUserReqValidate(t *testing.T) { - cases := []struct { - desc string - req viewUserReq - err error - }{ - { - desc: "valid request", - req: viewUserReq{ - id: validID, - }, - err: nil, - }, - { - desc: "empty id", - req: viewUserReq{ - id: "", - }, - err: apiutil.ErrMissingID, - }, - } - for _, c := range cases { - err := c.req.validate() - assert.Equal(t, c.err, err, "%s: expected %s got %s\n", c.desc, c.err, err) - } -} - -func TestListUsersReqValidate(t *testing.T) { - cases := []struct { - desc string - req listUsersReq - err error - }{ - { - desc: "valid request", - req: listUsersReq{ - limit: 10, - }, - err: nil, - }, - { - desc: "limit too big", - req: listUsersReq{ - limit: api.MaxLimitSize + 1, - }, - err: apiutil.ErrLimitSize, - }, - { - desc: "limit too small", - req: listUsersReq{ - limit: 0, - }, - err: apiutil.ErrLimitSize, - }, - { - desc: "invalid direction", - req: listUsersReq{ - limit: 10, - dir: "invalid", - }, - err: apiutil.ErrInvalidDirection, - }, - } - for _, c := range cases { - err := c.req.validate() - assert.Equal(t, c.err, err, "%s: expected %s got %s\n", c.desc, c.err, err) - } -} - -func TestSearchUsersReqValidate(t *testing.T) { - cases := []struct { - desc string - req searchUsersReq - err error - }{ - { - desc: "valid request", - req: searchUsersReq{ - Username: name, - }, - err: nil, - }, - { - desc: "empty query", - req: searchUsersReq{}, - err: apiutil.ErrEmptySearchQuery, - }, - } - for _, c := range cases { - err := c.req.validate() - assert.Equal(t, c.err, err) - } -} - -func TestUpdateUserReqValidate(t *testing.T) { - cases := []struct { - desc string - req updateUserReq - err error - }{ - { - desc: "valid request", - req: updateUserReq{ - id: validID, - }, - err: nil, - }, - { - desc: "empty id", - req: updateUserReq{ - id: "", - }, - err: apiutil.ErrMissingID, - }, - } - for _, c := range cases { - err := c.req.validate() - assert.Equal(t, c.err, err, "%s: expected %s got %s\n", c.desc, c.err, err) - } -} - -func TestUpdateUserTagsReqValidate(t *testing.T) { - tags := []string{"tag1", "tag2"} - cases := []struct { - desc string - req updateUserTagsReq - err error - }{ - { - desc: "valid request", - req: updateUserTagsReq{ - id: validID, - Tags: &tags, - }, - err: nil, - }, - { - desc: "empty id", - req: updateUserTagsReq{ - id: "", - Tags: &tags, - }, - err: apiutil.ErrMissingID, - }, - } - for _, c := range cases { - err := c.req.validate() - assert.Equal(t, c.err, err, "%s: expected %s got %s\n", c.desc, c.err, err) - } -} - -func TestUpdateUsernameReqValidate(t *testing.T) { - cases := []struct { - desc string - req updateUsernameReq - err error - }{ - { - desc: "valid request", - req: updateUsernameReq{ - id: validID, - Username: "validUsername", - }, - err: nil, - }, - { - desc: "missing user ID", - req: updateUsernameReq{ - id: "", - Username: "validUsername", - }, - err: apiutil.ErrMissingID, - }, - { - desc: "name too long", - req: updateUsernameReq{ - id: validID, - Username: strings.Repeat("a", api.MaxNameSize+1), - }, - err: apiutil.ErrNameSize, - }, - } - for _, tc := range cases { - err := tc.req.validate() - assert.Equal(t, tc.err, err, "%s: expected %s got %s\n", tc.desc, tc.err, err) - } -} - -func TestUpdateProfilePictureReqValidate(t *testing.T) { - base64EncodedString := "https://example.com/profile.jpg" - - parsedURL, err := url.Parse(base64EncodedString) - if err != nil { - t.Fatalf("Error parsing URL: %v", err) - } - url := parsedURL.String() - cases := []struct { - desc string - req updateProfilePictureReq - err error - }{ - { - desc: "valid request", - req: updateProfilePictureReq{ - id: validID, - ProfilePicture: &url, - }, - err: nil, - }, - { - desc: "empty ID", - req: updateProfilePictureReq{ - id: "", - ProfilePicture: &url, - }, - err: apiutil.ErrMissingID, - }, - } - for _, tc := range cases { - err := tc.req.validate() - assert.Equal(t, tc.err, err, "%s: expected %s got %s\n", tc.desc, tc.err, err) - } -} - -func TestUpdateUserRoleReqValidate(t *testing.T) { - cases := []struct { - desc string - req updateUserRoleReq - err error - }{ - { - desc: "valid request", - req: updateUserRoleReq{ - id: validID, - Role: "admin", - }, - err: nil, - }, - { - desc: "empty id", - req: updateUserRoleReq{ - id: "", - Role: "admin", - }, - err: apiutil.ErrMissingID, - }, - } - for _, c := range cases { - err := c.req.validate() - assert.Equal(t, c.err, err, "%s: expected %s got %s\n", c.desc, c.err, err) - } -} - -func TestUpdateUserEmailReqValidate(t *testing.T) { - cases := []struct { - desc string - req updateEmailReq - err error - }{ - { - desc: "valid request", - req: updateEmailReq{ - id: validID, - Email: "example@example.com", - }, - err: nil, - }, - { - desc: "empty id", - req: updateEmailReq{ - id: "", - Email: "example@example.com", - }, - err: apiutil.ErrMissingID, - }, - } - for _, c := range cases { - err := c.req.validate() - assert.Equal(t, c.err, err, "%s: expected %s got %s\n", c.desc, c.err, err) - } -} - -func TestUpdateUserSecretReqValidate(t *testing.T) { - cases := []struct { - desc string - req updateUserSecretReq - err error - }{ - { - desc: "valid request", - req: updateUserSecretReq{ - OldSecret: secret, - NewSecret: secret, - }, - err: nil, - }, - { - desc: "missing old secret", - req: updateUserSecretReq{ - OldSecret: "", - NewSecret: secret, - }, - err: apiutil.ErrMissingPass, - }, - { - desc: "missing new secret", - req: updateUserSecretReq{ - OldSecret: secret, - NewSecret: "", - }, - err: apiutil.ErrMissingPass, - }, - { - desc: "invalid new secret", - req: updateUserSecretReq{ - OldSecret: secret, - NewSecret: "invalid", - }, - err: apiutil.ErrPasswordFormat, - }, - } - for _, c := range cases { - err := c.req.validate() - assert.Equal(t, c.err, err) - } -} - -func TestChangeUserStatusReqValidate(t *testing.T) { - cases := []struct { - desc string - req changeUserStatusReq - err error - }{ - { - desc: "valid request", - req: changeUserStatusReq{ - id: validID, - }, - err: nil, - }, - { - desc: "empty id", - req: changeUserStatusReq{ - id: "", - }, - err: apiutil.ErrMissingID, - }, - } - for _, c := range cases { - err := c.req.validate() - assert.Equal(t, c.err, err, "%s: expected %s got %s\n", c.desc, c.err, err) - } -} - -func TestLoginUserReqValidate(t *testing.T) { - cases := []struct { - desc string - req loginUserReq - err error - }{ - { - desc: "valid request with identity", - req: loginUserReq{ - Username: "example", - Password: secret, - }, - err: nil, - }, - { - desc: "empty identity", - req: loginUserReq{ - Username: "", - Password: secret, - }, - err: apiutil.ErrMissingUsernameEmail, - }, - { - desc: "empty secret", - req: loginUserReq{ - Password: "", - Username: "example", - }, - err: apiutil.ErrMissingPass, - }, - } - for _, c := range cases { - err := c.req.validate() - assert.Equal(t, c.err, err, "%s: expected %s got %s\n", c.desc, c.err, err) - } -} - -func TestTokenReqValidate(t *testing.T) { - cases := []struct { - desc string - req tokenReq - err error - }{ - { - desc: "valid request", - req: tokenReq{ - RefreshToken: valid, - }, - err: nil, - }, - { - desc: "empty token", - req: tokenReq{ - RefreshToken: "", - }, - err: apiutil.ErrBearerToken, - }, - } - for _, c := range cases { - err := c.req.validate() - assert.Equal(t, c.err, err, "%s: expected %s got %s\n", c.desc, c.err, err) - } -} - -func TestPasswResetReqValidate(t *testing.T) { - cases := []struct { - desc string - req passResetReq - err error - }{ - { - desc: "valid request", - req: passResetReq{ - Email: "example@example.com", - }, - err: nil, - }, - { - desc: "empty email", - req: passResetReq{ - Email: "", - }, - err: apiutil.ErrMissingEmail, - }, - } - for _, c := range cases { - err := c.req.validate() - assert.Equal(t, c.err, err, "%s: expected %s got %s\n", c.desc, c.err, err) - } -} - -func TestResetTokenReqValidate(t *testing.T) { - cases := []struct { - desc string - req resetTokenReq - err error - }{ - { - desc: "valid request", - req: resetTokenReq{ - Password: secret, - ConfPass: secret, - }, - err: nil, - }, - { - desc: "empty password", - req: resetTokenReq{ - Password: "", - ConfPass: secret, - }, - err: apiutil.ErrMissingPass, - }, - { - desc: "empty confpass", - req: resetTokenReq{ - Password: secret, - ConfPass: "", - }, - err: apiutil.ErrMissingConfPass, - }, - { - desc: "mismatching password and confpass", - req: resetTokenReq{ - Password: "secret", - ConfPass: secret, - }, - err: apiutil.ErrInvalidResetPass, - }, - } - for _, c := range cases { - err := c.req.validate() - assert.Equal(t, c.err, err) - } -} diff --git a/users/api/responses.go b/users/api/responses.go deleted file mode 100644 index fac085408..000000000 --- a/users/api/responses.go +++ /dev/null @@ -1,256 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package api - -import ( - "fmt" - "net/http" - - "github.com/absmach/magistrala" - grpcTokenV1 "github.com/absmach/magistrala/api/grpc/token/v1" - "github.com/absmach/magistrala/users" -) - -// MailSent message response when link is sent. -const MailSent = "Email with reset link is sent" - -var ( - _ magistrala.Response = (*tokenRes)(nil) - _ magistrala.Response = (*sendVerificationRes)(nil) - _ magistrala.Response = (*verifyEmailRes)(nil) - _ magistrala.Response = (*viewUserRes)(nil) - _ magistrala.Response = (*createUserRes)(nil) - _ magistrala.Response = (*changeUserStatusRes)(nil) - _ magistrala.Response = (*usersPageRes)(nil) - _ magistrala.Response = (*passResetReqRes)(nil) - _ magistrala.Response = (*passChangeRes)(nil) - _ magistrala.Response = (*updateUserRes)(nil) - _ magistrala.Response = (*revokeRes)(nil) - _ magistrala.Response = (*deleteUserRes)(nil) - _ magistrala.Response = (*listRefreshTokensRes)(nil) -) - -type pageRes struct { - Limit uint64 `json:"limit,omitempty"` - Offset uint64 `json:"offset,omitempty"` - Total uint64 `json:"total"` -} - -type createUserRes struct { - users.User - created bool -} - -func (res createUserRes) Code() int { - if res.created { - return http.StatusCreated - } - - return http.StatusOK -} - -func (res createUserRes) Headers() map[string]string { - if res.created { - return map[string]string{ - "Location": fmt.Sprintf("/users/%s", res.ID), - } - } - - return map[string]string{} -} - -func (res createUserRes) Empty() bool { - return false -} - -type tokenRes struct { - AccessToken string `json:"access_token,omitempty"` - RefreshToken string `json:"refresh_token,omitempty"` - AccessType string `json:"access_type,omitempty"` -} - -func (res tokenRes) Code() int { - return http.StatusCreated -} - -func (res tokenRes) Headers() map[string]string { - return map[string]string{} -} - -func (res tokenRes) Empty() bool { - return res.AccessToken == "" || res.RefreshToken == "" -} - -type revokeRes struct{} - -func (res revokeRes) Code() int { - return http.StatusNoContent -} - -func (res revokeRes) Headers() map[string]string { - return map[string]string{} -} - -func (res revokeRes) Empty() bool { - return true -} - -type listRefreshTokensRes struct { - RefreshTokens []*grpcTokenV1.RefreshToken `json:"refresh_tokens"` -} - -func (res listRefreshTokensRes) Code() int { - return http.StatusOK -} - -func (res listRefreshTokensRes) Headers() map[string]string { - return map[string]string{} -} - -func (res listRefreshTokensRes) Empty() bool { - return false -} - -type sendVerificationRes struct{} - -func (res sendVerificationRes) Code() int { - return http.StatusOK -} - -func (res sendVerificationRes) Headers() map[string]string { - return map[string]string{} -} - -func (res sendVerificationRes) Empty() bool { - return true -} - -type verifyEmailRes struct{} - -func (res verifyEmailRes) Code() int { - return http.StatusOK -} - -func (res verifyEmailRes) Headers() map[string]string { - return map[string]string{} -} - -func (res verifyEmailRes) Empty() bool { - return true -} - -type updateUserRes struct { - users.User `json:",inline"` -} - -func (res updateUserRes) Code() int { - return http.StatusOK -} - -func (res updateUserRes) Headers() map[string]string { - return map[string]string{} -} - -func (res updateUserRes) Empty() bool { - return false -} - -type viewUserRes struct { - users.User `json:",inline"` -} - -func (res viewUserRes) Code() int { - return http.StatusOK -} - -func (res viewUserRes) Headers() map[string]string { - return map[string]string{} -} - -func (res viewUserRes) Empty() bool { - return false -} - -type usersPageRes struct { - pageRes - Users []viewUserRes `json:"users"` -} - -func (res usersPageRes) Code() int { - return http.StatusOK -} - -func (res usersPageRes) Headers() map[string]string { - return map[string]string{} -} - -func (res usersPageRes) Empty() bool { - return false -} - -type changeUserStatusRes struct { - users.User `json:",inline"` -} - -func (res changeUserStatusRes) Code() int { - return http.StatusOK -} - -func (res changeUserStatusRes) Headers() map[string]string { - return map[string]string{} -} - -func (res changeUserStatusRes) Empty() bool { - return false -} - -type passResetReqRes struct { - Msg string `json:"msg"` -} - -func (res passResetReqRes) Code() int { - return http.StatusCreated -} - -func (res passResetReqRes) Headers() map[string]string { - return map[string]string{} -} - -func (res passResetReqRes) Empty() bool { - return false -} - -type passChangeRes struct{} - -func (res passChangeRes) Code() int { - return http.StatusCreated -} - -func (res passChangeRes) Headers() map[string]string { - return map[string]string{} -} - -func (res passChangeRes) Empty() bool { - return false -} - -type deleteUserRes struct { - deleted bool -} - -func (res deleteUserRes) Code() int { - if res.deleted { - return http.StatusNoContent - } - - return http.StatusOK -} - -func (res deleteUserRes) Headers() map[string]string { - return map[string]string{} -} - -func (res deleteUserRes) Empty() bool { - return true -} diff --git a/users/api/transport.go b/users/api/transport.go deleted file mode 100644 index 2c2f2f3a5..000000000 --- a/users/api/transport.go +++ /dev/null @@ -1,28 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package api - -import ( - "log/slog" - "net/http" - "regexp" - - "github.com/absmach/magistrala" - grpcTokenV1 "github.com/absmach/magistrala/api/grpc/token/v1" - smqauthn "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/oauth2" - "github.com/absmach/magistrala/users" - "github.com/go-chi/chi/v5" - "github.com/prometheus/client_golang/prometheus/promhttp" -) - -// MakeHandler returns a HTTP handler for Users and Groups API endpoints. -func MakeHandler(cls users.Service, authn smqauthn.AuthNMiddleware, tokensvc grpcTokenV1.TokenServiceClient, selfRegister bool, mux *chi.Mux, logger *slog.Logger, instanceID string, pr *regexp.Regexp, idp magistrala.IDProvider, providers ...oauth2.Provider) http.Handler { - mux = usersHandler(cls, authn, tokensvc, selfRegister, mux, logger, pr, idp, providers...) - - mux.Get("/health", magistrala.Health("users", instanceID)) - mux.Handle("/metrics", promhttp.Handler()) - - return mux -} diff --git a/users/api/users.go b/users/api/users.go deleted file mode 100644 index ca35bbbf1..000000000 --- a/users/api/users.go +++ /dev/null @@ -1,681 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package api - -import ( - "context" - "encoding/json" - "log/slog" - "net/http" - "regexp" - "strings" - "time" - - "github.com/absmach/magistrala" - grpcTokenV1 "github.com/absmach/magistrala/api/grpc/token/v1" - api "github.com/absmach/magistrala/api/http" - apiutil "github.com/absmach/magistrala/api/http/util" - smqauth "github.com/absmach/magistrala/auth" - smqauthn "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/errors" - "github.com/absmach/magistrala/pkg/oauth2" - "github.com/absmach/magistrala/users" - "github.com/go-chi/chi/v5" - kithttp "github.com/go-kit/kit/transport/http" - "go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp" -) - -var passRegex = regexp.MustCompile("^.{8,}$") - -// usersHandler returns a HTTP handler for API endpoints. -func usersHandler(svc users.Service, authn smqauthn.AuthNMiddleware, tokenClient grpcTokenV1.TokenServiceClient, selfRegister bool, r *chi.Mux, logger *slog.Logger, pr *regexp.Regexp, idp magistrala.IDProvider, providers ...oauth2.Provider) *chi.Mux { - passRegex = pr - - opts := []kithttp.ServerOption{ - kithttp.ServerErrorEncoder(apiutil.LoggingErrorEncoder(logger, api.EncodeError)), - } - - // All endpoints in users service don't required Domain check - authn = authn.WithOptions(smqauthn.WithDomainCheck(false)) - r.Route("/users", func(r chi.Router) { - r.Use(api.RequestIDMiddleware(idp)) - - switch selfRegister { - case true: - r.Post("/", otelhttp.NewHandler(kithttp.NewServer( - registrationEndpoint(svc, selfRegister), - decodeCreateUserReq, - api.EncodeResponse, - opts..., - ), "register_user").ServeHTTP) - default: - r.With(authn.Middleware()).Post("/", otelhttp.NewHandler(kithttp.NewServer( - registrationEndpoint(svc, selfRegister), - decodeCreateUserReq, - api.EncodeResponse, - opts..., - ), "register_user").ServeHTTP) - } - // Endpoints which are allowed for unverified user - r.Group(func(r chi.Router) { - r.Use(authn.WithOptions(smqauthn.WithAllowUnverifiedUser(true)).Middleware()) - r.Post("/send-verification", otelhttp.NewHandler(kithttp.NewServer( - sendVerificationEndpoint(svc), - decodeSendVerification, - api.EncodeResponse, - opts..., - ), "send_verification").ServeHTTP) - - r.Get("/profile", otelhttp.NewHandler(kithttp.NewServer( - viewProfileEndpoint(svc), - decodeViewProfile, - api.EncodeResponse, - opts..., - ), "view_profile").ServeHTTP) - r.Post("/tokens/refresh", otelhttp.NewHandler(kithttp.NewServer( - refreshTokenEndpoint(svc), - decodeRefreshToken, - api.EncodeResponse, - opts..., - ), "refresh_token").ServeHTTP) - r.Post("/tokens/revoke", otelhttp.NewHandler(kithttp.NewServer( - revokeRefreshTokenEndpoint(svc), - decodeRevokeRefreshToken, - api.EncodeResponse, - opts..., - ), "revoke_refresh_token").ServeHTTP) - r.Get("/tokens/refresh-tokens", otelhttp.NewHandler(kithttp.NewServer( - listActiveRefreshTokensEndpoint(svc), - decodeListActiveRefreshTokens, - api.EncodeResponse, - opts..., - ), "list_active_refresh_tokens").ServeHTTP) - r.Patch("/{id}/email", otelhttp.NewHandler(kithttp.NewServer( - updateEmailEndpoint(svc), - decodeUpdateUserEmail, - api.EncodeResponse, - opts..., - ), "update_user_email").ServeHTTP) - }) - - r.Group(func(r chi.Router) { - r.Use(authn.Middleware()) - - r.Get("/{id}", otelhttp.NewHandler(kithttp.NewServer( - viewEndpoint(svc), - decodeViewUser, - api.EncodeResponse, - opts..., - ), "view_user").ServeHTTP) - - r.Get("/", otelhttp.NewHandler(kithttp.NewServer( - listUsersEndpoint(svc), - decodeListUsers, - api.EncodeResponse, - opts..., - ), "list_users").ServeHTTP) - - r.Get("/search", otelhttp.NewHandler(kithttp.NewServer( - searchUsersEndpoint(svc), - decodeSearchUsers, - api.EncodeResponse, - opts..., - ), "search_users").ServeHTTP) - - r.Patch("/secret", otelhttp.NewHandler(kithttp.NewServer( - updateSecretEndpoint(svc), - decodeUpdateUserSecret, - api.EncodeResponse, - opts..., - ), "update_user_secret").ServeHTTP) - - r.Patch("/{id}", otelhttp.NewHandler(kithttp.NewServer( - updateEndpoint(svc), - decodeUpdateUser, - api.EncodeResponse, - opts..., - ), "update_user").ServeHTTP) - - r.Patch("/{id}/username", otelhttp.NewHandler(kithttp.NewServer( - updateUsernameEndpoint(svc), - decodeUpdateUsername, - api.EncodeResponse, - opts..., - ), "update_username").ServeHTTP) - - r.Patch("/{id}/picture", otelhttp.NewHandler(kithttp.NewServer( - updateProfilePictureEndpoint(svc), - decodeUpdateUserProfilePicture, - api.EncodeResponse, - opts..., - ), "update_profile_picture").ServeHTTP) - - r.Patch("/{id}/tags", otelhttp.NewHandler(kithttp.NewServer( - updateTagsEndpoint(svc), - decodeUpdateUserTags, - api.EncodeResponse, - opts..., - ), "update_user_tags").ServeHTTP) - - r.Patch("/{id}/role", otelhttp.NewHandler(kithttp.NewServer( - updateRoleEndpoint(svc), - decodeUpdateUserRole, - api.EncodeResponse, - opts..., - ), "update_user_role").ServeHTTP) - - r.Post("/{id}/enable", otelhttp.NewHandler(kithttp.NewServer( - enableEndpoint(svc), - decodeChangeUserStatus, - api.EncodeResponse, - opts..., - ), "enable_user").ServeHTTP) - - r.Post("/{id}/disable", otelhttp.NewHandler(kithttp.NewServer( - disableEndpoint(svc), - decodeChangeUserStatus, - api.EncodeResponse, - opts..., - ), "disable_user").ServeHTTP) - - r.Delete("/{id}", otelhttp.NewHandler(kithttp.NewServer( - deleteEndpoint(svc), - decodeChangeUserStatus, - api.EncodeResponse, - opts..., - ), "delete_user").ServeHTTP) - }) - }) - - r.Group(func(r chi.Router) { - r.Use(authn.WithOptions(smqauthn.WithAllowUnverifiedUser(true)).Middleware()) - r.Put("/password/reset", otelhttp.NewHandler(kithttp.NewServer( - passwordResetEndpoint(svc), - decodePasswordReset, - api.EncodeResponse, - opts..., - ), "password_reset").ServeHTTP) - }) - - r.Post("/users/tokens/issue", otelhttp.NewHandler(kithttp.NewServer( - issueTokenEndpoint(svc), - decodeCredentials, - api.EncodeResponse, - opts..., - ), "issue_token").ServeHTTP) - - r.Post("/password/reset-request", otelhttp.NewHandler(kithttp.NewServer( - passwordResetRequestEndpoint(svc), - decodePasswordResetRequest, - api.EncodeResponse, - opts..., - ), "password_reset_req").ServeHTTP) - - r.Get("/verify-email", otelhttp.NewHandler(kithttp.NewServer( - verifyEmailEndpoint(svc), - decodeVerifyEmail, - api.EncodeResponse, - opts..., - ), "verify_email").ServeHTTP) - - for _, provider := range providers { - r.HandleFunc("/oauth/callback/"+provider.Name(), oauth2CallbackHandler(provider, svc, tokenClient)) - } - - return r -} - -func decodeSendVerification(_ context.Context, r *http.Request) (any, error) { - req := sendVerificationReq{} - return req, nil -} - -func decodeVerifyEmail(_ context.Context, r *http.Request) (any, error) { - token, err := apiutil.ReadStringQuery(r, api.TokenKey, "") - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - return verifyEmailReq{ - token: token, - }, nil -} - -func decodeViewUser(_ context.Context, r *http.Request) (any, error) { - req := viewUserReq{ - id: chi.URLParam(r, "id"), - } - - return req, nil -} - -func decodeViewProfile(_ context.Context, r *http.Request) (any, error) { - return nil, nil -} - -func decodeListUsers(_ context.Context, r *http.Request) (any, error) { - s, err := apiutil.ReadStringQuery(r, api.StatusKey, api.DefUserStatus) - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - o, err := apiutil.ReadNumQuery[uint64](r, api.OffsetKey, api.DefOffset) - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - l, err := apiutil.ReadNumQuery[uint64](r, api.LimitKey, api.DefLimit) - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - m, err := apiutil.ReadMetadataQuery(r, api.MetadataKey, nil) - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - n, err := apiutil.ReadStringQuery(r, api.UsernameKey, "") - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - d, err := apiutil.ReadStringQuery(r, api.EmailKey, "") - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - i, err := apiutil.ReadStringQuery(r, api.FirstNameKey, "") - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - f, err := apiutil.ReadStringQuery(r, api.LastNameKey, "") - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - t, err := apiutil.ReadStringQuery(r, api.TagsKey, "") - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - var tq users.TagsQuery - if t != "" { - tq = users.ToTagsQuery(t) - } - order, err := apiutil.ReadStringQuery(r, api.OrderKey, api.DefOrder) - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - dir, err := apiutil.ReadStringQuery(r, api.DirKey, api.DefDir) - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - id, err := apiutil.ReadStringQuery(r, api.IDOrder, "") - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - st, err := users.ToStatus(s) - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - ot, err := apiutil.ReadBoolQuery(r, api.OnlyTotal, false) - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - cfrom, err := apiutil.ReadStringQuery(r, "created_from", "") - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - cto, err := apiutil.ReadStringQuery(r, "created_to", "") - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - var createdFrom, createdTo time.Time - if cfrom != "" { - if createdFrom, err = time.Parse(time.RFC3339, cfrom); err != nil { - return nil, errors.Wrap(apiutil.ErrInvalidQueryParams, err) - } - } - if cto != "" { - if createdTo, err = time.Parse(time.RFC3339, cto); err != nil { - return nil, errors.Wrap(apiutil.ErrInvalidQueryParams, err) - } - } - - req := listUsersReq{ - status: st, - offset: o, - limit: l, - onlyTotal: ot, - metadata: m, - userName: n, - firstName: i, - lastName: f, - tags: tq, - order: order, - dir: dir, - id: id, - email: d, - createdFrom: createdFrom, - createdTo: createdTo, - } - - return req, nil -} - -func decodeSearchUsers(_ context.Context, r *http.Request) (any, error) { - o, err := apiutil.ReadNumQuery[uint64](r, api.OffsetKey, api.DefOffset) - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - l, err := apiutil.ReadNumQuery[uint64](r, api.LimitKey, api.DefLimit) - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - n, err := apiutil.ReadStringQuery(r, api.UsernameKey, "") - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - f, err := apiutil.ReadStringQuery(r, api.FirstNameKey, "") - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - e, err := apiutil.ReadStringQuery(r, api.LastNameKey, "") - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - id, err := apiutil.ReadStringQuery(r, api.IDOrder, "") - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - order, err := apiutil.ReadStringQuery(r, api.OrderKey, api.DefOrder) - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - dir, err := apiutil.ReadStringQuery(r, api.DirKey, api.DefDir) - if err != nil { - return nil, errors.Wrap(apiutil.ErrValidation, err) - } - - req := searchUsersReq{ - Offset: o, - Limit: l, - Username: n, - FirstName: f, - LastName: e, - Id: id, - Order: order, - Dir: dir, - } - - for _, field := range []string{req.Username, req.Id} { - if field != "" && len(field) < 3 { - req = searchUsersReq{} - return req, errors.Wrap(apiutil.ErrLenSearchQuery, apiutil.ErrValidation) - } - } - - return req, nil -} - -func decodeUpdateUser(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - - req := updateUserReq{ - id: chi.URLParam(r, "id"), - } - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - - return req, nil -} - -func decodeUpdateUserTags(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - - req := updateUserTagsReq{ - id: chi.URLParam(r, "id"), - } - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - - return req, nil -} - -func decodeUpdateUserEmail(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - - req := updateEmailReq{ - id: chi.URLParam(r, "id"), - } - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - - return req, nil -} - -func decodeUpdateUserSecret(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - - req := updateUserSecretReq{} - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - - return req, nil -} - -func decodeUpdateUsername(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - - req := updateUsernameReq{ - id: chi.URLParam(r, "id"), - } - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - - return req, nil -} - -func decodeUpdateUserProfilePicture(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - - req := updateProfilePictureReq{ - id: chi.URLParam(r, "id"), - } - - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - - return req, nil -} - -func decodePasswordResetRequest(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, apiutil.ErrUnsupportedContentType - } - - var req passResetReq - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - - return req, nil -} - -func decodePasswordReset(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - - var req resetTokenReq - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - - return req, nil -} - -func decodeUpdateUserRole(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - - req := updateUserRoleReq{ - id: chi.URLParam(r, "id"), - } - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - var err error - req.role, err = users.ToRole(req.Role) - return req, err -} - -func decodeCredentials(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - - req := loginUserReq{} - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - - return req, nil -} - -func decodeRefreshToken(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - req := tokenReq{RefreshToken: apiutil.ExtractBearerToken(r)} - - return req, nil -} - -func decodeRevokeRefreshToken(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - - var req revokeTokenReq - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - - return req, nil -} - -func decodeListActiveRefreshTokens(_ context.Context, r *http.Request) (any, error) { - return nil, nil -} - -func decodeCreateUserReq(_ context.Context, r *http.Request) (any, error) { - if !strings.Contains(r.Header.Get("Content-Type"), api.ContentType) { - return nil, errors.Wrap(apiutil.ErrValidation, apiutil.ErrUnsupportedContentType) - } - - var req createUserReq - if err := json.NewDecoder(r.Body).Decode(&req); err != nil { - return nil, errors.Wrap(apiutil.ErrMalformedRequestBody, err) - } - - return req, nil -} - -func decodeChangeUserStatus(_ context.Context, r *http.Request) (any, error) { - req := changeUserStatusReq{ - id: chi.URLParam(r, "id"), - } - - return req, nil -} - -// oauth2CallbackHandler is a http.HandlerFunc that handles OAuth2 callbacks. -func oauth2CallbackHandler(oauth oauth2.Provider, svc users.Service, tokenClient grpcTokenV1.TokenServiceClient) http.HandlerFunc { - return func(w http.ResponseWriter, r *http.Request) { - if !oauth.IsEnabled() { - http.Redirect(w, r, oauth.ErrorURL()+"?error=oauth%20provider%20is%20disabled", http.StatusSeeOther) - return - } - state := r.FormValue("state") - if state != oauth.State() { - http.Redirect(w, r, oauth.ErrorURL()+"?error=invalid%20state", http.StatusSeeOther) - return - } - - if code := r.FormValue("code"); code != "" { - token, err := oauth.Exchange(r.Context(), code) - if err != nil { - http.Redirect(w, r, oauth.ErrorURL()+"?error="+err.Error(), http.StatusSeeOther) - return - } - - user, err := oauth.UserInfo(token.AccessToken) - if err != nil { - http.Redirect(w, r, oauth.ErrorURL()+"?error="+err.Error(), http.StatusSeeOther) - return - } - - user.AuthProvider = oauth.Name() - if user.AuthProvider == "" { - user.AuthProvider = "oauth" - } - user, err = svc.OAuthCallback(r.Context(), user) - if err != nil { - http.Redirect(w, r, oauth.ErrorURL()+"?error="+err.Error(), http.StatusSeeOther) - return - } - if err := svc.OAuthAddUserPolicy(r.Context(), user); err != nil { - http.Redirect(w, r, oauth.ErrorURL()+"?error="+err.Error(), http.StatusSeeOther) - return - } - - jwt, err := tokenClient.Issue(r.Context(), &grpcTokenV1.IssueReq{ - UserId: user.ID, - Type: uint32(smqauth.AccessKey), - UserRole: uint32(smqauth.UserRole), - Verified: !user.VerifiedAt.IsZero(), - }) - if err != nil { - http.Redirect(w, r, oauth.ErrorURL()+"?error="+err.Error(), http.StatusSeeOther) - return - } - - http.SetCookie(w, &http.Cookie{ - Name: "access_token", - Value: jwt.GetAccessToken(), - Path: "/", - HttpOnly: true, - Secure: true, - }) - http.SetCookie(w, &http.Cookie{ - Name: "refresh_token", - Value: jwt.GetRefreshToken(), - Path: "/", - HttpOnly: true, - Secure: true, - }) - - http.Redirect(w, r, oauth.RedirectURL(), http.StatusFound) - return - } - - http.Redirect(w, r, oauth.ErrorURL()+"?error=empty%20code", http.StatusSeeOther) - } -} diff --git a/users/delete_handler.go b/users/delete_handler.go deleted file mode 100644 index 28dcd8987..000000000 --- a/users/delete_handler.go +++ /dev/null @@ -1,109 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// The DeleteHandler is a cron job that runs periodically to delete users that have been marked as deleted -// for a certain period of time together with the user's policies from the auth service. -// The handler runs in a separate goroutine and checks for users that have been marked as deleted for a certain period of time. -// If the user has been marked as deleted for more than the specified period, -// the handler deletes the user's policies from the auth service and deletes the user from the database. - -package users - -import ( - "context" - "log/slog" - "time" - - grpcDomainsV1 "github.com/absmach/magistrala/api/grpc/domains/v1" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/pkg/policies" -) - -const defLimit = uint64(100) - -type handler struct { - users Repository - domains grpcDomainsV1.DomainsServiceClient - policies policies.Service - checkInterval time.Duration - deleteAfter time.Duration - logger *slog.Logger -} - -func NewDeleteHandler(ctx context.Context, users Repository, policyService policies.Service, domainsClient grpcDomainsV1.DomainsServiceClient, defCheckInterval, deleteAfter time.Duration, logger *slog.Logger) { - handler := &handler{ - users: users, - domains: domainsClient, - policies: policyService, - checkInterval: defCheckInterval, - deleteAfter: deleteAfter, - logger: logger, - } - - go func() { - ticker := time.NewTicker(handler.checkInterval) - defer ticker.Stop() - - for { - select { - case <-ctx.Done(): - return - case <-ticker.C: - handler.handle(ctx) - } - } - }() -} - -func (h *handler) handle(ctx context.Context) { - pm := Page{Limit: defLimit, Offset: 0, Status: DeletedStatus} - - for { - dbUsers, err := h.users.RetrieveAll(ctx, pm) - if err != nil { - h.logger.Error("failed to retrieve users", slog.Any("error", err)) - break - } - if dbUsers.Total == 0 { - break - } - - for _, u := range dbUsers.Users { - if time.Since(u.UpdatedAt) < h.deleteAfter { - continue - } - - deletedRes, err := h.domains.DeleteUserFromDomains(ctx, &grpcDomainsV1.DeleteUserReq{ - Id: u.ID, - }) - if err != nil { - h.logger.Error("failed to delete user from domains", slog.Any("error", err)) - continue - } - if !deletedRes.Deleted { - h.logger.Error("failed to delete user from domains", slog.Any("error", svcerr.ErrAuthorization)) - continue - } - - req := policies.Policy{ - Subject: u.ID, - SubjectType: policies.UserType, - } - if err := h.policies.DeletePolicyFilter(ctx, req); err != nil { - h.logger.Error("failed to delete user policies", slog.Any("error", err)) - continue - } - - if err := h.users.Delete(ctx, u.ID); err != nil { - h.logger.Error("failed to delete user", slog.Any("error", err)) - continue - } - - h.logger.Info("user deleted", slog.Group("user", - slog.String("id", u.ID), - slog.String("first_name", u.FirstName), - slog.String("last_name", u.LastName), - )) - } - } -} diff --git a/users/doc.go b/users/doc.go deleted file mode 100644 index 242071154..000000000 --- a/users/doc.go +++ /dev/null @@ -1,11 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package users contains the domain concept definitions needed to -// support Magistrala users service functionality. -// -// This package defines the core domain concepts and types necessary to -// handle users in the context of a Magistrala users service. It abstracts -// the underlying complexities of user management and provides a structured -// approach to working with users. -package users diff --git a/users/emailer.go b/users/emailer.go deleted file mode 100644 index 4b93df933..000000000 --- a/users/emailer.go +++ /dev/null @@ -1,13 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package users - -// Emailer wrapper around the email. -type Emailer interface { - // SendPasswordReset sends an email to the user with a link to reset the password. - SendPasswordReset(To []string, user, token string) error - - // SendVerification sends an email to the user with a verification token. - SendVerification(To []string, user, verificationToken string) error -} diff --git a/users/emailer/doc.go b/users/emailer/doc.go deleted file mode 100644 index 4db3fb1c8..000000000 --- a/users/emailer/doc.go +++ /dev/null @@ -1,6 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package emailer contains the domain concept definitions needed to support -// Magistrala users email service functionality. -package emailer diff --git a/users/emailer/emailer.go b/users/emailer/emailer.go deleted file mode 100644 index d5b6e66f7..000000000 --- a/users/emailer/emailer.go +++ /dev/null @@ -1,50 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package emailer - -import ( - "fmt" - - "github.com/absmach/magistrala/internal/email" - "github.com/absmach/magistrala/users" -) - -var _ users.Emailer = (*emailer)(nil) - -type emailer struct { - resetURL string - verificationURL string - resetAgent *email.Agent - verifyAgent *email.Agent -} - -// New creates new emailer utility. -func New(resetURL, verificationURL string, resetConfig, verifyConfig *email.Config) (users.Emailer, error) { - resetAgent, err := email.New(resetConfig) - if err != nil { - return nil, err - } - - verifyAgent, err := email.New(verifyConfig) - if err != nil { - return nil, err - } - - return &emailer{ - resetURL: resetURL, - verificationURL: verificationURL, - resetAgent: resetAgent, - verifyAgent: verifyAgent, - }, nil -} - -func (e *emailer) SendPasswordReset(to []string, user, token string) error { - url := fmt.Sprintf("%s?token=%s", e.resetURL, token) - return e.resetAgent.Send(to, "", "Password Reset Request", "", user, url, "", nil) -} - -func (e *emailer) SendVerification(to []string, user, verificationToken string) error { - url := fmt.Sprintf("%s?token=%s", e.verificationURL, verificationToken) - return e.verifyAgent.Send(to, "", "Email Verification", "", user, url, "", nil) -} diff --git a/users/events/doc.go b/users/events/doc.go deleted file mode 100644 index 86f9918a2..000000000 --- a/users/events/doc.go +++ /dev/null @@ -1,6 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package events provides the domain concept definitions needed to -// support Magistrala users service functionality. -package events diff --git a/users/events/events.go b/users/events/events.go deleted file mode 100644 index 0252085f5..000000000 --- a/users/events/events.go +++ /dev/null @@ -1,581 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package events - -import ( - "time" - - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/events" - "github.com/absmach/magistrala/users" -) - -const ( - userPrefix = "user." - userCreate = userPrefix + "create" - userSendVerification = userPrefix + "send_verification" - userVerifyEmail = userPrefix + "verify_email" - userUpdate = userPrefix + "update" - userUpdateRole = userPrefix + "update_role" - userUpdateTags = userPrefix + "update_tags" - userUpdateSecret = userPrefix + "update_secret" - userUpdateUsername = userPrefix + "update_username" - userUpdateProfilePicture = userPrefix + "update_profile_picture" - userUpdateEmail = userPrefix + "update_email" - userEnable = userPrefix + "enable" - userDisable = userPrefix + "disable" - userView = userPrefix + "view" - profileView = userPrefix + "view_profile" - userList = userPrefix + "list" - userSearch = userPrefix + "search" - userIdentify = userPrefix + "identify" - issueToken = userPrefix + "issue_token" - refreshToken = userPrefix + "refresh_token" - revokeRefreshToken = userPrefix + "revoke_refresh_token" - resetSecret = userPrefix + "reset_secret" - sendPasswordReset = userPrefix + "send_password_reset" - oauthCallback = userPrefix + "oauth_callback" - addClientPolicy = userPrefix + "add_policy" - deleteUser = userPrefix + "delete" -) - -var ( - _ events.Event = (*createUserEvent)(nil) - _ events.Event = (*sendVerificationEvent)(nil) - _ events.Event = (*verifyEmailEvent)(nil) - _ events.Event = (*updateUserEvent)(nil) - _ events.Event = (*updateProfilePictureEvent)(nil) - _ events.Event = (*updateUsernameEvent)(nil) - _ events.Event = (*changeUserStatusEvent)(nil) - _ events.Event = (*viewUserEvent)(nil) - _ events.Event = (*viewProfileEvent)(nil) - _ events.Event = (*listUserEvent)(nil) - _ events.Event = (*searchUserEvent)(nil) - _ events.Event = (*identifyUserEvent)(nil) - _ events.Event = (*issueTokenEvent)(nil) - _ events.Event = (*refreshTokenEvent)(nil) - _ events.Event = (*revokeRefreshTokenEvent)(nil) - _ events.Event = (*resetSecretEvent)(nil) - _ events.Event = (*sendPasswordResetEvent)(nil) - _ events.Event = (*oauthCallbackEvent)(nil) - _ events.Event = (*deleteUserEvent)(nil) - _ events.Event = (*addUserPolicyEvent)(nil) -) - -type createUserEvent struct { - users.User - authn.Session - requestID string -} - -func (uce createUserEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": userCreate, - "id": uce.ID, - "status": uce.Status.String(), - "created_at": uce.CreatedAt, - "token_type": uce.Type.String(), - "super_admin": uce.SuperAdmin, - "request_id": uce.requestID, - } - - if uce.FirstName != "" { - val["first_name"] = uce.FirstName - } - if uce.LastName != "" { - val["last_name"] = uce.LastName - } - if len(uce.Tags) > 0 { - val["tags"] = uce.Tags - } - if uce.Metadata != nil { - val["metadata"] = uce.Metadata - } - if uce.PrivateMetadata != nil { - val["private_metadata"] = uce.PrivateMetadata - } - if uce.Credentials.Username != "" { - val["username"] = uce.Credentials.Username - } - if uce.Email != "" { - val["email"] = uce.Email - } - - return val, nil -} - -type sendVerificationEvent struct { - authn.Session - requestID string -} - -func (sve sendVerificationEvent) Encode() (map[string]any, error) { - return map[string]any{ - "operation": userSendVerification, - "user_id": sve.UserID, - "token_type": sve.Type.String(), - "request_id": sve.requestID, - }, nil -} - -type verifyEmailEvent struct { - requestID string - email string - userID string - verifiedAt time.Time -} - -func (vee verifyEmailEvent) Encode() (map[string]any, error) { - return map[string]any{ - "operation": userVerifyEmail, - "request_id": vee.requestID, - "email": vee.email, - "user_id": vee.userID, - "verified_at": vee.verifiedAt, - }, nil -} - -type updateUserEvent struct { - users.User - operation string - authn.Session - requestID string -} - -func (uce updateUserEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": uce.operation, - "updated_at": uce.UpdatedAt, - "updated_by": uce.UpdatedBy, - "token_type": uce.Type.String(), - "super_admin": uce.SuperAdmin, - "request_id": uce.requestID, - } - - if uce.ID != "" { - val["id"] = uce.ID - } - if uce.FirstName != "" { - val["first_name"] = uce.FirstName - } - if uce.LastName != "" { - val["last_name"] = uce.LastName - } - if len(uce.Tags) > 0 { - val["tags"] = uce.Tags - } - if uce.Credentials.Username != "" { - val["username"] = uce.Credentials.Username - } - if uce.Email != "" { - val["email"] = uce.Email - } - if uce.Metadata != nil { - val["metadata"] = uce.Metadata - } - if uce.PrivateMetadata != nil { - val["private_metadata"] = uce.PrivateMetadata - } - if !uce.CreatedAt.IsZero() { - val["created_at"] = uce.CreatedAt - } - if uce.Status.String() != "" { - val["status"] = uce.Status.String() - } - - return val, nil -} - -type updateUsernameEvent struct { - users.User - authn.Session - requestID string -} - -func (une updateUsernameEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": userUpdateUsername, - "updated_at": une.UpdatedAt, - "updated_by": une.UpdatedBy, - "token_type": une.Type.String(), - "super_admin": une.SuperAdmin, - "request_id": une.requestID, - } - - if une.ID != "" { - val["id"] = une.ID - } - if une.FirstName != "" { - val["first_name"] = une.FirstName - } - if une.LastName != "" { - val["last_name"] = une.LastName - } - if une.Credentials.Username != "" { - val["username"] = une.Credentials.Username - } - - return val, nil -} - -type updateProfilePictureEvent struct { - users.User - authn.Session - requestID string -} - -func (req updateProfilePictureEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": userUpdateProfilePicture, - "updated_at": req.UpdatedAt, - "updated_by": req.UpdatedBy, - "token_type": req.Type.String(), - "super_admin": req.SuperAdmin, - "request_id": req.requestID, - } - - if req.ID != "" { - val["id"] = req.ID - } - if req.ProfilePicture != "" { - val["profile_picture"] = req.ProfilePicture - } - - return val, nil -} - -type changeUserStatusEvent struct { - id string - operation string - status string - updatedAt time.Time - updatedBy string - authn.Session - requestID string -} - -func (rce changeUserStatusEvent) Encode() (map[string]any, error) { - return map[string]any{ - "operation": rce.operation, - "id": rce.id, - "status": rce.status, - "updated_at": rce.updatedAt, - "updated_by": rce.updatedBy, - "token_type": rce.Type.String(), - "super_admin": rce.SuperAdmin, - "request_id": rce.requestID, - }, nil -} - -type viewUserEvent struct { - users.User - authn.Session - requestID string -} - -func (vue viewUserEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": userView, - "id": vue.ID, - "token_type": vue.Type.String(), - "super_admin": vue.SuperAdmin, - "request_id": vue.requestID, - } - - if vue.LastName != "" { - val["last_name"] = vue.LastName - } - if vue.FirstName != "" { - val["first_name"] = vue.FirstName - } - if len(vue.Tags) > 0 { - val["tags"] = vue.Tags - } - if vue.Email != "" { - val["email"] = vue.Email - } - if vue.Credentials.Username != "" { - val["username"] = vue.Credentials.Username - } - if vue.Metadata != nil { - val["metadata"] = vue.Metadata - } - if vue.PrivateMetadata != nil { - val["private_metadata"] = vue.PrivateMetadata - } - if !vue.CreatedAt.IsZero() { - val["created_at"] = vue.CreatedAt - } - if !vue.UpdatedAt.IsZero() { - val["updated_at"] = vue.UpdatedAt - } - if vue.UpdatedBy != "" { - val["updated_by"] = vue.UpdatedBy - } - if vue.Status.String() != "" { - val["status"] = vue.Status.String() - } - - return val, nil -} - -type viewProfileEvent struct { - users.User - authn.Session - requestID string -} - -func (vpe viewProfileEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": profileView, - "id": vpe.ID, - "token_type": vpe.Type.String(), - "super_admin": vpe.SuperAdmin, - "request_id": vpe.requestID, - } - - if vpe.FirstName != "" { - val["first_name"] = vpe.FirstName - } - if len(vpe.Tags) > 0 { - val["tags"] = vpe.Tags - } - if vpe.Credentials.Username != "" { - val["username"] = vpe.Credentials.Username - } - if vpe.Metadata != nil { - val["metadata"] = vpe.Metadata - } - if vpe.PrivateMetadata != nil { - val["private_metadata"] = vpe.PrivateMetadata - } - if !vpe.CreatedAt.IsZero() { - val["created_at"] = vpe.CreatedAt - } - if !vpe.UpdatedAt.IsZero() { - val["updated_at"] = vpe.UpdatedAt - } - if vpe.UpdatedBy != "" { - val["updated_by"] = vpe.UpdatedBy - } - if vpe.Status.String() != "" { - val["status"] = vpe.Status.String() - } - if vpe.Email != "" { - val["email"] = vpe.Email - } - - return val, nil -} - -type listUserEvent struct { - users.Page - authn.Session - requestID string -} - -func (lue listUserEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": userList, - "total": lue.Total, - "offset": lue.Offset, - "limit": lue.Limit, - "token_type": lue.Type.String(), - "super_admin": lue.SuperAdmin, - "request_id": lue.requestID, - } - - if lue.FirstName != "" { - val["first_name"] = lue.FirstName - } - if lue.LastName != "" { - val["last_name"] = lue.LastName - } - if lue.Order != "" { - val["order"] = lue.Order - } - if lue.Dir != "" { - val["dir"] = lue.Dir - } - if lue.Metadata != nil { - val["metadata"] = lue.Metadata - } - if lue.Domain != "" { - val["domain"] = lue.Domain - } - if len(lue.Tags.Elements) > 0 { - val["tags"] = lue.Tags.Elements - } - if lue.Permission != "" { - val["permission"] = lue.Permission - } - if lue.Status.String() != "" { - val["status"] = lue.Status.String() - } - if lue.Username != "" { - val["username"] = lue.Username - } - if lue.Email != "" { - val["email"] = lue.Email - } - - return val, nil -} - -type searchUserEvent struct { - users.Page - requestID string -} - -func (sce searchUserEvent) Encode() (map[string]any, error) { - val := map[string]any{ - "operation": userSearch, - "total": sce.Total, - "offset": sce.Offset, - "limit": sce.Limit, - "request_id": sce.requestID, - } - if sce.Username != "" { - val["username"] = sce.Username - } - if sce.FirstName != "" { - val["first_name"] = sce.FirstName - } - if sce.LastName != "" { - val["last_name"] = sce.LastName - } - if sce.Email != "" { - val["email"] = sce.Email - } - if sce.Id != "" { - val["id"] = sce.Id - } - - return val, nil -} - -type identifyUserEvent struct { - userID string - requestID string -} - -func (ise identifyUserEvent) Encode() (map[string]any, error) { - return map[string]any{ - "operation": userIdentify, - "id": ise.userID, - "request_id": ise.requestID, - }, nil -} - -type issueTokenEvent struct { - username string - requestID string -} - -func (ite issueTokenEvent) Encode() (map[string]any, error) { - return map[string]any{ - "operation": issueToken, - "username": ite.username, - "request_id": ite.requestID, - }, nil -} - -type refreshTokenEvent struct { - requestID string -} - -func (rte refreshTokenEvent) Encode() (map[string]any, error) { - return map[string]any{ - "operation": refreshToken, - "request_id": rte.requestID, - }, nil -} - -type revokeRefreshTokenEvent struct { - tokenID string - requestID string -} - -func (rrte revokeRefreshTokenEvent) Encode() (map[string]any, error) { - return map[string]any{ - "operation": revokeRefreshToken, - "token_id": rrte.tokenID, - "request_id": rrte.requestID, - }, nil -} - -type resetSecretEvent struct { - requestID string -} - -func (rse resetSecretEvent) Encode() (map[string]any, error) { - return map[string]any{ - "operation": resetSecret, - "request_id": rse.requestID, - }, nil -} - -type sendPasswordResetEvent struct { - host string - email string - user string - requestID string -} - -func (req sendPasswordResetEvent) Encode() (map[string]any, error) { - return map[string]any{ - "operation": sendPasswordReset, - "host": req.host, - "email": req.email, - "user": req.user, - "request_id": req.requestID, - }, nil -} - -type oauthCallbackEvent struct { - userID string - requestID string -} - -func (oce oauthCallbackEvent) Encode() (map[string]any, error) { - return map[string]any{ - "operation": oauthCallback, - "user_id": oce.userID, - "request_id": oce.requestID, - }, nil -} - -type deleteUserEvent struct { - id string - authn.Session - requestID string -} - -func (dce deleteUserEvent) Encode() (map[string]any, error) { - return map[string]any{ - "operation": deleteUser, - "id": dce.id, - "token_type": dce.Type.String(), - "super_admin": dce.SuperAdmin, - "request_id": dce.requestID, - }, nil -} - -type addUserPolicyEvent struct { - id string - role string - authn.Session - requestID string -} - -func (req addUserPolicyEvent) Encode() (map[string]any, error) { - return map[string]any{ - "operation": addClientPolicy, - "id": req.id, - "role": req.role, - "token_type": req.Type.String(), - "super_admin": req.SuperAdmin, - "request_id": req.requestID, - }, nil -} diff --git a/users/events/streams.go b/users/events/streams.go deleted file mode 100644 index 0b29338e0..000000000 --- a/users/events/streams.go +++ /dev/null @@ -1,463 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package events - -import ( - "context" - - grpcTokenV1 "github.com/absmach/magistrala/api/grpc/token/v1" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/events" - "github.com/absmach/magistrala/pkg/events/store" - "github.com/absmach/magistrala/users" - "github.com/go-chi/chi/v5/middleware" -) - -const ( - magistralaPrefix = "magistrala." - createStream = magistralaPrefix + userCreate - sendVerificationStream = magistralaPrefix + userSendVerification - verifyEmailStream = magistralaPrefix + userVerifyEmail - updateStream = magistralaPrefix + userUpdate - updateRoleStream = magistralaPrefix + userUpdateRole - updateTagsStream = magistralaPrefix + userUpdateTags - updateSecretStream = magistralaPrefix + userUpdateSecret - updateUsernameStream = magistralaPrefix + userUpdateUsername - updatePictureStream = magistralaPrefix + userUpdateProfilePicture - UpdateEmailStream = magistralaPrefix + userUpdateEmail - enableStream = magistralaPrefix + userEnable - disableStream = magistralaPrefix + userDisable - viewStream = magistralaPrefix + userView - viewProfileStream = magistralaPrefix + profileView - listStream = magistralaPrefix + userList - searchStream = magistralaPrefix + userSearch - identifyStream = magistralaPrefix + userIdentify - issueTokenStream = magistralaPrefix + issueToken - refreshTokenStream = magistralaPrefix + refreshToken - revokeRefreshTokenStream = magistralaPrefix + revokeRefreshToken - resetSecretStream = magistralaPrefix + resetSecret - sendPasswordResetStream = magistralaPrefix + sendPasswordReset - oauthStream = magistralaPrefix + oauthCallback - addPolicyStream = magistralaPrefix + addClientPolicy - deleteStream = magistralaPrefix + deleteUser -) - -var _ users.Service = (*eventStore)(nil) - -type eventStore struct { - events.Publisher - svc users.Service -} - -// NewEventStoreMiddleware returns wrapper around users service that sends -// events to event store. -func NewEventStoreMiddleware(ctx context.Context, svc users.Service, url string) (users.Service, error) { - publisher, err := store.NewPublisher(ctx, url, "users-es-pub") - if err != nil { - return nil, err - } - - return &eventStore{ - svc: svc, - Publisher: publisher, - }, nil -} - -func (es *eventStore) Register(ctx context.Context, session authn.Session, user users.User, selfRegister bool) (users.User, error) { - user, err := es.svc.Register(ctx, session, user, selfRegister) - if err != nil { - return user, err - } - - event := createUserEvent{ - user, - session, - middleware.GetReqID(ctx), - } - - if err := es.Publish(ctx, createStream, event); err != nil { - return user, err - } - - return user, nil -} - -func (es *eventStore) SendVerification(ctx context.Context, session authn.Session) error { - err := es.svc.SendVerification(ctx, session) - if err != nil { - return err - } - - event := sendVerificationEvent{ - session, - middleware.GetReqID(ctx), - } - - return es.Publish(ctx, sendVerificationStream, event) -} - -func (es *eventStore) VerifyEmail(ctx context.Context, verificationToken string) (users.User, error) { - user, err := es.svc.VerifyEmail(ctx, verificationToken) - if err != nil { - return user, err - } - - event := verifyEmailEvent{ - email: user.Email, - userID: user.ID, - verifiedAt: user.VerifiedAt, - requestID: middleware.GetReqID(ctx), - } - if err := es.Publish(ctx, verifyEmailStream, event); err != nil { - return user, err - } - return user, nil -} - -func (es *eventStore) Update(ctx context.Context, session authn.Session, id string, usr users.UserReq) (users.User, error) { - user, err := es.svc.Update(ctx, session, id, usr) - if err != nil { - return user, err - } - - return es.update(ctx, session, userUpdate, updateStream, user) -} - -func (es *eventStore) UpdateRole(ctx context.Context, session authn.Session, user users.User) (users.User, error) { - user, err := es.svc.UpdateRole(ctx, session, user) - if err != nil { - return user, err - } - - return es.update(ctx, session, userUpdateRole, updateRoleStream, user) -} - -func (es *eventStore) UpdateTags(ctx context.Context, session authn.Session, id string, usr users.UserReq) (users.User, error) { - user, err := es.svc.UpdateTags(ctx, session, id, usr) - if err != nil { - return user, err - } - - return es.update(ctx, session, userUpdateTags, updateTagsStream, user) -} - -func (es *eventStore) UpdateSecret(ctx context.Context, session authn.Session, oldSecret, newSecret string) (users.User, error) { - user, err := es.svc.UpdateSecret(ctx, session, oldSecret, newSecret) - if err != nil { - return user, err - } - - return es.update(ctx, session, userUpdateSecret, updateSecretStream, user) -} - -func (es *eventStore) UpdateUsername(ctx context.Context, session authn.Session, id, username string) (users.User, error) { - user, err := es.svc.UpdateUsername(ctx, session, id, username) - if err != nil { - return user, err - } - - event := updateUsernameEvent{ - user, - session, - middleware.GetReqID(ctx), - } - - if err := es.Publish(ctx, updateUsernameStream, event); err != nil { - return user, err - } - - return user, nil -} - -func (es *eventStore) UpdateProfilePicture(ctx context.Context, session authn.Session, id string, usr users.UserReq) (users.User, error) { - user, err := es.svc.UpdateProfilePicture(ctx, session, id, usr) - if err != nil { - return user, err - } - - event := updateProfilePictureEvent{ - user, - session, - middleware.GetReqID(ctx), - } - - if err := es.Publish(ctx, updatePictureStream, event); err != nil { - return user, err - } - - return user, nil -} - -func (es *eventStore) UpdateEmail(ctx context.Context, session authn.Session, id, email string) (users.User, error) { - user, err := es.svc.UpdateEmail(ctx, session, id, email) - if err != nil { - return user, err - } - - return es.update(ctx, session, userUpdateEmail, UpdateEmailStream, user) -} - -func (es *eventStore) update(ctx context.Context, session authn.Session, operation, stream string, user users.User) (users.User, error) { - event := updateUserEvent{ - user, operation, session, middleware.GetReqID(ctx), - } - - if err := es.Publish(ctx, stream, event); err != nil { - return user, err - } - - return user, nil -} - -func (es *eventStore) View(ctx context.Context, session authn.Session, id string) (users.User, error) { - user, err := es.svc.View(ctx, session, id) - if err != nil { - return user, err - } - - event := viewUserEvent{ - user, - session, - middleware.GetReqID(ctx), - } - - if err := es.Publish(ctx, viewStream, event); err != nil { - return user, err - } - - return user, nil -} - -func (es *eventStore) ViewProfile(ctx context.Context, session authn.Session) (users.User, error) { - user, err := es.svc.ViewProfile(ctx, session) - if err != nil { - return user, err - } - - event := viewProfileEvent{ - user, - session, - middleware.GetReqID(ctx), - } - - if err := es.Publish(ctx, viewProfileStream, event); err != nil { - return user, err - } - - return user, nil -} - -func (es *eventStore) ListUsers(ctx context.Context, session authn.Session, pm users.Page) (users.UsersPage, error) { - cp, err := es.svc.ListUsers(ctx, session, pm) - if err != nil { - return cp, err - } - event := listUserEvent{ - pm, - session, - middleware.GetReqID(ctx), - } - - if err := es.Publish(ctx, listStream, event); err != nil { - return cp, err - } - - return cp, nil -} - -func (es *eventStore) SearchUsers(ctx context.Context, pm users.Page) (users.UsersPage, error) { - cp, err := es.svc.SearchUsers(ctx, pm) - if err != nil { - return cp, err - } - event := searchUserEvent{ - pm, - middleware.GetReqID(ctx), - } - - if err := es.Publish(ctx, searchStream, event); err != nil { - return cp, err - } - - return cp, nil -} - -func (es *eventStore) Enable(ctx context.Context, session authn.Session, id string) (users.User, error) { - user, err := es.svc.Enable(ctx, session, id) - if err != nil { - return user, err - } - - return es.changeStatus(ctx, session, userEnable, enableStream, user) -} - -func (es *eventStore) Disable(ctx context.Context, session authn.Session, id string) (users.User, error) { - user, err := es.svc.Disable(ctx, session, id) - if err != nil { - return user, err - } - - return es.changeStatus(ctx, session, userDisable, disableStream, user) -} - -func (es *eventStore) changeStatus(ctx context.Context, session authn.Session, operation, stream string, user users.User) (users.User, error) { - event := changeUserStatusEvent{ - id: user.ID, - operation: operation, - updatedAt: user.UpdatedAt, - updatedBy: user.UpdatedBy, - status: user.Status.String(), - Session: session, - requestID: middleware.GetReqID(ctx), - } - - if err := es.Publish(ctx, stream, event); err != nil { - return user, err - } - - return user, nil -} - -func (es *eventStore) Identify(ctx context.Context, session authn.Session) (string, error) { - userID, err := es.svc.Identify(ctx, session) - if err != nil { - return userID, err - } - - event := identifyUserEvent{ - userID: userID, - requestID: middleware.GetReqID(ctx), - } - - if err := es.Publish(ctx, identifyStream, event); err != nil { - return userID, err - } - - return userID, nil -} - -func (es *eventStore) SendPasswordReset(ctx context.Context, email string) error { - err := es.svc.SendPasswordReset(ctx, email) - if err != nil { - return err - } - - event := sendPasswordResetEvent{ - email: email, - requestID: middleware.GetReqID(ctx), - } - - return es.Publish(ctx, sendPasswordResetStream, event) -} - -func (es *eventStore) IssueToken(ctx context.Context, username, secret, description string) (*grpcTokenV1.Token, error) { - token, err := es.svc.IssueToken(ctx, username, secret, description) - if err != nil { - return token, err - } - - event := issueTokenEvent{ - username: username, - requestID: middleware.GetReqID(ctx), - } - - if err := es.Publish(ctx, issueTokenStream, event); err != nil { - return token, err - } - - return token, nil -} - -func (es *eventStore) RefreshToken(ctx context.Context, session authn.Session, refreshToken string) (*grpcTokenV1.Token, error) { - token, err := es.svc.RefreshToken(ctx, session, refreshToken) - if err != nil { - return token, err - } - - event := refreshTokenEvent{ - requestID: middleware.GetReqID(ctx), - } - - if err := es.Publish(ctx, refreshTokenStream, event); err != nil { - return token, err - } - - return token, nil -} - -func (es *eventStore) RevokeRefreshToken(ctx context.Context, session authn.Session, tokenID string) error { - err := es.svc.RevokeRefreshToken(ctx, session, tokenID) - if err != nil { - return err - } - - event := revokeRefreshTokenEvent{ - tokenID: tokenID, - requestID: middleware.GetReqID(ctx), - } - - return es.Publish(ctx, revokeRefreshTokenStream, event) -} - -func (es *eventStore) ListActiveRefreshTokens(ctx context.Context, session authn.Session) (*grpcTokenV1.ListUserRefreshTokensRes, error) { - return es.svc.ListActiveRefreshTokens(ctx, session) -} - -func (es *eventStore) ResetSecret(ctx context.Context, session authn.Session, secret string) error { - if err := es.svc.ResetSecret(ctx, session, secret); err != nil { - return err - } - - event := resetSecretEvent{ - requestID: middleware.GetReqID(ctx), - } - - return es.Publish(ctx, resetSecretStream, event) -} - -func (es *eventStore) OAuthCallback(ctx context.Context, user users.User) (users.User, error) { - token, err := es.svc.OAuthCallback(ctx, user) - if err != nil { - return token, err - } - - event := oauthCallbackEvent{ - userID: user.ID, - requestID: middleware.GetReqID(ctx), - } - - if err := es.Publish(ctx, oauthStream, event); err != nil { - return token, err - } - - return token, nil -} - -func (es *eventStore) Delete(ctx context.Context, session authn.Session, id string) error { - if err := es.svc.Delete(ctx, session, id); err != nil { - return err - } - - event := deleteUserEvent{ - id: id, - Session: session, - requestID: middleware.GetReqID(ctx), - } - - return es.Publish(ctx, deleteStream, event) -} - -func (es *eventStore) OAuthAddUserPolicy(ctx context.Context, user users.User) error { - if err := es.svc.OAuthAddUserPolicy(ctx, user); err != nil { - return err - } - - event := addUserPolicyEvent{ - id: user.ID, - role: user.Role.String(), - requestID: middleware.GetReqID(ctx), - } - - return es.Publish(ctx, addPolicyStream, event) -} diff --git a/users/events/streams_test.go b/users/events/streams_test.go deleted file mode 100644 index 179ba5d9d..000000000 --- a/users/events/streams_test.go +++ /dev/null @@ -1,1258 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package events_test - -import ( - "context" - "fmt" - "os" - "testing" - "time" - - grpcTokenV1 "github.com/absmach/magistrala/api/grpc/token/v1" - "github.com/absmach/magistrala/internal/testsutil" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/users" - "github.com/absmach/magistrala/users/events" - "github.com/absmach/magistrala/users/mocks" - "github.com/go-chi/chi/v5/middleware" - "github.com/redis/go-redis/v9" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" -) - -var ( - storeClient *redis.Client - storeURL string - validSession = authn.Session{ - UserID: testsutil.GenerateUUID(&testing.T{}), - } - validUser = generateTestUser(&testing.T{}) - validUsersPage = users.UsersPage{ - Page: users.Page{ - Limit: 10, - Offset: 0, - Total: 1, - }, - Users: []users.User{validUser}, - } -) - -func newEventStoreMiddleware(t *testing.T) (*mocks.Service, users.Service) { - svc := new(mocks.Service) - nsvc, err := events.NewEventStoreMiddleware(context.Background(), svc, storeURL) - require.Nil(t, err, fmt.Sprintf("create events store middleware failed with unexpected error: %s", err)) - - return svc, nsvc -} - -func TestMain(m *testing.M) { - code := testsutil.RunRedisTest(m, &storeClient, &storeURL) - os.Exit(code) -} - -func TestRegister(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validID := testsutil.GenerateUUID(t) - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, validID) - - cases := []struct { - desc string - session authn.Session - user users.User - selfRegister bool - svcRes users.User - svcErr error - resp users.User - err error - }{ - { - desc: "publish successfully", - session: validSession, - user: validUser, - selfRegister: true, - svcRes: validUser, - svcErr: nil, - resp: validUser, - err: nil, - }, - { - desc: "failed to pusblish with service error", - session: validSession, - user: validUser, - selfRegister: true, - svcRes: users.User{}, - svcErr: svcerr.ErrCreateEntity, - resp: users.User{}, - err: svcerr.ErrCreateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("Register", validCtx, tc.session, tc.user, tc.selfRegister).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.Register(validCtx, tc.session, tc.user, tc.selfRegister) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestSendVerification(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - cases := []struct { - desc string - session authn.Session - userID string - svcErr error - err error - }{ - { - desc: "publish successfully", - session: validSession, - userID: validUser.ID, - svcErr: nil, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - userID: validUser.ID, - svcErr: svcerr.ErrCreateEntity, - err: svcerr.ErrCreateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("SendVerification", validCtx, tc.session).Return(tc.svcErr) - err := nsvc.SendVerification(validCtx, tc.session) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - svcCall.Unset() - }) - } -} - -func TestVerifyEmail(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - validToken := "validVerificationToken" - cases := []struct { - desc string - verificationToken string - svcRes users.User - svcErr error - resp users.User - err error - }{ - { - desc: "publish successfully", - verificationToken: validToken, - svcRes: validUser, - svcErr: nil, - resp: validUser, - err: nil, - }, - { - desc: "failed to publish with service error", - verificationToken: validToken, - svcRes: users.User{}, - svcErr: svcerr.ErrCreateEntity, - resp: users.User{}, - err: svcerr.ErrCreateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("VerifyEmail", validCtx, tc.verificationToken).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.VerifyEmail(validCtx, tc.verificationToken) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestUpdate(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - updatedUser := validUser - updatedUser.FirstName = "updatedFirstName" - - cases := []struct { - desc string - session authn.Session - userID string - userReq users.UserReq - svcRes users.User - svcErr error - resp users.User - err error - }{ - { - desc: "publish successfully", - session: validSession, - userID: validSession.UserID, - userReq: users.UserReq{ - FirstName: &updatedUser.FirstName, - }, - svcRes: updatedUser, - svcErr: nil, - resp: updatedUser, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - userID: validSession.UserID, - userReq: users.UserReq{ - FirstName: &updatedUser.FirstName, - }, - svcRes: users.User{}, - svcErr: svcerr.ErrUpdateEntity, - resp: users.User{}, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("Update", validCtx, tc.session, tc.userID, tc.userReq).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.Update(validCtx, tc.session, tc.userID, tc.userReq) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestUpdateRole(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - updatedUser := validUser - updatedUser.Role = users.AdminRole - - cases := []struct { - desc string - session authn.Session - user users.User - svcRes users.User - svcErr error - resp users.User - err error - }{ - { - desc: "publish successfully", - session: validSession, - user: updatedUser, - svcRes: updatedUser, - svcErr: nil, - resp: updatedUser, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - user: updatedUser, - svcRes: users.User{}, - svcErr: svcerr.ErrUpdateEntity, - resp: users.User{}, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("UpdateRole", validCtx, tc.session, tc.user).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.UpdateRole(validCtx, tc.session, tc.user) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestUpdateTags(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - updatedUser := validUser - updatedUser.Tags = []string{"newTag1", "newTag2"} - - cases := []struct { - desc string - session authn.Session - userID string - userReq users.UserReq - svcRes users.User - svcErr error - resp users.User - err error - }{ - { - desc: "publish successfully", - session: validSession, - userID: validSession.UserID, - userReq: users.UserReq{ - Tags: &updatedUser.Tags, - }, - svcRes: updatedUser, - svcErr: nil, - resp: updatedUser, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - userID: validSession.UserID, - userReq: users.UserReq{ - Tags: &updatedUser.Tags, - }, - svcRes: users.User{}, - svcErr: svcerr.ErrUpdateEntity, - resp: users.User{}, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("UpdateTags", validCtx, tc.session, tc.userID, tc.userReq).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.UpdateTags(validCtx, tc.session, tc.userID, tc.userReq) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestUpdateSecret(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - updatedUser := validUser - updatedUser.Credentials.Secret = "newSecret" - - cases := []struct { - desc string - session authn.Session - oldSecret string - newSecret string - svcRes users.User - svcErr error - resp users.User - err error - }{ - { - desc: "publish successfully", - session: validSession, - oldSecret: "secret", - newSecret: "newSecret", - svcRes: updatedUser, - svcErr: nil, - resp: updatedUser, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - oldSecret: "secret", - newSecret: "newSecret", - svcRes: users.User{}, - svcErr: svcerr.ErrUpdateEntity, - resp: users.User{}, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("UpdateSecret", validCtx, tc.session, tc.oldSecret, tc.newSecret).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.UpdateSecret(validCtx, tc.session, tc.oldSecret, tc.newSecret) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestUpdateUsername(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - updatedUser := validUser - updatedUser.Credentials.Username = "newUsername" - - cases := []struct { - desc string - session authn.Session - userID string - newUsername string - svcRes users.User - svcErr error - resp users.User - err error - }{ - { - desc: "publish successfully", - session: validSession, - userID: validSession.UserID, - newUsername: "newUsername", - svcRes: updatedUser, - svcErr: nil, - resp: updatedUser, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - userID: validSession.UserID, - newUsername: "newUsername", - svcRes: users.User{}, - svcErr: svcerr.ErrUpdateEntity, - resp: users.User{}, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("UpdateUsername", validCtx, tc.session, tc.userID, tc.newUsername).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.UpdateUsername(validCtx, tc.session, tc.userID, tc.newUsername) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestUpdateProfilePicture(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - updatedUser := validUser - updatedUser.ProfilePicture = "https://example.com/newprofilepic.jpg" - - cases := []struct { - desc string - session authn.Session - userID string - userReq users.UserReq - svcRes users.User - svcErr error - resp users.User - err error - }{ - { - desc: "publish successfully", - session: validSession, - userID: validSession.UserID, - userReq: users.UserReq{ - ProfilePicture: &updatedUser.ProfilePicture, - }, - svcRes: updatedUser, - svcErr: nil, - resp: updatedUser, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - userID: validSession.UserID, - userReq: users.UserReq{ - ProfilePicture: &updatedUser.ProfilePicture, - }, - svcRes: users.User{}, - svcErr: svcerr.ErrUpdateEntity, - resp: users.User{}, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("UpdateProfilePicture", validCtx, tc.session, tc.userID, tc.userReq).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.UpdateProfilePicture(validCtx, tc.session, tc.userID, tc.userReq) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestUpdateEmail(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - updatedUser := validUser - updatedUser.Email = "updatedemail@example.com" - - cases := []struct { - desc string - session authn.Session - userID string - newEmail string - svcRes users.User - svcErr error - resp users.User - err error - }{ - { - desc: "publish successfully", - session: validSession, - userID: validSession.UserID, - newEmail: "updatedemail@example.com", - svcRes: updatedUser, - svcErr: nil, - resp: updatedUser, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - userID: validSession.UserID, - newEmail: "updatedemail@example.com", - svcRes: users.User{}, - svcErr: svcerr.ErrUpdateEntity, - resp: users.User{}, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("UpdateEmail", validCtx, tc.session, tc.userID, tc.newEmail).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.UpdateEmail(validCtx, tc.session, tc.userID, tc.newEmail) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestView(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - userID string - svcRes users.User - svcErr error - resp users.User - err error - }{ - { - desc: "publish successfully", - session: validSession, - userID: validSession.UserID, - svcRes: validUser, - svcErr: nil, - resp: validUser, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - userID: validSession.UserID, - svcRes: users.User{}, - svcErr: svcerr.ErrViewEntity, - resp: users.User{}, - err: svcerr.ErrViewEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("View", validCtx, tc.session, tc.userID).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.View(validCtx, tc.session, tc.userID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestViewProfile(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - svcRes users.User - svcErr error - resp users.User - err error - }{ - { - desc: "publish successfully", - session: validSession, - svcRes: validUser, - svcErr: nil, - resp: validUser, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - svcRes: users.User{}, - svcErr: svcerr.ErrViewEntity, - resp: users.User{}, - err: svcerr.ErrViewEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("ViewProfile", validCtx, tc.session).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.ViewProfile(validCtx, tc.session) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestListUsers(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - pageMeta users.Page - svcRes users.UsersPage - svcErr error - resp users.UsersPage - err error - }{ - { - desc: "publish successfully", - session: validSession, - pageMeta: users.Page{ - Limit: 10, - Offset: 0, - }, - svcRes: validUsersPage, - svcErr: nil, - resp: validUsersPage, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - pageMeta: users.Page{ - Limit: 10, - Offset: 0, - }, - svcRes: users.UsersPage{}, - svcErr: svcerr.ErrViewEntity, - resp: users.UsersPage{}, - err: svcerr.ErrViewEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("ListUsers", validCtx, tc.session, tc.pageMeta).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.ListUsers(validCtx, tc.session, tc.pageMeta) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestSearchUsers(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - pageMeta users.Page - svcRes users.UsersPage - svcErr error - resp users.UsersPage - err error - }{ - { - desc: "publish successfully", - pageMeta: users.Page{ - Limit: 10, - Offset: 0, - }, - svcRes: validUsersPage, - svcErr: nil, - resp: validUsersPage, - err: nil, - }, - { - desc: "failed to publish with service error", - pageMeta: users.Page{ - Limit: 10, - Offset: 0, - }, - svcRes: users.UsersPage{}, - svcErr: svcerr.ErrViewEntity, - resp: users.UsersPage{}, - err: svcerr.ErrViewEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("SearchUsers", validCtx, tc.pageMeta).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.SearchUsers(validCtx, tc.pageMeta) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestEnable(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - userID string - svcRes users.User - svcErr error - resp users.User - err error - }{ - { - desc: "publish successfully", - session: validSession, - userID: validSession.UserID, - svcRes: validUser, - svcErr: nil, - resp: validUser, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - userID: validSession.UserID, - svcRes: users.User{}, - svcErr: svcerr.ErrUpdateEntity, - resp: users.User{}, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("Enable", validCtx, tc.session, tc.userID).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.Enable(validCtx, tc.session, tc.userID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestDisable(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - cases := []struct { - desc string - session authn.Session - userID string - svcRes users.User - svcErr error - resp users.User - err error - }{ - { - desc: "publish successfully", - session: validSession, - userID: validSession.UserID, - svcRes: validUser, - svcErr: nil, - resp: validUser, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - userID: validSession.UserID, - svcRes: users.User{}, - svcErr: svcerr.ErrUpdateEntity, - resp: users.User{}, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("Disable", validCtx, tc.session, tc.userID).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.Disable(validCtx, tc.session, tc.userID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestIdentify(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - svcRes string - svcErr error - resp string - err error - }{ - { - desc: "publish successfully", - session: validSession, - svcRes: validUser.ID, - svcErr: nil, - resp: validUser.ID, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - svcRes: "", - svcErr: svcerr.ErrViewEntity, - resp: "", - err: svcerr.ErrViewEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("Identify", validCtx, tc.session).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.Identify(validCtx, tc.session) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestSendPasswordReset(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - cases := []struct { - desc string - email string - svcErr error - err error - }{ - { - desc: "publish successfully", - email: validUser.Email, - svcErr: nil, - err: nil, - }, - { - desc: "failed to publish with service error", - email: validUser.Email, - svcErr: svcerr.ErrCreateEntity, - err: svcerr.ErrCreateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("SendPasswordReset", validCtx, tc.email).Return(tc.svcErr) - err := nsvc.SendPasswordReset(validCtx, tc.email) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - svcCall.Unset() - }) - } -} - -func TestIssueToken(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - validToken := &grpcTokenV1.Token{ - AccessToken: "validAccessToken", - } - - cases := []struct { - desc string - username string - secret string - description string - svcRes *grpcTokenV1.Token - svcErr error - resp *grpcTokenV1.Token - err error - }{ - { - desc: "publish successfully", - username: validUser.Credentials.Username, - secret: validUser.Credentials.Secret, - description: "valid token", - svcRes: validToken, - svcErr: nil, - resp: validToken, - err: nil, - }, - { - desc: "failed to publish with service error", - username: validUser.Credentials.Username, - secret: validUser.Credentials.Secret, - svcRes: nil, - svcErr: svcerr.ErrCreateEntity, - resp: nil, - err: svcerr.ErrCreateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("IssueToken", validCtx, tc.username, tc.secret, tc.description).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.IssueToken(validCtx, tc.username, tc.secret, tc.description) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestRefreshToken(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - validRefreshToken := "validRefreshToken" - validToken := &grpcTokenV1.Token{ - AccessToken: "validAccessToken", - } - - cases := []struct { - desc string - session authn.Session - refreshToken string - svcRes *grpcTokenV1.Token - svcErr error - resp *grpcTokenV1.Token - err error - }{ - { - desc: "publish successfully", - session: validSession, - refreshToken: validRefreshToken, - svcRes: validToken, - svcErr: nil, - resp: validToken, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - refreshToken: validRefreshToken, - svcRes: nil, - svcErr: svcerr.ErrCreateEntity, - resp: nil, - err: svcerr.ErrCreateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("RefreshToken", validCtx, tc.session, tc.refreshToken).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.RefreshToken(validCtx, tc.session, tc.refreshToken) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestRevokeRefreshToken(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - validTokenID := "validTokenID" - - cases := []struct { - desc string - session authn.Session - tokenID string - svcErr error - err error - }{ - { - desc: "publish successfully", - session: validSession, - tokenID: validTokenID, - svcErr: nil, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - tokenID: validTokenID, - svcErr: svcerr.ErrUpdateEntity, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("RevokeRefreshToken", validCtx, tc.session, tc.tokenID).Return(tc.svcErr) - err := nsvc.RevokeRefreshToken(validCtx, tc.session, tc.tokenID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - svcCall.Unset() - }) - } -} - -func TestListActiveRefreshTokens(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - validTokensList := &grpcTokenV1.ListUserRefreshTokensRes{ - RefreshTokens: []*grpcTokenV1.RefreshToken{ - {Id: "token1", Description: "token1"}, - {Id: "token2", Description: "token2"}, - }, - } - - cases := []struct { - desc string - session authn.Session - svcRes *grpcTokenV1.ListUserRefreshTokensRes - svcErr error - resp *grpcTokenV1.ListUserRefreshTokensRes - err error - }{ - { - desc: "publish successfully", - session: validSession, - svcRes: validTokensList, - svcErr: nil, - resp: validTokensList, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - svcRes: nil, - svcErr: svcerr.ErrViewEntity, - resp: nil, - err: svcerr.ErrViewEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("ListActiveRefreshTokens", validCtx, tc.session).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.ListActiveRefreshTokens(validCtx, tc.session) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestResetSecret(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - newSecret := "newSecret" - - cases := []struct { - desc string - session authn.Session - secret string - svcErr error - err error - }{ - { - desc: "publish successfully", - session: validSession, - secret: newSecret, - svcErr: nil, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - secret: newSecret, - svcErr: svcerr.ErrUpdateEntity, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("ResetSecret", validCtx, tc.session, tc.secret).Return(tc.svcErr) - err := nsvc.ResetSecret(validCtx, tc.session, tc.secret) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - svcCall.Unset() - }) - } -} - -func TestOAuthCallback(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - user users.User - svcRes users.User - svcErr error - resp users.User - err error - }{ - { - desc: "publish successfully", - user: validUser, - svcRes: validUser, - svcErr: nil, - resp: validUser, - err: nil, - }, - { - desc: "failed to publish with service error", - user: validUser, - svcRes: users.User{}, - svcErr: svcerr.ErrCreateEntity, - resp: users.User{}, - err: svcerr.ErrCreateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("OAuthCallback", validCtx, tc.user).Return(tc.svcRes, tc.svcErr) - resp, err := nsvc.OAuthCallback(validCtx, tc.user) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.resp, resp, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.resp, resp)) - svcCall.Unset() - }) - } -} - -func TestDelete(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - session authn.Session - userID string - svcErr error - err error - }{ - { - desc: "publish successfully", - session: validSession, - userID: validSession.UserID, - svcErr: nil, - err: nil, - }, - { - desc: "failed to publish with service error", - session: validSession, - userID: validSession.UserID, - svcErr: svcerr.ErrRemoveEntity, - err: svcerr.ErrRemoveEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("Delete", validCtx, tc.session, tc.userID).Return(tc.svcErr) - err := nsvc.Delete(validCtx, tc.session, tc.userID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - svcCall.Unset() - }) - } -} - -func TestOAuthAddUserPolicy(t *testing.T) { - svc, nsvc := newEventStoreMiddleware(t) - - validCtx := context.WithValue(context.Background(), middleware.RequestIDKey, testsutil.GenerateUUID(t)) - - cases := []struct { - desc string - user users.User - svcErr error - err error - }{ - { - desc: "publish successfully", - user: validUser, - svcErr: nil, - err: nil, - }, - { - desc: "failed to publish with service error", - user: validUser, - svcErr: svcerr.ErrCreateEntity, - err: svcerr.ErrCreateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - svcCall := svc.On("OAuthAddUserPolicy", validCtx, tc.user).Return(tc.svcErr) - err := nsvc.OAuthAddUserPolicy(validCtx, tc.user) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - svcCall.Unset() - }) - } -} - -func generateTestUser(t *testing.T) users.User { - createdAt, err := time.Parse(time.RFC3339, "2024-01-01T00:00:00Z") - assert.Nil(t, err, fmt.Sprintf("Unexpected error parsing time: %v", err)) - return users.User{ - ID: testsutil.GenerateUUID(t), - FirstName: "userfirstname", - LastName: "userlastname", - Email: "useremail@example.com", - Credentials: users.Credentials{ - Username: "username", - Secret: "secret", - }, - Tags: []string{"tag1", "tag2"}, - PrivateMetadata: users.Metadata{ - "key1": "value1", - "key2": "value2", - }, - CreatedAt: createdAt, - UpdatedAt: createdAt, - Status: users.EnabledStatus, - Role: users.UserRole, - } -} diff --git a/users/hasher.go b/users/hasher.go deleted file mode 100644 index 073ecfbbc..000000000 --- a/users/hasher.go +++ /dev/null @@ -1,15 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package users - -// Hasher specifies an API for generating hashes of an arbitrary textual -// content. -type Hasher interface { - // Hash generates the hashed string from plain-text. - Hash(string) (string, error) - - // Compare compares plain-text version to the hashed one. An error should - // indicate failed comparison. - Compare(string, string) error -} diff --git a/users/hasher/doc.go b/users/hasher/doc.go deleted file mode 100644 index 98be99226..000000000 --- a/users/hasher/doc.go +++ /dev/null @@ -1,6 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package hasher contains the domain concept definitions needed to -// support Magistrala users password hasher sub-service functionality. -package hasher diff --git a/users/hasher/hasher.go b/users/hasher/hasher.go deleted file mode 100644 index dc2084a70..000000000 --- a/users/hasher/hasher.go +++ /dev/null @@ -1,43 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package hasher - -import ( - "github.com/absmach/magistrala/pkg/errors" - "github.com/absmach/magistrala/users" - "golang.org/x/crypto/bcrypt" -) - -const cost int = 10 - -var ( - errHashPassword = errors.NewServiceError("generate hash from password failed") - errComparePassword = errors.NewServiceError("compare hash and password failed") -) - -var _ users.Hasher = (*bcryptHasher)(nil) - -type bcryptHasher struct{} - -// New instantiates a bcrypt-based hasher implementation. -func New() users.Hasher { - return &bcryptHasher{} -} - -func (bh *bcryptHasher) Hash(pwd string) (string, error) { - hash, err := bcrypt.GenerateFromPassword([]byte(pwd), cost) - if err != nil { - return "", errors.Wrap(errHashPassword, err) - } - - return string(hash), nil -} - -func (bh *bcryptHasher) Compare(plain, hashed string) error { - if err := bcrypt.CompareHashAndPassword([]byte(hashed), []byte(plain)); err != nil { - return errors.Wrap(errComparePassword, err) - } - - return nil -} diff --git a/users/middleware/authorization.go b/users/middleware/authorization.go deleted file mode 100644 index e68c3845f..000000000 --- a/users/middleware/authorization.go +++ /dev/null @@ -1,280 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package middleware - -import ( - "context" - - grpcTokenV1 "github.com/absmach/magistrala/api/grpc/token/v1" - "github.com/absmach/magistrala/auth" - "github.com/absmach/magistrala/pkg/authn" - smqauthz "github.com/absmach/magistrala/pkg/authz" - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/pkg/policies" - "github.com/absmach/magistrala/users" -) - -var _ users.Service = (*authorizationMiddleware)(nil) - -type authorizationMiddleware struct { - svc users.Service - authz smqauthz.Authorization - selfRegister bool -} - -// NewAuthorization adds authorization to the users service. -func NewAuthorization(svc users.Service, authz smqauthz.Authorization, selfRegister bool) users.Service { - return &authorizationMiddleware{svc: svc, authz: authz, selfRegister: selfRegister} -} - -func (am *authorizationMiddleware) SendVerification(ctx context.Context, session authn.Session) error { - return am.svc.SendVerification(ctx, session) -} - -func (am *authorizationMiddleware) VerifyEmail(ctx context.Context, verificationToken string) (users.User, error) { - return am.svc.VerifyEmail(ctx, verificationToken) -} - -func (am *authorizationMiddleware) Register(ctx context.Context, session authn.Session, user users.User, selfRegister bool) (users.User, error) { - if selfRegister { - switch err := am.checkSuperAdmin(ctx, session); { - case err == nil: - session.SuperAdmin = true - case errors.Contains(err, svcerr.ErrSuperAdminAction): - default: - return users.User{}, err - } - } - - return am.svc.Register(ctx, session, user, selfRegister) -} - -func (am *authorizationMiddleware) View(ctx context.Context, session authn.Session, id string) (users.User, error) { - switch err := am.checkSuperAdmin(ctx, session); { - case err == nil: - session.SuperAdmin = true - case errors.Contains(err, svcerr.ErrSuperAdminAction): - default: - return users.User{}, err - } - - return am.svc.View(ctx, session, id) -} - -func (am *authorizationMiddleware) ViewProfile(ctx context.Context, session authn.Session) (users.User, error) { - return am.svc.ViewProfile(ctx, session) -} - -func (am *authorizationMiddleware) ListUsers(ctx context.Context, session authn.Session, pm users.Page) (users.UsersPage, error) { - switch err := am.checkSuperAdmin(ctx, session); { - case err == nil: - session.SuperAdmin = true - case errors.Contains(err, svcerr.ErrSuperAdminAction): - default: - return users.UsersPage{}, err - } - - return am.svc.ListUsers(ctx, session, pm) -} - -func (am *authorizationMiddleware) SearchUsers(ctx context.Context, pm users.Page) (users.UsersPage, error) { - return am.svc.SearchUsers(ctx, pm) -} - -func (am *authorizationMiddleware) Update(ctx context.Context, session authn.Session, id string, user users.UserReq) (users.User, error) { - switch err := am.checkSuperAdmin(ctx, session); { - case err == nil: - session.SuperAdmin = true - case errors.Contains(err, svcerr.ErrSuperAdminAction): - default: - return users.User{}, err - } - - return am.svc.Update(ctx, session, id, user) -} - -func (am *authorizationMiddleware) UpdateTags(ctx context.Context, session authn.Session, id string, user users.UserReq) (users.User, error) { - switch err := am.checkSuperAdmin(ctx, session); { - case err == nil: - session.SuperAdmin = true - case errors.Contains(err, svcerr.ErrSuperAdminAction): - default: - return users.User{}, err - } - - return am.svc.UpdateTags(ctx, session, id, user) -} - -func (am *authorizationMiddleware) UpdateEmail(ctx context.Context, session authn.Session, id, email string) (users.User, error) { - switch err := am.checkSuperAdmin(ctx, session); { - case err == nil: - session.SuperAdmin = true - case errors.Contains(err, svcerr.ErrSuperAdminAction): - default: - return users.User{}, err - } - - return am.svc.UpdateEmail(ctx, session, id, email) -} - -func (am *authorizationMiddleware) UpdateUsername(ctx context.Context, session authn.Session, id, username string) (users.User, error) { - switch err := am.checkSuperAdmin(ctx, session); { - case err == nil: - session.SuperAdmin = true - case errors.Contains(err, svcerr.ErrSuperAdminAction): - default: - return users.User{}, err - } - - return am.svc.UpdateUsername(ctx, session, id, username) -} - -func (am *authorizationMiddleware) UpdateProfilePicture(ctx context.Context, session authn.Session, id string, usr users.UserReq) (users.User, error) { - switch err := am.checkSuperAdmin(ctx, session); { - case err == nil: - session.SuperAdmin = true - case errors.Contains(err, svcerr.ErrSuperAdminAction): - default: - return users.User{}, err - } - - return am.svc.UpdateProfilePicture(ctx, session, id, usr) -} - -func (am *authorizationMiddleware) SendPasswordReset(ctx context.Context, email string) error { - return am.svc.SendPasswordReset(ctx, email) -} - -func (am *authorizationMiddleware) UpdateSecret(ctx context.Context, session authn.Session, oldSecret, newSecret string) (users.User, error) { - return am.svc.UpdateSecret(ctx, session, oldSecret, newSecret) -} - -func (am *authorizationMiddleware) ResetSecret(ctx context.Context, session authn.Session, secret string) error { - return am.svc.ResetSecret(ctx, session, secret) -} - -func (am *authorizationMiddleware) UpdateRole(ctx context.Context, session authn.Session, user users.User) (users.User, error) { - if err := am.checkSuperAdmin(ctx, session); err != nil { - return users.User{}, err - } - session.SuperAdmin = true - if err := am.authorize(ctx, session, "", policies.UserType, policies.UsersKind, user.ID, policies.MembershipPermission, policies.PlatformType, policies.MagistralaObject); err != nil { - return users.User{}, err - } - - return am.svc.UpdateRole(ctx, session, user) -} - -func (am *authorizationMiddleware) Enable(ctx context.Context, session authn.Session, id string) (users.User, error) { - switch err := am.checkSuperAdmin(ctx, session); { - case err == nil: - session.SuperAdmin = true - case errors.Contains(err, svcerr.ErrSuperAdminAction): - default: - return users.User{}, err - } - - return am.svc.Enable(ctx, session, id) -} - -func (am *authorizationMiddleware) Disable(ctx context.Context, session authn.Session, id string) (users.User, error) { - switch err := am.checkSuperAdmin(ctx, session); { - case err == nil: - session.SuperAdmin = true - case errors.Contains(err, svcerr.ErrSuperAdminAction): - default: - return users.User{}, err - } - - return am.svc.Disable(ctx, session, id) -} - -func (am *authorizationMiddleware) Delete(ctx context.Context, session authn.Session, id string) error { - switch err := am.checkSuperAdmin(ctx, session); { - case err == nil: - session.SuperAdmin = true - case errors.Contains(err, svcerr.ErrSuperAdminAction): - default: - return err - } - - return am.svc.Delete(ctx, session, id) -} - -func (am *authorizationMiddleware) Identify(ctx context.Context, session authn.Session) (string, error) { - return am.svc.Identify(ctx, session) -} - -func (am *authorizationMiddleware) IssueToken(ctx context.Context, username, secret, description string) (*grpcTokenV1.Token, error) { - return am.svc.IssueToken(ctx, username, secret, description) -} - -func (am *authorizationMiddleware) RefreshToken(ctx context.Context, session authn.Session, refreshToken string) (*grpcTokenV1.Token, error) { - return am.svc.RefreshToken(ctx, session, refreshToken) -} - -func (am *authorizationMiddleware) RevokeRefreshToken(ctx context.Context, session authn.Session, tokenID string) error { - return am.svc.RevokeRefreshToken(ctx, session, tokenID) -} - -func (am *authorizationMiddleware) ListActiveRefreshTokens(ctx context.Context, session authn.Session) (*grpcTokenV1.ListUserRefreshTokensRes, error) { - return am.svc.ListActiveRefreshTokens(ctx, session) -} - -func (am *authorizationMiddleware) OAuthCallback(ctx context.Context, user users.User) (users.User, error) { - return am.svc.OAuthCallback(ctx, user) -} - -func (am *authorizationMiddleware) OAuthAddUserPolicy(ctx context.Context, user users.User) error { - if err := am.authorize(ctx, authn.Session{}, "", policies.UserType, policies.UsersKind, user.ID, policies.MembershipPermission, policies.PlatformType, policies.MagistralaObject); err == nil { - return nil - } - return am.svc.OAuthAddUserPolicy(ctx, user) -} - -func (am *authorizationMiddleware) checkSuperAdmin(ctx context.Context, session authn.Session) error { - if session.Role != authn.SuperAdminRole { - return svcerr.ErrSuperAdminAction - } - if err := am.authz.Authorize(ctx, smqauthz.PolicyReq{ - SubjectType: policies.UserType, - Subject: session.UserID, - Permission: policies.AdminPermission, - ObjectType: policies.PlatformType, - Object: policies.MagistralaObject, - }, nil); err != nil { - return err - } - return nil -} - -func (am *authorizationMiddleware) authorize(ctx context.Context, session authn.Session, domain, subjType, subjKind, subj, perm, objType, obj string) error { - req := smqauthz.PolicyReq{ - Domain: domain, - SubjectType: subjType, - SubjectKind: subjKind, - Subject: subj, - Permission: perm, - ObjectType: objType, - Object: obj, - } - - var pat *smqauthz.PATReq - if session.PatID != "" { - pat = &smqauthz.PATReq{ - UserID: session.UserID, - PatID: session.PatID, - EntityID: subj, - EntityType: auth.UsersType.String(), - Operation: perm, - Domain: domain, - } - } - - if err := am.authz.Authorize(ctx, req, pat); err != nil { - return err - } - return nil -} diff --git a/users/middleware/doc.go b/users/middleware/doc.go deleted file mode 100644 index 12db13874..000000000 --- a/users/middleware/doc.go +++ /dev/null @@ -1,9 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package middleware provides authorization, logging, metrics and tracing middleware -// for Magistrala Users Service. -// -// For more details about tracing instrumentation for Magistrala refer to the -// documentation at https://magistrala.absmach.eu/docs/. -package middleware diff --git a/users/middleware/logging.go b/users/middleware/logging.go deleted file mode 100644 index 5913697fe..000000000 --- a/users/middleware/logging.go +++ /dev/null @@ -1,562 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package middleware - -import ( - "context" - "log/slog" - "time" - - grpcTokenV1 "github.com/absmach/magistrala/api/grpc/token/v1" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/users" - "github.com/go-chi/chi/v5/middleware" -) - -var _ users.Service = (*loggingMiddleware)(nil) - -type loggingMiddleware struct { - logger *slog.Logger - svc users.Service -} - -// NewLogging adds logging facilities to the users service. -func NewLogging(svc users.Service, logger *slog.Logger) users.Service { - return &loggingMiddleware{logger, svc} -} - -// Register logs the user request. It logs the user id and the time it took to complete the request. -// If the request fails, it logs the error. -func (lm *loggingMiddleware) Register(ctx context.Context, session authn.Session, user users.User, selfRegister bool) (u users.User, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.Group("user", - slog.String("username", user.Credentials.Username), - slog.String("first_name", user.FirstName), - slog.String("last_name", user.LastName), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Register user failed", args...) - return - } - args = append(args, slog.String("user_id", u.ID)) - lm.logger.Info("Register user completed successfully", args...) - }(time.Now()) - return lm.svc.Register(ctx, session, user, selfRegister) -} - -// SendVerification logs the send_verification request. It logs the time it took to complete the request. -// If the request fails, it logs the error. -func (lm *loggingMiddleware) SendVerification(ctx context.Context, session authn.Session) (err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("request_id", middleware.GetReqID(ctx)), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Send verification failed", args...) - return - } - lm.logger.Info("Send verification completed successfully", args...) - }(time.Now()) - return lm.svc.SendVerification(ctx, session) -} - -// VerifyEmail logs the verify_email request. It logs the time it took to complete the request. -// If the request fails, it logs the error. -func (lm *loggingMiddleware) VerifyEmail(ctx context.Context, verificationToken string) (user users.User, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("user_id", user.ID), - slog.String("request_id", middleware.GetReqID(ctx)), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Verify email failed", args...) - return - } - lm.logger.Info("Verify email completed successfully", args...) - }(time.Now()) - return lm.svc.VerifyEmail(ctx, verificationToken) -} - -// IssueToken logs the issue_token request. It logs the username type and the time it took to complete the request. -// If the request fails, it logs the error. -func (lm *loggingMiddleware) IssueToken(ctx context.Context, username, secret, description string) (t *grpcTokenV1.Token, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("request_id", middleware.GetReqID(ctx)), - } - if t.AccessType != "" { - args = append(args, slog.String("access_type", t.AccessType)) - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Issue token failed", args...) - return - } - lm.logger.Info("Issue token completed successfully", args...) - }(time.Now()) - return lm.svc.IssueToken(ctx, username, secret, description) -} - -// RefreshToken logs the refresh_token request. It logs the refreshtoken, token type and the time it took to complete the request. -// If the request fails, it logs the error. -func (lm *loggingMiddleware) RefreshToken(ctx context.Context, session authn.Session, refreshToken string) (t *grpcTokenV1.Token, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("request_id", middleware.GetReqID(ctx)), - } - if t.AccessType != "" { - args = append(args, slog.String("access_type", t.AccessType)) - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Refresh token failed", args...) - return - } - lm.logger.Info("Refresh token completed successfully", args...) - }(time.Now()) - return lm.svc.RefreshToken(ctx, session, refreshToken) -} - -// RevokeRefreshToken logs the revoke_refresh_token request. It logs the time it took to complete the request. -// If the request fails, it logs the error. -func (lm *loggingMiddleware) RevokeRefreshToken(ctx context.Context, session authn.Session, tokenID string) (err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("request_id", middleware.GetReqID(ctx)), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Revoke refresh token failed", args...) - return - } - lm.logger.Info("Revoke refresh token completed successfully", args...) - }(time.Now()) - return lm.svc.RevokeRefreshToken(ctx, session, tokenID) -} - -// ListActiveRefreshTokens logs the list_active_refresh_tokens request. It logs the time it took to complete the request. -// If the request fails, it logs the error. -func (lm *loggingMiddleware) ListActiveRefreshTokens(ctx context.Context, session authn.Session) (tokens *grpcTokenV1.ListUserRefreshTokensRes, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("request_id", middleware.GetReqID(ctx)), - } - if tokens != nil { - args = append(args, slog.Int("tokens_count", len(tokens.GetRefreshTokens()))) - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("List active refresh tokens failed", args...) - return - } - lm.logger.Info("List active refresh tokens completed successfully", args...) - }(time.Now()) - return lm.svc.ListActiveRefreshTokens(ctx, session) -} - -// View logs the view_user request. It logs the user id and the time it took to complete the request. -// If the request fails, it logs the error. -func (lm *loggingMiddleware) View(ctx context.Context, session authn.Session, id string) (c users.User, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.Group("user", - slog.String("id", id), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("View user failed", args...) - return - } - lm.logger.Info("View user completed successfully", args...) - }(time.Now()) - return lm.svc.View(ctx, session, id) -} - -// ViewProfile logs the view_profile request. It logs the user id and the time it took to complete the request. -// If the request fails, it logs the error. -func (lm *loggingMiddleware) ViewProfile(ctx context.Context, session authn.Session) (c users.User, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.Group("user", - slog.String("id", c.ID), - slog.String("username", c.Credentials.Username), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("View profile failed", args...) - return - } - lm.logger.Info("View profile completed successfully", args...) - }(time.Now()) - return lm.svc.ViewProfile(ctx, session) -} - -// ListUsers logs the list_users request. It logs the page metadata and the time it took to complete the request. -// If the request fails, it logs the error. -func (lm *loggingMiddleware) ListUsers(ctx context.Context, session authn.Session, pm users.Page) (cp users.UsersPage, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.Group("page", - slog.Uint64("limit", pm.Limit), - slog.Uint64("offset", pm.Offset), - slog.Uint64("total", cp.Total), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("List users failed", args...) - return - } - lm.logger.Info("List users completed successfully", args...) - }(time.Now()) - return lm.svc.ListUsers(ctx, session, pm) -} - -// SearchUsers logs the search_users request. It logs the page metadata and the time it took to complete the request. -func (lm *loggingMiddleware) SearchUsers(ctx context.Context, cp users.Page) (mp users.UsersPage, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.Group("page", - slog.Uint64("limit", cp.Limit), - slog.Uint64("offset", cp.Offset), - slog.Uint64("total", mp.Total), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Search users failed to complete successfully", args...) - return - } - lm.logger.Info("Search users completed successfully", args...) - }(time.Now()) - return lm.svc.SearchUsers(ctx, cp) -} - -// Update logs the update_user request. It logs the user id and the time it took to complete the request. -// If the request fails, it logs the error. -func (lm *loggingMiddleware) Update(ctx context.Context, session authn.Session, id string, user users.UserReq) (u users.User, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.Group("user", - slog.String("id", u.ID), - slog.String("username", u.Credentials.Username), - slog.String("first_name", u.FirstName), - slog.String("last_name", u.LastName), - slog.Any("metadata", u.Metadata), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Update user failed", args...) - return - } - lm.logger.Info("Update user completed successfully", args...) - }(time.Now()) - return lm.svc.Update(ctx, session, id, user) -} - -// UpdateTags logs the update_user_tags request. It logs the user id and the time it took to complete the request. -// If the request fails, it logs the error. -func (lm *loggingMiddleware) UpdateTags(ctx context.Context, session authn.Session, id string, user users.UserReq) (c users.User, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.Group("user", - slog.String("id", c.ID), - slog.Any("tags", c.Tags), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Update user tags failed", args...) - return - } - lm.logger.Info("Update user tags completed successfully", args...) - }(time.Now()) - return lm.svc.UpdateTags(ctx, session, id, user) -} - -// UpdateEmail logs the update_user_email request. It logs the user id and the time it took to complete the request. -// If the request fails, it logs the error. -func (lm *loggingMiddleware) UpdateEmail(ctx context.Context, session authn.Session, id, email string) (c users.User, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.Group("user", - slog.String("id", c.ID), - slog.String("email", c.Email), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Update user email failed", args...) - return - } - lm.logger.Info("Update user email completed successfully", args...) - }(time.Now()) - return lm.svc.UpdateEmail(ctx, session, id, email) -} - -// UpdateSecret logs the update_user_secret request. It logs the user id and the time it took to complete the request. -// If the request fails, it logs the error. -func (lm *loggingMiddleware) UpdateSecret(ctx context.Context, session authn.Session, oldSecret, newSecret string) (c users.User, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.Group("user", - slog.String("id", c.ID), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Update user secret failed", args...) - return - } - lm.logger.Info("Update user secret completed successfully", args...) - }(time.Now()) - return lm.svc.UpdateSecret(ctx, session, oldSecret, newSecret) -} - -// UpdateUsername logs the update_usernames request. It logs the user id and the time it took to complete the request. -// If the request fails, it logs the error. -func (lm *loggingMiddleware) UpdateUsername(ctx context.Context, session authn.Session, id, username string) (u users.User, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.Group("user", - slog.String("id", u.ID), - slog.String("username", u.Credentials.Username), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Update user names failed", args...) - return - } - lm.logger.Info("Update user names completed successfully", args...) - }(time.Now()) - return lm.svc.UpdateUsername(ctx, session, id, username) -} - -// UpdateProfilePicture logs the update_profile_picture request. It logs the user id and the time it took to complete the request. -// If the request fails, it logs the error. -func (lm *loggingMiddleware) UpdateProfilePicture(ctx context.Context, session authn.Session, id string, user users.UserReq) (u users.User, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.Group("user", - slog.String("id", u.ID), - slog.String("profile_picture", u.ProfilePicture), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Update profile picture failed", args...) - return - } - lm.logger.Info("Update profile picture completed successfully", args...) - }(time.Now()) - return lm.svc.UpdateProfilePicture(ctx, session, id, user) -} - -// SendPasswordReset logs the send_password_reset request. It logs the time it took to complete the request. -// If the request fails, it logs the error. -func (lm *loggingMiddleware) SendPasswordReset(ctx context.Context, email string) (err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("request_id", middleware.GetReqID(ctx)), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Generate reset token failed", args...) - return - } - lm.logger.Info("Send password reset completed successfully", args...) - }(time.Now()) - return lm.svc.SendPasswordReset(ctx, email) -} - -// ResetSecret logs the reset_secret request. It logs the time it took to complete the request. -// If the request fails, it logs the error. -func (lm *loggingMiddleware) ResetSecret(ctx context.Context, session authn.Session, secret string) (err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("request_id", middleware.GetReqID(ctx)), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Reset secret failed", args...) - return - } - lm.logger.Info("Reset secret completed successfully", args...) - }(time.Now()) - return lm.svc.ResetSecret(ctx, session, secret) -} - -// UpdateRole logs the update_user_role request. It logs the user id and the time it took to complete the request. -// If the request fails, it logs the error. -func (lm *loggingMiddleware) UpdateRole(ctx context.Context, session authn.Session, user users.User) (c users.User, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.Group("user", - slog.String("id", user.ID), - slog.String("role", user.Role.String()), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Update user role failed", args...) - return - } - lm.logger.Info("Update user role completed successfully", args...) - }(time.Now()) - return lm.svc.UpdateRole(ctx, session, user) -} - -// Enable logs the enable_user request. It logs the user id and the time it took to complete the request. -// If the request fails, it logs the error. -func (lm *loggingMiddleware) Enable(ctx context.Context, session authn.Session, id string) (c users.User, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.Group("user", - slog.String("id", id), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Enable user failed", args...) - return - } - lm.logger.Info("Enable user completed successfully", args...) - }(time.Now()) - return lm.svc.Enable(ctx, session, id) -} - -// Disable logs the disable_user request. It logs the user id and the time it took to complete the request. -// If the request fails, it logs the error. -func (lm *loggingMiddleware) Disable(ctx context.Context, session authn.Session, id string) (c users.User, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.Group("user", - slog.String("id", id), - ), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Disable user failed", args...) - return - } - lm.logger.Info("Disable user completed successfully", args...) - }(time.Now()) - return lm.svc.Disable(ctx, session, id) -} - -// Identify logs the identify request. It logs the time it took to complete the request. -func (lm *loggingMiddleware) Identify(ctx context.Context, session authn.Session) (id string, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.String("user_id", id), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Identify user failed", args...) - return - } - lm.logger.Info("Identify user completed successfully", args...) - }(time.Now()) - return lm.svc.Identify(ctx, session) -} - -func (lm *loggingMiddleware) OAuthCallback(ctx context.Context, user users.User) (c users.User, err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.String("user_id", user.ID), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("OAuth callback failed", args...) - return - } - lm.logger.Info("OAuth callback completed successfully", args...) - }(time.Now()) - return lm.svc.OAuthCallback(ctx, user) -} - -// Delete logs the delete_user request. It logs the user id and token and the time it took to complete the request. -func (lm *loggingMiddleware) Delete(ctx context.Context, session authn.Session, id string) (err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.String("user_id", id), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Delete user failed to complete successfully", args...) - return - } - lm.logger.Info("Delete user completed successfully", args...) - }(time.Now()) - return lm.svc.Delete(ctx, session, id) -} - -// OAuthAddUserPolicy logs the add_user_policy request. It logs the user id and the time it took to complete the request. -func (lm *loggingMiddleware) OAuthAddUserPolicy(ctx context.Context, user users.User) (err error) { - defer func(begin time.Time) { - args := []any{ - slog.String("duration", time.Since(begin).String()), - slog.String("request_id", middleware.GetReqID(ctx)), - slog.String("user_id", user.ID), - } - if err != nil { - args = append(args, slog.String("error", err.Error())) - lm.logger.Warn("Add user policy failed", args...) - return - } - lm.logger.Info("Add user policy completed successfully", args...) - }(time.Now()) - return lm.svc.OAuthAddUserPolicy(ctx, user) -} diff --git a/users/middleware/metrics.go b/users/middleware/metrics.go deleted file mode 100644 index 35cba56f8..000000000 --- a/users/middleware/metrics.go +++ /dev/null @@ -1,265 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package middleware - -import ( - "context" - "time" - - grpcTokenV1 "github.com/absmach/magistrala/api/grpc/token/v1" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/users" - "github.com/go-kit/kit/metrics" -) - -var _ users.Service = (*metricsMiddleware)(nil) - -type metricsMiddleware struct { - counter metrics.Counter - latency metrics.Histogram - svc users.Service -} - -// NewMetrics instruments policies service by tracking request count and latency. -func NewMetrics(svc users.Service, counter metrics.Counter, latency metrics.Histogram) users.Service { - return &metricsMiddleware{ - counter: counter, - latency: latency, - svc: svc, - } -} - -// Register instruments Register method with metrics. -func (ms *metricsMiddleware) Register(ctx context.Context, session authn.Session, user users.User, selfRegister bool) (users.User, error) { - defer func(begin time.Time) { - ms.counter.With("method", "register_user").Add(1) - ms.latency.With("method", "register_user").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.Register(ctx, session, user, selfRegister) -} - -// SendVerification instruments SendVerification method with metrics. -func (ms *metricsMiddleware) SendVerification(ctx context.Context, session authn.Session) error { - defer func(begin time.Time) { - ms.counter.With("method", "send_verification").Add(1) - ms.latency.With("method", "send_verification").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.SendVerification(ctx, session) -} - -// VerifyEmail instruments VerifyEmail method with metrics. -func (ms *metricsMiddleware) VerifyEmail(ctx context.Context, verificationToken string) (users.User, error) { - defer func(begin time.Time) { - ms.counter.With("method", "verify_email").Add(1) - ms.latency.With("method", "verify_email").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.VerifyEmail(ctx, verificationToken) -} - -// IssueToken instruments IssueToken method with metrics. -func (ms *metricsMiddleware) IssueToken(ctx context.Context, username, secret, description string) (*grpcTokenV1.Token, error) { - defer func(begin time.Time) { - ms.counter.With("method", "issue_token").Add(1) - ms.latency.With("method", "issue_token").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.IssueToken(ctx, username, secret, description) -} - -// RefreshToken instruments RefreshToken method with metrics. -func (ms *metricsMiddleware) RefreshToken(ctx context.Context, session authn.Session, refreshToken string) (token *grpcTokenV1.Token, err error) { - defer func(begin time.Time) { - ms.counter.With("method", "refresh_token").Add(1) - ms.latency.With("method", "refresh_token").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.RefreshToken(ctx, session, refreshToken) -} - -// RevokeRefreshToken instruments RevokeRefreshToken method with metrics. -func (ms *metricsMiddleware) RevokeRefreshToken(ctx context.Context, session authn.Session, tokenID string) error { - defer func(begin time.Time) { - ms.counter.With("method", "revoke_refresh_token").Add(1) - ms.latency.With("method", "revoke_refresh_token").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.RevokeRefreshToken(ctx, session, tokenID) -} - -// ListActiveRefreshTokens instruments ListActiveRefreshTokens method with metrics. -func (ms *metricsMiddleware) ListActiveRefreshTokens(ctx context.Context, session authn.Session) (*grpcTokenV1.ListUserRefreshTokensRes, error) { - defer func(begin time.Time) { - ms.counter.With("method", "list_active_refresh_tokens").Add(1) - ms.latency.With("method", "list_active_refresh_tokens").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.ListActiveRefreshTokens(ctx, session) -} - -// View instruments View method with metrics. -func (ms *metricsMiddleware) View(ctx context.Context, session authn.Session, id string) (users.User, error) { - defer func(begin time.Time) { - ms.counter.With("method", "view_user").Add(1) - ms.latency.With("method", "view_user").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.View(ctx, session, id) -} - -// ViewProfile instruments ViewProfile method with metrics. -func (ms *metricsMiddleware) ViewProfile(ctx context.Context, session authn.Session) (users.User, error) { - defer func(begin time.Time) { - ms.counter.With("method", "view_profile").Add(1) - ms.latency.With("method", "view_profile").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.ViewProfile(ctx, session) -} - -// ListUsers instruments ListUsers method with metrics. -func (ms *metricsMiddleware) ListUsers(ctx context.Context, session authn.Session, pm users.Page) (users.UsersPage, error) { - defer func(begin time.Time) { - ms.counter.With("method", "list_users").Add(1) - ms.latency.With("method", "list_users").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.ListUsers(ctx, session, pm) -} - -// SearchUsers instruments SearchUsers method with metrics. -func (ms *metricsMiddleware) SearchUsers(ctx context.Context, pm users.Page) (mp users.UsersPage, err error) { - defer func(begin time.Time) { - ms.counter.With("method", "search_users").Add(1) - ms.latency.With("method", "search_users").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.SearchUsers(ctx, pm) -} - -// Update instruments Update method with metrics. -func (ms *metricsMiddleware) Update(ctx context.Context, session authn.Session, id string, user users.UserReq) (users.User, error) { - defer func(begin time.Time) { - ms.counter.With("method", "update_user").Add(1) - ms.latency.With("method", "update_user").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.Update(ctx, session, id, user) -} - -// UpdateTags instruments UpdateTags method with metrics. -func (ms *metricsMiddleware) UpdateTags(ctx context.Context, session authn.Session, id string, user users.UserReq) (users.User, error) { - defer func(begin time.Time) { - ms.counter.With("method", "update_user_tags").Add(1) - ms.latency.With("method", "update_user_tags").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.UpdateTags(ctx, session, id, user) -} - -// UpdateEmail instruments UpdateEmail method with metrics. -func (ms *metricsMiddleware) UpdateEmail(ctx context.Context, session authn.Session, id, email string) (users.User, error) { - defer func(begin time.Time) { - ms.counter.With("method", "update_user_email").Add(1) - ms.latency.With("method", "update_user_email").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.UpdateEmail(ctx, session, id, email) -} - -// UpdateSecret instruments UpdateSecret method with metrics. -func (ms *metricsMiddleware) UpdateSecret(ctx context.Context, session authn.Session, oldSecret, newSecret string) (users.User, error) { - defer func(begin time.Time) { - ms.counter.With("method", "update_user_secret").Add(1) - ms.latency.With("method", "update_user_secret").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.UpdateSecret(ctx, session, oldSecret, newSecret) -} - -// UpdateUsername instruments UpdateUsername method with metrics. -func (ms *metricsMiddleware) UpdateUsername(ctx context.Context, session authn.Session, id, username string) (users.User, error) { - defer func(begin time.Time) { - ms.counter.With("method", "update_usernames").Add(1) - ms.latency.With("method", "update_usernames").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.UpdateUsername(ctx, session, id, username) -} - -// UpdateProfilePicture instruments UpdateProfilePicture method with metrics. -func (ms *metricsMiddleware) UpdateProfilePicture(ctx context.Context, session authn.Session, id string, user users.UserReq) (users.User, error) { - defer func(begin time.Time) { - ms.counter.With("method", "update_profile_picture").Add(1) - ms.latency.With("method", "update_profile_picture").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.UpdateProfilePicture(ctx, session, id, user) -} - -// SendPasswordReset instruments SendPasswordReset method with metrics. -func (ms *metricsMiddleware) SendPasswordReset(ctx context.Context, email string) error { - defer func(begin time.Time) { - ms.counter.With("method", "send_password_reset").Add(1) - ms.latency.With("method", "send_password_reset").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.SendPasswordReset(ctx, email) -} - -// ResetSecret instruments ResetSecret method with metrics. -func (ms *metricsMiddleware) ResetSecret(ctx context.Context, session authn.Session, secret string) error { - defer func(begin time.Time) { - ms.counter.With("method", "reset_secret").Add(1) - ms.latency.With("method", "reset_secret").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.ResetSecret(ctx, session, secret) -} - -// UpdateRole instruments UpdateRole method with metrics. -func (ms *metricsMiddleware) UpdateRole(ctx context.Context, session authn.Session, user users.User) (users.User, error) { - defer func(begin time.Time) { - ms.counter.With("method", "update_user_role").Add(1) - ms.latency.With("method", "update_user_role").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.UpdateRole(ctx, session, user) -} - -// Enable instruments Enable method with metrics. -func (ms *metricsMiddleware) Enable(ctx context.Context, session authn.Session, id string) (users.User, error) { - defer func(begin time.Time) { - ms.counter.With("method", "enable_user").Add(1) - ms.latency.With("method", "enable_user").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.Enable(ctx, session, id) -} - -// Disable instruments Disable method with metrics. -func (ms *metricsMiddleware) Disable(ctx context.Context, session authn.Session, id string) (users.User, error) { - defer func(begin time.Time) { - ms.counter.With("method", "disable_user").Add(1) - ms.latency.With("method", "disable_user").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.Disable(ctx, session, id) -} - -// Identify instruments Identify method with metrics. -func (ms *metricsMiddleware) Identify(ctx context.Context, session authn.Session) (string, error) { - defer func(begin time.Time) { - ms.counter.With("method", "identify").Add(1) - ms.latency.With("method", "identify").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.Identify(ctx, session) -} - -// OAuthCallback instruments OAuthCallback method with metrics. -func (ms *metricsMiddleware) OAuthCallback(ctx context.Context, user users.User) (users.User, error) { - defer func(begin time.Time) { - ms.counter.With("method", "oauth_callback").Add(1) - ms.latency.With("method", "oauth_callback").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.OAuthCallback(ctx, user) -} - -// Delete instruments Delete method with metrics. -func (ms *metricsMiddleware) Delete(ctx context.Context, session authn.Session, id string) error { - defer func(begin time.Time) { - ms.counter.With("method", "delete_user").Add(1) - ms.latency.With("method", "delete_user").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.Delete(ctx, session, id) -} - -// OAuthAddUserPolicy instruments OAuthAddUserPolicy method with metrics. -func (ms *metricsMiddleware) OAuthAddUserPolicy(ctx context.Context, user users.User) error { - defer func(begin time.Time) { - ms.counter.With("method", "add_user_policy").Add(1) - ms.latency.With("method", "add_user_policy").Observe(time.Since(begin).Seconds()) - }(time.Now()) - return ms.svc.OAuthAddUserPolicy(ctx, user) -} diff --git a/users/middleware/tracing.go b/users/middleware/tracing.go deleted file mode 100644 index 12dc93394..000000000 --- a/users/middleware/tracing.go +++ /dev/null @@ -1,266 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package middleware - -import ( - "context" - - grpcTokenV1 "github.com/absmach/magistrala/api/grpc/token/v1" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/tracing" - users "github.com/absmach/magistrala/users" - "go.opentelemetry.io/otel/attribute" - "go.opentelemetry.io/otel/trace" -) - -var _ users.Service = (*tracingMiddleware)(nil) - -type tracingMiddleware struct { - tracer trace.Tracer - svc users.Service -} - -// NewTracing returns a new users service with tracing capabilities. -func NewTracing(svc users.Service, tracer trace.Tracer) users.Service { - return &tracingMiddleware{tracer, svc} -} - -// Register traces the "Register" operation of the wrapped users.Service. -func (tm *tracingMiddleware) Register(ctx context.Context, session authn.Session, user users.User, selfRegister bool) (users.User, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_register_user", trace.WithAttributes(attribute.String("email", user.Email))) - defer span.End() - - return tm.svc.Register(ctx, session, user, selfRegister) -} - -// SendVerification traces the "SendVerification" operation of the wrapped users.Service. -func (tm *tracingMiddleware) SendVerification(ctx context.Context, session authn.Session) error { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_send_verification") - defer span.End() - - return tm.svc.SendVerification(ctx, session) -} - -// VerifyEmail traces the "VerifyEmail" operation of the wrapped users.Service. -func (tm *tracingMiddleware) VerifyEmail(ctx context.Context, verificationToken string) (users.User, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_verify_email") - defer span.End() - return tm.svc.VerifyEmail(ctx, verificationToken) -} - -// IssueToken traces the "IssueToken" operation of the wrapped users.Service. -func (tm *tracingMiddleware) IssueToken(ctx context.Context, username, secret, description string) (*grpcTokenV1.Token, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_issue_token", trace.WithAttributes(attribute.String("username", username))) - defer span.End() - - return tm.svc.IssueToken(ctx, username, secret, description) -} - -// RefreshToken traces the "RefreshToken" operation of the wrapped users.Service. -func (tm *tracingMiddleware) RefreshToken(ctx context.Context, session authn.Session, refreshToken string) (*grpcTokenV1.Token, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_refresh_token", trace.WithAttributes(attribute.String("refresh_token", refreshToken))) - defer span.End() - - return tm.svc.RefreshToken(ctx, session, refreshToken) -} - -// RevokeRefreshToken traces the "RevokeRefreshToken" operation of the wrapped users.Service. -func (tm *tracingMiddleware) RevokeRefreshToken(ctx context.Context, session authn.Session, tokenID string) error { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_revoke_refresh_token") - defer span.End() - - return tm.svc.RevokeRefreshToken(ctx, session, tokenID) -} - -// ListActiveRefreshTokens traces the "ListActiveRefreshTokens" operation of the wrapped users.Service. -func (tm *tracingMiddleware) ListActiveRefreshTokens(ctx context.Context, session authn.Session) (*grpcTokenV1.ListUserRefreshTokensRes, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_list_active_refresh_tokens") - defer span.End() - - return tm.svc.ListActiveRefreshTokens(ctx, session) -} - -// View traces the "View" operation of the wrapped users.Service. -func (tm *tracingMiddleware) View(ctx context.Context, session authn.Session, id string) (users.User, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_view_user", trace.WithAttributes(attribute.String("id", id))) - defer span.End() - - return tm.svc.View(ctx, session, id) -} - -// ListUsers traces the "ListUsers" operation of the wrapped users.Service. -func (tm *tracingMiddleware) ListUsers(ctx context.Context, session authn.Session, pm users.Page) (users.UsersPage, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_list_users", trace.WithAttributes( - attribute.Int64("offset", int64(pm.Offset)), - attribute.Int64("limit", int64(pm.Limit)), - attribute.String("direction", pm.Dir), - attribute.String("order", pm.Order), - )) - - defer span.End() - - return tm.svc.ListUsers(ctx, session, pm) -} - -// SearchUsers traces the "SearchUsers" operation of the wrapped users.Service. -func (tm *tracingMiddleware) SearchUsers(ctx context.Context, pm users.Page) (users.UsersPage, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_search_users", trace.WithAttributes( - attribute.Int64("offset", int64(pm.Offset)), - attribute.Int64("limit", int64(pm.Limit)), - attribute.String("direction", pm.Dir), - attribute.String("order", pm.Order), - )) - defer span.End() - - return tm.svc.SearchUsers(ctx, pm) -} - -// Update traces the "Update" operation of the wrapped users.Service. -func (tm *tracingMiddleware) Update(ctx context.Context, session authn.Session, id string, user users.UserReq) (users.User, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_update_user", trace.WithAttributes( - attribute.String("id", id), - )) - defer span.End() - - return tm.svc.Update(ctx, session, id, user) -} - -// UpdateTags traces the "UpdateTags" operation of the wrapped users.Service. -func (tm *tracingMiddleware) UpdateTags(ctx context.Context, session authn.Session, id string, user users.UserReq) (users.User, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_update_user_tags", trace.WithAttributes( - attribute.String("id", id), - )) - defer span.End() - - return tm.svc.UpdateTags(ctx, session, id, user) -} - -// UpdateEmail traces the "UpdateEmail" operation of the wrapped users.Service. -func (tm *tracingMiddleware) UpdateEmail(ctx context.Context, session authn.Session, id, email string) (users.User, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_update_user_email", trace.WithAttributes( - attribute.String("id", id), - attribute.String("email", email), - )) - defer span.End() - - return tm.svc.UpdateEmail(ctx, session, id, email) -} - -// UpdateSecret traces the "UpdateSecret" operation of the wrapped users.Service. -func (tm *tracingMiddleware) UpdateSecret(ctx context.Context, session authn.Session, oldSecret, newSecret string) (users.User, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_update_user_secret") - defer span.End() - - return tm.svc.UpdateSecret(ctx, session, oldSecret, newSecret) -} - -// UpdateUsername traces the "UpdateUsername" operation of the wrapped users.Service. -func (tm *tracingMiddleware) UpdateUsername(ctx context.Context, session authn.Session, id, username string) (users.User, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_update_usernames", trace.WithAttributes( - attribute.String("id", id), - attribute.String("username", username), - )) - defer span.End() - - return tm.svc.UpdateUsername(ctx, session, id, username) -} - -// UpdateProfilePicture traces the "UpdateProfilePicture" operation of the wrapped users.Service. -func (tm *tracingMiddleware) UpdateProfilePicture(ctx context.Context, session authn.Session, id string, usr users.UserReq) (users.User, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_update_profile_picture", trace.WithAttributes( - attribute.String("id", id), - )) - defer span.End() - - return tm.svc.UpdateProfilePicture(ctx, session, id, usr) -} - -// SendPasswordReset traces the "SendPasswordReset" operation of the wrapped users.Service. -func (tm *tracingMiddleware) SendPasswordReset(ctx context.Context, email string) error { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_send_password_reset", trace.WithAttributes( - attribute.String("email", email), - )) - defer span.End() - - return tm.svc.SendPasswordReset(ctx, email) -} - -// ResetSecret traces the "ResetSecret" operation of the wrapped users.Service. -func (tm *tracingMiddleware) ResetSecret(ctx context.Context, session authn.Session, secret string) error { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_reset_secret") - defer span.End() - - return tm.svc.ResetSecret(ctx, session, secret) -} - -// ViewProfile traces the "ViewProfile" operation of the wrapped users.Service. -func (tm *tracingMiddleware) ViewProfile(ctx context.Context, session authn.Session) (users.User, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_view_profile") - defer span.End() - - return tm.svc.ViewProfile(ctx, session) -} - -// UpdateRole traces the "UpdateRole" operation of the wrapped users.Service. -func (tm *tracingMiddleware) UpdateRole(ctx context.Context, session authn.Session, cli users.User) (users.User, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_update_user_role", trace.WithAttributes( - attribute.String("id", cli.ID), - attribute.StringSlice("tags", cli.Tags), - )) - defer span.End() - - return tm.svc.UpdateRole(ctx, session, cli) -} - -// Enable traces the "Enable" operation of the wrapped users.Service. -func (tm *tracingMiddleware) Enable(ctx context.Context, session authn.Session, id string) (users.User, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_enable_user", trace.WithAttributes(attribute.String("id", id))) - defer span.End() - - return tm.svc.Enable(ctx, session, id) -} - -// Disable traces the "Disable" operation of the wrapped users.Service. -func (tm *tracingMiddleware) Disable(ctx context.Context, session authn.Session, id string) (users.User, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_disable_user", trace.WithAttributes(attribute.String("id", id))) - defer span.End() - - return tm.svc.Disable(ctx, session, id) -} - -// Identify traces the "Identify" operation of the wrapped users.Service. -func (tm *tracingMiddleware) Identify(ctx context.Context, session authn.Session) (string, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_identify", trace.WithAttributes(attribute.String("user_id", session.UserID))) - defer span.End() - - return tm.svc.Identify(ctx, session) -} - -// OAuthCallback traces the "OAuthCallback" operation of the wrapped users.Service. -func (tm *tracingMiddleware) OAuthCallback(ctx context.Context, user users.User) (users.User, error) { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_oauth_callback", trace.WithAttributes( - attribute.String("user_id", user.ID), - )) - defer span.End() - - return tm.svc.OAuthCallback(ctx, user) -} - -// Delete traces the "Delete" operation of the wrapped users.Service. -func (tm *tracingMiddleware) Delete(ctx context.Context, session authn.Session, id string) error { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_delete_user", trace.WithAttributes(attribute.String("id", id))) - defer span.End() - - return tm.svc.Delete(ctx, session, id) -} - -// OAuthAddUserPolicy traces the "OAuthAddUserPolicy" operation of the wrapped users.Service. -func (tm *tracingMiddleware) OAuthAddUserPolicy(ctx context.Context, user users.User) error { - ctx, span := tracing.StartSpan(ctx, tm.tracer, "svc_add_user_policy", trace.WithAttributes( - attribute.String("id", user.ID), - )) - defer span.End() - - return tm.svc.OAuthAddUserPolicy(ctx, user) -} diff --git a/users/mocks/doc.go b/users/mocks/doc.go deleted file mode 100644 index 16ed198af..000000000 --- a/users/mocks/doc.go +++ /dev/null @@ -1,5 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package mocks contains mocks for testing purposes. -package mocks diff --git a/users/mocks/emailer.go b/users/mocks/emailer.go deleted file mode 100644 index 5ebeea684..000000000 --- a/users/mocks/emailer.go +++ /dev/null @@ -1,166 +0,0 @@ -// Copyright (c) Abstract Machines - -// SPDX-License-Identifier: Apache-2.0 - -// Code generated by mockery; DO NOT EDIT. -// github.com/vektra/mockery -// template: testify - -package mocks - -import ( - mock "github.com/stretchr/testify/mock" -) - -// NewEmailer creates a new instance of Emailer. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. -// The first argument is typically a *testing.T value. -func NewEmailer(t interface { - mock.TestingT - Cleanup(func()) -}) *Emailer { - mock := &Emailer{} - mock.Mock.Test(t) - - t.Cleanup(func() { mock.AssertExpectations(t) }) - - return mock -} - -// Emailer is an autogenerated mock type for the Emailer type -type Emailer struct { - mock.Mock -} - -type Emailer_Expecter struct { - mock *mock.Mock -} - -func (_m *Emailer) EXPECT() *Emailer_Expecter { - return &Emailer_Expecter{mock: &_m.Mock} -} - -// SendPasswordReset provides a mock function for the type Emailer -func (_mock *Emailer) SendPasswordReset(To []string, user string, token string) error { - ret := _mock.Called(To, user, token) - - if len(ret) == 0 { - panic("no return value specified for SendPasswordReset") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func([]string, string, string) error); ok { - r0 = returnFunc(To, user, token) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Emailer_SendPasswordReset_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'SendPasswordReset' -type Emailer_SendPasswordReset_Call struct { - *mock.Call -} - -// SendPasswordReset is a helper method to define mock.On call -// - To []string -// - user string -// - token string -func (_e *Emailer_Expecter) SendPasswordReset(To interface{}, user interface{}, token interface{}) *Emailer_SendPasswordReset_Call { - return &Emailer_SendPasswordReset_Call{Call: _e.mock.On("SendPasswordReset", To, user, token)} -} - -func (_c *Emailer_SendPasswordReset_Call) Run(run func(To []string, user string, token string)) *Emailer_SendPasswordReset_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 []string - if args[0] != nil { - arg0 = args[0].([]string) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Emailer_SendPasswordReset_Call) Return(err error) *Emailer_SendPasswordReset_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Emailer_SendPasswordReset_Call) RunAndReturn(run func(To []string, user string, token string) error) *Emailer_SendPasswordReset_Call { - _c.Call.Return(run) - return _c -} - -// SendVerification provides a mock function for the type Emailer -func (_mock *Emailer) SendVerification(To []string, user string, verificationToken string) error { - ret := _mock.Called(To, user, verificationToken) - - if len(ret) == 0 { - panic("no return value specified for SendVerification") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func([]string, string, string) error); ok { - r0 = returnFunc(To, user, verificationToken) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Emailer_SendVerification_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'SendVerification' -type Emailer_SendVerification_Call struct { - *mock.Call -} - -// SendVerification is a helper method to define mock.On call -// - To []string -// - user string -// - verificationToken string -func (_e *Emailer_Expecter) SendVerification(To interface{}, user interface{}, verificationToken interface{}) *Emailer_SendVerification_Call { - return &Emailer_SendVerification_Call{Call: _e.mock.On("SendVerification", To, user, verificationToken)} -} - -func (_c *Emailer_SendVerification_Call) Run(run func(To []string, user string, verificationToken string)) *Emailer_SendVerification_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 []string - if args[0] != nil { - arg0 = args[0].([]string) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Emailer_SendVerification_Call) Return(err error) *Emailer_SendVerification_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Emailer_SendVerification_Call) RunAndReturn(run func(To []string, user string, verificationToken string) error) *Emailer_SendVerification_Call { - _c.Call.Return(run) - return _c -} diff --git a/users/mocks/hasher.go b/users/mocks/hasher.go deleted file mode 100644 index ab4964780..000000000 --- a/users/mocks/hasher.go +++ /dev/null @@ -1,157 +0,0 @@ -// Copyright (c) Abstract Machines - -// SPDX-License-Identifier: Apache-2.0 - -// Code generated by mockery; DO NOT EDIT. -// github.com/vektra/mockery -// template: testify - -package mocks - -import ( - mock "github.com/stretchr/testify/mock" -) - -// NewHasher creates a new instance of Hasher. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. -// The first argument is typically a *testing.T value. -func NewHasher(t interface { - mock.TestingT - Cleanup(func()) -}) *Hasher { - mock := &Hasher{} - mock.Mock.Test(t) - - t.Cleanup(func() { mock.AssertExpectations(t) }) - - return mock -} - -// Hasher is an autogenerated mock type for the Hasher type -type Hasher struct { - mock.Mock -} - -type Hasher_Expecter struct { - mock *mock.Mock -} - -func (_m *Hasher) EXPECT() *Hasher_Expecter { - return &Hasher_Expecter{mock: &_m.Mock} -} - -// Compare provides a mock function for the type Hasher -func (_mock *Hasher) Compare(s string, s1 string) error { - ret := _mock.Called(s, s1) - - if len(ret) == 0 { - panic("no return value specified for Compare") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(string, string) error); ok { - r0 = returnFunc(s, s1) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Hasher_Compare_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Compare' -type Hasher_Compare_Call struct { - *mock.Call -} - -// Compare is a helper method to define mock.On call -// - s string -// - s1 string -func (_e *Hasher_Expecter) Compare(s interface{}, s1 interface{}) *Hasher_Compare_Call { - return &Hasher_Compare_Call{Call: _e.mock.On("Compare", s, s1)} -} - -func (_c *Hasher_Compare_Call) Run(run func(s string, s1 string)) *Hasher_Compare_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 string - if args[0] != nil { - arg0 = args[0].(string) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Hasher_Compare_Call) Return(err error) *Hasher_Compare_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Hasher_Compare_Call) RunAndReturn(run func(s string, s1 string) error) *Hasher_Compare_Call { - _c.Call.Return(run) - return _c -} - -// Hash provides a mock function for the type Hasher -func (_mock *Hasher) Hash(s string) (string, error) { - ret := _mock.Called(s) - - if len(ret) == 0 { - panic("no return value specified for Hash") - } - - var r0 string - var r1 error - if returnFunc, ok := ret.Get(0).(func(string) (string, error)); ok { - return returnFunc(s) - } - if returnFunc, ok := ret.Get(0).(func(string) string); ok { - r0 = returnFunc(s) - } else { - r0 = ret.Get(0).(string) - } - if returnFunc, ok := ret.Get(1).(func(string) error); ok { - r1 = returnFunc(s) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Hasher_Hash_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Hash' -type Hasher_Hash_Call struct { - *mock.Call -} - -// Hash is a helper method to define mock.On call -// - s string -func (_e *Hasher_Expecter) Hash(s interface{}) *Hasher_Hash_Call { - return &Hasher_Hash_Call{Call: _e.mock.On("Hash", s)} -} - -func (_c *Hasher_Hash_Call) Run(run func(s string)) *Hasher_Hash_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 string - if args[0] != nil { - arg0 = args[0].(string) - } - run( - arg0, - ) - }) - return _c -} - -func (_c *Hasher_Hash_Call) Return(s1 string, err error) *Hasher_Hash_Call { - _c.Call.Return(s1, err) - return _c -} - -func (_c *Hasher_Hash_Call) RunAndReturn(run func(s string) (string, error)) *Hasher_Hash_Call { - _c.Call.Return(run) - return _c -} diff --git a/users/mocks/repository.go b/users/mocks/repository.go deleted file mode 100644 index 245006473..000000000 --- a/users/mocks/repository.go +++ /dev/null @@ -1,1273 +0,0 @@ -// Copyright (c) Abstract Machines - -// SPDX-License-Identifier: Apache-2.0 - -// Code generated by mockery; DO NOT EDIT. -// github.com/vektra/mockery -// template: testify - -package mocks - -import ( - "context" - - "github.com/absmach/magistrala/users" - mock "github.com/stretchr/testify/mock" -) - -// NewRepository creates a new instance of Repository. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. -// The first argument is typically a *testing.T value. -func NewRepository(t interface { - mock.TestingT - Cleanup(func()) -}) *Repository { - mock := &Repository{} - mock.Mock.Test(t) - - t.Cleanup(func() { mock.AssertExpectations(t) }) - - return mock -} - -// Repository is an autogenerated mock type for the Repository type -type Repository struct { - mock.Mock -} - -type Repository_Expecter struct { - mock *mock.Mock -} - -func (_m *Repository) EXPECT() *Repository_Expecter { - return &Repository_Expecter{mock: &_m.Mock} -} - -// AddUserVerification provides a mock function for the type Repository -func (_mock *Repository) AddUserVerification(ctx context.Context, uv users.UserVerification) error { - ret := _mock.Called(ctx, uv) - - if len(ret) == 0 { - panic("no return value specified for AddUserVerification") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, users.UserVerification) error); ok { - r0 = returnFunc(ctx, uv) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_AddUserVerification_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'AddUserVerification' -type Repository_AddUserVerification_Call struct { - *mock.Call -} - -// AddUserVerification is a helper method to define mock.On call -// - ctx context.Context -// - uv users.UserVerification -func (_e *Repository_Expecter) AddUserVerification(ctx interface{}, uv interface{}) *Repository_AddUserVerification_Call { - return &Repository_AddUserVerification_Call{Call: _e.mock.On("AddUserVerification", ctx, uv)} -} - -func (_c *Repository_AddUserVerification_Call) Run(run func(ctx context.Context, uv users.UserVerification)) *Repository_AddUserVerification_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 users.UserVerification - if args[1] != nil { - arg1 = args[1].(users.UserVerification) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_AddUserVerification_Call) Return(err error) *Repository_AddUserVerification_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_AddUserVerification_Call) RunAndReturn(run func(ctx context.Context, uv users.UserVerification) error) *Repository_AddUserVerification_Call { - _c.Call.Return(run) - return _c -} - -// ChangeStatus provides a mock function for the type Repository -func (_mock *Repository) ChangeStatus(ctx context.Context, user users.User) (users.User, error) { - ret := _mock.Called(ctx, user) - - if len(ret) == 0 { - panic("no return value specified for ChangeStatus") - } - - var r0 users.User - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, users.User) (users.User, error)); ok { - return returnFunc(ctx, user) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, users.User) users.User); ok { - r0 = returnFunc(ctx, user) - } else { - r0 = ret.Get(0).(users.User) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, users.User) error); ok { - r1 = returnFunc(ctx, user) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_ChangeStatus_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ChangeStatus' -type Repository_ChangeStatus_Call struct { - *mock.Call -} - -// ChangeStatus is a helper method to define mock.On call -// - ctx context.Context -// - user users.User -func (_e *Repository_Expecter) ChangeStatus(ctx interface{}, user interface{}) *Repository_ChangeStatus_Call { - return &Repository_ChangeStatus_Call{Call: _e.mock.On("ChangeStatus", ctx, user)} -} - -func (_c *Repository_ChangeStatus_Call) Run(run func(ctx context.Context, user users.User)) *Repository_ChangeStatus_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 users.User - if args[1] != nil { - arg1 = args[1].(users.User) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_ChangeStatus_Call) Return(user1 users.User, err error) *Repository_ChangeStatus_Call { - _c.Call.Return(user1, err) - return _c -} - -func (_c *Repository_ChangeStatus_Call) RunAndReturn(run func(ctx context.Context, user users.User) (users.User, error)) *Repository_ChangeStatus_Call { - _c.Call.Return(run) - return _c -} - -// CheckSuperAdmin provides a mock function for the type Repository -func (_mock *Repository) CheckSuperAdmin(ctx context.Context, adminID string) error { - ret := _mock.Called(ctx, adminID) - - if len(ret) == 0 { - panic("no return value specified for CheckSuperAdmin") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) error); ok { - r0 = returnFunc(ctx, adminID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_CheckSuperAdmin_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'CheckSuperAdmin' -type Repository_CheckSuperAdmin_Call struct { - *mock.Call -} - -// CheckSuperAdmin is a helper method to define mock.On call -// - ctx context.Context -// - adminID string -func (_e *Repository_Expecter) CheckSuperAdmin(ctx interface{}, adminID interface{}) *Repository_CheckSuperAdmin_Call { - return &Repository_CheckSuperAdmin_Call{Call: _e.mock.On("CheckSuperAdmin", ctx, adminID)} -} - -func (_c *Repository_CheckSuperAdmin_Call) Run(run func(ctx context.Context, adminID string)) *Repository_CheckSuperAdmin_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_CheckSuperAdmin_Call) Return(err error) *Repository_CheckSuperAdmin_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_CheckSuperAdmin_Call) RunAndReturn(run func(ctx context.Context, adminID string) error) *Repository_CheckSuperAdmin_Call { - _c.Call.Return(run) - return _c -} - -// Delete provides a mock function for the type Repository -func (_mock *Repository) Delete(ctx context.Context, id string) error { - ret := _mock.Called(ctx, id) - - if len(ret) == 0 { - panic("no return value specified for Delete") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) error); ok { - r0 = returnFunc(ctx, id) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_Delete_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Delete' -type Repository_Delete_Call struct { - *mock.Call -} - -// Delete is a helper method to define mock.On call -// - ctx context.Context -// - id string -func (_e *Repository_Expecter) Delete(ctx interface{}, id interface{}) *Repository_Delete_Call { - return &Repository_Delete_Call{Call: _e.mock.On("Delete", ctx, id)} -} - -func (_c *Repository_Delete_Call) Run(run func(ctx context.Context, id string)) *Repository_Delete_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_Delete_Call) Return(err error) *Repository_Delete_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_Delete_Call) RunAndReturn(run func(ctx context.Context, id string) error) *Repository_Delete_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveAll provides a mock function for the type Repository -func (_mock *Repository) RetrieveAll(ctx context.Context, pm users.Page) (users.UsersPage, error) { - ret := _mock.Called(ctx, pm) - - if len(ret) == 0 { - panic("no return value specified for RetrieveAll") - } - - var r0 users.UsersPage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, users.Page) (users.UsersPage, error)); ok { - return returnFunc(ctx, pm) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, users.Page) users.UsersPage); ok { - r0 = returnFunc(ctx, pm) - } else { - r0 = ret.Get(0).(users.UsersPage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, users.Page) error); ok { - r1 = returnFunc(ctx, pm) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveAll_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveAll' -type Repository_RetrieveAll_Call struct { - *mock.Call -} - -// RetrieveAll is a helper method to define mock.On call -// - ctx context.Context -// - pm users.Page -func (_e *Repository_Expecter) RetrieveAll(ctx interface{}, pm interface{}) *Repository_RetrieveAll_Call { - return &Repository_RetrieveAll_Call{Call: _e.mock.On("RetrieveAll", ctx, pm)} -} - -func (_c *Repository_RetrieveAll_Call) Run(run func(ctx context.Context, pm users.Page)) *Repository_RetrieveAll_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 users.Page - if args[1] != nil { - arg1 = args[1].(users.Page) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RetrieveAll_Call) Return(usersPage users.UsersPage, err error) *Repository_RetrieveAll_Call { - _c.Call.Return(usersPage, err) - return _c -} - -func (_c *Repository_RetrieveAll_Call) RunAndReturn(run func(ctx context.Context, pm users.Page) (users.UsersPage, error)) *Repository_RetrieveAll_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveAllByIDs provides a mock function for the type Repository -func (_mock *Repository) RetrieveAllByIDs(ctx context.Context, pm users.Page) (users.UsersPage, error) { - ret := _mock.Called(ctx, pm) - - if len(ret) == 0 { - panic("no return value specified for RetrieveAllByIDs") - } - - var r0 users.UsersPage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, users.Page) (users.UsersPage, error)); ok { - return returnFunc(ctx, pm) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, users.Page) users.UsersPage); ok { - r0 = returnFunc(ctx, pm) - } else { - r0 = ret.Get(0).(users.UsersPage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, users.Page) error); ok { - r1 = returnFunc(ctx, pm) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveAllByIDs_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveAllByIDs' -type Repository_RetrieveAllByIDs_Call struct { - *mock.Call -} - -// RetrieveAllByIDs is a helper method to define mock.On call -// - ctx context.Context -// - pm users.Page -func (_e *Repository_Expecter) RetrieveAllByIDs(ctx interface{}, pm interface{}) *Repository_RetrieveAllByIDs_Call { - return &Repository_RetrieveAllByIDs_Call{Call: _e.mock.On("RetrieveAllByIDs", ctx, pm)} -} - -func (_c *Repository_RetrieveAllByIDs_Call) Run(run func(ctx context.Context, pm users.Page)) *Repository_RetrieveAllByIDs_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 users.Page - if args[1] != nil { - arg1 = args[1].(users.Page) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RetrieveAllByIDs_Call) Return(usersPage users.UsersPage, err error) *Repository_RetrieveAllByIDs_Call { - _c.Call.Return(usersPage, err) - return _c -} - -func (_c *Repository_RetrieveAllByIDs_Call) RunAndReturn(run func(ctx context.Context, pm users.Page) (users.UsersPage, error)) *Repository_RetrieveAllByIDs_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveByEmail provides a mock function for the type Repository -func (_mock *Repository) RetrieveByEmail(ctx context.Context, email string) (users.User, error) { - ret := _mock.Called(ctx, email) - - if len(ret) == 0 { - panic("no return value specified for RetrieveByEmail") - } - - var r0 users.User - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) (users.User, error)); ok { - return returnFunc(ctx, email) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string) users.User); ok { - r0 = returnFunc(ctx, email) - } else { - r0 = ret.Get(0).(users.User) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = returnFunc(ctx, email) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveByEmail_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveByEmail' -type Repository_RetrieveByEmail_Call struct { - *mock.Call -} - -// RetrieveByEmail is a helper method to define mock.On call -// - ctx context.Context -// - email string -func (_e *Repository_Expecter) RetrieveByEmail(ctx interface{}, email interface{}) *Repository_RetrieveByEmail_Call { - return &Repository_RetrieveByEmail_Call{Call: _e.mock.On("RetrieveByEmail", ctx, email)} -} - -func (_c *Repository_RetrieveByEmail_Call) Run(run func(ctx context.Context, email string)) *Repository_RetrieveByEmail_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RetrieveByEmail_Call) Return(user users.User, err error) *Repository_RetrieveByEmail_Call { - _c.Call.Return(user, err) - return _c -} - -func (_c *Repository_RetrieveByEmail_Call) RunAndReturn(run func(ctx context.Context, email string) (users.User, error)) *Repository_RetrieveByEmail_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveByID provides a mock function for the type Repository -func (_mock *Repository) RetrieveByID(ctx context.Context, id string) (users.User, error) { - ret := _mock.Called(ctx, id) - - if len(ret) == 0 { - panic("no return value specified for RetrieveByID") - } - - var r0 users.User - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) (users.User, error)); ok { - return returnFunc(ctx, id) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string) users.User); ok { - r0 = returnFunc(ctx, id) - } else { - r0 = ret.Get(0).(users.User) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = returnFunc(ctx, id) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveByID_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveByID' -type Repository_RetrieveByID_Call struct { - *mock.Call -} - -// RetrieveByID is a helper method to define mock.On call -// - ctx context.Context -// - id string -func (_e *Repository_Expecter) RetrieveByID(ctx interface{}, id interface{}) *Repository_RetrieveByID_Call { - return &Repository_RetrieveByID_Call{Call: _e.mock.On("RetrieveByID", ctx, id)} -} - -func (_c *Repository_RetrieveByID_Call) Run(run func(ctx context.Context, id string)) *Repository_RetrieveByID_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RetrieveByID_Call) Return(user users.User, err error) *Repository_RetrieveByID_Call { - _c.Call.Return(user, err) - return _c -} - -func (_c *Repository_RetrieveByID_Call) RunAndReturn(run func(ctx context.Context, id string) (users.User, error)) *Repository_RetrieveByID_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveByUsername provides a mock function for the type Repository -func (_mock *Repository) RetrieveByUsername(ctx context.Context, username string) (users.User, error) { - ret := _mock.Called(ctx, username) - - if len(ret) == 0 { - panic("no return value specified for RetrieveByUsername") - } - - var r0 users.User - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) (users.User, error)); ok { - return returnFunc(ctx, username) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string) users.User); ok { - r0 = returnFunc(ctx, username) - } else { - r0 = ret.Get(0).(users.User) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = returnFunc(ctx, username) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveByUsername_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveByUsername' -type Repository_RetrieveByUsername_Call struct { - *mock.Call -} - -// RetrieveByUsername is a helper method to define mock.On call -// - ctx context.Context -// - username string -func (_e *Repository_Expecter) RetrieveByUsername(ctx interface{}, username interface{}) *Repository_RetrieveByUsername_Call { - return &Repository_RetrieveByUsername_Call{Call: _e.mock.On("RetrieveByUsername", ctx, username)} -} - -func (_c *Repository_RetrieveByUsername_Call) Run(run func(ctx context.Context, username string)) *Repository_RetrieveByUsername_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_RetrieveByUsername_Call) Return(user users.User, err error) *Repository_RetrieveByUsername_Call { - _c.Call.Return(user, err) - return _c -} - -func (_c *Repository_RetrieveByUsername_Call) RunAndReturn(run func(ctx context.Context, username string) (users.User, error)) *Repository_RetrieveByUsername_Call { - _c.Call.Return(run) - return _c -} - -// RetrieveUserVerification provides a mock function for the type Repository -func (_mock *Repository) RetrieveUserVerification(ctx context.Context, userID string, email string) (users.UserVerification, error) { - ret := _mock.Called(ctx, userID, email) - - if len(ret) == 0 { - panic("no return value specified for RetrieveUserVerification") - } - - var r0 users.UserVerification - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) (users.UserVerification, error)); ok { - return returnFunc(ctx, userID, email) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string) users.UserVerification); ok { - r0 = returnFunc(ctx, userID, email) - } else { - r0 = ret.Get(0).(users.UserVerification) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, string) error); ok { - r1 = returnFunc(ctx, userID, email) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_RetrieveUserVerification_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveUserVerification' -type Repository_RetrieveUserVerification_Call struct { - *mock.Call -} - -// RetrieveUserVerification is a helper method to define mock.On call -// - ctx context.Context -// - userID string -// - email string -func (_e *Repository_Expecter) RetrieveUserVerification(ctx interface{}, userID interface{}, email interface{}) *Repository_RetrieveUserVerification_Call { - return &Repository_RetrieveUserVerification_Call{Call: _e.mock.On("RetrieveUserVerification", ctx, userID, email)} -} - -func (_c *Repository_RetrieveUserVerification_Call) Run(run func(ctx context.Context, userID string, email string)) *Repository_RetrieveUserVerification_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_RetrieveUserVerification_Call) Return(userVerification users.UserVerification, err error) *Repository_RetrieveUserVerification_Call { - _c.Call.Return(userVerification, err) - return _c -} - -func (_c *Repository_RetrieveUserVerification_Call) RunAndReturn(run func(ctx context.Context, userID string, email string) (users.UserVerification, error)) *Repository_RetrieveUserVerification_Call { - _c.Call.Return(run) - return _c -} - -// Save provides a mock function for the type Repository -func (_mock *Repository) Save(ctx context.Context, user users.User) (users.User, error) { - ret := _mock.Called(ctx, user) - - if len(ret) == 0 { - panic("no return value specified for Save") - } - - var r0 users.User - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, users.User) (users.User, error)); ok { - return returnFunc(ctx, user) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, users.User) users.User); ok { - r0 = returnFunc(ctx, user) - } else { - r0 = ret.Get(0).(users.User) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, users.User) error); ok { - r1 = returnFunc(ctx, user) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_Save_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Save' -type Repository_Save_Call struct { - *mock.Call -} - -// Save is a helper method to define mock.On call -// - ctx context.Context -// - user users.User -func (_e *Repository_Expecter) Save(ctx interface{}, user interface{}) *Repository_Save_Call { - return &Repository_Save_Call{Call: _e.mock.On("Save", ctx, user)} -} - -func (_c *Repository_Save_Call) Run(run func(ctx context.Context, user users.User)) *Repository_Save_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 users.User - if args[1] != nil { - arg1 = args[1].(users.User) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_Save_Call) Return(user1 users.User, err error) *Repository_Save_Call { - _c.Call.Return(user1, err) - return _c -} - -func (_c *Repository_Save_Call) RunAndReturn(run func(ctx context.Context, user users.User) (users.User, error)) *Repository_Save_Call { - _c.Call.Return(run) - return _c -} - -// SearchUsers provides a mock function for the type Repository -func (_mock *Repository) SearchUsers(ctx context.Context, pm users.Page) (users.UsersPage, error) { - ret := _mock.Called(ctx, pm) - - if len(ret) == 0 { - panic("no return value specified for SearchUsers") - } - - var r0 users.UsersPage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, users.Page) (users.UsersPage, error)); ok { - return returnFunc(ctx, pm) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, users.Page) users.UsersPage); ok { - r0 = returnFunc(ctx, pm) - } else { - r0 = ret.Get(0).(users.UsersPage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, users.Page) error); ok { - r1 = returnFunc(ctx, pm) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_SearchUsers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'SearchUsers' -type Repository_SearchUsers_Call struct { - *mock.Call -} - -// SearchUsers is a helper method to define mock.On call -// - ctx context.Context -// - pm users.Page -func (_e *Repository_Expecter) SearchUsers(ctx interface{}, pm interface{}) *Repository_SearchUsers_Call { - return &Repository_SearchUsers_Call{Call: _e.mock.On("SearchUsers", ctx, pm)} -} - -func (_c *Repository_SearchUsers_Call) Run(run func(ctx context.Context, pm users.Page)) *Repository_SearchUsers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 users.Page - if args[1] != nil { - arg1 = args[1].(users.Page) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_SearchUsers_Call) Return(usersPage users.UsersPage, err error) *Repository_SearchUsers_Call { - _c.Call.Return(usersPage, err) - return _c -} - -func (_c *Repository_SearchUsers_Call) RunAndReturn(run func(ctx context.Context, pm users.Page) (users.UsersPage, error)) *Repository_SearchUsers_Call { - _c.Call.Return(run) - return _c -} - -// Update provides a mock function for the type Repository -func (_mock *Repository) Update(ctx context.Context, id string, user users.UserReq) (users.User, error) { - ret := _mock.Called(ctx, id, user) - - if len(ret) == 0 { - panic("no return value specified for Update") - } - - var r0 users.User - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, users.UserReq) (users.User, error)); ok { - return returnFunc(ctx, id, user) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, users.UserReq) users.User); ok { - r0 = returnFunc(ctx, id, user) - } else { - r0 = ret.Get(0).(users.User) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, users.UserReq) error); ok { - r1 = returnFunc(ctx, id, user) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_Update_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Update' -type Repository_Update_Call struct { - *mock.Call -} - -// Update is a helper method to define mock.On call -// - ctx context.Context -// - id string -// - user users.UserReq -func (_e *Repository_Expecter) Update(ctx interface{}, id interface{}, user interface{}) *Repository_Update_Call { - return &Repository_Update_Call{Call: _e.mock.On("Update", ctx, id, user)} -} - -func (_c *Repository_Update_Call) Run(run func(ctx context.Context, id string, user users.UserReq)) *Repository_Update_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 users.UserReq - if args[2] != nil { - arg2 = args[2].(users.UserReq) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Repository_Update_Call) Return(user1 users.User, err error) *Repository_Update_Call { - _c.Call.Return(user1, err) - return _c -} - -func (_c *Repository_Update_Call) RunAndReturn(run func(ctx context.Context, id string, user users.UserReq) (users.User, error)) *Repository_Update_Call { - _c.Call.Return(run) - return _c -} - -// UpdateEmail provides a mock function for the type Repository -func (_mock *Repository) UpdateEmail(ctx context.Context, user users.User) (users.User, error) { - ret := _mock.Called(ctx, user) - - if len(ret) == 0 { - panic("no return value specified for UpdateEmail") - } - - var r0 users.User - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, users.User) (users.User, error)); ok { - return returnFunc(ctx, user) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, users.User) users.User); ok { - r0 = returnFunc(ctx, user) - } else { - r0 = ret.Get(0).(users.User) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, users.User) error); ok { - r1 = returnFunc(ctx, user) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_UpdateEmail_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateEmail' -type Repository_UpdateEmail_Call struct { - *mock.Call -} - -// UpdateEmail is a helper method to define mock.On call -// - ctx context.Context -// - user users.User -func (_e *Repository_Expecter) UpdateEmail(ctx interface{}, user interface{}) *Repository_UpdateEmail_Call { - return &Repository_UpdateEmail_Call{Call: _e.mock.On("UpdateEmail", ctx, user)} -} - -func (_c *Repository_UpdateEmail_Call) Run(run func(ctx context.Context, user users.User)) *Repository_UpdateEmail_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 users.User - if args[1] != nil { - arg1 = args[1].(users.User) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_UpdateEmail_Call) Return(user1 users.User, err error) *Repository_UpdateEmail_Call { - _c.Call.Return(user1, err) - return _c -} - -func (_c *Repository_UpdateEmail_Call) RunAndReturn(run func(ctx context.Context, user users.User) (users.User, error)) *Repository_UpdateEmail_Call { - _c.Call.Return(run) - return _c -} - -// UpdateRole provides a mock function for the type Repository -func (_mock *Repository) UpdateRole(ctx context.Context, user users.User) (users.User, error) { - ret := _mock.Called(ctx, user) - - if len(ret) == 0 { - panic("no return value specified for UpdateRole") - } - - var r0 users.User - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, users.User) (users.User, error)); ok { - return returnFunc(ctx, user) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, users.User) users.User); ok { - r0 = returnFunc(ctx, user) - } else { - r0 = ret.Get(0).(users.User) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, users.User) error); ok { - r1 = returnFunc(ctx, user) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_UpdateRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateRole' -type Repository_UpdateRole_Call struct { - *mock.Call -} - -// UpdateRole is a helper method to define mock.On call -// - ctx context.Context -// - user users.User -func (_e *Repository_Expecter) UpdateRole(ctx interface{}, user interface{}) *Repository_UpdateRole_Call { - return &Repository_UpdateRole_Call{Call: _e.mock.On("UpdateRole", ctx, user)} -} - -func (_c *Repository_UpdateRole_Call) Run(run func(ctx context.Context, user users.User)) *Repository_UpdateRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 users.User - if args[1] != nil { - arg1 = args[1].(users.User) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_UpdateRole_Call) Return(user1 users.User, err error) *Repository_UpdateRole_Call { - _c.Call.Return(user1, err) - return _c -} - -func (_c *Repository_UpdateRole_Call) RunAndReturn(run func(ctx context.Context, user users.User) (users.User, error)) *Repository_UpdateRole_Call { - _c.Call.Return(run) - return _c -} - -// UpdateSecret provides a mock function for the type Repository -func (_mock *Repository) UpdateSecret(ctx context.Context, user users.User) (users.User, error) { - ret := _mock.Called(ctx, user) - - if len(ret) == 0 { - panic("no return value specified for UpdateSecret") - } - - var r0 users.User - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, users.User) (users.User, error)); ok { - return returnFunc(ctx, user) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, users.User) users.User); ok { - r0 = returnFunc(ctx, user) - } else { - r0 = ret.Get(0).(users.User) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, users.User) error); ok { - r1 = returnFunc(ctx, user) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_UpdateSecret_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateSecret' -type Repository_UpdateSecret_Call struct { - *mock.Call -} - -// UpdateSecret is a helper method to define mock.On call -// - ctx context.Context -// - user users.User -func (_e *Repository_Expecter) UpdateSecret(ctx interface{}, user interface{}) *Repository_UpdateSecret_Call { - return &Repository_UpdateSecret_Call{Call: _e.mock.On("UpdateSecret", ctx, user)} -} - -func (_c *Repository_UpdateSecret_Call) Run(run func(ctx context.Context, user users.User)) *Repository_UpdateSecret_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 users.User - if args[1] != nil { - arg1 = args[1].(users.User) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_UpdateSecret_Call) Return(user1 users.User, err error) *Repository_UpdateSecret_Call { - _c.Call.Return(user1, err) - return _c -} - -func (_c *Repository_UpdateSecret_Call) RunAndReturn(run func(ctx context.Context, user users.User) (users.User, error)) *Repository_UpdateSecret_Call { - _c.Call.Return(run) - return _c -} - -// UpdateUserVerification provides a mock function for the type Repository -func (_mock *Repository) UpdateUserVerification(ctx context.Context, uv users.UserVerification) error { - ret := _mock.Called(ctx, uv) - - if len(ret) == 0 { - panic("no return value specified for UpdateUserVerification") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, users.UserVerification) error); ok { - r0 = returnFunc(ctx, uv) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Repository_UpdateUserVerification_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateUserVerification' -type Repository_UpdateUserVerification_Call struct { - *mock.Call -} - -// UpdateUserVerification is a helper method to define mock.On call -// - ctx context.Context -// - uv users.UserVerification -func (_e *Repository_Expecter) UpdateUserVerification(ctx interface{}, uv interface{}) *Repository_UpdateUserVerification_Call { - return &Repository_UpdateUserVerification_Call{Call: _e.mock.On("UpdateUserVerification", ctx, uv)} -} - -func (_c *Repository_UpdateUserVerification_Call) Run(run func(ctx context.Context, uv users.UserVerification)) *Repository_UpdateUserVerification_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 users.UserVerification - if args[1] != nil { - arg1 = args[1].(users.UserVerification) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_UpdateUserVerification_Call) Return(err error) *Repository_UpdateUserVerification_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Repository_UpdateUserVerification_Call) RunAndReturn(run func(ctx context.Context, uv users.UserVerification) error) *Repository_UpdateUserVerification_Call { - _c.Call.Return(run) - return _c -} - -// UpdateUsername provides a mock function for the type Repository -func (_mock *Repository) UpdateUsername(ctx context.Context, user users.User) (users.User, error) { - ret := _mock.Called(ctx, user) - - if len(ret) == 0 { - panic("no return value specified for UpdateUsername") - } - - var r0 users.User - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, users.User) (users.User, error)); ok { - return returnFunc(ctx, user) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, users.User) users.User); ok { - r0 = returnFunc(ctx, user) - } else { - r0 = ret.Get(0).(users.User) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, users.User) error); ok { - r1 = returnFunc(ctx, user) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_UpdateUsername_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateUsername' -type Repository_UpdateUsername_Call struct { - *mock.Call -} - -// UpdateUsername is a helper method to define mock.On call -// - ctx context.Context -// - user users.User -func (_e *Repository_Expecter) UpdateUsername(ctx interface{}, user interface{}) *Repository_UpdateUsername_Call { - return &Repository_UpdateUsername_Call{Call: _e.mock.On("UpdateUsername", ctx, user)} -} - -func (_c *Repository_UpdateUsername_Call) Run(run func(ctx context.Context, user users.User)) *Repository_UpdateUsername_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 users.User - if args[1] != nil { - arg1 = args[1].(users.User) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_UpdateUsername_Call) Return(user1 users.User, err error) *Repository_UpdateUsername_Call { - _c.Call.Return(user1, err) - return _c -} - -func (_c *Repository_UpdateUsername_Call) RunAndReturn(run func(ctx context.Context, user users.User) (users.User, error)) *Repository_UpdateUsername_Call { - _c.Call.Return(run) - return _c -} - -// UpdateVerifiedAt provides a mock function for the type Repository -func (_mock *Repository) UpdateVerifiedAt(ctx context.Context, user users.User) (users.User, error) { - ret := _mock.Called(ctx, user) - - if len(ret) == 0 { - panic("no return value specified for UpdateVerifiedAt") - } - - var r0 users.User - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, users.User) (users.User, error)); ok { - return returnFunc(ctx, user) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, users.User) users.User); ok { - r0 = returnFunc(ctx, user) - } else { - r0 = ret.Get(0).(users.User) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, users.User) error); ok { - r1 = returnFunc(ctx, user) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Repository_UpdateVerifiedAt_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateVerifiedAt' -type Repository_UpdateVerifiedAt_Call struct { - *mock.Call -} - -// UpdateVerifiedAt is a helper method to define mock.On call -// - ctx context.Context -// - user users.User -func (_e *Repository_Expecter) UpdateVerifiedAt(ctx interface{}, user interface{}) *Repository_UpdateVerifiedAt_Call { - return &Repository_UpdateVerifiedAt_Call{Call: _e.mock.On("UpdateVerifiedAt", ctx, user)} -} - -func (_c *Repository_UpdateVerifiedAt_Call) Run(run func(ctx context.Context, user users.User)) *Repository_UpdateVerifiedAt_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 users.User - if args[1] != nil { - arg1 = args[1].(users.User) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Repository_UpdateVerifiedAt_Call) Return(user1 users.User, err error) *Repository_UpdateVerifiedAt_Call { - _c.Call.Return(user1, err) - return _c -} - -func (_c *Repository_UpdateVerifiedAt_Call) RunAndReturn(run func(ctx context.Context, user users.User) (users.User, error)) *Repository_UpdateVerifiedAt_Call { - _c.Call.Return(run) - return _c -} diff --git a/users/mocks/service.go b/users/mocks/service.go deleted file mode 100644 index 1792e32f1..000000000 --- a/users/mocks/service.go +++ /dev/null @@ -1,1863 +0,0 @@ -// Copyright (c) Abstract Machines - -// SPDX-License-Identifier: Apache-2.0 - -// Code generated by mockery; DO NOT EDIT. -// github.com/vektra/mockery -// template: testify - -package mocks - -import ( - "context" - - "github.com/absmach/magistrala/api/grpc/token/v1" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/users" - mock "github.com/stretchr/testify/mock" -) - -// NewService creates a new instance of Service. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. -// The first argument is typically a *testing.T value. -func NewService(t interface { - mock.TestingT - Cleanup(func()) -}) *Service { - mock := &Service{} - mock.Mock.Test(t) - - t.Cleanup(func() { mock.AssertExpectations(t) }) - - return mock -} - -// Service is an autogenerated mock type for the Service type -type Service struct { - mock.Mock -} - -type Service_Expecter struct { - mock *mock.Mock -} - -func (_m *Service) EXPECT() *Service_Expecter { - return &Service_Expecter{mock: &_m.Mock} -} - -// Delete provides a mock function for the type Service -func (_mock *Service) Delete(ctx context.Context, session authn.Session, id string) error { - ret := _mock.Called(ctx, session, id) - - if len(ret) == 0 { - panic("no return value specified for Delete") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) error); ok { - r0 = returnFunc(ctx, session, id) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_Delete_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Delete' -type Service_Delete_Call struct { - *mock.Call -} - -// Delete is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - id string -func (_e *Service_Expecter) Delete(ctx interface{}, session interface{}, id interface{}) *Service_Delete_Call { - return &Service_Delete_Call{Call: _e.mock.On("Delete", ctx, session, id)} -} - -func (_c *Service_Delete_Call) Run(run func(ctx context.Context, session authn.Session, id string)) *Service_Delete_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_Delete_Call) Return(err error) *Service_Delete_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_Delete_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, id string) error) *Service_Delete_Call { - _c.Call.Return(run) - return _c -} - -// Disable provides a mock function for the type Service -func (_mock *Service) Disable(ctx context.Context, session authn.Session, id string) (users.User, error) { - ret := _mock.Called(ctx, session, id) - - if len(ret) == 0 { - panic("no return value specified for Disable") - } - - var r0 users.User - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) (users.User, error)); ok { - return returnFunc(ctx, session, id) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) users.User); ok { - r0 = returnFunc(ctx, session, id) - } else { - r0 = ret.Get(0).(users.User) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string) error); ok { - r1 = returnFunc(ctx, session, id) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_Disable_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Disable' -type Service_Disable_Call struct { - *mock.Call -} - -// Disable is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - id string -func (_e *Service_Expecter) Disable(ctx interface{}, session interface{}, id interface{}) *Service_Disable_Call { - return &Service_Disable_Call{Call: _e.mock.On("Disable", ctx, session, id)} -} - -func (_c *Service_Disable_Call) Run(run func(ctx context.Context, session authn.Session, id string)) *Service_Disable_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_Disable_Call) Return(user users.User, err error) *Service_Disable_Call { - _c.Call.Return(user, err) - return _c -} - -func (_c *Service_Disable_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, id string) (users.User, error)) *Service_Disable_Call { - _c.Call.Return(run) - return _c -} - -// Enable provides a mock function for the type Service -func (_mock *Service) Enable(ctx context.Context, session authn.Session, id string) (users.User, error) { - ret := _mock.Called(ctx, session, id) - - if len(ret) == 0 { - panic("no return value specified for Enable") - } - - var r0 users.User - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) (users.User, error)); ok { - return returnFunc(ctx, session, id) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) users.User); ok { - r0 = returnFunc(ctx, session, id) - } else { - r0 = ret.Get(0).(users.User) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string) error); ok { - r1 = returnFunc(ctx, session, id) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_Enable_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Enable' -type Service_Enable_Call struct { - *mock.Call -} - -// Enable is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - id string -func (_e *Service_Expecter) Enable(ctx interface{}, session interface{}, id interface{}) *Service_Enable_Call { - return &Service_Enable_Call{Call: _e.mock.On("Enable", ctx, session, id)} -} - -func (_c *Service_Enable_Call) Run(run func(ctx context.Context, session authn.Session, id string)) *Service_Enable_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_Enable_Call) Return(user users.User, err error) *Service_Enable_Call { - _c.Call.Return(user, err) - return _c -} - -func (_c *Service_Enable_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, id string) (users.User, error)) *Service_Enable_Call { - _c.Call.Return(run) - return _c -} - -// Identify provides a mock function for the type Service -func (_mock *Service) Identify(ctx context.Context, session authn.Session) (string, error) { - ret := _mock.Called(ctx, session) - - if len(ret) == 0 { - panic("no return value specified for Identify") - } - - var r0 string - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session) (string, error)); ok { - return returnFunc(ctx, session) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session) string); ok { - r0 = returnFunc(ctx, session) - } else { - r0 = ret.Get(0).(string) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session) error); ok { - r1 = returnFunc(ctx, session) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_Identify_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Identify' -type Service_Identify_Call struct { - *mock.Call -} - -// Identify is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -func (_e *Service_Expecter) Identify(ctx interface{}, session interface{}) *Service_Identify_Call { - return &Service_Identify_Call{Call: _e.mock.On("Identify", ctx, session)} -} - -func (_c *Service_Identify_Call) Run(run func(ctx context.Context, session authn.Session)) *Service_Identify_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Service_Identify_Call) Return(s string, err error) *Service_Identify_Call { - _c.Call.Return(s, err) - return _c -} - -func (_c *Service_Identify_Call) RunAndReturn(run func(ctx context.Context, session authn.Session) (string, error)) *Service_Identify_Call { - _c.Call.Return(run) - return _c -} - -// IssueToken provides a mock function for the type Service -func (_mock *Service) IssueToken(ctx context.Context, identity string, secret string, description string) (*v1.Token, error) { - ret := _mock.Called(ctx, identity, secret, description) - - if len(ret) == 0 { - panic("no return value specified for IssueToken") - } - - var r0 *v1.Token - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string, string) (*v1.Token, error)); ok { - return returnFunc(ctx, identity, secret, description) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string, string, string) *v1.Token); ok { - r0 = returnFunc(ctx, identity, secret, description) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).(*v1.Token) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string, string, string) error); ok { - r1 = returnFunc(ctx, identity, secret, description) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_IssueToken_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'IssueToken' -type Service_IssueToken_Call struct { - *mock.Call -} - -// IssueToken is a helper method to define mock.On call -// - ctx context.Context -// - identity string -// - secret string -// - description string -func (_e *Service_Expecter) IssueToken(ctx interface{}, identity interface{}, secret interface{}, description interface{}) *Service_IssueToken_Call { - return &Service_IssueToken_Call{Call: _e.mock.On("IssueToken", ctx, identity, secret, description)} -} - -func (_c *Service_IssueToken_Call) Run(run func(ctx context.Context, identity string, secret string, description string)) *Service_IssueToken_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_IssueToken_Call) Return(token *v1.Token, err error) *Service_IssueToken_Call { - _c.Call.Return(token, err) - return _c -} - -func (_c *Service_IssueToken_Call) RunAndReturn(run func(ctx context.Context, identity string, secret string, description string) (*v1.Token, error)) *Service_IssueToken_Call { - _c.Call.Return(run) - return _c -} - -// ListActiveRefreshTokens provides a mock function for the type Service -func (_mock *Service) ListActiveRefreshTokens(ctx context.Context, session authn.Session) (*v1.ListUserRefreshTokensRes, error) { - ret := _mock.Called(ctx, session) - - if len(ret) == 0 { - panic("no return value specified for ListActiveRefreshTokens") - } - - var r0 *v1.ListUserRefreshTokensRes - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session) (*v1.ListUserRefreshTokensRes, error)); ok { - return returnFunc(ctx, session) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session) *v1.ListUserRefreshTokensRes); ok { - r0 = returnFunc(ctx, session) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).(*v1.ListUserRefreshTokensRes) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session) error); ok { - r1 = returnFunc(ctx, session) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_ListActiveRefreshTokens_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ListActiveRefreshTokens' -type Service_ListActiveRefreshTokens_Call struct { - *mock.Call -} - -// ListActiveRefreshTokens is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -func (_e *Service_Expecter) ListActiveRefreshTokens(ctx interface{}, session interface{}) *Service_ListActiveRefreshTokens_Call { - return &Service_ListActiveRefreshTokens_Call{Call: _e.mock.On("ListActiveRefreshTokens", ctx, session)} -} - -func (_c *Service_ListActiveRefreshTokens_Call) Run(run func(ctx context.Context, session authn.Session)) *Service_ListActiveRefreshTokens_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Service_ListActiveRefreshTokens_Call) Return(listUserRefreshTokensRes *v1.ListUserRefreshTokensRes, err error) *Service_ListActiveRefreshTokens_Call { - _c.Call.Return(listUserRefreshTokensRes, err) - return _c -} - -func (_c *Service_ListActiveRefreshTokens_Call) RunAndReturn(run func(ctx context.Context, session authn.Session) (*v1.ListUserRefreshTokensRes, error)) *Service_ListActiveRefreshTokens_Call { - _c.Call.Return(run) - return _c -} - -// ListUsers provides a mock function for the type Service -func (_mock *Service) ListUsers(ctx context.Context, session authn.Session, pm users.Page) (users.UsersPage, error) { - ret := _mock.Called(ctx, session, pm) - - if len(ret) == 0 { - panic("no return value specified for ListUsers") - } - - var r0 users.UsersPage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, users.Page) (users.UsersPage, error)); ok { - return returnFunc(ctx, session, pm) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, users.Page) users.UsersPage); ok { - r0 = returnFunc(ctx, session, pm) - } else { - r0 = ret.Get(0).(users.UsersPage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, users.Page) error); ok { - r1 = returnFunc(ctx, session, pm) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_ListUsers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ListUsers' -type Service_ListUsers_Call struct { - *mock.Call -} - -// ListUsers is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - pm users.Page -func (_e *Service_Expecter) ListUsers(ctx interface{}, session interface{}, pm interface{}) *Service_ListUsers_Call { - return &Service_ListUsers_Call{Call: _e.mock.On("ListUsers", ctx, session, pm)} -} - -func (_c *Service_ListUsers_Call) Run(run func(ctx context.Context, session authn.Session, pm users.Page)) *Service_ListUsers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 users.Page - if args[2] != nil { - arg2 = args[2].(users.Page) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_ListUsers_Call) Return(usersPage users.UsersPage, err error) *Service_ListUsers_Call { - _c.Call.Return(usersPage, err) - return _c -} - -func (_c *Service_ListUsers_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, pm users.Page) (users.UsersPage, error)) *Service_ListUsers_Call { - _c.Call.Return(run) - return _c -} - -// OAuthAddUserPolicy provides a mock function for the type Service -func (_mock *Service) OAuthAddUserPolicy(ctx context.Context, user users.User) error { - ret := _mock.Called(ctx, user) - - if len(ret) == 0 { - panic("no return value specified for OAuthAddUserPolicy") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, users.User) error); ok { - r0 = returnFunc(ctx, user) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_OAuthAddUserPolicy_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'OAuthAddUserPolicy' -type Service_OAuthAddUserPolicy_Call struct { - *mock.Call -} - -// OAuthAddUserPolicy is a helper method to define mock.On call -// - ctx context.Context -// - user users.User -func (_e *Service_Expecter) OAuthAddUserPolicy(ctx interface{}, user interface{}) *Service_OAuthAddUserPolicy_Call { - return &Service_OAuthAddUserPolicy_Call{Call: _e.mock.On("OAuthAddUserPolicy", ctx, user)} -} - -func (_c *Service_OAuthAddUserPolicy_Call) Run(run func(ctx context.Context, user users.User)) *Service_OAuthAddUserPolicy_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 users.User - if args[1] != nil { - arg1 = args[1].(users.User) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Service_OAuthAddUserPolicy_Call) Return(err error) *Service_OAuthAddUserPolicy_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_OAuthAddUserPolicy_Call) RunAndReturn(run func(ctx context.Context, user users.User) error) *Service_OAuthAddUserPolicy_Call { - _c.Call.Return(run) - return _c -} - -// OAuthCallback provides a mock function for the type Service -func (_mock *Service) OAuthCallback(ctx context.Context, user users.User) (users.User, error) { - ret := _mock.Called(ctx, user) - - if len(ret) == 0 { - panic("no return value specified for OAuthCallback") - } - - var r0 users.User - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, users.User) (users.User, error)); ok { - return returnFunc(ctx, user) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, users.User) users.User); ok { - r0 = returnFunc(ctx, user) - } else { - r0 = ret.Get(0).(users.User) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, users.User) error); ok { - r1 = returnFunc(ctx, user) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_OAuthCallback_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'OAuthCallback' -type Service_OAuthCallback_Call struct { - *mock.Call -} - -// OAuthCallback is a helper method to define mock.On call -// - ctx context.Context -// - user users.User -func (_e *Service_Expecter) OAuthCallback(ctx interface{}, user interface{}) *Service_OAuthCallback_Call { - return &Service_OAuthCallback_Call{Call: _e.mock.On("OAuthCallback", ctx, user)} -} - -func (_c *Service_OAuthCallback_Call) Run(run func(ctx context.Context, user users.User)) *Service_OAuthCallback_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 users.User - if args[1] != nil { - arg1 = args[1].(users.User) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Service_OAuthCallback_Call) Return(user1 users.User, err error) *Service_OAuthCallback_Call { - _c.Call.Return(user1, err) - return _c -} - -func (_c *Service_OAuthCallback_Call) RunAndReturn(run func(ctx context.Context, user users.User) (users.User, error)) *Service_OAuthCallback_Call { - _c.Call.Return(run) - return _c -} - -// RefreshToken provides a mock function for the type Service -func (_mock *Service) RefreshToken(ctx context.Context, session authn.Session, refreshToken string) (*v1.Token, error) { - ret := _mock.Called(ctx, session, refreshToken) - - if len(ret) == 0 { - panic("no return value specified for RefreshToken") - } - - var r0 *v1.Token - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) (*v1.Token, error)); ok { - return returnFunc(ctx, session, refreshToken) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) *v1.Token); ok { - r0 = returnFunc(ctx, session, refreshToken) - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).(*v1.Token) - } - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string) error); ok { - r1 = returnFunc(ctx, session, refreshToken) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_RefreshToken_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RefreshToken' -type Service_RefreshToken_Call struct { - *mock.Call -} - -// RefreshToken is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - refreshToken string -func (_e *Service_Expecter) RefreshToken(ctx interface{}, session interface{}, refreshToken interface{}) *Service_RefreshToken_Call { - return &Service_RefreshToken_Call{Call: _e.mock.On("RefreshToken", ctx, session, refreshToken)} -} - -func (_c *Service_RefreshToken_Call) Run(run func(ctx context.Context, session authn.Session, refreshToken string)) *Service_RefreshToken_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_RefreshToken_Call) Return(token *v1.Token, err error) *Service_RefreshToken_Call { - _c.Call.Return(token, err) - return _c -} - -func (_c *Service_RefreshToken_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, refreshToken string) (*v1.Token, error)) *Service_RefreshToken_Call { - _c.Call.Return(run) - return _c -} - -// Register provides a mock function for the type Service -func (_mock *Service) Register(ctx context.Context, session authn.Session, user users.User, selfRegister bool) (users.User, error) { - ret := _mock.Called(ctx, session, user, selfRegister) - - if len(ret) == 0 { - panic("no return value specified for Register") - } - - var r0 users.User - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, users.User, bool) (users.User, error)); ok { - return returnFunc(ctx, session, user, selfRegister) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, users.User, bool) users.User); ok { - r0 = returnFunc(ctx, session, user, selfRegister) - } else { - r0 = ret.Get(0).(users.User) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, users.User, bool) error); ok { - r1 = returnFunc(ctx, session, user, selfRegister) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_Register_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Register' -type Service_Register_Call struct { - *mock.Call -} - -// Register is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - user users.User -// - selfRegister bool -func (_e *Service_Expecter) Register(ctx interface{}, session interface{}, user interface{}, selfRegister interface{}) *Service_Register_Call { - return &Service_Register_Call{Call: _e.mock.On("Register", ctx, session, user, selfRegister)} -} - -func (_c *Service_Register_Call) Run(run func(ctx context.Context, session authn.Session, user users.User, selfRegister bool)) *Service_Register_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 users.User - if args[2] != nil { - arg2 = args[2].(users.User) - } - var arg3 bool - if args[3] != nil { - arg3 = args[3].(bool) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_Register_Call) Return(user1 users.User, err error) *Service_Register_Call { - _c.Call.Return(user1, err) - return _c -} - -func (_c *Service_Register_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, user users.User, selfRegister bool) (users.User, error)) *Service_Register_Call { - _c.Call.Return(run) - return _c -} - -// ResetSecret provides a mock function for the type Service -func (_mock *Service) ResetSecret(ctx context.Context, session authn.Session, secret string) error { - ret := _mock.Called(ctx, session, secret) - - if len(ret) == 0 { - panic("no return value specified for ResetSecret") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) error); ok { - r0 = returnFunc(ctx, session, secret) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_ResetSecret_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ResetSecret' -type Service_ResetSecret_Call struct { - *mock.Call -} - -// ResetSecret is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - secret string -func (_e *Service_Expecter) ResetSecret(ctx interface{}, session interface{}, secret interface{}) *Service_ResetSecret_Call { - return &Service_ResetSecret_Call{Call: _e.mock.On("ResetSecret", ctx, session, secret)} -} - -func (_c *Service_ResetSecret_Call) Run(run func(ctx context.Context, session authn.Session, secret string)) *Service_ResetSecret_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_ResetSecret_Call) Return(err error) *Service_ResetSecret_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_ResetSecret_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, secret string) error) *Service_ResetSecret_Call { - _c.Call.Return(run) - return _c -} - -// RevokeRefreshToken provides a mock function for the type Service -func (_mock *Service) RevokeRefreshToken(ctx context.Context, session authn.Session, tokenID string) error { - ret := _mock.Called(ctx, session, tokenID) - - if len(ret) == 0 { - panic("no return value specified for RevokeRefreshToken") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) error); ok { - r0 = returnFunc(ctx, session, tokenID) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_RevokeRefreshToken_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RevokeRefreshToken' -type Service_RevokeRefreshToken_Call struct { - *mock.Call -} - -// RevokeRefreshToken is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - tokenID string -func (_e *Service_Expecter) RevokeRefreshToken(ctx interface{}, session interface{}, tokenID interface{}) *Service_RevokeRefreshToken_Call { - return &Service_RevokeRefreshToken_Call{Call: _e.mock.On("RevokeRefreshToken", ctx, session, tokenID)} -} - -func (_c *Service_RevokeRefreshToken_Call) Run(run func(ctx context.Context, session authn.Session, tokenID string)) *Service_RevokeRefreshToken_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_RevokeRefreshToken_Call) Return(err error) *Service_RevokeRefreshToken_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_RevokeRefreshToken_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, tokenID string) error) *Service_RevokeRefreshToken_Call { - _c.Call.Return(run) - return _c -} - -// SearchUsers provides a mock function for the type Service -func (_mock *Service) SearchUsers(ctx context.Context, pm users.Page) (users.UsersPage, error) { - ret := _mock.Called(ctx, pm) - - if len(ret) == 0 { - panic("no return value specified for SearchUsers") - } - - var r0 users.UsersPage - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, users.Page) (users.UsersPage, error)); ok { - return returnFunc(ctx, pm) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, users.Page) users.UsersPage); ok { - r0 = returnFunc(ctx, pm) - } else { - r0 = ret.Get(0).(users.UsersPage) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, users.Page) error); ok { - r1 = returnFunc(ctx, pm) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_SearchUsers_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'SearchUsers' -type Service_SearchUsers_Call struct { - *mock.Call -} - -// SearchUsers is a helper method to define mock.On call -// - ctx context.Context -// - pm users.Page -func (_e *Service_Expecter) SearchUsers(ctx interface{}, pm interface{}) *Service_SearchUsers_Call { - return &Service_SearchUsers_Call{Call: _e.mock.On("SearchUsers", ctx, pm)} -} - -func (_c *Service_SearchUsers_Call) Run(run func(ctx context.Context, pm users.Page)) *Service_SearchUsers_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 users.Page - if args[1] != nil { - arg1 = args[1].(users.Page) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Service_SearchUsers_Call) Return(usersPage users.UsersPage, err error) *Service_SearchUsers_Call { - _c.Call.Return(usersPage, err) - return _c -} - -func (_c *Service_SearchUsers_Call) RunAndReturn(run func(ctx context.Context, pm users.Page) (users.UsersPage, error)) *Service_SearchUsers_Call { - _c.Call.Return(run) - return _c -} - -// SendPasswordReset provides a mock function for the type Service -func (_mock *Service) SendPasswordReset(ctx context.Context, email string) error { - ret := _mock.Called(ctx, email) - - if len(ret) == 0 { - panic("no return value specified for SendPasswordReset") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) error); ok { - r0 = returnFunc(ctx, email) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_SendPasswordReset_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'SendPasswordReset' -type Service_SendPasswordReset_Call struct { - *mock.Call -} - -// SendPasswordReset is a helper method to define mock.On call -// - ctx context.Context -// - email string -func (_e *Service_Expecter) SendPasswordReset(ctx interface{}, email interface{}) *Service_SendPasswordReset_Call { - return &Service_SendPasswordReset_Call{Call: _e.mock.On("SendPasswordReset", ctx, email)} -} - -func (_c *Service_SendPasswordReset_Call) Run(run func(ctx context.Context, email string)) *Service_SendPasswordReset_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Service_SendPasswordReset_Call) Return(err error) *Service_SendPasswordReset_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_SendPasswordReset_Call) RunAndReturn(run func(ctx context.Context, email string) error) *Service_SendPasswordReset_Call { - _c.Call.Return(run) - return _c -} - -// SendVerification provides a mock function for the type Service -func (_mock *Service) SendVerification(ctx context.Context, session authn.Session) error { - ret := _mock.Called(ctx, session) - - if len(ret) == 0 { - panic("no return value specified for SendVerification") - } - - var r0 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session) error); ok { - r0 = returnFunc(ctx, session) - } else { - r0 = ret.Error(0) - } - return r0 -} - -// Service_SendVerification_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'SendVerification' -type Service_SendVerification_Call struct { - *mock.Call -} - -// SendVerification is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -func (_e *Service_Expecter) SendVerification(ctx interface{}, session interface{}) *Service_SendVerification_Call { - return &Service_SendVerification_Call{Call: _e.mock.On("SendVerification", ctx, session)} -} - -func (_c *Service_SendVerification_Call) Run(run func(ctx context.Context, session authn.Session)) *Service_SendVerification_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Service_SendVerification_Call) Return(err error) *Service_SendVerification_Call { - _c.Call.Return(err) - return _c -} - -func (_c *Service_SendVerification_Call) RunAndReturn(run func(ctx context.Context, session authn.Session) error) *Service_SendVerification_Call { - _c.Call.Return(run) - return _c -} - -// Update provides a mock function for the type Service -func (_mock *Service) Update(ctx context.Context, session authn.Session, id string, user users.UserReq) (users.User, error) { - ret := _mock.Called(ctx, session, id, user) - - if len(ret) == 0 { - panic("no return value specified for Update") - } - - var r0 users.User - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, users.UserReq) (users.User, error)); ok { - return returnFunc(ctx, session, id, user) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, users.UserReq) users.User); ok { - r0 = returnFunc(ctx, session, id, user) - } else { - r0 = ret.Get(0).(users.User) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, users.UserReq) error); ok { - r1 = returnFunc(ctx, session, id, user) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_Update_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Update' -type Service_Update_Call struct { - *mock.Call -} - -// Update is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - id string -// - user users.UserReq -func (_e *Service_Expecter) Update(ctx interface{}, session interface{}, id interface{}, user interface{}) *Service_Update_Call { - return &Service_Update_Call{Call: _e.mock.On("Update", ctx, session, id, user)} -} - -func (_c *Service_Update_Call) Run(run func(ctx context.Context, session authn.Session, id string, user users.UserReq)) *Service_Update_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 users.UserReq - if args[3] != nil { - arg3 = args[3].(users.UserReq) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_Update_Call) Return(user1 users.User, err error) *Service_Update_Call { - _c.Call.Return(user1, err) - return _c -} - -func (_c *Service_Update_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, id string, user users.UserReq) (users.User, error)) *Service_Update_Call { - _c.Call.Return(run) - return _c -} - -// UpdateEmail provides a mock function for the type Service -func (_mock *Service) UpdateEmail(ctx context.Context, session authn.Session, id string, email string) (users.User, error) { - ret := _mock.Called(ctx, session, id, email) - - if len(ret) == 0 { - panic("no return value specified for UpdateEmail") - } - - var r0 users.User - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) (users.User, error)); ok { - return returnFunc(ctx, session, id, email) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) users.User); ok { - r0 = returnFunc(ctx, session, id, email) - } else { - r0 = ret.Get(0).(users.User) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string) error); ok { - r1 = returnFunc(ctx, session, id, email) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_UpdateEmail_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateEmail' -type Service_UpdateEmail_Call struct { - *mock.Call -} - -// UpdateEmail is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - id string -// - email string -func (_e *Service_Expecter) UpdateEmail(ctx interface{}, session interface{}, id interface{}, email interface{}) *Service_UpdateEmail_Call { - return &Service_UpdateEmail_Call{Call: _e.mock.On("UpdateEmail", ctx, session, id, email)} -} - -func (_c *Service_UpdateEmail_Call) Run(run func(ctx context.Context, session authn.Session, id string, email string)) *Service_UpdateEmail_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_UpdateEmail_Call) Return(user users.User, err error) *Service_UpdateEmail_Call { - _c.Call.Return(user, err) - return _c -} - -func (_c *Service_UpdateEmail_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, id string, email string) (users.User, error)) *Service_UpdateEmail_Call { - _c.Call.Return(run) - return _c -} - -// UpdateProfilePicture provides a mock function for the type Service -func (_mock *Service) UpdateProfilePicture(ctx context.Context, session authn.Session, id string, usr users.UserReq) (users.User, error) { - ret := _mock.Called(ctx, session, id, usr) - - if len(ret) == 0 { - panic("no return value specified for UpdateProfilePicture") - } - - var r0 users.User - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, users.UserReq) (users.User, error)); ok { - return returnFunc(ctx, session, id, usr) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, users.UserReq) users.User); ok { - r0 = returnFunc(ctx, session, id, usr) - } else { - r0 = ret.Get(0).(users.User) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, users.UserReq) error); ok { - r1 = returnFunc(ctx, session, id, usr) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_UpdateProfilePicture_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateProfilePicture' -type Service_UpdateProfilePicture_Call struct { - *mock.Call -} - -// UpdateProfilePicture is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - id string -// - usr users.UserReq -func (_e *Service_Expecter) UpdateProfilePicture(ctx interface{}, session interface{}, id interface{}, usr interface{}) *Service_UpdateProfilePicture_Call { - return &Service_UpdateProfilePicture_Call{Call: _e.mock.On("UpdateProfilePicture", ctx, session, id, usr)} -} - -func (_c *Service_UpdateProfilePicture_Call) Run(run func(ctx context.Context, session authn.Session, id string, usr users.UserReq)) *Service_UpdateProfilePicture_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 users.UserReq - if args[3] != nil { - arg3 = args[3].(users.UserReq) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_UpdateProfilePicture_Call) Return(user users.User, err error) *Service_UpdateProfilePicture_Call { - _c.Call.Return(user, err) - return _c -} - -func (_c *Service_UpdateProfilePicture_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, id string, usr users.UserReq) (users.User, error)) *Service_UpdateProfilePicture_Call { - _c.Call.Return(run) - return _c -} - -// UpdateRole provides a mock function for the type Service -func (_mock *Service) UpdateRole(ctx context.Context, session authn.Session, user users.User) (users.User, error) { - ret := _mock.Called(ctx, session, user) - - if len(ret) == 0 { - panic("no return value specified for UpdateRole") - } - - var r0 users.User - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, users.User) (users.User, error)); ok { - return returnFunc(ctx, session, user) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, users.User) users.User); ok { - r0 = returnFunc(ctx, session, user) - } else { - r0 = ret.Get(0).(users.User) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, users.User) error); ok { - r1 = returnFunc(ctx, session, user) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_UpdateRole_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateRole' -type Service_UpdateRole_Call struct { - *mock.Call -} - -// UpdateRole is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - user users.User -func (_e *Service_Expecter) UpdateRole(ctx interface{}, session interface{}, user interface{}) *Service_UpdateRole_Call { - return &Service_UpdateRole_Call{Call: _e.mock.On("UpdateRole", ctx, session, user)} -} - -func (_c *Service_UpdateRole_Call) Run(run func(ctx context.Context, session authn.Session, user users.User)) *Service_UpdateRole_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 users.User - if args[2] != nil { - arg2 = args[2].(users.User) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_UpdateRole_Call) Return(user1 users.User, err error) *Service_UpdateRole_Call { - _c.Call.Return(user1, err) - return _c -} - -func (_c *Service_UpdateRole_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, user users.User) (users.User, error)) *Service_UpdateRole_Call { - _c.Call.Return(run) - return _c -} - -// UpdateSecret provides a mock function for the type Service -func (_mock *Service) UpdateSecret(ctx context.Context, session authn.Session, oldSecret string, newSecret string) (users.User, error) { - ret := _mock.Called(ctx, session, oldSecret, newSecret) - - if len(ret) == 0 { - panic("no return value specified for UpdateSecret") - } - - var r0 users.User - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) (users.User, error)); ok { - return returnFunc(ctx, session, oldSecret, newSecret) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) users.User); ok { - r0 = returnFunc(ctx, session, oldSecret, newSecret) - } else { - r0 = ret.Get(0).(users.User) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string) error); ok { - r1 = returnFunc(ctx, session, oldSecret, newSecret) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_UpdateSecret_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateSecret' -type Service_UpdateSecret_Call struct { - *mock.Call -} - -// UpdateSecret is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - oldSecret string -// - newSecret string -func (_e *Service_Expecter) UpdateSecret(ctx interface{}, session interface{}, oldSecret interface{}, newSecret interface{}) *Service_UpdateSecret_Call { - return &Service_UpdateSecret_Call{Call: _e.mock.On("UpdateSecret", ctx, session, oldSecret, newSecret)} -} - -func (_c *Service_UpdateSecret_Call) Run(run func(ctx context.Context, session authn.Session, oldSecret string, newSecret string)) *Service_UpdateSecret_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_UpdateSecret_Call) Return(user users.User, err error) *Service_UpdateSecret_Call { - _c.Call.Return(user, err) - return _c -} - -func (_c *Service_UpdateSecret_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, oldSecret string, newSecret string) (users.User, error)) *Service_UpdateSecret_Call { - _c.Call.Return(run) - return _c -} - -// UpdateTags provides a mock function for the type Service -func (_mock *Service) UpdateTags(ctx context.Context, session authn.Session, id string, user users.UserReq) (users.User, error) { - ret := _mock.Called(ctx, session, id, user) - - if len(ret) == 0 { - panic("no return value specified for UpdateTags") - } - - var r0 users.User - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, users.UserReq) (users.User, error)); ok { - return returnFunc(ctx, session, id, user) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, users.UserReq) users.User); ok { - r0 = returnFunc(ctx, session, id, user) - } else { - r0 = ret.Get(0).(users.User) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, users.UserReq) error); ok { - r1 = returnFunc(ctx, session, id, user) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_UpdateTags_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateTags' -type Service_UpdateTags_Call struct { - *mock.Call -} - -// UpdateTags is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - id string -// - user users.UserReq -func (_e *Service_Expecter) UpdateTags(ctx interface{}, session interface{}, id interface{}, user interface{}) *Service_UpdateTags_Call { - return &Service_UpdateTags_Call{Call: _e.mock.On("UpdateTags", ctx, session, id, user)} -} - -func (_c *Service_UpdateTags_Call) Run(run func(ctx context.Context, session authn.Session, id string, user users.UserReq)) *Service_UpdateTags_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 users.UserReq - if args[3] != nil { - arg3 = args[3].(users.UserReq) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_UpdateTags_Call) Return(user1 users.User, err error) *Service_UpdateTags_Call { - _c.Call.Return(user1, err) - return _c -} - -func (_c *Service_UpdateTags_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, id string, user users.UserReq) (users.User, error)) *Service_UpdateTags_Call { - _c.Call.Return(run) - return _c -} - -// UpdateUsername provides a mock function for the type Service -func (_mock *Service) UpdateUsername(ctx context.Context, session authn.Session, id string, username string) (users.User, error) { - ret := _mock.Called(ctx, session, id, username) - - if len(ret) == 0 { - panic("no return value specified for UpdateUsername") - } - - var r0 users.User - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) (users.User, error)); ok { - return returnFunc(ctx, session, id, username) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string, string) users.User); ok { - r0 = returnFunc(ctx, session, id, username) - } else { - r0 = ret.Get(0).(users.User) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string, string) error); ok { - r1 = returnFunc(ctx, session, id, username) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_UpdateUsername_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UpdateUsername' -type Service_UpdateUsername_Call struct { - *mock.Call -} - -// UpdateUsername is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - id string -// - username string -func (_e *Service_Expecter) UpdateUsername(ctx interface{}, session interface{}, id interface{}, username interface{}) *Service_UpdateUsername_Call { - return &Service_UpdateUsername_Call{Call: _e.mock.On("UpdateUsername", ctx, session, id, username)} -} - -func (_c *Service_UpdateUsername_Call) Run(run func(ctx context.Context, session authn.Session, id string, username string)) *Service_UpdateUsername_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - var arg3 string - if args[3] != nil { - arg3 = args[3].(string) - } - run( - arg0, - arg1, - arg2, - arg3, - ) - }) - return _c -} - -func (_c *Service_UpdateUsername_Call) Return(user users.User, err error) *Service_UpdateUsername_Call { - _c.Call.Return(user, err) - return _c -} - -func (_c *Service_UpdateUsername_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, id string, username string) (users.User, error)) *Service_UpdateUsername_Call { - _c.Call.Return(run) - return _c -} - -// VerifyEmail provides a mock function for the type Service -func (_mock *Service) VerifyEmail(ctx context.Context, verificationToken string) (users.User, error) { - ret := _mock.Called(ctx, verificationToken) - - if len(ret) == 0 { - panic("no return value specified for VerifyEmail") - } - - var r0 users.User - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, string) (users.User, error)); ok { - return returnFunc(ctx, verificationToken) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, string) users.User); ok { - r0 = returnFunc(ctx, verificationToken) - } else { - r0 = ret.Get(0).(users.User) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, string) error); ok { - r1 = returnFunc(ctx, verificationToken) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_VerifyEmail_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'VerifyEmail' -type Service_VerifyEmail_Call struct { - *mock.Call -} - -// VerifyEmail is a helper method to define mock.On call -// - ctx context.Context -// - verificationToken string -func (_e *Service_Expecter) VerifyEmail(ctx interface{}, verificationToken interface{}) *Service_VerifyEmail_Call { - return &Service_VerifyEmail_Call{Call: _e.mock.On("VerifyEmail", ctx, verificationToken)} -} - -func (_c *Service_VerifyEmail_Call) Run(run func(ctx context.Context, verificationToken string)) *Service_VerifyEmail_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 string - if args[1] != nil { - arg1 = args[1].(string) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Service_VerifyEmail_Call) Return(user users.User, err error) *Service_VerifyEmail_Call { - _c.Call.Return(user, err) - return _c -} - -func (_c *Service_VerifyEmail_Call) RunAndReturn(run func(ctx context.Context, verificationToken string) (users.User, error)) *Service_VerifyEmail_Call { - _c.Call.Return(run) - return _c -} - -// View provides a mock function for the type Service -func (_mock *Service) View(ctx context.Context, session authn.Session, id string) (users.User, error) { - ret := _mock.Called(ctx, session, id) - - if len(ret) == 0 { - panic("no return value specified for View") - } - - var r0 users.User - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) (users.User, error)); ok { - return returnFunc(ctx, session, id) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session, string) users.User); ok { - r0 = returnFunc(ctx, session, id) - } else { - r0 = ret.Get(0).(users.User) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session, string) error); ok { - r1 = returnFunc(ctx, session, id) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_View_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'View' -type Service_View_Call struct { - *mock.Call -} - -// View is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -// - id string -func (_e *Service_Expecter) View(ctx interface{}, session interface{}, id interface{}) *Service_View_Call { - return &Service_View_Call{Call: _e.mock.On("View", ctx, session, id)} -} - -func (_c *Service_View_Call) Run(run func(ctx context.Context, session authn.Session, id string)) *Service_View_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - var arg2 string - if args[2] != nil { - arg2 = args[2].(string) - } - run( - arg0, - arg1, - arg2, - ) - }) - return _c -} - -func (_c *Service_View_Call) Return(user users.User, err error) *Service_View_Call { - _c.Call.Return(user, err) - return _c -} - -func (_c *Service_View_Call) RunAndReturn(run func(ctx context.Context, session authn.Session, id string) (users.User, error)) *Service_View_Call { - _c.Call.Return(run) - return _c -} - -// ViewProfile provides a mock function for the type Service -func (_mock *Service) ViewProfile(ctx context.Context, session authn.Session) (users.User, error) { - ret := _mock.Called(ctx, session) - - if len(ret) == 0 { - panic("no return value specified for ViewProfile") - } - - var r0 users.User - var r1 error - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session) (users.User, error)); ok { - return returnFunc(ctx, session) - } - if returnFunc, ok := ret.Get(0).(func(context.Context, authn.Session) users.User); ok { - r0 = returnFunc(ctx, session) - } else { - r0 = ret.Get(0).(users.User) - } - if returnFunc, ok := ret.Get(1).(func(context.Context, authn.Session) error); ok { - r1 = returnFunc(ctx, session) - } else { - r1 = ret.Error(1) - } - return r0, r1 -} - -// Service_ViewProfile_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'ViewProfile' -type Service_ViewProfile_Call struct { - *mock.Call -} - -// ViewProfile is a helper method to define mock.On call -// - ctx context.Context -// - session authn.Session -func (_e *Service_Expecter) ViewProfile(ctx interface{}, session interface{}) *Service_ViewProfile_Call { - return &Service_ViewProfile_Call{Call: _e.mock.On("ViewProfile", ctx, session)} -} - -func (_c *Service_ViewProfile_Call) Run(run func(ctx context.Context, session authn.Session)) *Service_ViewProfile_Call { - _c.Call.Run(func(args mock.Arguments) { - var arg0 context.Context - if args[0] != nil { - arg0 = args[0].(context.Context) - } - var arg1 authn.Session - if args[1] != nil { - arg1 = args[1].(authn.Session) - } - run( - arg0, - arg1, - ) - }) - return _c -} - -func (_c *Service_ViewProfile_Call) Return(user users.User, err error) *Service_ViewProfile_Call { - _c.Call.Return(user, err) - return _c -} - -func (_c *Service_ViewProfile_Call) RunAndReturn(run func(ctx context.Context, session authn.Session) (users.User, error)) *Service_ViewProfile_Call { - _c.Call.Return(run) - return _c -} diff --git a/users/postgres/doc.go b/users/postgres/doc.go deleted file mode 100644 index b4f616d7d..000000000 --- a/users/postgres/doc.go +++ /dev/null @@ -1,5 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -// Package postgres contains the database implementation of users repository layer. -package postgres diff --git a/users/postgres/errors.go b/users/postgres/errors.go deleted file mode 100644 index 87ebe9af5..000000000 --- a/users/postgres/errors.go +++ /dev/null @@ -1,26 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package postgres - -import "github.com/absmach/magistrala/pkg/errors" - -var _ errors.Mapper = (*duplicateErrors)(nil) - -type duplicateErrors struct{} - -// GetError maps constraint names to known errors. -func (d duplicateErrors) GetError(constraint string) (error, bool) { - switch constraint { - case "clients_email_key": - return errors.NewRequestError("email id already registered"), true - case "clients_username_key": - return errors.NewRequestError("username not available"), true - default: - return nil, false - } -} - -func NewDuplicateErrors() errors.Mapper { - return duplicateErrors{} -} diff --git a/users/postgres/init.go b/users/postgres/init.go deleted file mode 100644 index ab3d385a0..000000000 --- a/users/postgres/init.go +++ /dev/null @@ -1,178 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package postgres - -import ( - _ "github.com/jackc/pgx/v5/stdlib" // required for SQL access - migrate "github.com/rubenv/sql-migrate" -) - -// Migration of Users service. -func Migration() *migrate.MemoryMigrationSource { - return &migrate.MemoryMigrationSource{ - Migrations: []*migrate.Migration{ - { - Id: "clients_01", - // VARCHAR(36) for column with IDs as UUIDS have a maximum of 36 characters - // STATUS 0 to imply enabled and 1 to imply disabled - // Role 0 to imply user role and 1 to imply admin role - Up: []string{ - `CREATE TABLE IF NOT EXISTS clients ( - id VARCHAR(36) PRIMARY KEY, - name VARCHAR(254) NOT NULL UNIQUE, - domain_id VARCHAR(36), - identity VARCHAR(254) NOT NULL UNIQUE, - secret TEXT NOT NULL, - tags TEXT[], - metadata JSONB, - created_at TIMESTAMP, - updated_at TIMESTAMP, - updated_by VARCHAR(254), - status SMALLINT NOT NULL DEFAULT 0 CHECK (status >= 0), - role SMALLINT DEFAULT 0 CHECK (status >= 0) - )`, - }, - Down: []string{ - `DROP TABLE IF EXISTS clients`, - }, - }, - { - // To support creation of clients from Oauth2 provider - Id: "clients_02", - Up: []string{ - `ALTER TABLE clients ALTER COLUMN secret DROP NOT NULL`, - }, - Down: []string{}, - }, - { - Id: "clients_03", - Up: []string{ - `ALTER TABLE clients - ADD COLUMN username VARCHAR(254) UNIQUE, - ADD COLUMN first_name VARCHAR(254) NOT NULL DEFAULT '', - ADD COLUMN last_name VARCHAR(254) NOT NULL DEFAULT '', - ADD COLUMN profile_picture TEXT`, - `ALTER TABLE clients RENAME COLUMN identity TO email`, - `ALTER TABLE clients DROP COLUMN name`, - }, - Down: []string{ - `ALTER TABLE clients - DROP COLUMN username, - DROP COLUMN first_name, - DROP COLUMN last_name, - DROP COLUMN profile_picture`, - `ALTER TABLE clients RENAME COLUMN email TO identity`, - `ALTER TABLE clients ADD COLUMN name VARCHAR(254) NOT NULL UNIQUE`, - }, - }, - { - Id: "clients_04", - Up: []string{ - `ALTER TABLE IF EXISTS clients RENAME TO users`, - }, - Down: []string{ - `ALTER TABLE IF EXISTS users RENAME TO clients`, - }, - }, - { - Id: "clients_05", - Up: []string{ - `ALTER TABLE users ALTER COLUMN first_name DROP DEFAULT`, - `ALTER TABLE users ALTER COLUMN last_name DROP DEFAULT`, - }, - Down: []string{ - `ALTER TABLE users ALTER COLUMN first_name SET DEFAULT ''`, - `ALTER TABLE users ALTER COLUMN last_name SET DEFAULT ''`, - }, - }, - { - Id: "clients_06", - Up: []string{ - `ALTER TABLE users ALTER COLUMN created_at TYPE TIMESTAMPTZ;`, - `ALTER TABLE users ALTER COLUMN updated_at TYPE TIMESTAMPTZ;`, - }, - Down: []string{ - `ALTER TABLE users ALTER COLUMN created_at TYPE TIMESTAMP;`, - `ALTER TABLE users ALTER COLUMN updated_at TYPE TIMESTAMP;`, - }, - }, - { - Id: "clients_07", - Up: []string{ - `ALTER TABLE users ADD COLUMN verified_at TIMESTAMPTZ DEFAULT NULL;`, - `CREATE TABLE users_verifications ( - user_id VARCHAR(36) NOT NULL, - email VARCHAR(254) NOT NULL, - otp VARCHAR(255), - created_at TIMESTAMPTZ, - expires_at TIMESTAMPTZ, - used_at TIMESTAMPTZ, - FOREIGN KEY (user_id) REFERENCES users (id) ON DELETE CASCADE - ); - CREATE INDEX idx_users_verifications_lookup ON users_verifications (user_id, email, created_at DESC); - `, - }, - Down: []string{ - `ALTER TABLE users DROP COLUMN verified_at;`, - `DROP TABLE users_verifications;`, - }, - }, - { - Id: "clients_08", - Up: []string{ - `ALTER TABLE users RENAME CONSTRAINT clients_identity_key TO clients_email_key;`, - }, - Down: []string{ - `ALTER TABLE users RENAME CONSTRAINT clients_email_key TO clients_identity_key;`, - }, - }, - { - Id: "clients_09", - Up: []string{ - `ALTER TABLE users ADD COLUMN auth_provider VARCHAR(254);`, - }, - Down: []string{ - `ALTER TABLE users DROP COLUMN auth_provider`, - }, - }, - { - Id: "clients_10", - Up: []string{ - `ALTER TABLE users ADD COLUMN private_metadata JSONB;`, - }, - Down: []string{ - `ALTER TABLE users DROP COLUMN private_metadata;`, - }, - }, - { - Id: "clients_11", - Up: []string{ - `UPDATE users - SET metadata = (COALESCE(metadata, '{}'::jsonb) || COALESCE(metadata->'ui', '{}'::jsonb)) - 'ui' - WHERE metadata ? 'ui' AND jsonb_typeof(metadata->'ui') = 'object'`, - `UPDATE users - SET private_metadata = (COALESCE(private_metadata, '{}'::jsonb) || COALESCE(private_metadata->'ui', '{}'::jsonb)) - 'ui' - WHERE private_metadata ? 'ui' AND jsonb_typeof(private_metadata->'ui') = 'object'`, - }, - Down: []string{ - `SELECT 1`, - }, - }, - { - Id: "clients_12", - Up: []string{ - `UPDATE users - SET metadata = (COALESCE(metadata, '{}'::jsonb) || COALESCE(metadata->'admin', '{}'::jsonb)) - 'admin' - WHERE metadata ? 'admin' AND jsonb_typeof(metadata->'admin') = 'object'`, - `UPDATE users - SET private_metadata = (COALESCE(private_metadata, '{}'::jsonb) || COALESCE(private_metadata->'admin', '{}'::jsonb)) - 'admin' - WHERE private_metadata ? 'admin' AND jsonb_typeof(private_metadata->'admin') = 'object'`, - }, - Down: []string{ - `SELECT 1`, - }, - }, - }, - } -} diff --git a/users/postgres/setup_test.go b/users/postgres/setup_test.go deleted file mode 100644 index a8cd27f56..000000000 --- a/users/postgres/setup_test.go +++ /dev/null @@ -1,93 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package postgres_test - -import ( - "database/sql" - "fmt" - "log" - "os" - "testing" - "time" - - pgclient "github.com/absmach/magistrala/pkg/postgres" - upostgres "github.com/absmach/magistrala/users/postgres" - "github.com/jmoiron/sqlx" - "github.com/ory/dockertest/v3" - "github.com/ory/dockertest/v3/docker" - "go.opentelemetry.io/otel" -) - -var ( - db *sqlx.DB - database pgclient.Database - tracer = otel.Tracer("repo_tests") -) - -func TestMain(m *testing.M) { - pool, err := dockertest.NewPool("") - if err != nil { - log.Fatalf("Could not connect to docker: %s", err) - } - - container, err := pool.RunWithOptions(&dockertest.RunOptions{ - Repository: "postgres", - Tag: "16.2-alpine", - Env: []string{ - "POSTGRES_USER=test", - "POSTGRES_PASSWORD=test", - "POSTGRES_DB=test", - "listen_addresses = '*'", - }, - }, func(config *docker.HostConfig) { - config.AutoRemove = true - config.RestartPolicy = docker.RestartPolicy{Name: "no"} - }) - if err != nil { - log.Fatalf("Could not start container: %s", err) - } - - port := container.GetPort("5432/tcp") - - // exponential backoff-retry, because the application in the container might not be ready to accept connections yet - pool.MaxWait = 120 * time.Second - if err := pool.Retry(func() error { - url := fmt.Sprintf("host=localhost port=%s user=test dbname=test password=test sslmode=disable", port) - db, err := sql.Open("pgx", url) - if err != nil { - return err - } - return db.Ping() - }); err != nil { - log.Fatalf("Could not connect to docker: %s", err) - } - - dbConfig := pgclient.Config{ - Host: "localhost", - Port: port, - User: "test", - Pass: "test", - Name: "test", - SSLMode: "disable", - SSLCert: "", - SSLKey: "", - SSLRootCert: "", - } - - if db, err = pgclient.Setup(dbConfig, *upostgres.Migration()); err != nil { - log.Fatalf("Could not setup test DB connection: %s", err) - } - - database = pgclient.NewDatabase(db, dbConfig, tracer) - - code := m.Run() - - // Defers will not be run when using os.Exit - db.Close() - if err := pool.Purge(container); err != nil { - log.Fatalf("Could not purge container: %s", err) - } - - os.Exit(code) -} diff --git a/users/postgres/users.go b/users/postgres/users.go deleted file mode 100644 index 39942e4a9..000000000 --- a/users/postgres/users.go +++ /dev/null @@ -1,771 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package postgres - -import ( - "context" - "database/sql" - "encoding/json" - "fmt" - "strings" - "time" - - api "github.com/absmach/magistrala/api/http" - "github.com/absmach/magistrala/groups" - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - "github.com/absmach/magistrala/pkg/postgres" - "github.com/absmach/magistrala/users" - "github.com/jackc/pgtype" - "github.com/lib/pq" -) - -type userRepo struct { - Repository users.UserRepository - eh errors.Handler -} - -func NewRepository(db postgres.Database) users.Repository { - errHandlerOptions := []errors.HandlerOption{ - postgres.WithDuplicateErrors(NewDuplicateErrors()), - } - return &userRepo{ - Repository: users.UserRepository{DB: db}, - eh: postgres.NewErrorHandler(errHandlerOptions...), - } -} - -func (repo *userRepo) Save(ctx context.Context, c users.User) (users.User, error) { - q := `INSERT INTO users (id, tags, email, secret, metadata, private_metadata, created_at, status, role, first_name, last_name, username, profile_picture, auth_provider) - VALUES (:id, :tags, :email, :secret, :metadata, :private_metadata, :created_at, :status, :role, :first_name, :last_name, :username, :profile_picture, :auth_provider) - RETURNING id, tags, email, metadata, private_metadata, created_at, status, role, first_name, last_name, username, profile_picture, verified_at, auth_provider` - - dbu, err := toDBUser(c) - if err != nil { - return users.User{}, repo.eh.HandleError(repoerr.ErrMarshalBDEntity, err) - } - - row, err := repo.Repository.DB.NamedQueryContext(ctx, q, dbu) - if err != nil { - return users.User{}, repo.eh.HandleError(repoerr.ErrCreateEntity, err) - } - - defer row.Close() - - row.Next() - - dbu = DBUser{} - if err := row.StructScan(&dbu); err != nil { - return users.User{}, repo.eh.HandleError(repoerr.ErrFailedOpDB, err) - } - - user, err := ToUser(dbu) - if err != nil { - return users.User{}, repo.eh.HandleError(repoerr.ErrUnmarshalBDEntity, err) - } - - return user, nil -} - -func (repo *userRepo) CheckSuperAdmin(ctx context.Context, adminID string) error { - q := "SELECT 1 FROM users WHERE id = $1 AND role = $2" - rows, err := repo.Repository.DB.QueryContext(ctx, q, adminID, users.AdminRole) - if err != nil { - return repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - defer rows.Close() - - if rows.Next() { - if err := rows.Err(); err != nil { - return repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - return nil - } - - return repoerr.ErrNotFound -} - -func (repo *userRepo) RetrieveByID(ctx context.Context, id string) (users.User, error) { - q := `SELECT id, tags, email, secret, metadata, private_metadata, created_at, updated_at, updated_by, status, role, first_name, last_name, username, profile_picture, verified_at, auth_provider - FROM users WHERE id = :id` - - dbu := DBUser{ - ID: id, - } - - rows, err := repo.Repository.DB.NamedQueryContext(ctx, q, dbu) - if err != nil { - return users.User{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - defer rows.Close() - - dbu = DBUser{} - if !rows.Next() { - return users.User{}, repoerr.ErrNotFound - } - - if err = rows.StructScan(&dbu); err != nil { - return users.User{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - - user, err := ToUser(dbu) - if err != nil { - return users.User{}, repo.eh.HandleError(repoerr.ErrUnmarshalBDEntity, err) - } - - return user, nil -} - -func (repo *userRepo) RetrieveAll(ctx context.Context, pm users.Page) (users.UsersPage, error) { - query, err := PageQuery(pm) - if err != nil { - return users.UsersPage{}, repo.eh.HandleError(repoerr.ErrParseQueryParams, err) - } - - dbPage, err := ToDBUsersPage(pm) - if err != nil { - return users.UsersPage{}, repo.eh.HandleError(repoerr.ErrMarshalBDEntity, err) - } - - if pm.OnlyTotal { - cq := fmt.Sprintf(`SELECT COUNT(*) FROM users u %s;`, query) - total, err := postgres.Total(ctx, repo.Repository.DB, cq, dbPage) - if err != nil { - return users.UsersPage{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - return users.UsersPage{ - Page: users.Page{Total: total, Offset: pm.Offset, Limit: pm.Limit}, - }, nil - } - - squery := applyOrdering(query, pm) - - q := fmt.Sprintf(`SELECT u.id, u.tags, u.email, u.metadata, u.status, u.role, u.first_name, u.last_name, u.username, - u.created_at, u.updated_at, u.profile_picture, COALESCE(u.updated_by, '') AS updated_by, u.verified_at, - COUNT(*) OVER() AS total_count - FROM users u %s LIMIT :limit OFFSET :offset;`, squery) - - rows, err := repo.Repository.DB.NamedQueryContext(ctx, q, dbPage) - if err != nil { - return users.UsersPage{}, repo.eh.HandleError(repoerr.ErrRetrieveAllUsers, err) - } - defer rows.Close() - - var total uint64 - var items []users.User - for rows.Next() { - dbu := DBUser{} - if err := rows.StructScan(&dbu); err != nil { - return users.UsersPage{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - total = dbu.TotalCount - - c, err := ToUser(dbu) - if err != nil { - return users.UsersPage{}, repo.eh.HandleError(repoerr.ErrUnmarshalBDEntity, err) - } - - items = append(items, c) - } - - if len(items) == 0 { - cq := fmt.Sprintf(`SELECT COUNT(*) FROM users u %s;`, query) - total, err = postgres.Total(ctx, repo.Repository.DB, cq, dbPage) - if err != nil { - return users.UsersPage{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - } - - return users.UsersPage{ - Page: users.Page{Total: total, Offset: pm.Offset, Limit: pm.Limit}, - Users: items, - }, nil -} - -func (repo *userRepo) UpdateUsername(ctx context.Context, user users.User) (users.User, error) { - q := `UPDATE users SET username = :username, updated_at = :updated_at, updated_by = :updated_by - WHERE id = :id AND status = :status - RETURNING id, tags, metadata, private_metadata, status, created_at, updated_at, updated_by, first_name, last_name, username, email, role, verified_at` - - return repo.update(ctx, user, q) -} - -func (repo *userRepo) Update(ctx context.Context, id string, ur users.UserReq) (users.User, error) { - var query []string - var upq string - u := users.User{ID: id} - if ur.FirstName != nil && *ur.FirstName != "" { - query = append(query, "first_name = :first_name") - u.FirstName = *ur.FirstName - } - if ur.LastName != nil && *ur.LastName != "" { - query = append(query, "last_name = :last_name") - u.LastName = *ur.LastName - } - if ur.Metadata != nil { - query = append(query, "metadata = :metadata") - u.Metadata = *ur.Metadata - } - if ur.PrivateMetadata != nil { - query = append(query, "private_metadata = :private_metadata") - u.PrivateMetadata = *ur.PrivateMetadata - } - if ur.Tags != nil { - query = append(query, "tags = :tags") - u.Tags = *ur.Tags - } - if ur.ProfilePicture != nil { - query = append(query, "profile_picture = :profile_picture") - u.ProfilePicture = *ur.ProfilePicture - } - u.UpdatedAt = time.Now().UTC() - if ur.UpdatedAt != nil { - query = append(query, "updated_at = :updated_at") - u.UpdatedAt = *ur.UpdatedAt - } - if ur.UpdatedBy != nil { - query = append(query, "updated_by = :updated_by") - u.UpdatedBy = *ur.UpdatedBy - } - - if len(query) > 0 { - upq = strings.Join(query, ", ") - } - - q := fmt.Sprintf(`UPDATE users SET %s - WHERE id = :id AND status = :status - RETURNING id, tags, metadata, private_metadata, status, created_at, updated_at, updated_by, last_name, first_name, username, profile_picture, email, role, verified_at`, upq) - - u.Status = users.EnabledStatus - return repo.update(ctx, u, q) -} - -func (repo *userRepo) update(ctx context.Context, user users.User, query string) (users.User, error) { - dbu, err := toDBUser(user) - if err != nil { - return users.User{}, repo.eh.HandleError(repoerr.ErrMarshalBDEntity, err) - } - - row, err := repo.Repository.DB.NamedQueryContext(ctx, query, dbu) - if err != nil { - return users.User{}, repo.eh.HandleError(repoerr.ErrUpdateEntity, err) - } - defer row.Close() - - dbu = DBUser{} - if !row.Next() { - return users.User{}, repoerr.ErrNotFound - } - - if err := row.StructScan(&dbu); err != nil { - return users.User{}, repo.eh.HandleError(repoerr.ErrUnmarshalBDEntity, err) - } - - return ToUser(dbu) -} - -func (repo *userRepo) UpdateEmail(ctx context.Context, user users.User) (users.User, error) { - q := `UPDATE users SET email = :email, verified_at = NULL, updated_at = :updated_at, updated_by = :updated_by - WHERE id = :id AND status = :status - RETURNING id, tags, email, metadata, private_metadata, status, created_at, updated_at, updated_by, first_name, last_name, username, role, verified_at` - user.Status = users.EnabledStatus - return repo.update(ctx, user, q) -} - -func (repo *userRepo) UpdateRole(ctx context.Context, user users.User) (users.User, error) { - q := `UPDATE users SET role = :role, updated_at = :updated_at, updated_by = :updated_by - WHERE id = :id AND status = :status - RETURNING id, tags, email, metadata, private_metadata, status, created_at, updated_at, updated_by, first_name, last_name, username, role, verified_at` - user.Status = users.EnabledStatus - return repo.update(ctx, user, q) -} - -func (repo *userRepo) UpdateSecret(ctx context.Context, user users.User) (users.User, error) { - q := `UPDATE users SET secret = :secret, updated_at = :updated_at, updated_by = :updated_by - WHERE id = :id AND status = :status - RETURNING id, tags, email, metadata, private_metadata, status, created_at, updated_at, updated_by, first_name, last_name, username, role, verified_at` - user.Status = users.EnabledStatus - return repo.update(ctx, user, q) -} - -func (repo *userRepo) ChangeStatus(ctx context.Context, user users.User) (users.User, error) { - q := `UPDATE users SET status = :status, updated_at = :updated_at, updated_by = :updated_by - WHERE id = :id - RETURNING id, tags, email, metadata, private_metadata, status, created_at, updated_at, updated_by, first_name, last_name, username, role, verified_at` - - return repo.update(ctx, user, q) -} - -func (repo *userRepo) UpdateVerifiedAt(ctx context.Context, user users.User) (users.User, error) { - q := `UPDATE users SET verified_at = :verified_at - WHERE id = :id and email = :email - RETURNING id, tags, email, metadata, private_metadata, status, created_at, updated_at, updated_by, first_name, last_name, username, role, verified_at` - - return repo.update(ctx, user, q) -} - -func (repo *userRepo) Delete(ctx context.Context, id string) error { - q := "DELETE FROM users AS u WHERE u.id = $1 ;" - - result, err := repo.Repository.DB.ExecContext(ctx, q, id) - if err != nil { - return repo.eh.HandleError(repoerr.ErrRemoveEntity, err) - } - if rows, _ := result.RowsAffected(); rows == 0 { - return repoerr.ErrNotFound - } - - return nil -} - -func (repo *userRepo) SearchUsers(ctx context.Context, pm users.Page) (users.UsersPage, error) { - query, err := PageQuery(pm) - if err != nil { - return users.UsersPage{}, repo.eh.HandleError(repoerr.ErrParseQueryParams, err) - } - - squery := applyOrdering(query, pm) - - q := fmt.Sprintf(`SELECT u.id, u.username, u.metadata, u.first_name, u.last_name, u.created_at, u.updated_at, - COUNT(*) OVER() AS total_count FROM users u %s LIMIT :limit OFFSET :offset;`, squery) - - dbPage, err := ToDBUsersPage(pm) - if err != nil { - return users.UsersPage{}, repo.eh.HandleError(repoerr.ErrMarshalBDEntity, err) - } - - rows, err := repo.Repository.DB.NamedQueryContext(ctx, q, dbPage) - if err != nil { - return users.UsersPage{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - defer rows.Close() - - var total uint64 - var items []users.User - for rows.Next() { - dbu := DBUser{} - if err := rows.StructScan(&dbu); err != nil { - return users.UsersPage{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - total = dbu.TotalCount - - c, err := ToUser(dbu) - if err != nil { - return users.UsersPage{}, err - } - - items = append(items, c) - } - - if len(items) == 0 { - cq := fmt.Sprintf(`SELECT COUNT(*) FROM users u %s;`, query) - total, err = postgres.Total(ctx, repo.Repository.DB, cq, dbPage) - if err != nil { - return users.UsersPage{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - } - - return users.UsersPage{ - Users: items, - Page: users.Page{Total: total, Offset: pm.Offset, Limit: pm.Limit}, - }, nil -} - -func (repo *userRepo) RetrieveAllByIDs(ctx context.Context, pm users.Page) (users.UsersPage, error) { - if (len(pm.IDs) == 0) && (pm.Domain == "") { - return users.UsersPage{ - Page: users.Page{Total: pm.Total, Offset: pm.Offset, Limit: pm.Limit}, - }, nil - } - query, err := PageQuery(pm) - if err != nil { - return users.UsersPage{}, repo.eh.HandleError(repoerr.ErrParseQueryParams, err) - } - squery := applyOrdering(query, pm) - - q := fmt.Sprintf(`SELECT u.id, u.username, u.tags, u.email, u.metadata, u.status, u.role, u.first_name, u.last_name, - u.created_at, u.updated_at, COALESCE(u.updated_by, '') AS updated_by, - COUNT(*) OVER() AS total_count FROM users u %s LIMIT :limit OFFSET :offset;`, squery) - dbPage, err := ToDBUsersPage(pm) - if err != nil { - return users.UsersPage{}, repo.eh.HandleError(repoerr.ErrMarshalBDEntity, err) - } - rows, err := repo.Repository.DB.NamedQueryContext(ctx, q, dbPage) - if err != nil { - return users.UsersPage{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - defer rows.Close() - - var total uint64 - var items []users.User - for rows.Next() { - dbu := DBUser{} - if err := rows.StructScan(&dbu); err != nil { - return users.UsersPage{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - total = dbu.TotalCount - - c, err := ToUser(dbu) - if err != nil { - return users.UsersPage{}, repo.eh.HandleError(repoerr.ErrUnmarshalBDEntity, err) - } - - items = append(items, c) - } - - if len(items) == 0 { - cq := fmt.Sprintf(`SELECT COUNT(*) FROM users u %s;`, query) - total, err = postgres.Total(ctx, repo.Repository.DB, cq, dbPage) - if err != nil { - return users.UsersPage{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - } - - return users.UsersPage{ - Users: items, - Page: users.Page{Total: total, Offset: pm.Offset, Limit: pm.Limit}, - }, nil -} - -func (repo *userRepo) RetrieveByEmail(ctx context.Context, email string) (users.User, error) { - q := `SELECT id, tags, email, secret, metadata, private_metadata, created_at, updated_at, updated_by, status, role, first_name, last_name, username, verified_at, auth_provider - FROM users WHERE email = :email AND status = :status` - - dbu := DBUser{ - Email: email, - Status: users.EnabledStatus, - } - - row, err := repo.Repository.DB.NamedQueryContext(ctx, q, dbu) - if err != nil { - return users.User{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - defer row.Close() - - dbu = DBUser{} - if row.Next() { - if err := row.StructScan(&dbu); err != nil { - return users.User{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - - return ToUser(dbu) - } - - return users.User{}, repoerr.ErrNotFound -} - -func (repo *userRepo) RetrieveByUsername(ctx context.Context, username string) (users.User, error) { - q := `SELECT id, tags, email, secret, metadata, private_metadata, created_at, updated_at, updated_by, status, role, first_name, last_name, username, verified_at, auth_provider - FROM users WHERE username = :username AND status = :status` - - dbu := DBUser{ - Username: sql.NullString{String: username, Valid: username != ""}, - Status: users.EnabledStatus, - } - - row, err := repo.Repository.DB.NamedQueryContext(ctx, q, dbu) - if err != nil { - return users.User{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - defer row.Close() - - dbu = DBUser{} - if row.Next() { - if err := row.StructScan(&dbu); err != nil { - return users.User{}, repo.eh.HandleError(repoerr.ErrViewEntity, err) - } - - return ToUser(dbu) - } - - return users.User{}, repoerr.ErrNotFound -} - -type DBUser struct { - ID string `db:"id"` - Domain string `db:"domain_id"` - Secret string `db:"secret"` - Metadata []byte `db:"metadata,omitempty"` - PrivateMetadata []byte `db:"private_metadata,omitempty"` - Tags pgtype.TextArray `db:"tags,omitempty"` // Tags - CreatedAt time.Time `db:"created_at,omitempty"` - UpdatedAt sql.NullTime `db:"updated_at,omitempty"` - UpdatedBy *string `db:"updated_by,omitempty"` - Groups []groups.Group `db:"groups,omitempty"` - Status users.Status `db:"status,omitempty"` - Role *users.Role `db:"role,omitempty"` - Username sql.NullString `db:"username, omitempty"` - FirstName sql.NullString `db:"first_name, omitempty"` - LastName sql.NullString `db:"last_name, omitempty"` - ProfilePicture sql.NullString `db:"profile_picture, omitempty"` - Email string `db:"email,omitempty"` - VerifiedAt sql.NullTime `db:"verified_at,omitempty"` - AuthProvider sql.NullString `db:"auth_provider,omitempty"` - TotalCount uint64 `db:"total_count"` -} - -func toDBUser(u users.User) (DBUser, error) { - metadata := []byte("{}") - if len(u.Metadata) > 0 { - b, err := json.Marshal(u.Metadata) - if err != nil { - return DBUser{}, errors.Wrap(repoerr.ErrMalformedEntity, err) - } - metadata = b - } - privateMetadata := []byte("{}") - if len(u.PrivateMetadata) > 0 { - b, err := json.Marshal(u.PrivateMetadata) - if err != nil { - return DBUser{}, errors.Wrap(repoerr.ErrMalformedEntity, err) - } - privateMetadata = b - } - var tags pgtype.TextArray - if err := tags.Set(u.Tags); err != nil { - return DBUser{}, err - } - var updatedBy *string - if u.UpdatedBy != "" { - updatedBy = &u.UpdatedBy - } - var updatedAt sql.NullTime - if u.UpdatedAt != (time.Time{}) { - updatedAt = sql.NullTime{Time: u.UpdatedAt, Valid: true} - } - var verifiedAt sql.NullTime - if u.VerifiedAt != (time.Time{}) { - verifiedAt = sql.NullTime{Time: u.VerifiedAt, Valid: true} - } - - var authProvider sql.NullString - if u.AuthProvider != "" { - authProvider = sql.NullString{String: u.AuthProvider, Valid: true} - } - - return DBUser{ - ID: u.ID, - Tags: tags, - Secret: u.Credentials.Secret, - Metadata: metadata, - PrivateMetadata: privateMetadata, - CreatedAt: u.CreatedAt, - UpdatedAt: updatedAt, - UpdatedBy: updatedBy, - Status: u.Status, - Role: &u.Role, - LastName: stringToNullString(u.LastName), - FirstName: stringToNullString(u.FirstName), - Username: stringToNullString(u.Credentials.Username), - ProfilePicture: stringToNullString(u.ProfilePicture), - Email: u.Email, - VerifiedAt: verifiedAt, - AuthProvider: authProvider, - }, nil -} - -func ToUser(dbu DBUser) (users.User, error) { - var metadata, privateMetadata users.Metadata - if dbu.Metadata != nil { - if err := json.Unmarshal([]byte(dbu.Metadata), &metadata); err != nil { - return users.User{}, errors.Wrap(repoerr.ErrMalformedEntity, err) - } - } - if dbu.PrivateMetadata != nil { - if err := json.Unmarshal([]byte(dbu.PrivateMetadata), &privateMetadata); err != nil { - return users.User{}, errors.Wrap(repoerr.ErrMalformedEntity, err) - } - } - var tags []string - for _, e := range dbu.Tags.Elements { - tags = append(tags, e.String) - } - var updatedBy string - if dbu.UpdatedBy != nil { - updatedBy = *dbu.UpdatedBy - } - var updatedAt time.Time - if dbu.UpdatedAt.Valid { - updatedAt = dbu.UpdatedAt.Time.UTC() - } - var verifiedAt time.Time - if dbu.VerifiedAt.Valid { - verifiedAt = dbu.VerifiedAt.Time.UTC() - } - - var authProvider string - if dbu.AuthProvider.Valid { - authProvider = dbu.AuthProvider.String - } - - user := users.User{ - ID: dbu.ID, - FirstName: nullStringString(dbu.FirstName), - LastName: nullStringString(dbu.LastName), - Credentials: users.Credentials{ - Username: nullStringString(dbu.Username), - Secret: dbu.Secret, - }, - Email: dbu.Email, - Metadata: metadata, - PrivateMetadata: privateMetadata, - CreatedAt: dbu.CreatedAt.UTC(), - UpdatedAt: updatedAt, - UpdatedBy: updatedBy, - Status: dbu.Status, - Tags: tags, - ProfilePicture: nullStringString(dbu.ProfilePicture), - VerifiedAt: verifiedAt, - AuthProvider: authProvider, - } - if dbu.Role != nil { - user.Role = *dbu.Role - } - return user, nil -} - -type DBUsersPage struct { - Total uint64 `db:"total"` - Limit uint64 `db:"limit"` - Offset uint64 `db:"offset"` - FirstName string `db:"first_name"` - LastName string `db:"last_name"` - Username string `db:"username"` - Id string `db:"id"` - Email string `db:"email"` - Metadata []byte `db:"metadata"` - Tags pgtype.TextArray `db:"tags"` - GroupID string `db:"group_id"` - Role users.Role `db:"role"` - Status users.Status `db:"status"` - IDs pq.StringArray `db:"ids"` - CreatedFrom time.Time `db:"created_from"` - CreatedTo time.Time `db:"created_to"` -} - -func ToDBUsersPage(pm users.Page) (DBUsersPage, error) { - _, data, err := postgres.CreateMetadataQuery("", pm.Metadata) - if err != nil { - return DBUsersPage{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - - var tags pgtype.TextArray - if err := tags.Set(pm.Tags.Elements); err != nil { - return DBUsersPage{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - - return DBUsersPage{ - FirstName: pm.FirstName, - LastName: pm.LastName, - Username: pm.Username, - Email: pm.Email, - Id: pm.Id, - Metadata: data, - Total: pm.Total, - Offset: pm.Offset, - Limit: pm.Limit, - Status: pm.Status, - Tags: tags, - Role: pm.Role, - IDs: pq.StringArray(pm.IDs), - CreatedFrom: pm.CreatedFrom, - CreatedTo: pm.CreatedTo, - }, nil -} - -func PageQuery(pm users.Page) (string, error) { - var query []string - if pm.FirstName != "" { - query = append(query, "first_name ILIKE '%' || :first_name || '%'") - } - if pm.LastName != "" { - query = append(query, "last_name ILIKE '%' || :last_name || '%'") - } - if pm.Username != "" { - query = append(query, "username ILIKE '%' || :username || '%'") - } - if pm.Email != "" { - query = append(query, "email ILIKE '%' || :email || '%'") - } - if pm.Id != "" { - query = append(query, "id ILIKE '%' || :id || '%'") - } - if len(pm.Tags.Elements) > 0 { - switch pm.Tags.Operator { - case users.AndOp: - query = append(query, "tags @> :tags") - default: // OR - query = append(query, "tags && :tags") - } - } - if pm.Role != users.AllRole { - query = append(query, "u.role = :role") - } - if len(pm.Metadata) > 0 { - query = append(query, "metadata @> :metadata") - } - if len(pm.IDs) != 0 { - query = append(query, "id = ANY(:ids)") - } - if pm.Status != users.AllStatus { - query = append(query, "u.status = :status") - } - if !pm.CreatedFrom.IsZero() { - query = append(query, "created_at >= :created_from") - } - if !pm.CreatedTo.IsZero() { - query = append(query, "created_at <= :created_to") - } - - var emq string - if len(query) > 0 { - emq = fmt.Sprintf("WHERE %s", strings.Join(query, " AND ")) - } - - return emq, nil -} - -func applyOrdering(emq string, pm users.Page) string { - col := "COALESCE(u.updated_at, u.created_at)" - - switch pm.Order { - case "username": - col = "u.username" - case "first_name": - col = "u.first_name" - case "last_name": - col = "u.last_name" - case "email": - col = "u.email" - case "created_at": - col = "u.created_at" - case "updated_at", "": - col = "COALESCE(u.updated_at, u.created_at)" - } - - dir := pm.Dir - if dir != api.AscDir && dir != api.DescDir { - dir = api.DescDir - } - - return fmt.Sprintf("%s ORDER BY %s %s, u.id %s", emq, col, dir, dir) -} - -func stringToNullString(s string) sql.NullString { - if s == "" { - return sql.NullString{} - } - - return sql.NullString{ - String: s, - Valid: true, - } -} - -func nullStringString(ns sql.NullString) string { - if ns.Valid { - return ns.String - } - return "" -} diff --git a/users/postgres/users_test.go b/users/postgres/users_test.go deleted file mode 100644 index ca05d3aa7..000000000 --- a/users/postgres/users_test.go +++ /dev/null @@ -1,2582 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package postgres_test - -import ( - "context" - "fmt" - "strings" - "testing" - "time" - - "github.com/0x6flab/namegenerator" - "github.com/absmach/magistrala/internal/testsutil" - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - "github.com/absmach/magistrala/users" - cpostgres "github.com/absmach/magistrala/users/postgres" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" -) - -const ( - maxNameSize = 254 - defOrder = "created_at" - defDir = "asc" -) - -var ( - invalidName = strings.Repeat("m", maxNameSize+10) - password = "$tr0ngPassw0rd" - namesgen = namegenerator.NewGenerator() - emailSuffix = "@example.com" - validTimestamp = time.Date(2023, 1, 1, 0, 0, 0, 0, time.UTC) - ascDir = "asc" - descDir = "desc" -) - -func TestUsersSave(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM users") - require.Nil(t, err, fmt.Sprintf("clean users unexpected error: %s", err)) - }) - - repo := cpostgres.NewRepository(database) - - uid := testsutil.GenerateUUID(t) - - first_name := namesgen.Generate() - last_name := namesgen.Generate() - username := namesgen.Generate() - - email := first_name + "@example.com" - - externalUser := users.User{ - ID: testsutil.GenerateUUID(t), - FirstName: namesgen.Generate(), - LastName: namesgen.Generate(), - PrivateMetadata: users.Metadata{}, - Metadata: users.Metadata{}, - Credentials: users.Credentials{ - Username: namesgen.Generate(), - }, - Email: namesgen.Generate() + "@example.com", - AuthProvider: "external", - } - cases := []struct { - desc string - user users.User - err error - }{ - { - desc: "add new user successfully", - user: users.User{ - ID: uid, - FirstName: first_name, - LastName: last_name, - Email: email, - Credentials: users.Credentials{ - Username: username, - Secret: password, - }, - PrivateMetadata: users.Metadata{ - "organization": namesgen.Generate(), - }, - Metadata: users.Metadata{ - "address": namesgen.Generate(), - }, - Status: users.EnabledStatus, - }, - err: nil, - }, - { - desc: "add new external user successfully", - user: externalUser, - err: nil, - }, - { - desc: "add user with duplicate user email", - user: users.User{ - ID: testsutil.GenerateUUID(t), - FirstName: first_name, - LastName: last_name, - Email: email, - Credentials: users.Credentials{ - Username: namesgen.Generate(), - Secret: password, - }, - PrivateMetadata: users.Metadata{ - "organization": namesgen.Generate(), - }, - Metadata: users.Metadata{ - "address": namesgen.Generate(), - }, - Status: users.EnabledStatus, - }, - err: errors.ErrEmailAlreadyExists, - }, - { - desc: "add user with duplicate user name", - user: users.User{ - ID: testsutil.GenerateUUID(t), - FirstName: namesgen.Generate(), - LastName: last_name, - Email: namesgen.Generate() + "@example.com", - Credentials: users.Credentials{ - Username: username, - Secret: password, - }, - PrivateMetadata: users.Metadata{ - "organization": namesgen.Generate(), - }, - Metadata: users.Metadata{ - "address": namesgen.Generate(), - }, - Status: users.EnabledStatus, - }, - err: errors.ErrUsernameNotAvailable, - }, - { - desc: "add user with invalid user id", - user: users.User{ - ID: invalidName, - FirstName: namesgen.Generate(), - LastName: namesgen.Generate(), - Email: namesgen.Generate() + "@example.com", - Credentials: users.Credentials{ - Username: username, - Secret: password, - }, - PrivateMetadata: users.Metadata{ - "organization": namesgen.Generate(), - }, - Metadata: users.Metadata{ - "address": namesgen.Generate(), - }, - Status: users.EnabledStatus, - }, - err: repoerr.ErrCreateEntity, - }, - { - desc: "add user with invalid user name", - user: users.User{ - ID: testsutil.GenerateUUID(t), - FirstName: first_name, - LastName: last_name, - Email: namesgen.Generate() + "@example.com", - Credentials: users.Credentials{ - Username: invalidName, - Secret: password, - }, - PrivateMetadata: users.Metadata{ - "organization": namesgen.Generate(), - }, - Metadata: users.Metadata{ - "address": namesgen.Generate(), - }, - Status: users.EnabledStatus, - }, - err: repoerr.ErrCreateEntity, - }, - { - desc: "add user with a missing username", - user: users.User{ - ID: testsutil.GenerateUUID(t), - FirstName: first_name, - LastName: last_name, - Email: namesgen.Generate() + "@example.com", - Credentials: users.Credentials{ - Secret: password, - }, - PrivateMetadata: users.Metadata{ - "organization": namesgen.Generate(), - }, - Metadata: users.Metadata{ - "address": namesgen.Generate(), - }, - }, - err: nil, - }, - { - desc: "add user with a missing user secret", - user: users.User{ - ID: testsutil.GenerateUUID(t), - FirstName: namesgen.Generate(), - LastName: namesgen.Generate(), - Email: namesgen.Generate() + "@example.com", - Credentials: users.Credentials{ - Username: namesgen.Generate(), - }, - PrivateMetadata: users.Metadata{ - "organization": namesgen.Generate(), - }, - Metadata: users.Metadata{ - "address": namesgen.Generate(), - }, - }, - err: nil, - }, - { - desc: "add a user with invalid metadata", - user: users.User{ - ID: testsutil.GenerateUUID(t), - FirstName: namesgen.Generate(), - Email: namesgen.Generate() + "@example.com", - Credentials: users.Credentials{ - Username: username, - Secret: password, - }, - PrivateMetadata: map[string]any{ - "key": make(chan int), - }, - }, - err: errors.ErrMalformedEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - rUser, err := repo.Save(context.Background(), tc.user) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - if err == nil { - rUser.Credentials.Secret = tc.user.Credentials.Secret - assert.Equal(t, tc.user, rUser, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.user, rUser)) - } - }) - } -} - -func TestIsPlatformAdmin(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM users") - require.Nil(t, err, fmt.Sprintf("clean users unexpected error: %s", err)) - }) - - repo := cpostgres.NewRepository(database) - - first_name := namesgen.Generate() - last_name := namesgen.Generate() - username := namesgen.Generate() - email := first_name + "@example.com" - - cases := []struct { - desc string - user users.User - err error - }{ - { - desc: "authorize check for super user", - user: users.User{ - ID: testsutil.GenerateUUID(t), - FirstName: first_name, - LastName: last_name, - Email: email, - Credentials: users.Credentials{ - Username: username, - Secret: password, - }, - PrivateMetadata: users.Metadata{}, - Metadata: users.Metadata{}, - Status: users.EnabledStatus, - Role: users.AdminRole, - }, - err: nil, - }, - { - desc: "unauthorize user", - user: users.User{ - ID: testsutil.GenerateUUID(t), - FirstName: first_name, - LastName: last_name, - Email: namesgen.Generate() + "@example.com", - Credentials: users.Credentials{ - Username: namesgen.Generate(), - Secret: password, - }, - PrivateMetadata: users.Metadata{}, - Metadata: users.Metadata{}, - Status: users.EnabledStatus, - Role: users.UserRole, - }, - err: repoerr.ErrNotFound, - }, - } - - for _, tc := range cases { - _, err := repo.Save(context.Background(), tc.user) - require.Nil(t, err, fmt.Sprintf("%s: save user unexpected error: %s", tc.desc, err)) - err = repo.CheckSuperAdmin(context.Background(), tc.user.ID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.err, err)) - } -} - -func TestRetrieveByID(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM users") - require.Nil(t, err, fmt.Sprintf("clean users unexpected error: %s", err)) - }) - - repo := cpostgres.NewRepository(database) - - user := users.User{ - ID: testsutil.GenerateUUID(t), - FirstName: namesgen.Generate(), - LastName: namesgen.Generate(), - Email: namesgen.Generate() + "@example.com", - Credentials: users.Credentials{ - Username: namesgen.Generate(), - Secret: password, - }, - PrivateMetadata: users.Metadata{ - "organization": namesgen.Generate(), - }, - Metadata: users.Metadata{ - "address": namesgen.Generate(), - }, - Status: users.EnabledStatus, - } - - _, err := repo.Save(context.Background(), user) - require.Nil(t, err, fmt.Sprintf("failed to save users %s", user.ID)) - - externalUser := users.User{ - ID: testsutil.GenerateUUID(t), - FirstName: namesgen.Generate(), - LastName: namesgen.Generate(), - PrivateMetadata: users.Metadata{}, - Metadata: users.Metadata{}, - Credentials: users.Credentials{ - Username: namesgen.Generate(), - }, - Email: namesgen.Generate() + "@example.com", - AuthProvider: "external", - } - - _, err = repo.Save(context.Background(), externalUser) - require.Nil(t, err, fmt.Sprintf("failed to save users %s", user.ID)) - - cases := []struct { - desc string - userID string - user users.User - err error - }{ - { - desc: "retrieve existing user", - userID: user.ID, - user: user, - err: nil, - }, - - { - desc: "retrieve existing oauth user", - userID: externalUser.ID, - user: externalUser, - err: nil, - }, - { - desc: "retrieve non-existing user", - userID: invalidName, - err: repoerr.ErrNotFound, - }, - { - desc: "retrieve with empty user id", - userID: "", - err: repoerr.ErrNotFound, - }, - } - - for _, tc := range cases { - rUser, err := repo.RetrieveByID(context.Background(), tc.userID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.err, err)) - if err == nil { - assert.Equal(t, tc.user, rUser, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.user, rUser)) - } - } -} - -func TestRetrieveAll(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM users") - require.Nil(t, err, fmt.Sprintf("clean users unexpected error: %s", err)) - }) - - repo := cpostgres.NewRepository(database) - - num := 200 - var items, enabledUsers []users.User - baseTime := time.Now().UTC().Truncate(time.Millisecond) - for i := 0; i < num; i++ { - user := users.User{ - ID: testsutil.GenerateUUID(t), - FirstName: namesgen.Generate(), - LastName: namesgen.Generate(), - Email: namesgen.Generate() + "@example.com", - Credentials: users.Credentials{ - Username: namesgen.Generate(), - Secret: "", - }, - Metadata: users.Metadata{}, - Status: users.EnabledStatus, - Tags: []string{"tag1"}, - CreatedAt: baseTime.Add(time.Duration(i) * time.Millisecond), - UpdatedAt: baseTime.Add(time.Duration(i) * time.Millisecond), - } - if i%50 == 0 { - user.Metadata = map[string]any{ - "key": "value", - } - user.Role = users.AdminRole - user.Status = users.DisabledStatus - } - if i%99 == 0 { - user.Tags = []string{"tag1", "tag2"} - } - _, err := repo.Save(context.Background(), user) - require.Nil(t, err, fmt.Sprintf("failed to save user %s", user.ID)) - items = append(items, user) - if user.Status == users.EnabledStatus { - enabledUsers = append(enabledUsers, user) - } - } - - reversedUsers := []users.User{} - for i := len(items) - 1; i >= 0; i-- { - reversedUsers = append(reversedUsers, items[i]) - } - - cases := []struct { - desc string - pageMeta users.Page - page users.UsersPage - err error - }{ - { - desc: "retrieve first page of users", - pageMeta: users.Page{ - Offset: 0, - Limit: 1, - Role: users.AllRole, - Status: users.AllStatus, - Order: "created_at", - Dir: ascDir, - }, - page: users.UsersPage{ - Page: users.Page{ - Total: 200, - Offset: 0, - Limit: 1, - }, - Users: items[0:1], - }, - err: nil, - }, - { - desc: "retrieve second page of users", - pageMeta: users.Page{ - Offset: 50, - Limit: 200, - Role: users.AllRole, - Status: users.AllStatus, - Order: "created_at", - Dir: ascDir, - }, - page: users.UsersPage{ - Page: users.Page{ - Total: 200, - Offset: 50, - Limit: 200, - }, - Users: items[50:200], - }, - err: nil, - }, - { - desc: "retrieve users with limit", - pageMeta: users.Page{ - Offset: 0, - Limit: 50, - Role: users.AllRole, - Status: users.AllStatus, - Order: "created_at", - Dir: ascDir, - }, - page: users.UsersPage{ - Page: users.Page{ - Total: uint64(num), - Offset: 0, - Limit: 50, - }, - Users: items[:50], - }, - }, - { - desc: "retrieve with offset out of range", - pageMeta: users.Page{ - Offset: 1000, - Limit: 200, - Role: users.AllRole, - Status: users.AllStatus, - }, - page: users.UsersPage{ - Page: users.Page{ - Total: 200, - Offset: 1000, - Limit: 200, - }, - Users: []users.User{}, - }, - err: nil, - }, - { - desc: "retrieve with limit out of range", - pageMeta: users.Page{ - Offset: 0, - Limit: 1000, - Role: users.AllRole, - Status: users.AllStatus, - Order: "created_at", - Dir: ascDir, - }, - page: users.UsersPage{ - Page: users.Page{ - Total: 200, - Offset: 0, - Limit: 1000, - }, - Users: items, - }, - err: nil, - }, - { - desc: "retrieve with empty page", - pageMeta: users.Page{}, - page: users.UsersPage{ - Page: users.Page{ - Total: 196, // number of enabled users - Offset: 0, - Limit: 0, - }, - Users: []users.User{}, - }, - err: nil, - }, - { - desc: "retrieve with user id", - pageMeta: users.Page{ - IDs: []string{items[0].ID}, - Offset: 0, - Limit: 3, - Role: users.AllRole, - Status: users.AllStatus, - Order: "created_at", - Dir: ascDir, - }, - page: users.UsersPage{ - Page: users.Page{ - Total: 1, - Offset: 0, - Limit: 3, - }, - Users: []users.User{items[0]}, - }, - err: nil, - }, - { - desc: "retrieve with invalid user id", - pageMeta: users.Page{ - IDs: []string{invalidName}, - Offset: 0, - Limit: 3, - Role: users.AllRole, - Status: users.AllStatus, - }, - page: users.UsersPage{ - Page: users.Page{ - Total: 0, - Offset: 0, - Limit: 3, - }, - Users: []users.User{}, - }, - err: nil, - }, - { - desc: "retrieve with first name", - pageMeta: users.Page{ - FirstName: items[0].FirstName, - Offset: 0, - Limit: 3, - Role: users.AllRole, - Status: users.AllStatus, - Order: "created_at", - Dir: ascDir, - }, - page: users.UsersPage{ - Page: users.Page{ - Total: 1, - Offset: 0, - Limit: 3, - }, - Users: []users.User{items[0]}, - }, - err: nil, - }, - { - desc: "retrieve with username", - pageMeta: users.Page{ - Username: items[0].Credentials.Username, - Offset: 0, - Limit: 3, - Role: users.AllRole, - Status: users.AllStatus, - Order: "created_at", - Dir: ascDir, - }, - page: users.UsersPage{ - Page: users.Page{ - Total: 1, - Offset: 0, - Limit: 3, - }, - Users: []users.User{items[0]}, - }, - err: nil, - }, - { - desc: "retrieve with enabled status", - pageMeta: users.Page{ - Status: users.EnabledStatus, - Offset: 0, - Limit: 200, - Role: users.AllRole, - Order: "created_at", - Dir: ascDir, - }, - page: users.UsersPage{ - Page: users.Page{ - Total: 196, - Offset: 0, - Limit: 200, - }, - Users: enabledUsers, - }, - err: nil, - }, - { - desc: "retrieve with disabled status", - pageMeta: users.Page{ - Status: users.DisabledStatus, - Offset: 0, - Limit: 200, - Role: users.AllRole, - Order: "created_at", - Dir: ascDir, - }, - page: users.UsersPage{ - Page: users.Page{ - Total: 4, - Offset: 0, - Limit: 200, - }, - Users: []users.User{items[0], items[50], items[100], items[150]}, - }, - }, - { - desc: "retrieve with all status", - pageMeta: users.Page{ - Status: users.AllStatus, - Offset: 0, - Limit: 200, - Role: users.AllRole, - Order: "created_at", - Dir: ascDir, - }, - page: users.UsersPage{ - Page: users.Page{ - Total: 200, - Offset: 0, - Limit: 200, - }, - Users: items, - }, - }, - { - desc: "retrieve by tags with OR operator", - pageMeta: users.Page{ - Tags: users.TagsQuery{Operator: users.OrOp, Elements: []string{"tag1"}}, - Offset: 0, - Limit: 200, - Role: users.AllRole, - Status: users.AllStatus, - }, - page: users.UsersPage{ - Page: users.Page{ - Total: 200, - Offset: 0, - Limit: 200, - }, - Users: items, - }, - err: nil, - }, - { - desc: "retrieve by tags with OR operator no match", - pageMeta: users.Page{ - Tags: users.TagsQuery{Operator: users.OrOp, Elements: []string{"non-existing-tag"}}, - Offset: 0, - Limit: 200, - Role: users.AllRole, - Status: users.AllStatus, - }, - page: users.UsersPage{ - Page: users.Page{ - Total: 0, - Offset: 0, - Limit: 200, - }, - Users: []users.User{}, - }, - err: nil, - }, - { - desc: "retrieve by tags with AND operator", - pageMeta: users.Page{ - Tags: users.TagsQuery{Operator: users.AndOp, Elements: []string{"tag1", "tag2"}}, - Offset: 0, - Limit: 200, - Role: users.AllRole, - Status: users.AllStatus, - }, - page: users.UsersPage{ - Page: users.Page{ - Total: 3, - Offset: 0, - Limit: 200, - }, - Users: []users.User{items[0], items[99], items[198]}, - }, - err: nil, - }, - { - desc: "retrieve by tags with AND operator no match", - pageMeta: users.Page{ - Tags: users.TagsQuery{Operator: users.AndOp, Elements: []string{"tag1", "non-existing-tag"}}, - Offset: 0, - Limit: 200, - Role: users.AllRole, - Status: users.AllStatus, - }, - page: users.UsersPage{ - Page: users.Page{ - Total: 0, - Offset: 0, - Limit: 200, - }, - Users: []users.User{}, - }, - err: nil, - }, - { - desc: "retrieve with invalid first name", - pageMeta: users.Page{ - FirstName: invalidName, - Offset: 0, - Limit: 3, - Role: users.AllRole, - Status: users.AllStatus, - }, - page: users.UsersPage{ - Page: users.Page{ - Total: 0, - Offset: 0, - Limit: 3, - }, - Users: []users.User{}, - }, - }, - { - desc: "retrieve with metadata", - pageMeta: users.Page{ - Metadata: map[string]any{ - "key": "value", - }, - Offset: 0, - Limit: 200, - Role: users.AllRole, - Status: users.AllStatus, - Order: "created_at", - Dir: ascDir, - }, - page: users.UsersPage{ - Page: users.Page{ - Total: 4, - Offset: 0, - Limit: 200, - }, - Users: []users.User{items[0], items[50], items[100], items[150]}, - }, - err: nil, - }, - { - desc: "retrieve with invalid metadata", - pageMeta: users.Page{ - Metadata: map[string]any{ - "key": "value1", - }, - Offset: 0, - Limit: 200, - Role: users.AllRole, - Status: users.AllStatus, - }, - page: users.UsersPage{ - Page: users.Page{ - Total: 0, - Offset: 0, - Limit: 200, - }, - Users: []users.User{}, - }, - err: nil, - }, - { - desc: "retrieve with role", - pageMeta: users.Page{ - Role: users.AdminRole, - Offset: 0, - Limit: 200, - Status: users.AllStatus, - Order: "created_at", - Dir: ascDir, - }, - page: users.UsersPage{ - Page: users.Page{ - Total: 4, - Offset: 0, - Limit: 200, - }, - Users: []users.User{items[0], items[50], items[100], items[150]}, - }, - err: nil, - }, - { - desc: "retrieve with invalid role", - pageMeta: users.Page{ - Role: users.AdminRole + 2, - Offset: 0, - Limit: 200, - Status: users.AllStatus, - }, - page: users.UsersPage{ - Page: users.Page{ - Total: 0, - Offset: 0, - Limit: 200, - }, - Users: []users.User{}, - }, - err: nil, - }, - { - desc: "retrieve users with order by first_name ascending", - pageMeta: users.Page{ - Offset: 0, - Limit: 10, - Role: users.AllRole, - Status: users.AllStatus, - Order: "first_name", - Dir: ascDir, - }, - page: users.UsersPage{ - Page: users.Page{ - Total: uint64(num), - Offset: 0, - Limit: 10, - }, - }, - err: nil, - }, - { - desc: "retrieve users with order by first_name descending", - pageMeta: users.Page{ - Offset: 0, - Limit: 10, - Role: users.AllRole, - Status: users.AllStatus, - Order: "first_name", - Dir: descDir, - }, - page: users.UsersPage{ - Page: users.Page{ - Total: uint64(num), - Offset: 0, - Limit: 10, - }, - }, - err: nil, - }, - { - desc: "retrieve users with order by username ascending", - pageMeta: users.Page{ - Offset: 0, - Limit: 10, - Role: users.AllRole, - Status: users.AllStatus, - Order: "username", - Dir: ascDir, - }, - page: users.UsersPage{ - Page: users.Page{ - Total: uint64(num), - Offset: 0, - Limit: 10, - }, - }, - err: nil, - }, - { - desc: "retrieve users with order by username descending", - pageMeta: users.Page{ - Offset: 0, - Limit: 10, - Role: users.AllRole, - Status: users.AllStatus, - Order: "username", - Dir: descDir, - }, - page: users.UsersPage{ - Page: users.Page{ - Total: uint64(num), - Offset: 0, - Limit: 10, - }, - }, - err: nil, - }, - { - desc: "retrieve users with order by created_at ascending", - pageMeta: users.Page{ - Offset: 0, - Limit: 10, - Role: users.AllRole, - Status: users.AllStatus, - Order: "created_at", - Dir: ascDir, - }, - page: users.UsersPage{ - Page: users.Page{ - Total: uint64(num), - Offset: 0, - Limit: 10, - }, - Users: items[:10], - }, - err: nil, - }, - { - desc: "retrieve users with order by created_at descending", - pageMeta: users.Page{ - Offset: 0, - Limit: 10, - Role: users.AllRole, - Status: users.AllStatus, - Order: "created_at", - Dir: descDir, - }, - page: users.UsersPage{ - Page: users.Page{ - Total: uint64(num), - Offset: 0, - Limit: 10, - }, - Users: reversedUsers[:10], - }, - err: nil, - }, - { - desc: "retrieve users with order by updated_at ascending", - pageMeta: users.Page{ - Offset: 0, - Limit: 10, - Role: users.AllRole, - Status: users.AllStatus, - Order: "updated_at", - Dir: ascDir, - }, - page: users.UsersPage{ - Page: users.Page{ - Total: uint64(num), - Offset: 0, - Limit: 10, - }, - Users: items[:10], - }, - err: nil, - }, - { - desc: "retrieve users with order by updated_at descending", - pageMeta: users.Page{ - Offset: 0, - Limit: 10, - Role: users.AllRole, - Status: users.AllStatus, - Order: "updated_at", - Dir: descDir, - }, - page: users.UsersPage{ - Page: users.Page{ - Total: uint64(num), - Offset: 0, - Limit: 10, - }, - Users: reversedUsers[:10], - }, - err: nil, - }, - { - desc: "retrieve users created from specific time", - pageMeta: users.Page{ - CreatedFrom: baseTime.Add(50 * time.Millisecond), - Offset: 0, - Limit: 200, - Role: users.AllRole, - Status: users.AllStatus, - Order: "created_at", - Dir: ascDir, - }, - page: users.UsersPage{ - Page: users.Page{ - Total: 150, - Offset: 0, - Limit: 200, - }, - Users: items[50:200], - }, - err: nil, - }, - { - desc: "retrieve users created to specific time", - pageMeta: users.Page{ - CreatedTo: baseTime.Add(49 * time.Millisecond), - Offset: 0, - Limit: 200, - Role: users.AllRole, - Status: users.AllStatus, - Order: "created_at", - Dir: ascDir, - }, - page: users.UsersPage{ - Page: users.Page{ - Total: 50, - Offset: 0, - Limit: 200, - }, - Users: items[0:50], - }, - err: nil, - }, - { - desc: "retrieve users created within time range", - pageMeta: users.Page{ - CreatedFrom: baseTime.Add(50 * time.Millisecond), - CreatedTo: baseTime.Add(99 * time.Millisecond), - Offset: 0, - Limit: 200, - Role: users.AllRole, - Status: users.AllStatus, - Order: "created_at", - Dir: ascDir, - }, - page: users.UsersPage{ - Page: users.Page{ - Total: 50, - Offset: 0, - Limit: 200, - }, - Users: items[50:100], - }, - err: nil, - }, - { - desc: "retrieve users with time range outside of all records", - pageMeta: users.Page{ - CreatedFrom: baseTime.Add(300 * time.Millisecond), - CreatedTo: baseTime.Add(400 * time.Millisecond), - Offset: 0, - Limit: 200, - Role: users.AllRole, - Status: users.AllStatus, - }, - page: users.UsersPage{ - Page: users.Page{ - Total: 0, - Offset: 0, - Limit: 200, - }, - Users: []users.User{}, - }, - err: nil, - }, - } - - for _, tc := range cases { - page, err := repo.RetrieveAll(context.Background(), tc.pageMeta) - assert.Equal(t, tc.page.Total, page.Total, fmt.Sprintf("%s: expected %d got %d\n", tc.desc, tc.page.Total, page.Total)) - assert.Equal(t, tc.page.Offset, page.Offset, fmt.Sprintf("%s: expected %d got %d\n", tc.desc, tc.page.Offset, page.Offset)) - assert.Equal(t, tc.page.Limit, page.Limit, fmt.Sprintf("%s: expected %d got %d\n", tc.desc, tc.page.Limit, page.Limit)) - assert.Equal(t, tc.page.Page, page.Page, fmt.Sprintf("%s: expected %v, got %v", tc.desc, tc.page, page)) - if len(tc.page.Users) > 0 { - assert.ElementsMatch(t, stripUserDetails(tc.page.Users), stripUserDetails(page.Users), fmt.Sprintf("%s: expected %v, got %v", tc.desc, tc.page.Users, page.Users)) - } - verifyUsersOrdering(t, page.Users, tc.pageMeta.Order, tc.pageMeta.Dir) - assert.Equal(t, tc.err, err, fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - } -} - -func TestSearch(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM users") - require.Nil(t, err, fmt.Sprintf("clean users unexpected error: %s", err)) - }) - repo := cpostgres.NewRepository(database) - - nUsers := uint64(200) - expectedUsers := []users.User{} - baseTime := time.Now().UTC().Truncate(time.Millisecond) - for i := 0; i < int(nUsers); i++ { - user := generateUserWithTime(t, users.EnabledStatus, repo, baseTime.Add(time.Duration(i)*time.Millisecond)) - - expectedUsers = append(expectedUsers, users.User{ - ID: user.ID, - FirstName: user.FirstName, - LastName: user.LastName, - Credentials: users.Credentials{ - Username: user.Credentials.Username, - }, - Metadata: user.Metadata, - CreatedAt: user.CreatedAt, - }) - } - - page, err := repo.RetrieveAll(context.Background(), users.Page{Offset: 0, Limit: nUsers}) - require.Nil(t, err, fmt.Sprintf("retrieve all users unexpected error: %s", err)) - assert.Equal(t, nUsers, page.Total) - - cases := []struct { - desc string - page users.Page - response users.UsersPage - err error - }{ - { - desc: "with empty page", - page: users.Page{}, - response: users.UsersPage{ - Users: []users.User(nil), - Page: users.Page{ - Total: nUsers, - Offset: 0, - Limit: 0, - }, - }, - err: nil, - }, - { - desc: "with offset only", - page: users.Page{ - Offset: 50, - }, - response: users.UsersPage{ - Users: []users.User(nil), - Page: users.Page{ - Total: nUsers, - Offset: 50, - Limit: 0, - }, - }, - err: nil, - }, - { - desc: "with limit only", - page: users.Page{ - Limit: 10, - Order: "name", - Dir: ascDir, - }, - response: users.UsersPage{ - Users: expectedUsers[0:10], - Page: users.Page{ - Total: nUsers, - Offset: 0, - Limit: 10, - }, - }, - err: nil, - }, - { - desc: "retrieve all users", - page: users.Page{ - Offset: 0, - Limit: 10, - Order: defOrder, - Dir: defDir, - }, - response: users.UsersPage{ - Page: users.Page{ - Total: nUsers, - Offset: 0, - Limit: 10, - }, - Users: expectedUsers[:10], - }, - }, - { - desc: "with offset and limit", - page: users.Page{ - Offset: 10, - Limit: 10, - Order: "name", - Dir: ascDir, - }, - response: users.UsersPage{ - Users: expectedUsers[10:20], - Page: users.Page{ - Total: nUsers, - Offset: 10, - Limit: 10, - }, - }, - err: nil, - }, - { - desc: "with offset out of range and limit", - page: users.Page{ - Offset: 1000, - Limit: 50, - }, - response: users.UsersPage{ - Page: users.Page{ - Total: nUsers, - Offset: 1000, - Limit: 50, - }, - Users: []users.User(nil), - }, - }, - { - desc: "with offset and limit out of range", - page: users.Page{ - Offset: 190, - Limit: 50, - Order: "name", - Dir: ascDir, - }, - response: users.UsersPage{ - Page: users.Page{ - Total: nUsers, - Offset: 190, - Limit: 50, - }, - Users: expectedUsers[190:200], - }, - }, - { - desc: "with shorter name", - page: users.Page{ - FirstName: expectedUsers[0].FirstName[:4], - Offset: 0, - Limit: 10, - Order: "first_name", - Dir: ascDir, - }, - response: users.UsersPage{ - Users: findUsers(expectedUsers, expectedUsers[0].FirstName[:4], 0, 10), - Page: users.Page{ - Total: nUsers, - Offset: 0, - Limit: 10, - }, - }, - err: nil, - }, - { - desc: "with longer name", - page: users.Page{ - FirstName: expectedUsers[0].FirstName, - Offset: 0, - Limit: 10, - }, - response: users.UsersPage{ - Users: []users.User{expectedUsers[0]}, - Page: users.Page{ - Total: 1, - Offset: 0, - Limit: 10, - }, - }, - err: nil, - }, - { - desc: "with name SQL injected", - page: users.Page{ - FirstName: fmt.Sprintf("%s' OR '1'='1", expectedUsers[0].FirstName[:1]), - Offset: 0, - Limit: 10, - }, - response: users.UsersPage{ - Users: []users.User(nil), - Page: users.Page{ - Total: 0, - Offset: 0, - Limit: 10, - }, - }, - err: nil, - }, - { - desc: "with shorter email", - page: users.Page{ - Email: expectedUsers[0].FirstName[:4], - Offset: 0, - Limit: 10, - Order: "first_name", - Dir: ascDir, - }, - response: users.UsersPage{ - Users: findUsers(expectedUsers, expectedUsers[0].FirstName[:4], 0, 10), - Page: users.Page{ - Total: nUsers, - Offset: 0, - Limit: 10, - }, - }, - err: nil, - }, - { - desc: "with Identity SQL injected", - page: users.Page{ - Email: fmt.Sprintf("%s' OR '1'='1", expectedUsers[0].FirstName[:1]), - Offset: 0, - Limit: 10, - }, - response: users.UsersPage{ - Users: []users.User(nil), - Page: users.Page{ - Total: 0, - Offset: 0, - Limit: 10, - }, - }, - err: nil, - }, - { - desc: "with unknown name", - page: users.Page{ - FirstName: namesgen.Generate(), - Offset: 0, - Limit: 10, - }, - response: users.UsersPage{ - Users: []users.User(nil), - Page: users.Page{ - Total: 0, - Offset: 0, - Limit: 10, - }, - }, - err: nil, - }, - { - desc: "with unknown email", - page: users.Page{ - Email: namesgen.Generate(), - Offset: 0, - Limit: 10, - }, - response: users.UsersPage{ - Users: []users.User(nil), - Page: users.Page{ - Total: 0, - Offset: 0, - Limit: 10, - }, - }, - err: nil, - }, - { - desc: "with name in asc order", - page: users.Page{ - Order: "first_name", - Dir: ascDir, - FirstName: expectedUsers[0].FirstName[:1], - Offset: 0, - Limit: 10, - }, - response: users.UsersPage{}, - err: nil, - }, - { - desc: "with name in desc order", - page: users.Page{ - Order: "first_name", - Dir: descDir, - FirstName: expectedUsers[0].FirstName[:1], - Offset: 0, - Limit: 10, - }, - response: users.UsersPage{}, - err: nil, - }, - { - desc: "with last name in asc order", - page: users.Page{ - LastName: expectedUsers[0].LastName[:1], - Order: "last_name", - Dir: ascDir, - }, - response: users.UsersPage{ - Users: []users.User{expectedUsers[0]}, - Page: users.Page{ - Total: 1, - Offset: 0, - Limit: 1, - }, - }, - err: nil, - }, - { - desc: "with username in asc order", - page: users.Page{ - Username: expectedUsers[0].Credentials.Username[:1], - Order: "username", - Dir: ascDir, - }, - response: users.UsersPage{ - Users: []users.User{expectedUsers[0]}, - Page: users.Page{ - Total: 1, - Offset: 0, - Limit: 1, - }, - }, - err: nil, - }, - } - for _, c := range cases { - t.Run(c.desc, func(t *testing.T) { - switch response, err := repo.SearchUsers(context.Background(), c.page); { - case err == nil: - if c.page.Order != "" && c.page.Dir != "" { - c.response = response - } - assert.Nil(t, err) - assert.Equal(t, c.response.Total, response.Total) - assert.Equal(t, c.response.Limit, response.Limit) - assert.Equal(t, c.response.Offset, response.Offset) - assert.ElementsMatch(t, response.Users, c.response.Users, fmt.Sprintf("expected %v got %v\n", c.response.Users, response.Users)) - default: - assert.True(t, errors.Contains(err, c.err), fmt.Sprintf("expected %s to contain %s\n", err, c.err)) - } - }) - } -} - -func TestUpdateRole(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM users") - require.Nil(t, err, fmt.Sprintf("clean users unexpected error: %s", err)) - }) - repo := cpostgres.NewRepository(database) - user1 := generateUser(t, users.EnabledStatus, repo) - user2 := generateUser(t, users.DisabledStatus, repo) - adminRole := users.AdminRole - userRole := users.UserRole - - cases := []struct { - desc string - update string - userID string - userReq users.User - err error - }{ - { - desc: "update role of user to admin", - userReq: users.User{ - ID: user1.ID, - Role: adminRole, - }, - err: nil, - }, - { - desc: "update role of admin to user", - userReq: users.User{ - ID: user1.ID, - Role: userRole, - }, - err: nil, - }, - { - desc: "update role for disabled user", - userReq: users.User{ - ID: user2.ID, - Role: adminRole, - }, - err: repoerr.ErrNotFound, - }, - { - desc: "update role for invalid user", - userReq: users.User{ - ID: testsutil.GenerateUUID(t), - Role: adminRole, - }, - err: repoerr.ErrNotFound, - }, - } - - for _, c := range cases { - t.Run(c.desc, func(t *testing.T) { - updatedAt := time.Now().UTC().Truncate(time.Microsecond) - updatedBy := testsutil.GenerateUUID(t) - c.userReq.UpdatedAt = updatedAt - c.userReq.UpdatedBy = updatedBy - expected, err := repo.UpdateRole(context.Background(), c.userReq) - assert.True(t, errors.Contains(err, c.err), fmt.Sprintf("expected %s to contain %s\n", err, c.err)) - if err == nil { - assert.Equal(t, c.userReq.Role, expected.Role) - assert.Equal(t, c.userReq.UpdatedAt, expected.UpdatedAt) - assert.Equal(t, c.userReq.UpdatedBy, expected.UpdatedBy) - } - }) - } -} - -func TestUpdateEmail(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM users") - require.Nil(t, err, fmt.Sprintf("clean users unexpected error: %s", err)) - }) - repo := cpostgres.NewRepository(database) - - user1 := generateUser(t, users.EnabledStatus, repo) - user2 := generateUser(t, users.DisabledStatus, repo) - user3 := generateUser(t, users.EnabledStatus, repo) - - updatedEmail := namesgen.Generate() + emailSuffix - emptyName := "" - - cases := []struct { - desc string - update string - userReq users.User - err error - }{ - { - desc: "update email for enabled user", - userReq: users.User{ - ID: user1.ID, - Email: updatedEmail, - }, - - err: nil, - }, - { - desc: "update empty email for enabled user", - userReq: users.User{ - ID: user3.ID, - Email: emptyName, - }, - err: nil, - }, - { - desc: "update email for disabled user", - userReq: users.User{ - ID: user2.ID, - Email: updatedEmail, - }, - err: repoerr.ErrNotFound, - }, - { - desc: "update email for invalid user", - userReq: users.User{ - ID: testsutil.GenerateUUID(t), - Email: updatedEmail, - }, - err: repoerr.ErrNotFound, - }, - } - - for _, c := range cases { - t.Run(c.desc, func(t *testing.T) { - updatedAt := time.Now().UTC().Truncate(time.Microsecond) - updatedBy := testsutil.GenerateUUID(t) - c.userReq.UpdatedAt = updatedAt - c.userReq.UpdatedBy = updatedBy - expected, err := repo.UpdateEmail(context.Background(), c.userReq) - assert.True(t, errors.Contains(err, c.err), fmt.Sprintf("expected %s to contain %s\n", err, c.err)) - if err == nil { - assert.Equal(t, c.userReq.Email, expected.Email) - assert.Equal(t, c.userReq.UpdatedAt, expected.UpdatedAt) - assert.Equal(t, c.userReq.UpdatedBy, expected.UpdatedBy) - } - }) - } -} - -func TestUpdate(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM users") - require.Nil(t, err, fmt.Sprintf("clean users unexpected error: %s", err)) - }) - repo := cpostgres.NewRepository(database) - - user1 := generateUser(t, users.EnabledStatus, repo) - user2 := generateUser(t, users.DisabledStatus, repo) - user3 := generateUser(t, users.EnabledStatus, repo) - - updatedMetadata := users.Metadata{"update": namesgen.Generate()} - malformedMetadata := users.Metadata{"update": make(chan int)} - updatedLastName := namesgen.Generate() - updatedFirstName := namesgen.Generate() - updateTags := namesgen.GenerateMultiple(5) - updatedProfilePicture := namesgen.Generate() - emptyName := "" - emptyTags := []string{} - - cases := []struct { - desc string - update string - userID string - userReq users.UserReq - userRes users.User - err error - }{ - { - desc: "update metadata for enabled user", - update: "metadata", - userID: user1.ID, - userReq: users.UserReq{ - Metadata: &updatedMetadata, - }, - userRes: users.User{ - Metadata: updatedMetadata, - }, - err: nil, - }, - { - desc: "update private metadata for enabled user", - update: "private_metadata", - userID: user1.ID, - userReq: users.UserReq{ - PrivateMetadata: &updatedMetadata, - }, - userRes: users.User{ - PrivateMetadata: updatedMetadata, - }, - err: nil, - }, - { - desc: "update malformed private metadata for enabled user", - update: "private_metadata", - userID: user1.ID, - userReq: users.UserReq{ - PrivateMetadata: &malformedMetadata, - }, - err: repoerr.ErrMalformedEntity, - }, - { - desc: "update empty metadata for enabled user", - update: "metadata", - userID: user3.ID, - userReq: users.UserReq{ - Metadata: &users.Metadata{}, - }, - userRes: users.User{ - Metadata: users.Metadata{}, - }, - err: nil, - }, - { - desc: "update metadata for disabled user", - update: "metadata", - userID: user2.ID, - userReq: users.UserReq{ - Metadata: &updatedMetadata, - }, - err: repoerr.ErrNotFound, - }, - { - desc: "update first name for enabled user", - update: "first_name", - userID: user1.ID, - userReq: users.UserReq{ - FirstName: &updatedFirstName, - }, - userRes: users.User{ - FirstName: updatedFirstName, - }, - err: nil, - }, - { - desc: "update empty first name for enabled user", - update: "first_name", - userID: user3.ID, - userReq: users.UserReq{ - FirstName: &emptyName, - }, - userRes: user3, - err: nil, - }, - { - desc: "update first name for disabled user", - update: "first_name", - userID: user2.ID, - userReq: users.UserReq{ - FirstName: &updatedFirstName, - }, - err: repoerr.ErrNotFound, - }, - { - desc: "update private metadata for invalid user", - update: "private_metadata", - userID: testsutil.GenerateUUID(t), - userReq: users.UserReq{ - PrivateMetadata: &updatedMetadata, - }, - err: repoerr.ErrNotFound, - }, - { - desc: "update first name for empty user", - update: "first_name", - userID: "", - userReq: users.UserReq{ - FirstName: &updatedFirstName, - }, - err: repoerr.ErrNotFound, - }, - { - desc: "update last name for enabled user", - update: "last_name", - userID: user1.ID, - userReq: users.UserReq{ - LastName: &updatedLastName, - }, - userRes: users.User{ - LastName: updatedLastName, - }, - err: nil, - }, - { - desc: "update empty last name for enabled user", - update: "last_name", - userID: user3.ID, - userReq: users.UserReq{ - LastName: &emptyName, - }, - userRes: user3, - err: nil, - }, - { - desc: "update last name for disabled user", - update: "last_name", - userID: user2.ID, - userReq: users.UserReq{ - LastName: &updatedLastName, - }, - err: repoerr.ErrNotFound, - }, - { - desc: "update last name for invalid user", - update: "last_name", - userID: testsutil.GenerateUUID(t), - userReq: users.UserReq{ - LastName: &updatedLastName, - }, - err: repoerr.ErrNotFound, - }, - { - desc: "update tags for enabled user", - userID: user1.ID, - userReq: users.UserReq{ - Tags: &updateTags, - }, - userRes: users.User{ - Tags: updateTags, - }, - err: nil, - }, - { - desc: "update empty tags for enabled user", - userID: user3.ID, - userReq: users.UserReq{ - Tags: &emptyTags, - }, - userRes: users.User{ - Tags: []string{}, - }, - err: nil, - }, - { - desc: "update tags for disabled user", - userID: user2.ID, - userReq: users.UserReq{ - Tags: &updateTags, - }, - err: repoerr.ErrNotFound, - }, - { - desc: "update tags for invalid user", - userID: testsutil.GenerateUUID(t), - userReq: users.UserReq{ - Tags: &updateTags, - }, - err: repoerr.ErrNotFound, - }, - { - desc: "update profile picture for enabled user", - userID: user1.ID, - userReq: users.UserReq{ - ProfilePicture: &updatedProfilePicture, - }, - userRes: users.User{ - ProfilePicture: updatedProfilePicture, - }, - err: nil, - }, - { - desc: "update empty profile picture for enabled user", - userID: user3.ID, - userReq: users.UserReq{ - ProfilePicture: &emptyName, - }, - userRes: users.User{ - ProfilePicture: emptyName, - }, - err: nil, - }, - { - desc: "update profile picture for disabled user", - userID: user2.ID, - userReq: users.UserReq{ - ProfilePicture: &updatedProfilePicture, - }, - err: repoerr.ErrNotFound, - }, - { - desc: "update profile picture for invalid user", - userID: testsutil.GenerateUUID(t), - userReq: users.UserReq{ - ProfilePicture: &updatedProfilePicture, - }, - err: repoerr.ErrNotFound, - }, - } - - for _, c := range cases { - t.Run(c.desc, func(t *testing.T) { - updatedAt := time.Now().UTC().Truncate(time.Microsecond) - updatedBy := testsutil.GenerateUUID(t) - c.userReq.UpdatedAt = &updatedAt - c.userReq.UpdatedBy = &updatedBy - c.userRes.UpdatedAt = updatedAt - c.userRes.UpdatedBy = updatedBy - expected, err := repo.Update(context.Background(), c.userID, c.userReq) - assert.True(t, errors.Contains(err, c.err), fmt.Sprintf("expected %s to contain %s\n", err, c.err)) - if err == nil { - switch c.update { - case "private_metadata": - assert.Equal(t, c.userRes.PrivateMetadata, expected.PrivateMetadata) - case "metadata": - assert.Equal(t, c.userRes.Metadata, expected.Metadata) - case "first_name": - assert.Equal(t, c.userRes.FirstName, expected.FirstName) - case "last_name": - assert.Equal(t, c.userRes.LastName, expected.LastName) - case "tags": - assert.Equal(t, c.userRes.Tags, expected.Tags) - case "profile_picture": - assert.Equal(t, c.userRes.ProfilePicture, expected.ProfilePicture) - case "role": - assert.Equal(t, c.userRes.Role, expected.Role) - case "email": - assert.Equal(t, c.userRes.Email, expected.Email) - } - assert.Equal(t, c.userRes.UpdatedAt, expected.UpdatedAt) - assert.Equal(t, c.userRes.UpdatedBy, expected.UpdatedBy) - } - }) - } -} - -func TestUpdateUsername(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM users") - require.Nil(t, err, fmt.Sprintf("clean users unexpected error: %s", err)) - }) - repo := cpostgres.NewRepository(database) - - user1 := generateUser(t, users.EnabledStatus, repo) - user2 := generateUser(t, users.DisabledStatus, repo) - - cases := []struct { - desc string - user users.User - err error - }{ - { - desc: "for enabled user", - user: users.User{ - ID: user1.ID, - Credentials: users.Credentials{ - Username: namesgen.Generate(), - }, - }, - err: nil, - }, - { - desc: "for enabled user with existing username", - user: users.User{ - ID: user1.ID, - Credentials: users.Credentials{ - Username: user2.Credentials.Username, - }, - }, - err: errors.ErrUsernameNotAvailable, - }, - { - desc: "for disabled user", - user: users.User{ - ID: user2.ID, - Credentials: users.Credentials{ - Username: namesgen.Generate(), - }, - }, - err: repoerr.ErrNotFound, - }, - { - desc: "for invalid user", - user: users.User{ - ID: testsutil.GenerateUUID(t), - Credentials: users.Credentials{ - Username: namesgen.Generate(), - }, - }, - err: repoerr.ErrNotFound, - }, - { - desc: "for empty user", - user: users.User{}, - err: repoerr.ErrNotFound, - }, - } - for _, c := range cases { - t.Run(c.desc, func(t *testing.T) { - c.user.UpdatedAt = time.Now().UTC().Truncate(time.Microsecond) - c.user.UpdatedBy = testsutil.GenerateUUID(t) - expected, err := repo.UpdateUsername(context.Background(), c.user) - assert.True(t, errors.Contains(err, c.err), fmt.Sprintf("expected %s to contain %s\n", err, c.err)) - if err == nil { - assert.Equal(t, c.user.Credentials.Username, expected.Credentials.Username) - assert.Equal(t, c.user.UpdatedAt, expected.UpdatedAt) - assert.Equal(t, c.user.UpdatedBy, expected.UpdatedBy) - } - }) - } -} - -func TestUpdateSecret(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM users") - require.Nil(t, err, fmt.Sprintf("clean users unexpected error: %s", err)) - }) - repo := cpostgres.NewRepository(database) - - user1 := generateUser(t, users.EnabledStatus, repo) - user2 := generateUser(t, users.DisabledStatus, repo) - - cases := []struct { - desc string - user users.User - err error - }{ - { - desc: "for enabled user", - user: users.User{ - ID: user1.ID, - Credentials: users.Credentials{ - Secret: "newpassword", - }, - }, - err: nil, - }, - { - desc: "for disabled user", - user: users.User{ - ID: user2.ID, - Credentials: users.Credentials{ - Secret: "newpassword", - }, - }, - err: repoerr.ErrNotFound, - }, - { - desc: "for invalid user", - user: users.User{ - ID: testsutil.GenerateUUID(t), - Credentials: users.Credentials{ - Secret: "newpassword", - }, - }, - err: repoerr.ErrNotFound, - }, - { - desc: "for empty user", - user: users.User{}, - err: repoerr.ErrNotFound, - }, - } - for _, c := range cases { - t.Run(c.desc, func(t *testing.T) { - c.user.UpdatedAt = time.Now().UTC().Truncate(time.Microsecond) - c.user.UpdatedBy = testsutil.GenerateUUID(t) - _, err := repo.UpdateSecret(context.Background(), c.user) - assert.True(t, errors.Contains(err, c.err), fmt.Sprintf("expected %s to contain %s\n", err, c.err)) - if err == nil { - rc, err := repo.RetrieveByID(context.Background(), c.user.ID) - require.Nil(t, err, fmt.Sprintf("retrieve user by id during update of secret unexpected error: %s", err)) - assert.Equal(t, c.user.Credentials.Secret, rc.Credentials.Secret) - assert.Equal(t, c.user.UpdatedAt, rc.UpdatedAt) - assert.Equal(t, c.user.UpdatedBy, rc.UpdatedBy) - } - }) - } -} - -func TestChangeStatus(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM users") - require.Nil(t, err, fmt.Sprintf("clean users unexpected error: %s", err)) - }) - repo := cpostgres.NewRepository(database) - - user1 := generateUser(t, users.EnabledStatus, repo) - user2 := generateUser(t, users.DisabledStatus, repo) - - cases := []struct { - desc string - user users.User - err error - }{ - { - desc: "for an enabled user", - user: users.User{ - ID: user1.ID, - Status: users.DisabledStatus, - }, - err: nil, - }, - { - desc: "for a disabled user", - user: users.User{ - ID: user2.ID, - Status: users.EnabledStatus, - }, - err: nil, - }, - { - desc: "for invalid user", - user: users.User{ - ID: testsutil.GenerateUUID(t), - Status: users.DisabledStatus, - }, - err: repoerr.ErrNotFound, - }, - { - desc: "for empty user", - user: users.User{}, - err: repoerr.ErrNotFound, - }, - } - - for _, c := range cases { - t.Run(c.desc, func(t *testing.T) { - c.user.UpdatedAt = time.Now().UTC().Truncate(time.Microsecond) - c.user.UpdatedBy = testsutil.GenerateUUID(t) - expected, err := repo.ChangeStatus(context.Background(), c.user) - assert.True(t, errors.Contains(err, c.err), fmt.Sprintf("expected %s to contain %s\n", err, c.err)) - if err == nil { - assert.Equal(t, c.user.Status, expected.Status) - assert.Equal(t, c.user.UpdatedAt, expected.UpdatedAt) - assert.Equal(t, c.user.UpdatedBy, expected.UpdatedBy) - } - }) - } -} - -func TestDelete(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM users") - require.Nil(t, err, fmt.Sprintf("clean users unexpected error: %s", err)) - }) - repo := cpostgres.NewRepository(database) - - user := generateUser(t, users.EnabledStatus, repo) - - cases := []struct { - desc string - id string - err error - }{ - { - desc: "delete user successfully", - id: user.ID, - err: nil, - }, - { - desc: "delete user with invalid id", - id: testsutil.GenerateUUID(t), - err: repoerr.ErrNotFound, - }, - { - desc: "delete user with empty id", - id: "", - err: repoerr.ErrNotFound, - }, - } - - for _, tc := range cases { - err := repo.Delete(context.Background(), tc.id) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - } -} - -func TestRetrieveByIDs(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM users") - require.Nil(t, err, fmt.Sprintf("clean users unexpected error: %s", err)) - }) - repo := cpostgres.NewRepository(database) - - num := 200 - - var items []users.User - baseTime := time.Now().UTC().Truncate(time.Millisecond) - for i := 0; i < num; i++ { - user := generateUserWithTime(t, users.EnabledStatus, repo, baseTime.Add(time.Duration(i)*time.Millisecond)) - user.PrivateMetadata = nil - items = append(items, user) - } - - page, err := repo.RetrieveAll(context.Background(), users.Page{Offset: 0, Limit: uint64(num)}) - require.Nil(t, err, fmt.Sprintf("retrieve all users unexpected error: %s", err)) - assert.Equal(t, uint64(num), page.Total) - - cases := []struct { - desc string - page users.Page - response users.UsersPage - err error - }{ - { - desc: "successfully", - page: users.Page{ - Offset: 0, - Limit: 10, - IDs: getIDs(items[0:3]), - Order: "created_at", - Dir: ascDir, - }, - response: users.UsersPage{ - Page: users.Page{ - Total: 3, - Offset: 0, - Limit: 10, - }, - Users: items[0:3], - }, - err: nil, - }, - { - desc: "with empty ids", - page: users.Page{ - Offset: 0, - Limit: 10, - IDs: []string{}, - }, - response: users.UsersPage{ - Page: users.Page{ - Offset: 0, - Limit: 10, - }, - Users: []users.User(nil), - }, - err: nil, - }, - { - desc: "with offset only", - page: users.Page{ - Offset: 10, - IDs: getIDs(items[0:20]), - Order: "created_at", - Dir: ascDir, - }, - response: users.UsersPage{ - Page: users.Page{ - Total: 20, - Offset: 10, - Limit: 0, - }, - Users: []users.User(nil), - }, - err: nil, - }, - { - desc: "with limit only", - page: users.Page{ - Limit: 10, - IDs: getIDs(items[0:20]), - Order: "created_at", - Dir: ascDir, - }, - response: users.UsersPage{ - Page: users.Page{ - Total: 20, - Offset: 0, - Limit: 10, - }, - Users: items[0:10], - }, - err: nil, - }, - { - desc: "with offset out of range", - page: users.Page{ - Offset: 1000, - Limit: 50, - IDs: getIDs(items[0:20]), - }, - response: users.UsersPage{ - Page: users.Page{ - Total: 20, - Offset: 1000, - Limit: 50, - }, - Users: []users.User(nil), - }, - err: nil, - }, - { - desc: "with offset and limit out of range", - page: users.Page{ - Offset: 15, - Limit: 10, - IDs: getIDs(items[0:20]), - Order: "created_at", - Dir: ascDir, - }, - response: users.UsersPage{ - Page: users.Page{ - Total: 20, - Offset: 15, - Limit: 10, - }, - Users: items[15:20], - }, - err: nil, - }, - { - desc: "with limit out of range", - page: users.Page{ - Offset: 0, - Limit: 1000, - IDs: getIDs(items[0:20]), - }, - response: users.UsersPage{ - Page: users.Page{ - Total: 20, - Offset: 0, - Limit: 1000, - }, - Users: items[:20], - }, - err: nil, - }, - { - desc: "with first name", - page: users.Page{ - Offset: 0, - Limit: 10, - FirstName: items[0].FirstName, - IDs: getIDs(items[0:20]), - Order: "created_at", - Dir: ascDir, - }, - response: users.UsersPage{ - Page: users.Page{ - Total: 1, - Offset: 0, - Limit: 10, - }, - Users: []users.User{items[0]}, - }, - err: nil, - }, - { - desc: "with metadata", - page: users.Page{ - Offset: 0, - Limit: 10, - Metadata: items[0].Metadata, - IDs: getIDs(items[0:20]), - }, - response: users.UsersPage{ - Page: users.Page{ - Total: 1, - Offset: 0, - Limit: 10, - }, - Users: []users.User{items[0]}, - }, - err: nil, - }, - { - desc: "with invalid metadata", - page: users.Page{ - Offset: 0, - Limit: 10, - Metadata: map[string]any{ - "key": make(chan int), - }, - IDs: getIDs(items[0:20]), - }, - response: users.UsersPage{ - Page: users.Page{ - Total: 0, - Offset: 0, - Limit: 10, - }, - Users: []users.User(nil), - }, - err: repoerr.ErrViewEntity, - }, - } - - for _, c := range cases { - t.Run(c.desc, func(t *testing.T) { - switch response, err := repo.RetrieveAllByIDs(context.Background(), c.page); { - case err == nil: - assert.Nil(t, err, fmt.Sprintf("%s: expected %s got %s\n", c.desc, c.err, err)) - assert.Equal(t, c.response.Total, response.Total) - assert.Equal(t, c.response.Limit, response.Limit) - assert.Equal(t, c.response.Offset, response.Offset) - assert.ElementsMatch(t, response.Users, c.response.Users) - default: - assert.True(t, errors.Contains(err, c.err), fmt.Sprintf("expected %s to contain %s\n", err, c.err)) - } - }) - } -} - -func TestRetrieveByEmail(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM users") - require.Nil(t, err, fmt.Sprintf("clean users unexpected error: %s", err)) - }) - repo := cpostgres.NewRepository(database) - - user := generateUser(t, users.EnabledStatus, repo) - - cases := []struct { - desc string - email string - response users.User - err error - }{ - { - desc: "successfully", - email: user.Email, - response: user, - err: nil, - }, - { - desc: "with invalid user id", - email: testsutil.GenerateUUID(t), - response: users.User{}, - err: repoerr.ErrNotFound, - }, - { - desc: "with empty user id", - email: "", - response: users.User{}, - err: repoerr.ErrNotFound, - }, - } - for _, c := range cases { - t.Run(c.desc, func(t *testing.T) { - usr, err := repo.RetrieveByEmail(context.Background(), c.email) - assert.True(t, errors.Contains(err, c.err), fmt.Sprintf("expected %s got %s\n", c.err, err)) - if err == nil { - assert.Equal(t, user.ID, usr.ID) - assert.Equal(t, user.FirstName, usr.FirstName) - assert.Equal(t, user.LastName, usr.LastName) - assert.Equal(t, user.Metadata, usr.Metadata) - assert.Equal(t, user.PrivateMetadata, usr.PrivateMetadata) - assert.Equal(t, user.Email, usr.Email) - assert.Equal(t, user.Credentials.Username, usr.Credentials.Username) - assert.Equal(t, user.Status, usr.Status) - } - }) - } -} - -func TestRetrieveByUsername(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM users") - require.Nil(t, err, fmt.Sprintf("clean users unexpected error: %s", err)) - }) - repo := cpostgres.NewRepository(database) - - user := generateUser(t, users.EnabledStatus, repo) - - cases := []struct { - desc string - username string - response users.User - err error - }{ - { - desc: "successfully", - username: user.Credentials.Username, - response: user, - err: nil, - }, - { - desc: "with invalid user id", - username: testsutil.GenerateUUID(t), - response: users.User{}, - err: repoerr.ErrNotFound, - }, - { - desc: "with empty user id", - username: "", - response: users.User{}, - err: repoerr.ErrNotFound, - }, - } - for _, c := range cases { - t.Run(c.desc, func(t *testing.T) { - usr, err := repo.RetrieveByUsername(context.Background(), c.username) - assert.True(t, errors.Contains(err, c.err), fmt.Sprintf("expected %s got %s\n", c.err, err)) - if err == nil { - assert.Equal(t, user.ID, usr.ID) - assert.Equal(t, user.FirstName, usr.FirstName) - assert.Equal(t, user.LastName, usr.LastName) - assert.Equal(t, user.PrivateMetadata, usr.PrivateMetadata) - assert.Equal(t, user.Email, usr.Email) - assert.Equal(t, user.Credentials.Username, usr.Credentials.Username) - assert.Equal(t, user.Status, usr.Status) - } - }) - } -} - -func findUsers(usrs []users.User, query string, offset, limit uint64) []users.User { - rUsers := []users.User{} - for _, user := range usrs { - if strings.Contains(user.FirstName, query) { - rUsers = append(rUsers, user) - } - } - - if offset > uint64(len(rUsers)) { - return []users.User{} - } - - if limit > uint64(len(rUsers)) { - return rUsers[offset:] - } - - return rUsers[offset:limit] -} - -func generateUser(t *testing.T, status users.Status, repo users.Repository) users.User { - return generateUserWithTime(t, status, repo, time.Now().UTC().Truncate(time.Millisecond)) -} - -func generateUserWithTime(t *testing.T, status users.Status, repo users.Repository, createdAt time.Time) users.User { - usr := users.User{ - ID: testsutil.GenerateUUID(t), - FirstName: namesgen.Generate(), - LastName: namesgen.Generate(), - Email: namesgen.Generate() + emailSuffix, - Credentials: users.Credentials{ - Username: namesgen.Generate(), - Secret: testsutil.GenerateUUID(t), - }, - Tags: namesgen.GenerateMultiple(5), - PrivateMetadata: users.Metadata{ - "organization": namesgen.Generate(), - }, - Metadata: users.Metadata{ - "address": namesgen.Generate(), - }, - Status: status, - CreatedAt: createdAt, - } - user, err := repo.Save(context.Background(), usr) - require.Nil(t, err, fmt.Sprintf("add new user: expected nil got %s\n", err)) - - return user -} - -func getIDs(usrs []users.User) []string { - var ids []string - for _, user := range usrs { - ids = append(ids, user.ID) - } - - return ids -} - -func stripUserDetails(users []users.User) []users.User { - for i := range users { - users[i].CreatedAt = validTimestamp - users[i].UpdatedAt = validTimestamp - } - return users -} - -func verifyUsersOrdering(t *testing.T, users []users.User, order, dir string) { - if order == "" || len(users) <= 1 { - return - } - - for i := 0; i < len(users)-1; i++ { - switch order { - case "first_name": - if dir == ascDir { - assert.LessOrEqual(t, users[i].FirstName, users[i+1].FirstName, fmt.Sprintf("Users not ordered by first_name ascending at index %d: %s > %s", i, users[i].FirstName, users[i+1].FirstName)) - continue - } - assert.GreaterOrEqual(t, users[i].FirstName, users[i+1].FirstName, fmt.Sprintf("Users not ordered by first_name descending at index %d: %s < %s", i, users[i].FirstName, users[i+1].FirstName)) - case "username": - if dir == ascDir { - assert.LessOrEqual(t, users[i].Credentials.Username, users[i+1].Credentials.Username, fmt.Sprintf("Users not ordered by username ascending at index %d: %s > %s", i, users[i].Credentials.Username, users[i+1].Credentials.Username)) - continue - } - assert.GreaterOrEqual(t, users[i].Credentials.Username, users[i+1].Credentials.Username, fmt.Sprintf("Users not ordered by username descending at index %d: %s < %s", i, users[i].Credentials.Username, users[i+1].Credentials.Username)) - case "created_at": - if dir == ascDir { - assert.False(t, users[i].CreatedAt.After(users[i+1].CreatedAt), fmt.Sprintf("Users not ordered by created_at ascending at index %d: %v > %v", i, users[i].CreatedAt, users[i+1].CreatedAt)) - continue - } - assert.False(t, users[i].CreatedAt.Before(users[i+1].CreatedAt), fmt.Sprintf("Users not ordered by created_at descending at index %d: %v < %v", i, users[i].CreatedAt, users[i+1].CreatedAt)) - case "updated_at": - if dir == ascDir { - assert.False(t, users[i].UpdatedAt.After(users[i+1].UpdatedAt), fmt.Sprintf("Users not ordered by updated_at ascending at index %d: %v > %v", i, users[i].UpdatedAt, users[i+1].UpdatedAt)) - continue - } - assert.False(t, users[i].UpdatedAt.Before(users[i+1].UpdatedAt), fmt.Sprintf("Users not ordered by updated_at descending at index %d: %v < %v", i, users[i].UpdatedAt, users[i+1].UpdatedAt)) - } - } -} diff --git a/users/postgres/verfications.go b/users/postgres/verfications.go deleted file mode 100644 index 0b36c2f7b..000000000 --- a/users/postgres/verfications.go +++ /dev/null @@ -1,131 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package postgres - -import ( - "context" - "database/sql" - "time" - - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - "github.com/absmach/magistrala/users" -) - -// AddUserVerification adds new verification for given user id and email. -func (repo *userRepo) AddUserVerification(ctx context.Context, uv users.UserVerification) error { - q := `INSERT INTO users_verifications (user_id, email, otp, created_at, expires_at ) - VALUES (:user_id, :email, :otp, :created_at, :expires_at );` - dbuv := toDBUserVerification(uv) - if _, err := repo.Repository.DB.NamedExecContext(ctx, q, dbuv); err != nil { - return errors.Wrap(repoerr.ErrCreateEntity, err) - } - return nil -} - -// RetrieveUserVerification retrieves verification details of given user id and email. -func (repo *userRepo) RetrieveUserVerification(ctx context.Context, userID, email string) (users.UserVerification, error) { - dbuv := dbUserVerification{ - UserID: userID, - Email: email, - } - q := `SELECT user_id, email, otp, created_at, expires_at , used_at FROM users_verifications WHERE user_id = :user_id AND email = :email ORDER BY created_at DESC LIMIT 1 ` - - row, err := repo.Repository.DB.NamedQueryContext(ctx, q, dbuv) - if err != nil { - return users.UserVerification{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - if !row.Next() { - return users.UserVerification{}, repoerr.ErrNotFound - } - - defer row.Close() - - if err := row.StructScan(&dbuv); err != nil { - return users.UserVerification{}, errors.Wrap(repoerr.ErrViewEntity, err) - } - - return toUserVerification(dbuv), nil -} - -// UpdateUserVerification update user verification details for the given user id and email. -func (repo *userRepo) UpdateUserVerification(ctx context.Context, uv users.UserVerification) error { - q := `UPDATE users_verifications SET otp = :otp, used_at = :used_at WHERE user_id = :user_id AND email = :email` - dbuv := toDBUserVerification(uv) - res, err := repo.Repository.DB.NamedExecContext(ctx, q, dbuv) - if err != nil { - return errors.Wrap(repoerr.ErrUpdateEntity, err) - } - rows, err := res.RowsAffected() - if err != nil { - return errors.Wrap(repoerr.ErrUpdateEntity, err) - } - if rows == 0 { - return repoerr.ErrNotFound - } - return nil -} - -type dbUserVerification struct { - UserID string `db:"user_id"` - Email string `db:"email"` - OTP sql.NullString `db:"otp"` - CreatedAt sql.NullTime `db:"created_at"` - ExpiresAt sql.NullTime `db:"expires_at"` - UsedAt sql.NullTime `db:"used_at"` -} - -func toDBUserVerification(uv users.UserVerification) dbUserVerification { - var otp sql.NullString - if uv.OTP != "" { - otp = sql.NullString{String: uv.OTP, Valid: true} - } - var createdAt sql.NullTime - if !uv.CreatedAt.IsZero() { - createdAt = sql.NullTime{Time: uv.CreatedAt, Valid: true} - } - var expiresAt sql.NullTime - if !uv.ExpiresAt.IsZero() { - expiresAt = sql.NullTime{Time: uv.ExpiresAt, Valid: true} - } - var usedAt sql.NullTime - if !uv.UsedAt.IsZero() { - usedAt = sql.NullTime{Time: uv.UsedAt, Valid: true} - } - - return dbUserVerification{ - UserID: uv.UserID, - Email: uv.Email, - OTP: otp, - CreatedAt: createdAt, - ExpiresAt: expiresAt, - UsedAt: usedAt, - } -} - -func toUserVerification(dbuv dbUserVerification) users.UserVerification { - var createdAt time.Time - if dbuv.CreatedAt.Valid { - createdAt = dbuv.CreatedAt.Time.UTC() - } - - var expiresAt time.Time - if dbuv.ExpiresAt.Valid { - expiresAt = dbuv.ExpiresAt.Time.UTC() - } - - var usedAt time.Time - if dbuv.UsedAt.Valid { - usedAt = dbuv.UsedAt.Time.UTC() - } - - return users.UserVerification{ - UserID: dbuv.UserID, - Email: dbuv.Email, - OTP: dbuv.OTP.String, - CreatedAt: createdAt, - ExpiresAt: expiresAt, - UsedAt: usedAt, - } -} diff --git a/users/postgres/verifications_test.go b/users/postgres/verifications_test.go deleted file mode 100644 index 0c26d2eb7..000000000 --- a/users/postgres/verifications_test.go +++ /dev/null @@ -1,217 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package postgres_test - -import ( - "context" - "fmt" - "testing" - "time" - - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - "github.com/absmach/magistrala/users" - "github.com/absmach/magistrala/users/postgres" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" -) - -func TestAddUserVerification(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM users") - require.Nil(t, err, fmt.Sprintf("clean users unexpected error: %s", err)) - _, err = db.Exec("DELETE FROM users_verifications") - require.Nil(t, err, fmt.Sprintf("clean users_verifications unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - first_name := namesgen.Generate() - last_name := namesgen.Generate() - username := namesgen.Generate() - user := users.User{ - ID: "test-user-id", - Email: "test@example.com", - FirstName: first_name, - LastName: last_name, - Credentials: users.Credentials{ - Username: username, - }, - } - _, err := repo.Save(context.Background(), user) - require.Nil(t, err, fmt.Sprintf("saving user unexpected error: %s", err)) - - cases := []struct { - desc string - uv users.UserVerification - err error - }{ - { - desc: "add new user verification", - uv: users.UserVerification{ - UserID: user.ID, - Email: user.Email, - CreatedAt: time.Now().UTC(), - OTP: "123456", - ExpiresAt: time.Now().UTC().Add(time.Hour), - }, - err: nil, - }, - { - desc: "add user verification for non-existing user", - uv: users.UserVerification{ - UserID: "non-existing-user", - Email: "non-existing@example.com", - OTP: "654321", - CreatedAt: time.Now().UTC(), - ExpiresAt: time.Now().UTC().Add(time.Hour), - }, - err: repoerr.ErrCreateEntity, - }, - } - - for _, tc := range cases { - err := repo.AddUserVerification(context.Background(), tc.uv) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s", tc.desc, tc.err, err)) - } -} - -func TestRetrieveUserVerification(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM users") - require.Nil(t, err, fmt.Sprintf("clean users unexpected error: %s", err)) - _, err = db.Exec("DELETE FROM users_verifications") - require.Nil(t, err, fmt.Sprintf("clean users_verifications unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - first_name := namesgen.Generate() - last_name := namesgen.Generate() - username := namesgen.Generate() - user := users.User{ - ID: "test-user-id", - Email: "test@example.com", - FirstName: first_name, - LastName: last_name, - Credentials: users.Credentials{ - Username: username, - }, - } - _, err := repo.Save(context.Background(), user) - require.Nil(t, err, fmt.Sprintf("saving user unexpected error: %s", err)) - - uv := users.UserVerification{ - UserID: user.ID, - Email: user.Email, - OTP: "123456", - CreatedAt: time.Now(), - ExpiresAt: time.Now().Add(time.Hour), - } - err = repo.AddUserVerification(context.Background(), uv) - require.Nil(t, err, fmt.Sprintf("adding user verification unexpected error: %s", err)) - - cases := []struct { - desc string - userID string - email string - err error - }{ - { - desc: "retrieve existing user verification", - userID: user.ID, - email: user.Email, - err: nil, - }, - { - desc: "retrieve non-existing user verification", - userID: "non-existing-user", - email: "non-existing@example.com", - err: repoerr.ErrNotFound, - }, - } - - for _, tc := range cases { - retrieved, err := repo.RetrieveUserVerification(context.Background(), tc.userID, tc.email) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s", tc.desc, tc.err, err)) - if err == nil { - assert.Equal(t, uv.UserID, retrieved.UserID, fmt.Sprintf("%s: expected %v got %v", tc.desc, uv.UserID, retrieved.UserID)) - assert.Equal(t, uv.Email, retrieved.Email, fmt.Sprintf("%s: expected %v got %v", tc.desc, uv.Email, retrieved.Email)) - assert.Equal(t, uv.OTP, retrieved.OTP, fmt.Sprintf("%s: expected %v got %v", tc.desc, uv.OTP, retrieved.OTP)) - } - } -} - -func TestUpdateUserVerification(t *testing.T) { - t.Cleanup(func() { - _, err := db.Exec("DELETE FROM users") - require.Nil(t, err, fmt.Sprintf("clean users unexpected error: %s", err)) - _, err = db.Exec("DELETE FROM users_verifications") - require.Nil(t, err, fmt.Sprintf("clean users_verifications unexpected error: %s", err)) - }) - - repo := postgres.NewRepository(database) - - first_name := namesgen.Generate() - last_name := namesgen.Generate() - username := namesgen.Generate() - user := users.User{ - ID: "test-user-id", - Email: "test@example.com", - FirstName: first_name, - LastName: last_name, - Credentials: users.Credentials{ - Username: username, - }, - } - _, err := repo.Save(context.Background(), user) - require.Nil(t, err, fmt.Sprintf("saving user unexpected error: %s", err)) - - uv := users.UserVerification{ - UserID: user.ID, - Email: user.Email, - OTP: "123456", - CreatedAt: time.Now().UTC(), - ExpiresAt: time.Now().UTC().Add(time.Hour), - } - err = repo.AddUserVerification(context.Background(), uv) - require.Nil(t, err, fmt.Sprintf("adding user verification unexpected error: %s", err)) - - usedTime := time.Now() - cases := []struct { - desc string - uv users.UserVerification - err error - }{ - { - desc: "update existing user verification", - uv: users.UserVerification{ - UserID: user.ID, - Email: user.Email, - OTP: "654321", - UsedAt: usedTime, - }, - err: nil, - }, - { - desc: "update non-existing user verification", - uv: users.UserVerification{ - UserID: "non-existing-user", - Email: "non-existing@example.com", - }, - err: repoerr.ErrNotFound, - }, - } - - for _, tc := range cases { - err := repo.UpdateUserVerification(context.Background(), tc.uv) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s", tc.desc, tc.err, err)) - if err == nil { - retrieved, err := repo.RetrieveUserVerification(context.Background(), tc.uv.UserID, tc.uv.Email) - require.Nil(t, err, fmt.Sprintf("retrieving updated verification unexpected error: %s", err)) - assert.Equal(t, tc.uv.OTP, retrieved.OTP, fmt.Sprintf("%s: expected %v got %v", tc.desc, tc.uv.OTP, retrieved.OTP)) - assert.WithinDuration(t, tc.uv.UsedAt, retrieved.UsedAt, 10*time.Second, fmt.Sprintf("%s: expected %v got %v", tc.desc, tc.uv.UsedAt, retrieved.UsedAt)) - } - } -} diff --git a/users/private/service.go b/users/private/service.go deleted file mode 100644 index bcaa4202d..000000000 --- a/users/private/service.go +++ /dev/null @@ -1,51 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package private - -import ( - "context" - - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/users" -) - -type Service interface { - RetrieveByIDs(ctx context.Context, ids []string, offset, limit uint64) (users.UsersPage, error) -} - -var _ Service = (*service)(nil) - -func New(repo users.Repository) Service { - return service{ - repo: repo, - } -} - -type service struct { - repo users.Repository -} - -func (svc service) RetrieveByIDs(ctx context.Context, ids []string, offset, limit uint64) (users.UsersPage, error) { - if len(ids) == 0 { - return users.UsersPage{}, svcerr.ErrMalformedEntity - } - - if limit == 0 { - limit = uint64(len(ids)) - } - - pm := users.Page{ - IDs: ids, - Offset: offset, - Limit: limit, - } - - page, err := svc.repo.RetrieveAllByIDs(ctx, pm) - if err != nil { - return users.UsersPage{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - - return page, nil -} diff --git a/users/roles.go b/users/roles.go deleted file mode 100644 index 5cde293c8..000000000 --- a/users/roles.go +++ /dev/null @@ -1,71 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package users - -import ( - "encoding/json" - "strings" - - apiutil "github.com/absmach/magistrala/api/http/util" -) - -// Role represents User role. -type Role uint8 - -// Possible User role values. -const ( - UserRole Role = iota - AdminRole - - // AllRole is used for querying purposes to list users irrespective - // of their role - both admin and user. It is never stored in the - // database as the actual user role and should always be the largest - // value in this enumeration. - AllRole -) - -// String representation of the possible role values. -const ( - Admin = "admin" - user = "user" -) - -// String converts user role to string literal. -func (cs Role) String() string { - switch cs { - case AdminRole: - return Admin - case UserRole: - return user - case AllRole: - return All - default: - return Unknown - } -} - -// ToRole converts string value to a valid User role. -func ToRole(status string) (Role, error) { - switch status { - case "", user: - return UserRole, nil - case Admin: - return AdminRole, nil - case All: - return AllRole, nil - default: - return Role(0), apiutil.ErrInvalidRole - } -} - -func (r Role) MarshalJSON() ([]byte, error) { - return json.Marshal(r.String()) -} - -func (r *Role) UnmarshalJSON(data []byte) error { - str := strings.Trim(string(data), "\"") - val, err := ToRole(str) - *r = val - return err -} diff --git a/users/service.go b/users/service.go deleted file mode 100644 index 8217acdc1..000000000 --- a/users/service.go +++ /dev/null @@ -1,858 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package users - -import ( - "context" - "crypto/rand" - "encoding/hex" - "fmt" - "net/mail" - "regexp" - "strings" - "time" - - "github.com/absmach/magistrala" - grpcTokenV1 "github.com/absmach/magistrala/api/grpc/token/v1" - apiutil "github.com/absmach/magistrala/api/http/util" - smqauth "github.com/absmach/magistrala/auth" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - "github.com/absmach/magistrala/pkg/policies" - "github.com/gofrs/uuid/v5" -) - -const defaultUsernamePrefix = "user" - -var ( - errRecoveryToken = errors.NewServiceError("failed to generate password recovery token") - errLoginDisableUser = errors.NewAuthNError("failed to login in disabled user") - errMatchUserVerification = errors.NewRequestError("user verification does not match with stored verification") - errSimilarUpdateEmail = errors.NewRequestError("new email is similar to the current email") - - usernameRegExp = regexp.MustCompile(`^[a-z0-9][a-z0-9_-]{34}[a-z0-9]$`) -) - -type service struct { - token grpcTokenV1.TokenServiceClient - users Repository - idProvider magistrala.IDProvider - policies policies.Service - hasher Hasher - email Emailer -} - -// NewService returns a new Users service implementation. -func NewService(token grpcTokenV1.TokenServiceClient, urepo Repository, policyService policies.Service, emailer Emailer, hasher Hasher, idp magistrala.IDProvider) Service { - return service{ - token: token, - users: urepo, - policies: policyService, - hasher: hasher, - email: emailer, - idProvider: idp, - } -} - -func (svc service) Register(ctx context.Context, session authn.Session, u User, selfRegister bool) (uc User, err error) { - if !selfRegister { - if err := svc.checkSuperAdmin(ctx, session); err != nil { - return User{}, err - } - } - - userID, err := svc.idProvider.ID() - if err != nil { - return User{}, errors.Wrap(svcerr.ErrIssueProviderID, err) - } - - if u.Credentials.Secret != "" { - hash, err := svc.hasher.Hash(u.Credentials.Secret) - if err != nil { - return User{}, errors.Wrap(svcerr.ErrHashPassword, err) - } - u.Credentials.Secret = hash - } - - if u.Status != DisabledStatus && u.Status != EnabledStatus { - return User{}, svcerr.ErrInvalidStatus - } - if u.Role != UserRole && u.Role != AdminRole { - return User{}, svcerr.ErrInvalidRole - } - u.ID = userID - u.CreatedAt = time.Now().UTC() - - if err := svc.addUserPolicy(ctx, u.ID, u.Role); err != nil { - return User{}, errors.Wrap(svcerr.ErrAddPolicies, err) - } - defer func() { - if err != nil { - if errRollback := svc.addUserPolicyRollback(ctx, u.ID, u.Role); errRollback != nil { - err = errors.Wrap(errors.Wrap(apiutil.ErrRollbackTx, errRollback), err) - } - } - }() - user, err := svc.users.Save(ctx, u) - if err != nil { - return User{}, errors.Wrap(svcerr.ErrCreateEntity, err) - } - return user, nil -} - -func (svc service) SendVerification(ctx context.Context, session authn.Session) error { - dbUser, err := svc.users.RetrieveByID(ctx, session.UserID) - if err != nil { - return errors.Wrap(svcerr.ErrViewEntity, err) - } - - if !dbUser.VerifiedAt.IsZero() { - return svcerr.ErrUserAlreadyVerified - } - - uv, err := svc.users.RetrieveUserVerification(ctx, dbUser.ID, dbUser.Email) - if err != nil && err != repoerr.ErrNotFound { - return errors.Wrap(svcerr.ErrCreateEntity, err) - } - - if err = uv.Valid(); err != nil { - uv, err = NewUserVerification(dbUser.ID, dbUser.Email) - if err != nil { - return errors.Wrap(svcerr.ErrCreateEntity, err) - } - if err := svc.users.AddUserVerification(ctx, uv); err != nil { - return errors.Wrap(svcerr.ErrCreateEntity, err) - } - } - - uvs, err := uv.Encode() - if err != nil { - return errors.Wrap(svcerr.ErrCreateEntity, err) - } - - if err := svc.email.SendVerification([]string{dbUser.Email}, dbUser.Credentials.Username, uvs); err != nil { - return errors.Wrap(svcerr.ErrCreateEntity, err) - } - return nil -} - -func (svc service) VerifyEmail(ctx context.Context, token string) (User, error) { - var received UserVerification - if err := received.Decode(token); err != nil { - return User{}, errors.Wrap(svcerr.ErrInvalidUserVerification, err) - } - - stored, err := svc.users.RetrieveUserVerification(ctx, received.UserID, received.Email) - if err != nil { - return User{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - - if err := stored.Match(received); err != nil { - return User{}, errors.Wrap(errMatchUserVerification, err) - } - - if err := stored.Valid(); err != nil { - if err == svcerr.ErrUserVerificationExpired { - return User{}, err - } - return User{}, errors.Wrap(svcerr.ErrMalformedEntity, err) - } - - stored.UsedAt = time.Now().UTC() - if err = svc.users.UpdateUserVerification(ctx, stored); err != nil { - return User{}, errors.Wrap(svcerr.ErrUpdateEntity, err) - } - - user := User{ - ID: stored.UserID, - Email: stored.Email, - VerifiedAt: time.Now().UTC(), - } - user, err = svc.users.UpdateVerifiedAt(ctx, user) - if err == repoerr.ErrNotFound { - return User{}, svcerr.ErrInvalidUserVerification - } - if err != nil { - return User{}, errors.Wrap(svcerr.ErrUpdateEntity, err) - } - - return user, nil -} - -func (svc service) IssueToken(ctx context.Context, identity, secret, description string) (*grpcTokenV1.Token, error) { - var dbUser User - var err error - - if _, parseErr := mail.ParseAddress(identity); parseErr != nil { - dbUser, err = svc.users.RetrieveByUsername(ctx, identity) - } else { - dbUser, err = svc.users.RetrieveByEmail(ctx, identity) - } - - if err == repoerr.ErrNotFound { - return &grpcTokenV1.Token{}, errors.Wrap(svcerr.ErrLogin, err) - } - - if err != nil { - return &grpcTokenV1.Token{}, errors.Wrap(svcerr.ErrAuthentication, err) - } - - if err := svc.hasher.Compare(secret, dbUser.Credentials.Secret); err != nil { - return &grpcTokenV1.Token{}, errors.Wrap(svcerr.ErrLogin, err) - } - - token, err := svc.token.Issue(ctx, &grpcTokenV1.IssueReq{ - UserId: dbUser.ID, - UserRole: uint32(dbUser.Role + 1), - Type: uint32(smqauth.AccessKey), - Verified: !dbUser.VerifiedAt.IsZero(), - Description: description, - }) - if err != nil { - return &grpcTokenV1.Token{}, err - } - - return token, nil -} - -func (svc service) RefreshToken(ctx context.Context, session authn.Session, refreshToken string) (*grpcTokenV1.Token, error) { - dbUser, err := svc.users.RetrieveByID(ctx, session.UserID) - if err != nil { - return &grpcTokenV1.Token{}, errors.Wrap(svcerr.ErrAuthentication, err) - } - if dbUser.Status == DisabledStatus { - return &grpcTokenV1.Token{}, errors.Wrap(svcerr.ErrAuthentication, errLoginDisableUser) - } - token, err := svc.token.Refresh(ctx, &grpcTokenV1.RefreshReq{RefreshToken: refreshToken, Verified: !dbUser.VerifiedAt.IsZero()}) - if err != nil { - return &grpcTokenV1.Token{}, err - } - - return token, nil -} - -func (svc service) RevokeRefreshToken(ctx context.Context, session authn.Session, tokenID string) error { - dbUser, err := svc.users.RetrieveByID(ctx, session.UserID) - if err != nil { - return errors.Wrap(svcerr.ErrAuthentication, err) - } - if dbUser.Status == DisabledStatus { - return errors.Wrap(svcerr.ErrAuthentication, errLoginDisableUser) - } - _, err = svc.token.Revoke(ctx, &grpcTokenV1.RevokeReq{UserId: session.UserID, TokenId: tokenID}) - if err != nil { - if errors.Contains(err, svcerr.ErrNotFound) { - return errors.Wrap(svcerr.ErrNotFound, err) - } - return errors.Wrap(svcerr.ErrRemoveEntity, err) - } - - return nil -} - -func (svc service) ListActiveRefreshTokens(ctx context.Context, session authn.Session) (*grpcTokenV1.ListUserRefreshTokensRes, error) { - dbUser, err := svc.users.RetrieveByID(ctx, session.UserID) - if err != nil { - return nil, errors.Wrap(svcerr.ErrAuthentication, err) - } - if dbUser.Status == DisabledStatus { - return nil, errors.Wrap(svcerr.ErrAuthentication, errLoginDisableUser) - } - - refreshTokens, err := svc.token.ListUserRefreshTokens(ctx, &grpcTokenV1.ListUserRefreshTokensReq{UserId: session.UserID}) - if err != nil { - return nil, errors.Wrap(svcerr.ErrAuthentication, err) - } - - return refreshTokens, nil -} - -func (svc service) View(ctx context.Context, session authn.Session, id string) (User, error) { - user, err := svc.users.RetrieveByID(ctx, id) - if err != nil { - return User{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - - if session.UserID != id { - if err := svc.checkSuperAdmin(ctx, session); err != nil { - return User{ - FirstName: user.FirstName, - LastName: user.LastName, - ID: user.ID, - Metadata: user.Metadata, - Credentials: Credentials{Username: user.Credentials.Username}, - }, nil - } - } - - user.Credentials.Secret = "" - - return user, nil -} - -func (svc service) ViewProfile(ctx context.Context, session authn.Session) (User, error) { - user, err := svc.users.RetrieveByID(ctx, session.UserID) - if err != nil { - return User{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - user.Credentials.Secret = "" - - return user, nil -} - -func (svc service) ListUsers(ctx context.Context, session authn.Session, pm Page) (UsersPage, error) { - if err := svc.checkSuperAdmin(ctx, session); err != nil { - return UsersPage{}, err - } - - pm.Role = AllRole - pg, err := svc.users.RetrieveAll(ctx, pm) - if err != nil { - return UsersPage{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - return pg, err -} - -func (svc service) SearchUsers(ctx context.Context, pm Page) (UsersPage, error) { - page := Page{ - Offset: pm.Offset, - Limit: pm.Limit, - FirstName: pm.FirstName, - LastName: pm.LastName, - Username: pm.Username, - Id: pm.Id, - Role: UserRole, - } - - cp, err := svc.users.SearchUsers(ctx, page) - if err != nil { - return UsersPage{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - - return cp, nil -} - -func (svc service) Update(ctx context.Context, session authn.Session, id string, usr UserReq) (User, error) { - if session.UserID != id { - if err := svc.checkSuperAdmin(ctx, session); err != nil { - return User{}, err - } - } - u, err := svc.users.RetrieveByID(ctx, id) - if err != nil { - return User{}, errors.Wrap(svcerr.ErrUpdateEntity, err) - } - if u.AuthProvider != "" { - if changed(usr.FirstName, u.FirstName) || - changed(usr.LastName, u.LastName) || - changed(usr.ProfilePicture, u.ProfilePicture) { - return User{}, svcerr.ErrExternalAuthProviderCouldNotUpdate - } - } - updatedAt := time.Now().UTC() - usr.UpdatedAt = &updatedAt - usr.UpdatedBy = &session.UserID - - user, err := svc.users.Update(ctx, id, usr) - if err != nil { - return User{}, errors.Wrap(svcerr.ErrUpdateEntity, err) - } - return user, nil -} - -func (svc service) UpdateTags(ctx context.Context, session authn.Session, id string, usr UserReq) (User, error) { - if session.UserID != id { - if err := svc.checkSuperAdmin(ctx, session); err != nil { - return User{}, err - } - } - - updatedAt := time.Now().UTC() - usr.UpdatedAt = &updatedAt - usr.UpdatedBy = &session.UserID - - user, err := svc.users.Update(ctx, id, usr) - if err != nil { - return User{}, errors.Wrap(svcerr.ErrUpdateEntity, err) - } - - return user, nil -} - -func (svc service) UpdateProfilePicture(ctx context.Context, session authn.Session, id string, usr UserReq) (User, error) { - if session.UserID != id { - if err := svc.checkSuperAdmin(ctx, session); err != nil { - return User{}, err - } - } - - u, err := svc.users.RetrieveByID(ctx, id) - if err != nil { - return User{}, errors.Wrap(svcerr.ErrUpdateEntity, err) - } - if u.AuthProvider != "" { - return User{}, svcerr.ErrExternalAuthProviderCouldNotUpdate - } - - updatedAt := time.Now().UTC() - usr.UpdatedAt = &updatedAt - usr.UpdatedBy = &session.UserID - - user, err := svc.users.Update(ctx, id, usr) - if err != nil { - return User{}, errors.Wrap(svcerr.ErrUpdateEntity, err) - } - - return user, nil -} - -func (svc service) UpdateEmail(ctx context.Context, session authn.Session, userID, email string) (User, error) { - if session.UserID != userID { - if err := svc.checkSuperAdmin(ctx, session); err != nil { - return User{}, err - } - } - oldUsr, err := svc.users.RetrieveByID(ctx, userID) - if err != nil { - return User{}, errors.Wrap(svcerr.ErrUpdateEntity, err) - } - if oldUsr.AuthProvider != "" { - return User{}, svcerr.ErrExternalAuthProviderCouldNotUpdate - } - if oldUsr.Email == email { - return User{}, errSimilarUpdateEmail - } - - usr := User{ - ID: userID, - Email: email, - UpdatedAt: time.Now().UTC(), - UpdatedBy: session.UserID, - VerifiedAt: time.Time{}, - } - - user, err := svc.users.UpdateEmail(ctx, usr) - if err != nil { - return User{}, errors.Wrap(svcerr.ErrUpdateEntity, err) - } - return user, nil -} - -func (svc service) SendPasswordReset(ctx context.Context, email string) error { - user, err := svc.users.RetrieveByEmail(ctx, email) - if err != nil { - return errors.Wrap(svcerr.ErrViewEntity, err) - } - - issueReq := &grpcTokenV1.IssueReq{ - UserId: user.ID, - UserRole: uint32(user.Role + 1), - Type: uint32(smqauth.RecoveryKey), - } - token, err := svc.token.Issue(ctx, issueReq) - if err != nil { - return errors.Wrap(errRecoveryToken, err) - } - - if err := svc.email.SendPasswordReset([]string{email}, user.Credentials.Username, token.AccessToken); err != nil { - return errors.NewInternalErrorWithErr(err) - } - - return nil -} - -func (svc service) ResetSecret(ctx context.Context, session authn.Session, secret string) error { - u, err := svc.users.RetrieveByID(ctx, session.UserID) - if err != nil { - return errors.Wrap(svcerr.ErrViewEntity, err) - } - - secret, err = svc.hasher.Hash(secret) - if err != nil { - return errors.Wrap(svcerr.ErrMalformedEntity, err) - } - u = User{ - ID: u.ID, - Email: u.Email, - Credentials: Credentials{ - Secret: secret, - }, - UpdatedAt: time.Now().UTC(), - UpdatedBy: session.UserID, - } - if _, err := svc.users.UpdateSecret(ctx, u); err != nil { - return errors.Wrap(svcerr.ErrAuthorization, err) - } - return nil -} - -func (svc service) UpdateSecret(ctx context.Context, session authn.Session, oldSecret, newSecret string) (User, error) { - dbUser, err := svc.users.RetrieveByID(ctx, session.UserID) - if err != nil { - return User{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - if _, err := svc.IssueToken(ctx, dbUser.Credentials.Username, oldSecret, ""); err != nil { - return User{}, err - } - newSecret, err = svc.hasher.Hash(newSecret) - if err != nil { - return User{}, errors.Wrap(svcerr.ErrMalformedEntity, err) - } - dbUser.Credentials.Secret = newSecret - dbUser.UpdatedAt = time.Now().UTC() - dbUser.UpdatedBy = session.UserID - - dbUser, err = svc.users.UpdateSecret(ctx, dbUser) - if err != nil { - return User{}, errors.Wrap(svcerr.ErrUpdateEntity, err) - } - - return dbUser, nil -} - -func (svc service) UpdateUsername(ctx context.Context, session authn.Session, id, username string) (User, error) { - if session.UserID != id { - if err := svc.checkSuperAdmin(ctx, session); err != nil { - return User{}, err - } - } - - usr := User{ - ID: id, - Credentials: Credentials{ - Username: username, - }, - UpdatedAt: time.Now().UTC(), - UpdatedBy: session.UserID, - } - updatedUser, err := svc.users.UpdateUsername(ctx, usr) - if err != nil { - return User{}, errors.Wrap(svcerr.ErrUpdateEntity, err) - } - return updatedUser, nil -} - -func (svc service) UpdateRole(ctx context.Context, session authn.Session, usr User) (User, error) { - if err := svc.checkSuperAdmin(ctx, session); err != nil { - return User{}, err - } - usr = User{ - ID: usr.ID, - Role: usr.Role, - UpdatedAt: time.Now().UTC(), - UpdatedBy: session.UserID, - } - - if err := svc.updateUserPolicy(ctx, usr.ID, usr.Role); err != nil { - return User{}, err - } - - u, err := svc.users.UpdateRole(ctx, usr) - if err != nil { - // If failed to update role in DB, then revert back to platform admin policies in spicedb - if errRollback := svc.updateUserPolicy(ctx, usr.ID, UserRole); errRollback != nil { - return User{}, errors.Wrap(errRollback, err) - } - return User{}, errors.Wrap(svcerr.ErrUpdateEntity, err) - } - return u, nil -} - -func (svc service) Enable(ctx context.Context, session authn.Session, id string) (User, error) { - u := User{ - ID: id, - UpdatedAt: time.Now().UTC(), - Status: EnabledStatus, - } - user, err := svc.changeUserStatus(ctx, session, u) - if err != nil { - return User{}, errors.Wrap(svcerr.ErrEnableUser, err) - } - - return user, nil -} - -func (svc service) Disable(ctx context.Context, session authn.Session, id string) (User, error) { - user := User{ - ID: id, - UpdatedAt: time.Now().UTC(), - Status: DisabledStatus, - } - user, err := svc.changeUserStatus(ctx, session, user) - if err != nil { - return User{}, errors.Wrap(svcerr.ErrDisableUser, err) - } - - return user, nil -} - -func (svc service) changeUserStatus(ctx context.Context, session authn.Session, user User) (User, error) { - if session.UserID != user.ID { - if err := svc.checkSuperAdmin(ctx, session); err != nil { - return User{}, err - } - } - dbu, err := svc.users.RetrieveByID(ctx, user.ID) - if err != nil { - return User{}, errors.Wrap(svcerr.ErrViewEntity, err) - } - if dbu.Status == user.Status { - return User{}, svcerr.ErrStatusAlreadyAssigned - } - user.UpdatedBy = session.UserID - - user, err = svc.users.ChangeStatus(ctx, user) - if err != nil { - return User{}, errors.Wrap(svcerr.ErrUpdateEntity, err) - } - return user, nil -} - -func (svc service) Delete(ctx context.Context, session authn.Session, id string) error { - user := User{ - ID: id, - UpdatedAt: time.Now().UTC(), - Status: DeletedStatus, - } - - if _, err := svc.changeUserStatus(ctx, session, user); err != nil { - return err - } - - return nil -} - -func (svc *service) checkSuperAdmin(ctx context.Context, session authn.Session) error { - if !session.SuperAdmin { - if err := svc.users.CheckSuperAdmin(ctx, session.UserID); err != nil { - return errors.Wrap(svcerr.ErrAuthorization, err) - } - } - - return nil -} - -func (svc service) OAuthCallback(ctx context.Context, user User) (User, error) { - u, err := svc.users.RetrieveByEmail(ctx, user.Email) - - if errors.Contains(err, repoerr.ErrNotFound) { - user.Credentials.Username = generateUsername(user.Email) - u, err = svc.Register(ctx, authn.Session{}, user, true) - if err != nil { - if errors.Contains(err, errors.ErrUsernameNotAvailable) { - return User{}, errors.ErrTryAgain - } - return User{}, err - } - } - - if err != nil && !errors.Contains(err, repoerr.ErrNotFound) { - return User{}, err - } - - if u.VerifiedAt.IsZero() { - user.ID = u.ID - user.VerifiedAt = time.Now() - u, err = svc.users.UpdateVerifiedAt(ctx, user) - if err != nil { - return User{}, err - } - } - - return User{ID: u.ID, Role: u.Role, VerifiedAt: u.VerifiedAt}, nil -} - -func (svc service) OAuthAddUserPolicy(ctx context.Context, user User) error { - return svc.addUserPolicy(ctx, user.ID, user.Role) -} - -func (svc service) Identify(ctx context.Context, session authn.Session) (string, error) { - return session.UserID, nil -} - -func (svc service) addUserPolicy(ctx context.Context, userID string, role Role) error { - policyList := []policies.Policy{} - - policyList = append(policyList, policies.Policy{ - SubjectType: policies.UserType, - Subject: userID, - Relation: policies.MemberRelation, - ObjectType: policies.PlatformType, - Object: policies.MagistralaObject, - }) - - if role == AdminRole { - policyList = append(policyList, policies.Policy{ - SubjectType: policies.UserType, - Subject: userID, - Relation: policies.AdministratorRelation, - ObjectType: policies.PlatformType, - Object: policies.MagistralaObject, - }) - } - err := svc.policies.AddPolicies(ctx, policyList) - if err != nil { - return errors.Wrap(svcerr.ErrAddPolicies, err) - } - - return nil -} - -func (svc service) addUserPolicyRollback(ctx context.Context, userID string, role Role) error { - policyList := []policies.Policy{} - - policyList = append(policyList, policies.Policy{ - SubjectType: policies.UserType, - Subject: userID, - Relation: policies.MemberRelation, - ObjectType: policies.PlatformType, - Object: policies.MagistralaObject, - }) - - if role == AdminRole { - policyList = append(policyList, policies.Policy{ - SubjectType: policies.UserType, - Subject: userID, - Relation: policies.AdministratorRelation, - ObjectType: policies.PlatformType, - Object: policies.MagistralaObject, - }) - } - err := svc.policies.DeletePolicies(ctx, policyList) - if err != nil { - return errors.Wrap(svcerr.ErrDeletePolicies, err) - } - - return nil -} - -func (svc service) updateUserPolicy(ctx context.Context, userID string, role Role) error { - switch role { - case AdminRole: - err := svc.policies.AddPolicy(ctx, policies.Policy{ - SubjectType: policies.UserType, - Subject: userID, - Relation: policies.AdministratorRelation, - ObjectType: policies.PlatformType, - Object: policies.MagistralaObject, - }) - if err != nil { - return errors.Wrap(svcerr.ErrAddPolicies, err) - } - - return nil - case UserRole: - fallthrough - default: - err := svc.policies.DeletePolicyFilter(ctx, policies.Policy{ - SubjectType: policies.UserType, - Subject: userID, - Relation: policies.AdministratorRelation, - ObjectType: policies.PlatformType, - Object: policies.MagistralaObject, - }) - if err != nil { - return errors.Wrap(svcerr.ErrDeletePolicies, err) - } - - return nil - } -} - -func generateUsername(email string) string { - uniqueSuffix := generateRandomID() - emailPrefix := extractEmailPrefix(email) - return fmt.Sprintf("%s_%s", emailPrefix, uniqueSuffix) -} - -func extractEmailPrefix(email string) string { - parts := strings.Split(email, "@") - if len(parts) == 0 { - return defaultUsernamePrefix - } - - prefix := parts[0] - cleaned := usernameRegExp.ReplaceAllString(prefix, "") - - cleaned = sanitizeForUsername(cleaned, 15) - if cleaned == "" { - cleaned = defaultUsernamePrefix - } - - return cleaned -} - -func generateRandomID() string { - // Generate 8 random bytes (will result in 16 hex chars, truncated to 10) - randomBytes := make([]byte, 8) - if _, err := rand.Read(randomBytes); err != nil { - // Fallback: use UUID if crypto/rand fails (should never happen) - id, uuidErr := uuid.NewV4() - if uuidErr != nil { - // Last resort fallback - return fmt.Sprintf("%x", time.Now().UnixNano())[:10] - } - return hex.EncodeToString(id.Bytes())[:10] - } - return hex.EncodeToString(randomBytes)[:10] -} - -// sanitizeForUsername extracts and cleans a string for use in username generation. -// As per the username requirements: -// - It keeps only lowercase alphanumeric characters, hyphens, and underscores -// - ensures valid boundaries (no hyphens/underscores at start/end) -// - removes consecutive hyphens/underscores (to pass validation) -// and finally limits the result to maxLen characters. -func sanitizeForUsername(s string, maxLen int) string { - if s == "" { - return "" - } - - // Convert to lowercase - s = strings.ToLower(s) - - // Filter characters - keep only alphanumeric, hyphen, underscore - buf := make([]byte, 0, len(s)) - var lastChar byte - - for i := 0; i < len(s); i++ { - c := s[i] - - // Keep alphanumeric, hyphen, underscore - if (c >= 'a' && c <= 'z') || (c >= '0' && c <= '9') || c == '-' || c == '_' { - // Skip if current char is hyphen/underscore and same as last char - if (c == '-' || c == '_') && c == lastChar { - continue // Skip consecutive hyphens or underscores - } - buf = append(buf, c) - lastChar = c - } else { - lastChar = 0 // Reset on special char - } - } - - cleaned := string(buf) - - // Trim invalid boundary characters - cleaned = strings.Trim(cleaned, "-_") - - // Limit length - if len(cleaned) > maxLen { - cleaned = cleaned[:maxLen] - // Re-trim in case truncation exposed hyphen/underscore at end - cleaned = strings.TrimRight(cleaned, "-_") - } - - return cleaned -} - -func changed(updated *string, old string) bool { - if updated == nil { - return false - } - - return *updated != old -} diff --git a/users/service_test.go b/users/service_test.go deleted file mode 100644 index 2fbf11425..000000000 --- a/users/service_test.go +++ /dev/null @@ -1,2330 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package users_test - -import ( - "context" - "fmt" - "strings" - "testing" - "time" - - grpcTokenV1 "github.com/absmach/magistrala/api/grpc/token/v1" - smqauth "github.com/absmach/magistrala/auth" - authmocks "github.com/absmach/magistrala/auth/mocks" - "github.com/absmach/magistrala/internal/testsutil" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/errors" - repoerr "github.com/absmach/magistrala/pkg/errors/repository" - svcerr "github.com/absmach/magistrala/pkg/errors/service" - policymocks "github.com/absmach/magistrala/pkg/policies/mocks" - "github.com/absmach/magistrala/pkg/uuid" - "github.com/absmach/magistrala/users" - "github.com/absmach/magistrala/users/hasher" - "github.com/absmach/magistrala/users/mocks" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" -) - -var ( - idProvider = uuid.New() - phasher = hasher.New() - secret = "strongsecret" - validCMetadata = users.Metadata{"role": "user"} - userID = "d8dd12ef-aa2a-43fe-8ef2-2e4fe514360f" - user = users.User{ - ID: userID, - FirstName: "firstname", - LastName: "lastname", - Tags: []string{"tag1", "tag2"}, - Credentials: users.Credentials{Username: "username", Secret: secret}, - Email: "useremail@email.com", - Metadata: validCMetadata, - PrivateMetadata: validCMetadata, - Status: users.EnabledStatus, - } - basicUser = users.User{ - Credentials: users.Credentials{ - Username: "username", - }, - ID: userID, - FirstName: "firstname", - LastName: "lastname", - } - validToken = "token" - validID = "d4ebb847-5d0e-4e46-bdd9-b6aceaaa3a22" - wrongID = testsutil.GenerateUUID(&testing.T{}) - errHashPassword = errors.New("generate hash from password failed") -) - -func newService() (users.Service, *authmocks.TokenServiceClient, *mocks.Repository, *policymocks.Service, *mocks.Emailer) { - cRepo := new(mocks.Repository) - policies := new(policymocks.Service) - e := new(mocks.Emailer) - tokenClient := new(authmocks.TokenServiceClient) - return users.NewService(tokenClient, cRepo, policies, e, phasher, idProvider), tokenClient, cRepo, policies, e -} - -func newServiceMinimal() (users.Service, *mocks.Repository) { - cRepo := new(mocks.Repository) - policies := new(policymocks.Service) - e := new(mocks.Emailer) - tokenUser := new(authmocks.TokenServiceClient) - return users.NewService(tokenUser, cRepo, policies, e, phasher, idProvider), cRepo -} - -func TestRegister(t *testing.T) { - svc, _, cRepo, policies, _ := newService() - - cases := []struct { - desc string - user users.User - addPoliciesResponseErr error - deletePoliciesResponseErr error - saveErr error - err error - }{ - { - desc: "register new user successfully", - user: user, - err: nil, - }, - { - desc: "register existing user", - user: user, - saveErr: repoerr.ErrConflict, - err: repoerr.ErrConflict, - }, - { - desc: "register a new enabled user with name", - user: users.User{ - FirstName: "userWithName", - Email: "newuserwithname@example.com", - Credentials: users.Credentials{ - Secret: secret, - }, - Status: users.EnabledStatus, - }, - err: nil, - }, - { - desc: "register a new disabled user with name", - user: users.User{ - FirstName: "userWithName", - Email: "newuserwithname@example.com", - Credentials: users.Credentials{ - Secret: secret, - }, - }, - err: nil, - }, - { - desc: "register a new user with all fields", - user: users.User{ - FirstName: "newuserwithallfields", - Tags: []string{"tag1", "tag2"}, - Email: "newuserwithallfields@example.com", - Credentials: users.Credentials{ - Secret: secret, - }, - PrivateMetadata: users.Metadata{ - "name": "newuserwithallfields", - }, - Metadata: users.Metadata{ - "name": "newuserwithallfields", - }, - Status: users.EnabledStatus, - }, - err: nil, - }, - { - desc: "register a new user with missing email", - user: users.User{ - FirstName: "userWithMissingEmail", - Credentials: users.Credentials{ - Secret: secret, - }, - }, - saveErr: errors.ErrMalformedEntity, - err: errors.ErrMalformedEntity, - }, - { - desc: "register a new user with missing secret", - user: users.User{ - FirstName: "userWithMissingSecret", - Email: "userwithmissingsecret@example.com", - Credentials: users.Credentials{ - Secret: "", - }, - }, - err: nil, - }, - { - desc: " register a user with a secret that is too long", - user: users.User{ - FirstName: "userWithLongSecret", - Email: "userwithlongsecret@example.com", - Credentials: users.Credentials{ - Secret: strings.Repeat("a", 73), - }, - }, - err: errHashPassword, - }, - { - desc: "register a new user with invalid status", - user: users.User{ - FirstName: "userWithInvalidStatus", - Email: "user with invalid status", - Credentials: users.Credentials{ - Secret: secret, - }, - Status: users.AllStatus, - }, - err: svcerr.ErrInvalidStatus, - }, - { - desc: "register a new user with invalid role", - user: users.User{ - FirstName: "userWithInvalidRole", - Email: "userwithinvalidrole@example.com", - Credentials: users.Credentials{ - Secret: secret, - }, - Role: 2, - }, - err: svcerr.ErrInvalidRole, - }, - { - desc: "register a new user with failed to add policies with err", - user: users.User{ - FirstName: "userWithFailedToAddPolicies", - Email: "userwithfailedpolicies@example.com", - Credentials: users.Credentials{ - Secret: secret, - }, - Role: users.AdminRole, - }, - addPoliciesResponseErr: svcerr.ErrAddPolicies, - err: svcerr.ErrAddPolicies, - }, - { - desc: "register a new user with failed to delete policies with err", - user: users.User{ - FirstName: "userWithFailedToDeletePolicies", - Email: "userwithfailedtodelete@example.com", - Credentials: users.Credentials{ - Secret: secret, - }, - Role: users.AdminRole, - }, - deletePoliciesResponseErr: svcerr.ErrConflict, - saveErr: repoerr.ErrConflict, - err: svcerr.ErrConflict, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - policyCall := policies.On("AddPolicies", context.Background(), mock.Anything).Return(tc.addPoliciesResponseErr) - policyCall1 := policies.On("DeletePolicies", context.Background(), mock.Anything).Return(tc.deletePoliciesResponseErr) - repoCall := cRepo.On("Save", context.Background(), mock.Anything).Return(tc.user, tc.saveErr) - expected, err := svc.Register(context.Background(), authn.Session{}, tc.user, true) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - if err == nil { - tc.user.ID = expected.ID - tc.user.CreatedAt = expected.CreatedAt - tc.user.UpdatedAt = expected.UpdatedAt - tc.user.Credentials.Secret = expected.Credentials.Secret - tc.user.UpdatedBy = expected.UpdatedBy - assert.Equal(t, tc.user, expected, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.user, expected)) - ok := repoCall.Parent.AssertCalled(t, "Save", context.Background(), mock.Anything) - assert.True(t, ok, fmt.Sprintf("Save was not called on %s", tc.desc)) - } - repoCall.Unset() - policyCall.Unset() - policyCall1.Unset() - }) - } - - svc, _, cRepo, policies, _ = newService() - - cases2 := []struct { - desc string - user users.User - session authn.Session - addPoliciesResponseErr error - deletePoliciesResponseErr error - saveErr error - checkSuperAdminErr error - err error - }{ - { - desc: "register new user successfully as admin", - user: user, - session: authn.Session{UserID: validID, SuperAdmin: true}, - err: nil, - }, - { - desc: "register a new user as admin with failed check on super admin", - user: user, - session: authn.Session{UserID: validID, SuperAdmin: false}, - checkSuperAdminErr: svcerr.ErrAuthorization, - err: svcerr.ErrAuthorization, - }, - } - for _, tc := range cases2 { - repoCall := cRepo.On("CheckSuperAdmin", context.Background(), mock.Anything).Return(tc.checkSuperAdminErr) - policyCall := policies.On("AddPolicies", context.Background(), mock.Anything).Return(tc.addPoliciesResponseErr) - policyCall1 := policies.On("DeletePolicies", context.Background(), mock.Anything).Return(tc.deletePoliciesResponseErr) - repoCall1 := cRepo.On("Save", context.Background(), mock.Anything).Return(tc.user, tc.saveErr) - expected, err := svc.Register(context.Background(), authn.Session{UserID: validID}, tc.user, false) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - if err == nil { - tc.user.ID = expected.ID - tc.user.CreatedAt = expected.CreatedAt - tc.user.UpdatedAt = expected.UpdatedAt - tc.user.Credentials.Secret = expected.Credentials.Secret - tc.user.UpdatedBy = expected.UpdatedBy - assert.Equal(t, tc.user, expected, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.user, expected)) - ok := repoCall1.Parent.AssertCalled(t, "Save", context.Background(), mock.Anything) - assert.True(t, ok, fmt.Sprintf("Save was not called on %s", tc.desc)) - } - repoCall1.Unset() - policyCall.Unset() - policyCall1.Unset() - repoCall.Unset() - } -} - -func TestViewUser(t *testing.T) { - svc, cRepo := newServiceMinimal() - - cases := []struct { - desc string - token string - reqUserID string - userID string - retrieveByIDResponse users.User - response users.User - identifyErr error - authorizeErr error - retrieveByIDErr error - checkSuperAdminErr error - err error - }{ - { - desc: "view user as normal user successfully", - retrieveByIDResponse: user, - response: user, - token: validToken, - reqUserID: user.ID, - userID: user.ID, - err: nil, - checkSuperAdminErr: svcerr.ErrAuthorization, - }, - { - desc: "view user as normal user with failed to retrieve user", - retrieveByIDResponse: users.User{}, - token: validToken, - reqUserID: user.ID, - userID: user.ID, - retrieveByIDErr: repoerr.ErrNotFound, - err: svcerr.ErrNotFound, - checkSuperAdminErr: svcerr.ErrAuthorization, - }, - { - desc: "view user as admin user successfully", - retrieveByIDResponse: user, - response: user, - token: validToken, - reqUserID: user.ID, - userID: user.ID, - err: nil, - }, - { - desc: "view user as admin user with failed check on super admin", - token: validToken, - retrieveByIDResponse: basicUser, - response: basicUser, - reqUserID: user.ID, - userID: "", - checkSuperAdminErr: svcerr.ErrAuthorization, - err: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := cRepo.On("CheckSuperAdmin", context.Background(), mock.Anything).Return(tc.checkSuperAdminErr) - repoCall1 := cRepo.On("RetrieveByID", context.Background(), tc.userID).Return(tc.retrieveByIDResponse, tc.retrieveByIDErr) - rUser, err := svc.View(context.Background(), authn.Session{UserID: tc.reqUserID}, tc.userID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - tc.response.Credentials.Secret = "" - assert.Equal(t, tc.response, rUser, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.response, rUser)) - if tc.err == nil { - ok := repoCall1.Parent.AssertCalled(t, "RetrieveByID", context.Background(), tc.userID) - assert.True(t, ok, fmt.Sprintf("RetrieveByID was not called on %s", tc.desc)) - } - repoCall1.Unset() - repoCall.Unset() - }) - } -} - -func TestListUsers(t *testing.T) { - svc, cRepo := newServiceMinimal() - - cases := []struct { - desc string - token string - page users.Page - retrieveAllResponse users.UsersPage - response users.UsersPage - size uint64 - retrieveAllErr error - superAdminErr error - err error - }{ - { - desc: "list clients as admin successfully", - page: users.Page{ - Total: 1, - }, - retrieveAllResponse: users.UsersPage{ - Page: users.Page{ - Total: 1, - }, - Users: []users.User{user}, - }, - response: users.UsersPage{ - Page: users.Page{ - Total: 1, - }, - Users: []users.User{user}, - }, - token: validToken, - err: nil, - }, - { - desc: "list clients as admin with failed to retrieve clients", - page: users.Page{ - Total: 1, - }, - retrieveAllResponse: users.UsersPage{}, - token: validToken, - retrieveAllErr: repoerr.ErrNotFound, - err: svcerr.ErrViewEntity, - }, - { - desc: "list clients as admin with failed check on super admin", - page: users.Page{ - Total: 1, - }, - token: validToken, - superAdminErr: svcerr.ErrAuthorization, - err: svcerr.ErrAuthorization, - }, - { - desc: "list clients as normal user with failed to retrieve clients", - page: users.Page{ - Total: 1, - }, - retrieveAllResponse: users.UsersPage{}, - token: validToken, - retrieveAllErr: repoerr.ErrNotFound, - err: svcerr.ErrViewEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := cRepo.On("CheckSuperAdmin", context.Background(), mock.Anything).Return(tc.superAdminErr) - repoCall1 := cRepo.On("RetrieveAll", context.Background(), mock.Anything).Return(tc.retrieveAllResponse, tc.retrieveAllErr) - page, err := svc.ListUsers(context.Background(), authn.Session{UserID: user.ID}, tc.page) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.response, page, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.response, page)) - if tc.err == nil { - ok := repoCall1.Parent.AssertCalled(t, "RetrieveAll", context.Background(), mock.Anything) - assert.True(t, ok, fmt.Sprintf("RetrieveAll was not called on %s", tc.desc)) - } - repoCall.Unset() - repoCall1.Unset() - }) - } -} - -func TestSearchUsers(t *testing.T) { - svc, cRepo := newServiceMinimal() - cases := []struct { - desc string - token string - page users.Page - response users.UsersPage - responseErr error - err error - }{ - { - desc: "search clients with valid token", - token: validToken, - page: users.Page{Offset: 0, FirstName: "username", Limit: 100}, - response: users.UsersPage{ - Page: users.Page{Total: 1, Offset: 0, Limit: 100}, - Users: []users.User{user}, - }, - }, - { - desc: "search clients with id", - token: validToken, - page: users.Page{Offset: 0, Id: "d8dd12ef-aa2a-43fe-8ef2-2e4fe514360f", Limit: 100}, - response: users.UsersPage{ - Page: users.Page{Total: 1, Offset: 0, Limit: 100}, - Users: []users.User{user}, - }, - }, - { - desc: "search clients with random name", - token: validToken, - page: users.Page{Offset: 0, FirstName: "randomname", Limit: 100}, - response: users.UsersPage{ - Page: users.Page{Total: 0, Offset: 0, Limit: 100}, - Users: []users.User{}, - }, - }, - { - desc: "search clients with repo failed", - token: validToken, - page: users.Page{Offset: 0, FirstName: "randomname", Limit: 100}, - response: users.UsersPage{ - Page: users.Page{Total: 0, Offset: 0, Limit: 0}, - }, - responseErr: repoerr.ErrViewEntity, - err: svcerr.ErrViewEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := cRepo.On("SearchUsers", context.Background(), mock.Anything).Return(tc.response, tc.responseErr) - page, err := svc.SearchUsers(context.Background(), tc.page) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.response, page, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.response, page)) - repoCall.Unset() - }) - } -} - -func TestUpdateUser(t *testing.T) { - svc, cRepo := newServiceMinimal() - - user1 := user - user2 := user - updateFirstName := "Updated user" - user1.FirstName = updateFirstName - updatedMetadata := users.Metadata{"role": "test"} - invalidMetadata := users.Metadata{"role": make(chan int)} - user2.PrivateMetadata = updatedMetadata - user2.Metadata = updatedMetadata - adminID := testsutil.GenerateUUID(t) - - cases := []struct { - desc string - userID string - userReq users.UserReq - session authn.Session - updateResponse users.User - retrieveByIDResp users.User - retrieveByIDErr error - token string - updateErr error - checkSuperAdminErr error - err error - }{ - { - desc: "update user name successfully as normal user", - userID: user1.ID, - userReq: users.UserReq{ - FirstName: &updateFirstName, - }, - session: authn.Session{UserID: user1.ID}, - updateResponse: user1, - retrieveByIDResp: user1, - token: validToken, - err: nil, - }, - { - desc: "update private metadata successfully as normal user", - userID: user2.ID, - userReq: users.UserReq{ - PrivateMetadata: &updatedMetadata, - }, - session: authn.Session{UserID: user2.ID}, - updateResponse: user2, - token: validToken, - err: nil, - }, - { - desc: "update private metadata with repo error", - userID: user2.ID, - userReq: users.UserReq{ - PrivateMetadata: &invalidMetadata, - }, - session: authn.Session{UserID: user2.ID}, - updateResponse: users.User{}, - token: validToken, - updateErr: errors.ErrMalformedEntity, - err: svcerr.ErrUpdateEntity, - }, - { - desc: "update metadata successfully as normal user", - userID: user2.ID, - userReq: users.UserReq{ - Metadata: &updatedMetadata, - }, - session: authn.Session{UserID: user2.ID}, - updateResponse: user2, - retrieveByIDResp: user2, - token: validToken, - err: nil, - }, - { - desc: "update metadata with repo error", - userID: user2.ID, - userReq: users.UserReq{ - Metadata: &invalidMetadata, - }, - session: authn.Session{UserID: user2.ID}, - updateResponse: users.User{}, - retrieveByIDResp: user2, - token: validToken, - updateErr: errors.ErrMalformedEntity, - err: svcerr.ErrUpdateEntity, - }, - { - desc: "update user name as normal user with repo error on update", - userID: user1.ID, - userReq: users.UserReq{ - FirstName: &updateFirstName, - }, - session: authn.Session{UserID: user1.ID}, - updateResponse: users.User{}, - retrieveByIDResp: user1, - token: validToken, - updateErr: errors.ErrMalformedEntity, - err: svcerr.ErrUpdateEntity, - }, - { - desc: "update user name as admin successfully", - userID: user1.ID, - userReq: users.UserReq{ - FirstName: &updateFirstName, - }, - session: authn.Session{UserID: adminID, SuperAdmin: true}, - updateResponse: user1, - retrieveByIDResp: user1, - token: validToken, - err: nil, - }, - { - desc: "update user private metadata as admin successfully", - userID: user2.ID, - userReq: users.UserReq{ - PrivateMetadata: &updatedMetadata, - }, - session: authn.Session{UserID: adminID, SuperAdmin: true}, - updateResponse: user2, - retrieveByIDResp: user2, - token: validToken, - err: nil, - }, - { - desc: "update user with failed check on super admin", - userID: user1.ID, - userReq: users.UserReq{ - FirstName: &updateFirstName, - }, - session: authn.Session{UserID: adminID}, - token: validToken, - checkSuperAdminErr: svcerr.ErrAuthorization, - err: svcerr.ErrAuthorization, - }, - { - desc: "update user name as admin with repo error on update", - userID: user1.ID, - userReq: users.UserReq{ - FirstName: &updateFirstName, - }, - session: authn.Session{UserID: adminID, SuperAdmin: true}, - updateResponse: users.User{}, - retrieveByIDResp: user1, - token: validToken, - updateErr: errors.ErrMalformedEntity, - err: svcerr.ErrUpdateEntity, - }, - { - desc: "update user first name with external auth provider should fail", - userID: user1.ID, - userReq: users.UserReq{ - FirstName: &updateFirstName, - }, - session: authn.Session{UserID: user1.ID}, - retrieveByIDResp: users.User{ - ID: user1.ID, - AuthProvider: "google", - }, - token: validToken, - err: svcerr.ErrExternalAuthProviderCouldNotUpdate, - }, - { - desc: "update user last name with external auth provider should fail", - userID: user1.ID, - userReq: users.UserReq{ - LastName: &updateFirstName, - }, - session: authn.Session{UserID: user1.ID}, - retrieveByIDResp: users.User{ - ID: user1.ID, - AuthProvider: "google", - }, - token: validToken, - err: svcerr.ErrExternalAuthProviderCouldNotUpdate, - }, - { - desc: "update user privatemetadata with external auth provider should succeed", - userID: user2.ID, - userReq: users.UserReq{ - PrivateMetadata: &updatedMetadata, - }, - session: authn.Session{UserID: user2.ID}, - retrieveByIDResp: users.User{ - ID: user2.ID, - AuthProvider: "google", - PrivateMetadata: updatedMetadata, - }, - updateResponse: users.User{ - ID: user2.ID, - AuthProvider: "google", - PrivateMetadata: updatedMetadata, - }, - token: validToken, - err: nil, - }, - { - desc: "update user with retrieve by id error", - userID: user1.ID, - userReq: users.UserReq{ - FirstName: &updateFirstName, - }, - session: authn.Session{UserID: user1.ID}, - retrieveByIDErr: repoerr.ErrNotFound, - token: validToken, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := cRepo.On("CheckSuperAdmin", context.Background(), mock.Anything).Return(tc.checkSuperAdminErr) - repoCall1 := cRepo.On("RetrieveByID", context.Background(), tc.userID).Return(tc.retrieveByIDResp, tc.retrieveByIDErr) - repoCall2 := cRepo.On("Update", context.Background(), tc.userID, mock.Anything).Return(tc.updateResponse, tc.updateErr) - updatedUser, err := svc.Update(context.Background(), tc.session, tc.userID, tc.userReq) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.updateResponse, updatedUser, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.updateResponse, updatedUser)) - if tc.err == nil { - ok := repoCall2.Parent.AssertCalled(t, "Update", context.Background(), tc.userID, mock.Anything) - assert.True(t, ok, fmt.Sprintf("Update was not called on %s", tc.desc)) - } - repoCall.Unset() - repoCall1.Unset() - repoCall2.Unset() - }) - } -} - -func TestUpdateTags(t *testing.T) { - svc, cRepo := newServiceMinimal() - - updateTags := []string{"tag1", "tag2"} - user.Tags = updateTags - adminID := testsutil.GenerateUUID(t) - - cases := []struct { - desc string - userID string - userReq users.UserReq - session authn.Session - updateUserTagsResponse users.User - updateUserTagsErr error - checkSuperAdminErr error - err error - }{ - { - desc: "update user tags as normal user successfully", - userID: user.ID, - userReq: users.UserReq{Tags: &updateTags}, - session: authn.Session{UserID: user.ID}, - updateUserTagsResponse: user, - err: nil, - }, - { - desc: "update user tags as normal user with repo error on update", - userID: user.ID, - userReq: users.UserReq{Tags: &updateTags}, - session: authn.Session{UserID: user.ID}, - updateUserTagsResponse: users.User{}, - updateUserTagsErr: errors.ErrMalformedEntity, - err: svcerr.ErrUpdateEntity, - }, - { - desc: "update user tags as admin successfully", - userID: user.ID, - userReq: users.UserReq{Tags: &updateTags}, - session: authn.Session{UserID: adminID, SuperAdmin: true}, - err: nil, - }, - { - desc: "update user tags as admin with failed check on super admin", - userID: user.ID, - userReq: users.UserReq{Tags: &updateTags}, - session: authn.Session{UserID: adminID}, - checkSuperAdminErr: svcerr.ErrAuthorization, - err: svcerr.ErrAuthorization, - }, - { - desc: "update user tags as admin with repo error on update", - userID: user.ID, - userReq: users.UserReq{Tags: &updateTags}, - session: authn.Session{UserID: adminID, SuperAdmin: true}, - updateUserTagsResponse: users.User{}, - updateUserTagsErr: errors.ErrMalformedEntity, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := cRepo.On("CheckSuperAdmin", context.Background(), mock.Anything).Return(tc.checkSuperAdminErr) - repoCall1 := cRepo.On("Update", context.Background(), tc.userID, mock.Anything).Return(tc.updateUserTagsResponse, tc.updateUserTagsErr) - updatedUser, err := svc.UpdateTags(context.Background(), tc.session, tc.userID, tc.userReq) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.updateUserTagsResponse, updatedUser, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.updateUserTagsResponse, updatedUser)) - - if tc.err == nil { - ok := repoCall1.Parent.AssertCalled(t, "Update", context.Background(), tc.userID, mock.Anything) - assert.True(t, ok, fmt.Sprintf("Update was not called on %s", tc.desc)) - } - repoCall.Unset() - repoCall1.Unset() - }) - } -} - -func TestUpdateRole(t *testing.T) { - svc, _, cRepo, policies, _ := newService() - - user2 := user - user.Role = users.AdminRole - user2.Role = users.UserRole - - cases := []struct { - desc string - user users.User - session authn.Session - updateRoleResponse users.User - deletePolicyErr error - addPolicyErr error - updateRoleErr error - checkSuperAdminErr error - err error - }{ - { - desc: "update user role successfully", - user: user, - session: authn.Session{UserID: validID, SuperAdmin: true}, - updateRoleResponse: user, - err: nil, - }, - { - desc: "update user role with failed check on super admin", - user: user, - session: authn.Session{UserID: validID, SuperAdmin: false}, - checkSuperAdminErr: svcerr.ErrAuthorization, - err: svcerr.ErrAuthorization, - }, - { - desc: "update user role with failed to add policies", - user: user, - session: authn.Session{UserID: validID, SuperAdmin: true}, - addPolicyErr: errors.ErrMalformedEntity, - err: svcerr.ErrAddPolicies, - }, - { - desc: "update user role to user role successfully ", - user: user2, - session: authn.Session{UserID: validID, SuperAdmin: true}, - updateRoleResponse: user2, - err: nil, - }, - { - desc: "update user role to user role with failed to delete policies", - user: user2, - session: authn.Session{UserID: validID, SuperAdmin: true}, - deletePolicyErr: svcerr.ErrAuthorization, - err: svcerr.ErrAuthorization, - }, - { - desc: "update user role to user role with failed to delete policies with error", - user: user2, - session: authn.Session{UserID: validID, SuperAdmin: true}, - deletePolicyErr: svcerr.ErrMalformedEntity, - err: svcerr.ErrDeletePolicies, - }, - { - desc: "Update user with failed repo update and roll back", - user: user, - session: authn.Session{UserID: validID, SuperAdmin: true}, - updateRoleErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "Update user with failed repo update and failedroll back", - user: user, - session: authn.Session{UserID: validID, SuperAdmin: true}, - deletePolicyErr: svcerr.ErrAuthorization, - updateRoleErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := cRepo.On("CheckSuperAdmin", context.Background(), mock.Anything).Return(tc.checkSuperAdminErr) - policyCall := policies.On("AddPolicy", context.Background(), mock.Anything).Return(tc.addPolicyErr) - policyCall1 := policies.On("DeletePolicyFilter", context.Background(), mock.Anything).Return(tc.deletePolicyErr) - repoCall1 := cRepo.On("UpdateRole", context.Background(), mock.Anything).Return(tc.updateRoleResponse, tc.updateRoleErr) - - updatedUser, err := svc.UpdateRole(context.Background(), tc.session, tc.user) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.updateRoleResponse, updatedUser, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.updateRoleResponse, updatedUser)) - if tc.err == nil { - ok := repoCall1.Parent.AssertCalled(t, "UpdateRole", context.Background(), mock.Anything, mock.Anything) - assert.True(t, ok, fmt.Sprintf("Update was not called on %s", tc.desc)) - } - repoCall.Unset() - policyCall.Unset() - policyCall1.Unset() - repoCall1.Unset() - }) - } -} - -func TestUpdateSecret(t *testing.T) { - svc, authUser, cRepo, _, _ := newService() - - newSecret := "newstrongSecret" - rUser := user - rUser.Credentials.Secret, _ = phasher.Hash(user.Credentials.Secret) - responseUser := user - responseUser.Credentials.Secret = newSecret - - cases := []struct { - desc string - oldSecret string - newSecret string - session authn.Session - retrieveByIDResponse users.User - retrieveByEmailResponse users.User - updateSecretResponse users.User - issueResponse *grpcTokenV1.Token - response users.User - retrieveByIDErr error - retrieveByEmailErr error - updateSecretErr error - issueErr error - err error - }{ - { - desc: "update user secret with valid token", - oldSecret: user.Credentials.Secret, - newSecret: newSecret, - session: authn.Session{UserID: user.ID}, - retrieveByEmailResponse: rUser, - retrieveByIDResponse: user, - updateSecretResponse: responseUser, - issueResponse: &grpcTokenV1.Token{AccessToken: validToken}, - response: responseUser, - err: nil, - }, - { - desc: "update user secret with failed to retrieve user by ID", - oldSecret: user.Credentials.Secret, - newSecret: newSecret, - session: authn.Session{UserID: user.ID}, - retrieveByIDResponse: users.User{}, - retrieveByIDErr: repoerr.ErrNotFound, - err: repoerr.ErrNotFound, - }, - { - desc: "update user secret with failed to retrieve user by email", - oldSecret: user.Credentials.Secret, - newSecret: newSecret, - session: authn.Session{UserID: user.ID}, - retrieveByIDResponse: user, - retrieveByEmailResponse: users.User{}, - retrieveByEmailErr: repoerr.ErrNotFound, - err: repoerr.ErrNotFound, - }, - { - desc: "update user secret with invalid old secret", - oldSecret: "invalid", - newSecret: newSecret, - session: authn.Session{UserID: user.ID}, - retrieveByIDResponse: user, - retrieveByEmailResponse: rUser, - err: svcerr.ErrLogin, - }, - { - desc: "update user secret with too long new secret", - oldSecret: user.Credentials.Secret, - newSecret: strings.Repeat("a", 73), - session: authn.Session{UserID: user.ID}, - retrieveByIDResponse: user, - retrieveByEmailResponse: rUser, - err: errHashPassword, - }, - { - desc: "update user secret with failed to update secret", - oldSecret: user.Credentials.Secret, - newSecret: newSecret, - session: authn.Session{UserID: user.ID}, - retrieveByIDResponse: user, - retrieveByEmailResponse: rUser, - updateSecretResponse: users.User{}, - updateSecretErr: repoerr.ErrMalformedEntity, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := cRepo.On("RetrieveByID", context.Background(), user.ID).Return(tc.retrieveByIDResponse, tc.retrieveByIDErr) - repoCall1 := cRepo.On("RetrieveByUsername", context.Background(), user.Credentials.Username).Return(tc.retrieveByEmailResponse, tc.retrieveByEmailErr) - repoCall2 := cRepo.On("UpdateSecret", context.Background(), mock.Anything).Return(tc.updateSecretResponse, tc.updateSecretErr) - authCall := authUser.On("Issue", context.Background(), mock.Anything).Return(tc.issueResponse, tc.issueErr) - updatedUser, err := svc.UpdateSecret(context.Background(), tc.session, tc.oldSecret, tc.newSecret) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.response, updatedUser, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.response, updatedUser)) - if tc.err == nil { - ok := repoCall.Parent.AssertCalled(t, "RetrieveByID", context.Background(), tc.response.ID) - assert.True(t, ok, fmt.Sprintf("RetrieveByID was not called on %s", tc.desc)) - ok = repoCall1.Parent.AssertCalled(t, "RetrieveByUsername", context.Background(), tc.response.Credentials.Username) - assert.True(t, ok, fmt.Sprintf("RetrieveByUsername was not called on %s", tc.desc)) - ok = repoCall2.Parent.AssertCalled(t, "UpdateSecret", context.Background(), mock.Anything) - assert.True(t, ok, fmt.Sprintf("UpdateSecret was not called on %s", tc.desc)) - } - repoCall.Unset() - repoCall1.Unset() - repoCall2.Unset() - authCall.Unset() - }) - } -} - -func TestUpdateEmail(t *testing.T) { - svc, cRepo := newServiceMinimal() - - user2 := user - user2.Email = "user2@example.com" - - cases := []struct { - desc string - email string - token string - reqUserID string - id string - updateEmailResponse users.User - updateEmailErr error - checkSuperAdminErr error - err error - }{ - { - desc: "update user as normal user successfully", - email: "user2-update-1@example.com", - token: validToken, - reqUserID: user.ID, - id: user.ID, - updateEmailResponse: user2, - err: nil, - }, - { - desc: "update to same email as normal user successfully", - email: "user2-update-1@example.com", - token: validToken, - reqUserID: user.ID, - id: user.ID, - updateEmailResponse: user2, - err: nil, - }, - - { - desc: "update user email as normal user with repo error on update", - email: "user2-update-2@example.com", - token: validToken, - reqUserID: user.ID, - id: user.ID, - updateEmailResponse: users.User{}, - updateEmailErr: errors.ErrMalformedEntity, - err: svcerr.ErrUpdateEntity, - }, - { - desc: "update user email as admin successfully", - email: "user2-update-3@example.com", - token: validToken, - id: user.ID, - err: nil, - }, - { - desc: "update user email as admin with repo error on update", - email: "user2-update-4@exmaple.com", - token: validToken, - reqUserID: user.ID, - id: user.ID, - updateEmailResponse: users.User{}, - updateEmailErr: errors.ErrMalformedEntity, - err: svcerr.ErrUpdateEntity, - }, - { - desc: "update user as admin user with failed check on super admin", - email: "user2-update-5@exmaple.com", - token: validToken, - reqUserID: user.ID, - id: "", - updateEmailResponse: users.User{}, - updateEmailErr: errors.ErrMalformedEntity, - checkSuperAdminErr: svcerr.ErrAuthorization, - err: svcerr.ErrAuthorization, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := cRepo.On("CheckSuperAdmin", context.Background(), mock.Anything).Return(tc.checkSuperAdminErr) - repocall2 := cRepo.On("RetrieveByID", context.Background(), mock.Anything).Return(tc.updateEmailResponse, tc.updateEmailErr) - repoCall1 := cRepo.On("UpdateEmail", context.Background(), mock.Anything).Return(tc.updateEmailResponse, tc.updateEmailErr) - updatedUser, err := svc.UpdateEmail(context.Background(), authn.Session{DomainUserID: tc.reqUserID, UserID: validID, DomainID: validID}, tc.id, tc.email) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.updateEmailResponse, updatedUser, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.updateEmailResponse, updatedUser)) - if tc.err == nil && user2.Email != tc.email { - ok := repoCall1.Parent.AssertCalled(t, "UpdateEmail", context.Background(), mock.Anything, mock.Anything) - assert.True(t, ok, fmt.Sprintf("Update was not called on %s", tc.desc)) - user2.Email = tc.email - } - repoCall.Unset() - repocall2.Unset() - repoCall1.Unset() - }) - } -} - -func TestUpdateProfilePicture(t *testing.T) { - svc, cRepo := newServiceMinimal() - - updatedPicture := "https://example.com/profile.jpg" - user.ProfilePicture = updatedPicture - adminID := testsutil.GenerateUUID(t) - - cases := []struct { - desc string - userID string - userReq users.UserReq - session authn.Session - updateProfilePicResponse users.User - retrieveByIDResp users.User - retrieveByIDErr error - updateProfilePicErr error - checkSuperAdminErr error - err error - }{ - { - desc: "update profile picture as normal user successfully", - userID: user.ID, - userReq: users.UserReq{ProfilePicture: &updatedPicture}, - session: authn.Session{UserID: user.ID}, - updateProfilePicResponse: user, - retrieveByIDResp: user, - err: nil, - }, - { - desc: "update profile picture as normal user with repo error on update", - userID: user.ID, - userReq: users.UserReq{ProfilePicture: &updatedPicture}, - session: authn.Session{UserID: user.ID}, - updateProfilePicResponse: users.User{}, - retrieveByIDResp: user, - updateProfilePicErr: errors.ErrMalformedEntity, - err: svcerr.ErrUpdateEntity, - }, - { - desc: "update profile picture as admin successfully", - userID: user.ID, - userReq: users.UserReq{ProfilePicture: &updatedPicture}, - session: authn.Session{UserID: adminID, SuperAdmin: true}, - retrieveByIDResp: user, - err: nil, - }, - { - desc: "update profile picture as admin with failed check on super admin", - userID: user.ID, - userReq: users.UserReq{ProfilePicture: &updatedPicture}, - session: authn.Session{UserID: adminID}, - checkSuperAdminErr: svcerr.ErrAuthorization, - err: svcerr.ErrAuthorization, - }, - { - desc: "update profile picture as admin with repo error on update", - userID: user.ID, - userReq: users.UserReq{ProfilePicture: &updatedPicture}, - session: authn.Session{UserID: adminID, SuperAdmin: true}, - updateProfilePicResponse: users.User{}, - retrieveByIDResp: user, - updateProfilePicErr: errors.ErrMalformedEntity, - err: svcerr.ErrUpdateEntity, - }, - { - desc: "update profile picture with external auth provider", - userID: user.ID, - userReq: users.UserReq{ProfilePicture: &updatedPicture}, - session: authn.Session{UserID: user.ID}, - retrieveByIDResp: users.User{ - ID: user.ID, - AuthProvider: "google", - }, - err: svcerr.ErrExternalAuthProviderCouldNotUpdate, - }, - { - desc: "update profile picture with retrieve by id error", - userID: user.ID, - userReq: users.UserReq{ProfilePicture: &updatedPicture}, - session: authn.Session{UserID: user.ID}, - retrieveByIDErr: repoerr.ErrNotFound, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := cRepo.On("CheckSuperAdmin", context.Background(), mock.Anything).Return(tc.checkSuperAdminErr) - repoCall1 := cRepo.On("RetrieveByID", context.Background(), tc.userID).Return(tc.retrieveByIDResp, tc.retrieveByIDErr) - repoCall2 := cRepo.On("Update", context.Background(), tc.userID, mock.Anything).Return(tc.updateProfilePicResponse, tc.updateProfilePicErr) - updatedUser, err := svc.UpdateProfilePicture(context.Background(), tc.session, tc.userID, tc.userReq) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.updateProfilePicResponse, updatedUser, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.updateProfilePicResponse, updatedUser)) - if tc.err == nil { - ok := repoCall2.Parent.AssertCalled(t, "Update", context.Background(), tc.userID, mock.Anything) - assert.True(t, ok, fmt.Sprintf("Update was not called on %s", tc.desc)) - } - repoCall.Unset() - repoCall1.Unset() - repoCall2.Unset() - }) - } -} - -func TestUpdateUsername(t *testing.T) { - svc, cRepo := newServiceMinimal() - - nuser := user - nuser.Credentials.Username = "newusername" - adminID := testsutil.GenerateUUID(t) - - cases := []struct { - desc string - user users.User - session authn.Session - updateUsernameResponse users.User - updateUsernameErr error - checkSuperAdminErr error - err error - }{ - { - desc: "update username as normal user successfully", - user: user, - session: authn.Session{UserID: user.ID}, - updateUsernameResponse: nuser, - err: nil, - }, - { - desc: "update username as normal user with repo error on update", - user: user, - session: authn.Session{UserID: user.ID}, - updateUsernameResponse: users.User{}, - updateUsernameErr: errors.ErrMalformedEntity, - err: svcerr.ErrUpdateEntity, - }, - { - desc: "update username as admin successfully", - user: user, - session: authn.Session{UserID: adminID, SuperAdmin: true}, - updateUsernameResponse: nuser, - err: nil, - }, - { - desc: "update username as admin with failed check on super admin", - user: user, - session: authn.Session{UserID: adminID}, - checkSuperAdminErr: svcerr.ErrAuthorization, - err: svcerr.ErrAuthorization, - }, - { - desc: "update username as admin with repo error on update", - user: user, - session: authn.Session{UserID: adminID, SuperAdmin: true}, - updateUsernameResponse: users.User{}, - updateUsernameErr: errors.ErrMalformedEntity, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := cRepo.On("CheckSuperAdmin", context.Background(), mock.Anything).Return(tc.checkSuperAdminErr) - repoCall1 := cRepo.On("UpdateUsername", context.Background(), mock.Anything).Return(tc.updateUsernameResponse, tc.updateUsernameErr) - updatedUser, err := svc.UpdateUsername(context.Background(), tc.session, tc.user.ID, tc.user.Credentials.Username) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - assert.Equal(t, tc.updateUsernameResponse, updatedUser, fmt.Sprintf("%s: expected %v got %v\n", tc.desc, tc.updateUsernameResponse, updatedUser)) - if tc.err == nil { - ok := repoCall1.Parent.AssertCalled(t, "UpdateUsername", context.Background(), mock.Anything) - assert.True(t, ok, fmt.Sprintf("UpdateUsername was not called on %s", tc.desc)) - } - repoCall.Unset() - repoCall1.Unset() - }) - } -} - -func TestEnableUser(t *testing.T) { - svc, cRepo := newServiceMinimal() - - enabledUser1 := users.User{ID: testsutil.GenerateUUID(t), Credentials: users.Credentials{Username: "user1@example.com", Secret: "password"}, Status: users.EnabledStatus} - disabledUser1 := users.User{ID: testsutil.GenerateUUID(t), Credentials: users.Credentials{Username: "user3@example.com", Secret: "password"}, Status: users.DisabledStatus} - endisabledUser1 := disabledUser1 - endisabledUser1.Status = users.EnabledStatus - - cases := []struct { - desc string - id string - user users.User - retrieveByIDResponse users.User - changeStatusResponse users.User - response users.User - retrieveByIDErr error - changeStatusErr error - checkSuperAdminErr error - err error - }{ - { - desc: "enable disabled user", - id: disabledUser1.ID, - user: disabledUser1, - retrieveByIDResponse: disabledUser1, - changeStatusResponse: endisabledUser1, - response: endisabledUser1, - err: nil, - }, - { - desc: "enable disabled user with normal user token", - id: disabledUser1.ID, - user: disabledUser1, - checkSuperAdminErr: svcerr.ErrAuthorization, - err: svcerr.ErrAuthorization, - }, - { - desc: "enable disabled user with failed to retrieve user by ID", - id: disabledUser1.ID, - user: disabledUser1, - retrieveByIDResponse: users.User{}, - retrieveByIDErr: repoerr.ErrNotFound, - err: repoerr.ErrNotFound, - }, - { - desc: "enable already enabled user", - id: enabledUser1.ID, - user: enabledUser1, - retrieveByIDResponse: enabledUser1, - err: svcerr.ErrStatusAlreadyAssigned, - }, - { - desc: "enable disabled user with failed to change status", - id: disabledUser1.ID, - user: disabledUser1, - retrieveByIDResponse: disabledUser1, - changeStatusResponse: users.User{}, - changeStatusErr: repoerr.ErrMalformedEntity, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := cRepo.On("CheckSuperAdmin", context.Background(), mock.Anything).Return(tc.checkSuperAdminErr) - repoCall1 := cRepo.On("RetrieveByID", context.Background(), tc.id).Return(tc.retrieveByIDResponse, tc.retrieveByIDErr) - repoCall2 := cRepo.On("ChangeStatus", context.Background(), mock.Anything).Return(tc.changeStatusResponse, tc.changeStatusErr) - - _, err := svc.Enable(context.Background(), authn.Session{}, tc.id) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - if tc.err == nil { - ok := repoCall1.Parent.AssertCalled(t, "RetrieveByID", context.Background(), tc.id) - assert.True(t, ok, fmt.Sprintf("RetrieveByID was not called on %s", tc.desc)) - ok = repoCall2.Parent.AssertCalled(t, "ChangeStatus", context.Background(), mock.Anything) - assert.True(t, ok, fmt.Sprintf("ChangeStatus was not called on %s", tc.desc)) - } - repoCall.Unset() - repoCall1.Unset() - repoCall2.Unset() - }) - } -} - -func TestDisableUser(t *testing.T) { - svc, cRepo := newServiceMinimal() - - enabledUser1 := users.User{ID: testsutil.GenerateUUID(t), Credentials: users.Credentials{Username: "user1@example.com", Secret: "password"}, Status: users.EnabledStatus} - disabledUser1 := users.User{ID: testsutil.GenerateUUID(t), Credentials: users.Credentials{Username: "user3@example.com", Secret: "password"}, Status: users.DisabledStatus} - disenabledUser1 := enabledUser1 - disenabledUser1.Status = users.DisabledStatus - - cases := []struct { - desc string - id string - user users.User - retrieveByIDResponse users.User - changeStatusResponse users.User - response users.User - retrieveByIDErr error - changeStatusErr error - checkSuperAdminErr error - err error - }{ - { - desc: "disable enabled user", - id: enabledUser1.ID, - user: enabledUser1, - retrieveByIDResponse: enabledUser1, - changeStatusResponse: disenabledUser1, - response: disenabledUser1, - err: nil, - }, - { - desc: "disable enabled user with normal user token", - id: enabledUser1.ID, - user: enabledUser1, - checkSuperAdminErr: svcerr.ErrAuthorization, - err: svcerr.ErrAuthorization, - }, - { - desc: "disable enabled user with failed to retrieve user by ID", - id: enabledUser1.ID, - user: enabledUser1, - retrieveByIDResponse: users.User{}, - retrieveByIDErr: repoerr.ErrNotFound, - err: repoerr.ErrNotFound, - }, - { - desc: "disable already disabled user", - id: disabledUser1.ID, - user: disabledUser1, - retrieveByIDResponse: disabledUser1, - err: svcerr.ErrStatusAlreadyAssigned, - }, - { - desc: "disable enabled user with failed to change status", - id: enabledUser1.ID, - user: enabledUser1, - changeStatusResponse: users.User{}, - changeStatusErr: repoerr.ErrMalformedEntity, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := cRepo.On("CheckSuperAdmin", context.Background(), mock.Anything).Return(tc.checkSuperAdminErr) - repoCall1 := cRepo.On("RetrieveByID", context.Background(), tc.id).Return(tc.retrieveByIDResponse, tc.retrieveByIDErr) - repoCall2 := cRepo.On("ChangeStatus", context.Background(), mock.Anything).Return(tc.changeStatusResponse, tc.changeStatusErr) - - _, err := svc.Disable(context.Background(), authn.Session{}, tc.id) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - if tc.err == nil { - ok := repoCall1.Parent.AssertCalled(t, "RetrieveByID", context.Background(), tc.id) - assert.True(t, ok, fmt.Sprintf("RetrieveByID was not called on %s", tc.desc)) - ok = repoCall2.Parent.AssertCalled(t, "ChangeStatus", context.Background(), mock.Anything) - assert.True(t, ok, fmt.Sprintf("ChangeStatus was not called on %s", tc.desc)) - } - repoCall.Unset() - repoCall1.Unset() - repoCall2.Unset() - }) - } -} - -func TestDeleteUser(t *testing.T) { - svc, cRepo := newServiceMinimal() - - enabledUser1 := users.User{ID: testsutil.GenerateUUID(t), Credentials: users.Credentials{Username: "user1@example.com", Secret: "password"}, Status: users.EnabledStatus} - deletedUser1 := users.User{ID: testsutil.GenerateUUID(t), Credentials: users.Credentials{Username: "user3@example.com", Secret: "password"}, Status: users.DeletedStatus} - disenabledUser1 := enabledUser1 - disenabledUser1.Status = users.DeletedStatus - - cases := []struct { - desc string - id string - session authn.Session - user users.User - retrieveByIDResponse users.User - changeStatusResponse users.User - response users.User - retrieveByIDErr error - changeStatusErr error - checkSuperAdminErr error - err error - }{ - { - desc: "delete enabled user", - id: enabledUser1.ID, - user: enabledUser1, - session: authn.Session{UserID: validID, SuperAdmin: true}, - retrieveByIDResponse: enabledUser1, - changeStatusResponse: disenabledUser1, - response: disenabledUser1, - err: nil, - }, - { - desc: "delete enabled user with failed to retrieve user by ID", - id: enabledUser1.ID, - user: enabledUser1, - session: authn.Session{UserID: validID, SuperAdmin: true}, - retrieveByIDResponse: users.User{}, - retrieveByIDErr: repoerr.ErrNotFound, - err: repoerr.ErrNotFound, - }, - { - desc: "delete already deleted user", - id: deletedUser1.ID, - user: deletedUser1, - session: authn.Session{UserID: validID, SuperAdmin: true}, - retrieveByIDResponse: deletedUser1, - err: svcerr.ErrStatusAlreadyAssigned, - }, - { - desc: "delete enabled user with failed to change status", - id: enabledUser1.ID, - user: enabledUser1, - session: authn.Session{UserID: validID, SuperAdmin: true}, - retrieveByIDResponse: enabledUser1, - changeStatusResponse: users.User{}, - changeStatusErr: repoerr.ErrMalformedEntity, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall2 := cRepo.On("CheckSuperAdmin", context.Background(), mock.Anything).Return(tc.checkSuperAdminErr) - repoCall3 := cRepo.On("RetrieveByID", context.Background(), tc.id).Return(tc.retrieveByIDResponse, tc.retrieveByIDErr) - repoCall4 := cRepo.On("ChangeStatus", context.Background(), mock.Anything).Return(tc.changeStatusResponse, tc.changeStatusErr) - err := svc.Delete(context.Background(), tc.session, tc.id) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - if tc.err == nil { - ok := repoCall3.Parent.AssertCalled(t, "RetrieveByID", context.Background(), tc.id) - assert.True(t, ok, fmt.Sprintf("RetrieveByID was not called on %s", tc.desc)) - ok = repoCall4.Parent.AssertCalled(t, "ChangeStatus", context.Background(), mock.Anything) - assert.True(t, ok, fmt.Sprintf("ChangeStatus was not called on %s", tc.desc)) - } - repoCall2.Unset() - repoCall3.Unset() - repoCall4.Unset() - }) - } -} - -func TestIssueToken(t *testing.T) { - svc, auth, cRepo, _, _ := newService() - - rUser := user - rUser2 := user - rUser3 := user - rUser.Credentials.Secret, _ = phasher.Hash(user.Credentials.Secret) - rUser2.Credentials.Secret = "wrongsecret" - rUser3.Credentials.Secret, _ = phasher.Hash("wrongsecret") - - cases := []struct { - desc string - user users.User - retrieveByUsernameResponse users.User - issueResponse *grpcTokenV1.Token - retrieveByUsernameErr error - issueErr error - err error - }{ - { - desc: "issue token for an existing user", - user: user, - retrieveByUsernameResponse: rUser, - issueResponse: &grpcTokenV1.Token{AccessToken: validToken, RefreshToken: &validToken, AccessType: "3"}, - err: nil, - }, - { - desc: "issue token for non-empty domain id", - user: user, - retrieveByUsernameResponse: rUser, - issueResponse: &grpcTokenV1.Token{AccessToken: validToken, RefreshToken: &validToken, AccessType: "3"}, - err: nil, - }, - { - desc: "issue token for a non-existing user", - user: user, - retrieveByUsernameResponse: users.User{}, - retrieveByUsernameErr: repoerr.ErrNotFound, - err: repoerr.ErrNotFound, - }, - { - desc: "issue token for a user with wrong secret", - user: user, - retrieveByUsernameResponse: rUser3, - err: svcerr.ErrLogin, - }, - { - desc: "issue token with empty domain id", - user: user, - retrieveByUsernameResponse: rUser, - issueResponse: &grpcTokenV1.Token{}, - issueErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - { - desc: "issue token with grpc error", - user: user, - retrieveByUsernameResponse: rUser, - issueResponse: &grpcTokenV1.Token{}, - issueErr: svcerr.ErrAuthentication, - err: svcerr.ErrAuthentication, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := cRepo.On("RetrieveByUsername", context.Background(), tc.user.Credentials.Username).Return(tc.retrieveByUsernameResponse, tc.retrieveByUsernameErr) - authCall := auth.On("Issue", context.Background(), &grpcTokenV1.IssueReq{UserId: tc.user.ID, UserRole: uint32(tc.user.Role + 1), Type: uint32(smqauth.AccessKey)}).Return(tc.issueResponse, tc.issueErr) - token, err := svc.IssueToken(context.Background(), tc.user.Credentials.Username, tc.user.Credentials.Secret, "") - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - if err == nil { - assert.NotEmpty(t, token.GetAccessToken(), fmt.Sprintf("%s: expected %s not to be empty\n", tc.desc, token.GetAccessToken())) - assert.NotEmpty(t, token.GetRefreshToken(), fmt.Sprintf("%s: expected %s not to be empty\n", tc.desc, token.GetRefreshToken())) - ok := repoCall.Parent.AssertCalled(t, "RetrieveByUsername", context.Background(), tc.user.Credentials.Username) - assert.True(t, ok, fmt.Sprintf("RetrieveByUsername was not called on %s", tc.desc)) - ok = authCall.Parent.AssertCalled(t, "Issue", context.Background(), &grpcTokenV1.IssueReq{UserId: tc.user.ID, UserRole: uint32(tc.user.Role + 1), Type: uint32(smqauth.AccessKey)}) - assert.True(t, ok, fmt.Sprintf("Issue was not called on %s", tc.desc)) - } - authCall.Unset() - repoCall.Unset() - }) - } -} - -func TestRefreshToken(t *testing.T) { - svc, authsvc, crepo, _, _ := newService() - - rUser := user - rUser.Credentials.Secret, _ = phasher.Hash(user.Credentials.Secret) - - cases := []struct { - desc string - session authn.Session - refreshResp *grpcTokenV1.Token - refresErr error - repoResp users.User - repoErr error - err error - }{ - { - desc: "refresh token with refresh token for an existing user", - session: authn.Session{DomainUserID: validID, UserID: validID, DomainID: validID}, - refreshResp: &grpcTokenV1.Token{AccessToken: validToken, RefreshToken: &validToken, AccessType: "3"}, - repoResp: rUser, - err: nil, - }, - { - desc: "refresh token with refresh token for empty domain id", - session: authn.Session{UserID: validID}, - refreshResp: &grpcTokenV1.Token{AccessToken: validToken, RefreshToken: &validToken, AccessType: "3"}, - repoResp: rUser, - err: nil, - }, - { - desc: "refresh token with access token for an existing user", - session: authn.Session{DomainUserID: validID, UserID: validID, DomainID: validID}, - refreshResp: &grpcTokenV1.Token{}, - refresErr: svcerr.ErrAuthentication, - repoResp: rUser, - err: svcerr.ErrAuthentication, - }, - { - desc: "refresh token with refresh token for a non-existing client", - session: authn.Session{DomainUserID: validID, UserID: validID, DomainID: validID}, - repoErr: repoerr.ErrNotFound, - err: repoerr.ErrNotFound, - }, - { - desc: "refresh token with refresh token for a disable user", - session: authn.Session{DomainUserID: validID, UserID: validID, DomainID: validID}, - repoResp: users.User{Status: users.DisabledStatus}, - err: svcerr.ErrAuthentication, - }, - { - desc: "refresh token with empty domain id", - session: authn.Session{DomainUserID: validID, UserID: validID, DomainID: validID}, - refreshResp: &grpcTokenV1.Token{}, - refresErr: svcerr.ErrAuthentication, - repoResp: rUser, - err: svcerr.ErrAuthentication, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - authCall := authsvc.On("Refresh", context.Background(), &grpcTokenV1.RefreshReq{RefreshToken: validToken}).Return(tc.refreshResp, tc.refresErr) - repoCall := crepo.On("RetrieveByID", context.Background(), tc.session.UserID).Return(tc.repoResp, tc.repoErr) - token, err := svc.RefreshToken(context.Background(), tc.session, validToken) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - if err == nil { - assert.NotEmpty(t, token.GetAccessToken(), fmt.Sprintf("%s: expected %s not to be empty\n", tc.desc, token.GetAccessToken())) - assert.NotEmpty(t, token.GetRefreshToken(), fmt.Sprintf("%s: expected %s not to be empty\n", tc.desc, token.GetRefreshToken())) - ok := authCall.Parent.AssertCalled(t, "Refresh", context.Background(), &grpcTokenV1.RefreshReq{RefreshToken: validToken}) - assert.True(t, ok, fmt.Sprintf("Refresh was not called on %s", tc.desc)) - ok = repoCall.Parent.AssertCalled(t, "RetrieveByID", context.Background(), tc.session.UserID) - assert.True(t, ok, fmt.Sprintf("RetrieveByID was not called on %s", tc.desc)) - } - authCall.Unset() - repoCall.Unset() - }) - } -} - -func TestRevokeRefreshToken(t *testing.T) { - svc, authsvc, crepo, _, _ := newService() - - rUser := user - rUser.Credentials.Secret, _ = phasher.Hash(user.Credentials.Secret) - - cases := []struct { - desc string - session authn.Session - tokenID string - revokeResp *grpcTokenV1.RevokeRes - revokeErr error - repoResp users.User - repoErr error - err error - }{ - { - desc: "revoke refresh token successfully", - session: authn.Session{DomainUserID: validID, UserID: validID, DomainID: validID}, - tokenID: validToken, - revokeResp: &grpcTokenV1.RevokeRes{}, - repoResp: rUser, - err: nil, - }, - { - desc: "revoke refresh token with empty domain id", - session: authn.Session{UserID: validID}, - tokenID: validToken, - revokeResp: &grpcTokenV1.RevokeRes{}, - repoResp: rUser, - err: nil, - }, - { - desc: "revoke refresh token for non-existing user", - session: authn.Session{DomainUserID: validID, UserID: validID, DomainID: validID}, - tokenID: validToken, - repoErr: repoerr.ErrNotFound, - err: repoerr.ErrNotFound, - }, - { - desc: "revoke refresh token for disabled user", - session: authn.Session{DomainUserID: validID, UserID: validID, DomainID: validID}, - tokenID: validToken, - repoResp: users.User{Status: users.DisabledStatus}, - err: svcerr.ErrAuthentication, - }, - { - desc: "revoke refresh token with revoke service error", - session: authn.Session{DomainUserID: validID, UserID: validID, DomainID: validID}, - tokenID: validToken, - revokeResp: &grpcTokenV1.RevokeRes{}, - revokeErr: svcerr.ErrAuthorization, - repoResp: rUser, - err: svcerr.ErrAuthorization, - }, - { - desc: "revoke refresh token not found", - session: authn.Session{DomainUserID: validID, UserID: validID, DomainID: validID}, - tokenID: validToken, - revokeResp: &grpcTokenV1.RevokeRes{}, - revokeErr: svcerr.ErrNotFound, - repoResp: rUser, - err: svcerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := crepo.On("RetrieveByID", context.Background(), tc.session.UserID).Return(tc.repoResp, tc.repoErr) - authCall := authsvc.On("Revoke", context.Background(), &grpcTokenV1.RevokeReq{UserId: tc.session.UserID, TokenId: tc.tokenID}).Return(tc.revokeResp, tc.revokeErr) - err := svc.RevokeRefreshToken(context.Background(), tc.session, tc.tokenID) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - if err == nil { - ok := repoCall.Parent.AssertCalled(t, "RetrieveByID", context.Background(), tc.session.UserID) - assert.True(t, ok, fmt.Sprintf("RetrieveByID was not called on %s", tc.desc)) - ok = authCall.Parent.AssertCalled(t, "Revoke", context.Background(), &grpcTokenV1.RevokeReq{UserId: tc.session.UserID, TokenId: tc.tokenID}) - assert.True(t, ok, fmt.Sprintf("Revoke was not called on %s", tc.desc)) - } - repoCall.Unset() - authCall.Unset() - }) - } -} - -func TestListActiveRefreshTokens(t *testing.T) { - svc, authsvc, crepo, _, _ := newService() - - rUser := user - rUser.Credentials.Secret, _ = phasher.Hash(user.Credentials.Secret) - - cases := []struct { - desc string - session authn.Session - listResp *grpcTokenV1.ListUserRefreshTokensRes - listErr error - repoResp users.User - repoErr error - expectedTokens int - err error - }{ - { - desc: "list active refresh tokens successfully", - session: authn.Session{DomainUserID: validID, UserID: validID, DomainID: validID}, - listResp: &grpcTokenV1.ListUserRefreshTokensRes{ - RefreshTokens: []*grpcTokenV1.RefreshToken{ - {Id: "token1", Description: "token1"}, - {Id: "token2", Description: "token2"}, - }, - }, - repoResp: rUser, - expectedTokens: 2, - err: nil, - }, - { - desc: "list active refresh tokens with empty domain id", - session: authn.Session{UserID: validID}, - listResp: &grpcTokenV1.ListUserRefreshTokensRes{ - RefreshTokens: []*grpcTokenV1.RefreshToken{ - {Id: "token1", Description: "token1"}, - }, - }, - repoResp: rUser, - expectedTokens: 1, - err: nil, - }, - { - desc: "list active refresh tokens for non-existing user", - session: authn.Session{DomainUserID: validID, UserID: validID, DomainID: validID}, - repoErr: repoerr.ErrNotFound, - err: repoerr.ErrNotFound, - }, - { - desc: "list active refresh tokens for disabled user", - session: authn.Session{DomainUserID: validID, UserID: validID, DomainID: validID}, - repoResp: users.User{Status: users.DisabledStatus}, - err: svcerr.ErrAuthentication, - }, - { - desc: "list active refresh tokens with list service error", - session: authn.Session{DomainUserID: validID, UserID: validID, DomainID: validID}, - listResp: &grpcTokenV1.ListUserRefreshTokensRes{}, - listErr: svcerr.ErrAuthentication, - repoResp: rUser, - err: svcerr.ErrAuthentication, - }, - { - desc: "list active refresh tokens with empty list", - session: authn.Session{DomainUserID: validID, UserID: validID, DomainID: validID}, - listResp: &grpcTokenV1.ListUserRefreshTokensRes{RefreshTokens: []*grpcTokenV1.RefreshToken{}}, - repoResp: rUser, - expectedTokens: 0, - err: nil, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := crepo.On("RetrieveByID", context.Background(), tc.session.UserID).Return(tc.repoResp, tc.repoErr) - authCall := authsvc.On("ListUserRefreshTokens", context.Background(), &grpcTokenV1.ListUserRefreshTokensReq{UserId: tc.session.UserID}).Return(tc.listResp, tc.listErr) - tokens, err := svc.ListActiveRefreshTokens(context.Background(), tc.session) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - if err == nil { - assert.NotNil(t, tokens, fmt.Sprintf("%s: expected tokens not to be nil\n", tc.desc)) - assert.Equal(t, tc.expectedTokens, len(tokens.GetRefreshTokens()), fmt.Sprintf("%s: expected %d tokens got %d\n", tc.desc, tc.expectedTokens, len(tokens.GetRefreshTokens()))) - ok := repoCall.Parent.AssertCalled(t, "RetrieveByID", context.Background(), tc.session.UserID) - assert.True(t, ok, fmt.Sprintf("RetrieveByID was not called on %s", tc.desc)) - ok = authCall.Parent.AssertCalled(t, "ListUserRefreshTokens", context.Background(), &grpcTokenV1.ListUserRefreshTokensReq{UserId: tc.session.UserID}) - assert.True(t, ok, fmt.Sprintf("ListUserRefreshTokens was not called on %s", tc.desc)) - } - repoCall.Unset() - authCall.Unset() - }) - } -} - -func TestSendPasswordReset(t *testing.T) { - svc, auth, cRepo, _, e := newService() - - cases := []struct { - desc string - email string - retrieveByEmailResponse users.User - issueResponse *grpcTokenV1.Token - retrieveByEmailErr error - issueErr error - err error - }{ - { - desc: "generate reset token for existing user", - email: "existingemail@example.com", - retrieveByEmailResponse: user, - issueResponse: &grpcTokenV1.Token{AccessToken: validToken, RefreshToken: &validToken, AccessType: "3"}, - err: nil, - }, - { - desc: "generate reset token for user with non-existing user", - email: "example@example.com", - retrieveByEmailResponse: users.User{ - ID: testsutil.GenerateUUID(t), - Email: "", - }, - retrieveByEmailErr: repoerr.ErrNotFound, - err: repoerr.ErrNotFound, - }, - { - desc: "generate reset token with failed to issue token", - email: "existingemail@example.com", - retrieveByEmailResponse: user, - issueResponse: &grpcTokenV1.Token{}, - issueErr: svcerr.ErrAuthorization, - err: svcerr.ErrAuthorization, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := cRepo.On("RetrieveByEmail", context.Background(), tc.email).Return(tc.retrieveByEmailResponse, tc.retrieveByEmailErr) - authCall := auth.On("Issue", context.Background(), mock.Anything).Return(tc.issueResponse, tc.issueErr) - svcCall := e.On("SendPasswordReset", []string{tc.email}, user.Credentials.Username, validToken).Return(tc.err) - err := svc.SendPasswordReset(context.Background(), tc.email) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - repoCall.Parent.AssertCalled(t, "RetrieveByEmail", context.Background(), tc.email) - repoCall.Unset() - authCall.Unset() - svcCall.Unset() - }) - } -} - -func TestResetSecret(t *testing.T) { - svc, cRepo := newServiceMinimal() - - user := users.User{ - ID: "userID", - Email: "test@example.com", - Credentials: users.Credentials{ - Secret: "Strongsecret", - }, - } - - cases := []struct { - desc string - newSecret string - session authn.Session - retrieveByIDResponse users.User - updateSecretResponse users.User - retrieveByIDErr error - updateSecretErr error - err error - }{ - { - desc: "reset secret with successfully", - newSecret: "newStrongSecret", - session: authn.Session{UserID: validID, SuperAdmin: true}, - retrieveByIDResponse: user, - updateSecretResponse: users.User{ - ID: "userID", - Email: "test@example.com", - Credentials: users.Credentials{ - Secret: "newStrongSecret", - }, - }, - err: nil, - }, - { - desc: "reset secret with invalid ID", - newSecret: "newStrongSecret", - session: authn.Session{UserID: validID, SuperAdmin: true}, - retrieveByIDResponse: users.User{}, - retrieveByIDErr: repoerr.ErrNotFound, - err: repoerr.ErrNotFound, - }, - { - desc: "reset secret with empty email", - session: authn.Session{UserID: validID, SuperAdmin: true}, - newSecret: "newStrongSecret", - retrieveByIDResponse: users.User{ - ID: "userID", - Email: "", - }, - err: nil, - }, - { - desc: "reset secret with failed to update secret", - newSecret: "newStrongSecret", - session: authn.Session{UserID: validID, SuperAdmin: true}, - retrieveByIDResponse: user, - updateSecretResponse: users.User{}, - updateSecretErr: svcerr.ErrUpdateEntity, - err: svcerr.ErrAuthorization, - }, - { - desc: "reset secret with a too long secret", - newSecret: strings.Repeat("strongSecret", 10), - session: authn.Session{UserID: validID, SuperAdmin: true}, - retrieveByIDResponse: user, - err: errHashPassword, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := cRepo.On("RetrieveByID", context.Background(), mock.Anything).Return(tc.retrieveByIDResponse, tc.retrieveByIDErr) - repoCall1 := cRepo.On("UpdateSecret", context.Background(), mock.Anything).Return(tc.updateSecretResponse, tc.updateSecretErr) - err := svc.ResetSecret(context.Background(), tc.session, tc.newSecret) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - if tc.err == nil { - repoCall1.Parent.AssertCalled(t, "UpdateSecret", context.Background(), mock.Anything) - repoCall.Parent.AssertCalled(t, "RetrieveByID", context.Background(), validID) - } - repoCall1.Unset() - repoCall.Unset() - }) - } -} - -func TestViewProfile(t *testing.T) { - svc, cRepo := newServiceMinimal() - - user := users.User{ - ID: "userID", - Email: "existingEmail", - Credentials: users.Credentials{ - Secret: "Strongsecret", - }, - } - cases := []struct { - desc string - user users.User - session authn.Session - retrieveByIDResponse users.User - retrieveByIDErr error - err error - }{ - { - desc: "view profile successfully", - user: user, - session: authn.Session{UserID: validID}, - retrieveByIDResponse: user, - err: nil, - }, - { - desc: "view profile with invalid ID", - user: user, - session: authn.Session{UserID: wrongID}, - retrieveByIDResponse: users.User{}, - retrieveByIDErr: repoerr.ErrNotFound, - err: repoerr.ErrNotFound, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := cRepo.On("RetrieveByID", context.Background(), mock.Anything).Return(tc.retrieveByIDResponse, tc.retrieveByIDErr) - _, err := svc.ViewProfile(context.Background(), tc.session) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - repoCall.Parent.AssertCalled(t, "RetrieveByID", context.Background(), mock.Anything) - repoCall.Unset() - }) - } -} - -func TestOAuthCallback(t *testing.T) { - svc, _, cRepo, policies, _ := newService() - - cases := []struct { - desc string - user users.User - retrieveByEmailResponse users.User - retrieveByEmailErr error - saveResponse users.User - addPoliciesErr error - updateVerifiedAtErr error - err error - }{ - { - desc: "oauth signin callback with already existing user", - user: users.User{ - Email: "test@example.com", - }, - retrieveByEmailResponse: users.User{ - ID: testsutil.GenerateUUID(t), - Role: users.UserRole, - VerifiedAt: time.Now(), - }, - err: nil, - }, - { - desc: "oauth signup callback with user not found", - user: users.User{ - Email: "test@example.com", - }, - retrieveByEmailErr: repoerr.ErrNotFound, - saveResponse: users.User{ - ID: testsutil.GenerateUUID(t), - Role: users.UserRole, - }, - err: nil, - }, - { - desc: "oauth signup callback with malformed entity", - user: users.User{ - Email: "test@example.com", - }, - retrieveByEmailErr: repoerr.ErrMalformedEntity, - err: repoerr.ErrMalformedEntity, - }, - { - desc: "oauth signup callback with failed to register user", - user: users.User{ - Email: "test@example.com", - }, - addPoliciesErr: svcerr.ErrAuthorization, - retrieveByEmailErr: repoerr.ErrNotFound, - err: svcerr.ErrAuthorization, - }, - { - desc: "oauth signin callback with user not in the platform", - user: users.User{ - Email: "test@example.com", - }, - retrieveByEmailResponse: users.User{ - ID: testsutil.GenerateUUID(t), - Role: users.UserRole, - }, - err: nil, - }, - { - desc: "oauth signin callback with failed update verified at", - user: users.User{ - Email: "test@example.com", - }, - retrieveByEmailResponse: users.User{ - ID: testsutil.GenerateUUID(t), - Role: users.UserRole, - }, - updateVerifiedAtErr: svcerr.ErrUpdateEntity, - err: svcerr.ErrUpdateEntity, - }, - } - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := cRepo.On("RetrieveByEmail", context.Background(), tc.user.Email).Return(tc.retrieveByEmailResponse, tc.retrieveByEmailErr) - repoCall1 := cRepo.On("Save", context.Background(), mock.Anything).Return(tc.saveResponse, nil) - repoCall2 := cRepo.On("UpdateVerifiedAt", context.Background(), mock.MatchedBy(func(u users.User) bool { - assert.NotEmpty(t, u.ID, "UpdateVerifiedAt must be called with non-empty user ID") - return u.ID != "" - })).Return(tc.retrieveByEmailResponse, tc.updateVerifiedAtErr) - policyCall := policies.On("AddPolicies", context.Background(), mock.Anything).Return(tc.addPoliciesErr) - _, err := svc.OAuthCallback(context.Background(), tc.user) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - repoCall.Parent.AssertCalled(t, "RetrieveByEmail", context.Background(), tc.user.Email) - repoCall.Unset() - repoCall1.Unset() - policyCall.Unset() - _ = repoCall2 - cRepo.ExpectedCalls = nil - policies.ExpectedCalls = nil - }) - } -} - -func TestSendVerification(t *testing.T) { - svc, _, cRepo, _, e := newService() - - verifiedAt := time.Now().UTC() - cases := []struct { - desc string - session authn.Session - retrieveByIDResponse users.User - retrieveByIDError error - retrieveUserVerResponse users.UserVerification - retrieveUserVerError error - addUserVerError error - sendVerificationEmailError error - err error - }{ - { - desc: "send verification email successfully", - session: authn.Session{UserID: user.ID}, - retrieveByIDResponse: user, - retrieveUserVerError: repoerr.ErrNotFound, - sendVerificationEmailError: nil, - err: nil, - }, - { - desc: "send verification email for already verified user", - session: authn.Session{UserID: user.ID}, - retrieveByIDResponse: users.User{VerifiedAt: verifiedAt}, - err: svcerr.ErrUserAlreadyVerified, - }, - { - desc: "send verification email for non-existing user", - session: authn.Session{UserID: wrongID}, - retrieveByIDError: repoerr.ErrNotFound, - err: repoerr.ErrNotFound, - }, - { - desc: "send verification email with failed to retrieve user verification", - session: authn.Session{UserID: user.ID}, - retrieveByIDResponse: user, - retrieveUserVerError: svcerr.ErrViewEntity, - err: svcerr.ErrViewEntity, - }, - { - desc: "send verification email with failed to add user verification", - session: authn.Session{UserID: user.ID}, - retrieveByIDResponse: user, - retrieveUserVerError: repoerr.ErrNotFound, - addUserVerError: svcerr.ErrCreateEntity, - err: svcerr.ErrCreateEntity, - }, - { - desc: "send verification email with failed to send email", - session: authn.Session{UserID: user.ID}, - retrieveByIDResponse: user, - retrieveUserVerError: repoerr.ErrNotFound, - sendVerificationEmailError: svcerr.ErrCreateEntity, - err: svcerr.ErrCreateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := cRepo.On("RetrieveByID", context.Background(), tc.session.UserID).Return(tc.retrieveByIDResponse, tc.retrieveByIDError) - repoCall1 := cRepo.On("RetrieveUserVerification", context.Background(), mock.Anything, mock.Anything).Return(tc.retrieveUserVerResponse, tc.retrieveUserVerError) - repoCall2 := cRepo.On("AddUserVerification", context.Background(), mock.Anything).Return(tc.addUserVerError) - emailCall := e.On("SendVerification", []string{user.Email}, user.Credentials.Username, mock.Anything).Return(tc.sendVerificationEmailError) - - err := svc.SendVerification(context.Background(), tc.session) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - - repoCall.Unset() - repoCall1.Unset() - repoCall2.Unset() - emailCall.Unset() - }) - } -} - -func TestVerifyEmail(t *testing.T) { - //nolint:dogsled - svc, _, cRepo, _, _ := newService() - uv, err := users.NewUserVerification(user.ID, user.Email) - assert.Nil(t, err, fmt.Sprintf("failed to create user verification: %v", err)) - uvs, err := uv.Encode() - assert.Nil(t, err, fmt.Sprintf("failed to encode user verification: %v", err)) - createdAt := time.Now().Add(-5 * users.VerificationExpiryDuration).UTC() - expiresdAt := time.Now().Add(-users.VerificationExpiryDuration).UTC() - cases := []struct { - desc string - uvs string - retrieveUserVerResponse users.UserVerification - retrieveUserVerError error - updateUserVerError error - updateVerifiedAtError error - err error - }{ - { - desc: "verify email successfully", - uvs: uvs, - retrieveUserVerResponse: uv, - err: nil, - }, - { - desc: "verify email with malformed token", - uvs: "invalid", - err: svcerr.ErrInvalidUserVerification, - }, - { - desc: "verify email with non-existing user verification", - uvs: uvs, - retrieveUserVerError: repoerr.ErrNotFound, - err: svcerr.ErrViewEntity, - }, - { - desc: "verify email with expired token", - uvs: uvs, - retrieveUserVerResponse: users.UserVerification{ - UserID: uv.UserID, - Email: uv.Email, - OTP: uv.OTP, - ExpiresAt: expiresdAt, - CreatedAt: createdAt, - UsedAt: uv.UsedAt, - }, - err: svcerr.ErrUserVerificationExpired, - }, - { - desc: "verify email with failed to update user verification", - uvs: uvs, - retrieveUserVerResponse: uv, - updateUserVerError: svcerr.ErrUpdateEntity, - err: svcerr.ErrUpdateEntity, - }, - { - desc: "verify email with failed to update verified at", - uvs: uvs, - retrieveUserVerResponse: uv, - updateVerifiedAtError: svcerr.ErrUpdateEntity, - err: svcerr.ErrUpdateEntity, - }, - } - - for _, tc := range cases { - t.Run(tc.desc, func(t *testing.T) { - repoCall := cRepo.On("RetrieveUserVerification", context.Background(), mock.Anything, mock.Anything).Return(tc.retrieveUserVerResponse, tc.retrieveUserVerError) - repoCall1 := cRepo.On("UpdateUserVerification", context.Background(), mock.Anything).Return(tc.updateUserVerError) - repoCall2 := cRepo.On("UpdateVerifiedAt", context.Background(), mock.Anything).Return(users.User{}, tc.updateVerifiedAtError) - _, err := svc.VerifyEmail(context.Background(), tc.uvs) - assert.True(t, errors.Contains(err, tc.err), fmt.Sprintf("%s: expected %s got %s\n", tc.desc, tc.err, err)) - - repoCall.Unset() - repoCall1.Unset() - repoCall2.Unset() - }) - } -} diff --git a/users/status.go b/users/status.go deleted file mode 100644 index 974cec227..000000000 --- a/users/status.go +++ /dev/null @@ -1,83 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package users - -import ( - "encoding/json" - "strings" - - svcerr "github.com/absmach/magistrala/pkg/errors/service" -) - -// Status represents User status. -type Status uint8 - -// Possible User status values. -const ( - // EnabledStatus represents enabled User. - EnabledStatus Status = iota - // DisabledStatus represents disabled User. - DisabledStatus - // DeletedStatus represents a user that will be deleted. - DeletedStatus - - // AllStatus is used for querying purposes to list users irrespective - // of their status - both enabled and disabled. It is never stored in the - // database as the actual User status and should always be the largest - // value in this enumeration. - AllStatus -) - -// String representation of the possible status values. -const ( - Disabled = "disabled" - Enabled = "enabled" - Deleted = "deleted" - All = "all" - Unknown = "unknown" -) - -// String converts user/group status to string literal. -func (s Status) String() string { - switch s { - case DisabledStatus: - return Disabled - case EnabledStatus: - return Enabled - case DeletedStatus: - return Deleted - case AllStatus: - return All - default: - return Unknown - } -} - -// ToStatus converts string value to a valid User/Group status. -func ToStatus(status string) (Status, error) { - switch status { - case "", Enabled: - return EnabledStatus, nil - case Disabled: - return DisabledStatus, nil - case Deleted: - return DeletedStatus, nil - case All: - return AllStatus, nil - } - return Status(0), svcerr.ErrInvalidStatus -} - -// Custom Marshaller for Uesr/Groups. -func (s Status) MarshalJSON() ([]byte, error) { - return json.Marshal(s.String()) -} - -// Custom Unmarshaler for User/Groups. -func (s *Status) UnmarshalJSON(data []byte) error { - str := strings.Trim(string(data), "\"") - val, err := ToStatus(str) - *s = val - return err -} diff --git a/users/users.go b/users/users.go deleted file mode 100644 index a8ff152c1..000000000 --- a/users/users.go +++ /dev/null @@ -1,288 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package users - -import ( - "context" - "net/mail" - "strings" - "time" - - grpcTokenV1 "github.com/absmach/magistrala/api/grpc/token/v1" - "github.com/absmach/magistrala/pkg/authn" - "github.com/absmach/magistrala/pkg/errors" - "github.com/absmach/magistrala/pkg/postgres" -) - -type User struct { - ID string `json:"id"` - FirstName string `json:"first_name,omitempty"` - LastName string `json:"last_name,omitempty"` - Tags []string `json:"tags,omitempty"` - Metadata Metadata `json:"metadata,omitempty"` - PrivateMetadata Metadata `json:"private_metadata,omitempty"` - Status Status `json:"status"` // 0 for enabled, 1 for disabled - Role Role `json:"role"` // 0 for normal user, 1 for admin - ProfilePicture string `json:"profile_picture,omitempty"` // profile picture URL - Credentials Credentials `json:"credentials,omitempty"` - Permissions []string `json:"permissions,omitempty"` - Email string `json:"email,omitempty"` - CreatedAt time.Time `json:"created_at,omitempty"` - UpdatedAt time.Time `json:"updated_at,omitempty"` - UpdatedBy string `json:"updated_by,omitempty"` - VerifiedAt time.Time `json:"verified_at,omitempty"` - AuthProvider string `json:"auth_provider,omitempty"` -} - -type Credentials struct { - Username string `json:"username,omitempty"` // username or profile name - Secret string `json:"secret,omitempty"` // password or token -} - -type UsersPage struct { - Page - Users []User -} - -// Metadata represents arbitrary JSON. -type Metadata map[string]any - -type UserReq struct { - FirstName *string `json:"first_name,omitempty"` - LastName *string `json:"last_name,omitempty"` - Metadata *Metadata `json:"metadata,omitempty"` - PrivateMetadata *Metadata `json:"private_metadata,omitempty"` - Tags *[]string `json:"tags,omitempty"` - ProfilePicture *string `json:"profile_picture,omitempty"` - UpdatedBy *string `json:"updated_by,omitempty"` - UpdatedAt *time.Time `json:"updated_at,omitempty"` -} - -// MembersPage contains page related metadata as well as list of members that -// belong to this page. -type MembersPage struct { - Page - Members []User -} - -// UserRepository struct implements the Repository interface. -type UserRepository struct { - DB postgres.Database -} - -type Repository interface { - // RetrieveByID retrieves user by their unique ID. - RetrieveByID(ctx context.Context, id string) (User, error) - - // RetrieveAll retrieves all users. - RetrieveAll(ctx context.Context, pm Page) (UsersPage, error) - - // RetrieveByEmail retrieves user by its unique credentials. - RetrieveByEmail(ctx context.Context, email string) (User, error) - - // RetrieveByUsername retrieves user by its unique credentials. - RetrieveByUsername(ctx context.Context, username string) (User, error) - - // Update updates the user name and metadata. - Update(ctx context.Context, id string, user UserReq) (User, error) - - // UpdateUsername updates the User's names. - UpdateUsername(ctx context.Context, user User) (User, error) - - // UpdateSecret updates secret for user with given email. - UpdateSecret(ctx context.Context, user User) (User, error) - - // UpdateEmail updates email for user with given id. - UpdateEmail(ctx context.Context, user User) (User, error) - - // UpdateRole updates role for user with given id. - UpdateRole(ctx context.Context, user User) (User, error) - - // UpdateVerifiedAt updates the verified time for user with given id. - UpdateVerifiedAt(ctx context.Context, user User) (User, error) - - // ChangeStatus changes user status to enabled or disabled - ChangeStatus(ctx context.Context, user User) (User, error) - - // Delete deletes user with given id - Delete(ctx context.Context, id string) error - - // Searchusers retrieves users based on search criteria. - SearchUsers(ctx context.Context, pm Page) (UsersPage, error) - - // RetrieveAllByIDs retrieves for given user IDs . - RetrieveAllByIDs(ctx context.Context, pm Page) (UsersPage, error) - - CheckSuperAdmin(ctx context.Context, adminID string) error - - // Save persists the user account. A non-nil error is returned to indicate - // operation failure. - Save(ctx context.Context, user User) (User, error) - - // AddUserVerification adds new verification for given user id and email - AddUserVerification(ctx context.Context, uv UserVerification) error - - // RetrieveVerificationToken retrieves verification token of given user id and email. - RetrieveUserVerification(ctx context.Context, userID, email string) (UserVerification, error) - - // UpdateUserVerificationDetails update verification details for the given user id and email. - UpdateUserVerification(ctx context.Context, uv UserVerification) error -} - -// Validate returns an error if user representation is invalid. -func (u User) Validate() error { - if !isEmail(u.Email) { - return errors.ErrMalformedEntity - } - return nil -} - -func isEmail(email string) bool { - _, err := mail.ParseAddress(email) - return err == nil -} - -type Operator uint8 - -const ( - OrOp Operator = iota - AndOp -) - -type TagsQuery struct { - Elements []string - Operator Operator -} - -func ToTagsQuery(s string) TagsQuery { - switch { - case strings.Contains(s, "+"): - elements := strings.Split(s, "+") - for i := range elements { - elements[i] = strings.TrimSpace(elements[i]) - } - return TagsQuery{Elements: elements, Operator: AndOp} - case strings.Contains(s, ","): - elements := strings.Split(s, ",") - for i := range elements { - elements[i] = strings.TrimSpace(elements[i]) - } - return TagsQuery{Elements: elements, Operator: OrOp} - default: - return TagsQuery{Elements: []string{s}, Operator: OrOp} - } -} - -// Page contains page metadata that helps navigation. -type Page struct { - Total uint64 `json:"total"` - Offset uint64 `json:"offset"` - Limit uint64 `json:"limit"` - OnlyTotal bool `json:"only_total"` - Id string `json:"id,omitempty"` - Order string `json:"order,omitempty"` - Dir string `json:"dir,omitempty"` - Metadata Metadata `json:"metadata,omitempty"` - Domain string `json:"domain,omitempty"` - Tags TagsQuery `json:"tag,omitempty"` - Permission string `json:"permission,omitempty"` - Status Status `json:"status,omitempty"` - IDs []string `json:"ids,omitempty"` - Role Role `json:"-"` - ListPerms bool `json:"-"` - Username string `json:"username,omitempty"` - FirstName string `json:"first_name,omitempty"` - LastName string `json:"last_name,omitempty"` - Email string `json:"email,omitempty"` - Verified bool `json:"verified,omitempty"` - CreatedFrom time.Time `json:"created_from,omitempty"` - CreatedTo time.Time `json:"created_to,omitempty"` -} - -// Service specifies an API that must be fullfiled by the domain service -// implementation, and all of its decorators (e.g. logging & metrics). -type Service interface { - // Register creates new user. In case of the failed registration, a - // non-nil error value is returned. - Register(ctx context.Context, session authn.Session, user User, selfRegister bool) (User, error) - - // SendVerification sends a verification email to the user. - SendVerification(ctx context.Context, session authn.Session) error - - // VerifyEmail verifies user's email using the verification token. - VerifyEmail(ctx context.Context, verificationToken string) (User, error) - - // View retrieves user info for a given user ID and an authorized token. - View(ctx context.Context, session authn.Session, id string) (User, error) - - // ViewProfile retrieves user info for a given token. - ViewProfile(ctx context.Context, session authn.Session) (User, error) - - // ListUsers retrieves users list for a valid auth token. - ListUsers(ctx context.Context, session authn.Session, pm Page) (UsersPage, error) - - // SearchUsers searches for users with provided filters for a valid auth token. - SearchUsers(ctx context.Context, pm Page) (UsersPage, error) - - // Update updates the user's name and metadata. - Update(ctx context.Context, session authn.Session, id string, user UserReq) (User, error) - - // UpdateTags updates the user's tags. - UpdateTags(ctx context.Context, session authn.Session, id string, user UserReq) (User, error) - - // UpdateEmail updates the user's email. - UpdateEmail(ctx context.Context, session authn.Session, id, email string) (User, error) - - // UpdateUsername updates the user's username. - UpdateUsername(ctx context.Context, session authn.Session, id, username string) (User, error) - - // UpdateProfilePicture updates the user's profile picture. - UpdateProfilePicture(ctx context.Context, session authn.Session, id string, usr UserReq) (User, error) - - // SendPasswordReset generates reset password link and sends it to the user via email. - SendPasswordReset(ctx context.Context, email string) error - - // UpdateSecret updates the user's secret. - UpdateSecret(ctx context.Context, session authn.Session, oldSecret, newSecret string) (User, error) - - // ResetSecret change users secret in reset flow. - // token can be authentication token or secret reset token. - ResetSecret(ctx context.Context, session authn.Session, secret string) error - - // UpdateRole updates the user's Role. - UpdateRole(ctx context.Context, session authn.Session, user User) (User, error) - - // Enable logically enables the user identified with the provided ID. - Enable(ctx context.Context, session authn.Session, id string) (User, error) - - // Disable logically disables the user identified with the provided ID. - Disable(ctx context.Context, session authn.Session, id string) (User, error) - - // Delete deletes user with given ID. - Delete(ctx context.Context, session authn.Session, id string) error - - // Identify returns the user id from the given token. - Identify(ctx context.Context, session authn.Session) (string, error) - - // IssueToken issues a new access and refresh token when provided with either a username or email. - IssueToken(ctx context.Context, identity, secret, description string) (*grpcTokenV1.Token, error) - - // RefreshToken refreshes expired access tokens. - // After an access token expires, the refresh token is used to get - // a new pair of access and refresh tokens. - RefreshToken(ctx context.Context, session authn.Session, refreshToken string) (*grpcTokenV1.Token, error) - - // RevokeRefreshToken revokes a refresh token by its ID. - RevokeRefreshToken(ctx context.Context, session authn.Session, tokenID string) error - - // ListActiveRefreshTokens lists all active refresh tokens for the authenticated user. - ListActiveRefreshTokens(ctx context.Context, session authn.Session) (*grpcTokenV1.ListUserRefreshTokensRes, error) - - // OAuthCallback handles the callback from any supported OAuth provider. - // It processes the OAuth tokens and either signs in or signs up the user based on the provided state. - OAuthCallback(ctx context.Context, user User) (User, error) - - // OAuthAddUserPolicy adds a policy to the user for an OAuth request. - OAuthAddUserPolicy(ctx context.Context, user User) error -} diff --git a/users/verification.go b/users/verification.go deleted file mode 100644 index a7e03d593..000000000 --- a/users/verification.go +++ /dev/null @@ -1,121 +0,0 @@ -// Copyright (c) Abstract Machines -// SPDX-License-Identifier: Apache-2.0 - -package users - -import ( - "crypto/rand" - "encoding/base64" - "encoding/json" - "time" - - "github.com/absmach/magistrala/pkg/errors" - svcerr "github.com/absmach/magistrala/pkg/errors/service" -) - -const VerificationExpiryDuration = 24 * time.Hour - -var ( - errFailedToCreateUserVerification = errors.New("failed to create new user verification") - errFailedToEncodeUserVerification = errors.New("failed to encode user verification") - errFailedToDecodeUserVerification = errors.New("failed to decode user verification") -) - -// UserVerification OTP is sent to the user's email as base64 encoded with UserID, Email and OTP. It should not be exposed via API. -type UserVerification struct { - UserID string `json:"user_id"` - Email string `json:"email"` - OTP string `json:"otp"` - CreatedAt time.Time `json:"-"` - ExpiresAt time.Time `json:"-"` - UsedAt time.Time `json:"-"` -} - -func NewUserVerification(userID, email string) (UserVerification, error) { - randomBytes := make([]byte, 32) - if _, err := rand.Read(randomBytes); err != nil { - return UserVerification{}, errors.Wrap(errFailedToCreateUserVerification, err) - } - - return UserVerification{ - UserID: userID, - Email: email, - OTP: base64.URLEncoding.EncodeToString(randomBytes), - CreatedAt: time.Now().UTC(), - ExpiresAt: time.Now().Add(VerificationExpiryDuration).UTC(), - }, nil -} - -func (u UserVerification) Encode() (string, error) { - jsonBytes, err := json.Marshal(u) - if err != nil { - return "", errors.Wrap(errFailedToEncodeUserVerification, err) - } - - return base64.URLEncoding.EncodeToString(jsonBytes), nil -} - -func (u *UserVerification) Decode(data string) error { - decodedPayload, err := base64.URLEncoding.DecodeString(data) - if err != nil { - return errors.Wrap(errFailedToDecodeUserVerification, err) - } - - if err := json.Unmarshal(decodedPayload, u); err != nil { - return errors.Wrap(errFailedToDecodeUserVerification, err) - } - - if u.UserID == "" || u.Email == "" || u.OTP == "" { - return svcerr.ErrInvalidUserVerification - } - - return nil -} - -func (u UserVerification) Valid() error { - if u.UserID == "" || u.Email == "" || u.OTP == "" { - return svcerr.ErrInvalidUserVerification - } - - // Verification should have created time. - if u.CreatedAt.IsZero() { - return svcerr.ErrInvalidUserVerification - } - - // Verification should have expiry time. - if u.ExpiresAt.IsZero() { - return svcerr.ErrInvalidUserVerification - } - - // Expiry time should not be before Created time - if u.ExpiresAt.Before(u.CreatedAt) { - return svcerr.ErrInvalidUserVerification - } - - // Verification should be not be Expired. - if time.Now().After(u.ExpiresAt) { - return svcerr.ErrUserVerificationExpired - } - - // Verification should not be used. - if !u.UsedAt.IsZero() { - return svcerr.ErrUserVerificationExpired - } - - return nil -} - -func (u UserVerification) Match(ruv UserVerification) error { - if u.UserID != ruv.UserID { - return svcerr.ErrInvalidUserVerification - } - - if u.Email != ruv.Email { - return svcerr.ErrInvalidUserVerification - } - - if u.OTP != ruv.OTP { - return svcerr.ErrInvalidUserVerification - } - return nil -}