mirror of
https://github.com/Ed1s0nZ/CyberStrikeAI.git
synced 2026-09-15 22:25:33 +02:00
fix: normalize SSE error responses for Eino models
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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.
|
||||
}
|
||||
@@ -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", "<html>bad gateway</html>", "<html>bad gateway</html>", 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")
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user