From ca2b7b9bb8feadd03c74b0c35759193d69b9fa50 Mon Sep 17 00:00:00 2001 From: htoday1 <2773851718@qq.com> Date: Sat, 11 Jul 2026 01:04:46 +0800 Subject: [PATCH] fix: apply matched provider group billing multiplier --- src/app/v1/_lib/proxy/provider-selector.ts | 68 ++++--- src/lib/utils/provider-group.test.ts | 30 +++ src/lib/utils/provider-group.ts | 30 +++ ...r-selector-select-provider-by-type.test.ts | 171 ++++++++++++++++++ 4 files changed, 278 insertions(+), 21 deletions(-) diff --git a/src/app/v1/_lib/proxy/provider-selector.ts b/src/app/v1/_lib/proxy/provider-selector.ts index 4dc42b9f0..90224ccf5 100644 --- a/src/app/v1/_lib/proxy/provider-selector.ts +++ b/src/app/v1/_lib/proxy/provider-selector.ts @@ -4,7 +4,11 @@ import { PROVIDER_GROUP } from "@/lib/constants/provider.constants"; import { logger } from "@/lib/logger"; import { RateLimitService } from "@/lib/rate-limit"; import { SessionManager } from "@/lib/session-manager"; -import { parseProviderGroups, resolveProviderGroupsWithDefault } from "@/lib/utils/provider-group"; +import { + parseProviderGroups, + resolveBillingProviderGroups, + resolveProviderGroupsWithDefault, +} from "@/lib/utils/provider-group"; import { isProviderActiveNow } from "@/lib/utils/provider-schedule"; import { resolveSystemTimezone } from "@/lib/utils/timezone"; import { isVendorTypeCircuitOpen } from "@/lib/vendor-type-circuit-breaker"; @@ -63,6 +67,46 @@ function checkProviderGroupMatch(providerGroupTag: string | null, userGroups: st return providerTags.some((tag) => groups.includes(tag)); } +async function resolveGroupCostMultiplierForProvider(session: ProxySession): Promise { + const effectiveGroup = getEffectiveProviderGroup(session); + const provider = session.provider; + + if (!effectiveGroup || !provider) { + session.setGroupCostMultiplier(1.0); + return; + } + + const billingGroups = resolveBillingProviderGroups(provider.groupTag, effectiveGroup); + if (billingGroups.length === 0) { + logger.warn( + "[ProviderResolver] Selected provider has no billing group intersection, falling back to 1.0", + { + providerId: provider.id, + providerName: provider.name, + providerGroups: provider.groupTag, + effectiveGroup, + } + ); + session.setGroupCostMultiplier(1.0); + return; + } + + const billingGroup = billingGroups.join(","); + + try { + const multiplier = await getGroupCostMultiplier(billingGroup); + session.setGroupCostMultiplier(multiplier); + } catch (error) { + logger.warn("[ProviderResolver] Failed to resolve group cost multiplier, falling back to 1.0", { + billingGroup, + effectiveGroup, + providerId: provider.id, + error: error instanceof Error ? error.message : String(error), + }); + session.setGroupCostMultiplier(1.0); + } +} + /** * 检查供应商是否支持指定模型(用于调度器匹配) * @@ -186,26 +230,6 @@ export class ProxyProviderResolver { session.setLastSelectionContext(context); // 保存用于后续记录 } - // === Resolve group cost multiplier === - // Fail soft: if the lookup throws (Redis/DB hiccup), fall back to 1.0 so - // request handling proceeds without billing disruption. - const effectiveGroup = getEffectiveProviderGroup(session); - if (effectiveGroup) { - try { - const multiplier = await getGroupCostMultiplier(effectiveGroup); - session.setGroupCostMultiplier(multiplier); - } catch (error) { - logger.warn( - "[ProviderResolver] Failed to resolve group cost multiplier, falling back to 1.0", - { - effectiveGroup, - error: error instanceof Error ? error.message : String(error), - } - ); - session.setGroupCostMultiplier(1.0); - } - } - // === 故障转移循环 === let attemptCount = 0; while (true) { @@ -341,11 +365,13 @@ export class ProxyProviderResolver { // 修复:延迟到 forwarder 请求成功后统一更新(见 forwarder.ts:75-80) // void SessionManager.updateSessionProvider(...); // ❌ 已移除 + await resolveGroupCostMultiplierForProvider(session); return null; // 成功 } // sessionId 为空的情况(理论上不应该发生) logger.warn("ProviderSelector: sessionId is null, skipping concurrent check"); + await resolveGroupCostMultiplierForProvider(session); return null; } diff --git a/src/lib/utils/provider-group.test.ts b/src/lib/utils/provider-group.test.ts index 90a68b940..828ddf979 100644 --- a/src/lib/utils/provider-group.test.ts +++ b/src/lib/utils/provider-group.test.ts @@ -3,6 +3,7 @@ import { normalizeProviderGroup, normalizeProviderGroupTag, parseProviderGroups, + resolveBillingProviderGroups, resolveProviderGroupsWithDefault, } from "./provider-group"; @@ -33,4 +34,33 @@ describe("provider-group utils", () => { expect(parseProviderGroups(null)).toEqual([]); expect(parseProviderGroups(" ")).toEqual([]); }); + + test("计费分组应取用户分组与已选供应商标签的交集", () => { + expect( + resolveBillingProviderGroups("cus_gpt,gpt_test", "cus_claude_pro,cus_grok,gpt_test,mimo") + ).toEqual(["gpt_test"]); + }); + + test("计费分组应保留用户分组声明顺序", () => { + expect(resolveBillingProviderGroups("group-b,group-a", "group-a,group-b")).toEqual([ + "group-a", + "group-b", + ]); + }); + + test("通配分组应按供应商标签解析倍率", () => { + expect(resolveBillingProviderGroups("group-b,group-a", "*")).toEqual(["group-b", "group-a"]); + }); + + test("显式匹配分组应优先于通配分组", () => { + expect(resolveBillingProviderGroups("group-b,group-a", "*,group-a")).toEqual(["group-a"]); + }); + + test("未分组供应商在通配访问下应使用 default 计费分组", () => { + expect(resolveBillingProviderGroups(null, "*")).toEqual(["default"]); + }); + + test("无交集且无通配权限时不应选择无关分组倍率", () => { + expect(resolveBillingProviderGroups("group-b", "group-a")).toEqual([]); + }); }); diff --git a/src/lib/utils/provider-group.ts b/src/lib/utils/provider-group.ts index 6ed403fc1..496cf348c 100644 --- a/src/lib/utils/provider-group.ts +++ b/src/lib/utils/provider-group.ts @@ -59,3 +59,33 @@ export function resolveProviderGroupsWithDefault(value: unknown): string[] { return groups; } + +/** + * Resolve the provider groups that should participate in billing. + * + * Explicit user/key groups are intersected with the selected provider's tags, + * preserving the user/key declaration order. A wildcard grants access to all + * providers but only falls back to the provider's own tag order when there is + * no explicit matching group. + */ +export function resolveBillingProviderGroups( + providerGroupTag: unknown, + userGroupsValue: unknown +): string[] { + const providerGroups = resolveProviderGroupsWithDefault(providerGroupTag); + const userGroups = parseProviderGroups(userGroupsValue); + const providerGroupSet = new Set(providerGroups); + + const explicitMatches = userGroups.filter( + (group) => group !== PROVIDER_GROUP.ALL && providerGroupSet.has(group) + ); + if (explicitMatches.length > 0) { + return explicitMatches; + } + + if (userGroups.includes(PROVIDER_GROUP.ALL)) { + return providerGroups; + } + + return []; +} diff --git a/tests/unit/proxy/provider-selector-select-provider-by-type.test.ts b/tests/unit/proxy/provider-selector-select-provider-by-type.test.ts index 85caea6cb..284bae735 100644 --- a/tests/unit/proxy/provider-selector-select-provider-by-type.test.ts +++ b/tests/unit/proxy/provider-selector-select-provider-by-type.test.ts @@ -3,6 +3,8 @@ import type { Provider } from "@/types/provider"; import { ProxyProviderResolver } from "@/app/v1/_lib/proxy/provider-selector"; const findAllProvidersMock = vi.hoisted(() => vi.fn<[], Promise>()); +const getGroupCostMultiplierMock = vi.hoisted(() => vi.fn()); +const checkAndTrackProviderSessionMock = vi.hoisted(() => vi.fn()); vi.mock("@/repository/provider", () => { return { @@ -11,6 +13,16 @@ vi.mock("@/repository/provider", () => { }; }); +vi.mock("@/repository/provider-groups", () => ({ + getGroupCostMultiplier: getGroupCostMultiplierMock, +})); + +vi.mock("@/lib/rate-limit", () => ({ + RateLimitService: { + checkAndTrackProviderSession: checkAndTrackProviderSessionMock, + }, +})); + describe("ProxyProviderResolver.selectProviderByType - /v1/models 分组隔离", () => { beforeEach(() => { vi.clearAllMocks(); @@ -89,3 +101,162 @@ describe("ProxyProviderResolver.selectProviderByType - /v1/models 分组隔离", expect(provider?.id).toBe(inGroup.id); }); }); + +describe("ProxyProviderResolver.ensure - 分组倍率", () => { + beforeEach(() => { + vi.clearAllMocks(); + }); + + test("按当前供应商与 Key 分组交集解析倍率", async () => { + const provider = { + id: 56, + name: "gpt-test-provider", + isEnabled: true, + providerType: "openai-compatible", + groupTag: "cus_gpt,gpt_test", + weight: 1, + priority: 0, + costMultiplier: 1, + limitConcurrentSessions: 0, + } as unknown as Provider; + + getGroupCostMultiplierMock.mockResolvedValueOnce(10); + + const findReusableSpy = vi + .spyOn(ProxyProviderResolver as never, "findReusable" as never) + .mockResolvedValue(null as never); + const pickRandomProviderSpy = vi + .spyOn(ProxyProviderResolver as never, "pickRandomProvider" as never) + .mockResolvedValue({ + provider, + context: { + totalProviders: 1, + enabledProviders: 1, + targetType: "openai-compatible", + requestedModel: "gpt-5.5", + groupFilterApplied: true, + userGroup: "cus_claude_pro,cus_grok,gpt_test,mimo", + beforeHealthCheck: 1, + afterHealthCheck: 1, + priorityLevels: [0], + selectedPriority: 0, + candidatesAtPriority: [], + }, + } as never); + + const setGroupCostMultiplier = vi.fn(); + const session = { + provider: null as Provider | null, + sessionId: null, + authState: { + user: { providerGroup: "cus_claude_pro,cus_grok,gpt_test,mimo" }, + key: { providerGroup: "cus_claude_pro,cus_grok,gpt_test,mimo" }, + }, + setProvider(selected: Provider | null) { + this.provider = selected; + }, + setLastSelectionContext: vi.fn(), + getLastSelectionContext: vi.fn(() => null), + setGroupCostMultiplier, + addProviderToChain: vi.fn(), + getOriginalModel: vi.fn(() => "gpt-5.5"), + } as unknown as Parameters[0]; + + try { + await expect(ProxyProviderResolver.ensure(session)).resolves.toBeNull(); + } finally { + findReusableSpy.mockRestore(); + pickRandomProviderSpy.mockRestore(); + } + + expect(getGroupCostMultiplierMock).toHaveBeenCalledWith("gpt_test"); + expect(setGroupCostMultiplier).toHaveBeenCalledWith(10); + }); + + test("故障切换后应按最终供应商重新解析倍率", async () => { + const firstProvider = { + id: 1, + name: "group-a-provider", + providerType: "openai-compatible", + groupTag: "group-a", + weight: 1, + priority: 0, + costMultiplier: 1, + limitConcurrentSessions: 1, + } as unknown as Provider; + const fallbackProvider = { + ...firstProvider, + id: 2, + name: "group-b-provider", + groupTag: "group-b", + limitConcurrentSessions: 0, + } as Provider; + + getGroupCostMultiplierMock.mockResolvedValueOnce(10); + checkAndTrackProviderSessionMock + .mockResolvedValueOnce({ + allowed: false, + count: 1, + referenced: false, + reason: "limit reached", + }) + .mockResolvedValueOnce({ + allowed: true, + count: 1, + referenced: false, + }); + + const context = { + totalProviders: 2, + enabledProviders: 2, + targetType: "openai-compatible", + requestedModel: "gpt-5.5", + groupFilterApplied: true, + userGroup: "group-a,group-b", + beforeHealthCheck: 2, + afterHealthCheck: 2, + priorityLevels: [0], + selectedPriority: 0, + candidatesAtPriority: [], + }; + + const findReusableSpy = vi + .spyOn(ProxyProviderResolver as never, "findReusable" as never) + .mockResolvedValue(null as never); + const pickRandomProviderSpy = vi + .spyOn(ProxyProviderResolver as never, "pickRandomProvider" as never) + .mockResolvedValueOnce({ provider: firstProvider, context } as never) + .mockResolvedValueOnce({ provider: fallbackProvider, context } as never); + + const setGroupCostMultiplier = vi.fn(); + const session = { + provider: null as Provider | null, + sessionId: "session-1", + authState: { + user: { providerGroup: "group-a,group-b" }, + key: { providerGroup: "group-a,group-b" }, + }, + setProvider(selected: Provider | null) { + this.provider = selected; + }, + setLastSelectionContext: vi.fn(), + getLastSelectionContext: vi.fn(() => context), + setGroupCostMultiplier, + addProviderToChain: vi.fn(), + getOriginalModel: vi.fn(() => "gpt-5.5"), + recordProviderSessionRef: vi.fn(), + } as unknown as Parameters[0]; + + try { + await expect(ProxyProviderResolver.ensure(session)).resolves.toBeNull(); + } finally { + findReusableSpy.mockRestore(); + pickRandomProviderSpy.mockRestore(); + } + + expect(checkAndTrackProviderSessionMock).toHaveBeenCalledTimes(2); + expect(getGroupCostMultiplierMock).toHaveBeenCalledTimes(1); + expect(getGroupCostMultiplierMock).toHaveBeenCalledWith("group-b"); + expect(setGroupCostMultiplier).toHaveBeenCalledWith(10); + }); +});