126 lines
4.7 KiB
Go
126 lines
4.7 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
|
|
response *http.Response
|
|
}
|
|
|
|
func (s *grokMediaContentUpstreamStub) Do(req *http.Request, _ string, _ int64, _ int) (*http.Response, error) {
|
|
s.request = req
|
|
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 TestForwardGrokMediaContentUsesUpstreamCredentialAndStreamsRange(t *testing.T) {
|
|
upstream := &grokMediaContentUpstreamStub{
|
|
response: &http.Response{
|
|
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.Equal(t, "https://relay.example/v1/videos/task-1/content", upstream.request.URL.String())
|
|
require.Equal(t, "Bearer upstream-key", upstream.request.Header.Get("Authorization"))
|
|
require.Equal(t, "bytes=0-12", upstream.request.Header.Get("Range"))
|
|
require.Equal(t, "*/*", upstream.request.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"))
|
|
}
|
|
|
|
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")
|
|
}
|