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
16 changes: 11 additions & 5 deletions providers/anthropic/anthropic.go
Original file line number Diff line number Diff line change
Expand Up @@ -1225,8 +1225,10 @@ func decodeToolCallInputMap(toolCall fantasy.ToolCallPart) (map[string]any, *fan
if strings.TrimSpace(toolCall.Input) == "" {
return map[string]any{}, nil
}
var inputMap map[string]any
if err := json.Unmarshal([]byte(toolCall.Input), &inputMap); err != nil {
// Keep each value as raw JSON. Generic decoding turns large integers
// into float64, and the SDK encodes json.Number as a string.
var rawMap map[string]json.RawMessage
if err := json.Unmarshal([]byte(toolCall.Input), &rawMap); err != nil {
return map[string]any{}, &fantasy.CallWarning{
Type: fantasy.CallWarningTypeOther,
Message: fmt.Sprintf(
Expand All @@ -1235,8 +1237,9 @@ func decodeToolCallInputMap(toolCall fantasy.ToolCallPart) (map[string]any, *fan
),
}
}
if inputMap == nil {
return map[string]any{}, nil
inputMap := make(map[string]any, len(rawMap))
for key, value := range rawMap {
inputMap[key] = value
}
return inputMap, nil
}
Expand All @@ -1248,7 +1251,7 @@ func decodeToolCallInputAny(toolCall fantasy.ToolCallPart) (any, *fantasy.CallWa
if strings.TrimSpace(toolCall.Input) == "" {
return nil, nil
}
var inputAny any
var inputAny json.RawMessage
if err := json.Unmarshal([]byte(toolCall.Input), &inputAny); err != nil {
return nil, &fantasy.CallWarning{
Type: fantasy.CallWarningTypeOther,
Expand All @@ -1258,6 +1261,9 @@ func decodeToolCallInputAny(toolCall fantasy.ToolCallPart) (any, *fantasy.CallWa
),
}
}
if string(inputAny) == "null" {
return nil, nil
}
return inputAny, nil
}

Expand Down
27 changes: 27 additions & 0 deletions providers/anthropic/anthropic_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,33 @@ var noopComputerRun = func(_ context.Context, _ fantasy.ToolCall) (fantasy.ToolR
return fantasy.ToolResponse{}, nil
}

func TestDecodeToolCallInputPreservesNumbers(t *testing.T) {
t.Parallel()
const input = `{"a":2,"b":-3.5,"value":9007199254740993,"nested":{"n":1e3}}`
toolCall := fantasy.ToolCallPart{ToolCallID: "call_1", ToolName: "echo", Input: input}

// Marshal through the SDK params: the SDK encoder, not encoding/json,
// decides how each value is written to the request body.
inputMap, warning := decodeToolCallInputMap(toolCall)
require.Nil(t, warning)
raw, err := json.Marshal(anthropic.ToolUseBlockParam{ID: "call_1", Name: "echo", Input: inputMap})
require.NoError(t, err)
var block struct {
Input json.RawMessage `json:"input"`
}
require.NoError(t, json.Unmarshal(raw, &block))
require.JSONEq(t, input, string(block.Input))
require.Contains(t, string(block.Input), "9007199254740993")

inputAny, warning := decodeToolCallInputAny(toolCall)
require.Nil(t, warning)
raw, err = json.Marshal(anthropic.ServerToolUseBlockParam{ID: "call_1", Name: "web_search", Input: inputAny})
require.NoError(t, err)
require.NoError(t, json.Unmarshal(raw, &block))
require.JSONEq(t, input, string(block.Input))
require.Contains(t, string(block.Input), "9007199254740993")
}

func TestToPrompt_DropsEmptyMessages(t *testing.T) {
t.Parallel()

Expand Down
18 changes: 18 additions & 0 deletions providers/openai/language_model.go
Original file line number Diff line number Diff line change
Expand Up @@ -37,11 +37,21 @@ type languageModel struct {
streamProviderMetadataFunc LanguageModelStreamProviderMetadataFunc
headerFunc LanguageModelHeaderFunc
toPromptFunc LanguageModelToPromptFunc
requireFinishReason bool
}

// LanguageModelOption is a function that configures a languageModel.
type LanguageModelOption = func(*languageModel)

// WithLanguageModelRequireFinishReason rejects Chat Completions streams that end
// without a finish_reason, even when the tool arguments form valid JSON.
// The default retains inference of tool-call completion for compatible providers.
func WithLanguageModelRequireFinishReason() LanguageModelOption {
return func(l *languageModel) {
l.requireFinishReason = true
}
}

// WithLanguageModelPrepareCallFunc sets the prepare call function for the language model.
func WithLanguageModelPrepareCallFunc(fn LanguageModelPrepareCallFunc) LanguageModelOption {
return func(l *languageModel) {
Expand Down Expand Up @@ -628,6 +638,14 @@ func (o languageModel) Stream(ctx context.Context, call fantasy.Call) (fantasy.S
// Emitting partial tool calls causes agents to dispatch them with
// invalid arguments before seeing the terminal reason.
mappedFinishReason := o.mapFinishReasonFunc(finishReason)
if o.requireFinishReason && finishReason == "" {
err := ctx.Err()
if err == nil {
err = fantasy.NewIncompleteStreamError()
}
yield(fantasy.StreamPart{Type: fantasy.StreamPartTypeError, Error: err})
return
}

// "Tool calls were seen" is not proof of a complete turn. Infer a
// tool-call turn only when the upstream said tool_calls/function_call
Expand Down
15 changes: 7 additions & 8 deletions providers/openai/openai_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -4661,6 +4661,7 @@ func TestResponsesToPrompt_ReasoningWithStore(t *testing.T) {
ItemID: reasoningItemID,
EncryptedContent: &encryptedContent,
Summary: []string{},
Finalized: true,
},
},
}
Expand Down Expand Up @@ -4704,19 +4705,17 @@ func TestResponsesToPrompt_ReasoningWithStore(t *testing.T) {
}
})

t.Run("store false skips reasoning", func(t *testing.T) {
t.Run("store false replays encrypted reasoning", func(t *testing.T) {
t.Parallel()

input, warnings := toResponsesPrompt(prompt, "system", false)
require.Empty(t, warnings)

// With store=false: user, assistant text, follow-up user.
require.Len(t, input, 3)

for _, item := range input {
require.Nil(t, item.OfReasoning,
"reasoning items must not appear when store=false")
}
require.Len(t, input, 4)
require.NotNil(t, input[1].OfReasoning)
require.Equal(t, reasoningItemID, input[1].OfReasoning.ID)
require.Equal(t, encryptedContent, input[1].OfReasoning.EncryptedContent.Value)
require.Empty(t, input[1].OfReasoning.Summary)
})
}

Expand Down
73 changes: 73 additions & 0 deletions providers/openai/responses_errors.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,73 @@
package openai

import (
"encoding/json"
"errors"
"fmt"
"net/http"

"charm.land/fantasy"
"github.com/charmbracelet/openai-go"
"github.com/charmbracelet/openai-go/packages/ssestream"
)

// ResponsesError keeps the provider error fields for errors.As.
// The embedded ProviderError contains HTTP details and is not serialized.
type ResponsesError struct {
Code string `json:"code"`
Type string `json:"type"`
*fantasy.ProviderError `json:"-"`
}

// Error returns the message and code, without request data.
func (e *ResponsesError) Error() string {
if e.Code == "" {
return e.ProviderError.Error()
}
return fmt.Sprintf("%s (code: %s)", e.ProviderError.Error(), e.Code)
}

// Unwrap lets callers inspect the provider and SDK errors.
func (e *ResponsesError) Unwrap() error { return e.ProviderError }

func responsesStreamFailureError(title, raw string, response *http.Response) *ResponsesError {
var payload struct {
Code string `json:"code"`
Type string `json:"type"`
Message string `json:"message"`
StatusCode int `json:"status_code"`
Error json.RawMessage `json:"error"`
}
_ = json.Unmarshal([]byte(raw), &payload)
if len(payload.Error) > 0 {
nested := payload.Error
_ = json.Unmarshal(nested, &payload)
}
status := payload.StatusCode
if response != nil {
status = response.StatusCode
}
provider := &fantasy.ProviderError{Title: title, Message: payload.Message, StatusCode: status}
parseContextTooLargeError(payload.Message, provider)
return &ResponsesError{Code: payload.Code, Type: payload.Type, ProviderError: provider}
}

func toResponsesProviderErr(err error, response *http.Response) error {
converted := toProviderErr(err)
var provider *fantasy.ProviderError
if !errors.As(converted, &provider) {
return converted
}
var api *openai.Error
if errors.As(err, &api) {
return &ResponsesError{Code: api.Code, Type: api.Type, ProviderError: provider}
}
var stream *ssestream.StreamError
if errors.As(err, &stream) {
typed := responsesStreamFailureError("stream error", string(stream.Event.Data), response)
provider.StatusCode = typed.StatusCode
typed.ProviderError = provider
return typed
}
return converted
}
132 changes: 132 additions & 0 deletions providers/openai/responses_errors_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,132 @@
package openai

import (
"encoding/json"
"errors"
"net/http"
"net/http/httptest"
"testing"

"charm.land/fantasy"
"github.com/stretchr/testify/require"
)

func TestResponsesTypedStreamErrors(t *testing.T) {
t.Parallel()
for _, event := range []struct{ name, data, code, kind, message string }{
{"response.failed", `{"type":"response.failed","response":{"id":"resp_1","status":"failed","error":{"code":"server_error","type":"api_error","message":"boom"}}}`, "server_error", "api_error", "boom"},
{"error", `{"type":"error","code":"invalid_prompt","message":"bad prompt"}`, "invalid_prompt", "error", "bad prompt"},
{"error", `{"error":{"type":"invalid_request_error","code":"bad_argument","message":"bad argument"}}`, "bad_argument", "invalid_request_error", "bad argument"},
{"error", `{"error":{"type":"invalid_request_error","code":"bad_argument","message":"bad argument","status_code":429}}`, "bad_argument", "invalid_request_error", "bad argument"},
} {
t.Run(event.data, func(t *testing.T) {
t.Parallel()
server := newStreamingMockServer()
defer server.close()
server.chunks = []string{
responsesSSEEvent("response.output_text.delta", `{"type":"response.output_text.delta","item_id":"msg_1","delta":"{\"answer\":\"yes\"}"}`),
responsesSSEEvent(event.name, event.data),
}
model := newResponsesProvider(t, server.server.URL)
stream, err := model.Stream(t.Context(), fantasy.Call{Prompt: testPrompt})
require.NoError(t, err)
var failure error
for part := range stream {
if part.Error != nil {
failure = part.Error
}
}
require.Error(t, failure)
var provider *fantasy.ProviderError
require.True(t, errors.As(failure, &provider))
require.Equal(t, 200, provider.StatusCode)
var typed *ResponsesError
require.ErrorAs(t, failure, &typed)
require.Equal(t, event.code, typed.Code)
require.Equal(t, event.kind, typed.Type)
require.Equal(t, event.message, typed.Message)
schema := fantasy.Schema{Type: "object", Properties: map[string]*fantasy.Schema{"answer": {Type: "string"}}, Required: []string{"answer"}}
objects, err := model.StreamObject(t.Context(), fantasy.ObjectCall{Prompt: testPrompt, Schema: schema})
require.NoError(t, err)
failure = nil
for part := range objects {
if part.Error != nil {
failure = part.Error
}
}
require.True(t, errors.As(failure, &provider))
require.Equal(t, 200, provider.StatusCode)
require.NotContains(t, failure.Error(), "test-api-key")
require.NotContains(t, failure.Error(), "Authorization")
})
}
}

func TestResponsesHTTPErrorFields(t *testing.T) {
t.Parallel()
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(http.StatusBadRequest)
_, _ = w.Write([]byte(`{"error":{"code":"invalid_prompt","type":"invalid_request_error","message":"bad prompt"}}`))
}))
defer server.Close()
model := newResponsesProvider(t, server.URL)
schema := fantasy.Schema{Type: "object", Properties: map[string]*fantasy.Schema{"answer": {Type: "string"}}}
_, err := model.Generate(t.Context(), fantasy.Call{Prompt: testPrompt})
checkResponsesHTTPError(t, err)
_, err = model.GenerateObject(t.Context(), fantasy.ObjectCall{Prompt: testPrompt, Schema: schema})
checkResponsesHTTPError(t, err)
stream, err := model.Stream(t.Context(), fantasy.Call{Prompt: testPrompt})
require.NoError(t, err)
seen := false
for part := range stream {
if part.Error != nil {
seen = true
checkResponsesHTTPError(t, part.Error)
}
}
require.True(t, seen)
objects, err := model.StreamObject(t.Context(), fantasy.ObjectCall{Prompt: testPrompt, Schema: schema})
require.NoError(t, err)
seen = false
for part := range objects {
if part.Error != nil {
seen = true
checkResponsesHTTPError(t, part.Error)
}
}
require.True(t, seen)
}

func checkResponsesHTTPError(t *testing.T, err error) {
t.Helper()
var typed *ResponsesError
require.ErrorAs(t, err, &typed)
require.Equal(t, "invalid_prompt", typed.Code)
require.Equal(t, "invalid_request_error", typed.Type)
require.Equal(t, "bad prompt", typed.Message)
require.Equal(t, http.StatusBadRequest, typed.StatusCode)
require.NotContains(t, err.Error(), "test-api-key")
require.NotContains(t, err.Error(), "Authorization")
// The existing SDK error keeps a request dump. Do not serialize it.
require.NotEmpty(t, typed.RequestBody)
data, marshalErr := json.Marshal(typed)
require.NoError(t, marshalErr)
require.NotContains(t, string(data), "test-api-key")
require.NotContains(t, string(data), "RequestBody")
}

func TestResponsesGenerateBodyErrorFields(t *testing.T) {
t.Parallel()
server := newMockServer()
defer server.close()
server.response = map[string]any{"id": "resp_1", "status": "failed", "error": map[string]any{"code": "server_error", "type": "api_error", "message": "boom"}}
model := newResponsesProvider(t, server.server.URL)
_, err := model.Generate(t.Context(), fantasy.Call{Prompt: testPrompt})
var typed *ResponsesError
require.ErrorAs(t, err, &typed)
require.Equal(t, "server_error", typed.Code)
require.Equal(t, "api_error", typed.Type)
require.Equal(t, "boom", typed.Message)
require.Equal(t, http.StatusOK, typed.StatusCode)
}
Loading
Loading