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
376 lines
13 KiB
Go
376 lines
13 KiB
Go
//go:build unit
|
|
|
|
package xai
|
|
|
|
import (
|
|
"net/url"
|
|
"testing"
|
|
|
|
"github.com/Wei-Shaw/sub2api/internal/util/urlvalidator"
|
|
"github.com/stretchr/testify/require"
|
|
)
|
|
|
|
func TestParseAuthorizationInput(t *testing.T) {
|
|
t.Parallel()
|
|
|
|
tests := []struct {
|
|
name string
|
|
raw string
|
|
wantCode string
|
|
wantState string
|
|
wantRequiresState bool
|
|
}{
|
|
{
|
|
name: "full callback url",
|
|
raw: "http://127.0.0.1:56121/callback?code=abc123&state=state456",
|
|
wantCode: "abc123",
|
|
wantState: "state456",
|
|
wantRequiresState: true,
|
|
},
|
|
{
|
|
name: "query string",
|
|
raw: "?code=abc123&state=state456",
|
|
wantCode: "abc123",
|
|
wantState: "state456",
|
|
wantRequiresState: true,
|
|
},
|
|
{
|
|
name: "full callback url missing state",
|
|
raw: "http://127.0.0.1:56121/callback?code=abc123",
|
|
wantCode: "abc123",
|
|
wantRequiresState: true,
|
|
},
|
|
{
|
|
name: "query string missing state",
|
|
raw: "code=abc123",
|
|
wantCode: "abc123",
|
|
wantRequiresState: true,
|
|
},
|
|
{
|
|
name: "bare code",
|
|
raw: "abc123",
|
|
wantCode: "abc123",
|
|
},
|
|
}
|
|
|
|
for _, tt := range tests {
|
|
t.Run(tt.name, func(t *testing.T) {
|
|
t.Parallel()
|
|
got := ParseAuthorizationInput(tt.raw)
|
|
require.Equal(t, tt.wantCode, got.Code)
|
|
require.Equal(t, tt.wantState, got.State)
|
|
require.Equal(t, tt.wantRequiresState, got.RequiresState)
|
|
})
|
|
}
|
|
}
|
|
|
|
func TestBuildAuthorizationURLIncludesHermesCompatibleParameters(t *testing.T) {
|
|
t.Setenv(EnvAuthorizeURL, "https://auth.example.test/oauth2/authorize")
|
|
t.Setenv(EnvClientID, "client-id")
|
|
t.Setenv(EnvScope, "openid profile offline_access api:access")
|
|
t.Setenv(EnvAllowUnsafeURLOverrides, "true")
|
|
|
|
authURL, err := BuildAuthorizationURL("state", "challenge", "http://127.0.0.1:56121/callback", "nonce")
|
|
require.NoError(t, err)
|
|
parsed, err := url.Parse(authURL)
|
|
require.NoError(t, err)
|
|
|
|
values := parsed.Query()
|
|
require.Equal(t, "https", parsed.Scheme)
|
|
require.Equal(t, "auth.example.test", parsed.Host)
|
|
require.Equal(t, "/oauth2/authorize", parsed.Path)
|
|
require.Equal(t, "code", values.Get("response_type"))
|
|
require.Equal(t, "client-id", values.Get("client_id"))
|
|
require.Equal(t, "http://127.0.0.1:56121/callback", values.Get("redirect_uri"))
|
|
require.Equal(t, "openid profile offline_access api:access", values.Get("scope"))
|
|
require.Equal(t, "state", values.Get("state"))
|
|
require.Equal(t, "nonce", values.Get("nonce"))
|
|
require.Equal(t, "challenge", values.Get("code_challenge"))
|
|
require.Equal(t, "S256", values.Get("code_challenge_method"))
|
|
require.Equal(t, "generic", values.Get("plan"))
|
|
require.Equal(t, "sub2api", values.Get("referrer"))
|
|
}
|
|
|
|
func TestValidateXAIURLsAllowOfficialOAuthAndGatewayHosts(t *testing.T) {
|
|
authorizeURL, err := ValidateOAuthEndpointURL(DefaultAuthorizeURL)
|
|
require.NoError(t, err)
|
|
require.Equal(t, DefaultAuthorizeURL, authorizeURL)
|
|
|
|
tokenURL, err := ValidateOAuthEndpointURL(DefaultTokenURL)
|
|
require.NoError(t, err)
|
|
require.Equal(t, DefaultTokenURL, tokenURL)
|
|
|
|
baseURL, err := ValidateBaseURL(DefaultBaseURL)
|
|
require.NoError(t, err)
|
|
require.Equal(t, DefaultBaseURL, baseURL)
|
|
|
|
cliBaseURL, err := ValidateBaseURL(DefaultCLIBaseURL)
|
|
require.NoError(t, err)
|
|
require.Equal(t, DefaultCLIBaseURL, cliBaseURL)
|
|
|
|
baseURLNoPath, err := ValidateBaseURL("https://api.x.ai")
|
|
require.NoError(t, err)
|
|
require.Equal(t, DefaultBaseURL, baseURLNoPath)
|
|
|
|
chatURL, err := BuildChatCompletionsURL(DefaultCLIBaseURL + "/")
|
|
require.NoError(t, err)
|
|
require.Equal(t, DefaultCLIBaseURL+"/chat/completions", chatURL)
|
|
}
|
|
|
|
func TestBuildGrokMediaURLs(t *testing.T) {
|
|
imagesURL, err := BuildImagesGenerationsURL(DefaultBaseURL + "/")
|
|
require.NoError(t, err)
|
|
require.Equal(t, DefaultBaseURL+"/images/generations", imagesURL)
|
|
|
|
editsURL, err := BuildImagesEditsURL(DefaultBaseURL)
|
|
require.NoError(t, err)
|
|
require.Equal(t, DefaultBaseURL+"/images/edits", editsURL)
|
|
|
|
videosURL, err := BuildVideosGenerationsURL(DefaultBaseURL)
|
|
require.NoError(t, err)
|
|
require.Equal(t, DefaultBaseURL+"/videos/generations", videosURL)
|
|
|
|
videoEditsURL, err := BuildVideosEditsURL(DefaultBaseURL)
|
|
require.NoError(t, err)
|
|
require.Equal(t, DefaultBaseURL+"/videos/edits", videoEditsURL)
|
|
|
|
videoExtensionsURL, err := BuildVideosExtensionsURL(DefaultBaseURL)
|
|
require.NoError(t, err)
|
|
require.Equal(t, DefaultBaseURL+"/videos/extensions", videoExtensionsURL)
|
|
|
|
videoURL, err := BuildVideoURL(DefaultBaseURL, "req 123")
|
|
require.NoError(t, err)
|
|
require.Equal(t, DefaultBaseURL+"/videos/req%20123", videoURL)
|
|
|
|
_, err = BuildVideoURL(DefaultBaseURL, " ")
|
|
require.Error(t, err)
|
|
}
|
|
|
|
func TestValidateXAIURLsRejectUntrustedOAuthAndUnsafeBaseURLsByDefault(t *testing.T) {
|
|
_, err := ValidateOAuthEndpointURL("https://auth.example.test/oauth2/token")
|
|
require.Error(t, err)
|
|
|
|
_, err = ValidateBaseURL("http://127.0.0.1:8080/v1")
|
|
require.Error(t, err)
|
|
|
|
_, err = ValidateBaseURL("https://api.x.ai/custom")
|
|
require.Error(t, err)
|
|
}
|
|
|
|
func TestValidateBaseURLAllowsPublicThirdPartyGrokAPI(t *testing.T) {
|
|
baseURL, err := ValidateBaseURL("https://grok.example.test/v1/")
|
|
require.NoError(t, err)
|
|
require.Equal(t, "https://grok.example.test/v1", baseURL)
|
|
|
|
_, err = ValidateTrustedBaseURL("https://grok.example.test/v1")
|
|
require.Error(t, err)
|
|
}
|
|
|
|
func TestValidateBaseURLPathPrefixPolicy(t *testing.T) {
|
|
// 非官方主机保留管理员配置的任意 path 前缀。
|
|
prefixed, err := ValidateBaseURL("https://relay.example.test/xai/v1/")
|
|
require.NoError(t, err)
|
|
require.Equal(t, "https://relay.example.test/xai/v1", prefixed)
|
|
|
|
deepPrefixed, err := ValidateBaseURL("https://relay.example.test/tenant-a/proxy")
|
|
require.NoError(t, err)
|
|
require.Equal(t, "https://relay.example.test/tenant-a/proxy", deepPrefixed)
|
|
|
|
// 空 path 仍按惯例补 /v1,保持既有配置兼容。
|
|
rootOnly, err := ValidateBaseURL("https://relay.example.test")
|
|
require.NoError(t, err)
|
|
require.Equal(t, "https://relay.example.test/v1", rootOnly)
|
|
|
|
// 官方主机固定 /v1 前缀。
|
|
_, err = ValidateBaseURL("https://api.x.ai/xai/v1")
|
|
require.Error(t, err)
|
|
_, err = ValidateBaseURL("https://cli-chat-proxy.grok.com/other")
|
|
require.Error(t, err)
|
|
}
|
|
|
|
func TestIsOfficialBaseURL(t *testing.T) {
|
|
official := []string{
|
|
"",
|
|
" ",
|
|
DefaultBaseURL,
|
|
DefaultCLIBaseURL,
|
|
"https://api.x.ai",
|
|
"HTTPS://API.X.AI:443/",
|
|
"https://api.x.ai:0443/v1",
|
|
"https://api.x.ai/%76%31",
|
|
"https://api.x.ai:8443/v1",
|
|
"HTTPS://CLI-CHAT-PROXY.GROK.COM:443/%76%31/",
|
|
"::invalid::url", // 无法解析的值按官方处理,回落默认端点
|
|
}
|
|
for _, raw := range official {
|
|
require.True(t, IsOfficialBaseURL(raw), "expected official: %q", raw)
|
|
}
|
|
|
|
custom := []string{
|
|
"https://relay.example.test/v1",
|
|
"https://relay.example.test/xai/v1",
|
|
"http://relay.example.test/v1",
|
|
"https://grok.com.evil.example.test/v1",
|
|
"https://api.x.ai.evil.example.test/v1", // 后缀伪装不属于 *.api.x.ai
|
|
}
|
|
for _, raw := range custom {
|
|
require.False(t, IsOfficialBaseURL(raw), "expected custom: %q", raw)
|
|
}
|
|
}
|
|
|
|
func TestRegionalAPIEndpointsAreOfficialAndTrusted(t *testing.T) {
|
|
regional := []string{
|
|
"https://us-east-1.api.x.ai/v1",
|
|
"https://us-west-2.api.x.ai/v1",
|
|
"https://eu-west-1.api.x.ai/v1",
|
|
}
|
|
for _, raw := range regional {
|
|
require.True(t, IsOfficialBaseURL(raw), "expected official: %q", raw)
|
|
|
|
validated, err := ValidateTrustedBaseURL(raw)
|
|
require.NoError(t, err, "trusted validation should accept regional endpoint %q", raw)
|
|
require.Equal(t, raw, validated)
|
|
}
|
|
|
|
// 区域端点作为官方主机同样强制 /v1 path
|
|
_, err := ValidateTrustedBaseURL("https://us-east-1.api.x.ai/other")
|
|
require.Error(t, err)
|
|
}
|
|
|
|
func TestValidateBaseURLsRejectEmptyQueryDelimiter(t *testing.T) {
|
|
_, err := ValidateBaseURL("https://grok.example.test/v1?")
|
|
require.Error(t, err)
|
|
|
|
_, err = ValidateTrustedBaseURL("https://api.x.ai/v1?")
|
|
require.Error(t, err)
|
|
}
|
|
|
|
func TestBuildResponsesURLWithValidatorUsesCallerPolicy(t *testing.T) {
|
|
validator := func(raw string) (string, error) {
|
|
return urlvalidator.ValidateURLFormat(raw, true)
|
|
}
|
|
|
|
target, err := BuildResponsesURLWithValidator("http://grok.example.test/v1/", validator)
|
|
require.NoError(t, err)
|
|
require.Equal(t, "http://grok.example.test/v1/responses", target)
|
|
}
|
|
|
|
func TestValidateTrustedBaseURLAcceptsOfficialRegionalHosts(t *testing.T) {
|
|
for _, raw := range []string{DefaultUSEast1BaseURL, DefaultUSWest2BaseURL, DefaultEUWest1BaseURL} {
|
|
got, err := ValidateTrustedBaseURL(raw)
|
|
require.NoError(t, err, raw)
|
|
require.Equal(t, raw, got)
|
|
}
|
|
}
|
|
|
|
func TestBuildResponsesURLPreservesUnsafeOverrideCustomPath(t *testing.T) {
|
|
t.Setenv(EnvAllowUnsafeURLOverrides, "true")
|
|
|
|
target, err := BuildResponsesURL("http://localhost:8080/custom")
|
|
require.NoError(t, err)
|
|
require.Equal(t, "http://localhost:8080/custom/responses", target)
|
|
}
|
|
|
|
func TestBuildResponsesURLWithValidatorRejectsBaseURLComponents(t *testing.T) {
|
|
permissive := func(raw string) (string, error) { return raw, nil }
|
|
tests := []struct {
|
|
name string
|
|
raw string
|
|
}{
|
|
{name: "userinfo", raw: "https://user:secret@grok.example.test/v1"},
|
|
{name: "query", raw: "https://grok.example.test/v1?token=secret"},
|
|
{name: "empty query delimiter", raw: "https://grok.example.test/v1?"},
|
|
{name: "fragment", raw: "https://grok.example.test/v1#secret"},
|
|
}
|
|
|
|
for _, tt := range tests {
|
|
t.Run(tt.name, func(t *testing.T) {
|
|
_, err := BuildResponsesURLWithValidator(tt.raw, permissive)
|
|
require.Error(t, err)
|
|
require.NotContains(t, err.Error(), "secret")
|
|
})
|
|
}
|
|
}
|
|
|
|
func TestValidateXAIURLsAllowUnsafeDevOverride(t *testing.T) {
|
|
t.Setenv(EnvAllowUnsafeURLOverrides, "true")
|
|
|
|
tokenURL, err := ValidateOAuthEndpointURL("http://127.0.0.1:8080/oauth2/token")
|
|
require.NoError(t, err)
|
|
require.Equal(t, "http://127.0.0.1:8080/oauth2/token", tokenURL)
|
|
|
|
baseURL, err := ValidateBaseURL("http://127.0.0.1:8080/v1/")
|
|
require.NoError(t, err)
|
|
require.Equal(t, "http://127.0.0.1:8080/v1", baseURL)
|
|
}
|
|
|
|
func TestRuntimeSanityReportsSafeDefaults(t *testing.T) {
|
|
t.Setenv(EnvBaseURL, "")
|
|
t.Setenv(EnvAuthorizeURL, "")
|
|
t.Setenv(EnvTokenURL, "")
|
|
t.Setenv(EnvRedirectURI, "")
|
|
t.Setenv(EnvAllowUnsafeURLOverrides, "")
|
|
t.Setenv(EnvUnsafeAllowHighConcurrency, "")
|
|
|
|
report := RuntimeSanity()
|
|
require.True(t, report.BaseURL.Valid)
|
|
require.Equal(t, DefaultBaseURL, report.BaseURL.Value)
|
|
require.True(t, report.BaseURL.IsDefault)
|
|
require.True(t, report.OAuthAuthorizeURL.Valid)
|
|
require.True(t, report.OAuthTokenURL.Valid)
|
|
require.True(t, report.OAuthRedirectURI.Valid)
|
|
require.False(t, report.UnsafeURLOverrides)
|
|
require.False(t, report.UnsafeHighConcurrency)
|
|
require.Equal(t, "responses_only", report.PublicGatewayScope)
|
|
require.Contains(t, report.ProxyPolicy, "account_proxy_optional")
|
|
require.Contains(t, report.ProxyPolicy, "API-key base URLs require public HTTPS")
|
|
}
|
|
|
|
func TestRuntimeSanityReportsInvalidOverridesWithoutSecrets(t *testing.T) {
|
|
t.Setenv(EnvBaseURL, "http://127.0.0.1:8080/v1?access_token=secret")
|
|
t.Setenv(EnvAuthorizeURL, "https://auth.example.test/oauth2/authorize")
|
|
t.Setenv(EnvTokenURL, "https://auth.example.test/oauth2/token")
|
|
t.Setenv(EnvRedirectURI, "not a url")
|
|
t.Setenv(EnvClientID, "client-secret-like-value")
|
|
t.Setenv(EnvAllowUnsafeURLOverrides, "")
|
|
|
|
report := RuntimeSanity()
|
|
require.False(t, report.BaseURL.Valid)
|
|
require.False(t, report.BaseURL.IsDefault)
|
|
require.Contains(t, report.BaseURL.Error, "invalid url")
|
|
require.NotContains(t, report.BaseURL.Value, "secret")
|
|
require.False(t, report.OAuthAuthorizeURL.Valid)
|
|
require.False(t, report.OAuthTokenURL.Valid)
|
|
require.False(t, report.OAuthRedirectURI.Valid)
|
|
require.NotContains(t, report.ProxyPolicy, "client-secret-like-value")
|
|
}
|
|
|
|
func TestDefaultModelMappingIncludesGrokAliases(t *testing.T) {
|
|
original := RuntimeModelMappingOptions()
|
|
t.Cleanup(func() { SetRuntimeModelMappingOptions(original) })
|
|
SetRuntimeModelMappingOptions(ModelMappingOptions{})
|
|
mapping := DefaultModelMapping()
|
|
require.Equal(t, "grok-4.5", mapping["grok"])
|
|
require.Equal(t, "grok-4.5", mapping["grok-latest"])
|
|
require.Equal(t, "grok-4.6", mapping["grok-4.6"])
|
|
require.Equal(t, "grok-4.6", mapping["grok-4.6-latest"])
|
|
require.Equal(t, "grok-4.5", mapping["grok-4.5"])
|
|
require.Equal(t, "grok-4.5", mapping["grok-4.5-latest"])
|
|
require.Equal(t, "grok-build-0.1", mapping["grok-build"])
|
|
require.Equal(t, "grok-4.5", mapping["grok-build-latest"])
|
|
require.Equal(t, "grok-composer-2.5-fast", mapping["grok-composer"])
|
|
require.Equal(t, "grok-composer-2.5-fast", mapping["composer-2.5"])
|
|
require.Equal(t, "grok-4.20-0309-reasoning", mapping["grok-4.20-reasoning"])
|
|
require.Equal(t, "grok-4.20-0309-non-reasoning", mapping["grok-4.20-non-reasoning"])
|
|
require.Equal(t, "grok-4.20-multi-agent-0309", mapping["grok-4.20-multi-agent-0309"])
|
|
require.Equal(t, DefaultImagineImageQualityModel, mapping["grok-imagine"])
|
|
require.Equal(t, DefaultImagineImageFastModel, mapping["grok-imagine-image"])
|
|
require.Equal(t, DefaultImagineImageQualityModel, mapping["grok-imagine-image-quality"])
|
|
require.Equal(t, DefaultImagineImageQualityModel, mapping["grok-imagine-edit"])
|
|
require.Equal(t, DefaultImagineVideoModel, mapping["grok-imagine-video"])
|
|
require.Equal(t, DefaultImagineVideo15LegacyModel, mapping["grok-imagine-video-1.5"])
|
|
require.Equal(t, DefaultImagineVideo15Model, mapping["grok-imagine-video-1.5-preview"])
|
|
_, hasGPT := mapping["gpt-*"]
|
|
require.False(t, hasGPT, "cross-client wildcards must be opt-in")
|
|
}
|