superagent请求问题修复

This commit is contained in:
andy committed 2026-09-06 01:40:15 +08:00
1 parent ee9d2d7cf1
commit 7243319bbb
19 files changed
+388 -74

No files matched your search

+10 -5
View File
@@ -13,6 +13,7 @@ import (
"net"
"net/http"
"net/url"
"strconv"
"strings"
"time"
)
@@ -34,6 +35,7 @@ type HTTPClient struct {
enabled bool
baseURL *url.URL
apiKey string
includeTrace bool
httpClient *http.Client
recoveryMaxAttempts int
recoveryInitialBackoff time.Duration
@@ -88,9 +90,10 @@ func NewHTTPClient(cfg Config) (*HTTPClient, error) {
}).DialContext
return &HTTPClient{
enabled: true,
baseURL: baseURL,
apiKey: cfg.APIKey,
enabled: true,
baseURL: baseURL,
apiKey: cfg.APIKey,
includeTrace: cfg.IncludeTrace,
httpClient: &http.Client{
Transport: transport,
CheckRedirect: func(_ *http.Request, _ []*http.Request) error {
@@ -198,7 +201,9 @@ func (c *HTTPClient) StreamMessage(
}
streamURL := c.endpoint("/api/open/agent-sessions/" + request.SessionID + "/messages/stream")
streamURL.RawQuery = "include_trace=true"
streamQuery := streamURL.Query()
streamQuery.Set("include_trace", strconv.FormatBool(c.includeTrace))
streamURL.RawQuery = streamQuery.Encode()
httpRequest, err := c.newRequest(
ctx,
http.MethodPost,
@@ -229,7 +234,7 @@ func (c *HTTPClient) StreamMessage(
return Result{}, err
}
state := newStreamState(request.SessionID)
state := newStreamState(request.SessionID, c.includeTrace)
if contentLocation := strings.TrimSpace(response.Header.Get("Content-Location")); contentLocation != "" {
runURL, locationErr := c.resolveRunURL(httpRequest.URL, request.SessionID, contentLocation)
if locationErr != nil {
@@ -105,6 +105,39 @@ func TestHTTPClientCreatesSessionAndStreamsMessage(t *testing.T) {
}
}
func TestHTTPClientStreamsWithoutTraceWhenDisabled(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
if got := request.URL.Query().Get("include_trace"); got != "false" {
t.Errorf("include_trace = %q, want false", got)
}
writer.Header().Set("Content-Type", "text/event-stream")
fmt.Fprint(writer, "event: messages\ndata: [{\"type\":\"AIMessageChunk\",\"content\":\"\",\"response_metadata\":{\"finish_reason\":\"stop\"}}]\n\n")
fmt.Fprint(writer, "event: message.final\ndata: {\"run_id\":\"run-no-trace\",\"text\":\"no-trace answer\"}\n\n")
fmt.Fprint(writer, "event: end\n\n")
}))
defer server.Close()
client, err := NewHTTPClient(Config{
Enabled: true,
BaseURL: server.URL,
APIKey: testAPIKey,
IncludeTrace: false,
ConnectTimeout: time.Second,
RecoveryInitialBackoff: time.Millisecond,
MaxMessageBytes: 1024,
})
if err != nil {
t.Fatalf("NewHTTPClient() error = %v", err)
}
result, err := client.StreamMessage(context.Background(), validMessageRequest(), nil)
if err != nil {
t.Fatalf("StreamMessage() error = %v", err)
}
if result.Answer != "no-trace answer" {
t.Fatalf("Answer = %q, want no-trace answer", result.Answer)
}
}
func TestHTTPClientRecoversWithoutRepostingMessage(t *testing.T) {
var messagePosts atomic.Int32
var runQueries atomic.Int32
@@ -405,6 +438,7 @@ func newTestHTTPClient(t *testing.T, baseURL string, recoveryAttempts int, backo
Enabled: true,
BaseURL: baseURL,
APIKey: testAPIKey,
IncludeTrace: true,
ConnectTimeout: time.Second,
RecoveryMaxAttempts: recoveryAttempts,
RecoveryInitialBackoff: backoff,
+54 -23
View File
@@ -28,13 +28,16 @@ var (
)
type streamState struct {
sessionID string
sessionID string
requireRunCompleted bool
endSeen bool
runCompleted bool
runID string
runURL string
lastEventID string
endSeen bool
runCompleted bool
finalAIStopSeen bool
topLevelFinalSeen bool
runID string
runURL string
lastEventID string
profileID string
profileVersionID string
@@ -56,11 +59,12 @@ type streamState struct {
streamBytes int64
}
func newStreamState(sessionID string) *streamState {
func newStreamState(sessionID string, requireRunCompleted bool) *streamState {
return &streamState{
sessionID: sessionID,
processedIDs: make(map[string]struct{}),
eventTypeSet: make(map[string]struct{}),
sessionID: sessionID,
requireRunCompleted: requireRunCompleted,
processedIDs: make(map[string]struct{}),
eventTypeSet: make(map[string]struct{}),
}
}
@@ -191,6 +195,21 @@ func (s *streamState) consumeFrame(frame sseFrame, traceHandler TraceHandler) er
return err
}
return s.consumeMessages(payload)
case "message.final":
payload, err := decodeJSONObject(data)
if err != nil {
return err
}
text := stringValue(payload, "text")
if strings.TrimSpace(text) == "" {
return nil
}
if len(text) > maxAnswerBytes {
return fmt.Errorf("%w: final answer exceeds size limit", ErrProtocol)
}
s.runID = valueOrExisting(safeValue(stringValue(payload, "run_id")), s.runID)
s.finalContent = text
s.topLevelFinalSeen = true
case "values":
payload, err := decodeJSONObject(data)
if err != nil {
@@ -277,16 +296,10 @@ func (s *streamState) consumeMessages(value any) error {
if nested, exists := current["messages"]; exists {
return s.consumeMessages(nested)
}
if stringValue(current, "type") != "ai" {
messageType := stringValue(current, "type")
if messageType != "ai" && messageType != "AIMessageChunk" {
return nil
}
content := contentText(current["content"])
if strings.TrimSpace(content) == "" {
return nil
}
if len(content) > maxAnswerBytes {
return fmt.Errorf("%w: final answer exceeds size limit", ErrProtocol)
}
responseMetadata, _ := current["response_metadata"].(map[string]any)
usageMetadata, _ := current["usage_metadata"].(map[string]any)
modelName := valueOrExisting(safeTrace(stringValue(responseMetadata, "model_name")), s.modelName)
@@ -295,10 +308,22 @@ func (s *streamState) consumeMessages(value any) error {
Output: int64Value(usageMetadata, "output_tokens"),
Total: int64Value(usageMetadata, "total_tokens"),
}
if stringValue(responseMetadata, "finish_reason") == "stop" {
s.finalContent = content
stopSeen := stringValue(responseMetadata, "finish_reason") == "stop"
if stopSeen {
s.finalAIStopSeen = true
s.modelName = modelName
s.usage = usage
}
content := contentText(current["content"])
if strings.TrimSpace(content) == "" {
return nil
}
if len(content) > maxAnswerBytes {
return fmt.Errorf("%w: final answer exceeds size limit", ErrProtocol)
}
if stopSeen {
s.finalContent = content
return nil
}
s.fallbackContent = content
@@ -315,14 +340,20 @@ func (s *streamState) result() (Result, error) {
if !s.endSeen {
return Result{}, fmt.Errorf("%w: missing end event", ErrProtocol)
}
if !s.runCompleted {
if s.requireRunCompleted && !s.runCompleted {
return Result{}, fmt.Errorf("%w: missing successful run.completed event", ErrProtocol)
}
if !s.requireRunCompleted && !s.finalAIStopSeen {
return Result{}, fmt.Errorf("%w: missing AI stop finish reason", ErrProtocol)
}
if !s.requireRunCompleted && !s.topLevelFinalSeen {
return Result{}, fmt.Errorf("%w: missing message.final event", ErrProtocol)
}
answer := s.finalContent
if strings.TrimSpace(answer) == "" {
if s.requireRunCompleted && strings.TrimSpace(answer) == "" {
answer = s.deltaContent.String()
}
if strings.TrimSpace(answer) == "" {
if s.requireRunCompleted && strings.TrimSpace(answer) == "" {
answer = s.fallbackContent
s.modelName = s.fallbackModel
s.usage = s.fallbackUsage
+98 -5
View File
@@ -7,7 +7,7 @@ import (
)
func TestStreamStateConsumesCurrentTraceProtocol(t *testing.T) {
state := newStreamState("session-1")
state := newStreamState("session-1", true)
var traces []TraceEvent
stream := strings.Join([]string{
"event: metadata\nid: 1\ndata: {\"run_id\":\"run-1\",\"resolved_profile_id\":\"profile-1\",\"resolved_profile_version_id\":\"version-1\"}\n",
@@ -39,7 +39,7 @@ func TestStreamStateConsumesCurrentTraceProtocol(t *testing.T) {
}
func TestStreamStateConsumesLegacyValuesAndStructuredContent(t *testing.T) {
state := newStreamState("session-legacy")
state := newStreamState("session-legacy", true)
stream := "event: values\n" +
"data: {\"messages\":[{\"type\":\"human\",\"content\":\"ignored\"},{\"type\":\"ai\",\"content\":[{\"text\":\"first\"},{\"content\":\"second\"}],\"response_metadata\":{\"finish_reason\":\"stop\",\"model_name\":\"model-1\"},\"usage_metadata\":{\"input_tokens\":2,\"output_tokens\":3,\"total_tokens\":5}}]}\n\n" +
"event: trace\ndata: {\"event\":\"run.completed\",\"status\":\"success\"}\n\n" +
@@ -60,8 +60,101 @@ func TestStreamStateConsumesLegacyValuesAndStructuredContent(t *testing.T) {
}
}
func TestStreamStateConsumesNoTraceFinalMessage(t *testing.T) {
state := newStreamState("session-no-trace", false)
stream := "event: messages\n" +
"data: [{\"type\":\"AIMessageChunk\",\"content\":\"\",\"response_metadata\":{\"finish_reason\":\"stop\",\"model_name\":\"model-1\"},\"usage_metadata\":{\"input_tokens\":2,\"output_tokens\":3,\"total_tokens\":5}}]\n\n" +
"event: message.final\n" +
"data: {\"run_id\":\"run-no-trace\",\"text\":\"No trace answer\"}\n\n" +
"event: end\n\n"
if err := state.consume(strings.NewReader(stream), nil); err != nil {
t.Fatalf("consume() error = %v", err)
}
result, err := state.result()
if err != nil {
t.Fatalf("result() error = %v", err)
}
if result.Answer != "No trace answer" || result.ModelName != "model-1" || result.RunID != "run-no-trace" {
t.Fatalf("unexpected result: %#v", result)
}
if result.Usage != (TokenUsage{Input: 2, Output: 3, Total: 5}) {
t.Fatalf("Usage = %#v", result.Usage)
}
}
func TestStreamStateNoTraceRequiresStopAndEnd(t *testing.T) {
tests := []struct {
name string
stream string
wantErr error
result bool
}{
{
name: "missing stop",
stream: "event: messages\n" +
"data: [{\"type\":\"AIMessageChunk\",\"content\":\"partial\",\"response_metadata\":{\"finish_reason\":null}}]\n\n" +
"event: message.final\ndata: {\"run_id\":\"run-1\",\"text\":\"complete\"}\n\n" +
"event: end\n\n",
wantErr: ErrProtocol,
result: true,
},
{
name: "missing message final",
stream: "event: messages\n" +
"data: [{\"type\":\"AIMessageChunk\",\"content\":\"\",\"response_metadata\":{\"finish_reason\":\"stop\"}}]\n\n" +
"event: end\n\n",
wantErr: ErrProtocol,
result: true,
},
{
name: "missing end",
stream: "event: messages\n" +
"data: [{\"type\":\"AIMessageChunk\",\"content\":\"\",\"response_metadata\":{\"finish_reason\":\"stop\"}}]\n\n" +
"event: message.final\ndata: {\"run_id\":\"run-1\",\"text\":\"complete\"}\n\n",
wantErr: ErrStreamIncomplete,
},
{
name: "top-level error",
stream: "event: error\ndata: {\"code\":\"provider_error\",\"message\":\"do not expose me\"}\n\n",
wantErr: ErrRunFailed,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
state := newStreamState("session-no-trace", false)
err := state.consume(strings.NewReader(tt.stream), nil)
if tt.result && err == nil {
_, err = state.result()
}
if !errors.Is(err, tt.wantErr) {
t.Fatalf("error = %v, want errors.Is(..., %v)", err, tt.wantErr)
}
if err != nil && strings.Contains(strings.ToLower(err.Error()), "do not expose") {
t.Fatalf("error leaked provider message: %v", err)
}
})
}
}
func TestStreamStateTraceStillRequiresRunCompleted(t *testing.T) {
state := newStreamState("session-trace", true)
stream := "event: values\n" +
"data: {\"messages\":[{\"type\":\"ai\",\"content\":\"complete\",\"response_metadata\":{\"finish_reason\":\"stop\"}}]}\n\n" +
"event: end\n\n"
if err := state.consume(strings.NewReader(stream), nil); err != nil {
t.Fatalf("consume() error = %v", err)
}
_, err := state.result()
if !errors.Is(err, ErrProtocol) || !strings.Contains(err.Error(), "run.completed") {
t.Fatalf("result() error = %v, want missing run.completed protocol error", err)
}
}
func TestStreamStateSupportsMultilineDataAndDeduplicatesIDs(t *testing.T) {
state := newStreamState("session-1")
state := newStreamState("session-1", true)
stream := ": heartbeat\n\n" +
"event: trace\nid: delta-1\ndata: {\"event\":\"message.delta\",\ndata: \"text\":\"A\"}\n\n" +
"event: trace\nid: delta-1\ndata: {\"event\":\"message.delta\",\"text\":\"duplicate\"}\n\n" +
@@ -128,7 +221,7 @@ func TestStreamStateRejectsIncompleteOrFailedStreams(t *testing.T) {
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
state := newStreamState("session-1")
state := newStreamState("session-1", true)
err := state.consume(strings.NewReader(tt.stream), nil)
if tt.result && err == nil {
_, err = state.result()
@@ -144,7 +237,7 @@ func TestStreamStateRejectsIncompleteOrFailedStreams(t *testing.T) {
}
func TestTraceProjectionRedactsAndBoundsText(t *testing.T) {
state := newStreamState("session-1")
state := newStreamState("session-1", true)
secretText := "token=highly-secret " + strings.Repeat("x", maxTraceText+100)
stream := "event: trace\ndata: {\"event\":\"message.final\",\"text\":" + quotedJSON(secretText) + "}\n\n" +
"event: trace\ndata: {\"event\":\"run.completed\",\"status\":\"success\"}\n\n" +
+1
View File
@@ -35,6 +35,7 @@ type Config struct {
Enabled bool
BaseURL string
APIKey string
IncludeTrace bool
ConnectTimeout time.Duration
RecoveryMaxAttempts int
RecoveryInitialBackoff time.Duration