544 lines
14 KiB
Go
544 lines
14 KiB
Go
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, "<22>"))
|
||
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
|
||
}
|