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
36 changes: 34 additions & 2 deletions internal/modelprovider/openai.go
Original file line number Diff line number Diff line change
Expand Up @@ -259,7 +259,7 @@ func CheckChatCompletionsAPIWithClient(ctx context.Context, client *http.Client,
"messages": []map[string]string{
{"role": "user", "content": "ping"},
},
"stream": false,
"stream": true,
"max_tokens": 16,
}
body, err := json.Marshal(payload)
Expand All @@ -271,6 +271,7 @@ func CheckChatCompletionsAPIWithClient(ctx context.Context, client *http.Client,
return fmt.Errorf("build chat completions probe request: %w", err)
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Accept", "text/event-stream")
if apiKey = strings.TrimSpace(apiKey); apiKey != "" {
req.Header.Set("Authorization", "Bearer "+apiKey)
}
Expand All @@ -297,11 +298,42 @@ func CheckChatCompletionsAPIWithClient(ctx context.Context, client *http.Client,
}

var probe openAIChatCompletionsProbeResponse
if err := json.NewDecoder(resp.Body).Decode(&probe); err != nil {
mediaType := strings.ToLower(strings.TrimSpace(strings.Split(resp.Header.Get("Content-Type"), ";")[0]))
if mediaType == "text/event-stream" {
probe, err = decodeOpenAIChatCompletionsProbeStream(resp.Body)
} else {
err = json.NewDecoder(resp.Body).Decode(&probe)
}
if err != nil {
return fmt.Errorf("decode chat completions probe from %s: %w", baseURL, err)
}
if len(probe.Choices) == 0 {
return fmt.Errorf("chat completions probe from %s returned no choices", baseURL)
}
return nil
}

func decodeOpenAIChatCompletionsProbeStream(r io.Reader) (openAIChatCompletionsProbeResponse, error) {
scanner := bufio.NewScanner(r)
for scanner.Scan() {
line := strings.TrimSpace(scanner.Text())
if !strings.HasPrefix(line, "data:") {
continue
}
data := strings.TrimSpace(strings.TrimPrefix(line, "data:"))
if data == "" || data == "[DONE]" {
continue
}
var probe openAIChatCompletionsProbeResponse
if err := json.Unmarshal([]byte(data), &probe); err != nil {
return openAIChatCompletionsProbeResponse{}, err
}
if len(probe.Choices) > 0 {
return probe, nil
}
}
if err := scanner.Err(); err != nil {
return openAIChatCompletionsProbeResponse{}, err
}
return openAIChatCompletionsProbeResponse{}, fmt.Errorf("stream returned no chat completion chunk")
}
25 changes: 25 additions & 0 deletions internal/modelprovider/openai_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -140,6 +140,28 @@ func TestCheckResponsesAPIWithClientClassifiesUnsupportedEndpoint(t *testing.T)
}
}

func TestCheckChatCompletionsAPIWithClientAcceptsStreamingResponse(t *testing.T) {
var gotPayload map[string]any
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if got := r.Header.Get("Accept"); got != "text/event-stream" {
t.Fatalf("Accept = %q, want text/event-stream", got)
}
if err := json.NewDecoder(r.Body).Decode(&gotPayload); err != nil {
t.Fatalf("Decode() error = %v", err)
}
w.Header().Set("Content-Type", "text/event-stream; charset=utf-8")
_, _ = w.Write([]byte("data: {\"id\":\"chatcmpl-test\",\"object\":\"chat.completion.chunk\",\"choices\":[{\"index\":0,\"delta\":{\"role\":\"assistant\",\"content\":\"pong\"},\"finish_reason\":null}]}\n\n"))
}))
defer srv.Close()

if err := CheckChatCompletionsAPIWithClient(context.Background(), srv.Client(), srv.URL+"/v1", "sk-test", "gpt-test", nil); err != nil {
t.Fatalf("CheckChatCompletionsAPIWithClient() error = %v", err)
}
if gotPayload["stream"] != true {
t.Fatalf("stream = %#v, want true", gotPayload["stream"])
}
}

func TestCheckResponsesOrChatCompletionsAPIWithClientFallsBackToChat(t *testing.T) {
var chatPayload map[string]any
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
Expand All @@ -165,6 +187,9 @@ func TestCheckResponsesOrChatCompletionsAPIWithClientFallsBackToChat(t *testing.
if chatPayload["model"] != "gpt-test" {
t.Fatalf("chat model = %#v, want gpt-test", chatPayload["model"])
}
if chatPayload["stream"] != true {
t.Fatalf("chat stream = %#v, want true", chatPayload["stream"])
}
messages, ok := chatPayload["messages"].([]any)
if !ok || len(messages) != 1 {
t.Fatalf("chat messages = %#v, want one probe message", chatPayload["messages"])
Expand Down
Loading