mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-21 14:19:18 +08:00
fix(openai): fallback PAT alpha search to responses web_search
This commit is contained in:
@@ -59,6 +59,14 @@ func (s *OpenAIGatewayService) ForwardAlphaSearch(ctx context.Context, c *gin.Co
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Codex Personal Access Token(at-...)目前可访问 ChatGPT Codex
|
||||
// /responses,但会被 standalone /alpha/search 的 access enforcement
|
||||
// 拒绝为 no_matching_rule。对 PAT 账号使用等价的 hosted web_search
|
||||
// Responses 路径兜底,避免把可用账号误判为搜索不可用。
|
||||
if account.IsOpenAIPersonalAccessToken() {
|
||||
return s.forwardAlphaSearchViaResponsesWebSearch(ctx, c, account, body, token, proxyURL, requestedModel, upstreamModel)
|
||||
}
|
||||
|
||||
req, err := s.buildOpenAIAlphaSearchRequest(ctx, c, account, body, token)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -119,6 +127,211 @@ func (s *OpenAIGatewayService) ForwardAlphaSearch(ctx context.Context, c *gin.Co
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s *OpenAIGatewayService) forwardAlphaSearchViaResponsesWebSearch(
|
||||
ctx context.Context,
|
||||
c *gin.Context,
|
||||
account *Account,
|
||||
alphaBody []byte,
|
||||
token string,
|
||||
proxyURL string,
|
||||
requestedModel string,
|
||||
upstreamModel string,
|
||||
) (*OpenAIForwardResult, error) {
|
||||
if upstreamModel == "" {
|
||||
upstreamModel = requestedModel
|
||||
}
|
||||
responsesBody, err := buildOpenAIAlphaSearchResponsesWebSearchBody(alphaBody, upstreamModel)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
req, err := s.buildOpenAIAlphaSearchResponsesWebSearchRequest(ctx, c, account, alphaBody, responsesBody, token)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
SetActualOpenAIUpstreamEndpoint(c, "/v1/responses")
|
||||
|
||||
upstreamStart := time.Now()
|
||||
resp, err := s.httpUpstream.Do(req, proxyURL, account.ID, account.Concurrency)
|
||||
SetOpsLatencyMs(c, OpsUpstreamLatencyMsKey, time.Since(upstreamStart).Milliseconds())
|
||||
if err != nil {
|
||||
return nil, s.handleOpenAIUpstreamTransportError(ctx, c, account, err, true)
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
|
||||
respBody, err := ReadUpstreamResponseBody(resp.Body, s.cfg, c, openAITooLargeError)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("read alpha search responses fallback response: %w", err)
|
||||
}
|
||||
|
||||
if resp.StatusCode >= http.StatusBadRequest {
|
||||
upstreamMessage := sanitizeUpstreamErrorMessage(strings.TrimSpace(extractUpstreamErrorMessage(respBody)))
|
||||
if s.shouldFailoverOpenAIUpstreamResponse(resp.StatusCode, upstreamMessage, respBody) {
|
||||
resp.Body = io.NopCloser(bytes.NewReader(respBody))
|
||||
// 仍按 alpha/search 工具请求处理:PAT 的工具链路失败不能直接永久置错。
|
||||
if shouldApplyOpenAIAlphaSearchAccountErrorSideEffects(resp.StatusCode) {
|
||||
s.handleFailoverSideEffects(ctx, resp, account, respBody, upstreamModel)
|
||||
}
|
||||
return nil, &UpstreamFailoverError{
|
||||
StatusCode: resp.StatusCode,
|
||||
ResponseBody: respBody,
|
||||
RetryableOnSameAccount: account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if resp.StatusCode < http.StatusOK || resp.StatusCode >= http.StatusMultipleChoices {
|
||||
writeOpenAIPassthroughResponseHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter)
|
||||
contentType := resp.Header.Get("Content-Type")
|
||||
if contentType == "" {
|
||||
contentType = "application/json"
|
||||
}
|
||||
c.Data(resp.StatusCode, contentType, respBody)
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
if !account.IsShadow() {
|
||||
s.UpdateCodexUsageSnapshotFromHeaders(ctx, account.ID, resp.Header)
|
||||
}
|
||||
alphaRespBody, err := openAIAlphaSearchResponseFromResponsesSSE(respBody)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
c.Data(http.StatusOK, "application/json", alphaRespBody)
|
||||
return &OpenAIForwardResult{
|
||||
RequestID: strings.TrimSpace(resp.Header.Get("x-request-id")),
|
||||
Model: requestedModel,
|
||||
UpstreamModel: upstreamModel,
|
||||
UpstreamEndpoint: "/v1/responses",
|
||||
ResponseHeaders: resp.Header.Clone(),
|
||||
Duration: time.Since(upstreamStart),
|
||||
WebSearchCalls: 1,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s *OpenAIGatewayService) buildOpenAIAlphaSearchResponsesWebSearchRequest(ctx context.Context, c *gin.Context, account *Account, alphaBody []byte, body []byte, token string) (*http.Request, error) {
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, chatgptCodexURL, bytes.NewReader(body))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
req = req.WithContext(WithHTTPUpstreamProfile(req.Context(), HTTPUpstreamProfileOpenAI))
|
||||
|
||||
authHeaders, err := s.buildOpenAIAuthenticationHeaders(ctx, account, token)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("build openai authentication headers: %w", err)
|
||||
}
|
||||
for key, values := range authHeaders {
|
||||
for _, value := range values {
|
||||
req.Header.Add(key, value)
|
||||
}
|
||||
}
|
||||
req.Host = "chatgpt.com"
|
||||
if err := resolveAndSetOpenAIChatGPTAccountHeaders(ctx, s.accountRepo, req.Header, account); err != nil {
|
||||
return nil, fmt.Errorf("resolve chatgpt account headers: %w", err)
|
||||
}
|
||||
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Accept", "text/event-stream")
|
||||
req.Header.Set("OpenAI-Beta", "responses=experimental")
|
||||
if turnMetadata := openAIAlphaSearchInboundHeader(c, "X-Codex-Turn-Metadata"); turnMetadata != "" {
|
||||
req.Header.Set("X-Codex-Turn-Metadata", turnMetadata)
|
||||
}
|
||||
if version := openAIAlphaSearchInboundHeader(c, "Version"); version != "" {
|
||||
req.Header.Set("Version", version)
|
||||
} else {
|
||||
req.Header.Set("Version", codexCLIVersion)
|
||||
}
|
||||
if originator := openAIAlphaSearchInboundHeader(c, "Originator"); originator != "" {
|
||||
req.Header.Set("Originator", originator)
|
||||
} else {
|
||||
req.Header.Set("Originator", "codex_cli_rs")
|
||||
}
|
||||
if customUA := account.GetOpenAIUserAgent(); customUA != "" {
|
||||
req.Header.Set("User-Agent", customUA)
|
||||
} else if userAgent := openAIAlphaSearchInboundHeader(c, "User-Agent"); userAgent != "" {
|
||||
req.Header.Set("User-Agent", userAgent)
|
||||
} else {
|
||||
req.Header.Set("User-Agent", codexCLIUserAgent)
|
||||
}
|
||||
if s.cfg != nil && s.cfg.Gateway.ForceCodexCLI {
|
||||
req.Header.Set("User-Agent", codexCLIUserAgent)
|
||||
}
|
||||
apiKeyID := getAPIKeyIDFromContext(c)
|
||||
if sessionID := strings.TrimSpace(gjson.GetBytes(alphaBody, "id").String()); sessionID != "" {
|
||||
isolated := isolateOpenAISessionID(apiKeyID, sessionID)
|
||||
req.Header.Set("Session_ID", isolated)
|
||||
req.Header.Set("Conversation_ID", isolated)
|
||||
}
|
||||
s.overrideBrowserUserAgent(ctx, account, req)
|
||||
enforceCodexIdentityHeaders(req.Header)
|
||||
account.ApplyHeaderOverrides(req.Header)
|
||||
return req, nil
|
||||
}
|
||||
|
||||
func buildOpenAIAlphaSearchResponsesWebSearchBody(alphaBody []byte, model string) ([]byte, error) {
|
||||
if strings.TrimSpace(model) == "" {
|
||||
return nil, fmt.Errorf("model is required")
|
||||
}
|
||||
tool := map[string]any{"type": "web_search"}
|
||||
if contextSize := strings.TrimSpace(gjson.GetBytes(alphaBody, "settings.search_context_size").String()); contextSize != "" {
|
||||
tool["search_context_size"] = contextSize
|
||||
}
|
||||
if userLocation := gjson.GetBytes(alphaBody, "settings.user_location"); userLocation.IsObject() {
|
||||
var loc map[string]any
|
||||
if err := json.Unmarshal([]byte(userLocation.Raw), &loc); err == nil && len(loc) > 0 {
|
||||
tool["user_location"] = loc
|
||||
}
|
||||
}
|
||||
payload := map[string]any{
|
||||
"model": model,
|
||||
"stream": true,
|
||||
"store": false,
|
||||
"input": []any{
|
||||
map[string]any{
|
||||
"role": "user",
|
||||
"content": []any{
|
||||
map[string]any{
|
||||
"type": "input_text",
|
||||
"text": openAIAlphaSearchResponsesWebSearchPrompt(alphaBody),
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
"tools": []any{tool},
|
||||
}
|
||||
return json.Marshal(payload)
|
||||
}
|
||||
|
||||
func openAIAlphaSearchResponsesWebSearchPrompt(alphaBody []byte) string {
|
||||
var b strings.Builder
|
||||
b.WriteString("Execute this Codex standalone web.run request for another model.\n")
|
||||
b.WriteString("Use the hosted web_search tool when web/current information is needed.\n")
|
||||
b.WriteString("Return concise source-backed results. Include titles, URLs, dates, and direct answers when available.\n")
|
||||
if commands := strings.TrimSpace(gjson.GetBytes(alphaBody, "commands").Raw); commands != "" {
|
||||
b.WriteString("\nCommands JSON:\n")
|
||||
b.WriteString(truncateOpenAIAlphaSearchPromptJSON(commands, 12000))
|
||||
}
|
||||
if settings := strings.TrimSpace(gjson.GetBytes(alphaBody, "settings").Raw); settings != "" {
|
||||
b.WriteString("\n\nSearch settings JSON:\n")
|
||||
b.WriteString(truncateOpenAIAlphaSearchPromptJSON(settings, 4000))
|
||||
}
|
||||
if input := strings.TrimSpace(gjson.GetBytes(alphaBody, "input").Raw); input != "" {
|
||||
b.WriteString("\n\nRecent conversation/input JSON:\n")
|
||||
b.WriteString(truncateOpenAIAlphaSearchPromptJSON(input, 8000))
|
||||
}
|
||||
if b.Len() == 0 {
|
||||
return "Execute the requested web search and return concise source-backed results."
|
||||
}
|
||||
return b.String()
|
||||
}
|
||||
|
||||
func truncateOpenAIAlphaSearchPromptJSON(value string, limit int) string {
|
||||
value = strings.TrimSpace(value)
|
||||
if limit <= 0 || len(value) <= limit {
|
||||
return value
|
||||
}
|
||||
return value[:limit] + "\n...<truncated>"
|
||||
}
|
||||
|
||||
func (s *OpenAIGatewayService) buildOpenAIAlphaSearchRequest(ctx context.Context, c *gin.Context, account *Account, body []byte, token string) (*http.Request, error) {
|
||||
targetURL, err := s.openAIAlphaSearchURL(account)
|
||||
if err != nil {
|
||||
@@ -296,6 +509,120 @@ func shouldApplyOpenAIAlphaSearchAccountErrorSideEffects(statusCode int) bool {
|
||||
return statusCode != http.StatusUnauthorized
|
||||
}
|
||||
|
||||
func openAIAlphaSearchResponseFromResponsesSSE(body []byte) ([]byte, error) {
|
||||
output, results := parseOpenAIResponsesSSEForAlphaSearch(body)
|
||||
resp := map[string]any{
|
||||
"output": output,
|
||||
}
|
||||
if len(results) > 0 {
|
||||
resp["results"] = results
|
||||
}
|
||||
return json.Marshal(resp)
|
||||
}
|
||||
|
||||
func parseOpenAIResponsesSSEForAlphaSearch(body []byte) (string, []any) {
|
||||
text := strings.ReplaceAll(string(body), "\r\n", "\n")
|
||||
var output strings.Builder
|
||||
var completedResponse any
|
||||
results := make([]any, 0)
|
||||
seenURLs := make(map[string]struct{})
|
||||
|
||||
for _, block := range strings.Split(text, "\n\n") {
|
||||
data := openAIAlphaSearchSSEData(block)
|
||||
if data == "" || data == "[DONE]" {
|
||||
continue
|
||||
}
|
||||
var event map[string]any
|
||||
if err := json.Unmarshal([]byte(data), &event); err != nil {
|
||||
continue
|
||||
}
|
||||
if delta, _ := event["delta"].(string); delta != "" && event["type"] == "response.output_text.delta" {
|
||||
output.WriteString(delta)
|
||||
}
|
||||
if event["type"] == "response.completed" {
|
||||
completedResponse = event["response"]
|
||||
}
|
||||
collectOpenAIAlphaSearchURLCitations(event, &results, seenURLs)
|
||||
}
|
||||
|
||||
out := output.String()
|
||||
if strings.TrimSpace(out) == "" && completedResponse != nil {
|
||||
out = extractOpenAIResponsesCompletedText(completedResponse)
|
||||
collectOpenAIAlphaSearchURLCitations(completedResponse, &results, seenURLs)
|
||||
}
|
||||
return out, results
|
||||
}
|
||||
|
||||
func openAIAlphaSearchSSEData(block string) string {
|
||||
var lines []string
|
||||
for _, line := range strings.Split(block, "\n") {
|
||||
line = strings.TrimRight(line, "\r")
|
||||
if !strings.HasPrefix(line, "data:") {
|
||||
continue
|
||||
}
|
||||
lines = append(lines, strings.TrimSpace(strings.TrimPrefix(line, "data:")))
|
||||
}
|
||||
return strings.TrimSpace(strings.Join(lines, "\n"))
|
||||
}
|
||||
|
||||
func extractOpenAIResponsesCompletedText(response any) string {
|
||||
resp, ok := response.(map[string]any)
|
||||
if !ok {
|
||||
return ""
|
||||
}
|
||||
outputItems, _ := resp["output"].([]any)
|
||||
var b strings.Builder
|
||||
for _, item := range outputItems {
|
||||
itemMap, ok := item.(map[string]any)
|
||||
if !ok || itemMap["type"] != "message" {
|
||||
continue
|
||||
}
|
||||
contentItems, _ := itemMap["content"].([]any)
|
||||
for _, content := range contentItems {
|
||||
contentMap, ok := content.(map[string]any)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if contentMap["type"] == "output_text" {
|
||||
if text, _ := contentMap["text"].(string); text != "" {
|
||||
b.WriteString(text)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
return b.String()
|
||||
}
|
||||
|
||||
func collectOpenAIAlphaSearchURLCitations(value any, results *[]any, seen map[string]struct{}) {
|
||||
switch typed := value.(type) {
|
||||
case map[string]any:
|
||||
if typed["type"] == "url_citation" {
|
||||
if urlValue, _ := typed["url"].(string); strings.TrimSpace(urlValue) != "" {
|
||||
urlValue = strings.TrimSpace(urlValue)
|
||||
if _, exists := seen[urlValue]; !exists {
|
||||
seen[urlValue] = struct{}{}
|
||||
result := map[string]any{
|
||||
"type": "text_result",
|
||||
"ref_id": fmt.Sprintf("turn0search%d", len(*results)),
|
||||
"url": urlValue,
|
||||
}
|
||||
if title, _ := typed["title"].(string); strings.TrimSpace(title) != "" {
|
||||
result["title"] = strings.TrimSpace(title)
|
||||
}
|
||||
*results = append(*results, result)
|
||||
}
|
||||
}
|
||||
}
|
||||
for _, child := range typed {
|
||||
collectOpenAIAlphaSearchURLCitations(child, results, seen)
|
||||
}
|
||||
case []any:
|
||||
for _, child := range typed {
|
||||
collectOpenAIAlphaSearchURLCitations(child, results, seen)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (s *OpenAIGatewayService) openAIAlphaSearchURL(account *Account) (string, error) {
|
||||
if account == nil {
|
||||
return "", fmt.Errorf("account is required")
|
||||
|
||||
@@ -6,6 +6,7 @@ import (
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
@@ -28,6 +29,15 @@ func (r *alphaSearchAccountStateRepo) SetError(_ context.Context, _ int64, error
|
||||
return nil
|
||||
}
|
||||
|
||||
func alphaSearchResponsesSSE(output string) string {
|
||||
return "event: response.output_text.delta\n" +
|
||||
`data: {"type":"response.output_text.delta","delta":` + strconv.Quote(output) + `}` + "\n\n" +
|
||||
"event: response.output_text.annotation.added\n" +
|
||||
`data: {"type":"response.output_text.annotation.added","annotation":{"type":"url_citation","url":"https://example.com/news","title":"Example News"}}` + "\n\n" +
|
||||
"event: response.completed\n" +
|
||||
`data: {"type":"response.completed","response":{"output":[{"type":"message","content":[{"type":"output_text","text":` + strconv.Quote(output) + `}]}]}}` + "\n\n"
|
||||
}
|
||||
|
||||
func TestForwardAlphaSearchOAuthPreservesWire(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
body := []byte(`{
|
||||
@@ -83,7 +93,7 @@ func TestForwardAlphaSearchOAuthPreservesWire(t *testing.T) {
|
||||
require.JSONEq(t, string(body), string(upstream.lastBody))
|
||||
}
|
||||
|
||||
func TestForwardAlphaSearchPATUsesStandaloneHeaderShape(t *testing.T) {
|
||||
func TestForwardAlphaSearchPATUsesResponsesWebSearchFallback(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
body := []byte(`{
|
||||
"id":"search-session",
|
||||
@@ -111,8 +121,8 @@ func TestForwardAlphaSearchPATUsesStandaloneHeaderShape(t *testing.T) {
|
||||
|
||||
upstream := &httpUpstreamRecorder{resp: &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: http.Header{"Content-Type": []string{"application/json"}},
|
||||
Body: io.NopCloser(strings.NewReader(`{"output":"search result"}`)),
|
||||
Header: http.Header{"Content-Type": []string{"text/event-stream"}, "x-request-id": []string{"req-search"}},
|
||||
Body: io.NopCloser(strings.NewReader(alphaSearchResponsesSSE("search result"))),
|
||||
}}
|
||||
service := &OpenAIGatewayService{cfg: &config.Config{}, httpUpstream: upstream}
|
||||
account := &Account{
|
||||
@@ -132,17 +142,20 @@ func TestForwardAlphaSearchPATUsesStandaloneHeaderShape(t *testing.T) {
|
||||
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, result)
|
||||
require.Equal(t, 1, result.WebSearchCalls)
|
||||
require.Equal(t, "/v1/responses", result.UpstreamEndpoint)
|
||||
require.Equal(t, http.StatusOK, recorder.Code)
|
||||
require.JSONEq(t, `{"output":"search result","results":[{"type":"text_result","ref_id":"turn0search0","url":"https://example.com/news","title":"Example News"}]}`, recorder.Body.String())
|
||||
require.Equal(t, chatgptCodexURL, upstream.lastReq.URL.String())
|
||||
require.Equal(t, "Bearer at-test-token", upstream.lastReq.Header.Get("Authorization"))
|
||||
require.Equal(t, "chatgpt-account", upstream.lastReq.Header.Get("ChatGPT-Account-ID"))
|
||||
require.Equal(t, "true", upstream.lastReq.Header.Get("X-OpenAI-Fedramp"))
|
||||
require.Equal(t, "application/json", upstream.lastReq.Header.Get("Content-Type"))
|
||||
require.Equal(t, "application/json", upstream.lastReq.Header.Get("Accept"))
|
||||
require.Equal(t, "text/event-stream", upstream.lastReq.Header.Get("Accept"))
|
||||
require.Equal(t, "responses=experimental", upstream.lastReq.Header.Get("OpenAI-Beta"))
|
||||
require.Equal(t, "0.144.1", upstream.lastReq.Header.Get("Version"))
|
||||
require.Equal(t, `{"turn_id":"turn-1"}`, upstream.lastReq.Header.Get("X-Codex-Turn-Metadata"))
|
||||
require.Equal(t, "codex_cli_rs", upstream.lastReq.Header.Get("Originator"))
|
||||
require.Empty(t, upstream.lastReq.Header.Get("OpenAI-Beta"))
|
||||
require.Empty(t, upstream.lastReq.Header.Get("Session_ID"))
|
||||
require.Empty(t, upstream.lastReq.Header.Get("Conversation_ID"))
|
||||
require.Empty(t, upstream.lastReq.Header.Get("X-Codex-Beta-Features"))
|
||||
require.Empty(t, upstream.lastReq.Header.Get("X-Codex-Turn-State"))
|
||||
require.Empty(t, upstream.lastReq.Header.Get(responsesLiteHeaderKey))
|
||||
@@ -150,6 +163,10 @@ func TestForwardAlphaSearchPATUsesStandaloneHeaderShape(t *testing.T) {
|
||||
require.False(t, gjson.GetBytes(upstream.lastBody, "prompt_cache_key").Exists())
|
||||
require.False(t, gjson.GetBytes(upstream.lastBody, "prompt_cache_retention").Exists())
|
||||
require.Equal(t, "gpt-5.6-sol", gjson.GetBytes(upstream.lastBody, "model").String())
|
||||
require.True(t, gjson.GetBytes(upstream.lastBody, "stream").Bool())
|
||||
require.False(t, gjson.GetBytes(upstream.lastBody, "store").Bool())
|
||||
require.Equal(t, "web_search", gjson.GetBytes(upstream.lastBody, "tools.0.type").String())
|
||||
require.Contains(t, gjson.GetBytes(upstream.lastBody, "input.0.content.0.text").String(), `"search_query"`)
|
||||
}
|
||||
|
||||
func TestForwardAlphaSearchPATBackfillsMissingChatGPTAccountMetadata(t *testing.T) {
|
||||
@@ -330,6 +347,52 @@ func TestForwardAlphaSearchUnauthorizedDoesNotMarkAccountError(t *testing.T) {
|
||||
require.False(t, c.Writer.Written())
|
||||
}
|
||||
|
||||
func TestForwardAlphaSearchPATResponsesFallbackUnauthorizedDoesNotMarkAccountError(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
body := []byte(`{"id":"search-session","model":"gpt-5.6-sol","commands":{"search_query":[{"q":"news"}]}}`)
|
||||
recorder := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(recorder)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/alpha/search", bytes.NewReader(body))
|
||||
|
||||
upstream := &httpUpstreamRecorder{resp: &http.Response{
|
||||
StatusCode: http.StatusUnauthorized,
|
||||
Header: http.Header{"Content-Type": []string{"application/json"}},
|
||||
Body: io.NopCloser(strings.NewReader(`{"detail":"Unauthorized"}`)),
|
||||
}}
|
||||
repo := &alphaSearchAccountStateRepo{}
|
||||
cfg := &config.Config{}
|
||||
service := &OpenAIGatewayService{
|
||||
cfg: cfg,
|
||||
httpUpstream: upstream,
|
||||
accountRepo: repo,
|
||||
rateLimitService: NewRateLimitService(repo, nil, cfg, nil, nil),
|
||||
}
|
||||
account := &Account{
|
||||
ID: 46,
|
||||
Platform: PlatformOpenAI,
|
||||
Type: AccountTypeOAuth,
|
||||
Concurrency: 1,
|
||||
Credentials: map[string]any{
|
||||
"access_token": "at-test-token",
|
||||
"auth_mode": OpenAIAuthModePersonalAccessToken,
|
||||
"chatgpt_account_id": "chatgpt-account",
|
||||
},
|
||||
}
|
||||
|
||||
result, err := service.ForwardAlphaSearch(context.Background(), c, account, body)
|
||||
|
||||
require.Nil(t, result)
|
||||
var failoverErr *UpstreamFailoverError
|
||||
require.ErrorAs(t, err, &failoverErr)
|
||||
require.Equal(t, http.StatusUnauthorized, failoverErr.StatusCode)
|
||||
require.Equal(t, chatgptCodexURL, upstream.lastReq.URL.String())
|
||||
require.Equal(t, "text/event-stream", upstream.lastReq.Header.Get("Accept"))
|
||||
require.Equal(t, "responses=experimental", upstream.lastReq.Header.Get("OpenAI-Beta"))
|
||||
require.Zero(t, repo.setErrorCalls)
|
||||
require.Empty(t, repo.lastError)
|
||||
require.False(t, c.Writer.Written())
|
||||
}
|
||||
|
||||
func TestShouldApplyOpenAIAlphaSearchAccountErrorSideEffects(t *testing.T) {
|
||||
require.False(t, shouldApplyOpenAIAlphaSearchAccountErrorSideEffects(http.StatusUnauthorized))
|
||||
require.True(t, shouldApplyOpenAIAlphaSearchAccountErrorSideEffects(http.StatusForbidden))
|
||||
|
||||
Reference in New Issue
Block a user