mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat(chatd): add provider-native web search tools to chats (#22909)
## What
Adds provider-native web search tools to the chat system. Anthropic,
OpenAI, and Google all offer server-side web search — this wires them up
as opt-in per-model config options using the existing
`ChatModelProviderOptions` JSONB column (no migration).
Web search is **off by default**.
## Config
Set `web_search_enabled: true` in the model config provider options:
```json
{
"provider_options": {
"anthropic": {
"web_search_enabled": true,
"allowed_domains": ["docs.coder.com", "github.com"]
}
}
}
```
Available options per provider:
- **Anthropic**: `web_search_enabled`, `allowed_domains`,
`blocked_domains`
- **OpenAI**: `web_search_enabled`, `search_context_size`
(`low`/`medium`/`high`), `allowed_domains`
- **Google**: `web_search_enabled`
## Backend
- `codersdk/chats.go` — new fields on the per-provider option structs
- `coderd/chatd/chatd.go` — `buildProviderTools()` reads config, creates
`ProviderDefinedTool` entries (uses `anthropic.WebSearchTool()` helper
from fantasy)
- `coderd/chatd/chatloop/chatloop.go` — `ProviderTools` on `RunOptions`,
merged into `Call.Tools`. Provider-executed tool calls skip local
execution. `StreamPartTypeToolResult` with `ProviderExecuted: true` is
accumulated inline (matching fantasy's own agent.go pattern) instead of
post-stream synthesis.
- `coderd/chatd/chatprompt/` — `MarshalToolResult` carries
`ProviderMetadata` through DB persistence so multi-turn round-trips work
(Anthropic needs `encrypted_content` back)
## Frontend
- Source citations render **inline** at the tool-call position (not
bottom-of-message), using `ToolCollapsible` so they look like other tool
cards — collapsed "Searched N results" with globe icon, expand to see
source pills
- Provider-executed tool calls/results are hidden from the normal tool
card UI
- Tool-role messages with only provider-executed results return `null`
(no empty bubble)
- Both persisted (messageParsing.ts) and streaming (streamState.ts)
paths group consecutive `source` parts into a single `{ type: "sources"
}` render block
## Fantasy changes
The fantasy fork (`kylecarbs/fantasy` branch `cj/go1.25`) has the
Anthropic tool code merged in, but will hopefully go upstream from:
https://github.com/charmbracelet/fantasy/pull/163
This commit is contained in:
+51
-2
@@ -12,6 +12,7 @@ import (
|
||||
"time"
|
||||
|
||||
"charm.land/fantasy"
|
||||
"charm.land/fantasy/providers/anthropic"
|
||||
"github.com/google/uuid"
|
||||
"github.com/sqlc-dev/pqtype"
|
||||
"golang.org/x/xerrors"
|
||||
@@ -2124,7 +2125,7 @@ func (p *Server) runChat(
|
||||
p.maybeGenerateChatTitle(context.WithoutCancel(ctx), chat, messages, model, providerKeys, logger)
|
||||
}()
|
||||
|
||||
prompt, err := chatprompt.ConvertMessagesWithFiles(ctx, messages, p.chatFileResolver())
|
||||
prompt, err := chatprompt.ConvertMessagesWithFiles(ctx, messages, p.chatFileResolver(), logger)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("build chat prompt: %w", err)
|
||||
}
|
||||
@@ -2492,6 +2493,13 @@ func (p *Server) runChat(
|
||||
})...)
|
||||
}
|
||||
|
||||
// Build provider-native tools (e.g., web search) based on
|
||||
// the model configuration.
|
||||
var providerTools []fantasy.Tool
|
||||
if callConfig.ProviderOptions != nil {
|
||||
providerTools = buildProviderTools(model.Provider(), callConfig.ProviderOptions)
|
||||
}
|
||||
|
||||
err = chatloop.Run(ctx, chatloop.RunOptions{
|
||||
Model: model,
|
||||
Messages: prompt,
|
||||
@@ -2500,6 +2508,7 @@ func (p *Server) runChat(
|
||||
|
||||
ModelConfig: callConfig,
|
||||
ProviderOptions: chatprovider.ProviderOptionsFromChatModelConfig(model, callConfig.ProviderOptions),
|
||||
ProviderTools: providerTools,
|
||||
|
||||
ContextLimitFallback: modelConfigContextLimit,
|
||||
|
||||
@@ -2516,7 +2525,7 @@ func (p *Server) runChat(
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("reload chat messages: %w", err)
|
||||
}
|
||||
reloadedPrompt, err := chatprompt.ConvertMessagesWithFiles(reloadCtx, reloadedMsgs, p.chatFileResolver())
|
||||
reloadedPrompt, err := chatprompt.ConvertMessagesWithFiles(reloadCtx, reloadedMsgs, p.chatFileResolver(), logger)
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("convert reloaded messages: %w", err)
|
||||
}
|
||||
@@ -2564,6 +2573,44 @@ func (p *Server) runChat(
|
||||
return err
|
||||
}
|
||||
|
||||
// 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) []fantasy.Tool {
|
||||
var tools []fantasy.Tool
|
||||
|
||||
if options.Anthropic != nil && options.Anthropic.WebSearchEnabled != nil && *options.Anthropic.WebSearchEnabled {
|
||||
tools = append(tools, anthropic.WebSearchTool(&anthropic.WebSearchToolOptions{
|
||||
AllowedDomains: options.Anthropic.AllowedDomains,
|
||||
BlockedDomains: options.Anthropic.BlockedDomains,
|
||||
}))
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
tools = append(tools, fantasy.ProviderDefinedTool{
|
||||
ID: "web_search",
|
||||
Name: "web_search",
|
||||
Args: args,
|
||||
})
|
||||
}
|
||||
|
||||
if options.Google != nil && options.Google.WebSearchEnabled != nil && *options.Google.WebSearchEnabled {
|
||||
tools = append(tools, fantasy.ProviderDefinedTool{
|
||||
ID: "web_search",
|
||||
Name: "web_search",
|
||||
})
|
||||
}
|
||||
|
||||
return tools
|
||||
}
|
||||
|
||||
// persistChatContextSummary persists a chat context summary to the database.
|
||||
// This is invoked via the chat loop's compaction callback.
|
||||
func (p *Server) persistChatContextSummary(
|
||||
@@ -2618,6 +2665,8 @@ func (p *Server) persistChatContextSummary(
|
||||
"chat_summarized",
|
||||
summaryResult,
|
||||
false,
|
||||
false,
|
||||
nil,
|
||||
)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("encode summary tool result: %w", err)
|
||||
|
||||
@@ -63,6 +63,12 @@ type RunOptions struct {
|
||||
// of the provider, which lives in chatd, not chatloop.
|
||||
ProviderOptions fantasy.ProviderOptions
|
||||
|
||||
// ProviderTools are provider-native tools (like web search)
|
||||
// that are passed directly to the provider API alongside
|
||||
// function tool definitions. These are not necessarily
|
||||
// executed server-side; handling is provider-specific.
|
||||
ProviderTools []fantasy.Tool
|
||||
|
||||
PersistStep func(context.Context, PersistedStep) error
|
||||
PublishMessagePart func(
|
||||
role fantasy.MessageRole,
|
||||
@@ -153,9 +159,10 @@ func (r stepResult) toResponseMessages() []fantasy.Message {
|
||||
continue
|
||||
}
|
||||
toolParts = append(toolParts, fantasy.ToolResultPart{
|
||||
ToolCallID: result.ToolCallID,
|
||||
Output: result.Result,
|
||||
ProviderOptions: fantasy.ProviderOptions(result.ProviderMetadata),
|
||||
ToolCallID: result.ToolCallID,
|
||||
Output: result.Result,
|
||||
ProviderExecuted: result.ProviderExecuted,
|
||||
ProviderOptions: fantasy.ProviderOptions(result.ProviderMetadata),
|
||||
})
|
||||
default:
|
||||
continue
|
||||
@@ -205,7 +212,7 @@ func Run(ctx context.Context, opts RunOptions) error {
|
||||
opts.PublishMessagePart(role, part)
|
||||
}
|
||||
|
||||
tools := buildToolDefinitions(opts.Tools, opts.ActiveTools)
|
||||
tools := buildToolDefinitions(opts.Tools, opts.ActiveTools, opts.ProviderTools)
|
||||
applyAnthropicCaching := shouldApplyAnthropicPromptCaching(opts.Model)
|
||||
|
||||
messages := opts.Messages
|
||||
@@ -316,7 +323,6 @@ func Run(ctx context.Context, opts RunOptions) error {
|
||||
Valid: true,
|
||||
}
|
||||
}
|
||||
|
||||
// Persist the step — errors propagate directly.
|
||||
if err := opts.PersistStep(ctx, PersistedStep{
|
||||
Content: result.content,
|
||||
@@ -494,17 +500,19 @@ func processStepStream(
|
||||
}
|
||||
|
||||
case fantasy.StreamPartTypeToolInputDelta:
|
||||
var providerExecuted bool
|
||||
if toolCall, exists := activeToolCalls[part.ID]; exists {
|
||||
toolCall.Input += part.Delta
|
||||
providerExecuted = toolCall.ProviderExecuted
|
||||
}
|
||||
toolName := toolNames[part.ID]
|
||||
publishMessagePart(fantasy.MessageRoleAssistant, codersdk.ChatMessagePart{
|
||||
Type: codersdk.ChatMessagePartTypeToolCall,
|
||||
ToolCallID: part.ID,
|
||||
ToolName: toolName,
|
||||
ArgsDelta: part.Delta,
|
||||
Type: codersdk.ChatMessagePartTypeToolCall,
|
||||
ToolCallID: part.ID,
|
||||
ToolName: toolName,
|
||||
ArgsDelta: part.Delta,
|
||||
ProviderExecuted: providerExecuted,
|
||||
})
|
||||
|
||||
case fantasy.StreamPartTypeToolInputEnd:
|
||||
// No callback needed; the full tool call arrives in
|
||||
// StreamPartTypeToolCall.
|
||||
@@ -544,6 +552,24 @@ func processStepStream(
|
||||
chatprompt.PartFromContent(sourceContent),
|
||||
)
|
||||
|
||||
case fantasy.StreamPartTypeToolResult:
|
||||
// Provider-executed tool results (e.g. web search)
|
||||
// are emitted by the provider and added directly
|
||||
// to the step content for multi-turn round-tripping.
|
||||
// This mirrors fantasy's agent.go accumulation logic.
|
||||
if part.ProviderExecuted {
|
||||
tr := fantasy.ToolResultContent{
|
||||
ToolCallID: part.ID,
|
||||
ToolName: part.ToolCallName,
|
||||
ProviderExecuted: part.ProviderExecuted,
|
||||
ProviderMetadata: part.ProviderMetadata,
|
||||
}
|
||||
result.content = append(result.content, tr)
|
||||
publishMessagePart(
|
||||
fantasy.MessageRoleTool,
|
||||
chatprompt.PartFromContent(tr),
|
||||
)
|
||||
}
|
||||
case fantasy.StreamPartTypeFinish:
|
||||
result.usage = part.Usage
|
||||
result.finishReason = part.FinishReason
|
||||
@@ -571,7 +597,14 @@ func processStepStream(
|
||||
}
|
||||
}
|
||||
|
||||
result.shouldContinue = len(result.toolCalls) > 0 &&
|
||||
hasLocalToolCalls := false
|
||||
for _, tc := range result.toolCalls {
|
||||
if !tc.ProviderExecuted {
|
||||
hasLocalToolCalls = true
|
||||
break
|
||||
}
|
||||
}
|
||||
result.shouldContinue = hasLocalToolCalls &&
|
||||
result.finishReason == fantasy.FinishReasonToolCalls
|
||||
return result, nil
|
||||
}
|
||||
@@ -590,15 +623,29 @@ func executeTools(
|
||||
return nil
|
||||
}
|
||||
|
||||
// Filter out provider-executed tool calls. These were
|
||||
// handled server-side by the LLM provider (e.g., web
|
||||
// search) and their results are already in the stream
|
||||
// content.
|
||||
localToolCalls := make([]fantasy.ToolCallContent, 0, len(toolCalls))
|
||||
for _, tc := range toolCalls {
|
||||
if !tc.ProviderExecuted {
|
||||
localToolCalls = append(localToolCalls, tc)
|
||||
}
|
||||
}
|
||||
if len(localToolCalls) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
toolMap := make(map[string]fantasy.AgentTool, len(allTools))
|
||||
for _, t := range allTools {
|
||||
toolMap[t.Info().Name] = t
|
||||
}
|
||||
|
||||
results := make([]fantasy.ToolResultContent, len(toolCalls))
|
||||
results := make([]fantasy.ToolResultContent, len(localToolCalls))
|
||||
var wg sync.WaitGroup
|
||||
wg.Add(len(toolCalls))
|
||||
for i, tc := range toolCalls {
|
||||
wg.Add(len(localToolCalls))
|
||||
for i, tc := range localToolCalls {
|
||||
go func(i int, tc fantasy.ToolCallContent) {
|
||||
defer wg.Done()
|
||||
defer func() {
|
||||
@@ -770,8 +817,9 @@ func persistInterruptedStep(
|
||||
continue
|
||||
}
|
||||
content = append(content, fantasy.ToolResultContent{
|
||||
ToolCallID: tc.ToolCallID,
|
||||
ToolName: tc.ToolName,
|
||||
ToolCallID: tc.ToolCallID,
|
||||
ToolName: tc.ToolName,
|
||||
ProviderExecuted: tc.ProviderExecuted,
|
||||
Result: fantasy.ToolResultOutputContentError{
|
||||
Error: xerrors.New(interruptedToolResultErrorMessage),
|
||||
},
|
||||
@@ -791,9 +839,10 @@ func persistInterruptedStep(
|
||||
|
||||
// buildToolDefinitions converts AgentTool definitions into the
|
||||
// fantasy.Tool slice expected by fantasy.Call. When activeTools
|
||||
// is non-empty, only tools whose name appears in the list are
|
||||
// included. This mirrors fantasy's agent.prepareTools filtering.
|
||||
func buildToolDefinitions(tools []fantasy.AgentTool, activeTools []string) []fantasy.Tool {
|
||||
// is non-empty, only function tools whose name appears in the
|
||||
// list are included. Provider tools bypass this filter and are
|
||||
// always appended unconditionally.
|
||||
func buildToolDefinitions(tools []fantasy.AgentTool, activeTools []string, providerTools []fantasy.Tool) []fantasy.Tool {
|
||||
prepared := make([]fantasy.Tool, 0, len(tools))
|
||||
for _, tool := range tools {
|
||||
info := tool.Info()
|
||||
@@ -813,6 +862,7 @@ func buildToolDefinitions(tools []fantasy.AgentTool, activeTools []string) []fan
|
||||
ProviderOptions: tool.ProviderOptions(),
|
||||
})
|
||||
}
|
||||
prepared = append(prepared, providerTools...)
|
||||
return prepared
|
||||
}
|
||||
|
||||
|
||||
@@ -11,6 +11,7 @@ import (
|
||||
"github.com/sqlc-dev/pqtype"
|
||||
"golang.org/x/xerrors"
|
||||
|
||||
"cdr.dev/slog/v3"
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
)
|
||||
@@ -110,7 +111,7 @@ func patchFileContent(
|
||||
func ConvertMessages(
|
||||
messages []database.ChatMessage,
|
||||
) ([]fantasy.Message, error) {
|
||||
return ConvertMessagesWithFiles(context.Background(), messages, nil)
|
||||
return ConvertMessagesWithFiles(context.Background(), messages, nil, slog.Logger{})
|
||||
}
|
||||
|
||||
// ConvertMessagesWithFiles converts persisted chat messages into LLM
|
||||
@@ -121,6 +122,7 @@ func ConvertMessagesWithFiles(
|
||||
ctx context.Context,
|
||||
messages []database.ChatMessage,
|
||||
resolver FileResolver,
|
||||
logger slog.Logger,
|
||||
) ([]fantasy.Message, error) {
|
||||
// Phase 1: Pre-scan user messages for file_id references.
|
||||
var allFileIDs []uuid.UUID
|
||||
@@ -229,7 +231,7 @@ func ConvertMessagesWithFiles(
|
||||
if row.ToolCallID != "" && row.ToolName != "" {
|
||||
toolNameByCallID[sanitizeToolCallID(row.ToolCallID)] = row.ToolName
|
||||
}
|
||||
parts = append(parts, row.toToolResultPart())
|
||||
parts = append(parts, row.toToolResultPart(logger))
|
||||
}
|
||||
prompt = append(prompt, fantasy.Message{
|
||||
Role: fantasy.MessageRoleTool,
|
||||
@@ -359,10 +361,12 @@ func ParseContent(role string, raw pqtype.NullRawMessage) ([]fantasy.Content, er
|
||||
// result row. We intentionally avoid a strict Go struct so that
|
||||
// historical shapes are never rejected.
|
||||
type toolResultRaw struct {
|
||||
ToolCallID string `json:"tool_call_id"`
|
||||
ToolName string `json:"tool_name"`
|
||||
Result json.RawMessage `json:"result"`
|
||||
IsError bool `json:"is_error,omitempty"`
|
||||
ToolCallID string `json:"tool_call_id"`
|
||||
ToolName string `json:"tool_name"`
|
||||
Result json.RawMessage `json:"result"`
|
||||
IsError bool `json:"is_error,omitempty"`
|
||||
ProviderExecuted bool `json:"provider_executed,omitempty"`
|
||||
ProviderMetadata json.RawMessage `json:"provider_metadata,omitempty"`
|
||||
}
|
||||
|
||||
// parseToolResultRows decodes persisted tool result rows.
|
||||
@@ -378,7 +382,7 @@ func parseToolResultRows(raw pqtype.NullRawMessage) ([]toolResultRaw, error) {
|
||||
return rows, nil
|
||||
}
|
||||
|
||||
func (r toolResultRaw) toToolResultPart() fantasy.ToolResultPart {
|
||||
func (r toolResultRaw) toToolResultPart(logger slog.Logger) fantasy.ToolResultPart {
|
||||
toolCallID := sanitizeToolCallID(r.ToolCallID)
|
||||
resultText := string(r.Result)
|
||||
if resultText == "" || resultText == "null" {
|
||||
@@ -391,7 +395,9 @@ func (r toolResultRaw) toToolResultPart() fantasy.ToolResultPart {
|
||||
message = extracted
|
||||
}
|
||||
return fantasy.ToolResultPart{
|
||||
ToolCallID: toolCallID,
|
||||
ToolCallID: toolCallID,
|
||||
ProviderExecuted: r.ProviderExecuted,
|
||||
ProviderOptions: r.providerOptions(logger),
|
||||
Output: fantasy.ToolResultOutputContentError{
|
||||
Error: xerrors.New(message),
|
||||
},
|
||||
@@ -399,13 +405,43 @@ func (r toolResultRaw) toToolResultPart() fantasy.ToolResultPart {
|
||||
}
|
||||
|
||||
return fantasy.ToolResultPart{
|
||||
ToolCallID: toolCallID,
|
||||
ToolCallID: toolCallID,
|
||||
ProviderExecuted: r.ProviderExecuted,
|
||||
ProviderOptions: r.providerOptions(logger),
|
||||
Output: fantasy.ToolResultOutputContentText{
|
||||
Text: resultText,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// providerOptions deserializes the stored provider metadata
|
||||
// JSON into a ProviderOptions map using the fantasy type
|
||||
// registry. Returns nil when no metadata is stored.
|
||||
func (r toolResultRaw) providerOptions(logger slog.Logger) fantasy.ProviderOptions {
|
||||
if len(r.ProviderMetadata) == 0 {
|
||||
return nil
|
||||
}
|
||||
var raw map[string]json.RawMessage
|
||||
if err := json.Unmarshal(r.ProviderMetadata, &raw); err != nil {
|
||||
logger.Warn(context.Background(),
|
||||
"failed to unmarshal provider metadata JSON",
|
||||
slog.F("tool_call_id", r.ToolCallID),
|
||||
slog.Error(err),
|
||||
)
|
||||
return nil
|
||||
}
|
||||
opts, err := fantasy.UnmarshalProviderOptions(raw)
|
||||
if err != nil {
|
||||
logger.Warn(context.Background(),
|
||||
"failed to deserialize provider metadata",
|
||||
slog.F("tool_call_id", r.ToolCallID),
|
||||
slog.Error(err),
|
||||
)
|
||||
return nil
|
||||
}
|
||||
return opts
|
||||
}
|
||||
|
||||
// extractErrorString pulls the "error" field from a JSON object if
|
||||
// present, returning it as a string. Returns "" if the field is
|
||||
// missing or the input is not an object.
|
||||
@@ -609,12 +645,22 @@ func injectFileID(encoded json.RawMessage, fileID uuid.UUID) (json.RawMessage, e
|
||||
// MarshalToolResult encodes a single tool result for persistence as
|
||||
// an opaque JSON blob. The stored shape is
|
||||
// [{"tool_call_id":…,"tool_name":…,"result":…,"is_error":…}].
|
||||
func MarshalToolResult(toolCallID, toolName string, result json.RawMessage, isError bool) (pqtype.NullRawMessage, error) {
|
||||
func MarshalToolResult(toolCallID, toolName string, result json.RawMessage, isError bool, providerExecuted bool, providerMetadata fantasy.ProviderMetadata) (pqtype.NullRawMessage, error) {
|
||||
var metaJSON json.RawMessage
|
||||
if len(providerMetadata) > 0 {
|
||||
var err error
|
||||
metaJSON, err = json.Marshal(providerMetadata)
|
||||
if err != nil {
|
||||
return pqtype.NullRawMessage{}, xerrors.Errorf("encode provider metadata: %w", err)
|
||||
}
|
||||
}
|
||||
row := toolResultRaw{
|
||||
ToolCallID: toolCallID,
|
||||
ToolName: toolName,
|
||||
Result: result,
|
||||
IsError: isError,
|
||||
ToolCallID: toolCallID,
|
||||
ToolName: toolName,
|
||||
Result: result,
|
||||
IsError: isError,
|
||||
ProviderExecuted: providerExecuted,
|
||||
ProviderMetadata: metaJSON,
|
||||
}
|
||||
data, err := json.Marshal([]toolResultRaw{row})
|
||||
if err != nil {
|
||||
@@ -653,7 +699,7 @@ func MarshalToolResultContent(content fantasy.ToolResultContent) (pqtype.NullRaw
|
||||
result = []byte(`{}`)
|
||||
}
|
||||
|
||||
return MarshalToolResult(content.ToolCallID, content.ToolName, result, isError)
|
||||
return MarshalToolResult(content.ToolCallID, content.ToolName, result, isError, content.ProviderExecuted, content.ProviderMetadata)
|
||||
}
|
||||
|
||||
// PartFromContent converts fantasy content into a SDK chat message part.
|
||||
@@ -681,17 +727,19 @@ func PartFromContent(block fantasy.Content) codersdk.ChatMessagePart {
|
||||
}
|
||||
case fantasy.ToolCallContent:
|
||||
return codersdk.ChatMessagePart{
|
||||
Type: codersdk.ChatMessagePartTypeToolCall,
|
||||
ToolCallID: value.ToolCallID,
|
||||
ToolName: value.ToolName,
|
||||
Args: []byte(value.Input),
|
||||
Type: codersdk.ChatMessagePartTypeToolCall,
|
||||
ToolCallID: value.ToolCallID,
|
||||
ToolName: value.ToolName,
|
||||
Args: []byte(value.Input),
|
||||
ProviderExecuted: value.ProviderExecuted,
|
||||
}
|
||||
case *fantasy.ToolCallContent:
|
||||
return codersdk.ChatMessagePart{
|
||||
Type: codersdk.ChatMessagePartTypeToolCall,
|
||||
ToolCallID: value.ToolCallID,
|
||||
ToolName: value.ToolName,
|
||||
Args: []byte(value.Input),
|
||||
Type: codersdk.ChatMessagePartTypeToolCall,
|
||||
ToolCallID: value.ToolCallID,
|
||||
ToolName: value.ToolName,
|
||||
Args: []byte(value.Input),
|
||||
ProviderExecuted: value.ProviderExecuted,
|
||||
}
|
||||
case fantasy.SourceContent:
|
||||
return codersdk.ChatMessagePart{
|
||||
@@ -771,7 +819,9 @@ func toolResultContentToPart(content fantasy.ToolResultContent) codersdk.ChatMes
|
||||
result = []byte(`{}`)
|
||||
}
|
||||
|
||||
return ToolResultToPart(content.ToolCallID, content.ToolName, result, isError)
|
||||
part := ToolResultToPart(content.ToolCallID, content.ToolName, result, isError)
|
||||
part.ProviderExecuted = content.ProviderExecuted
|
||||
return part
|
||||
}
|
||||
|
||||
func injectMissingToolResults(prompt []fantasy.Message) []fantasy.Message {
|
||||
|
||||
@@ -10,6 +10,7 @@ import (
|
||||
"github.com/sqlc-dev/pqtype"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"cdr.dev/slog/v3/sloggers/slogtest"
|
||||
"github.com/coder/coder/v2/coderd/chatd/chatprompt"
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
)
|
||||
@@ -63,6 +64,8 @@ func TestConvertMessages_NormalizesAssistantToolCallInput(t *testing.T) {
|
||||
"execute",
|
||||
json.RawMessage(`{"error":"tool call was interrupted before it produced a result"}`),
|
||||
true,
|
||||
false,
|
||||
nil,
|
||||
)
|
||||
require.NoError(t, err)
|
||||
|
||||
@@ -134,6 +137,7 @@ func TestConvertMessagesWithFiles_ResolvesFileData(t *testing.T) {
|
||||
},
|
||||
},
|
||||
resolver,
|
||||
slogtest.Make(t, nil),
|
||||
)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, prompt, 1)
|
||||
@@ -175,6 +179,7 @@ func TestConvertMessagesWithFiles_BackwardCompat(t *testing.T) {
|
||||
},
|
||||
},
|
||||
nil, // No resolver.
|
||||
slogtest.Make(t, nil),
|
||||
)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, prompt, 1)
|
||||
|
||||
@@ -1193,11 +1193,12 @@ func chatMessageParts(role string, raw pqtype.NullRawMessage) ([]codersdk.ChatMe
|
||||
parts := make([]codersdk.ChatMessagePart, 0, len(results))
|
||||
for _, result := range results {
|
||||
parts = append(parts, codersdk.ChatMessagePart{
|
||||
Type: codersdk.ChatMessagePartTypeToolResult,
|
||||
ToolCallID: result.ToolCallID,
|
||||
ToolName: result.ToolName,
|
||||
Result: result.Result,
|
||||
IsError: result.IsError,
|
||||
Type: codersdk.ChatMessagePartTypeToolResult,
|
||||
ToolCallID: result.ToolCallID,
|
||||
ToolName: result.ToolName,
|
||||
Result: result.Result,
|
||||
IsError: result.IsError,
|
||||
ProviderExecuted: result.ProviderExecuted,
|
||||
})
|
||||
}
|
||||
return parts, nil
|
||||
@@ -1251,10 +1252,11 @@ func parseContentBlocks(role string, raw pqtype.NullRawMessage) ([]fantasy.Conte
|
||||
// toolResultRow is used only for extracting top-level fields from
|
||||
// persisted tool result JSON. The result payload is kept as raw JSON.
|
||||
type toolResultRow struct {
|
||||
ToolCallID string `json:"tool_call_id"`
|
||||
ToolName string `json:"tool_name"`
|
||||
Result json.RawMessage `json:"result"`
|
||||
IsError bool `json:"is_error,omitempty"`
|
||||
ToolCallID string `json:"tool_call_id"`
|
||||
ToolName string `json:"tool_name"`
|
||||
Result json.RawMessage `json:"result"`
|
||||
IsError bool `json:"is_error,omitempty"`
|
||||
ProviderExecuted bool `json:"provider_executed,omitempty"`
|
||||
}
|
||||
|
||||
func parseToolResults(raw pqtype.NullRawMessage) ([]toolResultRow, error) {
|
||||
@@ -1293,17 +1295,19 @@ func contentBlockToPart(block fantasy.Content) codersdk.ChatMessagePart {
|
||||
}
|
||||
case fantasy.ToolCallContent:
|
||||
return codersdk.ChatMessagePart{
|
||||
Type: codersdk.ChatMessagePartTypeToolCall,
|
||||
ToolCallID: value.ToolCallID,
|
||||
ToolName: value.ToolName,
|
||||
Args: []byte(value.Input),
|
||||
Type: codersdk.ChatMessagePartTypeToolCall,
|
||||
ToolCallID: value.ToolCallID,
|
||||
ToolName: value.ToolName,
|
||||
Args: []byte(value.Input),
|
||||
ProviderExecuted: value.ProviderExecuted,
|
||||
}
|
||||
case *fantasy.ToolCallContent:
|
||||
return codersdk.ChatMessagePart{
|
||||
Type: codersdk.ChatMessagePartTypeToolCall,
|
||||
ToolCallID: value.ToolCallID,
|
||||
ToolName: value.ToolName,
|
||||
Args: []byte(value.Input),
|
||||
Type: codersdk.ChatMessagePartTypeToolCall,
|
||||
ToolCallID: value.ToolCallID,
|
||||
ToolName: value.ToolName,
|
||||
Args: []byte(value.Input),
|
||||
ProviderExecuted: value.ProviderExecuted,
|
||||
}
|
||||
case fantasy.SourceContent:
|
||||
return codersdk.ChatMessagePart{
|
||||
|
||||
Reference in New Issue
Block a user