iop/apps/edge/internal/openai/stream.go

425 lines
13 KiB
Go

package openai
import (
"encoding/json"
"fmt"
"net/http"
"strings"
"time"
"unicode"
"go.uber.org/zap"
edgeservice "iop/apps/edge/internal/service"
)
func (s *Server) streamChatCompletion(w http.ResponseWriter, r *http.Request, req chatCompletionRequest, submitReq edgeservice.SubmitRunRequest, handle edgeservice.RunResult, outputPolicy strictOutputPolicy, validation toolValidationContract) {
flusher, ok := w.(http.Flusher)
if !ok {
handle.Close()
writeError(w, http.StatusInternalServerError, "streaming_not_supported", "response writer does not support streaming")
return
}
w.Header().Set("Content-Type", "text/event-stream")
w.Header().Set("Cache-Control", "no-cache")
w.Header().Set("Connection", "keep-alive")
// Strict buffered streams collect and validate the full response before
// emitting any user-visible chunk, so they own handle lifecycle, retries,
// and the initial role chunk internally.
if outputPolicy.Strict && outputPolicy.StreamBuffer {
s.streamBufferedChatCompletion(w, r, req, submitReq, handle, flusher, outputPolicy, validation)
return
}
// Live SSE may emit content deltas before the terminal event, so runtime
// tool validation is excluded upstream; write the role chunk immediately.
defer handle.Close()
created := time.Now().Unix()
id := "chatcmpl-" + handle.Dispatch().RunID
model := responseModel(req.Model, handle.Dispatch().Target)
writeSSE(w, flusher, chatCompletionChunk{
ID: id,
Object: "chat.completion.chunk",
Created: created,
Model: model,
Choices: []chatCompletionChunkChoice{{
Index: 0,
Delta: chatDelta{Role: "assistant"},
}},
})
var contentBuilder strings.Builder
var reasoningBuilder strings.Builder
var emittedContent strings.Builder
var toolTextFilter *streamToolTextFilter
if len(req.Tools) > 0 {
toolTextFilter = &streamToolTextFilter{}
}
defer func() {
s.logger.Info("openai chat completion stream closed",
zap.String("run_id", handle.Dispatch().RunID),
zap.Bool("strict_output", outputPolicy.Strict),
zap.Bool("strict_stream_buffer", outputPolicy.StreamBuffer),
zap.Int("content_len", contentBuilder.Len()),
zap.Int("reasoning_len", reasoningBuilder.Len()),
zap.String("content_preview", previewString(contentBuilder.String(), 1000)),
zap.String("reasoning_preview", previewString(reasoningBuilder.String(), 1000)),
)
}()
stream := handle.Stream()
if stream.Events == nil {
writeSSEError(w, flusher, "run stream unavailable")
return
}
for {
select {
case <-r.Context().Done():
s.cancelRunOnHTTPGiveUp(handle.Dispatch(), r.Context().Err())
return
case nodeEvent, ok := <-stream.NodeEvents:
if !ok {
stream.NodeEvents = nil
continue
}
if edgeservice.IsNodeDisconnected(nodeEvent) {
writeSSEError(w, flusher, "node disconnected")
return
}
case event, ok := <-stream.Events:
if !ok {
writeSSEError(w, flusher, "run stream closed")
return
}
if event == nil {
continue
}
switch event.GetType() {
case "delta":
if event.GetDelta() == "" {
continue
}
delta := event.GetDelta()
contentBuilder.WriteString(delta)
if toolTextFilter != nil {
delta = toolTextFilter.Append(delta)
}
if delta == "" {
continue
}
emittedContent.WriteString(delta)
writeContentDeltaSSE(w, flusher, id, created, model, delta)
case "reasoning_delta":
if event.GetDelta() == "" {
continue
}
reasoningBuilder.WriteString(event.GetDelta())
if outputPolicy.Strict {
continue
}
writeSSE(w, flusher, chatCompletionChunk{
ID: id,
Object: "chat.completion.chunk",
Created: created,
Model: model,
Choices: []chatCompletionChunkChoice{{
Index: 0,
Delta: chatDelta{ReasoningContent: event.GetDelta()},
}},
})
case "complete":
metadata := event.GetMetadata()
finishReason := metadata["finish_reason"]
if finishReason == "" {
finishReason = "stop"
}
nativeToolCalls := toolCallsFromRunMetadata(metadata)
if len(nativeToolCalls) > 0 {
if toolTextFilter != nil {
pending := toolTextFilter.Flush()
if pending != "" {
emittedContent.WriteString(pending)
writeContentDeltaSSE(w, flusher, id, created, model, pending)
}
}
nativeToolCalls = normalizeNativeToolCallsForResponse(nativeToolCalls, req.Tools)
writeToolCallsDeltaSSE(w, flusher, id, created, model, nativeToolCalls)
finishReason = "tool_calls"
} else if toolTextFilter != nil && shouldSynthesizeTextToolCalls(handle.Dispatch(), isTextToolFallback(metadata)) {
cleaned, toolCalls := synthesizeToolCallsFromText(contentBuilder.String(), req.Tools, handle.Dispatch().RunID)
if len(toolCalls) > 0 {
if remaining := unstreamedCleanedText(cleaned, emittedContent.String()); remaining != "" {
emittedContent.WriteString(remaining)
writeContentDeltaSSE(w, flusher, id, created, model, remaining)
}
writeToolCallsDeltaSSE(w, flusher, id, created, model, toolCalls)
finishReason = "tool_calls"
} else if pending := toolTextFilter.Flush(); pending != "" {
emittedContent.WriteString(pending)
writeContentDeltaSSE(w, flusher, id, created, model, pending)
}
} else if toolTextFilter != nil {
if pending := toolTextFilter.Flush(); pending != "" {
emittedContent.WriteString(pending)
writeContentDeltaSSE(w, flusher, id, created, model, pending)
}
}
writeSSE(w, flusher, chatCompletionChunk{
ID: id,
Object: "chat.completion.chunk",
Created: created,
Model: model,
Choices: []chatCompletionChunkChoice{{
Index: 0,
Delta: chatDelta{},
FinishReason: finishReason,
}},
})
fmt.Fprint(w, "data: [DONE]\n\n")
flusher.Flush()
return
case "error", "cancelled":
msg := event.GetError()
if msg == "" {
msg = event.GetMessage()
}
if msg == "" {
msg = "run failed"
}
writeSSEError(w, flusher, msg)
return
}
case <-time.After(handle.WaitTimeout()):
s.cancelRunOnHTTPGiveUp(handle.Dispatch(), errRunTimedOut)
writeSSEError(w, flusher, errRunTimedOut.Error())
return
}
}
}
type streamToolTextFilter struct {
pending string
}
func (f *streamToolTextFilter) Append(delta string) string {
f.pending += delta
if idx := firstStreamToolTextCandidateIndex(f.pending); idx >= 0 {
out := strings.TrimRightFunc(f.pending[:idx], unicode.IsSpace)
f.pending = f.pending[idx:]
return out
}
flushLen := streamToolTextSafeFlushLen(f.pending)
out := strings.TrimRightFunc(f.pending[:flushLen], unicode.IsSpace)
f.pending = f.pending[len(out):]
return out
}
func (f *streamToolTextFilter) Flush() string {
out := f.pending
f.pending = ""
return out
}
func firstStreamToolTextCandidateIndex(s string) int {
xml := strings.Index(strings.ToLower(s), "<tool_call")
mustache := strings.Index(s, textMustacheToolCallOpenHint)
switch {
case xml < 0:
return mustache
case mustache < 0:
return xml
case xml < mustache:
return xml
default:
return mustache
}
}
func streamToolTextSafeFlushLen(s string) int {
if s == "" {
return 0
}
const xmlHint = "<tool_call"
lower := strings.ToLower(s)
start := len(s) - len(xmlHint) + 1
if start < 0 {
start = 0
}
for i := start; i < len(s); i++ {
if strings.HasPrefix(xmlHint, lower[i:]) ||
strings.HasPrefix(textMustacheToolCallOpenHint, s[i:]) {
return i
}
}
return len(s)
}
func unstreamedCleanedText(cleaned, emitted string) string {
if cleaned == "" {
return ""
}
if emitted == "" {
return cleaned
}
if strings.HasPrefix(cleaned, emitted) {
return cleaned[len(emitted):]
}
cleaned = strings.TrimSpace(cleaned)
emitted = strings.TrimSpace(emitted)
if cleaned == "" || cleaned == emitted {
return ""
}
if emitted != "" && strings.HasPrefix(cleaned, emitted) {
return strings.TrimSpace(strings.TrimPrefix(cleaned, emitted))
}
return ""
}
func writeContentDeltaSSE(w http.ResponseWriter, flusher http.Flusher, id string, created int64, model string, content string) {
if content == "" {
return
}
writeSSE(w, flusher, chatCompletionChunk{
ID: id,
Object: "chat.completion.chunk",
Created: created,
Model: model,
Choices: []chatCompletionChunkChoice{{
Index: 0,
Delta: chatDelta{Content: content},
}},
})
}
func writeToolCallsDeltaSSE(w http.ResponseWriter, flusher http.Flusher, id string, created int64, model string, toolCalls []any) {
if len(toolCalls) == 0 {
return
}
writeSSE(w, flusher, chatCompletionChunk{
ID: id,
Object: "chat.completion.chunk",
Created: created,
Model: model,
Choices: []chatCompletionChunkChoice{{
Index: 0,
Delta: chatDelta{ToolCalls: toolCallsForStreamDelta(toolCalls)},
}},
})
}
// streamBufferedChatCompletion collects the full run output, validates tool
// calls against the request schema before emitting anything, and — when
// validation fails on a run that has not yet written a user-visible chunk —
// resubmits the same request as a bounded exact replay. Only a validated (or
// validation-disabled) result reaches the SSE writer; a malformed tool call
// that survives the retry budget is surfaced as an SSE error instead of a
// successful tool_calls chunk.
func (s *Server) streamBufferedChatCompletion(w http.ResponseWriter, r *http.Request, req chatCompletionRequest, submitReq edgeservice.SubmitRunRequest, handle edgeservice.RunResult, flusher http.Flusher, outputPolicy strictOutputPolicy, validation toolValidationContract) {
attempt := 1
for {
result, err := collectChatCompletionOutput(r.Context(), req, handle, outputPolicy)
if err != nil {
s.cancelRunOnHTTPGiveUp(handle.Dispatch(), err)
handle.Close()
writeSSEError(w, flusher, err.Error())
return
}
if verr := validateToolCallResponse(validation, result.toolCalls, result.toolCallOrigin); verr != nil {
failedRunID := handle.Dispatch().RunID
if attempt < maxToolValidationAttempts {
handle.Close()
attempt++
retryReq := submitReq
retryReq.Metadata = toolValidationAttemptMetadata(submitReq.Metadata, attempt, failedRunID, verr.Error())
s.logger.Warn("openai chat completion stream tool validation retry",
zap.String("run_id", failedRunID),
zap.Int("attempt", attempt),
zap.String("reason", verr.Error()),
)
next, submitErr := s.service.SubmitRun(r.Context(), retryReq)
if submitErr != nil {
writeSSEErrorWithType(w, flusher, "tool_validation_retry_error", submitErr.Error())
return
}
handle = next
continue
}
handle.Close()
s.logger.Warn("openai chat completion stream tool validation failed",
zap.String("run_id", failedRunID),
zap.Int("attempt", attempt),
zap.String("reason", verr.Error()),
)
writeSSEErrorWithType(w, flusher, "tool_validation_error", verr.Error())
return
}
handle.Close()
s.writeBufferedStreamOutput(w, flusher, req, handle.Dispatch(), result, outputPolicy)
return
}
}
// writeBufferedStreamOutput emits a validated buffered result as SSE chunks:
// the role chunk, an optional content chunk, an optional tool_calls chunk, and
// the terminal finish chunk.
func (s *Server) writeBufferedStreamOutput(w http.ResponseWriter, flusher http.Flusher, req chatCompletionRequest, dispatch edgeservice.RunDispatch, result chatCompletionOutput, outputPolicy strictOutputPolicy) {
created := time.Now().Unix()
id := "chatcmpl-" + dispatch.RunID
model := responseModel(req.Model, dispatch.Target)
content := result.message.Content
s.logger.Info("openai chat completion stream closed",
zap.String("run_id", dispatch.RunID),
zap.Bool("strict_output", outputPolicy.Strict),
zap.Bool("strict_stream_buffer", outputPolicy.StreamBuffer),
zap.String("xml_completion_tool", outputPolicy.XMLCompletionTool),
zap.Bool("normalized", result.normalized),
zap.Int("content_len", len(content)),
zap.Int("reasoning_len", result.reasoningLen),
zap.String("content_preview", previewString(content, 1000)),
)
writeSSE(w, flusher, chatCompletionChunk{
ID: id,
Object: "chat.completion.chunk",
Created: created,
Model: model,
Choices: []chatCompletionChunkChoice{{
Index: 0,
Delta: chatDelta{Role: "assistant"},
}},
})
writeContentDeltaSSE(w, flusher, id, created, model, content)
writeToolCallsDeltaSSE(w, flusher, id, created, model, result.toolCalls)
writeSSE(w, flusher, chatCompletionChunk{
ID: id,
Object: "chat.completion.chunk",
Created: created,
Model: model,
Choices: []chatCompletionChunkChoice{{
Index: 0,
Delta: chatDelta{},
FinishReason: result.finishReason,
}},
})
fmt.Fprint(w, "data: [DONE]\n\n")
flusher.Flush()
}
func writeSSE(w http.ResponseWriter, flusher http.Flusher, v any) {
b, err := json.Marshal(v)
if err != nil {
return
}
fmt.Fprintf(w, "data: %s\n\n", b)
flusher.Flush()
}
func writeSSEError(w http.ResponseWriter, flusher http.Flusher, message string) {
writeSSEErrorWithType(w, flusher, "run_error", message)
}
func writeSSEErrorWithType(w http.ResponseWriter, flusher http.Flusher, errType, message string) {
writeSSE(w, flusher, errorResponse{Error: errorBody{Type: errType, Message: message}})
fmt.Fprint(w, "data: [DONE]\n\n")
flusher.Flush()
}