diff --git a/coderd/aitasks.go b/coderd/aitasks.go index 8023917f68..0e0e93eb72 100644 --- a/coderd/aitasks.go +++ b/coderd/aitasks.go @@ -3,6 +3,7 @@ package coderd import ( "context" "database/sql" + "encoding/json" "errors" "fmt" "net" @@ -12,10 +13,12 @@ import ( "strings" "time" + "github.com/go-chi/chi/v5" "github.com/google/uuid" "golang.org/x/xerrors" - aiagentapi "github.com/coder/agentapi-sdk-go" + "cdr.dev/slog/v3" + agentapisdk "github.com/coder/agentapi-sdk-go" "github.com/coder/coder/v2/coderd/audit" "github.com/coder/coder/v2/coderd/database" "github.com/coder/coder/v2/coderd/database/dbtime" @@ -740,7 +743,7 @@ func (api *API) taskSend(rw http.ResponseWriter, r *http.Request) { } if err := api.authAndDoWithTaskAppClient(r, task, func(ctx context.Context, client *http.Client, appURL *url.URL) error { - agentAPIClient, err := aiagentapi.NewClient(appURL.String(), aiagentapi.WithHTTPClient(client)) + agentAPIClient, err := agentapisdk.NewClient(appURL.String(), agentapisdk.WithHTTPClient(client)) if err != nil { return httperror.NewResponseError(http.StatusBadGateway, codersdk.Response{ Message: "Failed to create agentapi client.", @@ -756,16 +759,16 @@ func (api *API) taskSend(rw http.ResponseWriter, r *http.Request) { }) } - if statusResp.Status != aiagentapi.StatusStable { + if statusResp.Status != agentapisdk.StatusStable { return httperror.NewResponseError(http.StatusBadGateway, codersdk.Response{ Message: "Task app is not ready to accept input.", Detail: fmt.Sprintf("Status: %s", statusResp.Status), }) } - _, err = agentAPIClient.PostMessage(ctx, aiagentapi.PostMessageParams{ + _, err = agentAPIClient.PostMessage(ctx, agentapisdk.PostMessageParams{ Content: req.Input, - Type: aiagentapi.MessageTypeUser, + Type: agentapisdk.MessageTypeUser, }) if err != nil { return httperror.NewResponseError(http.StatusBadGateway, codersdk.Response{ @@ -798,7 +801,7 @@ func (api *API) taskLogs(rw http.ResponseWriter, r *http.Request) { var out codersdk.TaskLogsResponse if err := api.authAndDoWithTaskAppClient(r, task, func(ctx context.Context, client *http.Client, appURL *url.URL) error { - agentAPIClient, err := aiagentapi.NewClient(appURL.String(), aiagentapi.WithHTTPClient(client)) + agentAPIClient, err := agentapisdk.NewClient(appURL.String(), agentapisdk.WithHTTPClient(client)) if err != nil { return httperror.NewResponseError(http.StatusBadGateway, codersdk.Response{ Message: "Failed to create agentapi client.", @@ -818,9 +821,9 @@ func (api *API) taskLogs(rw http.ResponseWriter, r *http.Request) { for _, m := range messagesResp.Messages { var typ codersdk.TaskLogType switch m.Role { - case aiagentapi.RoleUser: + case agentapisdk.RoleUser: typ = codersdk.TaskLogTypeInput - case aiagentapi.RoleAgent: + case agentapisdk.RoleAgent: typ = codersdk.TaskLogTypeOutput default: return httperror.NewResponseError(http.StatusBadGateway, codersdk.Response{ @@ -950,3 +953,165 @@ func (api *API) authAndDoWithTaskAppClient( } return do(ctx, client, parsedURL) } + +const ( + // taskSnapshotMaxSize is the maximum size for task log snapshots (64KB). + // Protects against excessive memory usage and database payload sizes. + taskSnapshotMaxSize = 64 * 1024 +) + +// TaskLogSnapshotEnvelope wraps a task log snapshot with format metadata. +type TaskLogSnapshotEnvelope struct { + Format string `json:"format"` + Data any `json:"data"` +} + +// @Summary Upload task log snapshot +// @ID upload-task-log-snapshot +// @Security CoderSessionToken +// @Accept json +// @Tags Tasks +// @Param task path string true "Task ID" format(uuid) +// @Param format query string true "Snapshot format" enums(agentapi) +// @Param request body object true "Raw snapshot payload (structure depends on format parameter)" +// @Success 204 +// @Router /workspaceagents/me/tasks/{task}/log-snapshot [post] +func (api *API) postWorkspaceAgentTaskLogSnapshot(rw http.ResponseWriter, r *http.Request) { + var ( + ctx = r.Context() + latestBuild = httpmw.LatestBuild(r) + ) + + // Parse task ID from path. + taskIDStr := chi.URLParam(r, "task") + taskID, err := uuid.Parse(taskIDStr) + if err != nil { + httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{ + Message: "Invalid task ID format.", + Detail: err.Error(), + }) + return + } + + // Validate format parameter (required). + p := httpapi.NewQueryParamParser().RequiredNotEmpty("format") + format := p.String(r.URL.Query(), "", "format") + p.ErrorExcessParams(r.URL.Query()) + if len(p.Errors) > 0 { + httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{ + Message: "Invalid query parameters.", + Validations: p.Errors, + }) + return + } + if format != "agentapi" { + httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{ + Message: "Invalid format parameter.", + Detail: fmt.Sprintf(`Only "agentapi" format is currently supported, got %q.`, format), + }) + return + } + + // Verify task exists before reading the potentially large payload. + // This prevents DoS attacks where attackers spam large payloads for + // non-existent or deleted tasks, forcing us to read 64KB into memory + // and do expensive JSON operations before the database rejects it. + // The UpsertTaskSnapshot will re-fetch for RBAC validation, but this + // early check protects against malicious load. + task, err := api.Database.GetTaskByID(ctx, taskID) + if err != nil { + if httpapi.Is404Error(err) { + httpapi.ResourceNotFound(rw) + return + } + httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{ + Message: "Internal error fetching task.", + Detail: err.Error(), + }) + return + } + + // Reject deleted tasks early. + if task.DeletedAt.Valid { + httpapi.ResourceNotFound(rw) + return + } + + // Verify task belongs to this agent's workspace. + if !task.WorkspaceID.Valid || task.WorkspaceID.UUID != latestBuild.WorkspaceID { + httpapi.ResourceNotFound(rw) + return + } + + // Limit payload size to avoid excessive memory or data usage. + r.Body = http.MaxBytesReader(rw, r.Body, taskSnapshotMaxSize) + + // Create envelope to store validated payload. + envelope := TaskLogSnapshotEnvelope{ + Format: format, + } + + switch format { + case "agentapi": + var payload agentapisdk.GetMessagesResponse + if err := json.NewDecoder(r.Body).Decode(&payload); err != nil { + httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{ + Message: "Failed to decode request payload.", + Detail: err.Error(), + }) + return + } + // Verify messages field exists (can be empty array). + if payload.Messages == nil { + httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{ + Message: "Invalid agentapi payload structure.", + Detail: `Missing required "messages" field.`, + }) + return + } + envelope.Data = payload + default: + // Defensive branch, we already validated "agentapi" format but may add + // more formats in the future. + httpapi.Write(ctx, rw, http.StatusBadRequest, codersdk.Response{ + Message: "Invalid format parameter.", + Detail: fmt.Sprintf(`Only "agentapi" format is currently supported, got %q.`, format), + }) + return + } + + // Marshal envelope with validated payload in a single pass. + snapshotJSON, err := json.Marshal(envelope) + if err != nil { + httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{ + Message: "Failed to create snapshot envelope.", + Detail: err.Error(), + }) + return + } + + // Upsert to database using agent's RBAC context. + err = api.Database.UpsertTaskSnapshot(ctx, database.UpsertTaskSnapshotParams{ + TaskID: task.ID, + LogSnapshot: json.RawMessage(snapshotJSON), + LogSnapshotCreatedAt: dbtime.Time(api.Clock.Now()), + }) + if err != nil { + if httpapi.IsUnauthorizedError(err) { + httpapi.ResourceNotFound(rw) + return + } + httpapi.Write(ctx, rw, http.StatusInternalServerError, codersdk.Response{ + Message: "Internal error storing snapshot.", + Detail: err.Error(), + }) + return + } + + api.Logger.Debug(ctx, "stored task log snapshot", + slog.F("task_id", task.ID), + slog.F("workspace_id", latestBuild.WorkspaceID), + slog.F("snapshot_size_bytes", len(snapshotJSON))) + + rw.WriteHeader(http.StatusNoContent) +} diff --git a/coderd/aitasks_test.go b/coderd/aitasks_test.go index 2c6a8de7ea..7940bc5272 100644 --- a/coderd/aitasks_test.go +++ b/coderd/aitasks_test.go @@ -1,12 +1,14 @@ package coderd_test import ( + "bytes" "context" "database/sql" "encoding/json" "io" "net/http" "net/http/httptest" + "strings" "testing" "time" @@ -17,6 +19,7 @@ import ( agentapisdk "github.com/coder/agentapi-sdk-go" "github.com/coder/coder/v2/agent" "github.com/coder/coder/v2/agent/agenttest" + "github.com/coder/coder/v2/coderd" "github.com/coder/coder/v2/coderd/coderdtest" "github.com/coder/coder/v2/coderd/database" "github.com/coder/coder/v2/coderd/database/dbauthz" @@ -1657,3 +1660,271 @@ func TestTasksNotification(t *testing.T) { }) } } + +func TestPostWorkspaceAgentTaskSnapshot(t *testing.T) { + t.Parallel() + + // Shared coderd with mock clock for all tests. + clock := quartz.NewMock(t) + ownerClient, db := coderdtest.NewWithDatabase(t, &coderdtest.Options{ + Clock: clock, + }) + owner := coderdtest.CreateFirstUser(t, ownerClient) + + createTaskWorkspace := func(t *testing.T, agentToken string) (taskID uuid.UUID, workspaceID uuid.UUID) { + t.Helper() + workspaceBuild := dbfake.WorkspaceBuild(t, db, database.WorkspaceTable{ + OrganizationID: owner.OrganizationID, + OwnerID: owner.UserID, + }).WithTask(database.TaskTable{ + Prompt: "test prompt", + }, &proto.App{ + Slug: "task-app", + Url: "http://localhost:8080", + }).WithAgent(func(agents []*proto.Agent) []*proto.Agent { + agents[0].Auth = &proto.Agent_Token{Token: agentToken} + return agents + }).Do() + return workspaceBuild.Task.ID, workspaceBuild.Workspace.ID + } + + makePayload := func(t *testing.T, content string) []byte { + t.Helper() + data := agentapisdk.GetMessagesResponse{ + Messages: []agentapisdk.Message{ + {Id: 0, Role: "agent", Content: content, Time: time.Now()}, + }, + } + b, err := json.Marshal(data) + require.NoError(t, err) + return b + } + + makeRequest := func(t *testing.T, taskID uuid.UUID, agentToken string, payload []byte, format string) *http.Response { + t.Helper() + ctx := testutil.Context(t, testutil.WaitShort) + + url := ownerClient.URL.JoinPath("/api/v2/workspaceagents/me/tasks", taskID.String(), "log-snapshot").String() + if format != "" { + url += "?format=" + format + } + + req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(payload)) + require.NoError(t, err) + req.Header.Set(codersdk.SessionTokenHeader, agentToken) + res, err := http.DefaultClient.Do(req) + require.NoError(t, err) + return res + } + + unmarshalSnapshot := func(t *testing.T, snapshotJSON json.RawMessage) agentapisdk.GetMessagesResponse { + t.Helper() + // Pre-populate Data with the correct type so json.Unmarshal decodes + // directly into it instead of creating a map[string]any. + envelope := coderd.TaskLogSnapshotEnvelope{ + Data: &agentapisdk.GetMessagesResponse{}, + } + err := json.Unmarshal(snapshotJSON, &envelope) + require.NoError(t, err) + require.Equal(t, "agentapi", envelope.Format) + + return *envelope.Data.(*agentapisdk.GetMessagesResponse) + } + + t.Run("Success", func(t *testing.T) { + t.Parallel() + agentToken := uuid.NewString() + taskID, _ := createTaskWorkspace(t, agentToken) + ctx := testutil.Context(t, testutil.WaitShort) + + res := makeRequest(t, taskID, agentToken, makePayload(t, "test"), "agentapi") + defer res.Body.Close() + require.Equal(t, http.StatusNoContent, res.StatusCode) + + snapshot, err := db.GetTaskSnapshot(dbauthz.AsSystemRestricted(ctx), taskID) + require.NoError(t, err) + + data := unmarshalSnapshot(t, snapshot.LogSnapshot) + require.Len(t, data.Messages, 1) + require.Equal(t, "test", data.Messages[0].Content) + }) + + //nolint:paralleltest // Not parallel, advances shared clock. + t.Run("Overwrite", func(t *testing.T) { + agentToken := uuid.NewString() + taskID, _ := createTaskWorkspace(t, agentToken) + ctx := testutil.Context(t, testutil.WaitShort) + + // First snapshot. + res1 := makeRequest(t, taskID, agentToken, makePayload(t, "first"), "agentapi") + res1.Body.Close() + require.Equal(t, http.StatusNoContent, res1.StatusCode) + + snapshot1, err := db.GetTaskSnapshot(dbauthz.AsSystemRestricted(ctx), taskID) + require.NoError(t, err) + firstTime := snapshot1.LogSnapshotCreatedAt + + // Advance clock to ensure timestamp differs. + clock.Advance(time.Second) + + // Second snapshot. + res2 := makeRequest(t, taskID, agentToken, makePayload(t, "second"), "agentapi") + res2.Body.Close() + require.Equal(t, http.StatusNoContent, res2.StatusCode) + + snapshot2, err := db.GetTaskSnapshot(dbauthz.AsSystemRestricted(ctx), taskID) + require.NoError(t, err) + require.True(t, snapshot2.LogSnapshotCreatedAt.After(firstTime)) + + // Verify data was overwritten. + data := unmarshalSnapshot(t, snapshot2.LogSnapshot) + require.Len(t, data.Messages, 1) + require.Equal(t, "second", data.Messages[0].Content) + }) + + t.Run("MissingFormat", func(t *testing.T) { + t.Parallel() + agentToken := uuid.NewString() + taskID, _ := createTaskWorkspace(t, agentToken) + + res := makeRequest(t, taskID, agentToken, makePayload(t, "test"), "") + defer res.Body.Close() + require.Equal(t, http.StatusBadRequest, res.StatusCode) + + var errResp codersdk.Response + json.NewDecoder(res.Body).Decode(&errResp) + require.Contains(t, errResp.Message, "Invalid query parameters") + require.Len(t, errResp.Validations, 1) + require.Equal(t, "format", errResp.Validations[0].Field) + require.Contains(t, errResp.Validations[0].Detail, "required and cannot be empty") + }) + + t.Run("InvalidFormat", func(t *testing.T) { + t.Parallel() + agentToken := uuid.NewString() + taskID, _ := createTaskWorkspace(t, agentToken) + + res := makeRequest(t, taskID, agentToken, makePayload(t, "test"), "unknown") + defer res.Body.Close() + require.Equal(t, http.StatusBadRequest, res.StatusCode) + + var errResp codersdk.Response + json.NewDecoder(res.Body).Decode(&errResp) + require.Contains(t, errResp.Message, "Invalid format parameter") + }) + + t.Run("PayloadTooLarge", func(t *testing.T) { + t.Parallel() + agentToken := uuid.NewString() + taskID, _ := createTaskWorkspace(t, agentToken) + + largeContent := strings.Repeat("x", 65*1024) + payload := makePayload(t, largeContent) + + res := makeRequest(t, taskID, agentToken, payload, "agentapi") + require.Equal(t, http.StatusBadRequest, res.StatusCode) + res.Body.Close() + }) + + t.Run("InvalidTaskID", func(t *testing.T) { + t.Parallel() + agentToken := uuid.NewString() + createTaskWorkspace(t, agentToken) + ctx := testutil.Context(t, testutil.WaitShort) + + url := ownerClient.URL.JoinPath("/api/v2/workspaceagents/me/tasks", "not-a-uuid", "log-snapshot").String() + "?format=agentapi" + req, _ := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(makePayload(t, "test"))) + req.Header.Set(codersdk.SessionTokenHeader, agentToken) + res, err := http.DefaultClient.Do(req) + require.NoError(t, err) + defer res.Body.Close() + require.Equal(t, http.StatusBadRequest, res.StatusCode) + + var errResp codersdk.Response + json.NewDecoder(res.Body).Decode(&errResp) + require.Contains(t, errResp.Message, "Invalid task ID format") + }) + + t.Run("TaskNotFound", func(t *testing.T) { + t.Parallel() + agentToken := uuid.NewString() + createTaskWorkspace(t, agentToken) + + res := makeRequest(t, uuid.New(), agentToken, makePayload(t, "test"), "agentapi") + defer res.Body.Close() + require.Equal(t, http.StatusNotFound, res.StatusCode) + }) + + t.Run("WrongWorkspace", func(t *testing.T) { + t.Parallel() + agent1Token := uuid.NewString() + agent2Token := uuid.NewString() + taskID1, _ := createTaskWorkspace(t, agent1Token) + taskID2, _ := createTaskWorkspace(t, agent2Token) + + // Try to POST snapshot for task2 using agent1's token. + res := makeRequest(t, taskID2, agent1Token, makePayload(t, "test"), "agentapi") + defer res.Body.Close() + require.Equal(t, http.StatusNotFound, res.StatusCode) + + // Verify we CAN post for our own task. + res2 := makeRequest(t, taskID1, agent1Token, makePayload(t, "test"), "agentapi") + defer res2.Body.Close() + require.Equal(t, http.StatusNoContent, res2.StatusCode) + }) + + t.Run("Unauthorized", func(t *testing.T) { + t.Parallel() + agentToken := uuid.NewString() + taskID, _ := createTaskWorkspace(t, agentToken) + + res := makeRequest(t, taskID, "", makePayload(t, "test"), "agentapi") + defer res.Body.Close() + require.Equal(t, http.StatusUnauthorized, res.StatusCode) + }) + + t.Run("MalformedJSON", func(t *testing.T) { + t.Parallel() + agentToken := uuid.NewString() + taskID, _ := createTaskWorkspace(t, agentToken) + + res := makeRequest(t, taskID, agentToken, []byte("{invalid json"), "agentapi") + defer res.Body.Close() + require.Equal(t, http.StatusBadRequest, res.StatusCode) + + var errResp codersdk.Response + json.NewDecoder(res.Body).Decode(&errResp) + require.Contains(t, errResp.Message, "Failed to decode request payload") + }) + + t.Run("InvalidAgentAPIPayload", func(t *testing.T) { + t.Parallel() + agentToken := uuid.NewString() + taskID, _ := createTaskWorkspace(t, agentToken) + + // Missing required "messages" field. + res := makeRequest(t, taskID, agentToken, []byte(`{"truncated":false,"total_count":0}`), "agentapi") + defer res.Body.Close() + require.Equal(t, http.StatusBadRequest, res.StatusCode) + + var errResp codersdk.Response + json.NewDecoder(res.Body).Decode(&errResp) + require.Contains(t, errResp.Message, "Invalid agentapi payload structure") + }) + + t.Run("DeletedTask", func(t *testing.T) { + t.Parallel() + agentToken := uuid.NewString() + taskID, _ := createTaskWorkspace(t, agentToken) + ctx := testutil.Context(t, testutil.WaitShort) + + // Delete the task. + err := ownerClient.DeleteTask(ctx, owner.UserID.String(), taskID) + require.NoError(t, err) + + res := makeRequest(t, taskID, agentToken, makePayload(t, "test"), "agentapi") + defer res.Body.Close() + // Agent token becomes invalid after task deletion. + require.Equal(t, http.StatusUnauthorized, res.StatusCode) + }) +} diff --git a/coderd/apidoc/docs.go b/coderd/apidoc/docs.go index e0152f9ed9..7c3a1130db 100644 --- a/coderd/apidoc/docs.go +++ b/coderd/apidoc/docs.go @@ -9556,6 +9556,57 @@ const docTemplate = `{ } } }, + "/workspaceagents/me/tasks/{task}/log-snapshot": { + "post": { + "security": [ + { + "CoderSessionToken": [] + } + ], + "consumes": [ + "application/json" + ], + "tags": [ + "Tasks" + ], + "summary": "Upload task log snapshot", + "operationId": "upload-task-log-snapshot", + "parameters": [ + { + "type": "string", + "format": "uuid", + "description": "Task ID", + "name": "task", + "in": "path", + "required": true + }, + { + "enum": [ + "agentapi" + ], + "type": "string", + "description": "Snapshot format", + "name": "format", + "in": "query", + "required": true + }, + { + "description": "Raw snapshot payload (structure depends on format parameter)", + "name": "request", + "in": "body", + "required": true, + "schema": { + "type": "object" + } + } + ], + "responses": { + "204": { + "description": "No Content" + } + } + } + }, "/workspaceagents/{workspaceagent}": { "get": { "security": [ diff --git a/coderd/apidoc/swagger.json b/coderd/apidoc/swagger.json index 6eb0ad5868..1c4438bfaa 100644 --- a/coderd/apidoc/swagger.json +++ b/coderd/apidoc/swagger.json @@ -8449,6 +8449,51 @@ } } }, + "/workspaceagents/me/tasks/{task}/log-snapshot": { + "post": { + "security": [ + { + "CoderSessionToken": [] + } + ], + "consumes": ["application/json"], + "tags": ["Tasks"], + "summary": "Upload task log snapshot", + "operationId": "upload-task-log-snapshot", + "parameters": [ + { + "type": "string", + "format": "uuid", + "description": "Task ID", + "name": "task", + "in": "path", + "required": true + }, + { + "enum": ["agentapi"], + "type": "string", + "description": "Snapshot format", + "name": "format", + "in": "query", + "required": true + }, + { + "description": "Raw snapshot payload (structure depends on format parameter)", + "name": "request", + "in": "body", + "required": true, + "schema": { + "type": "object" + } + } + ], + "responses": { + "204": { + "description": "No Content" + } + } + } + }, "/workspaceagents/{workspaceagent}": { "get": { "security": [ diff --git a/coderd/coderd.go b/coderd/coderd.go index b53f78e56b..eeda351b52 100644 --- a/coderd/coderd.go +++ b/coderd/coderd.go @@ -1448,6 +1448,9 @@ func New(options *Options) *API { r.Get("/gitsshkey", api.agentGitSSHKey) r.Post("/log-source", api.workspaceAgentPostLogSource) r.Get("/reinit", api.workspaceAgentReinit) + r.Route("/tasks/{task}", func(r chi.Router) { + r.Post("/log-snapshot", api.postWorkspaceAgentTaskLogSnapshot) + }) }) r.Route("/{workspaceagent}", func(r chi.Router) { r.Use( diff --git a/coderd/database/dbauthz/dbauthz.go b/coderd/database/dbauthz/dbauthz.go index 3083914b4c..a427d32afb 100644 --- a/coderd/database/dbauthz/dbauthz.go +++ b/coderd/database/dbauthz/dbauthz.go @@ -3055,6 +3055,25 @@ func (q *querier) GetTaskByWorkspaceID(ctx context.Context, workspaceID uuid.UUI return fetch(q.log, q.auth, q.db.GetTaskByWorkspaceID)(ctx, workspaceID) } +func (q *querier) GetTaskSnapshot(ctx context.Context, taskID uuid.UUID) (database.TaskSnapshot, error) { + // Fetch task to build RBAC object for authorization. + task, err := q.GetTaskByID(ctx, taskID) + if err != nil { + return database.TaskSnapshot{}, err + } + + obj := rbac.ResourceTask. + WithID(task.ID). + WithOwner(task.OwnerID.String()). + InOrg(task.OrganizationID) + + if err := q.authorizeContext(ctx, policy.ActionRead, obj); err != nil { + return database.TaskSnapshot{}, err + } + + return q.db.GetTaskSnapshot(ctx, taskID) +} + func (q *querier) GetTelemetryItem(ctx context.Context, key string) (database.TelemetryItem, error) { if err := q.authorizeContext(ctx, policy.ActionRead, rbac.ResourceSystem); err != nil { return database.TelemetryItem{}, err @@ -6024,6 +6043,25 @@ func (q *querier) UpsertTailnetTunnel(ctx context.Context, arg database.UpsertTa return q.db.UpsertTailnetTunnel(ctx, arg) } +func (q *querier) UpsertTaskSnapshot(ctx context.Context, arg database.UpsertTaskSnapshotParams) error { + // Fetch task to build RBAC object for authorization. + task, err := q.GetTaskByID(ctx, arg.TaskID) + if err != nil { + return err + } + + obj := rbac.ResourceTask. + WithID(task.ID). + WithOwner(task.OwnerID.String()). + InOrg(task.OrganizationID) + + if err := q.authorizeContext(ctx, policy.ActionUpdate, obj); err != nil { + return err + } + + return q.db.UpsertTaskSnapshot(ctx, arg) +} + func (q *querier) UpsertTaskWorkspaceApp(ctx context.Context, arg database.UpsertTaskWorkspaceAppParams) (database.TaskWorkspaceApp, error) { // Fetch the task to derive the RBAC object and authorize update on it. task, err := q.db.GetTaskByID(ctx, arg.TaskID) diff --git a/coderd/database/dbauthz/dbauthz_test.go b/coderd/database/dbauthz/dbauthz_test.go index b69b0493b4..08ce4ef2f3 100644 --- a/coderd/database/dbauthz/dbauthz_test.go +++ b/coderd/database/dbauthz/dbauthz_test.go @@ -2547,6 +2547,24 @@ func (s *MethodTestSuite) TestTasks() { dbm.EXPECT().ListTasks(gomock.Any(), gomock.Any()).Return([]database.Task{t1, t2}, nil).AnyTimes() check.Args(database.ListTasksParams{}).Asserts(t1, policy.ActionRead, t2, policy.ActionRead).Returns([]database.Task{t1, t2}) })) + s.Run("GetTaskSnapshot", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { + task := testutil.Fake(s.T(), faker, database.Task{}) + snapshot := testutil.Fake(s.T(), faker, database.TaskSnapshot{TaskID: task.ID}) + dbm.EXPECT().GetTaskByID(gomock.Any(), task.ID).Return(task, nil).AnyTimes() + dbm.EXPECT().GetTaskSnapshot(gomock.Any(), task.ID).Return(snapshot, nil).AnyTimes() + check.Args(task.ID).Asserts(task, policy.ActionRead, task, policy.ActionRead).Returns(snapshot) + })) + s.Run("UpsertTaskSnapshot", s.Mocked(func(dbm *dbmock.MockStore, faker *gofakeit.Faker, check *expects) { + task := testutil.Fake(s.T(), faker, database.Task{}) + arg := database.UpsertTaskSnapshotParams{ + TaskID: task.ID, + LogSnapshot: []byte(`{"format":"agentapi","data":[]}`), + LogSnapshotCreatedAt: dbtime.Now(), + } + dbm.EXPECT().GetTaskByID(gomock.Any(), task.ID).Return(task, nil).AnyTimes() + dbm.EXPECT().UpsertTaskSnapshot(gomock.Any(), arg).Return(nil).AnyTimes() + check.Args(arg).Asserts(task, policy.ActionRead, task, policy.ActionUpdate).Returns() + })) } func (s *MethodTestSuite) TestProvisionerKeys() { diff --git a/coderd/database/dbmetrics/querymetrics.go b/coderd/database/dbmetrics/querymetrics.go index 65e0b38a28..97469853fc 100644 --- a/coderd/database/dbmetrics/querymetrics.go +++ b/coderd/database/dbmetrics/querymetrics.go @@ -1758,6 +1758,14 @@ func (m queryMetricsStore) GetTaskByWorkspaceID(ctx context.Context, workspaceID return r0, r1 } +func (m queryMetricsStore) GetTaskSnapshot(ctx context.Context, taskID uuid.UUID) (database.TaskSnapshot, error) { + start := time.Now() + r0, r1 := m.s.GetTaskSnapshot(ctx, taskID) + m.queryLatencies.WithLabelValues("GetTaskSnapshot").Observe(time.Since(start).Seconds()) + m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "GetTaskSnapshot").Inc() + return r0, r1 +} + func (m queryMetricsStore) GetTelemetryItem(ctx context.Context, key string) (database.TelemetryItem, error) { start := time.Now() r0, r1 := m.s.GetTelemetryItem(ctx, key) @@ -4189,6 +4197,14 @@ func (m queryMetricsStore) UpsertTailnetTunnel(ctx context.Context, arg database return r0, r1 } +func (m queryMetricsStore) UpsertTaskSnapshot(ctx context.Context, arg database.UpsertTaskSnapshotParams) error { + start := time.Now() + r0 := m.s.UpsertTaskSnapshot(ctx, arg) + m.queryLatencies.WithLabelValues("UpsertTaskSnapshot").Observe(time.Since(start).Seconds()) + m.queryCounts.WithLabelValues(httpmw.ExtractHTTPRoute(ctx), httpmw.ExtractHTTPMethod(ctx), "UpsertTaskSnapshot").Inc() + return r0 +} + func (m queryMetricsStore) UpsertTaskWorkspaceApp(ctx context.Context, arg database.UpsertTaskWorkspaceAppParams) (database.TaskWorkspaceApp, error) { start := time.Now() r0, r1 := m.s.UpsertTaskWorkspaceApp(ctx, arg) diff --git a/coderd/database/dbmock/dbmock.go b/coderd/database/dbmock/dbmock.go index cf989e216f..7e329cfe5f 100644 --- a/coderd/database/dbmock/dbmock.go +++ b/coderd/database/dbmock/dbmock.go @@ -3254,6 +3254,21 @@ func (mr *MockStoreMockRecorder) GetTaskByWorkspaceID(ctx, workspaceID any) *gom return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetTaskByWorkspaceID", reflect.TypeOf((*MockStore)(nil).GetTaskByWorkspaceID), ctx, workspaceID) } +// GetTaskSnapshot mocks base method. +func (m *MockStore) GetTaskSnapshot(ctx context.Context, taskID uuid.UUID) (database.TaskSnapshot, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetTaskSnapshot", ctx, taskID) + ret0, _ := ret[0].(database.TaskSnapshot) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetTaskSnapshot indicates an expected call of GetTaskSnapshot. +func (mr *MockStoreMockRecorder) GetTaskSnapshot(ctx, taskID any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetTaskSnapshot", reflect.TypeOf((*MockStore)(nil).GetTaskSnapshot), ctx, taskID) +} + // GetTelemetryItem mocks base method. func (m *MockStore) GetTelemetryItem(ctx context.Context, key string) (database.TelemetryItem, error) { m.ctrl.T.Helper() @@ -7817,6 +7832,20 @@ func (mr *MockStoreMockRecorder) UpsertTailnetTunnel(ctx, arg any) *gomock.Call return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UpsertTailnetTunnel", reflect.TypeOf((*MockStore)(nil).UpsertTailnetTunnel), ctx, arg) } +// UpsertTaskSnapshot mocks base method. +func (m *MockStore) UpsertTaskSnapshot(ctx context.Context, arg database.UpsertTaskSnapshotParams) error { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "UpsertTaskSnapshot", ctx, arg) + ret0, _ := ret[0].(error) + return ret0 +} + +// UpsertTaskSnapshot indicates an expected call of UpsertTaskSnapshot. +func (mr *MockStoreMockRecorder) UpsertTaskSnapshot(ctx, arg any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UpsertTaskSnapshot", reflect.TypeOf((*MockStore)(nil).UpsertTaskSnapshot), ctx, arg) +} + // UpsertTaskWorkspaceApp mocks base method. func (m *MockStore) UpsertTaskWorkspaceApp(ctx context.Context, arg database.UpsertTaskWorkspaceAppParams) (database.TaskWorkspaceApp, error) { m.ctrl.T.Helper() diff --git a/coderd/database/querier.go b/coderd/database/querier.go index 79088d846d..d4c3e04a4c 100644 --- a/coderd/database/querier.go +++ b/coderd/database/querier.go @@ -353,6 +353,7 @@ type sqlcQuerier interface { GetTaskByID(ctx context.Context, id uuid.UUID) (Task, error) GetTaskByOwnerIDAndName(ctx context.Context, arg GetTaskByOwnerIDAndNameParams) (Task, error) GetTaskByWorkspaceID(ctx context.Context, workspaceID uuid.UUID) (Task, error) + GetTaskSnapshot(ctx context.Context, taskID uuid.UUID) (TaskSnapshot, error) GetTelemetryItem(ctx context.Context, key string) (TelemetryItem, error) GetTelemetryItems(ctx context.Context) ([]TelemetryItem, error) // GetTemplateAppInsights returns the aggregate usage of each app in a given @@ -772,6 +773,7 @@ type sqlcQuerier interface { UpsertTailnetCoordinator(ctx context.Context, id uuid.UUID) (TailnetCoordinator, error) UpsertTailnetPeer(ctx context.Context, arg UpsertTailnetPeerParams) (TailnetPeer, error) UpsertTailnetTunnel(ctx context.Context, arg UpsertTailnetTunnelParams) (TailnetTunnel, error) + UpsertTaskSnapshot(ctx context.Context, arg UpsertTaskSnapshotParams) error UpsertTaskWorkspaceApp(ctx context.Context, arg UpsertTaskWorkspaceAppParams) (TaskWorkspaceApp, error) UpsertTelemetryItem(ctx context.Context, arg UpsertTelemetryItemParams) error // This query aggregates the workspace_agent_stats and workspace_app_stats data diff --git a/coderd/database/queries.sql.go b/coderd/database/queries.sql.go index 7d84db116c..8a1188a2e0 100644 --- a/coderd/database/queries.sql.go +++ b/coderd/database/queries.sql.go @@ -13183,6 +13183,22 @@ func (q *sqlQuerier) GetTaskByWorkspaceID(ctx context.Context, workspaceID uuid. return i, err } +const getTaskSnapshot = `-- name: GetTaskSnapshot :one +SELECT + task_id, log_snapshot, log_snapshot_created_at +FROM + task_snapshots +WHERE + task_id = $1 +` + +func (q *sqlQuerier) GetTaskSnapshot(ctx context.Context, taskID uuid.UUID) (TaskSnapshot, error) { + row := q.db.QueryRowContext(ctx, getTaskSnapshot, taskID) + var i TaskSnapshot + err := row.Scan(&i.TaskID, &i.LogSnapshot, &i.LogSnapshotCreatedAt) + return i, err +} + const insertTask = `-- name: InsertTask :one INSERT INTO tasks (id, organization_id, owner_id, name, display_name, workspace_id, template_version_id, template_parameters, prompt, created_at) @@ -13373,6 +13389,29 @@ func (q *sqlQuerier) UpdateTaskWorkspaceID(ctx context.Context, arg UpdateTaskWo return i, err } +const upsertTaskSnapshot = `-- name: UpsertTaskSnapshot :exec +INSERT INTO + task_snapshots (task_id, log_snapshot, log_snapshot_created_at) +VALUES + ($1, $2, $3) +ON CONFLICT + (task_id) +DO UPDATE SET + log_snapshot = EXCLUDED.log_snapshot, + log_snapshot_created_at = EXCLUDED.log_snapshot_created_at +` + +type UpsertTaskSnapshotParams struct { + TaskID uuid.UUID `db:"task_id" json:"task_id"` + LogSnapshot json.RawMessage `db:"log_snapshot" json:"log_snapshot"` + LogSnapshotCreatedAt time.Time `db:"log_snapshot_created_at" json:"log_snapshot_created_at"` +} + +func (q *sqlQuerier) UpsertTaskSnapshot(ctx context.Context, arg UpsertTaskSnapshotParams) error { + _, err := q.db.ExecContext(ctx, upsertTaskSnapshot, arg.TaskID, arg.LogSnapshot, arg.LogSnapshotCreatedAt) + return err +} + const upsertTaskWorkspaceApp = `-- name: UpsertTaskWorkspaceApp :one INSERT INTO task_workspace_apps (task_id, workspace_build_number, workspace_agent_id, workspace_app_id) diff --git a/coderd/database/queries/tasks.sql b/coderd/database/queries/tasks.sql index 52e259953f..8deda80a2b 100644 --- a/coderd/database/queries/tasks.sql +++ b/coderd/database/queries/tasks.sql @@ -75,3 +75,22 @@ WHERE id = @id::uuid AND deleted_at IS NULL RETURNING *; + +-- name: UpsertTaskSnapshot :exec +INSERT INTO + task_snapshots (task_id, log_snapshot, log_snapshot_created_at) +VALUES + ($1, $2, $3) +ON CONFLICT + (task_id) +DO UPDATE SET + log_snapshot = EXCLUDED.log_snapshot, + log_snapshot_created_at = EXCLUDED.log_snapshot_created_at; + +-- name: GetTaskSnapshot :one +SELECT + * +FROM + task_snapshots +WHERE + task_id = $1; diff --git a/docs/reference/api/tasks.md b/docs/reference/api/tasks.md index 7a85fccefb..1952e8f023 100644 --- a/docs/reference/api/tasks.md +++ b/docs/reference/api/tasks.md @@ -399,3 +399,44 @@ curl -X POST http://coder-server:8080/api/v2/tasks/{user}/{task}/send \ | 204 | [No Content](https://tools.ietf.org/html/rfc7231#section-6.3.5) | No Content | | To perform this operation, you must be authenticated. [Learn more](authentication.md). + +## Upload task log snapshot + +### Code samples + +```shell +# Example request using curl +curl -X POST http://coder-server:8080/api/v2/workspaceagents/me/tasks/{task}/log-snapshot?format=agentapi \ + -H 'Content-Type: application/json' \ + -H 'Coder-Session-Token: API_KEY' +``` + +`POST /workspaceagents/me/tasks/{task}/log-snapshot` + +> Body parameter + +```json +{} +``` + +### Parameters + +| Name | In | Type | Required | Description | +|----------|-------|--------------|----------|--------------------------------------------------------------| +| `task` | path | string(uuid) | true | Task ID | +| `format` | query | string | true | Snapshot format | +| `body` | body | object | true | Raw snapshot payload (structure depends on format parameter) | + +#### Enumerated Values + +| Parameter | Value(s) | +|-----------|------------| +| `format` | `agentapi` | + +### Responses + +| Status | Meaning | Description | Schema | +|--------|-----------------------------------------------------------------|-------------|--------| +| 204 | [No Content](https://tools.ietf.org/html/rfc7231#section-6.3.5) | No Content | | + +To perform this operation, you must be authenticated. [Learn more](authentication.md).