diff --git a/internal/openai/eino_http.go b/internal/openai/eino_http.go index 7c2a8bc5..a4d761e0 100644 --- a/internal/openai/eino_http.go +++ b/internal/openai/eino_http.go @@ -6,7 +6,8 @@ import ( "cyberstrike-ai/internal/config" ) -// NewEinoHTTPClient adds OpenAI-compatible request fixes and SSE sanitation. +// NewEinoHTTPClient adds OpenAI-compatible request fixes, SSE error decoding, +// and successful-stream heartbeat sanitation. // Claude channels use Eino's native agenticclaude model and never enter here. func NewEinoHTTPClient(cfg *config.OpenAIConfig, base *http.Client) *http.Client { if base == nil { @@ -18,6 +19,7 @@ func NewEinoHTTPClient(cfg *config.OpenAIConfig, base *http.Client) *http.Client transport = http.DefaultTransport } transport = &reasoningToolChoiceCompatRoundTripper{base: transport, cfg: cfg} + transport = &einoSSEErrorRoundTripper{base: transport} transport = &einoSSESanitizingRoundTripper{base: transport} cloned.Transport = transport return &cloned diff --git a/internal/openai/eino_sse_error.go b/internal/openai/eino_sse_error.go new file mode 100644 index 00000000..5437cada --- /dev/null +++ b/internal/openai/eino_sse_error.go @@ -0,0 +1,94 @@ +package openai + +import ( + "bytes" + "encoding/json" + "io" + "net/http" + "strconv" +) + +const einoSSEErrorMaxBytes = 64 * 1024 + +// Some gateways send SSE even for HTTP errors (occasionally labeled JSON). +// The SDK decodes non-2xx responses as JSON. Unwrap only a validated error +// event, preserving HTTP status and error fields for retry classification. +type einoSSEErrorRoundTripper struct { + base http.RoundTripper +} + +func (rt *einoSSEErrorRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) { + resp, err := rt.base.RoundTrip(req) + if err != nil || resp == nil || resp.Body == nil || resp.StatusCode < 400 { + return resp, err + } + upstream := resp.Body + body, readErr := io.ReadAll(io.LimitReader(upstream, einoSSEErrorMaxBytes+1)) + // Replay all consumed bytes and any read failure if normalization is unsafe. + readers := []io.Reader{bytes.NewReader(body)} + if readErr != nil { + readers = append(readers, sseErrorReadFailure{readErr}) + } + readers = append(readers, upstream) + resp.Body = &sseErrorReplayBody{Reader: io.MultiReader(readers...), Closer: upstream} + if readErr != nil || len(body) > einoSSEErrorMaxBytes { + return resp, nil + } + payload := extractSSEError(body) + if payload == nil { + return resp, nil + } + resp.Body = &sseErrorReplayBody{Reader: bytes.NewReader(payload), Closer: upstream} + resp.Header = resp.Header.Clone() + if resp.Header == nil { + resp.Header = make(http.Header) + } + resp.Header.Set("Content-Type", "application/json") + resp.Header.Set("Content-Length", strconv.Itoa(len(payload))) + resp.Header.Del("Transfer-Encoding") + resp.ContentLength = int64(len(payload)) + resp.TransferEncoding = nil + return resp, nil +} + +type sseErrorReplayBody struct { + io.Reader + io.Closer +} + +type sseErrorReadFailure struct{ err error } + +func (r sseErrorReadFailure) Read([]byte) (int, error) { return 0, r.err } + +// Honor SSE event boundaries and multiline data. Unknown/plain/invalid bodies +// are left intact, rather than replacing useful diagnostics with guesses. +func extractSSEError(body []byte) []byte { + var data []byte + validate := func() []byte { + var envelope struct { + Error *struct { + Message json.RawMessage `json:"message"` + } `json:"error"` + } + if json.Unmarshal(data, &envelope) == nil && envelope.Error != nil && len(envelope.Error.Message) > 0 && !bytes.Equal(envelope.Error.Message, []byte("null")) { + return bytes.TrimSpace(data) + } + return nil + } + for _, line := range bytes.Split(body, []byte("\n")) { + line = bytes.TrimSuffix(line, []byte("\r")) + if len(line) == 0 { + if payload := validate(); payload != nil { + return payload + } + data = nil + continue + } + if bytes.HasPrefix(line, []byte("data:")) { + value := bytes.TrimPrefix(line[5:], []byte(" ")) + data = append(data, value...) + data = append(data, '\n') + } + } + return validate() // Gateways sometimes omit the final blank line. +} diff --git a/internal/openai/eino_sse_error_test.go b/internal/openai/eino_sse_error_test.go new file mode 100644 index 00000000..1673fbb7 --- /dev/null +++ b/internal/openai/eino_sse_error_test.go @@ -0,0 +1,112 @@ +package openai + +import ( + "context" + "errors" + "io" + "net/http" + "strings" + "testing" + + einoopenai "github.com/cloudwego/eino-ext/components/model/openai" + "github.com/cloudwego/eino/schema" +) + +func TestSSEErrorThroughEinoModel(t *testing.T) { + for _, status := range []int{400, 429, 503} { + for _, ct := range []string{"text/event-stream", "application/json", ""} { + t.Run(http.StatusText(status)+"/"+ct, func(t *testing.T) { + srv := newSSEServer(t, `data:{"error":{"code":"400","message":"Invalid request parameters","param":"tool_choice","type":"BadRequestError"}}`, ct, status) + defer srv.Close() + m, err := einoopenai.NewChatModel(context.Background(), &einoopenai.ChatModelConfig{APIKey: "local-test", BaseURL: srv.URL, Model: "local-test", HTTPClient: NewEinoHTTPClient(nil, srv.Client())}) + if err != nil { + t.Fatal(err) + } + _, err = m.Stream(context.Background(), []*schema.Message{schema.UserMessage("test")}) + var apiErr *einoopenai.APIError + if !errors.As(err, &apiErr) { + t.Fatalf("expected structured API error, got %v", err) + } + if apiErr.HTTPStatusCode != status || apiErr.Message != "Invalid request parameters" || apiErr.Type != "BadRequestError" || apiErr.Code != "400" || apiErr.Param == nil || *apiErr.Param != "tool_choice" { + t.Fatalf("error fields changed: %+v", apiErr) + } + }) + } + } +} + +func TestSSEErrorBodyCompatibility(t *testing.T) { + payload := `{"error":{"message":"rejected","extra":"preserved"}}` + for _, tc := range []struct { + name, body, want string + status int + }{ + {"plain JSON", payload, payload, 400}, + {"HTML", "bad gateway", "bad gateway", 502}, + {"malformed", "data:{broken}", "data:{broken}", 400}, + {"not error", "data:{\"choices\":[]}", "data:{\"choices\":[]}", 400}, + {"null error", "data:{\"error\":null}", "data:{\"error\":null}", 400}, + {"SSE", "data:" + payload, payload, 400}, + {"CRLF heartbeat", ": ping\r\nevent: error\r\ndata: " + payload + "\r\n\r\ndata: [DONE]\r\n", payload, 400}, + {"multiline", "data: {\"error\":\ndata: {\"message\":\"rejected\",\"extra\":\"preserved\"}}\n\n", "{\"error\":\n{\"message\":\"rejected\",\"extra\":\"preserved\"}}", 400}, + {"success", "data:" + payload, "data:" + payload, 200}, + {"oversized", "data:" + payload + strings.Repeat(" ", einoSSEErrorMaxBytes), "data:" + payload + strings.Repeat(" ", einoSSEErrorMaxBytes), 400}, + } { + t.Run(tc.name, func(t *testing.T) { + srv := newSSEServer(t, tc.body, "text/event-stream", tc.status) + defer srv.Close() + client := &http.Client{Transport: &einoSSEErrorRoundTripper{base: http.DefaultTransport}} + resp, err := client.Get(srv.URL) + if err != nil { + t.Fatal(err) + } + got := readAll(t, resp.Body) + if got != tc.want { + t.Fatalf("body differs: got %q want %q", got, tc.want) + } + if resp.StatusCode != tc.status { + t.Fatal("status changed") + } + if tc.want != tc.body && (resp.Header.Get("Content-Type") != "application/json" || resp.ContentLength != int64(len(got))) { + t.Fatal("incorrect normalized headers") + } + }) + } +} + +type errorResponseTransport struct{ body io.ReadCloser } + +func (rt errorResponseTransport) RoundTrip(*http.Request) (*http.Response, error) { + return &http.Response{StatusCode: 400, Header: http.Header{"X-Request-Id": []string{"test-id"}}, Body: rt.body}, nil +} + +type failingErrorBody struct { + closed bool + err error +} + +func (b *failingErrorBody) Read(p []byte) (int, error) { + return copy(p, `data:{"error":{"message":"partial"}}`), b.err +} +func (b *failingErrorBody) Close() error { b.closed = true; return nil } +func TestSSEErrorPreservesReadFailureAndClose(t *testing.T) { + sentinel := errors.New("upstream disconnected") + body := &failingErrorBody{err: sentinel} + rt := &einoSSEErrorRoundTripper{base: errorResponseTransport{body: body}} + req, _ := http.NewRequest("GET", "http://local.test", nil) + resp, err := rt.RoundTrip(req) + if err != nil { + t.Fatal(err) + } + got, err := io.ReadAll(resp.Body) + if !errors.Is(err, sentinel) || string(got) != `data:{"error":{"message":"partial"}}` { + t.Fatalf("read failure lost: %q %v", got, err) + } + if resp.Header.Get("X-Request-Id") != "test-id" { + t.Fatal("request ID lost") + } + resp.Body.Close() + if !body.closed { + t.Fatal("upstream not closed") + } +} diff --git a/internal/openai/eino_sse_sanitizer.go b/internal/openai/eino_sse_sanitizer.go index 43e07d5b..340b3a04 100644 --- a/internal/openai/eino_sse_sanitizer.go +++ b/internal/openai/eino_sse_sanitizer.go @@ -54,7 +54,7 @@ func (rt *einoSSESanitizingRoundTripper) RoundTrip(req *http.Request) (*http.Res } // isSSEResponse 仅对 200 + text/event-stream 的响应做清洗; -// 错误响应 (4xx/5xx 通常是 application/json) 不动, 由 SDK 走原错误路径。 +// 错误响应由独立的 einoSSEErrorRoundTripper 处理,此层不改动。 func isSSEResponse(resp *http.Response) bool { if resp.StatusCode != http.StatusOK { return false