Files
sub2api/backend/internal/service/channel_plaza_test.go
T
李建琦 6d655c9903
Release / update-version (push) Has been cancelled
Release / build-frontend (push) Has been cancelled
Release / release (push) Has been cancelled
Release / sync-version-file (push) Has been cancelled
CI / shell (push) Canceled after 0s
CI / test (push) Canceled after 0s
CI / frontend (push) Canceled after 0s
CI / golangci-lint (push) Canceled after 0s
Security Scan / backend-security (push) Canceled after 0s
Security Scan / frontend-security (push) Canceled after 0s
Sub2API v1.0 - AI API 网关(二开初始版本,基于上游 Wei-Shaw/sub2api)
2026-08-21 18:30:13 +08:00

331 lines
13 KiB
Go

//go:build unit
package service
import (
"context"
"errors"
"testing"
"github.com/stretchr/testify/require"
)
// newPlazaChannelService 构造 ListPlazaGroups 测试用的 ChannelService。
func newPlazaChannelService(channels []Channel, groups []Group, pricing *PricingService) *ChannelService {
repo := &mockChannelRepository{
listAllFn: func(ctx context.Context) ([]Channel, error) { return channels, nil },
}
svc := NewChannelService(repo, &stubGroupRepoForAvailable{activeGroups: groups}, nil, nil)
svc.pricingService = pricing
return svc
}
func plazaPricedChannel(id int64, name string, groupIDs []int64, platform string, models ...string) Channel {
return Channel{
ID: id,
Name: name,
Status: StatusActive,
GroupIDs: groupIDs,
ModelPricing: []ChannelModelPricing{{
Platform: platform,
Models: models,
BillingMode: BillingModeToken,
InputPrice: testPtrFloat64(3e-6),
OutputPrice: testPtrFloat64(1.5e-5),
}},
}
}
func TestListPlazaGroups_GroupCentricAggregation(t *testing.T) {
// 两个渠道挂同一分组:模型并入同一 PlazaGroup;无模型的分组不返回。
channels := []Channel{
plazaPricedChannel(1, "chA", []int64{10}, "anthropic", "claude-sonnet"),
plazaPricedChannel(2, "chB", []int64{10}, "anthropic", "claude-opus"),
}
groups := []Group{
{ID: 10, Name: "g-main", Description: "desc", Platform: "anthropic", RateMultiplier: 1},
{ID: 20, Name: "g-empty", Platform: "anthropic", RateMultiplier: 0.5},
}
svc := newPlazaChannelService(channels, groups, nil)
out, err := svc.ListPlazaGroups(context.Background())
require.NoError(t, err)
require.Len(t, out, 1, "无模型的分组不应返回")
require.Equal(t, int64(10), out[0].ID)
require.Equal(t, "desc", out[0].Description)
require.Len(t, out[0].Models, 2)
// 组内模型按名称排序
require.Equal(t, "claude-opus", out[0].Models[0].Name)
require.Equal(t, "claude-sonnet", out[0].Models[1].Name)
}
func TestListPlazaGroups_DedupFirstWinsWithPricingUpgrade(t *testing.T) {
// 同名模型:先见者胜;仅当已存条目无定价而新条目有定价时升级替换。
unpriced := Channel{
ID: 1, Name: "alpha", Status: StatusActive, GroupIDs: []int64{10},
// mapping-only → SupportedModels 产出无定价条目
ModelMapping: map[string]map[string]string{
"anthropic": {"claude-sonnet": "claude-sonnet"},
},
}
priced := plazaPricedChannel(2, "beta", []int64{10}, "anthropic", "claude-sonnet")
groups := []Group{{ID: 10, Name: "g", Platform: "anthropic", RateMultiplier: 1}}
// alpha(无价)按名称序先于 beta(有价):先见者无价,应被有价条目升级。
svc := newPlazaChannelService([]Channel{priced, unpriced}, groups, nil)
out, err := svc.ListPlazaGroups(context.Background())
require.NoError(t, err)
require.Len(t, out, 1)
require.Len(t, out[0].Models, 1)
require.NotNil(t, out[0].Models[0].Pricing, "无价条目应被有价条目升级")
require.NotNil(t, out[0].Models[0].Pricing.InputPrice)
}
func TestListPlazaGroups_PlatformIsolation(t *testing.T) {
// 渠道同时有 anthropic/openai 定价,anthropic 分组只应看到 anthropic 模型。
ch := Channel{
ID: 1, Name: "multi", Status: StatusActive, GroupIDs: []int64{10, 20},
ModelPricing: []ChannelModelPricing{
{Platform: "anthropic", Models: []string{"claude-sonnet"}, InputPrice: testPtrFloat64(3e-6)},
{Platform: "openai", Models: []string{"gpt-5"}, InputPrice: testPtrFloat64(2e-6)},
},
}
groups := []Group{
{ID: 10, Name: "g-claude", Platform: "anthropic", RateMultiplier: 1},
{ID: 20, Name: "g-gpt", Platform: "openai", RateMultiplier: 1},
}
svc := newPlazaChannelService([]Channel{ch}, groups, nil)
out, err := svc.ListPlazaGroups(context.Background())
require.NoError(t, err)
require.Len(t, out, 2)
byName := map[string][]PlazaModel{}
for _, g := range out {
byName[g.Name] = g.Models
}
require.Len(t, byName["g-claude"], 1)
require.Equal(t, "claude-sonnet", byName["g-claude"][0].Name)
require.Len(t, byName["g-gpt"], 1)
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"
groups := []Group{{ID: 10, Name: "g", Platform: "anthropic", RateMultiplier: 1}}
svc := newPlazaChannelService([]Channel{inactive}, groups, nil)
out, err := svc.ListPlazaGroups(context.Background())
require.NoError(t, err)
require.Empty(t, out)
}
func TestListPlazaGroups_SortedByRateMultiplierAsc(t *testing.T) {
channels := []Channel{
plazaPricedChannel(1, "ch", []int64{10, 20, 30}, "anthropic", "claude-sonnet"),
}
groups := []Group{
{ID: 10, Name: "b-standard", Platform: "anthropic", RateMultiplier: 1},
{ID: 20, Name: "a-standard", Platform: "anthropic", RateMultiplier: 1},
{ID: 30, Name: "cheap", Platform: "anthropic", RateMultiplier: 0.5},
}
svc := newPlazaChannelService(channels, groups, nil)
out, err := svc.ListPlazaGroups(context.Background())
require.NoError(t, err)
require.Len(t, out, 3)
require.Equal(t, "cheap", out[0].Name, "倍率低者在前")
require.Equal(t, "a-standard", out[1].Name, "同倍率按名称")
require.Equal(t, "b-standard", out[2].Name)
}
func TestListPlazaGroups_OfficialPricingFill(t *testing.T) {
pricingSvc := newStubPricingServiceFromMap(map[string]*LiteLLMModelPricing{
"claude-sonnet": {
Mode: "chat",
InputCostPerToken: 3e-6,
OutputCostPerToken: 1.5e-5,
CacheCreationInputTokenCost: 3.75e-6,
CacheCreationInputTokenCostAbove1hr: 6e-6,
CacheReadInputTokenCost: 3e-7,
},
"token-absent": {Mode: "image_generation", TokenPricingAbsent: true, OutputCostPerImage: 0.04},
})
channels := []Channel{
plazaPricedChannel(1, "ch", []int64{10}, "anthropic", "claude-sonnet", "unknown-model", "token-absent"),
}
groups := []Group{{ID: 10, Name: "g", Platform: "anthropic", RateMultiplier: 1}}
svc := newPlazaChannelService(channels, groups, pricingSvc)
out, err := svc.ListPlazaGroups(context.Background())
require.NoError(t, err)
require.Len(t, out, 1)
require.Len(t, out[0].Models, 3)
byName := map[string]PlazaModel{}
for _, m := range out[0].Models {
byName[m.Name] = m
}
// 命中:填充完整官方价(含 1h 缓存写入)
official := byName["claude-sonnet"].OfficialPricing
require.NotNil(t, official)
require.InDelta(t, 3e-6, *official.InputPrice, 1e-12)
require.InDelta(t, 6e-6, *official.CacheWrite1hPrice, 1e-12)
require.InDelta(t, 3e-7, *official.CacheReadPrice, 1e-12)
// 未命中:nil(GetModelPricing 的 claude 系列模糊匹配对非 claude 名不生效)
require.Nil(t, byName["unknown-model"].OfficialPricing)
// TokenPricingAbsent 条目不作为官方 token 价展示
require.Nil(t, byName["token-absent"].OfficialPricing)
}
func TestListPlazaGroups_GroupImagePriceOverridesChannelPricing(t *testing.T) {
// 图片计费模型:档位价按实收口径合成(分组图片价 > 渠道档位价 > 渠道默认按次价),
// 分组独立倍率字段透传;未配图片价的分组保持渠道定价原样。
perReq := 0.2
tier4K := 0.3
imgPrice := 0.02
channels := []Channel{{
ID: 1, Name: "img-ch", Status: StatusActive, GroupIDs: []int64{10, 20},
ModelPricing: []ChannelModelPricing{{
Platform: "openai",
Models: []string{"gpt-image-2"},
BillingMode: BillingModeImage,
PerRequestPrice: &perReq,
Intervals: []PricingInterval{{TierLabel: "4K", PerRequestPrice: &tier4K}},
}},
}}
groups := []Group{
{ID: 10, Name: "g-media", Platform: "openai", RateMultiplier: 1,
ImagePrice1K: &imgPrice, ImageRateIndependent: true, ImageRateMultiplier: 1},
{ID: 20, Name: "g-plain", Platform: "openai", RateMultiplier: 0.1},
}
svc := newPlazaChannelService(channels, groups, nil)
out, err := svc.ListPlazaGroups(context.Background())
require.NoError(t, err)
require.Len(t, out, 2)
byName := map[string]PlazaGroup{}
for _, g := range out {
byName[g.Name] = g
}
media := byName["g-media"]
require.True(t, media.ImageRateIndependent)
require.InDelta(t, 1.0, media.ImageRateMultiplier, 1e-9)
require.Len(t, media.Models, 1)
p := media.Models[0].Pricing
require.NotNil(t, p)
require.Len(t, p.Intervals, 3)
tierPrices := map[string]float64{}
for _, iv := range p.Intervals {
require.NotNil(t, iv.PerRequestPrice)
tierPrices[iv.TierLabel] = *iv.PerRequestPrice
}
require.InDelta(t, 0.02, tierPrices["1K"], 1e-9, "1K 用分组图片价")
require.InDelta(t, 0.2, tierPrices["2K"], 1e-9, "2K 分组未配,回落渠道默认按次价")
require.InDelta(t, 0.3, tierPrices["4K"], 1e-9, "4K 分组未配,回落渠道档位价")
plain := byName["g-plain"]
require.False(t, plain.ImageRateIndependent)
require.Len(t, plain.Models, 1)
pp := plain.Models[0].Pricing
require.NotNil(t, pp)
require.Len(t, pp.Intervals, 1, "未配分组图片价:渠道定价原样")
require.InDelta(t, 0.2, *pp.PerRequestPrice, 1e-9)
// 合成为克隆,渠道原始定价不被修改
require.Len(t, channels[0].ModelPricing[0].Intervals, 1)
}
func TestListPlazaGroups_GroupImagePriceIgnoredForNonImageModes(t *testing.T) {
// token 模式定价不受分组图片价影响。
imgPrice := 0.02
channels := []Channel{plazaPricedChannel(1, "ch", []int64{10}, "openai", "gpt-5")}
groups := []Group{{ID: 10, Name: "g", Platform: "openai", RateMultiplier: 1, ImagePrice1K: &imgPrice}}
svc := newPlazaChannelService(channels, groups, nil)
out, err := svc.ListPlazaGroups(context.Background())
require.NoError(t, err)
require.Len(t, out, 1)
p := out[0].Models[0].Pricing
require.NotNil(t, p)
require.Empty(t, p.Intervals)
require.NotNil(t, p.InputPrice)
require.Nil(t, p.PerRequestPrice)
}
func TestListPlazaGroups_RepoErrorsPropagate(t *testing.T) {
sentinel := errors.New("boom")
repo := &mockChannelRepository{
listAllFn: func(ctx context.Context) ([]Channel, error) { return nil, sentinel },
}
svc := NewChannelService(repo, &stubGroupRepoForAvailable{}, nil, nil)
out, err := svc.ListPlazaGroups(context.Background())
require.Nil(t, out)
require.ErrorIs(t, err, sentinel)
svc2 := NewChannelService(
&mockChannelRepository{listAllFn: func(ctx context.Context) ([]Channel, error) { return nil, nil }},
&stubGroupRepoForAvailable{listActiveErr: sentinel},
nil, nil,
)
out2, err2 := svc2.ListPlazaGroups(context.Background())
require.Nil(t, out2)
require.ErrorIs(t, err2, sentinel)
}