Files
sub2api/backend/internal/handler/grok_media_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

180 lines
6.4 KiB
Go

package handler
import (
"context"
"errors"
"strings"
"testing"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/stretchr/testify/require"
)
type grokMediaEligibilityProberStub struct {
eligible bool
reason string
err error
calls int
}
func (s *grokMediaEligibilityProberStub) ProbeMediaEligibility(context.Context, int64) (bool, string, error) {
s.calls++
return s.eligible, s.reason, s.err
}
func TestShouldRecordGrokMediaUsage(t *testing.T) {
tests := []struct {
name string
endpoint service.GrokMediaEndpoint
model string
want bool
}{
{
name: "image generation records usage",
endpoint: service.GrokMediaEndpointImagesGenerations,
model: "grok-imagine",
want: true,
},
{
name: "image edit records usage",
endpoint: service.GrokMediaEndpointImagesEdits,
model: "grok-imagine-edit",
want: true,
},
{
name: "video generation defers usage until status",
endpoint: service.GrokMediaEndpointVideosGenerations,
model: "grok-imagine-video-1.5",
want: false,
},
{
name: "video status skips immediate helper (status path claims separately)",
endpoint: service.GrokMediaEndpointVideoStatus,
model: "",
want: false,
},
{
name: "video content skips usage",
endpoint: service.GrokMediaEndpointVideoContent,
model: "",
want: false,
},
{
name: "generation skips usage without model",
endpoint: service.GrokMediaEndpointImagesGenerations,
model: " ",
want: false,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
// Nil result must never bill.
require.False(t, shouldRecordGrokMediaUsage(tt.endpoint, tt.model, nil))
// Immediate helper only bills image generation (async video bills on status).
result := &service.OpenAIForwardResult{ImageCount: 1, VideoCount: 0}
if tt.endpoint.IsGenerationRequest() && !isGrokVideoCreateEndpoint(tt.endpoint) && 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))
})
}
}
func TestGrokMediaRequiredCapability(t *testing.T) {
tests := []struct {
name string
endpoint service.GrokMediaEndpoint
want service.OpenAIEndpointCapability
}{
{name: "image generation", endpoint: service.GrokMediaEndpointImagesGenerations, want: service.OpenAIEndpointCapabilityGrokMediaGeneration},
{name: "image edit", endpoint: service.GrokMediaEndpointImagesEdits, want: service.OpenAIEndpointCapabilityGrokMediaGeneration},
{name: "video generation", endpoint: service.GrokMediaEndpointVideosGenerations, want: service.OpenAIEndpointCapabilityGrokMediaGeneration},
{name: "video edit", endpoint: service.GrokMediaEndpointVideosEdits, want: service.OpenAIEndpointCapabilityGrokMediaGeneration},
{name: "video extension", endpoint: service.GrokMediaEndpointVideosExtensions, want: service.OpenAIEndpointCapabilityGrokMediaGeneration},
{name: "video status preserves lookup", endpoint: service.GrokMediaEndpointVideoStatus, want: ""},
{name: "video content preserves lookup", endpoint: service.GrokMediaEndpointVideoContent, want: ""},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
require.Equal(t, tt.want, grokMediaRequiredCapability(tt.endpoint))
})
}
}
func TestGrokMediaScheduleModelUsesNormalizedMappedUpstream(t *testing.T) {
account := &service.Account{
Platform: service.PlatformGrok,
Credentials: map[string]any{
"model_mapping": map[string]any{
"grok-imagine-video-1.5": "wrong-raw-model",
"grok-imagine-video": "mapped-video-model",
},
},
}
require.Equal(t, "mapped-video-model", grokMediaScheduleModel(account, "grok-imagine-video", nil))
require.Equal(t, "actual-upstream-model", grokMediaScheduleModel(account, "grok-imagine-video", &service.OpenAIForwardResult{
UpstreamModel: "actual-upstream-model",
}))
require.Equal(t, "mapped-video-model", grokMediaScheduleModel(account, "grok-imagine-video", &service.OpenAIForwardResult{}))
require.Equal(t, "grok-imagine-video", grokMediaScheduleModel(nil, " grok-imagine-video ", nil))
}
func TestEnsureGrokMediaAccountEligibility(t *testing.T) {
t.Run("non oauth account does not probe", func(t *testing.T) {
prober := &grokMediaEligibilityProberStub{}
h := &OpenAIGatewayHandler{grokMediaEligibilityProber: prober}
account := &service.Account{Platform: service.PlatformGrok, Type: service.AccountTypeAPIKey}
eligible, reason, err := h.ensureGrokMediaAccountEligibility(context.Background(), account)
require.NoError(t, err)
require.True(t, eligible)
require.Equal(t, "non_oauth", reason)
require.Zero(t, prober.calls)
})
t.Run("unobserved oauth is probed before forwarding", func(t *testing.T) {
prober := &grokMediaEligibilityProberStub{eligible: true, reason: "eligible"}
h := &OpenAIGatewayHandler{grokMediaEligibilityProber: prober}
account := &service.Account{ID: 7, Platform: service.PlatformGrok, Type: service.AccountTypeOAuth}
eligible, reason, err := h.ensureGrokMediaAccountEligibility(context.Background(), account)
require.NoError(t, err)
require.True(t, eligible)
require.Equal(t, "eligible", reason)
require.Equal(t, 1, prober.calls)
})
t.Run("missing prober fails closed", func(t *testing.T) {
h := &OpenAIGatewayHandler{}
account := &service.Account{ID: 8, Platform: service.PlatformGrok, Type: service.AccountTypeOAuth}
eligible, reason, err := h.ensureGrokMediaAccountEligibility(context.Background(), account)
require.Error(t, err)
require.False(t, eligible)
require.Equal(t, "billing_probe_unavailable", reason)
})
t.Run("probe failure fails closed", func(t *testing.T) {
probeErr := errors.New("probe failed")
prober := &grokMediaEligibilityProberStub{reason: "billing_unobserved", err: probeErr}
h := &OpenAIGatewayHandler{grokMediaEligibilityProber: prober}
account := &service.Account{ID: 9, Platform: service.PlatformGrok, Type: service.AccountTypeOAuth}
eligible, reason, err := h.ensureGrokMediaAccountEligibility(context.Background(), account)
require.ErrorIs(t, err, probeErr)
require.False(t, eligible)
require.Equal(t, "billing_unobserved", reason)
})
}