发布 0.1.0:通知渠道、告警规则、令牌版本与安全加固
CI / test (push) Successful in 30s
Release / release (push) Successful in 49s

This commit is contained in:
Wang Defa
2026-07-10 17:38:34 +08:00
parent 4af6a0ca92
commit dbba1f4905
78 changed files with 6898 additions and 551 deletions
+5 -4
View File
@@ -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)
+38 -10
View File
@@ -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 组内) ----
+16 -8
View File
@@ -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
}
+122
View File
@@ -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 侧链路状态与关键事件清单
+14 -2
View File
@@ -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 {
+26 -9
View File
@@ -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{
+54
View File
@@ -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)
}
})
}
}
+1 -1
View File
@@ -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)
+2 -1
View File
@@ -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)
+102
View File
@@ -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)
}
}
})
}
}
+1 -1
View File
@@ -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)
+4 -3
View File
@@ -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) })
}
+1
View File
@@ -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)
+3
View File
@@ -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)
+4
View File
@@ -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)
+76
View File
@@ -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 返回全部通知事件开关
+53 -22
View File
@@ -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
+1 -1
View File
@@ -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
}
+17 -2
View File
@@ -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{},
)
+36 -7
View File
@@ -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 {
+90
View File
@@ -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
+50
View File
@@ -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: `<html>`, 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)
}
})
}
}
+113 -1
View File
@@ -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)。
+45
View File
@@ -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)
}
})
}
}
+16 -1
View File
@@ -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) {
+15
View File
@@ -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)
}
}
+31 -22
View File
@@ -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)
// FederationSAML 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)
+88 -25
View File
@@ -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
}
+25 -20
View File
@@ -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
}
+107
View File
@@ -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)
}
})
}
}
@@ -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)
}
+12 -12
View File
@@ -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)
+6 -6
View File
@@ -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
}
+147 -28
View File
@@ -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
}
+1 -1
View File
@@ -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,
},
+36 -2
View File
@@ -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)
}
}
+1
View File
@@ -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 {
+305
View File
@@ -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)
}
}
+345
View File
@@ -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)
}
}
+65 -50
View File
@@ -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
}
+111 -78
View File
@@ -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)
}
}
+58 -12
View File
@@ -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
}
+65 -8
View File
@@ -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)
}
}
+34 -16
View File
@@ -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) {
+4 -4
View File
@@ -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 {
+23 -23
View File
@@ -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)
}
+133 -40
View File
@@ -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 先删过期记录,再对超量部分删最旧;阈值参数化便于测试。
+148 -4
View File
@@ -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()
+96
View File
@@ -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()
+105
View File
@@ -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("最早锁定条目未被驱逐")
}
}
+134 -18
View File
@@ -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 等待全部在途异步通知发完,供退出前收尾与测试同步。
+218
View File
@@ -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()
}
+339
View File
@@ -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{}
}
+23 -3
View File
@@ -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 通知"
}
+9 -3
View File
@@ -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)
}
+55 -14
View File
@@ -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)
+55
View File
@@ -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)
}
}
+1 -1
View File
@@ -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++
}
}
+6
View File
@@ -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
}
+309
View File
@@ -0,0 +1,309 @@
package service
import (
"context"
"fmt"
"strconv"
"strings"
)
// 通知渠道配置键,前缀 notify_ch_<type>_;标注「密文」的键 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
}
+240 -12
View File
@@ -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:
+103 -1
View File
@@ -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)
}
}
}
+37 -28
View File
@@ -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)
}
+457
View File
@@ -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)
}
}
+609
View File
@@ -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)
}
}
+4 -2
View File
@@ -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 校验停用凭证:验证码或密码任一通过即可。
+1 -1
View File
@@ -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)
}
// 未绑定身份拒绝登录