From dbba1f4905deb02c8892dfdef994ae5f6862e345 Mon Sep 17 00:00:00 2001 From: Wang Defa <61809431+wangdefaa@users.noreply.github.com> Date: Fri, 10 Jul 2026 17:38:34 +0800 Subject: [PATCH] =?UTF-8?q?=E5=8F=91=E5=B8=83=200.1.0:=E9=80=9A=E7=9F=A5?= =?UTF-8?q?=E6=B8=A0=E9=81=93=E3=80=81=E5=91=8A=E8=AD=A6=E8=A7=84=E5=88=99?= =?UTF-8?q?=E3=80=81=E4=BB=A4=E7=89=8C=E7=89=88=E6=9C=AC=E4=B8=8E=E5=AE=89?= =?UTF-8?q?=E5=85=A8=E5=8A=A0=E5=9B=BA?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .gitea/workflows/release.yml | 16 +- .github/workflows/release.yml | 13 +- .vscode/launch.json | 3 +- CHANGELOG.md | 28 + DASH_VERSION | 1 + README.md | 85 ++- cmd/server/main.go | 41 +- docker-compose.yml | 8 +- docs/docs.go | 483 +++++++++++++- docs/swagger.json | 483 +++++++++++++- docs/swagger.yaml | 304 ++++++++- internal/api/auth.go | 9 +- internal/api/authx.go | 48 +- internal/api/federation.go | 24 +- internal/api/logevent.go | 122 ++++ internal/api/middleware.go | 16 +- internal/api/ociconfig.go | 35 +- internal/api/ociconfig_test.go | 54 ++ internal/api/region.go | 2 +- internal/api/router.go | 3 +- internal/api/router_test.go | 102 +++ internal/api/routes_ai.go | 2 +- internal/api/routes_auth.go | 7 +- internal/api/routes_oci.go | 1 + internal/api/routes_settings.go | 3 + internal/api/routes_tasks.go | 4 + internal/api/settings.go | 76 +++ internal/api/tenant.go | 75 ++- internal/api/webconsole.go | 2 +- internal/database/database.go | 19 +- internal/model/models.go | 43 +- internal/oci/account.go | 90 +++ internal/oci/account_test.go | 50 ++ internal/oci/audit.go | 114 +++- internal/oci/audit_test.go | 45 ++ internal/oci/cached.go | 17 +- internal/oci/cached_test.go | 15 + internal/oci/client.go | 53 +- internal/oci/domainsettings.go | 113 +++- internal/oci/federation.go | 45 +- internal/oci/iddomain_test.go | 107 +++ .../oci/real_federation_integration_test.go | 26 +- internal/oci/real_tenant_integration_test.go | 24 +- internal/oci/signon.go | 12 +- internal/oci/tenantuser.go | 175 ++++- internal/oci/tenantuser_test.go | 2 +- internal/service/aigateway.go | 38 +- internal/service/aigateway_test.go | 1 + internal/service/alertrule.go | 305 +++++++++ internal/service/alertrule_test.go | 345 ++++++++++ internal/service/audit.go | 115 ++-- internal/service/audit_test.go | 189 +++--- internal/service/auth.go | 70 +- internal/service/auth_test.go | 73 ++- internal/service/credentials.go | 50 +- internal/service/credentials_test.go | 8 +- internal/service/federation.go | 46 +- internal/service/logevent.go | 173 +++-- internal/service/logevent_test.go | 152 ++++- internal/service/loginguard.go | 96 +++ internal/service/loginguard_test.go | 105 +++ internal/service/notify.go | 152 ++++- internal/service/notifychannel.go | 218 +++++++ internal/service/notifychannel_test.go | 339 ++++++++++ internal/service/notifytpl.go | 26 +- internal/service/oauth.go | 12 +- internal/service/ociconfig.go | 69 +- internal/service/ociconfig_test.go | 55 ++ internal/service/overview.go | 2 +- internal/service/scopecache.go | 6 + internal/service/settingnotify.go | 309 +++++++++ internal/service/task.go | 252 +++++++- internal/service/task_test.go | 104 ++- internal/service/tenantadmin.go | 65 +- internal/service/tenantdelete.go | 457 +++++++++++++ internal/service/tenantdelete_test.go | 609 ++++++++++++++++++ internal/service/totp.go | 6 +- internal/service/totp_oauth_test.go | 2 +- 78 files changed, 6898 insertions(+), 551 deletions(-) create mode 100644 DASH_VERSION create mode 100644 internal/api/ociconfig_test.go create mode 100644 internal/oci/iddomain_test.go create mode 100644 internal/service/alertrule.go create mode 100644 internal/service/alertrule_test.go create mode 100644 internal/service/loginguard_test.go create mode 100644 internal/service/notifychannel.go create mode 100644 internal/service/notifychannel_test.go create mode 100644 internal/service/settingnotify.go create mode 100644 internal/service/tenantdelete.go create mode 100644 internal/service/tenantdelete_test.go diff --git a/.gitea/workflows/release.yml b/.gitea/workflows/release.yml index 92f1e71..7c69282 100644 --- a/.gitea/workflows/release.yml +++ b/.gitea/workflows/release.yml @@ -23,14 +23,15 @@ jobs: echo "IMAGE=${REGISTRY}/${{ gitea.repository }}" >> "$GITHUB_ENV" echo "BUILD_TIME=$(date -u +%Y-%m-%dT%H:%M:%SZ)" >> "$GITHUB_ENV" - - name: 下载前端最新 dist.zip 解压进嵌入目录 + - name: 下载 DASH_VERSION 固定版本的前端 dist.zip,校验后解压进嵌入目录 env: TOKEN: ${{ secrets.BUILD_TOKEN }} run: | - DASH_TAG=$(curl -fsSL -H "Authorization: token ${TOKEN}" \ - "${{ gitea.server_url }}/api/v1/repos/${DASH_REPO}/releases/latest" | jq -r '.tag_name') - curl -fsSL -H "Authorization: token ${TOKEN}" \ - -o dist.zip "${{ gitea.server_url }}/${DASH_REPO}/releases/download/${DASH_TAG}/dist.zip" + DASH_TAG=$(head -1 DASH_VERSION) + BASE="${{ gitea.server_url }}/${DASH_REPO}/releases/download/${DASH_TAG}" + curl -fsSL -H "Authorization: token ${TOKEN}" -o dist.zip "${BASE}/dist.zip" + curl -fsSL -H "Authorization: token ${TOKEN}" -o dist.zip.sha256 "${BASE}/dist.zip.sha256" + sha256sum -c dist.zip.sha256 rm -rf internal/webui/dist && mkdir -p internal/webui/dist unzip -q dist.zip -d internal/webui/dist test -f internal/webui/dist/index.html @@ -46,8 +47,9 @@ jobs: mkdir -p assets mv output/linux_amd64/oci-portal-server assets/oci-portal-server-linux-amd64 mv output/linux_arm64/oci-portal-server assets/oci-portal-server-linux-arm64 + (cd assets && sha256sum * > SHA256SUMS) - - name: 构建并推送多平台镜像(与上一步共享构建缓存) + - name: 构建并推送多平台镜像 env: TOKEN: ${{ secrets.BUILD_TOKEN }} run: | @@ -65,7 +67,7 @@ jobs: awk -v ver="${TAG#v}" '$0 ~ "^## \\[" ver "\\]" {flag=1; next} /^## \[/ && flag {exit} flag {print}' CHANGELOG.md > release-notes.md [ -s release-notes.md ] || echo "Release ${TAG}" > release-notes.md - - name: 创建 Release 并上传二进制附件(重跑幂等,同名附件自动替换) + - name: 创建 Release 并上传二进制附件 uses: https://gitea.com/actions/gitea-release-action@v1 with: token: ${{ secrets.BUILD_TOKEN }} diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index a7029df..4cd7e8a 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -10,7 +10,7 @@ permissions: env: BUILDX_NO_DEFAULT_ATTESTATIONS: "1" - DASH_DOWNLOAD: https://github.com/wangdefaa/oci-portal-dash/releases/latest/download/dist.zip + DASH_REPO: wangdefaa/oci-portal-dash IMAGE: ghcr.io/wangdefaa/oci-portal jobs: @@ -21,14 +21,18 @@ jobs: - uses: docker/setup-buildx-action@v3 - - name: 下载前端最新 dist.zip 解压进嵌入目录 + - name: 下载前端 dist.zip run: | - curl -fsSL -o dist.zip "$DASH_DOWNLOAD" + DASH_TAG=$(head -1 DASH_VERSION) + BASE="https://github.com/${DASH_REPO}/releases/download/${DASH_TAG}" + curl -fsSL -o dist.zip "${BASE}/dist.zip" + curl -fsSL -o dist.zip.sha256 "${BASE}/dist.zip.sha256" + sha256sum -c dist.zip.sha256 rm -rf internal/webui/dist && mkdir -p internal/webui/dist unzip -q dist.zip -d internal/webui/dist test -f internal/webui/dist/index.html - - name: 构建双架构二进制(artifact stage 导出) + - name: 构建双架构二进制 run: | BUILD_TIME=$(date -u +%Y-%m-%dT%H:%M:%SZ) echo "BUILD_TIME=${BUILD_TIME}" >> "$GITHUB_ENV" @@ -41,6 +45,7 @@ jobs: mkdir -p assets mv output/linux_amd64/oci-portal-server assets/oci-portal-server-linux-amd64 mv output/linux_arm64/oci-portal-server assets/oci-portal-server-linux-arm64 + (cd assets && sha256sum * > SHA256SUMS) ls -lh assets/ - name: 提取 CHANGELOG 版本段作为 Release 描述 diff --git a/.vscode/launch.json b/.vscode/launch.json index 648cd50..ea4b06c 100644 --- a/.vscode/launch.json +++ b/.vscode/launch.json @@ -11,7 +11,8 @@ "env": { "DATA_KEY": "dev-data-key", "JWT_SECRET": "dev-jwt-secret", - "ADMIN_PASSWORD": "admin123" + "ADMIN_PASSWORD": "admin123", + "GIN_MODE": "debug" } }, { diff --git a/CHANGELOG.md b/CHANGELOG.md index b497338..67fd6e8 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -2,6 +2,34 @@ 格式遵循 [Keep a Changelog](https://keepachangelog.com/zh-CN/1.1.0/),版本号遵循语义化版本。 +## [0.1.0] - 2026-07-10 + +### Added + +- 通知渠道扩展:Webhook(Slack / 飞书 / 钉钉 / 企微预设)、ntfy、Bark、SMTP 邮件与 Telegram 五渠道并存,逐渠道推送互不影响,密钥加密落库 +- 审计告警规则:按租户 / 事件 / 资源 / 来源 IP(CIDR in|notin)组合匹配,阈值窗口计数与冷却,命中即推送模板通知 +- 撤销全部会话:`POST /auth/revoke-sessions` 使全部旧令牌立即失效,响应携带操作者的新令牌 +- 租户暂停状态识别:测活经账户 capabilities 判定 `suspended`,独立于存活 / 失效展示 +- 抢机成功通知补全实例明细:租户别名、区域别名、类型与配置、镜像、公网 IPv4 与 root 密码 + +### Changed + +- 敏感身份操作(修改凭据、TOTP 开关、外部身份绑定 / 解绑、密码登录开关)成功后令牌版本递增:全部旧 JWT 立即失效,响应返回新令牌无感续期(原 204 改为 200 + token) +- 管理员徽标与管理员组识别支持双候选组名(Administrators / OCI_Administrators) +- Release 流程:前端产物版本由 `DASH_VERSION` 显式固定并做 sha256 校验,二进制附 `SHA256SUMS`;docker-compose 示例默认只绑定 `127.0.0.1` +- README 新增「反向代理」小节(Caddy / nginx / Traefik 示例与路由前缀清单) + +### Fixed + +- 删除租户改为单事务清理全部关联本地数据(任务及日志、测活 / 成本快照、回传事件、告警规则、AI 渠道、Webhook 密钥),等待在途任务执行结束,迟到写入不再残留孤儿数据 + +### Security + +- 实例 root 密码收敛:实例详情默认掩码、点击查看,串口控制台不再提供「复制 root 密码」 +- JWT 令牌版本机制:凭据 / TOTP / 外部身份 / 登录策略变更即失效全部旧令牌 +- HTTP 入口加固:显式读写超时与优雅关闭、请求体上限(面板 1MB / AI 网关 10MB)、登录字段长度上限、登录守卫条目数有界 +- 出错响应脱敏:内部错误只返回固定文案 + requestId,完整原因写服务端日志;GORM SQL 日志参数化输出,不再记录实参 + ## [0.0.1] - 2026-07-09 首个版本:自托管 OCI 多租户管理面板后端。 diff --git a/DASH_VERSION b/DASH_VERSION new file mode 100644 index 0000000..b82608c --- /dev/null +++ b/DASH_VERSION @@ -0,0 +1 @@ +v0.1.0 diff --git a/README.md b/README.md index f450613..5fb9cb3 100644 --- a/README.md +++ b/README.md @@ -35,10 +35,10 @@ - **抢机任务**:cron 周期尝试创建实例直到成功,支持熔断与通知 - **网页控制台**:实例串行控制台(xterm)与 VNC(noVNC),两跳 SSH 隧道 - **AI 网关**:OpenAI / Anthropic 兼容端点转发 OCI GenAI(对话/Responses/Embeddings),号池渠道加权负载均衡、熔断探测、密钥管理与用量日志 -- **日志回传**:OCI Audit 事件经 Connector Hub → Notifications HTTPS 订阅回传入库,一键创建链路 -- **租户治理**:IAM 用户/MFA/API Key 管理、密码策略、身份提供商(SAML)、通知收件人 -- **安全**:JWT + bcrypt、TOTP 两步验证、OIDC/GitHub 外部登录、登录锁定、IP 限速、真实 IP 头可配、系统操作审计 -- **通知**:Telegram 模板化推送(测活/抢机/成本/锁定等事件) +- **日志回传**:OCI Audit 事件经 Connector Hub → Notifications HTTPS 订阅回传入库,一键创建链路;自定义告警规则(事件类型/来源 IP 白名单/资源/频率阈值)命中即推送 +- **租户治理**:IAM 用户/MFA/API Key 管理、密码策略、身份提供商(SAML)、通知收件人;多 Identity Domain 租户可按域切换管理 +- **安全**:JWT + bcrypt(凭据变更旧令牌立即失效,可一键撤销全部会话)、TOTP 两步验证、OIDC/GitHub 外部登录、登录锁定、IP 限速、真实 IP 头可配、请求体/超时防护、系统操作审计 +- **通知**:模板化推送,五渠道并存(Telegram / Webhook / ntfy / Bark / SMTP),Webhook 可对接飞书、钉钉、Slack、企业微信机器人 ## 快速开始 @@ -55,9 +55,13 @@ DATA_KEY=$(openssl rand -hex 32) JWT_SECRET=$(openssl rand -hex 32) ADMIN_PASSWO ### Docker Compose ```bash +# 镜像以 distroless nonroot(uid 65532)运行,数据卷需可写,否则 SQLite 报 unable to open database file +mkdir -p ./data && sudo chown 65532:65532 ./data DATA_KEY=$(openssl rand -hex 32) JWT_SECRET=$(openssl rand -hex 32) ADMIN_PASSWORD=<初始密码> docker compose up -d ``` +默认只监听 `127.0.0.1:18888`;公网访问须置于 TLS 反向代理之后,见下文「反向代理」。 + ### 源码构建(单文件,含前端) ```bash @@ -72,6 +76,78 @@ CGO_ENABLED=0 go build -trimpath \ > 注意:上述解压会覆盖仓库占位文件 `internal/webui/dist/index.html`,提交代码前勿把真实产物带入版本库。 +### 反向代理(公网部署必读) + +面板自身只提供 HTTP,管理员口令、JWT 与租户 API 私钥都会经明文承载。除本机试用外,应让面板仅监听回环地址(compose 示例已默认 `127.0.0.1:18888`),由支持 TLS 的反向代理对外提供 HTTPS。 + +Caddy 最小示例(整站反代,自动申请并续期 Let's Encrypt 证书,WebSocket 自动透传): + +```caddyfile +portal.example.com { + reverse_proxy 127.0.0.1:18888 +} +``` + +nginx 示例(证书自备;Web Console 的 WebSocket 升级与长超时必须显式配置,否则串行终端连不上或空闲即断): + +```nginx +# http 块内:按请求是否升级为 WebSocket 决定 Connection 头 +map $http_upgrade $connection_upgrade { + default upgrade; + "" close; +} + +server { + listen 443 ssl; + server_name portal.example.com; + + ssl_certificate /etc/nginx/certs/portal.example.com.crt; + ssl_certificate_key /etc/nginx/certs/portal.example.com.key; + + # 放行到最大入口(AI 网关 10MB);各入口更细的上限由后端自身执行 + client_max_body_size 10m; + + location / { + proxy_pass http://127.0.0.1:18888; + proxy_http_version 1.1; + proxy_set_header Upgrade $http_upgrade; + proxy_set_header Connection $connection_upgrade; + proxy_set_header Host $host; + proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for; + # WebSocket 空闲与 AI 流式响应都需要长读超时(默认 60s 会断) + proxy_read_timeout 1h; + proxy_send_timeout 1h; + } +} +``` + +Traefik 示例(已运行 Traefik 的 Docker 环境,面板容器加 label 接入即可,WebSocket 透明支持): + +```yaml +services: + oci-portal: + image: ghcr.io/wangdefaa/oci-portal:latest + # ……环境变量与数据卷同上文 compose 示例;与 Traefik 同一 docker 网络, + # Traefik 经容器网络直达 8080,无需(也不应)对外映射 ports + labels: + - traefik.enable=true + - traefik.http.routers.oci-portal.rule=Host(`portal.example.com`) + - traefik.http.routers.oci-portal.entrypoints=websecure + - traefik.http.routers.oci-portal.tls.certresolver=le # 换成你的 certResolver 名 + - traefik.http.services.oci-portal.loadbalancer.server.port=8080 +``` + +若前端静态文件由反代直接伺服(不走内嵌页面),则按前缀代理,缺一不可: + +| 前缀 | 内容 | 何时需要 | +| --- | --- | --- | +| `/api/*` | 面板 REST、Web Console WebSocket、日志回传 webhook | 始终 | +| `/ai/*` | AI 网关(OpenAI / Claude 兼容端点,独立密钥鉴权) | 启用 AI 网关时 | +| 其余路径 | 前端 SPA 静态文件(404 回退 `index.html`) | 静态分离形态 | + +- 反代默认追加的 `X-Forwarded-For` 用于还原真实客户端 IP,系统日志留痕、登录锁定与 IP 限速都依赖它 +- 请求体上限建议与后端一致:`/api/*` 1MB、`/ai/*` 10MB(Caddy 用 `request_body` 按前缀分层) + ## 环境变量 | 变量 | 必填 | 默认值 | 说明 | @@ -88,6 +164,7 @@ CGO_ENABLED=0 go build -trimpath \ | `HTTPS_PROXY` | 否 | 无 | 出站代理(如 Telegram 通知走代理) | | `TZ` | 否 | 系统 | cron 表达式解释时区(容器内建议显式设置,二进制已嵌 tzdata) | | `SWAGGER` | 否 | 关 | `1` 时开放 `/swagger/index.html` API 文档(生产建议按需临时开启) | +| `GIN_MODE` | 否 | `release` | `debug` / `release`,影响日志与调试输出 | ## 开发 diff --git a/cmd/server/main.go b/cmd/server/main.go index e27a059..a27b92f 100644 --- a/cmd/server/main.go +++ b/cmd/server/main.go @@ -9,12 +9,15 @@ import ( "net" "net/http" "os" + "os/signal" + "syscall" "time" // 纯 Go 时区数据(+~450KB):distroless/scratch 容器无 tzdata, // 嵌入后 TZ 环境变量即可让 cron 按本地时区解释 _ "time/tzdata" + "github.com/gin-gonic/gin" "gorm.io/gorm" // swag 生成的 OpenAPI spec(init 时注册,SWAGGER=1 的 /swagger 路由消费) @@ -35,6 +38,12 @@ import ( // @in header // @name Authorization // @description 登录接口返回的 JWT,格式: Bearer +func init() { + if os.Getenv("GIN_MODE") == "" { + gin.SetMode(gin.ReleaseMode) + } +} + func main() { healthcheck := flag.Bool("healthcheck", false, "探活本机服务后退出:健康 0,异常 1(容器 HEALTHCHECK 用)") flag.Parse() @@ -126,6 +135,7 @@ func run() error { defer aiGateway.Wait() defer stopCleanup() tasks := service.NewTaskService(db, ociConfigs, notifier, settings) + ociConfigs.SetTenantCleanupDeps(tasks, logEvents) tasks.AttachAiGateway(aiGateway) aiGateway.SetOnChannelsChanged(tasks.SyncAiProbeTask) if err := tasks.Start(); err != nil { @@ -137,5 +147,34 @@ func run() error { console := service.NewConsoleService(ociConfigs) proxies := service.NewProxyService(db, cipher) defer proxies.Wait() - return api.NewRouter(auth, oauth, ociConfigs, tasks, console, settings, notifier, systemLogs, logEvents, proxies, aiGateway).Run(cfg.Addr) + router := api.NewRouter(auth, oauth, ociConfigs, tasks, console, settings, notifier, systemLogs, logEvents, proxies, aiGateway) + return serveHTTP(cfg.Addr, router) +} + +// serveHTTP 以显式超时与请求头上限启动 HTTP 服务,并在 SIGINT/SIGTERM 时 +// 优雅关闭(10 秒排空);防慢连接(Slowloris)与超大头耗尽资源。 +// WriteTimeout 置 0:AI 网关流式响应无固定上界,WebSocket 为 hijack 连接 +// 本就不受 net/http 超时管控;读侧超时已覆盖慢请求风险。 +func serveHTTP(addr string, handler http.Handler) error { + srv := &http.Server{ + Addr: addr, + Handler: handler, + ReadHeaderTimeout: 5 * time.Second, + ReadTimeout: 60 * time.Second, + IdleTimeout: 2 * time.Minute, + MaxHeaderBytes: 64 << 10, + } + ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM) + defer stop() + errCh := make(chan error, 1) + go func() { errCh <- srv.ListenAndServe() }() + select { + case err := <-errCh: + return err + case <-ctx.Done(): + log.Println("[info] shutting down http server") + shutdownCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + return srv.Shutdown(shutdownCtx) + } } diff --git a/docker-compose.yml b/docker-compose.yml index 687351e..6fa4d28 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -1,9 +1,10 @@ services: oci-portal: image: ghcr.io/wangdefaa/oci-portal:latest + container_name: oci-portal restart: unless-stopped ports: - - "18888:8080" + - "127.0.0.1:18888:8080" environment: DB_PATH: /data/oci-portal.db DATA_KEY: ${DATA_KEY:?DATA_KEY is required} @@ -11,7 +12,4 @@ services: ADMIN_PASSWORD: ${ADMIN_PASSWORD:-} TZ: Asia/Shanghai volumes: - - oci-portal-data:/data - -volumes: - oci-portal-data: + - ./data:/data diff --git a/docs/docs.go b/docs/docs.go index 85322d3..0596abf 100644 --- a/docs/docs.go +++ b/docs/docs.go @@ -871,6 +871,30 @@ const docTemplate = `{ } } }, + "/api/v1/auth/revoke-sessions": { + "post": { + "security": [ + { + "BearerAuth": [] + } + ], + "tags": [ + "认证" + ], + "summary": "撤销全部会话", + "responses": { + "200": { + "description": "新 token 与 expiresAt", + "schema": { + "type": "object", + "additionalProperties": { + "type": "string" + } + } + } + } + } + }, "/api/v1/auth/totp": { "get": { "security": [ @@ -1078,6 +1102,124 @@ const docTemplate = `{ } } }, + "/api/v1/log-events/alert-rules": { + "get": { + "security": [ + { + "BearerAuth": [] + } + ], + "tags": [ + "任务与日志回传" + ], + "summary": "告警规则列表", + "responses": { + "200": { + "description": "OK", + "schema": { + "type": "object", + "additionalProperties": true + } + } + } + }, + "post": { + "security": [ + { + "BearerAuth": [] + } + ], + "tags": [ + "任务与日志回传" + ], + "summary": "创建告警规则", + "parameters": [ + { + "description": "请求体", + "name": "body", + "in": "body", + "required": true, + "schema": { + "$ref": "#/definitions/internal_api.alertRuleRequest" + } + } + ], + "responses": { + "201": { + "description": "Created", + "schema": { + "type": "object", + "additionalProperties": true + } + } + } + } + }, + "/api/v1/log-events/alert-rules/{ruleId}": { + "put": { + "security": [ + { + "BearerAuth": [] + } + ], + "tags": [ + "任务与日志回传" + ], + "summary": "更新告警规则", + "parameters": [ + { + "type": "integer", + "description": "规则 ID", + "name": "ruleId", + "in": "path", + "required": true + }, + { + "description": "请求体", + "name": "body", + "in": "body", + "required": true, + "schema": { + "$ref": "#/definitions/internal_api.alertRuleRequest" + } + } + ], + "responses": { + "200": { + "description": "OK", + "schema": { + "type": "object", + "additionalProperties": true + } + } + } + }, + "delete": { + "security": [ + { + "BearerAuth": [] + } + ], + "tags": [ + "任务与日志回传" + ], + "summary": "删除告警规则", + "parameters": [ + { + "type": "integer", + "description": "规则 ID", + "name": "ruleId", + "in": "path", + "required": true + } + ], + "responses": { + "204": { + "description": "无内容" + } + } + } + }, "/api/v1/oci-configs": { "get": { "security": [ @@ -1205,6 +1347,7 @@ const docTemplate = `{ "BearerAuth": [] } ], + "description": "删除租户配置及其本地关联任务、快照、回传、告警和 AI 渠道数据,并撤销 Webhook 密钥;不删除 OCI 云端资源。", "tags": [ "租户配置" ], @@ -1235,7 +1378,7 @@ const docTemplate = `{ "tags": [ "租户 IAM" ], - "summary": "实时查询租户 OCI 审计事件:hours 首查", + "summary": "批式懒加载查询租户 OCI 审计事件", "parameters": [ { "type": "integer", @@ -1243,6 +1386,18 @@ const docTemplate = `{ "name": "id", "in": "path", "required": true + }, + { + "type": "string", + "description": "续查游标(上次响应原样带回)", + "name": "cursor", + "in": "query" + }, + { + "type": "integer", + "description": "单批目标条数,缺省 100,上限 200", + "name": "limit", + "in": "query" } ], "responses": { @@ -1656,6 +1811,37 @@ const docTemplate = `{ } } }, + "/api/v1/oci-configs/{id}/domains": { + "get": { + "security": [ + { + "BearerAuth": [] + } + ], + "tags": [ + "租户 IAM" + ], + "summary": "租户身份域列表", + "parameters": [ + { + "type": "integer", + "description": "配置 ID", + "name": "id", + "in": "path", + "required": true + } + ], + "responses": { + "200": { + "description": "OK", + "schema": { + "type": "object", + "additionalProperties": true + } + } + } + } + }, "/api/v1/oci-configs/{id}/identity-providers": { "get": { "security": [ @@ -1674,6 +1860,12 @@ const docTemplate = `{ "name": "id", "in": "path", "required": true + }, + { + "type": "string", + "description": "身份域 OCID(缺省 Default 域)", + "name": "domainId", + "in": "query" } ], "responses": { @@ -1704,6 +1896,12 @@ const docTemplate = `{ "in": "path", "required": true }, + { + "type": "string", + "description": "身份域 OCID(缺省 Default 域)", + "name": "domainId", + "in": "query" + }, { "description": "请求体", "name": "body", @@ -1744,6 +1942,12 @@ const docTemplate = `{ "in": "path", "required": true }, + { + "type": "string", + "description": "身份域 OCID(缺省 Default 域)", + "name": "domainId", + "in": "query" + }, { "type": "string", "description": "idpId", @@ -1778,6 +1982,12 @@ const docTemplate = `{ "in": "path", "required": true }, + { + "type": "string", + "description": "身份域 OCID(缺省 Default 域)", + "name": "domainId", + "in": "query" + }, { "type": "string", "description": "idpId", @@ -1824,6 +2034,12 @@ const docTemplate = `{ "name": "id", "in": "path", "required": true + }, + { + "type": "string", + "description": "身份域 OCID(缺省 Default 域)", + "name": "domainId", + "in": "query" } ], "responses": { @@ -1854,6 +2070,12 @@ const docTemplate = `{ "in": "path", "required": true }, + { + "type": "string", + "description": "身份域 OCID(缺省 Default 域)", + "name": "domainId", + "in": "query" + }, { "description": "请求体(见接口说明)", "name": "body", @@ -3015,6 +3237,12 @@ const docTemplate = `{ "name": "id", "in": "path", "required": true + }, + { + "type": "string", + "description": "身份域 OCID(缺省 Default 域)", + "name": "domainId", + "in": "query" } ], "responses": { @@ -3045,6 +3273,12 @@ const docTemplate = `{ "in": "path", "required": true }, + { + "type": "string", + "description": "身份域 OCID(缺省 Default 域)", + "name": "domainId", + "in": "query" + }, { "description": "请求体(见接口说明)", "name": "body", @@ -3084,6 +3318,12 @@ const docTemplate = `{ "name": "id", "in": "path", "required": true + }, + { + "type": "string", + "description": "身份域 OCID(缺省 Default 域)", + "name": "domainId", + "in": "query" } ], "responses": { @@ -3116,6 +3356,12 @@ const docTemplate = `{ "in": "path", "required": true }, + { + "type": "string", + "description": "身份域 OCID(缺省 Default 域)", + "name": "domainId", + "in": "query" + }, { "type": "string", "description": "policyId", @@ -3231,6 +3477,12 @@ const docTemplate = `{ "name": "id", "in": "path", "required": true + }, + { + "type": "string", + "description": "身份域 OCID(缺省 Default 域)", + "name": "domainId", + "in": "query" } ], "responses": { @@ -3478,6 +3730,12 @@ const docTemplate = `{ "in": "path", "required": true }, + { + "type": "string", + "description": "身份域 OCID(缺省 Default 域)", + "name": "domainId", + "in": "query" + }, { "description": "请求体(见接口说明)", "name": "body", @@ -3518,6 +3776,12 @@ const docTemplate = `{ "in": "path", "required": true }, + { + "type": "string", + "description": "身份域 OCID(缺省 Default 域)", + "name": "domainId", + "in": "query" + }, { "type": "string", "description": "ruleId", @@ -3551,6 +3815,12 @@ const docTemplate = `{ "name": "id", "in": "path", "required": true + }, + { + "type": "string", + "description": "身份域 OCID(缺省 Default 域)", + "name": "domainId", + "in": "query" } ], "responses": { @@ -3835,6 +4105,12 @@ const docTemplate = `{ "name": "id", "in": "path", "required": true + }, + { + "type": "string", + "description": "身份域 OCID(缺省 Default 域)", + "name": "domainId", + "in": "query" } ], "responses": { @@ -3865,6 +4141,12 @@ const docTemplate = `{ "in": "path", "required": true }, + { + "type": "string", + "description": "身份域 OCID(缺省 Default 域)", + "name": "domainId", + "in": "query" + }, { "description": "请求体", "name": "body", @@ -3905,6 +4187,12 @@ const docTemplate = `{ "in": "path", "required": true }, + { + "type": "string", + "description": "身份域 OCID(缺省 Default 域)", + "name": "domainId", + "in": "query" + }, { "type": "string", "description": "userId", @@ -3941,6 +4229,12 @@ const docTemplate = `{ "in": "path", "required": true }, + { + "type": "string", + "description": "身份域 OCID(缺省 Default 域)", + "name": "domainId", + "in": "query" + }, { "type": "string", "description": "userId", @@ -3986,6 +4280,12 @@ const docTemplate = `{ "in": "path", "required": true }, + { + "type": "string", + "description": "身份域 OCID(缺省 Default 域)", + "name": "domainId", + "in": "query" + }, { "type": "string", "description": "userId", @@ -4058,6 +4358,12 @@ const docTemplate = `{ "in": "path", "required": true }, + { + "type": "string", + "description": "身份域 OCID(缺省 Default 域)", + "name": "domainId", + "in": "query" + }, { "type": "string", "description": "userId", @@ -4096,6 +4402,12 @@ const docTemplate = `{ "in": "path", "required": true }, + { + "type": "string", + "description": "身份域 OCID(缺省 Default 域)", + "name": "domainId", + "in": "query" + }, { "type": "string", "description": "userId", @@ -4749,6 +5061,95 @@ const docTemplate = `{ } } }, + "/api/v1/settings/notify-channels": { + "get": { + "security": [ + { + "BearerAuth": [] + } + ], + "tags": [ + "设置" + ], + "summary": "返回全部通知渠道的脱敏视图", + "responses": { + "200": { + "description": "OK", + "schema": { + "type": "object", + "additionalProperties": true + } + } + } + } + }, + "/api/v1/settings/notify-channels/{type}": { + "put": { + "security": [ + { + "BearerAuth": [] + } + ], + "tags": [ + "设置" + ], + "summary": "保存单渠道配置并返回全渠道最新视图", + "parameters": [ + { + "type": "string", + "description": "渠道类型(webhook/ntfy/bark/smtp)", + "name": "type", + "in": "path", + "required": true + }, + { + "description": "请求体", + "name": "body", + "in": "body", + "required": true, + "schema": { + "$ref": "#/definitions/internal_api.updateNotifyChannelRequest" + } + } + ], + "responses": { + "200": { + "description": "OK", + "schema": { + "type": "object", + "additionalProperties": true + } + } + } + } + }, + "/api/v1/settings/notify-channels/{type}/test": { + "post": { + "security": [ + { + "BearerAuth": [] + } + ], + "tags": [ + "设置" + ], + "summary": "向指定渠道发送测试消息", + "parameters": [ + { + "type": "string", + "description": "渠道类型(webhook/ntfy/bark/smtp)", + "name": "type", + "in": "path", + "required": true + } + ], + "responses": { + "204": { + "description": "无内容" + } + } + } + }, "/api/v1/settings/notify-events": { "get": { "security": [ @@ -5403,6 +5804,38 @@ const docTemplate = `{ } }, "definitions": { + "internal_api.alertRuleRequest": { + "type": "object", + "properties": { + "enabled": { + "type": "boolean" + }, + "eventTypes": { + "type": "string" + }, + "name": { + "type": "string" + }, + "ociConfigId": { + "type": "integer" + }, + "resourceMatch": { + "type": "string" + }, + "sourceIpMode": { + "type": "string" + }, + "sourceIps": { + "type": "string" + }, + "threshold": { + "type": "integer" + }, + "windowMinutes": { + "type": "integer" + } + } + }, "internal_api.attachBootVolumeRequest": { "type": "object", "required": [ @@ -5817,14 +6250,18 @@ const docTemplate = `{ ], "properties": { "password": { - "type": "string" + "type": "string", + "maxLength": 128 }, "totpCode": { "description": "两步验证码;账号已启用 TOTP 时必填,缺失返回 428", - "type": "string" + "type": "string", + "maxLength": 8 }, "username": { - "type": "string" + "description": "字段长度上限防超长输入撑爆登录守卫与 bcrypt(S-03)", + "type": "string", + "maxLength": 64 } } }, @@ -5891,6 +6328,44 @@ const docTemplate = `{ } } }, + "internal_api.updateNotifyChannelRequest": { + "type": "object", + "properties": { + "bodyTemplate": { + "type": "string" + }, + "enabled": { + "type": "boolean" + }, + "from": { + "type": "string" + }, + "host": { + "type": "string" + }, + "port": { + "type": "integer" + }, + "secret": { + "type": "string" + }, + "server": { + "type": "string" + }, + "to": { + "type": "string" + }, + "topic": { + "type": "string" + }, + "url": { + "type": "string" + }, + "username": { + "type": "string" + } + } + }, "internal_api.updateSecurityListRequest": { "type": "object", "properties": { diff --git a/docs/swagger.json b/docs/swagger.json index 716c4b1..e1539c5 100644 --- a/docs/swagger.json +++ b/docs/swagger.json @@ -864,6 +864,30 @@ } } }, + "/api/v1/auth/revoke-sessions": { + "post": { + "security": [ + { + "BearerAuth": [] + } + ], + "tags": [ + "认证" + ], + "summary": "撤销全部会话", + "responses": { + "200": { + "description": "新 token 与 expiresAt", + "schema": { + "type": "object", + "additionalProperties": { + "type": "string" + } + } + } + } + } + }, "/api/v1/auth/totp": { "get": { "security": [ @@ -1071,6 +1095,124 @@ } } }, + "/api/v1/log-events/alert-rules": { + "get": { + "security": [ + { + "BearerAuth": [] + } + ], + "tags": [ + "任务与日志回传" + ], + "summary": "告警规则列表", + "responses": { + "200": { + "description": "OK", + "schema": { + "type": "object", + "additionalProperties": true + } + } + } + }, + "post": { + "security": [ + { + "BearerAuth": [] + } + ], + "tags": [ + "任务与日志回传" + ], + "summary": "创建告警规则", + "parameters": [ + { + "description": "请求体", + "name": "body", + "in": "body", + "required": true, + "schema": { + "$ref": "#/definitions/internal_api.alertRuleRequest" + } + } + ], + "responses": { + "201": { + "description": "Created", + "schema": { + "type": "object", + "additionalProperties": true + } + } + } + } + }, + "/api/v1/log-events/alert-rules/{ruleId}": { + "put": { + "security": [ + { + "BearerAuth": [] + } + ], + "tags": [ + "任务与日志回传" + ], + "summary": "更新告警规则", + "parameters": [ + { + "type": "integer", + "description": "规则 ID", + "name": "ruleId", + "in": "path", + "required": true + }, + { + "description": "请求体", + "name": "body", + "in": "body", + "required": true, + "schema": { + "$ref": "#/definitions/internal_api.alertRuleRequest" + } + } + ], + "responses": { + "200": { + "description": "OK", + "schema": { + "type": "object", + "additionalProperties": true + } + } + } + }, + "delete": { + "security": [ + { + "BearerAuth": [] + } + ], + "tags": [ + "任务与日志回传" + ], + "summary": "删除告警规则", + "parameters": [ + { + "type": "integer", + "description": "规则 ID", + "name": "ruleId", + "in": "path", + "required": true + } + ], + "responses": { + "204": { + "description": "无内容" + } + } + } + }, "/api/v1/oci-configs": { "get": { "security": [ @@ -1198,6 +1340,7 @@ "BearerAuth": [] } ], + "description": "删除租户配置及其本地关联任务、快照、回传、告警和 AI 渠道数据,并撤销 Webhook 密钥;不删除 OCI 云端资源。", "tags": [ "租户配置" ], @@ -1228,7 +1371,7 @@ "tags": [ "租户 IAM" ], - "summary": "实时查询租户 OCI 审计事件:hours 首查", + "summary": "批式懒加载查询租户 OCI 审计事件", "parameters": [ { "type": "integer", @@ -1236,6 +1379,18 @@ "name": "id", "in": "path", "required": true + }, + { + "type": "string", + "description": "续查游标(上次响应原样带回)", + "name": "cursor", + "in": "query" + }, + { + "type": "integer", + "description": "单批目标条数,缺省 100,上限 200", + "name": "limit", + "in": "query" } ], "responses": { @@ -1649,6 +1804,37 @@ } } }, + "/api/v1/oci-configs/{id}/domains": { + "get": { + "security": [ + { + "BearerAuth": [] + } + ], + "tags": [ + "租户 IAM" + ], + "summary": "租户身份域列表", + "parameters": [ + { + "type": "integer", + "description": "配置 ID", + "name": "id", + "in": "path", + "required": true + } + ], + "responses": { + "200": { + "description": "OK", + "schema": { + "type": "object", + "additionalProperties": true + } + } + } + } + }, "/api/v1/oci-configs/{id}/identity-providers": { "get": { "security": [ @@ -1667,6 +1853,12 @@ "name": "id", "in": "path", "required": true + }, + { + "type": "string", + "description": "身份域 OCID(缺省 Default 域)", + "name": "domainId", + "in": "query" } ], "responses": { @@ -1697,6 +1889,12 @@ "in": "path", "required": true }, + { + "type": "string", + "description": "身份域 OCID(缺省 Default 域)", + "name": "domainId", + "in": "query" + }, { "description": "请求体", "name": "body", @@ -1737,6 +1935,12 @@ "in": "path", "required": true }, + { + "type": "string", + "description": "身份域 OCID(缺省 Default 域)", + "name": "domainId", + "in": "query" + }, { "type": "string", "description": "idpId", @@ -1771,6 +1975,12 @@ "in": "path", "required": true }, + { + "type": "string", + "description": "身份域 OCID(缺省 Default 域)", + "name": "domainId", + "in": "query" + }, { "type": "string", "description": "idpId", @@ -1817,6 +2027,12 @@ "name": "id", "in": "path", "required": true + }, + { + "type": "string", + "description": "身份域 OCID(缺省 Default 域)", + "name": "domainId", + "in": "query" } ], "responses": { @@ -1847,6 +2063,12 @@ "in": "path", "required": true }, + { + "type": "string", + "description": "身份域 OCID(缺省 Default 域)", + "name": "domainId", + "in": "query" + }, { "description": "请求体(见接口说明)", "name": "body", @@ -3008,6 +3230,12 @@ "name": "id", "in": "path", "required": true + }, + { + "type": "string", + "description": "身份域 OCID(缺省 Default 域)", + "name": "domainId", + "in": "query" } ], "responses": { @@ -3038,6 +3266,12 @@ "in": "path", "required": true }, + { + "type": "string", + "description": "身份域 OCID(缺省 Default 域)", + "name": "domainId", + "in": "query" + }, { "description": "请求体(见接口说明)", "name": "body", @@ -3077,6 +3311,12 @@ "name": "id", "in": "path", "required": true + }, + { + "type": "string", + "description": "身份域 OCID(缺省 Default 域)", + "name": "domainId", + "in": "query" } ], "responses": { @@ -3109,6 +3349,12 @@ "in": "path", "required": true }, + { + "type": "string", + "description": "身份域 OCID(缺省 Default 域)", + "name": "domainId", + "in": "query" + }, { "type": "string", "description": "policyId", @@ -3224,6 +3470,12 @@ "name": "id", "in": "path", "required": true + }, + { + "type": "string", + "description": "身份域 OCID(缺省 Default 域)", + "name": "domainId", + "in": "query" } ], "responses": { @@ -3471,6 +3723,12 @@ "in": "path", "required": true }, + { + "type": "string", + "description": "身份域 OCID(缺省 Default 域)", + "name": "domainId", + "in": "query" + }, { "description": "请求体(见接口说明)", "name": "body", @@ -3511,6 +3769,12 @@ "in": "path", "required": true }, + { + "type": "string", + "description": "身份域 OCID(缺省 Default 域)", + "name": "domainId", + "in": "query" + }, { "type": "string", "description": "ruleId", @@ -3544,6 +3808,12 @@ "name": "id", "in": "path", "required": true + }, + { + "type": "string", + "description": "身份域 OCID(缺省 Default 域)", + "name": "domainId", + "in": "query" } ], "responses": { @@ -3828,6 +4098,12 @@ "name": "id", "in": "path", "required": true + }, + { + "type": "string", + "description": "身份域 OCID(缺省 Default 域)", + "name": "domainId", + "in": "query" } ], "responses": { @@ -3858,6 +4134,12 @@ "in": "path", "required": true }, + { + "type": "string", + "description": "身份域 OCID(缺省 Default 域)", + "name": "domainId", + "in": "query" + }, { "description": "请求体", "name": "body", @@ -3898,6 +4180,12 @@ "in": "path", "required": true }, + { + "type": "string", + "description": "身份域 OCID(缺省 Default 域)", + "name": "domainId", + "in": "query" + }, { "type": "string", "description": "userId", @@ -3934,6 +4222,12 @@ "in": "path", "required": true }, + { + "type": "string", + "description": "身份域 OCID(缺省 Default 域)", + "name": "domainId", + "in": "query" + }, { "type": "string", "description": "userId", @@ -3979,6 +4273,12 @@ "in": "path", "required": true }, + { + "type": "string", + "description": "身份域 OCID(缺省 Default 域)", + "name": "domainId", + "in": "query" + }, { "type": "string", "description": "userId", @@ -4051,6 +4351,12 @@ "in": "path", "required": true }, + { + "type": "string", + "description": "身份域 OCID(缺省 Default 域)", + "name": "domainId", + "in": "query" + }, { "type": "string", "description": "userId", @@ -4089,6 +4395,12 @@ "in": "path", "required": true }, + { + "type": "string", + "description": "身份域 OCID(缺省 Default 域)", + "name": "domainId", + "in": "query" + }, { "type": "string", "description": "userId", @@ -4742,6 +5054,95 @@ } } }, + "/api/v1/settings/notify-channels": { + "get": { + "security": [ + { + "BearerAuth": [] + } + ], + "tags": [ + "设置" + ], + "summary": "返回全部通知渠道的脱敏视图", + "responses": { + "200": { + "description": "OK", + "schema": { + "type": "object", + "additionalProperties": true + } + } + } + } + }, + "/api/v1/settings/notify-channels/{type}": { + "put": { + "security": [ + { + "BearerAuth": [] + } + ], + "tags": [ + "设置" + ], + "summary": "保存单渠道配置并返回全渠道最新视图", + "parameters": [ + { + "type": "string", + "description": "渠道类型(webhook/ntfy/bark/smtp)", + "name": "type", + "in": "path", + "required": true + }, + { + "description": "请求体", + "name": "body", + "in": "body", + "required": true, + "schema": { + "$ref": "#/definitions/internal_api.updateNotifyChannelRequest" + } + } + ], + "responses": { + "200": { + "description": "OK", + "schema": { + "type": "object", + "additionalProperties": true + } + } + } + } + }, + "/api/v1/settings/notify-channels/{type}/test": { + "post": { + "security": [ + { + "BearerAuth": [] + } + ], + "tags": [ + "设置" + ], + "summary": "向指定渠道发送测试消息", + "parameters": [ + { + "type": "string", + "description": "渠道类型(webhook/ntfy/bark/smtp)", + "name": "type", + "in": "path", + "required": true + } + ], + "responses": { + "204": { + "description": "无内容" + } + } + } + }, "/api/v1/settings/notify-events": { "get": { "security": [ @@ -5396,6 +5797,38 @@ } }, "definitions": { + "internal_api.alertRuleRequest": { + "type": "object", + "properties": { + "enabled": { + "type": "boolean" + }, + "eventTypes": { + "type": "string" + }, + "name": { + "type": "string" + }, + "ociConfigId": { + "type": "integer" + }, + "resourceMatch": { + "type": "string" + }, + "sourceIpMode": { + "type": "string" + }, + "sourceIps": { + "type": "string" + }, + "threshold": { + "type": "integer" + }, + "windowMinutes": { + "type": "integer" + } + } + }, "internal_api.attachBootVolumeRequest": { "type": "object", "required": [ @@ -5810,14 +6243,18 @@ ], "properties": { "password": { - "type": "string" + "type": "string", + "maxLength": 128 }, "totpCode": { "description": "两步验证码;账号已启用 TOTP 时必填,缺失返回 428", - "type": "string" + "type": "string", + "maxLength": 8 }, "username": { - "type": "string" + "description": "字段长度上限防超长输入撑爆登录守卫与 bcrypt(S-03)", + "type": "string", + "maxLength": 64 } } }, @@ -5884,6 +6321,44 @@ } } }, + "internal_api.updateNotifyChannelRequest": { + "type": "object", + "properties": { + "bodyTemplate": { + "type": "string" + }, + "enabled": { + "type": "boolean" + }, + "from": { + "type": "string" + }, + "host": { + "type": "string" + }, + "port": { + "type": "integer" + }, + "secret": { + "type": "string" + }, + "server": { + "type": "string" + }, + "to": { + "type": "string" + }, + "topic": { + "type": "string" + }, + "url": { + "type": "string" + }, + "username": { + "type": "string" + } + } + }, "internal_api.updateSecurityListRequest": { "type": "object", "properties": { diff --git a/docs/swagger.yaml b/docs/swagger.yaml index 0a6b900..887aa21 100644 --- a/docs/swagger.yaml +++ b/docs/swagger.yaml @@ -1,5 +1,26 @@ basePath: / definitions: + internal_api.alertRuleRequest: + properties: + enabled: + type: boolean + eventTypes: + type: string + name: + type: string + ociConfigId: + type: integer + resourceMatch: + type: string + sourceIpMode: + type: string + sourceIps: + type: string + threshold: + type: integer + windowMinutes: + type: integer + type: object internal_api.attachBootVolumeRequest: properties: bootVolumeId: @@ -277,11 +298,15 @@ definitions: internal_api.loginRequest: properties: password: + maxLength: 128 type: string totpCode: description: 两步验证码;账号已启用 TOTP 时必填,缺失返回 428 + maxLength: 8 type: string username: + description: 字段长度上限防超长输入撑爆登录守卫与 bcrypt(S-03) + maxLength: 64 type: string required: - password @@ -328,6 +353,31 @@ definitions: description: 更换 shape,运行中实例会自动重启 type: string type: object + internal_api.updateNotifyChannelRequest: + properties: + bodyTemplate: + type: string + enabled: + type: boolean + from: + type: string + host: + type: string + port: + type: integer + secret: + type: string + server: + type: string + to: + type: string + topic: + type: string + url: + type: string + username: + type: string + type: object internal_api.updateSecurityListRequest: properties: displayName: @@ -1078,6 +1128,20 @@ paths: summary: 密码登录开关 tags: - 认证 + /api/v1/auth/revoke-sessions: + post: + responses: + "200": + description: 新 token 与 expiresAt + schema: + additionalProperties: + type: string + type: object + security: + - BearerAuth: [] + summary: 撤销全部会话 + tags: + - 认证 /api/v1/auth/totp: get: responses: @@ -1203,6 +1267,78 @@ paths: summary: 回传日志事件列表 tags: - 任务与日志回传 + /api/v1/log-events/alert-rules: + get: + responses: + "200": + description: OK + schema: + additionalProperties: true + type: object + security: + - BearerAuth: [] + summary: 告警规则列表 + tags: + - 任务与日志回传 + post: + parameters: + - description: 请求体 + in: body + name: body + required: true + schema: + $ref: '#/definitions/internal_api.alertRuleRequest' + responses: + "201": + description: Created + schema: + additionalProperties: true + type: object + security: + - BearerAuth: [] + summary: 创建告警规则 + tags: + - 任务与日志回传 + /api/v1/log-events/alert-rules/{ruleId}: + delete: + parameters: + - description: 规则 ID + in: path + name: ruleId + required: true + type: integer + responses: + "204": + description: 无内容 + security: + - BearerAuth: [] + summary: 删除告警规则 + tags: + - 任务与日志回传 + put: + parameters: + - description: 规则 ID + in: path + name: ruleId + required: true + type: integer + - description: 请求体 + in: body + name: body + required: true + schema: + $ref: '#/definitions/internal_api.alertRuleRequest' + responses: + "200": + description: OK + schema: + additionalProperties: true + type: object + security: + - BearerAuth: [] + summary: 更新告警规则 + tags: + - 任务与日志回传 /api/v1/oci-configs: get: responses: @@ -1237,6 +1373,7 @@ paths: - 租户配置 /api/v1/oci-configs/{id}: delete: + description: 删除租户配置及其本地关联任务、快照、回传、告警和 AI 渠道数据,并撤销 Webhook 密钥;不删除 OCI 云端资源。 parameters: - description: 配置 ID in: path @@ -1301,6 +1438,14 @@ paths: name: id required: true type: integer + - description: 续查游标(上次响应原样带回) + in: query + name: cursor + type: string + - description: 单批目标条数,缺省 100,上限 200 + in: query + name: limit + type: integer responses: "200": description: OK @@ -1309,7 +1454,7 @@ paths: type: object security: - BearerAuth: [] - summary: 实时查询租户 OCI 审计事件:hours 首查 + summary: 批式懒加载查询租户 OCI 审计事件 tags: - 租户 IAM /api/v1/oci-configs/{id}/audit-events/detail: @@ -1560,6 +1705,25 @@ paths: summary: 配置成本快照 tags: - 成本 + /api/v1/oci-configs/{id}/domains: + get: + parameters: + - description: 配置 ID + in: path + name: id + required: true + type: integer + responses: + "200": + description: OK + schema: + additionalProperties: true + type: object + security: + - BearerAuth: [] + summary: 租户身份域列表 + tags: + - 租户 IAM /api/v1/oci-configs/{id}/identity-providers: get: parameters: @@ -1568,6 +1732,10 @@ paths: name: id required: true type: integer + - description: 身份域 OCID(缺省 Default 域) + in: query + name: domainId + type: string responses: "200": description: OK @@ -1586,6 +1754,10 @@ paths: name: id required: true type: integer + - description: 身份域 OCID(缺省 Default 域) + in: query + name: domainId + type: string - description: 请求体 in: body name: body @@ -1611,6 +1783,10 @@ paths: name: id required: true type: integer + - description: 身份域 OCID(缺省 Default 域) + in: query + name: domainId + type: string - description: idpId in: path name: idpId @@ -1632,6 +1808,10 @@ paths: name: id required: true type: integer + - description: 身份域 OCID(缺省 Default 域) + in: query + name: domainId + type: string - description: idpId in: path name: idpId @@ -1662,6 +1842,10 @@ paths: name: id required: true type: integer + - description: 身份域 OCID(缺省 Default 域) + in: query + name: domainId + type: string responses: "200": description: OK @@ -1680,6 +1864,10 @@ paths: name: id required: true type: integer + - description: 身份域 OCID(缺省 Default 域) + in: query + name: domainId + type: string - description: 请求体(见接口说明) in: body name: body @@ -2412,6 +2600,10 @@ paths: name: id required: true type: integer + - description: 身份域 OCID(缺省 Default 域) + in: query + name: domainId + type: string responses: "200": description: OK @@ -2430,6 +2622,10 @@ paths: name: id required: true type: integer + - description: 身份域 OCID(缺省 Default 域) + in: query + name: domainId + type: string - description: 请求体(见接口说明) in: body name: body @@ -2455,6 +2651,10 @@ paths: name: id required: true type: integer + - description: 身份域 OCID(缺省 Default 域) + in: query + name: domainId + type: string responses: "200": description: OK @@ -2474,6 +2674,10 @@ paths: name: id required: true type: integer + - description: 身份域 OCID(缺省 Default 域) + in: query + name: domainId + type: string - description: policyId in: path name: policyId @@ -2547,6 +2751,10 @@ paths: name: id required: true type: integer + - description: 身份域 OCID(缺省 Default 域) + in: query + name: domainId + type: string responses: "200": description: OK @@ -2701,6 +2909,10 @@ paths: name: id required: true type: integer + - description: 身份域 OCID(缺省 Default 域) + in: query + name: domainId + type: string - description: 请求体(见接口说明) in: body name: body @@ -2726,6 +2938,10 @@ paths: name: id required: true type: integer + - description: 身份域 OCID(缺省 Default 域) + in: query + name: domainId + type: string - description: ruleId in: path name: ruleId @@ -2747,6 +2963,10 @@ paths: name: id required: true type: integer + - description: 身份域 OCID(缺省 Default 域) + in: query + name: domainId + type: string responses: "200": description: OK @@ -2925,6 +3145,10 @@ paths: name: id required: true type: integer + - description: 身份域 OCID(缺省 Default 域) + in: query + name: domainId + type: string responses: "200": description: OK @@ -2943,6 +3167,10 @@ paths: name: id required: true type: integer + - description: 身份域 OCID(缺省 Default 域) + in: query + name: domainId + type: string - description: 请求体 in: body name: body @@ -2968,6 +3196,10 @@ paths: name: id required: true type: integer + - description: 身份域 OCID(缺省 Default 域) + in: query + name: domainId + type: string - description: userId in: path name: userId @@ -2988,6 +3220,10 @@ paths: name: id required: true type: integer + - description: 身份域 OCID(缺省 Default 域) + in: query + name: domainId + type: string - description: userId in: path name: userId @@ -3011,6 +3247,10 @@ paths: name: id required: true type: integer + - description: 身份域 OCID(缺省 Default 域) + in: query + name: domainId + type: string - description: userId in: path name: userId @@ -3065,6 +3305,10 @@ paths: name: id required: true type: integer + - description: 身份域 OCID(缺省 Default 域) + in: query + name: domainId + type: string - description: userId in: path name: userId @@ -3089,6 +3333,10 @@ paths: name: id required: true type: integer + - description: 身份域 OCID(缺省 Default 域) + in: query + name: domainId + type: string - description: userId in: path name: userId @@ -3498,6 +3746,60 @@ paths: summary: 返回本地维护的完整区域表 tags: - 租户配置 + /api/v1/settings/notify-channels: + get: + responses: + "200": + description: OK + schema: + additionalProperties: true + type: object + security: + - BearerAuth: [] + summary: 返回全部通知渠道的脱敏视图 + tags: + - 设置 + /api/v1/settings/notify-channels/{type}: + put: + parameters: + - description: 渠道类型(webhook/ntfy/bark/smtp) + in: path + name: type + required: true + type: string + - description: 请求体 + in: body + name: body + required: true + schema: + $ref: '#/definitions/internal_api.updateNotifyChannelRequest' + responses: + "200": + description: OK + schema: + additionalProperties: true + type: object + security: + - BearerAuth: [] + summary: 保存单渠道配置并返回全渠道最新视图 + tags: + - 设置 + /api/v1/settings/notify-channels/{type}/test: + post: + parameters: + - description: 渠道类型(webhook/ntfy/bark/smtp) + in: path + name: type + required: true + type: string + responses: + "204": + description: 无内容 + security: + - BearerAuth: [] + summary: 向指定渠道发送测试消息 + tags: + - 设置 /api/v1/settings/notify-events: get: responses: diff --git a/internal/api/auth.go b/internal/api/auth.go index 14fde19..5372215 100644 --- a/internal/api/auth.go +++ b/internal/api/auth.go @@ -18,10 +18,11 @@ type authHandler struct { } type loginRequest struct { - Username string `json:"username" binding:"required"` - Password string `json:"password" binding:"required"` + // 字段长度上限防超长输入撑爆登录守卫与 bcrypt(S-03) + Username string `json:"username" binding:"required,max=64"` + Password string `json:"password" binding:"required,max=128"` // 两步验证码;账号已启用 TOTP 时必填,缺失返回 428 - TotpCode string `json:"totpCode"` + TotpCode string `json:"totpCode" binding:"max=8"` } // login 校验用户名密码(与可选 TOTP)后签发 JWT。 @@ -64,7 +65,7 @@ func (h *authHandler) login(c *gin.Context) { return } if err != nil { - c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) + respondError(c, err) return } h.recordLogin(c, req.Username, http.StatusOK, start) diff --git a/internal/api/authx.go b/internal/api/authx.go index 0789c20..4eaba3b 100644 --- a/internal/api/authx.go +++ b/internal/api/authx.go @@ -81,7 +81,7 @@ func (h *authxHandler) totpActivate(c *gin.Context) { respondError(c, err) return } - c.Status(http.StatusNoContent) + h.respondFreshToken(c) } // totpDisable 停用两步验证;需当前验证码或登录密码任一确认。 @@ -110,7 +110,34 @@ func (h *authxHandler) totpDisable(c *gin.Context) { respondError(c, err) return } - c.Status(http.StatusNoContent) + h.respondFreshToken(c) +} + +// respondFreshToken 敏感变更后为操作者签发新令牌返回(版本已递增, +// 旧令牌全部失效);签发失败降级 204,前端按会话失效走重新登录。 +func (h *authxHandler) respondFreshToken(c *gin.Context) { + token, expires, err := h.auth.IssueToken(c.Request.Context(), c.GetString(usernameKey)) + if err != nil { + c.Status(http.StatusNoContent) + return + } + c.JSON(http.StatusOK, gin.H{"token": token, "expiresAt": expires}) +} + +// revokeSessions 撤销全部会话(令牌版本递增),并为当前操作者重签新令牌。 +// +// @Summary 撤销全部会话 +// @Tags 认证 +// @Success 200 {object} map[string]string "新 token 与 expiresAt" +// @Security BearerAuth +// @Router /api/v1/auth/revoke-sessions [post] +func (h *authxHandler) revokeSessions(c *gin.Context) { + token, expires, err := h.auth.RevokeSessions(c.Request.Context(), c.GetString(usernameKey)) + if err != nil { + respondError(c, err) + return + } + c.JSON(http.StatusOK, gin.H{"token": token, "expiresAt": expires}) } // ---- 登录凭据(JWT 组内) ---- @@ -149,7 +176,7 @@ func (h *authxHandler) updateCredentials(c *gin.Context) { c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) return } - err := h.auth.UpdateCredentials(c.Request.Context(), c.GetString(usernameKey), req) + finalName, err := h.auth.UpdateCredentials(c.Request.Context(), c.GetString(usernameKey), req) if errors.Is(err, service.ErrCredentialConfirm) { c.JSON(http.StatusUnauthorized, gin.H{"error": err.Error()}) return @@ -162,7 +189,8 @@ func (h *authxHandler) updateCredentials(c *gin.Context) { respondError(c, err) return } - c.Status(http.StatusNoContent) + c.Set(usernameKey, finalName) + h.respondFreshToken(c) } // updatePasswordLogin 保存密码登录禁用开关;开启需至少绑定一个外部身份。 @@ -191,7 +219,7 @@ func (h *authxHandler) updatePasswordLogin(c *gin.Context) { respondError(c, err) return } - c.Status(http.StatusNoContent) + h.respondFreshToken(c) } // ---- OAuth ---- @@ -254,7 +282,7 @@ func (h *authxHandler) bearerUser(c *gin.Context) (string, bool) { if len(header) < 8 || header[:7] != "Bearer " { return "", false } - username, err := h.auth.ParseToken(header[7:]) + username, err := h.auth.ParseToken(c.Request.Context(), header[7:]) if err != nil { return "", false } @@ -282,9 +310,9 @@ func (h *authxHandler) oauthCallback(c *gin.Context) { c.Redirect(http.StatusFound, target+"?oauthError="+url.QueryEscape(oauthErrText(err))) return } - if token == "" { - // 绑定模式:回设置页安全 tab - c.Redirect(http.StatusFound, "/settings?oauth=bound") + if mode == "bind" { + // 绑定模式:版本已递增,新 token 经 fragment 带回设置页无感换发 + c.Redirect(http.StatusFound, "/settings?oauth=bound#oauthToken="+url.QueryEscape(token)) return } // 登录模式:token 放 fragment,不进服务端日志与 Referer @@ -340,7 +368,7 @@ func (h *authxHandler) unbindIdentity(c *gin.Context) { respondError(c, err) return } - c.Status(http.StatusNoContent) + h.respondFreshToken(c) } // ---- OAuth provider 设置(JWT 组内) ---- diff --git a/internal/api/federation.go b/internal/api/federation.go index 4bfe1a0..6871862 100644 --- a/internal/api/federation.go +++ b/internal/api/federation.go @@ -40,6 +40,7 @@ func boolOr(v *bool, def bool) bool { // @Summary 身份提供商列表 // @Tags 租户 IAM // @Param id path int true "配置 ID" +// @Param domainId query string false "身份域 OCID(缺省 Default 域)" // @Success 200 {object} map[string]any // @Security BearerAuth // @Router /api/v1/oci-configs/{id}/identity-providers [get] @@ -48,7 +49,7 @@ func (h *ociConfigHandler) listIdentityProviders(c *gin.Context) { if !ok { return } - idps, err := h.svc.IdentityProviders(c.Request.Context(), id) + idps, err := h.svc.IdentityProviders(c.Request.Context(), id, c.Query("domainId")) if err != nil { respondError(c, err) return @@ -59,6 +60,7 @@ func (h *ociConfigHandler) listIdentityProviders(c *gin.Context) { // @Summary 创建身份提供商(SAML) // @Tags 租户 IAM // @Param id path int true "配置 ID" +// @Param domainId query string false "身份域 OCID(缺省 Default 域)" // @Param body body createIdpRequest true "请求体" // @Success 201 {object} map[string]any // @Security BearerAuth @@ -73,7 +75,7 @@ func (h *ociConfigHandler) createIdentityProvider(c *gin.Context) { c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) return } - idp, err := h.svc.CreateIdentityProvider(c.Request.Context(), id, oci.CreateIdpInput{ + idp, err := h.svc.CreateIdentityProvider(c.Request.Context(), id, c.Query("domainId"), oci.CreateIdpInput{ Name: req.Name, Metadata: req.Metadata, Description: req.Description, @@ -97,6 +99,7 @@ func (h *ociConfigHandler) createIdentityProvider(c *gin.Context) { // @Summary 激活身份提供商 // @Tags 租户 IAM // @Param id path int true "配置 ID" +// @Param domainId query string false "身份域 OCID(缺省 Default 域)" // @Param idpId path string true "idpId" // @Param body body object true "请求体(见接口说明)" // @Success 200 {object} map[string]any @@ -114,7 +117,7 @@ func (h *ociConfigHandler) activateIdentityProvider(c *gin.Context) { c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) return } - idp, err := h.svc.SetIdentityProviderEnabled(c.Request.Context(), id, c.Param("idpId"), *req.Enabled) + idp, err := h.svc.SetIdentityProviderEnabled(c.Request.Context(), id, c.Query("domainId"), c.Param("idpId"), *req.Enabled) if err != nil { respondError(c, err) return @@ -125,6 +128,7 @@ func (h *ociConfigHandler) activateIdentityProvider(c *gin.Context) { // @Summary 删除身份提供商 // @Tags 租户 IAM // @Param id path int true "配置 ID" +// @Param domainId query string false "身份域 OCID(缺省 Default 域)" // @Param idpId path string true "idpId" // @Success 204 "无内容" // @Security BearerAuth @@ -134,7 +138,7 @@ func (h *ociConfigHandler) deleteIdentityProvider(c *gin.Context) { if !ok { return } - if err := h.svc.DeleteIdentityProvider(c.Request.Context(), id, c.Param("idpId")); err != nil { + if err := h.svc.DeleteIdentityProvider(c.Request.Context(), id, c.Query("domainId"), c.Param("idpId")); err != nil { respondError(c, err) return } @@ -144,6 +148,7 @@ func (h *ociConfigHandler) deleteIdentityProvider(c *gin.Context) { // @Summary 下载 SAML 元数据 // @Tags 租户 IAM // @Param id path int true "配置 ID" +// @Param domainId query string false "身份域 OCID(缺省 Default 域)" // @Success 200 {object} map[string]any // @Security BearerAuth // @Router /api/v1/oci-configs/{id}/saml-metadata [get] @@ -152,7 +157,7 @@ func (h *ociConfigHandler) downloadSamlMetadata(c *gin.Context) { if !ok { return } - metadata, err := h.svc.DomainSamlMetadata(c.Request.Context(), id) + metadata, err := h.svc.DomainSamlMetadata(c.Request.Context(), id, c.Query("domainId")) if err != nil { respondError(c, err) return @@ -164,6 +169,7 @@ func (h *ociConfigHandler) downloadSamlMetadata(c *gin.Context) { // @Summary 单点登录规则列表 // @Tags 租户 IAM // @Param id path int true "配置 ID" +// @Param domainId query string false "身份域 OCID(缺省 Default 域)" // @Success 200 {object} map[string]any // @Security BearerAuth // @Router /api/v1/oci-configs/{id}/sign-on-rules [get] @@ -172,7 +178,7 @@ func (h *ociConfigHandler) listSignOnRules(c *gin.Context) { if !ok { return } - rules, err := h.svc.ConsoleSignOnRules(c.Request.Context(), id) + rules, err := h.svc.ConsoleSignOnRules(c.Request.Context(), id, c.Query("domainId")) if err != nil { respondError(c, err) return @@ -183,6 +189,7 @@ func (h *ociConfigHandler) listSignOnRules(c *gin.Context) { // @Summary 创建 MFA 豁免规则 // @Tags 租户 IAM // @Param id path int true "配置 ID" +// @Param domainId query string false "身份域 OCID(缺省 Default 域)" // @Param body body object true "请求体(见接口说明)" // @Success 201 {object} map[string]any // @Security BearerAuth @@ -199,7 +206,7 @@ func (h *ociConfigHandler) createMfaExemption(c *gin.Context) { c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) return } - rule, err := h.svc.CreateMfaExemption(c.Request.Context(), id, req.IdentityProviderID) + rule, err := h.svc.CreateMfaExemption(c.Request.Context(), id, c.Query("domainId"), req.IdentityProviderID) if err != nil { respondError(c, err) return @@ -210,6 +217,7 @@ func (h *ociConfigHandler) createMfaExemption(c *gin.Context) { // @Summary 删除 MFA 豁免规则 // @Tags 租户 IAM // @Param id path int true "配置 ID" +// @Param domainId query string false "身份域 OCID(缺省 Default 域)" // @Param ruleId path string true "ruleId" // @Success 204 "无内容" // @Security BearerAuth @@ -219,7 +227,7 @@ func (h *ociConfigHandler) deleteMfaExemption(c *gin.Context) { if !ok { return } - if err := h.svc.DeleteMfaExemption(c.Request.Context(), id, c.Param("ruleId")); err != nil { + if err := h.svc.DeleteMfaExemption(c.Request.Context(), id, c.Query("domainId"), c.Param("ruleId")); err != nil { respondError(c, err) return } diff --git a/internal/api/logevent.go b/internal/api/logevent.go index 98ebf8d..11d33dc 100644 --- a/internal/api/logevent.go +++ b/internal/api/logevent.go @@ -7,6 +7,7 @@ import ( "github.com/gin-gonic/gin" + "oci-portal/internal/model" "oci-portal/internal/service" ) @@ -106,6 +107,127 @@ func (h *logEventHandler) list(c *gin.Context) { c.JSON(http.StatusOK, gin.H{"items": items, "total": total}) } +// alertRuleRequest 是创建/更新告警规则的请求体。 +type alertRuleRequest struct { + Name string `json:"name"` + Enabled bool `json:"enabled"` + OciConfigID uint `json:"ociConfigId"` + EventTypes string `json:"eventTypes"` + SourceIPs string `json:"sourceIps"` + SourceIPMode string `json:"sourceIpMode"` + ResourceMatch string `json:"resourceMatch"` + Threshold int `json:"threshold"` + WindowMinutes int `json:"windowMinutes"` +} + +// toModel 转为规则模型(默认阈值 1)。 +func (r alertRuleRequest) toModel() model.AlertRule { + if r.Threshold == 0 { + r.Threshold = 1 + } + return model.AlertRule{ + Name: r.Name, Enabled: r.Enabled, OciConfigID: r.OciConfigID, + EventTypes: r.EventTypes, SourceIPs: r.SourceIPs, SourceIPMode: r.SourceIPMode, + ResourceMatch: r.ResourceMatch, Threshold: r.Threshold, WindowMinutes: r.WindowMinutes, + } +} + +// listAlertRules 返回全部告警规则。 +// +// @Summary 告警规则列表 +// @Tags 任务与日志回传 +// @Success 200 {object} map[string]any +// @Security BearerAuth +// @Router /api/v1/log-events/alert-rules [get] +func (h *logEventHandler) listAlertRules(c *gin.Context) { + rules, err := h.svc.ListAlertRules(c.Request.Context()) + if err != nil { + respondError(c, err) + return + } + c.JSON(http.StatusOK, gin.H{"items": rules}) +} + +// createAlertRule 创建告警规则;字段非法返回 400。 +// +// @Summary 创建告警规则 +// @Tags 任务与日志回传 +// @Param body body alertRuleRequest true "请求体" +// @Success 201 {object} map[string]any +// @Security BearerAuth +// @Router /api/v1/log-events/alert-rules [post] +func (h *logEventHandler) createAlertRule(c *gin.Context) { + var req alertRuleRequest + if err := c.ShouldBindJSON(&req); err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) + return + } + rule, err := h.svc.CreateAlertRule(c.Request.Context(), req.toModel()) + if err != nil { + respondAlertRuleError(c, err) + return + } + c.JSON(http.StatusCreated, rule) +} + +// updateAlertRule 整体覆盖告警规则(含启停)。 +// +// @Summary 更新告警规则 +// @Tags 任务与日志回传 +// @Param ruleId path int true "规则 ID" +// @Param body body alertRuleRequest true "请求体" +// @Success 200 {object} map[string]any +// @Security BearerAuth +// @Router /api/v1/log-events/alert-rules/{ruleId} [put] +func (h *logEventHandler) updateAlertRule(c *gin.Context) { + id, err := strconv.ParseUint(c.Param("ruleId"), 10, 64) + if err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": "规则 ID 非法"}) + return + } + var req alertRuleRequest + if err := c.ShouldBindJSON(&req); err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) + return + } + rule, err := h.svc.UpdateAlertRule(c.Request.Context(), uint(id), req.toModel()) + if err != nil { + respondAlertRuleError(c, err) + return + } + c.JSON(http.StatusOK, rule) +} + +// deleteAlertRule 删除告警规则及其命中记录。 +// +// @Summary 删除告警规则 +// @Tags 任务与日志回传 +// @Param ruleId path int true "规则 ID" +// @Success 204 "无内容" +// @Security BearerAuth +// @Router /api/v1/log-events/alert-rules/{ruleId} [delete] +func (h *logEventHandler) deleteAlertRule(c *gin.Context) { + id, err := strconv.ParseUint(c.Param("ruleId"), 10, 64) + if err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": "规则 ID 非法"}) + return + } + if err := h.svc.DeleteAlertRule(c.Request.Context(), uint(id)); err != nil { + respondError(c, err) + return + } + c.Status(http.StatusNoContent) +} + +// respondAlertRuleError 把字段校验错误映射为 400。 +func respondAlertRuleError(c *gin.Context, err error) { + if errors.Is(err, service.ErrInvalidAlertRule) { + c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) + return + } + respondError(c, err) +} + // getRelay 查询 OCI 侧链路状态与关键事件清单。 // // @Summary 查询 OCI 侧链路状态与关键事件清单 diff --git a/internal/api/middleware.go b/internal/api/middleware.go index 7364b31..e87e1a8 100644 --- a/internal/api/middleware.go +++ b/internal/api/middleware.go @@ -15,7 +15,8 @@ import ( // usernameKey 是鉴权通过后写入 gin.Context 的用户名键。 const usernameKey = "username" -// RequireAuth 校验 Authorization: Bearer 令牌,通过后把用户名放进上下文。 +// RequireAuth 校验 Authorization: Bearer 令牌(签名、有效期与令牌版本), +// 通过后把用户名放进上下文。 func RequireAuth(auth *service.AuthService) gin.HandlerFunc { return func(c *gin.Context) { token, ok := strings.CutPrefix(c.GetHeader("Authorization"), "Bearer ") @@ -23,7 +24,7 @@ func RequireAuth(auth *service.AuthService) gin.HandlerFunc { c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error": "missing bearer token"}) return } - username, err := auth.ParseToken(token) + username, err := auth.ParseToken(c.Request.Context(), token) if err != nil { c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error": "invalid or expired token"}) return @@ -33,6 +34,17 @@ func RequireAuth(auth *service.AuthService) gin.HandlerFunc { } } +// bodyLimit 给请求体套上限(S-03):超限读取由 MaxBytesReader 截断报错, +// 绑定失败返回 400;webhook 入口另有独立 256KiB 自限,不经此中间件。 +func bodyLimit(n int64) gin.HandlerFunc { + return func(c *gin.Context) { + if c.Request.Body != nil { + c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, n) + } + c.Next() + } +} + // systemLogMiddleware 把写请求(POST/PUT/DELETE)异步记入系统日志,只读请求跳过。 // 只记方法、路径等元数据,绝不读请求体——请求体可能含私钥 / 口令等敏感数据。 func systemLogMiddleware(logs *service.SystemLogService) gin.HandlerFunc { diff --git a/internal/api/ociconfig.go b/internal/api/ociconfig.go index c9868d0..2aa38b6 100644 --- a/internal/api/ociconfig.go +++ b/internal/api/ociconfig.go @@ -1,7 +1,10 @@ package api import ( + "crypto/rand" + "encoding/hex" "errors" + "log" "net/http" "strconv" @@ -74,7 +77,7 @@ func (h *ociConfigHandler) create(c *gin.Context) { func (h *ociConfigHandler) list(c *gin.Context) { items, err := h.svc.List(c.Request.Context()) if err != nil { - c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) + respondError(c, err) return } c.JSON(http.StatusOK, items) @@ -168,6 +171,7 @@ func (h *ociConfigHandler) update(c *gin.Context) { } // @Summary 删除租户配置 +// @Description 删除租户配置及其本地关联任务、快照、回传、告警和 AI 渠道数据,并撤销 Webhook 密钥;不删除 OCI 云端资源。 // @Tags 租户配置 // @Param id path int true "配置 ID" // @Success 204 "无内容" @@ -258,17 +262,29 @@ func pathID(c *gin.Context) (uint, bool) { return uint(id), true } +// respondError 统一出错响应边界(S-08):OCI 服务错误经脱敏透出,记录缺失 +// 映射 404;其余按内部错误处理——响应只含固定文案与关联 ID,完整原因 +// (可能携带 SQL / DSN / 路径等内部细节)仅写服务端日志,经同一 ID 对应。 func respondError(c *gin.Context, err error) { - status := http.StatusInternalServerError - if errors.Is(err, gorm.ErrRecordNotFound) { - status = http.StatusNotFound - } var svcErr common.ServiceError if errors.As(err, &svcErr) { respondOCIError(c, err, svcErr) return } - c.JSON(status, gin.H{"error": err.Error()}) + if errors.Is(err, gorm.ErrRecordNotFound) { + c.JSON(http.StatusNotFound, gin.H{"error": "资源不存在"}) + return + } + id := newRequestID() + log.Printf("[ERR %s] %s %s: %v", id, c.Request.Method, requestPath(c), err) + c.JSON(http.StatusInternalServerError, gin.H{"error": "服务器内部错误", "requestId": id}) +} + +// newRequestID 生成错误关联 ID(8 字节随机 hex),响应与服务端日志据此对应。 +func newRequestID() string { + b := make([]byte, 8) + _, _ = rand.Read(b) + return hex.EncodeToString(b) } // respondOCIError 提炼 OCI 服务错误:透传 HTTP 状态码, @@ -278,9 +294,10 @@ func respondOCIError(c *gin.Context, err error, svcErr common.ServiceError) { if status < http.StatusBadRequest { status = http.StatusBadGateway } - // 面板的 401 专属本地 JWT 失效(前端收到即登出); - // 上游 OCI 的 401 是代理调用被拒,改用 502 透出 - if status == http.StatusUnauthorized { + // 面板的 401 专属本地 JWT 失效(前端收到即登出)、429 专属本地 IP 限速 + // (前端收到即跳 /blocked);上游 OCI 的同码是代理调用被拒/云端限流, + // 均改用 502 透出,body 保留 ociCode 供辨识 + if status == http.StatusUnauthorized || status == http.StatusTooManyRequests { status = http.StatusBadGateway } body := gin.H{ diff --git a/internal/api/ociconfig_test.go b/internal/api/ociconfig_test.go new file mode 100644 index 0000000..5f10c04 --- /dev/null +++ b/internal/api/ociconfig_test.go @@ -0,0 +1,54 @@ +package api + +import ( + "fmt" + "net/http/httptest" + "testing" + + "github.com/gin-gonic/gin" +) + +// fakeServiceError 模拟 oci-go-sdk 的 common.ServiceError,驱动 respondOCIError 分支。 +type fakeServiceError struct { + status int + code string + message string +} + +func (e fakeServiceError) Error() string { + return fmt.Sprintf("Error returned by service. Http Status Code: %d. Error Code: %s. Message: %s", + e.status, e.code, e.message) +} + +func (e fakeServiceError) GetHTTPStatusCode() int { return e.status } +func (e fakeServiceError) GetMessage() string { return e.message } +func (e fakeServiceError) GetCode() string { return e.code } +func (e fakeServiceError) GetOpcRequestID() string { return "req-1" } + +// TestRespondErrorOCIStatus 覆盖上游状态码改写规则: +// 401/429 在前端有本地专属语义(登出 / 跳 /blocked),上游同码须改写 502。 +func TestRespondErrorOCIStatus(t *testing.T) { + gin.SetMode(gin.TestMode) + cases := []struct { + name string + upstream int + wantStatus int + }{ + {"404 原样透传", 404, 404}, + {"409 原样透传", 409, 409}, + {"上游 401 改写 502", 401, 502}, + {"上游 429 改写 502", 429, 502}, + {"非错误码兜底 502", 200, 502}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + err := fmt.Errorf("获取流量: %w", fakeServiceError{tc.upstream, "TooManyRequests", "slow down"}) + respondError(c, err) + if rec.Code != tc.wantStatus { + t.Fatalf("status = %d, want %d", rec.Code, tc.wantStatus) + } + }) + } +} diff --git a/internal/api/region.go b/internal/api/region.go index b2949b2..5d746ff 100644 --- a/internal/api/region.go +++ b/internal/api/region.go @@ -19,7 +19,7 @@ import ( func listRegions(c *gin.Context) { regions, err := oci.AllRegions() if err != nil { - c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) + respondError(c, err) return } c.JSON(http.StatusOK, regions) diff --git a/internal/api/router.go b/internal/api/router.go index 512d6d0..3744f7e 100644 --- a/internal/api/router.go +++ b/internal/api/router.go @@ -21,7 +21,8 @@ func NewRouter(auth *service.AuthService, oauth *service.OAuthService, ociConfig accessLog := gin.LoggerWithConfig(gin.LoggerConfig{SkipQueryString: true}) r.Use(accessLog, RealIPMiddleware(settings), IPRateMiddleware(settings), gin.Recovery()) - v1 := r.Group("/api/v1") + // 面板 API 请求体统一 1MB 上限(S-03);AI 网关按对话体量单独 10MB + v1 := r.Group("/api/v1", bodyLimit(1<<20)) registerAuthPublic(v1, auth, oauth, systemLogs) // AI 网关对外端点:独立密钥鉴权,挂在全局 Use 之后,仍受 IP 限速与 Recovery 保护 registerAiGateway(r, aiGateway) diff --git a/internal/api/router_test.go b/internal/api/router_test.go index 3cab282..b110225 100644 --- a/internal/api/router_test.go +++ b/internal/api/router_test.go @@ -2,6 +2,8 @@ package api import ( "context" + "encoding/json" + "fmt" "net/http" "net/http/httptest" "strings" @@ -183,3 +185,103 @@ func TestSystemLogsEndpoint(t *testing.T) { }) } } + +// TestLoginFieldLimits 超长登录字段被绑定校验拒绝(S-03 高基数防护第一道)。 +func TestLoginFieldLimits(t *testing.T) { + r, _, _ := newTestRouter(t) + longName := strings.Repeat("a", 65) + longPass := strings.Repeat("b", 129) + tests := []struct { + name string + body string + }{ + {name: "用户名超长", body: `{"username":"` + longName + `","password":"x"}`}, + {name: "密码超长", body: `{"username":"admin","password":"` + longPass + `"}`}, + {name: "验证码超长", body: `{"username":"admin","password":"x","totpCode":"123456789"}`}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + w := doRequest(t, r, http.MethodPost, "/api/v1/auth/login", "", tt.body) + if w.Code != http.StatusBadRequest { + t.Errorf("status = %d, want 400", w.Code) + } + }) + } +} + +// TestBodyLimitRejectsHuge 面板 API 请求体超 1MB 被拒(S-03)。 +func TestBodyLimitRejectsHuge(t *testing.T) { + r, _, _ := newTestRouter(t) + huge := `{"username":"admin","password":"` + strings.Repeat("x", 2<<20) + `"}` + w := doRequest(t, r, http.MethodPost, "/api/v1/auth/login", "", huge) + if w.Code != http.StatusBadRequest { + t.Errorf("status = %d, want 400(body too large)", w.Code) + } +} + +// TestRevokeSessionsEndpoint 撤销全部会话:直接调用即生效, +// 响应带新 token,旧 token 随即失效。 +func TestRevokeSessionsEndpoint(t *testing.T) { + r, _, _ := newTestRouter(t) + login := doRequest(t, r, http.MethodPost, "/api/v1/auth/login", "", `{"username":"admin","password":"pass123"}`) + var sess struct{ Token string } + if err := json.Unmarshal(login.Body.Bytes(), &sess); err != nil || sess.Token == "" { + t.Fatalf("login: %s", login.Body.String()) + } + w := doRequest(t, r, http.MethodPost, "/api/v1/auth/revoke-sessions", sess.Token, "") + if w.Code != http.StatusOK || !strings.Contains(w.Body.String(), "token") { + t.Fatalf("revoke = %d %s", w.Code, w.Body.String()) + } + w = doRequest(t, r, http.MethodGet, "/api/v1/auth/credentials", sess.Token, "") + if w.Code != http.StatusUnauthorized { + t.Errorf("撤销后旧 token 访问 = %d, want 401", w.Code) + } +} + +// TestRespondErrorSanitizesInternal 内部错误边界(S-08):wrap 链细节不透传, +// 响应为固定文案 + requestId;记录缺失映射 404 固定文案。 +func TestRespondErrorSanitizesInternal(t *testing.T) { + gin.SetMode(gin.TestMode) + cases := []struct { + name string + err error + wantStatus int + wantBody string // 必须出现 + leaks []string // 不得出现 + }{ + { + name: "内部错误脱敏", + err: fmt.Errorf("query users: dial tcp 10.0.0.5:3306: dsn user:secretpw"), + wantStatus: http.StatusInternalServerError, + wantBody: "requestId", + leaks: []string{"secretpw", "10.0.0.5", "dsn", "dial tcp"}, + }, + { + name: "记录缺失固定文案", + err: fmt.Errorf("find oci config 5: %w", gorm.ErrRecordNotFound), + wantStatus: http.StatusNotFound, + wantBody: "资源不存在", + leaks: []string{"find oci config"}, + }, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodGet, "/api/v1/x", nil) + respondError(c, tc.err) + if rec.Code != tc.wantStatus { + t.Fatalf("status = %d, want %d", rec.Code, tc.wantStatus) + } + body := rec.Body.String() + if !strings.Contains(body, tc.wantBody) { + t.Errorf("body %q 应包含 %q", body, tc.wantBody) + } + for _, leak := range tc.leaks { + if strings.Contains(body, leak) { + t.Errorf("body %q 泄露内部细节 %q", body, leak) + } + } + }) + } +} diff --git a/internal/api/routes_ai.go b/internal/api/routes_ai.go index b899ac3..8738129 100644 --- a/internal/api/routes_ai.go +++ b/internal/api/routes_ai.go @@ -10,7 +10,7 @@ import ( // 挂在全局中间件之后,仍受 IP 限速与 Recovery 保护;高频调用自有 AiCallLog。 func registerAiGateway(r *gin.Engine, aiGateway *service.AiGatewayService) { aih := &aiGatewayHandler{gw: aiGateway} - ai := r.Group("/ai/v1", aih.auth) + ai := r.Group("/ai/v1", bodyLimit(10<<20), aih.auth) ai.POST("/chat/completions", aih.chatCompletions) ai.POST("/responses", aih.responses) ai.POST("/messages", aih.messages) diff --git a/internal/api/routes_auth.go b/internal/api/routes_auth.go index b55594b..8b0441d 100644 --- a/internal/api/routes_auth.go +++ b/internal/api/routes_auth.go @@ -26,13 +26,14 @@ func registerAuthSecured(secured *gin.RouterGroup, auth *service.AuthService, oa ax := &authxHandler{auth: auth, oauth: oauth} secured.GET("/auth/totp", ax.totpStatus) secured.POST("/auth/totp/setup", ax.totpSetup) - secured.POST("/auth/totp/activate", ax.totpActivate) - secured.POST("/auth/totp/disable", ax.totpDisable) secured.GET("/auth/credentials", ax.getCredentials) secured.PUT("/auth/credentials", ax.updateCredentials) - secured.PUT("/auth/password-login", ax.updatePasswordLogin) secured.GET("/auth/identities", ax.identities) + secured.POST("/auth/totp/activate", ax.totpActivate) + secured.POST("/auth/totp/disable", ax.totpDisable) + secured.PUT("/auth/password-login", ax.updatePasswordLogin) secured.DELETE("/auth/identities/:id", ax.unbindIdentity) + secured.POST("/auth/revoke-sessions", ax.revokeSessions) secured.GET("/settings/oauth", func(c *gin.Context) { ax.getOAuthSettings(c, settings) }) secured.PUT("/settings/oauth", func(c *gin.Context) { ax.updateOAuthSettings(c, settings) }) } diff --git a/internal/api/routes_oci.go b/internal/api/routes_oci.go index 2d2d008..421fdf1 100644 --- a/internal/api/routes_oci.go +++ b/internal/api/routes_oci.go @@ -109,6 +109,7 @@ func registerOciCost(g *gin.RouterGroup, h *ociConfigHandler) { func registerOciTenantIAM(g *gin.RouterGroup, h *ociConfigHandler) { g.GET("/oci-configs/:id/audit-events", h.getAuditEvents) g.GET("/oci-configs/:id/audit-events/detail", h.getAuditEventDetail) + g.GET("/oci-configs/:id/domains", h.listIdentityDomains) g.GET("/oci-configs/:id/users", h.listTenantUsers) g.POST("/oci-configs/:id/users", h.createTenantUser) g.GET("/oci-configs/:id/users/:userId", h.getTenantUserDetail) diff --git a/internal/api/routes_settings.go b/internal/api/routes_settings.go index 715c9dc..69b2550 100644 --- a/internal/api/routes_settings.go +++ b/internal/api/routes_settings.go @@ -15,6 +15,9 @@ func registerSettings(secured *gin.RouterGroup, settings *service.SettingService secured.GET("/settings/telegram", st.getTelegram) secured.PUT("/settings/telegram", st.updateTelegram) secured.POST("/settings/telegram/test", st.testTelegram) + secured.GET("/settings/notify-channels", st.listNotifyChannels) + secured.PUT("/settings/notify-channels/:type", st.updateNotifyChannel) + secured.POST("/settings/notify-channels/:type/test", st.testNotifyChannel) secured.GET("/settings/notify-events", st.getNotifyEvents) secured.PUT("/settings/notify-events", st.updateNotifyEvents) secured.GET("/settings/notify-templates", st.listNotifyTemplates) diff --git a/internal/api/routes_tasks.go b/internal/api/routes_tasks.go index 40851e2..6d0c30d 100644 --- a/internal/api/routes_tasks.go +++ b/internal/api/routes_tasks.go @@ -19,6 +19,10 @@ func registerTasksAndLogs(secured *gin.RouterGroup, tasks *service.TaskService, le := &logEventHandler{svc: logEvents} secured.GET("/log-events", le.list) + secured.GET("/log-events/alert-rules", le.listAlertRules) + secured.POST("/log-events/alert-rules", le.createAlertRule) + secured.PUT("/log-events/alert-rules/:ruleId", le.updateAlertRule) + secured.DELETE("/log-events/alert-rules/:ruleId", le.deleteAlertRule) secured.GET("/oci-configs/:id/log-webhook", le.getWebhook) secured.POST("/oci-configs/:id/log-webhook", le.ensureWebhook) secured.DELETE("/oci-configs/:id/log-webhook", le.revokeWebhook) diff --git a/internal/api/settings.go b/internal/api/settings.go index 7c65894..c3852c9 100644 --- a/internal/api/settings.go +++ b/internal/api/settings.go @@ -80,6 +80,82 @@ func (h *settingsHandler) testTelegram(c *gin.Context) { c.Status(http.StatusNoContent) } +// listNotifyChannels 返回全部通知渠道的脱敏视图(webhook/ntfy/bark/smtp,不含 telegram)。 +// +// @Summary 返回全部通知渠道的脱敏视图 +// @Tags 设置 +// @Success 200 {object} map[string]any +// @Security BearerAuth +// @Router /api/v1/settings/notify-channels [get] +func (h *settingsHandler) listNotifyChannels(c *gin.Context) { + items, err := h.svc.NotifyChannels(c.Request.Context()) + if err != nil { + respondError(c, err) + return + } + c.JSON(http.StatusOK, gin.H{"items": items}) +} + +// updateNotifyChannelRequest 是保存单渠道配置的请求体; +// secret 承载渠道密文字段(ntfy token / bark device key / smtp 密码), +// 缺省沿用已存值,空串清除。 +type updateNotifyChannelRequest struct { + Enabled bool `json:"enabled"` + URL string `json:"url"` + BodyTemplate string `json:"bodyTemplate"` + Server string `json:"server"` + Topic string `json:"topic"` + Host string `json:"host"` + Port int `json:"port"` + Username string `json:"username"` + From string `json:"from"` + To string `json:"to"` + Secret *string `json:"secret"` +} + +// updateNotifyChannel 保存单渠道配置并返回全渠道最新视图;未知类型或校验失败返回 400。 +// +// @Summary 保存单渠道配置并返回全渠道最新视图 +// @Tags 设置 +// @Param type path string true "渠道类型(webhook/ntfy/bark/smtp)" +// @Param body body updateNotifyChannelRequest true "请求体" +// @Success 200 {object} map[string]any +// @Security BearerAuth +// @Router /api/v1/settings/notify-channels/{type} [put] +func (h *settingsHandler) updateNotifyChannel(c *gin.Context) { + var req updateNotifyChannelRequest + if err := c.ShouldBindJSON(&req); err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) + return + } + err := h.svc.UpdateNotifyChannel(c.Request.Context(), c.Param("type"), service.UpdateNotifyChannelInput{ + Enabled: req.Enabled, URL: req.URL, BodyTemplate: req.BodyTemplate, + Server: req.Server, Topic: req.Topic, Host: req.Host, Port: req.Port, + Username: req.Username, From: req.From, To: req.To, Secret: req.Secret, + }) + if err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) + return + } + h.listNotifyChannels(c) +} + +// testNotifyChannel 用已保存配置向指定渠道同步发送测试消息。 +// +// @Summary 向指定渠道发送测试消息 +// @Tags 设置 +// @Param type path string true "渠道类型(webhook/ntfy/bark/smtp)" +// @Success 204 "无内容" +// @Security BearerAuth +// @Router /api/v1/settings/notify-channels/{type}/test [post] +func (h *settingsHandler) testNotifyChannel(c *gin.Context) { + if err := h.notifier.SendChannelTest(c.Request.Context(), c.Param("type")); err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) + return + } + c.Status(http.StatusNoContent) +} + // getNotifyEvents 返回全部通知事件开关;键缺省(存量部署)视为开启。 // // @Summary 返回全部通知事件开关 diff --git a/internal/api/tenant.go b/internal/api/tenant.go index 30781b1..56fd9e7 100644 --- a/internal/api/tenant.go +++ b/internal/api/tenant.go @@ -67,12 +67,14 @@ func (h *ociConfigHandler) costs(c *gin.Context) { // ---- 租户审计日志 ---- -// getAuditEvents 实时查询租户 OCI 审计事件:hours 首查(缺省 24); -// start/end/page 为截断后的同窗续查;窗口越界或非法返回 400。 +// getAuditEvents 批式懒加载查询审计事件:cursor 为空自当前时刻首查, +// 非空从上次响应游标继续向更早回溯;limit 单批目标条数(缺省 100,上限 200)。 // -// @Summary 实时查询租户 OCI 审计事件:hours 首查 +// @Summary 批式懒加载查询租户 OCI 审计事件 // @Tags 租户 IAM // @Param id path int true "配置 ID" +// @Param cursor query string false "续查游标(上次响应原样带回)" +// @Param limit query int false "单批目标条数,缺省 100,上限 200" // @Success 200 {object} map[string]any // @Security BearerAuth // @Router /api/v1/oci-configs/{id}/audit-events [get] @@ -81,13 +83,10 @@ func (h *ociConfigHandler) getAuditEvents(c *gin.Context) { if !ok { return } - hours, _ := strconv.Atoi(c.DefaultQuery("hours", "24")) - q := service.AuditQuery{ - Region: c.Query("region"), Hours: hours, - Start: c.Query("start"), End: c.Query("end"), Page: c.Query("page"), - } + limit, _ := strconv.Atoi(c.Query("limit")) + q := service.AuditQuery{Region: c.Query("region"), Cursor: c.Query("cursor"), Limit: limit} result, err := h.svc.AuditEvents(c.Request.Context(), id, q) - if errors.Is(err, service.ErrInvalidAuditHours) || errors.Is(err, service.ErrInvalidAuditWindow) { + if errors.Is(err, service.ErrInvalidAuditCursor) { c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) return } @@ -143,9 +142,29 @@ type createTenantUserRequest struct { AddToAdminGroup bool `json:"addToAdminGroup"` } +// @Summary 租户身份域列表 +// @Tags 租户 IAM +// @Param id path int true "配置 ID" +// @Success 200 {object} map[string]any +// @Security BearerAuth +// @Router /api/v1/oci-configs/{id}/domains [get] +func (h *ociConfigHandler) listIdentityDomains(c *gin.Context) { + id, ok := pathID(c) + if !ok { + return + } + domains, err := h.svc.IdentityDomains(c.Request.Context(), id) + if err != nil { + respondError(c, err) + return + } + c.JSON(http.StatusOK, domains) +} + // @Summary 租户 IAM 用户列表 // @Tags 租户 IAM // @Param id path int true "配置 ID" +// @Param domainId query string false "身份域 OCID(缺省 Default 域)" // @Success 200 {object} map[string]any // @Security BearerAuth // @Router /api/v1/oci-configs/{id}/users [get] @@ -154,7 +173,7 @@ func (h *ociConfigHandler) listTenantUsers(c *gin.Context) { if !ok { return } - users, err := h.svc.TenantUsers(c.Request.Context(), id) + users, err := h.svc.TenantUsers(c.Request.Context(), id, c.Query("domainId")) if err != nil { respondError(c, err) return @@ -165,6 +184,7 @@ func (h *ociConfigHandler) listTenantUsers(c *gin.Context) { // @Summary 租户用户详情 // @Tags 租户 IAM // @Param id path int true "配置 ID" +// @Param domainId query string false "身份域 OCID(缺省 Default 域)" // @Param userId path string true "userId" // @Success 200 {object} map[string]any // @Security BearerAuth @@ -174,7 +194,7 @@ func (h *ociConfigHandler) getTenantUserDetail(c *gin.Context) { if !ok { return } - detail, err := h.svc.TenantUserDetail(c.Request.Context(), id, c.Param("userId")) + detail, err := h.svc.TenantUserDetail(c.Request.Context(), id, c.Query("domainId"), c.Param("userId")) if err != nil { respondError(c, err) return @@ -185,6 +205,7 @@ func (h *ociConfigHandler) getTenantUserDetail(c *gin.Context) { // @Summary 创建租户 IAM 用户 // @Tags 租户 IAM // @Param id path int true "配置 ID" +// @Param domainId query string false "身份域 OCID(缺省 Default 域)" // @Param body body createTenantUserRequest true "请求体" // @Success 201 {object} map[string]any // @Security BearerAuth @@ -199,7 +220,7 @@ func (h *ociConfigHandler) createTenantUser(c *gin.Context) { c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) return } - user, err := h.svc.CreateTenantUser(c.Request.Context(), id, oci.CreateTenantUserInput{ + user, err := h.svc.CreateTenantUser(c.Request.Context(), id, c.Query("domainId"), oci.CreateTenantUserInput{ Name: req.Name, Description: req.Description, Email: req.Email, @@ -229,6 +250,7 @@ type updateTenantUserRequest struct { // @Summary 更新租户 IAM 用户 // @Tags 租户 IAM // @Param id path int true "配置 ID" +// @Param domainId query string false "身份域 OCID(缺省 Default 域)" // @Param userId path string true "userId" // @Param body body updateTenantUserRequest true "请求体" // @Success 200 {object} map[string]any @@ -244,7 +266,7 @@ func (h *ociConfigHandler) updateTenantUser(c *gin.Context) { c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) return } - user, err := h.svc.UpdateTenantUser(c.Request.Context(), id, c.Param("userId"), oci.UpdateTenantUserInput{ + user, err := h.svc.UpdateTenantUser(c.Request.Context(), id, c.Query("domainId"), c.Param("userId"), oci.UpdateTenantUserInput{ Description: req.Description, Email: req.Email, GivenName: req.GivenName, @@ -262,6 +284,7 @@ func (h *ociConfigHandler) updateTenantUser(c *gin.Context) { // @Summary 删除租户 IAM 用户 // @Tags 租户 IAM // @Param id path int true "配置 ID" +// @Param domainId query string false "身份域 OCID(缺省 Default 域)" // @Param userId path string true "userId" // @Success 204 "无内容" // @Security BearerAuth @@ -271,7 +294,7 @@ func (h *ociConfigHandler) deleteTenantUser(c *gin.Context) { if !ok { return } - if err := h.svc.DeleteTenantUser(c.Request.Context(), id, c.Param("userId")); err != nil { + if err := h.svc.DeleteTenantUser(c.Request.Context(), id, c.Query("domainId"), c.Param("userId")); err != nil { respondError(c, err) return } @@ -281,6 +304,7 @@ func (h *ociConfigHandler) deleteTenantUser(c *gin.Context) { // @Summary 重置租户用户密码 // @Tags 租户 IAM // @Param id path int true "配置 ID" +// @Param domainId query string false "身份域 OCID(缺省 Default 域)" // @Param userId path string true "userId" // @Success 200 {object} map[string]any // @Security BearerAuth @@ -290,7 +314,7 @@ func (h *ociConfigHandler) resetTenantUserPassword(c *gin.Context) { if !ok { return } - password, err := h.svc.ResetTenantUserPassword(c.Request.Context(), id, c.Param("userId")) + password, err := h.svc.ResetTenantUserPassword(c.Request.Context(), id, c.Query("domainId"), c.Param("userId")) if err != nil { respondError(c, err) return @@ -301,6 +325,7 @@ func (h *ociConfigHandler) resetTenantUserPassword(c *gin.Context) { // @Summary 清除用户 MFA 设备 // @Tags 租户 IAM // @Param id path int true "配置 ID" +// @Param domainId query string false "身份域 OCID(缺省 Default 域)" // @Param userId path string true "userId" // @Success 200 {object} map[string]any // @Security BearerAuth @@ -310,7 +335,7 @@ func (h *ociConfigHandler) deleteTenantUserMfa(c *gin.Context) { if !ok { return } - deleted, err := h.svc.DeleteTenantUserMfa(c.Request.Context(), id, c.Param("userId")) + deleted, err := h.svc.DeleteTenantUserMfa(c.Request.Context(), id, c.Query("domainId"), c.Param("userId")) if err != nil { respondError(c, err) return @@ -344,6 +369,7 @@ func (h *ociConfigHandler) deleteTenantUserApiKeys(c *gin.Context) { // @Summary ---- 域通知收件人与密码策略 ---- // @Tags 租户 IAM // @Param id path int true "配置 ID" +// @Param domainId query string false "身份域 OCID(缺省 Default 域)" // @Success 200 {object} map[string]any // @Security BearerAuth // @Router /api/v1/oci-configs/{id}/notification-recipients [get] @@ -352,7 +378,7 @@ func (h *ociConfigHandler) getNotificationRecipients(c *gin.Context) { if !ok { return } - recipients, err := h.svc.NotificationRecipients(c.Request.Context(), id) + recipients, err := h.svc.NotificationRecipients(c.Request.Context(), id, c.Query("domainId")) if err != nil { respondError(c, err) return @@ -363,6 +389,7 @@ func (h *ociConfigHandler) getNotificationRecipients(c *gin.Context) { // @Summary 更新通知收件人 // @Tags 租户 IAM // @Param id path int true "配置 ID" +// @Param domainId query string false "身份域 OCID(缺省 Default 域)" // @Param body body object true "请求体(见接口说明)" // @Success 200 {object} map[string]any // @Security BearerAuth @@ -379,7 +406,7 @@ func (h *ociConfigHandler) updateNotificationRecipients(c *gin.Context) { c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) return } - recipients, err := h.svc.UpdateNotificationRecipients(c.Request.Context(), id, req.Recipients) + recipients, err := h.svc.UpdateNotificationRecipients(c.Request.Context(), id, c.Query("domainId"), req.Recipients) if err != nil { respondError(c, err) return @@ -390,6 +417,7 @@ func (h *ociConfigHandler) updateNotificationRecipients(c *gin.Context) { // @Summary 密码策略列表 // @Tags 租户 IAM // @Param id path int true "配置 ID" +// @Param domainId query string false "身份域 OCID(缺省 Default 域)" // @Success 200 {object} map[string]any // @Security BearerAuth // @Router /api/v1/oci-configs/{id}/password-policies [get] @@ -398,7 +426,7 @@ func (h *ociConfigHandler) listPasswordPolicies(c *gin.Context) { if !ok { return } - policies, err := h.svc.PasswordPolicies(c.Request.Context(), id) + policies, err := h.svc.PasswordPolicies(c.Request.Context(), id, c.Query("domainId")) if err != nil { respondError(c, err) return @@ -409,6 +437,7 @@ func (h *ociConfigHandler) listPasswordPolicies(c *gin.Context) { // @Summary 更新密码策略 // @Tags 租户 IAM // @Param id path int true "配置 ID" +// @Param domainId query string false "身份域 OCID(缺省 Default 域)" // @Param policyId path string true "policyId" // @Param body body object true "请求体(见接口说明)" // @Success 200 {object} map[string]any @@ -428,7 +457,7 @@ func (h *ociConfigHandler) updatePasswordPolicy(c *gin.Context) { c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) return } - policy, err := h.svc.UpdatePasswordPolicy(c.Request.Context(), id, c.Param("policyId"), oci.UpdatePasswordPolicyInput{ + policy, err := h.svc.UpdatePasswordPolicy(c.Request.Context(), id, c.Query("domainId"), c.Param("policyId"), oci.UpdatePasswordPolicyInput{ PasswordExpiresAfter: req.PasswordExpiresAfter, MinLength: req.MinLength, NumPasswordsInHistory: req.NumPasswordsInHistory, @@ -443,6 +472,7 @@ func (h *ociConfigHandler) updatePasswordPolicy(c *gin.Context) { // @Summary 身份域设置 // @Tags 租户 IAM // @Param id path int true "配置 ID" +// @Param domainId query string false "身份域 OCID(缺省 Default 域)" // @Success 200 {object} map[string]any // @Security BearerAuth // @Router /api/v1/oci-configs/{id}/identity-settings [get] @@ -451,7 +481,7 @@ func (h *ociConfigHandler) getIdentitySetting(c *gin.Context) { if !ok { return } - setting, err := h.svc.IdentitySetting(c.Request.Context(), id) + setting, err := h.svc.IdentitySetting(c.Request.Context(), id, c.Query("domainId")) if err != nil { respondError(c, err) return @@ -462,6 +492,7 @@ func (h *ociConfigHandler) getIdentitySetting(c *gin.Context) { // @Summary 更新身份域设置 // @Tags 租户 IAM // @Param id path int true "配置 ID" +// @Param domainId query string false "身份域 OCID(缺省 Default 域)" // @Param body body object true "请求体(见接口说明)" // @Success 200 {object} map[string]any // @Security BearerAuth @@ -478,7 +509,7 @@ func (h *ociConfigHandler) updateIdentitySetting(c *gin.Context) { c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()}) return } - setting, err := h.svc.UpdateIdentitySetting(c.Request.Context(), id, *req.PrimaryEmailRequired) + setting, err := h.svc.UpdateIdentitySetting(c.Request.Context(), id, c.Query("domainId"), *req.PrimaryEmailRequired) if err != nil { respondError(c, err) return diff --git a/internal/api/webconsole.go b/internal/api/webconsole.go index 486e668..270d337 100644 --- a/internal/api/webconsole.go +++ b/internal/api/webconsole.go @@ -79,7 +79,7 @@ var wsUpgrader = websocket.Upgrader{ // @Success 101 "升级为 WebSocket" // @Router /api/v1/console-sessions/{sessionId}/ws [get] func (h *consoleHandler) ws(c *gin.Context) { - if _, err := h.auth.ParseToken(c.Query("token")); err != nil { + if _, err := h.auth.ParseToken(c.Request.Context(), c.Query("token")); err != nil { c.JSON(http.StatusUnauthorized, gin.H{"error": "invalid or expired token"}) return } diff --git a/internal/database/database.go b/internal/database/database.go index feab8b0..9b351ee 100644 --- a/internal/database/database.go +++ b/internal/database/database.go @@ -2,6 +2,9 @@ package database import ( "fmt" + "log" + "os" + "time" "github.com/glebarez/sqlite" "gorm.io/driver/mysql" @@ -12,6 +15,18 @@ import ( "oci-portal/internal/model" ) +// newDBLogger 构造脱敏的 GORM 日志器(S-08):SQL 一律参数化输出, +// 实参(密码哈希 / token / 任务 payload)不落日志;业务正常路径的 +// record-not-found 不打 SQL。 +func newDBLogger() logger.Interface { + return logger.New(log.New(os.Stdout, "\r\n", log.LstdFlags), logger.Config{ + SlowThreshold: 200 * time.Millisecond, + LogLevel: logger.Warn, + IgnoreRecordNotFoundError: true, + ParameterizedQueries: true, + }) +} + // Open 按驱动打开数据库并自动迁移全部模型。 // driver 取值 sqlite(默认)/mysql/postgres;sqlite 用 path,其余用 dsn。 func Open(driver, dsn, path string) (*gorm.DB, error) { @@ -20,7 +35,7 @@ func Open(driver, dsn, path string) (*gorm.DB, error) { return nil, err } db, err := gorm.Open(dialector, &gorm.Config{ - Logger: logger.Default.LogMode(logger.Warn), + Logger: newDBLogger(), }) if err != nil { return nil, fmt.Errorf("open %s database: %w", dialector.Name(), err) @@ -52,7 +67,7 @@ func autoMigrate(db *gorm.DB) error { &model.User{}, &model.UserIdentity{}, &model.OciConfig{}, &model.Task{}, &model.TaskLog{}, &model.CheckSnapshot{}, &model.CostSnapshot{}, &model.RegionCache{}, &model.CompartmentCache{}, &model.Setting{}, - &model.SystemLog{}, &model.LogEvent{}, &model.Proxy{}, + &model.SystemLog{}, &model.LogEvent{}, &model.AlertRule{}, &model.AlertRuleHit{}, &model.Proxy{}, &model.AiKey{}, &model.AiChannel{}, &model.AiModelCache{}, &model.AiCallLog{}, &model.AiContentLog{}, ) diff --git a/internal/model/models.go b/internal/model/models.go index df34735..42bd36c 100644 --- a/internal/model/models.go +++ b/internal/model/models.go @@ -10,11 +10,12 @@ const ( AccountTypeUnknown = "unknown" ) -// 测活状态取值。 +// 测活状态取值;suspended 来自账户能力接口的云端暂停标记(区别于失联)。 const ( - AliveStatusAlive = "alive" - AliveStatusDead = "dead" - AliveStatusUnknown = "unknown" + AliveStatusAlive = "alive" + AliveStatusDead = "dead" + AliveStatusSuspended = "suspended" + AliveStatusUnknown = "unknown" ) // Setting 是系统级键值配置;敏感值(如 telegram_bot_token)以 AES-GCM 密文存储。 @@ -64,9 +65,12 @@ type User struct { Username string `gorm:"uniqueIndex;size:64" json:"username"` PasswordHash string `json:"-"` // TOTP 共享密钥 AES-GCM 密文;空串表示未启用两步验证 - TotpSecretEnc string `json:"-"` - CreatedAt time.Time `json:"createdAt"` - UpdatedAt time.Time `json:"updatedAt"` + TotpSecretEnc string `json:"-"` + // TokenVersion 是 JWT 版本号:凭据/TOTP/外部身份/登录策略变更或撤销会话时 + // 原子递增,旧版本令牌随即失效(存量令牌无 ver 视为 0,与零值兼容) + TokenVersion uint `json:"-"` + CreatedAt time.Time `json:"createdAt"` + UpdatedAt time.Time `json:"updatedAt"` } // UserIdentity 是管理员绑定的外部登录身份(OIDC / GitHub); @@ -207,6 +211,31 @@ type LogEvent struct { ReceivedAt time.Time `json:"receivedAt"` } +// AlertRule 是回传事件的自定义告警规则;条件间 AND 关系,空条件视为任意。 +type AlertRule struct { + ID uint `gorm:"primaryKey" json:"id"` + Name string `gorm:"size:64" json:"name"` + Enabled bool `json:"enabled"` + OciConfigID uint `json:"ociConfigId"` // 0=全部租户 + EventTypes string `gorm:"size:512" json:"eventTypes"` // 逗号分隔事件短名,空=全部 + SourceIPs string `gorm:"size:512" json:"sourceIps"` // 逗号分隔 IP/CIDR,空=任意 + // in:来源命中列表才告警(默认);notin:不在列表才告警(白名单场景) + SourceIPMode string `gorm:"size:8" json:"sourceIpMode"` + ResourceMatch string `gorm:"size:128" json:"resourceMatch"` // 资源名子串,空=任意 + Threshold int `json:"threshold"` // 触发阈值,默认 1(即时) + WindowMinutes int `json:"windowMinutes"` // 聚合窗口,Threshold>1 时必填 + CreatedAt time.Time `json:"createdAt"` + UpdatedAt time.Time `json:"updatedAt"` +} + +// AlertRuleHit 记录规则命中,供阈值窗口计数;随周期清理删除过期行。 +type AlertRuleHit struct { + ID uint `gorm:"primaryKey" json:"id"` + RuleID uint `gorm:"index" json:"ruleId"` + LogEventID uint `json:"logEventId"` + HitAt time.Time `gorm:"index" json:"hitAt"` +} + // RegionCache 是开启多区域支持的配置缓存的订阅区域(每配置多行,整组覆盖)。 // Status 非 READY(新订阅进行中)时读取接口会实时刷新,直到全部 READY。 type RegionCache struct { diff --git a/internal/oci/account.go b/internal/oci/account.go index 3194e22..c5c97bb 100644 --- a/internal/oci/account.go +++ b/internal/oci/account.go @@ -1,9 +1,16 @@ package oci import ( + "context" + "encoding/json" + "fmt" + "io" + "net/http" "strings" "time" + "github.com/oracle/oci-go-sdk/v65/common" + "oci-portal/internal/model" ) @@ -19,6 +26,89 @@ type AccountProfile struct { PromotionExpires *time.Time } +// AccountCapabilities 是账户能力接口的关键字段。来源为 Limits 服务 +// 20230301 版 `/compartments/{tenancy}/capabilities?accountInfoOnly=true` +// (与控制台同源;oci-go-sdk v65 未收录,自签名裸调)。 +// 该接口对暂停/已终止租户仍可访问,借此把「暂停」从「失联」中区分出来。 +type AccountCapabilities struct { + Suspended bool `json:"suspended"` + AccountStatus string `json:"accountStatus"` + FreeTierEnabled bool `json:"freeTierEnabled"` + FreeTierOnly bool `json:"freeTierOnly"` + PromotionStatus string `json:"promotionStatus"` + ProgramType string `json:"programType"` + HasSaasSubscription bool `json:"hasSaasSubscription"` + Deletable bool `json:"deletable"` +} + +// capabilitiesEnvelope 对应接口原始响应:顶层与 accountInfo 各有一个 +// suspended,任一为 true 即视为暂停。 +type capabilitiesEnvelope struct { + Capabilities struct { + Suspended bool `json:"suspended"` + AccountInfo struct { + FreeTierEnabled bool `json:"freeTierEnabled"` + FreeTierOnly bool `json:"freeTierOnly"` + PromotionStatus string `json:"promotionStatus"` + Suspended bool `json:"suspended"` + ProgramType string `json:"programType"` + HasSaasSubscription bool `json:"hasSaasSubscription"` + AccountStatus string `json:"accountStatus"` + Deletable bool `json:"deletable"` + } `json:"accountInfo"` + } `json:"capabilities"` +} + +// parseAccountCapabilities 把原始响应体压平为关键字段视图。 +func parseAccountCapabilities(body []byte) (AccountCapabilities, error) { + var env capabilitiesEnvelope + if err := json.Unmarshal(body, &env); err != nil { + return AccountCapabilities{}, fmt.Errorf("decode account capabilities: %w", err) + } + info := env.Capabilities.AccountInfo + return AccountCapabilities{ + Suspended: env.Capabilities.Suspended || info.Suspended, + AccountStatus: info.AccountStatus, + FreeTierEnabled: info.FreeTierEnabled, + FreeTierOnly: info.FreeTierOnly, + PromotionStatus: info.PromotionStatus, + ProgramType: info.ProgramType, + HasSaasSubscription: info.HasSaasSubscription, + Deletable: info.Deletable, + }, nil +} + +// GetAccountCapabilities 实现 Client:查询账户能力(暂停/账户状态/免费层)。 +// region 应传 home region;请求带 API Key 签名,并沿用租户出站代理。 +func (c *RealClient) GetAccountCapabilities(ctx context.Context, cred Credentials, region string) (AccountCapabilities, error) { + url := fmt.Sprintf("https://limits.%s.oci.oraclecloud.com/20230301/compartments/%s/capabilities?accountInfoOnly=true", + normalizeRegion(region), cred.TenancyOCID) + req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil) + if err != nil { + return AccountCapabilities{}, fmt.Errorf("new capabilities request: %w", err) + } + if err := common.DefaultRequestSigner(provider(cred)).Sign(req); err != nil { + return AccountCapabilities{}, fmt.Errorf("sign capabilities request: %w", err) + } + hc := proxyHTTPClient(cred.Proxy) + if hc == nil { + hc = &http.Client{Timeout: proxyClientTimeout} + } + resp, err := hc.Do(req) + if err != nil { + return AccountCapabilities{}, fmt.Errorf("get account capabilities: %w", err) + } + defer resp.Body.Close() + body, err := io.ReadAll(io.LimitReader(resp.Body, 1<<20)) + if err != nil { + return AccountCapabilities{}, fmt.Errorf("read account capabilities: %w", err) + } + if resp.StatusCode != http.StatusOK { + return AccountCapabilities{}, fmt.Errorf("get account capabilities: status %d", resp.StatusCode) + } + return parseAccountCapabilities(body) +} + // ClassifyAccount 按订阅字段判定账户类别: // - paymentModel 非空且不是 FREE_TRIAL:付费; // - paymentModel 为 FREE_TRIAL 时,promotion ACTIVE 或 tier FREE_AND_TRIAL diff --git a/internal/oci/account_test.go b/internal/oci/account_test.go index 98c5d54..84a7f8f 100644 --- a/internal/oci/account_test.go +++ b/internal/oci/account_test.go @@ -36,3 +36,53 @@ func TestClassifyAccount(t *testing.T) { }) } } + +// TestParseAccountCapabilities 用控制台同源接口的真实响应样例驱动。 +func TestParseAccountCapabilities(t *testing.T) { + cases := []struct { + name string + body string + wantSusp bool + wantStatus string + wantErr bool + }{ + { + name: "已终止租户(顶层与 accountInfo 双 suspended)", + body: `{"compartmentId":"ocid1.tenancy.oc1..x","capabilities":{"suspended":true,` + + `"accountInfo":{"freeTierEnabled":false,"freeTierOnly":false,"promotionStatus":"none",` + + `"suspended":true,"programType":"default","limitsProvisioned":true,"intentToPay":false,` + + `"hasSaasSubscription":true,"accountStatus":"terminated","accountFlags":0,"deletable":true,` + + `"limitIncrease":false,"orgProperties":"0"}}}`, + wantSusp: true, + wantStatus: "terminated", + }, + { + name: "正常租户", + body: `{"capabilities":{"suspended":false,"accountInfo":{"accountStatus":"active","freeTierEnabled":true}}}`, + wantSusp: false, wantStatus: "active", + }, + { + name: "仅 accountInfo 标记 suspended 也算暂停", + body: `{"capabilities":{"suspended":false,"accountInfo":{"suspended":true}}}`, + wantSusp: true, + }, + {name: "非 JSON 报错", body: ``, wantErr: true}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + got, err := parseAccountCapabilities([]byte(tc.body)) + if tc.wantErr { + if err == nil { + t.Fatal("expected error") + } + return + } + if err != nil { + t.Fatalf("parse: %v", err) + } + if got.Suspended != tc.wantSusp || got.AccountStatus != tc.wantStatus { + t.Fatalf("got %+v, want suspended=%v status=%q", got, tc.wantSusp, tc.wantStatus) + } + }) + } +} diff --git a/internal/oci/audit.go b/internal/oci/audit.go index fb30f50..266f4bd 100644 --- a/internal/oci/audit.go +++ b/internal/oci/audit.go @@ -93,6 +93,11 @@ func (c *RealClient) ListAuditEvents(ctx context.Context, cred Credentials, regi // appendAuditEvents 过滤噪声后追加一页事件;原始事件只对保留条目序列化。 func appendAuditEvents(result *AuditEventsResult, items []audit.AuditEvent) { + result.Items = appendKeptAuditEvents(result.Items, items) +} + +// appendKeptAuditEvents 是过滤追加的通用形态,窗口式与批式查询共用。 +func appendKeptAuditEvents(dst []AuditEvent, items []audit.AuditEvent) []AuditEvent { for _, ev := range items { out := toAuditEvent(ev) if !keepAuditEvent(out) { @@ -101,8 +106,115 @@ func appendAuditEvents(result *AuditEventsResult, items []audit.AuditEvent) { if raw, mErr := json.Marshal(ev); mErr == nil { out.Raw = raw } - result.Items = append(result.Items, out) + dst = append(dst, out) } + return dst +} + +// ---- 批式懒加载查询:分窗回溯 + 游标续查 ---- + +// 批式查询参数:单批 OCI 翻页预算沿用 maxAuditPages;首窗 24h, +// 连续空窗倍增(上限 30 天)加速跨越闲置期;回溯下限为事件保留期 365 天。 +const ( + auditWindowHours = 24 + auditWindowMaxHours = 720 + auditRetentionDays = 365 +) + +// AuditCursor 是批式查询的续查位置:当前时间窗、窗内 OCI 翻页游标 +// 与当前窗宽(小时,空窗倍增的记忆)。序列化为不透明 cursor 由 service 层负责。 +type AuditCursor struct { + Start time.Time `json:"s"` + End time.Time `json:"e"` + Page string `json:"p,omitempty"` + WindowHours int `json:"w"` +} + +// NewAuditCursor 构造首查游标:自 now 起回溯第一个 24h 窗。 +func NewAuditCursor(now time.Time) AuditCursor { + end := now.UTC().Truncate(time.Minute) + return AuditCursor{Start: end.Add(-auditWindowHours * time.Hour), End: end, WindowHours: auditWindowHours} +} + +// advance 推进到紧邻更早的窗;empty 表示刚结束的窗无保留事件,窗宽倍增, +// 否则重置 24h。done 为 true 表示已越过保留期尽头。 +func (cur AuditCursor) advance(now time.Time, empty bool) (AuditCursor, bool) { + w := cur.WindowHours + if w <= 0 { + w = auditWindowHours + } + if empty { + if w *= 2; w > auditWindowMaxHours { + w = auditWindowMaxHours + } + } else { + w = auditWindowHours + } + end := cur.Start + if end.Before(now.UTC().AddDate(0, 0, -auditRetentionDays)) { + return cur, true + } + return AuditCursor{Start: end.Add(-time.Duration(w) * time.Hour), End: end, WindowHours: w}, false +} + +// AuditBatchResult 是一批懒加载结果;Cursor 为 nil 且 Exhausted 为 true +// 表示已回溯到保留期尽头,无更早数据。 +type AuditBatchResult struct { + Items []AuditEvent + Cursor *AuditCursor + Exhausted bool +} + +// ListAuditEventsBatch 实现 Client:从 cur 位置向更早方向收集约 limit 条 +// 保留事件;单批最多消费 maxAuditPages 页 OCI 调用,不足额也按预算返回, +// 由前端按需续查。窗口不重叠 + 窗内游标续翻保证跨批不重不漏。 +func (c *RealClient) ListAuditEventsBatch(ctx context.Context, cred Credentials, region string, cur AuditCursor, limit int) (AuditBatchResult, error) { + ac, err := c.auditClient(cred, region) + if err != nil { + return AuditBatchResult{}, err + } + res := AuditBatchResult{Items: []AuditEvent{}} + windowHasKept := false + for budget := maxAuditPages; budget > 0 && len(res.Items) < limit; budget-- { + items, next, err := listAuditPage(ctx, ac, cred.TenancyOCID, cur) + if err != nil { + return AuditBatchResult{}, err + } + before := len(res.Items) + res.Items = appendKeptAuditEvents(res.Items, items) + windowHasKept = windowHasKept || len(res.Items) > before + if next != "" { + cur.Page = next + continue + } + nextCur, done := cur.advance(time.Now(), !windowHasKept) + if done { + res.Exhausted = true + sortAuditEvents(res.Items) + return res, nil + } + cur, windowHasKept = nextCur, false + } + sortAuditEvents(res.Items) + res.Cursor = &cur + return res, nil +} + +// listAuditPage 拉取当前游标位置的一页原始事件。 +func listAuditPage(ctx context.Context, ac audit.AuditClient, tenancyOCID string, cur AuditCursor) ([]audit.AuditEvent, string, error) { + req := audit.ListEventsRequest{ + CompartmentId: &tenancyOCID, + StartTime: &common.SDKTime{Time: cur.Start.UTC().Truncate(time.Minute)}, + EndTime: &common.SDKTime{Time: cur.End.UTC().Truncate(time.Minute)}, + } + if cur.Page != "" { + req.Page = &cur.Page + } + resp, err := ac.ListEvents(ctx, req) + if err != nil { + return nil, "", fmt.Errorf("list audit events: %w", err) + } + return resp.Items, deref(resp.OpcNextPage), nil } // auditInternalCIDRs 是 OCI 服务内部互调的发起方网段(RFC1918 + CGNAT)。 diff --git a/internal/oci/audit_test.go b/internal/oci/audit_test.go index d986f78..3b23a93 100644 --- a/internal/oci/audit_test.go +++ b/internal/oci/audit_test.go @@ -142,3 +142,48 @@ func TestKeepAuditEvent(t *testing.T) { }) } } + +func TestAuditCursorAdvance(t *testing.T) { + now := time.Date(2026, 7, 10, 12, 0, 0, 0, time.UTC) + base := AuditCursor{ + Start: now.Add(-24 * time.Hour), + End: now, + WindowHours: 24, + } + cases := []struct { + name string + cur AuditCursor + empty bool + wantHours int + wantDone bool + }{ + {"有事件重置 24h 窗", AuditCursor{Start: base.Start, End: base.End, WindowHours: 96}, false, 24, false}, + {"空窗倍增", base, true, 48, false}, + {"倍增封顶 720h", AuditCursor{Start: base.Start, End: base.End, WindowHours: 512}, true, 720, false}, + {"窗宽缺省按 24h 起算", AuditCursor{Start: base.Start, End: base.End}, true, 48, false}, + {"越过保留期即尽头", AuditCursor{Start: now.AddDate(0, 0, -366), End: now.AddDate(0, 0, -365), WindowHours: 24}, false, 0, true}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + next, done := tc.cur.advance(now, tc.empty) + if done != tc.wantDone { + t.Fatalf("done = %v, want %v", done, tc.wantDone) + } + if done { + return + } + if next.WindowHours != tc.wantHours { + t.Fatalf("WindowHours = %d, want %d", next.WindowHours, tc.wantHours) + } + if !next.End.Equal(tc.cur.Start) { + t.Fatalf("新窗 End = %v, 应紧邻上窗 Start %v", next.End, tc.cur.Start) + } + if got := next.End.Sub(next.Start); got != time.Duration(tc.wantHours)*time.Hour { + t.Fatalf("窗宽 = %v, want %dh", got, tc.wantHours) + } + if next.Page != "" { + t.Fatalf("新窗应清空窗内游标, got %q", next.Page) + } + }) + } +} diff --git a/internal/oci/cached.go b/internal/oci/cached.go index 1d21e91..627b86d 100644 --- a/internal/oci/cached.go +++ b/internal/oci/cached.go @@ -39,10 +39,18 @@ func ckey(cred Credentials, parts ...string) string { // bust 写操作成功后失效该租户全部读缓存。 func (c *CachedClient) bust(err error, cred Credentials) { if err == nil { - c.store.DeletePrefix(cred.TenancyOCID + "|") + c.InvalidateTenancy(cred.TenancyOCID) } } +// InvalidateTenancy 主动失效指定 tenancy 的全部 OCI 读缓存。 +func (c *CachedClient) InvalidateTenancy(tenancyOCID string) { + if tenancyOCID == "" { + return + } + c.store.DeletePrefix(tenancyOCID + "|") +} + // cachedList 缓存切片结果;命中返回浅拷贝,防调用方排序共享底层数组。 func cachedList[T any](c *CachedClient, key string, ttl time.Duration, fn func() ([]T, error)) ([]T, error) { v, err := cache.Do(c.store, key, ttl, fn) @@ -120,6 +128,13 @@ func (c *CachedClient) ListVolumeAttachments(ctx context.Context, cred Credentia }) } +// ListIdentityDomains 域列表准静态且是所有按域操作的 URL 解析源,长 TTL 缓存。 +func (c *CachedClient) ListIdentityDomains(ctx context.Context, cred Credentials, region string) ([]IdentityDomain, error) { + return cachedList(c, ckey(cred, "iddomains", region), cacheTTLStatic, func() ([]IdentityDomain, error) { + return c.Client.ListIdentityDomains(ctx, cred, region) + }) +} + // ---- 写:直通,成功后按租户失效 ---- func (c *CachedClient) LaunchInstance(ctx context.Context, cred Credentials, in CreateInstanceInput) (Instance, error) { diff --git a/internal/oci/cached_test.go b/internal/oci/cached_test.go index 1d0bf6c..119e4bb 100644 --- a/internal/oci/cached_test.go +++ b/internal/oci/cached_test.go @@ -86,3 +86,18 @@ func TestCachedClientReturnsClone(t *testing.T) { t.Errorf("缓存底层数组被调用方污染: second[0]=%s, want i-1", second[0].ID) } } + +func TestCachedClientInvalidateTenancy(t *testing.T) { + inner := &countingClient{} + c := NewCachedClient(inner) + ctx := context.Background() + + _, _ = c.ListInstances(ctx, testCred("t1"), "r1") + _, _ = c.ListInstances(ctx, testCred("t2"), "r1") + c.InvalidateTenancy("t1") + _, _ = c.ListInstances(ctx, testCred("t1"), "r1") + _, _ = c.ListInstances(ctx, testCred("t2"), "r1") + if inner.instCalls != 3 { + t.Fatalf("invalidate tenancy calls = %d, want 3", inner.instCalls) + } +} diff --git a/internal/oci/client.go b/internal/oci/client.go index de1070d..efece79 100644 --- a/internal/oci/client.go +++ b/internal/oci/client.go @@ -122,32 +122,41 @@ type Client interface { SummarizeCosts(ctx context.Context, cred Credentials, q CostQuery) ([]CostItem, error) // ListAuditEvents 实时查询租户根 compartment 的审计事件(最多翻 maxAuditPages 页, // 到限截断并回传续查游标);page 传上次结果的 NextPage 可从断点续翻。 + // 批式懒加载走 ListAuditEventsBatch,本方法保留给详情小窗反查。 ListAuditEvents(ctx context.Context, cred Credentials, region string, start, end time.Time, page string) (AuditEventsResult, error) - // 租户用户管理:经典 IAM 为主,新式租户自动回退 Identity Domains。 - ListTenantUsers(ctx context.Context, cred Credentials) ([]TenantUser, error) - GetTenantUserDetail(ctx context.Context, cred Credentials, homeRegion, userID string) (TenantUserDetail, error) - CreateTenantUser(ctx context.Context, cred Credentials, homeRegion string, in CreateTenantUserInput) (TenantUser, error) - UpdateTenantUser(ctx context.Context, cred Credentials, homeRegion, userID string, in UpdateTenantUserInput) (TenantUser, error) - DeleteTenantUser(ctx context.Context, cred Credentials, homeRegion, userID string) error - ResetTenantUserPassword(ctx context.Context, cred Credentials, homeRegion, userID string) (string, error) - DeleteTenantUserMfaDevices(ctx context.Context, cred Credentials, homeRegion, userID string) (int, error) + // ListAuditEventsBatch 从游标位置向更早方向收集约 limit 条保留事件, + // 分窗回溯(空窗倍增),跨批不重不漏;到 365 天保留期尽头置 Exhausted。 + ListAuditEventsBatch(ctx context.Context, cred Credentials, region string, cur AuditCursor, limit int) (AuditBatchResult, error) + // 身份域:租户 ACTIVE 域枚举(多域租户选择器数据源)。 + ListIdentityDomains(ctx context.Context, cred Credentials, region string) ([]IdentityDomain, error) + // 账户能力:暂停标记/账户状态/免费层信息(对暂停租户仍可访问,测活据此判「暂停」)。 + GetAccountCapabilities(ctx context.Context, cred Credentials, region string) (AccountCapabilities, error) + // 租户用户管理:domainID 非空按域走 Identity Domains SCIM, + // 为空保持经典 IAM 为主、新式租户自动回退的旧行为。 + ListTenantUsers(ctx context.Context, cred Credentials, homeRegion, domainID string) ([]TenantUser, error) + GetTenantUserDetail(ctx context.Context, cred Credentials, homeRegion, domainID, userID string) (TenantUserDetail, error) + CreateTenantUser(ctx context.Context, cred Credentials, homeRegion, domainID string, in CreateTenantUserInput) (TenantUser, error) + UpdateTenantUser(ctx context.Context, cred Credentials, homeRegion, domainID, userID string, in UpdateTenantUserInput) (TenantUser, error) + DeleteTenantUser(ctx context.Context, cred Credentials, homeRegion, domainID, userID string) error + ResetTenantUserPassword(ctx context.Context, cred Credentials, homeRegion, domainID, userID string) (string, error) + DeleteTenantUserMfaDevices(ctx context.Context, cred Credentials, homeRegion, domainID, userID string) (int, error) DeleteTenantUserApiKeys(ctx context.Context, cred Credentials, homeRegion, userID string, includeCurrent bool) (int, error) // 域设置:通知收件人、密码策略、身份设置。 - GetNotificationRecipients(ctx context.Context, cred Credentials, region string) (NotificationRecipients, error) - UpdateNotificationRecipients(ctx context.Context, cred Credentials, region string, emails []string) (NotificationRecipients, error) - ListPasswordPolicies(ctx context.Context, cred Credentials, region string) ([]PasswordPolicyInfo, error) - UpdatePasswordPolicy(ctx context.Context, cred Credentials, region, policyID string, in UpdatePasswordPolicyInput) (PasswordPolicyInfo, error) - GetIdentitySetting(ctx context.Context, cred Credentials, region string) (IdentitySettingInfo, error) - UpdateIdentitySetting(ctx context.Context, cred Credentials, region string, primaryEmailRequired bool) (IdentitySettingInfo, error) + GetNotificationRecipients(ctx context.Context, cred Credentials, region, domainID string) (NotificationRecipients, error) + UpdateNotificationRecipients(ctx context.Context, cred Credentials, region, domainID string, emails []string) (NotificationRecipients, error) + ListPasswordPolicies(ctx context.Context, cred Credentials, region, domainID string) ([]PasswordPolicyInfo, error) + UpdatePasswordPolicy(ctx context.Context, cred Credentials, region, domainID, policyID string, in UpdatePasswordPolicyInput) (PasswordPolicyInfo, error) + GetIdentitySetting(ctx context.Context, cred Credentials, region, domainID string) (IdentitySettingInfo, error) + UpdateIdentitySetting(ctx context.Context, cred Credentials, region, domainID string, primaryEmailRequired bool) (IdentitySettingInfo, error) // Federation:SAML IdP 管理、域 SP 元数据、sign-on 免 MFA 规则。 - ListIdentityProviders(ctx context.Context, cred Credentials, region string) ([]IdentityProviderInfo, error) - CreateSamlIdentityProvider(ctx context.Context, cred Credentials, region string, in CreateIdpInput) (IdentityProviderInfo, error) - SetIdentityProviderEnabled(ctx context.Context, cred Credentials, region, idpID string, enabled bool) (IdentityProviderInfo, error) - DeleteIdentityProvider(ctx context.Context, cred Credentials, region, idpID string) error - DownloadDomainSamlMetadata(ctx context.Context, cred Credentials, region string) ([]byte, error) - ListConsoleSignOnRules(ctx context.Context, cred Credentials, region string) ([]SignOnRuleInfo, error) - CreateMfaExemptionRule(ctx context.Context, cred Credentials, region, idpID, ruleName string) (SignOnRuleInfo, error) - DeleteMfaExemptionRule(ctx context.Context, cred Credentials, region, ruleID string) error + ListIdentityProviders(ctx context.Context, cred Credentials, region, domainID string) ([]IdentityProviderInfo, error) + CreateSamlIdentityProvider(ctx context.Context, cred Credentials, region, domainID string, in CreateIdpInput) (IdentityProviderInfo, error) + SetIdentityProviderEnabled(ctx context.Context, cred Credentials, region, domainID, idpID string, enabled bool) (IdentityProviderInfo, error) + DeleteIdentityProvider(ctx context.Context, cred Credentials, region, domainID, idpID string) error + DownloadDomainSamlMetadata(ctx context.Context, cred Credentials, region, domainID string) ([]byte, error) + ListConsoleSignOnRules(ctx context.Context, cred Credentials, region, domainID string) ([]SignOnRuleInfo, error) + CreateMfaExemptionRule(ctx context.Context, cred Credentials, region, domainID, idpID, ruleName string) (SignOnRuleInfo, error) + DeleteMfaExemptionRule(ctx context.Context, cred Credentials, region, domainID, ruleID string) error // 日志回传链路(方案A):Topic/CUSTOM_HTTPS 订阅/IAM Policy/Service Connector // 的幂等创建、状态聚合与销毁;IAM 写操作发往 homeRegion。 EnsureRelayTopic(ctx context.Context, cred Credentials) (RelayResource, error) diff --git a/internal/oci/domainsettings.go b/internal/oci/domainsettings.go index 64b40ac..3970aa8 100644 --- a/internal/oci/domainsettings.go +++ b/internal/oci/domainsettings.go @@ -47,26 +47,82 @@ type IdentitySettingInfo struct { PrimaryEmailRequired bool `json:"primaryEmailRequired"` } -// domainURL 解析租户 Default Identity Domain 的 endpoint URL。 -func (c *RealClient) domainURL(ctx context.Context, cred Credentials, region string) (string, error) { +// IdentityDomain 是租户下一个 ACTIVE 身份域的列表视图; +// URL 仅供后端解析 SCIM 端点,不下发给前端(防面板被引导签名任意地址)。 +type IdentityDomain struct { + ID string `json:"id"` + DisplayName string `json:"displayName"` + URL string `json:"-"` + HomeRegion string `json:"homeRegion"` + Type string `json:"type"` + LicenseType string `json:"licenseType"` +} + +// ListIdentityDomains 实现 Client:列出租户全部 ACTIVE 身份域。 +// 无域老租户返回空列表(非错误),调用方据此回退经典 IAM 路径。 +func (c *RealClient) ListIdentityDomains(ctx context.Context, cred Credentials, region string) ([]IdentityDomain, error) { + items, err := c.listDomainSummaries(ctx, cred, region) + if err != nil { + return nil, err + } + out := make([]IdentityDomain, 0, len(items)) + for _, d := range items { + if d.LifecycleState != identity.DomainLifecycleStateActive { + continue + } + out = append(out, IdentityDomain{ + ID: deref(d.Id), + DisplayName: deref(d.DisplayName), + URL: deref(d.Url), + HomeRegion: deref(d.HomeRegion), + Type: string(d.Type), + LicenseType: deref(d.LicenseType), + }) + } + return out, nil +} + +// listDomainSummaries 拉取租户全部域(数量极少,仍按游标翻全以防万一)。 +func (c *RealClient) listDomainSummaries(ctx context.Context, cred Credentials, region string) ([]identity.DomainSummary, error) { ic, err := c.identityClientAt(cred, region) + if err != nil { + return nil, err + } + var items []identity.DomainSummary + var page *string + for { + resp, err := ic.ListDomains(ctx, identity.ListDomainsRequest{CompartmentId: &cred.TenancyOCID, Page: page}) + if err != nil { + return nil, fmt.Errorf("list identity domains: %w", err) + } + items = append(items, resp.Items...) + if resp.OpcNextPage == nil { + return items, nil + } + page = resp.OpcNextPage + } +} + +// domainURL 解析身份域的 endpoint URL:domainID 非空按 OCID 精确匹配 +// (URL 只能来自租户自己的域列表,不信任外部输入),为空回退 Default 优先。 +func (c *RealClient) domainURL(ctx context.Context, cred Credentials, region, domainID string) (string, error) { + items, err := c.listDomainSummaries(ctx, cred, region) if err != nil { return "", err } - resp, err := ic.ListDomains(ctx, identity.ListDomainsRequest{CompartmentId: &cred.TenancyOCID}) - if err != nil { - return "", fmt.Errorf("list identity domains: %w", err) - } - url := pickDomainURL(resp.Items) + url := pickDomainURL(items, domainID) if url == "" { + if domainID != "" { + return "", fmt.Errorf("identity domain %s not found or inactive", domainID) + } return "", fmt.Errorf("no active identity domain found") } return url, nil } -// domainsClient 解析租户的 Default Identity Domain 并构造 SCIM 客户端。 -func (c *RealClient) domainsClient(ctx context.Context, cred Credentials, region string) (identitydomains.IdentityDomainsClient, error) { - url, err := c.domainURL(ctx, cred, region) +// domainsClient 解析身份域(domainID 为空取 Default)并构造 SCIM 客户端。 +func (c *RealClient) domainsClient(ctx context.Context, cred Credentials, region, domainID string) (identitydomains.IdentityDomainsClient, error) { + url, err := c.domainURL(ctx, cred, region, domainID) if err != nil { return identitydomains.IdentityDomainsClient{}, err } @@ -80,13 +136,20 @@ func (c *RealClient) domainsClient(ctx context.Context, cred Credentials, region return dc, nil } -// pickDomainURL 优先返回名为 Default 的域,否则取第一个 ACTIVE 的域。 -func pickDomainURL(domains []identity.DomainSummary) string { +// pickDomainURL 在 ACTIVE 域中选取:domainID 非空按 OCID 匹配; +// 为空优先名为 Default 的域,否则取第一个 ACTIVE 的域。 +func pickDomainURL(domains []identity.DomainSummary, domainID string) string { fallback := "" for _, d := range domains { if d.LifecycleState != identity.DomainLifecycleStateActive { continue } + if domainID != "" { + if deref(d.Id) == domainID { + return deref(d.Url) + } + continue + } if deref(d.DisplayName) == "Default" { return deref(d.Url) } @@ -98,8 +161,8 @@ func pickDomainURL(domains []identity.DomainSummary) string { } // GetNotificationRecipients 实现 Client:查询域通知收件人设置。 -func (c *RealClient) GetNotificationRecipients(ctx context.Context, cred Credentials, region string) (NotificationRecipients, error) { - dc, err := c.domainsClient(ctx, cred, region) +func (c *RealClient) GetNotificationRecipients(ctx context.Context, cred Credentials, region, domainID string) (NotificationRecipients, error) { + dc, err := c.domainsClient(ctx, cred, region, domainID) if err != nil { return NotificationRecipients{}, err } @@ -115,8 +178,8 @@ func (c *RealClient) GetNotificationRecipients(ctx context.Context, cred Credent // UpdateNotificationRecipients 实现 Client:把域通知改为只发给指定收件人; // emails 为空时关闭 test mode 恢复默认发送。其余字段原样保留。 -func (c *RealClient) UpdateNotificationRecipients(ctx context.Context, cred Credentials, region string, emails []string) (NotificationRecipients, error) { - dc, err := c.domainsClient(ctx, cred, region) +func (c *RealClient) UpdateNotificationRecipients(ctx context.Context, cred Credentials, region, domainID string, emails []string) (NotificationRecipients, error) { + dc, err := c.domainsClient(ctx, cred, region, domainID) if err != nil { return NotificationRecipients{}, err } @@ -153,8 +216,8 @@ func getNotificationSetting(ctx context.Context, dc identitydomains.IdentityDoma } // ListPasswordPolicies 实现 Client:列出域内全部密码策略。 -func (c *RealClient) ListPasswordPolicies(ctx context.Context, cred Credentials, region string) ([]PasswordPolicyInfo, error) { - dc, err := c.domainsClient(ctx, cred, region) +func (c *RealClient) ListPasswordPolicies(ctx context.Context, cred Credentials, region, domainID string) ([]PasswordPolicyInfo, error) { + dc, err := c.domainsClient(ctx, cred, region, domainID) if err != nil { return nil, err } @@ -171,12 +234,12 @@ func (c *RealClient) ListPasswordPolicies(ctx context.Context, cred Credentials, // UpdatePasswordPolicy 实现 Client:以 SCIM Patch 修改密码策略的指定字段。 // 仅 Custom 强度的策略可修改,Simple/Standard 内置策略 OCI 返回 403。 -func (c *RealClient) UpdatePasswordPolicy(ctx context.Context, cred Credentials, region, policyID string, in UpdatePasswordPolicyInput) (PasswordPolicyInfo, error) { +func (c *RealClient) UpdatePasswordPolicy(ctx context.Context, cred Credentials, region, domainID, policyID string, in UpdatePasswordPolicyInput) (PasswordPolicyInfo, error) { ops := buildPolicyPatchOps(in) if len(ops) == 0 { return PasswordPolicyInfo{}, fmt.Errorf("update password policy: nothing to update") } - dc, err := c.domainsClient(ctx, cred, region) + dc, err := c.domainsClient(ctx, cred, region, domainID) if err != nil { return PasswordPolicyInfo{}, err } @@ -233,8 +296,8 @@ func removeOp(path string) identitydomains.Operations { } // GetIdentitySetting 实现 Client:读取域身份设置(IdentitySettings 单例资源)。 -func (c *RealClient) GetIdentitySetting(ctx context.Context, cred Credentials, region string) (IdentitySettingInfo, error) { - dc, err := c.domainsClient(ctx, cred, region) +func (c *RealClient) GetIdentitySetting(ctx context.Context, cred Credentials, region, domainID string) (IdentitySettingInfo, error) { + dc, err := c.domainsClient(ctx, cred, region, domainID) if err != nil { return IdentitySettingInfo{}, err } @@ -254,12 +317,12 @@ func (c *RealClient) GetIdentitySetting(ctx context.Context, cred Credentials, r } // UpdateIdentitySetting 实现 Client:修改「用户需要提供主电子邮件地址」开关。 -func (c *RealClient) UpdateIdentitySetting(ctx context.Context, cred Credentials, region string, primaryEmailRequired bool) (IdentitySettingInfo, error) { - current, err := c.GetIdentitySetting(ctx, cred, region) +func (c *RealClient) UpdateIdentitySetting(ctx context.Context, cred Credentials, region, domainID string, primaryEmailRequired bool) (IdentitySettingInfo, error) { + current, err := c.GetIdentitySetting(ctx, cred, region, domainID) if err != nil { return IdentitySettingInfo{}, err } - dc, err := c.domainsClient(ctx, cred, region) + dc, err := c.domainsClient(ctx, cred, region, domainID) if err != nil { return IdentitySettingInfo{}, err } diff --git a/internal/oci/federation.go b/internal/oci/federation.go index db6e86c..3d18e29 100644 --- a/internal/oci/federation.go +++ b/internal/oci/federation.go @@ -56,8 +56,8 @@ type CreateIdpInput struct { } // ListIdentityProviders 实现 Client:列出域内 SAML 身份提供者。 -func (c *RealClient) ListIdentityProviders(ctx context.Context, cred Credentials, region string) ([]IdentityProviderInfo, error) { - dc, err := c.domainsClient(ctx, cred, region) +func (c *RealClient) ListIdentityProviders(ctx context.Context, cred Credentials, region, domainID string) ([]IdentityProviderInfo, error) { + dc, err := c.domainsClient(ctx, cred, region, domainID) if err != nil { return nil, err } @@ -76,8 +76,8 @@ func (c *RealClient) ListIdentityProviders(ctx context.Context, cred Credentials } // CreateSamlIdentityProvider 实现 Client:按输入创建禁用态 SAML IdP 并配置 JIT。 -func (c *RealClient) CreateSamlIdentityProvider(ctx context.Context, cred Credentials, region string, in CreateIdpInput) (IdentityProviderInfo, error) { - dc, err := c.domainsClient(ctx, cred, region) +func (c *RealClient) CreateSamlIdentityProvider(ctx context.Context, cred Credentials, region, domainID string, in CreateIdpInput) (IdentityProviderInfo, error) { + dc, err := c.domainsClient(ctx, cred, region, domainID) if err != nil { return IdentityProviderInfo{}, err } @@ -139,22 +139,27 @@ func buildSamlIdp(in CreateIdpInput, adminGroup *identitydomains.IdentityProvide return idp } -// adminGroupRef 查询域内 Administrators 组作为 JIT 静态分配目标;不存在时返回 nil。 +// adminGroupRef 查询域内管理员组作为 JIT 静态分配与用户授权目标; +// 按 adminGroupNames 优先序取第一个存在的组,都不存在返回 nil。 func adminGroupRef(ctx context.Context, dc identitydomains.IdentityDomainsClient) (*identitydomains.IdentityProviderJitUserProvAssignedGroups, error) { - filter := fmt.Sprintf("displayName eq %q", administratorsGroupName) - count := 1 + filter := fmt.Sprintf("displayName eq %q or displayName eq %q", adminGroupNames[0], adminGroupNames[1]) + count := len(adminGroupNames) resp, err := dc.ListGroups(ctx, identitydomains.ListGroupsRequest{ Filter: &filter, Attributes: common.String("id,displayName"), Count: &count, }) if err != nil { - return nil, fmt.Errorf("find Administrators group: %w", err) + return nil, fmt.Errorf("find administrators group: %w", err) } - if len(resp.Resources) == 0 || resp.Resources[0].Id == nil { - return nil, nil + for _, name := range adminGroupNames { + for _, g := range resp.Resources { + if deref(g.DisplayName) == name && g.Id != nil { + return &identitydomains.IdentityProviderJitUserProvAssignedGroups{Value: g.Id}, nil + } + } } - return &identitydomains.IdentityProviderJitUserProvAssignedGroups{Value: resp.Resources[0].Id}, nil + return nil, nil } // jitMappings 生成 JIT 属性映射:来源为 NameID(或指定断言属性), @@ -207,8 +212,8 @@ func ensureJitAttributeMappings(ctx context.Context, dc identitydomains.Identity // SetIdentityProviderEnabled 实现 Client:启用/停用 IdP 并同步登录页显示。 // 顺序有讲究:启用后才能上登录页,停用前必须先从登录页移除。 -func (c *RealClient) SetIdentityProviderEnabled(ctx context.Context, cred Credentials, region, idpID string, enabled bool) (IdentityProviderInfo, error) { - dc, err := c.domainsClient(ctx, cred, region) +func (c *RealClient) SetIdentityProviderEnabled(ctx context.Context, cred Credentials, region, domainID, idpID string, enabled bool) (IdentityProviderInfo, error) { + dc, err := c.domainsClient(ctx, cred, region, domainID) if err != nil { return IdentityProviderInfo{}, err } @@ -249,8 +254,8 @@ func patchIdpEnabled(ctx context.Context, dc identitydomains.IdentityDomainsClie } // DeleteIdentityProvider 实现 Client:删除 IdP;先移出登录页并停用。 -func (c *RealClient) DeleteIdentityProvider(ctx context.Context, cred Credentials, region, idpID string) error { - dc, err := c.domainsClient(ctx, cred, region) +func (c *RealClient) DeleteIdentityProvider(ctx context.Context, cred Credentials, region, domainID, idpID string) error { + dc, err := c.domainsClient(ctx, cred, region, domainID) if err != nil { return err } @@ -346,8 +351,8 @@ func toggleJSONList(list, item string, add bool) (string, bool, error) { // DownloadDomainSamlMetadata 实现 Client:下载域的 SP SAML 元数据 XML。 // 域设置「访问签名证书」未开启时端点不可用,自动开启后重试 // (联邦配置要求 IdP 侧能读取该元数据,开启公开访问是官方标准做法)。 -func (c *RealClient) DownloadDomainSamlMetadata(ctx context.Context, cred Credentials, region string) ([]byte, error) { - url, err := c.domainURL(ctx, cred, region) +func (c *RealClient) DownloadDomainSamlMetadata(ctx context.Context, cred Credentials, region, domainID string) ([]byte, error) { + url, err := c.domainURL(ctx, cred, region, domainID) if err != nil { return nil, err } @@ -358,7 +363,7 @@ func (c *RealClient) DownloadDomainSamlMetadata(ctx context.Context, cred Creden if status == http.StatusOK { return body, nil } - if err := c.enableSigningCertPublicAccess(ctx, cred, region); err != nil { + if err := c.enableSigningCertPublicAccess(ctx, cred, region, domainID); err != nil { return nil, err } body, status, err = fetchSamlMetadata(ctx, url) @@ -390,8 +395,8 @@ func fetchSamlMetadata(ctx context.Context, url string) ([]byte, int, error) { } // enableSigningCertPublicAccess 开启域设置「访问签名证书」公开访问。 -func (c *RealClient) enableSigningCertPublicAccess(ctx context.Context, cred Credentials, region string) error { - dc, err := c.domainsClient(ctx, cred, region) +func (c *RealClient) enableSigningCertPublicAccess(ctx context.Context, cred Credentials, region, domainID string) error { + dc, err := c.domainsClient(ctx, cred, region, domainID) if err != nil { return err } diff --git a/internal/oci/iddomain_test.go b/internal/oci/iddomain_test.go new file mode 100644 index 0000000..9c25007 --- /dev/null +++ b/internal/oci/iddomain_test.go @@ -0,0 +1,107 @@ +package oci + +import ( + "testing" + "time" + + "github.com/oracle/oci-go-sdk/v65/common" + "github.com/oracle/oci-go-sdk/v65/identity" + "github.com/oracle/oci-go-sdk/v65/identitydomains" +) + +func domainSummary(id, name, url string, state identity.DomainLifecycleStateEnum) identity.DomainSummary { + return identity.DomainSummary{Id: &id, DisplayName: &name, Url: &url, LifecycleState: state} +} + +func TestPickDomainURL(t *testing.T) { + domains := []identity.DomainSummary{ + domainSummary("ocid1.domain.idcs", "OracleIdentityCloudService", "https://idcs-a.example.com", identity.DomainLifecycleStateActive), + domainSummary("ocid1.domain.default", "Default", "https://idcs-b.example.com", identity.DomainLifecycleStateActive), + domainSummary("ocid1.domain.off", "Off", "https://idcs-c.example.com", identity.DomainLifecycleStateInactive), + } + cases := []struct { + name string + domains []identity.DomainSummary + domainID string + want string + }{ + {"缺省优先 Default 域", domains, "", "https://idcs-b.example.com"}, + {"按 OCID 精确匹配", domains, "ocid1.domain.idcs", "https://idcs-a.example.com"}, + {"OCID 不存在返回空", domains, "ocid1.domain.miss", ""}, + {"非 ACTIVE 域不可选", domains, "ocid1.domain.off", ""}, + {"无 Default 时取第一个 ACTIVE", domains[:1], "", "https://idcs-a.example.com"}, + {"空列表返回空", nil, "", ""}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + if got := pickDomainURL(tc.domains, tc.domainID); got != tc.want { + t.Fatalf("pickDomainURL() = %q, want %q", got, tc.want) + } + }) + } +} + +func TestToDomainListUser(t *testing.T) { + created := "2024-12-12T14:52:00Z" + lastLogin := "2026-07-01T08:00:00.123Z" + yes, no := true, false + u := identitydomains.User{ + Id: common.String("scim-1"), + Ocid: common.String("ocid1.user.oc1..a"), + UserName: common.String("oci"), + Active: &yes, + Meta: &identitydomains.Meta{Created: &created}, + Emails: []identitydomains.UserEmails{{ + Value: common.String("a@b.c"), Primary: &yes, Verified: &yes, + Type: identitydomains.UserEmailsTypeWork, + }}, + UrnIetfParamsScimSchemasOracleIdcsExtensionMfaUser: &identitydomains.ExtensionMfaUser{ + MfaStatus: identitydomains.ExtensionMfaUserMfaStatusEnrolled, + }, + UrnIetfParamsScimSchemasOracleIdcsExtensionUserStateUser: &identitydomains.ExtensionUserStateUser{ + LastSuccessfulLoginDate: &lastLogin, + }, + } + got := toDomainListUser(u, "ocid1.user.oc1..a") + if got.ID != "ocid1.user.oc1..a" || !got.IsCurrentUser { + t.Fatalf("ID/IsCurrentUser 映射错误: %+v", got) + } + if !got.MfaActivated || !got.EmailVerified || got.Email != "a@b.c" { + t.Fatalf("MFA/邮箱映射错误: %+v", got) + } + if got.LifecycleState != "ACTIVE" || got.TimeCreated == nil || got.LastLoginTime == nil { + t.Fatalf("状态/时间映射错误: %+v", got) + } + + inactive := identitydomains.User{Id: common.String("scim-2"), Ocid: common.String("ocid1.user.oc1..b"), Active: &no} + if s := toDomainListUser(inactive, "x").LifecycleState; s != "INACTIVE" { + t.Fatalf("inactive 用户 LifecycleState = %q, want INACTIVE", s) + } +} + +func TestParseScimTime(t *testing.T) { + valid := "2026-07-10T01:02:03Z" + bad := "not-a-time" + empty := "" + cases := []struct { + name string + in *string + want *time.Time + }{ + {"合法 RFC3339", &valid, func() *time.Time { t0, _ := time.Parse(time.RFC3339, valid); return &t0 }()}, + {"nil 返回 nil", nil, nil}, + {"空串返回 nil", &empty, nil}, + {"非法格式返回 nil", &bad, nil}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + got := parseScimTime(tc.in) + if (got == nil) != (tc.want == nil) { + t.Fatalf("parseScimTime() = %v, want %v", got, tc.want) + } + if got != nil && !got.Equal(*tc.want) { + t.Fatalf("parseScimTime() = %v, want %v", got, tc.want) + } + }) + } +} diff --git a/internal/oci/real_federation_integration_test.go b/internal/oci/real_federation_integration_test.go index 06f4191..39e3bbd 100644 --- a/internal/oci/real_federation_integration_test.go +++ b/internal/oci/real_federation_integration_test.go @@ -25,7 +25,7 @@ func TestRealOCIFederation(t *testing.T) { client := NewClient() region := cred.Region - metadata, err := client.DownloadDomainSamlMetadata(ctx, cred, region) + metadata, err := client.DownloadDomainSamlMetadata(ctx, cred, region, "") if err != nil { t.Fatalf("DownloadDomainSamlMetadata: %v", err) } @@ -34,7 +34,7 @@ func TestRealOCIFederation(t *testing.T) { } t.Logf("domain saml metadata: %d bytes", len(metadata)) - before, err := client.ListConsoleSignOnRules(ctx, cred, region) + before, err := client.ListConsoleSignOnRules(ctx, cred, region, "") if err != nil { t.Fatalf("ListConsoleSignOnRules: %v", err) } @@ -43,7 +43,7 @@ func TestRealOCIFederation(t *testing.T) { idp := createTestIdp(ctx, t, client, cred, region) verifyIdpDefaults(ctx, t, client, cred, region, idp.ID) - activated, err := client.SetIdentityProviderEnabled(ctx, cred, region, idp.ID, true) + activated, err := client.SetIdentityProviderEnabled(ctx, cred, region, "", idp.ID, true) if err != nil { t.Fatalf("activate idp: %v", err) } @@ -54,7 +54,7 @@ func TestRealOCIFederation(t *testing.T) { exemptionRoundTrip(ctx, t, client, cred, region, idp.ID, len(before)) - deactivated, err := client.SetIdentityProviderEnabled(ctx, cred, region, idp.ID, false) + deactivated, err := client.SetIdentityProviderEnabled(ctx, cred, region, "", idp.ID, false) if err != nil { t.Fatalf("deactivate idp: %v", err) } @@ -71,7 +71,7 @@ func createTestIdp(ctx context.Context, t *testing.T, client *RealClient, cred C t.Fatalf("read test-idp.xml: %v", err) } name := fmt.Sprintf("oci-portal-e2e-idp-%d", time.Now().Unix()) - idp, err := client.CreateSamlIdentityProvider(ctx, cred, region, CreateIdpInput{ + idp, err := client.CreateSamlIdentityProvider(ctx, cred, region, "", CreateIdpInput{ Name: name, Metadata: string(metadata), Description: "oci-portal e2e temporary idp", @@ -86,7 +86,7 @@ func createTestIdp(ctx context.Context, t *testing.T, client *RealClient, cred C t.Cleanup(func() { cctx, ccancel := context.WithTimeout(context.Background(), 2*time.Minute) defer ccancel() - if err := client.DeleteIdentityProvider(cctx, cred, region, idp.ID); err != nil { + if err := client.DeleteIdentityProvider(cctx, cred, region, "", idp.ID); err != nil { t.Errorf("cleanup: delete idp: %v", err) } else { t.Log("cleanup: idp deleted") @@ -98,7 +98,7 @@ func createTestIdp(ctx context.Context, t *testing.T, client *RealClient, cred C // verifyIdpDefaults 回读 IdP 断言控制台默认值与 JIT 属性映射两条。 func verifyIdpDefaults(ctx context.Context, t *testing.T, client *RealClient, cred Credentials, region, idpID string) { t.Helper() - dc, err := client.domainsClient(ctx, cred, region) + dc, err := client.domainsClient(ctx, cred, region, "") if err != nil { t.Fatalf("domainsClient: %v", err) } @@ -151,7 +151,7 @@ func verifyJitMappings(ctx context.Context, t *testing.T, dc identitydomains.Ide // assertLoginPage 校验 DefaultIDPRule 的 SamlIDPs 是否包含目标 IdP。 func assertLoginPage(ctx context.Context, t *testing.T, client *RealClient, cred Credentials, region, idpID string, want bool) { t.Helper() - dc, err := client.domainsClient(ctx, cred, region) + dc, err := client.domainsClient(ctx, cred, region, "") if err != nil { t.Fatalf("domainsClient: %v", err) } @@ -175,7 +175,7 @@ func assertLoginPage(ctx context.Context, t *testing.T, client *RealClient, cred // exemptionRoundTrip 创建免 MFA 规则断言置顶与顺延,删除后断言完全复位。 func exemptionRoundTrip(ctx context.Context, t *testing.T, client *RealClient, cred Credentials, region, idpID string, beforeCount int) { t.Helper() - rule, err := client.CreateMfaExemptionRule(ctx, cred, region, idpID, "oci-portal-e2e-skip-mfa") + rule, err := client.CreateMfaExemptionRule(ctx, cred, region, "", idpID, "oci-portal-e2e-skip-mfa") if err != nil { t.Fatalf("CreateMfaExemptionRule: %v", err) } @@ -185,12 +185,12 @@ func exemptionRoundTrip(ctx context.Context, t *testing.T, client *RealClient, c if deleted { return } - if err := client.DeleteMfaExemptionRule(ctx, cred, region, rule.ID); err != nil { + if err := client.DeleteMfaExemptionRule(ctx, cred, region, "", rule.ID); err != nil { t.Errorf("cleanup: delete exemption rule: %v", err) } }() - after, err := client.ListConsoleSignOnRules(ctx, cred, region) + after, err := client.ListConsoleSignOnRules(ctx, cred, region, "") if err != nil { t.Fatalf("ListConsoleSignOnRules after create: %v", err) } @@ -208,11 +208,11 @@ func exemptionRoundTrip(ctx context.Context, t *testing.T, client *RealClient, c t.Errorf("condition value %q does not reference idp", after[0].ConditionValue) } - if err := client.DeleteMfaExemptionRule(ctx, cred, region, rule.ID); err != nil { + if err := client.DeleteMfaExemptionRule(ctx, cred, region, "", rule.ID); err != nil { t.Fatalf("DeleteMfaExemptionRule: %v", err) } deleted = true - restored, err := client.ListConsoleSignOnRules(ctx, cred, region) + restored, err := client.ListConsoleSignOnRules(ctx, cred, region, "") if err != nil { t.Fatalf("ListConsoleSignOnRules after delete: %v", err) } diff --git a/internal/oci/real_tenant_integration_test.go b/internal/oci/real_tenant_integration_test.go index b3c37d6..1445603 100644 --- a/internal/oci/real_tenant_integration_test.go +++ b/internal/oci/real_tenant_integration_test.go @@ -21,7 +21,7 @@ func TestRealOCITenantUsers(t *testing.T) { client := NewClient() homeRegion := cred.Region - users, err := client.ListTenantUsers(ctx, cred) + users, err := client.ListTenantUsers(ctx, cred, homeRegion, "") if err != nil { t.Fatalf("ListTenantUsers: %v", err) } @@ -40,7 +40,7 @@ func TestRealOCITenantUsers(t *testing.T) { t.Logf("tenant users: %d", len(users)) name := fmt.Sprintf("oci-portal-e2e-user-%d", time.Now().Unix()) - user, err := client.CreateTenantUser(ctx, cred, homeRegion, CreateTenantUserInput{ + user, err := client.CreateTenantUser(ctx, cred, homeRegion, "", CreateTenantUserInput{ Name: name, Description: "e2e temp user", Email: name + "@example.com", GivenName: "E2E", FamilyName: "Temp", }) @@ -49,7 +49,7 @@ func TestRealOCITenantUsers(t *testing.T) { } t.Logf("created user %s (%s)", user.Name, user.ID) defer func() { - if err := client.DeleteTenantUser(ctx, cred, homeRegion, user.ID); err != nil { + if err := client.DeleteTenantUser(ctx, cred, homeRegion, "", user.ID); err != nil { t.Errorf("cleanup: DeleteTenantUser: %v", err) } else { t.Log("cleanup: user deleted") @@ -57,7 +57,7 @@ func TestRealOCITenantUsers(t *testing.T) { }() newDesc, newEmail, newGiven := "e2e updated", name+"+upd@example.com", "E2E2" - updated, err := client.UpdateTenantUser(ctx, cred, homeRegion, user.ID, UpdateTenantUserInput{ + updated, err := client.UpdateTenantUser(ctx, cred, homeRegion, "", user.ID, UpdateTenantUserInput{ Description: &newDesc, Email: &newEmail, GivenName: &newGiven, }) if err != nil { @@ -68,7 +68,7 @@ func TestRealOCITenantUsers(t *testing.T) { } t.Logf("updated user desc=%q email=%q", updated.Description, updated.Email) - password, err := client.ResetTenantUserPassword(ctx, cred, homeRegion, user.ID) + password, err := client.ResetTenantUserPassword(ctx, cred, homeRegion, "", user.ID) if err != nil { t.Fatalf("ResetTenantUserPassword: %v", err) } @@ -77,7 +77,7 @@ func TestRealOCITenantUsers(t *testing.T) { } t.Logf("password reset ok (len=%d)", len(password)) - mfaDeleted, err := client.DeleteTenantUserMfaDevices(ctx, cred, homeRegion, user.ID) + mfaDeleted, err := client.DeleteTenantUserMfaDevices(ctx, cred, homeRegion, "", user.ID) if err != nil { t.Fatalf("DeleteTenantUserMfaDevices: %v", err) } @@ -101,12 +101,12 @@ func TestRealOCIDomainSettings(t *testing.T) { client := NewClient() region := cred.Region - original, err := client.GetNotificationRecipients(ctx, cred, region) + original, err := client.GetNotificationRecipients(ctx, cred, region, "") if err != nil { t.Fatalf("GetNotificationRecipients: %v", err) } t.Logf("original recipients=%v testMode=%v", original.Recipients, original.TestModeEnabled) - updated, err := client.UpdateNotificationRecipients(ctx, cred, region, []string{"e2e@example.com"}) + updated, err := client.UpdateNotificationRecipients(ctx, cred, region, "", []string{"e2e@example.com"}) if err != nil { t.Fatalf("UpdateNotificationRecipients: %v", err) } @@ -117,13 +117,13 @@ func TestRealOCIDomainSettings(t *testing.T) { if !original.TestModeEnabled { restoreRecipients = nil } - if _, err := client.UpdateNotificationRecipients(ctx, cred, region, restoreRecipients); err != nil { + if _, err := client.UpdateNotificationRecipients(ctx, cred, region, "", restoreRecipients); err != nil { t.Errorf("restore recipients: %v", err) } else { t.Log("recipients restored") } - policies, err := client.ListPasswordPolicies(ctx, cred, region) + policies, err := client.ListPasswordPolicies(ctx, cred, region, "") if err != nil { t.Fatalf("ListPasswordPolicies: %v", err) } @@ -133,7 +133,7 @@ func TestRealOCIDomainSettings(t *testing.T) { } t.Logf("policies: %d, patch target %s strength=%s expires=%v", len(policies), target.Name, target.PasswordStrength, target.PasswordExpiresAfter) days := 350 - patched, err := client.UpdatePasswordPolicy(ctx, cred, region, target.ID, UpdatePasswordPolicyInput{PasswordExpiresAfter: &days}) + patched, err := client.UpdatePasswordPolicy(ctx, cred, region, "", target.ID, UpdatePasswordPolicyInput{PasswordExpiresAfter: &days}) if err != nil { t.Fatalf("UpdatePasswordPolicy: %v", err) } @@ -144,7 +144,7 @@ func TestRealOCIDomainSettings(t *testing.T) { if target.PasswordExpiresAfter != nil { restore = *target.PasswordExpiresAfter } - if _, err := client.UpdatePasswordPolicy(ctx, cred, region, target.ID, UpdatePasswordPolicyInput{PasswordExpiresAfter: &restore}); err != nil { + if _, err := client.UpdatePasswordPolicy(ctx, cred, region, "", target.ID, UpdatePasswordPolicyInput{PasswordExpiresAfter: &restore}); err != nil { t.Errorf("restore password policy: %v", err) } else { t.Logf("policy restored to %d", restore) diff --git a/internal/oci/signon.go b/internal/oci/signon.go index a508a5a..2bf4a3d 100644 --- a/internal/oci/signon.go +++ b/internal/oci/signon.go @@ -29,8 +29,8 @@ type SignOnRuleInfo struct { } // ListConsoleSignOnRules 实现 Client:按优先级列出 OCI Console sign-on 策略的规则。 -func (c *RealClient) ListConsoleSignOnRules(ctx context.Context, cred Credentials, region string) ([]SignOnRuleInfo, error) { - dc, err := c.domainsClient(ctx, cred, region) +func (c *RealClient) ListConsoleSignOnRules(ctx context.Context, cred Credentials, region, domainID string) ([]SignOnRuleInfo, error) { + dc, err := c.domainsClient(ctx, cred, region, domainID) if err != nil { return nil, err } @@ -99,8 +99,8 @@ func fillRuleCondition(ctx context.Context, dc identitydomains.IdentityDomainsCl // CreateMfaExemptionRule 实现 Client:为指定 IdP 创建免 MFA sign-on 规则并置顶。 // 规则语义与控制台一致:subject.authenticatedBy in [idpID] → authenticationFactor=IDP。 -func (c *RealClient) CreateMfaExemptionRule(ctx context.Context, cred Credentials, region, idpID, ruleName string) (SignOnRuleInfo, error) { - dc, err := c.domainsClient(ctx, cred, region) +func (c *RealClient) CreateMfaExemptionRule(ctx context.Context, cred Credentials, region, domainID, idpID, ruleName string) (SignOnRuleInfo, error) { + dc, err := c.domainsClient(ctx, cred, region, domainID) if err != nil { return SignOnRuleInfo{}, err } @@ -239,11 +239,11 @@ func patchPolicyRules(ctx context.Context, dc identitydomains.IdentityDomainsCli // DeleteMfaExemptionRule 实现 Client:删除免 MFA 规则并连带清理其条件。 // 拒绝删除 Oracle 预置规则。 -func (c *RealClient) DeleteMfaExemptionRule(ctx context.Context, cred Credentials, region, ruleID string) error { +func (c *RealClient) DeleteMfaExemptionRule(ctx context.Context, cred Credentials, region, domainID, ruleID string) error { if strings.HasPrefix(ruleID, "OciConsole") { return fmt.Errorf("delete sign-on rule: %s is a built-in rule and cannot be deleted", ruleID) } - dc, err := c.domainsClient(ctx, cred, region) + dc, err := c.domainsClient(ctx, cred, region, domainID) if err != nil { return err } diff --git a/internal/oci/tenantuser.go b/internal/oci/tenantuser.go index 04d5dad..8f55896 100644 --- a/internal/oci/tenantuser.go +++ b/internal/oci/tenantuser.go @@ -6,6 +6,7 @@ import ( "errors" "fmt" "math/big" + "slices" "strings" "time" @@ -28,11 +29,18 @@ type TenantUser struct { LastLoginTime *time.Time `json:"lastLoginTime,omitempty"` } -// Oracle 预置的管理员角色与管理员组显示名。 -const ( - domainAdminRoleName = "Identity Domain Administrator" - administratorsGroupName = "Administrators" -) +// Oracle 预置的「身份域管理员」应用角色显示名。 +const domainAdminRoleName = "Identity Domain Administrator" + +// adminGroupNames 是租户管理员组候选名(按优先序):原生域为 Administrators, +// IDCS 迁移域(idcs foundation,如 OracleIdentityCloudService)为 OCI_Administrators; +// 迁移域另有 IDCS_Administrators 属身份域管理范畴,由上面的应用角色概念覆盖,不算在内。 +var adminGroupNames = []string{"Administrators", "OCI_Administrators"} + +// isAdminGroupName 判断组显示名是否为租户管理员组。 +func isAdminGroupName(name string) bool { + return slices.Contains(adminGroupNames, name) +} // scimAppRolesAttr 是 SCIM 用户 appRoles 扩展属性的完整路径。 const scimAppRolesAttr = "urn:ietf:params:scim:schemas:oracle:idcs:extension:user:User:appRoles" @@ -48,8 +56,8 @@ type TenantUserDetail struct { } // GetTenantUserDetail 实现 Client:查用户域档案与管理员状态(编辑表单预填充用)。 -func (c *RealClient) GetTenantUserDetail(ctx context.Context, cred Credentials, homeRegion, userID string) (TenantUserDetail, error) { - dc, err := c.domainsClient(ctx, cred, homeRegion) +func (c *RealClient) GetTenantUserDetail(ctx context.Context, cred Credentials, homeRegion, domainID, userID string) (TenantUserDetail, error) { + dc, err := c.domainsClient(ctx, cred, homeRegion, domainID) if err != nil { return TenantUserDetail{}, err } @@ -68,7 +76,7 @@ func toTenantUserDetail(u identitydomains.User) TenantUserDetail { d.FamilyName = deref(u.Name.FamilyName) } for _, g := range u.Groups { - if deref(g.Display) == administratorsGroupName { + if isAdminGroupName(deref(g.Display)) { d.InAdminGroup = true } } @@ -94,8 +102,13 @@ func (c *RealClient) identityClientAt(cred Credentials, region string) (identity return ic, nil } -// ListTenantUsers 实现 Client:列出租户全部 IAM 用户。 -func (c *RealClient) ListTenantUsers(ctx context.Context, cred Credentials) ([]TenantUser, error) { +// ListTenantUsers 实现 Client:列出租户 IAM 用户。 +// domainID 非空时经 Identity Domains SCIM 只列该域用户; +// 为空时走经典 IAM(跨域拍平视图,兼容无域老租户)。 +func (c *RealClient) ListTenantUsers(ctx context.Context, cred Credentials, homeRegion, domainID string) ([]TenantUser, error) { + if domainID != "" { + return c.listDomainUsers(ctx, cred, homeRegion, domainID) + } ic, err := c.identityClientAt(cred, "") if err != nil { return nil, err @@ -117,6 +130,78 @@ func (c *RealClient) ListTenantUsers(ctx context.Context, cred Credentials) ([]T } } +// scimListAttrs 是按域列用户时请求的属性集:列表字段 + MFA 与最近登录扩展。 +const scimListAttrs = "id,ocid,userName,description,emails,active,meta.created," + + "urn:ietf:params:scim:schemas:oracle:idcs:extension:mfa:User:mfaStatus," + + "urn:ietf:params:scim:schemas:oracle:idcs:extension:userState:User:lastSuccessfulLoginDate" + +// maxScimUserPages 限制按域列用户的翻页数(每页 200,上限 5000 人)。 +const maxScimUserPages = 25 + +// listDomainUsers 经 SCIM 按 startIndex 翻页列出域内全部用户。 +func (c *RealClient) listDomainUsers(ctx context.Context, cred Credentials, homeRegion, domainID string) ([]TenantUser, error) { + dc, err := c.domainsClient(ctx, cred, homeRegion, domainID) + if err != nil { + return nil, err + } + var users []TenantUser + attrs := scimListAttrs + count := 200 + start := 1 + for range maxScimUserPages { + req := identitydomains.ListUsersRequest{Attributes: &attrs, Count: &count, StartIndex: &start} + resp, err := dc.ListUsers(ctx, req) + if err != nil { + return nil, fmt.Errorf("list domain users: %w", err) + } + for _, u := range resp.Resources { + users = append(users, toDomainListUser(u, cred.UserOCID)) + } + start += len(resp.Resources) + if resp.TotalResults == nil || start > *resp.TotalResults || len(resp.Resources) == 0 { + return orEmpty(users), nil + } + } + return orEmpty(users), nil +} + +// toDomainListUser 把 SCIM 用户映射为列表 DTO;ID 用 IAM OCID +// 保持与经典路径同一 id 语义(IsCurrentUser 比对、后续操作路由)。 +func toDomainListUser(u identitydomains.User, currentUserOCID string) TenantUser { + out := toDomainTenantUser(u) + out.IsCurrentUser = out.ID == currentUserOCID + if u.Active != nil && !*u.Active { + out.LifecycleState = "INACTIVE" + } + if u.Meta != nil { + out.TimeCreated = parseScimTime(u.Meta.Created) + } + if ext := u.UrnIetfParamsScimSchemasOracleIdcsExtensionMfaUser; ext != nil { + out.MfaActivated = ext.MfaStatus == identitydomains.ExtensionMfaUserMfaStatusEnrolled + } + if ext := u.UrnIetfParamsScimSchemasOracleIdcsExtensionUserStateUser; ext != nil { + out.LastLoginTime = parseScimTime(ext.LastSuccessfulLoginDate) + } + for _, e := range u.Emails { + if e.Primary != nil && *e.Primary { + out.EmailVerified = e.Verified != nil && *e.Verified + } + } + return out +} + +// parseScimTime 解析 SCIM 的 RFC3339 时间串;空值或格式异常返回 nil。 +func parseScimTime(s *string) *time.Time { + if s == nil || *s == "" { + return nil + } + t, err := time.Parse(time.RFC3339, *s) + if err != nil { + return nil + } + return &t +} + // CreateTenantUserInput 是创建用户的输入;GivenName/FamilyName 走 // Identity Domains SCIM(经典 IAM 无对应字段)。两个管理员选项同时勾选时 // 先授身份域管理员角色,再加入 Administrators 组。 @@ -133,12 +218,13 @@ type CreateTenantUserInput struct { // CreateTenantUser 实现 Client:创建用户(须发往 home region)。 // 优先 Identity Domains SCIM(支持姓名与管理员授权);SCIM 不可用时回退 // 经典 IAM API,此时管理员选项无法执行、姓名并入描述。 -func (c *RealClient) CreateTenantUser(ctx context.Context, cred Credentials, homeRegion string, in CreateTenantUserInput) (TenantUser, error) { - user, err := c.createUserViaDomain(ctx, cred, homeRegion, in) +// domainID 非空表示用户显式选择了目标域,SCIM 失败不再回退(避免建错域)。 +func (c *RealClient) CreateTenantUser(ctx context.Context, cred Credentials, homeRegion, domainID string, in CreateTenantUserInput) (TenantUser, error) { + user, err := c.createUserViaDomain(ctx, cred, homeRegion, domainID, in) if err == nil { return user, nil } - if in.GrantDomainAdmin || in.AddToAdminGroup { + if domainID != "" || in.GrantDomainAdmin || in.AddToAdminGroup { return TenantUser{}, fmt.Errorf("create user via identity domain: %w", err) } return c.createUserClassic(ctx, cred, homeRegion, in) @@ -170,8 +256,8 @@ func (c *RealClient) createUserClassic(ctx context.Context, cred Credentials, ho } // createUserViaDomain 通过 Identity Domains SCIM 创建用户并按选项授权。 -func (c *RealClient) createUserViaDomain(ctx context.Context, cred Credentials, homeRegion string, in CreateTenantUserInput) (TenantUser, error) { - dc, err := c.domainsClient(ctx, cred, homeRegion) +func (c *RealClient) createUserViaDomain(ctx context.Context, cred Credentials, homeRegion, domainID string, in CreateTenantUserInput) (TenantUser, error) { + dc, err := c.domainsClient(ctx, cred, homeRegion, domainID) if err != nil { return TenantUser{}, err } @@ -373,9 +459,9 @@ type UpdateTenantUserInput struct { // UpdateTenantUser 实现 Client:编辑用户资料(须发往 home region)。 // 优先 Identity Domains SCIM(支持姓名);用户不在域中时回退经典 -// IAM API(仅备注与邮箱,姓名变更报错)。 -func (c *RealClient) UpdateTenantUser(ctx context.Context, cred Credentials, homeRegion, userID string, in UpdateTenantUserInput) (TenantUser, error) { - dc, err := c.domainsClient(ctx, cred, homeRegion) +// IAM API(仅备注与邮箱,姓名变更报错);domainID 非空时不回退。 +func (c *RealClient) UpdateTenantUser(ctx context.Context, cred Credentials, homeRegion, domainID, userID string, in UpdateTenantUserInput) (TenantUser, error) { + dc, err := c.domainsClient(ctx, cred, homeRegion, domainID) if err != nil { return TenantUser{}, err } @@ -394,6 +480,9 @@ func (c *RealClient) UpdateTenantUser(ctx context.Context, cred Credentials, hom } return user, nil } + if domainID != "" { + return TenantUser{}, fmt.Errorf("update user %s: user not found in identity domain", userID) + } if in.GivenName != nil || in.FamilyName != nil { return TenantUser{}, fmt.Errorf("update user %s: 姓名仅身份域用户可修改(该用户不在身份域中)", userID) } @@ -495,7 +584,11 @@ func emailPatchOp(u identitydomains.User, email string) identitydomains.Operatio } // DeleteTenantUser 实现 Client:删除 IAM 用户(须发往 home region)。 -func (c *RealClient) DeleteTenantUser(ctx context.Context, cred Credentials, homeRegion, userID string) error { +// domainID 非空时经该域 SCIM 删除(经典 API 对非默认域用户不生效)。 +func (c *RealClient) DeleteTenantUser(ctx context.Context, cred Credentials, homeRegion, domainID, userID string) error { + if domainID != "" { + return c.deleteDomainUser(ctx, cred, homeRegion, domainID, userID) + } ic, err := c.identityClientAt(cred, homeRegion) if err != nil { return err @@ -506,9 +599,31 @@ func (c *RealClient) DeleteTenantUser(ctx context.Context, cred Credentials, hom return nil } +// deleteDomainUser 经 SCIM 强制删除域用户(forceDelete 连带其授权与组成员关系)。 +func (c *RealClient) deleteDomainUser(ctx context.Context, cred Credentials, homeRegion, domainID, userID string) error { + dc, err := c.domainsClient(ctx, cred, homeRegion, domainID) + if err != nil { + return err + } + scimID, err := scimUserID(ctx, dc, userID) + if err != nil { + return err + } + force := true + _, err = dc.DeleteUser(ctx, identitydomains.DeleteUserRequest{UserId: &scimID, ForceDelete: &force}) + if err != nil { + return fmt.Errorf("delete domain user %s: %w", userID, err) + } + return nil +} + // ResetTenantUserPassword 实现 Client:重置用户控制台密码并返回一次性新密码。 -// 优先经典 IAM API;新式租户(用户只存在于 Identity Domain)自动回退 SCIM。 -func (c *RealClient) ResetTenantUserPassword(ctx context.Context, cred Credentials, homeRegion, userID string) (string, error) { +// domainID 非空直走该域 SCIM;为空优先经典 IAM API, +// 新式租户(用户只存在于 Identity Domain)自动回退 SCIM。 +func (c *RealClient) ResetTenantUserPassword(ctx context.Context, cred Credentials, homeRegion, domainID, userID string) (string, error) { + if domainID != "" { + return c.resetPasswordViaDomain(ctx, cred, homeRegion, domainID, userID) + } ic, err := c.identityClientAt(cred, homeRegion) if err != nil { return "", err @@ -520,12 +635,12 @@ func (c *RealClient) ResetTenantUserPassword(ctx context.Context, cred Credentia if !isClassicUnsupported(err) { return "", fmt.Errorf("reset ui password of %s: %w", userID, err) } - return c.resetPasswordViaDomain(ctx, cred, homeRegion, userID) + return c.resetPasswordViaDomain(ctx, cred, homeRegion, "", userID) } // resetPasswordViaDomain 通过 Identity Domains SCIM API 强设随机新密码。 -func (c *RealClient) resetPasswordViaDomain(ctx context.Context, cred Credentials, homeRegion, userID string) (string, error) { - dc, err := c.domainsClient(ctx, cred, homeRegion) +func (c *RealClient) resetPasswordViaDomain(ctx context.Context, cred Credentials, homeRegion, domainID, userID string) (string, error) { + dc, err := c.domainsClient(ctx, cred, homeRegion, domainID) if err != nil { return "", err } @@ -558,7 +673,11 @@ func (c *RealClient) resetPasswordViaDomain(ctx context.Context, cred Credential // DeleteTenantUserMfaDevices 实现 Client:清除用户全部 MFA。 // 先删经典 TOTP 设备,再通过 Identity Domains 移除全部认证因子(尽力)。 -func (c *RealClient) DeleteTenantUserMfaDevices(ctx context.Context, cred Credentials, homeRegion, userID string) (int, error) { +func (c *RealClient) DeleteTenantUserMfaDevices(ctx context.Context, cred Credentials, homeRegion, domainID, userID string) (int, error) { + // 显式指定域时用户的认证因子只在域内,经典 TOTP 清理不适用 + if domainID != "" { + return 0, c.removeDomainFactors(ctx, cred, homeRegion, domainID, userID) + } ic, err := c.identityClientAt(cred, homeRegion) if err != nil { return 0, err @@ -567,7 +686,7 @@ func (c *RealClient) DeleteTenantUserMfaDevices(ctx context.Context, cred Creden if err != nil { return deleted, err } - if err := c.removeDomainFactors(ctx, cred, homeRegion, userID); err != nil { + if err := c.removeDomainFactors(ctx, cred, homeRegion, "", userID); err != nil { return deleted, err } return deleted, nil @@ -594,8 +713,8 @@ func deleteClassicTotpDevices(ctx context.Context, ic identity.IdentityClient, u // removeDomainFactors 通过 SCIM AuthenticationFactorsRemover 移除用户 // 在 Identity Domain 注册的全部认证因子;用户不在域中时视为无事可做。 -func (c *RealClient) removeDomainFactors(ctx context.Context, cred Credentials, homeRegion, userID string) error { - dc, err := c.domainsClient(ctx, cred, homeRegion) +func (c *RealClient) removeDomainFactors(ctx context.Context, cred Credentials, homeRegion, domainID, userID string) error { + dc, err := c.domainsClient(ctx, cred, homeRegion, domainID) if err != nil { return err } diff --git a/internal/oci/tenantuser_test.go b/internal/oci/tenantuser_test.go index 5c30210..5251f0d 100644 --- a/internal/oci/tenantuser_test.go +++ b/internal/oci/tenantuser_test.go @@ -47,7 +47,7 @@ func TestToTenantUserDetail(t *testing.T) { }, Groups: []identitydomains.UserGroups{ {Value: common.String("g1"), Display: common.String("Readers")}, - {Value: common.String("g2"), Display: common.String(administratorsGroupName)}, + {Value: common.String("g2"), Display: common.String(adminGroupNames[0])}, }, UrnIetfParamsScimSchemasOracleIdcsExtensionUserUser: adminRoles, }, diff --git a/internal/service/aigateway.go b/internal/service/aigateway.go index aa53abf..beacd98 100644 --- a/internal/service/aigateway.go +++ b/internal/service/aigateway.go @@ -14,6 +14,7 @@ import ( "time" "gorm.io/gorm" + "gorm.io/gorm/clause" "oci-portal/internal/aiwire" "oci-portal/internal/model" @@ -506,13 +507,36 @@ func (s *AiGatewayService) ProbeAll(ctx context.Context) (string, error) { // LogCall 落一条调用日志(仅元数据与用量,永不含请求 / 响应正文),返回落库 ID 供内容日志关联(失败为 0)。 func (s *AiGatewayService) LogCall(entry model.AiCallLog) uint { entry.ErrMsg = truncateErr(entry.ErrMsg) - if err := s.db.Create(&entry).Error; err != nil { + stored := false + err := s.db.Transaction(func(tx *gorm.DB) error { + ok, err := lockAiLogParent(tx, &model.AiChannel{}, entry.ChannelID) + if err != nil || !ok { + return err + } + stored = true + return tx.Create(&entry).Error + }) + if err != nil { log.Printf("ai call log: %v", err) return 0 } + if !stored { + return 0 + } return entry.ID } +func lockAiLogParent(tx *gorm.DB, value any, id uint) (bool, error) { + if id == 0 { + return true, nil + } + err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Select("id").First(value, id).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + return false, nil + } + return err == nil, err +} + // CallLogs 分页查询调用日志。 func (s *AiGatewayService) CallLogs(ctx context.Context, page, size int) ([]model.AiCallLog, int64, error) { if page < 1 { @@ -555,9 +579,19 @@ func (s *AiGatewayService) UpdateKeyContentLog(ctx context.Context, id uint, hou // LogContent 写一条内容日志(调用方已确认密钥开启且未过期);正文截断至 64KB。 func (s *AiGatewayService) LogContent(entry model.AiContentLog) { + if entry.CallLogID == 0 { + return + } entry.RequestBody = truncateBody(entry.RequestBody) entry.ResponseBody = truncateBody(entry.ResponseBody) - if err := s.db.Create(&entry).Error; err != nil { + err := s.db.Transaction(func(tx *gorm.DB) error { + ok, err := lockAiLogParent(tx, &model.AiCallLog{}, entry.CallLogID) + if err != nil || !ok { + return err + } + return tx.Create(&entry).Error + }) + if err != nil { log.Printf("ai content log: %v", err) } } diff --git a/internal/service/aigateway_test.go b/internal/service/aigateway_test.go index 59a199a..65e70f1 100644 --- a/internal/service/aigateway_test.go +++ b/internal/service/aigateway_test.go @@ -488,6 +488,7 @@ func TestAiContentLogSwitch(t *testing.T) { t.Error("超过 7 天上限应被拒绝") } // 写入与截断(带调用日志关联) + mustCreate(t, gw.db, &model.AiCallLog{ID: 42, KeyID: key.ID, ChannelID: 0}) gw.LogContent(model.AiContentLog{CallLogID: 42, KeyID: key.ID, KeyName: "k1", Endpoint: "openai", Model: "m", RequestBody: strings.Repeat("x", 70*1024)}) rows, total, err := gw.ContentLogs(ctx, key.ID, 0, 1, 20) if err != nil || total != 1 || len(rows[0].RequestBody) != 64*1024 { diff --git a/internal/service/alertrule.go b/internal/service/alertrule.go new file mode 100644 index 0000000..d3c35e6 --- /dev/null +++ b/internal/service/alertrule.go @@ -0,0 +1,305 @@ +package service + +import ( + "context" + "errors" + "fmt" + "log" + "net/netip" + "slices" + "strings" + "time" + + "gorm.io/gorm" + "gorm.io/gorm/clause" + + "oci-portal/internal/model" +) + +// 告警规则约束与命中记录保留期(窗口计数之外多留几天便于排查)。 +const ( + alertMaxThreshold = 100 + alertMaxWindowMin = 1440 + alertHitRetention = 7 * 24 * time.Hour + alertSourceIPIn = "in" + alertSourceIPNotIn = "notin" +) + +// ErrInvalidAlertRule 标记规则字段非法,api 层映射 400。 +var ErrInvalidAlertRule = fmt.Errorf("告警规则字段非法") + +// ListAlertRules 返回全部告警规则(创建顺序)。 +func (s *LogEventService) ListAlertRules(ctx context.Context) ([]model.AlertRule, error) { + var rules []model.AlertRule + if err := s.db.WithContext(ctx).Order("id").Find(&rules).Error; err != nil { + return nil, fmt.Errorf("list alert rules: %w", err) + } + return rules, nil +} + +// CreateAlertRule 校验并创建规则。 +func (s *LogEventService) CreateAlertRule(ctx context.Context, rule model.AlertRule) (model.AlertRule, error) { + if err := validateAlertRule(&rule); err != nil { + return model.AlertRule{}, err + } + rule.ID = 0 + if err := s.db.WithContext(ctx).Create(&rule).Error; err != nil { + return model.AlertRule{}, fmt.Errorf("create alert rule: %w", err) + } + return rule, nil +} + +// UpdateAlertRule 校验并整体覆盖规则(含启停)。 +func (s *LogEventService) UpdateAlertRule(ctx context.Context, id uint, rule model.AlertRule) (model.AlertRule, error) { + if err := validateAlertRule(&rule); err != nil { + return model.AlertRule{}, err + } + var cur model.AlertRule + if err := s.db.WithContext(ctx).First(&cur, id).Error; err != nil { + return model.AlertRule{}, fmt.Errorf("find alert rule %d: %w", id, err) + } + rule.ID, rule.CreatedAt = cur.ID, cur.CreatedAt + if err := s.db.WithContext(ctx).Save(&rule).Error; err != nil { + return model.AlertRule{}, fmt.Errorf("update alert rule: %w", err) + } + return rule, nil +} + +// DeleteAlertRule 删除规则及其命中记录。 +func (s *LogEventService) DeleteAlertRule(ctx context.Context, id uint) error { + if err := s.db.WithContext(ctx).Delete(&model.AlertRule{}, id).Error; err != nil { + return fmt.Errorf("delete alert rule: %w", err) + } + if err := s.db.WithContext(ctx).Where("rule_id = ?", id).Delete(&model.AlertRuleHit{}).Error; err != nil { + return fmt.Errorf("delete alert rule hits: %w", err) + } + return nil +} + +// validateAlertRule 校验字段并归一化;非法时返回含具体原因的 ErrInvalidAlertRule 包装。 +func validateAlertRule(rule *model.AlertRule) error { + rule.Name = strings.TrimSpace(rule.Name) + if rule.Name == "" { + return fmt.Errorf("%w: 名称必填", ErrInvalidAlertRule) + } + if rule.SourceIPMode == "" { + rule.SourceIPMode = alertSourceIPIn + } + if rule.SourceIPMode != alertSourceIPIn && rule.SourceIPMode != alertSourceIPNotIn { + return fmt.Errorf("%w: 来源 IP 模式须为 in/notin", ErrInvalidAlertRule) + } + if rule.Threshold < 1 || rule.Threshold > alertMaxThreshold { + return fmt.Errorf("%w: 阈值须在 1-%d 之间", ErrInvalidAlertRule, alertMaxThreshold) + } + if rule.Threshold > 1 && (rule.WindowMinutes < 1 || rule.WindowMinutes > alertMaxWindowMin) { + return fmt.Errorf("%w: 阈值>1 时窗口须在 1-%d 分钟之间", ErrInvalidAlertRule, alertMaxWindowMin) + } + if rule.EventTypes != "" { + rule.EventTypes = normalizeCSV(rule.EventTypes) + } + return validateAlertRuleIPs(rule) +} + +// validateAlertRuleIPs 归一化并校验来源 IP 列表(裸 IP 或 CIDR)。 +func validateAlertRuleIPs(rule *model.AlertRule) error { + if rule.SourceIPs == "" { + return nil + } + rule.SourceIPs = normalizeCSV(rule.SourceIPs) + for _, item := range strings.Split(rule.SourceIPs, ",") { + if _, err := parseIPMatcher(item); err != nil { + return fmt.Errorf("%w: 来源 IP %q 不是合法的 IP 或 CIDR", ErrInvalidAlertRule, item) + } + } + return nil +} + +// normalizeCSV 去除各项空白与空项后重组逗号分隔串。 +func normalizeCSV(s string) string { + parts := strings.Split(s, ",") + out := parts[:0] + for _, p := range parts { + if p = strings.TrimSpace(p); p != "" { + out = append(out, p) + } + } + return strings.Join(out, ",") +} + +// parseIPMatcher 把裸 IP 或 CIDR 解析为前缀(裸 IP 视为单地址前缀)。 +func parseIPMatcher(item string) (netip.Prefix, error) { + if strings.Contains(item, "/") { + return netip.ParsePrefix(item) + } + addr, err := netip.ParseAddr(item) + if err != nil { + return netip.Prefix{}, err + } + return netip.PrefixFrom(addr, addr.BitLen()), nil +} + +// ipListMatch 报告 ip 是否命中列表中的任一前缀;ip 解析失败视为未命中。 +func ipListMatch(list, ip string) bool { + addr, err := netip.ParseAddr(ip) + if err != nil { + return false + } + for _, item := range strings.Split(list, ",") { + if p, err := parseIPMatcher(item); err == nil && p.Contains(addr) { + return true + } + } + return false +} + +// ruleHits 报告事件是否命中规则的全部条件(AND 语义,空条件视为任意)。 +func ruleHits(rule model.AlertRule, e *model.LogEvent, p parsedEvent) bool { + if rule.OciConfigID != 0 && rule.OciConfigID != e.OciConfigID { + return false + } + name := relayEventShortName(p.EventType) + if rule.EventTypes != "" && !slices.Contains(strings.Split(rule.EventTypes, ","), name) { + return false + } + if rule.ResourceMatch != "" && !strings.Contains(p.ResourceName, rule.ResourceMatch) { + return false + } + return ruleIPHits(rule, p.SourceIP) +} + +// ruleIPHits 按模式判定来源 IP 条件:in 命中列表告警;notin 不在列表才告警, +// 事件缺 IP 字段时 notin 不告警(避免解析缺字段导致白名单误报)。 +func ruleIPHits(rule model.AlertRule, ip string) bool { + if rule.SourceIPs == "" { + return true + } + if rule.SourceIPMode == alertSourceIPNotIn { + return ip != "" && !ipListMatch(rule.SourceIPs, ip) + } + return ipListMatch(rule.SourceIPs, ip) +} + +// matchAlertRules 对一条已解析事件执行全部启用规则;任何内部错误只记日志,不影响解析主流程。 +func (s *LogEventService) matchAlertRules(ctx context.Context, rules []model.AlertRule, e *model.LogEvent, p parsedEvent) { + if s.notifier == nil { + return + } + for _, rule := range rules { + if !rule.Enabled || !ruleHits(rule, e, p) { + continue + } + count, ok := s.recordAlertHit(ctx, rule, e) + if !ok || count < rule.Threshold || !s.alertCooldownPass(rule) { + continue + } + s.notifier.SendTemplateAsync("audit_alert", map[string]string{ + "rule": rule.Name, "tenant": s.configAlias(ctx, e.OciConfigID), + "event": relayEventShortName(p.EventType), "resource": p.ResourceName, + "ip": p.SourceIP, "count": fmt.Sprint(count), + }) + } +} + +// recordAlertHit 落一条命中并返回窗口内累计次数;阈值 1 的规则免计数直接触发。 +func (s *LogEventService) recordAlertHit(ctx context.Context, rule model.AlertRule, e *model.LogEvent) (int, bool) { + count, err := s.recordAlertHitTx(ctx, rule, e) + if err != nil { + if !errors.Is(err, gorm.ErrRecordNotFound) { + log.Printf("alert rule hit record: %v", err) + } + return 0, false + } + return count, true +} + +func (s *LogEventService) recordAlertHitTx(ctx context.Context, rule model.AlertRule, event *model.LogEvent) (int, error) { + count := 0 + err := s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + var err error + count, err = insertAlertHit(tx, rule, event) + return err + }) + return count, err +} + +// insertAlertHit 按 rule→event 锁顺序确认引用存在后插入并统计窗口命中。 +func insertAlertHit(tx *gorm.DB, rule model.AlertRule, event *model.LogEvent) (int, error) { + if err := lockAlertRefs(tx, rule.ID, event.ID); err != nil { + return 0, err + } + if rule.Threshold <= 1 { + return 1, nil + } + now := time.Now() + hit := model.AlertRuleHit{RuleID: rule.ID, LogEventID: event.ID, HitAt: now} + if err := tx.Create(&hit).Error; err != nil { + return 0, fmt.Errorf("create alert rule hit: %w", err) + } + var count int64 + cutoff := now.Add(-time.Duration(rule.WindowMinutes) * time.Minute) + err := tx.Model(&model.AlertRuleHit{}). + Where("rule_id = ? AND hit_at >= ?", rule.ID, cutoff).Count(&count).Error + if err != nil { + return 0, fmt.Errorf("count alert rule hits: %w", err) + } + return int(count), nil +} + +func lockAlertRefs(tx *gorm.DB, ruleID, eventID uint) error { + var rule model.AlertRule + if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Select("id").First(&rule, ruleID).Error; err != nil { + return fmt.Errorf("lock alert rule %d: %w", ruleID, err) + } + var event model.LogEvent + if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Select("id").First(&event, eventID).Error; err != nil { + return fmt.Errorf("lock log event %d: %w", eventID, err) + } + return nil +} + +// alertCooldownPass 报告规则是否已过冷却窗口;通过即记录本次发送时刻。 +// 阈值 1 的规则无冷却(每次命中即时告警,与既有云端事件通知一致)。 +func (s *LogEventService) alertCooldownPass(rule model.AlertRule) bool { + if rule.Threshold <= 1 { + return true + } + s.alertMu.Lock() + defer s.alertMu.Unlock() + window := time.Duration(rule.WindowMinutes) * time.Minute + if last, ok := s.alertSentAt[rule.ID]; ok && time.Since(last) < window { + return false + } + if s.alertSentAt == nil { + s.alertSentAt = map[uint]time.Time{} + } + s.alertSentAt[rule.ID] = time.Now() + return true +} + +// ClearAlertCooldown 清除已删除租户规则的进程内冷却状态。 +func (s *LogEventService) ClearAlertCooldown(ruleIDs []uint) { + s.alertMu.Lock() + defer s.alertMu.Unlock() + for _, id := range ruleIDs { + delete(s.alertSentAt, id) + } +} + +// loadEnabledAlertRules 载入启用中的规则;失败时返回空集并记日志(解析主流程照常)。 +func (s *LogEventService) loadEnabledAlertRules(ctx context.Context) []model.AlertRule { + var rules []model.AlertRule + err := s.db.WithContext(ctx).Where("enabled = ?", true).Order("id").Find(&rules).Error + if err != nil { + log.Printf("load alert rules: %v", err) + return nil + } + return rules +} + +// cleanupAlertHits 删除保留期外的命中记录(随 cleanupOnce 周期执行)。 +func (s *LogEventService) cleanupAlertHits(ctx context.Context) { + cutoff := time.Now().Add(-alertHitRetention) + if err := s.db.WithContext(ctx).Where("hit_at < ?", cutoff).Delete(&model.AlertRuleHit{}).Error; err != nil { + log.Printf("cleanup alert hits: %v", err) + } +} diff --git a/internal/service/alertrule_test.go b/internal/service/alertrule_test.go new file mode 100644 index 0000000..dcbe5d3 --- /dev/null +++ b/internal/service/alertrule_test.go @@ -0,0 +1,345 @@ +package service + +import ( + "context" + "fmt" + "strings" + "testing" + "time" + + "oci-portal/internal/crypto" + "oci-portal/internal/model" +) + +func TestValidateAlertRule(t *testing.T) { + tests := []struct { + name string + rule model.AlertRule + wantErr string + check func(t *testing.T, r model.AlertRule) + }{ + {name: "名称必填", rule: model.AlertRule{Threshold: 1}, wantErr: "名称"}, + {name: "模式非法", rule: model.AlertRule{Name: "r", Threshold: 1, SourceIPMode: "any"}, wantErr: "in/notin"}, + {name: "阈值越界", rule: model.AlertRule{Name: "r", Threshold: 101}, wantErr: "阈值"}, + {name: "阈值>1须带窗口", rule: model.AlertRule{Name: "r", Threshold: 3}, wantErr: "窗口"}, + {name: "IP 非法", rule: model.AlertRule{Name: "r", Threshold: 1, SourceIPs: "300.1.1.1"}, wantErr: "IP"}, + {name: "CIDR 合法", rule: model.AlertRule{Name: "r", Threshold: 1, SourceIPs: "10.0.0.0/8, 1.2.3.4"}, + check: func(t *testing.T, r model.AlertRule) { + if r.SourceIPs != "10.0.0.0/8,1.2.3.4" { + t.Errorf("SourceIPs = %q, 应去空白归一化", r.SourceIPs) + } + if r.SourceIPMode != alertSourceIPIn { + t.Errorf("SourceIPMode = %q, 应默认 in", r.SourceIPMode) + } + }}, + {name: "事件清单归一化", rule: model.AlertRule{Name: "r", Threshold: 1, EventTypes: " TerminateInstance , CreateApiKey ,"}, + check: func(t *testing.T, r model.AlertRule) { + if r.EventTypes != "TerminateInstance,CreateApiKey" { + t.Errorf("EventTypes = %q", r.EventTypes) + } + }}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + rule := tt.rule + err := validateAlertRule(&rule) + if tt.wantErr == "" { + if err != nil { + t.Fatalf("validateAlertRule: %v", err) + } + if tt.check != nil { + tt.check(t, rule) + } + return + } + if err == nil || !strings.Contains(err.Error(), tt.wantErr) { + t.Fatalf("err = %v, want contains %q", err, tt.wantErr) + } + }) + } +} + +func TestRuleHits(t *testing.T) { + base := model.AlertRule{Name: "r", Threshold: 1, SourceIPMode: alertSourceIPIn} + ev := &model.LogEvent{OciConfigID: 7} + parsed := parsedEvent{ + EventType: "com.oraclecloud.ComputeApi.TerminateInstance", + SourceIP: "203.0.113.8", + ResourceName: "web-server-1", + } + tests := []struct { + name string + mod func(r *model.AlertRule) + p *parsedEvent + want bool + }{ + {name: "空条件全命中", mod: func(r *model.AlertRule) {}, want: true}, + {name: "租户匹配", mod: func(r *model.AlertRule) { r.OciConfigID = 7 }, want: true}, + {name: "租户不匹配", mod: func(r *model.AlertRule) { r.OciConfigID = 8 }, want: false}, + {name: "事件短名命中", mod: func(r *model.AlertRule) { r.EventTypes = "LaunchInstance,TerminateInstance" }, want: true}, + {name: "事件不在清单", mod: func(r *model.AlertRule) { r.EventTypes = "CreateUser" }, want: false}, + {name: "资源子串命中", mod: func(r *model.AlertRule) { r.ResourceMatch = "web-" }, want: true}, + {name: "资源不含", mod: func(r *model.AlertRule) { r.ResourceMatch = "db-" }, want: false}, + {name: "IP in 命中 CIDR", mod: func(r *model.AlertRule) { r.SourceIPs = "203.0.113.0/24" }, want: true}, + {name: "IP in 未命中", mod: func(r *model.AlertRule) { r.SourceIPs = "10.0.0.0/8" }, want: false}, + {name: "IP notin 白名单外告警", mod: func(r *model.AlertRule) { + r.SourceIPs, r.SourceIPMode = "10.0.0.0/8", alertSourceIPNotIn + }, want: true}, + {name: "IP notin 白名单内不告警", mod: func(r *model.AlertRule) { + r.SourceIPs, r.SourceIPMode = "203.0.113.8", alertSourceIPNotIn + }, want: false}, + {name: "notin 事件缺 IP 不告警", mod: func(r *model.AlertRule) { + r.SourceIPs, r.SourceIPMode = "10.0.0.0/8", alertSourceIPNotIn + }, p: &parsedEvent{EventType: parsed.EventType}, want: false}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + rule := base + tt.mod(&rule) + p := parsed + if tt.p != nil { + p = *tt.p + } + if got := ruleHits(rule, ev, p); got != tt.want { + t.Errorf("ruleHits = %v, want %v", got, tt.want) + } + }) + } +} + +func TestAlertRuleCRUD(t *testing.T) { + svc, _, _ := newLogEventEnv(t) + ctx := context.Background() + + created, err := svc.CreateAlertRule(ctx, model.AlertRule{Name: "非白名单终止", Enabled: true, Threshold: 1}) + if err != nil { + t.Fatalf("create: %v", err) + } + if created.ID == 0 { + t.Fatal("create 未回填 ID") + } + if _, err := svc.CreateAlertRule(ctx, model.AlertRule{Threshold: 1}); err == nil { + t.Fatal("空名称应校验失败") + } + + created.Enabled = false + created.EventTypes = "TerminateInstance" + updated, err := svc.UpdateAlertRule(ctx, created.ID, created) + if err != nil { + t.Fatalf("update: %v", err) + } + if updated.Enabled || updated.EventTypes != "TerminateInstance" { + t.Fatalf("update 未生效: %+v", updated) + } + + rules, err := svc.ListAlertRules(ctx) + if err != nil || len(rules) != 1 { + t.Fatalf("list = %v, %v", rules, err) + } + if err := svc.DeleteAlertRule(ctx, created.ID); err != nil { + t.Fatalf("delete: %v", err) + } + if rules, _ := svc.ListAlertRules(ctx); len(rules) != 0 { + t.Fatalf("delete 后仍有 %d 条", len(rules)) + } +} + +// auditEventPayload 构造一条含资源与来源 IP 的 CloudEvents 审计消息。 +func auditEventPayload(event, resource, ip string) string { + return fmt.Sprintf(`{"eventType":"com.oraclecloud.ComputeApi.%s","source":"ComputeApi",`+ + `"eventTime":"2026-07-10T08:00:00Z","data":{"resourceName":%q,"identity":{"ipAddress":%q}}}`, + event, resource, ip) +} + +// newAlertNotifyEnv 组装带假 Telegram 通道的告警测试环境。 +func newAlertNotifyEnv(t *testing.T) (*LogEventService, *telegramCapture, func()) { + t.Helper() + svc, db, _ := newLogEventEnv(t) + srv, rec := newFakeTelegram(t, `{"ok":true}`) + cipher, err := crypto.NewCipher("test-data-key") + if err != nil { + t.Fatalf("new cipher: %v", err) + } + settings := NewSettingService(db, cipher) + token := "123456:AAfake" + if err := settings.UpdateTelegram(context.Background(), + UpdateTelegramInput{Enabled: true, BotToken: &token, ChatID: "42"}); err != nil { + t.Fatalf("update telegram: %v", err) + } + n := NewNotifier(settings) + n.base = srv.URL + svc.SetNotifier(n, settings) + return svc, rec, n.Wait +} + +func TestMatchAlertRulesNotify(t *testing.T) { + svc, rec, wait := newAlertNotifyEnv(t) + ctx := context.Background() + _, err := svc.CreateAlertRule(ctx, model.AlertRule{ + Name: "白名单外终止", Enabled: true, Threshold: 1, + EventTypes: "TerminateInstance", SourceIPs: "10.0.0.0/8", SourceIPMode: alertSourceIPNotIn, + }) + if err != nil { + t.Fatalf("create rule: %v", err) + } + // 命中:白名单外 IP;不命中:白名单内 IP + mustIngest(t, svc, "m1", auditEventPayload("TerminateInstance", "web-1", "203.0.113.8")) + mustIngest(t, svc, "m2", auditEventPayload("TerminateInstance", "web-2", "10.1.2.3")) + svc.parseOnce(ctx) + wait() + + alerts := auditAlerts(rec.snapshot()) + joined := strings.Join(alerts, "\n---\n") + if !strings.Contains(joined, "白名单外终止") || !strings.Contains(joined, "web-1") { + t.Fatalf("应收到含规则名与资源的告警,got %q", joined) + } + if strings.Contains(joined, "web-2") { + t.Fatalf("白名单内事件不应告警,got %q", joined) + } +} + +// auditAlerts 过滤出审计告警推送(排除既有 notifyCritical 的云端事件通知)。 +func auditAlerts(texts []string) []string { + var out []string + for _, s := range texts { + if strings.Contains(s, "审计告警") { + out = append(out, s) + } + } + return out +} + +func TestAlertThresholdWindow(t *testing.T) { + svc, rec, wait := newAlertNotifyEnv(t) + ctx := context.Background() + _, err := svc.CreateAlertRule(ctx, model.AlertRule{ + Name: "登录风暴", Enabled: true, Threshold: 3, WindowMinutes: 5, EventTypes: "InteractiveLogin", + }) + if err != nil { + t.Fatalf("create rule: %v", err) + } + for i := 1; i <= 4; i++ { + mustIngest(t, svc, fmt.Sprint("login-", i), + auditEventPayload("InteractiveLogin", "user@x.com", "203.0.113.8")) + } + svc.parseOnce(ctx) + wait() + + alerts := auditAlerts(rec.snapshot()) + if len(alerts) != 1 { + t.Fatalf("窗口内 4 次命中应只告警 1 次(第 3 次触发后冷却),got %d 条: %v", len(alerts), alerts) + } + if !strings.Contains(alerts[0], "3 次") { + t.Errorf("告警文案应含累计次数,got %q", alerts[0]) + } +} + +// TestAlertRuleBadDataDoesNotBlockParse 验证规则表异常不影响解析主流程。 +func TestAlertRuleBadDataDoesNotBlockParse(t *testing.T) { + svc, db, _ := newLogEventEnv(t) + ctx := context.Background() + // 直插一条绕过校验的坏规则(IP 列表非法) + bad := model.AlertRule{Name: "bad", Enabled: true, Threshold: 1, SourceIPs: "not-an-ip"} + if err := db.Create(&bad).Error; err != nil { + t.Fatalf("insert bad rule: %v", err) + } + mustIngest(t, svc, "m1", auditEventPayload("TerminateInstance", "web-1", "1.2.3.4")) + svc.parseOnce(ctx) + + var e model.LogEvent + if err := db.First(&e, "message_id = ?", "m1").Error; err != nil { + t.Fatalf("find event: %v", err) + } + if !e.Processed { + t.Fatal("坏规则不应阻塞事件解析") + } +} + +// mustIngest 落一条回传事件,失败即终止测试。 +func mustIngest(t *testing.T, svc *LogEventService, msgID, payload string) { + t.Helper() + if err := svc.Ingest(context.Background(), 1, msgID, []byte(payload), false); err != nil { + t.Fatalf("ingest %s: %v", msgID, err) + } +} + +// TestCleanupAlertHits 验证过期命中记录随清理删除。 +func TestCleanupAlertHits(t *testing.T) { + svc, db, _ := newLogEventEnv(t) + old := model.AlertRuleHit{RuleID: 1, HitAt: time.Now().Add(-8 * 24 * time.Hour)} + fresh := model.AlertRuleHit{RuleID: 1, HitAt: time.Now()} + if err := db.Create(&old).Error; err != nil { + t.Fatalf("insert: %v", err) + } + if err := db.Create(&fresh).Error; err != nil { + t.Fatalf("insert: %v", err) + } + svc.cleanupAlertHits(context.Background()) + var count int64 + db.Model(&model.AlertRuleHit{}).Count(&count) + if count != 1 { + t.Fatalf("清理后应剩 1 条,got %d", count) + } +} + +func TestClearAlertCooldown(t *testing.T) { + svc := NewLogEventService(nil) + rule1 := model.AlertRule{ID: 1, Threshold: 2, WindowMinutes: 10} + rule2 := model.AlertRule{ID: 2, Threshold: 2, WindowMinutes: 10} + if !svc.alertCooldownPass(rule1) || !svc.alertCooldownPass(rule2) { + t.Fatal("首次命中应通过冷却检查") + } + + svc.ClearAlertCooldown([]uint{rule1.ID}) + if !svc.alertCooldownPass(rule1) { + t.Fatal("已清理规则应重新通过冷却检查") + } + if svc.alertCooldownPass(rule2) { + t.Fatal("未清理规则不应通过冷却检查") + } +} + +func TestRecordAlertHitRejectsMissingRefs(t *testing.T) { + tests := []struct { + name string + deleteRule bool + }{ + {name: "规则已删除", deleteRule: true}, + {name: "事件已删除"}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + testRecordAlertHitMissingRef(t, tt.deleteRule) + }) + } +} + +func testRecordAlertHitMissingRef(t *testing.T, deleteRule bool) { + t.Helper() + svc, db, cfgID := newLogEventEnv(t) + rule := model.AlertRule{Name: "r", OciConfigID: cfgID, Threshold: 2, WindowMinutes: 5} + event := model.LogEvent{OciConfigID: cfgID, MessageID: "m"} + if err := db.Create(&rule).Error; err != nil { + t.Fatalf("create rule: %v", err) + } + if err := db.Create(&event).Error; err != nil { + t.Fatalf("create event: %v", err) + } + var err error + if deleteRule { + err = db.Delete(&rule).Error + } else { + err = db.Delete(&event).Error + } + if err != nil { + t.Fatalf("delete reference: %v", err) + } + if count, ok := svc.recordAlertHit(context.Background(), rule, &event); ok || count != 0 { + t.Fatalf("record missing refs = (%d,%v), want (0,false)", count, ok) + } + var hits int64 + db.Model(&model.AlertRuleHit{}).Count(&hits) + if hits != 0 { + t.Fatalf("orphan alert hits = %d, want 0", hits) + } +} diff --git a/internal/service/audit.go b/internal/service/audit.go index 4dbf04f..7f16b45 100644 --- a/internal/service/audit.go +++ b/internal/service/audit.go @@ -2,17 +2,19 @@ package service import ( "context" + "encoding/base64" "encoding/json" "errors" + "strconv" "time" "oci-portal/internal/oci" ) -// 审计查询时间窗(小时)上下限。 +// 批式懒加载的单批条数上下限与缺省值。 const ( - minAuditHours = 1 - maxAuditHours = 72 + auditLimitDefault = 100 + auditLimitMax = 200 ) // 审计原始事件缓存:约 3KB/条,上限 2000 条 ≈ 6MB;TTL 内详情秒开。 @@ -21,39 +23,32 @@ const ( auditRawTTL = 10 * time.Minute ) -// ErrInvalidAuditHours 表示审计时间窗参数越界,handler 据此返回 400。 -var ErrInvalidAuditHours = errors.New("audit events: hours must be between 1 and 72") - -// ErrInvalidAuditWindow 表示续查的绝对时间窗非法(格式/顺序/跨度),handler 返回 400。 -var ErrInvalidAuditWindow = errors.New("audit events: invalid start/end window") +// ErrInvalidAuditCursor 表示续查游标不可解析(过期格式/被篡改),handler 返回 400。 +var ErrInvalidAuditCursor = errors.New("audit events: invalid cursor, refresh to restart") // ErrAuditEventGone 表示原始事件缓存过期且小窗重查未找回;handler 映射 404。 var ErrAuditEventGone = errors.New("原始事件已不可取回,请刷新列表后重试") -// AuditQuery 是审计查询参数:Hours 为首查(相对当前时刻); -// Start/End(RFC3339)为续查绝对窗,配合 Page 从截断游标断点续翻。 +// AuditQuery 是批式懒加载查询参数:Cursor 为空表示自当前时刻首查, +// 非空则从上次响应的游标位置继续向更早回溯;Limit 为单批目标条数。 type AuditQuery struct { Region string - Hours int - Start string - End string - Page string + Cursor string + Limit int } -// AuditEventsView 是审计查询响应:列表不含 raw(详情接口取回); -// Start/End 回传本次实际使用的绝对窗,截断续查必须原样带回(游标绑定查询参数)。 +// AuditEventsView 是批式查询响应:列表不含 raw(详情接口取回); +// Cursor 供下一批续查原样带回,空且 Exhausted 表示已到 365 天保留期尽头。 type AuditEventsView struct { Items []oci.AuditEvent `json:"items"` - Truncated bool `json:"truncated"` - NextPage string `json:"nextPage,omitempty"` - Start time.Time `json:"start"` - End time.Time `json:"end"` + Cursor string `json:"cursor,omitempty"` + Exhausted bool `json:"exhausted"` } // AuditEvents 实时查询租户 OCI 审计事件,纯透传不入库;region 为空时用配置 // 默认区域。原始事件剥离进进程内缓存,响应体瘦身约 10 倍(调研④方案 A1+B1)。 func (s *OciConfigService) AuditEvents(ctx context.Context, id uint, q AuditQuery) (AuditEventsView, error) { - start, end, err := auditWindow(q) + cur, err := decodeAuditCursor(q.Cursor) if err != nil { return AuditEventsView{}, err } @@ -61,53 +56,73 @@ func (s *OciConfigService) AuditEvents(ctx context.Context, id uint, q AuditQuer if err != nil { return AuditEventsView{}, err } - res, err := s.client.ListAuditEvents(ctx, cred, q.Region, start, end, q.Page) + res, err := s.client.ListAuditEventsBatch(ctx, cred, q.Region, cur, auditLimit(q.Limit)) if err != nil { return AuditEventsView{}, err } - return AuditEventsView{ - Items: s.stripAuditRaw(res.Items), - Truncated: res.Truncated, - NextPage: res.NextPage, - // 与实际请求同粒度(分钟),续查回传时窗口逐字节一致 - Start: start.UTC().Truncate(time.Minute), - End: end.UTC().Truncate(time.Minute), - }, nil + view := AuditEventsView{Items: s.stripAuditRaw(id, res.Items), Exhausted: res.Exhausted} + if res.Cursor != nil { + view.Cursor = encodeAuditCursor(*res.Cursor) + } + return view, nil } -// auditWindow 解析查询窗口:带 start/end/page 走绝对窗校验,否则按 hours 相对窗。 -func auditWindow(q AuditQuery) (time.Time, time.Time, error) { - if q.Start != "" || q.End != "" || q.Page != "" { - start, err1 := time.Parse(time.RFC3339, q.Start) - end, err2 := time.Parse(time.RFC3339, q.End) - if err1 != nil || err2 != nil || !start.Before(end) || - end.Sub(start) > maxAuditHours*time.Hour+time.Minute { - return time.Time{}, time.Time{}, ErrInvalidAuditWindow - } - return start, end, nil +// auditLimit 归一单批条数:缺省 100,上限 200(响应体量与页预算的平衡)。 +func auditLimit(limit int) int { + if limit <= 0 { + return auditLimitDefault } - if q.Hours < minAuditHours || q.Hours > maxAuditHours { - return time.Time{}, time.Time{}, ErrInvalidAuditHours - } - end := time.Now().UTC() - return end.Add(-time.Duration(q.Hours) * time.Hour), end, nil + return min(limit, auditLimitMax) } -// stripAuditRaw 把每条原始事件按 eventId 放进缓存并从列表剥离。 -func (s *OciConfigService) stripAuditRaw(items []oci.AuditEvent) []oci.AuditEvent { +// encodeAuditCursor 把游标序列化为不透明字符串(base64url JSON)。 +// 内容仅时间窗与 OCI 翻页令牌,伪造只影响自己的查询范围,无需签名。 +func encodeAuditCursor(cur oci.AuditCursor) string { + b, _ := json.Marshal(cur) + return base64.RawURLEncoding.EncodeToString(b) +} + +// decodeAuditCursor 解析续查游标;空串返回自当前时刻起的首查游标。 +func decodeAuditCursor(s string) (oci.AuditCursor, error) { + if s == "" { + return oci.NewAuditCursor(time.Now()), nil + } + b, err := base64.RawURLEncoding.DecodeString(s) + if err != nil { + return oci.AuditCursor{}, ErrInvalidAuditCursor + } + var cur oci.AuditCursor + if err := json.Unmarshal(b, &cur); err != nil || cur.Start.IsZero() || !cur.Start.Before(cur.End) { + return oci.AuditCursor{}, ErrInvalidAuditCursor + } + return cur, nil +} + +// auditRawKey 组装租户隔离的原始事件缓存键。 +func auditRawKey(configID uint, eventID string) string { + return strconv.FormatUint(uint64(configID), 10) + ":" + eventID +} + +// stripAuditRaw 把每条原始事件按租户与 eventId 放进缓存并从列表剥离。 +func (s *OciConfigService) stripAuditRaw(configID uint, items []oci.AuditEvent) []oci.AuditEvent { for i := range items { if items[i].EventId != "" && items[i].Raw != nil { - s.auditRaw.Set(items[i].EventId, items[i].Raw, auditRawTTL) + s.auditRaw.Set(auditRawKey(configID, items[i].EventId), items[i].Raw, auditRawTTL) } items[i].Raw = nil } return items } +// InvalidateAuditCache 删除指定租户的全部审计原始事件缓存。 +func (s *OciConfigService) InvalidateAuditCache(configID uint) { + s.auditRaw.DeletePrefix(strconv.FormatUint(uint64(configID), 10) + ":") +} + // AuditEventDetail 取回单条事件的原始 JSON:缓存命中即回;miss 按事件时间 // 所在分钟起 2 分钟小窗重查兜底(调研④方案 A2),仍未命中报 ErrAuditEventGone。 func (s *OciConfigService) AuditEventDetail(ctx context.Context, id uint, region, eventID string, eventTime time.Time) (json.RawMessage, error) { - if raw, ok := s.auditRaw.Get(eventID); ok { + if raw, ok := s.auditRaw.Get(auditRawKey(id, eventID)); ok { return raw.(json.RawMessage), nil } cred, err := s.credentialsByID(ctx, id) @@ -124,7 +139,7 @@ func (s *OciConfigService) AuditEventDetail(ctx context.Context, id uint, region if ev.EventId == "" || ev.Raw == nil { continue } - s.auditRaw.Set(ev.EventId, ev.Raw, auditRawTTL) + s.auditRaw.Set(auditRawKey(id, ev.EventId), ev.Raw, auditRawTTL) if ev.EventId == eventID { found = ev.Raw } diff --git a/internal/service/audit_test.go b/internal/service/audit_test.go index a7284aa..50951be 100644 --- a/internal/service/audit_test.go +++ b/internal/service/audit_test.go @@ -13,115 +13,124 @@ import ( // auditStubClient 覆写审计查询并记录透传参数,其余行为沿用 fakeClient。 type auditStubClient struct { *fakeClient + batch oci.AuditBatchResult result oci.AuditEventsResult gotRegion string + gotCursor oci.AuditCursor + gotLimit int gotStart time.Time gotEnd time.Time - gotPage string calls int } +func (f *auditStubClient) ListAuditEventsBatch(ctx context.Context, cred oci.Credentials, region string, cur oci.AuditCursor, limit int) (oci.AuditBatchResult, error) { + f.calls++ + f.gotRegion, f.gotCursor, f.gotLimit = region, cur, limit + return f.batch, nil +} + func (f *auditStubClient) ListAuditEvents(ctx context.Context, cred oci.Credentials, region string, start, end time.Time, page string) (oci.AuditEventsResult, error) { f.calls++ - f.gotRegion, f.gotStart, f.gotEnd, f.gotPage = region, start, end, page + f.gotRegion, f.gotStart, f.gotEnd = region, start, end return f.result, nil } -func TestAuditEventsHoursValidation(t *testing.T) { - tests := []struct { - name string - hours int - wantErr bool - }{ - {name: "下界 1 小时", hours: 1}, - {name: "默认 24 小时", hours: 24}, - {name: "上界 72 小时", hours: 72}, - {name: "0 越界", hours: 0, wantErr: true}, - {name: "负数越界", hours: -3, wantErr: true}, - {name: "100 越界", hours: 100, wantErr: true}, +func TestAuditCursorCodec(t *testing.T) { + cur := oci.AuditCursor{ + Start: time.Date(2026, 7, 9, 10, 0, 0, 0, time.UTC), + End: time.Date(2026, 7, 10, 10, 0, 0, 0, time.UTC), + Page: "tok-1", + WindowHours: 48, } - client := &auditStubClient{ - fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}}, - result: oci.AuditEventsResult{Items: []oci.AuditEvent{{EventName: "GetInstance"}}}, - } - svc := newTestService(t, client) - cfg := importAliveConfig(t, svc) - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - got, err := svc.AuditEvents(context.Background(), cfg.ID, AuditQuery{Hours: tt.hours}) - if tt.wantErr { - if !errors.Is(err, ErrInvalidAuditHours) { - t.Fatalf("AuditEvents(hours=%d) error = %v, want ErrInvalidAuditHours", tt.hours, err) - } - return - } - if err != nil { - t.Fatalf("AuditEvents(hours=%d): %v", tt.hours, err) - } - if len(got.Items) != 1 { - t.Errorf("items = %d, want 1", len(got.Items)) - } - }) - } -} - -func TestAuditEventsWindowAndRegion(t *testing.T) { - client := &auditStubClient{fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}}} - svc := newTestService(t, client) - cfg := importAliveConfig(t, svc) - - got, err := svc.AuditEvents(context.Background(), cfg.ID, AuditQuery{Region: "ap-tokyo-1", Hours: 6}) + got, err := decodeAuditCursor(encodeAuditCursor(cur)) if err != nil { - t.Fatalf("AuditEvents: %v", err) + t.Fatalf("roundtrip: %v", err) } - if client.calls != 1 { - t.Fatalf("calls = %d, want 1", client.calls) + if !got.Start.Equal(cur.Start) || !got.End.Equal(cur.End) || got.Page != cur.Page || got.WindowHours != cur.WindowHours { + t.Fatalf("roundtrip = %+v, want %+v", got, cur) } - if client.gotRegion != "ap-tokyo-1" { - t.Errorf("region = %q, want %q", client.gotRegion, "ap-tokyo-1") + + // 空串 → 自当前时刻首查:24h 窗、分钟粒度、无窗内游标 + first, err := decodeAuditCursor("") + if err != nil { + t.Fatalf("first cursor: %v", err) } - if window := client.gotEnd.Sub(client.gotStart); window != 6*time.Hour { - t.Errorf("window = %v, want %v", window, 6*time.Hour) + if first.End.Sub(first.Start) != 24*time.Hour || first.Page != "" || first.End.Second() != 0 { + t.Fatalf("first cursor = %+v, 应为 24h 分钟粒度首窗", first) } - // 响应回传分钟粒度的绝对窗,续查据此原样带回 - if got.Start.Second() != 0 || !got.End.After(got.Start) { - t.Errorf("响应窗口 = [%v, %v), 应为分钟粒度且有序", got.Start, got.End) + + bad := []string{"!!!", "bm90LWpzb24", encodeAuditCursor(oci.AuditCursor{})} + for i, s := range bad { + if _, err := decodeAuditCursor(s); !errors.Is(err, ErrInvalidAuditCursor) { + t.Errorf("bad[%d] err = %v, want ErrInvalidAuditCursor", i, err) + } } } -func TestAuditEventsResume(t *testing.T) { +func TestAuditEventsBatchParams(t *testing.T) { + next := oci.AuditCursor{ + Start: time.Date(2026, 7, 8, 10, 0, 0, 0, time.UTC), + End: time.Date(2026, 7, 9, 10, 0, 0, 0, time.UTC), + WindowHours: 24, + } client := &auditStubClient{ fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}}, - result: oci.AuditEventsResult{Items: []oci.AuditEvent{}, Truncated: true, NextPage: "tok-2"}, + batch: oci.AuditBatchResult{ + Items: []oci.AuditEvent{{EventId: "e1", EventName: "GetInstance"}}, + Cursor: &next, + }, } svc := newTestService(t, client) cfg := importAliveConfig(t, svc) ctx := context.Background() - // 续查:绝对窗 + 游标透传;NextPage 原样回传 - got, err := svc.AuditEvents(ctx, cfg.ID, AuditQuery{ - Start: "2026-07-08T10:00:00Z", End: "2026-07-09T10:00:00Z", Page: "tok-1", - }) + // limit 归一:0 → 100;超上限截到 200;region 透传 + cases := []struct { + name string + limit int + wantLimit int + }{ + {"缺省 100", 0, 100}, + {"正常透传", 50, 50}, + {"超限截断", 999, 200}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + got, err := svc.AuditEvents(ctx, cfg.ID, AuditQuery{Region: "ap-tokyo-1", Limit: tc.limit}) + if err != nil { + t.Fatalf("AuditEvents: %v", err) + } + if client.gotLimit != tc.wantLimit || client.gotRegion != "ap-tokyo-1" { + t.Fatalf("limit = %d(want %d), region = %q", client.gotLimit, tc.wantLimit, client.gotRegion) + } + if got.Cursor == "" || got.Exhausted { + t.Fatalf("响应应携带续查游标: %+v", got) + } + }) + } + + // 续查:响应游标原样带回可解析,并透传到 oci 层 + got, err := svc.AuditEvents(ctx, cfg.ID, AuditQuery{}) if err != nil { + t.Fatalf("首查: %v", err) + } + if _, err := svc.AuditEvents(ctx, cfg.ID, AuditQuery{Cursor: got.Cursor}); err != nil { t.Fatalf("续查: %v", err) } - if client.gotPage != "tok-1" || !got.Truncated || got.NextPage != "tok-2" { - t.Errorf("游标透传 page=%q next=%q truncated=%v", client.gotPage, got.NextPage, got.Truncated) + if !client.gotCursor.Start.Equal(next.Start) || client.gotCursor.WindowHours != 24 { + t.Fatalf("续查游标透传 = %+v, want %+v", client.gotCursor, next) } - if !client.gotStart.Equal(time.Date(2026, 7, 8, 10, 0, 0, 0, time.UTC)) { - t.Errorf("续查未用绝对窗: start=%v", client.gotStart) + + // 尽头:Cursor 为 nil → 响应空游标 + exhausted + client.batch = oci.AuditBatchResult{Items: []oci.AuditEvent{}, Exhausted: true} + end, err := svc.AuditEvents(ctx, cfg.ID, AuditQuery{}) + if err != nil || end.Cursor != "" || !end.Exhausted { + t.Fatalf("尽头响应 = %+v, %v", end, err) } - // 非法窗口逐项拒绝 - bad := []AuditQuery{ - {Start: "not-a-time", End: "2026-07-09T10:00:00Z"}, - {Start: "2026-07-09T10:00:00Z", End: "2026-07-08T10:00:00Z"}, // 倒序 - {Start: "2026-07-01T00:00:00Z", End: "2026-07-09T10:00:00Z"}, // 超 72h - {Page: "tok-only"}, // 带游标缺窗口 - } - for i, q := range bad { - if _, err := svc.AuditEvents(ctx, cfg.ID, q); !errors.Is(err, ErrInvalidAuditWindow) { - t.Errorf("bad[%d] err = %v, want ErrInvalidAuditWindow", i, err) - } + + // 非法游标 → ErrInvalidAuditCursor + if _, err := svc.AuditEvents(ctx, cfg.ID, AuditQuery{Cursor: "!!!"}); !errors.Is(err, ErrInvalidAuditCursor) { + t.Fatalf("非法游标 err = %v", err) } } @@ -130,7 +139,7 @@ func TestAuditRawStrippedAndDetail(t *testing.T) { eventTime := time.Date(2026, 7, 9, 10, 30, 40, 0, time.UTC) client := &auditStubClient{ fakeClient: &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}}, - result: oci.AuditEventsResult{Items: []oci.AuditEvent{ + batch: oci.AuditBatchResult{Items: []oci.AuditEvent{ {EventId: "evt-1", EventTime: &eventTime, EventName: "TerminateInstance", Raw: raw}, }}, } @@ -138,7 +147,7 @@ func TestAuditRawStrippedAndDetail(t *testing.T) { cfg := importAliveConfig(t, svc) ctx := context.Background() - got, err := svc.AuditEvents(ctx, cfg.ID, AuditQuery{Hours: 1}) + got, err := svc.AuditEvents(ctx, cfg.ID, AuditQuery{}) if err != nil { t.Fatalf("AuditEvents: %v", err) } @@ -184,3 +193,27 @@ func TestAuditDetailRequeryFallback(t *testing.T) { t.Errorf("未找回 err = %v, want ErrAuditEventGone", err) } } + +func TestAuditRawCacheTenantIsolationAndInvalidation(t *testing.T) { + svc := newTestService(t, &fakeClient{}) + raw1 := json.RawMessage(`{"tenant":1}`) + raw2 := json.RawMessage(`{"tenant":2}`) + svc.stripAuditRaw(1, []oci.AuditEvent{{EventId: "same", Raw: raw1}}) + svc.stripAuditRaw(2, []oci.AuditEvent{{EventId: "same", Raw: raw2}}) + + assertAuditRaw(t, svc, 1, raw1) + assertAuditRaw(t, svc, 2, raw2) + svc.InvalidateAuditCache(1) + if _, ok := svc.auditRaw.Get(auditRawKey(1, "same")); ok { + t.Fatal("tenant 1 raw cache still exists after invalidation") + } + assertAuditRaw(t, svc, 2, raw2) +} + +func assertAuditRaw(t *testing.T, svc *OciConfigService, configID uint, want json.RawMessage) { + t.Helper() + got, ok := svc.auditRaw.Get(auditRawKey(configID, "same")) + if !ok || string(got.(json.RawMessage)) != string(want) { + t.Fatalf("tenant %d raw = %v, want %s", configID, got, want) + } +} diff --git a/internal/service/auth.go b/internal/service/auth.go index de8d2ea..760001b 100644 --- a/internal/service/auth.go +++ b/internal/service/auth.go @@ -28,6 +28,13 @@ var ErrLoginLocked = errors.New("too many failed attempts, try again later") // tokenTTL 是登录令牌有效期。 const tokenTTL = 24 * time.Hour +// authClaims 在标准声明外携带令牌版本;版本落后于账号当前值即失效。 +// 存量令牌无 ver 字段解析为 0,与存量账号的零值版本兼容(升级不强制登出)。 +type authClaims struct { + jwt.RegisteredClaims + Ver uint `json:"ver"` +} + // dummyBcryptHash 是恒定失败的占位哈希("dummy-password"), // 用户不存在分支比对它以对齐耗时,防用户枚举与时序侧信道。 const dummyBcryptHash = "$2a$10$N9qo8uLOickgx2ZMRZoMye3xW1Wq8p1zEIfQpXCXbXyE3xY5C6P6W" @@ -106,7 +113,7 @@ func (s *AuthService) createUser(username, password string) error { // 锁定期内一律 ErrLoginLocked(正确密码同样拒绝);阈值与时长取安全设置。 // 已启用两步验证时:密码通过但缺验证码返回 ErrTotpRequired(不计失败),验证码错误计入守卫。 func (s *AuthService) Login(ctx context.Context, username, password, clientIP, totpCode string) (string, time.Time, error) { - key := clientIP + "|" + username + key := guardKey(clientIP, username) now := time.Now() sec := securityOf(s.settings) lockFor := time.Duration(sec.LoginLockMinutes) * time.Minute @@ -138,7 +145,7 @@ func (s *AuthService) Login(ctx context.Context, username, password, clientIP, t } } s.guard.success(key) - return s.signToken(user.Username) + return s.signToken(user.Username, user.TokenVersion) } // failLogin 记失败;达到阈值转锁定并推送告警(开关 login_lock,缺省开)。 @@ -169,15 +176,18 @@ func (s *AuthService) notifyLock(username, clientIP string, sec SecuritySettings }) } -func (s *AuthService) signToken(username string) (string, time.Time, error) { +func (s *AuthService) signToken(username string, ver uint) (string, time.Time, error) { now := time.Now() expires := now.Add(tokenTTL) - claims := jwt.RegisteredClaims{ - Subject: username, - IssuedAt: jwt.NewNumericDate(now), - ExpiresAt: jwt.NewNumericDate(expires), - // jti:同一秒签发的令牌若无唯一 ID 字节全同,登出一个会连坐全部 - ID: newTokenID(), + claims := authClaims{ + RegisteredClaims: jwt.RegisteredClaims{ + Subject: username, + IssuedAt: jwt.NewNumericDate(now), + ExpiresAt: jwt.NewNumericDate(expires), + // jti:同一秒签发的令牌若无唯一 ID 字节全同,登出一个会连坐全部 + ID: newTokenID(), + }, + Ver: ver, } token, err := jwt.NewWithClaims(jwt.SigningMethodHS256, claims).SignedString(s.jwtSecret) if err != nil { @@ -186,6 +196,34 @@ func (s *AuthService) signToken(username string) (string, time.Time, error) { return token, expires, nil } +// IssueToken 按账号当前令牌版本签发新 JWT;敏感操作递增版本后用它为 +// 操作者重签,避免操作者自身会话中断。 +func (s *AuthService) IssueToken(ctx context.Context, username string) (string, time.Time, error) { + user, err := s.findUser(ctx, username) + if err != nil { + return "", time.Time{}, err + } + return s.signToken(user.Username, user.TokenVersion) +} + +// bumpTokenVersion 原子递增账号令牌版本,使所有已签发令牌立即失效。 +func (s *AuthService) bumpTokenVersion(ctx context.Context, username string) error { + err := s.db.WithContext(ctx).Model(&model.User{}).Where("username = ?", username). + UpdateColumn("token_version", gorm.Expr("token_version + 1")).Error + if err != nil { + return fmt.Errorf("bump token version: %w", err) + } + return nil +} + +// RevokeSessions 撤销账号全部会话(版本递增),并为操作者重签新令牌。 +func (s *AuthService) RevokeSessions(ctx context.Context, username string) (string, time.Time, error) { + if err := s.bumpTokenVersion(ctx, username); err != nil { + return "", time.Time{}, err + } + return s.IssueToken(ctx, username) +} + // newTokenID 生成 128 位随机令牌 ID(crypto/rand 自 Go 1.24 起不会失败)。 func newTokenID() string { b := make([]byte, 16) @@ -193,9 +231,10 @@ func newTokenID() string { return hex.EncodeToString(b) } -// ParseToken 验证 JWT 签名与有效期,返回其中的用户名。 -func (s *AuthService) ParseToken(tokenString string) (string, error) { - claims := &jwt.RegisteredClaims{} +// ParseToken 验证 JWT 签名、有效期与令牌版本,返回其中的用户名。 +// 版本落后于账号当前值(凭据等已变更)按无效处理,不区分具体原因。 +func (s *AuthService) ParseToken(ctx context.Context, tokenString string) (string, error) { + claims := &authClaims{} _, err := jwt.ParseWithClaims(tokenString, claims, func(t *jwt.Token) (any, error) { if _, ok := t.Method.(*jwt.SigningMethodHMAC); !ok { return nil, fmt.Errorf("unexpected signing method %v", t.Header["alg"]) @@ -208,6 +247,13 @@ func (s *AuthService) ParseToken(tokenString string) (string, error) { if _, hit := s.revoked.Get(tokenHash(tokenString)); hit { return "", errors.New("token revoked") } + user, err := s.findUser(ctx, claims.Subject) + if err != nil { + return "", fmt.Errorf("token subject: %w", err) + } + if claims.Ver != user.TokenVersion { + return "", errors.New("token version outdated") + } return claims.Subject, nil } diff --git a/internal/service/auth_test.go b/internal/service/auth_test.go index 3d880cf..3aa9367 100644 --- a/internal/service/auth_test.go +++ b/internal/service/auth_test.go @@ -109,7 +109,7 @@ func TestTokenRoundTrip(t *testing.T) { if expires.IsZero() { t.Error("expires is zero, want future time") } - username, err := auth.ParseToken(token) + username, err := auth.ParseToken(context.Background(), token) if err != nil { t.Fatalf("ParseToken: %v", err) } @@ -129,35 +129,92 @@ func TestParseTokenRejectsForged(t *testing.T) { if err != nil { t.Fatalf("Login: %v", err) } - if _, err := auth.ParseToken(forged); err == nil { + if _, err := auth.ParseToken(context.Background(), forged); err == nil { t.Error("ParseToken(forged): got nil error, want failure") } - if _, err := auth.ParseToken("not.a.token"); err == nil { + if _, err := auth.ParseToken(context.Background(), "not.a.token"); err == nil { t.Error("ParseToken(garbage): got nil error, want failure") } } func TestLogoutRevokesToken(t *testing.T) { auth := newTestAuth(t) - token, _, err := auth.signToken("admin") + // ParseToken 现校验令牌版本,须存在对应账号 + if err := auth.EnsureAdmin("admin", "pass123"); err != nil { + t.Fatalf("EnsureAdmin: %v", err) + } + token, _, err := auth.signToken("admin", 0) if err != nil { t.Fatalf("signToken: %v", err) } - if _, err := auth.ParseToken(token); err != nil { + if _, err := auth.ParseToken(context.Background(), token); err != nil { t.Fatalf("ParseToken before logout: %v", err) } auth.Logout(token) - if _, err := auth.ParseToken(token); err == nil { + if _, err := auth.ParseToken(context.Background(), token); err == nil { t.Error("ParseToken after logout: got nil error, want revoked") } // 幂等:重复登出与无效令牌登出都不应 panic,也不影响其他令牌 auth.Logout(token) auth.Logout("not.a.token") - fresh, _, err := auth.signToken("admin") + fresh, _, err := auth.signToken("admin", 0) if err != nil { t.Fatalf("signToken fresh: %v", err) } - if _, err := auth.ParseToken(fresh); err != nil { + if _, err := auth.ParseToken(context.Background(), fresh); err != nil { t.Errorf("ParseToken(fresh) after revoking old: %v", err) } } + +// TestTokenVersionInvalidatesOldToken 凭据变更递增令牌版本,旧 JWT 立即失效, +// 重签的新令牌可用(审计 S-02 动态复现序列的反向断言)。 +func TestTokenVersionInvalidatesOldToken(t *testing.T) { + auth := newTestAuth(t) + if err := auth.EnsureAdmin("admin", "pass123"); err != nil { + t.Fatalf("EnsureAdmin: %v", err) + } + ctx := context.Background() + old, _, err := auth.Login(ctx, "admin", "pass123", "127.0.0.1", "") + if err != nil { + t.Fatalf("Login: %v", err) + } + finalName, err := auth.UpdateCredentials(ctx, "admin", UpdateCredentialsInput{ + NewPassword: "changed-456", CurrentPassword: "pass123", + }) + if err != nil { + t.Fatalf("UpdateCredentials: %v", err) + } + if _, err := auth.ParseToken(ctx, old); err == nil { + t.Error("旧 token 在凭据变更后仍有效, want 失效") + } + fresh, _, err := auth.IssueToken(ctx, finalName) + if err != nil { + t.Fatalf("IssueToken: %v", err) + } + if name, err := auth.ParseToken(ctx, fresh); err != nil || name != "admin" { + t.Errorf("重签 token 应有效: name=%q err=%v", name, err) + } +} + +// TestRevokeSessions 撤销全部会话:旧 token 失效,返回的新 token 有效。 +func TestRevokeSessions(t *testing.T) { + auth := newTestAuth(t) + if err := auth.EnsureAdmin("admin", "pass123"); err != nil { + t.Fatalf("EnsureAdmin: %v", err) + } + ctx := context.Background() + old, _, err := auth.Login(ctx, "admin", "pass123", "127.0.0.1", "") + if err != nil { + t.Fatalf("Login: %v", err) + } + fresh, _, err := auth.RevokeSessions(ctx, "admin") + if err != nil { + t.Fatalf("RevokeSessions: %v", err) + } + if _, err := auth.ParseToken(ctx, old); err == nil { + t.Error("撤销后旧 token 仍有效, want 失效") + } + if name, err := auth.ParseToken(ctx, fresh); err != nil || name != "admin" { + t.Errorf("撤销后新 token 应有效: name=%q err=%v", name, err) + } +} diff --git a/internal/service/credentials.go b/internal/service/credentials.go index aa302ce..0b1d810 100644 --- a/internal/service/credentials.go +++ b/internal/service/credentials.go @@ -7,6 +7,7 @@ import ( "strings" "golang.org/x/crypto/bcrypt" + "gorm.io/gorm" "oci-portal/internal/model" ) @@ -33,42 +34,55 @@ type UpdateCredentialsInput struct { CurrentPassword string `json:"currentPassword" binding:"required"` } -// UpdateCredentials 修改用户名 / 密码:当前密码必验;改名后旧 JWT 的 -// sub 不再命中账号,前端应强制重新登录。 -func (s *AuthService) UpdateCredentials(ctx context.Context, username string, in UpdateCredentialsInput) error { +// UpdateCredentials 修改用户名 / 密码:当前密码必验;成功后令牌版本递增 +// (全部旧 JWT 立即失效),返回最终用户名供调用方为操作者重签新令牌。 +func (s *AuthService) UpdateCredentials(ctx context.Context, username string, in UpdateCredentialsInput) (string, error) { user, err := s.findUser(ctx, username) if err != nil { - return err + return "", err } if bcrypt.CompareHashAndPassword([]byte(user.PasswordHash), []byte(in.CurrentPassword)) != nil { - return ErrCredentialConfirm + return "", ErrCredentialConfirm } newName := strings.TrimSpace(in.NewUsername) if err := validateCredentialChange(user, newName, in.NewPassword); err != nil { - return err + return "", err } + updates, finalName, err := s.credentialUpdates(ctx, user, newName, in.NewPassword) + if err != nil { + return "", err + } + // 同一条 UPDATE 里递增令牌版本,与凭据变更保持原子 + updates["token_version"] = gorm.Expr("token_version + 1") + if err := s.db.WithContext(ctx).Model(&model.User{}).Where("id = ?", user.ID).Updates(updates).Error; err != nil { + return "", fmt.Errorf("update credentials: %w", err) + } + return finalName, nil +} + +// credentialUpdates 组装凭据变更字段并返回最终用户名。 +func (s *AuthService) credentialUpdates(ctx context.Context, user *model.User, newName, newPassword string) (map[string]any, string, error) { updates := map[string]any{} + finalName := user.Username if newName != "" && newName != user.Username { taken, err := s.usernameTaken(ctx, newName, user.ID) if err != nil { - return err + return nil, "", err } if taken { - return fmt.Errorf("用户名已被占用: %w", ErrCredentialInvalid) + return nil, "", fmt.Errorf("用户名已被占用: %w", ErrCredentialInvalid) } updates["username"] = newName + finalName = newName } - if in.NewPassword != "" { - hash, err := bcrypt.GenerateFromPassword([]byte(in.NewPassword), bcrypt.DefaultCost) + if newPassword != "" { + hash, err := bcrypt.GenerateFromPassword([]byte(newPassword), bcrypt.DefaultCost) if err != nil { - return fmt.Errorf("hash password: %w", err) + return nil, "", fmt.Errorf("hash password: %w", err) } updates["password_hash"] = string(hash) } - if err := s.db.WithContext(ctx).Model(&model.User{}).Where("id = ?", user.ID).Updates(updates).Error; err != nil { - return fmt.Errorf("update credentials: %w", err) - } - return nil + return updates, finalName, nil } // validateCredentialChange 校验改名 / 改密输入;两者均无实际变更时报非法。 @@ -122,7 +136,11 @@ func (s *AuthService) SetPasswordLoginDisabled(ctx context.Context, username str return ErrNeedIdentity } } - return s.settings.SetPasswordLoginDisabled(ctx, disabled) + if err := s.settings.SetPasswordLoginDisabled(ctx, disabled); err != nil { + return err + } + // 登录策略属敏感变更:递增令牌版本,已签发会话全部失效 + return s.bumpTokenVersion(ctx, username) } func (s *AuthService) identityCount(ctx context.Context, userID uint) (int64, error) { diff --git a/internal/service/credentials_test.go b/internal/service/credentials_test.go index edf36e2..42af90e 100644 --- a/internal/service/credentials_test.go +++ b/internal/service/credentials_test.go @@ -26,22 +26,22 @@ func TestUpdateCredentials(t *testing.T) { ctx := context.Background() // 当前密码错误 - err := auth.UpdateCredentials(ctx, "admin", UpdateCredentialsInput{NewPassword: "newpass-123", CurrentPassword: "wrong"}) + _, err := auth.UpdateCredentials(ctx, "admin", UpdateCredentialsInput{NewPassword: "newpass-123", CurrentPassword: "wrong"}) if !errors.Is(err, ErrCredentialConfirm) { t.Fatalf("wrong current password: err = %v, want ErrCredentialConfirm", err) } // 短密码 / 无变更均拒绝 - err = auth.UpdateCredentials(ctx, "admin", UpdateCredentialsInput{NewPassword: "short", CurrentPassword: "pass123"}) + _, err = auth.UpdateCredentials(ctx, "admin", UpdateCredentialsInput{NewPassword: "short", CurrentPassword: "pass123"}) if !errors.Is(err, ErrCredentialInvalid) { t.Fatalf("short password: err = %v, want ErrCredentialInvalid", err) } - err = auth.UpdateCredentials(ctx, "admin", UpdateCredentialsInput{NewUsername: "admin", CurrentPassword: "pass123"}) + _, err = auth.UpdateCredentials(ctx, "admin", UpdateCredentialsInput{NewUsername: "admin", CurrentPassword: "pass123"}) if !errors.Is(err, ErrCredentialInvalid) { t.Fatalf("no-op change: err = %v, want ErrCredentialInvalid", err) } // 同时改名改密 - err = auth.UpdateCredentials(ctx, "admin", UpdateCredentialsInput{ + _, err = auth.UpdateCredentials(ctx, "admin", UpdateCredentialsInput{ NewUsername: "root", NewPassword: "newpass-123", CurrentPassword: "pass123", }) if err != nil { diff --git a/internal/service/federation.go b/internal/service/federation.go index 039c370..df61916 100644 --- a/internal/service/federation.go +++ b/internal/service/federation.go @@ -9,17 +9,17 @@ import ( ) // IdentityProviders 列出域内 SAML 身份提供者。 -func (s *OciConfigService) IdentityProviders(ctx context.Context, id uint) ([]oci.IdentityProviderInfo, error) { +func (s *OciConfigService) IdentityProviders(ctx context.Context, id uint, domainID string) ([]oci.IdentityProviderInfo, error) { cred, homeRegion, err := s.credentialsAndHomeRegion(ctx, id) if err != nil { return nil, err } - return s.client.ListIdentityProviders(ctx, cred, homeRegion) + return s.client.ListIdentityProviders(ctx, cred, homeRegion, domainID) } // CreateIdentityProvider 创建禁用态 SAML IdP,激活另走 activate。 // 映射与 JIT 字段缺省时填控制台默认值(normalizeIdpInput)。 -func (s *OciConfigService) CreateIdentityProvider(ctx context.Context, id uint, in oci.CreateIdpInput) (oci.IdentityProviderInfo, error) { +func (s *OciConfigService) CreateIdentityProvider(ctx context.Context, id uint, domainID string, in oci.CreateIdpInput) (oci.IdentityProviderInfo, error) { if strings.TrimSpace(in.Name) == "" { return oci.IdentityProviderInfo{}, fmt.Errorf("create identity provider: name is required") } @@ -34,7 +34,7 @@ func (s *OciConfigService) CreateIdentityProvider(ctx context.Context, id uint, if err != nil { return oci.IdentityProviderInfo{}, err } - return s.client.CreateSamlIdentityProvider(ctx, cred, homeRegion, normalizeIdpInput(in)) + return s.client.CreateSamlIdentityProvider(ctx, cred, homeRegion, domainID, normalizeIdpInput(in)) } // normalizeIdpInput 填充映射字段的控制台默认值:名称 ID 格式「无」、 @@ -50,30 +50,30 @@ func normalizeIdpInput(in oci.CreateIdpInput) oci.CreateIdpInput { } // SetIdentityProviderEnabled 激活/停用 IdP 并同步登录页显示。 -func (s *OciConfigService) SetIdentityProviderEnabled(ctx context.Context, id uint, idpID string, enabled bool) (oci.IdentityProviderInfo, error) { +func (s *OciConfigService) SetIdentityProviderEnabled(ctx context.Context, id uint, domainID, idpID string, enabled bool) (oci.IdentityProviderInfo, error) { cred, homeRegion, err := s.credentialsAndHomeRegion(ctx, id) if err != nil { return oci.IdentityProviderInfo{}, err } - return s.client.SetIdentityProviderEnabled(ctx, cred, homeRegion, idpID, enabled) + return s.client.SetIdentityProviderEnabled(ctx, cred, homeRegion, domainID, idpID, enabled) } // DeleteIdentityProvider 级联删除 IdP:先删关联的免 MFA 规则,再从登录页 // 移除、停用、删除 IdP 本体。 -func (s *OciConfigService) DeleteIdentityProvider(ctx context.Context, id uint, idpID string) error { +func (s *OciConfigService) DeleteIdentityProvider(ctx context.Context, id uint, domainID, idpID string) error { cred, homeRegion, err := s.credentialsAndHomeRegion(ctx, id) if err != nil { return err } - if err := s.deleteIdpExemptions(ctx, cred, homeRegion, idpID); err != nil { + if err := s.deleteIdpExemptions(ctx, cred, homeRegion, domainID, idpID); err != nil { return err } - return s.client.DeleteIdentityProvider(ctx, cred, homeRegion, idpID) + return s.client.DeleteIdentityProvider(ctx, cred, homeRegion, domainID, idpID) } // deleteIdpExemptions 删除 sign-on 策略中引用目标 IdP 的免 MFA 规则。 -func (s *OciConfigService) deleteIdpExemptions(ctx context.Context, cred oci.Credentials, region, idpID string) error { - rules, err := s.client.ListConsoleSignOnRules(ctx, cred, region) +func (s *OciConfigService) deleteIdpExemptions(ctx context.Context, cred oci.Credentials, region, domainID, idpID string) error { + rules, err := s.client.ListConsoleSignOnRules(ctx, cred, region, domainID) if err != nil { return fmt.Errorf("list sign-on rules before delete idp: %w", err) } @@ -81,7 +81,7 @@ func (s *OciConfigService) deleteIdpExemptions(ctx context.Context, cred oci.Cre if r.BuiltIn || !strings.Contains(r.ConditionValue, idpID) { continue } - if err := s.client.DeleteMfaExemptionRule(ctx, cred, region, r.ID); err != nil { + if err := s.client.DeleteMfaExemptionRule(ctx, cred, region, domainID, r.ID); err != nil { return fmt.Errorf("delete exemption rule %s: %w", r.Name, err) } } @@ -89,25 +89,25 @@ func (s *OciConfigService) deleteIdpExemptions(ctx context.Context, cred oci.Cre } // DomainSamlMetadata 下载域的 SP SAML 元数据 XML(提供给 IdP 侧配置)。 -func (s *OciConfigService) DomainSamlMetadata(ctx context.Context, id uint) ([]byte, error) { +func (s *OciConfigService) DomainSamlMetadata(ctx context.Context, id uint, domainID string) ([]byte, error) { cred, homeRegion, err := s.credentialsAndHomeRegion(ctx, id) if err != nil { return nil, err } - return s.client.DownloadDomainSamlMetadata(ctx, cred, homeRegion) + return s.client.DownloadDomainSamlMetadata(ctx, cred, homeRegion, domainID) } // ConsoleSignOnRules 按优先级列出 OCI Console sign-on 策略的规则。 -func (s *OciConfigService) ConsoleSignOnRules(ctx context.Context, id uint) ([]oci.SignOnRuleInfo, error) { +func (s *OciConfigService) ConsoleSignOnRules(ctx context.Context, id uint, domainID string) ([]oci.SignOnRuleInfo, error) { cred, homeRegion, err := s.credentialsAndHomeRegion(ctx, id) if err != nil { return nil, err } - return s.client.ListConsoleSignOnRules(ctx, cred, homeRegion) + return s.client.ListConsoleSignOnRules(ctx, cred, homeRegion, domainID) } // CreateMfaExemption 为指定 IdP 创建免 MFA 规则并置为最高优先级。 -func (s *OciConfigService) CreateMfaExemption(ctx context.Context, id uint, idpID string) (oci.SignOnRuleInfo, error) { +func (s *OciConfigService) CreateMfaExemption(ctx context.Context, id uint, domainID, idpID string) (oci.SignOnRuleInfo, error) { if strings.TrimSpace(idpID) == "" { return oci.SignOnRuleInfo{}, fmt.Errorf("create mfa exemption: identityProviderId is required") } @@ -115,16 +115,16 @@ func (s *OciConfigService) CreateMfaExemption(ctx context.Context, id uint, idpI if err != nil { return oci.SignOnRuleInfo{}, err } - name, err := s.exemptionRuleName(ctx, cred, homeRegion, idpID) + name, err := s.exemptionRuleName(ctx, cred, homeRegion, domainID, idpID) if err != nil { return oci.SignOnRuleInfo{}, err } - return s.client.CreateMfaExemptionRule(ctx, cred, homeRegion, idpID, name) + return s.client.CreateMfaExemptionRule(ctx, cred, homeRegion, domainID, idpID, name) } // exemptionRuleName 校验 IdP 存在并用其名称生成规则名。 -func (s *OciConfigService) exemptionRuleName(ctx context.Context, cred oci.Credentials, region, idpID string) (string, error) { - idps, err := s.client.ListIdentityProviders(ctx, cred, region) +func (s *OciConfigService) exemptionRuleName(ctx context.Context, cred oci.Credentials, region, domainID, idpID string) (string, error) { + idps, err := s.client.ListIdentityProviders(ctx, cred, region, domainID) if err != nil { return "", err } @@ -137,10 +137,10 @@ func (s *OciConfigService) exemptionRuleName(ctx context.Context, cred oci.Crede } // DeleteMfaExemption 删除免 MFA 规则并恢复其余规则优先级。 -func (s *OciConfigService) DeleteMfaExemption(ctx context.Context, id uint, ruleID string) error { +func (s *OciConfigService) DeleteMfaExemption(ctx context.Context, id uint, domainID, ruleID string) error { cred, homeRegion, err := s.credentialsAndHomeRegion(ctx, id) if err != nil { return err } - return s.client.DeleteMfaExemptionRule(ctx, cred, homeRegion, ruleID) + return s.client.DeleteMfaExemptionRule(ctx, cred, homeRegion, domainID, ruleID) } diff --git a/internal/service/logevent.go b/internal/service/logevent.go index 8a19917..b66be80 100644 --- a/internal/service/logevent.go +++ b/internal/service/logevent.go @@ -57,6 +57,9 @@ type LogEventService struct { relayPollTick time.Duration // 订阅确认轮询间隔,零值用默认 relayPollTimeout time.Duration // 订阅确认轮询上限,零值用默认 + + alertMu sync.Mutex // 保护告警规则冷却表 + alertSentAt map[uint]time.Time // 规则 ID → 上次告警时刻(阈值型规则冷却) } // NewLogEventService 组装依赖;调用 StartParser / StartCleanup 后台协程后生效。 @@ -80,37 +83,62 @@ type LogWebhookInfo struct { // EnsureSecret 为配置生成(或幂等返回)回传 secret。 func (s *LogEventService) EnsureSecret(ctx context.Context, cfgID uint) (LogWebhookInfo, error) { - if err := s.requireConfig(ctx, cfgID); err != nil { + secret, err := generateWebhookSecret() + if err != nil { return LogWebhookInfo{}, err } - if info, ok, err := s.SecretInfo(ctx, cfgID); err != nil || ok { - return info, err + var info LogWebhookInfo + err = s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + var err error + info, err = ensureSecretTx(tx, cfgID, secret) + return err + }) + if err != nil { + return LogWebhookInfo{}, fmt.Errorf("ensure webhook secret: %w", err) } + return info, nil +} + +func generateWebhookSecret() (string, error) { buf := make([]byte, 32) if _, err := rand.Read(buf); err != nil { - return LogWebhookInfo{}, fmt.Errorf("generate webhook secret: %w", err) + return "", fmt.Errorf("generate webhook secret: %w", err) + } + return hex.EncodeToString(buf), nil +} + +func ensureSecretTx(tx *gorm.DB, cfgID uint, secret string) (LogWebhookInfo, error) { + if err := lockOciConfig(tx, cfgID); err != nil { + return LogWebhookInfo{}, err + } + if info, ok, err := secretInfoTx(tx, cfgID); err != nil || ok { + return info, err } - secret := hex.EncodeToString(buf) st := model.Setting{Key: secretKey(cfgID), Value: secret, UpdatedAt: time.Now()} - if err := s.db.WithContext(ctx).Save(&st).Error; err != nil { + if err := tx.Create(&st).Error; err != nil { return LogWebhookInfo{}, fmt.Errorf("save webhook secret: %w", err) } return webhookInfo(secret, st.UpdatedAt), nil } -// requireConfig 校验配置存在,不存在透传 gorm.ErrRecordNotFound(api 层映射 404)。 -func (s *LogEventService) requireConfig(ctx context.Context, cfgID uint) error { - var cfg model.OciConfig - if err := s.db.WithContext(ctx).Select("id").First(&cfg, cfgID).Error; err != nil { - return fmt.Errorf("find oci config %d: %w", cfgID, err) - } - return nil -} - // SecretInfo 查询配置是否已生成 secret;未生成时 ok 为 false。 func (s *LogEventService) SecretInfo(ctx context.Context, cfgID uint) (LogWebhookInfo, bool, error) { + var info LogWebhookInfo + var ok bool + err := s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + if err := lockOciConfig(tx, cfgID); err != nil { + return err + } + var err error + info, ok, err = secretInfoTx(tx, cfgID) + return err + }) + return info, ok, err +} + +func secretInfoTx(tx *gorm.DB, cfgID uint) (LogWebhookInfo, bool, error) { var st model.Setting - err := s.db.WithContext(ctx).First(&st, "key = ?", secretKey(cfgID)).Error + err := tx.First(&st, "key = ?", secretKey(cfgID)).Error if err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return LogWebhookInfo{}, false, nil @@ -154,19 +182,61 @@ func (s *LogEventService) ResolveSecret(ctx context.Context, secret string) (uin return 0, false } for _, row := range rows { - if subtle.ConstantTimeCompare([]byte(row.Value), []byte(secret)) == 1 { - id, err := strconv.ParseUint(strings.TrimPrefix(row.Key, logWebhookSecretPrefix), 10, 64) - if err != nil { - return 0, false - } - return uint(id), true + if subtle.ConstantTimeCompare([]byte(row.Value), []byte(secret)) != 1 { + continue } + id, ok := webhookConfigID(row.Key) + if !ok || !s.configExists(ctx, id) { + return 0, false + } + return id, true } return 0, false } +// webhookConfigID 从 secret Setting 键解析租户配置 ID。 +func webhookConfigID(key string) (uint, bool) { + id, err := strconv.ParseUint(strings.TrimPrefix(key, logWebhookSecretPrefix), 10, 64) + return uint(id), err == nil && id > 0 +} + +// configExists 确认 secret 对应租户仍存在;数据库异常按认证失败处理。 +func (s *LogEventService) configExists(ctx context.Context, id uint) bool { + var cfg model.OciConfig + err := s.db.WithContext(ctx).Select("id").First(&cfg, id).Error + if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) { + log.Printf("resolve webhook config: %v", err) + } + return err == nil +} + // Ingest 落一条回传事件;MessageID 唯一索引冲突即静默忽略(at-least-once 幂等)。 func (s *LogEventService) Ingest(ctx context.Context, cfgID uint, messageID string, payload []byte, truncated bool) error { + err := s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + if err := lockOciConfig(tx, cfgID); err != nil { + return err + } + return createLogEvent(tx, cfgID, messageID, payload, truncated) + }) + if err != nil { + return fmt.Errorf("ingest log event: %w", err) + } + return nil +} + +// lockOciConfig 与租户删除共用行锁,避免并发清理后写入孤儿数据。 +func lockOciConfig(tx *gorm.DB, cfgID uint) error { + var cfg model.OciConfig + err := tx.Clauses(clause.Locking{Strength: "UPDATE"}). + Select("id").First(&cfg, cfgID).Error + if err != nil { + return fmt.Errorf("lock oci config %d: %w", cfgID, err) + } + return nil +} + +// createLogEvent 在已锁定租户的事务内幂等写入事件。 +func createLogEvent(tx *gorm.DB, cfgID uint, messageID string, payload []byte, truncated bool) error { event := model.LogEvent{ OciConfigID: cfgID, MessageID: messageID, @@ -174,13 +244,9 @@ func (s *LogEventService) Ingest(ctx context.Context, cfgID uint, messageID stri Truncated: truncated, ReceivedAt: time.Now(), } - err := s.db.WithContext(ctx). - Clauses(clause.OnConflict{Columns: []clause.Column{{Name: "message_id"}}, DoNothing: true}). - Create(&event).Error - if err != nil { - return fmt.Errorf("ingest log event: %w", err) - } - return nil + return tx.Clauses(clause.OnConflict{ + Columns: []clause.Column{{Name: "message_id"}}, DoNothing: true, + }).Create(&event).Error } // LogEventQuery 是回传事件查询参数;CfgID 为 0 表示全部租户。 @@ -304,18 +370,44 @@ func (s *LogEventService) parseOnce(ctx context.Context) { log.Printf("log event parse load: %v", err) return } - for i := range events { - e := &events[i] - parsed := parseLogEvent([]byte(e.Payload)) - e.EventType, e.Source, e.SourceIP, e.EventTime = - parsed.EventType, parsed.Source, parsed.SourceIP, parsed.EventTime - e.Processed = true - if err := s.db.WithContext(ctx).Save(e).Error; err != nil { - log.Printf("log event parse save %d: %v", e.ID, err) - return - } - s.notifyCritical(ctx, e, parsed) + if len(events) == 0 { + return } + rules := s.loadEnabledAlertRules(ctx) + for i := range events { + s.processLogEvent(ctx, &events[i], rules) + } +} + +// processLogEvent 仅在条件更新命中原行后触发通知,删除并发胜出时静默跳过。 +func (s *LogEventService) processLogEvent(ctx context.Context, event *model.LogEvent, rules []model.AlertRule) { + parsed := parseLogEvent([]byte(event.Payload)) + if !s.updateParsedEvent(ctx, event, parsed) { + return + } + s.notifyCritical(ctx, event, parsed) + s.matchAlertRules(ctx, rules, event, parsed) +} + +// updateParsedEvent 用条件 UPDATE 禁止 Save 在删除后隐式重建事件。 +func (s *LogEventService) updateParsedEvent(ctx context.Context, event *model.LogEvent, parsed parsedEvent) bool { + updates := map[string]any{ + "event_type": parsed.EventType, "source": parsed.Source, + "source_ip": parsed.SourceIP, "event_time": parsed.EventTime, "processed": true, + } + res := s.db.WithContext(ctx).Model(&model.LogEvent{}). + Where("id = ? AND processed = ?", event.ID, false).Updates(updates) + if res.Error != nil { + log.Printf("log event parse update %d: %v", event.ID, res.Error) + return false + } + if res.RowsAffected != 1 { + return false + } + event.EventType, event.Source, event.SourceIP, event.EventTime = + parsed.EventType, parsed.Source, parsed.SourceIP, parsed.EventTime + event.Processed = true + return true } // onsEnvelope 覆盖 ONS 消息与 CloudEvents 审计事件的常见字段; @@ -443,6 +535,7 @@ func (s *LogEventService) cleanupOnce(ctx context.Context) { if err := s.cleanup(ctx, logEventRetention, logEventMaxRows); err != nil { log.Printf("log event cleanup: %v", err) } + s.cleanupAlertHits(ctx) } // cleanup 先删过期记录,再对超量部分删最旧;阈值参数化便于测试。 diff --git a/internal/service/logevent_test.go b/internal/service/logevent_test.go index 4767267..465823f 100644 --- a/internal/service/logevent_test.go +++ b/internal/service/logevent_test.go @@ -2,7 +2,9 @@ package service import ( "context" + "errors" "fmt" + "path/filepath" "testing" "time" @@ -28,7 +30,8 @@ func newLogEventEnv(t *testing.T) (*LogEventService, *gorm.DB, uint) { t.Fatalf("db handle: %v", err) } sqlDB.SetMaxOpenConns(1) - if err := db.AutoMigrate(&model.Setting{}, &model.OciConfig{}, &model.LogEvent{}); err != nil { + if err := db.AutoMigrate(&model.Setting{}, &model.OciConfig{}, &model.LogEvent{}, + &model.AlertRule{}, &model.AlertRuleHit{}); err != nil { t.Fatalf("auto migrate: %v", err) } cfg := model.OciConfig{Alias: "测试租户"} @@ -98,9 +101,103 @@ func TestEnsureAndResolveSecret(t *testing.T) { } func TestEnsureSecretRejectsUnknownConfig(t *testing.T) { - svc, _, _ := newLogEventEnv(t) - if _, err := svc.EnsureSecret(context.Background(), 9999); err == nil { - t.Fatal("ensure secret for unknown config succeeded, want error") + svc, db, _ := newLogEventEnv(t) + if _, err := svc.EnsureSecret(context.Background(), 9999); !errors.Is(err, gorm.ErrRecordNotFound) { + t.Fatalf("ensure unknown config error = %v, want record not found", err) + } + var count int64 + db.Model(&model.Setting{}).Where("key = ?", secretKey(9999)).Count(&count) + if count != 0 { + t.Fatalf("unknown config secret rows = %d, want 0", count) + } +} + +func TestEnsureSecretConcurrentDeleteNoOrphan(t *testing.T) { + svc, db := newConcurrentSecretEnv(t) + for i := 0; i < 8; i++ { + cfg := model.OciConfig{Alias: fmt.Sprintf("race-%d", i)} + if err := db.Create(&cfg).Error; err != nil { + t.Fatalf("create config: %v", err) + } + ensureErr, deleteErr := runEnsureDeleteRace(svc, db, cfg.ID) + if deleteErr != nil { + t.Fatalf("delete config %d: %v", cfg.ID, deleteErr) + } + if ensureErr != nil && !errors.Is(ensureErr, gorm.ErrRecordNotFound) { + t.Fatalf("ensure config %d: %v", cfg.ID, ensureErr) + } + assertNoSecretSetting(t, db, cfg.ID) + } +} + +func newConcurrentSecretEnv(t *testing.T) (*LogEventService, *gorm.DB) { + t.Helper() + dsn := filepath.Join(t.TempDir(), "secret.db") + "?_txlock=immediate&_pragma=busy_timeout%3d5000" + db, err := gorm.Open(sqlite.Open(dsn), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)}) + if err != nil { + t.Fatalf("open sqlite: %v", err) + } + sqlDB, err := db.DB() + if err != nil { + t.Fatalf("db handle: %v", err) + } + sqlDB.SetMaxOpenConns(4) + t.Cleanup(func() { _ = sqlDB.Close() }) + if err := db.AutoMigrate(&model.Setting{}, &model.OciConfig{}); err != nil { + t.Fatalf("auto migrate: %v", err) + } + return NewLogEventService(db), db +} + +func runEnsureDeleteRace(svc *LogEventService, db *gorm.DB, cfgID uint) (error, error) { + start := make(chan struct{}) + ensureDone, deleteDone := make(chan error, 1), make(chan error, 1) + go func() { + <-start + _, err := svc.EnsureSecret(context.Background(), cfgID) + ensureDone <- err + }() + go func() { + <-start + deleteDone <- deleteConfigAndSecret(db, cfgID) + }() + close(start) + return <-ensureDone, <-deleteDone +} + +func deleteConfigAndSecret(db *gorm.DB, cfgID uint) error { + return db.Transaction(func(tx *gorm.DB) error { + if err := lockOciConfig(tx, cfgID); err != nil { + return err + } + if err := tx.Delete(&model.Setting{}, "key = ?", secretKey(cfgID)).Error; err != nil { + return err + } + return tx.Delete(&model.OciConfig{}, cfgID).Error + }) +} + +func assertNoSecretSetting(t *testing.T, db *gorm.DB, cfgID uint) { + t.Helper() + var count int64 + db.Model(&model.Setting{}).Where("key = ?", secretKey(cfgID)).Count(&count) + if count != 0 { + t.Fatalf("orphan secret rows for config %d = %d", cfgID, count) + } +} + +func TestResolveSecretRejectsOrphanSetting(t *testing.T) { + svc, db, _ := newLogEventEnv(t) + orphan := model.Setting{Key: secretKey(9999), Value: "orphan-secret"} + if err := db.Create(&orphan).Error; err != nil { + t.Fatalf("create orphan secret: %v", err) + } + if id, ok := svc.ResolveSecret(context.Background(), orphan.Value); ok || id != 0 { + t.Fatalf("resolve orphan secret = (%d,%v), want (0,false)", id, ok) + } + info, ok, err := svc.SecretInfo(context.Background(), 9999) + if !errors.Is(err, gorm.ErrRecordNotFound) || ok || info.Secret != "" { + t.Fatalf("orphan secret info = (%q,%v,%v), want empty,false,record not found", info.Secret, ok, err) } } @@ -127,6 +224,19 @@ func TestIngestIdempotent(t *testing.T) { } } +func TestIngestRejectsUnknownConfig(t *testing.T) { + svc, db, _ := newLogEventEnv(t) + err := svc.Ingest(context.Background(), 9999, "orphan", []byte(`{}`), false) + if !errors.Is(err, gorm.ErrRecordNotFound) { + t.Fatalf("ingest unknown config error = %v, want record not found", err) + } + var count int64 + db.Model(&model.LogEvent{}).Where("message_id = ?", "orphan").Count(&count) + if count != 0 { + t.Fatalf("orphan event rows = %d, want 0", count) + } +} + func TestParseLogEvent(t *testing.T) { ts := "2026-07-07T08:00:00Z" tests := []struct { @@ -209,6 +319,40 @@ func TestParseOnceMarksProcessed(t *testing.T) { } } +func TestUpdateParsedEventDoesNotRecreateDeleted(t *testing.T) { + svc, db, cfgID := newLogEventEnv(t) + event := model.LogEvent{OciConfigID: cfgID, MessageID: "deleted", Payload: `{}`} + if err := db.Create(&event).Error; err != nil { + t.Fatalf("create event: %v", err) + } + if err := db.Delete(&model.LogEvent{}, event.ID).Error; err != nil { + t.Fatalf("delete event: %v", err) + } + if svc.updateParsedEvent(context.Background(), &event, parsedEvent{EventType: "x"}) { + t.Fatal("deleted event update reported success") + } + var count int64 + db.Model(&model.LogEvent{}).Where("id = ?", event.ID).Count(&count) + if count != 0 { + t.Fatalf("deleted event was recreated, rows = %d", count) + } +} + +func TestUpdateParsedEventOnlyOnce(t *testing.T) { + svc, db, cfgID := newLogEventEnv(t) + event := model.LogEvent{OciConfigID: cfgID, MessageID: "once", Payload: `{}`} + if err := db.Create(&event).Error; err != nil { + t.Fatalf("create event: %v", err) + } + parsed := parsedEvent{EventType: "x"} + if !svc.updateParsedEvent(context.Background(), &event, parsed) { + t.Fatal("first conditional update should succeed") + } + if svc.updateParsedEvent(context.Background(), &event, parsed) { + t.Fatal("processed event should not update twice") + } +} + func TestLogEventCleanup(t *testing.T) { svc, db, cfgID := newLogEventEnv(t) ctx := context.Background() diff --git a/internal/service/loginguard.go b/internal/service/loginguard.go index 1943e22..00150b2 100644 --- a/internal/service/loginguard.go +++ b/internal/service/loginguard.go @@ -1,10 +1,18 @@ package service import ( + "strings" "sync" "time" ) +// maxGuardEntries 是 failures / lockedAt 各自的条目上限(S-03): +// 攻击者持续提交唯一用户名时先清过期、仍满驱逐最旧,内存有界。 +const maxGuardEntries = 4096 + +// guardUserMax 是计入守卫 key 的用户名字节上限,超长截断防高基数撑爆。 +const guardUserMax = 32 + // loginGuard 按「IP|用户名」双维度做滑动窗口失败计数与锁定; // 窗口与锁定时长同值、阈值由调用方按安全设置传入(探索文档主题三); // 内存实现(单体面板),重启清零可接受。 @@ -21,6 +29,16 @@ func newLoginGuard() *loginGuard { } } +// guardKey 规范化守卫键:用户名去首尾空白并按字节截断, +// 高基数/超长输入不会产生无界的独立条目。 +func guardKey(clientIP, username string) string { + name := strings.TrimSpace(username) + if len(name) > guardUserMax { + name = name[:guardUserMax] + } + return clientIP + "|" + name +} + // locked 判定 key 是否处于锁定期;过期锁惰性清除。 func (g *loginGuard) locked(key string, now time.Time, lockFor time.Duration) bool { g.mu.Lock() @@ -37,9 +55,13 @@ func (g *loginGuard) locked(key string, now time.Time, lockFor time.Duration) bo } // fail 记一次失败并裁剪窗口;达到阈值转入锁定并返回 true(仅触发那一次)。 +// 新建条目前先保证容量(过期清理 + 最旧驱逐),防高基数 key 单向增长。 func (g *loginGuard) fail(key string, now time.Time, limit int, window time.Duration) bool { g.mu.Lock() defer g.mu.Unlock() + if _, exists := g.failures[key]; !exists { + g.ensureCapacity(now, window) + } kept := g.failures[key][:0] for _, t := range g.failures[key] { if now.Sub(t) < window { @@ -49,6 +71,9 @@ func (g *loginGuard) fail(key string, now time.Time, limit int, window time.Dura kept = append(kept, now) if len(kept) >= limit { delete(g.failures, key) + // 转入锁定同样要保证容量:该写入发生在已有 failures 条目的 + // 第 N 次失败,不经过新建条目路径的 ensureCapacity + g.ensureLockCapacity(now, window) g.lockedAt[key] = now return true } @@ -56,6 +81,77 @@ func (g *loginGuard) fail(key string, now time.Time, limit int, window time.Dura return false } +// ensureLockCapacity 在写入 lockedAt 前保证容量:先清已过锁定期的锁, +// 仍满驱逐锁定最早的条目(其剩余锁定时间最短,提前解锁的影响最小)。 +func (g *loginGuard) ensureLockCapacity(now time.Time, lockFor time.Duration) { + if len(g.lockedAt) < maxGuardEntries { + return + } + for k, at := range g.lockedAt { + if now.Sub(at) >= lockFor { + delete(g.lockedAt, k) + } + } + if len(g.lockedAt) >= maxGuardEntries { + g.evictOldestLock() + } +} + +// evictOldestLock 驱逐锁定时间最早的 lockedAt 条目(调用方须持锁)。 +func (g *loginGuard) evictOldestLock() { + var oldestKey string + var oldestAt time.Time + for k, at := range g.lockedAt { + if oldestKey == "" || at.Before(oldestAt) { + oldestKey, oldestAt = k, at + } + } + if oldestKey != "" { + delete(g.lockedAt, oldestKey) + } +} + +// ensureCapacity 在新建 failures 条目前保证容量:先清出窗口外的失败记录 +// 与已过锁定期的锁,仍满则驱逐最后失败时间最早的条目(调用方须持锁)。 +func (g *loginGuard) ensureCapacity(now time.Time, window time.Duration) { + if len(g.failures) < maxGuardEntries && len(g.lockedAt) < maxGuardEntries { + return + } + for k, times := range g.failures { + if len(times) == 0 || now.Sub(times[len(times)-1]) >= window { + delete(g.failures, k) + } + } + for k, at := range g.lockedAt { + if now.Sub(at) >= window { + delete(g.lockedAt, k) + } + } + if len(g.failures) >= maxGuardEntries { + g.evictOldest() + } +} + +// evictOldest 驱逐最后失败时间最早的 failures 条目(调用方须持锁); +// 仅在容量兜底时触发,O(n) 扫描可接受。 +func (g *loginGuard) evictOldest() { + var oldestKey string + var oldestAt time.Time + for k, times := range g.failures { + if len(times) == 0 { + delete(g.failures, k) + continue + } + last := times[len(times)-1] + if oldestKey == "" || last.Before(oldestAt) { + oldestKey, oldestAt = k, last + } + } + if oldestKey != "" { + delete(g.failures, oldestKey) + } +} + // success 清空 key 的失败计数(成功登录重置窗口)。 func (g *loginGuard) success(key string) { g.mu.Lock() diff --git a/internal/service/loginguard_test.go b/internal/service/loginguard_test.go new file mode 100644 index 0000000..5eb476d --- /dev/null +++ b/internal/service/loginguard_test.go @@ -0,0 +1,105 @@ +package service + +import ( + "fmt" + "strings" + "testing" + "time" +) + +// TestGuardKeyNormalization 用户名去空白并按字节截断,高基数超长输入收敛。 +func TestGuardKeyNormalization(t *testing.T) { + tests := []struct { + name string + ip string + user string + want string + }{ + {name: "常规", ip: "1.2.3.4", user: "admin", want: "1.2.3.4|admin"}, + {name: "去空白", ip: "1.2.3.4", user: " admin ", want: "1.2.3.4|admin"}, + {name: "超长截断", ip: "1.2.3.4", user: strings.Repeat("a", 100), want: "1.2.3.4|" + strings.Repeat("a", guardUserMax)}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := guardKey(tt.ip, tt.user); got != tt.want { + t.Errorf("guardKey = %q, want %q", got, tt.want) + } + }) + } +} + +// TestGuardCapacityBounded 高基数唯一 key 持续失败,条目数不超过上限。 +func TestGuardCapacityBounded(t *testing.T) { + g := newLoginGuard() + now := time.Now() + for i := 0; i < maxGuardEntries*2; i++ { + g.fail(fmt.Sprintf("1.2.3.4|user-%d", i), now, 100, time.Hour) + } + if n := len(g.failures); n > maxGuardEntries { + t.Errorf("failures 条目 = %d, want ≤ %d", n, maxGuardEntries) + } +} + +// TestGuardCapacitySweepExpired 容量满时优先清理窗口外条目,而非驱逐活跃条目。 +func TestGuardCapacitySweepExpired(t *testing.T) { + g := newLoginGuard() + base := time.Now() + window := time.Hour + // 填满:全部为「早已过窗」的旧失败 + for i := 0; i < maxGuardEntries; i++ { + g.fail(fmt.Sprintf("old|%d", i), base.Add(-2*window), 100, window) + } + // 新失败触发清理:过期条目被清,新条目正常记录 + g.fail("fresh|user", base, 100, window) + if _, ok := g.failures["fresh|user"]; !ok { + t.Error("新条目未记录") + } + if n := len(g.failures); n > maxGuardEntries { + t.Errorf("清理后条目 = %d, want ≤ %d", n, maxGuardEntries) + } +} + +// TestGuardEvictOldestKeepsNewest 无过期可清时驱逐最旧,最新条目保留。 +func TestGuardEvictOldestKeepsNewest(t *testing.T) { + g := newLoginGuard() + base := time.Now() + window := 24 * time.Hour + for i := 0; i < maxGuardEntries; i++ { + // 均在窗口内,时间递增;user-0 最旧 + g.fail(fmt.Sprintf("ip|user-%d", i), base.Add(time.Duration(i)*time.Second), 100, window) + } + g.fail("ip|newest", base.Add(time.Duration(maxGuardEntries+1)*time.Second), 100, window) + if _, ok := g.failures["ip|newest"]; !ok { + t.Error("最新条目未记录") + } + if _, ok := g.failures["ip|user-0"]; ok { + t.Error("最旧条目未被驱逐") + } +} + +// TestGuardLockCapacityBounded 高基数唯一 key 全部打满阈值转入锁定, +// 锁定期内(无可清理的过期锁)条目数仍不超过上限,最新锁保留、最早锁被驱逐。 +func TestGuardLockCapacityBounded(t *testing.T) { + g := newLoginGuard() + base := time.Now() + window := 15 * time.Minute + total := maxGuardEntries + 100 + for i := 0; i < total; i++ { + key := fmt.Sprintf("1.2.3.4|user-%d", i) + // 每个 key 打满 2 次(阈值)即转入锁定;时间递增,i 越小锁定越早 + at := base.Add(time.Duration(i) * time.Millisecond) + g.fail(key, at, 2, window) + if !g.fail(key, at, 2, window) { + t.Fatalf("key %s 未按阈值转入锁定", key) + } + } + if n := len(g.lockedAt); n > maxGuardEntries { + t.Errorf("lockedAt 条目 = %d, want ≤ %d", n, maxGuardEntries) + } + if _, ok := g.lockedAt[fmt.Sprintf("1.2.3.4|user-%d", total-1)]; !ok { + t.Error("最新锁定条目未保留") + } + if _, ok := g.lockedAt["1.2.3.4|user-0"]; ok { + t.Error("最早锁定条目未被驱逐") + } +} diff --git a/internal/service/notify.go b/internal/service/notify.go index b5935d1..69f985e 100644 --- a/internal/service/notify.go +++ b/internal/service/notify.go @@ -24,12 +24,12 @@ const notifyTextLimit = 3800 // telegramAPIBase 是 Telegram Bot API 官方根地址。 const telegramAPIBase = "https://api.telegram.org" -// Notifier 通过 Telegram Bot API 推送通知;未启用或未配置时静默跳过。 -// 云外 HTTP 调用,不属于 internal/oci 的职责范围。 +// Notifier 向 Telegram 与全部启用的通知渠道(webhook/ntfy/bark/smtp)推送; +// 未启用或未配置的渠道静默跳过。云外 HTTP 调用,不属于 internal/oci 的职责范围。 type Notifier struct { settings *SettingService client *http.Client - base string // API 根地址,测试时注入 httptest 服务 + base string // Telegram API 根地址,测试时注入 httptest 服务 wg sync.WaitGroup } @@ -42,8 +42,28 @@ func NewNotifier(settings *SettingService) *Notifier { } } -// Send 读取配置并发送一条消息;未启用或未配置时直接返回 nil。 +// Send 向全部启用渠道发送一条纯文本消息;无渠道可用时返回 nil。 func (n *Notifier) Send(ctx context.Context, text string) error { + return n.dispatch(ctx, notifyMsg{Title: "oci-portal 通知", Text: text}) +} + +// dispatch 把消息投递到 Telegram 与全部启用渠道;渠道间互不影响,错误合并返回供日志。 +func (n *Notifier) dispatch(ctx context.Context, msg notifyMsg) error { + msg.Text = truncateText(msg.Text, notifyTextLimit) + var errs []error + if err := n.sendTelegram(ctx, msg); err != nil { + errs = append(errs, fmt.Errorf("telegram: %w", err)) + } + for _, chType := range NotifyChannelTypes { + if err := n.sendChannel(ctx, chType, msg); err != nil { + errs = append(errs, fmt.Errorf("%s: %w", chType, err)) + } + } + return errors.Join(errs...) +} + +// sendTelegram 走既有 Telegram 路径;未启用或未配置时静默跳过。 +func (n *Notifier) sendTelegram(ctx context.Context, msg notifyMsg) error { cfg, err := n.settings.TelegramConfig(ctx) if err != nil { return err @@ -51,7 +71,38 @@ func (n *Notifier) Send(ctx context.Context, text string) error { if !cfg.Enabled || cfg.BotToken == "" || cfg.ChatID == "" { return nil } - return n.post(ctx, cfg, text) + return n.postHTML(ctx, cfg, mdToTelegramHTML(msg.Text)) +} + +// sendChannel 向单个新增渠道投递;未启用或关键字段缺失时静默跳过。 +func (n *Notifier) sendChannel(ctx context.Context, chType string, msg notifyMsg) error { + switch chType { + case ChannelWebhook: + cfg, err := n.settings.webhookChannel(ctx) + if err != nil || !cfg.Enabled || cfg.URL == "" { + return err + } + return sendWebhook(ctx, n.client, cfg, msg) + case ChannelNtfy: + cfg, err := n.settings.ntfyChannel(ctx) + if err != nil || !cfg.Enabled || cfg.Topic == "" { + return err + } + return sendNtfy(ctx, n.client, cfg, msg) + case ChannelBark: + cfg, err := n.settings.barkChannel(ctx) + if err != nil || !cfg.Enabled || cfg.DeviceKey == "" { + return err + } + return sendBark(ctx, n.client, cfg, msg) + case ChannelSMTP: + cfg, err := n.settings.smtpChannel(ctx) + if err != nil || !cfg.Enabled || cfg.Host == "" { + return err + } + return sendSMTP(ctx, cfg, msg) + } + return nil } // SendAsync 异步发送并自带超时,失败只记日志,绝不影响调用链路。 @@ -62,13 +113,13 @@ func (n *Notifier) SendAsync(text string) { ctx, cancel := context.WithTimeout(context.Background(), notifySendTimeout) defer cancel() if err := n.Send(ctx, text); err != nil { - log.Printf("telegram notify: %v", err) + log.Printf("notify: %v", err) } }() } // SendTemplateAsync 按 kind 的生效模板(自定义或默认)渲染变量后异步发送; -// 模板支持轻量 Markdown,按 Telegram HTML 语义投递。 +// 模板支持轻量 Markdown,Telegram 按 HTML 语义投递,其余渠道按纯文本投递。 func (n *Notifier) SendTemplateAsync(kind string, vars map[string]string) { n.wg.Add(1) go func() { @@ -76,7 +127,7 @@ func (n *Notifier) SendTemplateAsync(kind string, vars map[string]string) { ctx, cancel := context.WithTimeout(context.Background(), notifySendTimeout) defer cancel() if err := n.sendTemplate(ctx, kind, "", vars); err != nil { - log.Printf("telegram notify %s: %v", kind, err) + log.Printf("notify %s: %v", kind, err) } }() } @@ -88,12 +139,12 @@ func (n *Notifier) SendTemplateTest(ctx context.Context, kind, tplOverride strin if !ok { return fmt.Errorf("未知通知类型 %q", kind) } - cfg, err := n.settings.TelegramConfig(ctx) + enabled, err := n.anyChannelEnabled(ctx) if err != nil { return err } - if !cfg.Enabled || cfg.BotToken == "" || cfg.ChatID == "" { - return fmt.Errorf("Telegram 通知未启用或未配置,请先在「通知方式」保存配置") + if !enabled { + return fmt.Errorf("没有已启用的通知渠道,请先在「通知方式」配置并启用") } vars := map[string]string{} for k, v := range def.Sample { @@ -102,15 +153,29 @@ func (n *Notifier) SendTemplateTest(ctx context.Context, kind, tplOverride strin return n.sendTemplate(ctx, kind, tplOverride, vars) } -// sendTemplate 渲染并发送:取生效模板 → 变量替换 → 截断 → Markdown 转 Telegram HTML。 -func (n *Notifier) sendTemplate(ctx context.Context, kind, tplOverride string, vars map[string]string) error { - cfg, err := n.settings.TelegramConfig(ctx) +// anyChannelEnabled 报告是否存在至少一个已启用且配置齐全的渠道(含 Telegram)。 +func (n *Notifier) anyChannelEnabled(ctx context.Context) (bool, error) { + tg, err := n.settings.TelegramConfig(ctx) if err != nil { - return err + return false, err } - if !cfg.Enabled || cfg.BotToken == "" || cfg.ChatID == "" { - return nil + if tg.Enabled && tg.BotToken != "" && tg.ChatID != "" { + return true, nil } + views, err := n.settings.NotifyChannels(ctx) + if err != nil { + return false, err + } + for _, v := range views { + if v.Enabled { + return true, nil + } + } + return false, nil +} + +// sendTemplate 渲染并分发:取生效模板 → 变量替换 → 投递到全部启用渠道。 +func (n *Notifier) sendTemplate(ctx context.Context, kind, tplOverride string, vars map[string]string) error { tpl := tplOverride if strings.TrimSpace(tpl) == "" { custom, err := n.settings.NotifyTemplate(ctx, kind) @@ -120,7 +185,58 @@ func (n *Notifier) sendTemplate(ctx context.Context, kind, tplOverride string, v tpl = notifyTplText(custom, kind) } text := renderNotifyTemplate(tpl, vars) - return n.postHTML(ctx, cfg, mdToTelegramHTML(truncateText(text, notifyTextLimit))) + return n.dispatch(ctx, notifyMsg{Title: notifyTplLabel(kind), Text: text}) +} + +// SendChannelTest 向指定渠道同步发送测试消息;不看启用开关,便于保存后即时验证。 +func (n *Notifier) SendChannelTest(ctx context.Context, chType string) error { + msg := notifyMsg{Title: "oci-portal 测试消息", Text: "✅ 通知渠道配置成功:" + chType} + switch chType { + case ChannelWebhook: + cfg, err := n.settings.webhookChannel(ctx) + if err != nil { + return err + } + if cfg.URL == "" { + return fmt.Errorf("请先保存 Webhook URL") + } + return sendWebhook(ctx, n.client, cfg, msg) + case ChannelNtfy: + cfg, err := n.settings.ntfyChannel(ctx) + if err != nil { + return err + } + if cfg.Topic == "" { + return fmt.Errorf("请先保存 ntfy topic") + } + return sendNtfy(ctx, n.client, cfg, msg) + } + return n.sendChannelTestRest(ctx, chType, msg) +} + +// sendChannelTestRest 承接 SendChannelTest 的 bark/smtp 分支(拆分控制函数长度)。 +func (n *Notifier) sendChannelTestRest(ctx context.Context, chType string, msg notifyMsg) error { + switch chType { + case ChannelBark: + cfg, err := n.settings.barkChannel(ctx) + if err != nil { + return err + } + if cfg.DeviceKey == "" { + return fmt.Errorf("请先保存 Bark device key") + } + return sendBark(ctx, n.client, cfg, msg) + case ChannelSMTP: + cfg, err := n.settings.smtpChannel(ctx) + if err != nil { + return err + } + if cfg.Host == "" || cfg.From == "" || cfg.To == "" { + return fmt.Errorf("请先保存 SMTP 主机、发件人与收件人") + } + return sendSMTP(ctx, cfg, msg) + } + return ErrUnknownNotifyChannel } // Wait 等待全部在途异步通知发完,供退出前收尾与测试同步。 diff --git a/internal/service/notifychannel.go b/internal/service/notifychannel.go new file mode 100644 index 0000000..0e36680 --- /dev/null +++ b/internal/service/notifychannel.go @@ -0,0 +1,218 @@ +package service + +import ( + "bytes" + "context" + "crypto/tls" + "encoding/json" + "fmt" + "mime" + "net" + "net/http" + "net/smtp" + "net/url" + "strings" +) + +// 通知渠道类型;telegram 走既有独立路径,不纳入此枚举。 +const ( + ChannelWebhook = "webhook" + ChannelNtfy = "ntfy" + ChannelBark = "bark" + ChannelSMTP = "smtp" +) + +// NotifyChannelTypes 是新增渠道的稳定顺序(视图输出与遍历发送共用)。 +var NotifyChannelTypes = []string{ChannelWebhook, ChannelNtfy, ChannelBark, ChannelSMTP} + +// notifyMsg 是渲染完成的渠道无关载荷:Title 取通知类型显示名,Text 为模板渲染文本。 +type notifyMsg struct { + Title string + Text string +} + +// defaultWebhookBody 是通用 Webhook 的缺省 body 模板(兼容 Slack incoming webhook)。 +const defaultWebhookBody = `{"text": "{{title}}\n{{text}}"}` + +// 渠道服务端缺省地址。 +const ( + defaultNtfyServer = "https://ntfy.sh" + defaultBarkServer = "https://api.day.app" +) + +// WebhookChannel 是通用 Webhook 渠道配置;BodyTemplate 空串按缺省模板发送。 +type WebhookChannel struct { + Enabled bool + URL string + BodyTemplate string +} + +// NtfyChannel 是 ntfy 渠道配置;Server 空串用官方 ntfy.sh,Token 可选。 +type NtfyChannel struct { + Enabled bool + Server string + Topic string + Token string +} + +// BarkChannel 是 Bark(iOS)渠道配置;Server 空串用官方 api.day.app。 +type BarkChannel struct { + Enabled bool + Server string + DeviceKey string +} + +// SMTPChannel 是邮件渠道配置;465 端口走隐式 TLS,其余端口 STARTTLS,Username 空则免认证。 +type SMTPChannel struct { + Enabled bool + Host string + Port int + Username string + Password string + From string + To string +} + +// validNotifyURL 校验渠道出站地址:仅接受 http/https 且 host 非空。 +func validNotifyURL(raw string) error { + u, err := url.Parse(raw) + if err != nil || (u.Scheme != "http" && u.Scheme != "https") || u.Host == "" { + return fmt.Errorf("地址须为 http(s):// 开头的合法 URL") + } + return nil +} + +// renderWebhookBody 渲染 body 模板:占位值先做 JSON 字符串转义(去掉包裹引号)再替换, +// 保证换行/引号嵌入 JSON 字符串后仍合法。 +func renderWebhookBody(tpl string, msg notifyMsg) string { + if strings.TrimSpace(tpl) == "" { + tpl = defaultWebhookBody + } + r := strings.NewReplacer("{{title}}", jsonEscape(msg.Title), "{{text}}", jsonEscape(msg.Text)) + return r.Replace(tpl) +} + +// jsonEscape 返回 s 作为 JSON 字符串字面量的内容(不含首尾引号)。 +func jsonEscape(s string) string { + b, _ := json.Marshal(s) + return string(b[1 : len(b)-1]) +} + +// sendWebhook 向自定义 endpoint POST 渲染后的 JSON body。 +func sendWebhook(ctx context.Context, client *http.Client, cfg WebhookChannel, msg notifyMsg) error { + body := renderWebhookBody(cfg.BodyTemplate, msg) + return postJSON(ctx, client, cfg.URL, "", []byte(body)) +} + +// sendNtfy 以 JSON publish 模式向 ntfy 服务端根路径发布(UTF-8 标题无兼容问题)。 +func sendNtfy(ctx context.Context, client *http.Client, cfg NtfyChannel, msg notifyMsg) error { + server := cfg.Server + if server == "" { + server = defaultNtfyServer + } + body, err := json.Marshal(map[string]string{"topic": cfg.Topic, "title": msg.Title, "message": msg.Text}) + if err != nil { + return fmt.Errorf("ntfy send: %w", err) + } + return postJSON(ctx, client, strings.TrimRight(server, "/"), cfg.Token, body) +} + +// sendBark 调用 Bark v2 push 接口。 +func sendBark(ctx context.Context, client *http.Client, cfg BarkChannel, msg notifyMsg) error { + server := cfg.Server + if server == "" { + server = defaultBarkServer + } + body, err := json.Marshal(map[string]string{"title": msg.Title, "body": msg.Text, "device_key": cfg.DeviceKey}) + if err != nil { + return fmt.Errorf("bark send: %w", err) + } + return postJSON(ctx, client, strings.TrimRight(server, "/")+"/push", "", body) +} + +// postJSON 发送 JSON POST;token 非空时附 Bearer 认证,非 2xx 一律视为失败。 +func postJSON(ctx context.Context, client *http.Client, u, token string, body []byte) error { + req, err := http.NewRequestWithContext(ctx, http.MethodPost, u, bytes.NewReader(body)) + if err != nil { + return sanitizeURLError(err) + } + req.Header.Set("Content-Type", "application/json") + if token != "" { + req.Header.Set("Authorization", "Bearer "+token) + } + resp, err := client.Do(req) + if err != nil { + return sanitizeURLError(err) + } + defer resp.Body.Close() + if resp.StatusCode < 200 || resp.StatusCode >= 300 { + return fmt.Errorf("http %d", resp.StatusCode) + } + return nil +} + +// sendSMTP 组装邮件报文并投递;主题按 RFC2047 编码,正文 UTF-8 纯文本。 +func sendSMTP(ctx context.Context, cfg SMTPChannel, msg notifyMsg) error { + var b strings.Builder + fmt.Fprintf(&b, "From: %s\r\n", cfg.From) + fmt.Fprintf(&b, "To: %s\r\n", cfg.To) + fmt.Fprintf(&b, "Subject: %s\r\n", mime.QEncoding.Encode("utf-8", msg.Title)) + b.WriteString("MIME-Version: 1.0\r\nContent-Type: text/plain; charset=utf-8\r\n\r\n") + b.WriteString(strings.ReplaceAll(msg.Text, "\n", "\r\n")) + return smtpDeliver(ctx, cfg, []byte(b.String())) +} + +// smtpDeliver 建立 SMTP 会话:465 端口隐式 TLS,其余端口连通后尝试 STARTTLS。 +func smtpDeliver(ctx context.Context, cfg SMTPChannel, mail []byte) error { + addr := net.JoinHostPort(cfg.Host, fmt.Sprint(cfg.Port)) + d := &net.Dialer{} + conn, err := d.DialContext(ctx, "tcp", addr) + if err != nil { + return fmt.Errorf("smtp dial: %w", err) + } + if cfg.Port == 465 { + conn = tls.Client(conn, &tls.Config{ServerName: cfg.Host}) + } + c, err := smtp.NewClient(conn, cfg.Host) + if err != nil { + conn.Close() + return fmt.Errorf("smtp handshake: %w", err) + } + defer c.Close() + return smtpWrite(c, cfg, mail) +} + +// smtpWrite 在既有会话上完成 STARTTLS/认证/收发件与报文写入。 +func smtpWrite(c *smtp.Client, cfg SMTPChannel, mail []byte) error { + if cfg.Port != 465 { + if ok, _ := c.Extension("STARTTLS"); ok { + if err := c.StartTLS(&tls.Config{ServerName: cfg.Host}); err != nil { + return fmt.Errorf("smtp starttls: %w", err) + } + } + } + if cfg.Username != "" { + auth := smtp.PlainAuth("", cfg.Username, cfg.Password, cfg.Host) + if err := c.Auth(auth); err != nil { + return fmt.Errorf("smtp auth: %w", err) + } + } + if err := c.Mail(cfg.From); err != nil { + return fmt.Errorf("smtp from: %w", err) + } + if err := c.Rcpt(cfg.To); err != nil { + return fmt.Errorf("smtp rcpt: %w", err) + } + w, err := c.Data() + if err != nil { + return fmt.Errorf("smtp data: %w", err) + } + if _, err := w.Write(mail); err != nil { + w.Close() + return fmt.Errorf("smtp write: %w", err) + } + if err := w.Close(); err != nil { + return fmt.Errorf("smtp close: %w", err) + } + return c.Quit() +} diff --git a/internal/service/notifychannel_test.go b/internal/service/notifychannel_test.go new file mode 100644 index 0000000..dfd8ca9 --- /dev/null +++ b/internal/service/notifychannel_test.go @@ -0,0 +1,339 @@ +package service + +import ( + "bufio" + "context" + "encoding/json" + "fmt" + "net" + "net/http" + "net/http/httptest" + "strings" + "sync" + "testing" +) + +func TestRenderWebhookBody(t *testing.T) { + msg := notifyMsg{Title: "任务失败", Text: "line1\nline\"2\""} + tests := []struct { + name string + tpl string + want string + }{ + { + name: "空模板用缺省 Slack 形状", + tpl: "", + want: `{"text": "任务失败\nline1\nline\"2\""}`, + }, + { + name: "自定义模板占位替换并转义", + tpl: `{"msgtype":"text","text":{"content":"{{text}}"}}`, + want: `{"msgtype":"text","text":{"content":"line1\nline\"2\""}}`, + }, + { + name: "无占位原样输出", + tpl: `{"fixed":1}`, + want: `{"fixed":1}`, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := renderWebhookBody(tt.tpl, msg) + if got != tt.want { + t.Errorf("renderWebhookBody = %s, want %s", got, tt.want) + } + if !json.Valid([]byte(got)) { + t.Errorf("渲染结果不是合法 JSON: %s", got) + } + }) + } +} + +// jsonCapture 收集一次 HTTP 请求的路径、认证头与 JSON body。 +type jsonCapture struct { + mu sync.Mutex + path string + auth string + fields map[string]string +} + +// newCaptureServer 起假渠道服务端,固定返回 status。 +func newCaptureServer(t *testing.T, status int) (*httptest.Server, *jsonCapture) { + t.Helper() + rec := &jsonCapture{} + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + rec.mu.Lock() + defer rec.mu.Unlock() + rec.path = r.URL.Path + rec.auth = r.Header.Get("Authorization") + rec.fields = map[string]string{} + _ = json.NewDecoder(r.Body).Decode(&rec.fields) + w.WriteHeader(status) + })) + t.Cleanup(srv.Close) + return srv, rec +} + +func TestSendNtfy(t *testing.T) { + srv, rec := newCaptureServer(t, http.StatusOK) + cfg := NtfyChannel{Server: srv.URL, Topic: "oci", Token: "tk-1"} + msg := notifyMsg{Title: "标题", Text: "正文"} + if err := sendNtfy(context.Background(), srv.Client(), cfg, msg); err != nil { + t.Fatalf("sendNtfy: %v", err) + } + if rec.auth != "Bearer tk-1" { + t.Errorf("Authorization = %q, want Bearer tk-1", rec.auth) + } + want := map[string]string{"topic": "oci", "title": "标题", "message": "正文"} + for k, v := range want { + if rec.fields[k] != v { + t.Errorf("body[%s] = %q, want %q", k, rec.fields[k], v) + } + } +} + +func TestSendBark(t *testing.T) { + srv, rec := newCaptureServer(t, http.StatusOK) + cfg := BarkChannel{Server: srv.URL, DeviceKey: "dk-1"} + if err := sendBark(context.Background(), srv.Client(), cfg, notifyMsg{Title: "标题", Text: "正文"}); err != nil { + t.Fatalf("sendBark: %v", err) + } + if rec.path != "/push" { + t.Errorf("path = %q, want /push", rec.path) + } + want := map[string]string{"title": "标题", "body": "正文", "device_key": "dk-1"} + for k, v := range want { + if rec.fields[k] != v { + t.Errorf("body[%s] = %q, want %q", k, rec.fields[k], v) + } + } +} + +func TestPostJSONNon2xx(t *testing.T) { + srv, _ := newCaptureServer(t, http.StatusBadGateway) + err := postJSON(context.Background(), srv.Client(), srv.URL, "", []byte(`{}`)) + if err == nil || !strings.Contains(err.Error(), "502") { + t.Fatalf("postJSON err = %v, want contains 502", err) + } +} + +func TestValidNotifyURL(t *testing.T) { + tests := []struct { + raw string + wantErr bool + }{ + {"https://example.com/hook", false}, + {"http://10.0.0.1:8080/x", false}, + {"ftp://example.com", true}, + {"not-a-url", true}, + {"", true}, + } + for _, tt := range tests { + if err := validNotifyURL(tt.raw); (err != nil) != tt.wantErr { + t.Errorf("validNotifyURL(%q) err = %v, wantErr %v", tt.raw, err, tt.wantErr) + } + } +} + +// fakeSMTP 起一个最小 SMTP 服务端,按标准会话应答并收集 DATA 报文。 +func fakeSMTP(t *testing.T) (addr string, got *strings.Builder) { + t.Helper() + ln, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatalf("listen: %v", err) + } + t.Cleanup(func() { ln.Close() }) + got = &strings.Builder{} + go func() { + conn, err := ln.Accept() + if err != nil { + return + } + defer conn.Close() + fmt.Fprint(conn, "220 fake ESMTP\r\n") + r := bufio.NewReader(conn) + for { + line, err := r.ReadString('\n') + if err != nil { + return + } + cmd := strings.ToUpper(strings.TrimSpace(line)) + switch { + case strings.HasPrefix(cmd, "EHLO"), strings.HasPrefix(cmd, "HELO"): + fmt.Fprint(conn, "250-fake\r\n250 AUTH PLAIN\r\n") + case strings.HasPrefix(cmd, "AUTH"): + fmt.Fprint(conn, "235 ok\r\n") + case strings.HasPrefix(cmd, "MAIL"), strings.HasPrefix(cmd, "RCPT"): + fmt.Fprint(conn, "250 ok\r\n") + case cmd == "DATA": + fmt.Fprint(conn, "354 go\r\n") + readSMTPData(r, got) + fmt.Fprint(conn, "250 ok\r\n") + case cmd == "QUIT": + fmt.Fprint(conn, "221 bye\r\n") + return + default: + fmt.Fprint(conn, "250 ok\r\n") + } + } + }() + return ln.Addr().String(), got +} + +// readSMTPData 读取 DATA 段直到单独一行 "."。 +func readSMTPData(r *bufio.Reader, got *strings.Builder) { + for { + line, err := r.ReadString('\n') + if err != nil || strings.TrimRight(line, "\r\n") == "." { + return + } + got.WriteString(line) + } +} + +func TestSendSMTP(t *testing.T) { + addr, got := fakeSMTP(t) + host, portStr, _ := net.SplitHostPort(addr) + var port int + fmt.Sscanf(portStr, "%d", &port) + cfg := SMTPChannel{ + Host: host, Port: port, + Username: "u", Password: "p", + From: "a@x.com", To: "b@y.com", + } + err := sendSMTP(context.Background(), cfg, notifyMsg{Title: "任务失败", Text: "第一行\n第二行"}) + if err != nil { + t.Fatalf("sendSMTP: %v", err) + } + mail := got.String() + for _, want := range []string{"From: a@x.com", "To: b@y.com", "Subject: =?utf-8?q?", "第一行\r\n第二行"} { + if !strings.Contains(mail, want) { + t.Errorf("mail 缺少 %q:\n%s", want, mail) + } + } +} + +// TestDispatchFanout 验证多渠道并发投递:一个渠道失败不影响其他渠道送达。 +func TestDispatchFanout(t *testing.T) { + okSrv, okRec := newCaptureServer(t, http.StatusOK) + badSrv, badRec := newCaptureServer(t, http.StatusInternalServerError) + settings, _ := newSettingEnv(t) + ctx := context.Background() + mustUpdateChannel(t, settings, ChannelWebhook, UpdateNotifyChannelInput{Enabled: true, URL: badSrv.URL}) + mustUpdateChannel(t, settings, ChannelNtfy, UpdateNotifyChannelInput{Enabled: true, Server: okSrv.URL, Topic: "oci"}) + n := NewNotifier(settings) + + err := n.dispatch(ctx, notifyMsg{Title: "T", Text: "body"}) + if err == nil || !strings.Contains(err.Error(), "webhook") { + t.Fatalf("dispatch err = %v, want webhook 失败", err) + } + if badRec.fields == nil { + t.Error("webhook 渠道未收到请求") + } + if okRec.fields["message"] != "body" { + t.Errorf("ntfy 渠道未正常送达: %v", okRec.fields) + } +} + +// TestDispatchSkipsDisabled 验证未启用/未配置的渠道静默跳过且不报错。 +func TestDispatchSkipsDisabled(t *testing.T) { + srv, rec := newCaptureServer(t, http.StatusOK) + settings, _ := newSettingEnv(t) + mustUpdateChannel(t, settings, ChannelWebhook, UpdateNotifyChannelInput{Enabled: false, URL: srv.URL}) + n := NewNotifier(settings) + if err := n.dispatch(context.Background(), notifyMsg{Title: "T", Text: "x"}); err != nil { + t.Fatalf("dispatch: %v", err) + } + if rec.fields != nil { + t.Error("未启用渠道不应收到请求") + } +} + +// mustUpdateChannel 保存渠道配置,失败即终止测试。 +func mustUpdateChannel(t *testing.T, s *SettingService, chType string, in UpdateNotifyChannelInput) { + t.Helper() + if err := s.UpdateNotifyChannel(context.Background(), chType, in); err != nil { + t.Fatalf("update channel %s: %v", chType, err) + } +} + +func TestUpdateNotifyChannelValidation(t *testing.T) { + secret := "sec-123456" + tests := []struct { + name string + chType string + in UpdateNotifyChannelInput + wantErr string + }{ + {name: "未知类型", chType: "pigeon", wantErr: "未知通知渠道"}, + {name: "启用 webhook 缺 URL", chType: ChannelWebhook, in: UpdateNotifyChannelInput{Enabled: true}, wantErr: "URL"}, + {name: "webhook 非法协议", chType: ChannelWebhook, in: UpdateNotifyChannelInput{Enabled: true, URL: "ftp://x"}, wantErr: "http(s)"}, + {name: "启用 ntfy 缺 topic", chType: ChannelNtfy, in: UpdateNotifyChannelInput{Enabled: true}, wantErr: "topic"}, + {name: "启用 smtp 缺主机", chType: ChannelSMTP, in: UpdateNotifyChannelInput{Enabled: true, Port: 587}, wantErr: "SMTP"}, + {name: "smtp 端口越界", chType: ChannelSMTP, in: UpdateNotifyChannelInput{Enabled: true, Host: "h", Port: 70000, From: "a@x", To: "b@y"}, wantErr: "端口"}, + {name: "未启用可存草稿", chType: ChannelNtfy, in: UpdateNotifyChannelInput{Enabled: false}}, + {name: "配置齐全可启用", chType: ChannelBark, in: UpdateNotifyChannelInput{Enabled: true, Secret: &secret}}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + settings, _ := newSettingEnv(t) + err := settings.UpdateNotifyChannel(context.Background(), tt.chType, tt.in) + if tt.wantErr == "" && err != nil { + t.Fatalf("UpdateNotifyChannel: %v", err) + } + if tt.wantErr != "" && (err == nil || !strings.Contains(err.Error(), tt.wantErr)) { + t.Fatalf("err = %v, want contains %q", err, tt.wantErr) + } + }) + } +} + +// TestNotifyChannelSecretLifecycle 验证密文字段的存取语义: +// 保存后视图只回 set/tail,nil 沿用旧值,空串清除。 +func TestNotifyChannelSecretLifecycle(t *testing.T) { + settings, _ := newSettingEnv(t) + ctx := context.Background() + secret := "ntfy-token-9876" + mustUpdateChannel(t, settings, ChannelNtfy, UpdateNotifyChannelInput{Topic: "t", Secret: &secret}) + + view := mustChannelView(t, settings, ChannelNtfy) + if !view.SecretSet || view.SecretTail != "9876" { + t.Fatalf("view = %+v, want SecretSet 且尾号 9876", view) + } + cfg, err := settings.ntfyChannel(ctx) + if err != nil || cfg.Token != secret { + t.Fatalf("ntfyChannel token = %q err = %v, want 原文", cfg.Token, err) + } + + // nil 沿用 + mustUpdateChannel(t, settings, ChannelNtfy, UpdateNotifyChannelInput{Topic: "t2"}) + if cfg, _ := settings.ntfyChannel(ctx); cfg.Token != secret || cfg.Topic != "t2" { + t.Fatalf("nil Secret 应沿用旧 token,got %+v", cfg) + } + + // 空串清除 + empty := "" + mustUpdateChannel(t, settings, ChannelNtfy, UpdateNotifyChannelInput{Topic: "t2", Secret: &empty}) + if cfg, _ := settings.ntfyChannel(ctx); cfg.Token != "" { + t.Fatalf("空串 Secret 应清除,got %q", cfg.Token) + } + if view := mustChannelView(t, settings, ChannelNtfy); view.SecretSet { + t.Fatal("清除后 SecretSet 应为 false") + } +} + +// mustChannelView 取单渠道脱敏视图。 +func mustChannelView(t *testing.T, s *SettingService, chType string) NotifyChannelView { + t.Helper() + views, err := s.NotifyChannels(context.Background()) + if err != nil { + t.Fatalf("NotifyChannels: %v", err) + } + for _, v := range views { + if v.Type == chType { + return v + } + } + t.Fatalf("view %s not found", chType) + return NotifyChannelView{} +} diff --git a/internal/service/notifytpl.go b/internal/service/notifytpl.go index 4dbc74f..0275c79 100644 --- a/internal/service/notifytpl.go +++ b/internal/service/notifytpl.go @@ -21,6 +21,7 @@ var notifyTplOrder = []string{ "task_fail", "task_recover", "snatch_success", "tenant_dead", "task_stop", "login_lock", "model_deprecated", "log_event_instance", "log_event_identity", "log_event_policy", "log_event_region", "log_event_login", + "audit_alert", } var notifyTplDefs = map[string]notifyTplDef{ @@ -38,9 +39,13 @@ var notifyTplDefs = map[string]notifyTplDef{ }, "snatch_success": { Label: "抢机成功", - Default: "🎉 抢机成功:{{task_name}}\n{{message}}", - Vars: []string{"task_name", "message"}, - Sample: map[string]string{"task_name": "示例抢机任务", "message": "已创建 1/1 台实例"}, + Default: "🎉 抢机成功:{{task_name}}\n租户:{{tenant}}\n区域:{{region}}\n类型:{{shape}}\n配置:{{spec}}\n镜像:{{image}}\nIP:{{ip}}\n密码:{{root_password}}", + Vars: []string{"task_name", "tenant", "region", "shape", "spec", "image", "ip", "root_password", "message"}, + Sample: map[string]string{ + "task_name": "示例抢机任务", "tenant": "免费01", "region": "Japan Central (Osaka)", + "shape": "VM.Standard.A1.Flex", "spec": "4C/24G", "image": "Canonical-Ubuntu-24.04-aarch64-2025.05.20-0", + "ip": "129.150.32.10", "root_password": "Xy3#kP9m-2Qw@7Zn", "message": "created 1: ocid1.instance…", + }, }, "tenant_dead": { Label: "租户失联", @@ -96,6 +101,13 @@ var notifyTplDefs = map[string]notifyTplDef{ Vars: []string{"tenant", "event", "resource"}, Sample: map[string]string{"tenant": "免费01", "event": "InteractiveLogin", "resource": "demo@example.com"}, }, + "audit_alert": { + Label: "审计告警", + Default: "🔔 审计告警:{{rule}}\n租户 {{tenant}} · {{event}} {{resource}}\n来源 {{ip}} · 窗口内 {{count}} 次", + Vars: []string{"rule", "tenant", "event", "resource", "ip", "count"}, + Sample: map[string]string{"rule": "非白名单终止实例", "tenant": "免费01", + "event": "TerminateInstance", "resource": "web-server-1", "ip": "203.0.113.8", "count": "1"}, + }, } var notifyVarPattern = regexp.MustCompile(`\{\{(\w+)\}\}`) @@ -161,3 +173,11 @@ func notifyTplText(custom, kind string) string { } return notifyTplDefs[kind].Default } + +// notifyTplLabel 返回 kind 的显示名,供渠道消息标题(ntfy/bark/邮件主题)使用。 +func notifyTplLabel(kind string) string { + if def, ok := notifyTplDefs[kind]; ok { + return def.Label + } + return "oci-portal 通知" +} diff --git a/internal/service/oauth.go b/internal/service/oauth.go index 5a7c5eb..0c9a141 100644 --- a/internal/service/oauth.go +++ b/internal/service/oauth.go @@ -197,7 +197,12 @@ func (o *OAuthService) HandleCallback(ctx context.Context, provider, state, code return "", "", p.mode, err } if p.mode == "bind" { - return "", ident.Display, p.mode, o.bind(ctx, p.username, provider, ident) + if err := o.bind(ctx, p.username, provider, ident); err != nil { + return "", ident.Display, p.mode, err + } + // 绑定属敏感变更:版本递增使旧令牌失效,同时为操作者签新令牌随回跳带回 + token, _, err := o.auth.RevokeSessions(ctx, p.username) + return token, ident.Display, p.mode, err } token, display, err = o.loginByIdentity(ctx, provider, ident) return token, display, p.mode, err @@ -305,7 +310,7 @@ func (o *OAuthService) loginByIdentity(ctx context.Context, provider string, ide if err := o.db.WithContext(ctx).First(&user, row.UserID).Error; err != nil { return "", "", fmt.Errorf("find bound user: %w", err) } - token, _, err := o.auth.signToken(user.Username) + token, _, err := o.auth.signToken(user.Username, user.TokenVersion) return token, ident.Display, err } @@ -345,5 +350,6 @@ func (o *OAuthService) Unbind(ctx context.Context, username string, id uint) err if res.RowsAffected == 0 { return gorm.ErrRecordNotFound } - return nil + // 解绑属敏感变更:递增令牌版本,已签发会话全部失效 + return o.auth.bumpTokenVersion(ctx, username) } diff --git a/internal/service/ociconfig.go b/internal/service/ociconfig.go index e612b66..1717faf 100644 --- a/internal/service/ociconfig.go +++ b/internal/service/ociconfig.go @@ -3,6 +3,7 @@ package service import ( "context" "fmt" + "strings" "time" "gorm.io/gorm" @@ -15,10 +16,12 @@ import ( // OciConfigService 管理 API Key 配置的导入、测活与快照同步。 type OciConfigService struct { - db *gorm.DB - cipher *crypto.Cipher - client oci.Client - // auditRaw 暂存审计原始事件(eventId → raw,TTL 10 分钟), + db *gorm.DB + cipher *crypto.Cipher + client oci.Client + cleanupTasks *TaskService + cleanupEvents *LogEventService + // auditRaw 按 configId:eventId 暂存审计原始事件(TTL 10 分钟), // 列表响应剥离 raw 后详情接口据此秒开;miss 走小窗重查兜底。 auditRaw *cache.Cache } @@ -28,6 +31,12 @@ func NewOciConfigService(db *gorm.DB, cipher *crypto.Cipher, client oci.Client) return &OciConfigService{db: db, cipher: cipher, client: client, auditRaw: cache.New(auditRawMax)} } +// SetTenantCleanupDeps 注入租户删除提交后的任务与告警内存状态同步依赖。 +func (s *OciConfigService) SetTenantCleanupDeps(tasks *TaskService, events *LogEventService) { + s.cleanupTasks = tasks + s.cleanupEvents = events +} + // ImportInput 是导入一份 API Key 的输入: // ConfigINI 与显式字段二选一(ConfigINI 优先),私钥内容必填。 type ImportInput struct { @@ -263,17 +272,18 @@ func (s *OciConfigService) Get(ctx context.Context, id uint) (*model.OciConfig, return &cfg, nil } -// Delete 删除配置及其区域 / 区间缓存。 +// Delete 在单一事务中删除配置及全部租户级本地关联数据。 func (s *OciConfigService) Delete(ctx context.Context, id uint) error { - res := s.db.WithContext(ctx).Delete(&model.OciConfig{}, id) - if res.Error != nil { - return fmt.Errorf("delete oci config %d: %w", id, res.Error) + unlock := func() {} + if s.cleanupTasks != nil { + unlock = s.cleanupTasks.lockTenantCleanup() } - if res.RowsAffected == 0 { - return fmt.Errorf("delete oci config %d: %w", id, gorm.ErrRecordNotFound) + defer unlock() + result, err := s.deleteTenant(ctx, id) + if err != nil { + return err } - s.db.WithContext(ctx).Where("oci_config_id = ?", id).Delete(&model.RegionCache{}) - s.db.WithContext(ctx).Where("oci_config_id = ?", id).Delete(&model.CompartmentCache{}) + s.afterTenantDelete(ctx, result) return nil } @@ -294,12 +304,43 @@ func (s *OciConfigService) refresh(ctx context.Context, cfg *model.OciConfig, cr s.applyProfile(ctx, cfg, cred) s.syncScopeCaches(ctx, cfg, cred) } - if err := s.db.WithContext(ctx).Save(cfg).Error; err != nil { - return fmt.Errorf("save oci config snapshot: %w", err) + s.applySuspension(ctx, cfg, cred) + return s.saveSnapshot(ctx, cfg) +} + +// saveSnapshot 只更新当前版本,避免在途测活把已删租户重新插入。 +func (s *OciConfigService) saveSnapshot(ctx context.Context, cfg *model.OciConfig) error { + res := s.db.WithContext(ctx).Model(&model.OciConfig{}). + Where("id = ? AND updated_at = ?", cfg.ID, cfg.UpdatedAt). + Select("*").Omit("id", "created_at").Updates(cfg) + if res.Error != nil { + return fmt.Errorf("save oci config snapshot: %w", res.Error) + } + if res.RowsAffected != 1 { + return fmt.Errorf("save oci config snapshot: %w", gorm.ErrRecordNotFound) } return nil } +// applySuspension 查询账户能力接口:云端标记暂停的租户覆盖测活结论为 suspended。 +// 该接口对暂停/已终止租户仍可访问(常规 API 此时多被拒,会误判失联);查询失败不改变结论。 +func (s *OciConfigService) applySuspension(ctx context.Context, cfg *model.OciConfig, cred oci.Credentials) { + caps, err := s.client.GetAccountCapabilities(ctx, cred, homeRegionOf(cfg)) + if err != nil || !caps.Suspended { + return + } + cfg.AliveStatus = model.AliveStatusSuspended + cfg.LastError = strings.TrimSpace("账户已被云端暂停 " + capsStatusNote(caps)) +} + +// capsStatusNote 生成暂停详情备注(账户状态,如 terminated)。 +func capsStatusNote(caps oci.AccountCapabilities) string { + if caps.AccountStatus == "" { + return "" + } + return "(account status: " + caps.AccountStatus + ")" +} + // applyProfile 拉取订阅信息;失败只把类别降级为 unknown,不影响测活结论。 func (s *OciConfigService) applyProfile(ctx context.Context, cfg *model.OciConfig, cred oci.Credentials) { profile, err := s.client.FetchAccountProfile(ctx, cred) diff --git a/internal/service/ociconfig_test.go b/internal/service/ociconfig_test.go index d88f2ea..6e39017 100644 --- a/internal/service/ociconfig_test.go +++ b/internal/service/ociconfig_test.go @@ -25,6 +25,8 @@ type fakeClient struct { tenancyErr error profile oci.AccountProfile profileErr error + caps oci.AccountCapabilities + capsErr error regionSubs []oci.RegionSubscription regionSubsErr error @@ -43,6 +45,10 @@ type fakeClient struct { instancesErr error costItems []oci.CostItem costErr error + instance oci.Instance + instanceErr error + image oci.Image + imageErr error subscribedHomeRegion string subscribedKey string @@ -55,6 +61,10 @@ func (f *fakeClient) ValidateKey(ctx context.Context, cred oci.Credentials) (oci return f.tenancy, f.tenancyErr } +func (f *fakeClient) GetAccountCapabilities(ctx context.Context, cred oci.Credentials, region string) (oci.AccountCapabilities, error) { + return f.caps, f.capsErr +} + func (f *fakeClient) FetchAccountProfile(ctx context.Context, cred oci.Credentials) (oci.AccountProfile, error) { return f.profile, f.profileErr } @@ -107,6 +117,14 @@ func (f *fakeClient) ListInstances(ctx context.Context, cred oci.Credentials, re return f.instances, f.instancesErr } +func (f *fakeClient) GetInstance(ctx context.Context, cred oci.Credentials, region, instanceID string) (oci.Instance, error) { + return f.instance, f.instanceErr +} + +func (f *fakeClient) GetImage(ctx context.Context, cred oci.Credentials, region, imageID string) (oci.Image, error) { + return f.image, f.imageErr +} + func (f *fakeClient) SummarizeCosts(ctx context.Context, cred oci.Credentials, q oci.CostQuery) ([]oci.CostItem, error) { return f.costItems, f.costErr } @@ -312,6 +330,7 @@ func TestVerifyMissingConfig(t *testing.T) { func TestDelete(t *testing.T) { svc := newTestService(t, &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}}) + migrateTenantDeleteModels(t, svc.db) cfg, err := svc.Import(context.Background(), trialImportInput()) if err != nil { t.Fatalf("Import: %v", err) @@ -458,3 +477,39 @@ func TestListReturnsSummaries(t *testing.T) { t.Errorf("items[1] alias/group = %q/%q, want 生产号/empty", items[1].Alias, items[1].Group) } } + +// TestVerifySuspendedOverridesStatus 云端暂停标记应覆盖测活结论为 suspended; +// 能力接口失败或未标记时不改变原结论。 +func TestVerifySuspendedOverridesStatus(t *testing.T) { + client := &fakeClient{tenancy: oci.TenancyInfo{Name: "t"}} + svc := newTestService(t, client) + cfg := importAliveConfig(t, svc) + ctx := context.Background() + + // 标记暂停:即使 ValidateKey 成功也应判 suspended + client.caps = oci.AccountCapabilities{Suspended: true, AccountStatus: "terminated"} + got, _, err := svc.Verify(ctx, cfg.ID) + if err != nil { + t.Fatalf("Verify: %v", err) + } + if got.AliveStatus != model.AliveStatusSuspended { + t.Fatalf("AliveStatus = %q, want suspended", got.AliveStatus) + } + if got.LastError == "" { + t.Error("暂停时 LastError 应携带账户状态备注") + } + + // 解除暂停:恢复 alive + client.caps = oci.AccountCapabilities{} + got, _, err = svc.Verify(ctx, cfg.ID) + if err != nil || got.AliveStatus != model.AliveStatusAlive { + t.Fatalf("解除暂停后 = %q, %v; want alive", got.AliveStatus, err) + } + + // 能力接口失败:不影响测活结论 + client.capsErr = errors.New("boom") + got, _, err = svc.Verify(ctx, cfg.ID) + if err != nil || got.AliveStatus != model.AliveStatusAlive { + t.Fatalf("能力接口失败后 = %q, %v; want alive", got.AliveStatus, err) + } +} diff --git a/internal/service/overview.go b/internal/service/overview.go index 0d82916..867f187 100644 --- a/internal/service/overview.go +++ b/internal/service/overview.go @@ -89,7 +89,7 @@ func (s *OciConfigService) overviewTenants(ctx context.Context, out *Overview) e switch cfg.AliveStatus { case model.AliveStatusAlive: out.Tenants.Alive++ - case model.AliveStatusDead: + case model.AliveStatusDead, model.AliveStatusSuspended: out.Tenants.Dead++ } } diff --git a/internal/service/scopecache.go b/internal/service/scopecache.go index af7869e..093ddd2 100644 --- a/internal/service/scopecache.go +++ b/internal/service/scopecache.go @@ -39,6 +39,9 @@ func (s *OciConfigService) saveRegionCache(ctx context.Context, cfgID uint, subs }) } return s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + if err := lockOciConfig(tx, cfgID); err != nil { + return err + } if err := tx.Where("oci_config_id = ?", cfgID).Delete(&model.RegionCache{}).Error; err != nil { return err } @@ -58,6 +61,9 @@ func (s *OciConfigService) saveCompartmentCache(ctx context.Context, cfgID uint, }) } return s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + if err := lockOciConfig(tx, cfgID); err != nil { + return err + } if err := tx.Where("oci_config_id = ?", cfgID).Delete(&model.CompartmentCache{}).Error; err != nil { return err } diff --git a/internal/service/settingnotify.go b/internal/service/settingnotify.go new file mode 100644 index 0000000..c8ae35f --- /dev/null +++ b/internal/service/settingnotify.go @@ -0,0 +1,309 @@ +package service + +import ( + "context" + "fmt" + "strconv" + "strings" +) + +// 通知渠道配置键,前缀 notify_ch__;标注「密文」的键 AES-GCM 加密落库。 +const ( + settingChWebhookEnabled = "notify_ch_webhook_enabled" // "1"/"0" + settingChWebhookURL = "notify_ch_webhook_url" + settingChWebhookBodyTpl = "notify_ch_webhook_body_tpl" + settingChNtfyEnabled = "notify_ch_ntfy_enabled" + settingChNtfyServer = "notify_ch_ntfy_server" + settingChNtfyTopic = "notify_ch_ntfy_topic" + settingChNtfyToken = "notify_ch_ntfy_token" // 密文 + settingChBarkEnabled = "notify_ch_bark_enabled" + settingChBarkServer = "notify_ch_bark_server" + settingChBarkDeviceKey = "notify_ch_bark_device_key" // 密文 + settingChSMTPEnabled = "notify_ch_smtp_enabled" + settingChSMTPHost = "notify_ch_smtp_host" + settingChSMTPPort = "notify_ch_smtp_port" + settingChSMTPUsername = "notify_ch_smtp_username" + settingChSMTPPassword = "notify_ch_smtp_password" // 密文 + settingChSMTPFrom = "notify_ch_smtp_from" + settingChSMTPTo = "notify_ch_smtp_to" +) + +// ErrUnknownNotifyChannel 标记未知渠道类型,api 层映射 400。 +var ErrUnknownNotifyChannel = fmt.Errorf("未知通知渠道类型") + +// NotifyChannelView 是单渠道的脱敏视图:密文字段只回是否已配置与尾 4 位。 +type NotifyChannelView struct { + Type string `json:"type"` + Enabled bool `json:"enabled"` + // webhook + URL string `json:"url,omitempty"` + BodyTemplate string `json:"bodyTemplate,omitempty"` + // ntfy / bark 共用 Server + Server string `json:"server,omitempty"` + Topic string `json:"topic,omitempty"` + // smtp + Host string `json:"host,omitempty"` + Port int `json:"port,omitempty"` + Username string `json:"username,omitempty"` + From string `json:"from,omitempty"` + To string `json:"to,omitempty"` + // 密文字段状态(ntfy token / bark device key / smtp password) + SecretSet bool `json:"secretSet"` + SecretTail string `json:"secretTail,omitempty"` +} + +// UpdateNotifyChannelInput 是单渠道更新入参;Secret 为 nil 沿用已存密文,空串清除。 +type UpdateNotifyChannelInput struct { + Enabled bool + URL string + BodyTemplate string + Server string + Topic string + Host string + Port int + Username string + From string + To string + Secret *string +} + +// notifyChannelKeys 返回渠道的(启用键, 明文键列表, 密文键);未知类型返回启用键空串。 +func notifyChannelKeys(chType string) (enabled string, plain []string, secret string) { + switch chType { + case ChannelWebhook: + return settingChWebhookEnabled, []string{settingChWebhookURL, settingChWebhookBodyTpl}, "" + case ChannelNtfy: + return settingChNtfyEnabled, []string{settingChNtfyServer, settingChNtfyTopic}, settingChNtfyToken + case ChannelBark: + return settingChBarkEnabled, []string{settingChBarkServer}, settingChBarkDeviceKey + case ChannelSMTP: + return settingChSMTPEnabled, []string{settingChSMTPHost, settingChSMTPPort, + settingChSMTPUsername, settingChSMTPFrom, settingChSMTPTo}, settingChSMTPPassword + } + return "", nil, "" +} + +// NotifyChannels 返回全部新增渠道的脱敏视图(不含 telegram,其沿用独立接口)。 +func (s *SettingService) NotifyChannels(ctx context.Context) ([]NotifyChannelView, error) { + out := make([]NotifyChannelView, 0, len(NotifyChannelTypes)) + for _, t := range NotifyChannelTypes { + v, err := s.notifyChannelView(ctx, t) + if err != nil { + return nil, err + } + out = append(out, v) + } + return out, nil +} + +// notifyChannelView 组装单渠道脱敏视图。 +func (s *SettingService) notifyChannelView(ctx context.Context, chType string) (NotifyChannelView, error) { + enabledKey, plain, secretKey := notifyChannelKeys(chType) + vals, err := s.getMany(ctx, append(plain, enabledKey)...) + if err != nil { + return NotifyChannelView{}, err + } + view := NotifyChannelView{Type: chType, Enabled: vals[enabledKey] == "1"} + fillChannelView(&view, chType, vals) + if secretKey == "" { + return view, nil + } + secret, err := s.decryptSetting(ctx, secretKey) + if err != nil { + return NotifyChannelView{}, err + } + view.SecretSet = secret != "" + view.SecretTail = tokenTail(secret) + return view, nil +} + +// fillChannelView 把明文键值填入视图对应字段。 +func fillChannelView(v *NotifyChannelView, chType string, vals map[string]string) { + switch chType { + case ChannelWebhook: + v.URL, v.BodyTemplate = vals[settingChWebhookURL], vals[settingChWebhookBodyTpl] + case ChannelNtfy: + v.Server, v.Topic = vals[settingChNtfyServer], vals[settingChNtfyTopic] + case ChannelBark: + v.Server = vals[settingChBarkServer] + case ChannelSMTP: + v.Host, v.Username = vals[settingChSMTPHost], vals[settingChSMTPUsername] + v.From, v.To = vals[settingChSMTPFrom], vals[settingChSMTPTo] + v.Port, _ = strconv.Atoi(vals[settingChSMTPPort]) + } +} + +// UpdateNotifyChannel 保存单渠道配置;字段先经 validateChannelInput 校验。 +func (s *SettingService) UpdateNotifyChannel(ctx context.Context, chType string, in UpdateNotifyChannelInput) error { + enabledKey, _, secretKey := notifyChannelKeys(chType) + if enabledKey == "" { + return ErrUnknownNotifyChannel + } + if err := validateChannelInput(chType, in); err != nil { + return err + } + pairs := channelPlainPairs(chType, in) + pairs[enabledKey] = boolSetting(in.Enabled) + for k, v := range pairs { + if err := s.set(ctx, k, v); err != nil { + return err + } + } + if secretKey == "" || in.Secret == nil { + return nil + } + return s.encryptSetting(ctx, secretKey, *in.Secret) +} + +// channelPlainPairs 返回渠道明文键与入参值的对应关系。 +func channelPlainPairs(chType string, in UpdateNotifyChannelInput) map[string]string { + switch chType { + case ChannelWebhook: + return map[string]string{settingChWebhookURL: in.URL, settingChWebhookBodyTpl: in.BodyTemplate} + case ChannelNtfy: + return map[string]string{settingChNtfyServer: in.Server, settingChNtfyTopic: in.Topic} + case ChannelBark: + return map[string]string{settingChBarkServer: in.Server} + case ChannelSMTP: + return map[string]string{settingChSMTPHost: in.Host, settingChSMTPPort: strconv.Itoa(in.Port), + settingChSMTPUsername: in.Username, settingChSMTPFrom: in.From, settingChSMTPTo: in.To} + } + return nil +} + +// validateChannelInput 按渠道校验入参;仅在启用时强校验必填,便于「先存草稿后启用」。 +func validateChannelInput(chType string, in UpdateNotifyChannelInput) error { + if !in.Enabled { + return validateChannelURLs(chType, in) + } + switch chType { + case ChannelWebhook: + if strings.TrimSpace(in.URL) == "" { + return fmt.Errorf("启用 Webhook 须填写 URL") + } + case ChannelNtfy: + if strings.TrimSpace(in.Topic) == "" { + return fmt.Errorf("启用 ntfy 须填写 topic") + } + case ChannelSMTP: + if in.Host == "" || in.Port <= 0 || in.Port > 65535 || in.From == "" || in.To == "" { + return fmt.Errorf("启用 SMTP 须填写主机、端口(1-65535)、发件人与收件人") + } + } + return validateChannelURLs(chType, in) +} + +// validateChannelURLs 校验渠道内的 URL 形字段(非空才校验)。 +func validateChannelURLs(chType string, in UpdateNotifyChannelInput) error { + if chType == ChannelWebhook && in.URL != "" { + return validNotifyURL(in.URL) + } + if (chType == ChannelNtfy || chType == ChannelBark) && in.Server != "" { + return validNotifyURL(in.Server) + } + return nil +} + +// boolSetting 把布尔序列化为设置值 "1"/"0"。 +func boolSetting(b bool) string { + if b { + return "1" + } + return "0" +} + +// encryptSetting 加密保存设置值;空串直接落空值表示清除。 +func (s *SettingService) encryptSetting(ctx context.Context, key, value string) error { + if value == "" { + return s.set(ctx, key, "") + } + enc, err := s.cipher.EncryptString(value) + if err != nil { + return fmt.Errorf("encrypt setting %s: %w", key, err) + } + return s.set(ctx, key, enc) +} + +// decryptSetting 读取并解密设置值;未配置返回空串。 +func (s *SettingService) decryptSetting(ctx context.Context, key string) (string, error) { + enc, err := s.get(ctx, key) + if err != nil || enc == "" { + return "", err + } + val, err := s.cipher.DecryptString(enc) + if err != nil { + return "", fmt.Errorf("decrypt setting %s: %w", key, err) + } + return val, nil +} + +// webhookChannel 读取解密后的 Webhook 配置,供发送端消费。 +func (s *SettingService) webhookChannel(ctx context.Context) (WebhookChannel, error) { + vals, err := s.getMany(ctx, settingChWebhookEnabled, settingChWebhookURL, settingChWebhookBodyTpl) + if err != nil { + return WebhookChannel{}, err + } + return WebhookChannel{ + Enabled: vals[settingChWebhookEnabled] == "1", + URL: vals[settingChWebhookURL], + BodyTemplate: vals[settingChWebhookBodyTpl], + }, nil +} + +// ntfyChannel 读取解密后的 ntfy 配置。 +func (s *SettingService) ntfyChannel(ctx context.Context) (NtfyChannel, error) { + vals, err := s.getMany(ctx, settingChNtfyEnabled, settingChNtfyServer, settingChNtfyTopic) + if err != nil { + return NtfyChannel{}, err + } + token, err := s.decryptSetting(ctx, settingChNtfyToken) + if err != nil { + return NtfyChannel{}, err + } + return NtfyChannel{ + Enabled: vals[settingChNtfyEnabled] == "1", + Server: vals[settingChNtfyServer], + Topic: vals[settingChNtfyTopic], + Token: token, + }, nil +} + +// barkChannel 读取解密后的 Bark 配置。 +func (s *SettingService) barkChannel(ctx context.Context) (BarkChannel, error) { + vals, err := s.getMany(ctx, settingChBarkEnabled, settingChBarkServer) + if err != nil { + return BarkChannel{}, err + } + key, err := s.decryptSetting(ctx, settingChBarkDeviceKey) + if err != nil { + return BarkChannel{}, err + } + return BarkChannel{ + Enabled: vals[settingChBarkEnabled] == "1", + Server: vals[settingChBarkServer], + DeviceKey: key, + }, nil +} + +// smtpChannel 读取解密后的 SMTP 配置。 +func (s *SettingService) smtpChannel(ctx context.Context) (SMTPChannel, error) { + vals, err := s.getMany(ctx, settingChSMTPEnabled, settingChSMTPHost, settingChSMTPPort, + settingChSMTPUsername, settingChSMTPFrom, settingChSMTPTo) + if err != nil { + return SMTPChannel{}, err + } + password, err := s.decryptSetting(ctx, settingChSMTPPassword) + if err != nil { + return SMTPChannel{}, err + } + port, _ := strconv.Atoi(vals[settingChSMTPPort]) + return SMTPChannel{ + Enabled: vals[settingChSMTPEnabled] == "1", + Host: vals[settingChSMTPHost], + Port: port, + Username: vals[settingChSMTPUsername], + Password: password, + From: vals[settingChSMTPFrom], + To: vals[settingChSMTPTo], + }, nil +} diff --git a/internal/service/task.go b/internal/service/task.go index a422f8e..0c99b31 100644 --- a/internal/service/task.go +++ b/internal/service/task.go @@ -4,6 +4,7 @@ import ( "context" "encoding/json" "fmt" + "log" "strings" "sync" "time" @@ -22,6 +23,13 @@ const taskRunTimeout = 10 * time.Minute // taskLogKeep 是每个任务保留的执行日志条数。 const taskLogKeep = 100 +// defaultSnatchIPWait 是抢机成功后等待公网 IPv4 分配的总预算(全部实例共享, +// 实例从创建到 VNIC 就绪通常 1~2 分钟);taskRunTimeout 余量内,超时降级不阻塞通知。 +const defaultSnatchIPWait = 3 * time.Minute + +// snatchIPPollInterval 是等待公网 IP 的轮询间隔。 +const snatchIPPollInterval = 10 * time.Second + // TaskService 管理后台任务的存储、cron 调度与执行。 type TaskService struct { db *gorm.DB @@ -31,21 +39,32 @@ type TaskService struct { cron *cron.Cron // aiGateway 供 AI 探测任务执行渠道探测,由 main 装配(可为 nil) aiGateway *AiGatewayService + // snatchIPWait 是抢机成功后等待公网 IP 的预算,测试置 0 跳过等待 + snatchIPWait time.Duration mu sync.Mutex entries map[uint]cron.EntryID + // runMu 让租户删除等待在途执行结束,并阻止新执行读取删除前的任务快照。 + runMu sync.RWMutex + + // snatchVarsMu 保护 snatchVars:runSnatch 抢满时写入成功通知变量, + // execute 组装执行后快照时取走(一次性) + snatchVarsMu sync.Mutex + snatchVars map[uint]map[string]string } // NewTaskService 组装依赖;notifier 传 nil 表示整体关闭通知, // settings 供发送前按事件类型过滤(nil 视为全开)。调用 Start 后开始调度。 func NewTaskService(db *gorm.DB, configs *OciConfigService, notifier *Notifier, settings *SettingService) *TaskService { return &TaskService{ - db: db, - configs: configs, - notifier: notifier, - settings: settings, - cron: cron.New(), - entries: map[uint]cron.EntryID{}, + db: db, + configs: configs, + notifier: notifier, + settings: settings, + cron: cron.New(), + entries: map[uint]cron.EntryID{}, + snatchIPWait: defaultSnatchIPWait, + snatchVars: map[uint]map[string]string{}, } } @@ -371,9 +390,37 @@ func (s *TaskService) unschedule(taskID uint) { } } +// ApplyTenantCleanup 在租户删除事务提交后注销已删除任务并对齐 AI 探测任务。 +func (s *TaskService) ApplyTenantCleanup(ctx context.Context, taskIDs []uint, channelsGone bool) { + for _, id := range taskIDs { + s.unschedule(id) + s.forgetSnatchVars(id) + } + if channelsGone { + s.SyncAiProbeTask(ctx) + } +} + +func (s *TaskService) forgetSnatchVars(taskID uint) { + s.snatchVarsMu.Lock() + delete(s.snatchVars, taskID) + s.snatchVarsMu.Unlock() +} + +func (s *TaskService) lockTenantCleanup() func() { + s.runMu.Lock() + return s.runMu.Unlock +} + // execute 执行一次任务:加载 → 分派 → 更新任务状态并写日志, // 前后状态交给通知判定,只在状态变化时推送。 func (s *TaskService) execute(taskID uint) *model.TaskLog { + s.runMu.RLock() + defer s.runMu.RUnlock() + return s.executeLocked(taskID) +} + +func (s *TaskService) executeLocked(taskID uint) *model.TaskLog { ctx, cancel := context.WithTimeout(context.Background(), taskRunTimeout) defer cancel() task, err := s.GetTask(ctx, taskID) @@ -391,12 +438,40 @@ func (s *TaskService) execute(taskID uint) *model.TaskLog { task.LastError = oci.CompactError(runErr) message = task.LastError } - s.db.Save(task) - cur := taskSnapshot{Name: task.Name, Status: task.Status, LastError: task.LastError, Message: message} + if !s.storeTaskRun(ctx, task) { + return nil + } + cur := taskSnapshot{Name: task.Name, Status: task.Status, LastError: task.LastError, + Message: message, SnatchVars: s.takeSnatchVars(task.ID)} s.notify(notifyEvents(prev, cur)) return s.appendLog(task.ID, runErr == nil, message, time.Since(start)) } +func (s *TaskService) storeTaskRun(ctx context.Context, task *model.Task) bool { + stored, err := s.persistTaskRun(ctx, task) + if err != nil { + log.Printf("persist task run: %v", err) + return false + } + return stored +} + +func (s *TaskService) persistTaskRun(ctx context.Context, task *model.Task) (bool, error) { + updates := map[string]any{ + "status": task.Status, "last_run_at": task.LastRunAt, + "last_error": task.LastError, "run_count": task.RunCount, + } + if task.Type == model.TaskTypeSnatch { + updates["payload"] = task.Payload + } + res := s.db.WithContext(ctx).Model(&model.Task{}). + Where("id = ? AND updated_at = ?", task.ID, task.UpdatedAt).Updates(updates) + if res.Error != nil { + return false, fmt.Errorf("persist task %d run: %w", task.ID, res.Error) + } + return res.RowsAffected == 1, nil +} + // run 按类型分派任务执行。 func (s *TaskService) run(ctx context.Context, task *model.Task) (string, error) { switch task.Type { @@ -650,12 +725,158 @@ func (s *TaskService) runSnatch(ctx context.Context, task *model.Task) (string, writeSnatchPayload(task, &p) return fmt.Sprintf("created %d (%s)%s, %d remaining", len(instances), strings.Join(ids, ","), adNote, remaining), nil } - task.Status = model.TaskStatusSucceeded - writeSnatchPayload(task, &p) - s.unschedule(task.ID) + s.snatchSucceed(ctx, task, &p, in, instances) return fmt.Sprintf("created %d: %s%s", len(instances), strings.Join(ids, ","), adNote), nil } +// snatchSucceed 处理抢满收尾:标记成功、停止调度,并组装成功通知变量暂存, +// 供 execute 发通知时合并进 snatch_success 事件。 +func (s *TaskService) snatchSucceed(ctx context.Context, task *model.Task, p *snatchPayload, in oci.CreateInstanceInput, instances []oci.Instance) { + task.Status = model.TaskStatusSucceeded + writeSnatchPayload(task, p) + s.unschedule(task.ID) + s.storeSnatchVars(task.ID, s.snatchSuccessVars(ctx, p, in, instances)) +} + +// snatchSuccessVars 组装抢机成功通知的模板变量;单项查询失败以占位值回退, +// 绝不阻塞通知发送。密码只进通知变量,不落任务日志。 +func (s *TaskService) snatchSuccessVars(ctx context.Context, p *snatchPayload, in oci.CreateInstanceInput, instances []oci.Instance) map[string]string { + tenant, region := "-", in.Region + if cfg, err := s.configs.Get(ctx, p.OciConfigID); err == nil { + tenant = cfg.Alias + if region == "" { + region = cfg.Region + } + } + ips, pwds := s.snatchInstanceNet(ctx, p.OciConfigID, region, in, instances) + return map[string]string{ + "tenant": tenant, + "region": regionLabel(region), + "shape": in.Shape, + "spec": specLabel(in), + "image": s.snatchImageLabel(ctx, p.OciConfigID, region, in), + "ip": ips, + "root_password": pwds, + } +} + +// snatchInstanceNet 逐台等待公网 IPv4 并取回 root 密码;全部实例共享 +// snatchIPWait 总预算,多台结果按创建顺序以「、」连接。 +func (s *TaskService) snatchInstanceNet(ctx context.Context, cfgID uint, region string, in oci.CreateInstanceInput, instances []oci.Instance) (string, string) { + deadline := time.Now().Add(s.snatchIPWait) + ips := make([]string, 0, len(instances)) + pwds := make([]string, 0, len(instances)) + for _, inst := range instances { + ips = append(ips, s.waitPublicIPv4(ctx, cfgID, region, inst.ID, deadline)) + pwds = append(pwds, rootPasswordOf(in, inst)) + } + return strings.Join(ips, "、"), strings.Join(pwds, "、") +} + +// waitPublicIPv4 轮询实例公网 IPv4 直至分配或超过 deadline; +// 超时回退内网 IPv4(标注内网),再无则「待分配」。 +func (s *TaskService) waitPublicIPv4(ctx context.Context, cfgID uint, region, instanceID string, deadline time.Time) string { + private := "" + for { + inst, err := s.configs.Instance(ctx, cfgID, region, instanceID) + if err == nil { + if inst.PublicIP != "" { + return inst.PublicIP + } + if inst.PrivateIP != "" { + private = inst.PrivateIP + } + } + if time.Now().After(deadline) { + break + } + select { + case <-ctx.Done(): + return waitIPFallback(private) + case <-time.After(snatchIPPollInterval): + } + } + return waitIPFallback(private) +} + +// waitIPFallback 是公网 IP 等待超时后的展示回退。 +func waitIPFallback(private string) string { + if private != "" { + return private + "(内网)" + } + return "待分配" +} + +// rootPasswordOf 取单台实例的 root 密码:随机生成模式从创建时写入的 +// FreeformTags 回读(每台独立),显式密码模式各台相同;SSH 密钥模式无密码。 +func rootPasswordOf(in oci.CreateInstanceInput, inst oci.Instance) string { + if pwd := inst.FreeformTags["RootPassword"]; pwd != "" { + return pwd + } + if in.RootPassword != "" { + return in.RootPassword + } + return "—(SSH 密钥登录)" +} + +// specLabel 拼 Flex 规格,如 4C/24G;固定规格 shape(无 ocpus 参数)为 -。 +func specLabel(in oci.CreateInstanceInput) string { + if in.Ocpus <= 0 { + return "-" + } + return fmt.Sprintf("%gC/%gG", in.Ocpus, in.MemoryInGBs) +} + +// regionLabel 把区域标识换成控制台风格别名展示,如「Japan Central (Osaka)」; +// 区域名与三字码均可匹配,表中未收录时回退原值。 +func regionLabel(region string) string { + if region == "" { + return "-" + } + regions, err := oci.AllRegions() + if err != nil { + return region + } + for _, r := range regions { + if strings.EqualFold(r.Name, region) || strings.EqualFold(r.Key, region) { + return r.Alias + } + } + return region +} + +// snatchImageLabel 返回启动源展示名:镜像查显示名(失败回退 OCID), +// 引导卷启动源固定文案。 +func (s *TaskService) snatchImageLabel(ctx context.Context, cfgID uint, region string, in oci.CreateInstanceInput) string { + if in.BootVolumeID != "" { + return "引导卷" + } + if in.ImageID == "" { + return "-" + } + img, err := s.configs.Image(ctx, cfgID, region, in.ImageID) + if err != nil || img.DisplayName == "" { + return in.ImageID + } + return img.DisplayName +} + +// storeSnatchVars 暂存抢机成功通知变量,由 takeSnatchVars 取走。 +func (s *TaskService) storeSnatchVars(taskID uint, vars map[string]string) { + s.snatchVarsMu.Lock() + defer s.snatchVarsMu.Unlock() + s.snatchVars[taskID] = vars +} + +// takeSnatchVars 取走并清除暂存的成功通知变量(无则 nil)。 +func (s *TaskService) takeSnatchVars(taskID uint) map[string]string { + s.snatchVarsMu.Lock() + defer s.snatchVarsMu.Unlock() + vars := s.snatchVars[taskID] + delete(s.snatchVars, taskID) + return vars +} + // snatchInstanceInput 组装本次创建参数:可用域显式指定时原样使用; // 留空(自动)时按执行序号轮询区域全部可用域——ad-1、ad-2、ad-3 依次循环, // 分摊单可用域容量不足。附加说明串供执行日志展示本次所用可用域。 @@ -735,6 +956,9 @@ type taskSnapshot struct { Status string LastError string Message string // 本次执行结果摘要,仅执行后快照填写 + // SnatchVars 是抢机成功通知的补充变量(租户/区域/实例明细), + // 仅抢满那次执行后快照携带 + SnatchVars map[string]string } // notifyKind 是通知事件类型,与设置页「通知管理」开关一一对应。 @@ -760,7 +984,11 @@ type notifyEvent struct { func notifyEvents(prev, cur taskSnapshot) []notifyEvent { var events []notifyEvent if prev.Status != model.TaskStatusSucceeded && cur.Status == model.TaskStatusSucceeded { - events = append(events, notifyEvent{notifySnatchSuccess, map[string]string{"task_name": cur.Name, "message": cur.Message}}) + vars := map[string]string{"task_name": cur.Name, "message": cur.Message} + for k, v := range cur.SnatchVars { + vars[k] = v + } + events = append(events, notifyEvent{notifySnatchSuccess, vars}) } switch { case prev.Status != model.TaskStatusFailed && cur.Status == model.TaskStatusFailed: diff --git a/internal/service/task_test.go b/internal/service/task_test.go index ee73d1e..8a63cd4 100644 --- a/internal/service/task_test.go +++ b/internal/service/task_test.go @@ -80,7 +80,10 @@ func newTaskEnv(t *testing.T, client oci.Client) (*TaskService, *OciConfigServic t.Fatalf("new cipher: %v", err) } configs := NewOciConfigService(db, cipher, client) - return NewTaskService(db, configs, nil, NewSettingService(db, cipher)), configs, db + tasks := NewTaskService(db, configs, nil, NewSettingService(db, cipher)) + // 测试不等待公网 IP 分配:预算置 0,首查后立即回退 + tasks.snatchIPWait = 0 + return tasks, configs, db } func TestCreateTaskValidation(t *testing.T) { @@ -699,3 +702,102 @@ func TestUpdateTaskFailedToActiveResetsAuthFail(t *testing.T) { t.Error("重新启用后任务未回到调度") } } + +// TestSnatchSuccessVars 锁定抢机成功通知变量的组装:租户/区域/类型/镜像, +// 以及 root 密码三种凭据模式与 IP 逐级回退;单项查询失败不阻塞其余字段。 +func TestSnatchSuccessVars(t *testing.T) { + tests := []struct { + name string + client *fakeClient + in oci.CreateInstanceInput + insts []oci.Instance + want map[string]string + }{ + { + name: "生成密码模式全字段", + client: &fakeClient{ + tenancy: oci.TenancyInfo{Name: "mytenancy", HomeRegionKey: "FRA"}, + instance: oci.Instance{PublicIP: "203.0.113.7"}, + image: oci.Image{DisplayName: "Ubuntu-24.04"}, + }, + in: oci.CreateInstanceInput{Region: "ap-osaka-1", Shape: "VM.Standard.A1.Flex", + Ocpus: 4, MemoryInGBs: 24, ImageID: "ocid1.image.oc1..img"}, + insts: []oci.Instance{{ID: "i1", FreeformTags: map[string]string{"RootPassword": "pw-from-tag"}}}, + want: map[string]string{ + "tenant": "试用期", + "region": "Japan Central (Osaka)", + "shape": "VM.Standard.A1.Flex", + "spec": "4C/24G", + "image": "Ubuntu-24.04", + "ip": "203.0.113.7", + "root_password": "pw-from-tag", + }, + }, + { + name: "显式密码_未知区域_IP回退内网_镜像失败回退OCID", + client: &fakeClient{ + instance: oci.Instance{PrivateIP: "10.0.0.5"}, + imageErr: fmt.Errorf("boom"), + }, + in: oci.CreateInstanceInput{Region: "xx-nowhere-1", Shape: "VM.Standard.E2.1.Micro", + ImageID: "ocid1.image.oc1..img", RootPassword: "explicit-pw"}, + insts: []oci.Instance{{ID: "i1"}}, + want: map[string]string{ + "tenant": "试用期", + "region": "xx-nowhere-1", + "shape": "VM.Standard.E2.1.Micro", + "spec": "-", + "image": "ocid1.image.oc1..img", + "ip": "10.0.0.5(内网)", + "root_password": "explicit-pw", + }, + }, + { + name: "SSH模式_引导卷_区域取配置默认_多台待分配", + client: &fakeClient{instanceErr: fmt.Errorf("boom")}, + in: oci.CreateInstanceInput{Shape: "s", + BootVolumeID: "ocid1.bootvolume.oc1..bv", SSHPublicKey: "ssh-ed25519 AAAA"}, + insts: []oci.Instance{{ID: "i1"}, {ID: "i2"}}, + want: map[string]string{ + "region": "Germany Central (Frankfurt)", + "image": "引导卷", + "ip": "待分配、待分配", + "root_password": "—(SSH 密钥登录)、—(SSH 密钥登录)", + }, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + tasks, configs, _ := newTaskEnv(t, tt.client) + cfg, err := configs.Import(context.Background(), trialImportInput()) + if err != nil { + t.Fatalf("import: %v", err) + } + got := tasks.snatchSuccessVars(context.Background(), + &snatchPayload{OciConfigID: cfg.ID}, tt.in, tt.insts) + for k, want := range tt.want { + if got[k] != want { + t.Errorf("%s = %q, want %q", k, got[k], want) + } + } + }) + } +} + +// TestNotifyEventsSnatchVarsMerged 抢满成功事件应合并暂存的补充变量。 +func TestNotifyEventsSnatchVarsMerged(t *testing.T) { + prev := taskSnapshot{Name: "抢机", Status: model.TaskStatusActive} + cur := taskSnapshot{Name: "抢机", Status: model.TaskStatusSucceeded, Message: "created 1", + SnatchVars: map[string]string{"tenant": "免费01·t", "ip": "203.0.113.7", "root_password": "pw"}} + events := notifyEvents(prev, cur) + if len(events) != 1 || events[0].Kind != notifySnatchSuccess { + t.Fatalf("events = %+v, want 1 条 snatch_success", events) + } + want := map[string]string{"task_name": "抢机", "message": "created 1", + "tenant": "免费01·t", "ip": "203.0.113.7", "root_password": "pw"} + for k, v := range want { + if events[0].Vars[k] != v { + t.Errorf("vars[%s] = %q, want %q", k, events[0].Vars[k], v) + } + } +} diff --git a/internal/service/tenantadmin.go b/internal/service/tenantadmin.go index f022094..da3ef1b 100644 --- a/internal/service/tenantadmin.go +++ b/internal/service/tenantadmin.go @@ -31,27 +31,36 @@ func (s *OciConfigService) credentialsAndHomeRegion(ctx context.Context, id uint return cred, homeRegionOf(cfg), nil } -// TenantUsers 列出租户 IAM 用户。 -func (s *OciConfigService) TenantUsers(ctx context.Context, id uint) ([]oci.TenantUser, error) { - cred, err := s.credentialsByID(ctx, id) +// IdentityDomains 列出租户全部 ACTIVE 身份域(域选择器数据源)。 +func (s *OciConfigService) IdentityDomains(ctx context.Context, id uint) ([]oci.IdentityDomain, error) { + cred, homeRegion, err := s.credentialsAndHomeRegion(ctx, id) if err != nil { return nil, err } - return s.client.ListTenantUsers(ctx, cred) + return s.client.ListIdentityDomains(ctx, cred, homeRegion) +} + +// TenantUsers 列出租户 IAM 用户;domainID 非空只列该身份域。 +func (s *OciConfigService) TenantUsers(ctx context.Context, id uint, domainID string) ([]oci.TenantUser, error) { + cred, homeRegion, err := s.credentialsAndHomeRegion(ctx, id) + if err != nil { + return nil, err + } + return s.client.ListTenantUsers(ctx, cred, homeRegion, domainID) } // TenantUserDetail 查用户的域档案与管理员状态(编辑表单预填充用)。 -func (s *OciConfigService) TenantUserDetail(ctx context.Context, id uint, userID string) (oci.TenantUserDetail, error) { +func (s *OciConfigService) TenantUserDetail(ctx context.Context, id uint, domainID, userID string) (oci.TenantUserDetail, error) { cred, homeRegion, err := s.credentialsAndHomeRegion(ctx, id) if err != nil { return oci.TenantUserDetail{}, err } - return s.client.GetTenantUserDetail(ctx, cred, homeRegion, userID) + return s.client.GetTenantUserDetail(ctx, cred, homeRegion, domainID, userID) } // CreateTenantUser 新增用户;名字与姓氏必填(域用户档案要求), // 描述缺省用「名 姓」。 -func (s *OciConfigService) CreateTenantUser(ctx context.Context, id uint, in oci.CreateTenantUserInput) (oci.TenantUser, error) { +func (s *OciConfigService) CreateTenantUser(ctx context.Context, id uint, domainID string, in oci.CreateTenantUserInput) (oci.TenantUser, error) { if strings.TrimSpace(in.Name) == "" { return oci.TenantUser{}, fmt.Errorf("create tenant user: name is required") } @@ -65,11 +74,11 @@ func (s *OciConfigService) CreateTenantUser(ctx context.Context, id uint, in oci if err != nil { return oci.TenantUser{}, err } - return s.client.CreateTenantUser(ctx, cred, homeRegion, in) + return s.client.CreateTenantUser(ctx, cred, homeRegion, domainID, in) } // UpdateTenantUser 编辑用户资料;nil 字段不修改,全 nil 报错。 -func (s *OciConfigService) UpdateTenantUser(ctx context.Context, id uint, userID string, in oci.UpdateTenantUserInput) (oci.TenantUser, error) { +func (s *OciConfigService) UpdateTenantUser(ctx context.Context, id uint, domainID, userID string, in oci.UpdateTenantUserInput) (oci.TenantUser, error) { hasField := in.Description != nil || in.Email != nil || in.GivenName != nil || in.FamilyName != nil hasAdmin := in.GrantDomainAdmin != nil || in.AddToAdminGroup != nil if !hasField && !hasAdmin { @@ -79,29 +88,29 @@ func (s *OciConfigService) UpdateTenantUser(ctx context.Context, id uint, userID if err != nil { return oci.TenantUser{}, err } - return s.client.UpdateTenantUser(ctx, cred, homeRegion, userID, in) + return s.client.UpdateTenantUser(ctx, cred, homeRegion, domainID, userID, in) } // IdentitySetting 读取域身份设置(当前只透出主邮箱必填开关)。 -func (s *OciConfigService) IdentitySetting(ctx context.Context, id uint) (oci.IdentitySettingInfo, error) { +func (s *OciConfigService) IdentitySetting(ctx context.Context, id uint, domainID string) (oci.IdentitySettingInfo, error) { cred, homeRegion, err := s.credentialsAndHomeRegion(ctx, id) if err != nil { return oci.IdentitySettingInfo{}, err } - return s.client.GetIdentitySetting(ctx, cred, homeRegion) + return s.client.GetIdentitySetting(ctx, cred, homeRegion, domainID) } // UpdateIdentitySetting 修改「用户需要提供主电子邮件地址」开关。 -func (s *OciConfigService) UpdateIdentitySetting(ctx context.Context, id uint, primaryEmailRequired bool) (oci.IdentitySettingInfo, error) { +func (s *OciConfigService) UpdateIdentitySetting(ctx context.Context, id uint, domainID string, primaryEmailRequired bool) (oci.IdentitySettingInfo, error) { cred, homeRegion, err := s.credentialsAndHomeRegion(ctx, id) if err != nil { return oci.IdentitySettingInfo{}, err } - return s.client.UpdateIdentitySetting(ctx, cred, homeRegion, primaryEmailRequired) + return s.client.UpdateIdentitySetting(ctx, cred, homeRegion, domainID, primaryEmailRequired) } // DeleteTenantUser 删除 IAM 用户;拒绝删除当前配置正在使用的用户。 -func (s *OciConfigService) DeleteTenantUser(ctx context.Context, id uint, userID string) error { +func (s *OciConfigService) DeleteTenantUser(ctx context.Context, id uint, domainID, userID string) error { cred, homeRegion, err := s.credentialsAndHomeRegion(ctx, id) if err != nil { return err @@ -109,25 +118,25 @@ func (s *OciConfigService) DeleteTenantUser(ctx context.Context, id uint, userID if userID == cred.UserOCID { return fmt.Errorf("delete tenant user: refusing to delete the user this config signs requests with") } - return s.client.DeleteTenantUser(ctx, cred, homeRegion, userID) + return s.client.DeleteTenantUser(ctx, cred, homeRegion, domainID, userID) } // ResetTenantUserPassword 重置用户控制台密码,返回一次性新密码。 -func (s *OciConfigService) ResetTenantUserPassword(ctx context.Context, id uint, userID string) (string, error) { +func (s *OciConfigService) ResetTenantUserPassword(ctx context.Context, id uint, domainID, userID string) (string, error) { cred, homeRegion, err := s.credentialsAndHomeRegion(ctx, id) if err != nil { return "", err } - return s.client.ResetTenantUserPassword(ctx, cred, homeRegion, userID) + return s.client.ResetTenantUserPassword(ctx, cred, homeRegion, domainID, userID) } // DeleteTenantUserMfa 清除用户全部 MFA,返回删除的经典 TOTP 设备数。 -func (s *OciConfigService) DeleteTenantUserMfa(ctx context.Context, id uint, userID string) (int, error) { +func (s *OciConfigService) DeleteTenantUserMfa(ctx context.Context, id uint, domainID, userID string) (int, error) { cred, homeRegion, err := s.credentialsAndHomeRegion(ctx, id) if err != nil { return 0, err } - return s.client.DeleteTenantUserMfaDevices(ctx, cred, homeRegion, userID) + return s.client.DeleteTenantUserMfaDevices(ctx, cred, homeRegion, domainID, userID) } // DeleteTenantUserApiKeys 删除用户的 API Key;默认保留当前配置使用的指纹。 @@ -140,17 +149,17 @@ func (s *OciConfigService) DeleteTenantUserApiKeys(ctx context.Context, id uint, } // NotificationRecipients 查询域通知收件人设置。 -func (s *OciConfigService) NotificationRecipients(ctx context.Context, id uint) (oci.NotificationRecipients, error) { +func (s *OciConfigService) NotificationRecipients(ctx context.Context, id uint, domainID string) (oci.NotificationRecipients, error) { cred, homeRegion, err := s.credentialsAndHomeRegion(ctx, id) if err != nil { return oci.NotificationRecipients{}, err } - return s.client.GetNotificationRecipients(ctx, cred, homeRegion) + return s.client.GetNotificationRecipients(ctx, cred, homeRegion, domainID) } // UpdateNotificationRecipients 把域通知改为只发给指定收件人; // 收件人为空表示关闭 test mode 恢复默认发送。 -func (s *OciConfigService) UpdateNotificationRecipients(ctx context.Context, id uint, emails []string) (oci.NotificationRecipients, error) { +func (s *OciConfigService) UpdateNotificationRecipients(ctx context.Context, id uint, domainID string, emails []string) (oci.NotificationRecipients, error) { if err := validateEmails(emails); err != nil { return oci.NotificationRecipients{}, err } @@ -158,7 +167,7 @@ func (s *OciConfigService) UpdateNotificationRecipients(ctx context.Context, id if err != nil { return oci.NotificationRecipients{}, err } - return s.client.UpdateNotificationRecipients(ctx, cred, homeRegion, emails) + return s.client.UpdateNotificationRecipients(ctx, cred, homeRegion, domainID, emails) } func validateEmails(emails []string) error { @@ -171,16 +180,16 @@ func validateEmails(emails []string) error { } // PasswordPolicies 列出域密码策略。 -func (s *OciConfigService) PasswordPolicies(ctx context.Context, id uint) ([]oci.PasswordPolicyInfo, error) { +func (s *OciConfigService) PasswordPolicies(ctx context.Context, id uint, domainID string) ([]oci.PasswordPolicyInfo, error) { cred, homeRegion, err := s.credentialsAndHomeRegion(ctx, id) if err != nil { return nil, err } - return s.client.ListPasswordPolicies(ctx, cred, homeRegion) + return s.client.ListPasswordPolicies(ctx, cred, homeRegion, domainID) } // UpdatePasswordPolicy 修改域密码策略的过期天数或最小长度。 -func (s *OciConfigService) UpdatePasswordPolicy(ctx context.Context, id uint, policyID string, in oci.UpdatePasswordPolicyInput) (oci.PasswordPolicyInfo, error) { +func (s *OciConfigService) UpdatePasswordPolicy(ctx context.Context, id uint, domainID, policyID string, in oci.UpdatePasswordPolicyInput) (oci.PasswordPolicyInfo, error) { if in.PasswordExpiresAfter == nil && in.MinLength == nil { return oci.PasswordPolicyInfo{}, fmt.Errorf("update password policy: nothing to update") } @@ -188,5 +197,5 @@ func (s *OciConfigService) UpdatePasswordPolicy(ctx context.Context, id uint, po if err != nil { return oci.PasswordPolicyInfo{}, err } - return s.client.UpdatePasswordPolicy(ctx, cred, homeRegion, policyID, in) + return s.client.UpdatePasswordPolicy(ctx, cred, homeRegion, domainID, policyID, in) } diff --git a/internal/service/tenantdelete.go b/internal/service/tenantdelete.go new file mode 100644 index 0000000..4b5053d --- /dev/null +++ b/internal/service/tenantdelete.go @@ -0,0 +1,457 @@ +package service + +import ( + "context" + "encoding/json" + "fmt" + "strings" + + "gorm.io/gorm" + "gorm.io/gorm/clause" + + "oci-portal/internal/model" +) + +type tenantDeleteResult struct { + config model.OciConfig + deletedTaskIDs []uint + alertRuleIDs []uint + channelsGone bool +} + +type tenantTaskAction struct { + task model.Task + deleteTask bool + payload string +} + +type tenancyCacheInvalidator interface { + InvalidateTenancy(tenancyOCID string) +} + +func (s *OciConfigService) deleteTenant(ctx context.Context, id uint) (*tenantDeleteResult, error) { + result := &tenantDeleteResult{} + err := s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + return s.deleteTenantInTx(tx, id, result) + }) + if err != nil { + return nil, fmt.Errorf("delete oci config %d: %w", id, err) + } + return result, nil +} + +func (s *OciConfigService) deleteTenantInTx(tx *gorm.DB, id uint, result *tenantDeleteResult) error { + if err := lockTenant(tx, id, &result.config); err != nil { + return err + } + if err := s.deleteTenantTasks(tx, id, result); err != nil { + return err + } + if err := deleteTenantEvents(tx, id, result); err != nil { + return err + } + if err := deleteTenantAI(tx, id, result); err != nil { + return err + } + if err := deleteTenantSnapshots(tx, id); err != nil { + return err + } + return deleteTenantConfig(tx, id) +} + +func lockTenant(tx *gorm.DB, id uint, cfg *model.OciConfig) error { + err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).First(cfg, id).Error + if err != nil { + return fmt.Errorf("load tenant: %w", err) + } + return nil +} + +func (s *OciConfigService) deleteTenantTasks(tx *gorm.DB, id uint, result *tenantDeleteResult) error { + var tasks []model.Task + err := tx.Where("type IN ?", []string{ + model.TaskTypeSnatch, model.TaskTypeHealthCheck, model.TaskTypeCost, + }).Find(&tasks).Error + if err != nil { + return fmt.Errorf("load tenant tasks: %w", err) + } + for i := range tasks { + action, ok, err := planTenantTask(tasks[i], id) + if err != nil { + return err + } + if !ok { + continue + } + if err := applyTenantTask(tx, action, result); err != nil { + return err + } + } + return nil +} + +func planTenantTask(task model.Task, id uint) (tenantTaskAction, bool, error) { + switch task.Type { + case model.TaskTypeSnatch: + return planSnatchTask(task, id) + case model.TaskTypeHealthCheck, model.TaskTypeCost: + return planMultiTenantTask(task, id) + default: + return tenantTaskAction{}, false, nil + } +} + +func planSnatchTask(task model.Task, id uint) (tenantTaskAction, bool, error) { + var payload snatchPayload + if err := json.Unmarshal([]byte(task.Payload), &payload); err != nil { + return tenantTaskAction{}, false, fmt.Errorf("parse task %d payload: %w", task.ID, err) + } + if payload.OciConfigID != id { + return tenantTaskAction{}, false, nil + } + return tenantTaskAction{task: task, deleteTask: true}, true, nil +} + +func planMultiTenantTask(task model.Task, id uint) (tenantTaskAction, bool, error) { + payload, err := decodeConfigIDs(task) + if err != nil { + return tenantTaskAction{}, false, err + } + if len(payload.OciConfigIDs) == 0 { + return tenantTaskAction{task: task, payload: task.Payload}, true, nil + } + remaining, found := removeConfigID(payload.OciConfigIDs, id) + if !found { + return tenantTaskAction{}, false, nil + } + if len(remaining) == 0 { + return tenantTaskAction{task: task, deleteTask: true}, true, nil + } + payload.OciConfigIDs = remaining + raw, err := json.Marshal(payload) + if err != nil { + return tenantTaskAction{}, false, fmt.Errorf("encode task %d payload: %w", task.ID, err) + } + return tenantTaskAction{task: task, payload: string(raw)}, true, nil +} + +func decodeConfigIDs(task model.Task) (healthCheckPayload, error) { + var payload healthCheckPayload + if strings.TrimSpace(task.Payload) == "" { + return payload, nil + } + if err := json.Unmarshal([]byte(task.Payload), &payload); err != nil { + return payload, fmt.Errorf("parse task %d payload: %w", task.ID, err) + } + return payload, nil +} + +func removeConfigID(ids []uint, id uint) ([]uint, bool) { + remaining := make([]uint, 0, len(ids)) + found := false + for _, candidate := range ids { + if candidate == id { + found = true + continue + } + remaining = append(remaining, candidate) + } + return remaining, found +} + +func applyTenantTask(tx *gorm.DB, action tenantTaskAction, result *tenantDeleteResult) error { + if err := deleteTaskHistory(tx, action.task.ID); err != nil { + return err + } + if action.deleteTask { + if err := tx.Delete(&model.Task{}, action.task.ID).Error; err != nil { + return fmt.Errorf("delete task %d: %w", action.task.ID, err) + } + result.deletedTaskIDs = append(result.deletedTaskIDs, action.task.ID) + return nil + } + return resetTenantTask(tx, action) +} + +func deleteTaskHistory(tx *gorm.DB, taskID uint) error { + if err := tx.Where("task_id = ?", taskID).Delete(&model.TaskLog{}).Error; err != nil { + return fmt.Errorf("delete task %d logs: %w", taskID, err) + } + key := deadAliasKey(taskID) + if err := tx.Where("key = ?", key).Delete(&model.Setting{}).Error; err != nil { + return fmt.Errorf("delete task %d state: %w", taskID, err) + } + return nil +} + +func resetTenantTask(tx *gorm.DB, action tenantTaskAction) error { + updates := map[string]any{ + "payload": action.payload, "last_run_at": nil, + "last_error": "", "run_count": 0, + } + err := tx.Model(&model.Task{}).Where("id = ?", action.task.ID).Updates(updates).Error + if err != nil { + return fmt.Errorf("reset task %d: %w", action.task.ID, err) + } + return nil +} + +func deleteTenantEvents(tx *gorm.DB, id uint, result *tenantDeleteResult) error { + ruleIDs, eventIDs, affectedRules, err := loadTenantEventRefs(tx, id) + if err != nil { + return err + } + if err := deleteAlertHits(tx, ruleIDs, eventIDs); err != nil { + return err + } + result.alertRuleIDs = mergeIDs(ruleIDs, affectedRules) + if err := deleteWhere(tx, &model.AlertRule{}, "oci_config_id = ?", id); err != nil { + return fmt.Errorf("delete tenant alert rules: %w", err) + } + if err := deleteWhere(tx, &model.LogEvent{}, "oci_config_id = ?", id); err != nil { + return fmt.Errorf("delete tenant log events: %w", err) + } + return nil +} + +func loadTenantEventRefs(tx *gorm.DB, id uint) ([]uint, []uint, []uint, error) { + ruleIDs, err := lockedTenantRuleIDs(tx, id) + if err != nil { + return nil, nil, nil, fmt.Errorf("load tenant alert rules: %w", err) + } + eventIDs, err := lockedTenantEventIDs(tx, id) + if err != nil { + return nil, nil, nil, fmt.Errorf("load tenant log events: %w", err) + } + affected, err := alertHitRuleIDs(tx, eventIDs) + if err != nil { + return nil, nil, nil, err + } + return ruleIDs, eventIDs, affected, nil +} + +func lockedTenantRuleIDs(tx *gorm.DB, id uint) ([]uint, error) { + var rows []model.AlertRule + err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Select("id"). + Where("oci_config_id = ?", id).Order("id").Find(&rows).Error + ids := make([]uint, 0, len(rows)) + for _, row := range rows { + ids = append(ids, row.ID) + } + return ids, err +} + +func lockedTenantEventIDs(tx *gorm.DB, id uint) ([]uint, error) { + var rows []model.LogEvent + err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Select("id"). + Where("oci_config_id = ?", id).Order("id").Find(&rows).Error + ids := make([]uint, 0, len(rows)) + for _, row := range rows { + ids = append(ids, row.ID) + } + return ids, err +} + +func alertHitRuleIDs(tx *gorm.DB, eventIDs []uint) ([]uint, error) { + if len(eventIDs) == 0 { + return nil, nil + } + ids := make([]uint, 0) + err := tx.Model(&model.AlertRuleHit{}).Where("log_event_id IN ?", eventIDs). + Distinct().Pluck("rule_id", &ids).Error + if err != nil { + return nil, fmt.Errorf("load affected alert rules: %w", err) + } + return ids, nil +} + +func mergeIDs(groups ...[]uint) []uint { + seen := make(map[uint]struct{}) + out := make([]uint, 0) + for _, ids := range groups { + for _, id := range ids { + if _, ok := seen[id]; ok { + continue + } + seen[id] = struct{}{} + out = append(out, id) + } + } + return out +} + +func deleteAlertHits(tx *gorm.DB, ruleIDs, eventIDs []uint) error { + query := tx.Model(&model.AlertRuleHit{}) + switch { + case len(ruleIDs) > 0 && len(eventIDs) > 0: + query = query.Where("rule_id IN ? OR log_event_id IN ?", ruleIDs, eventIDs) + case len(ruleIDs) > 0: + query = query.Where("rule_id IN ?", ruleIDs) + case len(eventIDs) > 0: + query = query.Where("log_event_id IN ?", eventIDs) + default: + return nil + } + if err := query.Delete(&model.AlertRuleHit{}).Error; err != nil { + return fmt.Errorf("delete tenant alert hits: %w", err) + } + return nil +} + +func deleteTenantAI(tx *gorm.DB, id uint, result *tenantDeleteResult) error { + channelIDs, err := lockedAiChannelIDs(tx, id) + if err != nil { + return fmt.Errorf("load tenant AI channels: %w", err) + } + if len(channelIDs) == 0 { + return nil + } + callIDs, err := lockedAiCallIDs(tx, channelIDs) + if err != nil { + return fmt.Errorf("load tenant AI calls: %w", err) + } + if err := deleteTenantAIRows(tx, channelIDs, callIDs); err != nil { + return err + } + result.channelsGone = true + return reconcileAiProbeRows(tx, result) +} + +func lockedAiChannelIDs(tx *gorm.DB, id uint) ([]uint, error) { + var rows []model.AiChannel + err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Select("id"). + Where("oci_config_id = ?", id).Order("id").Find(&rows).Error + ids := make([]uint, 0, len(rows)) + for _, row := range rows { + ids = append(ids, row.ID) + } + return ids, err +} + +func lockedAiCallIDs(tx *gorm.DB, channelIDs []uint) ([]uint, error) { + var rows []model.AiCallLog + err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Select("id"). + Where("channel_id IN ?", channelIDs).Order("id").Find(&rows).Error + ids := make([]uint, 0, len(rows)) + for _, row := range rows { + ids = append(ids, row.ID) + } + return ids, err +} + +func deleteTenantAIRows(tx *gorm.DB, channelIDs, callIDs []uint) error { + if len(callIDs) > 0 { + if err := deleteWhere(tx, &model.AiContentLog{}, "call_log_id IN ?", callIDs); err != nil { + return fmt.Errorf("delete tenant AI content logs: %w", err) + } + } + steps := []struct { + name string + value any + }{ + {"AI call logs", &model.AiCallLog{}}, + {"AI model cache", &model.AiModelCache{}}, + {"AI channels", &model.AiChannel{}}, + } + for _, step := range steps { + if err := deleteAIChannelRows(tx, step.value, channelIDs); err != nil { + return fmt.Errorf("delete tenant %s: %w", step.name, err) + } + } + return nil +} + +func deleteAIChannelRows(tx *gorm.DB, value any, channelIDs []uint) error { + column := "channel_id" + if _, ok := value.(*model.AiChannel); ok { + column = "id" + } + return deleteWhere(tx, value, column+" IN ?", channelIDs) +} + +func reconcileAiProbeRows(tx *gorm.DB, result *tenantDeleteResult) error { + var channelCount int64 + if err := tx.Model(&model.AiChannel{}).Count(&channelCount).Error; err != nil { + return fmt.Errorf("count remaining AI channels: %w", err) + } + var ids []uint + err := tx.Model(&model.Task{}).Where("type = ?", model.TaskTypeAiProbe).Pluck("id", &ids).Error + if err != nil { + return fmt.Errorf("load AI probe task: %w", err) + } + for _, id := range ids { + if err := deleteTaskHistory(tx, id); err != nil { + return err + } + if channelCount == 0 { + if err := tx.Delete(&model.Task{}, id).Error; err != nil { + return fmt.Errorf("delete AI probe task %d: %w", id, err) + } + result.deletedTaskIDs = append(result.deletedTaskIDs, id) + continue + } + if err := resetTaskHistoryFields(tx, id); err != nil { + return err + } + } + return nil +} + +func resetTaskHistoryFields(tx *gorm.DB, id uint) error { + updates := map[string]any{"last_run_at": nil, "last_error": "", "run_count": 0} + if err := tx.Model(&model.Task{}).Where("id = ?", id).Updates(updates).Error; err != nil { + return fmt.Errorf("reset task %d history: %w", id, err) + } + return nil +} + +func deleteTenantSnapshots(tx *gorm.DB, id uint) error { + steps := []struct { + name string + value any + }{ + {"check snapshots", &model.CheckSnapshot{}}, + {"cost snapshots", &model.CostSnapshot{}}, + {"region cache", &model.RegionCache{}}, + {"compartment cache", &model.CompartmentCache{}}, + } + for _, step := range steps { + if err := deleteWhere(tx, step.value, "oci_config_id = ?", id); err != nil { + return fmt.Errorf("delete tenant %s: %w", step.name, err) + } + } + if err := tx.Where("key = ?", secretKey(id)).Delete(&model.Setting{}).Error; err != nil { + return fmt.Errorf("delete tenant webhook secret: %w", err) + } + return nil +} + +func deleteTenantConfig(tx *gorm.DB, id uint) error { + res := tx.Delete(&model.OciConfig{}, id) + if res.Error != nil { + return fmt.Errorf("delete tenant row: %w", res.Error) + } + if res.RowsAffected == 0 { + return gorm.ErrRecordNotFound + } + return nil +} + +func deleteWhere(tx *gorm.DB, value any, query string, args ...any) error { + return tx.Where(query, args...).Delete(value).Error +} + +func (s *OciConfigService) afterTenantDelete(ctx context.Context, result *tenantDeleteResult) { + // 数据已提交,客户端断开不应中断必需的运行时对齐。 + ctx = context.WithoutCancel(ctx) + s.InvalidateAuditCache(result.config.ID) + if client, ok := s.client.(tenancyCacheInvalidator); ok { + client.InvalidateTenancy(result.config.TenancyOCID) + } + if s.cleanupEvents != nil { + s.cleanupEvents.ClearAlertCooldown(result.alertRuleIDs) + } + if s.cleanupTasks != nil { + s.cleanupTasks.ApplyTenantCleanup(ctx, result.deletedTaskIDs, result.channelsGone) + } +} diff --git a/internal/service/tenantdelete_test.go b/internal/service/tenantdelete_test.go new file mode 100644 index 0000000..25eb58b --- /dev/null +++ b/internal/service/tenantdelete_test.go @@ -0,0 +1,609 @@ +package service + +import ( + "context" + "errors" + "fmt" + "testing" + "time" + + "github.com/glebarez/sqlite" + "gorm.io/gorm" + "gorm.io/gorm/logger" + + "oci-portal/internal/crypto" + "oci-portal/internal/model" + "oci-portal/internal/oci" +) + +type invalidatingClient struct { + *fakeClient + invalidated []string +} + +type blockingClient struct { + *fakeClient + started chan struct{} + release chan struct{} +} + +func (c *blockingClient) ValidateKey(context.Context, oci.Credentials) (oci.TenancyInfo, error) { + c.started <- struct{}{} + <-c.release + return oci.TenancyInfo{Name: "target"}, nil +} + +func (c *invalidatingClient) InvalidateTenancy(id string) { + c.invalidated = append(c.invalidated, id) +} + +func newTenantDeleteEnv(t *testing.T, client oci.Client) (*OciConfigService, *TaskService, *gorm.DB) { + t.Helper() + db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)}) + if err != nil { + t.Fatalf("open database: %v", err) + } + sqlDB, err := db.DB() + if err != nil { + t.Fatalf("database handle: %v", err) + } + sqlDB.SetMaxOpenConns(1) + migrateTenantDeleteModels(t, db) + cipher, err := crypto.NewCipher("test-data-key") + if err != nil { + t.Fatalf("new cipher: %v", err) + } + configs := NewOciConfigService(db, cipher, client) + tasks := NewTaskService(db, configs, nil, nil) + configs.SetTenantCleanupDeps(tasks, nil) + return configs, tasks, db +} + +func migrateTenantDeleteModels(t *testing.T, db *gorm.DB) { + t.Helper() + err := db.AutoMigrate( + &model.OciConfig{}, &model.Task{}, &model.TaskLog{}, &model.Setting{}, + &model.CheckSnapshot{}, &model.CostSnapshot{}, &model.RegionCache{}, &model.CompartmentCache{}, + &model.LogEvent{}, &model.AlertRule{}, &model.AlertRuleHit{}, + &model.AiChannel{}, &model.AiModelCache{}, &model.AiCallLog{}, &model.AiContentLog{}, + &model.Proxy{}, &model.AiKey{}, &model.SystemLog{}, + ) + if err != nil { + t.Fatalf("auto migrate: %v", err) + } +} + +func TestPlanTenantTask(t *testing.T) { + for _, tt := range tenantTaskCases() { + t.Run(tt.name, func(t *testing.T) { + action, ok, err := planTenantTask(tt.task, 1) + if (err != nil) != tt.wantErr { + t.Fatalf("error = %v, wantErr %v", err, tt.wantErr) + } + if ok != tt.wantOK || action.deleteTask != tt.wantDelete { + t.Errorf("result = (ok=%v, delete=%v), want (%v, %v)", ok, action.deleteTask, tt.wantOK, tt.wantDelete) + } + if action.payload != tt.wantPayload { + t.Errorf("payload = %q, want %q", action.payload, tt.wantPayload) + } + }) + } +} + +type tenantTaskCase struct { + name string + task model.Task + wantOK bool + wantDelete bool + wantPayload string + wantErr bool +} + +func tenantTaskCases() []tenantTaskCase { + return []tenantTaskCase{ + {name: "抢机命中", task: taskOf(model.TaskTypeSnatch, `{"ociConfigId":1}`), wantOK: true, wantDelete: true}, + {name: "抢机未命中", task: taskOf(model.TaskTypeSnatch, `{"ociConfigId":2}`)}, + {name: "测活全局", task: taskOf(model.TaskTypeHealthCheck, `{"ociConfigIds":[]}`), wantOK: true, wantPayload: `{"ociConfigIds":[]}`}, + {name: "测活单租户", task: taskOf(model.TaskTypeHealthCheck, `{"ociConfigIds":[1]}`), wantOK: true, wantDelete: true}, + {name: "测活多租户", task: taskOf(model.TaskTypeHealthCheck, `{"ociConfigIds":[1,2]}`), wantOK: true, wantPayload: `{"ociConfigIds":[2]}`}, + {name: "成本去重命中", task: taskOf(model.TaskTypeCost, `{"ociConfigIds":[1,1,2]}`), wantOK: true, wantPayload: `{"ociConfigIds":[2]}`}, + {name: "成本未命中", task: taskOf(model.TaskTypeCost, `{"ociConfigIds":[2]}`)}, + {name: "非法 JSON", task: taskOf(model.TaskTypeCost, `{`), wantErr: true}, + } +} + +func taskOf(taskType, payload string) model.Task { + return model.Task{ID: 10, Type: taskType, Payload: payload} +} + +func TestDeleteTenantCleansRelatedRows(t *testing.T) { + client := &invalidatingClient{fakeClient: &fakeClient{}} + configs, _, db := newTenantDeleteEnv(t, client) + target, other := seedDeleteTenants(t, db) + seedTenantSnapshots(t, db, target.ID, other.ID) + seedTenantEvents(t, db, target.ID, other.ID) + seedTenantAI(t, db, target.ID, other.ID) + seedRetainedGlobals(t, db) + if err := configs.Delete(context.Background(), target.ID); err != nil { + t.Fatalf("delete tenant: %v", err) + } + assertTenantRowsGone(t, db, target.ID) + assertOtherTenantRowsRemain(t, db, other.ID) + assertRetainedGlobals(t, db) + if len(client.invalidated) != 1 || client.invalidated[0] != target.TenancyOCID { + t.Errorf("invalidated = %v, want [%s]", client.invalidated, target.TenancyOCID) + } +} + +func seedDeleteTenants(t *testing.T, db *gorm.DB) (model.OciConfig, model.OciConfig) { + t.Helper() + target := model.OciConfig{Alias: "target", TenancyOCID: "ocid1.tenancy.target"} + other := model.OciConfig{Alias: "other", TenancyOCID: "ocid1.tenancy.other"} + mustCreate(t, db, &target) + mustCreate(t, db, &other) + return target, other +} + +func seedTenantSnapshots(t *testing.T, db *gorm.DB, target, other uint) { + t.Helper() + for _, id := range []uint{target, other} { + mustCreate(t, db, &model.CheckSnapshot{OciConfigID: id}) + mustCreate(t, db, &model.CostSnapshot{OciConfigID: id, Day: "2026-07-10"}) + mustCreate(t, db, &model.RegionCache{OciConfigID: id, Key: "PHX"}) + mustCreate(t, db, &model.CompartmentCache{OciConfigID: id, OCID: fmt.Sprintf("comp-%d", id)}) + mustCreate(t, db, &model.Setting{Key: secretKey(id), Value: fmt.Sprintf("secret-%d", id)}) + } +} + +func seedTenantEvents(t *testing.T, db *gorm.DB, target, other uint) { + t.Helper() + targetEvent := model.LogEvent{OciConfigID: target, MessageID: "target-event"} + otherEvent := model.LogEvent{OciConfigID: other, MessageID: "other-event"} + targetRule := model.AlertRule{Name: "target-rule", OciConfigID: target} + globalRule := model.AlertRule{Name: "global-rule", OciConfigID: 0} + for _, value := range []any{&targetEvent, &otherEvent, &targetRule, &globalRule} { + mustCreate(t, db, value) + } + hits := []model.AlertRuleHit{ + {RuleID: targetRule.ID, LogEventID: otherEvent.ID}, + {RuleID: globalRule.ID, LogEventID: targetEvent.ID}, + {RuleID: globalRule.ID, LogEventID: otherEvent.ID}, + } + for i := range hits { + mustCreate(t, db, &hits[i]) + } +} + +func seedTenantAI(t *testing.T, db *gorm.DB, target, other uint) { + t.Helper() + for _, id := range []uint{target, other} { + channel := model.AiChannel{Name: fmt.Sprintf("channel-%d", id), OciConfigID: id, Region: "us-phoenix-1"} + mustCreate(t, db, &channel) + mustCreate(t, db, &model.AiModelCache{ChannelID: channel.ID, ModelOcid: fmt.Sprintf("model-%d", id)}) + call := model.AiCallLog{ChannelID: channel.ID, ChannelName: channel.Name} + mustCreate(t, db, &call) + mustCreate(t, db, &model.AiContentLog{CallLogID: call.ID, RequestBody: "sensitive"}) + } +} + +func seedRetainedGlobals(t *testing.T, db *gorm.DB) { + t.Helper() + mustCreate(t, db, &model.Proxy{Name: "shared", Type: "http"}) + mustCreate(t, db, &model.AiKey{Name: "global-key", KeyHash: "hash", Tail: "hash"}) + mustCreate(t, db, &model.SystemLog{Method: "DELETE", Path: "/api/v1/oci-configs/1"}) + mustCreate(t, db, &model.Setting{Key: "notify_channels", Value: "[]"}) +} + +func assertTenantRowsGone(t *testing.T, db *gorm.DB, id uint) { + t.Helper() + rows := []any{ + &model.OciConfig{}, &model.CheckSnapshot{}, &model.CostSnapshot{}, + &model.RegionCache{}, &model.CompartmentCache{}, &model.LogEvent{}, + &model.AlertRule{}, &model.AiChannel{}, + } + for _, value := range rows { + column := "oci_config_id" + if _, ok := value.(*model.OciConfig); ok { + column = "id" + } + assertCount(t, db, value, column+" = ?", []any{id}, 0) + } + assertCount(t, db, &model.Setting{}, "key = ?", []any{secretKey(id)}, 0) + assertCount(t, db, &model.AlertRuleHit{}, "", nil, 1) + assertCount(t, db, &model.AiModelCache{}, "", nil, 1) + assertCount(t, db, &model.AiCallLog{}, "", nil, 1) + assertCount(t, db, &model.AiContentLog{}, "", nil, 1) +} + +func assertOtherTenantRowsRemain(t *testing.T, db *gorm.DB, id uint) { + t.Helper() + for _, value := range []any{ + &model.OciConfig{}, &model.CheckSnapshot{}, &model.CostSnapshot{}, + &model.RegionCache{}, &model.CompartmentCache{}, &model.LogEvent{}, + &model.AiChannel{}, + } { + column := "oci_config_id" + if _, ok := value.(*model.OciConfig); ok { + column = "id" + } + assertCount(t, db, value, column+" = ?", []any{id}, 1) + } + assertCount(t, db, &model.Setting{}, "key = ?", []any{secretKey(id)}, 1) + assertRemainingIndirectRows(t, db, id) +} + +func assertRemainingIndirectRows(t *testing.T, db *gorm.DB, otherID uint) { + t.Helper() + var channel model.AiChannel + if err := db.Where("oci_config_id = ?", otherID).First(&channel).Error; err != nil { + t.Fatalf("load other AI channel: %v", err) + } + assertCount(t, db, &model.AiModelCache{}, "channel_id = ?", []any{channel.ID}, 1) + assertCount(t, db, &model.AiCallLog{}, "channel_id = ?", []any{channel.ID}, 1) + var call model.AiCallLog + if err := db.Where("channel_id = ?", channel.ID).First(&call).Error; err != nil { + t.Fatalf("load other AI call: %v", err) + } + assertCount(t, db, &model.AiContentLog{}, "call_log_id = ?", []any{call.ID}, 1) + assertRemainingAlertHit(t, db, otherID) +} + +func assertRemainingAlertHit(t *testing.T, db *gorm.DB, otherID uint) { + t.Helper() + var hit model.AlertRuleHit + if err := db.First(&hit).Error; err != nil { + t.Fatalf("load remaining alert hit: %v", err) + } + var rule model.AlertRule + var event model.LogEvent + if err := db.First(&rule, hit.RuleID).Error; err != nil { + t.Fatalf("load remaining rule: %v", err) + } + if err := db.First(&event, hit.LogEventID).Error; err != nil { + t.Fatalf("load remaining event: %v", err) + } + if rule.OciConfigID != 0 || event.OciConfigID != otherID { + t.Errorf("remaining hit = rule cfg %d/event cfg %d, want global/other", rule.OciConfigID, event.OciConfigID) + } +} + +func assertRetainedGlobals(t *testing.T, db *gorm.DB) { + t.Helper() + assertCount(t, db, &model.AlertRule{}, "oci_config_id = 0", nil, 1) + assertCount(t, db, &model.Proxy{}, "", nil, 1) + assertCount(t, db, &model.AiKey{}, "", nil, 1) + assertCount(t, db, &model.SystemLog{}, "", nil, 1) + assertCount(t, db, &model.Setting{}, "key = ?", []any{"notify_channels"}, 1) +} + +func TestDeleteTenantRewritesTasksAndCron(t *testing.T) { + configs, tasks, db := newTenantDeleteEnv(t, &fakeClient{}) + target, other := seedDeleteTenants(t, db) + created := seedTenantTasks(t, tasks, target.ID, other.ID) + if err := configs.Delete(context.Background(), target.ID); err != nil { + t.Fatalf("delete tenant: %v", err) + } + assertDeletedTasks(t, db, tasks, created[:2]) + assertRewrittenTasks(t, db, created[2:]) +} + +func TestDeleteTenantReconcilesAiProbe(t *testing.T) { + for _, keepOther := range []bool{false, true} { + name := fmt.Sprintf("keepOther=%v", keepOther) + t.Run(name, func(t *testing.T) { + configs, tasks, db := newTenantDeleteEnv(t, &fakeClient{}) + target, other := seedDeleteTenants(t, db) + seedProbeChannels(t, db, target.ID, other.ID, keepOther) + tasks.SyncAiProbeTask(context.Background()) + probe := loadAiProbe(t, db) + seedTaskHistory(t, db, probe.ID) + if err := configs.Delete(context.Background(), target.ID); err != nil { + t.Fatalf("delete tenant: %v", err) + } + assertAiProbeResult(t, db, tasks, probe.ID, keepOther) + }) + } +} + +func seedProbeChannels(t *testing.T, db *gorm.DB, target, other uint, keepOther bool) { + t.Helper() + mustCreate(t, db, &model.AiChannel{Name: "target", OciConfigID: target, Region: "r1"}) + if keepOther { + mustCreate(t, db, &model.AiChannel{Name: "other", OciConfigID: other, Region: "r1"}) + } +} + +func loadAiProbe(t *testing.T, db *gorm.DB) model.Task { + t.Helper() + var task model.Task + if err := db.Where("type = ?", model.TaskTypeAiProbe).First(&task).Error; err != nil { + t.Fatalf("load AI probe: %v", err) + } + return task +} + +func seedTaskHistory(t *testing.T, db *gorm.DB, taskID uint) { + t.Helper() + updates := map[string]any{"last_error": "old", "run_count": 2, "last_run_at": time.Now()} + if err := db.Model(&model.Task{}).Where("id = ?", taskID).Updates(updates).Error; err != nil { + t.Fatalf("seed task history: %v", err) + } + mustCreate(t, db, &model.TaskLog{TaskID: taskID, Message: "old"}) + mustCreate(t, db, &model.Setting{Key: deadAliasKey(taskID), Value: "old"}) +} + +func assertAiProbeResult(t *testing.T, db *gorm.DB, tasks *TaskService, id uint, keep bool) { + t.Helper() + want := int64(0) + if keep { + want = 1 + } + assertCount(t, db, &model.Task{}, "id = ?", []any{id}, want) + assertCount(t, db, &model.TaskLog{}, "task_id = ?", []any{id}, 0) + assertCount(t, db, &model.Setting{}, "key = ?", []any{deadAliasKey(id)}, 0) + if !keep { + if _, ok := tasks.entries[id]; ok { + t.Errorf("AI probe %d still scheduled", id) + } + return + } + probe := loadAiProbe(t, db) + if probe.RunCount != 0 || probe.LastRunAt != nil || probe.LastError != "" { + t.Errorf("AI probe history not reset: %+v", probe) + } +} + +func seedTenantTasks(t *testing.T, tasks *TaskService, target, other uint) []model.Task { + t.Helper() + inputs := []CreateTaskInput{ + {Name: "snatch", Type: model.TaskTypeSnatch, CronExpr: "0 0 * * *", Payload: []byte(fmt.Sprintf(`{"ociConfigId":%d,"instance":{"displayName":"vm","region":"r","availabilityDomain":"a","subnetId":"s","shape":"x","imageId":"i"}}`, target))}, + {Name: "single", Type: model.TaskTypeCost, CronExpr: "0 0 * * *", Payload: []byte(fmt.Sprintf(`{"ociConfigIds":[%d]}`, target))}, + {Name: "mixed", Type: model.TaskTypeHealthCheck, CronExpr: "0 0 * * *", Payload: []byte(fmt.Sprintf(`{"ociConfigIds":[%d,%d]}`, target, other))}, + {Name: "global", Type: model.TaskTypeCost, CronExpr: "0 0 * * *", Payload: []byte(`{"ociConfigIds":[]}`)}, + {Name: "other", Type: model.TaskTypeHealthCheck, CronExpr: "0 0 * * *", Payload: []byte(fmt.Sprintf(`{"ociConfigIds":[%d]}`, other))}, + } + return createTasksWithHistory(t, tasks, inputs) +} + +func createTasksWithHistory(t *testing.T, tasks *TaskService, inputs []CreateTaskInput) []model.Task { + t.Helper() + out := make([]model.Task, 0, len(inputs)) + for _, input := range inputs { + task, err := tasks.CreateTask(context.Background(), input) + if err != nil { + t.Fatalf("create task %s: %v", input.Name, err) + } + tasks.db.Model(task).Updates(map[string]any{"last_error": "old", "run_count": 3, "last_run_at": time.Now()}) + mustCreate(t, tasks.db, &model.TaskLog{TaskID: task.ID, Message: "old"}) + mustCreate(t, tasks.db, &model.Setting{Key: deadAliasKey(task.ID), Value: `["target"]`}) + out = append(out, *task) + } + return out +} + +func assertDeletedTasks(t *testing.T, db *gorm.DB, tasks *TaskService, deleted []model.Task) { + t.Helper() + for _, task := range deleted { + assertCount(t, db, &model.Task{}, "id = ?", []any{task.ID}, 0) + assertCount(t, db, &model.TaskLog{}, "task_id = ?", []any{task.ID}, 0) + assertCount(t, db, &model.Setting{}, "key = ?", []any{deadAliasKey(task.ID)}, 0) + if _, ok := tasks.entries[task.ID]; ok { + t.Errorf("task %d still scheduled", task.ID) + } + } +} + +func assertRewrittenTasks(t *testing.T, db *gorm.DB, tasks []model.Task) { + t.Helper() + wantPayload := []string{`{"ociConfigIds":[2]}`, `{"ociConfigIds":[]}`, `{"ociConfigIds":[2]}`} + for i, original := range tasks { + var got model.Task + if err := db.First(&got, original.ID).Error; err != nil { + t.Fatalf("load task %d: %v", original.ID, err) + } + if got.Payload != wantPayload[i] { + t.Errorf("task %d payload = %s, want %s", got.ID, got.Payload, wantPayload[i]) + } + wantHistory := int64(0) + if original.Name == "other" { + wantHistory = 1 + } else if got.RunCount != 0 || got.LastRunAt != nil || got.LastError != "" { + t.Errorf("task %d history fields not reset: %+v", got.ID, got) + } + assertCount(t, db, &model.TaskLog{}, "task_id = ?", []any{got.ID}, wantHistory) + assertCount(t, db, &model.Setting{}, "key = ?", []any{deadAliasKey(got.ID)}, wantHistory) + } +} + +func TestDeleteTenantRollback(t *testing.T) { + configs, tasks, db := newTenantDeleteEnv(t, &fakeClient{}) + target, _ := seedDeleteTenants(t, db) + mustCreate(t, db, &model.CheckSnapshot{OciConfigID: target.ID}) + task := createHealthTask(t, tasks, target.ID) + seedTaskHistory(t, db, task.ID) + registerDeleteFailure(t, db, "cost_snapshots") + err := configs.Delete(context.Background(), target.ID) + if err == nil || !errors.Is(err, errInjectedTenantDelete) { + t.Fatalf("delete error = %v, want injected failure", err) + } + assertCount(t, db, &model.OciConfig{}, "id = ?", []any{target.ID}, 1) + assertCount(t, db, &model.CheckSnapshot{}, "oci_config_id = ?", []any{target.ID}, 1) + assertCount(t, db, &model.Task{}, "id = ?", []any{task.ID}, 1) + assertCount(t, db, &model.TaskLog{}, "task_id = ?", []any{task.ID}, 1) + assertCount(t, db, &model.Setting{}, "key = ?", []any{deadAliasKey(task.ID)}, 1) +} + +func TestDeleteTenantWaitsForRunningTask(t *testing.T) { + client := &blockingClient{fakeClient: &fakeClient{}, started: make(chan struct{}, 1), release: make(chan struct{})} + configs, tasks, db := newTenantDeleteEnv(t, client) + target, _ := seedDeleteTenants(t, db) + setTenantPrivateKey(t, configs, target.ID) + task := createHealthTask(t, tasks, target.ID) + executed := make(chan struct{}) + go func() { + tasks.execute(task.ID) + close(executed) + }() + <-client.started + deleted := make(chan error, 1) + go func() { deleted <- configs.Delete(context.Background(), target.ID) }() + assertDeleteBlocked(t, deleted) + close(client.release) + <-executed + if err := <-deleted; err != nil { + t.Fatalf("delete tenant: %v", err) + } + assertCount(t, db, &model.Task{}, "id = ?", []any{task.ID}, 0) + assertCount(t, db, &model.TaskLog{}, "task_id = ?", []any{task.ID}, 0) +} + +func TestVerifyDoesNotResurrectDeletedTenant(t *testing.T) { + client := &blockingClient{fakeClient: &fakeClient{}, started: make(chan struct{}, 1), release: make(chan struct{})} + configs, _, db := newTenantDeleteEnv(t, client) + target, _ := seedDeleteTenants(t, db) + setTenantPrivateKey(t, configs, target.ID) + verified := make(chan error, 1) + go func() { + _, _, err := configs.Verify(context.Background(), target.ID) + verified <- err + }() + <-client.started + if err := configs.Delete(context.Background(), target.ID); err != nil { + t.Fatalf("delete tenant: %v", err) + } + close(client.release) + if err := <-verified; !errors.Is(err, gorm.ErrRecordNotFound) { + t.Fatalf("verify error = %v, want record not found", err) + } + assertCount(t, db, &model.OciConfig{}, "id = ?", []any{target.ID}, 0) +} + +func TestScopeCacheRejectsDeletedTenant(t *testing.T) { + configs, _, db := newTenantDeleteEnv(t, &fakeClient{}) + target, _ := seedDeleteTenants(t, db) + if err := configs.Delete(context.Background(), target.ID); err != nil { + t.Fatalf("delete tenant: %v", err) + } + err := configs.saveRegionCache(context.Background(), target.ID, + []oci.RegionSubscription{{Key: "PHX", Name: "us-phoenix-1"}}) + if !errors.Is(err, gorm.ErrRecordNotFound) { + t.Fatalf("save region cache error = %v, want record not found", err) + } + err = configs.saveCompartmentCache(context.Background(), target.ID, + []oci.Compartment{{ID: "compartment", Name: "deleted"}}) + if !errors.Is(err, gorm.ErrRecordNotFound) { + t.Fatalf("save compartment cache error = %v, want record not found", err) + } + assertCount(t, db, &model.RegionCache{}, "oci_config_id = ?", []any{target.ID}, 0) + assertCount(t, db, &model.CompartmentCache{}, "oci_config_id = ?", []any{target.ID}, 0) +} + +func setTenantPrivateKey(t *testing.T, configs *OciConfigService, id uint) { + t.Helper() + encrypted, err := configs.cipher.EncryptString("private-key") + if err != nil { + t.Fatalf("encrypt private key: %v", err) + } + if err := configs.db.Model(&model.OciConfig{}).Where("id = ?", id). + Update("private_key_enc", encrypted).Error; err != nil { + t.Fatalf("set private key: %v", err) + } +} + +func createHealthTask(t *testing.T, tasks *TaskService, cfgID uint) *model.Task { + t.Helper() + payload := []byte(fmt.Sprintf(`{"ociConfigIds":[%d]}`, cfgID)) + task, err := tasks.CreateTask(context.Background(), CreateTaskInput{ + Name: "running", Type: model.TaskTypeHealthCheck, + CronExpr: "0 0 * * *", Payload: payload, + }) + if err != nil { + t.Fatalf("create task: %v", err) + } + return task +} + +func assertDeleteBlocked(t *testing.T, deleted <-chan error) { + t.Helper() + select { + case err := <-deleted: + t.Fatalf("delete returned before running task finished: %v", err) + case <-time.After(50 * time.Millisecond): + } +} + +func TestPersistTaskRunDoesNotResurrectDeletedTask(t *testing.T) { + _, tasks, db := newTenantDeleteEnv(t, &fakeClient{}) + task := &model.Task{Name: "stale", Type: model.TaskTypeCost, Status: model.TaskStatusActive} + mustCreate(t, db, task) + stale := *task + if err := db.Delete(task).Error; err != nil { + t.Fatalf("delete task: %v", err) + } + stored, err := tasks.persistTaskRun(context.Background(), &stale) + if err != nil || stored { + t.Fatalf("persist stale task = (%v, %v), want (false, nil)", stored, err) + } + assertCount(t, db, &model.Task{}, "id = ?", []any{task.ID}, 0) +} + +func TestAiLogsRejectDeletedTenantParents(t *testing.T) { + configs, _, db := newTenantDeleteEnv(t, &fakeClient{}) + target, _ := seedDeleteTenants(t, db) + channel := model.AiChannel{Name: "target", OciConfigID: target.ID, Region: "r1"} + mustCreate(t, db, &channel) + gw := NewAiGatewayService(db, configs, &fakeClient{}) + callID := gw.LogCall(model.AiCallLog{ChannelID: channel.ID, ChannelName: channel.Name}) + if callID == 0 { + t.Fatal("initial call log was not created") + } + if err := configs.Delete(context.Background(), target.ID); err != nil { + t.Fatalf("delete tenant: %v", err) + } + lateID := gw.LogCall(model.AiCallLog{ChannelID: channel.ID, ChannelName: channel.Name}) + if lateID != 0 { + t.Errorf("late call ID = %d, want 0", lateID) + } + gw.LogContent(model.AiContentLog{CallLogID: callID, RequestBody: "late"}) + gw.LogContent(model.AiContentLog{CallLogID: 0, RequestBody: "orphan"}) + assertCount(t, db, &model.AiCallLog{}, "channel_id = ?", []any{channel.ID}, 0) + assertCount(t, db, &model.AiContentLog{}, "", nil, 0) +} + +var errInjectedTenantDelete = errors.New("injected tenant delete failure") + +func registerDeleteFailure(t *testing.T, db *gorm.DB, table string) { + t.Helper() + name := "test:tenant_delete_failure" + err := db.Callback().Delete().Before("gorm:delete").Register(name, func(tx *gorm.DB) { + if tx.Statement.Table == table { + tx.AddError(errInjectedTenantDelete) + } + }) + if err != nil { + t.Fatalf("register callback: %v", err) + } + t.Cleanup(func() { _ = db.Callback().Delete().Remove(name) }) +} + +func mustCreate(t *testing.T, db *gorm.DB, value any) { + t.Helper() + if err := db.Create(value).Error; err != nil { + t.Fatalf("create %T: %v", value, err) + } +} + +func assertCount(t *testing.T, db *gorm.DB, value any, query string, args []any, want int64) { + t.Helper() + q := db.Model(value) + if query != "" { + q = q.Where(query, args...) + } + var got int64 + if err := q.Count(&got).Error; err != nil { + t.Fatalf("count %T: %v", value, err) + } + if got != want { + t.Errorf("count %T = %d, want %d", value, got, want) + } +} diff --git a/internal/service/totp.go b/internal/service/totp.go index 05d2e10..2145599 100644 --- a/internal/service/totp.go +++ b/internal/service/totp.go @@ -102,7 +102,8 @@ func (s *AuthService) ActivateTotp(ctx context.Context, username, code string) e s.totpMu.Lock() delete(s.totpPending, username) s.totpMu.Unlock() - return nil + // 两步验证形态变更:旧令牌全部失效 + return s.bumpTokenVersion(ctx, username) } // DisableTotp 停用两步验证;需当前验证码或登录密码任一确认。 @@ -122,7 +123,8 @@ func (s *AuthService) DisableTotp(ctx context.Context, username, password, code if err != nil { return fmt.Errorf("clear totp secret: %w", err) } - return nil + // 两步验证形态变更:旧令牌全部失效 + return s.bumpTokenVersion(ctx, username) } // confirmDisable 校验停用凭证:验证码或密码任一通过即可。 diff --git a/internal/service/totp_oauth_test.go b/internal/service/totp_oauth_test.go index fdf835a..16e8dc5 100644 --- a/internal/service/totp_oauth_test.go +++ b/internal/service/totp_oauth_test.go @@ -176,7 +176,7 @@ func TestOAuthBindLoginUnbind(t *testing.T) { if display != "octocat" { t.Errorf("display = %q", display) } - if username, err := auth.ParseToken(token); err != nil || username != "admin" { + if username, err := auth.ParseToken(context.Background(), token); err != nil || username != "admin" { t.Errorf("token 应属 admin, got %q (%v)", username, err) } // 未绑定身份拒绝登录