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

599 lines
19 KiB
Go

package openai
import (
"encoding/json"
"fmt"
"net/http"
"os"
"strings"
"time"
"unicode"
"go.uber.org/zap"
edgeservice "iop/apps/edge/internal/service"
)
const (
streamTraceMetadataKey = "iop_trace_stream"
streamTraceEnvKey = "IOP_OPENAI_COMPAT_TRACE_STREAM"
)
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)
traceStream := openAICompatTraceStreamEnabled(submitReq.Metadata)
traceSeq := 0
writeTracedContentDelta := func(content, source string, filtered bool) {
if content == "" {
return
}
traceSeq++
if traceStream {
logOpenAICompatStreamOutput(s.logger, handle.Dispatch().RunID, traceSeq, "content", content,
zap.Int("source_len", len(source)),
zap.Bool("filtered", filtered),
zap.Bool("source_has_open_think", hasOpenThinkTag(source)),
zap.Bool("source_has_close_think", hasCloseThinkTag(source)),
zap.String("source_preview", previewString(source, 2000)),
)
}
writeContentDeltaSSE(w, flusher, id, created, model, content)
}
writeTracedReasoningDelta := func(reasoning string) {
if reasoning == "" {
return
}
traceSeq++
if traceStream {
logOpenAICompatStreamOutput(s.logger, handle.Dispatch().RunID, traceSeq, "reasoning", reasoning)
}
writeSSE(w, flusher, chatCompletionChunk{
ID: id,
Object: "chat.completion.chunk",
Created: created,
Model: model,
Choices: []chatCompletionChunkChoice{{
Index: 0,
Delta: chatDelta{ReasoningContent: reasoning},
}},
})
}
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
contentSentinelFilter := &streamSentinelFilter{}
reasoningSentinelFilter := &streamSentinelFilter{}
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
}
processContentDelta := func(sourceDelta string) {
if sourceDelta == "" {
return
}
delta := contentSentinelFilter.Append(sourceDelta)
if delta == "" {
return
}
contentBuilder.WriteString(delta)
filteredDelta := delta
if toolTextFilter != nil {
filteredDelta = toolTextFilter.Append(filteredDelta)
}
if filteredDelta == "" {
if traceStream {
s.logger.Info("openai chat completion stream output suppressed",
zap.String("run_id", handle.Dispatch().RunID),
zap.String("type", "content"),
zap.Int("source_len", len(sourceDelta)),
zap.Bool("source_has_open_think", hasOpenThinkTag(sourceDelta)),
zap.Bool("source_has_close_think", hasCloseThinkTag(sourceDelta)),
zap.String("source_preview", previewString(sourceDelta, 2000)),
)
}
return
}
emittedContent.WriteString(filteredDelta)
writeTracedContentDelta(filteredDelta, sourceDelta, filteredDelta != sourceDelta)
}
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":
processContentDelta(event.GetDelta())
case "reasoning_delta":
if event.GetDelta() == "" {
continue
}
reasoning := reasoningSentinelFilter.Append(event.GetDelta())
if reasoning == "" {
continue
}
reasoningBuilder.WriteString(reasoning)
if outputPolicy.Strict || !req.includeReasoning() {
continue
}
writeTracedReasoningDelta(reasoning)
case "complete":
processContentDelta(contentSentinelFilter.Flush())
if reasoning := reasoningSentinelFilter.Flush(); reasoning != "" {
reasoningBuilder.WriteString(reasoning)
if !outputPolicy.Strict && req.includeReasoning() {
writeTracedReasoningDelta(reasoning)
}
}
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)
writeTracedContentDelta(pending, pending, false)
}
}
nativeToolCalls = normalizeNativeToolCallsForResponse(nativeToolCalls, req.Tools)
if traceStream {
s.logger.Info("openai chat completion stream output chunk",
zap.String("run_id", handle.Dispatch().RunID),
zap.Int("seq", traceSeq+1),
zap.String("type", "tool_calls"),
zap.Int("tool_call_count", len(nativeToolCalls)),
)
}
writeToolCallsDeltaSSE(w, flusher, id, created, model, nativeToolCalls)
finishReason = "tool_calls"
} else if len(req.Tools) > 0 {
res := synthesizeToolCallsFromTextResult(contentBuilder.String(), req.Tools, handle.Dispatch().RunID)
if res.validationErr != nil {
writeSSEErrorWithType(w, flusher, "tool_validation_error", res.validationErr.Error())
return
}
if len(res.toolCalls) > 0 {
if remaining := unstreamedCleanedText(res.cleaned, emittedContent.String()); remaining != "" {
emittedContent.WriteString(remaining)
writeTracedContentDelta(remaining, remaining, false)
}
if traceStream {
s.logger.Info("openai chat completion stream output chunk",
zap.String("run_id", handle.Dispatch().RunID),
zap.Int("seq", traceSeq+1),
zap.String("type", "tool_calls"),
zap.Int("tool_call_count", len(res.toolCalls)),
)
}
writeToolCallsDeltaSSE(w, flusher, id, created, model, res.toolCalls)
finishReason = "tool_calls"
} else if !res.candidateFound {
if toolTextFilter != nil {
if pending := toolTextFilter.Flush(); pending != "" {
emittedContent.WriteString(pending)
writeTracedContentDelta(pending, pending, false)
}
}
}
} else if toolTextFilter != nil {
if pending := toolTextFilter.Flush(); pending != "" {
emittedContent.WriteString(pending)
writeTracedContentDelta(pending, pending, false)
}
}
if !outputPolicy.Strict &&
strings.TrimSpace(contentBuilder.String()) == "" &&
strings.TrimSpace(emittedContent.String()) == "" &&
strings.TrimSpace(reasoningBuilder.String()) != "" &&
len(nativeToolCalls) == 0 {
fallback := hiddenReasoningFallbackContent(finishReason)
if req.includeReasoning() {
fallback = reasoningOnlyFallbackContent(reasoningBuilder.String(), finishReason)
}
emittedContent.WriteString(fallback)
writeTracedContentDelta(fallback, reasoningBuilder.String(), false)
}
if traceStream {
s.logger.Info("openai chat completion stream output done",
zap.String("run_id", handle.Dispatch().RunID),
zap.Int("seq", traceSeq+1),
zap.String("finish_reason", finishReason),
zap.Int("content_len", contentBuilder.Len()),
zap.Int("reasoning_len", reasoningBuilder.Len()),
)
}
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
}
verr := result.toolValidationErr
if verr == nil {
verr = validateToolCallResponse(validation, result.toolCalls, result.toolCallOrigin)
}
if 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, openAICompatTraceStreamEnabled(submitReq.Metadata))
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, traceStream bool) {
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"},
}},
})
if traceStream {
logOpenAICompatStreamOutput(s.logger, dispatch.RunID, 1, "content", content)
}
writeContentDeltaSSE(w, flusher, id, created, model, content)
if traceStream && len(result.toolCalls) > 0 {
s.logger.Info("openai chat completion stream output chunk",
zap.String("run_id", dispatch.RunID),
zap.Int("seq", 2),
zap.String("type", "tool_calls"),
zap.Int("tool_call_count", len(result.toolCalls)),
)
}
writeToolCallsDeltaSSE(w, flusher, id, created, model, result.toolCalls)
if traceStream {
s.logger.Info("openai chat completion stream output done",
zap.String("run_id", dispatch.RunID),
zap.Int("seq", 3),
zap.String("finish_reason", result.finishReason),
zap.Int("content_len", len(content)),
zap.Int("reasoning_len", result.reasoningLen),
)
}
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()
}
func openAICompatTraceStreamEnabled(metadata map[string]string) bool {
if traceBool(os.Getenv(streamTraceEnvKey)) {
return true
}
if metadata == nil {
return false
}
return traceBool(metadata[streamTraceMetadataKey])
}
func traceBool(v string) bool {
switch strings.ToLower(strings.TrimSpace(v)) {
case "1", "true", "yes", "on":
return true
default:
return false
}
}
func logOpenAICompatStreamOutput(logger *zap.Logger, runID string, seq int, typ string, delta string, extra ...zap.Field) {
fields := []zap.Field{
zap.String("run_id", runID),
zap.Int("seq", seq),
zap.String("type", typ),
zap.Int("delta_len", len(delta)),
zap.Bool("delta_has_open_think", hasOpenThinkTag(delta)),
zap.Bool("delta_has_close_think", hasCloseThinkTag(delta)),
zap.String("delta_preview", previewString(delta, 2000)),
}
fields = append(fields, extra...)
logger.Info("openai chat completion stream output chunk", fields...)
}
func hasOpenThinkTag(s string) bool {
return strings.Contains(strings.ToLower(s), "<think")
}
func hasCloseThinkTag(s string) bool {
return strings.Contains(strings.ToLower(s), "</think")
}