superagent请求问题修复
This commit is contained in:
1 parent
ee9d2d7cf1
commit
7243319bbb
19 files changed
+388
-74
No files matched your search
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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" +
|
||||
|
||||
@@ -35,6 +35,7 @@ type Config struct {
|
||||
Enabled bool
|
||||
BaseURL string
|
||||
APIKey string
|
||||
IncludeTrace bool
|
||||
ConnectTimeout time.Duration
|
||||
RecoveryMaxAttempts int
|
||||
RecoveryInitialBackoff time.Duration
|
||||
|
||||
Reference in new issue
Block a user