diff --git a/coderd/apidoc/docs.go b/coderd/apidoc/docs.go index 27ed0964ed..bc7ecbda7c 100644 --- a/coderd/apidoc/docs.go +++ b/coderd/apidoc/docs.go @@ -374,6 +374,91 @@ const docTemplate = `{ ] } }, + "/api/experimental/chats/files/{file}/download": { + "get": { + "description": "Experimental: this endpoint is subject to change.", + "produces": [ + "image/png", + "image/jpeg", + "image/gif", + "image/webp", + "text/plain", + "text/markdown", + "text/csv", + "application/json", + "application/pdf" + ], + "tags": [ + "Chats" + ], + "summary": "Download chat file with signed token", + "operationId": "download-chat-file", + "parameters": [ + { + "type": "string", + "format": "uuid", + "description": "File ID", + "name": "file", + "in": "path", + "required": true + }, + { + "type": "string", + "description": "Signed download token", + "name": "token", + "in": "query", + "required": true + } + ], + "responses": { + "200": { + "description": "OK" + } + }, + "x-apidocgen": { + "skip": true + } + } + }, + "/api/experimental/chats/files/{file}/download-url": { + "post": { + "description": "Experimental: this endpoint is subject to change.", + "produces": [ + "application/json" + ], + "tags": [ + "Chats" + ], + "summary": "Create chat file download URL", + "operationId": "create-chat-file-download-url", + "parameters": [ + { + "type": "string", + "format": "uuid", + "description": "File ID", + "name": "file", + "in": "path", + "required": true + } + ], + "responses": { + "200": { + "description": "OK", + "schema": { + "$ref": "#/definitions/codersdk.ChatFileDownloadURLResponse" + } + } + }, + "security": [ + { + "CoderSessionToken": [] + } + ], + "x-apidocgen": { + "skip": true + } + } + }, "/api/experimental/chats/models": { "get": { "description": "Experimental: this endpoint is subject to change.", @@ -17906,6 +17991,31 @@ const docTemplate = `{ "ChatErrorKindHookDenied" ] }, + "codersdk.ChatFileDownloadURLResponse": { + "type": "object", + "properties": { + "expires_at": { + "type": "string", + "format": "date-time" + }, + "mime_type": { + "type": "string" + }, + "name": { + "type": "string" + }, + "sha256": { + "type": "string" + }, + "size_bytes": { + "type": "integer" + }, + "url": { + "type": "string", + "format": "uri" + } + } + }, "codersdk.ChatFileMetadata": { "type": "object", "properties": { @@ -17930,6 +18040,9 @@ const docTemplate = `{ "owner_id": { "type": "string", "format": "uuid" + }, + "size_bytes": { + "type": "integer" } } }, @@ -19858,6 +19971,7 @@ const docTemplate = `{ "workspace_apps_api_key", "workspace_apps_token", "oidc_convert", + "chat_files_token", "tailnet_resume", "nats_ca" ], @@ -19865,6 +19979,7 @@ const docTemplate = `{ "CryptoKeyFeatureWorkspaceAppsAPIKey", "CryptoKeyFeatureWorkspaceAppsToken", "CryptoKeyFeatureOIDCConvert", + "CryptoKeyFeatureChatFilesToken", "CryptoKeyFeatureTailnetResume", "CryptoKeyFeatureNATSCA" ] diff --git a/coderd/apidoc/swagger.json b/coderd/apidoc/swagger.json index 53df566fc8..638c27ec24 100644 --- a/coderd/apidoc/swagger.json +++ b/coderd/apidoc/swagger.json @@ -327,6 +327,85 @@ ] } }, + "/api/experimental/chats/files/{file}/download": { + "get": { + "description": "Experimental: this endpoint is subject to change.", + "produces": [ + "image/png", + "image/jpeg", + "image/gif", + "image/webp", + "text/plain", + "text/markdown", + "text/csv", + "application/json", + "application/pdf" + ], + "tags": ["Chats"], + "summary": "Download chat file with signed token", + "operationId": "download-chat-file", + "parameters": [ + { + "type": "string", + "format": "uuid", + "description": "File ID", + "name": "file", + "in": "path", + "required": true + }, + { + "type": "string", + "description": "Signed download token", + "name": "token", + "in": "query", + "required": true + } + ], + "responses": { + "200": { + "description": "OK" + } + }, + "x-apidocgen": { + "skip": true + } + } + }, + "/api/experimental/chats/files/{file}/download-url": { + "post": { + "description": "Experimental: this endpoint is subject to change.", + "produces": ["application/json"], + "tags": ["Chats"], + "summary": "Create chat file download URL", + "operationId": "create-chat-file-download-url", + "parameters": [ + { + "type": "string", + "format": "uuid", + "description": "File ID", + "name": "file", + "in": "path", + "required": true + } + ], + "responses": { + "200": { + "description": "OK", + "schema": { + "$ref": "#/definitions/codersdk.ChatFileDownloadURLResponse" + } + } + }, + "security": [ + { + "CoderSessionToken": [] + } + ], + "x-apidocgen": { + "skip": true + } + } + }, "/api/experimental/chats/models": { "get": { "description": "Experimental: this endpoint is subject to change.", @@ -16120,6 +16199,31 @@ "ChatErrorKindHookDenied" ] }, + "codersdk.ChatFileDownloadURLResponse": { + "type": "object", + "properties": { + "expires_at": { + "type": "string", + "format": "date-time" + }, + "mime_type": { + "type": "string" + }, + "name": { + "type": "string" + }, + "sha256": { + "type": "string" + }, + "size_bytes": { + "type": "integer" + }, + "url": { + "type": "string", + "format": "uri" + } + } + }, "codersdk.ChatFileMetadata": { "type": "object", "properties": { @@ -16144,6 +16248,9 @@ "owner_id": { "type": "string", "format": "uuid" + }, + "size_bytes": { + "type": "integer" } } }, @@ -17987,6 +18094,7 @@ "workspace_apps_api_key", "workspace_apps_token", "oidc_convert", + "chat_files_token", "tailnet_resume", "nats_ca" ], @@ -17994,6 +18102,7 @@ "CryptoKeyFeatureWorkspaceAppsAPIKey", "CryptoKeyFeatureWorkspaceAppsToken", "CryptoKeyFeatureOIDCConvert", + "CryptoKeyFeatureChatFilesToken", "CryptoKeyFeatureTailnetResume", "CryptoKeyFeatureNATSCA" ] diff --git a/coderd/coderd.go b/coderd/coderd.go index c7effe22b5..3f9d3babbf 100644 --- a/coderd/coderd.go +++ b/coderd/coderd.go @@ -325,6 +325,7 @@ type Options struct { AppSigningKeyCache cryptokeys.SigningKeycache AppEncryptionKeyCache cryptokeys.EncryptionKeycache OIDCConvertKeyCache cryptokeys.SigningKeycache + ChatFileTokenKeyCache cryptokeys.SigningKeycache // NATSCACache serves the NATS cluster mTLS CA via the generic signing key // cache for the nats_ca feature. SigningKey returns the active CA // (a *NATSCA); VerifyingKey returns a specific CA by sequence. The key @@ -595,6 +596,17 @@ func New(options *Options) *API { } } + if options.ChatFileTokenKeyCache == nil { + options.ChatFileTokenKeyCache, err = cryptokeys.NewSigningCache(ctx, + options.Logger.Named("chat_file_token_keycache"), + fetcher, + codersdk.CryptoKeyFeatureChatFilesToken, + ) + if err != nil { + options.Logger.Fatal(ctx, "failed to properly instantiate chat file token signing cache", slog.Error(err)) + } + } + if options.AppSigningKeyCache == nil { options.AppSigningKeyCache, err = cryptokeys.NewSigningCache(ctx, options.Logger.Named("app_signing_keycache"), @@ -1362,6 +1374,10 @@ func New(options *Options) *API { r.Delete("/", api.deleteUserAIProviderKey) }) }) + r.Group(func(r chi.Router) { + r.Use(httpmw.RateLimit(options.FilesRateLimit, time.Minute)) + r.Get("/chats/files/{file}/download", api.downloadChatFile) + }) r.Route("/chats", func(r chi.Router) { r.Use( apiKeyMiddleware, @@ -1374,6 +1390,7 @@ func New(options *Options) *API { r.Route("/files", func(r chi.Router) { r.Use(httpmw.RateLimit(options.FilesRateLimit, time.Minute)) r.Post("/", api.postChatFile) + r.Post("/{file}/download-url", api.postChatFileDownloadURL) r.Get("/{file}", api.chatFileByID) }) r.Route("/config", func(r chi.Router) { @@ -2496,6 +2513,7 @@ func (api *API) Close() error { } _ = api.NetworkTelemetryBatcher.Close() _ = api.OIDCConvertKeyCache.Close() + _ = api.ChatFileTokenKeyCache.Close() _ = api.AppSigningKeyCache.Close() _ = api.AppEncryptionKeyCache.Close() if api.NATSCACache != nil { diff --git a/coderd/coderdtest/coderdtest.go b/coderd/coderdtest/coderdtest.go index 5377658a35..101e7efbed 100644 --- a/coderd/coderdtest/coderdtest.go +++ b/coderd/coderdtest/coderdtest.go @@ -204,6 +204,7 @@ type Options struct { NotificationsEnqueuer notifications.Enqueuer APIKeyEncryptionCache cryptokeys.EncryptionKeycache OIDCConvertKeyCache cryptokeys.SigningKeycache + ChatFileTokenKeyCache cryptokeys.SigningKeycache Clock quartz.Clock Acquirer *provisionerdserver.Acquirer TelemetryReporter telemetry.Reporter @@ -693,6 +694,7 @@ func NewOptions(t testing.TB, options *Options) (func(http.Handler), context.Can Acquirer: options.Acquirer, AppEncryptionKeyCache: options.APIKeyEncryptionCache, OIDCConvertKeyCache: options.OIDCConvertKeyCache, + ChatFileTokenKeyCache: options.ChatFileTokenKeyCache, ProvisionerdServerMetrics: options.ProvisionerdServerMetrics, WorkspaceBuilderMetrics: options.WorkspaceBuilderMetrics, } diff --git a/coderd/cryptokeys/cache.go b/coderd/cryptokeys/cache.go index 75269d9839..503b715ed4 100644 --- a/coderd/cryptokeys/cache.go +++ b/coderd/cryptokeys/cache.go @@ -233,7 +233,7 @@ func isEncryptionKeyFeature(feature codersdk.CryptoKeyFeature) bool { func isSigningKeyFeature(feature codersdk.CryptoKeyFeature) bool { switch feature { - case codersdk.CryptoKeyFeatureTailnetResume, codersdk.CryptoKeyFeatureOIDCConvert, codersdk.CryptoKeyFeatureWorkspaceAppsToken, codersdk.CryptoKeyFeatureNATSCA: + case codersdk.CryptoKeyFeatureTailnetResume, codersdk.CryptoKeyFeatureOIDCConvert, codersdk.CryptoKeyFeatureChatFilesToken, codersdk.CryptoKeyFeatureWorkspaceAppsToken, codersdk.CryptoKeyFeatureNATSCA: return true default: return false diff --git a/coderd/cryptokeys/rotate.go b/coderd/cryptokeys/rotate.go index 775131185b..486f2e6c6b 100644 --- a/coderd/cryptokeys/rotate.go +++ b/coderd/cryptokeys/rotate.go @@ -20,6 +20,9 @@ import ( const ( WorkspaceAppsTokenDuration = time.Minute OIDCConvertTokenDuration = time.Minute * 5 + // ChatFilesTokenDuration is also the lifetime of minted chat file + // download URLs, keeping key retention aligned with token expiry. + ChatFilesTokenDuration = time.Minute * 5 TailnetResumeTokenDuration = time.Hour * 24 // NATSCAOverlap is how long a NATS cluster CA certificate stays valid past // the end of its active-signing window (startsAt + keyDuration). The next CA @@ -46,6 +49,7 @@ var defaultRotatedFeatures = []database.CryptoKeyFeature{ database.CryptoKeyFeatureWorkspaceAppsToken, database.CryptoKeyFeatureWorkspaceAppsAPIKey, database.CryptoKeyFeatureOIDCConvert, + database.CryptoKeyFeatureChatFilesToken, database.CryptoKeyFeatureTailnetResume, } @@ -273,6 +277,8 @@ func generateNewSecret(feature database.CryptoKeyFeature, startsAt time.Time, ke return generateKey(64) case database.CryptoKeyFeatureOIDCConvert: return generateKey(64) + case database.CryptoKeyFeatureChatFilesToken: + return generateKey(64) case database.CryptoKeyFeatureTailnetResume: return generateKey(64) case database.CryptoKeyFeatureNATSCA: @@ -298,6 +304,8 @@ func tokenDuration(feature database.CryptoKeyFeature) time.Duration { return WorkspaceAppsTokenDuration case database.CryptoKeyFeatureOIDCConvert: return OIDCConvertTokenDuration + case database.CryptoKeyFeatureChatFilesToken: + return ChatFilesTokenDuration case database.CryptoKeyFeatureTailnetResume: return TailnetResumeTokenDuration case database.CryptoKeyFeatureNATSCA: diff --git a/coderd/cryptokeys/rotate_internal_test.go b/coderd/cryptokeys/rotate_internal_test.go index 89216cf708..a96281990c 100644 --- a/coderd/cryptokeys/rotate_internal_test.go +++ b/coderd/cryptokeys/rotate_internal_test.go @@ -513,7 +513,7 @@ func Test_rotateKeys(t *testing.T) { keys, err := db.GetCryptoKeys(ctx) require.NoError(t, err) - require.Len(t, keys, 5) + require.Len(t, keys, 6) kbf := keysByFeature(keys, defaultRotatedFeatures) @@ -525,13 +525,16 @@ func Test_rotateKeys(t *testing.T) { // caused a key to be inserted. require.Len(t, kbf[database.CryptoKeyFeatureTailnetResume], 1) require.Len(t, kbf[database.CryptoKeyFeatureWorkspaceAppsToken], 1) + require.Len(t, kbf[database.CryptoKeyFeatureChatFilesToken], 1) oidcKey := kbf[database.CryptoKeyFeatureOIDCConvert][0] tailnetKey := kbf[database.CryptoKeyFeatureTailnetResume][0] appTokenKey := kbf[database.CryptoKeyFeatureWorkspaceAppsToken][0] + chatFileTokenKey := kbf[database.CryptoKeyFeatureChatFilesToken][0] requireKey(t, oidcKey, database.CryptoKeyFeatureOIDCConvert, now, nullTime, validKey.Sequence) requireKey(t, tailnetKey, database.CryptoKeyFeatureTailnetResume, now, nullTime, deletedKey.Sequence+1) requireKey(t, appTokenKey, database.CryptoKeyFeatureWorkspaceAppsToken, now, nullTime, 1) + requireKey(t, chatFileTokenKey, database.CryptoKeyFeatureChatFilesToken, now, nullTime, 1) newKey := kbf[database.CryptoKeyFeatureWorkspaceAppsAPIKey][0] oldKey := kbf[database.CryptoKeyFeatureWorkspaceAppsAPIKey][1] if newKey.Sequence == rotatedKey.Sequence { @@ -705,6 +708,8 @@ func requireKey(t *testing.T, key database.CryptoKey, feature database.CryptoKey switch key.Feature { case database.CryptoKeyFeatureOIDCConvert: require.Len(t, secret, 64) + case database.CryptoKeyFeatureChatFilesToken: + require.Len(t, secret, 64) case database.CryptoKeyFeatureWorkspaceAppsToken: require.Len(t, secret, 64) case database.CryptoKeyFeatureWorkspaceAppsAPIKey: diff --git a/coderd/database/db2sdk/db2sdk.go b/coderd/database/db2sdk/db2sdk.go index 943c2bb955..2dbcc4248c 100644 --- a/coderd/database/db2sdk/db2sdk.go +++ b/coderd/database/db2sdk/db2sdk.go @@ -1903,6 +1903,7 @@ func Chat(c database.Chat, diffStatus *database.ChatDiffStatus, files []database OrganizationID: row.OrganizationID, Name: row.Name, MimeType: row.Mimetype, + SizeBytes: row.SizeBytes, CreatedAt: row.CreatedAt, }) } diff --git a/coderd/database/db2sdk/db2sdk_test.go b/coderd/database/db2sdk/db2sdk_test.go index 6a9ac0f0d6..20b2b55637 100644 --- a/coderd/database/db2sdk/db2sdk_test.go +++ b/coderd/database/db2sdk/db2sdk_test.go @@ -848,6 +848,7 @@ func TestChat_FileMetadataConversion(t *testing.T) { OrganizationID: orgID, Name: "screenshot.png", Mimetype: "image/png", + SizeBytes: 1234, CreatedAt: now, }, } @@ -861,6 +862,7 @@ func TestChat_FileMetadataConversion(t *testing.T) { require.Equal(t, orgID, f.OrganizationID, "OrganizationID must be mapped from DB row") require.Equal(t, "screenshot.png", f.Name) require.Equal(t, "image/png", f.MimeType) + require.Equal(t, int64(1234), f.SizeBytes) require.Equal(t, now, f.CreatedAt) // Verify JSON serialization uses snake_case for mime_type. diff --git a/coderd/database/dbauthz/dbauthz_test.go b/coderd/database/dbauthz/dbauthz_test.go index 86665079cd..38d8435e71 100644 --- a/coderd/database/dbauthz/dbauthz_test.go +++ b/coderd/database/dbauthz/dbauthz_test.go @@ -945,6 +945,7 @@ func (s *MethodTestSuite) TestChats() { ID: file.ID, Name: file.Name, Mimetype: file.Mimetype, + SizeBytes: int64(len(file.Data)), CreatedAt: file.CreatedAt, OwnerID: file.OwnerID, OrganizationID: file.OrganizationID, diff --git a/coderd/database/dbgen/dbgen.go b/coderd/database/dbgen/dbgen.go index ff20709991..df4f2cc2af 100644 --- a/coderd/database/dbgen/dbgen.go +++ b/coderd/database/dbgen/dbgen.go @@ -2229,6 +2229,8 @@ func newCryptoKeySecret(feature database.CryptoKeyFeature) (string, error) { return generateCryptoKey(64) case database.CryptoKeyFeatureTailnetResume: return generateCryptoKey(64) + case database.CryptoKeyFeatureChatFilesToken: + return generateCryptoKey(64) case database.CryptoKeyFeatureNATSCA: return generateCACryptoKeySecret() } diff --git a/coderd/database/dump.sql b/coderd/database/dump.sql index 7365b8821b..92f2f77b97 100644 --- a/coderd/database/dump.sql +++ b/coderd/database/dump.sql @@ -398,7 +398,8 @@ CREATE TYPE crypto_key_feature AS ENUM ( 'workspace_apps_api_key', 'oidc_convert', 'tailnet_resume', - 'nats_ca' + 'nats_ca', + 'chat_files_token' ); CREATE TYPE display_app AS ENUM ( diff --git a/coderd/database/migrations/000572_chat_files_token_crypto_key_feature.down.sql b/coderd/database/migrations/000572_chat_files_token_crypto_key_feature.down.sql new file mode 100644 index 0000000000..9df35cdaa9 --- /dev/null +++ b/coderd/database/migrations/000572_chat_files_token_crypto_key_feature.down.sql @@ -0,0 +1 @@ +-- PostgreSQL does not support removing enum values safely. diff --git a/coderd/database/migrations/000572_chat_files_token_crypto_key_feature.up.sql b/coderd/database/migrations/000572_chat_files_token_crypto_key_feature.up.sql new file mode 100644 index 0000000000..fccfaec783 --- /dev/null +++ b/coderd/database/migrations/000572_chat_files_token_crypto_key_feature.up.sql @@ -0,0 +1 @@ +ALTER TYPE crypto_key_feature ADD VALUE IF NOT EXISTS 'chat_files_token'; diff --git a/coderd/database/models.go b/coderd/database/models.go index c80a23665c..22aaada3fb 100644 --- a/coderd/database/models.go +++ b/coderd/database/models.go @@ -2043,6 +2043,7 @@ const ( CryptoKeyFeatureOIDCConvert CryptoKeyFeature = "oidc_convert" CryptoKeyFeatureTailnetResume CryptoKeyFeature = "tailnet_resume" CryptoKeyFeatureNATSCA CryptoKeyFeature = "nats_ca" + CryptoKeyFeatureChatFilesToken CryptoKeyFeature = "chat_files_token" ) func (e *CryptoKeyFeature) Scan(src interface{}) error { @@ -2086,7 +2087,8 @@ func (e CryptoKeyFeature) Valid() bool { CryptoKeyFeatureWorkspaceAppsAPIKey, CryptoKeyFeatureOIDCConvert, CryptoKeyFeatureTailnetResume, - CryptoKeyFeatureNATSCA: + CryptoKeyFeatureNATSCA, + CryptoKeyFeatureChatFilesToken: return true } return false @@ -2099,6 +2101,7 @@ func AllCryptoKeyFeatureValues() []CryptoKeyFeature { CryptoKeyFeatureOIDCConvert, CryptoKeyFeatureTailnetResume, CryptoKeyFeatureNATSCA, + CryptoKeyFeatureChatFilesToken, } } diff --git a/coderd/database/queries.sql.go b/coderd/database/queries.sql.go index 2649038095..bf7063bbb6 100644 --- a/coderd/database/queries.sql.go +++ b/coderd/database/queries.sql.go @@ -5930,7 +5930,8 @@ func (q *sqlQuerier) GetChatFileDataPrefixesByIDs(ctx context.Context, arg GetCh } const getChatFileMetadataByChatID = `-- name: GetChatFileMetadataByChatID :many -SELECT cf.id, cf.owner_id, cf.organization_id, cf.name, cf.mimetype, cf.created_at +SELECT cf.id, cf.owner_id, cf.organization_id, cf.name, cf.mimetype, cf.created_at, + octet_length(cf.data)::bigint AS size_bytes FROM chat_files cf JOIN chat_file_links cfl ON cfl.file_id = cf.id WHERE cfl.chat_id = $1::uuid @@ -5944,6 +5945,7 @@ type GetChatFileMetadataByChatIDRow struct { Name string `db:"name" json:"name"` Mimetype string `db:"mimetype" json:"mimetype"` CreatedAt time.Time `db:"created_at" json:"created_at"` + SizeBytes int64 `db:"size_bytes" json:"size_bytes"` } // GetChatFileMetadataByChatID returns lightweight file metadata for @@ -5965,6 +5967,7 @@ func (q *sqlQuerier) GetChatFileMetadataByChatID(ctx context.Context, chatID uui &i.Name, &i.Mimetype, &i.CreatedAt, + &i.SizeBytes, ); err != nil { return nil, err } diff --git a/coderd/database/queries/chatfiles.sql b/coderd/database/queries/chatfiles.sql index 6bf3f8813b..34a60bd61f 100644 --- a/coderd/database/queries/chatfiles.sql +++ b/coderd/database/queries/chatfiles.sql @@ -21,7 +21,8 @@ WHERE id = ANY(@ids::uuid[]); -- GetChatFileMetadataByChatID returns lightweight file metadata for -- all files linked to a chat. The data column is excluded to avoid -- loading file content. -SELECT cf.id, cf.owner_id, cf.organization_id, cf.name, cf.mimetype, cf.created_at +SELECT cf.id, cf.owner_id, cf.organization_id, cf.name, cf.mimetype, cf.created_at, + octet_length(cf.data)::bigint AS size_bytes FROM chat_files cf JOIN chat_file_links cfl ON cfl.file_id = cf.id WHERE cfl.chat_id = @chat_id::uuid diff --git a/coderd/exp_chats.go b/coderd/exp_chats.go index 5dcaad9e5c..f98a22d11f 100644 --- a/coderd/exp_chats.go +++ b/coderd/exp_chats.go @@ -3,7 +3,9 @@ package coderd import ( "cmp" "context" + "crypto/sha256" "database/sql" + "encoding/hex" "encoding/json" "errors" "fmt" @@ -11,6 +13,7 @@ import ( "mime" "net/http" "net/http/httptest" + "net/url" "slices" "strconv" "strings" @@ -19,6 +22,7 @@ import ( "unicode/utf8" "github.com/go-chi/chi/v5" + "github.com/go-jose/go-jose/v4/jwt" "github.com/google/uuid" "github.com/sqlc-dev/pqtype" "golang.org/x/xerrors" @@ -26,6 +30,7 @@ import ( "cdr.dev/slog/v3" "github.com/coder/coder/v2/agent/agentssh" "github.com/coder/coder/v2/coderd/audit" + "github.com/coder/coder/v2/coderd/cryptokeys" "github.com/coder/coder/v2/coderd/database" "github.com/coder/coder/v2/coderd/database/db2sdk" "github.com/coder/coder/v2/coderd/database/dbauthz" @@ -36,6 +41,7 @@ import ( "github.com/coder/coder/v2/coderd/httpapi" "github.com/coder/coder/v2/coderd/httpapi/httperror" "github.com/coder/coder/v2/coderd/httpmw" + "github.com/coder/coder/v2/coderd/jwtutils" "github.com/coder/coder/v2/coderd/pubsub" "github.com/coder/coder/v2/coderd/rbac" "github.com/coder/coder/v2/coderd/rbac/policy" @@ -5944,6 +5950,145 @@ func (api *API) postChatFile(rw http.ResponseWriter, r *http.Request) { }) } +// ChatFileDownloadClaims are the signed claims embedded in a chat file +// download URL token. +type ChatFileDownloadClaims struct { + jwtutils.RegisteredClaims + FileID uuid.UUID `json:"file_id"` + UserID uuid.UUID `json:"user_id"` +} + +func (c ChatFileDownloadClaims) Validate(expected jwt.Expected) error { + if c.FileID == uuid.Nil { + return xerrors.New("file ID is required") + } + if c.UserID == uuid.Nil { + return xerrors.New("user ID is required") + } + return c.RegisteredClaims.Validate(expected) +} + +// EXPERIMENTAL: this endpoint is experimental and is subject to change. +// +// @Summary Create chat file download URL +// @ID create-chat-file-download-url +// @Security CoderSessionToken +// @Tags Chats +// @Produce json +// @Param file path string true "File ID" format(uuid) +// @Success 200 {object} codersdk.ChatFileDownloadURLResponse +// @Router /api/experimental/chats/files/{file}/download-url [post] +// @x-apidocgen {"skip": true} +// @Description Experimental: this endpoint is subject to change. +func (api *API) postChatFileDownloadURL(rw http.ResponseWriter, r *http.Request) { + ctx := r.Context() + fileID, err := uuid.Parse(chi.URLParam(r, "file")) + if err != nil { + httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{ + Message: "Invalid file ID.", + }) + return + } + + chatFile, err := api.Database.GetChatFileByID(ctx, fileID) + if err != nil { + if httpapi.Is404Error(err) { + httpapi.ResourceNotFound(rw) + return + } + httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{ + Message: "Failed to get chat file.", + Detail: err.Error(), + }) + return + } + + now := time.Now() + // Truncate to whole seconds so the advertised expiry matches the JWT + // exp claim, which jwt.NewNumericDate stores at second precision. + expiresAt := now.Add(cryptokeys.ChatFilesTokenDuration).Truncate(time.Second) + claims := ChatFileDownloadClaims{ + RegisteredClaims: jwtutils.RegisteredClaims{ + Expiry: jwt.NewNumericDate(expiresAt), + IssuedAt: jwt.NewNumericDate(now), + }, + FileID: fileID, + UserID: httpmw.APIKey(r).UserID, + } + token, err := jwtutils.Sign(ctx, api.ChatFileTokenKeyCache, claims) + if err != nil { + httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{ + Message: "Failed to create chat file download URL.", + Detail: err.Error(), + }) + return + } + + downloadURL := api.AccessURL.JoinPath("api", "experimental", "chats", "files", fileID.String(), "download") + downloadURL.RawQuery = url.Values{"token": {token}}.Encode() + digest := sha256.Sum256(chatFile.Data) + httpapi.Write(ctx, rw, http.StatusOK, codersdk.ChatFileDownloadURLResponse{ + URL: downloadURL.String(), + ExpiresAt: expiresAt, + SHA256: hex.EncodeToString(digest[:]), + SizeBytes: int64(len(chatFile.Data)), + Name: chatFile.Name, + MimeType: chatFile.Mimetype, + }) +} + +// EXPERIMENTAL: this endpoint is experimental and is subject to change. +// +// @Summary Download chat file with signed token +// @ID download-chat-file +// @Tags Chats +// @Produce image/png,image/jpeg,image/gif,image/webp,text/plain,text/markdown,text/csv,application/json,application/pdf +// @Param file path string true "File ID" format(uuid) +// @Param token query string true "Signed download token" +// @Success 200 +// @Router /api/experimental/chats/files/{file}/download [get] +// @x-apidocgen {"skip": true} +// @Description Experimental: this endpoint is subject to change. +func (api *API) downloadChatFile(rw http.ResponseWriter, r *http.Request) { + ctx := r.Context() + fileID, err := uuid.Parse(chi.URLParam(r, "file")) + if err != nil { + httpapi.ResourceNotFound(rw) + return + } + + var claims ChatFileDownloadClaims + if err := jwtutils.Verify(ctx, api.ChatFileTokenKeyCache, r.URL.Query().Get("token"), &claims); err != nil || claims.FileID != fileID { + httpapi.ResourceNotFound(rw) + return + } + + subject, status, err := httpmw.UserRBACSubject(ctx, api.Database, claims.UserID, rbac.ScopeAll) + if err != nil || status != database.UserStatusActive { + httpapi.ResourceNotFound(rw) + return + } + + chatFile, err := api.Database.GetChatFileByID(dbauthz.As(ctx, subject), fileID) + if err != nil { + if httpapi.Is404Error(err) { + httpapi.ResourceNotFound(rw) + return + } + // This endpoint is reachable without a session token, so internal + // error details stay in the logs rather than the response body. + api.Logger.Error(ctx, "failed to get chat file for signed download", slog.Error(err)) + httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{ + Message: "Failed to get chat file.", + }) + return + } + // Never let private caches replay a signed URL past its expiry; + // revocation is re-checked only when the request reaches coderd. + rw.Header().Set("Cache-Control", "no-store") + api.serveChatFile(ctx, rw, chatFile) +} + // EXPERIMENTAL: this endpoint is experimental and is subject to change. // // @Summary Get chat file @@ -5980,6 +6125,14 @@ func (api *API) chatFileByID(rw http.ResponseWriter, r *http.Request) { return } + rw.Header().Set("Cache-Control", "private, max-age=31536000, immutable") + api.serveChatFile(r.Context(), rw, chatFile) +} + +// serveChatFile writes the file body and content headers; callers set +// Cache-Control because signed and session-authenticated downloads have +// different caching requirements. +func (api *API) serveChatFile(ctx context.Context, rw http.ResponseWriter, chatFile database.ChatFile) { rw.Header().Set("Content-Type", chatFile.Mimetype) disposition := "attachment" if chatfiles.IsInlineRenderableStoredMediaType(chatFile.Mimetype) { @@ -5990,7 +6143,6 @@ func (api *API) chatFileByID(rw http.ResponseWriter, r *http.Request) { } else { rw.Header().Set("Content-Disposition", disposition) } - rw.Header().Set("Cache-Control", "private, max-age=31536000, immutable") rw.Header().Set("Content-Length", strconv.Itoa(len(chatFile.Data))) rw.WriteHeader(http.StatusOK) if _, err := rw.Write(chatFile.Data); err != nil { diff --git a/coderd/exp_chats_test.go b/coderd/exp_chats_test.go index da87e26776..e57d4f089b 100644 --- a/coderd/exp_chats_test.go +++ b/coderd/exp_chats_test.go @@ -3,6 +3,7 @@ package coderd_test import ( "bytes" "context" + "crypto/sha256" "database/sql" "encoding/json" stderrors "errors" @@ -10,6 +11,7 @@ import ( "io" "mime" "net/http" + "net/url" "regexp" "slices" "strconv" @@ -18,6 +20,7 @@ import ( "testing" "time" + "github.com/go-jose/go-jose/v4/jwt" "github.com/google/uuid" "github.com/modelcontextprotocol/go-sdk/mcp" "github.com/sqlc-dev/pqtype" @@ -41,6 +44,7 @@ import ( "github.com/coder/coder/v2/coderd/database/dbtime" dbpubsub "github.com/coder/coder/v2/coderd/database/pubsub" "github.com/coder/coder/v2/coderd/externalauth" + "github.com/coder/coder/v2/coderd/jwtutils" "github.com/coder/coder/v2/coderd/rbac" "github.com/coder/coder/v2/coderd/rbac/policy" "github.com/coder/coder/v2/coderd/util/ptr" @@ -8071,6 +8075,7 @@ func TestChatMessageWithFiles(t *testing.T) { require.Equal(t, firstUser.UserID, f.OwnerID) require.NotEqual(t, uuid.Nil, f.OrganizationID) require.Equal(t, "image/png", f.MimeType) + require.Equal(t, int64(len(pngData)), f.SizeBytes) require.Equal(t, "test.png", f.Name) require.NotZero(t, f.CreatedAt) }) @@ -11272,6 +11277,214 @@ func TestGetChatFile(t *testing.T) { }) } +func TestChatFileDownloadURL(t *testing.T) { + t.Parallel() + + newClient := func(t *testing.T) (*codersdk.ExperimentalClient, jwtutils.StaticKey) { + t.Helper() + key := jwtutils.StaticKey{ID: "1", Key: []byte(strings.Repeat("k", 64))} + client := newChatClient(t, func(options *coderdtest.Options) { + options.ChatFileTokenKeyCache = key + }) + return client, key + } + uploadPNG := func(t *testing.T, ctx context.Context, client *codersdk.ExperimentalClient, organizationID uuid.UUID, name string) (codersdk.UploadChatFileResponse, []byte) { + t.Helper() + data := append([]byte{0x89, 0x50, 0x4E, 0x47, 0x0D, 0x0A, 0x1A, 0x0A}, make([]byte, 64)...) + uploaded, err := client.UploadChatFile(ctx, organizationID, "image/png", name, bytes.NewReader(data)) + require.NoError(t, err) + return uploaded, data + } + get := func(t *testing.T, ctx context.Context, rawURL string) *http.Response { + t.Helper() + req, err := http.NewRequestWithContext(ctx, http.MethodGet, rawURL, nil) + require.NoError(t, err) + res, err := http.DefaultClient.Do(req) + require.NoError(t, err) + return res + } + + t.Run("MintAndRedeem", func(t *testing.T) { + t.Parallel() + ctx := testutil.Context(t, testutil.WaitLong) + client, _ := newClient(t) + firstUser := coderdtest.CreateFirstUser(t, client.Client) + uploaded, data := uploadPNG(t, ctx, client, firstUser.OrganizationID, "evidence.png") + + download, err := client.ChatFileDownloadURL(ctx, uploaded.ID) + require.NoError(t, err) + require.Equal(t, int64(len(data)), download.SizeBytes) + require.Equal(t, fmt.Sprintf("%x", sha256.Sum256(data)), download.SHA256) + require.Equal(t, "evidence.png", download.Name) + require.Equal(t, "image/png", download.MimeType) + require.True(t, download.ExpiresAt.After(time.Now())) + // expires_at must match the JWT exp claim, which is stored at + // second precision; sub-second drift would overstate validity. + require.True(t, download.ExpiresAt.Equal(download.ExpiresAt.Truncate(time.Second))) + + res := get(t, ctx, download.URL) + defer res.Body.Close() + require.Equal(t, http.StatusOK, res.StatusCode) + require.Equal(t, "image/png", res.Header.Get("Content-Type")) + require.Equal(t, "no-store", res.Header.Get("Cache-Control")) + got, err := io.ReadAll(res.Body) + require.NoError(t, err) + require.Equal(t, data, got) + }) + + t.Run("ExpiredToken", func(t *testing.T) { + t.Parallel() + ctx := testutil.Context(t, testutil.WaitLong) + client, key := newClient(t) + firstUser := coderdtest.CreateFirstUser(t, client.Client) + uploaded, _ := uploadPNG(t, ctx, client, firstUser.OrganizationID, "expired.png") + token, err := jwtutils.Sign(ctx, key, coderd.ChatFileDownloadClaims{ + RegisteredClaims: jwtutils.RegisteredClaims{ + Expiry: jwt.NewNumericDate(time.Now().Add(-time.Minute)), + IssuedAt: jwt.NewNumericDate(time.Now().Add(-2 * time.Minute)), + }, + FileID: uploaded.ID, + UserID: firstUser.UserID, + }) + require.NoError(t, err) + downloadURL := client.URL.JoinPath("api", "experimental", "chats", "files", uploaded.ID.String(), "download") + query := downloadURL.Query() + query.Set("token", token) + downloadURL.RawQuery = query.Encode() + + res := get(t, ctx, downloadURL.String()) + defer res.Body.Close() + require.Equal(t, http.StatusNotFound, res.StatusCode) + }) + + t.Run("TamperedToken", func(t *testing.T) { + t.Parallel() + ctx := testutil.Context(t, testutil.WaitLong) + client, _ := newClient(t) + firstUser := coderdtest.CreateFirstUser(t, client.Client) + uploaded, _ := uploadPNG(t, ctx, client, firstUser.OrganizationID, "tampered.png") + download, err := client.ChatFileDownloadURL(ctx, uploaded.ID) + require.NoError(t, err) + downloadURL, err := url.Parse(download.URL) + require.NoError(t, err) + token := downloadURL.Query().Get("token") + require.NotEmpty(t, token) + first := byte('A') + if token[0] == first { + first = 'B' + } + query := downloadURL.Query() + query.Set("token", string(first)+token[1:]) + downloadURL.RawQuery = query.Encode() + + res := get(t, ctx, downloadURL.String()) + defer res.Body.Close() + require.Equal(t, http.StatusNotFound, res.StatusCode) + }) + + t.Run("TokenFileMismatch", func(t *testing.T) { + t.Parallel() + ctx := testutil.Context(t, testutil.WaitLong) + client, _ := newClient(t) + firstUser := coderdtest.CreateFirstUser(t, client.Client) + fileA, _ := uploadPNG(t, ctx, client, firstUser.OrganizationID, "a.png") + fileB, _ := uploadPNG(t, ctx, client, firstUser.OrganizationID, "b.png") + download, err := client.ChatFileDownloadURL(ctx, fileA.ID) + require.NoError(t, err) + downloadURL, err := url.Parse(download.URL) + require.NoError(t, err) + downloadURL.Path = fmt.Sprintf("/api/experimental/chats/files/%s/download", fileB.ID) + + res := get(t, ctx, downloadURL.String()) + defer res.Body.Close() + require.Equal(t, http.StatusNotFound, res.StatusCode) + }) + + t.Run("PlainGetRequiresAuthentication", func(t *testing.T) { + t.Parallel() + ctx := testutil.Context(t, testutil.WaitLong) + client, _ := newClient(t) + anonymous := codersdk.NewExperimentalClient(codersdk.New(client.URL)) + + _, _, err := anonymous.GetChatFile(ctx, uuid.New()) + requireSDKError(t, err, http.StatusUnauthorized) + }) + + t.Run("NonOwnerCannotMint", func(t *testing.T) { + t.Parallel() + ctx := testutil.Context(t, testutil.WaitLong) + client, _ := newClient(t) + firstUser := coderdtest.CreateFirstUser(t, client.Client) + uploaded, _ := uploadPNG(t, ctx, client, firstUser.OrganizationID, "owner.png") + otherClientRaw, _ := coderdtest.CreateAnotherUser(t, client.Client, firstUser.OrganizationID, rbac.ScopedRoleAgentsAccess(firstUser.OrganizationID)) + otherClient := codersdk.NewExperimentalClient(otherClientRaw) + + _, err := otherClient.ChatFileDownloadURL(ctx, uploaded.ID) + requireSDKError(t, err, http.StatusNotFound) + }) + + t.Run("RevokedShareCannotRedeem", func(t *testing.T) { + t.Parallel() + ctx := testutil.Context(t, testutil.WaitLong) + client, _ := newClient(t) + firstUser := coderdtest.CreateFirstUser(t, client.Client) + _ = createChatModelConfig(t, client) + uploaded, _ := uploadPNG(t, ctx, client, firstUser.OrganizationID, "shared.png") + chat, err := client.CreateChat(ctx, codersdk.CreateChatRequest{ + OrganizationID: firstUser.OrganizationID, + Content: []codersdk.ChatInputPart{ + {Type: codersdk.ChatInputPartTypeText, Text: "shared evidence"}, + {Type: codersdk.ChatInputPartTypeFile, FileID: uploaded.ID}, + }, + }) + require.NoError(t, err) + + memberRaw, member := coderdtest.CreateAnotherUser(t, client.Client, firstUser.OrganizationID, rbac.ScopedRoleAgentsAccess(firstUser.OrganizationID)) + memberClient := codersdk.NewExperimentalClient(memberRaw) + err = client.UpdateChatACL(ctx, chat.ID, codersdk.UpdateChatACL{ + UserRoles: map[string]codersdk.ChatRole{member.ID.String(): codersdk.ChatRoleRead}, + }) + require.NoError(t, err) + + // The member can mint while the share is active. + download, err := memberClient.ChatFileDownloadURL(ctx, uploaded.ID) + require.NoError(t, err) + + // Revoking the share invalidates the already-minted URL because + // redemption rechecks the minting user's access live. + err = client.UpdateChatACL(ctx, chat.ID, codersdk.UpdateChatACL{ + UserRoles: map[string]codersdk.ChatRole{member.ID.String(): codersdk.ChatRoleDeleted}, + }) + require.NoError(t, err) + + res := get(t, ctx, download.URL) + defer res.Body.Close() + require.Equal(t, http.StatusNotFound, res.StatusCode) + }) + + t.Run("SuspendedUserCannotRedeem", func(t *testing.T) { + t.Parallel() + ctx := testutil.Context(t, testutil.WaitLong) + client, _ := newClient(t) + firstUser := coderdtest.CreateFirstUser(t, client.Client) + memberRaw, member := coderdtest.CreateAnotherUser(t, client.Client, firstUser.OrganizationID, rbac.ScopedRoleAgentsAccess(firstUser.OrganizationID)) + memberClient := codersdk.NewExperimentalClient(memberRaw) + uploaded, _ := uploadPNG(t, ctx, memberClient, firstUser.OrganizationID, "suspended.png") + + download, err := memberClient.ChatFileDownloadURL(ctx, uploaded.ID) + require.NoError(t, err) + + // Suspending the minting user invalidates the URL because + // redemption requires the token's user to still be active. + _, err = client.Client.UpdateUserStatus(ctx, member.ID.String(), codersdk.UserStatusSuspended) + require.NoError(t, err) + + res := get(t, ctx, download.URL) + defer res.Body.Close() + require.Equal(t, http.StatusNotFound, res.StatusCode) + }) +} + // seedChatGatewayRequest records one finished Coder Agents gateway request // under sessionChatID, mirroring aibridged: the session ID is the spawning // chat, and each usage is one provider response within that one request. diff --git a/coderd/tracing/httpmw.go b/coderd/tracing/httpmw.go index 6c62ece2dd..657f5a11a5 100644 --- a/coderd/tracing/httpmw.go +++ b/coderd/tracing/httpmw.go @@ -42,8 +42,11 @@ func Middleware(tracerProvider trace.TracerProvider) func(http.Handler) http.Han } // Start span with default span name. Span name will be updated to - // "method route" format once request finishes. - r, span := StartHTTPSpan(tracer, rw, r, fmt.Sprintf("%s %s", r.Method, r.RequestURI)) + // "method route" format once request finishes. The initial name + // excludes the query string because span names are exported to + // tracing backends at span start and some endpoints accept + // bearer credentials as query parameters. + r, span := StartHTTPSpan(tracer, rw, r, fmt.Sprintf("%s %s", r.Method, r.URL.Path)) defer span.End() sw, ok := rw.(*StatusWriter) diff --git a/coderd/tracing/httpmw_test.go b/coderd/tracing/httpmw_test.go index 0f3611717e..52d45cabeb 100644 --- a/coderd/tracing/httpmw_test.go +++ b/coderd/tracing/httpmw_test.go @@ -5,11 +5,14 @@ import ( "net/http" "net/http/httptest" "strings" + "sync" "sync/atomic" "testing" "github.com/go-chi/chi/v5" "github.com/stretchr/testify/require" + sdktrace "go.opentelemetry.io/otel/sdk/trace" + "go.opentelemetry.io/otel/sdk/trace/tracetest" "go.opentelemetry.io/otel/trace" "go.opentelemetry.io/otel/trace/noop" @@ -43,6 +46,24 @@ func (f *fakeTracer) Start(ctx context.Context, _ string, _ ...trace.SpanStartOp return ctx, tracing.NoopSpan } +// startNameRecorder captures span names as they are at span start, before +// EndHTTPSpan renames them, because span names are exported to tracing +// backends at span start. +type startNameRecorder struct { + mu sync.Mutex + names []string +} + +func (r *startNameRecorder) OnStart(_ context.Context, s sdktrace.ReadWriteSpan) { + r.mu.Lock() + defer r.mu.Unlock() + r.names = append(r.names, s.Name()) +} + +func (*startNameRecorder) OnEnd(sdktrace.ReadOnlySpan) {} +func (*startNameRecorder) Shutdown(context.Context) error { return nil } +func (*startNameRecorder) ForceFlush(context.Context) error { return nil } + func Test_Middleware(t *testing.T) { t.Parallel() @@ -99,4 +120,47 @@ func Test_Middleware(t *testing.T) { }) } }) + + t.Run("QueryCredentialsNotExported", func(t *testing.T) { + t.Parallel() + + // Some endpoints accept bearer credentials as query parameters + // (e.g. signed chat file download tokens). No exported span name + // or attribute may carry the query string. + const token = "super-secret-download-token" + + startNames := &startNameRecorder{} + recorder := tracetest.NewSpanRecorder() + provider := sdktrace.NewTracerProvider( + sdktrace.WithSpanProcessor(startNames), + sdktrace.WithSpanProcessor(recorder), + ) + + rw := &tracing.StatusWriter{ResponseWriter: httptest.NewRecorder()} + r := httptest.NewRequest("GET", "/api/experimental/chats/files/abc/download?token="+token, nil) + + ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitLong) + defer cancel() + ctx = context.WithValue(ctx, chi.RouteCtxKey, chi.NewRouteContext()) + r = r.WithContext(ctx) + + tracing.Middleware(provider)(http.HandlerFunc(func(rw http.ResponseWriter, r *http.Request) { + rw.WriteHeader(http.StatusOK) + })).ServeHTTP(rw, r) + + require.NoError(t, provider.ForceFlush(ctx)) + require.NotEmpty(t, startNames.names) + for _, name := range startNames.names { + require.NotContains(t, name, token) + } + spans := recorder.Ended() + require.NotEmpty(t, spans) + for _, span := range spans { + require.NotContains(t, span.Name(), token) + for _, attr := range span.Attributes() { + require.NotContains(t, attr.Value.Emit(), token, + "span attribute %s must not carry query credentials", attr.Key) + } + } + }) } diff --git a/codersdk/chats.go b/codersdk/chats.go index de4f7696d0..ea03c48a0b 100644 --- a/codersdk/chats.go +++ b/codersdk/chats.go @@ -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) diff --git a/codersdk/deployment.go b/codersdk/deployment.go index 1817412829..9cf291bfc1 100644 --- a/codersdk/deployment.go +++ b/codersdk/deployment.go @@ -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 diff --git a/codersdk/toolsdk/chats.go b/codersdk/toolsdk/chats.go index 786d0eac59..947076570f 100644 --- a/codersdk/toolsdk/chats.go +++ b/codersdk/toolsdk/chats.go @@ -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 "". 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: (full-text, cannot combine with title, pr_title, or pr), title:, repo:, pr:, pr_title:, pr_status:, diff_url:, archived:, has_unread:, source:. 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 }, diff --git a/codersdk/toolsdk/chats_test.go b/codersdk/toolsdk/chats_test.go index 3d94666cb1..c4dbd0c5a9 100644 --- a/codersdk/toolsdk/chats_test.go +++ b/codersdk/toolsdk/chats_test.go @@ -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(), diff --git a/codersdk/toolsdk/toolsdk.go b/codersdk/toolsdk/toolsdk.go index f8ca4bcb1b..c808bc1b10 100644 --- a/codersdk/toolsdk/toolsdk.go +++ b/codersdk/toolsdk/toolsdk.go @@ -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 { diff --git a/codersdk/toolsdk/toolsdk_test.go b/codersdk/toolsdk/toolsdk_test.go index bd4949baaa..44414d85fe 100644 --- a/codersdk/toolsdk/toolsdk_test.go +++ b/codersdk/toolsdk/toolsdk_test.go @@ -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) } }) diff --git a/docs/reference/api/chats.md b/docs/reference/api/chats.md index 82567f1166..41400771f6 100644 --- a/docs/reference/api/chats.md +++ b/docs/reference/api/chats.md @@ -97,7 +97,8 @@ Experimental: this endpoint is subject to change. "mime_type": "string", "name": "string", "organization_id": "7c60d51f-b44e-4682-87d6-449835ea4de6", - "owner_id": "8826ee2e-7933-4665-aef2-2393f84a0d05" + "owner_id": "8826ee2e-7933-4665-aef2-2393f84a0d05", + "size_bytes": 0 } ], "has_unread": true, @@ -203,6 +204,7 @@ Status Code **200** | `»» name` | string | false | | | | `»» organization_id` | string(uuid) | false | | | | `»» owner_id` | string(uuid) | false | | | +| `»» size_bytes` | integer | false | | | | `» has_unread` | boolean | false | | Has unread is true when assistant messages exist beyond the owner's read cursor, which updates on stream connect and disconnect. | | `» id` | string(uuid) | false | | | | `» labels` | object | false | | | @@ -376,7 +378,8 @@ Experimental: this endpoint is subject to change. "mime_type": "string", "name": "string", "organization_id": "7c60d51f-b44e-4682-87d6-449835ea4de6", - "owner_id": "8826ee2e-7933-4665-aef2-2393f84a0d05" + "owner_id": "8826ee2e-7933-4665-aef2-2393f84a0d05", + "size_bytes": 0 } ], "has_unread": true, @@ -471,7 +474,8 @@ Experimental: this endpoint is subject to change. "mime_type": "string", "name": "string", "organization_id": "7c60d51f-b44e-4682-87d6-449835ea4de6", - "owner_id": "8826ee2e-7933-4665-aef2-2393f84a0d05" + "owner_id": "8826ee2e-7933-4665-aef2-2393f84a0d05", + "size_bytes": 0 } ], "has_unread": true, @@ -723,7 +727,8 @@ Experimental: this endpoint is subject to change. "mime_type": "string", "name": "string", "organization_id": "7c60d51f-b44e-4682-87d6-449835ea4de6", - "owner_id": "8826ee2e-7933-4665-aef2-2393f84a0d05" + "owner_id": "8826ee2e-7933-4665-aef2-2393f84a0d05", + "size_bytes": 0 } ], "has_unread": true, @@ -872,7 +877,8 @@ Experimental: this endpoint is subject to change. "mime_type": "string", "name": "string", "organization_id": "7c60d51f-b44e-4682-87d6-449835ea4de6", - "owner_id": "8826ee2e-7933-4665-aef2-2393f84a0d05" + "owner_id": "8826ee2e-7933-4665-aef2-2393f84a0d05", + "size_bytes": 0 } ], "has_unread": true, @@ -967,7 +973,8 @@ Experimental: this endpoint is subject to change. "mime_type": "string", "name": "string", "organization_id": "7c60d51f-b44e-4682-87d6-449835ea4de6", - "owner_id": "8826ee2e-7933-4665-aef2-2393f84a0d05" + "owner_id": "8826ee2e-7933-4665-aef2-2393f84a0d05", + "size_bytes": 0 } ], "has_unread": true, @@ -1153,7 +1160,8 @@ Experimental: this endpoint is subject to change. "mime_type": "string", "name": "string", "organization_id": "7c60d51f-b44e-4682-87d6-449835ea4de6", - "owner_id": "8826ee2e-7933-4665-aef2-2393f84a0d05" + "owner_id": "8826ee2e-7933-4665-aef2-2393f84a0d05", + "size_bytes": 0 } ], "has_unread": true, @@ -1248,7 +1256,8 @@ Experimental: this endpoint is subject to change. "mime_type": "string", "name": "string", "organization_id": "7c60d51f-b44e-4682-87d6-449835ea4de6", - "owner_id": "8826ee2e-7933-4665-aef2-2393f84a0d05" + "owner_id": "8826ee2e-7933-4665-aef2-2393f84a0d05", + "size_bytes": 0 } ], "has_unread": true, @@ -1484,7 +1493,8 @@ Experimental: this endpoint is subject to change. "mime_type": "string", "name": "string", "organization_id": "7c60d51f-b44e-4682-87d6-449835ea4de6", - "owner_id": "8826ee2e-7933-4665-aef2-2393f84a0d05" + "owner_id": "8826ee2e-7933-4665-aef2-2393f84a0d05", + "size_bytes": 0 } ], "has_unread": true, @@ -1579,7 +1589,8 @@ Experimental: this endpoint is subject to change. "mime_type": "string", "name": "string", "organization_id": "7c60d51f-b44e-4682-87d6-449835ea4de6", - "owner_id": "8826ee2e-7933-4665-aef2-2393f84a0d05" + "owner_id": "8826ee2e-7933-4665-aef2-2393f84a0d05", + "size_bytes": 0 } ], "has_unread": true, @@ -2501,7 +2512,8 @@ Experimental: this endpoint is subject to change. "mime_type": "string", "name": "string", "organization_id": "7c60d51f-b44e-4682-87d6-449835ea4de6", - "owner_id": "8826ee2e-7933-4665-aef2-2393f84a0d05" + "owner_id": "8826ee2e-7933-4665-aef2-2393f84a0d05", + "size_bytes": 0 } ], "has_unread": true, @@ -2596,7 +2608,8 @@ Experimental: this endpoint is subject to change. "mime_type": "string", "name": "string", "organization_id": "7c60d51f-b44e-4682-87d6-449835ea4de6", - "owner_id": "8826ee2e-7933-4665-aef2-2393f84a0d05" + "owner_id": "8826ee2e-7933-4665-aef2-2393f84a0d05", + "size_bytes": 0 } ], "has_unread": true, @@ -3105,7 +3118,8 @@ Experimental: this endpoint is subject to change. "mime_type": "string", "name": "string", "organization_id": "7c60d51f-b44e-4682-87d6-449835ea4de6", - "owner_id": "8826ee2e-7933-4665-aef2-2393f84a0d05" + "owner_id": "8826ee2e-7933-4665-aef2-2393f84a0d05", + "size_bytes": 0 } ], "has_unread": true, @@ -3200,7 +3214,8 @@ Experimental: this endpoint is subject to change. "mime_type": "string", "name": "string", "organization_id": "7c60d51f-b44e-4682-87d6-449835ea4de6", - "owner_id": "8826ee2e-7933-4665-aef2-2393f84a0d05" + "owner_id": "8826ee2e-7933-4665-aef2-2393f84a0d05", + "size_bytes": 0 } ], "has_unread": true, diff --git a/docs/reference/api/schemas.md b/docs/reference/api/schemas.md index e5ce947f4e..6383f70233 100644 --- a/docs/reference/api/schemas.md +++ b/docs/reference/api/schemas.md @@ -2250,7 +2250,8 @@ AuthorizationObject can represent a "set" of objects, such as: all workspaces in "mime_type": "string", "name": "string", "organization_id": "7c60d51f-b44e-4682-87d6-449835ea4de6", - "owner_id": "8826ee2e-7933-4665-aef2-2393f84a0d05" + "owner_id": "8826ee2e-7933-4665-aef2-2393f84a0d05", + "size_bytes": 0 } ], "has_unread": true, @@ -2345,7 +2346,8 @@ AuthorizationObject can represent a "set" of objects, such as: all workspaces in "mime_type": "string", "name": "string", "organization_id": "7c60d51f-b44e-4682-87d6-449835ea4de6", - "owner_id": "8826ee2e-7933-4665-aef2-2393f84a0d05" + "owner_id": "8826ee2e-7933-4665-aef2-2393f84a0d05", + "size_bytes": 0 } ], "has_unread": true, @@ -2793,6 +2795,30 @@ AuthorizationObject can represent a "set" of objects, such as: all workspaces in |----------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------| | `auth`, `config`, `content_filter`, `generic`, `hook_denied`, `hook_dispatch_failed`, `missing_key`, `overloaded`, `provider_disabled`, `rate_limit`, `stream_silence_timeout`, `timeout`, `usage_limit` | +## codersdk.ChatFileDownloadURLResponse + +```json +{ + "expires_at": "2019-08-24T14:15:22Z", + "mime_type": "string", + "name": "string", + "sha256": "string", + "size_bytes": 0, + "url": "http://example.com" +} +``` + +### Properties + +| Name | Type | Required | Restrictions | Description | +|--------------|---------|----------|--------------|-------------| +| `expires_at` | string | false | | | +| `mime_type` | string | false | | | +| `name` | string | false | | | +| `sha256` | string | false | | | +| `size_bytes` | integer | false | | | +| `url` | string | false | | | + ## codersdk.ChatFileMetadata ```json @@ -2802,20 +2828,22 @@ AuthorizationObject can represent a "set" of objects, such as: all workspaces in "mime_type": "string", "name": "string", "organization_id": "7c60d51f-b44e-4682-87d6-449835ea4de6", - "owner_id": "8826ee2e-7933-4665-aef2-2393f84a0d05" + "owner_id": "8826ee2e-7933-4665-aef2-2393f84a0d05", + "size_bytes": 0 } ``` ### Properties -| Name | Type | Required | Restrictions | Description | -|-------------------|--------|----------|--------------|-------------| -| `created_at` | string | false | | | -| `id` | string | false | | | -| `mime_type` | string | false | | | -| `name` | string | false | | | -| `organization_id` | string | false | | | -| `owner_id` | string | false | | | +| Name | Type | Required | Restrictions | Description | +|-------------------|---------|----------|--------------|-------------| +| `created_at` | string | false | | | +| `id` | string | false | | | +| `mime_type` | string | false | | | +| `name` | string | false | | | +| `organization_id` | string | false | | | +| `owner_id` | string | false | | | +| `size_bytes` | integer | false | | | ## codersdk.ChatGroup @@ -4176,7 +4204,8 @@ AuthorizationObject can represent a "set" of objects, such as: all workspaces in "mime_type": "string", "name": "string", "organization_id": "7c60d51f-b44e-4682-87d6-449835ea4de6", - "owner_id": "8826ee2e-7933-4665-aef2-2393f84a0d05" + "owner_id": "8826ee2e-7933-4665-aef2-2393f84a0d05", + "size_bytes": 0 } ], "has_unread": true, @@ -5580,9 +5609,9 @@ CreateWorkspaceRequest provides options for creating a new workspace. Only one o #### Enumerated Values -| Value(s) | -|-----------------------------------------------------------------------------------------------| -| `nats_ca`, `oidc_convert`, `tailnet_resume`, `workspace_apps_api_key`, `workspace_apps_token` | +| Value(s) | +|-------------------------------------------------------------------------------------------------------------------| +| `chat_files_token`, `nats_ca`, `oidc_convert`, `tailnet_resume`, `workspace_apps_api_key`, `workspace_apps_token` | ## codersdk.CustomNotificationContent diff --git a/site/src/api/typesGenerated.ts b/site/src/api/typesGenerated.ts index e532692338..54a1e5dfd6 100644 --- a/site/src/api/typesGenerated.ts +++ b/site/src/api/typesGenerated.ts @@ -2478,6 +2478,19 @@ export const ChatErrorKinds: ChatErrorKind[] = [ "usage_limit", ]; +// From codersdk/chats.go +/** + * ChatFileDownloadURLResponse contains a short-lived URL for downloading a chat file. + */ +export interface ChatFileDownloadURLResponse { + readonly url: string; + readonly expires_at: string; + readonly sha256: string; + readonly size_bytes: number; + readonly name: string; + readonly mime_type: string; +} + // From codersdk/chats.go /** * ChatFileMetadata contains lightweight metadata about a file @@ -2489,6 +2502,7 @@ export interface ChatFileMetadata { readonly organization_id: string; readonly name: string; readonly mime_type: string; + readonly size_bytes: number; readonly created_at: string; } @@ -4421,6 +4435,7 @@ export interface CryptoKey { // From codersdk/deployment.go export type CryptoKeyFeature = + | "chat_files_token" | "nats_ca" | "oidc_convert" | "tailnet_resume" @@ -4428,6 +4443,7 @@ export type CryptoKeyFeature = | "workspace_apps_token"; export const CryptoKeyFeatures: CryptoKeyFeature[] = [ + "chat_files_token", "nats_ca", "oidc_convert", "tailnet_resume",