Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
197 changes: 117 additions & 80 deletions core/relay/adaptor/openai/chat.go
Original file line number Diff line number Diff line change
Expand Up @@ -23,13 +23,73 @@ import (

// chatCompletionStreamState manages state for ChatCompletion stream conversion
type chatCompletionStreamState struct {
messageID string
meta *meta.Meta
c *gin.Context
currentToolCall *relaymodel.ToolCall
currentToolCallID string
toolCallArgs string
hasToolCall bool
messageID string
created int64
meta *meta.Meta
c *gin.Context
toolCallIndexByItemID map[string]int
toolCallIndexByOutputIndex map[int]int
nextToolCallIndex int
hasToolCall bool
}

func (s *chatCompletionStreamState) createdAt() int64 {
if s.created == 0 {
s.created = time.Now().Unix()
}

return s.created
}

func (s *chatCompletionStreamState) registerToolCall(
event *relaymodel.ResponseStreamEvent,
) int {
if event.Item.ID != "" {
if index, ok := s.toolCallIndexByItemID[event.Item.ID]; ok {
return index
}
}

if event.OutputIndex != nil {
if index, ok := s.toolCallIndexByOutputIndex[*event.OutputIndex]; ok {
return index
}
}

index := s.nextToolCallIndex
s.nextToolCallIndex++

if event.Item.ID != "" {
s.toolCallIndexByItemID[event.Item.ID] = index
}

if event.OutputIndex != nil {
s.toolCallIndexByOutputIndex[*event.OutputIndex] = index
}

return index
}

func (s *chatCompletionStreamState) toolCallIndex(
event *relaymodel.ResponseStreamEvent,
) (int, bool) {
if event.ItemID != "" {
index, ok := s.toolCallIndexByItemID[event.ItemID]
if ok {
return index, true
}
}

if event.OutputIndex != nil {
index, ok := s.toolCallIndexByOutputIndex[*event.OutputIndex]
return index, ok
}

if event.ItemID == "" && s.nextToolCallIndex > 0 {
return s.nextToolCallIndex - 1, true
}

return 0, false
}

func responseModelName(meta *meta.Meta) string {
Expand Down Expand Up @@ -132,11 +192,14 @@ func (s *chatCompletionStreamState) handleResponseCreated(
}

s.messageID = event.Response.ID
if event.Response.CreatedAt != 0 {
s.created = event.Response.CreatedAt
}

return &relaymodel.ChatCompletionsStreamResponse{
ID: s.messageID,
Object: relaymodel.ChatCompletionChunkObject,
Created: event.Response.CreatedAt,
Created: s.createdAt(),
Model: responseModelName(s.meta),
Choices: []*relaymodel.ChatCompletionsStreamResponseChoice{
{
Expand All @@ -160,7 +223,7 @@ func (s *chatCompletionStreamState) handleOutputTextDelta(
return &relaymodel.ChatCompletionsStreamResponse{
ID: s.messageID,
Object: relaymodel.ChatCompletionChunkObject,
Created: time.Now().Unix(),
Created: s.createdAt(),
Model: responseModelName(s.meta),
Choices: []*relaymodel.ChatCompletionsStreamResponseChoice{
{
Expand All @@ -183,7 +246,7 @@ func (s *chatCompletionStreamState) handleReasoningSummaryTextDelta(
return &relaymodel.ChatCompletionsStreamResponse{
ID: s.messageID,
Object: relaymodel.ChatCompletionChunkObject,
Created: time.Now().Unix(),
Created: s.createdAt(),
Model: responseModelName(s.meta),
Choices: []*relaymodel.ChatCompletionsStreamResponseChoice{
{
Expand All @@ -207,30 +270,21 @@ func (s *chatCompletionStreamState) handleOutputItemAdded(
// Track function calls
if event.Item.Type == relaymodel.InputItemTypeFunctionCall {
s.hasToolCall = true
s.currentToolCallID = event.Item.ID
s.currentToolCall = &relaymodel.ToolCall{
ID: event.Item.CallID,
Type: relaymodel.ToolChoiceTypeFunction,
Function: relaymodel.Function{
Name: event.Item.Name,
Arguments: "",
},
}
s.toolCallArgs = ""
toolCallIndex := s.registerToolCall(event)

// Send tool call start
return &relaymodel.ChatCompletionsStreamResponse{
ID: s.messageID,
Object: relaymodel.ChatCompletionChunkObject,
Created: time.Now().Unix(),
Created: s.createdAt(),
Model: responseModelName(s.meta),
Choices: []*relaymodel.ChatCompletionsStreamResponseChoice{
{
Index: 0,
Delta: relaymodel.Message{
ToolCalls: []relaymodel.ToolCall{
{
Index: 0,
Index: toolCallIndex,
ID: event.Item.CallID,
Type: relaymodel.ToolChoiceTypeFunction,
Function: relaymodel.Function{
Expand All @@ -249,7 +303,7 @@ func (s *chatCompletionStreamState) handleOutputItemAdded(
return &relaymodel.ChatCompletionsStreamResponse{
ID: s.messageID,
Object: relaymodel.ChatCompletionChunkObject,
Created: time.Now().Unix(),
Created: s.createdAt(),
Model: responseModelName(s.meta),
Choices: []*relaymodel.ChatCompletionsStreamResponseChoice{
{
Expand All @@ -269,26 +323,28 @@ func (s *chatCompletionStreamState) handleOutputItemAdded(
func (s *chatCompletionStreamState) handleFunctionCallArgumentsDelta(
event *relaymodel.ResponseStreamEvent,
) *relaymodel.ChatCompletionsStreamResponse {
if event.Delta == "" || s.currentToolCall == nil {
if event.Delta == "" {
return nil
}

// Accumulate arguments
s.toolCallArgs += event.Delta
toolCallIndex, ok := s.toolCallIndex(event)
if !ok {
return nil
}

// Send delta
return &relaymodel.ChatCompletionsStreamResponse{
ID: s.messageID,
Object: relaymodel.ChatCompletionChunkObject,
Created: time.Now().Unix(),
Created: s.createdAt(),
Model: responseModelName(s.meta),
Choices: []*relaymodel.ChatCompletionsStreamResponseChoice{
{
Index: 0,
Delta: relaymodel.Message{
ToolCalls: []relaymodel.ToolCall{
{
Index: 0,
Index: toolCallIndex,
Function: relaymodel.Function{
Arguments: event.Delta,
},
Expand All @@ -300,32 +356,6 @@ func (s *chatCompletionStreamState) handleFunctionCallArgumentsDelta(
}
}

// handleOutputItemDone handles response.output_item.done event for ChatCompletion
func (s *chatCompletionStreamState) handleOutputItemDone(
event *relaymodel.ResponseStreamEvent,
) {
if event.Item == nil {
return
}

// Handle function call completion
if event.Item.Type == relaymodel.InputItemTypeFunctionCall && s.currentToolCall != nil &&
event.Item.ID == s.currentToolCallID {
// Update with final arguments
if s.toolCallArgs != "" {
s.currentToolCall.Function.Arguments = s.toolCallArgs
}

// Reset state
s.currentToolCall = nil
s.currentToolCallID = ""
s.toolCallArgs = ""

// No need to send another chunk - arguments already streamed
return
}
}

// handleResponseCompleted handles response.completed/done event for ChatCompletion
func (s *chatCompletionStreamState) handleResponseCompleted(
event *relaymodel.ResponseStreamEvent,
Expand All @@ -344,7 +374,7 @@ func (s *chatCompletionStreamState) handleResponseCompleted(
return &relaymodel.ChatCompletionsStreamResponse{
ID: s.messageID,
Object: relaymodel.ChatCompletionChunkObject,
Created: time.Now().Unix(),
Created: s.createdAt(),
Model: responseModelName(s.meta),
Choices: []*relaymodel.ChatCompletionsStreamResponseChoice{
{
Expand Down Expand Up @@ -1341,6 +1371,8 @@ func ConvertResponsesToChatCompletionResponse(

reasonContent := responseReasoningSummaryText(&responsesResp)

var toolCallChoice *relaymodel.TextResponseChoice

// Convert output items to choices
for _, outputItem := range responsesResp.Output {
switch outputItem.Type {
Expand Down Expand Up @@ -1380,31 +1412,36 @@ func ConvertResponsesToChatCompletionResponse(
toolCallID = outputItem.ID
}

finishReason := responseToChatFinishReason(&responsesResp)
if finishReason == relaymodel.FinishReasonStop {
finishReason = relaymodel.FinishReasonToolCalls
if toolCallChoice == nil {
finishReason := responseToChatFinishReason(&responsesResp)
if finishReason == relaymodel.FinishReasonStop {
finishReason = relaymodel.FinishReasonToolCalls
}

toolCallChoice = &relaymodel.TextResponseChoice{
Index: len(chatResp.Choices),
Message: relaymodel.Message{
Role: relaymodel.RoleAssistant,
ReasoningContent: reasonContent,
},
FinishReason: finishReason,
}
chatResp.Choices = append(chatResp.Choices, toolCallChoice)
reasonContent = ""
}

chatResp.Choices = append(chatResp.Choices, &relaymodel.TextResponseChoice{
Index: len(chatResp.Choices),
Message: relaymodel.Message{
Role: relaymodel.RoleAssistant,
ReasoningContent: reasonContent,
ToolCalls: []relaymodel.ToolCall{
{
Index: 0,
ID: toolCallID,
Type: relaymodel.ToolChoiceTypeFunction,
Function: relaymodel.Function{
Name: outputItem.Name,
Arguments: outputItem.Arguments.String(),
},
},
toolCallChoice.Message.ToolCalls = append(
toolCallChoice.Message.ToolCalls,
relaymodel.ToolCall{
Index: len(toolCallChoice.Message.ToolCalls),
ID: toolCallID,
Type: relaymodel.ToolChoiceTypeFunction,
Function: relaymodel.Function{
Name: outputItem.Name,
Arguments: outputItem.Arguments.String(),
},
},
FinishReason: finishReason,
})
reasonContent = ""
)

default:
continue
Expand Down Expand Up @@ -1556,8 +1593,10 @@ func ConvertResponsesToChatCompletionStreamResponse(
errorState := responsesStreamErrorState{}

state := &chatCompletionStreamState{
meta: meta,
c: c,
meta: meta,
c: c,
toolCallIndexByItemID: make(map[string]int),
toolCallIndexByOutputIndex: make(map[int]int),
}
stopStream := false

Expand Down Expand Up @@ -1635,8 +1674,6 @@ func ConvertResponsesToChatCompletionStreamResponse(
chatStreamResp = state.handleOutputItemAdded(&event)
case relaymodel.EventFunctionCallArgumentsDelta:
chatStreamResp = state.handleFunctionCallArgumentsDelta(&event)
case relaymodel.EventOutputItemDone:
state.handleOutputItemDone(&event)
case relaymodel.EventResponseCompleted,
relaymodel.EventResponseIncomplete,
relaymodel.EventResponseDone:
Expand Down
Loading
Loading