Files
sub2api/backend/internal/pkg/xai/oauth_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

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")
}