Skip to content
Open
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
55 changes: 49 additions & 6 deletions meat/openai.go
Original file line number Diff line number Diff line change
Expand Up @@ -78,8 +78,8 @@ type openAIReq struct {
Instructions string `json:"instructions,omitempty"`
Input []json.RawMessage `json:"input"`
Tools []openAITool `json:"tools,omitempty"`
Reasoning openAIReasoning `json:"reasoning"`
Include []string `json:"include"`
Reasoning *openAIReasoning `json:"reasoning,omitempty"`
Include []string `json:"include,omitempty"`
Store bool `json:"store"`
Stream bool `json:"stream"`
MaxOutputTokens int `json:"max_output_tokens"`
Expand Down Expand Up @@ -125,6 +125,7 @@ type openAIStreamEvent struct {
Response *openAIResp `json:"response"`
Code string `json:"code"`
Message string `json:"message"`
Error *openAIError `json:"error"`
}

type openAIOutputItem struct {
Expand Down Expand Up @@ -154,17 +155,20 @@ func (m *OpenAIModel) Generate(ctx context.Context, system string, messages []Me
if err != nil {
return nil, err
}
model := cmpOr(m.Model, DefaultOpenAIModel)
reqBody := openAIReq{
Model: cmpOr(m.Model, DefaultOpenAIModel),
Model: model,
Instructions: system,
Input: input,
Tools: toOpenAITools(tools),
Reasoning: openAIReasoning{Effort: cmpOr(m.ReasoningEffort, DefaultReasoningEffort)},
Include: []string{"reasoning.encrypted_content"},
Store: false,
Stream: true,
MaxOutputTokens: maxOpenAIOutputTokens,
}
if isOpenAIReasoningModel(model) {
reqBody.Reasoning = &openAIReasoning{Effort: cmpOr(m.ReasoningEffort, DefaultReasoningEffort)}
reqBody.Include = []string{"reasoning.encrypted_content"}
}
body, err := json.Marshal(reqBody)
if err != nil {
return nil, err
Expand Down Expand Up @@ -224,6 +228,16 @@ func openAIResponsesURL(base string) string {
return base + "/v1/responses"
}

func isOpenAIReasoningModel(model string) bool {
model = strings.ToLower(strings.TrimSpace(model))
for _, family := range []string{"gpt-5", "o1", "o3", "o4"} {
if model == family || strings.HasPrefix(model, family+"-") || strings.HasPrefix(model, family+".") {
return true
}
}
return false
}

func toOpenAITools(tools []Tool) []openAITool {
out := make([]openAITool, 0, len(tools))
for _, t := range tools {
Expand Down Expand Up @@ -410,7 +424,7 @@ func decodeOpenAIResponse(raw []byte) (openAIResp, error) {
final = &copyResp
}
case "error":
streamErr = fmt.Errorf("openai stream error %s: %s", event.Code, event.Message)
streamErr = openAIStreamError(event)
}
return nil
}
Expand Down Expand Up @@ -454,3 +468,32 @@ func decodeOpenAIResponse(raw []byte) (openAIResp, error) {
}
return *final, nil
}

func openAIStreamError(event openAIStreamEvent) error {
code := event.Code
message := event.Message
var errorType string
if event.Error != nil {
errorType = event.Error.Type
code = cmpOr(code, event.Error.Code)
message = cmpOr(message, event.Error.Message)
}

var details []string
if errorType != "" {
details = append(details, "type: "+errorType)
}
if code != "" {
details = append(details, "code: "+code)
}
if len(details) == 0 && message == "" {
return fmt.Errorf("openai stream error: provider returned no error details")
}
if len(details) == 0 {
return fmt.Errorf("openai stream error: %s", message)
}
if message == "" {
return fmt.Errorf("openai stream error (%s)", strings.Join(details, ", "))
}
return fmt.Errorf("openai stream error (%s): %s", strings.Join(details, ", "), message)
}
133 changes: 133 additions & 0 deletions meat/openai_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -4,13 +4,30 @@ import (
"context"
"encoding/json"
"fmt"
"io"
"net/http"
"net/http/httptest"
"strings"
"sync/atomic"
"testing"
)

type roundTripFunc func(*http.Request) (*http.Response, error)

func (f roundTripFunc) RoundTrip(r *http.Request) (*http.Response, error) {
return f(r)
}

func openAIStreamClient(body string) *http.Client {
return &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) {
return &http.Response{
StatusCode: http.StatusOK,
Body: io.NopCloser(strings.NewReader(body)),
Header: make(http.Header),
}, nil
})}
}

