refactor(seams): Phase-2 接缝——payment 注册表化 + 网关 pre-flight 钩子链

设计经双视角架构评审裁决(零行为变更,特征化测试先行锁定):
- payment/provider:factory switch → 包内私有构造器注册表(init 自注册);
  unknown-key 文案与 ApplicationError 透传逐字节等价(前端 i18n 依赖)
- internal/gatewayhook:协议无关 pre-flight 钩子链(panic 隔离、error 默认
  fail-open、首 Blocked 短路、空链零分配);Decision→格式化保留在各调用点
- 8 个 HTTP moderation 调用点收敛经链(各点入参表达式/格式化函数逐一保留),
  WebSocket 两点按裁决保持原路径;moderation 为 Wire 装配的核心钩子
- 新增 7 个拦截格式特征化测试(5 种 HTTP 格式差异 + WS turn-2 + fail-open)

实施后对抗审计零 P0/P1,突变实验 4/4 被测试捕获;bench allocs 全部零增加。
This commit is contained in:
shaw
2026-06-11 22:59:47 +08:00
parent 9be516448e
commit 1cdab6f858
22 changed files with 1426 additions and 45 deletions
+82
View File
@@ -0,0 +1,82 @@
package gatewayhook
import (
"context"
"github.com/Wei-Shaw/sub2api/internal/pkg/logger"
"go.uber.org/zap"
)
// Chain 按固定顺序执行 pre-flight 钩子。
//
// 执行语义(SEAM-DESIGN 裁决 H 链层三件必补):
// 1. 每钩子 recover():panic 记 error 日志后按 fail-open 继续下一钩子;
// 2. 钩子返回 error 默认 fail-open:记日志后继续下一钩子(与内容审核现状一致);
// 3. 首个 Blocked Decision 即短路返回;非 Blocked 的 Decision 视为放行并继续。
//
// Run 永不向调用方返回 error:返回 nil 即放行。
type Chain struct {
hooks []PreFlightHook
log *zap.Logger
}
// NewChain 构造钩子链。log 传 nil 时延迟回退到进程全局 logger;
// nil 钩子被忽略。
func NewChain(log *zap.Logger, hooks ...PreFlightHook) *Chain {
filtered := make([]PreFlightHook, 0, len(hooks))
for _, hook := range hooks {
if hook != nil {
filtered = append(filtered, hook)
}
}
return &Chain{hooks: filtered, log: log}
}
// IsEmpty 报告链上是否没有任何钩子(nil 链视为空链)。
// 调用侧可据此零开销跳过钩子请求的构造。
func (c *Chain) IsEmpty() bool {
return c == nil || len(c.hooks) == 0
}
// Run 按序执行钩子,返回首个 Blocked Decision;全部放行(或链为空)返回 nil。
func (c *Chain) Run(ctx context.Context, req *Request) *Decision {
if c.IsEmpty() {
return nil
}
for _, hook := range c.hooks {
if decision := c.runHook(ctx, hook, req); decision != nil && decision.Blocked {
return decision
}
}
return nil
}
// runHook 执行单个钩子并隔离其 panic / error(均按 fail-open 降级为放行)。
func (c *Chain) runHook(ctx context.Context, hook PreFlightHook, req *Request) (decision *Decision) {
hookID := ""
defer func() {
if r := recover(); r != nil {
c.logger().Error("gatewayhook.hook_panic",
zap.String("hook_id", hookID),
zap.Any("panic_value", r),
zap.Stack("stack"))
decision = nil // fail-open:panic 钩子视为放行
}
}()
hookID = hook.HookID()
d, err := hook.CheckPreFlight(ctx, req)
if err != nil {
c.logger().Warn("gatewayhook.hook_error",
zap.String("hook_id", hookID),
zap.Error(err))
return nil // fail-open:钩子自身故障视为放行
}
return d
}
func (c *Chain) logger() *zap.Logger {
if c != nil && c.log != nil {
return c.log
}
return logger.L()
}
+130
View File
@@ -0,0 +1,130 @@
package gatewayhook
import (
"context"
"errors"
"testing"
"github.com/stretchr/testify/require"
"go.uber.org/zap"
"go.uber.org/zap/zaptest/observer"
)
// stubHook 是可编程的测试钩子,记录自己是否被执行。
type stubHook struct {
id string
decision *Decision
err error
panicVal any
calls *[]string
}
func (h *stubHook) HookID() string { return h.id }
func (h *stubHook) CheckPreFlight(ctx context.Context, req *Request) (*Decision, error) {
if h.calls != nil {
*h.calls = append(*h.calls, h.id)
}
if h.panicVal != nil {
panic(h.panicVal)
}
return h.decision, h.err
}
func TestChainRun_ExecutesHooksInOrder(t *testing.T) {
var calls []string
chain := NewChain(zap.NewNop(),
&stubHook{id: "first", calls: &calls},
&stubHook{id: "second", calls: &calls},
&stubHook{id: "third", calls: &calls},
)
decision := chain.Run(context.Background(), &Request{})
require.Nil(t, decision, "全部放行时 Run 必须返回 nil")
require.Equal(t, []string{"first", "second", "third"}, calls, "钩子必须按注册顺序执行")
}
func TestChainRun_FirstBlockedShortCircuits(t *testing.T) {
var calls []string
blocked := &Decision{Blocked: true, StatusCode: 403, ErrorType: "content_policy_violation", Message: "blocked"}
chain := NewChain(zap.NewNop(),
&stubHook{id: "pass", calls: &calls},
&stubHook{id: "block", decision: blocked, calls: &calls},
&stubHook{id: "never", calls: &calls},
)
decision := chain.Run(context.Background(), &Request{})
require.Same(t, blocked, decision, "必须原样返回首个 Blocked Decision")
require.Equal(t, []string{"pass", "block"}, calls, "Blocked 之后的钩子不得执行")
}
func TestChainRun_NonBlockedDecisionContinues(t *testing.T) {
var calls []string
chain := NewChain(zap.NewNop(),
&stubHook{id: "flag-only", decision: &Decision{Blocked: false, Message: "flagged"}, calls: &calls},
&stubHook{id: "after", calls: &calls},
)
decision := chain.Run(context.Background(), &Request{})
require.Nil(t, decision, "非 Blocked 的 Decision 视为放行")
require.Equal(t, []string{"flag-only", "after"}, calls)
}
func TestChainRun_PanicIsolatedAndFailOpen(t *testing.T) {
core, logs := observer.New(zap.ErrorLevel)
var calls []string
blocked := &Decision{Blocked: true, StatusCode: 403, Message: "blocked"}
chain := NewChain(zap.New(core),
&stubHook{id: "boomer", panicVal: "boom", calls: &calls},
&stubHook{id: "blocker", decision: blocked, calls: &calls},
)
decision := chain.Run(context.Background(), &Request{})
require.Same(t, blocked, decision, "panic 钩子 fail-open 后链必须继续执行后续钩子")
require.Equal(t, []string{"boomer", "blocker"}, calls)
entries := logs.FilterMessage("gatewayhook.hook_panic").All()
require.Len(t, entries, 1, "panic 必须记一条 error 日志")
fields := entries[0].ContextMap()
require.Equal(t, "boomer", fields["hook_id"])
require.Equal(t, "boom", fields["panic_value"])
// 仅 panic 钩子的链:整体放行。
onlyPanic := NewChain(zap.New(core), &stubHook{id: "boomer2", panicVal: errors.New("kaboom")})
require.Nil(t, onlyPanic.Run(context.Background(), &Request{}))
}
func TestChainRun_HookErrorFailOpen(t *testing.T) {
core, logs := observer.New(zap.WarnLevel)
var calls []string
chain := NewChain(zap.New(core),
// 即使 error 钩子同时返回了 Blocked Decision,error 优先按 fail-open 丢弃。
&stubHook{id: "broken", decision: &Decision{Blocked: true}, err: errors.New("hook exploded"), calls: &calls},
&stubHook{id: "after", calls: &calls},
)
decision := chain.Run(context.Background(), &Request{})
require.Nil(t, decision, "钩子 error 必须 fail-open 放行")
require.Equal(t, []string{"broken", "after"}, calls, "error 钩子之后链必须继续")
entries := logs.FilterMessage("gatewayhook.hook_error").All()
require.Len(t, entries, 1)
require.Equal(t, "broken", entries[0].ContextMap()["hook_id"])
}
func TestChainRun_EmptyChain(t *testing.T) {
require.Nil(t, NewChain(zap.NewNop()).Run(context.Background(), &Request{}))
require.True(t, NewChain(zap.NewNop()).IsEmpty())
var nilChain *Chain
require.True(t, nilChain.IsEmpty(), "nil 链视为空链")
require.Nil(t, nilChain.Run(context.Background(), &Request{}), "nil 链 Run 必须安全放行")
// nil 钩子在构造时被过滤。
require.True(t, NewChain(zap.NewNop(), nil, nil).IsEmpty())
}
+96
View File
@@ -0,0 +1,96 @@
// Package gatewayhook 定义网关 pre-flight 钩子链(Phase-2 SEAM-DESIGN 裁决 H)。
//
// 设计要点:
// - 独立包:不依赖 internal/handler 与 internal/server/middleware(CallerInfo
// 即为与 middleware.AuthSubject 解耦的最小身份视图);
// - 链只产 Decision,错误格式化留在各调用点(5 种 HTTP 拦截格式差异由
// 特征化测试锁定,与协议常量无映射关系);
// - 核心钩子(如内容审核)由 Wire 装配;gateway.hook.* 命名空间的模块收集
// 机制留待 Phase-3 与平台 Provider 需求合并设计,本包不做 Runtime 收集。
package gatewayhook
import (
"context"
"net/http"
"github.com/Wei-Shaw/sub2api/internal/service"
)
// CallerInfo 是调用方身份的最小视图(与 middleware.AuthSubject 解耦)。
type CallerInfo struct {
// UserID 为认证后的用户 ID。
UserID int64
// GroupID 为 API Key 所属分组 ID(可能为 nil)。
GroupID *int64
// KeyID 为平台 API Key ID。
KeyID int64
}
// RequestHeaders 是入站请求头的只读访问器(不暴露 gin.Context / http.Header 本体)。
type RequestHeaders struct {
header http.Header
}
// NewRequestHeaders 包装 http.Header 为只读访问器;传 nil 得到空访问器。
func NewRequestHeaders(header http.Header) RequestHeaders {
return RequestHeaders{header: header}
}
// Get 返回指定 key 的首个 header 值(语义同 http.Header.Get)。
func (r RequestHeaders) Get(key string) string {
if r.header == nil {
return ""
}
return r.header.Get(key)
}
// Values 返回指定 key 的全部 header 值(语义同 http.Header.Values)。
func (r RequestHeaders) Values(key string) []string {
if r.header == nil {
return nil
}
return r.header.Values(key)
}
// Request 是 pre-flight 钩子的协议无关只读入参。
type Request struct {
// Protocol 复用 service.ContentModerationProtocol* 现有常量值
// (如 "anthropic_messages" / "openai_chat_completions"),禁止新造枚举。
Protocol string
// Model 为请求的模型名(未做 TrimSpace,由钩子按需规范化)。
Model string
// Body 为送检请求体。各调用点保留自己的入参表达式
// (如 images 调用点传 parsed.ModerationBody()),不强行统一。
Body []byte
// Caller 为调用方最小身份视图。
Caller CallerInfo
// APIKey 为完整平台 API Key 上下文(audit 类钩子需要;可能为 nil)。
APIKey *service.APIKey
// Headers 为入站请求头的只读访问器。
Headers RequestHeaders
// Path 为入站端点路径(调用侧按 inbound endpoint 归一化后传入)。
Path string
}
// Decision 是钩子的协议无关决策;nil 表示放行。
// ErrorType 对 gemini 调用点无效(googleError 格式化器不消费该字段)。
type Decision struct {
Blocked bool
StatusCode int
ErrorType string
Message string
}
// PreFlightHook 是网关转发前的拦截钩子。
//
// 契约(链层强制,见 chain.go):
// - 返回 (nil, nil) 表示放行;
// - 返回 Blocked Decision 表示拦截(链短路返回);
// - 返回 error 表示钩子自身故障,链层按 fail-open 处理(记日志后继续下一钩子);
// - 钩子内 panic 同样被链层隔离并按 fail-open 降级。
type PreFlightHook interface {
// HookID 返回钩子的稳定标识(用于日志与排障)。
HookID() string
// CheckPreFlight 对请求执行转发前检查。
CheckPreFlight(ctx context.Context, req *Request) (*Decision, error)
}
@@ -2,7 +2,6 @@ package handler
import (
"context"
"net/http"
"strings"
"github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey"
@@ -12,19 +11,10 @@ import (
"go.uber.org/zap"
)
func (h *GatewayHandler) checkContentModeration(c *gin.Context, reqLog *zap.Logger, apiKey *service.APIKey, subject middleware2.AuthSubject, protocol string, model string, body []byte) *service.ContentModerationDecision {
if h == nil || h.contentModerationService == nil {
return nil
}
return runContentModeration(c, reqLog, h.contentModerationService, apiKey, subject, protocol, model, body)
}
func contentModerationStatus(decision *service.ContentModerationDecision) int {
if decision == nil || decision.StatusCode < 400 || decision.StatusCode > 599 {
return http.StatusForbidden
}
return decision.StatusCode
}
// 本文件保留 WebSocket 调用点(openai_gateway_handler.go ResponsesWebSocket 的
// turn-1 首帧与 turn≥2 每消息审核)仍在使用的内容审核 helper。
// HTTP 调用点已迁移至 pre-flight 钩子链(见 gateway_preflight.go,SEAM-DESIGN 裁决 H:
// WS 两点因每消息语义 + WS 帧错误格式排除在链改造之外,保持现状)。
func contentModerationErrorCode(decision *service.ContentModerationDecision) string {
return "content_policy_violation"
+6 -5
View File
@@ -15,6 +15,7 @@ import (
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/Wei-Shaw/sub2api/internal/domain"
"github.com/Wei-Shaw/sub2api/internal/gatewayhook"
"github.com/Wei-Shaw/sub2api/internal/pkg/antigravity"
"github.com/Wei-Shaw/sub2api/internal/pkg/claude"
"github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey"
@@ -47,7 +48,7 @@ type GatewayHandler struct {
apiKeyService *service.APIKeyService
usageRecordWorkerPool *service.UsageRecordWorkerPool
errorPassthroughService *service.ErrorPassthroughService
contentModerationService *service.ContentModerationService
preFlightHooks *gatewayhook.Chain
concurrencyHelper *ConcurrencyHelper
userMsgQueueHelper *UserMsgQueueHelper
maxAccountSwitches int
@@ -68,7 +69,7 @@ func NewGatewayHandler(
apiKeyService *service.APIKeyService,
usageRecordWorkerPool *service.UsageRecordWorkerPool,
errorPassthroughService *service.ErrorPassthroughService,
contentModerationService *service.ContentModerationService,
preFlightHooks *gatewayhook.Chain,
userMsgQueueService *service.UserMessageQueueService,
cfg *config.Config,
settingService *service.SettingService,
@@ -102,7 +103,7 @@ func NewGatewayHandler(
apiKeyService: apiKeyService,
usageRecordWorkerPool: usageRecordWorkerPool,
errorPassthroughService: errorPassthroughService,
contentModerationService: contentModerationService,
preFlightHooks: preFlightHooks,
concurrencyHelper: NewConcurrencyHelper(concurrencyService, SSEPingFormatClaude, pingInterval),
userMsgQueueHelper: umqHelper,
maxAccountSwitches: maxAccountSwitches,
@@ -195,8 +196,8 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
return
}
if decision := h.checkContentModeration(c, reqLog, apiKey, subject, service.ContentModerationProtocolAnthropicMessages, reqModel, body); decision != nil && decision.Blocked {
h.errorResponse(c, contentModerationStatus(decision), contentModerationErrorCode(decision), decision.Message)
if decision := h.runPreFlightHooks(c, reqLog, apiKey, subject, service.ContentModerationProtocolAnthropicMessages, reqModel, body); decision != nil && decision.Blocked {
h.errorResponse(c, preFlightStatus(decision), preFlightErrorCode(decision), decision.Message)
return
}
@@ -95,8 +95,8 @@ func (h *GatewayHandler) ChatCompletions(c *gin.Context) {
return
}
if decision := h.checkContentModeration(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIChat, reqModel, body); decision != nil && decision.Blocked {
h.chatCompletionsErrorResponse(c, contentModerationStatus(decision), contentModerationErrorCode(decision), decision.Message)
if decision := h.runPreFlightHooks(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIChat, reqModel, body); decision != nil && decision.Blocked {
h.chatCompletionsErrorResponse(c, preFlightStatus(decision), preFlightErrorCode(decision), decision.Message)
return
}
@@ -104,8 +104,8 @@ func (h *GatewayHandler) Responses(c *gin.Context) {
return
}
if decision := h.checkContentModeration(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIResponses, reqModel, body); decision != nil && decision.Blocked {
h.responsesErrorResponse(c, contentModerationStatus(decision), contentModerationErrorCode(decision), decision.Message)
if decision := h.runPreFlightHooks(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIResponses, reqModel, body); decision != nil && decision.Blocked {
h.responsesErrorResponse(c, preFlightStatus(decision), preFlightErrorCode(decision), decision.Message)
return
}
@@ -0,0 +1,479 @@
//go:build unit
// Phase-2 TASK-001 特征化测试:按"格式类"锁定内容审核拦截后的客户端可观测输出
// (钩子链改造 TASK-003 的实施前置 gate,调用点清单与格式类定义见
// issues/plugin-refactor/phase-2_pilots/CALLSITE-INVENTORY.md)。
//
// 覆盖(已有测试不重复):
// - B 格式 chat_completions(无顶层 type) —— gateway_handler_chat_completions.go:98
// - C 格式 responses(error.code 字符串字段) —— gateway_handler_responses.go:107
// - D 格式 gemini googleError(code int + status 串)—— gemini_v1beta_handler.go:190
// - A' 格式 openai 网关 anthropic(顶层 type:"error")—— openai_gateway_handler.go:676(格式化 :934)
// - B' 格式 openai images(与 B 字节级同形,归并覆盖 openai_gateway_handler.go:244 / openai_chat_completions.go:88,
// 三点共用 openai_gateway_handler.go:1920 同一格式化函数)—— openai_images.go:89
// - fail-open 非 anthropic 协议(chat_completions 审核服务 500 → 放行至账号调度阶段)
// - WS turn-2 每消息审核 close-error 路径 —— openai_gateway_handler.go:1426
//
// 特征化纪律:所有断言均为先驱动真实链路跑出实际行为后固化,不写应然值。
// 复用 gateway_intercept_characterization_test.go 的 passChar* 夹具与
// openai_gateway_handler_test.go 的 WS 测试基建;新增辅助统一 p2Char 前缀。
package handler
import (
"bytes"
"context"
"io"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey"
middleware "github.com/Wei-Shaw/sub2api/internal/server/middleware"
"github.com/Wei-Shaw/sub2api/internal/service"
coderws "github.com/coder/websocket"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
"github.com/tidwall/gjson"
)
// p2CharFlagAllModerationServer 返回对任意输入都判定命中的 mock 审核 API。
func p2CharFlagAllModerationServer(t *testing.T) *httptest.Server {
t.Helper()
return httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
_, _ = w.Write([]byte(`{"results":[{"flagged":true,"category_scores":{"sexual":0.99}}]}`))
}))
}
// p2CharNewContext 构造带认证上下文的网关请求(任意路径/分组/请求体)。
func p2CharNewContext(t *testing.T, path string, group *service.Group, body []byte) (*gin.Context, *httptest.ResponseRecorder) {
t.Helper()
gin.SetMode(gin.TestMode)
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
req := httptest.NewRequest(http.MethodPost, path, bytes.NewReader(body))
req.Header.Set("Content-Type", "application/json")
req = req.WithContext(context.WithValue(req.Context(), ctxkey.Group, group))
c.Request = req
apiKey := &service.APIKey{
ID: 7311,
UserID: 7411,
GroupID: &group.ID,
Status: service.StatusActive,
User: &service.User{
ID: 7411,
Concurrency: 10,
Balance: 100,
},
Group: group,
}
c.Set(string(middleware.ContextKeyAPIKey), apiKey)
c.Set(string(middleware.ContextKeyUser), middleware.AuthSubject{UserID: apiKey.UserID, Concurrency: 10})
return c, rec
}
// p2CharAnthropicGroup 返回 anthropic 平台分组(GatewayHandler 系端点用)。
func p2CharAnthropicGroup() *service.Group {
return &service.Group{
ID: 7001,
Hydrated: true,
Platform: service.PlatformAnthropic,
Status: service.StatusActive,
}
}
// TestP2Characterization_ModerationBlock_ChatCompletionsFormat 锁定 B 格式:
// GatewayHandler.ChatCompletions(/v1/chat/completions)审核命中时返回 403 +
// {"error":{"type":"content_policy_violation","message":<拦截文案>}}——无顶层 type 字段。
func TestP2Characterization_ModerationBlock_ChatCompletionsFormat(t *testing.T) {
moderationSrv := p2CharFlagAllModerationServer(t)
defer moderationSrv.Close()
group := p2CharAnthropicGroup()
h, cleanup := newTestGatewayHandler(t, group, nil)
defer cleanup()
h.preFlightHooks = ProvideGatewayHookChain(passCharModerationService(t, moderationSrv.URL))
body := []byte(`{"model":"claude-sonnet-4-5","stream":false,"messages":[{"role":"user","content":"bad words"}]}`)
c, rec := p2CharNewContext(t, "/v1/chat/completions", group, body)
h.ChatCompletions(c)
require.Equal(t, http.StatusForbidden, rec.Code)
require.JSONEq(t,
`{"error":{"type":"content_policy_violation","message":"内容审计命中风险规则,请调整输入后重试"}}`,
rec.Body.String())
// 关键差异点:chat_completions 格式没有顶层 type:"error"(区别于 anthropic 格式)。
require.False(t, gjson.GetBytes(rec.Body.Bytes(), "type").Exists(),
"chat_completions 拦截体不应有顶层 type 字段")
}
// TestP2Characterization_ModerationBlock_ResponsesFormat 锁定 C 格式:
// GatewayHandler.Responses(/v1/responses)审核命中时返回 403 +
// {"error":{"code":"content_policy_violation","message":<拦截文案>}}——
// error 内是 code 字段(字符串)而非 type 字段。
func TestP2Characterization_ModerationBlock_ResponsesFormat(t *testing.T) {
moderationSrv := p2CharFlagAllModerationServer(t)
defer moderationSrv.Close()
group := p2CharAnthropicGroup()
h, cleanup := newTestGatewayHandler(t, group, nil)
defer cleanup()
h.preFlightHooks = ProvideGatewayHookChain(passCharModerationService(t, moderationSrv.URL))
body := []byte(`{"model":"claude-sonnet-4-5","input":[{"type":"message","role":"user","content":[{"type":"input_text","text":"bad words"}]}]}`)
c, rec := p2CharNewContext(t, "/v1/responses", group, body)
h.Responses(c)
require.Equal(t, http.StatusForbidden, rec.Code)
require.JSONEq(t,
`{"error":{"code":"content_policy_violation","message":"内容审计命中风险规则,请调整输入后重试"}}`,
rec.Body.String())
// 关键差异点:responses 格式用 error.code(字符串),无 error.type、无顶层 type。
require.False(t, gjson.GetBytes(rec.Body.Bytes(), "type").Exists())
require.False(t, gjson.GetBytes(rec.Body.Bytes(), "error.type").Exists(),
"responses 拦截体 error 内不应有 type 字段(用 code)")
require.Equal(t, gjson.String, gjson.GetBytes(rec.Body.Bytes(), "error.code").Type,
"responses 拦截体 error.code 应为字符串")
}
// TestP2Characterization_ModerationBlock_GeminiGoogleErrorFormat 锁定 D 格式:
// GatewayHandler.GeminiV1BetaModels(/v1beta/models/{model}:{action})审核命中时返回 403 +
// {"error":{"code":403,"message":<拦截文案>,"status":"PERMISSION_DENIED"}}。
// 注意:该调用点不消费 contentModerationErrorCode(content_policy_violation 被丢弃)。
func TestP2Characterization_ModerationBlock_GeminiGoogleErrorFormat(t *testing.T) {
moderationSrv := p2CharFlagAllModerationServer(t)
defer moderationSrv.Close()
group := &service.Group{
ID: 7002,
Hydrated: true,
Platform: service.PlatformGemini,
Status: service.StatusActive,
}
h, cleanup := newTestGatewayHandler(t, group, nil)
defer cleanup()
h.preFlightHooks = ProvideGatewayHookChain(passCharModerationService(t, moderationSrv.URL))
body := []byte(`{"contents":[{"role":"user","parts":[{"text":"bad words"}]}]}`)
c, rec := p2CharNewContext(t, "/v1beta/models/gemini-2.5-pro:generateContent", group, body)
c.Params = gin.Params{{Key: "modelAction", Value: "/gemini-2.5-pro:generateContent"}}
h.GeminiV1BetaModels(c)
require.Equal(t, http.StatusForbidden, rec.Code)
require.JSONEq(t,
`{"error":{"code":403,"message":"内容审计命中风险规则,请调整输入后重试","status":"PERMISSION_DENIED"}}`,
rec.Body.String())
// 关键差异点:googleError 三字段——code 是 int HTTP 状态码、status 是 google 状态串;
// 无顶层 type、无 error.type,审核错误码 content_policy_violation 不出现。
parsed := gjson.ParseBytes(rec.Body.Bytes())
require.False(t, parsed.Get("type").Exists(), "googleError 不应有顶层 type 字段")
require.False(t, parsed.Get("error.type").Exists(), "googleError 的 error 内不应有 type 字段")
require.Equal(t, gjson.Number, parsed.Get("error.code").Type, "googleError 的 error.code 应为数字")
require.NotContains(t, rec.Body.String(), "content_policy_violation",
"gemini 调用点不消费 contentModerationErrorCode")
}
// p2CharOpenAIGatewayHandler 构造满足 ensureResponsesDependencies 的最小 OpenAIGatewayHandler。
// 审核拦截发生在账号调度/转发之前,故各 service 用零值即可。
// HTTP 调用点走 preFlightHooks 链;contentModerationService 同时装配以镜像生产接线
// (WS 调用点仍直接消费该字段)。
func p2CharOpenAIGatewayHandler(t *testing.T, moderationBaseURL string) *OpenAIGatewayHandler {
t.Helper()
moderationSvc := passCharModerationService(t, moderationBaseURL)
return &OpenAIGatewayHandler{
gatewayService: &service.OpenAIGatewayService{},
billingCacheService: &service.BillingCacheService{},
apiKeyService: &service.APIKeyService{},
contentModerationService: moderationSvc,
preFlightHooks: ProvideGatewayHookChain(moderationSvc),
concurrencyHelper: NewConcurrencyHelper(service.NewConcurrencyService(&concurrencyCacheMock{}), SSEPingFormatNone, time.Second),
}
}
// TestP2Characterization_ModerationBlock_OpenAIGatewayAnthropicFormat 锁定 A' 格式:
// OpenAIGatewayHandler.Messages(openai 平台分组的 /v1/messages dispatch)审核命中时返回 403 +
// {"type":"error","error":{"type":"content_policy_violation","message":<拦截文案>}}
//(openai_gateway_handler.go:934 anthropicErrorResponse 一族,与 anthropic 主格式字节级同形)。
func TestP2Characterization_ModerationBlock_OpenAIGatewayAnthropicFormat(t *testing.T) {
moderationSrv := p2CharFlagAllModerationServer(t)
defer moderationSrv.Close()
group := &service.Group{
ID: 7003,
Hydrated: true,
Platform: service.PlatformOpenAI,
Status: service.StatusActive,
AllowMessagesDispatch: true,
}
h := p2CharOpenAIGatewayHandler(t, moderationSrv.URL)
body := []byte(`{"model":"claude-sonnet-4-5","max_tokens":64,"messages":[{"role":"user","content":[{"type":"text","text":"bad words"}]}]}`)
c, rec := p2CharNewContext(t, "/v1/messages", group, body)
h.Messages(c)
require.Equal(t, http.StatusForbidden, rec.Code)
require.JSONEq(t,
`{"type":"error","error":{"type":"content_policy_violation","message":"内容审计命中风险规则,请调整输入后重试"}}`,
rec.Body.String())
// 关键差异点:openai 网关 anthropic 格式保留顶层 type:"error"。
require.Equal(t, "error", gjson.GetBytes(rec.Body.Bytes(), "type").String())
}
// TestP2Characterization_ModerationBlock_OpenAIImagesFormat 锁定 B' 格式:
// OpenAIGatewayHandler.Images(/v1/images/generations,审核入参为 parsed.ModerationBody())
// 审核命中时返回 403 + {"error":{"type":"content_policy_violation","message":<拦截文案>}}。
// 该格式与 B(chat_completions)字节级同形;openai 网关 /v1/responses(openai_gateway_handler.go:244)
// 与 /v1/chat/completions(openai_chat_completions.go:88)共用同一格式化函数
//(openai_gateway_handler.go:1920),由本测试归并覆盖(见 CALLSITE-INVENTORY.md)。
func TestP2Characterization_ModerationBlock_OpenAIImagesFormat(t *testing.T) {
moderationSrv := p2CharFlagAllModerationServer(t)
defer moderationSrv.Close()
group := &service.Group{
ID: 7004,
Hydrated: true,
Platform: service.PlatformOpenAI,
Status: service.StatusActive,
AllowImageGeneration: true,
}
h := p2CharOpenAIGatewayHandler(t, moderationSrv.URL)
body := []byte(`{"model":"gpt-image-1","prompt":"bad words"}`)
c, rec := p2CharNewContext(t, "/v1/images/generations", group, body)
h.Images(c)
require.Equal(t, http.StatusForbidden, rec.Code)
require.JSONEq(t,
`{"error":{"type":"content_policy_violation","message":"内容审计命中风险规则,请调整输入后重试"}}`,
rec.Body.String())
require.False(t, gjson.GetBytes(rec.Body.Bytes(), "type").Exists(),
"openai 网关通用错误格式不应有顶层 type 字段")
}
// TestP2Characterization_ModerationFailOpen_ChatCompletionsFormat 锁定非 anthropic 协议的
// fail-open 行为:chat_completions 链路上审核 API 自身故障(HTTP 500)时请求必须放行——
// 请求穿过审核 gate 继续进入账号调度阶段(夹具无可用账号,故观测到 503 No available accounts,
// 而不是 403 content_policy_violation)。
func TestP2Characterization_ModerationFailOpen_ChatCompletionsFormat(t *testing.T) {
moderationSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
http.Error(w, "moderation backend exploded", http.StatusInternalServerError)
}))
defer moderationSrv.Close()
group := p2CharAnthropicGroup()
h, cleanup := newTestGatewayHandler(t, group, nil)
defer cleanup()
h.preFlightHooks = ProvideGatewayHookChain(passCharModerationService(t, moderationSrv.URL))
body := []byte(`{"model":"claude-sonnet-4-5","stream":false,"messages":[{"role":"user","content":"bad words"}]}`)
c, rec := p2CharNewContext(t, "/v1/chat/completions", group, body)
h.ChatCompletions(c)
// fail-open:未被审核拦截(非 403 + content_policy_violation),
// 请求推进到账号调度阶段(夹具无账号 → 503 api_error)。
require.Equal(t, http.StatusServiceUnavailable, rec.Code)
require.NotContains(t, rec.Body.String(), "content_policy_violation")
require.Equal(t, "api_error", gjson.GetBytes(rec.Body.Bytes(), "error.type").String())
require.Contains(t, gjson.GetBytes(rec.Body.Bytes(), "error.message").String(), "No available accounts")
}
// p2CharConditionalModerationServer 返回仅当审核请求中含 marker 才判定命中的 mock 审核 API。
func p2CharConditionalModerationServer(t *testing.T, marker string) *httptest.Server {
t.Helper()
return httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
raw, _ := io.ReadAll(r.Body)
w.Header().Set("Content-Type", "application/json")
if strings.Contains(string(raw), marker) {
_, _ = w.Write([]byte(`{"results":[{"flagged":true,"category_scores":{"sexual":0.99}}]}`))
return
}
_, _ = w.Write([]byte(`{"results":[{"flagged":false,"category_scores":{"sexual":0.0}}]}`))
}))
}
// TestP2Characterization_OpenAIResponsesWSTurn2ModerationCloseError 锁定 E 格式 turn-2 路径
// (openai_gateway_handler.go:1426,OpenAIWSIngressHooks.BeforeRequest 的 turn≥2 每消息审核):
// turn-1 干净消息正常转发并完成;turn-2 消息审核命中后客户端先收到
// writeContentModerationWSError 错误帧,随后连接以 close(1008 StatusPolicyViolation,
// reason=<拦截文案>) 关闭,且 turn-2 消息不会到达上游。
// 复用 TestOpenAIResponsesWebSocket_ContentModerationBlocksFirstFrame 的 WS 基建
//(newOpenAIWSHandlerTestServer 路由形态 + passthrough 上游夹具)。
func TestP2Characterization_OpenAIResponsesWSTurn2ModerationCloseError(t *testing.T) {
gin.SetMode(gin.TestMode)
const turn2Marker = "p2char-turn2-bad"
moderationSrv := p2CharConditionalModerationServer(t, turn2Marker)
defer moderationSrv.Close()
upstreamFrames := make(chan []byte, 4)
upstreamServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
conn, err := coderws.Accept(w, r, &coderws.AcceptOptions{CompressionMode: coderws.CompressionContextTakeover})
if err != nil {
return
}
defer func() { _ = conn.CloseNow() }()
for {
readCtx, cancelRead := context.WithTimeout(r.Context(), 5*time.Second)
_, payload, readErr := conn.Read(readCtx)
cancelRead()
if readErr != nil {
return
}
upstreamFrames <- payload
writeCtx, cancelWrite := context.WithTimeout(r.Context(), 3*time.Second)
writeErr := conn.Write(writeCtx, coderws.MessageText, []byte(
`{"type":"response.completed","response":{"id":"resp_p2char_turn1","model":"gpt-5.1","usage":{"input_tokens":1,"output_tokens":1}}}`,
))
cancelWrite()
if writeErr != nil {
return
}
}
}))
defer upstreamServer.Close()
account := service.Account{
ID: 9911,
Name: "p2char-ws-turn2",
Platform: service.PlatformOpenAI,
Type: service.AccountTypeAPIKey,
Status: service.StatusActive,
Schedulable: true,
Concurrency: 1,
Credentials: map[string]any{
"api_key": "sk-p2char",
"base_url": upstreamServer.URL,
},
Extra: map[string]any{
"openai_apikey_responses_websockets_v2_enabled": true,
"openai_apikey_responses_websockets_v2_mode": service.OpenAIWSIngressModePassthrough,
},
}
cfg := &config.Config{}
cfg.RunMode = config.RunModeSimple
cfg.Default.RateMultiplier = 1
cfg.Security.URLAllowlist.Enabled = false
cfg.Security.URLAllowlist.AllowInsecureHTTP = true
cfg.Gateway.OpenAIWS.Enabled = true
cfg.Gateway.OpenAIWS.APIKeyEnabled = true
cfg.Gateway.OpenAIWS.ResponsesWebsocketsV2 = true
cfg.Gateway.OpenAIWS.ModeRouterV2Enabled = true
cfg.Gateway.OpenAIWS.DialTimeoutSeconds = 3
cfg.Gateway.OpenAIWS.ReadTimeoutSeconds = 5
cfg.Gateway.OpenAIWS.WriteTimeoutSeconds = 3
accountRepo := &openAIWSUsageHandlerAccountRepoStub{account: account}
usageRepo := &openAIWSUsageHandlerUsageLogRepoStub{created: make(chan *service.UsageLog, 4)}
billingCacheSvc := service.NewBillingCacheService(nil, nil, nil, nil, nil, nil, cfg, nil)
gatewaySvc := service.NewOpenAIGatewayService(
accountRepo,
usageRepo,
nil,
nil,
nil,
nil,
nil,
cfg,
nil,
nil,
service.NewBillingService(cfg, nil),
nil,
billingCacheSvc,
nil,
&service.DeferredService{},
nil,
nil,
nil,
nil,
nil,
nil,
)
cache := &concurrencyCacheMock{
acquireUserSlotFn: func(ctx context.Context, userID int64, maxConcurrency int, requestID string) (bool, error) {
return true, nil
},
acquireAccountSlotFn: func(ctx context.Context, accountID int64, maxConcurrency int, requestID string) (bool, error) {
return true, nil
},
}
h := &OpenAIGatewayHandler{
gatewayService: gatewaySvc,
billingCacheService: billingCacheSvc,
apiKeyService: &service.APIKeyService{},
contentModerationService: passCharModerationService(t, moderationSrv.URL),
concurrencyHelper: NewConcurrencyHelper(service.NewConcurrencyService(cache), SSEPingFormatNone, time.Second),
}
wsServer := newOpenAIWSHandlerTestServer(t, h, middleware.AuthSubject{UserID: 1, Concurrency: 1})
defer wsServer.Close()
dialCtx, cancelDial := context.WithTimeout(context.Background(), 3*time.Second)
clientConn, _, err := coderws.Dial(
dialCtx,
"ws"+strings.TrimPrefix(wsServer.URL, "http")+"/openai/v1/responses",
&coderws.DialOptions{CompressionMode: coderws.CompressionContextTakeover},
)
cancelDial()
require.NoError(t, err)
defer func() { _ = clientConn.CloseNow() }()
// turn-1:干净消息正常转发,收到上游 response.completed。
writeCtx, cancelWrite := context.WithTimeout(context.Background(), 3*time.Second)
err = clientConn.Write(writeCtx, coderws.MessageText, []byte(
`{"type":"response.create","model":"gpt-5.1","stream":false,"input":[{"type":"message","role":"user","content":[{"type":"input_text","text":"clean prompt"}]}]}`,
))
cancelWrite()
require.NoError(t, err)
readCtx, cancelRead := context.WithTimeout(context.Background(), 5*time.Second)
_, turn1Event, err := clientConn.Read(readCtx)
cancelRead()
require.NoError(t, err)
require.Equal(t, "response.completed", gjson.GetBytes(turn1Event, "type").String())
select {
case <-upstreamFrames:
case <-time.After(3 * time.Second):
t.Fatal("等待上游收到 turn-1 帧超时")
}
// turn-2:命中审核的消息被 BeforeRequest 拦截。
writeCtx, cancelWrite = context.WithTimeout(context.Background(), 3*time.Second)
err = clientConn.Write(writeCtx, coderws.MessageText, []byte(
`{"type":"response.create","model":"gpt-5.1","stream":false,"input":[{"type":"message","role":"user","content":[{"type":"input_text","text":"`+turn2Marker+`"}]}]}`,
))
cancelWrite()
require.NoError(t, err)
// 实际行为:先收到 writeContentModerationWSError 错误帧,再收到 close(1008) 关闭。
readCtx, cancelRead = context.WithTimeout(context.Background(), 5*time.Second)
_, errFrame, readErr := clientConn.Read(readCtx)
cancelRead()
require.NoError(t, readErr, "turn-2 拦截后应先收到错误帧")
require.JSONEq(t,
`{"event_id":"evt_content_moderation_blocked","type":"error","error":{"type":"invalid_request_error","code":"content_policy_violation","message":"内容审计命中风险规则,请调整输入后重试"}}`,
string(errFrame))
readCtx, cancelRead = context.WithTimeout(context.Background(), 5*time.Second)
_, _, closeReadErr := clientConn.Read(readCtx)
cancelRead()
require.Error(t, closeReadErr)
var closeErr coderws.CloseError
require.ErrorAs(t, closeReadErr, &closeErr)
require.Equal(t, coderws.StatusPolicyViolation, closeErr.Code)
require.Equal(t, "内容审计命中风险规则,请调整输入后重试", closeErr.Reason)
// turn-2 消息不得到达上游。
select {
case frame := <-upstreamFrames:
t.Fatalf("turn-2 被拦截的消息不应转发到上游,实际收到: %s", frame)
case <-time.After(300 * time.Millisecond):
}
}
@@ -0,0 +1,215 @@
package handler
import (
"context"
"net/http"
"strings"
"github.com/Wei-Shaw/sub2api/internal/gatewayhook"
"github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey"
middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/gin-gonic/gin"
"go.uber.org/zap"
)
// ProvideGatewayHookChain 装配网关 pre-flight 钩子链(Wire 注入的核心钩子;
// 裁决 H-C:只建链,不做 gateway.hook.* 命名空间的 Runtime 收集——该机制留待
// Phase-3 与平台 Provider 真实需求合并设计)。当前链 = [content_moderation]。
func ProvideGatewayHookChain(contentModerationService *service.ContentModerationService) *gatewayhook.Chain {
if contentModerationService == nil {
return gatewayhook.NewChain(nil)
}
return gatewayhook.NewChain(nil, &contentModerationPreFlightHook{svc: contentModerationService})
}
// preFlightLoggerKey 在钩子执行期间通过 ctx 传递请求级 logger(reqLog),
// 保留其已累积的 component/model/stream 等字段,保证审核日志字段等价。
type preFlightLoggerKey struct{}
func withPreFlightLogger(ctx context.Context, log *zap.Logger) context.Context {
if log == nil {
return ctx
}
return context.WithValue(ctx, preFlightLoggerKey{}, log)
}
func preFlightLoggerFromContext(ctx context.Context) *zap.Logger {
if ctx == nil {
return nil
}
log, _ := ctx.Value(preFlightLoggerKey{}).(*zap.Logger)
return log
}
// newGatewayHookRequest 从 gin 上下文构造钩子请求(协议无关只读视图)。
// Path 的计算与原 buildContentModerationInput 的 Endpoint 等价:
// GetInboundEndpoint 优先,空值回退原始 URL path。
func newGatewayHookRequest(c *gin.Context, apiKey *service.APIKey, subject middleware2.AuthSubject, protocol string, model string, body []byte) *gatewayhook.Request {
caller := gatewayhook.CallerInfo{UserID: subject.UserID}
if apiKey != nil {
caller.KeyID = apiKey.ID
caller.GroupID = apiKey.GroupID
}
path := GetInboundEndpoint(c)
if path == "" && c.Request.URL != nil {
path = c.Request.URL.Path
}
return &gatewayhook.Request{
Protocol: protocol,
Model: model,
Body: body,
Caller: caller,
APIKey: apiKey,
Headers: gatewayhook.NewRequestHeaders(c.Request.Header),
Path: path,
}
}
// runGatewayPreFlight 构造钩子请求并执行 pre-flight 链。
// 等价性约束(对应原 runContentModeration 的调用面):
// - 链为空(审核服务未装配/测试夹具)时直接放行,且不构造任何对象(零分配);
// - c / c.Request 为 nil 时直接放行(与原 helper 的 nil 防御一致)。
func runGatewayPreFlight(chain *gatewayhook.Chain, c *gin.Context, reqLog *zap.Logger, apiKey *service.APIKey, subject middleware2.AuthSubject, protocol string, model string, body []byte) *gatewayhook.Decision {
if chain.IsEmpty() || c == nil || c.Request == nil {
return nil
}
req := newGatewayHookRequest(c, apiKey, subject, protocol, model, body)
return chain.Run(withPreFlightLogger(c.Request.Context(), reqLog), req)
}
// runPreFlightHooks 执行网关 pre-flight 钩子链(替代原 checkContentModeration 的 HTTP 调用面)。
func (h *GatewayHandler) runPreFlightHooks(c *gin.Context, reqLog *zap.Logger, apiKey *service.APIKey, subject middleware2.AuthSubject, protocol string, model string, body []byte) *gatewayhook.Decision {
if h == nil {
return nil
}
return runGatewayPreFlight(h.preFlightHooks, c, reqLog, apiKey, subject, protocol, model, body)
}
// runPreFlightHooks 执行网关 pre-flight 钩子链(HTTP 调用点专用;
// WebSocket 两调用点仍走 checkContentModeration,按裁决排除在链改造之外)。
func (h *OpenAIGatewayHandler) runPreFlightHooks(c *gin.Context, reqLog *zap.Logger, apiKey *service.APIKey, subject middleware2.AuthSubject, protocol string, model string, body []byte) *gatewayhook.Decision {
if h == nil {
return nil
}
return runGatewayPreFlight(h.preFlightHooks, c, reqLog, apiKey, subject, protocol, model, body)
}
// preFlightStatus 与原 contentModerationStatus 等价:非 4xx/5xx 状态码钳制为 403。
func preFlightStatus(decision *gatewayhook.Decision) int {
if decision == nil || decision.StatusCode < 400 || decision.StatusCode > 599 {
return http.StatusForbidden
}
return decision.StatusCode
}
// preFlightErrorCode 与原 contentModerationErrorCode 等价:
// 拦截错误码固定回退 content_policy_violation(gemini 调用点不消费该值)。
func preFlightErrorCode(decision *gatewayhook.Decision) string {
if decision == nil || strings.TrimSpace(decision.ErrorType) == "" {
return "content_policy_violation"
}
return decision.ErrorType
}
// contentModerationPreFlightHook 是内容审核的核心钩子 adapter
// (裁决 H:moderation 本体依赖面超出 Host 能力,不模块化,仅包装现有 Check)。
//
// 等价性硬约束(与原 runContentModeration 逐项对应,WS 路径仍直接使用后者):
// - gateway_check_start(11 字段)/ gateway_check_done(8 字段)结构化日志
// 的事件名、字段名、字段来源逐一保留,且经请求级 logger(reqLog)输出;
// - Check 返回 error 时记 content_moderation.check_failed 后放行(fail-open,
// 由本钩子自行吞错,不依赖链层的通用 error 降级,保持原日志事件名);
// - svc 为 nil 时直接放行。
type contentModerationPreFlightHook struct {
svc *service.ContentModerationService
}
func (h *contentModerationPreFlightHook) HookID() string { return "content_moderation" }
func (h *contentModerationPreFlightHook) CheckPreFlight(ctx context.Context, req *gatewayhook.Request) (*gatewayhook.Decision, error) {
if h == nil || h.svc == nil || req == nil {
return nil, nil
}
reqLog := preFlightLoggerFromContext(ctx)
input := moderationInputFromHookRequest(ctx, req)
if reqLog != nil {
reqLog.Info("content_moderation.gateway_check_start",
zap.String("request_id", input.RequestID),
zap.Int64("user_id", input.UserID),
zap.Int64("api_key_id", input.APIKeyID),
zap.String("api_key_name", input.APIKeyName),
zap.Int64p("group_id", input.GroupID),
zap.String("group_name", input.GroupName),
zap.String("endpoint", input.Endpoint),
zap.String("provider", input.Provider),
zap.String("protocol", input.Protocol),
zap.String("model", input.Model),
zap.Int("body_bytes", len(req.Body)),
)
}
decision, err := h.svc.Check(ctx, input)
if err != nil {
if reqLog != nil {
reqLog.Warn("content_moderation.check_failed", zap.Error(err))
}
return nil, nil
}
if reqLog != nil && decision != nil {
reqLog.Info("content_moderation.gateway_check_done",
zap.String("request_id", input.RequestID),
zap.Bool("allowed", decision.Allowed),
zap.Bool("blocked", decision.Blocked),
zap.Bool("flagged", decision.Flagged),
zap.String("action", decision.Action),
zap.Int("status_code", decision.StatusCode),
zap.String("highest_category", decision.HighestCategory),
zap.Float64("highest_score", decision.HighestScore),
)
}
if decision == nil || !decision.Blocked {
return nil, nil
}
return &gatewayhook.Decision{
Blocked: true,
StatusCode: decision.StatusCode,
ErrorType: "content_policy_violation",
Message: decision.Message,
}, nil
}
// moderationInputFromHookRequest 从钩子请求构造审核输入,与 buildContentModerationInput
// (WS 路径仍在使用)逐字段等价:
// - RequestID 来自 ctx(调用侧传 c.Request.Context());
// - Endpoint 即 req.Path(调用侧已按 GetInboundEndpoint + URL.Path 兜底预计算);
// - 强制平台覆盖改读 ctxkey.ForcePlatform(middleware.ForcePlatform 同步写入
// gin key 与 request context,两者在审核执行点恒一致)。
func moderationInputFromHookRequest(ctx context.Context, req *gatewayhook.Request) service.ContentModerationCheckInput {
input := service.ContentModerationCheckInput{
RequestID: contentModerationRequestID(ctx),
UserID: req.Caller.UserID,
Endpoint: req.Path,
Provider: contentModerationProvider(req.APIKey),
Model: strings.TrimSpace(req.Model),
Protocol: req.Protocol,
Body: req.Body,
}
if forcedPlatform, ok := ctx.Value(ctxkey.ForcePlatform).(string); ok {
input.Provider = strings.TrimSpace(forcedPlatform)
}
if req.APIKey != nil {
input.APIKeyID = req.APIKey.ID
input.APIKeyName = req.APIKey.Name
if req.APIKey.User != nil {
input.UserEmail = req.APIKey.User.Email
}
if req.APIKey.GroupID != nil {
groupID := *req.APIKey.GroupID
input.GroupID = &groupID
}
if req.APIKey.Group != nil {
input.GroupName = req.APIKey.Group.Name
}
}
return input
}
@@ -0,0 +1,193 @@
//go:build unit
// Phase-2 TASK-003 钩子链等价性单测:
// - moderationInputFromHookRequest 与 buildContentModerationInput(WS 路径仍用)逐字段等价;
// - moderation 核心钩子的 gateway_check_start(11 字段)/ gateway_check_done(8 字段)
// 结构化日志契约;
// - 空链 / nil 防御路径零分配(moderation 未装配时热路径无新增开销)。
package handler
import (
"bytes"
"context"
"net/http"
"net/http/httptest"
"testing"
"github.com/Wei-Shaw/sub2api/internal/gatewayhook"
"github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey"
middleware "github.com/Wei-Shaw/sub2api/internal/server/middleware"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
"go.uber.org/zap"
"go.uber.org/zap/zaptest/observer"
)
// preFlightTestGinContext 构造带请求上下文的 gin.Context。
// forcePlatform 非空时按 middleware.ForcePlatform 的真实行为同时写入
// request context(ctxkey.ForcePlatform)与 gin key。
func preFlightTestGinContext(t *testing.T, path string, forcePlatform string) *gin.Context {
t.Helper()
gin.SetMode(gin.TestMode)
c, _ := gin.CreateTestContext(httptest.NewRecorder())
req := httptest.NewRequest(http.MethodPost, path, bytes.NewReader([]byte(`{}`)))
ctx := context.WithValue(req.Context(), ctxkey.RequestID, "req-pf-001")
if forcePlatform != "" {
ctx = context.WithValue(ctx, ctxkey.ForcePlatform, forcePlatform)
}
c.Request = req.WithContext(ctx)
if forcePlatform != "" {
c.Set(string(middleware.ContextKeyForcePlatform), forcePlatform)
}
return c
}
func preFlightTestAPIKey() *service.APIKey {
groupID := int64(9001)
return &service.APIKey{
ID: 9101,
UserID: 9201,
Name: "pf-key",
GroupID: &groupID,
User: &service.User{ID: 9201, Email: "pf@example.com"},
Group: &service.Group{ID: groupID, Name: "pf-group", Platform: service.PlatformAnthropic},
}
}
// TestPreFlightModerationInput_EquivalentToLegacyBuilder 锁定钩子侧审核输入
// 与 WS 路径仍在使用的 buildContentModerationInput 逐字段等价。
func TestPreFlightModerationInput_EquivalentToLegacyBuilder(t *testing.T) {
subject := middleware.AuthSubject{UserID: 9201, Concurrency: 5}
body := []byte(`{"model":"claude-sonnet-4-5"}`)
cases := []struct {
name string
forcePlatform string
apiKey *service.APIKey
model string
}{
{name: "常规请求", apiKey: preFlightTestAPIKey(), model: " claude-sonnet-4-5 "},
{name: "强制平台覆盖provider", forcePlatform: "antigravity", apiKey: preFlightTestAPIKey(), model: "claude-sonnet-4-5"},
{name: "nil apiKey防御", apiKey: nil, model: "gpt-5.1"},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
c := preFlightTestGinContext(t, "/v1/messages", tc.forcePlatform)
legacy := buildContentModerationInput(c, tc.apiKey, subject, service.ContentModerationProtocolAnthropicMessages, tc.model, body)
req := newGatewayHookRequest(c, tc.apiKey, subject, service.ContentModerationProtocolAnthropicMessages, tc.model, body)
hookInput := moderationInputFromHookRequest(c.Request.Context(), req)
require.Equal(t, legacy, hookInput, "钩子侧审核输入必须与原 builder 逐字段等价")
})
}
}
// TestContentModerationPreFlightHook_CheckLogContract 锁定核心钩子的
// gateway_check_start / gateway_check_done 日志字段契约与拦截决策映射。
func TestContentModerationPreFlightHook_CheckLogContract(t *testing.T) {
moderationSrv := p2CharFlagAllModerationServer(t)
defer moderationSrv.Close()
core, logs := observer.New(zap.InfoLevel)
reqLog := zap.New(core)
hook := &contentModerationPreFlightHook{svc: passCharModerationService(t, moderationSrv.URL)}
c := preFlightTestGinContext(t, "/v1/messages", "")
apiKey := preFlightTestAPIKey()
subject := middleware.AuthSubject{UserID: 9201, Concurrency: 5}
body := []byte(`{"model":"claude-sonnet-4-5","messages":[{"role":"user","content":"bad words"}]}`)
req := newGatewayHookRequest(c, apiKey, subject, service.ContentModerationProtocolAnthropicMessages, "claude-sonnet-4-5", body)
ctx := withPreFlightLogger(c.Request.Context(), reqLog)
decision, err := hook.CheckPreFlight(ctx, req)
require.NoError(t, err)
require.NotNil(t, decision)
require.True(t, decision.Blocked)
require.Equal(t, "content_policy_violation", decision.ErrorType)
require.NotEmpty(t, decision.Message)
start := logs.FilterMessage("content_moderation.gateway_check_start").All()
require.Len(t, start, 1, "必须记一条 gateway_check_start")
startFields := start[0].ContextMap()
require.Equal(t, []string{
"request_id", "user_id", "api_key_id", "api_key_name", "group_id",
"group_name", "endpoint", "provider", "protocol", "model", "body_bytes",
}, preFlightLogFieldKeys(start[0]), "gateway_check_start 11 字段的名称与顺序必须保留")
require.Equal(t, "req-pf-001", startFields["request_id"])
require.Equal(t, int64(9201), startFields["user_id"])
require.Equal(t, int64(9101), startFields["api_key_id"])
require.Equal(t, "pf-key", startFields["api_key_name"])
require.Equal(t, int64(9001), startFields["group_id"])
require.Equal(t, "pf-group", startFields["group_name"])
require.Equal(t, "/v1/messages", startFields["endpoint"])
require.Equal(t, service.PlatformAnthropic, startFields["provider"])
require.Equal(t, service.ContentModerationProtocolAnthropicMessages, startFields["protocol"])
require.Equal(t, "claude-sonnet-4-5", startFields["model"])
require.Equal(t, int64(len(body)), startFields["body_bytes"])
done := logs.FilterMessage("content_moderation.gateway_check_done").All()
require.Len(t, done, 1, "必须记一条 gateway_check_done")
require.Equal(t, []string{
"request_id", "allowed", "blocked", "flagged", "action",
"status_code", "highest_category", "highest_score",
}, preFlightLogFieldKeys(done[0]), "gateway_check_done 8 字段的名称与顺序必须保留")
doneFields := done[0].ContextMap()
require.Equal(t, true, doneFields["blocked"])
require.Equal(t, "req-pf-001", doneFields["request_id"])
}
func preFlightLogFieldKeys(entry observer.LoggedEntry) []string {
keys := make([]string, 0, len(entry.Context))
for _, f := range entry.Context {
keys = append(keys, f.Key)
}
return keys
}
// TestRunGatewayPreFlight_EmptyChainZeroAlloc 锁定审核未装配(空链/nil 链)时
// pre-flight 调用面零分配直接放行——热路径无新增开销。
func TestRunGatewayPreFlight_EmptyChainZeroAlloc(t *testing.T) {
c := preFlightTestGinContext(t, "/v1/messages", "")
apiKey := preFlightTestAPIKey()
subject := middleware.AuthSubject{UserID: 9201}
body := []byte(`{}`)
reqLog := zap.NewNop()
for name, chain := range map[string]*gatewayhook.Chain{
"nil链": nil,
"空链": ProvideGatewayHookChain(nil),
} {
t.Run(name, func(t *testing.T) {
require.Nil(t, runGatewayPreFlight(chain, c, reqLog, apiKey, subject, service.ContentModerationProtocolAnthropicMessages, "m", body))
allocs := testing.AllocsPerRun(100, func() {
_ = runGatewayPreFlight(chain, c, reqLog, apiKey, subject, service.ContentModerationProtocolAnthropicMessages, "m", body)
})
require.Zero(t, allocs, "空链路径必须零分配")
})
}
}
// TestContentModerationPreFlightHook_NilGuards 锁定 nil 防御:
// svc/req 为 nil 或 ctx 无请求级 logger 时不 panic 且放行。
func TestContentModerationPreFlightHook_NilGuards(t *testing.T) {
var nilHook *contentModerationPreFlightHook
d, err := nilHook.CheckPreFlight(context.Background(), &gatewayhook.Request{})
require.NoError(t, err)
require.Nil(t, d)
hookNoSvc := &contentModerationPreFlightHook{}
d, err = hookNoSvc.CheckPreFlight(context.Background(), &gatewayhook.Request{})
require.NoError(t, err)
require.Nil(t, d)
moderationSrv := p2CharFlagAllModerationServer(t)
defer moderationSrv.Close()
hook := &contentModerationPreFlightHook{svc: passCharModerationService(t, moderationSrv.URL)}
d, err = hook.CheckPreFlight(context.Background(), nil)
require.NoError(t, err)
require.Nil(t, d)
}
@@ -187,8 +187,8 @@ func (h *GatewayHandler) GeminiV1BetaModels(c *gin.Context) {
setOpsRequestContext(c, modelName, stream)
setOpsEndpointContext(c, "", int16(service.RequestTypeFromLegacy(stream, false)))
if decision := h.checkContentModeration(c, reqLog, apiKey, authSubject, service.ContentModerationProtocolGemini, modelName, body); decision != nil && decision.Blocked {
googleError(c, contentModerationStatus(decision), decision.Message)
if decision := h.runPreFlightHooks(c, reqLog, apiKey, authSubject, service.ContentModerationProtocolGemini, modelName, body); decision != nil && decision.Blocked {
googleError(c, preFlightStatus(decision), decision.Message)
return
}
@@ -85,8 +85,8 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) {
setOpsRequestContext(c, reqModel, reqStream)
setOpsEndpointContext(c, "", int16(service.RequestTypeFromLegacy(reqStream, false)))
if decision := h.checkContentModeration(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIChat, reqModel, body); decision != nil && decision.Blocked {
h.errorResponse(c, contentModerationStatus(decision), contentModerationErrorCode(decision), decision.Message)
if decision := h.runPreFlightHooks(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIChat, reqModel, body); decision != nil && decision.Blocked {
h.errorResponse(c, preFlightStatus(decision), preFlightErrorCode(decision), decision.Message)
return
}
@@ -12,6 +12,7 @@ import (
"time"
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/Wei-Shaw/sub2api/internal/gatewayhook"
"github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey"
pkghttputil "github.com/Wei-Shaw/sub2api/internal/pkg/httputil"
"github.com/Wei-Shaw/sub2api/internal/pkg/ip"
@@ -34,6 +35,7 @@ type OpenAIGatewayHandler struct {
usageRecordWorkerPool *service.UsageRecordWorkerPool
errorPassthroughService *service.ErrorPassthroughService
contentModerationService *service.ContentModerationService
preFlightHooks *gatewayhook.Chain
concurrencyHelper *ConcurrencyHelper
imageLimiter *imageConcurrencyLimiter
maxAccountSwitches int
@@ -105,6 +107,7 @@ func NewOpenAIGatewayHandler(
usageRecordWorkerPool *service.UsageRecordWorkerPool,
errorPassthroughService *service.ErrorPassthroughService,
contentModerationService *service.ContentModerationService,
preFlightHooks *gatewayhook.Chain,
cfg *config.Config,
) *OpenAIGatewayHandler {
pingInterval := time.Duration(0)
@@ -122,6 +125,7 @@ func NewOpenAIGatewayHandler(
usageRecordWorkerPool: usageRecordWorkerPool,
errorPassthroughService: errorPassthroughService,
contentModerationService: contentModerationService,
preFlightHooks: preFlightHooks,
concurrencyHelper: NewConcurrencyHelper(concurrencyService, SSEPingFormatComment, pingInterval),
imageLimiter: &imageConcurrencyLimiter{},
maxAccountSwitches: maxAccountSwitches,
@@ -241,8 +245,8 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
setOpsRequestContext(c, reqModel, reqStream)
setOpsEndpointContext(c, "", int16(service.RequestTypeFromLegacy(reqStream, false)))
if decision := h.checkContentModeration(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIResponses, reqModel, body); decision != nil && decision.Blocked {
h.errorResponse(c, contentModerationStatus(decision), contentModerationErrorCode(decision), decision.Message)
if decision := h.runPreFlightHooks(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIResponses, reqModel, body); decision != nil && decision.Blocked {
h.errorResponse(c, preFlightStatus(decision), preFlightErrorCode(decision), decision.Message)
return
}
@@ -673,8 +677,8 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) {
setOpsRequestContext(c, reqModel, reqStream)
setOpsEndpointContext(c, "", int16(service.RequestTypeFromLegacy(reqStream, false)))
if decision := h.checkContentModeration(c, reqLog, apiKey, subject, service.ContentModerationProtocolAnthropicMessages, reqModel, body); decision != nil && decision.Blocked {
h.anthropicErrorResponse(c, contentModerationStatus(decision), contentModerationErrorCode(decision), decision.Message)
if decision := h.runPreFlightHooks(c, reqLog, apiKey, subject, service.ContentModerationProtocolAnthropicMessages, reqModel, body); decision != nil && decision.Blocked {
h.anthropicErrorResponse(c, preFlightStatus(decision), preFlightErrorCode(decision), decision.Message)
return
}
+2 -2
View File
@@ -86,8 +86,8 @@ func (h *OpenAIGatewayHandler) Images(c *gin.Context) {
h.errorResponse(c, http.StatusForbidden, "permission_error", service.ImageGenerationPermissionMessage())
return
}
if decision := h.checkContentModeration(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIImages, requestModel, parsed.ModerationBody()); decision != nil && decision.Blocked {
h.errorResponse(c, contentModerationStatus(decision), contentModerationErrorCode(decision), decision.Message)
if decision := h.runPreFlightHooks(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIImages, requestModel, parsed.ModerationBody()); decision != nil && decision.Blocked {
h.errorResponse(c, preFlightStatus(decision), preFlightErrorCode(decision), decision.Message)
return
}
imageReleaseFunc, acquired := h.acquireImageGenerationSlot(c, streamStarted)
@@ -55,6 +55,12 @@ type airwallexTokenState struct {
var airwallexAccessTokens sync.Map
func init() {
register(payment.TypeAirwallex, func(instanceID string, config map[string]string) (payment.Provider, error) {
return NewAirwallex(instanceID, config)
})
}
func NewAirwallex(instanceID string, config map[string]string) (*Airwallex, error) {
for _, k := range []string{"clientId", "apiKey", "webhookSecret", "apiBase"} {
if strings.TrimSpace(config[k]) == "" {
@@ -48,6 +48,12 @@ type Alipay struct {
client *alipay.Client
}
func init() {
register(payment.TypeAlipay, func(instanceID string, config map[string]string) (payment.Provider, error) {
return NewAlipay(instanceID, config)
})
}
// NewAlipay creates a new Alipay provider instance.
func NewAlipay(instanceID string, config map[string]string) (*Alipay, error) {
required := []string{"appId", "privateKey"}
@@ -41,6 +41,12 @@ type EasyPay struct {
// NewEasyPay creates a new EasyPay provider.
// config keys: pid, pkey, apiBase, notifyUrl, returnUrl, cid, cidAlipay, cidWxpay
func init() {
register(payment.TypeEasyPay, func(instanceID string, config map[string]string) (payment.Provider, error) {
return NewEasyPay(instanceID, config)
})
}
func NewEasyPay(instanceID string, config map[string]string) (*EasyPay, error) {
for _, k := range []string{"pid", "pkey", "apiBase", "notifyUrl", "returnUrl"} {
if strings.TrimSpace(config[k]) == "" {
+6 -12
View File
@@ -7,19 +7,13 @@ import (
)
// CreateProvider creates a Provider from a provider key, instance ID and decrypted config.
// Constructor errors (including *infraerrors.ApplicationError from wxpay) are returned
// as-is, never wrapped: the "_validate_" path relies on the structured error for
// frontend i18n (see service/payment_config_providers.go).
func CreateProvider(providerKey string, instanceID string, config map[string]string) (payment.Provider, error) {
switch providerKey {
case payment.TypeEasyPay:
return NewEasyPay(instanceID, config)
case payment.TypeAlipay:
return NewAlipay(instanceID, config)
case payment.TypeWxpay:
return NewWxpay(instanceID, config)
case payment.TypeStripe:
return NewStripe(instanceID, config)
case payment.TypeAirwallex:
return NewAirwallex(instanceID, config)
default:
fn, ok := constructors[providerKey]
if !ok {
return nil, fmt.Errorf("unknown provider key: %s", providerKey)
}
return fn(instanceID, config)
}
@@ -0,0 +1,26 @@
package provider
import (
"fmt"
"github.com/Wei-Shaw/sub2api/internal/payment"
)
// ConstructorFunc builds a payment.Provider from an instance ID and decrypted config.
type ConstructorFunc func(instanceID string, config map[string]string) (payment.Provider, error)
// constructors is the package-private provider constructor registry.
// It is populated exclusively by each provider file's init() via register()
// and is read-only afterwards, so concurrent CreateProvider calls
// (e.g. RefreshProviders' clear-and-reload loop) are safe and side-effect free.
var constructors = map[string]ConstructorFunc{}
// register adds a provider constructor under the given key.
// A duplicate key panics: that is an init-time instrumentation error
// which must surface as early as possible.
func register(key string, fn ConstructorFunc) {
if _, exists := constructors[key]; exists {
panic(fmt.Sprintf("payment/provider: register: provider key %q is already registered (duplicate registration)", key))
}
constructors[key] = fn
}
@@ -0,0 +1,141 @@
package provider
import (
"errors"
"fmt"
"strings"
"testing"
"github.com/Wei-Shaw/sub2api/internal/payment"
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
)
// TestCreateProviderUnknownKeyMessage 锁定 unknown-key 错误文案(等价性硬约束 B2):
// 必须与改造前 factory.go switch default 分支的文案逐字节一致。
func TestCreateProviderUnknownKeyMessage(t *testing.T) {
prov, err := CreateProvider("nope", "inst-1", map[string]string{})
if prov != nil {
t.Fatalf("expected nil provider for unknown key, got %T", prov)
}
if err == nil {
t.Fatal("expected error for unknown provider key")
}
const want = "unknown provider key: nope"
if got := err.Error(); got != want {
t.Fatalf("unknown-key error message changed:\n got: %q\nwant: %q", got, want)
}
}
// TestRegistryContainsExactlyAllProviderKeys 验证 5 个 provider 的 init 自注册全部生效,
// 且没有意外多注册的 key。
func TestRegistryContainsExactlyAllProviderKeys(t *testing.T) {
want := []string{
payment.TypeEasyPay,
payment.TypeAlipay,
payment.TypeWxpay,
payment.TypeStripe,
payment.TypeAirwallex,
}
if len(constructors) != len(want) {
t.Fatalf("expected %d registered provider keys, got %d: %v", len(want), len(constructors), registeredKeys())
}
for _, key := range want {
if _, ok := constructors[key]; !ok {
t.Errorf("provider key %q not registered", key)
}
}
}
func registeredKeys() []string {
keys := make([]string, 0, len(constructors))
for k := range constructors {
keys = append(keys, k)
}
return keys
}
// TestCreateProviderDispatchesToConstructors 验证每个已注册 key 都能分发到对应构造器:
// 空 config 下返回的错误必须是各 provider 自身的校验错误(原样透传,等价性硬约束 B1),
// 而不是 unknown-key 错误或任何包裹后的错误。
func TestCreateProviderDispatchesToConstructors(t *testing.T) {
cases := []struct {
key string
wantErr string
}{
{payment.TypeEasyPay, "easypay config missing required key: pid"},
{payment.TypeAlipay, "alipay config missing required key: appId"},
{payment.TypeStripe, "stripe config missing required key: secretKey"},
{payment.TypeAirwallex, "airwallex config missing required key: clientId"},
}
for _, tc := range cases {
_, err := CreateProvider(tc.key, "_validate_", map[string]string{})
if err == nil {
t.Errorf("key %q: expected constructor validation error, got nil", tc.key)
continue
}
if got := err.Error(); got != tc.wantErr {
t.Errorf("key %q: constructor error changed:\n got: %q\nwant: %q", tc.key, got, tc.wantErr)
}
}
}
// TestCreateProviderPassesThroughApplicationError 验证 wxpay 构造器返回的结构化
// *infraerrors.ApplicationError 被原样透传不包裹(等价性硬约束 B1):
// "_validate_" 路径依赖该结构化错误做前端 i18n(service/payment_config_providers.go)。
func TestCreateProviderPassesThroughApplicationError(t *testing.T) {
_, err := CreateProvider(payment.TypeWxpay, "_validate_", map[string]string{})
if err == nil {
t.Fatal("expected wxpay constructor validation error, got nil")
}
appErr, ok := err.(*infraerrors.ApplicationError)
if !ok {
var as *infraerrors.ApplicationError
if errors.As(err, &as) {
t.Fatalf("ApplicationError was wrapped instead of passed through as-is: %T", err)
}
t.Fatalf("expected *infraerrors.ApplicationError, got %T: %v", err, err)
}
if appErr.Reason != "WXPAY_CONFIG_MISSING_KEY" {
t.Errorf("unexpected reason: got %q, want %q", appErr.Reason, "WXPAY_CONFIG_MISSING_KEY")
}
if appErr.Metadata["key"] != "appId" {
t.Errorf("unexpected metadata key: got %q, want %q", appErr.Metadata["key"], "appId")
}
}
// TestCreateProviderSuccessPath 验证命中 key 且 config 合法时返回可用的 Provider 实例。
func TestCreateProviderSuccessPath(t *testing.T) {
prov, err := CreateProvider(payment.TypeStripe, "inst-42", map[string]string{"secretKey": "sk_test_123"})
if err != nil {
t.Fatalf("expected success, got error: %v", err)
}
if prov == nil {
t.Fatal("expected non-nil provider")
}
if got := prov.ProviderKey(); got != payment.TypeStripe {
t.Fatalf("unexpected provider key: got %q, want %q", got, payment.TypeStripe)
}
}
// TestRegisterDuplicateKeyPanics 验证重复注册同一 key 在 init 期立即 panic,
// 使插装错误尽早暴露。复用已注册的 stripe key,panic 发生在写入之前,
// 不会污染包级注册表状态(等价性硬约束 B3)。
func TestRegisterDuplicateKeyPanics(t *testing.T) {
before := len(constructors)
defer func() {
r := recover()
if r == nil {
t.Fatal("expected panic on duplicate registration")
}
msg := fmt.Sprint(r)
if !strings.Contains(msg, fmt.Sprintf("%q", payment.TypeStripe)) {
t.Fatalf("panic message should contain the duplicate key %q, got: %s", payment.TypeStripe, msg)
}
if len(constructors) != before {
t.Fatalf("duplicate registration mutated registry state: %d -> %d", before, len(constructors))
}
}()
register(payment.TypeStripe, func(string, map[string]string) (payment.Provider, error) {
return nil, nil
})
}
@@ -28,6 +28,12 @@ type Stripe struct {
sc *stripe.Client
}
func init() {
register(payment.TypeStripe, func(instanceID string, config map[string]string) (payment.Provider, error) {
return NewStripe(instanceID, config)
})
}
// NewStripe creates a new Stripe provider instance.
func NewStripe(instanceID string, config map[string]string) (*Stripe, error) {
if config["secretKey"] == "" {
@@ -82,6 +82,12 @@ type Wxpay struct {
const wxpayAPIv3KeyLength = 32
func init() {
register(payment.TypeWxpay, func(instanceID string, config map[string]string) (payment.Provider, error) {
return NewWxpay(instanceID, config)
})
}
func NewWxpay(instanceID string, config map[string]string) (*Wxpay, error) {
// All fields are required. Platform-certificate mode is intentionally unsupported —
// WeChat has been migrating all merchants to the pubkey verifier since 2024-10,