diff --git a/.github/actions/start-services/action.yml b/.github/actions/start-services/action.yml index a2f5be1288..08ac484931 100644 --- a/.github/actions/start-services/action.yml +++ b/.github/actions/start-services/action.yml @@ -128,7 +128,6 @@ runs: # Client-proxy config API_INTERNAL_GRPC_ADDRESS: "localhost:5009" DEFAULT_PERSISTENT_VOLUME_TYPE: "test-volume-type" - SANDBOX_STORAGE_BACKEND: "redis" COMPRESS_ENABLED: ${{ inputs.compress_enabled }} COMPRESS_TYPE: ${{ inputs.compress_type }} COMPRESS_LEVEL: ${{ inputs.compress_level }} diff --git a/iac/modules/job-api/jobs/api.hcl b/iac/modules/job-api/jobs/api.hcl index 93fa309012..7aa55bfe2e 100644 --- a/iac/modules/job-api/jobs/api.hcl +++ b/iac/modules/job-api/jobs/api.hcl @@ -185,8 +185,6 @@ job "api" { REDIS_TLS_CA_BASE64 = "${redis_tls_ca_base64}" REDIS_URL = "${redis_url}" - SANDBOX_STORAGE_BACKEND = "${sandbox_storage_backend}" - %{ if launch_darkly_api_key != "" } LAUNCH_DARKLY_API_KEY = "${launch_darkly_api_key}" %{ endif } diff --git a/iac/modules/job-api/main.tf b/iac/modules/job-api/main.tf index 088b7db195..1338d85537 100644 --- a/iac/modules/job-api/main.tf +++ b/iac/modules/job-api/main.tf @@ -47,7 +47,6 @@ resource "nomad_job" "api" { clickhouse_connection_string = var.clickhouse_connection_string loki_url = var.loki_url sandbox_access_token_hash_seed = var.sandbox_access_token_hash_seed - sandbox_storage_backend = var.sandbox_storage_backend db_migrator_docker_image = var.db_migrator_docker_image launch_darkly_api_key = trimspace(var.launch_darkly_api_key) default_persistent_volume_type = var.default_persistent_volume_type diff --git a/iac/modules/job-api/variables.tf b/iac/modules/job-api/variables.tf index 6bae861bb6..7736d1edb7 100644 --- a/iac/modules/job-api/variables.tf +++ b/iac/modules/job-api/variables.tf @@ -120,11 +120,6 @@ variable "sandbox_access_token_hash_seed" { sensitive = true } -variable "sandbox_storage_backend" { - type = string - default = "memory" -} - variable "redis_url" { type = string sensitive = true diff --git a/iac/provider-gcp/Makefile b/iac/provider-gcp/Makefile index 5a12239b98..ef8f1dfa74 100644 --- a/iac/provider-gcp/Makefile +++ b/iac/provider-gcp/Makefile @@ -40,7 +40,6 @@ tf_vars := \ $(call tfvar, GCP_ZONE) \ $(call tfvar, DOMAIN_NAME) \ $(call tfvar, PREFIX) \ - $(call tfvar, SANDBOX_STORAGE_BACKEND) \ $(call tfvar, ORCHESTRATOR_ENABLED) \ $(call tfvar, ALLOW_SANDBOX_INTERNET) \ $(call tfvar, API_INTERNAL_GRPC_PORT) \ diff --git a/iac/provider-gcp/main.tf b/iac/provider-gcp/main.tf index 5ee34982e7..312753da6c 100644 --- a/iac/provider-gcp/main.tf +++ b/iac/provider-gcp/main.tf @@ -258,7 +258,6 @@ module "nomad" { redis_cluster_url_secret_version = module.init.redis_cluster_url_secret_version redis_tls_ca_base64_secret_version = module.init.redis_tls_ca_base64_secret_version sandbox_access_token_hash_seed = random_password.sandbox_access_token_hash_seed.result - sandbox_storage_backend = var.sandbox_storage_backend db_max_open_connections = var.db_max_open_connections db_min_idle_connections = var.db_min_idle_connections auth_db_max_open_connections = var.auth_db_max_open_connections diff --git a/iac/provider-gcp/nomad/main.tf b/iac/provider-gcp/nomad/main.tf index 7180de588d..3f6a7ef0ce 100644 --- a/iac/provider-gcp/nomad/main.tf +++ b/iac/provider-gcp/nomad/main.tf @@ -140,7 +140,6 @@ module "api" { clickhouse_connection_string = local.clickhouse_connection_string loki_url = local.loki_url sandbox_access_token_hash_seed = var.sandbox_access_token_hash_seed - sandbox_storage_backend = var.sandbox_storage_backend db_max_open_connections = var.db_max_open_connections db_min_idle_connections = var.db_min_idle_connections auth_db_max_open_connections = var.auth_db_max_open_connections diff --git a/iac/provider-gcp/nomad/variables.tf b/iac/provider-gcp/nomad/variables.tf index 17286ee38f..cfc3270d39 100644 --- a/iac/provider-gcp/nomad/variables.tf +++ b/iac/provider-gcp/nomad/variables.tf @@ -113,11 +113,6 @@ variable "sandbox_access_token_hash_seed" { type = string } -variable "sandbox_storage_backend" { - type = string - default = "memory" -} - variable "db_max_open_connections" { type = number } diff --git a/iac/provider-gcp/variables.tf b/iac/provider-gcp/variables.tf index 72bdfb8642..4f99b1c14c 100644 --- a/iac/provider-gcp/variables.tf +++ b/iac/provider-gcp/variables.tf @@ -715,12 +715,6 @@ variable "loki_boot_disk_type" { default = "pd-ssd" } -variable "sandbox_storage_backend" { - description = "The sandbox storage backend to use. Valid values: 'memory', 'redis'." - type = string - default = "" -} - variable "db_max_open_connections" { type = number default = 40 diff --git a/packages/api/go.mod b/packages/api/go.mod index 38e87fb575..d995b33a73 100644 --- a/packages/api/go.mod +++ b/packages/api/go.mod @@ -41,7 +41,6 @@ require ( github.com/launchdarkly/go-server-sdk/v7 v7.13.0 github.com/oapi-codegen/gin-middleware v1.0.2 github.com/oapi-codegen/runtime v1.4.0 - github.com/orcaman/concurrent-map/v2 v2.0.1 github.com/posthog/posthog-go v0.0.0-20230801140217-d607812dee69 github.com/redis/go-redis/v9 v9.17.3 github.com/stretchr/testify v1.11.1 @@ -295,6 +294,7 @@ require ( github.com/opentracing-contrib/go-grpc v0.1.2 // indirect github.com/opentracing-contrib/go-stdlib v1.1.0 // indirect github.com/opentracing/opentracing-go v1.2.1-0.20220228012449-10b1cf09e00b // indirect + github.com/orcaman/concurrent-map/v2 v2.0.1 // indirect github.com/patrickmn/go-cache v2.1.0+incompatible // indirect github.com/paulmach/orb v0.11.1 // indirect github.com/pb33f/jsonpath v0.8.2 // indirect diff --git a/packages/api/internal/cfg/model.go b/packages/api/internal/cfg/model.go index eea6a87fd0..d3c24d7d91 100644 --- a/packages/api/internal/cfg/model.go +++ b/packages/api/internal/cfg/model.go @@ -16,11 +16,6 @@ import ( ) const ( - // SandboxStorageBackendMemory will use memory backend as a primary storage for sandbox data. - // It will also keep redis populated to allow for seamless migration to redis. - SandboxStorageBackendMemory = "memory" - SandboxStorageBackendRedis = "redis" - // ServiceDiscoveryProviderNomad queries Nomad's HTTP API (the original / Nomad-based deploy). ServiceDiscoveryProviderNomad = "nomad" // ServiceDiscoveryProviderKubernetes queries the in-cluster K8s API (the K8s deploy). @@ -92,10 +87,6 @@ type Config struct { DefaultPersistentVolumeType string `env:"DEFAULT_PERSISTENT_VOLUME_TYPE"` - // SandboxStorageBackend selects the sandbox storage implementation. - // "redis" uses Redis directly; "populate_redis" uses in-memory with Redis shadow writes. - SandboxStorageBackend string `env:"SANDBOX_STORAGE_BACKEND" envDefault:"memory"` - DomainName string `env:"DOMAIN_NAME" envDefault:""` } @@ -165,10 +156,6 @@ func Parse() (Config, error) { config.AuthDBConnectionString = config.PostgresConnectionString } - if !slices.Contains([]string{SandboxStorageBackendMemory, SandboxStorageBackendRedis}, config.SandboxStorageBackend) { - return config, fmt.Errorf("invalid sandbox storage backend: %s", config.SandboxStorageBackend) - } - if !slices.Contains([]string{ServiceDiscoveryProviderNomad, ServiceDiscoveryProviderKubernetes, ServiceDiscoveryProviderLocal}, config.ServiceDiscoveryProvider) { return config, fmt.Errorf("invalid service discovery provider: %s", config.ServiceDiscoveryProvider) } diff --git a/packages/api/internal/cfg/model_test.go b/packages/api/internal/cfg/model_test.go index 745b91edf2..e5a65b96f4 100644 --- a/packages/api/internal/cfg/model_test.go +++ b/packages/api/internal/cfg/model_test.go @@ -51,13 +51,6 @@ func TestParse(t *testing.T) { require.NoError(t, err) assert.Equal(t, content, result.VolumesToken.SigningKey) }) - - t.Run("test sandbox backend empty string", func(t *testing.T) { - t.Setenv("SANDBOX_STORAGE_BACKEND", "") - result, err := Parse() - require.NoError(t, err) - assert.Equal(t, SandboxStorageBackendMemory, result.SandboxStorageBackend) - }) } // removeEnv was mostly copied from the implementation of t.Setenv diff --git a/packages/api/internal/orchestrator/autoresume_test.go b/packages/api/internal/orchestrator/autoresume_test.go index 01d134d50e..249ec51041 100644 --- a/packages/api/internal/orchestrator/autoresume_test.go +++ b/packages/api/internal/orchestrator/autoresume_test.go @@ -9,19 +9,29 @@ import ( "github.com/google/uuid" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "go.opentelemetry.io/otel/metric/noop" "github.com/e2b-dev/infra/packages/api/internal/orchestrator/nodemanager" "github.com/e2b-dev/infra/packages/api/internal/sandbox" - "github.com/e2b-dev/infra/packages/api/internal/sandbox/reservations" - sandboxmemory "github.com/e2b-dev/infra/packages/api/internal/sandbox/storage/memory" + redisreservations "github.com/e2b-dev/infra/packages/api/internal/sandbox/reservations/redis" + sandboxredis "github.com/e2b-dev/infra/packages/api/internal/sandbox/storage/redis" + redis_utils "github.com/e2b-dev/infra/packages/shared/pkg/redis" "github.com/e2b-dev/infra/packages/shared/pkg/smap" ) -func newTestAutoResumeOrchestrator() *Orchestrator { +func newTestAutoResumeOrchestrator(t *testing.T) *Orchestrator { + t.Helper() + + client := redis_utils.SetupInstance(t) + storage, err := sandboxredis.NewStorage(client, noop.NewMeterProvider()) + require.NoError(t, err) + go storage.Start(t.Context()) + t.Cleanup(func() { storage.Close(context.WithoutCancel(t.Context())) }) + return &Orchestrator{ sandboxStore: sandbox.NewStore( - sandboxmemory.NewStorage(), - reservations.NewReservationStorage(), + storage, + redisreservations.NewReservationStorage(client, storage.Notifier()), sandbox.Callbacks{ AddSandboxToRoutingTable: func(context.Context, sandbox.Sandbox) {}, AsyncNewlyCreatedSandbox: func(context.Context, sandbox.Sandbox, sandbox.CreationMetadata) {}, @@ -64,7 +74,7 @@ func TestHandleExistingSandboxAutoResume(t *testing.T) { t.Run("running sandbox returns node ip immediately", func(t *testing.T) { t.Parallel() - o := newTestAutoResumeOrchestrator() + o := newTestAutoResumeOrchestrator(t) sbx := testSandboxForAutoResume(sandbox.StateRunning) registerNode(o, sbx, "10.0.0.1") @@ -77,7 +87,7 @@ func TestHandleExistingSandboxAutoResume(t *testing.T) { t.Run("running sandbox with empty ip returns error", func(t *testing.T) { t.Parallel() - o := newTestAutoResumeOrchestrator() + o := newTestAutoResumeOrchestrator(t) sbx := testSandboxForAutoResume(sandbox.StateRunning) registerNode(o, sbx, "") @@ -90,7 +100,7 @@ func TestHandleExistingSandboxAutoResume(t *testing.T) { t.Run("snapshotting sandbox waits and routes when transition finishes", func(t *testing.T) { t.Parallel() - o := newTestAutoResumeOrchestrator() + o := newTestAutoResumeOrchestrator(t) sbx := testSandboxForAutoResume(sandbox.StateRunning) addSandbox(t, o, sbx) registerNode(o, sbx, "10.0.0.2") @@ -118,7 +128,7 @@ func TestHandleExistingSandboxAutoResume(t *testing.T) { t.Run("pausing sandbox returns still transitioning after retries", func(t *testing.T) { t.Parallel() - o := newTestAutoResumeOrchestrator() + o := newTestAutoResumeOrchestrator(t) sbx := testSandboxForAutoResume(sandbox.StateRunning) addSandbox(t, o, sbx) @@ -138,10 +148,10 @@ func TestHandleExistingSandboxAutoResume(t *testing.T) { assert.ErrorIs(t, err, ErrSandboxStillTransitioning) }) - t.Run("pausing sandbox wait failure returns internal error", func(t *testing.T) { + t.Run("pausing sandbox wait failure returns still transitioning", func(t *testing.T) { t.Parallel() - o := newTestAutoResumeOrchestrator() + o := newTestAutoResumeOrchestrator(t) sbx := testSandboxForAutoResume(sandbox.StateRunning) addSandbox(t, o, sbx) @@ -157,13 +167,13 @@ func TestHandleExistingSandboxAutoResume(t *testing.T) { _, handled, err := o.HandleExistingSandboxAutoResume(t.Context(), sbx.TeamID, sbx.SandboxID, pausingSandbox, time.Minute) require.Error(t, err) assert.False(t, handled) - assert.EqualError(t, err, "error waiting for sandbox to pause") + assert.ErrorIs(t, err, ErrSandboxStillTransitioning) }) t.Run("pausing sandbox wait timeout returns failed precondition", func(t *testing.T) { t.Parallel() - o := newTestAutoResumeOrchestrator() + o := newTestAutoResumeOrchestrator(t) sbx := testSandboxForAutoResume(sandbox.StateRunning) addSandbox(t, o, sbx) @@ -183,7 +193,7 @@ func TestHandleExistingSandboxAutoResume(t *testing.T) { t.Run("killing sandbox returns not found", func(t *testing.T) { t.Parallel() - o := newTestAutoResumeOrchestrator() + o := newTestAutoResumeOrchestrator(t) sbx := testSandboxForAutoResume(sandbox.StateKilling) _, handled, err := o.HandleExistingSandboxAutoResume(t.Context(), sbx.TeamID, sbx.SandboxID, sbx, time.Minute) @@ -195,7 +205,7 @@ func TestHandleExistingSandboxAutoResume(t *testing.T) { t.Run("unknown sandbox state returns internal error", func(t *testing.T) { t.Parallel() - o := newTestAutoResumeOrchestrator() + o := newTestAutoResumeOrchestrator(t) sbx := testSandboxForAutoResume(sandbox.State("mystery")) _, handled, err := o.HandleExistingSandboxAutoResume(t.Context(), sbx.TeamID, sbx.SandboxID, sbx, time.Minute) diff --git a/packages/api/internal/orchestrator/create_instance_events_test.go b/packages/api/internal/orchestrator/create_instance_events_test.go index abf41440dc..f8d57d6e26 100644 --- a/packages/api/internal/orchestrator/create_instance_events_test.go +++ b/packages/api/internal/orchestrator/create_instance_events_test.go @@ -18,9 +18,10 @@ import ( "github.com/e2b-dev/infra/packages/api/internal/orchestrator/nodemanager" "github.com/e2b-dev/infra/packages/api/internal/orchestrator/placement" "github.com/e2b-dev/infra/packages/api/internal/sandbox" - "github.com/e2b-dev/infra/packages/api/internal/sandbox/reservations" - sandboxmemory "github.com/e2b-dev/infra/packages/api/internal/sandbox/storage/memory" + redisreservations "github.com/e2b-dev/infra/packages/api/internal/sandbox/reservations/redis" + sandboxredis "github.com/e2b-dev/infra/packages/api/internal/sandbox/storage/redis" "github.com/e2b-dev/infra/packages/shared/pkg/featureflags" + redis_utils "github.com/e2b-dev/infra/packages/shared/pkg/redis" "github.com/e2b-dev/infra/packages/shared/pkg/smap" ) @@ -55,9 +56,15 @@ func newOrchestratorWithCounter(t *testing.T) (*Orchestrator, *eventCounter) { ec := &eventCounter{} + client := redis_utils.SetupInstance(t) + storage, err := sandboxredis.NewStorage(client, noop.NewMeterProvider()) + require.NoError(t, err) + go storage.Start(t.Context()) + t.Cleanup(func() { storage.Close(context.WithoutCancel(t.Context())) }) + store := sandbox.NewStore( - sandboxmemory.NewStorage(), - reservations.NewReservationStorage(), + storage, + redisreservations.NewReservationStorage(client, storage.Notifier()), sandbox.Callbacks{ AddSandboxToRoutingTable: func(context.Context, sandbox.Sandbox) {}, AsyncNewlyCreatedSandbox: ec.callback(), diff --git a/packages/api/internal/orchestrator/create_instance_test.go b/packages/api/internal/orchestrator/create_instance_test.go index 383b8c83b8..d0630a274e 100644 --- a/packages/api/internal/orchestrator/create_instance_test.go +++ b/packages/api/internal/orchestrator/create_instance_test.go @@ -15,12 +15,13 @@ import ( "github.com/e2b-dev/infra/packages/api/internal/orchestrator/nodemanager" "github.com/e2b-dev/infra/packages/api/internal/orchestrator/placement" "github.com/e2b-dev/infra/packages/api/internal/sandbox" - "github.com/e2b-dev/infra/packages/api/internal/sandbox/reservations" - sandboxmemory "github.com/e2b-dev/infra/packages/api/internal/sandbox/storage/memory" + redisreservations "github.com/e2b-dev/infra/packages/api/internal/sandbox/reservations/redis" + sandboxredis "github.com/e2b-dev/infra/packages/api/internal/sandbox/storage/redis" teamtypes "github.com/e2b-dev/infra/packages/auth/pkg/types" authqueries "github.com/e2b-dev/infra/packages/db/pkg/auth/queries" "github.com/e2b-dev/infra/packages/db/queries" "github.com/e2b-dev/infra/packages/shared/pkg/featureflags" + redis_utils "github.com/e2b-dev/infra/packages/shared/pkg/redis" "github.com/e2b-dev/infra/packages/shared/pkg/smap" ) @@ -45,9 +46,15 @@ func testBuild() queries.EnvBuild { func newCreateSandboxTestOrchestrator(t *testing.T) (*Orchestrator, *nodemanager.Node) { t.Helper() + client := redis_utils.SetupInstance(t) + storage, err := sandboxredis.NewStorage(client, noop.NewMeterProvider()) + require.NoError(t, err) + go storage.Start(t.Context()) + t.Cleanup(func() { storage.Close(context.WithoutCancel(t.Context())) }) + store := sandbox.NewStore( - sandboxmemory.NewStorage(), - reservations.NewReservationStorage(), + storage, + redisreservations.NewReservationStorage(client, storage.Notifier()), sandbox.Callbacks{ AddSandboxToRoutingTable: func(context.Context, sandbox.Sandbox) {}, AsyncNewlyCreatedSandbox: func(context.Context, sandbox.Sandbox, sandbox.CreationMetadata) {}, @@ -57,8 +64,8 @@ func newCreateSandboxTestOrchestrator(t *testing.T) (*Orchestrator, *nodemanager meter := noop.NewMeterProvider().Meter("github.com/e2b-dev/infra/packages/api/internal/orchestrator") counter, _ := meter.Int64Counter("test-created-sandboxes") - ffClient, err := featureflags.NewClientWithDatasource(ldtestdata.DataSource()) - require.NoError(t, err) + ffClient, ffErr := featureflags.NewClientWithDatasource(ldtestdata.DataSource()) + require.NoError(t, ffErr) algo := placement.NewBestOfK(placement.DefaultBestOfKConfig()).(*placement.BestOfK) diff --git a/packages/api/internal/orchestrator/orchestrator.go b/packages/api/internal/orchestrator/orchestrator.go index 934f228d5f..472fb34128 100644 --- a/packages/api/internal/orchestrator/orchestrator.go +++ b/packages/api/internal/orchestrator/orchestrator.go @@ -21,10 +21,7 @@ import ( "github.com/e2b-dev/infra/packages/api/internal/orchestrator/nodemanager" "github.com/e2b-dev/infra/packages/api/internal/orchestrator/placement" "github.com/e2b-dev/infra/packages/api/internal/sandbox" - "github.com/e2b-dev/infra/packages/api/internal/sandbox/reservations" redisreservations "github.com/e2b-dev/infra/packages/api/internal/sandbox/reservations/redis" - "github.com/e2b-dev/infra/packages/api/internal/sandbox/storage/memory" - "github.com/e2b-dev/infra/packages/api/internal/sandbox/storage/populate_redis" redisbackend "github.com/e2b-dev/infra/packages/api/internal/sandbox/storage/redis" sqlcdb "github.com/e2b-dev/infra/packages/db/client" "github.com/e2b-dev/infra/packages/shared/pkg/env" @@ -156,27 +153,9 @@ func New( snapshotUpsertSem: snapshotUpsertSem, } - var reservationStorage sandbox.ReservationStorage - var sandboxStorage sandbox.Storage - - switch config.SandboxStorageBackend { - case cfg.SandboxStorageBackendMemory: - reservationStorage = reservations.NewReservationStorage() - sandboxStorage = populate_redis.NewStorage(memory.NewStorage(), redisStorage) - logger.L().Info(ctx, "Using populate_redis sandbox storage backend") - - go redisbackend.NewCleaner(redisStorage).Start(ctx) - case cfg.SandboxStorageBackendRedis: - reservationStorage = redisreservations.NewReservationStorage(redisClient, redisStorage.Notifier()) - sandboxStorage = redisStorage - logger.L().Info(ctx, "Using redis sandbox storage backend") - default: - return nil, fmt.Errorf("invalid sandbox storage backend: %s", config.SandboxStorageBackend) - } - o.sandboxStore = sandbox.NewStore( - sandboxStorage, - reservationStorage, + redisStorage, + redisreservations.NewReservationStorage(redisClient, redisStorage.Notifier()), sandbox.Callbacks{ AddSandboxToRoutingTable: o.addSandboxToRoutingTable, AsyncNewlyCreatedSandbox: o.handleNewlyCreatedSandbox, diff --git a/packages/api/internal/sandbox/reservations/reservation.go b/packages/api/internal/sandbox/reservations/reservation.go deleted file mode 100644 index 870a8c6422..0000000000 --- a/packages/api/internal/sandbox/reservations/reservation.go +++ /dev/null @@ -1,103 +0,0 @@ -package reservations - -import ( - "context" - - "github.com/google/uuid" - "go.uber.org/zap" - - "github.com/e2b-dev/infra/packages/api/internal/sandbox/sandboxtypes" - "github.com/e2b-dev/infra/packages/shared/pkg/logger" - "github.com/e2b-dev/infra/packages/shared/pkg/smap" - "github.com/e2b-dev/infra/packages/shared/pkg/utils" -) - -type sandboxReservation struct { - start *utils.SetOnce[sandboxtypes.Sandbox] -} - -func newSandboxReservation(start *utils.SetOnce[sandboxtypes.Sandbox]) *sandboxReservation { - return &sandboxReservation{ - start: start, - } -} - -type TeamSandboxes map[string]*sandboxReservation - -type ReservationStorage struct { - reservations *smap.Map[TeamSandboxes] -} - -var _ sandboxtypes.ReservationStorage = &ReservationStorage{} - -func NewReservationStorage() *ReservationStorage { - return &ReservationStorage{ - reservations: smap.New[TeamSandboxes](), - } -} - -func (s *ReservationStorage) Reserve(ctx context.Context, teamID uuid.UUID, sandboxID string, limit int) (finishStart func(sandboxtypes.Sandbox, error), waitForStart func(ctx context.Context) (sandboxtypes.Sandbox, error), err error) { - alreadyPresent := false - limitExceeded := false - var startResult *utils.SetOnce[sandboxtypes.Sandbox] - - teamIDStr := teamID.String() - s.reservations.Upsert(teamIDStr, nil, func(exist bool, teamSandboxes, _ TeamSandboxes) TeamSandboxes { - if !exist { - teamSandboxes = make(map[string]*sandboxReservation) - } - - if sbx, ok := teamSandboxes[sandboxID]; ok { - alreadyPresent = true - startResult = sbx.start - - return teamSandboxes - } - - if limit >= 0 && len(teamSandboxes) >= limit { - limitExceeded = true - - return teamSandboxes - } - - startResult = utils.NewSetOnce[sandboxtypes.Sandbox]() - teamSandboxes[sandboxID] = newSandboxReservation(startResult) - - return teamSandboxes - }) - - if limitExceeded { - return nil, nil, &sandboxtypes.LimitExceededError{TeamID: teamID} - } - - if alreadyPresent { - return nil, startResult.WaitWithContext, nil - } - - return func(sbx sandboxtypes.Sandbox, err error) { - setErr := startResult.SetResult(sbx, err) - if setErr != nil { - logger.L().Error(ctx, "failed to set the result of the reservation", zap.Error(setErr), logger.WithSandboxID(sandboxID)) - } - - // Remove the reservation if the sandbox creation failed - if err != nil { - _ = s.Release(ctx, teamID, sandboxID) - } - }, nil, nil -} - -func (s *ReservationStorage) Release(_ context.Context, teamID uuid.UUID, sandboxID string) error { - teamIDStr := teamID.String() - s.reservations.RemoveCb(teamIDStr, func(_ string, ts TeamSandboxes, exists bool) bool { - if !exists { - return true - } - - delete(ts, sandboxID) - - return len(ts) == 0 - }) - - return nil -} diff --git a/packages/api/internal/sandbox/reservations/reservation_test.go b/packages/api/internal/sandbox/reservations/reservation_test.go deleted file mode 100644 index a0aa828c9c..0000000000 --- a/packages/api/internal/sandbox/reservations/reservation_test.go +++ /dev/null @@ -1,553 +0,0 @@ -package reservations - -import ( - "context" - "errors" - "fmt" - "sync" - "sync/atomic" - "testing" - "time" - - "github.com/google/uuid" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" - "golang.org/x/sync/errgroup" - - "github.com/e2b-dev/infra/packages/api/internal/sandbox/sandboxtypes" - "github.com/e2b-dev/infra/packages/shared/pkg/consts" -) - -const ( - sandboxID = "test-sandbox-id" -) - -var teamID = uuid.New() - -func newReservationStorage() *ReservationStorage { - cache := NewReservationStorage() - - return cache -} - -func TestReservation(t *testing.T) { - t.Parallel() - cache := newReservationStorage() - - _, _, err := cache.Reserve(t.Context(), teamID, sandboxID, 1) - assert.NoError(t, err) -} - -func TestReservation_Exceeded(t *testing.T) { - t.Parallel() - cache := newReservationStorage() - - _, _, err := cache.Reserve(t.Context(), teamID, sandboxID, 1) - require.NoError(t, err) - _, _, err = cache.Reserve(t.Context(), teamID, "sandbox-2", 1) - require.ErrorAs(t, err, new(&sandboxtypes.LimitExceededError{})) -} - -func TestReservation_SameSandbox(t *testing.T) { - t.Parallel() - cache := newReservationStorage() - - _, _, err := cache.Reserve(t.Context(), teamID, sandboxID, 1) - require.NoError(t, err) - - _, waitForStart, err := cache.Reserve(t.Context(), teamID, sandboxID, 1) - require.NoError(t, err) - assert.NotNil(t, waitForStart) -} - -func TestReservation_Release(t *testing.T) { - t.Parallel() - cache := newReservationStorage() - - _, _, err := cache.Reserve(t.Context(), teamID, sandboxID, 1) - require.NoError(t, err) - err = cache.Release(t.Context(), teamID, sandboxID) - require.NoError(t, err) - - _, _, err = cache.Reserve(t.Context(), teamID, sandboxID, 1) - assert.NoError(t, err) -} - -func TestReservation_ResumeAlreadyRunningSandbox(t *testing.T) { - t.Parallel() - cache := newReservationStorage() - - _, _, err := cache.Reserve(t.Context(), teamID, sandboxID, 1) - require.NoError(t, err) - - _, waitForStart, err := cache.Reserve(t.Context(), teamID, sandboxID, 1) - require.NoError(t, err) - assert.NotNil(t, waitForStart) -} - -func TestReservation_WaitForStart(t *testing.T) { - t.Parallel() - cache := newReservationStorage() - - finishStart, _, err := cache.Reserve(t.Context(), teamID, sandboxID, 10) - require.NoError(t, err) - require.NotNil(t, finishStart) - - // Second call should return waitForStart - _, waitForStart, err := cache.Reserve(t.Context(), teamID, sandboxID, 10) - require.NoError(t, err) - require.NotNil(t, waitForStart) - - // Finish the start operation - expectedSbx := sandboxtypes.Sandbox{ - ClientID: consts.ClientID, - SandboxID: sandboxID, - TemplateID: "test", - TeamID: teamID, - StartTime: time.Now(), - EndTime: time.Now().Add(time.Hour), - MaxInstanceLength: time.Hour, - } - finishStart(expectedSbx, nil) - - // Wait should now complete and return the sandbox - ctx := t.Context() - result, err := waitForStart(ctx) - require.NoError(t, err) - assert.Equal(t, expectedSbx.SandboxID, result.SandboxID) - assert.Equal(t, expectedSbx.TemplateID, result.TemplateID) -} - -func TestReservation_WaitForStartError(t *testing.T) { - t.Parallel() - cache := newReservationStorage() - - finishStart, _, err := cache.Reserve(t.Context(), teamID, sandboxID, 10) - require.NoError(t, err) - require.NotNil(t, finishStart) - - // Second call should return waitForStart - _, waitForStart, err := cache.Reserve(t.Context(), teamID, sandboxID, 10) - require.NoError(t, err) - require.NotNil(t, waitForStart) - - // Finish with an error - expectedErr := assert.AnError - finishStart(sandboxtypes.Sandbox{}, expectedErr) - - // Wait should return the error - ctx := t.Context() - _, err = waitForStart(ctx) - require.Error(t, err) - assert.Equal(t, expectedErr, err) -} - -func TestReservation_MultipleWaiters(t *testing.T) { - t.Parallel() - cache := newReservationStorage() - - finishStart, _, err := cache.Reserve(t.Context(), teamID, sandboxID, 10) - require.NoError(t, err) - require.NotNil(t, finishStart) - - // Multiple calls should all return waitForStart - _, waitForStart1, err := cache.Reserve(t.Context(), teamID, sandboxID, 10) - require.NoError(t, err) - require.NotNil(t, waitForStart1) - - _, waitForStart2, err := cache.Reserve(t.Context(), teamID, sandboxID, 10) - require.NoError(t, err) - require.NotNil(t, waitForStart2) - - // Finish the start operation - expectedSbx := sandboxtypes.Sandbox{ - ClientID: consts.ClientID, - SandboxID: sandboxID, - TemplateID: "test", - TeamID: teamID, - StartTime: time.Now(), - EndTime: time.Now().Add(time.Hour), - MaxInstanceLength: time.Hour, - } - finishStart(expectedSbx, nil) - - // All waiters should get the result - ctx := t.Context() - result1, err := waitForStart1(ctx) - require.NoError(t, err) - assert.Equal(t, expectedSbx.SandboxID, result1.SandboxID) - - result2, err := waitForStart2(ctx) - require.NoError(t, err) - assert.Equal(t, expectedSbx.SandboxID, result2.SandboxID) -} - -func TestReservation_Remove(t *testing.T) { - t.Parallel() - cache := newReservationStorage() - - finishStart, _, err := cache.Reserve(t.Context(), teamID, sandboxID, 1) - require.NoError(t, err) - require.NotNil(t, finishStart) - - expectedSbx := sandboxtypes.Sandbox{ - ClientID: consts.ClientID, - SandboxID: sandboxID, - TemplateID: "test", - TeamID: teamID, - StartTime: time.Now(), - EndTime: time.Now().Add(time.Hour), - MaxInstanceLength: time.Hour, - } - finishStart(expectedSbx, nil) - - // Remove the reservation - err = cache.Release(t.Context(), teamID, sandboxID) - require.NoError(t, err) - - // Should be able to reserve again - finishStart2, _, err := cache.Reserve(t.Context(), teamID, sandboxID, 1) - require.NoError(t, err) - require.NotNil(t, finishStart2) -} - -func TestReservation_MultipleTeams(t *testing.T) { - t.Parallel() - cache := newReservationStorage() - - team1 := uuid.New() - team2 := uuid.New() - sandbox1 := "sandbox-1" - sandbox2 := "sandbox-2" - - // Reserve for team1 - _, _, err := cache.Reserve(t.Context(), team1, sandbox1, 1) - require.NoError(t, err) - - // Should not affect team2's limit - _, _, err = cache.Reserve(t.Context(), team2, sandbox2, 1) - require.NoError(t, err) - - // team1 should be at limit - _, _, err = cache.Reserve(t.Context(), team1, "sandbox-3", 1) - require.ErrorAs(t, err, new(&sandboxtypes.LimitExceededError{})) - - // team2 should also be at limit - _, _, err = cache.Reserve(t.Context(), team2, "sandbox-4", 1) - require.ErrorAs(t, err, new(&sandboxtypes.LimitExceededError{})) -} - -func TestReservation_FailedStart(t *testing.T) { - t.Parallel() - cache := newReservationStorage() - team := uuid.New() - sbxID := "failed-sandbox" - - // Reserve sandbox - finishStart, _, err := cache.Reserve(t.Context(), team, sbxID, 10) - require.NoError(t, err) - require.NotNil(t, finishStart) - - // Finish with an error - expectedErr := errors.New("start failed") - finishStart(sandboxtypes.Sandbox{}, expectedErr) - - // After failed start, should be able to reserve again - finishStart2, _, err := cache.Reserve(t.Context(), team, sbxID, 10) - require.NoError(t, err) - require.NotNil(t, finishStart2) -} - -func TestReservation_FailedStartWithWaiters(t *testing.T) { - t.Parallel() - cache := newReservationStorage() - team := uuid.New() - sbxID := "failed-with-waiters" - numWaiters := 10 - - // First reservation - finishStart, _, err := cache.Reserve(t.Context(), team, sbxID, 100) - require.NoError(t, err) - require.NotNil(t, finishStart) - - var wg errgroup.Group - waiters := make([]func(ctx context.Context) (sandboxtypes.Sandbox, error), numWaiters) - - // Multiple waiters - for i := range numWaiters { - wg.Go(func() error { - _, waitForStart, err := cache.Reserve(t.Context(), team, sbxID, 100) - if err != nil { - return err - } - - if waitForStart == nil { - return errors.New("waitForStart should not be nil") - } - waiters[i] = waitForStart - - return nil - }) - } - - wg.Wait() - - // Finish with an error - expectedErr := errors.New("start failed") - finishStart(sandboxtypes.Sandbox{}, expectedErr) - - // All waiters should receive the error - var wg2 sync.WaitGroup - var errorCount atomic.Int32 - - for _, waiter := range waiters { - wg2.Add(1) - go func(w func(ctx context.Context) (sandboxtypes.Sandbox, error)) { - defer wg2.Done() - _, err := w(t.Context()) - if err != nil { - errorCount.Add(1) - } - }(waiter) - } - - wg2.Wait() - assert.Equal(t, int32(numWaiters), errorCount.Load()) -} - -func TestReservation_ConcurrentReservations(t *testing.T) { - t.Parallel() - cache := newReservationStorage() - team := uuid.New() - concurrency := 100 - limit := 50 - - var wg sync.WaitGroup - var successCount atomic.Int32 - var limitExceededCount atomic.Int32 - - for i := range concurrency { - wg.Go(func() { - sandboxID := fmt.Sprintf("sandbox-%d", i) - _, _, err := cache.Reserve(t.Context(), team, sandboxID, limit) - if err == nil { - successCount.Add(1) - } else { - var limitExceededError *sandboxtypes.LimitExceededError - if errors.As(err, &limitExceededError) { - limitExceededCount.Add(1) - } - } - }) - } - - wg.Wait() - - // Should have exactly 50 successful reservations and 50 limit exceeded errors - assert.Equal(t, int32(limit), successCount.Load()) - assert.Equal(t, int32(concurrency)-int32(limit), limitExceededCount.Load()) -} - -func TestReservation_ConcurrentSameSandbox(t *testing.T) { - t.Parallel() - cache := newReservationStorage() - team := uuid.New() - sbxID := "concurrent-sandbox" - concurrency := 50 - - var wg errgroup.Group - var finishStartCount atomic.Int32 - var waitForStartCount atomic.Int32 - - // Multiple goroutines try to reserve the same sandbox - for range concurrency { - wg.Go(func() error { - finishStart, waitForStart, err := cache.Reserve(t.Context(), team, sbxID, 10) - if err != nil { - return err - } - - if finishStart != nil { - finishStartCount.Add(1) - } - if waitForStart != nil { - waitForStartCount.Add(1) - } - - return nil - }) - } - - wg.Wait() - - // Only one should get finishStart, all others should get waitForStart - assert.Equal(t, int32(1), finishStartCount.Load()) - assert.Equal(t, int32(concurrency-1), waitForStartCount.Load()) -} - -func TestReservation_ConcurrentWaitAndFinish(t *testing.T) { - t.Parallel() - cache := newReservationStorage() - team := uuid.New() - sbxID := "wait-finish-sandbox" - numWaiters := 20 - - // First goroutine reserves - finishStart, _, err := cache.Reserve(t.Context(), team, sbxID, 1) - require.NoError(t, err) - require.NotNil(t, finishStart) - - var wg errgroup.Group - waiters := make([]func(ctx context.Context) (sandboxtypes.Sandbox, error), numWaiters) - - // Multiple waiters - for i := range numWaiters { - wg.Go(func() error { - _, waitForStart, err := cache.Reserve(t.Context(), team, sbxID, 1) - if err != nil { - return err - } - - if waitForStart == nil { - return errors.New("waitForStart should not be nil") - } - - waiters[i] = waitForStart - - return nil - }) - } - - wg.Wait() - - // Finish the start operation - expectedSbx := sandboxtypes.Sandbox{ - ClientID: consts.ClientID, - SandboxID: sbxID, - TemplateID: "test", - TeamID: team, - StartTime: time.Now(), - EndTime: time.Now().Add(time.Hour), - MaxInstanceLength: time.Hour, - } - finishStart(expectedSbx, nil) - - // All waiters should receive the result - var wg2 sync.WaitGroup - var successCount atomic.Int32 - - for _, waiter := range waiters { - wg2.Add(1) - go func(w func(ctx context.Context) (sandboxtypes.Sandbox, error)) { - defer wg2.Done() - result, err := w(t.Context()) - if err == nil && result.SandboxID == sbxID { - successCount.Add(1) - } - }(waiter) - } - - wg2.Wait() - assert.Equal(t, int32(numWaiters), successCount.Load()) -} - -func TestReservation_ConcurrentRemove(t *testing.T) { - t.Parallel() - cache := newReservationStorage() - team := uuid.New() - concurrency := 50 - - var wg errgroup.Group - - // Concurrently reserve and remove sandboxes - for i := range concurrency { - wg.Go(func() error { - sbxID := fmt.Sprintf("sandbox-%d", i) - - // Reserve - _, _, err := cache.Reserve(t.Context(), team, sbxID, 100) - if err != nil { - return err - } - - // Remove - err = cache.Release(t.Context(), team, sbxID) - if err != nil { - return err - } - - // Should be able to reserve again - _, _, err = cache.Reserve(t.Context(), team, sbxID, 100) - if err != nil { - return err - } - - return nil - }) - } - - err := wg.Wait() - require.NoError(t, err) -} - -func TestReservation_RaceConditionStressTest(t *testing.T) { - t.Parallel() - cache := newReservationStorage() - team := uuid.New() - numOperations := 2000 - numSandboxes := 100 - limit := 5 - - var wg sync.WaitGroup - var operationCount atomic.Int32 - - // Mix of reserve, remove, and finish operations - for i := range numOperations { - wg.Go(func() { - sbxID := fmt.Sprintf("sandbox-%d", i%numSandboxes) - - switch i % 3 { - case 0: - // Reserve - finishStart, waitForStart, err := cache.Reserve(t.Context(), team, sbxID, limit) - if err == nil { - operationCount.Add(1) - if finishStart != nil { - // Immediately finish - go func() { - time.Sleep(time.Millisecond) - finishStart(sandboxtypes.Sandbox{ - SandboxID: sbxID, - TeamID: team, - }, nil) - }() - } - if waitForStart != nil { - // Try to wait - go func() { - _, _ = waitForStart(t.Context()) - }() - } - } else { - var limitExceededError *sandboxtypes.LimitExceededError - if errors.As(err, &limitExceededError) { - operationCount.Add(1) - } - } - case 1: - // Remove - _ = cache.Release(t.Context(), team, sbxID) - - operationCount.Add(1) - case 2: - // Reserve again - _, _, _ = cache.Reserve(t.Context(), team, sbxID, limit) - operationCount.Add(1) - } - }) - } - - wg.Wait() - - assert.Equal(t, operationCount.Load(), int32(numOperations)) -} diff --git a/packages/api/internal/sandbox/sandboxtypes/storage.go b/packages/api/internal/sandbox/sandboxtypes/storage.go index 4b706e2507..f01f93e1ce 100644 --- a/packages/api/internal/sandbox/sandboxtypes/storage.go +++ b/packages/api/internal/sandbox/sandboxtypes/storage.go @@ -7,16 +7,11 @@ import ( ) const ( - StorageNameMemory = "memory" - StorageNameRedis = "redis" - StorageNamePopulateRedis = "populate_redis" + StorageNameRedis = "redis" ) -// Storage is the persistence interface implemented by the memory and redis backends. -// -// TODO [ENG-3514]: Remove Name() and Sync() and nolint once migrated to Redis -type Storage interface { //nolint: interfacebloat - Name() string +// Storage is the persistence interface implemented by the redis backend. +type Storage interface { Add(ctx context.Context, sandbox Sandbox) error Get(ctx context.Context, teamID uuid.UUID, sandboxID string) (Sandbox, error) Remove(ctx context.Context, teamID uuid.UUID, sandboxID string) error diff --git a/packages/api/internal/sandbox/storage/memory/main.go b/packages/api/internal/sandbox/storage/memory/main.go deleted file mode 100644 index 4431948693..0000000000 --- a/packages/api/internal/sandbox/storage/memory/main.go +++ /dev/null @@ -1,23 +0,0 @@ -package memory - -import ( - cmap "github.com/orcaman/concurrent-map/v2" - - "github.com/e2b-dev/infra/packages/api/internal/sandbox/sandboxtypes" -) - -var _ sandboxtypes.Storage = (*Storage)(nil) - -type Storage struct { - items cmap.ConcurrentMap[string, *memorySandbox] -} - -func (s *Storage) Name() string { return sandboxtypes.StorageNameMemory } - -func NewStorage() *Storage { - instanceCache := &Storage{ - items: cmap.New[*memorySandbox](), - } - - return instanceCache -} diff --git a/packages/api/internal/sandbox/storage/memory/operations.go b/packages/api/internal/sandbox/storage/memory/operations.go deleted file mode 100644 index 1bd2f86ae0..0000000000 --- a/packages/api/internal/sandbox/storage/memory/operations.go +++ /dev/null @@ -1,280 +0,0 @@ -package memory - -import ( - "context" - "errors" - "fmt" - "slices" - "time" - - "github.com/google/uuid" - "go.uber.org/zap" - - "github.com/e2b-dev/infra/packages/api/internal/sandbox/sandboxtypes" - "github.com/e2b-dev/infra/packages/shared/pkg/logger" - "github.com/e2b-dev/infra/packages/shared/pkg/middleware/otel/joined" - "github.com/e2b-dev/infra/packages/shared/pkg/utils" -) - -// Add the sandbox to the cache -func (s *Storage) Add(_ context.Context, sbx sandboxtypes.Sandbox) error { - added := s.items.SetIfAbsent(sbx.SandboxID, newMemorySandbox(sbx)) - if !added { - return sandboxtypes.ErrAlreadyExists - } - - return nil -} - -// exists check if the sandbox exists in the cache or is being evicted. -func (s *Storage) exists(sandboxID string) bool { - return s.items.Has(sandboxID) -} - -// Get the item from the cache. -func (s *Storage) get(sandboxID string) (*memorySandbox, error) { - item, ok := s.items.Get(sandboxID) - if !ok { - return nil, fmt.Errorf("sandbox \"%s\" doesn't exist", sandboxID) - } - - return item, nil -} - -// Get the item from the cache. -func (s *Storage) Get(_ context.Context, teamID uuid.UUID, sandboxID string) (sandboxtypes.Sandbox, error) { - item, ok := s.items.Get(sandboxID) - if !ok { - return sandboxtypes.Sandbox{}, fmt.Errorf("sandbox %q: %w", sandboxID, sandboxtypes.ErrNotFound) - } - - data := item.Data() - if data.TeamID != teamID { - return sandboxtypes.Sandbox{}, fmt.Errorf("sandbox %q: %w", sandboxID, sandboxtypes.ErrNotFound) - } - - return data, nil -} - -func (s *Storage) Remove(_ context.Context, _ uuid.UUID, sandboxID string) error { - s.items.Remove(sandboxID) - - return nil -} - -func (s *Storage) getItems(teamID *uuid.UUID, states []sandboxtypes.State) []sandboxtypes.Sandbox { - items := make([]sandboxtypes.Sandbox, 0) - - s.items.IterCb(func(_ string, item *memorySandbox) { - data := item.Data() - - if teamID != nil && *teamID != data.TeamID { - return - } - - if len(states) > 0 && !slices.Contains(states, data.State) { - return - } - - items = append(items, data) - }) - - return items -} - -func (s *Storage) TeamItems(_ context.Context, teamID uuid.UUID, states []sandboxtypes.State) ([]sandboxtypes.Sandbox, error) { - return s.getItems(&teamID, states), nil -} - -func (s *Storage) TeamsWithSandboxCount(_ context.Context) (map[uuid.UUID]int64, error) { - teams := make(map[uuid.UUID]int64) - - s.items.IterCb(func(_ string, item *memorySandbox) { - teams[item._data.TeamID]++ - }) - - return teams, nil -} - -func (s *Storage) ExpiredItems(_ context.Context) ([]sandboxtypes.Sandbox, error) { - now := time.Now() - expired := make([]sandboxtypes.Sandbox, 0) - - s.items.IterCb(func(_ string, item *memorySandbox) { - sbx := item.Data() - if sbx.State != sandboxtypes.StateRunning { - return - } - - if sbx.IsExpired(now) { - expired = append(expired, sbx) - } - }) - - return expired, nil -} - -func (s *Storage) Update(_ context.Context, teamID uuid.UUID, sandboxID string, updateFunc func(sandboxtypes.Sandbox) (sandboxtypes.Sandbox, error)) (sandboxtypes.Sandbox, error) { - item, ok := s.items.Get(sandboxID) - if !ok { - return sandboxtypes.Sandbox{}, fmt.Errorf("sandbox %q: %w", sandboxID, sandboxtypes.ErrNotFound) - } - - item.mu.Lock() - defer item.mu.Unlock() - - if item._data.TeamID != teamID { - return sandboxtypes.Sandbox{}, fmt.Errorf("sandbox %q: %w", sandboxID, sandboxtypes.ErrNotFound) - } - - sbx, err := updateFunc(item._data) - if err != nil { - return sandboxtypes.Sandbox{}, err - } - - item._data = sbx - - return sbx, nil -} - -func (s *Storage) StartRemoving(ctx context.Context, teamID uuid.UUID, sandboxID string, opts sandboxtypes.RemoveOpts) (sandboxtypes.Sandbox, bool, func(context.Context, error), error) { - sbx, err := s.get(sandboxID) - if err != nil { - return sandboxtypes.Sandbox{}, false, nil, fmt.Errorf("sandbox %q: %w", sandboxID, sandboxtypes.ErrNotFound) - } - - data := sbx.Data() - if data.TeamID != teamID { - return sandboxtypes.Sandbox{}, false, nil, fmt.Errorf("sandbox %q: %w", sandboxID, sandboxtypes.ErrNotFound) - } - - alreadyDone, callback, err := startRemoving(ctx, sbx, opts) - - return sbx.Data(), alreadyDone, callback, err -} - -func startRemoving(ctx context.Context, sbx *memorySandbox, opts sandboxtypes.RemoveOpts) (alreadyDone bool, callback func(ctx context.Context, err error), err error) { - sbx.mu.Lock() - transition := sbx.transition - - // Resolve eviction under the lock + re-check expiry - if opts.Eviction { - // If there's a transition already in place, don't evict. - if transition != nil { - sbx.mu.Unlock() - - return false, nil, sandboxtypes.ErrEvictionInProgress - } - - // If sandbox isn't expired (e.g. race condition with KeepAliveFor), skip. - if !sbx._data.IsExpired(time.Now()) { - sbx.mu.Unlock() - - return false, nil, sandboxtypes.ErrEvictionNotNeeded - } - } - - newState := opts.Action.TargetState - - if transition != nil { - currentState := sbx._data.State - sbx.mu.Unlock() - - if currentState != newState && !sandboxtypes.AllowedTransitions[currentState][newState] { - return false, nil, &sandboxtypes.InvalidStateTransitionError{CurrentState: currentState, TargetState: newState} - } - - if currentState == newState { - // The caller will inherit the in-flight transition's result - // without doing the work itself: this is a joiner. Mark before - // waiting so the request stays tagged even if the inherited - // transition fails. - joined.Mark(ctx) - } - - logger.L().Debug(ctx, "State transition already in progress to the same state, waiting", logger.WithSandboxID(sbx.SandboxID()), zap.String("state", string(newState))) - err = transition.WaitWithContext(ctx) - if err != nil { - return false, nil, fmt.Errorf("sandbox is in failed state: %w", err) - } - - // If the transition is to the same state just wait - switch { - case currentState == newState: - return true, func(context.Context, error) {}, nil - case sandboxtypes.AllowedTransitions[currentState][newState]: - return startRemoving(ctx, sbx, sandboxtypes.RemoveOpts{Action: opts.Action}) - default: - return false, nil, errors.New("unexpected state transition") - } - } - - defer sbx.mu.Unlock() - if sbx._data.State == newState { - logger.L().Debug(ctx, "Already in the same state", logger.WithSandboxID(sbx.SandboxID()), zap.String("state", string(newState))) - - return true, func(context.Context, error) {}, nil - } - - if _, ok := sandboxtypes.AllowedTransitions[sbx._data.State][newState]; !ok { - return false, nil, &sandboxtypes.InvalidStateTransitionError{CurrentState: sbx._data.State, TargetState: newState} - } - - if opts.Action.Effect == sandboxtypes.TransitionExpires { - sbx.setExpired() - } - - sbx._data.State = newState - sbx.transition = utils.NewErrorOnce() - - callback = func(ctx context.Context, err error) { - logger.L().Debug(ctx, "Transition complete", logger.WithSandboxID(sbx.SandboxID()), zap.String("state", string(newState)), zap.Error(err)) - sbx.mu.Lock() - defer sbx.mu.Unlock() - - if opts.Action.Effect == sandboxtypes.TransitionTransient { - if err == nil && sbx._data.State == newState { - sbx._data.State = sandboxtypes.StateRunning - } - - // Signal nil to waiters so concurrent callers (e.g. kill) - // are unblocked and can proceed with their own transition. - err = nil - } - - setErr := sbx.transition.SetError(err) - if setErr != nil { - logger.L().Warn(ctx, "Failed to set transition result", logger.WithSandboxID(sbx.SandboxID()), zap.Error(setErr)) - } - - if err != nil { - // Keep the transition in place so the error stays - return - } - - // The transition is completed and the next transition can be started - sbx.transition = nil - } - - return false, callback, nil -} - -func (s *Storage) WaitForStateChange(ctx context.Context, _ uuid.UUID, sandboxID string) error { - sbx, err := s.get(sandboxID) - if err != nil { - return fmt.Errorf("failed to get sandbox: %w", err) - } - - return waitForStateChange(ctx, sbx) -} - -func waitForStateChange(ctx context.Context, sbx *memorySandbox) error { - sbx.mu.RLock() - transition := sbx.transition - sbx.mu.RUnlock() - if transition == nil { - return nil - } - - return transition.WaitWithContext(ctx) -} diff --git a/packages/api/internal/sandbox/storage/memory/operations_benchmark_test.go b/packages/api/internal/sandbox/storage/memory/operations_benchmark_test.go deleted file mode 100644 index 01406afa41..0000000000 --- a/packages/api/internal/sandbox/storage/memory/operations_benchmark_test.go +++ /dev/null @@ -1,136 +0,0 @@ -package memory - -import ( - "fmt" - "testing" - "time" - - "github.com/google/uuid" - - "github.com/e2b-dev/infra/packages/api/internal/sandbox/sandboxtypes" -) - -type benchFixture struct { - storage *Storage - teamIDs []uuid.UUID - runningTeam uuid.UUID - syncNodeID string - syncInput []sandboxtypes.Sandbox -} - -func buildFixture(total int) benchFixture { - now := time.Now() - s := NewStorage() - - const teamCount = 128 - teamIDs := make([]uuid.UUID, teamCount) - for i := range teamCount { - teamIDs[i] = uuid.New() - } - - const nodeCount = 16 - nodeIDs := make([]string, nodeCount) - for i := range nodeCount { - nodeIDs[i] = fmt.Sprintf("node-%02d", i) - } - - runningTeam := teamIDs[0] - syncNodeID := nodeIDs[0] - syncInput := make([]sandboxtypes.Sandbox, 0, total/nodeCount+1) - - for i := range total { - teamID := teamIDs[i%teamCount] - nodeID := nodeIDs[i%nodeCount] - - state := sandboxtypes.StateRunning - if i%5 == 0 { - state = sandboxtypes.StatePausing - } - - endTime := now.Add(1 * time.Hour) - // Keep a stable 5% expired-running subset for ExpiredItems benchmarks. - if i%20 == 0 { - state = sandboxtypes.StateRunning - endTime = now.Add(-1 * time.Minute) - } - - sbx := sandboxtypes.Sandbox{ - SandboxID: fmt.Sprintf("sbx-%06d", i), - TeamID: teamID, - NodeID: nodeID, - State: state, - StartTime: now.Add(-1 * time.Hour), - EndTime: endTime, - } - - s.items.Set(sbx.SandboxID, newMemorySandbox(sbx)) - - // Feed Sync with full coverage for one node to avoid mutations during bench. - if nodeID == syncNodeID { - syncInput = append(syncInput, sbx) - } - } - - return benchFixture{ - storage: s, - teamIDs: teamIDs, - runningTeam: runningTeam, - syncNodeID: syncNodeID, - syncInput: syncInput, - } -} - -func benchmarkSizes(b *testing.B, fn func(b *testing.B, f benchFixture)) { - b.Helper() - - for _, size := range []int{5000, 10000, 25000, 50000} { - b.Run(fmt.Sprintf("items=%d", size), func(b *testing.B) { - fixture := buildFixture(size) - b.ReportAllocs() - b.ResetTimer() - fn(b, fixture) - }) - } -} - -func BenchmarkStorageGetItemsRunningByTeam(b *testing.B) { - benchmarkSizes(b, func(b *testing.B, f benchFixture) { - b.Helper() - - for range b.N { - _ = f.storage.getItems(&f.runningTeam, []sandboxtypes.State{sandboxtypes.StateRunning}) - } - }) -} - -func BenchmarkStorageExpiredItems(b *testing.B) { - ctx := b.Context() - benchmarkSizes(b, func(b *testing.B, f benchFixture) { - b.Helper() - - for range b.N { - _, _ = f.storage.ExpiredItems(ctx) - } - }) -} - -func BenchmarkStorageTeamsWithSandboxCount(b *testing.B) { - ctx := b.Context() - benchmarkSizes(b, func(b *testing.B, f benchFixture) { - b.Helper() - - for range b.N { - _, _ = f.storage.TeamsWithSandboxCount(ctx) - } - }) -} - -func BenchmarkStorageSyncRemoveScan(b *testing.B) { - benchmarkSizes(b, func(b *testing.B, f benchFixture) { - b.Helper() - - for range b.N { - _ = f.storage.Reconcile(b.Context(), f.syncInput, f.syncNodeID) - } - }) -} diff --git a/packages/api/internal/sandbox/storage/memory/operations_test.go b/packages/api/internal/sandbox/storage/memory/operations_test.go deleted file mode 100644 index c5f8861523..0000000000 --- a/packages/api/internal/sandbox/storage/memory/operations_test.go +++ /dev/null @@ -1,896 +0,0 @@ -package memory - -import ( - "context" - "errors" - "math/rand" - "sync" - "sync/atomic" - "testing" - "time" - - "github.com/google/uuid" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" - - "github.com/e2b-dev/infra/packages/api/internal/sandbox/sandboxtypes" -) - -func createTestSandbox() *memorySandbox { - return newMemorySandbox(sandboxtypes.Sandbox{ - SandboxID: "test-sandbox", - TemplateID: "test-template", - ClientID: "test-client", - TeamID: uuid.New(), - StartTime: time.Now(), - EndTime: time.Now().Add(time.Hour), - MaxInstanceLength: time.Hour, - State: sandboxtypes.StateRunning, - }) -} - -// Test basic state transitions -func TestStartRemoving_BasicTransitions(t *testing.T) { - t.Parallel() - - tests := []struct { - name string - fromState sandboxtypes.State - stateAction sandboxtypes.StateAction - expState sandboxtypes.State - shouldError bool - }{ - {"Running to Paused", sandboxtypes.StateRunning, sandboxtypes.StateActionPause, sandboxtypes.StatePausing, false}, - {"Running to Killed", sandboxtypes.StateRunning, sandboxtypes.StateActionKill, sandboxtypes.StateKilling, false}, - {"Running to Snapshotting", sandboxtypes.StateRunning, sandboxtypes.StateActionSnapshot, sandboxtypes.StateSnapshotting, false}, - {"Paused to Killed", sandboxtypes.StatePausing, sandboxtypes.StateActionKill, sandboxtypes.StateKilling, false}, - {"Killed to Paused (invalid)", sandboxtypes.StateKilling, sandboxtypes.StateActionPause, sandboxtypes.StatePausing, true}, - {"Killed to Killed (same)", sandboxtypes.StateKilling, sandboxtypes.StateActionKill, sandboxtypes.StateKilling, false}, - {"Paused to Paused (same)", sandboxtypes.StatePausing, sandboxtypes.StateActionPause, sandboxtypes.StatePausing, false}, - {"Snapshotting to Killed", sandboxtypes.StateSnapshotting, sandboxtypes.StateActionKill, sandboxtypes.StateKilling, false}, - {"Snapshotting to Paused", sandboxtypes.StateSnapshotting, sandboxtypes.StateActionPause, sandboxtypes.StatePausing, false}, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - t.Parallel() - - sbx := createTestSandbox() - sbx._data.State = tt.fromState - ctx := t.Context() - - alreadyDone, finish, err := startRemoving(ctx, sbx, sandboxtypes.RemoveOpts{Action: tt.stateAction}) - - switch { - case tt.shouldError: - require.Error(t, err) - assert.False(t, alreadyDone) - assert.Nil(t, finish) - assert.Equal(t, tt.fromState, sbx.State()) // State unchanged - case tt.fromState == tt.expState: - require.NoError(t, err) - assert.True(t, alreadyDone) - assert.NotNil(t, finish) - assert.Equal(t, tt.fromState, sbx.State()) - default: - require.NoError(t, err) - assert.False(t, alreadyDone) - assert.NotNil(t, finish) - assert.Equal(t, tt.expState, sbx.State()) // State changed immediately - finish(ctx, nil) // Complete the transition - } - }) - } -} - -func TestStartRemoving_PauseThenKill(t *testing.T) { - t.Parallel() - sbx := createTestSandbox() - ctx := t.Context() - - // Simulate a pause operation that takes time - alreadyDone, finish, err := startRemoving(ctx, sbx, sandboxtypes.RemoveOpts{Action: sandboxtypes.StateActionPause}) - require.NoError(t, err) - assert.False(t, alreadyDone) - require.NotNil(t, finish) - - // The state should be changed immediately - assert.Equal(t, sandboxtypes.StatePausing, sbx.State()) - - // Simulate the actual pause operation taking time - started := make(chan struct{}) - - go func() { - started <- struct{}{} - time.Sleep(100 * time.Millisecond) - // The state should still be Paused - assert.Equal(t, sandboxtypes.StatePausing, sbx.State()) - finish(ctx, nil) - }() - - // Meanwhile, another request tries to kill the sandbox - <-started // Ensure the pause operation has started - - start := time.Now() - alreadyDone2, finish2, err2 := startRemoving(ctx, sbx, sandboxtypes.RemoveOpts{Action: sandboxtypes.StateActionKill}) - elapsed := time.Since(start) - - // Should have waited for the pause to complete - assert.Greater(t, elapsed, 80*time.Millisecond) - require.NoError(t, err2) - assert.False(t, alreadyDone2) - assert.NotNil(t, finish2) - assert.Equal(t, sandboxtypes.StateKilling, sbx.State()) - - // Complete the kill operation - finish2(ctx, nil) - assert.Equal(t, sandboxtypes.StateKilling, sbx.State()) -} - -// Test concurrent requests to transition to the same state (idempotency) -func TestStartRemoving_ConcurrentSameState(t *testing.T) { - t.Parallel() - - sbx := createTestSandbox() - ctx := t.Context() - - results := make(chan struct { - alreadyDone bool - worked bool - }, 3) - - // Three concurrent requests to pause the sandbox - for range 3 { - go func() { - alreadyDone, finish, err := startRemoving(ctx, sbx, sandboxtypes.RemoveOpts{Action: sandboxtypes.StateActionPause}) - if err == nil { - if alreadyDone { - // Already alreadyDone (waited for another transition) - results <- struct { - alreadyDone bool - worked bool - }{alreadyDone, false} - } else { - // We got to perform the transition - time.Sleep(10 * time.Millisecond) - finish(ctx, nil) - results <- struct { - alreadyDone bool - worked bool - }{alreadyDone, true} - } - } else { - results <- struct { - alreadyDone bool - worked bool - }{false, false} - } - }() - } - - // Collect results - performedCount := 0 - alreadyDoneCount := 0 - for range 3 { - result := <-results - if result.worked { - performedCount++ - } - if result.alreadyDone { - alreadyDoneCount++ - } - } - - // Only one should have actually performed the transition (worked) - // But others waiting should get alreadyDone=true after the transition completes - assert.Equal(t, 1, performedCount, "Only one request should actually perform the transition") - assert.Equal(t, 2, alreadyDoneCount, "Two concurrent requests should see it's already alreadyDone") - assert.Equal(t, sandboxtypes.StatePausing, sbx.State()) -} - -// Test transition fails and subsequent request handles it -func TestStartRemoving_Error(t *testing.T) { - t.Parallel() - - sbx := createTestSandbox() - ctx := t.Context() - - // First attempt to pause - alreadyDone1, finish1, err := startRemoving(ctx, sbx, sandboxtypes.RemoveOpts{Action: sandboxtypes.StateActionPause}) - require.NoError(t, err) - assert.False(t, alreadyDone1) - require.NotNil(t, finish1) - - // Start a concurrent request that will wait for the first transition - var alreadyDone2 bool - var err2 error - var finish2 func(context.Context, error) - completed := make(chan bool) - - go func() { - // This should wait for the first transition, then try to go to Killed - alreadyDone2, finish2, err2 = startRemoving(ctx, sbx, sandboxtypes.RemoveOpts{Action: sandboxtypes.StateActionKill}) - completed <- true - }() - - // Give the goroutine time to start waiting - time.Sleep(10 * time.Millisecond) - - // Complete first transition with error - failureErr := errors.New("network timeout") - finish1(ctx, failureErr) - - // Wait for second request to complete - <-completed - - // The waiting request should have received the error from the first transition - require.Error(t, err2) - assert.Contains(t, err2.Error(), failureErr.Error()) - assert.False(t, alreadyDone2) - assert.Nil(t, finish2) - - // From Failed state, no transitions are allowed - alreadyDone3, finish3, err3 := startRemoving(ctx, sbx, sandboxtypes.RemoveOpts{Action: sandboxtypes.StateActionPause}) - require.Error(t, err3) - require.ErrorIs(t, err3, failureErr) - assert.False(t, alreadyDone3) - assert.Nil(t, finish3) - - // Trying to transition to Killed should also fail - alreadyDone4, finish4, err4 := startRemoving(ctx, sbx, sandboxtypes.RemoveOpts{Action: sandboxtypes.StateActionKill}) - require.Error(t, err4) - require.ErrorIs(t, err4, failureErr) - assert.False(t, alreadyDone4) - assert.Nil(t, finish4) -} - -// Test context timeout during wait -func TestStartRemoving_ContextTimeout(t *testing.T) { - t.Parallel() - - sbx := createTestSandbox() - - // Start a long-running transition - alreadyDone1, finish1, err := startRemoving(t.Context(), sbx, sandboxtypes.RemoveOpts{Action: sandboxtypes.StateActionPause}) - require.NoError(t, err) - assert.False(t, alreadyDone1) - require.NotNil(t, finish1) - - // Another request with a short timeout - ctx, cancel := context.WithTimeout(t.Context(), 20*time.Millisecond) - defer cancel() - - start := time.Now() - _, _, err2 := startRemoving(ctx, sbx, sandboxtypes.RemoveOpts{Action: sandboxtypes.StateActionKill}) - elapsed := time.Since(start) - - // Should timeout after about 20ms - require.ErrorIs(t, err2, context.DeadlineExceeded) - assert.Greater(t, elapsed, 15*time.Millisecond) - assert.Less(t, elapsed, 500*time.Millisecond) - - // Clean up - finish1(ctx, nil) - assert.Equal(t, sandboxtypes.StatePausing, sbx.State()) -} - -func TestWaitForStateChange_NoTransition(t *testing.T) { - t.Parallel() - sbx := createTestSandbox() - ctx := t.Context() - - // Should work even with canceled context - no wait needed - ctx, cancel := context.WithCancel(ctx) - cancel() - - // No transition in progress, no need to wait - err := waitForStateChange(ctx, sbx) - require.NoError(t, err) -} - -func TestWaitForStateChange_WaitForCompletion(t *testing.T) { - t.Parallel() - sbx := createTestSandbox() - ctx := t.Context() - - // Start a transition - alreadyalreadyDone, finish, err := startRemoving(ctx, sbx, sandboxtypes.RemoveOpts{Action: sandboxtypes.StateActionPause}) - require.NoError(t, err) - assert.False(t, alreadyalreadyDone) - require.NotNil(t, finish) - - // Wait for state change in a goroutine - var waitErr error - alreadyDone := make(chan bool) - - go func() { - waitErr = waitForStateChange(ctx, sbx) - alreadyDone <- true - }() - - // Give the goroutine time to start waiting - time.Sleep(10 * time.Millisecond) - - // Complete the transition - finish(ctx, nil) - - // Wait should complete - <-alreadyDone - require.NoError(t, waitErr) -} - -func TestWaitForStateChange_WaitWithError(t *testing.T) { - t.Parallel() - sbx := createTestSandbox() - ctx := t.Context() - - // Start a transition - alreadyalreadyDone, finish, err := startRemoving(ctx, sbx, sandboxtypes.RemoveOpts{Action: sandboxtypes.StateActionPause}) - require.NoError(t, err) - assert.False(t, alreadyalreadyDone) - require.NotNil(t, finish) - - // Wait for state change in a goroutine - var waitErr error - alreadyDone := make(chan bool) - - go func() { - waitErr = waitForStateChange(ctx, sbx) - alreadyDone <- true - }() - - // Give the goroutine time to start waiting - time.Sleep(10 * time.Millisecond) - - // Complete the transition with error - testErr := assert.AnError - finish(ctx, testErr) - - // Wait should complete with error - <-alreadyDone - require.Error(t, waitErr) - assert.Equal(t, testErr, waitErr) -} - -func TestWaitForStateChange_ContextCancellation(t *testing.T) { - t.Parallel() - sbx := createTestSandbox() - ctx, cancel := context.WithCancel(t.Context()) - - // Start a transition - alreadyalreadyDone, finish, err := startRemoving(ctx, sbx, sandboxtypes.RemoveOpts{Action: sandboxtypes.StateActionPause}) - require.NoError(t, err) - assert.False(t, alreadyalreadyDone) - require.NotNil(t, finish) - - // Wait for state change in a goroutine - var waitErr error - alreadyDone := make(chan bool) - - go func() { - waitErr = waitForStateChange(ctx, sbx) - alreadyDone <- true - }() - - // Give the goroutine time to start waiting - time.Sleep(10 * time.Millisecond) - - // Cancel the context - cancel() - - // Wait should complete with context error - <-alreadyDone - require.Error(t, waitErr) - assert.Equal(t, context.Canceled, waitErr) - - // Clean up - complete the transition - finish(ctx, nil) -} - -func TestWaitForStateChange_MultipleWaiters(t *testing.T) { - t.Parallel() - sbx := createTestSandbox() - ctx := t.Context() - - // Start a transition - alreadyalreadyDone, finish, err := startRemoving(ctx, sbx, sandboxtypes.RemoveOpts{Action: sandboxtypes.StateActionPause}) - require.NoError(t, err) - assert.False(t, alreadyalreadyDone) - require.NotNil(t, finish) - - // Start multiple waiters - numWaiters := 5 - errs := make([]error, numWaiters) - var wg sync.WaitGroup - - for i := range numWaiters { - wg.Add(1) - go func(idx int) { - defer wg.Done() - errs[idx] = waitForStateChange(ctx, sbx) - }(i) - } - - // Give the goroutines time to start waiting - time.Sleep(10 * time.Millisecond) - - // Complete the transition - finish(ctx, nil) - - // Wait for all waiters to complete - wg.Wait() - - // All waiters should complete successfully - for i := range numWaiters { - require.NoError(t, errs[i]) - } -} - -func TestStartRemoving_DuringSnapshotting(t *testing.T) { - t.Parallel() - - t.Run("pause waits for snapshotting then succeeds", func(t *testing.T) { - t.Parallel() - - ctx := t.Context() - storage := NewStorage() - - sbx := sandboxtypes.Sandbox{ - SandboxID: "snap-pause-test", - TemplateID: "test-template", - ClientID: "test-client", - TeamID: uuid.New(), - StartTime: time.Now(), - EndTime: time.Now().Add(time.Hour), - MaxInstanceLength: time.Hour, - State: sandboxtypes.StateRunning, - } - - err := storage.Add(ctx, sbx) - require.NoError(t, err) - - _, snapAlreadyDone, finishSnap, err := storage.StartRemoving(ctx, sbx.TeamID, sbx.SandboxID, sandboxtypes.RemoveOpts{Action: sandboxtypes.StateActionSnapshot}) - require.NoError(t, err) - assert.False(t, snapAlreadyDone) - require.NotNil(t, finishSnap) - - pauseDone := make(chan struct{}) - var pauseErr error - var pauseAlreadyDone bool - var pauseFinish func(context.Context, error) - - go func() { - defer close(pauseDone) - _, pauseAlreadyDone, pauseFinish, pauseErr = storage.StartRemoving(ctx, sbx.TeamID, sbx.SandboxID, sandboxtypes.RemoveOpts{Action: sandboxtypes.StateActionPause}) - }() - - time.Sleep(50 * time.Millisecond) - - got, getErr := storage.Get(ctx, sbx.TeamID, sbx.SandboxID) - require.NoError(t, getErr) - assert.Equal(t, sandboxtypes.StateSnapshotting, got.State) - - finishSnap(ctx, nil) - - <-pauseDone - - require.NoError(t, pauseErr) - assert.False(t, pauseAlreadyDone) - require.NotNil(t, pauseFinish) - - got, getErr = storage.Get(ctx, sbx.TeamID, sbx.SandboxID) - require.NoError(t, getErr) - assert.Equal(t, sandboxtypes.StatePausing, got.State) - - pauseFinish(ctx, nil) - }) - - t.Run("kill waits for snapshotting then succeeds", func(t *testing.T) { - t.Parallel() - - ctx := t.Context() - storage := NewStorage() - - sbx := sandboxtypes.Sandbox{ - SandboxID: "snap-kill-test", - TemplateID: "test-template", - ClientID: "test-client", - TeamID: uuid.New(), - StartTime: time.Now(), - EndTime: time.Now().Add(time.Hour), - MaxInstanceLength: time.Hour, - State: sandboxtypes.StateRunning, - } - - err := storage.Add(ctx, sbx) - require.NoError(t, err) - - _, snapAlreadyDone, finishSnap, err := storage.StartRemoving(ctx, sbx.TeamID, sbx.SandboxID, sandboxtypes.RemoveOpts{Action: sandboxtypes.StateActionSnapshot}) - require.NoError(t, err) - assert.False(t, snapAlreadyDone) - - // Kill blocks waiting for snapshotting to complete - killDone := make(chan struct{}) - var killErr error - var killAlreadyDone bool - var killFinish func(context.Context, error) - - go func() { - defer close(killDone) - _, killAlreadyDone, killFinish, killErr = storage.StartRemoving(ctx, sbx.TeamID, sbx.SandboxID, sandboxtypes.RemoveOpts{Action: sandboxtypes.StateActionKill}) - }() - - // Give the kill goroutine time to start waiting - time.Sleep(50 * time.Millisecond) - - // Verify the sandbox is still snapshotting (kill is blocked) - got, getErr := storage.Get(ctx, sbx.TeamID, sbx.SandboxID) - require.NoError(t, getErr) - assert.Equal(t, sandboxtypes.StateSnapshotting, got.State) - - // Complete snapshotting — this unblocks the kill - finishSnap(ctx, nil) - - <-killDone - - require.NoError(t, killErr) - assert.False(t, killAlreadyDone) - require.NotNil(t, killFinish) - - got, getErr = storage.Get(ctx, sbx.TeamID, sbx.SandboxID) - require.NoError(t, getErr) - assert.Equal(t, sandboxtypes.StateKilling, got.State) - - killFinish(ctx, nil) - }) - - t.Run("kill after failed snapshotting proceeds from Snapshotting state", func(t *testing.T) { - t.Parallel() - - ctx := t.Context() - storage := NewStorage() - - sbx := sandboxtypes.Sandbox{ - SandboxID: "snap-fail-kill-test", - TemplateID: "test-template", - ClientID: "test-client", - TeamID: uuid.New(), - StartTime: time.Now(), - EndTime: time.Now().Add(time.Hour), - MaxInstanceLength: time.Hour, - State: sandboxtypes.StateRunning, - } - - err := storage.Add(ctx, sbx) - require.NoError(t, err) - - _, _, finishSnap, err := storage.StartRemoving(ctx, sbx.TeamID, sbx.SandboxID, sandboxtypes.RemoveOpts{Action: sandboxtypes.StateActionSnapshot}) - require.NoError(t, err) - - // Finish with error — state stays Snapshotting, transition cleared - finishSnap(ctx, errors.New("checkpoint failed")) - - got, getErr := storage.Get(ctx, sbx.TeamID, sbx.SandboxID) - require.NoError(t, getErr) - assert.Equal(t, sandboxtypes.StateSnapshotting, got.State) - - // Kill proceeds immediately — no active transition, Snapshotting→Killing is allowed - _, killAlreadyDone, killFinish, killErr := storage.StartRemoving(ctx, sbx.TeamID, sbx.SandboxID, sandboxtypes.RemoveOpts{Action: sandboxtypes.StateActionKill}) - require.NoError(t, killErr) - assert.False(t, killAlreadyDone) - require.NotNil(t, killFinish) - - got, getErr = storage.Get(ctx, sbx.TeamID, sbx.SandboxID) - require.NoError(t, getErr) - assert.Equal(t, sandboxtypes.StateKilling, got.State) - - killFinish(ctx, nil) - }) - - t.Run("resume sees snapshotting state", func(t *testing.T) { - t.Parallel() - - ctx := t.Context() - storage := NewStorage() - - sbx := sandboxtypes.Sandbox{ - SandboxID: "snap-resume-test", - TemplateID: "test-template", - ClientID: "test-client", - TeamID: uuid.New(), - StartTime: time.Now(), - EndTime: time.Now().Add(time.Hour), - MaxInstanceLength: time.Hour, - State: sandboxtypes.StateRunning, - } - - err := storage.Add(ctx, sbx) - require.NoError(t, err) - - _, snapAlreadyDone, finishSnap, err := storage.StartRemoving(ctx, sbx.TeamID, sbx.SandboxID, sandboxtypes.RemoveOpts{Action: sandboxtypes.StateActionSnapshot}) - require.NoError(t, err) - assert.False(t, snapAlreadyDone) - - got, getErr := storage.Get(ctx, sbx.TeamID, sbx.SandboxID) - require.NoError(t, getErr) - assert.Equal(t, sandboxtypes.StateSnapshotting, got.State) - - finishSnap(ctx, nil) - }) -} - -// Eviction-specific tests: verify the Eviction flag in RemoveOpts -// re-checks expiry and transition state under the lock. -func TestStartRemoving_Eviction(t *testing.T) { - t.Parallel() - - t.Run("expired sandbox with no transition is evicted", func(t *testing.T) { - t.Parallel() - - ctx := t.Context() - storage := NewStorage() - - sbx := sandboxtypes.Sandbox{ - SandboxID: "evict-ok", - TemplateID: "test-template", - ClientID: "test-client", - TeamID: uuid.New(), - StartTime: time.Now().Add(-2 * time.Hour), - EndTime: time.Now().Add(-time.Second), // already expired - MaxInstanceLength: time.Hour, - State: sandboxtypes.StateRunning, - } - - err := storage.Add(ctx, sbx) - require.NoError(t, err) - - _, alreadyDone, finish, err := storage.StartRemoving(ctx, sbx.TeamID, sbx.SandboxID, sandboxtypes.RemoveOpts{Action: sandboxtypes.StateActionKill, Eviction: true}) - require.NoError(t, err) - assert.False(t, alreadyDone) - require.NotNil(t, finish) - - got, getErr := storage.Get(ctx, sbx.TeamID, sbx.SandboxID) - require.NoError(t, getErr) - assert.Equal(t, sandboxtypes.StateKilling, got.State) - - finish(ctx, nil) - }) - - t.Run("non-expired sandbox returns eviction not needed", func(t *testing.T) { - t.Parallel() - - ctx := t.Context() - storage := NewStorage() - - sbx := sandboxtypes.Sandbox{ - SandboxID: "evict-not-expired", - TemplateID: "test-template", - ClientID: "test-client", - TeamID: uuid.New(), - StartTime: time.Now(), - EndTime: time.Now().Add(time.Hour), // not expired - MaxInstanceLength: time.Hour, - State: sandboxtypes.StateRunning, - } - - err := storage.Add(ctx, sbx) - require.NoError(t, err) - - _, alreadyDone, finish, err := storage.StartRemoving(ctx, sbx.TeamID, sbx.SandboxID, sandboxtypes.RemoveOpts{Action: sandboxtypes.StateActionKill, Eviction: true}) - require.ErrorIs(t, err, sandboxtypes.ErrEvictionNotNeeded) - assert.False(t, alreadyDone) - assert.Nil(t, finish) - - // State must remain Running — sandbox was not touched. - got, getErr := storage.Get(ctx, sbx.TeamID, sbx.SandboxID) - require.NoError(t, getErr) - assert.Equal(t, sandboxtypes.StateRunning, got.State) - }) - - t.Run("expired sandbox with active transition returns eviction in progress", func(t *testing.T) { - t.Parallel() - - ctx := t.Context() - storage := NewStorage() - - sbx := sandboxtypes.Sandbox{ - SandboxID: "evict-in-transition", - TemplateID: "test-template", - ClientID: "test-client", - TeamID: uuid.New(), - StartTime: time.Now().Add(-2 * time.Hour), - EndTime: time.Now().Add(-time.Second), // expired - MaxInstanceLength: time.Hour, - State: sandboxtypes.StateRunning, - } - - err := storage.Add(ctx, sbx) - require.NoError(t, err) - - // Start a non-eviction pause transition to occupy the transition slot. - _, _, pauseFinish, err := storage.StartRemoving(ctx, sbx.TeamID, sbx.SandboxID, sandboxtypes.RemoveOpts{Action: sandboxtypes.StateActionPause}) - require.NoError(t, err) - require.NotNil(t, pauseFinish) - - // Eviction should be rejected immediately (not block). - start := time.Now() - _, alreadyDone, finish, evictErr := storage.StartRemoving(ctx, sbx.TeamID, sbx.SandboxID, sandboxtypes.RemoveOpts{Action: sandboxtypes.StateActionKill, Eviction: true}) - elapsed := time.Since(start) - - require.ErrorIs(t, evictErr, sandboxtypes.ErrEvictionInProgress) - assert.False(t, alreadyDone) - assert.Nil(t, finish) - assert.Less(t, elapsed, 50*time.Millisecond, "eviction should return immediately, not wait for the transition") - - // Clean up - pauseFinish(ctx, nil) - }) - - t.Run("expired sandbox evicted with auto-pause action", func(t *testing.T) { - t.Parallel() - - ctx := t.Context() - storage := NewStorage() - - sbx := sandboxtypes.Sandbox{ - SandboxID: "evict-autopause", - TemplateID: "test-template", - ClientID: "test-client", - TeamID: uuid.New(), - StartTime: time.Now().Add(-2 * time.Hour), - EndTime: time.Now().Add(-time.Second), // expired - MaxInstanceLength: time.Hour, - AutoPause: true, - State: sandboxtypes.StateRunning, - } - - err := storage.Add(ctx, sbx) - require.NoError(t, err) - - _, alreadyDone, finish, err := storage.StartRemoving(ctx, sbx.TeamID, sbx.SandboxID, sandboxtypes.RemoveOpts{Action: sandboxtypes.StateActionPause, Eviction: true}) - require.NoError(t, err) - assert.False(t, alreadyDone) - require.NotNil(t, finish) - - got, getErr := storage.Get(ctx, sbx.TeamID, sbx.SandboxID) - require.NoError(t, getErr) - assert.Equal(t, sandboxtypes.StatePausing, got.State) - - finish(ctx, nil) - }) - - t.Run("eviction flag is not propagated on retry after waiting", func(t *testing.T) { - t.Parallel() - - // This tests that when a non-eviction request waits for an - // existing transition and retries, the retry path works correctly - // even when the sandbox EndTime has been extended mid-flight. - // An eviction in the same situation would bail out, but a regular - // kill must proceed regardless of expiry. - ctx := t.Context() - - sbx := createTestSandbox() - // Make sandbox expired so the initial state matches what the - // evictor would have seen before adding it to the eviction list. - sbx._data.EndTime = time.Now().Add(-time.Second) - - // Start a non-eviction pause. - alreadyDone, pauseFinish, err := startRemoving(ctx, sbx, sandboxtypes.RemoveOpts{Action: sandboxtypes.StateActionPause}) - require.NoError(t, err) - assert.False(t, alreadyDone) - require.NotNil(t, pauseFinish) - - // A non-eviction kill will wait for the pause, then retry. - killDone := make(chan struct{}) - var killErr error - var killAlreadyDone bool - var killFinish func(context.Context, error) - - go func() { - defer close(killDone) - killAlreadyDone, killFinish, killErr = startRemoving(ctx, sbx, sandboxtypes.RemoveOpts{Action: sandboxtypes.StateActionKill}) - }() - - time.Sleep(50 * time.Millisecond) - - // Extend the sandbox timeout while the pause is in progress - // (simulating KeepAliveFor extending EndTime). - sbx.mu.Lock() - sbx._data.EndTime = time.Now().Add(time.Hour) - sbx.mu.Unlock() - - // Complete the pause. - pauseFinish(ctx, nil) - - <-killDone - - // The kill should succeed because it's NOT an eviction — the - // non-expired EndTime doesn't block a regular kill. - require.NoError(t, killErr) - assert.False(t, killAlreadyDone) - require.NotNil(t, killFinish) - assert.Equal(t, sandboxtypes.StateKilling, sbx.State()) - - killFinish(ctx, nil) - }) -} - -// Stress test with random operations -func TestConcurrency_StressTest(t *testing.T) { - t.Parallel() - - if testing.Short() { - t.Skip("Skipping stress test in short mode") - } - - sbx := createTestSandbox() - - duration := 100 * time.Millisecond - deadline := time.Now().Add(duration) - - var wg sync.WaitGroup - stopCh := make(chan struct{}) - - // Metrics - var opsCompleted atomic.Uint64 - var errorCount atomic.Uint64 - - // Launch workers that continuously perform random operations - for i := range 200 { - wg.Add(1) - go func(workerID int) { - defer wg.Done() - - for { - select { - case <-stopCh: - return - default: - // Random operation - switch workerID % 4 { - case 0: // State transitions - stateActions := []sandboxtypes.StateAction{sandboxtypes.StateActionPause, sandboxtypes.StateActionKill} - stateAction := stateActions[rand.Intn(len(stateActions))] - - alreadyDone, finish, err := startRemoving(t.Context(), sbx, sandboxtypes.RemoveOpts{Action: stateAction}) - if err == nil && (finish != nil || alreadyDone) { - if finish != nil { - finish(t.Context(), nil) - } - opsCompleted.Add(1) - } else if err != nil { - errorCount.Add(1) - } - case 1: // Read state - _ = sbx.State() - opsCompleted.Add(1) - case 2: // Wait with timeout - waitCtx, cancel := context.WithTimeout(t.Context(), time.Microsecond*10) - _ = waitForStateChange(waitCtx, sbx) - cancel() - opsCompleted.Add(1) - case 3: // Read _data - _ = sbx.Data() - opsCompleted.Add(1) - } - } - - if time.Now().After(deadline) { - return - } - } - }(i) - } - - // Let it run - time.Sleep(duration) - close(stopCh) - wg.Wait() - - finalOps := opsCompleted.Load() - finalErrors := errorCount.Load() - t.Logf("Stress test completed: %d operations, %d errors", finalOps, finalErrors) - - // Should have completed many operations without panic - assert.Greater(t, finalOps, uint64(100), "Should complete many operations") -} diff --git a/packages/api/internal/sandbox/storage/memory/sandbox.go b/packages/api/internal/sandbox/storage/memory/sandbox.go deleted file mode 100644 index 6ea424c501..0000000000 --- a/packages/api/internal/sandbox/storage/memory/sandbox.go +++ /dev/null @@ -1,62 +0,0 @@ -package memory - -import ( - "sync" - "time" - - "github.com/google/uuid" - - "github.com/e2b-dev/infra/packages/api/internal/sandbox/sandboxtypes" - "github.com/e2b-dev/infra/packages/shared/pkg/utils" -) - -type memorySandbox struct { - _data sandboxtypes.Sandbox - - transition *utils.ErrorOnce - mu sync.RWMutex -} - -func newMemorySandbox(data sandboxtypes.Sandbox) *memorySandbox { - return &memorySandbox{ - _data: data, - } -} - -func (i *memorySandbox) SetExpired() { - i.mu.Lock() - defer i.mu.Unlock() - - i.setExpired() -} - -func (i *memorySandbox) setExpired() { - now := time.Now() - if !i._data.IsExpired(now) { - i._data.EndTime = now - } -} - -func (i *memorySandbox) Data() sandboxtypes.Sandbox { - i.mu.RLock() - defer i.mu.RUnlock() - - return i._data -} - -func (i *memorySandbox) State() sandboxtypes.State { - i.mu.RLock() - defer i.mu.RUnlock() - - return i._data.State -} - -// SandboxID returns the sandbox ID, safe to use without lock, it's immutable -func (i *memorySandbox) SandboxID() string { - return i._data.SandboxID -} - -// TeamID returns the team ID, safe to use without lock, it's immutable -func (i *memorySandbox) TeamID() uuid.UUID { - return i._data.TeamID -} diff --git a/packages/api/internal/sandbox/storage/memory/sync.go b/packages/api/internal/sandbox/storage/memory/sync.go deleted file mode 100644 index dcce53feb3..0000000000 --- a/packages/api/internal/sandbox/storage/memory/sync.go +++ /dev/null @@ -1,71 +0,0 @@ -package memory - -import ( - "context" - "time" - - "github.com/e2b-dev/infra/packages/api/internal/sandbox/sandboxtypes" - "github.com/e2b-dev/infra/packages/shared/pkg/logger" -) - -// TODO: this should be removed once we have a better way to handle node sync -// Don't remove sandboxes that were started in the grace period on node sync -// This is to prevent remove instances that are still being started -const syncSandboxRemoveGracePeriod = 10 * time.Second - -func (s *Storage) Reconcile(ctx context.Context, sandboxes []sandboxtypes.Sandbox, nodeID string) []sandboxtypes.Sandbox { - sandboxMap := make(map[string]sandboxtypes.Sandbox) - now := time.Now() - - // Use a map for faster lookup - for _, sbx := range sandboxes { - sandboxMap[sbx.SandboxID] = sbx - } - - // Remove sandboxes that are not in Orchestrator anymore - s.items.IterCb(func(_ string, item *memorySandbox) { - data := item.Data() - if data.IsExpired(now) { - return - } - - if data.NodeID != nodeID { - return - } - - if time.Since(data.StartTime) <= syncSandboxRemoveGracePeriod { - return - } - - _, found := sandboxMap[data.SandboxID] - if !found { - logger.L().Debug( - ctx, - "sync expiring sandbox missing from node report", - logger.WithSandboxID(data.SandboxID), - logger.WithTeamID(data.TeamID.String()), - logger.WithNodeID(nodeID), - ) - item.SetExpired() - } - }) - - toBeAdded := make([]sandboxtypes.Sandbox, 0, len(sandboxes)) - // Add sandboxes that are not in the cache with the default TTL - for _, sbx := range sandboxes { - if s.exists(sbx.SandboxID) { - continue - } - - logger.L().Debug( - ctx, - "sync discovered sandbox missing from cache", - logger.WithSandboxID(sbx.SandboxID), - logger.WithTeamID(sbx.TeamID.String()), - logger.WithNodeID(nodeID), - ) - toBeAdded = append(toBeAdded, sbx) - } - - return toBeAdded -} diff --git a/packages/api/internal/sandbox/storage/populate_redis/main.go b/packages/api/internal/sandbox/storage/populate_redis/main.go deleted file mode 100644 index 2e3b9c1d1a..0000000000 --- a/packages/api/internal/sandbox/storage/populate_redis/main.go +++ /dev/null @@ -1,102 +0,0 @@ -package populate_redis - -import ( - "context" - - "github.com/google/uuid" - "go.uber.org/zap" - - "github.com/e2b-dev/infra/packages/api/internal/sandbox/sandboxtypes" - "github.com/e2b-dev/infra/packages/api/internal/sandbox/storage/memory" - "github.com/e2b-dev/infra/packages/api/internal/sandbox/storage/redis" - "github.com/e2b-dev/infra/packages/shared/pkg/logger" -) - -var _ sandboxtypes.Storage = (*PopulateRedisStorage)(nil) - -type PopulateRedisStorage struct { - memoryBackend *memory.Storage - redisBackend *redis.Storage -} - -func (m *PopulateRedisStorage) Name() string { return sandboxtypes.StorageNamePopulateRedis } - -func (m *PopulateRedisStorage) Add(ctx context.Context, sandbox sandboxtypes.Sandbox) error { - err := m.memoryBackend.Add(ctx, sandbox) - if err != nil { - return err - } - - err = m.redisBackend.Add(ctx, sandbox) - if err != nil { - logger.L().Error(ctx, "failed to add sandbox to redis", zap.Error(err)) - } - - return nil -} - -func (m *PopulateRedisStorage) Get(ctx context.Context, teamID uuid.UUID, sandboxID string) (sandboxtypes.Sandbox, error) { - return m.memoryBackend.Get(ctx, teamID, sandboxID) -} - -func (m *PopulateRedisStorage) Remove(ctx context.Context, teamID uuid.UUID, sandboxID string) error { - err := m.memoryBackend.Remove(ctx, teamID, sandboxID) - if err != nil { - return err - } - - err = m.redisBackend.Remove(ctx, teamID, sandboxID) - if err != nil { - logger.L().Error(ctx, "failed to remove sandbox from redis", zap.Error(err), logger.WithSandboxID(sandboxID)) - } - - return nil -} - -func (m *PopulateRedisStorage) TeamItems(ctx context.Context, teamID uuid.UUID, states []sandboxtypes.State) ([]sandboxtypes.Sandbox, error) { - return m.memoryBackend.TeamItems(ctx, teamID, states) -} - -func (m *PopulateRedisStorage) ExpiredItems(ctx context.Context) ([]sandboxtypes.Sandbox, error) { - return m.memoryBackend.ExpiredItems(ctx) -} - -func (m *PopulateRedisStorage) TeamsWithSandboxCount(ctx context.Context) (map[uuid.UUID]int64, error) { - return m.memoryBackend.TeamsWithSandboxCount(ctx) -} - -func (m *PopulateRedisStorage) Update(ctx context.Context, teamID uuid.UUID, sandboxID string, updateFunc func(sandbox sandboxtypes.Sandbox) (sandboxtypes.Sandbox, error)) (sandboxtypes.Sandbox, error) { - sbx, err := m.memoryBackend.Update(ctx, teamID, sandboxID, updateFunc) - if err != nil { - return sandboxtypes.Sandbox{}, err - } - - _, err = m.redisBackend.Update(ctx, teamID, sandboxID, updateFunc) - if err != nil { - logger.L().Error(ctx, "failed to update sandbox in redis", zap.Error(err), logger.WithSandboxID(sandboxID)) - } - - return sbx, nil -} - -func (m *PopulateRedisStorage) StartRemoving(ctx context.Context, teamID uuid.UUID, sandboxID string, opts sandboxtypes.RemoveOpts) (sandboxtypes.Sandbox, bool, func(context.Context, error), error) { - return m.memoryBackend.StartRemoving(ctx, teamID, sandboxID, opts) -} - -func (m *PopulateRedisStorage) WaitForStateChange(ctx context.Context, teamID uuid.UUID, sandboxID string) error { - return m.memoryBackend.WaitForStateChange(ctx, teamID, sandboxID) -} - -func (m *PopulateRedisStorage) Reconcile(ctx context.Context, sandboxes []sandboxtypes.Sandbox, nodeID string) []sandboxtypes.Sandbox { - return m.memoryBackend.Reconcile(ctx, sandboxes, nodeID) -} - -func NewStorage( - memoryStorage *memory.Storage, - redisStorage *redis.Storage, -) *PopulateRedisStorage { - return &PopulateRedisStorage{ - memoryBackend: memoryStorage, - redisBackend: redisStorage, - } -} diff --git a/packages/api/internal/sandbox/storage/redis/cleaner.go b/packages/api/internal/sandbox/storage/redis/cleaner.go deleted file mode 100644 index 89c662d390..0000000000 --- a/packages/api/internal/sandbox/storage/redis/cleaner.go +++ /dev/null @@ -1,101 +0,0 @@ -package redis - -import ( - "context" - "errors" - "fmt" - "time" - - "go.uber.org/zap" - - "github.com/e2b-dev/infra/packages/api/internal/sandbox/sandboxtypes" - "github.com/e2b-dev/infra/packages/shared/pkg/logger" -) - -const cleanerInterval = time.Minute - -// TODO: Remove once fully migrated to Redis -// -// Cleaner: -// - prunes stale entries from the two Redis sandbox indexes (`globalExpirationSet` and `globalTeamsSet`). -// - removes expired sandboxes -// -// Multi-pod safety: every operation the Cleaner triggers (ZREM/SREM of -// possibly-absent members) is idempotent. Concurrent Cleaners across pods -// produce duplicate Redis traffic, not incorrect state, so we do not take -// a distributed lock. -type Cleaner struct { - storage *Storage - tick time.Duration -} - -func NewCleaner(storage *Storage) *Cleaner { - return &Cleaner{ - storage: storage, - tick: cleanerInterval, - } -} - -// Start blocks until ctx is cancelled -func (c *Cleaner) Start(ctx context.Context) { - t := time.NewTicker(c.tick) - defer t.Stop() - - for { - select { - case <-ctx.Done(): - return - case <-t.C: - if err := c.RunOnce(ctx); err != nil && !errors.Is(err, context.Canceled) { - logger.L().Warn(ctx, "redis storage cleanup cycle failed", zap.Error(err)) - } - } - } -} - -// RunOnce performs one cleanup pass. Each sub-step is independent; a failure -// in one is logged but does not abort the other. -func (c *Cleaner) RunOnce(ctx context.Context) error { - var errs []error - - // 1. globalExpirationSet: ExpiredItems internally ZREMs members whose sandbox JSON is gone. - // 2. evictExpired removes sandboxes whose EndTime is older than StaleCutoff; - // recently expired ones are left to the evictor to avoid racing it. - expired, err := c.storage.ExpiredItems(ctx) - if err != nil { - errs = append(errs, fmt.Errorf("expiration index sweep: %w", err)) - } else { - c.evictExpired(ctx, expired) - } - - // 3. globalTeamsSet: TeamsWithSandboxCount internally ZREMs teams whose - // per-team SCARD is 0 AND whose score is older than StaleCutoff - // (operations.go:268-288). Discard the returned counts. - if _, err := c.storage.TeamsWithSandboxCount(ctx); err != nil { - errs = append(errs, fmt.Errorf("teams index sweep: %w", err)) - } - - return errors.Join(errs...) -} - -func (c *Cleaner) evictExpired(ctx context.Context, expired []sandboxtypes.Sandbox) { - if len(expired) == 0 { - return - } - - logger.L().Info(ctx, "Cleaner found expired sandboxes", zap.Int("count", len(expired))) - - for _, sbx := range expired { - if time.Since(sbx.EndTime) < sandboxtypes.StaleCutoff { - continue - } - - if rmErr := c.storage.Remove(context.WithoutCancel(ctx), sbx.TeamID, sbx.SandboxID); rmErr != nil { - logger.L().Error(ctx, "Cleaner failed to remove stale expired sandbox", - zap.Error(rmErr), - logger.WithSandboxID(sbx.SandboxID), - logger.WithTeamID(sbx.TeamID.String()), - ) - } - } -} diff --git a/packages/api/internal/sandbox/storage/redis/cleaner_test.go b/packages/api/internal/sandbox/storage/redis/cleaner_test.go deleted file mode 100644 index 0ac961c963..0000000000 --- a/packages/api/internal/sandbox/storage/redis/cleaner_test.go +++ /dev/null @@ -1,282 +0,0 @@ -package redis - -import ( - "context" - "fmt" - "sync" - "testing" - "time" - - "github.com/google/uuid" - "github.com/redis/go-redis/v9" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" - - "github.com/e2b-dev/infra/packages/api/internal/sandbox/sandboxtypes" -) - -// TestCleaner_PrunesOrphanedExpirationEntry is the smoke test for the whole -// feature: if a member sits in globalExpirationSet without a matching sandbox -// JSON key, the cleaner must ZREM it. This is the dominant leak source — -// an Add that wrote globalExpirationSet (step 1) but failed before the -// atomic SET+SADD (step 2 via addSandboxScript). -func TestCleaner_PrunesOrphanedExpirationEntry(t *testing.T) { - t.Parallel() - - storage, client := setupTestStorage(t) - ctx := t.Context() - - teamID := uuid.New().String() - sandboxID := "ghost-" + uuid.NewString() - member := expirationMember(teamID, sandboxID) - - // Plant an orphan: ZSET entry with old score, no sandbox JSON key. - require.NoError(t, client.ZAdd(ctx, globalExpirationSet, redis.Z{ - Score: float64(time.Now().Add(-time.Hour).UnixMilli()), - Member: member, - }).Err()) - - cleaner := NewCleaner(storage) - require.NoError(t, cleaner.RunOnce(ctx)) - - _, err := client.ZScore(ctx, globalExpirationSet, member).Result() - require.ErrorIs(t, err, redis.Nil, "orphan should have been ZREM'd from globalExpirationSet") -} - -// TestCleaner_PreservesLiveEntries is the negative regression guard: a real -// sandbox added through Storage.Add must survive a RunOnce in all indexes -// plus its JSON key. -func TestCleaner_PreservesLiveEntries(t *testing.T) { - t.Parallel() - - storage, client := setupTestStorage(t) - ctx := t.Context() - - sbx := createTestSandbox("live-" + uuid.NewString()) - require.NoError(t, storage.Add(ctx, sbx)) - - cleaner := NewCleaner(storage) - require.NoError(t, cleaner.RunOnce(ctx)) - - // globalExpirationSet still has it. - _, err := client.ZScore(ctx, globalExpirationSet, - expirationMember(sbx.TeamID.String(), sbx.SandboxID)).Result() - require.NoError(t, err, "live entry should remain in globalExpirationSet") - - // globalTeamsSet still has the team. - _, err = client.ZScore(ctx, globalTeamsSet, sbx.TeamID.String()).Result() - require.NoError(t, err, "live team should remain in globalTeamsSet") - - // Per-team SET still has the sandbox ID. - isMember, err := client.SIsMember(ctx, - GetSandboxStorageTeamIndexKey(sbx.TeamID.String()), sbx.SandboxID).Result() - require.NoError(t, err) - require.True(t, isMember, "live sandbox should remain in per-team index") - - // Sandbox JSON itself is untouched. - got, err := storage.Get(ctx, sbx.TeamID, sbx.SandboxID) - require.NoError(t, err) - require.Equal(t, sbx.SandboxID, got.SandboxID) -} - -// TestCleaner_PrunesEmptyOldTeam plants a stale entry in globalTeamsSet whose -// per-team SET is empty; TeamsWithSandboxCount must ZREM it. -func TestCleaner_PrunesEmptyOldTeam(t *testing.T) { - t.Parallel() - - storage, client := setupTestStorage(t) - ctx := t.Context() - - staleTeamID := uuid.New().String() - oldScore := float64(time.Now().Add(-time.Hour).Unix()) - - require.NoError(t, client.ZAdd(ctx, globalTeamsSet, redis.Z{ - Score: oldScore, - Member: staleTeamID, - }).Err()) - - cleaner := NewCleaner(storage) - require.NoError(t, cleaner.RunOnce(ctx)) - - _, err := client.ZScore(ctx, globalTeamsSet, staleTeamID).Result() - require.ErrorIs(t, err, redis.Nil, "stale empty team should have been ZREM'd") -} - -// TestCleaner_PreservesYoungEmptyTeam validates the StaleCutoff race guard: -// an empty team that was added recently must NOT be pruned, because Add -// writes the team entry before the per-team SET SADD (operations.go:33-47) -// so an empty SET on a fresh team can be a transient in-flight Add. -func TestCleaner_PreservesYoungEmptyTeam(t *testing.T) { - t.Parallel() - - storage, client := setupTestStorage(t) - ctx := t.Context() - - youngTeamID := uuid.New().String() - - require.NoError(t, client.ZAdd(ctx, globalTeamsSet, redis.Z{ - Score: float64(time.Now().Unix()), - Member: youngTeamID, - }).Err()) - - cleaner := NewCleaner(storage) - require.NoError(t, cleaner.RunOnce(ctx)) - - _, err := client.ZScore(ctx, globalTeamsSet, youngTeamID).Result() - require.NoError(t, err, "young empty team should be preserved (race guard)") -} - -// TestCleaner_ConcurrentRunsConverge proves the multi-pod safety contract: -// running RunOnce from N goroutines against the same Redis must produce -// the same final state as one run, with no errors. -func TestCleaner_ConcurrentRunsConverge(t *testing.T) { - t.Parallel() - - storage, client := setupTestStorage(t) - ctx := t.Context() - - const n = 50 - for i := range n { - teamID := uuid.New().String() - sandboxID := fmt.Sprintf("ghost-%d", i) - require.NoError(t, client.ZAdd(ctx, globalExpirationSet, redis.Z{ - Score: float64(time.Now().Add(-time.Hour).UnixMilli()), - Member: expirationMember(teamID, sandboxID), - }).Err()) - } - - initial, err := client.ZCard(ctx, globalExpirationSet).Result() - require.NoError(t, err) - require.EqualValues(t, n, initial) - - cleaner := NewCleaner(storage) - - var wg sync.WaitGroup - for range 4 { - wg.Go(func() { - assert.NoError(t, cleaner.RunOnce(ctx)) - }) - } - wg.Wait() - - final, err := client.ZCard(ctx, globalExpirationSet).Result() - require.NoError(t, err) - require.Zero(t, final, "all orphans should be pruned despite concurrent runs") -} - -// TestCleaner_StartExitsOnContextCancel guards against a goroutine leak — -// the Start loop must return promptly when its context is cancelled. -func TestCleaner_StartExitsOnContextCancel(t *testing.T) { - t.Parallel() - - storage, _ := setupTestStorage(t) - cleaner := NewCleaner(storage) - cleaner.tick = 10 * time.Millisecond // tighten so a tick happens during the test - - ctx, cancel := context.WithCancel(t.Context()) - done := make(chan struct{}) - go func() { - cleaner.Start(ctx) - close(done) - }() - - // Let at least one tick fire, then cancel. - time.Sleep(50 * time.Millisecond) - cancel() - - select { - case <-done: - // success - case <-time.After(2 * time.Second): - t.Fatal("Cleaner.Start did not exit on ctx.Done()") - } -} - -// TestCleaner_PreservesFutureScoredExpirationEntry guards the score-based -// filter inside ExpiredItems: a member whose score is in the future -// (sandbox still running) but whose JSON is briefly missing must not be -// touched, because ExpiredItems' ZRangeByScore filters by score <= now. -func TestCleaner_PreservesFutureScoredExpirationEntry(t *testing.T) { - t.Parallel() - - storage, client := setupTestStorage(t) - ctx := t.Context() - - teamID := uuid.New().String() - sandboxID := "future-" + uuid.NewString() - member := expirationMember(teamID, sandboxID) - - // Score in the future — outside the ZRangeByScore window in ExpiredItems. - require.NoError(t, client.ZAdd(ctx, globalExpirationSet, redis.Z{ - Score: float64(time.Now().Add(time.Hour).UnixMilli()), - Member: member, - }).Err()) - - cleaner := NewCleaner(storage) - require.NoError(t, cleaner.RunOnce(ctx)) - - _, err := client.ZScore(ctx, globalExpirationSet, member).Result() - require.NoError(t, err, "future-scored entry must not be pruned") -} - -// TestCleaner_EvictsStaleExpiredSandbox covers the new evictExpired path: -// a sandbox whose EndTime is older than StaleCutoff must be Remove()'d by -// the cleaner so its JSON key, per-team index entry, and globalExpirationSet -// member all disappear. -func TestCleaner_EvictsStaleExpiredSandbox(t *testing.T) { - t.Parallel() - - storage, client := setupTestStorage(t) - ctx := t.Context() - - sbx := createTestSandbox("stale-expired-" + uuid.NewString()) - sbx.EndTime = time.Now().Add(-sandboxtypes.StaleCutoff - time.Minute) - require.NoError(t, storage.Add(ctx, sbx)) - - cleaner := NewCleaner(storage) - require.NoError(t, cleaner.RunOnce(ctx)) - - _, err := storage.Get(ctx, sbx.TeamID, sbx.SandboxID) - require.ErrorIs(t, err, sandboxtypes.ErrNotFound, "stale expired sandbox JSON should be removed") - - _, err = client.ZScore(ctx, globalExpirationSet, - expirationMember(sbx.TeamID.String(), sbx.SandboxID)).Result() - require.ErrorIs(t, err, redis.Nil, "stale expired sandbox should be removed from globalExpirationSet") - - isMember, err := client.SIsMember(ctx, - GetSandboxStorageTeamIndexKey(sbx.TeamID.String()), sbx.SandboxID).Result() - require.NoError(t, err) - require.False(t, isMember, "stale expired sandbox should be removed from per-team index") -} - -// TestCleaner_PreservesRecentlyExpiredSandbox guards the StaleCutoff window -// inside evictExpired: a sandbox that has just expired (EndTime in the past -// but newer than StaleCutoff) is still the evictor's responsibility — the -// cleaner must leave it alone so we don't race the evictor. -func TestCleaner_PreservesRecentlyExpiredSandbox(t *testing.T) { - t.Parallel() - - storage, _ := setupTestStorage(t) - ctx := t.Context() - - sbx := createTestSandbox("fresh-expired-" + uuid.NewString()) - sbx.EndTime = time.Now().Add(-time.Second) - require.NoError(t, storage.Add(ctx, sbx)) - - cleaner := NewCleaner(storage) - require.NoError(t, cleaner.RunOnce(ctx)) - - got, err := storage.Get(ctx, sbx.TeamID, sbx.SandboxID) - require.NoError(t, err, "recently expired sandbox must survive — eviction is the evictor's job") - require.Equal(t, sbx.SandboxID, got.SandboxID) -} - -// Compile-time guard so future refactors of sandboxtypes.StaleCutoff get noticed -// here: the cleaner's correctness depends on it being > 0. -var _ = func() bool { - if sandboxtypes.StaleCutoff <= 0 { - panic("sandboxtypes.StaleCutoff must be positive for cleaner race guards to hold") - } - - return true -}() diff --git a/packages/api/internal/sandbox/storage/redis/main.go b/packages/api/internal/sandbox/storage/redis/main.go index cb6ba8a6c9..57d98a46eb 100644 --- a/packages/api/internal/sandbox/storage/redis/main.go +++ b/packages/api/internal/sandbox/storage/redis/main.go @@ -37,8 +37,6 @@ type Storage struct { publisher *publisher } -func (s *Storage) Name() string { return sandboxtypes.StorageNameRedis } - const meterScope = "github.com/e2b-dev/infra/packages/api/internal/sandbox/storage/redis" func NewStorage( diff --git a/packages/api/internal/sandbox/storage/redis/state_change.go b/packages/api/internal/sandbox/storage/redis/state_change.go index add88e6574..4a50293f94 100644 --- a/packages/api/internal/sandbox/storage/redis/state_change.go +++ b/packages/api/internal/sandbox/storage/redis/state_change.go @@ -174,8 +174,7 @@ func (s *Storage) createCallback(teamID uuid.UUID, sandboxID, transitionKey, res // Determine result value for waiters: // - Restore failure: propagate so callers know state is inconsistent - // - Transient original failure: signal success so concurrent ops - // (e.g. kill) can proceed — matching the memory implementation + // - Transient original failure: signal success so concurrent ops (e.g. kill) can proceed // - Non-transient failure: propagate the error resultValue := "" if restoreErr != nil { diff --git a/packages/api/internal/sandbox/store.go b/packages/api/internal/sandbox/store.go index d412ddbd48..2a69e5eb7b 100644 --- a/packages/api/internal/sandbox/store.go +++ b/packages/api/internal/sandbox/store.go @@ -30,18 +30,10 @@ type ( const sbxRemoveTimeout = 10 * time.Second -// Storage names are re-exported from sandboxtypes for callers using this package. -const ( - StorageNameMemory = sandboxtypes.StorageNameMemory - StorageNameRedis = sandboxtypes.StorageNameRedis - StorageNamePopulateRedis = sandboxtypes.StorageNamePopulateRedis -) - // Storage and ReservationStorage are re-exported from sandboxtypes so external // callers can continue to use sandbox.Storage / sandbox.ReservationStorage. -// They live in sandboxtypes (a leaf package) so storage backends like -// sandbox/storage/memory can implement them without creating an import cycle -// back into package sandbox. +// They live in sandboxtypes (a leaf package) so storage backends can implement +// them without creating an import cycle back into package sandbox. type ( Storage = sandboxtypes.Storage ReservationStorage = sandboxtypes.ReservationStorage @@ -92,32 +84,10 @@ func (s *Store) Add(ctx context.Context, sandbox Sandbox, creation *CreationMeta } err := s.storage.Add(ctx, sandbox) - if err == nil { - // Count only newly added sandboxes to the store - s.callbacks.AddSandboxToRoutingTable(ctx, sandbox) - } else { - // TODO [ENG-3514]: Remove once migrated to Redis - // There's a race condition when the sandbox is added from node sync - // This should be fixed once the sync is improved - if !errors.Is(err, ErrAlreadyExists) { - return err - } - - logger.L().Warn(ctx, "Sandbox already exists in cache", logger.WithSandboxID(sandbox.SandboxID)) - } - - // TODO [ENG-3514]: Simplify once migrated to Redis - // Ensure the team reservation is set - no limit. - if s.storage.Name() != StorageNameRedis { - finishStart, _, err := s.reservations.Reserve(ctx, sandbox.TeamID, sandbox.SandboxID, -1) - if err != nil { - logger.L().Error(ctx, "Failed to reserve sandbox", zap.Error(err), logger.WithSandboxID(sandbox.SandboxID)) - } - - if finishStart != nil { - finishStart(sandbox, nil) - } + if err != nil { + return err } + s.callbacks.AddSandboxToRoutingTable(ctx, sandbox) if creation != nil { meta := *creation @@ -168,31 +138,20 @@ func (s *Store) WaitForStateChange(ctx context.Context, teamID uuid.UUID, sandbo } func (s *Store) Reconcile(ctx context.Context, sandboxes []Sandbox, nodeID string) { - sbxsToBeSynced := s.storage.Reconcile(ctx, sandboxes, nodeID) - - if s.storage.Name() == StorageNameRedis { - // Redis is the source of truth — divergent sandboxes are orphans running - // on the node but not present in the store. Kill them. - wg := sync.WaitGroup{} - for _, sbx := range sbxsToBeSynced { - wg.Go(func() { - ctx, cancel := context.WithTimeout(context.WithoutCancel(ctx), sbxRemoveTimeout) - defer cancel() - s.callbacks.RemoveSandboxFromNode(ctx, sbx) - }) - } - - wg.Wait() - } else { - // Memory backend — divergent sandboxes are ones discovered on the node - // that aren't in the local cache yet. Re-add them. - for _, sbx := range sbxsToBeSynced { - err := s.Add(ctx, sbx, nil) - if err != nil { - logger.L().Error(ctx, "Failed to re-add sandbox during sync", zap.Error(err), logger.WithSandboxID(sbx.SandboxID)) - } - } + // Redis is the source of truth — divergent sandboxes are orphans running + // on the node but not present in the store. Kill them. + orphans := s.storage.Reconcile(ctx, sandboxes, nodeID) + + wg := sync.WaitGroup{} + for _, sbx := range orphans { + wg.Go(func() { + ctx, cancel := context.WithTimeout(context.WithoutCancel(ctx), sbxRemoveTimeout) + defer cancel() + s.callbacks.RemoveSandboxFromNode(ctx, sbx) + }) } + + wg.Wait() } func (s *Store) Reserve(ctx context.Context, teamID uuid.UUID, sandboxID string, limit int) (finishStart func(Sandbox, error), waitForStart func(ctx context.Context) (Sandbox, error), err error) { diff --git a/packages/api/internal/sandbox/store_test.go b/packages/api/internal/sandbox/store_test.go index f7d0c7d0bc..a40b323b46 100644 --- a/packages/api/internal/sandbox/store_test.go +++ b/packages/api/internal/sandbox/store_test.go @@ -12,11 +12,27 @@ import ( "github.com/google/uuid" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "go.opentelemetry.io/otel/metric/noop" - "github.com/e2b-dev/infra/packages/api/internal/sandbox/storage/memory" + sandboxredis "github.com/e2b-dev/infra/packages/api/internal/sandbox/storage/redis" "github.com/e2b-dev/infra/packages/shared/pkg/consts" + redis_utils "github.com/e2b-dev/infra/packages/shared/pkg/redis" ) +// newTestStorage spins up a redis testcontainer and returns a fresh redis-backed +// sandbox.Storage. The container and storage are cleaned up via t.Cleanup. +func newTestStorage(t *testing.T) Storage { + t.Helper() + + client := redis_utils.SetupInstance(t) + storage, err := sandboxredis.NewStorage(client, noop.NewMeterProvider()) + require.NoError(t, err) + go storage.Start(t.Context()) + t.Cleanup(func() { storage.Close(context.WithoutCancel(t.Context())) }) + + return storage +} + // ============================================================================= // Test Helpers // ============================================================================= @@ -212,7 +228,7 @@ func TestAdd_NewSandbox(t *testing.T) { ctx := t.Context() // Setup - storage := memory.NewStorage() + storage := newTestStorage(t) reservations := &NoOpReservationStorage{} tracker := NewCallbackTracker(2) // Expect 2 callbacks @@ -244,92 +260,13 @@ func TestAdd_NewSandbox(t *testing.T) { }) } -func TestAdd_AlreadyInCache(t *testing.T) { - t.Parallel() - t.Run("newlyCreated=true - only AsyncNewlyCreatedSandbox called when already in cache", func(t *testing.T) { - t.Parallel() - ctx := t.Context() - - storage := memory.NewStorage() - reservations := &NoOpReservationStorage{} - - // First add with all 2 callbacks - tracker1 := NewCallbackTracker(2) - callbacks1 := Callbacks{ - AddSandboxToRoutingTable: tracker1.Track("AddSandboxToRoutingTable"), - AsyncNewlyCreatedSandbox: tracker1.TrackCreation("AsyncNewlyCreatedSandbox"), - } - store1 := NewStore(storage, reservations, callbacks1) - sbx := createTestSandbox() - - err := store1.Add(ctx, sbx, &CreationMetadata{}) - tracker1.WaitForCalls(t, 2*time.Second) - require.NoError(t, err) - - // Second add with newlyCreated=true, only AsyncNewlyCreatedSandbox callback - // (AddSandboxToRoutingTable is NOT called because already in cache) - tracker2 := NewCallbackTracker(1) - callbacks2 := Callbacks{ - AddSandboxToRoutingTable: tracker2.Track("AddSandboxToRoutingTable"), - AsyncNewlyCreatedSandbox: tracker2.TrackCreation("AsyncNewlyCreatedSandbox"), - } - store2 := NewStore(storage, reservations, callbacks2) - - err = store2.Add(ctx, sbx, &CreationMetadata{}) - tracker2.WaitForCalls(t, 2*time.Second) - - require.NoError(t, err) - tracker2.AssertNotCalled(t, "AddSandboxToRoutingTable") - tracker2.AssertCallCount(t, "AsyncNewlyCreatedSandbox", 1) - }) - - t.Run("newlyCreated=false - no callbacks called when already in cache", func(t *testing.T) { - t.Parallel() - ctx := t.Context() - - storage := memory.NewStorage() - reservations := &NoOpReservationStorage{} - - // First add with newlyCreated=true - tracker1 := NewCallbackTracker(2) - callbacks1 := Callbacks{ - AddSandboxToRoutingTable: tracker1.Track("AddSandboxToRoutingTable"), - AsyncNewlyCreatedSandbox: tracker1.TrackCreation("AsyncNewlyCreatedSandbox"), - } - store1 := NewStore(storage, reservations, callbacks1) - sbx := createTestSandbox() - - err := store1.Add(ctx, sbx, &CreationMetadata{}) - tracker1.WaitForCalls(t, 2*time.Second) - require.NoError(t, err) - - // Second add with newlyCreated=false, no callbacks expected - // (No callbacks called because already in cache) - tracker2 := NewCallbackTracker(0) - callbacks2 := Callbacks{ - AddSandboxToRoutingTable: tracker2.Track("AddSandboxToRoutingTable"), - AsyncNewlyCreatedSandbox: tracker2.TrackCreation("AsyncNewlyCreatedSandbox"), - } - store2 := NewStore(storage, reservations, callbacks2) - - err = store2.Add(ctx, sbx, nil) - require.NoError(t, err) - - // Give a small delay for any async callbacks (there should be none) - time.Sleep(100 * time.Millisecond) - - tracker2.AssertNotCalled(t, "AddSandboxToRoutingTable") - tracker2.AssertNotCalled(t, "AsyncNewlyCreatedSandbox") - }) -} - func TestAdd_NotNewlyCreated(t *testing.T) { t.Parallel() t.Run("not in cache - AddSandboxToRoutingTable called", func(t *testing.T) { t.Parallel() ctx := t.Context() - storage := memory.NewStorage() + storage := newTestStorage(t) reservations := &NoOpReservationStorage{} // Add with newlyCreated=false, expect 1 callback @@ -348,44 +285,6 @@ func TestAdd_NotNewlyCreated(t *testing.T) { tracker.AssertCallCount(t, "AddSandboxToRoutingTable", 1) tracker.AssertNotCalled(t, "AsyncNewlyCreatedSandbox") }) - - t.Run("already in cache - no callbacks called", func(t *testing.T) { - t.Parallel() - ctx := t.Context() - - storage := memory.NewStorage() - reservations := &NoOpReservationStorage{} - - // First add - tracker1 := NewCallbackTracker(1) - callbacks1 := Callbacks{ - AddSandboxToRoutingTable: tracker1.Track("AddSandboxToRoutingTable"), - AsyncNewlyCreatedSandbox: tracker1.TrackCreation("AsyncNewlyCreatedSandbox"), - } - store1 := NewStore(storage, reservations, callbacks1) - sbx := createTestSandbox() - - err := store1.Add(ctx, sbx, nil) - tracker1.WaitForCalls(t, 2*time.Second) - require.NoError(t, err) - - // Second add with same sandbox, newlyCreated=false, no callbacks expected - tracker2 := NewCallbackTracker(0) - callbacks2 := Callbacks{ - AddSandboxToRoutingTable: tracker2.Track("AddSandboxToRoutingTable"), - AsyncNewlyCreatedSandbox: tracker2.TrackCreation("AsyncNewlyCreatedSandbox"), - } - store2 := NewStore(storage, reservations, callbacks2) - - err = store2.Add(ctx, sbx, nil) - require.NoError(t, err) - - // Give a small delay for any async callbacks (there should be none) - time.Sleep(100 * time.Millisecond) - - tracker2.AssertNotCalled(t, "AddSandboxToRoutingTable") - tracker2.AssertNotCalled(t, "AsyncNewlyCreatedSandbox") - }) } func TestAdd_StorageErrors(t *testing.T) { @@ -394,7 +293,7 @@ func TestAdd_StorageErrors(t *testing.T) { t.Parallel() ctx := t.Context() - storage := memory.NewStorage() + storage := newTestStorage(t) mockStorage := NewMockStorage(storage) customErr := errors.New("storage failure") mockStorage.SetAddError(customErr) @@ -431,7 +330,7 @@ func TestAdd_ConcurrentCalls(t *testing.T) { t.Parallel() ctx := t.Context() - storage := memory.NewStorage() + storage := newTestStorage(t) reservations := &NoOpReservationStorage{} numGoroutines := 100 @@ -486,56 +385,4 @@ func TestAdd_ConcurrentCalls(t *testing.T) { assert.NoError(t, err, "expected sandbox %s to be in storage", sandboxID) } }) - - t.Run("concurrent adds for same sandbox", func(t *testing.T) { - t.Parallel() - ctx := t.Context() - - storage := memory.NewStorage() - reservations := &NoOpReservationStorage{} - - numGoroutines := 10 - sbx := createTestSandbox() - sbx.SandboxID = "concurrent-same-sandbox" - - // One will succeed with all 2 callbacks, rest will get ErrAlreadyExists with only AsyncNewlyCreatedSandbox callback - // Total: 2 + 9 = 11 callbacks (AddSandboxToRoutingTable: 1, AsyncNewlyCreatedSandbox: 10) - tracker := NewCallbackTracker(1 + numGoroutines) - - callbacks := Callbacks{ - AddSandboxToRoutingTable: tracker.Track("AddSandboxToRoutingTable"), - AsyncNewlyCreatedSandbox: tracker.TrackCreation("AsyncNewlyCreatedSandbox"), - } - store := NewStore(storage, reservations, callbacks) - - var wg sync.WaitGroup - successCount := atomic.Int32{} - - // Launch concurrent adds for the same sandbox - for range numGoroutines { - wg.Go(func() { - err := store.Add(ctx, sbx, &CreationMetadata{}) - if err == nil { - successCount.Add(1) - } - }) - } - - wg.Wait() - - // All should succeed (Add returns nil even for ErrAlreadyExists) - assert.Equal(t, int32(numGoroutines), successCount.Load()) - - // Wait for all callbacks - tracker.WaitForCalls(t, 5*time.Second) - - // Verify callbacks - tracker.AssertCallCount(t, "AddSandboxToRoutingTable", 1) // Only called once (first successful add) - tracker.AssertCallCount(t, "AsyncNewlyCreatedSandbox", numGoroutines) // All calls have newlyCreated=true - - // Verify sandbox exists in storage - stored, err := storage.Get(ctx, sbx.TeamID, sbx.SandboxID) - require.NoError(t, err) - assert.Equal(t, sbx.SandboxID, stored.SandboxID) - }) }