func writeOpenAIEvent(w http.ResponseWriter, event any) {
raw, err := json.Marshal(event)
if err != nil {
Expand Down Expand Up @@ -194,6 +211,112 @@ func TestOpenAIGenerate_StreamingText(t *testing.T) {
}
}

func TestOpenAIGenerate_ReasoningFieldsMatchModelFamily(t *testing.T) {
for _, tt := range []struct {
model string
wantReasoning bool
}{
{model: "gpt-5.6-sol", wantReasoning: true},
{model: "gpt-4o", wantReasoning: false},
{model: "gpt-4.1", wantReasoning: false},
} {
t.Run(tt.model, func(t *testing.T) {
var request map[string]json.RawMessage
client := &http.Client{Transport: roundTripFunc(func(r *http.Request) (*http.Response, error) {
if err := json.NewDecoder(r.Body).Decode(&request); err != nil {
return nil, err
}
body := `data: {"type":"response.completed","response":{"id":"resp_model","status":"completed","output":[],"usage":{"input_tokens":1,"output_tokens":1}}}` + "\n\n"
return &http.Response{
StatusCode: http.StatusOK,
Body: io.NopCloser(strings.NewReader(body)),
Header: make(http.Header),
}, nil
})}

m := &OpenAIModel{
Model: tt.model,
BaseURL: "https://openai.example",
ReasoningEffort: "high",
HTTPC: client,
}
if _, err := m.Generate(context.Background(), "sys", []Message{{Role: RoleUser, Content: []Block{textBlock("hi")}}}, nil); err != nil {
t.Fatal(err)
}

for _, field := range []string{"model", "instructions", "input", "store", "stream", "max_output_tokens"} {
if _, ok := request[field]; !ok {
t.Errorf("common request field %q is missing", field)
}
}
_, hasReasoning := request["reasoning"]
_, hasInclude := request["include"]
if hasReasoning != tt.wantReasoning || hasInclude != tt.wantReasoning {
t.Fatalf("reasoning/include presence = %v/%v, want %v/%v; request = %s", hasReasoning, hasInclude, tt.wantReasoning, tt.wantReasoning, request)
}
if tt.wantReasoning {
if got := string(request["reasoning"]); got != `{"effort":"high"}` {
t.Errorf("reasoning = %s, want configured effort", got)
}
if got := string(request["include"]); got != `["reasoning.encrypted_content"]` {
t.Errorf("include = %s", got)
}
}
})
}
}

func TestDecodeOpenAIResponse_StreamErrorDetails(t *testing.T) {
for _, tt := range []struct {
name string
event map[string]any
want []string
}{
{
name: "direct",
event: map[string]any{
"type": "error",
"code": "unsupported_parameter",
"message": "reasoning.effort is not supported",
},
want: []string{"unsupported_parameter", "reasoning.effort is not supported"},
},
{
name: "nested",
event: map[string]any{
"type": "error",
"error": map[string]any{
"type": "invalid_request_error",
"code": "model_not_found",
"message": "The requested model does not exist",
},
},
want: []string{"invalid_request_error", "model_not_found", "The requested model does not exist"},
},
{
name: "malformed",
event: map[string]any{"type": "error"},
want: []string{"openai stream error", "no error details"},
},
} {
t.Run(tt.name, func(t *testing.T) {
raw, err := json.Marshal(tt.event)
if err != nil {
t.Fatal(err)
}
_, err = decodeOpenAIResponse([]byte("data: " + string(raw) + "\n\n"))
if err == nil {
t.Fatal("error = nil, want stream error")
}
for _, want := range tt.want {
if !strings.Contains(err.Error(), want) {
t.Errorf("error = %q, want it to contain %q", err, want)
}
}
})
}
}

func TestOpenAIGenerate_IncompleteIsError(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("content-type", "text/event-stream")
Expand All @@ -216,6 +339,16 @@ func TestOpenAIGenerate_IncompleteIsError(t *testing.T) {
}
}

func TestOpenAIGenerate_FailedResponseIsError(t *testing.T) {
body := `data: {"type":"response.failed","response":{"id":"resp_failed","status":"failed","output":[],"error":{"type":"invalid_request_error","code":"bad_request","message":"request failed"}}}` + "\n\n"
m := &OpenAIModel{BaseURL: "https://openai.example", HTTPC: openAIStreamClient(body)}

_, err := m.Generate(context.Background(), "sys", []Message{{Role: RoleUser, Content: []Block{textBlock("hi")}}}, nil)
if err == nil || !strings.Contains(err.Error(), "bad_request") || !strings.Contains(err.Error(), "request failed") {
t.Fatalf("error = %v, want failed response details", err)
}
}

func TestOpenAIResponsesURL(t *testing.T) {
for base, want := range map[string]string{
"https://api.openai.com": "https://api.openai.com/v1/responses",
Expand Down