mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
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:
+17
-283
@@ -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,
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -0,0 +1,499 @@
|
||||
package chatopenai_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"charm.land/fantasy"
|
||||
fantasyazure "charm.land/fantasy/providers/azure"
|
||||
fantasyopenai "charm.land/fantasy/providers/openai"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/x/chatd/chatopenai"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
)
|
||||
|
||||
func TestProviderOptionsFromChatConfigLegacy(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
store := false
|
||||
logProbs := true
|
||||
topLogProbs := int64(3)
|
||||
parallelToolCalls := true
|
||||
maxCompletionTokens := int64(4096)
|
||||
structuredOutputs := true
|
||||
options := &codersdk.ChatModelOpenAIProviderOptions{
|
||||
LogitBias: map[string]int64{
|
||||
"50256": -10,
|
||||
},
|
||||
LogProbs: &logProbs,
|
||||
TopLogProbs: &topLogProbs,
|
||||
ParallelToolCalls: ¶llelToolCalls,
|
||||
User: ptr(" user-1 "),
|
||||
ReasoningEffort: ptr(" HIGH "),
|
||||
MaxCompletionTokens: &maxCompletionTokens,
|
||||
TextVerbosity: ptr(" High "),
|
||||
Prediction: map[string]any{
|
||||
"type": "content",
|
||||
},
|
||||
Store: &store,
|
||||
Metadata: map[string]any{"feature": "chat"},
|
||||
PromptCacheKey: ptr(" cache-key "),
|
||||
SafetyIdentifier: ptr(" safety-id "),
|
||||
ServiceTier: ptr(" priority "),
|
||||
StructuredOutputs: &structuredOutputs,
|
||||
}
|
||||
|
||||
got := chatopenai.ProviderOptionsFromChatConfig(
|
||||
fakeLanguageModel{provider: fantasyopenai.Name, model: "gpt-3.5-turbo-instruct"},
|
||||
options,
|
||||
)
|
||||
|
||||
providerOptions, ok := got.(*fantasyopenai.ProviderOptions)
|
||||
require.True(t, ok)
|
||||
require.Equal(t, options.LogitBias, providerOptions.LogitBias)
|
||||
require.Same(t, options.LogProbs, providerOptions.LogProbs)
|
||||
require.Same(t, options.TopLogProbs, providerOptions.TopLogProbs)
|
||||
require.Same(t, options.ParallelToolCalls, providerOptions.ParallelToolCalls)
|
||||
require.Equal(t, "user-1", requireStringPointerValue(t, providerOptions.User))
|
||||
require.Equal(t, fantasyopenai.ReasoningEffortHigh, requireReasoningEffortPointerValue(t, providerOptions.ReasoningEffort))
|
||||
require.Same(t, options.MaxCompletionTokens, providerOptions.MaxCompletionTokens)
|
||||
require.Equal(t, "High", requireStringPointerValue(t, providerOptions.TextVerbosity))
|
||||
require.Equal(t, options.Prediction, providerOptions.Prediction)
|
||||
require.Same(t, options.Store, providerOptions.Store)
|
||||
require.Equal(t, false, requireBoolPointerValue(t, providerOptions.Store))
|
||||
require.Equal(t, options.Metadata, providerOptions.Metadata)
|
||||
require.Equal(t, "cache-key", requireStringPointerValue(t, providerOptions.PromptCacheKey))
|
||||
require.Equal(t, "safety-id", requireStringPointerValue(t, providerOptions.SafetyIdentifier))
|
||||
require.Equal(t, "priority", requireStringPointerValue(t, providerOptions.ServiceTier))
|
||||
require.Same(t, options.StructuredOutputs, providerOptions.StructuredOutputs)
|
||||
}
|
||||
|
||||
func TestProviderOptionsFromChatConfigResponses(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
topLogProbs := int64(5)
|
||||
maxToolCalls := int64(8)
|
||||
parallelToolCalls := false
|
||||
strictJSONSchema := true
|
||||
options := &codersdk.ChatModelOpenAIProviderOptions{
|
||||
Include: []string{
|
||||
string(fantasyopenai.IncludeFileSearchCallResults),
|
||||
"unsupported",
|
||||
},
|
||||
Instructions: ptr(" instructions "),
|
||||
LogProbs: ptr(true),
|
||||
TopLogProbs: &topLogProbs,
|
||||
MaxToolCalls: &maxToolCalls,
|
||||
Metadata: map[string]any{"scope": "unit"},
|
||||
ParallelToolCalls: ¶llelToolCalls,
|
||||
PromptCacheKey: ptr(" prompt-cache "),
|
||||
ReasoningEffort: ptr(" minimal "),
|
||||
ReasoningSummary: ptr(" auto "),
|
||||
SafetyIdentifier: ptr(" safety "),
|
||||
ServiceTier: ptr(" FLEX "),
|
||||
StrictJSONSchema: &strictJSONSchema,
|
||||
TextVerbosity: ptr(" MEDIUM "),
|
||||
User: ptr(" user-2 "),
|
||||
}
|
||||
|
||||
got := chatopenai.ProviderOptionsFromChatConfig(
|
||||
fakeLanguageModel{provider: fantasyopenai.Name, model: "gpt-4.1"},
|
||||
options,
|
||||
)
|
||||
|
||||
providerOptions, ok := got.(*fantasyopenai.ResponsesProviderOptions)
|
||||
require.True(t, ok)
|
||||
require.Equal(t, []fantasyopenai.IncludeType{
|
||||
fantasyopenai.IncludeFileSearchCallResults,
|
||||
fantasyopenai.IncludeReasoningEncryptedContent,
|
||||
}, providerOptions.Include)
|
||||
require.Equal(t, "instructions", requireStringPointerValue(t, providerOptions.Instructions))
|
||||
require.Equal(t, int64(5), providerOptions.Logprobs)
|
||||
require.Same(t, options.MaxToolCalls, providerOptions.MaxToolCalls)
|
||||
require.Equal(t, options.Metadata, providerOptions.Metadata)
|
||||
require.Same(t, options.ParallelToolCalls, providerOptions.ParallelToolCalls)
|
||||
require.Equal(t, "prompt-cache", requireStringPointerValue(t, providerOptions.PromptCacheKey))
|
||||
require.Equal(t, fantasyopenai.ReasoningEffortMinimal, requireReasoningEffortPointerValue(t, providerOptions.ReasoningEffort))
|
||||
require.Equal(t, "auto", requireStringPointerValue(t, providerOptions.ReasoningSummary))
|
||||
require.Equal(t, "safety", requireStringPointerValue(t, providerOptions.SafetyIdentifier))
|
||||
require.Equal(t, fantasyopenai.ServiceTierFlex, requireServiceTierPointerValue(t, providerOptions.ServiceTier))
|
||||
require.Same(t, options.StrictJSONSchema, providerOptions.StrictJSONSchema)
|
||||
require.NotNil(t, providerOptions.Store)
|
||||
require.True(t, *providerOptions.Store)
|
||||
require.Equal(t, fantasyopenai.TextVerbosityMedium, requireTextVerbosityPointerValue(t, providerOptions.TextVerbosity))
|
||||
require.Equal(t, "user-2", requireStringPointerValue(t, providerOptions.User))
|
||||
}
|
||||
|
||||
func TestTextVerbosityFromChat(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
value *string
|
||||
want *fantasyopenai.TextVerbosity
|
||||
}{
|
||||
{name: "Nil"},
|
||||
{name: "Empty", value: ptr(" ")},
|
||||
{name: "Low", value: ptr(" low "), want: ptr(fantasyopenai.TextVerbosityLow)},
|
||||
{name: "MediumCase", value: ptr(" MEDIUM "), want: ptr(fantasyopenai.TextVerbosityMedium)},
|
||||
{name: "High", value: ptr("high"), want: ptr(fantasyopenai.TextVerbosityHigh)},
|
||||
{name: "Invalid", value: ptr("verbose")},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
got := chatopenai.TextVerbosityFromChat(tt.value)
|
||||
if tt.want == nil {
|
||||
require.Nil(t, got)
|
||||
return
|
||||
}
|
||||
require.NotNil(t, got)
|
||||
require.Equal(t, *tt.want, *got)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestIncludeFromChat(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
values []string
|
||||
want []fantasyopenai.IncludeType
|
||||
}{
|
||||
{name: "Nil"},
|
||||
{name: "Empty", values: []string{}, want: []fantasyopenai.IncludeType{}},
|
||||
{
|
||||
name: "ValidAndInvalid",
|
||||
values: []string{
|
||||
" " + string(fantasyopenai.IncludeReasoningEncryptedContent) + " ",
|
||||
string(fantasyopenai.IncludeFileSearchCallResults),
|
||||
"unsupported",
|
||||
string(fantasyopenai.IncludeMessageOutputTextLogprobs),
|
||||
},
|
||||
want: []fantasyopenai.IncludeType{
|
||||
fantasyopenai.IncludeReasoningEncryptedContent,
|
||||
fantasyopenai.IncludeFileSearchCallResults,
|
||||
fantasyopenai.IncludeMessageOutputTextLogprobs,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
got := chatopenai.IncludeFromChat(tt.values)
|
||||
require.Equal(t, tt.want, got)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestEnsureResponseIncludes(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
values []fantasyopenai.IncludeType
|
||||
want []fantasyopenai.IncludeType
|
||||
}{
|
||||
{
|
||||
name: "NilAddsRequired",
|
||||
want: []fantasyopenai.IncludeType{fantasyopenai.IncludeReasoningEncryptedContent},
|
||||
},
|
||||
{
|
||||
name: "EmptyAddsRequired",
|
||||
values: []fantasyopenai.IncludeType{},
|
||||
want: []fantasyopenai.IncludeType{fantasyopenai.IncludeReasoningEncryptedContent},
|
||||
},
|
||||
{
|
||||
name: "AddsRequiredAfterExistingValues",
|
||||
values: []fantasyopenai.IncludeType{
|
||||
fantasyopenai.IncludeFileSearchCallResults,
|
||||
},
|
||||
want: []fantasyopenai.IncludeType{
|
||||
fantasyopenai.IncludeFileSearchCallResults,
|
||||
fantasyopenai.IncludeReasoningEncryptedContent,
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "DoesNotDuplicateRequired",
|
||||
values: []fantasyopenai.IncludeType{
|
||||
fantasyopenai.IncludeReasoningEncryptedContent,
|
||||
fantasyopenai.IncludeFileSearchCallResults,
|
||||
},
|
||||
want: []fantasyopenai.IncludeType{
|
||||
fantasyopenai.IncludeReasoningEncryptedContent,
|
||||
fantasyopenai.IncludeFileSearchCallResults,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
got := chatopenai.EnsureResponseIncludes(tt.values)
|
||||
require.Equal(t, tt.want, got)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestUsesResponsesOptions(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
model fantasy.LanguageModel
|
||||
want bool
|
||||
}{
|
||||
{name: "Nil"},
|
||||
{
|
||||
name: "OpenAIResponsesModel",
|
||||
model: fakeLanguageModel{provider: fantasyopenai.Name, model: "gpt-4.1"},
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "AzureResponsesModel",
|
||||
model: fakeLanguageModel{provider: fantasyazure.Name, model: "gpt-4.1"},
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "OpenAINonResponsesModel",
|
||||
model: fakeLanguageModel{provider: fantasyopenai.Name, model: "gpt-3.5-turbo-instruct"},
|
||||
},
|
||||
{
|
||||
name: "NonOpenAIProvider",
|
||||
model: fakeLanguageModel{provider: "other", model: "gpt-4.1"},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
got := chatopenai.UsesResponsesOptions(tt.model)
|
||||
require.Equal(t, tt.want, got)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestReasoningEffortFromChat(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
value *string
|
||||
want *fantasyopenai.ReasoningEffort
|
||||
}{
|
||||
{name: "Nil"},
|
||||
{name: "Empty", value: ptr(" ")},
|
||||
{name: "Minimal", value: ptr(" minimal "), want: ptr(fantasyopenai.ReasoningEffortMinimal)},
|
||||
{name: "LowCase", value: ptr(" LOW "), want: ptr(fantasyopenai.ReasoningEffortLow)},
|
||||
{name: "Medium", value: ptr("medium"), want: ptr(fantasyopenai.ReasoningEffortMedium)},
|
||||
{name: "High", value: ptr("high"), want: ptr(fantasyopenai.ReasoningEffortHigh)},
|
||||
{name: "XHigh", value: ptr("xhigh"), want: ptr(fantasyopenai.ReasoningEffortXHigh)},
|
||||
{name: "NoneUnsupported", value: ptr("none")},
|
||||
{name: "Invalid", value: ptr("max")},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
got := chatopenai.ReasoningEffortFromChat(tt.value)
|
||||
if tt.want == nil {
|
||||
require.Nil(t, got)
|
||||
return
|
||||
}
|
||||
require.NotNil(t, got)
|
||||
require.Equal(t, *tt.want, *got)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestServiceTierFromChat(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
value *string
|
||||
want *fantasyopenai.ServiceTier
|
||||
}{
|
||||
{name: "Nil"},
|
||||
{name: "Empty", value: ptr(" ")},
|
||||
{name: "Auto", value: ptr(" auto "), want: ptr(fantasyopenai.ServiceTierAuto)},
|
||||
{name: "FlexCase", value: ptr(" FLEX "), want: ptr(fantasyopenai.ServiceTierFlex)},
|
||||
{name: "Priority", value: ptr("priority"), want: ptr(fantasyopenai.ServiceTierPriority)},
|
||||
{name: "DefaultUnsupported", value: ptr("default")},
|
||||
{name: "Invalid", value: ptr("fast")},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
got := chatopenai.ServiceTierFromChat(tt.value)
|
||||
if tt.want == nil {
|
||||
require.Nil(t, got)
|
||||
return
|
||||
}
|
||||
require.NotNil(t, got)
|
||||
require.Equal(t, *tt.want, *got)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestResponsesLogProbsFromChatConfig(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
logProbs := true
|
||||
topLogProbs := int64(4)
|
||||
tests := []struct {
|
||||
name string
|
||||
options *codersdk.ChatModelOpenAIProviderOptions
|
||||
want any
|
||||
}{
|
||||
{name: "Nil"},
|
||||
{
|
||||
name: "Empty",
|
||||
options: &codersdk.ChatModelOpenAIProviderOptions{},
|
||||
},
|
||||
{
|
||||
name: "LogProbs",
|
||||
options: &codersdk.ChatModelOpenAIProviderOptions{
|
||||
LogProbs: &logProbs,
|
||||
},
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "TopLogProbs",
|
||||
options: &codersdk.ChatModelOpenAIProviderOptions{
|
||||
TopLogProbs: &topLogProbs,
|
||||
},
|
||||
want: int64(4),
|
||||
},
|
||||
{
|
||||
name: "TopLogProbsPrecedence",
|
||||
options: &codersdk.ChatModelOpenAIProviderOptions{
|
||||
LogProbs: &logProbs,
|
||||
TopLogProbs: &topLogProbs,
|
||||
},
|
||||
want: int64(4),
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
got := chatopenai.ResponsesLogProbsFromChatConfig(tt.options)
|
||||
require.Equal(t, tt.want, got)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsReasoningModel(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
model string
|
||||
want bool
|
||||
}{
|
||||
{model: ""},
|
||||
{model: "o"},
|
||||
{model: "o1", want: true},
|
||||
{model: "o1-mini", want: true},
|
||||
{model: "o3.5", want: true},
|
||||
{model: "o10-preview", want: true},
|
||||
{model: "oabc"},
|
||||
{model: "ox"},
|
||||
{model: "o1preview"},
|
||||
{model: "gpt-5"},
|
||||
{model: "O1"},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.model, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
got := chatopenai.IsReasoningModel(tt.model)
|
||||
require.Equal(t, tt.want, got)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func requireStringPointerValue(t *testing.T, value *string) string {
|
||||
t.Helper()
|
||||
require.NotNil(t, value)
|
||||
return *value
|
||||
}
|
||||
|
||||
func requireBoolPointerValue(t *testing.T, value *bool) bool {
|
||||
t.Helper()
|
||||
require.NotNil(t, value)
|
||||
return *value
|
||||
}
|
||||
|
||||
func requireReasoningEffortPointerValue(
|
||||
t *testing.T,
|
||||
value *fantasyopenai.ReasoningEffort,
|
||||
) fantasyopenai.ReasoningEffort {
|
||||
t.Helper()
|
||||
require.NotNil(t, value)
|
||||
return *value
|
||||
}
|
||||
|
||||
func requireServiceTierPointerValue(
|
||||
t *testing.T,
|
||||
value *fantasyopenai.ServiceTier,
|
||||
) fantasyopenai.ServiceTier {
|
||||
t.Helper()
|
||||
require.NotNil(t, value)
|
||||
return *value
|
||||
}
|
||||
|
||||
func requireTextVerbosityPointerValue(
|
||||
t *testing.T,
|
||||
value *fantasyopenai.TextVerbosity,
|
||||
) fantasyopenai.TextVerbosity {
|
||||
t.Helper()
|
||||
require.NotNil(t, value)
|
||||
return *value
|
||||
}
|
||||
|
||||
func ptr[T any](value T) *T {
|
||||
return &value
|
||||
}
|
||||
|
||||
type fakeLanguageModel struct {
|
||||
provider string
|
||||
model string
|
||||
}
|
||||
|
||||
func (fakeLanguageModel) Generate(context.Context, fantasy.Call) (*fantasy.Response, error) {
|
||||
panic("not implemented")
|
||||
}
|
||||
|
||||
func (fakeLanguageModel) Stream(context.Context, fantasy.Call) (fantasy.StreamResponse, error) {
|
||||
panic("not implemented")
|
||||
}
|
||||
|
||||
func (fakeLanguageModel) GenerateObject(context.Context, fantasy.ObjectCall) (*fantasy.ObjectResponse, error) {
|
||||
panic("not implemented")
|
||||
}
|
||||
|
||||
func (fakeLanguageModel) StreamObject(context.Context, fantasy.ObjectCall) (fantasy.ObjectStreamResponse, error) {
|
||||
panic("not implemented")
|
||||
}
|
||||
|
||||
func (f fakeLanguageModel) Provider() string {
|
||||
return f.provider
|
||||
}
|
||||
|
||||
func (f fakeLanguageModel) Model() string {
|
||||
return f.model
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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,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
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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},
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
Reference in New Issue
Block a user