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

465 lines
16 KiB
Go

package openai
import (
"bytes"
"context"
"encoding/json"
"fmt"
"net/http"
"strings"
"sync"
"time"
"go.uber.org/zap"
edgeservice "iop/apps/edge/internal/service"
"iop/packages/go/streamgate"
)
type openAIResponsesAttemptResult struct {
text string
reasoning string
toolCalls []any
usage *openAIUsage
dispatch edgeservice.RunDispatch
collectErr error
}
type openAIResponsesResultHolder struct {
mu sync.Mutex
result openAIResponsesAttemptResult
set bool
}
func (h *openAIResponsesResultHolder) beginAttempt() {
h.mu.Lock()
h.result = openAIResponsesAttemptResult{}
h.set = false
h.mu.Unlock()
}
func (h *openAIResponsesResultHolder) store(result openAIResponsesAttemptResult) {
h.mu.Lock()
h.result = result
h.set = true
h.mu.Unlock()
}
func (h *openAIResponsesResultHolder) get() (openAIResponsesAttemptResult, bool) {
h.mu.Lock()
defer h.mu.Unlock()
return h.result, h.set
}
type openAIResponsesAttemptContext struct {
mu sync.Mutex
dc *responsesDispatchContext
}
func (s *openAIResponsesAttemptContext) set(dc *responsesDispatchContext) {
s.mu.Lock()
s.dc = dc
s.mu.Unlock()
}
func (s *openAIResponsesAttemptContext) get() *responsesDispatchContext {
s.mu.Lock()
defer s.mu.Unlock()
return s.dc
}
type openAIResponsesEventSource struct {
dc *responsesDispatchContext
handle edgeservice.RunResult
holder *openAIResponsesResultHolder
usage *openAIStreamGateUsageHolder
mu sync.Mutex
started bool
loaded bool
pending []streamgate.NormalizedEvent
}
func newOpenAIResponsesEventSource(dc *responsesDispatchContext, handle edgeservice.RunResult, holder *openAIResponsesResultHolder, usage *openAIStreamGateUsageHolder) *openAIResponsesEventSource {
return &openAIResponsesEventSource{dc: dc, handle: handle, holder: holder, usage: usage}
}
func (s *openAIResponsesEventSource) NextEvent(ctx context.Context) (streamgate.NormalizedEvent, error) {
s.mu.Lock()
if !s.started {
s.started = true
s.holder.beginAttempt()
s.mu.Unlock()
return streamgate.NewResponseStartEvent(streamGateChannelDefault, http.StatusOK, map[string]string{"Content-Type": "application/json"}, time.Now())
}
if len(s.pending) > 0 {
event := s.pending[0]
s.pending = s.pending[1:]
s.mu.Unlock()
return event, nil
}
if s.loaded {
s.mu.Unlock()
return newOpenAIProviderErrorEvent(streamGateErrorStreamClosed)
}
s.loaded = true
s.mu.Unlock()
text, reasoning, _, toolCalls, usage, _, err := collectRunResult(ctx, s.handle.Stream(), s.handle.WaitTimeout())
if err != nil {
s.holder.store(openAIResponsesAttemptResult{dispatch: s.handle.Dispatch(), collectErr: err})
return newOpenAIProviderErrorEvent(streamGateErrorRunFailed)
}
text, reasoning, _ = normalizeCompletionOutput(s.dc.outputPolicy, text, reasoning, false)
result := openAIResponsesAttemptResult{text: text, reasoning: reasoning, toolCalls: toolCalls, usage: usage, dispatch: s.handle.Dispatch()}
s.holder.store(result)
if s.usage != nil {
s.usage.set(usageObservationFromOpenAIUsage(usage, len(reasoning)))
}
var events []streamgate.NormalizedEvent
if text != "" {
event, eventErr := streamgate.NewTextDeltaEvent(streamGateChannelDefault, text, time.Now())
if eventErr != nil {
return streamgate.NormalizedEvent{}, eventErr
}
events = append(events, event)
}
if reasoning != "" {
event, eventErr := streamgate.NewReasoningDeltaEvent(streamGateChannelDefault, reasoning, time.Now())
if eventErr != nil {
return streamgate.NormalizedEvent{}, eventErr
}
events = append(events, event)
}
for i, raw := range toolCalls {
call, ok := decodeResponsesToolCall(raw, i)
if !ok || call.Arguments == "" {
continue
}
event, eventErr := streamgate.NewToolCallFragmentEvent(streamGateChannelDefault, call.CallID, call.Name, call.Arguments, time.Now())
if eventErr != nil {
return streamgate.NormalizedEvent{}, eventErr
}
events = append(events, event)
}
terminal, terminalErr := streamgate.NewTerminalEvent(streamGateChannelDefault, time.Now())
if terminalErr != nil {
return streamgate.NormalizedEvent{}, terminalErr
}
events = append(events, terminal)
s.mu.Lock()
s.pending = append(s.pending, events...)
event := s.pending[0]
s.pending = s.pending[1:]
s.mu.Unlock()
return event, nil
}
var _ streamgate.NormalizedEventSource = (*openAIResponsesEventSource)(nil)
type openAIResponsesToolCall struct {
ID string
CallID string
Name string
Arguments string
}
func decodeResponsesToolCall(raw any, index int) (openAIResponsesToolCall, bool) {
encoded, err := json.Marshal(raw)
if err != nil {
return openAIResponsesToolCall{}, false
}
var value struct {
ID string `json:"id"`
CallID string `json:"call_id"`
Name string `json:"name"`
Arguments string `json:"arguments"`
Function struct {
Name string `json:"name"`
Arguments string `json:"arguments"`
} `json:"function"`
}
if json.Unmarshal(encoded, &value) != nil {
return openAIResponsesToolCall{}, false
}
name := value.Name
if name == "" {
name = value.Function.Name
}
args := value.Arguments
if args == "" {
args = value.Function.Arguments
}
callID := value.CallID
if callID == "" {
callID = value.ID
}
if callID == "" {
callID = fmt.Sprintf("call-%d", index)
}
id := value.ID
if id == "" {
id = fmt.Sprintf("fc-%d", index)
}
if name == "" {
name = "function"
}
return openAIResponsesToolCall{ID: id, CallID: callID, Name: name, Arguments: args}, true
}
func responsesOutputItems(text string, toolCalls []any) []responsesOutputItem {
items := []responsesOutputItem{{
Type: "message",
Role: "assistant",
Content: []responsesContentItem{{Type: "output_text", Text: text}},
}}
for i, raw := range toolCalls {
call, ok := decodeResponsesToolCall(raw, i)
if !ok {
continue
}
items = append(items, responsesOutputItem{
Type: "function_call", ID: call.ID, CallID: call.CallID,
Name: call.Name, Arguments: call.Arguments,
})
}
return items
}
type openAIResponsesReleaseSink struct {
server *Server
w http.ResponseWriter
req responsesRequest
holder *openAIResponsesResultHolder
recoveryAdmission *openAIRecoveryAdmissionState
mu sync.Mutex
terminalCommitted bool
terminalSuccess bool
}
func (s *openAIResponsesReleaseSink) setRecoveryAdmissionState(state *openAIRecoveryAdmissionState) {
s.mu.Lock()
s.recoveryAdmission = state
s.mu.Unlock()
}
func newOpenAIResponsesReleaseSink(server *Server, w http.ResponseWriter, dc *responsesDispatchContext, holder *openAIResponsesResultHolder) *openAIResponsesReleaseSink {
return &openAIResponsesReleaseSink{server: server, w: w, req: dc.req, holder: holder}
}
func (s *openAIResponsesReleaseSink) terminalStatus() (bool, bool) {
s.mu.Lock()
defer s.mu.Unlock()
return s.terminalCommitted, s.terminalSuccess
}
func (s *openAIResponsesReleaseSink) CommitResponseStart(context.Context, streamgate.ResponseStart) (streamgate.CommitState, error) {
return streamgate.CommitStateStreamOpen, nil
}
func (s *openAIResponsesReleaseSink) Release(_ context.Context, event streamgate.ReleaseEvent) (streamgate.CommitState, error) {
switch event.Kind() {
case streamgate.EventKindTextDelta, streamgate.EventKindReasoningDelta, streamgate.EventKindToolCallFragment:
return streamgate.CommitStateStreamOpen, nil
default:
return streamgate.CommitStateStreamOpen, fmt.Errorf("openai stream gate: responses sink does not support %q", event.Kind())
}
}
func (s *openAIResponsesReleaseSink) CommitTerminal(_ context.Context, terminal streamgate.TerminalResult) (streamgate.CommitState, error) {
s.mu.Lock()
defer s.mu.Unlock()
s.terminalCommitted = true
s.terminalSuccess = terminal.Success()
result, ok := s.holder.get()
if !terminal.Success() && s.recoveryAdmission.rejected() {
writeError(s.w, http.StatusBadRequest, "invalid_request_error", openAIStreamGateCandidateRejectedMessage)
return streamgate.CommitStateTerminalCommitted, nil
}
if !terminal.Success() || !ok || result.collectErr != nil {
message := openAIStreamGateErrorMessage(terminal)
status := http.StatusBadGateway
if ok && result.collectErr != nil {
message = result.collectErr.Error()
status = httpStatusForRunError(result.collectErr)
}
writeError(s.w, status, "run_error", message)
return streamgate.CommitStateTerminalCommitted, nil
}
var usage openAIUsage
if result.usage != nil {
usage = *result.usage
}
s.server.logger.Info("openai responses output",
zap.String("run_id", result.dispatch.RunID),
zap.Int("content_len", len(result.text)),
zap.Int("reasoning_len", len(result.reasoning)),
)
writeJSON(s.w, http.StatusOK, responsesResponse{
ID: "resp-" + result.dispatch.RunID, Object: "response", CreatedAt: time.Now().Unix(),
Model: responseModel(s.req.Model, result.dispatch.Target), OutputText: result.text,
Output: responsesOutputItems(result.text, result.toolCalls), Usage: usage,
})
return streamgate.CommitStateTerminalCommitted, nil
}
var _ openAIStreamGateSink = (*openAIResponsesReleaseSink)(nil)
func newOpenAIResponsesRecoveryAdmissionBuilder(server *Server, initial *responsesDispatchContext, state *openAIResponsesAttemptContext) openAIAttemptAdmissionBuilder {
return func(ctx context.Context, request streamgate.RebuiltRequest, body []byte) (openAIAttemptAdmission, error) {
var req responsesRequest
if err := decodeResponsesRequest(json.NewDecoder(bytes.NewReader(body)), &req); err != nil {
return openAIAttemptAdmission{}, err
}
dc, err := server.newResponsesDispatchContext(initial.responsesRequestContext, req)
if err != nil {
return openAIAttemptAdmission{}, err
}
state.set(dc)
if initial.poolDispatch == nil {
return openAIAttemptAdmission{kind: openAIAdmissionRun, run: dc.submitReq}, nil
}
pool := *initial.poolDispatch
pool.Run = dc.submitReq
pool.Run.ProviderPool = true
pool.Tunnel.BuildBody = func(target string) ([]byte, error) {
return rewriteResponsesModel(body, target)
}
pool.PrepareRun = func(runReq edgeservice.SubmitRunRequest) (edgeservice.SubmitRunRequest, error) {
runReq.Prompt = dc.submitReq.Prompt
runReq.Input = dc.submitReq.Input
runReq.Metadata = dc.submitReq.Metadata
runReq.EstimatedInputTokens = dc.submitReq.EstimatedInputTokens
runReq.ContextClass = dc.submitReq.ContextClass
return runReq, nil
}
return openAIAttemptAdmission{kind: openAIAdmissionPool, pool: pool}, nil
}
}
func (s *Server) buildOpenAIResponsesStreamGateRuntime(dc *responsesDispatchContext, handle edgeservice.RunResult, sink openAIStreamGateSink, registry streamgate.FilterRegistrySnapshot) (*streamgate.RequestRuntime, *openAIStreamGateUsageHolder, error) {
holderSink, ok := sink.(*openAIResponsesReleaseSink)
if !ok {
if composite, compositeOK := sink.(*openAICompositeReleaseSink); compositeOK {
holderSink, _ = composite.normalized.(*openAIResponsesReleaseSink)
}
}
if holderSink == nil {
return nil, nil, fmt.Errorf("openai responses stream gate: normalized sink is required")
}
holder := holderSink.holder
usage := &openAIStreamGateUsageHolder{}
state := &openAIResponsesAttemptContext{dc: dc}
rebuilder, err := newOpenAIRequestRebuilder(dc.ingress, openAIRebuildEndpointResponses)
if err != nil {
return nil, nil, err
}
selector := newOpenAIStreamGateCodecSelector(openAIStreamGateCodecNormalized)
if composite, ok := sink.(*openAICompositeReleaseSink); ok {
selector = composite.selector
}
factory := func(transport openAIAttemptTransport) (streamgate.NormalizedEventSource, error) {
switch transport.path {
case openAIAdmissionRun:
selector.set(openAIStreamGateCodecNormalized)
attemptDC := state.get()
if attemptDC == nil || transport.run == nil {
return nil, fmt.Errorf("openai responses normalized attempt is incomplete")
}
return newOpenAIResponsesEventSource(attemptDC, transport.run, holder, usage), nil
case openAIAdmissionTunnel:
selector.set(openAIStreamGateCodecTunnel)
codecState := openAITunnelCodecStateForSink(sink)
codecState.reset()
assembler := &providerChatAssembler{streaming: dc.req.Stream}
rewriter := newProviderModelRewriter(dc.req.Stream, "")
return newOpenAITunnelEndpointEventSource(transport.tunnel.Stream(), transport.tunnel.WaitTimeout(), rewriter, assembler, openAIRebuildEndpointResponses, codecState), nil
default:
return nil, fmt.Errorf("openai responses unsupported attempt path %q", transport.path)
}
}
dispatcher, err := newOpenAIAttemptDispatcher(s.service, rebuilder.RebuiltStore(), newOpenAIResponsesRecoveryAdmissionBuilder(s, dc, state), factory)
if err != nil {
return nil, nil, err
}
bindOpenAIRecoveryAdmissionState(sink, dispatcher.admissionState())
initialSource := newOpenAIResponsesEventSource(dc, handle, holder, usage)
dispatch := handle.Dispatch()
controller := &openAIAttemptController{service: s.service, dispatch: dispatch, closeTransport: handle.Close}
binding, err := streamgate.NewAttemptBinding(
openAIStreamGateSafeToken("attempt", dispatch.RunID), actualOpenAIModel(dispatch), actualOpenAIProvider(dispatch),
actualOpenAIExecutionPath(dispatch, openAIAdmissionRun), initialSource, controller,
)
if err != nil {
return nil, nil, err
}
opts, err := s.streamGateRuntimeOptions()
if err != nil {
return nil, nil, err
}
snapRef, err := dc.ingress.recoveryRef()
if err != nil {
return nil, nil, err
}
snapshot, err := streamgate.NewRequestRuntimeSnapshot(
openAIStreamGateSafeToken("req", dispatch.RunID), streamGateConfigGeneration, s.streamGateConfig().EffectiveEnvironment(),
openAIRebuildEndpointResponses, openAIRebuildFamily, opts, registry, nil, snapRef, dispatcher, rebuilder, nil, nil, sink,
)
if err != nil {
return nil, nil, err
}
snapshot = snapshot.WithObservationSink(s.observationSink())
modelGroup := strings.TrimSpace(dc.req.Model)
if modelGroup == "" {
modelGroup = actualOpenAIModel(dispatch)
}
runtime, err := streamgate.NewRequestRuntime(snapshot, modelGroup, binding)
if err != nil {
return nil, nil, err
}
return runtime, usage, nil
}
func (s *Server) runOpenAIResponsesStreamGate(w http.ResponseWriter, dc *responsesDispatchContext, handle edgeservice.RunResult) {
holder := &openAIResponsesResultHolder{}
normalized := newOpenAIResponsesReleaseSink(s, w, dc, holder)
selector := newOpenAIStreamGateCodecSelector(openAIStreamGateCodecNormalized)
var sink openAIStreamGateSink = normalized
if dc.poolDispatch != nil {
tunnel := newOpenAIBufferedTunnelReleaseSink(w, nil, "")
sink = newOpenAICompositeReleaseSink(selector, normalized, tunnel)
}
fctx, err := s.openAIResponsesOutputFilterContext(dc.responsesRequestContext)
if err != nil {
handle.Close()
writeError(w, http.StatusInternalServerError, "run_error", "stream gate runtime unavailable")
return
}
registry, err := openAIStreamGateRegistrySnapshotFor(s.streamGateConfig(), fctx)
if err != nil {
handle.Close()
writeError(w, http.StatusInternalServerError, "run_error", "stream gate runtime unavailable")
return
}
runtime, usage, err := s.buildOpenAIResponsesStreamGateRuntime(dc, handle, sink, registry)
if err != nil {
handle.Close()
writeError(w, http.StatusInternalServerError, "run_error", "stream gate runtime unavailable")
return
}
runErr := runtime.Run(dc.r.Context())
committed, success := sink.terminalStatus()
_ = runtime.CloseRequestResources(context.Background(), runErr == nil && committed && success)
labels := dc.usageLabels(s, responseModeNormalized)
if composite, ok := sink.(*openAICompositeReleaseSink); ok && composite.resolvedCodec() == openAIStreamGateCodecTunnel {
labels = dc.usageLabels(s, responseModePassthrough)
}
status := streamGateUsageStatus(runErr, committed, success)
if status == usageStatusSuccess {
emitUsageMetrics(labels, status, usage.get())
return
}
emitUsageMetrics(labels, status, usageObservation{})
}