diff --git a/.github/cicd-images.json b/.github/cicd-images.json index 37103166..bbad0524 100644 --- a/.github/cicd-images.json +++ b/.github/cicd-images.json @@ -1,6 +1,6 @@ { "project": "basaltpass", - "namespace": "basaltpass", + "namespace": "beancs-system", "images": [ { "name": "backend", diff --git a/.github/workflows/deploy-tag.yml b/.github/workflows/deploy-tag.yml deleted file mode 100644 index 195fbc8a..00000000 --- a/.github/workflows/deploy-tag.yml +++ /dev/null @@ -1,317 +0,0 @@ -name: Build And Deploy From Tag - -on: - push: - tags: - - "deploy-prod-v*.*.*" - -permissions: - contents: read - packages: write - deployments: write - -jobs: - build-and-deploy: - runs-on: ubuntu-latest - env: - DEPLOY_HOST: ${{ vars.DEPLOY_HOST }} - DEPLOY_PORT: ${{ vars.DEPLOY_PORT || '22' }} - DEPLOY_USER: ${{ vars.DEPLOY_USER }} - DEPLOY_PATH: ${{ vars.DEPLOY_PATH || '/opt/basaltpass' }} - GHCR_NAMESPACE: ${{ vars.GHCR_NAMESPACE || github.repository_owner }} - SERVICE_URL: ${{ vars.SERVICE_URL }} - BASALTPASS_DATABASE_DRIVER: ${{ vars.BASALTPASS_DATABASE_DRIVER || 'sqlite' }} - BASALTPASS_CORS_ALLOW_ORIGINS: ${{ vars.BASALTPASS_CORS_ALLOW_ORIGINS }} - BASALTPASS_EMAIL_PROVIDER: ${{ vars.BASALTPASS_EMAIL_PROVIDER }} - BASALTPASS_EMAIL_FROM: ${{ vars.BASALTPASS_EMAIL_FROM }} - BASALTPASS_EMAIL_SMTP_PORT: ${{ vars.BASALTPASS_EMAIL_SMTP_PORT }} - BASALTPASS_EMAIL_SMTP_USE_TLS: ${{ vars.BASALTPASS_EMAIL_SMTP_USE_TLS }} - BASALTPASS_EMAIL_SMTP_USE_SSL: ${{ vars.BASALTPASS_EMAIL_SMTP_USE_SSL }} - BASALTPASS_EMAIL_SMTP_SKIP_CERT_VERIFY: ${{ vars.BASALTPASS_EMAIL_SMTP_SKIP_CERT_VERIFY }} - BASALTPASS_EMAIL_AWS_SES_REGION: ${{ vars.BASALTPASS_EMAIL_AWS_SES_REGION }} - steps: - - name: Checkout source repo - uses: actions/checkout@v4 - - - name: Validate deployment configuration - run: | - for required_var in DEPLOY_HOST DEPLOY_USER DEPLOY_PATH GHCR_NAMESPACE SERVICE_URL BASALTPASS_DATABASE_DRIVER BASALTPASS_CORS_ALLOW_ORIGINS; do - value="$(printenv "$required_var" || true)" - if [ -z "$value" ]; then - echo "Missing repository variable: $required_var" - exit 1 - fi - done - - for required_secret in DEPLOY_GHCR_USERNAME DEPLOY_GHCR_TOKEN PROD_JWT_SECRET BASALTPASS_DATABASE_DSN; do - case "$required_secret" in - DEPLOY_GHCR_USERNAME) - value="${{ secrets.DEPLOY_GHCR_USERNAME }}" - ;; - DEPLOY_GHCR_TOKEN) - value="${{ secrets.DEPLOY_GHCR_TOKEN }}" - ;; - PROD_JWT_SECRET) - value="${{ secrets.PROD_JWT_SECRET }}" - ;; - BASALTPASS_DATABASE_DSN) - value="${{ secrets.BASALTPASS_DATABASE_DSN }}" - ;; - esac - - if [ -z "$value" ]; then - echo "Missing repository secret: $required_secret" - exit 1 - fi - done - - if [ -z "${{ secrets.DEPLOY_SSH_PRIVATE_KEY }}" ] && [ -z "${{ secrets.DEPLOY_SSH_PASSWORD }}" ]; then - echo "Either DEPLOY_SSH_PRIVATE_KEY or DEPLOY_SSH_PASSWORD is required." - exit 1 - fi - - - name: Parse deploy version - id: version - run: | - RAW_TAG="${GITHUB_REF_NAME}" - VERSION="${RAW_TAG#deploy-prod-v}" - echo "raw_tag=${RAW_TAG}" >> "$GITHUB_OUTPUT" - echo "version=${VERSION}" >> "$GITHUB_OUTPUT" - echo "backend_image=ghcr.io/${GHCR_NAMESPACE}/basaltpass-backend:${VERSION}" >> "$GITHUB_OUTPUT" - echo "frontend_image=ghcr.io/${GHCR_NAMESPACE}/basaltpass-frontend:${VERSION}" >> "$GITHUB_OUTPUT" - - - name: Create GitHub deployment - id: create_deployment - uses: actions/github-script@v7 - with: - script: | - const response = await github.rest.repos.createDeployment({ - owner: context.repo.owner, - repo: context.repo.repo, - ref: context.ref, - environment: 'production', - auto_merge: false, - required_contexts: [] - }); - core.setOutput('deployment_id', String(response.data.id)); - - - name: Set up Docker Buildx - uses: docker/setup-buildx-action@v3 - - - name: Log in to GHCR - uses: docker/login-action@v3 - with: - registry: ghcr.io - username: ${{ github.actor }} - password: ${{ secrets.GITHUB_TOKEN }} - - - name: Build and push backend image - uses: docker/build-push-action@v6 - with: - context: . - file: ./backend.Dockerfile - push: true - tags: | - ${{ steps.version.outputs.backend_image }} - ghcr.io/${{ env.GHCR_NAMESPACE }}/basaltpass-backend:latest - - - name: Build and push frontend image - uses: docker/build-push-action@v6 - with: - context: . - file: ./frontend.Dockerfile - push: true - build-args: | - VITE_API_BASE=${{ vars.SERVICE_URL }} - VITE_APP_VERSION=${{ steps.version.outputs.raw_tag }} - VITE_CONSOLE_USER_URL=${{ vars.SERVICE_URL }} - VITE_CONSOLE_TENANT_URL=${{ vars.SERVICE_URL }} - VITE_CONSOLE_ADMIN_URL=${{ vars.SERVICE_URL }} - tags: | - ${{ steps.version.outputs.frontend_image }} - ghcr.io/${{ env.GHCR_NAMESPACE }}/basaltpass-frontend:latest - - - name: Create deployment env file - run: | - cat > .deploy.env <> .deploy.env - fi - if [ -n "${BASALTPASS_EMAIL_FROM}" ]; then - echo "BASALTPASS_EMAIL_FROM=${BASALTPASS_EMAIL_FROM}" >> .deploy.env - fi - if [ -n "${{ secrets.BASALTPASS_EMAIL_SMTP_HOST }}" ]; then - echo "BASALTPASS_EMAIL_SMTP_HOST=${{ secrets.BASALTPASS_EMAIL_SMTP_HOST }}" >> .deploy.env - fi - if [ -n "${BASALTPASS_EMAIL_SMTP_PORT}" ]; then - echo "BASALTPASS_EMAIL_SMTP_PORT=${BASALTPASS_EMAIL_SMTP_PORT}" >> .deploy.env - fi - if [ -n "${{ secrets.BASALTPASS_EMAIL_SMTP_USERNAME }}" ]; then - echo "BASALTPASS_EMAIL_SMTP_USERNAME=${{ secrets.BASALTPASS_EMAIL_SMTP_USERNAME }}" >> .deploy.env - fi - if [ -n "${{ secrets.BASALTPASS_EMAIL_SMTP_PASSWORD }}" ]; then - echo "BASALTPASS_EMAIL_SMTP_PASSWORD=${{ secrets.BASALTPASS_EMAIL_SMTP_PASSWORD }}" >> .deploy.env - fi - if [ -n "${BASALTPASS_EMAIL_SMTP_USE_TLS}" ]; then - echo "BASALTPASS_EMAIL_SMTP_USE_TLS=${BASALTPASS_EMAIL_SMTP_USE_TLS}" >> .deploy.env - fi - if [ -n "${BASALTPASS_EMAIL_SMTP_USE_SSL}" ]; then - echo "BASALTPASS_EMAIL_SMTP_USE_SSL=${BASALTPASS_EMAIL_SMTP_USE_SSL}" >> .deploy.env - fi - if [ -n "${BASALTPASS_EMAIL_SMTP_SKIP_CERT_VERIFY}" ]; then - echo "BASALTPASS_EMAIL_SMTP_SKIP_CERT_VERIFY=${BASALTPASS_EMAIL_SMTP_SKIP_CERT_VERIFY}" >> .deploy.env - fi - if [ -n "${BASALTPASS_EMAIL_AWS_SES_REGION}" ]; then - echo "BASALTPASS_EMAIL_AWS_SES_REGION=${BASALTPASS_EMAIL_AWS_SES_REGION}" >> .deploy.env - fi - if [ -n "${{ secrets.BASALTPASS_EMAIL_AWS_SES_ACCESS_KEY_ID }}" ]; then - echo "BASALTPASS_EMAIL_AWS_SES_ACCESS_KEY_ID=${{ secrets.BASALTPASS_EMAIL_AWS_SES_ACCESS_KEY_ID }}" >> .deploy.env - fi - if [ -n "${{ secrets.BASALTPASS_EMAIL_AWS_SES_SECRET_ACCESS_KEY }}" ]; then - echo "BASALTPASS_EMAIL_AWS_SES_SECRET_ACCESS_KEY=${{ secrets.BASALTPASS_EMAIL_AWS_SES_SECRET_ACCESS_KEY }}" >> .deploy.env - fi - if [ -n "${{ secrets.BASALTPASS_EMAIL_AWS_SES_CONFIGURATION_SET }}" ]; then - echo "BASALTPASS_EMAIL_AWS_SES_CONFIGURATION_SET=${{ secrets.BASALTPASS_EMAIL_AWS_SES_CONFIGURATION_SET }}" >> .deploy.env - fi - - - name: Prepare SSH access - run: | - mkdir -p ~/.ssh - chmod 700 ~/.ssh - ssh-keyscan -p "${DEPLOY_PORT}" -H "${DEPLOY_HOST}" >> ~/.ssh/known_hosts - - if [ -n "${{ secrets.DEPLOY_SSH_PRIVATE_KEY }}" ]; then - printf '%s\n' "${{ secrets.DEPLOY_SSH_PRIVATE_KEY }}" > ~/.ssh/id_ed25519 - chmod 600 ~/.ssh/id_ed25519 - else - sudo apt-get update - sudo apt-get install -y sshpass - fi - - - name: Upload compose and env files - env: - DEPLOY_SSH_PASSWORD: ${{ secrets.DEPLOY_SSH_PASSWORD }} - run: | - if [ -n "${{ secrets.DEPLOY_SSH_PRIVATE_KEY }}" ]; then - ssh -i ~/.ssh/id_ed25519 -p "${DEPLOY_PORT}" "${DEPLOY_USER}@${DEPLOY_HOST}" "mkdir -p '${DEPLOY_PATH}'" - scp -i ~/.ssh/id_ed25519 -P "${DEPLOY_PORT}" ./deploy/docker-compose.prod.yml "${DEPLOY_USER}@${DEPLOY_HOST}:${DEPLOY_PATH}/docker-compose.yml" - scp -i ~/.ssh/id_ed25519 -P "${DEPLOY_PORT}" ./.deploy.env "${DEPLOY_USER}@${DEPLOY_HOST}:${DEPLOY_PATH}/.env" - else - export SSHPASS="${DEPLOY_SSH_PASSWORD}" - sshpass -e ssh -p "${DEPLOY_PORT}" -o StrictHostKeyChecking=yes "${DEPLOY_USER}@${DEPLOY_HOST}" "mkdir -p '${DEPLOY_PATH}'" - sshpass -e scp -P "${DEPLOY_PORT}" -o StrictHostKeyChecking=yes ./deploy/docker-compose.prod.yml "${DEPLOY_USER}@${DEPLOY_HOST}:${DEPLOY_PATH}/docker-compose.yml" - sshpass -e scp -P "${DEPLOY_PORT}" -o StrictHostKeyChecking=yes ./.deploy.env "${DEPLOY_USER}@${DEPLOY_HOST}:${DEPLOY_PATH}/.env" - fi - - - name: Deploy on Docker server - env: - BACKEND_IMAGE: ${{ steps.version.outputs.backend_image }} - FRONTEND_IMAGE: ${{ steps.version.outputs.frontend_image }} - DEPLOY_GHCR_USERNAME: ${{ secrets.DEPLOY_GHCR_USERNAME }} - DEPLOY_GHCR_TOKEN: ${{ secrets.DEPLOY_GHCR_TOKEN }} - DEPLOY_SSH_PASSWORD: ${{ secrets.DEPLOY_SSH_PASSWORD }} - run: | - BACKEND_IMAGE_ESCAPED="$(printf '%q' "${BACKEND_IMAGE}")" - FRONTEND_IMAGE_ESCAPED="$(printf '%q' "${FRONTEND_IMAGE}")" - DEPLOY_GHCR_USERNAME_ESCAPED="$(printf '%q' "${DEPLOY_GHCR_USERNAME}")" - DEPLOY_GHCR_TOKEN_ESCAPED="$(printf '%q' "${DEPLOY_GHCR_TOKEN}")" - - REMOTE_COMMAND="export BACKEND_IMAGE=${BACKEND_IMAGE_ESCAPED} FRONTEND_IMAGE=${FRONTEND_IMAGE_ESCAPED} DEPLOY_GHCR_USERNAME=${DEPLOY_GHCR_USERNAME_ESCAPED} DEPLOY_GHCR_TOKEN=${DEPLOY_GHCR_TOKEN_ESCAPED} && cd '${DEPLOY_PATH}' && echo \"\$DEPLOY_GHCR_TOKEN\" | docker login ghcr.io -u \"\$DEPLOY_GHCR_USERNAME\" --password-stdin && docker compose pull && docker compose up -d --remove-orphans" - - if [ -n "${{ secrets.DEPLOY_SSH_PRIVATE_KEY }}" ]; then - ssh -i ~/.ssh/id_ed25519 -p "${DEPLOY_PORT}" "${DEPLOY_USER}@${DEPLOY_HOST}" "${REMOTE_COMMAND}" - else - export SSHPASS="${DEPLOY_SSH_PASSWORD}" - sshpass -e ssh -p "${DEPLOY_PORT}" -o StrictHostKeyChecking=yes "${DEPLOY_USER}@${DEPLOY_HOST}" "${REMOTE_COMMAND}" - fi - - - name: Collect remote diagnostics on deploy failure - if: failure() - env: - DEPLOY_SSH_PASSWORD: ${{ secrets.DEPLOY_SSH_PASSWORD }} - run: | - REMOTE_DEBUG_COMMAND="cd '${DEPLOY_PATH}' && docker compose ps -a && echo '--- backend logs ---' && docker compose logs backend --tail 200" - - if [ -n "${{ secrets.DEPLOY_SSH_PRIVATE_KEY }}" ]; then - ssh -i ~/.ssh/id_ed25519 -p "${DEPLOY_PORT}" "${DEPLOY_USER}@${DEPLOY_HOST}" "${REMOTE_DEBUG_COMMAND}" || true - else - export SSHPASS="${DEPLOY_SSH_PASSWORD}" - sshpass -e ssh -p "${DEPLOY_PORT}" -o StrictHostKeyChecking=yes "${DEPLOY_USER}@${DEPLOY_HOST}" "${REMOTE_DEBUG_COMMAND}" || true - fi - - - name: Verify health check on server - env: - DEPLOY_SSH_PASSWORD: ${{ secrets.DEPLOY_SSH_PASSWORD }} - run: | - DEPLOY_PATH_ESCAPED="$(printf '%q' "${DEPLOY_PATH}")" - REMOTE_HEALTHCHECK_COMMAND="$(cat < 0)' "${CONFIG_FILE}" >/dev/null + - name: Set up Go + uses: actions/setup-go@v5 + with: + go-version-file: basaltpass-backend/go.mod + cache-dependency-path: basaltpass-backend/go.sum - - name: Build matrix - id: matrix - run: | - MATRIX_JSON=$(jq -c '{include: [.images[] | {name, context, dockerfile, build_args: (.build_args // "")}]}' "${CONFIG_FILE}") - echo "matrix=${MATRIX_JSON}" >> "$GITHUB_OUTPUT" + - name: Test backend + working-directory: basaltpass-backend + env: + JWT_SECRET: test-secret-for-ci + run: go test ./... build_and_push: - needs: prepare + needs: test runs-on: ubuntu-latest - strategy: - fail-fast: false - matrix: ${{ fromJSON(needs.prepare.outputs.matrix) }} + outputs: + backend_image: ${{ steps.meta.outputs.backend_image }} + frontend_image: ${{ steps.meta.outputs.frontend_image }} steps: - name: Checkout uses: actions/checkout@v4 - - name: Resolve image prefix + - name: Resolve image tags id: meta shell: bash run: | - PROJECT_SLUG=$(jq -r '.project' "${CONFIG_FILE}") - OWNER_LC=$(echo "${GITHUB_REPOSITORY_OWNER}" | tr '[:upper:]' '[:lower:]') - echo "image_prefix=ghcr.io/${OWNER_LC}/${PROJECT_SLUG}" >> "$GITHUB_OUTPUT" + OWNER_LC="$(echo "${GITHUB_REPOSITORY_OWNER}" | tr '[:upper:]' '[:lower:]')" + BACKEND_IMAGE="ghcr.io/${OWNER_LC}/basaltpass-backend" + FRONTEND_IMAGE="ghcr.io/${OWNER_LC}/basaltpass-frontend" + echo "backend_image=${BACKEND_IMAGE}" >> "$GITHUB_OUTPUT" + echo "frontend_image=${FRONTEND_IMAGE}" >> "$GITHUB_OUTPUT" + { + echo "backend_tags<> "$GITHUB_OUTPUT" - name: Set up Docker Buildx uses: docker/setup-buildx-action@v3 @@ -62,20 +80,31 @@ jobs: username: ${{ github.repository_owner }} password: ${{ secrets.GITHUB_TOKEN }} - - name: Build and push ${{ matrix.name }} + - name: Build and push backend + uses: docker/build-push-action@v6 + with: + context: . + file: ./backend.Dockerfile + push: true + tags: ${{ steps.meta.outputs.backend_tags }} + + - name: Build and push frontend uses: docker/build-push-action@v6 with: - context: ${{ matrix.context }} - file: ${{ matrix.dockerfile }} + context: . + file: ./frontend.Dockerfile push: true - build-args: ${{ matrix.build_args }} - tags: | - ${{ steps.meta.outputs.image_prefix }}-${{ matrix.name }}:${{ github.sha }} - ${{ steps.meta.outputs.image_prefix }}-${{ matrix.name }}:latest + build-args: | + VITE_API_BASE=${{ vars.BASALTPASS_PUBLIC_URL }} + VITE_APP_VERSION=${{ github.ref_name }} + VITE_CONSOLE_USER_URL=${{ vars.BASALTPASS_PUBLIC_URL }} + VITE_CONSOLE_TENANT_URL=${{ vars.BASALTPASS_PUBLIC_URL }}/tenant + VITE_CONSOLE_ADMIN_URL=${{ vars.BASALTPASS_PUBLIC_URL }}/admin + tags: ${{ steps.meta.outputs.frontend_tags }} - deploy_k3s: + deploy_beancs_system: needs: build_and_push - if: startsWith(github.ref, 'refs/tags/v') + if: startsWith(github.ref, 'refs/tags/beancs-') runs-on: ubuntu-latest steps: - name: Checkout @@ -87,9 +116,9 @@ jobs: shell: bash run: | git fetch origin deploy --depth=1 - TAG_SHA=$(git rev-list -n 1 "${GITHUB_REF}") + TAG_SHA="$(git rev-list -n 1 "${GITHUB_REF}")" if ! git merge-base --is-ancestor "${TAG_SHA}" "origin/deploy"; then - echo "Tag commit is not contained in origin/deploy, skip deployment." + echo "Tag commit is not contained in origin/deploy." exit 1 fi @@ -104,47 +133,208 @@ jobs: printf '%s' "${{ secrets.KUBE_CONFIG }}" > ~/.kube/config chmod 600 ~/.kube/config - - name: Deploy workloads + - name: Ensure namespace and pull secret shell: bash run: | - PROJECT_SLUG=$(jq -r '.project' "${CONFIG_FILE}") - NAMESPACE=$(jq -r '.namespace // .project' "${CONFIG_FILE}") - OWNER_LC=$(echo "${GITHUB_REPOSITORY_OWNER}" | tr '[:upper:]' '[:lower:]') - GHCR_PASSWORD="${{ secrets.GITHUB_TOKEN }}" - + OWNER_LC="$(echo "${GITHUB_REPOSITORY_OWNER}" | tr '[:upper:]' '[:lower:]')" kubectl get namespace "${NAMESPACE}" >/dev/null 2>&1 || kubectl create namespace "${NAMESPACE}" - kubectl -n "${NAMESPACE}" create secret docker-registry ghcr-creds \ --docker-server=ghcr.io \ --docker-username="${OWNER_LC}" \ - --docker-password="${GHCR_PASSWORD}" \ + --docker-password="${{ secrets.GHCR_PAT || secrets.GITHUB_TOKEN }}" \ --dry-run=client -o yaml | kubectl apply -f - - jq -c '.images[]' "${CONFIG_FILE}" | while read -r item; do - NAME=$(jq -r '.name' <<< "${item}") - KIND=$(jq -r '.kind // "deployment"' <<< "${item}") - PORT=$(jq -r '.port // empty' <<< "${item}") - REPLICAS=$(jq -r '.replicas // 1' <<< "${item}") - IMAGE="ghcr.io/${OWNER_LC}/${PROJECT_SLUG}-${NAME}:${GITHUB_SHA}" - - if [ "${KIND}" = "job" ]; then - kubectl -n "${NAMESPACE}" delete job "${PROJECT_SLUG}-${NAME}" --ignore-not-found=true - kubectl -n "${NAMESPACE}" create job "${PROJECT_SLUG}-${NAME}" \ - --image="${IMAGE}" \ - --dry-run=client -o yaml | kubectl apply -f - - continue - fi + - name: Configure BasaltPass backend secret + shell: bash + run: | + kubectl -n "${NAMESPACE}" create secret generic basaltpass-env \ + --from-literal=JWT_SECRET="${{ secrets.PROD_JWT_SECRET }}" \ + --from-literal=BASALTPASS_ENV=production \ + --from-literal=BASALTPASS_SERVER_ADDRESS=:8101 \ + --from-literal=BASALTPASS_DATABASE_DRIVER="${{ vars.BASALTPASS_DATABASE_DRIVER || 'postgres' }}" \ + --from-literal=BASALTPASS_DATABASE_DSN="${{ secrets.BASALTPASS_DATABASE_DSN }}" \ + --from-literal=BASALTPASS_CORS_ALLOW_ORIGINS="${{ vars.BASALTPASS_CORS_ALLOW_ORIGINS }}" \ + --from-literal=BASALTPASS_EMAIL_PROVIDER="${{ vars.BASALTPASS_EMAIL_PROVIDER }}" \ + --from-literal=BASALTPASS_EMAIL_FROM="${{ vars.BASALTPASS_EMAIL_FROM }}" \ + --from-literal=BASALTPASS_EMAIL_AWS_SES_REGION="${{ secrets.BASALTPASS_EMAIL_AWS_SES_REGION }}" \ + --from-literal=BASALTPASS_EMAIL_AWS_SES_ACCESS_KEY_ID="${{ secrets.BASALTPASS_EMAIL_AWS_SES_ACCESS_KEY_ID }}" \ + --from-literal=BASALTPASS_EMAIL_AWS_SES_SECRET_ACCESS_KEY="${{ secrets.BASALTPASS_EMAIL_AWS_SES_SECRET_ACCESS_KEY }}" \ + --dry-run=client -o yaml | kubectl apply -f - - kubectl -n "${NAMESPACE}" create deployment "${PROJECT_SLUG}-${NAME}" \ - --image="${IMAGE}" \ - --dry-run=client -o yaml | kubectl apply -f - - kubectl -n "${NAMESPACE}" scale deployment/${PROJECT_SLUG}-${NAME} --replicas="${REPLICAS}" + - name: Deploy BasaltPass workloads + shell: bash + run: | + BACKEND_IMAGE="${{ needs.build_and_push.outputs.backend_image }}:${GITHUB_SHA}" + FRONTEND_IMAGE="${{ needs.build_and_push.outputs.frontend_image }}:${GITHUB_SHA}" - if [ -n "${PORT}" ]; then - kubectl -n "${NAMESPACE}" create service clusterip "${PROJECT_SLUG}-${NAME}" \ - --tcp="${PORT}:${PORT}" \ - --dry-run=client -o yaml | kubectl apply -f - - fi + cat < 0 { + q.Set("app_tenant_id", strconv.FormatUint(uint64(appTenantID), 10)) + } + } + if decision != nil { + if decision.JoinRequired { + q.Set("current_user_join_required", "true") + } + if decision.TenantID > 0 { + q.Set("decision_tenant_id", strconv.FormatUint(uint64(decision.TenantID), 10)) + } + } if client != nil { // Prefer app display fields when available. if strings.TrimSpace(client.App.Name) != "" { @@ -218,6 +235,8 @@ func ConsentHandler(c *fiber.Ctx) error { scope := c.FormValue("scope") codeChallenge := c.FormValue("code_challenge") codeChallengeMethod := c.FormValue("code_challenge_method") + selectedAccessToken := strings.TrimSpace(c.FormValue("selected_access_token")) + joinTenant := strings.EqualFold(strings.TrimSpace(c.FormValue("join_tenant")), "true") || c.FormValue("join_tenant") == "1" action := c.FormValue("action") // "allow" 或 "deny" if action != "allow" { @@ -225,6 +244,17 @@ func ConsentHandler(c *fiber.Ctx) error { return redirectWithErrorIfAllowed(c, clientID, redirectURI, "access_denied", state) } + if selectedAccessToken != "" { + selectedUserID, ok := parseUserIDFromJWT(selectedAccessToken) + if !ok || selectedUserID == 0 { + return c.Status(fiber.StatusUnauthorized).JSON(fiber.Map{ + "error": "invalid_selected_session", + "error_description": "Selected account session is invalid or expired", + }) + } + userID = selectedUserID + } + // 构建授权请求 req := &AuthorizeRequest{ ClientID: clientID, @@ -242,14 +272,36 @@ func ConsentHandler(c *fiber.Ctx) error { return redirectWithErrorIfAllowed(c, clientID, redirectURI, err.Error(), state) } - // 验证用户是否属于该租户 - if err := oauthServerService.ValidateUserTenant(userID, client); err != nil { + decision, err := oauthServerService.EvaluateUserTenantAuthorization(userID, client) + if err != nil { return c.Status(fiber.StatusForbidden).JSON(fiber.Map{ - "error": "tenant_mismatch", - "error_description": "User does not belong to the tenant of this application", + "error": "tenant_context_error", + "error_description": err.Error(), }) } + if !decision.Allowed { + if decision.JoinRequired { + if !joinTenant { + return c.Status(fiber.StatusForbidden).JSON(fiber.Map{ + "error": "join_confirmation_required", + "error_description": "Global account must confirm tenant join before authorization", + }) + } + if err := oauthServerService.EnsureUserTenantIdentity(userID, decision.TenantID, model.TenantRoleMember); err != nil { + return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{ + "error": "identity_join_failed", + "error_description": "Failed to join tenant identity", + }) + } + } else { + return c.Status(fiber.StatusForbidden).JSON(fiber.Map{ + "error": "tenant_mismatch", + "error_description": "User does not belong to the tenant of this application", + }) + } + } + // 生成授权码 code, err := oauthServerService.GenerateAuthorizationCode(userID, req, client) if err != nil { diff --git a/basaltpass-backend/internal/handler/public/oauth/server_service.go b/basaltpass-backend/internal/handler/public/oauth/server_service.go index 2da97463..6de31878 100644 --- a/basaltpass-backend/internal/handler/public/oauth/server_service.go +++ b/basaltpass-backend/internal/handler/public/oauth/server_service.go @@ -23,6 +23,13 @@ type OAuthServerService struct { db *gorm.DB } +type UserTenantAuthorizationDecision struct { + Allowed bool + JoinRequired bool + TenantID uint + Reason string +} + // NewOAuthServerService 创建新的OAuth2服务器服务 func NewOAuthServerService() *OAuthServerService { return &OAuthServerService{ @@ -597,29 +604,90 @@ func (s *OAuthServerService) RevokeToken(token string) error { return nil } -// ValidateUserTenant 验证用户是否属于应用所在的租户 -func (s *OAuthServerService) ValidateUserTenant(userID uint, client *model.OAuthClient) error { - // 获取用户信息 - var user model.User - if err := s.db.Select("id", "tenant_id", "is_system_admin").First(&user, userID).Error; err != nil { - return errors.New("user_not_found") +func (s *OAuthServerService) EnsureUserTenantIdentity(userID, tenantID uint, role model.TenantRole) error { + if userID == 0 || tenantID == 0 { + return errors.New("invalid_identity_context") + } + if role == "" { + role = model.TenantRoleMember } - // 系统管理员可以访问所有租户的应用 - if user.IsSystemAdmin != nil && *user.IsSystemAdmin { + var existing model.TenantUser + err := s.db.Where("user_id = ? AND tenant_id = ?", userID, tenantID).First(&existing).Error + if err == nil { return nil } + if !errors.Is(err, gorm.ErrRecordNotFound) { + return err + } + + return s.db.Create(&model.TenantUser{ + UserID: userID, + TenantID: tenantID, + Role: role, + }).Error +} + +func (s *OAuthServerService) EvaluateUserTenantAuthorization(userID uint, client *model.OAuthClient) (*UserTenantAuthorizationDecision, error) { + var user model.User + if err := s.db.Select("id", "tenant_id", "is_system_admin").First(&user, userID).Error; err != nil { + return nil, errors.New("user_not_found") + } - // 获取应用所属租户ID(支持历史脏数据的兜底恢复) tenantID := s.resolveClientTenantID(client) if tenantID == 0 { - return errors.New("app_tenant_not_found") + return nil, errors.New("app_tenant_not_found") } - // 验证用户的tenant_id是否与应用的tenant_id匹配 - if user.TenantID != tenantID { - return errors.New("tenant_mismatch") + if user.IsSystemAdmin != nil && *user.IsSystemAdmin { + return &UserTenantAuthorizationDecision{Allowed: true, TenantID: tenantID}, nil } - return nil + if user.TenantID == tenantID { + return &UserTenantAuthorizationDecision{Allowed: true, TenantID: tenantID}, nil + } + + if user.TenantID == 0 { + var membershipCount int64 + if err := s.db.Model(&model.TenantUser{}). + Where("user_id = ? AND tenant_id = ?", userID, tenantID). + Count(&membershipCount).Error; err != nil { + return nil, err + } + if membershipCount > 0 { + return &UserTenantAuthorizationDecision{Allowed: true, TenantID: tenantID}, nil + } + + return &UserTenantAuthorizationDecision{ + Allowed: false, + JoinRequired: true, + TenantID: tenantID, + Reason: "join_required", + }, nil + } + + return &UserTenantAuthorizationDecision{ + Allowed: false, + JoinRequired: false, + TenantID: tenantID, + Reason: "tenant_mismatch", + }, nil +} + +// ValidateUserTenant 验证用户是否属于应用所在的租户 +func (s *OAuthServerService) ValidateUserTenant(userID uint, client *model.OAuthClient) error { + decision, err := s.EvaluateUserTenantAuthorization(userID, client) + if err != nil { + return err + } + if decision.Allowed { + return nil + } + if decision.JoinRequired { + return errors.New("join_required") + } + if decision.Reason != "" { + return errors.New(decision.Reason) + } + return errors.New("tenant_mismatch") } diff --git a/basaltpass-backend/internal/handler/public/oauth/server_tenant_authorization_test.go b/basaltpass-backend/internal/handler/public/oauth/server_tenant_authorization_test.go new file mode 100644 index 00000000..9324f31c --- /dev/null +++ b/basaltpass-backend/internal/handler/public/oauth/server_tenant_authorization_test.go @@ -0,0 +1,114 @@ +package oauth + +import ( + "testing" + "time" + + "basaltpass-backend/internal/common" + "basaltpass-backend/internal/model" + + "github.com/glebarez/sqlite" + "github.com/golang-jwt/jwt/v5" + "gorm.io/gorm" +) + +func setupOAuthTenantTestDB(t *testing.T) *gorm.DB { + t.Helper() + + db, err := gorm.Open(sqlite.Open("file::memory:?cache=shared"), &gorm.Config{}) + if err != nil { + t.Fatalf("open sqlite failed: %v", err) + } + + if err := db.AutoMigrate(&model.User{}, &model.Tenant{}, &model.TenantUser{}, &model.App{}, &model.OAuthClient{}); err != nil { + t.Fatalf("auto migrate failed: %v", err) + } + + common.SetDBForTest(db) + return db +} + +func TestEvaluateUserTenantAuthorizationJoinThenAllow(t *testing.T) { + db := setupOAuthTenantTestDB(t) + + tenant := model.Tenant{Name: "Tenant-A", Code: "tenant-a", Status: model.TenantStatusActive} + if err := db.Create(&tenant).Error; err != nil { + t.Fatalf("create tenant failed: %v", err) + } + + globalUser := model.User{TenantID: 0, Email: "global-join@example.com", PasswordHash: "x"} + if err := db.Create(&globalUser).Error; err != nil { + t.Fatalf("create user failed: %v", err) + } + + app := model.App{TenantID: tenant.ID, Name: "Join App", Status: model.AppStatusActive} + if err := db.Create(&app).Error; err != nil { + t.Fatalf("create app failed: %v", err) + } + + client := model.OAuthClient{ + AppID: app.ID, + ClientID: "client-join", + ClientSecret: "secret", + RedirectURIs: "https://example.com/callback", + IsActive: true, + CreatedBy: globalUser.ID, + } + if err := db.Create(&client).Error; err != nil { + t.Fatalf("create client failed: %v", err) + } + + svc := NewOAuthServerService() + + decision, err := svc.EvaluateUserTenantAuthorization(globalUser.ID, &client) + if err != nil { + t.Fatalf("evaluate authorization failed: %v", err) + } + if decision.Allowed { + t.Fatalf("expected not allowed before join") + } + if !decision.JoinRequired { + t.Fatalf("expected join_required=true for global user") + } + if decision.TenantID != tenant.ID { + t.Fatalf("expected tenant id %d, got %d", tenant.ID, decision.TenantID) + } + + if err := svc.EnsureUserTenantIdentity(globalUser.ID, tenant.ID, model.TenantRoleMember); err != nil { + t.Fatalf("ensure tenant identity failed: %v", err) + } + + decision, err = svc.EvaluateUserTenantAuthorization(globalUser.ID, &client) + if err != nil { + t.Fatalf("re-evaluate authorization failed: %v", err) + } + if !decision.Allowed { + t.Fatalf("expected allowed after join") + } + if decision.JoinRequired { + t.Fatalf("expected join_required=false after membership creation") + } +} + +func TestParseUserIDFromJWTForSelectedAccessToken(t *testing.T) { + token := jwt.NewWithClaims(jwt.SigningMethodHS256, jwt.MapClaims{ + "sub": float64(42), + "exp": time.Now().Add(5 * time.Minute).Unix(), + }) + signed, err := token.SignedString(common.MustJWTSecret()) + if err != nil { + t.Fatalf("sign token failed: %v", err) + } + + uid, ok := parseUserIDFromJWT(signed) + if !ok { + t.Fatalf("expected parse success for valid selected_access_token") + } + if uid != 42 { + t.Fatalf("expected user id 42, got %d", uid) + } + + if _, ok := parseUserIDFromJWT("not-a-token"); ok { + t.Fatalf("expected parse failure for invalid token") + } +} diff --git a/basaltpass-backend/internal/handler/public/order/handler.go b/basaltpass-backend/internal/handler/public/order/handler.go index fb6c0238..401d11e0 100644 --- a/basaltpass-backend/internal/handler/public/order/handler.go +++ b/basaltpass-backend/internal/handler/public/order/handler.go @@ -5,6 +5,7 @@ import ( "strconv" "basaltpass-backend/internal/common" + "basaltpass-backend/internal/model" "github.com/gofiber/fiber/v2" ) @@ -20,11 +21,46 @@ func InitOrderHandler() { } } +func resolveOrderTenantID(c *fiber.Ctx) (uint64, error) { + if tenantLocal, ok := c.Locals("tenantID").(uint); ok && tenantLocal > 0 { + return uint64(tenantLocal), nil + } + + userID, ok := c.Locals("userID").(uint) + if !ok || userID == 0 { + return 0, fiber.NewError(fiber.StatusUnauthorized, "未登录") + } + + var user model.User + if err := common.DB().Select("id", "tenant_id").First(&user, userID).Error; err != nil { + return 0, fiber.NewError(fiber.StatusForbidden, "无法识别当前租户") + } + + if user.TenantID > 0 { + return uint64(user.TenantID), nil + } + + var tenantUser model.TenantUser + if err := common.DB().Select("tenant_id").Where("user_id = ?", userID).Order("created_at ASC").First(&tenantUser).Error; err != nil { + return 0, fiber.NewError(fiber.StatusForbidden, "当前用户没有关联租户") + } + + if tenantUser.TenantID == 0 { + return 0, fiber.NewError(fiber.StatusForbidden, "当前用户没有关联租户") + } + + return uint64(tenantUser.TenantID), nil +} + // CreateOrderHandler POST /orders - 创建订单 func CreateOrderHandler(c *fiber.Ctx) error { InitOrderHandler() userID := c.Locals("userID").(uint) + tenantID, tenantErr := resolveOrderTenantID(c) + if tenantErr != nil { + return c.Status(fiber.StatusForbidden).JSON(fiber.Map{"error": tenantErr.Error()}) + } var req order2.CreateOrderRequest if err := c.BodyParser(&req); err != nil { @@ -40,7 +76,7 @@ func CreateOrderHandler(c *fiber.Ctx) error { }) } - order, err := orderService.CreateOrder(&req) + order, err := orderService.CreateOrder(&req, tenantID) if err != nil { return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{ "error": err.Error(), @@ -59,6 +95,10 @@ func GetOrderHandler(c *fiber.Ctx) error { InitOrderHandler() userID := c.Locals("userID").(uint) + tenantID, tenantErr := resolveOrderTenantID(c) + if tenantErr != nil { + return c.Status(fiber.StatusForbidden).JSON(fiber.Map{"error": tenantErr.Error()}) + } orderIDStr := c.Params("id") orderID, err := strconv.ParseUint(orderIDStr, 10, 32) @@ -70,7 +110,7 @@ func GetOrderHandler(c *fiber.Ctx) error { activate := c.Query("activate") == "1" || c.Query("activate") == "true" - order, err := orderService.GetOrder(userID, uint(orderID), activate) + order, err := orderService.GetOrder(userID, uint(orderID), activate, tenantID) if err != nil { return c.Status(fiber.StatusNotFound).JSON(fiber.Map{ "error": err.Error(), @@ -88,9 +128,13 @@ func GetOrderByNumberHandler(c *fiber.Ctx) error { InitOrderHandler() userID := c.Locals("userID").(uint) + tenantID, tenantErr := resolveOrderTenantID(c) + if tenantErr != nil { + return c.Status(fiber.StatusForbidden).JSON(fiber.Map{"error": tenantErr.Error()}) + } orderNumber := c.Params("number") - order, err := orderService.GetOrderByNumber(userID, orderNumber) + order, err := orderService.GetOrderByNumber(userID, orderNumber, tenantID) if err != nil { return c.Status(fiber.StatusNotFound).JSON(fiber.Map{ "error": err.Error(), @@ -108,6 +152,10 @@ func ListOrdersHandler(c *fiber.Ctx) error { InitOrderHandler() userID := c.Locals("userID").(uint) + tenantID, tenantErr := resolveOrderTenantID(c) + if tenantErr != nil { + return c.Status(fiber.StatusForbidden).JSON(fiber.Map{"error": tenantErr.Error()}) + } limitStr := c.Query("limit", "20") limit, _ := strconv.Atoi(limitStr) @@ -115,7 +163,7 @@ func ListOrdersHandler(c *fiber.Ctx) error { limit = 100 } - orders, err := orderService.ListUserOrders(userID, limit) + orders, err := orderService.ListUserOrders(userID, limit, tenantID) if err != nil { return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{ "error": err.Error(), diff --git a/basaltpass-backend/internal/handler/public/payment/handler.go b/basaltpass-backend/internal/handler/public/payment/handler.go index eba5f8a4..22ffccd9 100644 --- a/basaltpass-backend/internal/handler/public/payment/handler.go +++ b/basaltpass-backend/internal/handler/public/payment/handler.go @@ -36,9 +36,35 @@ func requireSuperAdmin(c *fiber.Ctx) error { return nil } +func resolvePaymentTenantID(c *fiber.Ctx) uint { + if tid, ok := c.Locals("tenantID").(uint); ok && tid > 0 { + return tid + } + + userID, ok := c.Locals("userID").(uint) + if !ok || userID == 0 { + return 0 + } + + var user model.User + if err := common.DB().Select("id", "tenant_id").First(&user, userID).Error; err == nil { + if user.TenantID > 0 { + return user.TenantID + } + } + + var membership model.TenantUser + if err := common.DB().Select("tenant_id").Where("user_id = ?", userID).Order("created_at ASC").First(&membership).Error; err == nil { + return membership.TenantID + } + + return 0 +} + // CreatePaymentIntentHandler POST /payment/intents - 创建支付意图 func CreatePaymentIntentHandler(c *fiber.Ctx) error { userID := c.Locals("userID").(uint) + activeTenantID := resolvePaymentTenantID(c) var req payment2.CreatePaymentIntentRequest if err := c.BodyParser(&req); err != nil { @@ -54,7 +80,7 @@ func CreatePaymentIntentHandler(c *fiber.Ctx) error { }) } - paymentIntent, mockResponse, err := payment2.CreatePaymentIntent(userID, req) + paymentIntent, mockResponse, err := payment2.CreatePaymentIntentForTenant(userID, activeTenantID, req) if err != nil { if errors.Is(err, payment2.ErrTenantStripeNotConfigured) { return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{ @@ -80,6 +106,7 @@ func CreatePaymentIntentHandler(c *fiber.Ctx) error { // CreatePaymentSessionHandler POST /payment/sessions - 创建支付会话 func CreatePaymentSessionHandler(c *fiber.Ctx) error { userID := c.Locals("userID").(uint) + activeTenantID := resolvePaymentTenantID(c) var req payment2.CreatePaymentSessionRequest if err := c.BodyParser(&req); err != nil { @@ -88,7 +115,7 @@ func CreatePaymentSessionHandler(c *fiber.Ctx) error { }) } - session, mockResponse, err := payment2.CreatePaymentSession(userID, req) + session, mockResponse, err := payment2.CreatePaymentSessionForTenant(userID, activeTenantID, req) if err != nil { if errors.Is(err, payment2.ErrTenantStripeNotConfigured) { return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{ @@ -109,6 +136,7 @@ func CreatePaymentSessionHandler(c *fiber.Ctx) error { // GetPaymentIntentHandler GET /payment/intents/:id - 获取支付意图 func GetPaymentIntentHandler(c *fiber.Ctx) error { userID := c.Locals("userID").(uint) + activeTenantID := resolvePaymentTenantID(c) paymentIntentIDStr := c.Params("id") paymentIntentID, err := strconv.ParseUint(paymentIntentIDStr, 10, 32) @@ -118,7 +146,7 @@ func GetPaymentIntentHandler(c *fiber.Ctx) error { }) } - paymentIntent, err := payment2.GetPaymentIntent(userID, uint(paymentIntentID)) + paymentIntent, err := payment2.GetPaymentIntentForTenant(userID, uint(paymentIntentID), activeTenantID) if err != nil { return c.Status(fiber.StatusNotFound).JSON(fiber.Map{ "error": "[GetPaymentIntentHandler] Payment intent not found", @@ -131,9 +159,10 @@ func GetPaymentIntentHandler(c *fiber.Ctx) error { // GetPaymentSessionHandler GET /payment/sessions/:session_id - 获取支付会话 func GetPaymentSessionHandler(c *fiber.Ctx) error { userID := c.Locals("userID").(uint) + activeTenantID := resolvePaymentTenantID(c) sessionID := c.Params("session_id") - session, err := payment2.GetPaymentSession(userID, sessionID) + session, err := payment2.GetPaymentSessionForTenant(userID, sessionID, activeTenantID) if err != nil { return c.Status(fiber.StatusNotFound).JSON(fiber.Map{ "error": "Payment session not found", @@ -146,6 +175,7 @@ func GetPaymentSessionHandler(c *fiber.Ctx) error { // ListPaymentIntentsHandler GET /payment/intents - 获取支付意图列表 func ListPaymentIntentsHandler(c *fiber.Ctx) error { userID := c.Locals("userID").(uint) + activeTenantID := resolvePaymentTenantID(c) limitStr := c.Query("limit", "20") limit, _ := strconv.Atoi(limitStr) @@ -153,7 +183,7 @@ func ListPaymentIntentsHandler(c *fiber.Ctx) error { limit = 100 // 限制最大查询数量 } - paymentIntents, err := payment2.ListPaymentIntents(userID, limit) + paymentIntents, err := payment2.ListPaymentIntentsForTenant(userID, activeTenantID, limit) if err != nil { return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{ "error": err.Error(), diff --git a/basaltpass-backend/internal/handler/public/subscription/checkout.go b/basaltpass-backend/internal/handler/public/subscription/checkout.go index 89111eb3..06d10080 100644 --- a/basaltpass-backend/internal/handler/public/subscription/checkout.go +++ b/basaltpass-backend/internal/handler/public/subscription/checkout.go @@ -13,12 +13,13 @@ import ( // CheckoutRequest 订阅结账请求 type CheckoutRequest struct { - UserID uint `json:"user_id" validate:"required"` - PriceID uint `json:"price_id" validate:"required"` - Quantity float64 `json:"quantity,omitempty"` - CouponCode *string `json:"coupon_code,omitempty"` - SuccessURL string `json:"success_url" validate:"required"` - CancelURL string `json:"cancel_url" validate:"required"` + UserID uint `json:"user_id" validate:"required"` + PriceID uint `json:"price_id" validate:"required"` + Quantity float64 `json:"quantity,omitempty"` + CouponCode *string `json:"coupon_code,omitempty"` + SuccessURL string `json:"success_url" validate:"required"` + CancelURL string `json:"cancel_url" validate:"required"` + ActiveTenantID uint64 `json:"-"` } // CheckoutResponse 订阅结账响应 @@ -59,11 +60,30 @@ func (s *CheckoutService) CreateCheckout(req *CheckoutRequest) (*CheckoutRespons } return nil, fmt.Errorf("查询客户失败: %w", err) } - if user.TenantID == 0 { + userTenantID := req.ActiveTenantID + if userTenantID == 0 && user.TenantID > 0 { + userTenantID = uint64(user.TenantID) + } + if userTenantID == 0 { + var membership model.TenantUser + if err := tx.Select("tenant_id").Where("user_id = ?", req.UserID).Order("created_at ASC").First(&membership).Error; err == nil { + userTenantID = uint64(membership.TenantID) + } + } + if userTenantID == 0 { tx.Rollback() return nil, fmt.Errorf("用户未绑定租户") } - userTenantID := uint64(user.TenantID) + + hasIdentity, err := userHasTenantIdentity(tx, req.UserID, userTenantID) + if err != nil { + tx.Rollback() + return nil, fmt.Errorf("校验用户租户身份失败: %w", err) + } + if !hasIdentity { + tx.Rollback() + return nil, fmt.Errorf("用户不属于当前租户") + } // 步骤2: 验证价格 var price model.Price @@ -74,7 +94,7 @@ func (s *CheckoutService) CreateCheckout(req *CheckoutRequest) (*CheckoutRespons } return nil, fmt.Errorf("查询价格失败: %w", err) } - if price.TenantID != nil && *price.TenantID != userTenantID { + if price.TenantID == nil || *price.TenantID != userTenantID { tx.Rollback() return nil, fmt.Errorf("价格不属于当前租户") } @@ -86,7 +106,7 @@ func (s *CheckoutService) CreateCheckout(req *CheckoutRequest) (*CheckoutRespons if req.CouponCode != nil && *req.CouponCode != "" { var c model.Coupon - if err := tx.Where("code = ? AND is_active = true AND (tenant_id IS NULL OR tenant_id = ?)", *req.CouponCode, userTenantID).First(&c).Error; err != nil { + if err := tx.Where("code = ? AND is_active = true AND tenant_id = ?", *req.CouponCode, userTenantID).First(&c).Error; err != nil { tx.Rollback() if errors.Is(err, gorm.ErrRecordNotFound) { return nil, fmt.Errorf("优惠券不存在或已失效") @@ -269,10 +289,11 @@ func (s *CheckoutService) CreateCheckout(req *CheckoutRequest) (*CheckoutRespons "invoice_id": invoice.ID, "payment_id": paymentRecord.ID, "source": "subscription_checkout", + "tenant_id": userTenantID, }, } - paymentIntent, _, err := payment.CreatePaymentIntent(req.UserID, paymentIntentReq) + paymentIntent, _, err := payment.CreatePaymentIntentForTenant(req.UserID, uint(userTenantID), paymentIntentReq) if err != nil { return nil, fmt.Errorf("创建支付意图失败: %w", err) } @@ -290,7 +311,7 @@ func (s *CheckoutService) CreateCheckout(req *CheckoutRequest) (*CheckoutRespons UserEmail: user.Email, } - paymentSession, sessionStripeResponse, err := payment.CreatePaymentSession(req.UserID, sessionReq) + paymentSession, sessionStripeResponse, err := payment.CreatePaymentSessionForTenant(req.UserID, uint(userTenantID), sessionReq) if err != nil { return nil, fmt.Errorf("创建支付会话失败: %w", err) } diff --git a/basaltpass-backend/internal/handler/public/subscription/checkout_handler.go b/basaltpass-backend/internal/handler/public/subscription/checkout_handler.go index baf64062..ff6579ce 100644 --- a/basaltpass-backend/internal/handler/public/subscription/checkout_handler.go +++ b/basaltpass-backend/internal/handler/public/subscription/checkout_handler.go @@ -20,6 +20,10 @@ func CheckoutHandler(c *fiber.Ctx) error { InitCheckoutHandler() userID := c.Locals("userID").(uint) + tenantID, tenantErr := resolveCurrentUserTenantID(c) + if tenantErr != nil { + return c.Status(fiber.StatusForbidden).JSON(fiber.Map{"error": tenantErr.Error()}) + } var req CheckoutRequest if err := c.BodyParser(&req); err != nil { @@ -48,6 +52,10 @@ func CheckoutHandler(c *fiber.Ctx) error { }) } + if tenantID != nil { + req.ActiveTenantID = *tenantID + } + response, err := checkoutService.CreateCheckout(&req) if err != nil { return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{ @@ -67,6 +75,10 @@ func QuickCheckoutHandler(c *fiber.Ctx) error { InitCheckoutHandler() userID := c.Locals("userID").(uint) + tenantID, tenantErr := resolveCurrentUserTenantID(c) + if tenantErr != nil { + return c.Status(fiber.StatusForbidden).JSON(fiber.Map{"error": tenantErr.Error()}) + } var req struct { PriceID uint `json:"price_id" validate:"required"` @@ -89,6 +101,9 @@ func QuickCheckoutHandler(c *fiber.Ctx) error { SuccessURL: "http://localhost:5101/subscriptions?payment=success", CancelURL: "http://localhost:5101/subscriptions?payment=canceled", } + if tenantID != nil { + checkoutReq.ActiveTenantID = *tenantID + } response, err := checkoutService.CreateCheckout(&checkoutReq) if err != nil { diff --git a/basaltpass-backend/internal/handler/public/subscription/handler.go b/basaltpass-backend/internal/handler/public/subscription/handler.go index 18c63a58..70ade4d0 100644 --- a/basaltpass-backend/internal/handler/public/subscription/handler.go +++ b/basaltpass-backend/internal/handler/public/subscription/handler.go @@ -447,7 +447,14 @@ func (h *Handler) ValidateCoupon(c *fiber.Ctx) error { return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{"error": "优惠券代码不能为空"}) } - coupon, err := h.service.GetCouponByCode(code) + tenantID, tenantErr := resolveCurrentUserTenantID(c) + var coupon *model.Coupon + var err error + if tenantErr == nil && tenantID != nil { + coupon, err = h.service.GetCouponByCodeForTenant(code, tenantID) + } else { + coupon, err = h.service.GetCouponByCode(code) + } if err != nil { return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{"error": err.Error()}) } diff --git a/basaltpass-backend/internal/handler/public/subscription/service.go b/basaltpass-backend/internal/handler/public/subscription/service.go index b8a4dcbe..06a1b22f 100644 --- a/basaltpass-backend/internal/handler/public/subscription/service.go +++ b/basaltpass-backend/internal/handler/public/subscription/service.go @@ -5,6 +5,8 @@ import ( paymentservice "basaltpass-backend/internal/service/payment" "errors" "fmt" + "strconv" + "strings" "time" "basaltpass-backend/internal/model" @@ -20,6 +22,57 @@ func NewService(db *gorm.DB) *Service { return &Service{db: db} } +func parseTenantIDFromMetadata(metadata map[string]interface{}) uint64 { + if metadata == nil { + return 0 + } + v, ok := metadata["tenant_id"] + if !ok { + return 0 + } + switch value := v.(type) { + case string: + n, err := strconv.ParseUint(strings.TrimSpace(value), 10, 64) + if err == nil { + return n + } + case float64: + return uint64(value) + case uint64: + return value + case uint: + return uint64(value) + case int: + if value >= 0 { + return uint64(value) + } + } + return 0 +} + +func userHasTenantIdentity(tx *gorm.DB, userID uint, tenantID uint64) (bool, error) { + if userID == 0 || tenantID == 0 { + return false, nil + } + + var user model.User + if err := tx.Select("tenant_id").First(&user, userID).Error; err != nil { + return false, err + } + if user.TenantID > 0 && uint64(user.TenantID) == tenantID { + return true, nil + } + + var membershipCount int64 + if err := tx.Model(&model.TenantUser{}). + Where("user_id = ? AND tenant_id = ?", userID, tenantID). + Count(&membershipCount).Error; err != nil { + return false, err + } + + return membershipCount > 0, nil +} + // ========== 产品管理 ========== // CreateProduct 创建产品 @@ -472,8 +525,17 @@ func (s *Service) CreateCoupon(req *subdto.CreateCouponRequest) (*model.Coupon, // GetCouponByCode 根据代码获取优惠券 func (s *Service) GetCouponByCode(code string) (*model.Coupon, error) { + return s.GetCouponByCodeForTenant(code, nil) +} + +func (s *Service) GetCouponByCodeForTenant(code string, tenantID *uint64) (*model.Coupon, error) { var coupon model.Coupon - if err := s.db.Where("code = ?", code).First(&coupon).Error; err != nil { + query := s.db.Where("code = ?", code) + if tenantID != nil { + query = query.Where("tenant_id = ?", *tenantID) + } + + if err := query.First(&coupon).Error; err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return nil, fmt.Errorf("优惠券不存在") } @@ -597,11 +659,25 @@ func (s *Service) CreateSubscription(req *subdto.CreateSubscriptionRequest) (*mo } return nil, fmt.Errorf("查询用户失败: %w", err) } - if user.TenantID == 0 { + + tenantID := parseTenantIDFromMetadata(req.Metadata) + if tenantID == 0 && user.TenantID > 0 { + tenantID = uint64(user.TenantID) + } + if tenantID == 0 { tx.Rollback() return nil, fmt.Errorf("用户未绑定租户") } - userTenantID := uint64(user.TenantID) + + hasIdentity, err := userHasTenantIdentity(tx, req.UserID, tenantID) + if err != nil { + tx.Rollback() + return nil, fmt.Errorf("校验用户租户身份失败: %w", err) + } + if !hasIdentity { + tx.Rollback() + return nil, fmt.Errorf("用户不属于当前租户") + } // 验证价格存在 var price model.Price @@ -612,7 +688,7 @@ func (s *Service) CreateSubscription(req *subdto.CreateSubscriptionRequest) (*mo } return nil, fmt.Errorf("查询价格失败: %w", err) } - if price.TenantID != nil && *price.TenantID != userTenantID { + if price.TenantID == nil || *price.TenantID != tenantID { tx.Rollback() return nil, fmt.Errorf("价格不属于当前租户") } @@ -620,15 +696,11 @@ func (s *Service) CreateSubscription(req *subdto.CreateSubscriptionRequest) (*mo // 处理优惠券 var couponID *uint if req.CouponCode != nil { - coupon, err := s.GetCouponByCode(*req.CouponCode) + coupon, err := s.GetCouponByCodeForTenant(*req.CouponCode, &tenantID) if err != nil { tx.Rollback() return nil, err } - if coupon.TenantID != nil && *coupon.TenantID != userTenantID { - tx.Rollback() - return nil, fmt.Errorf("优惠券不属于当前租户") - } couponID = &coupon.ID } @@ -656,7 +728,7 @@ func (s *Service) CreateSubscription(req *subdto.CreateSubscriptionRequest) (*mo } subscription := &model.Subscription{ - TenantID: &userTenantID, + TenantID: &tenantID, UserID: req.UserID, Status: status, CurrentPriceID: req.PriceID, diff --git a/basaltpass-backend/internal/handler/s2s/user_handler.go b/basaltpass-backend/internal/handler/s2s/user_handler.go index fbc8aa52..4c6cfa4b 100644 --- a/basaltpass-backend/internal/handler/s2s/user_handler.go +++ b/basaltpass-backend/internal/handler/s2s/user_handler.go @@ -24,6 +24,7 @@ func unifiedResponse(c *fiber.Ctx, status int, data interface{}, errObj interfac func userSummary(u model.User) fiber.Map { return fiber.Map{ "id": u.ID, + "user_uuid": u.UserUUID, "email": u.Email, "nickname": u.Nickname, "avatar_url": u.AvatarURL, @@ -332,9 +333,8 @@ func LookupUsersHandler(c *fiber.Ctx) error { } db := common.DB().Model(&model.User{}). - Joins("JOIN system_auth_user_roles ur ON ur.user_id = system_auth_users.id"). - Joins("JOIN system_auth_roles r ON r.id = ur.role_id"). - Where("r.tenant_id = ?", tenantID) + Joins("JOIN tenant_users tu ON tu.user_id = system_auth_users.id"). + Where("tu.tenant_id = ?", tenantID) page, _ := strconv.Atoi(c.Query("page", "1")) pageSize, _ := strconv.Atoi(c.Query("page_size", "20")) diff --git a/basaltpass-backend/internal/handler/tenant/user_handler.go b/basaltpass-backend/internal/handler/tenant/user_handler.go index 388a9e76..3ca8df8c 100644 --- a/basaltpass-backend/internal/handler/tenant/user_handler.go +++ b/basaltpass-backend/internal/handler/tenant/user_handler.go @@ -18,12 +18,14 @@ import ( "strings" "github.com/gofiber/fiber/v2" + "github.com/google/uuid" "gorm.io/gorm" ) // TenantUserResponse 租户用户响应 type TenantUserResponse struct { ID uint `json:"id"` + UserUUID string `json:"user_uuid"` Email string `json:"email"` Nickname string `json:"nickname"` Avatar string `json:"avatar"` @@ -55,6 +57,21 @@ type InviteTenantUserRequest struct { Message string `json:"message,omitempty"` } +// GlobalUserCandidateResponse 全局用户候选响应 +type GlobalUserCandidateResponse struct { + ID uint `json:"id"` + UserUUID string `json:"user_uuid"` + Email string `json:"email"` + Nickname string `json:"nickname"` + Avatar string `json:"avatar"` + CreatedAt time.Time `json:"created_at"` +} + +// AuthorizeGlobalUserRequest 授权全局用户加入当前租户 +type AuthorizeGlobalUserRequest struct { + Role string `json:"role,omitempty"` +} + func normalizeTenantInviteRole(raw string) (model.TenantRole, string, bool) { switch strings.ToLower(strings.TrimSpace(raw)) { case string(model.TenantRoleAdmin): @@ -107,6 +124,52 @@ func loadTenantMembership(userID, tenantID uint) (model.User, *model.TenantUser, return user, tenantUserPtr, true, nil } +func loadTenantMembershipByUUID(userUUID string, tenantID uint) (model.User, *model.TenantUser, bool, error) { + normalizedUUID := strings.ToLower(strings.TrimSpace(userUUID)) + if normalizedUUID == "" { + return model.User{}, nil, false, gorm.ErrRecordNotFound + } + + var user model.User + if err := common.DB().Unscoped().Where("LOWER(user_uuid) = ?", normalizedUUID).First(&user).Error; err != nil { + return model.User{}, nil, false, err + } + + return loadTenantMembership(user.ID, tenantID) +} + +func buildTenantUserResponse(user model.User, tenantUser *model.TenantUser) TenantUserResponse { + role := string(model.TenantRoleUser) + createdAt := user.CreatedAt + updatedAt := user.UpdatedAt + if tenantUser != nil { + role = string(tenantUser.Role) + createdAt = tenantUser.CreatedAt + updatedAt = tenantUser.UpdatedAt + } + + status := "active" + if role == "baned" { + status = "baned" + } else if !user.EmailVerified { + status = "inactive" + } else if user.DeletedAt.Valid { + status = "suspended" + } + + return TenantUserResponse{ + ID: user.ID, + UserUUID: user.UserUUID, + Email: user.Email, + Nickname: user.Nickname, + Avatar: user.AvatarURL, + Role: role, + Status: status, + CreatedAt: createdAt, + UpdatedAt: updatedAt, + } +} + // GetTenantUsersHandler 获取租户用户列表 // GET /api/v1/tenant/users func GetTenantUsersHandler(c *fiber.Ctx) error { @@ -173,6 +236,7 @@ func GetTenantUsersHandler(c *fiber.Ctx) error { // 注意:SQLite 中时间字段需要特殊处理 type tenantUserRow struct { ID uint `json:"id"` + UserUUID string `json:"user_uuid"` Email string `json:"email"` Nickname string `json:"nickname"` Avatar string `json:"avatar"` @@ -185,6 +249,7 @@ func GetTenantUsersHandler(c *fiber.Ctx) error { var rows []tenantUserRow listQuery := base.Select(` u.id, + u.user_uuid, u.email, u.nickname, COALESCE(u.avatar_url, '') as avatar, @@ -211,6 +276,7 @@ func GetTenantUsersHandler(c *fiber.Ctx) error { for _, r := range rows { user := TenantUserResponse{ ID: r.ID, + UserUUID: r.UserUUID, Email: r.Email, Nickname: r.Nickname, Avatar: r.Avatar, @@ -246,6 +312,7 @@ func GetTenantUsersHandler(c *fiber.Ctx) error { // TenantAppLinkedUserResponse 授权了当前租户任一应用的用户(从 app_users 聚合) type TenantAppLinkedUserResponse struct { ID uint `json:"id"` + UserUUID string `json:"user_uuid"` Email string `json:"email"` Nickname string `json:"nickname"` Avatar string `json:"avatar"` @@ -298,13 +365,14 @@ func GetTenantAppLinkedUsersHandler(c *fiber.Ctx) error { var users []TenantAppLinkedUserResponse listQuery := base.Select(` system_auth_users.id, + system_auth_users.user_uuid, system_auth_users.email, system_auth_users.nickname, system_auth_users.avatar_url as avatar, COUNT(DISTINCT au.app_id) as app_count, MAX(au.last_authorized_at) as last_authorized_at, MAX(au.last_active_at) as last_active_at`). - Group("system_auth_users.id, system_auth_users.email, system_auth_users.nickname, system_auth_users.avatar_url"). + Group("system_auth_users.id, system_auth_users.user_uuid, system_auth_users.email, system_auth_users.nickname, system_auth_users.avatar_url"). Order("MAX(au.last_authorized_at) DESC"). Offset(offset).Limit(limit) @@ -369,6 +437,187 @@ func GetTenantUserStatsHandler(c *fiber.Ctx) error { }) } +// GetGlobalUserCandidatesHandler 获取可授权加入当前租户的全局用户候选 +// GET /api/v1/tenant/users/global-candidates +func GetGlobalUserCandidatesHandler(c *fiber.Ctx) error { + if err := requireTenantAdminRole(c); err != nil { + return err + } + + page, _ := strconv.Atoi(c.Query("page", "1")) + limit, _ := strconv.Atoi(c.Query("limit", "20")) + search := strings.TrimSpace(c.Query("search", "")) + + if page < 1 { + page = 1 + } + if limit < 1 || limit > 100 { + limit = 20 + } + offset := (page - 1) * limit + + query := common.DB().Table("system_auth_users"). + Where("tenant_id = ?", 0). + Where("deleted_at IS NULL"). + Where("COALESCE(is_system_admin, 0) = 0") + + if search != "" { + query = query.Where("LOWER(email) LIKE LOWER(?) OR LOWER(nickname) LIKE LOWER(?)", "%"+search+"%", "%"+search+"%") + } + + var total int64 + if err := query.Count(&total).Error; err != nil { + return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{ + "error": "查询全局用户失败", + }) + } + + var users []GlobalUserCandidateResponse + if err := query.Select("id, user_uuid, email, nickname, COALESCE(avatar_url, '') as avatar, created_at"). + Order("created_at DESC"). + Offset(offset). + Limit(limit). + Scan(&users).Error; err != nil { + return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{ + "error": "查询全局用户失败", + }) + } + + return c.JSON(fiber.Map{ + "users": users, + "pagination": fiber.Map{ + "page": page, + "limit": limit, + "total": total, + }, + }) +} + +// AuthorizeGlobalUserToTenantHandler 授权全局用户加入当前租户 +// POST /api/v1/tenant/users/global-candidates/:id/authorize +func AuthorizeGlobalUserToTenantHandler(c *fiber.Ctx) error { + if err := requireTenantAdminRole(c); err != nil { + return err + } + + tenantID := c.Locals("tenantID").(uint) + userIDUint64, err := strconv.ParseUint(c.Params("id"), 10, 32) + if err != nil { + return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{ + "error": "无效的用户ID", + }) + } + userID := uint(userIDUint64) + + var req AuthorizeGlobalUserRequest + if err := c.BodyParser(&req); err != nil { + return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{ + "error": "请求参数错误", + }) + } + + assignedRole := model.TenantRoleMember + if strings.TrimSpace(req.Role) != "" { + normalizedRole, _, ok := normalizeTenantInviteRole(req.Role) + if !ok { + return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{ + "error": "无效的角色类型", + }) + } + assignedRole = normalizedRole + } + + var user model.User + if err := common.DB().First(&user, userID).Error; err != nil { + if err == gorm.ErrRecordNotFound { + return c.Status(fiber.StatusNotFound).JSON(fiber.Map{ + "error": "用户不存在", + }) + } + return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{ + "error": "查询用户失败", + }) + } + + if user.IsSuperAdmin() { + return c.Status(fiber.StatusForbidden).JSON(fiber.Map{ + "error": "不支持授权系统管理员加入租户", + }) + } + + var membershipCount int64 + if err := common.DB().Model(&model.TenantUser{}).Where("user_id = ?", user.ID).Count(&membershipCount).Error; err != nil { + return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{ + "error": "查询用户失败", + }) + } + + if user.TenantID != 0 && user.TenantID != tenantID && membershipCount == 0 { + return c.Status(fiber.StatusConflict).JSON(fiber.Map{ + "error": "用户已属于其他租户", + }) + } + + tx := common.DB().Begin() + defer func() { + if r := recover(); r != nil { + tx.Rollback() + } + }() + + // Keep global user identity in users.tenant_id=0 and rely on tenant_users for tenant perspective. + // For legacy drifted data, normalize back to 0 when membership exists. + if user.TenantID != 0 && membershipCount > 0 { + if err := tx.Model(&model.User{}).Where("id = ?", user.ID).Update("tenant_id", 0).Error; err != nil { + tx.Rollback() + return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{ + "error": "授权加入租户失败", + }) + } + } + + var tenantUser model.TenantUser + err = tx.Where("tenant_id = ? AND user_id = ?", tenantID, user.ID).First(&tenantUser).Error + if err != nil { + if err != gorm.ErrRecordNotFound { + tx.Rollback() + return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{ + "error": "授权加入租户失败", + }) + } + + newTenantUser := model.TenantUser{ + UserID: user.ID, + TenantID: tenantID, + Role: assignedRole, + } + if err := tx.Create(&newTenantUser).Error; err != nil { + tx.Rollback() + return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{ + "error": "授权加入租户失败", + }) + } + } else if strings.TrimSpace(req.Role) != "" && tenantUser.Role != model.TenantRoleOwner { + tenantUser.Role = assignedRole + if err := tx.Save(&tenantUser).Error; err != nil { + tx.Rollback() + return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{ + "error": "授权加入租户失败", + }) + } + } + + if err := tx.Commit().Error; err != nil { + return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{ + "error": "授权加入租户失败", + }) + } + + return c.JSON(fiber.Map{ + "message": "用户已加入当前租户", + }) +} + // UpdateTenantUserHandler 更新租户用户 // PUT /api/v1/tenant/users/:id func UpdateTenantUserHandler(c *fiber.Ctx) error { @@ -819,34 +1068,43 @@ func GetTenantUserHandler(c *fiber.Ctx) error { }) } - role := string(model.TenantRoleUser) - createdAt := user.CreatedAt - updatedAt := user.UpdatedAt - if tenantUser != nil { - role = string(tenantUser.Role) - createdAt = tenantUser.CreatedAt - updatedAt = tenantUser.UpdatedAt - } + resp := buildTenantUserResponse(user, tenantUser) - status := "active" - if role == "baned" { - status = "baned" - } else if !user.EmailVerified { - status = "inactive" - } else if user.DeletedAt.Valid { - status = "suspended" + return c.JSON(fiber.Map{ + "user": resp, + }) +} + +// GetTenantUserByUUIDHandler 获取租户用户详情(通过 user_uuid) +// GET /api/v1/tenant/users/by-uuid/:user_uuid +func GetTenantUserByUUIDHandler(c *fiber.Ctx) error { + tenantID := c.Locals("tenantID").(uint) + userUUID := strings.TrimSpace(c.Params("user_uuid")) + + if _, err := uuid.Parse(userUUID); err != nil { + return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{ + "error": "无效的用户UUID", + }) } - resp := TenantUserResponse{ - ID: user.ID, - Email: user.Email, - Nickname: user.Nickname, - Avatar: user.AvatarURL, - Role: role, - Status: status, - CreatedAt: createdAt, - UpdatedAt: updatedAt, + user, tenantUser, belongs, err := loadTenantMembershipByUUID(userUUID, tenantID) + if err != nil { + if err == gorm.ErrRecordNotFound { + return c.Status(fiber.StatusNotFound).JSON(fiber.Map{ + "error": "用户不存在", + }) + } + return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{ + "error": "查询用户失败", + }) } + if !belongs { + return c.Status(fiber.StatusNotFound).JSON(fiber.Map{ + "error": "用户不存在或不属于当前租户", + }) + } + + resp := buildTenantUserResponse(user, tenantUser) return c.JSON(fiber.Map{ "user": resp, diff --git a/basaltpass-backend/internal/handler/tenant/user_handler_uuid_test.go b/basaltpass-backend/internal/handler/tenant/user_handler_uuid_test.go new file mode 100644 index 00000000..af40a870 --- /dev/null +++ b/basaltpass-backend/internal/handler/tenant/user_handler_uuid_test.go @@ -0,0 +1,127 @@ +package tenant + +import ( + "encoding/json" + "fmt" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "basaltpass-backend/internal/common" + "basaltpass-backend/internal/model" + + "github.com/glebarez/sqlite" + "github.com/gofiber/fiber/v2" + "gorm.io/gorm" +) + +func setupTenantUserUUIDHandlerTestDB(t *testing.T) *gorm.DB { + t.Helper() + + dsn := fmt.Sprintf("file:%s?mode=memory&cache=shared", strings.ReplaceAll(t.Name(), "/", "_")) + db, err := gorm.Open(sqlite.Open(dsn), &gorm.Config{}) + if err != nil { + t.Fatalf("failed to open sqlite: %v", err) + } + + if err := db.AutoMigrate(&model.User{}, &model.TenantUser{}); err != nil { + t.Fatalf("auto migrate failed: %v", err) + } + + common.SetDBForTest(db) + return db +} + +func newTenantUserUUIDTestApp(tenantID uint) *fiber.App { + app := fiber.New() + app.Get("/tenant/users/by-uuid/:user_uuid", func(c *fiber.Ctx) error { + c.Locals("tenantID", tenantID) + return GetTenantUserByUUIDHandler(c) + }) + return app +} + +func TestGetTenantUserByUUIDHandler_Success(t *testing.T) { + db := setupTenantUserUUIDHandlerTestDB(t) + tenantID := uint(2001) + + u := model.User{ + TenantID: tenantID, + Email: "tenant-uuid-user@example.com", + PasswordHash: "x", + Nickname: "tenant-user", + EmailVerified: true, + } + if err := db.Create(&u).Error; err != nil { + t.Fatalf("create user failed: %v", err) + } + + app := newTenantUserUUIDTestApp(tenantID) + req := httptest.NewRequest(http.MethodGet, "/tenant/users/by-uuid/"+u.UserUUID, nil) + resp, err := app.Test(req, -1) + if err != nil { + t.Fatalf("request failed: %v", err) + } + defer resp.Body.Close() + + if resp.StatusCode != fiber.StatusOK { + t.Fatalf("expected 200, got %d", resp.StatusCode) + } + + var payload struct { + User TenantUserResponse `json:"user"` + } + if err := json.NewDecoder(resp.Body).Decode(&payload); err != nil { + t.Fatalf("decode response failed: %v", err) + } + if payload.User.ID != u.ID { + t.Fatalf("expected id %d, got %d", u.ID, payload.User.ID) + } + if payload.User.UserUUID != u.UserUUID { + t.Fatalf("expected user_uuid %q, got %q", u.UserUUID, payload.User.UserUUID) + } +} + +func TestGetTenantUserByUUIDHandler_InvalidUUID(t *testing.T) { + setupTenantUserUUIDHandlerTestDB(t) + app := newTenantUserUUIDTestApp(1) + + req := httptest.NewRequest(http.MethodGet, "/tenant/users/by-uuid/not-a-uuid", nil) + resp, err := app.Test(req, -1) + if err != nil { + t.Fatalf("request failed: %v", err) + } + defer resp.Body.Close() + + if resp.StatusCode != fiber.StatusBadRequest { + t.Fatalf("expected 400, got %d", resp.StatusCode) + } +} + +func TestGetTenantUserByUUIDHandler_WrongTenant(t *testing.T) { + db := setupTenantUserUUIDHandlerTestDB(t) + ownerTenantID := uint(3001) + otherTenantID := uint(3002) + + u := model.User{ + TenantID: ownerTenantID, + Email: "other-tenant-user@example.com", + PasswordHash: "x", + } + if err := db.Create(&u).Error; err != nil { + t.Fatalf("create user failed: %v", err) + } + + app := newTenantUserUUIDTestApp(otherTenantID) + req := httptest.NewRequest(http.MethodGet, "/tenant/users/by-uuid/"+u.UserUUID, nil) + resp, err := app.Test(req, -1) + if err != nil { + t.Fatalf("request failed: %v", err) + } + defer resp.Body.Close() + + if resp.StatusCode != fiber.StatusNotFound { + t.Fatalf("expected 404, got %d", resp.StatusCode) + } +} diff --git a/basaltpass-backend/internal/handler/user/handler.go b/basaltpass-backend/internal/handler/user/handler.go index 0f56cd33..394dba6b 100644 --- a/basaltpass-backend/internal/handler/user/handler.go +++ b/basaltpass-backend/internal/handler/user/handler.go @@ -5,9 +5,12 @@ import ( userdto "basaltpass-backend/internal/dto/user" "basaltpass-backend/internal/model" tenant2 "basaltpass-backend/internal/service/tenant" + "errors" "strconv" + "strings" "github.com/gofiber/fiber/v2" + "gorm.io/gorm" ) var svc = Service{} @@ -16,7 +19,11 @@ var tenantSvc = tenant2.NewTenantService() // GetProfileHandler handles GET /user/profile func GetProfileHandler(c *fiber.Ctx) error { uid := c.Locals("userID").(uint) - profile, err := svc.GetProfile(uid) + var activeTenantID uint + if tid, ok := c.Locals("tenantID").(uint); ok { + activeTenantID = tid + } + profile, err := svc.GetProfile(uid, activeTenantID) if err != nil { return c.Status(fiber.StatusNotFound).JSON(fiber.Map{"error": err.Error()}) } @@ -40,6 +47,107 @@ func GetUserTenantsHandler(c *fiber.Ctx) error { }) } +// JoinTenantByCodeHandler allows a logged-in user to join a tenant via tenant code. +// POST /user/tenants/join-by-code/:code +func JoinTenantByCodeHandler(c *fiber.Ctx) error { + userID := c.Locals("userID").(uint) + tenantCode := strings.TrimSpace(c.Params("code")) + if tenantCode == "" { + return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{ + "error": "tenant code is required", + }) + } + + db := common.DB() + + var user model.User + if err := db.Select("id", "tenant_id", "is_system_admin").First(&user, userID).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return c.Status(fiber.StatusNotFound).JSON(fiber.Map{"error": "user not found"}) + } + return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{"error": "failed to load user"}) + } + + if user.IsSuperAdmin() { + return c.Status(fiber.StatusForbidden).JSON(fiber.Map{ + "error": "system admin cannot join tenant via join link", + }) + } + + var membershipCount int64 + if err := db.Model(&model.TenantUser{}).Where("user_id = ?", user.ID).Count(&membershipCount).Error; err != nil { + return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{"error": "failed to load user membership"}) + } + + var tenant model.Tenant + if err := db.Where("code = ? AND status = ?", tenantCode, model.TenantStatusActive). + Select("id", "name", "code", "status"). + First(&tenant).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return c.Status(fiber.StatusNotFound).JSON(fiber.Map{"error": "tenant not found"}) + } + return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{"error": "failed to load tenant"}) + } + + if user.TenantID != 0 && user.TenantID != tenant.ID && membershipCount == 0 { + return c.Status(fiber.StatusConflict).JSON(fiber.Map{ + "error": "user already belongs to another tenant", + }) + } + + tx := db.Begin() + defer func() { + if r := recover(); r != nil { + tx.Rollback() + } + }() + + joinedNow := false + var membership model.TenantUser + err := tx.Where("tenant_id = ? AND user_id = ?", tenant.ID, user.ID).First(&membership).Error + if err != nil { + if !errors.Is(err, gorm.ErrRecordNotFound) { + tx.Rollback() + return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{"error": "failed to join tenant"}) + } + + membership = model.TenantUser{ + TenantID: tenant.ID, + UserID: user.ID, + Role: model.TenantRoleMember, + } + if err := tx.Create(&membership).Error; err != nil { + tx.Rollback() + return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{"error": "failed to join tenant"}) + } + joinedNow = true + } + + // Keep global account identity at tenant_id=0 and use tenant_users for tenant perspective switching. + // For legacy data drift (tenant_id changed after join), normalize it back to 0 when membership exists. + if user.TenantID != 0 && membershipCount > 0 { + if err := tx.Model(&model.User{}).Where("id = ?", user.ID).Update("tenant_id", 0).Error; err != nil { + tx.Rollback() + return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{"error": "failed to join tenant"}) + } + } + + if err := tx.Commit().Error; err != nil { + return c.Status(fiber.StatusInternalServerError).JSON(fiber.Map{"error": "failed to join tenant"}) + } + + return c.JSON(fiber.Map{ + "message": "joined tenant successfully", + "joined": joinedNow, + "already_joined": !joinedNow, + "tenant": fiber.Map{ + "id": tenant.ID, + "name": tenant.Name, + "code": tenant.Code, + }, + }) +} + // UpdateProfileHandler handles PUT /user/profile func UpdateProfileHandler(c *fiber.Ctx) error { uid := c.Locals("userID").(uint) diff --git a/basaltpass-backend/internal/handler/user/service.go b/basaltpass-backend/internal/handler/user/service.go index 19a193f4..e9c3c94e 100644 --- a/basaltpass-backend/internal/handler/user/service.go +++ b/basaltpass-backend/internal/handler/user/service.go @@ -12,7 +12,7 @@ import ( type Service struct{} // GetProfile returns user profile by ID. -func (s Service) GetProfile(userID uint) (userdto.ProfileResponse, error) { +func (s Service) GetProfile(userID uint, activeTenantID uint) (userdto.ProfileResponse, error) { var u model.User if err := common.DB().First(&u, userID).Error; err != nil { return userdto.ProfileResponse{}, err @@ -26,21 +26,29 @@ func (s Service) GetProfile(userID uint) (userdto.ProfileResponse, error) { tenantRole string ) - // 普通租户用户始终以 users.tenant_id 作为其默认租户上下文。 - if u.TenantID > 0 { - tid := u.TenantID - tenantID = &tid - hasTenant = true + // activeTenantID 来自当前 token 的 tid,可支持全局用户切换到某个 tenant identity。 + resolvedTenantID := activeTenantID + if resolvedTenantID == 0 { + if u.TenantID > 0 { + resolvedTenantID = u.TenantID + } } - // 平台管理员(tenant_id=0)不应因为 tenant_users 历史记录而自动获得租户控制台上下文。 - if u.TenantID > 0 { - var ta model.TenantUser + if resolvedTenantID > 0 { + var membership model.TenantUser if err := common.DB(). - Where("user_id = ? AND tenant_id = ?", userID, u.TenantID). + Where("user_id = ? AND tenant_id = ?", userID, resolvedTenantID). Order("created_at ASC"). - First(&ta).Error; err == nil { - tenantRole = string(ta.Role) + First(&membership).Error; err == nil { + hasTenant = true + tid := resolvedTenantID + tenantID = &tid + tenantRole = string(membership.Role) + } else if u.TenantID == resolvedTenantID { + hasTenant = true + tid := resolvedTenantID + tenantID = &tid + tenantRole = string(model.TenantRoleUser) } } diff --git a/basaltpass-backend/internal/handler/user/service_test.go b/basaltpass-backend/internal/handler/user/service_test.go index 5163cf22..ae642973 100644 --- a/basaltpass-backend/internal/handler/user/service_test.go +++ b/basaltpass-backend/internal/handler/user/service_test.go @@ -48,7 +48,7 @@ func TestGetProfileUsesUserTenantIDForRegularUser(t *testing.T) { } svc := Service{} - profile, err := svc.GetProfile(u.ID) + profile, err := svc.GetProfile(u.ID, 0) if err != nil { t.Fatalf("GetProfile failed: %v", err) } @@ -94,7 +94,7 @@ func TestGetProfileDoesNotDeriveTenantContextFromTenantUsersForPlatformAdmin(t * } svc := Service{} - profile, err := svc.GetProfile(u.ID) + profile, err := svc.GetProfile(u.ID, 0) if err != nil { t.Fatalf("GetProfile failed: %v", err) } diff --git a/basaltpass-backend/internal/handler/user/wallet_handler.go b/basaltpass-backend/internal/handler/user/wallet_handler.go index 7ee2855c..6d28e9df 100644 --- a/basaltpass-backend/internal/handler/user/wallet_handler.go +++ b/basaltpass-backend/internal/handler/user/wallet_handler.go @@ -13,9 +13,13 @@ import ( // GetWalletBalanceHandler GET /wallet/balance?currency=USD func GetWalletBalanceHandler(c *fiber.Ctx) error { uid := c.Locals("userID").(uint) + activeTenantID, _ := c.Locals("tenantID").(uint) currency := c.Query("currency", "USD") - w, err := wallet.GetBalanceByCode(uid, currency) + w, err := wallet.GetBalanceByCodeWithTenant(uid, activeTenantID, currency) if err != nil { + if errors.Is(err, wallet.ErrNoTenantIdentity) { + return c.Status(fiber.StatusNotFound).JSON(fiber.Map{"error": "wallet is unavailable for users without tenant"}) + } return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{"error": err.Error()}) } return c.JSON(fiber.Map{"balance": w.Balance, "currency_id": w.CurrencyID, "tenant_id": w.TenantID}) @@ -24,6 +28,7 @@ func GetWalletBalanceHandler(c *fiber.Ctx) error { // RechargeWalletHandler POST /wallet/recharge {currency, amount} func RechargeWalletHandler(c *fiber.Ctx) error { uid := c.Locals("userID").(uint) + activeTenantID, _ := c.Locals("tenantID").(uint) var body struct { Currency string `json:"currency"` Amount int64 `json:"amount"` @@ -31,7 +36,10 @@ func RechargeWalletHandler(c *fiber.Ctx) error { if err := c.BodyParser(&body); err != nil { return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{"error": err.Error()}) } - if err := wallet.RechargeByCode(uid, body.Currency, body.Amount); err != nil { + if err := wallet.RechargeByCodeWithTenant(uid, activeTenantID, body.Currency, body.Amount); err != nil { + if errors.Is(err, wallet.ErrNoTenantIdentity) { + return c.Status(fiber.StatusForbidden).JSON(fiber.Map{"error": "wallet is unavailable for users without tenant"}) + } if errors.Is(err, wallet.ErrWalletRechargeWithdrawDisabled) { return c.Status(fiber.StatusForbidden).JSON(fiber.Map{"error": err.Error()}) } @@ -43,6 +51,7 @@ func RechargeWalletHandler(c *fiber.Ctx) error { // WithdrawWalletHandler POST /wallet/withdraw {currency, amount} func WithdrawWalletHandler(c *fiber.Ctx) error { uid := c.Locals("userID").(uint) + activeTenantID, _ := c.Locals("tenantID").(uint) var body struct { Currency string `json:"currency"` Amount int64 `json:"amount"` @@ -50,7 +59,10 @@ func WithdrawWalletHandler(c *fiber.Ctx) error { if err := c.BodyParser(&body); err != nil { return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{"error": err.Error()}) } - if err := wallet.WithdrawByCode(uid, body.Currency, body.Amount); err != nil { + if err := wallet.WithdrawByCodeWithTenant(uid, activeTenantID, body.Currency, body.Amount); err != nil { + if errors.Is(err, wallet.ErrNoTenantIdentity) { + return c.Status(fiber.StatusForbidden).JSON(fiber.Map{"error": "wallet is unavailable for users without tenant"}) + } if errors.Is(err, wallet.ErrWalletRechargeWithdrawDisabled) { return c.Status(fiber.StatusForbidden).JSON(fiber.Map{"error": err.Error()}) } @@ -62,17 +74,21 @@ func WithdrawWalletHandler(c *fiber.Ctx) error { // WalletHistoryHandler GET /wallet/history?currency=USD&limit=20 func WalletHistoryHandler(c *fiber.Ctx) error { uid := c.Locals("userID").(uint) + activeTenantID, _ := c.Locals("tenantID").(uint) currency := c.Query("currency", "") limitStr := c.Query("limit", "20") limit, _ := strconv.Atoi(limitStr) var txs interface{} var err error if currency == "" || currency == "all" { - txs, err = wallet.HistoryAllByUser(uid, limit) + txs, err = wallet.HistoryAllByUserWithTenant(uid, activeTenantID, limit) } else { - txs, err = wallet.HistoryByCode(uid, currency, limit) + txs, err = wallet.HistoryByCodeWithTenant(uid, activeTenantID, currency, limit) } if err != nil { + if errors.Is(err, wallet.ErrNoTenantIdentity) { + return c.JSON([]interface{}{}) + } return c.Status(fiber.StatusBadRequest).JSON(fiber.Map{"error": err.Error()}) } return c.JSON(txs) diff --git a/basaltpass-backend/internal/migration/migrate.go b/basaltpass-backend/internal/migration/migrate.go index 83260750..686277e1 100644 --- a/basaltpass-backend/internal/migration/migrate.go +++ b/basaltpass-backend/internal/migration/migrate.go @@ -16,6 +16,7 @@ import ( "strings" "time" + "github.com/google/uuid" "golang.org/x/crypto/bcrypt" ) @@ -254,6 +255,7 @@ func RunMigrations() { // 在自动迁移后处理特殊的迁移情况 handleSpecialMigrations() ensureUserTenantScopedUniqueIndexes() + ensureUserUUIDBackfillAndUniqueIndex() ensureIsSystemAdminNonUniqueIndex() seedSystemApps() @@ -1462,6 +1464,100 @@ func ensureUserTenantScopedUniqueIndexes() { } } +// ensureUserUUIDBackfillAndUniqueIndex guarantees every user has a stable UUID and enforces uniqueness. +func ensureUserUUIDBackfillAndUniqueIndex() { + db := common.DB() + log.Println("[Migration] Ensuring immutable UUID for every user...") + + type userRow struct { + ID uint + UserUUID string + } + var users []userRow + if err := db.Model(&model.User{}).Select("id", "user_uuid").Find(&users).Error; err != nil { + log.Printf("[Migration] Failed to load users for UUID backfill: %v", err) + return + } + + filled := 0 + for _, u := range users { + if strings.TrimSpace(u.UserUUID) != "" { + continue + } + if err := db.Model(&model.User{}).Where("id = ?", u.ID).Update("user_uuid", uuid.NewString()).Error; err != nil { + log.Printf("[Migration] Failed to backfill user_uuid for user %d: %v", u.ID, err) + continue + } + filled++ + } + if filled > 0 { + log.Printf("[Migration] Backfilled user_uuid for %d users", filled) + } + + type duplicateUUID struct { + UserUUID string + Count int64 + } + var duplicates []duplicateUUID + if err := db.Table("system_auth_users"). + Select("user_uuid, COUNT(1) AS count"). + Where("user_uuid IS NOT NULL AND TRIM(user_uuid) <> ''"). + Group("user_uuid"). + Having("COUNT(1) > 1"). + Scan(&duplicates).Error; err != nil { + log.Printf("[Migration] Failed to scan duplicate user_uuid values: %v", err) + } + + deduped := 0 + for _, dup := range duplicates { + type idRow struct { + ID uint + } + var rows []idRow + if err := db.Table("system_auth_users"). + Select("id"). + Where("user_uuid = ?", dup.UserUUID). + Order("id ASC"). + Scan(&rows).Error; err != nil { + log.Printf("[Migration] Failed to load duplicate users for user_uuid %s: %v", dup.UserUUID, err) + continue + } + + for i := 1; i < len(rows); i++ { + if err := db.Model(&model.User{}).Where("id = ?", rows[i].ID).Update("user_uuid", uuid.NewString()).Error; err != nil { + log.Printf("[Migration] Failed to deduplicate user_uuid for user %d: %v", rows[i].ID, err) + continue + } + deduped++ + } + } + if deduped > 0 { + log.Printf("[Migration] Re-assigned duplicated user_uuid for %d users", deduped) + } + + dropIndexIfExists("system_auth_users", "idx_users_uuid") + dropIndexIfExists("system_auth_users", "idx_users_user_uuid") + + dialect := strings.ToLower(db.Dialector.Name()) + if dialect == "sqlite" || dialect == "postgres" || dialect == "postgresql" { + if err := createIndexIfMissing( + "system_auth_users", + "idx_users_uuid", + "CREATE UNIQUE INDEX idx_users_uuid ON system_auth_users (user_uuid) WHERE user_uuid IS NOT NULL AND user_uuid != ''", + ); err != nil { + log.Printf("[Migration] Failed to create idx_users_uuid: %v", err) + } + } else { + if err := createIndexIfMissing( + "system_auth_users", + "idx_users_uuid", + "CREATE UNIQUE INDEX idx_users_uuid ON system_auth_users (user_uuid)", + ); err != nil { + log.Printf("[Migration] Failed to create idx_users_uuid: %v", err) + } + } +} + // ensureIsSystemAdminNonUniqueIndex removes legacy unique index on users.is_system_admin // and keeps a normal non-unique index for query performance. func ensureIsSystemAdminNonUniqueIndex() { diff --git a/basaltpass-backend/internal/model/user.go b/basaltpass-backend/internal/model/user.go index f79e53b1..481dc365 100644 --- a/basaltpass-backend/internal/model/user.go +++ b/basaltpass-backend/internal/model/user.go @@ -3,10 +3,12 @@ package model import ( "basaltpass-backend/internal/common" "encoding/base64" + "strings" "time" "github.com/go-webauthn/webauthn/protocol" "github.com/go-webauthn/webauthn/webauthn" + "github.com/google/uuid" "gorm.io/gorm" ) @@ -23,6 +25,7 @@ type User struct { // (email, tenant_id) 与 (phone, tenant_id) // 这样同一个邮箱/手机号可以在不同租户下注册不同账户。 // 注意:这里不声明 unique,避免 AutoMigrate 误建成单列唯一索引。 + UserUUID string `gorm:"column:user_uuid;size:36;<-:create" json:"user_uuid"` Email string `gorm:"size:128;index;default:null" json:"email"` Phone string `gorm:"size:32;index;default:null" json:"phone"` PasswordHash string `gorm:"size:255"` @@ -64,6 +67,14 @@ func (User) TableName() string { return "system_auth_users" } +// BeforeCreate guarantees an immutable, globally unique user UUID at creation time. +func (u *User) BeforeCreate(_ *gorm.DB) error { + if strings.TrimSpace(u.UserUUID) == "" { + u.UserUUID = uuid.NewString() + } + return nil +} + // IsSuperAdmin 检查用户是否为系统最高管理员,以 is_system_admin 字段为唯一依据。 func (u *User) IsSuperAdmin() bool { return u.IsSystemAdmin != nil && *u.IsSystemAdmin diff --git a/basaltpass-backend/internal/model/user_uuid_test.go b/basaltpass-backend/internal/model/user_uuid_test.go new file mode 100644 index 00000000..98aa1a48 --- /dev/null +++ b/basaltpass-backend/internal/model/user_uuid_test.go @@ -0,0 +1,68 @@ +package model + +import ( + "strings" + "testing" + + "github.com/glebarez/sqlite" + "github.com/google/uuid" + "gorm.io/gorm" +) + +func newUserUUIDTestDB(t *testing.T) *gorm.DB { + t.Helper() + + db, err := gorm.Open(sqlite.Open("file::memory:?cache=shared"), &gorm.Config{}) + if err != nil { + t.Fatalf("open sqlite failed: %v", err) + } + if err := db.AutoMigrate(&User{}); err != nil { + t.Fatalf("auto migrate user failed: %v", err) + } + return db +} + +func TestUserUUIDGeneratedOnCreate(t *testing.T) { + db := newUserUUIDTestDB(t) + + u := User{Email: "uuid-gen@example.com", PasswordHash: "x"} + if err := db.Create(&u).Error; err != nil { + t.Fatalf("create user failed: %v", err) + } + + if strings.TrimSpace(u.UserUUID) == "" { + t.Fatal("expected non-empty user_uuid after create") + } + if _, err := uuid.Parse(u.UserUUID); err != nil { + t.Fatalf("expected valid uuid format, got %q: %v", u.UserUUID, err) + } +} + +func TestUserUUIDIsCreateOnly(t *testing.T) { + db := newUserUUIDTestDB(t) + + u := User{Email: "uuid-immutable@example.com", PasswordHash: "x"} + if err := db.Create(&u).Error; err != nil { + t.Fatalf("create user failed: %v", err) + } + originalUUID := u.UserUUID + + if err := db.Model(&User{}).Where("id = ?", u.ID).Updates(map[string]interface{}{ + "user_uuid": "manually-overridden", + "nickname": "updated-nickname", + }).Error; err != nil { + t.Fatalf("update user failed: %v", err) + } + + var loaded User + if err := db.First(&loaded, u.ID).Error; err != nil { + t.Fatalf("reload user failed: %v", err) + } + + if loaded.Nickname != "updated-nickname" { + t.Fatalf("expected nickname to be updated, got %q", loaded.Nickname) + } + if loaded.UserUUID != originalUUID { + t.Fatalf("expected user_uuid to remain immutable, got %q want %q", loaded.UserUUID, originalUUID) + } +} diff --git a/basaltpass-backend/internal/service/access/access_service.go b/basaltpass-backend/internal/service/access/access_service.go index a3db2b42..bb9b6703 100644 --- a/basaltpass-backend/internal/service/access/access_service.go +++ b/basaltpass-backend/internal/service/access/access_service.go @@ -35,16 +35,34 @@ func (s *Service) ResolveTenantContext(userID uint, requestedTenantID uint) (uin return 0, "", err } - if user.TenantID == 0 { + tenantID := requestedTenantID + if tenantID == 0 { + if user.TenantID > 0 { + tenantID = user.TenantID + } else { + var membership model.TenantUser + if err := s.db.Select("tenant_id").Where("user_id = ?", userID).Order("created_at ASC").First(&membership).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return 0, "", ErrNoTenantAssociation + } + return 0, "", err + } + tenantID = membership.TenantID + } + } + + if tenantID == 0 { return 0, "", ErrNoTenantAssociation } - tenantID := user.TenantID - if requestedTenantID > 0 { - if requestedTenantID != user.TenantID { + if user.TenantID != tenantID { + var membershipCount int64 + if err := s.db.Model(&model.TenantUser{}).Where("user_id = ? AND tenant_id = ?", userID, tenantID).Count(&membershipCount).Error; err != nil { + return 0, "", err + } + if membershipCount == 0 { return 0, "", ErrInvalidTenantAssociation } - tenantID = requestedTenantID } var tenant model.Tenant @@ -72,6 +90,10 @@ func (s *Service) ResolveTenantContext(userID uint, requestedTenantID uint) (uin // GetTenantRole returns user's role in a tenant. func (s *Service) GetTenantRole(userID, tenantID uint) (model.TenantRole, error) { + if tenantID == 0 { + return "", ErrTenantMembershipNotFound + } + var user model.User if err := s.db.Select("id", "tenant_id").First(&user, userID).Error; err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { @@ -79,8 +101,15 @@ func (s *Service) GetTenantRole(userID, tenantID uint) (model.TenantRole, error) } return "", err } - if user.TenantID != tenantID || tenantID == 0 { - return "", ErrTenantMembershipNotFound + + if user.TenantID != tenantID { + var membershipCount int64 + if err := s.db.Model(&model.TenantUser{}).Where("user_id = ? AND tenant_id = ?", userID, tenantID).Count(&membershipCount).Error; err != nil { + return "", err + } + if membershipCount == 0 { + return "", ErrTenantMembershipNotFound + } } var tenantUser model.TenantUser diff --git a/basaltpass-backend/internal/service/auth/dto.go b/basaltpass-backend/internal/service/auth/dto.go index 93874655..e3d430f3 100644 --- a/basaltpass-backend/internal/service/auth/dto.go +++ b/basaltpass-backend/internal/service/auth/dto.go @@ -13,6 +13,7 @@ type LoginRequest struct { EmailOrPhone string `json:"identifier"` Password string `json:"password"` TenantID uint `json:"tenant_id"` // 租户ID,用于识别用户属于哪个租户 + Scope string `json:"-"` } // Verify2FARequest defines input for 2FA verification. diff --git a/basaltpass-backend/internal/service/auth/jwt.go b/basaltpass-backend/internal/service/auth/jwt.go index 04e01035..71a405b7 100644 --- a/basaltpass-backend/internal/service/auth/jwt.go +++ b/basaltpass-backend/internal/service/auth/jwt.go @@ -129,16 +129,14 @@ func resolveTokenTenantID(userID uint, claimedTenantID uint, scope string) (uint return claimedTenantID, nil } - if scope == ConsoleScopeTenant { - var membershipCount int64 - if err := common.DB().Model(&model.TenantUser{}). - Where("user_id = ? AND tenant_id = ?", userID, claimedTenantID). - Count(&membershipCount).Error; err != nil { - return 0, err - } - if membershipCount > 0 { - return claimedTenantID, nil - } + var membershipCount int64 + if err := common.DB().Model(&model.TenantUser{}). + Where("user_id = ? AND tenant_id = ?", userID, claimedTenantID). + Count(&membershipCount).Error; err != nil { + return 0, err + } + if membershipCount > 0 { + return claimedTenantID, nil } } diff --git a/basaltpass-backend/internal/service/auth/service.go b/basaltpass-backend/internal/service/auth/service.go index 39512b0a..b0ba1613 100644 --- a/basaltpass-backend/internal/service/auth/service.go +++ b/basaltpass-backend/internal/service/auth/service.go @@ -9,6 +9,7 @@ import ( "context" "errors" "fmt" + "strings" "time" "basaltpass-backend/internal/common" @@ -26,6 +27,7 @@ var ( ErrMissingCredentials = errors.New("identifier and password required") ErrInvalidCredentials = errors.New("invalid email or password") ErrPlatformAdminOnly = errors.New("only administrators can login to platform") + ErrTenantAccountOnly = errors.New("tenant account must login via tenant portal") ErrTenantLoginDisabled = errors.New("tenant login is disabled") ErrServiceUnavailable = errors.New("authentication service temporarily unavailable") ) @@ -44,6 +46,9 @@ func normalizeLoginQueryError(err error) error { // Register creates a new user with hashed password. func (s Service) Register(req RegisterRequest) (*model.User, error) { + if !settingssvc.GetBool("auth.enable_register", true) { + return nil, errors.New("registration is disabled") + } if req.Email == "" && req.Phone == "" { return nil, errors.New("email or phone required") } @@ -74,21 +79,48 @@ func (s Service) Register(req RegisterRequest) (*model.User, error) { } isFirstUser := userCount == 0 - // 检查用户是否已存在(同一个租户下的邮箱/手机号) - // 注意:admin用户(tenant_id=0)可以与普通用户使用相同的邮箱/手机号 + // 全局注册(tenant_id=0)要求邮箱/手机号全局唯一; + // 租户注册(tenant_id>0)仅要求同租户唯一,同时不能与全局账号冲突。 db := common.DB() - var existingUser model.User + normalizedEmail := strings.ToLower(strings.TrimSpace(req.Email)) - // 构建检查条件 - checkQuery := db.Where("tenant_id = ?", req.TenantID) - if req.Email != "" { - checkQuery = checkQuery.Where("email = ?", req.Email) - } else if normalizedPhone != "" { - checkQuery = checkQuery.Where("phone = ?", normalizedPhone) - } + if req.TenantID == 0 { + var conflict int64 + query := db.Model(&model.User{}) + if normalizedEmail != "" { + query = query.Where("email = ?", normalizedEmail) + } + if normalizedPhone != "" { + if normalizedEmail != "" { + query = query.Or("phone = ?", normalizedPhone) + } else { + query = query.Where("phone = ?", normalizedPhone) + } + } + if err := query.Count(&conflict).Error; err != nil { + return nil, err + } + if conflict > 0 { + return nil, errors.New("user already exists") + } + } else { + if normalizedEmail != "" { + var sameTenant int64 + if err := db.Model(&model.User{}).Where("email = ? AND tenant_id = ?", normalizedEmail, req.TenantID).Count(&sameTenant).Error; err != nil { + return nil, err + } + if sameTenant > 0 { + return nil, errors.New("user already exists in this tenant") + } - if err := checkQuery.First(&existingUser).Error; err == nil { - return nil, errors.New("user already exists in this tenant") + var globalConflict int64 + if err := db.Model(&model.User{}).Where("email = ? AND tenant_id = 0", normalizedEmail).Count(&globalConflict).Error; err != nil { + return nil, err + } + if globalConflict > 0 { + return nil, errors.New("email already registered as a global account") + } + } } // 开始事务 @@ -100,7 +132,7 @@ func (s Service) Register(req RegisterRequest) (*model.User, error) { }() user := &model.User{ - Email: req.Email, + Email: normalizedEmail, Phone: normalizedPhone, PasswordHash: string(hash), Nickname: "New User", @@ -156,6 +188,10 @@ func (s Service) LoginV2(req LoginRequest) (LoginResult, error) { if req.EmailOrPhone == "" || req.Password == "" { return LoginResult{}, ErrMissingCredentials } + if req.Scope == "" { + req.Scope = ConsoleScopeUser + } + identifier := strings.TrimSpace(req.EmailOrPhone) if req.TenantID > 0 { allowed, err := tenantservice.IsTenantLoginAllowed(req.TenantID) @@ -177,28 +213,68 @@ func (s Service) LoginV2(req LoginRequest) (LoginResult, error) { db := common.DB().WithContext(ctx) // 构建查询条件:email/phone + tenant_id - // 规则: - // 1. tenant_id = 0 表示平台登录,只允许 is_system_admin = true 的账号登录 - // 2. tenant_id > 0 表示租户登录,只允许该租户下 (users.tenant_id = tenant_id) 的账号登录 - query := db.Preload("Passkeys").Where("email = ? OR phone = ?", req.EmailOrPhone, req.EmailOrPhone) + query := db.Preload("Passkeys").Where("email = ? OR phone = ?", identifier, identifier) if req.TenantID == 0 { - // 平台登录:只允许全局管理员账号 + // 全局登录:仅允许 tenant_id=0 账户。 + query = query.Where("tenant_id = 0") if err := query.First(&user).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + // 兼容历史数据:若全局账号曾被错误迁移到 tenant_id>0,但具备 tenant_users 关系, + // 仍允许从全局入口登录(token tid=0)。 + var legacyUser model.User + if legacyErr := db.Where("email = ? OR phone = ?", identifier, identifier). + Order("id ASC"). + First(&legacyUser).Error; legacyErr == nil { + var membershipCount int64 + if countErr := db.Model(&model.TenantUser{}). + Where("user_id = ?", legacyUser.ID). + Count(&membershipCount).Error; countErr == nil && membershipCount > 0 { + if preloadErr := db.Preload("Passkeys").First(&user, legacyUser.ID).Error; preloadErr == nil { + goto LOGIN_USER_FOUND + } + } + + if legacyUser.TenantID > 0 { + return LoginResult{}, ErrTenantAccountOnly + } + } + } return LoginResult{}, normalizeLoginQueryError(err) } - - if !user.IsSuperAdmin() { - return LoginResult{}, ErrPlatformAdminOnly - } } else { - // 租户登录:查询指定租户下的用户 + // 租户登录:优先查询指定租户下的本地账户 query = query.Where("tenant_id = ?", req.TenantID) if err := query.First(&user).Error; err != nil { - return LoginResult{}, normalizeLoginQueryError(err) + if !errors.Is(err, gorm.ErrRecordNotFound) { + return LoginResult{}, normalizeLoginQueryError(err) + } + + // 允许全局账号(tenant_id=0)通过 tenant_users 成员关系登录租户入口。 + var globalUser model.User + if gErr := db.Where("(email = ? OR phone = ?) AND tenant_id = 0", identifier, identifier). + First(&globalUser).Error; gErr != nil { + return LoginResult{}, normalizeLoginQueryError(err) + } + + var membershipCount int64 + if mErr := db.Model(&model.TenantUser{}). + Where("user_id = ? AND tenant_id = ?", globalUser.ID, req.TenantID). + Count(&membershipCount).Error; mErr != nil { + return LoginResult{}, fmt.Errorf("%w: %v", ErrServiceUnavailable, mErr) + } + if membershipCount == 0 { + return LoginResult{}, ErrInvalidCredentials + } + + if pErr := db.Preload("Passkeys").First(&user, globalUser.ID).Error; pErr != nil { + return LoginResult{}, normalizeLoginQueryError(pErr) + } } } +LOGIN_USER_FOUND: + if err := bcrypt.CompareHashAndPassword([]byte(user.PasswordHash), []byte(req.Password)); err != nil { return LoginResult{}, ErrInvalidCredentials } @@ -274,7 +350,13 @@ func (s Service) LoginV2(req LoginRequest) (LoginResult, error) { } // 不需要二次验证,直接登录 - tokens, err := GenerateTokenPair(user.ID) + tokenScope := ConsoleScopeUser + tokenTenantID := req.TenantID + if req.Scope == ConsoleScopeAdmin { + tokenScope = ConsoleScopeAdmin + tokenTenantID = 0 + } + tokens, err := GenerateTokenPairWithTenantAndScope(user.ID, tokenTenantID, tokenScope) if err != nil { return LoginResult{}, fmt.Errorf("%w: %v", ErrServiceUnavailable, err) } @@ -370,7 +452,7 @@ func (s Service) Verify2FA(req Verify2FARequest) (TokenPair, error) { default: return TokenPair{}, errors.New("unsupported 2FA type") } - return GenerateTokenPair(user.ID) + return GenerateTokenPairWithTenantAndScope(user.ID, tenantID, ConsoleScopeUser) } // setupFirstUserAsGlobalAdmin 设置第一个用户为全局管理员。 diff --git a/basaltpass-backend/internal/service/auth/service_test.go b/basaltpass-backend/internal/service/auth/service_test.go index a97257a5..2e9e65c0 100644 --- a/basaltpass-backend/internal/service/auth/service_test.go +++ b/basaltpass-backend/internal/service/auth/service_test.go @@ -1,8 +1,17 @@ package auth import ( + "errors" "os" "testing" + + "basaltpass-backend/internal/common" + "basaltpass-backend/internal/model" + + "github.com/glebarez/sqlite" + "github.com/golang-jwt/jwt/v5" + "golang.org/x/crypto/bcrypt" + "gorm.io/gorm" ) func TestGenerateTokenPair(t *testing.T) { @@ -15,3 +24,126 @@ func TestGenerateTokenPair(t *testing.T) { t.Fatalf("token pair invalid %v", err) } } + +func setupAuthLoginTestDB(t *testing.T) *gorm.DB { + t.Helper() + os.Setenv("JWT_SECRET", "test-secret-for-unit-tests") + + db, err := gorm.Open(sqlite.Open("file::memory:?cache=shared"), &gorm.Config{}) + if err != nil { + t.Fatalf("open sqlite failed: %v", err) + } + + if err := db.AutoMigrate(&model.User{}, &model.Passkey{}, &model.TenantUser{}, &model.UserTenantTOTP{}); err != nil { + t.Fatalf("auto migrate failed: %v", err) + } + + common.SetDBForTest(db) + return db +} + +func mustPasswordHash(t *testing.T, raw string) string { + t.Helper() + h, err := bcrypt.GenerateFromPassword([]byte(raw), bcrypt.DefaultCost) + if err != nil { + t.Fatalf("generate password hash failed: %v", err) + } + return string(h) +} + +func TestLoginV2GlobalPortalRejectsTenantOnlyAccount(t *testing.T) { + db := setupAuthLoginTestDB(t) + + user := model.User{ + TenantID: 100, + Email: "tenant-only@example.com", + PasswordHash: mustPasswordHash(t, "pass-123"), + Nickname: "tenant-user", + EmailVerified: true, + } + if err := db.Create(&user).Error; err != nil { + t.Fatalf("create tenant user failed: %v", err) + } + + _, err := Service{}.LoginV2(LoginRequest{ + EmailOrPhone: user.Email, + Password: "pass-123", + TenantID: 0, + Scope: ConsoleScopeUser, + }) + + if !errors.Is(err, ErrTenantAccountOnly) { + t.Fatalf("expected ErrTenantAccountOnly, got %v", err) + } +} + +func TestLoginV2GlobalPortalAllowsGlobalAccount(t *testing.T) { + db := setupAuthLoginTestDB(t) + + user := model.User{ + TenantID: 0, + Email: "global@example.com", + PasswordHash: mustPasswordHash(t, "pass-456"), + Nickname: "global-user", + EmailVerified: true, + } + if err := db.Create(&user).Error; err != nil { + t.Fatalf("create global user failed: %v", err) + } + + res, err := Service{}.LoginV2(LoginRequest{ + EmailOrPhone: user.Email, + Password: "pass-456", + TenantID: 0, + Scope: ConsoleScopeUser, + }) + if err != nil { + t.Fatalf("global login should succeed, got error: %v", err) + } + if res.UserID != user.ID { + t.Fatalf("expected user id %d, got %d", user.ID, res.UserID) + } +} + +func TestLoginV2GlobalPortalAllowsRegularUserInAdminScope(t *testing.T) { + db := setupAuthLoginTestDB(t) + + user := model.User{ + TenantID: 0, + Email: "regular-admin-scope@example.com", + PasswordHash: mustPasswordHash(t, "pass-789"), + Nickname: "regular-user", + EmailVerified: true, + } + if err := db.Create(&user).Error; err != nil { + t.Fatalf("create user failed: %v", err) + } + + res, err := Service{}.LoginV2(LoginRequest{ + EmailOrPhone: user.Email, + Password: "pass-789", + TenantID: 0, + Scope: ConsoleScopeAdmin, + }) + if err != nil { + t.Fatalf("admin-scope login for regular global user should succeed, got error: %v", err) + } + if res.UserID != user.ID { + t.Fatalf("expected user id %d, got %d", user.ID, res.UserID) + } + + accessToken, parseErr := ParseToken(res.AccessToken) + if parseErr != nil || accessToken == nil || !accessToken.Valid { + t.Fatalf("access token invalid: %v", parseErr) + } + claims, ok := accessToken.Claims.(jwt.MapClaims) + if !ok { + t.Fatalf("unexpected claims type: %T", accessToken.Claims) + } + if scope, _ := claims["scp"].(string); scope != ConsoleScopeAdmin { + t.Fatalf("expected scope %q, got %q", ConsoleScopeAdmin, scope) + } + if tenantID, _ := claims["tid"].(float64); uint(tenantID) != 0 { + t.Fatalf("expected tenant id 0, got %v", claims["tid"]) + } +} diff --git a/basaltpass-backend/internal/service/order/service.go b/basaltpass-backend/internal/service/order/service.go index 285677cb..58583304 100644 --- a/basaltpass-backend/internal/service/order/service.go +++ b/basaltpass-backend/internal/service/order/service.go @@ -64,7 +64,7 @@ func (s *OrderService) generateOrderNumber() string { } // CreateOrder 创建订单 -func (s *OrderService) CreateOrder(req *CreateOrderRequest) (*OrderResponse, error) { +func (s *OrderService) CreateOrder(req *CreateOrderRequest, tenantID uint64) (*OrderResponse, error) { var result *OrderResponse err := s.db.Transaction(func(tx *gorm.DB) error { // 验证用户 @@ -78,7 +78,11 @@ func (s *OrderService) CreateOrder(req *CreateOrderRequest) (*OrderResponse, err // 验证价格 var price model.Price - if err := tx.Preload("Plan.Product").First(&price, req.PriceID).Error; err != nil { + priceQuery := tx.Preload("Plan.Product") + if tenantID > 0 { + priceQuery = priceQuery.Where("id = ? AND tenant_id = ?", req.PriceID, tenantID) + } + if err := priceQuery.First(&price).Error; err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return fmt.Errorf("价格不存在") } @@ -92,7 +96,11 @@ func (s *OrderService) CreateOrder(req *CreateOrderRequest) (*OrderResponse, err if req.CouponCode != nil && *req.CouponCode != "" { var c model.Coupon - if err := tx.Where("code = ? AND is_active = true", *req.CouponCode).First(&c).Error; err != nil { + couponQuery := tx.Where("code = ? AND is_active = true", *req.CouponCode) + if tenantID > 0 { + couponQuery = couponQuery.Where("tenant_id = ?", tenantID) + } + if err := couponQuery.First(&c).Error; err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return fmt.Errorf("优惠券不存在或已失效") } @@ -134,6 +142,11 @@ func (s *OrderService) CreateOrder(req *CreateOrderRequest) (*OrderResponse, err } // 创建订单 + metadata := model.JSONB{"source": "web_order"} + if tenantID > 0 { + metadata["tenant_id"] = tenantID + } + order := &model.Order{ OrderNumber: s.generateOrderNumber(), UserID: req.UserID, @@ -147,7 +160,7 @@ func (s *OrderService) CreateOrder(req *CreateOrderRequest) (*OrderResponse, err Currency: price.Currency, Description: fmt.Sprintf("订阅:%s - %s", price.Plan.Product.Name, price.Plan.DisplayName), ExpiresAt: time.Now().Add(30 * time.Minute), // 30分钟内支付 - Metadata: model.JSONB{"source": "web_order"}, + Metadata: metadata, } if err := tx.Create(order).Error; err != nil { @@ -185,10 +198,16 @@ func (s *OrderService) CreateOrder(req *CreateOrderRequest) (*OrderResponse, err } // GetOrder 获取订单 -func (s *OrderService) GetOrder(userID uint, orderID uint, activate bool) (*OrderResponse, error) { +func (s *OrderService) GetOrder(userID uint, orderID uint, activate bool, tenantID uint64) (*OrderResponse, error) { + query := s.db.Preload("Price.Plan.Product").Preload("Coupon").Preload("PaymentSession"). + Where("market_orders.id = ? AND market_orders.user_id = ?", orderID, userID) + if tenantID > 0 { + query = query.Joins("JOIN market_prices ON market_prices.id = market_orders.price_id"). + Where("market_prices.tenant_id = ?", tenantID) + } + var order model.Order - if err := s.db.Preload("Price.Plan.Product").Preload("Coupon").Preload("PaymentSession"). - Where("id = ? AND user_id = ?", orderID, userID).First(&order).Error; err != nil { + if err := query.First(&order).Error; err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return nil, fmt.Errorf("订单不存在") } @@ -197,17 +216,32 @@ func (s *OrderService) GetOrder(userID uint, orderID uint, activate bool) (*Orde if activate && order.Status == model.OrderStatusPending && order.PaymentSession != nil { _ = paymentservice.FinalizeOrderPaymentBySessionForUser(userID, order.PaymentSession.StripeSessionID) - _ = s.db.Preload("Price.Plan.Product").Preload("Coupon").Preload("PaymentSession"). - Where("id = ? AND user_id = ?", orderID, userID).First(&order).Error + reloadQuery := s.db.Preload("Price.Plan.Product").Preload("Coupon").Preload("PaymentSession"). + Where("market_orders.id = ? AND market_orders.user_id = ?", orderID, userID) + if tenantID > 0 { + reloadQuery = reloadQuery.Joins("JOIN market_prices ON market_prices.id = market_orders.price_id"). + Where("market_prices.tenant_id = ?", tenantID) + } + _ = reloadQuery.First(&order).Error } if activate && order.Status == model.OrderStatusPending { _ = paymentservice.ReconcileUserOrderPaymentsFromStripe(userID) - _ = s.db.Preload("Price.Plan.Product").Preload("Coupon").Preload("PaymentSession"). - Where("id = ? AND user_id = ?", orderID, userID).First(&order).Error + reloadQuery := s.db.Preload("Price.Plan.Product").Preload("Coupon").Preload("PaymentSession"). + Where("market_orders.id = ? AND market_orders.user_id = ?", orderID, userID) + if tenantID > 0 { + reloadQuery = reloadQuery.Joins("JOIN market_prices ON market_prices.id = market_orders.price_id"). + Where("market_prices.tenant_id = ?", tenantID) + } + _ = reloadQuery.First(&order).Error if order.Status == model.OrderStatusPending && order.PaymentSession != nil { _ = paymentservice.FinalizeOrderPaymentBySessionForUser(userID, order.PaymentSession.StripeSessionID) - _ = s.db.Preload("Price.Plan.Product").Preload("Coupon").Preload("PaymentSession"). - Where("id = ? AND user_id = ?", orderID, userID).First(&order).Error + reloadQuery = s.db.Preload("Price.Plan.Product").Preload("Coupon").Preload("PaymentSession"). + Where("market_orders.id = ? AND market_orders.user_id = ?", orderID, userID) + if tenantID > 0 { + reloadQuery = reloadQuery.Joins("JOIN market_prices ON market_prices.id = market_orders.price_id"). + Where("market_prices.tenant_id = ?", tenantID) + } + _ = reloadQuery.First(&order).Error } } @@ -234,10 +268,16 @@ func (s *OrderService) GetOrder(userID uint, orderID uint, activate bool) (*Orde } // GetOrderByNumber 根据订单号获取订单 -func (s *OrderService) GetOrderByNumber(userID uint, orderNumber string) (*OrderResponse, error) { +func (s *OrderService) GetOrderByNumber(userID uint, orderNumber string, tenantID uint64) (*OrderResponse, error) { + query := s.db.Preload("Price.Plan.Product").Preload("Coupon"). + Where("market_orders.order_number = ? AND market_orders.user_id = ?", orderNumber, userID) + if tenantID > 0 { + query = query.Joins("JOIN market_prices ON market_prices.id = market_orders.price_id"). + Where("market_prices.tenant_id = ?", tenantID) + } + var order model.Order - if err := s.db.Preload("Price.Plan.Product").Preload("Coupon"). - Where("order_number = ? AND user_id = ?", orderNumber, userID).First(&order).Error; err != nil { + if err := query.First(&order).Error; err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return nil, fmt.Errorf("订单不存在") } @@ -282,10 +322,15 @@ func (s *OrderService) UpdateOrderStatus(orderID uint, status model.OrderStatus) } // ListUserOrders 获取用户订单列表 -func (s *OrderService) ListUserOrders(userID uint, limit int) ([]*OrderResponse, error) { +func (s *OrderService) ListUserOrders(userID uint, limit int, tenantID uint64) ([]*OrderResponse, error) { var orders []model.Order query := s.db.Preload("Price.Plan.Product").Preload("Coupon"). - Where("user_id = ?", userID).Order("created_at desc") + Where("market_orders.user_id = ?", userID).Order("market_orders.created_at desc") + + if tenantID > 0 { + query = query.Joins("JOIN market_prices ON market_prices.id = market_orders.price_id"). + Where("market_prices.tenant_id = ?", tenantID) + } if limit > 0 { query = query.Limit(limit) diff --git a/basaltpass-backend/internal/service/order/service_tenant_test.go b/basaltpass-backend/internal/service/order/service_tenant_test.go new file mode 100644 index 00000000..ea76c295 --- /dev/null +++ b/basaltpass-backend/internal/service/order/service_tenant_test.go @@ -0,0 +1,150 @@ +package order + +import ( + "testing" + "time" + + "basaltpass-backend/internal/common" + "basaltpass-backend/internal/model" + + "github.com/glebarez/sqlite" + "gorm.io/gorm" +) + +func setupOrderServiceTestDB(t *testing.T) *gorm.DB { + t.Helper() + + db, err := gorm.Open(sqlite.Open("file::memory:?cache=shared"), &gorm.Config{ + DisableForeignKeyConstraintWhenMigrating: true, + }) + if err != nil { + t.Fatalf("open sqlite failed: %v", err) + } + + if err := db.AutoMigrate( + &model.User{}, + &model.Product{}, + &model.Plan{}, + &model.Price{}, + &model.Coupon{}, + &model.PaymentIntent{}, + &model.PaymentSession{}, + &model.Order{}, + ); err != nil { + t.Fatalf("auto migrate failed: %v", err) + } + + common.SetDBForTest(db) + return db +} + +func createPriceForTenant(t *testing.T, db *gorm.DB, tenantID uint64, codeSuffix string, amount int64) model.Price { + t.Helper() + + product := model.Product{ + TenantID: &tenantID, + Code: "prod-" + codeSuffix, + Name: "Product " + codeSuffix, + IsActive: true, + } + if err := db.Create(&product).Error; err != nil { + t.Fatalf("create product failed: %v", err) + } + + plan := model.Plan{ + TenantID: &tenantID, + ProductID: product.ID, + Code: "plan-" + codeSuffix, + DisplayName: "Plan " + codeSuffix, + } + if err := db.Create(&plan).Error; err != nil { + t.Fatalf("create plan failed: %v", err) + } + + price := model.Price{ + TenantID: &tenantID, + PlanID: plan.ID, + Currency: "USD", + AmountCents: amount, + BillingPeriod: model.BillingPeriodMonth, + BillingInterval: 1, + UsageType: model.UsageTypeLicense, + } + if err := db.Create(&price).Error; err != nil { + t.Fatalf("create price failed: %v", err) + } + + return price +} + +func TestOrderServiceTenantIsolation(t *testing.T) { + db := setupOrderServiceTestDB(t) + + user := model.User{TenantID: 0, Email: "order-user@example.com", PasswordHash: "x"} + if err := db.Create(&user).Error; err != nil { + t.Fatalf("create user failed: %v", err) + } + + tenantA := uint64(111) + tenantB := uint64(222) + priceA := createPriceForTenant(t, db, tenantA, "a", 1000) + priceB := createPriceForTenant(t, db, tenantB, "b", 2000) + + orderA := model.Order{ + OrderNumber: "ORD-TENANT-A", + UserID: user.ID, + PriceID: priceA.ID, + Status: model.OrderStatusPending, + Quantity: 1, + BaseAmount: 1000, + DiscountAmount: 0, + TotalAmount: 1000, + Currency: "USD", + Description: "tenant-a-order", + ExpiresAt: time.Now().Add(30 * time.Minute), + } + orderB := model.Order{ + OrderNumber: "ORD-TENANT-B", + UserID: user.ID, + PriceID: priceB.ID, + Status: model.OrderStatusPending, + Quantity: 1, + BaseAmount: 2000, + DiscountAmount: 0, + TotalAmount: 2000, + Currency: "USD", + Description: "tenant-b-order", + ExpiresAt: time.Now().Add(30 * time.Minute), + } + if err := db.Create(&orderA).Error; err != nil { + t.Fatalf("create order A failed: %v", err) + } + if err := db.Create(&orderB).Error; err != nil { + t.Fatalf("create order B failed: %v", err) + } + + svc := NewOrderService(db) + + orders, err := svc.ListUserOrders(user.ID, 20, tenantA) + if err != nil { + t.Fatalf("list orders failed: %v", err) + } + if len(orders) != 1 { + t.Fatalf("expected 1 order under tenant A, got %d", len(orders)) + } + if orders[0].OrderNumber != orderA.OrderNumber { + t.Fatalf("expected tenant A order %s, got %s", orderA.OrderNumber, orders[0].OrderNumber) + } + + if _, err := svc.GetOrder(user.ID, orderB.ID, false, tenantA); err == nil { + t.Fatalf("expected tenant-isolated get to reject tenant B order under tenant A context") + } + + got, err := svc.GetOrder(user.ID, orderA.ID, false, tenantA) + if err != nil { + t.Fatalf("expected get order in matching tenant to succeed, got %v", err) + } + if got.OrderNumber != orderA.OrderNumber { + t.Fatalf("unexpected order returned: %s", got.OrderNumber) + } +} diff --git a/basaltpass-backend/internal/service/payment/service.go b/basaltpass-backend/internal/service/payment/service.go index 484fe0e2..ef9b2cd7 100644 --- a/basaltpass-backend/internal/service/payment/service.go +++ b/basaltpass-backend/internal/service/payment/service.go @@ -81,18 +81,13 @@ func generateStripeID(prefix string) string { return fmt.Sprintf("%s_%s", prefix, hex.EncodeToString(bytes)) } -func resolveTenantStripeConfigByUser(db *gorm.DB, userID uint) (*tenantStripeConfig, error) { - var user model.User - if err := db.Select("id", "tenant_id").First(&user, userID).Error; err != nil { - return nil, err - } - - if user.TenantID == 0 { +func resolveTenantStripeConfigByTenantID(db *gorm.DB, tenantID uint) (*tenantStripeConfig, error) { + if tenantID == 0 { return &tenantStripeConfig{SecretKey: "sk_test_legacy_placeholder", TenantID: 0}, nil } var tenant model.Tenant - if err := db.Select("id", "metadata").First(&tenant, user.TenantID).Error; err != nil { + if err := db.Select("id", "metadata").First(&tenant, tenantID).Error; err != nil { return nil, err } @@ -125,6 +120,15 @@ func resolveTenantStripeConfigByUser(db *gorm.DB, userID uint) (*tenantStripeCon return config, nil } +func resolveTenantStripeConfigByUser(db *gorm.DB, userID uint) (*tenantStripeConfig, error) { + var user model.User + if err := db.Select("id", "tenant_id").First(&user, userID).Error; err != nil { + return nil, err + } + + return resolveTenantStripeConfigByTenantID(db, user.TenantID) +} + func parseString(v interface{}) string { s, _ := v.(string) return strings.TrimSpace(s) @@ -437,6 +441,20 @@ func parseUIntFromAny(v interface{}) uint { return 0 } +func parseTenantIDFromRawMetadata(raw string) uint { + raw = strings.TrimSpace(raw) + if raw == "" { + return 0 + } + + var metadata map[string]interface{} + if err := json.Unmarshal([]byte(raw), &metadata); err != nil { + return 0 + } + + return parseUIntFromAny(metadata["tenant_id"]) +} + func processStripeCheckoutSessionEvent(tx *gorm.DB, eventType string, eventObject map[string]interface{}) error { stripeSessionID := parseString(eventObject["id"]) if stripeSessionID == "" { @@ -472,7 +490,8 @@ func processStripeCheckoutSessionEvent(tx *gorm.DB, eventType string, eventObjec } if !wasComplete { - if err := wallet.RechargeByCode(session.UserID, session.Currency, session.Amount); err != nil { + tenantID := parseTenantIDFromRawMetadata(session.PaymentIntent.Metadata) + if err := wallet.RechargeByCodeWithTenant(session.UserID, tenantID, session.Currency, session.Amount); err != nil { return fmt.Errorf("failed to update wallet: %w", err) } } @@ -687,8 +706,21 @@ func GetWebhookEventStatus(eventID string) (*WebhookEventStatus, error) { // CreatePaymentIntent 创建支付意图 func CreatePaymentIntent(userID uint, req CreatePaymentIntentRequest) (*model.PaymentIntent, *MockStripeResponse, error) { + return CreatePaymentIntentForTenant(userID, 0, req) +} + +// CreatePaymentIntentForTenant creates payment intent in explicit tenant context. +func CreatePaymentIntentForTenant(userID uint, tenantID uint, req CreatePaymentIntentRequest) (*model.PaymentIntent, *MockStripeResponse, error) { db := common.DB() - stripeConfig, err := resolveTenantStripeConfigByUser(db, userID) + var ( + stripeConfig *tenantStripeConfig + err error + ) + if tenantID > 0 { + stripeConfig, err = resolveTenantStripeConfigByTenantID(db, tenantID) + } else { + stripeConfig, err = resolveTenantStripeConfigByUser(db, userID) + } if err != nil { return nil, nil, err } @@ -715,8 +747,12 @@ func CreatePaymentIntent(userID uint, req CreatePaymentIntentRequest) (*model.Pa if req.Metadata == nil { req.Metadata = map[string]interface{}{} } - if stripeConfig.TenantID > 0 { - req.Metadata["tenant_id"] = strconv.FormatUint(uint64(stripeConfig.TenantID), 10) + effectiveTenantID := stripeConfig.TenantID + if tenantID > 0 { + effectiveTenantID = tenantID + } + if effectiveTenantID > 0 { + req.Metadata["tenant_id"] = strconv.FormatUint(uint64(effectiveTenantID), 10) } req.Metadata["user_id"] = strconv.FormatUint(uint64(userID), 10) @@ -766,7 +802,7 @@ func CreatePaymentIntent(userID uint, req CreatePaymentIntentRequest) (*model.Pa }, RequestBody: req, Response: stripePIBody, - Timestamp: time.Now(), + Timestamp: time.Now(), } return &paymentIntent, mockResponse, nil @@ -774,8 +810,21 @@ func CreatePaymentIntent(userID uint, req CreatePaymentIntentRequest) (*model.Pa // CreatePaymentSession 创建支付会话 func CreatePaymentSession(userID uint, req CreatePaymentSessionRequest) (*model.PaymentSession, *MockStripeResponse, error) { + return CreatePaymentSessionForTenant(userID, 0, req) +} + +// CreatePaymentSessionForTenant creates payment session in explicit tenant context. +func CreatePaymentSessionForTenant(userID uint, tenantID uint, req CreatePaymentSessionRequest) (*model.PaymentSession, *MockStripeResponse, error) { db := common.DB() - stripeConfig, err := resolveTenantStripeConfigByUser(db, userID) + var ( + stripeConfig *tenantStripeConfig + err error + ) + if tenantID > 0 { + stripeConfig, err = resolveTenantStripeConfigByTenantID(db, tenantID) + } else { + stripeConfig, err = resolveTenantStripeConfigByUser(db, userID) + } if err != nil { return nil, nil, err } @@ -786,6 +835,12 @@ func CreatePaymentSession(userID uint, req CreatePaymentSessionRequest) (*model. if err := db.Where("id = ? AND user_id = ?", req.PaymentIntentID, userID).First(&paymentIntent).Error; err != nil { return nil, nil, errors.New("[CreatePaymentSession] payment intent not found. userID: " + fmt.Sprintf("%d", userID) + " paymentIntentID: " + fmt.Sprintf("%d", req.PaymentIntentID)) } + if tenantID > 0 { + intentTenantID := parseTenantIDFromRawMetadata(paymentIntent.Metadata) + if intentTenantID != tenantID { + return nil, nil, errors.New("payment intent tenant mismatch") + } + } stripeSessionBody, err := stripeRequest(stripeConfig.SecretKey, "https://api.stripe.com/v1/checkout/sessions", buildStripeCheckoutSessionForm(&paymentIntent, req)) if err != nil { @@ -842,7 +897,7 @@ func CreatePaymentSession(userID uint, req CreatePaymentSessionRequest) (*model. }, RequestBody: req, Response: stripeSessionBody, - Timestamp: time.Now(), + Timestamp: time.Now(), } return &session, mockResponse, nil @@ -932,7 +987,8 @@ func SimulatePayment(sessionID string, success bool) (*MockStripeResponse, error // 如果支付成功,更新用户钱包 if success { if !wasComplete { - if err := wallet.RechargeByCode(session.UserID, session.Currency, session.Amount); err != nil { + tenantID := parseTenantIDFromRawMetadata(session.PaymentIntent.Metadata) + if err := wallet.RechargeByCodeWithTenant(session.UserID, tenantID, session.Currency, session.Amount); err != nil { return nil, fmt.Errorf("failed to update wallet: %w", err) } } @@ -1500,24 +1556,42 @@ func calculatePeriodEnd(start time.Time, price *model.Price) time.Time { // GetPaymentIntent 获取支付意图 func GetPaymentIntent(userID uint, paymentIntentID uint) (*model.PaymentIntent, error) { + return GetPaymentIntentForTenant(userID, paymentIntentID, 0) +} + +func GetPaymentIntentForTenant(userID uint, paymentIntentID uint, tenantID uint) (*model.PaymentIntent, error) { db := common.DB() var paymentIntent model.PaymentIntent if err := db.Where("id = ? AND user_id = ?", paymentIntentID, userID).First(&paymentIntent).Error; err != nil { return nil, err } + if tenantID > 0 { + if parseTenantIDFromRawMetadata(paymentIntent.Metadata) != tenantID { + return nil, gorm.ErrRecordNotFound + } + } return &paymentIntent, nil } // GetPaymentSession 获取支付会话 func GetPaymentSession(userID uint, sessionID string) (*model.PaymentSession, error) { + return GetPaymentSessionForTenant(userID, sessionID, 0) +} + +func GetPaymentSessionForTenant(userID uint, sessionID string, tenantID uint) (*model.PaymentSession, error) { db := common.DB() var session model.PaymentSession if err := db.Preload("PaymentIntent").Where("stripe_session_id = ? AND user_id = ?", sessionID, userID).First(&session).Error; err != nil { return nil, err } + if tenantID > 0 { + if parseTenantIDFromRawMetadata(session.PaymentIntent.Metadata) != tenantID { + return nil, gorm.ErrRecordNotFound + } + } return &session, nil } @@ -1536,6 +1610,10 @@ func GetPaymentSessionByStripeID(sessionID string) (*model.PaymentSession, error // ListPaymentIntents 获取用户的支付意图列表 func ListPaymentIntents(userID uint, limit int) ([]model.PaymentIntent, error) { + return ListPaymentIntentsForTenant(userID, 0, limit) +} + +func ListPaymentIntentsForTenant(userID uint, tenantID uint, limit int) ([]model.PaymentIntent, error) { db := common.DB() var paymentIntents []model.PaymentIntent @@ -1547,6 +1625,16 @@ func ListPaymentIntents(userID uint, limit int) ([]model.PaymentIntent, error) { if err := query.Find(&paymentIntents).Error; err != nil { return nil, err } + if tenantID == 0 { + return paymentIntents, nil + } + + filtered := make([]model.PaymentIntent, 0, len(paymentIntents)) + for _, item := range paymentIntents { + if parseTenantIDFromRawMetadata(item.Metadata) == tenantID { + filtered = append(filtered, item) + } + } - return paymentIntents, nil + return filtered, nil } diff --git a/basaltpass-backend/internal/service/tenant/tenant_service.go b/basaltpass-backend/internal/service/tenant/tenant_service.go index 386eda86..90f5bb42 100644 --- a/basaltpass-backend/internal/service/tenant/tenant_service.go +++ b/basaltpass-backend/internal/service/tenant/tenant_service.go @@ -376,33 +376,65 @@ func (s *TenantService) ListTenants(page, pageSize int, search string) ([]*Tenan func (s *TenantService) GetUserTenants(userID uint) ([]*TenantResponse, error) { var user model.User if err := s.db.Select("tenant_id").First(&user, userID).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return []*TenantResponse{}, nil + } return nil, err } - if user.TenantID == 0 { - return []*TenantResponse{}, nil + + var memberships []model.TenantUser + if err := s.db.Where("user_id = ?", userID).Order("created_at ASC").Find(&memberships).Error; err != nil { + return nil, err } - var tenant model.Tenant - if err := s.db.First(&tenant, user.TenantID).Error; err != nil { - if errors.Is(err, gorm.ErrRecordNotFound) { - return []*TenantResponse{}, nil + tenantOrder := make([]uint, 0, len(memberships)+1) + rolesByTenant := make(map[uint]model.TenantRole) + + if user.TenantID > 0 { + tenantOrder = append(tenantOrder, user.TenantID) + rolesByTenant[user.TenantID] = model.TenantRoleUser + } + + for _, membership := range memberships { + if membership.TenantID == 0 { + continue } - return nil, err + if _, ok := rolesByTenant[membership.TenantID]; !ok { + tenantOrder = append(tenantOrder, membership.TenantID) + } + rolesByTenant[membership.TenantID] = membership.Role } - resp := s.tenantToResponse(&tenant, nil) - if resp.Metadata == nil { - resp.Metadata = make(map[string]interface{}) + if len(tenantOrder) == 0 { + return []*TenantResponse{}, nil } - var tenantUser model.TenantUser - if err := s.db.Where("user_id = ? AND tenant_id = ?", userID, user.TenantID).First(&tenantUser).Error; err == nil { - resp.Metadata["user_role"] = string(tenantUser.Role) - } else { - resp.Metadata["user_role"] = string(model.TenantRoleUser) + tenantsByID := make(map[uint]*model.Tenant) + responses := make([]*TenantResponse, 0, len(tenantOrder)) + + for _, tenantID := range tenantOrder { + tenant, ok := tenantsByID[tenantID] + if !ok { + loaded := &model.Tenant{} + if err := s.db.First(loaded, tenantID).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + continue + } + return nil, err + } + tenantsByID[tenantID] = loaded + tenant = loaded + } + + resp := s.tenantToResponse(tenant, nil) + if resp.Metadata == nil { + resp.Metadata = make(map[string]interface{}) + } + resp.Metadata["user_role"] = string(rolesByTenant[tenantID]) + responses = append(responses, resp) } - return []*TenantResponse{resp}, nil + return responses, nil } // UpdateTenant 更新租户 diff --git a/basaltpass-backend/internal/service/verification/service.go b/basaltpass-backend/internal/service/verification/service.go index 65071d59..cfc8d3b4 100644 --- a/basaltpass-backend/internal/service/verification/service.go +++ b/basaltpass-backend/internal/service/verification/service.go @@ -5,6 +5,7 @@ import ( "basaltpass-backend/internal/config" "basaltpass-backend/internal/model" emailservice "basaltpass-backend/internal/service/email" + settingssvc "basaltpass-backend/internal/service/settings" tenantservice "basaltpass-backend/internal/service/tenant" "basaltpass-backend/internal/service/wallet" "basaltpass-backend/internal/utils" @@ -87,6 +88,9 @@ type CompleteSignupRequest struct { // StartSignup 开始注册流程 func (s *Service) StartSignup(req StartSignupRequest) (*StartSignupResponse, error) { + if !settingssvc.GetBool("auth.enable_register", true) { + return nil, errors.New("registration is disabled") + } if req.Email == "" && req.Phone == "" { return nil, errors.New("email or phone required") } @@ -123,6 +127,10 @@ func (s *Service) StartSignup(req StartSignupRequest) (*StartSignupResponse, err normalizedPhone = normalized } + if err := s.ensureSignupEmailAvailable(common.DB(), normalizedEmail, req.TenantID); err != nil { + return nil, err + } + // 风控评估 riskLevel := s.assessRisk(req.IP, req.UserAgent, normalizedEmail) @@ -345,6 +353,10 @@ func (s *Service) ChangeEmail(req ChangeEmailRequest) error { // CompleteSignup 完成注册 func (s *Service) CompleteSignup(req CompleteSignupRequest) (*model.User, error) { + if !settingssvc.GetBool("auth.enable_register", true) { + return nil, errors.New("registration is disabled") + } + var pendingSignup model.PendingSignup if err := common.DB().Where("id = ? AND status = ?", req.SignupID, model.SignupStatusCompleted).First(&pendingSignup).Error; err != nil { @@ -375,11 +387,9 @@ func (s *Service) CompleteSignup(req CompleteSignupRequest) (*model.User, error) } }() - // 检查邮箱和租户组合是否已被注册(同一邮箱可以在不同租户注册) - var existingUser model.User - if err := tx.Where("email = ? AND tenant_id = ?", pendingSignup.Email, pendingSignup.TenantID).First(&existingUser).Error; err == nil { + if err := s.ensureSignupEmailAvailable(tx, strings.ToLower(strings.TrimSpace(pendingSignup.Email)), pendingSignup.TenantID); err != nil { tx.Rollback() - return nil, errors.New("email already registered in this tenant") + return nil, err } // 检查是否是第一个用户 @@ -415,11 +425,6 @@ func (s *Service) CompleteSignup(req CompleteSignupRequest) (*model.User, error) return nil, err } - if err := wallet.EnsureUserCreditWalletTx(tx, user.ID); err != nil { - tx.Rollback() - return nil, err - } - // 自动处理租户邀请(如果用户存在未处理的邀请记录) var invitations []model.TenantInvitation if err := tx.Where("email = ? AND status = ?", user.Email, "pending").Find(&invitations).Error; err == nil && len(invitations) > 0 { @@ -448,6 +453,14 @@ func (s *Service) CompleteSignup(req CompleteSignupRequest) (*model.User, error) } } + // 钱包初始化放在邀请处理之后: + // 1) 若用户因邀请获得了租户身份,可正常创建租户上下文钱包; + // 2) 对纯平台用户(无租户身份)跳过钱包初始化,不阻塞注册。 + if err := wallet.EnsureUserCreditWalletTx(tx, user.ID); err != nil && !errors.Is(err, wallet.ErrNoTenantIdentity) { + tx.Rollback() + return nil, err + } + // 清理:标记注册会话为已完成,使相关挑战失效 tx.Model(&pendingSignup).Update("status", model.SignupStatusCompleted) tx.Model(&model.VerificationChallenge{}).Where("signup_id = ? AND status = ?", @@ -471,6 +484,42 @@ func (s *Service) generateSignupID() (string, error) { return hex.EncodeToString(bytes), nil } +func (s *Service) ensureSignupEmailAvailable(tx *gorm.DB, email string, tenantID uint) error { + email = strings.ToLower(strings.TrimSpace(email)) + if email == "" { + return nil + } + + if tenantID == 0 { + var exists int64 + if err := tx.Model(&model.User{}).Where("email = ?", email).Count(&exists).Error; err != nil { + return errors.New("failed to check existing account") + } + if exists > 0 { + return errors.New("email already registered") + } + return nil + } + + var sameTenant int64 + if err := tx.Model(&model.User{}).Where("email = ? AND tenant_id = ?", email, tenantID).Count(&sameTenant).Error; err != nil { + return errors.New("failed to check existing account") + } + if sameTenant > 0 { + return errors.New("email already registered in this tenant") + } + + var globalExists int64 + if err := tx.Model(&model.User{}).Where("email = ? AND tenant_id = 0", email).Count(&globalExists).Error; err != nil { + return errors.New("failed to check existing account") + } + if globalExists > 0 { + return errors.New("email already registered as a global account") + } + + return nil +} + // assessRisk 评估风险等级 func (s *Service) assessRisk(ip, userAgent, email string) string { // 这里实现简单的风险评估逻辑 diff --git a/basaltpass-backend/internal/service/wallet/service.go b/basaltpass-backend/internal/service/wallet/service.go index fda22d95..a6e7c461 100644 --- a/basaltpass-backend/internal/service/wallet/service.go +++ b/basaltpass-backend/internal/service/wallet/service.go @@ -17,6 +17,11 @@ const ( fallbackCreditCode2 = "USD" ) +var ( + ErrNoTenantIdentity = errors.New("user has no tenant identity") + ErrUserNotInTenantContext = errors.New("user does not belong to requested tenant") +) + func resolveCreditCurrency(tx *gorm.DB) (model.Currency, error) { var curr model.Currency if err := tx.Where("code = ? AND is_active = ?", CreditCurrencyCode, true).First(&curr).Error; err == nil { @@ -32,6 +37,10 @@ func resolveCreditCurrency(tx *gorm.DB) (model.Currency, error) { } func resolveUserTenantID(tx *gorm.DB, userID uint) (uint, error) { + return resolveEffectiveTenantID(tx, userID, 0) +} + +func resolveEffectiveTenantID(tx *gorm.DB, userID uint, requestedTenantID uint) (uint, error) { if userID == 0 { return 0, errors.New("invalid user id") } @@ -40,7 +49,37 @@ func resolveUserTenantID(tx *gorm.DB, userID uint) (uint, error) { if err := tx.Select("tenant_id").First(&user, userID).Error; err != nil { return 0, err } - return user.TenantID, nil + + if requestedTenantID == 0 { + if user.TenantID > 0 { + return user.TenantID, nil + } + + var membership model.TenantUser + if err := tx.Select("tenant_id").Where("user_id = ?", userID).Order("created_at ASC").First(&membership).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return 0, ErrNoTenantIdentity + } + return 0, err + } + return membership.TenantID, nil + } + + if user.TenantID == requestedTenantID { + return requestedTenantID, nil + } + + var membershipCount int64 + if err := tx.Model(&model.TenantUser{}). + Where("user_id = ? AND tenant_id = ?", userID, requestedTenantID). + Count(&membershipCount).Error; err != nil { + return 0, err + } + if membershipCount == 0 { + return 0, ErrUserNotInTenantContext + } + + return requestedTenantID, nil } // EnsureUserCreditWalletTx ensures one credit wallet exists for user under current transaction. @@ -107,6 +146,11 @@ func EnsureCreditWalletsForAllUsers() (int64, error) { // GetBalance returns wallet balance for user+currency (creates row if absent) func GetBalance(userID uint, currencyID uint) (model.Wallet, error) { + return GetBalanceWithTenant(userID, currencyID, 0) +} + +// GetBalanceWithTenant returns wallet balance in specified tenant context. +func GetBalanceWithTenant(userID uint, currencyID uint, tenantID uint) (model.Wallet, error) { // 验证货币是否存在 _, err := currency.GetCurrencyByID(currencyID) if err != nil { @@ -114,14 +158,14 @@ func GetBalance(userID uint, currencyID uint) (model.Wallet, error) { } db := common.DB() - tenantID, err := resolveUserTenantID(db, userID) + effectiveTenantID, err := resolveEffectiveTenantID(db, userID, tenantID) if err != nil { return model.Wallet{}, err } var w model.Wallet - if err := db.Where("user_id = ? AND currency_id = ? AND tenant_id = ?", userID, currencyID, tenantID).First(&w).Error; err != nil { - w = model.Wallet{TenantID: tenantID, UserID: &userID, CurrencyID: ¤cyID} + if err := db.Where("user_id = ? AND currency_id = ? AND tenant_id = ?", userID, currencyID, effectiveTenantID).First(&w).Error; err != nil { + w = model.Wallet{TenantID: effectiveTenantID, UserID: &userID, CurrencyID: ¤cyID} if err := db.Create(&w).Error; err != nil { return model.Wallet{}, err } @@ -131,16 +175,26 @@ func GetBalance(userID uint, currencyID uint) (model.Wallet, error) { // GetBalanceByCode returns wallet balance for user+currency code (convenience function) func GetBalanceByCode(userID uint, currencyCode string) (model.Wallet, error) { + return GetBalanceByCodeWithTenant(userID, 0, currencyCode) +} + +// GetBalanceByCodeWithTenant returns wallet balance by code in tenant context. +func GetBalanceByCodeWithTenant(userID uint, tenantID uint, currencyCode string) (model.Wallet, error) { // 根据代码获取货币ID curr, err := currency.GetCurrencyByCode(currencyCode) if err != nil { return model.Wallet{}, errors.New("invalid currency code") } - return GetBalance(userID, curr.ID) + return GetBalanceWithTenant(userID, curr.ID, tenantID) } // Recharge adds amount to balance and creates transaction (mock auto success) func Recharge(userID uint, currencyID uint, amount int64) error { + return RechargeWithTenant(userID, 0, currencyID, amount) +} + +// RechargeWithTenant adds amount under specified tenant context. +func RechargeWithTenant(userID uint, tenantID uint, currencyID uint, amount int64) error { if !RechargeWithdrawEnabled() { return ErrWalletRechargeWithdrawDisabled } @@ -155,15 +209,15 @@ func Recharge(userID uint, currencyID uint, amount int64) error { } db := common.DB() - tenantID, err := resolveUserTenantID(db, userID) + effectiveTenantID, err := resolveEffectiveTenantID(db, userID, tenantID) if err != nil { return err } return db.Transaction(func(tx *gorm.DB) error { var w model.Wallet - if err := tx.Where("user_id = ? AND currency_id = ? AND tenant_id = ?", userID, currencyID, tenantID). - FirstOrCreate(&w, model.Wallet{TenantID: tenantID, UserID: &userID, CurrencyID: ¤cyID}).Error; err != nil { + if err := tx.Where("user_id = ? AND currency_id = ? AND tenant_id = ?", userID, currencyID, effectiveTenantID). + FirstOrCreate(&w, model.Wallet{TenantID: effectiveTenantID, UserID: &userID, CurrencyID: ¤cyID}).Error; err != nil { return err } w.Balance += amount @@ -177,15 +231,23 @@ func Recharge(userID uint, currencyID uint, amount int64) error { // RechargeByCode adds amount to balance using currency code (convenience function) func RechargeByCode(userID uint, currencyCode string, amount int64) error { + return RechargeByCodeWithTenant(userID, 0, currencyCode, amount) +} + +func RechargeByCodeWithTenant(userID uint, tenantID uint, currencyCode string, amount int64) error { curr, err := currency.GetCurrencyByCode(currencyCode) if err != nil { return errors.New("invalid currency code") } - return Recharge(userID, curr.ID, amount) + return RechargeWithTenant(userID, tenantID, curr.ID, amount) } // Withdraw deducts amount (mock immediate success) func Withdraw(userID uint, currencyID uint, amount int64) error { + return WithdrawWithTenant(userID, 0, currencyID, amount) +} + +func WithdrawWithTenant(userID uint, tenantID uint, currencyID uint, amount int64) error { if !RechargeWithdrawEnabled() { return ErrWalletRechargeWithdrawDisabled } @@ -200,14 +262,14 @@ func Withdraw(userID uint, currencyID uint, amount int64) error { } db := common.DB() - tenantID, err := resolveUserTenantID(db, userID) + effectiveTenantID, err := resolveEffectiveTenantID(db, userID, tenantID) if err != nil { return err } return db.Transaction(func(tx *gorm.DB) error { var w model.Wallet - if err := tx.Where("user_id = ? AND currency_id = ? AND tenant_id = ?", userID, currencyID, tenantID).First(&w).Error; err != nil { + if err := tx.Where("user_id = ? AND currency_id = ? AND tenant_id = ?", userID, currencyID, effectiveTenantID).First(&w).Error; err != nil { return err } if w.Balance < amount { @@ -224,15 +286,23 @@ func Withdraw(userID uint, currencyID uint, amount int64) error { // WithdrawByCode deducts amount using currency code (convenience function) func WithdrawByCode(userID uint, currencyCode string, amount int64) error { + return WithdrawByCodeWithTenant(userID, 0, currencyCode, amount) +} + +func WithdrawByCodeWithTenant(userID uint, tenantID uint, currencyCode string, amount int64) error { curr, err := currency.GetCurrencyByCode(currencyCode) if err != nil { return errors.New("invalid currency code") } - return Withdraw(userID, curr.ID, amount) + return WithdrawWithTenant(userID, tenantID, curr.ID, amount) } // AdjustByCode changes wallet balance by delta in smallest unit and records a transaction. func AdjustByCode(userID uint, currencyCode string, delta int64, txType string, reference string) (model.Wallet, error) { + return AdjustByCodeWithTenant(userID, 0, currencyCode, delta, txType, reference) +} + +func AdjustByCodeWithTenant(userID uint, tenantID uint, currencyCode string, delta int64, txType string, reference string) (model.Wallet, error) { if delta == 0 { return model.Wallet{}, errors.New("amount must not be zero") } @@ -243,7 +313,7 @@ func AdjustByCode(userID uint, currencyCode string, delta int64, txType string, } db := common.DB() - tenantID, err := resolveUserTenantID(db, userID) + effectiveTenantID, err := resolveEffectiveTenantID(db, userID, tenantID) if err != nil { return model.Wallet{}, err } @@ -251,8 +321,8 @@ func AdjustByCode(userID uint, currencyCode string, delta int64, txType string, var updated model.Wallet err = db.Transaction(func(tx *gorm.DB) error { var w model.Wallet - if err := tx.Where("user_id = ? AND currency_id = ? AND tenant_id = ?", userID, curr.ID, tenantID). - FirstOrCreate(&w, model.Wallet{TenantID: tenantID, UserID: &userID, CurrencyID: &curr.ID}).Error; err != nil { + if err := tx.Where("user_id = ? AND currency_id = ? AND tenant_id = ?", userID, curr.ID, effectiveTenantID). + FirstOrCreate(&w, model.Wallet{TenantID: effectiveTenantID, UserID: &userID, CurrencyID: &curr.ID}).Error; err != nil { return err } @@ -297,14 +367,18 @@ func AdjustByCode(userID uint, currencyCode string, delta int64, txType string, // History returns last n transactions func History(userID uint, currencyID uint, limit int) ([]model.WalletTx, error) { + return HistoryWithTenant(userID, 0, currencyID, limit) +} + +func HistoryWithTenant(userID uint, tenantID uint, currencyID uint, limit int) ([]model.WalletTx, error) { var w model.Wallet db := common.DB() - tenantID, err := resolveUserTenantID(db, userID) + effectiveTenantID, err := resolveEffectiveTenantID(db, userID, tenantID) if err != nil { return nil, err } - if err := db.Where("user_id = ? AND currency_id = ? AND tenant_id = ?", userID, currencyID, tenantID).First(&w).Error; err != nil { + if err := db.Where("user_id = ? AND currency_id = ? AND tenant_id = ?", userID, currencyID, effectiveTenantID).First(&w).Error; err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return []model.WalletTx{}, nil } @@ -317,28 +391,36 @@ func History(userID uint, currencyID uint, limit int) ([]model.WalletTx, error) // HistoryByCode returns last n transactions using currency code (convenience function) func HistoryByCode(userID uint, currencyCode string, limit int) ([]model.WalletTx, error) { + return HistoryByCodeWithTenant(userID, 0, currencyCode, limit) +} + +func HistoryByCodeWithTenant(userID uint, tenantID uint, currencyCode string, limit int) ([]model.WalletTx, error) { curr, err := currency.GetCurrencyByCode(currencyCode) if err != nil { return nil, errors.New("invalid currency code") } - return History(userID, curr.ID, limit) + return HistoryWithTenant(userID, tenantID, curr.ID, limit) } // HistoryAllByUser returns last n transactions across all wallets owned by the user. func HistoryAllByUser(userID uint, limit int) ([]model.WalletTx, error) { + return HistoryAllByUserWithTenant(userID, 0, limit) +} + +func HistoryAllByUserWithTenant(userID uint, tenantID uint, limit int) ([]model.WalletTx, error) { if limit <= 0 { limit = 20 } db := common.DB() - tenantID, err := resolveUserTenantID(db, userID) + effectiveTenantID, err := resolveEffectiveTenantID(db, userID, tenantID) if err != nil { return nil, err } var walletIDs []uint if err := db.Model(&model.Wallet{}). - Where("user_id = ? AND tenant_id = ?", userID, tenantID). + Where("user_id = ? AND tenant_id = ?", userID, effectiveTenantID). Pluck("id", &walletIDs).Error; err != nil { return nil, err } diff --git a/basaltpass-backend/internal/service/wallet/service_tenant_test.go b/basaltpass-backend/internal/service/wallet/service_tenant_test.go new file mode 100644 index 00000000..5146f27b --- /dev/null +++ b/basaltpass-backend/internal/service/wallet/service_tenant_test.go @@ -0,0 +1,118 @@ +package wallet + +import ( + "errors" + "fmt" + "strings" + "testing" + + "basaltpass-backend/internal/common" + "basaltpass-backend/internal/model" + + "github.com/glebarez/sqlite" + "gorm.io/gorm" +) + +func setupWalletServiceTestDB(t *testing.T) *gorm.DB { + t.Helper() + + dsn := fmt.Sprintf("file:%s?mode=memory&cache=shared", strings.ReplaceAll(t.Name(), "/", "_")) + db, err := gorm.Open(sqlite.Open(dsn), &gorm.Config{}) + if err != nil { + t.Fatalf("open sqlite failed: %v", err) + } + + if err := db.AutoMigrate(&model.User{}, &model.TenantUser{}, &model.Currency{}, &model.Wallet{}, &model.WalletTx{}); err != nil { + t.Fatalf("auto migrate failed: %v", err) + } + + common.SetDBForTest(db) + return db +} + +func TestWalletTenantIsolationByActiveTenant(t *testing.T) { + db := setupWalletServiceTestDB(t) + + curr := model.Currency{Code: "USD", Name: "US Dollar", IsActive: true, DecimalPlaces: 2} + if err := db.Create(&curr).Error; err != nil { + t.Fatalf("create currency failed: %v", err) + } + + user := model.User{TenantID: 0, Email: "wallet-user@example.com", PasswordHash: "x"} + if err := db.Create(&user).Error; err != nil { + t.Fatalf("create user failed: %v", err) + } + + tenantA := uint(101) + tenantB := uint(202) + memberships := []model.TenantUser{ + {UserID: user.ID, TenantID: tenantA, Role: model.TenantRoleMember}, + {UserID: user.ID, TenantID: tenantB, Role: model.TenantRoleMember}, + } + if err := db.Create(&memberships).Error; err != nil { + t.Fatalf("create memberships failed: %v", err) + } + + walletA := model.Wallet{TenantID: tenantA, UserID: &user.ID, CurrencyID: &curr.ID, Balance: 1000} + walletB := model.Wallet{TenantID: tenantB, UserID: &user.ID, CurrencyID: &curr.ID, Balance: 3000} + if err := db.Create(&walletA).Error; err != nil { + t.Fatalf("create tenant A wallet failed: %v", err) + } + if err := db.Create(&walletB).Error; err != nil { + t.Fatalf("create tenant B wallet failed: %v", err) + } + + wa, err := GetBalanceByCodeWithTenant(user.ID, tenantA, "USD") + if err != nil { + t.Fatalf("get tenant A balance failed: %v", err) + } + if wa.TenantID != tenantA || wa.Balance != 1000 { + t.Fatalf("unexpected tenant A wallet: tenant=%d balance=%d", wa.TenantID, wa.Balance) + } + + wb, err := GetBalanceByCodeWithTenant(user.ID, tenantB, "USD") + if err != nil { + t.Fatalf("get tenant B balance failed: %v", err) + } + if wb.TenantID != tenantB || wb.Balance != 3000 { + t.Fatalf("unexpected tenant B wallet: tenant=%d balance=%d", wb.TenantID, wb.Balance) + } + + if _, err := AdjustByCodeWithTenant(user.ID, tenantA, "USD", 250, "test_adjust", ""); err != nil { + t.Fatalf("adjust tenant A wallet failed: %v", err) + } + + var afterA model.Wallet + if err := db.Where("id = ?", walletA.ID).First(&afterA).Error; err != nil { + t.Fatalf("reload tenant A wallet failed: %v", err) + } + if afterA.Balance != 1250 { + t.Fatalf("expected tenant A balance 1250, got %d", afterA.Balance) + } + + var afterB model.Wallet + if err := db.Where("id = ?", walletB.ID).First(&afterB).Error; err != nil { + t.Fatalf("reload tenant B wallet failed: %v", err) + } + if afterB.Balance != 3000 { + t.Fatalf("tenant B balance changed unexpectedly: %d", afterB.Balance) + } +} + +func TestWalletRequiresTenantIdentity(t *testing.T) { + db := setupWalletServiceTestDB(t) + + curr := model.Currency{Code: "USD", Name: "US Dollar", IsActive: true, DecimalPlaces: 2} + if err := db.Create(&curr).Error; err != nil { + t.Fatalf("create currency failed: %v", err) + } + + user := model.User{TenantID: 0, Email: "wallet-no-tenant@example.com", PasswordHash: "x"} + if err := db.Create(&user).Error; err != nil { + t.Fatalf("create user failed: %v", err) + } + + if _, err := GetBalanceByCodeWithTenant(user.ID, 0, "USD"); !errors.Is(err, ErrNoTenantIdentity) { + t.Fatalf("expected ErrNoTenantIdentity, got %v", err) + } +} diff --git a/basaltpass-backend/sdk/go/client.go b/basaltpass-backend/sdk/go/client.go index 93ea1695..ab658a82 100644 --- a/basaltpass-backend/sdk/go/client.go +++ b/basaltpass-backend/sdk/go/client.go @@ -45,6 +45,7 @@ func (e *ApiError) Error() string { type S2SUser struct { ID int64 `json:"id"` + UserUUID string `json:"user_uuid"` Email string `json:"email"` Nickname string `json:"nickname"` AvatarURL string `json:"avatar_url"` diff --git a/basaltpass-backend/sdk/python/basaltpass_s2s/__pycache__/client.cpython-314.pyc b/basaltpass-backend/sdk/python/basaltpass_s2s/__pycache__/client.cpython-314.pyc index 046c3836..36c0e0d5 100644 Binary files a/basaltpass-backend/sdk/python/basaltpass_s2s/__pycache__/client.cpython-314.pyc and b/basaltpass-backend/sdk/python/basaltpass_s2s/__pycache__/client.cpython-314.pyc differ diff --git a/basaltpass-backend/sdk/python/basaltpass_s2s/__pycache__/models.cpython-314.pyc b/basaltpass-backend/sdk/python/basaltpass_s2s/__pycache__/models.cpython-314.pyc index 5d56c5f4..d7b14521 100644 Binary files a/basaltpass-backend/sdk/python/basaltpass_s2s/__pycache__/models.cpython-314.pyc and b/basaltpass-backend/sdk/python/basaltpass_s2s/__pycache__/models.cpython-314.pyc differ diff --git a/basaltpass-backend/sdk/python/basaltpass_s2s/models.py b/basaltpass-backend/sdk/python/basaltpass_s2s/models.py index 501fb034..da02377d 100644 --- a/basaltpass-backend/sdk/python/basaltpass_s2s/models.py +++ b/basaltpass-backend/sdk/python/basaltpass_s2s/models.py @@ -4,6 +4,7 @@ @dataclass class S2SUser: id: int + user_uuid: Optional[str] = None email: Optional[str] = None nickname: Optional[str] = None avatar_url: Optional[str] = None diff --git a/basaltpass-frontend/apps/admin/dist/index.html b/basaltpass-frontend/apps/admin/dist/index.html index 71a4fb29..040bc032 100644 --- a/basaltpass-frontend/apps/admin/dist/index.html +++ b/basaltpass-frontend/apps/admin/dist/index.html @@ -25,8 +25,8 @@ .catch(function () {}) })() - - + +
diff --git a/basaltpass-frontend/apps/tenant/dist/index.html b/basaltpass-frontend/apps/tenant/dist/index.html index 377e35a0..0d6e4158 100644 --- a/basaltpass-frontend/apps/tenant/dist/index.html +++ b/basaltpass-frontend/apps/tenant/dist/index.html @@ -25,8 +25,8 @@ .catch(function () {}) })() - - + +
diff --git a/basaltpass-frontend/apps/tenant/src/router.tsx b/basaltpass-frontend/apps/tenant/src/router.tsx index 714c8a3a..dbf1c1dd 100644 --- a/basaltpass-frontend/apps/tenant/src/router.tsx +++ b/basaltpass-frontend/apps/tenant/src/router.tsx @@ -6,6 +6,7 @@ import { useI18n } from '../../../src/shared/i18n' import TenantDashboard from '../../../src/features/tenant/Dashboard' import TenantInfo from '../../../src/features/tenant/TenantInfo' +import TenantSettings from '../../../src/features/tenant/TenantSettings' import TenantApps from '../../../src/features/tenant/app/Apps' import CreateApp from '../../../src/features/tenant/app/CreateApp' import AppDetail from '../../../src/features/tenant/app/AppDetail' @@ -13,7 +14,9 @@ import AppSettings from '../../../src/features/tenant/app/AppSettings' import AppStats from '../../../src/features/tenant/app/AppStats' import TenantOAuthClients from '../../../src/features/tenant/app/OAuthClients' import TenantAutomationTokens from '../../../src/features/tenant/security/AutomationTokens' +import CrossAppTrustManagement from '../../../src/features/tenant/security/CrossAppTrustManagement' import TenantUserManagement from '../../../src/features/tenant/user/UserManagement' +import TenantGlobalUserAuthorization from '../../../src/features/tenant/user/GlobalUserAuthorization' import TenantTeamsPage from '../../../src/features/tenant/team/Teams' import TenantWalletManagement from '../../../src/features/tenant/wallet/WalletManagement' import GiftCardManagement from '../../../src/features/tenant/wallet/GiftCardManagement' @@ -86,6 +89,7 @@ export default function AppRouter() { } /> } /> + } /> } /> } /> } /> @@ -95,9 +99,11 @@ export default function AppRouter() { } /> } /> } /> + } /> } /> } /> + } /> } /> } /> } /> diff --git a/basaltpass-frontend/apps/user/dist/index.html b/basaltpass-frontend/apps/user/dist/index.html index c12ec9d2..259a98e5 100644 --- a/basaltpass-frontend/apps/user/dist/index.html +++ b/basaltpass-frontend/apps/user/dist/index.html @@ -25,8 +25,8 @@ .catch(function () {}) })() - - + +
diff --git a/basaltpass-frontend/apps/user/src/router.tsx b/basaltpass-frontend/apps/user/src/router.tsx index 055c209b..d866d10d 100644 --- a/basaltpass-frontend/apps/user/src/router.tsx +++ b/basaltpass-frontend/apps/user/src/router.tsx @@ -5,6 +5,7 @@ import Login from '../../../src/features/auth/Login' import Register from '../../../src/features/auth/Register' import TenantLogin from '../../../src/features/auth/TenantLogin' import TenantRegister from '../../../src/features/auth/TenantRegister' +import TenantJoin from '../../../src/features/auth/TenantJoin' import OauthSuccess from '../../../src/features/auth/OauthSuccess' import OAuthConsent from '../../../src/features/auth/OAuthConsent' import ResetPassword from '../../../src/features/auth/ResetPassword' @@ -90,6 +91,14 @@ export default function AppRouter() { } /> + + + + } + /> {/* Backward compatibility */} } /> + + + + } + /> } /> } /> } /> @@ -181,11 +198,11 @@ export default function AppRouter() { } /> {/* Wallet */} - } /> - } /> - } /> - } /> - } /> + } /> + } /> + } /> + } /> + } /> } /> {/* Security */} diff --git a/basaltpass-frontend/build.ps1 b/basaltpass-frontend/build.ps1 new file mode 100644 index 00000000..7bf1c401 --- /dev/null +++ b/basaltpass-frontend/build.ps1 @@ -0,0 +1,2 @@ +cd C:\Users\Administrator\Desktop\WorkPlace\BasaltPass\basaltpass-frontend +npm run build:tenant diff --git a/basaltpass-frontend/fix.py b/basaltpass-frontend/fix.py new file mode 100644 index 00000000..ee25e703 --- /dev/null +++ b/basaltpass-frontend/fix.py @@ -0,0 +1,76 @@ + +import re + +files = [ + 'C:/Users/Administrator/Desktop/WorkPlace/BasaltPass/basaltpass-frontend/src/features/tenant/components/TenantLayout.tsx', + 'C:/Users/Administrator/Desktop/WorkPlace/BasaltPass/basaltpass-frontend/src/features/admin/components/AdminLayout.tsx' +] + +for file in files: + with open(file, 'r', encoding='utf-8') as f: + content = f.read() + + headerRegex = re.compile(r'( \{/\* : NotificationProvider .*?)(\r?\n \r?\n \r?\n \r?\n )', re.DOTALL) + match = headerRegex.search(content) + if not match: + print('Not found in', file) + continue + + block = match.group(1) + block = block.replace('absolute right-0 z-50 mt-2 w-56 origin-top-right', 'absolute left-0 z-50 mb-2 bottom-full ml-4 w-56 origin-bottom-left') + + footer = '''
+
+''' + block + ''' +
+
''' + + content = content[:match.start()] + match.group(2) + content[match.end():] + + content = content.replace(' \\n ', ' \\n \\n' + footer) + content = content.replace(' \\r\\n ', ' \\r\\n \\r\\n' + footer) + + if 'TenantNavigation' in file: + content = content.replace('
\\n \\n
\\n ', '
\\n \\n
\\n \\n' + footer) + content = content.replace('
\\r\\n \\r\\n
\\r\\n ', '
\\r\\n \\r\\n
\\r\\n \\r\\n' + footer) + else: + content = content.replace('
\\n \\n
\\n ', '
\\n \\n
\\n \\n' + footer) + content = content.replace('
\\r\\n \\r\\n
\\r\\n ', '
\\r\\n \\r\\n
\\r\\n \\r\\n' + footer) + + with open(file, 'w', encoding='utf-8') as f: + f.write(content) + print('done', file) + diff --git a/basaltpass-frontend/runfix.js b/basaltpass-frontend/runfix.js new file mode 100644 index 00000000..5c084250 --- /dev/null +++ b/basaltpass-frontend/runfix.js @@ -0,0 +1,89 @@ +const fs = require('fs'); +const path = require('path'); + +const paths = [ + 'C:\\Users\\Administrator\\Desktop\\WorkPlace\\BasaltPass\\basaltpass-frontend\\src\\features\\tenant\\components\\TenantLayout.tsx', + 'C:\\Users\\Administrator\\Desktop\\WorkPlace\\BasaltPass\\basaltpass-frontend\\src\\features\\admin\\components\\AdminLayout.tsx', + 'C:\\Users\\Administrator\\Desktop\\WorkPlace\\BasaltPass\\basaltpass-frontend\\src\\features\\user\\components\\Layout.tsx' +]; + +function formatBlock(blockStr) { + // The classes for dropdown: "absolute left-0 z-50 mb-2 bottom-full ml-4 w-56 origin-bottom-left" + // Remove the old absolute positioning classes and replace them + let newBlock = blockStr.replace(/className="absolute [^"]+origin-[^"]+"/g, + 'className="absolute left-0 z-50 mb-2 bottom-full ml-4 w-56 origin-bottom-left overflow-hidden rounded-xl bg-white pt-1 shadow-lg ring-1 ring-black ring-opacity-5 focus:outline-none"' + ); + + return ` +
+
+${newBlock} +
+
`; +} + +paths.forEach(p => { + if (!fs.existsSync(p)) return; + + let content = fs.readFileSync(p, 'utf-8'); + + // Find the block in header + // Start with:
(Notification) + // Then:
(User Menu) + + let regex = /(\s*(?:\{\/\*\s*Notification[^\*]*\*\/\}\s*)?
[\s\S]*?(?:]*\/>|]*>[\s\S]*?[\s\S]*?)<\/div>\s*(?:\{\/\*\s*User dropdown[^\*]*\*\/\}\s*)?
[\s\S]*?<\/div>\s*\)\}\s*<\/div>)/; + + let match = content.match(regex); + if (!match) { + // Also match layout which has different text + regex = /(\s*\{\/\*\s*(?::|NotificationProvider|[a-zA-Z\s]*).*?\*\/\}\s*
[\s\S]*?]*\/>\s*<\/div>[\s\S]*?
[\s\S]*?<\/div>\s*\)\}\s*<\/div>)/; + match = content.match(regex); + } + + if (!match) { + console.log(`Block not found in ${p}`); + return; + } + + let block = match[1]; + + // Safety check: Ensure it's in the header + let headerMatch = content.match(//); + if (headerMatch && headerMatch[0].includes(block.trim())) { + content = content.replace(block, ''); + console.log(`Removed from header in ${p}`); + } else { + console.log(`Not removing from header (maybe already moved) in ${p}`); + return; // Already moved possibly + } + + let wrappedBlock = formatBlock(block); + + // Mobile insertion: after \s*
inside mobile Panel or absolute div + let navs = [...content.matchAll(/<\/nav>\s*<\/div>/g)]; + if (navs.length > 0) { + let idx = navs[0].index + navs[0][0].length; + content = content.slice(0, idx) + wrappedBlock + content.slice(idx); + } + + // Desktop insertion + // Look for
or
before the main
or matching typical Sidebar layouts. + // E.g., \n
\n
+ // or
+ let desktopRegex = /(<(?:Admin|Tenant|User)Navigation\b[^>]*>[\s\S]*?<\/div>\s*<\/div>)/; + let desktopMatch = content.match(desktopRegex); + if (desktopMatch) { + let idx = desktopMatch.index + desktopMatch[0].length; + content = content.slice(0, idx) + wrappedBlock + content.slice(idx); + } else { + // Fallback for User Layout + let usrRegex = /(
+ # There are multiple , we can insert after every instance except we know the first is mobile, the second is desktop + # Actually wait. + # In AdminLayout mobile: + + nav_closes = list(re.finditer(r'<\/nav>\s*<\/div>', content)) + if nav_closes: + # First one is mobile usually + idx = nav_closes[0].end() + content = content[:idx] + wrapped + content[idx:] + print('Inserted mobile') + + # In AdminLayout desktop: + #
+ #
+ #
+ # + #
+ #
+ #
+ # we want to insert right before the last + + # Find AdminNavigation/TenantNavigation + # Let's just search for it manually using a very specific regex + desktop_nav_re = list(re.finditer(r'(<(?:Admin|Tenant|User)Navigation\b[^>]*>[\s\S]*?<\/div>\s*<\/div>)', content)) + if desktop_nav_re: + # wait! Layout.tsx might not have Navigation component, it might have just `\r\n \r\n' + footerComponent + ); + content = content.replace( + ' \n \n', + ' \n \n' + footerComponent + ); + + // 3. Insert into desktop sidebar (Tenant) + content = content.replace( + '
\r\n \r\n
\r\n \r\n', + '
\r\n \r\n
\r\n \r\n' + footerComponent + ); + content = content.replace( + '
\n \n
\n \n', + '
\n \n
\n \n' + footerComponent + ); + + // 3. Insert into desktop sidebar (Admin) + content = content.replace( + '
\r\n \r\n
\r\n \r\n', + '
\r\n \r\n
\r\n \r\n' + footerComponent + ); + content = content.replace( + '
\n \n
\n \n', + '
\n \n
\n \n' + footerComponent + ); + + fs.writeFileSync(file, content, 'utf-8'); + console.log('Saved', file); + } +} + +processFile('C:/Users/Administrator/Desktop/WorkPlace/BasaltPass/basaltpass-frontend/src/features/tenant/components/TenantLayout.tsx', false); +processFile('C:/Users/Administrator/Desktop/WorkPlace/BasaltPass/basaltpass-frontend/src/features/admin/components/AdminLayout.tsx', false); diff --git a/basaltpass-frontend/scripts/rewrite_tenant.py b/basaltpass-frontend/scripts/rewrite_tenant.py new file mode 100644 index 00000000..a1c68e54 --- /dev/null +++ b/basaltpass-frontend/scripts/rewrite_tenant.py @@ -0,0 +1,37 @@ +import re + +with open(r'C:\Users\Administrator\Desktop\WorkPlace\BasaltPass\basaltpass-frontend\src\features\tenant\components\TenantLayout.tsx', 'r', encoding='utf-8') as f: + content = f.read() + +match = re.search(r'( \{\/\* .*?: NotificationProvider.*\n)(
.*?\n
\n \)}\n \n)', content, re.DOTALL) +if not match: + print('Pattern not found') + exit(1) + +extracted = match.group(0) +new_content = content.replace(extracted, '') + +# Change pop-down to pop-up +extracted = extracted.replace('origin-top-right', 'origin-bottom-left bottom-full mb-2 ml-12') +extracted = extracted.replace('right-0 mt-2', 'left-0') + +footer_block = """ +
+
+""" + extracted + """ +
+
+""" + +side_mobile_match = re.search(r'(\s*\s*)', new_content) +if side_mobile_match: + new_content = new_content.replace(side_mobile_match.group(1), side_mobile_match.group(1) + footer_block) + +side_desktop_match = re.search(r'(\s*
\s*\s*
\s*)', new_content) +if side_desktop_match: + new_content = new_content.replace(side_desktop_match.group(1), side_desktop_match.group(1) + footer_block) + +with open(r'C:\Users\Administrator\Desktop\WorkPlace\BasaltPass\basaltpass-frontend\src\features\tenant\components\TenantLayout.tsx', 'w', encoding='utf-8') as f: + f.write(new_content) + +print("Done") diff --git a/basaltpass-frontend/scripts/rewrite_v3.py b/basaltpass-frontend/scripts/rewrite_v3.py new file mode 100644 index 00000000..9c8da7cb --- /dev/null +++ b/basaltpass-frontend/scripts/rewrite_v3.py @@ -0,0 +1,66 @@ +import sys + +def process(filepath): + with open(filepath, 'r', encoding='utf-8') as f: + content = f.read() + + # We know the anchor is: + start_anchor = " {/* : NotificationProvider , */}" + if start_anchor not in content: + start_anchor = " {/* NotificationProvider */}" + if start_anchor not in content: + print(f"Skipping {filepath}, start anchor not found") + return + + # Find the end of the userMenuRef block + end_anchor = " )}\n \n \n \n \n " + if end_anchor not in content: + print(f"end anchor not found in {filepath}") + return + + start_idx = content.find(start_anchor) + end_idx = content.find(end_anchor) + len(" )}\n \n") + + extracted = content[start_idx:end_idx] + + # Create new layout by wiping extracted + content = content[:start_idx] + content[end_idx:] + + # Modify extracted: + # 1. Flip the dropdown origin + extracted = extracted.replace('absolute right-0 z-50 mt-2 w-56 origin-top-right', 'absolute left-0 z-50 mb-2 bottom-full ml-4 w-56 origin-bottom-left') + + # Render it into a flex component + # They want the layout to be Notification right-aligned next to Avatar? Actually the avatar has the username and email. + # To make it "avatar with label on the left, notification on the right": + # Let's put both into a space-x-2 block + footer_var = f''' +
+
+{extracted} +
+
+''' + # Wait, the dropdown relies on userMenuRef which must be valid. + + # Mobile sidebar + mob_nav_end = ' \n ' + if mob_nav_end in content: + content = content.replace(mob_nav_end, mob_nav_end + footer_var, 1) + + # Desktop sidebar + desk_nav_end1 = '
\n \n
\n ' + desk_nav_end2 = '
\n \n
\n ' + + if desk_nav_end1 in content: + content = content.replace(desk_nav_end1, desk_nav_end1 + footer_var, 1) + if desk_nav_end2 in content: + content = content.replace(desk_nav_end2, desk_nav_end2 + footer_var, 1) + + with open(filepath, 'w', encoding='utf-8') as f: + f.write(content) + print(f"Successfully processed {filepath}") + +process(r'C:\Users\Administrator\Desktop\WorkPlace\BasaltPass\basaltpass-frontend\src\features\tenant\components\TenantLayout.tsx') +process(r'C:\Users\Administrator\Desktop\WorkPlace\BasaltPass\basaltpass-frontend\src\features\admin\components\AdminLayout.tsx') + diff --git a/basaltpass-frontend/scripts/safe_rewrite.py b/basaltpass-frontend/scripts/safe_rewrite.py new file mode 100644 index 00000000..13b95c69 --- /dev/null +++ b/basaltpass-frontend/scripts/safe_rewrite.py @@ -0,0 +1,52 @@ +import re + +def rewrite(file_path, scope='tenant'): + with open(file_path, 'r', encoding='utf-8') as f: + content = f.read() + + match = re.search(r'(\s+\{\/\*.*?NotificationProvider.*?)\s+\n \n \n \n ', content, re.DOTALL) + if not match: + print("Could not find header extraction block in " + file_path) + return + + extracted_raw = match.group(1) + + content = content.replace(extracted_raw, '') + + # Process extracted block + replaced_block = extracted_raw.replace('absolute right-0 z-50 mt-2 w-56 origin-top-right', 'absolute left-0 z-50 mb-2 ml-4 w-56 origin-bottom-left bottom-full') + + footer_tmpl = f""" +
+
+{replaced_block} +
+
+""" + + if scope == 'tenant': + content = content.replace( + ' \n \n \n \n )}', + ' \n \n' + footer_tmpl + ' \n \n )}' + ) + content = content.replace( + '
\n \n
\n \n \n ', + '
\n \n
\n \n' + footer_tmpl + ' \n ' + ) + elif scope == 'admin': + content = content.replace( + ' \n \n \n \n )}', + ' \n \n' + footer_tmpl + ' \n \n )}' + ) + content = content.replace( + '
\n \n
\n \n \n ', + '
\n \n
\n \n' + footer_tmpl + ' \n ' + ) + + with open(file_path, 'w', encoding='utf-8') as f: + f.write(content) + +rewrite(r'C:\Users\Administrator\Desktop\WorkPlace\BasaltPass\basaltpass-frontend\src\features\tenant\components\TenantLayout.tsx', 'tenant') +rewrite(r'C:\Users\Administrator\Desktop\WorkPlace\BasaltPass\basaltpass-frontend\src\features\admin\components\AdminLayout.tsx', 'admin') + +print("Done rewrite") diff --git a/basaltpass-frontend/src/features/admin/Dashboard.tsx b/basaltpass-frontend/src/features/admin/Dashboard.tsx index f51a00ed..711b0419 100644 --- a/basaltpass-frontend/src/features/admin/Dashboard.tsx +++ b/basaltpass-frontend/src/features/admin/Dashboard.tsx @@ -3,22 +3,18 @@ import { Link } from 'react-router-dom' import { UsersIcon, WalletIcon, - BuildingOfficeIcon, - ChartBarIcon, ArrowUpIcon, - ArrowDownIcon, CurrencyDollarIcon, ClockIcon, DocumentTextIcon, - BellIcon, CreditCardIcon, ShoppingCartIcon, CubeIcon, - KeyIcon + ChevronRightIcon } from '@heroicons/react/24/outline' import { getDashboardStats, getRecentActivities, triggerAdminLivenessCheck } from '@api/admin/admin' import AdminLayout from '@features/admin/components/AdminLayout' -import { PSkeleton, PButton, PCard, PPageHeader } from '@ui' +import { PSkeleton, PButton, PCard } from '@ui' import { ROUTES } from '@constants' import { useI18n } from '@shared/i18n' @@ -75,42 +71,42 @@ export default function AdminDashboard() { { name: t('adminDashboard.quickActions.users.name'), description: t('adminDashboard.quickActions.users.description'), - href: '/admin/users', + href: ROUTES.admin.users, icon: UsersIcon, color: 'bg-blue-500' }, { name: t('adminDashboard.quickActions.wallets.name'), description: t('adminDashboard.quickActions.wallets.description'), - href: '/admin/wallets', + href: ROUTES.admin.wallets, icon: WalletIcon, color: 'bg-green-500' }, { name: t('adminDashboard.quickActions.subscriptions.name'), description: t('adminDashboard.quickActions.subscriptions.description'), - href: '/admin/subscriptions', + href: ROUTES.admin.subscriptions, icon: CreditCardIcon, color: 'bg-indigo-500' }, { name: t('adminDashboard.quickActions.apps.name'), description: t('adminDashboard.quickActions.apps.description'), - href: '/admin/apps', + href: ROUTES.admin.apps, icon: CubeIcon, color: 'bg-indigo-500' }, { name: t('adminDashboard.quickActions.products.name'), description: t('adminDashboard.quickActions.products.description'), - href: '/admin/products', + href: ROUTES.admin.products, icon: ShoppingCartIcon, color: 'bg-yellow-500' }, { name: t('adminDashboard.quickActions.logs.name'), description: t('adminDashboard.quickActions.logs.description'), - href: '/admin/logs', + href: ROUTES.admin.logs, icon: DocumentTextIcon, color: 'bg-gray-500' } @@ -254,6 +250,75 @@ export default function AdminDashboard() { } } + const activeUserRate = stats.totalUsers > 0 ? (stats.activeUsers / stats.totalUsers) * 100 : 0 + const activeSubscriptionRate = stats.totalSubscriptions > 0 ? (stats.activeSubscriptions / stats.totalSubscriptions) * 100 : 0 + + const summaryCards = [ + { + title: t('adminDashboard.stats.totalUsers'), + value: stats.totalUsers.toLocaleString(locale), + sub: `${t('adminDashboard.stats.activeUsers')}: ${stats.activeUsers.toLocaleString(locale)}`, + trend: `${activeUserRate.toFixed(1)}%`, + href: ROUTES.admin.users, + icon: UsersIcon, + iconClass: 'bg-blue-50 text-blue-700', + trendClass: 'text-blue-700', + }, + { + title: t('adminDashboard.stats.totalWallets'), + value: stats.totalWallets.toLocaleString(locale), + sub: `${t('adminDashboard.stats.todayRevenue')}: ${formatCurrency(stats.todayRevenue)}`, + trend: formatCurrency(stats.todayRevenue), + href: ROUTES.admin.wallets, + icon: WalletIcon, + iconClass: 'bg-emerald-50 text-emerald-700', + trendClass: 'text-emerald-700', + }, + { + title: t('adminDashboard.stats.totalRevenue'), + value: formatCurrency(stats.totalRevenue), + sub: `${t('adminDashboard.stats.today')}: ${formatCurrency(stats.todayRevenue)}`, + trend: formatCurrency(stats.todayRevenue), + href: ROUTES.admin.subscriptions, + icon: CurrencyDollarIcon, + iconClass: 'bg-amber-50 text-amber-700', + trendClass: 'text-amber-700', + }, + { + title: t('adminDashboard.stats.totalSubscriptions'), + value: stats.totalSubscriptions.toLocaleString(locale), + sub: `${t('adminDashboard.stats.active')}: ${stats.activeSubscriptions.toLocaleString(locale)}`, + trend: `${activeSubscriptionRate.toFixed(1)}%`, + href: ROUTES.admin.subscriptions, + icon: CreditCardIcon, + iconClass: 'bg-indigo-50 text-indigo-700', + trendClass: 'text-indigo-700', + }, + ] + + const systemStatusItems = [ + { + label: t('adminDashboard.systemStatus.availability'), + value: '99.9%', + valueClass: 'text-emerald-600', + }, + { + label: t('adminDashboard.systemStatus.totalApps'), + value: stats.totalApplications.toLocaleString(locale), + valueClass: 'text-indigo-600', + }, + { + label: t('adminDashboard.systemStatus.pendingTasks'), + value: '0', + valueClass: 'text-amber-600', + }, + { + label: t('adminDashboard.systemStatus.serviceStatus'), + value: t('adminDashboard.systemStatus.normal'), + valueClass: 'text-sky-600', + }, + ] + if (isLoading) { return ( @@ -276,215 +341,139 @@ export default function AdminDashboard() { } return ( - + {t('adminDashboard.actions.livenessCheck')}} + >
- {t('adminDashboard.actions.livenessCheck')}} - /> - {livenessTip &&

{livenessTip}

} - -
- -
-
-
- -
-
-
-
{t('adminDashboard.stats.totalUsers')}
-
{stats.totalUsers.toLocaleString(locale)}
-
-
-
-
-
- {t('adminDashboard.stats.activeUsers')}: {stats.activeUsers.toLocaleString(locale)} - - - {((stats.activeUsers / stats.totalUsers) * 100).toFixed(1)}% + +
+
+
+
+ + {t('adminDashboard.layoutTitle')} +

{t('adminDashboard.title')}

+

{t('adminDashboard.description')}

+ {livenessTip && ( +
+ {livenessTip} +
+ )}
-
-
- - - -
-
-
- -
-
-
-
{t('adminDashboard.stats.totalWallets')}
-
{stats.totalWallets.toLocaleString(locale)}
-
-
-
-
-
- {t('adminDashboard.stats.todayRevenue')}: {formatCurrency(stats.todayRevenue)} -
-
-
-
- - -
-
-
- -
-
-
-
{t('adminDashboard.stats.totalRevenue')}
-
{formatCurrency(stats.totalRevenue)}
-
-
-
-
-
- - {t('adminDashboard.stats.today')}: {formatCurrency(stats.todayRevenue)} +
+
+

{t('adminDashboard.stats.totalRevenue')}

+

{formatCurrency(stats.totalRevenue)}

+
+
+

{t('adminDashboard.stats.todayRevenue')}

+

{formatCurrency(stats.todayRevenue)}

+
- -
-
-
- -
-
-
-
{t('adminDashboard.stats.totalSubscriptions')}
-
{stats.totalSubscriptions.toLocaleString(locale)}
-
+
+ {summaryCards.map((card) => ( + +
+
+

{card.title}

+

{card.value}

+
+
+ +
-
-
-
- {t('adminDashboard.stats.active')}: {stats.activeSubscriptions.toLocaleString(locale)} - - - {((stats.activeSubscriptions / stats.totalSubscriptions) * 100).toFixed(1)}% +
+

{card.sub}

+ + + {card.trend}
-
-
- -
+ + ))} +
-
-
- -
-

{t('adminDashboard.quickActions.title')}

+
+ +
+

{t('adminDashboard.quickActions.title')}

{t('adminDashboard.quickActions.subtitle')}

-
-
- {quickActions.map((action) => ( - -
- - - -
-
-

- {action.name} -

-

{action.description}

-
- - ))} -
+
+ {quickActions.map((action) => ( + + + + +
+

{action.name}

+

{action.description}

+
+ + + ))}
-
-
- -
-

{t('adminDashboard.recentActivities.title')}

+ +
+

{t('adminDashboard.recentActivities.title')}

{t('adminDashboard.recentActivities.subtitle')}

-
-
- {recentActivities.map((activity) => ( -
-
- {getActivityIcon(activity.type)} -
-
-
- {activity.description} - {activity.user && ( - - {activity.user} - )} -
- {activity.amount && ( -
- {t('adminDashboard.recentActivities.amount')}: {formatCurrency(activity.amount)} -
- )} -
- - {activity.timestamp} -
-
+
+ {recentActivities.map((activity) => ( +
+
+ {getActivityIcon(activity.type)}
- ))} -
-
- +
+

{activity.description}

+ {activity.user &&

{activity.user}

} + {activity.amount && ( +

+ {t('adminDashboard.recentActivities.amount')}: {formatCurrency(activity.amount)} +

+ )} +

+ + {activity.timestamp} +

+
+
+ ))} +
+ {t('adminDashboard.recentActivities.viewAll')}
-
- -
-

{t('adminDashboard.systemStatus.title')}

-

{t('adminDashboard.systemStatus.subtitle')}

-
-
-
-
-
99.9%
-
{t('adminDashboard.systemStatus.availability')}
-
-
-
{stats.totalApplications.toLocaleString(locale)}
-
{t('adminDashboard.systemStatus.totalApps')}
-
-
-
0
-
{t('adminDashboard.systemStatus.pendingTasks')}
-
-
-
{t('adminDashboard.systemStatus.normal')}
-
{t('adminDashboard.systemStatus.serviceStatus')}
-
+ +
+

{t('adminDashboard.systemStatus.title')}

+

{t('adminDashboard.systemStatus.subtitle')}

-
- +
+ {systemStatusItems.map((item) => ( +
+

{item.label}

+

{item.value}

+
+ ))} +
+
) diff --git a/basaltpass-frontend/src/features/admin/components/AdminLayout.tsx b/basaltpass-frontend/src/features/admin/components/AdminLayout.tsx index 8eb781fe..97967609 100644 --- a/basaltpass-frontend/src/features/admin/components/AdminLayout.tsx +++ b/basaltpass-frontend/src/features/admin/components/AdminLayout.tsx @@ -27,7 +27,8 @@ export default function AdminLayout({ children, title, actions }: AdminLayoutPro const [isUserMenuOpen, setIsUserMenuOpen] = useState(false) const [sidebarOpen, setSidebarOpen] = useState(false) const [showAccountSwitcher, setShowAccountSwitcher] = useState(false) - const userMenuRef = useRef(null) + const mobileUserMenuRef = useRef(null) + const desktopUserMenuRef = useRef(null) const currentSessionKey = `${user?.id || 0}:${Number(user?.tenant_id || 0)}` const handleLogout = () => { @@ -51,7 +52,9 @@ export default function AdminLayout({ children, title, actions }: AdminLayoutPro const handlePointerDown = (event: MouseEvent) => { const target = event.target as Node | null if (!target) return - if (userMenuRef.current?.contains(target)) return + const clickedMobileMenu = mobileUserMenuRef.current?.contains(target) + const clickedDesktopMenu = desktopUserMenuRef.current?.contains(target) + if (clickedMobileMenu || clickedDesktopMenu) return setIsUserMenuOpen(false) } @@ -97,91 +100,72 @@ export default function AdminLayout({ children, title, actions }: AdminLayoutPro } return 'U' } + + const userDisplayName = user?.nickname || user?.email || t('common.user') return (
-
-
-
-
- - -
- {siteInitial} -
- {siteName} - - {title && ( - <> - / -

{title}

- - )} -
- -
- {actions} - - {isAdminPath && canAccessTenant && ( - + +
+ {sidebarOpen && ( +
+
setSidebarOpen(false)} /> +
+
+ - )} - - {isAdminPath && ( - - )} - -
- {t('common.viewNotifications')} - + {siteName} +
+
- -
+
+
+
setIsUserMenuOpen(!isUserMenuOpen)} - className="flex items-center rounded-full bg-white p-1 text-sm focus:ring-indigo-500 focus:ring-offset-2 hover:bg-gray-50" + className="flex items-center justify-start rounded-lg bg-white px-1 py-1 text-sm focus:ring-indigo-500 focus:ring-offset-2 hover:bg-gray-50" > {t('common.openUserMenu')} {user?.avatar_url ? ( {user.nickname ) : ( -
- {getUserInitial()} +
+ {getUserInitial()}
)} + {userDisplayName} {isUserMenuOpen && ( -
+

{user?.nickname || t('common.user')} @@ -208,6 +192,34 @@ export default function AdminLayout({ children, title, actions }: AdminLayoutPro {t('common.settings')} + + {isAdminPath && canAccessTenant && ( + { + setIsUserMenuOpen(false) + void switchToTenant() + }} + className="flex w-full items-center justify-start rounded-none px-4 py-2 text-sm font-medium text-indigo-600 transition-colors hover:bg-indigo-50 hover:text-indigo-700" + > + + {t('adminLayout.switchToTenantLabel')} + + )} + + {isAdminPath && ( + { + setIsUserMenuOpen(false) + switchToUser() + }} + className="flex w-full items-center justify-start rounded-none px-4 py-2 text-sm font-medium text-green-600 transition-colors hover:bg-green-50 hover:text-green-700" + > + + {t('adminLayout.switchToUserLabel')} + + )} {t('common.switchAccount')} @@ -226,7 +238,7 @@ export default function AdminLayout({ children, title, actions }: AdminLayoutPro {t('common.logout')} @@ -234,55 +246,164 @@ export default function AdminLayout({ children, title, actions }: AdminLayoutPro

)}
-
-
-
-
-
- {sidebarOpen && ( -
-
setSidebarOpen(false)} /> -
-
- +
+ {t('common.viewNotifications')} +
-
-
-
- {siteInitial} -
- {siteName} -
-
+
)} -
-
-
-
+
+
+
+
+
+ {siteInitial} +
+

{siteName}

+
+
-
+ +
+
+
+ setIsUserMenuOpen(!isUserMenuOpen)} + className="flex items-center justify-start rounded-lg bg-white px-1 py-1 text-sm focus:ring-indigo-500 focus:ring-offset-2 hover:bg-gray-50" + > + {t('common.openUserMenu')} + {user?.avatar_url ? ( + {user.nickname + ) : ( +
+ {getUserInitial()} +
+ )} + {userDisplayName} + +
+ + {isUserMenuOpen && ( +
+
+

+ {user?.nickname || t('common.user')} +

+

+ {user?.email} +

+
+ + setIsUserMenuOpen(false)} + > + + {t('common.profile')} + + + setIsUserMenuOpen(false)} + > + + {t('common.settings')} + + + {isAdminPath && canAccessTenant && ( + { + setIsUserMenuOpen(false) + void switchToTenant() + }} + className="flex w-full items-center justify-start rounded-none px-4 py-2 text-sm font-medium text-indigo-600 transition-colors hover:bg-indigo-50 hover:text-indigo-700" + > + + {t('adminLayout.switchToTenantLabel')} + + )} + + {isAdminPath && ( + { + setIsUserMenuOpen(false) + switchToUser() + }} + className="flex w-full items-center justify-start rounded-none px-4 py-2 text-sm font-medium text-green-600 transition-colors hover:bg-green-50 hover:text-green-700" + > + + {t('adminLayout.switchToUserLabel')} + + )} + + { + setShowAccountSwitcher(true) + setIsUserMenuOpen(false) + }} + className="flex w-full items-center justify-start rounded-none px-4 py-2 text-sm font-medium text-blue-600 transition-colors hover:bg-blue-50 hover:text-blue-700" + > + + {t('common.switchAccount')} + + +
+ + + + {t('common.logout')} + +
+ )} +
+ +
+ {t('common.viewNotifications')} + +
+
+
-
+
+ {(title || actions) && ( +
+ {title ?

{title}

:
} + {actions ?
{actions}
: null} +
+ )} {children}
diff --git a/basaltpass-frontend/src/features/admin/components/AdminNavigation.tsx b/basaltpass-frontend/src/features/admin/components/AdminNavigation.tsx index 431c6f15..06522585 100644 --- a/basaltpass-frontend/src/features/admin/components/AdminNavigation.tsx +++ b/basaltpass-frontend/src/features/admin/components/AdminNavigation.tsx @@ -1,4 +1,4 @@ -import { useState } from 'react' +import { useEffect, useState } from 'react' import { Link, useLocation } from 'react-router-dom' import { BuildingOfficeIcon, @@ -7,7 +7,6 @@ import { UserGroupIcon, Cog6ToothIcon, ChevronDownIcon, - ChevronRightIcon, WalletIcon, CreditCardIcon, GiftIcon, @@ -131,7 +130,10 @@ export default function AdminNavigation() { const { t } = useI18n() const location = useLocation() const { marketEnabled } = useConfig() - const [expandedSections, setExpandedSections] = useState(['adminNav.systemManagement']) + + const isCurrentPath = (href: string) => { + return location.pathname === href || location.pathname.startsWith(href + '/') + } const navigation = navigationItems.filter(item => { if (item.requiresMarket && !marketEnabled) { @@ -140,20 +142,25 @@ export default function AdminNavigation() { return true }) - const toggleSection = (sectionKey: string) => { - setExpandedSections(prev => - prev.includes(sectionKey) - ? prev.filter(name => name !== sectionKey) - : [...prev, sectionKey] - ) - } + const activeSectionKey = + navigation.find(item => item.children?.some(child => child.href && isCurrentPath(child.href)))?.key ?? null - const isCurrentPath = (href: string) => { - return location.pathname === href || location.pathname.startsWith(href + '/') + const [expandedSection, setExpandedSection] = useState( + activeSectionKey ?? 'adminNav.systemManagement' + ) + + useEffect(() => { + if (activeSectionKey) { + setExpandedSection(activeSectionKey) + } + }, [activeSectionKey]) + + const toggleSection = (sectionKey: string) => { + setExpandedSection(prev => (prev === sectionKey ? null : sectionKey)) } const renderNavigationItem = (item: NavigationItem, depth = 0) => { - const isExpanded = expandedSections.includes(item.key) + const isExpanded = expandedSection === item.key const isCurrent = item.href ? isCurrentPath(item.href) : false const hasCurrentChild = item.children?.some(child => child.href && isCurrentPath(child.href)) const sharedStateClass = hasCurrentChild || isCurrent @@ -167,22 +174,23 @@ export default function AdminNavigation() { onClick={() => toggleSection(item.key)} className={`w-full flex items-center justify-between px-3 py-2 text-left text-sm font-medium rounded-lg transition-colors ${sharedStateClass}`} style={{ paddingLeft: `${0.75 + depth * 1}rem` }} + aria-expanded={isExpanded} >
{t(item.key)}
- {isExpanded ? ( - - ) : ( - - )} + - {isExpanded && ( -
- {item.children.map(child => renderNavigationItem(child, depth + 1))} -
- )} +
+ {item.children.map(child => renderNavigationItem(child, depth + 1))} +
) } diff --git a/basaltpass-frontend/src/features/admin/team/Teams.tsx b/basaltpass-frontend/src/features/admin/team/Teams.tsx index ebc717a1..9e2c205c 100644 --- a/basaltpass-frontend/src/features/admin/team/Teams.tsx +++ b/basaltpass-frontend/src/features/admin/team/Teams.tsx @@ -129,24 +129,24 @@ export default function AdminTeamsPage() {
- {teams.map(t=> ( - + {teams.map(team=> ( +
-

{t.name}

- toggleActive(t)}> - {t.is_active? t('adminTeams.actions.deactivate') : t('adminTeams.actions.activate')} +

{team.name}

+ toggleActive(team)}> + {team.is_active? t('adminTeams.actions.deactivate') : t('adminTeams.actions.activate')}
-

{t.description}

+

{team.description}

- {t('adminTeams.card.meta', { count: t.member_count, date: new Date(t.created_at).toLocaleDateString(locale) })} + {t('adminTeams.card.meta', { count: team.member_count, date: new Date(team.created_at).toLocaleDateString(locale) })}
- } onClick={()=>openMembers(t)}>{t('adminTeams.actions.members')} - } onClick={()=>openEdit(t)}>{t('adminTeams.actions.edit')} - } onClick={()=>removeTeam(t)}>{t('adminTeams.actions.delete')} + } onClick={()=>openMembers(team)}>{t('adminTeams.actions.members')} + } onClick={()=>openEdit(team)}>{t('adminTeams.actions.edit')} + } onClick={()=>removeTeam(team)}>{t('adminTeams.actions.delete')}
))} diff --git a/basaltpass-frontend/src/features/admin/tenant/TenantDetail.tsx b/basaltpass-frontend/src/features/admin/tenant/TenantDetail.tsx index abffdf46..4f2a5741 100644 --- a/basaltpass-frontend/src/features/admin/tenant/TenantDetail.tsx +++ b/basaltpass-frontend/src/features/admin/tenant/TenantDetail.tsx @@ -5,6 +5,8 @@ import { BuildingOfficeIcon, DocumentTextIcon, CogIcon, + LinkIcon, + ClipboardDocumentIcon, ExclamationTriangleIcon, CheckCircleIcon, UserIcon, @@ -38,6 +40,7 @@ const TenantDetail: React.FC = () => { const [authLoading, setAuthLoading] = useState(false) const [authSaving, setAuthSaving] = useState(false) const [authError, setAuthError] = useState(null) + const [copiedField, setCopiedField] = useState(null) const [formData, setFormData] = useState({ name: '', description: '', @@ -100,7 +103,7 @@ const TenantDetail: React.FC = () => { setAuthSettings(settings) } catch (err: any) { console.error('Failed to fetch tenant auth settings', err) - setAuthError(err.response?.data?.error || '加载租户认证开关失败') + setAuthError(err.response?.data?.error || t('adminTenantDetail.auth.loadFailed')) setAuthSettings(null) } finally { setAuthLoading(false) @@ -130,10 +133,10 @@ const TenantDetail: React.FC = () => { allow_login: authSettings.allow_login, }) setAuthSettings(updated) - uiAlert('租户认证开关已更新') + uiAlert(t('adminTenantDetail.auth.updateSuccess')) } catch (err: any) { console.error('Failed to save tenant auth settings', err) - const message = err.response?.data?.error || '更新租户认证开关失败' + const message = err.response?.data?.error || t('adminTenantDetail.auth.updateFailed') setAuthError(message) } finally { setAuthSaving(false) @@ -223,15 +226,47 @@ const TenantDetail: React.FC = () => { return new Intl.NumberFormat(locale).format(value ?? 0) } + const copyToClipboard = async (text: string, field: string) => { + try { + await navigator.clipboard.writeText(text) + setCopiedField(field) + setTimeout(() => setCopiedField(null), 2000) + } catch { + setCopiedField(null) + } + } + + const userConsoleBaseUrl = useMemo(() => { + const configured = (import.meta as any).env?.VITE_CONSOLE_USER_URL + if (configured && typeof configured === 'string') { + return configured.replace(/\/+$/, '') + } + if (typeof window !== 'undefined') { + return window.location.origin + } + return '' + }, []) + const tenantLoginUrl = useMemo(() => { - if (!tenant?.code) { + if (!tenant?.code || !userConsoleBaseUrl) { + return '' + } + return `${userConsoleBaseUrl}/auth/tenant/${tenant.code}/login` + }, [tenant?.code, userConsoleBaseUrl]) + + const tenantRegisterUrl = useMemo(() => { + if (!tenant?.code || !userConsoleBaseUrl) { return '' } - if (typeof window === 'undefined') { - return `/auth/tenant/${tenant.code}/login` + return `${userConsoleBaseUrl}/auth/tenant/${tenant.code}/register` + }, [tenant?.code, userConsoleBaseUrl]) + + const tenantJoinUrl = useMemo(() => { + if (!tenant?.code || !userConsoleBaseUrl) { + return '' } - return `${window.location.origin}/auth/tenant/${tenant.code}/login` - }, [tenant?.code]) + return `${userConsoleBaseUrl}/auth/tenant/${tenant.code}/join` + }, [tenant?.code, userConsoleBaseUrl]) if (loading) { return ( @@ -379,16 +414,53 @@ const TenantDetail: React.FC = () => {
{tenant.owner_email}
-
{t('adminTenantDetail.meta.loginUrl')}
-
- - {tenantLoginUrl} - +
+ + {t('adminTenantDetail.meta.accessLinks')} +
+
+
+
{t('adminTenantDetail.meta.joinUrl')}
+
+ + copyToClipboard(tenantJoinUrl, 'join')} + title={t('adminTenantDetail.actions.copyLink')} + > + {copiedField === 'join' ? : } + +
+
+
+
{t('adminTenantDetail.meta.loginUrl')}
+
+ + copyToClipboard(tenantLoginUrl, 'login')} + title={t('adminTenantDetail.actions.copyLink')} + > + {copiedField === 'login' ? : } + +
+
+
+
{t('adminTenantDetail.meta.registerUrl')}
+
+ + copyToClipboard(tenantRegisterUrl, 'register')} + title={t('adminTenantDetail.actions.copyLink')} + > + {copiedField === 'register' ? : } + +
+
@@ -528,13 +600,13 @@ const TenantDetail: React.FC = () => {

- 认证访问控制 + {t('adminTenantDetail.auth.title')}

{authError && ( @@ -549,53 +621,53 @@ const TenantDetail: React.FC = () => {
handleAuthSettingChange('allow_registration', (e.target as HTMLInputElement).checked)} disabled={authSaving} />

- 关闭后,所有新用户将无法通过租户注册页创建账号。 + {t('adminTenantDetail.auth.registrationHint')}

handleAuthSettingChange('allow_login', (e.target as HTMLInputElement).checked)} disabled={authSaving} />

- 关闭后,租户下账号将无法完成登录与令牌刷新。 + {t('adminTenantDetail.auth.loginHint')}

{(!authSettings.allow_registration || !authSettings.allow_login) && (
- 当前租户存在认证限制,请确保业务方已知晓影响范围。 + {t('adminTenantDetail.auth.warning')}
)}
- 注册 + {t('adminTenantDetail.auth.registration')} - {authSettings.allow_registration ? '开启' : '关闭'} + {authSettings.allow_registration ? t('adminTenantDetail.auth.enabled') : t('adminTenantDetail.auth.disabled')} - 登录 + {t('adminTenantDetail.auth.login')} - {authSettings.allow_login ? '开启' : '关闭'} + {authSettings.allow_login ? t('adminTenantDetail.auth.enabled') : t('adminTenantDetail.auth.disabled')}
- 保存认证开关 + {t('adminTenantDetail.auth.save')}
) : ( -
暂无认证开关数据
+
{t('adminTenantDetail.auth.empty')}
)}
diff --git a/basaltpass-frontend/src/features/auth/OAuthConsent.tsx b/basaltpass-frontend/src/features/auth/OAuthConsent.tsx index 7dec0f2d..96c511d7 100644 --- a/basaltpass-frontend/src/features/auth/OAuthConsent.tsx +++ b/basaltpass-frontend/src/features/auth/OAuthConsent.tsx @@ -1,8 +1,18 @@ -import { useState, useEffect, useRef } from 'react' +import { useState, useEffect, useMemo, useRef } from 'react' import { useSearchParams } from 'react-router-dom' import client from '@api/client' import { ShieldCheckIcon, ExclamationTriangleIcon } from '@heroicons/react/24/outline' import { useI18n } from '@shared/i18n' +import { getAccessToken } from '@utils/auth' +import { + pruneExpiredUserConsoleSessions, + type UserConsoleSession, +} from '@utils/userSessions' + +type SessionOption = UserConsoleSession & { + joinRequired: boolean + tenantMismatch: boolean +} export default function OAuthConsent() { const { t } = useI18n() @@ -23,11 +33,57 @@ export default function OAuthConsent() { const privacyPolicyUrl = searchParams.get('privacy_policy_url') || '' const termsOfServiceUrl = searchParams.get('terms_of_service_url') || '' const isVerified = searchParams.get('is_verified') === 'true' + const appTenantId = Number(searchParams.get('app_tenant_id') || 0) + const currentUserJoinRequired = searchParams.get('current_user_join_required') === 'true' + + const sessionOptions = useMemo(() => { + return pruneExpiredUserConsoleSessions().map((session) => { + const tenantId = Number(session.tenant_id || 0) + const joinRequired = appTenantId > 0 && tenantId === 0 + const tenantMismatch = appTenantId > 0 && tenantId > 0 && tenantId !== appTenantId + return { + ...session, + joinRequired, + tenantMismatch, + } + }) + }, [appTenantId]) + + const [selectedSessionKey, setSelectedSessionKey] = useState('') + const [confirmJoinTenant, setConfirmJoinTenant] = useState(false) + + useEffect(() => { + if (!sessionOptions.length) { + setSelectedSessionKey('') + return + } + + const currentToken = getAccessToken() + const currentSession = currentToken + ? sessionOptions.find((session) => session.token === currentToken) + : null + + if (currentSession) { + setSelectedSessionKey(currentSession.key) + return + } + + setSelectedSessionKey((prev) => { + if (prev && sessionOptions.some((session) => session.key === prev)) { + return prev + } + return sessionOptions[0].key + }) + }, [sessionOptions]) + + const selectedSession = useMemo(() => { + return sessionOptions.find((session) => session.key === selectedSessionKey) || null + }, [sessionOptions, selectedSessionKey]) const apiBase = client.defaults.baseURL || (import.meta as any).env?.VITE_API_BASE || 'http://localhost:8101' const consentEndpoint = String(apiBase).replace(/\/$/, '') + '/api/v1/oauth/consent' - const submitConsentForm = (action: 'allow' | 'deny') => { + const submitConsentForm = (action: 'allow' | 'deny', opts?: { selectedToken?: string; joinTenant?: boolean }) => { if (submittedRef.current) return submittedRef.current = true @@ -50,6 +106,8 @@ export default function OAuthConsent() { if (state) append('state', state) if (codeChallenge) append('code_challenge', codeChallenge) if (codeChallengeMethod) append('code_challenge_method', codeChallengeMethod) + if (opts?.selectedToken) append('selected_access_token', opts.selectedToken) + if (opts?.joinTenant) append('join_tenant', 'true') document.body.appendChild(form) form.submit() @@ -60,7 +118,21 @@ export default function OAuthConsent() { setLoading(true) setError('') try { - submitConsentForm('allow') + const needsJoin = !!selectedSession?.joinRequired || (!selectedSession && currentUserJoinRequired) + if (selectedSession?.tenantMismatch) { + setError('Selected account does not belong to this application tenant. Please choose another account.') + setLoading(false) + return + } + if (needsJoin && !confirmJoinTenant) { + setError('Please confirm joining this tenant before continuing.') + setLoading(false) + return + } + submitConsentForm('allow', { + selectedToken: selectedSession?.token, + joinTenant: needsJoin && confirmJoinTenant, + }) } catch (err: any) { submittedRef.current = false setError(err.message || t('auth.oauthConsent.errors.authorizeFailed')) @@ -87,6 +159,10 @@ export default function OAuthConsent() { } }, [clientId, redirectUri, t]) + useEffect(() => { + setConfirmJoinTenant(false) + }, [selectedSessionKey]) + const getScopeDisplayName = (scope: string) => { const scopeNames: Record = { openid: t('auth.oauthConsent.scopes.openid.name'), @@ -221,6 +297,56 @@ export default function OAuthConsent() {

{t('auth.oauthConsent.permissions.title')}

+ + {sessionOptions.length > 0 ? ( +
+ + + {selectedSession?.tenantMismatch ? ( +

+ This account cannot authorize this app because its tenant does not match. +

+ ) : null} +
+ ) : null} + + {(selectedSession?.joinRequired || (!selectedSession && currentUserJoinRequired)) ? ( + + ) : null} +
{scopes.map((scope) => (
diff --git a/basaltpass-frontend/src/features/auth/TenantJoin.tsx b/basaltpass-frontend/src/features/auth/TenantJoin.tsx new file mode 100644 index 00000000..684bc348 --- /dev/null +++ b/basaltpass-frontend/src/features/auth/TenantJoin.tsx @@ -0,0 +1,172 @@ +import { useCallback, useState } from 'react' +import { Link, useParams } from 'react-router-dom' +import client from '@api/client' +import PSkeleton from '@ui/PSkeleton' +import { PAlert, PButton } from '@ui' +import { ROUTES } from '@constants' +import { useAuth } from '@contexts/AuthContext' +import { useConfig } from '@contexts/ConfigContext' +import { useI18n } from '@shared/i18n' +import { TenantLoginShell } from './tenant-login/TenantLoginShell' +import { useTenantInfo } from './tenant-login/useTenantInfo' + +function TenantJoin() { + const { tenantCode } = useParams<{ tenantCode: string }>() + const { isAuthenticated, isLoading: isAuthLoading, checkAuth } = useAuth() + const { setPageTitle } = useConfig() + const { t } = useI18n() + + const [tenantLoadError, setTenantLoadError] = useState('') + const [joinError, setJoinError] = useState('') + const [joining, setJoining] = useState(false) + const [joinSuccess, setJoinSuccess] = useState(false) + const [alreadyJoined, setAlreadyJoined] = useState(false) + + const { tenantInfo, loadingTenant } = useTenantInfo({ + tenantCode, + setPageTitle, + onError: setTenantLoadError, + }) + + const loginRedirectHref = (() => { + if (typeof window === 'undefined') { + return ROUTES.user.login + } + const fullPath = `${window.location.pathname}${window.location.search}` + return `${ROUTES.user.login}?redirect=${encodeURIComponent(fullPath)}` + })() + + const joinTenant = useCallback(async () => { + if (!tenantCode || joining) { + return + } + + setJoining(true) + setJoinError('') + + try { + const res = await client.post(`/api/v1/user/tenants/join-by-code/${encodeURIComponent(tenantCode)}`) + setAlreadyJoined(Boolean(res?.data?.already_joined)) + setJoinSuccess(true) + await checkAuth() + } catch (err: any) { + setJoinSuccess(false) + setJoinError(err?.response?.data?.error || t('auth.tenantJoin.errors.joinFailed')) + } finally { + setJoining(false) + } + }, [checkAuth, joining, t, tenantCode]) + + const handleConfirmJoin = () => { + if (!tenantInfo) { + return + } + const ok = window.confirm(t('auth.tenantJoin.confirm.prompt', { tenantName: tenantInfo.name })) + if (!ok) { + return + } + void joinTenant() + } + + if (loadingTenant || isAuthLoading) { + return + } + + if (!tenantInfo && !loadingTenant) { + return ( + {tenantLoadError || t('auth.tenant.notFoundDescription')}} + tenantCode={tenantCode} + tenantInfo={tenantInfo} + > +
+ + {t('auth.tenant.backToPlatformLogin')} + +
+
+ ) + } + + if (!isAuthenticated) { + return ( + + {t('auth.tenantJoin.loginRequired.descriptionPrefix')} {tenantInfo?.name} {t('auth.tenantJoin.loginRequired.descriptionSuffix')} + + } + tenantCode={tenantCode} + tenantInfo={tenantInfo} + > +
+ + + {t('auth.tenantJoin.actions.loginToJoin')} + + +
+ + {t('auth.tenantJoin.actions.goTenantLogin')} + +
+
+
+ ) + } + + return ( + + {t('auth.tenantJoin.success.alreadyJoinedPrefix')} {tenantInfo?.name} {t('auth.tenantJoin.success.alreadyJoinedSuffix')} + + ) : ( + <> + {t('auth.tenantJoin.success.joinedPrefix')} {tenantInfo?.name} {t('auth.tenantJoin.success.joinedSuffix')} + + ) + ) : ( + <> + {t('auth.tenantJoin.descriptionPrefix')} {tenantInfo?.name} {t('auth.tenantJoin.descriptionSuffix')} + + ) + } + tenantCode={tenantCode} + tenantInfo={tenantInfo} + > +
+ {joinError ? : null} + + {!joinSuccess ? ( + + {joining ? t('auth.tenantJoin.actions.joining') : t('auth.tenantJoin.actions.confirmJoin')} + + ) : null} + +
+ + + {t('auth.tenantJoin.actions.goTenantLogin')} + + + + + {t('auth.tenantJoin.actions.goUserDashboard')} + + +
+
+
+ ) +} + +export default TenantJoin diff --git a/basaltpass-frontend/src/features/tenant/Dashboard.tsx b/basaltpass-frontend/src/features/tenant/Dashboard.tsx index 5cb8f4c4..f59ef46a 100644 --- a/basaltpass-frontend/src/features/tenant/Dashboard.tsx +++ b/basaltpass-frontend/src/features/tenant/Dashboard.tsx @@ -175,15 +175,20 @@ export default function TenantDashboard() { } const getLoginUrl = () => { - const baseUrl = (import.meta as any).env?.VITE_CONSOLE_USER_URL || 'http://localhost:5101' + const baseUrl = (import.meta as any).env?.VITE_CONSOLE_USER_URL || window.location.origin return `${baseUrl}/auth/tenant/${tenantCode}/login` } const getRegisterUrl = () => { - const baseUrl = (import.meta as any).env?.VITE_CONSOLE_USER_URL || 'http://localhost:5101' + const baseUrl = (import.meta as any).env?.VITE_CONSOLE_USER_URL || window.location.origin return `${baseUrl}/auth/tenant/${tenantCode}/register` } + const getJoinUrl = () => { + const baseUrl = (import.meta as any).env?.VITE_CONSOLE_USER_URL || window.location.origin + return `${baseUrl}/auth/tenant/${tenantCode}/join` + } + const handleLivenessCheck = async () => { try { setIsCheckingLiveness(true) @@ -287,98 +292,127 @@ export default function TenantDashboard() { ))}
- {/* */} - {tenantCode && ( - -
-
- -

{t('tenantDashboardPage.userAccessLinks.title')}

+
+ {/* */} + {tenantCode && ( + +
+
+ +

{t('tenantDashboardPage.userAccessLinks.title')}

+
+

{t('tenantDashboardPage.userAccessLinks.description')}

-

{t('tenantDashboardPage.userAccessLinks.description')}

-
-
-
- {/* */} -
- -
-
- +
+
+ {/* */} +
+ +
+
+ +
+ copyToClipboard(getLoginUrl(), 'login')} + title={t('tenantDashboardPage.actions.copyLink')} + > + {copiedField === 'login' ? ( + + ) : ( + + )} + +
+
+ {/* 加入链接 */} +
+ +
+
+ +
+ copyToClipboard(getJoinUrl(), 'join')} + title={t('tenantDashboardPage.actions.copyLink')} + > + {copiedField === 'join' ? ( + + ) : ( + + )} +
- copyToClipboard(getLoginUrl(), 'login')} - title={t('tenantDashboardPage.actions.copyLink')} - > - {copiedField === 'login' ? ( - - ) : ( - - )} -
-
- {/* */} -
- -
-
- + {/* */} +
+ +
+
+ +
+ copyToClipboard(getRegisterUrl(), 'register')} + title={t('tenantDashboardPage.actions.copyLink')} + > + {copiedField === 'register' ? ( + + ) : ( + + )} +
- copyToClipboard(getRegisterUrl(), 'register')} - title={t('tenantDashboardPage.actions.copyLink')} - > - {copiedField === 'register' ? ( - - ) : ( - - )} -
-
- {/* */} -
-
- -
-

{t('tenantDashboardPage.userAccessLinks.tip')}

+ {/* */} +
+
+ +
+

{t('tenantDashboardPage.userAccessLinks.tip')}

+
-
- - )} + + )} - -
-

{t('tenantDashboardPage.quickActions.title')}

-
-
- {quickActions - .filter(action => !action.requiresMarket || marketEnabled) - .map((action) => ( + +
+

{t('tenantDashboardPage.quickActions.title')}

+
+
+ {quickActions + .filter(action => !action.requiresMarket || marketEnabled) + .map((action) => ( ))} -
-
+
+
+
) diff --git a/basaltpass-frontend/src/features/tenant/TenantInfo.tsx b/basaltpass-frontend/src/features/tenant/TenantInfo.tsx index 333efaf1..4542fb85 100644 --- a/basaltpass-frontend/src/features/tenant/TenantInfo.tsx +++ b/basaltpass-frontend/src/features/tenant/TenantInfo.tsx @@ -1,5 +1,6 @@ -import { useState, useEffect, type ChangeEvent } from 'react' -import { +import { useState, useEffect } from 'react' +import { Link } from 'react-router-dom' +import { BuildingOffice2Icon, CubeIcon, UsersIcon, @@ -12,13 +13,13 @@ import { InformationCircleIcon, LinkIcon, ClipboardDocumentIcon, - CheckIcon + CheckIcon, } from '@heroicons/react/24/outline' import TenantLayout from '@features/tenant/components/TenantLayout' -import { tenantApi, TenantAuthSettings, TenantInfo } from '@api/tenant/tenant' -import { PSkeleton, PBadge, PButton, PCard, PCheckbox, PInput, PPageHeader } from '@ui' +import { tenantApi, type TenantInfo } from '@api/tenant/tenant' +import { ROUTES } from '@constants' +import { PSkeleton, PBadge, PButton, PCard, PInput, PPageHeader } from '@ui' import { useI18n } from '@shared/i18n' -import { uiAlert } from '@contexts/DialogContext' export default function TenantInfoPage() { const { t, locale } = useI18n() @@ -27,30 +28,9 @@ export default function TenantInfoPage() { const [error, setError] = useState('') const [debugInfo, setDebugInfo] = useState(null) const [copiedField, setCopiedField] = useState(null) - const [stripeLoading, setStripeLoading] = useState(false) - const [stripeSaving, setStripeSaving] = useState(false) - const [stripeError, setStripeError] = useState('') - const [authLoading, setAuthLoading] = useState(false) - const [authSaving, setAuthSaving] = useState(false) - const [authError, setAuthError] = useState('') - const [authSettings, setAuthSettings] = useState(null) - const [stripeMeta, setStripeMeta] = useState({ - has_secret_key: false, - secret_key_masked: '', - has_webhook_secret: false, - webhook_secret_masked: '' - }) - const [stripeForm, setStripeForm] = useState({ - enabled: false, - publishable_key: '', - secret_key: '', - webhook_secret: '' - }) useEffect(() => { fetchTenantInfo() - fetchStripeConfig() - fetchAuthSettings() }, []) const fetchTenantInfo = async () => { @@ -61,8 +41,7 @@ export default function TenantInfoPage() { } catch (err: any) { console.error(t('tenantInfoPage.logs.fetchTenantInfoFailed'), err) setError(err.response?.data?.error || t('tenantInfoPage.errors.fetchTenantInfoFailed')) - - // , + try { const debugResponse = await tenantApi.debugUserStatus() setDebugInfo(debugResponse) @@ -74,173 +53,6 @@ export default function TenantInfoPage() { } } - const fetchStripeConfig = async () => { - try { - setStripeLoading(true) - setStripeError('') - const response = await tenantApi.getStripeConfig() - const data = response.data - setStripeForm(prev => ({ - ...prev, - enabled: data.enabled, - publishable_key: data.publishable_key || '', - secret_key: '', - webhook_secret: '' - })) - setStripeMeta({ - has_secret_key: data.has_secret_key, - secret_key_masked: data.secret_key_masked || '', - has_webhook_secret: data.has_webhook_secret, - webhook_secret_masked: data.webhook_secret_masked || '' - }) - } catch (err: any) { - console.error('Failed to fetch tenant stripe config', err) - setStripeError(err.response?.data?.error || '加载 Stripe 配置失败') - } finally { - setStripeLoading(false) - } - } - - const fetchAuthSettings = async () => { - try { - setAuthLoading(true) - setAuthError('') - const response = await tenantApi.getAuthSettings() - setAuthSettings(response.data) - } catch (err: any) { - console.error('Failed to fetch tenant auth settings', err) - setAuthError(err.response?.data?.error || '加载认证开关失败') - setAuthSettings(null) - } finally { - setAuthLoading(false) - } - } - - const handleAuthSwitchChange = (key: 'allow_registration' | 'allow_login', checked: boolean) => { - setAuthSettings(prev => { - if (!prev) { - return prev - } - return { - ...prev, - [key]: checked, - } - }) - } - - const saveAuthSettings = async () => { - if (!authSettings) { - return - } - - try { - setAuthSaving(true) - setAuthError('') - const response = await tenantApi.updateAuthSettings({ - allow_registration: authSettings.allow_registration, - allow_login: authSettings.allow_login, - }) - setAuthSettings(response.data) - uiAlert('认证开关已保存') - } catch (err: any) { - console.error('Failed to save tenant auth settings', err) - const message = err.response?.data?.error || '保存认证开关失败' - setAuthError(message) - uiAlert(message) - } finally { - setAuthSaving(false) - } - } - - const saveStripeConfig = async () => { - const publishableKey = stripeForm.publishable_key.trim() - const secretKey = stripeForm.secret_key.trim() - const webhookSecret = stripeForm.webhook_secret.trim() - - if (publishableKey && !publishableKey.startsWith('pk_')) { - setStripeError('Publishable Key 必须以 pk_ 开头') - return - } - if (secretKey && !secretKey.startsWith('sk_')) { - setStripeError('Secret Key 必须以 sk_ 开头') - return - } - if (stripeForm.enabled) { - if (!publishableKey) { - setStripeError('启用 Stripe 前请先填写 Publishable Key') - return - } - if (!secretKey && !stripeMeta.has_secret_key) { - setStripeError('启用 Stripe 前请先填写 Secret Key') - return - } - } - - try { - setStripeSaving(true) - setStripeError('') - await tenantApi.updateStripeConfig({ - enabled: stripeForm.enabled, - publishable_key: publishableKey, - secret_key: secretKey || undefined, - webhook_secret: webhookSecret || undefined - }) - uiAlert('Stripe 配置已保存') - await fetchStripeConfig() - setStripeForm(prev => ({ - ...prev, - secret_key: '', - webhook_secret: '' - })) - } catch (err: any) { - console.error('Failed to save tenant stripe config', err) - const message = err.response?.data?.error || '保存 Stripe 配置失败' - setStripeError(message) - uiAlert(message) - } finally { - setStripeSaving(false) - } - } - - const clearSecretKey = async () => { - try { - setStripeSaving(true) - setStripeError('') - await tenantApi.updateStripeConfig({ - clear_secret_key: true, - enabled: false - }) - uiAlert('Secret Key 已清除') - await fetchStripeConfig() - setStripeForm(prev => ({ ...prev, enabled: false, secret_key: '' })) - } catch (err: any) { - const message = err.response?.data?.error || '清除 Secret Key 失败' - setStripeError(message) - uiAlert(message) - } finally { - setStripeSaving(false) - } - } - - const clearWebhookSecret = async () => { - try { - setStripeSaving(true) - setStripeError('') - await tenantApi.updateStripeConfig({ - clear_webhook_secret: true - }) - uiAlert('Webhook Secret 已清除') - await fetchStripeConfig() - setStripeForm(prev => ({ ...prev, webhook_secret: '' })) - } catch (err: any) { - const message = err.response?.data?.error || '清除 Webhook Secret 失败' - setStripeError(message) - uiAlert(message) - } finally { - setStripeSaving(false) - } - } - const formatDate = (dateString: string) => { return new Date(dateString).toLocaleString(locale) } @@ -249,7 +61,7 @@ export default function TenantInfoPage() { const planNames = { free: t('tenantInfoPage.plan.free'), pro: t('tenantInfoPage.plan.pro'), - enterprise: t('tenantInfoPage.plan.enterprise') + enterprise: t('tenantInfoPage.plan.enterprise'), } return planNames[plan as keyof typeof planNames] || plan } @@ -258,7 +70,7 @@ export default function TenantInfoPage() { const planVariants: Record = { free: 'default', pro: 'info', - enterprise: 'info' + enterprise: 'info', } return planVariants[plan] || 'default' } @@ -267,7 +79,7 @@ export default function TenantInfoPage() { const statusVariants: Record = { active: 'success', suspended: 'warning', - deleted: 'error' + deleted: 'error', } return statusVariants[status] || 'default' } @@ -276,7 +88,7 @@ export default function TenantInfoPage() { const statusNames = { active: t('tenantInfoPage.status.active'), suspended: t('tenantInfoPage.status.suspended'), - deleted: t('tenantInfoPage.status.deleted') + deleted: t('tenantInfoPage.status.deleted'), } return statusNames[status as keyof typeof statusNames] || status } @@ -292,17 +104,19 @@ export default function TenantInfoPage() { } const getLoginUrl = () => { - const baseUrl = (import.meta as any).env?.VITE_CONSOLE_USER_URL || 'http://localhost:5101' + const baseUrl = (import.meta as any).env?.VITE_CONSOLE_USER_URL || window.location.origin return `${baseUrl}/auth/tenant/${tenantInfo?.code}/login` } const getRegisterUrl = () => { - const baseUrl = (import.meta as any).env?.VITE_CONSOLE_USER_URL || 'http://localhost:5101' + const baseUrl = (import.meta as any).env?.VITE_CONSOLE_USER_URL || window.location.origin return `${baseUrl}/auth/tenant/${tenantInfo?.code}/register` } - const loginEnabled = authSettings?.allow_login ?? true - const registrationEnabled = authSettings?.allow_registration ?? true + const getJoinUrl = () => { + const baseUrl = (import.meta as any).env?.VITE_CONSOLE_USER_URL || window.location.origin + return `${baseUrl}/auth/tenant/${tenantInfo?.code}/join` + } if (loading) { return ( @@ -321,8 +135,7 @@ export default function TenantInfoPage() {

{error}

- - {/* */} + {debugInfo && (

{t('tenantInfoPage.debug.title')}

@@ -344,7 +157,7 @@ export default function TenantInfoPage() {
)} - +
{t('tenantInfoPage.actions.retry')}
@@ -368,17 +181,23 @@ export default function TenantInfoPage() { } return ( - -
+ +
} /> +
+ {t('tenantInfoPage.notice.readOnlySettings')} + + {t('tenantInfoPage.notice.openSettings')} + +
+
- {/* */} -
+

{t('tenantInfoPage.sections.basicInfo')}

@@ -402,7 +221,9 @@ export default function TenantInfoPage() {
{t('tenantInfoPage.fields.status')}
- }>{getStatusDisplayName(tenantInfo.status)} + }> + {getStatusDisplayName(tenantInfo.status)} +
@@ -429,137 +250,6 @@ export default function TenantInfoPage() {
- -
-

Stripe 收款配置

- {stripeLoading && 加载中...} -
-
- {stripeError && ( -
- {stripeError} -
- )} - - - - ) => setStripeForm(prev => ({ ...prev, publishable_key: e.target.value }))} - placeholder="pk_test_..." - /> - - ) => setStripeForm(prev => ({ ...prev, secret_key: e.target.value }))} - placeholder="sk_test_..." - /> - {stripeMeta.has_secret_key && ( -
-

当前已保存 Secret Key:{stripeMeta.secret_key_masked}

- - 清除 Secret Key - -
- )} - - ) => setStripeForm(prev => ({ ...prev, webhook_secret: e.target.value }))} - placeholder="whsec_..." - /> - {stripeMeta.has_webhook_secret && ( -
-

当前已保存 Webhook Secret:{stripeMeta.webhook_secret_masked}

- - 清除 Webhook Secret - -
- )} - -
- - {stripeSaving ? '保存中...' : '保存 Stripe 配置'} - -
-
-
- - -
-

认证访问控制

- {authLoading && 加载中...} -
-
- {authError && ( -
- {authError} -
- )} - - {!authLoading && authSettings && ( - <> -
- handleAuthSwitchChange('allow_registration', (e.target as HTMLInputElement).checked)} - disabled={authSaving} - /> -

关闭后,租户注册页将拒绝新账号创建。

-
- -
- handleAuthSwitchChange('allow_login', (e.target as HTMLInputElement).checked)} - disabled={authSaving} - /> -

关闭后,租户用户无法登录和刷新令牌。

-
- - {(!authSettings.allow_registration || !authSettings.allow_login) && ( -
- 当前租户已启用部分认证限制,建议先通知业务方后再执行。 -
- )} - -
-
- 注册 - {registrationEnabled ? '开启' : '关闭'} - 登录 - {loginEnabled ? '开启' : '关闭'} -
- - {authSaving ? '保存中...' : '保存认证开关'} - -
- - )} -
-
-
- - {/* */} -
- {/* */}
@@ -570,85 +260,41 @@ export default function TenantInfoPage() {
- {/* */}
-
- - {loginEnabled ? '可访问' : '已禁用'} -
+
- - copyToClipboard(getLoginUrl(), 'login')} - title={t('tenantInfoPage.actions.copyLink')} - > - {copiedField === 'login' ? ( - - ) : ( - - )} + + copyToClipboard(getJoinUrl(), 'join')} title={t('tenantInfoPage.actions.copyLink')}> + {copiedField === 'join' ? : }
- {/* */}
-
- - {registrationEnabled ? '可访问' : '已禁用'} -
+
- - copyToClipboard(getRegisterUrl(), 'register')} - title={t('tenantInfoPage.actions.copyLink')} - > - {copiedField === 'register' ? ( - - ) : ( - - )} + + copyToClipboard(getLoginUrl(), 'login')} title={t('tenantInfoPage.actions.copyLink')}> + {copiedField === 'login' ? : }
- {/* */} -
-
- -
-

{t('tenantInfoPage.userAccessLinks.tip')}

- {(!loginEnabled || !registrationEnabled) && ( -

- 当前有认证开关被关闭,部分链接将无法正常使用。 -

- )} -
+
+ +
+ + copyToClipboard(getRegisterUrl(), 'register')} title={t('tenantInfoPage.actions.copyLink')}> + {copiedField === 'register' ? : } +
+
- {/* */} +

{t('tenantInfoPage.sections.usageStats')}

@@ -660,55 +306,44 @@ export default function TenantInfoPage() { {t('tenantInfoPage.stats.totalUsers')}
- - {tenantInfo.stats.total_users} - + {tenantInfo.stats.total_users}
- +
{t('tenantInfoPage.stats.totalApps')}
- - {tenantInfo.stats.total_apps} - + {tenantInfo.stats.total_apps}
- +
{t('tenantInfoPage.stats.activeApps')}
- - {tenantInfo.stats.active_apps} - + {tenantInfo.stats.active_apps}
- +
{t('tenantInfoPage.stats.oauthClients')}
- - {tenantInfo.stats.total_clients} - + {tenantInfo.stats.total_clients}
- +
{t('tenantInfoPage.stats.activeTokens')}
- - {tenantInfo.stats.active_tokens} - + {tenantInfo.stats.active_tokens}
- {/* */} {tenantInfo.quota && (
@@ -724,15 +359,13 @@ export default function TenantInfoPage() {
-
- +
{t('tenantInfoPage.quota.users')} @@ -741,25 +374,21 @@ export default function TenantInfoPage() {
-
- +
{t('tenantInfoPage.quota.tokensPerHour')} - - {tenantInfo.quota.max_tokens_per_hour.toLocaleString()} - + {tenantInfo.quota.max_tokens_per_hour.toLocaleString()}
-
+
)}
diff --git a/basaltpass-frontend/src/features/tenant/TenantSettings.tsx b/basaltpass-frontend/src/features/tenant/TenantSettings.tsx new file mode 100644 index 00000000..c42f2090 --- /dev/null +++ b/basaltpass-frontend/src/features/tenant/TenantSettings.tsx @@ -0,0 +1,481 @@ +import { useEffect, useState, type ChangeEvent } from 'react' +import { Link } from 'react-router-dom' +import { + BuildingOffice2Icon, + ClipboardDocumentIcon, + CheckIcon, + ExclamationTriangleIcon, + InformationCircleIcon, + LinkIcon, +} from '@heroicons/react/24/outline' +import TenantLayout from '@features/tenant/components/TenantLayout' +import { tenantApi, type TenantAuthSettings, type TenantInfo } from '@api/tenant/tenant' +import { uiAlert } from '@contexts/DialogContext' +import { ROUTES } from '@constants' +import { PBadge, PButton, PCard, PCheckbox, PInput, PPageHeader } from '@ui' +import { useI18n } from '@shared/i18n' + +export default function TenantSettingsPage() { + const { t } = useI18n() + const [tenantInfo, setTenantInfo] = useState(null) + const [tenantLoading, setTenantLoading] = useState(true) + const [tenantError, setTenantError] = useState('') + + const [copiedField, setCopiedField] = useState(null) + + const [stripeLoading, setStripeLoading] = useState(false) + const [stripeSaving, setStripeSaving] = useState(false) + const [stripeError, setStripeError] = useState('') + + const [authLoading, setAuthLoading] = useState(false) + const [authSaving, setAuthSaving] = useState(false) + const [authError, setAuthError] = useState('') + const [authSettings, setAuthSettings] = useState(null) + + const [stripeMeta, setStripeMeta] = useState({ + has_secret_key: false, + secret_key_masked: '', + has_webhook_secret: false, + webhook_secret_masked: '', + }) + + const [stripeForm, setStripeForm] = useState({ + enabled: false, + publishable_key: '', + secret_key: '', + webhook_secret: '', + }) + + useEffect(() => { + fetchTenantInfo() + fetchStripeConfig() + fetchAuthSettings() + }, []) + + const fetchTenantInfo = async () => { + try { + setTenantLoading(true) + setTenantError('') + const response = await tenantApi.getTenantInfo() + setTenantInfo(response.data) + } catch (err: any) { + setTenantError(err.response?.data?.error || t('tenantSettingsPage.errors.loadTenantInfoFailed')) + } finally { + setTenantLoading(false) + } + } + + const fetchStripeConfig = async () => { + try { + setStripeLoading(true) + setStripeError('') + const response = await tenantApi.getStripeConfig() + const data = response.data + setStripeForm(prev => ({ + ...prev, + enabled: data.enabled, + publishable_key: data.publishable_key || '', + secret_key: '', + webhook_secret: '', + })) + setStripeMeta({ + has_secret_key: data.has_secret_key, + secret_key_masked: data.secret_key_masked || '', + has_webhook_secret: data.has_webhook_secret, + webhook_secret_masked: data.webhook_secret_masked || '', + }) + } catch (err: any) { + setStripeError(err.response?.data?.error || t('tenantSettingsPage.errors.loadStripeConfigFailed')) + } finally { + setStripeLoading(false) + } + } + + const fetchAuthSettings = async () => { + try { + setAuthLoading(true) + setAuthError('') + const response = await tenantApi.getAuthSettings() + setAuthSettings(response.data) + } catch (err: any) { + setAuthError(err.response?.data?.error || t('tenantSettingsPage.errors.loadAuthSettingsFailed')) + setAuthSettings(null) + } finally { + setAuthLoading(false) + } + } + + const handleAuthSwitchChange = (key: 'allow_registration' | 'allow_login', checked: boolean) => { + setAuthSettings(prev => { + if (!prev) return prev + return { + ...prev, + [key]: checked, + } + }) + } + + const saveAuthSettings = async () => { + if (!authSettings) return + + try { + setAuthSaving(true) + setAuthError('') + const response = await tenantApi.updateAuthSettings({ + allow_registration: authSettings.allow_registration, + allow_login: authSettings.allow_login, + }) + setAuthSettings(response.data) + uiAlert(t('tenantSettingsPage.messages.authSaved')) + } catch (err: any) { + const message = err.response?.data?.error || t('tenantSettingsPage.errors.saveAuthSettingsFailed') + setAuthError(message) + uiAlert(message) + } finally { + setAuthSaving(false) + } + } + + const saveStripeConfig = async () => { + const publishableKey = stripeForm.publishable_key.trim() + const secretKey = stripeForm.secret_key.trim() + const webhookSecret = stripeForm.webhook_secret.trim() + + if (publishableKey && !publishableKey.startsWith('pk_')) { + setStripeError(t('tenantSettingsPage.errors.publishableKeyPrefix')) + return + } + + if (secretKey && !secretKey.startsWith('sk_')) { + setStripeError(t('tenantSettingsPage.errors.secretKeyPrefix')) + return + } + + if (stripeForm.enabled) { + if (!publishableKey) { + setStripeError(t('tenantSettingsPage.errors.publishableKeyRequired')) + return + } + + if (!secretKey && !stripeMeta.has_secret_key) { + setStripeError(t('tenantSettingsPage.errors.secretKeyRequired')) + return + } + } + + try { + setStripeSaving(true) + setStripeError('') + await tenantApi.updateStripeConfig({ + enabled: stripeForm.enabled, + publishable_key: publishableKey, + secret_key: secretKey || undefined, + webhook_secret: webhookSecret || undefined, + }) + uiAlert(t('tenantSettingsPage.messages.stripeSaved')) + await fetchStripeConfig() + setStripeForm(prev => ({ + ...prev, + secret_key: '', + webhook_secret: '', + })) + } catch (err: any) { + const message = err.response?.data?.error || t('tenantSettingsPage.errors.saveStripeConfigFailed') + setStripeError(message) + uiAlert(message) + } finally { + setStripeSaving(false) + } + } + + const clearSecretKey = async () => { + try { + setStripeSaving(true) + setStripeError('') + await tenantApi.updateStripeConfig({ + clear_secret_key: true, + enabled: false, + }) + uiAlert(t('tenantSettingsPage.messages.secretKeyCleared')) + await fetchStripeConfig() + setStripeForm(prev => ({ ...prev, enabled: false, secret_key: '' })) + } catch (err: any) { + const message = err.response?.data?.error || t('tenantSettingsPage.errors.clearSecretKeyFailed') + setStripeError(message) + uiAlert(message) + } finally { + setStripeSaving(false) + } + } + + const clearWebhookSecret = async () => { + try { + setStripeSaving(true) + setStripeError('') + await tenantApi.updateStripeConfig({ + clear_webhook_secret: true, + }) + uiAlert(t('tenantSettingsPage.messages.webhookSecretCleared')) + await fetchStripeConfig() + setStripeForm(prev => ({ ...prev, webhook_secret: '' })) + } catch (err: any) { + const message = err.response?.data?.error || t('tenantSettingsPage.errors.clearWebhookSecretFailed') + setStripeError(message) + uiAlert(message) + } finally { + setStripeSaving(false) + } + } + + const copyToClipboard = async (text: string, field: string) => { + try { + await navigator.clipboard.writeText(text) + setCopiedField(field) + setTimeout(() => setCopiedField(null), 2000) + } catch { + uiAlert(t('tenantSettingsPage.errors.copyFailed')) + } + } + + const getLoginUrl = () => { + const baseUrl = (import.meta as any).env?.VITE_CONSOLE_USER_URL || window.location.origin + return `${baseUrl}/auth/tenant/${tenantInfo?.code}/login` + } + + const getRegisterUrl = () => { + const baseUrl = (import.meta as any).env?.VITE_CONSOLE_USER_URL || window.location.origin + return `${baseUrl}/auth/tenant/${tenantInfo?.code}/register` + } + + const getJoinUrl = () => { + const baseUrl = (import.meta as any).env?.VITE_CONSOLE_USER_URL || window.location.origin + return `${baseUrl}/auth/tenant/${tenantInfo?.code}/join` + } + + const loginEnabled = authSettings?.allow_login ?? true + const registrationEnabled = authSettings?.allow_registration ?? true + + if (tenantLoading) { + return ( + +
{t('tenantSettingsPage.values.loading')}
+
+ ) + } + + if (tenantError || !tenantInfo) { + return ( + +
+
+ {tenantError || t('tenantSettingsPage.errors.tenantNotFound')} +
+
+
+ ) + } + + return ( + +
+ } + /> + +
+
+ + {t('tenantSettingsPage.notice.readOnlyTenantInfo')} +
+ + {t('tenantSettingsPage.notice.openTenantInfo')} + +
+ +
+ +
+

{t('tenantSettingsPage.stripe.title')}

+ {stripeLoading && {t('tenantSettingsPage.values.loading')}} +
+
+ {stripeError && ( +
+ + {stripeError} +
+ )} + + + + ) => setStripeForm(prev => ({ ...prev, publishable_key: e.target.value }))} + placeholder="pk_test_..." + /> + + ) => setStripeForm(prev => ({ ...prev, secret_key: e.target.value }))} + placeholder="sk_test_..." + /> + {stripeMeta.has_secret_key && ( +
+

{t('tenantSettingsPage.stripe.currentSecretKey', { value: stripeMeta.secret_key_masked })}

+ + {t('tenantSettingsPage.stripe.clearSecretKey')} + +
+ )} + + ) => setStripeForm(prev => ({ ...prev, webhook_secret: e.target.value }))} + placeholder="whsec_..." + /> + {stripeMeta.has_webhook_secret && ( +
+

{t('tenantSettingsPage.stripe.currentWebhookSecret', { value: stripeMeta.webhook_secret_masked })}

+ + {t('tenantSettingsPage.stripe.clearWebhookSecret')} + +
+ )} + +
+ + {stripeSaving ? t('tenantSettingsPage.values.saving') : t('tenantSettingsPage.stripe.save')} + +
+
+
+ + +
+

{t('tenantSettingsPage.auth.title')}

+ {authLoading && {t('tenantSettingsPage.values.loading')}} +
+
+ {authError && ( +
+ {authError} +
+ )} + + {!authLoading && authSettings && ( + <> +
+ handleAuthSwitchChange('allow_registration', (e.target as HTMLInputElement).checked)} + disabled={authSaving} + /> +

{t('tenantSettingsPage.auth.registrationHint')}

+
+ +
+ handleAuthSwitchChange('allow_login', (e.target as HTMLInputElement).checked)} + disabled={authSaving} + /> +

{t('tenantSettingsPage.auth.loginHint')}

+
+ + {(!authSettings.allow_registration || !authSettings.allow_login) && ( +
+ {t('tenantSettingsPage.auth.warning')} +
+ )} + +
+
+ {t('tenantSettingsPage.auth.registration')} + {registrationEnabled ? t('tenantSettingsPage.auth.enabled') : t('tenantSettingsPage.auth.disabled')} + {t('tenantSettingsPage.auth.login')} + {loginEnabled ? t('tenantSettingsPage.auth.enabled') : t('tenantSettingsPage.auth.disabled')} +
+ + {authSaving ? t('tenantSettingsPage.values.saving') : t('tenantSettingsPage.auth.save')} + +
+ + )} +
+
+
+ +
+ +
+
+ +

{t('tenantSettingsPage.accessLinks.title')}

+
+

{t('tenantSettingsPage.accessLinks.description')}

+
+
+
+
+ + {registrationEnabled ? t('tenantSettingsPage.accessLinks.accessible') : t('tenantSettingsPage.accessLinks.disabled')} +
+
+ + copyToClipboard(getJoinUrl(), 'join')}> + {copiedField === 'join' ? : } + +
+
+ +
+
+ + {loginEnabled ? t('tenantSettingsPage.accessLinks.accessible') : t('tenantSettingsPage.accessLinks.disabled')} +
+
+ + copyToClipboard(getLoginUrl(), 'login')}> + {copiedField === 'login' ? : } + +
+
+ +
+
+ + {registrationEnabled ? t('tenantSettingsPage.accessLinks.accessible') : t('tenantSettingsPage.accessLinks.disabled')} +
+
+ + copyToClipboard(getRegisterUrl(), 'register')}> + {copiedField === 'register' ? : } + +
+
+
+
+
+
+
+ ) +} diff --git a/basaltpass-frontend/src/features/tenant/app/AppDetail.tsx b/basaltpass-frontend/src/features/tenant/app/AppDetail.tsx index b7460fee..8fbea553 100644 --- a/basaltpass-frontend/src/features/tenant/app/AppDetail.tsx +++ b/basaltpass-frontend/src/features/tenant/app/AppDetail.tsx @@ -167,7 +167,7 @@ export default function AppDetail() { title={app.name} description={app.description || t('tenantAppDetail.header.fallbackDescription', { date: new Date(app.created_at).toLocaleDateString(locale), id: app.id })} /> -
+
{getStatusText(app.status)} @@ -176,7 +176,7 @@ export default function AppDetail() {
-
+
-
+
() const navigate = useNavigate() + const pageSize = 20 const [app, setApp] = useState(null) const [permissions, setPermissions] = useState([]) @@ -151,35 +153,93 @@ export default function AppPermissionManagement() { } } - // - const getCategories = () => { - const categories = Array.from(new Set(permissions.map(p => p.category))) - return categories.sort() - } + const categories = useMemo(() => { + return Array.from(new Set(permissions.map(p => p.category))).sort() + }, [permissions]) // - const filteredPermissions = permissions.filter(permission => { - const matchesSearch = debouncedSearchTerm === '' || - permission.name.toLowerCase().includes(debouncedSearchTerm.toLowerCase()) || - permission.code.toLowerCase().includes(debouncedSearchTerm.toLowerCase()) || - (permission.description && permission.description.toLowerCase().includes(debouncedSearchTerm.toLowerCase())) - - const matchesCategory = selectedCategory === '' || permission.category === selectedCategory - - return matchesSearch && matchesCategory - }) + const filteredPermissions = useMemo(() => { + return permissions.filter(permission => { + const matchesSearch = debouncedSearchTerm === '' || + permission.name.toLowerCase().includes(debouncedSearchTerm.toLowerCase()) || + permission.code.toLowerCase().includes(debouncedSearchTerm.toLowerCase()) || + (permission.description && permission.description.toLowerCase().includes(debouncedSearchTerm.toLowerCase())) - // - const getPermissionsByCategory = () => { - const categories: { [key: string]: Permission[] } = {} - filteredPermissions.forEach(permission => { - if (!categories[permission.category]) { - categories[permission.category] = [] - } - categories[permission.category].push(permission) + const matchesCategory = selectedCategory === '' || permission.category === selectedCategory + return matchesSearch && matchesCategory }) - return categories - } + }, [permissions, debouncedSearchTerm, selectedCategory]) + const { + currentPage, + pageItems: paginatedPermissions, + totalItems, + setPage, + resetPage, + } = useClientPagination(filteredPermissions, pageSize) + + const paginationBar = useManagedPaginationBar({ + currentPage, + pageSize, + totalItems, + onPageChange: setPage, + summary: ({ start, end, total }) => t('tenantRoleManagement.pagination.summary', { start, end, total }), + pageInfo: ({ currentPage: page, totalPages }) => t('tenantRoleManagement.pagination.pageInfo', { current: page, total: totalPages }), + }) + + useEffect(() => { + resetPage() + }, [debouncedSearchTerm, selectedCategory, resetPage]) + + const columns: PTableColumn[] = [ + { + key: 'code', + title: t('tenantAppPermissionManagement.table.permissionCode'), + sortable: true, + render: (row) =>
{row.code}
+ }, + { + key: 'name', + title: t('tenantAppPermissionManagement.table.permissionName'), + sortable: true, + render: (row) =>
{row.name}
+ }, + { + key: 'category', + title: t('tenantAppPermissionManagement.table.category'), + sortable: true, + render: (row) => }>{row.category} + }, + { + key: 'description', + title: t('tenantAppPermissionManagement.table.description'), + render: (row) =>
{row.description || '-'}
+ }, + { + key: 'created_at', + title: t('tenantAppPermissionManagement.table.createdAt'), + align: 'right', + sortable: true, + sorter: (a, b) => new Date(a.created_at).getTime() - new Date(b.created_at).getTime(), + render: (row) => new Date(row.created_at).toLocaleString(locale) + } + ] + + const actions: PTableAction[] = [ + { + key: 'edit', + label: t('tenantAppPermissionManagement.actions.edit'), + icon: , + variant: 'secondary', + onClick: (row) => handleEditPermission(row) + }, + { + key: 'delete', + label: t('tenantAppPermissionManagement.actions.delete'), + icon: , + variant: 'danger', + onClick: (row) => handleDeletePermission(row) + } + ] if (loading) { return ( @@ -209,117 +269,61 @@ export default function AppPermissionManagement() { return ( -
- {/* */} - } - actions={ -
- navigate(`/tenant/apps/${appId}/roles`)}>{t('tenantAppPermissionManagement.actions.roleManagement')} - navigate(`/tenant/apps/${appId}/users`)}>{t('tenantAppPermissionManagement.actions.userManagement')} - }>{t('tenantAppPermissionManagement.actions.createPermission')} -
- } - /> - - {/* */} -
-
-
- setSearchTerm((e.target as HTMLInputElement).value)} - icon={} - /> -
-
- setSelectedCategory((e.target as HTMLSelectElement).value)} - > - - {getCategories().map(category => ( - - ))} - -
-
-
- - {/* ( PTable)*/} -
- - data={filteredPermissions} - rowKey={(row) => String(row.id)} - loading={loading} - emptyText={searchTerm || selectedCategory ? t('tenantAppPermissionManagement.empty.searchNoResult') : t('tenantAppPermissionManagement.empty.noPermission')} - emptyContent={!searchTerm && !selectedCategory ? ( - }>{t('tenantAppPermissionManagement.actions.createPermission')} - ) : undefined} - size="md" - striped - defaultSort={{ key: 'name', order: 'asc' }} - columns={[ - { - key: 'name', - title: t('tenantAppPermissionManagement.table.permissionName'), - dataIndex: 'name', - sortable: true, - render: (row) => ( -
- -
-
{row.name}
-
{row.code}
-
-
- ) - }, - { - key: 'category', - title: t('tenantAppPermissionManagement.table.category'), - dataIndex: 'category', - sortable: true, - }, - { - key: 'description', - title: t('tenantAppPermissionManagement.table.description'), - dataIndex: 'description', - className: 'max-w-xl truncate', - }, - { - key: 'created_at', - title: t('tenantAppPermissionManagement.table.createdAt'), - dataIndex: 'created_at', - align: 'right', - sortable: true, - sorter: (a, b) => new Date(a.created_at).getTime() - new Date(b.created_at).getTime(), - render: (row) => new Date(row.created_at).toLocaleString(locale) - } - ]} - actions={[ - { - key: 'edit', - label: t('tenantAppPermissionManagement.actions.edit'), - icon: , - variant: 'secondary', - onClick: (row) => handleEditPermission(row) - }, - { - key: 'delete', - label: t('tenantAppPermissionManagement.actions.delete'), - icon: , - variant: 'danger', - confirm: t('tenantAppPermissionManagement.deleteActionConfirm'), - onClick: (row) => handleDeletePermission(row) - } - ]} + } + actions={ +
+ navigate(`/tenant/apps/${appId}/roles`)}>{t('tenantAppPermissionManagement.actions.roleManagement')} + navigate(`/tenant/apps/${appId}/users`)}>{t('tenantAppPermissionManagement.actions.userManagement')} + }>{t('tenantAppPermissionManagement.actions.createPermission')} +
+ } /> -
-
+ } + filter={ + + + setSelectedCategory((e.target as HTMLSelectElement).value)} + className="pl-10" + > + + {categories.map(category => ( + + ))} + +
+ } + /> + } + > + + data={paginatedPermissions} + columns={columns} + actions={actions} + rowKey={(row) => String(row.id)} + loading={loading} + emptyText={searchTerm || selectedCategory ? t('tenantAppPermissionManagement.empty.searchNoResult') : t('tenantAppPermissionManagement.empty.noPermission')} + emptyContent={!searchTerm && !selectedCategory ? ( + }>{t('tenantAppPermissionManagement.actions.createPermission')} + ) : undefined} + size="md" + striped + defaultSort={{ key: 'created_at', order: 'desc' }} + pagination={paginationBar} + /> + {/* */} {showCreateModal && ( @@ -352,8 +356,13 @@ export default function AppPermissionManagement() { // const PermissionModal: React.FC<{ title: string - formData: any - setFormData: (data: any) => void + formData: { + code: string + name: string + description: string + category: string + } + setFormData: (data: { code: string; name: string; description: string; category: string }) => void submitting: boolean onSubmit: () => void onClose: () => void diff --git a/basaltpass-frontend/src/features/tenant/app/AppRoleManagement.tsx b/basaltpass-frontend/src/features/tenant/app/AppRoleManagement.tsx index 5111d092..a274e328 100644 --- a/basaltpass-frontend/src/features/tenant/app/AppRoleManagement.tsx +++ b/basaltpass-frontend/src/features/tenant/app/AppRoleManagement.tsx @@ -1,12 +1,11 @@ -import React, { useCallback, useContext, useEffect, useMemo, useRef, useState } from 'react' -import { uiAlert, uiConfirm, uiPrompt } from '@contexts/DialogContext' +import React, { useEffect, useMemo, useState } from 'react' +import { uiAlert, uiConfirm } from '@contexts/DialogContext' import { useParams, useNavigate } from 'react-router-dom' import { ShieldCheckIcon, PlusIcon, PencilIcon, TrashIcon, - MagnifyingGlassIcon, XMarkIcon, KeyIcon, ExclamationTriangleIcon @@ -15,13 +14,17 @@ import TenantLayout from '@features/tenant/components/TenantLayout' import { tenantAppApi } from '@api/tenant/tenantApp' import { userPermissionsApi, type Permission, type Role } from '@api/tenant/appPermissions' import useDebounce from '@hooks/useDebounce' -import { PSkeleton, PPageHeader, PEmptyState, PButton, PInput, PTextarea, PBadge } from '@ui' +import useClientPagination from '@hooks/useClientPagination' +import useManagedPaginationBar from '@hooks/useManagedPaginationBar' +import { PSkeleton, PPageHeader, PEmptyState, PButton, PInput, PTextarea, PBadge, PManagementFilterCard, PManagedTableSection, PManagementPageContainer } from '@ui' +import { type PTableAction, type PTableColumn } from '@ui/PTable' import { useI18n } from '@shared/i18n' export default function AppRoleManagement() { const { t, locale } = useI18n() const { id: appId } = useParams<{ id: string }>() const navigate = useNavigate() + const pageSize = 10 const [app, setApp] = useState(null) const [roles, setRoles] = useState([]) @@ -153,12 +156,80 @@ export default function AppRoleManagement() { } } - const filteredRoles = roles.filter(role => - debouncedSearchTerm === '' || - role.name.toLowerCase().includes(debouncedSearchTerm.toLowerCase()) || - role.code.toLowerCase().includes(debouncedSearchTerm.toLowerCase()) || - (role.description && role.description.toLowerCase().includes(debouncedSearchTerm.toLowerCase())) - ) + const filteredRoles = useMemo(() => { + return roles.filter(role => + debouncedSearchTerm === '' || + role.name.toLowerCase().includes(debouncedSearchTerm.toLowerCase()) || + role.code.toLowerCase().includes(debouncedSearchTerm.toLowerCase()) || + (role.description && role.description.toLowerCase().includes(debouncedSearchTerm.toLowerCase())) + ) + }, [roles, debouncedSearchTerm]) + const { + currentPage, + pageItems: paginatedRoles, + totalItems, + setPage, + resetPage, + } = useClientPagination(filteredRoles, pageSize) + + const paginationBar = useManagedPaginationBar({ + currentPage, + pageSize, + totalItems, + onPageChange: setPage, + summary: ({ start, end, total }) => t('tenantRoleManagement.pagination.summary', { start, end, total }), + pageInfo: ({ currentPage: page, totalPages }) => t('tenantRoleManagement.pagination.pageInfo', { current: page, total: totalPages }), + }) + + useEffect(() => { + resetPage() + }, [debouncedSearchTerm, resetPage]) + + const columns: PTableColumn[] = [ + { + key: 'role_info', + title: t('tenantAppRoleManagement.table.roleInfo'), + render: (role) => ( +
+
{role.name}
+
{role.code}
+ {role.description ?
{role.description}
: null} +
+ ) + }, + { + key: 'permission_count', + title: t('tenantAppRoleManagement.table.permissionCount'), + align: 'center', + render: (role) => ( + {t('tenantAppRoleManagement.fields.permissionCount', { count: role.permissions?.length || 0 })} + ) + }, + { + key: 'created_at', + title: t('tenantAppRoleManagement.table.createdAt'), + sortable: true, + sorter: (a, b) => new Date(a.created_at).getTime() - new Date(b.created_at).getTime(), + render: (role) => new Date(role.created_at).toLocaleDateString(locale) + } + ] + + const actions: PTableAction[] = [ + { + key: 'edit', + label: t('tenantAppRoleManagement.actions.edit'), + icon: , + variant: 'secondary', + onClick: (role) => handleEditRole(role) + }, + { + key: 'delete', + label: t('tenantAppRoleManagement.actions.delete'), + icon: , + variant: 'danger', + onClick: (role) => handleDeleteRole(role) + } + ] const getPermissionsByCategory = () => { const categories: { [key: string]: Permission[] } = {} @@ -199,117 +270,56 @@ export default function AppRoleManagement() { return ( -
- {/* */} - } - actions={ -
- navigate(`/tenant/apps/${appId}/permissions`)} leftIcon={}>{t('tenantAppRoleManagement.actions.permissionManagement')} - navigate(`/tenant/apps/${appId}/users`)}>{t('tenantAppRoleManagement.actions.userManagement')} - }>{t('tenantAppRoleManagement.actions.createRole')} -
- } - /> - - {/* */} -
-
- - setSearchTerm(e.target.value)} - className="block w-full pl-10 pr-3 py-2 border border-gray-300 rounded-md leading-5 bg-white placeholder-gray-500 focus:outline-none focus:placeholder-gray-400 focus:ring-1 focus:ring-blue-500 focus:border-blue-500" - /> -
-
- - {/* */} -
-
-
-

- {t('tenantAppRoleManagement.list.title', { count: filteredRoles.length })} -

-
- - {filteredRoles.length === 0 ? ( + } + actions={ +
+ navigate(`/tenant/apps/${appId}/permissions`)} leftIcon={}>{t('tenantAppRoleManagement.actions.permissionManagement')} + navigate(`/tenant/apps/${appId}/users`)}>{t('tenantAppRoleManagement.actions.userManagement')} + }>{t('tenantAppRoleManagement.actions.createRole')} +
+ } + /> + } + filter={ + + } + > +
+ {filteredRoles.length === 0 && !loading ? ( +
- {!searchTerm && ( - }>{t('tenantAppRoleManagement.actions.createRole')} - )} + {!searchTerm ? }>{t('tenantAppRoleManagement.actions.createRole')} : null} - ) : ( -
- {filteredRoles.map((role) => ( -
-
-
- -
-

{role.name}

-

{t('tenantAppRoleManagement.fields.code')}: {role.code}

-
-
-
- handleEditRole(role)} - variant="ghost" - size="sm" - className="px-2 text-blue-600 hover:text-blue-800" - > - - - handleDeleteRole(role)} - variant="ghost" - size="sm" - className="px-2 text-red-600 hover:text-red-800" - > - - -
-
- - {role.description && ( -

{role.description}

- )} - -
-
- - {t('tenantAppRoleManagement.fields.permissionCount', { count: role.permissions?.length || 0 })} -
-
- {new Date(role.created_at).toLocaleDateString(locale)} -
-
- - {role.permissions && role.permissions.length > 0 && ( -
- {role.permissions.slice(0, 3).map((permission) => ( - {permission.name} - ))} - {role.permissions.length > 3 && ( - {t('tenantAppRoleManagement.fields.moreCount', { count: role.permissions.length - 3 })} - )} -
- )} -
- ))} -
- )} -
+
+ ) : ( + <> + + data={paginatedRoles} + columns={columns} + actions={actions} + rowKey={(row) => String(row.id)} + loading={loading} + emptyText={t('tenantAppRoleManagement.empty.noRole')} + defaultSort={{ key: 'created_at', order: 'desc' }} + pagination={paginationBar} + /> + + )}
-
+ {/* */} {showCreateModal && ( diff --git a/basaltpass-frontend/src/features/tenant/app/AppStats.tsx b/basaltpass-frontend/src/features/tenant/app/AppStats.tsx index b2ca4d8c..89c62e8b 100644 --- a/basaltpass-frontend/src/features/tenant/app/AppStats.tsx +++ b/basaltpass-frontend/src/features/tenant/app/AppStats.tsx @@ -181,7 +181,7 @@ export default function AppStats() {
{/* */}
-