From 3fe7b4da78715389a60c07bcf1ba02da044f50cc Mon Sep 17 00:00:00 2001 From: cat Date: Tue, 14 Jul 2026 17:01:53 +0800 Subject: [PATCH] =?UTF-8?q?fix(openai):=20=E9=97=AD=E5=90=88=E7=9B=B4?= =?UTF-8?q?=E9=80=9A=20WS=20=E6=81=A2=E5=A4=8D=E4=B8=8E=E9=94=99=E8=AF=AF?= =?UTF-8?q?=E8=84=B1=E6=95=8F?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../internal/service/openai_agent_identity.go | 22 +++++++++++++++++++ .../service/openai_agent_identity_test.go | 20 +++++++++++++---- .../service/openai_gateway_upstream_errors.go | 1 + backend/internal/service/openai_images.go | 1 + .../openai_ws_v2_passthrough_adapter.go | 9 ++++++-- 5 files changed, 47 insertions(+), 6 deletions(-) diff --git a/backend/internal/service/openai_agent_identity.go b/backend/internal/service/openai_agent_identity.go index 8383a10e9c..a5ea9b7612 100644 --- a/backend/internal/service/openai_agent_identity.go +++ b/backend/internal/service/openai_agent_identity.go @@ -261,6 +261,25 @@ func ensureAgentIdentityTaskForAccount(ctx context.Context, repo AccountReposito } sharedTaskMu.Lock() defer sharedTaskMu.Unlock() + // Re-read inside the shared lock. Different request paths often receive + // independent repository snapshots; checking only the caller's snapshot + // would allow sequential duplicate registrations after the first writer + // has already persisted a new task. + if repo != nil && credAccount.ID > 0 { + if refreshed, refreshErr := repo.GetByID(ctx, credAccount.ID); refreshErr == nil && refreshed != nil { + if refreshed.IsShadow() { + if resolved, resolveErr := resolveCredentialAccount(ctx, repo, refreshed); resolveErr == nil && resolved != nil { + refreshed = resolved + } + } + if refreshed.IsOpenAIAgentIdentity() { + credAccount = refreshed + if !account.IsShadow() { + account.Credentials = shallowCopyMap(credAccount.Credentials) + } + } + } + } currentTaskID = strings.TrimSpace(credAccount.GetCredential("task_id")) if currentTaskID != "" && (expectedTaskID == "" || currentTaskID != expectedTaskID) { return nil @@ -277,6 +296,9 @@ func ensureAgentIdentityTaskForAccount(ctx context.Context, repo AccountReposito if err := persistAccountCredentials(ctx, repo, credAccount, credentials); err != nil { return err } + if !account.IsShadow() && account != credAccount { + account.Credentials = shallowCopyMap(credAccount.Credentials) + } if pool != nil { pool.ClearAccount(credAccount.ID) } diff --git a/backend/internal/service/openai_agent_identity_test.go b/backend/internal/service/openai_agent_identity_test.go index 6e94f53abe..12f5a7766a 100644 --- a/backend/internal/service/openai_agent_identity_test.go +++ b/backend/internal/service/openai_agent_identity_test.go @@ -163,7 +163,7 @@ func TestEnsureAgentIdentityTaskSharesLockAcrossServicesForSameAccount(t *testin "agent_runtime_id": key.runtimeID, "agent_private_key": privateKey, }} - repo := &agentIdentityCredentialsRepo{} + repo := &agentIdentityCredentialsRepo{account: account} registerCalls := 0 var registerMu sync.Mutex server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { @@ -179,10 +179,11 @@ func TestEnsureAgentIdentityTaskSharesLockAcrossServicesForSameAccount(t *testin start := make(chan struct{}) errors := make(chan error, 2) - for range 2 { + requests := []*Account{cloneAgentIdentityTestAccount(account), cloneAgentIdentityTestAccount(account)} + for _, request := range requests { go func() { <-start - errors <- ensureAgentIdentityTaskForAccount(context.Background(), repo, nil, &sync.Mutex{}, account, "") + errors <- ensureAgentIdentityTaskForAccount(context.Background(), repo, nil, &sync.Mutex{}, request, "") }() } close(start) @@ -191,15 +192,26 @@ func TestEnsureAgentIdentityTaskSharesLockAcrossServicesForSameAccount(t *testin registerMu.Lock() defer registerMu.Unlock() require.Equal(t, 1, registerCalls) - require.Equal(t, "task-shared", account.GetCredential("task_id")) + require.Equal(t, "task-shared", repo.account.GetCredential("task_id")) +} + +func cloneAgentIdentityTestAccount(account *Account) *Account { + copy := *account + copy.Credentials = shallowCopyMap(account.Credentials) + return © } type agentIdentityCredentialsRepo struct { AccountRepository credentials map[string]any + account *Account mu sync.Mutex } +func (r *agentIdentityCredentialsRepo) GetByID(_ context.Context, _ int64) (*Account, error) { + return r.account, nil +} + func (r *agentIdentityCredentialsRepo) UpdateCredentials(_ context.Context, _ int64, credentials map[string]any) error { r.mu.Lock() defer r.mu.Unlock() diff --git a/backend/internal/service/openai_gateway_upstream_errors.go b/backend/internal/service/openai_gateway_upstream_errors.go index bc9bf91fef..d8823fe4ed 100644 --- a/backend/internal/service/openai_gateway_upstream_errors.go +++ b/backend/internal/service/openai_gateway_upstream_errors.go @@ -471,6 +471,7 @@ func (s *OpenAIGatewayService) handleCompatErrorResponse( requestedModel ...string, ) (*OpenAIForwardResult, error) { body := s.readUpstreamErrorBody(resp) + body = s.redactAgentIdentitySensitiveBody(context.Background(), account, body) // cyber_policy:兼容路径(Chat Completions / Anthropic)以各自格式回写错误, // 不原样透传 responses 格式的 cyber body(否则对下游格式不合法)。cyber 是上游网络 diff --git a/backend/internal/service/openai_images.go b/backend/internal/service/openai_images.go index b411cb3104..081b5e741b 100644 --- a/backend/internal/service/openai_images.go +++ b/backend/internal/service/openai_images.go @@ -632,6 +632,7 @@ func (s *OpenAIGatewayService) forwardOpenAIImagesAPIKey( if resp.StatusCode >= 400 { respBody := s.readUpstreamErrorBody(resp) _ = resp.Body.Close() + respBody = s.redactAgentIdentitySensitiveBody(upstreamCtx, account, respBody) resp.Body = io.NopCloser(bytes.NewReader(respBody)) upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(respBody)) upstreamMsg = sanitizeUpstreamErrorMessage(upstreamMsg) diff --git a/backend/internal/service/openai_ws_v2_passthrough_adapter.go b/backend/internal/service/openai_ws_v2_passthrough_adapter.go index bb097e540b..75df7191ac 100644 --- a/backend/internal/service/openai_ws_v2_passthrough_adapter.go +++ b/backend/internal/service/openai_ws_v2_passthrough_adapter.go @@ -374,8 +374,13 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough( if err == nil { break } - var dialErr *openAIWSDialError - if s.isAgentIdentityAccount(ctx, account) && errors.As(err, &dialErr) && isAgentIdentityTaskInvalidWSDialError(dialErr) && !agentTaskRecoveryTried { + var handshakeErr *openAIWSHandshakeError + responseBody := []byte(nil) + if errors.As(err, &handshakeErr) && handshakeErr != nil { + responseBody = handshakeErr.Body + } + dialErr := &openAIWSDialError{StatusCode: statusCode, ResponseBody: responseBody, Err: err} + if s.isAgentIdentityAccount(ctx, account) && isAgentIdentityTaskInvalidWSDialError(dialErr) && !agentTaskRecoveryTried { agentTaskRecoveryTried = true if recoveryErr := s.recoverAgentIdentityTask(ctx, account, account.GetCredential("task_id")); recoveryErr != nil { return fmt.Errorf("agent identity task recovery failed: %w", recoveryErr)