refactor(coderd/x/chatd): extract OpenAI logic into chatopenai package (#24788)

Extracts OpenAI-specific logic from `coderd/x/chatd` into
`coderd/x/chatd/chatopenai` so the main chat path no longer references
`fantasyopenai` directly for chain mode info, response IDs, web search
tooling, or option mapping.

Structural refactor. The only deliberate behavioral narrowing is
consolidating Responses store checks and related keyed option or
metadata access on `opts[fantasyopenai.Name]`. That is documented by
`TestIsResponsesStoreEnabledIgnoresMalformedNonOpenAIKey` and is
unreachable in production where Responses options always live under
`fantasyopenai.Name`.

Summary:

- Moves OpenAI Responses chain mode info, response ID helpers, web
search tool construction, and provider option conversion into
`chatopenai`.
- Keeps Anthropic, Google, OpenRouter, and Vercel provider branches as
thin, existing code paths.
- `chatopenai` only imports `chatprompt` from chatd subpackages. It does
not import `chatd`, `chatloop`, `chatprovider`, or `chaterror`.
- Follow-up review fixes align helper names, keyed provider option
access, map cloning behavior, and PR documentation with the extracted
package boundary.
- Final sweep trims unused chain-mode state, removes a duplicate
store-check test case, drops an unused provider-tool parameter, and
shares the chat-message test helper through `chattest`.

> Mux is updating this PR on Mike's behalf.
This commit is contained in:
Michael Suchacz
2026-05-04 11:17:19 +02:00
committed by GitHub
parent 761adfa62a
commit 203b0a9df8
13 changed files with 2468 additions and 1277 deletions
+17 -283
View File
@@ -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,
})
}
+27 -648
View File
@@ -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
+5 -84
View File
@@ -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.
+228
View File
@@ -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
}
+499
View File
@@ -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: &parallelToolCalls,
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: &parallelToolCalls,
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
}
+409
View File
@@ -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
}
+993
View File
@@ -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
}
+29
View File
@@ -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
}
+116
View File
@@ -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)
})
}
}
+19 -262
View File
@@ -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
}
}
+19
View File
@@ -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},
}
}
+28
View File
@@ -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
}
+79
View File
@@ -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
}