Skip to content
6 changes: 6 additions & 0 deletions packages/envd/internal/services/filesystem/service_test.go
Original file line number Diff line number Diff line change
@@ -1,12 +1,18 @@
package filesystem

import (
"github.com/rs/zerolog"

"github.com/e2b-dev/infra/packages/envd/internal/execcontext"
"github.com/e2b-dev/infra/packages/envd/internal/utils"
)

func mockService() Service {
logger := zerolog.Nop()

return Service{
logger: &logger,
watchers: utils.NewMap[string, *FileWatcher](),
defaults: &execcontext.Defaults{
EnvVars: utils.NewEnvVars(),
},
Expand Down
37 changes: 37 additions & 0 deletions packages/envd/internal/services/filesystem/utils.go
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@ import (
"time"

"connectrpc.com/connect"
"github.com/rs/zerolog"
"google.golang.org/protobuf/types/known/timestamppb"

rpc "github.com/e2b-dev/infra/packages/envd/internal/services/spec/filesystem"
Expand Down Expand Up @@ -66,6 +67,42 @@ func entryInfo(path string) (*rpc.EntryInfo, error) {
}, nil
}

// opCarriesEntry reports whether a filesystem event of the given type refers to an entry
// that is expected to still exist at the path, and may therefore carry EntryInfo. Remove
// and rename events refer to a path whose original entry is gone, so they must never carry
// entry info: stat-ing the path could otherwise attach a replacement entry that was created
// at the same path before the event was handled.
func opCarriesEntry(op rpc.EventType) bool {
switch op {
case rpc.EventType_EVENT_TYPE_CREATE,
rpc.EventType_EVENT_TYPE_WRITE,
rpc.EventType_EVENT_TYPE_CHMOD:
return true
default:
return false
}
}

// eventEntryInfo returns the EntryInfo for the path that triggered a filesystem event.
// It must only be called for events that carry an entry (see opCarriesEntry).
//
// Entry info is best-effort: a nil entry is returned (and the watch keeps running) when it
// cannot be retrieved. A NotFound result is treated as a benign race (the entry was removed
// between the event and the stat) and is not logged; any other failure is logged at warn level.
func eventEntryInfo(logger *zerolog.Logger, path string) *rpc.EntryInfo {
entry, err := entryInfo(path)
Comment thread
mishushakov marked this conversation as resolved.
if err != nil {
// NotFound is a benign race: the entry was removed before we could stat it.
if connect.CodeOf(err) != connect.CodeNotFound {
logger.Warn().Err(err).Str("path", path).Msg("failed to get entry info for filesystem event")
}

return nil
}

return entry
}

func toTimestamp(time time.Time) *timestamppb.Timestamp {
if time.IsZero() {
return nil
Expand Down
16 changes: 10 additions & 6 deletions packages/envd/internal/services/filesystem/watch.go
Original file line number Diff line number Diff line change
Expand Up @@ -129,15 +129,19 @@ func (s Service) watchHandler(ctx context.Context, req *connect.Request[rpc.Watc
return connect.NewError(connect.CodeInternal, fmt.Errorf("error getting relative path: %w", nameErr))
}

filesystemEvent := &rpc.WatchDirResponse_Filesystem{
Filesystem: &rpc.FilesystemEvent{
Name: name,
Type: op,
},
filesystemEvent := &rpc.FilesystemEvent{
Name: name,
Type: op,
}

if req.Msg.GetIncludeEntry() && opCarriesEntry(op) {
filesystemEvent.Entry = eventEntryInfo(s.logger, e.Name)
}

event := &rpc.WatchDirResponse{
Event: filesystemEvent,
Event: &rpc.WatchDirResponse_Filesystem{
Filesystem: filesystemEvent,
},
}

if streamErr := stream.Send(event); streamErr != nil {
Expand Down
39 changes: 28 additions & 11 deletions packages/envd/internal/services/filesystem/watch_sync.go
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@ import (

"connectrpc.com/connect"
"github.com/e2b-dev/fsnotify"
"github.com/rs/zerolog"

"github.com/e2b-dev/infra/packages/envd/internal/permissions"
rpc "github.com/e2b-dev/infra/packages/envd/internal/services/spec/filesystem"
Expand All @@ -26,7 +27,7 @@ type FileWatcher struct {
Lock sync.Mutex
}

func CreateFileWatcher(ctx context.Context, watchPath string, recursive bool) (*FileWatcher, error) {
func CreateFileWatcher(ctx context.Context, logger *zerolog.Logger, watchPath string, recursive bool, includeEntryInfo bool) (*FileWatcher, error) {
w, err := fsnotify.NewWatcher()
if err != nil {
return nil, connect.NewError(connect.CodeInternal, fmt.Errorf("error creating watcher: %w", err))
Expand Down Expand Up @@ -56,17 +57,17 @@ func CreateFileWatcher(ctx context.Context, watchPath string, recursive bool) (*
return
case chErr, ok := <-w.Errors:
if !ok {
fw.Error = connect.NewError(connect.CodeInternal, errors.New("watcher error channel closed"))
fw.setError(connect.NewError(connect.CodeInternal, errors.New("watcher error channel closed")))

return
}

fw.Error = connect.NewError(connect.CodeInternal, fmt.Errorf("watcher error: %w", chErr))
fw.setError(connect.NewError(connect.CodeInternal, fmt.Errorf("watcher error: %w", chErr)))

return
case e, ok := <-w.Events:
if !ok {
fw.Error = connect.NewError(connect.CodeInternal, errors.New("watcher event channel closed"))
fw.setError(connect.NewError(connect.CodeInternal, errors.New("watcher event channel closed")))

return
}
Expand Down Expand Up @@ -97,16 +98,22 @@ func CreateFileWatcher(ctx context.Context, watchPath string, recursive bool) (*
for _, op := range ops {
name, nameErr := filepath.Rel(watchPath, e.Name)
if nameErr != nil {
fw.Error = connect.NewError(connect.CodeInternal, fmt.Errorf("error getting relative path: %w", nameErr))
fw.setError(connect.NewError(connect.CodeInternal, fmt.Errorf("error getting relative path: %w", nameErr)))

return
}

fw.Lock.Lock()
fw.Events = append(fw.Events, &rpc.FilesystemEvent{
filesystemEvent := &rpc.FilesystemEvent{
Name: name,
Type: op,
})
}

if includeEntryInfo && opCarriesEntry(op) {
filesystemEvent.Entry = eventEntryInfo(logger, e.Name)
}

fw.Lock.Lock()
fw.Events = append(fw.Events, filesystemEvent)
fw.Lock.Unlock()
}
}
Expand All @@ -121,6 +128,15 @@ func (fw *FileWatcher) Close() {
fw.cancel()
}

// setError records a terminal watcher error. Guarded by Lock so it is safe against
// concurrent GetWatcherEvents reads.
func (fw *FileWatcher) setError(err error) {
fw.Lock.Lock()
defer fw.Lock.Unlock()

fw.Error = err
}

func (s Service) CreateWatcher(ctx context.Context, req *connect.Request[rpc.CreateWatcherRequest]) (*connect.Response[rpc.CreateWatcherResponse], error) {
u, err := permissions.GetAuthUser(ctx, s.defaults.User)
if err != nil {
Expand Down Expand Up @@ -156,7 +172,7 @@ func (s Service) CreateWatcher(ctx context.Context, req *connect.Request[rpc.Cre

watcherId := "w" + id.Generate()

w, err := CreateFileWatcher(ctx, watchPath, req.Msg.GetRecursive())
w, err := CreateFileWatcher(ctx, s.logger, watchPath, req.Msg.GetRecursive(), req.Msg.GetIncludeEntry())
if err != nil {
return nil, err
}
Expand All @@ -176,12 +192,13 @@ func (s Service) GetWatcherEvents(_ context.Context, req *connect.Request[rpc.Ge
return nil, connect.NewError(connect.CodeNotFound, fmt.Errorf("watcher with id %s not found", watcherId))
}

w.Lock.Lock()
defer w.Lock.Unlock()

if w.Error != nil {
return nil, w.Error
}

w.Lock.Lock()
defer w.Lock.Unlock()
events := w.Events
w.Events = []*rpc.FilesystemEvent{}

Expand Down
166 changes: 166 additions & 0 deletions packages/envd/internal/services/filesystem/watch_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,166 @@
package filesystem

import (
"context"
"os"
"os/user"
"path/filepath"
"testing"
"time"

"connectrpc.com/authn"
"connectrpc.com/connect"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"

"github.com/e2b-dev/infra/packages/envd/internal/services/spec/filesystem"
)

// collectEvents polls GetWatcherEvents until at least one event is returned or the
// deadline is reached. fsnotify delivers events asynchronously, so we can't assume
// they are available immediately after the filesystem operation.
func collectEvents(t *testing.T, ctx context.Context, svc Service, watcherID string) []*filesystem.FilesystemEvent {
t.Helper()

deadline := time.Now().Add(2 * time.Second)
for time.Now().Before(deadline) {
resp, err := svc.GetWatcherEvents(ctx, connect.NewRequest(&filesystem.GetWatcherEventsRequest{
WatcherId: watcherID,
}))
require.NoError(t, err)

if len(resp.Msg.GetEvents()) > 0 {
return resp.Msg.GetEvents()
}

time.Sleep(20 * time.Millisecond)
}

return nil
}

func TestWatcherIncludeEntryInfo(t *testing.T) {
t.Parallel()

u, err := user.Current()
require.NoError(t, err)

tests := []struct {
name string
includeEntryInfo bool
wantEntry bool
}{
{name: "entry info included when requested", includeEntryInfo: true, wantEntry: true},
{name: "entry info omitted when not requested", includeEntryInfo: false, wantEntry: false},
}

for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
t.Parallel()

root := t.TempDir()
svc := mockService()
ctx := authn.SetInfo(t.Context(), u)

created, err := svc.CreateWatcher(ctx, connect.NewRequest(&filesystem.CreateWatcherRequest{
Path: root,
IncludeEntry: tt.includeEntryInfo,
}))
require.NoError(t, err)
watcherID := created.Msg.GetWatcherId()
t.Cleanup(func() {
_, _ = svc.RemoveWatcher(ctx, connect.NewRequest(&filesystem.RemoveWatcherRequest{
WatcherId: watcherID,
}))
})

// Trigger an event that leaves a stat-able entry behind.
filePath := filepath.Join(root, "file.txt")
require.NoError(t, os.WriteFile(filePath, []byte("hello"), 0o644))

events := collectEvents(t, ctx, svc, watcherID)
require.NotEmpty(t, events, "expected at least one filesystem event")

for _, e := range events {
assert.Equal(t, "file.txt", e.GetName())

if tt.wantEntry {
require.NotNil(t, e.GetEntry(), "expected entry info on event %s", e.GetType())
assert.Equal(t, "file.txt", e.GetEntry().GetName())
assert.Equal(t, filePath, e.GetEntry().GetPath())
assert.Equal(t, filesystem.FileType_FILE_TYPE_FILE, e.GetEntry().GetType())
} else {
assert.Nil(t, e.GetEntry(), "expected no entry info on event %s", e.GetType())
}
}
})
}
}

// TestWatcherIncludeEntryInfo_RemoveDoesNotCarryReplacement guards against a TOCTOU race:
// if an entry is removed and a different entry is created at the same path before the
// watcher handles the remove event, stat-ing the path would succeed and could attach the
// replacement's info to the remove event. Remove/rename events must never carry entry info.
func TestWatcherIncludeEntryInfo_RemoveDoesNotCarryReplacement(t *testing.T) {
t.Parallel()

u, err := user.Current()
require.NoError(t, err)

root := t.TempDir()

// File exists before we start watching, so its removal is what we observe.
filePath := filepath.Join(root, "file.txt")
require.NoError(t, os.WriteFile(filePath, []byte("hello"), 0o644))

svc := mockService()
ctx := authn.SetInfo(t.Context(), u)

created, err := svc.CreateWatcher(ctx, connect.NewRequest(&filesystem.CreateWatcherRequest{
Path: root,
IncludeEntry: true,
}))
require.NoError(t, err)
watcherID := created.Msg.GetWatcherId()
t.Cleanup(func() {
_, _ = svc.RemoveWatcher(ctx, connect.NewRequest(&filesystem.RemoveWatcherRequest{
WatcherId: watcherID,
}))
})

require.NoError(t, os.Remove(filePath))
// Recreate a different entry at the same path before the watcher handles the remove
// event, so the path is occupied (and stat-able) by the time the event is processed.
require.NoError(t, os.WriteFile(filePath, []byte("replacement"), 0o644))

// Accumulate events until we have observed both the removal and the replacement, or time out.
var removeEvent, replacementEvent *filesystem.FilesystemEvent
deadline := time.Now().Add(2 * time.Second)
for time.Now().Before(deadline) && (removeEvent == nil || replacementEvent == nil) {
resp, eventsErr := svc.GetWatcherEvents(ctx, connect.NewRequest(&filesystem.GetWatcherEventsRequest{
WatcherId: watcherID,
}))
require.NoError(t, eventsErr)

for _, e := range resp.Msg.GetEvents() {
switch e.GetType() {
case filesystem.EventType_EVENT_TYPE_REMOVE:
removeEvent = e
case filesystem.EventType_EVENT_TYPE_CREATE, filesystem.EventType_EVENT_TYPE_WRITE:
if replacementEvent == nil {
replacementEvent = e
}
}
}

time.Sleep(20 * time.Millisecond)
}

require.NotNil(t, removeEvent, "expected a remove event")
assert.Nil(t, removeEvent.GetEntry(), "remove event must not carry entry info even when a new entry occupies the path")

// Sanity check: the replacement is stat-able, so the nil above is by design — not just
// because the path happens to be empty.
require.NotNil(t, replacementEvent, "expected a create/write event for the replacement")
assert.NotNil(t, replacementEvent.GetEntry(), "event for the existing replacement should carry entry info")
}
Loading
Loading