Merge remote-tracking branch 'origin/main' into feature/batch-image-foundation

# Conflicts:
#	.github/audit-exceptions.yml
This commit is contained in:
Turtle_Li
2026-07-07 10:53:17 +08:00
33 changed files with 2201 additions and 337 deletions
+2 -2
View File
@@ -5,14 +5,14 @@ exceptions:
severity: high
reason: "Admin export only; switched to dynamic import to reduce exposure (CVE-2023-30533)"
mitigation: "Load only on export; restrict export permissions and data scope"
expires_on: "2026-08-06"
expires_on: "2026-10-06"
owner: "security@your-domain"
- package: xlsx
advisory: "GHSA-5pgg-2g8v-p4x9"
severity: high
reason: "Admin export only; switched to dynamic import to reduce exposure (CVE-2024-22363)"
mitigation: "Load only on export; restrict export permissions and data scope"
expires_on: "2026-08-06"
expires_on: "2026-10-06"
owner: "security@your-domain"
- package: lodash
advisory: "GHSA-r5fr-rjxr-66jc"
+109 -7
View File
@@ -19,6 +19,7 @@ const (
EndpointChatCompletions = "/v1/chat/completions"
EndpointEmbeddings = "/v1/embeddings"
EndpointResponses = "/v1/responses"
EndpointResponsesCompact = "/v1/responses/compact"
EndpointImagesGenerations = "/v1/images/generations"
EndpointImagesEdits = "/v1/images/edits"
EndpointVideosGenerations = "/v1/videos/generations"
@@ -42,6 +43,33 @@ const (
// "/v1/chat/completions" → "/v1/chat/completions"
// "/openai/v1/responses/foo" → "/v1/responses"
// "/v1beta/models/gemini:gen" → "/v1beta/models"
//
// The OpenAI Responses API is also exposed via a few bare/alias
// routes that do not carry a "/v1/" prefix (top-level bare route and
// the Codex direct route). "/responses/compact" (and "/backend-api/
// codex/responses/compact") is a distinct client endpoint — the
// "compact" client — and is normalized to its OWN canonical inbound
// endpoint, EndpointResponsesCompact, rather than being folded into
// the root Responses endpoint. Any other subpath under the bare/alias
// roots (i.e. not "compact" itself or nested under it) remains a
// subresource suffix of the root Responses endpoint:
//
// "/v1/responses/compact" → EndpointResponsesCompact
// "/v1/responses/compact/detail" → EndpointResponsesCompact
// "/openai/v1/responses/compact" → EndpointResponsesCompact
// "/openai/v1/responses/compact/detail" → EndpointResponsesCompact
// "/responses/compact" → EndpointResponsesCompact
// "/responses/compact/detail" → EndpointResponsesCompact
// "/backend-api/codex/responses/compact" → EndpointResponsesCompact
// "/backend-api/codex/responses/compact/detail" → EndpointResponsesCompact
// "/v1/responses" → EndpointResponses
// "/openai/v1/responses" → EndpointResponses
// "/responses" → EndpointResponses
// "/backend-api/codex/responses" → EndpointResponses
//
// The compact check MUST be evaluated before the root Responses check,
// otherwise "/v1/responses" (a prefix of "/v1/responses/compact")
// would erroneously match first.
func NormalizeInboundEndpoint(path string) string {
path = strings.TrimSpace(path)
switch {
@@ -59,7 +87,9 @@ func NormalizeInboundEndpoint(path string) string {
return EndpointVideosGenerations
case strings.Contains(path, EndpointVideos) || strings.Contains(path, "/videos/"):
return EndpointVideos
case strings.Contains(path, EndpointResponses):
case strings.Contains(path, EndpointResponsesCompact) || isResponsesCompactAliasPath(path):
return EndpointResponsesCompact
case strings.Contains(path, EndpointResponses) || isResponsesRootAliasPath(path):
return EndpointResponses
case strings.Contains(path, EndpointGeminiModels):
return EndpointGeminiModels
@@ -68,6 +98,59 @@ func NormalizeInboundEndpoint(path string) string {
}
}
// isResponsesCompactAliasPath reports whether path is the bare/alias
// "compact" client endpoint — i.e. it is rooted at "/responses/compact"
// or "/backend-api/codex/responses/compact" (bare routes that serve
// the OpenAI Responses API "compact" client without a "/v1/" prefix),
// or any subpath nested under either of those roots:
//
// - "/responses/compact" (bare route, compact client)
// - "/responses/compact/*subpath" (nested, e.g. "/responses/compact/detail")
// - "/backend-api/codex/responses/compact" (Codex direct route, compact client)
// - "/backend-api/codex/responses/compact/*subpath" (nested, e.g.
// "/backend-api/codex/responses/compact/detail")
//
// This MUST be checked before isResponsesRootAliasPath, since
// "/responses" is a prefix of "/responses/compact".
func isResponsesCompactAliasPath(path string) bool {
trimmed := strings.TrimRight(strings.TrimSpace(path), "/")
if trimmed == "" {
return false
}
return isBareOrSubpathOf(trimmed, "/responses/compact") || isBareOrSubpathOf(trimmed, "/backend-api/codex/responses/compact")
}
// isResponsesRootAliasPath reports whether path is one of the bare/alias
// routes that serve the root OpenAI Responses API without a "/v1/"
// prefix, or any non-"compact" subpath registered under them:
//
// - "/responses" (top-level bare route)
// - "/responses/*subpath" (any subpath other than "compact",
// since "compact" is its own distinct inbound endpoint)
// - "/backend-api/codex/responses" (Codex direct route)
// - "/backend-api/codex/responses/*subpath" (any subpath other than
// "compact")
//
// Only the top-level bare route and the Codex direct route (and their
// subpaths) are recognized here — this deliberately does NOT generalize
// to any path merely ending in "/responses" (e.g. an unrelated
// "/foo/responses" must not match).
func isResponsesRootAliasPath(path string) bool {
trimmed := strings.TrimRight(strings.TrimSpace(path), "/")
if trimmed == "" {
return false
}
return isBareOrSubpathOf(trimmed, "/responses") || isBareOrSubpathOf(trimmed, "/backend-api/codex/responses")
}
// isBareOrSubpathOf reports whether path is exactly root, or a subpath
// rooted at root (i.e. root followed by "/"). This anchors the match
// at the start of path so it cannot match paths where root appears
// nested under some other unrelated prefix.
func isBareOrSubpathOf(path, root string) bool {
return path == root || strings.HasPrefix(path, root+"/")
}
// DeriveUpstreamEndpoint determines the upstream endpoint from the
// account platform and the normalized inbound endpoint.
//
@@ -88,10 +171,20 @@ func DeriveUpstreamEndpoint(inbound, rawRequestPath, platform string) string {
return inbound
}
// OpenAI forwards everything to the Responses API.
// Preserve subresource suffix (e.g. /v1/responses/compact).
// Preserve subresource suffix (e.g. /v1/responses/compact,
// /v1/responses/compact/detail) as derived from the raw path.
if suffix := responsesSubpathSuffix(rawRequestPath); suffix != "" {
return EndpointResponses + suffix
}
// The raw path carried no derivable suffix (e.g. it was already
// normalized upstream, or the caller only has the canonical
// inbound endpoint available) — fall back to the canonical
// compact endpoint when that's what the inbound request was
// recognized as, so it isn't silently treated as the root
// Responses endpoint.
if inbound == EndpointResponsesCompact {
return EndpointResponsesCompact
}
return EndpointResponses
case service.PlatformAnthropic:
@@ -142,10 +235,13 @@ func responsesSubpathSuffix(rawPath string) string {
// Apply this middleware to all gateway route groups.
func InboundEndpointMiddleware() gin.HandlerFunc {
return func(c *gin.Context) {
path := c.FullPath()
if path == "" && c.Request != nil && c.Request.URL != nil {
path := ""
if c.Request != nil && c.Request.URL != nil {
path = c.Request.URL.Path
}
if path == "" {
path = c.FullPath()
}
c.Set(ctxKeyInboundEndpoint, NormalizeInboundEndpoint(path))
c.Next()
}
@@ -158,7 +254,11 @@ func InboundEndpointMiddleware() gin.HandlerFunc {
// GetInboundEndpoint returns the canonical inbound endpoint stored by
// InboundEndpointMiddleware. If the middleware did not run (e.g. in
// tests), it falls back to normalizing c.FullPath() on the fly.
// tests), it falls back to normalizing c.Request.URL.Path on the fly
// (preferring the raw request path over c.FullPath(), which collapses
// wildcard route patterns such as "/v1/responses/*subpath" and would
// otherwise mis-normalize concrete requests like "/v1/responses/compact"
// to the root Responses endpoint).
func GetInboundEndpoint(c *gin.Context) string {
if v, ok := c.Get(ctxKeyInboundEndpoint); ok {
if s, ok := v.(string); ok && s != "" {
@@ -168,10 +268,12 @@ func GetInboundEndpoint(c *gin.Context) string {
// Fallback: normalize on the fly.
path := ""
if c != nil {
path = c.FullPath()
if path == "" && c.Request != nil && c.Request.URL != nil {
if c.Request != nil && c.Request.URL != nil {
path = c.Request.URL.Path
}
if path == "" {
path = c.FullPath()
}
}
return NormalizeInboundEndpoint(path)
}
+178 -8
View File
@@ -26,23 +26,42 @@ func TestNormalizeInboundEndpoint(t *testing.T) {
{"/v1/chat/completions", EndpointChatCompletions},
{"/v1/embeddings", EndpointEmbeddings},
{"/v1/responses", EndpointResponses},
{"/v1/responses/compact", EndpointResponsesCompact},
{"/v1/responses/compact/detail", EndpointResponsesCompact},
{"/v1/images/generations", EndpointImagesGenerations},
{"/v1/images/edits", EndpointImagesEdits},
{"/v1/videos/generations", EndpointVideosGenerations},
{"/v1/videos/req_123", EndpointVideos},
{"/v1beta/models", EndpointGeminiModels},
// Prefixed paths (antigravity, openai).
// Prefixed paths (antigravity, openai) — root Responses.
{"/antigravity/v1/messages", EndpointMessages},
{"/openai/v1/responses", EndpointResponses},
{"/openai/v1/responses/compact", EndpointResponses},
{"/openai/v1/images/generations", EndpointImagesGenerations},
{"/openai/v1/images/edits", EndpointImagesEdits},
{"/antigravity/v1beta/models/gemini:generateContent", EndpointGeminiModels},
// Gin route patterns with wildcards.
{"/v1beta/models/*modelAction", EndpointGeminiModels},
{"/v1/responses/*subpath", EndpointResponses},
// Prefixed paths — "/responses/compact" is its OWN distinct
// inbound endpoint, not folded into the root Responses endpoint.
{"/openai/v1/responses/compact", EndpointResponsesCompact},
{"/openai/v1/responses/compact/detail", EndpointResponsesCompact},
// Bare top-level alias route "/responses" — root vs. compact.
{"/responses", EndpointResponses},
{"/responses/compact", EndpointResponsesCompact},
{"/responses/compact/detail", EndpointResponsesCompact},
// Bare Codex direct alias route — root vs. compact.
{"/backend-api/codex/responses", EndpointResponses},
{"/backend-api/codex/responses/compact", EndpointResponsesCompact},
{"/backend-api/codex/responses/compact/detail", EndpointResponsesCompact},
// Must NOT generalize to arbitrary paths merely ending in
// "/responses" (or "/responses/compact") that are unrelated to
// the two known bare alias roots, unless they already carry a
// supported "/v1/responses..." prefix form.
{"/foo/responses", "/foo/responses"},
{"/foo/responses/compact", "/foo/responses/compact"},
// Unknown path is returned as-is.
{"/v1/embeddings", "/v1/embeddings"},
@@ -74,10 +93,29 @@ func TestDeriveUpstreamEndpoint(t *testing.T) {
// Gemini.
{"gemini models", EndpointGeminiModels, "/v1beta/models/gemini:gen", service.PlatformGemini, EndpointGeminiModels},
// OpenAI — always /v1/responses.
// OpenAI — root Responses.
{"openai responses root", EndpointResponses, "/v1/responses", service.PlatformOpenAI, EndpointResponses},
{"openai responses compact", EndpointResponses, "/openai/v1/responses/compact", service.PlatformOpenAI, "/v1/responses/compact"},
{"openai responses nested", EndpointResponses, "/openai/v1/responses/compact/detail", service.PlatformOpenAI, "/v1/responses/compact/detail"},
// OpenAI — compact, raw path carries the derivable "/compact"
// (or nested) suffix, which must be preserved on the upstream
// endpoint.
{"openai responses compact", EndpointResponsesCompact, "/openai/v1/responses/compact", service.PlatformOpenAI, "/v1/responses/compact"},
{"openai responses nested", EndpointResponsesCompact, "/openai/v1/responses/compact/detail", service.PlatformOpenAI, "/v1/responses/compact/detail"},
{"openai bare responses compact", EndpointResponsesCompact, "/responses/compact", service.PlatformOpenAI, "/v1/responses/compact"},
{"openai bare responses compact detail", EndpointResponsesCompact, "/responses/compact/detail", service.PlatformOpenAI, "/v1/responses/compact/detail"},
{"openai codex direct responses compact", EndpointResponsesCompact, "/backend-api/codex/responses/compact", service.PlatformOpenAI, "/v1/responses/compact"},
{"openai codex direct responses compact detail", EndpointResponsesCompact, "/backend-api/codex/responses/compact/detail", service.PlatformOpenAI, "/v1/responses/compact/detail"},
// OpenAI — bare root alias routes normalize to root Responses.
{"openai bare responses", EndpointResponses, "/responses", service.PlatformOpenAI, EndpointResponses},
{"openai codex direct responses", EndpointResponses, "/backend-api/codex/responses", service.PlatformOpenAI, EndpointResponses},
// OpenAI — inbound is already the canonical compact endpoint but
// the raw path carries no derivable "/responses..." suffix (e.g.
// it was already normalized upstream). Must not silently fall
// back to the root Responses endpoint.
{"openai responses compact inbound only, unrelated raw path", EndpointResponsesCompact, "/v1/messages", service.PlatformOpenAI, EndpointResponsesCompact},
{"openai from messages", EndpointMessages, "/v1/messages", service.PlatformOpenAI, EndpointResponses},
{"openai from completions", EndpointChatCompletions, "/v1/chat/completions", service.PlatformOpenAI, EndpointResponses},
{"openai embeddings", EndpointEmbeddings, "/v1/embeddings", service.PlatformOpenAI, EndpointEmbeddings},
@@ -113,6 +151,12 @@ func TestResponsesSubpathSuffix(t *testing.T) {
{"/v1/responses/", ""},
{"/v1/responses/compact", "/compact"},
{"/openai/v1/responses/compact/detail", "/compact/detail"},
{"/responses", ""},
{"/responses/compact", "/compact"},
{"/responses/compact/detail", "/compact/detail"},
{"/backend-api/codex/responses", ""},
{"/backend-api/codex/responses/compact", "/compact"},
{"/backend-api/codex/responses/compact/detail", "/compact/detail"},
{"/v1/messages", ""},
{"", ""},
}
@@ -154,6 +198,132 @@ func TestGetInboundEndpoint_FallbackWithoutMiddleware(t *testing.T) {
require.Equal(t, EndpointMessages, got)
}
// TestInboundEndpointMiddleware_WildcardRoutes verifies that, when a
// gateway route is registered with a Gin wildcard pattern (e.g.
// "/v1/responses/*subpath"), InboundEndpointMiddleware normalizes based
// on the concrete request path (c.Request.URL.Path) rather than the
// route pattern (c.FullPath()). Using c.FullPath() here would collapse
// every request under the wildcard — including "/v1/responses/compact"
// — down to the literal pattern string, which never matches the
// "compact" alias detection and would incorrectly normalize to the root
// Responses endpoint.
func TestInboundEndpointMiddleware_WildcardRoutes(t *testing.T) {
tests := []struct {
name string
routePath string
requestPath string
want string
}{
{
name: "v1 responses wildcard route, compact request",
routePath: "/v1/responses/*subpath",
requestPath: "/v1/responses/compact",
want: EndpointResponsesCompact,
},
{
name: "bare responses wildcard route, compact request",
routePath: "/responses/*subpath",
requestPath: "/responses/compact",
want: EndpointResponsesCompact,
},
{
name: "codex direct wildcard route, compact request",
routePath: "/backend-api/codex/responses/*subpath",
requestPath: "/backend-api/codex/responses/compact",
want: EndpointResponsesCompact,
},
{
name: "v1 responses wildcard route, non-compact subpath request",
routePath: "/v1/responses/*subpath",
requestPath: "/v1/responses/foo",
want: EndpointResponses,
},
{
name: "bare responses wildcard route, non-compact subpath request",
routePath: "/responses/*subpath",
requestPath: "/responses/foo",
want: EndpointResponses,
},
{
name: "codex direct wildcard route, non-compact subpath request",
routePath: "/backend-api/codex/responses/*subpath",
requestPath: "/backend-api/codex/responses/foo",
want: EndpointResponses,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
router := gin.New()
router.Use(InboundEndpointMiddleware())
var captured string
router.POST(tt.routePath, func(c *gin.Context) {
captured = GetInboundEndpoint(c)
c.Status(http.StatusOK)
})
req := httptest.NewRequest(http.MethodPost, tt.requestPath, nil)
rec := httptest.NewRecorder()
router.ServeHTTP(rec, req)
require.Equal(t, http.StatusOK, rec.Code)
require.Equal(t, tt.want, captured)
})
}
}
// TestInboundEndpointMiddleware_GeminiWildcardRoute verifies that a Gemini
// wildcard route (e.g. "/v1beta/models/*modelAction", used to capture the
// ":generateContent"-style action suffix embedded in the path) is normalized
// to EndpointGeminiModels via InboundEndpointMiddleware, using the same real
// Gin routing path as TestInboundEndpointMiddleware_WildcardRoutes above.
func TestInboundEndpointMiddleware_GeminiWildcardRoute(t *testing.T) {
router := gin.New()
router.Use(InboundEndpointMiddleware())
var captured string
router.POST("/v1beta/models/*modelAction", func(c *gin.Context) {
captured = GetInboundEndpoint(c)
c.Status(http.StatusOK)
})
req := httptest.NewRequest(http.MethodPost, "/v1beta/models/gemini-2.5-pro:generateContent", nil)
rec := httptest.NewRecorder()
router.ServeHTTP(rec, req)
require.Equal(t, http.StatusOK, rec.Code)
require.Equal(t, EndpointGeminiModels, captured)
}
// TestGetInboundEndpoint_FallbackWildcardRouteWithoutMiddleware verifies
// that when InboundEndpointMiddleware did NOT run (so no value is stored
// in gin.Context), the GetInboundEndpoint fallback path still prefers
// c.Request.URL.Path over c.FullPath(). This guards against the fallback
// regressing to prefer c.FullPath() again, which would misnormalize
// concrete requests matched by a wildcard route pattern (e.g.
// "/v1/responses/*subpath" matching "/v1/responses/compact") down to
// the root Responses endpoint.
func TestGetInboundEndpoint_FallbackWildcardRouteWithoutMiddleware(t *testing.T) {
router := gin.New()
// Deliberately do NOT register InboundEndpointMiddleware.
var captured string
router.POST("/v1/responses/*subpath", func(c *gin.Context) {
// Sanity check: FullPath returns the route pattern, not the
// concrete request path, when a wildcard route matches.
require.Equal(t, "/v1/responses/*subpath", c.FullPath())
captured = GetInboundEndpoint(c)
c.Status(http.StatusOK)
})
req := httptest.NewRequest(http.MethodPost, "/v1/responses/compact", nil)
rec := httptest.NewRecorder()
router.ServeHTTP(rec, req)
require.Equal(t, http.StatusOK, rec.Code)
require.Equal(t, EndpointResponsesCompact, captured)
}
func TestGetUpstreamEndpoint_FullFlow(t *testing.T) {
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
+12 -7
View File
@@ -85,22 +85,27 @@ func writeResponsesFailedSSE(c *gin.Context, errType, message string) bool {
return true
}
// inboundIsResponses 判断当前请求是否落在任何 /responses 路由上。
// inboundIsResponses 判断当前请求是否落在任意 Responses 路由上
// (不区分 root 还是 compact 变体)。
//
// 不能直接用 GetInboundEndpoint(c) == EndpointResponses 比较,因为
// NormalizeInboundEndpoint 只识别包含 "/v1/responses" 子串的路径;
// 项目里实际注册了多组路由(gateway_v1、top-level bare、codex direct),
// 其中 r.POST("/responses", ...) 和 codexDirect.POST("/responses", ...)
// 的 c.FullPath() 不含 "/v1/" 前缀,会被归一化为原始路径,
// 导致协议合规终止事件没法发出去。
// GetInboundEndpoint/NormalizeInboundEndpoint 会把 compact 变体归一化为
// 单独的 EndpointResponsesCompact(而不是 EndpointResponses),
// 而本函数在这里只关心“是不是 Responses 家族的请求”,
// 不需要区分 root/compact,所以不能用那个等值比较。
//
// 这里用 FullPath 的后缀判断,覆盖所有变体:
// 这里改用 FullPath 的后缀/子串判断,一次性覆盖 root 和 compact 的所有变体:
// - /v1/responses
// - /v1/responses/compact
// - /responses
// - /responses/compact
// - /backend-api/codex/responses
// - /backend-api/codex/responses/compact
//
// 对于通配路由(如 "/v1/responses/*action")注册的 FullPath 本身就带有
// "/responses/" 子串(例如 "/v1/responses/*action"),所以下面的
// strings.Contains(p, "/responses/") 分支同样能覆盖这些通配路由,
// 不需要额外处理通配符本身。
func inboundIsResponses(c *gin.Context) bool {
if c == nil {
return false
+461 -132
View File
@@ -6,6 +6,7 @@ import (
"fmt"
"strconv"
"github.com/Wei-Shaw/sub2api/internal/pkg/logger"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/redis/go-redis/v9"
)
@@ -36,6 +37,19 @@ const (
// 默认槽位过期时间(分钟),可通过配置覆盖
defaultSlotTTLMinutes = 15
// 活跃索引用来替代后台任务全量 SCAN 槽位键。
// member 是账号/用户 ID,score 是“预计仍需关注到”的 Redis Unix 秒时间戳。
accountActiveIndexKey = "concurrency:account:active_index" // ZSET member=accountID, score=expireAtUnixSeconds
userActiveIndexKey = "concurrency:user:active_index" // ZSET member=userID, score=expireAtUnixSeconds
// 后台清理只按批处理索引候选,避免单次任务占用 Redis 太久。
activeIndexCleanupBatchSize = 1000
activeIndexPipelineChunkSize = 500
// 一次性迁移 marker:活跃索引机制上线前遗留的等待计数键无法被索引发现,
// 且有流量时 TTL 会被不断刷新,必须清扫一次。marker 存在即代表已完成。
legacyWaitSweepMarkerKey = "concurrency:startup:legacy_wait_sweep:v1"
)
var (
@@ -45,6 +59,7 @@ var (
// ARGV[1] = maxConcurrency
// ARGV[2] = TTL(秒)
// ARGV[3] = requestID
// 返回 {是否成功, Redis 当前秒},Go 侧复用同一时间源写活跃索引,省去额外 TIME 往返。
acquireScript = redis.NewScript(`
-- Redis 3.2-4.x compat: opt into effects replication so redis.call('TIME')
-- replicates correctly. No-op on Redis 5.0+ (effects replication is default).
@@ -67,7 +82,7 @@ var (
if exists ~= false then
redis.call('ZADD', key, now, requestID)
redis.call('EXPIRE', key, ttl)
return 1
return {1, now}
end
-- 检查是否达到并发上限
@@ -75,10 +90,10 @@ var (
if count < maxConcurrency then
redis.call('ZADD', key, now, requestID)
redis.call('EXPIRE', key, ttl)
return 1
return {1, now}
end
return 0
return {0, now}
`)
// getCountScript 统计有序集合中的槽位数量并清理过期条目
@@ -127,46 +142,56 @@ var (
// KEYS[1] = wait queue key
// ARGV[1] = maxWait
// ARGV[2] = TTL in seconds
// 返回 {是否成功, Redis 当前秒},供 Go 侧免额外 TIME 往返写活跃索引。
incrementWaitScript = redis.NewScript(`
-- Redis 3.2-4.x compat: opt into effects replication so redis.call('TIME')
-- replicates correctly. No-op on Redis 5.0+ (effects replication is default).
redis.replicate_commands()
local current = redis.call('GET', KEYS[1])
if current == false then
current = 0
else
current = tonumber(current)
end
local now = tonumber(redis.call('TIME')[1])
if current >= tonumber(ARGV[1]) then
return 0
return {0, now}
end
local newVal = redis.call('INCR', KEYS[1])
redis.call('INCR', KEYS[1])
-- Refresh TTL so long-running traffic doesn't expire active queue counters.
redis.call('EXPIRE', KEYS[1], ARGV[2])
return 1
`)
return {1, now}
`)
// incrementAccountWaitScript - account-level wait queue count (refresh TTL on each increment)
// 返回值同 incrementWaitScript:{是否成功, Redis 当前秒}。
incrementAccountWaitScript = redis.NewScript(`
local current = redis.call('GET', KEYS[1])
if current == false then
current = 0
else
current = tonumber(current)
end
-- Redis 3.2-4.x compat: opt into effects replication so redis.call('TIME')
-- replicates correctly. No-op on Redis 5.0+ (effects replication is default).
redis.replicate_commands()
local current = redis.call('GET', KEYS[1])
if current == false then
current = 0
else
current = tonumber(current)
end
local now = tonumber(redis.call('TIME')[1])
if current >= tonumber(ARGV[1]) then
return 0
end
if current >= tonumber(ARGV[1]) then
return {0, now}
end
local newVal = redis.call('INCR', KEYS[1])
redis.call('INCR', KEYS[1])
-- Refresh TTL so long-running traffic doesn't expire active queue counters.
redis.call('EXPIRE', KEYS[1], ARGV[2])
-- Refresh TTL so long-running traffic doesn't expire active queue counters.
redis.call('EXPIRE', KEYS[1], ARGV[2])
return 1
`)
return {1, now}
`)
// decrementWaitScript - same as before
decrementWaitScript = redis.NewScript(`
@@ -198,51 +223,27 @@ var (
return 1
`)
// cleanupExpiredSlotKeysScript 批量清理实际存在的账号槽位键,避免后台任务从数据库加载全量账号。
// KEYS = 有序集合键列表,ARGV[1] = TTL(秒)。
cleanupExpiredSlotKeysScript = redis.NewScript(`
-- Redis 3.2-4.x compat: opt into effects replication so redis.call('TIME')
-- replicates correctly. No-op on Redis 5.0+ (effects replication is default).
redis.replicate_commands()
local ttl = tonumber(ARGV[1])
local timeResult = redis.call('TIME')
local now = tonumber(timeResult[1])
local expireBefore = now - ttl
local removed = 0
for i = 1, #KEYS do
local key = KEYS[i]
removed = removed + redis.call('ZREMRANGEBYSCORE', key, '-inf', expireBefore)
if redis.call('ZCARD', key) == 0 then
redis.call('DEL', key)
else
redis.call('EXPIRE', key, ttl)
end
end
return removed
`)
// startupCleanupScript 清理非当前进程前缀的槽位成员。
// KEYS 是有序集合键列表,ARGV[1] 是当前进程前缀,ARGV[2] 是槽位 TTL。
// 遍历每个 KEYS[i],移除前缀不匹配的成员,清空后删 key,否则刷新 EXPIRE。
startupCleanupScript = redis.NewScript(`
// startupCleanupSlotScript 清理单个槽位 key 中非当前进程前缀的成员,避免 Redis Cluster CROSSSLOT。
// KEYS[1] 是有序集合键,ARGV[1] 是当前进程前缀,ARGV[2] 是槽位 TTL。
// 返回 {清除数量, 剩余成员数},Go 侧据剩余数决定索引 member 去留,无需再回读槽位。
startupCleanupSlotScript = redis.NewScript(`
local key = KEYS[1]
local activePrefix = ARGV[1]
local slotTTL = tonumber(ARGV[2])
local removed = 0
for i = 1, #KEYS do
local key = KEYS[i]
local members = redis.call('ZRANGE', key, 0, -1)
for _, member in ipairs(members) do
if string.sub(member, 1, string.len(activePrefix)) ~= activePrefix then
removed = removed + redis.call('ZREM', key, member)
end
end
if redis.call('ZCARD', key) == 0 then
redis.call('DEL', key)
else
redis.call('EXPIRE', key, slotTTL)
local members = redis.call('ZRANGE', key, 0, -1)
for _, member in ipairs(members) do
if string.sub(member, 1, string.len(activePrefix)) ~= activePrefix then
removed = removed + redis.call('ZREM', key, member)
end
end
return removed
local remaining = redis.call('ZCARD', key)
if remaining == 0 then
redis.call('DEL', key)
else
redis.call('EXPIRE', key, slotTTL)
end
return {removed, remaining}
`)
)
@@ -290,21 +291,242 @@ func accountWaitKey(accountID int64) string {
return fmt.Sprintf("%s%d", accountWaitKeyPrefix, accountID)
}
// redisUnixSeconds 统一使用 Redis 服务器时间,避免多实例本地时钟漂移导致索引提前/延后过期。
func (c *concurrencyCache) redisUnixSeconds(ctx context.Context) (int64, error) {
now, err := c.rdb.Time(ctx).Result()
if err != nil {
return 0, fmt.Errorf("redis TIME: %w", err)
}
return now.Unix(), nil
}
// slotIndexSpec 描述一个活跃索引及其对应的槽位/等待键构造方式。
// 用具名字段避免把 slotKey/waitKey 两个同签名函数按位置传参时写反。
type slotIndexSpec struct {
indexKey string
slotKey func(int64) string
waitKey func(int64) string
}
var (
accountSlotIndex = slotIndexSpec{indexKey: accountActiveIndexKey, slotKey: accountSlotKey, waitKey: accountWaitKey}
userSlotIndex = slotIndexSpec{indexKey: userActiveIndexKey, slotKey: userSlotKey, waitKey: waitQueueKey}
)
// touchActiveIndexAt 是写路径上的轻量标记:主操作已成功时,尽力把 ID 放入活跃索引,
// score 为给定的绝对过期时间(Redis Unix 秒)。索引失败不影响并发槽位/等待队列本身,
// 后续释放或清理会再次校正,因此只记日志不上抛。
func (c *concurrencyCache) touchActiveIndexAt(ctx context.Context, indexKey string, id int64, expireAt int64) {
if c == nil || c.rdb == nil || id <= 0 || expireAt <= 0 {
return
}
if err := c.rdb.ZAdd(ctx, indexKey, redis.Z{
Score: float64(expireAt),
Member: strconv.FormatInt(id, 10),
}).Err(); err != nil {
logger.LegacyPrintf("repository.concurrency", "Warning: touch active index %s for %d failed: %v", indexKey, id, err)
}
}
func (c *concurrencyCache) refreshAccountActiveIndex(ctx context.Context, accountID int64) {
c.refreshActiveIndex(ctx, accountActiveIndexKey, accountID, accountSlotKey(accountID), accountWaitKey(accountID))
}
func (c *concurrencyCache) refreshUserActiveIndex(ctx context.Context, userID int64) {
c.refreshActiveIndex(ctx, userActiveIndexKey, userID, userSlotKey(userID), waitQueueKey(userID))
}
// refreshActiveIndex 以 Redis 中的真实槽位/等待数为准重建索引状态。
// 释放槽位、等待计数减少、清理过期成员后都会调用它,防止索引残留。
// 索引维护是 best-effort:失败只记日志,不影响主流程。
func (c *concurrencyCache) refreshActiveIndex(ctx context.Context, indexKey string, id int64, slotKey, waitKey string) {
if c == nil || c.rdb == nil || id <= 0 {
return
}
now, err := c.redisUnixSeconds(ctx)
if err != nil {
logger.LegacyPrintf("repository.concurrency", "Warning: refresh active index %s for %d failed: %v", indexKey, id, err)
return
}
load, err := c.readActiveLoadForKey(ctx, id, slotKey, waitKey, now)
if err != nil {
logger.LegacyPrintf("repository.concurrency", "Warning: refresh active index %s for %d failed: %v", indexKey, id, err)
return
}
member := strconv.FormatInt(id, 10)
if load.slotCount == 0 && load.waitCount <= 0 {
if err := c.rdb.ZRem(ctx, indexKey, member).Err(); err != nil {
logger.LegacyPrintf("repository.concurrency", "Warning: remove active index member %s from %s failed: %v", member, indexKey, err)
}
return
}
ttlSeconds := c.activeIndexTTL(load.slotCount, load.waitCount)
if ttlSeconds <= 0 {
return
}
c.touchActiveIndexAt(ctx, indexKey, id, now+int64(ttlSeconds))
}
type activeIndexLoad struct {
id int64
member string
slotCount int
waitCount int
}
// activeIndexTTL 取槽位 TTL 与等待队列 TTL 中仍然需要关注的较大值。
// 只要并发槽位或等待计数还有负载,就保留索引;两者都为 0 时调用方会删除索引。
func (c *concurrencyCache) activeIndexTTL(slotCount int, waitCount int) int {
ttlSeconds := 0
if slotCount > 0 {
ttlSeconds = c.slotTTLSeconds
}
if waitCount > 0 && c.waitQueueTTLSeconds > ttlSeconds {
ttlSeconds = c.waitQueueTTLSeconds
}
return ttlSeconds
}
// readActiveLoadForKey 读取单个 ID 的当前负载,并顺手清理该槽位集合中的过期成员。
func (c *concurrencyCache) readActiveLoadForKey(ctx context.Context, id int64, slotKey, waitKey string, now int64) (activeIndexLoad, error) {
cutoffTime := now - int64(c.slotTTLSeconds)
pipe := c.rdb.Pipeline()
pipe.ZRemRangeByScore(ctx, slotKey, "-inf", strconv.FormatInt(cutoffTime, 10))
zcardCmd := pipe.ZCard(ctx, slotKey)
getCmd := pipe.Get(ctx, waitKey)
if _, err := pipe.Exec(ctx); err != nil && !errors.Is(err, redis.Nil) {
return activeIndexLoad{}, fmt.Errorf("pipeline exec: %w", err)
}
waitCount := 0
if v, err := getCmd.Int(); err == nil && v > 0 {
waitCount = v
}
return activeIndexLoad{
id: id,
member: strconv.FormatInt(id, 10),
slotCount: int(zcardCmd.Val()),
waitCount: waitCount,
}, nil
}
// readIndexLoads 批量读取索引候选的真实负载(账号/用户通用)。
// 分块 Pipeline 可以减少 Redis 往返,同时避免一次 Pipeline 塞入过多命令。
func (c *concurrencyCache) readIndexLoads(ctx context.Context, spec slotIndexSpec, members []string, now int64) ([]activeIndexLoad, []string, error) {
loads := make([]activeIndexLoad, 0, len(members))
staleMembers := make([]string, 0)
candidates := make([]activeIndexLoad, 0, len(members))
for _, member := range members {
id, err := strconv.ParseInt(member, 10, 64)
if err != nil || id <= 0 {
staleMembers = append(staleMembers, member)
continue
}
candidates = append(candidates, activeIndexLoad{id: id, member: member})
}
cutoffTime := now - int64(c.slotTTLSeconds)
for start := 0; start < len(candidates); start += activeIndexPipelineChunkSize {
end := start + activeIndexPipelineChunkSize
if end > len(candidates) {
end = len(candidates)
}
chunk := candidates[start:end]
pipe := c.rdb.Pipeline()
type loadCmd struct {
activeIndexLoad
zcardCmd *redis.IntCmd
getCmd *redis.StringCmd
}
cmds := make([]loadCmd, 0, len(chunk))
for _, candidate := range chunk {
slotKey := spec.slotKey(candidate.id)
waitKey := spec.waitKey(candidate.id)
pipe.ZRemRangeByScore(ctx, slotKey, "-inf", strconv.FormatInt(cutoffTime, 10))
cmds = append(cmds, loadCmd{
activeIndexLoad: candidate,
zcardCmd: pipe.ZCard(ctx, slotKey),
getCmd: pipe.Get(ctx, waitKey),
})
}
if _, err := pipe.Exec(ctx); err != nil && !errors.Is(err, redis.Nil) {
return nil, nil, fmt.Errorf("pipeline exec: %w", err)
}
for _, cmd := range cmds {
waitCount := 0
if v, err := cmd.getCmd.Int(); err == nil && v > 0 {
waitCount = v
}
loads = append(loads, activeIndexLoad{
id: cmd.id,
member: cmd.member,
slotCount: int(cmd.zcardCmd.Val()),
waitCount: waitCount,
})
}
}
return loads, staleMembers, nil
}
// removeActiveIndexMembers 清理无效 member;这是辅助索引的维护动作,调用方无需因为失败中断主流程。
func (c *concurrencyCache) removeActiveIndexMembers(ctx context.Context, indexKey string, members []string) {
if len(members) == 0 {
return
}
args := make([]any, 0, len(members))
for _, member := range members {
args = append(args, member)
}
if err := c.rdb.ZRem(ctx, indexKey, args...).Err(); err != nil {
logger.LegacyPrintf("repository.concurrency", "Warning: remove %d active index members from %s failed: %v", len(members), indexKey, err)
}
}
// runScriptInt64Pair 执行返回两元素整数数组的 Lua 脚本并解析(如 {result, now}、{removed, remaining})。
func runScriptInt64Pair(ctx context.Context, rdb *redis.Client, script *redis.Script, keys []string, args ...any) (int64, int64, error) {
raw, err := script.Run(ctx, rdb, keys, args...).Result()
if err != nil {
return 0, 0, err
}
first, err := redisScriptInt64At(raw, 0)
if err != nil {
return 0, 0, fmt.Errorf("parse script value 0: %w", err)
}
second, err := redisScriptInt64At(raw, 1)
if err != nil {
return 0, 0, fmt.Errorf("parse script value 1: %w", err)
}
return first, second, nil
}
// Account slot operations
func (c *concurrencyCache) AcquireAccountSlot(ctx context.Context, accountID int64, maxConcurrency int, requestID string) (bool, error) {
key := accountSlotKey(accountID)
// 时间戳在 Lua 脚本内使用 Redis TIME 命令获取,确保多实例时钟一致
result, err := acquireScript.Run(ctx, c.rdb, []string{key}, maxConcurrency, c.slotTTLSeconds, requestID).Int()
result, now, err := runScriptInt64Pair(ctx, c.rdb, acquireScript, []string{key}, maxConcurrency, c.slotTTLSeconds, requestID)
if err != nil {
return false, err
}
if result == 1 {
// 成功占槽后标记活跃账号,后台清理即可从索引定位候选账号。
c.touchActiveIndexAt(ctx, accountActiveIndexKey, accountID, now+int64(c.slotTTLSeconds))
}
return result == 1, nil
}
func (c *concurrencyCache) ReleaseAccountSlot(ctx context.Context, accountID int64, requestID string) error {
key := accountSlotKey(accountID)
return c.rdb.ZRem(ctx, key, requestID).Err()
if err := c.rdb.ZRem(ctx, key, requestID).Err(); err != nil {
return err
}
// 释放后用真实负载刷新索引;若没有槽位和等待计数,会移除索引 member。
c.refreshAccountActiveIndex(ctx, accountID)
return nil
}
func (c *concurrencyCache) GetAccountConcurrency(ctx context.Context, accountID int64) (int, error) {
@@ -359,16 +581,25 @@ func (c *concurrencyCache) GetAccountConcurrencyBatch(ctx context.Context, accou
func (c *concurrencyCache) AcquireUserSlot(ctx context.Context, userID int64, maxConcurrency int, requestID string) (bool, error) {
key := userSlotKey(userID)
// 时间戳在 Lua 脚本内使用 Redis TIME 命令获取,确保多实例时钟一致
result, err := acquireScript.Run(ctx, c.rdb, []string{key}, maxConcurrency, c.slotTTLSeconds, requestID).Int()
result, now, err := runScriptInt64Pair(ctx, c.rdb, acquireScript, []string{key}, maxConcurrency, c.slotTTLSeconds, requestID)
if err != nil {
return false, err
}
if result == 1 {
// 成功占槽后标记活跃用户,避免启动清理依赖全量 SCAN。
c.touchActiveIndexAt(ctx, userActiveIndexKey, userID, now+int64(c.slotTTLSeconds))
}
return result == 1, nil
}
func (c *concurrencyCache) ReleaseUserSlot(ctx context.Context, userID int64, requestID string) error {
key := userSlotKey(userID)
return c.rdb.ZRem(ctx, key, requestID).Err()
if err := c.rdb.ZRem(ctx, key, requestID).Err(); err != nil {
return err
}
// 释放后按 Redis 中剩余负载修正索引状态。
c.refreshUserActiveIndex(ctx, userID)
return nil
}
func (c *concurrencyCache) GetUserConcurrency(ctx context.Context, userID int64) (int, error) {
@@ -433,16 +664,24 @@ func (c *concurrencyCache) GetAPIKeyConcurrencyBatch(ctx context.Context, apiKey
func (c *concurrencyCache) IncrementWaitCount(ctx context.Context, userID int64, maxWait int) (bool, error) {
key := waitQueueKey(userID)
result, err := incrementWaitScript.Run(ctx, c.rdb, []string{key}, maxWait, c.waitQueueTTLSeconds).Int()
result, now, err := runScriptInt64Pair(ctx, c.rdb, incrementWaitScript, []string{key}, maxWait, c.waitQueueTTLSeconds)
if err != nil {
return false, err
}
if result == 1 {
// 等待队列也会让用户保持“活跃”,否则槽位为 0 时后台任务可能漏看等待计数。
c.touchActiveIndexAt(ctx, userActiveIndexKey, userID, now+int64(c.waitQueueTTLSeconds))
}
return result == 1, nil
}
func (c *concurrencyCache) DecrementWaitCount(ctx context.Context, userID int64) error {
key := waitQueueKey(userID)
_, err := decrementWaitScript.Run(ctx, c.rdb, []string{key}).Result()
if err == nil {
// 等待数减少后重新判断是否还需要保留索引。
c.refreshUserActiveIndex(ctx, userID)
}
return err
}
@@ -450,16 +689,24 @@ func (c *concurrencyCache) DecrementWaitCount(ctx context.Context, userID int64)
func (c *concurrencyCache) IncrementAccountWaitCount(ctx context.Context, accountID int64, maxWait int) (bool, error) {
key := accountWaitKey(accountID)
result, err := incrementAccountWaitScript.Run(ctx, c.rdb, []string{key}, maxWait, c.waitQueueTTLSeconds).Int()
result, now, err := runScriptInt64Pair(ctx, c.rdb, incrementAccountWaitScript, []string{key}, maxWait, c.waitQueueTTLSeconds)
if err != nil {
return false, err
}
if result == 1 {
// 账号级等待队列同样写入账号活跃索引,供负载查询和清理任务使用。
c.touchActiveIndexAt(ctx, accountActiveIndexKey, accountID, now+int64(c.waitQueueTTLSeconds))
}
return result == 1, nil
}
func (c *concurrencyCache) DecrementAccountWaitCount(ctx context.Context, accountID int64) error {
key := accountWaitKey(accountID)
_, err := decrementWaitScript.Run(ctx, c.rdb, []string{key}).Result()
if err == nil {
// 等待计数归零后索引需要同步删除,避免后台任务反复处理空账号。
c.refreshAccountActiveIndex(ctx, accountID)
}
return err
}
@@ -599,101 +846,183 @@ func (c *concurrencyCache) GetUsersLoadBatch(ctx context.Context, users []servic
func (c *concurrencyCache) CleanupExpiredAccountSlots(ctx context.Context, accountID int64) error {
key := accountSlotKey(accountID)
_, err := cleanupExpiredSlotsScript.Run(ctx, c.rdb, []string{key}, c.slotTTLSeconds).Result()
if err == nil {
// 单账号清理后同步索引,保持后台批量清理的候选集准确。
c.refreshAccountActiveIndex(ctx, accountID)
}
return err
}
// CleanupExpiredAccountSlotKeys 处理账号与用户两个活跃索引中已到期的候选。
// (方法名中的 Account 是历史遗留,保留以避免接口变更;实际同时回收两个索引,
// 否则 user 索引的过期成员没有任何清理路径,会无界累积。)
func (c *concurrencyCache) CleanupExpiredAccountSlotKeys(ctx context.Context) error {
return c.cleanupExpiredSlotKeysByPattern(ctx, accountSlotKeyPrefix+"*")
if err := c.reconcileExpiredIndexCandidates(ctx, accountSlotIndex); err != nil {
return err
}
return c.reconcileExpiredIndexCandidates(ctx, userSlotIndex)
}
// reconcileExpiredIndexCandidates 处理单个活跃索引中 score 已到期的候选:
// 无真实负载则移除 member;仍有负载则按真实负载批量刷新 score。
func (c *concurrencyCache) reconcileExpiredIndexCandidates(ctx context.Context, spec slotIndexSpec) error {
now, err := c.redisUnixSeconds(ctx)
if err != nil {
return err
}
members, err := c.rdb.ZRangeByScore(ctx, spec.indexKey, &redis.ZRangeBy{
Min: "-inf",
Max: strconv.FormatInt(now, 10),
Count: activeIndexCleanupBatchSize,
}).Result()
if err != nil {
return fmt.Errorf("read expired index %s: %w", spec.indexKey, err)
}
loads, staleMembers, err := c.readIndexLoads(ctx, spec, members, now)
if err != nil {
return err
}
refreshed := make([]redis.Z, 0, len(loads))
for _, load := range loads {
if load.slotCount == 0 && load.waitCount <= 0 {
// 真实槽位和等待数都为空,说明这个索引 member 已经完成使命。
staleMembers = append(staleMembers, load.member)
continue
}
refreshed = append(refreshed, redis.Z{
Score: float64(now + int64(c.activeIndexTTL(load.slotCount, load.waitCount))),
Member: load.member,
})
}
if len(refreshed) > 0 {
if err := c.rdb.ZAdd(ctx, spec.indexKey, refreshed...).Err(); err != nil {
logger.LegacyPrintf("repository.concurrency", "Warning: refresh %d active index members in %s failed: %v", len(refreshed), spec.indexKey, err)
}
}
c.removeActiveIndexMembers(ctx, spec.indexKey, staleMembers)
return nil
}
// CleanupStaleProcessSlots 启动时清理非当前进程前缀的槽位。
// 清理范围来自活跃索引(含 score 已过期的成员——它们往往正是崩溃进程留下的残留),
// 避免在 Redis 上 SCAN 全部 concurrency:* 键;另有一次性迁移清扫兜底索引机制上线前的遗留等待计数。
// API Key 槽位(concurrency:api_key:*)是 stats-only 数据:每次 Track/读取都会按分数
// 裁剪过期成员,key 自带 TTL,可在一个 slot TTL 内自愈,因此不参与启动清理。
func (c *concurrencyCache) CleanupStaleProcessSlots(ctx context.Context, activeRequestPrefix string) error {
if activeRequestPrefix == "" {
return nil
}
// 1. 清理有序集合中非当前进程前缀的成员
slotPatterns := []string{accountSlotKeyPrefix + "*", userSlotKeyPrefix + "*", apiKeySlotKeyPrefix + "*"}
for _, pattern := range slotPatterns {
if err := c.cleanupSlotsByPattern(ctx, pattern, activeRequestPrefix); err != nil {
return err
}
if err := c.sweepLegacyWaitKeysOnce(ctx); err != nil {
return err
}
now, err := c.redisUnixSeconds(ctx)
if err != nil {
return err
}
// 2. 删除所有等待队列计数器(重启后计数器失效)
waitPatterns := []string{accountWaitKeyPrefix + "*", waitQueueKeyPrefix + "*"}
for _, pattern := range waitPatterns {
if err := c.deleteKeysByPattern(ctx, pattern); err != nil {
return err
}
accountMembers, err := c.allIndexMembers(ctx, accountActiveIndexKey)
if err != nil {
return err
}
if err := c.cleanupStaleProcessSlotsForIndex(ctx, accountSlotIndex, accountMembers, activeRequestPrefix, now); err != nil {
return err
}
return nil
userMembers, err := c.allIndexMembers(ctx, userActiveIndexKey)
if err != nil {
return err
}
return c.cleanupStaleProcessSlotsForIndex(ctx, userSlotIndex, userMembers, activeRequestPrefix, now)
}
// cleanupExpiredSlotKeysByPattern 扫描实际存在的账号槽位键并批量清理过期成员。
func (c *concurrencyCache) cleanupExpiredSlotKeysByPattern(ctx context.Context, pattern string) error {
const scanCount = 200
var cursor uint64
for {
keys, nextCursor, err := c.rdb.Scan(ctx, cursor, pattern, scanCount).Result()
if err != nil {
return fmt.Errorf("scan %s: %w", pattern, err)
}
if len(keys) > 0 {
_, err := cleanupExpiredSlotKeysScript.Run(ctx, c.rdb, keys, c.slotTTLSeconds).Result()
// sweepLegacyWaitKeysOnce 一次性清扫活跃索引机制上线前遗留的等待计数键。
// 等待计数在有流量时会不断刷新 TTL、无法自然过期,而索引不认识旧键,
// 因此这里例外地做一次 SCAN,用 marker 键保证整个 Redis 数据生命周期内只执行一次。
// 先清扫后写 marker:清扫失败时下次启动会重试;并发实例重复清扫是幂等的。
func (c *concurrencyCache) sweepLegacyWaitKeysOnce(ctx context.Context) error {
exists, err := c.rdb.Exists(ctx, legacyWaitSweepMarkerKey).Result()
if err != nil {
return fmt.Errorf("check legacy wait sweep marker: %w", err)
}
if exists > 0 {
return nil
}
for _, pattern := range []string{accountWaitKeyPrefix + "*", waitQueueKeyPrefix + "*"} {
var cursor uint64
for {
keys, next, err := c.rdb.Scan(ctx, cursor, pattern, 200).Result()
if err != nil {
return fmt.Errorf("cleanup expired slots %s: %w", pattern, err)
return fmt.Errorf("scan legacy wait keys %s: %w", pattern, err)
}
if len(keys) > 0 {
if err := c.rdb.Del(ctx, keys...).Err(); err != nil {
return fmt.Errorf("delete legacy wait keys: %w", err)
}
}
cursor = next
if cursor == 0 {
break
}
}
cursor = nextCursor
if cursor == 0 {
break
}
}
if err := c.rdb.Set(ctx, legacyWaitSweepMarkerKey, "1", 0).Err(); err != nil {
return fmt.Errorf("set legacy wait sweep marker: %w", err)
}
return nil
}
// cleanupSlotsByPattern 扫描匹配 pattern 的有序集合键,批量调用 Lua 脚本清理非当前进程成员。
func (c *concurrencyCache) cleanupSlotsByPattern(ctx context.Context, pattern, activePrefix string) error {
const scanCount = 200
var cursor uint64
for {
keys, nextCursor, err := c.rdb.Scan(ctx, cursor, pattern, scanCount).Result()
if err != nil {
return fmt.Errorf("scan %s: %w", pattern, err)
}
if len(keys) > 0 {
_, err := startupCleanupScript.Run(ctx, c.rdb, keys, activePrefix, c.slotTTLSeconds).Result()
if err != nil {
return fmt.Errorf("cleanup slots %s: %w", pattern, err)
}
}
cursor = nextCursor
if cursor == 0 {
break
}
// allIndexMembers 返回索引中全部 member(含 score 已过期的)。
// 启动清理必须覆盖过期成员:长时间停机后 score 过期的候选恰恰最可能持有死进程残留。
func (c *concurrencyCache) allIndexMembers(ctx context.Context, indexKey string) ([]string, error) {
members, err := c.rdb.ZRange(ctx, indexKey, 0, -1).Result()
if err != nil {
return nil, fmt.Errorf("read active index %s: %w", indexKey, err)
}
return nil
return members, nil
}
// deleteKeysByPattern 扫描匹配 pattern 的键并删除。
func (c *concurrencyCache) deleteKeysByPattern(ctx context.Context, pattern string) error {
const scanCount = 200
var cursor uint64
for {
keys, nextCursor, err := c.rdb.Scan(ctx, cursor, pattern, scanCount).Result()
// cleanupStaleProcessSlotsForIndex 逐个处理索引中的账号/用户。
// Lua 脚本一次只碰一个槽位 key,兼容 Redis Cluster,随后删除重启后已失效的等待计数;
// 索引 member 的去留由脚本返回的剩余槽位数决定,最后批量写回。
func (c *concurrencyCache) cleanupStaleProcessSlotsForIndex(
ctx context.Context,
spec slotIndexSpec,
members []string,
activeRequestPrefix string,
now int64,
) error {
staleMembers := make([]string, 0)
refreshed := make([]redis.Z, 0)
for _, member := range members {
id, err := strconv.ParseInt(member, 10, 64)
if err != nil || id <= 0 {
staleMembers = append(staleMembers, member)
continue
}
_, remaining, err := runScriptInt64Pair(ctx, c.rdb, startupCleanupSlotScript, []string{spec.slotKey(id)}, activeRequestPrefix, c.slotTTLSeconds)
if err != nil {
return fmt.Errorf("scan %s: %w", pattern, err)
return fmt.Errorf("cleanup stale process slots %s: %w", spec.slotKey(id), err)
}
if len(keys) > 0 {
if err := c.rdb.Del(ctx, keys...).Err(); err != nil {
return fmt.Errorf("del %s: %w", pattern, err)
}
// 等待计数属于已死进程,直接删除;剩余槽位(当前进程前缀)决定索引 member 去留。
if err := c.rdb.Del(ctx, spec.waitKey(id)).Err(); err != nil {
return fmt.Errorf("delete stale wait key %s: %w", spec.waitKey(id), err)
}
cursor = nextCursor
if cursor == 0 {
break
if remaining > 0 {
refreshed = append(refreshed, redis.Z{
Score: float64(now + int64(c.slotTTLSeconds)),
Member: member,
})
} else {
staleMembers = append(staleMembers, member)
}
}
if len(refreshed) > 0 {
if err := c.rdb.ZAdd(ctx, spec.indexKey, refreshed...).Err(); err != nil {
logger.LegacyPrintf("repository.concurrency", "Warning: refresh %d active index members in %s failed: %v", len(refreshed), spec.indexKey, err)
}
}
c.removeActiveIndexMembers(ctx, spec.indexKey, staleMembers)
return nil
}
@@ -6,6 +6,7 @@ import (
"context"
"errors"
"fmt"
"strconv"
"testing"
"time"
@@ -23,7 +24,8 @@ var testSlotTTL = time.Duration(testSlotTTLMinutes) * time.Minute
type ConcurrencyCacheSuite struct {
IntegrationRedisSuite
cache service.ConcurrencyCache
cache service.ConcurrencyCache
rawCache *concurrencyCache
}
func TestConcurrencyCacheSuite(t *testing.T) {
@@ -32,7 +34,8 @@ func TestConcurrencyCacheSuite(t *testing.T) {
func (s *ConcurrencyCacheSuite) SetupTest() {
s.IntegrationRedisSuite.SetupTest()
s.cache = NewConcurrencyCache(s.rdb, testSlotTTLMinutes, int(testSlotTTL.Seconds()))
s.rawCache = NewConcurrencyCache(s.rdb, testSlotTTLMinutes, int(testSlotTTL.Seconds())).(*concurrencyCache)
s.cache = s.rawCache
}
type apiKeyConcurrencyCacheForTest interface {
@@ -74,6 +77,63 @@ func (s *ConcurrencyCacheSuite) TestAccountSlot_AcquireAndRelease() {
require.Equal(s.T(), 1, cur, "expected 1 after release")
}
func (s *ConcurrencyCacheSuite) TestAccountActiveIndex_AcquireAndRelease() {
accountID := int64(610)
member := strconv.FormatInt(accountID, 10)
reqID := "active-index-req"
now, err := s.rawCache.redisUnixSeconds(s.ctx)
require.NoError(s.T(), err)
ok, err := s.cache.AcquireAccountSlot(s.ctx, accountID, 2, reqID)
require.NoError(s.T(), err)
require.True(s.T(), ok)
score, err := s.rdb.ZScore(s.ctx, accountActiveIndexKey, member).Result()
require.NoError(s.T(), err)
require.Greater(s.T(), int64(score), now, "index score should be a future expiry")
require.NoError(s.T(), s.cache.ReleaseAccountSlot(s.ctx, accountID, reqID))
_, err = s.rdb.ZScore(s.ctx, accountActiveIndexKey, member).Result()
require.ErrorIs(s.T(), err, redis.Nil, "index member should be removed after load drops to zero")
}
func (s *ConcurrencyCacheSuite) TestAccountActiveIndex_WaitLifecycle() {
accountID := int64(611)
member := strconv.FormatInt(accountID, 10)
ok, err := s.cache.IncrementAccountWaitCount(s.ctx, accountID, 2)
require.NoError(s.T(), err)
require.True(s.T(), ok)
_, err = s.rdb.ZScore(s.ctx, accountActiveIndexKey, member).Result()
require.NoError(s.T(), err, "wait increment should register index member")
require.NoError(s.T(), s.cache.DecrementAccountWaitCount(s.ctx, accountID))
_, err = s.rdb.ZScore(s.ctx, accountActiveIndexKey, member).Result()
require.ErrorIs(s.T(), err, redis.Nil, "index member should be removed after wait drops to zero")
}
func (s *ConcurrencyCacheSuite) TestUserActiveIndex_AcquireAndRelease() {
userID := int64(612)
member := strconv.FormatInt(userID, 10)
reqID := "user-active-index-req"
ok, err := s.cache.AcquireUserSlot(s.ctx, userID, 2, reqID)
require.NoError(s.T(), err)
require.True(s.T(), ok)
_, err = s.rdb.ZScore(s.ctx, userActiveIndexKey, member).Result()
require.NoError(s.T(), err, "acquire should register user index member")
require.NoError(s.T(), s.cache.ReleaseUserSlot(s.ctx, userID, reqID))
_, err = s.rdb.ZScore(s.ctx, userActiveIndexKey, member).Result()
require.ErrorIs(s.T(), err, redis.Nil, "user index member should be removed after release")
}
func (s *ConcurrencyCacheSuite) TestAccountSlot_TTL() {
accountID := int64(11)
reqID := "req_ttl_test"
@@ -293,16 +353,22 @@ func (s *ConcurrencyCacheSuite) TestAccountWaitQueue_IncrementAndDecrement() {
}
func (s *ConcurrencyCacheSuite) TestCleanupStaleProcessSlots() {
// 预置迁移 marker,隔离一次性清扫,只验证索引驱动的清理路径。
require.NoError(s.T(), s.rdb.Set(s.ctx, legacyWaitSweepMarkerKey, "1", 0).Err())
accountID := int64(901)
userID := int64(902)
apiKeyID := int64(903)
unindexedAccountID := int64(1901)
accountKey := fmt.Sprintf("%s%d", accountSlotKeyPrefix, accountID)
userKey := fmt.Sprintf("%s%d", userSlotKeyPrefix, userID)
apiKeyKey := fmt.Sprintf("%s%d", apiKeySlotKeyPrefix, apiKeyID)
unindexedAccountKey := fmt.Sprintf("%s%d", accountSlotKeyPrefix, unindexedAccountID)
userWaitKey := fmt.Sprintf("%s%d", waitQueueKeyPrefix, userID)
accountWaitKey := fmt.Sprintf("%s%d", accountWaitKeyPrefix, accountID)
unindexedAccountWaitKey := fmt.Sprintf("%s%d", accountWaitKeyPrefix, unindexedAccountID)
now := time.Now().Unix()
now, err := s.rawCache.redisUnixSeconds(s.ctx)
require.NoError(s.T(), err)
require.NoError(s.T(), s.rdb.ZAdd(s.ctx, accountKey,
redis.Z{Score: float64(now), Member: "oldproc-1"},
redis.Z{Score: float64(now), Member: "keep-1"},
@@ -311,12 +377,24 @@ func (s *ConcurrencyCacheSuite) TestCleanupStaleProcessSlots() {
redis.Z{Score: float64(now), Member: "oldproc-2"},
redis.Z{Score: float64(now), Member: "keep-2"},
).Err())
require.NoError(s.T(), s.rdb.ZAdd(s.ctx, unindexedAccountKey,
redis.Z{Score: float64(now), Member: "oldproc-unindexed"},
).Err())
require.NoError(s.T(), s.rdb.ZAdd(s.ctx, apiKeyKey,
redis.Z{Score: float64(now), Member: "oldproc-3"},
redis.Z{Score: float64(now), Member: "keep-3"},
).Err())
require.NoError(s.T(), s.rdb.Set(s.ctx, userWaitKey, 3, time.Minute).Err())
require.NoError(s.T(), s.rdb.Set(s.ctx, accountWaitKey, 2, time.Minute).Err())
require.NoError(s.T(), s.rdb.Set(s.ctx, unindexedAccountWaitKey, 2, time.Minute).Err())
require.NoError(s.T(), s.rdb.ZAdd(s.ctx, accountActiveIndexKey, redis.Z{
Score: float64(now + 60),
Member: strconv.FormatInt(accountID, 10),
}).Err())
require.NoError(s.T(), s.rdb.ZAdd(s.ctx, userActiveIndexKey, redis.Z{
Score: float64(now + 60),
Member: strconv.FormatInt(userID, 10),
}).Err())
require.NoError(s.T(), s.cache.CleanupStaleProcessSlots(s.ctx, "keep-"))
@@ -328,15 +406,22 @@ func (s *ConcurrencyCacheSuite) TestCleanupStaleProcessSlots() {
require.NoError(s.T(), err)
require.Equal(s.T(), []string{"keep-2"}, userMembers)
// API Key 槽位(stats-only)不在启动清理范围内,靠分数裁剪与 key TTL 自愈。
apiKeyMembers, err := s.rdb.ZRange(s.ctx, apiKeyKey, 0, -1).Result()
require.NoError(s.T(), err)
require.Equal(s.T(), []string{"keep-3"}, apiKeyMembers)
require.ElementsMatch(s.T(), []string{"keep-3", "oldproc-3"}, apiKeyMembers)
_, err = s.rdb.Get(s.ctx, userWaitKey).Result()
require.True(s.T(), errors.Is(err, redis.Nil))
_, err = s.rdb.Get(s.ctx, accountWaitKey).Result()
require.True(s.T(), errors.Is(err, redis.Nil))
unindexedMembers, err := s.rdb.ZRange(s.ctx, unindexedAccountKey, 0, -1).Result()
require.NoError(s.T(), err)
require.Equal(s.T(), []string{"oldproc-unindexed"}, unindexedMembers)
_, err = s.rdb.Get(s.ctx, unindexedAccountWaitKey).Result()
require.NoError(s.T(), err)
}
func (s *ConcurrencyCacheSuite) TestGetAccountConcurrency_Missing() {
@@ -487,11 +572,13 @@ func (s *ConcurrencyCacheSuite) TestCleanupExpiredAccountSlots_NoExpired() {
}
func (s *ConcurrencyCacheSuite) TestCleanupExpiredAccountSlotKeys() {
now := time.Now().Unix()
now, err := s.rawCache.redisUnixSeconds(s.ctx)
require.NoError(s.T(), err)
expiredTime := now - int64(testSlotTTL.Seconds()) - 10
accountKeyWithFresh := fmt.Sprintf("%s%d", accountSlotKeyPrefix, 301)
accountKeyExpiredOnly := fmt.Sprintf("%s%d", accountSlotKeyPrefix, 302)
userKey := fmt.Sprintf("%s%d", userSlotKeyPrefix, 303)
unindexedAccountKey := fmt.Sprintf("%s%d", accountSlotKeyPrefix, 304)
require.NoError(s.T(), s.rdb.ZAdd(s.ctx, accountKeyWithFresh,
redis.Z{Score: float64(expiredTime), Member: "expired"},
@@ -503,6 +590,13 @@ func (s *ConcurrencyCacheSuite) TestCleanupExpiredAccountSlotKeys() {
require.NoError(s.T(), s.rdb.ZAdd(s.ctx, userKey,
redis.Z{Score: float64(expiredTime), Member: "user-expired"},
).Err())
require.NoError(s.T(), s.rdb.ZAdd(s.ctx, unindexedAccountKey,
redis.Z{Score: float64(expiredTime), Member: "unindexed-expired"},
).Err())
require.NoError(s.T(), s.rdb.ZAdd(s.ctx, accountActiveIndexKey,
redis.Z{Score: float64(now), Member: "301"},
redis.Z{Score: float64(now), Member: "302"},
).Err())
require.NoError(s.T(), s.cache.CleanupExpiredAccountSlotKeys(s.ctx))
@@ -517,9 +611,125 @@ func (s *ConcurrencyCacheSuite) TestCleanupExpiredAccountSlotKeys() {
userMembers, err := s.rdb.ZRange(s.ctx, userKey, 0, -1).Result()
require.NoError(s.T(), err)
require.Equal(s.T(), []string{"user-expired"}, userMembers)
unindexedMembers, err := s.rdb.ZRange(s.ctx, unindexedAccountKey, 0, -1).Result()
require.NoError(s.T(), err)
require.Equal(s.T(), []string{"unindexed-expired"}, unindexedMembers)
score, err := s.rdb.ZScore(s.ctx, accountActiveIndexKey, "301").Result()
require.NoError(s.T(), err)
require.Greater(s.T(), int64(score), now)
_, err = s.rdb.ZScore(s.ctx, accountActiveIndexKey, "302").Result()
require.ErrorIs(s.T(), err, redis.Nil)
}
func (s *ConcurrencyCacheSuite) TestCleanupExpiredAccountSlotKeys_ReapsUserIndex() {
now, err := s.rawCache.redisUnixSeconds(s.ctx)
require.NoError(s.T(), err)
expiredScore := float64(now - 10)
userKeyWithFresh := fmt.Sprintf("%s%d", userSlotKeyPrefix, 401)
// 401 有真实负载但索引 score 已过期:应刷新而不是删除。
require.NoError(s.T(), s.rdb.ZAdd(s.ctx, userKeyWithFresh,
redis.Z{Score: float64(now), Member: "fresh"},
).Err())
// 402 无任何负载:过期索引 member 应被回收。
// 非法 member 也应随过期候选一并清除。
require.NoError(s.T(), s.rdb.ZAdd(s.ctx, userActiveIndexKey,
redis.Z{Score: expiredScore, Member: "401"},
redis.Z{Score: expiredScore, Member: "402"},
redis.Z{Score: expiredScore, Member: "not-a-user-id"},
).Err())
require.NoError(s.T(), s.cache.CleanupExpiredAccountSlotKeys(s.ctx))
score, err := s.rdb.ZScore(s.ctx, userActiveIndexKey, "401").Result()
require.NoError(s.T(), err)
require.Greater(s.T(), int64(score), now, "loaded user should be re-scheduled, not dropped")
_, err = s.rdb.ZScore(s.ctx, userActiveIndexKey, "402").Result()
require.ErrorIs(s.T(), err, redis.Nil, "idle expired user member should be reaped")
_, err = s.rdb.ZScore(s.ctx, userActiveIndexKey, "not-a-user-id").Result()
require.ErrorIs(s.T(), err, redis.Nil, "invalid member should be reaped")
}
func (s *ConcurrencyCacheSuite) TestCleanupStaleProcessSlots_LegacyWaitSweepRunsOnce() {
unindexedAccountWaitKey := fmt.Sprintf("%s%d", accountWaitKeyPrefix, 2901)
unindexedUserWaitKey := fmt.Sprintf("%s%d", waitQueueKeyPrefix, 2902)
require.NoError(s.T(), s.rdb.Set(s.ctx, unindexedAccountWaitKey, 5, time.Minute).Err())
require.NoError(s.T(), s.rdb.Set(s.ctx, unindexedUserWaitKey, 3, time.Minute).Err())
// 首次运行:marker 不存在,一次性清扫删除所有遗留等待计数(含未入索引的)。
require.NoError(s.T(), s.cache.CleanupStaleProcessSlots(s.ctx, "keep-"))
_, err := s.rdb.Get(s.ctx, unindexedAccountWaitKey).Result()
require.ErrorIs(s.T(), err, redis.Nil, "legacy account wait key should be swept on first startup")
_, err = s.rdb.Get(s.ctx, unindexedUserWaitKey).Result()
require.ErrorIs(s.T(), err, redis.Nil, "legacy user wait key should be swept on first startup")
exists, err := s.rdb.Exists(s.ctx, legacyWaitSweepMarkerKey).Result()
require.NoError(s.T(), err)
require.EqualValues(s.T(), 1, exists, "sweep marker should be set after first run")
// 再次运行:marker 已存在,未入索引的等待计数不再被触碰。
require.NoError(s.T(), s.rdb.Set(s.ctx, unindexedAccountWaitKey, 5, time.Minute).Err())
require.NoError(s.T(), s.cache.CleanupStaleProcessSlots(s.ctx, "keep-"))
val, err := s.rdb.Get(s.ctx, unindexedAccountWaitKey).Int()
require.NoError(s.T(), err, "sweep must not run twice")
require.Equal(s.T(), 5, val)
}
func (s *ConcurrencyCacheSuite) TestCleanupStaleProcessSlots_ProcessesExpiredIndexMembers() {
// score 已过期的索引成员往往正是崩溃进程留下的残留,启动清理必须覆盖它们。
require.NoError(s.T(), s.rdb.Set(s.ctx, legacyWaitSweepMarkerKey, "1", 0).Err())
accountID := int64(3901)
userID := int64(3902)
accountKey := fmt.Sprintf("%s%d", accountSlotKeyPrefix, accountID)
userKey := fmt.Sprintf("%s%d", userSlotKeyPrefix, userID)
accountWaitKey := fmt.Sprintf("%s%d", accountWaitKeyPrefix, accountID)
now, err := s.rawCache.redisUnixSeconds(s.ctx)
require.NoError(s.T(), err)
require.NoError(s.T(), s.rdb.ZAdd(s.ctx, accountKey,
redis.Z{Score: float64(now), Member: "oldproc-1"},
).Err())
require.NoError(s.T(), s.rdb.ZAdd(s.ctx, userKey,
redis.Z{Score: float64(now), Member: "oldproc-2"},
).Err())
require.NoError(s.T(), s.rdb.Set(s.ctx, accountWaitKey, 4, time.Minute).Err())
// 索引 score 设为过去时刻,模拟长时间停机后索引已“过期”。
require.NoError(s.T(), s.rdb.ZAdd(s.ctx, accountActiveIndexKey, redis.Z{
Score: float64(now - 100),
Member: strconv.FormatInt(accountID, 10),
}).Err())
require.NoError(s.T(), s.rdb.ZAdd(s.ctx, userActiveIndexKey, redis.Z{
Score: float64(now - 100),
Member: strconv.FormatInt(userID, 10),
}).Err())
require.NoError(s.T(), s.cache.CleanupStaleProcessSlots(s.ctx, "keep-"))
exists, err := s.rdb.Exists(s.ctx, accountKey).Result()
require.NoError(s.T(), err)
require.EqualValues(s.T(), 0, exists, "stale slot key of expired index member should be purged")
exists, err = s.rdb.Exists(s.ctx, userKey).Result()
require.NoError(s.T(), err)
require.EqualValues(s.T(), 0, exists)
_, err = s.rdb.Get(s.ctx, accountWaitKey).Result()
require.ErrorIs(s.T(), err, redis.Nil, "wait counter of expired index member should be deleted")
_, err = s.rdb.ZScore(s.ctx, accountActiveIndexKey, strconv.FormatInt(accountID, 10)).Result()
require.ErrorIs(s.T(), err, redis.Nil, "emptied member should be removed from index")
_, err = s.rdb.ZScore(s.ctx, userActiveIndexKey, strconv.FormatInt(userID, 10)).Result()
require.ErrorIs(s.T(), err, redis.Nil)
}
func (s *ConcurrencyCacheSuite) TestCleanupStaleProcessSlots_RemovesOldPrefixesAndWaitCounters() {
// 预置迁移 marker,确保等待计数删除来自索引驱动路径而非一次性清扫。
require.NoError(s.T(), s.rdb.Set(s.ctx, legacyWaitSweepMarkerKey, "1", 0).Err())
accountID := int64(901)
userID := int64(902)
accountSlotKey := fmt.Sprintf("%s%d", accountSlotKeyPrefix, accountID)
@@ -527,19 +737,28 @@ func (s *ConcurrencyCacheSuite) TestCleanupStaleProcessSlots_RemovesOldPrefixesA
userWaitKey := fmt.Sprintf("%s%d", waitQueueKeyPrefix, userID)
accountWaitKey := fmt.Sprintf("%s%d", accountWaitKeyPrefix, accountID)
now := float64(time.Now().Unix())
now, err := s.rawCache.redisUnixSeconds(s.ctx)
require.NoError(s.T(), err)
require.NoError(s.T(), s.rdb.ZAdd(s.ctx, accountSlotKey,
redis.Z{Score: now, Member: "oldproc-1"},
redis.Z{Score: now, Member: "activeproc-1"},
redis.Z{Score: float64(now), Member: "oldproc-1"},
redis.Z{Score: float64(now), Member: "activeproc-1"},
).Err())
require.NoError(s.T(), s.rdb.Expire(s.ctx, accountSlotKey, testSlotTTL).Err())
require.NoError(s.T(), s.rdb.ZAdd(s.ctx, userSlotKey,
redis.Z{Score: now, Member: "oldproc-2"},
redis.Z{Score: now, Member: "activeproc-2"},
redis.Z{Score: float64(now), Member: "oldproc-2"},
redis.Z{Score: float64(now), Member: "activeproc-2"},
).Err())
require.NoError(s.T(), s.rdb.Expire(s.ctx, userSlotKey, testSlotTTL).Err())
require.NoError(s.T(), s.rdb.Set(s.ctx, userWaitKey, 3, testSlotTTL).Err())
require.NoError(s.T(), s.rdb.Set(s.ctx, accountWaitKey, 2, testSlotTTL).Err())
require.NoError(s.T(), s.rdb.ZAdd(s.ctx, accountActiveIndexKey, redis.Z{
Score: float64(now + 60),
Member: strconv.FormatInt(accountID, 10),
}).Err())
require.NoError(s.T(), s.rdb.ZAdd(s.ctx, userActiveIndexKey, redis.Z{
Score: float64(now + 60),
Member: strconv.FormatInt(userID, 10),
}).Err())
require.NoError(s.T(), s.cache.CleanupStaleProcessSlots(s.ctx, "activeproc-"))
@@ -560,8 +779,14 @@ func (s *ConcurrencyCacheSuite) TestCleanupStaleProcessSlots_RemovesOldPrefixesA
func (s *ConcurrencyCacheSuite) TestCleanupStaleProcessSlots_DeletesEmptySlotKeys() {
accountID := int64(903)
accountSlotKey := fmt.Sprintf("%s%d", accountSlotKeyPrefix, accountID)
require.NoError(s.T(), s.rdb.ZAdd(s.ctx, accountSlotKey, redis.Z{Score: float64(time.Now().Unix()), Member: "oldproc-1"}).Err())
now, err := s.rawCache.redisUnixSeconds(s.ctx)
require.NoError(s.T(), err)
require.NoError(s.T(), s.rdb.ZAdd(s.ctx, accountSlotKey, redis.Z{Score: float64(now), Member: "oldproc-1"}).Err())
require.NoError(s.T(), s.rdb.Expire(s.ctx, accountSlotKey, testSlotTTL).Err())
require.NoError(s.T(), s.rdb.ZAdd(s.ctx, accountActiveIndexKey, redis.Z{
Score: float64(now + 60),
Member: strconv.FormatInt(accountID, 10),
}).Err())
require.NoError(s.T(), s.cache.CleanupStaleProcessSlots(s.ctx, "activeproc-"))
@@ -5,9 +5,8 @@ import (
"errors"
"fmt"
"strconv"
"strings"
"time"
"github.com/Wei-Shaw/sub2api/internal/pkg/logger"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/redis/go-redis/v9"
)
@@ -18,18 +17,41 @@ const (
umqKeyPrefix = "umq:"
umqLockSuffix = ":lock" // STRING (requestID), PX lockTtlMs
umqLastSuffix = ":last" // STRING (毫秒时间戳), EX 60s
// 锁索引用来替代后台清理对 umq:*:lock 的全量 SCAN。
// member 是 accountID,score 是锁预计过期的 Redis Unix 毫秒时间戳。
umqLockIndexKey = "umq:lock:index" // ZSET member=accountID, score=lockExpireAtUnixMs
umqLockIndexCleanupBatchSize = 1000
)
// Lua 脚本:原子获取串行锁(SET NX PX + 重入安全)
// 返回 {是否获取成功, 锁预计过期时间毫秒},让 Go 侧用同一 Redis 时间源更新索引。
// 获取失败(锁被他人持有)时也返回观测到的到期时间,供 Go 侧回填锁索引:
// 这让升级窗口遗留、索引写失败、释放竞态误删索引的存量锁在下一次被争用时自动重新入索引,
// 是替代旧 SCAN 兜底的自愈机制。PTTL == -1 的异常锁返回当前时间,使其立即成为 reconcile 候选。
var acquireLockScript = redis.NewScript(`
redis.replicate_commands()
local cur = redis.call('GET', KEYS[1])
local ttl = tonumber(ARGV[2])
if cur == ARGV[1] then
redis.call('PEXPIRE', KEYS[1], tonumber(ARGV[2]))
return 1
redis.call('PEXPIRE', KEYS[1], ttl)
local t = redis.call('TIME')
local ms = tonumber(t[1])*1000 + math.floor(tonumber(t[2])/1000)
return {1, ms + ttl}
end
if cur ~= false then return 0 end
redis.call('SET', KEYS[1], ARGV[1], 'PX', tonumber(ARGV[2]))
return 1
if cur ~= false then
local t = redis.call('TIME')
local ms = tonumber(t[1])*1000 + math.floor(tonumber(t[2])/1000)
local pttl = redis.call('PTTL', KEYS[1])
if pttl and pttl > 0 then
return {0, ms + pttl}
end
return {0, ms}
end
redis.call('SET', KEYS[1], ARGV[1], 'PX', ttl)
local t = redis.call('TIME')
local ms = tonumber(t[1])*1000 + math.floor(tonumber(t[2])/1000)
return {1, ms + ttl}
`)
// Lua 脚本:原子释放锁 + 记录完成时间(使用 Redis TIME 避免时钟偏差)
@@ -48,14 +70,18 @@ end
return 0
`)
// Lua 脚本:原子清理孤儿锁(仅在 PTTL == -1 时删除,避免 TOCTOU 竞态误删合法锁)
var forceReleaseLockScript = redis.NewScript(`
// Lua 脚本:校验锁 TTL 状态,PTTL == -1 时原子删除异常锁。
// 返回状态: -2=锁不存在,-1=无 TTL 的异常锁已删除,1=锁仍存活并返回剩余 PTTL。
var reconcileLockScript = redis.NewScript(`
local pttl = redis.call('PTTL', KEYS[1])
if pttl == -2 then
return {-2, 0}
end
if pttl == -1 then
redis.call('DEL', KEYS[1])
return 1
return {-1, 0}
end
return 0
return {1, pttl}
`)
type userMsgQueueCache struct {
@@ -77,22 +103,36 @@ func umqLastKey(accountID int64) string {
return umqKeyPrefix + "{" + strconv.FormatInt(accountID, 10) + "}" + umqLastSuffix
}
// umqScanPattern 用于 SCAN 扫描锁 key
func umqScanPattern() string {
return umqKeyPrefix + "{*}" + umqLockSuffix
}
// AcquireLock 尝试获取账号级串行锁
// 无论成功与否都尽力写入锁索引:成功时登记自己的锁,失败时回填观测到的持有者锁,
// 保证任何被争用的锁都能被后台 reconcile 发现,无需扫描所有锁 key。
func (c *userMsgQueueCache) AcquireLock(ctx context.Context, accountID int64, requestID string, lockTtlMs int) (bool, error) {
key := umqLockKey(accountID)
result, err := acquireLockScript.Run(ctx, c.rdb, []string{key}, requestID, lockTtlMs).Int()
result, err := acquireLockScript.Run(ctx, c.rdb, []string{key}, requestID, lockTtlMs).Result()
if err != nil {
return false, fmt.Errorf("umq acquire lock: %w", err)
}
return result == 1, nil
acquired, err := redisScriptInt64At(result, 0)
if err != nil {
return false, fmt.Errorf("umq parse acquire lock result: %w", err)
}
expireAtMs, err := redisScriptInt64At(result, 1)
if err != nil {
return false, fmt.Errorf("umq parse acquire lock expire: %w", err)
}
if expireAtMs > 0 {
if err := c.rdb.ZAdd(ctx, umqLockIndexKey, redis.Z{
Score: float64(expireAtMs),
Member: strconv.FormatInt(accountID, 10),
}).Err(); err != nil {
logger.LegacyPrintf("repository.umq", "Warning: update lock index for account %d failed: %v", accountID, err)
}
}
return acquired == 1, nil
}
// ReleaseLock 释放锁并记录完成时间
// 只有 requestID 匹配时才删除锁索引,避免误删其他请求重入后写入的新锁。
func (c *userMsgQueueCache) ReleaseLock(ctx context.Context, accountID int64, requestID string) (bool, error) {
lockKey := umqLockKey(accountID)
lastKey := umqLastKey(accountID)
@@ -100,6 +140,13 @@ func (c *userMsgQueueCache) ReleaseLock(ctx context.Context, accountID int64, re
if err != nil {
return false, fmt.Errorf("umq release lock: %w", err)
}
if result == 1 {
// 与下一个 AcquireLock 的 ZAdd 存在竞态:可能误删新持有者刚写入的索引项。
// 该锁下次被争用时 AcquireLock 的回填路径会重新登记,无需在此加锁。
if err := c.rdb.ZRem(ctx, umqLockIndexKey, strconv.FormatInt(accountID, 10)).Err(); err != nil {
logger.LegacyPrintf("repository.umq", "Warning: remove lock index for account %d failed: %v", accountID, err)
}
}
return result == 1, nil
}
@@ -120,65 +167,6 @@ func (c *userMsgQueueCache) GetLastCompletedMs(ctx context.Context, accountID in
return ms, nil
}
// ForceReleaseLock 原子清理孤儿锁(仅在 PTTL == -1 时删除,防止 TOCTOU 竞态误删合法锁)
func (c *userMsgQueueCache) ForceReleaseLock(ctx context.Context, accountID int64) error {
key := umqLockKey(accountID)
_, err := forceReleaseLockScript.Run(ctx, c.rdb, []string{key}).Result()
if err != nil && !errors.Is(err, redis.Nil) {
return fmt.Errorf("umq force release lock: %w", err)
}
return nil
}
// ScanLockKeys 扫描所有锁 key,仅返回 PTTL == -1(无过期时间)的孤儿锁 accountID 列表
// 正常的锁都有 PX 过期时间,PTTL == -1 表示异常状态(如 Redis 故障恢复后丢失 TTL)
func (c *userMsgQueueCache) ScanLockKeys(ctx context.Context, maxCount int) ([]int64, error) {
var accountIDs []int64
var cursor uint64
pattern := umqScanPattern()
for {
keys, nextCursor, err := c.rdb.Scan(ctx, cursor, pattern, 100).Result()
if err != nil {
return nil, fmt.Errorf("umq scan lock keys: %w", err)
}
for _, key := range keys {
// 检查 PTTL:只清理 PTTL == -1(无过期时间)的异常锁
pttl, err := c.rdb.PTTL(ctx, key).Result()
if err != nil {
continue
}
// PTTL 返回值:-2 = key 不存在,-1 = 无过期时间,>0 = 剩余毫秒
// go-redis 对哨兵值 -1/-2 不乘精度系数,直接返回 time.Duration(-1)/-2
// 只删除 -1(无过期时间的异常锁),跳过正常持有的锁
if pttl != time.Duration(-1) {
continue
}
// 从 key 中提取 accountID: umq:{123}:lock → 提取 {} 内的数字
openBrace := strings.IndexByte(key, '{')
closeBrace := strings.IndexByte(key, '}')
if openBrace < 0 || closeBrace <= openBrace+1 {
continue
}
idStr := key[openBrace+1 : closeBrace]
id, err := strconv.ParseInt(idStr, 10, 64)
if err != nil {
continue
}
accountIDs = append(accountIDs, id)
if len(accountIDs) >= maxCount {
return accountIDs, nil
}
}
cursor = nextCursor
if cursor == 0 {
break
}
}
return accountIDs, nil
}
// GetCurrentTimeMs 通过 Redis TIME 命令获取当前服务器时间(毫秒),确保与锁记录的时间源一致
func (c *userMsgQueueCache) GetCurrentTimeMs(ctx context.Context) (int64, error) {
t, err := c.rdb.Time(ctx).Result()
@@ -187,3 +175,94 @@ func (c *userMsgQueueCache) GetCurrentTimeMs(ctx context.Context) (int64, error)
}
return t.UnixMilli(), nil
}
// ReconcileExpiredLockCandidates 只处理索引里已经到期的候选锁。
// 候选到期不等于锁一定失效:可能是续租后索引滞后,所以必须再用 PTTL 二次确认。
func (c *userMsgQueueCache) ReconcileExpiredLockCandidates(ctx context.Context, maxCount int) (int, error) {
if maxCount <= 0 {
maxCount = umqLockIndexCleanupBatchSize
}
nowMs, err := c.GetCurrentTimeMs(ctx)
if err != nil {
return 0, err
}
members, err := c.rdb.ZRangeByScore(ctx, umqLockIndexKey, &redis.ZRangeBy{
Min: "-inf",
Max: strconv.FormatInt(nowMs, 10),
Count: int64(maxCount),
}).Result()
if err != nil {
return 0, fmt.Errorf("umq read lock index: %w", err)
}
cleaned := 0
for _, member := range members {
accountID, err := strconv.ParseInt(member, 10, 64)
if err != nil || accountID <= 0 {
c.removeLockIndexMember(ctx, member)
continue
}
result, err := reconcileLockScript.Run(ctx, c.rdb, []string{umqLockKey(accountID)}).Result()
if err != nil && !errors.Is(err, redis.Nil) {
return cleaned, fmt.Errorf("umq reconcile lock: %w", err)
}
status, err := redisScriptInt64At(result, 0)
if err != nil {
return cleaned, fmt.Errorf("umq parse reconcile status: %w", err)
}
pttl, err := redisScriptInt64At(result, 1)
if err != nil {
return cleaned, fmt.Errorf("umq parse reconcile pttl: %w", err)
}
switch status {
case -2:
// 锁自然过期或已释放,只需移除索引残留。
c.removeLockIndexMember(ctx, member)
case -1:
// 无 TTL 的锁会永久阻塞队列,Lua 已原子删除它,这里统计一次清理。
c.removeLockIndexMember(ctx, member)
cleaned++
case 1:
// 锁仍存活,说明索引过期时间滞后;按剩余 PTTL 重新排期。
if err := c.rdb.ZAdd(ctx, umqLockIndexKey, redis.Z{
Score: float64(nowMs + pttl),
Member: member,
}).Err(); err != nil {
logger.LegacyPrintf("repository.umq", "Warning: reschedule lock index member %s failed: %v", member, err)
}
}
}
return cleaned, nil
}
// removeLockIndexMember 移除锁索引残留;索引维护是 best-effort,失败只记日志。
func (c *userMsgQueueCache) removeLockIndexMember(ctx context.Context, member string) {
if err := c.rdb.ZRem(ctx, umqLockIndexKey, member).Err(); err != nil {
logger.LegacyPrintf("repository.umq", "Warning: remove lock index member %s failed: %v", member, err)
}
}
// redisScriptInt64At 兼容 go-redis 对 Lua 数组元素的不同返回类型。
func redisScriptInt64At(result any, index int) (int64, error) {
values, ok := result.([]any)
if !ok {
return 0, fmt.Errorf("expected redis script array, got %T", result)
}
if index < 0 || index >= len(values) {
return 0, fmt.Errorf("redis script array missing index %d", index)
}
switch v := values[index].(type) {
case int64:
return v, nil
case int:
return int64(v), nil
case string:
return strconv.ParseInt(v, 10, 64)
case []byte:
return strconv.ParseInt(string(v), 10, 64)
default:
return 0, fmt.Errorf("unexpected redis script value %T", v)
}
}
@@ -0,0 +1,177 @@
//go:build integration
package repository
import (
"errors"
"testing"
"time"
"github.com/redis/go-redis/v9"
"github.com/stretchr/testify/require"
"github.com/stretchr/testify/suite"
)
type UserMsgQueueCacheSuite struct {
IntegrationRedisSuite
cache *userMsgQueueCache
}
func TestUserMsgQueueCacheSuite(t *testing.T) {
suite.Run(t, new(UserMsgQueueCacheSuite))
}
func (s *UserMsgQueueCacheSuite) SetupTest() {
s.IntegrationRedisSuite.SetupTest()
s.cache = NewUserMsgQueueCache(s.rdb).(*userMsgQueueCache)
}
func (s *UserMsgQueueCacheSuite) TestAcquireLockWritesIndexAndReleaseRemovesIt() {
accountID := int64(701)
nowMs, err := s.cache.GetCurrentTimeMs(s.ctx)
require.NoError(s.T(), err)
acquired, err := s.cache.AcquireLock(s.ctx, accountID, "req-701", 10_000)
require.NoError(s.T(), err)
require.True(s.T(), acquired)
score, err := s.rdb.ZScore(s.ctx, umqLockIndexKey, "701").Result()
require.NoError(s.T(), err)
require.Greater(s.T(), int64(score), nowMs)
released, err := s.cache.ReleaseLock(s.ctx, accountID, "req-701")
require.NoError(s.T(), err)
require.True(s.T(), released)
_, err = s.rdb.ZScore(s.ctx, umqLockIndexKey, "701").Result()
require.ErrorIs(s.T(), err, redis.Nil)
}
func (s *UserMsgQueueCacheSuite) TestReconcileExpiredLockCandidatesRemovesNaturallyExpiredLockIndex() {
accountID := int64(702)
acquired, err := s.cache.AcquireLock(s.ctx, accountID, "req-702", 20)
require.NoError(s.T(), err)
require.True(s.T(), acquired)
score, err := s.rdb.ZScore(s.ctx, umqLockIndexKey, "702").Result()
require.NoError(s.T(), err)
require.Eventually(s.T(), func() bool {
nowMs, err := s.cache.GetCurrentTimeMs(s.ctx)
return err == nil && nowMs >= int64(score)
}, time.Second, 10*time.Millisecond)
cleaned, err := s.cache.ReconcileExpiredLockCandidates(s.ctx, 1000)
require.NoError(s.T(), err)
require.Equal(s.T(), 0, cleaned)
_, err = s.rdb.ZScore(s.ctx, umqLockIndexKey, "702").Result()
require.ErrorIs(s.T(), err, redis.Nil)
}
func (s *UserMsgQueueCacheSuite) TestReconcileExpiredLockCandidatesRefreshesLiveLockIndex() {
accountID := int64(703)
nowMs, err := s.cache.GetCurrentTimeMs(s.ctx)
require.NoError(s.T(), err)
require.NoError(s.T(), s.rdb.Set(s.ctx, umqLockKey(accountID), "req-703", time.Minute).Err())
require.NoError(s.T(), s.rdb.ZAdd(s.ctx, umqLockIndexKey, redis.Z{
Score: float64(nowMs - 1),
Member: "703",
}).Err())
cleaned, err := s.cache.ReconcileExpiredLockCandidates(s.ctx, 1000)
require.NoError(s.T(), err)
require.Equal(s.T(), 0, cleaned)
score, err := s.rdb.ZScore(s.ctx, umqLockIndexKey, "703").Result()
require.NoError(s.T(), err)
require.Greater(s.T(), int64(score), nowMs)
exists, err := s.rdb.Exists(s.ctx, umqLockKey(accountID)).Result()
require.NoError(s.T(), err)
require.EqualValues(s.T(), 1, exists)
}
func (s *UserMsgQueueCacheSuite) TestReconcileExpiredLockCandidatesDeletesNoTTLLock() {
accountID := int64(704)
nowMs, err := s.cache.GetCurrentTimeMs(s.ctx)
require.NoError(s.T(), err)
require.NoError(s.T(), s.rdb.Set(s.ctx, umqLockKey(accountID), "req-704", 0).Err())
require.NoError(s.T(), s.rdb.ZAdd(s.ctx, umqLockIndexKey, redis.Z{
Score: float64(nowMs),
Member: "704",
}).Err())
cleaned, err := s.cache.ReconcileExpiredLockCandidates(s.ctx, 1000)
require.NoError(s.T(), err)
require.Equal(s.T(), 1, cleaned)
exists, err := s.rdb.Exists(s.ctx, umqLockKey(accountID)).Result()
require.NoError(s.T(), err)
require.EqualValues(s.T(), 0, exists)
_, err = s.rdb.ZScore(s.ctx, umqLockIndexKey, "704").Result()
require.ErrorIs(s.T(), err, redis.Nil)
}
func (s *UserMsgQueueCacheSuite) TestReconcileExpiredLockCandidatesRemovesInvalidMember() {
nowMs, err := s.cache.GetCurrentTimeMs(s.ctx)
require.NoError(s.T(), err)
require.NoError(s.T(), s.rdb.ZAdd(s.ctx, umqLockIndexKey, redis.Z{
Score: float64(nowMs),
Member: "not-an-account-id",
}).Err())
cleaned, err := s.cache.ReconcileExpiredLockCandidates(s.ctx, 1000)
require.NoError(s.T(), err)
require.Equal(s.T(), 0, cleaned)
_, err = s.rdb.ZScore(s.ctx, umqLockIndexKey, "not-an-account-id").Result()
require.True(s.T(), errors.Is(err, redis.Nil))
}
func (s *UserMsgQueueCacheSuite) TestAcquireLockBusyPathReindexesUnindexedLiveLock() {
// 模拟索引丢失的存量锁(升级窗口/索引写失败/释放竞态误删):
// 锁存在且有 TTL,但索引里没有对应 member。
accountID := int64(705)
nowMs, err := s.cache.GetCurrentTimeMs(s.ctx)
require.NoError(s.T(), err)
require.NoError(s.T(), s.rdb.Set(s.ctx, umqLockKey(accountID), "holder-705", time.Minute).Err())
// 另一个请求争锁失败,应顺手把观测到的持有者锁回填进索引。
acquired, err := s.cache.AcquireLock(s.ctx, accountID, "contender-705", 10_000)
require.NoError(s.T(), err)
require.False(s.T(), acquired)
score, err := s.rdb.ZScore(s.ctx, umqLockIndexKey, "705").Result()
require.NoError(s.T(), err, "busy acquire should re-index the observed live lock")
require.Greater(s.T(), int64(score), nowMs)
// 锁本身不应被争锁方改动。
val, err := s.rdb.Get(s.ctx, umqLockKey(accountID)).Result()
require.NoError(s.T(), err)
require.Equal(s.T(), "holder-705", val)
}
func (s *UserMsgQueueCacheSuite) TestAcquireLockBusyPathMakesNoTTLLockReconcilable() {
// PTTL == -1 的异常锁若不在索引中,永远不会被 reconcile 发现;
// 争锁失败路径必须以“已到期候选”的 score 回填它,形成自愈闭环。
accountID := int64(706)
require.NoError(s.T(), s.rdb.Set(s.ctx, umqLockKey(accountID), "holder-706", 0).Err())
acquired, err := s.cache.AcquireLock(s.ctx, accountID, "contender-706", 10_000)
require.NoError(s.T(), err)
require.False(s.T(), acquired)
nowMs, err := s.cache.GetCurrentTimeMs(s.ctx)
require.NoError(s.T(), err)
score, err := s.rdb.ZScore(s.ctx, umqLockIndexKey, "706").Result()
require.NoError(s.T(), err, "busy acquire should index the anomalous lock")
require.LessOrEqual(s.T(), int64(score), nowMs, "anomalous lock should be an immediately-expired candidate")
cleaned, err := s.cache.ReconcileExpiredLockCandidates(s.ctx, 1000)
require.NoError(s.T(), err)
require.Equal(s.T(), 1, cleaned, "reconcile should delete the no-TTL lock")
exists, err := s.rdb.Exists(s.ctx, umqLockKey(accountID)).Result()
require.NoError(s.T(), err)
require.EqualValues(s.T(), 0, exists, "queue is unblocked after reconcile")
_, err = s.rdb.ZScore(s.ctx, umqLockIndexKey, "706").Result()
require.ErrorIs(s.T(), err, redis.Nil)
}
@@ -603,6 +603,13 @@ func (s *AccountTestService) testOpenAIAccountConnection(c *gin.Context, account
if isOAuth {
req.Host = "chatgpt.com"
req.Header.Set("accept", "text/event-stream")
req.Header.Set("OpenAI-Beta", "responses=experimental")
req.Header.Set("Originator", "codex_cli_rs")
if customUA := strings.TrimSpace(credentialAccount.GetOpenAIUserAgent()); customUA != "" {
req.Header.Set("User-Agent", customUA)
} else {
req.Header.Set("User-Agent", codexCLIUserAgent)
}
setOpenAIChatGPTAccountHeaders(req.Header, credentialAccount)
}
@@ -37,7 +37,7 @@ type ConcurrencyCache interface {
ReleaseUserSlot(ctx context.Context, userID int64, requestID string) error
GetUserConcurrency(ctx context.Context, userID int64) (int, error)
// 等待队列计数(只在首次创建时设置 TTL)
// 等待队列计数(每次入队都会刷新 TTL,避免长时间排队时计数提前过期)
IncrementWaitCount(ctx context.Context, userID int64, maxWait int) (bool, error)
DecrementWaitCount(ctx context.Context, userID int64) error
+1
View File
@@ -483,6 +483,7 @@ func grokMediaUsageFromResponse(endpoint GrokMediaEndpoint, requestInfo GrokMedi
meta.ImageOutputSizes = collectOpenAIResponseImageOutputSizesFromJSONBytes(responseBody)
case GrokMediaEndpointVideosGenerations:
meta.ResponseID = extractGrokMediaVideoRequestID(responseBody)
// Video generation is one billable media unit; the legacy usage schema stores it in ImageCount.
meta.ImageCount = 1
meta.ImageSize = requestInfo.SizeTier
meta.ImageInputSize = requestInfo.Size
@@ -1,6 +1,7 @@
package service
import (
"fmt"
"net/http"
"github.com/Wei-Shaw/sub2api/internal/config"
@@ -8,6 +9,11 @@ import (
"github.com/gin-gonic/gin"
)
// CodexOfficialClientsOnlyMessage 是 codex_cli_only 拒绝时面向客户端的通用兜底文案。
// 仅当拒绝原因不是「可解析版本但越界」(VersionTooLow/VersionTooHigh)时使用:
// 未命中官方/黑名单/缺指纹/版本无法识别都沿用这句(避免向伪装客户端泄露门控细节)。
const CodexOfficialClientsOnlyMessage = "This account only allows Codex official clients"
const (
// CodexClientRestrictionReasonDisabled 表示账号未开启 codex_cli_only。
CodexClientRestrictionReasonDisabled = "codex_cli_only_disabled"
@@ -51,6 +57,13 @@ type CodexClientRestrictionDetectionResult struct {
Enabled bool
Matched bool
Reason string
// DetectedVersion 是从官方 UA 解析出的 Codex 引擎版本;仅在版本门拒绝
// (VersionTooLow / VersionTooHigh) 时填充,供面向客户端的差异化文案使用。
DetectedVersion string
// MinCodexVersion 是触发 VersionTooLow 时的最低要求版本(来自策略快照)。
MinCodexVersion string
// MaxCodexVersion 是触发 VersionTooHigh 时的最高允许版本(来自策略快照)。
MaxCodexVersion string
}
// CodexClientRestrictionDetector 定义 codex_cli_only 统一检测入口。
@@ -127,10 +140,22 @@ func (d *OpenAICodexClientRestrictionDetector) Detect(c *gin.Context, account *A
return CodexClientRestrictionDetectionResult{Enabled: true, Matched: false, Reason: CodexClientRestrictionReasonVersionUndetectable}
}
if policy.MinCodexVersion != "" && CompareVersions(ver, policy.MinCodexVersion) < 0 {
return CodexClientRestrictionDetectionResult{Enabled: true, Matched: false, Reason: CodexClientRestrictionReasonVersionTooLow}
return CodexClientRestrictionDetectionResult{
Enabled: true,
Matched: false,
Reason: CodexClientRestrictionReasonVersionTooLow,
DetectedVersion: ver,
MinCodexVersion: policy.MinCodexVersion,
}
}
if policy.MaxCodexVersion != "" && CompareVersions(ver, policy.MaxCodexVersion) > 0 {
return CodexClientRestrictionDetectionResult{Enabled: true, Matched: false, Reason: CodexClientRestrictionReasonVersionTooHigh}
return CodexClientRestrictionDetectionResult{
Enabled: true,
Matched: false,
Reason: CodexClientRestrictionReasonVersionTooHigh,
DetectedVersion: ver,
MaxCodexVersion: policy.MaxCodexVersion,
}
}
}
@@ -145,3 +170,22 @@ func (d *OpenAICodexClientRestrictionDetector) Detect(c *gin.Context, account *A
return CodexClientRestrictionDetectionResult{Enabled: true, Matched: true, Reason: reason}
}
// CodexClientRestrictionMessage 把检测结果映射为面向客户端的 403 文案。
// 仅版本越界(VersionTooLow/VersionTooHigh)给出带实际版本号与边界的差异化提示——
// 这类请求其实已被识别为官方 Codex(命中官方 UA/originator),再回「只允许官方客户端」会误导;
// 其余拒绝原因统一沿用通用兜底句,不暴露门控细节。
func CodexClientRestrictionMessage(r CodexClientRestrictionDetectionResult) string {
switch r.Reason {
case CodexClientRestrictionReasonVersionTooLow:
return fmt.Sprintf(
"Your Codex version (%s) is below the minimum required version (%s). Please update Codex.",
r.DetectedVersion, r.MinCodexVersion)
case CodexClientRestrictionReasonVersionTooHigh:
return fmt.Sprintf(
"Your Codex version (%s) exceeds the maximum allowed version (%s). Please downgrade Codex to %s or lower.",
r.DetectedVersion, r.MaxCodexVersion, r.MaxCodexVersion)
default:
return CodexOfficialClientsOnlyMessage
}
}
@@ -284,6 +284,66 @@ func TestDetect_V3_AppServerAndSkipAndVersionScope(t *testing.T) {
})
}
func TestDetect_VersionGateCarriesVersionFields(t *testing.T) {
gin.SetMode(gin.TestMode)
d := NewOpenAICodexClientRestrictionDetector(nil)
acc := func() *Account {
return &Account{Platform: PlatformOpenAI, Type: AccountTypeOAuth, Extra: map[string]any{"codex_cli_only": true}}
}
t.Run("版本太低:携带 DetectedVersion + MinCodexVersion", func(t *testing.T) {
c := newCodexDetectorTestContext("codex_cli_rs/0.39.0 (x)", "")
r := d.Detect(c, acc(), CodexRestrictionPolicy{MinCodexVersion: "0.42.0"}, nil)
require.False(t, r.Matched)
require.Equal(t, CodexClientRestrictionReasonVersionTooLow, r.Reason)
require.Equal(t, "0.39.0", r.DetectedVersion)
require.Equal(t, "0.42.0", r.MinCodexVersion)
})
t.Run("版本太高:携带 DetectedVersion + MaxCodexVersion", func(t *testing.T) {
c := newCodexDetectorTestContext("codex_cli_rs/0.45.0 (x)", "")
r := d.Detect(c, acc(), CodexRestrictionPolicy{MaxCodexVersion: "0.42.0"}, nil)
require.False(t, r.Matched)
require.Equal(t, CodexClientRestrictionReasonVersionTooHigh, r.Reason)
require.Equal(t, "0.45.0", r.DetectedVersion)
require.Equal(t, "0.42.0", r.MaxCodexVersion)
})
}
func TestCodexClientRestrictionMessage(t *testing.T) {
t.Run("版本太低:带实际版本与最低要求", func(t *testing.T) {
msg := CodexClientRestrictionMessage(CodexClientRestrictionDetectionResult{
Reason: CodexClientRestrictionReasonVersionTooLow,
DetectedVersion: "0.39.0",
MinCodexVersion: "0.42.0",
})
require.Equal(t, "Your Codex version (0.39.0) is below the minimum required version (0.42.0). Please update Codex.", msg)
})
t.Run("版本太高:带实际版本与最高允许", func(t *testing.T) {
msg := CodexClientRestrictionMessage(CodexClientRestrictionDetectionResult{
Reason: CodexClientRestrictionReasonVersionTooHigh,
DetectedVersion: "0.45.0",
MaxCodexVersion: "0.42.0",
})
require.Equal(t, "Your Codex version (0.45.0) exceeds the maximum allowed version (0.42.0). Please downgrade Codex to 0.42.0 or lower.", msg)
})
t.Run("无法识别版本:保持原通用句", func(t *testing.T) {
msg := CodexClientRestrictionMessage(CodexClientRestrictionDetectionResult{
Reason: CodexClientRestrictionReasonVersionUndetectable,
})
require.Equal(t, "This account only allows Codex official clients", msg)
})
t.Run("未命中官方:保持原通用句", func(t *testing.T) {
msg := CodexClientRestrictionMessage(CodexClientRestrictionDetectionResult{
Reason: CodexClientRestrictionReasonNotMatchedUA,
})
require.Equal(t, "This account only allows Codex official clients", msg)
})
}
func TestDetect_EngineFingerprintSignals(t *testing.T) {
gin.SetMode(gin.TestMode)
det := NewOpenAICodexClientRestrictionDetector(&config.Config{})
@@ -1803,6 +1803,52 @@ func TestOpenAIGatewayServiceRecordUsage_ImageIndependentMultiplierUsesImageRate
require.Equal(t, string(BillingModeImage), *usageRepo.lastLog.BillingMode)
}
func TestGrokVideoMediaBillingUsesImageRateMultiplier(t *testing.T) {
mediaPrice2K := 0.4
groupID := int64(126)
usageRepo := &openAIRecordUsageLogRepoStub{inserted: true}
svc := newOpenAIRecordUsageServiceForTest(usageRepo, &openAIRecordUsageUserRepoStub{}, &openAIRecordUsageSubRepoStub{}, nil)
err := svc.RecordUsage(context.Background(), &OpenAIRecordUsageInput{
Result: &OpenAIForwardResult{
RequestID: "video-request-123",
ResponseID: "video-request-123",
Model: "grok-imagine-video-1.5",
BillingModel: "grok-imagine-video-1.5",
// The usage schema has no separate video count; video generation is billed as one media unit.
ImageCount: 1,
ImageSize: ImageBillingSize2K,
Duration: time.Second,
},
APIKey: &APIKey{
ID: 10126,
GroupID: i64p(groupID),
Group: &Group{
ID: groupID,
Platform: PlatformGrok,
RateMultiplier: 0.15,
ImageRateIndependent: true,
ImageRateMultiplier: 0.5,
ImagePrice2K: &mediaPrice2K,
},
},
User: &User{ID: 20126},
Account: &Account{ID: 30126, Platform: PlatformGrok},
})
require.NoError(t, err)
require.NotNil(t, usageRepo.lastLog)
require.Equal(t, "grok-imagine-video-1.5", usageRepo.lastLog.Model)
require.Equal(t, 1, usageRepo.lastLog.ImageCount)
require.Equal(t, ImageBillingSize2K, *usageRepo.lastLog.ImageSize)
require.InDelta(t, 0.4, usageRepo.lastLog.TotalCost, 1e-12)
require.InDelta(t, 0.2, usageRepo.lastLog.ActualCost, 1e-12)
require.InDelta(t, 0.5, usageRepo.lastLog.RateMultiplier, 1e-12)
require.NotNil(t, usageRepo.lastLog.BillingMode)
require.Equal(t, string(BillingModeImage), *usageRepo.lastLog.BillingMode)
}
func TestOpenAIGatewayServiceRecordUsage_ChannelImageBillingUsesImageCountAndSharedMultiplier(t *testing.T) {
groupID := int64(123)
usageRepo := &openAIRecordUsageLogRepoStub{inserted: true}
@@ -2617,7 +2617,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
c.JSON(http.StatusForbidden, gin.H{
"error": gin.H{
"type": "forbidden_error",
"message": "This account only allows Codex official clients",
"message": CodexClientRestrictionMessage(restrictionResult),
},
})
return nil, errors.New("codex_cli_only restriction: only codex official clients are allowed")
@@ -59,6 +59,52 @@ func TestOpenAIGatewayService_GetCodexClientRestrictionDetector(t *testing.T) {
})
}
func TestOpenAIGatewayService_Forward_VersionGateMessage(t *testing.T) {
gin.SetMode(gin.TestMode)
newCtx := func() (*httptest.ResponseRecorder, *gin.Context) {
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(nil))
return rec, c
}
account := func() *Account {
return &Account{Platform: PlatformOpenAI, Type: AccountTypeOAuth, Extra: map[string]any{"codex_cli_only": true}}
}
body := []byte(`{"model":"gpt-5.1-codex"}`)
t.Run("版本太低:返回带版本号的差异化文案", func(t *testing.T) {
rec, c := newCtx()
svc := &OpenAIGatewayService{codexDetector: &stubCodexRestrictionDetector{result: CodexClientRestrictionDetectionResult{
Enabled: true,
Matched: false,
Reason: CodexClientRestrictionReasonVersionTooLow,
DetectedVersion: "0.39.0",
MinCodexVersion: "0.42.0",
}}}
_, err := svc.Forward(context.Background(), c, account(), body)
require.Error(t, err)
require.Equal(t, http.StatusForbidden, rec.Code)
require.Contains(t, rec.Body.String(), "Your Codex version (0.39.0) is below the minimum required version (0.42.0)")
require.NotContains(t, rec.Body.String(), "This account only allows Codex official clients")
})
t.Run("未命中官方:仍返回通用兜底文案", func(t *testing.T) {
rec, c := newCtx()
svc := &OpenAIGatewayService{codexDetector: &stubCodexRestrictionDetector{result: CodexClientRestrictionDetectionResult{
Enabled: true,
Matched: false,
Reason: CodexClientRestrictionReasonNotMatchedUA,
}}}
_, err := svc.Forward(context.Background(), c, account(), body)
require.Error(t, err)
require.Equal(t, http.StatusForbidden, rec.Code)
require.Contains(t, rec.Body.String(), "This account only allows Codex official clients")
})
}
func TestGetAPIKeyIDFromContext(t *testing.T) {
gin.SetMode(gin.TestMode)
+1 -8
View File
@@ -391,14 +391,7 @@ func buildV1ModelsURL(base string) string {
}
func buildOpenAIModelsURL(base string) string {
normalized := strings.TrimRight(strings.TrimSpace(base), "/")
if strings.HasSuffix(normalized, "/v1/models") {
return normalized
}
if strings.HasSuffix(normalized, "/v1") {
return normalized + "/models"
}
return normalized + "/v1/models"
return buildOpenAIEndpointURL(base, "/v1/models")
}
func buildGeminiModelsURL(base string) string {
@@ -29,6 +29,61 @@ func TestBuildV1ModelsURL(t *testing.T) {
require.Equal(t, "https://gateway.example.com/antigravity/v1/models", buildV1ModelsURL("https://gateway.example.com/antigravity/"))
}
func TestBuildOpenAIModelsURL(t *testing.T) {
t.Parallel()
tests := []struct {
name string
base string
want string
}{
{
name: "zhipu v4 coding base url",
base: "https://open.bigmodel.cn/api/coding/paas/v4",
want: "https://open.bigmodel.cn/api/coding/paas/v4/models",
},
{
name: "openai v1 base url",
base: "https://api.openai.com/v1",
want: "https://api.openai.com/v1/models",
},
{
name: "models url unchanged",
base: "https://api.openai.com/v1/models",
want: "https://api.openai.com/v1/models",
},
{
name: "host fallback uses v1",
base: "https://api.openai.com",
want: "https://api.openai.com/v1/models",
},
{
name: "trailing slash on v4",
base: "https://open.bigmodel.cn/api/coding/paas/v4/",
want: "https://open.bigmodel.cn/api/coding/paas/v4/models",
},
{
name: "v2 base url",
base: "https://gateway.example.com/openai/v2",
want: "https://gateway.example.com/openai/v2/models",
},
{
name: "v3 base url",
base: "https://gateway.example.com/openai/v3",
want: "https://gateway.example.com/openai/v3/models",
},
}
for _, tt := range tests {
tt := tt
t.Run(tt.name, func(t *testing.T) {
t.Parallel()
require.Equal(t, tt.want, buildOpenAIModelsURL(tt.base))
})
}
}
func TestBuildGeminiModelsURL(t *testing.T) {
t.Parallel()
@@ -25,10 +25,8 @@ type UserMsgQueueCache interface {
GetLastCompletedMs(ctx context.Context, accountID int64) (int64, error)
// GetCurrentTimeMs 获取 Redis 服务器当前时间(毫秒),与 ReleaseLock 记录的时间源一致
GetCurrentTimeMs(ctx context.Context) (int64, error)
// ForceReleaseLock 强制释放锁(孤儿锁清理)
ForceReleaseLock(ctx context.Context, accountID int64) error
// ScanLockKeys 扫描 PTTL == -1 的孤儿锁 key,返回 accountID 列表
ScanLockKeys(ctx context.Context, maxCount int) ([]int64, error)
// ReconcileExpiredLockCandidates 处理锁索引中的到期候选,按真实 PTTL 清理或刷新索引
ReconcileExpiredLockCandidates(ctx context.Context, maxCount int) (cleaned int, err error)
}
// QueueLockResult 锁获取结果
@@ -246,8 +244,8 @@ func (s *UserMessageQueueService) CalculateRPMAwareDelay(ctx context.Context, ac
return applyJitter(baseDelay, 0.15)
}
// StartCleanupWorker 启动孤儿锁清理 worker
// 定期 SCAN umq:*:lock 并清理 PTTL == -1 的异常锁(PTTL 检查在 cache.ScanLockKeys 内完成)
// StartCleanupWorker 启动孤儿锁清理 worker。
// worker 只处理锁索引中的到期候选,真正删除前由 cache 层再次校验锁 PTTL。
func (s *UserMessageQueueService) StartCleanupWorker(interval time.Duration) {
if s == nil || s.cache == nil || interval <= 0 {
return
@@ -257,23 +255,13 @@ func (s *UserMessageQueueService) StartCleanupWorker(interval time.Duration) {
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
accountIDs, err := s.cache.ScanLockKeys(ctx, 1000)
// 每轮限制处理数量,避免清理任务在大量过期候选时长时间占用 Redis。
cleaned, err := s.cache.ReconcileExpiredLockCandidates(ctx, 1000)
if err != nil {
logger.LegacyPrintf("service.umq", "Cleanup scan failed: %v", err)
logger.LegacyPrintf("service.umq", "Cleanup reconcile failed: %v", err)
return
}
cleaned := 0
for _, accountID := range accountIDs {
cleanCtx, cleanCancel := context.WithTimeout(context.Background(), 2*time.Second)
if err := s.cache.ForceReleaseLock(cleanCtx, accountID); err != nil {
logger.LegacyPrintf("service.umq", "Cleanup force release failed for account %d: %v", accountID, err)
} else {
cleaned++
}
cleanCancel()
}
if cleaned > 0 {
logger.LegacyPrintf("service.umq", "Cleanup completed: released %d orphaned locks", cleaned)
}
@@ -0,0 +1,54 @@
//go:build unit
package service
import (
"context"
"sync/atomic"
"testing"
"time"
"github.com/stretchr/testify/require"
)
type cleanupWorkerUserMsgQueueCache struct {
reconcileCalls atomic.Int64
maxCount atomic.Int64
}
var _ UserMsgQueueCache = (*cleanupWorkerUserMsgQueueCache)(nil)
func (c *cleanupWorkerUserMsgQueueCache) AcquireLock(context.Context, int64, string, int) (bool, error) {
return true, nil
}
func (c *cleanupWorkerUserMsgQueueCache) ReleaseLock(context.Context, int64, string) (bool, error) {
return true, nil
}
func (c *cleanupWorkerUserMsgQueueCache) GetLastCompletedMs(context.Context, int64) (int64, error) {
return 0, nil
}
func (c *cleanupWorkerUserMsgQueueCache) GetCurrentTimeMs(context.Context) (int64, error) {
return time.Now().UnixMilli(), nil
}
func (c *cleanupWorkerUserMsgQueueCache) ReconcileExpiredLockCandidates(_ context.Context, maxCount int) (int, error) {
c.reconcileCalls.Add(1)
c.maxCount.Store(int64(maxCount))
return 1, nil
}
func TestStartCleanupWorker_ReconcilesExpiredLockCandidates(t *testing.T) {
cache := &cleanupWorkerUserMsgQueueCache{}
svc := NewUserMessageQueueService(cache, nil, nil)
defer svc.Stop()
svc.StartCleanupWorker(time.Millisecond)
require.Eventually(t, func() bool {
return cache.reconcileCalls.Load() > 0
}, time.Second, 10*time.Millisecond)
require.EqualValues(t, 1000, cache.maxCount.Load())
}
@@ -1,14 +1,16 @@
import { describe, it, expect, vi, beforeEach } from 'vitest'
import { mount } from '@vue/test-utils'
import { flushPromises, mount } from '@vue/test-utils'
import ImportDataModal from '@/components/admin/account/ImportDataModal.vue'
const showError = vi.fn()
const showSuccess = vi.fn()
const showWarning = vi.fn()
vi.mock('@/stores/app', () => ({
useAppStore: () => ({
showError,
showSuccess
showSuccess,
showWarning
})
}))
@@ -26,49 +28,187 @@ vi.mock('vue-i18n', () => ({
})
}))
const mountModal = () =>
mount(ImportDataModal, {
props: { show: true },
global: {
stubs: {
BaseDialog: { template: '<div><slot /><slot name="footer" /></div>' }
}
}
})
const makeJsonFile = (name: string, content: string, type = 'application/json') => {
const file = new File([content], name, { type })
Object.defineProperty(file, 'text', {
value: () => Promise.resolve(content)
})
return file
}
const setInputFiles = (element: Element, files: File[]) => {
Object.defineProperty(element, 'files', {
value: files,
configurable: true
})
}
describe('ImportDataModal', () => {
beforeEach(() => {
beforeEach(async () => {
showError.mockReset()
showSuccess.mockReset()
showWarning.mockReset()
const { adminAPI } = await import('@/api/admin')
vi.mocked(adminAPI.accounts.importData).mockReset()
})
it('未选择文件时提示错误', async () => {
const wrapper = mount(ImportDataModal, {
props: { show: true },
global: {
stubs: {
BaseDialog: { template: '<div><slot /><slot name="footer" /></div>' }
}
}
})
const wrapper = mountModal()
await wrapper.find('form').trigger('submit')
expect(showError).toHaveBeenCalledWith('admin.accounts.dataImportSelectFile')
})
it('无效 JSON 时提示解析失败', async () => {
const wrapper = mount(ImportDataModal, {
props: { show: true },
global: {
stubs: {
BaseDialog: { template: '<div><slot /><slot name="footer" /></div>' }
}
}
})
it('无效 JSON 时按文件名提示解析失败', async () => {
const { adminAPI } = await import('@/api/admin')
const wrapper = mountModal()
const input = wrapper.find('input[type="file"]')
const file = new File(['invalid json'], 'data.json', { type: 'application/json' })
Object.defineProperty(file, 'text', {
value: () => Promise.resolve('invalid json')
})
Object.defineProperty(input.element, 'files', {
value: [file]
})
setInputFiles(input.element, [makeJsonFile('data.json', 'invalid json')])
await input.trigger('change')
await wrapper.find('form').trigger('submit')
await Promise.resolve()
await flushPromises()
expect(showError).toHaveBeenCalledWith('admin.accounts.dataImportParseFailed')
expect(showError).toHaveBeenCalledWith('admin.accounts.dataImportParseFailedFile')
expect(adminAPI.accounts.importData).not.toHaveBeenCalled()
})
it('不是导出数据的 JSON 按文件名拒绝', async () => {
const { adminAPI } = await import('@/api/admin')
const wrapper = mountModal()
const input = wrapper.find('input[type="file"]')
setInputFiles(input.element, [makeJsonFile('random.json', JSON.stringify({ name: 'test' }))])
await input.trigger('change')
await wrapper.find('form').trigger('submit')
await flushPromises()
expect(showError).toHaveBeenCalledWith('admin.accounts.dataImportInvalidFile')
expect(adminAPI.accounts.importData).not.toHaveBeenCalled()
})
it('无有效 JSON 的选择不清空已有选择', async () => {
const { adminAPI } = await import('@/api/admin')
vi.mocked(adminAPI.accounts.importData).mockResolvedValue({
proxy_created: 0,
proxy_reused: 0,
proxy_failed: 0,
account_created: 1,
account_failed: 0
})
const wrapper = mountModal()
const input = wrapper.find('input[type="file"]')
const valid = makeJsonFile(
'valid.json',
JSON.stringify({ exported_at: '2026-07-05T00:00:00Z', proxies: [], accounts: [{ name: 'a' }] })
)
setInputFiles(input.element, [valid])
await input.trigger('change')
setInputFiles(input.element, [new File(['hello'], 'notes.txt', { type: 'text/plain' })])
await input.trigger('change')
expect(showError).toHaveBeenCalledWith('admin.accounts.dataImportSelectFile')
await wrapper.find('form').trigger('submit')
await flushPromises()
expect(adminAPI.accounts.importData).toHaveBeenCalledWith({
data: expect.objectContaining({
accounts: [{ name: 'a' }]
}),
skip_default_group_bind: true
})
})
it('merges multiple selected JSON files before importing', async () => {
const { adminAPI } = await import('@/api/admin')
vi.mocked(adminAPI.accounts.importData).mockResolvedValue({
proxy_created: 0,
proxy_reused: 0,
proxy_failed: 0,
account_created: 2,
account_failed: 0
})
const wrapper = mountModal()
const input = wrapper.find('input[type="file"]')
const first = makeJsonFile(
'first.json',
JSON.stringify({ exported_at: '2026-07-05T00:00:00Z', proxies: [], accounts: [{ name: 'a' }] })
)
const second = makeJsonFile(
'second.json',
JSON.stringify({
exported_at: '2026-07-05T00:00:01Z',
proxies: [{ proxy_key: 'p' }],
accounts: [{ name: 'b' }]
})
)
setInputFiles(input.element, [first, second])
await input.trigger('change')
await wrapper.find('form').trigger('submit')
await flushPromises()
expect(adminAPI.accounts.importData).toHaveBeenCalledWith({
data: expect.objectContaining({
proxies: [{ proxy_key: 'p' }],
accounts: [{ name: 'a' }, { name: 'b' }]
}),
skip_default_group_bind: true
})
expect(showSuccess).toHaveBeenCalledWith('admin.accounts.dataImportSuccess')
})
it('部分成功时关闭弹窗仍通知父组件刷新', async () => {
const { adminAPI } = await import('@/api/admin')
vi.mocked(adminAPI.accounts.importData).mockResolvedValue({
proxy_created: 0,
proxy_reused: 0,
proxy_failed: 0,
account_created: 1,
account_failed: 1
})
const wrapper = mountModal()
const input = wrapper.find('input[type="file"]')
setInputFiles(input.element, [
makeJsonFile(
'mixed.json',
JSON.stringify({
exported_at: '2026-07-05T00:00:00Z',
proxies: [],
accounts: [{ name: 'a' }, { name: 'b' }]
})
)
])
await input.trigger('change')
await wrapper.find('form').trigger('submit')
await flushPromises()
expect(showError).toHaveBeenCalledWith('admin.accounts.dataImportCompletedWithErrors')
expect(wrapper.emitted('imported')).toBeUndefined()
// 第二个 btn-secondary 是 footer 的取消按钮(第一个是选择文件)
await wrapper.findAll('button.btn-secondary')[1]!.trigger('click')
expect(wrapper.emitted('imported')).toHaveLength(1)
expect(wrapper.emitted('close')).toHaveLength(1)
})
})
+2
View File
@@ -25,6 +25,7 @@ export interface AdminPaymentConfig {
balance_disabled: boolean
balance_recharge_multiplier: number
subscription_usd_to_cny_rate: number
recharge_fee_rate: number
load_balance_strategy: string
product_name_prefix: string
product_name_suffix: string
@@ -44,6 +45,7 @@ export interface UpdatePaymentConfigRequest {
balance_disabled?: boolean
balance_recharge_multiplier?: number
subscription_usd_to_cny_rate?: number
recharge_fee_rate?: number
load_balance_strategy?: string
product_name_prefix?: string
product_name_suffix?: string
@@ -55,6 +55,17 @@
/>
</div>
<div v-if="isOpenAIAccount" class="space-y-1.5">
<label class="text-sm font-medium text-gray-700 dark:text-gray-300">
{{ t('admin.accounts.openai.testMode') }}
</label>
<Select
v-model="testMode"
:options="openAITestModeOptions"
:disabled="status === 'connecting'"
/>
</div>
<div v-if="supportsImageTest" class="space-y-1.5">
<TextArea
v-model="testPrompt"
@@ -276,6 +287,12 @@ const loadingModels = ref(false)
let abortController: AbortController | null = null
const generatedImages = ref<PreviewImage[]>([])
const previewImageUrl = ref('')
const testMode = ref<'default' | 'compact'>('default')
const isOpenAIAccount = computed(() => props.account?.platform === 'openai')
const openAITestModeOptions = computed(() => [
{ value: 'default', label: t('admin.accounts.openai.testModeDefault') },
{ value: 'compact', label: t('admin.accounts.openai.testModeCompact') }
])
const prioritizedGeminiModels = ['gemini-3.1-flash-image', 'gemini-2.5-flash-image', 'gemini-3.5-flash', 'gemini-2.5-flash', 'gemini-2.5-pro', 'gemini-3-flash-preview', 'gemini-3-pro-preview', 'gemini-2.0-flash']
const supportsGeminiImageTest = computed(() => {
const modelID = selectedModelId.value.toLowerCase()
@@ -309,6 +326,7 @@ watch(
async (newVal) => {
if (newVal && props.account) {
testPrompt.value = ''
testMode.value = 'default'
resetState()
await loadAvailableModels()
} else {
@@ -400,6 +418,18 @@ const startTest = async () => {
abortController = new AbortController()
try {
const requestBody: {
model_id: string
prompt: string
mode?: 'default' | 'compact'
} = {
model_id: selectedModelId.value,
prompt: supportsImageTest.value ? testPrompt.value.trim() : ''
}
if (isOpenAIAccount.value) {
requestBody.mode = testMode.value
}
// Use the configured API base; EventSource does not support POST.
const url = buildApiUrl(`/admin/accounts/${props.account.id}/test`)
@@ -410,10 +440,7 @@ const startTest = async () => {
Authorization: `Bearer ${localStorage.getItem('auth_token')}`,
'Content-Type': 'application/json'
},
body: JSON.stringify({
model_id: selectedModelId.value,
prompt: supportsImageTest.value ? testPrompt.value.trim() : ''
}),
body: JSON.stringify(requestBody),
signal: abortController.signal
})
@@ -505,6 +532,12 @@ const handleEvent = (event: {
}
break
case 'status':
if (event.text) {
addLine(event.text, 'text-cyan-300')
}
break
case 'test_complete':
// Move streaming content to output lines
if (streamingContent.value) {
@@ -19,13 +19,23 @@
<div>
<label class="input-label">{{ t('admin.accounts.dataImportFile') }}</label>
<div
class="flex items-center justify-between gap-3 rounded-lg border border-dashed border-gray-300 bg-gray-50 px-4 py-3 dark:border-dark-600 dark:bg-dark-800"
class="flex items-center justify-between gap-3 rounded-lg border border-dashed px-4 py-3 transition-colors"
:class="dragActive
? 'border-primary-400 bg-primary-50/70 dark:border-primary-500 dark:bg-primary-900/20'
: 'border-gray-300 bg-gray-50 dark:border-dark-600 dark:bg-dark-800'"
@dragenter.prevent="handleDragEnter"
@dragover.prevent
@dragleave.prevent="handleDragLeave"
@drop.prevent="handleDrop"
>
<div class="min-w-0">
<div class="truncate text-sm text-gray-700 dark:text-dark-200">
{{ fileName || t('admin.accounts.dataImportSelectFile') }}
<div class="truncate text-sm text-gray-700 dark:text-dark-200" :title="fileListTitle">
{{ selectedFilesLabel || t('admin.accounts.dataImportSelectFile') }}
</div>
<div class="text-xs text-gray-500 dark:text-dark-400">
JSON (.json)
<span v-if="files.length > 1"> · {{ fileListTitle }}</span>
</div>
<div class="text-xs text-gray-500 dark:text-dark-400">JSON (.json)</div>
</div>
<button type="button" class="btn btn-secondary shrink-0" @click="openFilePicker">
{{ t('common.chooseFile') }}
@@ -36,6 +46,7 @@
type="file"
class="hidden"
accept="application/json,.json"
multiple
@change="handleFileChange"
/>
</div>
@@ -90,7 +101,7 @@ import { useI18n } from 'vue-i18n'
import BaseDialog from '@/components/common/BaseDialog.vue'
import { adminAPI } from '@/api/admin'
import { useAppStore } from '@/stores/app'
import type { AdminDataImportResult } from '@/types'
import type { AdminDataImportResult, AdminDataPayload } from '@/types'
interface Props {
show: boolean
@@ -108,11 +119,19 @@ const { t } = useI18n()
const appStore = useAppStore()
const importing = ref(false)
const file = ref<File | null>(null)
const files = ref<File[]>([])
const dragDepth = ref(0)
const dragActive = computed(() => dragDepth.value > 0)
const hasCreatedData = ref(false)
const result = ref<AdminDataImportResult | null>(null)
const fileInput = ref<HTMLInputElement | null>(null)
const fileName = computed(() => file.value?.name || '')
const selectedFilesLabel = computed(() => {
if (files.value.length === 0) return ''
if (files.value.length === 1) return files.value[0]?.name || ''
return t('admin.accounts.selectedCount', { count: files.value.length })
})
const fileListTitle = computed(() => files.value.map((item) => item.name).join(', '))
const errorItems = computed(() => result.value?.errors || [])
@@ -120,7 +139,9 @@ watch(
() => props.show,
(open) => {
if (open) {
file.value = null
files.value = []
dragDepth.value = 0
hasCreatedData.value = false
result.value = null
if (fileInput.value) {
fileInput.value.value = ''
@@ -135,14 +156,56 @@ const openFilePicker = () => {
const handleFileChange = (event: Event) => {
const target = event.target as HTMLInputElement
file.value = target.files?.[0] || null
setSelectedFiles(target.files)
target.value = ''
}
const handleClose = () => {
if (importing.value) return
if (hasCreatedData.value) {
hasCreatedData.value = false
emit('imported')
}
emit('close')
}
const isJsonFile = (sourceFile: File) => {
const name = sourceFile.name.toLowerCase()
return name.endsWith('.json') || sourceFile.type === 'application/json'
}
const setSelectedFiles = (sourceFiles: FileList | File[] | null | undefined) => {
if (importing.value) return
const incoming = Array.from(sourceFiles || [])
const picked = incoming.filter(isJsonFile)
if (!picked.length) {
appStore.showError(t('admin.accounts.dataImportSelectFile'))
return
}
if (picked.length < incoming.length) {
appStore.showWarning(
t('admin.accounts.dataImportIgnoredFiles', { count: incoming.length - picked.length })
)
}
files.value = picked
result.value = null
}
const handleDragEnter = () => {
if (importing.value) return
dragDepth.value += 1
}
const handleDragLeave = () => {
dragDepth.value = Math.max(0, dragDepth.value - 1)
}
const handleDrop = (event: DragEvent) => {
dragDepth.value = 0
if (importing.value) return
setSelectedFiles(event.dataTransfer?.files)
}
const readFileAsText = async (sourceFile: File): Promise<string> => {
if (typeof sourceFile.text === 'function') {
return sourceFile.text()
@@ -161,16 +224,74 @@ const readFileAsText = async (sourceFile: File): Promise<string> => {
})
}
const SUPPORTED_DATA_TYPES = ['sub2api-data', 'sub2api-bundle']
const SUPPORTED_DATA_VERSION = 1
// 与后端 validateDataHeader 对齐:合并前逐文件校验,避免坏文件混入合并 payload 后
// 报错无法定位来源,或绕过后端本会对单文件做的 type/version 检查。
const isValidDataPayload = (payload: unknown): payload is AdminDataPayload => {
if (!payload || typeof payload !== 'object' || Array.isArray(payload)) return false
const candidate = payload as Record<string, unknown>
if (
candidate.type !== undefined &&
candidate.type !== '' &&
!SUPPORTED_DATA_TYPES.includes(candidate.type as string)
) {
return false
}
if (
candidate.version !== undefined &&
candidate.version !== 0 &&
candidate.version !== SUPPORTED_DATA_VERSION
) {
return false
}
return Array.isArray(candidate.proxies) && Array.isArray(candidate.accounts)
}
const mergeDataPayloads = (payloads: AdminDataPayload[]): AdminDataPayload => {
const [firstPayload] = payloads
if (payloads.length === 1 && firstPayload) return firstPayload
return {
type: payloads.find((item) => typeof item.type === 'string')?.type,
version: payloads.find((item) => typeof item.version === 'number')?.version,
exported_at: new Date().toISOString(),
proxies: payloads.flatMap((item) => item.proxies),
accounts: payloads.flatMap((item) => item.accounts),
skipped_shadows: payloads.reduce((sum, item) => {
const count = Number(item.skipped_shadows || 0)
return Number.isFinite(count) ? sum + count : sum
}, 0)
}
}
const handleImport = async () => {
if (!file.value) {
if (files.value.length === 0) {
appStore.showError(t('admin.accounts.dataImportSelectFile'))
return
}
importing.value = true
try {
const text = await readFileAsText(file.value)
const dataPayload = JSON.parse(text)
const dataPayloads: AdminDataPayload[] = []
for (const sourceFile of files.value) {
let parsed: unknown
try {
parsed = JSON.parse(await readFileAsText(sourceFile))
} catch {
appStore.showError(
t('admin.accounts.dataImportParseFailedFile', { name: sourceFile.name })
)
return
}
if (!isValidDataPayload(parsed)) {
appStore.showError(t('admin.accounts.dataImportInvalidFile', { name: sourceFile.name }))
return
}
dataPayloads.push(parsed)
}
const dataPayload = mergeDataPayloads(dataPayloads)
const res = await adminAPI.accounts.importData({
data: dataPayload,
@@ -187,17 +308,17 @@ const handleImport = async () => {
proxy_failed: res.proxy_failed,
}
if (res.account_failed > 0 || res.proxy_failed > 0) {
// 部分成功也创建了数据;弹窗关闭时通过 imported 通知父组件刷新列表
if (res.account_created > 0 || res.proxy_created > 0) {
hasCreatedData.value = true
}
appStore.showError(t('admin.accounts.dataImportCompletedWithErrors', msgParams))
} else {
appStore.showSuccess(t('admin.accounts.dataImportSuccess', msgParams))
emit('imported')
}
} catch (error: any) {
if (error instanceof SyntaxError) {
appStore.showError(t('admin.accounts.dataImportParseFailed'))
} else {
appStore.showError(error?.message || t('admin.accounts.dataImportFailed'))
}
appStore.showError(error?.message || t('admin.accounts.dataImportFailed'))
} finally {
importing.value = false
}
@@ -182,4 +182,38 @@ describe('AccountTestModal', () => {
prompt: ''
})
})
it('OpenAI Compact 探测会携带 compact 测试模式', async () => {
getAvailableModels.mockResolvedValue([
{ id: 'gpt-5.4', display_name: 'GPT-5.4' }
])
global.fetch = vi.fn().mockResolvedValue(
createStreamResponse([
'data: {"type":"test_complete","success":true}\n'
])
) as any
const wrapper = mountModal({
id: 42,
name: 'OpenAI OAuth',
platform: 'openai',
type: 'oauth',
status: 'active'
})
await wrapper.setProps({ show: true })
await flushPromises()
;(wrapper.vm as any).selectedModelId = 'gpt-5.4'
;(wrapper.vm as any).testMode = 'compact'
await (wrapper.vm as any).startTest()
await flushPromises()
expect(global.fetch).toHaveBeenCalledTimes(1)
const [, request] = (global.fetch as any).mock.calls[0]
expect(JSON.parse(request.body)).toMatchObject({
model_id: 'gpt-5.4',
prompt: '',
mode: 'compact'
})
})
})
+5
View File
@@ -3172,6 +3172,9 @@ export default {
dataImporting: 'Importing...',
dataImportSelectFile: 'Please select a data file',
dataImportParseFailed: 'Failed to parse data file',
dataImportParseFailedFile: 'Failed to parse {name}',
dataImportInvalidFile: '{name} is not a supported data export file',
dataImportIgnoredFiles: 'Ignored {count} non-JSON file(s)',
dataImportFailed: 'Data import failed',
dataImportResult: 'Import Result',
dataImportResultSummary: 'Proxies created {proxy_created}, reused {proxy_reused}, failed {proxy_failed}; Accounts created {account_created}, failed {account_failed}',
@@ -7486,6 +7489,8 @@ export default {
deletePlanConfirm: 'Are you sure you want to delete this plan?',
originalPrice: 'Original Price',
price: 'Price',
subscriptionCnyPayPreview: 'CNY channel charge preview: {amount}',
subscriptionCnyPayPreviewWithFee: '({feeRate}% fee included: {total})',
validityDays: 'Validity (days)',
validityUnit: 'Validity Unit',
sortOrder: 'Sort Order',
+5
View File
@@ -3246,6 +3246,9 @@ export default {
dataImporting: '导入中...',
dataImportSelectFile: '请选择数据文件',
dataImportParseFailed: '数据解析失败',
dataImportParseFailedFile: '文件 {name} 解析失败',
dataImportInvalidFile: '文件 {name} 不是受支持的导出数据文件',
dataImportIgnoredFiles: '已忽略 {count} 个非 JSON 文件',
dataImportFailed: '数据导入失败',
dataImportResult: '导入结果',
dataImportResultSummary: '代理创建 {proxy_created},复用 {proxy_reused},失败 {proxy_failed};账号创建 {account_created},失败 {account_failed}',
@@ -7663,6 +7666,8 @@ export default {
deletePlanConfirm: '确定要删除此套餐吗?',
originalPrice: '原价',
price: '价格',
subscriptionCnyPayPreview: 'CNY 通道实扣预览:{amount}',
subscriptionCnyPayPreviewWithFee: '(含 {feeRate}% 手续费:{total})',
validityDays: '有效期(天)',
validityUnit: '有效期单位',
sortOrder: '排序',
+3 -10
View File
@@ -789,11 +789,7 @@
<!-- 图片生成计费配置 -->
<div
v-if="
createForm.platform === 'antigravity' ||
createForm.platform === 'gemini' ||
createForm.platform === 'openai'
"
v-if="supportsImagePricingPlatform(createForm.platform)"
class="border-t pt-4"
>
<label
@@ -2180,11 +2176,7 @@
<!-- 图片生成计费配置 -->
<div
v-if="
editForm.platform === 'antigravity' ||
editForm.platform === 'gemini' ||
editForm.platform === 'openai'
"
v-if="supportsImagePricingPlatform(editForm.platform)"
class="border-t pt-4"
>
<label
@@ -3320,6 +3312,7 @@ import {
} from "./groupsModelsList";
import { createModelsListCandidatesTracker } from "./groupsModelsListCandidates";
import { normalizeSupportedModelScopesForPlatform } from "./groupsSupportedModelScopes";
import { supportsImagePricingPlatform } from "./groupsImagePricing";
const { t } = useI18n();
const appStore = useAppStore();
@@ -0,0 +1,17 @@
import { describe, expect, it } from "vitest";
import {
imagePricingPlatforms,
supportsImagePricingPlatform,
} from "../groupsImagePricing";
describe("groups image pricing platform support", () => {
it("includes Grok media groups", () => {
expect(supportsImagePricingPlatform("grok")).toBe(true);
expect(imagePricingPlatforms.has("grok")).toBe(true);
});
it("keeps non-media group platforms out of the image pricing controls", () => {
expect(supportsImagePricingPlatform("anthropic")).toBe(false);
});
});
@@ -0,0 +1,9 @@
export const imagePricingPlatforms = new Set([
"antigravity",
"gemini",
"grok",
"openai",
]);
export const supportsImagePricingPlatform = (platform: string): boolean =>
imagePricingPlatforms.has(platform);
@@ -67,7 +67,7 @@
</div>
<!-- Plan Edit Dialog -->
<PlanEditDialog :show="showPlanDialog" :plan="editingPlan" :groups="groups" @close="showPlanDialog = false" @saved="loadPlans" />
<PlanEditDialog :show="showPlanDialog" :plan="editingPlan" :groups="groups" :payment-config="paymentConfig" @close="showPlanDialog = false" @saved="loadPlans" />
<ConfirmDialog :show="showDeletePlanDialog" :title="t('payment.admin.deletePlan')" :message="t('payment.admin.deletePlanConfirm')" :confirm-text="t('common.delete')" danger @confirm="handleDeletePlan" @cancel="showDeletePlanDialog = false" />
</AppLayout>
@@ -78,6 +78,7 @@ import { ref, computed, onMounted } from 'vue'
import { useI18n } from 'vue-i18n'
import { useAppStore } from '@/stores/app'
import { adminPaymentAPI } from '@/api/admin/payment'
import type { AdminPaymentConfig } from '@/api/admin/payment'
import { extractI18nErrorMessage } from '@/utils/apiError'
import adminAPI from '@/api/admin'
import type { SubscriptionPlan } from '@/types/payment'
@@ -97,6 +98,7 @@ const appStore = useAppStore()
// ==================== Groups ====================
const groups = ref<AdminGroup[]>([])
const paymentConfig = ref<AdminPaymentConfig | null>(null)
async function loadGroups() {
try {
@@ -104,6 +106,13 @@ async function loadGroups() {
} catch { /* ignore */ }
}
async function loadPaymentConfig() {
try {
const res = await adminPaymentAPI.getConfig()
paymentConfig.value = res.data
} catch { /* preview only */ }
}
function getGroup(id: number): AdminGroup | undefined {
return groups.value.find(g => g.id === id)
}
@@ -181,6 +190,7 @@ async function handleDeletePlan() {
onMounted(() => {
loadGroups()
loadPaymentConfig()
loadPlans()
})
</script>
@@ -35,7 +35,16 @@
<div><label class="input-label">{{ t('payment.admin.planDescription') }} <span class="text-red-500">*</span></label><textarea v-model="planForm.description" rows="2" class="input" required></textarea></div>
<div class="grid grid-cols-2 gap-4">
<div><label class="input-label">{{ t('payment.admin.price') }} <span class="text-red-500">*</span></label><input v-model.number="planForm.price" type="number" step="0.01" min="0.01" class="input" required /></div>
<div>
<label class="input-label">{{ t('payment.admin.price') }} <span class="text-red-500">*</span></label>
<input v-model.number="planForm.price" type="number" step="0.01" min="0.01" class="input" required />
<p v-if="subscriptionCnyPreview" class="mt-1 text-xs font-medium text-primary-600 dark:text-primary-400">
{{ t('payment.admin.subscriptionCnyPayPreview', { amount: subscriptionCnyPreview.amount }) }}
<span v-if="subscriptionCnyPreview.feeRate > 0">
{{ t('payment.admin.subscriptionCnyPayPreviewWithFee', { feeRate: subscriptionCnyPreview.feeRate, total: subscriptionCnyPreview.total }) }}
</span>
</p>
</div>
<div><label class="input-label">{{ t('payment.admin.originalPrice') }}</label><input v-model.number="planForm.original_price" type="number" step="0.01" min="0" class="input" /></div>
</div>
<div class="grid grid-cols-2 gap-4">
@@ -81,7 +90,9 @@ import { ref, reactive, computed, watch } from 'vue'
import { useI18n } from 'vue-i18n'
import { useAppStore } from '@/stores/app'
import { adminPaymentAPI } from '@/api/admin/payment'
import type { AdminPaymentConfig } from '@/api/admin/payment'
import { extractApiErrorMessage } from '@/utils/apiError'
import { formatPaymentAmount } from '@/components/payment/currency'
import type { SubscriptionPlan } from '@/types/payment'
import type { AdminGroup } from '@/types'
import BaseDialog from '@/components/common/BaseDialog.vue'
@@ -94,6 +105,7 @@ const props = defineProps<{
show: boolean
plan: SubscriptionPlan | null
groups: AdminGroup[]
paymentConfig?: AdminPaymentConfig | null
}>()
const emit = defineEmits<{
@@ -129,6 +141,31 @@ const selectedGroupInfo = computed(() => {
return props.groups.find(g => g.id === planForm.group_id) || null
})
function roundCnyAmount(value: number): number {
return Math.round(value * 100) / 100
}
function ceilCnyAmount(value: number): number {
return Math.ceil(value * 100) / 100
}
const subscriptionCnyPreview = computed(() => {
const price = Number(planForm.price) || 0
const rate = Number(props.paymentConfig?.subscription_usd_to_cny_rate) || 0
if (price <= 0 || rate <= 0) return null
const amount = roundCnyAmount(price * rate)
const feeRate = Number(props.paymentConfig?.recharge_fee_rate) || 0
const fee = feeRate > 0 ? ceilCnyAmount((amount * feeRate) / 100) : 0
const total = feeRate > 0 ? roundCnyAmount(amount + fee) : amount
return {
amount: formatPaymentAmount(amount, 'CNY'),
feeRate,
total: formatPaymentAmount(total, 'CNY'),
}
})
// Reset form when dialog opens
watch(() => props.show, (visible) => {
if (!visible) return
@@ -0,0 +1,77 @@
import { describe, expect, it, vi } from 'vitest'
import { mount } from '@vue/test-utils'
import PlanEditDialog from '../PlanEditDialog.vue'
vi.mock('vue-i18n', () => ({
useI18n: () => ({
t: (key: string, params?: Record<string, unknown>) => {
if (key === 'payment.admin.subscriptionCnyPayPreview') return `preview ${params?.amount}`
if (key === 'payment.admin.subscriptionCnyPayPreviewWithFee') return `fee ${params?.feeRate} ${params?.total}`
return key
},
}),
}))
vi.mock('@/stores/app', () => ({
useAppStore: () => ({
showError: vi.fn(),
showSuccess: vi.fn(),
}),
}))
vi.mock('@/api/admin/payment', () => ({
adminPaymentAPI: {
createPlan: vi.fn(),
updatePlan: vi.fn(),
},
}))
function mountDialog(paymentConfig: Record<string, unknown> | null) {
return mount(PlanEditDialog, {
props: {
show: true,
plan: null,
groups: [],
paymentConfig,
},
global: {
stubs: {
BaseDialog: {
props: ['show'],
template: '<div v-if="show"><slot /><slot name="footer" /></div>',
},
Select: true,
Icon: true,
GroupBadge: true,
},
},
})
}
describe('PlanEditDialog subscription CNY payment preview', () => {
it('shows CNY channel charge using the configured subscription rate and fee', async () => {
const wrapper = mountDialog({
subscription_usd_to_cny_rate: 7.15,
recharge_fee_rate: 2.5,
})
await wrapper.find('input[type="number"]').setValue('9.99')
expect(wrapper.text()).toContain('preview')
expect(wrapper.text()).toContain('¥71.43')
expect(wrapper.text()).toContain('fee 2.5')
expect(wrapper.text()).toContain('¥73.22')
})
it('hides the preview when the subscription rate is not configured', async () => {
const wrapper = mountDialog({
subscription_usd_to_cny_rate: 0,
recharge_fee_rate: 2.5,
})
await wrapper.find('input[type="number"]').setValue('9.99')
expect(wrapper.text()).not.toContain('preview')
expect(wrapper.text()).not.toContain('¥71.43')
})
})