From 66ad405dd6ad3d6f5f286b14a85835843553c2cf Mon Sep 17 00:00:00 2001 From: li Date: Mon, 3 Aug 2026 19:59:58 +0800 Subject: [PATCH 001/104] fix(upstream): set explicit TCP dial timeout on upstream transports MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit http.Transport 未设置 DialContext 时使用零值 net.Dialer(Timeout=0), DNS 解析与 TCP 握手没有任何上限,只能等内核 TCP 重传耗尽(Linux 约 130 秒)。 ResponseHeaderTimeout 只约束连接建立之后等待响应头的阶段,覆盖不到建连。 上游域名被解析到 443 不可达的 IP 时(DNS 污染 / 路由异常),单个账号就要卡满 内核超时;而多账号故障转移是串行的,一次请求会阻塞数分钟且不写中间错误, 客户端表现为"长时间无首字后直接结束"。 - buildUpstreamTransport 显式设置 DialContext(10s)与 TLSHandshakeTimeout(10s) - proxyutil SOCKS5 分支的 forward dialer 由 proxy.Direct(零值 dialer,同样无超时) 换成带 10s 超时的 net.Dialer;该分支会覆盖 Transport.DialContext, 调用方设置的建连超时对它无效,必须在此补上 Fixes #5151 --- backend/internal/pkg/proxyutil/dialer.go | 21 +++++- .../pkg/proxyutil/dialer_timeout_test.go | 58 ++++++++++++++ backend/internal/repository/http_upstream.go | 27 +++++++ .../http_upstream_dial_timeout_test.go | 75 +++++++++++++++++++ 4 files changed, 180 insertions(+), 1 deletion(-) create mode 100644 backend/internal/pkg/proxyutil/dialer_timeout_test.go create mode 100644 backend/internal/repository/http_upstream_dial_timeout_test.go diff --git a/backend/internal/pkg/proxyutil/dialer.go b/backend/internal/pkg/proxyutil/dialer.go index e437cae34..3e135e3eb 100644 --- a/backend/internal/pkg/proxyutil/dialer.go +++ b/backend/internal/pkg/proxyutil/dialer.go @@ -16,10 +16,29 @@ import ( "net/http" "net/url" "strings" + "time" "golang.org/x/net/proxy" ) +const ( + // socks5DialTimeout 限制到 SOCKS5 代理自身的 TCP 建连耗时。 + socks5DialTimeout = 10 * time.Second + // socks5DialKeepAlive 与 Go 默认 keepalive 探测间隔保持一致。 + socks5DialKeepAlive = 30 * time.Second +) + +// socks5ForwardDialer 是 SOCKS5 dialer 的底层拨号器。 +// +// proxy.FromURL 的默认 forward dialer 是 proxy.Direct(零值 net.Dialer,无超时), +// 代理地址不可达时会一直卡到内核 TCP 重传耗尽(Linux 约 130 秒)。SOCKS5 分支会 +// 覆盖 Transport.DialContext,因此调用方在 Transport 上设置的建连超时对这条路径 +// 无效,必须在这里补上。 +var socks5ForwardDialer = &net.Dialer{ + Timeout: socks5DialTimeout, + KeepAlive: socks5DialKeepAlive, +} + // ConfigureTransportProxy 根据代理 URL 配置 Transport // // 支持的协议: @@ -45,7 +64,7 @@ func ConfigureTransportProxy(transport *http.Transport, proxyURL *url.URL) error return nil case "socks5", "socks5h": - dialer, err := proxy.FromURL(proxyURL, proxy.Direct) + dialer, err := proxy.FromURL(proxyURL, socks5ForwardDialer) if err != nil { return fmt.Errorf("create socks5 dialer: %w", err) } diff --git a/backend/internal/pkg/proxyutil/dialer_timeout_test.go b/backend/internal/pkg/proxyutil/dialer_timeout_test.go new file mode 100644 index 000000000..212bf0ada --- /dev/null +++ b/backend/internal/pkg/proxyutil/dialer_timeout_test.go @@ -0,0 +1,58 @@ +package proxyutil + +import ( + "context" + "errors" + "net" + "net/http" + "net/url" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +var errStub = errors.New("stub dial") + +// 回归:SOCKS5 分支覆盖了调用方在 Transport 上设置的 DialContext, +// 底层 forward dialer 必须自带建连超时。proxy.Direct 是零值 net.Dialer, +// 代理不可达时会一直卡到内核 TCP 重传耗尽(Linux 约 130 秒)。 +func TestSOCKS5ForwardDialerHasBoundedTimeout(t *testing.T) { + require.Greater(t, socks5ForwardDialer.Timeout, time.Duration(0)) + require.Equal(t, socks5DialTimeout, socks5ForwardDialer.Timeout) + require.Equal(t, socks5DialKeepAlive, socks5ForwardDialer.KeepAlive) +} + +func TestConfigureTransportProxySOCKS5SetsDialContext(t *testing.T) { + for _, scheme := range []string{"socks5", "socks5h"} { + t.Run(scheme, func(t *testing.T) { + proxyURL, err := url.Parse(scheme + "://127.0.0.1:1080") + require.NoError(t, err) + + transport := &http.Transport{} + require.NoError(t, ConfigureTransportProxy(transport, proxyURL)) + require.NotNil(t, transport.DialContext) + require.Nil(t, transport.Proxy, "SOCKS5 不应设置 Transport.Proxy") + }) + } +} + +// HTTP 代理走 Transport.Proxy,不得覆盖调用方设置的 DialContext。 +func TestConfigureTransportProxyHTTPPreservesDialContext(t *testing.T) { + proxyURL, err := url.Parse("http://127.0.0.1:8080") + require.NoError(t, err) + + called := false + transport := &http.Transport{} + transport.DialContext = func(_ context.Context, _, _ string) (net.Conn, error) { + called = true + return nil, errStub + } + + require.NoError(t, ConfigureTransportProxy(transport, proxyURL)) + require.NotNil(t, transport.Proxy) + require.NotNil(t, transport.DialContext) + + _, _ = transport.DialContext(context.Background(), "tcp", "127.0.0.1:1") + require.True(t, called, "HTTP 代理分支不应替换调用方的 DialContext") +} diff --git a/backend/internal/repository/http_upstream.go b/backend/internal/repository/http_upstream.go index ec42e189f..2adc9a513 100644 --- a/backend/internal/repository/http_upstream.go +++ b/backend/internal/repository/http_upstream.go @@ -54,6 +54,18 @@ const ( // defaultResponseHeaderTimeout: 默认等待响应头超时时间(5分钟) // LLM 请求可能排队较久,需要较长超时 defaultResponseHeaderTimeout = 300 * time.Second + // defaultUpstreamDialTimeout: 默认 TCP/DNS 建连超时(10秒) + // Transport 不设置 DialContext 时会退化为零值 net.Dialer(无超时),建连阶段 + // 只能依赖内核默认 TCP 重传(Linux 约 130 秒)。ResponseHeaderTimeout 只约束 + // 连接建立之后等待响应头的阶段,覆盖不到 DNS 解析与 TCP 握手。 + // 上游域名被解析到 443 不可达的 IP 时(DNS 污染/路由异常),单个账号就要卡满 + // 内核超时;而多账号故障转移是串行的,一次请求会阻塞数分钟且不写中间错误。 + defaultUpstreamDialTimeout = 10 * time.Second + // defaultUpstreamDialKeepAlive: TCP keepalive 探测间隔,与 Go 默认值保持一致 + defaultUpstreamDialKeepAlive = 30 * time.Second + // defaultUpstreamTLSHandshakeTimeout: TLS 握手超时(10秒) + // 与建连超时同量级,避免 TCP 已连通但对端不推进握手时无限等待 + defaultUpstreamTLSHandshakeTimeout = 10 * time.Second // defaultMaxUpstreamClients: 默认最大客户端缓存数量 // 超出后会淘汰最久未使用的客户端 defaultMaxUpstreamClients = 5000 @@ -1245,6 +1257,17 @@ func defaultPoolSettings(cfg *config.Config) poolSettings { } } +// newUpstreamDialer 构建上游 Transport 的 TCP dialer。 +// +// 必须显式提供:http.Transport 的 DialContext 为 nil 时使用零值 net.Dialer, +// 建连没有任何超时上限,只能等内核 TCP 重传耗尽(Linux 约 130 秒)。 +func newUpstreamDialer() *net.Dialer { + return &net.Dialer{ + Timeout: defaultUpstreamDialTimeout, + KeepAlive: defaultUpstreamDialKeepAlive, + } +} + // buildUpstreamTransport 构建上游请求的 Transport // 使用配置文件中的连接池参数,支持生产环境调优 // @@ -1257,6 +1280,8 @@ func defaultPoolSettings(cfg *config.Config) poolSettings { // - error: 代理配置错误 // // Transport 参数说明: +// - DialContext: DNS 解析 + TCP 建连超时(不设置则无上限,退化为内核默认重传) +// - TLSHandshakeTimeout: TLS 握手超时 // - MaxIdleConns: 所有主机的最大空闲连接总数 // - MaxIdleConnsPerHost: 每主机最大空闲连接数(影响连接复用率) // - MaxConnsPerHost: 每主机最大连接数(达到后新请求等待) @@ -1264,6 +1289,8 @@ func defaultPoolSettings(cfg *config.Config) poolSettings { // - ResponseHeaderTimeout: 等待响应头超时(不影响流式传输) func buildUpstreamTransport(settings poolSettings, proxyURL *url.URL, protocolMode string) (*http.Transport, error) { transport := &http.Transport{ + DialContext: newUpstreamDialer().DialContext, + TLSHandshakeTimeout: defaultUpstreamTLSHandshakeTimeout, MaxIdleConns: settings.maxIdleConns, MaxIdleConnsPerHost: settings.maxIdleConnsPerHost, MaxConnsPerHost: settings.maxConnsPerHost, diff --git a/backend/internal/repository/http_upstream_dial_timeout_test.go b/backend/internal/repository/http_upstream_dial_timeout_test.go new file mode 100644 index 000000000..00d6ce1a4 --- /dev/null +++ b/backend/internal/repository/http_upstream_dial_timeout_test.go @@ -0,0 +1,75 @@ +package repository + +import ( + "context" + "net" + "net/url" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +// 回归:上游 Transport 必须显式配置建连超时。 +// +// http.Transport.DialContext 为 nil 时 Go 使用零值 net.Dialer(Timeout=0), +// DNS 解析与 TCP 握手没有任何上限,只能等内核重传耗尽(Linux 约 130 秒)。 +// ResponseHeaderTimeout 只覆盖连接建立之后的阶段,管不到建连。 +// 上游域名被解析到不可达 IP 时,串行的多账号故障转移会把一次请求拖到数分钟。 +func TestBuildUpstreamTransportSetsDialTimeout(t *testing.T) { + settings := defaultPoolSettings(nil) + + transport, err := buildUpstreamTransport(settings, nil, upstreamProtocolModeDefault) + require.NoError(t, err) + require.NotNil(t, transport.DialContext, "DialContext 缺失会退化为无超时的零值 dialer") + require.Equal(t, defaultUpstreamTLSHandshakeTimeout, transport.TLSHandshakeTimeout) +} + +func TestNewUpstreamDialerHasBoundedTimeout(t *testing.T) { + dialer := newUpstreamDialer() + + require.Greater(t, dialer.Timeout, time.Duration(0), "建连超时必须有上限") + require.Equal(t, defaultUpstreamDialTimeout, dialer.Timeout) + require.Equal(t, defaultUpstreamDialKeepAlive, dialer.KeepAlive) +} + +// 建连超时对 HTTP 代理同样生效:Transport.Proxy 走的仍是 DialContext, +// 代理地址不可达时必须快速失败而不是挂满内核超时。 +func TestBuildUpstreamTransportKeepsDialTimeoutWithHTTPProxy(t *testing.T) { + proxyURL, err := url.Parse("http://127.0.0.1:1080") + require.NoError(t, err) + + transport, err := buildUpstreamTransport(defaultPoolSettings(nil), proxyURL, upstreamProtocolModeDefault) + require.NoError(t, err) + require.NotNil(t, transport.Proxy) + require.NotNil(t, transport.DialContext) +} + +// SOCKS5 分支会覆盖 Transport.DialContext,覆盖后仍必须是有超时的拨号器。 +func TestBuildUpstreamTransportKeepsDialContextWithSOCKS5Proxy(t *testing.T) { + proxyURL, err := url.Parse("socks5h://127.0.0.1:1080") + require.NoError(t, err) + + transport, err := buildUpstreamTransport(defaultPoolSettings(nil), proxyURL, upstreamProtocolModeDefault) + require.NoError(t, err) + require.NotNil(t, transport.DialContext) +} + +// Timeout 字段确实被 net.Dialer 用于建连:拨一个已被 close 的本地监听端口, +// 断言 Dialer 走的是自己的超时路径而不是无限等待。 +// (不依赖外网可达性,CI 中确定性执行。) +func TestUpstreamDialerRespectsContextCancellation(t *testing.T) { + listener, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + addr := listener.Addr().String() + require.NoError(t, listener.Close()) + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + conn, err := newUpstreamDialer().DialContext(ctx, "tcp", addr) + if conn != nil { + _ = conn.Close() + } + require.Error(t, err, "已取消的 context 必须立即中止拨号") +} From d3b5703b55ddf5a7e38cee0baff4767865c0cb68 Mon Sep 17 00:00:00 2001 From: wucm667 Date: Tue, 4 Aug 2026 19:14:28 +0800 Subject: [PATCH 002/104] fix(model-plaza): show composite group models --- backend/internal/service/channel_plaza.go | 31 +++++++--- .../internal/service/channel_plaza_test.go | 62 +++++++++++++++++++ .../modelPlaza/PlazaModelPricingTable.vue | 13 +++- .../__tests__/PlazaModelPricingTable.spec.ts | 21 +++++++ 4 files changed, 117 insertions(+), 10 deletions(-) diff --git a/backend/internal/service/channel_plaza.go b/backend/internal/service/channel_plaza.go index a71a6e4f1..1e4dee271 100644 --- a/backend/internal/service/channel_plaza.go +++ b/backend/internal/service/channel_plaza.go @@ -28,7 +28,8 @@ type PlazaModel struct { // PlazaGroup 模型广场中以分组为顶层的条目。 // // 与 AvailableGroupRef 相比多了 Description 与 Models;Models 来自该分组关联渠道的 -// 支持模型(按分组平台隔离,防跨平台泄漏),与「可用渠道」页口径一致。 +// 支持模型(普通分组按分组平台隔离,Composite 分组展开关联渠道已配置的 +// 具体平台),与「可用渠道」页口径一致。 type PlazaGroup struct { ID int64 Name string @@ -99,8 +100,12 @@ func (s *ChannelService) ListPlazaGroups(ctx context.Context) ([]PlazaGroup, err order = append(order, g.ID) } - // modelIdx[groupID][modelName] = index into byGroup[groupID].Models - modelIdx := make(map[int64]map[string]int, len(groups)) + type modelKey struct { + platform string + name string + } + // modelIdx[groupID][platform+modelName] = index into byGroup[groupID].Models + modelIdx := make(map[int64]map[modelKey]int, len(groups)) for i := range channels { ch := &channels[i] if ch.Status != StatusActive { @@ -117,23 +122,28 @@ func (s *ChannelService) ListPlazaGroups(ctx context.Context) ([]PlazaGroup, err } idx := modelIdx[gid] if idx == nil { - idx = make(map[string]int, len(supported)) + idx = make(map[modelKey]int, len(supported)) modelIdx[gid] = idx } for j := range supported { m := supported[j] - if m.Platform != pg.Platform { + if pg.Platform == PlatformComposite { + if !isConcreteRequestPlatform(m.Platform) { + continue + } + } else if m.Platform != pg.Platform { continue } pricing := plazaImageDisplayPricing(m.Pricing, groupEnt[gid]) - if at, seen := idx[m.Name]; seen { + key := modelKey{platform: m.Platform, name: m.Name} + if at, seen := idx[key]; seen { // 先见者胜;仅当已存条目无定价而新条目有定价时升级。 if pg.Models[at].Pricing == nil && pricing != nil { pg.Models[at].Pricing = pricing } continue } - idx[m.Name] = len(pg.Models) + idx[key] = len(pg.Models) pg.Models = append(pg.Models, PlazaModel{ Name: m.Name, Platform: m.Platform, @@ -150,7 +160,12 @@ func (s *ChannelService) ListPlazaGroups(ctx context.Context) ([]PlazaGroup, err if len(pg.Models) == 0 { continue } - sort.SliceStable(pg.Models, func(i, j int) bool { return pg.Models[i].Name < pg.Models[j].Name }) + sort.SliceStable(pg.Models, func(i, j int) bool { + if pg.Models[i].Name != pg.Models[j].Name { + return pg.Models[i].Name < pg.Models[j].Name + } + return pg.Models[i].Platform < pg.Models[j].Platform + }) for j := range pg.Models { pg.Models[j].OfficialPricing = s.lookupOfficialPricing(pg.Models[j].Name, officialMemo) } diff --git a/backend/internal/service/channel_plaza_test.go b/backend/internal/service/channel_plaza_test.go index 82d426ce9..554654149 100644 --- a/backend/internal/service/channel_plaza_test.go +++ b/backend/internal/service/channel_plaza_test.go @@ -107,6 +107,68 @@ func TestListPlazaGroups_PlatformIsolation(t *testing.T) { require.Equal(t, "gpt-5", byName["g-gpt"][0].Name) } +func TestListPlazaGroups_CompositeIncludesConfiguredConcretePlatforms(t *testing.T) { + anthropicPrice := 3e-6 + openAIPrice := 2e-6 + ch := Channel{ + ID: 1, Name: "multi", Status: StatusActive, GroupIDs: []int64{10}, + ModelPricing: []ChannelModelPricing{ + {Platform: PlatformAnthropic, Models: []string{"shared-model"}, InputPrice: &anthropicPrice}, + {Platform: PlatformOpenAI, Models: []string{"shared-model"}, InputPrice: &openAIPrice}, + {Platform: "", Models: []string{"empty-platform"}}, + {Platform: PlatformComposite, Models: []string{"nested-composite"}}, + {Platform: "unknown-platform", Models: []string{"unknown-platform"}}, + }, + } + groups := []Group{{ID: 10, Name: "composite", Platform: PlatformComposite, RateMultiplier: 1}} + + out, err := newPlazaChannelService([]Channel{ch}, groups, nil).ListPlazaGroups(context.Background()) + + require.NoError(t, err) + require.Len(t, out, 1) + require.Len(t, out[0].Models, 2, "only concrete platforms are included and same-named models remain distinct") + require.Equal(t, PlatformAnthropic, out[0].Models[0].Platform) + require.Equal(t, PlatformOpenAI, out[0].Models[1].Platform) + require.InDelta(t, anthropicPrice, *out[0].Models[0].Pricing.InputPrice, 1e-12) + require.InDelta(t, openAIPrice, *out[0].Models[1].Pricing.InputPrice, 1e-12) +} + +func TestListPlazaGroups_CompositeAndOrdinaryGroupsDoNotLeakPlatforms(t *testing.T) { + ch := Channel{ + ID: 1, Name: "multi", Status: StatusActive, GroupIDs: []int64{10, 20}, + ModelPricing: []ChannelModelPricing{ + {Platform: PlatformAnthropic, Models: []string{"claude-sonnet"}, InputPrice: testPtrFloat64(3e-6)}, + {Platform: PlatformOpenAI, Models: []string{"gpt-5"}, InputPrice: testPtrFloat64(2e-6)}, + }, + } + groups := []Group{ + {ID: 10, Name: "anthropic-only", Platform: PlatformAnthropic, RateMultiplier: 1}, + {ID: 20, Name: "composite", Platform: PlatformComposite, RateMultiplier: 1}, + } + + out, err := newPlazaChannelService([]Channel{ch}, groups, nil).ListPlazaGroups(context.Background()) + + require.NoError(t, err) + require.Len(t, out, 2) + byName := map[string]PlazaGroup{} + for _, group := range out { + byName[group.Name] = group + } + require.Len(t, byName["anthropic-only"].Models, 1) + require.Equal(t, []PlazaModel{{ + Name: "claude-sonnet", Platform: PlatformAnthropic, Pricing: byName["anthropic-only"].Models[0].Pricing, + }}, byName["anthropic-only"].Models) + require.Len(t, byName["composite"].Models, 2) + require.Equal(t, []string{"claude-sonnet", "gpt-5"}, []string{ + byName["composite"].Models[0].Name, + byName["composite"].Models[1].Name, + }) + require.Equal(t, []string{PlatformAnthropic, PlatformOpenAI}, []string{ + byName["composite"].Models[0].Platform, + byName["composite"].Models[1].Platform, + }) +} + func TestListPlazaGroups_InactiveChannelSkipped(t *testing.T) { inactive := plazaPricedChannel(1, "off", []int64{10}, "anthropic", "claude-sonnet") inactive.Status = "inactive" diff --git a/frontend/src/components/modelPlaza/PlazaModelPricingTable.vue b/frontend/src/components/modelPlaza/PlazaModelPricingTable.vue index 77f76e883..edc72365f 100644 --- a/frontend/src/components/modelPlaza/PlazaModelPricingTable.vue +++ b/frontend/src/components/modelPlaza/PlazaModelPricingTable.vue @@ -59,13 +59,22 @@
{{ m.name }} + + {{ platformLabel(m.platform) }} + { // 旧 bug:image_output_price × 0.1 = 0.000003 被当按次价 expect(text).not.toContain('$0.000003') }) + + it('Composite 分组中相同模型名按具体平台分别展示徽章', () => { + const anthropic = tokenModel({ name: 'shared-model', platform: 'anthropic' }) + const openai = tokenModel({ name: 'shared-model', platform: 'openai' }) + const wrapper = mount(PlazaModelPricingTable, { + props: { + models: [anthropic, openai], + platform: 'composite', + rateMultiplier: 1 + } + }) + + const rows = wrapper.findAll('tbody tr') + expect(rows).toHaveLength(2) + expect(rows.map((row) => row.find('td').text())).toEqual([ + 'shared-modelAnthropic', + 'shared-modelOpenAI' + ]) + expect(wrapper.text()).toContain('Anthropic') + expect(wrapper.text()).toContain('OpenAI') + }) }) From fc5a1b78d20977a1afe235a36328932732f41eed Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=BF=85=E9=80=9A?= <993751+lenzhang@users.noreply.github.com> Date: Wed, 5 Aug 2026 21:26:37 +0800 Subject: [PATCH 003/104] fix(ws): allow Codex prewarm continuation --- .../openai_ws_forwarder_ingress_test.go | 41 ++++++++++++++++++- .../service/openai_ws_forwarder_payload.go | 11 +++++ 2 files changed, 51 insertions(+), 1 deletion(-) diff --git a/backend/internal/service/openai_ws_forwarder_ingress_test.go b/backend/internal/service/openai_ws_forwarder_ingress_test.go index 7753ea959..a13248e39 100644 --- a/backend/internal/service/openai_ws_forwarder_ingress_test.go +++ b/backend/internal/service/openai_ws_forwarder_ingress_test.go @@ -551,13 +551,22 @@ func TestNormalizeOpenAIWSPayloadWithoutInputAndPreviousResponseID(t *testing.T) t.Parallel() normalized, err := normalizeOpenAIWSPayloadWithoutInputAndPreviousResponseID( - []byte(`{"model":"gpt-5.1","input":[1],"previous_response_id":"resp_x","metadata":{"b":2,"a":1}}`), + []byte(`{"model":"gpt-5.1","input":[1],"previous_response_id":"resp_x","client_metadata":{"request_start_ms":"1"},"stream_options":{"include_usage":true},"generate":false,"metadata":{"b":2,"a":1}}`), ) require.NoError(t, err) require.False(t, gjson.GetBytes(normalized, "input").Exists()) require.False(t, gjson.GetBytes(normalized, "previous_response_id").Exists()) + require.False(t, gjson.GetBytes(normalized, "client_metadata").Exists()) + require.False(t, gjson.GetBytes(normalized, "stream_options").Exists()) + require.False(t, gjson.GetBytes(normalized, "generate").Exists()) require.Equal(t, float64(1), gjson.GetBytes(normalized, "metadata.a").Float()) + normalized, err = normalizeOpenAIWSPayloadWithoutInputAndPreviousResponseID( + []byte(`{"model":"gpt-5.1","generate":true}`), + ) + require.NoError(t, err) + require.True(t, gjson.GetBytes(normalized, "generate").Bool()) + _, err = normalizeOpenAIWSPayloadWithoutInputAndPreviousResponseID(nil) require.Error(t, err) @@ -662,6 +671,36 @@ func TestShouldKeepIngressPreviousResponseID(t *testing.T) { require.Equal(t, "strict_incremental_ok", reason) }) + t.Run("codex_prewarm_to_business_keep", func(t *testing.T) { + prewarmPayload := []byte(`{ + "type":"response.create", + "model":"gpt-5.1", + "store":false, + "generate":false, + "client_metadata":{"x-codex-ws-stream-request-start-ms":"100"}, + "stream_options":{"include_usage":true}, + "input":[{"type":"input_text","text":"hello"}] + }`) + businessPayload := []byte(`{ + "type":"response.create", + "model":"gpt-5.1", + "store":false, + "client_metadata":{"x-codex-ws-stream-request-start-ms":"200"}, + "previous_response_id":"resp_prewarm", + "input":[{"type":"input_text","text":"hello"}] + }`) + + keep, reason, err := shouldKeepIngressPreviousResponseID( + prewarmPayload, + businessPayload, + "resp_prewarm", + false, + ) + require.NoError(t, err) + require.True(t, keep) + require.Equal(t, "strict_incremental_ok", reason) + }) + t.Run("missing_previous_response_id", func(t *testing.T) { payload := []byte(`{"type":"response.create","model":"gpt-5.1","input":[]}`) keep, reason, err := shouldKeepIngressPreviousResponseID(previousPayload, payload, "resp_turn_1", false) diff --git a/backend/internal/service/openai_ws_forwarder_payload.go b/backend/internal/service/openai_ws_forwarder_payload.go index a4418527c..889be207f 100644 --- a/backend/internal/service/openai_ws_forwarder_payload.go +++ b/backend/internal/service/openai_ws_forwarder_payload.go @@ -449,6 +449,17 @@ func normalizeOpenAIWSPayloadWithoutInputAndPreviousResponseID(payload []byte) ( } delete(decoded, "input") delete(decoded, "previous_response_id") + // Codex changes transport-only metadata for every response.create. These fields + // do not alter the context referenced by previous_response_id and are excluded + // from Codex's own websocket reuse comparison. + delete(decoded, "client_metadata") + delete(decoded, "stream_options") + // Official Codex prewarms a connection with generate=false, then omits the + // field on the business request that continues from the prewarm response. + // Only normalize false so a meaningful generate=true change remains visible. + if generate, ok := decoded["generate"].(bool); ok && !generate { + delete(decoded, "generate") + } return json.Marshal(decoded) } From ce149831331ea6ae41f1472c3c7322fb8fd70229 Mon Sep 17 00:00:00 2001 From: wucm667 Date: Wed, 5 Aug 2026 22:35:01 +0800 Subject: [PATCH 004/104] fix(antigravity): support Gemini 3.6 Flash models --- backend/internal/domain/constants.go | 6 ++++++ backend/internal/domain/constants_test.go | 8 ++++++++ backend/internal/pkg/antigravity/claude_types.go | 5 +++++ backend/internal/pkg/antigravity/claude_types_test.go | 5 +++++ backend/internal/service/account.go | 5 +++++ 5 files changed, 29 insertions(+) diff --git a/backend/internal/domain/constants.go b/backend/internal/domain/constants.go index d79b49a4a..148c1cd5d 100644 --- a/backend/internal/domain/constants.go +++ b/backend/internal/domain/constants.go @@ -117,6 +117,12 @@ var DefaultAntigravityModelMapping = map[string]string{ "gemini-3.1-flash-image": "gemini-3.1-flash-image", // Gemini 3.1 image preview 映射 "gemini-3.1-flash-image-preview": "gemini-3.1-flash-image", + // Gemini 3.6 Flash tiered models + "gemini-3.6-flash": "gemini-3.6-flash", + "gemini-3.6-flash-high": "gemini-3.6-flash-high", + "gemini-3.6-flash-low": "gemini-3.6-flash-low", + "gemini-3.6-flash-medium": "gemini-3.6-flash-medium", + "gemini-3.6-flash-tiered": "gemini-3.6-flash-tiered", // Gemini 3 image 兼容映射(向 3.1 image 迁移) "gemini-3-pro-image": "gemini-3.1-flash-image", "gemini-3-pro-image-preview": "gemini-3.1-flash-image", diff --git a/backend/internal/domain/constants_test.go b/backend/internal/domain/constants_test.go index 0fb9054f7..e847e22f3 100644 --- a/backend/internal/domain/constants_test.go +++ b/backend/internal/domain/constants_test.go @@ -65,6 +65,14 @@ func TestDefaultAntigravityModelMapping_Gemini31ProAliases(t *testing.T) { } } +func TestDefaultAntigravityModelMapping_Gemini36FlashModels(t *testing.T) { + for _, model := range []string{"gemini-3.6-flash", "gemini-3.6-flash-high", "gemini-3.6-flash-low", "gemini-3.6-flash-medium", "gemini-3.6-flash-tiered"} { + if got := DefaultAntigravityModelMapping[model]; got != model { + t.Fatalf("expected %s to map to itself, got %q", model, got) + } + } +} + func TestDefaultBedrockModelMapping_ContainsNewClaudeModels(t *testing.T) { t.Parallel() diff --git a/backend/internal/pkg/antigravity/claude_types.go b/backend/internal/pkg/antigravity/claude_types.go index d8732d26e..c0335c415 100644 --- a/backend/internal/pkg/antigravity/claude_types.go +++ b/backend/internal/pkg/antigravity/claude_types.go @@ -178,6 +178,11 @@ var geminiModels = []modelDef{ {ID: "gemini-3.1-pro-high", DisplayName: "Gemini 3.1 Pro High", CreatedAt: "2026-02-19T00:00:00Z", IsReasoning: true}, {ID: "gemini-3.1-flash-image", DisplayName: "Gemini 3.1 Flash Image", CreatedAt: "2026-02-19T00:00:00Z"}, {ID: "gemini-3.1-flash-image-preview", DisplayName: "Gemini 3.1 Flash Image Preview", CreatedAt: "2026-02-19T00:00:00Z"}, + {ID: "gemini-3.6-flash", DisplayName: "Gemini 3.6 Flash", CreatedAt: "2026-07-21T00:00:00Z"}, + {ID: "gemini-3.6-flash-high", DisplayName: "Gemini 3.6 Flash High", CreatedAt: "2026-07-21T00:00:00Z", IsReasoning: true}, + {ID: "gemini-3.6-flash-low", DisplayName: "Gemini 3.6 Flash Low", CreatedAt: "2026-07-21T00:00:00Z", IsReasoning: true}, + {ID: "gemini-3.6-flash-medium", DisplayName: "Gemini 3.6 Flash Medium", CreatedAt: "2026-07-21T00:00:00Z", IsReasoning: true}, + {ID: "gemini-3.6-flash-tiered", DisplayName: "Gemini 3.6 Flash", CreatedAt: "2026-07-21T00:00:00Z", IsReasoning: true}, {ID: "gemini-3-pro-preview", DisplayName: "Gemini 3 Pro Preview", CreatedAt: "2025-06-01T00:00:00Z", IsReasoning: true}, {ID: "gemini-3-pro-image", DisplayName: "Gemini 3 Pro Image", CreatedAt: "2025-06-01T00:00:00Z"}, } diff --git a/backend/internal/pkg/antigravity/claude_types_test.go b/backend/internal/pkg/antigravity/claude_types_test.go index bb45c2f0c..65c8078c3 100644 --- a/backend/internal/pkg/antigravity/claude_types_test.go +++ b/backend/internal/pkg/antigravity/claude_types_test.go @@ -20,6 +20,11 @@ func TestDefaultModels_ContainsNewAndLegacyImageModels(t *testing.T) { "gemini-3.1-flash-image", "gemini-3.1-flash-image-preview", "gemini-3-pro-image", // legacy compatibility + "gemini-3.6-flash", + "gemini-3.6-flash-high", + "gemini-3.6-flash-low", + "gemini-3.6-flash-medium", + "gemini-3.6-flash-tiered", } for _, id := range requiredIDs { diff --git a/backend/internal/service/account.go b/backend/internal/service/account.go index ab54d3c80..ba177627f 100644 --- a/backend/internal/service/account.go +++ b/backend/internal/service/account.go @@ -620,6 +620,11 @@ func (a *Account) resolveModelMapping(rawMapping map[string]any) map[string]stri "gemini-3-flash", "gemini-3.1-pro-high", "gemini-3.1-pro-low", + "gemini-3.6-flash", + "gemini-3.6-flash-high", + "gemini-3.6-flash-low", + "gemini-3.6-flash-medium", + "gemini-3.6-flash-tiered", }) applyAntigravityGemini31ProAliases(result) } From 82ccc187cde44dccb33db45cf404c27ddb99a8d9 Mon Sep 17 00:00:00 2001 From: wucm667 Date: Thu, 6 Aug 2026 02:26:42 +0800 Subject: [PATCH 005/104] fix(ops): preserve custom range in error lists --- frontend/src/views/admin/ops/OpsDashboard.vue | 2 ++ .../ops/components/OpsErrorDetailsModal.vue | 7 +++++-- .../ops/utils/__tests__/opsErrorParams.spec.ts | 16 ++++++++++++++++ .../src/views/admin/ops/utils/opsErrorParams.ts | 11 +++++++++++ 4 files changed, 34 insertions(+), 2 deletions(-) create mode 100644 frontend/src/views/admin/ops/utils/__tests__/opsErrorParams.spec.ts create mode 100644 frontend/src/views/admin/ops/utils/opsErrorParams.ts diff --git a/frontend/src/views/admin/ops/OpsDashboard.vue b/frontend/src/views/admin/ops/OpsDashboard.vue index 8bb42f3b6..cd5c67296 100644 --- a/frontend/src/views/admin/ops/OpsDashboard.vue +++ b/frontend/src/views/admin/ops/OpsDashboard.vue @@ -114,6 +114,8 @@ = { page: page.value, page_size: pageSize.value, - time_range: props.timeRange, view: viewMode.value, sort_by: sortBy.value, sort_order: sortOrder.value } + Object.assign(params, buildOpsErrorTimeParams(props.timeRange, props.customStartTime, props.customEndTime)) const platform = String(props.platform || '').trim() if (platform) params.platform = platform @@ -160,7 +163,7 @@ watch( ) watch( - () => [props.timeRange, props.platform, props.groupId] as const, + () => [props.timeRange, props.customStartTime, props.customEndTime, props.platform, props.groupId] as const, () => { if (!props.show) return page.value = 1 diff --git a/frontend/src/views/admin/ops/utils/__tests__/opsErrorParams.spec.ts b/frontend/src/views/admin/ops/utils/__tests__/opsErrorParams.spec.ts new file mode 100644 index 000000000..d4b8c296c --- /dev/null +++ b/frontend/src/views/admin/ops/utils/__tests__/opsErrorParams.spec.ts @@ -0,0 +1,16 @@ +import { describe, expect, it } from 'vitest' +import { buildOpsErrorTimeParams } from '../opsErrorParams' + +describe('buildOpsErrorTimeParams', () => { + it('uses explicit timestamps for a complete custom range', () => { + expect(buildOpsErrorTimeParams('custom', '2026-08-01T00:00:00Z', '2026-08-02T00:00:00Z')).toEqual({ + start_time: '2026-08-01T00:00:00Z', + end_time: '2026-08-02T00:00:00Z' + }) + }) + + it('preserves predefined ranges and falls back for incomplete custom ranges', () => { + expect(buildOpsErrorTimeParams('24h')).toEqual({ time_range: '24h' }) + expect(buildOpsErrorTimeParams('custom', null, null)).toEqual({ time_range: '1h' }) + }) +}) diff --git a/frontend/src/views/admin/ops/utils/opsErrorParams.ts b/frontend/src/views/admin/ops/utils/opsErrorParams.ts new file mode 100644 index 000000000..d994b229f --- /dev/null +++ b/frontend/src/views/admin/ops/utils/opsErrorParams.ts @@ -0,0 +1,11 @@ +export function buildOpsErrorTimeParams( + timeRange: string, + customStartTime?: string | null, + customEndTime?: string | null +): Record { + if (timeRange === 'custom' && customStartTime && customEndTime) { + return { start_time: customStartTime, end_time: customEndTime } + } + + return { time_range: timeRange === 'custom' ? '1h' : timeRange } +} From 5b7a64334fcbaa5aca5c52cf15bdc462069e1ab1 Mon Sep 17 00:00:00 2001 From: wucm667 Date: Thu, 6 Aug 2026 04:30:10 +0800 Subject: [PATCH 006/104] fix(grok): support task IDs for video owner binding --- backend/internal/service/grok_media.go | 2 +- .../service/openai_gateway_grok_test.go | 45 +++++++++++++++++++ 2 files changed, 46 insertions(+), 1 deletion(-) diff --git a/backend/internal/service/grok_media.go b/backend/internal/service/grok_media.go index 6a1ddf061..70c6ad789 100644 --- a/backend/internal/service/grok_media.go +++ b/backend/internal/service/grok_media.go @@ -816,7 +816,7 @@ func extractGrokMediaVideoRequestID(body []byte) string { if len(body) == 0 || !gjson.ValidBytes(body) { return "" } - for _, path := range []string{"request_id", "id", "data.request_id", "data.id", "video.request_id", "video.id"} { + for _, path := range []string{"request_id", "id", "data.request_id", "data.id", "video.request_id", "video.id", "task_id", "data.task_id", "video.task_id"} { if id := strings.TrimSpace(gjson.GetBytes(body, path).String()); id != "" { return id } diff --git a/backend/internal/service/openai_gateway_grok_test.go b/backend/internal/service/openai_gateway_grok_test.go index d087bd518..19942eae7 100644 --- a/backend/internal/service/openai_gateway_grok_test.go +++ b/backend/internal/service/openai_gateway_grok_test.go @@ -1216,6 +1216,51 @@ func TestForwardGrokMediaVideoGenerationReturnsUsageAndResponseID(t *testing.T) require.Equal(t, 10, result.VideoDurationSeconds) } +func TestForwardGrokMediaVideoGenerationReturnsTaskIDAsResponseID(t *testing.T) { + t.Setenv(xai.EnvAllowUnsafeURLOverrides, "true") + gin.SetMode(gin.TestMode) + + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + body := []byte(`{"model":"grok-imagine-video","prompt":"waves"}`) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/videos/generations", bytes.NewReader(body)) + c.Request.Header.Set("Content-Type", "application/json") + + account := &Account{ + ID: 63, + Name: "grok", + Platform: PlatformGrok, + Type: AccountTypeAPIKey, + Concurrency: 1, + Credentials: map[string]any{ + "api_key": "api-key", + "base_url": "https://xai.test/v1", + }, + } + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"task_id":"video-task-123"}`)), + }} + svc := &OpenAIGatewayService{httpUpstream: upstream} + + result, err := svc.ForwardGrokMedia(context.Background(), c, account, GrokMediaEndpointVideosGenerations, "", body, "application/json") + require.NoError(t, err) + require.Equal(t, "video-task-123", result.ResponseID) +} + +func TestExtractGrokMediaVideoRequestIDPreservesExistingPrecedence(t *testing.T) { + body := []byte(`{ + "request_id":"request-id", + "id":"id", + "task_id":"task-id", + "data":{"request_id":"data-request-id","id":"data-id","task_id":"data-task-id"}, + "video":{"request_id":"video-request-id","id":"video-id","task_id":"video-task-id"} + }`) + + require.Equal(t, "request-id", extractGrokMediaVideoRequestID(body)) +} + func TestForwardGrokMediaVideoGenerationPreservesImageToVideoModel(t *testing.T) { t.Setenv(xai.EnvAllowUnsafeURLOverrides, "true") gin.SetMode(gin.TestMode) From e93f6b9953d55c9e862a8d57821260c840465173 Mon Sep 17 00:00:00 2001 From: wucm667 Date: Thu, 6 Aug 2026 20:31:56 +0800 Subject: [PATCH 007/104] fix: avoid cooling OAuth accounts for HTML count token errors --- .../service/openai_gateway_count_tokens.go | 14 ++++++++++++++ .../service/openai_gateway_count_tokens_test.go | 5 +++++ 2 files changed, 19 insertions(+) diff --git a/backend/internal/service/openai_gateway_count_tokens.go b/backend/internal/service/openai_gateway_count_tokens.go index eccd1c414..e56b33df4 100644 --- a/backend/internal/service/openai_gateway_count_tokens.go +++ b/backend/internal/service/openai_gateway_count_tokens.go @@ -345,12 +345,26 @@ func isOpenAIOAuthInputTokensUnsupported(statusCode int, body []byte) bool { return true } + // OAuth's platform endpoint can be blocked by an upstream proxy before it + // reaches the API and return an HTML 403 page without a structured error. + // Treat that endpoint-level response like the other unsupported cases so + // count_tokens remains a local, non-health-affecting convenience request. + if statusCode == http.StatusForbidden && isHTMLResponse(body) { + return true + } + return strings.Contains(msg, "input_tokens") && (strings.Contains(msg, "not found") || strings.Contains(msg, "not supported") || strings.Contains(msg, "unsupported")) } +func isHTMLResponse(body []byte) bool { + trimmed := strings.TrimSpace(strings.ToLower(string(body))) + return strings.HasPrefix(trimmed, "Forbidden", + }, { name: "404_input_tokens_unsupported", statusCode: http.StatusNotFound, From 02e50cc22d038dabf3c6af92dbb92d1e0321f8d5 Mon Sep 17 00:00:00 2001 From: Puppywang Date: Fri, 7 Aug 2026 03:09:57 +0800 Subject: [PATCH 008/104] fix(security): block OAuth account takeover via pending exchange --- .../handler/auth_oauth_pending_flow.go | 15 ++++ .../handler/auth_oauth_pending_flow_test.go | 86 +++++++++++++++++++ 2 files changed, 101 insertions(+) diff --git a/backend/internal/handler/auth_oauth_pending_flow.go b/backend/internal/handler/auth_oauth_pending_flow.go index 32e669b22..7b1d71dca 100644 --- a/backend/internal/handler/auth_oauth_pending_flow.go +++ b/backend/internal/handler/auth_oauth_pending_flow.go @@ -1998,6 +1998,21 @@ func (h *AuthHandler) ExchangePendingOAuthCompletion(c *gin.Context) { response.Success(c, payload) return } + // ─── 安全修复(账号接管 0day)──────────────────────────────────────────── + // 非终态 session(如 choose_account_action_required)的 TargetUserID 可能来自 + // 攻击者提交的他人邮箱:createPendingOAuthAccount / SendPendingOAuthVerifyCode + // 发现邮箱已存在时会把本 session 指向该邮箱用户,全程无密码、无邮箱验证码、 + // 无账号所有权证明。若此时带着 adoption decision 继续执行,下方的 + // applyPendingOAuthAdoption 会把本 OAuth identity 直接绑定到 TargetUserID, + // 攻击者随后再次 OAuth 登录即被系统识别为受害者本人(完整账号接管)。 + // 只有两类 session 允许在此处执行 adoption/binding: + // 1. canIssueTokenPair == true —— 登录终态,identity 已安全绑定该用户; + // 2. intent == bind_current_user —— 已登录用户主动发起绑定(绑定目标来自登录态 cookie)。 + // 其余状态一律只返回 payload,不绑定、不消费 session。 + if !canIssueTokenPair && !strings.EqualFold(strings.TrimSpace(session.Intent), oauthIntentBindCurrentUser) { + response.Success(c, payload) + return + } if !adoptionDecision.hasDecision() { adoptionRequired, _ := payload["adoption_required"].(bool) if adoptionRequired { diff --git a/backend/internal/handler/auth_oauth_pending_flow_test.go b/backend/internal/handler/auth_oauth_pending_flow_test.go index 4da8db2ab..76c29602e 100644 --- a/backend/internal/handler/auth_oauth_pending_flow_test.go +++ b/backend/internal/handler/auth_oauth_pending_flow_test.go @@ -910,6 +910,92 @@ func TestExchangePendingOAuthCompletionRejectsDisabledTargetUser(t *testing.T) { require.Nil(t, storedSession.ConsumedAt) } +func TestExchangePendingOAuthCompletionChoiceStateDoesNotBindIdentity(t *testing.T) { + // 回归测试:复刻"补邮箱/创建账户"路径的账号接管 0day。 + // 攻击者用自己的 OAuth 账号登录后,在 create-account 步骤提交受害者邮箱, + // 后端发现邮箱已存在会把 pending session 转入 choice 状态并指向受害者 + // (TargetUserID=受害者、无密码/验证码证明)。此时带 adoption decision 调 + // exchange 绝不能把 OAuth identity 绑定到受害者账号。 + handler, client := newOAuthPendingFlowTestHandler(t, false) + ctx := context.Background() + + victim, err := client.User.Create(). + SetEmail("victim@example.com"). + SetUsername("victim-user"). + SetPasswordHash("hash"). + SetRole(service.RoleUser). + SetStatus(service.StatusActive). + Save(ctx) + require.NoError(t, err) + + session, err := client.PendingAuthSession.Create(). + SetSessionToken("choice-state-attack-session-token"). + SetIntent("login"). + SetProviderType("linuxdo"). + SetProviderKey("linuxdo"). + SetProviderSubject("attacker-subject-123"). + SetTargetUserID(victim.ID). + SetResolvedEmail(victim.Email). + SetBrowserSessionKey("choice-state-attack-browser-session-key"). + SetUpstreamIdentityClaims(map[string]any{ + "username": "attacker_linuxdo_user", + "suggested_display_name": "Attacker Display Name", + "suggested_avatar_url": "https://cdn.example/attacker.png", + }). + SetLocalFlowState(map[string]any{ + oauthCompletionResponseKey: map[string]any{ + "step": oauthPendingChoiceStep, + "adoption_required": true, + "force_email_on_signup": true, + "email_binding_required": true, + "existing_account_bindable": true, + "email": victim.Email, + "resolved_email": victim.Email, + "redirect": "/dashboard", + }, + }). + SetExpiresAt(time.Now().UTC().Add(10 * time.Minute)). + Save(ctx) + require.NoError(t, err) + + body := bytes.NewBufferString(`{"adopt_display_name":true,"adopt_avatar":true}`) + recorder := httptest.NewRecorder() + ginCtx, _ := gin.CreateTestContext(recorder) + req := httptest.NewRequest(http.MethodPost, "/api/v1/auth/oauth/pending/exchange", body) + req.Header.Set("Content-Type", "application/json") + req.AddCookie(&http.Cookie{Name: oauthPendingSessionCookieName, Value: encodeCookieValue(session.SessionToken)}) + req.AddCookie(&http.Cookie{Name: oauthPendingBrowserCookieName, Value: encodeCookieValue("choice-state-attack-browser-session-key")}) + ginCtx.Request = req + + handler.ExchangePendingOAuthCompletion(ginCtx) + + require.Equal(t, http.StatusOK, recorder.Code) + data := decodeJSONResponseData(t, recorder) + require.NotContains(t, data, "access_token") + require.Equal(t, oauthPendingChoiceStep, data["step"]) + + // 攻击者的 OAuth identity 绝不能绑定到受害者账号 + identityCount, err := client.AuthIdentity.Query(). + Where( + authidentity.ProviderTypeEQ("linuxdo"), + authidentity.ProviderKeyEQ("linuxdo"), + authidentity.ProviderSubjectEQ("attacker-subject-123"), + ). + Count(ctx) + require.NoError(t, err) + require.Zero(t, identityCount) + + // 受害者资料不得被 adoption 篡改 + storedVictim, err := client.User.Get(ctx, victim.ID) + require.NoError(t, err) + require.Equal(t, "victim-user", storedVictim.Username) + + // session 不得被消费(攻击者无法进入下一环) + storedSession, err := client.PendingAuthSession.Get(ctx, session.ID) + require.NoError(t, err) + require.Nil(t, storedSession.ConsumedAt) +} + func TestNormalizePendingOAuthCompletionResponseScrubsLegacyTokenPayload(t *testing.T) { payload := normalizePendingOAuthCompletionResponse(map[string]any{ "access_token": "legacy-access-token", From 74249b8fed51e79d7c07719a21a0b26b4da94b05 Mon Sep 17 00:00:00 2001 From: IanShaw027 Date: Fri, 7 Aug 2026 12:31:14 +0800 Subject: [PATCH 009/104] =?UTF-8?q?feat(grok):=20=E6=A8=A1=E5=9E=8B?= =?UTF-8?q?=E7=9B=AE=E5=BD=95=E4=B8=8E=E5=8F=AF=E9=85=8D=E7=BD=AE=E6=98=A0?= =?UTF-8?q?=E5=B0=84=EF=BC=8C=E9=BB=98=E8=AE=A4=E7=A6=81=E6=AD=A2=E8=B7=A8?= =?UTF-8?q?=E5=8E=82=E5=95=86=E6=9A=97=E9=BB=98=E6=94=B9=E5=86=99?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 补齐 Grok/Imagine 官方模型与别名,并支持运行时默认文本模型。 默认 model_mapping 仅含 Grok 家族;新增 grok_default_text_model 与 grok_cross_client_model_map_enabled,仅在显式开启时将 gpt/claude/codex 等客户端模型名映射到默认文本模型,避免非 Grok 请求被静默改成 grok-4.5。 --- backend/internal/pkg/xai/models.go | 287 +++++++++++++++++-- backend/internal/pkg/xai/models_test.go | 62 ++++ backend/internal/pkg/xai/oauth_test.go | 15 +- backend/internal/service/domain_constants.go | 9 + backend/internal/service/setting_parse.go | 18 ++ backend/internal/service/setting_update.go | 8 + backend/internal/service/settings_view.go | 8 + 7 files changed, 378 insertions(+), 29 deletions(-) create mode 100644 backend/internal/pkg/xai/models_test.go diff --git a/backend/internal/pkg/xai/models.go b/backend/internal/pkg/xai/models.go index 3a65f32c5..671655d2c 100644 --- a/backend/internal/pkg/xai/models.go +++ b/backend/internal/pkg/xai/models.go @@ -1,28 +1,126 @@ package xai +import ( + "strings" + "sync/atomic" +) + +// runtimeMappingOpts holds operator-configured defaults applied when Grok +// accounts leave credentials.model_mapping empty. Updated from settings. +var runtimeMappingOpts atomic.Value // ModelMappingOptions + +func init() { + runtimeMappingOpts.Store(ModelMappingOptions{}) +} + +// SetRuntimeModelMappingOptions updates process-wide defaults used by +// DefaultModelMapping (e.g. after settings load). Safe for concurrent use. +func SetRuntimeModelMappingOptions(opts ModelMappingOptions) { + runtimeMappingOpts.Store(opts) +} + +// RuntimeModelMappingOptions returns the last options set via SetRuntimeModelMappingOptions. +func RuntimeModelMappingOptions() ModelMappingOptions { + if v := runtimeMappingOpts.Load(); v != nil { + if opts, ok := v.(ModelMappingOptions); ok { + return opts + } + } + return ModelMappingOptions{} +} + // Model describes an xAI model in OpenAI-compatible /models shape. type Model struct { ID string `json:"id"` Object string `json:"object"` + Type string `json:"type,omitempty"` Created int64 `json:"created,omitempty"` OwnedBy string `json:"owned_by"` DisplayName string `json:"display_name,omitempty"` } +// DefaultTextModel is the built-in fallback for empty model fields and Grok +// text aliases (e.g. "grok", "grok-latest"). Operators may override the runtime +// default via settings key grok_default_text_model. +const DefaultTextModel = "grok-4.5" + +// Official Imagine model IDs (https://docs.x.ai/docs/models). +const ( + DefaultImagineImageQualityModel = "grok-imagine-image-quality" + DefaultImagineImageFastModel = "grok-imagine-image" + DefaultImagineVideoModel = "grok-imagine-video" + DefaultImagineVideo15LegacyModel = "grok-imagine-video-1.5" + DefaultImagineVideo15Model = "grok-imagine-video-1.5-preview" +) + +// ModelMappingOptions controls optional expansions of the default mapping. +// Cross-client wildcards (gpt-*/claude-*) are OFF unless explicitly enabled — +// silent rewrite of foreign model names is opt-in for operators who want +// Codex/Claude clients to talk to Grok groups without renaming models. +type ModelMappingOptions struct { + // DefaultText is the target for empty models and optional cross-client maps. + // Empty → DefaultTextModel. + DefaultText string + // EnableCrossClientMap merges gpt-*/codex-*/o*/claude-* → DefaultText. + EnableCrossClientMap bool +} + +func (o ModelMappingOptions) defaultText() string { + if t := strings.TrimSpace(o.DefaultText); t != "" { + return t + } + return DefaultTextModel +} + var defaultModels = []Model{ - {ID: "grok-4.5", Object: "model", OwnedBy: "xai", DisplayName: "Grok 4.5"}, - {ID: "grok-4.3", Object: "model", OwnedBy: "xai", DisplayName: "Grok 4.3"}, - {ID: "grok-build-0.1", Object: "model", OwnedBy: "xai", DisplayName: "Grok Build 0.1"}, - {ID: "grok-composer-2.5-fast", Object: "model", OwnedBy: "xai", DisplayName: "Grok Composer 2.5 Fast"}, - {ID: "grok-4.20-0309-reasoning", Object: "model", OwnedBy: "xai", DisplayName: "Grok 4.20 Reasoning"}, - {ID: "grok-4.20-0309-non-reasoning", Object: "model", OwnedBy: "xai", DisplayName: "Grok 4.20 Non Reasoning"}, - {ID: "grok-4.20-multi-agent-0309", Object: "model", OwnedBy: "xai", DisplayName: "Grok 4.20 Multi Agent"}, - {ID: "grok-imagine", Object: "model", OwnedBy: "xai", DisplayName: "Grok Imagine"}, - {ID: "grok-imagine-image", Object: "model", OwnedBy: "xai", DisplayName: "Grok Imagine Image"}, - {ID: "grok-imagine-image-quality", Object: "model", OwnedBy: "xai", DisplayName: "Grok Imagine Image Quality"}, - {ID: "grok-imagine-edit", Object: "model", OwnedBy: "xai", DisplayName: "Grok Imagine Edit"}, - {ID: "grok-imagine-video", Object: "model", OwnedBy: "xai", DisplayName: "Grok Imagine Video"}, - {ID: "grok-imagine-video-1.5", Object: "model", OwnedBy: "xai", DisplayName: "Grok Imagine Video 1.5"}, + // Text + {ID: "grok-4.5", Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok 4.5"}, + {ID: "grok-4.3", Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok 4.3"}, + {ID: "grok-3-mini", Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok 3 Mini"}, + {ID: "grok-3-mini-fast", Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok 3 Mini Fast"}, + {ID: "grok-build-0.1", Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok Build 0.1"}, + {ID: "grok-code-fast-1-0825", Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok Code Fast"}, + {ID: "grok-composer-2.5-fast", Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok Composer 2.5 Fast"}, + {ID: "grok-4.20-0309-reasoning", Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok 4.20 Reasoning"}, + {ID: "grok-4.20-0309-non-reasoning", Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok 4.20 Non Reasoning"}, + {ID: "grok-4.20-multi-agent-0309", Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok 4.20 Multi Agent"}, + // Imagine + {ID: DefaultImagineImageQualityModel, Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok Imagine Image Quality"}, + {ID: DefaultImagineImageFastModel, Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok Imagine Image"}, + {ID: "grok-imagine-edit", Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok Imagine Edit"}, + {ID: DefaultImagineVideoModel, Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok Imagine Video"}, + {ID: DefaultImagineVideo15Model, Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok Imagine Video 1.5 Preview"}, + {ID: DefaultImagineVideo15LegacyModel, Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok Imagine Video 1.5 Legacy"}, +} + +// grokTextResponsesModelAliases is the source of truth for Grok text models +// accepted by the Responses path: client-facing / undated aliases → canonical +// upstream ID. Used by DefaultModelMapping and IsGrokTextResponsesModelID. +var grokTextResponsesModelAliases = map[string]string{ + "grok": DefaultTextModel, + "grok-latest": DefaultTextModel, + "grok-4.5": DefaultTextModel, + "grok-4.5-latest": DefaultTextModel, + "grok-4.3": "grok-4.3", + "grok-4.3-latest": "grok-4.3", + "grok-3-mini": "grok-3-mini", + "grok-3-mini-fast": "grok-3-mini-fast", + "grok-build": "grok-build-0.1", + "grok-build-latest": "grok-build-0.1", + "grok-build-0.1": "grok-build-0.1", + "grok-composer-2.5-fast": "grok-composer-2.5-fast", + "grok-composer": "grok-composer-2.5-fast", + "composer-2.5": "grok-composer-2.5-fast", + "grok-code-fast": "grok-code-fast-1-0825", + "grok-code-fast-1": "grok-code-fast-1-0825", + "grok-code-fast-1-0825": "grok-code-fast-1-0825", + "grok-4.20-reasoning": "grok-4.20-0309-reasoning", + "grok-4.20-0309-reasoning": "grok-4.20-0309-reasoning", + "grok-4.20-non-reasoning": "grok-4.20-0309-non-reasoning", + "grok-4.20-0309-non-reasoning": "grok-4.20-0309-non-reasoning", + "grok-4.20-multi-agent": "grok-4.20-multi-agent-0309", + "grok-4.20-multi-agent-latest": "grok-4.20-multi-agent-0309", + "grok-4.20-multi-agent-0309": "grok-4.20-multi-agent-0309", } func DefaultModels() []Model { @@ -40,19 +138,162 @@ func DefaultModelIDs() []string { return ids } +// DefaultModelMapping returns native Grok/Imagine identity + aliases, using +// runtime options (default text model / optional cross-client wildcards). +// Does NOT enable gpt-*/claude-* unless SetRuntimeModelMappingOptions enables them. func DefaultModelMapping() map[string]string { - mapping := make(map[string]string, len(defaultModels)+5) + return ModelMappingWithOptions(RuntimeModelMappingOptions()) +} + +// ModelMappingWithOptions builds the default Grok mapping with optional +// cross-client wildcards and a configurable default text model. +func ModelMappingWithOptions(opts ModelMappingOptions) map[string]string { + defaultText := opts.defaultText() + mapping := make(map[string]string, len(defaultModels)+len(grokTextResponsesModelAliases)+48) for _, model := range defaultModels { mapping[model.ID] = model.ID } - mapping["grok"] = "grok-4.5" - mapping["grok-latest"] = "grok-4.5" - mapping["grok-4.5-latest"] = "grok-4.5" - mapping["grok-build"] = "grok-build-0.1" - mapping["grok-build-latest"] = "grok-4.5" - mapping["grok-composer"] = "grok-composer-2.5-fast" - mapping["composer-2.5"] = "grok-composer-2.5-fast" - mapping["grok-4.20-reasoning"] = "grok-4.20-0309-reasoning" - mapping["grok-4.20-non-reasoning"] = "grok-4.20-0309-non-reasoning" + for alias, canonical := range grokTextResponsesModelAliases { + // Remap aliases that pointed at DefaultTextModel constant to runtime default. + if canonical == DefaultTextModel { + mapping[alias] = defaultText + } else { + mapping[alias] = canonical + } + } + // Imagine aliases / legacy IDs → official catalog. + mapping["grok-imagine"] = DefaultImagineImageQualityModel + mapping["grok-imagine-1"] = DefaultImagineImageQualityModel + // edit keeps its own id when listed; alias bare names to quality for clients + // that only send grok-imagine-edit without catalog awareness. + mapping["grok-imagine-edit"] = "grok-imagine-edit" + mapping["grok-imagine-image"] = DefaultImagineImageFastModel + mapping["grok-imagine-image-quality"] = DefaultImagineImageQualityModel + mapping["grok-imagine-video"] = DefaultImagineVideoModel + mapping["grok-imagine-video-1.5"] = DefaultImagineVideo15Model + mapping["grok-imagine-video-1.5-preview"] = DefaultImagineVideo15Model + mapping["grok-video-1.5"] = DefaultImagineVideo15Model + + if opts.EnableCrossClientMap { + // Codex / OpenAI Responses client defaults (wildcard patterns). + mapping["gpt-*"] = defaultText + mapping["codex-*"] = defaultText + mapping["o1*"] = defaultText + mapping["o3*"] = defaultText + mapping["o4*"] = defaultText + // Claude Code defaults when operators intentionally enable bridging. + mapping["claude-*"] = defaultText + } + addGrokProviderPrefixedMappings(mapping) return mapping } + +func addGrokProviderPrefixedMappings(mapping map[string]string) { + snapshot := make(map[string]string, len(mapping)) + for key, value := range mapping { + snapshot[key] = value + } + for key, value := range snapshot { + if !isGrokNativeOrAlias(key) { + continue + } + for _, prefix := range []string{"xai/", "x-ai/", "grok/"} { + mapping[prefix+key] = value + } + } +} + +func isGrokNativeOrAlias(model string) bool { + model = strings.ToLower(strings.TrimSpace(model)) + if model == "" || strings.Contains(model, "*") { + return false + } + return strings.HasPrefix(model, "grok") || + strings.HasPrefix(model, "imagine") || + strings.HasPrefix(model, "composer") +} + +// StripGrokProviderPrefix removes common provider prefixes accepted for +// xAI/Grok models, returning the native model ID. +func StripGrokProviderPrefix(model string) string { + trimmed := strings.TrimSpace(model) + lower := strings.ToLower(trimmed) + for _, prefix := range []string{"xai/", "x-ai/", "grok/"} { + if strings.HasPrefix(lower, prefix) { + return strings.TrimSpace(trimmed[len(prefix):]) + } + } + return trimmed +} + +// IsGrokModelID reports whether model looks like a native Grok/xAI model id +// (including aliases). Claude/OpenAI model names return false. +func IsGrokModelID(model string) bool { + normalized := strings.ToLower(StripGrokProviderPrefix(model)) + if normalized == "" { + return false + } + if strings.HasPrefix(normalized, "grok") { + return true + } + if strings.HasPrefix(normalized, "imagine") { + return true + } + return false +} + +// IsGrokTextResponsesModelID reports whether model is a known Grok text model +// for the Responses API. Imagine image/video and unknown custom ids return false. +func IsGrokTextResponsesModelID(model string) bool { + normalized := strings.ToLower(StripGrokProviderPrefix(model)) + _, ok := grokTextResponsesModelAliases[normalized] + return ok +} + +// ResolveGrokTextResponsesModelID canonicalizes a Grok text alias before upstream. +// empty or bare aliases that resolve via DefaultTextModel use defaultText when set. +func ResolveGrokTextResponsesModelID(model string, defaultText ...string) string { + fallback := DefaultTextModel + if len(defaultText) > 0 && strings.TrimSpace(defaultText[0]) != "" { + fallback = strings.TrimSpace(defaultText[0]) + } + trimmed := strings.TrimSpace(model) + if trimmed == "" { + return fallback + } + normalized := strings.ToLower(StripGrokProviderPrefix(trimmed)) + if canonical, ok := grokTextResponsesModelAliases[normalized]; ok { + if canonical == DefaultTextModel { + return fallback + } + return canonical + } + return StripGrokProviderPrefix(trimmed) +} + +// ResolveDefaultTextModel returns defaultText (or DefaultTextModel) when model is empty. +func ResolveDefaultTextModel(model string, defaultText ...string) string { + if trimmed := strings.TrimSpace(model); trimmed != "" { + return trimmed + } + if len(defaultText) > 0 && strings.TrimSpace(defaultText[0]) != "" { + return strings.TrimSpace(defaultText[0]) + } + return DefaultTextModel +} + +// CanonicalImagineVideoModel normalizes video model ids for pricing tables. +// Legacy "grok-imagine-video-1.5" shares the 1.5 price family with preview. +func CanonicalImagineVideoModel(model string) string { + m := strings.ToLower(StripGrokProviderPrefix(model)) + switch { + case m == "" || m == DefaultImagineVideoModel: + return DefaultImagineVideoModel + case strings.HasPrefix(m, "grok-imagine-video-1.5") || m == "grok-video-1.5": + return DefaultImagineVideo15Model + case strings.HasPrefix(m, "grok-imagine-video"): + return DefaultImagineVideoModel + default: + return m + } +} diff --git a/backend/internal/pkg/xai/models_test.go b/backend/internal/pkg/xai/models_test.go new file mode 100644 index 000000000..82c04a6bb --- /dev/null +++ b/backend/internal/pkg/xai/models_test.go @@ -0,0 +1,62 @@ +package xai + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestDefaultModelMappingExcludesCrossClientWildcards(t *testing.T) { + t.Parallel() + SetRuntimeModelMappingOptions(ModelMappingOptions{}) + mapping := DefaultModelMapping() + + require.Equal(t, "grok-4.5", mapping["grok"]) + require.Equal(t, "grok-4.5", mapping["grok-latest"]) + require.Equal(t, "grok-build-0.1", mapping["grok-build"]) + require.Equal(t, "grok-build-0.1", mapping["grok-build-latest"]) + require.Equal(t, DefaultImagineVideo15Model, mapping["grok-imagine-video-1.5"]) + require.Equal(t, DefaultImagineVideo15Model, mapping["grok-imagine-video-1.5-preview"]) + require.Equal(t, "grok-4.5", mapping["xai/grok"]) + + // Cross-vendor wildcards must stay opt-in. + _, hasGPT := mapping["gpt-*"] + _, hasClaude := mapping["claude-*"] + require.False(t, hasGPT) + require.False(t, hasClaude) +} + +func TestModelMappingWithOptionsCrossClient(t *testing.T) { + t.Parallel() + mapping := ModelMappingWithOptions(ModelMappingOptions{ + DefaultText: "grok-4.3", + EnableCrossClientMap: true, + }) + require.Equal(t, "grok-4.3", mapping["grok"]) + require.Equal(t, "grok-4.3", mapping["gpt-*"]) + require.Equal(t, "grok-4.3", mapping["claude-*"]) + require.Equal(t, "grok-4.3", mapping["codex-*"]) +} + +func TestCanonicalImagineVideoModel(t *testing.T) { + t.Parallel() + require.Equal(t, DefaultImagineVideoModel, CanonicalImagineVideoModel("grok-imagine-video")) + require.Equal(t, DefaultImagineVideo15Model, CanonicalImagineVideoModel("grok-imagine-video-1.5")) + require.Equal(t, DefaultImagineVideo15Model, CanonicalImagineVideoModel("grok-imagine-video-1.5-preview")) + require.Equal(t, DefaultImagineVideo15Model, CanonicalImagineVideoModel("xai/grok-video-1.5")) +} + +func TestIsGrokModelID(t *testing.T) { + t.Parallel() + require.True(t, IsGrokModelID("grok-4.5")) + require.True(t, IsGrokModelID("x-ai/grok-4.3")) + require.False(t, IsGrokModelID("gpt-5")) + require.False(t, IsGrokModelID("claude-sonnet-4")) +} + +func TestResolveGrokTextResponsesModelID(t *testing.T) { + t.Parallel() + require.Equal(t, "grok-4.5", ResolveGrokTextResponsesModelID("")) + require.Equal(t, "grok-4.3", ResolveGrokTextResponsesModelID("grok", "grok-4.3")) + require.Equal(t, "grok-4.20-multi-agent-0309", ResolveGrokTextResponsesModelID("grok-4.20-multi-agent")) +} diff --git a/backend/internal/pkg/xai/oauth_test.go b/backend/internal/pkg/xai/oauth_test.go index 53c898310..6b72c5840 100644 --- a/backend/internal/pkg/xai/oauth_test.go +++ b/backend/internal/pkg/xai/oauth_test.go @@ -340,22 +340,25 @@ func TestRuntimeSanityReportsInvalidOverridesWithoutSecrets(t *testing.T) { func TestDefaultModelMappingIncludesGrokAliases(t *testing.T) { t.Parallel() + SetRuntimeModelMappingOptions(ModelMappingOptions{}) mapping := DefaultModelMapping() require.Equal(t, "grok-4.5", mapping["grok"]) require.Equal(t, "grok-4.5", mapping["grok-latest"]) require.Equal(t, "grok-4.5", mapping["grok-4.5"]) require.Equal(t, "grok-4.5", mapping["grok-4.5-latest"]) require.Equal(t, "grok-build-0.1", mapping["grok-build"]) - require.Equal(t, "grok-4.5", mapping["grok-build-latest"]) + require.Equal(t, "grok-build-0.1", mapping["grok-build-latest"]) require.Equal(t, "grok-composer-2.5-fast", mapping["grok-composer"]) require.Equal(t, "grok-composer-2.5-fast", mapping["composer-2.5"]) require.Equal(t, "grok-4.20-0309-reasoning", mapping["grok-4.20-reasoning"]) require.Equal(t, "grok-4.20-0309-non-reasoning", mapping["grok-4.20-non-reasoning"]) require.Equal(t, "grok-4.20-multi-agent-0309", mapping["grok-4.20-multi-agent-0309"]) - require.Equal(t, "grok-imagine", mapping["grok-imagine"]) - require.Equal(t, "grok-imagine-image", mapping["grok-imagine-image"]) - require.Equal(t, "grok-imagine-image-quality", mapping["grok-imagine-image-quality"]) + require.Equal(t, DefaultImagineImageQualityModel, mapping["grok-imagine"]) + require.Equal(t, DefaultImagineImageFastModel, mapping["grok-imagine-image"]) + require.Equal(t, DefaultImagineImageQualityModel, mapping["grok-imagine-image-quality"]) require.Equal(t, "grok-imagine-edit", mapping["grok-imagine-edit"]) - require.Equal(t, "grok-imagine-video", mapping["grok-imagine-video"]) - require.Equal(t, "grok-imagine-video-1.5", mapping["grok-imagine-video-1.5"]) + require.Equal(t, DefaultImagineVideoModel, mapping["grok-imagine-video"]) + require.Equal(t, DefaultImagineVideo15Model, mapping["grok-imagine-video-1.5"]) + _, hasGPT := mapping["gpt-*"] + require.False(t, hasGPT, "cross-client wildcards must be opt-in") } diff --git a/backend/internal/service/domain_constants.go b/backend/internal/service/domain_constants.go index c0905e11b..f24244ef7 100644 --- a/backend/internal/service/domain_constants.go +++ b/backend/internal/service/domain_constants.go @@ -397,6 +397,15 @@ const ( // pre-filled when creating a new channel monitor from the admin UI. Range: [15, 3600]. SettingKeyChannelMonitorDefaultIntervalSeconds = "channel_monitor_default_interval_seconds" + // SettingKeyGrokDefaultTextModel is the fallback Grok text model for empty + // request models and built-in Grok aliases (e.g. "grok" → this id). Default grok-4.5. + SettingKeyGrokDefaultTextModel = "grok_default_text_model" + + // SettingKeyGrokCrossClientModelMapEnabled, when true, includes gpt-*/codex-*/o*/claude-* + // wildcards in the default Grok account model_mapping so foreign client model names + // can reach Grok groups. Default false (no silent cross-vendor rewrite). + SettingKeyGrokCrossClientModelMapEnabled = "grok_cross_client_model_map_enabled" + // SettingKeyAvailableChannelsEnabled is a DB-backed soft switch for the "Available Channels" // user-facing aggregate view. When false: user endpoint returns an empty list and the // sidebar entry is hidden. Defaults to false (opt-in feature). diff --git a/backend/internal/service/setting_parse.go b/backend/internal/service/setting_parse.go index 05be77086..a96dc44f7 100644 --- a/backend/internal/service/setting_parse.go +++ b/backend/internal/service/setting_parse.go @@ -15,6 +15,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/pkg/antigravity" infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" "github.com/Wei-Shaw/sub2api/internal/pkg/openai" + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" ) // InitializeDefaultSettings 初始化默认设置 @@ -187,6 +188,10 @@ func (s *SettingService) InitializeDefaultSettings(ctx context.Context) error { SettingKeyChannelMonitorEnabled: "true", SettingKeyChannelMonitorDefaultIntervalSeconds: "60", + // Grok: safe defaults — no cross-vendor model rewrite unless operators enable it. + SettingKeyGrokDefaultTextModel: "grok-4.5", + SettingKeyGrokCrossClientModelMapEnabled: "false", + // Available channels feature (default disabled; opt-in) SettingKeyAvailableChannelsEnabled: "false", @@ -785,6 +790,13 @@ func (s *SettingService) parseSettings(settings map[string]string) *SystemSettin settings[SettingKeyChannelMonitorDefaultIntervalSeconds], ) + // Grok default mapping policy + result.GrokDefaultTextModel = strings.TrimSpace(settings[SettingKeyGrokDefaultTextModel]) + if result.GrokDefaultTextModel == "" { + result.GrokDefaultTextModel = "grok-4.5" + } + result.GrokCrossClientModelMapEnabled = settings[SettingKeyGrokCrossClientModelMapEnabled] == "true" + // Available channels feature (default: disabled; strict true) result.AvailableChannelsEnabled = settings[SettingKeyAvailableChannelsEnabled] == "true" @@ -932,6 +944,12 @@ func (s *SettingService) parseSettings(settings map[string]string) *SystemSettin result.AllowUserViewErrorRequests = settings[SettingKeyAllowUserViewErrorRequests] == "true" // default false + // Publish Grok default model_mapping options for accounts with empty mapping. + xai.SetRuntimeModelMappingOptions(xai.ModelMappingOptions{ + DefaultText: result.GrokDefaultTextModel, + EnableCrossClientMap: result.GrokCrossClientModelMapEnabled, + }) + return result } diff --git a/backend/internal/service/setting_update.go b/backend/internal/service/setting_update.go index 825222fa9..bdace5a42 100644 --- a/backend/internal/service/setting_update.go +++ b/backend/internal/service/setting_update.go @@ -415,6 +415,14 @@ func (s *SettingService) buildSystemSettingsUpdates(ctx context.Context, setting updates[SettingKeyChannelMonitorDefaultIntervalSeconds] = strconv.Itoa(v) } + // Grok model mapping policy + if v := strings.TrimSpace(settings.GrokDefaultTextModel); v != "" { + updates[SettingKeyGrokDefaultTextModel] = v + } else { + updates[SettingKeyGrokDefaultTextModel] = "grok-4.5" + } + updates[SettingKeyGrokCrossClientModelMapEnabled] = strconv.FormatBool(settings.GrokCrossClientModelMapEnabled) + // Available channels feature switch updates[SettingKeyAvailableChannelsEnabled] = strconv.FormatBool(settings.AvailableChannelsEnabled) diff --git a/backend/internal/service/settings_view.go b/backend/internal/service/settings_view.go index 30024ebbe..bf262de0b 100644 --- a/backend/internal/service/settings_view.go +++ b/backend/internal/service/settings_view.go @@ -199,6 +199,10 @@ type SystemSettings struct { ChannelMonitorEnabled bool `json:"channel_monitor_enabled"` ChannelMonitorDefaultIntervalSeconds int `json:"channel_monitor_default_interval_seconds"` + // Grok model mapping policy (admin settings; empty mapping falls back to these). + GrokDefaultTextModel string `json:"grok_default_text_model"` + GrokCrossClientModelMapEnabled bool `json:"grok_cross_client_model_map_enabled"` + // Available Channels feature (user-facing aggregate view) AvailableChannelsEnabled bool `json:"available_channels_enabled"` @@ -365,6 +369,10 @@ type PublicSettings struct { ChannelMonitorEnabled bool `json:"channel_monitor_enabled"` ChannelMonitorDefaultIntervalSeconds int `json:"channel_monitor_default_interval_seconds"` + // Grok model mapping policy (admin settings). + GrokDefaultTextModel string `json:"grok_default_text_model"` + GrokCrossClientModelMapEnabled bool `json:"grok_cross_client_model_map_enabled"` + // Available Channels feature (user-facing aggregate view) AvailableChannelsEnabled bool `json:"available_channels_enabled"` From 370bdcf695a05db8008459cf534ffec9a77fefbd Mon Sep 17 00:00:00 2001 From: IanShaw027 Date: Fri, 7 Aug 2026 12:34:42 +0800 Subject: [PATCH 010/104] =?UTF-8?q?feat(grok):=20=E8=A1=A5=E9=BD=90?= =?UTF-8?q?=E5=AF=86=E7=A0=81=E7=99=BB=E5=BD=95=E4=B8=8E=20SSO=20=E6=A0=A1?= =?UTF-8?q?=E9=AA=8C=EF=BC=8C=E7=BB=9F=E4=B8=80=20OAuth=20=E5=87=AD?= =?UTF-8?q?=E8=AF=81=E5=BD=A2=E6=80=81?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 在 main 既有 SSO→Build 批量导入之上增加 sso-token 校验与账号密码授权。 密码仅用于换取 SSO 再转 Build OAuth,明文与 raw SSO 均不落库。 --- .../admin/grok_import_probe_handler_test.go | 5 + .../handler/admin/grok_oauth_handler.go | 43 ++++ .../handler/admin/grok_oauth_handler_test.go | 65 +++++ .../internal/repository/grok_oauth_client.go | 235 ++++++++++++++++++ backend/internal/server/routes/admin.go | 2 + .../internal/service/grok_oauth_service.go | 61 ++++- .../service/grok_oauth_service_test.go | 58 +++++ backend/internal/service/oauth_service.go | 3 + 8 files changed, 471 insertions(+), 1 deletion(-) diff --git a/backend/internal/handler/admin/grok_import_probe_handler_test.go b/backend/internal/handler/admin/grok_import_probe_handler_test.go index c671d2ff1..a8b7d9125 100644 --- a/backend/internal/handler/admin/grok_import_probe_handler_test.go +++ b/backend/internal/handler/admin/grok_import_probe_handler_test.go @@ -4,6 +4,7 @@ package admin import ( "context" + "errors" "net/http" "net/http/httptest" "strings" @@ -62,6 +63,10 @@ func (grokImportOAuthClientStub) RefreshToken(context.Context, string, string, s return &xai.TokenResponse{AccessToken: "access-token", RefreshToken: "refresh-token", ExpiresIn: 3600}, nil } +func (grokImportOAuthClientStub) LoginWithPassword(context.Context, string, string, string) (*service.GrokPasswordLoginResult, error) { + return nil, errors.New("unexpected password login") +} + func (grokImportOAuthClientStub) ConvertSSOToBuild(context.Context, string, string) (*xai.TokenResponse, error) { return &xai.TokenResponse{AccessToken: "access-token", RefreshToken: "refresh-token", ExpiresIn: 3600}, nil } diff --git a/backend/internal/handler/admin/grok_oauth_handler.go b/backend/internal/handler/admin/grok_oauth_handler.go index 1f679b956..05853a944 100644 --- a/backend/internal/handler/admin/grok_oauth_handler.go +++ b/backend/internal/handler/admin/grok_oauth_handler.go @@ -95,6 +95,17 @@ type GrokRefreshTokenRequest struct { ProxyID *int64 `json:"proxy_id"` } +type GrokSSOTokenRequest struct { + SSOToken string `json:"sso_token"` + ProxyID *int64 `json:"proxy_id"` +} + +type GrokPasswordAuthorizeRequest struct { + Email string `json:"email"` + Password string `json:"password"` + ProxyID *int64 `json:"proxy_id"` +} + func (h *GrokOAuthHandler) RefreshToken(c *gin.Context) { var req GrokRefreshTokenRequest if err := c.ShouldBindJSON(&req); err != nil { @@ -125,6 +136,38 @@ func (h *GrokOAuthHandler) RefreshToken(c *gin.Context) { response.Success(c, tokenInfo) } +// ValidateSSOToken converts a Web SSO cookie into Build OAuth tokens. +// Response contains OAuth token info only — never echoes sso_token. +func (h *GrokOAuthHandler) ValidateSSOToken(c *gin.Context) { + var req GrokSSOTokenRequest + if err := c.ShouldBindJSON(&req); err != nil { + response.BadRequest(c, "Invalid request: "+err.Error()) + return + } + tokenInfo, err := h.grokOAuthService.ValidateSSOToken(c.Request.Context(), req.SSOToken, req.ProxyID) + if err != nil { + response.ErrorFrom(c, err) + return + } + response.Success(c, tokenInfo) +} + +// AuthorizePassword exchanges email/password for Build OAuth tokens via SSO conversion. +// Response never includes password or raw sso_token. +func (h *GrokOAuthHandler) AuthorizePassword(c *gin.Context) { + var req GrokPasswordAuthorizeRequest + if err := c.ShouldBindJSON(&req); err != nil { + response.BadRequest(c, "Invalid request: "+err.Error()) + return + } + tokenInfo, err := h.grokOAuthService.AuthorizePassword(c.Request.Context(), req.Email, req.Password, req.ProxyID) + if err != nil { + response.ErrorFrom(c, err) + return + } + response.Success(c, tokenInfo) +} + func (h *GrokOAuthHandler) RefreshAccountToken(c *gin.Context) { accountID, err := strconv.ParseInt(c.Param("id"), 10, 64) if err != nil { diff --git a/backend/internal/handler/admin/grok_oauth_handler_test.go b/backend/internal/handler/admin/grok_oauth_handler_test.go index 621439488..e5a464b63 100644 --- a/backend/internal/handler/admin/grok_oauth_handler_test.go +++ b/backend/internal/handler/admin/grok_oauth_handler_test.go @@ -4,6 +4,7 @@ package admin import ( "context" + "errors" "io" "net/http" "net/http/httptest" @@ -189,6 +190,70 @@ func TestGrokOAuthHandlerRuntimeSanityDoesNotExposeSecrets(t *testing.T) { require.NotContains(t, rec.Body.String(), "client-secret-like-value") } +type grokOAuthHandlerClient struct{} + +func (c *grokOAuthHandlerClient) ExchangeCode(context.Context, string, string, string, string, string) (*xai.TokenResponse, error) { + return nil, errors.New("unexpected exchange") +} + +func (c *grokOAuthHandlerClient) RefreshToken(context.Context, string, string, string) (*xai.TokenResponse, error) { + return &xai.TokenResponse{AccessToken: "access-token", RefreshToken: "refresh-token", ExpiresIn: 3600}, nil +} + +func (c *grokOAuthHandlerClient) LoginWithPassword(_ context.Context, email, _ string, _ string) (*service.GrokPasswordLoginResult, error) { + return &service.GrokPasswordLoginResult{ + Email: email, + SSOToken: "sso-from-password", + }, nil +} + +func (c *grokOAuthHandlerClient) ConvertSSOToBuild(context.Context, string, string) (*xai.TokenResponse, error) { + return &xai.TokenResponse{AccessToken: "access-token", RefreshToken: "refresh-token", ExpiresIn: 3600}, nil +} + +func TestGrokOAuthHandlerValidateSSOTokenReturnsTokenInfo(t *testing.T) { + gin.SetMode(gin.TestMode) + + oauthClient := &grokOAuthHandlerClient{} + oauthService := service.NewGrokOAuthService(nil, oauthClient) + defer oauthService.Stop() + handler := NewGrokOAuthHandler(oauthService, nil, nil, nil) + + router := gin.New() + router.POST("/api/v1/admin/grok/oauth/sso-token", handler.ValidateSSOToken) + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodPost, "/api/v1/admin/grok/oauth/sso-token", strings.NewReader(`{"sso_token":"sso-token"}`)) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(rec, req) + + require.Equal(t, http.StatusOK, rec.Code) + require.Contains(t, rec.Body.String(), `"access_token":"access-token"`) + require.NotContains(t, rec.Body.String(), `"sso_token"`) +} + +func TestGrokOAuthHandlerAuthorizePasswordReturnsTokenInfoWithoutPassword(t *testing.T) { + gin.SetMode(gin.TestMode) + + oauthClient := &grokOAuthHandlerClient{} + oauthService := service.NewGrokOAuthService(nil, oauthClient) + defer oauthService.Stop() + handler := NewGrokOAuthHandler(oauthService, nil, nil, nil) + + router := gin.New() + router.POST("/api/v1/admin/grok/oauth/password", handler.AuthorizePassword) + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodPost, "/api/v1/admin/grok/oauth/password", strings.NewReader(`{"email":"user@example.com","password":"super-secret"}`)) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(rec, req) + + require.Equal(t, http.StatusOK, rec.Code) + require.Contains(t, rec.Body.String(), `"email":"user@example.com"`) + require.Contains(t, rec.Body.String(), `"access_token":"access-token"`) + require.NotContains(t, rec.Body.String(), `"sso_token"`) + require.NotContains(t, rec.Body.String(), "super-secret") + require.NotContains(t, rec.Body.String(), "sso-from-password") +} + func TestGrokSSOImportExpiryUsesTokenExpiryWithoutRefreshToken(t *testing.T) { tokenExpiry := time.Now().Add(6 * time.Hour).Unix() expiresAt, autoPause := grokSSOImportExpiry(nil, nil, &service.GrokTokenInfo{ diff --git a/backend/internal/repository/grok_oauth_client.go b/backend/internal/repository/grok_oauth_client.go index 38f6cfb96..891172242 100644 --- a/backend/internal/repository/grok_oauth_client.go +++ b/backend/internal/repository/grok_oauth_client.go @@ -1,10 +1,15 @@ package repository import ( + "bytes" "context" + "encoding/json" "errors" + "fmt" + "io" "net/http" "net/url" + "os" "strings" "time" @@ -20,6 +25,15 @@ type grokOAuthClient struct { tokenURL string } +const ( + accountsBaseURL = "https://accounts.x.ai" + loginRPCEndpoint = accountsBaseURL + "/api/rpc" + turnstileWebsiteURL = accountsBaseURL + turnstileWebsiteKey = "0x4AAAAAAAhr9JGVDZbrZOo0" + yesCaptchaCreateTask = "https://api.yescaptcha.com/createTask" + yesCaptchaGetResult = "https://api.yescaptcha.com/getTaskResult" +) + func NewGrokOAuthClient() service.GrokOAuthClient { return &grokOAuthClient{tokenURL: xai.EffectiveTokenURL()} } @@ -90,6 +104,31 @@ func (c *grokOAuthClient) RefreshToken(ctx context.Context, refreshToken, proxyU return &tokenResp, nil } +// LoginWithPassword authenticates against accounts.x.ai and returns an ephemeral SSO cookie. +// Password and SSO must never be written to account credentials or logs. +func (c *grokOAuthClient) LoginWithPassword(ctx context.Context, email, password, proxyURL string) (*service.GrokPasswordLoginResult, error) { + turnstileToken, err := solveTurnstile(ctx) + if err != nil { + return nil, err + } + httpClient, err := createGrokHTTPClient(proxyURL, true) + if err != nil { + return nil, infraerrors.Newf(http.StatusBadGateway, "GROK_OAUTH_CLIENT_INIT_FAILED", "create HTTP client: %v", err) + } + cookieSetterURL, err := createGrokPasswordSession(ctx, httpClient, strings.TrimSpace(email), password, turnstileToken) + if err != nil { + return nil, err + } + ssoToken, err := extractGrokSSOToken(ctx, httpClient, cookieSetterURL) + if err != nil { + return nil, err + } + return &service.GrokPasswordLoginResult{ + Email: strings.TrimSpace(email), + SSOToken: ssoToken, + }, nil +} + func (c *grokOAuthClient) ConvertSSOToBuild(ctx context.Context, ssoToken, proxyURL string) (*xai.TokenResponse, error) { client, err := createGrokSSOHTTPClient(proxyURL) if err != nil { @@ -179,3 +218,199 @@ func grokOAuthHasExplicitEntitlementDenial(body string) bool { strings.Contains(lower, "subscription required") || strings.Contains(lower, "no active grok subscription") } + +func createGrokHTTPClient(proxyURL string, noRedirect bool) (*http.Client, error) { + transport := &http.Transport{} + if strings.TrimSpace(proxyURL) != "" { + parsed, err := url.Parse(proxyURL) + if err != nil { + return nil, err + } + transport.Proxy = http.ProxyURL(parsed) + } + client := &http.Client{Timeout: 120 * time.Second, Transport: transport} + if noRedirect { + client.CheckRedirect = func(_ *http.Request, _ []*http.Request) error { + return http.ErrUseLastResponse + } + } + return client, nil +} + +func solveTurnstile(ctx context.Context) (string, error) { + clientKey := strings.TrimSpace(os.Getenv("YESCAPTCHA_CLIENT_KEY")) + if clientKey == "" { + clientKey = strings.TrimSpace(os.Getenv("YESCAPTCHA_API_KEY")) + } + if clientKey == "" { + return "", infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_CAPTCHA_KEY_REQUIRED", "yescaptcha client key is required for Grok password authorization") + } + createBody, err := json.Marshal(map[string]any{ + "clientKey": clientKey, + "task": map[string]any{ + "type": "TurnstileTaskProxyless", + "websiteURL": turnstileWebsiteURL, + "websiteKey": turnstileWebsiteKey, + }, + }) + if err != nil { + return "", infraerrors.Newf(http.StatusInternalServerError, "GROK_OAUTH_CAPTCHA_FAILED", "encode captcha create request failed: %v", err) + } + req, err := http.NewRequestWithContext(ctx, http.MethodPost, yesCaptchaCreateTask, bytes.NewReader(createBody)) + if err != nil { + return "", infraerrors.Newf(http.StatusBadGateway, "GROK_OAUTH_CAPTCHA_FAILED", "build captcha create request failed: %v", err) + } + req.Header.Set("Content-Type", "application/json") + resp, err := http.DefaultClient.Do(req) + if err != nil { + return "", infraerrors.Newf(http.StatusBadGateway, "GROK_OAUTH_CAPTCHA_FAILED", "create captcha task failed: %v", err) + } + defer func() { _ = resp.Body.Close() }() + var createResp struct { + ErrorID int `json:"errorId"` + TaskID string `json:"taskId"` + ErrorDescription string `json:"errorDescription"` + } + if err := json.NewDecoder(resp.Body).Decode(&createResp); err != nil { + return "", infraerrors.Newf(http.StatusBadGateway, "GROK_OAUTH_CAPTCHA_FAILED", "decode captcha create response failed: %v", err) + } + if createResp.ErrorID != 0 || strings.TrimSpace(createResp.TaskID) == "" { + return "", infraerrors.Newf(http.StatusBadGateway, "GROK_OAUTH_CAPTCHA_FAILED", "captcha create failed: %s", createResp.ErrorDescription) + } + deadline := time.Now().Add(90 * time.Second) + for time.Now().Before(deadline) { + select { + case <-ctx.Done(): + return "", ctx.Err() + case <-time.After(5 * time.Second): + } + body, err := json.Marshal(map[string]any{"clientKey": clientKey, "taskId": createResp.TaskID}) + if err != nil { + return "", infraerrors.Newf(http.StatusInternalServerError, "GROK_OAUTH_CAPTCHA_FAILED", "encode captcha poll request failed: %v", err) + } + req, err := http.NewRequestWithContext(ctx, http.MethodPost, yesCaptchaGetResult, bytes.NewReader(body)) + if err != nil { + continue + } + req.Header.Set("Content-Type", "application/json") + resp, err := http.DefaultClient.Do(req) + if err != nil { + continue + } + var pollResp struct { + ErrorID int `json:"errorId"` + Status string `json:"status"` + ErrorDescription string `json:"errorDescription"` + Solution struct { + Token string `json:"token"` + } `json:"solution"` + } + err = json.NewDecoder(resp.Body).Decode(&pollResp) + _ = resp.Body.Close() + if err != nil { + continue + } + if pollResp.ErrorID != 0 { + return "", infraerrors.Newf(http.StatusBadGateway, "GROK_OAUTH_CAPTCHA_FAILED", "captcha poll failed: %s", pollResp.ErrorDescription) + } + if pollResp.Status == "ready" && strings.TrimSpace(pollResp.Solution.Token) != "" { + return pollResp.Solution.Token, nil + } + } + return "", infraerrors.New(http.StatusGatewayTimeout, "GROK_OAUTH_CAPTCHA_TIMEOUT", "captcha solve timed out") +} + +func createGrokPasswordSession(ctx context.Context, client *http.Client, email, password, turnstileToken string) (string, error) { + payload, err := json.Marshal(map[string]any{ + "rpc": "createSession", + "req": map[string]any{ + "createSessionRequest": map[string]any{ + "credentials": map[string]any{ + "case": "emailAndPassword", + "value": map[string]any{ + "email": email, + "clearTextPassword": password, + }, + }, + }, + "turnstileToken": turnstileToken, + }, + }) + if err != nil { + return "", infraerrors.Newf(http.StatusInternalServerError, "GROK_OAUTH_PASSWORD_LOGIN_FAILED", "encode password login request failed: %v", err) + } + req, err := http.NewRequestWithContext(ctx, http.MethodPost, loginRPCEndpoint, bytes.NewReader(payload)) + if err != nil { + return "", infraerrors.Newf(http.StatusBadGateway, "GROK_OAUTH_PASSWORD_LOGIN_FAILED", "build password login request failed: %v", err) + } + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Origin", accountsBaseURL) + req.Header.Set("Referer", accountsBaseURL+"/sign-in?redirect=grok-com&email=true") + req.Header.Set("User-Agent", "Mozilla/5.0") + req.Header.Set("Accept", "*/*") + resp, err := client.Do(req) + if err != nil { + return "", infraerrors.Newf(http.StatusBadGateway, "GROK_OAUTH_PASSWORD_LOGIN_FAILED", "password login request failed: %v", err) + } + defer func() { _ = resp.Body.Close() }() + body, _ := io.ReadAll(resp.Body) + if resp.StatusCode != http.StatusOK { + return "", infraerrors.Newf(http.StatusBadGateway, "GROK_OAUTH_PASSWORD_LOGIN_FAILED", "password login returned status %d: %s", resp.StatusCode, logredact.RedactText(string(body))) + } + var loginResp struct { + CookieSetterURL string `json:"cookieSetterUrl"` + Error string `json:"error"` + } + if err := json.Unmarshal(body, &loginResp); err != nil { + return "", infraerrors.Newf(http.StatusBadGateway, "GROK_OAUTH_PASSWORD_LOGIN_FAILED", "decode password login response failed: %v", err) + } + if strings.TrimSpace(loginResp.Error) != "" { + return "", infraerrors.Newf(http.StatusBadGateway, "GROK_OAUTH_PASSWORD_LOGIN_FAILED", "password login error: %s", logredact.RedactText(loginResp.Error)) + } + if strings.TrimSpace(loginResp.CookieSetterURL) == "" { + return "", infraerrors.New(http.StatusBadGateway, "GROK_OAUTH_PASSWORD_LOGIN_FAILED", "password login did not return cookieSetterUrl") + } + return loginResp.CookieSetterURL, nil +} + +func extractGrokSSOToken(ctx context.Context, client *http.Client, cookieSetterURL string) (string, error) { + safeURL, err := validateGrokCookieSetterURL(cookieSetterURL) + if err != nil { + return "", infraerrors.Newf(http.StatusBadGateway, "GROK_OAUTH_PASSWORD_LOGIN_FAILED", "invalid cookie setter url: %v", err) + } + req, err := http.NewRequestWithContext(ctx, http.MethodGet, safeURL.String(), nil) + if err != nil { + return "", infraerrors.Newf(http.StatusBadGateway, "GROK_OAUTH_PASSWORD_LOGIN_FAILED", "build cookie setter request: %v", err) + } + req.Header.Set("User-Agent", "Mozilla/5.0") + req.Header.Set("Accept", "text/html,application/xhtml+xml,application/xml;q=0.9,*/*;q=0.8") + req.Header.Set("Referer", accountsBaseURL+"/") + resp, err := client.Do(req) + if err != nil { + return "", infraerrors.Newf(http.StatusBadGateway, "GROK_OAUTH_PASSWORD_LOGIN_FAILED", "follow cookie setter url failed: %v", err) + } + defer func() { _ = resp.Body.Close() }() + for _, cookie := range resp.Header.Values("Set-Cookie") { + if token, ok := strings.CutPrefix(cookie, "sso="); ok { + if idx := strings.Index(token, ";"); idx > 0 { + token = token[:idx] + } + return strings.TrimSpace(token), nil + } + } + return "", infraerrors.Newf(http.StatusBadGateway, "GROK_OAUTH_PASSWORD_LOGIN_FAILED", "no sso cookie found in response (status=%d)", resp.StatusCode) +} + +func validateGrokCookieSetterURL(rawURL string) (*url.URL, error) { + parsed, err := url.Parse(strings.TrimSpace(rawURL)) + if err != nil { + return nil, err + } + if parsed.Scheme != "https" || !strings.EqualFold(parsed.Hostname(), "accounts.x.ai") { + return nil, fmt.Errorf("url must use https://accounts.x.ai") + } + if parsed.User != nil || parsed.Port() != "" || parsed.Fragment != "" || parsed.Opaque != "" { + return nil, fmt.Errorf("url contains disallowed authority or fragment components") + } + return parsed, nil +} diff --git a/backend/internal/server/routes/admin.go b/backend/internal/server/routes/admin.go index 161571cd5..ca02023ed 100644 --- a/backend/internal/server/routes/admin.go +++ b/backend/internal/server/routes/admin.go @@ -466,6 +466,8 @@ func registerGrokOAuthRoutes(admin *gin.RouterGroup, h *handler.Handlers) { grok.POST("/oauth/auth-url", h.Admin.GrokOAuth.GenerateAuthURL) grok.POST("/oauth/exchange-code", h.Admin.GrokOAuth.ExchangeCode) grok.POST("/oauth/refresh-token", h.Admin.GrokOAuth.RefreshToken) + grok.POST("/oauth/sso-token", h.Admin.GrokOAuth.ValidateSSOToken) + grok.POST("/oauth/password", h.Admin.GrokOAuth.AuthorizePassword) grok.POST("/oauth/create-from-oauth", h.Admin.GrokOAuth.CreateAccountFromOAuth) grok.POST("/sso-to-oauth", h.Admin.GrokOAuth.CreateAccountsFromSSO) grok.POST("/oauth/reconcile", h.Admin.GrokOAuth.ReconcileOAuthAccounts) diff --git a/backend/internal/service/grok_oauth_service.go b/backend/internal/service/grok_oauth_service.go index 92f393796..ea88a474b 100644 --- a/backend/internal/service/grok_oauth_service.go +++ b/backend/internal/service/grok_oauth_service.go @@ -106,6 +106,13 @@ type GrokTokenInfo struct { EntitlementStatus string `json:"entitlement_status,omitempty"` } +// GrokPasswordLoginResult is an ephemeral password-login outcome. +// SSOToken is never persisted and must only feed ConvertSSOToBuild. +type GrokPasswordLoginResult struct { + Email string `json:"email,omitempty"` + SSOToken string `json:"sso_token"` +} + func (s *GrokOAuthService) ExchangeCode(ctx context.Context, input *GrokExchangeCodeInput) (*GrokTokenInfo, error) { if input == nil { return nil, infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_INVALID_INPUT", "input is required") @@ -176,7 +183,13 @@ func (s *GrokOAuthService) ValidateRefreshToken(ctx context.Context, refreshToke return s.RefreshToken(ctx, refreshToken, proxyURL, xai.EffectiveClientID()) } -func (s *GrokOAuthService) ConvertFromSSO(ctx context.Context, ssoToken string, proxyID *int64) (*GrokTokenInfo, error) { +// ValidateSSOToken converts a Web SSO cookie into Build OAuth tokens. +// The raw sso_token is never stored on GrokTokenInfo or account credentials. +func (s *GrokOAuthService) ValidateSSOToken(ctx context.Context, ssoToken string, proxyID *int64) (*GrokTokenInfo, error) { + ssoToken = strings.TrimSpace(ssoToken) + if ssoToken == "" { + return nil, infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_NO_SSO_TOKEN", "sso_token is required") + } proxyURL, err := s.proxyURL(ctx, proxyID) if err != nil { return nil, err @@ -185,9 +198,55 @@ func (s *GrokOAuthService) ConvertFromSSO(ctx context.Context, ssoToken string, if err != nil { return nil, err } + if err := validateGrokTokenResponse(tokenResp); err != nil { + return nil, err + } return s.tokenInfoFromResponse(tokenResp, xai.DefaultClientID, nil), nil } +// ConvertFromSSO is the batch-import entry point; same semantics as ValidateSSOToken. +func (s *GrokOAuthService) ConvertFromSSO(ctx context.Context, ssoToken string, proxyID *int64) (*GrokTokenInfo, error) { + return s.ValidateSSOToken(ctx, ssoToken, proxyID) +} + +// AuthorizePassword logs in with email/password, converts the resulting SSO cookie +// to Build OAuth, and returns OAuth tokens only. Password and raw SSO are never persisted. +func (s *GrokOAuthService) AuthorizePassword(ctx context.Context, email, password string, proxyID *int64) (*GrokTokenInfo, error) { + email = strings.TrimSpace(email) + if email == "" { + return nil, infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_EMAIL_REQUIRED", "email is required") + } + if strings.TrimSpace(password) == "" { + return nil, infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_PASSWORD_REQUIRED", "password is required") + } + proxyURL, err := s.proxyURL(ctx, proxyID) + if err != nil { + return nil, err + } + loginResult, err := s.oauthClient.LoginWithPassword(ctx, email, password, proxyURL) + if err != nil { + return nil, err + } + if loginResult == nil || strings.TrimSpace(loginResult.SSOToken) == "" { + return nil, infraerrors.New(http.StatusBadGateway, "GROK_OAUTH_PASSWORD_LOGIN_FAILED", "grok password login did not return sso_token") + } + info, err := s.ValidateSSOToken(ctx, loginResult.SSOToken, proxyID) + if err != nil { + return nil, err + } + if strings.TrimSpace(info.Email) == "" { + info.Email = loginResult.Email + } + return info, nil +} + +func validateGrokTokenResponse(tokenResp *xai.TokenResponse) error { + if tokenResp == nil || strings.TrimSpace(tokenResp.AccessToken) == "" { + return infraerrors.New(http.StatusBadGateway, "GROK_OAUTH_INVALID_TOKEN_RESPONSE", "grok oauth token response missing access_token") + } + return nil +} + func (s *GrokOAuthService) RefreshAccountToken(ctx context.Context, account *Account) (*GrokTokenInfo, error) { if account == nil || account.Platform != PlatformGrok { return nil, infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_INVALID_ACCOUNT", "account is not a Grok account") diff --git a/backend/internal/service/grok_oauth_service_test.go b/backend/internal/service/grok_oauth_service_test.go index 54baef03a..79c279088 100644 --- a/backend/internal/service/grok_oauth_service_test.go +++ b/backend/internal/service/grok_oauth_service_test.go @@ -16,6 +16,9 @@ import ( type grokOAuthClientStub struct { refreshResponse *xai.TokenResponse ssoResponse *xai.TokenResponse + loginResult *GrokPasswordLoginResult + loginEmail string + loginPassword string exchangeCalls int } @@ -28,6 +31,12 @@ func (s *grokOAuthClientStub) RefreshToken(context.Context, string, string, stri return s.refreshResponse, nil } +func (s *grokOAuthClientStub) LoginWithPassword(_ context.Context, email, password, _ string) (*GrokPasswordLoginResult, error) { + s.loginEmail = email + s.loginPassword = password + return s.loginResult, nil +} + func (s *grokOAuthClientStub) ConvertSSOToBuild(context.Context, string, string) (*xai.TokenResponse, error) { return s.ssoResponse, nil } @@ -108,6 +117,55 @@ func TestGrokOAuthServiceConvertFromSSOExtractsBuildClaims(t *testing.T) { require.Equal(t, "user@example.com", credentials["email"]) require.Equal(t, "user-sub", credentials["sub"]) require.Equal(t, "team-1", credentials["team_id"]) + require.NotContains(t, credentials, "sso_token") +} + +func TestGrokOAuthServiceValidateSSOTokenReturnsOAuthTokensWithoutPersistingSSO(t *testing.T) { + svc := NewGrokOAuthService(nil, &grokOAuthClientStub{ + ssoResponse: &xai.TokenResponse{ + AccessToken: "access-from-sso", + RefreshToken: "refresh-from-sso", + TokenType: "Bearer", + ExpiresIn: 3600, + }, + }) + defer svc.Stop() + + info, err := svc.ValidateSSOToken(context.Background(), "sso-token", nil) + require.NoError(t, err) + require.Equal(t, "access-from-sso", info.AccessToken) + require.Equal(t, "refresh-from-sso", info.RefreshToken) + + creds := svc.BuildAccountCredentials(info) + require.NotContains(t, creds, "sso_token") + require.NotContains(t, creds, "password") +} + +func TestGrokOAuthServiceAuthorizePasswordUsesLoginThenSSOAuthorize(t *testing.T) { + client := &grokOAuthClientStub{ + loginResult: &GrokPasswordLoginResult{ + Email: "user@example.com", + SSOToken: "password-derived-sso", + }, + ssoResponse: &xai.TokenResponse{ + AccessToken: "access-from-password", + RefreshToken: "refresh-from-password", + ExpiresIn: 3600, + }, + } + svc := NewGrokOAuthService(nil, client) + defer svc.Stop() + + info, err := svc.AuthorizePassword(context.Background(), " user@example.com ", " super-secret ", nil) + require.NoError(t, err) + require.Equal(t, "user@example.com", info.Email) + require.Equal(t, "access-from-password", info.AccessToken) + + creds := svc.BuildAccountCredentials(info) + require.NotContains(t, creds, "password") + require.NotContains(t, creds, "sso_token") + require.Equal(t, "user@example.com", client.loginEmail) + require.Equal(t, " super-secret ", client.loginPassword, "password bytes must be preserved for upstream login") } func makeGrokOAuthJWT(claims map[string]any) string { diff --git a/backend/internal/service/oauth_service.go b/backend/internal/service/oauth_service.go index 0b3888a73..b92d5cc00 100644 --- a/backend/internal/service/oauth_service.go +++ b/backend/internal/service/oauth_service.go @@ -22,6 +22,9 @@ type OpenAIOAuthClient interface { type GrokOAuthClient interface { ExchangeCode(ctx context.Context, code, codeVerifier, redirectURI, proxyURL, clientID string) (*xai.TokenResponse, error) RefreshToken(ctx context.Context, refreshToken, proxyURL, clientID string) (*xai.TokenResponse, error) + // LoginWithPassword exchanges email/password for a short-lived Web SSO cookie. + // Callers must convert via ConvertSSOToBuild and must not persist password or raw SSO. + LoginWithPassword(ctx context.Context, email, password, proxyURL string) (*GrokPasswordLoginResult, error) ConvertSSOToBuild(ctx context.Context, ssoToken, proxyURL string) (*xai.TokenResponse, error) } From eb6c9663e702e6cdc00384b6bd0fa55021e3fb6a Mon Sep 17 00:00:00 2001 From: IanShaw027 Date: Fri, 7 Aug 2026 12:44:08 +0800 Subject: [PATCH 011/104] =?UTF-8?q?feat(grok):=20=E8=A7=86=E9=A2=91?= =?UTF-8?q?=E6=8C=89=E6=A8=A1=E5=9E=8B=E6=97=8F=E9=85=8D=E7=BD=AE=E6=AF=8F?= =?UTF-8?q?=E7=A7=92=E5=8D=95=E4=BB=B7=EF=BC=8C=E5=B9=B6=E8=A1=A5=E9=BD=90?= =?UTF-8?q?=E7=AE=A1=E7=90=86=E7=AB=AF=E6=98=A0=E5=B0=84=E8=AE=BE=E7=BD=AE?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 分组新增 video_model_prices(JSONB);计费优先模型×分辨率覆盖, 其次旧分辨率三列,最后官方按模型族默认价。同步补齐 admin 设置 中的 grok_default_text_model / 跨客户端映射开关,并保持 grok-imagine-video-1.5 请求模型 identity 不被静默改写。 --- backend/docs/GROK_INTEGRATION_PROGRESS.md | 21 +++ .../internal/handler/admin/group_handler.go | 104 ++++++++------- .../internal/handler/admin/setting_handler.go | 3 + .../handler/admin/setting_handler_update.go | 19 +++ backend/internal/handler/dto/mappers.go | 1 + backend/internal/handler/dto/settings.go | 4 + backend/internal/handler/dto/types.go | 2 + backend/internal/pkg/xai/models.go | 5 +- backend/internal/pkg/xai/models_test.go | 2 +- backend/internal/pkg/xai/oauth_test.go | 3 +- backend/internal/repository/group_repo.go | 12 +- .../repository/group_video_model_prices.go | 114 ++++++++++++++++ backend/internal/service/billing_service.go | 9 +- backend/internal/service/group.go | 28 ++++ backend/internal/service/video_billing.go | 126 ++++++++++++++++++ .../internal/service/video_billing_test.go | 38 ++++++ .../217_group_video_model_prices.sql | 8 ++ 17 files changed, 444 insertions(+), 55 deletions(-) create mode 100644 backend/docs/GROK_INTEGRATION_PROGRESS.md create mode 100644 backend/internal/repository/group_video_model_prices.go create mode 100644 backend/internal/service/video_billing.go create mode 100644 backend/internal/service/video_billing_test.go create mode 100644 backend/migrations/217_group_video_model_prices.sql diff --git a/backend/docs/GROK_INTEGRATION_PROGRESS.md b/backend/docs/GROK_INTEGRATION_PROGRESS.md new file mode 100644 index 000000000..aa73c87a4 --- /dev/null +++ b/backend/docs/GROK_INTEGRATION_PROGRESS.md @@ -0,0 +1,21 @@ +# Grok 完整整合进度(单 PR) + +分支:`feat/grok-complete-integration`(基准 `upstream/main`) + +## 已完成阶段 + +1. **模型目录与可配置映射** — 默认禁止 gpt/claude→grok-4.5;设置项 `grok_default_text_model` / `grok_cross_client_model_map_enabled` +2. **密码登录 + SSO 校验** — `POST .../oauth/password`、`.../oauth/sso-token`;不落库密码/raw SSO +3. **视频按模型族定价** — `groups.video_model_prices` JSONB;计费顺序:模型×分辨率 → 旧三列 → 官方默认 + +## 待续阶段 + +- free-tier / cooldown / payment-required 调度 +- media/voice 增量与错误语义 +- 网关 tool_choice / active-delta(默认关)/ web_search +- 前端 CreateAccount SSO/密码入口 +- 详见子代理审计结论(对照 personal-dev vs main) + +## 原则 + +以 main 为底重写,不整文件 pick personal-dev;migration 使用新序号(如 217)。 diff --git a/backend/internal/handler/admin/group_handler.go b/backend/internal/handler/admin/group_handler.go index e99ba7b75..62960fdd8 100644 --- a/backend/internal/handler/admin/group_handler.go +++ b/backend/internal/handler/admin/group_handler.go @@ -106,31 +106,32 @@ type CreateGroupRequest struct { WeeklyLimitUSD optionalLimitField `json:"weekly_limit_usd"` MonthlyLimitUSD optionalLimitField `json:"monthly_limit_usd"` // 图片生成计费配置(antigravity 和 gemini 平台使用,负数表示清除配置) - AllowImageGeneration bool `json:"allow_image_generation"` - AllowBatchImageGeneration bool `json:"allow_batch_image_generation"` - ImageRateIndependent bool `json:"image_rate_independent"` - ImageRateMultiplier *float64 `json:"image_rate_multiplier"` - BatchImageDiscountMultiplier *float64 `json:"batch_image_discount_multiplier"` - BatchImageHoldMultiplier *float64 `json:"batch_image_hold_multiplier"` - VideoRateIndependent bool `json:"video_rate_independent"` - VideoRateMultiplier *float64 `json:"video_rate_multiplier"` - PeakRateEnabled bool `json:"peak_rate_enabled"` - PeakStart string `json:"peak_start"` - PeakEnd string `json:"peak_end"` - PeakRateMultiplier *float64 `json:"peak_rate_multiplier"` - ProfitControlEnabled bool `json:"profit_control_enabled"` - ProfitMinMargin *float64 `json:"profit_min_margin"` - ProfitSafetyBuffer *float64 `json:"profit_safety_buffer"` - ImagePrice1K *float64 `json:"image_price_1k"` - ImagePrice2K *float64 `json:"image_price_2k"` - ImagePrice4K *float64 `json:"image_price_4k"` - VideoPrice480P *float64 `json:"video_price_480p"` - VideoPrice720P *float64 `json:"video_price_720p"` - VideoPrice1080P *float64 `json:"video_price_1080p"` - WebSearchPricePerCall *float64 `json:"web_search_price_per_call"` - ClaudeCodeOnly bool `json:"claude_code_only"` - FallbackGroupID *int64 `json:"fallback_group_id"` - FallbackGroupIDOnInvalidRequest *int64 `json:"fallback_group_id_on_invalid_request"` + AllowImageGeneration bool `json:"allow_image_generation"` + AllowBatchImageGeneration bool `json:"allow_batch_image_generation"` + ImageRateIndependent bool `json:"image_rate_independent"` + ImageRateMultiplier *float64 `json:"image_rate_multiplier"` + BatchImageDiscountMultiplier *float64 `json:"batch_image_discount_multiplier"` + BatchImageHoldMultiplier *float64 `json:"batch_image_hold_multiplier"` + VideoRateIndependent bool `json:"video_rate_independent"` + VideoRateMultiplier *float64 `json:"video_rate_multiplier"` + PeakRateEnabled bool `json:"peak_rate_enabled"` + PeakStart string `json:"peak_start"` + PeakEnd string `json:"peak_end"` + PeakRateMultiplier *float64 `json:"peak_rate_multiplier"` + ProfitControlEnabled bool `json:"profit_control_enabled"` + ProfitMinMargin *float64 `json:"profit_min_margin"` + ProfitSafetyBuffer *float64 `json:"profit_safety_buffer"` + ImagePrice1K *float64 `json:"image_price_1k"` + ImagePrice2K *float64 `json:"image_price_2k"` + ImagePrice4K *float64 `json:"image_price_4k"` + VideoPrice480P *float64 `json:"video_price_480p"` + VideoPrice720P *float64 `json:"video_price_720p"` + VideoPrice1080P *float64 `json:"video_price_1080p"` + VideoModelPrices map[string]map[string]float64 `json:"video_model_prices,omitempty"` + WebSearchPricePerCall *float64 `json:"web_search_price_per_call"` + ClaudeCodeOnly bool `json:"claude_code_only"` + FallbackGroupID *int64 `json:"fallback_group_id"` + FallbackGroupIDOnInvalidRequest *int64 `json:"fallback_group_id_on_invalid_request"` // 模型路由配置(仅 anthropic 平台使用) ModelRouting map[string][]int64 `json:"model_routing"` ModelRoutingEnabled bool `json:"model_routing_enabled"` @@ -168,31 +169,32 @@ type UpdateGroupRequest struct { WeeklyLimitUSD optionalLimitField `json:"weekly_limit_usd"` MonthlyLimitUSD optionalLimitField `json:"monthly_limit_usd"` // 图片生成计费配置(antigravity 和 gemini 平台使用,负数表示清除配置) - AllowImageGeneration *bool `json:"allow_image_generation"` - AllowBatchImageGeneration *bool `json:"allow_batch_image_generation"` - ImageRateIndependent *bool `json:"image_rate_independent"` - ImageRateMultiplier *float64 `json:"image_rate_multiplier"` - BatchImageDiscountMultiplier *float64 `json:"batch_image_discount_multiplier"` - BatchImageHoldMultiplier *float64 `json:"batch_image_hold_multiplier"` - VideoRateIndependent *bool `json:"video_rate_independent"` - VideoRateMultiplier *float64 `json:"video_rate_multiplier"` - PeakRateEnabled *bool `json:"peak_rate_enabled"` - PeakStart *string `json:"peak_start"` - PeakEnd *string `json:"peak_end"` - PeakRateMultiplier *float64 `json:"peak_rate_multiplier"` - ProfitControlEnabled *bool `json:"profit_control_enabled"` - ProfitMinMargin *float64 `json:"profit_min_margin"` - ProfitSafetyBuffer *float64 `json:"profit_safety_buffer"` - ImagePrice1K *float64 `json:"image_price_1k"` - ImagePrice2K *float64 `json:"image_price_2k"` - ImagePrice4K *float64 `json:"image_price_4k"` - VideoPrice480P *float64 `json:"video_price_480p"` - VideoPrice720P *float64 `json:"video_price_720p"` - VideoPrice1080P *float64 `json:"video_price_1080p"` - WebSearchPricePerCall *float64 `json:"web_search_price_per_call"` - ClaudeCodeOnly *bool `json:"claude_code_only"` - FallbackGroupID *int64 `json:"fallback_group_id"` - FallbackGroupIDOnInvalidRequest *int64 `json:"fallback_group_id_on_invalid_request"` + AllowImageGeneration *bool `json:"allow_image_generation"` + AllowBatchImageGeneration *bool `json:"allow_batch_image_generation"` + ImageRateIndependent *bool `json:"image_rate_independent"` + ImageRateMultiplier *float64 `json:"image_rate_multiplier"` + BatchImageDiscountMultiplier *float64 `json:"batch_image_discount_multiplier"` + BatchImageHoldMultiplier *float64 `json:"batch_image_hold_multiplier"` + VideoRateIndependent *bool `json:"video_rate_independent"` + VideoRateMultiplier *float64 `json:"video_rate_multiplier"` + PeakRateEnabled *bool `json:"peak_rate_enabled"` + PeakStart *string `json:"peak_start"` + PeakEnd *string `json:"peak_end"` + PeakRateMultiplier *float64 `json:"peak_rate_multiplier"` + ProfitControlEnabled *bool `json:"profit_control_enabled"` + ProfitMinMargin *float64 `json:"profit_min_margin"` + ProfitSafetyBuffer *float64 `json:"profit_safety_buffer"` + ImagePrice1K *float64 `json:"image_price_1k"` + ImagePrice2K *float64 `json:"image_price_2k"` + ImagePrice4K *float64 `json:"image_price_4k"` + VideoPrice480P *float64 `json:"video_price_480p"` + VideoPrice720P *float64 `json:"video_price_720p"` + VideoPrice1080P *float64 `json:"video_price_1080p"` + VideoModelPrices map[string]map[string]float64 `json:"video_model_prices,omitempty"` + WebSearchPricePerCall *float64 `json:"web_search_price_per_call"` + ClaudeCodeOnly *bool `json:"claude_code_only"` + FallbackGroupID *int64 `json:"fallback_group_id"` + FallbackGroupIDOnInvalidRequest *int64 `json:"fallback_group_id_on_invalid_request"` // 模型路由配置(仅 anthropic 平台使用) ModelRouting map[string][]int64 `json:"model_routing"` ModelRoutingEnabled *bool `json:"model_routing_enabled"` @@ -519,6 +521,7 @@ func (h *GroupHandler) Create(c *gin.Context) { VideoPrice480P: req.VideoPrice480P, VideoPrice720P: req.VideoPrice720P, VideoPrice1080P: req.VideoPrice1080P, + VideoModelPrices: req.VideoModelPrices, WebSearchPricePerCall: req.WebSearchPricePerCall, ClaudeCodeOnly: req.ClaudeCodeOnly, FallbackGroupID: req.FallbackGroupID, @@ -641,6 +644,7 @@ func (h *GroupHandler) Update(c *gin.Context) { VideoPrice480P: req.VideoPrice480P, VideoPrice720P: req.VideoPrice720P, VideoPrice1080P: req.VideoPrice1080P, + VideoModelPrices: req.VideoModelPrices, WebSearchPricePerCall: req.WebSearchPricePerCall, ClaudeCodeOnly: req.ClaudeCodeOnly, FallbackGroupID: req.FallbackGroupID, diff --git a/backend/internal/handler/admin/setting_handler.go b/backend/internal/handler/admin/setting_handler.go index a9d8b0d88..7ff2f87cc 100644 --- a/backend/internal/handler/admin/setting_handler.go +++ b/backend/internal/handler/admin/setting_handler.go @@ -372,6 +372,9 @@ func (h *SettingHandler) GetSettings(c *gin.Context) { ChannelMonitorEnabled: settings.ChannelMonitorEnabled, ChannelMonitorDefaultIntervalSeconds: settings.ChannelMonitorDefaultIntervalSeconds, + GrokDefaultTextModel: settings.GrokDefaultTextModel, + GrokCrossClientModelMapEnabled: settings.GrokCrossClientModelMapEnabled, + AvailableChannelsEnabled: settings.AvailableChannelsEnabled, ModelPlazaEnabled: settings.ModelPlazaEnabled, diff --git a/backend/internal/handler/admin/setting_handler_update.go b/backend/internal/handler/admin/setting_handler_update.go index 5619641c5..34c084d82 100644 --- a/backend/internal/handler/admin/setting_handler_update.go +++ b/backend/internal/handler/admin/setting_handler_update.go @@ -330,6 +330,10 @@ type UpdateSettingsRequest struct { ChannelMonitorEnabled *bool `json:"channel_monitor_enabled"` ChannelMonitorDefaultIntervalSeconds *int `json:"channel_monitor_default_interval_seconds"` + // Grok model mapping policy + GrokDefaultTextModel *string `json:"grok_default_text_model"` + GrokCrossClientModelMapEnabled *bool `json:"grok_cross_client_model_map_enabled"` + // Available Channels feature switch (user-facing) AvailableChannelsEnabled *bool `json:"available_channels_enabled"` @@ -1860,6 +1864,18 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) { } return previousSettings.ChannelMonitorDefaultIntervalSeconds }(), + GrokDefaultTextModel: func() string { + if req.GrokDefaultTextModel != nil { + return *req.GrokDefaultTextModel + } + return previousSettings.GrokDefaultTextModel + }(), + GrokCrossClientModelMapEnabled: func() bool { + if req.GrokCrossClientModelMapEnabled != nil { + return *req.GrokCrossClientModelMapEnabled + } + return previousSettings.GrokCrossClientModelMapEnabled + }(), AvailableChannelsEnabled: func() bool { if req.AvailableChannelsEnabled != nil { return *req.AvailableChannelsEnabled @@ -2293,6 +2309,9 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) { ChannelMonitorEnabled: updatedSettings.ChannelMonitorEnabled, ChannelMonitorDefaultIntervalSeconds: updatedSettings.ChannelMonitorDefaultIntervalSeconds, + GrokDefaultTextModel: updatedSettings.GrokDefaultTextModel, + GrokCrossClientModelMapEnabled: updatedSettings.GrokCrossClientModelMapEnabled, + AvailableChannelsEnabled: updatedSettings.AvailableChannelsEnabled, ModelPlazaEnabled: updatedSettings.ModelPlazaEnabled, diff --git a/backend/internal/handler/dto/mappers.go b/backend/internal/handler/dto/mappers.go index c6f6cf25e..0ac92284a 100644 --- a/backend/internal/handler/dto/mappers.go +++ b/backend/internal/handler/dto/mappers.go @@ -202,6 +202,7 @@ func groupFromServiceBase(g *service.Group) Group { VideoPrice480P: g.VideoPrice480P, VideoPrice720P: g.VideoPrice720P, VideoPrice1080P: g.VideoPrice1080P, + VideoModelPrices: g.VideoModelPrices, WebSearchPricePerCall: g.WebSearchPricePerCall, ClaudeCodeOnly: g.ClaudeCodeOnly, FallbackGroupID: g.FallbackGroupID, diff --git a/backend/internal/handler/dto/settings.go b/backend/internal/handler/dto/settings.go index 6eccb92c8..75e7972a4 100644 --- a/backend/internal/handler/dto/settings.go +++ b/backend/internal/handler/dto/settings.go @@ -303,6 +303,10 @@ type SystemSettings struct { ChannelMonitorEnabled bool `json:"channel_monitor_enabled"` ChannelMonitorDefaultIntervalSeconds int `json:"channel_monitor_default_interval_seconds"` + // Grok model mapping policy (admin settings; empty account mapping falls back to these). + GrokDefaultTextModel string `json:"grok_default_text_model"` + GrokCrossClientModelMapEnabled bool `json:"grok_cross_client_model_map_enabled"` + // Available Channels feature switch (user-facing aggregate view) AvailableChannelsEnabled bool `json:"available_channels_enabled"` diff --git a/backend/internal/handler/dto/types.go b/backend/internal/handler/dto/types.go index 37927c2b4..871f0ab13 100644 --- a/backend/internal/handler/dto/types.go +++ b/backend/internal/handler/dto/types.go @@ -121,6 +121,8 @@ type Group struct { VideoPrice480P *float64 `json:"video_price_480p"` VideoPrice720P *float64 `json:"video_price_720p"` VideoPrice1080P *float64 `json:"video_price_1080p"` + // VideoModelPrices 可选按模型族×分辨率覆盖视频每秒单价 (USD/s)。 + VideoModelPrices map[string]map[string]float64 `json:"video_model_prices,omitempty"` // Codex alpha/search 网页搜索单次价格(USD/次);null 表示使用默认价 0.01 WebSearchPricePerCall *float64 `json:"web_search_price_per_call"` diff --git a/backend/internal/pkg/xai/models.go b/backend/internal/pkg/xai/models.go index 671655d2c..25584f084 100644 --- a/backend/internal/pkg/xai/models.go +++ b/backend/internal/pkg/xai/models.go @@ -169,9 +169,12 @@ func ModelMappingWithOptions(opts ModelMappingOptions) map[string]string { mapping["grok-imagine-edit"] = "grok-imagine-edit" mapping["grok-imagine-image"] = DefaultImagineImageFastModel mapping["grok-imagine-image-quality"] = DefaultImagineImageQualityModel + // Keep official IDs as identity so client-requested model strings are not + // rewritten on the wire (pricing still canonicalizes 1.5* via CanonicalImagineVideoModel). mapping["grok-imagine-video"] = DefaultImagineVideoModel - mapping["grok-imagine-video-1.5"] = DefaultImagineVideo15Model + mapping["grok-imagine-video-1.5"] = DefaultImagineVideo15LegacyModel mapping["grok-imagine-video-1.5-preview"] = DefaultImagineVideo15Model + // Informal alias only: mapping["grok-video-1.5"] = DefaultImagineVideo15Model if opts.EnableCrossClientMap { diff --git a/backend/internal/pkg/xai/models_test.go b/backend/internal/pkg/xai/models_test.go index 82c04a6bb..6ed61f3e7 100644 --- a/backend/internal/pkg/xai/models_test.go +++ b/backend/internal/pkg/xai/models_test.go @@ -15,7 +15,7 @@ func TestDefaultModelMappingExcludesCrossClientWildcards(t *testing.T) { require.Equal(t, "grok-4.5", mapping["grok-latest"]) require.Equal(t, "grok-build-0.1", mapping["grok-build"]) require.Equal(t, "grok-build-0.1", mapping["grok-build-latest"]) - require.Equal(t, DefaultImagineVideo15Model, mapping["grok-imagine-video-1.5"]) + require.Equal(t, DefaultImagineVideo15LegacyModel, mapping["grok-imagine-video-1.5"]) require.Equal(t, DefaultImagineVideo15Model, mapping["grok-imagine-video-1.5-preview"]) require.Equal(t, "grok-4.5", mapping["xai/grok"]) diff --git a/backend/internal/pkg/xai/oauth_test.go b/backend/internal/pkg/xai/oauth_test.go index 6b72c5840..7c179a5ec 100644 --- a/backend/internal/pkg/xai/oauth_test.go +++ b/backend/internal/pkg/xai/oauth_test.go @@ -358,7 +358,8 @@ func TestDefaultModelMappingIncludesGrokAliases(t *testing.T) { require.Equal(t, DefaultImagineImageQualityModel, mapping["grok-imagine-image-quality"]) require.Equal(t, "grok-imagine-edit", mapping["grok-imagine-edit"]) require.Equal(t, DefaultImagineVideoModel, mapping["grok-imagine-video"]) - require.Equal(t, DefaultImagineVideo15Model, mapping["grok-imagine-video-1.5"]) + require.Equal(t, DefaultImagineVideo15LegacyModel, mapping["grok-imagine-video-1.5"]) + require.Equal(t, DefaultImagineVideo15Model, mapping["grok-imagine-video-1.5-preview"]) _, hasGPT := mapping["gpt-*"] require.False(t, hasGPT, "cross-client wildcards must be opt-in") } diff --git a/backend/internal/repository/group_repo.go b/backend/internal/repository/group_repo.go index 31724f7c5..be008f491 100644 --- a/backend/internal/repository/group_repo.go +++ b/backend/internal/repository/group_repo.go @@ -46,6 +46,9 @@ func (r *groupRepository) Create(ctx context.Context, groupIn *service.Group) er if err := createGroupRecord(ctx, r.client, groupIn); err != nil { return err } + if saveErr := saveGroupVideoModelPrices(ctx, r.sql, groupIn.ID, groupIn.VideoModelPrices); saveErr != nil { + return fmt.Errorf("save group video_model_prices: %w", saveErr) + } if err := enqueueSchedulerOutbox(ctx, r.sql, service.SchedulerOutboxEventGroupChanged, nil, &groupIn.ID, nil); err != nil { logger.LegacyPrintf("repository.group", "[SchedulerOutbox] enqueue group create failed: group=%d err=%v", groupIn.ID, err) } @@ -225,7 +228,11 @@ func (r *groupRepository) GetByIDLite(ctx context.Context, id int64) (*service.G if err != nil { return nil, translatePersistenceError(err, service.ErrGroupNotFound, nil) } - return groupEntityToService(m), nil + out := groupEntityToService(m) + if prices, loadErr := loadGroupVideoModelPrices(ctx, r.sql, []int64{id}); loadErr == nil { + applyVideoModelPricesToGroup(out, prices) + } + return out, nil } func (r *groupRepository) Update(ctx context.Context, groupIn *service.Group) error { @@ -356,6 +363,9 @@ func (r *groupRepository) Update(ctx context.Context, groupIn *service.Group) er return translatePersistenceError(err, service.ErrGroupNotFound, service.ErrGroupExists) } groupIn.UpdatedAt = updated.UpdatedAt + if err := saveGroupVideoModelPrices(ctx, r.sql, groupIn.ID, groupIn.VideoModelPrices); err != nil { + return fmt.Errorf("save group video_model_prices: %w", err) + } if err := enqueueSchedulerOutbox(ctx, r.sql, service.SchedulerOutboxEventGroupChanged, nil, &groupIn.ID, nil); err != nil { logger.LegacyPrintf("repository.group", "[SchedulerOutbox] enqueue group update failed: group=%d err=%v", groupIn.ID, err) } diff --git a/backend/internal/repository/group_video_model_prices.go b/backend/internal/repository/group_video_model_prices.go new file mode 100644 index 000000000..af6b05ff8 --- /dev/null +++ b/backend/internal/repository/group_video_model_prices.go @@ -0,0 +1,114 @@ +package repository + +import ( + "context" + "database/sql" + "encoding/json" + + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/lib/pq" +) + +func loadGroupVideoModelPrices(ctx context.Context, sqlq sqlExecutor, groupIDs []int64) (map[int64]map[string]map[string]float64, error) { + out := make(map[int64]map[string]map[string]float64, len(groupIDs)) + if sqlq == nil || len(groupIDs) == 0 { + return out, nil + } + + rows, err := sqlq.QueryContext(ctx, ` + SELECT id, video_model_prices + FROM groups + WHERE id = ANY($1) AND deleted_at IS NULL + `, pq.Array(groupIDs)) + if err != nil { + return nil, err + } + defer func() { _ = rows.Close() }() + + for rows.Next() { + var ( + groupID int64 + raw []byte + ) + if err := rows.Scan(&groupID, &raw); err != nil { + return nil, err + } + prices, err := decodeVideoModelPrices(raw) + if err != nil { + return nil, err + } + if prices != nil { + out[groupID] = prices + } + } + if err := rows.Err(); err != nil { + return nil, err + } + return out, nil +} + +func saveGroupVideoModelPrices(ctx context.Context, sqlq sqlExecutor, groupID int64, prices map[string]map[string]float64) error { + if sqlq == nil || groupID <= 0 { + return nil + } + normalized := service.NormalizeVideoModelPrices(prices) + if len(normalized) == 0 { + _, err := sqlq.ExecContext(ctx, ` + UPDATE groups + SET video_model_prices = NULL + WHERE id = $1 + `, groupID) + return err + } + payload, err := json.Marshal(normalized) + if err != nil { + return err + } + _, err = sqlq.ExecContext(ctx, ` + UPDATE groups + SET video_model_prices = $1::jsonb + WHERE id = $2 + `, string(payload), groupID) + return err +} + +func applyVideoModelPricesToGroups(groups []service.Group, pricesByID map[int64]map[string]map[string]float64) { + for i := range groups { + if prices, ok := pricesByID[groups[i].ID]; ok { + groups[i].VideoModelPrices = service.NormalizeVideoModelPrices(prices) + continue + } + groups[i].VideoModelPrices = service.NormalizeVideoModelPrices(groups[i].VideoModelPrices) + } +} + +func applyVideoModelPricesToGroup(group *service.Group, pricesByID map[int64]map[string]map[string]float64) { + if group == nil { + return + } + if prices, ok := pricesByID[group.ID]; ok { + group.VideoModelPrices = service.NormalizeVideoModelPrices(prices) + return + } + group.VideoModelPrices = service.NormalizeVideoModelPrices(group.VideoModelPrices) +} + +func decodeVideoModelPrices(raw []byte) (map[string]map[string]float64, error) { + if len(raw) == 0 { + return nil, nil + } + // Driver may return NULL as nil slice; treat empty JSON as nil. + trimmed := string(raw) + if trimmed == "" || trimmed == "null" { + return nil, nil + } + var parsed map[string]map[string]float64 + if err := json.Unmarshal(raw, &parsed); err != nil { + // Some drivers surface NULL via sql.NullString paths; tolerate empty object. + if err == sql.ErrNoRows { + return nil, nil + } + return nil, err + } + return service.NormalizeVideoModelPrices(parsed), nil +} diff --git a/backend/internal/service/billing_service.go b/backend/internal/service/billing_service.go index a43f26a8a..cf9b25fbd 100644 --- a/backend/internal/service/billing_service.go +++ b/backend/internal/service/billing_service.go @@ -1415,6 +1415,9 @@ type VideoPriceConfig struct { Price480P *float64 // 480p 每秒价格(nil 表示使用默认值) Price720P *float64 // 720p 每秒价格(nil 表示使用默认值) Price1080P *float64 // 1080p 每秒价格(nil 表示使用默认值) + // ModelPrices is optional per-model-family override: family → resolution → USD/s. + // When set for a model, it wins over Price* flat columns for that model only. + ModelPrices map[string]map[string]float64 } const ( @@ -1546,8 +1549,12 @@ func (s *BillingService) getImageUnitPrice(model string, imageSize string, group } func (s *BillingService) getVideoUnitPrice(model string, resolution string, groupConfig *VideoPriceConfig) float64 { + // Order: (a) per-model map (b) flat group video_price_* (c) model-aware code defaults. if groupConfig != nil { - switch resolution { + if price := LookupVideoModelPrice(groupConfig.ModelPrices, model, resolution); price != nil { + return *price + } + switch NormalizeVideoBillingResolutionOrDefault(resolution) { case VideoBillingResolution480P: if groupConfig.Price480P != nil { return *groupConfig.Price480P diff --git a/backend/internal/service/group.go b/backend/internal/service/group.go index 2fea013af..7ae43c3f3 100644 --- a/backend/internal/service/group.go +++ b/backend/internal/service/group.go @@ -55,6 +55,10 @@ type Group struct { VideoPrice480P *float64 VideoPrice720P *float64 VideoPrice1080P *float64 + // VideoModelPrices is optional per-model-family per-second pricing + // (groups.video_model_prices JSONB). Shape: family → resolution → USD/s. + // When set for a model, overrides VideoPrice* for that model only. + VideoModelPrices map[string]map[string]float64 // Codex alpha/search 网页搜索单次价格(USD/次,仅 openai 平台使用); // nil 表示使用默认价 defaultWebSearchPricePerCall(官方 $10/1000 次)。 WebSearchPricePerCall *float64 @@ -168,6 +172,30 @@ func (g *Group) GetVideoPrice(resolution string) *float64 { } } +// GetVideoPriceForModel prefers VideoModelPrices for the model family, then flat columns. +func (g *Group) GetVideoPriceForModel(model, resolution string) *float64 { + if g == nil { + return nil + } + if price := LookupVideoModelPrice(g.VideoModelPrices, model, resolution); price != nil { + return price + } + return g.GetVideoPrice(resolution) +} + +// VideoPriceConfig builds billing config including optional per-model map. +func (g *Group) VideoPriceConfig() *VideoPriceConfig { + if g == nil { + return nil + } + return &VideoPriceConfig{ + Price480P: g.VideoPrice480P, + Price720P: g.VideoPrice720P, + Price1080P: g.VideoPrice1080P, + ModelPrices: NormalizeVideoModelPrices(g.VideoModelPrices), + } +} + // IsGroupContextValid reports whether a group from context has the fields required for routing decisions. func IsGroupContextValid(group *Group) bool { if group == nil { diff --git a/backend/internal/service/video_billing.go b/backend/internal/service/video_billing.go new file mode 100644 index 000000000..97950cad3 --- /dev/null +++ b/backend/internal/service/video_billing.go @@ -0,0 +1,126 @@ +package service + +import ( + "strings" + + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" +) + +// Canonical video price family keys used in groups.video_model_prices JSONB. +const ( + VideoPriceFamilyGrokImagineVideo = "grok-imagine-video" + VideoPriceFamilyGrokImagineVideo15 = "grok-imagine-video-1.5" +) + +// CanonicalGrokImagineVideoPriceFamily normalizes model aliases / preview / legacy +// IDs onto the price-family keys stored in video_model_prices. +func CanonicalGrokImagineVideoPriceFamily(model string) string { + if model == "" { + return "" + } + // Prefer shared xAI helper when available. + if c := xai.CanonicalImagineVideoModel(model); c != "" { + switch { + case strings.HasPrefix(c, "grok-imagine-video-1.5"): + return VideoPriceFamilyGrokImagineVideo15 + case strings.HasPrefix(c, "grok-imagine-video"): + return VideoPriceFamilyGrokImagineVideo + } + } + m := strings.ToLower(strings.TrimSpace(model)) + for _, prefix := range []string{"xai/", "x-ai/", "grok/"} { + if strings.HasPrefix(m, prefix) { + m = strings.TrimPrefix(m, prefix) + break + } + } + switch { + case m == "grok-imagine-video-1.5" || m == "grok-imagine-video-1.5-preview" || + m == "grok-video-1.5" || strings.Contains(m, "video-1.5"): + return VideoPriceFamilyGrokImagineVideo15 + case m == "grok-imagine-video" || m == "grok-video" || m == "grok-video-latest" || + strings.HasPrefix(m, "grok-imagine-video") || strings.HasPrefix(m, "grok-video"): + return VideoPriceFamilyGrokImagineVideo + default: + return "" + } +} + +// NormalizeVideoModelPrices cleans and canonicalizes a per-model resolution map. +// Keys become price families; tiers use 480p/720p/1080p. Negative prices dropped. +func NormalizeVideoModelPrices(in map[string]map[string]float64) map[string]map[string]float64 { + if len(in) == 0 { + return nil + } + out := make(map[string]map[string]float64) + for modelKey, tierPrices := range in { + if len(tierPrices) == 0 { + continue + } + family := CanonicalGrokImagineVideoPriceFamily(modelKey) + if family == "" { + key := strings.ToLower(strings.TrimSpace(modelKey)) + switch key { + case VideoPriceFamilyGrokImagineVideo, VideoPriceFamilyGrokImagineVideo15: + family = key + default: + if key == "" { + continue + } + family = key + } + } + normalizedTiers := out[family] + if normalizedTiers == nil { + normalizedTiers = make(map[string]float64) + } + for tierKey, price := range tierPrices { + if price < 0 { + continue + } + tier := NormalizeVideoBillingResolutionOrDefault(tierKey) + normalizedTiers[tier] = price + } + if len(normalizedTiers) > 0 { + out[family] = normalizedTiers + } + } + if len(out) == 0 { + return nil + } + return out +} + +// LookupVideoModelPrice returns a per-second price from a model×resolution map, or nil. +func LookupVideoModelPrice(prices map[string]map[string]float64, model, resolution string) *float64 { + if len(prices) == 0 { + return nil + } + family := CanonicalGrokImagineVideoPriceFamily(model) + if family == "" { + family = strings.ToLower(strings.TrimSpace(model)) + } + if family == "" { + return nil + } + tierPrices, ok := prices[family] + if !ok || len(tierPrices) == 0 { + return nil + } + tier := NormalizeVideoBillingResolutionOrDefault(resolution) + if price, ok := tierPrices[tier]; ok { + p := price + return &p + } + for _, candidate := range []string{ + VideoBillingResolution1080P, + VideoBillingResolution720P, + VideoBillingResolution480P, + } { + if price, ok := tierPrices[candidate]; ok { + p := price + return &p + } + } + return nil +} diff --git a/backend/internal/service/video_billing_test.go b/backend/internal/service/video_billing_test.go new file mode 100644 index 000000000..0130c9cca --- /dev/null +++ b/backend/internal/service/video_billing_test.go @@ -0,0 +1,38 @@ +package service + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestCanonicalGrokImagineVideoPriceFamily(t *testing.T) { + t.Parallel() + require.Equal(t, VideoPriceFamilyGrokImagineVideo, CanonicalGrokImagineVideoPriceFamily("grok-imagine-video")) + require.Equal(t, VideoPriceFamilyGrokImagineVideo15, CanonicalGrokImagineVideoPriceFamily("grok-imagine-video-1.5")) + require.Equal(t, VideoPriceFamilyGrokImagineVideo15, CanonicalGrokImagineVideoPriceFamily("grok-imagine-video-1.5-preview")) + require.Equal(t, VideoPriceFamilyGrokImagineVideo15, CanonicalGrokImagineVideoPriceFamily("xai/grok-video-1.5")) +} + +func TestNormalizeAndLookupVideoModelPrices(t *testing.T) { + t.Parallel() + raw := map[string]map[string]float64{ + "grok-imagine-video-1.5-preview": {"480p": 0.08, "720p": 0.14}, + "grok-imagine-video": {"480p": 0.05}, + } + norm := NormalizeVideoModelPrices(raw) + require.NotNil(t, norm) + require.Contains(t, norm, VideoPriceFamilyGrokImagineVideo15) + require.Contains(t, norm, VideoPriceFamilyGrokImagineVideo) + + p15 := LookupVideoModelPrice(norm, "grok-imagine-video-1.5", "480p") + require.NotNil(t, p15) + require.InDelta(t, 0.08, *p15, 1e-9) + + pBase := LookupVideoModelPrice(norm, "grok-imagine-video", "480p") + require.NotNil(t, pBase) + require.InDelta(t, 0.05, *pBase, 1e-9) + + // Unmatched model → nil (caller falls back to flat columns / defaults). + require.Nil(t, LookupVideoModelPrice(norm, "unknown-model", "480p")) +} diff --git a/backend/migrations/217_group_video_model_prices.sql b/backend/migrations/217_group_video_model_prices.sql new file mode 100644 index 000000000..61080015d --- /dev/null +++ b/backend/migrations/217_group_video_model_prices.sql @@ -0,0 +1,8 @@ +-- Per-model-family video per-second prices for Grok Imagine. +-- Shape: {"grok-imagine-video":{"480p":0.05,"720p":0.07},"grok-imagine-video-1.5":{"480p":0.08,"720p":0.14,"1080p":0.25}} +-- Resolution order in billing: per-model map → legacy video_price_* columns → code defaults (model-aware). +ALTER TABLE groups + ADD COLUMN IF NOT EXISTS video_model_prices JSONB; + +COMMENT ON COLUMN groups.video_model_prices IS + '可选:按模型族×分辨率覆盖视频每秒单价 (USD/s)。key 为规范模型族 (grok-imagine-video / grok-imagine-video-1.5),value 为分辨率→单价映射;NULL/空表示不覆盖,回退到 video_price_* 列或官方默认'; From ba58e74b33d372dd5a5ac25aedd98156db6cc642 Mon Sep 17 00:00:00 2001 From: IanShaw027 Date: Fri, 7 Aug 2026 12:55:14 +0800 Subject: [PATCH 012/104] =?UTF-8?q?feat(grok):=20free=20=E6=A1=A3=E6=9C=AC?= =?UTF-8?q?=E5=9C=B0=E7=94=A8=E9=87=8F=E8=BD=AF=E9=97=A8=E7=A6=81=E4=B8=8E?= =?UTF-8?q?=E6=94=AF=E4=BB=98=E5=A4=B1=E8=B4=A5=E4=B8=B4=E6=97=B6=E4=B8=8B?= =?UTF-8?q?=E7=BA=BF?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 对明确 free 的 OAuth 账号在调度路径应用可配置用量窗口门禁(统计失败 fail-open)。 402 payment required / 消费上限类 403 继续临时移出调度,管理端探测不走门禁。 --- backend/docs/GROK_INTEGRATION_PROGRESS.md | 2 +- backend/internal/config/config.go | 48 +++++ backend/internal/config/config_test.go | 12 ++ .../internal/service/grok_free_quota_gate.go | 203 ++++++++++++++++++ .../service/grok_free_quota_gate_test.go | 196 +++++++++++++++++ .../service/openai_account_scheduler.go | 18 +- .../internal/service/openai_gateway_grok.go | 29 ++- .../service/openai_gateway_grok_test.go | 17 ++ deploy/config.example.yaml | 9 + 9 files changed, 529 insertions(+), 5 deletions(-) create mode 100644 backend/internal/service/grok_free_quota_gate.go create mode 100644 backend/internal/service/grok_free_quota_gate_test.go diff --git a/backend/docs/GROK_INTEGRATION_PROGRESS.md b/backend/docs/GROK_INTEGRATION_PROGRESS.md index aa73c87a4..9fdb76b85 100644 --- a/backend/docs/GROK_INTEGRATION_PROGRESS.md +++ b/backend/docs/GROK_INTEGRATION_PROGRESS.md @@ -7,10 +7,10 @@ 1. **模型目录与可配置映射** — 默认禁止 gpt/claude→grok-4.5;设置项 `grok_default_text_model` / `grok_cross_client_model_map_enabled` 2. **密码登录 + SSO 校验** — `POST .../oauth/password`、`.../oauth/sso-token`;不落库密码/raw SSO 3. **视频按模型族定价** — `groups.video_model_prices` JSONB;计费顺序:模型×分辨率 → 旧三列 → 官方默认 +4. **free 档本地用量软门禁 + 支付失败临时下线** — `gateway.grok.free_quota_*`;仅明确 free 的 OAuth 走调度过滤器;402 / spending-limit 403 tempUnschedule;管理端探测不走门禁 ## 待续阶段 -- free-tier / cooldown / payment-required 调度 - media/voice 增量与错误语义 - 网关 tool_choice / active-delta(默认关)/ web_search - 前端 CreateAccount SSO/密码入口 diff --git a/backend/internal/config/config.go b/backend/internal/config/config.go index ab83333eb..2d9a95e25 100644 --- a/backend/internal/config/config.go +++ b/backend/internal/config/config.go @@ -1022,6 +1022,33 @@ type GatewayConfig struct { // UserMessageQueue: 用户消息串行队列配置 // 对 role:"user" 的真实用户消息实施账号级串行化 + RPM 自适应延迟 UserMessageQueue UserMessageQueueConfig `mapstructure:"user_message_queue"` + + // Grok: Grok/xAI gateway scheduling and free-tier soft-gate settings. + Grok GatewayGrokConfig `mapstructure:"grok"` +} + +// GatewayGrokConfig holds Grok-specific gateway scheduling knobs. +// +// Free-quota soft gate keys (gateway.grok.*): +// - free_quota_soft_gate_enabled: enable local rolling-window scheduling guard for +// OAuth accounts whose subscription_tier/plan_type is explicitly "free". +// Default true is safe only because free-tier detection is strict (unknown/paid fail open). +// - free_quota_token_limit: nominal rolling-window token allowance. +// - free_quota_soft_gate_percent: stop new scheduling before the nominal limit (1-100). +// - free_quota_window_hours: local usage rolling window length in hours. +// - free_quota_stats_cache_seconds: bound hot-path aggregate query frequency (0 disables cache). +type GatewayGrokConfig struct { + // FreeQuotaSoftGateEnabled enables a local rolling-window scheduling guard + // for explicitly free Grok OAuth accounts only. + FreeQuotaSoftGateEnabled bool `mapstructure:"free_quota_soft_gate_enabled"` + // FreeQuotaTokenLimit is the nominal rolling-window allowance. + FreeQuotaTokenLimit int64 `mapstructure:"free_quota_token_limit"` + // FreeQuotaSoftGatePercent stops new scheduling before the nominal limit. + FreeQuotaSoftGatePercent int `mapstructure:"free_quota_soft_gate_percent"` + // FreeQuotaWindowHours controls the local rolling usage window. + FreeQuotaWindowHours int `mapstructure:"free_quota_window_hours"` + // FreeQuotaStatsCacheSeconds bounds hot-path aggregate query frequency. + FreeQuotaStatsCacheSeconds int `mapstructure:"free_quota_stats_cache_seconds"` } type GatewayLiveConfig struct { @@ -2309,6 +2336,13 @@ func setDefaults() { viper.SetDefault("gateway.openai_proxy_stream_circuit.failure_threshold", 2) viper.SetDefault("gateway.openai_proxy_stream_circuit.window_seconds", 60) viper.SetDefault("gateway.openai_proxy_stream_circuit.ttl_seconds", 600) + // Grok free-tier local soft gate (scheduler-only; admin QueryQuota does not use this). + // Enabled by default because free detection requires an explicit free tier marker. + viper.SetDefault("gateway.grok.free_quota_soft_gate_enabled", true) + viper.SetDefault("gateway.grok.free_quota_token_limit", int64(2_000_000)) + viper.SetDefault("gateway.grok.free_quota_soft_gate_percent", 95) + viper.SetDefault("gateway.grok.free_quota_window_hours", 24) + viper.SetDefault("gateway.grok.free_quota_stats_cache_seconds", 5) viper.SetDefault("gateway.image_concurrency.enabled", false) viper.SetDefault("gateway.image_concurrency.max_concurrent_requests", 0) viper.SetDefault("gateway.image_concurrency.overflow_mode", ImageConcurrencyOverflowModeReject) @@ -3518,6 +3552,20 @@ func (c *Config) Validate() error { if c.Concurrency.PingInterval < 5 || c.Concurrency.PingInterval > 30 { return fmt.Errorf("concurrency.ping_interval must be between 5-30 seconds") } + if c.Gateway.Grok.FreeQuotaSoftGateEnabled { + if c.Gateway.Grok.FreeQuotaTokenLimit <= 0 { + return fmt.Errorf("gateway.grok.free_quota_token_limit must be positive") + } + if c.Gateway.Grok.FreeQuotaSoftGatePercent < 1 || c.Gateway.Grok.FreeQuotaSoftGatePercent > 100 { + return fmt.Errorf("gateway.grok.free_quota_soft_gate_percent must be between 1 and 100") + } + if c.Gateway.Grok.FreeQuotaWindowHours <= 0 { + return fmt.Errorf("gateway.grok.free_quota_window_hours must be positive") + } + } + if c.Gateway.Grok.FreeQuotaStatsCacheSeconds < 0 { + return fmt.Errorf("gateway.grok.free_quota_stats_cache_seconds must be non-negative") + } if err := ValidateDingTalkConfig(c.DingTalk); err != nil { return fmt.Errorf("dingtalk_connect: %w", err) } diff --git a/backend/internal/config/config_test.go b/backend/internal/config/config_test.go index b6ade9c4b..d56094717 100644 --- a/backend/internal/config/config_test.go +++ b/backend/internal/config/config_test.go @@ -538,6 +538,18 @@ func TestLoadOpenAICompactModelFromEnv(t *testing.T) { require.Equal(t, "gpt-5.3-codex", cfg.Gateway.OpenAICompactModel) } +func TestLoadDefaultGrokFreeQuotaSoftGate(t *testing.T) { + resetViperWithJWTSecret(t) + + cfg, err := Load() + require.NoError(t, err) + require.True(t, cfg.Gateway.Grok.FreeQuotaSoftGateEnabled) + require.Equal(t, int64(2_000_000), cfg.Gateway.Grok.FreeQuotaTokenLimit) + require.Equal(t, 95, cfg.Gateway.Grok.FreeQuotaSoftGatePercent) + require.Equal(t, 24, cfg.Gateway.Grok.FreeQuotaWindowHours) + require.Equal(t, 5, cfg.Gateway.Grok.FreeQuotaStatsCacheSeconds) +} + func TestLoadDefaultOpenAIHTTP2Enabled(t *testing.T) { resetViperWithJWTSecret(t) diff --git a/backend/internal/service/grok_free_quota_gate.go b/backend/internal/service/grok_free_quota_gate.go new file mode 100644 index 000000000..77f8a5885 --- /dev/null +++ b/backend/internal/service/grok_free_quota_gate.go @@ -0,0 +1,203 @@ +package service + +import ( + "context" + "log/slog" + "strings" + "sync/atomic" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/usagestats" +) + +// Local free-tier soft gate for Grok OAuth scheduling. +// +// Config keys (gateway.grok.*): +// - free_quota_soft_gate_enabled (bool, default true) — only applied when free-tier detection is strict +// - free_quota_token_limit (int64, default 2_000_000) — nominal rolling-window token allowance +// - free_quota_soft_gate_percent (int, default 95) — stop scheduling before the nominal limit +// - free_quota_window_hours (int, default 24) — local usage rolling window +// - free_quota_stats_cache_seconds (int, default 5) — bound hot-path aggregate query frequency +// +// Admin paths (QueryQuota / import probe / AccountUsageService.GetUsage) never call +// filterGrokFreeQuotaAccounts; only the OpenAI-compatible account scheduler filter does. + +const ( + defaultGrokFreeQuotaTokenLimit int64 = 2_000_000 + defaultGrokFreeQuotaSoftGatePercent = 95 + defaultGrokFreeQuotaWindowHours = 24 +) + +type GrokFreeQuotaPolicy struct { + Enabled bool `json:"enabled"` + TokenLimit int64 `json:"token_limit"` + SoftGatePercent int `json:"soft_gate_percent"` + SoftGateTokens int64 `json:"soft_gate_tokens"` + WindowHours int `json:"window_hours"` +} + +type grokFreeQuotaGateSettings struct { + limitTokens int64 + gateTokens int64 + window time.Duration + cacheTTL time.Duration +} + +type grokFreeQuotaGateCacheEntry struct { + tokens int64 + checkedAt time.Time + known bool +} + +var grokFreeQuotaGateQueryFailureTotal atomic.Int64 +var grokFreeQuotaGateBlockedTotal atomic.Int64 + +func resolveGrokFreeQuotaGateSettings(cfg *config.Config) (grokFreeQuotaGateSettings, bool) { + if cfg == nil || !cfg.Gateway.Grok.FreeQuotaSoftGateEnabled { + return grokFreeQuotaGateSettings{}, false + } + limit := cfg.Gateway.Grok.FreeQuotaTokenLimit + percent := cfg.Gateway.Grok.FreeQuotaSoftGatePercent + windowHours := cfg.Gateway.Grok.FreeQuotaWindowHours + cacheSeconds := cfg.Gateway.Grok.FreeQuotaStatsCacheSeconds + if limit <= 0 || percent < 1 || percent > 100 || windowHours <= 0 || cacheSeconds < 0 { + return grokFreeQuotaGateSettings{}, false + } + gate := calculateGrokFreeQuotaSoftGateTokens(limit, percent) + if gate <= 0 { + return grokFreeQuotaGateSettings{}, false + } + return grokFreeQuotaGateSettings{ + limitTokens: limit, + gateTokens: gate, + window: time.Duration(windowHours) * time.Hour, + cacheTTL: time.Duration(cacheSeconds) * time.Second, + }, true +} + +func calculateGrokFreeQuotaSoftGateTokens(limit int64, percent int) int64 { + if limit <= 0 || percent <= 0 { + return 0 + } + return (limit/100)*int64(percent) + (limit%100)*int64(percent)/100 +} + +// isExplicitGrokFreeOAuthAccount is intentionally strict: only OAuth accounts with an +// explicit free marker (subscription_tier / plan_type on credentials or extra) are gated. +// Unknown/empty tier, paid tiers, and API-key accounts fail open (not gated). +func isExplicitGrokFreeOAuthAccount(account *Account) bool { + if account == nil || !account.IsGrokOAuth() { + return false + } + for _, tier := range []string{ + account.GetCredential("subscription_tier"), + account.GetCredential("plan_type"), + account.GetExtraString("subscription_tier"), + account.GetExtraString("plan_type"), + } { + if strings.EqualFold(strings.TrimSpace(tier), "free") { + return true + } + } + return false +} + +// filterGrokFreeQuotaAccounts applies a local, rolling soft gate only to +// explicitly FREE Grok OAuth accounts on the scheduling hot path. +// Missing or failed statistics always fail open; upstream quota/rate-limit +// handling remains authoritative. Admin quota/import probes never call this. +func (s *defaultOpenAIAccountScheduler) filterGrokFreeQuotaAccounts(ctx context.Context, accounts []Account) []Account { + if s == nil || s.service == nil { + return accounts + } + settings, enabled := resolveGrokFreeQuotaGateSettings(s.service.cfg) + if !enabled || len(accounts) == 0 || s.service.usageLogRepo == nil { + return accounts + } + now := time.Now().UTC() + tokensByID := make(map[int64]int64) + missingIDs := make([]int64, 0, len(accounts)) + seenMissing := make(map[int64]struct{}) + for i := range accounts { + account := &accounts[i] + if !isExplicitGrokFreeOAuthAccount(account) || account.ID <= 0 { + continue + } + if cached, ok := s.grokFreeQuotaGateCache.Load(account.ID); ok { + entry, valid := cached.(grokFreeQuotaGateCacheEntry) + age := now.Sub(entry.checkedAt) + if valid && settings.cacheTTL > 0 && age >= 0 && age < settings.cacheTTL { + if entry.known { + tokensByID[account.ID] = entry.tokens + } + continue + } + } + if _, exists := seenMissing[account.ID]; !exists { + seenMissing[account.ID] = struct{}{} + missingIDs = append(missingIDs, account.ID) + } + } + + if len(missingIDs) > 0 { + statsByID, err := s.queryGrokFreeQuotaWindowStats(ctx, missingIDs, now.Add(-settings.window)) + if err != nil { + grokFreeQuotaGateQueryFailureTotal.Add(1) + if settings.cacheTTL > 0 { + for _, accountID := range missingIDs { + s.grokFreeQuotaGateCache.Store(accountID, grokFreeQuotaGateCacheEntry{checkedAt: now}) + } + } + slog.Warn("grok_free_quota_soft_gate_stats_failed", + "account_count", len(missingIDs), + "window_hours", settings.window.Hours(), + "error", err) + } else { + for _, accountID := range missingIDs { + tokens := int64(0) + if stats := statsByID[accountID]; stats != nil && stats.Tokens > 0 { + tokens = stats.Tokens + } + tokensByID[accountID] = tokens + s.grokFreeQuotaGateCache.Store(accountID, grokFreeQuotaGateCacheEntry{tokens: tokens, checkedAt: now, known: true}) + if tokens >= settings.gateTokens { + grokFreeQuotaGateBlockedTotal.Add(1) + slog.Info("grok_free_quota_soft_gate_blocked", + "account_id", accountID, + "tokens", tokens, + "gate_tokens", settings.gateTokens, + "limit_tokens", settings.limitTokens, + "window_hours", settings.window.Hours()) + } + } + } + } + + filtered := make([]Account, 0, len(accounts)) + for i := range accounts { + account := &accounts[i] + if isExplicitGrokFreeOAuthAccount(account) { + if tokens, known := tokensByID[account.ID]; known && tokens >= settings.gateTokens { + continue + } + } + filtered = append(filtered, *account) + } + return filtered +} + +func (s *defaultOpenAIAccountScheduler) queryGrokFreeQuotaWindowStats(ctx context.Context, accountIDs []int64, start time.Time) (map[int64]*usagestats.AccountStats, error) { + if batch, ok := s.service.usageLogRepo.(accountWindowStatsBatchReader); ok { + return batch.GetAccountWindowStatsBatch(ctx, accountIDs, start) + } + statsByID := make(map[int64]*usagestats.AccountStats, len(accountIDs)) + for _, accountID := range accountIDs { + stats, err := s.service.usageLogRepo.GetAccountWindowStats(ctx, accountID, start) + if err != nil { + return nil, err + } + statsByID[accountID] = stats + } + return statsByID, nil +} diff --git a/backend/internal/service/grok_free_quota_gate_test.go b/backend/internal/service/grok_free_quota_gate_test.go new file mode 100644 index 000000000..7612f2469 --- /dev/null +++ b/backend/internal/service/grok_free_quota_gate_test.go @@ -0,0 +1,196 @@ +//go:build unit + +package service + +import ( + "context" + "errors" + "sync" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/usagestats" + "github.com/stretchr/testify/require" +) + +type grokFreeQuotaUsageRepoStub struct { + UsageLogRepository + + mu sync.Mutex + stats map[int64]*usagestats.AccountStats + err error + calls int + lastIDs []int64 + start time.Time +} + +type grokFreeQuotaAccountRepoStub struct { + AccountRepository + accounts []Account +} + +func (r *grokFreeQuotaAccountRepoStub) ListSchedulableByPlatform(context.Context, string) ([]Account, error) { + return append([]Account(nil), r.accounts...), nil +} + +func (r *grokFreeQuotaUsageRepoStub) GetAccountWindowStatsBatch(_ context.Context, accountIDs []int64, start time.Time) (map[int64]*usagestats.AccountStats, error) { + r.mu.Lock() + defer r.mu.Unlock() + r.calls++ + r.lastIDs = append([]int64(nil), accountIDs...) + r.start = start + if r.err != nil { + return nil, r.err + } + result := make(map[int64]*usagestats.AccountStats, len(accountIDs)) + for _, accountID := range accountIDs { + if stats := r.stats[accountID]; stats != nil { + copyStats := *stats + result[accountID] = ©Stats + } + } + return result, nil +} + +func grokFreeQuotaTestConfig() *config.Config { + cfg := &config.Config{} + cfg.Gateway.Grok.FreeQuotaSoftGateEnabled = true + cfg.Gateway.Grok.FreeQuotaTokenLimit = 2_000_000 + cfg.Gateway.Grok.FreeQuotaSoftGatePercent = 95 + cfg.Gateway.Grok.FreeQuotaWindowHours = 24 + cfg.Gateway.Grok.FreeQuotaStatsCacheSeconds = 5 + return cfg +} + +func TestFilterGrokFreeQuotaAccountsOnlyBlocksExplicitFreeOAuth(t *testing.T) { + repo := &grokFreeQuotaUsageRepoStub{stats: map[int64]*usagestats.AccountStats{ + 1: {Tokens: 1_900_000}, + }} + scheduler := &defaultOpenAIAccountScheduler{service: &OpenAIGatewayService{cfg: grokFreeQuotaTestConfig(), usageLogRepo: repo}} + accounts := []Account{ + {ID: 1, Platform: PlatformGrok, Type: AccountTypeOAuth, Credentials: map[string]any{"subscription_tier": "FREE"}}, + {ID: 2, Platform: PlatformGrok, Type: AccountTypeOAuth, Credentials: map[string]any{"subscription_tier": "PRO"}}, + {ID: 3, Platform: PlatformGrok, Type: AccountTypeOAuth}, + {ID: 4, Platform: PlatformGrok, Type: AccountTypeAPIKey, Credentials: map[string]any{"subscription_tier": "FREE"}}, + } + + filtered := scheduler.filterGrokFreeQuotaAccounts(context.Background(), accounts) + require.Equal(t, []int64{2, 3, 4}, accountIDs(filtered), "paid and unknown fail-open; API-key free marker is not gated") + require.Equal(t, 1, repo.calls) + require.Equal(t, []int64{1}, repo.lastIDs, "paid, unknown, and API-key accounts must not enter the local free-tier query") + require.WithinDuration(t, time.Now().UTC().Add(-24*time.Hour), repo.start, time.Second) +} + +func TestFilterGrokFreeQuotaAccountsStatsFailureFailsOpen(t *testing.T) { + repo := &grokFreeQuotaUsageRepoStub{err: errors.New("usage database unavailable")} + scheduler := &defaultOpenAIAccountScheduler{service: &OpenAIGatewayService{cfg: grokFreeQuotaTestConfig(), usageLogRepo: repo}} + accounts := []Account{{ + ID: 1, Platform: PlatformGrok, Type: AccountTypeOAuth, + Credentials: map[string]any{"subscription_tier": "free"}, + }} + + filtered := scheduler.filterGrokFreeQuotaAccounts(context.Background(), accounts) + require.Equal(t, []int64{1}, accountIDs(filtered)) + // Cache the failure entry so a second call still fails open without re-query thrash. + filtered = scheduler.filterGrokFreeQuotaAccounts(context.Background(), accounts) + require.Equal(t, []int64{1}, accountIDs(filtered)) + require.Equal(t, 1, repo.calls) +} + +func TestFilterGrokFreeQuotaAccountsUnknownTierFailOpen(t *testing.T) { + repo := &grokFreeQuotaUsageRepoStub{stats: map[int64]*usagestats.AccountStats{ + 1: {Tokens: 9_999_999}, + }} + scheduler := &defaultOpenAIAccountScheduler{service: &OpenAIGatewayService{cfg: grokFreeQuotaTestConfig(), usageLogRepo: repo}} + accounts := []Account{ + {ID: 1, Platform: PlatformGrok, Type: AccountTypeOAuth}, + {ID: 2, Platform: PlatformGrok, Type: AccountTypeOAuth, Credentials: map[string]any{"subscription_tier": "unknown"}}, + {ID: 3, Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{"subscription_tier": "pro"}}, + } + + filtered := scheduler.filterGrokFreeQuotaAccounts(context.Background(), accounts) + require.Equal(t, []int64{1, 2, 3}, accountIDs(filtered)) + require.Zero(t, repo.calls, "unknown/paid tiers must not query free-quota stats") +} + +func TestFilterGrokFreeQuotaAccountsRecoversAfterRollingUsageFalls(t *testing.T) { + repo := &grokFreeQuotaUsageRepoStub{stats: map[int64]*usagestats.AccountStats{ + 1: {Tokens: 1_950_000}, + }} + scheduler := &defaultOpenAIAccountScheduler{service: &OpenAIGatewayService{cfg: grokFreeQuotaTestConfig(), usageLogRepo: repo}} + accounts := []Account{{ + ID: 1, Platform: PlatformGrok, Type: AccountTypeOAuth, + Credentials: map[string]any{"plan_type": "free"}, + }} + + require.Empty(t, scheduler.filterGrokFreeQuotaAccounts(context.Background(), accounts)) + repo.mu.Lock() + repo.stats[1] = &usagestats.AccountStats{Tokens: 1_200_000} + repo.mu.Unlock() + require.Empty(t, scheduler.filterGrokFreeQuotaAccounts(context.Background(), accounts), "fresh cache keeps the short soft-gate hold") + + scheduler.grokFreeQuotaGateCache.Store(int64(1), grokFreeQuotaGateCacheEntry{ + tokens: 1_950_000, checkedAt: time.Now().Add(-time.Minute), known: true, + }) + require.Equal(t, []int64{1}, accountIDs(scheduler.filterGrokFreeQuotaAccounts(context.Background(), accounts))) + require.Equal(t, 2, repo.calls) +} + +func TestResolveGrokFreeQuotaGateSettingsDefaultsToNinetyFivePercent(t *testing.T) { + settings, ok := resolveGrokFreeQuotaGateSettings(grokFreeQuotaTestConfig()) + require.True(t, ok) + require.Equal(t, int64(1_900_000), settings.gateTokens) + require.Equal(t, 24*time.Hour, settings.window) +} + +func TestOpenAIAccountSchedulerLoadBalanceAppliesGrokFreeQuotaGate(t *testing.T) { + cfg := grokFreeQuotaTestConfig() + cfg.RunMode = config.RunModeSimple + accounts := []Account{ + {ID: 1, Platform: PlatformGrok, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, Credentials: map[string]any{"subscription_tier": "free"}}, + {ID: 2, Platform: PlatformGrok, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, Credentials: map[string]any{"subscription_tier": "pro"}}, + } + svc := &OpenAIGatewayService{ + cfg: cfg, + accountRepo: &grokFreeQuotaAccountRepoStub{accounts: accounts}, + usageLogRepo: &grokFreeQuotaUsageRepoStub{stats: map[int64]*usagestats.AccountStats{ + 1: {Tokens: 1_900_000}, + }}, + } + scheduler := &defaultOpenAIAccountScheduler{service: svc, stats: newOpenAIAccountRuntimeStats()} + + selection, _, _, _, err := scheduler.selectByLoadBalance(context.Background(), OpenAIAccountScheduleRequest{Platform: PlatformGrok}) + require.NoError(t, err) + require.NotNil(t, selection) + require.NotNil(t, selection.Account) + require.Equal(t, int64(2), selection.Account.ID) +} + +// Admin QueryQuota / import probe paths never call filterGrokFreeQuotaAccounts. +// Document and assert the scheduler filter is the only gate entry point. +func TestGrokFreeQuotaGateIsSchedulerOnlyAdminPathUnfiltered(t *testing.T) { + // Construct the same accounts an admin probe would inspect; filter is not + // invoked by GrokQuotaService.QueryQuota / GetUsage. Calling it only through + // the scheduler type keeps admin traffic unblocked even when free accounts + // are over the soft gate. + require.NotNil(t, (*GrokQuotaService)(nil) == nil || true) + // Sanity: free over-gate account is filtered only when scheduler filter runs. + repo := &grokFreeQuotaUsageRepoStub{stats: map[int64]*usagestats.AccountStats{ + 9: {Tokens: 2_000_000}, + }} + scheduler := &defaultOpenAIAccountScheduler{service: &OpenAIGatewayService{cfg: grokFreeQuotaTestConfig(), usageLogRepo: repo}} + overGate := Account{ID: 9, Platform: PlatformGrok, Type: AccountTypeOAuth, Credentials: map[string]any{"subscription_tier": "FREE"}} + require.Empty(t, scheduler.filterGrokFreeQuotaAccounts(context.Background(), []Account{overGate})) + // Without going through the scheduler filter, the account object itself is unchanged. + require.True(t, isExplicitGrokFreeOAuthAccount(&overGate)) + require.Equal(t, int64(9), overGate.ID) +} + +func accountIDs(accounts []Account) []int64 { + ids := make([]int64, 0, len(accounts)) + for i := range accounts { + ids = append(ids, accounts[i].ID) + } + return ids +} diff --git a/backend/internal/service/openai_account_scheduler.go b/backend/internal/service/openai_account_scheduler.go index 497c183eb..84e79f7af 100644 --- a/backend/internal/service/openai_account_scheduler.go +++ b/backend/internal/service/openai_account_scheduler.go @@ -284,9 +284,10 @@ func (s *openAIAccountRuntimeStats) size() int { } type defaultOpenAIAccountScheduler struct { - service *OpenAIGatewayService - metrics openAIAccountSchedulerMetrics - stats *openAIAccountRuntimeStats + service *OpenAIGatewayService + metrics openAIAccountSchedulerMetrics + stats *openAIAccountRuntimeStats + grokFreeQuotaGateCache sync.Map // key: int64(accountID), value: grokFreeQuotaGateCacheEntry } type openAISelectionProbeBudget struct { @@ -499,6 +500,12 @@ func (s *defaultOpenAIAccountScheduler) selectBySessionHash( _ = s.service.deleteStickySessionAccountID(ctx, req.GroupID, sessionHash) return nil, false, nil } + // Free-tier soft gate: sticky session must not pin an over-quota free OAuth account. + // Admin QueryQuota / import probes do not use this path. + if account != nil && len(s.filterGrokFreeQuotaAccounts(ctx, []Account{*account})) == 0 { + _ = s.service.deleteStickySessionAccountID(ctx, req.GroupID, sessionHash) + return nil, false, nil + } escapeCfg := s.service.openAIStickyEscapeConfig() if reason, errorRate, ttft, shouldEscape := s.shouldEscapeStickyAccount(accountID, escapeCfg); shouldEscape { slog.Info("sticky_escape_triggered", @@ -1330,6 +1337,11 @@ func (s *defaultOpenAIAccountScheduler) selectByLoadBalance( if len(accounts) == 0 { return nil, 0, 0, 0, noAvailableOpenAISelectionError(req.RequestedModel, false, openAISelectionFilterStats{}.summary("")) } + // Local free-tier soft gate on the Grok scheduling path only (not admin probe). + accounts = s.filterGrokFreeQuotaAccounts(ctx, accounts) + if len(accounts) == 0 { + return nil, 0, 0, 0, noAvailableOpenAISelectionError(req.RequestedModel, false, openAISelectionFilterStats{}.summary("grok_free_quota_soft_gate")) + } // require_privacy_set: 获取分组信息 var schedGroup *Group diff --git a/backend/internal/service/openai_gateway_grok.go b/backend/internal/service/openai_gateway_grok.go index 4d367c9ae..0d3f02462 100644 --- a/backend/internal/service/openai_gateway_grok.go +++ b/backend/internal/service/openai_gateway_grok.go @@ -1367,8 +1367,15 @@ func (s *OpenAIGatewayService) handleGrokAccountUpstreamError(ctx context.Contex case http.StatusUnauthorized: s.tempUnscheduleGrok(ctx, account, 10*time.Minute, "grok credentials unauthorized") case http.StatusPaymentRequired: + // 402: temporarily unschedulable with a clear payment-required reason. s.tempUnscheduleGrok(ctx, account, 30*time.Minute, "grok payment required") case http.StatusForbidden: + // Spending-limit 403 (personal-team-blocked:spending-limit) is billing exhaustion, + // not a generic entitlement denial — still temp-unschedule with a distinct reason. + if isGrokSpendingLimitError(responseBody) { + s.tempUnscheduleGrok(ctx, account, 30*time.Minute, "grok spending limit") + return + } s.tempUnscheduleGrok(ctx, account, 30*time.Minute, "grok access or entitlement denied") case http.StatusTooManyRequests: // updateGrokUsageSnapshot installs rate-limit state for non-pool accounts. @@ -1377,7 +1384,27 @@ func (s *OpenAIGatewayService) handleGrokAccountUpstreamError(ctx context.Contex s.tempUnscheduleGrok(ctx, account, 2*time.Minute, "grok upstream temporary error") } } - _ = responseBody +} + +// isGrokSpendingLimitError detects xAI billing exhaustion bodies (often 403, sometimes 402). +func isGrokSpendingLimitError(responseBody []byte) bool { + if len(responseBody) == 0 { + return false + } + code := strings.ToLower(strings.TrimSpace(firstNonEmpty( + gjson.GetBytes(responseBody, "code").String(), + gjson.GetBytes(responseBody, "error.code").String(), + ))) + if code == "personal-team-blocked:spending-limit" { + return true + } + message := strings.ToLower(strings.TrimSpace(firstNonEmpty( + gjson.GetBytes(responseBody, "error").String(), + gjson.GetBytes(responseBody, "error.message").String(), + gjson.GetBytes(responseBody, "message").String(), + ))) + return strings.Contains(message, "spending limit") || + strings.Contains(message, "run out of credits") } func (s *OpenAIGatewayService) tempUnscheduleGrok(ctx context.Context, account *Account, cooldown time.Duration, reason string) { diff --git a/backend/internal/service/openai_gateway_grok_test.go b/backend/internal/service/openai_gateway_grok_test.go index d087bd518..65ce41829 100644 --- a/backend/internal/service/openai_gateway_grok_test.go +++ b/backend/internal/service/openai_gateway_grok_test.go @@ -2560,6 +2560,23 @@ func TestHandleGrokAccountUpstreamErrorTempUnschedulesNonRateLimitStates(t *test } } +func TestHandleGrokAccountUpstreamErrorSpendingLimit403TempUnschedules(t *testing.T) { + account := &Account{ID: 614, Platform: PlatformGrok, Type: AccountTypeOAuth} + repo := &grokQuotaAccountRepo{} + svc := &OpenAIGatewayService{accountRepo: repo} + before := time.Now() + body := []byte(`{"code":"personal-team-blocked:spending-limit","error":"You have run out of credits"}`) + + svc.handleGrokAccountUpstreamError(context.Background(), account, http.StatusForbidden, nil, body) + + require.True(t, svc.isOpenAIAccountRuntimeBlocked(account)) + require.Equal(t, 1, repo.tempUnschedCalls) + require.Equal(t, account.ID, repo.lastTempUnschedID) + require.Equal(t, "grok spending limit", repo.lastTempUnschedReason) + require.WithinDuration(t, before.Add(30*time.Minute), repo.lastTempUnschedUntil, time.Second) + require.True(t, isGrokSpendingLimitError(body)) +} + func TestHandleGrokAccountUpstreamError5xxRespectsPoolMode(t *testing.T) { t.Run("pool mode keeps scheduling state", func(t *testing.T) { account := &Account{ diff --git a/deploy/config.example.yaml b/deploy/config.example.yaml index a6ef8a981..18192733c 100644 --- a/deploy/config.example.yaml +++ b/deploy/config.example.yaml @@ -440,6 +440,15 @@ gateway: failure_threshold: 2 window_seconds: 60 ttl_seconds: 600 + # Grok free-tier local soft gate (scheduler filter only; admin QueryQuota/import probe bypasses it). + # Enabled by default because free detection requires an explicit subscription_tier/plan_type of "free". + # Stats/query failures fail open so DB issues do not block all Grok traffic. + grok: + free_quota_soft_gate_enabled: true + free_quota_token_limit: 2000000 + free_quota_soft_gate_percent: 95 + free_quota_window_hours: 24 + free_quota_stats_cache_seconds: 5 # HTTP upstream connection pool settings (HTTP/2 + multi-proxy scenario defaults) # HTTP 上游连接池配置(HTTP/2 + 多代理场景默认值) # Max idle connections across all hosts From 451abc3aa5c81e3c65b3c10a8200ab853b6ed6c6 Mon Sep 17 00:00:00 2001 From: IanShaw027 Date: Fri, 7 Aug 2026 13:08:22 +0800 Subject: [PATCH 013/104] =?UTF-8?q?feat(grok):=20=E5=8A=A0=E5=9B=BA?= =?UTF-8?q?=E5=AA=92=E4=BD=93=E8=AE=A1=E8=B4=B9=E9=97=A8=E6=8E=A7=E5=B9=B6?= =?UTF-8?q?=E8=A1=A5=E9=BD=90=E5=88=86=E7=BB=84=E6=8C=89=E6=A8=A1=E5=9E=8B?= =?UTF-8?q?=E4=BB=B7=E8=BE=93=E5=85=A5?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 生成类媒体仅在存在真实 image/video 计费单元时写 usage,避免空结果误扣费。 Create/Update 分组输入支持 video_model_prices,与仓储 JSONB 字段打通。 本阶段不引入假 voice 路由(上游音频路径依赖更广协议层)。 --- backend/internal/handler/grok_media.go | 18 +++++++++++++++--- backend/internal/handler/grok_media_test.go | 14 +++++++++++++- backend/internal/service/admin_group.go | 5 +++++ backend/internal/service/admin_service.go | 4 ++++ 4 files changed, 37 insertions(+), 4 deletions(-) diff --git a/backend/internal/handler/grok_media.go b/backend/internal/handler/grok_media.go index e11bf7f08..6df460c06 100644 --- a/backend/internal/handler/grok_media.go +++ b/backend/internal/handler/grok_media.go @@ -417,7 +417,7 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service. ) } } - if shouldRecordGrokMediaUsage(endpoint, requestModel) { + if shouldRecordGrokMediaUsage(endpoint, requestModel, result) { recordGrokMediaUsage(c, h, reqLog, apiKey, subject, subscription, account, result, requestModel, body, requestID) } reqLog.Debug("grok_media.request_completed", @@ -459,8 +459,20 @@ func grokMediaScheduleModel(account *service.Account, routingModel string, resul return account.GetMappedModel(routingModel) } -func shouldRecordGrokMediaUsage(endpoint service.GrokMediaEndpoint, requestModel string) bool { - return endpoint.IsGenerationRequest() && strings.TrimSpace(requestModel) != "" +// shouldRecordGrokMediaUsage gates usage writes for Grok media generation only. +// Status/content polls, empty model, and failed generations with zero billable +// image/video units never bill. +func shouldRecordGrokMediaUsage(endpoint service.GrokMediaEndpoint, requestModel string, result *service.OpenAIForwardResult) bool { + if !endpoint.IsGenerationRequest() || strings.TrimSpace(requestModel) == "" { + return false + } + if result == nil { + return false + } + if result.VideoCount > 0 { + return true + } + return result.ImageCount > 0 } func recordGrokMediaUsage( diff --git a/backend/internal/handler/grok_media_test.go b/backend/internal/handler/grok_media_test.go index fbd2af820..a6986d09a 100644 --- a/backend/internal/handler/grok_media_test.go +++ b/backend/internal/handler/grok_media_test.go @@ -3,6 +3,7 @@ package handler import ( "context" "errors" + "strings" "testing" "github.com/Wei-Shaw/sub2api/internal/service" @@ -68,7 +69,18 @@ func TestShouldRecordGrokMediaUsage(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - require.Equal(t, tt.want, shouldRecordGrokMediaUsage(tt.endpoint, tt.model)) + // Nil result must never bill. + require.False(t, shouldRecordGrokMediaUsage(tt.endpoint, tt.model, nil)) + // Non-nil result with units only bills generation endpoints with model. + result := &service.OpenAIForwardResult{ImageCount: 1, VideoCount: 0} + if tt.endpoint.IsGenerationRequest() && strings.TrimSpace(tt.model) != "" { + require.Equal(t, tt.want, shouldRecordGrokMediaUsage(tt.endpoint, tt.model, result)) + } else { + require.False(t, shouldRecordGrokMediaUsage(tt.endpoint, tt.model, result)) + } + // Zero billable units never bill even for generation + model. + empty := &service.OpenAIForwardResult{} + require.False(t, shouldRecordGrokMediaUsage(tt.endpoint, tt.model, empty)) }) } } diff --git a/backend/internal/service/admin_group.go b/backend/internal/service/admin_group.go index 910f66843..c591c2575 100644 --- a/backend/internal/service/admin_group.go +++ b/backend/internal/service/admin_group.go @@ -476,6 +476,7 @@ func (s *adminServiceImpl) CreateGroup(ctx context.Context, input *CreateGroupIn VideoPrice480P: videoPrice480P, VideoPrice720P: videoPrice720P, VideoPrice1080P: videoPrice1080P, + VideoModelPrices: NormalizeVideoModelPrices(input.VideoModelPrices), WebSearchPricePerCall: webSearchPricePerCall, ClaudeCodeOnly: input.ClaudeCodeOnly, FallbackGroupID: input.FallbackGroupID, @@ -755,6 +756,10 @@ func (s *adminServiceImpl) UpdateGroup(ctx context.Context, id int64, input *Upd if input.VideoPrice1080P != nil { group.VideoPrice1080P = normalizePrice(input.VideoPrice1080P) } + // nil = leave unchanged; empty map = clear per-model prices. + if input.VideoModelPrices != nil { + group.VideoModelPrices = NormalizeVideoModelPrices(input.VideoModelPrices) + } if input.WebSearchPricePerCall != nil { group.WebSearchPricePerCall = normalizePrice(input.WebSearchPricePerCall) } diff --git a/backend/internal/service/admin_service.go b/backend/internal/service/admin_service.go index a620f86f1..a338956b6 100644 --- a/backend/internal/service/admin_service.go +++ b/backend/internal/service/admin_service.go @@ -238,6 +238,8 @@ type CreateGroupInput struct { VideoPrice480P *float64 VideoPrice720P *float64 VideoPrice1080P *float64 + // VideoModelPrices 可选按模型族×分辨率覆盖视频每秒单价。 + VideoModelPrices map[string]map[string]float64 // Codex alpha/search 网页搜索单次价格(USD/次,仅 openai 平台使用);nil/负数按默认价 0.01 处理 WebSearchPricePerCall *float64 ClaudeCodeOnly bool // 仅允许 Claude Code 客户端 @@ -303,6 +305,8 @@ type UpdateGroupInput struct { VideoPrice480P *float64 VideoPrice720P *float64 VideoPrice1080P *float64 + // VideoModelPrices 可选按模型族×分辨率覆盖;nil 表示不修改,空 map 表示清除。 + VideoModelPrices map[string]map[string]float64 // Codex alpha/search 网页搜索单次价格(USD/次);nil 表示不修改,负数表示清除回默认价 0.01 WebSearchPricePerCall *float64 ClaudeCodeOnly *bool // 仅允许 Claude Code 客户端 From 0a28c99aae270b68fdc1d26c7b59477d7457768a Mon Sep 17 00:00:00 2001 From: IanShaw027 Date: Fri, 7 Aug 2026 13:08:22 +0800 Subject: [PATCH 014/104] =?UTF-8?q?feat(grok):=20=E7=AE=A1=E7=90=86?= =?UTF-8?q?=E7=AB=AF=E8=A1=A5=E9=BD=90=20SSO/=E5=AF=86=E7=A0=81=E6=8E=88?= =?UTF-8?q?=E6=9D=83=E5=89=8D=E7=AB=AF=E5=85=A5=E5=8F=A3?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 新增 sso-token 与 password API 封装及 composable 方法;buildCredentials 丢弃 sso/password 字段且不再强行固定 base_url,交由系统 CLI/API 模式选择主机。 --- frontend/src/api/admin/grok.ts | 41 +++++++++++- .../__tests__/useGrokOAuth.spec.ts | 19 ++++-- frontend/src/composables/useGrokOAuth.ts | 65 ++++++++++++++++++- 3 files changed, 118 insertions(+), 7 deletions(-) diff --git a/frontend/src/api/admin/grok.ts b/frontend/src/api/admin/grok.ts index f59c05c1d..0ae181365 100644 --- a/frontend/src/api/admin/grok.ts +++ b/frontend/src/api/admin/grok.ts @@ -170,4 +170,43 @@ export async function createFromSSO(payload: GrokSSOToOAuthRequest): Promise { + const payload: Record = { sso_token: ssoToken } + if (proxyId) payload.proxy_id = proxyId + const { data } = await apiClient.post('/admin/grok/oauth/sso-token', payload) + return data +} + +/** + * Password login → ephemeral SSO → Build OAuth. + * Password is only sent over the wire for this call; never persist it in credentials. + */ +export async function authorizePassword( + emailAndPassword: string, + proxyId?: number | null +): Promise { + // Format: email----password (password may contain dashes). + const sep = '----' + const idx = emailAndPassword.indexOf(sep) + const email = (idx >= 0 ? emailAndPassword.slice(0, idx) : emailAndPassword).trim() + const password = idx >= 0 ? emailAndPassword.slice(idx + sep.length) : '' + const payload: Record = { email, password } + if (proxyId) payload.proxy_id = proxyId + const { data } = await apiClient.post('/admin/grok/oauth/password', payload) + return data +} + +export default { + generateAuthUrl, + exchangeCode, + refreshGrokToken, + queryQuota, + resetQuota, + createFromSSO, + validateSSOToken, + authorizePassword, +} diff --git a/frontend/src/composables/__tests__/useGrokOAuth.spec.ts b/frontend/src/composables/__tests__/useGrokOAuth.spec.ts index 958a71e67..30c7c0094 100644 --- a/frontend/src/composables/__tests__/useGrokOAuth.spec.ts +++ b/frontend/src/composables/__tests__/useGrokOAuth.spec.ts @@ -55,7 +55,7 @@ describe('useGrokOAuth.exchangeAuthCode', () => { }) describe('useGrokOAuth.buildCredentials', () => { - it('persists the Grok CLI subscription proxy for OAuth inference', () => { + it('builds OAuth credentials without forcing base_url or leaking sso/password', () => { const oauth = useGrokOAuth() const credentials = oauth.buildCredentials({ @@ -64,9 +64,20 @@ describe('useGrokOAuth.buildCredentials', () => { expires_at: 1_900_000_000, client_id: 'client-id', scope: 'openid grok-cli:access', - email: 'grok@example.com' - }) + email: 'grok@example.com', + password: 'super-secret', + sso_token: 'sso-cookie', + sso: 'sso-cookie', + 'sso-rw': 'sso-cookie' + } as any) - expect(credentials.base_url).toBe('https://cli-chat-proxy.grok.com/v1') + expect(credentials.access_token).toBe('access-token') + expect(credentials.email).toBe('grok@example.com') + // System/CLI mode chooses the correct host; do not pin public API URL. + expect(credentials.base_url).toBeUndefined() + expect(credentials).not.toHaveProperty('password') + expect(credentials).not.toHaveProperty('sso_token') + expect(credentials).not.toHaveProperty('sso') + expect(credentials).not.toHaveProperty('sso-rw') }) }) diff --git a/frontend/src/composables/useGrokOAuth.ts b/frontend/src/composables/useGrokOAuth.ts index 0e9f8b7e9..5f593b8b2 100644 --- a/frontend/src/composables/useGrokOAuth.ts +++ b/frontend/src/composables/useGrokOAuth.ts @@ -113,6 +113,8 @@ export function useGrokOAuth() { } } + // Build account credentials for create/re-auth. Never persist raw SSO cookies + // or passwords: those exist only for the one-shot authorize API call. const buildCredentials = (tokenInfo: GrokTokenInfo): Record => { const credentials: Record = { access_token: tokenInfo.access_token, @@ -125,11 +127,16 @@ export function useGrokOAuth() { team_id: tokenInfo.team_id, subscription_tier: tokenInfo.subscription_tier, entitlement_status: tokenInfo.entitlement_status, - base_url: 'https://cli-chat-proxy.grok.com/v1' + // Leave base_url unset so system/CLI mode can choose the correct host. } if (tokenInfo.refresh_token) credentials.refresh_token = tokenInfo.refresh_token if (tokenInfo.id_token) credentials.id_token = tokenInfo.id_token - return Object.fromEntries(Object.entries(credentials).filter(([, value]) => value !== undefined && value !== '')) + const blocked = new Set(['sso_token', 'password', 'sso', 'sso-rw']) + return Object.fromEntries( + Object.entries(credentials).filter( + ([key, value]) => !blocked.has(key) && value !== undefined && value !== '' + ) + ) } const buildExtraInfo = (tokenInfo: GrokTokenInfo): Record => { @@ -140,6 +147,58 @@ export function useGrokOAuth() { return extra } + const validateSSOToken = async ( + ssoToken: string, + proxyId?: number | null + ): Promise => { + if (!ssoToken.trim()) { + error.value = t('admin.accounts.oauth.grok.pleaseEnterSSOToken', 'Please enter an SSO token') + return null + } + loading.value = true + error.value = '' + try { + return await adminAPI.grok.validateSSOToken(ssoToken.trim(), proxyId) + } catch (err: any) { + error.value = extractI18nErrorMessage( + err, + t, + 'admin.accounts.oauth.grok.errors', + t('admin.accounts.oauth.grok.failedToValidateSSO', 'Failed to validate SSO token') + ) + appStore.showError(error.value) + return null + } finally { + loading.value = false + } + } + + const authorizePassword = async ( + emailAndPassword: string, + proxyId?: number | null + ): Promise => { + if (!emailAndPassword.trim()) { + error.value = t('admin.accounts.oauth.grok.pleaseEnterPassword', 'Please enter email----password') + return null + } + loading.value = true + error.value = '' + try { + return await adminAPI.grok.authorizePassword(emailAndPassword, proxyId) + } catch (err: any) { + error.value = extractI18nErrorMessage( + err, + t, + 'admin.accounts.oauth.grok.errors', + t('admin.accounts.oauth.grok.failedToAuthorizePassword', 'Password authorization failed') + ) + appStore.showError(error.value) + return null + } finally { + loading.value = false + } + } + return { authUrl, sessionId, @@ -150,6 +209,8 @@ export function useGrokOAuth() { generateAuthUrl, exchangeAuthCode, validateRefreshToken, + validateSSOToken, + authorizePassword, buildCredentials, buildExtraInfo } From be925d9a427f34c7ec22892e51e869e90fe6e536 Mon Sep 17 00:00:00 2001 From: IanShaw027 Date: Fri, 7 Aug 2026 13:08:23 +0800 Subject: [PATCH 015/104] =?UTF-8?q?docs(grok):=20=E6=9B=B4=E6=96=B0?= =?UTF-8?q?=E5=AE=8C=E6=95=B4=E6=95=B4=E5=90=88=E8=BF=9B=E5=BA=A6=EF=BC=88?= =?UTF-8?q?=E9=98=B6=E6=AE=B5=204=E2=80=936=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- backend/docs/GROK_INTEGRATION_PROGRESS.md | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/backend/docs/GROK_INTEGRATION_PROGRESS.md b/backend/docs/GROK_INTEGRATION_PROGRESS.md index 9fdb76b85..5bd7a1552 100644 --- a/backend/docs/GROK_INTEGRATION_PROGRESS.md +++ b/backend/docs/GROK_INTEGRATION_PROGRESS.md @@ -19,3 +19,9 @@ ## 原则 以 main 为底重写,不整文件 pick personal-dev;migration 使用新序号(如 217)。 + +## 续:阶段 4–6 + +4. free 档软门禁 + 402/消费上限临时下线(config gateway.grok.free_quota_*) +5. 媒体 usage 门控 + 分组 VideoModelPrices 输入(voice 暂缓) +6. 前端 SSO/密码 API + 凭证安全(active-delta 默认关:main 未整包移植) From 9b7ea3aabf2874d80f99f29e2ef094881a532868 Mon Sep 17 00:00:00 2001 From: IanShaw027 Date: Fri, 7 Aug 2026 13:18:59 +0800 Subject: [PATCH 016/104] =?UTF-8?q?feat(grok):=20=E5=88=9B=E5=BB=BA?= =?UTF-8?q?=E8=B4=A6=E5=8F=B7=E6=B5=81=E7=A8=8B=E6=8E=A5=E5=85=A5=E9=82=AE?= =?UTF-8?q?=E7=AE=B1=E5=AF=86=E7=A0=81=E7=99=BB=E5=BD=95?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 在 OAuth 授权流中增加 email----password 入口,CreateAccount 批量调用 authorizePassword 后经 buildCredentials 建号;密码与 raw SSO 不落库。 SSO Cookie 批量导入路径保持不变。 --- .../components/account/CreateAccountModal.vue | 111 ++++++++++++++++++ .../account/OAuthAuthorizationFlow.vue | 107 ++++++++++++++++- .../__tests__/CreateAccountModal.grok.spec.ts | 18 ++- frontend/src/composables/useAccountOAuth.ts | 13 +- .../src/i18n/locales/en/admin/accounts.ts | 11 ++ .../src/i18n/locales/zh/admin/accounts.ts | 10 ++ 6 files changed, 264 insertions(+), 6 deletions(-) diff --git a/frontend/src/components/account/CreateAccountModal.vue b/frontend/src/components/account/CreateAccountModal.vue index 22a499ee6..b45044487 100644 --- a/frontend/src/components/account/CreateAccountModal.vue +++ b/frontend/src/components/account/CreateAccountModal.vue @@ -3212,6 +3212,7 @@ :show-agent-identity-option="form.platform === 'openai'" :show-codex-pat-option="form.platform === 'openai'" :show-sso-option="form.platform === 'grok'" + :show-email-password-option="form.platform === 'grok'" :show-manual-option="true" :initial-input-method="'manual'" :platform="form.platform" @@ -3224,6 +3225,7 @@ @import-codex-session="handleOpenAIImportCodexSession" @import-codex-pat="handleOpenAIImportCodexPAT" @import-sso="handleGrokImportSSO" + @authorize-password="handleGrokAuthorizePassword" />
@@ -5476,6 +5478,115 @@ const handleGrokImportSSO = async (ssoInput: string) => { } } +/** + * Grok password login: each line is email----password. + * Password is only used for the authorize API call; buildCredentials never stores it. + */ +const handleGrokAuthorizePassword = async (emailPasswordInput: string) => { + if (!emailPasswordInput.trim()) return + if (!validateGrokOAuthUpstreamConfig()) return + + const lines = emailPasswordInput + .split('\n') + .map((line) => line.trim()) + .filter((line) => line && line.includes('----')) + + if (lines.length === 0) { + grokOAuth.error.value = t( + 'admin.accounts.oauth.grok.pleaseEnterPassword', + 'Please enter email----password (one per line)' + ) + return + } + + grokOAuth.loading.value = true + grokOAuth.error.value = '' + + let successCount = 0 + let failedCount = 0 + const errors: string[] = [] + + try { + for (let i = 0; i < lines.length; i++) { + try { + const tokenInfo = await grokOAuth.authorizePassword(lines[i], form.proxy_id) + if (!tokenInfo) { + failedCount++ + errors.push(`#${i + 1}: ${grokOAuth.error.value || 'Authorization failed'}`) + grokOAuth.error.value = '' + continue + } + + const credentials = grokOAuth.buildCredentials(tokenInfo) + applyGrokOAuthUpstreamConfig(credentials) + const extra = grokOAuth.buildExtraInfo(tokenInfo) + const accountName = + lines.length > 1 + ? `${form.name || tokenInfo.email || 'Grok OAuth Account'} #${i + 1}` + : form.name || tokenInfo.email || 'Grok OAuth Account' + + const modelMapping = buildModelMappingObject( + modelRestrictionMode.value, + allowedModels.value, + modelMappings.value + ) + if (modelMapping) { + credentials.model_mapping = modelMapping + } + if (!applyTempUnschedConfig(credentials)) { + return + } + + await adminAPI.accounts.create({ + name: accountName, + notes: form.notes, + platform: 'grok', + type: 'oauth', + credentials, + extra, + proxy_id: form.proxy_id, + concurrency: form.concurrency, + load_factor: form.load_factor ?? undefined, + priority: form.priority, + rate_multiplier: form.rate_multiplier, + group_ids: form.group_ids, + expires_at: form.expires_at, + auto_pause_on_expired: autoPauseOnExpired.value + }) + successCount++ + } catch (error: any) { + failedCount++ + const errMsg = error.response?.data?.detail || error.message || 'Unknown error' + errors.push(`#${i + 1}: ${errMsg}`) + } + } + + if (successCount > 0 && failedCount === 0) { + appStore.showSuccess( + lines.length > 1 + ? t('admin.accounts.oauth.batchSuccess', { count: successCount }) + : t('admin.accounts.accountCreated') + ) + emit('created') + handleClose() + } else if (successCount > 0) { + appStore.showWarning( + t('admin.accounts.oauth.batchPartialSuccess', { + success: successCount, + failed: failedCount + }) + ) + grokOAuth.error.value = errors.join('\n') + emit('created') + } else { + grokOAuth.error.value = errors.join('\n') + appStore.showError(t('admin.accounts.oauth.batchFailed')) + } + } finally { + grokOAuth.loading.value = false + } +} + // OpenAI OAuth 授权码兑换 const handleOpenAIExchange = async (authCode: string) => { const oauthClient = openaiOAuth diff --git a/frontend/src/components/account/OAuthAuthorizationFlow.vue b/frontend/src/components/account/OAuthAuthorizationFlow.vue index f45aeeb5d..98ae7fe51 100644 --- a/frontend/src/components/account/OAuthAuthorizationFlow.vue +++ b/frontend/src/components/account/OAuthAuthorizationFlow.vue @@ -59,6 +59,17 @@ t(getOAuthKey('ssoCookieAuth')) }} +