diff --git a/coderd/x/chatd/chatd.go b/coderd/x/chatd/chatd.go index d9183bb4cb..ad8ff0c896 100644 --- a/coderd/x/chatd/chatd.go +++ b/coderd/x/chatd/chatd.go @@ -41,6 +41,7 @@ import ( "github.com/coder/coder/v2/coderd/x/chatd/chatdebug" "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/chatopenai" "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/chatretry" @@ -3457,262 +3458,6 @@ func (m chatMessage) withProviderResponseID(id string) chatMessage { return m } -// chainModeInfo holds the information needed to determine whether -// a follow-up turn can use OpenAI's previous_response_id chaining -// instead of replaying full conversation history. -type chainModeInfo struct { - // previousResponseID is the provider response ID from the last - // assistant message, if any. - previousResponseID string - // modelConfigID is the model configuration used to produce the - // assistant message referenced by previousResponseID. - modelConfigID uuid.UUID - // trailingUserCount is the number of contiguous user messages - // at the end of the conversation that form the current turn. - trailingUserCount int - // 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 - // providerMissingToolResults is true when the assistant message - // has local tool calls with local results, but no follow-up - // assistant message exists to confirm the results were sent - // back to the provider. This happens when StopAfterTool - // terminates a turn before the results are round-tripped. - providerMissingToolResults bool -} - -func userMessageContributesToChainMode(msg database.ChatMessage) bool { - parts, err := chatprompt.ParseContent(msg) - if err != nil { - return false - } - for _, part := range parts { - switch part.Type { - case codersdk.ChatMessagePartTypeText, - codersdk.ChatMessagePartTypeReasoning: - if strings.TrimSpace(part.Text) != "" { - return true - } - case codersdk.ChatMessagePartTypeFile, - codersdk.ChatMessagePartTypeFileReference: - return true - case codersdk.ChatMessagePartTypeContextFile: - if part.ContextFileContent != "" { - return true - } - } - } - 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) -} - -// providerHasMissingToolResults reports whether the assistant message -// at assistantIdx has local tool calls whose results exist in the DB -// but were never sent back to the provider. This is detected by the -// absence of a follow-up assistant message after the tool results: -// in normal flow the LLM processes tool results and produces a -// follow-up response, but StopAfterTool skips that round-trip. -func providerHasMissingToolResults( - messages []database.ChatMessage, - assistantIdx int, -) bool { - if assistantIdx < 0 || assistantIdx >= len(messages) { - return false - } - - parts, err := chatprompt.ParseContent(messages[assistantIdx]) - if err != nil { - // Parsing errors are already handled by - // assistantHasUnresolvedLocalToolCalls. - return false - } - - if !slices.ContainsFunc(parts, func(p codersdk.ChatMessagePart) bool { - return p.Type == codersdk.ChatMessagePartTypeToolCall && !p.ProviderExecuted - }) { - return false - } - - // Scan forward past tool messages. If the first non-tool message - // is not an assistant, the tool results were never round-tripped - // to the provider. - for i := assistantIdx + 1; i < len(messages); i++ { - switch messages[i].Role { - case database.ChatMessageRoleTool: - continue - case database.ChatMessageRoleAssistant: - // A follow-up assistant exists; results were sent. - return false - default: - // User or system message with no follow-up assistant. - return true - } - } - // Reached end of messages without a follow-up assistant. - return true -} - -// 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, -// complete local tool state, and confirmation that tool results were -// actually sent to the provider (not just persisted locally). -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 && - !info.providerMissingToolResults -} - -// 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. -func resolveChainMode(messages []database.ChatMessage) chainModeInfo { - var info chainModeInfo - i := len(messages) - 1 - for ; i >= 0; i-- { - if messages[i].Role != database.ChatMessageRoleUser { - break - } - info.trailingUserCount++ - if userMessageContributesToChainMode(messages[i]) { - info.contributingTrailingUserCount++ - } - } - for ; i >= 0; i-- { - switch messages[i].Role { - case database.ChatMessageRoleAssistant: - if messages[i].ProviderResponseID.Valid && - messages[i].ProviderResponseID.String != "" { - info.previousResponseID = messages[i].ProviderResponseID.String - if messages[i].ModelConfigID.Valid { - info.modelConfigID = messages[i].ModelConfigID.UUID - } - info.hasUnresolvedLocalToolCalls = assistantHasUnresolvedLocalToolCalls(messages, i) - if !info.hasUnresolvedLocalToolCalls { - info.providerMissingToolResults = providerHasMissingToolResults(messages, i) - } - return info - } - return info - case database.ChatMessageRoleTool: - continue - default: - return info - } - } - return info -} - -// filterPromptForChainMode keeps only system messages and the trailing -// user messages that still contribute model-visible content to the -// current turn. Assistant and tool messages are dropped because the -// provider already has them via the previous_response_id chain. -func filterPromptForChainMode( - prompt []fantasy.Message, - info chainModeInfo, -) []fantasy.Message { - if info.contributingTrailingUserCount <= 0 { - return prompt - } - - totalUsers := 0 - for _, msg := range prompt { - if msg.Role == "user" { - totalUsers++ - } - } - - // Prompt construction already drops user turns with no model-visible - // content, such as skill-only sentinel messages. That means the user - // count here stays aligned with contributingTrailingUserCount even - // when non-contributing DB turns are interleaved in the trailing - // block. - usersToSkip := totalUsers - info.contributingTrailingUserCount - if usersToSkip < 0 { - usersToSkip = 0 - } - - filtered := make([]fantasy.Message, 0, len(prompt)) - usersSeen := 0 - for _, msg := range prompt { - switch msg.Role { - case "system": - filtered = append(filtered, msg) - case "user": - usersSeen++ - if usersSeen > usersToSkip { - filtered = append(filtered, msg) - } - } - } - - return filtered -} - // appendChatMessage appends a single message to the batch insert params. func appendChatMessage( params *database.InsertChatMessagesParams, @@ -6346,7 +6091,7 @@ func (p *Server) runChat( advisorPromptSnapshot = slices.Clone(msgs) } - chainInfo := resolveChainMode(messages) + chainInfo := chatopenai.ResolveChainMode(messages) result.PushSummaryModel = model result.ProviderKeys = providerKeys result.FallbackProvider = modelConfig.Provider @@ -7123,7 +6868,7 @@ func (p *Server) runChat( // blocked for all Explore chats. var providerTools []chatloop.ProviderTool if !isPlanModeTurn && callConfig.ProviderOptions != nil { - providerTools = buildProviderTools(model.Provider(), callConfig.ProviderOptions) + providerTools = buildProviderTools(callConfig.ProviderOptions) if isExploreSubagent { if !chat.ParentChatID.Valid { providerTools = nil @@ -7162,28 +6907,28 @@ 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 := shouldActivateChainMode( + chainModeActive := chatopenai.ShouldActivateChainMode( providerOptions, chainInfo, modelConfig.ID, isPlanModeTurn, ) - if !chainModeActive && chainInfo.previousResponseID != "" { + if !chainModeActive && chainInfo.PreviousResponseID() != "" { logger.Debug(ctx, "chain mode disabled", - slog.F("has_unresolved_local_tool_calls", chainInfo.hasUnresolvedLocalToolCalls), - slog.F("provider_missing_tool_results", chainInfo.providerMissingToolResults), + slog.F("has_unresolved_local_tool_calls", chainInfo.HasUnresolvedLocalToolCalls()), + slog.F("provider_missing_tool_results", chainInfo.ProviderMissingToolResults()), slog.F("is_plan_mode_turn", isPlanModeTurn), - slog.F("model_config_match", chainInfo.modelConfigID == modelConfig.ID), - slog.F("store_enabled", chatprovider.IsResponsesStoreEnabled(providerOptions)), - slog.F("contributing_trailing_user_count", chainInfo.contributingTrailingUserCount), + slog.F("model_config_match", chainInfo.ModelConfigID() == modelConfig.ID), + slog.F("store_enabled", chatopenai.IsResponsesStoreEnabled(providerOptions)), + slog.F("contributing_trailing_user_count", chainInfo.ContributingTrailingUserCount()), ) } if chainModeActive { - providerOptions = chatprovider.CloneWithPreviousResponseID( + providerOptions = chatopenai.WithPreviousResponseID( providerOptions, - chainInfo.previousResponseID, + chainInfo.PreviousResponseID(), ) - prompt = filterPromptForChainMode(prompt, chainInfo) + prompt = chatopenai.FilterPromptForChainMode(prompt, chainInfo) } activeToolNames := activeToolNamesForTurn( tools, @@ -7338,7 +7083,7 @@ func (p *Server) runChat( // history is unavailable. setAdvisorPromptSnapshot(reloadedPrompt) if chainModeActive { - reloadedPrompt = filterPromptForChainMode( + reloadedPrompt = chatopenai.FilterPromptForChainMode( reloadedPrompt, chainInfo, ) @@ -7421,7 +7166,7 @@ func (p *Server) runChat( // buildProviderTools creates provider-native tool definitions // (like web search) based on the model configuration. These // tools are executed server-side by the LLM provider. -func buildProviderTools(_ string, options *codersdk.ChatModelProviderOptions) []chatloop.ProviderTool { +func buildProviderTools(options *codersdk.ChatModelProviderOptions) []chatloop.ProviderTool { var tools []chatloop.ProviderTool if options == nil { @@ -7437,20 +7182,9 @@ func buildProviderTools(_ string, options *codersdk.ChatModelProviderOptions) [] }) } - if options.OpenAI != nil && options.OpenAI.WebSearchEnabled != nil && *options.OpenAI.WebSearchEnabled { - args := map[string]any{} - if options.OpenAI.SearchContextSize != nil && *options.OpenAI.SearchContextSize != "" { - args["search_context_size"] = *options.OpenAI.SearchContextSize - } - if len(options.OpenAI.AllowedDomains) > 0 { - args["allowed_domains"] = options.OpenAI.AllowedDomains - } + if tool, ok := chatopenai.WebSearchTool(options.OpenAI); ok { tools = append(tools, chatloop.ProviderTool{ - Definition: fantasy.ProviderDefinedTool{ - ID: "web_search", - Name: "web_search", - Args: args, - }, + Definition: tool, }) } diff --git a/coderd/x/chatd/chatd_internal_test.go b/coderd/x/chatd/chatd_internal_test.go index e76c5ae24f..c22a4785a5 100644 --- a/coderd/x/chatd/chatd_internal_test.go +++ b/coderd/x/chatd/chatd_internal_test.go @@ -9,7 +9,6 @@ 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" @@ -24,7 +23,6 @@ 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" @@ -2603,7 +2601,7 @@ func TestSkillsFromParts(t *testing.T) { t.Run("NoSkillParts", func(t *testing.T) { t.Parallel() msgs := []database.ChatMessage{ - chatMessageWithParts([]codersdk.ChatMessagePart{ + chattest.ChatMessageWithParts([]codersdk.ChatMessagePart{ {Type: codersdk.ChatMessagePartTypeText, Text: "hello"}, }), } @@ -2614,7 +2612,7 @@ func TestSkillsFromParts(t *testing.T) { t.Run("SingleSkill", func(t *testing.T) { t.Parallel() msgs := []database.ChatMessage{ - chatMessageWithParts([]codersdk.ChatMessagePart{ + chattest.ChatMessageWithParts([]codersdk.ChatMessagePart{ { Type: codersdk.ChatMessagePartTypeSkill, SkillName: "deep-review", @@ -2633,14 +2631,14 @@ func TestSkillsFromParts(t *testing.T) { t.Run("MultipleSkillsAcrossMessages", func(t *testing.T) { t.Parallel() msgs := []database.ChatMessage{ - chatMessageWithParts([]codersdk.ChatMessagePart{ + chattest.ChatMessageWithParts([]codersdk.ChatMessagePart{ { Type: codersdk.ChatMessagePartTypeSkill, SkillName: "pull-requests", SkillDir: "/home/coder/.agents/skills/pull-requests", }, }), - chatMessageWithParts([]codersdk.ChatMessagePart{ + chattest.ChatMessageWithParts([]codersdk.ChatMessagePart{ { Type: codersdk.ChatMessagePartTypeSkill, SkillName: "deep-review", @@ -2657,7 +2655,7 @@ func TestSkillsFromParts(t *testing.T) { t.Run("MixedPartTypes", func(t *testing.T) { t.Parallel() msgs := []database.ChatMessage{ - chatMessageWithParts([]codersdk.ChatMessagePart{ + chattest.ChatMessageWithParts([]codersdk.ChatMessagePart{ { Type: codersdk.ChatMessagePartTypeContextFile, ContextFilePath: "/home/coder/.coder/AGENTS.md", @@ -2669,7 +2667,7 @@ func TestSkillsFromParts(t *testing.T) { }, }), // A text-only message should be skipped entirely. - chatMessageWithParts([]codersdk.ChatMessagePart{ + chattest.ChatMessageWithParts([]codersdk.ChatMessagePart{ {Type: codersdk.ChatMessagePartTypeText, Text: "user turn"}, }), } @@ -2682,7 +2680,7 @@ func TestSkillsFromParts(t *testing.T) { t.Run("OptionalDescriptionOmitted", func(t *testing.T) { t.Parallel() msgs := []database.ChatMessage{ - chatMessageWithParts([]codersdk.ChatMessagePart{ + chattest.ChatMessageWithParts([]codersdk.ChatMessagePart{ { Type: codersdk.ChatMessagePartTypeSkill, SkillName: "refine-plan", @@ -2730,7 +2728,7 @@ func TestSkillsFromParts(t *testing.T) { ContextFileAgentID: uuid.NullUUID{UUID: agentID, Valid: true}, }) } - msgs := []database.ChatMessage{chatMessageWithParts(parts)} + msgs := []database.ChatMessage{chattest.ChatMessageWithParts(parts)} got := skillsFromParts(msgs) require.Len(t, got, len(want)) for i, w := range want { @@ -2754,7 +2752,7 @@ func TestContextFileAgentID(t *testing.T) { t.Run("NoContextFileParts", func(t *testing.T) { t.Parallel() msgs := []database.ChatMessage{ - chatMessageWithParts([]codersdk.ChatMessagePart{ + chattest.ChatMessageWithParts([]codersdk.ChatMessagePart{ {Type: codersdk.ChatMessagePartTypeText, Text: "hello"}, }), } @@ -2767,7 +2765,7 @@ func TestContextFileAgentID(t *testing.T) { t.Parallel() agentID := uuid.New() msgs := []database.ChatMessage{ - chatMessageWithParts([]codersdk.ChatMessagePart{ + chattest.ChatMessageWithParts([]codersdk.ChatMessagePart{ { Type: codersdk.ChatMessagePartTypeContextFile, ContextFilePath: "/some/path", @@ -2785,14 +2783,14 @@ func TestContextFileAgentID(t *testing.T) { agentID1 := uuid.New() agentID2 := uuid.New() msgs := []database.ChatMessage{ - chatMessageWithParts([]codersdk.ChatMessagePart{ + chattest.ChatMessageWithParts([]codersdk.ChatMessagePart{ { Type: codersdk.ChatMessagePartTypeContextFile, ContextFilePath: "/first/path", ContextFileAgentID: uuid.NullUUID{UUID: agentID1, Valid: true}, }, }), - chatMessageWithParts([]codersdk.ChatMessagePart{ + chattest.ChatMessageWithParts([]codersdk.ChatMessagePart{ { Type: codersdk.ChatMessagePartTypeContextFile, ContextFilePath: "/second/path", @@ -2810,12 +2808,12 @@ func TestContextFileAgentID(t *testing.T) { instructionAgentID := uuid.New() sentinelAgentID := uuid.New() msgs := []database.ChatMessage{ - chatMessageWithParts([]codersdk.ChatMessagePart{{ + chattest.ChatMessageWithParts([]codersdk.ChatMessagePart{{ Type: codersdk.ChatMessagePartTypeContextFile, ContextFilePath: "/workspace/AGENTS.md", ContextFileAgentID: uuid.NullUUID{UUID: instructionAgentID, Valid: true}, }}), - chatMessageWithParts([]codersdk.ChatMessagePart{{ + chattest.ChatMessageWithParts([]codersdk.ChatMessagePart{{ Type: codersdk.ChatMessagePartTypeContextFile, ContextFilePath: AgentChatContextSentinelPath, ContextFileAgentID: uuid.NullUUID{ @@ -2832,7 +2830,7 @@ func TestContextFileAgentID(t *testing.T) { t.Run("SentinelWithoutAgentID", func(t *testing.T) { t.Parallel() msgs := []database.ChatMessage{ - chatMessageWithParts([]codersdk.ChatMessagePart{ + chattest.ChatMessageWithParts([]codersdk.ChatMessagePart{ { Type: codersdk.ChatMessagePartTypeContextFile, ContextFileAgentID: uuid.NullUUID{Valid: false}, @@ -2852,7 +2850,7 @@ func TestHasPersistedInstructionFiles(t *testing.T) { t.Parallel() agentID := uuid.New() msgs := []database.ChatMessage{ - chatMessageWithParts([]codersdk.ChatMessagePart{{ + chattest.ChatMessageWithParts([]codersdk.ChatMessagePart{{ Type: codersdk.ChatMessagePartTypeContextFile, ContextFilePath: AgentChatContextSentinelPath, ContextFileAgentID: uuid.NullUUID{ @@ -2868,7 +2866,7 @@ func TestHasPersistedInstructionFiles(t *testing.T) { t.Parallel() agentID := uuid.New() msgs := []database.ChatMessage{ - chatMessageWithParts([]codersdk.ChatMessagePart{{ + chattest.ChatMessageWithParts([]codersdk.ChatMessagePart{{ Type: codersdk.ChatMessagePartTypeContextFile, ContextFilePath: "/workspace/AGENTS.md", ContextFileContent: "repo instructions", @@ -2885,7 +2883,7 @@ func TestInstructionFromContextFilesUsesLatestContextAgent(t *testing.T) { oldAgentID := uuid.New() newAgentID := uuid.New() msgs := []database.ChatMessage{ - chatMessageWithParts([]codersdk.ChatMessagePart{{ + chattest.ChatMessageWithParts([]codersdk.ChatMessagePart{{ Type: codersdk.ChatMessagePartTypeContextFile, ContextFilePath: "/old/AGENTS.md", ContextFileContent: "old instructions", @@ -2893,7 +2891,7 @@ func TestInstructionFromContextFilesUsesLatestContextAgent(t *testing.T) { ContextFileDirectory: "/old", ContextFileAgentID: uuid.NullUUID{UUID: oldAgentID, Valid: true}, }}), - chatMessageWithParts([]codersdk.ChatMessagePart{{ + chattest.ChatMessageWithParts([]codersdk.ChatMessagePart{{ Type: codersdk.ChatMessagePartTypeContextFile, ContextFilePath: "/new/AGENTS.md", ContextFileContent: "new instructions", @@ -2917,12 +2915,12 @@ func TestInstructionFromContextFilesKeepsLegacyUnstampedParts(t *testing.T) { oldAgentID := uuid.New() newAgentID := uuid.New() msgs := []database.ChatMessage{ - chatMessageWithParts([]codersdk.ChatMessagePart{{ + chattest.ChatMessageWithParts([]codersdk.ChatMessagePart{{ Type: codersdk.ChatMessagePartTypeContextFile, ContextFilePath: "/legacy/AGENTS.md", ContextFileContent: "legacy instructions", }}), - chatMessageWithParts([]codersdk.ChatMessagePart{{ + chattest.ChatMessageWithParts([]codersdk.ChatMessagePart{{ Type: codersdk.ChatMessagePartTypeContextFile, ContextFilePath: "/old/AGENTS.md", ContextFileContent: "old instructions", @@ -2930,7 +2928,7 @@ func TestInstructionFromContextFilesKeepsLegacyUnstampedParts(t *testing.T) { ContextFileDirectory: "/old", ContextFileAgentID: uuid.NullUUID{UUID: oldAgentID, Valid: true}, }}), - chatMessageWithParts([]codersdk.ChatMessagePart{{ + chattest.ChatMessageWithParts([]codersdk.ChatMessagePart{{ Type: codersdk.ChatMessagePartTypeContextFile, ContextFilePath: "/new/AGENTS.md", ContextFileContent: "new instructions", @@ -2955,12 +2953,12 @@ func TestSkillsFromPartsKeepsLegacyUnstampedParts(t *testing.T) { oldAgentID := uuid.New() newAgentID := uuid.New() msgs := []database.ChatMessage{ - chatMessageWithParts([]codersdk.ChatMessagePart{{ + chattest.ChatMessageWithParts([]codersdk.ChatMessagePart{{ Type: codersdk.ChatMessagePartTypeSkill, SkillName: "repo-helper-legacy", SkillDir: "/skills/repo-helper-legacy", }}), - chatMessageWithParts([]codersdk.ChatMessagePart{ + chattest.ChatMessageWithParts([]codersdk.ChatMessagePart{ { Type: codersdk.ChatMessagePartTypeContextFile, ContextFilePath: "/old/AGENTS.md", @@ -2973,7 +2971,7 @@ func TestSkillsFromPartsKeepsLegacyUnstampedParts(t *testing.T) { ContextFileAgentID: uuid.NullUUID{UUID: oldAgentID, Valid: true}, }, }), - chatMessageWithParts([]codersdk.ChatMessagePart{ + chattest.ChatMessageWithParts([]codersdk.ChatMessagePart{ { Type: codersdk.ChatMessagePartTypeContextFile, ContextFilePath: AgentChatContextSentinelPath, @@ -3004,7 +3002,7 @@ func TestSkillsFromPartsUsesLatestContextAgent(t *testing.T) { oldAgentID := uuid.New() newAgentID := uuid.New() msgs := []database.ChatMessage{ - chatMessageWithParts([]codersdk.ChatMessagePart{ + chattest.ChatMessageWithParts([]codersdk.ChatMessagePart{ { Type: codersdk.ChatMessagePartTypeContextFile, ContextFilePath: "/old/AGENTS.md", @@ -3017,7 +3015,7 @@ func TestSkillsFromPartsUsesLatestContextAgent(t *testing.T) { ContextFileAgentID: uuid.NullUUID{UUID: oldAgentID, Valid: true}, }, }), - chatMessageWithParts([]codersdk.ChatMessagePart{ + chattest.ChatMessageWithParts([]codersdk.ChatMessagePart{ { Type: codersdk.ChatMessagePartTypeContextFile, ContextFilePath: AgentChatContextSentinelPath, @@ -3113,625 +3111,6 @@ func TestSelectSkillMetasForInstructionRefresh(t *testing.T) { }) } -func TestResolveChainModeIgnoresSkillOnlySentinelMessages(t *testing.T) { - t.Parallel() - - modelConfigID := uuid.New() - assistant := database.ChatMessage{ - Role: database.ChatMessageRoleAssistant, - ProviderResponseID: sql.NullString{String: "resp-123", Valid: true}, - ModelConfigID: uuid.NullUUID{UUID: modelConfigID, Valid: true}, - } - skillOnly := chatMessageWithParts([]codersdk.ChatMessagePart{ - { - Type: codersdk.ChatMessagePartTypeContextFile, - ContextFilePath: AgentChatContextSentinelPath, - ContextFileAgentID: uuid.NullUUID{ - UUID: uuid.New(), - Valid: true, - }, - }, - { - Type: codersdk.ChatMessagePartTypeSkill, - SkillName: "repo-helper", - SkillDir: "/skills/repo-helper", - }, - }) - skillOnly.Role = database.ChatMessageRoleUser - user := chatMessageWithParts([]codersdk.ChatMessagePart{{ - Type: codersdk.ChatMessagePartTypeText, - Text: "latest user message", - }}) - user.Role = database.ChatMessageRoleUser - - got := resolveChainMode([]database.ChatMessage{assistant, skillOnly, user}) - require.Equal(t, "resp-123", got.previousResponseID) - require.Equal(t, modelConfigID, got.modelConfigID) - require.Equal(t, 2, got.trailingUserCount) - 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.False(t, chainInfo.providerMissingToolResults) - 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, - ) - - // A follow-up assistant after the tool result confirms the - // result was sent back to the provider. Chain mode should - // activate from the follow-up assistant's response ID. - // Use a distinct response ID on the follow-up assistant - // so the assertion verifies resolveChainMode selects the - // follow-up (last assistant), not the original tool-caller. - followUp := chainModeAssistantMessage(modelConfigID, nil) - followUp.ProviderResponseID = sql.NullString{String: "resp-follow-up", Valid: true} - - chainInfo := resolveChainMode([]database.ChatMessage{ - chainModeSystemMessage(), - chainModeUserMessage("prior user message"), - chainModeAssistantMessage(modelConfigID, []codersdk.ChatMessagePart{toolCall}), - chainModeToolMessage([]codersdk.ChatMessagePart{toolResult}), - followUp, - chainModeUserMessage("latest user message"), - }) - - require.Equal(t, "resp-follow-up", chainInfo.previousResponseID) - require.False(t, chainInfo.hasUnresolvedLocalToolCalls) - require.False(t, chainInfo.providerMissingToolResults) - 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, - )) -} - -// Tests for providerMissingToolResults detection. -// These cover the StopAfterTool + chain mode desync bug where local -// tool results exist in the DB but were never sent back to the -// provider, leaving an unresolved function_call in the stored chain. - -func TestResolveChainMode_BlocksWhenToolResultNeverSentToProvider(t *testing.T) { - t.Parallel() - - modelConfigID := uuid.New() - toolCall := codersdk.ChatMessageToolCall( - "call-local", - "propose_plan", - json.RawMessage(`{"path":"plan.md"}`), - ) - toolResult := codersdk.ChatMessageToolResult( - "call-local", - "propose_plan", - json.RawMessage(`{"ok":true}`), - false, - false, - ) - - chainInfo := resolveChainMode([]database.ChatMessage{ - chainModeSystemMessage(), - chainModeUserMessage("make a plan"), - chainModeAssistantMessage(modelConfigID, []codersdk.ChatMessagePart{toolCall}), - chainModeToolMessage([]codersdk.ChatMessagePart{toolResult}), - // No follow-up assistant: StopAfterTool fired, tool result - // was persisted locally but never sent back to the provider. - chainModeUserMessage("implement the plan"), - }) - - require.Equal(t, "resp-123", chainInfo.previousResponseID) - // Local tool calls are resolved (result exists in DB). - require.False(t, chainInfo.hasUnresolvedLocalToolCalls) - // But the provider never received the result. - require.True(t, chainInfo.providerMissingToolResults) - // Chain mode must NOT activate. - require.False(t, shouldActivateChainMode( - chainModeProviderOptions(), - chainInfo, - modelConfigID, - false, - )) -} - -func TestResolveChainMode_BlocksProviderMissingWithMultipleToolCalls(t *testing.T) { - t.Parallel() - - modelConfigID := uuid.New() - call1 := codersdk.ChatMessageToolCall( - "call-1", "propose_plan", - json.RawMessage(`{"path":"plan.md"}`), - ) - call2 := codersdk.ChatMessageToolCall( - "call-2", "write_file", - json.RawMessage(`{"path":"foo.go"}`), - ) - result1 := codersdk.ChatMessageToolResult( - "call-1", "propose_plan", - json.RawMessage(`{"ok":true}`), false, false, - ) - result2 := codersdk.ChatMessageToolResult( - "call-2", "write_file", - json.RawMessage(`{"ok":true}`), false, false, - ) - - chainInfo := resolveChainMode([]database.ChatMessage{ - chainModeSystemMessage(), - chainModeUserMessage("do it"), - chainModeAssistantMessage(modelConfigID, []codersdk.ChatMessagePart{call1, call2}), - chainModeToolMessage([]codersdk.ChatMessagePart{result1, result2}), - chainModeUserMessage("next"), - }) - - require.False(t, chainInfo.hasUnresolvedLocalToolCalls) - require.True(t, chainInfo.providerMissingToolResults) - require.False(t, shouldActivateChainMode( - chainModeProviderOptions(), chainInfo, modelConfigID, false, - )) -} - -func TestResolveChainMode_AllowsWhenNoToolCalls(t *testing.T) { - t.Parallel() - - modelConfigID := uuid.New() - - chainInfo := resolveChainMode([]database.ChatMessage{ - chainModeSystemMessage(), - chainModeUserMessage("hello"), - chainModeAssistantMessage(modelConfigID, nil), - chainModeUserMessage("thanks"), - }) - - require.Equal(t, "resp-123", chainInfo.previousResponseID) - require.False(t, chainInfo.hasUnresolvedLocalToolCalls) - require.False(t, chainInfo.providerMissingToolResults) - require.True(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() - - modelConfigID := uuid.New() - priorUser := chatMessageWithParts([]codersdk.ChatMessagePart{{ - Type: codersdk.ChatMessagePartTypeText, - Text: "prior user message", - }}) - priorUser.Role = database.ChatMessageRoleUser - assistant := database.ChatMessage{ - Role: database.ChatMessageRoleAssistant, - ProviderResponseID: sql.NullString{String: "resp-123", Valid: true}, - ModelConfigID: uuid.NullUUID{UUID: modelConfigID, Valid: true}, - } - firstTrailingUser := chatMessageWithParts([]codersdk.ChatMessagePart{{ - Type: codersdk.ChatMessagePartTypeText, - Text: "first trailing user", - }}) - firstTrailingUser.Role = database.ChatMessageRoleUser - skillOnly := chatMessageWithParts([]codersdk.ChatMessagePart{ - { - Type: codersdk.ChatMessagePartTypeContextFile, - ContextFilePath: AgentChatContextSentinelPath, - ContextFileAgentID: uuid.NullUUID{ - UUID: uuid.New(), - Valid: true, - }, - }, - { - Type: codersdk.ChatMessagePartTypeSkill, - SkillName: "repo-helper", - SkillDir: "/skills/repo-helper", - }, - }) - skillOnly.Role = database.ChatMessageRoleUser - lastTrailingUser := chatMessageWithParts([]codersdk.ChatMessagePart{{ - Type: codersdk.ChatMessagePartTypeText, - Text: "last trailing user", - }}) - lastTrailingUser.Role = database.ChatMessageRoleUser - - chainInfo := resolveChainMode([]database.ChatMessage{ - priorUser, - assistant, - firstTrailingUser, - skillOnly, - lastTrailingUser, - }) - require.Equal(t, 3, chainInfo.trailingUserCount) - require.Equal(t, 2, chainInfo.contributingTrailingUserCount) - - prompt := []fantasy.Message{ - { - Role: fantasy.MessageRoleSystem, - Content: []fantasy.MessagePart{ - fantasy.TextPart{Text: "system instruction"}, - }, - }, - { - Role: fantasy.MessageRoleUser, - Content: []fantasy.MessagePart{ - fantasy.TextPart{Text: "prior user message"}, - }, - }, - { - Role: fantasy.MessageRoleAssistant, - Content: []fantasy.MessagePart{ - fantasy.TextPart{Text: "assistant reply"}, - }, - }, - { - Role: fantasy.MessageRoleUser, - Content: []fantasy.MessagePart{ - fantasy.TextPart{Text: "first trailing user"}, - }, - }, - { - Role: fantasy.MessageRoleUser, - Content: []fantasy.MessagePart{ - fantasy.TextPart{Text: "last trailing user"}, - }, - }, - } - - got := filterPromptForChainMode(prompt, chainInfo) - require.Len(t, got, 3) - require.Equal(t, fantasy.MessageRoleSystem, got[0].Role) - require.Equal(t, fantasy.MessageRoleUser, got[1].Role) - require.Equal(t, fantasy.MessageRoleUser, got[2].Role) - - firstPart, ok := fantasy.AsMessagePart[fantasy.TextPart](got[1].Content[0]) - require.True(t, ok) - require.Equal(t, "first trailing user", firstPart.Text) - lastPart, ok := fantasy.AsMessagePart[fantasy.TextPart](got[2].Content[0]) - require.True(t, ok) - require.Equal(t, "last trailing user", lastPart.Text) -} - -func TestFilterPromptForChainModeUsesContributingTrailingUsers(t *testing.T) { - t.Parallel() - - modelConfigID := uuid.New() - priorUser := chatMessageWithParts([]codersdk.ChatMessagePart{{ - Type: codersdk.ChatMessagePartTypeText, - Text: "prior user message", - }}) - priorUser.Role = database.ChatMessageRoleUser - assistant := database.ChatMessage{ - Role: database.ChatMessageRoleAssistant, - ProviderResponseID: sql.NullString{String: "resp-123", Valid: true}, - ModelConfigID: uuid.NullUUID{UUID: modelConfigID, Valid: true}, - } - skillOnly := chatMessageWithParts([]codersdk.ChatMessagePart{ - { - Type: codersdk.ChatMessagePartTypeContextFile, - ContextFilePath: AgentChatContextSentinelPath, - ContextFileAgentID: uuid.NullUUID{ - UUID: uuid.New(), - Valid: true, - }, - }, - { - Type: codersdk.ChatMessagePartTypeSkill, - SkillName: "repo-helper", - SkillDir: "/skills/repo-helper", - }, - }) - skillOnly.Role = database.ChatMessageRoleUser - latestUser := chatMessageWithParts([]codersdk.ChatMessagePart{{ - Type: codersdk.ChatMessagePartTypeText, - Text: "latest user message", - }}) - latestUser.Role = database.ChatMessageRoleUser - - chainInfo := resolveChainMode([]database.ChatMessage{ - priorUser, - assistant, - skillOnly, - latestUser, - }) - require.Equal(t, 2, chainInfo.trailingUserCount) - require.Equal(t, 1, chainInfo.contributingTrailingUserCount) - - prompt := []fantasy.Message{ - { - Role: fantasy.MessageRoleSystem, - Content: []fantasy.MessagePart{ - fantasy.TextPart{Text: "system instruction"}, - }, - }, - { - Role: fantasy.MessageRoleUser, - Content: []fantasy.MessagePart{ - fantasy.TextPart{Text: "prior user message"}, - }, - }, - { - Role: fantasy.MessageRoleAssistant, - Content: []fantasy.MessagePart{ - fantasy.TextPart{Text: "assistant reply"}, - }, - }, - { - Role: fantasy.MessageRoleUser, - Content: []fantasy.MessagePart{ - fantasy.TextPart{Text: "latest user message"}, - }, - }, - } - - got := filterPromptForChainMode(prompt, chainInfo) - require.Len(t, got, 2) - require.Equal(t, fantasy.MessageRoleSystem, got[0].Role) - require.Equal(t, fantasy.MessageRoleUser, got[1].Role) - - part, ok := fantasy.AsMessagePart[fantasy.TextPart](got[1].Content[0]) - require.True(t, ok) - require.Equal(t, "latest user message", part.Text) -} - -func chatMessageWithParts(parts []codersdk.ChatMessagePart) database.ChatMessage { - raw, _ := json.Marshal(parts) - return database.ChatMessage{ - Content: pqtype.NullRawMessage{RawMessage: raw, Valid: true}, - } -} - // TestProcessChat_IgnoresStaleControlNotification verifies that // processChat is not interrupted by a "pending" notification // published before processing begins. This is the race that caused diff --git a/coderd/x/chatd/chatloop/chatloop.go b/coderd/x/chatd/chatloop/chatloop.go index 34de2e4e74..bda2167dca 100644 --- a/coderd/x/chatd/chatloop/chatloop.go +++ b/coderd/x/chatd/chatloop/chatloop.go @@ -16,7 +16,6 @@ import ( "charm.land/fantasy" fantasyanthropic "charm.land/fantasy/providers/anthropic" - fantasyopenai "charm.land/fantasy/providers/openai" "charm.land/fantasy/schema" "golang.org/x/xerrors" @@ -24,6 +23,7 @@ import ( "github.com/coder/coder/v2/coderd/database/dbtime" "github.com/coder/coder/v2/coderd/x/chatd/chatdebug" "github.com/coder/coder/v2/coderd/x/chatd/chaterror" + "github.com/coder/coder/v2/coderd/x/chatd/chatopenai" "github.com/coder/coder/v2/coderd/x/chatd/chatprompt" "github.com/coder/coder/v2/coderd/x/chatd/chatretry" "github.com/coder/coder/v2/coderd/x/chatd/chatsanitize" @@ -512,7 +512,7 @@ func Run(ctx context.Context, opts RunOptions) error { Content: result.content, Usage: result.usage, ContextLimit: contextLimit, - ProviderResponseID: extractOpenAIResponseIDIfStored(opts.ProviderOptions, result.providerMetadata), + ProviderResponseID: chatopenai.ExtractResponseIDIfStored(opts.ProviderOptions, result.providerMetadata), Runtime: time.Since(stepStart), ToolCallCreatedAt: result.toolCallCreatedAt, ToolResultCreatedAt: result.toolResultCreatedAt, @@ -538,8 +538,8 @@ func Run(ctx context.Context, opts RunOptions) error { // when previous_response_id is set, so we must leave chain // mode and reload the full history before the next call. stepMessages := result.toResponseMessages() - if hasPreviousResponseID(opts.ProviderOptions) { - clearPreviousResponseID(opts.ProviderOptions) + if chatopenai.HasPreviousResponseID(opts.ProviderOptions) { + opts.ProviderOptions = chatopenai.ClearPreviousResponseID(opts.ProviderOptions) if opts.DisableChainMode != nil { opts.DisableChainMode() } @@ -1233,7 +1233,7 @@ func persistPendingDynamicStep( Content: result.content, Usage: result.usage, ContextLimit: contextLimit, - ProviderResponseID: extractOpenAIResponseIDIfStored(opts.ProviderOptions, result.providerMetadata), + ProviderResponseID: chatopenai.ExtractResponseIDIfStored(opts.ProviderOptions, result.providerMetadata), Runtime: time.Since(stepStart), PendingDynamicToolCalls: pending, }); err != nil { @@ -1709,85 +1709,6 @@ func addAnthropicPromptCaching(messages []fantasy.Message) { } } -// hasPreviousResponseID checks whether the provider options contain -// an OpenAI Responses entry with a non-empty PreviousResponseID. -func hasPreviousResponseID(providerOptions fantasy.ProviderOptions) bool { - if providerOptions == nil { - return false - } - - for _, entry := range providerOptions { - if options, ok := entry.(*fantasyopenai.ResponsesProviderOptions); ok { - return options.PreviousResponseID != nil && - *options.PreviousResponseID != "" - } - } - - return false -} - -// clearPreviousResponseID removes PreviousResponseID from the OpenAI -// Responses provider options entry, if present. -func clearPreviousResponseID(providerOptions fantasy.ProviderOptions) { - if providerOptions == nil { - return - } - - for _, entry := range providerOptions { - if options, ok := entry.(*fantasyopenai.ResponsesProviderOptions); ok { - options.PreviousResponseID = nil - } - } -} - -// extractOpenAIResponseID extracts the OpenAI Responses API response -// ID from provider metadata. Returns an empty string if no OpenAI -// Responses metadata is present. -func extractOpenAIResponseID(metadata fantasy.ProviderMetadata) string { - if len(metadata) == 0 { - return "" - } - - for _, entry := range metadata { - if providerMetadata, ok := entry.(*fantasyopenai.ResponsesProviderMetadata); ok && providerMetadata != nil { - return providerMetadata.ResponseID - } - } - - return "" -} - -// extractOpenAIResponseIDIfStored returns the OpenAI response ID -// only when the provider options indicate store=true. Response IDs -// from store=false turns are not persisted server-side and cannot -// be used for chaining. -func extractOpenAIResponseIDIfStored( - providerOptions fantasy.ProviderOptions, - metadata fantasy.ProviderMetadata, -) string { - if !isResponsesStoreEnabled(providerOptions) { - return "" - } - - return extractOpenAIResponseID(metadata) -} - -// isResponsesStoreEnabled checks whether the OpenAI Responses -// provider options explicitly enable store=true. -func isResponsesStoreEnabled(providerOptions fantasy.ProviderOptions) bool { - if providerOptions == nil { - return false - } - - for _, entry := range providerOptions { - if options, ok := entry.(*fantasyopenai.ResponsesProviderOptions); ok { - return options.Store != nil && *options.Store - } - } - - return false -} - // recordToolResultTimestamp lazily initializes the // toolResultCreatedAt map on the stepResult and records // the completion timestamp for the given tool-call ID. diff --git a/coderd/x/chatd/chatopenai/options.go b/coderd/x/chatd/chatopenai/options.go new file mode 100644 index 0000000000..91d87fe582 --- /dev/null +++ b/coderd/x/chatd/chatopenai/options.go @@ -0,0 +1,228 @@ +package chatopenai + +import ( + "slices" + "strings" + + "charm.land/fantasy" + fantasyazure "charm.land/fantasy/providers/azure" + fantasyopenai "charm.land/fantasy/providers/openai" + + "github.com/coder/coder/v2/coderd/x/chatd/chatutil" + "github.com/coder/coder/v2/codersdk" +) + +// ProviderOptionsFromChatConfig converts chat model OpenAI options to fantasy +// provider options used for inference calls. +func ProviderOptionsFromChatConfig( + model fantasy.LanguageModel, + options *codersdk.ChatModelOpenAIProviderOptions, +) fantasy.ProviderOptionsData { + reasoningEffort := ReasoningEffortFromChat(options.ReasoningEffort) + if UsesResponsesOptions(model) { + include := EnsureResponseIncludes(IncludeFromChat(options.Include)) + providerOptions := &fantasyopenai.ResponsesProviderOptions{ + Include: include, + Instructions: chatutil.NormalizedStringPointer(options.Instructions), + Logprobs: ResponsesLogProbsFromChatConfig(options), + MaxToolCalls: options.MaxToolCalls, + Metadata: options.Metadata, + ParallelToolCalls: options.ParallelToolCalls, + PromptCacheKey: chatutil.NormalizedStringPointer(options.PromptCacheKey), + ReasoningEffort: reasoningEffort, + ReasoningSummary: chatutil.NormalizedStringPointer(options.ReasoningSummary), + SafetyIdentifier: chatutil.NormalizedStringPointer(options.SafetyIdentifier), + ServiceTier: ServiceTierFromChat(options.ServiceTier), + StrictJSONSchema: options.StrictJSONSchema, + Store: boolPtrOrDefault(options.Store, true), + TextVerbosity: TextVerbosityFromChat(options.TextVerbosity), + User: chatutil.NormalizedStringPointer(options.User), + } + return providerOptions + } + + return &fantasyopenai.ProviderOptions{ + LogitBias: options.LogitBias, + LogProbs: options.LogProbs, + TopLogProbs: options.TopLogProbs, + ParallelToolCalls: options.ParallelToolCalls, + User: chatutil.NormalizedStringPointer(options.User), + ReasoningEffort: reasoningEffort, + MaxCompletionTokens: options.MaxCompletionTokens, + TextVerbosity: chatutil.NormalizedStringPointer(options.TextVerbosity), + Prediction: options.Prediction, + Store: boolPtrOrDefault(options.Store, true), + Metadata: options.Metadata, + PromptCacheKey: chatutil.NormalizedStringPointer(options.PromptCacheKey), + SafetyIdentifier: chatutil.NormalizedStringPointer(options.SafetyIdentifier), + ServiceTier: chatutil.NormalizedStringPointer(options.ServiceTier), + StructuredOutputs: options.StructuredOutputs, + } +} + +// TextVerbosityFromChat normalizes chat-config text verbosity values for +// OpenAI and returns the canonical provider verbosity value. +func TextVerbosityFromChat(value *string) *fantasyopenai.TextVerbosity { + if value == nil { + return nil + } + + normalized := strings.ToLower(strings.TrimSpace(*value)) + if normalized == "" { + return nil + } + + verbosity := chatutil.NormalizedEnumValue( + normalized, + string(fantasyopenai.TextVerbosityLow), + string(fantasyopenai.TextVerbosityMedium), + string(fantasyopenai.TextVerbosityHigh), + ) + if verbosity == nil { + return nil + } + valueCopy := fantasyopenai.TextVerbosity(*verbosity) + return &valueCopy +} + +// IncludeFromChat converts chat-config include values to OpenAI Responses +// include values and ignores unsupported entries. +func IncludeFromChat(values []string) []fantasyopenai.IncludeType { + if values == nil { + return nil + } + + result := make([]fantasyopenai.IncludeType, 0, len(values)) + for _, value := range values { + switch strings.TrimSpace(value) { + case string(fantasyopenai.IncludeReasoningEncryptedContent): + result = append(result, fantasyopenai.IncludeReasoningEncryptedContent) + case string(fantasyopenai.IncludeFileSearchCallResults): + result = append(result, fantasyopenai.IncludeFileSearchCallResults) + case string(fantasyopenai.IncludeMessageOutputTextLogprobs): + result = append(result, fantasyopenai.IncludeMessageOutputTextLogprobs) + } + } + return result +} + +// EnsureResponseIncludes adds the OpenAI encrypted reasoning include required +// for Responses API reasoning continuity when it is not already present. +func EnsureResponseIncludes( + values []fantasyopenai.IncludeType, +) []fantasyopenai.IncludeType { + const required = fantasyopenai.IncludeReasoningEncryptedContent + + if slices.Contains(values, required) { + return values + } + return append(values, required) +} + +// UsesResponsesOptions reports whether the model should use OpenAI Responses +// API provider options. +func UsesResponsesOptions(model fantasy.LanguageModel) bool { + if model == nil { + return false + } + switch model.Provider() { + case fantasyopenai.Name, fantasyazure.Name: + return fantasyopenai.IsResponsesModel(model.Model()) + default: + return false + } +} + +// ReasoningEffortFromChat normalizes chat-config reasoning effort values for +// OpenAI and returns the canonical provider effort value. +func ReasoningEffortFromChat(value *string) *fantasyopenai.ReasoningEffort { + if value == nil { + return nil + } + + normalized := strings.ToLower(strings.TrimSpace(*value)) + if normalized == "" { + return nil + } + + effort := chatutil.NormalizedEnumValue( + normalized, + string(fantasyopenai.ReasoningEffortMinimal), + string(fantasyopenai.ReasoningEffortLow), + string(fantasyopenai.ReasoningEffortMedium), + string(fantasyopenai.ReasoningEffortHigh), + string(fantasyopenai.ReasoningEffortXHigh), + ) + if effort == nil { + return nil + } + valueCopy := fantasyopenai.ReasoningEffort(*effort) + return &valueCopy +} + +// ServiceTierFromChat normalizes chat-config service tier values for OpenAI +// Responses API and returns the canonical provider service tier value. +func ServiceTierFromChat(value *string) *fantasyopenai.ServiceTier { + normalized := chatutil.NormalizedStringPointer(value) + if normalized == nil { + return nil + } + switch strings.ToLower(*normalized) { + case string(fantasyopenai.ServiceTierAuto): + serviceTier := fantasyopenai.ServiceTierAuto + return &serviceTier + case string(fantasyopenai.ServiceTierFlex): + serviceTier := fantasyopenai.ServiceTierFlex + return &serviceTier + case string(fantasyopenai.ServiceTierPriority): + serviceTier := fantasyopenai.ServiceTierPriority + return &serviceTier + default: + return nil + } +} + +// ResponsesLogProbsFromChatConfig maps chat-config log probability options to the +// value expected by OpenAI Responses provider options. +func ResponsesLogProbsFromChatConfig( + options *codersdk.ChatModelOpenAIProviderOptions, +) any { + if options == nil { + return nil + } + if options.TopLogProbs != nil { + return *options.TopLogProbs + } + if options.LogProbs != nil { + return *options.LogProbs + } + return nil +} + +// IsReasoningModel reports whether a model ID follows OpenAI reasoning model +// naming conventions. +func IsReasoningModel(modelID string) bool { + if len(modelID) < 2 || modelID[0] != 'o' { + return false + } + + index := 1 + for index < len(modelID) && modelID[index] >= '0' && modelID[index] <= '9' { + index++ + } + if index == 1 { + return false + } + + if index == len(modelID) { + return true + } + return modelID[index] == '-' || modelID[index] == '.' +} + +func boolPtrOrDefault(value *bool, def bool) *bool { + if value != nil { + return value + } + return &def +} diff --git a/coderd/x/chatd/chatopenai/options_test.go b/coderd/x/chatd/chatopenai/options_test.go new file mode 100644 index 0000000000..1320300b11 --- /dev/null +++ b/coderd/x/chatd/chatopenai/options_test.go @@ -0,0 +1,499 @@ +package chatopenai_test + +import ( + "context" + "testing" + + "charm.land/fantasy" + fantasyazure "charm.land/fantasy/providers/azure" + fantasyopenai "charm.land/fantasy/providers/openai" + "github.com/stretchr/testify/require" + + "github.com/coder/coder/v2/coderd/x/chatd/chatopenai" + "github.com/coder/coder/v2/codersdk" +) + +func TestProviderOptionsFromChatConfigLegacy(t *testing.T) { + t.Parallel() + + store := false + logProbs := true + topLogProbs := int64(3) + parallelToolCalls := true + maxCompletionTokens := int64(4096) + structuredOutputs := true + options := &codersdk.ChatModelOpenAIProviderOptions{ + LogitBias: map[string]int64{ + "50256": -10, + }, + LogProbs: &logProbs, + TopLogProbs: &topLogProbs, + ParallelToolCalls: ¶llelToolCalls, + User: ptr(" user-1 "), + ReasoningEffort: ptr(" HIGH "), + MaxCompletionTokens: &maxCompletionTokens, + TextVerbosity: ptr(" High "), + Prediction: map[string]any{ + "type": "content", + }, + Store: &store, + Metadata: map[string]any{"feature": "chat"}, + PromptCacheKey: ptr(" cache-key "), + SafetyIdentifier: ptr(" safety-id "), + ServiceTier: ptr(" priority "), + StructuredOutputs: &structuredOutputs, + } + + got := chatopenai.ProviderOptionsFromChatConfig( + fakeLanguageModel{provider: fantasyopenai.Name, model: "gpt-3.5-turbo-instruct"}, + options, + ) + + providerOptions, ok := got.(*fantasyopenai.ProviderOptions) + require.True(t, ok) + require.Equal(t, options.LogitBias, providerOptions.LogitBias) + require.Same(t, options.LogProbs, providerOptions.LogProbs) + require.Same(t, options.TopLogProbs, providerOptions.TopLogProbs) + require.Same(t, options.ParallelToolCalls, providerOptions.ParallelToolCalls) + require.Equal(t, "user-1", requireStringPointerValue(t, providerOptions.User)) + require.Equal(t, fantasyopenai.ReasoningEffortHigh, requireReasoningEffortPointerValue(t, providerOptions.ReasoningEffort)) + require.Same(t, options.MaxCompletionTokens, providerOptions.MaxCompletionTokens) + require.Equal(t, "High", requireStringPointerValue(t, providerOptions.TextVerbosity)) + require.Equal(t, options.Prediction, providerOptions.Prediction) + require.Same(t, options.Store, providerOptions.Store) + require.Equal(t, false, requireBoolPointerValue(t, providerOptions.Store)) + require.Equal(t, options.Metadata, providerOptions.Metadata) + require.Equal(t, "cache-key", requireStringPointerValue(t, providerOptions.PromptCacheKey)) + require.Equal(t, "safety-id", requireStringPointerValue(t, providerOptions.SafetyIdentifier)) + require.Equal(t, "priority", requireStringPointerValue(t, providerOptions.ServiceTier)) + require.Same(t, options.StructuredOutputs, providerOptions.StructuredOutputs) +} + +func TestProviderOptionsFromChatConfigResponses(t *testing.T) { + t.Parallel() + + topLogProbs := int64(5) + maxToolCalls := int64(8) + parallelToolCalls := false + strictJSONSchema := true + options := &codersdk.ChatModelOpenAIProviderOptions{ + Include: []string{ + string(fantasyopenai.IncludeFileSearchCallResults), + "unsupported", + }, + Instructions: ptr(" instructions "), + LogProbs: ptr(true), + TopLogProbs: &topLogProbs, + MaxToolCalls: &maxToolCalls, + Metadata: map[string]any{"scope": "unit"}, + ParallelToolCalls: ¶llelToolCalls, + PromptCacheKey: ptr(" prompt-cache "), + ReasoningEffort: ptr(" minimal "), + ReasoningSummary: ptr(" auto "), + SafetyIdentifier: ptr(" safety "), + ServiceTier: ptr(" FLEX "), + StrictJSONSchema: &strictJSONSchema, + TextVerbosity: ptr(" MEDIUM "), + User: ptr(" user-2 "), + } + + got := chatopenai.ProviderOptionsFromChatConfig( + fakeLanguageModel{provider: fantasyopenai.Name, model: "gpt-4.1"}, + options, + ) + + providerOptions, ok := got.(*fantasyopenai.ResponsesProviderOptions) + require.True(t, ok) + require.Equal(t, []fantasyopenai.IncludeType{ + fantasyopenai.IncludeFileSearchCallResults, + fantasyopenai.IncludeReasoningEncryptedContent, + }, providerOptions.Include) + require.Equal(t, "instructions", requireStringPointerValue(t, providerOptions.Instructions)) + require.Equal(t, int64(5), providerOptions.Logprobs) + require.Same(t, options.MaxToolCalls, providerOptions.MaxToolCalls) + require.Equal(t, options.Metadata, providerOptions.Metadata) + require.Same(t, options.ParallelToolCalls, providerOptions.ParallelToolCalls) + require.Equal(t, "prompt-cache", requireStringPointerValue(t, providerOptions.PromptCacheKey)) + require.Equal(t, fantasyopenai.ReasoningEffortMinimal, requireReasoningEffortPointerValue(t, providerOptions.ReasoningEffort)) + require.Equal(t, "auto", requireStringPointerValue(t, providerOptions.ReasoningSummary)) + require.Equal(t, "safety", requireStringPointerValue(t, providerOptions.SafetyIdentifier)) + require.Equal(t, fantasyopenai.ServiceTierFlex, requireServiceTierPointerValue(t, providerOptions.ServiceTier)) + require.Same(t, options.StrictJSONSchema, providerOptions.StrictJSONSchema) + require.NotNil(t, providerOptions.Store) + require.True(t, *providerOptions.Store) + require.Equal(t, fantasyopenai.TextVerbosityMedium, requireTextVerbosityPointerValue(t, providerOptions.TextVerbosity)) + require.Equal(t, "user-2", requireStringPointerValue(t, providerOptions.User)) +} + +func TestTextVerbosityFromChat(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + value *string + want *fantasyopenai.TextVerbosity + }{ + {name: "Nil"}, + {name: "Empty", value: ptr(" ")}, + {name: "Low", value: ptr(" low "), want: ptr(fantasyopenai.TextVerbosityLow)}, + {name: "MediumCase", value: ptr(" MEDIUM "), want: ptr(fantasyopenai.TextVerbosityMedium)}, + {name: "High", value: ptr("high"), want: ptr(fantasyopenai.TextVerbosityHigh)}, + {name: "Invalid", value: ptr("verbose")}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + got := chatopenai.TextVerbosityFromChat(tt.value) + if tt.want == nil { + require.Nil(t, got) + return + } + require.NotNil(t, got) + require.Equal(t, *tt.want, *got) + }) + } +} + +func TestIncludeFromChat(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + values []string + want []fantasyopenai.IncludeType + }{ + {name: "Nil"}, + {name: "Empty", values: []string{}, want: []fantasyopenai.IncludeType{}}, + { + name: "ValidAndInvalid", + values: []string{ + " " + string(fantasyopenai.IncludeReasoningEncryptedContent) + " ", + string(fantasyopenai.IncludeFileSearchCallResults), + "unsupported", + string(fantasyopenai.IncludeMessageOutputTextLogprobs), + }, + want: []fantasyopenai.IncludeType{ + fantasyopenai.IncludeReasoningEncryptedContent, + fantasyopenai.IncludeFileSearchCallResults, + fantasyopenai.IncludeMessageOutputTextLogprobs, + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + got := chatopenai.IncludeFromChat(tt.values) + require.Equal(t, tt.want, got) + }) + } +} + +func TestEnsureResponseIncludes(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + values []fantasyopenai.IncludeType + want []fantasyopenai.IncludeType + }{ + { + name: "NilAddsRequired", + want: []fantasyopenai.IncludeType{fantasyopenai.IncludeReasoningEncryptedContent}, + }, + { + name: "EmptyAddsRequired", + values: []fantasyopenai.IncludeType{}, + want: []fantasyopenai.IncludeType{fantasyopenai.IncludeReasoningEncryptedContent}, + }, + { + name: "AddsRequiredAfterExistingValues", + values: []fantasyopenai.IncludeType{ + fantasyopenai.IncludeFileSearchCallResults, + }, + want: []fantasyopenai.IncludeType{ + fantasyopenai.IncludeFileSearchCallResults, + fantasyopenai.IncludeReasoningEncryptedContent, + }, + }, + { + name: "DoesNotDuplicateRequired", + values: []fantasyopenai.IncludeType{ + fantasyopenai.IncludeReasoningEncryptedContent, + fantasyopenai.IncludeFileSearchCallResults, + }, + want: []fantasyopenai.IncludeType{ + fantasyopenai.IncludeReasoningEncryptedContent, + fantasyopenai.IncludeFileSearchCallResults, + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + got := chatopenai.EnsureResponseIncludes(tt.values) + require.Equal(t, tt.want, got) + }) + } +} + +func TestUsesResponsesOptions(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + model fantasy.LanguageModel + want bool + }{ + {name: "Nil"}, + { + name: "OpenAIResponsesModel", + model: fakeLanguageModel{provider: fantasyopenai.Name, model: "gpt-4.1"}, + want: true, + }, + { + name: "AzureResponsesModel", + model: fakeLanguageModel{provider: fantasyazure.Name, model: "gpt-4.1"}, + want: true, + }, + { + name: "OpenAINonResponsesModel", + model: fakeLanguageModel{provider: fantasyopenai.Name, model: "gpt-3.5-turbo-instruct"}, + }, + { + name: "NonOpenAIProvider", + model: fakeLanguageModel{provider: "other", model: "gpt-4.1"}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + got := chatopenai.UsesResponsesOptions(tt.model) + require.Equal(t, tt.want, got) + }) + } +} + +func TestReasoningEffortFromChat(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + value *string + want *fantasyopenai.ReasoningEffort + }{ + {name: "Nil"}, + {name: "Empty", value: ptr(" ")}, + {name: "Minimal", value: ptr(" minimal "), want: ptr(fantasyopenai.ReasoningEffortMinimal)}, + {name: "LowCase", value: ptr(" LOW "), want: ptr(fantasyopenai.ReasoningEffortLow)}, + {name: "Medium", value: ptr("medium"), want: ptr(fantasyopenai.ReasoningEffortMedium)}, + {name: "High", value: ptr("high"), want: ptr(fantasyopenai.ReasoningEffortHigh)}, + {name: "XHigh", value: ptr("xhigh"), want: ptr(fantasyopenai.ReasoningEffortXHigh)}, + {name: "NoneUnsupported", value: ptr("none")}, + {name: "Invalid", value: ptr("max")}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + got := chatopenai.ReasoningEffortFromChat(tt.value) + if tt.want == nil { + require.Nil(t, got) + return + } + require.NotNil(t, got) + require.Equal(t, *tt.want, *got) + }) + } +} + +func TestServiceTierFromChat(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + value *string + want *fantasyopenai.ServiceTier + }{ + {name: "Nil"}, + {name: "Empty", value: ptr(" ")}, + {name: "Auto", value: ptr(" auto "), want: ptr(fantasyopenai.ServiceTierAuto)}, + {name: "FlexCase", value: ptr(" FLEX "), want: ptr(fantasyopenai.ServiceTierFlex)}, + {name: "Priority", value: ptr("priority"), want: ptr(fantasyopenai.ServiceTierPriority)}, + {name: "DefaultUnsupported", value: ptr("default")}, + {name: "Invalid", value: ptr("fast")}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + got := chatopenai.ServiceTierFromChat(tt.value) + if tt.want == nil { + require.Nil(t, got) + return + } + require.NotNil(t, got) + require.Equal(t, *tt.want, *got) + }) + } +} + +func TestResponsesLogProbsFromChatConfig(t *testing.T) { + t.Parallel() + + logProbs := true + topLogProbs := int64(4) + tests := []struct { + name string + options *codersdk.ChatModelOpenAIProviderOptions + want any + }{ + {name: "Nil"}, + { + name: "Empty", + options: &codersdk.ChatModelOpenAIProviderOptions{}, + }, + { + name: "LogProbs", + options: &codersdk.ChatModelOpenAIProviderOptions{ + LogProbs: &logProbs, + }, + want: true, + }, + { + name: "TopLogProbs", + options: &codersdk.ChatModelOpenAIProviderOptions{ + TopLogProbs: &topLogProbs, + }, + want: int64(4), + }, + { + name: "TopLogProbsPrecedence", + options: &codersdk.ChatModelOpenAIProviderOptions{ + LogProbs: &logProbs, + TopLogProbs: &topLogProbs, + }, + want: int64(4), + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + got := chatopenai.ResponsesLogProbsFromChatConfig(tt.options) + require.Equal(t, tt.want, got) + }) + } +} + +func TestIsReasoningModel(t *testing.T) { + t.Parallel() + + tests := []struct { + model string + want bool + }{ + {model: ""}, + {model: "o"}, + {model: "o1", want: true}, + {model: "o1-mini", want: true}, + {model: "o3.5", want: true}, + {model: "o10-preview", want: true}, + {model: "oabc"}, + {model: "ox"}, + {model: "o1preview"}, + {model: "gpt-5"}, + {model: "O1"}, + } + + for _, tt := range tests { + t.Run(tt.model, func(t *testing.T) { + t.Parallel() + + got := chatopenai.IsReasoningModel(tt.model) + require.Equal(t, tt.want, got) + }) + } +} + +func requireStringPointerValue(t *testing.T, value *string) string { + t.Helper() + require.NotNil(t, value) + return *value +} + +func requireBoolPointerValue(t *testing.T, value *bool) bool { + t.Helper() + require.NotNil(t, value) + return *value +} + +func requireReasoningEffortPointerValue( + t *testing.T, + value *fantasyopenai.ReasoningEffort, +) fantasyopenai.ReasoningEffort { + t.Helper() + require.NotNil(t, value) + return *value +} + +func requireServiceTierPointerValue( + t *testing.T, + value *fantasyopenai.ServiceTier, +) fantasyopenai.ServiceTier { + t.Helper() + require.NotNil(t, value) + return *value +} + +func requireTextVerbosityPointerValue( + t *testing.T, + value *fantasyopenai.TextVerbosity, +) fantasyopenai.TextVerbosity { + t.Helper() + require.NotNil(t, value) + return *value +} + +func ptr[T any](value T) *T { + return &value +} + +type fakeLanguageModel struct { + provider string + model string +} + +func (fakeLanguageModel) Generate(context.Context, fantasy.Call) (*fantasy.Response, error) { + panic("not implemented") +} + +func (fakeLanguageModel) Stream(context.Context, fantasy.Call) (fantasy.StreamResponse, error) { + panic("not implemented") +} + +func (fakeLanguageModel) GenerateObject(context.Context, fantasy.ObjectCall) (*fantasy.ObjectResponse, error) { + panic("not implemented") +} + +func (fakeLanguageModel) StreamObject(context.Context, fantasy.ObjectCall) (fantasy.ObjectStreamResponse, error) { + panic("not implemented") +} + +func (f fakeLanguageModel) Provider() string { + return f.provider +} + +func (f fakeLanguageModel) Model() string { + return f.model +} diff --git a/coderd/x/chatd/chatopenai/responses.go b/coderd/x/chatd/chatopenai/responses.go new file mode 100644 index 0000000000..2c3cad1b09 --- /dev/null +++ b/coderd/x/chatd/chatopenai/responses.go @@ -0,0 +1,409 @@ +package chatopenai + +import ( + "maps" + "slices" + "strings" + + "charm.land/fantasy" + fantasyopenai "charm.land/fantasy/providers/openai" + "github.com/google/uuid" + + "github.com/coder/coder/v2/coderd/database" + "github.com/coder/coder/v2/coderd/x/chatd/chatprompt" + "github.com/coder/coder/v2/codersdk" +) + +// ChainModeInfo holds the information needed to determine whether a follow-up turn +// can use OpenAI's previous_response_id chaining instead of replaying full +// conversation history. +type ChainModeInfo struct { + // previousResponseID is the provider response ID from the last assistant + // message, if any. + previousResponseID string + // modelConfigID is the model configuration used to produce the assistant + // message referenced by previousResponseID. + modelConfigID uuid.UUID + // 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 + // providerMissingToolResults is true when the assistant message has local + // tool calls with local results, but no follow-up assistant message exists to + // confirm the results were sent back to the provider. This happens when + // StopAfterTool terminates a turn before the results are round-tripped. + providerMissingToolResults bool +} + +// PreviousResponseID returns the provider response ID from the last assistant +// message, if any. +func (c ChainModeInfo) PreviousResponseID() string { + return c.previousResponseID +} + +// ModelConfigID returns the model configuration used to produce the assistant +// message referenced by PreviousResponseID. +func (c ChainModeInfo) ModelConfigID() uuid.UUID { + return c.modelConfigID +} + +// ContributingTrailingUserCount returns the number of trailing user messages +// that materially change the provider input. +func (c ChainModeInfo) ContributingTrailingUserCount() int { + return c.contributingTrailingUserCount +} + +// HasUnresolvedLocalToolCalls reports whether PreviousResponseID points at an +// assistant message with pending local tool calls. +func (c ChainModeInfo) HasUnresolvedLocalToolCalls() bool { + return c.hasUnresolvedLocalToolCalls +} + +// ProviderMissingToolResults reports whether PreviousResponseID points at an +// assistant message with local tool results, but no follow-up assistant message +// confirms those tool results were sent to the provider (not just persisted +// locally). +func (c ChainModeInfo) ProviderMissingToolResults() bool { + return c.providerMissingToolResults +} + +// IsResponsesStoreEnabled checks if the OpenAI Responses provider options are +// present and have Store set to true. When true, the provider stores +// conversation history server-side, enabling follow-up chaining via +// PreviousResponseID. +func IsResponsesStoreEnabled(opts fantasy.ProviderOptions) bool { + if opts == nil { + return false + } + raw, ok := opts[fantasyopenai.Name] + if !ok { + return false + } + respOpts, ok := raw.(*fantasyopenai.ResponsesProviderOptions) + if !ok || respOpts == nil { + return false + } + return respOpts.Store != nil && *respOpts.Store +} + +// WithPreviousResponseID shallow-clones the provider options map and the OpenAI +// Responses entry, setting PreviousResponseID on the clone. The original map +// and entry are not mutated. +func WithPreviousResponseID( + opts fantasy.ProviderOptions, + previousResponseID string, +) fantasy.ProviderOptions { + cloned := maps.Clone(opts) + if cloned == nil { + cloned = fantasy.ProviderOptions{} + } + if raw, ok := cloned[fantasyopenai.Name]; ok { + if respOpts, ok := raw.(*fantasyopenai.ResponsesProviderOptions); ok && respOpts != nil { + clone := *respOpts + clone.PreviousResponseID = &previousResponseID + cloned[fantasyopenai.Name] = &clone + } + } + return cloned +} + +// HasPreviousResponseID checks whether the provider options contain an OpenAI +// Responses entry with a non-empty PreviousResponseID. +func HasPreviousResponseID(providerOptions fantasy.ProviderOptions) bool { + if len(providerOptions) == 0 { + return false + } + + entry, ok := providerOptions[fantasyopenai.Name] + if !ok { + return false + } + options, ok := entry.(*fantasyopenai.ResponsesProviderOptions) + return ok && options != nil && options.PreviousResponseID != nil && + *options.PreviousResponseID != "" +} + +// ClearPreviousResponseID returns a clone of providerOptions with +// PreviousResponseID cleared on the OpenAI Responses options. The original +// providerOptions is not modified. +func ClearPreviousResponseID(providerOptions fantasy.ProviderOptions) fantasy.ProviderOptions { + cloned := maps.Clone(providerOptions) + if cloned == nil { + return fantasy.ProviderOptions{} + } + + entry, ok := cloned[fantasyopenai.Name] + if !ok { + return cloned + } + options, ok := entry.(*fantasyopenai.ResponsesProviderOptions) + if !ok || options == nil { + return cloned + } + optionsClone := *options + optionsClone.PreviousResponseID = nil + cloned[fantasyopenai.Name] = &optionsClone + return cloned +} + +// extractResponseID extracts the OpenAI Responses API response ID from provider +// metadata. Returns an empty string if no OpenAI Responses metadata is present. +func extractResponseID(metadata fantasy.ProviderMetadata) string { + if len(metadata) == 0 { + return "" + } + + entry, ok := metadata[fantasyopenai.Name] + if !ok { + return "" + } + providerMetadata, ok := entry.(*fantasyopenai.ResponsesProviderMetadata) + if !ok || providerMetadata == nil { + return "" + } + return providerMetadata.ResponseID +} + +// ExtractResponseIDIfStored returns the OpenAI response ID only when the +// provider options indicate store=true. Response IDs from store=false turns are +// not persisted server-side and cannot be used for chaining. +func ExtractResponseIDIfStored( + providerOptions fantasy.ProviderOptions, + metadata fantasy.ProviderMetadata, +) string { + if !IsResponsesStoreEnabled(providerOptions) { + return "" + } + + return extractResponseID(metadata) +} + +// 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, +// complete local tool state, and confirmation that tool results were sent to +// the provider. +func ShouldActivateChainMode( + providerOptions fantasy.ProviderOptions, + info ChainModeInfo, + modelConfigID uuid.UUID, + isPlanModeTurn bool, +) bool { + return IsResponsesStoreEnabled(providerOptions) && + info.previousResponseID != "" && + info.contributingTrailingUserCount > 0 && + info.modelConfigID == modelConfigID && + !isPlanModeTurn && + !info.hasUnresolvedLocalToolCalls && + !info.providerMissingToolResults +} + +// ResolveChainMode scans DB messages from the end to inspect the current +// trailing user turn and detect whether the immediately preceding assistant/tool +// block can chain from a provider response ID. +func ResolveChainMode(messages []database.ChatMessage) ChainModeInfo { + var info ChainModeInfo + i := len(messages) - 1 + for ; i >= 0; i-- { + if messages[i].Role != database.ChatMessageRoleUser { + break + } + if userMessageContributesToChainMode(messages[i]) { + info.contributingTrailingUserCount++ + } + } + for ; i >= 0; i-- { + switch messages[i].Role { + case database.ChatMessageRoleAssistant: + if messages[i].ProviderResponseID.Valid && + messages[i].ProviderResponseID.String != "" { + info.previousResponseID = messages[i].ProviderResponseID.String + if messages[i].ModelConfigID.Valid { + info.modelConfigID = messages[i].ModelConfigID.UUID + } + info.hasUnresolvedLocalToolCalls = assistantHasUnresolvedLocalToolCalls(messages, i) + if !info.hasUnresolvedLocalToolCalls { + info.providerMissingToolResults = providerHasMissingToolResults(messages, i) + } + return info + } + return info + case database.ChatMessageRoleTool: + continue + default: + return info + } + } + return info +} + +// FilterPromptForChainMode keeps only system messages and the trailing user +// messages that still contribute model-visible content to the current turn. +// Assistant and tool messages are dropped because the provider already has +// them via the previous_response_id chain. +func FilterPromptForChainMode( + prompt []fantasy.Message, + info ChainModeInfo, +) []fantasy.Message { + if info.contributingTrailingUserCount <= 0 { + return prompt + } + + totalUsers := 0 + for _, msg := range prompt { + if msg.Role == "user" { + totalUsers++ + } + } + + // Prompt construction already drops user turns with no model-visible + // content, such as skill-only sentinel messages. That means the user + // count here stays aligned with contributingTrailingUserCount even + // when non-contributing DB turns are interleaved in the trailing + // block. + usersToSkip := totalUsers - info.contributingTrailingUserCount + if usersToSkip < 0 { + usersToSkip = 0 + } + + filtered := make([]fantasy.Message, 0, len(prompt)) + usersSeen := 0 + for _, msg := range prompt { + switch msg.Role { + case "system": + filtered = append(filtered, msg) + case "user": + usersSeen++ + if usersSeen > usersToSkip { + filtered = append(filtered, msg) + } + } + } + + return filtered +} + +func userMessageContributesToChainMode(msg database.ChatMessage) bool { + parts, err := chatprompt.ParseContent(msg) + if err != nil { + return false + } + for _, part := range parts { + switch part.Type { + case codersdk.ChatMessagePartTypeText, + codersdk.ChatMessagePartTypeReasoning: + if strings.TrimSpace(part.Text) != "" { + return true + } + case codersdk.ChatMessagePartTypeFile, + codersdk.ChatMessagePartTypeFileReference: + return true + case codersdk.ChatMessagePartTypeContextFile: + if part.ContextFileContent != "" { + return true + } + } + } + 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) +} + +// providerHasMissingToolResults reports whether the assistant message at +// assistantIdx has local tool calls whose results exist in the database but +// were never sent back to the provider. This is detected by the absence of a +// follow-up assistant message after the tool results. In normal flow the LLM +// processes tool results and produces a follow-up response, but StopAfterTool +// skips that round-trip. +func providerHasMissingToolResults( + messages []database.ChatMessage, + assistantIdx int, +) bool { + if assistantIdx < 0 || assistantIdx >= len(messages) { + return false + } + + parts, err := chatprompt.ParseContent(messages[assistantIdx]) + if err != nil { + // Parsing errors are already handled by + // assistantHasUnresolvedLocalToolCalls. + return false + } + + if !slices.ContainsFunc(parts, func(p codersdk.ChatMessagePart) bool { + return p.Type == codersdk.ChatMessagePartTypeToolCall && !p.ProviderExecuted + }) { + return false + } + + // Scan forward past tool messages. If the first non-tool message is not an + // assistant, the tool results were never round-tripped to the provider. + for i := assistantIdx + 1; i < len(messages); i++ { + switch messages[i].Role { + case database.ChatMessageRoleTool: + continue + case database.ChatMessageRoleAssistant: + // A follow-up assistant exists, so results were sent. + return false + default: + // User or system message with no follow-up assistant. + return true + } + } + + // Reached end of messages without a follow-up assistant. + return true +} diff --git a/coderd/x/chatd/chatopenai/responses_test.go b/coderd/x/chatd/chatopenai/responses_test.go new file mode 100644 index 0000000000..5a6e3b9596 --- /dev/null +++ b/coderd/x/chatd/chatopenai/responses_test.go @@ -0,0 +1,993 @@ +package chatopenai_test + +import ( + "database/sql" + "encoding/json" + "testing" + + "charm.land/fantasy" + fantasyopenai "charm.land/fantasy/providers/openai" + "github.com/google/uuid" + "github.com/sqlc-dev/pqtype" + "github.com/stretchr/testify/require" + + "github.com/coder/coder/v2/coderd/database" + "github.com/coder/coder/v2/coderd/x/chatd/chatopenai" + "github.com/coder/coder/v2/coderd/x/chatd/chatprompt" + "github.com/coder/coder/v2/coderd/x/chatd/chattest" + "github.com/coder/coder/v2/codersdk" +) + +func TestIsResponsesStoreEnabled(t *testing.T) { + t.Parallel() + + storeTrue := true + storeFalse := false + + tests := []struct { + name string + opts fantasy.ProviderOptions + want bool + }{ + { + name: "NilOptions", + }, + { + name: "NonOpenAIKeysOnly", + opts: fantasy.ProviderOptions{ + "other": &fantasyopenai.ProviderOptions{}, + }, + }, + { + name: "OpenAIKeyWithNonResponsesOptions", + opts: fantasy.ProviderOptions{ + fantasyopenai.Name: &fantasyopenai.ProviderOptions{}, + }, + }, + { + name: "OpenAIKeyWithNilStore", + opts: fantasy.ProviderOptions{ + fantasyopenai.Name: &fantasyopenai.ResponsesProviderOptions{}, + }, + }, + { + name: "OpenAIKeyWithFalseStore", + opts: fantasy.ProviderOptions{ + fantasyopenai.Name: &fantasyopenai.ResponsesProviderOptions{Store: &storeFalse}, + }, + }, + { + name: "OpenAIKeyWithTrueStore", + opts: fantasy.ProviderOptions{ + fantasyopenai.Name: &fantasyopenai.ResponsesProviderOptions{Store: &storeTrue}, + }, + want: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + got := chatopenai.IsResponsesStoreEnabled(tt.opts) + require.Equal(t, tt.want, got) + }) + } +} + +func TestIsResponsesStoreEnabledIgnoresMalformedNonOpenAIKey(t *testing.T) { + t.Parallel() + + store := true + // This intentionally documents the only synthetic mismatch from the old + // chatloop value scan: a malformed map with OpenAI Responses options under a + // non-OpenAI key is not treated as enabled. + opts := fantasy.ProviderOptions{ + "not-openai": &fantasyopenai.ResponsesProviderOptions{Store: &store}, + } + + require.False(t, chatopenai.IsResponsesStoreEnabled(opts)) +} + +func TestShouldActivateChainMode(t *testing.T) { + t.Parallel() + + modelConfigID := uuid.New() + baseInfo := chatopenai.ResolveChainMode([]database.ChatMessage{ + chainModeSystemMessage(), + chainModeUserMessage("prior user message"), + chainModeAssistantMessage(modelConfigID, nil), + chainModeUserMessage("latest user message"), + }) + + localCall := codersdk.ChatMessageToolCall( + "call-local", + "read_file", + json.RawMessage(`{"path":"main.go"}`), + ) + unresolvedLocalInfo := chatopenai.ResolveChainMode([]database.ChatMessage{ + chainModeSystemMessage(), + chainModeUserMessage("prior user message"), + chainModeAssistantMessage(modelConfigID, []codersdk.ChatMessagePart{localCall}), + chainModeUserMessage("latest user message"), + }) + localResult := codersdk.ChatMessageToolResult( + "call-local", + "read_file", + json.RawMessage(`{"ok":true}`), + false, + false, + ) + missingToolResultsInfo := chatopenai.ResolveChainMode([]database.ChatMessage{ + chainModeSystemMessage(), + chainModeUserMessage("prior user message"), + chainModeAssistantMessage(modelConfigID, []codersdk.ChatMessagePart{localCall}), + chainModeToolMessage([]codersdk.ChatMessagePart{localResult}), + chainModeUserMessage("latest user message"), + }) + skillOnlyInfo := chatopenai.ResolveChainMode([]database.ChatMessage{ + chainModeSystemMessage(), + chainModeUserMessage("prior user message"), + chainModeAssistantMessage(modelConfigID, nil), + chainModeSkillOnlyUserMessage(), + }) + missingResponseInfo := chatopenai.ResolveChainMode([]database.ChatMessage{ + chainModeSystemMessage(), + chainModeUserMessage("prior user message"), + chainModeAssistantMessageWithoutResponse(modelConfigID), + chainModeUserMessage("latest user message"), + }) + + tests := []struct { + name string + providerOpts fantasy.ProviderOptions + info chatopenai.ChainModeInfo + modelConfigID uuid.UUID + isPlanModeTurn bool + want bool + }{ + { + name: "StoreDisabled", + providerOpts: chainModeProviderOptions(false), + info: baseInfo, + modelConfigID: modelConfigID, + }, + { + name: "MissingPreviousResponseID", + providerOpts: chainModeProviderOptions(true), + info: missingResponseInfo, + modelConfigID: modelConfigID, + }, + { + name: "MismatchedModelConfigID", + providerOpts: chainModeProviderOptions(true), + info: baseInfo, + modelConfigID: uuid.New(), + }, + { + name: "PlanMode", + providerOpts: chainModeProviderOptions(true), + info: baseInfo, + modelConfigID: modelConfigID, + isPlanModeTurn: true, + }, + { + name: "NoContributingTrailingUser", + providerOpts: chainModeProviderOptions(true), + info: skillOnlyInfo, + modelConfigID: modelConfigID, + }, + { + name: "UnresolvedLocalToolCalls", + providerOpts: chainModeProviderOptions(true), + info: unresolvedLocalInfo, + modelConfigID: modelConfigID, + }, + { + name: "ProviderMissingToolResults", + providerOpts: chainModeProviderOptions(true), + info: missingToolResultsInfo, + modelConfigID: modelConfigID, + }, + { + name: "AllConditionsMet", + providerOpts: chainModeProviderOptions(true), + info: baseInfo, + modelConfigID: modelConfigID, + want: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + got := chatopenai.ShouldActivateChainMode( + tt.providerOpts, + tt.info, + tt.modelConfigID, + tt.isPlanModeTurn, + ) + require.Equal(t, tt.want, got) + }) + } +} + +func TestWithPreviousResponseID(t *testing.T) { + t.Parallel() + + store := true + originalResponses := &fantasyopenai.ResponsesProviderOptions{Store: &store} + otherOptions := &fantasyopenai.ProviderOptions{} + opts := fantasy.ProviderOptions{ + fantasyopenai.Name: originalResponses, + "other": otherOptions, + } + + got := chatopenai.WithPreviousResponseID(opts, "resp-next") + + gotOtherOptions, ok := got["other"].(*fantasyopenai.ProviderOptions) + require.True(t, ok) + require.True(t, otherOptions == gotOtherOptions) + gotOriginalResponses, ok := opts[fantasyopenai.Name].(*fantasyopenai.ResponsesProviderOptions) + require.True(t, ok) + require.True(t, originalResponses == gotOriginalResponses) + require.Nil(t, originalResponses.PreviousResponseID) + + clonedResponses, ok := got[fantasyopenai.Name].(*fantasyopenai.ResponsesProviderOptions) + require.True(t, ok) + require.NotSame(t, originalResponses, clonedResponses) + require.NotNil(t, clonedResponses.PreviousResponseID) + require.Equal(t, "resp-next", *clonedResponses.PreviousResponseID) + require.True(t, originalResponses.Store == clonedResponses.Store) + + got["new"] = otherOptions + require.NotContains(t, opts, "new") +} + +func TestWithPreviousResponseIDNilInput(t *testing.T) { + t.Parallel() + + got := chatopenai.WithPreviousResponseID(nil, "resp-next") + + require.NotNil(t, got) + require.Empty(t, got) +} + +func TestHasPreviousResponseID(t *testing.T) { + t.Parallel() + + emptyID := "" + responseID := "resp-123" + + tests := []struct { + name string + opts fantasy.ProviderOptions + want bool + }{ + { + name: "NilOptions", + }, + { + name: "EmptyID", + opts: fantasy.ProviderOptions{ + fantasyopenai.Name: &fantasyopenai.ResponsesProviderOptions{ + PreviousResponseID: &emptyID, + }, + }, + }, + { + name: "NonEmptyID", + opts: fantasy.ProviderOptions{ + fantasyopenai.Name: &fantasyopenai.ResponsesProviderOptions{ + PreviousResponseID: &responseID, + }, + }, + want: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + got := chatopenai.HasPreviousResponseID(tt.opts) + require.Equal(t, tt.want, got) + }) + } +} + +func TestClearPreviousResponseID(t *testing.T) { + t.Parallel() + + responseID := "resp-123" + options := &fantasyopenai.ResponsesProviderOptions{ + PreviousResponseID: &responseID, + } + otherOptions := &fantasyopenai.ProviderOptions{} + opts := fantasy.ProviderOptions{ + fantasyopenai.Name: options, + "other": otherOptions, + } + + got := chatopenai.ClearPreviousResponseID(opts) + + got["new"] = otherOptions + require.NotContains(t, opts, "new") + require.NotNil(t, options.PreviousResponseID) + require.Equal(t, "resp-123", *options.PreviousResponseID) + + gotOtherOptions, ok := got["other"].(*fantasyopenai.ProviderOptions) + require.True(t, ok) + require.True(t, otherOptions == gotOtherOptions) + clonedOptions, ok := got[fantasyopenai.Name].(*fantasyopenai.ResponsesProviderOptions) + require.True(t, ok) + require.NotSame(t, options, clonedOptions) + require.Nil(t, clonedOptions.PreviousResponseID) + + require.NotPanics(t, func() { + got := chatopenai.ClearPreviousResponseID(nil) + require.NotNil(t, got) + chatopenai.ClearPreviousResponseID(fantasy.ProviderOptions{ + fantasyopenai.Name: &fantasyopenai.ProviderOptions{}, + }) + }) +} + +func TestExtractResponseIDIfStoredMetadata(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + metadata fantasy.ProviderMetadata + want string + }{ + { + name: "NilMetadata", + }, + { + name: "NoResponsesMetadata", + metadata: fantasy.ProviderMetadata{ + "other": &fantasyopenai.ProviderOptions{}, + }, + }, + { + name: "ResponsesMetadataUnderNonOpenAIKey", + metadata: fantasy.ProviderMetadata{ + "other": &fantasyopenai.ResponsesProviderMetadata{ + ResponseID: "resp-123", + }, + }, + }, + { + name: "ResponsesMetadata", + metadata: fantasy.ProviderMetadata{ + fantasyopenai.Name: &fantasyopenai.ResponsesProviderMetadata{ + ResponseID: "resp-123", + }, + }, + want: "resp-123", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + got := chatopenai.ExtractResponseIDIfStored( + chainModeProviderOptions(true), + tt.metadata, + ) + require.Equal(t, tt.want, got) + }) + } +} + +func TestExtractResponseIDIfStored(t *testing.T) { + t.Parallel() + + metadata := fantasy.ProviderMetadata{ + fantasyopenai.Name: &fantasyopenai.ResponsesProviderMetadata{ + ResponseID: "resp-123", + }, + } + + require.Empty(t, chatopenai.ExtractResponseIDIfStored( + chainModeProviderOptions(false), + metadata, + )) + require.Equal(t, "resp-123", chatopenai.ExtractResponseIDIfStored( + chainModeProviderOptions(true), + metadata, + )) +} + +func TestResolveChainModeIgnoresSkillOnlySentinelMessages(t *testing.T) { + t.Parallel() + + modelConfigID := uuid.New() + assistant := database.ChatMessage{ + Role: database.ChatMessageRoleAssistant, + ProviderResponseID: sql.NullString{String: "resp-123", Valid: true}, + ModelConfigID: uuid.NullUUID{UUID: modelConfigID, Valid: true}, + } + skillOnly := chainModeSkillOnlyUserMessage() + user := chattest.ChatMessageWithParts([]codersdk.ChatMessagePart{{ + Type: codersdk.ChatMessagePartTypeText, + Text: "latest user message", + }}) + user.Role = database.ChatMessageRoleUser + + got := chatopenai.ResolveChainMode([]database.ChatMessage{assistant, skillOnly, user}) + require.Equal(t, "resp-123", got.PreviousResponseID()) + require.Equal(t, modelConfigID, got.ModelConfigID()) + 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 := chatopenai.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, chatopenai.ShouldActivateChainMode( + chainModeProviderOptions(true), + chainInfo, + modelConfigID, + false, + )) +} + +func TestResolveChainMode_BlocksWhenAssistantContentCannotParse(t *testing.T) { + t.Parallel() + + modelConfigID := uuid.New() + chainInfo := chatopenai.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, chatopenai.ShouldActivateChainMode( + chainModeProviderOptions(true), + 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 := chatopenai.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, chatopenai.ShouldActivateChainMode( + chainModeProviderOptions(true), + 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 := chatopenai.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.False(t, chainInfo.ProviderMissingToolResults()) + require.True(t, chatopenai.ShouldActivateChainMode( + chainModeProviderOptions(true), + 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 := chatopenai.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, chatopenai.ShouldActivateChainMode( + chainModeProviderOptions(true), + 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, + ) + followUp := chainModeAssistantMessage(modelConfigID, nil) + followUp.ProviderResponseID = sql.NullString{String: "resp-follow-up", Valid: true} + + chainInfo := chatopenai.ResolveChainMode([]database.ChatMessage{ + chainModeSystemMessage(), + chainModeUserMessage("prior user message"), + chainModeAssistantMessage(modelConfigID, []codersdk.ChatMessagePart{toolCall}), + chainModeToolMessage([]codersdk.ChatMessagePart{toolResult}), + followUp, + chainModeUserMessage("latest user message"), + }) + + require.Equal(t, "resp-follow-up", chainInfo.PreviousResponseID()) + require.False(t, chainInfo.HasUnresolvedLocalToolCalls()) + require.False(t, chainInfo.ProviderMissingToolResults()) + require.True(t, chatopenai.ShouldActivateChainMode( + chainModeProviderOptions(true), + 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 := chatopenai.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, chatopenai.ShouldActivateChainMode( + chainModeProviderOptions(true), + chainInfo, + modelConfigID, + false, + )) +} + +func TestResolveChainMode_BlocksWhenToolResultNeverSentToProvider(t *testing.T) { + t.Parallel() + + modelConfigID := uuid.New() + toolCall := codersdk.ChatMessageToolCall( + "call-local", + "propose_plan", + json.RawMessage(`{"path":"plan.md"}`), + ) + toolResult := codersdk.ChatMessageToolResult( + "call-local", + "propose_plan", + json.RawMessage(`{"ok":true}`), + false, + false, + ) + + chainInfo := chatopenai.ResolveChainMode([]database.ChatMessage{ + chainModeSystemMessage(), + chainModeUserMessage("make a plan"), + chainModeAssistantMessage(modelConfigID, []codersdk.ChatMessagePart{toolCall}), + chainModeToolMessage([]codersdk.ChatMessagePart{toolResult}), + chainModeUserMessage("implement the plan"), + }) + + require.Equal(t, "resp-123", chainInfo.PreviousResponseID()) + require.False(t, chainInfo.HasUnresolvedLocalToolCalls()) + require.True(t, chainInfo.ProviderMissingToolResults()) + require.False(t, chatopenai.ShouldActivateChainMode( + chainModeProviderOptions(true), + chainInfo, + modelConfigID, + false, + )) +} + +func TestResolveChainMode_BlocksProviderMissingWithMultipleToolCalls(t *testing.T) { + t.Parallel() + + modelConfigID := uuid.New() + call1 := codersdk.ChatMessageToolCall( + "call-1", + "propose_plan", + json.RawMessage(`{"path":"plan.md"}`), + ) + call2 := codersdk.ChatMessageToolCall( + "call-2", + "write_file", + json.RawMessage(`{"path":"foo.go"}`), + ) + result1 := codersdk.ChatMessageToolResult( + "call-1", + "propose_plan", + json.RawMessage(`{"ok":true}`), + false, + false, + ) + result2 := codersdk.ChatMessageToolResult( + "call-2", + "write_file", + json.RawMessage(`{"ok":true}`), + false, + false, + ) + + chainInfo := chatopenai.ResolveChainMode([]database.ChatMessage{ + chainModeSystemMessage(), + chainModeUserMessage("do it"), + chainModeAssistantMessage(modelConfigID, []codersdk.ChatMessagePart{call1, call2}), + chainModeToolMessage([]codersdk.ChatMessagePart{result1, result2}), + chainModeUserMessage("next"), + }) + + require.False(t, chainInfo.HasUnresolvedLocalToolCalls()) + require.True(t, chainInfo.ProviderMissingToolResults()) + require.False(t, chatopenai.ShouldActivateChainMode( + chainModeProviderOptions(true), + chainInfo, + modelConfigID, + false, + )) +} + +func TestResolveChainMode_AllowsWhenNoToolCalls(t *testing.T) { + t.Parallel() + + modelConfigID := uuid.New() + + chainInfo := chatopenai.ResolveChainMode([]database.ChatMessage{ + chainModeSystemMessage(), + chainModeUserMessage("hello"), + chainModeAssistantMessage(modelConfigID, nil), + chainModeUserMessage("thanks"), + }) + + require.Equal(t, "resp-123", chainInfo.PreviousResponseID()) + require.False(t, chainInfo.HasUnresolvedLocalToolCalls()) + require.False(t, chainInfo.ProviderMissingToolResults()) + require.True(t, chatopenai.ShouldActivateChainMode( + chainModeProviderOptions(true), + chainInfo, + modelConfigID, + false, + )) +} + +func TestFilterPromptForChainModeKeepsContributingUsersAcrossSkippedSentinelTurns(t *testing.T) { + t.Parallel() + + modelConfigID := uuid.New() + priorUser := chattest.ChatMessageWithParts([]codersdk.ChatMessagePart{{ + Type: codersdk.ChatMessagePartTypeText, + Text: "prior user message", + }}) + priorUser.Role = database.ChatMessageRoleUser + assistant := database.ChatMessage{ + Role: database.ChatMessageRoleAssistant, + ProviderResponseID: sql.NullString{String: "resp-123", Valid: true}, + ModelConfigID: uuid.NullUUID{UUID: modelConfigID, Valid: true}, + } + firstTrailingUser := chattest.ChatMessageWithParts([]codersdk.ChatMessagePart{{ + Type: codersdk.ChatMessagePartTypeText, + Text: "first trailing user", + }}) + firstTrailingUser.Role = database.ChatMessageRoleUser + skillOnly := chainModeSkillOnlyUserMessage() + lastTrailingUser := chattest.ChatMessageWithParts([]codersdk.ChatMessagePart{{ + Type: codersdk.ChatMessagePartTypeText, + Text: "last trailing user", + }}) + lastTrailingUser.Role = database.ChatMessageRoleUser + + chainInfo := chatopenai.ResolveChainMode([]database.ChatMessage{ + priorUser, + assistant, + firstTrailingUser, + skillOnly, + lastTrailingUser, + }) + require.Equal(t, 2, chainInfo.ContributingTrailingUserCount()) + + prompt := []fantasy.Message{ + { + Role: fantasy.MessageRoleSystem, + Content: []fantasy.MessagePart{ + fantasy.TextPart{Text: "system instruction"}, + }, + }, + { + Role: fantasy.MessageRoleUser, + Content: []fantasy.MessagePart{ + fantasy.TextPart{Text: "prior user message"}, + }, + }, + { + Role: fantasy.MessageRoleAssistant, + Content: []fantasy.MessagePart{ + fantasy.TextPart{Text: "assistant reply"}, + }, + }, + { + Role: fantasy.MessageRoleUser, + Content: []fantasy.MessagePart{ + fantasy.TextPart{Text: "first trailing user"}, + }, + }, + { + Role: fantasy.MessageRoleUser, + Content: []fantasy.MessagePart{ + fantasy.TextPart{Text: "last trailing user"}, + }, + }, + } + + got := chatopenai.FilterPromptForChainMode(prompt, chainInfo) + require.Len(t, got, 3) + require.Equal(t, fantasy.MessageRoleSystem, got[0].Role) + require.Equal(t, fantasy.MessageRoleUser, got[1].Role) + require.Equal(t, fantasy.MessageRoleUser, got[2].Role) + + firstPart, ok := fantasy.AsMessagePart[fantasy.TextPart](got[1].Content[0]) + require.True(t, ok) + require.Equal(t, "first trailing user", firstPart.Text) + lastPart, ok := fantasy.AsMessagePart[fantasy.TextPart](got[2].Content[0]) + require.True(t, ok) + require.Equal(t, "last trailing user", lastPart.Text) +} + +func TestFilterPromptForChainModeUsesContributingTrailingUsers(t *testing.T) { + t.Parallel() + + modelConfigID := uuid.New() + priorUser := chattest.ChatMessageWithParts([]codersdk.ChatMessagePart{{ + Type: codersdk.ChatMessagePartTypeText, + Text: "prior user message", + }}) + priorUser.Role = database.ChatMessageRoleUser + assistant := database.ChatMessage{ + Role: database.ChatMessageRoleAssistant, + ProviderResponseID: sql.NullString{String: "resp-123", Valid: true}, + ModelConfigID: uuid.NullUUID{UUID: modelConfigID, Valid: true}, + } + skillOnly := chainModeSkillOnlyUserMessage() + latestUser := chattest.ChatMessageWithParts([]codersdk.ChatMessagePart{{ + Type: codersdk.ChatMessagePartTypeText, + Text: "latest user message", + }}) + latestUser.Role = database.ChatMessageRoleUser + + chainInfo := chatopenai.ResolveChainMode([]database.ChatMessage{ + priorUser, + assistant, + skillOnly, + latestUser, + }) + require.Equal(t, 1, chainInfo.ContributingTrailingUserCount()) + + prompt := []fantasy.Message{ + { + Role: fantasy.MessageRoleSystem, + Content: []fantasy.MessagePart{ + fantasy.TextPart{Text: "system instruction"}, + }, + }, + { + Role: fantasy.MessageRoleUser, + Content: []fantasy.MessagePart{ + fantasy.TextPart{Text: "prior user message"}, + }, + }, + { + Role: fantasy.MessageRoleAssistant, + Content: []fantasy.MessagePart{ + fantasy.TextPart{Text: "assistant reply"}, + }, + }, + { + Role: fantasy.MessageRoleUser, + Content: []fantasy.MessagePart{ + fantasy.TextPart{Text: "latest user message"}, + }, + }, + } + + got := chatopenai.FilterPromptForChainMode(prompt, chainInfo) + require.Len(t, got, 2) + require.Equal(t, fantasy.MessageRoleSystem, got[0].Role) + require.Equal(t, fantasy.MessageRoleUser, got[1].Role) + + part, ok := fantasy.AsMessagePart[fantasy.TextPart](got[1].Content[0]) + require.True(t, ok) + require.Equal(t, "latest user message", part.Text) +} + +func chainModeProviderOptions(store bool) fantasy.ProviderOptions { + 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 := chattest.ChatMessageWithParts([]codersdk.ChatMessagePart{ + codersdk.ChatMessageText(text), + }) + msg.Role = database.ChatMessageRoleUser + return msg +} + +func chainModeSkillOnlyUserMessage() database.ChatMessage { + msg := chattest.ChatMessageWithParts([]codersdk.ChatMessagePart{ + { + Type: codersdk.ChatMessagePartTypeContextFile, + // Keep this in sync with chatd.AgentChatContextSentinelPath. + ContextFilePath: ".coder/agent-chat-context-sentinel", + ContextFileAgentID: uuid.NullUUID{ + UUID: uuid.New(), + Valid: true, + }, + }, + { + Type: codersdk.ChatMessagePartTypeSkill, + SkillName: "repo-helper", + SkillDir: "/skills/repo-helper", + }, + }) + msg.Role = database.ChatMessageRoleUser + return msg +} + +func chainModeAssistantMessage( + modelConfigID uuid.UUID, + parts []codersdk.ChatMessagePart, +) database.ChatMessage { + msg := chattest.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 chainModeAssistantMessageWithoutResponse( + modelConfigID uuid.UUID, +) database.ChatMessage { + msg := chattest.ChatMessageWithParts(nil) + msg.Role = database.ChatMessageRoleAssistant + 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 := chattest.ChatMessageWithParts(parts) + msg.Role = database.ChatMessageRoleTool + return msg +} diff --git a/coderd/x/chatd/chatopenai/tools.go b/coderd/x/chatd/chatopenai/tools.go new file mode 100644 index 0000000000..325463c435 --- /dev/null +++ b/coderd/x/chatd/chatopenai/tools.go @@ -0,0 +1,29 @@ +package chatopenai + +import ( + "charm.land/fantasy" + + "github.com/coder/coder/v2/codersdk" +) + +// WebSearchTool returns the OpenAI provider-native web search tool when +// enabled by the model provider options. +func WebSearchTool(options *codersdk.ChatModelOpenAIProviderOptions) (fantasy.Tool, bool) { + if options == nil || options.WebSearchEnabled == nil || !*options.WebSearchEnabled { + return nil, false + } + + args := map[string]any{} + if options.SearchContextSize != nil && *options.SearchContextSize != "" { + args["search_context_size"] = *options.SearchContextSize + } + if len(options.AllowedDomains) > 0 { + args["allowed_domains"] = options.AllowedDomains + } + + return fantasy.ProviderDefinedTool{ + ID: "web_search", + Name: "web_search", + Args: args, + }, true +} diff --git a/coderd/x/chatd/chatopenai/tools_test.go b/coderd/x/chatd/chatopenai/tools_test.go new file mode 100644 index 0000000000..b8be793419 --- /dev/null +++ b/coderd/x/chatd/chatopenai/tools_test.go @@ -0,0 +1,116 @@ +package chatopenai_test + +import ( + "testing" + + "charm.land/fantasy" + "github.com/stretchr/testify/require" + + "github.com/coder/coder/v2/coderd/x/chatd/chatopenai" + "github.com/coder/coder/v2/codersdk" +) + +func TestWebSearchToolDisabled(t *testing.T) { + t.Parallel() + + disabled := false + + tests := []struct { + name string + options *codersdk.ChatModelOpenAIProviderOptions + }{ + { + name: "NilOptions", + }, + { + name: "NilWebSearchEnabled", + options: &codersdk.ChatModelOpenAIProviderOptions{}, + }, + { + name: "WebSearchDisabled", + options: &codersdk.ChatModelOpenAIProviderOptions{ + WebSearchEnabled: &disabled, + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + tool, ok := chatopenai.WebSearchTool(tt.options) + require.False(t, ok) + require.Nil(t, tool) + }) + } +} + +func TestWebSearchTool(t *testing.T) { + t.Parallel() + + enabled := true + searchContextSize := "high" + allowedDomains := []string{"example.com", "coder.com"} + + tests := []struct { + name string + options *codersdk.ChatModelOpenAIProviderOptions + want map[string]any + }{ + { + name: "NoExtraFields", + options: &codersdk.ChatModelOpenAIProviderOptions{ + WebSearchEnabled: &enabled, + }, + want: map[string]any{}, + }, + { + name: "SearchContextSize", + options: &codersdk.ChatModelOpenAIProviderOptions{ + WebSearchEnabled: &enabled, + SearchContextSize: &searchContextSize, + }, + want: map[string]any{ + "search_context_size": searchContextSize, + }, + }, + { + name: "AllowedDomains", + options: &codersdk.ChatModelOpenAIProviderOptions{ + WebSearchEnabled: &enabled, + AllowedDomains: allowedDomains, + }, + want: map[string]any{ + "allowed_domains": allowedDomains, + }, + }, + { + name: "BothFields", + options: &codersdk.ChatModelOpenAIProviderOptions{ + WebSearchEnabled: &enabled, + SearchContextSize: &searchContextSize, + AllowedDomains: allowedDomains, + }, + want: map[string]any{ + "search_context_size": searchContextSize, + "allowed_domains": allowedDomains, + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + tool, ok := chatopenai.WebSearchTool(tt.options) + require.True(t, ok) + + providerTool, ok := tool.(fantasy.ProviderDefinedTool) + require.True(t, ok) + require.Equal(t, "web_search", providerTool.ID) + require.Equal(t, "web_search", providerTool.Name) + require.NotNil(t, providerTool.Args) + require.Equal(t, tt.want, providerTool.Args) + }) + } +} diff --git a/coderd/x/chatd/chatprovider/chatprovider.go b/coderd/x/chatd/chatprovider/chatprovider.go index f0af44e038..6c019abcb2 100644 --- a/coderd/x/chatd/chatprovider/chatprovider.go +++ b/coderd/x/chatd/chatprovider/chatprovider.go @@ -19,6 +19,8 @@ import ( "golang.org/x/xerrors" "github.com/coder/coder/v2/coderd/database" + "github.com/coder/coder/v2/coderd/x/chatd/chatopenai" + "github.com/coder/coder/v2/coderd/x/chatd/chatutil" "github.com/coder/coder/v2/codersdk" ) @@ -637,7 +639,7 @@ func isChatModelForProvider(provider, modelID string) bool { case fantasyopenai.Name: return strings.HasPrefix(normalizedModel, "gpt-") || strings.HasPrefix(normalizedModel, "chatgpt-") || - isOpenAIReasoningModel(normalizedModel) + chatopenai.IsReasoningModel(normalizedModel) case fantasyanthropic.Name: return strings.HasPrefix(normalizedModel, "claude-") case fantasygoogle.Name: @@ -648,25 +650,6 @@ func isChatModelForProvider(provider, modelID string) bool { } } -func isOpenAIReasoningModel(modelID string) bool { - if len(modelID) < 2 || modelID[0] != 'o' { - return false - } - - index := 1 - for index < len(modelID) && modelID[index] >= '0' && modelID[index] <= '9' { - index++ - } - if index == 1 { - return false - } - - if index == len(modelID) { - return true - } - return modelID[index] == '-' || modelID[index] == '.' -} - // ReasoningEffortFromChat normalizes chat-config reasoning effort values for a // provider and returns the canonical provider effort value. func ReasoningEffortFromChat(provider string, value *string) *string { @@ -681,16 +664,14 @@ func ReasoningEffortFromChat(provider string, value *string) *string { switch NormalizeProvider(provider) { case fantasyopenai.Name: - return normalizedEnumValue( - normalized, - string(fantasyopenai.ReasoningEffortMinimal), - string(fantasyopenai.ReasoningEffortLow), - string(fantasyopenai.ReasoningEffortMedium), - string(fantasyopenai.ReasoningEffortHigh), - string(fantasyopenai.ReasoningEffortXHigh), - ) + effort := chatopenai.ReasoningEffortFromChat(value) + if effort == nil { + return nil + } + valueCopy := string(*effort) + return &valueCopy case fantasyanthropic.Name: - return normalizedEnumValue( + return chatutil.NormalizedEnumValue( normalized, string(fantasyanthropic.EffortLow), string(fantasyanthropic.EffortMedium), @@ -699,14 +680,14 @@ func ReasoningEffortFromChat(provider string, value *string) *string { string(fantasyanthropic.EffortMax), ) case fantasyopenrouter.Name: - return normalizedEnumValue( + return chatutil.NormalizedEnumValue( normalized, string(fantasyopenrouter.ReasoningEffortLow), string(fantasyopenrouter.ReasoningEffortMedium), string(fantasyopenrouter.ReasoningEffortHigh), ) case fantasyvercel.Name: - return normalizedEnumValue( + return chatutil.NormalizedEnumValue( normalized, string(fantasyvercel.ReasoningEffortNone), string(fantasyvercel.ReasoningEffortMinimal), @@ -859,41 +840,6 @@ func applyReasoningEffortDispatch( } } -// OpenAITextVerbosityFromChat normalizes chat-config text verbosity values for -// OpenAI and returns the canonical provider verbosity value. -func OpenAITextVerbosityFromChat(value *string) *fantasyopenai.TextVerbosity { - if value == nil { - return nil - } - - normalized := strings.ToLower(strings.TrimSpace(*value)) - if normalized == "" { - return nil - } - - verbosity := normalizedEnumValue( - normalized, - string(fantasyopenai.TextVerbosityLow), - string(fantasyopenai.TextVerbosityMedium), - string(fantasyopenai.TextVerbosityHigh), - ) - if verbosity == nil { - return nil - } - valueCopy := fantasyopenai.TextVerbosity(*verbosity) - return &valueCopy -} - -func normalizedEnumValue(value string, allowed ...string) *string { - for _, candidate := range allowed { - if value == strings.ToLower(candidate) { - match := candidate - return &match - } - } - return nil -} - // MergeMissingModelCostConfig fills unset pricing metadata from defaults. func MergeMissingModelCostConfig( dst **codersdk.ModelCostConfig, @@ -1487,7 +1433,7 @@ func ProviderOptionsFromChatModelConfig( result := fantasy.ProviderOptions{} if options.OpenAI != nil { - result[fantasyopenai.Name] = openAIProviderOptionsFromChatConfig( + result[fantasyopenai.Name] = chatopenai.ProviderOptionsFromChatConfig( model, options.OpenAI, ) @@ -1524,92 +1470,6 @@ func ProviderOptionsFromChatModelConfig( return result } -// IsResponsesStoreEnabled checks if the OpenAI Responses provider -// options are present and have Store set to true. When true, the -// provider stores conversation history server-side, enabling -// follow-up chaining via PreviousResponseID. -func IsResponsesStoreEnabled(opts fantasy.ProviderOptions) bool { - if opts == nil { - return false - } - raw, ok := opts[fantasyopenai.Name] - if !ok { - return false - } - respOpts, ok := raw.(*fantasyopenai.ResponsesProviderOptions) - if !ok || respOpts == nil { - return false - } - return respOpts.Store != nil && *respOpts.Store -} - -// CloneWithPreviousResponseID shallow-clones the provider options -// map and the OpenAI Responses entry, setting PreviousResponseID -// on the clone. The original map and entry are not mutated. -func CloneWithPreviousResponseID( - opts fantasy.ProviderOptions, - previousResponseID string, -) fantasy.ProviderOptions { - cloned := make(fantasy.ProviderOptions, len(opts)) - for k, v := range opts { - cloned[k] = v - } - if raw, ok := cloned[fantasyopenai.Name]; ok { - if respOpts, ok := raw.(*fantasyopenai.ResponsesProviderOptions); ok && respOpts != nil { - clone := *respOpts - clone.PreviousResponseID = &previousResponseID - cloned[fantasyopenai.Name] = &clone - } - } - return cloned -} - -func openAIProviderOptionsFromChatConfig( - model fantasy.LanguageModel, - options *codersdk.ChatModelOpenAIProviderOptions, -) fantasy.ProviderOptionsData { - reasoningEffort := openAIReasoningEffortFromChat(options.ReasoningEffort) - if useOpenAIResponsesOptions(model) { - include := ensureOpenAIResponseIncludes(openAIIncludeFromChat(options.Include)) - providerOptions := &fantasyopenai.ResponsesProviderOptions{ - Include: include, - Instructions: normalizedStringPointer(options.Instructions), - Logprobs: openAIResponsesLogProbsFromChat(options), - MaxToolCalls: options.MaxToolCalls, - Metadata: options.Metadata, - ParallelToolCalls: options.ParallelToolCalls, - PromptCacheKey: normalizedStringPointer(options.PromptCacheKey), - ReasoningEffort: reasoningEffort, - ReasoningSummary: normalizedStringPointer(options.ReasoningSummary), - SafetyIdentifier: normalizedStringPointer(options.SafetyIdentifier), - ServiceTier: openAIServiceTierFromChat(options.ServiceTier), - StrictJSONSchema: options.StrictJSONSchema, - Store: boolPtrOrDefault(options.Store, true), - TextVerbosity: OpenAITextVerbosityFromChat(options.TextVerbosity), - User: normalizedStringPointer(options.User), - } - return providerOptions - } - - return &fantasyopenai.ProviderOptions{ - LogitBias: options.LogitBias, - LogProbs: options.LogProbs, - TopLogProbs: options.TopLogProbs, - ParallelToolCalls: options.ParallelToolCalls, - User: normalizedStringPointer(options.User), - ReasoningEffort: reasoningEffort, - MaxCompletionTokens: options.MaxCompletionTokens, - TextVerbosity: normalizedStringPointer(options.TextVerbosity), - Prediction: options.Prediction, - Store: boolPtrOrDefault(options.Store, true), - Metadata: options.Metadata, - PromptCacheKey: normalizedStringPointer(options.PromptCacheKey), - SafetyIdentifier: normalizedStringPointer(options.SafetyIdentifier), - ServiceTier: normalizedStringPointer(options.ServiceTier), - StructuredOutputs: options.StructuredOutputs, - } -} - func anthropicProviderOptionsFromChatConfig( options *codersdk.ChatModelAnthropicProviderOptions, ) *fantasyanthropic.ProviderOptions { @@ -1659,8 +1519,8 @@ func openAICompatProviderOptionsFromChatConfig( options *codersdk.ChatModelOpenAICompatProviderOptions, ) *fantasyopenaicompat.ProviderOptions { return &fantasyopenaicompat.ProviderOptions{ - User: normalizedStringPointer(options.User), - ReasoningEffort: openAIReasoningEffortFromChat(options.ReasoningEffort), + User: chatutil.NormalizedStringPointer(options.User), + ReasoningEffort: chatopenai.ReasoningEffortFromChat(options.ReasoningEffort), } } @@ -1673,7 +1533,7 @@ func openRouterProviderOptionsFromChatConfig( LogitBias: options.LogitBias, LogProbs: options.LogProbs, ParallelToolCalls: options.ParallelToolCalls, - User: normalizedStringPointer(options.User), + User: chatutil.NormalizedStringPointer(options.User), } if options.Reasoning != nil { result.Reasoning = &fantasyopenrouter.ReasoningOptions{ @@ -1688,11 +1548,11 @@ func openRouterProviderOptionsFromChatConfig( Order: options.Provider.Order, AllowFallbacks: options.Provider.AllowFallbacks, RequireParameters: options.Provider.RequireParameters, - DataCollection: normalizedStringPointer(options.Provider.DataCollection), + DataCollection: chatutil.NormalizedStringPointer(options.Provider.DataCollection), Only: options.Provider.Only, Ignore: options.Provider.Ignore, Quantizations: options.Provider.Quantizations, - Sort: normalizedStringPointer(options.Provider.Sort), + Sort: chatutil.NormalizedStringPointer(options.Provider.Sort), } } return result @@ -1702,7 +1562,7 @@ func vercelProviderOptionsFromChatConfig( options *codersdk.ChatModelVercelProviderOptions, ) *fantasyvercel.ProviderOptions { result := &fantasyvercel.ProviderOptions{ - User: normalizedStringPointer(options.User), + User: chatutil.NormalizedStringPointer(options.User), LogitBias: options.LogitBias, LogProbs: options.LogProbs, TopLogProbs: options.TopLogProbs, @@ -1726,89 +1586,6 @@ func vercelProviderOptionsFromChatConfig( return result } -func openAIResponsesLogProbsFromChat( - options *codersdk.ChatModelOpenAIProviderOptions, -) any { - if options.TopLogProbs != nil { - return *options.TopLogProbs - } - if options.LogProbs != nil { - return *options.LogProbs - } - return nil -} - -func openAIIncludeFromChat(values []string) []fantasyopenai.IncludeType { - if values == nil { - return nil - } - - result := make([]fantasyopenai.IncludeType, 0, len(values)) - for _, value := range values { - switch strings.TrimSpace(value) { - case string(fantasyopenai.IncludeReasoningEncryptedContent): - result = append(result, fantasyopenai.IncludeReasoningEncryptedContent) - case string(fantasyopenai.IncludeFileSearchCallResults): - result = append(result, fantasyopenai.IncludeFileSearchCallResults) - case string(fantasyopenai.IncludeMessageOutputTextLogprobs): - result = append(result, fantasyopenai.IncludeMessageOutputTextLogprobs) - } - } - return result -} - -func ensureOpenAIResponseIncludes( - values []fantasyopenai.IncludeType, -) []fantasyopenai.IncludeType { - const required = fantasyopenai.IncludeReasoningEncryptedContent - - for _, value := range values { - if value == required { - return values - } - } - return append(values, required) -} - -func useOpenAIResponsesOptions(model fantasy.LanguageModel) bool { - if model == nil { - return false - } - switch model.Provider() { - case fantasyopenai.Name, fantasyazure.Name: - return fantasyopenai.IsResponsesModel(model.Model()) - default: - return false - } -} - -func boolPtrOrDefault(value *bool, def bool) *bool { - if value != nil { - return value - } - return &def -} - -func normalizedStringPointer(value *string) *string { - if value == nil { - return nil - } - trimmed := strings.TrimSpace(*value) - if trimmed == "" { - return nil - } - return &trimmed -} - -func openAIReasoningEffortFromChat(value *string) *fantasyopenai.ReasoningEffort { - effort := ReasoningEffortFromChat(fantasyopenai.Name, value) - if effort == nil { - return nil - } - valueCopy := fantasyopenai.ReasoningEffort(*effort) - return &valueCopy -} - func anthropicEffortFromChat(value *string) *fantasyanthropic.Effort { effort := ReasoningEffortFromChat(fantasyanthropic.Name, value) if effort == nil { @@ -1835,23 +1612,3 @@ func vercelReasoningEffortFromChat(value *string) *fantasyvercel.ReasoningEffort valueCopy := fantasyvercel.ReasoningEffort(*effort) return &valueCopy } - -func openAIServiceTierFromChat(value *string) *fantasyopenai.ServiceTier { - normalized := normalizedStringPointer(value) - if normalized == nil { - return nil - } - switch strings.ToLower(*normalized) { - case string(fantasyopenai.ServiceTierAuto): - serviceTier := fantasyopenai.ServiceTierAuto - return &serviceTier - case string(fantasyopenai.ServiceTierFlex): - serviceTier := fantasyopenai.ServiceTierFlex - return &serviceTier - case string(fantasyopenai.ServiceTierPriority): - serviceTier := fantasyopenai.ServiceTierPriority - return &serviceTier - default: - return nil - } -} diff --git a/coderd/x/chatd/chattest/messages.go b/coderd/x/chatd/chattest/messages.go new file mode 100644 index 0000000000..0833be109d --- /dev/null +++ b/coderd/x/chatd/chattest/messages.go @@ -0,0 +1,19 @@ +package chattest + +import ( + "encoding/json" + + "github.com/sqlc-dev/pqtype" + + "github.com/coder/coder/v2/coderd/database" + "github.com/coder/coder/v2/codersdk" +) + +// ChatMessageWithParts returns a database chat message whose content is the +// JSON encoding of the provided SDK message parts. +func ChatMessageWithParts(parts []codersdk.ChatMessagePart) database.ChatMessage { + raw, _ := json.Marshal(parts) + return database.ChatMessage{ + Content: pqtype.NullRawMessage{RawMessage: raw, Valid: true}, + } +} diff --git a/coderd/x/chatd/chatutil/chatutil.go b/coderd/x/chatd/chatutil/chatutil.go new file mode 100644 index 0000000000..9158fbb598 --- /dev/null +++ b/coderd/x/chatd/chatutil/chatutil.go @@ -0,0 +1,28 @@ +package chatutil + +import "strings" + +// NormalizedStringPointer trims a string pointer and returns nil for nil or +// empty values. +func NormalizedStringPointer(value *string) *string { + if value == nil { + return nil + } + trimmed := strings.TrimSpace(*value) + if trimmed == "" { + return nil + } + return &trimmed +} + +// NormalizedEnumValue returns the canonical allowed value matching value after +// case normalization, or nil when no value matches. +func NormalizedEnumValue(value string, allowed ...string) *string { + for _, candidate := range allowed { + if value == strings.ToLower(candidate) { + match := candidate + return &match + } + } + return nil +} diff --git a/coderd/x/chatd/chatutil/chatutil_test.go b/coderd/x/chatd/chatutil/chatutil_test.go new file mode 100644 index 0000000000..5bd7835f21 --- /dev/null +++ b/coderd/x/chatd/chatutil/chatutil_test.go @@ -0,0 +1,79 @@ +package chatutil_test + +import ( + "testing" + + "github.com/stretchr/testify/require" + + "github.com/coder/coder/v2/coderd/x/chatd/chatutil" +) + +func TestNormalizedStringPointer(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + value *string + want *string + }{ + {name: "Nil"}, + {name: "Empty", value: ptr("")}, + {name: "WhitespaceOnly", value: ptr(" \t\n ")}, + {name: "Trimmed", value: ptr(" value "), want: ptr("value")}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + got := chatutil.NormalizedStringPointer(tt.value) + if tt.want == nil { + require.Nil(t, got) + return + } + require.NotNil(t, got) + require.Equal(t, *tt.want, *got) + }) + } +} + +func TestNormalizedEnumValue(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + value string + allowed []string + want *string + }{ + { + name: "MatchFound", + value: "medium", + allowed: []string{"Low", "Medium", "High"}, + want: ptr("Medium"), + }, + { + name: "MatchMissing", + value: "maximum", + allowed: []string{"Low", "Medium", "High"}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + got := chatutil.NormalizedEnumValue(tt.value, tt.allowed...) + if tt.want == nil { + require.Nil(t, got) + return + } + require.NotNil(t, got) + require.Equal(t, *tt.want, *got) + }) + } +} + +func ptr[T any](value T) *T { + return &value +}