Files
sub2api/backend/internal/handler/grok_media_test.go
T
shaw 14608dc6d4 Merge origin/main into codex/secure-protected-video-content-4498
Reconcile with #4539 (grok media account model mapping), now on main:
- handler/grok_media.go non-failover error path keeps both changes —
  #4539's grokMediaScheduleModel(account, routingModel, nil) schedule
  attribution and this branch's IsResponseCommitted guard
- auto-merged sections verified: routing/classify use #4539's routingModel,
  video lookup owner-binding and no-failover semantics intact, ForwardGrokMedia
  keeps mapping block (skipped for lookup endpoints via RequiresRequestBody),
  empty-image failover, and video-status URL rewrite in order
2026-07-18 21:38:17 +08:00

168 lines
5.7 KiB
Go

package handler
import (
"context"
"errors"
"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 records usage",
endpoint: service.GrokMediaEndpointVideosGenerations,
model: "grok-imagine-video-1.5",
want: true,
},
{
name: "video status skips empty model usage",
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) {
require.Equal(t, tt.want, shouldRecordGrokMediaUsage(tt.endpoint, tt.model))
})
}
}
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)
})
}