mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-24 16:05:44 +08:00
feat(openai): cyber_policy 硬阻断全链路透传、审计与计费
上游对单次请求下发 error.code=cyber_policy 硬阻断时,网关在所有端点 (/v1/responses、/v1/chat/completions、/v1/messages、WebSocket)及流式/ 非流式路径下,将该结果原样透传给客户端,绝不 failover、换号或同步拦截; 命中后异步完成审计与计费: - 风控中心记录 cyber_policy 留痕并发送通知邮件,落库先于发信,SMTP 阻塞 不影响留痕 - ops 错误请求记录,状态码对齐客户端实际接收(流式 200 / 非流式 400) - 用量明细标记 request_type=cyber,按上游真实 token 计费,HTTP 与 WebSocket 计费口径统一,零 token 命中不误扣 - 会话级自动屏蔽(管理员开关,默认关):命中的会话在可配 TTL 内本地拦截 不再发往上游,仅屏蔽该会话不影响同 Key 其他会话 - 封号计数排除开关:可选让 cyber 命中不计入自动封号,命中当次不判定且 历史行在违规计数中一并排除 WebSocket 多轮连接下 cyber 标记按 turn 生命周期管理,逐轮独立检测与记录; 透传的错误响应不被兜底逻辑追加内容污染。 Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 4.8
parent
e34ad2b194
commit
b62b573f7f
@@ -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)
|
||||
|
||||
@@ -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 命中应跳过中间件落库")
|
||||
}
|
||||
Reference in New Issue
Block a user