diff --git a/packages/envd/internal/services/filesystem/handle_test.go b/packages/envd/internal/services/filesystem/handle_test.go new file mode 100644 index 0000000000..35ac4d5f4a --- /dev/null +++ b/packages/envd/internal/services/filesystem/handle_test.go @@ -0,0 +1,48 @@ +package filesystem + +import ( + "net/http/httptest" + "testing" + + "connectrpc.com/connect" + "github.com/go-chi/chi/v5" + "github.com/rs/zerolog" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/e2b-dev/infra/packages/envd/internal/execcontext" + rpc "github.com/e2b-dev/infra/packages/envd/internal/services/spec/filesystem" + spec "github.com/e2b-dev/infra/packages/envd/internal/services/spec/filesystem/filesystemconnect" + "github.com/e2b-dev/infra/packages/envd/internal/services/streaming" + "github.com/e2b-dev/infra/packages/envd/internal/utils" +) + +// TestHandleDisablesProxyBufferingForWatchDir checks the interceptor is wired +// into the service Handle mounts: WatchDir responses must tell reverse proxies +// not to buffer, while unary responses keep the default. +func TestHandleDisablesProxyBufferingForWatchDir(t *testing.T) { + t.Parallel() + + logger := zerolog.Nop() + mux := chi.NewRouter() + Handle(mux, &logger, &execcontext.Defaults{EnvVars: utils.NewEnvVars()}) + + srv := httptest.NewServer(mux) + t.Cleanup(srv.Close) + + client := spec.NewFilesystemClient(srv.Client(), srv.URL) + + // No user is configured, so the stream ends with an error; the header is + // set before the handler runs and must be present regardless. + stream, err := client.WatchDir(t.Context(), connect.NewRequest(&rpc.WatchDirRequest{Path: t.TempDir()})) + require.NoError(t, err) + t.Cleanup(func() { _ = stream.Close() }) + + assert.False(t, stream.Receive()) + assert.Equal(t, []string{"no"}, stream.ResponseHeader().Values(streaming.AccelBufferingHeader)) + + _, err = client.ListDir(t.Context(), connect.NewRequest(&rpc.ListDirRequest{Path: t.TempDir()})) + var connectErr *connect.Error + require.ErrorAs(t, err, &connectErr) + assert.Empty(t, connectErr.Meta().Values(streaming.AccelBufferingHeader)) +} diff --git a/packages/envd/internal/services/filesystem/service.go b/packages/envd/internal/services/filesystem/service.go index 6365f830bc..b29d1e1ef3 100644 --- a/packages/envd/internal/services/filesystem/service.go +++ b/packages/envd/internal/services/filesystem/service.go @@ -11,6 +11,7 @@ import ( "github.com/e2b-dev/infra/packages/envd/internal/logs" "github.com/e2b-dev/infra/packages/envd/internal/services/legacy" spec "github.com/e2b-dev/infra/packages/envd/internal/services/spec/filesystem/filesystemconnect" + "github.com/e2b-dev/infra/packages/envd/internal/services/streaming" "github.com/e2b-dev/infra/packages/envd/internal/utils" ) @@ -39,6 +40,7 @@ func Handle(server *chi.Mux, l *zerolog.Logger, defaults *execcontext.Defaults) interceptors := connect.WithInterceptors( logs.NewUnaryLogInterceptor(l), legacy.Convert(), + streaming.DisableProxyBuffering(), ) path, handler := spec.NewFilesystemHandler(service, interceptors) diff --git a/packages/envd/internal/services/process/handle_test.go b/packages/envd/internal/services/process/handle_test.go new file mode 100644 index 0000000000..70925683c8 --- /dev/null +++ b/packages/envd/internal/services/process/handle_test.go @@ -0,0 +1,55 @@ +package process + +import ( + "net/http/httptest" + "testing" + + "connectrpc.com/connect" + "github.com/go-chi/chi/v5" + "github.com/rs/zerolog" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/e2b-dev/infra/packages/envd/internal/execcontext" + "github.com/e2b-dev/infra/packages/envd/internal/services/cgroups" + rpc "github.com/e2b-dev/infra/packages/envd/internal/services/spec/process" + spec "github.com/e2b-dev/infra/packages/envd/internal/services/spec/process/processconnect" + "github.com/e2b-dev/infra/packages/envd/internal/services/streaming" + "github.com/e2b-dev/infra/packages/envd/internal/utils" +) + +// TestHandleDisablesProxyBufferingForConnect checks the interceptor is wired +// into the service Handle mounts: Connect responses must tell reverse proxies +// not to buffer, while unary responses keep the default. +func TestHandleDisablesProxyBufferingForConnect(t *testing.T) { + t.Parallel() + + cwd := t.TempDir() + logger := zerolog.Nop() + mux := chi.NewRouter() + Handle(mux, &logger, &execcontext.Defaults{ + EnvVars: utils.NewEnvVars(), + Workdir: &cwd, + }, cgroups.NewWorkloadFreezer(cgroups.NewNoopManager())) + + srv := httptest.NewServer(mux) + t.Cleanup(srv.Close) + + client := spec.NewProcessClient(srv.Client(), srv.URL) + + // No such process, so the stream ends with NotFound; the header is set + // before the handler runs and must be present regardless. + stream, err := client.Connect(t.Context(), connect.NewRequest(&rpc.ConnectRequest{ + Process: &rpc.ProcessSelector{Selector: &rpc.ProcessSelector_Pid{Pid: 1 << 30}}, + })) + require.NoError(t, err) + t.Cleanup(func() { _ = stream.Close() }) + + assert.False(t, stream.Receive()) + assert.Equal(t, connect.CodeNotFound, connect.CodeOf(stream.Err())) + assert.Equal(t, []string{"no"}, stream.ResponseHeader().Values(streaming.AccelBufferingHeader)) + + resp, err := client.List(t.Context(), connect.NewRequest(&rpc.ListRequest{})) + require.NoError(t, err) + assert.Empty(t, resp.Header().Values(streaming.AccelBufferingHeader)) +} diff --git a/packages/envd/internal/services/process/service.go b/packages/envd/internal/services/process/service.go index 889e484fd4..94207a029b 100644 --- a/packages/envd/internal/services/process/service.go +++ b/packages/envd/internal/services/process/service.go @@ -15,6 +15,7 @@ import ( "github.com/e2b-dev/infra/packages/envd/internal/services/process/handler" rpc "github.com/e2b-dev/infra/packages/envd/internal/services/spec/process" spec "github.com/e2b-dev/infra/packages/envd/internal/services/spec/process/processconnect" + "github.com/e2b-dev/infra/packages/envd/internal/services/streaming" "github.com/e2b-dev/infra/packages/envd/internal/utils" ) @@ -160,7 +161,10 @@ func (s *Service) clearTerminatedForTag(tag string) { func Handle(server *chi.Mux, l *zerolog.Logger, defaults *execcontext.Defaults, workloadFreezer *cgroups.WorkloadFreezer) *Service { service := newService(l, defaults, workloadFreezer) - interceptors := connect.WithInterceptors(logs.NewUnaryLogInterceptor(l)) + interceptors := connect.WithInterceptors( + logs.NewUnaryLogInterceptor(l), + streaming.DisableProxyBuffering(), + ) path, h := spec.NewProcessHandler(service, interceptors) diff --git a/packages/envd/internal/services/streaming/interceptor.go b/packages/envd/internal/services/streaming/interceptor.go new file mode 100644 index 0000000000..c83e070c5e --- /dev/null +++ b/packages/envd/internal/services/streaming/interceptor.go @@ -0,0 +1,48 @@ +// Package streaming holds connect interceptors that apply to envd's +// streaming RPCs. +package streaming + +import ( + "context" + + "connectrpc.com/connect" +) + +// AccelBufferingHeader is the response header reverse proxies such as nginx +// read on a per-response basis to decide whether to buffer the body. Setting +// it to "no" makes the proxy forward each chunk as soon as it is written. +const AccelBufferingHeader = "X-Accel-Buffering" + +// DisableProxyBuffering returns an interceptor that marks every response of a +// server-streaming (or bidi) RPC with "X-Accel-Buffering: no". +// +// Without it, a proxy with response buffering enabled (the nginx default) +// holds stream messages such as process output or watch events until its +// buffer fills, so an interactive client sees nothing until something forces a +// flush. Unary and client-streaming RPCs return a single message and are left +// untouched, so they keep the proxy's normal buffering behavior. +func DisableProxyBuffering() NoProxyBufferingInterceptor { + return NoProxyBufferingInterceptor{} +} + +type NoProxyBufferingInterceptor struct{} + +var _ connect.Interceptor = NoProxyBufferingInterceptor{} + +func (NoProxyBufferingInterceptor) WrapUnary(next connect.UnaryFunc) connect.UnaryFunc { + return next +} + +func (NoProxyBufferingInterceptor) WrapStreamingClient(next connect.StreamingClientFunc) connect.StreamingClientFunc { + return next +} + +func (NoProxyBufferingInterceptor) WrapStreamingHandler(next connect.StreamingHandlerFunc) connect.StreamingHandlerFunc { + return func(ctx context.Context, conn connect.StreamingHandlerConn) error { + if conn.Spec().StreamType&connect.StreamTypeServer != 0 { + conn.ResponseHeader().Set(AccelBufferingHeader, "no") + } + + return next(ctx, conn) + } +} diff --git a/packages/envd/internal/services/streaming/interceptor_test.go b/packages/envd/internal/services/streaming/interceptor_test.go new file mode 100644 index 0000000000..d331c7d2e4 --- /dev/null +++ b/packages/envd/internal/services/streaming/interceptor_test.go @@ -0,0 +1,136 @@ +package streaming + +import ( + "context" + "net/http/httptest" + "testing" + + "connectrpc.com/connect" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/mock" + "github.com/stretchr/testify/require" + + "github.com/e2b-dev/infra/packages/envd/internal/services/spec/filesystem" + "github.com/e2b-dev/infra/packages/envd/internal/services/spec/filesystem/filesystemconnect" + filesystemconnectmocks "github.com/e2b-dev/infra/packages/envd/internal/services/spec/filesystem/filesystemconnect/mocks" + "github.com/e2b-dev/infra/packages/envd/internal/services/spec/process" + "github.com/e2b-dev/infra/packages/envd/internal/services/spec/process/processconnect" +) + +func newFilesystemClient(t *testing.T, mockFS *filesystemconnectmocks.MockFilesystemHandler) filesystemconnect.FilesystemClient { + t.Helper() + + _, handler := filesystemconnect.NewFilesystemHandler(mockFS, connect.WithInterceptors(DisableProxyBuffering())) + srv := httptest.NewServer(handler) + t.Cleanup(srv.Close) + + return filesystemconnect.NewFilesystemClient(srv.Client(), srv.URL) +} + +func TestServerStreamDisablesProxyBuffering(t *testing.T) { + t.Parallel() + + mockFS := filesystemconnectmocks.NewMockFilesystemHandler(t) + mockFS.EXPECT(). + WatchDir(mock.Anything, mock.Anything, mock.Anything). + RunAndReturn(func(_ context.Context, _ *connect.Request[filesystem.WatchDirRequest], stream *connect.ServerStream[filesystem.WatchDirResponse]) error { + return stream.Send(&filesystem.WatchDirResponse{Event: &filesystem.WatchDirResponse_Start{Start: &filesystem.WatchDirResponse_StartEvent{}}}) + }) + + client := newFilesystemClient(t, mockFS) + + stream, err := client.WatchDir(t.Context(), connect.NewRequest(&filesystem.WatchDirRequest{Path: "/a"})) + require.NoError(t, err) + t.Cleanup(func() { _ = stream.Close() }) + + require.True(t, stream.Receive(), stream.Err()) + assert.Equal(t, []string{"no"}, stream.ResponseHeader().Values(AccelBufferingHeader)) +} + +func TestServerStreamErrorStillDisablesProxyBuffering(t *testing.T) { + t.Parallel() + + mockFS := filesystemconnectmocks.NewMockFilesystemHandler(t) + mockFS.EXPECT(). + WatchDir(mock.Anything, mock.Anything, mock.Anything). + Return(connect.NewError(connect.CodeNotFound, assert.AnError)) + + client := newFilesystemClient(t, mockFS) + + stream, err := client.WatchDir(t.Context(), connect.NewRequest(&filesystem.WatchDirRequest{Path: "/missing"})) + require.NoError(t, err) + t.Cleanup(func() { _ = stream.Close() }) + + assert.False(t, stream.Receive()) + assert.Equal(t, connect.CodeNotFound, connect.CodeOf(stream.Err())) + assert.Equal(t, []string{"no"}, stream.ResponseHeader().Values(AccelBufferingHeader)) +} + +func TestUnaryLeavesProxyBufferingAlone(t *testing.T) { + t.Parallel() + + mockFS := filesystemconnectmocks.NewMockFilesystemHandler(t) + mockFS.EXPECT(). + ListDir(mock.Anything, mock.Anything). + Return(connect.NewResponse(&filesystem.ListDirResponse{}), nil) + + client := newFilesystemClient(t, mockFS) + + resp, err := client.ListDir(t.Context(), connect.NewRequest(&filesystem.ListDirRequest{Path: "/a"})) + require.NoError(t, err) + assert.Empty(t, resp.Header().Values(AccelBufferingHeader)) +} + +// fakeProcess implements only the process RPCs these tests call. +type fakeProcess struct { + processconnect.UnimplementedProcessHandler +} + +func (fakeProcess) Start(_ context.Context, _ *connect.Request[process.StartRequest], stream *connect.ServerStream[process.StartResponse]) error { + return stream.Send(&process.StartResponse{Event: &process.ProcessEvent{}}) +} + +func (fakeProcess) StreamInput(_ context.Context, stream *connect.ClientStream[process.StreamInputRequest]) (*connect.Response[process.StreamInputResponse], error) { + for stream.Receive() { + } + + return connect.NewResponse(&process.StreamInputResponse{}), stream.Err() +} + +func newProcessClient(t *testing.T) processconnect.ProcessClient { + t.Helper() + + _, handler := processconnect.NewProcessHandler(fakeProcess{}, connect.WithInterceptors(DisableProxyBuffering())) + srv := httptest.NewUnstartedServer(handler) + srv.EnableHTTP2 = true + srv.StartTLS() + t.Cleanup(srv.Close) + + return processconnect.NewProcessClient(srv.Client(), srv.URL) +} + +func TestProcessStartDisablesProxyBuffering(t *testing.T) { + t.Parallel() + + client := newProcessClient(t) + + stream, err := client.Start(t.Context(), connect.NewRequest(&process.StartRequest{})) + require.NoError(t, err) + t.Cleanup(func() { _ = stream.Close() }) + + require.True(t, stream.Receive(), stream.Err()) + assert.Equal(t, []string{"no"}, stream.ResponseHeader().Values(AccelBufferingHeader)) +} + +func TestClientStreamLeavesProxyBufferingAlone(t *testing.T) { + t.Parallel() + + client := newProcessClient(t) + + stream := client.StreamInput(t.Context()) + require.NoError(t, stream.Send(&process.StreamInputRequest{})) + + resp, err := stream.CloseAndReceive() + require.NoError(t, err) + assert.Empty(t, resp.Header().Values(AccelBufferingHeader)) +}