From ff9ed9181135c8fd1aa79e67be81808e11847829 Mon Sep 17 00:00:00 2001 From: Asher Date: Fri, 16 Jan 2026 18:03:17 -0800 Subject: [PATCH] chore: move agent's file API into separate package (#21531) This makes it so we can test it directly without having to go through Tailnet, which appears to be causing flakes in CI where the requests time out and never make it to the agent. Takes inspiration from the container-related API endpoints. Would probably make sense to refactor the ls tests to also go through the API (rather than be internal tests like they are currently) but I left those alone for now to keep the diff minimal. --- agent/agent.go | 5 + agent/agentfiles/api.go | 36 ++++++ agent/{ => agentfiles}/files.go | 40 +++--- agent/{ => agentfiles}/files_test.go | 135 ++++++++++++--------- agent/{ => agentfiles}/ls.go | 6 +- agent/{ => agentfiles}/ls_internal_test.go | 2 +- agent/api.go | 6 +- 7 files changed, 143 insertions(+), 87 deletions(-) create mode 100644 agent/agentfiles/api.go rename agent/{ => agentfiles}/files.go (83%) rename agent/{ => agentfiles}/files_test.go (84%) rename agent/{ => agentfiles}/ls.go (97%) rename agent/{ => agentfiles}/ls_internal_test.go (99%) 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)