From b62b573f7f9fd9b859e4f4e4b85b89ab4a520ad4 Mon Sep 17 00:00:00 2001 From: DaydreamCoding <22166516+DaydreamCoding@users.noreply.github.com> Date: Thu, 11 Jun 2026 17:26:14 +0800 Subject: [PATCH] =?UTF-8?q?feat(openai):=20cyber=5Fpolicy=20=E7=A1=AC?= =?UTF-8?q?=E9=98=BB=E6=96=AD=E5=85=A8=E9=93=BE=E8=B7=AF=E9=80=8F=E4=BC=A0?= =?UTF-8?q?=E3=80=81=E5=AE=A1=E8=AE=A1=E4=B8=8E=E8=AE=A1=E8=B4=B9?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 上游对单次请求下发 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 --- backend/cmd/server/wire_gen.go | 2 +- .../admin/content_moderation_handler.go | 124 +++--- .../internal/handler/admin/setting_handler.go | 36 +- backend/internal/handler/dto/settings.go | 4 + .../handler/openai_chat_completions.go | 10 + .../handler/openai_gateway_cyber_test.go | 180 ++++++++ .../handler/openai_gateway_handler.go | 420 +++++++++++++++++- .../handler/openai_gateway_handler_test.go | 30 +- backend/internal/handler/ops_error_logger.go | 10 + .../handler/ops_error_logger_cyber_test.go | 22 + .../repository/content_moderation_repo.go | 14 +- .../content_moderation_repo_test.go | 22 +- backend/internal/repository/gateway_cache.go | 20 + .../repository/ops_error_where_test.go | 22 + backend/internal/repository/ops_repo.go | 6 +- backend/internal/server/api_contract_test.go | 4 + .../internal/service/content_moderation.go | 332 ++++++++++---- .../service/content_moderation_cyber_test.go | 302 +++++++++++++ .../service/content_moderation_email.go | 38 ++ .../service/content_moderation_test.go | 50 ++- backend/internal/service/domain_constants.go | 2 + .../service/notification_email_service.go | 39 ++ .../notification_email_service_test.go | 1 + .../internal/service/openai_cyber_policy.go | 88 ++++ .../service/openai_cyber_policy_test.go | 88 ++++ .../service/openai_cyber_session_block.go | 99 +++++ .../openai_cyber_session_block_test.go | 192 ++++++++ .../openai_gateway_chat_completions.go | 81 +++- .../openai_gateway_chat_completions_test.go | 56 +++ .../openai_gateway_compat_cyber_test.go | 136 ++++++ .../service/openai_gateway_messages.go | 75 ++++ .../openai_gateway_record_usage_test.go | 93 ++++ .../service/openai_gateway_service.go | 163 ++++++- .../service/openai_gateway_service_test.go | 106 +++++ .../openai_image_generation_controls_test.go | 28 ++ .../internal/service/openai_ws_forwarder.go | 26 ++ .../service/openai_ws_forwarder_test.go | 54 +++ backend/internal/service/ops_user_error.go | 4 + .../service/ops_user_error_cyber_test.go | 14 + backend/internal/service/setting_service.go | 89 ++++ backend/internal/service/settings_view.go | 2 + backend/internal/service/usage_log.go | 17 +- .../internal/service/usage_log_cyber_test.go | 25 ++ frontend/src/api/admin/riskControl.ts | 2 + frontend/src/api/admin/settings.ts | 10 + .../components/admin/usage/UsageFilters.vue | 3 +- .../src/components/admin/usage/UsageTable.vue | 2 + .../user/UserErrorRequestsTable.vue | 2 +- frontend/src/i18n/locales/en.ts | 10 +- frontend/src/i18n/locales/zh.ts | 10 +- frontend/src/types/index.ts | 2 +- frontend/src/utils/usageRequestType.ts | 5 +- frontend/src/views/admin/RiskControlView.vue | 14 +- frontend/src/views/admin/SettingsView.vue | 30 ++ frontend/src/views/admin/UsageView.vue | 1 + frontend/src/views/user/UsageView.vue | 3 + 56 files changed, 3036 insertions(+), 184 deletions(-) create mode 100644 backend/internal/handler/openai_gateway_cyber_test.go create mode 100644 backend/internal/handler/ops_error_logger_cyber_test.go create mode 100644 backend/internal/service/content_moderation_cyber_test.go create mode 100644 backend/internal/service/openai_cyber_policy.go create mode 100644 backend/internal/service/openai_cyber_policy_test.go create mode 100644 backend/internal/service/openai_cyber_session_block.go create mode 100644 backend/internal/service/openai_cyber_session_block_test.go create mode 100644 backend/internal/service/openai_gateway_compat_cyber_test.go create mode 100644 backend/internal/service/ops_user_error_cyber_test.go create mode 100644 backend/internal/service/usage_log_cyber_test.go diff --git a/backend/cmd/server/wire_gen.go b/backend/cmd/server/wire_gen.go index 879d56d22f..5111481558 100644 --- a/backend/cmd/server/wire_gen.go +++ b/backend/cmd/server/wire_gen.go @@ -246,7 +246,7 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { userMsgQueueCache := repository.NewUserMsgQueueCache(redisClient) userMessageQueueService := service.ProvideUserMessageQueueService(userMsgQueueCache, rpmCache, configConfig) 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) totpHandler := handler.NewTotpHandler(totpService) handlerPaymentHandler := handler.NewPaymentHandler(paymentService, paymentConfigService, channelService) diff --git a/backend/internal/handler/admin/content_moderation_handler.go b/backend/internal/handler/admin/content_moderation_handler.go index e11ea6eb21..00a8dca05f 100644 --- a/backend/internal/handler/admin/content_moderation_handler.go +++ b/backend/internal/handler/admin/content_moderation_handler.go @@ -20,36 +20,39 @@ func NewContentModerationHandler(svc *service.ContentModerationService) *Content } type contentModerationConfigRequest struct { - Enabled *bool `json:"enabled"` - Mode *string `json:"mode"` - BaseURL *string `json:"base_url"` - Model *string `json:"model"` - APIKey *string `json:"api_key"` - APIKeys *[]string `json:"api_keys"` - APIKeysMode string `json:"api_keys_mode"` - DeleteAPIKeyHashes *[]string `json:"delete_api_key_hashes"` - ClearAPIKey bool `json:"clear_api_key"` - TimeoutMS *int `json:"timeout_ms"` - SampleRate *int `json:"sample_rate"` - AllGroups *bool `json:"all_groups"` - GroupIDs *[]int64 `json:"group_ids"` - RecordNonHits *bool `json:"record_non_hits"` - Thresholds *map[string]float64 `json:"thresholds"` - WorkerCount *int `json:"worker_count"` - QueueSize *int `json:"queue_size"` - BlockStatus *int `json:"block_status"` - BlockMessage *string `json:"block_message"` - EmailOnHit *bool `json:"email_on_hit"` - AutoBanEnabled *bool `json:"auto_ban_enabled"` - BanThreshold *int `json:"ban_threshold"` - ViolationWindowHours *int `json:"violation_window_hours"` - RetryCount *int `json:"retry_count"` - HitRetentionDays *int `json:"hit_retention_days"` - NonHitRetentionDays *int `json:"non_hit_retention_days"` - PreHashCheckEnabled *bool `json:"pre_hash_check_enabled"` - BlockedKeywords *[]string `json:"blocked_keywords"` - KeywordBlockingMode *string `json:"keyword_blocking_mode"` - ModelFilter *service.ContentModerationModelFilter `json:"model_filter"` + Enabled *bool `json:"enabled"` + Mode *string `json:"mode"` + BaseURL *string `json:"base_url"` + Model *string `json:"model"` + APIKey *string `json:"api_key"` + APIKeys *[]string `json:"api_keys"` + APIKeysMode string `json:"api_keys_mode"` + DeleteAPIKeyHashes *[]string `json:"delete_api_key_hashes"` + ClearAPIKey bool `json:"clear_api_key"` + TimeoutMS *int `json:"timeout_ms"` + SampleRate *int `json:"sample_rate"` + AllGroups *bool `json:"all_groups"` + GroupIDs *[]int64 `json:"group_ids"` + RecordNonHits *bool `json:"record_non_hits"` + Thresholds *map[string]float64 `json:"thresholds"` + WorkerCount *int `json:"worker_count"` + QueueSize *int `json:"queue_size"` + BlockStatus *int `json:"block_status"` + BlockMessage *string `json:"block_message"` + EmailOnHit *bool `json:"email_on_hit"` + AutoBanEnabled *bool `json:"auto_ban_enabled"` + BanThreshold *int `json:"ban_threshold"` + 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"` + HitRetentionDays *int `json:"hit_retention_days"` + NonHitRetentionDays *int `json:"non_hit_retention_days"` + PreHashCheckEnabled *bool `json:"pre_hash_check_enabled"` + BlockedKeywords *[]string `json:"blocked_keywords"` + KeywordBlockingMode *string `json:"keyword_blocking_mode"` + ModelFilter *service.ContentModerationModelFilter `json:"model_filter"` } type contentModerationAPIKeyTestRequest struct { @@ -81,36 +84,37 @@ func (h *ContentModerationHandler) UpdateConfig(c *gin.Context) { return } cfg, err := h.service.UpdateConfig(c.Request.Context(), service.UpdateContentModerationConfigInput{ - Enabled: req.Enabled, - Mode: req.Mode, - BaseURL: req.BaseURL, - Model: req.Model, - APIKey: req.APIKey, - APIKeys: req.APIKeys, - APIKeysMode: req.APIKeysMode, - DeleteAPIKeyHashes: req.DeleteAPIKeyHashes, - ClearAPIKey: req.ClearAPIKey, - TimeoutMS: req.TimeoutMS, - SampleRate: req.SampleRate, - AllGroups: req.AllGroups, - GroupIDs: req.GroupIDs, - RecordNonHits: req.RecordNonHits, - Thresholds: req.Thresholds, - WorkerCount: req.WorkerCount, - QueueSize: req.QueueSize, - BlockStatus: req.BlockStatus, - BlockMessage: req.BlockMessage, - EmailOnHit: req.EmailOnHit, - AutoBanEnabled: req.AutoBanEnabled, - BanThreshold: req.BanThreshold, - ViolationWindowHours: req.ViolationWindowHours, - RetryCount: req.RetryCount, - HitRetentionDays: req.HitRetentionDays, - NonHitRetentionDays: req.NonHitRetentionDays, - PreHashCheckEnabled: req.PreHashCheckEnabled, - BlockedKeywords: req.BlockedKeywords, - KeywordBlockingMode: req.KeywordBlockingMode, - ModelFilter: req.ModelFilter, + Enabled: req.Enabled, + Mode: req.Mode, + BaseURL: req.BaseURL, + Model: req.Model, + APIKey: req.APIKey, + APIKeys: req.APIKeys, + APIKeysMode: req.APIKeysMode, + DeleteAPIKeyHashes: req.DeleteAPIKeyHashes, + ClearAPIKey: req.ClearAPIKey, + TimeoutMS: req.TimeoutMS, + SampleRate: req.SampleRate, + AllGroups: req.AllGroups, + GroupIDs: req.GroupIDs, + RecordNonHits: req.RecordNonHits, + Thresholds: req.Thresholds, + WorkerCount: req.WorkerCount, + QueueSize: req.QueueSize, + BlockStatus: req.BlockStatus, + BlockMessage: req.BlockMessage, + EmailOnHit: req.EmailOnHit, + AutoBanEnabled: req.AutoBanEnabled, + BanThreshold: req.BanThreshold, + ViolationWindowHours: req.ViolationWindowHours, + CyberPolicyExcludeFromBanCount: req.CyberPolicyExcludeFromBanCount, + RetryCount: req.RetryCount, + HitRetentionDays: req.HitRetentionDays, + NonHitRetentionDays: req.NonHitRetentionDays, + PreHashCheckEnabled: req.PreHashCheckEnabled, + BlockedKeywords: req.BlockedKeywords, + KeywordBlockingMode: req.KeywordBlockingMode, + ModelFilter: req.ModelFilter, }) if err != nil { response.ErrorFrom(c, err) diff --git a/backend/internal/handler/admin/setting_handler.go b/backend/internal/handler/admin/setting_handler.go index 0beb15d3a1..41896ac41a 100644 --- a/backend/internal/handler/admin/setting_handler.go +++ b/backend/internal/handler/admin/setting_handler.go @@ -228,6 +228,8 @@ func (h *SettingHandler) GetSettings(c *gin.Context) { DefaultConcurrency: settings.DefaultConcurrency, DefaultBalance: settings.DefaultBalance, RiskControlEnabled: settings.RiskControlEnabled, + CyberSessionBlockEnabled: settings.CyberSessionBlockEnabled, + CyberSessionBlockTTLSeconds: settings.CyberSessionBlockTTLSeconds, AffiliateRebateRate: settings.AffiliateRebateRate, AffiliateRebateFreezeHours: settings.AffiliateRebateFreezeHours, AffiliateRebateDurationDays: settings.AffiliateRebateDurationDays, @@ -646,6 +648,10 @@ type UpdateSettingsRequest struct { // 风控中心功能开关 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) 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{ // 系统全局 platform quota 默认值(整体替换语义) DefaultPlatformQuotas: req.DefaultPlatformQuotas, @@ -1769,6 +1781,18 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) { } 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 中的值); @@ -2090,8 +2114,10 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) { AffiliateEnabled: updatedSettings.AffiliateEnabled, - RiskControlEnabled: updatedSettings.RiskControlEnabled, - AllowUserViewErrorRequests: updatedSettings.AllowUserViewErrorRequests, + RiskControlEnabled: updatedSettings.RiskControlEnabled, + CyberSessionBlockEnabled: updatedSettings.CyberSessionBlockEnabled, + CyberSessionBlockTTLSeconds: updatedSettings.CyberSessionBlockTTLSeconds, + AllowUserViewErrorRequests: updatedSettings.AllowUserViewErrorRequests, } if fastPolicy, err := h.settingService.GetOpenAIFastPolicySettings(c.Request.Context()); err != nil { slog.Error("openai_fast_policy_settings_get_failed", "error", err) @@ -2572,6 +2598,12 @@ func diffSettings(before *service.SystemSettings, after *service.SystemSettings, if before.RiskControlEnabled != after.RiskControlEnabled { 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,整体比较) if !equalPlatformQuotaSettings(before.DefaultPlatformQuotas, after.DefaultPlatformQuotas) { changed = append(changed, service.SettingKeyDefaultPlatformQuotas) diff --git a/backend/internal/handler/dto/settings.go b/backend/internal/handler/dto/settings.go index da89ac237c..c49522de7c 100644 --- a/backend/internal/handler/dto/settings.go +++ b/backend/internal/handler/dto/settings.go @@ -244,6 +244,10 @@ type SystemSettings struct { // 风控中心功能开关 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 AffiliateEnabled bool `json:"affiliate_enabled"` diff --git a/backend/internal/handler/openai_chat_completions.go b/backend/internal/handler/openai_chat_completions.go index 64a0f97ba7..a59450f130 100644 --- a/backend/internal/handler/openai_chat_completions.go +++ b/backend/internal/handler/openai_chat_completions.go @@ -89,6 +89,9 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) { h.errorResponse(c, contentModerationStatus(decision), contentModerationErrorCode(decision), decision.Message) return } + if h.rejectIfCyberSessionBlocked(c, apiKey, body, reqModel, cyberBlockFormatChat) { + return + } // 解析渠道级模型映射 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, "") }() + 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() upstreamLatencyMs, _ := getContextInt64(c, service.OpsUpstreamLatencyMsKey) @@ -283,6 +291,7 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) { inboundEndpoint := GetInboundEndpoint(c) upstreamEndpoint := resolveRawCCUpstreamEndpoint(c, account) + cyberBlocked := service.GetOpsCyberPolicy(c) != nil h.submitOpenAIUsageRecordTask(c.Request.Context(), result, func(ctx context.Context) { if err := h.gatewayService.RecordUsage(ctx, &service.OpenAIRecordUsageInput{ Result: result, @@ -296,6 +305,7 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) { IPAddress: clientIP, APIKeyService: h.apiKeyService, ChannelUsageFields: channelMapping.ToUsageFields(reqModel, result.UpstreamModel), + CyberBlocked: cyberBlocked, }); err != nil { logger.L().With( zap.String("component", "handler.openai_gateway.chat_completions"), diff --git a/backend/internal/handler/openai_gateway_cyber_test.go b/backend/internal/handler/openai_gateway_cyber_test.go new file mode 100644 index 0000000000..b0060963db --- /dev/null +++ b/backend/internal/handler/openai_gateway_cyber_test.go @@ -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) + }) + } +} diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go index 65bf172ecd..85e8823c22 100644 --- a/backend/internal/handler/openai_gateway_handler.go +++ b/backend/internal/handler/openai_gateway_handler.go @@ -34,6 +34,7 @@ type OpenAIGatewayHandler struct { usageRecordWorkerPool *service.UsageRecordWorkerPool errorPassthroughService *service.ErrorPassthroughService contentModerationService *service.ContentModerationService + opsService *service.OpsService concurrencyHelper *ConcurrencyHelper imageLimiter *imageConcurrencyLimiter maxAccountSwitches int @@ -105,6 +106,7 @@ func NewOpenAIGatewayHandler( usageRecordWorkerPool *service.UsageRecordWorkerPool, errorPassthroughService *service.ErrorPassthroughService, contentModerationService *service.ContentModerationService, + opsService *service.OpsService, cfg *config.Config, ) *OpenAIGatewayHandler { pingInterval := time.Duration(0) @@ -122,6 +124,7 @@ func NewOpenAIGatewayHandler( usageRecordWorkerPool: usageRecordWorkerPool, errorPassthroughService: errorPassthroughService, contentModerationService: contentModerationService, + opsService: opsService, concurrencyHelper: NewConcurrencyHelper(concurrencyService, SSEPingFormatComment, pingInterval), imageLimiter: &imageConcurrencyLimiter{}, maxAccountSwitches: maxAccountSwitches, @@ -305,6 +308,9 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) { // Generate session hash (header first; fallback to prompt_cache_key) sessionHash := h.gatewayService.GenerateSessionHash(c, sessionHashBody) + if h.rejectIfCyberSessionBlocked(c, apiKey, sessionHashBody, reqModel, cyberBlockFormatResponses) { + return + } requireCompact := isOpenAIRemoteCompactPath(c) maxAccountSwitches := h.maxAccountSwitches @@ -387,6 +393,11 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) { }() 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() upstreamLatencyMs, _ := getContextInt64(c, service.OpsUpstreamLatencyMsKey) responseLatencyMs := forwardDurationMs @@ -488,6 +499,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) { upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform) // 使用量记录通过有界 worker 池提交,避免请求热路径创建无界 goroutine。 + cyberBlocked := service.GetOpsCyberPolicy(c) != nil h.submitOpenAIUsageRecordTask(c.Request.Context(), result, func(ctx context.Context) { if err := h.gatewayService.RecordUsage(ctx, &service.OpenAIRecordUsageInput{ Result: result, @@ -502,6 +514,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) { RequestPayloadHash: requestPayloadHash, APIKeyService: h.apiKeyService, ChannelUsageFields: channelMapping.ToUsageFields(reqModel, result.UpstreamModel), + CyberBlocked: cyberBlocked, }); err != nil { logger.L().With( zap.String("component", "handler.openai_gateway.responses"), @@ -713,6 +726,9 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) { sessionHash := h.gatewayService.GenerateSessionHash(c, body) promptCacheKey := h.gatewayService.ExtractSessionID(c, body) sessionHash, promptCacheKey = resolveOpenAIMessagesMetadataSession(sessionHash, promptCacheKey, reqModel, body) + if h.rejectIfCyberSessionBlocked(c, apiKey, body, reqModel, cyberBlockFormatAnthropic) { + return + } maxAccountSwitches := h.maxAccountSwitches 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) }() - + 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() upstreamLatencyMs, _ := getContextInt64(c, service.OpsUpstreamLatencyMsKey) responseLatencyMs := forwardDurationMs @@ -883,6 +903,7 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) { inboundEndpoint := GetInboundEndpoint(c) upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform) + cyberBlocked := service.GetOpsCyberPolicy(c) != nil h.submitOpenAIUsageRecordTask(c.Request.Context(), result, func(ctx context.Context) { if err := h.gatewayService.RecordUsage(ctx, &service.OpenAIRecordUsageInput{ Result: result, @@ -897,6 +918,7 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) { RequestPayloadHash: requestPayloadHash, APIKeyService: h.apiKeyService, ChannelUsageFields: channelMappingMsg.ToUsageFields(reqModel, result.UpstreamModel), + CyberBlocked: cyberBlocked, }); err != nil { logger.L().With( zap.String("component", "handler.openai_gateway.messages"), @@ -1259,6 +1281,17 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { 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) @@ -1430,6 +1463,10 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { return nil }, BeforeTurn: func(turn int) error { + // turn==1 的会话屏蔽已由握手层检查覆盖;连接内 flag 只拦截后续 turn。 + if cyberBlockedThisConn { + return service.NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, cyberSessionBlockedClientMsg, nil) + } if turn == 1 { return nil } @@ -1461,11 +1498,24 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { return nil }, AfterTurn: func(turn int, result *service.OpenAIForwardResult, turnErr error) { + // F1: cyber 标记按 turn 生命周期清理——defer 保证任意早返回路径都执行; + // CyberBlocked 必须在 submit 前同步预捕获(task 闭包由 worker 池异步执行, + // 届时 defer 已清除标记)。 + defer clearCyberPolicyTurnState(c) 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 result == nil || result.ImageCount <= 0 { return } + // cyber 命中时该 turn 的用量已由 recordCyberPolicyIfMarked(forwardErrored=true) + // 按真实 token 记录,这里不再走下方 RecordUsage,避免对同一 turn 双写/双扣费。 + if service.GetOpsCyberPolicy(c) != nil { + return + } reqLog.Warn("openai.websocket_partial_error_with_image_result", zap.Int64("account_id", account.ID), 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) inboundEndpoint := GetInboundEndpoint(c) upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform) + cyberBlocked := service.GetOpsCyberPolicy(c) != nil h.submitOpenAIUsageRecordTask(ctx, result, func(taskCtx context.Context) { if err := h.gatewayService.RecordUsage(taskCtx, &service.OpenAIRecordUsageInput{ Result: result, @@ -1495,6 +1546,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { RequestPayloadHash: requestPayloadHash, APIKeyService: h.apiKeyService, ChannelUsageFields: channelMappingWS.ToUsageFields(reqModel, result.UpstreamModel), + CyberBlocked: cyberBlocked, }); err != nil { reqLog.Error("openai.websocket_record_usage_failed", zap.Int64("account_id", account.ID), @@ -1904,6 +1956,14 @@ func openAIForwardErrorAlreadyCommunicated(c *gin.Context, writerSizeBeforeForwa 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()) for _, prefix := range []string{ "upstream response failed:", @@ -2017,6 +2077,364 @@ func writeContentModerationWSError(ctx context.Context, conn *coderws.Conn, deci _ = 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) { if err == nil { return "-", "-" diff --git a/backend/internal/handler/openai_gateway_handler_test.go b/backend/internal/handler/openai_gateway_handler_test.go index 7743d6ab62..fd3508837e 100644 --- a/backend/internal/handler/openai_gateway_handler_test.go +++ b/backend/internal/handler/openai_gateway_handler_test.go @@ -805,7 +805,7 @@ func (r *contentModerationHandlerTestRepo) ListLogs(ctx context.Context, filter 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 } @@ -813,6 +813,10 @@ func (r *contentModerationHandlerTestRepo) CleanupExpiredLogs(ctx context.Contex return &service.ContentModerationCleanupResult{}, nil } +func (r *contentModerationHandlerTestRepo) UpdateLogEmailSent(ctx context.Context, id int64, sent bool) error { + return nil +} + func TestOpenAIResponsesWebSocket_ContentModerationBlocksFirstFrame(t *testing.T) { gin.SetMode(gin.TestMode) @@ -1654,4 +1658,28 @@ data: {"type":"response.failed","error":{"message":"This content was flagged"}} 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"))) + }) } diff --git a/backend/internal/handler/ops_error_logger.go b/backend/internal/handler/ops_error_logger.go index b86c7f6991..8cd9d02ba0 100644 --- a/backend/internal/handler/ops_error_logger.go +++ b/backend/internal/handler/ops_error_logger.go @@ -549,6 +549,10 @@ func OpsErrorLoggerMiddleware(ops *service.OpsService) gin.HandlerFunc { return } + if shouldSkipOpsErrorLogForCyber(c) { + return + } + status := c.Writer.Status() if status < 400 { // 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 } + +// shouldSkipOpsErrorLogForCyber:cyber_policy 命中的请求由 recordCyberPolicyIfMarked +// 统一落一条 status=403 的错误请求,故中间件跳过自身落库,避免双写。 +func shouldSkipOpsErrorLogForCyber(c *gin.Context) bool { + return service.GetOpsCyberPolicy(c) != nil +} diff --git a/backend/internal/handler/ops_error_logger_cyber_test.go b/backend/internal/handler/ops_error_logger_cyber_test.go new file mode 100644 index 0000000000..07bd6c76eb --- /dev/null +++ b/backend/internal/handler/ops_error_logger_cyber_test.go @@ -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 命中应跳过中间件落库") +} diff --git a/backend/internal/repository/content_moderation_repo.go b/backend/internal/repository/content_moderation_repo.go index 9b19cce988..9a40b3c2fd 100644 --- a/backend/internal/repository/content_moderation_repo.go +++ b/backend/internal/repository/content_moderation_repo.go @@ -177,10 +177,11 @@ LIMIT $`+fmt.Sprint(len(queryArgs)-1)+` OFFSET $`+fmt.Sprint(len(queryArgs)), 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 { return 0, nil } + // SQL 中的 'cyber_policy' 字面量须与 service.ContentModerationActionCyberPolicy 保持一致。 var count int err := r.db.QueryRowContext(ctx, ` WITH last_auto_ban AS ( @@ -193,15 +194,24 @@ FROM content_moderation_logs WHERE user_id = $1 AND flagged = TRUE AND action <> 'hash_block' + AND ($3::bool IS FALSE OR action <> 'cyber_policy') AND created_at >= $2 AND created_at > COALESCE((SELECT at FROM last_auto_ban), '-infinity'::timestamptz) -`, userID, since).Scan(&count) +`, userID, since, excludeCyberPolicy).Scan(&count) if err != nil { return 0, fmt.Errorf("count user content moderation flagged logs: %w", err) } 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) { result := &service.ContentModerationCleanupResult{FinishedAt: time.Now()} if r == nil || r.db == nil { diff --git a/backend/internal/repository/content_moderation_repo_test.go b/backend/internal/repository/content_moderation_repo_test.go index 6d5faa121f..3f34f9d05f 100644 --- a/backend/internal/repository/content_moderation_repo_test.go +++ b/backend/internal/repository/content_moderation_repo_test.go @@ -29,12 +29,30 @@ func TestContentModerationRepositoryCountFlaggedByUserSince_ExcludesHashBlock(t repo := NewContentModerationRepository(db) since := time.Now().Add(-time.Hour) mock.ExpectQuery(regexp.QuoteMeta("AND action <> 'hash_block'")). - WithArgs(int64(1001), since). + WithArgs(int64(1001), since, false). 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.Equal(t, 2, count) 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()) +} diff --git a/backend/internal/repository/gateway_cache.go b/backend/internal/repository/gateway_cache.go index 58291b6652..b8229eed1a 100644 --- a/backend/internal/repository/gateway_cache.go +++ b/backend/internal/repository/gateway_cache.go @@ -51,3 +51,23 @@ func (c *gatewayCache) DeleteSessionAccountID(ctx context.Context, groupID int64 key := buildSessionKey(groupID, sessionHash) 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 +} diff --git a/backend/internal/repository/ops_error_where_test.go b/backend/internal/repository/ops_error_where_test.go index 9bebb15895..5b9d7ab1c3 100644 --- a/backend/internal/repository/ops_error_where_test.go +++ b/backend/internal/repository/ops_error_where_test.go @@ -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) { uid := int64(42) diff --git a/backend/internal/repository/ops_repo.go b/backend/internal/repository/ops_repo.go index f300a171e1..f272812512 100644 --- a/backend/internal/repository/ops_repo.go +++ b/backend/internal/repository/ops_repo.go @@ -919,8 +919,12 @@ func buildOpsErrorLogsWhere(filter *service.OpsErrorLogFilter) (string, []any) { resolvedFilter = filter.Resolved } // 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" { - 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() { diff --git a/backend/internal/server/api_contract_test.go b/backend/internal/server/api_contract_test.go index fb9b616adc..979914eb1a 100644 --- a/backend/internal/server/api_contract_test.go +++ b/backend/internal/server/api_contract_test.go @@ -884,6 +884,8 @@ func TestAPIContracts(t *testing.T) { "channel_monitor_default_interval_seconds": 60, "available_channels_enabled": false, "risk_control_enabled": false, + "cyber_session_block_enabled": false, + "cyber_session_block_ttl_seconds": 3600, "affiliate_enabled": false, "wechat_connect_enabled": false, "wechat_connect_app_id": "", @@ -1120,6 +1122,8 @@ func TestAPIContracts(t *testing.T) { "channel_monitor_default_interval_seconds": 60, "available_channels_enabled": false, "risk_control_enabled": false, + "cyber_session_block_enabled": false, + "cyber_session_block_ttl_seconds": 3600, "affiliate_enabled": false, "wechat_connect_enabled": true, "wechat_connect_app_id": "wx-open-config", diff --git a/backend/internal/service/content_moderation.go b/backend/internal/service/content_moderation.go index 42b909c90f..0ce4355fe7 100644 --- a/backend/internal/service/content_moderation.go +++ b/backend/internal/service/content_moderation.go @@ -37,6 +37,7 @@ const ( ContentModerationActionHashBlock = "hash_block" ContentModerationActionKeywordBlock = "keyword_block" ContentModerationActionError = "error" + ContentModerationActionCyberPolicy = "cyber_policy" // cyber_policy 硬阻断的风控日志 action(封号计数排除按此值过滤) contentModerationKeywordCategory = "keyword" @@ -160,39 +161,44 @@ type ContentModerationConfig struct { BlockedKeywords []string `json:"blocked_keywords"` KeywordBlockingMode string `json:"keyword_blocking_mode"` 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 { - Enabled bool `json:"enabled"` - Mode string `json:"mode"` - BaseURL string `json:"base_url"` - Model string `json:"model"` - APIKeyConfigured bool `json:"api_key_configured"` - APIKeyMasked string `json:"api_key_masked"` - APIKeyCount int `json:"api_key_count"` - APIKeyMasks []string `json:"api_key_masks"` - APIKeyStatuses []ContentModerationAPIKeyStatus `json:"api_key_statuses"` - TimeoutMS int `json:"timeout_ms"` - SampleRate int `json:"sample_rate"` - AllGroups bool `json:"all_groups"` - GroupIDs []int64 `json:"group_ids"` - RecordNonHits bool `json:"record_non_hits"` - Thresholds map[string]float64 `json:"thresholds"` - WorkerCount int `json:"worker_count"` - QueueSize int `json:"queue_size"` - BlockStatus int `json:"block_status"` - BlockMessage string `json:"block_message"` - EmailOnHit bool `json:"email_on_hit"` - AutoBanEnabled bool `json:"auto_ban_enabled"` - BanThreshold int `json:"ban_threshold"` - ViolationWindowHours int `json:"violation_window_hours"` - RetryCount int `json:"retry_count"` - HitRetentionDays int `json:"hit_retention_days"` - NonHitRetentionDays int `json:"non_hit_retention_days"` - PreHashCheckEnabled bool `json:"pre_hash_check_enabled"` - BlockedKeywords []string `json:"blocked_keywords"` - KeywordBlockingMode string `json:"keyword_blocking_mode"` - ModelFilter ContentModerationModelFilter `json:"model_filter"` + Enabled bool `json:"enabled"` + Mode string `json:"mode"` + BaseURL string `json:"base_url"` + Model string `json:"model"` + APIKeyConfigured bool `json:"api_key_configured"` + APIKeyMasked string `json:"api_key_masked"` + APIKeyCount int `json:"api_key_count"` + APIKeyMasks []string `json:"api_key_masks"` + APIKeyStatuses []ContentModerationAPIKeyStatus `json:"api_key_statuses"` + TimeoutMS int `json:"timeout_ms"` + SampleRate int `json:"sample_rate"` + AllGroups bool `json:"all_groups"` + GroupIDs []int64 `json:"group_ids"` + RecordNonHits bool `json:"record_non_hits"` + Thresholds map[string]float64 `json:"thresholds"` + WorkerCount int `json:"worker_count"` + QueueSize int `json:"queue_size"` + BlockStatus int `json:"block_status"` + BlockMessage string `json:"block_message"` + EmailOnHit bool `json:"email_on_hit"` + AutoBanEnabled bool `json:"auto_ban_enabled"` + BanThreshold int `json:"ban_threshold"` + ViolationWindowHours int `json:"violation_window_hours"` + RetryCount int `json:"retry_count"` + HitRetentionDays int `json:"hit_retention_days"` + NonHitRetentionDays int `json:"non_hit_retention_days"` + PreHashCheckEnabled bool `json:"pre_hash_check_enabled"` + BlockedKeywords []string `json:"blocked_keywords"` + KeywordBlockingMode string `json:"keyword_blocking_mode"` + ModelFilter ContentModerationModelFilter `json:"model_filter"` + CyberPolicyExcludeFromBanCount bool `json:"cyber_policy_exclude_from_ban_count"` } type ContentModerationAPIKeyStatus struct { @@ -250,36 +256,37 @@ type ContentModerationTestAuditResult struct { } type UpdateContentModerationConfigInput struct { - Enabled *bool `json:"enabled"` - Mode *string `json:"mode"` - BaseURL *string `json:"base_url"` - Model *string `json:"model"` - APIKey *string `json:"api_key"` - APIKeys *[]string `json:"api_keys"` - APIKeysMode string `json:"api_keys_mode"` - DeleteAPIKeyHashes *[]string `json:"delete_api_key_hashes"` - ClearAPIKey bool `json:"clear_api_key"` - TimeoutMS *int `json:"timeout_ms"` - SampleRate *int `json:"sample_rate"` - AllGroups *bool `json:"all_groups"` - GroupIDs *[]int64 `json:"group_ids"` - RecordNonHits *bool `json:"record_non_hits"` - Thresholds *map[string]float64 `json:"thresholds"` - WorkerCount *int `json:"worker_count"` - QueueSize *int `json:"queue_size"` - BlockStatus *int `json:"block_status"` - BlockMessage *string `json:"block_message"` - EmailOnHit *bool `json:"email_on_hit"` - AutoBanEnabled *bool `json:"auto_ban_enabled"` - BanThreshold *int `json:"ban_threshold"` - ViolationWindowHours *int `json:"violation_window_hours"` - RetryCount *int `json:"retry_count"` - HitRetentionDays *int `json:"hit_retention_days"` - NonHitRetentionDays *int `json:"non_hit_retention_days"` - PreHashCheckEnabled *bool `json:"pre_hash_check_enabled"` - BlockedKeywords *[]string `json:"blocked_keywords"` - KeywordBlockingMode *string `json:"keyword_blocking_mode"` - ModelFilter *ContentModerationModelFilter `json:"model_filter"` + Enabled *bool `json:"enabled"` + Mode *string `json:"mode"` + BaseURL *string `json:"base_url"` + Model *string `json:"model"` + APIKey *string `json:"api_key"` + APIKeys *[]string `json:"api_keys"` + APIKeysMode string `json:"api_keys_mode"` + DeleteAPIKeyHashes *[]string `json:"delete_api_key_hashes"` + ClearAPIKey bool `json:"clear_api_key"` + TimeoutMS *int `json:"timeout_ms"` + SampleRate *int `json:"sample_rate"` + AllGroups *bool `json:"all_groups"` + GroupIDs *[]int64 `json:"group_ids"` + RecordNonHits *bool `json:"record_non_hits"` + Thresholds *map[string]float64 `json:"thresholds"` + WorkerCount *int `json:"worker_count"` + QueueSize *int `json:"queue_size"` + BlockStatus *int `json:"block_status"` + BlockMessage *string `json:"block_message"` + EmailOnHit *bool `json:"email_on_hit"` + AutoBanEnabled *bool `json:"auto_ban_enabled"` + BanThreshold *int `json:"ban_threshold"` + ViolationWindowHours *int `json:"violation_window_hours"` + RetryCount *int `json:"retry_count"` + HitRetentionDays *int `json:"hit_retention_days"` + NonHitRetentionDays *int `json:"non_hit_retention_days"` + PreHashCheckEnabled *bool `json:"pre_hash_check_enabled"` + BlockedKeywords *[]string `json:"blocked_keywords"` + KeywordBlockingMode *string `json:"keyword_blocking_mode"` + ModelFilter *ContentModerationModelFilter `json:"model_filter"` + CyberPolicyExcludeFromBanCount *bool `json:"cyber_policy_exclude_from_ban_count"` } type ContentModerationModelFilter struct { @@ -461,8 +468,12 @@ type ContentModerationClearHashesResult struct { type ContentModerationRepository interface { CreateLog(ctx context.Context, log *ContentModerationLog) 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) + // UpdateLogEmailSent 回写邮件发送结果(F7:CreateLog 先行后补 EmailSent)。 + UpdateLogEmailSent(ctx context.Context, id int64, sent bool) error } type ContentModerationHashCache interface { @@ -648,6 +659,9 @@ func (s *ContentModerationService) UpdateConfig(ctx context.Context, input Updat if input.RecordNonHits != nil { cfg.RecordNonHits = *input.RecordNonHits } + if input.CyberPolicyExcludeFromBanCount != nil { + cfg.CyberPolicyExcludeFromBanCount = *input.CyberPolicyExcludeFromBanCount + } if input.Thresholds != nil { cfg.Thresholds = mergeContentModerationThresholds(ContentModerationDefaultThresholds(), *input.Thresholds) } @@ -1644,7 +1658,7 @@ func (s *ContentModerationService) applyFlaggedAccountSideEffects(ctx context.Co count := 1 if s.repo != nil && cfg.ViolationWindowHours > 0 { 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 } } @@ -1835,6 +1849,7 @@ func defaultContentModerationConfig() *ContentModerationConfig { Type: ContentModerationModelFilterAll, Models: []string{}, }, + CyberPolicyExcludeFromBanCount: false, } } @@ -2131,36 +2146,37 @@ func (s *ContentModerationService) configView(cfg *ContentModerationConfig) *Con apiKeyMasked = masks[0] } return &ContentModerationConfigView{ - Enabled: cfg.Enabled, - Mode: cfg.Mode, - BaseURL: cfg.BaseURL, - Model: cfg.Model, - APIKeyConfigured: len(keys) > 0, - APIKeyMasked: apiKeyMasked, - APIKeyCount: len(keys), - APIKeyMasks: masks, - APIKeyStatuses: s.apiKeyStatuses(keys), - TimeoutMS: cfg.TimeoutMS, - SampleRate: cfg.SampleRate, - AllGroups: cfg.AllGroups, - GroupIDs: append([]int64(nil), cfg.GroupIDs...), - RecordNonHits: cfg.RecordNonHits, - Thresholds: cloneFloatMap(cfg.Thresholds), - WorkerCount: cfg.WorkerCount, - QueueSize: cfg.QueueSize, - BlockStatus: cfg.BlockStatus, - BlockMessage: cfg.BlockMessage, - EmailOnHit: cfg.EmailOnHit, - AutoBanEnabled: cfg.AutoBanEnabled, - BanThreshold: cfg.BanThreshold, - ViolationWindowHours: cfg.ViolationWindowHours, - RetryCount: cfg.RetryCount, - HitRetentionDays: cfg.HitRetentionDays, - NonHitRetentionDays: cfg.NonHitRetentionDays, - PreHashCheckEnabled: cfg.PreHashCheckEnabled, - BlockedKeywords: append([]string(nil), cfg.BlockedKeywords...), - KeywordBlockingMode: cfg.KeywordBlockingMode, - ModelFilter: cloneContentModerationModelFilter(cfg.ModelFilter), + Enabled: cfg.Enabled, + Mode: cfg.Mode, + BaseURL: cfg.BaseURL, + Model: cfg.Model, + APIKeyConfigured: len(keys) > 0, + APIKeyMasked: apiKeyMasked, + APIKeyCount: len(keys), + APIKeyMasks: masks, + APIKeyStatuses: s.apiKeyStatuses(keys), + TimeoutMS: cfg.TimeoutMS, + SampleRate: cfg.SampleRate, + AllGroups: cfg.AllGroups, + GroupIDs: append([]int64(nil), cfg.GroupIDs...), + RecordNonHits: cfg.RecordNonHits, + Thresholds: cloneFloatMap(cfg.Thresholds), + WorkerCount: cfg.WorkerCount, + QueueSize: cfg.QueueSize, + BlockStatus: cfg.BlockStatus, + BlockMessage: cfg.BlockMessage, + EmailOnHit: cfg.EmailOnHit, + AutoBanEnabled: cfg.AutoBanEnabled, + BanThreshold: cfg.BanThreshold, + ViolationWindowHours: cfg.ViolationWindowHours, + RetryCount: cfg.RetryCount, + HitRetentionDays: cfg.HitRetentionDays, + NonHitRetentionDays: cfg.NonHitRetentionDays, + PreHashCheckEnabled: cfg.PreHashCheckEnabled, + BlockedKeywords: append([]string(nil), cfg.BlockedKeywords...), + KeywordBlockingMode: cfg.KeywordBlockingMode, + ModelFilter: cloneContentModerationModelFilter(cfg.ModelFilter), + CyberPolicyExcludeFromBanCount: cfg.CyberPolicyExcludeFromBanCount, } } @@ -2688,3 +2704,135 @@ func maskSecretTail(secret string) string { } 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)) +} diff --git a/backend/internal/service/content_moderation_cyber_test.go b/backend/internal/service/content_moderation_cyber_test.go new file mode 100644 index 0000000000..07ef2812d3 --- /dev/null +++ b/backend/internal/service/content_moderation_cyber_test.go @@ -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, "默认路径行为不变(现状回归)") +} diff --git a/backend/internal/service/content_moderation_email.go b/backend/internal/service/content_moderation_email.go index e462ff88f0..501f4341a1 100644 --- a/backend/internal/service/content_moderation_email.go +++ b/backend/internal/service/content_moderation_email.go @@ -115,3 +115,41 @@ func defaultContentModerationString(value string, fallback string) string { } 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(` + +
+
+
+
Risk Control / 网络安全策略
+

请求被网络安全策略拦截

+

尊敬的用户 %s,您的请求被上游网络安全策略(cyber policy)拦截。

+
+ + + + +
触发时间%s
模型%s
上游说明%s
+
+

如认为系误判,可调整请求措辞后重试,或申请获得授权的安全访问权限。

+

此邮件由 %s 自动发送,请勿回复。

+
+
+`, + 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), + ) +} diff --git a/backend/internal/service/content_moderation_test.go b/backend/internal/service/content_moderation_test.go index 9cfdc1e4b8..3336745647 100644 --- a/backend/internal/service/content_moderation_test.go +++ b/backend/internal/service/content_moderation_test.go @@ -94,7 +94,7 @@ func (r *contentModerationTestRepo) ListLogs(ctx context.Context, filter Content 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() defer r.mu.Unlock() 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 { continue } + if excludeCyberPolicy && log.Action == ContentModerationActionCyberPolicy { + continue + } if log.CreatedAt.IsZero() || log.CreatedAt.Before(since) { continue } @@ -114,6 +117,10 @@ func (r *contentModerationTestRepo) CleanupExpiredLogs(ctx context.Context, hitB return &ContentModerationCleanupResult{}, nil } +func (r *contentModerationTestRepo) UpdateLogEmailSent(ctx context.Context, id int64, sent bool) error { + return nil +} + func (r *contentModerationTestRepo) snapshotLogs() []ContentModerationLog { r.mu.Lock() defer r.mu.Unlock() @@ -1789,3 +1796,44 @@ func TestContentModerationUnbanUser_ActiveUserOnlyInvalidatesAuthCache(t *testin func contentModerationIntPtr(v int) *int { 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) +} diff --git a/backend/internal/service/domain_constants.go b/backend/internal/service/domain_constants.go index 11245d0016..54e32b8bd4 100644 --- a/backend/internal/service/domain_constants.go +++ b/backend/internal/service/domain_constants.go @@ -136,6 +136,8 @@ const ( SettingKeyAffiliateRebatePerInviteeCap = "affiliate_rebate_per_invitee_cap" // 单人返利上限(0=无上限) SettingKeyRiskControlEnabled = "risk_control_enabled" // 是否启用风控中心入口与审计链路 SettingKeyContentModerationConfig = "content_moderation_config" // 内容审计配置(JSON) + SettingKeyCyberSessionBlockEnabled = "cyber_session_block_enabled" // cyber 命中后会话级自动屏蔽总开关(默认关) + SettingKeyCyberSessionBlockTTLSeconds = "cyber_session_block_ttl_seconds" // 会话屏蔽 TTL 秒数(默认 3600) SettingKeyLoginAgreementEnabled = "login_agreement_enabled" // 登录前是否要求同意条款 SettingKeyLoginAgreementMode = "login_agreement_mode" // 条款确认展示模式:modal / checkbox SettingKeyLoginAgreementUpdatedAt = "login_agreement_updated_at" // 条款更新日期(展示用) diff --git a/backend/internal/service/notification_email_service.go b/backend/internal/service/notification_email_service.go index d21363b6e6..b7acccc7aa 100644 --- a/backend/internal/service/notification_email_service.go +++ b/backend/internal/service/notification_email_service.go @@ -30,6 +30,7 @@ const ( NotificationEmailEventAccountQuotaAlert = "account.quota_alert" NotificationEmailEventContentModerationViolation = "content_moderation.violation_notice" NotificationEmailEventContentModerationDisabled = "content_moderation.account_disabled" + NotificationEmailEventCyberPolicyNotice = "content_moderation.cyber_policy_notice" NotificationEmailEventOpsAlert = "ops.alert" NotificationEmailEventOpsScheduledReport = "ops.scheduled_report" @@ -947,6 +948,7 @@ var notificationEmailEventOrder = []string{ NotificationEmailEventAccountQuotaAlert, NotificationEmailEventContentModerationViolation, NotificationEmailEventContentModerationDisabled, + NotificationEmailEventCyberPolicyNotice, NotificationEmailEventOpsAlert, NotificationEmailEventOpsScheduledReport, } @@ -1035,6 +1037,15 @@ var notificationEmailEventDefinitions = map[string]NotificationEmailEventInfo{ Placeholders: append(append([]string{}, notificationEmailCommonPlaceholders...), "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: { Event: NotificationEmailEventOpsAlert, Label: "Ops alert", @@ -1277,6 +1288,34 @@ var notificationEmailOfficialTemplates = map[string]map[string]notificationEmail

如需申诉或恢复账号,请联系平台管理员处理。

`), }, }, + NotificationEmailEventCyberPolicyNotice: { + notificationEmailDefaultLocale: { + Subject: "[{{site_name}}] Cyber-security policy notice", + HTML: notificationEmailCard("#ef4444", "Cyber-security policy notice", ` +

Hello {{recipient_name}},

+

Your request was blocked by the upstream provider's cyber-security policy.

+ + + + + +
Triggered at{{triggered_at}}
Model{{model}}
Group{{group_name}}
Upstream message{{upstream_message}}
+

If you believe this is a mistake, try rephrasing your request, or apply for authorized security access.

`), + }, + notificationEmailLocaleChinese: { + Subject: "[{{site_name}}] 网络安全策略拦截提醒", + HTML: notificationEmailCard("#ef4444", "网络安全策略拦截提醒", ` +

{{recipient_name}},您好:

+

您的请求被上游服务商的网络安全策略(cyber policy)拦截。

+ + + + + +
触发时间{{triggered_at}}
模型{{model}}
所属分组{{group_name}}
上游说明{{upstream_message}}
+

如认为系误判,可调整请求措辞后重试,或申请获得授权的安全访问权限。

`), + }, + }, NotificationEmailEventOpsAlert: { notificationEmailDefaultLocale: { Subject: "[Ops Alert][{{severity}}] {{rule_name}}", diff --git a/backend/internal/service/notification_email_service_test.go b/backend/internal/service/notification_email_service_test.go index 38987f6f5a..9faeda7993 100644 --- a/backend/internal/service/notification_email_service_test.go +++ b/backend/internal/service/notification_email_service_test.go @@ -140,6 +140,7 @@ func TestNotificationEmailAdditionalEventsAreListedAndPreviewable(t *testing.T) {NotificationEmailEventAccountQuotaAlert, "account_name"}, {NotificationEmailEventContentModerationViolation, "moderation_category"}, {NotificationEmailEventContentModerationDisabled, "violation_count"}, + {NotificationEmailEventCyberPolicyNotice, "upstream_message"}, {NotificationEmailEventOpsAlert, "rule_name"}, {NotificationEmailEventOpsScheduledReport, "report_html"}, } diff --git a/backend/internal/service/openai_cyber_policy.go b/backend/internal/service/openai_cyber_policy.go new file mode 100644 index 0000000000..524bade30c --- /dev/null +++ b/backend/internal/service/openai_cyber_policy.go @@ -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) +} diff --git a/backend/internal/service/openai_cyber_policy_test.go b/backend/internal/service/openai_cyber_policy_test.go new file mode 100644 index 0000000000..1ee4805481 --- /dev/null +++ b/backend/internal/service/openai_cyber_policy_test.go @@ -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) + } + }) + } +} diff --git a/backend/internal/service/openai_cyber_session_block.go b/backend/internal/service/openai_cyber_session_block.go new file mode 100644 index 0000000000..ebca8b1967 --- /dev/null +++ b/backend/internal/service/openai_cyber_session_block.go @@ -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 +} diff --git a/backend/internal/service/openai_cyber_session_block_test.go b/backend/internal/service/openai_cyber_session_block_test.go new file mode 100644 index 0000000000..3957b824e5 --- /dev/null +++ b/backend/internal/service/openai_cyber_session_block_test.go @@ -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")) +} diff --git a/backend/internal/service/openai_gateway_chat_completions.go b/backend/internal/service/openai_gateway_chat_completions.go index 8b75097349..a1051afe5e 100644 --- a/backend/internal/service/openai_gateway_chat_completions.go +++ b/backend/internal/service/openai_gateway_chat_completions.go @@ -312,6 +312,15 @@ func (s *OpenAIGatewayService) ForwardAsChatCompletions( 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 if handleErr == nil && result != nil { if responsesReq.ServiceTier != "" { @@ -412,6 +421,24 @@ func (s *OpenAIGatewayService) handleChatBufferedStreamingResponse( } if strings.TrimSpace(finalResponse.Status) == "failed" { 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)) } @@ -550,8 +577,40 @@ func (s *OpenAIGatewayService) handleChatStreamingResponse( if strings.TrimSpace(event.Type) == "response.failed" { payloadBytes := []byte(payload) message := extractOpenAISSEErrorMessage(payloadBytes) - streamFailoverErr = s.newOpenAIStreamFailoverError(c, account, false, requestID, payloadBytes, message) - return true + 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) + return true + } } chunks := apicompat.ResponsesEventToChatChunks(&event, state) @@ -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" +} diff --git a/backend/internal/service/openai_gateway_chat_completions_test.go b/backend/internal/service/openai_gateway_chat_completions_test.go index 779565037d..071308a57a 100644 --- a/backend/internal/service/openai_gateway_chat_completions_test.go +++ b/backend/internal/service/openai_gateway_chat_completions_test.go @@ -316,6 +316,51 @@ func TestForwardAsChatCompletions_StreamResponseFailedTriggersFailoverBeforeFlus 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) { gin.SetMode(gin.TestMode) @@ -773,3 +818,14 @@ func TestForwardAsChatCompletions_UpstreamRequestIgnoresClientCancel(t *testing. require.NotNil(t, upstream.lastReq) 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()) +} diff --git a/backend/internal/service/openai_gateway_compat_cyber_test.go b/backend/internal/service/openai_gateway_compat_cyber_test.go new file mode 100644 index 0000000000..c416de5ba7 --- /dev/null +++ b/backend/internal/service/openai_gateway_compat_cyber_test.go @@ -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") +} diff --git a/backend/internal/service/openai_gateway_messages.go b/backend/internal/service/openai_gateway_messages.go index 99653fec45..7d07fcc808 100644 --- a/backend/internal/service/openai_gateway_messages.go +++ b/backend/internal/service/openai_gateway_messages.go @@ -366,6 +366,15 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic( 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 if handleErr == nil && result != nil { if compatContinuationEnabled && promptCacheKey != "" && result.ResponseID != "" { @@ -444,6 +453,28 @@ func (s *OpenAIGatewayService) handleAnthropicBufferedStreamingResponse( 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 // accumulated delta events so the client receives the full content. acc.SupplementResponseOutput(finalResponse) @@ -770,6 +801,32 @@ func (s *OpenAIGatewayService) handleAnthropicStreamingResponse( if event.Usage != nil { 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 @@ -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 { if usage == nil { return OpenAIUsage{} diff --git a/backend/internal/service/openai_gateway_record_usage_test.go b/backend/internal/service/openai_gateway_record_usage_test.go index 318c086113..1683fa4279 100644 --- a/backend/internal/service/openai_gateway_record_usage_test.go +++ b/backend/internal/service/openai_gateway_record_usage_test.go @@ -58,6 +58,75 @@ func TestOpenAIGatewayServiceRecordUsage_RejectsNilInput(t *testing.T) { 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 { UserRepository @@ -1761,6 +1830,30 @@ func TestGatewayServiceCalculateRecordUsageCost_ChannelImageBillingUsesSizeTier( 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) { groupID := int64(128) defaultPrice := 0.10 diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go index 6374ee91bd..e0cf567a9a 100644 --- a/backend/internal/service/openai_gateway_service.go +++ b/backend/internal/service/openai_gateway_service.go @@ -3555,6 +3555,19 @@ func (s *OpenAIGatewayService) handleErrorResponsePassthrough( MarkResponseCommitted(c) 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 = sanitizeUpstreamErrorMessage(upstreamMsg) upstreamDetail := "" @@ -3568,9 +3581,11 @@ func (s *OpenAIGatewayService) handleErrorResponsePassthrough( setOpsUpstreamError(c, resp.StatusCode, upstreamMsg, upstreamDetail) logOpenAIInstructionsRequiredDebug(ctx, c, account, resp.StatusCode, upstreamMsg, requestBody, body) // 透传模式保留原始上游错误响应,但运行态账号状态仍需更新, - // 避免粘性路由继续复用刚被限流的账号。 - reqModel, _, _ := extractOpenAIRequestMetaFromBody(requestBody) - _ = s.handleOpenAIAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, body, reqModel) + // 避免粘性路由继续复用刚被限流的账号。cyber 例外:不冷却账号。 + if !cyberHit { + reqModel, _, _ := extractOpenAIRequestMetaFromBody(requestBody) + _ = s.handleOpenAIAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, body, reqModel) + } appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ Platform: account.Platform, AccountID: account.ID, @@ -3852,7 +3867,20 @@ func (s *OpenAIGatewayService) handleStreamingResponsePassthrough( eventType := strings.TrimSpace(gjson.Get(trimmedData, "type").String()) if eventType == "response.failed" { 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(), s.newOpenAIStreamFailoverError(c, account, true, upstreamRequestID, dataBytes, failedMessage) } @@ -4257,6 +4285,29 @@ func (s *OpenAIGatewayService) handleErrorResponse( ) (*OpenAIForwardResult, error) { 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 = sanitizeUpstreamErrorMessage(upstreamMsg) upstreamDetail := "" @@ -4422,6 +4473,29 @@ func (s *OpenAIGatewayService) handleCompatErrorResponse( ) (*OpenAIForwardResult, error) { 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)) if upstreamMsg == "" { upstreamMsg = fmt.Sprintf("Upstream error: %d", resp.StatusCode) @@ -4744,7 +4818,20 @@ func (s *OpenAIGatewayService) handleStreamingResponse(ctx context.Context, resp forceFlushFailedEvent := false if eventType == "response.failed" { 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 streamFailoverErr = s.newOpenAIStreamFailoverError(c, account, false, upstreamRequestID, dataBytes, failedMessage) return @@ -5731,9 +5818,72 @@ type OpenAIRecordUsageInput struct { IPAddress string // 请求的客户端 IP 地址 RequestPayloadHash string APIKeyService APIKeyQuotaUpdater + // CyberBlocked 为 true 时把该用量行标记为 cyber(request_type=cyber),计费逻辑不变。 + CyberBlocked bool 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 func (s *OpenAIGatewayService) RecordUsage(ctx context.Context, input *OpenAIRecordUsageInput) error { if input == nil { @@ -5888,6 +6038,9 @@ func (s *OpenAIGatewayService) RecordUsage(ctx context.Context, input *OpenAIRec usageLog.AccountRateMultiplier = &accountRateMultiplier usageLog.BillingType = billingType usageLog.Stream = result.Stream + if input.CyberBlocked { + usageLog.RequestType = RequestTypeCyberBlocked + } usageLog.OpenAIWSMode = result.OpenAIWSMode usageLog.DurationMs = &durationMs usageLog.FirstTokenMs = result.FirstTokenMs diff --git a/backend/internal/service/openai_gateway_service_test.go b/backend/internal/service/openai_gateway_service_test.go index 9866ddeff4..51c528cb88 100644 --- a/backend/internal/service/openai_gateway_service_test.go +++ b/backend/internal/service/openai_gateway_service_test.go @@ -2668,3 +2668,109 @@ func TestOpenAICompatSSEFrameParserResetsEventTypeAtFrameBoundary(t *testing.T) require.Empty(t, frame.EventType) 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)) +} diff --git a/backend/internal/service/openai_image_generation_controls_test.go b/backend/internal/service/openai_image_generation_controls_test.go index 9ff8b5107b..ce66d6722b 100644 --- a/backend/internal/service/openai_image_generation_controls_test.go +++ b/backend/internal/service/openai_image_generation_controls_test.go @@ -320,6 +320,34 @@ func TestOpenAIGatewayServiceHandleResponsesImageOutputs_Streaming(t *testing.T) 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 { cfg := &config.Config{} return &OpenAIGatewayService{ diff --git a/backend/internal/service/openai_ws_forwarder.go b/backend/internal/service/openai_ws_forwarder.go index 43538e811a..f45dcfc694 100644 --- a/backend/internal/service/openai_ws_forwarder.go +++ b/backend/internal/service/openai_ws_forwarder.go @@ -2237,6 +2237,19 @@ func (s *OpenAIGatewayService) forwardOpenAIWSV2( } 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" { errCodeRaw, errTypeRaw, errMsgRaw := parseOpenAIWSErrorEventFields(message) s.persistOpenAIWSRateLimitSignal(ctx, account, lease.HandshakeHeaders(), message, errCodeRaw, errTypeRaw, errMsgRaw) @@ -3210,6 +3223,19 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( } 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 needModelReplace && len(mappedModelBytes) > 0 && openAIWSEventMayContainModel(eventType) && bytes.Contains(upstreamMessage, mappedModelBytes) { upstreamMessage = replaceOpenAIWSMessageModel(upstreamMessage, mappedModel, originalModel) diff --git a/backend/internal/service/openai_ws_forwarder_test.go b/backend/internal/service/openai_ws_forwarder_test.go index d817de6e89..b9944e92fa 100644 --- a/backend/internal/service/openai_ws_forwarder_test.go +++ b/backend/internal/service/openai_ws_forwarder_test.go @@ -1,8 +1,10 @@ package service import ( + "net/http/httptest" "testing" + "github.com/gin-gonic/gin" "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 事件集合与终止事件集合互斥」的不变量。 // firstTokenMs 的计算依赖于 isTokenEvent && !isTerminalEvent; // 若两者再次出现交集,则 issue #2651 描述的 latency 误报会重现。 diff --git a/backend/internal/service/ops_user_error.go b/backend/internal/service/ops_user_error.go index 817c139afe..7dd128afa7 100644 --- a/backend/internal/service/ops_user_error.go +++ b/backend/internal/service/ops_user_error.go @@ -48,6 +48,8 @@ func MapUserErrorCategory(phase, errType string) string { return "quota" case "invalid_request_error": return "invalid_request" + case "cyber_policy": + return "cyber" } } return "other" @@ -72,6 +74,8 @@ func CategoryToFilter(category string) (phases []string, errorTypes []string) { return nil, []string{"billing_error", "subscription_error"} case "invalid_request": return nil, []string{"invalid_request_error"} + case "cyber": + return []string{"request"}, []string{"cyber_policy"} default: return nil, nil } diff --git a/backend/internal/service/ops_user_error_cyber_test.go b/backend/internal/service/ops_user_error_cyber_test.go new file mode 100644 index 0000000000..e11f0aaa3c --- /dev/null +++ b/backend/internal/service/ops_user_error_cyber_test.go @@ -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) +} diff --git a/backend/internal/service/setting_service.go b/backend/internal/service/setting_service.go index 7043736a80..7179557228 100644 --- a/backend/internal/service/setting_service.go +++ b/backend/internal/service/setting_service.go @@ -157,6 +157,18 @@ const openAIAllowCodexPluginCacheTTL = 60 * time.Second const openAIAllowCodexPluginErrorTTL = 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 openAIQuotaAutoPauseSettingsErrorTTL = 5 * time.Second const openAIQuotaAutoPauseSettingsDBTimeout = 5 * time.Second @@ -188,6 +200,9 @@ type SettingService struct { openAIAllowCodexPluginCache atomic.Value // *cachedOpenAIAllowCodexPlugin openAIAllowCodexPluginSF singleflight.Group + cyberSessionBlockRuntimeCache atomic.Value // *cachedCyberSessionBlockRuntime + cyberSessionBlockRuntimeSF singleflight.Group + // openAIQuotaAutoPauseSettingsCache holds the most recently observed quota auto-pause // settings. GetOpenAIQuotaAutoPauseSettings reads this atomic.Value on the request hot // 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 } +// 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 获取公开设置(无需登录) func (s *SettingService) GetPublicSettings(ctx context.Context) (*PublicSettings, error) { keys := []string{ @@ -1895,6 +1966,12 @@ func (s *SettingService) buildSystemSettingsUpdates(ctx context.Context, setting // 风控中心功能开关 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 updates[SettingKeyMinClaudeCodeVersion] = settings.MinClaudeCodeVersion updates[SettingKeyMaxClaudeCodeVersion] = settings.MaxClaudeCodeVersion @@ -2819,6 +2896,10 @@ func (s *SettingService) InitializeDefaultSettings(ctx context.Context) error { // 风控中心功能(默认关闭,显式启用) SettingKeyRiskControlEnabled: "false", + // cyber 会话屏蔽(默认关闭,TTL 默认 3600s) + SettingKeyCyberSessionBlockEnabled: "false", + SettingKeyCyberSessionBlockTTLSeconds: "3600", + // Claude Code version check (default: empty = disabled) SettingKeyMinClaudeCodeVersion: "", SettingKeyMaxClaudeCodeVersion: "", @@ -3328,6 +3409,14 @@ func (s *SettingService) parseSettings(settings map[string]string) *SystemSettin // 风控中心功能(默认关闭,严格 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 result.MinClaudeCodeVersion = settings[SettingKeyMinClaudeCodeVersion] result.MaxClaudeCodeVersion = settings[SettingKeyMaxClaudeCodeVersion] diff --git a/backend/internal/service/settings_view.go b/backend/internal/service/settings_view.go index 9721720213..25d89b2d7b 100644 --- a/backend/internal/service/settings_view.go +++ b/backend/internal/service/settings_view.go @@ -145,6 +145,8 @@ type SystemSettings struct { DefaultConcurrency int DefaultBalance float64 RiskControlEnabled bool + CyberSessionBlockEnabled bool + CyberSessionBlockTTLSeconds int AffiliateEnabled bool AffiliateRebateRate float64 AffiliateRebateFreezeHours int diff --git a/backend/internal/service/usage_log.go b/backend/internal/service/usage_log.go index d63f47ccf0..66cdb62b2b 100644 --- a/backend/internal/service/usage_log.go +++ b/backend/internal/service/usage_log.go @@ -14,15 +14,16 @@ const ( type RequestType int16 const ( - RequestTypeUnknown RequestType = 0 - RequestTypeSync RequestType = 1 - RequestTypeStream RequestType = 2 - RequestTypeWSV2 RequestType = 3 + RequestTypeUnknown RequestType = 0 + RequestTypeSync RequestType = 1 + RequestTypeStream RequestType = 2 + RequestTypeWSV2 RequestType = 3 + RequestTypeCyberBlocked RequestType = 4 // cyber_policy 命中(透传但被上游安全策略拒绝) ) func (t RequestType) IsValid() bool { switch t { - case RequestTypeUnknown, RequestTypeSync, RequestTypeStream, RequestTypeWSV2: + case RequestTypeUnknown, RequestTypeSync, RequestTypeStream, RequestTypeWSV2, RequestTypeCyberBlocked: return true default: return false @@ -44,6 +45,8 @@ func (t RequestType) String() string { return "stream" case RequestTypeWSV2: return "ws_v2" + case RequestTypeCyberBlocked: + return "cyber" default: return "unknown" } @@ -63,8 +66,10 @@ func ParseUsageRequestType(value string) (RequestType, error) { return RequestTypeStream, nil case "ws_v2": return RequestTypeWSV2, nil + case "cyber": + return RequestTypeCyberBlocked, nil 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") } } diff --git a/backend/internal/service/usage_log_cyber_test.go b/backend/internal/service/usage_log_cyber_test.go new file mode 100644 index 0000000000..a6c6e845ad --- /dev/null +++ b/backend/internal/service/usage_log_cyber_test.go @@ -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 字段") +} diff --git a/frontend/src/api/admin/riskControl.ts b/frontend/src/api/admin/riskControl.ts index aefd1618b1..5cce4c1827 100644 --- a/frontend/src/api/admin/riskControl.ts +++ b/frontend/src/api/admin/riskControl.ts @@ -40,6 +40,7 @@ export interface ContentModerationConfig { blocked_keywords: string[] keyword_blocking_mode: KeywordBlockingMode model_filter: ContentModerationModelFilter + cyber_policy_exclude_from_ban_count: boolean } export type ContentModerationAPIKeyStatusValue = 'unknown' | 'ok' | 'error' | 'frozen' @@ -115,6 +116,7 @@ export interface UpdateContentModerationConfig { blocked_keywords?: string[] keyword_blocking_mode?: KeywordBlockingMode model_filter?: ContentModerationModelFilter + cyber_policy_exclude_from_ban_count?: boolean } export interface ContentModerationRuntimeStatus { diff --git a/frontend/src/api/admin/settings.ts b/frontend/src/api/admin/settings.ts index 5be6307609..21e712056e 100644 --- a/frontend/src/api/admin/settings.ts +++ b/frontend/src/api/admin/settings.ts @@ -566,6 +566,11 @@ export interface SystemSettings { // Payment configuration payment_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_max_amount: number; payment_daily_limit: number; @@ -800,6 +805,11 @@ export interface UpdateSettingsRequest { // Payment configuration payment_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_max_amount?: number; payment_daily_limit?: number; diff --git a/frontend/src/components/admin/usage/UsageFilters.vue b/frontend/src/components/admin/usage/UsageFilters.vue index a800f1903d..f16012ccff 100644 --- a/frontend/src/components/admin/usage/UsageFilters.vue +++ b/frontend/src/components/admin/usage/UsageFilters.vue @@ -233,7 +233,8 @@ const requestTypeOptions = ref([ { value: null, label: t('admin.usage.allTypes') }, { value: 'ws_v2', label: t('usage.ws') }, { 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([ diff --git a/frontend/src/components/admin/usage/UsageTable.vue b/frontend/src/components/admin/usage/UsageTable.vue index 93d3d8b88e..1103067c59 100644 --- a/frontend/src/components/admin/usage/UsageTable.vue +++ b/frontend/src/components/admin/usage/UsageTable.vue @@ -475,6 +475,7 @@ const tokenTooltipData = ref(null) const getRequestTypeLabel = (row: AdminUsageLog): string => { const requestType = resolveUsageRequestType(row) + if (requestType === 'cyber') return t('usage.cyber') if (requestType === 'ws_v2') return t('usage.ws') if (requestType === 'stream') return t('usage.stream') if (requestType === 'sync') return t('usage.sync') @@ -483,6 +484,7 @@ const getRequestTypeLabel = (row: AdminUsageLog): string => { const getRequestTypeBadgeClass = (row: AdminUsageLog): string => { 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 === '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' diff --git a/frontend/src/components/user/UserErrorRequestsTable.vue b/frontend/src/components/user/UserErrorRequestsTable.vue index 4bbb28cdf8..42cb93ce04 100644 --- a/frontend/src/components/user/UserErrorRequestsTable.vue +++ b/frontend/src/components/user/UserErrorRequestsTable.vue @@ -123,7 +123,7 @@ const localModel = ref('') const localCategory = ref('') const localApiKeyId = ref(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(() => [ { value: '', label: t('usage.errors.allCategories') }, diff --git a/frontend/src/i18n/locales/en.ts b/frontend/src/i18n/locales/en.ts index f71d62a19e..8dea4b58ab 100644 --- a/frontend/src/i18n/locales/en.ts +++ b/frontend/src/i18n/locales/en.ts @@ -923,6 +923,7 @@ export default { ws: 'WS', stream: 'Stream', sync: 'Sync', + cyber: 'Cyber', unknown: 'Unknown', in: 'In', out: 'Out', @@ -980,7 +981,7 @@ export default { categories: { auth: 'Auth failed', rate_limit: 'Rate limited', quota: 'Balance/Subscription', 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: { 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.', autoBan: 'Auto Ban User', 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', violationWindowHours: 'Count Window (hours)', hitRetentionDays: 'Hit Record Retention (days)', @@ -2771,6 +2775,7 @@ export default { action: { block: 'Blocked', keywordBlock: 'Keyword Blocked', + cyberPolicy: 'Cyber policy', error: 'Error', }, }, @@ -5438,6 +5443,9 @@ export default { configureLink: 'Configure content moderation in Risk Control', enabled: 'Enable Risk Control', 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: { title: 'Affiliate (Invite Rebate)', diff --git a/frontend/src/i18n/locales/zh.ts b/frontend/src/i18n/locales/zh.ts index 996b38d6d9..688dc95c35 100644 --- a/frontend/src/i18n/locales/zh.ts +++ b/frontend/src/i18n/locales/zh.ts @@ -927,6 +927,7 @@ export default { ws: 'WS', stream: '流式', sync: '同步', + cyber: '安全策略', unknown: '未知', in: '输入', out: '输出', @@ -984,7 +985,7 @@ export default { categories: { auth: '认证失败', rate_limit: '限流', quota: '余额/订阅', invalid_request: '参数错误', service_unavailable: '服务暂时不可用', - upstream: '上游错误', internal: '平台错误', other: '其他', + upstream: '上游错误', internal: '平台错误', other: '其他', cyber: '安全策略', }, detail: { title: '错误请求详情', @@ -2704,6 +2705,9 @@ export default { emailOnHitHint: '开启后每次达到阈值都会向用户发送风控提醒邮件;自动封禁通知始终发送。', autoBan: '自动封禁用户', autoBanHint: '命中次数达到阈值后将禁用用户账号、刷新认证缓存并发送封禁通知邮件。', + cyberPolicyExcludeBan: 'cyber_policy 不计入封号次数', + cyberPolicyExcludeBanHint: '开启后,cyber_policy 拦截不再计入自动封号的违规次数:当次不判定封号,历史累计亦排除。风控日志与通知邮件照常。', + violationNotCounted: '未计入封号', banThreshold: '封禁触发次数', violationWindowHours: '累计窗口(小时)', hitRetentionDays: '命中记录保留(天)', @@ -2848,6 +2852,7 @@ export default { action: { block: '拦截', keywordBlock: '关键词拦截', + cyberPolicy: '网络安全策略', error: '异常', }, }, @@ -5598,6 +5603,9 @@ export default { configureLink: '前往 风控中心 配置内容审计', enabled: '启用风控中心', enabledHint: '关闭后管理员侧边栏入口隐藏,网关内容审计不会执行。', + cyberSessionBlock: 'cyber 会话自动屏蔽', + cyberSessionBlockHint: '开启后,被上游网络安全策略(cyber_policy)拦截的会话将在 TTL 内被本地屏蔽,不再发往上游。仅屏蔽该会话,不影响同 Key 其他会话。', + cyberSessionBlockTTL: '屏蔽时长(秒)', }, affiliate: { title: '邀请返利', diff --git a/frontend/src/types/index.ts b/frontend/src/types/index.ts index 3fa9e0f573..04ae60b8ef 100644 --- a/frontend/src/types/index.ts +++ b/frontend/src/types/index.ts @@ -1203,7 +1203,7 @@ export interface CodexSessionImportResult { // ==================== Usage & Redeem Types ==================== 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 ImageSizeBreakdown = Record diff --git a/frontend/src/utils/usageRequestType.ts b/frontend/src/utils/usageRequestType.ts index bfdafb0750..14de5f043d 100644 --- a/frontend/src/utils/usageRequestType.ts +++ b/frontend/src/utils/usageRequestType.ts @@ -6,7 +6,7 @@ export interface UsageRequestTypeLike { openai_ws_mode?: boolean | null } -const VALID_REQUEST_TYPES = new Set(['unknown', 'sync', 'stream', 'ws_v2']) +const VALID_REQUEST_TYPES = new Set(['unknown', 'sync', 'stream', 'ws_v2', 'cyber']) export const isUsageRequestType = (value: unknown): value is 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 => { - if (!requestType || requestType === 'unknown') { + // cyber 与 stream 正交(cyber 可发生在 stream 或非 stream 请求),不映射到 legacy stream 维度。 + if (!requestType || requestType === 'unknown' || requestType === 'cyber') { return null } if (requestType === 'sync') { diff --git a/frontend/src/views/admin/RiskControlView.vue b/frontend/src/views/admin/RiskControlView.vue index b6d62767be..fbb33a2673 100644 --- a/frontend/src/views/admin/RiskControlView.vue +++ b/frontend/src/views/admin/RiskControlView.vue @@ -881,6 +881,13 @@ +
+
+

