Files
fire-safety-ymd/internal/integration/superagent/sse.go
2026-09-06 01:40:15 +08:00

544 lines
14 KiB
Go
Raw Blame History

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
}