mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
fix: prevent malformed OpenAI Responses continuations (#24725)
> Worked on by Mux on Mike's behalf. ## Summary - Disable OpenAI Responses `previous_response_id` chain mode when the prior assistant response has unresolved local tool calls, so the next request can include paired tool outputs instead of sending an incomplete continuation. - Update the fantasy pin to a Responses replay fix that preserves stored reasoning references, only replays web search references when paired with reasoning, and validates local function-call output pairing before send. - Add fake OpenAI Responses input validation for the two production 400 shapes and integration coverage for full-history reasoning plus web search replay. - Add sanitized diagnostics for the OpenAI Responses continuity errors. ## Tests - `go test ./providers/openai -run 'TestResponsesToPrompt_(ReasoningWithStore|ReasoningWithWebSearchCombined|WebSearchRequiresReasoningReference|ReasoningWithFunctionCallCombined|WebSearchProviderExecutedToolResults)|TestPrepareParams_(SkipsProviderExecutedToolReferences|ValidatesFunctionCallOutputPairing)|TestValidateResponsesInput_WebSearchReferenceRequiresReasoning' -count=1` - `go test ./providers/openai -count=1` - `GOWORK=off go test ./coderd/x/chatd/chattest -run TestValidateResponsesAPIInput -count=1` - `GOWORK=off go test ./coderd/x/chatd -run 'TestOpenAIResponses(NoStaleWebSearchReplay|FullReplayPairsReasoningAndWebSearch|ChainModeSkipsWhenLocalCallPending|ChainModeStillFiresForProviderExecutedOnly)$|TestResolveChainMode_' -count=1` - `GOWORK=off go test ./coderd/x/chatd/chatprompt -run 'TestInjectMissingToolResults_' -count=1` - `GOWORK=off go test ./coderd/x/chatd/chaterror -run TestClassify_OpenAIResponsesAPIDiagnostics -count=1` - `GOWORK=off go test ./coderd/x/chatd/... -count=1` - `git diff --check` - `git commit` pre-commit hook
This commit is contained in:
+81
-5
@@ -3260,6 +3260,9 @@ type chainModeInfo struct {
|
||||
// contributingTrailingUserCount counts the trailing user
|
||||
// messages that materially change the provider input.
|
||||
contributingTrailingUserCount int
|
||||
// hasUnresolvedLocalToolCalls is true when previousResponseID
|
||||
// points at an assistant message with pending local tool calls.
|
||||
hasUnresolvedLocalToolCalls bool
|
||||
}
|
||||
|
||||
func userMessageContributesToChainMode(msg database.ChatMessage) bool {
|
||||
@@ -3286,6 +3289,77 @@ func userMessageContributesToChainMode(msg database.ChatMessage) bool {
|
||||
return false
|
||||
}
|
||||
|
||||
// assistantHasUnresolvedLocalToolCalls reports whether the assistant message
|
||||
// at assistantIdx contains local tool calls that lack matching tool results.
|
||||
// It returns true when content parsing fails because full-history replay is
|
||||
// safer than chaining from state that cannot be inspected.
|
||||
func assistantHasUnresolvedLocalToolCalls(
|
||||
messages []database.ChatMessage,
|
||||
assistantIdx int,
|
||||
) bool {
|
||||
if assistantIdx < 0 || assistantIdx >= len(messages) {
|
||||
return false
|
||||
}
|
||||
|
||||
parts, err := chatprompt.ParseContent(messages[assistantIdx])
|
||||
if err != nil {
|
||||
// Use full replay when persisted assistant content cannot be parsed.
|
||||
return true
|
||||
}
|
||||
|
||||
localCallIDs := make(map[string]struct{})
|
||||
for _, part := range parts {
|
||||
if part.Type != codersdk.ChatMessagePartTypeToolCall ||
|
||||
part.ProviderExecuted {
|
||||
continue
|
||||
}
|
||||
localCallIDs[part.ToolCallID] = struct{}{}
|
||||
}
|
||||
if len(localCallIDs) == 0 {
|
||||
return false
|
||||
}
|
||||
|
||||
resolvedCallIDs := make(map[string]struct{})
|
||||
for i := assistantIdx + 1; i < len(messages); i++ {
|
||||
if messages[i].Role != database.ChatMessageRoleTool {
|
||||
break
|
||||
}
|
||||
parts, err := chatprompt.ParseContent(messages[i])
|
||||
if err != nil {
|
||||
// Use full replay when persisted tool content cannot be parsed.
|
||||
return true
|
||||
}
|
||||
for _, part := range parts {
|
||||
if part.Type != codersdk.ChatMessagePartTypeToolResult {
|
||||
continue
|
||||
}
|
||||
if _, ok := localCallIDs[part.ToolCallID]; ok {
|
||||
resolvedCallIDs[part.ToolCallID] = struct{}{}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return len(resolvedCallIDs) != len(localCallIDs)
|
||||
}
|
||||
|
||||
// shouldActivateChainMode reports whether a follow-up turn can use
|
||||
// previous_response_id instead of replaying history. It requires store=true,
|
||||
// a matching model config, meaningful trailing user input, non-plan mode, and
|
||||
// complete local tool state so the provider has all required outputs.
|
||||
func shouldActivateChainMode(
|
||||
providerOptions fantasy.ProviderOptions,
|
||||
info chainModeInfo,
|
||||
modelConfigID uuid.UUID,
|
||||
isPlanModeTurn bool,
|
||||
) bool {
|
||||
return chatprovider.IsResponsesStoreEnabled(providerOptions) &&
|
||||
info.previousResponseID != "" &&
|
||||
info.contributingTrailingUserCount > 0 &&
|
||||
info.modelConfigID == modelConfigID &&
|
||||
!isPlanModeTurn &&
|
||||
!info.hasUnresolvedLocalToolCalls
|
||||
}
|
||||
|
||||
// resolveChainMode scans DB messages from the end to count trailing user
|
||||
// messages for the current turn and detect whether the immediately
|
||||
// preceding assistant/tool block can chain from a provider response ID.
|
||||
@@ -3310,6 +3384,7 @@ func resolveChainMode(messages []database.ChatMessage) chainModeInfo {
|
||||
if messages[i].ModelConfigID.Valid {
|
||||
info.modelConfigID = messages[i].ModelConfigID.UUID
|
||||
}
|
||||
info.hasUnresolvedLocalToolCalls = assistantHasUnresolvedLocalToolCalls(messages, i)
|
||||
return info
|
||||
}
|
||||
return info
|
||||
@@ -6680,11 +6755,12 @@ func (p *Server) runChat(
|
||||
// we set previous_response_id and send only system instructions
|
||||
// plus the new user input, avoiding redundant replay of prior
|
||||
// assistant and tool messages that the provider already has.
|
||||
chainModeActive := chatprovider.IsResponsesStoreEnabled(providerOptions) &&
|
||||
chainInfo.previousResponseID != "" &&
|
||||
chainInfo.contributingTrailingUserCount > 0 &&
|
||||
chainInfo.modelConfigID == modelConfig.ID &&
|
||||
!isPlanModeTurn
|
||||
chainModeActive := shouldActivateChainMode(
|
||||
providerOptions,
|
||||
chainInfo,
|
||||
modelConfig.ID,
|
||||
isPlanModeTurn,
|
||||
)
|
||||
if chainModeActive {
|
||||
providerOptions = chatprovider.CloneWithPreviousResponseID(
|
||||
providerOptions,
|
||||
|
||||
@@ -9,6 +9,7 @@ import (
|
||||
"time"
|
||||
|
||||
"charm.land/fantasy"
|
||||
fantasyopenai "charm.land/fantasy/providers/openai"
|
||||
"github.com/google/uuid"
|
||||
"github.com/sqlc-dev/pqtype"
|
||||
"github.com/stretchr/testify/require"
|
||||
@@ -23,6 +24,7 @@ import (
|
||||
coderdpubsub "github.com/coder/coder/v2/coderd/pubsub"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chaterror"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatloop"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatprompt"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatprovider"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chattest"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chattool"
|
||||
@@ -2913,6 +2915,286 @@ func TestResolveChainModeIgnoresSkillOnlySentinelMessages(t *testing.T) {
|
||||
require.Equal(t, 1, got.contributingTrailingUserCount)
|
||||
}
|
||||
|
||||
func TestResolveChainMode_BlocksOnUnresolvedLocalToolCall(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
modelConfigID := uuid.New()
|
||||
toolCall := codersdk.ChatMessageToolCall(
|
||||
"call-local",
|
||||
"read_file",
|
||||
json.RawMessage(`{"path":"main.go"}`),
|
||||
)
|
||||
|
||||
chainInfo := resolveChainMode([]database.ChatMessage{
|
||||
chainModeSystemMessage(),
|
||||
chainModeUserMessage("prior user message"),
|
||||
chainModeAssistantMessage(modelConfigID, []codersdk.ChatMessagePart{toolCall}),
|
||||
chainModeUserMessage("latest user message"),
|
||||
})
|
||||
|
||||
require.Equal(t, "resp-123", chainInfo.previousResponseID)
|
||||
require.True(t, chainInfo.hasUnresolvedLocalToolCalls)
|
||||
require.False(t, shouldActivateChainMode(
|
||||
chainModeProviderOptions(),
|
||||
chainInfo,
|
||||
modelConfigID,
|
||||
false,
|
||||
))
|
||||
}
|
||||
|
||||
func TestResolveChainMode_BlocksWhenAssistantContentCannotParse(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
modelConfigID := uuid.New()
|
||||
chainInfo := resolveChainMode([]database.ChatMessage{
|
||||
chainModeSystemMessage(),
|
||||
chainModeUserMessage("prior user message"),
|
||||
chainModeCorruptAssistantMessage(modelConfigID),
|
||||
chainModeUserMessage("latest user message"),
|
||||
})
|
||||
|
||||
require.Equal(t, "resp-123", chainInfo.previousResponseID)
|
||||
require.True(t, chainInfo.hasUnresolvedLocalToolCalls)
|
||||
require.False(t, shouldActivateChainMode(
|
||||
chainModeProviderOptions(),
|
||||
chainInfo,
|
||||
modelConfigID,
|
||||
false,
|
||||
))
|
||||
}
|
||||
|
||||
func TestResolveChainMode_BlocksWhenToolContentCannotParse(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
modelConfigID := uuid.New()
|
||||
toolCall := codersdk.ChatMessageToolCall(
|
||||
"call-local",
|
||||
"read_file",
|
||||
json.RawMessage(`{"path":"main.go"}`),
|
||||
)
|
||||
|
||||
chainInfo := resolveChainMode([]database.ChatMessage{
|
||||
chainModeSystemMessage(),
|
||||
chainModeUserMessage("prior user message"),
|
||||
chainModeAssistantMessage(modelConfigID, []codersdk.ChatMessagePart{toolCall}),
|
||||
chainModeCorruptToolMessage(),
|
||||
chainModeUserMessage("latest user message"),
|
||||
})
|
||||
|
||||
require.Equal(t, "resp-123", chainInfo.previousResponseID)
|
||||
require.True(t, chainInfo.hasUnresolvedLocalToolCalls)
|
||||
require.False(t, shouldActivateChainMode(
|
||||
chainModeProviderOptions(),
|
||||
chainInfo,
|
||||
modelConfigID,
|
||||
false,
|
||||
))
|
||||
}
|
||||
|
||||
func TestResolveChainMode_AllowsProviderExecutedOnly(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
modelConfigID := uuid.New()
|
||||
toolCall := codersdk.ChatMessageToolCall(
|
||||
"call-web-search",
|
||||
"web_search",
|
||||
json.RawMessage(`{"query":"coder docs"}`),
|
||||
)
|
||||
toolCall.ProviderExecuted = true
|
||||
|
||||
chainInfo := resolveChainMode([]database.ChatMessage{
|
||||
chainModeSystemMessage(),
|
||||
chainModeUserMessage("prior user message"),
|
||||
chainModeAssistantMessage(modelConfigID, []codersdk.ChatMessagePart{toolCall}),
|
||||
chainModeUserMessage("latest user message"),
|
||||
})
|
||||
|
||||
require.Equal(t, "resp-123", chainInfo.previousResponseID)
|
||||
require.False(t, chainInfo.hasUnresolvedLocalToolCalls)
|
||||
require.True(t, shouldActivateChainMode(
|
||||
chainModeProviderOptions(),
|
||||
chainInfo,
|
||||
modelConfigID,
|
||||
false,
|
||||
))
|
||||
}
|
||||
|
||||
func TestResolveChainMode_BlocksOnMixedProviderExecutedAndUnresolvedLocalCall(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
modelConfigID := uuid.New()
|
||||
providerCall := codersdk.ChatMessageToolCall(
|
||||
"call-web-search",
|
||||
"web_search",
|
||||
json.RawMessage(`{"query":"coder docs"}`),
|
||||
)
|
||||
providerCall.ProviderExecuted = true
|
||||
localCall := codersdk.ChatMessageToolCall(
|
||||
"call-local",
|
||||
"read_file",
|
||||
json.RawMessage(`{"path":"main.go"}`),
|
||||
)
|
||||
|
||||
chainInfo := resolveChainMode([]database.ChatMessage{
|
||||
chainModeSystemMessage(),
|
||||
chainModeUserMessage("prior user message"),
|
||||
chainModeAssistantMessage(
|
||||
modelConfigID,
|
||||
[]codersdk.ChatMessagePart{providerCall, localCall},
|
||||
),
|
||||
chainModeUserMessage("latest user message"),
|
||||
})
|
||||
|
||||
require.Equal(t, "resp-123", chainInfo.previousResponseID)
|
||||
require.True(t, chainInfo.hasUnresolvedLocalToolCalls)
|
||||
require.False(t, shouldActivateChainMode(
|
||||
chainModeProviderOptions(),
|
||||
chainInfo,
|
||||
modelConfigID,
|
||||
false,
|
||||
))
|
||||
}
|
||||
|
||||
func TestResolveChainMode_AllowsResolvedLocalCall(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
modelConfigID := uuid.New()
|
||||
toolCall := codersdk.ChatMessageToolCall(
|
||||
"call-local",
|
||||
"read_file",
|
||||
json.RawMessage(`{"path":"main.go"}`),
|
||||
)
|
||||
toolResult := codersdk.ChatMessageToolResult(
|
||||
"call-local",
|
||||
"read_file",
|
||||
json.RawMessage(`{"ok":true}`),
|
||||
false,
|
||||
false,
|
||||
)
|
||||
|
||||
chainInfo := resolveChainMode([]database.ChatMessage{
|
||||
chainModeSystemMessage(),
|
||||
chainModeUserMessage("prior user message"),
|
||||
chainModeAssistantMessage(modelConfigID, []codersdk.ChatMessagePart{toolCall}),
|
||||
chainModeToolMessage([]codersdk.ChatMessagePart{toolResult}),
|
||||
chainModeUserMessage("latest user message"),
|
||||
})
|
||||
|
||||
require.Equal(t, "resp-123", chainInfo.previousResponseID)
|
||||
require.False(t, chainInfo.hasUnresolvedLocalToolCalls)
|
||||
require.True(t, shouldActivateChainMode(
|
||||
chainModeProviderOptions(),
|
||||
chainInfo,
|
||||
modelConfigID,
|
||||
false,
|
||||
))
|
||||
}
|
||||
|
||||
func TestResolveChainMode_BlocksOnMixedResolvedAndUnresolved(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
modelConfigID := uuid.New()
|
||||
firstCall := codersdk.ChatMessageToolCall(
|
||||
"call-first",
|
||||
"read_file",
|
||||
json.RawMessage(`{"path":"main.go"}`),
|
||||
)
|
||||
secondCall := codersdk.ChatMessageToolCall(
|
||||
"call-second",
|
||||
"read_file",
|
||||
json.RawMessage(`{"path":"README.md"}`),
|
||||
)
|
||||
toolResult := codersdk.ChatMessageToolResult(
|
||||
"call-first",
|
||||
"read_file",
|
||||
json.RawMessage(`{"ok":true}`),
|
||||
false,
|
||||
false,
|
||||
)
|
||||
|
||||
chainInfo := resolveChainMode([]database.ChatMessage{
|
||||
chainModeSystemMessage(),
|
||||
chainModeUserMessage("prior user message"),
|
||||
chainModeAssistantMessage(
|
||||
modelConfigID,
|
||||
[]codersdk.ChatMessagePart{firstCall, secondCall},
|
||||
),
|
||||
chainModeToolMessage([]codersdk.ChatMessagePart{toolResult}),
|
||||
chainModeUserMessage("latest user message"),
|
||||
})
|
||||
|
||||
require.Equal(t, "resp-123", chainInfo.previousResponseID)
|
||||
require.True(t, chainInfo.hasUnresolvedLocalToolCalls)
|
||||
require.False(t, shouldActivateChainMode(
|
||||
chainModeProviderOptions(),
|
||||
chainInfo,
|
||||
modelConfigID,
|
||||
false,
|
||||
))
|
||||
}
|
||||
|
||||
func chainModeProviderOptions() fantasy.ProviderOptions {
|
||||
store := true
|
||||
return fantasy.ProviderOptions{
|
||||
fantasyopenai.Name: &fantasyopenai.ResponsesProviderOptions{
|
||||
Store: &store,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func chainModeSystemMessage() database.ChatMessage {
|
||||
return database.ChatMessage{Role: database.ChatMessageRoleSystem}
|
||||
}
|
||||
|
||||
func chainModeUserMessage(text string) database.ChatMessage {
|
||||
msg := chatMessageWithParts([]codersdk.ChatMessagePart{
|
||||
codersdk.ChatMessageText(text),
|
||||
})
|
||||
msg.Role = database.ChatMessageRoleUser
|
||||
return msg
|
||||
}
|
||||
|
||||
func chainModeAssistantMessage(
|
||||
modelConfigID uuid.UUID,
|
||||
parts []codersdk.ChatMessagePart,
|
||||
) database.ChatMessage {
|
||||
msg := chatMessageWithParts(parts)
|
||||
msg.Role = database.ChatMessageRoleAssistant
|
||||
msg.ProviderResponseID = sql.NullString{String: "resp-123", Valid: true}
|
||||
msg.ModelConfigID = uuid.NullUUID{UUID: modelConfigID, Valid: true}
|
||||
return msg
|
||||
}
|
||||
|
||||
func chainModeCorruptAssistantMessage(modelConfigID uuid.UUID) database.ChatMessage {
|
||||
return database.ChatMessage{
|
||||
Role: database.ChatMessageRoleAssistant,
|
||||
ProviderResponseID: sql.NullString{String: "resp-123", Valid: true},
|
||||
ModelConfigID: uuid.NullUUID{UUID: modelConfigID, Valid: true},
|
||||
Content: pqtype.NullRawMessage{
|
||||
RawMessage: []byte("not json"),
|
||||
Valid: true,
|
||||
},
|
||||
ContentVersion: chatprompt.CurrentContentVersion,
|
||||
}
|
||||
}
|
||||
|
||||
func chainModeCorruptToolMessage() database.ChatMessage {
|
||||
return database.ChatMessage{
|
||||
Role: database.ChatMessageRoleTool,
|
||||
Content: pqtype.NullRawMessage{
|
||||
RawMessage: []byte("not json"),
|
||||
Valid: true,
|
||||
},
|
||||
ContentVersion: chatprompt.CurrentContentVersion,
|
||||
}
|
||||
}
|
||||
|
||||
func chainModeToolMessage(parts []codersdk.ChatMessagePart) database.ChatMessage {
|
||||
msg := chatMessageWithParts(parts)
|
||||
msg.Role = database.ChatMessageRoleTool
|
||||
return msg
|
||||
}
|
||||
|
||||
func TestFilterPromptForChainModeKeepsContributingUsersAcrossSkippedSentinelTurns(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
|
||||
@@ -22,6 +22,28 @@ type ClassifiedError struct {
|
||||
RetryAfter time.Duration
|
||||
}
|
||||
|
||||
const responsesAPIDiagnosticMessage = "The chat continuation failed due to an " +
|
||||
"internal state mismatch. This is not a configuration or billing issue."
|
||||
|
||||
type responsesAPIDiagnosticMatch struct {
|
||||
pattern string
|
||||
detail string
|
||||
}
|
||||
|
||||
// responsesAPIDiagnosticMatches maps provider error fragments to safe
|
||||
// diagnostics. Details must not include provider item IDs because they are
|
||||
// returned to clients and used by operators for grepping.
|
||||
var responsesAPIDiagnosticMatches = []responsesAPIDiagnosticMatch{
|
||||
{
|
||||
pattern: "no tool output found for function call",
|
||||
detail: "OpenAI Responses API request continuity diagnostic: match=function_call_output_missing.",
|
||||
},
|
||||
{
|
||||
pattern: "was provided without its required 'reasoning' item",
|
||||
detail: "OpenAI Responses API request continuity diagnostic: match=web_search_reasoning_missing.",
|
||||
},
|
||||
}
|
||||
|
||||
// WithProvider returns a copy of the classification using an explicit
|
||||
// provider hint. Explicit provider hints are trusted over provider names
|
||||
// heuristically parsed from the error text.
|
||||
@@ -102,6 +124,17 @@ func Classify(err error) ClassifiedError {
|
||||
})
|
||||
}
|
||||
|
||||
if detail, ok := responsesAPIDiagnostic(lower, structured.detail); ok {
|
||||
return normalizeClassification(ClassifiedError{
|
||||
Message: responsesAPIDiagnosticMessage,
|
||||
Detail: detail,
|
||||
Kind: KindGeneric,
|
||||
Provider: provider,
|
||||
StatusCode: statusCode,
|
||||
RetryAfter: structured.retryAfter,
|
||||
})
|
||||
}
|
||||
|
||||
deadline := errors.Is(err, context.DeadlineExceeded) || strings.Contains(lower, "context deadline exceeded")
|
||||
overloadedMatch := statusCode == 529 || containsAny(lower, overloadedPatterns...)
|
||||
authStrong := statusCode == 401 || containsAny(lower, authStrongPatterns...)
|
||||
@@ -183,6 +216,16 @@ func Classify(err error) ClassifiedError {
|
||||
})
|
||||
}
|
||||
|
||||
func responsesAPIDiagnostic(lowerMessage, detail string) (string, bool) {
|
||||
lowerDetail := strings.ToLower(detail)
|
||||
for _, match := range responsesAPIDiagnosticMatches {
|
||||
if strings.Contains(lowerMessage, match.pattern) || strings.Contains(lowerDetail, match.pattern) {
|
||||
return match.detail, true
|
||||
}
|
||||
}
|
||||
return "", false
|
||||
}
|
||||
|
||||
func normalizeClassification(classified ClassifiedError) ClassifiedError {
|
||||
classified.Message = strings.TrimSpace(classified.Message)
|
||||
classified.Detail = normalizeClassificationDetail(classified.Detail)
|
||||
|
||||
@@ -197,6 +197,82 @@ func TestClassify(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestClassify_OpenAIResponsesAPIDiagnostics(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
err string
|
||||
responseBody string
|
||||
wantDetail string
|
||||
forbidden []string
|
||||
}{
|
||||
{
|
||||
name: "FunctionCallOutputMissing",
|
||||
err: "No Tool Output Found For Function Call call_sensitive123",
|
||||
responseBody: `{"error":{"message":"No tool output found for function call call_sensitive123"}}`,
|
||||
wantDetail: "OpenAI Responses API request continuity diagnostic: match=function_call_output_missing.",
|
||||
forbidden: []string{"call_sensitive123"},
|
||||
},
|
||||
{
|
||||
name: "WebSearchReasoningMissing",
|
||||
err: "Item 'ws_sensitive123' of type 'web_search_call' WAS PROVIDED WITHOUT ITS REQUIRED 'reasoning' item: 'rs_sensitive123'",
|
||||
responseBody: `{"error":{"message":"Item 'ws_sensitive123' of type 'web_search_call' was provided without its required 'reasoning' item: 'rs_sensitive123'"}}`,
|
||||
wantDetail: "OpenAI Responses API request continuity diagnostic: match=web_search_reasoning_missing.",
|
||||
forbidden: []string{"ws_sensitive123", "rs_sensitive123"},
|
||||
},
|
||||
}
|
||||
|
||||
assertNoLeak := func(t *testing.T, classified chaterror.ClassifiedError, forbidden []string) {
|
||||
t.Helper()
|
||||
for _, value := range forbidden {
|
||||
require.NotContains(t, classified.Message, value)
|
||||
require.NotContains(t, classified.Detail, value)
|
||||
}
|
||||
}
|
||||
|
||||
assertDirectionalMessage := func(t *testing.T, message string) {
|
||||
t.Helper()
|
||||
require.Contains(t, message, "chat continuation")
|
||||
require.Contains(t, message, "internal state mismatch")
|
||||
require.Contains(t, message, "not a configuration or billing issue")
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name+"/BareString", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
classified := chaterror.Classify(xerrors.New(tt.err))
|
||||
require.Equal(t, chaterror.KindGeneric, classified.Kind)
|
||||
require.False(t, classified.Retryable)
|
||||
require.Zero(t, classified.StatusCode)
|
||||
assertDirectionalMessage(t, classified.Message)
|
||||
require.Equal(t, tt.wantDetail, classified.Detail)
|
||||
assertNoLeak(t, classified, tt.forbidden)
|
||||
})
|
||||
|
||||
t.Run(tt.name+"/WrappedProviderError", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
classified := chaterror.Classify(xerrors.Errorf(
|
||||
"provider request failed: %w",
|
||||
testProviderError(
|
||||
"",
|
||||
400,
|
||||
nil,
|
||||
testProviderResponseDump(tt.responseBody),
|
||||
),
|
||||
))
|
||||
require.Equal(t, chaterror.KindGeneric, classified.Kind)
|
||||
require.False(t, classified.Retryable)
|
||||
require.Equal(t, 400, classified.StatusCode)
|
||||
assertDirectionalMessage(t, classified.Message)
|
||||
require.Equal(t, tt.wantDetail, classified.Detail)
|
||||
assertNoLeak(t, classified, tt.forbidden)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestClassify_PatternCoverage(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
|
||||
@@ -759,6 +759,122 @@ func TestInjectMissingToolResults_SkipsProviderExecuted(t *testing.T) {
|
||||
require.Equal(t, "toolu_local", remainingToolCalls[0].ToolCallID)
|
||||
}
|
||||
|
||||
func TestInjectMissingToolResults_SkipsProviderExecutedAndInjectsLocal(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
providerCall := codersdk.ChatMessageToolCall(
|
||||
"srvtoolu_web_search",
|
||||
"web_search",
|
||||
json.RawMessage(`{"query":"coder"}`),
|
||||
)
|
||||
providerCall.ProviderExecuted = true
|
||||
localCall := codersdk.ChatMessageToolCall(
|
||||
"toolu_read",
|
||||
"read_file",
|
||||
json.RawMessage(`{"path":"main.go"}`),
|
||||
)
|
||||
assistantContent, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{
|
||||
providerCall,
|
||||
localCall,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
prompt := convertMessagesWithoutFiles(t, []database.ChatMessage{{
|
||||
Role: database.ChatMessageRoleAssistant,
|
||||
Visibility: database.ChatMessageVisibilityBoth,
|
||||
Content: assistantContent,
|
||||
ContentVersion: chatprompt.CurrentContentVersion,
|
||||
}})
|
||||
|
||||
require.Len(t, prompt, 2, "expected assistant plus local synthetic tool result")
|
||||
require.Equal(t, fantasy.MessageRoleAssistant, prompt[0].Role)
|
||||
require.Equal(t, fantasy.MessageRoleTool, prompt[1].Role)
|
||||
|
||||
toolCalls := chatprompt.ExtractToolCalls(prompt[0].Content)
|
||||
require.Len(t, toolCalls, 2)
|
||||
require.Equal(t, "srvtoolu_web_search", toolCalls[0].ToolCallID)
|
||||
require.True(t, toolCalls[0].ProviderExecuted)
|
||||
require.Equal(t, "toolu_read", toolCalls[1].ToolCallID)
|
||||
require.False(t, toolCalls[1].ProviderExecuted)
|
||||
|
||||
require.Equal(t, []string{"toolu_read"}, extractToolResultIDs(t, prompt[1]))
|
||||
require.Len(t, prompt[1].Content, 1)
|
||||
toolResult, ok := fantasy.AsMessagePart[fantasy.ToolResultPart](prompt[1].Content[0])
|
||||
require.True(t, ok, "expected synthetic ToolResultPart")
|
||||
require.Equal(t, "toolu_read", toolResult.ToolCallID)
|
||||
require.False(t, toolResult.ProviderExecuted)
|
||||
errOutput, ok := fantasy.AsToolResultOutputType[fantasy.ToolResultOutputContentError](toolResult.Output)
|
||||
require.True(t, ok, "expected synthetic error output")
|
||||
require.ErrorContains(t, errOutput.Error, "tool call was interrupted")
|
||||
}
|
||||
|
||||
func TestInjectMissingToolResults_AdjacentAssistantsInjectLocalResults(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
assistantAContent, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{
|
||||
codersdk.ChatMessageText("assistant a"),
|
||||
codersdk.ChatMessageToolCall(
|
||||
"toolu_a",
|
||||
"read_file",
|
||||
json.RawMessage(`{"path":"a.go"}`),
|
||||
),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assistantBContent, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{
|
||||
codersdk.ChatMessageText("assistant b"),
|
||||
codersdk.ChatMessageToolCall(
|
||||
"toolu_b",
|
||||
"read_file",
|
||||
json.RawMessage(`{"path":"b.go"}`),
|
||||
),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
userContent, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{
|
||||
codersdk.ChatMessageText("next user message"),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
prompt := convertMessagesWithoutFiles(t, []database.ChatMessage{
|
||||
{
|
||||
Role: database.ChatMessageRoleAssistant,
|
||||
Visibility: database.ChatMessageVisibilityBoth,
|
||||
Content: assistantAContent,
|
||||
ContentVersion: chatprompt.CurrentContentVersion,
|
||||
},
|
||||
{
|
||||
Role: database.ChatMessageRoleAssistant,
|
||||
Visibility: database.ChatMessageVisibilityBoth,
|
||||
Content: assistantBContent,
|
||||
ContentVersion: chatprompt.CurrentContentVersion,
|
||||
},
|
||||
{
|
||||
Role: database.ChatMessageRoleUser,
|
||||
Visibility: database.ChatMessageVisibilityBoth,
|
||||
Content: userContent,
|
||||
ContentVersion: chatprompt.CurrentContentVersion,
|
||||
},
|
||||
})
|
||||
|
||||
require.Len(t, prompt, 5)
|
||||
require.Equal(t, fantasy.MessageRoleAssistant, prompt[0].Role)
|
||||
require.Equal(t, fantasy.MessageRoleTool, prompt[1].Role)
|
||||
require.Equal(t, fantasy.MessageRoleAssistant, prompt[2].Role)
|
||||
require.Equal(t, fantasy.MessageRoleTool, prompt[3].Role)
|
||||
require.Equal(t, fantasy.MessageRoleUser, prompt[4].Role)
|
||||
require.Equal(t, []string{"toolu_a"}, extractToolResultIDs(t, prompt[1]))
|
||||
require.Equal(t, []string{"toolu_b"}, extractToolResultIDs(t, prompt[3]))
|
||||
|
||||
assistantAText, ok := fantasy.AsMessagePart[fantasy.TextPart](prompt[0].Content[0])
|
||||
require.True(t, ok, "expected assistant A text")
|
||||
require.Equal(t, "assistant a", assistantAText.Text)
|
||||
assistantBText, ok := fantasy.AsMessagePart[fantasy.TextPart](prompt[2].Content[0])
|
||||
require.True(t, ok, "expected assistant B text")
|
||||
require.Equal(t, "assistant b", assistantBText.Text)
|
||||
userText, ok := fantasy.AsMessagePart[fantasy.TextPart](prompt[4].Content[0])
|
||||
require.True(t, ok, "expected user text")
|
||||
require.Equal(t, "next user message", userText.Text)
|
||||
}
|
||||
|
||||
// TestInjectMissingToolUses_DropsProviderExecutedOrphans verifies that
|
||||
// provider-executed tool results that end up after the wrong assistant
|
||||
// message (because they were persisted in a later step) are dropped
|
||||
|
||||
@@ -50,13 +50,32 @@ type OpenAIRequest struct {
|
||||
Messages []OpenAIMessage `json:"messages"`
|
||||
Stream bool `json:"stream,omitempty"`
|
||||
Tools []OpenAITool `json:"tools,omitempty"`
|
||||
Prompt []interface{} `json:"prompt,omitempty"` // For responses API
|
||||
Prompt []interface{} `json:"prompt,omitempty"` // Responses API input or prompt.
|
||||
Store *bool `json:"store,omitempty"`
|
||||
PreviousResponseID *string `json:"previous_response_id,omitempty"`
|
||||
// TODO: encoding/json ignores inline tags. Add custom UnmarshalJSON to capture unknown keys.
|
||||
Options map[string]interface{} `json:",inline"` //nolint:revive
|
||||
}
|
||||
|
||||
func (r *OpenAIRequest) UnmarshalJSON(data []byte) error {
|
||||
type openAIRequest OpenAIRequest
|
||||
decoded := struct {
|
||||
*openAIRequest
|
||||
Input []interface{} `json:"input,omitempty"`
|
||||
}{
|
||||
openAIRequest: (*openAIRequest)(r),
|
||||
}
|
||||
if err := json.Unmarshal(data, &decoded); err != nil {
|
||||
return err
|
||||
}
|
||||
// The Responses API uses input, while older fake-server tests
|
||||
// inspected prompt. Keep exposing both shapes through Prompt.
|
||||
if r.Prompt == nil && decoded.Input != nil {
|
||||
r.Prompt = decoded.Input
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// OpenAIMessage represents a message in an OpenAI request.
|
||||
type OpenAIMessage struct {
|
||||
Role string `json:"role"`
|
||||
@@ -213,6 +232,13 @@ func (s *openAIServer) handleResponses(w http.ResponseWriter, r *http.Request) {
|
||||
s.request = &req
|
||||
s.mu.Unlock()
|
||||
|
||||
if req.Prompt != nil {
|
||||
if errResp := ValidateResponsesAPIInput(req.Prompt); errResp != nil {
|
||||
writeErrorResponse(s.t, w, errResp)
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
resp := s.handler(&req)
|
||||
s.writeResponsesAPIResponse(w, &req, resp)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,196 @@
|
||||
package chattest
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// ValidateResponsesAPIInput validates the Responses API item relationships
|
||||
// that OpenAI enforces but the fake test server would otherwise miss.
|
||||
func ValidateResponsesAPIInput(items []interface{}) *ErrorResponse {
|
||||
if err := validateResponsesWebSearchReasoning(items); err != nil {
|
||||
return err
|
||||
}
|
||||
return validateResponsesFunctionCallOutputs(items)
|
||||
}
|
||||
|
||||
type responsesInputKind int
|
||||
|
||||
const (
|
||||
responsesInputOther responsesInputKind = iota
|
||||
responsesInputReasoning
|
||||
responsesInputWebSearch
|
||||
responsesInputFunctionCall
|
||||
responsesInputFunctionCallOutput
|
||||
)
|
||||
|
||||
type responsesInputItem struct {
|
||||
kind responsesInputKind
|
||||
id string
|
||||
callID string
|
||||
}
|
||||
|
||||
func validateResponsesWebSearchReasoning(items []interface{}) *ErrorResponse {
|
||||
previousKind := responsesInputOther
|
||||
for _, raw := range items {
|
||||
item := classifyResponsesInputItem(raw)
|
||||
if item.kind == responsesInputWebSearch && previousKind != responsesInputReasoning {
|
||||
return openAIResponsesValidationError(fmt.Sprintf(
|
||||
"Item %q of type 'web_search_call' was provided without its required 'reasoning' item.",
|
||||
item.id,
|
||||
))
|
||||
}
|
||||
previousKind = item.kind
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateResponsesFunctionCallOutputs(items []interface{}) *ErrorResponse {
|
||||
type callState struct {
|
||||
calls int
|
||||
outputs int
|
||||
firstCall int
|
||||
firstOutput int
|
||||
}
|
||||
states := make(map[string]*callState)
|
||||
var callIDs []string
|
||||
var outputCallIDs []string
|
||||
|
||||
stateFor := func(callID string) *callState {
|
||||
state, ok := states[callID]
|
||||
if ok {
|
||||
return state
|
||||
}
|
||||
state = &callState{firstCall: -1, firstOutput: -1}
|
||||
states[callID] = state
|
||||
return state
|
||||
}
|
||||
|
||||
for index, raw := range items {
|
||||
item := classifyResponsesInputItem(raw)
|
||||
switch item.kind {
|
||||
case responsesInputFunctionCall:
|
||||
if item.callID == "" {
|
||||
continue
|
||||
}
|
||||
state := stateFor(item.callID)
|
||||
if state.calls == 0 {
|
||||
callIDs = append(callIDs, item.callID)
|
||||
state.firstCall = index
|
||||
}
|
||||
state.calls++
|
||||
case responsesInputFunctionCallOutput:
|
||||
if item.callID == "" {
|
||||
continue
|
||||
}
|
||||
state := stateFor(item.callID)
|
||||
if state.outputs == 0 {
|
||||
outputCallIDs = append(outputCallIDs, item.callID)
|
||||
state.firstOutput = index
|
||||
}
|
||||
state.outputs++
|
||||
}
|
||||
}
|
||||
|
||||
for _, callID := range callIDs {
|
||||
state := states[callID]
|
||||
if state.calls > 1 {
|
||||
return openAIResponsesValidationError(fmt.Sprintf(
|
||||
"Duplicate function call found for call_id %s.", callID,
|
||||
))
|
||||
}
|
||||
}
|
||||
for _, callID := range outputCallIDs {
|
||||
state := states[callID]
|
||||
if state.outputs > 1 {
|
||||
return openAIResponsesValidationError(fmt.Sprintf(
|
||||
"Duplicate tool output found for function call %s.", callID,
|
||||
))
|
||||
}
|
||||
}
|
||||
for _, callID := range outputCallIDs {
|
||||
state := states[callID]
|
||||
if state.calls == 0 || state.firstOutput < state.firstCall {
|
||||
return openAIResponsesValidationError(fmt.Sprintf(
|
||||
"Tool output found without preceding function call %s.", callID,
|
||||
))
|
||||
}
|
||||
}
|
||||
for _, callID := range callIDs {
|
||||
state := states[callID]
|
||||
if state.outputs == 0 {
|
||||
return openAIResponsesValidationError(fmt.Sprintf(
|
||||
"No tool output found for function call %s.", callID,
|
||||
))
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func classifyResponsesInputItem(raw interface{}) responsesInputItem {
|
||||
itemMap, ok := raw.(map[string]interface{})
|
||||
if !ok {
|
||||
return responsesInputItem{kind: responsesInputOther}
|
||||
}
|
||||
|
||||
itemType := StringResponseField(itemMap, "type")
|
||||
id := StringResponseField(itemMap, "id")
|
||||
callID := StringResponseField(itemMap, "call_id")
|
||||
|
||||
switch itemType {
|
||||
case "reasoning":
|
||||
return responsesInputItem{kind: responsesInputReasoning, id: id}
|
||||
case "web_search_call":
|
||||
return responsesInputItem{kind: responsesInputWebSearch, id: id}
|
||||
case "function_call":
|
||||
return responsesInputItem{kind: responsesInputFunctionCall, callID: callID}
|
||||
case "function_call_output":
|
||||
return responsesInputItem{kind: responsesInputFunctionCallOutput, callID: callID}
|
||||
case "item_reference":
|
||||
switch {
|
||||
case strings.HasPrefix(id, "rs_"):
|
||||
return responsesInputItem{kind: responsesInputReasoning, id: id}
|
||||
case strings.HasPrefix(id, "ws_"):
|
||||
return responsesInputItem{kind: responsesInputWebSearch, id: id}
|
||||
default:
|
||||
return responsesInputItem{kind: responsesInputOther, id: id}
|
||||
}
|
||||
}
|
||||
|
||||
// Some SDK encoders omit the type field for item references. Fall
|
||||
// back to stable OpenAI item ID prefixes so tests still catch an
|
||||
// invalid prompt shape.
|
||||
switch {
|
||||
case strings.HasPrefix(id, "rs_"):
|
||||
return responsesInputItem{kind: responsesInputReasoning, id: id}
|
||||
case strings.HasPrefix(id, "ws_"):
|
||||
return responsesInputItem{kind: responsesInputWebSearch, id: id}
|
||||
default:
|
||||
return responsesInputItem{kind: responsesInputOther, id: id, callID: callID}
|
||||
}
|
||||
}
|
||||
|
||||
// StringResponseField returns the string value for key from a decoded
|
||||
// Responses API item, or an empty string when the field is absent or not a
|
||||
// string.
|
||||
func StringResponseField(values map[string]interface{}, key string) string {
|
||||
value, ok := values[key]
|
||||
if !ok {
|
||||
return ""
|
||||
}
|
||||
text, ok := value.(string)
|
||||
if !ok {
|
||||
return ""
|
||||
}
|
||||
return text
|
||||
}
|
||||
|
||||
func openAIResponsesValidationError(message string) *ErrorResponse {
|
||||
return &ErrorResponse{
|
||||
StatusCode: http.StatusBadRequest,
|
||||
Type: "invalid_request_error",
|
||||
Message: message,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,100 @@
|
||||
package chattest_test
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chattest"
|
||||
)
|
||||
|
||||
func TestValidateResponsesAPIInput(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
t.Run("valid reasoning and web search references", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
errResp := chattest.ValidateResponsesAPIInput([]interface{}{
|
||||
map[string]interface{}{"type": "item_reference", "id": "rs_valid"},
|
||||
map[string]interface{}{"type": "item_reference", "id": "ws_valid"},
|
||||
})
|
||||
require.Nil(t, errResp)
|
||||
})
|
||||
|
||||
t.Run("rejects web search without reasoning", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
errResp := chattest.ValidateResponsesAPIInput([]interface{}{
|
||||
map[string]interface{}{"type": "item_reference", "id": "ws_orphan"},
|
||||
})
|
||||
require.NotNil(t, errResp)
|
||||
require.Equal(t, 400, errResp.StatusCode)
|
||||
require.Contains(t, errResp.Message, "web_search_call")
|
||||
require.Contains(t, errResp.Message, "reasoning")
|
||||
})
|
||||
|
||||
t.Run("valid function call and output", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
errResp := chattest.ValidateResponsesAPIInput([]interface{}{
|
||||
map[string]interface{}{"type": "function_call", "call_id": "call_valid"},
|
||||
map[string]interface{}{"type": "function_call_output", "call_id": "call_valid"},
|
||||
})
|
||||
require.Nil(t, errResp)
|
||||
})
|
||||
|
||||
t.Run("rejects function call without output", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
errResp := chattest.ValidateResponsesAPIInput([]interface{}{
|
||||
map[string]interface{}{"type": "function_call", "call_id": "call_orphan"},
|
||||
})
|
||||
require.NotNil(t, errResp)
|
||||
require.Contains(t, errResp.Message, "No tool output found for function call call_orphan")
|
||||
})
|
||||
|
||||
t.Run("rejects output before function call", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
errResp := chattest.ValidateResponsesAPIInput([]interface{}{
|
||||
map[string]interface{}{"type": "function_call_output", "call_id": "call_late"},
|
||||
map[string]interface{}{"type": "function_call", "call_id": "call_late"},
|
||||
})
|
||||
require.NotNil(t, errResp)
|
||||
require.Contains(t, errResp.Message, "Tool output found without preceding function call call_late")
|
||||
})
|
||||
|
||||
t.Run("rejects duplicate function call", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
errResp := chattest.ValidateResponsesAPIInput([]interface{}{
|
||||
map[string]interface{}{"type": "function_call", "call_id": "call_duplicate"},
|
||||
map[string]interface{}{"type": "function_call", "call_id": "call_duplicate"},
|
||||
map[string]interface{}{"type": "function_call_output", "call_id": "call_duplicate"},
|
||||
})
|
||||
require.NotNil(t, errResp)
|
||||
require.Contains(t, errResp.Message, "Duplicate function call found for call_id call_duplicate")
|
||||
})
|
||||
|
||||
t.Run("rejects duplicate function call output", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
errResp := chattest.ValidateResponsesAPIInput([]interface{}{
|
||||
map[string]interface{}{"type": "function_call", "call_id": "call_duplicate_output"},
|
||||
map[string]interface{}{"type": "function_call_output", "call_id": "call_duplicate_output"},
|
||||
map[string]interface{}{"type": "function_call_output", "call_id": "call_duplicate_output"},
|
||||
})
|
||||
require.NotNil(t, errResp)
|
||||
require.Contains(t, errResp.Message, "Duplicate tool output found for function call call_duplicate_output")
|
||||
})
|
||||
|
||||
t.Run("classifies item reference by prefix without type field", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
errResp := chattest.ValidateResponsesAPIInput([]interface{}{
|
||||
map[string]interface{}{"id": "rs_prefix_only"},
|
||||
map[string]interface{}{"id": "ws_prefix_only"},
|
||||
})
|
||||
require.Nil(t, errResp)
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,622 @@
|
||||
package chatd_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/database/dbtestutil"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatprompt"
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chattest"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
"github.com/coder/coder/v2/testutil"
|
||||
)
|
||||
|
||||
func TestOpenAIResponsesNoStaleWebSearchReplay(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
|
||||
const (
|
||||
reasoningID = "rs_no_stale_reasoning"
|
||||
webSearchID = "ws_no_stale_search"
|
||||
)
|
||||
var recorder responsesRequestRecorder
|
||||
openAIURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
|
||||
if !req.Stream {
|
||||
return chattest.OpenAINonStreamingResponse("title")
|
||||
}
|
||||
|
||||
requestNumber := recorder.record(req)
|
||||
switch requestNumber {
|
||||
case 1:
|
||||
resp := chattest.OpenAIStreamingResponse(
|
||||
chattest.OpenAITextChunks("search result summary")...,
|
||||
)
|
||||
resp.ResponseID = "resp_no_stale_first"
|
||||
resp.Reasoning = &chattest.OpenAIReasoningItem{
|
||||
ID: reasoningID,
|
||||
Summary: "checked provider-side search state",
|
||||
EncryptedContent: "encrypted-no-stale",
|
||||
}
|
||||
resp.WebSearch = &chattest.OpenAIWebSearchCall{
|
||||
ID: webSearchID,
|
||||
Query: "coder changelog",
|
||||
}
|
||||
return resp
|
||||
default:
|
||||
resp := chattest.OpenAIStreamingResponse(
|
||||
chattest.OpenAITextChunks("follow-up answer")...,
|
||||
)
|
||||
resp.ResponseID = "resp_no_stale_second"
|
||||
return resp
|
||||
}
|
||||
})
|
||||
|
||||
user, org, _ := seedChatDependenciesWithProvider(ctx, t, db, "openai", openAIURL)
|
||||
model := insertOpenAIResponsesModelConfig(ctx, t, db, user.ID, false, true)
|
||||
server := newActiveTestServer(t, db, ps)
|
||||
|
||||
chat, err := server.CreateChat(ctx, chatd.CreateOptions{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: user.ID,
|
||||
Title: uniqueResponsesTitle(t, "no-stale"),
|
||||
ModelConfigID: model.ID,
|
||||
InitialUserContent: []codersdk.ChatMessagePart{
|
||||
codersdk.ChatMessageText("search for the latest Coder docs"),
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
waitForChatProcessed(ctx, t, db, chat.ID, server)
|
||||
requireResponsesChatWaiting(ctx, t, db, chat.ID)
|
||||
require.Len(t, recorder.all(), 1)
|
||||
|
||||
_, err = server.SendMessage(ctx, chatd.SendMessageOptions{
|
||||
ChatID: chat.ID,
|
||||
CreatedBy: user.ID,
|
||||
ModelConfigID: model.ID,
|
||||
Content: []codersdk.ChatMessagePart{
|
||||
codersdk.ChatMessageText("summarize the result without searching again"),
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
waitForChatProcessed(ctx, t, db, chat.ID, server)
|
||||
requireResponsesChatWaiting(ctx, t, db, chat.ID)
|
||||
|
||||
requests := recorder.all()
|
||||
require.Len(t, requests, 2)
|
||||
followup := requests[1]
|
||||
require.NotNil(t, followup.Store)
|
||||
require.False(t, *followup.Store)
|
||||
require.Nil(t, followup.PreviousResponseID)
|
||||
require.NotEmpty(t, followup.Prompt)
|
||||
requireNoResponsesProviderItemReplay(t, followup.Prompt, reasoningID, webSearchID)
|
||||
require.NotContains(t, promptItemTypes(followup.Prompt), "web_search_call")
|
||||
}
|
||||
|
||||
func TestOpenAIResponsesFullReplayPairsReasoningAndWebSearch(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
|
||||
const (
|
||||
reasoningID = "rs_full_replay_reasoning"
|
||||
webSearchID = "ws_full_replay_search"
|
||||
)
|
||||
var recorder responsesRequestRecorder
|
||||
openAIURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
|
||||
if !req.Stream {
|
||||
return chattest.OpenAINonStreamingResponse("title")
|
||||
}
|
||||
requestNumber := recorder.record(req)
|
||||
switch requestNumber {
|
||||
case 1:
|
||||
resp := chattest.OpenAIStreamingResponse(
|
||||
chattest.OpenAITextChunks("search result summary")...,
|
||||
)
|
||||
resp.ResponseID = "resp_full_replay_first"
|
||||
resp.Reasoning = &chattest.OpenAIReasoningItem{
|
||||
ID: reasoningID,
|
||||
Summary: "checked provider-side search state",
|
||||
EncryptedContent: "encrypted-full-replay",
|
||||
}
|
||||
resp.WebSearch = &chattest.OpenAIWebSearchCall{
|
||||
ID: webSearchID,
|
||||
Query: "coder changelog",
|
||||
}
|
||||
return resp
|
||||
default:
|
||||
resp := chattest.OpenAIStreamingResponse(
|
||||
chattest.OpenAITextChunks("follow-up answer")...,
|
||||
)
|
||||
resp.ResponseID = "resp_full_replay_second"
|
||||
return resp
|
||||
}
|
||||
})
|
||||
|
||||
user, org, _ := seedChatDependenciesWithProvider(ctx, t, db, "openai", openAIURL)
|
||||
firstModel := insertOpenAIResponsesModelConfig(ctx, t, db, user.ID, true, true)
|
||||
secondModel := insertOpenAIResponsesModelConfig(ctx, t, db, user.ID, true, true)
|
||||
server := newActiveTestServer(t, db, ps)
|
||||
|
||||
chat, err := server.CreateChat(ctx, chatd.CreateOptions{
|
||||
OrganizationID: org.ID,
|
||||
OwnerID: user.ID,
|
||||
Title: uniqueResponsesTitle(t, "full-replay"),
|
||||
ModelConfigID: firstModel.ID,
|
||||
InitialUserContent: []codersdk.ChatMessagePart{
|
||||
codersdk.ChatMessageText("search for the latest Coder docs"),
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
waitForChatProcessed(ctx, t, db, chat.ID, server)
|
||||
requireResponsesChatWaiting(ctx, t, db, chat.ID)
|
||||
require.Len(t, recorder.all(), 1)
|
||||
|
||||
_, err = server.SendMessage(ctx, chatd.SendMessageOptions{
|
||||
ChatID: chat.ID,
|
||||
CreatedBy: user.ID,
|
||||
ModelConfigID: secondModel.ID,
|
||||
Content: []codersdk.ChatMessagePart{
|
||||
codersdk.ChatMessageText("summarize the result without searching again"),
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
waitForChatProcessed(ctx, t, db, chat.ID, server)
|
||||
requireResponsesChatWaiting(ctx, t, db, chat.ID)
|
||||
|
||||
requests := recorder.all()
|
||||
require.Len(t, requests, 2)
|
||||
followup := requests[1]
|
||||
require.NotNil(t, followup.Store)
|
||||
require.True(t, *followup.Store)
|
||||
require.Nil(t, followup.PreviousResponseID)
|
||||
require.NotEmpty(t, followup.Prompt)
|
||||
requirePromptItemReferenceOrder(t, followup.Prompt, reasoningID, webSearchID)
|
||||
}
|
||||
|
||||
func TestOpenAIResponsesChainModeSkipsWhenLocalCallPending(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
|
||||
var recorder responsesRequestRecorder
|
||||
openAIURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
|
||||
if !req.Stream {
|
||||
return chattest.OpenAINonStreamingResponse("title")
|
||||
}
|
||||
recorder.record(req)
|
||||
resp := chattest.OpenAIStreamingResponse(
|
||||
chattest.OpenAITextChunks("resolved after local call")...,
|
||||
)
|
||||
resp.ResponseID = "resp_local_pending_next"
|
||||
return resp
|
||||
})
|
||||
|
||||
user, org, _ := seedChatDependenciesWithProvider(ctx, t, db, "openai", openAIURL)
|
||||
model := insertOpenAIResponsesModelConfig(ctx, t, db, user.ID, true, false)
|
||||
chat := insertOpenAIResponsesChat(ctx, t, db, org.ID, user.ID, model.ID, "local-pending")
|
||||
|
||||
callID := fmt.Sprintf("call_local_%d", time.Now().UnixNano())
|
||||
localCall := codersdk.ChatMessageToolCall(
|
||||
callID,
|
||||
"read_file",
|
||||
json.RawMessage(`{"path":"README.md"}`),
|
||||
)
|
||||
insertOpenAIResponsesMessages(ctx, t, db, chat.ID, user.ID, model.ID,
|
||||
persistedResponsesMessage{
|
||||
role: database.ChatMessageRoleUser,
|
||||
parts: []codersdk.ChatMessagePart{
|
||||
codersdk.ChatMessageText("please inspect the README"),
|
||||
},
|
||||
},
|
||||
persistedResponsesMessage{
|
||||
role: database.ChatMessageRoleAssistant,
|
||||
parts: []codersdk.ChatMessagePart{localCall},
|
||||
providerResponseID: "resp_local_pending_prior",
|
||||
},
|
||||
)
|
||||
|
||||
server := newActiveTestServer(t, db, ps)
|
||||
_, err := server.SendMessage(ctx, chatd.SendMessageOptions{
|
||||
ChatID: chat.ID,
|
||||
CreatedBy: user.ID,
|
||||
ModelConfigID: model.ID,
|
||||
Content: []codersdk.ChatMessagePart{
|
||||
codersdk.ChatMessageText("continue after that tool call"),
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
waitForChatProcessed(ctx, t, db, chat.ID, server)
|
||||
requireResponsesChatWaiting(ctx, t, db, chat.ID)
|
||||
|
||||
requests := recorder.all()
|
||||
require.Len(t, requests, 1)
|
||||
request := requests[0]
|
||||
require.NotNil(t, request.Store)
|
||||
require.True(t, *request.Store)
|
||||
require.Nil(t, request.PreviousResponseID)
|
||||
require.NotEmpty(t, request.Prompt)
|
||||
requirePromptItemWithTypeAndCallID(t, request.Prompt, "function_call", callID)
|
||||
requirePromptItemWithTypeAndCallID(t, request.Prompt, "function_call_output", callID)
|
||||
}
|
||||
|
||||
func TestOpenAIResponsesChainModeStillFiresForProviderExecutedOnly(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
db, ps := dbtestutil.NewDB(t)
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
|
||||
var recorder responsesRequestRecorder
|
||||
openAIURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
|
||||
if !req.Stream {
|
||||
return chattest.OpenAINonStreamingResponse("title")
|
||||
}
|
||||
recorder.record(req)
|
||||
resp := chattest.OpenAIStreamingResponse(
|
||||
chattest.OpenAITextChunks("chained answer")...,
|
||||
)
|
||||
resp.ResponseID = "resp_provider_only_next"
|
||||
return resp
|
||||
})
|
||||
|
||||
user, org, _ := seedChatDependenciesWithProvider(ctx, t, db, "openai", openAIURL)
|
||||
model := insertOpenAIResponsesModelConfig(ctx, t, db, user.ID, true, true)
|
||||
chat := insertOpenAIResponsesChat(ctx, t, db, org.ID, user.ID, model.ID, "provider-only")
|
||||
|
||||
const (
|
||||
previousResponseID = "resp_provider_only_prior"
|
||||
webSearchID = "ws_provider_only_search"
|
||||
)
|
||||
webSearchCall := codersdk.ChatMessageToolCall(
|
||||
webSearchID,
|
||||
"web_search",
|
||||
json.RawMessage(`{"query":"coder docs"}`),
|
||||
)
|
||||
webSearchCall.ProviderExecuted = true
|
||||
webSearchResult := codersdk.ChatMessageToolResult(
|
||||
webSearchID,
|
||||
"web_search",
|
||||
json.RawMessage(`{"status":"completed"}`),
|
||||
false,
|
||||
false,
|
||||
)
|
||||
webSearchResult.ProviderExecuted = true
|
||||
insertOpenAIResponsesMessages(ctx, t, db, chat.ID, user.ID, model.ID,
|
||||
persistedResponsesMessage{
|
||||
role: database.ChatMessageRoleUser,
|
||||
parts: []codersdk.ChatMessagePart{
|
||||
codersdk.ChatMessageText("look up the docs"),
|
||||
},
|
||||
},
|
||||
persistedResponsesMessage{
|
||||
role: database.ChatMessageRoleAssistant,
|
||||
parts: []codersdk.ChatMessagePart{
|
||||
webSearchCall,
|
||||
webSearchResult,
|
||||
},
|
||||
providerResponseID: previousResponseID,
|
||||
},
|
||||
)
|
||||
|
||||
server := newActiveTestServer(t, db, ps)
|
||||
_, err := server.SendMessage(ctx, chatd.SendMessageOptions{
|
||||
ChatID: chat.ID,
|
||||
CreatedBy: user.ID,
|
||||
ModelConfigID: model.ID,
|
||||
Content: []codersdk.ChatMessagePart{
|
||||
codersdk.ChatMessageText("what did it find"),
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
waitForChatProcessed(ctx, t, db, chat.ID, server)
|
||||
requireResponsesChatWaiting(ctx, t, db, chat.ID)
|
||||
|
||||
requests := recorder.all()
|
||||
require.Len(t, requests, 1)
|
||||
request := requests[0]
|
||||
require.NotNil(t, request.Store)
|
||||
require.True(t, *request.Store)
|
||||
require.NotNil(t, request.PreviousResponseID)
|
||||
require.Equal(t, previousResponseID, *request.PreviousResponseID)
|
||||
require.NotEmpty(t, request.Prompt)
|
||||
requireNoResponsesProviderItemReplay(t, request.Prompt, webSearchID)
|
||||
require.NotContains(t, promptItemTypes(request.Prompt), "web_search_call")
|
||||
require.NotContains(t, promptItemRoles(request.Prompt), "assistant")
|
||||
}
|
||||
|
||||
type recordedResponsesRequest struct {
|
||||
Prompt []interface{}
|
||||
Store *bool
|
||||
PreviousResponseID *string
|
||||
}
|
||||
|
||||
type responsesRequestRecorder struct {
|
||||
mu sync.Mutex
|
||||
requests []recordedResponsesRequest
|
||||
}
|
||||
|
||||
func (r *responsesRequestRecorder) record(req *chattest.OpenAIRequest) int {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
|
||||
var store *bool
|
||||
if req.Store != nil {
|
||||
value := *req.Store
|
||||
store = &value
|
||||
}
|
||||
var previousResponseID *string
|
||||
if req.PreviousResponseID != nil {
|
||||
value := *req.PreviousResponseID
|
||||
previousResponseID = &value
|
||||
}
|
||||
r.requests = append(r.requests, recordedResponsesRequest{
|
||||
Prompt: append([]interface{}(nil), req.Prompt...),
|
||||
Store: store,
|
||||
PreviousResponseID: previousResponseID,
|
||||
})
|
||||
return len(r.requests)
|
||||
}
|
||||
|
||||
func (r *responsesRequestRecorder) all() []recordedResponsesRequest {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
return append([]recordedResponsesRequest(nil), r.requests...)
|
||||
}
|
||||
|
||||
type persistedResponsesMessage struct {
|
||||
role database.ChatMessageRole
|
||||
parts []codersdk.ChatMessagePart
|
||||
providerResponseID string
|
||||
}
|
||||
|
||||
func insertOpenAIResponsesModelConfig(
|
||||
ctx context.Context,
|
||||
t *testing.T,
|
||||
db database.Store,
|
||||
userID uuid.UUID,
|
||||
store bool,
|
||||
webSearchEnabled bool,
|
||||
) database.ChatModelConfig {
|
||||
t.Helper()
|
||||
return insertChatModelConfigWithCallConfig(
|
||||
ctx,
|
||||
t,
|
||||
db,
|
||||
userID,
|
||||
"openai",
|
||||
"gpt-4o",
|
||||
codersdk.ChatModelCallConfig{
|
||||
ProviderOptions: &codersdk.ChatModelProviderOptions{
|
||||
OpenAI: &codersdk.ChatModelOpenAIProviderOptions{
|
||||
Store: &store,
|
||||
WebSearchEnabled: &webSearchEnabled,
|
||||
},
|
||||
},
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
func insertOpenAIResponsesChat(
|
||||
ctx context.Context,
|
||||
t *testing.T,
|
||||
db database.Store,
|
||||
organizationID uuid.UUID,
|
||||
ownerID uuid.UUID,
|
||||
modelConfigID uuid.UUID,
|
||||
titlePrefix string,
|
||||
) database.Chat {
|
||||
t.Helper()
|
||||
chat, err := db.InsertChat(ctx, database.InsertChatParams{
|
||||
OrganizationID: organizationID,
|
||||
OwnerID: ownerID,
|
||||
LastModelConfigID: modelConfigID,
|
||||
Title: uniqueResponsesTitle(t, titlePrefix),
|
||||
Status: database.ChatStatusWaiting,
|
||||
MCPServerIDs: []uuid.UUID{},
|
||||
ClientType: database.ChatClientTypeApi,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
return chat
|
||||
}
|
||||
|
||||
func insertOpenAIResponsesMessages(
|
||||
ctx context.Context,
|
||||
t *testing.T,
|
||||
db database.Store,
|
||||
chatID uuid.UUID,
|
||||
createdBy uuid.UUID,
|
||||
modelConfigID uuid.UUID,
|
||||
messages ...persistedResponsesMessage,
|
||||
) {
|
||||
t.Helper()
|
||||
params := database.InsertChatMessagesParams{ChatID: chatID}
|
||||
for _, message := range messages {
|
||||
content, err := chatprompt.MarshalParts(message.parts)
|
||||
require.NoError(t, err)
|
||||
params.CreatedBy = append(params.CreatedBy, createdBy)
|
||||
params.ModelConfigID = append(params.ModelConfigID, modelConfigID)
|
||||
params.Role = append(params.Role, message.role)
|
||||
params.Content = append(params.Content, string(content.RawMessage))
|
||||
params.ContentVersion = append(params.ContentVersion, chatprompt.CurrentContentVersion)
|
||||
params.Visibility = append(params.Visibility, database.ChatMessageVisibilityBoth)
|
||||
params.InputTokens = append(params.InputTokens, 0)
|
||||
params.OutputTokens = append(params.OutputTokens, 0)
|
||||
params.TotalTokens = append(params.TotalTokens, 0)
|
||||
params.ReasoningTokens = append(params.ReasoningTokens, 0)
|
||||
params.CacheCreationTokens = append(params.CacheCreationTokens, 0)
|
||||
params.CacheReadTokens = append(params.CacheReadTokens, 0)
|
||||
params.ContextLimit = append(params.ContextLimit, 0)
|
||||
params.Compressed = append(params.Compressed, false)
|
||||
params.TotalCostMicros = append(params.TotalCostMicros, 0)
|
||||
params.RuntimeMs = append(params.RuntimeMs, 0)
|
||||
params.ProviderResponseID = append(params.ProviderResponseID, message.providerResponseID)
|
||||
}
|
||||
_, err := db.InsertChatMessages(ctx, params)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
func requireResponsesChatWaiting(
|
||||
ctx context.Context,
|
||||
t *testing.T,
|
||||
db database.Store,
|
||||
chatID uuid.UUID,
|
||||
) {
|
||||
t.Helper()
|
||||
chat, err := db.GetChatByID(ctx, chatID)
|
||||
require.NoError(t, err)
|
||||
if chat.Status == database.ChatStatusError {
|
||||
require.FailNowf(t, "chat failed", "last_error=%q", chat.LastError.String)
|
||||
}
|
||||
require.Equal(t, database.ChatStatusWaiting, chat.Status)
|
||||
}
|
||||
|
||||
func uniqueResponsesTitle(t *testing.T, prefix string) string {
|
||||
t.Helper()
|
||||
return fmt.Sprintf("%s-%s-%d", prefix, t.Name(), time.Now().UnixNano())
|
||||
}
|
||||
|
||||
func promptItemTypes(prompt []interface{}) []string {
|
||||
types := make([]string, 0, len(prompt))
|
||||
for _, item := range prompt {
|
||||
itemMap, ok := item.(map[string]interface{})
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if itemType := chattest.StringResponseField(itemMap, "type"); itemType != "" {
|
||||
types = append(types, itemType)
|
||||
}
|
||||
}
|
||||
return types
|
||||
}
|
||||
|
||||
func promptItemRoles(prompt []interface{}) []string {
|
||||
roles := make([]string, 0, len(prompt))
|
||||
for _, item := range prompt {
|
||||
itemMap, ok := item.(map[string]interface{})
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if role := chattest.StringResponseField(itemMap, "role"); role != "" {
|
||||
roles = append(roles, role)
|
||||
}
|
||||
}
|
||||
return roles
|
||||
}
|
||||
|
||||
func requirePromptItemWithTypeAndCallID(
|
||||
t *testing.T,
|
||||
prompt []interface{},
|
||||
itemType string,
|
||||
callID string,
|
||||
) map[string]interface{} {
|
||||
t.Helper()
|
||||
for _, item := range prompt {
|
||||
itemMap, ok := item.(map[string]interface{})
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if chattest.StringResponseField(itemMap, "type") == itemType &&
|
||||
chattest.StringResponseField(itemMap, "call_id") == callID {
|
||||
return itemMap
|
||||
}
|
||||
}
|
||||
promptJSON, err := json.Marshal(prompt)
|
||||
require.NoError(t, err)
|
||||
require.FailNowf(t, "prompt item missing",
|
||||
"missing type=%q call_id=%q in prompt %s", itemType, callID, promptJSON)
|
||||
return nil
|
||||
}
|
||||
|
||||
// requireNoResponsesProviderItemReplay rejects the explicit stale IDs and all
|
||||
// provider-managed Responses item IDs. Chain mode should rely on
|
||||
// previous_response_id, not replay rs_ or ws_ identifiers in prompt input.
|
||||
func requireNoResponsesProviderItemReplay(
|
||||
t *testing.T,
|
||||
prompt []interface{},
|
||||
staleIDs ...string,
|
||||
) {
|
||||
t.Helper()
|
||||
stale := make(map[string]struct{}, len(staleIDs))
|
||||
for _, id := range staleIDs {
|
||||
stale[id] = struct{}{}
|
||||
}
|
||||
for _, item := range prompt {
|
||||
assertNoResponsesProviderItemReplay(t, item, stale)
|
||||
}
|
||||
}
|
||||
|
||||
func assertNoResponsesProviderItemReplay(
|
||||
t *testing.T,
|
||||
value interface{},
|
||||
staleIDs map[string]struct{},
|
||||
) {
|
||||
t.Helper()
|
||||
switch typed := value.(type) {
|
||||
case map[string]interface{}:
|
||||
for key, raw := range typed {
|
||||
if text, ok := raw.(string); ok {
|
||||
if key == "type" && text == "web_search_call" {
|
||||
require.FailNow(t, "prompt replayed web_search_call provider item")
|
||||
}
|
||||
if key == "id" || key == "call_id" || key == "item_id" {
|
||||
if _, isStale := staleIDs[text]; isStale {
|
||||
require.FailNowf(t, "prompt replayed stale provider item ID",
|
||||
"field %q contained stale provider ID %q", key, text)
|
||||
}
|
||||
if strings.HasPrefix(text, "ws_") || strings.HasPrefix(text, "rs_") {
|
||||
require.FailNowf(t, "prompt replayed provider item ID",
|
||||
"field %q contained provider-managed ID %q", key, text)
|
||||
}
|
||||
}
|
||||
}
|
||||
assertNoResponsesProviderItemReplay(t, raw, staleIDs)
|
||||
}
|
||||
case []interface{}:
|
||||
for _, item := range typed {
|
||||
assertNoResponsesProviderItemReplay(t, item, staleIDs)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func requirePromptItemReferenceOrder(
|
||||
t *testing.T,
|
||||
prompt []interface{},
|
||||
firstID string,
|
||||
secondID string,
|
||||
) {
|
||||
t.Helper()
|
||||
firstIndex := -1
|
||||
secondIndex := -1
|
||||
for index, item := range prompt {
|
||||
itemMap, ok := item.(map[string]interface{})
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
itemID := chattest.StringResponseField(itemMap, "id")
|
||||
if itemID == "" {
|
||||
itemID = chattest.StringResponseField(itemMap, "item_id")
|
||||
}
|
||||
switch itemID {
|
||||
case firstID:
|
||||
firstIndex = index
|
||||
case secondID:
|
||||
secondIndex = index
|
||||
}
|
||||
}
|
||||
require.NotEqual(t, -1, firstIndex, "missing first item reference")
|
||||
require.NotEqual(t, -1, secondIndex, "missing second item reference")
|
||||
require.Less(t, firstIndex, secondIndex)
|
||||
}
|
||||
@@ -83,8 +83,11 @@ replace github.com/spf13/afero => github.com/aslilac/afero v0.0.0-20250403163713
|
||||
// 4) (anthropic-sdk-go) dannykopping's appendCompact performance fixes
|
||||
// 5) (anthropic-sdk-go) DirectEncoder to eliminate nested MarshalJSON allocation chain
|
||||
// 6) Anthropic EffortXHigh constant for Claude Opus 4.7
|
||||
// See: https://github.com/coder/fantasy/commits/5ab464a305f4
|
||||
replace charm.land/fantasy => github.com/coder/fantasy v0.0.0-20260424191546-5ab464a305f4
|
||||
// 7) coder/fantasy#mike/openai-responses-continuity, OpenAI Responses replay safety:
|
||||
// replay stored reasoning item references, only replay web_search references
|
||||
// when paired with reasoning, and validate function_call output pairing.
|
||||
// See: https://github.com/coder/fantasy/commits/f83367a4a205
|
||||
replace charm.land/fantasy => github.com/coder/fantasy v0.0.0-20260426185602-951a49c681df
|
||||
|
||||
// coder/coder uses a fork of charmbracelet's fork of the Anthropic Go SDK with some
|
||||
// additional performance improvements.
|
||||
|
||||
@@ -322,8 +322,8 @@ github.com/coder/bubbletea v1.2.2-0.20241212190825-007a1cdb2c41 h1:SBN/DA63+ZHwu
|
||||
github.com/coder/bubbletea v1.2.2-0.20241212190825-007a1cdb2c41/go.mod h1:I9ULxr64UaOSUv7hcb3nX4kowodJCVS7vt7VVJk/kW4=
|
||||
github.com/coder/clistat v1.2.1 h1:P9/10njXMyj5cWzIU5wkRsSy5LVQH49+tcGMsAgWX0w=
|
||||
github.com/coder/clistat v1.2.1/go.mod h1:m7SC0uj88eEERgvF8Kn6+w6XF21BeSr+15f7GoLAw0A=
|
||||
github.com/coder/fantasy v0.0.0-20260424191546-5ab464a305f4 h1:eL6f03ujlQiXs24FjAtdyD8TMbRiZtDWj6JdcSlMeUw=
|
||||
github.com/coder/fantasy v0.0.0-20260424191546-5ab464a305f4/go.mod h1:wZ0e3lEPqrM0XiIdAUQLvMKCLYhc3gi96MRX2wjbX44=
|
||||
github.com/coder/fantasy v0.0.0-20260426185602-951a49c681df h1:Xog/dBDcnXxr98lGZqRxOeFrCrhVZUBrFldtXH7v0EY=
|
||||
github.com/coder/fantasy v0.0.0-20260426185602-951a49c681df/go.mod h1:wZ0e3lEPqrM0XiIdAUQLvMKCLYhc3gi96MRX2wjbX44=
|
||||
github.com/coder/flog v1.1.0 h1:kbAes1ai8fIS5OeV+QAnKBQE22ty1jRF/mcAwHpLBa4=
|
||||
github.com/coder/flog v1.1.0/go.mod h1:UQlQvrkJBvnRGo69Le8E24Tcl5SJleAAR7gYEHzAmdQ=
|
||||
github.com/coder/go-httpstat v0.0.0-20230801153223-321c88088322 h1:m0lPZjlQ7vdVpRBPKfYIFlmgevoTkBxB10wv6l2gOaU=
|
||||
|
||||
Reference in New Issue
Block a user