feat: extend agent chat MCP tools for remote UAT evidence loops (#28233)

Extends the Agent-chat MCP tools so an unattended UAT evidence loop can
fetch artifacts, monitor long runs, and find prior runs without burning
model context.

## Backend

- New `chat_files_token` crypto key feature (migration 000571) with
rotator support and a dedicated signing keycache on coderd.
- `POST /api/experimental/chats/files/{file}/download-url`
(authenticated) mints a short-lived (5 min) signed URL and returns it
with `sha256`, `size_bytes`, `name`, `mime_type`, and `expires_at`.
- `GET /api/experimental/chats/files/{file}/download?token=` (no session
token) redeems the signed URL: verifies the JWS, requires the token's
`file_id` to match the path, and re-checks the minting user's RBAC
access live at redemption. Clients can `curl -o` artifacts with zero
credentials in the URL consumer.
- `ChatFileMetadata` gains `size_bytes` (via `octet_length`, no bytes
fetched).

## MCP tools (`codersdk/toolsdk`)

- `coder_download_chat_file`: by `file_id` or `chat_id`+`file_name`;
returns the signed URL plus checksum and size instead of base64.
- `coder_await_chat`: blocks (bounded `wait_secs`, 1-120) until a chat
leaves `running`/`interrupting`, using the existing watch stream with
subscribe-before-read.
- `coder_list_chats`: label, query, and limit filtering; chat
projections now include labels.
- `coder_get_chat_messages`: `after_id` forward cursor with
`next_after_id` (exact incremental reads), plus per-message `files`
metadata so artifact-bearing messages are identifiable.
- `coder_get_chat`: file listings now include `size_bytes` and
`created_at`.
- `coder_list_templates`: exposes `agents_allowed` for pre-flight
checks.

## Testing

- coderd: mint/redeem happy path with an unauthenticated client,
expired/tampered/file-mismatched tokens, auth still required on the
plain file endpoint, non-owner mint rejection.
- toolsdk: harness + integration coverage for all new/changed tools,
including signed-URL redemption with checksum verification,
forward-cursor exactness, await transition/timeout paths, and label
filtering.
- Remote dogfood UAT (dev.coder.com Coder Agent) passed all six
acceptance scenarios end to end over both MCP transports.

Note: `go test ./codersdk/toolsdk/` has a pre-existing goleak flake on
main (leaked `agentssh` non-PTY session goroutines from SSH exec tests;
reproduced 3/3 on clean `b4971bc49f1`). It is unrelated to this diff.

> Mux acted on Mike's behalf to create this PR.

<!-- mux-attribution: model=claude-sonnet-4-6 thinking=high -->
This commit is contained in:
Michael Suchacz
2026-08-18 19:15:30 +02:00
committed by GitHub
parent 7724ee281a
commit affeeaf9c8
30 changed files with 1840 additions and 91 deletions
+25
View File
@@ -242,6 +242,7 @@ type ChatFileMetadata struct {
OrganizationID uuid.UUID `json:"organization_id" format:"uuid"`
Name string `json:"name"`
MimeType string `json:"mime_type"`
SizeBytes int64 `json:"size_bytes"`
CreatedAt time.Time `json:"created_at" format:"date-time"`
}
@@ -677,6 +678,16 @@ type UploadChatFileResponse struct {
ID uuid.UUID `json:"id" format:"uuid"`
}
// ChatFileDownloadURLResponse contains a short-lived URL for downloading a chat file.
type ChatFileDownloadURLResponse struct {
URL string `json:"url" format:"uri"`
ExpiresAt time.Time `json:"expires_at" format:"date-time"`
SHA256 string `json:"sha256"`
SizeBytes int64 `json:"size_bytes"`
Name string `json:"name"`
MimeType string `json:"mime_type"`
}
// ChatMessagesResponse contains the messages and queued messages for a chat.
type ChatMessagesResponse struct {
Messages []ChatMessage `json:"messages"`
@@ -3167,6 +3178,20 @@ func (c *ExperimentalClient) UploadChatFile(ctx context.Context, organizationID
return resp, ReadBodyAsJSON(res, &resp)
}
// ChatFileDownloadURL creates a short-lived download URL for a chat file.
func (c *ExperimentalClient) ChatFileDownloadURL(ctx context.Context, fileID uuid.UUID) (ChatFileDownloadURLResponse, error) {
res, err := c.Request(ctx, http.MethodPost, fmt.Sprintf("/api/experimental/chats/files/%s/download-url", fileID), nil)
if err != nil {
return ChatFileDownloadURLResponse{}, err
}
defer res.Body.Close()
if res.StatusCode != http.StatusOK {
return ChatFileDownloadURLResponse{}, ReadBodyAsError(res)
}
var resp ChatFileDownloadURLResponse
return resp, ReadBodyAsJSON(res, &resp)
}
// GetChatFile retrieves a previously uploaded chat file by ID.
func (c *ExperimentalClient) GetChatFile(ctx context.Context, fileID uuid.UUID) ([]byte, string, error) {
res, err := c.Request(ctx, http.MethodGet, fmt.Sprintf("/api/experimental/chats/files/%s", fileID), nil)
+1
View File
@@ -5694,6 +5694,7 @@ const (
//nolint:gosec // This denotes a type of key, not a literal.
CryptoKeyFeatureWorkspaceAppsToken CryptoKeyFeature = "workspace_apps_token"
CryptoKeyFeatureOIDCConvert CryptoKeyFeature = "oidc_convert"
CryptoKeyFeatureChatFilesToken CryptoKeyFeature = "chat_files_token"
CryptoKeyFeatureTailnetResume CryptoKeyFeature = "tailnet_resume"
// CryptoKeyFeatureNATSCA is the CA that signs NATS cluster mTLS leaf
// certificates. Its secret is a PEM cert+key bundle (not a hex secret like
+437 -33
View File
@@ -4,6 +4,7 @@ import (
"context"
"errors"
"fmt"
"io"
"net/http"
"strings"
"time"
@@ -34,9 +35,11 @@ func parseChatID(chatID string) (uuid.UUID, error) {
}
type ChatToolFile struct {
ID string `json:"id"`
Name string `json:"name"`
MimeType string `json:"mime_type"`
ID string `json:"id"`
Name string `json:"name"`
MimeType string `json:"mime_type"`
SizeBytes int64 `json:"size_bytes"`
CreatedAt time.Time `json:"created_at"`
}
type ChatToolStatus struct {
@@ -48,6 +51,7 @@ type ChatToolStatus struct {
LastTurnSummary string `json:"last_turn_summary,omitempty"`
WorkspaceID string `json:"workspace_id,omitempty"`
URL string `json:"url"`
Labels map[string]string `json:"labels,omitempty"`
Files []ChatToolFile `json:"files,omitempty"`
}
@@ -59,6 +63,7 @@ func chatToolStatus(deps Deps, chat codersdk.Chat) ChatToolStatus {
Archived: chat.Archived,
LastError: chat.LastError,
URL: fmt.Sprintf("%s/agents/%s", deps.ServerURL(), chat.ID),
Labels: chat.Labels,
}
if chat.LastTurnSummary != nil {
resp.LastTurnSummary = *chat.LastTurnSummary
@@ -68,9 +73,11 @@ func chatToolStatus(deps Deps, chat codersdk.Chat) ChatToolStatus {
}
for _, file := range chat.Files {
resp.Files = append(resp.Files, ChatToolFile{
ID: file.ID.String(),
Name: file.Name,
MimeType: file.MimeType,
ID: file.ID.String(),
Name: file.Name,
MimeType: file.MimeType,
SizeBytes: file.SizeBytes,
CreatedAt: file.CreatedAt,
})
}
return resp
@@ -191,10 +198,345 @@ var GetChat = Tool[GetChatArgs, ChatToolStatus]{
},
}
type DownloadChatFileArgs struct {
FileID string `json:"file_id"`
ChatID string `json:"chat_id"`
FileName string `json:"file_name"`
}
type DownloadChatFileResponse struct {
FileID string `json:"file_id"`
Name string `json:"name"`
MimeType string `json:"mime_type"`
SizeBytes int64 `json:"size_bytes"`
SHA256 string `json:"sha256"`
URL string `json:"url"`
ExpiresAt time.Time `json:"expires_at"`
}
func chatFilesDescription(files []codersdk.ChatFileMetadata) string {
descriptions := make([]string, len(files))
for i, file := range files {
descriptions[i] = fmt.Sprintf("{id: %s, name: %q, mime_type: %q, size_bytes: %d}", file.ID, file.Name, file.MimeType, file.SizeBytes)
}
return "[" + strings.Join(descriptions, ", ") + "]"
}
var DownloadChatFile = Tool[DownloadChatFileArgs, DownloadChatFileResponse]{
Tool: aisdk.Tool{
Name: ToolNameDownloadChatFile,
Description: `Create a short-lived download URL for a file attached to a Coder Agents chat.
Address the file with file_id alone, or with chat_id and an exact file_name. The URL expires in about 5 minutes and needs no authentication header. Fetch it with curl -fSs -o <path> "<url>". Do not read binary contents into context.`,
Schema: aisdk.Schema{
Properties: map[string]any{
"file_id": map[string]any{
"type": "string",
"description": "Optional chat file UUID. Use this alone when the file ID is known.",
},
"chat_id": map[string]any{
"type": "string",
"description": "Optional chat UUID. Use together with file_name when the file ID is unknown.",
},
"file_name": map[string]any{
"type": "string",
"description": "Optional exact file name. Use together with chat_id.",
},
},
Required: []string{},
},
},
MCPAnnotations: mcpReadOnlyAnnotations,
Handler: func(ctx context.Context, deps Deps, args DownloadChatFileArgs) (DownloadChatFileResponse, error) {
fileIDMode := args.FileID != "" && args.ChatID == "" && args.FileName == ""
chatFileMode := args.FileID == "" && args.ChatID != "" && args.FileName != ""
if !fileIDMode && !chatFileMode {
return DownloadChatFileResponse{}, xerrors.New("provide exactly one addressing mode: file_id alone, or chat_id with file_name")
}
var fileID uuid.UUID
if fileIDMode {
var err error
fileID, err = uuid.Parse(args.FileID)
if err != nil {
return DownloadChatFileResponse{}, xerrors.New("file_id must be a valid UUID")
}
} else {
chatID, err := parseChatID(args.ChatID)
if err != nil {
return DownloadChatFileResponse{}, err
}
chat, err := codersdk.NewExperimentalClient(deps.coderClient).GetChat(ctx, chatID)
if err != nil {
return DownloadChatFileResponse{}, xerrors.Errorf("get chat: %w", err)
}
found := false
for _, file := range chat.Files {
if file.Name != args.FileName {
continue
}
if found {
return DownloadChatFileResponse{}, xerrors.Errorf("multiple chat files named %q; available files: %s", args.FileName, chatFilesDescription(chat.Files))
}
fileID = file.ID
found = true
}
if !found {
return DownloadChatFileResponse{}, xerrors.Errorf("no chat file named %q; available files: %s", args.FileName, chatFilesDescription(chat.Files))
}
}
download, err := codersdk.NewExperimentalClient(deps.coderClient).ChatFileDownloadURL(ctx, fileID)
if err != nil {
return DownloadChatFileResponse{}, xerrors.Errorf("create chat file download URL: %w", err)
}
return DownloadChatFileResponse{
FileID: fileID.String(),
Name: download.Name,
MimeType: download.MimeType,
SizeBytes: download.SizeBytes,
SHA256: download.SHA256,
URL: download.URL,
ExpiresAt: download.ExpiresAt,
}, nil
},
}
type AwaitChatArgs struct {
ChatID string `json:"chat_id"`
WaitSecs int `json:"wait_secs"`
}
type AwaitChatResponse struct {
TimedOut bool `json:"timed_out"`
Chat ChatToolStatus `json:"chat"`
}
func chatStatusBusy(status codersdk.ChatStatus) bool {
return status == codersdk.ChatStatusRunning || status == codersdk.ChatStatusInterrupting
}
var AwaitChat = Tool[AwaitChatArgs, AwaitChatResponse]{
Tool: aisdk.Tool{
Name: ToolNameAwaitChat,
Description: `Block until a Coder Agents chat stops generating or the wait times out. Waiting, error, and requires_action all end the wait. If timed_out is true, chat holds the last status observed inside the wait window; call this tool again to continue waiting.`,
Schema: aisdk.Schema{
Properties: map[string]any{
"chat_id": map[string]any{
"type": "string",
"description": chatIDDescription,
},
"wait_secs": map[string]any{
"type": "integer",
"description": "Maximum seconds to wait (1-120, default 60).",
},
},
Required: []string{"chat_id"},
},
},
MCPAnnotations: mcpReadOnlyAnnotations,
Handler: func(ctx context.Context, deps Deps, args AwaitChatArgs) (AwaitChatResponse, error) {
chatID, err := parseChatID(args.ChatID)
if err != nil {
return AwaitChatResponse{}, err
}
if args.WaitSecs < 0 || args.WaitSecs > 120 {
return AwaitChatResponse{}, xerrors.New("wait_secs must be between 1 and 120")
}
waitSecs := args.WaitSecs
if waitSecs == 0 {
waitSecs = 60
}
expClient := codersdk.NewExperimentalClient(deps.coderClient)
// Every request in the wait window runs under this deadline so a
// stalled websocket upgrade or REST call cannot extend the wait
// past wait_secs.
waitCtx, cancelWait := context.WithTimeout(ctx, time.Duration(waitSecs)*time.Second)
defer cancelWait()
// lastBusy is the most recent (busy) chat state observed inside
// the wait window; non-busy states return immediately instead.
var lastBusy *codersdk.Chat
// finalStatus reports the last busy state observed inside the
// wait window once it closes. Every caller runs after lastBusy
// is set, and no request runs after the window ends, so
// wait_secs stays a hard upper bound on the tool's duration.
finalStatus := func() (AwaitChatResponse, error) {
if ctx.Err() != nil {
return AwaitChatResponse{}, ctx.Err()
}
return AwaitChatResponse{TimedOut: true, Chat: chatToolStatus(deps, *lastBusy)}, nil
}
// Dial asynchronously: a slow or failed watch dial (e.g. a proxy
// stalling or rejecting upgrades) must not block the REST poller.
// The events channel stays nil until the dial succeeds; a missed
// transition in that window is caught by the next poll tick.
type watchDial struct {
events <-chan codersdk.ChatWatchEvent
closer io.Closer
}
dialed := make(chan watchDial, 1)
go func() {
events, closer, err := expClient.WatchChats(waitCtx)
if err != nil {
return
}
dialed <- watchDial{events: events, closer: closer}
}()
var (
events <-chan codersdk.ChatWatchEvent
watchCloser io.Closer
)
defer func() {
cancelWait()
if watchCloser == nil {
select {
case dial := <-dialed:
watchCloser = dial.closer
default:
}
}
if watchCloser != nil {
_ = watchCloser.Close()
}
}()
// The initial status request errors when it outlives the wait
// window: no state was observed, and a post-deadline fetch would
// break the wait_secs bound.
chat, err := expClient.GetChat(waitCtx, chatID)
if err != nil {
return AwaitChatResponse{}, xerrors.Errorf("get chat: %w", err)
}
if !chatStatusBusy(chat.Status) {
return AwaitChatResponse{Chat: chatToolStatus(deps, chat)}, nil
}
lastBusy = &chat
// Chat status events are published only on the owner's channel, so
// shared-chat callers need polling to observe transitions.
poller := time.NewTicker(5 * time.Second)
defer poller.Stop()
for {
select {
case <-waitCtx.Done():
return finalStatus()
case <-poller.C:
chat, err := expClient.GetChat(waitCtx, chatID)
if err != nil {
if waitCtx.Err() != nil {
return finalStatus()
}
return AwaitChatResponse{}, xerrors.Errorf("get chat while polling: %w", err)
}
if !chatStatusBusy(chat.Status) {
return AwaitChatResponse{Chat: chatToolStatus(deps, chat)}, nil
}
lastBusy = &chat
case dial := <-dialed:
events = dial.events
watchCloser = dial.closer
dialed = nil
case event, ok := <-events:
if !ok {
// A dropped watch stream must not end the wait early;
// the poll ticker keeps observing until the window closes.
events = nil
continue
}
if event.Chat.ID == chatID && !chatStatusBusy(event.Chat.Status) {
chat, err := expClient.GetChat(waitCtx, chatID)
if err != nil {
if waitCtx.Err() != nil {
return finalStatus()
}
return AwaitChatResponse{}, xerrors.Errorf("get chat after status change: %w", err)
}
// A new turn may start between the event and this
// confirmation; keep waiting if the chat is busy again.
if !chatStatusBusy(chat.Status) {
return AwaitChatResponse{Chat: chatToolStatus(deps, chat)}, nil
}
lastBusy = &chat
}
}
}
},
}
type ListChatsArgs struct {
Labels map[string]string `json:"labels"`
Query string `json:"query"`
Limit int `json:"limit"`
}
type ListChatsResponse struct {
Chats []ChatToolStatus `json:"chats"`
}
var ListChats = Tool[ListChatsArgs, ListChatsResponse]{
Tool: aisdk.Tool{
Name: ToolNameListChats,
Description: `List Coder Agents chats, optionally filtered by labels or a search query.`,
Schema: aisdk.Schema{
Properties: map[string]any{
"labels": map[string]any{
"type": "object",
"description": "Optional exact-match string key/value labels.",
"additionalProperties": map[string]any{"type": "string"},
},
"query": map[string]any{
"type": "string",
"description": "Optional chat search query using fielded terms; bare text is rejected. Supported fields: search:<text> (full-text, cannot combine with title, pr_title, or pr), title:<text>, repo:<owner/name>, pr:<number>, pr_title:<text>, pr_status:<draft|open|merged|closed>, diff_url:<url>, archived:<true|false>, has_unread:<true|false>, source:<created_by_me|shared_with_me>. Quote values containing spaces or colons (URLs always need quoting), e.g. search:\"failed deployment\" or diff_url:\"https://github.com/org/repo/pull/1\".",
},
"limit": map[string]any{
"type": "integer",
"description": "Maximum chats to return (1-100, default 25).",
},
},
Required: []string{},
},
},
MCPAnnotations: mcpReadOnlyAnnotations,
Handler: func(ctx context.Context, deps Deps, args ListChatsArgs) (ListChatsResponse, error) {
if args.Limit < 0 || args.Limit > 100 {
return ListChatsResponse{}, xerrors.New("limit must be between 1 and 100")
}
limit := args.Limit
if limit == 0 {
limit = 25
}
chats, err := codersdk.NewExperimentalClient(deps.coderClient).ListChats(ctx, &codersdk.ListChatsOptions{
Query: args.Query,
Labels: args.Labels,
Pagination: codersdk.Pagination{
Limit: limit,
},
})
if err != nil {
return ListChatsResponse{}, xerrors.Errorf("list chats: %w", err)
}
resp := ListChatsResponse{Chats: make([]ChatToolStatus, len(chats))}
for i, chat := range chats {
resp.Chats[i] = chatToolStatus(deps, chat)
}
return resp, nil
},
}
type GetChatMessagesArgs struct {
ChatID string `json:"chat_id"`
Limit int `json:"limit"`
BeforeID int64 `json:"before_id"`
AfterID int64 `json:"after_id"`
}
type ChatToolMessageFile struct {
ID string `json:"id"`
Name string `json:"name"`
MimeType string `json:"mime_type"`
}
type ChatToolMessage struct {
@@ -202,15 +544,15 @@ type ChatToolMessage struct {
Role codersdk.ChatMessageRole `json:"role"`
CreatedAt time.Time `json:"created_at"`
Text string `json:"text"`
Files []ChatToolMessageFile `json:"files,omitempty"`
}
type GetChatMessagesResponse struct {
Messages []ChatToolMessage `json:"messages"`
HasMore bool `json:"has_more"`
// NextBeforeID is the cursor for the next older page when HasMore is
// true. It is derived from the unfiltered API page, so it stays valid
// even when every message in this page was filtered out as non-text.
// Cursors come from the raw page so filtered pages remain traversable.
NextBeforeID int64 `json:"next_before_id,omitempty"`
NextAfterID int64 `json:"next_after_id,omitempty"`
// QueuedMessages is populated only on the initial page.
QueuedMessages []string `json:"queued_messages,omitempty"`
}
@@ -228,12 +570,34 @@ func userFacingText(parts []codersdk.ChatMessagePart) string {
return strings.Join(texts, "\n")
}
func chatToolMessage(msg codersdk.ChatMessage) (ChatToolMessage, bool) {
toolMessage := ChatToolMessage{
ID: msg.ID,
Role: msg.Role,
CreatedAt: msg.CreatedAt,
Text: userFacingText(msg.Content),
}
for _, part := range msg.Content {
if part.Type == codersdk.ChatMessagePartTypeFile && part.FileID.Valid {
toolMessage.Files = append(toolMessage.Files, ChatToolMessageFile{
ID: part.FileID.UUID.String(),
Name: part.Name,
MimeType: part.MediaType,
})
}
}
if toolMessage.Text == "" && len(toolMessage.Files) == 0 {
return ChatToolMessage{}, false
}
return toolMessage, true
}
var GetChatMessages = Tool[GetChatMessagesArgs, GetChatMessagesResponse]{
Tool: aisdk.Tool{
Name: ToolNameGetChatMessages,
Description: `Get the newest messages of a Coder Agents chat in chronological order.
Description: `Get messages from a Coder Agents chat in chronological order.
Only user-facing text content is returned (including lifecycle hook notices); tool calls and other internal parts are omitted. Prompts still queued behind a busy chat appear in queued_messages. When has_more is true, pass next_before_id as before_id to page through older messages.`,
Only user-facing text content is returned (including lifecycle hook notices); tool calls and other internal parts are omitted. Prompts still queued behind a busy chat appear in queued_messages. Use before_id with next_before_id to page backward from the newest messages, or after_id with next_after_id to page forward.`,
Schema: aisdk.Schema{
Properties: map[string]any{
"chat_id": map[string]any{
@@ -242,12 +606,16 @@ Only user-facing text content is returned (including lifecycle hook notices); to
},
"limit": map[string]any{
"type": "integer",
"description": "Maximum number of messages to fetch, from newest to oldest (1-200, default 50).",
"description": "Maximum number of messages per page (1-200, default 50). Pages are newest-first unless after_id is set, which pages forward in chronological order.",
},
"before_id": map[string]any{
"type": "integer",
"description": "Only fetch messages with an id lower than this cursor. Omit to fetch the newest messages.",
},
"after_id": map[string]any{
"type": "integer",
"description": "Only fetch messages with an id greater than this cursor, in chronological order. Cannot be combined with before_id.",
},
},
Required: []string{"chat_id"},
},
@@ -264,45 +632,80 @@ Only user-facing text content is returned (including lifecycle hook notices); to
if args.BeforeID < 0 {
return GetChatMessagesResponse{}, xerrors.New("before_id must be a positive message id")
}
if args.AfterID < 0 {
return GetChatMessagesResponse{}, xerrors.New("after_id must be a positive message id")
}
if args.BeforeID > 0 && args.AfterID > 0 {
return GetChatMessagesResponse{}, xerrors.New("before_id and after_id cannot be used together")
}
var opts *codersdk.ChatMessagesPaginationOptions
if args.Limit > 0 || args.BeforeID > 0 {
if args.Limit > 0 || args.BeforeID > 0 || args.AfterID > 0 {
opts = &codersdk.ChatMessagesPaginationOptions{
Limit: args.Limit,
BeforeID: args.BeforeID,
AfterID: args.AfterID,
}
}
resp, err := codersdk.NewExperimentalClient(deps.coderClient).GetChatMessages(ctx, chatID, opts)
if err != nil {
return GetChatMessagesResponse{}, xerrors.Errorf("get chat messages: %w", err)
}
// The API returns messages newest first; reverse into
// chronological order so the transcript reads naturally.
messages := make([]ChatToolMessage, 0, len(resp.Messages))
for i := len(resp.Messages) - 1; i >= 0; i-- {
msg := resp.Messages[i]
text := userFacingText(msg.Content)
if text == "" {
continue
if args.AfterID > 0 {
for _, msg := range resp.Messages {
if toolMessage, ok := chatToolMessage(msg); ok {
messages = append(messages, toolMessage)
}
}
} else {
for i := len(resp.Messages) - 1; i >= 0; i-- {
if toolMessage, ok := chatToolMessage(resp.Messages[i]); ok {
messages = append(messages, toolMessage)
}
}
messages = append(messages, ChatToolMessage{
ID: msg.ID,
Role: msg.Role,
CreatedAt: msg.CreatedAt,
Text: text,
})
}
var queued []string
for _, msg := range resp.QueuedMessages {
if text := userFacingText(msg.Content); text != "" {
text := userFacingText(msg.Content)
if text == "" {
// A queued prompt can carry only file parts; represent it
// by its attachments instead of dropping it.
var names []string
for _, part := range msg.Content {
if part.Type == codersdk.ChatMessagePartTypeFile && part.FileID.Valid {
name := part.Name
if name == "" {
name = part.FileID.UUID.String()
}
names = append(names, name)
}
}
if len(names) > 0 {
text = "(attached files: " + strings.Join(names, ", ") + ")"
}
}
if text != "" {
queued = append(queued, text)
}
}
var nextBeforeID int64
if resp.HasMore && len(resp.Messages) > 0 {
nextBeforeID = resp.Messages[0].ID
for _, msg := range resp.Messages {
if msg.ID < nextBeforeID {
nextBeforeID = msg.ID
var nextBeforeID, nextAfterID int64
if len(resp.Messages) > 0 {
if args.AfterID > 0 {
// Forward pollers need a cursor from every nonempty page,
// even the last one, or a page of filtered-out internal
// messages would leave them stuck replaying the same page.
nextAfterID = resp.Messages[0].ID
for _, msg := range resp.Messages {
if msg.ID > nextAfterID {
nextAfterID = msg.ID
}
}
} else if resp.HasMore {
nextBeforeID = resp.Messages[0].ID
for _, msg := range resp.Messages {
if msg.ID < nextBeforeID {
nextBeforeID = msg.ID
}
}
}
}
@@ -310,6 +713,7 @@ Only user-facing text content is returned (including lifecycle hook notices); to
Messages: messages,
HasMore: resp.HasMore,
NextBeforeID: nextBeforeID,
NextAfterID: nextAfterID,
QueuedMessages: queued,
}, nil
},
+526 -2
View File
@@ -1,8 +1,15 @@
package toolsdk_test
import (
"bytes"
"context"
"crypto/sha256"
"fmt"
"io"
"net/http"
"sync"
"testing"
"time"
"github.com/google/uuid"
"github.com/stretchr/testify/require"
@@ -12,6 +19,7 @@ import (
"github.com/coder/coder/v2/coderd/coderdtest"
"github.com/coder/coder/v2/coderd/database"
"github.com/coder/coder/v2/coderd/database/dbgen"
"github.com/coder/coder/v2/coderd/jwtutils"
"github.com/coder/coder/v2/coderd/rbac"
"github.com/coder/coder/v2/coderd/util/ptr"
"github.com/coder/coder/v2/coderd/x/chatd/chatprompt"
@@ -32,6 +40,36 @@ func (t *failPathTransport) RoundTrip(req *http.Request) (*http.Response, error)
return http.DefaultTransport.RoundTrip(req)
}
// stallTransport hangs every request until its context is canceled.
type stallTransport struct{}
func (stallTransport) RoundTrip(req *http.Request) (*http.Response, error) {
<-req.Context().Done()
return nil, req.Context().Err()
}
type signalPathTransport struct {
path string
seen chan struct{}
release chan struct{}
once sync.Once
}
func (t *signalPathTransport) RoundTrip(req *http.Request) (*http.Response, error) {
res, err := http.DefaultTransport.RoundTrip(req)
if err != nil || req.URL.Path != t.path {
return res, err
}
t.once.Do(func() { close(t.seen) })
select {
case <-t.release:
return res, nil
case <-req.Context().Done():
_ = res.Body.Close()
return nil, req.Context().Err()
}
}
// Chat tools need a chat-enabled coderd (provider keys, a default model
// config, and an AI bridge daemon), so they are tested separately from
// TestTools. Subtests run sequentially and share the deployment.
@@ -39,8 +77,9 @@ func (t *failPathTransport) RoundTrip(req *http.Request) (*http.Response, error)
func TestChatTools(t *testing.T) {
providerKeys := coderdtest.FakeOpenAICompatProviderAPIKeys(t)
client, _, api := coderdtest.NewWithAPI(t, &coderdtest.Options{
DeploymentValues: coderdtest.DeploymentValues(t),
ChatProviderAPIKeys: &providerKeys,
DeploymentValues: coderdtest.DeploymentValues(t),
ChatProviderAPIKeys: &providerKeys,
ChatFileTokenKeyCache: jwtutils.StaticKey{ID: "1", Key: bytes.Repeat([]byte("k"), 64)},
})
firstUser := coderdtest.CreateFirstUser(t, client)
expClient := codersdk.NewExperimentalClient(client)
@@ -164,6 +203,464 @@ func TestChatTools(t *testing.T) {
require.True(t, got.Archived)
})
t.Run("DownloadChatFile", func(t *testing.T) {
ctx := testutil.Context(t, testutil.WaitLong)
data := []byte("MCP chat UAT evidence\n")
uploaded, err := expClient.UploadChatFile(ctx, firstUser.OrganizationID, "text/plain", "evidence.txt", bytes.NewReader(data))
require.NoError(t, err)
chat, err := expClient.CreateChat(ctx, codersdk.CreateChatRequest{
OrganizationID: firstUser.OrganizationID,
Content: []codersdk.ChatInputPart{
{Type: codersdk.ChatInputPartTypeText, Text: "Review the evidence."},
{Type: codersdk.ChatInputPartTypeFile, FileID: uploaded.ID},
},
})
require.NoError(t, err)
assertDownload := func(t *testing.T, got toolsdk.DownloadChatFileResponse) {
t.Helper()
require.Equal(t, uploaded.ID.String(), got.FileID)
require.Equal(t, "evidence.txt", got.Name)
require.Equal(t, "text/plain", got.MimeType)
require.Equal(t, int64(len(data)), got.SizeBytes)
require.Equal(t, fmt.Sprintf("%x", sha256.Sum256(data)), got.SHA256)
require.NotEmpty(t, got.URL)
require.False(t, got.ExpiresAt.IsZero())
req, err := http.NewRequestWithContext(ctx, http.MethodGet, got.URL, nil)
require.NoError(t, err)
res, err := http.DefaultClient.Do(req)
require.NoError(t, err)
defer res.Body.Close()
require.Equal(t, http.StatusOK, res.StatusCode)
body, err := io.ReadAll(res.Body)
require.NoError(t, err)
require.Equal(t, data, body)
}
byID, err := testTool(t, toolsdk.DownloadChatFile, tb, toolsdk.DownloadChatFileArgs{FileID: uploaded.ID.String()})
require.NoError(t, err)
assertDownload(t, byID)
byName, err := testTool(t, toolsdk.DownloadChatFile, tb, toolsdk.DownloadChatFileArgs{
ChatID: chat.ID.String(),
FileName: "evidence.txt",
})
require.NoError(t, err)
assertDownload(t, byName)
status, err := testTool(t, toolsdk.GetChat, tb, toolsdk.GetChatArgs{ChatID: chat.ID.String()})
require.NoError(t, err)
require.Len(t, status.Files, 1)
require.Equal(t, int64(len(data)), status.Files[0].SizeBytes)
require.False(t, status.Files[0].CreatedAt.IsZero())
messages, err := testTool(t, toolsdk.GetChatMessages, tb, toolsdk.GetChatMessagesArgs{ChatID: chat.ID.String()})
require.NoError(t, err)
var attachingMessage *toolsdk.ChatToolMessage
for i := range messages.Messages {
if messages.Messages[i].Text == "Review the evidence." {
attachingMessage = &messages.Messages[i]
break
}
}
require.NotNil(t, attachingMessage)
require.Equal(t, []toolsdk.ChatToolMessageFile{{
ID: uploaded.ID.String(),
Name: "evidence.txt",
MimeType: "text/plain",
}}, attachingMessage.Files)
coderdtest.WaitForChatSettled(ctx, t, api, chat.ID)
fileOnly, err := expClient.UploadChatFile(ctx, firstUser.OrganizationID, "text/plain", "file-only.txt", bytes.NewReader([]byte("file only")))
require.NoError(t, err)
_, err = expClient.CreateChatMessage(ctx, chat.ID, codersdk.CreateChatMessageRequest{
Content: []codersdk.ChatInputPart{{Type: codersdk.ChatInputPartTypeFile, FileID: fileOnly.ID}},
})
require.NoError(t, err)
messages, err = testTool(t, toolsdk.GetChatMessages, tb, toolsdk.GetChatMessagesArgs{ChatID: chat.ID.String()})
require.NoError(t, err)
var fileOnlyMessage *toolsdk.ChatToolMessage
for i := range messages.Messages {
if len(messages.Messages[i].Files) == 1 && messages.Messages[i].Files[0].ID == fileOnly.ID.String() {
fileOnlyMessage = &messages.Messages[i]
break
}
}
require.NotNil(t, fileOnlyMessage)
require.Empty(t, fileOnlyMessage.Text)
require.Equal(t, []toolsdk.ChatToolMessageFile{{
ID: fileOnly.ID.String(),
Name: "file-only.txt",
MimeType: "text/plain",
}}, fileOnlyMessage.Files)
duplicateA, err := expClient.UploadChatFile(ctx, firstUser.OrganizationID, "text/plain", "duplicate.txt", bytes.NewReader([]byte("a")))
require.NoError(t, err)
duplicateB, err := expClient.UploadChatFile(ctx, firstUser.OrganizationID, "text/plain", "duplicate.txt", bytes.NewReader([]byte("bb")))
require.NoError(t, err)
ambiguousChat, err := expClient.CreateChat(ctx, codersdk.CreateChatRequest{
OrganizationID: firstUser.OrganizationID,
Content: []codersdk.ChatInputPart{
{Type: codersdk.ChatInputPartTypeText, Text: "Compare these files."},
{Type: codersdk.ChatInputPartTypeFile, FileID: duplicateA.ID},
{Type: codersdk.ChatInputPartTypeFile, FileID: duplicateB.ID},
},
})
require.NoError(t, err)
_, err = testTool(t, toolsdk.DownloadChatFile, tb, toolsdk.DownloadChatFileArgs{
ChatID: ambiguousChat.ID.String(),
FileName: "duplicate.txt",
})
require.ErrorContains(t, err, "multiple chat files")
require.ErrorContains(t, err, duplicateA.ID.String())
require.ErrorContains(t, err, duplicateB.ID.String())
require.ErrorContains(t, err, "mime_type")
require.ErrorContains(t, err, "size_bytes")
_, err = testTool(t, toolsdk.DownloadChatFile, tb, toolsdk.DownloadChatFileArgs{
ChatID: ambiguousChat.ID.String(),
FileName: "missing.txt",
})
require.ErrorContains(t, err, "no chat file")
require.ErrorContains(t, err, "duplicate.txt")
coderdtest.WaitForChatSettled(ctx, t, api, chat.ID)
coderdtest.WaitForChatSettled(ctx, t, api, ambiguousChat.ID)
})
t.Run("ForwardMessagePagination", func(t *testing.T) {
ctx := testutil.Context(t, testutil.WaitLong)
chat, err := expClient.CreateChat(ctx, codersdk.CreateChatRequest{
OrganizationID: firstUser.OrganizationID,
Content: []codersdk.ChatInputPart{{Type: codersdk.ChatInputPartTypeText, Text: "Create a pagination baseline."}},
})
require.NoError(t, err)
coderdtest.WaitForChatSettled(ctx, t, api, chat.ID)
existing, err := expClient.GetChatMessages(ctx, chat.ID, nil)
require.NoError(t, err)
var baselineID int64
for _, msg := range existing.Messages {
baselineID = max(baselineID, msg.ID)
}
require.Positive(t, baselineID)
textContent := func(text string) database.ChatMessage {
content, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{{Type: codersdk.ChatMessagePartTypeText, Text: text}})
require.NoError(t, err)
return database.ChatMessage{
ChatID: chat.ID,
ModelConfigID: uuid.NullUUID{UUID: defaultModelConfig.ID, Valid: true},
Role: database.ChatMessageRoleUser,
Content: content,
}
}
first := dbgen.ChatMessage(t, api.Database, textContent("forward one"))
toolContent, err := chatprompt.MarshalParts([]codersdk.ChatMessagePart{{
Type: codersdk.ChatMessagePartTypeToolCall,
ToolCallID: "forward-call",
ToolName: "execute",
}})
require.NoError(t, err)
toolOnly := dbgen.ChatMessage(t, api.Database, database.ChatMessage{
ChatID: chat.ID,
ModelConfigID: uuid.NullUUID{UUID: defaultModelConfig.ID, Valid: true},
Role: database.ChatMessageRoleAssistant,
Content: toolContent,
})
last := dbgen.ChatMessage(t, api.Database, textContent("forward two"))
firstPage, err := testTool(t, toolsdk.GetChatMessages, tb, toolsdk.GetChatMessagesArgs{
ChatID: chat.ID.String(),
AfterID: baselineID,
Limit: 2,
})
require.NoError(t, err)
require.True(t, firstPage.HasMore)
require.Equal(t, toolOnly.ID, firstPage.NextAfterID)
require.Zero(t, firstPage.NextBeforeID)
require.Len(t, firstPage.Messages, 1)
require.Equal(t, first.ID, firstPage.Messages[0].ID)
secondPage, err := testTool(t, toolsdk.GetChatMessages, tb, toolsdk.GetChatMessagesArgs{
ChatID: chat.ID.String(),
AfterID: firstPage.NextAfterID,
Limit: 2,
})
require.NoError(t, err)
require.False(t, secondPage.HasMore)
require.Equal(t, last.ID, secondPage.NextAfterID)
require.Len(t, secondPage.Messages, 1)
require.Equal(t, last.ID, secondPage.Messages[0].ID)
require.Greater(t, secondPage.Messages[0].ID, firstPage.Messages[0].ID)
finalPage, err := testTool(t, toolsdk.GetChatMessages, tb, toolsdk.GetChatMessagesArgs{
ChatID: chat.ID.String(),
AfterID: secondPage.NextAfterID,
Limit: 2,
})
require.NoError(t, err)
require.False(t, finalPage.HasMore)
require.Zero(t, finalPage.NextAfterID)
require.Empty(t, finalPage.Messages)
})
t.Run("ListChatsByLabel", func(t *testing.T) {
ctx := testutil.Context(t, testutil.WaitLong)
labelValue := uuid.NewString()
matching, err := expClient.CreateChat(ctx, codersdk.CreateChatRequest{
OrganizationID: firstUser.OrganizationID,
Content: []codersdk.ChatInputPart{{Type: codersdk.ChatInputPartTypeText, Text: "Matching chat."}},
Labels: map[string]string{"uat-evidence": labelValue},
})
require.NoError(t, err)
nonmatching, err := expClient.CreateChat(ctx, codersdk.CreateChatRequest{
OrganizationID: firstUser.OrganizationID,
Content: []codersdk.ChatInputPart{{Type: codersdk.ChatInputPartTypeText, Text: "Nonmatching chat."}},
Labels: map[string]string{"uat-evidence": uuid.NewString()},
})
require.NoError(t, err)
result, err := testTool(t, toolsdk.ListChats, tb, toolsdk.ListChatsArgs{
Labels: map[string]string{"uat-evidence": labelValue},
Limit: 100,
})
require.NoError(t, err)
require.Len(t, result.Chats, 1)
require.Equal(t, matching.ID.String(), result.Chats[0].ID)
require.Equal(t, map[string]string{"uat-evidence": labelValue}, result.Chats[0].Labels)
coderdtest.WaitForChatSettled(ctx, t, api, matching.ID)
coderdtest.WaitForChatSettled(ctx, t, api, nonmatching.ID)
})
t.Run("AwaitChat", func(t *testing.T) {
ctx := testutil.Context(t, testutil.WaitLong)
settled, err := expClient.CreateChat(ctx, codersdk.CreateChatRequest{
OrganizationID: firstUser.OrganizationID,
Content: []codersdk.ChatInputPart{{Type: codersdk.ChatInputPartTypeText, Text: "Settle immediately."}},
})
require.NoError(t, err)
coderdtest.WaitForChatSettled(ctx, t, api, settled.ID)
immediate, err := testTool(t, toolsdk.AwaitChat, tb, toolsdk.AwaitChatArgs{ChatID: settled.ID.String()})
require.NoError(t, err)
require.False(t, immediate.TimedOut)
require.Equal(t, codersdk.ChatStatusWaiting, immediate.Chat.Status)
streamStarted := make(chan struct{})
providerRelease := make(chan struct{})
var providerStartedOnce sync.Once
var providerReleaseOnce sync.Once
releaseProvider := func() { providerReleaseOnce.Do(func() { close(providerRelease) }) }
t.Cleanup(releaseProvider)
blockingURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
if req.Stream {
providerStartedOnce.Do(func() { close(streamStarted) })
select {
case <-providerRelease:
case <-req.Context().Done():
}
return chattest.OpenAIStreamingResponse(chattest.OpenAITextChunks("Released.")...)
}
return chattest.OpenAINonStreamingResponse(`{"title": "Await Test"}`)
})
blockingModel := coderdtest.CreateOpenAICompatChatModelConfig(t, expClient, blockingURL)
awaitFile, err := expClient.UploadChatFile(ctx, firstUser.OrganizationID, "text/plain", "await.txt", bytes.NewReader([]byte("await evidence")))
require.NoError(t, err)
running, err := expClient.CreateChat(ctx, codersdk.CreateChatRequest{
OrganizationID: firstUser.OrganizationID,
Content: []codersdk.ChatInputPart{
{Type: codersdk.ChatInputPartTypeText, Text: "Wait for release."},
{Type: codersdk.ChatInputPartTypeFile, FileID: awaitFile.ID},
},
ModelConfigID: &blockingModel.ID,
})
require.NoError(t, err)
select {
case <-streamStarted:
case <-ctx.Done():
t.Fatal(ctx.Err())
}
getSeen := make(chan struct{})
getRelease := make(chan struct{})
transport := &signalPathTransport{
path: "/api/experimental/chats/" + running.ID.String(),
seen: getSeen,
release: getRelease,
}
awaitClient := codersdk.New(client.URL)
awaitClient.SetSessionToken(client.SessionToken())
awaitClient.HTTPClient = &http.Client{Transport: transport}
t.Cleanup(awaitClient.HTTPClient.CloseIdleConnections)
awaitDeps, err := toolsdk.NewDeps(awaitClient)
require.NoError(t, err)
type awaitResult struct {
response toolsdk.AwaitChatResponse
err error
}
result := make(chan awaitResult, 1)
go func() {
response, err := toolsdk.AwaitChat.Handler(ctx, awaitDeps, toolsdk.AwaitChatArgs{
ChatID: running.ID.String(),
WaitSecs: 10,
})
result <- awaitResult{response: response, err: err}
}()
select {
case <-getSeen:
case <-ctx.Done():
t.Fatal(ctx.Err())
}
close(getRelease)
releaseProvider()
select {
case awaited := <-result:
require.NoError(t, awaited.err)
require.False(t, awaited.response.TimedOut)
require.Equal(t, codersdk.ChatStatusWaiting, awaited.response.Chat.Status)
require.Len(t, awaited.response.Chat.Files, 1)
require.Equal(t, awaitFile.ID.String(), awaited.response.Chat.Files[0].ID)
case <-ctx.Done():
t.Fatal(ctx.Err())
}
coderdtest.WaitForChatSettled(ctx, t, api, running.ID)
sharedClient, sharedUser := coderdtest.CreateAnotherUser(t, client, firstUser.OrganizationID)
sharedStreamStarted := make(chan struct{})
sharedProviderRelease := make(chan struct{})
var sharedProviderStartedOnce sync.Once
var sharedProviderReleaseOnce sync.Once
releaseSharedProvider := func() { sharedProviderReleaseOnce.Do(func() { close(sharedProviderRelease) }) }
t.Cleanup(releaseSharedProvider)
sharedBlockingURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
if req.Stream {
sharedProviderStartedOnce.Do(func() { close(sharedStreamStarted) })
select {
case <-sharedProviderRelease:
case <-req.Context().Done():
}
return chattest.OpenAIStreamingResponse(chattest.OpenAITextChunks("Released.")...)
}
return chattest.OpenAINonStreamingResponse(`{"title": "Shared Await Test"}`)
})
sharedBlockingModel := coderdtest.CreateOpenAICompatChatModelConfig(t, expClient, sharedBlockingURL)
sharedRunning, err := expClient.CreateChat(ctx, codersdk.CreateChatRequest{
OrganizationID: firstUser.OrganizationID,
Content: []codersdk.ChatInputPart{{Type: codersdk.ChatInputPartTypeText, Text: "Wait for shared release."}},
ModelConfigID: &sharedBlockingModel.ID,
})
require.NoError(t, err)
select {
case <-sharedStreamStarted:
case <-ctx.Done():
t.Fatal(ctx.Err())
}
err = expClient.UpdateChatACL(ctx, sharedRunning.ID, codersdk.UpdateChatACL{
UserRoles: map[string]codersdk.ChatRole{sharedUser.ID.String(): codersdk.ChatRoleRead},
})
require.NoError(t, err)
acl, err := expClient.GetChatACL(ctx, sharedRunning.ID)
require.NoError(t, err)
require.Len(t, acl.Users, 1)
require.Equal(t, sharedUser.ID, acl.Users[0].ID)
require.Equal(t, codersdk.ChatRoleRead, acl.Users[0].Role)
sharedGetSeen := make(chan struct{})
sharedGetRelease := make(chan struct{})
sharedAwaitClient := codersdk.New(sharedClient.URL)
sharedAwaitClient.SetSessionToken(sharedClient.SessionToken())
sharedAwaitClient.HTTPClient = &http.Client{Transport: &signalPathTransport{
path: "/api/experimental/chats/" + sharedRunning.ID.String(),
seen: sharedGetSeen,
release: sharedGetRelease,
}}
t.Cleanup(sharedAwaitClient.HTTPClient.CloseIdleConnections)
sharedAwaitDeps, err := toolsdk.NewDeps(sharedAwaitClient)
require.NoError(t, err)
sharedAwaitCtx := testutil.Context(t, testutil.WaitMedium)
sharedResult := make(chan awaitResult, 1)
go func() {
response, err := toolsdk.AwaitChat.Handler(sharedAwaitCtx, sharedAwaitDeps, toolsdk.AwaitChatArgs{
ChatID: sharedRunning.ID.String(),
WaitSecs: 20,
})
sharedResult <- awaitResult{response: response, err: err}
}()
select {
case <-sharedGetSeen:
case <-ctx.Done():
t.Fatal(ctx.Err())
}
close(sharedGetRelease)
releaseSharedProvider()
select {
case awaited := <-sharedResult:
require.NoError(t, awaited.err)
require.False(t, awaited.response.TimedOut)
require.Equal(t, codersdk.ChatStatusWaiting, awaited.response.Chat.Status)
case <-sharedAwaitCtx.Done():
t.Fatal(sharedAwaitCtx.Err())
}
coderdtest.WaitForChatSettled(ctx, t, api, sharedRunning.ID)
timeoutStarted := make(chan struct{})
timeoutRelease := make(chan struct{})
var timeoutStartedOnce sync.Once
var timeoutReleaseOnce sync.Once
releaseTimeout := func() { timeoutReleaseOnce.Do(func() { close(timeoutRelease) }) }
t.Cleanup(releaseTimeout)
timeoutURL := chattest.NewOpenAI(t, func(req *chattest.OpenAIRequest) chattest.OpenAIResponse {
if req.Stream {
timeoutStartedOnce.Do(func() { close(timeoutStarted) })
select {
case <-timeoutRelease:
case <-req.Context().Done():
}
return chattest.OpenAIStreamingResponse(chattest.OpenAITextChunks("Released.")...)
}
return chattest.OpenAINonStreamingResponse(`{"title": "Await Timeout Test"}`)
})
timeoutModel := coderdtest.CreateOpenAICompatChatModelConfig(t, expClient, timeoutURL)
busy, err := expClient.CreateChat(ctx, codersdk.CreateChatRequest{
OrganizationID: firstUser.OrganizationID,
Content: []codersdk.ChatInputPart{{Type: codersdk.ChatInputPartTypeText, Text: "Stay busy."}},
ModelConfigID: &timeoutModel.ID,
})
require.NoError(t, err)
select {
case <-timeoutStarted:
case <-ctx.Done():
t.Fatal(ctx.Err())
}
timedOut, err := testTool(t, toolsdk.AwaitChat, tb, toolsdk.AwaitChatArgs{
ChatID: busy.ID.String(),
WaitSecs: 1,
})
require.NoError(t, err)
require.True(t, timedOut.TimedOut)
require.Equal(t, codersdk.ChatStatusRunning, timedOut.Chat.Status)
releaseTimeout()
coderdtest.WaitForChatSettled(ctx, t, api, busy.ID)
// An initial status request that outlives the wait window
// errors within the wait_secs bound; the old post-deadline
// fallback fetch added up to 15 extra seconds.
stallClient := codersdk.New(client.URL)
stallClient.SetSessionToken(client.SessionToken())
stallClient.HTTPClient = &http.Client{Transport: stallTransport{}}
t.Cleanup(stallClient.HTTPClient.CloseIdleConnections)
stallDeps, err := toolsdk.NewDeps(stallClient)
require.NoError(t, err)
stallStart := time.Now()
_, err = toolsdk.AwaitChat.Handler(ctx, stallDeps, toolsdk.AwaitChatArgs{
ChatID: busy.ID.String(),
WaitSecs: 1,
})
require.ErrorIs(t, err, context.DeadlineExceeded)
require.Less(t, time.Since(stallStart), 10*time.Second)
})
t.Run("ListChatModelConfigsSkipsDisabledProviders", func(t *testing.T) {
ctx := testutil.Context(t, testutil.WaitLong)
@@ -231,9 +728,19 @@ func TestChatTools(t *testing.T) {
})
require.NoError(t, err)
require.True(t, sent.Queued)
// A queued prompt carrying only a file must surface in
// queued_messages rather than being dropped.
queuedFile, err := expClient.UploadChatFile(ctx, firstUser.OrganizationID, "text/plain", "queued-only.txt", bytes.NewReader([]byte("queued file")))
require.NoError(t, err)
_, err = expClient.CreateChatMessage(ctx, uuid.MustParse(created.ID), codersdk.CreateChatMessageRequest{
Content: []codersdk.ChatInputPart{{Type: codersdk.ChatInputPartTypeFile, FileID: queuedFile.ID}},
BusyBehavior: codersdk.ChatBusyBehaviorQueue,
})
require.NoError(t, err)
transcript, err := testTool(t, toolsdk.GetChatMessages, tb, toolsdk.GetChatMessagesArgs{ChatID: created.ID})
require.NoError(t, err)
require.Contains(t, transcript.QueuedMessages, "Queued while busy.")
require.Contains(t, transcript.QueuedMessages, "(attached files: queued-only.txt)")
interrupted, err := testTool(t, toolsdk.InterruptChat, tb, toolsdk.InterruptChatArgs{ChatID: created.ID})
require.NoError(t, err)
@@ -299,6 +806,23 @@ func TestChatTools(t *testing.T) {
})
require.ErrorContains(t, err, "busy_behavior")
_, err = testTool(t, toolsdk.DownloadChatFile, tb, toolsdk.DownloadChatFileArgs{})
require.ErrorContains(t, err, "exactly one addressing mode")
_, err = testTool(t, toolsdk.AwaitChat, tb, toolsdk.AwaitChatArgs{ChatID: "not-a-uuid"})
require.ErrorContains(t, err, "chat_id must be a valid UUID")
listed, err := testTool(t, toolsdk.ListChats, tb, toolsdk.ListChatsArgs{Limit: 1})
require.NoError(t, err)
require.LessOrEqual(t, len(listed.Chats), 1)
_, err = testTool(t, toolsdk.GetChatMessages, tb, toolsdk.GetChatMessagesArgs{
ChatID: uuid.NewString(),
BeforeID: 1,
AfterID: 2,
})
require.ErrorContains(t, err, "before_id and after_id cannot be used together")
for _, limit := range []int{-1, 201} {
_, err = testTool(t, toolsdk.GetChatMessages, tb, toolsdk.GetChatMessagesArgs{
ChatID: uuid.NewString(),
+23 -18
View File
@@ -60,6 +60,9 @@ const (
ToolNameGetTaskLogs = "coder_get_task_logs"
ToolNameCreateChat = "coder_create_chat"
ToolNameGetChat = "coder_get_chat"
ToolNameDownloadChatFile = "coder_download_chat_file"
ToolNameAwaitChat = "coder_await_chat"
ToolNameListChats = "coder_list_chats"
ToolNameGetChatMessages = "coder_get_chat_messages"
ToolNameSendChatMessage = "coder_send_chat_message"
ToolNameInterruptChat = "coder_interrupt_chat"
@@ -347,6 +350,9 @@ var All = []GenericTool{
GetTaskLogs.Generic(),
CreateChat.Generic(),
GetChat.Generic(),
DownloadChatFile.Generic(),
AwaitChat.Generic(),
ListChats.Generic(),
GetChatMessages.Generic(),
SendChatMessage.Generic(),
InterruptChat.Generic(),
@@ -632,10 +638,22 @@ var ListWorkspaces = Tool[ListWorkspacesArgs, []MinimalWorkspace]{
},
}
func minimalTemplate(template codersdk.Template) MinimalTemplate {
return MinimalTemplate{
DisplayName: template.DisplayName,
ID: template.ID.String(),
Name: template.Name,
Description: template.Description,
ActiveVersionID: template.ActiveVersionID,
ActiveUserCount: template.ActiveUserCount,
AgentsAllowed: template.AgentsAllowed,
}
}
var ListTemplates = Tool[NoArgs, []MinimalTemplate]{
Tool: aisdk.Tool{
Name: ToolNameListTemplates,
Description: "Lists templates for the authenticated user.",
Description: "Lists templates for the authenticated user. agents_allowed indicates whether Coder Agents (chats) may create workspaces from the template.",
Schema: aisdk.Schema{
Properties: map[string]any{},
Required: []string{},
@@ -649,14 +667,7 @@ var ListTemplates = Tool[NoArgs, []MinimalTemplate]{
}
minimalTemplates := make([]MinimalTemplate, len(templates))
for i, template := range templates {
minimalTemplates[i] = MinimalTemplate{
DisplayName: template.DisplayName,
ID: template.ID.String(),
Name: template.Name,
Description: template.Description,
ActiveVersionID: template.ActiveVersionID,
ActiveUserCount: template.ActiveUserCount,
}
minimalTemplates[i] = minimalTemplate(template)
}
return minimalTemplates, nil
},
@@ -786,15 +797,8 @@ When selecting a preset: if a preset is marked default and the user has not spec
return TemplateDetail{}, xerrors.Errorf("get template presets: %w", err)
}
detail := TemplateDetail{
MinimalTemplate: MinimalTemplate{
DisplayName: template.DisplayName,
ID: template.ID.String(),
Name: template.Name,
Description: template.Description,
ActiveVersionID: template.ActiveVersionID,
ActiveUserCount: template.ActiveUserCount,
},
Parameters: parameters,
MinimalTemplate: minimalTemplate(template),
Parameters: parameters,
}
for _, p := range presets {
detail.Presets = append(detail.Presets, toPresetView(p))
@@ -1735,6 +1739,7 @@ type MinimalTemplate struct {
Description string `json:"description"`
ActiveVersionID uuid.UUID `json:"active_version_id"`
ActiveUserCount int `json:"active_user_count"`
AgentsAllowed bool `json:"agents_allowed"`
}
type WorkspaceLSArgs struct {
+25
View File
@@ -108,6 +108,30 @@ func TestGenericToolMCPAnnotations(t *testing.T) {
idempotentHint: true,
openWorldHint: false,
},
{
name: "DownloadChatFileIsReadOnly",
toolName: toolsdk.ToolNameDownloadChatFile,
readOnlyHint: true,
destructiveHint: false,
idempotentHint: true,
openWorldHint: false,
},
{
name: "AwaitChatIsReadOnly",
toolName: toolsdk.ToolNameAwaitChat,
readOnlyHint: true,
destructiveHint: false,
idempotentHint: true,
openWorldHint: false,
},
{
name: "ListChatsIsReadOnly",
toolName: toolsdk.ToolNameListChats,
readOnlyHint: true,
destructiveHint: false,
idempotentHint: true,
openWorldHint: false,
},
{
name: "DestructiveTool",
toolName: toolsdk.ToolNameWorkspaceWriteFile,
@@ -309,6 +333,7 @@ func TestTools(t *testing.T) {
})
for i, template := range result {
require.Equal(t, expected[i].ID.String(), template.ID)
require.Equal(t, expected[i].AgentsAllowed, template.AgentsAllowed)
}
})