Bind OpenAI HTTP response IDs to selected accounts

This commit is contained in:
b605166577
2026-06-03 13:07:16 +08:00
parent aa69e3947d
commit 7513b7ea69
2 changed files with 79 additions and 0 deletions
@@ -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{}