feat(openai): cyber_policy 硬阻断全链路透传、审计与计费

上游对单次请求下发 error.code=cyber_policy 硬阻断时,网关在所有端点
(/v1/responses、/v1/chat/completions、/v1/messages、WebSocket)及流式/
非流式路径下,将该结果原样透传给客户端,绝不 failover、换号或同步拦截;
命中后异步完成审计与计费:

- 风控中心记录 cyber_policy 留痕并发送通知邮件,落库先于发信,SMTP 阻塞
  不影响留痕
- ops 错误请求记录,状态码对齐客户端实际接收(流式 200 / 非流式 400)
- 用量明细标记 request_type=cyber,按上游真实 token 计费,HTTP 与
  WebSocket 计费口径统一,零 token 命中不误扣
- 会话级自动屏蔽(管理员开关,默认关):命中的会话在可配 TTL 内本地拦截
  不再发往上游,仅屏蔽该会话不影响同 Key 其他会话
- 封号计数排除开关:可选让 cyber 命中不计入自动封号,命中当次不判定且
  历史行在违规计数中一并排除

WebSocket 多轮连接下 cyber 标记按 turn 生命周期管理,逐轮独立检测与记录;
透传的错误响应不被兜底逻辑追加内容污染。

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
This commit is contained in:
DaydreamCoding
2026-06-12 01:47:01 +08:00
co-authored by Claude Opus 4.8
parent e34ad2b194
commit b62b573f7f
56 changed files with 3036 additions and 184 deletions
@@ -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)
@@ -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)
+4
View File
@@ -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"`
@@ -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"),
@@ -0,0 +1,180 @@
package handler
import (
"net/http/httptest"
"strings"
"testing"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)
// newTestGinContext builds a bare gin.Context backed by an httptest recorder.
func newTestGinContext() *gin.Context {
gin.SetMode(gin.TestMode)
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
return c
}
// TestRecordCyberPolicyIfMarked_NoMark verifies that when no cyber mark is set,
// the function returns immediately and does NOT set the recorded flag.
func TestRecordCyberPolicyIfMarked_NoMark(t *testing.T) {
c := newTestGinContext()
h := &OpenAIGatewayHandler{}
h.recordCyberPolicyIfMarked(c, nil, nil, nil, "gpt-5", true, "", service.ChannelUsageFields{}, "")
// Flag must NOT be set when there was no mark.
require.False(t, c.GetBool(cyberPolicyRecordedKey),
"cyberPolicyRecordedKey must remain false when no cyber mark is present")
}
// TestRecordCyberPolicyIfMarked_WithMark verifies that:
// 1. When a cyber mark is present, the recorded flag is set (guard activated).
// 2. A second call is a no-op (idempotent guard).
// 3. Nil services do not panic.
func TestRecordCyberPolicyIfMarked_WithMark(t *testing.T) {
c := newTestGinContext()
service.MarkOpsCyberPolicy(c, service.CyberPolicyMark{
Message: "flagged",
Body: `{"error":{"code":"cyber_policy"}}`,
UpstreamStatus: 400,
})
h := &OpenAIGatewayHandler{} // nil services — must not panic
// First call: should set the flag.
require.NotPanics(t, func() {
h.recordCyberPolicyIfMarked(c, nil, nil, nil, "gpt-5", true, "", service.ChannelUsageFields{}, "")
})
require.True(t, c.GetBool(cyberPolicyRecordedKey),
"cyberPolicyRecordedKey must be true after first call with a mark")
// Second call: flag already set — must be a no-op (idempotent).
require.NotPanics(t, func() {
h.recordCyberPolicyIfMarked(c, nil, nil, nil, "gpt-5", false, "", service.ChannelUsageFields{}, "")
})
// Flag should still be true (not toggled or cleared).
require.True(t, c.GetBool(cyberPolicyRecordedKey),
"cyberPolicyRecordedKey must remain true after second call (guard)")
}
// TestRecordCyberPolicyIfMarked_ForwardSuccessSkipsUsageLog verifies the semantic:
// when forwardErrored=false the function still sets the guard flag (mark present),
// but the cyber usage row is NOT requested (only RecordCyberPolicyEvent fires).
// Since services are nil here we only verify the guard flag and no panic.
func TestRecordCyberPolicyIfMarked_ForwardSuccessSkipsUsageLog(t *testing.T) {
c := newTestGinContext()
service.MarkOpsCyberPolicy(c, service.CyberPolicyMark{
Message: "flagged",
UpstreamStatus: 200,
})
h := &OpenAIGatewayHandler{}
require.NotPanics(t, func() {
h.recordCyberPolicyIfMarked(c, nil, nil, nil, "gpt-5", false /* forwardErrored=false */, "", service.ChannelUsageFields{}, "")
})
require.True(t, c.GetBool(cyberPolicyRecordedKey))
}
// TestClearCyberPolicyTurnState verifies F1 at the handler level: after a turn
// is finalized, both the mark and the recorded guard are reset so the next WS
// turn detects/records independently.
func TestClearCyberPolicyTurnState(t *testing.T) {
c := newTestGinContext()
h := &OpenAIGatewayHandler{}
service.MarkOpsCyberPolicy(c, service.CyberPolicyMark{Message: "turn1", UpstreamStatus: 200})
h.recordCyberPolicyIfMarked(c, nil, nil, nil, "gpt-5", false, "", service.ChannelUsageFields{}, "")
require.True(t, c.GetBool(cyberPolicyRecordedKey))
clearCyberPolicyTurnState(c)
require.Nil(t, service.GetOpsCyberPolicy(c))
require.False(t, c.GetBool(cyberPolicyRecordedKey))
// turn2: a fresh cyber hit must be recordable again.
service.MarkOpsCyberPolicy(c, service.CyberPolicyMark{Message: "turn2", UpstreamStatus: 200})
h.recordCyberPolicyIfMarked(c, nil, nil, nil, "gpt-5", false, "", service.ChannelUsageFields{}, "")
require.True(t, c.GetBool(cyberPolicyRecordedKey))
require.Equal(t, "turn2", service.GetOpsCyberPolicy(c).Message)
}
// TestBuildCyberSessionBlockedOpsEntry verifies the locally-rejected request is
// auditable: 403 / phase=request / type=cyber_policy_session_blocked — distinct
// from upstream cyber_policy hits, and it must NOT touch moderation/violation.
func TestBuildCyberSessionBlockedOpsEntry(t *testing.T) {
entry := buildCyberSessionBlockedOpsEntry(cyberPolicyOpsErrorMeta{
RequestID: "req-9", Model: "gpt-5", RequestPath: "/openai/v1/responses",
})
require.Equal(t, 403, entry.StatusCode)
require.Equal(t, "cyber_policy_session_blocked", entry.ErrorType)
require.Equal(t, "request", entry.ErrorPhase)
require.True(t, entry.IsBusinessLimited)
require.Equal(t, "gateway_local", entry.ErrorSource)
require.Equal(t, "platform", entry.ErrorOwner)
require.Empty(t, entry.ErrorBody, "no session block key → ErrorBody must be empty")
entryWithKey := buildCyberSessionBlockedOpsEntry(cyberPolicyOpsErrorMeta{
RequestID: "req-9", Model: "gpt-5", RequestPath: "/openai/v1/responses",
SessionBlockKey: "abc123",
})
require.Equal(t, "session_block_key=abc123", entryWithKey.ErrorBody)
}
// TestRejectIfCyberSessionBlocked_FailOpen verifies fail-open paths: nil handler
// services, no explicit session signal, and (implicitly) disabled switch all
// pass the request through.
func TestRejectIfCyberSessionBlocked_FailOpen(t *testing.T) {
c := newTestGinContext()
c.Request = httptest.NewRequest("POST", "/openai/v1/responses", strings.NewReader(`{}`))
h := &OpenAIGatewayHandler{}
require.False(t, h.rejectIfCyberSessionBlocked(c, nil, []byte(`{}`), "gpt-5", cyberBlockFormatResponses), "nil apiKey → pass")
h2 := &OpenAIGatewayHandler{gatewayService: nil}
key := &service.APIKey{ID: 1}
require.False(t, h2.rejectIfCyberSessionBlocked(c, key, []byte(`{}`), "gpt-5", cyberBlockFormatResponses), "nil gateway service → pass")
}
// TestRecordCyberPolicyIfMarked_BlockKeyPlumbed verifies the 6th param is
// accepted and a non-empty key with nil gateway service does not panic
// (write-side guards live in the service layer).
func TestRecordCyberPolicyIfMarked_BlockKeyPlumbed(t *testing.T) {
c := newTestGinContext()
service.MarkOpsCyberPolicy(c, service.CyberPolicyMark{Message: "x", UpstreamStatus: 400})
h := &OpenAIGatewayHandler{}
require.NotPanics(t, func() {
h.recordCyberPolicyIfMarked(c, nil, nil, nil, "gpt-5", true, "deadbeef", service.ChannelUsageFields{}, "")
})
}
// TestBuildCyberPolicyOpsErrorEntry_StatusCode verifies F6: the ops error log
// records the status the codex client actually received (400 non-stream / 200 stream),
// not a hardcoded 403.
func TestBuildCyberPolicyOpsErrorEntry_StatusCode(t *testing.T) {
for _, tc := range []struct {
name string
upstreamStatus int
}{
{"non_stream_400", 400},
{"stream_200", 200},
{"zero_value", 0},
} {
t.Run(tc.name, func(t *testing.T) {
mark := &service.CyberPolicyMark{
Code: "cyber_policy",
Message: "blocked",
UpstreamStatus: tc.upstreamStatus,
}
entry := buildCyberPolicyOpsErrorEntry(cyberPolicyOpsErrorMeta{
RequestID: "req-1", Model: "gpt-5", RequestPath: "/openai/v1/responses",
}, mark)
require.Equal(t, tc.upstreamStatus, entry.StatusCode)
require.Equal(t, "cyber_policy", entry.ErrorType)
require.Equal(t, "request", entry.ErrorPhase)
})
}
}
@@ -34,6 +34,7 @@ type OpenAIGatewayHandler struct {
usageRecordWorkerPool *service.UsageRecordWorkerPool
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 "-", "-"
@@ -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")))
})
}
@@ -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
}
@@ -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 命中应跳过中间件落库")
}