iop/apps/node/internal/adapters/openai_compat/stream.go

446 lines
12 KiB
Go

package openai_compat
import (
"bufio"
"context"
"encoding/json"
"fmt"
"io"
"net/url"
"os"
"strings"
"time"
"go.uber.org/zap"
"iop/apps/node/internal/runtime"
)
const (
runtimeMetadataOpenAIToolCalls = "openai_tool_calls"
runtimeMetadataOpenAITextToolFallback = "openai_text_tool_fallback"
streamTraceMetadataKey = "iop_trace_stream"
streamTraceEnvKey = "IOP_OPENAI_COMPAT_TRACE_STREAM"
)
// Execute runs the OpenAI-compatible chat completions stream. It builds the
// request body, sends the streaming request, and emits RuntimeEvents for
// reasoning deltas, content deltas, and completion.
func (a *Adapter) Execute(ctx context.Context, spec runtime.ExecutionSpec, sink runtime.EventSink) error {
if a.endpoint == "" {
return fmt.Errorf("openai_compat adapter: endpoint is required")
}
model := strings.TrimSpace(spec.Target)
if model == "" {
model = stringInput(spec.Input, "model")
}
if model == "" {
return fmt.Errorf("openai_compat adapter: target/model is required")
}
messages := messagesFromInput(spec.Input)
if len(messages) == 0 {
return fmt.Errorf("openai_compat adapter: messages are required")
}
if err := sink.Emit(ctx, runtime.RuntimeEvent{
RunID: spec.RunID,
Type: runtime.EventTypeStart,
Timestamp: time.Now(),
}); err != nil {
return err
}
reqBody, err := a.buildRequestBody(model, messages, spec.Input)
if err != nil {
_ = emitError(ctx, sink, spec.RunID, fmt.Sprintf("openai_compat build request body failed: %v", err))
return fmt.Errorf("openai_compat adapter: build request body: %w", err)
}
body, err := json.Marshal(reqBody)
if err != nil {
return fmt.Errorf("openai_compat adapter: marshal request: %w", err)
}
textToolFallback := false
a.logger.Info("openai_compat adapter executing",
zap.String("run_id", spec.RunID),
zap.String("provider", a.provider),
zap.String("target", model),
zap.String("endpoint", a.endpoint),
)
resp, err := a.doChatCompletion(ctx, body)
if err != nil {
_ = emitError(ctx, sink, spec.RunID, fmt.Sprintf("openai_compat request failed: %v", err))
return fmt.Errorf("openai_compat adapter: request: %w", err)
}
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
msg := readLimited(resp.Body, 4096)
_ = resp.Body.Close()
if retryBody, ok := retryBodyWithForcedSingleTool(reqBody, msg); ok {
body, err = json.Marshal(retryBody)
if err != nil {
_ = emitError(ctx, sink, spec.RunID, fmt.Sprintf("openai_compat retry marshal failed: %v", err))
return fmt.Errorf("openai_compat adapter: retry marshal: %w", err)
}
resp, err = a.doChatCompletion(ctx, body)
if err != nil {
_ = emitError(ctx, sink, spec.RunID, fmt.Sprintf("openai_compat retry request failed: %v", err))
return fmt.Errorf("openai_compat adapter: retry request: %w", err)
}
if resp.StatusCode >= 200 && resp.StatusCode < 300 {
goto streamResponse
}
msg = readLimited(resp.Body, 4096)
_ = resp.Body.Close()
}
if retryBody, ok := retryBodyWithTextToolFallback(reqBody, msg); ok {
body, err = json.Marshal(retryBody)
if err != nil {
_ = emitError(ctx, sink, spec.RunID, fmt.Sprintf("openai_compat text tool fallback marshal failed: %v", err))
return fmt.Errorf("openai_compat adapter: text tool fallback marshal: %w", err)
}
resp, err = a.doChatCompletion(ctx, body)
if err != nil {
_ = emitError(ctx, sink, spec.RunID, fmt.Sprintf("openai_compat text tool fallback request failed: %v", err))
return fmt.Errorf("openai_compat adapter: text tool fallback request: %w", err)
}
if resp.StatusCode >= 200 && resp.StatusCode < 300 {
textToolFallback = true
goto streamResponse
}
msg = readLimited(resp.Body, 4096)
_ = resp.Body.Close()
}
if msg == "" {
msg = resp.Status
}
_ = emitError(ctx, sink, spec.RunID, fmt.Sprintf("openai_compat returned %s: %s", resp.Status, msg))
return fmt.Errorf("openai_compat adapter: non-2xx response: %s", resp.Status)
}
streamResponse:
defer resp.Body.Close()
scanner := bufio.NewScanner(resp.Body)
outputTokens := 0
finishReason := ""
var usage *runtime.UsageStats
var toolCalls openAIToolCallAccumulator
traceStream := openAICompatTraceStreamEnabled(spec.Metadata)
traceSeq := 0
for scanner.Scan() {
line := scanner.Text()
if !strings.HasPrefix(line, "data: ") {
continue
}
payload := strings.TrimPrefix(line, "data: ")
if payload == "[DONE]" {
if traceStream {
a.logger.Info("openai_compat provider stream done",
zap.String("run_id", spec.RunID),
zap.Int("seq", traceSeq+1),
)
}
return sink.Emit(ctx, completeEvent(spec.RunID, finishReason, usage, outputTokens, toolCalls.ToolCalls(), textToolFallback))
}
traceSeq++
if traceStream {
a.logger.Info("openai_compat provider stream raw chunk",
zap.String("run_id", spec.RunID),
zap.Int("seq", traceSeq),
zap.Int("raw_len", len(payload)),
zap.Bool("raw_has_open_think", hasOpenThinkTag(payload)),
zap.Bool("raw_has_close_think", hasCloseThinkTag(payload)),
zap.String("raw_preview", tracePreview(payload, 4000)),
)
}
var chunk chatChunk
if err := json.Unmarshal([]byte(payload), &chunk); err != nil {
if traceStream {
a.logger.Warn("openai_compat provider stream raw chunk ignored",
zap.String("run_id", spec.RunID),
zap.Int("seq", traceSeq),
zap.Error(err),
)
}
continue
}
if chunk.Usage != nil {
usage = &runtime.UsageStats{
InputTokens: chunk.Usage.PromptTokens,
OutputTokens: chunk.Usage.CompletionTokens,
}
if d := chunk.Usage.PromptTokensDetails; d != nil {
usage.CachedInputTokens = d.CachedTokens
}
if d := chunk.Usage.CompletionTokensDetails; d != nil {
usage.ReasoningTokens = d.ReasoningTokens
}
}
if len(chunk.Choices) == 0 {
continue
}
choice := chunk.Choices[0]
if choice.FinishReason != nil && *choice.FinishReason != "" {
finishReason = *choice.FinishReason
}
if len(choice.Delta.ToolCalls) > 0 {
toolCalls.AddDelta(choice.Delta.ToolCalls)
}
reasoning := choice.Delta.ReasoningText()
if traceStream {
a.logger.Info("openai_compat provider stream parsed chunk",
zap.String("run_id", spec.RunID),
zap.Int("seq", traceSeq),
zap.String("finish_reason", finishReason),
zap.Int("tool_call_delta_count", len(choice.Delta.ToolCalls)),
zap.Int("content_len", len(choice.Delta.Content)),
zap.Bool("content_has_open_think", hasOpenThinkTag(choice.Delta.Content)),
zap.Bool("content_has_close_think", hasCloseThinkTag(choice.Delta.Content)),
zap.String("content_preview", tracePreview(choice.Delta.Content, 2000)),
zap.Int("reasoning_len", len(reasoning)),
zap.Bool("reasoning_has_open_think", hasOpenThinkTag(reasoning)),
zap.Bool("reasoning_has_close_think", hasCloseThinkTag(reasoning)),
zap.String("reasoning_preview", tracePreview(reasoning, 2000)),
)
}
if reasoning != "" {
if err := sink.Emit(ctx, runtime.RuntimeEvent{
RunID: spec.RunID,
Type: runtime.EventTypeReasoningDelta,
Delta: reasoning,
Timestamp: time.Now(),
}); err != nil {
return err
}
}
if content := choice.Delta.Content; content != "" {
outputTokens += len(strings.Fields(content))
if err := sink.Emit(ctx, runtime.RuntimeEvent{
RunID: spec.RunID,
Type: runtime.EventTypeDelta,
Delta: content,
Timestamp: time.Now(),
}); err != nil {
return err
}
}
}
if err := scanner.Err(); err != nil {
_ = emitError(ctx, sink, spec.RunID, fmt.Sprintf("openai_compat stream scan failed: %v", err))
return fmt.Errorf("openai_compat adapter: scan stream: %w", err)
}
_ = emitError(ctx, sink, spec.RunID, "openai_compat stream ended without [DONE]")
return fmt.Errorf("openai_compat adapter: stream ended without [DONE]")
}
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 tracePreview(s string, max int) string {
if max <= 0 || len(s) <= max {
return s
}
return s[:max] + "...[truncated]"
}
func hasOpenThinkTag(s string) bool {
return strings.Contains(strings.ToLower(s), "<think")
}
func hasCloseThinkTag(s string) bool {
return strings.Contains(strings.ToLower(s), "</think")
}
// joinOpenAIPath appends an OpenAI-compatible path to the endpoint. When the
// endpoint already targets the /v1 root it strips the leading /v1 from the path
// so a configured ".../v1" endpoint does not produce ".../v1/v1/...".
func joinOpenAIPath(baseURL, path string) string {
u, err := url.Parse(baseURL)
if err != nil {
return strings.TrimRight(baseURL, "/") + path
}
basePath := strings.TrimRight(u.Path, "/")
if strings.HasSuffix(basePath, "/v1") && strings.HasPrefix(path, "/v1/") {
path = strings.TrimPrefix(path, "/v1")
}
u.Path = basePath + path
return u.String()
}
func readLimited(r io.Reader, limit int64) string {
b, _ := io.ReadAll(io.LimitReader(r, limit))
return strings.TrimSpace(string(b))
}
type chatMessage struct {
Role string `json:"role"`
Content string `json:"content"`
ToolCalls []any `json:"tool_calls,omitempty"`
ToolCallID string `json:"tool_call_id,omitempty"`
}
type chatChunk struct {
Choices []struct {
Delta chatChunkDelta `json:"delta"`
FinishReason *string `json:"finish_reason"`
} `json:"choices"`
Usage *struct {
PromptTokens int `json:"prompt_tokens"`
CompletionTokens int `json:"completion_tokens"`
PromptTokensDetails *struct {
CachedTokens int `json:"cached_tokens"`
} `json:"prompt_tokens_details"`
CompletionTokensDetails *struct {
ReasoningTokens int `json:"reasoning_tokens"`
} `json:"completion_tokens_details"`
} `json:"usage"`
}
type chatChunkDelta struct {
Content string `json:"content"`
ReasoningContent string `json:"reasoning_content"`
Reasoning string `json:"reasoning"`
ToolCalls []any `json:"tool_calls"`
}
func (d chatChunkDelta) ReasoningText() string {
if d.ReasoningContent != "" {
return d.ReasoningContent
}
return d.Reasoning
}
type openAIToolCallAccumulator struct {
calls map[int]map[string]any
order []int
}
func (a *openAIToolCallAccumulator) AddDelta(raw []any) {
for fallbackIndex, item := range raw {
call, ok := item.(map[string]any)
if !ok {
continue
}
index := fallbackIndex
if parsed, ok := numericIndex(call["index"]); ok {
index = parsed
}
dst := a.ensure(index)
for key, value := range call {
switch key {
case "index":
continue
case "function":
mergeToolCallFunction(dst, value)
default:
if !emptyToolCallValue(value) {
dst[key] = value
}
}
}
}
}
func (a *openAIToolCallAccumulator) ToolCalls() []any {
if len(a.order) == 0 {
return nil
}
out := make([]any, 0, len(a.order))
for _, index := range a.order {
call := a.calls[index]
if len(call) == 0 {
continue
}
copyCall := make(map[string]any, len(call))
for key, value := range call {
copyCall[key] = value
}
out = append(out, copyCall)
}
return out
}
func (a *openAIToolCallAccumulator) ensure(index int) map[string]any {
if a.calls == nil {
a.calls = make(map[int]map[string]any)
}
if call, ok := a.calls[index]; ok {
return call
}
call := make(map[string]any)
a.calls[index] = call
a.order = append(a.order, index)
return call
}
func mergeToolCallFunction(call map[string]any, value any) {
fn, ok := value.(map[string]any)
if !ok {
return
}
current, _ := call["function"].(map[string]any)
if current == nil {
current = make(map[string]any)
call["function"] = current
}
for key, item := range fn {
if key == "arguments" {
if part, ok := item.(string); ok {
current["arguments"] = currentString(current["arguments"]) + part
continue
}
}
if !emptyToolCallValue(item) {
current[key] = item
}
}
}
func numericIndex(value any) (int, bool) {
switch v := value.(type) {
case int:
return v, true
case float64:
return int(v), true
default:
return 0, false
}
}
func emptyToolCallValue(value any) bool {
if value == nil {
return true
}
if text, ok := value.(string); ok {
return text == ""
}
return false
}
func currentString(value any) string {
if text, ok := value.(string); ok {
return text
}
return ""
}
type modelsResponse struct {
Data []struct {
ID string `json:"id"`
} `json:"data"`
}