diff --git a/packages/envd/internal/permissions/keepalive.go b/packages/envd/internal/permissions/keepalive.go index 3d7a0461af..caf0e61cc6 100644 --- a/packages/envd/internal/permissions/keepalive.go +++ b/packages/envd/internal/permissions/keepalive.go @@ -1,6 +1,7 @@ package permissions import ( + "math" "strconv" "time" @@ -12,12 +13,10 @@ const defaultKeepAliveInterval = 90 * time.Second func GetKeepAliveTicker[T any](req *connect.Request[T]) (*time.Ticker, func()) { keepAliveIntervalHeader := req.Header().Get("Keepalive-Ping-Interval") - var interval time.Duration - - keepAliveIntervalInt, err := strconv.Atoi(keepAliveIntervalHeader) - if err != nil { - interval = defaultKeepAliveInterval - } else { + interval := defaultKeepAliveInterval + keepAliveIntervalInt, err := strconv.ParseInt(keepAliveIntervalHeader, 10, 64) + // Validate seconds before multiplication, which could overflow time.Duration. + if err == nil && keepAliveIntervalInt > 0 && keepAliveIntervalInt <= math.MaxInt64/int64(time.Second) { interval = time.Duration(keepAliveIntervalInt) * time.Second } diff --git a/packages/envd/internal/permissions/keepalive_test.go b/packages/envd/internal/permissions/keepalive_test.go new file mode 100644 index 0000000000..c753a407ac --- /dev/null +++ b/packages/envd/internal/permissions/keepalive_test.go @@ -0,0 +1,94 @@ +package permissions_test + +import ( + "testing" + "testing/synctest" + "time" + + "connectrpc.com/connect" + "github.com/stretchr/testify/require" + + "github.com/e2b-dev/infra/packages/envd/internal/permissions" +) + +func TestGetKeepAliveTicker(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + header string + interval time.Duration + }{ + {name: "missing", interval: 90 * time.Second}, + {name: "non-numeric", header: "invalid", interval: 90 * time.Second}, + {name: "zero", header: "0", interval: 90 * time.Second}, + {name: "negative", header: "-1", interval: 90 * time.Second}, + {name: "duration overflow", header: "9223372037", interval: 90 * time.Second}, + {name: "positive duration overflow", header: "18446744074", interval: 90 * time.Second}, + {name: "integer overflow", header: "9223372036854775808", interval: 90 * time.Second}, + {name: "one second", header: "1", interval: time.Second}, + {name: "SDK interval", header: "50", interval: 50 * time.Second}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + synctest.Test(t, func(t *testing.T) { + req := connect.NewRequest(&struct{}{}) + if tt.header != "" { + req.Header().Set("Keepalive-Ping-Interval", tt.header) + } + + var ticker *time.Ticker + var reset func() + require.NotPanics(t, func() { + ticker, reset = permissions.GetKeepAliveTicker(req) + }) + defer ticker.Stop() + + started := time.Now() + time.Sleep(tt.interval) + select { + case tick := <-ticker.C: + require.Equal(t, tt.interval, tick.Sub(started)) + default: + t.Fatal("keepalive did not tick at the expected interval") + } + + time.Sleep(tt.interval / 2) + reset() + resetAt := time.Now() + time.Sleep(tt.interval) + select { + case tick := <-ticker.C: + require.Equal(t, tt.interval, tick.Sub(resetAt), "reset should restart the same interval") + default: + t.Fatal("keepalive did not tick after reset") + } + }) + }) + } +} + +func TestGetKeepAliveTicker_MaximumInterval(t *testing.T) { + t.Parallel() + + synctest.Test(t, func(t *testing.T) { + req := connect.NewRequest(&struct{}{}) + req.Header().Set("Keepalive-Ping-Interval", "9223372036") + ticker, reset := permissions.GetKeepAliveTicker(req) + defer ticker.Stop() + + // The largest representable whole-second interval must not fall back to 90s. + // Observe a short window to avoid overflowing the virtual monotonic clock. + for range 2 { + time.Sleep(90 * time.Second) + select { + case <-ticker.C: + t.Fatal("maximum valid interval fell back to the default") + default: + } + reset() + } + }) +} diff --git a/packages/envd/pkg/version.go b/packages/envd/pkg/version.go index 271c8e0b50..3e943ccbdc 100644 --- a/packages/envd/pkg/version.go +++ b/packages/envd/pkg/version.go @@ -1,3 +1,3 @@ package pkg -var Version = "0.9.0" // x-release-please-version +var Version = "0.9.1" // x-release-please-version diff --git a/tests/integration/internal/tests/envd/keepalive_test.go b/tests/integration/internal/tests/envd/keepalive_test.go new file mode 100644 index 0000000000..711b24b3c7 --- /dev/null +++ b/tests/integration/internal/tests/envd/keepalive_test.go @@ -0,0 +1,134 @@ +package envd + +import ( + "context" + "strings" + "testing" + "time" + + "connectrpc.com/connect" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/e2b-dev/infra/packages/shared/pkg/grpc/envd/process" + "github.com/e2b-dev/infra/tests/integration/internal/setup" + "github.com/e2b-dev/infra/tests/integration/internal/utils" +) + +func TestCommandKeepaliveInterval(t *testing.T) { + t.Parallel() + + for _, tc := range []struct { + name string + interval string + }{ + {name: "valid", interval: "50"}, + {name: "zero", interval: "0"}, + } { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + // Each case owns its sandbox so an envd crash cannot affect other tests. + sbx := utils.SetupSandboxWithCleanup(t, setup.GetAPIClient(), utils.WithTimeout(60)) + ctx, cancel := context.WithTimeout(t.Context(), 15*time.Second) + defer cancel() + + client := setup.GetEnvdClient(t, ctx) + backgroundReq := connect.NewRequest(&process.StartRequest{ + Process: &process.ProcessConfig{ + Cmd: "/bin/sleep", + Args: []string{"60"}, + }, + }) + setup.SetSandboxHeader(t, backgroundReq.Header(), sbx.SandboxID) + setup.SetUserHeader(t, backgroundReq.Header(), "user") + background, err := client.ProcessClient.Start(ctx, backgroundReq) + require.NoError(t, err) + defer background.Close() + + require.True(t, background.Receive(), "background command must start: %v", background.Err()) + start := background.Msg().GetEvent().GetStart() + require.NotNil(t, start) + pid := start.GetPid() + + req := connect.NewRequest(&process.StartRequest{ + Process: &process.ProcessConfig{ + Cmd: "/bin/echo", + Args: []string{"ok"}, + }, + }) + setup.SetSandboxHeader(t, req.Header(), sbx.SandboxID) + setup.SetUserHeader(t, req.Header(), "user") + // A zero interval is invalid input. Like a non-numeric interval, + // it should fall back to the default without interrupting the command. + req.Header().Set("Keepalive-Ping-Interval", tc.interval) + + stream, err := client.ProcessClient.Start(ctx, req) + require.NoError(t, err) + defer stream.Close() + + var stdout strings.Builder + var end *process.ProcessEvent_EndEvent + for stream.Receive() { + event := stream.Msg().GetEvent() + stdout.Write(event.GetData().GetStdout()) + if event.GetEnd() != nil { + end = event.GetEnd() + } + } + + require.NoError(t, stream.Err(), "command stream must finish normally") + require.NotNil(t, end, "command must report its exit status") + assert.EqualValues(t, 0, end.GetExitCode()) + assert.Equal(t, "ok\n", stdout.String()) + + listReq := connect.NewRequest(&process.ListRequest{}) + setup.SetSandboxHeader(t, listReq.Header(), sbx.SandboxID) + setup.SetUserHeader(t, listReq.Header(), "user") + listed, err := client.ProcessClient.List(ctx, listReq) + require.NoError(t, err) + pids := make([]uint32, 0, len(listed.Msg.GetProcesses())) + for _, proc := range listed.Msg.GetProcesses() { + pids = append(pids, proc.GetPid()) + } + require.Contains(t, pids, pid, "envd must retain the existing command") + + selector := &process.ProcessSelector{Selector: &process.ProcessSelector_Pid{Pid: pid}} + connectReq := connect.NewRequest(&process.ConnectRequest{Process: selector}) + setup.SetSandboxHeader(t, connectReq.Header(), sbx.SandboxID) + setup.SetUserHeader(t, connectReq.Header(), "user") + connected, err := client.ProcessClient.Connect(ctx, connectReq) + require.NoError(t, err) + defer connected.Close() + require.True(t, connected.Receive(), "existing command must remain connectable: %v", connected.Err()) + require.Equal(t, pid, connected.Msg().GetEvent().GetStart().GetPid()) + + killReq := connect.NewRequest(&process.SendSignalRequest{ + Process: selector, + Signal: process.Signal_SIGNAL_SIGTERM, + }) + setup.SetSandboxHeader(t, killReq.Header(), sbx.SandboxID) + setup.SetUserHeader(t, killReq.Header(), "user") + _, err = client.ProcessClient.SendSignal(ctx, killReq) + require.NoError(t, err) + + var backgroundEnd *process.ProcessEvent_EndEvent + for background.Receive() { + if event := background.Msg().GetEvent().GetEnd(); event != nil { + backgroundEnd = event + } + } + require.NoError(t, background.Err(), "the original stream must survive the invalid interval") + require.NotNil(t, backgroundEnd, "the original stream must report termination") + + var connectedEnd *process.ProcessEvent_EndEvent + for connected.Receive() { + if event := connected.Msg().GetEvent().GetEnd(); event != nil { + connectedEnd = event + } + } + require.NoError(t, connected.Err()) + require.NotNil(t, connectedEnd, "the reconnected stream must report termination") + }) + } +}