mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
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:
@@ -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)
|
||||
|
||||
@@ -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
@@ -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
|
||||
},
|
||||
|
||||
@@ -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
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
|
||||
|
||||
Reference in New Issue
Block a user