{{ t('admin.riskControl.cyberPolicyExcludeBan') }}

+

{{ t('admin.riskControl.cyberPolicyExcludeBanHint') }}

+
+ +
@@ -1226,6 +1233,7 @@ const configForm = reactive({ block_message: '内容审计命中风险规则,请调整输入后重试', email_on_hit: true, auto_ban_enabled: true, + cyber_policy_exclude_from_ban_count: false, ban_threshold: 10, violation_window_hours: 720, hit_retention_days: 180, @@ -1702,6 +1710,7 @@ function applyConfig(config: ContentModerationConfig) { configForm.block_message = config.block_message || '内容审计命中风险规则,请调整输入后重试' configForm.email_on_hit = config.email_on_hit ?? 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.violation_window_hours = config.violation_window_hours || 720 configForm.hit_retention_days = config.hit_retention_days || 180 @@ -1782,6 +1791,7 @@ async function saveConfig() { block_message: configForm.block_message || '内容审计命中风险规则,请调整输入后重试', email_on_hit: configForm.email_on_hit, 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, violation_window_hours: Number(configForm.violation_window_hours) || 720, hit_retention_days: Number(configForm.hit_retention_days) || 180, @@ -2097,6 +2107,7 @@ function modeDescription(mode: ModerationMode): 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 === 'block') return t('admin.riskControl.action.block') 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 { - 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.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' @@ -2303,6 +2314,7 @@ function parseBlockedKeywords(value: string): string[] { function violationCountText(row: ContentModerationLog): string { if (!row.flagged) return '-' + if (row.violation_count === 0) return t('admin.riskControl.violationNotCounted') return t('admin.riskControl.violationCount', { count: row.violation_count || 1 }) } diff --git a/frontend/src/views/admin/SettingsView.vue b/frontend/src/views/admin/SettingsView.vue index c3e0f2b965..74a1328667 100644 --- a/frontend/src/views/admin/SettingsView.vue +++ b/frontend/src/views/admin/SettingsView.vue @@ -5297,6 +5297,31 @@
+ +
+
+ +

+ {{ t('admin.settings.features.riskControl.cyberSessionBlockHint') }} +

+
+ +
+ +
+ + +
@@ -7045,6 +7070,8 @@ const form = reactive({ hide_ccs_import_button: false, payment_enabled: false, risk_control_enabled: false, + cyber_session_block_enabled: false, + cyber_session_block_ttl_seconds: 3600, payment_min_amount: 1, payment_max_amount: 10000, payment_daily_limit: 50000, @@ -8316,6 +8343,9 @@ async function saveSettings() { // Payment configuration payment_enabled: form.payment_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_max_amount: Number(form.payment_max_amount) || 0, payment_daily_limit: Number(form.payment_daily_limit) || 0, diff --git a/frontend/src/views/admin/UsageView.vue b/frontend/src/views/admin/UsageView.vue index 4fb2caacb8..3dc1b59a43 100644 --- a/frontend/src/views/admin/UsageView.vue +++ b/frontend/src/views/admin/UsageView.vue @@ -490,6 +490,7 @@ const cancelExport = () => exportAbortController?.abort() const openCleanupDialog = () => { cleanupDialogVisible.value = true } const getRequestTypeLabel = (log: AdminUsageLog): string => { const requestType = resolveUsageRequestType(log) + if (requestType === 'cyber') return t('usage.cyber') if (requestType === 'ws_v2') return t('usage.ws') if (requestType === 'stream') return t('usage.stream') if (requestType === 'sync') return t('usage.sync') diff --git a/frontend/src/views/user/UsageView.vue b/frontend/src/views/user/UsageView.vue index 49e779a8b8..c6ef768318 100644 --- a/frontend/src/views/user/UsageView.vue +++ b/frontend/src/views/user/UsageView.vue @@ -773,6 +773,7 @@ const formatUserAgent = (ua: string): string => { const getRequestTypeLabel = (log: UsageLog): string => { const requestType = resolveUsageRequestType(log) + if (requestType === 'cyber') return t('usage.cyber') if (requestType === 'ws_v2') return t('usage.ws') if (requestType === 'stream') return t('usage.stream') if (requestType === 'sync') return t('usage.sync') @@ -781,6 +782,7 @@ const getRequestTypeLabel = (log: UsageLog): string => { const getRequestTypeBadgeClass = (log: UsageLog): string => { 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 === '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' @@ -790,6 +792,7 @@ const getRequestTypeBadgeClass = (log: UsageLog): string => { const getRequestTypeExportText = (log: UsageLog): string => { const requestType = resolveUsageRequestType(log) + if (requestType === 'cyber') return 'Cyber' if (requestType === 'ws_v2') return 'WS' if (requestType === 'stream') return 'Stream' if (requestType === 'sync') return 'Sync'