Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
94 changes: 88 additions & 6 deletions backend/internal/infra/provider/web/image.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
}
Expand Down Expand Up @@ -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
}
Expand Down Expand Up @@ -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 {
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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")
}
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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
Expand Down
194 changes: 194 additions & 0 deletions backend/internal/infra/provider/web/media_clearance_retry_test.go
Original file line number Diff line number Diff line change
@@ -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: "<!doctype html><title>Just a moment...</title>", 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("<!doctype html><title>Just a moment...</title>"))
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("<!doctype html><title>Just a moment...</title>"))
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
}
26 changes: 26 additions & 0 deletions backend/internal/infra/provider/web/video.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Loading