mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-21 14:19:18 +08:00
feat(openai): cyber_policy 硬阻断全链路透传、审计与计费
上游对单次请求下发 error.code=cyber_policy 硬阻断时,网关在所有端点 (/v1/responses、/v1/chat/completions、/v1/messages、WebSocket)及流式/ 非流式路径下,将该结果原样透传给客户端,绝不 failover、换号或同步拦截; 命中后异步完成审计与计费: - 风控中心记录 cyber_policy 留痕并发送通知邮件,落库先于发信,SMTP 阻塞 不影响留痕 - ops 错误请求记录,状态码对齐客户端实际接收(流式 200 / 非流式 400) - 用量明细标记 request_type=cyber,按上游真实 token 计费,HTTP 与 WebSocket 计费口径统一,零 token 命中不误扣 - 会话级自动屏蔽(管理员开关,默认关):命中的会话在可配 TTL 内本地拦截 不再发往上游,仅屏蔽该会话不影响同 Key 其他会话 - 封号计数排除开关:可选让 cyber 命中不计入自动封号,命中当次不判定且 历史行在违规计数中一并排除 WebSocket 多轮连接下 cyber 标记按 turn 生命周期管理,逐轮独立检测与记录; 透传的错误响应不被兜底逻辑追加内容污染。 Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 4.8
parent
e34ad2b194
commit
b62b573f7f
@@ -246,7 +246,7 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
|
|||||||
userMsgQueueCache := repository.NewUserMsgQueueCache(redisClient)
|
userMsgQueueCache := repository.NewUserMsgQueueCache(redisClient)
|
||||||
userMessageQueueService := service.ProvideUserMessageQueueService(userMsgQueueCache, rpmCache, configConfig)
|
userMessageQueueService := service.ProvideUserMessageQueueService(userMsgQueueCache, rpmCache, configConfig)
|
||||||
gatewayHandler := handler.NewGatewayHandler(gatewayService, geminiMessagesCompatService, antigravityGatewayService, userService, concurrencyService, billingCacheService, usageService, apiKeyService, usageRecordWorkerPool, errorPassthroughService, contentModerationService, userMessageQueueService, configConfig, settingService)
|
gatewayHandler := handler.NewGatewayHandler(gatewayService, geminiMessagesCompatService, antigravityGatewayService, userService, concurrencyService, billingCacheService, usageService, apiKeyService, usageRecordWorkerPool, errorPassthroughService, contentModerationService, userMessageQueueService, configConfig, settingService)
|
||||||
openAIGatewayHandler := handler.NewOpenAIGatewayHandler(openAIGatewayService, concurrencyService, billingCacheService, apiKeyService, usageRecordWorkerPool, errorPassthroughService, contentModerationService, configConfig)
|
openAIGatewayHandler := handler.NewOpenAIGatewayHandler(openAIGatewayService, concurrencyService, billingCacheService, apiKeyService, usageRecordWorkerPool, errorPassthroughService, contentModerationService, opsService, configConfig)
|
||||||
handlerSettingHandler := handler.ProvideSettingHandler(settingService, buildInfo, notificationEmailService)
|
handlerSettingHandler := handler.ProvideSettingHandler(settingService, buildInfo, notificationEmailService)
|
||||||
totpHandler := handler.NewTotpHandler(totpService)
|
totpHandler := handler.NewTotpHandler(totpService)
|
||||||
handlerPaymentHandler := handler.NewPaymentHandler(paymentService, paymentConfigService, channelService)
|
handlerPaymentHandler := handler.NewPaymentHandler(paymentService, paymentConfigService, channelService)
|
||||||
|
|||||||
@@ -43,6 +43,9 @@ type contentModerationConfigRequest struct {
|
|||||||
AutoBanEnabled *bool `json:"auto_ban_enabled"`
|
AutoBanEnabled *bool `json:"auto_ban_enabled"`
|
||||||
BanThreshold *int `json:"ban_threshold"`
|
BanThreshold *int `json:"ban_threshold"`
|
||||||
ViolationWindowHours *int `json:"violation_window_hours"`
|
ViolationWindowHours *int `json:"violation_window_hours"`
|
||||||
|
// cyber_policy 命中是否排除出自动封号计数;前端 RiskControlView 已发送该字段,
|
||||||
|
// service.UpdateContentModerationConfigInput 已支持,此前 handler 层缺透传导致开关静默失效。
|
||||||
|
CyberPolicyExcludeFromBanCount *bool `json:"cyber_policy_exclude_from_ban_count"`
|
||||||
RetryCount *int `json:"retry_count"`
|
RetryCount *int `json:"retry_count"`
|
||||||
HitRetentionDays *int `json:"hit_retention_days"`
|
HitRetentionDays *int `json:"hit_retention_days"`
|
||||||
NonHitRetentionDays *int `json:"non_hit_retention_days"`
|
NonHitRetentionDays *int `json:"non_hit_retention_days"`
|
||||||
@@ -104,6 +107,7 @@ func (h *ContentModerationHandler) UpdateConfig(c *gin.Context) {
|
|||||||
AutoBanEnabled: req.AutoBanEnabled,
|
AutoBanEnabled: req.AutoBanEnabled,
|
||||||
BanThreshold: req.BanThreshold,
|
BanThreshold: req.BanThreshold,
|
||||||
ViolationWindowHours: req.ViolationWindowHours,
|
ViolationWindowHours: req.ViolationWindowHours,
|
||||||
|
CyberPolicyExcludeFromBanCount: req.CyberPolicyExcludeFromBanCount,
|
||||||
RetryCount: req.RetryCount,
|
RetryCount: req.RetryCount,
|
||||||
HitRetentionDays: req.HitRetentionDays,
|
HitRetentionDays: req.HitRetentionDays,
|
||||||
NonHitRetentionDays: req.NonHitRetentionDays,
|
NonHitRetentionDays: req.NonHitRetentionDays,
|
||||||
|
|||||||
@@ -228,6 +228,8 @@ func (h *SettingHandler) GetSettings(c *gin.Context) {
|
|||||||
DefaultConcurrency: settings.DefaultConcurrency,
|
DefaultConcurrency: settings.DefaultConcurrency,
|
||||||
DefaultBalance: settings.DefaultBalance,
|
DefaultBalance: settings.DefaultBalance,
|
||||||
RiskControlEnabled: settings.RiskControlEnabled,
|
RiskControlEnabled: settings.RiskControlEnabled,
|
||||||
|
CyberSessionBlockEnabled: settings.CyberSessionBlockEnabled,
|
||||||
|
CyberSessionBlockTTLSeconds: settings.CyberSessionBlockTTLSeconds,
|
||||||
AffiliateRebateRate: settings.AffiliateRebateRate,
|
AffiliateRebateRate: settings.AffiliateRebateRate,
|
||||||
AffiliateRebateFreezeHours: settings.AffiliateRebateFreezeHours,
|
AffiliateRebateFreezeHours: settings.AffiliateRebateFreezeHours,
|
||||||
AffiliateRebateDurationDays: settings.AffiliateRebateDurationDays,
|
AffiliateRebateDurationDays: settings.AffiliateRebateDurationDays,
|
||||||
@@ -646,6 +648,10 @@ type UpdateSettingsRequest struct {
|
|||||||
// 风控中心功能开关
|
// 风控中心功能开关
|
||||||
RiskControlEnabled *bool `json:"risk_control_enabled"`
|
RiskControlEnabled *bool `json:"risk_control_enabled"`
|
||||||
|
|
||||||
|
// cyber 会话屏蔽开关 + TTL
|
||||||
|
CyberSessionBlockEnabled *bool `json:"cyber_session_block_enabled"`
|
||||||
|
CyberSessionBlockTTLSeconds *int `json:"cyber_session_block_ttl_seconds"`
|
||||||
|
|
||||||
// OpenAI fast/flex policy (optional, only updated when provided)
|
// OpenAI fast/flex policy (optional, only updated when provided)
|
||||||
OpenAIFastPolicySettings *dto.OpenAIFastPolicySettings `json:"openai_fast_policy_settings,omitempty"`
|
OpenAIFastPolicySettings *dto.OpenAIFastPolicySettings `json:"openai_fast_policy_settings,omitempty"`
|
||||||
|
|
||||||
@@ -1462,6 +1468,12 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// cyber 会话屏蔽 TTL 校验:提供时必须 > 0
|
||||||
|
if req.CyberSessionBlockTTLSeconds != nil && *req.CyberSessionBlockTTLSeconds <= 0 {
|
||||||
|
response.BadRequest(c, "cyber_session_block_ttl_seconds must be > 0")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
settings := &service.SystemSettings{
|
settings := &service.SystemSettings{
|
||||||
// 系统全局 platform quota 默认值(整体替换语义)
|
// 系统全局 platform quota 默认值(整体替换语义)
|
||||||
DefaultPlatformQuotas: req.DefaultPlatformQuotas,
|
DefaultPlatformQuotas: req.DefaultPlatformQuotas,
|
||||||
@@ -1769,6 +1781,18 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) {
|
|||||||
}
|
}
|
||||||
return previousSettings.RiskControlEnabled
|
return previousSettings.RiskControlEnabled
|
||||||
}(),
|
}(),
|
||||||
|
CyberSessionBlockEnabled: func() bool {
|
||||||
|
if req.CyberSessionBlockEnabled != nil {
|
||||||
|
return *req.CyberSessionBlockEnabled
|
||||||
|
}
|
||||||
|
return previousSettings.CyberSessionBlockEnabled
|
||||||
|
}(),
|
||||||
|
CyberSessionBlockTTLSeconds: func() int {
|
||||||
|
if req.CyberSessionBlockTTLSeconds != nil {
|
||||||
|
return *req.CyberSessionBlockTTLSeconds
|
||||||
|
}
|
||||||
|
return previousSettings.CyberSessionBlockTTLSeconds
|
||||||
|
}(),
|
||||||
}
|
}
|
||||||
|
|
||||||
// req.AuthSourceXxxPlatformQuotas 为 nil 表示本次请求未包含该 source 的 quota 配置(保留 previousAuthSourceDefaults 中的值);
|
// req.AuthSourceXxxPlatformQuotas 为 nil 表示本次请求未包含该 source 的 quota 配置(保留 previousAuthSourceDefaults 中的值);
|
||||||
@@ -2091,6 +2115,8 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) {
|
|||||||
AffiliateEnabled: updatedSettings.AffiliateEnabled,
|
AffiliateEnabled: updatedSettings.AffiliateEnabled,
|
||||||
|
|
||||||
RiskControlEnabled: updatedSettings.RiskControlEnabled,
|
RiskControlEnabled: updatedSettings.RiskControlEnabled,
|
||||||
|
CyberSessionBlockEnabled: updatedSettings.CyberSessionBlockEnabled,
|
||||||
|
CyberSessionBlockTTLSeconds: updatedSettings.CyberSessionBlockTTLSeconds,
|
||||||
AllowUserViewErrorRequests: updatedSettings.AllowUserViewErrorRequests,
|
AllowUserViewErrorRequests: updatedSettings.AllowUserViewErrorRequests,
|
||||||
}
|
}
|
||||||
if fastPolicy, err := h.settingService.GetOpenAIFastPolicySettings(c.Request.Context()); err != nil {
|
if fastPolicy, err := h.settingService.GetOpenAIFastPolicySettings(c.Request.Context()); err != nil {
|
||||||
@@ -2572,6 +2598,12 @@ func diffSettings(before *service.SystemSettings, after *service.SystemSettings,
|
|||||||
if before.RiskControlEnabled != after.RiskControlEnabled {
|
if before.RiskControlEnabled != after.RiskControlEnabled {
|
||||||
changed = append(changed, "risk_control_enabled")
|
changed = append(changed, "risk_control_enabled")
|
||||||
}
|
}
|
||||||
|
if before.CyberSessionBlockEnabled != after.CyberSessionBlockEnabled {
|
||||||
|
changed = append(changed, "cyber_session_block_enabled")
|
||||||
|
}
|
||||||
|
if before.CyberSessionBlockTTLSeconds != after.CyberSessionBlockTTLSeconds {
|
||||||
|
changed = append(changed, "cyber_session_block_ttl_seconds")
|
||||||
|
}
|
||||||
// Default platform quotas(JSON map,整体比较)
|
// Default platform quotas(JSON map,整体比较)
|
||||||
if !equalPlatformQuotaSettings(before.DefaultPlatformQuotas, after.DefaultPlatformQuotas) {
|
if !equalPlatformQuotaSettings(before.DefaultPlatformQuotas, after.DefaultPlatformQuotas) {
|
||||||
changed = append(changed, service.SettingKeyDefaultPlatformQuotas)
|
changed = append(changed, service.SettingKeyDefaultPlatformQuotas)
|
||||||
|
|||||||
@@ -244,6 +244,10 @@ type SystemSettings struct {
|
|||||||
// 风控中心功能开关
|
// 风控中心功能开关
|
||||||
RiskControlEnabled bool `json:"risk_control_enabled"`
|
RiskControlEnabled bool `json:"risk_control_enabled"`
|
||||||
|
|
||||||
|
// cyber 会话屏蔽开关 + TTL
|
||||||
|
CyberSessionBlockEnabled bool `json:"cyber_session_block_enabled"`
|
||||||
|
CyberSessionBlockTTLSeconds int `json:"cyber_session_block_ttl_seconds"`
|
||||||
|
|
||||||
// Affiliate (邀请返利) feature switch
|
// Affiliate (邀请返利) feature switch
|
||||||
AffiliateEnabled bool `json:"affiliate_enabled"`
|
AffiliateEnabled bool `json:"affiliate_enabled"`
|
||||||
|
|
||||||
|
|||||||
@@ -89,6 +89,9 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) {
|
|||||||
h.errorResponse(c, contentModerationStatus(decision), contentModerationErrorCode(decision), decision.Message)
|
h.errorResponse(c, contentModerationStatus(decision), contentModerationErrorCode(decision), decision.Message)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
if h.rejectIfCyberSessionBlocked(c, apiKey, body, reqModel, cyberBlockFormatChat) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
// 解析渠道级模型映射
|
// 解析渠道级模型映射
|
||||||
channelMapping, _ := h.gatewayService.ResolveChannelMappingAndRestrict(c.Request.Context(), apiKey.GroupID, reqModel)
|
channelMapping, _ := h.gatewayService.ResolveChannelMappingAndRestrict(c.Request.Context(), apiKey.GroupID, reqModel)
|
||||||
@@ -192,6 +195,11 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) {
|
|||||||
}()
|
}()
|
||||||
return h.gatewayService.ForwardAsChatCompletions(c.Request.Context(), c, account, forwardBody, promptCacheKey, "")
|
return h.gatewayService.ForwardAsChatCompletions(c.Request.Context(), c, account, forwardBody, promptCacheKey, "")
|
||||||
}()
|
}()
|
||||||
|
cyberBlockKeyChat := ""
|
||||||
|
if service.GetOpsCyberPolicy(c) != nil {
|
||||||
|
cyberBlockKeyChat = service.CyberSessionBlockKey(apiKey.ID, c, body)
|
||||||
|
}
|
||||||
|
h.recordCyberPolicyIfMarked(c, apiKey, account, subscription, reqModel, err != nil, cyberBlockKeyChat, channelMapping.ToUsageFields(reqModel, ""), service.HashUsageRequestPayload(body))
|
||||||
|
|
||||||
forwardDurationMs := time.Since(forwardStart).Milliseconds()
|
forwardDurationMs := time.Since(forwardStart).Milliseconds()
|
||||||
upstreamLatencyMs, _ := getContextInt64(c, service.OpsUpstreamLatencyMsKey)
|
upstreamLatencyMs, _ := getContextInt64(c, service.OpsUpstreamLatencyMsKey)
|
||||||
@@ -283,6 +291,7 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) {
|
|||||||
inboundEndpoint := GetInboundEndpoint(c)
|
inboundEndpoint := GetInboundEndpoint(c)
|
||||||
upstreamEndpoint := resolveRawCCUpstreamEndpoint(c, account)
|
upstreamEndpoint := resolveRawCCUpstreamEndpoint(c, account)
|
||||||
|
|
||||||
|
cyberBlocked := service.GetOpsCyberPolicy(c) != nil
|
||||||
h.submitOpenAIUsageRecordTask(c.Request.Context(), result, func(ctx context.Context) {
|
h.submitOpenAIUsageRecordTask(c.Request.Context(), result, func(ctx context.Context) {
|
||||||
if err := h.gatewayService.RecordUsage(ctx, &service.OpenAIRecordUsageInput{
|
if err := h.gatewayService.RecordUsage(ctx, &service.OpenAIRecordUsageInput{
|
||||||
Result: result,
|
Result: result,
|
||||||
@@ -296,6 +305,7 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) {
|
|||||||
IPAddress: clientIP,
|
IPAddress: clientIP,
|
||||||
APIKeyService: h.apiKeyService,
|
APIKeyService: h.apiKeyService,
|
||||||
ChannelUsageFields: channelMapping.ToUsageFields(reqModel, result.UpstreamModel),
|
ChannelUsageFields: channelMapping.ToUsageFields(reqModel, result.UpstreamModel),
|
||||||
|
CyberBlocked: cyberBlocked,
|
||||||
}); err != nil {
|
}); err != nil {
|
||||||
logger.L().With(
|
logger.L().With(
|
||||||
zap.String("component", "handler.openai_gateway.chat_completions"),
|
zap.String("component", "handler.openai_gateway.chat_completions"),
|
||||||
|
|||||||
@@ -0,0 +1,180 @@
|
|||||||
|
package handler
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/http/httptest"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
// newTestGinContext builds a bare gin.Context backed by an httptest recorder.
|
||||||
|
func newTestGinContext() *gin.Context {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(w)
|
||||||
|
return c
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestRecordCyberPolicyIfMarked_NoMark verifies that when no cyber mark is set,
|
||||||
|
// the function returns immediately and does NOT set the recorded flag.
|
||||||
|
func TestRecordCyberPolicyIfMarked_NoMark(t *testing.T) {
|
||||||
|
c := newTestGinContext()
|
||||||
|
h := &OpenAIGatewayHandler{}
|
||||||
|
|
||||||
|
h.recordCyberPolicyIfMarked(c, nil, nil, nil, "gpt-5", true, "", service.ChannelUsageFields{}, "")
|
||||||
|
|
||||||
|
// Flag must NOT be set when there was no mark.
|
||||||
|
require.False(t, c.GetBool(cyberPolicyRecordedKey),
|
||||||
|
"cyberPolicyRecordedKey must remain false when no cyber mark is present")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestRecordCyberPolicyIfMarked_WithMark verifies that:
|
||||||
|
// 1. When a cyber mark is present, the recorded flag is set (guard activated).
|
||||||
|
// 2. A second call is a no-op (idempotent guard).
|
||||||
|
// 3. Nil services do not panic.
|
||||||
|
func TestRecordCyberPolicyIfMarked_WithMark(t *testing.T) {
|
||||||
|
c := newTestGinContext()
|
||||||
|
service.MarkOpsCyberPolicy(c, service.CyberPolicyMark{
|
||||||
|
Message: "flagged",
|
||||||
|
Body: `{"error":{"code":"cyber_policy"}}`,
|
||||||
|
UpstreamStatus: 400,
|
||||||
|
})
|
||||||
|
|
||||||
|
h := &OpenAIGatewayHandler{} // nil services — must not panic
|
||||||
|
|
||||||
|
// First call: should set the flag.
|
||||||
|
require.NotPanics(t, func() {
|
||||||
|
h.recordCyberPolicyIfMarked(c, nil, nil, nil, "gpt-5", true, "", service.ChannelUsageFields{}, "")
|
||||||
|
})
|
||||||
|
require.True(t, c.GetBool(cyberPolicyRecordedKey),
|
||||||
|
"cyberPolicyRecordedKey must be true after first call with a mark")
|
||||||
|
|
||||||
|
// Second call: flag already set — must be a no-op (idempotent).
|
||||||
|
require.NotPanics(t, func() {
|
||||||
|
h.recordCyberPolicyIfMarked(c, nil, nil, nil, "gpt-5", false, "", service.ChannelUsageFields{}, "")
|
||||||
|
})
|
||||||
|
// Flag should still be true (not toggled or cleared).
|
||||||
|
require.True(t, c.GetBool(cyberPolicyRecordedKey),
|
||||||
|
"cyberPolicyRecordedKey must remain true after second call (guard)")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestRecordCyberPolicyIfMarked_ForwardSuccessSkipsUsageLog verifies the semantic:
|
||||||
|
// when forwardErrored=false the function still sets the guard flag (mark present),
|
||||||
|
// but the cyber usage row is NOT requested (only RecordCyberPolicyEvent fires).
|
||||||
|
// Since services are nil here we only verify the guard flag and no panic.
|
||||||
|
func TestRecordCyberPolicyIfMarked_ForwardSuccessSkipsUsageLog(t *testing.T) {
|
||||||
|
c := newTestGinContext()
|
||||||
|
service.MarkOpsCyberPolicy(c, service.CyberPolicyMark{
|
||||||
|
Message: "flagged",
|
||||||
|
UpstreamStatus: 200,
|
||||||
|
})
|
||||||
|
|
||||||
|
h := &OpenAIGatewayHandler{}
|
||||||
|
|
||||||
|
require.NotPanics(t, func() {
|
||||||
|
h.recordCyberPolicyIfMarked(c, nil, nil, nil, "gpt-5", false /* forwardErrored=false */, "", service.ChannelUsageFields{}, "")
|
||||||
|
})
|
||||||
|
require.True(t, c.GetBool(cyberPolicyRecordedKey))
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestClearCyberPolicyTurnState verifies F1 at the handler level: after a turn
|
||||||
|
// is finalized, both the mark and the recorded guard are reset so the next WS
|
||||||
|
// turn detects/records independently.
|
||||||
|
func TestClearCyberPolicyTurnState(t *testing.T) {
|
||||||
|
c := newTestGinContext()
|
||||||
|
h := &OpenAIGatewayHandler{}
|
||||||
|
|
||||||
|
service.MarkOpsCyberPolicy(c, service.CyberPolicyMark{Message: "turn1", UpstreamStatus: 200})
|
||||||
|
h.recordCyberPolicyIfMarked(c, nil, nil, nil, "gpt-5", false, "", service.ChannelUsageFields{}, "")
|
||||||
|
require.True(t, c.GetBool(cyberPolicyRecordedKey))
|
||||||
|
|
||||||
|
clearCyberPolicyTurnState(c)
|
||||||
|
require.Nil(t, service.GetOpsCyberPolicy(c))
|
||||||
|
require.False(t, c.GetBool(cyberPolicyRecordedKey))
|
||||||
|
|
||||||
|
// turn2: a fresh cyber hit must be recordable again.
|
||||||
|
service.MarkOpsCyberPolicy(c, service.CyberPolicyMark{Message: "turn2", UpstreamStatus: 200})
|
||||||
|
h.recordCyberPolicyIfMarked(c, nil, nil, nil, "gpt-5", false, "", service.ChannelUsageFields{}, "")
|
||||||
|
require.True(t, c.GetBool(cyberPolicyRecordedKey))
|
||||||
|
require.Equal(t, "turn2", service.GetOpsCyberPolicy(c).Message)
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestBuildCyberSessionBlockedOpsEntry verifies the locally-rejected request is
|
||||||
|
// auditable: 403 / phase=request / type=cyber_policy_session_blocked — distinct
|
||||||
|
// from upstream cyber_policy hits, and it must NOT touch moderation/violation.
|
||||||
|
func TestBuildCyberSessionBlockedOpsEntry(t *testing.T) {
|
||||||
|
entry := buildCyberSessionBlockedOpsEntry(cyberPolicyOpsErrorMeta{
|
||||||
|
RequestID: "req-9", Model: "gpt-5", RequestPath: "/openai/v1/responses",
|
||||||
|
})
|
||||||
|
require.Equal(t, 403, entry.StatusCode)
|
||||||
|
require.Equal(t, "cyber_policy_session_blocked", entry.ErrorType)
|
||||||
|
require.Equal(t, "request", entry.ErrorPhase)
|
||||||
|
require.True(t, entry.IsBusinessLimited)
|
||||||
|
require.Equal(t, "gateway_local", entry.ErrorSource)
|
||||||
|
require.Equal(t, "platform", entry.ErrorOwner)
|
||||||
|
require.Empty(t, entry.ErrorBody, "no session block key → ErrorBody must be empty")
|
||||||
|
|
||||||
|
entryWithKey := buildCyberSessionBlockedOpsEntry(cyberPolicyOpsErrorMeta{
|
||||||
|
RequestID: "req-9", Model: "gpt-5", RequestPath: "/openai/v1/responses",
|
||||||
|
SessionBlockKey: "abc123",
|
||||||
|
})
|
||||||
|
require.Equal(t, "session_block_key=abc123", entryWithKey.ErrorBody)
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestRejectIfCyberSessionBlocked_FailOpen verifies fail-open paths: nil handler
|
||||||
|
// services, no explicit session signal, and (implicitly) disabled switch all
|
||||||
|
// pass the request through.
|
||||||
|
func TestRejectIfCyberSessionBlocked_FailOpen(t *testing.T) {
|
||||||
|
c := newTestGinContext()
|
||||||
|
c.Request = httptest.NewRequest("POST", "/openai/v1/responses", strings.NewReader(`{}`))
|
||||||
|
|
||||||
|
h := &OpenAIGatewayHandler{}
|
||||||
|
require.False(t, h.rejectIfCyberSessionBlocked(c, nil, []byte(`{}`), "gpt-5", cyberBlockFormatResponses), "nil apiKey → pass")
|
||||||
|
|
||||||
|
h2 := &OpenAIGatewayHandler{gatewayService: nil}
|
||||||
|
key := &service.APIKey{ID: 1}
|
||||||
|
require.False(t, h2.rejectIfCyberSessionBlocked(c, key, []byte(`{}`), "gpt-5", cyberBlockFormatResponses), "nil gateway service → pass")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestRecordCyberPolicyIfMarked_BlockKeyPlumbed verifies the 6th param is
|
||||||
|
// accepted and a non-empty key with nil gateway service does not panic
|
||||||
|
// (write-side guards live in the service layer).
|
||||||
|
func TestRecordCyberPolicyIfMarked_BlockKeyPlumbed(t *testing.T) {
|
||||||
|
c := newTestGinContext()
|
||||||
|
service.MarkOpsCyberPolicy(c, service.CyberPolicyMark{Message: "x", UpstreamStatus: 400})
|
||||||
|
h := &OpenAIGatewayHandler{}
|
||||||
|
require.NotPanics(t, func() {
|
||||||
|
h.recordCyberPolicyIfMarked(c, nil, nil, nil, "gpt-5", true, "deadbeef", service.ChannelUsageFields{}, "")
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestBuildCyberPolicyOpsErrorEntry_StatusCode verifies F6: the ops error log
|
||||||
|
// records the status the codex client actually received (400 non-stream / 200 stream),
|
||||||
|
// not a hardcoded 403.
|
||||||
|
func TestBuildCyberPolicyOpsErrorEntry_StatusCode(t *testing.T) {
|
||||||
|
for _, tc := range []struct {
|
||||||
|
name string
|
||||||
|
upstreamStatus int
|
||||||
|
}{
|
||||||
|
{"non_stream_400", 400},
|
||||||
|
{"stream_200", 200},
|
||||||
|
{"zero_value", 0},
|
||||||
|
} {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
mark := &service.CyberPolicyMark{
|
||||||
|
Code: "cyber_policy",
|
||||||
|
Message: "blocked",
|
||||||
|
UpstreamStatus: tc.upstreamStatus,
|
||||||
|
}
|
||||||
|
entry := buildCyberPolicyOpsErrorEntry(cyberPolicyOpsErrorMeta{
|
||||||
|
RequestID: "req-1", Model: "gpt-5", RequestPath: "/openai/v1/responses",
|
||||||
|
}, mark)
|
||||||
|
require.Equal(t, tc.upstreamStatus, entry.StatusCode)
|
||||||
|
require.Equal(t, "cyber_policy", entry.ErrorType)
|
||||||
|
require.Equal(t, "request", entry.ErrorPhase)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -34,6 +34,7 @@ type OpenAIGatewayHandler struct {
|
|||||||
usageRecordWorkerPool *service.UsageRecordWorkerPool
|
usageRecordWorkerPool *service.UsageRecordWorkerPool
|
||||||
errorPassthroughService *service.ErrorPassthroughService
|
errorPassthroughService *service.ErrorPassthroughService
|
||||||
contentModerationService *service.ContentModerationService
|
contentModerationService *service.ContentModerationService
|
||||||
|
opsService *service.OpsService
|
||||||
concurrencyHelper *ConcurrencyHelper
|
concurrencyHelper *ConcurrencyHelper
|
||||||
imageLimiter *imageConcurrencyLimiter
|
imageLimiter *imageConcurrencyLimiter
|
||||||
maxAccountSwitches int
|
maxAccountSwitches int
|
||||||
@@ -105,6 +106,7 @@ func NewOpenAIGatewayHandler(
|
|||||||
usageRecordWorkerPool *service.UsageRecordWorkerPool,
|
usageRecordWorkerPool *service.UsageRecordWorkerPool,
|
||||||
errorPassthroughService *service.ErrorPassthroughService,
|
errorPassthroughService *service.ErrorPassthroughService,
|
||||||
contentModerationService *service.ContentModerationService,
|
contentModerationService *service.ContentModerationService,
|
||||||
|
opsService *service.OpsService,
|
||||||
cfg *config.Config,
|
cfg *config.Config,
|
||||||
) *OpenAIGatewayHandler {
|
) *OpenAIGatewayHandler {
|
||||||
pingInterval := time.Duration(0)
|
pingInterval := time.Duration(0)
|
||||||
@@ -122,6 +124,7 @@ func NewOpenAIGatewayHandler(
|
|||||||
usageRecordWorkerPool: usageRecordWorkerPool,
|
usageRecordWorkerPool: usageRecordWorkerPool,
|
||||||
errorPassthroughService: errorPassthroughService,
|
errorPassthroughService: errorPassthroughService,
|
||||||
contentModerationService: contentModerationService,
|
contentModerationService: contentModerationService,
|
||||||
|
opsService: opsService,
|
||||||
concurrencyHelper: NewConcurrencyHelper(concurrencyService, SSEPingFormatComment, pingInterval),
|
concurrencyHelper: NewConcurrencyHelper(concurrencyService, SSEPingFormatComment, pingInterval),
|
||||||
imageLimiter: &imageConcurrencyLimiter{},
|
imageLimiter: &imageConcurrencyLimiter{},
|
||||||
maxAccountSwitches: maxAccountSwitches,
|
maxAccountSwitches: maxAccountSwitches,
|
||||||
@@ -305,6 +308,9 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
|
|||||||
|
|
||||||
// Generate session hash (header first; fallback to prompt_cache_key)
|
// Generate session hash (header first; fallback to prompt_cache_key)
|
||||||
sessionHash := h.gatewayService.GenerateSessionHash(c, sessionHashBody)
|
sessionHash := h.gatewayService.GenerateSessionHash(c, sessionHashBody)
|
||||||
|
if h.rejectIfCyberSessionBlocked(c, apiKey, sessionHashBody, reqModel, cyberBlockFormatResponses) {
|
||||||
|
return
|
||||||
|
}
|
||||||
requireCompact := isOpenAIRemoteCompactPath(c)
|
requireCompact := isOpenAIRemoteCompactPath(c)
|
||||||
|
|
||||||
maxAccountSwitches := h.maxAccountSwitches
|
maxAccountSwitches := h.maxAccountSwitches
|
||||||
@@ -387,6 +393,11 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
|
|||||||
}()
|
}()
|
||||||
return h.gatewayService.Forward(c.Request.Context(), c, account, forwardBody)
|
return h.gatewayService.Forward(c.Request.Context(), c, account, forwardBody)
|
||||||
}()
|
}()
|
||||||
|
cyberBlockKeyHTTP := ""
|
||||||
|
if service.GetOpsCyberPolicy(c) != nil {
|
||||||
|
cyberBlockKeyHTTP = service.CyberSessionBlockKey(apiKey.ID, c, sessionHashBody)
|
||||||
|
}
|
||||||
|
h.recordCyberPolicyIfMarked(c, apiKey, account, subscription, reqModel, err != nil, cyberBlockKeyHTTP, channelMapping.ToUsageFields(reqModel, ""), service.HashUsageRequestPayload(body))
|
||||||
forwardDurationMs := time.Since(forwardStart).Milliseconds()
|
forwardDurationMs := time.Since(forwardStart).Milliseconds()
|
||||||
upstreamLatencyMs, _ := getContextInt64(c, service.OpsUpstreamLatencyMsKey)
|
upstreamLatencyMs, _ := getContextInt64(c, service.OpsUpstreamLatencyMsKey)
|
||||||
responseLatencyMs := forwardDurationMs
|
responseLatencyMs := forwardDurationMs
|
||||||
@@ -488,6 +499,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
|
|||||||
upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform)
|
upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform)
|
||||||
|
|
||||||
// 使用量记录通过有界 worker 池提交,避免请求热路径创建无界 goroutine。
|
// 使用量记录通过有界 worker 池提交,避免请求热路径创建无界 goroutine。
|
||||||
|
cyberBlocked := service.GetOpsCyberPolicy(c) != nil
|
||||||
h.submitOpenAIUsageRecordTask(c.Request.Context(), result, func(ctx context.Context) {
|
h.submitOpenAIUsageRecordTask(c.Request.Context(), result, func(ctx context.Context) {
|
||||||
if err := h.gatewayService.RecordUsage(ctx, &service.OpenAIRecordUsageInput{
|
if err := h.gatewayService.RecordUsage(ctx, &service.OpenAIRecordUsageInput{
|
||||||
Result: result,
|
Result: result,
|
||||||
@@ -502,6 +514,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
|
|||||||
RequestPayloadHash: requestPayloadHash,
|
RequestPayloadHash: requestPayloadHash,
|
||||||
APIKeyService: h.apiKeyService,
|
APIKeyService: h.apiKeyService,
|
||||||
ChannelUsageFields: channelMapping.ToUsageFields(reqModel, result.UpstreamModel),
|
ChannelUsageFields: channelMapping.ToUsageFields(reqModel, result.UpstreamModel),
|
||||||
|
CyberBlocked: cyberBlocked,
|
||||||
}); err != nil {
|
}); err != nil {
|
||||||
logger.L().With(
|
logger.L().With(
|
||||||
zap.String("component", "handler.openai_gateway.responses"),
|
zap.String("component", "handler.openai_gateway.responses"),
|
||||||
@@ -713,6 +726,9 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) {
|
|||||||
sessionHash := h.gatewayService.GenerateSessionHash(c, body)
|
sessionHash := h.gatewayService.GenerateSessionHash(c, body)
|
||||||
promptCacheKey := h.gatewayService.ExtractSessionID(c, body)
|
promptCacheKey := h.gatewayService.ExtractSessionID(c, body)
|
||||||
sessionHash, promptCacheKey = resolveOpenAIMessagesMetadataSession(sessionHash, promptCacheKey, reqModel, body)
|
sessionHash, promptCacheKey = resolveOpenAIMessagesMetadataSession(sessionHash, promptCacheKey, reqModel, body)
|
||||||
|
if h.rejectIfCyberSessionBlocked(c, apiKey, body, reqModel, cyberBlockFormatAnthropic) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
maxAccountSwitches := h.maxAccountSwitches
|
maxAccountSwitches := h.maxAccountSwitches
|
||||||
switchCount := 0
|
switchCount := 0
|
||||||
@@ -789,7 +805,11 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) {
|
|||||||
}()
|
}()
|
||||||
return h.gatewayService.ForwardAsAnthropic(c.Request.Context(), c, account, forwardBody, promptCacheKey, defaultMappedModel)
|
return h.gatewayService.ForwardAsAnthropic(c.Request.Context(), c, account, forwardBody, promptCacheKey, defaultMappedModel)
|
||||||
}()
|
}()
|
||||||
|
cyberBlockKeyMsg := ""
|
||||||
|
if service.GetOpsCyberPolicy(c) != nil {
|
||||||
|
cyberBlockKeyMsg = service.CyberSessionBlockKey(apiKey.ID, c, body)
|
||||||
|
}
|
||||||
|
h.recordCyberPolicyIfMarked(c, apiKey, account, subscription, reqModel, err != nil, cyberBlockKeyMsg, channelMappingMsg.ToUsageFields(reqModel, ""), service.HashUsageRequestPayload(body))
|
||||||
forwardDurationMs := time.Since(forwardStart).Milliseconds()
|
forwardDurationMs := time.Since(forwardStart).Milliseconds()
|
||||||
upstreamLatencyMs, _ := getContextInt64(c, service.OpsUpstreamLatencyMsKey)
|
upstreamLatencyMs, _ := getContextInt64(c, service.OpsUpstreamLatencyMsKey)
|
||||||
responseLatencyMs := forwardDurationMs
|
responseLatencyMs := forwardDurationMs
|
||||||
@@ -883,6 +903,7 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) {
|
|||||||
inboundEndpoint := GetInboundEndpoint(c)
|
inboundEndpoint := GetInboundEndpoint(c)
|
||||||
upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform)
|
upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform)
|
||||||
|
|
||||||
|
cyberBlocked := service.GetOpsCyberPolicy(c) != nil
|
||||||
h.submitOpenAIUsageRecordTask(c.Request.Context(), result, func(ctx context.Context) {
|
h.submitOpenAIUsageRecordTask(c.Request.Context(), result, func(ctx context.Context) {
|
||||||
if err := h.gatewayService.RecordUsage(ctx, &service.OpenAIRecordUsageInput{
|
if err := h.gatewayService.RecordUsage(ctx, &service.OpenAIRecordUsageInput{
|
||||||
Result: result,
|
Result: result,
|
||||||
@@ -897,6 +918,7 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) {
|
|||||||
RequestPayloadHash: requestPayloadHash,
|
RequestPayloadHash: requestPayloadHash,
|
||||||
APIKeyService: h.apiKeyService,
|
APIKeyService: h.apiKeyService,
|
||||||
ChannelUsageFields: channelMappingMsg.ToUsageFields(reqModel, result.UpstreamModel),
|
ChannelUsageFields: channelMappingMsg.ToUsageFields(reqModel, result.UpstreamModel),
|
||||||
|
CyberBlocked: cyberBlocked,
|
||||||
}); err != nil {
|
}); err != nil {
|
||||||
logger.L().With(
|
logger.L().With(
|
||||||
zap.String("component", "handler.openai_gateway.messages"),
|
zap.String("component", "handler.openai_gateway.messages"),
|
||||||
@@ -1259,6 +1281,17 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// F5a: 握手层会话屏蔽检查。WS 握手无 body,显式标识仅来自握手 header
|
||||||
|
// (session_id / conversation_id);无标识则放行,连接内仍有本地 flag 兜底。
|
||||||
|
cyberBlockKey := service.CyberSessionBlockKey(apiKey.ID, c, nil)
|
||||||
|
if cyberBlockKey != "" && h.gatewayService.IsCyberSessionBlocked(c.Request.Context(), cyberBlockKey) {
|
||||||
|
writeCyberSessionBlockedWSError(c.Request.Context(), wsConn)
|
||||||
|
closeOpenAIClientWS(wsConn, coderws.StatusPolicyViolation, "session blocked by cyber-security policy")
|
||||||
|
h.enqueueCyberSessionBlockedOpsEntry(c, apiKey, reqModel, cyberBlockKey)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
cyberBlockedThisConn := false
|
||||||
|
|
||||||
// 解析渠道级模型映射
|
// 解析渠道级模型映射
|
||||||
channelMappingWS, _ := h.gatewayService.ResolveChannelMappingAndRestrict(ctx, apiKey.GroupID, reqModel)
|
channelMappingWS, _ := h.gatewayService.ResolveChannelMappingAndRestrict(ctx, apiKey.GroupID, reqModel)
|
||||||
|
|
||||||
@@ -1430,6 +1463,10 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
|
|||||||
return nil
|
return nil
|
||||||
},
|
},
|
||||||
BeforeTurn: func(turn int) error {
|
BeforeTurn: func(turn int) error {
|
||||||
|
// turn==1 的会话屏蔽已由握手层检查覆盖;连接内 flag 只拦截后续 turn。
|
||||||
|
if cyberBlockedThisConn {
|
||||||
|
return service.NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, cyberSessionBlockedClientMsg, nil)
|
||||||
|
}
|
||||||
if turn == 1 {
|
if turn == 1 {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -1461,11 +1498,24 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
|
|||||||
return nil
|
return nil
|
||||||
},
|
},
|
||||||
AfterTurn: func(turn int, result *service.OpenAIForwardResult, turnErr error) {
|
AfterTurn: func(turn int, result *service.OpenAIForwardResult, turnErr error) {
|
||||||
|
// F1: cyber 标记按 turn 生命周期清理——defer 保证任意早返回路径都执行;
|
||||||
|
// CyberBlocked 必须在 submit 前同步预捕获(task 闭包由 worker 池异步执行,
|
||||||
|
// 届时 defer 已清除标记)。
|
||||||
|
defer clearCyberPolicyTurnState(c)
|
||||||
releaseTurnSlots()
|
releaseTurnSlots()
|
||||||
|
h.recordCyberPolicyIfMarked(c, apiKey, account, subscription, reqModel, turnErr != nil, cyberBlockKey, channelMappingWS.ToUsageFields(reqModel, ""), requestPayloadHash)
|
||||||
|
if service.GetOpsCyberPolicy(c) != nil {
|
||||||
|
cyberBlockedThisConn = true
|
||||||
|
}
|
||||||
if turnErr != nil {
|
if turnErr != nil {
|
||||||
if result == nil || result.ImageCount <= 0 {
|
if result == nil || result.ImageCount <= 0 {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
// cyber 命中时该 turn 的用量已由 recordCyberPolicyIfMarked(forwardErrored=true)
|
||||||
|
// 按真实 token 记录,这里不再走下方 RecordUsage,避免对同一 turn 双写/双扣费。
|
||||||
|
if service.GetOpsCyberPolicy(c) != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
reqLog.Warn("openai.websocket_partial_error_with_image_result",
|
reqLog.Warn("openai.websocket_partial_error_with_image_result",
|
||||||
zap.Int64("account_id", account.ID),
|
zap.Int64("account_id", account.ID),
|
||||||
zap.Int("image_count", result.ImageCount),
|
zap.Int("image_count", result.ImageCount),
|
||||||
@@ -1481,6 +1531,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
|
|||||||
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, true, result.FirstTokenMs)
|
h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, true, result.FirstTokenMs)
|
||||||
inboundEndpoint := GetInboundEndpoint(c)
|
inboundEndpoint := GetInboundEndpoint(c)
|
||||||
upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform)
|
upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform)
|
||||||
|
cyberBlocked := service.GetOpsCyberPolicy(c) != nil
|
||||||
h.submitOpenAIUsageRecordTask(ctx, result, func(taskCtx context.Context) {
|
h.submitOpenAIUsageRecordTask(ctx, result, func(taskCtx context.Context) {
|
||||||
if err := h.gatewayService.RecordUsage(taskCtx, &service.OpenAIRecordUsageInput{
|
if err := h.gatewayService.RecordUsage(taskCtx, &service.OpenAIRecordUsageInput{
|
||||||
Result: result,
|
Result: result,
|
||||||
@@ -1495,6 +1546,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
|
|||||||
RequestPayloadHash: requestPayloadHash,
|
RequestPayloadHash: requestPayloadHash,
|
||||||
APIKeyService: h.apiKeyService,
|
APIKeyService: h.apiKeyService,
|
||||||
ChannelUsageFields: channelMappingWS.ToUsageFields(reqModel, result.UpstreamModel),
|
ChannelUsageFields: channelMappingWS.ToUsageFields(reqModel, result.UpstreamModel),
|
||||||
|
CyberBlocked: cyberBlocked,
|
||||||
}); err != nil {
|
}); err != nil {
|
||||||
reqLog.Error("openai.websocket_record_usage_failed",
|
reqLog.Error("openai.websocket_record_usage_failed",
|
||||||
zap.Int64("account_id", account.ID),
|
zap.Int64("account_id", account.ID),
|
||||||
@@ -1904,6 +1956,14 @@ func openAIForwardErrorAlreadyCommunicated(c *gin.Context, writerSizeBeforeForwa
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// cyber_policy 命中时上游原始错误体已透传给客户端(非流式 c.Data 写出 400 body,
|
||||||
|
// 流式写出 response.failed 事件),不能再让 ensureForwardErrorResponse 追加
|
||||||
|
// fallback —— 否则在已写出的完整响应尾部追加 SSE(responses 端点尾随
|
||||||
|
// response.failed、chat 端点尾随 event:error),污染响应体。Size 已变化证明响应确已写出。
|
||||||
|
if service.GetOpsCyberPolicy(c) != nil {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
msg := strings.TrimSpace(err.Error())
|
msg := strings.TrimSpace(err.Error())
|
||||||
for _, prefix := range []string{
|
for _, prefix := range []string{
|
||||||
"upstream response failed:",
|
"upstream response failed:",
|
||||||
@@ -2017,6 +2077,364 @@ func writeContentModerationWSError(ctx context.Context, conn *coderws.Conn, deci
|
|||||||
_ = conn.Write(writeCtx, coderws.MessageText, payload)
|
_ = conn.Write(writeCtx, coderws.MessageText, payload)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// writeCyberSessionBlockedWSError sends an error frame telling the client this
|
||||||
|
// session is blocked by the cyber session block (F5a) before closing.
|
||||||
|
func writeCyberSessionBlockedWSError(ctx context.Context, conn *coderws.Conn) {
|
||||||
|
if conn == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if ctx == nil {
|
||||||
|
ctx = context.Background()
|
||||||
|
}
|
||||||
|
payload, err := json.Marshal(gin.H{
|
||||||
|
"event_id": "evt_cyber_session_blocked",
|
||||||
|
"type": "error",
|
||||||
|
"error": gin.H{
|
||||||
|
"type": "permission_error",
|
||||||
|
"code": "session_blocked_by_cyber_policy",
|
||||||
|
"message": cyberSessionBlockedClientMsg,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
payload = []byte(`{"event_id":"evt_cyber_session_blocked","type":"error","error":{"type":"permission_error","code":"session_blocked_by_cyber_policy","message":"This session is blocked by cyber-security policy, please start a new session"}}`)
|
||||||
|
}
|
||||||
|
writeCtx, cancel := context.WithTimeout(ctx, 2*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
_ = conn.Write(writeCtx, coderws.MessageText, payload)
|
||||||
|
}
|
||||||
|
|
||||||
|
// cyberPolicyRecordedKey guards against double-firing recordCyberPolicyIfMarked
|
||||||
|
// within one request (e.g. in a retry/failover loop).
|
||||||
|
const cyberPolicyRecordedKey = "ops_cyber_recorded"
|
||||||
|
|
||||||
|
// cyberPolicyOpsErrorMeta carries request-scoped fields captured outside the
|
||||||
|
// async goroutine for building the cyber ops_error_logs entry.
|
||||||
|
type cyberPolicyOpsErrorMeta struct {
|
||||||
|
RequestID string
|
||||||
|
ClientRequestID string
|
||||||
|
Platform string
|
||||||
|
Model string
|
||||||
|
RequestPath string
|
||||||
|
Stream bool
|
||||||
|
InboundEndpoint string
|
||||||
|
UserAgent string
|
||||||
|
APIKeyPrefix string
|
||||||
|
UserID int64
|
||||||
|
APIKeyID int64
|
||||||
|
AccountID int64
|
||||||
|
GroupID *int64
|
||||||
|
ClientIP string
|
||||||
|
CreatedAt time.Time
|
||||||
|
SessionBlockKey string
|
||||||
|
}
|
||||||
|
|
||||||
|
// buildCyberPolicyOpsErrorEntry builds the ops_error_logs entry for an upstream
|
||||||
|
// cyber_policy hit. StatusCode mirrors what the codex client actually received
|
||||||
|
// (400 non-stream / 200 stream), per F6.
|
||||||
|
func buildCyberPolicyOpsErrorEntry(meta cyberPolicyOpsErrorMeta, mark *service.CyberPolicyMark) *service.OpsInsertErrorLogInput {
|
||||||
|
rt := int16(service.RequestTypeCyberBlocked)
|
||||||
|
entry := &service.OpsInsertErrorLogInput{
|
||||||
|
RequestID: meta.RequestID,
|
||||||
|
ClientRequestID: meta.ClientRequestID,
|
||||||
|
Platform: meta.Platform,
|
||||||
|
Model: meta.Model,
|
||||||
|
RequestPath: meta.RequestPath,
|
||||||
|
Stream: meta.Stream,
|
||||||
|
InboundEndpoint: meta.InboundEndpoint,
|
||||||
|
RequestType: &rt,
|
||||||
|
UserAgent: meta.UserAgent,
|
||||||
|
APIKeyPrefix: meta.APIKeyPrefix,
|
||||||
|
ErrorPhase: "request",
|
||||||
|
ErrorType: "cyber_policy",
|
||||||
|
Severity: "P3",
|
||||||
|
StatusCode: mark.UpstreamStatus,
|
||||||
|
IsBusinessLimited: true,
|
||||||
|
ErrorMessage: "cyber_policy: " + mark.Message,
|
||||||
|
// 原始 body 直接入队;ops service 落库前统一走 sanitizeErrorBodyForStorage 脱敏与截断。
|
||||||
|
ErrorBody: mark.Body,
|
||||||
|
ErrorSource: "upstream_http",
|
||||||
|
ErrorOwner: "provider",
|
||||||
|
CreatedAt: meta.CreatedAt,
|
||||||
|
}
|
||||||
|
if meta.UserID > 0 {
|
||||||
|
entry.UserID = &meta.UserID
|
||||||
|
}
|
||||||
|
if meta.APIKeyID > 0 {
|
||||||
|
entry.APIKeyID = &meta.APIKeyID
|
||||||
|
}
|
||||||
|
if meta.AccountID > 0 {
|
||||||
|
entry.AccountID = &meta.AccountID
|
||||||
|
}
|
||||||
|
entry.GroupID = meta.GroupID
|
||||||
|
if meta.ClientIP != "" {
|
||||||
|
entry.ClientIP = &meta.ClientIP
|
||||||
|
}
|
||||||
|
return entry
|
||||||
|
}
|
||||||
|
|
||||||
|
// 双语单串:网关客户端面向中英用户,且本错误无 i18n 协商通道。
|
||||||
|
const cyberSessionBlockedClientMsg = "该会话已被网络安全策略屏蔽,请开启新会话 / This session is blocked by cyber-security policy, please start a new session"
|
||||||
|
|
||||||
|
// buildCyberSessionBlockedOpsEntry builds the ops_error_logs entry for a request
|
||||||
|
// rejected locally by the cyber session block (F5a). Distinct error_type from
|
||||||
|
// upstream `cyber_policy`; never feeds moderation logs / violation counting
|
||||||
|
// (the request never reached upstream — see spec).
|
||||||
|
func buildCyberSessionBlockedOpsEntry(meta cyberPolicyOpsErrorMeta) *service.OpsInsertErrorLogInput {
|
||||||
|
rt := int16(service.RequestTypeCyberBlocked)
|
||||||
|
entry := &service.OpsInsertErrorLogInput{
|
||||||
|
RequestID: meta.RequestID,
|
||||||
|
ClientRequestID: meta.ClientRequestID,
|
||||||
|
Platform: meta.Platform,
|
||||||
|
Model: meta.Model,
|
||||||
|
RequestPath: meta.RequestPath,
|
||||||
|
Stream: meta.Stream,
|
||||||
|
InboundEndpoint: meta.InboundEndpoint,
|
||||||
|
RequestType: &rt,
|
||||||
|
UserAgent: meta.UserAgent,
|
||||||
|
APIKeyPrefix: meta.APIKeyPrefix,
|
||||||
|
ErrorPhase: "request",
|
||||||
|
ErrorType: "cyber_policy_session_blocked",
|
||||||
|
Severity: "P3",
|
||||||
|
StatusCode: http.StatusForbidden,
|
||||||
|
IsBusinessLimited: true,
|
||||||
|
ErrorMessage: "cyber_policy_session_blocked: request rejected locally by session block",
|
||||||
|
ErrorSource: "gateway_local",
|
||||||
|
ErrorOwner: "platform",
|
||||||
|
CreatedAt: meta.CreatedAt,
|
||||||
|
// AccountID 有意不设:请求在账号选择前即被拒绝。
|
||||||
|
}
|
||||||
|
if meta.SessionBlockKey != "" {
|
||||||
|
entry.ErrorBody = "session_block_key=" + meta.SessionBlockKey
|
||||||
|
}
|
||||||
|
if meta.UserID > 0 {
|
||||||
|
entry.UserID = &meta.UserID
|
||||||
|
}
|
||||||
|
if meta.APIKeyID > 0 {
|
||||||
|
entry.APIKeyID = &meta.APIKeyID
|
||||||
|
}
|
||||||
|
entry.GroupID = meta.GroupID
|
||||||
|
if meta.ClientIP != "" {
|
||||||
|
entry.ClientIP = &meta.ClientIP
|
||||||
|
}
|
||||||
|
return entry
|
||||||
|
}
|
||||||
|
|
||||||
|
// cyberSessionBlockFormat selects the per-endpoint error envelope for a locally
|
||||||
|
// blocked session (用户决策:兼容路径各自格式).
|
||||||
|
type cyberSessionBlockFormat int
|
||||||
|
|
||||||
|
const (
|
||||||
|
cyberBlockFormatResponses cyberSessionBlockFormat = iota
|
||||||
|
cyberBlockFormatChat
|
||||||
|
cyberBlockFormatAnthropic
|
||||||
|
)
|
||||||
|
|
||||||
|
// rejectIfCyberSessionBlocked checks the session-block table BEFORE account
|
||||||
|
// selection. Returns true when the request was rejected (response already
|
||||||
|
// written + ops entry enqueued). Fail-open: disabled switch / empty key /
|
||||||
|
// store error → false.
|
||||||
|
func (h *OpenAIGatewayHandler) rejectIfCyberSessionBlocked(c *gin.Context, apiKey *service.APIKey, body []byte, model string, format cyberSessionBlockFormat) bool {
|
||||||
|
if h == nil || h.gatewayService == nil || apiKey == nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
// 开关默认关:先走 ~ns 级缓存开关检查,再付出 key 派生(gjson+sha256)成本。
|
||||||
|
if enabled, _ := h.gatewayService.CyberSessionBlockRuntime(c.Request.Context()); !enabled {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
key := service.CyberSessionBlockKey(apiKey.ID, c, body)
|
||||||
|
if key == "" {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if !h.gatewayService.IsCyberSessionBlocked(c.Request.Context(), key) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
switch format {
|
||||||
|
case cyberBlockFormatAnthropic:
|
||||||
|
c.JSON(http.StatusForbidden, gin.H{"type": "error", "error": gin.H{
|
||||||
|
"type": "permission_error",
|
||||||
|
"message": cyberSessionBlockedClientMsg,
|
||||||
|
}})
|
||||||
|
default: // cyberBlockFormatResponses 与 cyberBlockFormatChat:同构的 OpenAI error envelope
|
||||||
|
c.JSON(http.StatusForbidden, gin.H{"error": gin.H{
|
||||||
|
"type": "permission_error",
|
||||||
|
"code": "session_blocked_by_cyber_policy",
|
||||||
|
"message": cyberSessionBlockedClientMsg,
|
||||||
|
}})
|
||||||
|
}
|
||||||
|
h.enqueueCyberSessionBlockedOpsEntry(c, apiKey, model, key)
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// enqueueCyberSessionBlockedOpsEntry captures request meta and enqueues the
|
||||||
|
// ops_error_logs entry for a locally blocked request.
|
||||||
|
func (h *OpenAIGatewayHandler) enqueueCyberSessionBlockedOpsEntry(c *gin.Context, apiKey *service.APIKey, model string, sessionBlockKey string) {
|
||||||
|
if h.opsService == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
meta := cyberPolicyOpsErrorMeta{Model: model, InboundEndpoint: GetInboundEndpoint(c), CreatedAt: time.Now(), SessionBlockKey: sessionBlockKey}
|
||||||
|
meta.RequestID = c.Writer.Header().Get("X-Request-Id")
|
||||||
|
if c.Request != nil && c.Request.URL != nil {
|
||||||
|
meta.RequestPath = c.Request.URL.Path
|
||||||
|
}
|
||||||
|
if v, ok := c.Get(opsStreamKey); ok {
|
||||||
|
if b, ok := v.(bool); ok {
|
||||||
|
meta.Stream = b
|
||||||
|
}
|
||||||
|
}
|
||||||
|
meta.Platform = resolveOpsPlatform(apiKey, guessPlatformFromPath(meta.RequestPath))
|
||||||
|
if c.Request != nil {
|
||||||
|
meta.ClientRequestID, _ = c.Request.Context().Value(ctxkey.ClientRequestID).(string)
|
||||||
|
meta.UserAgent = c.GetHeader("User-Agent")
|
||||||
|
meta.ClientIP = strings.TrimSpace(ip.GetClientIP(c))
|
||||||
|
}
|
||||||
|
meta.APIKeyID = apiKey.ID
|
||||||
|
meta.GroupID = apiKey.GroupID
|
||||||
|
meta.APIKeyPrefix = keyPrefix(apiKey.Key, 8)
|
||||||
|
if apiKey.User != nil {
|
||||||
|
meta.UserID = apiKey.User.ID
|
||||||
|
}
|
||||||
|
enqueueOpsErrorLog(h.opsService, buildCyberSessionBlockedOpsEntry(meta))
|
||||||
|
}
|
||||||
|
|
||||||
|
// recordCyberPolicyIfMarked 在 gateway forward 返回后检查 cyber 标记,异步写风控日志/邮件,
|
||||||
|
// 并在 forward 返回错误时写一条 tokens=0 用量行。标记由 gateway 服务层在透传 cyber 后设置;
|
||||||
|
// 当前请求已发给用户,本方法只做事后记录,不影响响应。forwardErrored 为 true 时才写用量行,
|
||||||
|
// 避免与正常 RecordUsage(forward 成功路径)重复。每请求至多记录一次。
|
||||||
|
func (h *OpenAIGatewayHandler) recordCyberPolicyIfMarked(c *gin.Context, apiKey *service.APIKey, account *service.Account, subscription *service.UserSubscription, model string, forwardErrored bool, cyberBlockKey string, channelFields service.ChannelUsageFields, requestPayloadHash string) {
|
||||||
|
mark := service.GetOpsCyberPolicy(c)
|
||||||
|
if mark == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if c.GetBool(cyberPolicyRecordedKey) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
c.Set(cyberPolicyRecordedKey, true)
|
||||||
|
|
||||||
|
requestID := c.Writer.Header().Get("X-Request-Id")
|
||||||
|
var userID, apiKeyID int64
|
||||||
|
var userEmail, apiKeyName, groupName string
|
||||||
|
var groupID *int64
|
||||||
|
if apiKey != nil {
|
||||||
|
apiKeyID = apiKey.ID
|
||||||
|
apiKeyName = apiKey.Name
|
||||||
|
groupID = apiKey.GroupID
|
||||||
|
if apiKey.User != nil {
|
||||||
|
userID = apiKey.User.ID
|
||||||
|
userEmail = apiKey.User.Email
|
||||||
|
}
|
||||||
|
if apiKey.Group != nil {
|
||||||
|
groupName = apiKey.Group.Name
|
||||||
|
}
|
||||||
|
}
|
||||||
|
inboundEndpoint := GetInboundEndpoint(c)
|
||||||
|
upstreamEndpoint := ""
|
||||||
|
var accountID int64
|
||||||
|
if account != nil {
|
||||||
|
accountID = account.ID
|
||||||
|
upstreamEndpoint = GetUpstreamEndpoint(c, account.Platform)
|
||||||
|
}
|
||||||
|
stream := false
|
||||||
|
if v, ok := c.Get(opsStreamKey); ok {
|
||||||
|
if b, ok := v.(bool); ok {
|
||||||
|
stream = b
|
||||||
|
}
|
||||||
|
}
|
||||||
|
cmSvc := h.contentModerationService
|
||||||
|
gwSvc := h.gatewayService
|
||||||
|
opsSvc := h.opsService
|
||||||
|
apiKeySvc := h.apiKeyService
|
||||||
|
requestPath := ""
|
||||||
|
if c.Request != nil && c.Request.URL != nil {
|
||||||
|
requestPath = c.Request.URL.Path
|
||||||
|
}
|
||||||
|
platform := resolveOpsPlatform(apiKey, guessPlatformFromPath(requestPath))
|
||||||
|
var clientRequestID, userAgent, clientIPStr string
|
||||||
|
if c.Request != nil {
|
||||||
|
clientRequestID, _ = c.Request.Context().Value(ctxkey.ClientRequestID).(string)
|
||||||
|
userAgent = c.GetHeader("User-Agent")
|
||||||
|
clientIPStr = strings.TrimSpace(ip.GetClientIP(c))
|
||||||
|
}
|
||||||
|
apiKeyPrefix := ""
|
||||||
|
if apiKey != nil {
|
||||||
|
apiKeyPrefix = keyPrefix(apiKey.Key, 8)
|
||||||
|
}
|
||||||
|
opsMeta := cyberPolicyOpsErrorMeta{
|
||||||
|
RequestID: requestID,
|
||||||
|
ClientRequestID: clientRequestID,
|
||||||
|
Platform: platform,
|
||||||
|
Model: model,
|
||||||
|
RequestPath: requestPath,
|
||||||
|
Stream: stream,
|
||||||
|
InboundEndpoint: inboundEndpoint,
|
||||||
|
UserAgent: userAgent,
|
||||||
|
APIKeyPrefix: apiKeyPrefix,
|
||||||
|
UserID: userID,
|
||||||
|
APIKeyID: apiKeyID,
|
||||||
|
AccountID: accountID,
|
||||||
|
GroupID: groupID,
|
||||||
|
ClientIP: clientIPStr,
|
||||||
|
CreatedAt: time.Now(),
|
||||||
|
}
|
||||||
|
go func() {
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
if cmSvc != nil {
|
||||||
|
cmSvc.RecordCyberPolicyEvent(ctx, service.CyberPolicyRecordInput{
|
||||||
|
RequestID: requestID,
|
||||||
|
UserID: userID,
|
||||||
|
UserEmail: userEmail,
|
||||||
|
APIKeyID: apiKeyID,
|
||||||
|
APIKeyName: apiKeyName,
|
||||||
|
GroupID: groupID,
|
||||||
|
GroupName: groupName,
|
||||||
|
Endpoint: inboundEndpoint,
|
||||||
|
Model: model,
|
||||||
|
UpstreamMessage: mark.Message,
|
||||||
|
UpstreamBody: mark.Body,
|
||||||
|
UpstreamStatus: mark.UpstreamStatus,
|
||||||
|
UpstreamInTok: mark.UpstreamInTok,
|
||||||
|
UpstreamOutTok: mark.UpstreamOutTok,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
if forwardErrored && gwSvc != nil {
|
||||||
|
gwSvc.RecordCyberPolicyUsageLog(ctx, service.CyberPolicyUsageInput{
|
||||||
|
APIKey: apiKey,
|
||||||
|
Account: account,
|
||||||
|
Subscription: subscription,
|
||||||
|
RequestID: requestID,
|
||||||
|
Model: model,
|
||||||
|
Stream: stream,
|
||||||
|
InputTokens: mark.UpstreamInTok,
|
||||||
|
OutputTokens: mark.UpstreamOutTok,
|
||||||
|
InboundEndpoint: inboundEndpoint,
|
||||||
|
UpstreamEndpoint: upstreamEndpoint,
|
||||||
|
UserAgent: userAgent,
|
||||||
|
IPAddress: clientIPStr,
|
||||||
|
RequestPayloadHash: requestPayloadHash,
|
||||||
|
APIKeyService: apiKeySvc,
|
||||||
|
ChannelUsageFields: channelFields,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
if gwSvc != nil && cyberBlockKey != "" {
|
||||||
|
gwSvc.MarkCyberSessionBlocked(ctx, cyberBlockKey)
|
||||||
|
}
|
||||||
|
if opsSvc != nil {
|
||||||
|
enqueueOpsErrorLog(opsSvc, buildCyberPolicyOpsErrorEntry(opsMeta, mark))
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
|
||||||
|
// clearCyberPolicyTurnState resets the cyber mark and the per-request recorded
|
||||||
|
// guard. WS-only: called at the END of AfterTurn, after recordCyberPolicyIfMarked
|
||||||
|
// and RecordUsage (which reads CyberBlocked) have both consumed the mark.
|
||||||
|
func clearCyberPolicyTurnState(c *gin.Context) {
|
||||||
|
if c == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
service.ClearOpsCyberPolicy(c)
|
||||||
|
c.Set(cyberPolicyRecordedKey, false)
|
||||||
|
}
|
||||||
|
|
||||||
func summarizeWSCloseErrorForLog(err error) (string, string) {
|
func summarizeWSCloseErrorForLog(err error) (string, string) {
|
||||||
if err == nil {
|
if err == nil {
|
||||||
return "-", "-"
|
return "-", "-"
|
||||||
|
|||||||
@@ -805,7 +805,7 @@ func (r *contentModerationHandlerTestRepo) ListLogs(ctx context.Context, filter
|
|||||||
return nil, nil, nil
|
return nil, nil, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *contentModerationHandlerTestRepo) CountFlaggedByUserSince(ctx context.Context, userID int64, since time.Time) (int, error) {
|
func (r *contentModerationHandlerTestRepo) CountFlaggedByUserSince(ctx context.Context, userID int64, since time.Time, excludeCyberPolicy bool) (int, error) {
|
||||||
return 0, nil
|
return 0, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -813,6 +813,10 @@ func (r *contentModerationHandlerTestRepo) CleanupExpiredLogs(ctx context.Contex
|
|||||||
return &service.ContentModerationCleanupResult{}, nil
|
return &service.ContentModerationCleanupResult{}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (r *contentModerationHandlerTestRepo) UpdateLogEmailSent(ctx context.Context, id int64, sent bool) error {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
func TestOpenAIResponsesWebSocket_ContentModerationBlocksFirstFrame(t *testing.T) {
|
func TestOpenAIResponsesWebSocket_ContentModerationBlocksFirstFrame(t *testing.T) {
|
||||||
gin.SetMode(gin.TestMode)
|
gin.SetMode(gin.TestMode)
|
||||||
|
|
||||||
@@ -1654,4 +1658,28 @@ data: {"type":"response.failed","error":{"message":"This content was flagged"}}
|
|||||||
|
|
||||||
require.False(t, reported)
|
require.False(t, reported)
|
||||||
})
|
})
|
||||||
|
|
||||||
|
// H-2: cyber_policy 命中且响应已写出时,即便 err 前缀不在白名单(非流式 400 cyber
|
||||||
|
// 返回 "openai cyber_policy:"、透传账号返回 "upstream error:"),也须判定已透传,避免
|
||||||
|
// ensureForwardErrorResponse 在已写出的完整响应尾部追加 SSE 污染响应体。
|
||||||
|
t.Run("cyber policy hit after write is already communicated", func(t *testing.T) {
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(w)
|
||||||
|
c.Request = httptest.NewRequest(http.MethodPost, EndpointResponses, nil)
|
||||||
|
service.MarkOpsCyberPolicy(c, service.CyberPolicyMark{Message: "blocked", UpstreamStatus: 400})
|
||||||
|
before := c.Writer.Size()
|
||||||
|
_, _ = c.Writer.WriteString(`{"error":{"code":"cyber_policy","message":"blocked"}}`)
|
||||||
|
|
||||||
|
require.True(t, openAIForwardErrorAlreadyCommunicated(c, before, errors.New("openai cyber_policy: blocked")))
|
||||||
|
})
|
||||||
|
|
||||||
|
// Size 守卫优先于 cyber 短路:cyber 命中但未写出任何响应时仍需补写错误。
|
||||||
|
t.Run("cyber policy without write still needs fallback", func(t *testing.T) {
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(w)
|
||||||
|
c.Request = httptest.NewRequest(http.MethodPost, EndpointResponses, nil)
|
||||||
|
service.MarkOpsCyberPolicy(c, service.CyberPolicyMark{Message: "blocked", UpstreamStatus: 400})
|
||||||
|
|
||||||
|
require.False(t, openAIForwardErrorAlreadyCommunicated(c, c.Writer.Size(), errors.New("openai cyber_policy: blocked")))
|
||||||
|
})
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -549,6 +549,10 @@ func OpsErrorLoggerMiddleware(ops *service.OpsService) gin.HandlerFunc {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if shouldSkipOpsErrorLogForCyber(c) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
status := c.Writer.Status()
|
status := c.Writer.Status()
|
||||||
if status < 400 {
|
if status < 400 {
|
||||||
// Even when the client request succeeds, we still want to persist upstream error attempts
|
// Even when the client request succeeds, we still want to persist upstream error attempts
|
||||||
@@ -1467,3 +1471,9 @@ func shouldSkipOpsErrorLog(ctx context.Context, ops *service.OpsService, message
|
|||||||
|
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// shouldSkipOpsErrorLogForCyber:cyber_policy 命中的请求由 recordCyberPolicyIfMarked
|
||||||
|
// 统一落一条 status=403 的错误请求,故中间件跳过自身落库,避免双写。
|
||||||
|
func shouldSkipOpsErrorLogForCyber(c *gin.Context) bool {
|
||||||
|
return service.GetOpsCyberPolicy(c) != nil
|
||||||
|
}
|
||||||
|
|||||||
@@ -0,0 +1,22 @@
|
|||||||
|
package handler
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
// cyber mark 存在时,中间件必须跳过自身落库(由 recordCyberPolicyIfMarked 统一落 403)。
|
||||||
|
func TestOpsErrorLoggerMiddlewareSkipsCyber(t *testing.T) {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
c, _ := gin.CreateTestContext(httptest.NewRecorder())
|
||||||
|
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil)
|
||||||
|
service.MarkOpsCyberPolicy(c, service.CyberPolicyMark{Code: "cyber_policy", Message: "blocked", UpstreamStatus: http.StatusOK})
|
||||||
|
|
||||||
|
require.NotNil(t, service.GetOpsCyberPolicy(c), "前置:mark 已设置")
|
||||||
|
require.True(t, shouldSkipOpsErrorLogForCyber(c), "cyber mark 命中应跳过中间件落库")
|
||||||
|
}
|
||||||
@@ -177,10 +177,11 @@ LIMIT $`+fmt.Sprint(len(queryArgs)-1)+` OFFSET $`+fmt.Sprint(len(queryArgs)),
|
|||||||
return items, paginationResultFromTotal(total, params), nil
|
return items, paginationResultFromTotal(total, params), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *contentModerationRepository) CountFlaggedByUserSince(ctx context.Context, userID int64, since time.Time) (int, error) {
|
func (r *contentModerationRepository) CountFlaggedByUserSince(ctx context.Context, userID int64, since time.Time, excludeCyberPolicy bool) (int, error) {
|
||||||
if userID <= 0 {
|
if userID <= 0 {
|
||||||
return 0, nil
|
return 0, nil
|
||||||
}
|
}
|
||||||
|
// SQL 中的 'cyber_policy' 字面量须与 service.ContentModerationActionCyberPolicy 保持一致。
|
||||||
var count int
|
var count int
|
||||||
err := r.db.QueryRowContext(ctx, `
|
err := r.db.QueryRowContext(ctx, `
|
||||||
WITH last_auto_ban AS (
|
WITH last_auto_ban AS (
|
||||||
@@ -193,15 +194,24 @@ FROM content_moderation_logs
|
|||||||
WHERE user_id = $1
|
WHERE user_id = $1
|
||||||
AND flagged = TRUE
|
AND flagged = TRUE
|
||||||
AND action <> 'hash_block'
|
AND action <> 'hash_block'
|
||||||
|
AND ($3::bool IS FALSE OR action <> 'cyber_policy')
|
||||||
AND created_at >= $2
|
AND created_at >= $2
|
||||||
AND created_at > COALESCE((SELECT at FROM last_auto_ban), '-infinity'::timestamptz)
|
AND created_at > COALESCE((SELECT at FROM last_auto_ban), '-infinity'::timestamptz)
|
||||||
`, userID, since).Scan(&count)
|
`, userID, since, excludeCyberPolicy).Scan(&count)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return 0, fmt.Errorf("count user content moderation flagged logs: %w", err)
|
return 0, fmt.Errorf("count user content moderation flagged logs: %w", err)
|
||||||
}
|
}
|
||||||
return count, nil
|
return count, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (r *contentModerationRepository) UpdateLogEmailSent(ctx context.Context, id int64, sent bool) error {
|
||||||
|
_, err := r.db.ExecContext(ctx, `UPDATE content_moderation_logs SET email_sent = $1 WHERE id = $2`, sent, id)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("update content moderation log email_sent: %w", err)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
func (r *contentModerationRepository) CleanupExpiredLogs(ctx context.Context, hitBefore time.Time, nonHitBefore time.Time) (*service.ContentModerationCleanupResult, error) {
|
func (r *contentModerationRepository) CleanupExpiredLogs(ctx context.Context, hitBefore time.Time, nonHitBefore time.Time) (*service.ContentModerationCleanupResult, error) {
|
||||||
result := &service.ContentModerationCleanupResult{FinishedAt: time.Now()}
|
result := &service.ContentModerationCleanupResult{FinishedAt: time.Now()}
|
||||||
if r == nil || r.db == nil {
|
if r == nil || r.db == nil {
|
||||||
|
|||||||
@@ -29,12 +29,30 @@ func TestContentModerationRepositoryCountFlaggedByUserSince_ExcludesHashBlock(t
|
|||||||
repo := NewContentModerationRepository(db)
|
repo := NewContentModerationRepository(db)
|
||||||
since := time.Now().Add(-time.Hour)
|
since := time.Now().Add(-time.Hour)
|
||||||
mock.ExpectQuery(regexp.QuoteMeta("AND action <> 'hash_block'")).
|
mock.ExpectQuery(regexp.QuoteMeta("AND action <> 'hash_block'")).
|
||||||
WithArgs(int64(1001), since).
|
WithArgs(int64(1001), since, false).
|
||||||
WillReturnRows(sqlmock.NewRows([]string{"count"}).AddRow(2))
|
WillReturnRows(sqlmock.NewRows([]string{"count"}).AddRow(2))
|
||||||
|
|
||||||
count, err := repo.CountFlaggedByUserSince(context.Background(), 1001, since)
|
count, err := repo.CountFlaggedByUserSince(context.Background(), 1001, since, false)
|
||||||
|
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
require.Equal(t, 2, count)
|
require.Equal(t, 2, count)
|
||||||
require.NoError(t, mock.ExpectationsWereMet())
|
require.NoError(t, mock.ExpectationsWereMet())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestContentModerationRepositoryCountFlaggedByUserSince_ExcludesCyberPolicyWhenRequested(t *testing.T) {
|
||||||
|
db, mock, err := sqlmock.New()
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer func() { _ = db.Close() }()
|
||||||
|
|
||||||
|
repo := NewContentModerationRepository(db)
|
||||||
|
since := time.Now().Add(-time.Hour)
|
||||||
|
mock.ExpectQuery(regexp.QuoteMeta("AND ($3::bool IS FALSE OR action <> 'cyber_policy')")).
|
||||||
|
WithArgs(int64(1001), since, true).
|
||||||
|
WillReturnRows(sqlmock.NewRows([]string{"count"}).AddRow(3))
|
||||||
|
|
||||||
|
count, err := repo.CountFlaggedByUserSince(context.Background(), 1001, since, true)
|
||||||
|
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Equal(t, 3, count)
|
||||||
|
require.NoError(t, mock.ExpectationsWereMet())
|
||||||
|
}
|
||||||
|
|||||||
@@ -51,3 +51,23 @@ func (c *gatewayCache) DeleteSessionAccountID(ctx context.Context, groupID int64
|
|||||||
key := buildSessionKey(groupID, sessionHash)
|
key := buildSessionKey(groupID, sessionHash)
|
||||||
return c.rdb.Del(ctx, key).Err()
|
return c.rdb.Del(ctx, key).Err()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Compile-time assertion: gatewayCache must implement CyberSessionBlockStore.
|
||||||
|
var _ service.CyberSessionBlockStore = (*gatewayCache)(nil)
|
||||||
|
|
||||||
|
const cyberSessionBlockPrefix = "cyber_session_block:"
|
||||||
|
|
||||||
|
// SetCyberSessionBlocked 把被 cyber_policy 命中的会话写入屏蔽表(TTL 自动过期)。
|
||||||
|
// 存储值 "1" 作为存在标记(IsCyberSessionBlocked 只检查 key 是否存在,不读值)。
|
||||||
|
func (c *gatewayCache) SetCyberSessionBlocked(ctx context.Context, key string, ttl time.Duration) error {
|
||||||
|
return c.rdb.Set(ctx, cyberSessionBlockPrefix+key, "1", ttl).Err()
|
||||||
|
}
|
||||||
|
|
||||||
|
// IsCyberSessionBlocked 查询会话是否在屏蔽表中。
|
||||||
|
func (c *gatewayCache) IsCyberSessionBlocked(ctx context.Context, key string) (bool, error) {
|
||||||
|
n, err := c.rdb.Exists(ctx, cyberSessionBlockPrefix+key).Result()
|
||||||
|
if err != nil {
|
||||||
|
return false, err
|
||||||
|
}
|
||||||
|
return n > 0, nil
|
||||||
|
}
|
||||||
|
|||||||
@@ -70,6 +70,28 @@ func TestBuildOpsErrorLogsWhere_ModelFuzzy(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// TestBuildOpsErrorLogsWhere_CyberPolicyStatusExemption verifies that streaming
|
||||||
|
// cyber_policy hits (status_code=200) remain visible in admin + user error-request
|
||||||
|
// lists. The repository filter must emit an OR exemption for error_type='cyber_policy'
|
||||||
|
// so that stream-path cyber rows (upstream delivers 200 with a failed SSE event) are
|
||||||
|
// not silently excluded by the COALESCE(status_code,0) >= 400 guard.
|
||||||
|
func TestBuildOpsErrorLogsWhere_CyberPolicyStatusExemption(t *testing.T) {
|
||||||
|
// Default filter (no phase) must include the cyber_policy exemption.
|
||||||
|
where, _ := buildOpsErrorLogsWhere(&service.OpsErrorLogFilter{})
|
||||||
|
if !strings.Contains(where, "e.error_type = 'cyber_policy'") {
|
||||||
|
t.Fatalf("default filter must exempt cyber_policy from status >= 400 guard\nfull: %s", where)
|
||||||
|
}
|
||||||
|
if !strings.Contains(where, "COALESCE(e.status_code, 0) >= 400") {
|
||||||
|
t.Fatalf("default filter must still include the status >= 400 guard for non-cyber rows\nfull: %s", where)
|
||||||
|
}
|
||||||
|
|
||||||
|
// phase=upstream skips the status guard entirely — exemption is irrelevant there.
|
||||||
|
whereUpstream, _ := buildOpsErrorLogsWhere(&service.OpsErrorLogFilter{Phase: "upstream"})
|
||||||
|
if strings.Contains(whereUpstream, "status_code") {
|
||||||
|
t.Fatalf("upstream phase filter must not add any status_code clause\nfull: %s", whereUpstream)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestBuildOpsErrorLogsWhere_MatchDeletedKeyOwner(t *testing.T) {
|
func TestBuildOpsErrorLogsWhere_MatchDeletedKeyOwner(t *testing.T) {
|
||||||
uid := int64(42)
|
uid := int64(42)
|
||||||
|
|
||||||
|
|||||||
@@ -919,8 +919,12 @@ func buildOpsErrorLogsWhere(filter *service.OpsErrorLogFilter) (string, []any) {
|
|||||||
resolvedFilter = filter.Resolved
|
resolvedFilter = filter.Resolved
|
||||||
}
|
}
|
||||||
// Keep list endpoints scoped to client errors unless explicitly filtering upstream phase.
|
// Keep list endpoints scoped to client errors unless explicitly filtering upstream phase.
|
||||||
|
// cyber_policy is exempt from the status >= 400 guard: streaming cyber hits arrive with
|
||||||
|
// status 200 (the SSE stream opened successfully before upstream returned response.failed),
|
||||||
|
// but they are always client-visible blocked requests that belong in admin + user error
|
||||||
|
// lists. Without the exemption the entire streaming-path cyber sink would be invisible.
|
||||||
if phaseFilter != "upstream" {
|
if phaseFilter != "upstream" {
|
||||||
clauses = append(clauses, "COALESCE(e.status_code, 0) >= 400")
|
clauses = append(clauses, "(COALESCE(e.status_code, 0) >= 400 OR e.error_type = 'cyber_policy')")
|
||||||
}
|
}
|
||||||
|
|
||||||
if filter.StartTime != nil && !filter.StartTime.IsZero() {
|
if filter.StartTime != nil && !filter.StartTime.IsZero() {
|
||||||
|
|||||||
@@ -884,6 +884,8 @@ func TestAPIContracts(t *testing.T) {
|
|||||||
"channel_monitor_default_interval_seconds": 60,
|
"channel_monitor_default_interval_seconds": 60,
|
||||||
"available_channels_enabled": false,
|
"available_channels_enabled": false,
|
||||||
"risk_control_enabled": false,
|
"risk_control_enabled": false,
|
||||||
|
"cyber_session_block_enabled": false,
|
||||||
|
"cyber_session_block_ttl_seconds": 3600,
|
||||||
"affiliate_enabled": false,
|
"affiliate_enabled": false,
|
||||||
"wechat_connect_enabled": false,
|
"wechat_connect_enabled": false,
|
||||||
"wechat_connect_app_id": "",
|
"wechat_connect_app_id": "",
|
||||||
@@ -1120,6 +1122,8 @@ func TestAPIContracts(t *testing.T) {
|
|||||||
"channel_monitor_default_interval_seconds": 60,
|
"channel_monitor_default_interval_seconds": 60,
|
||||||
"available_channels_enabled": false,
|
"available_channels_enabled": false,
|
||||||
"risk_control_enabled": false,
|
"risk_control_enabled": false,
|
||||||
|
"cyber_session_block_enabled": false,
|
||||||
|
"cyber_session_block_ttl_seconds": 3600,
|
||||||
"affiliate_enabled": false,
|
"affiliate_enabled": false,
|
||||||
"wechat_connect_enabled": true,
|
"wechat_connect_enabled": true,
|
||||||
"wechat_connect_app_id": "wx-open-config",
|
"wechat_connect_app_id": "wx-open-config",
|
||||||
|
|||||||
@@ -37,6 +37,7 @@ const (
|
|||||||
ContentModerationActionHashBlock = "hash_block"
|
ContentModerationActionHashBlock = "hash_block"
|
||||||
ContentModerationActionKeywordBlock = "keyword_block"
|
ContentModerationActionKeywordBlock = "keyword_block"
|
||||||
ContentModerationActionError = "error"
|
ContentModerationActionError = "error"
|
||||||
|
ContentModerationActionCyberPolicy = "cyber_policy" // cyber_policy 硬阻断的风控日志 action(封号计数排除按此值过滤)
|
||||||
|
|
||||||
contentModerationKeywordCategory = "keyword"
|
contentModerationKeywordCategory = "keyword"
|
||||||
|
|
||||||
@@ -160,6 +161,10 @@ type ContentModerationConfig struct {
|
|||||||
BlockedKeywords []string `json:"blocked_keywords"`
|
BlockedKeywords []string `json:"blocked_keywords"`
|
||||||
KeywordBlockingMode string `json:"keyword_blocking_mode"`
|
KeywordBlockingMode string `json:"keyword_blocking_mode"`
|
||||||
ModelFilter ContentModerationModelFilter `json:"model_filter"`
|
ModelFilter ContentModerationModelFilter `json:"model_filter"`
|
||||||
|
// CyberPolicyExcludeFromBanCount 为 true 时,cyber_policy 命中不参与自动封号计数:
|
||||||
|
// 当次不判定封号,且历史 cyber 行在 CountFlaggedByUserSince 中被排除。
|
||||||
|
// 默认 false(计入,与历史行为一致;旧配置 JSON 无此字段时反序列化为 false)。
|
||||||
|
CyberPolicyExcludeFromBanCount bool `json:"cyber_policy_exclude_from_ban_count"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type ContentModerationConfigView struct {
|
type ContentModerationConfigView struct {
|
||||||
@@ -193,6 +198,7 @@ type ContentModerationConfigView struct {
|
|||||||
BlockedKeywords []string `json:"blocked_keywords"`
|
BlockedKeywords []string `json:"blocked_keywords"`
|
||||||
KeywordBlockingMode string `json:"keyword_blocking_mode"`
|
KeywordBlockingMode string `json:"keyword_blocking_mode"`
|
||||||
ModelFilter ContentModerationModelFilter `json:"model_filter"`
|
ModelFilter ContentModerationModelFilter `json:"model_filter"`
|
||||||
|
CyberPolicyExcludeFromBanCount bool `json:"cyber_policy_exclude_from_ban_count"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type ContentModerationAPIKeyStatus struct {
|
type ContentModerationAPIKeyStatus struct {
|
||||||
@@ -280,6 +286,7 @@ type UpdateContentModerationConfigInput struct {
|
|||||||
BlockedKeywords *[]string `json:"blocked_keywords"`
|
BlockedKeywords *[]string `json:"blocked_keywords"`
|
||||||
KeywordBlockingMode *string `json:"keyword_blocking_mode"`
|
KeywordBlockingMode *string `json:"keyword_blocking_mode"`
|
||||||
ModelFilter *ContentModerationModelFilter `json:"model_filter"`
|
ModelFilter *ContentModerationModelFilter `json:"model_filter"`
|
||||||
|
CyberPolicyExcludeFromBanCount *bool `json:"cyber_policy_exclude_from_ban_count"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type ContentModerationModelFilter struct {
|
type ContentModerationModelFilter struct {
|
||||||
@@ -461,8 +468,12 @@ type ContentModerationClearHashesResult struct {
|
|||||||
type ContentModerationRepository interface {
|
type ContentModerationRepository interface {
|
||||||
CreateLog(ctx context.Context, log *ContentModerationLog) error
|
CreateLog(ctx context.Context, log *ContentModerationLog) error
|
||||||
ListLogs(ctx context.Context, filter ContentModerationLogFilter) ([]ContentModerationLog, *pagination.PaginationResult, error)
|
ListLogs(ctx context.Context, filter ContentModerationLogFilter) ([]ContentModerationLog, *pagination.PaginationResult, error)
|
||||||
CountFlaggedByUserSince(ctx context.Context, userID int64, since time.Time) (int, error)
|
// CountFlaggedByUserSince 统计窗口内计入封号的违规次数(排除 hash_block;
|
||||||
|
// excludeCyberPolicy 为 true 时额外排除 cyber_policy 行)。
|
||||||
|
CountFlaggedByUserSince(ctx context.Context, userID int64, since time.Time, excludeCyberPolicy bool) (int, error)
|
||||||
CleanupExpiredLogs(ctx context.Context, hitBefore time.Time, nonHitBefore time.Time) (*ContentModerationCleanupResult, error)
|
CleanupExpiredLogs(ctx context.Context, hitBefore time.Time, nonHitBefore time.Time) (*ContentModerationCleanupResult, error)
|
||||||
|
// UpdateLogEmailSent 回写邮件发送结果(F7:CreateLog 先行后补 EmailSent)。
|
||||||
|
UpdateLogEmailSent(ctx context.Context, id int64, sent bool) error
|
||||||
}
|
}
|
||||||
|
|
||||||
type ContentModerationHashCache interface {
|
type ContentModerationHashCache interface {
|
||||||
@@ -648,6 +659,9 @@ func (s *ContentModerationService) UpdateConfig(ctx context.Context, input Updat
|
|||||||
if input.RecordNonHits != nil {
|
if input.RecordNonHits != nil {
|
||||||
cfg.RecordNonHits = *input.RecordNonHits
|
cfg.RecordNonHits = *input.RecordNonHits
|
||||||
}
|
}
|
||||||
|
if input.CyberPolicyExcludeFromBanCount != nil {
|
||||||
|
cfg.CyberPolicyExcludeFromBanCount = *input.CyberPolicyExcludeFromBanCount
|
||||||
|
}
|
||||||
if input.Thresholds != nil {
|
if input.Thresholds != nil {
|
||||||
cfg.Thresholds = mergeContentModerationThresholds(ContentModerationDefaultThresholds(), *input.Thresholds)
|
cfg.Thresholds = mergeContentModerationThresholds(ContentModerationDefaultThresholds(), *input.Thresholds)
|
||||||
}
|
}
|
||||||
@@ -1644,7 +1658,7 @@ func (s *ContentModerationService) applyFlaggedAccountSideEffects(ctx context.Co
|
|||||||
count := 1
|
count := 1
|
||||||
if s.repo != nil && cfg.ViolationWindowHours > 0 {
|
if s.repo != nil && cfg.ViolationWindowHours > 0 {
|
||||||
since := time.Now().Add(-time.Duration(cfg.ViolationWindowHours) * time.Hour)
|
since := time.Now().Add(-time.Duration(cfg.ViolationWindowHours) * time.Hour)
|
||||||
if n, err := s.repo.CountFlaggedByUserSince(ctx, *log.UserID, since); err == nil {
|
if n, err := s.repo.CountFlaggedByUserSince(ctx, *log.UserID, since, cfg.CyberPolicyExcludeFromBanCount); err == nil {
|
||||||
count = n + 1
|
count = n + 1
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -1835,6 +1849,7 @@ func defaultContentModerationConfig() *ContentModerationConfig {
|
|||||||
Type: ContentModerationModelFilterAll,
|
Type: ContentModerationModelFilterAll,
|
||||||
Models: []string{},
|
Models: []string{},
|
||||||
},
|
},
|
||||||
|
CyberPolicyExcludeFromBanCount: false,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -2161,6 +2176,7 @@ func (s *ContentModerationService) configView(cfg *ContentModerationConfig) *Con
|
|||||||
BlockedKeywords: append([]string(nil), cfg.BlockedKeywords...),
|
BlockedKeywords: append([]string(nil), cfg.BlockedKeywords...),
|
||||||
KeywordBlockingMode: cfg.KeywordBlockingMode,
|
KeywordBlockingMode: cfg.KeywordBlockingMode,
|
||||||
ModelFilter: cloneContentModerationModelFilter(cfg.ModelFilter),
|
ModelFilter: cloneContentModerationModelFilter(cfg.ModelFilter),
|
||||||
|
CyberPolicyExcludeFromBanCount: cfg.CyberPolicyExcludeFromBanCount,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -2688,3 +2704,135 @@ func maskSecretTail(secret string) string {
|
|||||||
}
|
}
|
||||||
return strings.Repeat("*", 8) + secret[len(secret)-4:]
|
return strings.Repeat("*", 8) + secret[len(secret)-4:]
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// CyberPolicyRecordInput 是一次 cyber_policy 硬阻断的风控记录入参。
|
||||||
|
type CyberPolicyRecordInput struct {
|
||||||
|
RequestID string
|
||||||
|
UserID int64
|
||||||
|
UserEmail string
|
||||||
|
APIKeyID int64
|
||||||
|
APIKeyName string
|
||||||
|
GroupID *int64
|
||||||
|
GroupName string
|
||||||
|
Endpoint string
|
||||||
|
Model string
|
||||||
|
UpstreamMessage string
|
||||||
|
UpstreamBody string
|
||||||
|
UpstreamStatus int
|
||||||
|
UpstreamInTok int
|
||||||
|
UpstreamOutTok int
|
||||||
|
}
|
||||||
|
|
||||||
|
// RecordCyberPolicyEvent 把一次 cyber_policy 硬阻断写入风控中心日志、计入违规计数、
|
||||||
|
// 并给用户发邮件。当前请求已由 gateway 透传给用户;本方法仅做事后记录/通知/计数。
|
||||||
|
// 仅受 risk_control_enabled 总开关约束(不受内容审核 Enabled/Mode/scope/sample 约束)。
|
||||||
|
func (s *ContentModerationService) RecordCyberPolicyEvent(ctx context.Context, in CyberPolicyRecordInput) {
|
||||||
|
if s == nil || s.repo == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if !s.isRiskControlEnabled(ctx) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
cfg, err := s.loadConfig(ctx)
|
||||||
|
if err != nil {
|
||||||
|
slog.Warn("content_moderation.cyber_load_config_failed", "error", err)
|
||||||
|
cfg = &ContentModerationConfig{}
|
||||||
|
}
|
||||||
|
var userID *int64
|
||||||
|
if in.UserID > 0 {
|
||||||
|
userID = &in.UserID
|
||||||
|
}
|
||||||
|
var apiKeyID *int64
|
||||||
|
if in.APIKeyID > 0 {
|
||||||
|
apiKeyID = &in.APIKeyID
|
||||||
|
}
|
||||||
|
errBody := strings.TrimSpace(in.UpstreamMessage)
|
||||||
|
if b := strings.TrimSpace(in.UpstreamBody); b != "" {
|
||||||
|
// 原始 body 不在此预脱敏;写入 log.Error 前由 redactContentModerationSecrets 统一脱敏。
|
||||||
|
errBody = strings.TrimSpace(errBody + "\n" + b)
|
||||||
|
}
|
||||||
|
if in.UpstreamInTok > 0 || in.UpstreamOutTok > 0 {
|
||||||
|
errBody = fmt.Sprintf("%s\nupstream_usage=in:%d,out:%d", errBody, in.UpstreamInTok, in.UpstreamOutTok)
|
||||||
|
}
|
||||||
|
log := &ContentModerationLog{
|
||||||
|
RequestID: in.RequestID,
|
||||||
|
UserID: userID,
|
||||||
|
UserEmail: in.UserEmail,
|
||||||
|
APIKeyID: apiKeyID,
|
||||||
|
APIKeyName: in.APIKeyName,
|
||||||
|
GroupID: cloneInt64Ptr(in.GroupID),
|
||||||
|
GroupName: in.GroupName,
|
||||||
|
Endpoint: in.Endpoint,
|
||||||
|
Provider: "openai",
|
||||||
|
Model: in.Model,
|
||||||
|
Mode: "post_upstream",
|
||||||
|
Action: ContentModerationActionCyberPolicy,
|
||||||
|
Flagged: true,
|
||||||
|
HighestCategory: "cyber_policy",
|
||||||
|
HighestScore: 1.0,
|
||||||
|
Error: trimRunes(redactContentModerationSecrets(errBody), maxModerationExcerptRunes*4),
|
||||||
|
CreatedAt: time.Now(),
|
||||||
|
}
|
||||||
|
// 开关开时 cyber_policy 不参与封号计数:当次不判定(此处跳过),
|
||||||
|
// 历史行由 CountFlaggedByUserSince 的 excludeCyberPolicy 排除。
|
||||||
|
autoBanned := false
|
||||||
|
if !cfg.CyberPolicyExcludeFromBanCount {
|
||||||
|
autoBanned = s.applyFlaggedAccountSideEffects(ctx, cfg, log)
|
||||||
|
}
|
||||||
|
log.EmailSent = false
|
||||||
|
logPersisted := true
|
||||||
|
if err := s.repo.CreateLog(ctx, log); err != nil {
|
||||||
|
logPersisted = false
|
||||||
|
slog.Warn("content_moderation.cyber_create_log_failed", "user_id", in.UserID, "error", err)
|
||||||
|
}
|
||||||
|
emailSent := false
|
||||||
|
if s.emailService != nil && strings.TrimSpace(log.UserEmail) != "" {
|
||||||
|
if err := s.sendCyberPolicyEmail(ctx, log); err != nil {
|
||||||
|
slog.Warn("content_moderation.cyber_email_failed", "user_id", in.UserID, "error", err)
|
||||||
|
} else {
|
||||||
|
emailSent = true
|
||||||
|
}
|
||||||
|
if autoBanned {
|
||||||
|
if err := s.sendAccountDisabledEmail(ctx, cfg, log); err != nil {
|
||||||
|
slog.Warn("content_moderation.cyber_ban_email_failed", "user_id", in.UserID, "error", err)
|
||||||
|
} else {
|
||||||
|
emailSent = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if logPersisted && emailSent {
|
||||||
|
if err := s.repo.UpdateLogEmailSent(ctx, log.ID, true); err != nil {
|
||||||
|
slog.Warn("content_moderation.cyber_update_email_sent_failed", "log_id", log.ID, "error", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *ContentModerationService) sendCyberPolicyEmail(ctx context.Context, log *ContentModerationLog) error {
|
||||||
|
siteName := s.siteName(ctx)
|
||||||
|
if s.emailService.notificationEmailService != nil {
|
||||||
|
variables := map[string]string{
|
||||||
|
"triggered_at": log.CreatedAt.UTC().Format(time.RFC3339),
|
||||||
|
"model": defaultContentModerationString(log.Model, "-"),
|
||||||
|
"group_name": defaultContentModerationString(log.GroupName, "-"),
|
||||||
|
"upstream_message": defaultContentModerationString(log.Error, "-"),
|
||||||
|
}
|
||||||
|
err := s.emailService.notificationEmailService.Send(ctx, NotificationEmailSendInput{
|
||||||
|
Event: NotificationEmailEventCyberPolicyNotice,
|
||||||
|
RecipientEmail: log.UserEmail,
|
||||||
|
RecipientName: emailRecipientName(log.UserEmail),
|
||||||
|
UserID: contentModerationEmailUserID(log),
|
||||||
|
SourceType: "content_moderation",
|
||||||
|
SourceID: contentModerationEmailSourceID(log),
|
||||||
|
Variables: variables,
|
||||||
|
})
|
||||||
|
if err == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if !shouldFallbackNotificationEmail(err) {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
slog.Warn("template cyber policy email failed; falling back", "err", err.Error())
|
||||||
|
}
|
||||||
|
subject := fmt.Sprintf("[%s] 网络安全策略拦截 / Cyber Policy Notice", sanitizeEmailHeader(siteName))
|
||||||
|
return s.emailService.SendEmail(ctx, log.UserEmail, subject, buildCyberPolicyNoticeEmailBody(siteName, log))
|
||||||
|
}
|
||||||
|
|||||||
@@ -0,0 +1,302 @@
|
|||||||
|
package service
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/Wei-Shaw/sub2api/internal/pkg/pagination"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
// cyberOrderingTestRepo records the sequence of repo calls to verify F7 ordering.
|
||||||
|
type cyberOrderingTestRepo struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
calls []string
|
||||||
|
emailSents []bool // EmailSent value captured at each CreateLog call
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *cyberOrderingTestRepo) CreateLog(ctx context.Context, log *ContentModerationLog) error {
|
||||||
|
r.mu.Lock()
|
||||||
|
defer r.mu.Unlock()
|
||||||
|
r.calls = append(r.calls, "create")
|
||||||
|
if log != nil {
|
||||||
|
r.emailSents = append(r.emailSents, log.EmailSent)
|
||||||
|
log.ID = 1 // simulate DB-assigned ID so UpdateLogEmailSent guard passes
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *cyberOrderingTestRepo) UpdateLogEmailSent(ctx context.Context, id int64, sent bool) error {
|
||||||
|
r.mu.Lock()
|
||||||
|
defer r.mu.Unlock()
|
||||||
|
r.calls = append(r.calls, "update_email_sent")
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *cyberOrderingTestRepo) ListLogs(ctx context.Context, filter ContentModerationLogFilter) ([]ContentModerationLog, *pagination.PaginationResult, error) {
|
||||||
|
return nil, nil, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *cyberOrderingTestRepo) CountFlaggedByUserSince(ctx context.Context, userID int64, since time.Time, excludeCyberPolicy bool) (int, error) {
|
||||||
|
return 0, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *cyberOrderingTestRepo) CleanupExpiredLogs(ctx context.Context, hitBefore time.Time, nonHitBefore time.Time) (*ContentModerationCleanupResult, error) {
|
||||||
|
return &ContentModerationCleanupResult{}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *cyberOrderingTestRepo) snapshot() []string {
|
||||||
|
r.mu.Lock()
|
||||||
|
defer r.mu.Unlock()
|
||||||
|
out := make([]string, len(r.calls))
|
||||||
|
copy(out, r.calls)
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *cyberOrderingTestRepo) snapshotEmailSents() []bool {
|
||||||
|
r.mu.Lock()
|
||||||
|
defer r.mu.Unlock()
|
||||||
|
out := make([]bool, len(r.emailSents))
|
||||||
|
copy(out, r.emailSents)
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRecordCyberPolicyEvent_DisabledWhenRiskControlOff(t *testing.T) {
|
||||||
|
repo := &contentModerationTestRepo{}
|
||||||
|
svc := NewContentModerationService(
|
||||||
|
&contentModerationTestSettingRepo{values: map[string]string{
|
||||||
|
SettingKeyRiskControlEnabled: "false",
|
||||||
|
}},
|
||||||
|
repo,
|
||||||
|
nil,
|
||||||
|
nil,
|
||||||
|
nil,
|
||||||
|
nil,
|
||||||
|
nil,
|
||||||
|
)
|
||||||
|
|
||||||
|
svc.RecordCyberPolicyEvent(context.Background(), CyberPolicyRecordInput{
|
||||||
|
UserID: 1,
|
||||||
|
UserEmail: "u@x.com",
|
||||||
|
Model: "gpt-5",
|
||||||
|
Endpoint: "/v1/responses",
|
||||||
|
UpstreamMessage: "flagged",
|
||||||
|
UpstreamBody: `{"error":{"code":"cyber_policy"}}`,
|
||||||
|
UpstreamStatus: 400,
|
||||||
|
})
|
||||||
|
|
||||||
|
require.Empty(t, repo.snapshotLogs(), "CreateLog must NOT be called when risk_control_enabled is off")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRecordCyberPolicyEvent_WritesLogWhenEnabled(t *testing.T) {
|
||||||
|
repo := &contentModerationTestRepo{}
|
||||||
|
svc := NewContentModerationService(
|
||||||
|
&contentModerationTestSettingRepo{values: map[string]string{
|
||||||
|
SettingKeyRiskControlEnabled: "true",
|
||||||
|
}},
|
||||||
|
repo,
|
||||||
|
nil,
|
||||||
|
nil,
|
||||||
|
nil,
|
||||||
|
nil,
|
||||||
|
nil, // emailService=nil: email path safely skipped
|
||||||
|
)
|
||||||
|
|
||||||
|
svc.RecordCyberPolicyEvent(context.Background(), CyberPolicyRecordInput{
|
||||||
|
UserID: 1,
|
||||||
|
UserEmail: "u@x.com",
|
||||||
|
Model: "gpt-5",
|
||||||
|
Endpoint: "/v1/responses",
|
||||||
|
UpstreamMessage: "flagged",
|
||||||
|
UpstreamBody: `{"error":{"code":"cyber_policy"}}`,
|
||||||
|
UpstreamStatus: 400,
|
||||||
|
})
|
||||||
|
|
||||||
|
logs := repo.snapshotLogs()
|
||||||
|
require.Len(t, logs, 1)
|
||||||
|
log := logs[0]
|
||||||
|
|
||||||
|
require.Equal(t, "cyber_policy", log.Action)
|
||||||
|
require.True(t, log.Flagged)
|
||||||
|
require.Equal(t, "cyber_policy", log.HighestCategory)
|
||||||
|
require.Contains(t, log.Error, "flagged")
|
||||||
|
require.False(t, log.AutoBanned)
|
||||||
|
// emailService is nil, so EmailSent must be false
|
||||||
|
require.False(t, log.EmailSent)
|
||||||
|
|
||||||
|
// UserID pointer must be set
|
||||||
|
require.NotNil(t, log.UserID)
|
||||||
|
require.Equal(t, int64(1), *log.UserID)
|
||||||
|
|
||||||
|
// score for cyber_policy is always 1.0
|
||||||
|
require.Equal(t, 1.0, log.HighestScore)
|
||||||
|
|
||||||
|
// mode must be post_upstream
|
||||||
|
require.Equal(t, "post_upstream", log.Mode)
|
||||||
|
|
||||||
|
// provider
|
||||||
|
require.Equal(t, "openai", log.Provider)
|
||||||
|
|
||||||
|
// model
|
||||||
|
require.Equal(t, "gpt-5", log.Model)
|
||||||
|
|
||||||
|
// endpoint
|
||||||
|
require.Equal(t, "/v1/responses", log.Endpoint)
|
||||||
|
|
||||||
|
// violation count >= 1 (side-effects ran)
|
||||||
|
require.GreaterOrEqual(t, log.ViolationCount, 1)
|
||||||
|
|
||||||
|
// Error field should also contain the upstream body JSON
|
||||||
|
require.True(t, strings.Contains(log.Error, "cyber_policy") || strings.Contains(log.Error, "flagged"),
|
||||||
|
"Error should mention flagged or cyber_policy")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestRecordCyberPolicyEvent_CreateLogBeforeEmail verifies F7: the moderation
|
||||||
|
// log is persisted BEFORE email delivery, and EmailSent is patched afterwards —
|
||||||
|
// SMTP hangs can no longer swallow the audit record.
|
||||||
|
//
|
||||||
|
// Note on email ordering: EmailService is a concrete type with no injectable
|
||||||
|
// send interface, so SMTP-success cannot be simulated in unit tests.
|
||||||
|
// With emailService=nil the email block is skipped and UpdateLogEmailSent is not
|
||||||
|
// called (correct: logPersisted && emailSent guard). The test therefore asserts
|
||||||
|
// the two invariants that ARE observable without real SMTP:
|
||||||
|
// 1. CreateLog runs first (calls[0]=="create").
|
||||||
|
// 2. The log is stored with EmailSent=false (not pre-set to true).
|
||||||
|
//
|
||||||
|
// The update_email_sent path is covered by integration/e2e tests where a real
|
||||||
|
// (or test-double) SMTP endpoint is available.
|
||||||
|
func TestRecordCyberPolicyEvent_CreateLogBeforeEmail(t *testing.T) {
|
||||||
|
repo := &cyberOrderingTestRepo{}
|
||||||
|
svc := NewContentModerationService(
|
||||||
|
&contentModerationTestSettingRepo{values: map[string]string{
|
||||||
|
SettingKeyRiskControlEnabled: "true",
|
||||||
|
}},
|
||||||
|
repo,
|
||||||
|
nil,
|
||||||
|
nil,
|
||||||
|
nil,
|
||||||
|
nil,
|
||||||
|
nil, // emailService=nil: email path safely skipped; see doc comment above
|
||||||
|
)
|
||||||
|
|
||||||
|
svc.RecordCyberPolicyEvent(context.Background(), CyberPolicyRecordInput{
|
||||||
|
RequestID: "req-1",
|
||||||
|
UserID: 7,
|
||||||
|
UserEmail: "u@example.com",
|
||||||
|
Model: "gpt-5",
|
||||||
|
UpstreamMessage: "blocked",
|
||||||
|
})
|
||||||
|
|
||||||
|
calls := repo.snapshot()
|
||||||
|
require.GreaterOrEqual(t, len(calls), 1, "CreateLog must be called")
|
||||||
|
require.Equal(t, "create", calls[0], "CreateLog must run first (F7: log-before-email)")
|
||||||
|
|
||||||
|
// EmailSent must be false when the log is first persisted (new code sets it
|
||||||
|
// false before CreateLog; email result is patched via UpdateLogEmailSent).
|
||||||
|
emailSents := repo.snapshotEmailSents()
|
||||||
|
require.NotEmpty(t, emailSents, "CreateLog must have captured EmailSent value")
|
||||||
|
require.False(t, emailSents[0], "log must be stored with EmailSent=false initially (F7)")
|
||||||
|
|
||||||
|
// With emailService=nil, no email is sent, so UpdateLogEmailSent must NOT
|
||||||
|
// be called (logPersisted && emailSent guard correctly suppresses the patch).
|
||||||
|
require.NotContains(t, calls, "update_email_sent",
|
||||||
|
"UpdateLogEmailSent must not be called when no email was sent")
|
||||||
|
}
|
||||||
|
|
||||||
|
// banCountArgsTestRepo 在 contentModerationTestRepo 基础上记录
|
||||||
|
// CountFlaggedByUserSince 收到的 excludeCyberPolicy 参数,供透传断言。
|
||||||
|
type banCountArgsTestRepo struct {
|
||||||
|
contentModerationTestRepo
|
||||||
|
argsMu sync.Mutex
|
||||||
|
countCalls []bool
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *banCountArgsTestRepo) CountFlaggedByUserSince(ctx context.Context, userID int64, since time.Time, excludeCyberPolicy bool) (int, error) {
|
||||||
|
r.argsMu.Lock()
|
||||||
|
r.countCalls = append(r.countCalls, excludeCyberPolicy)
|
||||||
|
r.argsMu.Unlock()
|
||||||
|
return r.contentModerationTestRepo.CountFlaggedByUserSince(ctx, userID, since, excludeCyberPolicy)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *banCountArgsTestRepo) snapshotCountCalls() []bool {
|
||||||
|
r.argsMu.Lock()
|
||||||
|
defer r.argsMu.Unlock()
|
||||||
|
out := make([]bool, len(r.countCalls))
|
||||||
|
copy(out, r.countCalls)
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestApplyFlaggedAccountSideEffects_PassesExcludeCyberFlag(t *testing.T) {
|
||||||
|
repo := &banCountArgsTestRepo{}
|
||||||
|
svc := NewContentModerationService(
|
||||||
|
&contentModerationTestSettingRepo{values: map[string]string{}},
|
||||||
|
repo, nil, nil, nil, nil, nil,
|
||||||
|
)
|
||||||
|
userID := int64(42)
|
||||||
|
|
||||||
|
cfgExclude := defaultContentModerationConfig()
|
||||||
|
cfgExclude.CyberPolicyExcludeFromBanCount = true
|
||||||
|
svc.applyFlaggedAccountSideEffects(context.Background(), cfgExclude, &ContentModerationLog{Flagged: true, UserID: &userID})
|
||||||
|
|
||||||
|
cfgDefault := defaultContentModerationConfig() // 默认 false
|
||||||
|
svc.applyFlaggedAccountSideEffects(context.Background(), cfgDefault, &ContentModerationLog{Flagged: true, UserID: &userID})
|
||||||
|
|
||||||
|
require.Equal(t, []bool{true, false}, repo.snapshotCountCalls(),
|
||||||
|
"applyFlaggedAccountSideEffects 必须把 cfg.CyberPolicyExcludeFromBanCount 透传给 COUNT 查询")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRecordCyberPolicyEvent_ExcludeFromBanCount_SkipsBanJudgment(t *testing.T) {
|
||||||
|
repo := &banCountArgsTestRepo{}
|
||||||
|
svc := NewContentModerationService(
|
||||||
|
&contentModerationTestSettingRepo{values: map[string]string{
|
||||||
|
SettingKeyRiskControlEnabled: "true",
|
||||||
|
SettingKeyContentModerationConfig: `{"cyber_policy_exclude_from_ban_count":true}`,
|
||||||
|
}},
|
||||||
|
repo, nil, nil, nil, nil, nil,
|
||||||
|
)
|
||||||
|
|
||||||
|
svc.RecordCyberPolicyEvent(context.Background(), CyberPolicyRecordInput{
|
||||||
|
UserID: 1,
|
||||||
|
UserEmail: "u@x.com",
|
||||||
|
Model: "gpt-5",
|
||||||
|
Endpoint: "/v1/responses",
|
||||||
|
UpstreamMessage: "flagged",
|
||||||
|
UpstreamStatus: 400,
|
||||||
|
})
|
||||||
|
|
||||||
|
require.Empty(t, repo.snapshotCountCalls(), "开关开时不得执行封号计数查询")
|
||||||
|
logs := repo.snapshotLogs()
|
||||||
|
require.Len(t, logs, 1, "风控日志必须照记")
|
||||||
|
require.True(t, logs[0].Flagged, "日志仍标记 Flagged=true(列表可见可筛)")
|
||||||
|
require.Equal(t, "cyber_policy", logs[0].Action)
|
||||||
|
require.Equal(t, 0, logs[0].ViolationCount, "不参与计数时 ViolationCount 保持 0")
|
||||||
|
require.False(t, logs[0].AutoBanned)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRecordCyberPolicyEvent_DefaultCountsTowardBan(t *testing.T) {
|
||||||
|
repo := &banCountArgsTestRepo{}
|
||||||
|
svc := NewContentModerationService(
|
||||||
|
&contentModerationTestSettingRepo{values: map[string]string{
|
||||||
|
SettingKeyRiskControlEnabled: "true",
|
||||||
|
}},
|
||||||
|
repo, nil, nil, nil, nil, nil,
|
||||||
|
)
|
||||||
|
|
||||||
|
svc.RecordCyberPolicyEvent(context.Background(), CyberPolicyRecordInput{
|
||||||
|
UserID: 1,
|
||||||
|
UserEmail: "u@x.com",
|
||||||
|
Model: "gpt-5",
|
||||||
|
Endpoint: "/v1/responses",
|
||||||
|
UpstreamMessage: "flagged",
|
||||||
|
UpstreamStatus: 400,
|
||||||
|
})
|
||||||
|
|
||||||
|
require.Equal(t, []bool{false}, repo.snapshotCountCalls(),
|
||||||
|
"默认配置必须执行计数查询且不排除 cyber 行")
|
||||||
|
logs := repo.snapshotLogs()
|
||||||
|
require.Len(t, logs, 1)
|
||||||
|
require.GreaterOrEqual(t, logs[0].ViolationCount, 1, "默认路径行为不变(现状回归)")
|
||||||
|
}
|
||||||
@@ -115,3 +115,41 @@ func defaultContentModerationString(value string, fallback string) string {
|
|||||||
}
|
}
|
||||||
return strings.TrimSpace(value)
|
return strings.TrimSpace(value)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// buildCyberPolicyNoticeEmailBody 是 cyber_policy 通知邮件的内置兜底正文,
|
||||||
|
// 当 notification email 模板渲染失败时使用(与 sendViolationEmail 的兜底同理)。
|
||||||
|
func buildCyberPolicyNoticeEmailBody(siteName string, log *ContentModerationLog) string {
|
||||||
|
if log == nil {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
userName := strings.TrimSpace(log.UserEmail)
|
||||||
|
if userName == "" && log.UserID != nil {
|
||||||
|
userName = fmt.Sprintf("UID %d", *log.UserID)
|
||||||
|
}
|
||||||
|
return fmt.Sprintf(`<!doctype html>
|
||||||
|
<html><body style="margin:0;padding:0;background:#f5f6fb;color:#222;font-family:-apple-system,BlinkMacSystemFont,'Segoe UI',Arial,sans-serif;">
|
||||||
|
<div style="max-width:680px;margin:0 auto;padding:32px 20px;">
|
||||||
|
<div style="height:8px;background:#ef4444;border-radius:14px 14px 0 0;"></div>
|
||||||
|
<div style="background:#fff;border-radius:0 0 14px 14px;padding:40px 48px;box-shadow:0 8px 28px rgba(15,23,42,.08);">
|
||||||
|
<div style="letter-spacing:4px;color:#999;font-size:14px;text-transform:uppercase;">Risk Control / 网络安全策略</div>
|
||||||
|
<h1 style="margin:20px 0 28px;font-size:30px;line-height:1.25;">请求被网络安全策略拦截</h1>
|
||||||
|
<p style="font-size:17px;line-height:1.9;margin:0 0 24px;">尊敬的用户 <strong>%s</strong>,您的请求被上游网络安全策略(cyber policy)拦截。</p>
|
||||||
|
<div style="background:#fff1f2;border:1px solid #fecdd3;border-radius:12px;padding:22px 28px;margin:28px 0;">
|
||||||
|
<table style="width:100%%;border-collapse:collapse;font-size:16px;">
|
||||||
|
<tr><td style="padding:12px 0;color:#888;border-bottom:1px solid #fee2e2;">触发时间</td><td style="padding:12px 0;border-bottom:1px solid #fee2e2;">%s</td></tr>
|
||||||
|
<tr><td style="padding:12px 0;color:#888;border-bottom:1px solid #fee2e2;">模型</td><td style="padding:12px 0;border-bottom:1px solid #fee2e2;">%s</td></tr>
|
||||||
|
<tr><td style="padding:12px 0;color:#888;">上游说明</td><td style="padding:12px 0;">%s</td></tr>
|
||||||
|
</table>
|
||||||
|
</div>
|
||||||
|
<p style="font-size:15px;line-height:1.8;color:#666;">如认为系误判,可调整请求措辞后重试,或申请获得授权的安全访问权限。</p>
|
||||||
|
<p style="font-size:14px;line-height:1.8;color:#777;margin-top:28px;">此邮件由 %s 自动发送,请勿回复。</p>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
</body></html>`,
|
||||||
|
html.EscapeString(userName),
|
||||||
|
html.EscapeString(log.CreatedAt.Format("2006-01-02 15:04:05")),
|
||||||
|
html.EscapeString(defaultContentModerationString(log.Model, "-")),
|
||||||
|
html.EscapeString(defaultContentModerationString(log.Error, "-")),
|
||||||
|
html.EscapeString(siteName),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|||||||
@@ -94,7 +94,7 @@ func (r *contentModerationTestRepo) ListLogs(ctx context.Context, filter Content
|
|||||||
return nil, nil, nil
|
return nil, nil, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *contentModerationTestRepo) CountFlaggedByUserSince(ctx context.Context, userID int64, since time.Time) (int, error) {
|
func (r *contentModerationTestRepo) CountFlaggedByUserSince(ctx context.Context, userID int64, since time.Time, excludeCyberPolicy bool) (int, error) {
|
||||||
r.mu.Lock()
|
r.mu.Lock()
|
||||||
defer r.mu.Unlock()
|
defer r.mu.Unlock()
|
||||||
count := 0
|
count := 0
|
||||||
@@ -102,6 +102,9 @@ func (r *contentModerationTestRepo) CountFlaggedByUserSince(ctx context.Context,
|
|||||||
if log.UserID == nil || *log.UserID != userID || !log.Flagged || log.Action == ContentModerationActionHashBlock {
|
if log.UserID == nil || *log.UserID != userID || !log.Flagged || log.Action == ContentModerationActionHashBlock {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
if excludeCyberPolicy && log.Action == ContentModerationActionCyberPolicy {
|
||||||
|
continue
|
||||||
|
}
|
||||||
if log.CreatedAt.IsZero() || log.CreatedAt.Before(since) {
|
if log.CreatedAt.IsZero() || log.CreatedAt.Before(since) {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
@@ -114,6 +117,10 @@ func (r *contentModerationTestRepo) CleanupExpiredLogs(ctx context.Context, hitB
|
|||||||
return &ContentModerationCleanupResult{}, nil
|
return &ContentModerationCleanupResult{}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (r *contentModerationTestRepo) UpdateLogEmailSent(ctx context.Context, id int64, sent bool) error {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
func (r *contentModerationTestRepo) snapshotLogs() []ContentModerationLog {
|
func (r *contentModerationTestRepo) snapshotLogs() []ContentModerationLog {
|
||||||
r.mu.Lock()
|
r.mu.Lock()
|
||||||
defer r.mu.Unlock()
|
defer r.mu.Unlock()
|
||||||
@@ -1789,3 +1796,44 @@ func TestContentModerationUnbanUser_ActiveUserOnlyInvalidatesAuthCache(t *testin
|
|||||||
func contentModerationIntPtr(v int) *int {
|
func contentModerationIntPtr(v int) *int {
|
||||||
return &v
|
return &v
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestContentModerationUpdateConfig_CyberPolicyExcludeFromBanCount(t *testing.T) {
|
||||||
|
settingRepo := &contentModerationTestSettingRepo{values: map[string]string{}}
|
||||||
|
svc := NewContentModerationService(settingRepo, nil, nil, nil, nil, nil, nil)
|
||||||
|
|
||||||
|
// 默认值必须是 false(计入,保持现状)
|
||||||
|
view, err := svc.GetConfig(context.Background())
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.False(t, view.CyberPolicyExcludeFromBanCount, "默认必须计入封号计数")
|
||||||
|
|
||||||
|
// 指针式部分更新为 true
|
||||||
|
exclude := true
|
||||||
|
view, err = svc.UpdateConfig(context.Background(), UpdateContentModerationConfigInput{
|
||||||
|
CyberPolicyExcludeFromBanCount: &exclude,
|
||||||
|
})
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.True(t, view.CyberPolicyExcludeFromBanCount)
|
||||||
|
|
||||||
|
// 持久化 JSON 含字段
|
||||||
|
var saved ContentModerationConfig
|
||||||
|
require.NoError(t, json.Unmarshal([]byte(settingRepo.values[SettingKeyContentModerationConfig]), &saved))
|
||||||
|
require.True(t, saved.CyberPolicyExcludeFromBanCount)
|
||||||
|
|
||||||
|
// 二次读取(从持久化 JSON 反序列化)roundtrip
|
||||||
|
view, err = svc.GetConfig(context.Background())
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.True(t, view.CyberPolicyExcludeFromBanCount)
|
||||||
|
|
||||||
|
// 不传该字段的更新不得改动它(指针 nil = 保留)
|
||||||
|
view, err = svc.UpdateConfig(context.Background(), UpdateContentModerationConfigInput{})
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.True(t, view.CyberPolicyExcludeFromBanCount)
|
||||||
|
|
||||||
|
// 主动回拨 false 必须生效(防止未来误加 if val 保护逻辑)
|
||||||
|
revert := false
|
||||||
|
view, err = svc.UpdateConfig(context.Background(), UpdateContentModerationConfigInput{
|
||||||
|
CyberPolicyExcludeFromBanCount: &revert,
|
||||||
|
})
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.False(t, view.CyberPolicyExcludeFromBanCount)
|
||||||
|
}
|
||||||
|
|||||||
@@ -136,6 +136,8 @@ const (
|
|||||||
SettingKeyAffiliateRebatePerInviteeCap = "affiliate_rebate_per_invitee_cap" // 单人返利上限(0=无上限)
|
SettingKeyAffiliateRebatePerInviteeCap = "affiliate_rebate_per_invitee_cap" // 单人返利上限(0=无上限)
|
||||||
SettingKeyRiskControlEnabled = "risk_control_enabled" // 是否启用风控中心入口与审计链路
|
SettingKeyRiskControlEnabled = "risk_control_enabled" // 是否启用风控中心入口与审计链路
|
||||||
SettingKeyContentModerationConfig = "content_moderation_config" // 内容审计配置(JSON)
|
SettingKeyContentModerationConfig = "content_moderation_config" // 内容审计配置(JSON)
|
||||||
|
SettingKeyCyberSessionBlockEnabled = "cyber_session_block_enabled" // cyber 命中后会话级自动屏蔽总开关(默认关)
|
||||||
|
SettingKeyCyberSessionBlockTTLSeconds = "cyber_session_block_ttl_seconds" // 会话屏蔽 TTL 秒数(默认 3600)
|
||||||
SettingKeyLoginAgreementEnabled = "login_agreement_enabled" // 登录前是否要求同意条款
|
SettingKeyLoginAgreementEnabled = "login_agreement_enabled" // 登录前是否要求同意条款
|
||||||
SettingKeyLoginAgreementMode = "login_agreement_mode" // 条款确认展示模式:modal / checkbox
|
SettingKeyLoginAgreementMode = "login_agreement_mode" // 条款确认展示模式:modal / checkbox
|
||||||
SettingKeyLoginAgreementUpdatedAt = "login_agreement_updated_at" // 条款更新日期(展示用)
|
SettingKeyLoginAgreementUpdatedAt = "login_agreement_updated_at" // 条款更新日期(展示用)
|
||||||
|
|||||||
@@ -30,6 +30,7 @@ const (
|
|||||||
NotificationEmailEventAccountQuotaAlert = "account.quota_alert"
|
NotificationEmailEventAccountQuotaAlert = "account.quota_alert"
|
||||||
NotificationEmailEventContentModerationViolation = "content_moderation.violation_notice"
|
NotificationEmailEventContentModerationViolation = "content_moderation.violation_notice"
|
||||||
NotificationEmailEventContentModerationDisabled = "content_moderation.account_disabled"
|
NotificationEmailEventContentModerationDisabled = "content_moderation.account_disabled"
|
||||||
|
NotificationEmailEventCyberPolicyNotice = "content_moderation.cyber_policy_notice"
|
||||||
NotificationEmailEventOpsAlert = "ops.alert"
|
NotificationEmailEventOpsAlert = "ops.alert"
|
||||||
NotificationEmailEventOpsScheduledReport = "ops.scheduled_report"
|
NotificationEmailEventOpsScheduledReport = "ops.scheduled_report"
|
||||||
|
|
||||||
@@ -947,6 +948,7 @@ var notificationEmailEventOrder = []string{
|
|||||||
NotificationEmailEventAccountQuotaAlert,
|
NotificationEmailEventAccountQuotaAlert,
|
||||||
NotificationEmailEventContentModerationViolation,
|
NotificationEmailEventContentModerationViolation,
|
||||||
NotificationEmailEventContentModerationDisabled,
|
NotificationEmailEventContentModerationDisabled,
|
||||||
|
NotificationEmailEventCyberPolicyNotice,
|
||||||
NotificationEmailEventOpsAlert,
|
NotificationEmailEventOpsAlert,
|
||||||
NotificationEmailEventOpsScheduledReport,
|
NotificationEmailEventOpsScheduledReport,
|
||||||
}
|
}
|
||||||
@@ -1035,6 +1037,15 @@ var notificationEmailEventDefinitions = map[string]NotificationEmailEventInfo{
|
|||||||
Placeholders: append(append([]string{}, notificationEmailCommonPlaceholders...),
|
Placeholders: append(append([]string{}, notificationEmailCommonPlaceholders...),
|
||||||
"triggered_at", "group_name", "moderation_category", "moderation_score", "violation_count", "ban_threshold"),
|
"triggered_at", "group_name", "moderation_category", "moderation_score", "violation_count", "ban_threshold"),
|
||||||
},
|
},
|
||||||
|
NotificationEmailEventCyberPolicyNotice: {
|
||||||
|
Event: NotificationEmailEventCyberPolicyNotice,
|
||||||
|
Label: "Cyber policy notice",
|
||||||
|
Description: "Sent to users when an upstream request is blocked by cyber-security policy (cyber_policy).",
|
||||||
|
Category: "risk_control",
|
||||||
|
Optional: false,
|
||||||
|
Placeholders: append(append([]string{}, notificationEmailCommonPlaceholders...),
|
||||||
|
"triggered_at", "model", "group_name", "upstream_message"),
|
||||||
|
},
|
||||||
NotificationEmailEventOpsAlert: {
|
NotificationEmailEventOpsAlert: {
|
||||||
Event: NotificationEmailEventOpsAlert,
|
Event: NotificationEmailEventOpsAlert,
|
||||||
Label: "Ops alert",
|
Label: "Ops alert",
|
||||||
@@ -1277,6 +1288,34 @@ var notificationEmailOfficialTemplates = map[string]map[string]notificationEmail
|
|||||||
<p>如需申诉或恢复账号,请联系平台管理员处理。</p>`),
|
<p>如需申诉或恢复账号,请联系平台管理员处理。</p>`),
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
|
NotificationEmailEventCyberPolicyNotice: {
|
||||||
|
notificationEmailDefaultLocale: {
|
||||||
|
Subject: "[{{site_name}}] Cyber-security policy notice",
|
||||||
|
HTML: notificationEmailCard("#ef4444", "Cyber-security policy notice", `
|
||||||
|
<p>Hello {{recipient_name}},</p>
|
||||||
|
<p>Your request was blocked by the upstream provider's cyber-security policy.</p>
|
||||||
|
<table style="width:100%;border-collapse:collapse;">
|
||||||
|
<tr><td>Triggered at</td><td>{{triggered_at}}</td></tr>
|
||||||
|
<tr><td>Model</td><td>{{model}}</td></tr>
|
||||||
|
<tr><td>Group</td><td>{{group_name}}</td></tr>
|
||||||
|
<tr><td>Upstream message</td><td>{{upstream_message}}</td></tr>
|
||||||
|
</table>
|
||||||
|
<p>If you believe this is a mistake, try rephrasing your request, or apply for authorized security access.</p>`),
|
||||||
|
},
|
||||||
|
notificationEmailLocaleChinese: {
|
||||||
|
Subject: "[{{site_name}}] 网络安全策略拦截提醒",
|
||||||
|
HTML: notificationEmailCard("#ef4444", "网络安全策略拦截提醒", `
|
||||||
|
<p>{{recipient_name}},您好:</p>
|
||||||
|
<p>您的请求被上游服务商的网络安全策略(cyber policy)拦截。</p>
|
||||||
|
<table style="width:100%;border-collapse:collapse;">
|
||||||
|
<tr><td>触发时间</td><td>{{triggered_at}}</td></tr>
|
||||||
|
<tr><td>模型</td><td>{{model}}</td></tr>
|
||||||
|
<tr><td>所属分组</td><td>{{group_name}}</td></tr>
|
||||||
|
<tr><td>上游说明</td><td>{{upstream_message}}</td></tr>
|
||||||
|
</table>
|
||||||
|
<p>如认为系误判,可调整请求措辞后重试,或申请获得授权的安全访问权限。</p>`),
|
||||||
|
},
|
||||||
|
},
|
||||||
NotificationEmailEventOpsAlert: {
|
NotificationEmailEventOpsAlert: {
|
||||||
notificationEmailDefaultLocale: {
|
notificationEmailDefaultLocale: {
|
||||||
Subject: "[Ops Alert][{{severity}}] {{rule_name}}",
|
Subject: "[Ops Alert][{{severity}}] {{rule_name}}",
|
||||||
|
|||||||
@@ -140,6 +140,7 @@ func TestNotificationEmailAdditionalEventsAreListedAndPreviewable(t *testing.T)
|
|||||||
{NotificationEmailEventAccountQuotaAlert, "account_name"},
|
{NotificationEmailEventAccountQuotaAlert, "account_name"},
|
||||||
{NotificationEmailEventContentModerationViolation, "moderation_category"},
|
{NotificationEmailEventContentModerationViolation, "moderation_category"},
|
||||||
{NotificationEmailEventContentModerationDisabled, "violation_count"},
|
{NotificationEmailEventContentModerationDisabled, "violation_count"},
|
||||||
|
{NotificationEmailEventCyberPolicyNotice, "upstream_message"},
|
||||||
{NotificationEmailEventOpsAlert, "rule_name"},
|
{NotificationEmailEventOpsAlert, "rule_name"},
|
||||||
{NotificationEmailEventOpsScheduledReport, "report_html"},
|
{NotificationEmailEventOpsScheduledReport, "report_html"},
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,88 @@
|
|||||||
|
package service
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/tidwall/gjson"
|
||||||
|
)
|
||||||
|
|
||||||
|
// opsCyberPolicyKey 在 gin context 中携带 cyber_policy 命中标记。
|
||||||
|
// 由 gateway 服务层在检测到上游 error.code=="cyber_policy" 时设置,
|
||||||
|
// handler 在 Forward 返回后读取以触发风控记录、邮件与 tokens=0 用量行。
|
||||||
|
const opsCyberPolicyKey = "ops_cyber_policy"
|
||||||
|
|
||||||
|
// errOpenAICyberPolicyForwarded 表示 cyber_policy 已按当前端点格式透传给客户端
|
||||||
|
// (error 已写出/下发)。compat 路径 ForwardAsChatCompletions / ForwardAsAnthropic 出口
|
||||||
|
// 据此丢弃 result 并返回该哨兵,使 handler 落入 tokens=0 免费用量行(对齐 /v1/responses),
|
||||||
|
// 既不计费、也不 failover、不重复写响应。
|
||||||
|
var errOpenAICyberPolicyForwarded = errors.New("openai cyber_policy forwarded to client")
|
||||||
|
|
||||||
|
// CyberPolicyMark 记录一次 cyber_policy 硬阻断的上游证据。
|
||||||
|
type CyberPolicyMark struct {
|
||||||
|
Code string // 固定 "cyber_policy"
|
||||||
|
Message string // 上游 error.message
|
||||||
|
Body string // 上游 response.failed / 400 原始 body(已截断;未脱敏,ops_error 落库由 sanitizeErrorBodyForStorage、风控日志由 redactContentModerationSecrets 统一脱敏)
|
||||||
|
UpstreamStatus int // 上游 HTTP 状态(流式=200,非流式=400)
|
||||||
|
UpstreamInTok int // 上游已报 input tokens(如有)
|
||||||
|
UpstreamOutTok int // 上游已报 output tokens(如有)
|
||||||
|
}
|
||||||
|
|
||||||
|
// MarkOpsCyberPolicy 记录 cyber 标记;首个写入生效,后续忽略(同一 turn 只记一次)。
|
||||||
|
// WS 多轮场景由 handler 在每个 turn 结束后调用 ClearOpsCyberPolicy 重置。
|
||||||
|
func MarkOpsCyberPolicy(c *gin.Context, mark CyberPolicyMark) {
|
||||||
|
if c == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if GetOpsCyberPolicy(c) != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
mark.Code = "cyber_policy"
|
||||||
|
mark.Message = strings.TrimSpace(mark.Message)
|
||||||
|
mark.Body = strings.TrimSpace(mark.Body)
|
||||||
|
c.Set(opsCyberPolicyKey, &mark)
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetOpsCyberPolicy 返回 cyber 标记,未命中(或已被 Clear)返回 nil。
|
||||||
|
func GetOpsCyberPolicy(c *gin.Context) *CyberPolicyMark {
|
||||||
|
if c == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if v, ok := c.Get(opsCyberPolicyKey); ok {
|
||||||
|
if m, ok := v.(*CyberPolicyMark); ok && m != nil {
|
||||||
|
return m
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ClearOpsCyberPolicy 清除 cyber 标记(typed-nil 覆盖;gin context 无并发安全的
|
||||||
|
// 删除原语,Set 走内部锁,与异步 GetOpsCyberPolicy 不构成 data race)。
|
||||||
|
// 仅 WS 多轮路径在 turn 收尾调用;HTTP 单请求路径不调用(context 随请求销毁,
|
||||||
|
// 且中间件 shouldSkipOpsErrorLogForCyber 依赖标记防双写)。
|
||||||
|
// WS 路径 clear 发生在中间件收尾之前,连接响应状态为 101,不触发中间件 status>=400
|
||||||
|
// 落库分支,故无双写/漏写。
|
||||||
|
func ClearOpsCyberPolicy(c *gin.Context) {
|
||||||
|
if c == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
c.Set(opsCyberPolicyKey, (*CyberPolicyMark)(nil))
|
||||||
|
}
|
||||||
|
|
||||||
|
// detectOpenAICyberPolicy 精确识别 cyber_policy(对齐 codex api_bridge.rs:145 /
|
||||||
|
// sse/responses.rs:529)。命中返回 (true, "cyber_policy", message)。
|
||||||
|
func detectOpenAICyberPolicy(payload []byte) (bool, string, string) {
|
||||||
|
code := gjson.GetBytes(payload, "error.code").String()
|
||||||
|
if code == "" {
|
||||||
|
code = gjson.GetBytes(payload, "response.error.code").String()
|
||||||
|
}
|
||||||
|
if !strings.EqualFold(strings.TrimSpace(code), "cyber_policy") {
|
||||||
|
return false, "", ""
|
||||||
|
}
|
||||||
|
msg := gjson.GetBytes(payload, "error.message").String()
|
||||||
|
if msg == "" {
|
||||||
|
msg = gjson.GetBytes(payload, "response.error.message").String()
|
||||||
|
}
|
||||||
|
return true, "cyber_policy", strings.TrimSpace(msg)
|
||||||
|
}
|
||||||
@@ -0,0 +1,88 @@
|
|||||||
|
package service
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/http/httptest"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestMarkAndGetOpsCyberPolicy(t *testing.T) {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(rec)
|
||||||
|
|
||||||
|
require.Nil(t, GetOpsCyberPolicy(c), "no mark initially")
|
||||||
|
|
||||||
|
MarkOpsCyberPolicy(c, CyberPolicyMark{
|
||||||
|
Code: "cyber_policy",
|
||||||
|
Message: "This request was flagged for cyber policy.",
|
||||||
|
Body: `{"error":{"code":"cyber_policy"}}`,
|
||||||
|
UpstreamStatus: 400,
|
||||||
|
})
|
||||||
|
|
||||||
|
got := GetOpsCyberPolicy(c)
|
||||||
|
require.NotNil(t, got)
|
||||||
|
require.Equal(t, "cyber_policy", got.Code)
|
||||||
|
require.Equal(t, 400, got.UpstreamStatus)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMarkOpsCyberPolicyFirstWins(t *testing.T) {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(rec)
|
||||||
|
MarkOpsCyberPolicy(c, CyberPolicyMark{Code: "cyber_policy", Message: "first"})
|
||||||
|
MarkOpsCyberPolicy(c, CyberPolicyMark{Code: "cyber_policy", Message: "second"})
|
||||||
|
require.Equal(t, "first", GetOpsCyberPolicy(c).Message, "first mark wins, later marks ignored")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMarkOpsCyberPolicyNilContext(t *testing.T) {
|
||||||
|
MarkOpsCyberPolicy(nil, CyberPolicyMark{Code: "cyber_policy"})
|
||||||
|
require.Nil(t, GetOpsCyberPolicy(nil))
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestClearOpsCyberPolicy_AllowsRemark verifies F1: after Clear, Get returns nil
|
||||||
|
// and a subsequent Mark takes effect (per-turn lifecycle in WS connections).
|
||||||
|
func TestClearOpsCyberPolicy_AllowsRemark(t *testing.T) {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
c, _ := gin.CreateTestContext(httptest.NewRecorder())
|
||||||
|
|
||||||
|
MarkOpsCyberPolicy(c, CyberPolicyMark{Message: "first", UpstreamStatus: 200})
|
||||||
|
require.NotNil(t, GetOpsCyberPolicy(c))
|
||||||
|
|
||||||
|
ClearOpsCyberPolicy(c)
|
||||||
|
require.Nil(t, GetOpsCyberPolicy(c), "mark must be invisible after Clear")
|
||||||
|
|
||||||
|
MarkOpsCyberPolicy(c, CyberPolicyMark{Message: "second", UpstreamStatus: 400})
|
||||||
|
got := GetOpsCyberPolicy(c)
|
||||||
|
require.NotNil(t, got, "re-mark after Clear must take effect")
|
||||||
|
require.Equal(t, "second", got.Message)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDetectOpenAICyberPolicy(t *testing.T) {
|
||||||
|
cases := []struct {
|
||||||
|
name string
|
||||||
|
payload string
|
||||||
|
hit bool
|
||||||
|
msg string
|
||||||
|
}{
|
||||||
|
{"top-level error", `{"error":{"code":"cyber_policy","message":"flagged"}}`, true, "flagged"},
|
||||||
|
{"response-wrapped", `{"response":{"error":{"code":"cyber_policy","message":" bad "}}}`, true, "bad"},
|
||||||
|
{"case-insensitive", `{"error":{"code":"Cyber_Policy"}}`, true, ""},
|
||||||
|
{"content_policy not cyber", `{"error":{"code":"content_policy","message":"x"}}`, false, ""},
|
||||||
|
{"safety message not cyber", `{"error":{"type":"safety_error","message":"high-risk cyber activity"}}`, false, ""},
|
||||||
|
{"empty", ``, false, ""},
|
||||||
|
{"upstream_error", `{"error":{"code":"upstream_error"}}`, false, ""},
|
||||||
|
}
|
||||||
|
for _, tc := range cases {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
hit, code, msg := detectOpenAICyberPolicy([]byte(tc.payload))
|
||||||
|
require.Equal(t, tc.hit, hit)
|
||||||
|
if tc.hit {
|
||||||
|
require.Equal(t, "cyber_policy", code)
|
||||||
|
require.Equal(t, tc.msg, msg)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,99 @@
|
|||||||
|
package service
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"crypto/sha256"
|
||||||
|
"encoding/hex"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/Wei-Shaw/sub2api/internal/pkg/logger"
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
)
|
||||||
|
|
||||||
|
// CyberSessionBlockStore 是 cyber 会话屏蔽表的存取接口。
|
||||||
|
// repository 层 gatewayCache 附带实现(类型断言探测接入,不改 GatewayCache
|
||||||
|
// 共享接口);测试 stub 不实现时屏蔽能力自动降级关闭。
|
||||||
|
type CyberSessionBlockStore interface {
|
||||||
|
SetCyberSessionBlocked(ctx context.Context, key string, ttl time.Duration) error
|
||||||
|
IsCyberSessionBlocked(ctx context.Context, key string) (bool, error)
|
||||||
|
}
|
||||||
|
|
||||||
|
// CyberSessionBlockKey 派生会话屏蔽 key:仅用显式会话标识(header
|
||||||
|
// session_id/conversation_id 或 body prompt_cache_key),混入 apiKeyID 隔离后
|
||||||
|
// sha256。无显式标识返回空串——调用方必须放行(粒度决策:不退化到
|
||||||
|
// user/apikey/内容派生)。
|
||||||
|
func CyberSessionBlockKey(apiKeyID int64, c *gin.Context, body []byte) string {
|
||||||
|
raw := explicitOpenAISessionID(c, body)
|
||||||
|
if raw == "" {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
isolated := isolateOpenAISessionID(apiKeyID, raw)
|
||||||
|
sum := sha256.Sum256([]byte(isolated))
|
||||||
|
return hex.EncodeToString(sum[:])
|
||||||
|
}
|
||||||
|
|
||||||
|
// cyberSessionBlockStore 探测 cache 是否具备屏蔽存储能力。
|
||||||
|
// 注意:若未来以装饰器包装 GatewayCache(如日志/指标装饰器),该装饰器必须同时实现
|
||||||
|
// CyberSessionBlockStore,否则会话屏蔽能力将静默降级关闭
|
||||||
|
// (编译断言 var _ service.CyberSessionBlockStore = (*gatewayCache)(nil) 只覆盖
|
||||||
|
// *gatewayCache 本体,无法覆盖其外层包装)。
|
||||||
|
func (s *OpenAIGatewayService) cyberSessionBlockStore() CyberSessionBlockStore {
|
||||||
|
if s == nil || s.cache == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
store, ok := s.cache.(CyberSessionBlockStore)
|
||||||
|
if !ok {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return store
|
||||||
|
}
|
||||||
|
|
||||||
|
// CyberSessionBlockRuntime 返回 (开关, TTL)。开关默认关。
|
||||||
|
// 委托给 SettingService.GetCyberSessionBlockRuntime,进程内缓存避免热路径 DB 往返。
|
||||||
|
func (s *OpenAIGatewayService) CyberSessionBlockRuntime(ctx context.Context) (bool, time.Duration) {
|
||||||
|
if s == nil || s.settingService == nil {
|
||||||
|
return false, time.Hour
|
||||||
|
}
|
||||||
|
return s.settingService.GetCyberSessionBlockRuntime(ctx)
|
||||||
|
}
|
||||||
|
|
||||||
|
// MarkCyberSessionBlocked 把会话写入屏蔽表(写入点:cyber 命中后)。
|
||||||
|
// 开关关闭、key 为空或存储不可用时静默跳过。
|
||||||
|
func (s *OpenAIGatewayService) MarkCyberSessionBlocked(ctx context.Context, key string) {
|
||||||
|
if key == "" {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
enabled, ttl := s.CyberSessionBlockRuntime(ctx)
|
||||||
|
if !enabled {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
store := s.cyberSessionBlockStore()
|
||||||
|
if store == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if err := store.SetCyberSessionBlocked(ctx, key, ttl); err != nil {
|
||||||
|
logger.LegacyPrintf("service.openai_gateway", "cyber session block write failed: err=%v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// IsCyberSessionBlocked 查询会话是否被屏蔽(拦截点)。开关关闭、key 为空、
|
||||||
|
// 存储不可用或查询出错时返回 false(fail-open:屏蔽是增强防护,不阻断主链路)。
|
||||||
|
func (s *OpenAIGatewayService) IsCyberSessionBlocked(ctx context.Context, key string) bool {
|
||||||
|
if key == "" {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
enabled, _ := s.CyberSessionBlockRuntime(ctx)
|
||||||
|
if !enabled {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
store := s.cyberSessionBlockStore()
|
||||||
|
if store == nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
blocked, err := store.IsCyberSessionBlocked(ctx, key)
|
||||||
|
if err != nil {
|
||||||
|
logger.LegacyPrintf("service.openai_gateway", "cyber session block read failed: err=%v", err)
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return blocked
|
||||||
|
}
|
||||||
@@ -0,0 +1,192 @@
|
|||||||
|
package service
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"net/http/httptest"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
func newCyberBlockTestCtx(headers map[string]string, body string) (*gin.Context, []byte) {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
c, _ := gin.CreateTestContext(httptest.NewRecorder())
|
||||||
|
req := httptest.NewRequest("POST", "/openai/v1/responses", strings.NewReader(body))
|
||||||
|
for k, v := range headers {
|
||||||
|
req.Header.Set(k, v)
|
||||||
|
}
|
||||||
|
c.Request = req
|
||||||
|
return c, []byte(body)
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestCyberSessionBlockKey verifies F5a key derivation: explicit session signals
|
||||||
|
// only (header session_id/conversation_id or body prompt_cache_key), apiKey
|
||||||
|
// isolated, and EMPTY when no explicit signal (no content-derived fallback —
|
||||||
|
// "不退化" decision).
|
||||||
|
func TestCyberSessionBlockKey(t *testing.T) {
|
||||||
|
c1, b1 := newCyberBlockTestCtx(map[string]string{"session_id": "sess-abc"}, `{}`)
|
||||||
|
k1 := CyberSessionBlockKey(101, c1, b1)
|
||||||
|
require.NotEmpty(t, k1)
|
||||||
|
|
||||||
|
// Same session, different apiKey → different key (isolation).
|
||||||
|
c2, b2 := newCyberBlockTestCtx(map[string]string{"session_id": "sess-abc"}, `{}`)
|
||||||
|
require.NotEqual(t, k1, CyberSessionBlockKey(202, c2, b2))
|
||||||
|
|
||||||
|
// Same session + same apiKey → stable key.
|
||||||
|
c3, b3 := newCyberBlockTestCtx(map[string]string{"session_id": "sess-abc"}, `{}`)
|
||||||
|
require.Equal(t, k1, CyberSessionBlockKey(101, c3, b3))
|
||||||
|
|
||||||
|
// prompt_cache_key in body counts as explicit.
|
||||||
|
c4, b4 := newCyberBlockTestCtx(nil, `{"prompt_cache_key":"pck-1"}`)
|
||||||
|
require.NotEmpty(t, CyberSessionBlockKey(101, c4, b4))
|
||||||
|
|
||||||
|
// No explicit signal → empty key → caller must skip blocking entirely.
|
||||||
|
c5, b5 := newCyberBlockTestCtx(nil, `{"input":"hello world"}`)
|
||||||
|
require.Empty(t, CyberSessionBlockKey(101, c5, b5))
|
||||||
|
|
||||||
|
// conversation_id header counts as explicit; key is stable and non-empty.
|
||||||
|
c6, b6 := newCyberBlockTestCtx(map[string]string{"conversation_id": "conv-xyz"}, `{}`)
|
||||||
|
k6 := CyberSessionBlockKey(101, c6, b6)
|
||||||
|
require.NotEmpty(t, k6)
|
||||||
|
c6b, b6b := newCyberBlockTestCtx(map[string]string{"conversation_id": "conv-xyz"}, `{}`)
|
||||||
|
require.Equal(t, k6, CyberSessionBlockKey(101, c6b, b6b), "conversation_id key must be stable")
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- fakes ---
|
||||||
|
|
||||||
|
type fakeCyberBlockStore struct {
|
||||||
|
blocked map[string]bool
|
||||||
|
}
|
||||||
|
|
||||||
|
var _ CyberSessionBlockStore = (*fakeCyberBlockStore)(nil)
|
||||||
|
|
||||||
|
func (f *fakeCyberBlockStore) SetCyberSessionBlocked(_ context.Context, key string, _ time.Duration) error {
|
||||||
|
if f.blocked == nil {
|
||||||
|
f.blocked = map[string]bool{}
|
||||||
|
}
|
||||||
|
f.blocked[key] = true
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *fakeCyberBlockStore) IsCyberSessionBlocked(_ context.Context, key string) (bool, error) {
|
||||||
|
return f.blocked[key], nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// fakeSettingRepo is a minimal SettingRepository stub for unit tests.
|
||||||
|
// Only GetValue is exercised by GetCyberSessionBlockRuntime; all other methods
|
||||||
|
// panic so accidental calls are caught immediately.
|
||||||
|
type fakeSettingRepo struct {
|
||||||
|
vals map[string]string
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *fakeSettingRepo) GetValue(_ context.Context, key string) (string, error) {
|
||||||
|
v, ok := r.vals[key]
|
||||||
|
if !ok {
|
||||||
|
return "", ErrSettingNotFound
|
||||||
|
}
|
||||||
|
return v, nil
|
||||||
|
}
|
||||||
|
func (r *fakeSettingRepo) Get(_ context.Context, _ string) (*Setting, error) {
|
||||||
|
panic("fakeSettingRepo.Get not implemented")
|
||||||
|
}
|
||||||
|
func (r *fakeSettingRepo) Set(_ context.Context, _, _ string) error {
|
||||||
|
panic("fakeSettingRepo.Set not implemented")
|
||||||
|
}
|
||||||
|
func (r *fakeSettingRepo) GetMultiple(_ context.Context, _ []string) (map[string]string, error) {
|
||||||
|
panic("fakeSettingRepo.GetMultiple not implemented")
|
||||||
|
}
|
||||||
|
func (r *fakeSettingRepo) SetMultiple(_ context.Context, _ map[string]string) error {
|
||||||
|
panic("fakeSettingRepo.SetMultiple not implemented")
|
||||||
|
}
|
||||||
|
func (r *fakeSettingRepo) GetAll(_ context.Context) (map[string]string, error) {
|
||||||
|
panic("fakeSettingRepo.GetAll not implemented")
|
||||||
|
}
|
||||||
|
func (r *fakeSettingRepo) Delete(_ context.Context, _ string) error {
|
||||||
|
panic("fakeSettingRepo.Delete not implemented")
|
||||||
|
}
|
||||||
|
|
||||||
|
var _ SettingRepository = (*fakeSettingRepo)(nil)
|
||||||
|
|
||||||
|
// comboCacheAndStore implements both GatewayCache (no-op stubs) and
|
||||||
|
// CyberSessionBlockStore (delegates to fakeCyberBlockStore) so it can be
|
||||||
|
// injected as s.cache and successfully type-asserted to CyberSessionBlockStore.
|
||||||
|
type comboCacheAndStore struct {
|
||||||
|
store fakeCyberBlockStore
|
||||||
|
}
|
||||||
|
|
||||||
|
var _ GatewayCache = (*comboCacheAndStore)(nil)
|
||||||
|
var _ CyberSessionBlockStore = (*comboCacheAndStore)(nil)
|
||||||
|
|
||||||
|
func (c *comboCacheAndStore) GetSessionAccountID(_ context.Context, _ int64, _ string) (int64, error) {
|
||||||
|
return 0, errors.New("stub")
|
||||||
|
}
|
||||||
|
func (c *comboCacheAndStore) SetSessionAccountID(_ context.Context, _ int64, _ string, _ int64, _ time.Duration) error {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
func (c *comboCacheAndStore) RefreshSessionTTL(_ context.Context, _ int64, _ string, _ time.Duration) error {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
func (c *comboCacheAndStore) DeleteSessionAccountID(_ context.Context, _ int64, _ string) error {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
func (c *comboCacheAndStore) SetCyberSessionBlocked(ctx context.Context, key string, ttl time.Duration) error {
|
||||||
|
return c.store.SetCyberSessionBlocked(ctx, key, ttl)
|
||||||
|
}
|
||||||
|
func (c *comboCacheAndStore) IsCyberSessionBlocked(ctx context.Context, key string) (bool, error) {
|
||||||
|
return c.store.IsCyberSessionBlocked(ctx, key)
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- tests ---
|
||||||
|
|
||||||
|
// TestIsCyberSessionBlocked_EmptyKeyAndNilService covers the fail-open paths:
|
||||||
|
// empty key, nil service, store missing → always false / no panic.
|
||||||
|
func TestIsCyberSessionBlocked_EmptyKeyAndNilService(t *testing.T) {
|
||||||
|
var nilSvc *OpenAIGatewayService
|
||||||
|
require.False(t, nilSvc.IsCyberSessionBlocked(context.Background(), "k"))
|
||||||
|
require.NotPanics(t, func() { nilSvc.MarkCyberSessionBlocked(context.Background(), "k") })
|
||||||
|
|
||||||
|
svc := &OpenAIGatewayService{}
|
||||||
|
require.False(t, svc.IsCyberSessionBlocked(context.Background(), ""))
|
||||||
|
require.False(t, svc.IsCyberSessionBlocked(context.Background(), "k"), "no store + no settings → fail-open false")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestCyberSessionBlock_RoundTrip exercises the type-assertion success path:
|
||||||
|
// mark a session blocked via a combo cache+store, then confirm IsCyberSessionBlocked
|
||||||
|
// returns true, and an unrelated key returns false.
|
||||||
|
func TestCyberSessionBlock_RoundTrip(t *testing.T) {
|
||||||
|
// SettingService with only settingRepo set — GetCyberSessionBlockRuntime needs
|
||||||
|
// nothing else (cfg/proxyRepo/etc. are not touched by this code path).
|
||||||
|
settingSvc := &SettingService{
|
||||||
|
settingRepo: &fakeSettingRepo{
|
||||||
|
vals: map[string]string{
|
||||||
|
SettingKeyCyberSessionBlockEnabled: "true",
|
||||||
|
SettingKeyCyberSessionBlockTTLSeconds: "60",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
combo := &comboCacheAndStore{}
|
||||||
|
svc := &OpenAIGatewayService{
|
||||||
|
cache: combo,
|
||||||
|
settingService: settingSvc,
|
||||||
|
}
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
const testKey = "deadbeef1234"
|
||||||
|
|
||||||
|
// Before marking: not blocked.
|
||||||
|
require.False(t, svc.IsCyberSessionBlocked(ctx, testKey))
|
||||||
|
|
||||||
|
// Mark as blocked.
|
||||||
|
svc.MarkCyberSessionBlocked(ctx, testKey)
|
||||||
|
|
||||||
|
// After marking: blocked.
|
||||||
|
require.True(t, svc.IsCyberSessionBlocked(ctx, testKey))
|
||||||
|
|
||||||
|
// Different key: still not blocked.
|
||||||
|
require.False(t, svc.IsCyberSessionBlocked(ctx, "other-key"))
|
||||||
|
}
|
||||||
@@ -312,6 +312,15 @@ func (s *OpenAIGatewayService) ForwardAsChatCompletions(
|
|||||||
result, handleErr = s.handleChatBufferedStreamingResponse(resp, c, account, originalModel, billingModel, upstreamModel, startTime)
|
result, handleErr = s.handleChatBufferedStreamingResponse(resp, c, account, originalModel, billingModel, upstreamModel, startTime)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// cyber_policy:标记已设、error 已按 Chat Completions 格式发给客户端。丢弃 result、
|
||||||
|
// 返回哨兵,使 handler 落入 tokens=0 免费用量行(对齐 /v1/responses),不计费、不 failover。
|
||||||
|
if GetOpsCyberPolicy(c) != nil {
|
||||||
|
if handleErr == nil {
|
||||||
|
handleErr = errOpenAICyberPolicyForwarded
|
||||||
|
}
|
||||||
|
return nil, handleErr
|
||||||
|
}
|
||||||
|
|
||||||
// Propagate ServiceTier and ReasoningEffort to result for billing
|
// Propagate ServiceTier and ReasoningEffort to result for billing
|
||||||
if handleErr == nil && result != nil {
|
if handleErr == nil && result != nil {
|
||||||
if responsesReq.ServiceTier != "" {
|
if responsesReq.ServiceTier != "" {
|
||||||
@@ -412,6 +421,24 @@ func (s *OpenAIGatewayService) handleChatBufferedStreamingResponse(
|
|||||||
}
|
}
|
||||||
if strings.TrimSpace(finalResponse.Status) == "failed" {
|
if strings.TrimSpace(finalResponse.Status) == "failed" {
|
||||||
payload, _ := json.Marshal(gin.H{"type": "response.failed", "response": finalResponse})
|
payload, _ := json.Marshal(gin.H{"type": "response.failed", "response": finalResponse})
|
||||||
|
// cyber_policy 致命不可重试:不 failover,以 Chat Completions 错误格式回写(F4),
|
||||||
|
// 标记供 handler 事后写风控/邮件/tokens=0 用量行。
|
||||||
|
if hit, code, msg := detectOpenAICyberPolicy(payload); hit {
|
||||||
|
MarkOpsCyberPolicy(c, CyberPolicyMark{
|
||||||
|
Code: code,
|
||||||
|
Message: msg,
|
||||||
|
Body: truncateString(string(payload), 4096),
|
||||||
|
UpstreamStatus: http.StatusOK,
|
||||||
|
UpstreamInTok: usage.InputTokens,
|
||||||
|
UpstreamOutTok: usage.OutputTokens,
|
||||||
|
})
|
||||||
|
clientMsg := msg
|
||||||
|
if clientMsg == "" {
|
||||||
|
clientMsg = "Request blocked by upstream cyber-security policy"
|
||||||
|
}
|
||||||
|
writeChatCompletionsError(c, http.StatusBadRequest, "invalid_request_error", clientMsg)
|
||||||
|
return nil, fmt.Errorf("openai cyber_policy: %s", msg)
|
||||||
|
}
|
||||||
return nil, s.newOpenAIStreamFailoverError(c, account, false, requestID, payload, openAICompatFailedResponseMessage(finalResponse))
|
return nil, s.newOpenAIStreamFailoverError(c, account, false, requestID, payload, openAICompatFailedResponseMessage(finalResponse))
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -550,9 +577,41 @@ func (s *OpenAIGatewayService) handleChatStreamingResponse(
|
|||||||
if strings.TrimSpace(event.Type) == "response.failed" {
|
if strings.TrimSpace(event.Type) == "response.failed" {
|
||||||
payloadBytes := []byte(payload)
|
payloadBytes := []byte(payload)
|
||||||
message := extractOpenAISSEErrorMessage(payloadBytes)
|
message := extractOpenAISSEErrorMessage(payloadBytes)
|
||||||
|
if hit, code, msg := detectOpenAICyberPolicy(payloadBytes); hit {
|
||||||
|
// cyber_policy 致命且不可重试:不 failover。下发标准 error chunk +
|
||||||
|
// [DONE],让程序化客户端可感知并停止重试(F4);标记供 handler 事后
|
||||||
|
// 写风控/邮件。
|
||||||
|
MarkOpsCyberPolicy(c, CyberPolicyMark{
|
||||||
|
Code: code,
|
||||||
|
Message: msg,
|
||||||
|
Body: truncateString(string(payloadBytes), 4096),
|
||||||
|
UpstreamStatus: http.StatusOK,
|
||||||
|
UpstreamInTok: usage.InputTokens,
|
||||||
|
UpstreamOutTok: usage.OutputTokens,
|
||||||
|
})
|
||||||
|
if !clientDisconnected {
|
||||||
|
// 被 refusal 检测扣留的 pendingSSE 有意丢弃——cyber 拦截优先于部分内容下发。
|
||||||
|
writeStreamHeaders()
|
||||||
|
clientMsg := msg
|
||||||
|
if clientMsg == "" {
|
||||||
|
clientMsg = "Request blocked by upstream cyber-security policy"
|
||||||
|
}
|
||||||
|
if _, err := fmt.Fprint(c.Writer, buildChatStreamErrorSSE(code, clientMsg)); err == nil {
|
||||||
|
_, _ = fmt.Fprint(c.Writer, "data: [DONE]\n\n")
|
||||||
|
if fl, ok := c.Writer.(http.Flusher); ok {
|
||||||
|
fl.Flush()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// 无条件置位:成功路径防 finalizeStream 重复 [DONE];写失败意味着连接已不可写,
|
||||||
|
// finalizeStream 的 [DONE] 同样发不出去,统一抑制。
|
||||||
|
clientDisconnected = true
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
} else {
|
||||||
streamFailoverErr = s.newOpenAIStreamFailoverError(c, account, false, requestID, payloadBytes, message)
|
streamFailoverErr = s.newOpenAIStreamFailoverError(c, account, false, requestID, payloadBytes, message)
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
}
|
||||||
|
|
||||||
chunks := apicompat.ResponsesEventToChatChunks(&event, state)
|
chunks := apicompat.ResponsesEventToChatChunks(&event, state)
|
||||||
if !clientDisconnected {
|
if !clientDisconnected {
|
||||||
@@ -861,3 +920,21 @@ func writeChatCompletionsError(c *gin.Context, statusCode int, errType, message
|
|||||||
},
|
},
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// buildChatStreamErrorSSE builds one SSE data frame carrying an OpenAI chat
|
||||||
|
// streaming error object. Used when the stream must terminate with a visible
|
||||||
|
// error (e.g. upstream cyber_policy), so programmatic clients stop retrying.
|
||||||
|
// Marshal 失败的兜底会丢弃 message 原文,仅保留 code 与固定提示。
|
||||||
|
func buildChatStreamErrorSSE(code, message string) string {
|
||||||
|
payload, err := json.Marshal(gin.H{
|
||||||
|
"error": gin.H{
|
||||||
|
"type": "invalid_request_error",
|
||||||
|
"code": code,
|
||||||
|
"message": message,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return "data: {\"error\":{\"type\":\"invalid_request_error\",\"code\":\"" + code + "\",\"message\":\"upstream error\"}}\n\n"
|
||||||
|
}
|
||||||
|
return "data: " + string(payload) + "\n\n"
|
||||||
|
}
|
||||||
|
|||||||
@@ -316,6 +316,51 @@ func TestForwardAsChatCompletions_StreamResponseFailedTriggersFailoverBeforeFlus
|
|||||||
require.False(t, c.Writer.Written())
|
require.False(t, c.Writer.Written())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestForwardAsChatCompletions_StreamCyberPolicyNoFailover(t *testing.T) {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(rec)
|
||||||
|
body := []byte(`{"model":"gpt-5.5","messages":[{"role":"user","content":"` + strings.Repeat("large prompt ", 6000) + `"}],"stream":true}`)
|
||||||
|
c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", bytes.NewReader(body))
|
||||||
|
c.Request.Header.Set("Content-Type", "application/json")
|
||||||
|
|
||||||
|
upstreamBody := strings.Join([]string{
|
||||||
|
`data: {"type":"response.created","response":{"id":"resp_cyber","model":"gpt-5.5","status":"in_progress","output":[]}}`,
|
||||||
|
"",
|
||||||
|
`event: response.failed`,
|
||||||
|
`data: {"type":"response.failed","response":{"id":"resp_cyber","object":"response","model":"gpt-5.5","status":"failed","output":[],"error":{"code":"cyber_policy","message":"flagged for cyber policy"}}}`,
|
||||||
|
"",
|
||||||
|
}, "\n")
|
||||||
|
upstream := &httpUpstreamRecorder{resp: &http.Response{
|
||||||
|
StatusCode: http.StatusOK,
|
||||||
|
Header: http.Header{"Content-Type": []string{"text/event-stream"}, "x-request-id": []string{"rid_chat_cyber"}},
|
||||||
|
Body: io.NopCloser(strings.NewReader(upstreamBody)),
|
||||||
|
}}
|
||||||
|
|
||||||
|
svc := &OpenAIGatewayService{httpUpstream: upstream}
|
||||||
|
account := &Account{
|
||||||
|
ID: 1,
|
||||||
|
Name: "openai-oauth",
|
||||||
|
Platform: PlatformOpenAI,
|
||||||
|
Type: AccountTypeOAuth,
|
||||||
|
Concurrency: 1,
|
||||||
|
Credentials: map[string]any{
|
||||||
|
"access_token": "oauth-token",
|
||||||
|
"chatgpt_account_id": "chatgpt-acc",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err := svc.ForwardAsChatCompletions(context.Background(), c, account, body, "", "gpt-5.5")
|
||||||
|
var failoverErr *UpstreamFailoverError
|
||||||
|
require.False(t, errors.As(err, &failoverErr), "cyber must NOT trigger failover")
|
||||||
|
require.NotNil(t, GetOpsCyberPolicy(c), "cyber mark must be set")
|
||||||
|
respBody := rec.Body.String()
|
||||||
|
require.Contains(t, respBody, `"error"`)
|
||||||
|
require.Contains(t, respBody, `"cyber_policy"`)
|
||||||
|
require.Contains(t, respBody, "data: [DONE]")
|
||||||
|
}
|
||||||
|
|
||||||
func TestForwardAsChatCompletions_StreamsUsageWithoutClientStreamOptions(t *testing.T) {
|
func TestForwardAsChatCompletions_StreamsUsageWithoutClientStreamOptions(t *testing.T) {
|
||||||
gin.SetMode(gin.TestMode)
|
gin.SetMode(gin.TestMode)
|
||||||
|
|
||||||
@@ -773,3 +818,14 @@ func TestForwardAsChatCompletions_UpstreamRequestIgnoresClientCancel(t *testing.
|
|||||||
require.NotNil(t, upstream.lastReq)
|
require.NotNil(t, upstream.lastReq)
|
||||||
require.NoError(t, upstream.lastReq.Context().Err())
|
require.NoError(t, upstream.lastReq.Context().Err())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// TestBuildChatStreamErrorSSE verifies F4: the error chunk payload follows the
|
||||||
|
// OpenAI chat streaming error convention so third-party clients stop retrying.
|
||||||
|
func TestBuildChatStreamErrorSSE(t *testing.T) {
|
||||||
|
got := buildChatStreamErrorSSE("cyber_policy", "blocked by policy")
|
||||||
|
require.True(t, strings.HasPrefix(got, "data: "), "must be an SSE data frame")
|
||||||
|
payload := strings.TrimSuffix(strings.TrimPrefix(got, "data: "), "\n\n")
|
||||||
|
require.Equal(t, "invalid_request_error", gjson.Get(payload, "error.type").String())
|
||||||
|
require.Equal(t, "cyber_policy", gjson.Get(payload, "error.code").String())
|
||||||
|
require.Equal(t, "blocked by policy", gjson.Get(payload, "error.message").String())
|
||||||
|
}
|
||||||
|
|||||||
@@ -0,0 +1,136 @@
|
|||||||
|
package service
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
// compatCyberOAuthAccount 是 compat cyber 测试共用的 OAuth 账号。
|
||||||
|
func compatCyberOAuthAccount() *Account {
|
||||||
|
return &Account{
|
||||||
|
ID: 1,
|
||||||
|
Name: "openai-oauth",
|
||||||
|
Platform: PlatformOpenAI,
|
||||||
|
Type: AccountTypeOAuth,
|
||||||
|
Concurrency: 1,
|
||||||
|
Credentials: map[string]any{
|
||||||
|
"access_token": "oauth-token",
|
||||||
|
"chatgpt_account_id": "chatgpt-acc",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// compatCyberUpstreamSSE 构造上游 responses SSE:response.created 后 response.failed(cyber_policy)。
|
||||||
|
func compatCyberUpstreamSSE() string {
|
||||||
|
return strings.Join([]string{
|
||||||
|
`data: {"type":"response.created","response":{"id":"resp_cyber","model":"gpt-5.5","status":"in_progress","output":[]}}`,
|
||||||
|
"",
|
||||||
|
`event: response.failed`,
|
||||||
|
`data: {"type":"response.failed","response":{"id":"resp_cyber","object":"response","model":"gpt-5.5","status":"failed","output":[],"error":{"code":"cyber_policy","message":"flagged for cyber policy"}}}`,
|
||||||
|
"",
|
||||||
|
}, "\n")
|
||||||
|
}
|
||||||
|
|
||||||
|
func compatCyberUpstreamRecorder() *httpUpstreamRecorder {
|
||||||
|
return &httpUpstreamRecorder{resp: &http.Response{
|
||||||
|
StatusCode: http.StatusOK,
|
||||||
|
Header: http.Header{"Content-Type": []string{"text/event-stream"}, "x-request-id": []string{"rid_cyber"}},
|
||||||
|
Body: io.NopCloser(strings.NewReader(compatCyberUpstreamSSE())),
|
||||||
|
}}
|
||||||
|
}
|
||||||
|
|
||||||
|
// C-1: chat completions 非流式客户端(buffered 路径)cyber 命中——不 failover、标记已设、
|
||||||
|
// 以 chat 错误格式回写、丢弃 result(使 handler 落入 tokens=0 免费用量行而非 RecordUsage 扣费)。
|
||||||
|
func TestForwardAsChatCompletions_BufferedCyberPolicyNoFailover(t *testing.T) {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(rec)
|
||||||
|
body := []byte(`{"model":"gpt-5.5","messages":[{"role":"user","content":"hi"}],"stream":false}`)
|
||||||
|
c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", bytes.NewReader(body))
|
||||||
|
c.Request.Header.Set("Content-Type", "application/json")
|
||||||
|
|
||||||
|
svc := &OpenAIGatewayService{httpUpstream: compatCyberUpstreamRecorder()}
|
||||||
|
|
||||||
|
result, err := svc.ForwardAsChatCompletions(context.Background(), c, compatCyberOAuthAccount(), body, "", "gpt-5.5")
|
||||||
|
require.Error(t, err)
|
||||||
|
require.Nil(t, result, "cyber must drop result so handler writes tokens=0 free row, not RecordUsage")
|
||||||
|
var failoverErr *UpstreamFailoverError
|
||||||
|
require.False(t, errors.As(err, &failoverErr), "cyber must NOT trigger failover")
|
||||||
|
mark := GetOpsCyberPolicy(c)
|
||||||
|
require.NotNil(t, mark, "cyber mark must be set for handler-side recording")
|
||||||
|
require.Equal(t, "cyber_policy", mark.Code)
|
||||||
|
require.True(t, c.Writer.Written(), "cyber error must be written to client (passthrough)")
|
||||||
|
}
|
||||||
|
|
||||||
|
// I-1: chat completions 流式客户端 cyber 命中——result 必须被丢弃(返回 nil),
|
||||||
|
// 使 handler forwardErrored 分支走 tokens=0 免费行,而非 RecordUsage(CyberBlocked) 扣费。
|
||||||
|
func TestForwardAsChatCompletions_StreamCyberPolicyDropsResult(t *testing.T) {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(rec)
|
||||||
|
body := []byte(`{"model":"gpt-5.5","messages":[{"role":"user","content":"hi"}],"stream":true}`)
|
||||||
|
c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", bytes.NewReader(body))
|
||||||
|
c.Request.Header.Set("Content-Type", "application/json")
|
||||||
|
|
||||||
|
svc := &OpenAIGatewayService{httpUpstream: compatCyberUpstreamRecorder()}
|
||||||
|
|
||||||
|
result, err := svc.ForwardAsChatCompletions(context.Background(), c, compatCyberOAuthAccount(), body, "", "gpt-5.5")
|
||||||
|
require.Error(t, err)
|
||||||
|
require.Nil(t, result, "cyber must drop result so handler does not bill via RecordUsage")
|
||||||
|
var failoverErr *UpstreamFailoverError
|
||||||
|
require.False(t, errors.As(err, &failoverErr), "cyber must NOT trigger failover")
|
||||||
|
require.NotNil(t, GetOpsCyberPolicy(c), "cyber mark must be set")
|
||||||
|
require.Contains(t, rec.Body.String(), "data: [DONE]", "stream must terminate with [DONE]")
|
||||||
|
}
|
||||||
|
|
||||||
|
// anthropic 非流式客户端(buffered 路径)cyber 命中——不 failover、标记已设、以 anthropic 错误格式回写、丢弃 result。
|
||||||
|
func TestForwardAsAnthropic_BufferedCyberPolicyNoFailover(t *testing.T) {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(rec)
|
||||||
|
body := []byte(`{"model":"gpt-5.5","max_tokens":1024,"messages":[{"role":"user","content":"hi"}],"stream":false}`)
|
||||||
|
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", bytes.NewReader(body))
|
||||||
|
c.Request.Header.Set("Content-Type", "application/json")
|
||||||
|
|
||||||
|
svc := &OpenAIGatewayService{httpUpstream: compatCyberUpstreamRecorder()}
|
||||||
|
|
||||||
|
result, err := svc.ForwardAsAnthropic(context.Background(), c, compatCyberOAuthAccount(), body, "", "gpt-5.5")
|
||||||
|
require.Error(t, err)
|
||||||
|
require.Nil(t, result, "cyber must drop result so handler writes tokens=0 free row")
|
||||||
|
var failoverErr *UpstreamFailoverError
|
||||||
|
require.False(t, errors.As(err, &failoverErr), "cyber must NOT trigger failover")
|
||||||
|
mark := GetOpsCyberPolicy(c)
|
||||||
|
require.NotNil(t, mark, "cyber mark must be set")
|
||||||
|
require.Equal(t, "cyber_policy", mark.Code)
|
||||||
|
require.True(t, c.Writer.Written(), "anthropic cyber error must be written to client")
|
||||||
|
require.Contains(t, rec.Body.String(), `"type":"error"`, "must use anthropic error envelope")
|
||||||
|
}
|
||||||
|
|
||||||
|
// anthropic 流式客户端 cyber 命中——不 failover、标记已设、下发 anthropic SSE error 事件、丢弃 result。
|
||||||
|
func TestForwardAsAnthropic_StreamCyberPolicyNoFailover(t *testing.T) {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(rec)
|
||||||
|
body := []byte(`{"model":"gpt-5.5","max_tokens":1024,"messages":[{"role":"user","content":"hi"}],"stream":true}`)
|
||||||
|
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", bytes.NewReader(body))
|
||||||
|
c.Request.Header.Set("Content-Type", "application/json")
|
||||||
|
|
||||||
|
svc := &OpenAIGatewayService{httpUpstream: compatCyberUpstreamRecorder()}
|
||||||
|
|
||||||
|
result, err := svc.ForwardAsAnthropic(context.Background(), c, compatCyberOAuthAccount(), body, "", "gpt-5.5")
|
||||||
|
require.Error(t, err)
|
||||||
|
require.Nil(t, result, "cyber must drop result so handler does not bill via RecordUsage")
|
||||||
|
var failoverErr *UpstreamFailoverError
|
||||||
|
require.False(t, errors.As(err, &failoverErr), "cyber must NOT trigger failover")
|
||||||
|
require.NotNil(t, GetOpsCyberPolicy(c), "cyber mark must be set")
|
||||||
|
require.Contains(t, rec.Body.String(), "event: error", "must emit anthropic SSE error event")
|
||||||
|
}
|
||||||
@@ -366,6 +366,15 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic(
|
|||||||
result, handleErr = s.handleAnthropicBufferedStreamingResponse(resp, c, originalModel, billingModel, upstreamModel, startTime)
|
result, handleErr = s.handleAnthropicBufferedStreamingResponse(resp, c, originalModel, billingModel, upstreamModel, startTime)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// cyber_policy:标记已设、error 已按 Anthropic 格式发给客户端。丢弃 result、返回哨兵,
|
||||||
|
// 使 handler 落入 tokens=0 免费用量行(对齐 /v1/responses),不计费、不 failover。
|
||||||
|
if GetOpsCyberPolicy(c) != nil {
|
||||||
|
if handleErr == nil {
|
||||||
|
handleErr = errOpenAICyberPolicyForwarded
|
||||||
|
}
|
||||||
|
return nil, handleErr
|
||||||
|
}
|
||||||
|
|
||||||
// Propagate ServiceTier and ReasoningEffort to result for billing
|
// Propagate ServiceTier and ReasoningEffort to result for billing
|
||||||
if handleErr == nil && result != nil {
|
if handleErr == nil && result != nil {
|
||||||
if compatContinuationEnabled && promptCacheKey != "" && result.ResponseID != "" {
|
if compatContinuationEnabled && promptCacheKey != "" && result.ResponseID != "" {
|
||||||
@@ -444,6 +453,28 @@ func (s *OpenAIGatewayService) handleAnthropicBufferedStreamingResponse(
|
|||||||
return nil, fmt.Errorf("upstream stream ended without terminal event")
|
return nil, fmt.Errorf("upstream stream ended without terminal event")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// cyber_policy:上游硬阻断(response.failed)。anthropic buffered 原对 failed 无特殊分支,
|
||||||
|
// 此处仅为 cyber 增加:以 Anthropic 错误格式回写,标记供 handler 事后写风控/邮件/tokens=0 用量行。
|
||||||
|
if strings.TrimSpace(finalResponse.Status) == "failed" {
|
||||||
|
payload, _ := json.Marshal(gin.H{"type": "response.failed", "response": finalResponse})
|
||||||
|
if hit, code, msg := detectOpenAICyberPolicy(payload); hit {
|
||||||
|
MarkOpsCyberPolicy(c, CyberPolicyMark{
|
||||||
|
Code: code,
|
||||||
|
Message: msg,
|
||||||
|
Body: truncateString(string(payload), 4096),
|
||||||
|
UpstreamStatus: http.StatusOK,
|
||||||
|
UpstreamInTok: usage.InputTokens,
|
||||||
|
UpstreamOutTok: usage.OutputTokens,
|
||||||
|
})
|
||||||
|
clientMsg := msg
|
||||||
|
if clientMsg == "" {
|
||||||
|
clientMsg = "Request blocked by upstream cyber-security policy"
|
||||||
|
}
|
||||||
|
writeAnthropicError(c, http.StatusBadRequest, "invalid_request_error", clientMsg)
|
||||||
|
return nil, fmt.Errorf("openai cyber_policy: %s", msg)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// When the terminal event has an empty output array, reconstruct from
|
// When the terminal event has an empty output array, reconstruct from
|
||||||
// accumulated delta events so the client receives the full content.
|
// accumulated delta events so the client receives the full content.
|
||||||
acc.SupplementResponseOutput(finalResponse)
|
acc.SupplementResponseOutput(finalResponse)
|
||||||
@@ -770,6 +801,32 @@ func (s *OpenAIGatewayService) handleAnthropicStreamingResponse(
|
|||||||
if event.Usage != nil {
|
if event.Usage != nil {
|
||||||
usage = copyOpenAIUsageFromResponsesUsage(event.Usage)
|
usage = copyOpenAIUsageFromResponsesUsage(event.Usage)
|
||||||
}
|
}
|
||||||
|
// cyber_policy 致命不可重试:标记供 handler 事后记录;以 Anthropic SSE error 事件
|
||||||
|
// 回写让客户端感知并停止重试(F4),丢弃后续转换输出。
|
||||||
|
if strings.TrimSpace(event.Type) == "response.failed" {
|
||||||
|
if hit, code, msg := detectOpenAICyberPolicy([]byte(payload)); hit {
|
||||||
|
MarkOpsCyberPolicy(c, CyberPolicyMark{
|
||||||
|
Code: code,
|
||||||
|
Message: msg,
|
||||||
|
Body: truncateString(payload, 4096),
|
||||||
|
UpstreamStatus: http.StatusOK,
|
||||||
|
UpstreamInTok: usage.InputTokens,
|
||||||
|
UpstreamOutTok: usage.OutputTokens,
|
||||||
|
})
|
||||||
|
if !clientDisconnected {
|
||||||
|
writeStreamHeaders()
|
||||||
|
clientMsg := msg
|
||||||
|
if clientMsg == "" {
|
||||||
|
clientMsg = "Request blocked by upstream cyber-security policy"
|
||||||
|
}
|
||||||
|
if _, err := fmt.Fprint(c.Writer, buildAnthropicStreamErrorSSE("invalid_request_error", clientMsg)); err == nil {
|
||||||
|
c.Writer.Flush()
|
||||||
|
}
|
||||||
|
clientDisconnected = true
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Convert to Anthropic events
|
// Convert to Anthropic events
|
||||||
@@ -1013,6 +1070,24 @@ func writeAnthropicError(c *gin.Context, statusCode int, errType, message string
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// buildAnthropicStreamErrorSSE builds one Anthropic SSE `error` event so a
|
||||||
|
// streaming response can terminate with a visible error (e.g. upstream
|
||||||
|
// cyber_policy) and programmatic clients stop retrying.
|
||||||
|
// Marshal 失败的兜底仅保留固定提示。
|
||||||
|
func buildAnthropicStreamErrorSSE(errType, message string) string {
|
||||||
|
payload, err := json.Marshal(gin.H{
|
||||||
|
"type": "error",
|
||||||
|
"error": gin.H{
|
||||||
|
"type": errType,
|
||||||
|
"message": message,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return "event: error\ndata: {\"type\":\"error\",\"error\":{\"type\":\"" + errType + "\",\"message\":\"upstream error\"}}\n\n"
|
||||||
|
}
|
||||||
|
return "event: error\ndata: " + string(payload) + "\n\n"
|
||||||
|
}
|
||||||
|
|
||||||
func copyOpenAIUsageFromResponsesUsage(usage *apicompat.ResponsesUsage) OpenAIUsage {
|
func copyOpenAIUsageFromResponsesUsage(usage *apicompat.ResponsesUsage) OpenAIUsage {
|
||||||
if usage == nil {
|
if usage == nil {
|
||||||
return OpenAIUsage{}
|
return OpenAIUsage{}
|
||||||
|
|||||||
@@ -58,6 +58,75 @@ func TestOpenAIGatewayServiceRecordUsage_RejectsNilInput(t *testing.T) {
|
|||||||
require.Error(t, svc.RecordUsage(context.Background(), &OpenAIRecordUsageInput{}))
|
require.Error(t, svc.RecordUsage(context.Background(), &OpenAIRecordUsageInput{}))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestRecordCyberPolicyUsageLog_BillsRealUpstreamTokens(t *testing.T) {
|
||||||
|
usageRepo := &openAIRecordUsageLogRepoStub{inserted: true}
|
||||||
|
userRepo := &openAIRecordUsageUserRepoStub{}
|
||||||
|
subRepo := &openAIRecordUsageSubRepoStub{}
|
||||||
|
svc := newOpenAIRecordUsageServiceForTest(usageRepo, userRepo, subRepo, nil)
|
||||||
|
usage := OpenAIUsage{InputTokens: 1200, OutputTokens: 300}
|
||||||
|
|
||||||
|
// 流式 cyber:上游 response.failed 报告了真实 token,须按真实 token 计费并扣费,
|
||||||
|
// 与 WS cyber / 正常请求口径一致(不再是 tokens=0 免费行)。
|
||||||
|
svc.RecordCyberPolicyUsageLog(context.Background(), CyberPolicyUsageInput{
|
||||||
|
APIKey: &APIKey{ID: 2, User: &User{ID: 1}},
|
||||||
|
Account: &Account{ID: 3},
|
||||||
|
RequestID: "rid-cyber-stream",
|
||||||
|
Model: "gpt-5.1",
|
||||||
|
Stream: true,
|
||||||
|
InputTokens: 1200,
|
||||||
|
OutputTokens: 300,
|
||||||
|
})
|
||||||
|
|
||||||
|
require.Equal(t, 1, usageRepo.calls)
|
||||||
|
require.NotNil(t, usageRepo.lastLog)
|
||||||
|
require.Equal(t, "gpt-5.1", usageRepo.lastLog.Model)
|
||||||
|
require.Equal(t, 1200, usageRepo.lastLog.InputTokens)
|
||||||
|
require.Equal(t, 300, usageRepo.lastLog.OutputTokens)
|
||||||
|
require.Equal(t, RequestTypeCyberBlocked, usageRepo.lastLog.RequestType, "cyber 行须标 request_type=cyber")
|
||||||
|
require.True(t, usageRepo.lastLog.Stream, "cyber 不覆盖真实 stream 字段")
|
||||||
|
|
||||||
|
expected := expectedOpenAICost(t, svc, "gpt-5.1", usage, 1.1)
|
||||||
|
require.Greater(t, usageRepo.lastLog.ActualCost, 0.0, "流式 cyber 有真实 token,须计费")
|
||||||
|
require.InDelta(t, expected.ActualCost, usageRepo.lastLog.ActualCost, 1e-12)
|
||||||
|
require.Equal(t, 1, userRepo.deductCalls, "按真实 token 扣费,与 WS/正常请求一致")
|
||||||
|
require.InDelta(t, expected.ActualCost, userRepo.lastAmount, 1e-12)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRecordCyberPolicyUsageLog_NonStreamZeroTokensZeroCost(t *testing.T) {
|
||||||
|
usageRepo := &openAIRecordUsageLogRepoStub{inserted: true}
|
||||||
|
userRepo := &openAIRecordUsageUserRepoStub{}
|
||||||
|
subRepo := &openAIRecordUsageSubRepoStub{}
|
||||||
|
svc := newOpenAIRecordUsageServiceForTest(usageRepo, userRepo, subRepo, nil)
|
||||||
|
|
||||||
|
// 非流式直接拒:上游未报 token,mark token 为 0 → cost 自然为 0,仍写一条 cyber 行(可见)。
|
||||||
|
svc.RecordCyberPolicyUsageLog(context.Background(), CyberPolicyUsageInput{
|
||||||
|
APIKey: &APIKey{ID: 2, User: &User{ID: 1}},
|
||||||
|
Account: &Account{ID: 3},
|
||||||
|
RequestID: "rid-cyber-400",
|
||||||
|
Model: "gpt-5.1",
|
||||||
|
Stream: false,
|
||||||
|
})
|
||||||
|
|
||||||
|
require.Equal(t, 1, usageRepo.calls)
|
||||||
|
require.NotNil(t, usageRepo.lastLog)
|
||||||
|
require.Equal(t, 0, usageRepo.lastLog.InputTokens)
|
||||||
|
require.Equal(t, 0, usageRepo.lastLog.OutputTokens)
|
||||||
|
require.Zero(t, usageRepo.lastLog.TotalCost)
|
||||||
|
require.Equal(t, RequestTypeCyberBlocked, usageRepo.lastLog.RequestType)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRecordCyberPolicyUsageLog_SkipsWhenIncomplete(t *testing.T) {
|
||||||
|
usageRepo := &openAIRecordUsageLogRepoStub{}
|
||||||
|
svc := newOpenAIRecordUsageServiceForTest(usageRepo, &openAIRecordUsageUserRepoStub{}, &openAIRecordUsageSubRepoStub{}, nil)
|
||||||
|
|
||||||
|
acct := &Account{ID: 3}
|
||||||
|
svc.RecordCyberPolicyUsageLog(context.Background(), CyberPolicyUsageInput{Account: acct, Model: "gpt-5"}) // APIKey nil
|
||||||
|
svc.RecordCyberPolicyUsageLog(context.Background(), CyberPolicyUsageInput{APIKey: &APIKey{ID: 2}, Account: acct, Model: "gpt-5"}) // User nil
|
||||||
|
svc.RecordCyberPolicyUsageLog(context.Background(), CyberPolicyUsageInput{APIKey: &APIKey{ID: 2, User: &User{ID: 1}}, Model: "gpt-5"}) // Account nil
|
||||||
|
svc.RecordCyberPolicyUsageLog(context.Background(), CyberPolicyUsageInput{APIKey: &APIKey{ID: 2, User: &User{ID: 1}}, Account: acct}) // Model 空
|
||||||
|
require.Equal(t, 0, usageRepo.calls, "APIKey/User/Account 缺失或 Model 空时跳过,不记不扣费")
|
||||||
|
}
|
||||||
|
|
||||||
type openAIRecordUsageUserRepoStub struct {
|
type openAIRecordUsageUserRepoStub struct {
|
||||||
UserRepository
|
UserRepository
|
||||||
|
|
||||||
@@ -1761,6 +1830,30 @@ func TestGatewayServiceCalculateRecordUsageCost_ChannelImageBillingUsesSizeTier(
|
|||||||
require.InDelta(t, 0.80, cost.ActualCost, 1e-12)
|
require.InDelta(t, 0.80, cost.ActualCost, 1e-12)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestRecordUsageMarksCyberRequestType(t *testing.T) {
|
||||||
|
logStub := &openAIRecordUsageLogRepoStub{inserted: true}
|
||||||
|
userStub := &openAIRecordUsageUserRepoStub{}
|
||||||
|
subStub := &openAIRecordUsageSubRepoStub{}
|
||||||
|
rateStub := &openAIUserGroupRateRepoStub{}
|
||||||
|
svc := newOpenAIRecordUsageServiceForTest(logStub, userStub, subStub, rateStub)
|
||||||
|
|
||||||
|
in := &OpenAIRecordUsageInput{
|
||||||
|
CyberBlocked: true,
|
||||||
|
Result: &OpenAIForwardResult{
|
||||||
|
Model: "gpt-5",
|
||||||
|
Duration: time.Second,
|
||||||
|
Usage: OpenAIUsage{InputTokens: 100, OutputTokens: 0},
|
||||||
|
},
|
||||||
|
APIKey: &APIKey{ID: 2, Group: &Group{RateMultiplier: 1}},
|
||||||
|
User: &User{ID: 1},
|
||||||
|
Account: &Account{ID: 3},
|
||||||
|
}
|
||||||
|
require.NoError(t, svc.RecordUsage(context.Background(), in))
|
||||||
|
require.NotNil(t, logStub.lastLog)
|
||||||
|
require.Equal(t, RequestTypeCyberBlocked, logStub.lastLog.RequestType)
|
||||||
|
require.Equal(t, 100, logStub.lastLog.InputTokens, "计费 token 不变(正常计费)")
|
||||||
|
}
|
||||||
|
|
||||||
func TestGatewayServiceCalculateRecordUsageCost_ChannelImageBillingNormalizesMissingSizeTier(t *testing.T) {
|
func TestGatewayServiceCalculateRecordUsageCost_ChannelImageBillingNormalizesMissingSizeTier(t *testing.T) {
|
||||||
groupID := int64(128)
|
groupID := int64(128)
|
||||||
defaultPrice := 0.10
|
defaultPrice := 0.10
|
||||||
|
|||||||
@@ -3555,6 +3555,19 @@ func (s *OpenAIGatewayService) handleErrorResponsePassthrough(
|
|||||||
MarkResponseCommitted(c)
|
MarkResponseCommitted(c)
|
||||||
body := s.readUpstreamErrorBody(resp)
|
body := s.readUpstreamErrorBody(resp)
|
||||||
|
|
||||||
|
// cyber_policy:透传账号本就把原始 body 回给客户端(下方 c.Data),此处仅打标记,
|
||||||
|
// 供 handler 事后写风控/邮件。cyber 是上游网络安全策略拦截,不冷却账号,
|
||||||
|
// 故下方跳过 handleOpenAIAccountUpstreamError(避免自定义 temp-unschedulable 规则误冷却)。
|
||||||
|
cyberHit, cyberCode, cyberMsg := detectOpenAICyberPolicy(body)
|
||||||
|
if cyberHit {
|
||||||
|
MarkOpsCyberPolicy(c, CyberPolicyMark{
|
||||||
|
Code: cyberCode,
|
||||||
|
Message: cyberMsg,
|
||||||
|
Body: truncateString(string(body), 4096),
|
||||||
|
UpstreamStatus: resp.StatusCode,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(body))
|
upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(body))
|
||||||
upstreamMsg = sanitizeUpstreamErrorMessage(upstreamMsg)
|
upstreamMsg = sanitizeUpstreamErrorMessage(upstreamMsg)
|
||||||
upstreamDetail := ""
|
upstreamDetail := ""
|
||||||
@@ -3568,9 +3581,11 @@ func (s *OpenAIGatewayService) handleErrorResponsePassthrough(
|
|||||||
setOpsUpstreamError(c, resp.StatusCode, upstreamMsg, upstreamDetail)
|
setOpsUpstreamError(c, resp.StatusCode, upstreamMsg, upstreamDetail)
|
||||||
logOpenAIInstructionsRequiredDebug(ctx, c, account, resp.StatusCode, upstreamMsg, requestBody, body)
|
logOpenAIInstructionsRequiredDebug(ctx, c, account, resp.StatusCode, upstreamMsg, requestBody, body)
|
||||||
// 透传模式保留原始上游错误响应,但运行态账号状态仍需更新,
|
// 透传模式保留原始上游错误响应,但运行态账号状态仍需更新,
|
||||||
// 避免粘性路由继续复用刚被限流的账号。
|
// 避免粘性路由继续复用刚被限流的账号。cyber 例外:不冷却账号。
|
||||||
|
if !cyberHit {
|
||||||
reqModel, _, _ := extractOpenAIRequestMetaFromBody(requestBody)
|
reqModel, _, _ := extractOpenAIRequestMetaFromBody(requestBody)
|
||||||
_ = s.handleOpenAIAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, body, reqModel)
|
_ = s.handleOpenAIAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, body, reqModel)
|
||||||
|
}
|
||||||
appendOpsUpstreamError(c, OpsUpstreamErrorEvent{
|
appendOpsUpstreamError(c, OpsUpstreamErrorEvent{
|
||||||
Platform: account.Platform,
|
Platform: account.Platform,
|
||||||
AccountID: account.ID,
|
AccountID: account.ID,
|
||||||
@@ -3852,7 +3867,20 @@ func (s *OpenAIGatewayService) handleStreamingResponsePassthrough(
|
|||||||
eventType := strings.TrimSpace(gjson.Get(trimmedData, "type").String())
|
eventType := strings.TrimSpace(gjson.Get(trimmedData, "type").String())
|
||||||
if eventType == "response.failed" {
|
if eventType == "response.failed" {
|
||||||
failedMessage = extractOpenAISSEErrorMessage(dataBytes)
|
failedMessage = extractOpenAISSEErrorMessage(dataBytes)
|
||||||
if !openAIStreamClientOutputStarted(c, clientOutputStarted) && openAIStreamFailedEventShouldFailover(dataBytes, failedMessage) {
|
// response.failed 自带上游已消耗的 usage(input token 通常已扣);必须先解析
|
||||||
|
// 再打 cyber 标记,否则 mark 记到的是解析前的 0,导致流式 cyber 按 0 token 计费
|
||||||
|
// 而漏记真实用量。对齐 WS V2 / Chat 流式路径(均先解析 usage 再 Mark)。
|
||||||
|
s.parseSSEUsageBytes(dataBytes, usage)
|
||||||
|
if hit, code, msg := detectOpenAICyberPolicy(dataBytes); hit {
|
||||||
|
MarkOpsCyberPolicy(c, CyberPolicyMark{
|
||||||
|
Code: code,
|
||||||
|
Message: msg,
|
||||||
|
Body: truncateString(string(dataBytes), 4096),
|
||||||
|
UpstreamStatus: http.StatusOK,
|
||||||
|
UpstreamInTok: usage.InputTokens,
|
||||||
|
UpstreamOutTok: usage.OutputTokens,
|
||||||
|
})
|
||||||
|
} else if !openAIStreamClientOutputStarted(c, clientOutputStarted) && openAIStreamFailedEventShouldFailover(dataBytes, failedMessage) {
|
||||||
return resultWithUsage(),
|
return resultWithUsage(),
|
||||||
s.newOpenAIStreamFailoverError(c, account, true, upstreamRequestID, dataBytes, failedMessage)
|
s.newOpenAIStreamFailoverError(c, account, true, upstreamRequestID, dataBytes, failedMessage)
|
||||||
}
|
}
|
||||||
@@ -4257,6 +4285,29 @@ func (s *OpenAIGatewayService) handleErrorResponse(
|
|||||||
) (*OpenAIForwardResult, error) {
|
) (*OpenAIForwardResult, error) {
|
||||||
body := s.readUpstreamErrorBody(resp)
|
body := s.readUpstreamErrorBody(resp)
|
||||||
|
|
||||||
|
// cyber_policy 硬阻断:透传上游原始错误体给客户端(不重包成通用 502),不冷却账号。
|
||||||
|
// 当前请求恒透传(需求1);标记供 handler 事后写风控/邮件。400 cyber 不可 failover
|
||||||
|
// (shouldFailoverUpstreamError(400)=false),故走到此处即可安全早返回。
|
||||||
|
if hit, code, cyberMsg := detectOpenAICyberPolicy(body); hit {
|
||||||
|
MarkOpsCyberPolicy(c, CyberPolicyMark{
|
||||||
|
Code: code,
|
||||||
|
Message: cyberMsg,
|
||||||
|
Body: truncateString(string(body), 4096),
|
||||||
|
UpstreamStatus: resp.StatusCode,
|
||||||
|
})
|
||||||
|
setOpsUpstreamError(c, resp.StatusCode, cyberMsg, truncateString(string(body), 2048))
|
||||||
|
writeOpenAIPassthroughResponseHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter)
|
||||||
|
contentType := resp.Header.Get("Content-Type")
|
||||||
|
if contentType == "" {
|
||||||
|
contentType = "application/json"
|
||||||
|
}
|
||||||
|
c.Data(resp.StatusCode, contentType, body)
|
||||||
|
if cyberMsg == "" {
|
||||||
|
return nil, fmt.Errorf("openai cyber_policy: %d", resp.StatusCode)
|
||||||
|
}
|
||||||
|
return nil, fmt.Errorf("openai cyber_policy: %s", cyberMsg)
|
||||||
|
}
|
||||||
|
|
||||||
upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(body))
|
upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(body))
|
||||||
upstreamMsg = sanitizeUpstreamErrorMessage(upstreamMsg)
|
upstreamMsg = sanitizeUpstreamErrorMessage(upstreamMsg)
|
||||||
upstreamDetail := ""
|
upstreamDetail := ""
|
||||||
@@ -4422,6 +4473,29 @@ func (s *OpenAIGatewayService) handleCompatErrorResponse(
|
|||||||
) (*OpenAIForwardResult, error) {
|
) (*OpenAIForwardResult, error) {
|
||||||
body := s.readUpstreamErrorBody(resp)
|
body := s.readUpstreamErrorBody(resp)
|
||||||
|
|
||||||
|
// cyber_policy:兼容路径(Chat Completions / Anthropic)以各自格式回写错误,
|
||||||
|
// 不原样透传 responses 格式的 cyber body(否则对下游格式不合法)。cyber 是上游网络
|
||||||
|
// 安全策略拦截,不冷却账号,故标记后直接以兼容格式回写错误并返回,跳过下方
|
||||||
|
// handleOpenAIAccountUpstreamError(避免自定义 temp-unschedulable 规则误冷却)。
|
||||||
|
if hit, code, cyberMsg := detectOpenAICyberPolicy(body); hit {
|
||||||
|
MarkOpsCyberPolicy(c, CyberPolicyMark{
|
||||||
|
Code: code,
|
||||||
|
Message: cyberMsg,
|
||||||
|
Body: truncateString(string(body), 4096),
|
||||||
|
UpstreamStatus: resp.StatusCode,
|
||||||
|
})
|
||||||
|
setOpsUpstreamError(c, resp.StatusCode, cyberMsg, truncateString(string(body), 2048))
|
||||||
|
clientMsg := cyberMsg
|
||||||
|
if clientMsg == "" {
|
||||||
|
clientMsg = "Request blocked by upstream cyber-security policy"
|
||||||
|
}
|
||||||
|
writeError(c, resp.StatusCode, "invalid_request_error", clientMsg)
|
||||||
|
if cyberMsg == "" {
|
||||||
|
return nil, fmt.Errorf("openai cyber_policy: %d", resp.StatusCode)
|
||||||
|
}
|
||||||
|
return nil, fmt.Errorf("openai cyber_policy: %s", cyberMsg)
|
||||||
|
}
|
||||||
|
|
||||||
upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(body))
|
upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(body))
|
||||||
if upstreamMsg == "" {
|
if upstreamMsg == "" {
|
||||||
upstreamMsg = fmt.Sprintf("Upstream error: %d", resp.StatusCode)
|
upstreamMsg = fmt.Sprintf("Upstream error: %d", resp.StatusCode)
|
||||||
@@ -4744,7 +4818,20 @@ func (s *OpenAIGatewayService) handleStreamingResponse(ctx context.Context, resp
|
|||||||
forceFlushFailedEvent := false
|
forceFlushFailedEvent := false
|
||||||
if eventType == "response.failed" {
|
if eventType == "response.failed" {
|
||||||
failedMessage = extractOpenAISSEErrorMessage(dataBytes)
|
failedMessage = extractOpenAISSEErrorMessage(dataBytes)
|
||||||
if !openAIStreamClientOutputStarted(c, clientOutputStarted) && openAIStreamFailedEventShouldFailover(dataBytes, failedMessage) {
|
// response.failed 自带上游已消耗的 usage(input token 通常已扣);必须先解析
|
||||||
|
// 再打 cyber 标记,否则 mark 记到的是解析前的 0,导致流式 cyber 按 0 token 计费
|
||||||
|
// 而漏记真实用量。对齐 WS V2 / Chat 流式路径(均先解析 usage 再 Mark)。
|
||||||
|
s.parseSSEUsageBytes(dataBytes, usage)
|
||||||
|
if hit, code, msg := detectOpenAICyberPolicy(dataBytes); hit {
|
||||||
|
MarkOpsCyberPolicy(c, CyberPolicyMark{
|
||||||
|
Code: code,
|
||||||
|
Message: msg,
|
||||||
|
Body: truncateString(string(dataBytes), 4096),
|
||||||
|
UpstreamStatus: http.StatusOK,
|
||||||
|
UpstreamInTok: usage.InputTokens,
|
||||||
|
UpstreamOutTok: usage.OutputTokens,
|
||||||
|
})
|
||||||
|
} else if !openAIStreamClientOutputStarted(c, clientOutputStarted) && openAIStreamFailedEventShouldFailover(dataBytes, failedMessage) {
|
||||||
sawFailedEvent = true
|
sawFailedEvent = true
|
||||||
streamFailoverErr = s.newOpenAIStreamFailoverError(c, account, false, upstreamRequestID, dataBytes, failedMessage)
|
streamFailoverErr = s.newOpenAIStreamFailoverError(c, account, false, upstreamRequestID, dataBytes, failedMessage)
|
||||||
return
|
return
|
||||||
@@ -5731,9 +5818,72 @@ type OpenAIRecordUsageInput struct {
|
|||||||
IPAddress string // 请求的客户端 IP 地址
|
IPAddress string // 请求的客户端 IP 地址
|
||||||
RequestPayloadHash string
|
RequestPayloadHash string
|
||||||
APIKeyService APIKeyQuotaUpdater
|
APIKeyService APIKeyQuotaUpdater
|
||||||
|
// CyberBlocked 为 true 时把该用量行标记为 cyber(request_type=cyber),计费逻辑不变。
|
||||||
|
CyberBlocked bool
|
||||||
ChannelUsageFields
|
ChannelUsageFields
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// CyberPolicyUsageInput 是 cyber 拒绝、未走正常 RecordUsage 的请求记录用量的入参。
|
||||||
|
// 用量按上游真实 token 计费,与 WS cyber 及正常请求口径一致(InputTokens/OutputTokens
|
||||||
|
// 取自上游 response.failed 报告的 usage,即 mark.UpstreamInTok/OutTok)。
|
||||||
|
type CyberPolicyUsageInput struct {
|
||||||
|
APIKey *APIKey
|
||||||
|
Account *Account
|
||||||
|
Subscription *UserSubscription
|
||||||
|
RequestID string
|
||||||
|
Model string
|
||||||
|
Stream bool
|
||||||
|
InputTokens int
|
||||||
|
OutputTokens int
|
||||||
|
// 渠道归因与请求级 meta,使 cyber 计费行与正常 RecordUsage 行口径一致
|
||||||
|
// (否则 cyber 行 channel_id 等为空,渠道维度统计会遗漏 cyber 命中)。
|
||||||
|
InboundEndpoint string
|
||||||
|
UpstreamEndpoint string
|
||||||
|
UserAgent string
|
||||||
|
IPAddress string
|
||||||
|
RequestPayloadHash string
|
||||||
|
APIKeyService APIKeyQuotaUpdater
|
||||||
|
ChannelUsageFields
|
||||||
|
}
|
||||||
|
|
||||||
|
// RecordCyberPolicyUsageLog 为被上游 cyber_policy 拒绝、未走正常 RecordUsage 的请求
|
||||||
|
// (HTTP forward 返回错误路径)记录用量并按上游真实 token 计费,使其与 WS cyber 路径、
|
||||||
|
// 与正常请求的计费口径统一(不再是 tokens=0 免费行)。token 取自上游 response.failed
|
||||||
|
// 报告的 usage(非流式直接拒通常为 0,cost 随之为 0)。复用 RecordUsage 完成成本计算、
|
||||||
|
// 扣费与用量行写入(request_type=cyber 由 CyberBlocked 置位)。仅 forward 返回错误的
|
||||||
|
// 路径由 handler 调用,避免与成功路径的正常 RecordUsage 重复。
|
||||||
|
func (s *OpenAIGatewayService) RecordCyberPolicyUsageLog(ctx context.Context, in CyberPolicyUsageInput) {
|
||||||
|
if s == nil || in.APIKey == nil || in.APIKey.User == nil || in.Account == nil || strings.TrimSpace(in.Model) == "" {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
result := &OpenAIForwardResult{
|
||||||
|
RequestID: in.RequestID,
|
||||||
|
Model: in.Model,
|
||||||
|
Stream: in.Stream,
|
||||||
|
Usage: OpenAIUsage{
|
||||||
|
InputTokens: in.InputTokens,
|
||||||
|
OutputTokens: in.OutputTokens,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
if err := s.RecordUsage(ctx, &OpenAIRecordUsageInput{
|
||||||
|
Result: result,
|
||||||
|
APIKey: in.APIKey,
|
||||||
|
User: in.APIKey.User,
|
||||||
|
Account: in.Account,
|
||||||
|
Subscription: in.Subscription,
|
||||||
|
InboundEndpoint: in.InboundEndpoint,
|
||||||
|
UpstreamEndpoint: in.UpstreamEndpoint,
|
||||||
|
UserAgent: in.UserAgent,
|
||||||
|
IPAddress: in.IPAddress,
|
||||||
|
RequestPayloadHash: in.RequestPayloadHash,
|
||||||
|
APIKeyService: in.APIKeyService,
|
||||||
|
ChannelUsageFields: in.ChannelUsageFields,
|
||||||
|
CyberBlocked: true,
|
||||||
|
}); err != nil {
|
||||||
|
logger.LegacyPrintf("service.openai_gateway", "cyber usage record failed: request_id=%s err=%v", in.RequestID, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// RecordUsage records usage and deducts balance
|
// RecordUsage records usage and deducts balance
|
||||||
func (s *OpenAIGatewayService) RecordUsage(ctx context.Context, input *OpenAIRecordUsageInput) error {
|
func (s *OpenAIGatewayService) RecordUsage(ctx context.Context, input *OpenAIRecordUsageInput) error {
|
||||||
if input == nil {
|
if input == nil {
|
||||||
@@ -5888,6 +6038,9 @@ func (s *OpenAIGatewayService) RecordUsage(ctx context.Context, input *OpenAIRec
|
|||||||
usageLog.AccountRateMultiplier = &accountRateMultiplier
|
usageLog.AccountRateMultiplier = &accountRateMultiplier
|
||||||
usageLog.BillingType = billingType
|
usageLog.BillingType = billingType
|
||||||
usageLog.Stream = result.Stream
|
usageLog.Stream = result.Stream
|
||||||
|
if input.CyberBlocked {
|
||||||
|
usageLog.RequestType = RequestTypeCyberBlocked
|
||||||
|
}
|
||||||
usageLog.OpenAIWSMode = result.OpenAIWSMode
|
usageLog.OpenAIWSMode = result.OpenAIWSMode
|
||||||
usageLog.DurationMs = &durationMs
|
usageLog.DurationMs = &durationMs
|
||||||
usageLog.FirstTokenMs = result.FirstTokenMs
|
usageLog.FirstTokenMs = result.FirstTokenMs
|
||||||
|
|||||||
@@ -2668,3 +2668,109 @@ func TestOpenAICompatSSEFrameParserResetsEventTypeAtFrameBoundary(t *testing.T)
|
|||||||
require.Empty(t, frame.EventType)
|
require.Empty(t, frame.EventType)
|
||||||
require.JSONEq(t, `{"delta":"ok"}`, frame.Data)
|
require.JSONEq(t, `{"delta":"ok"}`, frame.Data)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestStreamingPassthroughCyberPolicyMarksAndPassesThrough(t *testing.T) {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
cfg := &config.Config{Gateway: config.GatewayConfig{MaxLineSize: defaultMaxLineSize}}
|
||||||
|
svc := &OpenAIGatewayService{cfg: cfg}
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(rec)
|
||||||
|
c.Request = httptest.NewRequest(http.MethodPost, "/", nil)
|
||||||
|
|
||||||
|
resp := &http.Response{
|
||||||
|
StatusCode: http.StatusOK,
|
||||||
|
Body: io.NopCloser(strings.NewReader(strings.Join([]string{
|
||||||
|
"event: response.created",
|
||||||
|
`data: {"type":"response.created","response":{"id":"r1"}}`,
|
||||||
|
"",
|
||||||
|
"event: response.failed",
|
||||||
|
`data: {"type":"response.failed","response":{"error":{"code":"cyber_policy","message":"flagged for cyber policy"}}}`,
|
||||||
|
"",
|
||||||
|
}, "\n"))),
|
||||||
|
Header: http.Header{"X-Request-Id": []string{"rid-cyber"}},
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err := svc.handleStreamingResponsePassthrough(c.Request.Context(), resp, c, &Account{ID: 1, Platform: PlatformOpenAI, Name: "a"}, time.Now(), "m", "m")
|
||||||
|
require.Error(t, err)
|
||||||
|
var failoverErr *UpstreamFailoverError
|
||||||
|
require.False(t, errors.As(err, &failoverErr), "cyber must NOT failover")
|
||||||
|
require.Contains(t, rec.Body.String(), "cyber_policy", "response.failed passed through to client")
|
||||||
|
mark := GetOpsCyberPolicy(c)
|
||||||
|
require.NotNil(t, mark)
|
||||||
|
require.Equal(t, "flagged for cyber policy", mark.Message)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleStreamingResponseCyberPolicyMarks(t *testing.T) {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
cfg := &config.Config{Gateway: config.GatewayConfig{MaxLineSize: defaultMaxLineSize}}
|
||||||
|
svc := &OpenAIGatewayService{cfg: cfg}
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(rec)
|
||||||
|
c.Request = httptest.NewRequest(http.MethodPost, "/", nil)
|
||||||
|
resp := &http.Response{
|
||||||
|
StatusCode: http.StatusOK,
|
||||||
|
Body: io.NopCloser(strings.NewReader(strings.Join([]string{
|
||||||
|
"event: response.created",
|
||||||
|
`data: {"type":"response.created","response":{"id":"r1"}}`,
|
||||||
|
"",
|
||||||
|
"event: response.failed",
|
||||||
|
`data: {"type":"response.failed","error":{"code":"cyber_policy","message":"flagged"}}`,
|
||||||
|
"",
|
||||||
|
}, "\n"))),
|
||||||
|
Header: http.Header{"X-Request-Id": []string{"rid"}},
|
||||||
|
}
|
||||||
|
_, err := svc.handleStreamingResponse(c.Request.Context(), resp, c, &Account{ID: 1, Platform: PlatformOpenAI, Name: "a"}, time.Now(), "m", "m")
|
||||||
|
require.Error(t, err)
|
||||||
|
var fo *UpstreamFailoverError
|
||||||
|
require.False(t, errors.As(err, &fo))
|
||||||
|
require.NotNil(t, GetOpsCyberPolicy(c))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleErrorResponseCyberPolicyPassthrough(t *testing.T) {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
svc := &OpenAIGatewayService{cfg: &config.Config{}}
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(rec)
|
||||||
|
c.Request = httptest.NewRequest(http.MethodPost, "/", nil)
|
||||||
|
cyberBody := `{"error":{"code":"cyber_policy","message":"flagged for cyber policy"}}`
|
||||||
|
resp := &http.Response{
|
||||||
|
StatusCode: http.StatusBadRequest,
|
||||||
|
Header: http.Header{"Content-Type": []string{"application/json"}, "X-Request-Id": []string{"rid"}},
|
||||||
|
Body: io.NopCloser(strings.NewReader(cyberBody)),
|
||||||
|
}
|
||||||
|
_, err := svc.handleErrorResponse(context.Background(), resp, c, &Account{ID: 1, Platform: PlatformOpenAI, Name: "a"}, nil)
|
||||||
|
require.Error(t, err)
|
||||||
|
require.Equal(t, http.StatusBadRequest, rec.Code, "passthrough upstream 400, not rewrapped 502")
|
||||||
|
require.Contains(t, rec.Body.String(), "cyber_policy", "client sees original cyber body")
|
||||||
|
require.NotContains(t, rec.Body.String(), "Upstream request failed", "must not 502-rewrap")
|
||||||
|
mark := GetOpsCyberPolicy(c)
|
||||||
|
require.NotNil(t, mark)
|
||||||
|
require.Equal(t, http.StatusBadRequest, mark.UpstreamStatus)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleCompatErrorResponseCyberPolicyEarlyReturn(t *testing.T) {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
svc := &OpenAIGatewayService{cfg: &config.Config{}}
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(rec)
|
||||||
|
c.Request = httptest.NewRequest(http.MethodPost, "/", nil)
|
||||||
|
cyberBody := `{"error":{"code":"cyber_policy","message":"flagged for cyber policy"}}`
|
||||||
|
resp := &http.Response{
|
||||||
|
StatusCode: http.StatusBadRequest,
|
||||||
|
Header: http.Header{"Content-Type": []string{"application/json"}},
|
||||||
|
Body: io.NopCloser(strings.NewReader(cyberBody)),
|
||||||
|
}
|
||||||
|
var gotStatus int
|
||||||
|
var gotType, gotMsg string
|
||||||
|
writeError := func(_ *gin.Context, statusCode int, errType, message string) {
|
||||||
|
gotStatus, gotType, gotMsg = statusCode, errType, message
|
||||||
|
}
|
||||||
|
// cyber 命中应早返回(写兼容错误 + 不冷却账号),而非落到通用 "Upstream request failed"。
|
||||||
|
_, err := svc.handleCompatErrorResponse(resp, c, &Account{ID: 1, Platform: PlatformOpenAI, Name: "a"}, writeError)
|
||||||
|
require.Error(t, err)
|
||||||
|
require.Equal(t, http.StatusBadRequest, gotStatus)
|
||||||
|
require.Equal(t, "invalid_request_error", gotType)
|
||||||
|
require.Contains(t, gotMsg, "flagged for cyber policy")
|
||||||
|
require.NotContains(t, gotMsg, "Upstream request failed")
|
||||||
|
require.NotNil(t, GetOpsCyberPolicy(c))
|
||||||
|
}
|
||||||
|
|||||||
@@ -320,6 +320,34 @@ func TestOpenAIGatewayServiceHandleResponsesImageOutputs_Streaming(t *testing.T)
|
|||||||
require.Equal(t, 4, result.usage.ImageOutputTokens)
|
require.Equal(t, 4, result.usage.ImageOutputTokens)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// TestHandleStreamingResponse_CyberPolicyCapturesRealUpstreamTokens 锁定流式
|
||||||
|
// /v1/responses 命中 cyber_policy 的计费正确性:response.failed 自带的真实 usage
|
||||||
|
// 必须在打 cyber 标记前被解析进 mark;否则计费走 mark.UpstreamInTok 会按 0 token
|
||||||
|
// 漏记真实用量(该路径返回错误,handler 仅经 RecordCyberPolicyUsageLog 计费)。
|
||||||
|
func TestHandleStreamingResponse_CyberPolicyCapturesRealUpstreamTokens(t *testing.T) {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
|
||||||
|
svc := newOpenAIImageGenerationControlTestService(&httpUpstreamRecorder{})
|
||||||
|
c, _ := newOpenAIImageGenerationControlTestContext(false, "unit-test-agent/1.0")
|
||||||
|
resp := &http.Response{
|
||||||
|
StatusCode: http.StatusOK,
|
||||||
|
Header: http.Header{"Content-Type": []string{"text/event-stream"}},
|
||||||
|
Body: io.NopCloser(strings.NewReader(
|
||||||
|
"data: {\"type\":\"response.created\",\"response\":{\"id\":\"resp_cyber\"}}\n\n" +
|
||||||
|
"data: {\"type\":\"response.failed\",\"response\":{\"id\":\"resp_cyber\",\"error\":{\"code\":\"cyber_policy\",\"message\":\"blocked by network policy\"},\"usage\":{\"input_tokens\":1234,\"output_tokens\":7}}}\n\n",
|
||||||
|
)),
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err := svc.handleStreamingResponse(context.Background(), resp, c, &Account{ID: 1}, time.Now(), "gpt-5.5", "gpt-5.5")
|
||||||
|
require.Error(t, err, "cyber 命中的流式响应应返回错误(sawFailedEvent)")
|
||||||
|
|
||||||
|
mark := GetOpsCyberPolicy(c)
|
||||||
|
require.NotNil(t, mark, "必须打上 cyber 标记")
|
||||||
|
require.Equal(t, "cyber_policy", mark.Code)
|
||||||
|
require.Equal(t, 1234, mark.UpstreamInTok, "必须捕获 response.failed 自带真实 input token,而非解析前的 0")
|
||||||
|
require.Equal(t, 7, mark.UpstreamOutTok)
|
||||||
|
}
|
||||||
|
|
||||||
func newOpenAIImageGenerationControlTestService(upstream *httpUpstreamRecorder) *OpenAIGatewayService {
|
func newOpenAIImageGenerationControlTestService(upstream *httpUpstreamRecorder) *OpenAIGatewayService {
|
||||||
cfg := &config.Config{}
|
cfg := &config.Config{}
|
||||||
return &OpenAIGatewayService{
|
return &OpenAIGatewayService{
|
||||||
|
|||||||
@@ -2237,6 +2237,19 @@ func (s *OpenAIGatewayService) forwardOpenAIWSV2(
|
|||||||
}
|
}
|
||||||
imageCounter.AddSSEData(message)
|
imageCounter.AddSSEData(message)
|
||||||
|
|
||||||
|
if eventType == "response.failed" {
|
||||||
|
if hit, code, msg := detectOpenAICyberPolicy(message); hit {
|
||||||
|
MarkOpsCyberPolicy(c, CyberPolicyMark{
|
||||||
|
Code: code,
|
||||||
|
Message: msg,
|
||||||
|
Body: truncateString(string(message), 4096),
|
||||||
|
UpstreamStatus: http.StatusOK,
|
||||||
|
UpstreamInTok: usage.InputTokens,
|
||||||
|
UpstreamOutTok: usage.OutputTokens,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
if eventType == "error" {
|
if eventType == "error" {
|
||||||
errCodeRaw, errTypeRaw, errMsgRaw := parseOpenAIWSErrorEventFields(message)
|
errCodeRaw, errTypeRaw, errMsgRaw := parseOpenAIWSErrorEventFields(message)
|
||||||
s.persistOpenAIWSRateLimitSignal(ctx, account, lease.HandshakeHeaders(), message, errCodeRaw, errTypeRaw, errMsgRaw)
|
s.persistOpenAIWSRateLimitSignal(ctx, account, lease.HandshakeHeaders(), message, errCodeRaw, errTypeRaw, errMsgRaw)
|
||||||
@@ -3210,6 +3223,19 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
|
|||||||
}
|
}
|
||||||
imageCounter.AddSSEData(upstreamMessage)
|
imageCounter.AddSSEData(upstreamMessage)
|
||||||
|
|
||||||
|
if eventType == "response.failed" {
|
||||||
|
if hit, code, msg := detectOpenAICyberPolicy(upstreamMessage); hit {
|
||||||
|
MarkOpsCyberPolicy(c, CyberPolicyMark{
|
||||||
|
Code: code,
|
||||||
|
Message: msg,
|
||||||
|
Body: truncateString(string(upstreamMessage), 4096),
|
||||||
|
UpstreamStatus: http.StatusOK,
|
||||||
|
UpstreamInTok: usage.InputTokens,
|
||||||
|
UpstreamOutTok: usage.OutputTokens,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
if !clientDisconnected {
|
if !clientDisconnected {
|
||||||
if needModelReplace && len(mappedModelBytes) > 0 && openAIWSEventMayContainModel(eventType) && bytes.Contains(upstreamMessage, mappedModelBytes) {
|
if needModelReplace && len(mappedModelBytes) > 0 && openAIWSEventMayContainModel(eventType) && bytes.Contains(upstreamMessage, mappedModelBytes) {
|
||||||
upstreamMessage = replaceOpenAIWSMessageModel(upstreamMessage, mappedModel, originalModel)
|
upstreamMessage = replaceOpenAIWSMessageModel(upstreamMessage, mappedModel, originalModel)
|
||||||
|
|||||||
@@ -1,8 +1,10 @@
|
|||||||
package service
|
package service
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"net/http/httptest"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -53,6 +55,58 @@ func TestIsOpenAIWSTokenEvent_TerminalEventsExcluded(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// TestOpenAIWSCyberPolicyMark_ResponseFailed 验证 WS 路径 response.failed cyber_policy 标记逻辑。
|
||||||
|
//
|
||||||
|
// 全量转发循环(forwardOpenAIWSV2 / sendAndRelay)依赖真实 WebSocket 连接,
|
||||||
|
// 无法在单元测试中驱动。本测试通过直接调用转发循环内使用的两个函数
|
||||||
|
// detectOpenAICyberPolicy + MarkOpsCyberPolicy,覆盖「从 response.failed 帧
|
||||||
|
// 到 gin context 写入」的完整调用序列,等同于循环体内对应代码段的逻辑验证。
|
||||||
|
// 全量 WS 端到端覆盖由后续集成测试(Task 12 handler 编排)承担。
|
||||||
|
func TestOpenAIWSCyberPolicyMark_ResponseFailed(t *testing.T) {
|
||||||
|
// 构造一个真实的 response.failed 帧(cyber_policy 命中路径)。
|
||||||
|
payload := []byte(`{"type":"response.failed","response":{"id":"resp_abc","status":"failed","error":{"code":"cyber_policy","message":"Request blocked by content policy."}}}`)
|
||||||
|
|
||||||
|
// 验证 detectOpenAICyberPolicy 能从 response.error.code 路径识别。
|
||||||
|
hit, code, msg := detectOpenAICyberPolicy(payload)
|
||||||
|
require.True(t, hit, "detectOpenAICyberPolicy should return true for cyber_policy payload")
|
||||||
|
require.Equal(t, "cyber_policy", code)
|
||||||
|
require.Equal(t, "Request blocked by content policy.", msg)
|
||||||
|
|
||||||
|
// 构造 gin test context,模拟转发循环调用 MarkOpsCyberPolicy。
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(rec)
|
||||||
|
|
||||||
|
usage := OpenAIUsage{InputTokens: 42, OutputTokens: 7}
|
||||||
|
MarkOpsCyberPolicy(c, CyberPolicyMark{
|
||||||
|
Code: code,
|
||||||
|
Message: msg,
|
||||||
|
Body: truncateString(string(payload), 4096),
|
||||||
|
UpstreamStatus: 200,
|
||||||
|
UpstreamInTok: usage.InputTokens,
|
||||||
|
UpstreamOutTok: usage.OutputTokens,
|
||||||
|
})
|
||||||
|
|
||||||
|
mark := GetOpsCyberPolicy(c)
|
||||||
|
require.NotNil(t, mark, "GetOpsCyberPolicy should return non-nil after MarkOpsCyberPolicy")
|
||||||
|
require.Equal(t, "cyber_policy", mark.Code)
|
||||||
|
require.Equal(t, "Request blocked by content policy.", mark.Message)
|
||||||
|
require.Equal(t, 200, mark.UpstreamStatus)
|
||||||
|
require.Equal(t, 42, mark.UpstreamInTok)
|
||||||
|
require.Equal(t, 7, mark.UpstreamOutTok)
|
||||||
|
|
||||||
|
// 验证幂等性:再次标记不覆盖首个。
|
||||||
|
MarkOpsCyberPolicy(c, CyberPolicyMark{Code: "cyber_policy", Message: "second call"})
|
||||||
|
require.Equal(t, "Request blocked by content policy.", GetOpsCyberPolicy(c).Message, "second MarkOpsCyberPolicy call must not overwrite first")
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestOpenAIWSCyberPolicyMark_NonCyberPayload 验证非 cyber_policy 的 response.failed 不触发标记。
|
||||||
|
func TestOpenAIWSCyberPolicyMark_NonCyberPayload(t *testing.T) {
|
||||||
|
payload := []byte(`{"type":"response.failed","response":{"id":"resp_xyz","status":"failed","error":{"code":"server_error","message":"Internal error"}}}`)
|
||||||
|
|
||||||
|
hit, _, _ := detectOpenAICyberPolicy(payload)
|
||||||
|
require.False(t, hit, "detectOpenAICyberPolicy should return false for non-cyber_policy error code")
|
||||||
|
}
|
||||||
|
|
||||||
// TestIsOpenAIWSTokenEvent_DisjointWithTerminal 守护「token 事件集合与终止事件集合互斥」的不变量。
|
// TestIsOpenAIWSTokenEvent_DisjointWithTerminal 守护「token 事件集合与终止事件集合互斥」的不变量。
|
||||||
// firstTokenMs 的计算依赖于 isTokenEvent && !isTerminalEvent;
|
// firstTokenMs 的计算依赖于 isTokenEvent && !isTerminalEvent;
|
||||||
// 若两者再次出现交集,则 issue #2651 描述的 latency 误报会重现。
|
// 若两者再次出现交集,则 issue #2651 描述的 latency 误报会重现。
|
||||||
|
|||||||
@@ -48,6 +48,8 @@ func MapUserErrorCategory(phase, errType string) string {
|
|||||||
return "quota"
|
return "quota"
|
||||||
case "invalid_request_error":
|
case "invalid_request_error":
|
||||||
return "invalid_request"
|
return "invalid_request"
|
||||||
|
case "cyber_policy":
|
||||||
|
return "cyber"
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return "other"
|
return "other"
|
||||||
@@ -72,6 +74,8 @@ func CategoryToFilter(category string) (phases []string, errorTypes []string) {
|
|||||||
return nil, []string{"billing_error", "subscription_error"}
|
return nil, []string{"billing_error", "subscription_error"}
|
||||||
case "invalid_request":
|
case "invalid_request":
|
||||||
return nil, []string{"invalid_request_error"}
|
return nil, []string{"invalid_request_error"}
|
||||||
|
case "cyber":
|
||||||
|
return []string{"request"}, []string{"cyber_policy"}
|
||||||
default:
|
default:
|
||||||
return nil, nil
|
return nil, nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,14 @@
|
|||||||
|
package service
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestMapUserErrorCategoryCyber(t *testing.T) {
|
||||||
|
require.Equal(t, "cyber", MapUserErrorCategory("request", "cyber_policy"))
|
||||||
|
phases, types := CategoryToFilter("cyber")
|
||||||
|
require.Equal(t, []string{"request"}, phases)
|
||||||
|
require.Equal(t, []string{"cyber_policy"}, types)
|
||||||
|
}
|
||||||
@@ -157,6 +157,18 @@ const openAIAllowCodexPluginCacheTTL = 60 * time.Second
|
|||||||
const openAIAllowCodexPluginErrorTTL = 5 * time.Second
|
const openAIAllowCodexPluginErrorTTL = 5 * time.Second
|
||||||
const openAIAllowCodexPluginDBTimeout = 5 * time.Second
|
const openAIAllowCodexPluginDBTimeout = 5 * time.Second
|
||||||
|
|
||||||
|
// cachedCyberSessionBlockRuntime cyber 会话屏蔽开关+TTL 进程内缓存(60s TTL)。
|
||||||
|
// GetCyberSessionBlockRuntime 在网关请求热路径上被调用,避免每次访问 DB。
|
||||||
|
type cachedCyberSessionBlockRuntime struct {
|
||||||
|
enabled bool
|
||||||
|
ttl time.Duration
|
||||||
|
expiresAt int64 // unix nano
|
||||||
|
}
|
||||||
|
|
||||||
|
const cyberSessionBlockRuntimeCacheTTL = 60 * time.Second
|
||||||
|
const cyberSessionBlockRuntimeErrorTTL = 5 * time.Second
|
||||||
|
const cyberSessionBlockRuntimeDBTimeout = 5 * time.Second
|
||||||
|
|
||||||
const openAIQuotaAutoPauseSettingsCacheTTL = 60 * time.Second
|
const openAIQuotaAutoPauseSettingsCacheTTL = 60 * time.Second
|
||||||
const openAIQuotaAutoPauseSettingsErrorTTL = 5 * time.Second
|
const openAIQuotaAutoPauseSettingsErrorTTL = 5 * time.Second
|
||||||
const openAIQuotaAutoPauseSettingsDBTimeout = 5 * time.Second
|
const openAIQuotaAutoPauseSettingsDBTimeout = 5 * time.Second
|
||||||
@@ -188,6 +200,9 @@ type SettingService struct {
|
|||||||
openAIAllowCodexPluginCache atomic.Value // *cachedOpenAIAllowCodexPlugin
|
openAIAllowCodexPluginCache atomic.Value // *cachedOpenAIAllowCodexPlugin
|
||||||
openAIAllowCodexPluginSF singleflight.Group
|
openAIAllowCodexPluginSF singleflight.Group
|
||||||
|
|
||||||
|
cyberSessionBlockRuntimeCache atomic.Value // *cachedCyberSessionBlockRuntime
|
||||||
|
cyberSessionBlockRuntimeSF singleflight.Group
|
||||||
|
|
||||||
// openAIQuotaAutoPauseSettingsCache holds the most recently observed quota auto-pause
|
// openAIQuotaAutoPauseSettingsCache holds the most recently observed quota auto-pause
|
||||||
// settings. GetOpenAIQuotaAutoPauseSettings reads this atomic.Value on the request hot
|
// settings. GetOpenAIQuotaAutoPauseSettings reads this atomic.Value on the request hot
|
||||||
// path without ever blocking on the DB; when the cached entry expires, a background
|
// path without ever blocking on the DB; when the cached entry expires, a background
|
||||||
@@ -692,6 +707,62 @@ func (s *SettingService) GetFrontendURL(ctx context.Context) string {
|
|||||||
return s.cfg.Server.FrontendURL
|
return s.cfg.Server.FrontendURL
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// GetCyberSessionBlockRuntime 返回 (开关, TTL),进程内缓存 ~60s,
|
||||||
|
// 模式对齐 IsOpenAIAllowClaudeCodeCodexPluginEnabled(热路径零 DB 往返)。
|
||||||
|
// 两个 setting key 在单次 singleflight 里一起读取,减少 DB 往返。
|
||||||
|
// 默认值:开关 false,TTL 1h(与粘性会话对齐)。
|
||||||
|
func (s *SettingService) GetCyberSessionBlockRuntime(ctx context.Context) (bool, time.Duration) {
|
||||||
|
if cached, ok := s.cyberSessionBlockRuntimeCache.Load().(*cachedCyberSessionBlockRuntime); ok && cached != nil {
|
||||||
|
if time.Now().UnixNano() < cached.expiresAt {
|
||||||
|
return cached.enabled, cached.ttl
|
||||||
|
}
|
||||||
|
}
|
||||||
|
result, _, _ := s.cyberSessionBlockRuntimeSF.Do("cyber_session_block_runtime", func() (any, error) {
|
||||||
|
if cached, ok := s.cyberSessionBlockRuntimeCache.Load().(*cachedCyberSessionBlockRuntime); ok && cached != nil {
|
||||||
|
if time.Now().UnixNano() < cached.expiresAt {
|
||||||
|
return cached, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
dbCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), cyberSessionBlockRuntimeDBTimeout)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
enabledVal, enabledErr := s.settingRepo.GetValue(dbCtx, SettingKeyCyberSessionBlockEnabled)
|
||||||
|
ttlVal, ttlErr := s.settingRepo.GetValue(dbCtx, SettingKeyCyberSessionBlockTTLSeconds)
|
||||||
|
|
||||||
|
if enabledErr != nil && !errors.Is(enabledErr, ErrSettingNotFound) {
|
||||||
|
slog.Warn("failed to get cyber_session_block_enabled setting", "error", enabledErr)
|
||||||
|
entry := &cachedCyberSessionBlockRuntime{
|
||||||
|
enabled: false,
|
||||||
|
ttl: time.Hour,
|
||||||
|
expiresAt: time.Now().Add(cyberSessionBlockRuntimeErrorTTL).UnixNano(),
|
||||||
|
}
|
||||||
|
s.cyberSessionBlockRuntimeCache.Store(entry)
|
||||||
|
return entry, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
enabled := enabledErr == nil && strings.TrimSpace(enabledVal) == "true"
|
||||||
|
|
||||||
|
ttl := time.Hour
|
||||||
|
if ttlErr == nil {
|
||||||
|
if n, perr := strconv.Atoi(strings.TrimSpace(ttlVal)); perr == nil && n > 0 {
|
||||||
|
ttl = time.Duration(n) * time.Second
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
entry := &cachedCyberSessionBlockRuntime{
|
||||||
|
enabled: enabled,
|
||||||
|
ttl: ttl,
|
||||||
|
expiresAt: time.Now().Add(cyberSessionBlockRuntimeCacheTTL).UnixNano(),
|
||||||
|
}
|
||||||
|
s.cyberSessionBlockRuntimeCache.Store(entry)
|
||||||
|
return entry, nil
|
||||||
|
})
|
||||||
|
if entry, ok := result.(*cachedCyberSessionBlockRuntime); ok && entry != nil {
|
||||||
|
return entry.enabled, entry.ttl
|
||||||
|
}
|
||||||
|
return false, time.Hour
|
||||||
|
}
|
||||||
|
|
||||||
// GetPublicSettings 获取公开设置(无需登录)
|
// GetPublicSettings 获取公开设置(无需登录)
|
||||||
func (s *SettingService) GetPublicSettings(ctx context.Context) (*PublicSettings, error) {
|
func (s *SettingService) GetPublicSettings(ctx context.Context) (*PublicSettings, error) {
|
||||||
keys := []string{
|
keys := []string{
|
||||||
@@ -1895,6 +1966,12 @@ func (s *SettingService) buildSystemSettingsUpdates(ctx context.Context, setting
|
|||||||
// 风控中心功能开关
|
// 风控中心功能开关
|
||||||
updates[SettingKeyRiskControlEnabled] = strconv.FormatBool(settings.RiskControlEnabled)
|
updates[SettingKeyRiskControlEnabled] = strconv.FormatBool(settings.RiskControlEnabled)
|
||||||
|
|
||||||
|
// cyber 会话屏蔽开关 + TTL
|
||||||
|
updates[SettingKeyCyberSessionBlockEnabled] = strconv.FormatBool(settings.CyberSessionBlockEnabled)
|
||||||
|
if settings.CyberSessionBlockTTLSeconds > 0 {
|
||||||
|
updates[SettingKeyCyberSessionBlockTTLSeconds] = strconv.Itoa(settings.CyberSessionBlockTTLSeconds)
|
||||||
|
}
|
||||||
|
|
||||||
// Claude Code version check
|
// Claude Code version check
|
||||||
updates[SettingKeyMinClaudeCodeVersion] = settings.MinClaudeCodeVersion
|
updates[SettingKeyMinClaudeCodeVersion] = settings.MinClaudeCodeVersion
|
||||||
updates[SettingKeyMaxClaudeCodeVersion] = settings.MaxClaudeCodeVersion
|
updates[SettingKeyMaxClaudeCodeVersion] = settings.MaxClaudeCodeVersion
|
||||||
@@ -2819,6 +2896,10 @@ func (s *SettingService) InitializeDefaultSettings(ctx context.Context) error {
|
|||||||
// 风控中心功能(默认关闭,显式启用)
|
// 风控中心功能(默认关闭,显式启用)
|
||||||
SettingKeyRiskControlEnabled: "false",
|
SettingKeyRiskControlEnabled: "false",
|
||||||
|
|
||||||
|
// cyber 会话屏蔽(默认关闭,TTL 默认 3600s)
|
||||||
|
SettingKeyCyberSessionBlockEnabled: "false",
|
||||||
|
SettingKeyCyberSessionBlockTTLSeconds: "3600",
|
||||||
|
|
||||||
// Claude Code version check (default: empty = disabled)
|
// Claude Code version check (default: empty = disabled)
|
||||||
SettingKeyMinClaudeCodeVersion: "",
|
SettingKeyMinClaudeCodeVersion: "",
|
||||||
SettingKeyMaxClaudeCodeVersion: "",
|
SettingKeyMaxClaudeCodeVersion: "",
|
||||||
@@ -3328,6 +3409,14 @@ func (s *SettingService) parseSettings(settings map[string]string) *SystemSettin
|
|||||||
// 风控中心功能(默认关闭,严格 true 才启用)
|
// 风控中心功能(默认关闭,严格 true 才启用)
|
||||||
result.RiskControlEnabled = settings[SettingKeyRiskControlEnabled] == "true"
|
result.RiskControlEnabled = settings[SettingKeyRiskControlEnabled] == "true"
|
||||||
|
|
||||||
|
// cyber 会话屏蔽(默认关闭,TTL 默认 3600s)
|
||||||
|
result.CyberSessionBlockEnabled = settings[SettingKeyCyberSessionBlockEnabled] == "true"
|
||||||
|
if v, err := strconv.Atoi(strings.TrimSpace(settings[SettingKeyCyberSessionBlockTTLSeconds])); err == nil && v > 0 {
|
||||||
|
result.CyberSessionBlockTTLSeconds = v
|
||||||
|
} else {
|
||||||
|
result.CyberSessionBlockTTLSeconds = 3600
|
||||||
|
}
|
||||||
|
|
||||||
// Claude Code version check
|
// Claude Code version check
|
||||||
result.MinClaudeCodeVersion = settings[SettingKeyMinClaudeCodeVersion]
|
result.MinClaudeCodeVersion = settings[SettingKeyMinClaudeCodeVersion]
|
||||||
result.MaxClaudeCodeVersion = settings[SettingKeyMaxClaudeCodeVersion]
|
result.MaxClaudeCodeVersion = settings[SettingKeyMaxClaudeCodeVersion]
|
||||||
|
|||||||
@@ -145,6 +145,8 @@ type SystemSettings struct {
|
|||||||
DefaultConcurrency int
|
DefaultConcurrency int
|
||||||
DefaultBalance float64
|
DefaultBalance float64
|
||||||
RiskControlEnabled bool
|
RiskControlEnabled bool
|
||||||
|
CyberSessionBlockEnabled bool
|
||||||
|
CyberSessionBlockTTLSeconds int
|
||||||
AffiliateEnabled bool
|
AffiliateEnabled bool
|
||||||
AffiliateRebateRate float64
|
AffiliateRebateRate float64
|
||||||
AffiliateRebateFreezeHours int
|
AffiliateRebateFreezeHours int
|
||||||
|
|||||||
@@ -18,11 +18,12 @@ const (
|
|||||||
RequestTypeSync RequestType = 1
|
RequestTypeSync RequestType = 1
|
||||||
RequestTypeStream RequestType = 2
|
RequestTypeStream RequestType = 2
|
||||||
RequestTypeWSV2 RequestType = 3
|
RequestTypeWSV2 RequestType = 3
|
||||||
|
RequestTypeCyberBlocked RequestType = 4 // cyber_policy 命中(透传但被上游安全策略拒绝)
|
||||||
)
|
)
|
||||||
|
|
||||||
func (t RequestType) IsValid() bool {
|
func (t RequestType) IsValid() bool {
|
||||||
switch t {
|
switch t {
|
||||||
case RequestTypeUnknown, RequestTypeSync, RequestTypeStream, RequestTypeWSV2:
|
case RequestTypeUnknown, RequestTypeSync, RequestTypeStream, RequestTypeWSV2, RequestTypeCyberBlocked:
|
||||||
return true
|
return true
|
||||||
default:
|
default:
|
||||||
return false
|
return false
|
||||||
@@ -44,6 +45,8 @@ func (t RequestType) String() string {
|
|||||||
return "stream"
|
return "stream"
|
||||||
case RequestTypeWSV2:
|
case RequestTypeWSV2:
|
||||||
return "ws_v2"
|
return "ws_v2"
|
||||||
|
case RequestTypeCyberBlocked:
|
||||||
|
return "cyber"
|
||||||
default:
|
default:
|
||||||
return "unknown"
|
return "unknown"
|
||||||
}
|
}
|
||||||
@@ -63,8 +66,10 @@ func ParseUsageRequestType(value string) (RequestType, error) {
|
|||||||
return RequestTypeStream, nil
|
return RequestTypeStream, nil
|
||||||
case "ws_v2":
|
case "ws_v2":
|
||||||
return RequestTypeWSV2, nil
|
return RequestTypeWSV2, nil
|
||||||
|
case "cyber":
|
||||||
|
return RequestTypeCyberBlocked, nil
|
||||||
default:
|
default:
|
||||||
return RequestTypeUnknown, fmt.Errorf("invalid request_type, allowed values: unknown, sync, stream, ws_v2")
|
return RequestTypeUnknown, fmt.Errorf("invalid request_type, allowed values: unknown, sync, stream, ws_v2, cyber")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,25 @@
|
|||||||
|
package service
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestRequestTypeCyberBlocked(t *testing.T) {
|
||||||
|
require.True(t, RequestTypeCyberBlocked.IsValid())
|
||||||
|
require.Equal(t, "cyber", RequestTypeCyberBlocked.String())
|
||||||
|
|
||||||
|
rt, err := ParseUsageRequestType("cyber")
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Equal(t, RequestTypeCyberBlocked, rt)
|
||||||
|
|
||||||
|
// 显式 cyber 被 EffectiveRequestType 保留(不被 legacy 推导覆盖)
|
||||||
|
u := &UsageLog{RequestType: RequestTypeCyberBlocked, Stream: true}
|
||||||
|
require.Equal(t, RequestTypeCyberBlocked, u.EffectiveRequestType())
|
||||||
|
|
||||||
|
// Sync 保留 cyber 且不覆盖真实 stream
|
||||||
|
u.SyncRequestTypeAndLegacyFields()
|
||||||
|
require.Equal(t, RequestTypeCyberBlocked, u.RequestType)
|
||||||
|
require.True(t, u.Stream, "cyber 不应覆盖真实 stream 字段")
|
||||||
|
}
|
||||||
@@ -40,6 +40,7 @@ export interface ContentModerationConfig {
|
|||||||
blocked_keywords: string[]
|
blocked_keywords: string[]
|
||||||
keyword_blocking_mode: KeywordBlockingMode
|
keyword_blocking_mode: KeywordBlockingMode
|
||||||
model_filter: ContentModerationModelFilter
|
model_filter: ContentModerationModelFilter
|
||||||
|
cyber_policy_exclude_from_ban_count: boolean
|
||||||
}
|
}
|
||||||
|
|
||||||
export type ContentModerationAPIKeyStatusValue = 'unknown' | 'ok' | 'error' | 'frozen'
|
export type ContentModerationAPIKeyStatusValue = 'unknown' | 'ok' | 'error' | 'frozen'
|
||||||
@@ -115,6 +116,7 @@ export interface UpdateContentModerationConfig {
|
|||||||
blocked_keywords?: string[]
|
blocked_keywords?: string[]
|
||||||
keyword_blocking_mode?: KeywordBlockingMode
|
keyword_blocking_mode?: KeywordBlockingMode
|
||||||
model_filter?: ContentModerationModelFilter
|
model_filter?: ContentModerationModelFilter
|
||||||
|
cyber_policy_exclude_from_ban_count?: boolean
|
||||||
}
|
}
|
||||||
|
|
||||||
export interface ContentModerationRuntimeStatus {
|
export interface ContentModerationRuntimeStatus {
|
||||||
|
|||||||
@@ -566,6 +566,11 @@ export interface SystemSettings {
|
|||||||
// Payment configuration
|
// Payment configuration
|
||||||
payment_enabled: boolean;
|
payment_enabled: boolean;
|
||||||
risk_control_enabled: boolean;
|
risk_control_enabled: boolean;
|
||||||
|
|
||||||
|
// Cyber session block
|
||||||
|
cyber_session_block_enabled: boolean;
|
||||||
|
cyber_session_block_ttl_seconds: number;
|
||||||
|
|
||||||
payment_min_amount: number;
|
payment_min_amount: number;
|
||||||
payment_max_amount: number;
|
payment_max_amount: number;
|
||||||
payment_daily_limit: number;
|
payment_daily_limit: number;
|
||||||
@@ -800,6 +805,11 @@ export interface UpdateSettingsRequest {
|
|||||||
// Payment configuration
|
// Payment configuration
|
||||||
payment_enabled?: boolean;
|
payment_enabled?: boolean;
|
||||||
risk_control_enabled?: boolean;
|
risk_control_enabled?: boolean;
|
||||||
|
|
||||||
|
// Cyber session block
|
||||||
|
cyber_session_block_enabled?: boolean;
|
||||||
|
cyber_session_block_ttl_seconds?: number;
|
||||||
|
|
||||||
payment_min_amount?: number;
|
payment_min_amount?: number;
|
||||||
payment_max_amount?: number;
|
payment_max_amount?: number;
|
||||||
payment_daily_limit?: number;
|
payment_daily_limit?: number;
|
||||||
|
|||||||
@@ -233,7 +233,8 @@ const requestTypeOptions = ref<SelectOption[]>([
|
|||||||
{ value: null, label: t('admin.usage.allTypes') },
|
{ value: null, label: t('admin.usage.allTypes') },
|
||||||
{ value: 'ws_v2', label: t('usage.ws') },
|
{ value: 'ws_v2', label: t('usage.ws') },
|
||||||
{ value: 'stream', label: t('usage.stream') },
|
{ value: 'stream', label: t('usage.stream') },
|
||||||
{ value: 'sync', label: t('usage.sync') }
|
{ value: 'sync', label: t('usage.sync') },
|
||||||
|
{ value: 'cyber', label: t('usage.cyber') }
|
||||||
])
|
])
|
||||||
|
|
||||||
const billingTypeOptions = ref<SelectOption[]>([
|
const billingTypeOptions = ref<SelectOption[]>([
|
||||||
|
|||||||
@@ -475,6 +475,7 @@ const tokenTooltipData = ref<AdminUsageLog | null>(null)
|
|||||||
|
|
||||||
const getRequestTypeLabel = (row: AdminUsageLog): string => {
|
const getRequestTypeLabel = (row: AdminUsageLog): string => {
|
||||||
const requestType = resolveUsageRequestType(row)
|
const requestType = resolveUsageRequestType(row)
|
||||||
|
if (requestType === 'cyber') return t('usage.cyber')
|
||||||
if (requestType === 'ws_v2') return t('usage.ws')
|
if (requestType === 'ws_v2') return t('usage.ws')
|
||||||
if (requestType === 'stream') return t('usage.stream')
|
if (requestType === 'stream') return t('usage.stream')
|
||||||
if (requestType === 'sync') return t('usage.sync')
|
if (requestType === 'sync') return t('usage.sync')
|
||||||
@@ -483,6 +484,7 @@ const getRequestTypeLabel = (row: AdminUsageLog): string => {
|
|||||||
|
|
||||||
const getRequestTypeBadgeClass = (row: AdminUsageLog): string => {
|
const getRequestTypeBadgeClass = (row: AdminUsageLog): string => {
|
||||||
const requestType = resolveUsageRequestType(row)
|
const requestType = resolveUsageRequestType(row)
|
||||||
|
if (requestType === 'cyber') return 'bg-red-100 text-red-800 dark:bg-red-900 dark:text-red-200'
|
||||||
if (requestType === 'ws_v2') return 'bg-violet-100 text-violet-800 dark:bg-violet-900 dark:text-violet-200'
|
if (requestType === 'ws_v2') return 'bg-violet-100 text-violet-800 dark:bg-violet-900 dark:text-violet-200'
|
||||||
if (requestType === 'stream') return 'bg-blue-100 text-blue-800 dark:bg-blue-900 dark:text-blue-200'
|
if (requestType === 'stream') return 'bg-blue-100 text-blue-800 dark:bg-blue-900 dark:text-blue-200'
|
||||||
if (requestType === 'sync') return 'bg-gray-100 text-gray-800 dark:bg-gray-700 dark:text-gray-200'
|
if (requestType === 'sync') return 'bg-gray-100 text-gray-800 dark:bg-gray-700 dark:text-gray-200'
|
||||||
|
|||||||
@@ -123,7 +123,7 @@ const localModel = ref<string | null>('')
|
|||||||
const localCategory = ref<string>('')
|
const localCategory = ref<string>('')
|
||||||
const localApiKeyId = ref<number | null>(null)
|
const localApiKeyId = ref<number | null>(null)
|
||||||
|
|
||||||
const categoryCodes = ['auth', 'rate_limit', 'quota', 'invalid_request', 'service_unavailable', 'upstream', 'internal']
|
const categoryCodes = ['auth', 'rate_limit', 'quota', 'invalid_request', 'service_unavailable', 'upstream', 'internal', 'cyber']
|
||||||
|
|
||||||
const categoryOptions = computed(() => [
|
const categoryOptions = computed(() => [
|
||||||
{ value: '', label: t('usage.errors.allCategories') },
|
{ value: '', label: t('usage.errors.allCategories') },
|
||||||
|
|||||||
@@ -923,6 +923,7 @@ export default {
|
|||||||
ws: 'WS',
|
ws: 'WS',
|
||||||
stream: 'Stream',
|
stream: 'Stream',
|
||||||
sync: 'Sync',
|
sync: 'Sync',
|
||||||
|
cyber: 'Cyber',
|
||||||
unknown: 'Unknown',
|
unknown: 'Unknown',
|
||||||
in: 'In',
|
in: 'In',
|
||||||
out: 'Out',
|
out: 'Out',
|
||||||
@@ -980,7 +981,7 @@ export default {
|
|||||||
categories: {
|
categories: {
|
||||||
auth: 'Auth failed', rate_limit: 'Rate limited', quota: 'Balance/Subscription',
|
auth: 'Auth failed', rate_limit: 'Rate limited', quota: 'Balance/Subscription',
|
||||||
invalid_request: 'Invalid request', service_unavailable: 'Service unavailable',
|
invalid_request: 'Invalid request', service_unavailable: 'Service unavailable',
|
||||||
upstream: 'Upstream error', internal: 'Platform error', other: 'Other',
|
upstream: 'Upstream error', internal: 'Platform error', other: 'Other', cyber: 'Cyber policy',
|
||||||
},
|
},
|
||||||
detail: {
|
detail: {
|
||||||
title: 'Error Request Detail',
|
title: 'Error Request Detail',
|
||||||
@@ -2627,6 +2628,9 @@ export default {
|
|||||||
emailOnHitHint: 'When enabled, send a risk-control email on every hit; auto-ban notices are always sent.',
|
emailOnHitHint: 'When enabled, send a risk-control email on every hit; auto-ban notices are always sent.',
|
||||||
autoBan: 'Auto Ban User',
|
autoBan: 'Auto Ban User',
|
||||||
autoBanHint: 'Disable the user, invalidate auth cache, and send a ban notice after the hit threshold is reached.',
|
autoBanHint: 'Disable the user, invalidate auth cache, and send a ban notice after the hit threshold is reached.',
|
||||||
|
cyberPolicyExcludeBan: 'Exclude Cyber Policy Hits from Ban Count',
|
||||||
|
cyberPolicyExcludeBanHint: 'When enabled, cyber_policy hits no longer count toward auto-ban violations: no ban judgment on the hit itself, and history rows are excluded from the rolling count. Logs and notice emails are unaffected.',
|
||||||
|
violationNotCounted: 'Not counted',
|
||||||
banThreshold: 'Ban Threshold',
|
banThreshold: 'Ban Threshold',
|
||||||
violationWindowHours: 'Count Window (hours)',
|
violationWindowHours: 'Count Window (hours)',
|
||||||
hitRetentionDays: 'Hit Record Retention (days)',
|
hitRetentionDays: 'Hit Record Retention (days)',
|
||||||
@@ -2771,6 +2775,7 @@ export default {
|
|||||||
action: {
|
action: {
|
||||||
block: 'Blocked',
|
block: 'Blocked',
|
||||||
keywordBlock: 'Keyword Blocked',
|
keywordBlock: 'Keyword Blocked',
|
||||||
|
cyberPolicy: 'Cyber policy',
|
||||||
error: 'Error',
|
error: 'Error',
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
@@ -5438,6 +5443,9 @@ export default {
|
|||||||
configureLink: 'Configure content moderation in Risk Control',
|
configureLink: 'Configure content moderation in Risk Control',
|
||||||
enabled: 'Enable Risk Control',
|
enabled: 'Enable Risk Control',
|
||||||
enabledHint: 'When off, the admin sidebar entry is hidden and gateway moderation is skipped.',
|
enabledHint: 'When off, the admin sidebar entry is hidden and gateway moderation is skipped.',
|
||||||
|
cyberSessionBlock: 'Cyber session auto-block',
|
||||||
|
cyberSessionBlockHint: 'When enabled, sessions hit by upstream cyber_policy are blocked locally for the TTL and no longer forwarded. Only the offending session is blocked; other sessions on the same key are unaffected.',
|
||||||
|
cyberSessionBlockTTL: 'Block TTL (seconds)',
|
||||||
},
|
},
|
||||||
affiliate: {
|
affiliate: {
|
||||||
title: 'Affiliate (Invite Rebate)',
|
title: 'Affiliate (Invite Rebate)',
|
||||||
|
|||||||
@@ -927,6 +927,7 @@ export default {
|
|||||||
ws: 'WS',
|
ws: 'WS',
|
||||||
stream: '流式',
|
stream: '流式',
|
||||||
sync: '同步',
|
sync: '同步',
|
||||||
|
cyber: '安全策略',
|
||||||
unknown: '未知',
|
unknown: '未知',
|
||||||
in: '输入',
|
in: '输入',
|
||||||
out: '输出',
|
out: '输出',
|
||||||
@@ -984,7 +985,7 @@ export default {
|
|||||||
categories: {
|
categories: {
|
||||||
auth: '认证失败', rate_limit: '限流', quota: '余额/订阅',
|
auth: '认证失败', rate_limit: '限流', quota: '余额/订阅',
|
||||||
invalid_request: '参数错误', service_unavailable: '服务暂时不可用',
|
invalid_request: '参数错误', service_unavailable: '服务暂时不可用',
|
||||||
upstream: '上游错误', internal: '平台错误', other: '其他',
|
upstream: '上游错误', internal: '平台错误', other: '其他', cyber: '安全策略',
|
||||||
},
|
},
|
||||||
detail: {
|
detail: {
|
||||||
title: '错误请求详情',
|
title: '错误请求详情',
|
||||||
@@ -2704,6 +2705,9 @@ export default {
|
|||||||
emailOnHitHint: '开启后每次达到阈值都会向用户发送风控提醒邮件;自动封禁通知始终发送。',
|
emailOnHitHint: '开启后每次达到阈值都会向用户发送风控提醒邮件;自动封禁通知始终发送。',
|
||||||
autoBan: '自动封禁用户',
|
autoBan: '自动封禁用户',
|
||||||
autoBanHint: '命中次数达到阈值后将禁用用户账号、刷新认证缓存并发送封禁通知邮件。',
|
autoBanHint: '命中次数达到阈值后将禁用用户账号、刷新认证缓存并发送封禁通知邮件。',
|
||||||
|
cyberPolicyExcludeBan: 'cyber_policy 不计入封号次数',
|
||||||
|
cyberPolicyExcludeBanHint: '开启后,cyber_policy 拦截不再计入自动封号的违规次数:当次不判定封号,历史累计亦排除。风控日志与通知邮件照常。',
|
||||||
|
violationNotCounted: '未计入封号',
|
||||||
banThreshold: '封禁触发次数',
|
banThreshold: '封禁触发次数',
|
||||||
violationWindowHours: '累计窗口(小时)',
|
violationWindowHours: '累计窗口(小时)',
|
||||||
hitRetentionDays: '命中记录保留(天)',
|
hitRetentionDays: '命中记录保留(天)',
|
||||||
@@ -2848,6 +2852,7 @@ export default {
|
|||||||
action: {
|
action: {
|
||||||
block: '拦截',
|
block: '拦截',
|
||||||
keywordBlock: '关键词拦截',
|
keywordBlock: '关键词拦截',
|
||||||
|
cyberPolicy: '网络安全策略',
|
||||||
error: '异常',
|
error: '异常',
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
@@ -5598,6 +5603,9 @@ export default {
|
|||||||
configureLink: '前往 风控中心 配置内容审计',
|
configureLink: '前往 风控中心 配置内容审计',
|
||||||
enabled: '启用风控中心',
|
enabled: '启用风控中心',
|
||||||
enabledHint: '关闭后管理员侧边栏入口隐藏,网关内容审计不会执行。',
|
enabledHint: '关闭后管理员侧边栏入口隐藏,网关内容审计不会执行。',
|
||||||
|
cyberSessionBlock: 'cyber 会话自动屏蔽',
|
||||||
|
cyberSessionBlockHint: '开启后,被上游网络安全策略(cyber_policy)拦截的会话将在 TTL 内被本地屏蔽,不再发往上游。仅屏蔽该会话,不影响同 Key 其他会话。',
|
||||||
|
cyberSessionBlockTTL: '屏蔽时长(秒)',
|
||||||
},
|
},
|
||||||
affiliate: {
|
affiliate: {
|
||||||
title: '邀请返利',
|
title: '邀请返利',
|
||||||
|
|||||||
@@ -1203,7 +1203,7 @@ export interface CodexSessionImportResult {
|
|||||||
// ==================== Usage & Redeem Types ====================
|
// ==================== Usage & Redeem Types ====================
|
||||||
|
|
||||||
export type RedeemCodeType = 'balance' | 'concurrency' | 'subscription' | 'invitation'
|
export type RedeemCodeType = 'balance' | 'concurrency' | 'subscription' | 'invitation'
|
||||||
export type UsageRequestType = 'unknown' | 'sync' | 'stream' | 'ws_v2'
|
export type UsageRequestType = 'unknown' | 'sync' | 'stream' | 'ws_v2' | 'cyber'
|
||||||
export type ImageSizeSource = 'output' | 'input' | 'default' | 'legacy'
|
export type ImageSizeSource = 'output' | 'input' | 'default' | 'legacy'
|
||||||
export type ImageSizeBreakdown = Record<string, number>
|
export type ImageSizeBreakdown = Record<string, number>
|
||||||
|
|
||||||
|
|||||||
@@ -6,7 +6,7 @@ export interface UsageRequestTypeLike {
|
|||||||
openai_ws_mode?: boolean | null
|
openai_ws_mode?: boolean | null
|
||||||
}
|
}
|
||||||
|
|
||||||
const VALID_REQUEST_TYPES = new Set<UsageRequestType>(['unknown', 'sync', 'stream', 'ws_v2'])
|
const VALID_REQUEST_TYPES = new Set<UsageRequestType>(['unknown', 'sync', 'stream', 'ws_v2', 'cyber'])
|
||||||
|
|
||||||
export const isUsageRequestType = (value: unknown): value is UsageRequestType => {
|
export const isUsageRequestType = (value: unknown): value is UsageRequestType => {
|
||||||
return typeof value === 'string' && VALID_REQUEST_TYPES.has(value as UsageRequestType)
|
return typeof value === 'string' && VALID_REQUEST_TYPES.has(value as UsageRequestType)
|
||||||
@@ -23,7 +23,8 @@ export const resolveUsageRequestType = (value: UsageRequestTypeLike): UsageReque
|
|||||||
}
|
}
|
||||||
|
|
||||||
export const requestTypeToLegacyStream = (requestType?: UsageRequestType | null): boolean | null | undefined => {
|
export const requestTypeToLegacyStream = (requestType?: UsageRequestType | null): boolean | null | undefined => {
|
||||||
if (!requestType || requestType === 'unknown') {
|
// cyber 与 stream 正交(cyber 可发生在 stream 或非 stream 请求),不映射到 legacy stream 维度。
|
||||||
|
if (!requestType || requestType === 'unknown' || requestType === 'cyber') {
|
||||||
return null
|
return null
|
||||||
}
|
}
|
||||||
if (requestType === 'sync') {
|
if (requestType === 'sync') {
|
||||||
|
|||||||
@@ -881,6 +881,13 @@
|
|||||||
</div>
|
</div>
|
||||||
<Toggle v-model="configForm.auto_ban_enabled" />
|
<Toggle v-model="configForm.auto_ban_enabled" />
|
||||||
</div>
|
</div>
|
||||||
|
<div class="flex items-center justify-between rounded-lg border border-gray-100 p-4 dark:border-dark-700 lg:col-span-2">
|
||||||
|
<div>
|
||||||
|
<p class="text-sm font-medium text-gray-900 dark:text-white">{{ t('admin.riskControl.cyberPolicyExcludeBan') }}</p>
|
||||||
|
<p class="mt-1 text-xs text-gray-500 dark:text-gray-400">{{ t('admin.riskControl.cyberPolicyExcludeBanHint') }}</p>
|
||||||
|
</div>
|
||||||
|
<Toggle v-model="configForm.cyber_policy_exclude_from_ban_count" />
|
||||||
|
</div>
|
||||||
<div>
|
<div>
|
||||||
<label class="input-label">{{ t('admin.riskControl.banThreshold') }}</label>
|
<label class="input-label">{{ t('admin.riskControl.banThreshold') }}</label>
|
||||||
<input v-model.number="configForm.ban_threshold" type="number" min="1" max="1000" class="input" />
|
<input v-model.number="configForm.ban_threshold" type="number" min="1" max="1000" class="input" />
|
||||||
@@ -1226,6 +1233,7 @@ const configForm = reactive({
|
|||||||
block_message: '内容审计命中风险规则,请调整输入后重试',
|
block_message: '内容审计命中风险规则,请调整输入后重试',
|
||||||
email_on_hit: true,
|
email_on_hit: true,
|
||||||
auto_ban_enabled: true,
|
auto_ban_enabled: true,
|
||||||
|
cyber_policy_exclude_from_ban_count: false,
|
||||||
ban_threshold: 10,
|
ban_threshold: 10,
|
||||||
violation_window_hours: 720,
|
violation_window_hours: 720,
|
||||||
hit_retention_days: 180,
|
hit_retention_days: 180,
|
||||||
@@ -1702,6 +1710,7 @@ function applyConfig(config: ContentModerationConfig) {
|
|||||||
configForm.block_message = config.block_message || '内容审计命中风险规则,请调整输入后重试'
|
configForm.block_message = config.block_message || '内容审计命中风险规则,请调整输入后重试'
|
||||||
configForm.email_on_hit = config.email_on_hit ?? true
|
configForm.email_on_hit = config.email_on_hit ?? true
|
||||||
configForm.auto_ban_enabled = config.auto_ban_enabled ?? true
|
configForm.auto_ban_enabled = config.auto_ban_enabled ?? true
|
||||||
|
configForm.cyber_policy_exclude_from_ban_count = config.cyber_policy_exclude_from_ban_count ?? false
|
||||||
configForm.ban_threshold = config.ban_threshold || 10
|
configForm.ban_threshold = config.ban_threshold || 10
|
||||||
configForm.violation_window_hours = config.violation_window_hours || 720
|
configForm.violation_window_hours = config.violation_window_hours || 720
|
||||||
configForm.hit_retention_days = config.hit_retention_days || 180
|
configForm.hit_retention_days = config.hit_retention_days || 180
|
||||||
@@ -1782,6 +1791,7 @@ async function saveConfig() {
|
|||||||
block_message: configForm.block_message || '内容审计命中风险规则,请调整输入后重试',
|
block_message: configForm.block_message || '内容审计命中风险规则,请调整输入后重试',
|
||||||
email_on_hit: configForm.email_on_hit,
|
email_on_hit: configForm.email_on_hit,
|
||||||
auto_ban_enabled: configForm.auto_ban_enabled,
|
auto_ban_enabled: configForm.auto_ban_enabled,
|
||||||
|
cyber_policy_exclude_from_ban_count: configForm.cyber_policy_exclude_from_ban_count,
|
||||||
ban_threshold: Number(configForm.ban_threshold) || 10,
|
ban_threshold: Number(configForm.ban_threshold) || 10,
|
||||||
violation_window_hours: Number(configForm.violation_window_hours) || 720,
|
violation_window_hours: Number(configForm.violation_window_hours) || 720,
|
||||||
hit_retention_days: Number(configForm.hit_retention_days) || 180,
|
hit_retention_days: Number(configForm.hit_retention_days) || 180,
|
||||||
@@ -2097,6 +2107,7 @@ function modeDescription(mode: ModerationMode): string {
|
|||||||
}
|
}
|
||||||
|
|
||||||
function resultLabel(row: ContentModerationLog): string {
|
function resultLabel(row: ContentModerationLog): string {
|
||||||
|
if (row.action === 'cyber_policy') return t('admin.riskControl.action.cyberPolicy')
|
||||||
if (row.action === 'keyword_block') return t('admin.riskControl.action.keywordBlock')
|
if (row.action === 'keyword_block') return t('admin.riskControl.action.keywordBlock')
|
||||||
if (row.action === 'block') return t('admin.riskControl.action.block')
|
if (row.action === 'block') return t('admin.riskControl.action.block')
|
||||||
if (row.action === 'error' || row.error) return t('admin.riskControl.action.error')
|
if (row.action === 'error' || row.error) return t('admin.riskControl.action.error')
|
||||||
@@ -2105,7 +2116,7 @@ function resultLabel(row: ContentModerationLog): string {
|
|||||||
}
|
}
|
||||||
|
|
||||||
function resultBadgeClass(row: ContentModerationLog): string {
|
function resultBadgeClass(row: ContentModerationLog): string {
|
||||||
if (row.action === 'block' || row.action === 'keyword_block') return 'bg-red-100 text-red-700 dark:bg-red-900/30 dark:text-red-300'
|
if (row.action === 'block' || row.action === 'keyword_block' || row.action === 'cyber_policy') return 'bg-red-100 text-red-700 dark:bg-red-900/30 dark:text-red-300'
|
||||||
if (row.action === 'error' || row.error) return 'bg-amber-100 text-amber-700 dark:bg-amber-900/30 dark:text-amber-300'
|
if (row.action === 'error' || row.error) return 'bg-amber-100 text-amber-700 dark:bg-amber-900/30 dark:text-amber-300'
|
||||||
if (row.flagged) return 'bg-pink-100 text-pink-700 dark:bg-pink-900/30 dark:text-pink-300'
|
if (row.flagged) return 'bg-pink-100 text-pink-700 dark:bg-pink-900/30 dark:text-pink-300'
|
||||||
return 'bg-green-100 text-green-700 dark:bg-green-900/30 dark:text-green-300'
|
return 'bg-green-100 text-green-700 dark:bg-green-900/30 dark:text-green-300'
|
||||||
@@ -2303,6 +2314,7 @@ function parseBlockedKeywords(value: string): string[] {
|
|||||||
|
|
||||||
function violationCountText(row: ContentModerationLog): string {
|
function violationCountText(row: ContentModerationLog): string {
|
||||||
if (!row.flagged) return '-'
|
if (!row.flagged) return '-'
|
||||||
|
if (row.violation_count === 0) return t('admin.riskControl.violationNotCounted')
|
||||||
return t('admin.riskControl.violationCount', { count: row.violation_count || 1 })
|
return t('admin.riskControl.violationCount', { count: row.violation_count || 1 })
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -5297,6 +5297,31 @@
|
|||||||
</div>
|
</div>
|
||||||
<Toggle v-model="form.risk_control_enabled" />
|
<Toggle v-model="form.risk_control_enabled" />
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
|
<div class="flex items-center justify-between">
|
||||||
|
<div>
|
||||||
|
<label class="text-sm font-medium text-gray-700 dark:text-gray-300">
|
||||||
|
{{ t('admin.settings.features.riskControl.cyberSessionBlock') }}
|
||||||
|
</label>
|
||||||
|
<p class="mt-0.5 text-xs text-gray-500 dark:text-gray-400">
|
||||||
|
{{ t('admin.settings.features.riskControl.cyberSessionBlockHint') }}
|
||||||
|
</p>
|
||||||
|
</div>
|
||||||
|
<Toggle v-model="form.cyber_session_block_enabled" />
|
||||||
|
</div>
|
||||||
|
|
||||||
|
<div v-if="form.cyber_session_block_enabled">
|
||||||
|
<label class="input-label">
|
||||||
|
{{ t('admin.settings.features.riskControl.cyberSessionBlockTTL') }}
|
||||||
|
<span class="text-red-500">*</span>
|
||||||
|
</label>
|
||||||
|
<input
|
||||||
|
v-model.number="form.cyber_session_block_ttl_seconds"
|
||||||
|
type="number"
|
||||||
|
min="1"
|
||||||
|
class="input"
|
||||||
|
/>
|
||||||
|
</div>
|
||||||
</div>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
@@ -7045,6 +7070,8 @@ const form = reactive<SettingsForm>({
|
|||||||
hide_ccs_import_button: false,
|
hide_ccs_import_button: false,
|
||||||
payment_enabled: false,
|
payment_enabled: false,
|
||||||
risk_control_enabled: false,
|
risk_control_enabled: false,
|
||||||
|
cyber_session_block_enabled: false,
|
||||||
|
cyber_session_block_ttl_seconds: 3600,
|
||||||
payment_min_amount: 1,
|
payment_min_amount: 1,
|
||||||
payment_max_amount: 10000,
|
payment_max_amount: 10000,
|
||||||
payment_daily_limit: 50000,
|
payment_daily_limit: 50000,
|
||||||
@@ -8316,6 +8343,9 @@ async function saveSettings() {
|
|||||||
// Payment configuration
|
// Payment configuration
|
||||||
payment_enabled: form.payment_enabled,
|
payment_enabled: form.payment_enabled,
|
||||||
risk_control_enabled: form.risk_control_enabled,
|
risk_control_enabled: form.risk_control_enabled,
|
||||||
|
cyber_session_block_enabled: form.cyber_session_block_enabled,
|
||||||
|
cyber_session_block_ttl_seconds:
|
||||||
|
Number(form.cyber_session_block_ttl_seconds) || 3600,
|
||||||
payment_min_amount: Number(form.payment_min_amount) || 0,
|
payment_min_amount: Number(form.payment_min_amount) || 0,
|
||||||
payment_max_amount: Number(form.payment_max_amount) || 0,
|
payment_max_amount: Number(form.payment_max_amount) || 0,
|
||||||
payment_daily_limit: Number(form.payment_daily_limit) || 0,
|
payment_daily_limit: Number(form.payment_daily_limit) || 0,
|
||||||
|
|||||||
@@ -490,6 +490,7 @@ const cancelExport = () => exportAbortController?.abort()
|
|||||||
const openCleanupDialog = () => { cleanupDialogVisible.value = true }
|
const openCleanupDialog = () => { cleanupDialogVisible.value = true }
|
||||||
const getRequestTypeLabel = (log: AdminUsageLog): string => {
|
const getRequestTypeLabel = (log: AdminUsageLog): string => {
|
||||||
const requestType = resolveUsageRequestType(log)
|
const requestType = resolveUsageRequestType(log)
|
||||||
|
if (requestType === 'cyber') return t('usage.cyber')
|
||||||
if (requestType === 'ws_v2') return t('usage.ws')
|
if (requestType === 'ws_v2') return t('usage.ws')
|
||||||
if (requestType === 'stream') return t('usage.stream')
|
if (requestType === 'stream') return t('usage.stream')
|
||||||
if (requestType === 'sync') return t('usage.sync')
|
if (requestType === 'sync') return t('usage.sync')
|
||||||
|
|||||||
@@ -773,6 +773,7 @@ const formatUserAgent = (ua: string): string => {
|
|||||||
|
|
||||||
const getRequestTypeLabel = (log: UsageLog): string => {
|
const getRequestTypeLabel = (log: UsageLog): string => {
|
||||||
const requestType = resolveUsageRequestType(log)
|
const requestType = resolveUsageRequestType(log)
|
||||||
|
if (requestType === 'cyber') return t('usage.cyber')
|
||||||
if (requestType === 'ws_v2') return t('usage.ws')
|
if (requestType === 'ws_v2') return t('usage.ws')
|
||||||
if (requestType === 'stream') return t('usage.stream')
|
if (requestType === 'stream') return t('usage.stream')
|
||||||
if (requestType === 'sync') return t('usage.sync')
|
if (requestType === 'sync') return t('usage.sync')
|
||||||
@@ -781,6 +782,7 @@ const getRequestTypeLabel = (log: UsageLog): string => {
|
|||||||
|
|
||||||
const getRequestTypeBadgeClass = (log: UsageLog): string => {
|
const getRequestTypeBadgeClass = (log: UsageLog): string => {
|
||||||
const requestType = resolveUsageRequestType(log)
|
const requestType = resolveUsageRequestType(log)
|
||||||
|
if (requestType === 'cyber') return 'bg-red-100 text-red-800 dark:bg-red-900 dark:text-red-200'
|
||||||
if (requestType === 'ws_v2') return 'bg-violet-100 text-violet-800 dark:bg-violet-900 dark:text-violet-200'
|
if (requestType === 'ws_v2') return 'bg-violet-100 text-violet-800 dark:bg-violet-900 dark:text-violet-200'
|
||||||
if (requestType === 'stream') return 'bg-blue-100 text-blue-800 dark:bg-blue-900 dark:text-blue-200'
|
if (requestType === 'stream') return 'bg-blue-100 text-blue-800 dark:bg-blue-900 dark:text-blue-200'
|
||||||
if (requestType === 'sync') return 'bg-gray-100 text-gray-800 dark:bg-gray-700 dark:text-gray-200'
|
if (requestType === 'sync') return 'bg-gray-100 text-gray-800 dark:bg-gray-700 dark:text-gray-200'
|
||||||
@@ -790,6 +792,7 @@ const getRequestTypeBadgeClass = (log: UsageLog): string => {
|
|||||||
|
|
||||||
const getRequestTypeExportText = (log: UsageLog): string => {
|
const getRequestTypeExportText = (log: UsageLog): string => {
|
||||||
const requestType = resolveUsageRequestType(log)
|
const requestType = resolveUsageRequestType(log)
|
||||||
|
if (requestType === 'cyber') return 'Cyber'
|
||||||
if (requestType === 'ws_v2') return 'WS'
|
if (requestType === 'ws_v2') return 'WS'
|
||||||
if (requestType === 'stream') return 'Stream'
|
if (requestType === 'stream') return 'Stream'
|
||||||
if (requestType === 'sync') return 'Sync'
|
if (requestType === 'sync') return 'Sync'
|
||||||
|
|||||||
Reference in New Issue
Block a user