Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 5 additions & 6 deletions packages/envd/internal/permissions/keepalive.go
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
package permissions

import (
"math"
"strconv"
"time"

Expand All @@ -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
}

Expand Down
94 changes: 94 additions & 0 deletions packages/envd/internal/permissions/keepalive_test.go
Original file line number Diff line number Diff line change
@@ -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()
}
})
}
2 changes: 1 addition & 1 deletion packages/envd/pkg/version.go
Original file line number Diff line number Diff line change
@@ -1,3 +1,3 @@
package pkg

var Version = "0.9.0" // x-release-please-version
var Version = "0.9.1" // x-release-please-version
134 changes: 134 additions & 0 deletions tests/integration/internal/tests/envd/keepalive_test.go
Original file line number Diff line number Diff line change
@@ -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")
})
}
}