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

243 lines
8.0 KiB
Go

package service
import (
"context"
"encoding/json"
"errors"
"io"
"net/http"
"strings"
"testing"
"time"
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/Wei-Shaw/sub2api/internal/pkg/tlsfingerprint"
coderws "github.com/coder/websocket"
"github.com/stretchr/testify/require"
)
type liveHTTPUpstreamStub struct {
request *http.Request
body []byte
}
type liveAttestationStub struct {
header string
err error
}
func (s liveAttestationStub) Check(context.Context) error {
return s.err
}
func (s liveAttestationStub) Generate(context.Context) (string, error) {
return s.header, s.err
}
func (s *liveHTTPUpstreamStub) Do(
request *http.Request,
_ string,
_ int64,
_ int,
) (*http.Response, error) {
s.request = request
body, err := io.ReadAll(request.Body)
if err != nil {
return nil, err
}
s.body = body
return &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{
"Location": {"/backend-api/codex/call_test"},
},
Body: io.NopCloser(strings.NewReader("v=0\r\n")),
}, nil
}
func (s *liveHTTPUpstreamStub) DoWithTLS(
request *http.Request,
proxyURL string,
accountID int64,
accountConcurrency int,
_ *tlsfingerprint.Profile,
) (*http.Response, error) {
return s.Do(request, proxyURL, accountID, accountConcurrency)
}
func TestLiveCapabilityOnlyAllowsOpenAIOAuth(t *testing.T) {
require.True(t, (&Account{Platform: PlatformOpenAI, Type: AccountTypeOAuth}).SupportsOpenAIEndpointCapability(OpenAIEndpointCapabilityLive))
require.False(t, (&Account{Platform: PlatformOpenAI, Type: AccountTypeAPIKey}).SupportsOpenAIEndpointCapability(OpenAIEndpointCapabilityLive))
require.False(t, (&Account{Platform: PlatformGrok, Type: AccountTypeOAuth}).SupportsOpenAIEndpointCapability(OpenAIEndpointCapabilityLive))
require.False(t, (&Account{
Platform: PlatformOpenAI,
Type: AccountTypeOAuth,
Credentials: map[string]any{
openAIAuthModeCredentialKey: OpenAIAuthModePersonalAccessToken,
},
}).SupportsOpenAIEndpointCapability(OpenAIEndpointCapabilityLive))
require.False(t, (&Account{
Platform: PlatformOpenAI,
Type: AccountTypeOAuth,
Credentials: map[string]any{
openAIAuthModeCredentialKey: OpenAIAuthModeAgentIdentity,
},
}).SupportsOpenAIEndpointCapability(OpenAIEndpointCapabilityLive))
}
func TestValidateLiveCallRequestDoesNotRequireDelegation(t *testing.T) {
request := &LiveCallRequest{
SDP: "v=0\r\n",
Session: json.RawMessage(`{"model":"gpt-live-test","instructions":"hello"}`),
}
require.NoError(t, ValidateLiveCallRequest(request))
require.NotContains(t, string(request.Session), "delegation")
}
func TestCreateUpstreamLiveCallPreservesSession(t *testing.T) {
upstream := &liveHTTPUpstreamStub{}
service := &OpenAIGatewayService{
cfg: &config.Config{},
httpUpstream: upstream,
}
account := &Account{
ID: 7,
Platform: PlatformOpenAI,
Type: AccountTypeOAuth,
Concurrency: 2,
Credentials: map[string]any{
"access_token": "test-access-token",
"chatgpt_account_id": "acct_test",
},
}
session := json.RawMessage(`{
"model":"gpt-live-test",
"delegation":{"type":"client"},
"custom":{"keep":true}
}`)
created, err := service.createUpstreamLiveCall(context.Background(), account, &LiveCallRequest{
SDP: "v=offer\r\n",
Session: session,
}, `{"v":1,"s":0,"t":"v1.test"}`)
require.NoError(t, err)
require.Equal(t, "call_test", created.CallID)
require.Equal(t, []byte("v=0\r\n"), created.SDP)
var forwarded struct {
SDP string `json:"sdp"`
Session json.RawMessage `json:"session"`
}
require.NoError(t, json.Unmarshal(upstream.body, &forwarded))
require.Equal(t, "v=offer\r\n", forwarded.SDP)
require.JSONEq(t, string(session), string(forwarded.Session))
require.Equal(t, "Bearer test-access-token", upstream.request.Header.Get("Authorization"))
require.Equal(t, "acct_test", upstream.request.Header.Get("Chatgpt-Account-Id"))
require.Equal(t, "quicksilver=v2", upstream.request.Header.Get("OpenAI-Alpha"))
require.Equal(t, `{"v":1,"s":0,"t":"v1.test"}`, upstream.request.Header.Get(liveAttestationHeader))
require.NotEmpty(t, upstream.request.Header.Get("Session-Id"))
require.NotEmpty(t, upstream.request.Header.Get("Thread-Id"))
require.Empty(t, upstream.request.Header.Get("OpenAI-Beta"))
require.Equal(t, HTTPUpstreamProfileOpenAI, HTTPUpstreamProfileFromContext(upstream.request.Context()))
require.True(t, HTTPUpstreamRedirectsDisabled(upstream.request.Context()))
}
func TestLiveAttestationCipherRoundTripAndRejectsOtherInstanceKey(t *testing.T) {
first := newLiveAttestationCipher(&config.Config{
JWT: config.JWTConfig{Secret: "first-live-secret"},
})
second := newLiveAttestationCipher(&config.Config{
JWT: config.JWTConfig{Secret: "second-live-secret"},
})
require.NotNil(t, first)
require.NotNil(t, second)
ciphertext, err := first.Encrypt(`{"v":1,"s":0,"t":"v1.opaque"}`)
require.NoError(t, err)
require.NotContains(t, ciphertext, "opaque")
plaintext, err := first.Decrypt(ciphertext)
require.NoError(t, err)
require.Equal(t, `{"v":1,"s":0,"t":"v1.opaque"}`, plaintext)
_, err = second.Decrypt(ciphertext)
require.Error(t, err)
}
func TestPrepareLiveAttestationEncryptsHeaderAndReturnsExplicitProviderError(t *testing.T) {
cipher := newLiveAttestationCipher(&config.Config{
JWT: config.JWTConfig{Secret: "live-attestation-test-secret"},
})
service := &OpenAIGatewayService{
liveAttestation: liveAttestationStub{header: `{"v":1,"s":0,"t":"v1.test"}`},
liveAttestationCipher: cipher,
}
header, ciphertext, err := service.prepareLiveAttestation(context.Background())
require.NoError(t, err)
require.Equal(t, `{"v":1,"s":0,"t":"v1.test"}`, header)
require.NotContains(t, ciphertext, "v1.test")
decrypted, err := cipher.Decrypt(ciphertext)
require.NoError(t, err)
require.Equal(t, header, decrypted)
service.liveAttestation = liveAttestationStub{err: errors.New("macOS app missing")}
_, _, err = service.prepareLiveAttestation(context.Background())
var unavailable *LiveAttestationUnavailableError
require.ErrorAs(t, err, &unavailable)
require.Contains(t, unavailable.Error(), "macOS app missing")
}
func TestLiveMaxSessionDurationDefaultsAndOverrides(t *testing.T) {
require.Equal(t, defaultLiveMaxSessionDuration, (&OpenAIGatewayService{}).liveMaxSessionDuration())
require.Equal(
t,
90*time.Second,
(&OpenAIGatewayService{cfg: &config.Config{
Gateway: config.GatewayConfig{
Live: config.GatewayLiveConfig{MaxSessionDurationSeconds: 90},
},
}}).liveMaxSessionDuration(),
)
}
func TestLiveSidebandNormalCloseEndsCall(t *testing.T) {
normalClose := coderws.CloseError{Code: coderws.StatusNormalClosure}
require.ErrorIs(t, liveSidebandReadError(normalClose), ErrLiveCallNotFound)
abnormalClose := coderws.CloseError{Code: coderws.StatusInternalError}
require.Equal(t, abnormalClose, liveSidebandReadError(abnormalClose))
}
func TestLiveCreateFailoverUsesExistingOpenAIPolicy(t *testing.T) {
service := &OpenAIGatewayService{}
require.False(t, service.shouldFailoverLiveCreateError(&UpstreamFailoverError{
StatusCode: http.StatusBadRequest,
ResponseBody: []byte(`{"error":{"message":"invalid session"}}`),
}))
require.True(t, service.shouldFailoverLiveCreateError(&UpstreamFailoverError{
StatusCode: http.StatusForbidden,
}))
require.True(t, service.shouldFailoverLiveCreateError(&UpstreamFailoverError{
StatusCode: http.StatusBadGateway,
}))
require.True(t, service.shouldFailoverLiveCreateError(errors.New("transport failed")))
}
func TestLiveCallIDFromLocation(t *testing.T) {
callID, err := liveCallIDFromLocation("https://chatgpt.com/backend-api/codex/call_123?intent=quicksilver")
require.NoError(t, err)
require.Equal(t, "call_123", callID)
callID, err = liveCallIDFromLocation("/backend-api/codex/call_456")
require.NoError(t, err)
require.Equal(t, "call_456", callID)
}
func TestRequestTypeLive(t *testing.T) {
require.True(t, RequestTypeLive.IsValid())
require.Equal(t, "live", RequestTypeLive.String())
parsed, err := ParseUsageRequestType("live")
require.NoError(t, err)
require.Equal(t, RequestTypeLive, parsed)
}