package superagent import ( "bufio" "encoding/json" "fmt" "io" "regexp" "strconv" "strings" ) const ( maxSSELineBytes = 1024 * 1024 maxSSEEventBytes = 1024 * 1024 maxSSEStreamBytes = 32 * 1024 * 1024 maxAnswerBytes = 4 * 1024 * 1024 maxSSEEvents = 10_000 maxTraceEvents = 1_000 maxTraceText = 2_048 ) var ( safeProviderValue = regexp.MustCompile(`^[A-Za-z0-9._-]{1,128}$`) traceJSONSecret = regexp.MustCompile(`(?i)("(?:api[_-]?key|token|secret|password|authorization|cookie|csrf[_-]?token)"\s*:\s*")[^"]*(")`) traceHeaderSecret = regexp.MustCompile(`(?i)((?:authorization|cookie|x-csrf-token)\s*:\s*)[^\s,;]+`) traceSecretValue = regexp.MustCompile(`(?i)((?:api[_-]?key|token|secret|password|authorization|cookie|csrf[_-]?token)\s*[=:]\s*)[^\s,;}]+`) ) type streamState struct { sessionID string requireRunCompleted bool endSeen bool runCompleted bool finalAIStopSeen bool topLevelFinalSeen bool runID string runURL string lastEventID string profileID string profileVersionID string modelName string finalContent string deltaContent strings.Builder fallbackContent string usage TokenUsage fallbackUsage TokenUsage fallbackModel string failureCode string failureStatus string processedIDs map[string]struct{} eventTypeSet map[string]struct{} eventTypes []string traceEvents []TraceEvent eventCount int streamBytes int64 } func newStreamState(sessionID string, requireRunCompleted bool) *streamState { return &streamState{ sessionID: sessionID, requireRunCompleted: requireRunCompleted, processedIDs: make(map[string]struct{}), eventTypeSet: make(map[string]struct{}), } } type sseFrame struct { eventType string id string data []string dataBytes int } func (s *streamState) consume(reader io.Reader, traceHandler TraceHandler) error { if reader == nil { return fmt.Errorf("%w: empty SSE body", ErrProtocol) } scanner := bufio.NewScanner(reader) scanner.Buffer(make([]byte, 64*1024), maxSSELineBytes) frame := sseFrame{eventType: "message"} flush := func() error { if len(frame.data) == 0 && frame.eventType != "end" { frame = sseFrame{eventType: "message"} return nil } err := s.consumeFrame(frame, traceHandler) frame = sseFrame{eventType: "message"} return err } for scanner.Scan() { line := scanner.Text() s.streamBytes += int64(len(line) + 1) if s.streamBytes > maxSSEStreamBytes { return fmt.Errorf("%w: SSE stream exceeds size limit", ErrProtocol) } if line == "" { if err := flush(); err != nil { return err } if s.endSeen { return nil } continue } if strings.HasPrefix(line, ":") { continue } field, value, found := strings.Cut(line, ":") if !found { field, value = line, "" } else if strings.HasPrefix(value, " ") { value = value[1:] } switch field { case "event": if strings.TrimSpace(value) == "" { frame.eventType = "message" } else { frame.eventType = strings.TrimSpace(value) } case "id": if !strings.ContainsRune(value, '\x00') { frame.id = strings.TrimSpace(value) } case "data": frame.dataBytes += len(value) + 1 if frame.dataBytes > maxSSEEventBytes { return fmt.Errorf("%w: SSE event exceeds size limit", ErrProtocol) } frame.data = append(frame.data, value) } } if err := scanner.Err(); err != nil { return ErrStreamRead } if err := flush(); err != nil { return err } if !s.endSeen { return ErrStreamIncomplete } return nil } func (s *streamState) consumeFrame(frame sseFrame, traceHandler TraceHandler) error { s.eventCount++ if s.eventCount > maxSSEEvents { return fmt.Errorf("%w: too many SSE events", ErrProtocol) } if !validResourceID(frame.eventType) { return fmt.Errorf("%w: invalid SSE event type", ErrProtocol) } if frame.id != "" && !validOptionalHeader(frame.id) { return fmt.Errorf("%w: invalid SSE event ID", ErrProtocol) } if frame.id != "" { if _, duplicate := s.processedIDs[frame.id]; duplicate { return nil } s.processedIDs[frame.id] = struct{}{} s.lastEventID = frame.id } s.addEventType(frame.eventType) data := strings.Join(frame.data, "\n") switch frame.eventType { case "end": s.endSeen = true return nil case "error": payload, _ := decodeJSONObject(data) s.failureCode = safeValue(stringValue(payload, "code")) if s.failureCode == "" { s.failureCode = "stream_error" } return &RunError{Code: s.failureCode} case "metadata": payload, err := decodeJSONObject(data) if err != nil { return err } s.runID = valueOrExisting(safeValue(stringValue(payload, "run_id")), s.runID) s.profileID = valueOrExisting(safeValue(stringValue(payload, "resolved_profile_id")), s.profileID) s.profileVersionID = valueOrExisting(safeValue(stringValue(payload, "resolved_profile_version_id")), s.profileVersionID) case "messages": payload, err := decodeJSON(data) if err != nil { 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 { return err } return s.consumeMessages(payload["messages"]) case "trace": payload, err := decodeJSONObject(data) if err != nil { return err } return s.consumeTrace(payload, traceHandler) } return nil } func (s *streamState) consumeTrace(payload map[string]any, traceHandler TraceHandler) error { trace := payload if _, hasEvent := payload["event"]; !hasEvent { if nested, ok := payload["data"].(map[string]any); ok { trace = nested } } s.runID = valueOrExisting(safeValue(stringValue(trace, "run_id")), s.runID) eventName := stringValue(trace, "event") var protocolErr error switch eventName { case "message.delta": text := stringValue(trace, "text") if s.deltaContent.Len()+len(text) > maxAnswerBytes { return fmt.Errorf("%w: streamed answer exceeds size limit", ErrProtocol) } s.deltaContent.WriteString(text) case "message.final": if text := stringValue(trace, "text"); strings.TrimSpace(text) != "" { if len(text) > maxAnswerBytes { return fmt.Errorf("%w: final answer exceeds size limit", ErrProtocol) } s.finalContent = text } case "run.completed": status := safeValue(stringValue(trace, "status")) if status != "success" { s.failureStatus = valueOrExisting(status, "non_success") protocolErr = &RunError{Status: s.failureStatus} } else { s.runCompleted = true } case "run.failed": if providerError, ok := trace["error"].(map[string]any); ok { s.failureCode = safeValue(stringValue(providerError, "code")) } if s.failureCode == "" { s.failureCode = "run_failed" } protocolErr = &RunError{Code: s.failureCode} } toolCalls, _ := trace["tool_calls"].([]any) if len(toolCalls) == 0 { s.emitTrace(traceEvent(trace, trace), traceHandler) return protocolErr } for _, rawCall := range toolCalls { detail, ok := rawCall.(map[string]any) if !ok { continue } s.emitTrace(traceEvent(trace, detail), traceHandler) } return protocolErr } func (s *streamState) consumeMessages(value any) error { switch current := value.(type) { case []any: for _, item := range current { if err := s.consumeMessages(item); err != nil { return err } } case map[string]any: if nested, exists := current["messages"]; exists { return s.consumeMessages(nested) } messageType := stringValue(current, "type") if messageType != "ai" && messageType != "AIMessageChunk" { return nil } responseMetadata, _ := current["response_metadata"].(map[string]any) usageMetadata, _ := current["usage_metadata"].(map[string]any) modelName := valueOrExisting(safeTrace(stringValue(responseMetadata, "model_name")), s.modelName) usage := TokenUsage{ Input: int64Value(usageMetadata, "input_tokens"), Output: int64Value(usageMetadata, "output_tokens"), Total: int64Value(usageMetadata, "total_tokens"), } 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 s.fallbackModel = modelName s.fallbackUsage = usage } return nil } func (s *streamState) result() (Result, error) { if s.failureCode != "" || s.failureStatus != "" { return Result{}, &RunError{Code: s.failureCode, Status: s.failureStatus} } if !s.endSeen { return Result{}, fmt.Errorf("%w: missing end event", ErrProtocol) } 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 s.requireRunCompleted && strings.TrimSpace(answer) == "" { answer = s.deltaContent.String() } if s.requireRunCompleted && strings.TrimSpace(answer) == "" { answer = s.fallbackContent s.modelName = s.fallbackModel s.usage = s.fallbackUsage } if strings.TrimSpace(answer) == "" { return Result{}, fmt.Errorf("%w: missing final AI answer", ErrProtocol) } return Result{ SessionID: s.sessionID, RunID: s.runID, ProfileID: s.profileID, ProfileVersionID: s.profileVersionID, ModelName: s.modelName, Answer: answer, Usage: s.usage, EventTypes: append([]string(nil), s.eventTypes...), TraceEvents: append([]TraceEvent(nil), s.traceEvents...), LastEventID: s.lastEventID, }, nil } func (s *streamState) addEventType(eventType string) { if _, exists := s.eventTypeSet[eventType]; exists { return } s.eventTypeSet[eventType] = struct{}{} s.eventTypes = append(s.eventTypes, eventType) } func (s *streamState) emitTrace(event TraceEvent, traceHandler TraceHandler) { if event.Event == "" { return } if len(s.traceEvents) < maxTraceEvents { s.traceEvents = append(s.traceEvents, event) } if traceHandler != nil { traceHandler(event) } } func traceEvent(trace, detail map[string]any) TraceEvent { toolCallID := valueOrExisting(stringValue(detail, "tool_call_id"), stringValue(trace, "tool_call_id")) toolName := valueOrExisting(stringValue(detail, "name"), stringValue(trace, "name")) return TraceEvent{ Event: safeTrace(stringValue(trace, "event")), RunID: safeTrace(stringValue(trace, "run_id")), MessageID: safeTrace(stringValue(trace, "message_id")), ToolCallID: safeTrace(toolCallID), ToolName: safeTrace(toolName), Text: safeTrace(stringValue(trace, "text")), Status: safeTrace(stringValue(trace, "status")), Timestamp: safeTrace(scalarString(trace["ts"])), } } func decodeJSON(data string) (any, error) { if strings.TrimSpace(data) == "" { return nil, fmt.Errorf("%w: empty JSON event", ErrProtocol) } decoder := json.NewDecoder(strings.NewReader(data)) decoder.UseNumber() var value any if err := decoder.Decode(&value); err != nil { return nil, fmt.Errorf("%w: invalid JSON event", ErrProtocol) } if decoder.Decode(&struct{}{}) != io.EOF { return nil, fmt.Errorf("%w: multiple JSON values in one event", ErrProtocol) } return value, nil } func decodeJSONObject(data string) (map[string]any, error) { value, err := decodeJSON(data) if err != nil { return nil, err } object, ok := value.(map[string]any) if !ok { return nil, fmt.Errorf("%w: JSON event must be an object", ErrProtocol) } return object, nil } func contentText(value any) string { switch current := value.(type) { case string: return current case []any: parts := make([]string, 0, len(current)) for _, item := range current { if text := contentText(item); strings.TrimSpace(text) != "" { parts = append(parts, text) } } return strings.Join(parts, "\n") case map[string]any: if text := contentText(current["text"]); strings.TrimSpace(text) != "" { return text } return contentText(current["content"]) default: return "" } } func stringValue(object map[string]any, key string) string { if object == nil { return "" } value, ok := object[key].(string) if !ok { return "" } return value } func int64Value(object map[string]any, key string) int64 { if object == nil { return 0 } switch value := object[key].(type) { case json.Number: parsed, _ := value.Int64() return nonNegative(parsed) case float64: return nonNegative(int64(value)) case int64: return nonNegative(value) case string: parsed, _ := strconv.ParseInt(value, 10, 64) return nonNegative(parsed) default: return 0 } } func nonNegative(value int64) int64 { if value < 0 { return 0 } return value } func scalarString(value any) string { switch current := value.(type) { case string: return current case json.Number: return current.String() case float64: return strconv.FormatFloat(current, 'f', -1, 64) default: return "" } } func valueOrExisting(value, existing string) string { if value == "" { return existing } return value } func safeValue(value string) string { value = strings.TrimSpace(value) if !safeProviderValue.MatchString(value) { return "" } return value } func safeTrace(value string) string { value = strings.TrimSpace(strings.ToValidUTF8(value, "�")) if value == "" { return "" } value = traceJSONSecret.ReplaceAllString(value, "${1}***${2}") value = traceHeaderSecret.ReplaceAllString(value, "${1}***") value = traceSecretValue.ReplaceAllString(value, "${1}***") runes := []rune(value) if len(runes) > maxTraceText { return string(runes[:maxTraceText]) } return value }