diff --git a/backend/internal/infra/provider/web/image.go b/backend/internal/infra/provider/web/image.go index a94f198a1..b111a42bb 100644 --- a/backend/internal/infra/provider/web/image.go +++ b/backend/internal/infra/provider/web/image.go @@ -385,8 +385,14 @@ func (a *Adapter) generateLiteImageURL(ctx context.Context, credential account.C if upstream.StatusCode < 200 || upstream.StatusCode >= 300 { body, _ := io.ReadAll(io.LimitReader(upstream.Body, 1<<20)) _ = upstream.Body.Close() - if upstream.StatusCode == http.StatusForbidden { - if attempt == 0 && a.invalidateSignedStatsig(http.MethodPost, statsigTarget) { + if upstream.StatusCode == http.StatusForbidden && attempt == 0 { + upstreamErr := newWebMediaUpstreamError(upstream.StatusCode, body, false) + if isClearanceRefreshableMediaError(upstreamErr) { + // The failed WebSocket handshake invalidates the current browser + // session. Statsig is independent and must not gate reacquiring + // a fresh lease for the retry. + lease.InvalidateClearance() + _ = a.invalidateSignedStatsig(http.MethodPost, statsigTarget) lease.Release() continue } @@ -420,7 +426,11 @@ func (a *Adapter) generateLiteImageURL(ctx context.Context, credential account.C status := 0 if errors.Is(consumeErr, errWebAntiBot) { status = http.StatusForbidden - if attempt == 0 && a.invalidateSignedStatsig(http.MethodPost, statsigTarget) { + // A challenge can arrive inside an otherwise successful stream, + // so the handshake path cannot invalidate it for us. + lease.InvalidateClearance() + if attempt == 0 { + _ = a.invalidateSignedStatsig(http.MethodPost, statsigTarget) lease.Release() continue } @@ -569,6 +579,24 @@ func liteImageMarkdown(item map[string]any) string { } func (a *Adapter) generateWSImage(ctx context.Context, request provider.ImageGenerationRequest, count int, format, ratio, resolution string, modelConfig imagineModelConfig) (*provider.Response, error) { + for attempt := 0; attempt < 2; attempt++ { + response, err := a.generateWSImageAttempt(ctx, request, count, format, ratio, resolution, modelConfig) + if err == nil { + return response, nil + } + var upstreamErr *webMediaUpstreamError + if !errors.As(err, &upstreamErr) || !isClearanceRefreshableMediaError(upstreamErr) || attempt > 0 { + if errors.As(err, &upstreamErr) { + return upstreamErr.providerResponse(), nil + } + return nil, err + } + a.log().Warn("web_image_clearance_retry", "operation", "imagine", "status", upstreamErr.status, "body_kind", upstreamErr.bodyKind) + } + return nil, fmt.Errorf("Imagine WebSocket Clearance 重试耗尽") +} + +func (a *Adapter) generateWSImageAttempt(ctx context.Context, request provider.ImageGenerationRequest, count int, format, ratio, resolution string, modelConfig imagineModelConfig) (*provider.Response, error) { cfg := a.config() token, err := a.cipher.Decrypt(request.Credential.EncryptedAccessToken) if err != nil { @@ -602,6 +630,23 @@ func (a *Adapter) generateWSImage(ctx context.Context, request provider.ImageGen status = response.StatusCode } a.egress.Feedback(context.WithoutCancel(ctx), lease.NodeID, status, err) + if response != nil { + var body []byte + if response.Body != nil { + body, _ = io.ReadAll(io.LimitReader(response.Body, webMediaDiagnosticBodyLimit+1)) + _ = response.Body.Close() + } + truncated := len(body) > webMediaDiagnosticBodyLimit + if truncated { + body = body[:webMediaDiagnosticBodyLimit] + } + upstreamErr := newWebMediaUpstreamError(response.StatusCode, body, truncated) + a.logWebMediaUpstreamRejection("image_imagine_handshake", &http.Response{ + StatusCode: response.StatusCode, + Header: http.Header(response.Header).Clone(), + }, upstreamErr) + return nil, upstreamErr + } return nil, fmt.Errorf("连接 Imagine WebSocket: %w", err) } connectionOwned := true @@ -677,7 +722,28 @@ func (a *Adapter) generateWSImage(ctx context.Context, request provider.ImageGen return result, err } +// EditImage retries the complete browser media flow once after a challenge +// response. Reacquiring the lease is required because the failed lease keeps +// the immutable browser-session cookies that were rejected upstream. func (a *Adapter) EditImage(ctx context.Context, request provider.ImageEditRequest) (*provider.Response, error) { + for attempt := 0; attempt < 2; attempt++ { + response, err := a.editImageAttempt(ctx, request) + if err == nil { + return response, nil + } + var upstreamErr *webMediaUpstreamError + if !errors.As(err, &upstreamErr) || !isClearanceRefreshableMediaError(upstreamErr) || attempt > 0 { + if errors.As(err, &upstreamErr) { + return upstreamErr.providerResponse(), nil + } + return nil, err + } + a.log().Warn("web_image_clearance_retry", "operation", "edit", "status", upstreamErr.status, "body_kind", upstreamErr.bodyKind) + } + return nil, fmt.Errorf("图片编辑 Clearance 重试耗尽") +} + +func (a *Adapter) editImageAttempt(ctx context.Context, request provider.ImageEditRequest) (*provider.Response, error) { if strings.TrimSpace(request.Quality) != "" { return invalidImageRequest("Grok Web 图片模型不支持 quality") } @@ -763,9 +829,18 @@ func (a *Adapter) EditImage(ctx context.Context, request provider.ImageEditReque return nil, err } if response.StatusCode < 200 || response.StatusCode >= 300 { - body, _ := io.ReadAll(io.LimitReader(response.Body, 1<<20)) + body, _ := io.ReadAll(io.LimitReader(response.Body, webMediaDiagnosticBodyLimit+1)) _ = response.Body.Close() - return &provider.Response{StatusCode: response.StatusCode, Status: response.Status, Header: jsonHeaders(), Body: io.NopCloser(bytes.NewReader(body))}, nil + truncated := len(body) > webMediaDiagnosticBodyLimit + if truncated { + body = body[:webMediaDiagnosticBodyLimit] + } + upstreamErr := newWebMediaUpstreamError(response.StatusCode, body, truncated) + a.logWebMediaUpstreamRejection("image_edit_generate", response, upstreamErr) + if isClearanceRefreshableMediaError(upstreamErr) { + lease.InvalidateClearance() + } + return nil, upstreamErr } if request.Streaming { reader, writer := io.Pipe() @@ -1687,7 +1762,14 @@ func imagineURL(baseURL string) (string, error) { if err != nil { return "", err } - value.Scheme = "wss" + switch value.Scheme { + case "https": + value.Scheme = "wss" + case "http": + value.Scheme = "ws" + default: + return "", fmt.Errorf("Grok Web Base URL 协议无效") + } value.Path = "/ws/imagine/listen" value.RawQuery = "" return value.String(), nil diff --git a/backend/internal/infra/provider/web/media_clearance_retry_test.go b/backend/internal/infra/provider/web/media_clearance_retry_test.go new file mode 100644 index 000000000..dac2db2f9 --- /dev/null +++ b/backend/internal/infra/provider/web/media_clearance_retry_test.go @@ -0,0 +1,194 @@ +package web + +import ( + "context" + "encoding/base64" + "io" + "net/http" + "strings" + "sync/atomic" + "testing" + + fhttp "github.com/bogdanfinn/fhttp" + fhttptest "github.com/bogdanfinn/fhttp/httptest" + "github.com/bogdanfinn/websocket" + + "github.com/chenyme/grok2api/backend/internal/domain/account" + infraegress "github.com/chenyme/grok2api/backend/internal/infra/egress" + "github.com/chenyme/grok2api/backend/internal/infra/provider" + "github.com/chenyme/grok2api/backend/internal/infra/security" +) + +func TestIsClearanceRefreshableMediaError(t *testing.T) { + tests := []struct { + name string + body string + code int + want bool + }{ + {name: "empty challenge response", code: http.StatusForbidden, want: true}, + {name: "cloudflare html", code: http.StatusForbidden, body: "Just a moment...", want: true}, + {name: "structured moderation response", code: http.StatusForbidden, body: `{"error":{"code":"content-moderated","message":"rejected"}}`, want: false}, + {name: "server failure", code: http.StatusBadGateway, want: false}, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + err := newWebMediaUpstreamError(test.code, []byte(test.body), false) + if got := isClearanceRefreshableMediaError(err); got != test.want { + t.Fatalf("refreshable=%v, want %v (kind=%q challenge=%v)", got, test.want, err.bodyKind, err.cloudflareChallenge) + } + }) + } +} + +func TestWebMediaUpstreamErrorProviderResponseIsBounded(t *testing.T) { + err := newWebMediaUpstreamError(http.StatusForbidden, nil, false) + response := err.providerResponse() + if response == nil || response.StatusCode != http.StatusForbidden { + t.Fatalf("response=%#v", response) + } + body, readErr := io.ReadAll(response.Body) + _ = response.Body.Close() + if readErr != nil { + t.Fatal(readErr) + } + if !strings.Contains(string(body), "upstream_forbidden") || !strings.Contains(string(body), "Grok Web") { + t.Fatalf("body=%s", body) + } +} + +func TestGenerateWSImageReacquiresAfterChallengeHandshake(t *testing.T) { + var handshakes atomic.Int32 + server := fhttptest.NewServer(fhttp.HandlerFunc(func(writer fhttp.ResponseWriter, request *fhttp.Request) { + if request.URL.Path != "/ws/imagine/listen" { + fhttp.NotFound(writer, request) + return + } + if handshakes.Add(1) == 1 { + writer.Header().Set("Content-Type", "text/html") + writer.WriteHeader(http.StatusForbidden) + _, _ = writer.Write([]byte("Just a moment...")) + return + } + connection, err := (&websocket.Upgrader{CheckOrigin: func(*fhttp.Request) bool { return true }}).Upgrade(writer, request, nil) + if err != nil { + t.Errorf("upgrade Imagine WebSocket: %v", err) + return + } + defer connection.Close() + for range 2 { + var message map[string]any + if err := connection.ReadJSON(&message); err != nil { + t.Errorf("read Imagine request: %v", err) + return + } + } + _ = connection.WriteJSON(map[string]any{ + "type": "image", "id": "image-1", "blob": "aW1hZ2U=", "percentage_complete": 100, "grid_index": 0, + }) + _ = connection.WriteJSON(map[string]any{ + "type": "json", "id": "image-1", "current_status": "completed", "moderated": false, "order": 0, + }) + })) + defer server.Close() + + adapter, credential := testMediaAdapter(t, server.URL) + response, err := adapter.GenerateImage(context.Background(), provider.ImageGenerationRequest{ + Credential: credential, Model: "grok-imagine-image-quality", Prompt: "draw a teapot", Count: 1, ResponseFormat: "b64_json", + }) + if err != nil { + t.Fatal(err) + } + defer response.Body.Close() + body, readErr := io.ReadAll(response.Body) + if readErr != nil || response.StatusCode != http.StatusOK || !strings.Contains(string(body), `"b64_json"`) { + t.Fatalf("status=%d body=%s err=%v", response.StatusCode, body, readErr) + } + if got := handshakes.Load(); got != 2 { + t.Fatalf("handshakes=%d, want 2", got) + } +} + +func TestGenerateLiteImageReacquiresAfterChallengeHandshake(t *testing.T) { + var handshakes atomic.Int32 + server := fhttptest.NewServer(fhttp.HandlerFunc(func(writer fhttp.ResponseWriter, request *fhttp.Request) { + if request.URL.Path != "/ws/mgw/" { + fhttp.NotFound(writer, request) + return + } + if handshakes.Add(1) == 1 { + writer.Header().Set("Content-Type", "text/html") + writer.WriteHeader(http.StatusForbidden) + _, _ = writer.Write([]byte("Just a moment...")) + return + } + connection, err := (&websocket.Upgrader{CheckOrigin: func(*fhttp.Request) bool { return true }}).Upgrade(writer, request, nil) + if err != nil { + t.Errorf("upgrade Gateway WebSocket: %v", err) + return + } + defer connection.Close() + var initial map[string]any + if err := connection.ReadJSON(&initial); err != nil { + t.Errorf("read Gateway session: %v", err) + return + } + event, _ := initial["event"].(map[string]any) + eventID, _ := event["event_id"].(string) + _ = connection.WriteJSON(map[string]any{ + "session_id": "session-1", "event": map[string]any{"type": "session.created", "client_event_id": eventID}, + }) + _ = connection.WriteJSON(map[string]any{ + "session_id": "session-1", "event": map[string]any{"type": "conversation.attached", "conversation": map[string]any{"id": "session-1"}}, + }) + for range 2 { + var message map[string]any + if err := connection.ReadJSON(&message); err != nil { + t.Errorf("read Gateway turn: %v", err) + return + } + } + _ = connection.WriteJSON(map[string]any{ + "session_id": "session-1", + "event": map[string]any{ + "type": "response.grok.output", + "output": map[string]any{"card_attachment": map[string]any{"jsonData": map[string]any{ + "id": "card-1", "image_chunk": map[string]any{"progress": 100, "imageUrl": "users/test/generated/image.jpg", "moderated": false}, + }}}, + }, + }) + })) + defer server.Close() + + adapter, credential := testMediaAdapter(t, server.URL) + credential.UserID = "497f19f8-49d4-458a-bee4-43ec3dcaf8ca" + spec, ok := Resolve("grok-imagine-image") + if !ok { + t.Fatal("missing Lite image model") + } + rawURL, err := adapter.generateLiteImageURL(context.Background(), credential, spec, "draw a teapot") + if err != nil { + t.Fatal(err) + } + if rawURL != "https://assets.grok.com/users/test/generated/image.jpg" { + t.Fatalf("url=%q", rawURL) + } + if got := handshakes.Load(); got != 2 { + t.Fatalf("handshakes=%d, want 2", got) + } +} + +func testMediaAdapter(t *testing.T, baseURL string) (*Adapter, account.Credential) { + t.Helper() + cipher, err := security.NewCipher(base64.StdEncoding.EncodeToString(make([]byte, 32))) + if err != nil { + t.Fatal(err) + } + encrypted, err := cipher.Encrypt("test-sso") + if err != nil { + t.Fatal(err) + } + adapter := NewAdapter(Config{BaseURL: baseURL, StatsigMode: "manual", ChatTimeoutSeconds: 5}, infraegress.NewManager(egressRepositoryStub{}, cipher), cipher, nil, imageAssetStoreStub{}) + credential := account.Credential{ID: 1, Provider: account.ProviderWeb, EncryptedAccessToken: encrypted} + return adapter, credential +} diff --git a/backend/internal/infra/provider/web/video.go b/backend/internal/infra/provider/web/video.go index c0c502a67..81aa12792 100644 --- a/backend/internal/infra/provider/web/video.go +++ b/backend/internal/infra/provider/web/video.go @@ -44,6 +44,32 @@ func (e *webMediaUpstreamError) HTTPStatusCode() int { return e.status } +// isClearanceRefreshableMediaError distinguishes browser-session challenges +// from structured upstream policy responses such as content moderation. Empty +// and HTML 403 bodies are the forms returned by the media endpoints when the +// request is rejected before the application response is built. +func isClearanceRefreshableMediaError(e *webMediaUpstreamError) bool { + if e == nil || e.status != http.StatusForbidden { + return false + } + return e.cloudflareChallenge || e.bodyKind == "empty" || e.bodyKind == "html" +} + +func (e *webMediaUpstreamError) providerResponse() *provider.Response { + if e == nil { + return nil + } + code := "upstream_forbidden" + if e.status != http.StatusForbidden { + code = "upstream_unavailable" + } + return jsonProviderResponse(e.status, map[string]any{"error": map[string]any{ + "message": e.summary, + "type": "upstream_error", + "code": code, + }}) +} + const ( webMediaDiagnosticBodyLimit = 64 << 10 webMediaDiagnosticSummaryLimit = 256