Files
sub2api/backend/internal/service/grok_media_content_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

345 lines
14 KiB
Go

//go:build unit
package service
import (
"context"
"io"
"net/http"
"net/http/httptest"
"strings"
"testing"
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/Wei-Shaw/sub2api/internal/pkg/tlsfingerprint"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
"github.com/tidwall/gjson"
)
type grokMediaContentUpstreamStub struct {
request *http.Request
requests []*http.Request
response *http.Response
responses []*http.Response
}
func (s *grokMediaContentUpstreamStub) Do(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) {
s.request = req
s.requests = append(s.requests, req)
if len(s.responses) > 0 {
resp := s.responses[0]
s.responses = s.responses[1:]
return resp, nil
}
return s.response, nil
}
func (s *grokMediaContentUpstreamStub) DoWithTLS(req *http.Request, proxyURL string, accountID int64, accountConcurrency int, _ *tlsfingerprint.Profile) (*http.Response, error) {
return s.Do(req, proxyURL, accountID, accountConcurrency)
}
func grokMediaContentTestAccount() *Account {
return &Account{
ID: 9,
Platform: PlatformGrok,
Type: AccountTypeAPIKey,
Credentials: map[string]any{
"api_key": "upstream-key",
"base_url": "https://relay.example/v1",
},
}
}
func grokMediaContentTestContext(method, target string, headers map[string]string) (*gin.Context, *httptest.ResponseRecorder) {
gin.SetMode(gin.TestMode)
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
c.Request = httptest.NewRequest(method, target, nil)
for name, value := range headers {
c.Request.Header.Set(name, value)
}
return c, recorder
}
func grokMediaContentStatusResponse(body string) *http.Response {
return &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"application/json"}},
Body: io.NopCloser(strings.NewReader(body)),
}
}
func TestForwardGrokMediaContentUsesUpstreamCredentialAndStreamsRange(t *testing.T) {
upstream := &grokMediaContentUpstreamStub{
responses: []*http.Response{grokMediaContentStatusResponse(`{"status":"completed"}`), {
StatusCode: http.StatusPartialContent,
Header: http.Header{
"Content-Type": []string{"video/mp4"},
"Content-Length": []string{"13"},
"Content-Range": []string{"bytes 0-12/100"},
"Accept-Ranges": []string{"bytes"},
"Content-Disposition": []string{
`attachment; filename="task-1.mp4"`,
},
},
Body: io.NopCloser(strings.NewReader("video-payload")),
}},
}
svc := &OpenAIGatewayService{cfg: &config.Config{}, httpUpstream: upstream}
c, recorder := grokMediaContentTestContext(http.MethodGet, "https://api.example/v1/videos/task-1/content", map[string]string{
"Range": "bytes=0-12",
})
result, err := svc.ForwardGrokMedia(
context.Background(), c, grokMediaContentTestAccount(),
GrokMediaEndpointVideoContent, "task-1", nil, "",
)
require.NoError(t, err)
require.NotNil(t, result)
require.Equal(t, http.StatusPartialContent, recorder.Code)
require.Equal(t, "video-payload", recorder.Body.String())
require.Len(t, upstream.requests, 2)
require.Equal(t, "https://relay.example/v1/videos/task-1", upstream.requests[0].URL.String())
require.Equal(t, "Bearer upstream-key", upstream.requests[0].Header.Get("Authorization"))
require.Equal(t, "https://relay.example/v1/videos/task-1/content", upstream.requests[1].URL.String())
require.Equal(t, "Bearer upstream-key", upstream.requests[1].Header.Get("Authorization"))
require.Equal(t, "bytes=0-12", upstream.requests[1].Header.Get("Range"))
require.Equal(t, "*/*", upstream.requests[1].Header.Get("Accept"))
require.Equal(t, "video/mp4", recorder.Header().Get("Content-Type"))
require.Equal(t, "13", recorder.Header().Get("Content-Length"))
require.Equal(t, "bytes 0-12/100", recorder.Header().Get("Content-Range"))
require.Equal(t, "bytes", recorder.Header().Get("Accept-Ranges"))
require.Equal(t, `attachment; filename="task-1.mp4"`, recorder.Header().Get("Content-Disposition"))
require.True(t, IsResponseCommitted(c))
}
func TestForwardGrokMediaContentStreamsFullResponseWithSafeDefaults(t *testing.T) {
upstream := &grokMediaContentUpstreamStub{
responses: []*http.Response{grokMediaContentStatusResponse(`{"status":"completed"}`), {
StatusCode: http.StatusOK,
Header: http.Header{"Set-Cookie": []string{"secret=upstream"}, "X-Upstream-Secret": []string{"hidden"}},
Body: io.NopCloser(strings.NewReader("full-video")),
ContentLength: -1,
}},
}
svc := &OpenAIGatewayService{cfg: &config.Config{}, httpUpstream: upstream}
c, recorder := grokMediaContentTestContext(http.MethodGet, "https://api.example/v1/videos/task-1/content", nil)
_, err := svc.ForwardGrokMedia(
context.Background(), c, grokMediaContentTestAccount(),
GrokMediaEndpointVideoContent, "task-1", nil, "",
)
require.NoError(t, err)
require.Equal(t, http.StatusOK, recorder.Code)
require.Equal(t, "full-video", recorder.Body.String())
require.Len(t, upstream.requests, 2)
require.Empty(t, upstream.requests[1].Header.Get("Range"))
require.Equal(t, "application/octet-stream", recorder.Header().Get("Content-Type"))
require.Empty(t, recorder.Header().Get("Content-Length"))
require.Empty(t, recorder.Header().Get("Set-Cookie"))
require.Empty(t, recorder.Header().Get("X-Upstream-Secret"))
require.True(t, IsResponseCommitted(c))
}
func TestForwardGrokMediaContentPreservesRangeNotSatisfiable(t *testing.T) {
upstream := &grokMediaContentUpstreamStub{
responses: []*http.Response{grokMediaContentStatusResponse(`{"status":"completed"}`), {
StatusCode: http.StatusRequestedRangeNotSatisfiable,
Header: http.Header{
"Content-Type": []string{"text/plain"},
"Content-Length": []string{"11"},
"Content-Range": []string{"bytes */100"},
"Accept-Ranges": []string{"bytes"},
},
Body: io.NopCloser(strings.NewReader("bad-range!!")),
}},
}
svc := &OpenAIGatewayService{cfg: &config.Config{}, httpUpstream: upstream}
c, recorder := grokMediaContentTestContext(http.MethodGet, "https://api.example/v1/videos/task-1/content", map[string]string{
"Range": "bytes=500-600",
})
_, err := svc.ForwardGrokMedia(
context.Background(), c, grokMediaContentTestAccount(),
GrokMediaEndpointVideoContent, "task-1", nil, "",
)
require.NoError(t, err)
require.Equal(t, http.StatusRequestedRangeNotSatisfiable, recorder.Code)
require.Equal(t, "bad-range!!", recorder.Body.String())
require.Len(t, upstream.requests, 2)
require.Equal(t, "bytes=500-600", upstream.requests[1].Header.Get("Range"))
require.Equal(t, "bytes */100", recorder.Header().Get("Content-Range"))
require.Equal(t, "bytes", recorder.Header().Get("Accept-Ranges"))
require.True(t, IsResponseCommitted(c))
}
func TestForwardGrokMediaContentFetchesValidatedSignedURLWithoutCredentials(t *testing.T) {
upstream := &grokMediaContentUpstreamStub{
responses: []*http.Response{
grokMediaContentStatusResponse(`{"status":"done","video":{"url":"https://vidgen.x.ai/signed-token/xai-video-task-1.mp4"}}`),
{
StatusCode: http.StatusPartialContent,
Header: http.Header{
"Content-Type": []string{"video/mp4"},
"Content-Length": []string{"13"},
"Content-Range": []string{"bytes 0-12/100"},
},
Body: io.NopCloser(strings.NewReader("video-payload")),
},
},
}
account := grokMediaContentTestAccount()
account.Credentials[credKeyHeaderOverrideEnabled] = true
account.Credentials[credKeyHeaderOverrides] = map[string]any{"user-agent": "private-agent"}
svc := &OpenAIGatewayService{cfg: &config.Config{}, httpUpstream: upstream}
c, recorder := grokMediaContentTestContext(http.MethodGet, "https://api.example/v1/videos/task-1/content", map[string]string{
"Range": "bytes=0-12",
})
_, err := svc.ForwardGrokMedia(
context.Background(), c, account,
GrokMediaEndpointVideoContent, "task-1", nil, "",
)
require.NoError(t, err)
require.Equal(t, http.StatusPartialContent, recorder.Code)
require.Equal(t, "video-payload", recorder.Body.String())
require.Len(t, upstream.requests, 2)
require.Equal(t, "https://relay.example/v1/videos/task-1", upstream.requests[0].URL.String())
require.Equal(t, "Bearer upstream-key", upstream.requests[0].Header.Get("Authorization"))
require.Equal(t, "private-agent", upstream.requests[0].Header.Get("User-Agent"))
require.True(t, HTTPUpstreamRedirectsDisabled(upstream.requests[0].Context()))
require.Equal(t, "https://vidgen.x.ai/signed-token/xai-video-task-1.mp4", upstream.requests[1].URL.String())
require.Empty(t, upstream.requests[1].Header.Get("Authorization"))
require.Empty(t, upstream.requests[1].Header.Get("User-Agent"))
require.Equal(t, "bytes=0-12", upstream.requests[1].Header.Get("Range"))
require.True(t, HTTPUpstreamRedirectsDisabled(upstream.requests[1].Context()))
}
func TestForwardGrokMediaContentFollowsAuthenticatedSub2APIRelay(t *testing.T) {
for _, statusURL := range []string{
`/v1/videos/task-1/content`,
`https://relay.example/v1/videos/task-1/content`,
} {
t.Run(statusURL, func(t *testing.T) {
upstream := &grokMediaContentUpstreamStub{
responses: []*http.Response{
grokMediaContentStatusResponse(`{"status":"completed","video":{"url":"` + statusURL + `"}}`),
{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"video/mp4"}},
Body: io.NopCloser(strings.NewReader("video-payload")),
},
},
}
svc := &OpenAIGatewayService{cfg: &config.Config{}, httpUpstream: upstream}
c, recorder := grokMediaContentTestContext(http.MethodGet, "https://api.example/v1/videos/task-1/content", nil)
_, err := svc.ForwardGrokMedia(
context.Background(), c, grokMediaContentTestAccount(),
GrokMediaEndpointVideoContent, "task-1", nil, "",
)
require.NoError(t, err)
require.Equal(t, http.StatusOK, recorder.Code)
require.Equal(t, "video-payload", recorder.Body.String())
require.Len(t, upstream.requests, 2)
require.Equal(t, "https://relay.example/v1/videos/task-1/content", upstream.requests[1].URL.String())
require.Equal(t, "Bearer upstream-key", upstream.requests[1].Header.Get("Authorization"))
})
}
}
func TestForwardGrokMediaContentRejectsUntrustedSignedURL(t *testing.T) {
upstream := &grokMediaContentUpstreamStub{
responses: []*http.Response{
grokMediaContentStatusResponse(`{"status":"done","video":{"url":"http://169.` + `254.169.254/latest/meta-data"}}`),
},
}
svc := &OpenAIGatewayService{cfg: &config.Config{}, httpUpstream: upstream}
c, _ := grokMediaContentTestContext(http.MethodGet, "https://api.example/v1/videos/task-1/content", nil)
_, err := svc.ForwardGrokMedia(
context.Background(), c, grokMediaContentTestAccount(),
GrokMediaEndpointVideoContent, "task-1", nil, "",
)
require.ErrorContains(t, err, "unsupported video content URL")
require.Len(t, upstream.requests, 1)
}
func TestGrokMediaSignedVideoContentURLRejectsDeceptiveOrigins(t *testing.T) {
for _, rawURL := range []string{
"https://vidgen.x.ai.attacker.invalid/video.mp4",
"https://vidgen.x.ai" + "@attacker.invalid/video.mp4",
"https://vidgen.x.ai:444/video.mp4",
"http://vidgen.x.ai/video.mp4",
} {
t.Run(rawURL, func(t *testing.T) {
_, err := grokMediaSignedVideoContentURL([]byte(`{"video":{"url":"`+rawURL+`"}}`), "task-1")
require.ErrorContains(t, err, "unsupported video content URL")
})
}
}
func TestGrokMediaSignedVideoContentURLRejectsDifferentRelayTask(t *testing.T) {
_, err := grokMediaSignedVideoContentURL(
[]byte(`{"video":{"url":"/v1/videos/task-2/content"}}`),
"task-1",
)
require.ErrorContains(t, err, "unsupported video content URL")
}
func TestForwardGrokVideoStatusRewritesOnlyProtectedContentURL(t *testing.T) {
statusBody := `{"id":"task-1","status":"completed","url":"https://relay.example/v1/videos/task-1/content","download_url":"/v1/videos/task-1/content","video_url":"https://vidgen.x.ai/task-1.mp4","counter":9007199254740993}`
upstream := &grokMediaContentUpstreamStub{
response: &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"application/json"}},
Body: io.NopCloser(strings.NewReader(statusBody)),
},
}
svc := &OpenAIGatewayService{cfg: &config.Config{}, httpUpstream: upstream}
c, recorder := grokMediaContentTestContext(http.MethodGet, "https://api.example/v1/videos/task-1", map[string]string{
"X-Forwarded-Host": "malicious.invalid",
"X-Forwarded-Proto": "https",
})
_, err := svc.ForwardGrokMedia(
context.Background(), c, grokMediaContentTestAccount(),
GrokMediaEndpointVideoStatus, "task-1", nil, "",
)
require.NoError(t, err)
require.Equal(t, http.StatusOK, recorder.Code)
require.Equal(t, "/v1/videos/task-1/content", gjson.Get(recorder.Body.String(), "url").String())
require.Equal(t, "/v1/videos/task-1/content", gjson.Get(recorder.Body.String(), "download_url").String())
require.Equal(t, "https://vidgen.x.ai/task-1.mp4", gjson.Get(recorder.Body.String(), "video_url").String())
require.Equal(t, "9007199254740993", gjson.Get(recorder.Body.String(), "counter").String())
require.NotContains(t, recorder.Body.String(), "malicious.invalid")
}
func TestRewriteGrokMediaVideoContentURLsPreservesOtherIDsAndHandlesNestedEscapedID(t *testing.T) {
body := []byte(`{"nested":[{"url":"https://relay.example/v1/videos/task%2Fone/content"},{"url":"https://relay.example/v1/videos/task-two/content"}]}`)
rewritten := rewriteGrokMediaVideoContentURLs(body, "task/one", "/v1/videos/task%2Fone/content")
require.Equal(t, "/v1/videos/task%2Fone/content", gjson.GetBytes(rewritten, "nested.0.url").String())
require.Equal(t, "https://relay.example/v1/videos/task-two/content", gjson.GetBytes(rewritten, "nested.1.url").String())
}
func TestRewriteGrokMediaVideoContentURLsRewritesSignedVideoURL(t *testing.T) {
body := []byte(`{"status":"done","video":{"url":"https://vidgen.x.ai/signed-token/xai-video-request-1.mp4","duration":8}}`)
rewritten := rewriteGrokMediaVideoContentURLs(body, "request-1", "/v1/videos/request-1/content")
require.Equal(t, "/v1/videos/request-1/content", gjson.GetBytes(rewritten, "video.url").String())
require.Equal(t, "8", gjson.GetBytes(rewritten, "video.duration").String())
require.Equal(t, "done", gjson.GetBytes(rewritten, "status").String())
}