mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-24 16:05:44 +08:00
Bind OpenAI HTTP response IDs to selected accounts
This commit is contained in:
@@ -3040,6 +3040,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
|
||||
// Handle normal response
|
||||
var usage *OpenAIUsage
|
||||
var firstTokenMs *int
|
||||
responseID := ""
|
||||
imageCount := 0
|
||||
var imageOutputSizes []string
|
||||
if reqStream {
|
||||
@@ -3049,6 +3050,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
|
||||
}
|
||||
usage = streamResult.usage
|
||||
firstTokenMs = streamResult.firstTokenMs
|
||||
responseID = strings.TrimSpace(streamResult.responseID)
|
||||
imageCount = streamResult.imageCount
|
||||
imageOutputSizes = streamResult.imageOutputSizes
|
||||
} else {
|
||||
@@ -3057,9 +3059,11 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
|
||||
return nil, err
|
||||
}
|
||||
usage = nonStreamResult.usage
|
||||
responseID = strings.TrimSpace(nonStreamResult.responseID)
|
||||
imageCount = nonStreamResult.imageCount
|
||||
imageOutputSizes = nonStreamResult.imageOutputSizes
|
||||
}
|
||||
s.bindHTTPResponseAccount(ctx, c, account, responseID)
|
||||
|
||||
// Extract and save Codex usage snapshot from response headers (for OAuth accounts)
|
||||
if account.Type == AccountTypeOAuth {
|
||||
@@ -3074,6 +3078,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
|
||||
|
||||
forwardResult := &OpenAIForwardResult{
|
||||
RequestID: resp.Header.Get("x-request-id"),
|
||||
ResponseID: responseID,
|
||||
Usage: *usage,
|
||||
Model: originalModel,
|
||||
UpstreamModel: upstreamModel,
|
||||
@@ -3286,6 +3291,7 @@ func (s *OpenAIGatewayService) forwardOpenAIPassthrough(
|
||||
|
||||
var usage *OpenAIUsage
|
||||
var firstTokenMs *int
|
||||
responseID := ""
|
||||
imageCount := 0
|
||||
var imageOutputSizes []string
|
||||
if reqStream {
|
||||
@@ -3295,6 +3301,7 @@ func (s *OpenAIGatewayService) forwardOpenAIPassthrough(
|
||||
}
|
||||
usage = result.usage
|
||||
firstTokenMs = result.firstTokenMs
|
||||
responseID = strings.TrimSpace(result.responseID)
|
||||
imageCount = result.imageCount
|
||||
imageOutputSizes = result.imageOutputSizes
|
||||
} else {
|
||||
@@ -3303,9 +3310,11 @@ func (s *OpenAIGatewayService) forwardOpenAIPassthrough(
|
||||
return nil, err
|
||||
}
|
||||
usage = result.usage
|
||||
responseID = strings.TrimSpace(result.responseID)
|
||||
imageCount = result.imageCount
|
||||
imageOutputSizes = result.imageOutputSizes
|
||||
}
|
||||
s.bindHTTPResponseAccount(ctx, c, account, responseID)
|
||||
|
||||
if snapshot := ParseCodexRateLimitHeaders(resp.Header); snapshot != nil {
|
||||
s.updateCodexUsageSnapshot(ctx, account.ID, snapshot)
|
||||
@@ -3317,6 +3326,7 @@ func (s *OpenAIGatewayService) forwardOpenAIPassthrough(
|
||||
|
||||
forwardResult := &OpenAIForwardResult{
|
||||
RequestID: resp.Header.Get("x-request-id"),
|
||||
ResponseID: responseID,
|
||||
Usage: *usage,
|
||||
Model: reqModel,
|
||||
UpstreamModel: upstreamPassthroughModel,
|
||||
@@ -3631,6 +3641,7 @@ func collectOpenAIPassthroughTimeoutHeaders(h http.Header) []string {
|
||||
type openaiStreamingResultPassthrough struct {
|
||||
usage *OpenAIUsage
|
||||
firstTokenMs *int
|
||||
responseID string
|
||||
imageCount int
|
||||
imageOutputSizes []string
|
||||
}
|
||||
@@ -3638,6 +3649,7 @@ type openaiStreamingResultPassthrough struct {
|
||||
type openaiNonStreamingResultPassthrough struct {
|
||||
*OpenAIUsage
|
||||
usage *OpenAIUsage
|
||||
responseID string
|
||||
imageCount int
|
||||
imageOutputSizes []string
|
||||
}
|
||||
@@ -3781,6 +3793,7 @@ func (s *OpenAIGatewayService) handleStreamingResponsePassthrough(
|
||||
usage := &OpenAIUsage{}
|
||||
imageCounter := newOpenAIImageOutputCounter()
|
||||
var firstTokenMs *int
|
||||
responseID := ""
|
||||
clientDisconnected := false
|
||||
sawDone := false
|
||||
sawTerminalEvent := false
|
||||
@@ -3815,6 +3828,7 @@ func (s *OpenAIGatewayService) handleStreamingResponsePassthrough(
|
||||
return &openaiStreamingResultPassthrough{
|
||||
usage: usage,
|
||||
firstTokenMs: firstTokenMs,
|
||||
responseID: responseID,
|
||||
imageCount: imageCounter.Count(),
|
||||
imageOutputSizes: imageCounter.Sizes(),
|
||||
}
|
||||
@@ -3850,6 +3864,9 @@ func (s *OpenAIGatewayService) handleStreamingResponsePassthrough(
|
||||
if openAIStreamEventIsTerminal(trimmedData) {
|
||||
sawTerminalEvent = true
|
||||
}
|
||||
if responseID == "" {
|
||||
responseID = extractOpenAIResponseIDFromJSONBytes(dataBytes)
|
||||
}
|
||||
imageCounter.AddSSEData(dataBytes)
|
||||
lineStartsClientOutput = forceFlushFailedEvent || openAIStreamDataStartsClientOutput(trimmedData, eventType)
|
||||
if firstTokenMs == nil && lineStartsClientOutput && trimmedData != "[DONE]" {
|
||||
@@ -3976,6 +3993,7 @@ func (s *OpenAIGatewayService) handleNonStreamingResponsePassthrough(
|
||||
return &openaiNonStreamingResultPassthrough{
|
||||
OpenAIUsage: usage,
|
||||
usage: usage,
|
||||
responseID: extractOpenAIResponseIDFromJSONBytes(body),
|
||||
imageCount: countOpenAIResponseImageOutputsFromJSONBytes(body),
|
||||
imageOutputSizes: collectOpenAIResponseImageOutputSizesFromJSONBytes(body),
|
||||
}, nil
|
||||
@@ -4039,6 +4057,7 @@ func (s *OpenAIGatewayService) handlePassthroughSSEToJSON(resp *http.Response, c
|
||||
return &openaiNonStreamingResultPassthrough{
|
||||
OpenAIUsage: usage,
|
||||
usage: usage,
|
||||
responseID: extractOpenAIResponseIDFromJSONBytes(body),
|
||||
imageCount: countOpenAIImageOutputsFromSSEBody(bodyText),
|
||||
imageOutputSizes: collectOpenAIImageOutputSizesFromSSEBody(bodyText),
|
||||
}, nil
|
||||
@@ -4500,6 +4519,7 @@ func (s *OpenAIGatewayService) handleCompatErrorResponse(
|
||||
type openaiStreamingResult struct {
|
||||
usage *OpenAIUsage
|
||||
firstTokenMs *int
|
||||
responseID string
|
||||
imageCount int
|
||||
imageOutputSizes []string
|
||||
}
|
||||
@@ -4507,6 +4527,7 @@ type openaiStreamingResult struct {
|
||||
type openaiNonStreamingResult struct {
|
||||
*OpenAIUsage
|
||||
usage *OpenAIUsage
|
||||
responseID string
|
||||
imageCount int
|
||||
imageOutputSizes []string
|
||||
}
|
||||
@@ -4544,6 +4565,7 @@ func (s *OpenAIGatewayService) handleStreamingResponse(ctx context.Context, resp
|
||||
usage := &OpenAIUsage{}
|
||||
imageCounter := newOpenAIImageOutputCounter()
|
||||
var firstTokenMs *int
|
||||
responseID := ""
|
||||
scanner := bufio.NewScanner(resp.Body)
|
||||
maxLineSize := defaultMaxLineSize
|
||||
if s.cfg != nil && s.cfg.Gateway.MaxLineSize > 0 {
|
||||
@@ -4624,6 +4646,7 @@ func (s *OpenAIGatewayService) handleStreamingResponse(ctx context.Context, resp
|
||||
return &openaiStreamingResult{
|
||||
usage: usage,
|
||||
firstTokenMs: firstTokenMs,
|
||||
responseID: responseID,
|
||||
imageCount: imageCounter.Count(),
|
||||
imageOutputSizes: imageCounter.Sizes(),
|
||||
}
|
||||
@@ -4710,6 +4733,9 @@ func (s *OpenAIGatewayService) handleStreamingResponse(ctx context.Context, resp
|
||||
sawTerminalEvent = true
|
||||
}
|
||||
eventType := strings.TrimSpace(gjson.GetBytes(dataBytes, "type").String())
|
||||
if responseID == "" {
|
||||
responseID = extractOpenAIResponseIDFromJSONBytes(dataBytes)
|
||||
}
|
||||
forceFlushFailedEvent := false
|
||||
if eventType == "response.failed" {
|
||||
failedMessage = extractOpenAISSEErrorMessage(dataBytes)
|
||||
@@ -5047,6 +5073,33 @@ func extractOpenAIUsageFromJSONBytes(body []byte) (OpenAIUsage, bool) {
|
||||
return openAIUsageFromGJSON(gjson.GetBytes(body, "response.usage"))
|
||||
}
|
||||
|
||||
func extractOpenAIResponseIDFromJSONBytes(body []byte) string {
|
||||
if len(body) == 0 || !gjson.ValidBytes(body) {
|
||||
return ""
|
||||
}
|
||||
if id := strings.TrimSpace(gjson.GetBytes(body, "id").String()); id != "" {
|
||||
return id
|
||||
}
|
||||
return strings.TrimSpace(gjson.GetBytes(body, "response.id").String())
|
||||
}
|
||||
|
||||
func (s *OpenAIGatewayService) bindHTTPResponseAccount(ctx context.Context, c *gin.Context, account *Account, responseID string) {
|
||||
if s == nil || account == nil || account.ID <= 0 {
|
||||
return
|
||||
}
|
||||
responseID = strings.TrimSpace(responseID)
|
||||
if responseID == "" {
|
||||
return
|
||||
}
|
||||
store := s.getOpenAIWSStateStore()
|
||||
if store == nil {
|
||||
return
|
||||
}
|
||||
groupID := getOpenAIGroupIDFromContext(c)
|
||||
ttl := s.openAIWSResponseStickyTTL()
|
||||
logOpenAIWSBindResponseAccountWarn(groupID, account.ID, responseID, store.BindResponseAccount(ctx, groupID, responseID, account.ID, ttl))
|
||||
}
|
||||
|
||||
func openAIUsageFromGJSON(value gjson.Result) (OpenAIUsage, bool) {
|
||||
if !value.Exists() || !value.IsObject() {
|
||||
return OpenAIUsage{}, false
|
||||
@@ -5127,6 +5180,7 @@ func (s *OpenAIGatewayService) handleNonStreamingResponse(ctx context.Context, r
|
||||
return &openaiNonStreamingResult{
|
||||
OpenAIUsage: usage,
|
||||
usage: usage,
|
||||
responseID: extractOpenAIResponseIDFromJSONBytes(body),
|
||||
imageCount: countOpenAIResponseImageOutputsFromJSONBytes(body),
|
||||
imageOutputSizes: collectOpenAIResponseImageOutputSizesFromJSONBytes(body),
|
||||
}, nil
|
||||
@@ -5192,6 +5246,7 @@ func (s *OpenAIGatewayService) handleSSEToJSON(resp *http.Response, c *gin.Conte
|
||||
return &openaiNonStreamingResult{
|
||||
OpenAIUsage: usage,
|
||||
usage: usage,
|
||||
responseID: extractOpenAIResponseIDFromJSONBytes(body),
|
||||
imageCount: countOpenAIImageOutputsFromSSEBody(bodyText),
|
||||
imageOutputSizes: collectOpenAIImageOutputSizesFromSSEBody(bodyText),
|
||||
}, nil
|
||||
|
||||
@@ -227,6 +227,30 @@ func TestOpenAIGatewayService_GenerateSessionHash_AttachesLegacyHashToContext(t
|
||||
require.NotEmpty(t, openAILegacySessionHashFromContext(c.Request.Context()))
|
||||
}
|
||||
|
||||
func TestExtractOpenAIResponseIDFromJSONBytes(t *testing.T) {
|
||||
require.Equal(t, "resp_json", extractOpenAIResponseIDFromJSONBytes([]byte(`{"id":"resp_json"}`)))
|
||||
require.Equal(t, "resp_sse", extractOpenAIResponseIDFromJSONBytes([]byte(`{"type":"response.completed","response":{"id":"resp_sse"}}`)))
|
||||
require.Empty(t, extractOpenAIResponseIDFromJSONBytes([]byte(`{"response":{}}`)))
|
||||
require.Empty(t, extractOpenAIResponseIDFromJSONBytes([]byte(`not-json`)))
|
||||
}
|
||||
|
||||
func TestOpenAIGatewayService_BindHTTPResponseAccount(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
rec := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(rec)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil)
|
||||
groupID := int64(4201)
|
||||
c.Set("api_key", &APIKey{ID: 501, GroupID: &groupID})
|
||||
|
||||
svc := &OpenAIGatewayService{}
|
||||
account := &Account{ID: 37001, Platform: PlatformOpenAI, Type: AccountTypeAPIKey}
|
||||
svc.bindHTTPResponseAccount(context.Background(), c, account, "resp_http_001")
|
||||
|
||||
got, err := svc.getOpenAIWSStateStore().GetResponseAccount(context.Background(), groupID, "resp_http_001")
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, account.ID, got)
|
||||
}
|
||||
|
||||
func TestOpenAIGatewayService_GenerateExplicitSessionHash_SkipsContentFallback(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
svc := &OpenAIGatewayService{}
|
||||
|
||||
Reference in New Issue
Block a user