diff --git a/agent/agent.go b/agent/agent.go index 5386ef19a1..45d4bb0ac2 100644 --- a/agent/agent.go +++ b/agent/agent.go @@ -40,6 +40,7 @@ import ( "github.com/coder/clistat" "github.com/coder/coder/v2/agent/agentcontainers" "github.com/coder/coder/v2/agent/agentexec" + "github.com/coder/coder/v2/agent/agentfiles" "github.com/coder/coder/v2/agent/agentscripts" "github.com/coder/coder/v2/agent/agentsocket" "github.com/coder/coder/v2/agent/agentssh" @@ -295,6 +296,8 @@ type agent struct { containerAPIOptions []agentcontainers.Option containerAPI *agentcontainers.API + filesAPI *agentfiles.API + socketServerEnabled bool socketPath string socketServer *agentsocket.Server @@ -365,6 +368,8 @@ func (a *agent) init() { a.containerAPI = agentcontainers.NewAPI(a.logger.Named("containers"), containerAPIOpts...) + a.filesAPI = agentfiles.NewAPI(a.logger.Named("files"), a.filesystem) + a.reconnectingPTYServer = reconnectingpty.NewServer( a.logger.Named("reconnecting-pty"), a.sshServer, diff --git a/agent/agentfiles/api.go b/agent/agentfiles/api.go new file mode 100644 index 0000000000..b4535bfb11 --- /dev/null +++ b/agent/agentfiles/api.go @@ -0,0 +1,36 @@ +package agentfiles + +import ( + "net/http" + + "github.com/go-chi/chi/v5" + "github.com/spf13/afero" + + "cdr.dev/slog/v3" +) + +// API exposes file-related operations performed through the agent. +type API struct { + logger slog.Logger + filesystem afero.Fs +} + +func NewAPI(logger slog.Logger, filesystem afero.Fs) *API { + api := &API{ + logger: logger, + filesystem: filesystem, + } + return api +} + +// Routes returns the HTTP handler for file-related routes. +func (api *API) Routes() http.Handler { + r := chi.NewRouter() + + r.Post("/list-directory", api.HandleLS) + r.Get("/read-file", api.HandleReadFile) + r.Post("/write-file", api.HandleWriteFile) + r.Post("/edit-files", api.HandleEditFiles) + + return r +} diff --git a/agent/files.go b/agent/agentfiles/files.go similarity index 83% rename from agent/files.go rename to agent/agentfiles/files.go index b4df9b42fa..86d073dfd1 100644 --- a/agent/files.go +++ b/agent/agentfiles/files.go @@ -1,4 +1,4 @@ -package agent +package agentfiles import ( "context" @@ -25,7 +25,7 @@ import ( type HTTPResponseCode = int -func (a *agent) HandleReadFile(rw http.ResponseWriter, r *http.Request) { +func (api *API) HandleReadFile(rw http.ResponseWriter, r *http.Request) { ctx := r.Context() query := r.URL.Query() @@ -42,7 +42,7 @@ func (a *agent) HandleReadFile(rw http.ResponseWriter, r *http.Request) { return } - status, err := a.streamFile(ctx, rw, path, offset, limit) + status, err := api.streamFile(ctx, rw, path, offset, limit) if err != nil { httpapi.Write(ctx, rw, status, codersdk.Response{ Message: err.Error(), @@ -51,12 +51,12 @@ func (a *agent) HandleReadFile(rw http.ResponseWriter, r *http.Request) { } } -func (a *agent) streamFile(ctx context.Context, rw http.ResponseWriter, path string, offset, limit int64) (HTTPResponseCode, error) { +func (api *API) streamFile(ctx context.Context, rw http.ResponseWriter, path string, offset, limit int64) (HTTPResponseCode, error) { if !filepath.IsAbs(path) { return http.StatusBadRequest, xerrors.Errorf("file path must be absolute: %q", path) } - f, err := a.filesystem.Open(path) + f, err := api.filesystem.Open(path) if err != nil { status := http.StatusInternalServerError switch { @@ -97,13 +97,13 @@ func (a *agent) streamFile(ctx context.Context, rw http.ResponseWriter, path str reader := io.NewSectionReader(f, offset, bytesToRead) _, err = io.Copy(rw, reader) if err != nil && !errors.Is(err, io.EOF) && ctx.Err() == nil { - a.logger.Error(ctx, "workspace agent read file", slog.Error(err)) + api.logger.Error(ctx, "workspace agent read file", slog.Error(err)) } return 0, nil } -func (a *agent) HandleWriteFile(rw http.ResponseWriter, r *http.Request) { +func (api *API) HandleWriteFile(rw http.ResponseWriter, r *http.Request) { ctx := r.Context() query := r.URL.Query() @@ -118,7 +118,7 @@ func (a *agent) HandleWriteFile(rw http.ResponseWriter, r *http.Request) { return } - status, err := a.writeFile(ctx, r, path) + status, err := api.writeFile(ctx, r, path) if err != nil { httpapi.Write(ctx, rw, status, codersdk.Response{ Message: err.Error(), @@ -131,13 +131,13 @@ func (a *agent) HandleWriteFile(rw http.ResponseWriter, r *http.Request) { }) } -func (a *agent) writeFile(ctx context.Context, r *http.Request, path string) (HTTPResponseCode, error) { +func (api *API) writeFile(ctx context.Context, r *http.Request, path string) (HTTPResponseCode, error) { if !filepath.IsAbs(path) { return http.StatusBadRequest, xerrors.Errorf("file path must be absolute: %q", path) } dir := filepath.Dir(path) - err := a.filesystem.MkdirAll(dir, 0o755) + err := api.filesystem.MkdirAll(dir, 0o755) if err != nil { status := http.StatusInternalServerError switch { @@ -149,7 +149,7 @@ func (a *agent) writeFile(ctx context.Context, r *http.Request, path string) (HT return status, err } - f, err := a.filesystem.Create(path) + f, err := api.filesystem.Create(path) if err != nil { status := http.StatusInternalServerError switch { @@ -164,13 +164,13 @@ func (a *agent) writeFile(ctx context.Context, r *http.Request, path string) (HT _, err = io.Copy(f, r.Body) if err != nil && !errors.Is(err, io.EOF) && ctx.Err() == nil { - a.logger.Error(ctx, "workspace agent write file", slog.Error(err)) + api.logger.Error(ctx, "workspace agent write file", slog.Error(err)) } return 0, nil } -func (a *agent) HandleEditFiles(rw http.ResponseWriter, r *http.Request) { +func (api *API) HandleEditFiles(rw http.ResponseWriter, r *http.Request) { ctx := r.Context() var req workspacesdk.FileEditRequest @@ -188,7 +188,7 @@ func (a *agent) HandleEditFiles(rw http.ResponseWriter, r *http.Request) { var combinedErr error status := http.StatusOK for _, edit := range req.Files { - s, err := a.editFile(r.Context(), edit.Path, edit.Edits) + s, err := api.editFile(r.Context(), edit.Path, edit.Edits) // Keep the highest response status, so 500 will be preferred over 400, etc. if s > status { status = s @@ -210,7 +210,7 @@ func (a *agent) HandleEditFiles(rw http.ResponseWriter, r *http.Request) { }) } -func (a *agent) editFile(ctx context.Context, path string, edits []workspacesdk.FileEdit) (int, error) { +func (api *API) editFile(ctx context.Context, path string, edits []workspacesdk.FileEdit) (int, error) { if path == "" { return http.StatusBadRequest, xerrors.New("\"path\" is required") } @@ -223,7 +223,7 @@ func (a *agent) editFile(ctx context.Context, path string, edits []workspacesdk. return http.StatusBadRequest, xerrors.New("must specify at least one edit") } - f, err := a.filesystem.Open(path) + f, err := api.filesystem.Open(path) if err != nil { status := http.StatusInternalServerError switch { @@ -252,7 +252,7 @@ func (a *agent) editFile(ctx context.Context, path string, edits []workspacesdk. // Create an adjacent file to ensure it will be on the same device and can be // moved atomically. - tmpfile, err := afero.TempFile(a.filesystem, filepath.Dir(path), filepath.Base(path)) + tmpfile, err := afero.TempFile(api.filesystem, filepath.Dir(path), filepath.Base(path)) if err != nil { return http.StatusInternalServerError, err } @@ -260,13 +260,13 @@ func (a *agent) editFile(ctx context.Context, path string, edits []workspacesdk. _, err = io.Copy(tmpfile, replace.Chain(f, transforms...)) if err != nil { - if rerr := a.filesystem.Remove(tmpfile.Name()); rerr != nil { - a.logger.Warn(ctx, "unable to clean up temp file", slog.Error(rerr)) + if rerr := api.filesystem.Remove(tmpfile.Name()); rerr != nil { + api.logger.Warn(ctx, "unable to clean up temp file", slog.Error(rerr)) } return http.StatusInternalServerError, xerrors.Errorf("edit %s: %w", path, err) } - err = a.filesystem.Rename(tmpfile.Name(), path) + err = api.filesystem.Rename(tmpfile.Name(), path) if err != nil { return http.StatusInternalServerError, err } diff --git a/agent/files_test.go b/agent/agentfiles/files_test.go similarity index 84% rename from agent/files_test.go rename to agent/agentfiles/files_test.go index 969c9b053b..0038795ad8 100644 --- a/agent/files_test.go +++ b/agent/agentfiles/files_test.go @@ -1,11 +1,13 @@ -package agent_test +package agentfiles_test import ( "bytes" "context" + "encoding/json" "fmt" "io" "net/http" + "net/http/httptest" "os" "path/filepath" "runtime" @@ -16,10 +18,10 @@ import ( "github.com/stretchr/testify/require" "golang.org/x/xerrors" - "github.com/coder/coder/v2/agent" - "github.com/coder/coder/v2/agent/agenttest" - "github.com/coder/coder/v2/coderd/coderdtest" - "github.com/coder/coder/v2/codersdk/agentsdk" + "cdr.dev/slog/v3" + "cdr.dev/slog/v3/sloggers/slogtest" + "github.com/coder/coder/v2/agent/agentfiles" + "github.com/coder/coder/v2/codersdk" "github.com/coder/coder/v2/codersdk/workspacesdk" "github.com/coder/coder/v2/testutil" ) @@ -106,15 +108,15 @@ func TestReadFile(t *testing.T) { tmpdir := os.TempDir() noPermsFilePath := filepath.Join(tmpdir, "no-perms") - //nolint:dogsled - conn, _, _, fs, _ := setupAgent(t, agentsdk.Manifest{}, 0, func(_ *agenttest.Client, opts *agent.Options) { - opts.Filesystem = newTestFs(opts.Filesystem, func(call, file string) error { - if file == noPermsFilePath { - return os.ErrPermission - } - return nil - }) + + logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}).Leveled(slog.LevelDebug) + fs := newTestFs(afero.NewMemMapFs(), func(call, file string) error { + if file == noPermsFilePath { + return os.ErrPermission + } + return nil }) + api := agentfiles.NewAPI(logger, fs) dirPath := filepath.Join(tmpdir, "a-directory") err := fs.MkdirAll(dirPath, 0o755) @@ -260,19 +262,22 @@ func TestReadFile(t *testing.T) { ctx, cancel := context.WithTimeout(context.Background(), testutil.WaitLong) defer cancel() - reader, mimeType, err := conn.ReadFile(ctx, tt.path, tt.offset, tt.limit) + w := httptest.NewRecorder() + r := httptest.NewRequestWithContext(ctx, http.MethodGet, fmt.Sprintf("/read-file?path=%s&offset=%d&limit=%d", tt.path, tt.offset, tt.limit), nil) + api.Routes().ServeHTTP(w, r) + if tt.errCode != 0 { - require.Error(t, err) - cerr := coderdtest.SDKError(t, err) - require.Contains(t, cerr.Error(), tt.error) - require.Equal(t, tt.errCode, cerr.StatusCode()) - } else { + got := &codersdk.Error{} + err := json.NewDecoder(w.Body).Decode(got) require.NoError(t, err) - defer reader.Close() - bytes, err := io.ReadAll(reader) + require.ErrorContains(t, got, tt.error) + require.Equal(t, tt.errCode, w.Code) + } else { + bytes, err := io.ReadAll(w.Body) require.NoError(t, err) require.Equal(t, tt.bytes, bytes) - require.Equal(t, tt.mimeType, mimeType) + require.Equal(t, tt.mimeType, w.Header().Get("Content-Type")) + require.Equal(t, http.StatusOK, w.Code) } }) } @@ -284,15 +289,14 @@ func TestWriteFile(t *testing.T) { tmpdir := os.TempDir() noPermsFilePath := filepath.Join(tmpdir, "no-perms-file") noPermsDirPath := filepath.Join(tmpdir, "no-perms-dir") - //nolint:dogsled - conn, _, _, fs, _ := setupAgent(t, agentsdk.Manifest{}, 0, func(_ *agenttest.Client, opts *agent.Options) { - opts.Filesystem = newTestFs(opts.Filesystem, func(call, file string) error { - if file == noPermsFilePath || file == noPermsDirPath { - return os.ErrPermission - } - return nil - }) + logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}).Leveled(slog.LevelDebug) + fs := newTestFs(afero.NewMemMapFs(), func(call, file string) error { + if file == noPermsFilePath || file == noPermsDirPath { + return os.ErrPermission + } + return nil }) + api := agentfiles.NewAPI(logger, fs) dirPath := filepath.Join(tmpdir, "directory") err := fs.MkdirAll(dirPath, 0o755) @@ -371,17 +375,21 @@ func TestWriteFile(t *testing.T) { defer cancel() reader := bytes.NewReader(tt.bytes) - err := conn.WriteFile(ctx, tt.path, reader) + w := httptest.NewRecorder() + r := httptest.NewRequestWithContext(ctx, http.MethodPost, fmt.Sprintf("/write-file?path=%s", tt.path), reader) + api.Routes().ServeHTTP(w, r) + if tt.errCode != 0 { - require.Error(t, err) - cerr := coderdtest.SDKError(t, err) - require.Contains(t, cerr.Error(), tt.error) - require.Equal(t, tt.errCode, cerr.StatusCode()) + got := &codersdk.Error{} + err := json.NewDecoder(w.Body).Decode(got) + require.NoError(t, err) + require.ErrorContains(t, got, tt.error) + require.Equal(t, tt.errCode, w.Code) } else { + bytes, err := afero.ReadFile(fs, tt.path) require.NoError(t, err) - b, err := afero.ReadFile(fs, tt.path) - require.NoError(t, err) - require.Equal(t, tt.bytes, b) + require.Equal(t, tt.bytes, bytes) + require.Equal(t, http.StatusOK, w.Code) } }) } @@ -393,21 +401,20 @@ func TestEditFiles(t *testing.T) { tmpdir := os.TempDir() noPermsFilePath := filepath.Join(tmpdir, "no-perms-file") failRenameFilePath := filepath.Join(tmpdir, "fail-rename") - //nolint:dogsled - conn, _, _, fs, _ := setupAgent(t, agentsdk.Manifest{}, 0, func(_ *agenttest.Client, opts *agent.Options) { - opts.Filesystem = newTestFs(opts.Filesystem, func(call, file string) error { - if file == noPermsFilePath { - return &os.PathError{ - Op: call, - Path: file, - Err: os.ErrPermission, - } - } else if file == failRenameFilePath && call == "rename" { - return xerrors.New("rename failed") + logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}).Leveled(slog.LevelDebug) + fs := newTestFs(afero.NewMemMapFs(), func(call, file string) error { + if file == noPermsFilePath { + return &os.PathError{ + Op: call, + Path: file, + Err: os.ErrPermission, } - return nil - }) + } else if file == failRenameFilePath && call == "rename" { + return xerrors.New("rename failed") + } + return nil }) + api := agentfiles.NewAPI(logger, fs) dirPath := filepath.Join(tmpdir, "directory") err := fs.MkdirAll(dirPath, 0o755) @@ -701,16 +708,26 @@ func TestEditFiles(t *testing.T) { require.NoError(t, err) } - err := conn.EditFiles(ctx, workspacesdk.FileEditRequest{Files: tt.edits}) + buf := bytes.NewBuffer(nil) + enc := json.NewEncoder(buf) + enc.SetEscapeHTML(false) + err := enc.Encode(workspacesdk.FileEditRequest{Files: tt.edits}) + require.NoError(t, err) + + w := httptest.NewRecorder() + r := httptest.NewRequestWithContext(ctx, http.MethodPost, "/edit-files", buf) + api.Routes().ServeHTTP(w, r) + if tt.errCode != 0 { - require.Error(t, err) - cerr := coderdtest.SDKError(t, err) - for _, error := range tt.errors { - require.Contains(t, cerr.Error(), error) - } - require.Equal(t, tt.errCode, cerr.StatusCode()) - } else { + got := &codersdk.Error{} + err := json.NewDecoder(w.Body).Decode(got) require.NoError(t, err) + for _, error := range tt.errors { + require.ErrorContains(t, got, error) + } + require.Equal(t, tt.errCode, w.Code) + } else { + require.Equal(t, http.StatusOK, w.Code) } for path, expect := range tt.expected { b, err := afero.ReadFile(fs, path) diff --git a/agent/ls.go b/agent/agentfiles/ls.go similarity index 97% rename from agent/ls.go rename to agent/agentfiles/ls.go index f2e2b27ea7..77f88cdd98 100644 --- a/agent/ls.go +++ b/agent/agentfiles/ls.go @@ -1,4 +1,4 @@ -package agent +package agentfiles import ( "errors" @@ -21,7 +21,7 @@ import ( var WindowsDriveRegex = regexp.MustCompile(`^[a-zA-Z]:\\$`) -func (a *agent) HandleLS(rw http.ResponseWriter, r *http.Request) { +func (api *API) HandleLS(rw http.ResponseWriter, r *http.Request) { ctx := r.Context() // An absolute path may be optionally provided, otherwise a path split into an @@ -43,7 +43,7 @@ func (a *agent) HandleLS(rw http.ResponseWriter, r *http.Request) { return } - resp, err := listFiles(a.filesystem, path, req) + resp, err := listFiles(api.filesystem, path, req) if err != nil { status := http.StatusInternalServerError switch { diff --git a/agent/ls_internal_test.go b/agent/agentfiles/ls_internal_test.go similarity index 99% rename from agent/ls_internal_test.go rename to agent/agentfiles/ls_internal_test.go index 18b959e5f8..a8a2a0cdb0 100644 --- a/agent/ls_internal_test.go +++ b/agent/agentfiles/ls_internal_test.go @@ -1,4 +1,4 @@ -package agent +package agentfiles import ( "os" diff --git a/agent/api.go b/agent/api.go index a631286c40..476eca181c 100644 --- a/agent/api.go +++ b/agent/api.go @@ -27,6 +27,8 @@ func (a *agent) apiHandler() http.Handler { }) }) + r.Mount("/api/v0", a.filesAPI.Routes()) + if a.devcontainers { r.Mount("/api/v0/containers", a.containerAPI.Routes()) } else if manifest := a.manifest.Load(); manifest != nil && manifest.ParentID != uuid.Nil { @@ -49,10 +51,6 @@ func (a *agent) apiHandler() http.Handler { r.Get("/api/v0/listening-ports", a.listeningPortsHandler.handler) r.Get("/api/v0/netcheck", a.HandleNetcheck) - r.Post("/api/v0/list-directory", a.HandleLS) - r.Get("/api/v0/read-file", a.HandleReadFile) - r.Post("/api/v0/write-file", a.HandleWriteFile) - r.Post("/api/v0/edit-files", a.HandleEditFiles) r.Get("/debug/logs", a.HandleHTTPDebugLogs) r.Get("/debug/magicsock", a.HandleHTTPDebugMagicsock) r.Get("/debug/magicsock/debug-logging/{state}", a.HandleHTTPMagicsockDebugLoggingState)