diff --git a/cmd/rainier/lifecycle.go b/cmd/rainier/lifecycle.go index cbde1d1..f0a86ca 100644 --- a/cmd/rainier/lifecycle.go +++ b/cmd/rainier/lifecycle.go @@ -115,7 +115,7 @@ func eligibilityText(verdict, note string) string { // yet" from "never". func attachNote(s session) string { switch s.State { - case "queued", "creating": + case "queued", "creating", "resuming": return "waits for it to start" case "suspended_warm", "suspended_cold": return "resumes it first" @@ -141,7 +141,7 @@ func stopNote(s session) string { return "this build does not know this state; the server decides" } switch s.State { - case "queued", "creating": + case "queued", "creating", "resuming": return "only a running session can be stopped" case "suspended_warm", "suspended_cold": return "already stopped" diff --git a/cmd/rainier/sessionstate.go b/cmd/rainier/sessionstate.go index c295dac..d1f9b71 100644 --- a/cmd/rainier/sessionstate.go +++ b/cmd/rainier/sessionstate.go @@ -77,7 +77,7 @@ const ( // exited would be the collapse this file exists to prevent. func lifecycleOf(s session) string { switch s.State { - case "queued", "creating": + case "queued", "creating", "resuming": return lifecycleStarting case "running": // Regardless of child_exit_code. The sandbox is up. @@ -188,7 +188,7 @@ const ( // Everything else is refused by the endpoint. func canAttach(s session) string { switch s.State { - case "running", "queued", "creating", "suspended_warm", "suspended_cold": + case "running", "queued", "creating", "resuming", "suspended_warm", "suspended_cold": return eligibleYes case "failed": if s.Reachable { @@ -216,7 +216,7 @@ func canStop(s session) string { switch s.State { case "running": return eligibleYes - case "queued", "creating", "suspended_warm", "suspended_cold", + case "queued", "creating", "resuming", "suspended_warm", "suspended_cold", "failed", "dead", "canceled", "destroyed": return eligibleNo default: @@ -231,7 +231,7 @@ func canStop(s session) string { // a live container and deleting it is the cleanup. func canDelete(s session) string { switch s.State { - case "running", "queued", "suspended_warm", "suspended_cold", + case "running", "queued", "resuming", "suspended_warm", "suspended_cold", "failed", "dead", "canceled": return eligibleYes case "creating": diff --git a/cmd/rainier/sessionstate_test.go b/cmd/rainier/sessionstate_test.go index 1c45f05..904f4bb 100644 --- a/cmd/rainier/sessionstate_test.go +++ b/cmd/rainier/sessionstate_test.go @@ -28,6 +28,7 @@ func TestThreeDimensions(t *testing.T) { "creating", session{State: "creating", Reachable: true}, lifecycleStarting, processNone, connectionAvailable, }, + {"resuming", session{State: "resuming", Reachable: true}, lifecycleStarting, processNone, connectionAvailable}, { "running with a live child", session{State: "running", Reachable: true}, lifecycleRunning, processRunning, connectionAvailable, @@ -159,6 +160,7 @@ func TestActionEligibilityFollowsRawStates(t *testing.T) { // it with a conflict because a dispatch may be in flight. "creating", session{State: "creating"}, eligibleYes, eligibleNo, eligibleNo, }, + {"resuming", session{State: "resuming"}, eligibleYes, eligibleNo, eligibleYes}, {"suspended warm", session{State: "suspended_warm"}, eligibleYes, eligibleNo, eligibleYes}, {"suspended cold", session{State: "suspended_cold"}, eligibleYes, eligibleNo, eligibleYes}, { diff --git a/cmd/rainier/wait.go b/cmd/rainier/wait.go index d7d3de1..1f18645 100644 --- a/cmd/rainier/wait.go +++ b/cmd/rainier/wait.go @@ -45,7 +45,7 @@ func initialAttachGuidance(ctx context.Context, cfg cli.Config, id string, since } else { row := resp.Session switch row.State { - case "queued", "creating": + case "queued", "creating", "resuming": explanation = "session is " + row.State + " and not ready to attach" if row.QueueReason != "" { // Queue text is server-owned prose. Redact known credentials before diff --git a/cmd/runnerd/main.go b/cmd/runnerd/main.go index 293ac48..fa25476 100644 --- a/cmd/runnerd/main.go +++ b/cmd/runnerd/main.go @@ -34,6 +34,8 @@ func main() { "bearer token for the controld dial (required when --controld is set; or set RAINIER_RUNNER_TOKEN, which keeps it out of the process list)") driverFlag := flag.String("driver", envDefault("RAINIER_RUNNER_DRIVER", "docker"), "execution driver for session sandboxes: docker | microvm") + microvmGuestReconnect := flag.Bool("microvm-guest-reconnect", os.Getenv("RAINIER_MICROVM_GUEST_RECONNECT") == "1", + "opt in to authenticated recovery of surviving microVM guests after control connection or runner restart; requires a matching control plane and guest image") kernelPath := flag.String("kernel", envDefault("RAINIER_KERNEL_PATH", ""), "guest vmlinux kernel path for the microvm driver (required when --driver=microvm; the runner refuses to start without a readable one)") rootfsPath := flag.String("rootfs", envDefault("RAINIER_ROOTFS_PATH", ""), @@ -149,14 +151,15 @@ func main() { log.Fatalf("--driver=microvm: %v", err) } mvm, err := driver.NewMicrovm(driver.MicrovmOpts{ - Checkpoint: ckpt, - KernelPath: *kernelPath, - BaseRootfs: *rootfsPath, - StateDir: *microvmStateDir, - TotalSlots: *slots, - VCPU: *microvmVCPUs, - MemoryMiB: *microvmMemoryMiB, - ImageSource: imageSource, + GuestReconnect: *microvmGuestReconnect, + Checkpoint: ckpt, + KernelPath: *kernelPath, + BaseRootfs: *rootfsPath, + StateDir: *microvmStateDir, + TotalSlots: *slots, + VCPU: *microvmVCPUs, + MemoryMiB: *microvmMemoryMiB, + ImageSource: imageSource, SlotGuestCIDR: *microvmGuestCIDR, SlotUplinkCIDR: *microvmUplinkCIDR, diff --git a/cmd/sessiond/bootlaunch.go b/cmd/sessiond/bootlaunch.go new file mode 100644 index 0000000..3a5de2e --- /dev/null +++ b/cmd/sessiond/bootlaunch.go @@ -0,0 +1,31 @@ +package main + +import ( + "context" + "slices" + + "github.com/tokencanopy/rainier/internal/relay" + "github.com/tokencanopy/rainier/protocol/runner" +) + +// guestBoot is the launch state consumed by main before preparing the boot chain. +type guestBoot struct { + conn relay.Conn + config runner.BootConfig + argv []string + failure *bootFailure +} + +func bootGuest(ctx context.Context, dial dialSession, boots *bootstrapper, imageCommand []string) (guestBoot, error) { + conn, cfg, failure, err := bootOverVsock(ctx, dial, boots) + if err != nil { + return guestBoot{}, err + } + // The image command only supplies a fallback. A microVM's per-session + // command arrives on the authenticated boot channel, not in PID 1's argv. + command := imageCommand + if len(cfg.Cmd) > 0 { + command = cfg.Cmd + } + return guestBoot{conn: conn, config: cfg, argv: slices.Clone(command), failure: failure}, nil +} diff --git a/cmd/sessiond/bootlaunch_test.go b/cmd/sessiond/bootlaunch_test.go new file mode 100644 index 0000000..d26f477 --- /dev/null +++ b/cmd/sessiond/bootlaunch_test.go @@ -0,0 +1,55 @@ +package main + +import ( + "context" + "os/exec" + "testing" + "time" + + "github.com/tokencanopy/rainier/protocol/runner" +) + +func TestGuestBootLaunchesRequestedCommand(t *testing.T) { + for _, requested := range []bool{true, false} { + name := "image_fallback" + if requested { + name = "requested_command" + } + t.Run(name, func(t *testing.T) { + cleanEnv(t) + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + dial, hosts := fakeTransport(t) + type result struct { + boot guestBoot + err error + } + done := make(chan result, 1) + go func() { + boot, err := bootGuest(ctx, dial, &bootstrapper{}, []string{"/bin/sh", "-c", "printf image_default"}) + done <- result{boot, err} + }() + host := <-hosts + cfg := runner.BootConfig{Protocol: runner.SessionBootstrapProtocolVersion, SessionID: "session_command_test"} + want := "image_default" + if requested { + // Shell-looking text must remain one literal argument, not be re-parsed. + want = "requested ; $(false) ' quoted" + cfg.Cmd = []string{"/bin/sh", "-c", `printf '%s' "$1"`, "command-test", want} + } + host.sendBootConfig(cfg) + got := <-done + if got.err != nil || got.boot.failure != nil { + t.Fatalf("guest boot failed: %v %v", got.err, got.boot.failure) + } + defer got.boot.conn.Close() + output, err := exec.CommandContext(ctx, got.boot.argv[0], got.boot.argv[1:]...).Output() + if err != nil { + t.Fatal(err) + } + if string(output) != want { + t.Fatalf("launched output %q, want %q", output, want) + } + }) + } +} diff --git a/cmd/sessiond/main.go b/cmd/sessiond/main.go index af1ac30..ea4b571 100644 --- a/cmd/sessiond/main.go +++ b/cmd/sessiond/main.go @@ -112,7 +112,7 @@ func main() { if *transport == transportVsock { dialer = vsockTransport() preamble = func(ctx context.Context, c relay.Conn) error { return reBootstrap(ctx, c, boots) } - conn, cfg, failure, err := bootOverVsock(context.Background(), dialer, boots) + boot, err := bootGuest(context.Background(), dialer, boots, argv) if err != nil { // Deliberately fatal, and the same judgement prepareBoot's own // failure gets: a guest that could not read its configuration @@ -120,12 +120,13 @@ func main() { // egress goes. Dying is what makes the runner notice. log.Fatalf("microvm boot: %v", err) } - firstConn, overVsock = conn, true + firstConn, overVsock = boot.conn, true + argv = boot.argv if *sessionID == "" { - *sessionID = cfg.SessionID + *sessionID = boot.config.SessionID } - if failure != nil { - secretsFailure = failure.Error() + if boot.failure != nil { + secretsFailure = boot.failure.Error() } } else if *transport != transportWebSocket { log.Fatalf("unknown --transport %q (valid: %s, %s)", *transport, transportWebSocket, transportVsock) diff --git a/cmd/sessiond/reconnect.go b/cmd/sessiond/reconnect.go index b2804db..cf9278e 100644 --- a/cmd/sessiond/reconnect.go +++ b/cmd/sessiond/reconnect.go @@ -5,6 +5,7 @@ import ( "context" "crypto/ed25519" "crypto/rand" + "crypto/sha256" "encoding/base64" "encoding/json" "errors" @@ -29,6 +30,7 @@ type guestReconnectIdentity struct { key ed25519.PrivateKey enrolled bool epoch uint64 + launch [32]byte } func (b *bootstrapper) enrollGuest(ctx context.Context, c relay.Conn, cfg runner.BootConfig) (map[string]string, error) { @@ -38,7 +40,7 @@ func (b *bootstrapper) enrollGuest(ctx context.Context, c relay.Conn, cfg runner return nil, errGuestReconnect } // A failed or unsupported opt-in is sticky: later dials cannot downgrade. - g := &guestReconnectIdentity{session: cfg.SessionID} + g := &guestReconnectIdentity{session: cfg.SessionID, launch: guestLaunchIdentity(cfg)} b.guest = g b.mu.Unlock() if cfg.GuestReconnect != runner.GuestReconnectProtocol || !guestID(cfg.SessionID) { @@ -106,7 +108,7 @@ func (b *bootstrapper) reconnectGuest(ctx context.Context, c relay.Conn) error { delivery, cancel := context.WithTimeout(ctx, time.Duration(accepted.ExpiresInSec)*time.Second) defer cancel() cfg, err := readBootConfig(delivery, c) - if err != nil || cfg.SessionID != session || cfg.GuestReconnect != 1 || cfg.BootstrapToken != accepted.Token { + if err != nil || cfg.SessionID != session || cfg.GuestReconnect != 1 || cfg.BootstrapToken != accepted.Token || guestLaunchIdentity(cfg) != g.launch { return errGuestReconnect } b.mu.Lock() @@ -335,3 +337,22 @@ func guestReady(ctx context.Context, c relay.Conn, epoch uint64) error { } return nil } + +// The live boot retains only a digest of fields whose consumers run once. +// Repositories, git attribution, setup, proxy policy and agent custody paths +// cannot be refreshed merely by changing sessiond's environment. A changed +// launch requires a new boot; ordinary environment values and secret names may +// still refresh. The digest never reaches the host or durable storage. +func guestLaunchIdentity(cfg runner.BootConfig) [32]byte { + cfg.BootstrapToken = "" + cfg.SecretNames = nil + launchEnv := map[string]string{} + for key, value := range cfg.Env { + if key == "RAINIER_AGENTS_B64" || strings.HasPrefix(value, "/rainier/agents/") { + launchEnv[key] = value + } + } + cfg.Env = launchEnv + body, _ := json.Marshal(cfg) + return sha256.Sum256(body) +} diff --git a/cmd/sessiond/reconnect_launch_test.go b/cmd/sessiond/reconnect_launch_test.go new file mode 100644 index 0000000..11077a6 --- /dev/null +++ b/cmd/sessiond/reconnect_launch_test.go @@ -0,0 +1,40 @@ +package main + +import ( + "github.com/tokencanopy/rainier/protocol/runner" + "testing" +) + +func TestReconnectRejectsChangedBootArtifactsBeforeRedemption(t *testing.T) { + for _, field := range []string{"command", "repository", "git", "setup", "proxy", "agent-manifest"} { + t.Run(field, func(t *testing.T) { + b, ch := reconnectFixture(t) + guest, host, ctx := reconnectPair(t) + done := make(chan error, 1) + go func() { done <- reBootstrap(ctx, guest, b); guest.Close() }() + accepted := acceptProof(t, ctx, host, b, ch, 1) + cfg := runner.BootConfig{Protocol: 1, GuestReconnect: 1, SessionID: "session-test", BootstrapToken: accepted.Token} + switch field { + case "command": + cfg.Cmd = []string{"different_test"} + case "repository": + cfg.Repos = []runner.RepoSpec{{Owner: "example", Name: "changed"}} + case "git": + cfg.GitAuthorName = "Changed Test" + case "setup": + cfg.Setup = "echo changed_test" + case "proxy": + cfg.ProxyURL = "http://proxy.example.invalid:80" + case "agent-manifest": + cfg.Env = map[string]string{"RAINIER_AGENTS_B64": "changed_test"} + } + sendReconnectConfig(t, ctx, host, cfg) + if raw, err := host.Read(ctx); err == nil || len(raw) != 0 { + t.Fatal("changed launch artifacts reached token redemption") + } + if err := <-done; err != errGuestReconnect { + t.Fatal("changed launch artifacts accepted") + } + }) + } +} diff --git a/cmd/sessiond/reconnect_test.go b/cmd/sessiond/reconnect_test.go index b0f2e97..167ad0b 100644 --- a/cmd/sessiond/reconnect_test.go +++ b/cmd/sessiond/reconnect_test.go @@ -75,7 +75,7 @@ func reconnectFixture(t *testing.T) (*bootstrapper, runner.GuestReconnectChallen if err != nil { t.Fatal(err) } - return &bootstrapper{guest: &guestReconnectIdentity{session: "session-test", boot: "boot-test", key: key, enrolled: true}}, runner.GuestReconnectChallenge{Protocol: 1, SessionID: "session-test", BootEpoch: "boot-test", HostIncarnation: "host-test", AttemptID: "attempt-test", PlacementGeneration: 1, Challenge: base64.RawURLEncoding.EncodeToString(make([]byte, 32))} + return &bootstrapper{guest: &guestReconnectIdentity{session: "session-test", boot: "boot-test", key: key, enrolled: true, launch: guestLaunchIdentity(runner.BootConfig{Protocol: 1, GuestReconnect: 1, SessionID: "session-test"})}}, runner.GuestReconnectChallenge{Protocol: 1, SessionID: "session-test", BootEpoch: "boot-test", HostIncarnation: "host-test", AttemptID: "attempt-test", PlacementGeneration: 1, Challenge: base64.RawURLEncoding.EncodeToString(make([]byte, 32))} } func writeReconnect(t *testing.T, ctx context.Context, c relay.Conn, kind string, v any) { t.Helper() diff --git a/control/contract_test.go b/control/contract_test.go index 81d8845..39f4b29 100644 --- a/control/contract_test.go +++ b/control/contract_test.go @@ -185,6 +185,7 @@ func TestSessionStateVocabulary(t *testing.T) { values := map[control.SessionState]string{ control.StateQueued: "queued", control.StateCreating: "creating", + control.StateResuming: "resuming", control.StateRunning: "running", control.StateSuspendedWarm: "suspended_warm", control.StateSuspendedCold: "suspended_cold", @@ -200,7 +201,7 @@ func TestSessionStateVocabulary(t *testing.T) { } terminalStates := map[control.SessionState]bool{ - control.StateQueued: false, control.StateCreating: false, control.StateRunning: false, + control.StateQueued: false, control.StateCreating: false, control.StateResuming: false, control.StateRunning: false, control.StateSuspendedWarm: false, control.StateSuspendedCold: false, control.StateCanceled: true, control.StateFailed: true, control.StateDead: true, control.StateDestroyed: true, @@ -213,7 +214,7 @@ func TestSessionStateVocabulary(t *testing.T) { slotStates := map[control.SessionState]bool{ control.StateQueued: false, - control.StateCreating: true, control.StateRunning: true, + control.StateCreating: true, control.StateResuming: true, control.StateRunning: true, control.StateSuspendedWarm: true, control.StateSuspendedCold: false, control.StateCanceled: false, control.StateFailed: false, control.StateDead: false, control.StateDestroyed: false, @@ -225,7 +226,7 @@ func TestSessionStateVocabulary(t *testing.T) { } wantOrder := []control.SessionState{ - control.StateQueued, control.StateCreating, control.StateRunning, + control.StateQueued, control.StateCreating, control.StateResuming, control.StateRunning, control.StateSuspendedWarm, control.StateSuspendedCold, } if len(control.NonTerminal) != len(wantOrder) { diff --git a/control/fleet.go b/control/fleet.go index e430386..fca6105 100644 --- a/control/fleet.go +++ b/control/fleet.go @@ -35,7 +35,11 @@ type Runner struct { PoolID PoolID CapacityUsed int CapacityTotal int - Connected bool + // CapacityPlacements identifies slots counted in the same CapacityUsed + // observation, by exact placement. Missing entries remain reservations. + CapacityPlacements map[SessionID]uint64 + + Connected bool // Generation is monotonic per runner; a stale holder is rejected without // replacing the event shapes (see RunnerEvent). Generation uint64 @@ -69,14 +73,15 @@ type RunnerSession struct { // and session bindings here rather than accepting a user-created Scope. Every // field is adapter-derived; none is decoded directly from client JSON. type RunnerRegistration struct { - WorkspaceID WorkspaceID - PoolID PoolID - RunnerID RunnerID - Generation uint64 - CapacityUsed int - CapacityTotal int - Capabilities []string - Sessions []RunnerSession + WorkspaceID WorkspaceID + PoolID PoolID + RunnerID RunnerID + Generation uint64 + CapacityUsed int + CapacityTotal int + CapacityPlacements map[SessionID]uint64 + Capabilities []string + Sessions []RunnerSession } // RunnerRegistrationResult is the application's answer to a registration. @@ -92,13 +97,14 @@ type RunnerRegistrationResult struct { // reconciliation: its identity, generation, capacity, and the sessions it // holds. type RunnerSnapshot struct { - WorkspaceID WorkspaceID - PoolID PoolID - RunnerID RunnerID - Generation uint64 - CapacityUsed int - CapacityTotal int - Sessions []RunnerSession + WorkspaceID WorkspaceID + PoolID PoolID + RunnerID RunnerID + Generation uint64 + CapacityUsed int + CapacityTotal int + CapacityPlacements map[SessionID]uint64 + Sessions []RunnerSession } // ReconcileResult is the application's answer to a snapshot. @@ -199,3 +205,29 @@ type Event struct { // of which an audit reader needs to know that a command was run. Command string } + +// ValidateCapacityPlacements rejects an observation that cannot describe its +// aggregate. Keys come only from the authenticated runner's own inventory. +func ValidateCapacityPlacements(used int, placements map[SessionID]uint64) error { + if used < 0 || len(placements) > used { + return ErrInvalid + } + for id, generation := range placements { + if id == "" || generation == 0 || generation > uint64(1<<63-1) { + return ErrInvalid + } + } + return nil +} + +// AvailableRunnerSlots subtracts pending create/resume reservations. Reservations already present in this exact aggregate sample occupy one slot, +// not two. An old placement or a legacy sample cannot cancel a reservation. +func AvailableRunnerSlots(host Runner, pending []Session) int { + reserved := 0 + for _, row := range pending { + if row.PlacementGeneration == 0 || host.CapacityPlacements[row.ID] != row.PlacementGeneration { + reserved++ + } + } + return host.CapacityTotal - host.CapacityUsed - reserved +} diff --git a/control/ports.go b/control/ports.go index 07b326e..caa8950 100644 --- a/control/ports.go +++ b/control/ports.go @@ -20,15 +20,22 @@ import ( type Action string const ( - ActionCreate Action = "create" - ActionGet Action = "get" - ActionList Action = "list" - ActionUpdate Action = "update" - ActionDelete Action = "delete" - ActionSuspend Action = "suspend" - ActionResume Action = "resume" - ActionSnapshot Action = "snapshot" - ActionAttach Action = "attach" + // These are audit labels; reconnect is authorized as ActionResume. + ActionGuestEnroll Action = "guest_enroll" + ActionGuestReconnectBegin Action = "guest_reconnect_begin" + ActionGuestReconnectAccept Action = "guest_reconnect_accept" + ActionGuestReconnectConfigure Action = "guest_reconnect_configure" + ActionGuestBootstrapMint Action = "guest_bootstrap_mint" + ActionGuestBootstrapRedeem Action = "guest_bootstrap_redeem" + ActionCreate Action = "create" + ActionGet Action = "get" + ActionList Action = "list" + ActionUpdate Action = "update" + ActionDelete Action = "delete" + ActionSuspend Action = "suspend" + ActionResume Action = "resume" + ActionSnapshot Action = "snapshot" + ActionAttach Action = "attach" // ActionExec is the audit label for one command run inside a session's // sandbox. It is deliberately an EVENT label and not an authorization // verb: exec is authorized as ActionAttach, because it grants nothing a diff --git a/control/session.go b/control/session.go index 1f7df5f..688d014 100644 --- a/control/session.go +++ b/control/session.go @@ -16,6 +16,7 @@ const ( StateQueued SessionState = "queued" StateCreating SessionState = "creating" StateRunning SessionState = "running" + StateResuming SessionState = "resuming" StateSuspendedWarm SessionState = "suspended_warm" StateSuspendedCold SessionState = "suspended_cold" StateCanceled SessionState = "canceled" @@ -35,10 +36,10 @@ func (s SessionState) Terminal() bool { } // OccupiesSlot reports whether a session in state s counts against a -// runner's capacity: creating, running, or suspended_warm. +// runner's capacity: creating, resuming, running, or suspended_warm. func (s SessionState) OccupiesSlot() bool { switch s { - case StateCreating, StateRunning, StateSuspendedWarm: + case StateCreating, StateResuming, StateRunning, StateSuspendedWarm: return true } return false @@ -47,11 +48,15 @@ func (s SessionState) OccupiesSlot() bool { // NonTerminal lists every non-terminal state, in the order a session // normally progresses through them. Callers pass it as the from-list of a // guarded transition when any live state should be accepted. -var NonTerminal = []SessionState{StateQueued, StateCreating, StateRunning, StateSuspendedWarm, StateSuspendedCold} +var NonTerminal = []SessionState{StateQueued, StateCreating, StateResuming, StateRunning, StateSuspendedWarm, StateSuspendedCold} // TransitionOpts carries the columns a guarded session transition may update // alongside state. A nil field leaves that column unchanged. type TransitionOpts struct { + // ExpectedPlacementGeneration, when non-nil, makes the state transition a + // compare-and-swap against this exact placement, including lifecycle retries. + ExpectedPlacementGeneration *uint64 + RunnerID *RunnerID Error *string // Image, when non-nil, records the image this placement resolved for the diff --git a/controlapp/fleet.go b/controlapp/fleet.go index f788590..a6ea613 100644 --- a/controlapp/fleet.go +++ b/controlapp/fleet.go @@ -61,6 +61,11 @@ type FleetOptions struct { // the exposure this whole design removes. Bootstraps control.SessionBootstrapStore + // GuestReconnect enables negotiation only when the host composes the + // transactional current-authority dispatcher. Driver capability alone + // cannot establish that its peer implements enrollment and recovery. + GuestReconnect bool + // DefaultEgress replaces the built-in developer egress baseline // (DefaultDeveloperEgressHosts) that every dispatched session's allowlist // is unioned with. It is a POINTER because "leave it alone" and "make it @@ -108,7 +113,8 @@ type FleetService struct { // bootstraps mints and records a microVM session's bootstrap token // (FleetOptions.Bootstraps), composed once here over the same clock every // other decision in this service reads. - bootstraps SessionBootstrapMinter + bootstraps SessionBootstrapMinter + guestReconnect bool // defaultEgress is the host's developer egress baseline, resolved once at // construction (FleetOptions.DefaultEgress). Held as a plain slice because @@ -160,6 +166,7 @@ func NewFleetService(opts FleetOptions) (*FleetService, error) { uow: opts.UnitOfWork, checkpoints: opts.Checkpoints, bootstraps: SessionBootstrapMinter{Store: opts.Bootstraps, Clock: opts.Clock}, + guestReconnect: opts.GuestReconnect, defaultEgress: resolveDefaultEgress(opts.DefaultEgress), wake: make(chan control.PoolID, 64), known: make(map[control.PoolID]struct{}), @@ -247,14 +254,15 @@ func (s *FleetService) registerRunner(ctx context.Context, r control.RunnerRegis } runner := control.Runner{ - ID: r.RunnerID, - PoolID: r.PoolID, - CapacityUsed: r.CapacityUsed, - CapacityTotal: r.CapacityTotal, - Connected: true, - Generation: r.Generation, - Capabilities: slices.Clone(r.Capabilities), - LastSeenAt: s.clock.Now(), + ID: r.RunnerID, + PoolID: r.PoolID, + CapacityUsed: r.CapacityUsed, + CapacityTotal: r.CapacityTotal, + CapacityPlacements: r.CapacityPlacements, + Connected: true, + Generation: r.Generation, + Capabilities: slices.Clone(r.Capabilities), + LastSeenAt: s.clock.Now(), } if err := s.fleet.UpsertRunner(ctx, r.PoolID, runner); err != nil { if errors.Is(err, control.ErrStale) { @@ -290,6 +298,9 @@ func (s *FleetService) authoritativeGeneration(ctx context.Context, pool control // validateRegistration rejects a malformed or contradictory claim before any // port is touched. func validateRegistration(r control.RunnerRegistration, poolScoped bool) error { + if err := control.ValidateCapacityPlacements(r.CapacityUsed, r.CapacityPlacements); err != nil { + return err + } if (!poolScoped && r.WorkspaceID == "") || r.PoolID == "" || r.RunnerID == "" || r.Generation == 0 || r.CapacityUsed < 0 || r.CapacityTotal < 0 || r.CapacityUsed > r.CapacityTotal { return control.ErrInvalid @@ -500,6 +511,9 @@ func (s *FleetService) reconcileRunner(ctx context.Context, snap control.RunnerS // validateSnapshot rejects a malformed snapshot, including one that names a // session twice, before any port is touched. func validateSnapshot(snap control.RunnerSnapshot, poolScoped bool) error { + if err := control.ValidateCapacityPlacements(snap.CapacityUsed, snap.CapacityPlacements); err != nil { + return err + } if (!poolScoped && snap.WorkspaceID == "") || snap.PoolID == "" || snap.RunnerID == "" || snap.Generation == 0 { return control.ErrInvalid } @@ -527,14 +541,15 @@ func (s *FleetService) upsertSnapshotRunner(ctx context.Context, snap control.Ru caps = slices.Clone(existing.Capabilities) } return s.fleet.UpsertRunner(ctx, snap.PoolID, control.Runner{ - ID: snap.RunnerID, - PoolID: snap.PoolID, - CapacityUsed: snap.CapacityUsed, - CapacityTotal: snap.CapacityTotal, - Connected: true, - Generation: snap.Generation, - Capabilities: caps, - LastSeenAt: s.clock.Now(), + ID: snap.RunnerID, + PoolID: snap.PoolID, + CapacityUsed: snap.CapacityUsed, + CapacityTotal: snap.CapacityTotal, + CapacityPlacements: snap.CapacityPlacements, + Connected: true, + Generation: snap.Generation, + Capabilities: caps, + LastSeenAt: s.clock.Now(), }) } @@ -562,7 +577,7 @@ func (s *FleetService) recordSnapshotAuthority(ctx context.Context, snap control // same snapshot produce the same Destroy list and no additional mutation. func (s *FleetService) reconcileSessions(ctx context.Context, snap control.RunnerSnapshot, poolScoped bool) ([]control.SessionID, error) { states := []control.SessionState{ - control.StateCreating, control.StateRunning, + control.StateCreating, control.StateResuming, control.StateRunning, control.StateSuspendedWarm, control.StateSuspendedCold, } stored, err := s.fleet.SessionsOnRunner(ctx, snap.PoolID, snap.RunnerID, states) @@ -597,6 +612,11 @@ func (s *FleetService) reconcileSessions(ctx context.Context, snap control.Runne } continue } + // An unversioned inventory cannot settle a pending cold boot. It may + // predate dispatch, and absence never permits fresh-create requeue. + if row.State == control.StateResuming { + continue + } if !present { if row.State == control.StateCreating { // A create that never landed goes back on the queue. @@ -714,8 +734,8 @@ var eventTransitions = map[control.SessionState][]control.SessionState{ control.StateRunning: {control.StateCreating, control.StateRunning}, control.StateSuspendedWarm: {control.StateRunning, control.StateSuspendedWarm}, control.StateSuspendedCold: {control.StateRunning, control.StateSuspendedCold}, - control.StateFailed: {control.StateCreating, control.StateRunning}, - control.StateDead: {control.StateCreating, control.StateRunning, control.StateSuspendedWarm, control.StateSuspendedCold}, + control.StateFailed: {control.StateCreating, control.StateResuming, control.StateRunning}, + control.StateDead: {control.StateCreating, control.StateResuming, control.StateRunning, control.StateSuspendedWarm, control.StateSuspendedCold}, } // runnerReportedDead is the safe reason recorded when a runner reports a @@ -775,7 +795,7 @@ func (s *FleetService) ApplyRunnerEvent(ctx context.Context, event control.Runne // placement generation the row has moved past comes from a sandbox this // session no longer has, even when it arrives on the current connection. // Zero is "not carried" (an old runner) and fences nothing. - if event.PlacementGeneration != 0 && event.PlacementGeneration != row.PlacementGeneration { + if (row.State == control.StateResuming && event.PlacementGeneration == 0) || (event.PlacementGeneration != 0 && event.PlacementGeneration != row.PlacementGeneration) { return control.ErrStale } @@ -816,7 +836,7 @@ func (s *FleetService) ApplyRunnerEvent(ctx context.Context, event control.Runne // Already-applied identical event; idempotent success, no record. return nil } - var opts control.TransitionOpts + opts := control.TransitionOpts{ExpectedPlacementGeneration: &row.PlacementGeneration} if target == control.StateFailed { detail := boundDetail(event.Detail) opts.Error = &detail diff --git a/controlapp/fleet_test.go b/controlapp/fleet_test.go index 3a29e0e..ead438d 100644 --- a/controlapp/fleet_test.go +++ b/controlapp/fleet_test.go @@ -151,7 +151,7 @@ func (st *fleetFakeStore) transition(ws control.WorkspaceID, id control.SessionI if !ok { return control.ErrNotFound } - if !fleetContainsState(from, s.State) { + if !fleetContainsState(from, s.State) || (opts.ExpectedPlacementGeneration != nil && s.PlacementGeneration != *opts.ExpectedPlacementGeneration) { return control.ErrConflict } s.State = to @@ -1150,6 +1150,8 @@ func TestReconcileRunnerMatrix(t *testing.T) { wantDestroy bool }{ {"creating adopted running", control.StateCreating, &control.RunnerSession{SessionID: "sess_example", State: control.StateRunning}, control.StateRunning, false}, + {"resuming missing remains pending", control.StateResuming, nil, control.StateResuming, false}, + {"resuming ignores old inventory", control.StateResuming, &control.RunnerSession{SessionID: "sess_example", State: control.StateRunning}, control.StateResuming, false}, {"running missing becomes dead", control.StateRunning, nil, control.StateDead, false}, {"creating missing requeues", control.StateCreating, nil, control.StateQueued, false}, {"terminal announced is orphan", control.StateDestroyed, &control.RunnerSession{SessionID: "sess_example", State: control.StateRunning}, control.StateDestroyed, true}, diff --git a/controlapp/reconnect_configuration_test.go b/controlapp/reconnect_configuration_test.go new file mode 100644 index 0000000..e3f068d --- /dev/null +++ b/controlapp/reconnect_configuration_test.go @@ -0,0 +1,93 @@ +package controlapp + +import ( + "context" + "errors" + "github.com/tokencanopy/rainier/control" + "github.com/tokencanopy/rainier/protocol/runner" + "slices" + "testing" +) + +func TestGuestReconnectSpecResolvesCurrentMaterialWithoutMinting(t *testing.T) { + fx := newFleetFixture(t) + row := bootstrapRow() + fx.resolver.material = LaunchMaterial{Environment: map[string]string{"OLD_TEST": "secret_test"}} + first, err := fx.service.ResolveGuestReconnectSpec(fleetCtx, row, nil) + if err != nil { + t.Fatal(err) + } + fx.resolver.material = LaunchMaterial{Environment: map[string]string{"NEW_TEST": "other_secret_test"}, Repos: []runner.RepoSpec{{Owner: "example", Name: "synthetic"}}} + next, err := fx.service.ResolveGuestReconnectSpec(fleetCtx, row, nil) + if err != nil { + t.Fatal(err) + } + if !slices.Equal(first.SecretNames, []string{"OLD_TEST"}) || !slices.Equal(next.SecretNames, []string{"NEW_TEST"}) { + t.Fatal("configuration was cached") + } + if len(next.Repos) != 1 || next.Repos[0].Name != "synthetic" { + t.Fatal("current repositories omitted") + } + if first.BootstrapToken != "" || next.BootstrapToken != "" { + t.Fatal("resolver minted authority") + } + if _, ok := fx.bootstraps.minted(row.ID); ok { + t.Fatal("resolver replaced the accepted token") + } + if _, ok := next.Env["NEW_TEST"]; ok { + t.Fatal("secret reached runner configuration") + } + if len(next.Env) == 0 { + t.Fatal("agent configuration omitted") + } + next.Env["MUTATION_TEST"] = "test" + again, err := fx.service.ResolveGuestReconnectSpec(fleetCtx, row, nil) + if err != nil || again.Env["MUTATION_TEST"] != "" { + t.Fatal("configuration aliases prior result") + } +} + +func TestGuestReconnectSpecFailureReturnsNoConfiguration(t *testing.T) { + fx := newFleetFixture(t) + fx.resolver.err = errors.New("private_resolver_test") + spec, err := fx.service.ResolveGuestReconnectSpec(fleetCtx, bootstrapRow(), nil) + if spec != nil || !errors.Is(err, control.ErrUnavailable) { + t.Fatalf("failed resolution: %v", err) + } + ctx, cancel := context.WithCancel(context.Background()) + cancel() + fx.resolver.err = nil + spec, err = fx.service.ResolveGuestReconnectSpec(ctx, bootstrapRow(), nil) + if spec != nil || !errors.Is(err, control.ErrUnavailable) { + t.Fatal("canceled resolution returned configuration") + } +} + +func TestGuestReconnectCreateRequiresNegotiatedCapability(t *testing.T) { + for _, enabled := range []bool{false, true} { + fx := newFleetFixture(t) + fx.service.guestReconnect = true + caps := []string{runner.CapabilityMicrovmV1} + if enabled { + caps = append(caps, runner.CapabilityGuestReconnectV1) + } + spec, fail := fx.service.createSpec(fleetCtx, bootstrapRow(), nil, caps, 3) + if fail != "" { + t.Fatal(fail) + } + if (spec.GuestReconnect == 1) != enabled { + t.Fatal("guest opt-in does not match negotiation") + } + } +} + +func TestGuestReconnectCreateRequiresHostSupport(t *testing.T) { + fx := newFleetFixture(t) + spec, fail := fx.service.createSpec(fleetCtx, bootstrapRow(), nil, []string{runner.CapabilityMicrovmV1, runner.CapabilityGuestReconnectV1}, 3) + if fail != "" { + t.Fatal(fail) + } + if spec.GuestReconnect != 0 { + t.Fatal("negotiated reconnect without host authority support") + } +} diff --git a/controlapp/repotest/repotest.go b/controlapp/repotest/repotest.go index dfd825f..a45336b 100644 --- a/controlapp/repotest/repotest.go +++ b/controlapp/repotest/repotest.go @@ -76,6 +76,7 @@ func cases() []suiteCase { {"S11 a placement transition records the resolved image", caseTransitionImage}, {"S12 exactly one of two claims from one generation wins", caseControllerClaimRace}, {"S13 the controller lease is fenced by its generation and its holder", caseControllerLease}, + {"S14 cold resume claims and completions compare placement generation", caseResumePlacementClaim}, {"E1 environment round trip", caseEnvironmentRoundTrip}, {"E2 an environment name is unique per workspace", caseEnvironmentName}, @@ -1406,7 +1407,7 @@ func caseRunnerRoundTrip(t *testing.T, s Stores) { for _, id := range []control.RunnerID{"runner_b", "runner_a"} { if err := s.Fleet.UpsertRunner(ctx, PoolA, control.Runner{ ID: id, PoolID: PoolA, CapacityUsed: 1, CapacityTotal: 4, Connected: true, - Generation: 1, Capabilities: caps, LastSeenAt: baseTime(), + Generation: 1, CapacityPlacements: map[control.SessionID]uint64{"sess_capacity_test": 2}, Capabilities: caps, LastSeenAt: baseTime(), }); err != nil { t.Fatalf("upsert %s: %v", id, err) } @@ -1425,6 +1426,8 @@ func caseRunnerRoundTrip(t *testing.T, s Stores) { t.Fatalf("PoolID = %q, want %q", got.PoolID, PoolA) case got.CapacityUsed != 1 || got.CapacityTotal != 4: t.Fatalf("capacity = %d/%d, want 1/4", got.CapacityUsed, got.CapacityTotal) + case got.CapacityPlacements["sess_capacity_test"] != 2: + t.Fatal("exact capacity placement missing from aggregate observation") case !got.Connected: t.Fatalf("connected = false, want true") case got.Generation != 1: @@ -1644,3 +1647,39 @@ func caseFleetEmptyPool(t *testing.T, s Stores) { } } } + +func caseResumePlacementClaim(t *testing.T, s Stores) { + ctx := context.Background() + row := mustCreate(t, s, Alpha, control.Session{ID: "sess_resume", CreatorID: "act_a", State: control.StateSuspendedCold, PoolID: PoolA, RunnerID: "runner_a"}) + expected := row.PlacementGeneration + opts := control.TransitionOpts{RunnerID: &row.RunnerID, ExpectedPlacementGeneration: &expected} + results := make(chan error, 2) + for i := 0; i < 2; i++ { + go func() { + results <- s.Sessions.Transition(ctx, Alpha, row.ID, []control.SessionState{control.StateSuspendedCold}, control.StateResuming, opts) + }() + } + wins := 0 + for i := 0; i < 2; i++ { + err := <-results + if err == nil { + wins++ + } else if !errors.Is(err, control.ErrConflict) { + t.Fatal(err) + } + } + if wins != 1 { + t.Fatalf("claim winners=%d", wins) + } + current, err := s.Sessions.GetSession(ctx, Alpha, row.ID) + if err != nil || current.PlacementGeneration != expected+1 || current.State != control.StateResuming { + t.Fatal("claim did not preserve its generation") + } + if err := s.Sessions.Transition(ctx, Alpha, row.ID, []control.SessionState{control.StateResuming}, control.StateRunning, control.TransitionOpts{ExpectedPlacementGeneration: &expected}); !errors.Is(err, control.ErrConflict) { + t.Fatal("stale completion was accepted") + } + expected = current.PlacementGeneration + if err := s.Sessions.Transition(ctx, Alpha, row.ID, []control.SessionState{control.StateResuming}, control.StateRunning, control.TransitionOpts{ExpectedPlacementGeneration: &expected}); err != nil { + t.Fatal(err) + } +} diff --git a/controlapp/resume_claim_test.go b/controlapp/resume_claim_test.go new file mode 100644 index 0000000..52ea8f2 --- /dev/null +++ b/controlapp/resume_claim_test.go @@ -0,0 +1,64 @@ +package controlapp + +import ( + "context" + "errors" + "testing" + + "github.com/tokencanopy/rainier/control" + "github.com/tokencanopy/rainier/protocol/runner" +) + +func TestColdResumeClaimsPlacementBeforeDispatch(t *testing.T) { + f := newSessionFixtureFull(t) + f.fleet.runners = []control.Runner{{ID: "runner_a", PoolID: "pool_a", CapacityTotal: 4, Connected: true}} + f.repo.put(sessionInState(control.StateSuspendedCold)) + f.transport.onDispatch = func(m runner.ToRunner) { + row, err := f.repo.GetSession(context.Background(), "ws_example", "sess_example") + if err != nil || row.State != control.SessionState("resuming") || row.PlacementGeneration != 2 || m.PlacementGeneration != 2 { + t.Fatal("cold boot dispatched before its new placement was claimed") + } + _, err = f.svc.ResumeSession(context.Background(), sessionTestScope(), control.ResumeSession{ID: row.ID}) + if !errors.Is(err, control.ErrConflict) { + t.Fatal("duplicate resume opened another claim") + } + } + got, err := f.svc.ResumeSession(context.Background(), sessionTestScope(), control.ResumeSession{ID: "sess_example"}) + if err != nil || got.State != control.StateRunning || got.PlacementGeneration != 2 { + t.Fatalf("resume outcome state=%s generation=%d err=%v", got.State, got.PlacementGeneration, err) + } +} + +func TestColdResumeTimeoutRetainsClaim(t *testing.T) { + f := newSessionFixtureFull(t) + f.fleet.runners = []control.Runner{{ID: "runner_a", PoolID: "pool_a", CapacityTotal: 4, Connected: true}} + f.repo.put(sessionInState(control.StateSuspendedCold)) + f.transport.err = context.DeadlineExceeded + _, err := f.svc.ResumeSession(context.Background(), sessionTestScope(), control.ResumeSession{ID: "sess_example"}) + if !errors.Is(err, control.ErrUnavailable) { + t.Fatalf("error=%v", err) + } + row, _ := f.repo.GetSession(context.Background(), "ws_example", "sess_example") + if row.State != control.SessionState("resuming") || row.PlacementGeneration != 2 { + t.Fatal("uncertain launch lost its claimed placement") + } +} + +func TestColdResumeCompletionCannotChangeNewerPlacement(t *testing.T) { + f := newSessionFixtureFull(t) + f.fleet.runners = []control.Runner{{ID: "runner_a", PoolID: "pool_a", CapacityTotal: 4, Connected: true}} + f.repo.put(sessionInState(control.StateSuspendedCold)) + f.transport.onDispatch = func(runner.ToRunner) { + row := f.repo.rows["sess_example"] + row.State = control.SessionState("resuming") + row.PlacementGeneration = 3 + f.repo.put(row) + } + got, err := f.svc.ResumeSession(context.Background(), sessionTestScope(), control.ResumeSession{ID: "sess_example"}) + if err != nil || got.State != control.SessionState("resuming") || got.PlacementGeneration != 3 { + t.Fatal("late completion changed newer placement") + } + if len(f.events.events) != 0 { + t.Fatal("late completion recorded a resume event") + } +} diff --git a/controlapp/resume_reconcile.go b/controlapp/resume_reconcile.go new file mode 100644 index 0000000..6b758b4 --- /dev/null +++ b/controlapp/resume_reconcile.go @@ -0,0 +1,64 @@ +package controlapp + +import ( + "context" + "errors" + "time" + + "github.com/tokencanopy/rainier/control" + "github.com/tokencanopy/rainier/protocol/runner" +) + +// Reuse the scheduler safety pass to settle lost resume replies. One bounded, +// sequential pass cannot launch VMs or create another placement. +func (s *FleetService) reconcileResumes(parent context.Context, pool control.PoolID) { + ctx, cancel := context.WithTimeout(parent, 5*time.Second) + defer cancel() + hosts, err := s.fleet.ListRunners(ctx, pool) + if err != nil { + return + } + remaining := 64 + for _, host := range hosts { + if !host.Connected { + continue + } + rows, err := s.fleet.SessionsOnRunner(ctx, pool, host.ID, []control.SessionState{control.StateResuming}) + if err != nil { + return + } + for _, row := range rows { + // Let the normal dispatch round trip finish before attempting to + // cancel an unreceived command. A live launch still reports pending. + if !row.UpdatedAt.IsZero() && s.clock.Now().Sub(row.UpdatedAt) < 2*time.Minute { + continue + } + if remaining == 0 || ctx.Err() != nil { + return + } + remaining-- + response, err := s.transport.Dispatch(ctx, pool, host.ID, runner.ToRunner{Type: "resume_status", Session: string(row.ID), PlacementGeneration: row.PlacementGeneration}) + if err != nil || !response.OK || response.PlacementGeneration != row.PlacementGeneration || row.PlacementGeneration == 0 { + continue + } + target := control.SessionState(response.State) + if target != control.StateRunning && target != control.StateSuspendedCold { + continue + } + _ = s.uow.Run(ctx, func(ctx context.Context) error { + err := s.sessions.Transition(ctx, row.WorkspaceID, row.ID, []control.SessionState{control.StateResuming}, target, control.TransitionOpts{ExpectedPlacementGeneration: &row.PlacementGeneration}) + if errors.Is(err, control.ErrConflict) || errors.Is(err, control.ErrNotFound) { + return nil + } + if err != nil { + return err + } + eventID := s.ids.NewEventID() + if eventID == "" { + return control.ErrUnavailable + } + return s.recordLifecycleEvent(ctx, control.RunnerEvent{WorkspaceID: row.WorkspaceID, RunnerID: host.ID, SessionID: row.ID}, row, eventID) + }) + } + } +} diff --git a/controlapp/resume_reconcile_test.go b/controlapp/resume_reconcile_test.go new file mode 100644 index 0000000..74c2783 --- /dev/null +++ b/controlapp/resume_reconcile_test.go @@ -0,0 +1,81 @@ +package controlapp + +import ( + "context" + "testing" + + "github.com/tokencanopy/rainier/control" + "github.com/tokencanopy/rainier/protocol/runner" +) + +func TestSchedulerReconcilesOnlyExactResumePlacement(t *testing.T) { + for _, tc := range []struct { + name, state string + generation uint64 + want control.SessionState + }{ + {"completed", "running", 2, control.StateRunning}, + {"failed_before_launch", "suspended_cold", 2, control.StateSuspendedCold}, + {"old_reply", "running", 1, control.StateResuming}, + {"unversioned", "running", 0, control.StateResuming}, + {"still_launching", "resuming", 2, control.StateResuming}, + } { + t.Run(tc.name, func(t *testing.T) { + fx := newFleetFixture(t) + fx.st.seedRunner(fleetReconcileRunnerRow()) + fx.st.seedSession(control.Session{ID: "sess_example", WorkspaceID: "ws_example", PoolID: "pool_example", RunnerID: "runner_example", State: control.StateResuming, PlacementGeneration: 2}) + fx.transport.dispatchReplies = []runner.FromRunner{{OK: true, State: tc.state, PlacementGeneration: tc.generation}} + fx.service.drainPool(context.Background(), "pool_example") + got := fleetGetSessionState(t, fx, "ws_example", "sess_example") + if got.State != tc.want || got.PlacementGeneration != 2 { + t.Fatalf("state=%s generation=%d", got.State, got.PlacementGeneration) + } + if len(fx.transport.dispatched) != 1 || fx.transport.dispatched[0].Type != "resume_status" || fx.transport.dispatched[0].PlacementGeneration != 2 { + t.Fatal("pending claim was not reconciled against exact runner placement") + } + }) + } +} + +func TestResumingAcceptsOnlyCurrentFailureEvidence(t *testing.T) { + for _, target := range []control.SessionState{control.StateFailed, control.StateDead} { + for _, generation := range []uint64{0, 1, 2} { + fx := newFleetFixture(t) + fx.st.seedRunner(fleetEventRunner(1)) + fx.st.seedSession(control.Session{ID: "sess_example", WorkspaceID: "ws_example", PoolID: "pool_example", RunnerID: "runner_example", State: control.StateResuming, PlacementGeneration: 2}) + err := fx.service.ApplyRunnerEvent(context.Background(), control.RunnerEvent{WorkspaceID: "ws_example", PoolID: "pool_example", RunnerID: "runner_example", SessionID: "sess_example", Generation: 1, PlacementGeneration: generation, State: target}) + row := fleetGetSessionState(t, fx, "ws_example", "sess_example") + if generation == 2 { + if err != nil || row.State != target { + t.Fatalf("current %s lost: %v state=%s", target, err, row.State) + } + } else if err == nil || row.State != control.StateResuming { + t.Fatal("stale failure changed pending boot") + } + } + } +} + +func TestPendingResumeCapacityCountsEachPlacementOnce(t *testing.T) { + for _, tc := range []struct { + name string + observed uint64 + want int + }{{"same", 2, 1}, {"old", 1, 0}, {"absent", 0, 0}} { + t.Run(tc.name, func(t *testing.T) { + fx := newFleetFixture(t) + host := fleetReconcileRunnerRow() + host.CapacityTotal = 2 + host.CapacityUsed = 1 + if tc.observed != 0 { + host.CapacityPlacements = map[control.SessionID]uint64{"sess_example": tc.observed} + } + fx.st.seedRunner(host) + fx.st.seedSession(control.Session{ID: "sess_example", WorkspaceID: "ws_example", PoolID: "pool_example", RunnerID: "runner_example", State: control.StateResuming, PlacementGeneration: 2}) + views, err := fx.service.freeCapacity(context.Background(), "pool_example") + if err != nil || len(views) != 1 || views[0].free != tc.want { + t.Fatalf("capacity=%+v err=%v want free=%d", views, err, tc.want) + } + }) + } +} diff --git a/controlapp/scheduler.go b/controlapp/scheduler.go index b2b9145..117fc8a 100644 --- a/controlapp/scheduler.go +++ b/controlapp/scheduler.go @@ -72,6 +72,7 @@ type runnerView struct { // runs sequentially here; only the create dispatch that follows a successful // placement runs concurrently. func (s *FleetService) drainPool(ctx context.Context, pool control.PoolID) { + s.reconcileResumes(ctx, pool) rows, err := s.fleet.OldestQueued(ctx, pool) if err != nil { return @@ -133,13 +134,13 @@ func (s *FleetService) freeCapacity(ctx context.Context, pool control.PoolID) ([ if !r.Connected { continue } - creating, err := s.fleet.SessionsOnRunner(ctx, pool, r.ID, []control.SessionState{control.StateCreating}) + creating, err := s.fleet.SessionsOnRunner(ctx, pool, r.ID, []control.SessionState{control.StateCreating, control.StateResuming}) if err != nil { return nil, err } views = append(views, runnerView{ id: r.ID, - free: r.CapacityTotal - r.CapacityUsed - len(creating), + free: control.AvailableRunnerSlots(r, creating), caps: r.Capabilities, }) } @@ -417,6 +418,42 @@ func (s *FleetService) placedGeneration(ctx context.Context, row control.Session // Spec.Env, as they always have, or stay behind a single-use token the guest // exchanges for them after it boots (see §3 of the design note). func (s *FleetService) createSpec(ctx context.Context, row control.Session, env *control.Environment, runnerCaps []string, gen uint64) (*runner.Spec, string) { + withheld := slices.Contains(runnerCaps, runner.CapabilityMicrovmV1) + spec, fail := s.launchSpec(ctx, row, env, withheld) + if fail != "" { + return nil, fail + } + if s.guestReconnect && withheld && slices.Contains(runnerCaps, runner.CapabilityGuestReconnectV1) { + spec.GuestReconnect = runner.GuestReconnectProtocol + } + if withheld { + token, err := s.bootstraps.Mint(ctx, row.WorkspaceID, row.ID, gen) + if err != nil { + return nil, "could not mint this session's bootstrap token" + } + spec.BootstrapToken = token + } + return spec, "" +} + +// ResolveGuestReconnectSpec resolves current non-secret launch configuration +// without minting or spending bootstrap authority. The caller must authorize +// current membership, policy and the exact runner placement before calling and +// before delivery. This is not authorization or a cache; the accepted proof's +// fresh token is supplied separately. Launch invariants must be checked against +// the live guest before applying changes. Failure returns nil and a fixed error. +func (s *FleetService) ResolveGuestReconnectSpec(ctx context.Context, row control.Session, env *control.Environment) (*runner.Spec, error) { + if ctx.Err() != nil { + return nil, control.ErrUnavailable + } + spec, fail := s.launchSpec(ctx, row, env, true) + if fail != "" || ctx.Err() != nil { + return nil, control.ErrUnavailable + } + return spec, nil +} + +func (s *FleetService) launchSpec(ctx context.Context, row control.Session, env *control.Environment, withheld bool) (*runner.Spec, string) { spec := runner.Spec{ Name: row.Name, Image: row.Spec.Image, @@ -456,7 +493,6 @@ func (s *FleetService) createSpec(ctx context.Context, row control.Session, env // ends in a file; so the values stay here and the create carries their // NAMES and a token instead (ADR-0003 §2.7 item 1). Every other runner — // which is every runner today — is dispatched exactly what it always was. - withheld := slices.Contains(runnerCaps, runner.CapabilityMicrovmV1) if !withheld { spec.Env = cloneMap(material.Environment) } @@ -484,7 +520,7 @@ func (s *FleetService) createSpec(ctx context.Context, row control.Session, env } spec.Env = agentEnv } - // The names, and then the token, in that order: the names are derived + // Withhold values before a caller adds its separately authorized token. Names are derived // from the material this function already holds, and they are computed // AFTER the agent-home block so that the one rule that block states — // agent paths and the manifest are launch invariants a workspace's own @@ -494,15 +530,6 @@ func (s *FleetService) createSpec(ctx context.Context, row control.Session, env // guest then applies over its own agent home. if withheld { spec.SecretNames = withholdableNames(material.Environment, spec.Env) - token, err := s.bootstraps.Mint(ctx, row.WorkspaceID, row.ID, gen) - if err != nil { - // Fail closed. The alternative to refusing here is a session - // dispatched with neither its secrets nor a way to ask for them, - // which boots, reports healthy, and fails at whatever the first - // credential-shaped thing it does is. - return nil, "could not mint this session's bootstrap token" - } - spec.BootstrapToken = token } // Either the values or the token, never both — stated as a check rather // than as a comment, because it is the whole security claim of §3 and it @@ -514,7 +541,7 @@ func (s *FleetService) createSpec(ctx context.Context, row control.Session, env // spec.Env below the withholding branch, or reorders the two — at which // point a failed create is a much better answer than a secret on a // shared host's disk. - if spec.BootstrapToken != "" { + if withheld { for _, name := range spec.SecretNames { if _, both := spec.Env[name]; both { return nil, "could not resolve launch material" diff --git a/controlapp/sessions.go b/controlapp/sessions.go index cdf6958..de55e63 100644 --- a/controlapp/sessions.go +++ b/controlapp/sessions.go @@ -608,18 +608,13 @@ func (s *SessionService) ResumeSession(ctx context.Context, scope control.Scope, if free <= 0 { return control.Session{}, control.ErrConflict } + return s.resumeCold(ctx, scope, row) } if _, err := s.dispatch(ctx, row, runner.ToRunner{Type: "resume", Session: string(row.ID)}); err != nil { return control.Session{}, err } - // A cold resume starts a new sandbox on the runner that holds the volume, - // so its transition names that runner and the repository opens a new - // placement generation for it; a warm resume names none. opts := control.TransitionOpts{} - if row.State == control.StateSuspendedCold { - opts.RunnerID = &row.RunnerID - } if err := s.uow.Run(ctx, func(ctx context.Context) error { if err := s.sessions.Transition(ctx, scope.WorkspaceID, cmd.ID, []control.SessionState{row.State}, control.StateRunning, opts); err != nil { if !errors.Is(err, control.ErrConflict) && !errors.Is(err, control.ErrNotFound) { @@ -627,18 +622,7 @@ func (s *SessionService) ResumeSession(ctx context.Context, scope control.Scope, } return nil } - // A cold resume's transition named a runner, so the repository opened - // a new placement generation for the sandbox it starts. The event - // belongs to the generation the row has AFTER the mutation, which - // only a re-read inside this same unit knows. recorded := row - if opts.RunnerID != nil { - cur, err := s.sessions.GetSession(ctx, scope.WorkspaceID, cmd.ID) - if err != nil { - return control.ErrUnavailable - } - recorded = cur - } return recordEvent(ctx, s.ids, s.events, s.clock, scope, control.ActionResume, sessionResource(recorded), recorded.PlacementGeneration) }); err != nil { @@ -648,8 +632,48 @@ func (s *SessionService) ResumeSession(ctx context.Context, scope control.Scope, return s.authoritative(ctx, scope.WorkspaceID, cmd.ID) } -// coldResumeFree computes free capacity on the runner already holding a cold -// session's volume: CapacityTotal - CapacityUsed - len(creating). A runner the +// resumeCold claims a new placement before any VM or bootstrap side effect. +// A lost dispatch reply retains the claim: only versioned runner evidence can +// establish whether the VM started. Never enqueue this workspace as a create. +func (s *SessionService) resumeCold(ctx context.Context, scope control.Scope, row control.Session) (control.Session, error) { + if !s.transport.Connected(row.PoolID, row.RunnerID) { + return control.Session{}, control.ErrUnavailable + } + var claimed control.Session + err := s.uow.Run(ctx, func(ctx context.Context) error { + opts := control.TransitionOpts{RunnerID: &row.RunnerID, ExpectedPlacementGeneration: &row.PlacementGeneration} + if err := s.sessions.Transition(ctx, scope.WorkspaceID, row.ID, []control.SessionState{control.StateSuspendedCold}, control.StateResuming, opts); err != nil { + return portError(err) + } + var err error + claimed, err = s.sessions.GetSession(ctx, scope.WorkspaceID, row.ID) + return portError(err) + }) + if err != nil { + return control.Session{}, err + } + if _, err = s.dispatch(ctx, claimed, runner.ToRunner{Type: "resume", Session: string(row.ID), PlacementGeneration: claimed.PlacementGeneration}); err != nil { + return control.Session{}, err + } + err = s.uow.Run(ctx, func(ctx context.Context) error { + opts := control.TransitionOpts{ExpectedPlacementGeneration: &claimed.PlacementGeneration} + if err := s.sessions.Transition(ctx, scope.WorkspaceID, row.ID, []control.SessionState{control.StateResuming}, control.StateRunning, opts); err != nil { + if errors.Is(err, control.ErrConflict) || errors.Is(err, control.ErrNotFound) { + return nil + } + return control.ErrUnavailable + } + return recordEvent(ctx, s.ids, s.events, s.clock, scope, control.ActionResume, sessionResource(claimed), claimed.PlacementGeneration) + }) + if err != nil { + return control.Session{}, err + } + s.wake(row.PoolID) + return s.authoritative(ctx, scope.WorkspaceID, row.ID) +} + +// coldResumeFree computes conservative free capacity on the runner holding +// the volume, including creating and pending-resume reservations. A runner the // pool no longer lists yields zero (no slot). func (s *SessionService) coldResumeFree(ctx context.Context, row control.Session) (int, error) { runners, err := s.fleet.ListRunners(ctx, row.PoolID) @@ -660,11 +684,11 @@ func (s *SessionService) coldResumeFree(ctx context.Context, row control.Session if r.ID != row.RunnerID { continue } - creating, err := s.fleet.SessionsOnRunner(ctx, row.PoolID, row.RunnerID, []control.SessionState{control.StateCreating}) + creating, err := s.fleet.SessionsOnRunner(ctx, row.PoolID, row.RunnerID, []control.SessionState{control.StateCreating, control.StateResuming}) if err != nil { return 0, err } - return r.CapacityTotal - r.CapacityUsed - len(creating), nil + return control.AvailableRunnerSlots(r, creating), nil } return 0, nil } diff --git a/controlapp/sessions_test.go b/controlapp/sessions_test.go index f62821c..80bafa3 100644 --- a/controlapp/sessions_test.go +++ b/controlapp/sessions_test.go @@ -230,7 +230,7 @@ func (r *sessionStubSessionRepo) Transition(ctx context.Context, ws control.Work if !ok || s.WorkspaceID != ws { return control.ErrNotFound } - if !slices.Contains(from, s.State) { + if !slices.Contains(from, s.State) || (opts.ExpectedPlacementGeneration != nil && s.PlacementGeneration != *opts.ExpectedPlacementGeneration) { return control.ErrConflict } s.State = to @@ -415,6 +415,8 @@ func (r *sessionStubEventRecorder) Record(ctx context.Context, e control.Event) } type sessionStubTransport struct { + onDispatch func(runner.ToRunner) + log *sessionCallLog res runner.FromRunner err error @@ -424,6 +426,9 @@ type sessionStubTransport struct { func (t *sessionStubTransport) Dispatch(ctx context.Context, pool control.PoolID, id control.RunnerID, m runner.ToRunner) (runner.FromRunner, error) { t.log.add("transport:dispatch:" + m.Type) + if t.onDispatch != nil { + t.onDispatch(m) + } if t.err != nil { return runner.FromRunner{}, t.err } @@ -1693,7 +1698,7 @@ func TestARunnerSConflictIsAConflictNotAnOutage(t *testing.T) { } }) - t.Run("the row is left exactly where it was", func(t *testing.T) { + t.Run("the cold claim remains pending for reconciliation", func(t *testing.T) { f := newSessionFixtureFull(t) f.transport.res = runner.FromRunner{OK: false, Conflict: true} f.fleet.runners = coldResumeCapacity() @@ -1706,12 +1711,11 @@ func TestARunnerSConflictIsAConflictNotAnOutage(t *testing.T) { if err != nil { t.Fatalf("re-read: %v", err) } - if row.State != control.StateSuspendedCold { - t.Fatalf("state after a refused resume = %q, want suspended_cold — the refusal must not "+ - "move the row, or the session is left claiming to run on a container that is stopping", row.State) + if row.State != control.StateResuming || row.PlacementGeneration != 2 { + t.Fatal("refusal lost the committed cold-resume claim") } - if f.log.hasPrefix("sessions:transition") { - t.Fatal("a refused resume transitioned the row") + if f.log.hasPrefix("sessions:transition:running") { + t.Fatal("refusal marked the session running") } }) diff --git a/controlapp/uow_test.go b/controlapp/uow_test.go index c46db7c..e7cdf02 100644 --- a/controlapp/uow_test.go +++ b/controlapp/uow_test.go @@ -294,8 +294,8 @@ func TestResumeSessionCommitsTransitionAndEventTogether(t *testing.T) { if _, err := fx.svc.ResumeSession(uowCtx, sessionTestScope(), control.ResumeSession{ID: "sess_one"}); err != nil { t.Fatal(err) } - if fx.uow.runs != 1 || fx.repo.transitionDepth != 1 || fx.rec.recordDepth != 1 || fx.woke != 1 { - t.Fatalf("runs %d, transition at depth %d, record at depth %d, woke %d; want 1/1/1/1", + if fx.uow.runs != 2 || fx.repo.transitionDepth != 1 || fx.rec.recordDepth != 1 || fx.woke != 1 { + t.Fatalf("runs %d, transition at depth %d, record at depth %d, woke %d; want 2/1/1/1", fx.uow.runs, fx.repo.transitionDepth, fx.rec.recordDepth, fx.woke) } if ev := fx.rec.last(t); ev.PlacementGeneration != 4 { diff --git a/docs/cli-v0-contract.md b/docs/cli-v0-contract.md index 26f040d..2efacd3 100644 --- a/docs/cli-v0-contract.md +++ b/docs/cli-v0-contract.md @@ -121,7 +121,7 @@ session over a network blip. | API `state` | Displayed | | --- | --- | -| `queued`, `creating` | Starting | +| `queued`, `creating`, `resuming` | Starting | | `running` | Running — **regardless of `child_exit_code`** | | `suspended_warm`, `suspended_cold` | Stopped | | `failed`, `dead` | Failed | @@ -978,3 +978,12 @@ server event separately from creation time and process state. A last event is not proof that an agent is actively working. Status compute checks include `facts.status` and `facts.health` verbatim. Missing required facts (including an unresolved default environment or a failed workspace lookup) fail readiness. + +Cold resume commits `resuming` and its new placement generation before dispatch. +Concurrent resume requests return conflict. A lost or refused reply retains that +claim until exact-generation runner status resolves it; unversioned inventory +cannot requeue it as a fresh create. After a two-minute dispatch grace, the +existing scheduler safety pass checks status without launching a VM. A cold +status durably fences any delayed command for that generation. The existing +runner and workspace remain attached to the session throughout. Deletion is +allowed while resuming; attach waits for readiness and stop is unavailable. diff --git a/docs/design/2026-10-02-guest-reconnect-integration.md b/docs/design/2026-10-02-guest-reconnect-integration.md new file mode 100644 index 0000000..7e0517b --- /dev/null +++ b/docs/design/2026-10-02-guest-reconnect-integration.md @@ -0,0 +1,96 @@ +# Guest reconnect integration candidate + +This connects the shared reconnect protocol, runner admission, guest readiness, +and hosted and standalone authorization. B1 remains unqualified. The runner +opt-in is `--microvm-guest-reconnect` (or `RAINIER_MICROVM_GUEST_RECONNECT=1`), +disabled by default. Only a matching host dispatcher plus both runner capabilities +can negotiate enrollment. Volatile standalone stores never enable it. + +## Ownership and delivery + +A negotiated fresh boot consumes the initial configuration once and enrolls its +in-memory public key before publishing a relay. Subsequent connections retain a +bounded admission lease through proof, fresh configuration, token redemption, +and ready/ack. Accepted proof consumes a monotonically increasing connection +epoch, closes the old hub, drains admitted callbacks, and prevents queued old +callbacks from forwarding RPC. Replacement control connections do not inherit +an old connection's queue or authority. + +The runner-only `guest_reconnect_configuration` RPC returns a current non-secret +launch specification scoped to session and placement. It neither mints nor spends +a token. Request/proof frames remain capped at 4 KiB; the authorized configuration +response is capped at 4 MiB and rejects unknown fields, duplicate keys, nulls and +embedded bootstrap tokens. Hosted resolution checks current membership, policy, +and runner binding. The separately accepted proof supplies the narrow token. + +The guest refreshes ordinary environment values and secrets, including an empty +secret response. It refuses changes to the already-running launch (command, +repositories, setup/init, proxy, agent manifest/home layout) before redemption; +those changes require a fresh boot. It never relaunches the current agent or PTY +as part of reconnect. + +## Runner restart candidate + +Every current-version microVM driver holds an exclusive state-directory lock +before discovery and cleanup, including when reconnect is disabled. Older +binaries do not participate in this lock and must be drained before rollout. Fresh launches persist non-secret process identity: host boot ID, +process start time, and network namespace inode/device. Recovery requires current +control-plane placement authorization before binding a listener. Local checks +verify exact process arguments and UID/GID, persisted identity, jail and VMM +socket ownership, cgroup membership, and the adopted network slot. The host only +replaces an owned, refused Unix control socket. It never unlinks an active +listener or the surviving VMM's device socket, and never reloads a boot token. + +The recovery worker is sequential, belongs to an accepted control connection, +and retries failed candidates with fresh authorization. Failed, replaced or +canceled authority cannot install a placement. A recovered listener starts with +its initial delivery consumed, so every peer must prove its enrolled key. + +## Interrupted launches + +Fresh creation and cold resume durably own their namespace, disks and placement +before launch. An uncertain result retains a blocked record and consumes capacity; +retry and workspace deletion cannot reuse or remove that record's disks. + +The engine writes a synced launch marker before starting the jailer and publishes +its PID and process start time atomically. Signal authorization requires exact +native argv and the original process start time before every signal, including +shutdown escalation. Recovered-process waits recheck that lifetime. These checks +do not provide a Linux pidfd guarantee: numeric PID check/signal has an immediate +race window, though the delayed escalation window is closed. A missing or invalid PID on the same host boot is unknown, +not proof of exit. Recovery and teardown retain the jail and network resources +until the child is confirmed exited. If the process identity was lost in that +window, host inspection or a confirmed host reboot is required; an empty cgroup +or missing API socket cannot prove that a jailer has exited. Failed Stop keeps +process ownership and its single reaper for a later teardown attempt. + +## Remaining qualification gates + +- Qualify cold resume followed by runner restart against the explicit resuming + lifecycle, durable launch ownership and refreshed recovery identity. +- Exercise the built processes over real transports and review the integrated + branches independently and adversarially. +- Run the bounded B1 harness on a disposable KVM host: relay loss, gateway + replacement and runner restart, preserving PID/start time, PTY and waiting + child, followed by another real coding-agent tool turn. +- Record immutable artifacts, negative controls and independently verified + teardown. No RAM snapshot is part of this capability. + +Unit/race results and cross-compilation do not constitute that live qualification. + +## Standalone authority + +The PostgreSQL dispatcher locks the accepted runner generation, current placement, +bootstrap record and current owner before checking the configured allowlist and +resume policy. Enrollment/proof/token mutations commit with closed audit events. +Secrets are resolved after commit; resolution failure never restores a capability. +All responses fit a bounded 4 MiB budget, including JSON expansion of secret values. +The authority boundary refuses a caller's existing transaction so a successful +return really commits before delivery. + +Begin, accept, configuration and mint are runner-origin only. Legacy guest token +redemption remains guest-origin only at epoch zero; proof-issued redemption is +runner-origin only at a positive epoch. PostgreSQL bootstrap RPCs receive these +current-authority checks even before negotiation; this intentionally prevents a +legacy path from bypassing owner revocation or spending a proof-issued token. +The other legacy RPC methods retain their existing handlers. diff --git a/internal/controld/api.go b/internal/controld/api.go index 526b784..c5181b1 100644 --- a/internal/controld/api.go +++ b/internal/controld/api.go @@ -238,8 +238,8 @@ func (r *sessionRenderer) runnerHasRoom(name string) bool { } // freeSlots is free capacity per connected runner: its reported total less -// its reported use less the sessions it is currently creating, whose slots -// the runner has not counted yet. A store that cannot answer any part of it +// its reported use less pending creates/resumes absent from that same usage +// observation. A store that cannot answer any part of it // yields an empty map rather than a partial one — half a capacity picture is // not a smaller truth, it is a different fleet. func (r *sessionRenderer) freeSlots() map[string]int { @@ -254,12 +254,12 @@ func (r *sessionRenderer) freeSlots() map[string]int { if !row.Connected { continue } - creating, err := fleet.SessionsOnRunner(r.ctx, installPool, row.ID, []control.SessionState{control.StateCreating}) + creating, err := fleet.SessionsOnRunner(r.ctx, installPool, row.ID, []control.SessionState{control.StateCreating, control.StateResuming}) if err != nil { log.Printf("controld: rendering a session view: sessions creating on a runner: %v", err) return map[string]int{} } - free[string(row.ID)] = row.CapacityTotal - row.CapacityUsed - len(creating) + free[string(row.ID)] = control.AvailableRunnerSlots(row, creating) } return free } diff --git a/internal/controld/api_test.go b/internal/controld/api_test.go index 0156b4e..afa2a18 100644 --- a/internal/controld/api_test.go +++ b/internal/controld/api_test.go @@ -2190,16 +2190,15 @@ func TestResumeSession(t *testing.T) { t.Errorf("the runner's own words reached the client: %s", raw) } - // And the row is left exactly where it was: the refusal moves - // nothing, so the session is not left claiming to run on a container - // that is stopping. + // A refusal preserves the claimed placement for reconciliation; it + // must not advertise the session as running. after := doRequest(t, ts, http.MethodGet, "/v0/sessions/sess_res_stopping", tok, nil, nil) var body v0wire.SessionEnvelope if err := json.Unmarshal([]byte(readBody(t, after)), &body); err != nil { t.Fatalf("decode: %v", err) } - if body.Session.State != string(control.StateSuspendedCold) { - t.Errorf("state after the refused resume = %q, want suspended_cold", body.Session.State) + if body.Session.State != string(control.StateResuming) { + t.Errorf("state after the refused resume = %q, want resuming", body.Session.State) } }) diff --git a/internal/controld/capacity_render_test.go b/internal/controld/capacity_render_test.go new file mode 100644 index 0000000..145e3bc --- /dev/null +++ b/internal/controld/capacity_render_test.go @@ -0,0 +1,43 @@ +package controld + +import ( + "context" + "github.com/tokencanopy/rainier/control" + "testing" +) + +func TestRenderedCapacityUsesExactPendingPlacements(t *testing.T) { + for _, tc := range []struct { + name string + state control.SessionState + used, total int + observed uint64 + want int + }{ + {"unobserved_resume", control.StateResuming, 0, 1, 0, 0}, + {"counted_resume", control.StateResuming, 1, 2, 1, 1}, + {"stale_resume", control.StateResuming, 1, 2, 2, 0}, + {"legacy_resume", control.StateResuming, 1, 2, 0, 0}, + {"counted_create", control.StateCreating, 1, 2, 1, 1}, + } { + t.Run(tc.name, func(t *testing.T) { + srv, st, _ := newTestControld(t) + ctx := context.Background() + row, err := st.Sessions().CreateSession(ctx, installWorkspace, control.Session{ID: "capacity_test", PoolID: installPool, RunnerID: "runner_test", State: tc.state}) + if err != nil { + t.Fatal(err) + } + placements := map[control.SessionID]uint64{} + if tc.observed > 0 { + placements[row.ID] = tc.observed + } + if err := st.Fleet().UpsertRunner(ctx, installPool, control.Runner{ID: row.RunnerID, Connected: true, Generation: 1, CapacityUsed: tc.used, CapacityTotal: tc.total, CapacityPlacements: placements}); err != nil { + t.Fatal(err) + } + r := sessionRenderer{srv: srv, ctx: ctx} + if got := r.freeSlots()[string(row.RunnerID)]; got != tc.want { + t.Fatalf("free=%d, want %d", got, tc.want) + } + }) + } +} diff --git a/internal/controld/controld.go b/internal/controld/controld.go index 894431b..3800805 100644 --- a/internal/controld/controld.go +++ b/internal/controld/controld.go @@ -292,10 +292,12 @@ func (s *Server) compose() error { events control.EventRecorder = s.st uow control.UnitOfWork = s.st ) + _, guestReconnect := s.st.(GuestReconnectAuthorityStore) fleetSvc, err := controlapp.NewFleetService(controlapp.FleetOptions{ Authorizer: auth, Sessions: sessions, Environments: envs, Fleet: fleet, Pools: pools, Transport: s.transport, Events: events, Clock: clock, IDs: ids, SafetyInterval: fleetSafetyInterval, + GuestReconnect: guestReconnect, LaunchMaterial: launchMaterial{st: s.st, key: s.cfg.SecretsKey}, // The self-hosted bounds for a hook whose environment declares none, // exactly as the old scheduler applied them (api.go). diff --git a/internal/controld/memstore.go b/internal/controld/memstore.go index 4e8615a..23c1c4c 100644 --- a/internal/controld/memstore.go +++ b/internal/controld/memstore.go @@ -6,6 +6,7 @@ import ( "context" "encoding/base64" "fmt" + "maps" "slices" "sort" "strconv" @@ -249,6 +250,7 @@ func cloneControlEnvironment(e control.Environment) control.Environment { func cloneControlRunner(r control.Runner) control.Runner { cp := r cp.Capabilities = slices.Clone(r.Capabilities) + cp.CapacityPlacements = maps.Clone(r.CapacityPlacements) return cp } @@ -397,7 +399,7 @@ func (r memSessions) Transition(ctx context.Context, ws control.WorkspaceID, id if !ok { return control.ErrNotFound } - if !slices.Contains(from, s.State) { + if !slices.Contains(from, s.State) || (opts.ExpectedPlacementGeneration != nil && s.PlacementGeneration != *opts.ExpectedPlacementGeneration) { return control.ErrConflict } s.State = to @@ -793,6 +795,9 @@ type memFleet struct{ m *memStore } // connection changes nothing at all, rather than half-overwriting the // current one's view of its own capacity. func (r memFleet) UpsertRunner(ctx context.Context, pool control.PoolID, run control.Runner) error { + if err := control.ValidateCapacityPlacements(run.CapacityUsed, run.CapacityPlacements); err != nil { + return err + } if pool == "" { return control.ErrInvalid } diff --git a/internal/controld/pgstore/fleet.go b/internal/controld/pgstore/fleet.go index ecbdbe0..d5598b8 100644 --- a/internal/controld/pgstore/fleet.go +++ b/internal/controld/pgstore/fleet.go @@ -16,17 +16,18 @@ import ( // workspace that pool serves. type pgFleet struct{ s *Store } -const runnerCols = `name, pool_id, capacity_used, capacity_total, connected, generation, capabilities, last_seen_at` +const runnerCols = `name, pool_id, capacity_used, capacity_total, connected, generation, capabilities, last_seen_at, capacity_placements` func scanControlRunner(row rowScanner) (control.Runner, error) { var ( - r control.Runner - id, pool string - generation int64 - capBytes []byte - lastSeen *time.Time + r control.Runner + id, pool string + generation int64 + capBytes []byte + placementBytes []byte + lastSeen *time.Time ) - if err := row.Scan(&id, &pool, &r.CapacityUsed, &r.CapacityTotal, &r.Connected, &generation, &capBytes, &lastSeen); err != nil { + if err := row.Scan(&id, &pool, &r.CapacityUsed, &r.CapacityTotal, &r.Connected, &generation, &capBytes, &lastSeen, &placementBytes); err != nil { return control.Runner{}, err } r.ID = control.RunnerID(id) @@ -37,6 +38,11 @@ func scanControlRunner(row rowScanner) (control.Runner, error) { return control.Runner{}, err } } + if len(placementBytes) > 0 { + if err := json.Unmarshal(placementBytes, &r.CapacityPlacements); err != nil { + return control.Runner{}, err + } + } if lastSeen != nil { r.LastSeenAt = *lastSeen } @@ -48,6 +54,14 @@ func scanControlRunner(row rowScanner) (control.Runner, error) { // rather than half-overwriting the current connection's view of its own // capacity. func (r pgFleet) UpsertRunner(ctx context.Context, pool control.PoolID, run control.Runner) error { + if err := control.ValidateCapacityPlacements(run.CapacityUsed, run.CapacityPlacements); err != nil { + return err + } + placements, err := json.Marshal(run.CapacityPlacements) + if err != nil { + return control.ErrInvalid + } + if pool == "" { return control.ErrInvalid } @@ -60,15 +74,15 @@ func (r pgFleet) UpsertRunner(ctx context.Context, pool control.PoolID, run cont lastSeen = &run.LastSeenAt } ct, err := r.s.q(ctx).Exec(ctx, ` - INSERT INTO runners (pool_id, name, capacity_used, capacity_total, connected, generation, capabilities, last_seen_at) - VALUES ($1, $2, $3, $4, $5, $6, $7, $8) + INSERT INTO runners (pool_id, name, capacity_used, capacity_total, connected, generation, capabilities, last_seen_at, capacity_placements) + VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9) ON CONFLICT (pool_id, name) DO UPDATE SET capacity_used = EXCLUDED.capacity_used, capacity_total = EXCLUDED.capacity_total, connected = EXCLUDED.connected, generation = EXCLUDED.generation, - capabilities = EXCLUDED.capabilities, last_seen_at = EXCLUDED.last_seen_at + capabilities = EXCLUDED.capabilities, last_seen_at = EXCLUDED.last_seen_at, capacity_placements=EXCLUDED.capacity_placements WHERE runners.generation <= EXCLUDED.generation`, string(pool), string(run.ID), run.CapacityUsed, run.CapacityTotal, run.Connected, - int64(run.Generation), caps, lastSeen) + int64(run.Generation), caps, lastSeen, placements) if err != nil { return unavailable("upsert runner", err) } diff --git a/internal/controld/pgstore/migrations/0017_capacity_placements.sql b/internal/controld/pgstore/migrations/0017_capacity_placements.sql new file mode 100644 index 0000000..0bb32ea --- /dev/null +++ b/internal/controld/pgstore/migrations/0017_capacity_placements.sql @@ -0,0 +1,2 @@ +-- Exact placements counted in the same runner usage observation. +ALTER TABLE runners ADD COLUMN capacity_placements jsonb NOT NULL DEFAULT '{}'::jsonb; diff --git a/internal/controld/pgstore/pgstore_test.go b/internal/controld/pgstore/pgstore_test.go index 54c31bb..5866348 100644 --- a/internal/controld/pgstore/pgstore_test.go +++ b/internal/controld/pgstore/pgstore_test.go @@ -272,14 +272,14 @@ func TestMigrate0003To0004AddsColumnsToLegacyRows(t *testing.T) { if want := embeddedMigrationVersions(t); !slices.Equal(applied, want) { t.Fatalf("schema_migrations = %v, want every embedded migration in order %v", applied, want) } - // This release's head is 16: a database that stopped at 0003 runs the + // This release's head is 17: a database that stopped at 0003 runs the // expand step (0007), the contract step (0008), the events table // (0009), the agent credentials table (0010), the tombstone (0011), the // durable revoke fence (0012), the controller lease (0013), the exec // event's command name (0014), the session bootstrap token (0015), and // guest reconnect authorization (0016) in the same start. - if head := applied[len(applied)-1]; head != 16 { - t.Fatalf("head migration = %d, want 16", head) + if head := applied[len(applied)-1]; head != 17 { + t.Fatalf("head migration = %d, want 17", head) } // The legacy session survived, and its new columns read as "never exited" diff --git a/internal/controld/pgstore/reconnect.go b/internal/controld/pgstore/reconnect.go index f32ab00..b7c3af4 100644 --- a/internal/controld/pgstore/reconnect.go +++ b/internal/controld/pgstore/reconnect.go @@ -36,7 +36,7 @@ func (r pgGuestReconnects) lockScope(ctx context.Context, b control.GuestReconne if err != nil { return control.ErrUnavailable } - err = r.s.q(ctx).QueryRow(ctx, `SELECT 1 FROM sessions WHERE workspace_id=$1 AND id=$2 AND pool_id=$3 AND runner=$4 AND placement_generation=$5 AND state IN ('creating','running') FOR UPDATE`, string(b.WorkspaceID), string(b.SessionID), string(b.PoolID), string(b.RunnerID), int64(b.PlacementGeneration)).Scan(&found) + err = r.s.q(ctx).QueryRow(ctx, `SELECT 1 FROM sessions WHERE workspace_id=$1 AND id=$2 AND pool_id=$3 AND runner=$4 AND placement_generation=$5 AND state IN ('creating','resuming','running') FOR UPDATE`, string(b.WorkspaceID), string(b.SessionID), string(b.PoolID), string(b.RunnerID), int64(b.PlacementGeneration)).Scan(&found) if errors.Is(err, pgx.ErrNoRows) { return control.ErrReconnectFenced } diff --git a/internal/controld/pgstore/reconnect_authority.go b/internal/controld/pgstore/reconnect_authority.go new file mode 100644 index 0000000..486702a --- /dev/null +++ b/internal/controld/pgstore/reconnect_authority.go @@ -0,0 +1,55 @@ +package pgstore + +import ( + "context" + "errors" + + "github.com/jackc/pgx/v5" + "github.com/tokencanopy/rainier/control" + "github.com/tokencanopy/rainier/internal/controld" +) + +var _ controld.GuestReconnectAuthorityStore = (*Store)(nil) + +// WithGuestReconnectAuthority holds runner, placement, bootstrap and owner rows +// stable while the host checks current policy and mutates its capability. +func (s *Store) WithGuestReconnectAuthority(ctx context.Context, b control.GuestReconnectScope, fn func(context.Context, controld.GuestReconnectAuthority) error) error { + if _, nested := ctx.Value(txKey{}).(pgx.Tx); nested { + return control.ErrReconnectInvalid + } + if fn == nil { + return control.ErrReconnectInvalid + } + return s.Run(ctx, func(ctx context.Context) error { + if err := (pgGuestReconnects{s}).lockScope(ctx, b); err != nil { + return err + } + a := controld.GuestReconnectAuthority{} + var err error + a.Session, err = s.Sessions().GetSession(ctx, b.WorkspaceID, b.SessionID) + if err != nil { + return control.ErrUnavailable + } + if a.Session.CreatorID == "" { + return control.ErrReconnectFenced + } + if _, err = s.q(ctx).Exec(ctx, `LOCK TABLE session_bootstraps IN ROW EXCLUSIVE MODE`); err != nil { + return control.ErrUnavailable + } + err = s.q(ctx).QueryRow(ctx, `SELECT guest_connection_epoch FROM session_bootstraps WHERE workspace_id=$1 AND session_id=$2 FOR UPDATE`, string(b.WorkspaceID), string(b.SessionID)).Scan(&a.Epoch) + if err != nil && !errors.Is(err, pgx.ErrNoRows) { + return control.ErrUnavailable + } + err = s.q(ctx).QueryRow(ctx, `SELECT id,github_id,login,role,created_at FROM users WHERE id=$1 FOR SHARE`, string(a.Session.CreatorID)).Scan(&a.User.ID, &a.User.GitHubID, &a.User.Login, &a.User.Role, &a.User.CreatedAt) + if errors.Is(err, pgx.ErrNoRows) { + return control.ErrReconnectFenced + } + if err != nil { + return control.ErrUnavailable + } + if err = s.q(ctx).QueryRow(ctx, `SELECT clock_timestamp()`).Scan(&a.Now); err != nil { + return control.ErrUnavailable + } + return fn(ctx, a) + }) +} diff --git a/internal/controld/pgstore/reconnect_authority_test.go b/internal/controld/pgstore/reconnect_authority_test.go new file mode 100644 index 0000000..f6980fd --- /dev/null +++ b/internal/controld/pgstore/reconnect_authority_test.go @@ -0,0 +1,69 @@ +package pgstore + +import ( + "context" + "errors" + "testing" + + "github.com/tokencanopy/rainier/control" + "github.com/tokencanopy/rainier/internal/controld" +) + +func TestReconnectAuthorityUsesCurrentOwnerAndSocket(t *testing.T) { + ctx := context.Background() + st := freshStore(t, startPostgres(t), t.Name()) + u, err := st.UpsertUser(ctx, 42, "member_test", "member") + if err != nil { + t.Fatal(err) + } + b := control.GuestReconnectScope{WorkspaceID: "ws_self_hosted", PoolID: "pool_self_hosted", SessionID: "session_authority_test", RunnerID: "runner_test", PlacementGeneration: 1, ConnectionGeneration: 1} + if _, err := st.Sessions().CreateSession(ctx, b.WorkspaceID, control.Session{ID: b.SessionID, CreatorID: control.ActorID(u.ID), PoolID: b.PoolID, RunnerID: b.RunnerID, State: control.StateRunning}); err != nil { + t.Fatal(err) + } + if err := st.Fleet().UpsertRunner(ctx, b.PoolID, control.Runner{ID: b.RunnerID, Generation: 1, Connected: true}); err != nil { + t.Fatal(err) + } + if err := st.Run(ctx, func(inner context.Context) error { + err := st.WithGuestReconnectAuthority(inner, b, func(context.Context, controld.GuestReconnectAuthority) error { + t.Fatal("nested transaction reached delivery authority") + return nil + }) + if !errors.Is(err, control.ErrReconnectInvalid) { + t.Fatalf("nested authority: %v", err) + } + return nil + }); err != nil { + t.Fatal(err) + } + called := false + err = st.WithGuestReconnectAuthority(ctx, b, func(_ context.Context, a controld.GuestReconnectAuthority) error { + called = true + if a.User.ID != u.ID || a.Session.ID != b.SessionID || a.Epoch != 0 || a.Now.IsZero() { + t.Fatal("incorrect authority") + } + return nil + }) + if err != nil || !called { + t.Fatalf("current authority refused: %v", err) + } + stale := b + stale.ConnectionGeneration = 2 + if err := st.WithGuestReconnectAuthority(ctx, stale, func(context.Context, controld.GuestReconnectAuthority) error { + t.Fatal("stale socket authorized") + return nil + }); !errors.Is(err, control.ErrReconnectFenced) { + t.Fatalf("stale socket: %v", err) + } + if _, err := st.pool.Exec(ctx, `UPDATE users SET login='revoked_test' WHERE id=$1`, u.ID); err != nil { + t.Fatal(err) + } + err = st.WithGuestReconnectAuthority(ctx, b, func(_ context.Context, a controld.GuestReconnectAuthority) error { + if a.User.Login != "revoked_test" { + t.Fatal("cached owner authority") + } + return control.ErrReconnectFenced + }) + if !errors.Is(err, control.ErrReconnectFenced) { + t.Fatalf("current owner refusal: %v", err) + } +} diff --git a/internal/controld/pgstore/reconnect_host_test.go b/internal/controld/pgstore/reconnect_host_test.go new file mode 100644 index 0000000..9687ca5 --- /dev/null +++ b/internal/controld/pgstore/reconnect_host_test.go @@ -0,0 +1,235 @@ +package pgstore + +import ( + "context" + "crypto/ed25519" + "crypto/rand" + "encoding/base64" + "encoding/json" + "net" + "net/http" + "net/http/httptest" + "os" + "os/exec" + "path/filepath" + "testing" + "time" + + "github.com/coder/websocket" + "github.com/coder/websocket/wsjson" + "github.com/tokencanopy/rainier/control" + "github.com/tokencanopy/rainier/internal/controld" + "github.com/tokencanopy/rainier/protocol/runner" +) + +func TestStandaloneReconnectRoundTrip(t *testing.T) { standaloneReconnectRoundTrip(t, false) } +func TestStandaloneReconnectShippingProcess(t *testing.T) { + if testing.Short() { + t.Skip("shipping process integration") + } + standaloneReconnectRoundTrip(t, true) +} +func standaloneReconnectRoundTrip(t *testing.T, shipping bool) { + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + st := freshStore(t, startPostgres(t), t.Name()) + user, err := st.UpsertUser(ctx, 42, "member_test", "member") + if err != nil { + t.Fatal(err) + } + endpoint := standaloneReconnectEndpoint(t, st, shipping) + var conn *websocket.Conn + for { + conn, _, err = websocket.Dial(ctx, endpoint+"/v0/runners/connect", &websocket.DialOptions{HTTPHeader: http.Header{"Authorization": []string{"Bearer runner_token_test"}}}) + if err == nil { + break + } + select { + case <-ctx.Done(): + t.Fatal("standalone process did not accept runner") + case <-time.After(20 * time.Millisecond): + } + } + defer conn.CloseNow() + if err := wsjson.Write(ctx, conn, runner.FromRunner{Type: "announce", Proto: runner.ProtocolVersion, Runner: "runner_test", Total: 4, Capabilities: []string{runner.CapabilityMicrovmV1, runner.CapabilityGuestReconnectV1}}); err != nil { + t.Fatal(err) + } + var accepted runner.ToRunner + if err := wsjson.Read(ctx, conn, &accepted); err != nil || accepted.Type != "accept" { + t.Fatal("runner was not accepted") + } + var sequence uint64 + ask := func(method string, payload any, host bool) runner.RPCEnvelope { + t.Helper() + sequence++ + id := sequence + if host { + id |= 1 << 63 + } + raw, err := json.Marshal(payload) + if err != nil { + t.Fatal(err) + } + if err := wsjson.Write(ctx, conn, runner.FromRunner{Type: "session_req", Session: "session_host_test", Generation: accepted.Generation, Total: 4, RPC: &runner.RPCEnvelope{ID: id, Method: method, Payload: raw}}); err != nil { + t.Fatal(err) + } + for { + var msg runner.ToRunner + if err := wsjson.Read(ctx, conn, &msg); err != nil { + t.Fatal(err) + } + if msg.Type == "session_rpc" && msg.RPC != nil && msg.RPC.ID == id { + return *msg.RPC + } + } + } + // The response follows initial inventory reconciliation. + ask("unknown_test", map[string]int{}, true) + if _, err := st.Sessions().CreateSession(ctx, "ws_self_hosted", control.Session{ID: "session_host_test", CreatorID: control.ActorID(user.ID), PoolID: "pool_self_hosted", RunnerID: "runner_test", State: control.StateRunning}); err != nil { + t.Fatal(err) + } + protocol := map[string]uint64{"protocol": 1} + mint := ask(runner.MethodMintSessionBootstrap, protocol, true) + var token struct { + Token string `json:"token"` + } + if !mint.OK || json.Unmarshal(mint.Payload, &token) != nil || token.Token == "" { + t.Fatal("mint failed") + } + pub, priv, err := ed25519.GenerateKey(rand.Reader) + if err != nil { + t.Fatal(err) + } + if _, err := st.pool.Exec(ctx, `ALTER TABLE events ADD CONSTRAINT reject_enrollment_test CHECK (action <> 'guest_enroll')`); err != nil { + t.Fatal(err) + } + enrollment := runner.GuestReconnectEnrollRequest{Protocol: 1, Token: token.Token, BootEpoch: "boot_test", PublicKey: base64.RawURLEncoding.EncodeToString(pub)} + if ask(runner.MethodEnrollGuestReconnect, enrollment, false).OK { + t.Fatal("enrollment escaped failed audit") + } + if _, err := st.pool.Exec(ctx, `ALTER TABLE events DROP CONSTRAINT reject_enrollment_test`); err != nil { + t.Fatal(err) + } + enroll := ask(runner.MethodEnrollGuestReconnect, enrollment, false) + if !enroll.OK { + t.Fatal("enrollment refused") + } + if ask(runner.MethodBeginGuestReconnect, protocol, false).OK { + t.Fatal("guest originated challenge") + } + if ask(runner.MethodBeginGuestReconnect, json.RawMessage(`{"protocol":1,"protocol":1}`), true).OK { + t.Fatal("duplicate protocol accepted") + } + begin := ask(runner.MethodBeginGuestReconnect, protocol, true) + var challenge runner.GuestReconnectChallenge + if !begin.OK || json.Unmarshal(begin.Payload, &challenge) != nil { + t.Fatal("challenge refused") + } + message, err := challenge.SigningMessage() + if err != nil { + t.Fatal(err) + } + proof := runner.GuestReconnectAcceptRequest{Protocol: 1, AttemptID: challenge.AttemptID, Signature: base64.RawURLEncoding.EncodeToString(ed25519.Sign(priv, message))} + wrong := proof + _, otherKey, err := ed25519.GenerateKey(rand.Reader) + if err != nil { + t.Fatal(err) + } + wrong.Signature = base64.RawURLEncoding.EncodeToString(ed25519.Sign(otherKey, message)) + if ask(runner.MethodAcceptGuestReconnect, wrong, true).OK { + t.Fatal("wrong proof accepted") + } + answer := ask(runner.MethodAcceptGuestReconnect, proof, true) + var capability runner.GuestReconnectAcceptResponse + if !answer.OK || json.Unmarshal(answer.Payload, &capability) != nil || capability.Epoch != 1 { + t.Fatal("proof refused") + } + if ask(runner.MethodAcceptGuestReconnect, proof, true).OK { + t.Fatal("proof replay accepted") + } + if !ask(runner.MethodGuestReconnectConfiguration, protocol, true).OK { + t.Fatal("configuration refused") + } + redeem := map[string]any{"protocol": 1, "token": capability.Token} + if ask(runner.MethodFetchSessionSecrets, redeem, false).OK { + t.Fatal("guest spent proof-issued token") + } + if _, err := st.pool.Exec(ctx, `UPDATE users SET login='revoked_test' WHERE id=$1`, user.ID); err != nil { + t.Fatal(err) + } + if ask(runner.MethodFetchSessionSecrets, redeem, true).OK { + t.Fatal("revoked owner redeemed") + } + if _, err := st.pool.Exec(ctx, `UPDATE users SET login='member_test' WHERE id=$1`, user.ID); err != nil { + t.Fatal(err) + } + if !ask(runner.MethodFetchSessionSecrets, redeem, true).OK { + t.Fatal("authorized redemption refused") + } + if ask(runner.MethodFetchSessionSecrets, redeem, true).OK { + t.Fatal("token replay accepted") + } + if _, err := st.Sessions().CreateSession(ctx, "ws_self_hosted", control.Session{ID: "session_negotiated_test", CreatorID: control.ActorID(user.ID), PoolID: "pool_self_hosted", State: control.StateQueued}); err != nil { + t.Fatal(err) + } + for { + var message runner.ToRunner + if err := wsjson.Read(ctx, conn, &message); err != nil { + t.Fatal("negotiated placement was not dispatched") + } + if message.Type != "create" || message.Session != "session_negotiated_test" { + continue + } + if message.Spec == nil || message.Spec.GuestReconnect != runner.GuestReconnectProtocol { + t.Fatal("capable host did not negotiate recovery") + } + if err := wsjson.Write(ctx, conn, runner.FromRunner{Type: "result", ReqID: message.ReqID, OK: true, Generation: accepted.Generation, Total: 4}); err != nil { + t.Fatal(err) + } + break + } + +} + +func standaloneReconnectEndpoint(t *testing.T, st *Store, shipping bool) string { + t.Helper() + if !shipping { + server, err := controld.New(st, controld.Config{RunnerToken: "runner_token_test", SecretsKey: [32]byte{1}, Members: []string{"member_test"}, ExternalURL: "http://127.0.0.1:1"}) + if err != nil { + t.Fatal(err) + } + loopCtx, stop := context.WithCancel(context.Background()) + t.Cleanup(stop) + go server.Run(loopCtx) + serverHTTP := httptest.NewServer(server.Handler()) + t.Cleanup(serverHTTP.Close) + return serverHTTP.URL + } + dir := t.TempDir() + binary := filepath.Join(dir, "controld") + build := exec.Command("go", "build", "-o", binary, "../../../cmd/controld") + if output, err := build.CombinedOutput(); err != nil { + t.Fatalf("build controld: %v\n%s", err, output) + } + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + address := listener.Addr().String() + listener.Close() + endpoint := "http://" + address + child := exec.Command(binary, "--listen", address, "--external-url", endpoint, "--members", "member_test") + child.Env = append(os.Environ(), "RAINIER_DB="+st.pool.Config().ConnString(), "RAINIER_RUNNER_TOKEN=runner_token_test", "RAINIER_SECRETS_KEY=0100000000000000000000000000000000000000000000000000000000000000") + logFile, err := os.OpenFile(filepath.Join(dir, "controld.log"), os.O_CREATE|os.O_WRONLY, 0600) + if err != nil { + t.Fatal(err) + } + child.Stdout = logFile + child.Stderr = logFile + if err := child.Start(); err != nil { + logFile.Close() + t.Fatal(err) + } + t.Cleanup(func() { _ = child.Process.Kill(); _ = child.Wait(); _ = logFile.Close() }) + return endpoint +} diff --git a/internal/controld/pgstore/sessions.go b/internal/controld/pgstore/sessions.go index acf9987..c7eb289 100644 --- a/internal/controld/pgstore/sessions.go +++ b/internal/controld/pgstore/sessions.go @@ -315,8 +315,9 @@ func (r pgSessions) Transition(ctx context.Context, ws control.WorkspaceID, id c error = COALESCE($3::text, error), image = COALESCE($7::text, image), updated_at = now(), last_event_at = now() - WHERE workspace_id = $4 AND id = $5 AND state = ANY($6)`, - string(to), runner, opts.Error, string(ws), string(id), fromStrs, opts.Image) + WHERE workspace_id = $4 AND id = $5 AND state = ANY($6) + AND ($8::bigint IS NULL OR placement_generation = $8)`, + string(to), runner, opts.Error, string(ws), string(id), fromStrs, opts.Image, opts.ExpectedPlacementGeneration) if err != nil { return unavailable("transition", err) } diff --git a/internal/controld/reconnect_authority.go b/internal/controld/reconnect_authority.go new file mode 100644 index 0000000..775413c --- /dev/null +++ b/internal/controld/reconnect_authority.go @@ -0,0 +1,24 @@ +package controld + +import ( + "context" + "time" + + "github.com/tokencanopy/rainier/control" +) + +// GuestReconnectAuthority is the current, locked host authority for one request. +// Now is sampled after database locks; Epoch distinguishes proof-issued tokens. +type GuestReconnectAuthority struct { + Session control.Session + User User + Epoch uint64 + Now time.Time +} + +// GuestReconnectAuthorityStore is optional: volatile stores cannot negotiate +// guest recovery. The callback and all capability mutations commit together. +type GuestReconnectAuthorityStore interface { + GuestReconnects() control.GuestReconnectStore + WithGuestReconnectAuthority(context.Context, control.GuestReconnectScope, func(context.Context, GuestReconnectAuthority) error) error +} diff --git a/internal/controld/runners.go b/internal/controld/runners.go index 8984b14..1731118 100644 --- a/internal/controld/runners.go +++ b/internal/controld/runners.go @@ -108,6 +108,9 @@ func (h runnerHost) Aside(ctx context.Context, b runnerplane.Binding, gen uint64 // owner's authority (srpc.go). The plane sets the answer's id and method and // sends it back down. func (h runnerHost) SessionRequest(ctx context.Context, b runnerplane.Binding, id control.SessionID, env runner.RPCEnvelope) runner.RPCEnvelope { + if host, ok := h.srv.st.(GuestReconnectAuthorityStore); ok && reconnectMethod(env.Method) { + return h.srv.answerGuestReconnect(ctx, host, b, id, env) + } return h.srv.authorizeSessionRequest(ctx, string(b.RunnerID), string(id), env) } diff --git a/internal/controld/sched_test.go b/internal/controld/sched_test.go index ce8a1fe..40f52ba 100644 --- a/internal/controld/sched_test.go +++ b/internal/controld/sched_test.go @@ -215,9 +215,13 @@ func TestPlacementPinQueuesWhenTheRunnerHasNoRoom(t *testing.T) { // The unpinned session behind it places on vm1... wantState(t, st, "sess_behind", control.StateCreating) - if got := rec.snapshot(); !sameSet(got, []string{"sess_behind"}) { - t.Fatalf("vm1 received creates for %v, want only sess_behind", got) - } + // Creating is committed before asynchronous dispatch reaches the runner. + eventually(t, 3*time.Second, func() error { + if got := rec.snapshot(); !sameSet(got, []string{"sess_behind"}) { + return fmt.Errorf("vm1 received creates for %v, want only sess_behind", got) + } + return nil + }) // ...while the pinned one stays queued and unplaced. got := getSession(t, st, "sess_blocked") if got.State != control.StateQueued || got.RunnerID != "" { @@ -238,9 +242,13 @@ func TestPlacementPinQueuesWhenTheRunnerHasNoRoom(t *testing.T) { startRun(t, s) wantState(t, st, "sess_any", control.StateCreating) - if got := rec.snapshot(); !sameSet(got, []string{"sess_any"}) { - t.Fatalf("vm1 received creates for %v, want only sess_any", got) - } + // Creating is committed before asynchronous dispatch reaches the runner. + eventually(t, 3*time.Second, func() error { + if got := rec.snapshot(); !sameSet(got, []string{"sess_any"}) { + return fmt.Errorf("vm1 received creates for %v, want only sess_any", got) + } + return nil + }) if got := getSession(t, st, "sess_hw"); got.State != control.StateQueued || got.RunnerID != "" { t.Fatalf("pinned session = %q on %q, want still queued and unplaced", got.State, got.RunnerID) } diff --git a/internal/controld/srpc_reconnect.go b/internal/controld/srpc_reconnect.go new file mode 100644 index 0000000..0a7f369 --- /dev/null +++ b/internal/controld/srpc_reconnect.go @@ -0,0 +1,189 @@ +package controld + +import ( + "context" + "encoding/json" + "errors" + "time" + + "github.com/tokencanopy/rainier/control" + "github.com/tokencanopy/rainier/controlapp" + "github.com/tokencanopy/rainier/protocol/runner" + "github.com/tokencanopy/rainier/runnerplane" +) + +func reconnectMethod(method string) bool { + switch method { + case runner.MethodEnrollGuestReconnect, runner.MethodBeginGuestReconnect, runner.MethodAcceptGuestReconnect, runner.MethodGuestReconnectConfiguration, runner.MethodMintSessionBootstrap, runner.MethodFetchSessionSecrets: + return true + } + return false +} + +type authorityClock struct{ now time.Time } + +func (c authorityClock) Now() time.Time { return c.now } + +func reconnectRefusal(id uint64, err error) runner.RPCEnvelope { + code := "unavailable" + switch { + case errors.Is(err, control.ErrReconnectInvalid): + code = "invalid" + case errors.Is(err, control.ErrReconnectExpired): + code = "expired" + case errors.Is(err, control.ErrReconnectFenced), errors.Is(err, control.ErrDenied): + code = "fenced" + } + return rpcRefusal(id, code) +} + +// answerGuestReconnect derives placement from storage and socket generation from +// the accepted connection. No peer payload may nominate either authority. +func (s *Server) answerGuestReconnect(ctx context.Context, host GuestReconnectAuthorityStore, b runnerplane.Binding, id control.SessionID, env runner.RPCEnvelope) runner.RPCEnvelope { + ctx, cancel := context.WithTimeout(ctx, 5*time.Second) + defer cancel() + hostOrigin := env.ID&(uint64(1)<<63) != 0 + if env.Method != runner.MethodEnrollGuestReconnect && env.Method != runner.MethodFetchSessionSecrets && !hostOrigin { + return reconnectRefusal(env.ID, control.ErrReconnectInvalid) + } + row, err := s.st.Sessions().GetSession(ctx, b.WorkspaceID, id) + if err != nil { + return reconnectRefusal(env.ID, control.ErrReconnectFenced) + } + binding := control.GuestReconnectScope{WorkspaceID: b.WorkspaceID, PoolID: b.PoolID, RunnerID: b.RunnerID, SessionID: id, PlacementGeneration: row.PlacementGeneration, ConnectionGeneration: b.ConnectionGeneration} + var answer any + var secrets bool + err = host.WithGuestReconnectAuthority(ctx, binding, func(ctx context.Context, a GuestReconnectAuthority) error { + role, ok := s.roleFor(a.User.Login) + if !ok { + return control.ErrReconnectFenced + } + a.User.Role = role + resource := control.Resource{Kind: control.ResourceSession, ID: string(a.Session.ID), WorkspaceID: a.Session.WorkspaceID, CreatorID: a.Session.CreatorID} + if err := (ownerOrAdmin{}).Authorize(withUser(ctx, a.User), userScope(a.User), control.ActionResume, resource); err != nil { + return control.ErrReconnectFenced + } + clock := authorityClock{a.Now} + service := controlapp.GuestReconnect{Store: host.GuestReconnects(), Clock: clock} + var action control.Action + switch env.Method { + case runner.MethodEnrollGuestReconnect: + req, err := runner.DecodeGuestReconnectEnrollRequest(env.Payload) + if err != nil { + return control.ErrReconnectInvalid + } + if err := service.Enroll(ctx, binding, req.Token, req.BootEpoch, req.PublicKey); err != nil { + return err + } + secrets = true + action = control.ActionGuestEnroll + case runner.MethodBeginGuestReconnect: + if _, err := runner.DecodeGuestReconnectBeginRequest(env.Payload); err != nil { + return control.ErrReconnectInvalid + } + challenge, err := service.Begin(ctx, binding) + if err != nil { + return err + } + answer = challenge + action = control.ActionGuestReconnectBegin + case runner.MethodAcceptGuestReconnect: + req, err := runner.DecodeGuestReconnectAcceptRequest(env.Payload) + if err != nil { + return control.ErrReconnectInvalid + } + epoch, token, err := service.Accept(ctx, binding, req.AttemptID, req.Signature) + if err != nil { + return err + } + answer = runner.GuestReconnectAcceptResponse{Epoch: epoch, Token: token, ExpiresInSec: uint32(controlapp.SessionBootstrapTTL.Seconds())} + action = control.ActionGuestReconnectAccept + case runner.MethodGuestReconnectConfiguration: + if _, err := runner.DecodeGuestReconnectBeginRequest(env.Payload); err != nil { + return control.ErrReconnectInvalid + } + var environment *control.Environment + if a.Session.EnvironmentID != "" { + current, err := s.st.Environments().GetEnvironment(ctx, a.Session.WorkspaceID, a.Session.EnvironmentID) + if err != nil { + return control.ErrUnavailable + } + environment = ¤t + } + spec, err := s.fleet.ResolveGuestReconnectSpec(ctx, a.Session, environment) + if err != nil { + return control.ErrUnavailable + } + config := runner.GuestReconnectConfiguration{Protocol: runner.GuestReconnectProtocol, SessionID: string(a.Session.ID), PlacementGeneration: a.Session.PlacementGeneration, Spec: spec} + encoded, err := json.Marshal(config) + if err != nil { + return control.ErrUnavailable + } + if _, err := runner.DecodeGuestReconnectConfiguration(encoded); err != nil { + return control.ErrUnavailable + } + answer = config + action = control.ActionGuestReconnectConfigure + case runner.MethodMintSessionBootstrap: + if _, err := runner.DecodeGuestReconnectBeginRequest(env.Payload); err != nil { + return control.ErrReconnectInvalid + } + token, err := (controlapp.SessionBootstrapMinter{Store: s.st.Bootstraps(), Clock: clock}).Mint(ctx, a.Session.WorkspaceID, a.Session.ID, a.Session.PlacementGeneration) + if err != nil { + return err + } + answer = sessionBootstrapAnswer{Token: token, ExpiresInSec: int(controlapp.SessionBootstrapTTL.Seconds())} + action = control.ActionGuestBootstrapMint + case runner.MethodFetchSessionSecrets: + req, err := runner.DecodeSessionBootstrapRedeemRequest(env.Payload) + if err != nil { + return control.ErrReconnectInvalid + } + if hostOrigin != (a.Epoch > 0) { + return control.ErrReconnectFenced + } + if err := s.st.Bootstraps().ConsumeSessionBootstrap(ctx, a.Session.WorkspaceID, a.Session.ID, controlapp.HashSessionBootstrapToken(req.Token), a.Session.PlacementGeneration, a.Now); err != nil { + return control.ErrReconnectInvalid + } + secrets = true + action = control.ActionGuestBootstrapRedeem + default: + return control.ErrReconnectInvalid + } + row = a.Session + return s.st.Record(ctx, control.Event{ID: (idGenerator{}).NewEventID(), WorkspaceID: a.Session.WorkspaceID, ActorID: a.Session.CreatorID, Action: action, Resource: resource, At: a.Now, PlacementGeneration: a.Session.PlacementGeneration}) + }) + if err != nil { + return reconnectRefusal(env.ID, err) + } + // Enrollment and spends are committed before secret resolution. A resolver + // failure cannot restore a capability, and its error never crosses the wire. + if secrets { + vars, err := s.sessionSecrets(ctx, row) + if err != nil { + return reconnectRefusal(env.ID, control.ErrUnavailable) + } + if vars == nil { + vars = map[string]string{} + } + answer = runner.GuestReconnectEnrollResponse{Env: vars} + } + if ctx.Err() != nil { + return reconnectRefusal(env.ID, control.ErrUnavailable) + } + payload, err := encodeReconnectAnswer(answer) + if err != nil { + return reconnectRefusal(env.ID, control.ErrUnavailable) + } + return runner.RPCEnvelope{ID: env.ID, Method: "resp", OK: true, Payload: payload} +} + +// Bound even secret-bearing answers below the shared runner socket frame limit; +// one oversized environment must not disconnect every session on that runner. +func encodeReconnectAnswer(answer any) ([]byte, error) { + payload, err := json.Marshal(answer) + if err != nil || len(payload) > runner.GuestReconnectConfigurationLimit { + return nil, control.ErrUnavailable + } + return payload, nil +} diff --git a/internal/controld/srpc_reconnect_test.go b/internal/controld/srpc_reconnect_test.go new file mode 100644 index 0000000..e27d757 --- /dev/null +++ b/internal/controld/srpc_reconnect_test.go @@ -0,0 +1,20 @@ +package controld + +import ( + "strings" + "testing" + + "github.com/tokencanopy/rainier/protocol/runner" +) + +func TestReconnectSecretResponseCannotBreakRunnerFrame(t *testing.T) { + // JSON escaping can exceed the wire budget even below the source byte count. + for _, value := range []string{strings.Repeat("x", runner.GuestReconnectConfigurationLimit), strings.Repeat("\x00", runner.GuestReconnectConfigurationLimit/5)} { + if payload, err := encodeReconnectAnswer(runner.GuestReconnectEnrollResponse{Env: map[string]string{"SYNTHETIC_TEST": value}}); err == nil || len(payload) != 0 { + t.Fatal("oversized environment escaped response bound") + } + } + if _, err := encodeReconnectAnswer(runner.GuestReconnectEnrollResponse{Env: map[string]string{"SYNTHETIC_TEST": "value"}}); err != nil { + t.Fatal(err) + } +} diff --git a/internal/driver/driver.go b/internal/driver/driver.go index 5ef6cd6..03670c4 100644 --- a/internal/driver/driver.go +++ b/internal/driver/driver.go @@ -6,12 +6,15 @@ import ( ) type Spec struct { - Name string // human label - Image string // OCI ref (v0 default a bash image) - Cmd []string // entrypoint override; empty = image default - DialURL string // runnerd URL the container's sessiond dials (relay) - SessionID string // stable id runnerd assigns; sessiond registers with it - EgressAllow []string // hostnames the session may reach + PlacementGeneration uint64 + + GuestReconnect uint64 // negotiated guest reconnect protocol; zero keeps legacy boot + Name string // human label + Image string // OCI ref (v0 default a bash image) + Cmd []string // entrypoint override; empty = image default + DialURL string // runnerd URL the container's sessiond dials (relay) + SessionID string // stable id runnerd assigns; sessiond registers with it + EgressAllow []string // hostnames the session may reach // ProxyURL, when non-empty, is the egress proxy the session's outbound // traffic must route through. Injected as both cases of HTTP_PROXY/ // HTTPS_PROXY (tools disagree on which they read: BusyBox wget and curl @@ -195,8 +198,11 @@ type Snapshot struct { // Listed pairs a driver handle with the session id it belongs to, for List's // bulk view of every rainier-managed resource. type Listed struct { - SessionID string - Handle Handle + PlacementGeneration uint64 + + GuestReconnect bool + SessionID string + Handle Handle } // CapabilityDriver is a driver that names portable capability tokens the diff --git a/internal/driver/microvm.go b/internal/driver/microvm.go index 92fa9b0..f7b1924 100644 --- a/internal/driver/microvm.go +++ b/internal/driver/microvm.go @@ -96,13 +96,15 @@ const ( // simulated engine reports every session as running while nothing executes, // which is the single worst failure mode this component has. type MicrovmOpts struct { - BaseRootfs string // default base ext4 rootfs image path (required in production) - KernelPath string // guest vmlinux kernel path (required in production) - StateDir string // directory holding instance sockets, metadata, and disks (always required) - TotalSlots int // maximum simultaneous active slot capacity - VCPU int // vCPUs per session; 0 means defaultMicrovmVCPU - MemoryMiB int // memory per session in MiB; 0 means defaultMicrovmMemoryMiB - VMMPath string // path to the Firecracker executable; empty means "firecracker" on PATH + // GuestReconnect enables the negotiated live-recovery capability; default off. + GuestReconnect bool + BaseRootfs string // default base ext4 rootfs image path (required in production) + KernelPath string // guest vmlinux kernel path (required in production) + StateDir string // directory holding instance sockets, metadata, and disks (always required) + TotalSlots int // maximum simultaneous active slot capacity + VCPU int // vCPUs per session; 0 means defaultMicrovmVCPU + MemoryMiB int // memory per session in MiB; 0 means defaultMicrovmMemoryMiB + VMMPath string // path to the Firecracker executable; empty means "firecracker" on PATH // SlotGuestCIDR and SlotUplinkCIDR are the two host-local ranges a // session's addresses are carved from, one /30 per slot (ADR-0003 §5.2). @@ -321,13 +323,18 @@ type DiskFormatter interface { // instanceRecord is the persistent metadata stored on disk for each microVM // instance. type instanceRecord struct { - ID string `json:"id"` - SessionID string `json:"session_id"` - State State `json:"state"` - Cold bool `json:"cold"` - Volume string `json:"volume"` - PID int `json:"pid"` - Cfg VMMConfig `json:"cfg"` + PlacementGeneration uint64 `json:"placement_generation,omitempty"` + RecoveryBlocked bool `json:"recovery_blocked,omitempty"` + + Identity guestHostIdentity `json:"guest_host_identity,omitempty"` + Reconnect bool `json:"guest_reconnect,omitempty"` + ID string `json:"id"` + SessionID string `json:"session_id"` + State State `json:"state"` + Cold bool `json:"cold"` + Volume string `json:"volume"` + PID int `json:"pid"` + Cfg VMMConfig `json:"cfg"` // The portable workspace checkpoint this session has, if any. All three are // PERSISTED, and they have to be: a cold resume after a runnerd restart is @@ -386,7 +393,9 @@ type instanceRecord struct { // in flight. A second Resume for the same id is refused rather than run // beside it: see Resume for what two concurrent cold ones would do to // each other's socket. - resuming bool + resuming bool + resumeDone chan struct{} + destroying bool // checkpointing is the same kind of claim for the workspace checkpoint a // cold suspend takes. Two of them on one instance would each send the guest @@ -456,6 +465,7 @@ func (rec *instanceRecord) bump() { rec.epoch++ } // Microvm implements driver.Driver for hardware-isolated microVMs. type Microvm struct { + stateLock *os.File // held until process exit; never unlink its inode mu sync.Mutex opts MicrovmOpts engine MicrovmEngine @@ -492,6 +502,23 @@ func NewMicrovm(opts MicrovmOpts) (*Microvm, error) { if opts.StateDir == "" { return nil, errors.New("microvm: a state directory is required (--microvm-state-dir / RAINIER_MICROVM_STATE_DIR): it holds every session's workspace and agent-home disk image, and a temp-directory default puts a tenant's files somewhere the host reaps") } + var stateLock *os.File + constructed := false + { // Every current-version writer participates, independent of capability. + if err := os.MkdirAll(opts.StateDir, microvmDirMode); err != nil { + return nil, err + } + var err error + stateLock, err = lockGuestState(opts.StateDir) + if err != nil { + return nil, err + } + defer func() { + if !constructed { + stateLock.Close() + } + }() + } if opts.TotalSlots <= 0 { opts.TotalSlots = 16 } @@ -565,6 +592,7 @@ func NewMicrovm(opts MicrovmOpts) (*Microvm, error) { } m := &Microvm{ + stateLock: stateLock, opts: opts, engine: engine, slots: slots, @@ -581,6 +609,7 @@ func NewMicrovm(opts MicrovmOpts) (*Microvm, error) { m.reclaimNetworkSlots() m.reclaimOrphanRootfs() m.reclaimRestoreScratch() + constructed = true return m, nil } @@ -874,10 +903,7 @@ func (m *Microvm) saveRecord(rec instanceRecord) error { if err != nil { return fmt.Errorf("marshal instance record: %w", err) } - if err := os.WriteFile(m.instanceMetaPath(rec.ID), data, microvmFileMode); err != nil { - return fmt.Errorf("write instance record: %w", err) - } - return nil + return atomicMetadata(m.instanceMetaPath(rec.ID), data) } func (m *Microvm) deleteInstanceRecord(id string) { @@ -911,6 +937,12 @@ func (m *Microvm) recoverDiskInstances() { continue } if st, err := m.engine.State(context.Background(), id); err == nil { + if rec.RecoveryBlocked && (st == VMMStateGone || st == VMMStateStopped) { + // An interrupted launch with no surviving process is cold, + // not a lost workspace. Keep the recorded placement fence. + rec.Cold, rec.RecoveryBlocked = true, false + rec.PID, rec.Identity = 0, guestHostIdentity{} + } reconcileState(&rec, st) } m.reassociateSlot(&rec) @@ -955,7 +987,9 @@ func (m *Microvm) reassociateSlot(rec *instanceRecord) { return } rec.slot = slot - applySlot(&rec.Cfg, slot) + if !rec.Reconnect { + applySlot(&rec.Cfg, slot) + } } // detachIdleSlot takes the network slot off a record that has stopped @@ -1295,7 +1329,7 @@ func sanitizeRef(ref string) (string, error) { func (m *Microvm) usedLocked() int { used := m.pending for _, inst := range m.instances { - if inst.State == StateRunning || (inst.State == StateSuspended && !inst.Cold) { + if inst.resuming || inst.State == StateRunning || (inst.State == StateSuspended && !inst.Cold) { used++ } } @@ -1423,9 +1457,15 @@ func buildGuestEnv(spec Spec) map[string]string { // runs with the mutex RELEASED. Holding it across engine.Launch froze // Inspect, List, Capacity and Destroy for every session on the host behind // one Firecracker that had not yet opened its socket. -func (m *Microvm) reserveSlot() (string, error) { +func (m *Microvm) reserveSlot(sessionID string) (string, error) { m.mu.Lock() defer m.mu.Unlock() + for _, rec := range m.instances { + if sessionID != "" && rec.SessionID == sessionID && rec.RecoveryBlocked { + return "", errors.New("microvm: workspace retained by an uncertain launch") + } + } + if used := m.usedLocked(); used >= m.opts.TotalSlots { return "", fmt.Errorf("no capacity: %d/%d", used, m.opts.TotalSlots) } @@ -1455,11 +1495,19 @@ func (m *Microvm) Create(ctx context.Context, spec Spec) (Handle, error) { // Before anything with a side effect, and before a slot is even // reserved: a create this host must not perform is refused rather than // half-performed. See refuseUnwithheldEnv. + if spec.GuestReconnect != 0 && (spec.GuestReconnect != runner.GuestReconnectProtocol || !m.opts.GuestReconnect) { + return Handle{}, errors.New("microvm: guest reconnect is not enabled") + } + if spec.GuestReconnect != 0 { + if _, ok := m.engine.(guestIdentityVerifier); !ok { + return Handle{}, errors.New("microvm: guest reconnect requires local VM identity verification") + } + } if err := refuseUnwithheldEnv(spec); err != nil { return Handle{}, err } - id, err := m.reserveSlot() + id, err := m.reserveSlot(spec.SessionID) if err != nil { return Handle{}, err } @@ -1613,30 +1661,49 @@ func (m *Microvm) launch(ctx context.Context, id string, spec Spec) (*instanceRe } applySlot(&cfg, slot) - // The jailer may create the child before Launch fails. Register this - // first so a later rollback stops a successfully launched VMM before - // removing its cgroup; failed Launch already reaps its own process. - undo = append(undo, func() { m.removeCgroup(id, cfg.CgroupPath) }) + rec := &instanceRecord{ + ID: id, SessionID: spec.SessionID, PlacementGeneration: spec.PlacementGeneration, + State: StateRunning, Volume: workspaceVolume(spec.SessionID), RecoveryBlocked: true, + Reconnect: spec.GuestReconnect == runner.GuestReconnectProtocol, Cfg: cfg, + slot: slot, boot: bootCfg, channel: channel, bootLive: true, boots: 1, + } + // Own every resource before a child can exist. Any uncertain rollback + // returns a blocked record alongside its error; Create accounts for it. + if err := m.saveRecord(persistable(rec)); err != nil { + return nil, err + } + refuse := func(cause error) (*instanceRecord, error) { + channel.close() + stopCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 5*time.Second) + defer cancel() + if err := m.engine.Stop(stopCtx, id); err != nil { + state, stateErr := m.engine.State(stopCtx, id) + if stateErr != nil || (state != VMMStateGone && state != VMMStateStopped) { + rec.PID, rec.Identity, rec.channel = m.engine.PID(id), guestHostIdentity{}, nil + rec.RecoveryBlocked = true + // The committed prelaunch intent remains safe if this write fails. + _ = m.saveRecord(persistable(rec)) + done = true + return rec, errors.New("microvm: uncertain create retained for cleanup") + } + } + m.removeCgroup(id, cfg.CgroupPath) + return nil, cause + } if err := m.engine.Launch(ctx, cfg); err != nil { - return nil, fmt.Errorf("launch microvm %s: %w", id, err) + return refuse(fmt.Errorf("launch microvm %s: %w", id, err)) } - undo = append(undo, func() { _ = m.engine.Stop(context.WithoutCancel(ctx), id) }) - - rec := &instanceRecord{ - ID: id, - SessionID: spec.SessionID, - State: StateRunning, - Volume: workspaceVolume(spec.SessionID), - PID: m.engine.PID(id), - Cfg: cfg, - slot: slot, - boot: bootCfg, - channel: channel, - bootLive: true, - boots: 1, + rec.PID = m.engine.PID(id) + if rec.Reconnect { + identity, err := m.engine.(guestIdentityVerifier).guestIdentity(cfg, rec.PID) + if err != nil { + return refuse(err) + } + rec.Identity = identity } + rec.RecoveryBlocked = false if err := m.saveRecord(persistable(rec)); err != nil { - return nil, fmt.Errorf("save instance metadata %s: %w", id, err) + return refuse(fmt.Errorf("save instance metadata %s: %w", id, err)) } done = true @@ -1744,13 +1811,30 @@ func (m *Microvm) Suspend(ctx context.Context, id string, warm bool) error { return nil } -func (m *Microvm) Resume(ctx context.Context, id string) (bool, error) { +func (m *Microvm) Resume(ctx context.Context, id string) (bool, error) { return m.resume(ctx, id, 0) } + +func (m *Microvm) ResumePlacement(ctx context.Context, id string, generation uint64) (bool, error) { + if generation == 0 { + return false, errors.New("microvm: invalid resume placement") + } + return m.resume(ctx, id, generation) +} + +func (m *Microvm) resume(ctx context.Context, id string, generation uint64) (bool, error) { m.mu.Lock() inst, ok := m.instances[id] if !ok { m.mu.Unlock() return false, fmt.Errorf("no such id %s", id) } + if inst.RecoveryBlocked || inst.destroying { + m.mu.Unlock() + return false, errors.New("microvm: unverified VM requires cleanup") + } + if generation != 0 && (generation < inst.PlacementGeneration || (generation == inst.PlacementGeneration && inst.Cold)) { + m.mu.Unlock() + return false, errors.New("microvm: stale resume placement") + } running := inst.State == StateRunning cold, bootLive, cfg := inst.Cold, inst.bootLive, inst.Cfg bootCfg, sessionID := inst.boot, inst.SessionID @@ -1780,7 +1864,22 @@ func (m *Microvm) Resume(ctx context.Context, id string) (bool, error) { m.mu.Unlock() return false, fmt.Errorf("resume of %s: a resume is already in flight for this instance", id) } + if generation != 0 { + before := inst.PlacementGeneration + inst.PlacementGeneration = generation + if err := m.saveRecord(persistable(inst)); err != nil { + inst.PlacementGeneration = before + m.mu.Unlock() + return false, err + } + } + if cold && m.usedLocked() >= m.opts.TotalSlots { + m.mu.Unlock() + return false, errors.New("microvm: no capacity for cold resume") + } inst.resuming = true + inst.resumeDone = make(chan struct{}) + claimed := inst // The boot number is taken HERE, under the same lock, and it advances // even for an attempt that fails: a failed launch may have left a socket // behind at that path, and an attempt that reuses a number is an attempt @@ -1808,13 +1907,14 @@ func (m *Microvm) Resume(ctx context.Context, id string) (bool, error) { } defer func() { m.mu.Lock() - if e, ok := m.instances[id]; ok { - e.resuming = false - } + claimed.resuming = false + close(claimed.resumeDone) + claimed.resumeDone = nil m.mu.Unlock() }() restarted := false + var identity guestHostIdentity var ( channel *guestChannel slot *netslot.Slot @@ -1829,12 +1929,47 @@ func (m *Microvm) Resume(ctx context.Context, id string) (bool, error) { // given back when it does not get there. defer func() { if slot != nil { + m.mu.Lock() + current, owned := m.instances[id] + owned = owned && current == claimed + rec := persistable(claimed) + m.mu.Unlock() + // Failure retains the conservative blocked launch intent on disk. + // Recovered engine evidence may release it only after observing exit. + if owned { + _ = m.saveRecord(rec) + } _ = m.slots.Release(context.WithoutCancel(ctx), slot) } if clonedRootfs { m.removeSessionRootfs(id) } }() + stopUnverified := func() bool { + stopCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 5*time.Second) + defer cancel() + stopErr := m.engine.Stop(stopCtx, id) + if stopErr != nil { + state, stateErr := m.engine.State(stopCtx, id) + if stateErr != nil || (state != VMMStateGone && state != VMMStateStopped) { + m.mu.Lock() + claimed.State, claimed.Cold, claimed.RecoveryBlocked = StateRunning, false, true + claimed.Identity = guestHostIdentity{} + claimed.PID, claimed.Cfg, claimed.slot, claimed.channel = m.engine.PID(id), cfg, slot, nil + claimed.bump() + m.instances[id] = claimed + rec := persistable(claimed) + slot, clonedRootfs = nil, false + m.mu.Unlock() + // If this write fails, the earlier blocked launch intent still + // owns the new namespace and forbids guest admission on restart. + _ = m.saveRecord(rec) + return false + } + } + m.removeCgroup(id, cfg.CgroupPath) + return true + } if cold { // A cold resume is a fresh boot, and a fresh boot needs the session's // whole configuration. This driver holds that in memory only @@ -1921,10 +2056,37 @@ func (m *Microvm) Resume(ctx context.Context, id string) (bool, error) { return false, fmt.Errorf("cold resume of %s: allocating a network slot: %w", id, err) } applySlot(&cfg, slot) + // Persist ownership BEFORE the engine can start a process. On a crash + // during Launch, startup must retain this namespace and rootfs even + // though no verified process identity has been published yet. + m.mu.Lock() + intent := persistable(inst) + intent.State, intent.Cold, intent.RecoveryBlocked = StateRunning, false, true + intent.Cfg, intent.PID, intent.Identity = cfg, 0, guestHostIdentity{} + intent = persistable(&intent) + m.mu.Unlock() + if err := m.saveRecord(intent); err != nil { + channel.close() + return false, err + } if err := m.engine.Launch(ctx, cfg); err != nil { channel.close() + if !stopUnverified() { + return false, errors.New("microvm: uncertain launch retained for cleanup") + } return false, fmt.Errorf("relaunch cold microvm %s: %w", id, err) } + if inst.Reconnect { + var err error + identity, err = m.engine.(guestIdentityVerifier).guestIdentity(cfg, m.engine.PID(id)) + if err != nil { + channel.close() + if !stopUnverified() { + return false, errors.New("microvm: unverified VM retained for cleanup") + } + return false, errors.New("microvm: new guest identity verification failed") + } + } restarted = true } else if err := m.engine.Resume(ctx, id); err != nil { return false, err @@ -1943,19 +2105,17 @@ func (m *Microvm) Resume(ctx context.Context, id string) (bool, error) { // and a network slot; and the deferred release is about to take that // slot back out from under it. So it is stopped here rather than // left running, and only then does the slot go. - if restarted { - if err := m.engine.Stop(context.WithoutCancel(ctx), id); err != nil { - log.Printf("microvm: %s was resumed onto a record that no longer exists and could not be stopped: %v", id, err) - } else { - m.removeCgroup(id, cfg.CgroupPath) - } + if restarted && !stopUnverified() { + return restarted, errors.New("microvm: orphaned launch retained for cleanup") } return restarted, fmt.Errorf("no such id %s", id) } inst.State = StateRunning inst.Cold = false + inst.RecoveryBlocked = false if restarted { inst.PID = m.engine.PID(id) + inst.Identity = identity if inst.channel != nil { inst.channel.close() } @@ -2350,12 +2510,32 @@ func (m *Microvm) destroyUnrecorded(ctx context.Context, id string) error { } func (m *Microvm) DestroyContainer(ctx context.Context, id string) error { + // A launch owns disks and a new namespace before it publishes them on the + // record. Wait for that handoff before stopping or releasing anything. m.mu.Lock() inst, ok := m.instances[id] + for ok && inst.resuming { + done := inst.resumeDone + m.mu.Unlock() + select { + case <-ctx.Done(): + return ctx.Err() + case <-done: + } + m.mu.Lock() + inst, ok = m.instances[id] + } if !ok { m.mu.Unlock() return m.destroyUnrecorded(ctx, id) } + if inst.destroying { + m.mu.Unlock() + return errors.New("microvm: teardown already in flight") + } + inst.destroying = true + defer func() { m.mu.Lock(); inst.destroying = false; m.mu.Unlock() }() + slot := inst.slot inst.slot = nil channel := inst.channel @@ -2426,6 +2606,15 @@ func (m *Microvm) RemoveWorkspace(_ context.Context, sessionID string) error { if sessionID == "" { return nil } + m.mu.Lock() + for _, rec := range m.instances { + if rec.SessionID == sessionID && (rec.RecoveryBlocked || rec.resuming || rec.State == StateRunning || (rec.State == StateSuspended && !rec.Cold)) { + m.mu.Unlock() + return errors.New("microvm: workspace is still owned by a live or uncertain VM") + } + } + m.mu.Unlock() + // A teardown is the one path where an unchecked id does real damage: // os.Remove of whatever "../../something" resolved to. An id that cannot // name a workspace is an error, not a removal. @@ -2460,7 +2649,7 @@ func (m *Microvm) Inspect(ctx context.Context, id string) (Handle, error) { return Handle{ID: id, State: StateGone}, nil } var idle *netslot.Slot - if stErr == nil && inst.epoch == epoch { + if stErr == nil && inst.epoch == epoch && !inst.resuming && !inst.destroying { reconcileState(inst, st) // A VM that went away without this driver parking it has stopped // occupying the host, and its slot has to go back with it. @@ -2509,15 +2698,17 @@ func (m *Microvm) List(ctx context.Context) ([]Listed, error) { out := make([]Listed, 0, len(m.instances)) var idle []*netslot.Slot for id, inst := range m.instances { - if st, ok := observed[id]; ok && inst.epoch == epochs[id] { + if st, ok := observed[id]; ok && inst.epoch == epochs[id] && !inst.resuming && !inst.destroying { reconcileState(inst, st) if s := detachIdleSlot(inst); s != nil { idle = append(idle, s) } } out = append(out, Listed{ - SessionID: inst.SessionID, - Handle: Handle{ID: id, State: inst.State}, + SessionID: inst.SessionID, + GuestReconnect: inst.Reconnect, + PlacementGeneration: inst.PlacementGeneration, + Handle: Handle{ID: id, State: inst.State}, }) } m.mu.Unlock() @@ -2608,14 +2799,16 @@ type FirecrackerEngine struct { // It is a VALUE and not a pointer: an engine that failed to construct // carries the zero range, whose forSlot answers an error for every index, // so no code path can reach a nil allocator — see uidRange. - uids uidRange - starter processStarter + uids uidRange + starter processStarter + startTime func(int) (uint64, error) // chown and link are the two filesystem operations the jail needs that an // ordinary test process cannot perform. Production is os.Chown and // os.Link; see FirecrackerOpts. chown func(path string, uid, gid int) error link func(oldname, newname string) error procs map[string]vmmProcess + waits map[string]chan error initErr error // kvm is the "can this host run a VM at all" check, as a field so the // jail tests can run on a machine without /dev/kvm. Production never @@ -2746,6 +2939,7 @@ func NewFirecrackerEngine(opts FirecrackerOpts) *FirecrackerEngine { jail: jail, uids: uids, starter: starter, + startTime: processStartTime, chown: chown, link: link, procs: make(map[string]vmmProcess), @@ -2893,6 +3087,10 @@ func (f *FirecrackerEngine) Launch(ctx context.Context, cfg VMMConfig) error { return errors.New("rootfs image path is required for Firecracker launch") } + if _, err := os.Stat(f.launchMarkerPath(cfg.ID)); !errors.Is(err, os.ErrNotExist) { + return errors.New("microvm: prior launch must be settled before starting another process") + } + // The jail, built before anything is started: a chrooted Firecracker can // only see what is already inside it. spec, err := f.jailSpecFor(cfg) @@ -2913,37 +3111,34 @@ func (f *FirecrackerEngine) Launch(ctx context.Context, cfg VMMConfig) error { // enters the session's network namespace (`--netns`), which is where the // slot's TAP device is — the device does not exist in the host's // namespace at all. + if err := atomicMetadata(f.launchMarkerPath(cfg.ID), []byte(hostLaunchBoot())); err != nil { + return err + } proc, err := f.starter.Start(f.jailerPath, jailerArgs(spec)) if err != nil { _ = f.removeJail(cfg.ID) + _ = f.removeLaunchEvidence(cfg.ID) return fmt.Errorf("start jailed firecracker %s: %w", cfg.ID, err) } - - if pid := proc.Pid(); pid > 0 { - pidDir := filepath.Join(f.stateDir, "instances", cfg.ID) - _ = os.MkdirAll(pidDir, microvmDirMode) - _ = os.WriteFile(f.pidFilePath(cfg.ID), []byte(strconv.Itoa(pid)), microvmFileMode) - } - + f.mu.Lock() + f.procs[cfg.ID] = proc + f.mu.Unlock() var initSuccess bool defer func() { - if initSuccess { - return + if !initSuccess { + f.finishFailedLaunch(cfg.ID, proc) } - // A launch that got part-way leaves a live VMM and a jail full of - // hard links. Both go, in that order: the process first, because - // removing the jail under a running Firecracker is how a VMM ends up - // writing into a directory that has been unlinked. The uid needs no - // undoing — it is the slot's, and the slot is the caller's to give - // back (see uidRange). - _ = proc.Kill() - _ = proc.Wait() - _ = f.removeJail(cfg.ID) - // And the pid file this launch wrote, which outlives the jail - // because it is not in it. A stale one is what State and PID read on - // the next boot. - _ = os.Remove(f.pidFilePath(cfg.ID)) }() + if pid := proc.Pid(); pid > 0 { + if err := f.saveProcessIdentity(cfg.ID, pid); err != nil { + return err + } + if err := atomicMetadata(f.pidFilePath(cfg.ID), []byte(strconv.Itoa(pid))); err != nil { + return err + } + } else { + return errors.New("microvm: launched child has no process identity") + } fcClient := newFirecrackerClient(sockPath) if err := waitForSocket(ctx, sockPath, firecrackerSocketTimeout); err != nil { @@ -3111,7 +3306,6 @@ func (f *FirecrackerEngine) Resume(ctx context.Context, id string) error { func (f *FirecrackerEngine) Stop(ctx context.Context, id string) error { f.mu.Lock() proc, tracked := f.procs[id] - delete(f.procs, id) f.mu.Unlock() var pid int @@ -3121,6 +3315,24 @@ func (f *FirecrackerEngine) Stop(ctx context.Context, id string) error { pid, _ = strconv.Atoi(strings.TrimSpace(string(data))) } + gone, evidenceErr := f.launchEvidence(id, pid) + if evidenceErr != nil { + if !tracked { + return evidenceErr + } + // A completed Wait proves this original child exited, even if its + // numeric PID now names something else. Otherwise retain uncertainty. + select { + case <-f.waitChild(id, proc): + gone = true + default: + return evidenceErr + } + } + if gone { + pid = 0 + } + // waited is the channel this process's own Wait reports on. A VMM this // engine started is a CHILD: signalling it is not enough, because an // un-Waited child stays in the process table as a zombie, and a runner @@ -3131,8 +3343,7 @@ func (f *FirecrackerEngine) Stop(ctx context.Context, id string) error { // already reparented it to init, which reaps it. var waited chan error if tracked && pid > 0 { - waited = make(chan error, 1) - go func() { waited <- proc.Wait() }() + waited = f.waitChild(id, proc) } var stopErr error @@ -3146,11 +3357,13 @@ func (f *FirecrackerEngine) Stop(ctx context.Context, id string) error { // kernel rather than assuming it, so a pid recovered across a // runnerd restart — whose group this process knows nothing about — // is only ever signalled on its own. - if err := killProcessTree(f.signals, pid, syscall.SIGTERM); err != nil { + if err := f.signalVM(id, pid, syscall.SIGTERM); err != nil { stopErr = fmt.Errorf("sigterm pid %d: %w", pid, err) - } else if !awaitExit(ctx, waited, pid, firecrackerTermTimeout) { - _ = killProcessTree(f.signals, pid, syscall.SIGKILL) - if !awaitExit(ctx, waited, pid, firecrackerKillTimeout) && stopErr == nil { + } else if !f.awaitVMExit(ctx, id, waited, pid, firecrackerTermTimeout) { + if err := f.signalVM(id, pid, syscall.SIGKILL); err != nil { + return err + } + if !f.awaitVMExit(ctx, id, waited, pid, firecrackerKillTimeout) && stopErr == nil { stopErr = notExitedErr(ctx, pid) } } @@ -3165,11 +3378,16 @@ func (f *FirecrackerEngine) Stop(ctx context.Context, id string) error { // live process that failed the identity check will never exit on its // own, so Stop — and with it Destroy, Suspend and every caller // holding a session's teardown — waited forever. - if !awaitExit(ctx, waited, pid, firecrackerKillTimeout) { + if !f.awaitVMExit(ctx, id, waited, pid, firecrackerKillTimeout) { stopErr = fmt.Errorf("firecracker %s: pid %d is still running but does not identify as this VM's VMM, so it was left alone", id, pid) } } + if stopErr != nil { + return stopErr + } + f.forgetExited(id) + // The jail goes with the VM, and only AFTER it: removing a chroot out // from under a live Firecracker is how a VMM ends up writing into // unlinked files. What is removed is the directory and the hard links in @@ -3197,8 +3415,8 @@ func (f *FirecrackerEngine) Stop(ctx context.Context, id string) error { // exists. It lives outside the jail, so removeJail does not take it, and // a stale one is the input to State's and PID's identity check on the // next boot. - if err := os.Remove(f.pidFilePath(id)); err != nil && !errors.Is(err, os.ErrNotExist) && stopErr == nil { - stopErr = fmt.Errorf("remove the pid file for %s: %w", id, err) + if err := f.removeLaunchEvidence(id); err != nil && stopErr == nil { + stopErr = err } return stopErr } @@ -3238,7 +3456,7 @@ func awaitExit(ctx context.Context, waited chan error, pid int, timeout time.Dur poll := time.NewTicker(50 * time.Millisecond) defer poll.Stop() for { - if err := syscall.Kill(pid, 0); err != nil { + if err := syscall.Kill(pid, 0); errors.Is(err, syscall.ESRCH) { return true } select { @@ -3276,6 +3494,14 @@ func (f *FirecrackerEngine) State(ctx context.Context, id string) (VMMState, err pid, _ = strconv.Atoi(strings.TrimSpace(string(data))) } + gone, evidenceErr := f.launchEvidence(id, pid) + if evidenceErr != nil { + return "", evidenceErr + } + if gone { + return VMMStateGone, nil + } + // A terminated Firecracker leaves no process to ask, so this engine never // answers VMMStateStopped: a cold-parked session reads as Gone here, and // it is the DRIVER's own record that tells parked from vanished. See @@ -3362,7 +3588,7 @@ func hasKVM() bool { // the jailer it is the instance id, and it is on the command line twice over: // the jailer passes its own `--id` through to Firecracker, and the binary it // execs lives at /firecracker//root/firecracker, so the -// argv carries the id whichever way it is read. It used to be the API socket +// argv carries an exact --id value; path substrings are never authority. It used to be the API socket // path, which no longer distinguishes anything — every jailed VMM serves // /run/firecracker.socket, because every one of them has a root of its own. func isFirecrackerPID(pid int, marker string) bool { @@ -3373,21 +3599,59 @@ func isFirecrackerPID(pid int, marker string) bool { return false } - cmdlinePath := fmt.Sprintf("/proc/%d/cmdline", pid) - if data, err := os.ReadFile(cmdlinePath); err == nil { - // The cmdline is NUL-separated, so a substring search over it can - // match across argument boundaries. That is harmless here: both - // needles are whole arguments or parts of one path, and the check is - // "is this plausibly the VMM we started" rather than a parser. - cmdline := string(data) - return strings.Contains(cmdline, jailExecName) && strings.Contains(cmdline, marker) - } + args, err := processArguments(pid) + return err == nil && guestProcessArguments(args, marker) +} - cmd := exec.Command("ps", "-p", strconv.Itoa(pid), "-o", "command=") - if out, err := cmd.Output(); err == nil { - s := string(out) - return strings.Contains(s, jailExecName) && (marker == "" || strings.Contains(s, marker)) +// ResumeStatus observes a claim without launching. If its command never reached +// a cold VM, persisting the requested generation fences that delayed command: +// ResumePlacement refuses an already-recorded cold generation. +func (m *Microvm) ResumeStatus(ctx context.Context, id string, generation uint64) (string, uint64, error) { + if _, err := m.Inspect(ctx, id); err != nil { + return "", 0, err + } + m.mu.Lock() + defer m.mu.Unlock() + inst, ok := m.instances[id] + if !ok || generation == 0 || generation < inst.PlacementGeneration { + return "", 0, errors.New("microvm: stale resume status") + } + if inst.resuming || inst.RecoveryBlocked { + return "resuming", inst.PlacementGeneration, nil + } + if generation > inst.PlacementGeneration { + if !inst.Cold { + return "", 0, errors.New("microvm: resume placement is not recorded") + } + before := inst.PlacementGeneration + inst.PlacementGeneration = generation + if err := m.saveRecord(persistable(inst)); err != nil { + inst.PlacementGeneration = before + return "", 0, err + } + } + if inst.Cold { + return "suspended_cold", generation, nil + } + if inst.State == StateRunning { + return "running", generation, nil } + return "resuming", generation, nil +} - return false +// CapacitySnapshot captures the aggregate and its known placements under the +// same driver lock. Pending creates without a published record remain unnamed. +func (m *Microvm) CapacitySnapshot(ctx context.Context) (int, int, map[string]uint64, error) { + if err := ctx.Err(); err != nil { + return 0, 0, nil, err + } + m.mu.Lock() + defer m.mu.Unlock() + placements := map[string]uint64{} + for _, inst := range m.instances { + if inst.PlacementGeneration > 0 && (inst.resuming || inst.State == StateRunning || (inst.State == StateSuspended && !inst.Cold)) { + placements[inst.SessionID] = inst.PlacementGeneration + } + } + return m.usedLocked(), m.opts.TotalSlots, placements, nil } diff --git a/internal/driver/microvm_jailer_test.go b/internal/driver/microvm_jailer_test.go index 41f623e..3f91b2e 100644 --- a/internal/driver/microvm_jailer_test.go +++ b/internal/driver/microvm_jailer_test.go @@ -14,6 +14,7 @@ import ( "sync" "syscall" "testing" + "time" ) // fakeStarter is the exec seam: it records the argv a host would be asked to @@ -190,6 +191,7 @@ func jailTestEngineIn(t *testing.T, dir string, fs *jailFakeFS, starter processS // machine has. Every test here is about what the engine WOULD run and // what it puts on disk first; production never replaces this. fc.kvm = func() bool { return true } + fc.startTime = func(pid int) (uint64, error) { return uint64(pid), nil } return fc } @@ -919,3 +921,174 @@ func writeJailImage(t *testing.T, dir, name string, mode os.FileMode) string { } return path } + +// A crash can happen after Start but before the PID is durably published. +func TestPendingLaunchWithoutPIDRetainsJail(t *testing.T) { + fc, dir, _ := jailTestEngine(t, &fakeStarter{}) + id := "mvm-9" + marker := filepath.Join(dir, "instances", id, "launch.pending") + if err := os.MkdirAll(filepath.Dir(marker), 0700); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(marker, []byte("unknown"), 0600); err != nil { + t.Fatal(err) + } + jail := jailInstanceDir(dir, id) + if err := os.MkdirAll(jail, 0700); err != nil { + t.Fatal(err) + } + for _, pidText := range []string{"", "not-a-pid", strconv.Itoa(os.Getpid())} { + if pidText != "" { + if err := os.WriteFile(fc.pidFilePath(id), []byte(pidText), 0600); err != nil { + t.Fatal(err) + } + } + if _, err := fc.State(context.Background(), id); err == nil { + t.Error("missing PID was treated as proof of exit") + } + if err := fc.Stop(context.Background(), id); err == nil { + t.Error("uncertain launch teardown succeeded") + } + if _, err := os.Stat(jail); err != nil { + t.Error("uncertain launch lost its jail") + } + } +} + +type observingStarter struct { + check func() + fakeStarter +} + +func (s *observingStarter) Start(name string, args []string) (vmmProcess, error) { + s.check() + return s.fakeStarter.Start(name, args) +} +func TestLaunchPublishesIntentBeforeStartingChild(t *testing.T) { + starter := &observingStarter{} + fc, dir, _ := jailTestEngine(t, starter) + starter.check = func() { + data, err := os.ReadFile(fc.launchMarkerPath("mvm-7")) + if err != nil || len(data) == 0 { + t.Fatal("child started without durable launch evidence") + } + } + cfg := VMMConfig{ID: "mvm-7", SlotIndex: 7, KernelPath: writeJailImage(t, dir, "kernel", 0644), RootfsPath: writeJailImage(t, dir, "rootfs", 0644)} + ctx, cancel := context.WithCancel(context.Background()) + cancel() + if err := fc.Launch(ctx, cfg); err == nil { + t.Fatal("expected unreachable API") + } + if _, err := os.Stat(fc.launchMarkerPath(cfg.ID)); !errors.Is(err, os.ErrNotExist) { + t.Fatal("confirmed child exit did not clear launch intent") + } +} + +func TestLaunchEvidenceRejectsAnotherInstancePID(t *testing.T) { + fc, _, _ := jailTestEngine(t, &fakeStarter{}) + proc, pid := fakeFirecracker(t, "mvm-10") + defer func() { _ = proc.Kill(); _ = proc.Wait() }() + if err := atomicMetadata(fc.launchMarkerPath("mvm-1"), []byte(hostLaunchBoot())); err != nil { + t.Fatal(err) + } + fc.startTime = processStartTime + if err := fc.saveProcessIdentity("mvm-1", pid); err != nil { + t.Fatal(err) + } + if gone, err := fc.launchEvidence("mvm-1", pid); err == nil { + t.Fatalf("another instance PID accepted as current launch evidence: gone=%v", gone) + } + if err := atomicMetadata(fc.pidFilePath("mvm-1"), []byte(strconv.Itoa(pid))); err != nil { + t.Fatal(err) + } + signals := &fakeSignaller{} + fc.signals = signals + if err := fc.Stop(context.Background(), "mvm-1"); err == nil { + t.Fatal("another VM accepted for teardown") + } + if len(signals.sent) != 0 { + t.Fatal("mismatched process received a signal") + } + if !isFirecrackerPID(pid, "mvm-10") { + t.Fatal("exact-ID positive control failed") + } + +} + +func TestLaunchEvidenceRequiresOriginalProcessLifetime(t *testing.T) { + fc, _, _ := jailTestEngine(t, &fakeStarter{}) + proc, pid := fakeFirecracker(t, "mvm-11") + defer func() { _ = proc.Kill(); _ = proc.Wait() }() + if err := atomicMetadata(fc.launchMarkerPath("mvm-11"), []byte(hostLaunchBoot())); err != nil { + t.Fatal(err) + } + if err := atomicMetadata(fc.pidFilePath("mvm-11"), []byte(strconv.Itoa(pid))); err != nil { + t.Fatal(err) + } + // An exact command/instance ID alone does not prove this is the original + // process. Until its birth identity has been published, ownership is unknown. + if _, err := fc.launchEvidence("mvm-11", pid); err == nil { + t.Fatal("exact argv accepted without process lifetime evidence") + } + fc.startTime = processStartTime + if err := fc.saveProcessIdentity("mvm-11", pid); err != nil { + t.Fatal(err) + } + if gone, err := fc.launchEvidence("mvm-11", pid); err != nil || gone { + t.Fatal("original process positive control failed") + } + actual := fc.startTime + fc.startTime = func(pid int) (uint64, error) { v, err := actual(pid); return v + 1, err } + if _, err := fc.launchEvidence("mvm-11", pid); err == nil { + t.Fatal("reused process lifetime accepted") + } + signals := &fakeSignaller{} + fc.signals = signals + if err := fc.Stop(context.Background(), "mvm-11"); err == nil { + t.Fatal("recycled lifetime accepted for teardown") + } + if len(signals.sent) != 0 { + t.Fatal("recycled process received a signal") + } + +} + +type lifetimeChangeSignal struct { + onTerm func() + kills int +} + +func (s *lifetimeChangeSignal) Getpgid(pid int) (int, error) { return pid + 1, nil } +func (s *lifetimeChangeSignal) Kill(_ int, sig syscall.Signal) error { + if sig == syscall.SIGTERM { + s.onTerm() + } + if sig == syscall.SIGKILL { + s.kills++ + } + return nil +} +func TestStopRechecksLifetimeBeforeEscalation(t *testing.T) { + fc, _, _ := jailTestEngine(t, &fakeStarter{}) + proc, pid := fakeFirecracker(t, "mvm-12") + defer func() { _ = proc.Kill(); _ = proc.Wait() }() + fc.startTime = processStartTime + if err := atomicMetadata(fc.launchMarkerPath("mvm-12"), []byte(hostLaunchBoot())); err != nil { + t.Fatal(err) + } + if err := fc.saveProcessIdentity("mvm-12", pid); err != nil { + t.Fatal(err) + } + if err := atomicMetadata(fc.pidFilePath("mvm-12"), []byte(strconv.Itoa(pid))); err != nil { + t.Fatal(err) + } + originalStart := fc.startTime + sig := &lifetimeChangeSignal{onTerm: func() { fc.startTime = func(pid int) (uint64, error) { v, e := originalStart(pid); return v + 1, e } }} + fc.signals = sig + c, cancel := context.WithTimeout(context.Background(), 20*time.Millisecond) + defer cancel() + _ = fc.Stop(c, "mvm-12") + if sig.kills != 0 { + t.Fatal("SIGKILL authorized after original birth identity changed during wait") + } +} diff --git a/internal/driver/microvm_launch_evidence.go b/internal/driver/microvm_launch_evidence.go new file mode 100644 index 0000000..4bf6dfa --- /dev/null +++ b/internal/driver/microvm_launch_evidence.go @@ -0,0 +1,212 @@ +package driver + +import ( + "context" + "encoding/json" + "errors" + "os" + "path/filepath" + "strings" + "syscall" + "time" +) + +// atomicMetadata publishes a complete, durable version in the same directory. +func atomicMetadata(path string, data []byte) error { + dir := filepath.Dir(path) + if err := os.MkdirAll(dir, microvmDirMode); err != nil { + return err + } + tmp, err := os.CreateTemp(dir, ".metadata-*") + if err != nil { + return err + } + defer os.Remove(tmp.Name()) + if err = tmp.Chmod(microvmFileMode); err == nil { + _, err = tmp.Write(data) + } + if err == nil { + err = tmp.Sync() + } + closeErr := tmp.Close() + if err != nil { + return err + } + if closeErr != nil { + return closeErr + } + if err := os.Rename(tmp.Name(), path); err != nil { + return err + } + if err := syncDir(dir); err != nil { + return err + } + return syncDir(filepath.Dir(dir)) +} + +func (f *FirecrackerEngine) launchMarkerPath(id string) string { + return filepath.Join(f.stateDir, "instances", id, "launch.pending") +} +func hostLaunchBoot() string { + data, err := os.ReadFile("/proc/sys/kernel/random/boot_id") + if err != nil || len(strings.TrimSpace(string(data))) == 0 { + return "unknown" + } + return strings.TrimSpace(string(data)) +} + +// launchEvidence distinguishes missing bookkeeping from process exit. A marker +// is durable before Start, so a missing PID on the same host boot is ambiguous. +// Preserve the jail/resources for host inspection; a proven host reboot ends +// that process lifetime and permits cleanup. No /proc scan or empty cgroup can +// prove absence while a jailer might still be entering its namespace/cgroup. +func (f *FirecrackerEngine) launchEvidence(id string, pid int) (gone bool, err error) { + data, err := os.ReadFile(f.launchMarkerPath(id)) + if errors.Is(err, os.ErrNotExist) { + return false, nil + } + if err != nil { + return false, errors.New("microvm: cannot read launch evidence") + } + boot := strings.TrimSpace(string(data)) + current := hostLaunchBoot() + if boot != "" && boot != "unknown" && current != "unknown" && boot != current { + return true, nil + } + if pid <= 0 { + return false, errors.New("microvm: launch lacks process evidence; retain resources for host inspection") + } + var birth launchProcessIdentity + data, err = os.ReadFile(f.processIdentityPath(id)) + if err != nil || json.Unmarshal(data, &birth) != nil || birth.PID != pid || birth.StartTime == 0 { + return false, errors.New("microvm: launch process lifetime is unknown; retain resources") + } + if errors.Is(syscall.Kill(pid, 0), syscall.ESRCH) { + return true, nil + } + start, err := f.startTime(pid) + if err != nil || start != birth.StartTime { + return false, errors.New("microvm: process lifetime no longer matches launch") + } + + if !isFirecrackerPID(pid, id) { + if errors.Is(syscall.Kill(pid, 0), syscall.ESRCH) { + return true, nil + } + return false, errors.New("microvm: launch process identity is uncertain; retain resources") + } + return false, nil +} + +// Wait is started once per child and remains usable by later teardown attempts. +func (f *FirecrackerEngine) waitChild(id string, proc vmmProcess) chan error { + f.mu.Lock() + defer f.mu.Unlock() + if f.waits == nil { + f.waits = map[string]chan error{} + } + if c := f.waits[id]; c != nil { + return c + } + c := make(chan error, 1) + f.waits[id] = c + go func() { c <- proc.Wait(); close(c) }() + return c +} +func (f *FirecrackerEngine) forgetExited(id string) { + f.mu.Lock() + delete(f.procs, id) + delete(f.waits, id) + f.mu.Unlock() +} +func (f *FirecrackerEngine) removeLaunchEvidence(id string) error { + for _, path := range []string{f.pidFilePath(id), f.processIdentityPath(id), f.launchMarkerPath(id)} { + if err := os.Remove(path); err != nil && !errors.Is(err, os.ErrNotExist) { + return err + } + } + dir := filepath.Dir(f.pidFilePath(id)) + if _, err := os.Stat(dir); errors.Is(err, os.ErrNotExist) { + return nil + } + return syncDir(dir) +} + +// finishFailedLaunch may only remove the jail after its own child has exited. +func (f *FirecrackerEngine) finishFailedLaunch(id string, proc vmmProcess) { + ctx, cancel := context.WithTimeout(context.Background(), firecrackerKillTimeout) + defer cancel() + _ = proc.Kill() + if !awaitExit(ctx, f.waitChild(id, proc), proc.Pid(), firecrackerKillTimeout) { + return + } + f.forgetExited(id) + _ = f.removeJail(id) + _ = f.removeLaunchEvidence(id) +} + +// The birth timestamp distinguishes a reused PID even when the next process +// carries the same exact instance ID. The launch marker supplies host boot scope. +type launchProcessIdentity struct { + PID int `json:"pid"` + StartTime uint64 `json:"start_time"` +} + +func (f *FirecrackerEngine) processIdentityPath(id string) string { + return filepath.Join(f.stateDir, "instances", id, "process.json") +} +func (f *FirecrackerEngine) saveProcessIdentity(id string, pid int) error { + start, err := f.startTime(pid) + if err != nil || start == 0 { + return errors.New("microvm: cannot establish child birth identity") + } + data, err := json.Marshal(launchProcessIdentity{PID: pid, StartTime: start}) + if err != nil { + return err + } + return atomicMetadata(f.processIdentityPath(id), data) +} + +// Every signal requires current authority, including escalation after waiting. +// A changed or unreadable lifetime retains the jail for later reconciliation. +func (f *FirecrackerEngine) signalVM(id string, pid int, signal syscall.Signal) error { + gone, err := f.launchEvidence(id, pid) + if err != nil { + return err + } + if gone { + return nil + } + if !isFirecrackerPID(pid, id) { + return errors.New("microvm: signal target identity is uncertain") + } + return killProcessTree(f.signals, pid, signal) +} + +// Recovered VMs have no child Wait handle. Observe their durable lifetime on +// every poll so a replacement PID cannot extend the original VM's wait. +func (f *FirecrackerEngine) awaitVMExit(ctx context.Context, id string, waited chan error, pid int, timeout time.Duration) bool { + if waited != nil { + return awaitExit(ctx, waited, pid, timeout) + } + timer := time.NewTimer(timeout) + defer timer.Stop() + poll := time.NewTicker(50 * time.Millisecond) + defer poll.Stop() + for { + gone, err := f.launchEvidence(id, pid) + if err != nil { + return false + } + if gone || errors.Is(syscall.Kill(pid, 0), syscall.ESRCH) { + return true + } + select { + case <-poll.C: + case <-timer.C: + return false + case <-ctx.Done(): + return false + } + } +} diff --git a/internal/driver/microvm_netslot_test.go b/internal/driver/microvm_netslot_test.go index fe71fb5..72fef4a 100644 --- a/internal/driver/microvm_netslot_test.go +++ b/internal/driver/microvm_netslot_test.go @@ -461,6 +461,7 @@ func TestMicrovmARecoveredSessionsUIDIsNotHandedToTheNextCreate(t *testing.T) { // The restart: a second driver over the same state directory and the same // host, as a restarted runnerd would be. + stopTestMicrovmRunner(t, first) second, _, _ := testMicrovmNet(t, MicrovmOpts{TotalSlots: 4, StateDir: stateDir, Net: net}) fresh, err := second.Create(ctx, Spec{SessionID: "beta"}) if err != nil { @@ -558,6 +559,7 @@ func TestMicrovmRecoveryReassociatesLiveSlotsAndReclaimsTheRest(t *testing.T) { // A second driver over the same state directory, as a restarted runnerd // would be. The simulated engine is rebuilt from the same directory so // the recovered record still reads as running. + stopTestMicrovmRunner(t, first) second, _, _ := testMicrovmNet(t, MicrovmOpts{TotalSlots: 4, StateDir: stateDir, Net: net}) if got := net.Namespaces(); len(got) != 1 || got[0] != liveCfg.Netns { @@ -596,6 +598,7 @@ func TestMicrovmRecoveryDropsAColdSessionsSlotClaim(t *testing.T) { t.Fatalf("cold Suspend: %v", err) } + stopTestMicrovmRunner(t, first) second, _, _ := testMicrovmNet(t, MicrovmOpts{TotalSlots: 1, StateDir: stateDir, Net: net}) if _, err := second.Create(ctx, Spec{SessionID: "beta"}); err != nil { t.Fatalf("Create after recovery: %v; the parked session's slot claim was never dropped", err) diff --git a/internal/driver/microvm_restore_test.go b/internal/driver/microvm_restore_test.go index 93371b7..8e5438a 100644 --- a/internal/driver/microvm_restore_test.go +++ b/internal/driver/microvm_restore_test.go @@ -634,6 +634,7 @@ func TestAColdResumeAfterARestartSaysWhichThingItIsWaitingFor(t *testing.T) { t.Fatal(err) } + stopTestMicrovmRunner(t, m) again, format := restartedOver(t, stateDir, ckpt) _, err = again.Resume(ctx, h.ID) if err == nil { @@ -683,6 +684,7 @@ func TestAColdResumeAfterARestartWithNoCheckpointSaysThatToo(t *testing.T) { if err != nil { t.Fatal(err) } + stopTestMicrovmRunner(t, m) again, format := restartedOver(t, stateDir, &CheckpointOpts{ Store: checkpoint.NewMemoryStore(), Keys: keys, KeyRef: keys.Ref(), }) diff --git a/internal/driver/microvm_rootfs_test.go b/internal/driver/microvm_rootfs_test.go index f026f11..e6a7997 100644 --- a/internal/driver/microvm_rootfs_test.go +++ b/internal/driver/microvm_rootfs_test.go @@ -273,6 +273,7 @@ func TestMicrovmReclaimsOrphanRootfsOnStart(t *testing.T) { } // A new runner over the same state directory: the restart. + stopTestMicrovmRunner(t, m) restarted, _, _ := testMicrovmCloning(t, MicrovmOpts{ TotalSlots: 4, StateDir: stateDir, BaseRootfs: m.opts.BaseRootfs, }) diff --git a/internal/driver/microvm_teardown_test.go b/internal/driver/microvm_teardown_test.go index b605da0..dd697ed 100644 --- a/internal/driver/microvm_teardown_test.go +++ b/internal/driver/microvm_teardown_test.go @@ -74,6 +74,7 @@ func TestMicrovmDestroyRemovesTheWorkspaceOfARecordOnlyOnDisk(t *testing.T) { } // The restart proves it: nothing comes back. + stopTestMicrovmRunner(t, m) restarted, _, _ := testMicrovmCloning(t, MicrovmOpts{ TotalSlots: 2, StateDir: m.opts.StateDir, BaseRootfs: m.opts.BaseRootfs, }) diff --git a/internal/driver/microvm_test.go b/internal/driver/microvm_test.go index 59b59a0..8df79c8 100644 --- a/internal/driver/microvm_test.go +++ b/internal/driver/microvm_test.go @@ -10,6 +10,7 @@ import ( "net" "net/http" "os" + "os/exec" "path/filepath" "reflect" "slices" @@ -63,6 +64,7 @@ func testMicrovmNet(t *testing.T, opts MicrovmOpts) (*Microvm, *SimulatedEngine, if err != nil { t.Fatalf("NewMicrovm: %v", err) } + t.Cleanup(func() { m.stateLock.Close() }) return m, sim, net } @@ -301,6 +303,7 @@ func TestMicrovmColdSuspendedSessionIsNotGone(t *testing.T) { // And after a restart: a new driver over the same state directory, with a // new engine that has no memory of the VM at all. + stopTestMicrovmRunner(t, m) m2, _ := testMicrovm(t, MicrovmOpts{ TotalSlots: 4, StateDir: stateDir, @@ -363,6 +366,7 @@ func TestMicrovmRestartRecovery(t *testing.T) { // A complete runner restart: a new engine AND a new driver over the same // state directory. + stopTestMicrovmRunner(t, m1) m2, _ := testMicrovm(t, MicrovmOpts{TotalSlots: 4, StateDir: stateDir}) listed, err := m2.List(ctx) @@ -649,6 +653,7 @@ func TestMicrovmColdResumeAfterRestartRefuses(t *testing.T) { t.Fatal(err) } + stopTestMicrovmRunner(t, m1) m2, _ := testMicrovm(t, MicrovmOpts{TotalSlots: 4, StateDir: stateDir}) m2.SetHost(&stubMicrovmHost{}) if _, err := m2.Resume(ctx, h.ID); err == nil { @@ -1435,7 +1440,7 @@ func TestFirecrackerStateReadsInstanceInfo(t *testing.T) { // fakeFirecracker starts a child that passes isFirecrackerPID: a script named // `firecracker`, invoked with the same `--id` argument the jailer passes // through to the real one, so both the /proc cmdline and the `ps -o command=` -// fallback see the binary name and the instance id the check looks for. It +// native argv lookup see the binary name and exact instance ID. It // exits on SIGTERM within one tick of its loop. // // A stand-in that does NOT pass the check (plain `sleep`, say) exercises a @@ -1449,9 +1454,9 @@ func fakeFirecracker(t *testing.T, id string) (vmmProcess, int) { t.Helper() dir := t.TempDir() bin := filepath.Join(dir, "firecracker") - script := "#!/bin/sh\ntrap 'exit 0' TERM\nwhile :; do sleep 0.02; done\n" - if err := os.WriteFile(bin, []byte(script), 0o700); err != nil { - t.Fatal(err) + build := exec.Command("go", "build", "-o", bin, "./testdata/fakefirecracker") + if output, err := build.CombinedOutput(); err != nil { + t.Fatalf("build native process fixture: %v %s", err, output) } proc, err := execStarter{}.Start(bin, []string{"--id", id, "--api-sock", jailAPISocketPath}) if err != nil { @@ -1551,6 +1556,10 @@ func TestFirecrackerStopDoesNotHangOnAnUnidentifiableChild(t *testing.T) { fc.procs["mvm-stuck"] = proc fc.mu.Unlock() + jail := jailInstanceDir(dir, "mvm-stuck") + if err := os.MkdirAll(jail, 0700); err != nil { + t.Fatal(err) + } // A caller whose context is already cut short must be answered at once. ctx, cancel := context.WithTimeout(context.Background(), 200*time.Millisecond) defer cancel() @@ -1575,4 +1584,46 @@ func TestFirecrackerStopDoesNotHangOnAnUnidentifiableChild(t *testing.T) { if err := syscall.Kill(pid, 0); err != nil { t.Errorf("Stop signalled a process it could not identify as this VM's VMM: %v", err) } + if _, err := os.Stat(jail); err != nil { + t.Fatal("failed stop removed a live child's jail") + } + fc.mu.Lock() + retained := fc.procs["mvm-stuck"] == proc + fc.mu.Unlock() + if !retained { + t.Fatal("failed stop lost process ownership") + } + +} + +// Simulate process exit without stopping surviving VMs or deleting their UDS. +// A second fixture is a restart, never a concurrent writer of the same state. +func stopTestMicrovmRunner(t *testing.T, m *Microvm) { + t.Helper() + m.mu.Lock() + var channels []*guestChannel + for _, rec := range m.instances { + if rec.channel != nil { + channels = append(channels, rec.channel) + } + } + m.mu.Unlock() + for _, g := range channels { + g.mu.Lock() + g.closed = true + conns := make([]net.Conn, 0, len(g.conns)) + for c := range g.conns { + conns = append(conns, c) + } + g.mu.Unlock() + if g.listener != nil { + g.listener.Close() + } + for _, c := range conns { + g.drop(c) + } + } + if err := m.stateLock.Close(); err != nil { + t.Fatal(err) + } } diff --git a/internal/driver/microvm_vsock.go b/internal/driver/microvm_vsock.go index 143cdfa..87d954b 100644 --- a/internal/driver/microvm_vsock.go +++ b/internal/driver/microvm_vsock.go @@ -165,7 +165,13 @@ var ( // announced microvm.v1, and this driver refuses a create that carries them // anyway. The claim and the refusal are two halves of one promise, so they // are read from one place. -func (m *Microvm) Capabilities() []string { return []string{runner.CapabilityMicrovmV1} } +func (m *Microvm) Capabilities() []string { + caps := []string{runner.CapabilityMicrovmV1} + if m.opts.GuestReconnect { + caps = append(caps, runner.CapabilityGuestReconnectV1) + } + return caps +} // SetHost installs the runner above this driver. It is a setter rather than a // field on MicrovmOpts because the runner is composed OVER the driver @@ -247,20 +253,12 @@ type guestChannel struct { // failure would put a line in an operator's log for every session that // ends normally. closed bool - // served records that this boot generation's ONE guest connection has - // been taken. /dev/vsock is world-accessible inside an ordinary guest, - // so every process in the sandbox can dial (2, 1024) — and what the - // first frame carries is the session's whole configuration and a LIVE - // bootstrap token. Serving every connection would hand that to whoever - // asked, as many times as they asked, and let the last one become the - // session's hub. - // - // So the model is the design note's: one guest-initiated connection per - // boot (§4). The first is sessiond; every later one is refused and closed - // having received nothing at all. A sessiond that crashed and came back - // is a NEW boot generation — a new VM, a new socket, a new mint (open - // question 2) — and not something to re-serve this token to. - served bool + // served permanently consumes the initial boot delivery. Legacy guests + // get no second connection. Negotiated guests may subsequently prove their + // enrolled key; that path never reads boot or replays its bootstrap token. + served bool + reconnect bool + pending bool // conns are the connections this channel is still responsible for. It // holds at most the one served guest, and it exists so close() can end // it: a Destroy that closed only the listener would leave a wedged @@ -270,22 +268,41 @@ type guestChannel struct { conns map[net.Conn]struct{} } -// claim reserves this boot generation's one guest connection for c and takes -// responsibility for closing it. It reports false for a second connection and -// for one that arrived after teardown — in both cases the caller closes c -// without writing a byte to it. -func (g *guestChannel) claim(c net.Conn) bool { +// claim reserves one bounded admission and takes responsibility for c. +// Legacy guests retain their one-connection limit; refused peers receive no bytes. +func (g *guestChannel) claim(c net.Conn) bool { _, ok := g.admit(c); return ok } + +// admit reserves work inline, before any goroutine is launched. A reconnect +// never reuses the first boot path, even when a previous delivery failed. +func (g *guestChannel) admit(c net.Conn) (fresh, ok bool) { g.mu.Lock() defer g.mu.Unlock() - if g.closed || g.served { - return false + if g.closed || g.pending || (g.served && !g.reconnect) { + return false, false } + fresh = !g.served g.served = true + g.pending = true if g.conns == nil { g.conns = map[net.Conn]struct{}{} } g.conns[c] = struct{}{} - return true + return fresh, true +} +func (g *guestChannel) finishAdmission() { g.mu.Lock(); g.pending = false; g.mu.Unlock() } + +// Closing the relay drops its raw socket from listener ownership immediately. +// Otherwise a long-lived guest would retain every past reconnect in this map. +type guestTrackedConn struct { + net.Conn + channel *guestChannel +} + +func (c *guestTrackedConn) Close() error { + c.channel.mu.Lock() + delete(c.channel.conns, c.Conn) + c.channel.mu.Unlock() + return c.Conn.Close() } // drop closes a claimed connection this channel will not be handing on, and @@ -363,7 +380,7 @@ func (m *Microvm) openGuestChannel(sessionID, udsPath, listenPath string, cfg ru _ = os.Remove(listenPath) return nil, fmt.Errorf("restrict the guest control socket %s: %w", listenPath, err) } - g := &guestChannel{listener: ln, listenPath: listenPath, udsPath: udsPath, boot: cfg} + g := &guestChannel{listener: ln, listenPath: listenPath, udsPath: udsPath, boot: cfg, reconnect: cfg.GuestReconnect == runner.GuestReconnectProtocol} go m.acceptGuests(sessionID, g) return g, nil } @@ -381,11 +398,16 @@ func (m *Microvm) acceptGuests(sessionID string, g *guestChannel) { } return } - if !g.claim(c) { + fresh, ok := g.admit(c) + if !ok { _ = c.Close() continue } - go m.serveGuest(sessionID, g, c) + if fresh { + go m.serveGuest(sessionID, g, c) + } else { + go m.serveReconnectingGuest(sessionID, g, c) + } } } @@ -398,16 +420,15 @@ func (m *Microvm) acceptGuests(sessionID string, g *guestChannel) { // the stream" true: the guest reads its whole configuration off the same // connection it will then serve its terminal and its RPC over. // -// Every connection after the first is closed having received nothing — not -// the configuration, not the token, not a byte. See guestChannel.served: the -// socket is reachable by every process in the sandbox, and this is the one -// place that decides the first dial is the session's and no other is. +// Subsequent negotiated connections go through serveReconnectingGuest; +// legacy peers are refused. This function is never used for reconnect. // // A host that has not been installed closes the connection rather than // holding it. There is nothing to attach it to, and the claim stays spent: // this boot has had its one connection. func (m *Microvm) serveGuest(sessionID string, g *guestChannel, c net.Conn) { - conn := relay.NetConn(c) + defer g.finishAdmission() + conn := relay.NetConn(&guestTrackedConn{Conn: c, channel: g}) ctx, cancel := context.WithTimeout(context.Background(), bootConfigWriteTimeout) defer cancel() if err := writeBootConfig(ctx, conn, g.boot); err != nil { @@ -453,7 +474,7 @@ func writeBootConfig(ctx context.Context, conn relay.Conn, cfg runner.BootConfig return conn.Write(ctx, frame) } -// bootConfigFor composes what a guest is told about itself. +// GuestBootConfig composes what a guest is told about itself. // // Every field is one the create resolved, carried across unchanged. There are // two things deliberately NOT in it: any environment secret value (they are @@ -465,9 +486,10 @@ func writeBootConfig(ctx context.Context, conn relay.Conn, cfg runner.BootConfig // The proxy URL carries the session's identity as URL userinfo, composed here // with the same helper the Docker driver uses, so egressd reads one thing // whichever driver the session is on. -func bootConfigFor(spec Spec) runner.BootConfig { +func GuestBootConfig(spec Spec) runner.BootConfig { cfg := runner.BootConfig{ Protocol: runner.SessionBootstrapProtocolVersion, + GuestReconnect: spec.GuestReconnect, SessionID: spec.SessionID, Cmd: slices.Clone(spec.Cmd), EgressAllow: slices.Clone(spec.EgressAllow), @@ -536,3 +558,27 @@ func refuseUnwithheldEnv(spec Spec) error { "control plane to one that withholds them for a runner announcing %q", spec.SessionID, len(spec.Env), runner.CapabilityMicrovmV1) } + +// bootConfigFor keeps the driver-internal call sites on the same constructor. +func bootConfigFor(spec Spec) runner.BootConfig { return GuestBootConfig(spec) } + +// GuestReconnectHandoffHost owns the entire authenticated handoff. Unlike the +// proof-only port, success transfers the admitted stream into a guarded relay. +// The driver must validate local VM ownership and bound admission beforehand. +type GuestReconnectHandoffHost interface { + ReconnectGuest(context.Context, string, relay.Conn) error +} + +func (m *Microvm) serveReconnectingGuest(id string, g *guestChannel, c net.Conn) { + defer g.finishAdmission() + host, ok := m.currentHost().(GuestReconnectHandoffHost) + if !ok { + g.drop(c) + return + } + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Minute) + defer cancel() + if host.ReconnectGuest(ctx, id, relay.NetConn(&guestTrackedConn{Conn: c, channel: g})) != nil { + g.drop(c) + } +} diff --git a/internal/driver/process_arguments_darwin.go b/internal/driver/process_arguments_darwin.go new file mode 100644 index 0000000..47d266b --- /dev/null +++ b/internal/driver/process_arguments_darwin.go @@ -0,0 +1,58 @@ +package driver + +import ( + "bytes" + "encoding/binary" + "errors" + "golang.org/x/sys/unix" +) + +// Preserve argv boundaries. Flattened ps output is not process identity. +func processArguments(pid int) ([]byte, error) { + raw, err := unix.SysctlRaw("kern.procargs2", pid) + if err != nil { + return nil, err + } + return darwinProcessArguments(raw) +} +func darwinProcessArguments(raw []byte) ([]byte, error) { + invalid := errors.New("microvm: unavailable process arguments") + if len(raw) < 5 { + return nil, invalid + } + argc := int(binary.NativeEndian.Uint32(raw[:4])) + if argc < 1 || argc > 4096 { + return nil, invalid + } + raw = raw[4:] + end := bytes.IndexByte(raw, 0) + if end < 0 { + return nil, invalid + } + raw = raw[end+1:] + for len(raw) > 0 && raw[0] == 0 { + raw = raw[1:] + } + var out []byte + for i := 0; i < argc; i++ { + end = bytes.IndexByte(raw, 0) + if end < 0 { + return nil, invalid + } + out = append(out, raw[:end+1]...) + raw = raw[end+1:] + } + return out, nil +} + +func processStartTime(pid int) (uint64, error) { + info, err := unix.SysctlKinfoProc("kern.proc.pid", pid) + if err != nil { + return 0, err + } + t := info.Proc.P_starttime + if t.Sec <= 0 || t.Usec < 0 { + return 0, errors.New("microvm: unavailable process start time") + } + return uint64(t.Sec)*1000000 + uint64(t.Usec), nil +} diff --git a/internal/driver/process_arguments_linux.go b/internal/driver/process_arguments_linux.go new file mode 100644 index 0000000..cd74e21 --- /dev/null +++ b/internal/driver/process_arguments_linux.go @@ -0,0 +1,18 @@ +package driver + +import ( + "fmt" + "os" +) + +func processArguments(pid int) ([]byte, error) { + return os.ReadFile(fmt.Sprintf("/proc/%d/cmdline", pid)) +} + +func processStartTime(pid int) (uint64, error) { + data, err := os.ReadFile(fmt.Sprintf("/proc/%d/stat", pid)) + if err != nil { + return 0, err + } + return guestProcessStart(data, pid) +} diff --git a/internal/driver/process_arguments_other.go b/internal/driver/process_arguments_other.go new file mode 100644 index 0000000..e3b1d39 --- /dev/null +++ b/internal/driver/process_arguments_other.go @@ -0,0 +1,12 @@ +//go:build !linux && !darwin + +package driver + +import "errors" + +func processArguments(int) ([]byte, error) { + return nil, errors.New("microvm: process identity unsupported on this host") +} +func processStartTime(int) (uint64, error) { + return 0, errors.New("microvm: process identity unsupported on this host") +} diff --git a/internal/driver/reconnect_admission_test.go b/internal/driver/reconnect_admission_test.go new file mode 100644 index 0000000..56b0efe --- /dev/null +++ b/internal/driver/reconnect_admission_test.go @@ -0,0 +1,58 @@ +package driver + +import ( + "net" + "testing" +) + +func TestGuestReconnectAdmissionAllowsOnePendingPeer(t *testing.T) { + g := &guestChannel{reconnect: true} + first, peer := net.Pipe() + defer peer.Close() + defer first.Close() + defer g.close() + boot, ok := g.admit(first) + if !ok || !boot { + t.Fatal("initial boot refused") + } + next, other := net.Pipe() + defer next.Close() + defer other.Close() + if _, ok := g.admit(next); ok { + t.Fatal("peer admitted while initial delivery pending") + } + g.finishAdmission() + boot, ok = g.admit(next) + if !ok || boot { + t.Fatal("reconnect reused initial boot path") + } + excess, last := net.Pipe() + defer excess.Close() + defer last.Close() + if _, ok := g.admit(excess); ok { + t.Fatal("unbounded pending reconnects") + } + g.finishAdmission() + g.drop(next) + boot, ok = g.admit(excess) + if !ok || boot { + t.Fatal("failed attempt prevented fresh proof") + } +} + +func TestGuestReconnectTrackedCloseReleasesConnection(t *testing.T) { + g := &guestChannel{reconnect: true} + a, b := net.Pipe() + defer b.Close() + g.admit(a) + tracked := &guestTrackedConn{Conn: a, channel: g} + if err := tracked.Close(); err != nil { + t.Fatal(err) + } + g.mu.Lock() + count := len(g.conns) + g.mu.Unlock() + if count != 0 { + t.Fatal("closed connections accumulate in listener") + } +} diff --git a/internal/driver/reconnect_identity.go b/internal/driver/reconnect_identity.go new file mode 100644 index 0000000..5e437a8 --- /dev/null +++ b/internal/driver/reconnect_identity.go @@ -0,0 +1,47 @@ +package driver + +import ( + "bytes" + "errors" + "path/filepath" + "strconv" + "strings" +) + +// guestProcessStart extracts Linux stat field 22 without treating spaces or +// parentheses in comm as field separators. A missing identity fails closed. +func guestProcessStart(data []byte, pid int) (uint64, error) { + s := string(data) + open, end := strings.Index(s, " ("), strings.LastIndex(s, ") ") + if open < 1 || end <= open || s[:open] != strconv.Itoa(pid) || pid <= 0 { + return 0, errors.New("invalid guest process stat") + } + fields := strings.Fields(s[end+2:]) + if len(fields) < 20 { + return 0, errors.New("short guest process stat") + } + start, err := strconv.ParseUint(fields[19], 10, 64) + if err != nil || start == 0 { + return 0, errors.New("invalid guest process start time") + } + return start, nil +} + +func guestProcessArguments(data []byte, id string) bool { + args := bytes.Split(data, []byte{0}) + if len(args) < 3 || filepath.Base(string(args[0])) != jailExecName { + return false + } + found := false + for i := 1; i < len(args); i++ { + if string(args[i]) != "--id" { + continue + } + if found || i+1 >= len(args) || string(args[i+1]) != id { + return false + } + found = true + i++ + } + return found +} diff --git a/internal/driver/reconnect_identity_host.go b/internal/driver/reconnect_identity_host.go new file mode 100644 index 0000000..1543f5d --- /dev/null +++ b/internal/driver/reconnect_identity_host.go @@ -0,0 +1,162 @@ +package driver + +import ( + "errors" + "os" + "path/filepath" + "strconv" + "strings" + + "golang.org/x/sys/unix" +) + +// guestHostIdentity is non-secret host metadata, never a guest credential. +// StartTime alone is insufficient across host reboots or namespace replacement. +type guestHostIdentity struct { + BootID string `json:"boot_id"` + StartTime uint64 `json:"start_time"` + NamespaceDevice uint64 `json:"namespace_device"` + NamespaceInode uint64 `json:"namespace_inode"` +} + +func readGuestHostIdentity(procRoot, bootPath, namespace string, pid int, id string, uid, gid int) (guestHostIdentity, error) { + var zero guestHostIdentity + fail := errors.New("microvm: surviving guest identity unavailable") + if pid <= 0 || checkPathSegment("instance", id) != nil { + return zero, fail + } + dir := filepath.Join(procRoot, strconv.Itoa(pid)) + stat, err := os.ReadFile(filepath.Join(dir, "stat")) + if err != nil { + return zero, fail + } + start, err := guestProcessStart(stat, pid) + if err != nil { + return zero, fail + } + args, err := os.ReadFile(filepath.Join(dir, "cmdline")) + if err != nil || !guestProcessArguments(args, id) { + return zero, fail + } + status, err := os.ReadFile(filepath.Join(dir, "status")) + if err != nil { + return zero, fail + } + owners := map[string]int{"Uid:": uid, "Gid:": gid} + for _, line := range strings.Split(string(status), "\n") { + fields := strings.Fields(line) + if len(fields) == 0 { + continue + } + want, ok := owners[fields[0]] + if !ok { + continue + } + if want < 0 || len(fields) != 5 { + return zero, fail + } + for _, value := range fields[1:] { + if value != strconv.Itoa(want) { + return zero, fail + } + } + owners[fields[0]] = -1 + } + if owners["Uid:"] != -1 || owners["Gid:"] != -1 { + return zero, fail + } + boot, err := os.ReadFile(bootPath) + if err != nil { + return zero, fail + } + bootID := strings.TrimSpace(string(boot)) + if len(bootID) != 36 { + return zero, fail + } + for i, c := range bootID { + if i == 8 || i == 13 || i == 18 || i == 23 { + if c != '-' { + return zero, fail + } + continue + } + if !(c >= '0' && c <= '9' || c >= 'a' && c <= 'f') { + return zero, fail + } + } + var ns unix.Stat_t + if unix.Lstat(namespace, &ns) != nil || ns.Mode&unix.S_IFMT != unix.S_IFREG || ns.Ino == 0 { + return zero, fail + } + // Reject PID reuse during the multi-file inspection, too. + stat, err = os.ReadFile(filepath.Join(dir, "stat")) + if err != nil { + return zero, fail + } + after, err := guestProcessStart(stat, pid) + if err != nil || after != start { + return zero, fail + } + return guestHostIdentity{BootID: bootID, StartTime: start, NamespaceDevice: uint64(ns.Dev), NamespaceInode: uint64(ns.Ino)}, nil +} + +func guestRecoveryBoot(path string) (int, error) { + boot, err := strconv.Atoi(strings.TrimSuffix(strings.TrimPrefix(path, "/v"), ".sock")) + if err != nil || boot < 1 || path != vsockGuestPath(boot) { + return 0, errors.New("microvm: invalid recovered guest socket path") + } + return boot, nil +} + +// guestIdentityVerifier is deliberately optional: reconnect must fail closed +// on an engine that cannot prove local process ownership. +type guestIdentityVerifier interface { + guestIdentity(VMMConfig, int) (guestHostIdentity, error) +} + +func (f *FirecrackerEngine) guestIdentity(cfg VMMConfig, pid int) (guestHostIdentity, error) { + var zero guestHostIdentity + fail := errors.New("microvm: surviving guest ownership unavailable") + if checkPathSegment("instance", cfg.ID) != nil || checkPathSegment("namespace", cfg.Netns) != nil || pid <= 0 || f.PID(cfg.ID) != pid { + return zero, fail + } + uid, err := f.uids.forSlot(cfg.SlotIndex) + if err != nil { + return zero, fail + } + boot, err := guestRecoveryBoot(cfg.VsockUDSPath) + if err != nil { + return zero, fail + } + root := jailRootDir(f.stateDir, cfg.ID) + // The named jail and both VMM-owned sockets must still belong to this VM. + paths := []struct { + path string + kind uint32 + }{ + {root, unix.S_IFDIR}, {f.socketPath(cfg.ID), unix.S_IFSOCK}, {filepath.Join(root, vsockSocketName(boot)), unix.S_IFSOCK}, + } + for _, p := range paths { + var st unix.Stat_t + if unix.Lstat(p.path, &st) != nil || uint32(st.Mode)&unix.S_IFMT != p.kind || st.Uid != uint32(uid) { + return zero, fail + } + if p.kind == unix.S_IFDIR && (st.Gid != uint32(f.jail.RunnerGID) || os.FileMode(st.Mode&0777) != jailDirMode) { + return zero, fail + } + } + members, err := os.ReadFile(filepath.Join(cfg.CgroupPath, "cgroup.procs")) + if err != nil { + return zero, fail + } + found := false + for _, p := range strings.Fields(string(members)) { + if p == strconv.Itoa(pid) { + found = true + } + } + if !found { + return zero, fail + } + return readGuestHostIdentity("/proc", "/proc/sys/kernel/random/boot_id", filepath.Join(f.netnsDir, cfg.Netns), pid, cfg.ID, uid, uid) +} diff --git a/internal/driver/reconnect_identity_host_test.go b/internal/driver/reconnect_identity_host_test.go new file mode 100644 index 0000000..67ee8b7 --- /dev/null +++ b/internal/driver/reconnect_identity_host_test.go @@ -0,0 +1,88 @@ +package driver + +import ( + "context" + "os" + "path/filepath" + "strconv" + "testing" +) + +func TestGuestHostIdentityPinsProcessAndNamespace(t *testing.T) { + root := t.TempDir() + proc := filepath.Join(root, "proc") + if err := os.MkdirAll(filepath.Join(proc, "42"), 0700); err != nil { + t.Fatal(err) + } + write := func(path, value string) { + t.Helper() + if err := os.WriteFile(path, []byte(value), 0600); err != nil { + t.Fatal(err) + } + } + stat := "42 (firecracker) S 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 912345 0" + write(filepath.Join(proc, "42", "stat"), stat) + write(filepath.Join(proc, "42", "cmdline"), "/firecracker\x00--id\x00mvm-1\x00") + uid, gid := strconv.Itoa(os.Geteuid()), strconv.Itoa(os.Getegid()) + write(filepath.Join(proc, "42", "status"), "Uid:\t"+uid+"\t"+uid+"\t"+uid+"\t"+uid+"\nGid:\t"+gid+"\t"+gid+"\t"+gid+"\t"+gid+"\n") + ns := filepath.Join(root, "namespace") + write(ns, "synthetic namespace identity") + boot := filepath.Join(root, "boot-id") + write(boot, "11111111-2222-4333-8444-555555555555\n") + got, err := readGuestHostIdentity(proc, boot, ns, 42, "mvm-1", os.Geteuid(), os.Getegid()) + if err != nil { + t.Fatal(err) + } + if got.StartTime != 912345 || got.BootID != "11111111-2222-4333-8444-555555555555" || got.NamespaceInode == 0 { + t.Fatalf("bad identity: %+v", got) + } + // Replacement namespace at the same name must not compare as the original. + if err := os.Rename(ns, ns+".old"); err != nil { + t.Fatal(err) + } + write(ns, "replacement") + next, err := readGuestHostIdentity(proc, boot, ns, 42, "mvm-1", os.Geteuid(), os.Getegid()) + if err != nil || next == got { + t.Fatalf("namespace replacement not distinguished: %v", err) + } + write(filepath.Join(proc, "42", "cmdline"), "/firecracker\x00--id\x00mvm-10\x00") + if _, err := readGuestHostIdentity(proc, boot, ns, 42, "mvm-1", os.Geteuid(), os.Getegid()); err == nil { + t.Fatal("accepted another VM") + } + write(filepath.Join(proc, "42", "cmdline"), "/firecracker\x00--id\x00mvm-1\x00") + if _, err := readGuestHostIdentity(proc, boot, ns, 42, "mvm-1", os.Geteuid()+1, os.Getegid()); err == nil { + t.Fatal("accepted wrong process owner") + } +} + +func TestGuestRecoveryBootPathIsCanonical(t *testing.T) { + for _, tc := range []struct { + path string + boot int + }{ + {"/v1.sock", 1}, {"/v27.sock", 27}, {"/v0.sock", 0}, {"/v01.sock", 0}, {"/tmp/v1.sock", 0}, {"/v-1.sock", 0}, {"/v1.sock/../v2.sock", 0}, + } { + boot, err := guestRecoveryBoot(tc.path) + if tc.boot == 0 { + if err == nil { + t.Errorf("accepted %q", tc.path) + } + } else if err != nil || boot != tc.boot { + t.Errorf("%q: %d %v", tc.path, boot, err) + } + } +} + +func TestGuestReconnectRefusesUnverifiableEngine(t *testing.T) { + m, _ := testMicrovm(t, MicrovmOpts{GuestReconnect: true}) + defer m.stateLock.Close() + h, err := m.Create(context.Background(), Spec{SessionID: "session-test", GuestReconnect: 1}) + if err == nil { + m.Destroy(context.Background(), h.ID) + t.Fatal("enabled reconnect without a VM identity verifier") + } + used, _, _ := m.Capacity(context.Background()) + if used != 0 { + t.Fatal("refused reconnect reserved capacity") + } +} diff --git a/internal/driver/reconnect_identity_test.go b/internal/driver/reconnect_identity_test.go new file mode 100644 index 0000000..bd8bcc1 --- /dev/null +++ b/internal/driver/reconnect_identity_test.go @@ -0,0 +1,37 @@ +package driver + +import ( + "strings" + "testing" +) + +func TestGuestProcessIdentityParsing(t *testing.T) { + // Linux /proc/PID/stat: field 22 is starttime. The comm field can contain + // spaces and closing parentheses; it must not shift positional parsing. + stat := "42 (firecracker worker)) S 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 912345 0 0" + got, err := guestProcessStart([]byte(stat), 42) + if err != nil || got != 912345 { + t.Fatalf("start=%d err=%v", got, err) + } + for _, bad := range []string{"", strings.Replace(stat, "42 (", "43 (", 1), strings.Replace(stat, "912345", "0", 1), strings.Replace(stat, "912345", "invalid", 1), "42 (firecracker) S 1"} { + if _, err := guestProcessStart([]byte(bad), 42); err == nil { + t.Fatalf("accepted malformed process stat %q", bad) + } + } +} +func TestGuestProcessArgumentsRequireExactInstance(t *testing.T) { + for _, tc := range []struct { + args string + ok bool + }{ + {"/firecracker\x00--id\x00mvm-1\x00--api-sock\x00/run/firecracker.socket\x00", true}, + {"/firecracker\x00--id\x00mvm-10\x00", false}, + {"/not-firecracker\x00--id\x00mvm-1\x00", false}, + {"/firecracker\x00--log-path\x00mvm-1\x00", false}, + {"/firecracker\x00--id\x00mvm-1\x00--id\x00mvm-2\x00", false}, + } { + if got := guestProcessArguments([]byte(tc.args), "mvm-1"); got != tc.ok { + t.Errorf("args=%q got=%v", tc.args, got) + } + } +} diff --git a/internal/driver/reconnect_listener.go b/internal/driver/reconnect_listener.go new file mode 100644 index 0000000..37aa09c --- /dev/null +++ b/internal/driver/reconnect_listener.go @@ -0,0 +1,112 @@ +package driver + +import ( + "context" + "errors" + "net" + "os" + "syscall" + "time" + + "golang.org/x/sys/unix" +) + +// bindRecoveredGuestListener requires the state-directory ownership lock and +// verified VM/jail ownership from its caller. It never unlinks an active socket +// and touches only the host's control listener, never the surviving VMM's UDS. +func bindRecoveredGuestListener(path string, uid int) (*net.UnixListener, error) { + fail := errors.New("microvm: recovered guest listener is not exclusively owned") + before, err := os.Lstat(path) + if err != nil && !errors.Is(err, os.ErrNotExist) { + return nil, fail + } + if err == nil { + var st unix.Stat_t + if !before.Mode().IsRegular() && before.Mode()&os.ModeSocket == 0 { + return nil, fail + } + if unix.Lstat(path, &st) != nil || st.Mode&unix.S_IFMT != unix.S_IFSOCK || st.Uid != uint32(uid) { + return nil, fail + } + c, dialErr := net.DialTimeout("unix", path, 100*time.Millisecond) + if dialErr == nil { + c.Close() + return nil, fail + } + if !errors.Is(dialErr, syscall.ECONNREFUSED) { + return nil, fail + } + after, err := os.Lstat(path) + if err != nil || !os.SameFile(before, after) { + return nil, fail + } + if err := os.Remove(path); err != nil { + return nil, fail + } + } + return net.ListenUnix("unix", &net.UnixAddr{Name: path, Net: "unix"}) +} + +// RecoverGuest binds only after runner registration has reauthorized the exact +// session placement. It never reconstructs boot configuration from disk. +func (m *Microvm) RecoverGuest(ctx context.Context, id, sessionID string) error { + m.mu.Lock() + defer m.mu.Unlock() + fail := errors.New("microvm: guest recovery fenced") + rec, ok := m.instances[id] + if !ok || !m.opts.GuestReconnect || m.stateLock == nil || !rec.Reconnect || rec.SessionID != sessionID || rec.State != StateRunning || rec.Cold || rec.resuming || rec.slot == nil || rec.slot.Key != id || ctx.Err() != nil { + return fail + } + if rec.channel != nil { + return nil + } + if rec.Cfg.ID != id || rec.Cfg.SessionID != sessionID || rec.Cfg.CgroupPath != m.cgroupPathFor(id) { + return fail + } + expected := rec.Cfg + applySlot(&expected, rec.slot) + cfg := rec.Cfg + if cfg.SlotIndex != expected.SlotIndex || cfg.Netns != expected.Netns || cfg.TapDevice != expected.TapDevice || cfg.GuestIP != expected.GuestIP || cfg.GatewayIP != expected.GatewayIP || cfg.GuestNetmask != expected.GuestNetmask || cfg.GuestMAC != expected.GuestMAC { + return fail + } + f, ok := m.engine.(*FirecrackerEngine) + if !ok { + return fail + } + identity, err := f.guestIdentity(cfg, rec.PID) + if err != nil || identity != rec.Identity || identity.StartTime == 0 { + return fail + } + boot, err := guestRecoveryBoot(cfg.VsockUDSPath) + if err != nil { + return fail + } + uds, path, err := m.vsockPaths(id, boot) + if err != nil { + return fail + } + uid, err := f.uids.forSlot(cfg.SlotIndex) + if err != nil { + return fail + } + listener, err := bindRecoveredGuestListener(path, uid) + if err != nil { + return err + } + if err := f.chownJailPath(path, uid, f.jail.RunnerGID, jailFileMode); err != nil { + listener.Close() + return fail + } + // Verify again after binding and before publication. A failed recovery closes + // only its new listener, never Firecracker's device socket. + current, err := f.guestIdentity(cfg, rec.PID) + if err != nil || current != identity || ctx.Err() != nil { + listener.Close() + return fail + } + g := &guestChannel{listener: listener, listenPath: path, udsPath: uds, served: true, reconnect: true} + rec.channel = g + rec.boots = boot + go m.acceptGuests(sessionID, g) + return nil +} diff --git a/internal/driver/reconnect_listener_test.go b/internal/driver/reconnect_listener_test.go new file mode 100644 index 0000000..faaeba1 --- /dev/null +++ b/internal/driver/reconnect_listener_test.go @@ -0,0 +1,54 @@ +package driver + +import ( + "net" + "os" + "path/filepath" + "testing" +) + +func TestGuestRecoveryListenerRefusesActiveOwner(t *testing.T) { + path := filepath.Join(shortTempDir(t), "active.sock") + old, err := net.ListenUnix("unix", &net.UnixAddr{Name: path, Net: "unix"}) + if err != nil { + t.Fatal(err) + } + defer old.Close() + before, err := os.Lstat(path) + if err != nil { + t.Fatal(err) + } + if next, err := bindRecoveredGuestListener(path, os.Geteuid()); err == nil { + next.Close() + t.Fatal("replaced active listener") + } + after, err := os.Lstat(path) + if err != nil || !os.SameFile(before, after) { + t.Fatal("changed active listener inode") + } +} +func TestGuestRecoveryListenerAcceptsOnlyStaleOwnedSocket(t *testing.T) { + path := filepath.Join(shortTempDir(t), "stale.sock") + old, err := net.ListenUnix("unix", &net.UnixAddr{Name: path, Net: "unix"}) + if err != nil { + t.Fatal(err) + } + old.SetUnlinkOnClose(false) + old.Close() + next, err := bindRecoveredGuestListener(path, os.Geteuid()) + if err != nil { + t.Fatal(err) + } + next.Close() + if err := os.WriteFile(path, []byte("synthetic marker"), 0600); err != nil { + t.Fatal(err) + } + if next, err := bindRecoveredGuestListener(path, os.Geteuid()); err == nil { + next.Close() + t.Fatal("replaced a non-socket") + } + data, _ := os.ReadFile(path) + if string(data) != "synthetic marker" { + t.Fatal("changed non-socket") + } +} diff --git a/internal/driver/reconnect_lock.go b/internal/driver/reconnect_lock.go new file mode 100644 index 0000000..9208160 --- /dev/null +++ b/internal/driver/reconnect_lock.go @@ -0,0 +1,34 @@ +package driver + +import ( + "fmt" + "os" + "path/filepath" + + "golang.org/x/sys/unix" +) + +// lockGuestState excludes another current-version microVM runner before discovery or +// cleanup can touch surviving VMs. Keep the file in place: unlinking it would +// permit a second owner to lock a different inode at the same path. +func lockGuestState(dir string) (*os.File, error) { + fd, err := unix.Open(filepath.Join(dir, "guest-reconnect.lock"), unix.O_CREAT|unix.O_RDWR|unix.O_NOFOLLOW|unix.O_CLOEXEC|unix.O_NONBLOCK, 0600) + if err != nil { + return nil, fmt.Errorf("microvm: open recovery ownership: %w", err) + } + f := os.NewFile(uintptr(fd), "guest-reconnect.lock") + var st unix.Stat_t + if err := unix.Fstat(fd, &st); err != nil { + f.Close() + return nil, fmt.Errorf("microvm: inspect recovery ownership: %w", err) + } + if st.Mode&unix.S_IFMT != unix.S_IFREG || st.Uid != uint32(os.Geteuid()) || st.Mode&0777 != 0600 || st.Nlink != 1 { + f.Close() + return nil, fmt.Errorf("microvm: unsafe recovery ownership file") + } + if err := unix.Flock(fd, unix.LOCK_EX|unix.LOCK_NB); err != nil { + f.Close() + return nil, fmt.Errorf("microvm: recovery state already owned: %w", err) + } + return f, nil +} diff --git a/internal/driver/reconnect_lock_test.go b/internal/driver/reconnect_lock_test.go new file mode 100644 index 0000000..b16bc79 --- /dev/null +++ b/internal/driver/reconnect_lock_test.go @@ -0,0 +1,57 @@ +package driver + +import ( + "os" + "path/filepath" + "testing" +) + +func TestGuestStateOwnershipIsExclusive(t *testing.T) { + dir := t.TempDir() + first, err := lockGuestState(dir) + if err != nil { + t.Fatal(err) + } + if second, err := lockGuestState(dir); err == nil { + second.Close() + t.Fatal("two runners acquired recovery ownership") + } + first.Close() + next, err := lockGuestState(dir) + if err != nil { + t.Fatal("exited owner prevented recovery") + } + next.Close() +} +func TestGuestStateOwnershipRejectsSymlink(t *testing.T) { + dir := t.TempDir() + target := filepath.Join(dir, "target") + if err := os.WriteFile(target, []byte("unchanged_test"), 0600); err != nil { + t.Fatal(err) + } + if err := os.Symlink(target, filepath.Join(dir, "guest-reconnect.lock")); err != nil { + t.Fatal(err) + } + if f, err := lockGuestState(dir); err == nil { + f.Close() + t.Fatal("followed recovery ownership symlink") + } + data, _ := os.ReadFile(target) + if string(data) != "unchanged_test" { + t.Fatal("altered symlink target") + } +} + +func TestGuestStateOwnershipCannotBeBypassedByDisablingReconnect(t *testing.T) { + m, _ := testMicrovm(t, MicrovmOpts{GuestReconnect: true}) + defer m.stateLock.Close() + opts := m.opts + opts.GuestReconnect = false + other, err := NewMicrovm(opts) + if err == nil { + if other.stateLock != nil { + other.stateLock.Close() + } + t.Fatal("capability-off driver bypassed active ownership") + } +} diff --git a/internal/driver/resume_identity_test.go b/internal/driver/resume_identity_test.go new file mode 100644 index 0000000..81772c7 --- /dev/null +++ b/internal/driver/resume_identity_test.go @@ -0,0 +1,342 @@ +package driver + +import ( + "context" + "errors" + "os" + "testing" + "time" +) + +type coldIdentityEngine struct { + *SimulatedEngine + identity guestHostIdentity + failIdentity, failStop bool + failLaunch bool +} + +func (e *coldIdentityEngine) guestIdentity(VMMConfig, int) (guestHostIdentity, error) { + if e.failIdentity { + return guestHostIdentity{}, errors.New("synthetic identity refusal") + } + return e.identity, nil +} +func (e *coldIdentityEngine) Launch(ctx context.Context, cfg VMMConfig) error { + if err := e.SimulatedEngine.Launch(ctx, cfg); err != nil { + return err + } + if e.failLaunch { + return errors.New("synthetic ambiguous launch reply") + } + return nil +} +func (e *coldIdentityEngine) Stop(ctx context.Context, id string) error { + if e.failStop { + return errors.New("synthetic stop refusal") + } + return e.SimulatedEngine.Stop(ctx, id) +} +func TestColdResumeRefreshesVerifiedHostIdentity(t *testing.T) { + for _, outcome := range []string{"verified", "identity_refused", "stop_refused", "launch_stop_refused"} { + t.Run(outcome, func(t *testing.T) { + ctx := context.Background() + dir := shortTempDir(t) + engine := &coldIdentityEngine{SimulatedEngine: NewSimulatedEngineWithDir(dir), identity: guestHostIdentity{BootID: "boot_test", StartTime: 1, NamespaceDevice: 1, NamespaceInode: 1}} + m, _ := testMicrovm(t, MicrovmOpts{StateDir: dir, Engine: engine, GuestReconnect: true}) + m.SetHost(&stubMicrovmHost{}) + h, err := m.Create(ctx, Spec{SessionID: "session_test", GuestReconnect: 1}) + if err != nil { + t.Fatal(err) + } + if err = m.Suspend(ctx, h.ID, false); err != nil { + t.Fatal(err) + } + engine.identity.StartTime = 2 + engine.identity.NamespaceInode = 2 + engine.failIdentity = outcome != "verified" + engine.failStop = outcome == "stop_refused" || outcome == "launch_stop_refused" + engine.failLaunch = outcome == "launch_stop_refused" + _, err = m.Resume(ctx, h.ID) + if outcome == "verified" { + if err != nil { + t.Fatal(err) + } + if m.instances[h.ID].Identity != engine.identity { + t.Fatal("cold boot retained old process identity") + } + } else { + if err == nil { + t.Fatal("unverified cold boot accepted") + } + if engine.failStop { + if m.instances[h.ID].slot == nil || m.instances[h.ID].Identity != (guestHostIdentity{}) { + t.Fatal("possibly live VM lost resource ownership or kept trusted identity") + } + } else if state, _ := engine.State(ctx, h.ID); state == VMMStateRunning { + t.Fatal("unverified VM was not stopped") + } + } + engine.failStop = false + _ = m.Destroy(ctx, h.ID) + }) + } +} + +func TestResumeStatusDurablyFencesDelayedLaunch(t *testing.T) { + ctx := context.Background() + m, engine := testMicrovm(t, MicrovmOpts{}) + m.SetHost(&stubMicrovmHost{}) + h, err := m.Create(ctx, Spec{SessionID: "session_test", PlacementGeneration: 1}) + if err != nil { + t.Fatal(err) + } + if err = m.Suspend(ctx, h.ID, false); err != nil { + t.Fatal(err) + } + state, generation, err := m.ResumeStatus(ctx, h.ID, 2) + if err != nil || state != "suspended_cold" || generation != 2 { + t.Fatal("unreceived resume was not durably fenced") + } + if _, err = m.ResumePlacement(ctx, h.ID, 2); err == nil { + t.Fatal("delayed canceled claim launched a VM") + } + stopTestMicrovmRunner(t, m) + recovered, _ := testMicrovm(t, MicrovmOpts{StateDir: m.opts.StateDir, BaseRootfs: m.opts.BaseRootfs, Engine: engine, Net: m.opts.Net, Format: m.opts.Format}) + listed, err := recovered.List(ctx) + if err != nil || len(listed) != 1 || listed[0].PlacementGeneration != 2 { + t.Fatal("resume fence did not survive runner restart") + } + if _, err = recovered.ResumePlacement(ctx, h.ID, 2); err == nil { + t.Fatal("restart allowed canceled claim to launch") + } + _ = recovered.Destroy(ctx, h.ID) +} + +func TestColdResumeDeletionWaitsForLaunchOwnership(t *testing.T) { + engine := &parkedLaunchEngine{SimulatedEngine: NewSimulatedEngine(), entered: make(chan struct{}, 1), release: make(chan struct{})} + m, _ := testMicrovm(t, MicrovmOpts{Engine: engine}) + m.SetHost(&stubMicrovmHost{}) + ctx := context.Background() + h, err := m.Create(ctx, Spec{SessionID: "session_test", PlacementGeneration: 1}) + if err != nil { + t.Fatal(err) + } + if err = m.Suspend(ctx, h.ID, false); err != nil { + t.Fatal(err) + } + engine.park() + done := make(chan error, 1) + go func() { _, err := m.ResumePlacement(ctx, h.ID, 2); done <- err }() + select { + case <-engine.entered: + case <-time.After(5 * time.Second): + t.Fatal("launch not reached") + } + // Deletion cannot declare the old stopped VM gone while a new launch is + // still using its disks and allocated namespace. + bounded, cancel := context.WithTimeout(ctx, 50*time.Millisecond) + defer cancel() + err = m.Destroy(bounded, h.ID) + if !errors.Is(err, context.DeadlineExceeded) { + t.Errorf("destroy during launch = %v, want bounded wait", err) + } + workspace, _ := m.workspaceDiskPath("session_test") + if _, err := os.Stat(workspace); err != nil { + t.Error("workspace removed during launch") + } + close(engine.release) + if err := <-done; err != nil { + t.Errorf("launch failed: %v", err) + } + if err := m.Destroy(ctx, h.ID); err != nil { + t.Fatal(err) + } + if state, _ := engine.State(ctx, h.ID); state == VMMStateRunning { + t.Fatal("VM survived deletion") + } +} + +type crashWindowEngine struct { + *SimulatedEngine + before func(VMMConfig) +} + +func (e *crashWindowEngine) Launch(ctx context.Context, cfg VMMConfig) error { + if e.before != nil { + e.before(cfg) + } + return e.SimulatedEngine.Launch(ctx, cfg) +} +func TestColdResumeCrashBeforeFinalMetadataRetainsResources(t *testing.T) { + ctx := context.Background() + engine := &crashWindowEngine{SimulatedEngine: NewSimulatedEngine()} + m, _, network := testMicrovmNet(t, MicrovmOpts{Engine: engine}) + m.SetHost(&stubMicrovmHost{}) + h, err := m.Create(ctx, Spec{SessionID: "session_test", PlacementGeneration: 1}) + if err != nil { + t.Fatal(err) + } + if err = m.Suspend(ctx, h.ID, false); err != nil { + t.Fatal(err) + } + var launchRecord []byte + engine.before = func(VMMConfig) { + used, _, placements, capErr := m.CapacitySnapshot(ctx) + if capErr != nil || used != 1 || placements["session_test"] != 2 { + t.Fatal("pending launch lacks an atomic counted placement") + } + launchRecord, err = os.ReadFile(m.instanceMetaPath(h.ID)) + if err != nil { + t.Fatal(err) + } + } + if _, err = m.ResumePlacement(ctx, h.ID, 2); err != nil { + t.Fatal(err) + } + liveNamespace := m.instances[h.ID].Cfg.Netns + // Reproduce a crash after Launch but before its final metadata publication, + // using the exact bytes the driver had already committed at that boundary. + if err = os.WriteFile(m.instanceMetaPath(h.ID), launchRecord, 0600); err != nil { + t.Fatal(err) + } + stopTestMicrovmRunner(t, m) + recovered, _, _ := testMicrovmNet(t, MicrovmOpts{StateDir: m.opts.StateDir, BaseRootfs: m.opts.BaseRootfs, Engine: engine, Net: network, Format: m.opts.Format}) + defer recovered.Destroy(ctx, h.ID) + if state, _ := engine.State(ctx, h.ID); state != VMMStateRunning { + t.Fatal("surviving VM missing") + } + names, err := network.ListNetns(ctx) + if err != nil { + t.Fatal(err) + } + retained := false + for _, name := range names { + retained = retained || name == liveNamespace + } + if !retained { + t.Fatal("recovery reclaimed a live cold-resume namespace") + } + if rec := recovered.instances[h.ID]; rec.slot == nil || !rec.RecoveryBlocked || rec.PlacementGeneration != 2 { + t.Fatal("incomplete launch lost quarantined ownership") + } +} + +func TestColdResumeInterruptedLaunchWithoutVMCanSettle(t *testing.T) { + ctx := context.Background() + engine := &crashWindowEngine{SimulatedEngine: NewSimulatedEngine()} + m, _, network := testMicrovmNet(t, MicrovmOpts{Engine: engine}) + m.SetHost(&stubMicrovmHost{}) + h, err := m.Create(ctx, Spec{SessionID: "session_test", PlacementGeneration: 1}) + if err != nil { + t.Fatal(err) + } + if err = m.Suspend(ctx, h.ID, false); err != nil { + t.Fatal(err) + } + var launchRecord []byte + engine.before = func(VMMConfig) { + used, _, placements, capErr := m.CapacitySnapshot(ctx) + if capErr != nil || used != 1 || placements["session_test"] != 2 { + t.Fatal("pending launch lacks an atomic counted placement") + } + launchRecord, err = os.ReadFile(m.instanceMetaPath(h.ID)) + if err != nil { + t.Fatal(err) + } + } + if _, err = m.ResumePlacement(ctx, h.ID, 2); err != nil { + t.Fatal(err) + } + if err := engine.Stop(ctx, h.ID); err != nil { + t.Fatal(err) + } + // Reproduce a crash after Launch but before its final metadata publication, + // using the exact bytes the driver had already committed at that boundary. + if err = os.WriteFile(m.instanceMetaPath(h.ID), launchRecord, 0600); err != nil { + t.Fatal(err) + } + stopTestMicrovmRunner(t, m) + recovered, _, _ := testMicrovmNet(t, MicrovmOpts{StateDir: m.opts.StateDir, BaseRootfs: m.opts.BaseRootfs, Engine: engine, Net: network, Format: m.opts.Format}) + defer recovered.Destroy(ctx, h.ID) + state, generation, err := recovered.ResumeStatus(ctx, h.ID, 2) + if err != nil || state != "suspended_cold" || generation != 2 { + t.Fatalf("interrupted launch without a VM cannot settle: %s %d %v", state, generation, err) + } + if _, err := recovered.ResumePlacement(ctx, h.ID, 2); err == nil { + t.Fatal("canceled placement was replayable after restart") + } +} + +func TestInitialCreateRetainsUnstoppableVM(t *testing.T) { + ctx := context.Background() + engine := &coldIdentityEngine{SimulatedEngine: NewSimulatedEngine(), failIdentity: true, failStop: true} + m, _, network := testMicrovmNet(t, MicrovmOpts{Engine: engine, GuestReconnect: true, TotalSlots: 2}) + m.SetHost(&stubMicrovmHost{}) + if _, err := m.Create(ctx, Spec{SessionID: "session_test", PlacementGeneration: 1, GuestReconnect: 1}); err == nil { + t.Fatal("expected identity refusal") + } + if state, _ := engine.State(ctx, "mvm-1"); state != VMMStateRunning { + t.Fatal("fixture lost live VM") + } + names, err := network.ListNetns(ctx) + if err != nil { + t.Fatal(err) + } + if len(names) == 0 { + t.Fatal("Create rollback released namespace despite Stop failure and surviving VM") + } + if rec := m.instances["mvm-1"]; rec == nil || !rec.RecoveryBlocked { + t.Fatal("surviving failed Create has no blocked ownership record") + } + if err := m.RemoveWorkspace(ctx, "session_test"); err == nil { + t.Error("workspace removed under uncertain VM") + } + if _, err := m.Create(ctx, Spec{SessionID: "session_test", PlacementGeneration: 1, GuestReconnect: 1}); err == nil { + t.Error("uncertain workspace reused by another create") + } + if state, _ := engine.State(ctx, "mvm-2"); state == VMMStateRunning { + t.Error("retry launched another VM on an uncertain workspace") + } + engine.failStop = false + _ = m.Destroy(ctx, "mvm-1") +} + +func TestInitialCreateCrashBeforeMetadataRetainsResources(t *testing.T) { + ctx := context.Background() + engine := &crashWindowEngine{SimulatedEngine: NewSimulatedEngine()} + m, _, network := testMicrovmNet(t, MicrovmOpts{Engine: engine}) + m.SetHost(&stubMicrovmHost{}) + var intent []byte + engine.before = func(cfg VMMConfig) { + var err error + intent, err = os.ReadFile(m.instanceMetaPath(cfg.ID)) + if err != nil { + t.Fatal(err) + } + } + h, err := m.Create(ctx, Spec{SessionID: "create_test", PlacementGeneration: 1}) + if err != nil { + t.Fatal(err) + } + ns := m.instances[h.ID].Cfg.Netns + if err := os.WriteFile(m.instanceMetaPath(h.ID), intent, 0600); err != nil { + t.Fatal(err) + } + stopTestMicrovmRunner(t, m) + recovered, _, _ := testMicrovmNet(t, MicrovmOpts{StateDir: m.opts.StateDir, BaseRootfs: m.opts.BaseRootfs, Engine: engine, Net: network, Format: m.opts.Format}) + defer recovered.Destroy(ctx, h.ID) + if rec := recovered.instances[h.ID]; rec == nil || !rec.RecoveryBlocked || rec.slot == nil { + t.Fatal("initial launch lost blocked ownership on restart") + } + names, err := network.ListNetns(ctx) + if err != nil { + t.Fatal(err) + } + found := false + for _, name := range names { + found = found || name == ns + } + if !found { + t.Fatal("live initial namespace reclaimed") + } +} diff --git a/internal/driver/testdata/fakefirecracker/main.go b/internal/driver/testdata/fakefirecracker/main.go new file mode 100644 index 0000000..b92fd21 --- /dev/null +++ b/internal/driver/testdata/fakefirecracker/main.go @@ -0,0 +1,11 @@ +// A native child for process-identity and reaping tests. It accepts the test's +// argv unchanged and exits on TERM; it does not emulate a VM or API socket. +package main + +import ( + "os" + "os/signal" + "syscall" +) + +func main() { c := make(chan os.Signal, 1); signal.Notify(c, syscall.SIGTERM); <-c } diff --git a/internal/e2e/idlestop_conflict_test.go b/internal/e2e/idlestop_conflict_test.go index 8594d9d..6f896e8 100644 --- a/internal/e2e/idlestop_conflict_test.go +++ b/internal/e2e/idlestop_conflict_test.go @@ -68,8 +68,8 @@ func (h *heldColdSuspend) Suspend(ctx context.Context, id string, warm bool) err // dispatched into the window. // // What they get must be a CONFLICT: 409, code `conflict`, the handler's -// sentence, the row untouched, and the very same command succeeding once the -// stop lands. Before the fix it was 500 `internal`, "could not resume +// sentence, a pending resume claim, and the same command succeeding after +// versioned reconciliation proves the stop landed. Before the fix it was 500 `internal`, "could not resume // session", about a runner that was healthy and answering all along. func TestIdleAutoStopRefusesAResumeAsAConflict(t *testing.T) { f := newFleet(t) @@ -130,26 +130,18 @@ func TestIdleAutoStopRefusesAResumeAsAConflict(t *testing.T) { t.Errorf("message = %q, want the handler's own sentence", api.Message) } - // And nothing moved: the row still says stopped, so the session is not - // left claiming to run on a container that is being stopped. + // The committed claim stays pending until a versioned runner observation + // proves the refused command cannot launch later. An unversioned stop + // event cannot roll it back. after := f.list()[created.ID] - if after.State != "suspended_cold" { - t.Errorf("state after the refused resume = %q, want suspended_cold", after.State) + if after.State != "resuming" { + t.Fatalf("state after refused resume = %q, want resuming", after.State) } - // The stop lands, and the refusal proves to have meant "not yet": the same - // command, from the same client, brings the session back. - // - // Retried against the ROW rather than against the resume's own status, - // because two known races sit in this window and neither is what the scene - // is about. The resume can still arrive before `docker stop` has returned - // (refused again, correctly); and the auto-stop's own `suspended_cold` - // event — deliberately unfenced, since it carries no placement generation - // (internal/runnerd/agent.go, and the PR's follow-up 2) — can land AFTER - // the resume and put the row back to stopped over a container that is now - // running. Both converge: the event fires once, and the next pass resumes - // again. close(held.release) + // Fake.Suspend changes container state only; explicitly model the cold + // stop's process/socket death before asking for evidence of a cold guest. + f.dropSessiond(created.ID) resumed := func() bool { if f.list()[created.ID].State == "running" { return true @@ -158,5 +150,5 @@ func TestIdleAutoStopRefusesAResumeAsAConflict(t *testing.T) { _ = f.client().Do(http.MethodPost, "/v0/sessions/"+created.ID+"/resume", nil, nil) return f.list()[created.ID].State == "running" } - waitUntil(t, 60*time.Second, "the session to come back once the stop has landed", resumed) + waitUntil(t, 3*time.Minute, "versioned reconciliation to settle the refused claim and permit resume", resumed) } diff --git a/internal/runnerd/agent.go b/internal/runnerd/agent.go index 77302c5..08450ec 100644 --- a/internal/runnerd/agent.go +++ b/internal/runnerd/agent.go @@ -137,7 +137,8 @@ func (cfg AgentConfig) resetsBackoff(establishedAt time.Time) bool { // runner sends afterwards. Atomic because the accept is handled on the // reader while events fire from session goroutines. type agentSessionState struct { - generation atomic.Uint64 + recoveryStarted atomic.Bool + generation atomic.Uint64 } // jitter returns a random duration in [0, d/2) — timing spread, not security. @@ -263,16 +264,6 @@ func (s *Server) agentSession(ctx context.Context, cfg AgentConfig) (established defer s.reconnectControl.CompareAndSwap(rc, nil) send := func(m runner.FromRunner) { - cctx, ccancel := context.WithTimeout(ctx, 5*time.Second) - m.Used, m.Total, _ = s.drv.Capacity(cctx) // best-effort; piggybacked on every message - ccancel() - // The two counts that turn "no free capacity" into something a person - // can act on: how much of `used` is a working agent and how much is a - // sandbox whose agent has finished. They ride the same message the - // used/total pair already does, from the registry rather than the - // driver — docker cannot say whether a container's child is still - // running; only sessiond's report can, and this runner keeps it. - m.Active, m.IdleExited = s.reg.counts() // The two generations every report carries (D19), stamped in the one // place every report passes through. The runner's own is whatever // controld granted this connection; the session's is the one its @@ -281,29 +272,9 @@ func (s *Server) agentSession(ctx context.Context, cfg AgentConfig) (established switch m.Type { case "event": m.Generation = ag.generation.Load() - // Every event about a session echoes the placement generation its - // create carried — except the runner's own idle auto-stop, which - // must carry NONE. A cold resume opens a new placement generation - // on the control plane's row but sends the runner no new value, so - // this entry's is stale by construction from the first resume on; - // stamping it would guarantee the report is fenced as stale, and a - // fenced auto-stop leaves the row reading "running" over a - // container that is stopped — a session `rainier attach` then - // refuses to resume and cannot reach, until the runner happens to - // reconnect. Zero means "not carried" and fences nothing, which is - // safe here specifically because the report is about the sandbox - // this runner holds right now, and a session re-placed onto a - // DIFFERENT runner is still fenced by the runner identity the - // service checks first. Carrying the generation on `resume` is the - // real fix and is a separate change (protocol + control plane). - // - // The test is on the state rather than on which call site - // produced it, and that is right for both producers: reannounce - // renders the same word for the same registry state, and it is - // equally a report about the sandbox this runner holds right - // now, whose generation is equally unknowable to it. Anything - // NEW that fires this state would have to be one too. - if m.Session != "" && m.State != "suspended_cold" { + // Create and cold resume both install the committed placement + // before launch; every lifecycle event can retain that fence. + if m.Session != "" { m.PlacementGeneration = s.reg.placementGeneration(m.Session) } case "result": @@ -332,11 +303,9 @@ func (s *Server) agentSession(ctx context.Context, cfg AgentConfig) (established }) defer s.SetOnSessionRPC(nil) - used, total, _ := s.drv.Capacity(ctx) - active, idleExited := s.reg.counts() - ann := runner.FromRunner{Type: "announce", Proto: runner.ProtocolVersion, Runner: cfg.RunnerName, - Sessions: s.Announce(), Used: used, Total: total, Active: active, IdleExited: idleExited, - Capabilities: buildCapabilities(cfg.Capabilities, s.driverCapabilities()...)} + ann := runner.FromRunner{Type: "announce", Proto: runner.ProtocolVersion, Runner: cfg.RunnerName, Sessions: s.Announce(), Capabilities: buildCapabilities(cfg.Capabilities, s.driverCapabilities()...)} + s.stampCapacity(connCtx, &ann) + if err := wsjson.Write(connCtx, c, ann); err != nil { return false, err // nothing can have been accepted before the announce } @@ -356,6 +325,7 @@ func (s *Server) agentSession(ctx context.Context, cfg AgentConfig) (established for { select { case m := <-out: + s.stampCapacity(connCtx, &m) if err := wsjson.Write(connCtx, c, m); err != nil { cancel() // a dead write direction means this connection // is done; unblock the reader below too, not just this @@ -416,6 +386,15 @@ func (s *Server) agentSession(ctx context.Context, cfg AgentConfig) (established s.execute(ctx, m, send, cfg, ag) continue } + if m.Type == "session_rpc" { + target := s.guestRPCTarget(connCtx, m.Session, ag) + go s.execute(connCtx, m, send, cfg, ag, target) + continue + } + if m.Type == "resume" || m.Type == "resume_status" { + go s.execute(connCtx, m, send, cfg, ag) + continue + } go s.execute(ctx, m, send, cfg, ag) // ops are slow (docker); never block the reader } } @@ -425,7 +404,7 @@ func (s *Server) agentSession(ctx context.Context, cfg AgentConfig) (established // slow docker op never blocks the next command from being read — except for // the "accept", which the reader runs inline because it is negotiation, not // work, and everything read after it depends on it having happened. -func (s *Server) execute(ctx context.Context, m runner.ToRunner, send func(runner.FromRunner), cfg AgentConfig, ag *agentSessionState) { +func (s *Server) execute(ctx context.Context, m runner.ToRunner, send func(runner.FromRunner), cfg AgentConfig, ag *agentSessionState, received ...*guestRPCTarget) { switch m.Type { case "accept": // controld's answer to the announce, and the first thing it sends. @@ -434,6 +413,14 @@ func (s *Server) execute(ctx context.Context, m runner.ToRunner, send func(runne // informational — the set controld will schedule on, which is this // runner's own claims minus anything it refused. ag.generation.Store(m.Generation) + for _, capability := range m.Capabilities { + if capability == runner.CapabilityGuestReconnectV1 && ag.recoveryStarted.CompareAndSwap(false, true) { + if rc := s.reconnectControl.Load(); rc != nil && rc.state == ag { + go s.recoverGuests(rc) + } + break + } + } log.Printf("agent: accepted at generation %d with %d capabilities", m.Generation, len(m.Capabilities)) case "create": var spec driver.Spec @@ -446,22 +433,7 @@ func (s *Server) execute(ctx context.Context, m runner.ToRunner, send func(runne // carry — and this runner's only job is to hand it to the // container. Nothing here logs Env; its values are secrets as // often as not. - spec = driver.Spec{ - Name: m.Spec.Name, Image: m.Spec.Image, Cmd: m.Spec.Cmd, EgressAllow: m.Spec.EgressAllow, - Setup: m.Spec.Setup, SetupTimeoutSec: m.Spec.SetupTimeoutSec, Env: m.Spec.Env, - Repos: driverRepos(m.Spec.Repos), - Init: m.Spec.Init, InitTimeoutSec: m.Spec.InitTimeoutSec, - GitAuthorName: m.Spec.GitAuthorName, GitAuthorEmail: m.Spec.GitAuthorEmail, - Home: driverHome(m.Spec.Home), - // The microVM bootstrap pair, carried through like - // everything else. This runner does not read the token — it - // goes into the guest's boot configuration and is dropped — - // and it does not check the names against Env. The DRIVER - // decides what an unwithheld create means, because the answer - // is different for each one: Docker has always accepted the - // values and still does. - BootstrapToken: m.Spec.BootstrapToken, SecretNames: m.Spec.SecretNames, - } + spec = guestDriverSpec(*m.Spec) allow = m.Spec.EgressAllow } // Idempotency lives inside CreateWithID's own putIfAbsent now, not a @@ -489,8 +461,14 @@ func (s *Server) execute(ctx context.Context, m runner.ToRunner, send func(runne // controld settles the dispatch before it sees the state. s.reannounce(m.Session, send) } + case "resume_status": + send(s.resumeStatus(ctx, m)) case "suspend", "resume": - err := s.Op(ctx, m.Session, m.Type, m.Warm) + if ctx.Err() != nil { + send(runner.FromRunner{Type: "result", ReqID: m.ReqID}) + return + } + err := s.opAtPlacement(ctx, m.Session, m.Type, m.Warm, m.PlacementGeneration) // Conflict is what tells controld apart the two ways this can be // not-ok: a command that failed, and one the runner refused because // it is already stopping (or still creating) this sandbox. Without @@ -577,7 +555,19 @@ func (s *Server) execute(ctx context.Context, m runner.ToRunner, send func(runne // one) comes from inside that container, not from here. So there is no // result to send — m.ReqID is zero on this type — and correlation // lives entirely in the envelope's own id. - s.forwardSessionRPC(m, send) + if ctx.Err() != nil { + return + } + if rc := s.reconnectControl.Load(); rc != nil && (rc.state != ag || rc.ctx.Err() != nil) { + return + } + var target *guestRPCTarget + if len(received) > 0 { + target = received[0] + } else { + target = s.guestRPCTarget(ctx, m.Session, ag) + } + s.forwardSessionRPC(m, send, target) case "dial_attach": // Deliberately not in a goroutine of its own: agentSession's read // loop already runs one execute per inbound command precisely so a @@ -600,7 +590,7 @@ func (s *Server) execute(ctx context.Context, m runner.ToRunner, send func(runne // answer that was never coming. A RESPONSE that cannot be delivered is only // logged: answering an answer is meaningless, and the sandbox that asked has // already lost its own pending entry along with the conn. -func (s *Server) forwardSessionRPC(m runner.ToRunner, send func(runner.FromRunner)) { +func (s *Server) forwardSessionRPC(m runner.ToRunner, send func(runner.FromRunner), received ...*guestRPCTarget) { if m.RPC == nil { log.Printf("agent: session_rpc for %s carried no envelope; ignoring", m.Session) return @@ -628,7 +618,33 @@ func (s *Server) forwardSessionRPC(m runner.ToRunner, send func(runner.FromRunne } return } - err := s.sendSessionRPC(m.Session, env) + var err error + row, exists := s.reg.snapshot(m.Session) + var pending guestRPCPending + correlated := false + if env.Method == "resp" { + pending, correlated = s.guestForwards.take(env.ID, m.Session) + } + if correlated { + env.ID = pending.guestID + err = s.sendGuestRPC(&pending.target, env) + } else if exists && row.guestReconnect { + if env.Method == "resp" { + return + } + var target *guestRPCTarget + if len(received) > 0 { + target = received[0] + } + err = s.sendGuestRPC(target, env) + } else { + // A command captured for a negotiated guest never falls back to a + // replacement legacy entry or waits for a different hub. + if len(received) > 0 && received[0] != nil { + return + } + err = s.sendSessionRPC(m.Session, env) + } if err == nil { return } @@ -916,3 +932,38 @@ func driverHome(h *runner.HomeMount) *driver.HomeMount { } return &driver.HomeMount{Volume: h.Volume, Path: h.Path} } + +func guestDriverSpec(spec runner.Spec) driver.Spec { + return driver.Spec{ + GuestReconnect: spec.GuestReconnect, Name: spec.Name, Image: spec.Image, Cmd: spec.Cmd, EgressAllow: spec.EgressAllow, + Setup: spec.Setup, SetupTimeoutSec: spec.SetupTimeoutSec, Env: spec.Env, + Repos: driverRepos(spec.Repos), + Init: spec.Init, InitTimeoutSec: spec.InitTimeoutSec, + GitAuthorName: spec.GitAuthorName, GitAuthorEmail: spec.GitAuthorEmail, + Home: driverHome(spec.Home), + // The microVM bootstrap pair, carried through like + // everything else. This runner does not read the token — it + // goes into the guest's boot configuration and is dropped — + // and it does not check the names against Env. The DRIVER + // decides what an unwithheld create means, because the answer + // is different for each one: Docker has always accepted the + // values and still does. + BootstrapToken: spec.BootstrapToken, SecretNames: spec.SecretNames, + } +} + +// Stamp at the single writer, so every origin (including runner-local RPCs) +// has one ordered capacity sample and cannot overwrite a newer observation. +func (s *Server) stampCapacity(parent context.Context, m *runner.FromRunner) { + ctx, cancel := context.WithTimeout(parent, 5*time.Second) + defer cancel() + if d, ok := s.drv.(interface { + CapacitySnapshot(context.Context) (int, int, map[string]uint64, error) + }); ok { + m.Used, m.Total, m.CapacityPlacements, _ = d.CapacitySnapshot(ctx) + } else { + m.Used, m.Total, _ = s.drv.Capacity(ctx) + m.CapacityPlacements = nil + } + m.Active, m.IdleExited = s.reg.counts() +} diff --git a/internal/runnerd/idlestop_e2e_test.go b/internal/runnerd/idlestop_e2e_test.go index 05f5a4b..f4eb3f2 100644 --- a/internal/runnerd/idlestop_e2e_test.go +++ b/internal/runnerd/idlestop_e2e_test.go @@ -338,16 +338,9 @@ func TestAttachCountsBeforeTheFirstClientFrame(t *testing.T) { } } -// TestIdleStopEventCarriesNoPlacementGeneration is the review's other severe -// finding. A cold resume opens a NEW placement generation on the control -// plane's row but sends the runner no new value, so this runner's is stale -// from the first resume on — and the control plane fences an event whose -// placement generation is not the row's. A fenced auto-stop would leave the -// row reading "running" over a container that is stopped: a session `rainier -// attach` refuses to resume and cannot reach. So this one event carries none, -// where every other event about a session still carries the one its create -// did. -func TestIdleStopEventCarriesNoPlacementGeneration(t *testing.T) { +// Cold resume now carries the committed placement to the runner, so idle +// stop must preserve that fence rather than sending an unversioned event. +func TestIdleStopEventCarriesCurrentPlacementGeneration(t *testing.T) { const placementGen = 7 clk := newFakeClock() rd, srv, conn := runnerWithControld(t, func(s *Server) { s.now = clk.now }) @@ -385,9 +378,8 @@ func TestIdleStopEventCarriesNoPlacementGeneration(t *testing.T) { if ev.Session != id { t.Fatalf("suspended_cold for %q, want %q", ev.Session, id) } - if ev.PlacementGeneration != 0 { - t.Fatalf("suspended_cold placement generation = %d, want 0 — a stale generation would fence the park", - ev.PlacementGeneration) + if ev.PlacementGeneration != placementGen { + t.Fatalf("suspended_cold placement generation = %d, want %d", ev.PlacementGeneration, placementGen) } } diff --git a/internal/runnerd/microvm.go b/internal/runnerd/microvm.go index 161c7bb..d07e821 100644 --- a/internal/runnerd/microvm.go +++ b/internal/runnerd/microvm.go @@ -43,7 +43,8 @@ var _ driver.MicrovmHost = (*Server)(nil) // whole difference between this door and the WebSocket one, where `register` // believes a query parameter with no authentication on the hop at all. func (s *Server) GuestConnected(sessionID string, conn relay.Conn) { - if _, ok := s.reg.get(sessionID); !ok { + row, exists := s.reg.snapshot(sessionID) + if !exists { // A guest for a session this runner no longer holds: a destroy that // raced the boot. Closing it is what tells the guest to stop; leaving // it open would leak the conn and its goroutine with no registry @@ -52,6 +53,10 @@ func (s *Server) GuestConnected(sessionID string, conn relay.Conn) { _ = conn.Close() return } + if row.guestReconnect { + go func() { _ = s.enrollGuestConnection(context.Background(), sessionID, conn) }() + return + } go s.serveSessionConn(context.Background(), sessionID, conn) } @@ -105,7 +110,7 @@ func isRunnerOriginated(id uint64) bool { return id&runnerOriginatedIDBase != 0 // waiting on that number. func refuseSandboxOrigin(method string, id uint64) string { switch { - case method == runner.MethodBeginGuestReconnect || method == runner.MethodAcceptGuestReconnect: + case method == runner.MethodBeginGuestReconnect || method == runner.MethodAcceptGuestReconnect || method == runner.MethodGuestReconnectConfiguration: return "fenced" case method == runner.MethodMintSessionBootstrap: return "a sandbox may not mint its own bootstrap token; a fresh one is minted by the runner on a cold resume" diff --git a/internal/runnerd/reconnect.go b/internal/runnerd/reconnect.go index 1598f2c..9e5968a 100644 --- a/internal/runnerd/reconnect.go +++ b/internal/runnerd/reconnect.go @@ -38,27 +38,25 @@ func (s *Server) AuthorizeGuestReconnect(ctx context.Context, id string, prove d if prove == nil { return zero, errReconnectInvalid } - row, ok := s.reg.snapshot(id) - if !ok || row.placementGen == 0 || (row.state != "running" && row.state != "starting") { - return zero, errReconnectFenced - } - rc := s.reconnectControl.Load() - if rc == nil || rc.ctx.Err() != nil || rc.state.generation.Load() == 0 { - return zero, errReconnectUnavailable + lease, err := s.acquireGuestReconnect(id) + if err != nil { + return zero, err } - generation := rc.state.generation.Load() - if !s.claimReconnect(id) { - return zero, errReconnectUnavailable + defer lease.close() + return lease.authorize(ctx, prove) +} + +func (lease *guestReconnectLease) authorize(ctx context.Context, prove driver.GuestReconnectProof) (runner.GuestReconnectAcceptResponse, error) { + var zero runner.GuestReconnectAcceptResponse + if prove == nil { + return zero, errReconnectInvalid } - defer s.releaseReconnect(id) + s, id, row, rc, generation := lease.server, lease.row.id, lease.row, lease.control, lease.generation ctx, cancel := context.WithTimeout(ctx, reconnectBudget) defer cancel() stop := context.AfterFunc(rc.ctx, cancel) defer stop() - valid := func() bool { - current, exists := s.reg.snapshot(id) - return ctx.Err() == nil && rc.ctx.Err() == nil && s.reconnectControl.Load() == rc && rc.state.generation.Load() == generation && exists && current.placementGen == row.placementGen && current.boot == row.boot && current.handle == row.handle && (current.state == "running" || current.state == "starting") - } + valid := func() bool { return lease.valid(ctx) } if !valid() { return zero, errReconnectFenced } @@ -128,8 +126,6 @@ func (s *Server) reconnectCall(ctx context.Context, rc *reconnectControl, sessio id, ch := s.runnerRPC.begin(session) defer s.runnerRPC.end(id) msg := runner.FromRunner{Type: "session_req", Session: session, RPC: &runner.RPCEnvelope{ID: id, Method: method, Payload: body}} - msg.Used, msg.Total, _ = s.drv.Capacity(ctx) - msg.Active, msg.IdleExited = s.reg.counts() if ctx.Err() != nil || rc.ctx.Err() != nil { return nil, errReconnectUnavailable } diff --git a/internal/runnerd/reconnect_enrollment.go b/internal/runnerd/reconnect_enrollment.go new file mode 100644 index 0000000..c589ec9 --- /dev/null +++ b/internal/runnerd/reconnect_enrollment.go @@ -0,0 +1,94 @@ +package runnerd + +import ( + "context" + "time" + + "github.com/tokencanopy/rainier/internal/relay" + "github.com/tokencanopy/rainier/protocol/runner" +) + +// Fresh opted-in boots cannot use the legacy environment exchange. Holding +// publication until enrollment commits also prevents ordinary client traffic +// from interleaving with the guest's strict bootstrap response reader. +func (s *Server) enrollGuestConnection(ctx context.Context, id string, conn relay.Conn) error { + success := false + defer func() { + if !success { + conn.Close() + } + }() + ctx, cancel := context.WithTimeout(ctx, 30*time.Second) + defer cancel() + original, ok := s.reg.snapshot(id) + if !ok { + return errReconnectFenced + } + // The guest may dial while driver.Create is returning its handle. Retain the + // original boot while waiting; never adopt a replacement entry with this id. + ticker := time.NewTicker(10 * time.Millisecond) + defer ticker.Stop() + for { + current, exists := s.reg.snapshot(id) + if !exists || current.boot != original.boot { + return errReconnectFenced + } + if current.handle != "" && !current.resumePending { + break + } + select { + case <-ctx.Done(): + return errReconnectUnavailable + case <-ticker.C: + } + } + lease, err := s.acquireGuestReconnect(id) + if err != nil { + return err + } + defer lease.close() + if lease.row.boot != original.boot || lease.row.hub != nil || lease.row.guestEpoch != 0 { + return errReconnectFenced + } + stop := context.AfterFunc(lease.control.ctx, cancel) + defer stop() + event, err := readGuestBootstrapRequest(ctx, conn, runner.MethodEnrollGuestReconnect) + if err != nil { + return err + } + request, err := runner.DecodeGuestReconnectEnrollRequest(event.Payload) + if err != nil || !lease.valid(ctx) { + return errReconnectInvalid + } + payload, err := s.reconnectCall(ctx, lease.control, id, runner.MethodEnrollGuestReconnect, request) + if err != nil || !lease.valid(ctx) { + return errReconnectUnavailable + } + if writeGuestControl(ctx, conn, relay.ControlEvent{Kind: "resp", ID: event.ID, OK: true, Payload: payload}) != nil { + return errReconnectUnavailable + } + if err = lease.install(ctx, 0, conn); err != nil { + return err + } + success = true + return nil +} + +func readGuestBootstrapRequest(ctx context.Context, conn relay.Conn, method string) (relay.ControlEvent, error) { + var zero relay.ControlEvent + reader, ok := conn.(interface { + ReadLimited(context.Context, int) ([]byte, error) + }) + if !ok { + return zero, errReconnectInvalid + } + raw, err := reader.ReadLimited(ctx, runner.GuestReconnectPayloadLimit) + if err != nil { + return zero, errReconnectUnavailable + } + frame, err := relay.Decode(raw) + if err != nil || frame.Type != relay.FrameControl || frame.AttachID != 0 { + return zero, errReconnectInvalid + } + return decodeGuestBootstrapEvent(frame.Payload, method) +} diff --git a/internal/runnerd/reconnect_enrollment_test.go b/internal/runnerd/reconnect_enrollment_test.go new file mode 100644 index 0000000..56bd2e5 --- /dev/null +++ b/internal/runnerd/reconnect_enrollment_test.go @@ -0,0 +1,126 @@ +package runnerd + +import ( + "context" + "encoding/json" + "github.com/tokencanopy/rainier/internal/relay" + "github.com/tokencanopy/rainier/protocol/runner" + "net" + "strings" + "testing" + "time" +) + +func TestGuestEnrollmentRejectsLegacyExchangeBeforePublication(t *testing.T) { + s, _, _ := reconnectPeer(t) + a, b := net.Pipe() + defer a.Close() + defer b.Close() + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + done := make(chan error, 1) + go func() { done <- s.enrollGuestConnection(ctx, "session_test", relay.NetConn(a)) }() + event := relay.ControlEvent{Kind: "req:" + runner.MethodFetchSessionSecrets, ID: 1, Payload: json.RawMessage(`{"protocol":1,"token":"AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA"}`)} + if err := writeGuestControl(ctx, relay.NetConn(b), event); err != nil { + t.Fatal(err) + } + if err := <-done; err == nil { + t.Fatal("legacy exchange accepted for enrolled boot") + } + if _, ok := s.reg.hub("session_test"); ok { + t.Fatal("legacy guest published") + } +} + +func TestGuestEnrollmentPublishesOnlyAfterCommittedExchange(t *testing.T) { + s, _, control := reconnectPeer(t) + a, b := net.Pipe() + defer a.Close() + defer b.Close() + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + done := make(chan error, 1) + go func() { done <- s.enrollGuestConnection(ctx, "session_test", relay.NetConn(a)) }() + guest := relay.NetConn(b) + payload, _ := json.Marshal(runner.GuestReconnectEnrollRequest{Protocol: 1, Token: strings.Repeat("A", 43), BootEpoch: "boot_test", PublicKey: strings.Repeat("A", 43)}) + if err := writeGuestControl(ctx, guest, relay.ControlEvent{Kind: "req:" + runner.MethodEnrollGuestReconnect, ID: 1, Payload: payload}); err != nil { + t.Fatal(err) + } + req := control.readMsg(t) + if req.RPC.Method != runner.MethodEnrollGuestReconnect || !isRunnerOriginated(req.RPC.ID) { + t.Fatal("enrollment not scoped to host") + } + if _, ok := s.reg.hub("session_test"); ok { + t.Fatal("guest published before enrollment committed") + } + replyReconnect(t, control, req, true, []byte(`{"env":{}}`)) + raw, err := guest.Read(ctx) + if err != nil { + t.Fatal(err) + } + frame, _ := relay.Decode(raw) + var event relay.ControlEvent + json.Unmarshal(frame.Payload, &event) + if event.Kind != "resp" || event.ID != 1 || !event.OK { + t.Fatal("enrollment answer missing") + } + if err = <-done; err != nil { + t.Fatal(err) + } + hub, ok := s.reg.hub("session_test") + if !ok { + t.Fatal("enrolled guest not published") + } + hub.Close() +} + +func TestColdGuestEnrollmentWaitsForClaimedDriverResult(t *testing.T) { + s, _, control := reconnectPeer(t) + s.reg.mu.Lock() + row := s.reg.items["session_test"] + row.state = "suspended" + generation := row.placementGen + 1 + s.reg.mu.Unlock() + if _, err := s.claimResumePlacement("session_test", row.handle, generation); err != nil { + t.Fatal(err) + } + claimed, _ := s.reg.snapshot("session_test") + a, b := net.Pipe() + defer a.Close() + defer b.Close() + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + done := make(chan error, 1) + go func() { done <- s.enrollGuestConnection(ctx, "session_test", relay.NetConn(a)) }() + guest := relay.NetConn(b) + payload, _ := json.Marshal(runner.GuestReconnectEnrollRequest{Protocol: 1, Token: strings.Repeat("A", 43), BootEpoch: "boot_new", PublicKey: strings.Repeat("A", 43)}) + written := make(chan error, 1) + go func() { + written <- writeGuestControl(ctx, guest, relay.ControlEvent{Kind: "req:" + runner.MethodEnrollGuestReconnect, ID: 1, Payload: payload}) + }() + select { + case err := <-written: + t.Fatalf("pending launch read enrollment: %v", err) + case <-time.After(30 * time.Millisecond): + } + s.reg.resumed("session_test", true) + if err := <-written; err != nil { + t.Fatal(err) + } + req := control.readMsg(t) + if req.RPC == nil || req.RPC.Method != runner.MethodEnrollGuestReconnect { + t.Fatal("fresh guest did not enroll") + } + replyReconnect(t, control, req, true, []byte(`{"env":{}}`)) + if _, err := guest.Read(ctx); err != nil { + t.Fatal(err) + } + if err := <-done; err != nil { + t.Fatal(err) + } + current, _ := s.reg.snapshot("session_test") + if current.boot != claimed.boot || current.placementGen != generation || current.hub == nil || current.guestEpoch != 0 { + t.Fatal("enrollment published outside claimed fresh boot") + } + current.hub.Close() +} diff --git a/internal/runnerd/reconnect_handoff.go b/internal/runnerd/reconnect_handoff.go new file mode 100644 index 0000000..6d98936 --- /dev/null +++ b/internal/runnerd/reconnect_handoff.go @@ -0,0 +1,232 @@ +package runnerd + +import ( + "context" + "encoding/json" + "strings" + "time" + + "github.com/tokencanopy/rainier/internal/driver" + "github.com/tokencanopy/rainier/internal/relay" + "github.com/tokencanopy/rainier/protocol/runner" +) + +func (lease *guestReconnectLease) AuthorizeGuestReconnect(ctx context.Context, id string, prove driver.GuestReconnectProof) (runner.GuestReconnectAcceptResponse, error) { + if id != lease.row.id { + return runner.GuestReconnectAcceptResponse{}, errReconnectFenced + } + return lease.authorize(ctx, prove) +} + +// ReconnectGuest owns an admitted stream through proof, current configuration, +// redemption and readiness. Failure closes it with no relay publication; success +// transfers it to the registry. The caller must bound admission and verify local +// VM/socket ownership before calling. No cached boot token/config is accepted. +func (s *Server) ReconnectGuest(ctx context.Context, id string, conn relay.Conn) error { + success := false + defer func() { + if !success && conn != nil { + conn.Close() + } + }() + if conn == nil { + return errReconnectInvalid + } + lease, err := s.acquireGuestReconnect(id) + if err != nil { + return err + } + defer lease.close() + accepted, err := driver.AuthorizeGuestConnection(ctx, lease, id, conn) + if err != nil { + return err + } + ctx, cancel := context.WithTimeout(ctx, time.Duration(accepted.ExpiresInSec)*time.Second) + defer cancel() + stop := context.AfterFunc(lease.control.ctx, cancel) + defer stop() + if err = lease.fence(ctx, accepted.Epoch); err != nil { + return err + } + payload, err := s.reconnectCall(ctx, lease.control, id, runner.MethodGuestReconnectConfiguration, runner.GuestReconnectBeginRequest{Protocol: 1}) + if err != nil || !lease.valid(ctx) { + return errReconnectUnavailable + } + fresh, err := runner.DecodeGuestReconnectConfiguration(payload) + if err != nil || fresh.SessionID != id || fresh.PlacementGeneration != lease.row.placementGen { + return errReconnectInvalid + } + spec := guestDriverSpec(*fresh.Spec) + spec.SessionID = id + spec.ProxyURL = s.proxyURL + spec.BootstrapToken = accepted.Token + cfg := driver.GuestBootConfig(spec) + cfg.GuestReconnect = runner.GuestReconnectProtocol + body, _ := json.Marshal(accepted) + if !lease.valid(ctx) || relay.WriteGuestReconnectFrame(ctx, conn, relay.KindGuestReconnectAccepted, body) != nil { + return errReconnectUnavailable + } + writeCtx, writeCancel := context.WithTimeout(ctx, 10*time.Second) + writeErr := writeGuestControl(writeCtx, conn, relay.ControlEvent{Kind: relay.KindBootConfig, Payload: mustGuestJSON(cfg)}) + writeCancel() + if writeErr != nil { + return errReconnectUnavailable + } + if err = lease.redeem(ctx, conn, accepted.Token); err != nil { + return err + } + readyCtx, readyCancel := context.WithTimeout(ctx, reconnectBudget) + defer readyCancel() + event, err := relay.ReadGuestReconnectFrame(readyCtx, conn) + if err != nil || event.Kind != relay.KindGuestReconnectReady { + return errReconnectInvalid + } + ready, err := runner.DecodeGuestReconnectReady(event.Payload) + if err != nil || ready.Epoch != accepted.Epoch || !lease.valid(readyCtx) { + return errReconnectFenced + } + if relay.WriteGuestReconnectFrame(readyCtx, conn, relay.KindGuestReconnectReadyAck, event.Payload) != nil || !lease.valid(readyCtx) { + return errReconnectUnavailable + } + if err = lease.install(ctx, accepted.Epoch, conn); err != nil { + return err + } + success = true + return nil +} + +func (lease *guestReconnectLease) install(ctx context.Context, epoch uint64, conn relay.Conn) error { + s, id := lease.server, lease.row.id + // The hub reader waits until publication. No attachment can interleave with + // configuration, redemption or acknowledgment, and buffered bytes stay on conn. + gate := make(chan struct{}) + authority := &guestRelayAuthority{} + reg := s.reg.registration() + var hub *relay.Hub + hub = relay.NewHubWithControl(lease.control.ctx, &reconnectReadGate{Conn: conn, ready: gate}, func(payload []byte) { + go authority.run(func() { s.routeReconnectedControl(lease, hub, reg, payload) }) + }) + if !lease.publish(ctx, epoch, hub, authority) { + hub.Close() + return errReconnectFenced + } + close(gate) + go s.monitorSessionHub(id, hub) + s.sendOwnedGuestMessage(lease, hub, runner.FromRunner{Type: "event", Session: id, State: "running"}) + return nil +} + +func mustGuestJSON(v any) json.RawMessage { body, _ := json.Marshal(v); return body } +func writeGuestControl(ctx context.Context, c relay.Conn, event relay.ControlEvent) error { + body, err := json.Marshal(event) + if err != nil { + return errReconnectInvalid + } + raw, err := relay.Encode(relay.Frame{Type: relay.FrameControl, Payload: body}) + if err != nil { + return errReconnectInvalid + } + return c.Write(ctx, raw) +} + +type reconnectReadGate struct { + relay.Conn + ready <-chan struct{} +} + +func (c *reconnectReadGate) Read(ctx context.Context) ([]byte, error) { + select { + case <-ctx.Done(): + return nil, ctx.Err() + case <-c.ready: + return c.Conn.Read(ctx) + } +} + +func (lease *guestReconnectLease) redeem(ctx context.Context, c relay.Conn, token string) error { + ctx, cancel := context.WithTimeout(ctx, 30*time.Second) + defer cancel() + event, err := readGuestBootstrapRequest(ctx, c, runner.MethodFetchSessionSecrets) + if err != nil { + return err + } + req, err := decodeGuestRedemption(event.Payload) + if err != nil { + return err + } + if req.Token != token || !lease.valid(ctx) { + return errReconnectFenced + } + payload, err := lease.server.reconnectCall(ctx, lease.control, lease.row.id, runner.MethodFetchSessionSecrets, req) + if err != nil || !lease.valid(ctx) { + return errReconnectUnavailable + } + return writeGuestControl(ctx, c, relay.ControlEvent{Kind: "resp", ID: event.ID, OK: true, Payload: payload}) +} + +func (s *Server) routeReconnectedControl(lease *guestReconnectLease, hub *relay.Hub, reg uint64, payload []byte) { + var event relay.ControlEvent + if json.Unmarshal(payload, &event) != nil { + return + } + method, request := strings.CutPrefix(event.Kind, "req:") + if !request && event.Kind != "resp" { + // Lifecycle events retain their boot/registration ordering guards. + s.reg.mu.Lock() + row, ok := s.reg.items[lease.row.id] + valid := ok && row.boot == lease.row.boot && row.hub == hub && lease.control.ctx.Err() == nil && s.reconnectControl.Load() == lease.control && lease.control.state.generation.Load() == lease.generation + s.reg.mu.Unlock() + if valid { + s.routeControlWithEvents(lease.row.id, lease.row.boot, reg, payload, func(id, state, detail string) { + s.sendOwnedGuestMessage(lease, hub, runner.FromRunner{Type: "event", Session: id, State: state, Detail: detail}) + }) + } + return + } + if event.Kind == "resp" { + method = "resp" + } + if method == "" || event.ID == 0 || refuseSandboxOrigin(method, event.ID) != "" { + return + } + msg := runner.FromRunner{Type: "session_req", Session: lease.row.id, RPC: &runner.RPCEnvelope{ID: event.ID, Method: method, OK: event.OK, Payload: event.Payload}} + s.sendOwnedGuestMessage(lease, hub, msg) +} + +func (s *Server) sendOwnedGuestMessage(lease *guestReconnectLease, hub *relay.Hub, msg runner.FromRunner) { + ctx, cancel := context.WithTimeout(lease.control.ctx, reconnectBudget) + defer cancel() + msg.Used, msg.Total, _ = s.drv.Capacity(ctx) + msg.Active, msg.IdleExited = s.reg.counts() + // Check ownership at enqueue after all potentially blocking work. Never send + // an old relay's credentials request on a replacement control connection. + s.reg.mu.Lock() + defer s.reg.mu.Unlock() + row, ok := s.reg.items[lease.row.id] + if !ok || row.boot != lease.row.boot || row.hub != hub || row.placementGen != lease.row.placementGen || ctx.Err() != nil || s.reconnectControl.Load() != lease.control || lease.control.state.generation.Load() != lease.generation { + return + } + if msg.Type == "event" { + msg.Generation = lease.generation + msg.PlacementGeneration = lease.row.placementGen + } + var forwarded uint64 + if msg.Type == "session_req" && msg.RPC != nil && msg.RPC.Method != "resp" { + target := guestRPCTarget{ctx: lease.control.ctx, row: *row, control: lease.control, generation: lease.generation} + var ok bool + forwarded, ok = s.guestForwards.begin(target, msg.RPC.ID) + if !ok { + return + } + rpc := *msg.RPC + rpc.ID = forwarded + msg.RPC = &rpc + } + select { + case lease.control.out <- msg: + default: + if forwarded != 0 { + s.guestForwards.take(forwarded, lease.row.id) + } + } +} diff --git a/internal/runnerd/reconnect_handoff_test.go b/internal/runnerd/reconnect_handoff_test.go new file mode 100644 index 0000000..ce1dc05 --- /dev/null +++ b/internal/runnerd/reconnect_handoff_test.go @@ -0,0 +1,213 @@ +package runnerd + +import ( + "context" + "encoding/json" + "github.com/tokencanopy/rainier/internal/relay" + "github.com/tokencanopy/rainier/protocol/runner" + "net" + "strings" + "testing" + "time" +) + +func TestReconnectHandoffFencesBeforePublication(t *testing.T) { + s, _, _ := reconnectPeer(t) + old := relay.NewHub(context.Background(), &handoffIdleConn{done: make(chan struct{})}) + defer old.Close() + s.reg.setHub("session_test", old) + lease, err := s.acquireGuestReconnect("session_test") + if err != nil { + t.Fatal(err) + } + defer lease.close() + if err = lease.fence(context.Background(), 2); err != nil { + t.Fatal(err) + } + select { + case <-old.Done(): + default: + t.Fatal("old relay remained usable") + } + if _, ok := s.reg.hub("session_test"); ok { + t.Fatal("old relay still published") + } + if err = lease.fence(context.Background(), 2); err != errReconnectFenced { + t.Fatal("same epoch accepted twice") + } +} + +type handoffIdleConn struct{ done chan struct{} } + +func (c *handoffIdleConn) Read(ctx context.Context) ([]byte, error) { + select { + case <-ctx.Done(): + return nil, ctx.Err() + case <-c.done: + return nil, context.Canceled + } +} +func (c *handoffIdleConn) Write(ctx context.Context, _ []byte) error { return ctx.Err() } +func (c *handoffIdleConn) Close() error { return nil } + +func TestReconnectGuestDeliversConfigurationBeforeReadyPublication(t *testing.T) { + s, _, control := reconnectPeer(t) + a, b := net.Pipe() + defer a.Close() + defer b.Close() + host, guest := relay.NetConn(a), relay.NetConn(b) + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + done := make(chan error, 1) + go func() { done <- s.ReconnectGuest(ctx, "session_test", host) }() + req := control.readMsg(t) + replyReconnect(t, control, req, true, reconnectChallengeJSON()) + event, err := relay.ReadGuestReconnectFrame(ctx, guest) + if err != nil || event.Kind != relay.KindGuestReconnectChallenge { + t.Fatal("challenge missing") + } + body, _ := json.Marshal(runner.GuestReconnectAcceptRequest{Protocol: 1, AttemptID: "attempt_test", Signature: strings.Repeat("A", 86)}) + if err = relay.WriteGuestReconnectFrame(ctx, guest, relay.KindGuestReconnectProof, body); err != nil { + t.Fatal(err) + } + req = control.readMsg(t) + replyReconnect(t, control, req, true, []byte(`{"epoch":2,"token":"AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA","expires_in_sec":120}`)) + req = control.readMsg(t) + if req.RPC.Method != runner.MethodGuestReconnectConfiguration { + t.Fatal("no fresh configuration request") + } + body, _ = json.Marshal(runner.GuestReconnectConfiguration{Protocol: 1, SessionID: "session_test", PlacementGeneration: 3, Spec: &runner.Spec{Env: map[string]string{"CONFIG_TEST": "fresh"}}}) + replyReconnect(t, control, req, true, body) + event, err = relay.ReadGuestReconnectFrame(ctx, guest) + if err != nil || event.Kind != relay.KindGuestReconnectAccepted { + t.Fatal("acceptance missing") + } + raw, err := guest.Read(ctx) + if err != nil { + t.Fatal(err) + } + frame, _ := relay.Decode(raw) + if json.Unmarshal(frame.Payload, &event) != nil || event.Kind != relay.KindBootConfig { + t.Fatal("config missing") + } + var cfg runner.BootConfig + if json.Unmarshal(event.Payload, &cfg) != nil || cfg.Env["CONFIG_TEST"] != "fresh" || cfg.BootstrapToken != strings.Repeat("A", 43) { + t.Fatal("wrong configuration") + } + body, _ = json.Marshal(relay.ControlEvent{Kind: "req:" + runner.MethodFetchSessionSecrets, ID: 9, Payload: json.RawMessage(`{"protocol":1,"token":"AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA"}`)}) + raw, _ = relay.Encode(relay.Frame{Type: relay.FrameControl, Payload: body}) + if err = guest.Write(ctx, raw); err != nil { + t.Fatal(err) + } + req = control.readMsg(t) + if req.RPC.Method != runner.MethodFetchSessionSecrets || !isRunnerOriginated(req.RPC.ID) { + t.Fatal("redemption not scoped to host request") + } + replyReconnect(t, control, req, true, []byte(`{"env":{}}`)) + raw, err = guest.Read(ctx) + if err != nil { + t.Fatal(err) + } + frame, _ = relay.Decode(raw) + json.Unmarshal(frame.Payload, &event) + if event.ID != 9 || event.Kind != "resp" || !event.OK { + t.Fatal("redemption correlation lost") + } + if _, ok := s.reg.hub("session_test"); ok { + t.Fatal("relay published before ready") + } + body = []byte(`{"protocol":1,"epoch":2}`) + if err = relay.WriteGuestReconnectFrame(ctx, guest, relay.KindGuestReconnectReady, body); err != nil { + t.Fatal(err) + } + event, err = relay.ReadGuestReconnectFrame(ctx, guest) + if err != nil || event.Kind != relay.KindGuestReconnectReadyAck { + t.Fatal("ack missing") + } + if err = <-done; err != nil { + t.Fatal(err) + } + hub, ok := s.reg.hub("session_test") + if !ok { + t.Fatal("ready relay not published") + } + defer hub.Close() +} + +func TestReconnectGuestConfigurationFailureNeverPublishes(t *testing.T) { + for _, fault := range []string{"scope", "placement", "token", "refusal", "deleted", "canceled"} { + t.Run(fault, func(t *testing.T) { + s, _, control := reconnectPeer(t) + a, b := net.Pipe() + defer a.Close() + defer b.Close() + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + guest := relay.NetConn(b) + done := make(chan error, 1) + go func() { done <- s.ReconnectGuest(ctx, "session_test", relay.NetConn(a)) }() + replyReconnect(t, control, control.readMsg(t), true, reconnectChallengeJSON()) + if _, err := relay.ReadGuestReconnectFrame(ctx, guest); err != nil { + t.Fatal(err) + } + proof, _ := json.Marshal(runner.GuestReconnectAcceptRequest{Protocol: 1, AttemptID: "attempt_test", Signature: strings.Repeat("A", 86)}) + if err := relay.WriteGuestReconnectFrame(ctx, guest, relay.KindGuestReconnectProof, proof); err != nil { + t.Fatal(err) + } + replyReconnect(t, control, control.readMsg(t), true, []byte(`{"epoch":2,"token":"AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA","expires_in_sec":120}`)) + req := control.readMsg(t) + cfg := runner.GuestReconnectConfiguration{Protocol: 1, SessionID: "session_test", PlacementGeneration: 3, Spec: &runner.Spec{Env: map[string]string{"CONFIG_TEST": "must_not_deliver"}}} + switch fault { + case "scope": + cfg.SessionID = "other_test" + case "placement": + cfg.PlacementGeneration++ + case "token": + cfg.Spec.BootstrapToken = "replacement_test" + case "deleted": + s.reg.remove("session_test") + case "canceled": + cancel() + } + if fault == "refusal" { + replyReconnect(t, control, req, false, []byte(`{"error":"fenced"}`)) + } else { + body, _ := json.Marshal(cfg) + replyReconnect(t, control, req, true, body) + } + select { + case err := <-done: + if err == nil { + t.Fatal("failed handoff succeeded") + } + case <-time.After(3 * time.Second): + t.Fatal("handoff stuck") + } + if _, ok := s.reg.hub("session_test"); ok { + t.Fatal("failed handoff published a relay") + } + readCtx, stop := context.WithTimeout(context.Background(), time.Second) + defer stop() + if raw, err := guest.Read(readCtx); err == nil || len(raw) != 0 { + t.Fatal("failed resolution disclosed configuration") + } + }) + } +} + +func TestReconnectRelayFenceRejectsQueuedCallbacks(t *testing.T) { + authority := &guestRelayAuthority{} + admitted := make(chan struct{}) + release := make(chan struct{}) + finished := make(chan struct{}) + go authority.run(func() { close(admitted); <-release }) + <-admitted + go func() { authority.fence(); close(finished) }() + close(release) + <-finished + ran := false + authority.run(func() { ran = true }) + if ran { + t.Fatal("fenced callback ran") + } +} diff --git a/internal/runnerd/reconnect_lease.go b/internal/runnerd/reconnect_lease.go new file mode 100644 index 0000000..ed37d44 --- /dev/null +++ b/internal/runnerd/reconnect_lease.go @@ -0,0 +1,121 @@ +package runnerd + +import ( + "context" + "sync" + "sync/atomic" + + "github.com/tokencanopy/rainier/internal/relay" +) + +// guestReconnectLease retains admission and original local/control ownership +// through proof, current configuration delivery, and guarded relay publication. +// It is process-local authority only; durable proof authorization is still +// required. Never reconstruct one from a guest payload or acceptance alone. +type guestReconnectLease struct { + server *Server + row sessionEntry + control *reconnectControl + generation uint64 + once sync.Once + closed atomic.Bool +} + +func (s *Server) acquireGuestReconnect(id string) (*guestReconnectLease, error) { + row, ok := s.reg.snapshot(id) + if !ok || row.placementGen == 0 || (row.state != "running" && row.state != "starting") { + return nil, errReconnectFenced + } + rc := s.reconnectControl.Load() + if rc == nil || rc.ctx.Err() != nil || rc.state.generation.Load() == 0 { + return nil, errReconnectUnavailable + } + if !s.claimReconnect(id) { + return nil, errReconnectUnavailable + } + lease := &guestReconnectLease{server: s, row: row, control: rc, generation: rc.state.generation.Load()} + if !lease.valid(context.Background()) { + lease.close() + return nil, errReconnectFenced + } + return lease, nil +} + +func (lease *guestReconnectLease) valid(ctx context.Context) bool { + current, exists := lease.server.reg.snapshot(lease.row.id) + return exists && lease.matches(ctx, current) +} + +func (lease *guestReconnectLease) close() { + lease.once.Do(func() { lease.closed.Store(true); lease.server.releaseReconnect(lease.row.id) }) +} + +// fence consumes the observed epoch and removes the old relay before any new +// configuration or traffic is delivered. Previously started callbacks finish +// before fencing completes; queued callbacks cannot regain authority afterwards. +func (lease *guestReconnectLease) fence(ctx context.Context, epoch uint64) error { + s := lease.server + s.reg.mu.Lock() + row, ok := s.reg.items[lease.row.id] + if !ok || !lease.matches(ctx, *row) || epoch == 0 || epoch <= row.guestEpoch { + s.reg.mu.Unlock() + return errReconnectFenced + } + old, authority := row.hub, row.relayAuthority + row.guestEpoch = epoch + row.hub = nil + row.relayAuthority = nil + s.reg.mu.Unlock() + if old != nil { + s.guestForwards.discard(old) + old.Close() + } + if authority != nil { + authority.fence() + } + if !lease.valid(ctx) { + return errReconnectFenced + } + return nil +} + +// guestRelayAuthority drains already-admitted callbacks and rejects queued ones. +// It owns no network or registry lock, so closing a relay cannot deadlock a +// callback that must inspect the registry before it enqueues a request. +type guestRelayAuthority struct { + mu sync.RWMutex + closed bool +} + +func (a *guestRelayAuthority) run(f func()) { + a.mu.RLock() + defer a.mu.RUnlock() + if !a.closed { + f() + } +} +func (a *guestRelayAuthority) fence() { + a.mu.Lock() + a.closed = true + a.mu.Unlock() +} + +func (lease *guestReconnectLease) matches(ctx context.Context, current sessionEntry) bool { + return !lease.closed.Load() && ctx.Err() == nil && lease.control.ctx.Err() == nil && + lease.server.reconnectControl.Load() == lease.control && lease.control.state.generation.Load() == lease.generation && + current.placementGen == lease.row.placementGen && current.boot == lease.row.boot && current.handle == lease.row.handle && + (current.hub == nil || current.hub == lease.row.hub) && (current.state == "running" || current.state == "starting") +} + +func (lease *guestReconnectLease) publish(ctx context.Context, epoch uint64, hub *relay.Hub, authority *guestRelayAuthority) bool { + s := lease.server + s.reg.mu.Lock() + defer s.reg.mu.Unlock() + row, ok := s.reg.items[lease.row.id] + if !ok || !lease.matches(ctx, *row) || row.guestEpoch != epoch || row.hub != nil { + return false + } + row.hub = hub + row.relayAuthority = authority + return true +} diff --git a/internal/runnerd/reconnect_lease_test.go b/internal/runnerd/reconnect_lease_test.go new file mode 100644 index 0000000..d789883 --- /dev/null +++ b/internal/runnerd/reconnect_lease_test.go @@ -0,0 +1,73 @@ +package runnerd + +import ( + "context" + "testing" +) + +func TestReconnectLeaseRetainsAdmissionThroughDelivery(t *testing.T) { + s, _, _ := reconnectPeer(t) + lease, err := s.acquireGuestReconnect("session_test") + if err != nil { + t.Fatal(err) + } + defer lease.close() + if !lease.valid(context.Background()) { + t.Fatal("new lease is invalid") + } + if _, err := s.acquireGuestReconnect("session_test"); err != errReconnectUnavailable { + t.Fatalf("parallel delivery admitted: %v", err) + } + lease.close() + next, err := s.acquireGuestReconnect("session_test") + if err != nil { + t.Fatal(err) + } + defer next.close() + lease.close() + if _, err := s.acquireGuestReconnect("session_test"); err != errReconnectUnavailable { + t.Fatalf("old close released new lease: %v", err) + } + if lease.valid(context.Background()) { + t.Fatal("closed lease retained authority") + } +} + +func TestReconnectLeaseRejectsOwnershipChanges(t *testing.T) { + for _, change := range []string{"handle", "boot", "placement", "state", "deleted", "control", "generation", "caller"} { + t.Run(change, func(t *testing.T) { + s, _, _ := reconnectPeer(t) + lease, err := s.acquireGuestReconnect("session_test") + if err != nil { + t.Fatal(err) + } + defer lease.close() + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + s.reg.mu.Lock() + row := s.reg.items["session_test"] + switch change { + case "handle": + row.handle = "replacement_test" + case "boot": + row.boot++ + case "placement": + row.placementGen++ + case "state": + row.state = "suspending" + case "deleted": + delete(s.reg.items, "session_test") + case "control": + s.reconnectControl.Store(nil) + case "generation": + lease.control.state.generation.Add(1) + case "caller": + cancel() + } + s.reg.mu.Unlock() + if lease.valid(ctx) { + t.Fatal("changed ownership retained authority") + } + }) + } +} diff --git a/internal/runnerd/reconnect_preamble.go b/internal/runnerd/reconnect_preamble.go new file mode 100644 index 0000000..93d0f25 --- /dev/null +++ b/internal/runnerd/reconnect_preamble.go @@ -0,0 +1,73 @@ +package runnerd + +import ( + "bytes" + "encoding/json" + "io" + + "github.com/tokencanopy/rainier/internal/relay" + "github.com/tokencanopy/rainier/protocol/runner" +) + +// Exact spellings and one occurrence per field, including the outer envelope. +// Typed JSON decoding alone accepts case aliases and duplicate last values. +func guestPreambleObject(raw []byte, keys ...string) (map[string]json.RawMessage, error) { + if len(raw) > runner.GuestReconnectPayloadLimit { + return nil, errReconnectInvalid + } + decoder := json.NewDecoder(bytes.NewReader(raw)) + token, err := decoder.Token() + if err != nil || token != json.Delim('{') { + return nil, errReconnectInvalid + } + fields := make(map[string]json.RawMessage, len(keys)) + allowed := make(map[string]bool, len(keys)) + for _, key := range keys { + allowed[key] = true + } + for decoder.More() { + token, err := decoder.Token() + if err != nil { + return nil, errReconnectInvalid + } + key, ok := token.(string) + if !ok || !allowed[key] || fields[key] != nil { + return nil, errReconnectInvalid + } + var value json.RawMessage + if decoder.Decode(&value) != nil || bytes.Equal(bytes.TrimSpace(value), []byte("null")) { + return nil, errReconnectInvalid + } + fields[key] = value + } + if token, err = decoder.Token(); err != nil || token != json.Delim('}') || len(fields) != len(keys) { + return nil, errReconnectInvalid + } + if _, err = decoder.Token(); err != io.EOF { + return nil, errReconnectInvalid + } + return fields, nil +} +func decodeGuestBootstrapEvent(raw []byte, method string) (relay.ControlEvent, error) { + fields, err := guestPreambleObject(raw, "kind", "id", "payload") + var event relay.ControlEvent + if err != nil || json.Unmarshal(fields["kind"], &event.Kind) != nil || json.Unmarshal(fields["id"], &event.ID) != nil || event.Kind != "req:"+method || event.ID == 0 || isRunnerOriginated(event.ID) { + return relay.ControlEvent{}, errReconnectInvalid + } + event.Payload = fields["payload"] + return event, nil +} + +type guestRedemption struct { + Protocol uint64 `json:"protocol"` + Token string `json:"token"` +} + +func decodeGuestRedemption(raw []byte) (guestRedemption, error) { + fields, err := guestPreambleObject(raw, "protocol", "token") + var req guestRedemption + if err != nil || json.Unmarshal(fields["protocol"], &req.Protocol) != nil || json.Unmarshal(fields["token"], &req.Token) != nil || req.Protocol != runner.GuestReconnectProtocol || req.Token == "" { + return guestRedemption{}, errReconnectInvalid + } + return req, nil +} diff --git a/internal/runnerd/reconnect_preamble_test.go b/internal/runnerd/reconnect_preamble_test.go new file mode 100644 index 0000000..863b8ab --- /dev/null +++ b/internal/runnerd/reconnect_preamble_test.go @@ -0,0 +1,37 @@ +package runnerd + +import ( + "github.com/tokencanopy/rainier/protocol/runner" + "testing" +) + +func TestGuestBootstrapPreambleRejectsAmbiguousFields(t *testing.T) { + for _, raw := range []string{ + `{"kind":"req:fetch_session_secrets","id":1,"id":2,"payload":{"protocol":1,"token":"test"}}`, + `{"Kind":"req:fetch_session_secrets","id":1,"payload":{}}`, + `{"kind":"req:fetch_session_secrets","id":1,"extra":true,"payload":{}}`, + `{"kind":"req:fetch_session_secrets","id":1,"payload":{}} {}`, + } { + if _, err := decodeGuestBootstrapEvent([]byte(raw), runner.MethodFetchSessionSecrets); err == nil { + t.Fatal("accepted ambiguous envelope") + } + } + for _, raw := range []string{ + `{"protocol":1,"token":"first","token":"second"}`, + `{"protocol":1,"Token":"test"}`, + `{"protocol":1,"token":"test","session":"another_test"}`, + `{"protocol":1,"token":null}`, + } { + if _, err := decodeGuestRedemption([]byte(raw)); err == nil { + t.Fatal("accepted ambiguous redemption") + } + } + event, err := decodeGuestBootstrapEvent([]byte(`{"kind":"req:fetch_session_secrets","id":9,"payload":{"protocol":1,"token":"test"}}`), runner.MethodFetchSessionSecrets) + if err != nil || event.ID != 9 { + t.Fatal("valid envelope rejected") + } + req, err := decodeGuestRedemption(event.Payload) + if err != nil || req.Token != "test" { + t.Fatal("valid redemption rejected") + } +} diff --git a/internal/runnerd/reconnect_recovery.go b/internal/runnerd/reconnect_recovery.go new file mode 100644 index 0000000..5efa6ee --- /dev/null +++ b/internal/runnerd/reconnect_recovery.go @@ -0,0 +1,82 @@ +package runnerd + +import ( + "context" + "time" + + "github.com/tokencanopy/rainier/protocol/runner" +) + +func (s *Server) authorizeRecoveredGuest(ctx context.Context, rc *reconnectControl, row sessionEntry) error { + if rc == nil || !row.recovered || !row.guestReconnect || row.handle == "" || row.state != "running" || row.hub != nil { + return errReconnectFenced + } + generation := rc.state.generation.Load() + if generation == 0 || rc.ctx.Err() != nil || s.reconnectControl.Load() != rc { + return errReconnectFenced + } + ctx, cancel := context.WithTimeout(ctx, reconnectBudget) + defer cancel() + stop := context.AfterFunc(rc.ctx, cancel) + defer stop() + payload, err := s.reconnectCall(ctx, rc, row.id, runner.MethodGuestReconnectConfiguration, runner.GuestReconnectBeginRequest{Protocol: 1}) + if err != nil { + return err + } + fresh, err := runner.DecodeGuestReconnectConfiguration(payload) + if err != nil || fresh.SessionID != row.id || row.placementGen != 0 && fresh.PlacementGeneration != row.placementGen { + return errReconnectFenced + } + s.reg.mu.Lock() + defer s.reg.mu.Unlock() + current, ok := s.reg.items[row.id] + if !ok || ctx.Err() != nil || rc.ctx.Err() != nil || s.reconnectControl.Load() != rc || rc.state.generation.Load() != generation || current.boot != row.boot || current.handle != row.handle || current.placementGen != row.placementGen || !current.recovered || !current.guestReconnect || current.state != "running" || current.hub != nil { + return errReconnectFenced + } + current.placementGen = fresh.PlacementGeneration + return nil +} + +// One sequential worker per accepted connection bounds recovery independently +// of guest traffic. Failed candidates retry under the same current authority; +// cancellation discards that worker and the replacement registration starts anew. +func (s *Server) recoverGuests(rc *reconnectControl) { + drv, ok := s.drv.(interface { + RecoverGuest(context.Context, string, string) error + }) + if !ok { + return + } + recovered := map[string]string{} + for rc.ctx.Err() == nil && s.reconnectControl.Load() == rc { + s.reg.mu.Lock() + var rows []sessionEntry + for _, row := range s.reg.items { + if row.recovered && row.guestReconnect && row.state == "running" && row.hub == nil && recovered[row.id] != row.handle { + rows = append(rows, *row) + } + } + s.reg.mu.Unlock() + for _, row := range rows { + if rc.ctx.Err() != nil || s.reconnectControl.Load() != rc { + return + } + ctx, cancel := context.WithTimeout(rc.ctx, 2*reconnectBudget) + err := s.authorizeRecoveredGuest(ctx, rc, row) + if err == nil { + err = drv.RecoverGuest(ctx, row.handle, row.id) + } + cancel() + if err == nil { + recovered[row.id] = row.handle + } + } + timer := time.NewTimer(5 * time.Second) + select { + case <-rc.ctx.Done(): + timer.Stop() + return + case <-timer.C: + } + } +} diff --git a/internal/runnerd/reconnect_recovery_test.go b/internal/runnerd/reconnect_recovery_test.go new file mode 100644 index 0000000..60438b8 --- /dev/null +++ b/internal/runnerd/reconnect_recovery_test.go @@ -0,0 +1,53 @@ +package runnerd + +import ( + "context" + "encoding/json" + "testing" + + "github.com/tokencanopy/rainier/protocol/runner" +) + +func TestRecoveredGuestPlacementRequiresCurrentAuthorization(t *testing.T) { + for _, fault := range []string{"none", "scope", "replaced", "control", "refused"} { + t.Run(fault, func(t *testing.T) { + s, _, control := reconnectPeer(t) + s.reg.mu.Lock() + entry := s.reg.items["session_test"] + entry.placementGen = 0 + entry.recovered = true + entry.guestReconnect = true + row := *entry + s.reg.mu.Unlock() + rc := s.reconnectControl.Load() + done := make(chan error, 1) + go func() { done <- s.authorizeRecoveredGuest(context.Background(), rc, row) }() + req := control.readMsg(t) + if req.RPC.Method != runner.MethodGuestReconnectConfiguration { + t.Fatal("recovery did not reauthorize current placement") + } + cfg := runner.GuestReconnectConfiguration{Protocol: 1, SessionID: "session_test", PlacementGeneration: 3, Spec: &runner.Spec{}} + switch fault { + case "scope": + cfg.SessionID = "another_test" + case "replaced": + s.reg.mu.Lock() + s.reg.items[row.id].handle = "replacement_vm" + s.reg.mu.Unlock() + case "control": + rc.state.generation.Add(1) + } + payload, _ := json.Marshal(cfg) + replyReconnect(t, control, req, fault != "refused", payload) + err := <-done + after, _ := s.reg.snapshot(row.id) + if fault == "none" { + if err != nil || after.placementGen != 3 { + t.Fatalf("current authorization not installed: %v", err) + } + } else if err == nil || after.placementGen != 0 { + t.Fatal("stale authorization installed") + } + }) + } +} diff --git a/internal/runnerd/reconnect_rpc_fencing.go b/internal/runnerd/reconnect_rpc_fencing.go new file mode 100644 index 0000000..ead42ac --- /dev/null +++ b/internal/runnerd/reconnect_rpc_fencing.go @@ -0,0 +1,106 @@ +package runnerd + +import ( + "context" + "encoding/json" + "sync" + "time" + + "github.com/tokencanopy/rainier/internal/relay" + "github.com/tokencanopy/rainier/protocol/runner" +) + +// Captured on the control reader, before a command can queue behind a handoff. +type guestRPCTarget struct { + ctx context.Context + row sessionEntry + control *reconnectControl + generation uint64 +} + +func (s *Server) guestRPCTarget(ctx context.Context, id string, ag *agentSessionState) *guestRPCTarget { + row, ok := s.reg.snapshot(id) + if !ok || !row.guestReconnect { + return nil + } + rc := s.reconnectControl.Load() + target := &guestRPCTarget{ctx: ctx, row: row, control: rc} + if rc != nil && (ag == nil || rc.state == ag) { + target.generation = rc.state.generation.Load() + } + return target +} +func (s *Server) sendGuestRPC(target *guestRPCTarget, env runner.RPCEnvelope) error { + if target == nil || target.control == nil || target.generation == 0 || target.row.hub == nil { + return errReconnectFenced + } + row, ok := s.reg.snapshot(target.row.id) + if !ok || target.ctx.Err() != nil || target.control.ctx.Err() != nil || s.reconnectControl.Load() != target.control || target.control.state.generation.Load() != target.generation || row.hub != target.row.hub || row.boot != target.row.boot || row.handle != target.row.handle || row.placementGen != target.row.placementGen { + return errReconnectFenced + } + event := relay.ControlEvent{Kind: "req:" + env.Method, ID: env.ID, Payload: env.Payload} + if env.Method == "resp" { + event.Kind = "resp" + event.OK = env.OK + } + body, err := json.Marshal(event) + if err != nil { + return errReconnectInvalid + } + // Always the captured hub. A concurrent fence closes this hub; it cannot + // redirect this write to a newly published guest connection. + return target.row.hub.SendControl(body) +} + +// Guest muxes may restart their low ID sequence at each reconnect. Translate +// requests to a runner-lifetime sequence so an old reply cannot alias a new +// guest request, even when both guests used ID 1. +type guestRPCPending struct { + target guestRPCTarget + guestID uint64 + expires time.Time +} +type guestRPCForwards struct { + mu sync.Mutex + seq uint64 + pending map[uint64]guestRPCPending +} + +func (f *guestRPCForwards) begin(target guestRPCTarget, id uint64) (uint64, bool) { + f.mu.Lock() + defer f.mu.Unlock() + now := time.Now() + for key, entry := range f.pending { + if !now.Before(entry.expires) || entry.target.control.ctx.Err() != nil { + delete(f.pending, key) + } + } + if id == 0 || isRunnerOriginated(id) || len(f.pending) >= 1024 || f.seq >= runnerOriginatedIDBase-1 { + return 0, false + } + if f.pending == nil { + f.pending = make(map[uint64]guestRPCPending) + } + f.seq++ + f.pending[f.seq] = guestRPCPending{target: target, guestID: id, expires: now.Add(2 * time.Minute)} + return f.seq, true +} +func (f *guestRPCForwards) take(id uint64, session string) (guestRPCPending, bool) { + f.mu.Lock() + defer f.mu.Unlock() + entry, ok := f.pending[id] + if !ok || entry.target.row.id != session { + return guestRPCPending{}, false + } + delete(f.pending, id) + return entry, time.Now().Before(entry.expires) +} +func (f *guestRPCForwards) discard(hub *relay.Hub) { + f.mu.Lock() + defer f.mu.Unlock() + for id, entry := range f.pending { + if entry.target.row.hub == hub { + delete(f.pending, id) + } + } +} diff --git a/internal/runnerd/reconnect_rpc_fencing_test.go b/internal/runnerd/reconnect_rpc_fencing_test.go new file mode 100644 index 0000000..7fafe45 --- /dev/null +++ b/internal/runnerd/reconnect_rpc_fencing_test.go @@ -0,0 +1,159 @@ +package runnerd + +import ( + "context" + "encoding/json" + "github.com/tokencanopy/rainier/internal/relay" + "github.com/tokencanopy/rainier/protocol/runner" + "testing" + "time" +) + +type reconnectRPCCapture struct{ writes chan []byte } + +func (c *reconnectRPCCapture) Read(ctx context.Context) ([]byte, error) { + <-ctx.Done() + return nil, ctx.Err() +} +func (c *reconnectRPCCapture) Write(ctx context.Context, b []byte) error { c.writes <- b; return nil } +func (c *reconnectRPCCapture) Close() error { return nil } +func TestReconnectRejectsCanceledControlCommand(t *testing.T) { + s, _ := testMicrovmServer(t) + s.reg.put("session_test", &sessionEntry{id: "session_test", handle: "vm_test", state: "running", placementGen: 3, guestReconnect: true}) + c := &reconnectRPCCapture{make(chan []byte, 1)} + h := relay.NewHub(context.Background(), c) + defer h.Close() + s.reg.setHub("session_test", h) + ctx, cancel := context.WithCancel(context.Background()) + cancel() + s.execute(ctx, runner.ToRunner{Type: "session_rpc", Session: "session_test", RPC: &runner.RPCEnvelope{ID: 17, Method: "exec", Payload: []byte(`{"command":"synthetic"}`)}}, func(runner.FromRunner) {}, AgentConfig{}, &agentSessionState{}) + select { + case <-c.writes: + t.Fatal("canceled superseded control command reached current guest hub") + case <-time.After(50 * time.Millisecond): + } +} +func TestReconnectRejectsUncorrelatedResponse(t *testing.T) { + s, _ := testMicrovmServer(t) + s.reg.put("session_test", &sessionEntry{id: "session_test", handle: "vm_test", state: "running", placementGen: 3, guestReconnect: true}) + c := &reconnectRPCCapture{make(chan []byte, 1)} + h := relay.NewHub(context.Background(), c) + defer h.Close() + s.reg.setHub("session_test", h) + s.forwardSessionRPC(runner.ToRunner{Type: "session_rpc", Session: "session_test", RPC: &runner.RPCEnvelope{ID: 17, Method: "resp", OK: true, Payload: []byte(`{"token":"synthetic"}`)}}, func(runner.FromRunner) {}) + select { + case <-c.writes: + t.Fatal("uncorrelated old response reached replacement guest hub") + case <-time.After(50 * time.Millisecond): + } +} + +func TestReconnectRPCReplyKeepsOriginalHubAndGuestID(t *testing.T) { + s, _, control := reconnectPeer(t) + oldConn := &reconnectRPCCapture{writes: make(chan []byte, 2)} + oldHub := relay.NewHub(context.Background(), oldConn) + defer oldHub.Close() + s.reg.mu.Lock() + s.reg.items["session_test"].guestReconnect = true + s.reg.mu.Unlock() + s.reg.setHub("session_test", oldHub) + lease, err := s.acquireGuestReconnect("session_test") + if err != nil { + t.Fatal(err) + } + defer lease.close() + sendRequest := func(hub *relay.Hub) runner.FromRunner { + s.sendOwnedGuestMessage(lease, hub, runner.FromRunner{Type: "session_req", Session: "session_test", RPC: &runner.RPCEnvelope{ID: 1, Method: "credential_test"}}) + return control.readMsg(t) + } + first := sendRequest(oldHub) + // A same-owner answer retains the original guest ID. + s.forwardSessionRPC(runner.ToRunner{Session: "session_test", RPC: &runner.RPCEnvelope{ID: first.RPC.ID, Method: "resp", OK: true}}, func(runner.FromRunner) {}) + select { + case raw := <-oldConn.writes: + frame, _ := relay.Decode(raw) + var event relay.ControlEvent + if json.Unmarshal(frame.Payload, &event) != nil || event.ID != 1 || event.Kind != "resp" { + t.Fatal("guest correlation lost") + } + case <-time.After(time.Second): + t.Fatal("current reply dropped") + } + late := sendRequest(oldHub) + nextConn := &reconnectRPCCapture{writes: make(chan []byte, 2)} + nextHub := relay.NewHub(context.Background(), nextConn) + defer nextHub.Close() + s.reg.setHub("session_test", nextHub) + current := sendRequest(nextHub) + if current.RPC.ID == late.RPC.ID { + t.Fatal("reused forward ID across guest muxes") + } + s.forwardSessionRPC(runner.ToRunner{Session: "session_test", RPC: &runner.RPCEnvelope{ID: late.RPC.ID, Method: "resp", OK: true}}, func(runner.FromRunner) {}) + select { + case <-nextConn.writes: + t.Fatal("old reply reached replacement") + case <-time.After(50 * time.Millisecond): + } + s.forwardSessionRPC(runner.ToRunner{Session: "session_test", RPC: &runner.RPCEnvelope{ID: current.RPC.ID, Method: "resp", OK: true}}, func(runner.FromRunner) {}) + select { + case <-nextConn.writes: + case <-time.After(time.Second): + t.Fatal("new mux reply dropped") + } +} + +func TestReconnectCommandCannotAdoptHubAfterReceipt(t *testing.T) { + s, _, _ := reconnectPeer(t) + s.reg.mu.Lock() + s.reg.items["session_test"].guestReconnect = true + s.reg.mu.Unlock() + rc := s.reconnectControl.Load() + target := s.guestRPCTarget(rc.ctx, "session_test", rc.state) + conn := &reconnectRPCCapture{writes: make(chan []byte, 1)} + hub := relay.NewHub(context.Background(), conn) + defer hub.Close() + s.reg.setHub("session_test", hub) + s.execute(rc.ctx, runner.ToRunner{Type: "session_rpc", Session: "session_test", RPC: &runner.RPCEnvelope{ID: 17, Method: "exec", Payload: []byte(`{"command":"synthetic"}`)}}, func(runner.FromRunner) {}, AgentConfig{}, rc.state, target) + select { + case <-conn.writes: + t.Fatal("queued command adopted later hub") + case <-time.After(50 * time.Millisecond): + } +} + +func TestReconnectForwardedRequestsAreBoundedAndConnectionScoped(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + target := guestRPCTarget{ctx: ctx, row: sessionEntry{id: "session_test"}, control: &reconnectControl{ctx: ctx}} + var forwards guestRPCForwards + var first uint64 + for i := 0; i < 1024; i++ { + id, ok := forwards.begin(target, 1) + if !ok { + t.Fatal("refused within bounded capacity") + } + if i == 0 { + first = id + } + } + if _, ok := forwards.begin(target, 1); ok { + t.Fatal("unbounded forwarding table") + } + if _, ok := forwards.take(first, "another_test"); ok { + t.Fatal("cross-session correlation") + } + if _, ok := forwards.take(first, "session_test"); !ok { + t.Fatal("wrong-session response consumed real waiter") + } + if _, ok := forwards.take(first, "session_test"); ok { + t.Fatal("response replay accepted") + } + cancel() + replacement := context.Background() + target.ctx = replacement + target.control = &reconnectControl{ctx: replacement} + id, ok := forwards.begin(target, 1) + if !ok || id <= 1024 { + t.Fatal("connection loss leaked capacity or reused correlation IDs") + } +} diff --git a/internal/runnerd/reconnect_test.go b/internal/runnerd/reconnect_test.go index a6dffd9..32d6516 100644 --- a/internal/runnerd/reconnect_test.go +++ b/internal/runnerd/reconnect_test.go @@ -48,6 +48,9 @@ func TestReconnectHostRoundTrip(t *testing.T) { done <- err }() req := conn.readMsg(t) + if req.Total != 4 { + t.Fatalf("host-originated reconnect RPC lost capacity snapshot: total=%d", req.Total) + } if req.Type != "session_req" || req.Session != "session_test" || req.RPC == nil || req.RPC.Method != runner.MethodBeginGuestReconnect || !isRunnerOriginated(req.RPC.ID) { t.Fatal("wrong begin request") } diff --git a/internal/runnerd/registry.go b/internal/runnerd/registry.go index b687de3..1c67dad 100644 --- a/internal/runnerd/registry.go +++ b/internal/runnerd/registry.go @@ -28,8 +28,12 @@ type sessionEntry struct { // echoes it, which is what lets controld fence a report from a sandbox // the session has since been re-placed away from. Zero is "the create // carried none" — an old controld — and fences nothing. - placementGen uint64 - hub *relay.Hub // set when sessiond registers; nil until then + placementGen uint64 + guestEpoch uint64 + resumePending bool + guestReconnect bool + relayAuthority *guestRelayAuthority + hub *relay.Hub // set when sessiond registers; nil until then // attachments is the number of attachments currently open over this // session's hub — viewers and controllers alike, since a runner neither // knows nor needs to know which of them holds the controller lease. @@ -352,6 +356,10 @@ func (r *registry) list() []sessionEntry { // registry is not where a hub's lifetime is decided and because a displaced // hub is not immediately useless: see retireDisplacedHub. func (r *registry) setHub(id string, h *relay.Hub) (displaced *relay.Hub, ok bool) { + return r.setHubAuthority(id, h, nil) +} + +func (r *registry) setHubAuthority(id string, h *relay.Hub, authority *guestRelayAuthority) (displaced *relay.Hub, ok bool) { r.mu.Lock() defer r.mu.Unlock() e, ok := r.items[id] @@ -360,6 +368,7 @@ func (r *registry) setHub(id string, h *relay.Hub) (displaced *relay.Hub, ok boo } displaced = e.hub e.hub = h + e.relayAuthority = authority if displaced == h { displaced = nil } @@ -957,6 +966,10 @@ func (r *registry) resumed(id string, restarted bool) { if !ok { return } + r.resumedLocked(e, restarted) +} + +func (r *registry) resumedLocked(e *sessionEntry, restarted bool) { // Read before it is overwritten: only the resume that actually brings a // PARKED entry back can be the one that restarted it. `docker start` on an // an already-started container exits 0, so a second resume reports a @@ -1005,6 +1018,10 @@ func (r *registry) resumed(id string, restarted bool) { // the new one. The register that is about to arrive opens the next // epoch; until it does, frames from either side of the restart match // nothing. See sessionEntry.boot. - r.nextBoot++ - e.boot = r.nextBoot + if e.resumePending { + e.resumePending = false // claim already opened this boot before launch + } else { + r.nextBoot++ + e.boot = r.nextBoot + } } diff --git a/internal/runnerd/resume_placement.go b/internal/runnerd/resume_placement.go new file mode 100644 index 0000000..fc11b94 --- /dev/null +++ b/internal/runnerd/resume_placement.go @@ -0,0 +1,131 @@ +package runnerd + +import ( + "context" + "errors" + "time" + + "github.com/tokencanopy/rainier/internal/driver" + "github.com/tokencanopy/rainier/protocol/runner" +) + +var errResumePlacement = errors.New("resume placement is not current") + +// claimResumePlacement opens the local boot before driver launch can expose a +// fresh guest. Only the authenticated control command supplies generation. +func (s *Server) claimResumePlacement(id, handle string, generation uint64) (uint64, error) { + s.reg.mu.Lock() + row, ok := s.reg.items[id] + if !ok || row.handle != handle || generation == 0 || generation <= row.placementGen || row.resumePending || row.state == "running" || row.state == "starting" || row.state == "destroying" || row.stopsInFlight > 0 { + s.reg.mu.Unlock() + return 0, errResumePlacement + } + old, authority := row.hub, row.relayAuthority + row.hub, row.relayAuthority = nil, nil + row.guestEpoch = 0 + row.placementGen = generation + row.resumePending = true + row.state = "resuming" + s.reg.nextBoot++ + row.boot = s.reg.nextBoot + boot := row.boot + s.reg.mu.Unlock() + if old != nil { + s.guestForwards.discard(old) + old.Close() + } + if authority != nil { + authority.fence() + } + return boot, nil +} + +func (s *Server) failResumePlacement(ctx context.Context, id, handle string, generation, boot uint64) { + // Only observed suspension proves a failed boot can be retried. Inspection + // errors or a live VM retain an unresolved state and all resource ownership. + ctx, cancel := context.WithTimeout(ctx, 5*time.Second) + defer cancel() + observed, err := s.drv.Inspect(ctx, handle) + s.reg.mu.Lock() + defer s.reg.mu.Unlock() + row, ok := s.reg.items[id] + if !ok || row.handle != handle || row.placementGen != generation || row.boot != boot { + return + } + row.resumePending = false + if row.state == "resuming" && err == nil && observed.State == driver.StateSuspended { + row.state = "suspended" + } +} + +func (s *Server) resumeStatus(ctx context.Context, m runner.ToRunner) runner.FromRunner { + refused := runner.FromRunner{Type: "result", ReqID: m.ReqID} + row, ok := s.reg.snapshot(m.Session) + if !ok || ctx.Err() != nil || m.PlacementGeneration == 0 || m.PlacementGeneration < row.placementGen { + return refused + } + state := "resuming" + generation := row.placementGen + if !row.resumePending { + if drv, ok := s.drv.(interface { + ResumeStatus(context.Context, string, uint64) (string, uint64, error) + }); ok { + var err error + state, generation, err = drv.ResumeStatus(ctx, row.handle, m.PlacementGeneration) + if err != nil { + return refused + } + } else if generation < m.PlacementGeneration && row.state == "suspended" && row.hub == nil && row.stopsInFlight == 0 { + // Drivers without durable placement metadata can cancel a command + // only while the existing sandbox is observably stopped. The + // registry comparison below serializes this fence with any launch. + observed, err := s.drv.Inspect(ctx, row.handle) + if err != nil || observed.State != driver.StateSuspended { + return refused + } + state, generation = "suspended_cold", m.PlacementGeneration + } else if generation == m.PlacementGeneration { + if row.state == "running" { + state = "running" + } else if row.state == "suspended" && row.hub == nil { + state = "suspended_cold" + } + } + } + if generation != m.PlacementGeneration || ctx.Err() != nil { + return refused + } + s.reg.mu.Lock() + defer s.reg.mu.Unlock() + current, ok := s.reg.items[m.Session] + if !ok || current.boot != row.boot || current.handle != row.handle || current.placementGen != row.placementGen || current.resumePending != row.resumePending || current.state != row.state || current.stopsInFlight != 0 { + return refused + } + if generation > current.placementGen { + if state != "suspended_cold" { + return refused + } + current.placementGen = generation // driver's durable cancellation tombstone + } + if state == "running" && current.guestReconnect && current.hub == nil { + state = "resuming" + } + return runner.FromRunner{Type: "result", ReqID: m.ReqID, OK: true, State: state, PlacementGeneration: generation} +} + +// Only the command that claimed this boot may complete it. Deletion keeps +// state ownership, even when the driver finishes while teardown waits for it. +func (s *Server) completeResumePlacement(id, handle string, generation, boot uint64) error { + s.reg.mu.Lock() + defer s.reg.mu.Unlock() + row, ok := s.reg.items[id] + if !ok || row.handle != handle || row.placementGen != generation || row.boot != boot || !row.resumePending { + return errResumePlacement + } + if row.state != "resuming" { + row.resumePending = false + return errResumePlacement + } + s.reg.resumedLocked(row, true) + return nil +} diff --git a/internal/runnerd/resume_placement_test.go b/internal/runnerd/resume_placement_test.go new file mode 100644 index 0000000..e607f79 --- /dev/null +++ b/internal/runnerd/resume_placement_test.go @@ -0,0 +1,142 @@ +package runnerd + +import ( + "context" + "testing" + + "github.com/tokencanopy/rainier/internal/driver" + "github.com/tokencanopy/rainier/protocol/runner" +) + +type observedResumeDriver struct { + *driver.Fake + before func() +} + +func (d *observedResumeDriver) Resume(ctx context.Context, id string) (bool, error) { + d.before() + return d.Fake.Resume(ctx, id) +} +func TestResumeCommandFencesBootBeforeDriverLaunch(t *testing.T) { + h := newIdleHarness(t) + if err := h.rd.Op(context.Background(), h.id, "suspend", false); err != nil { + t.Fatal(err) + } + h.rd.reg.mu.Lock() + row := h.rd.reg.items[h.id] + row.placementGen = 1 + row.guestEpoch = 9 + row.guestReconnect = true + oldBoot := row.boot + h.rd.reg.mu.Unlock() + var claimedBoot uint64 + h.rd.drv = &observedResumeDriver{Fake: h.fd, before: func() { + current, _ := h.rd.reg.snapshot(h.id) + if current.placementGen != 2 || current.guestEpoch != 0 || current.boot == oldBoot || current.state != "resuming" { + t.Fatal("new guest can launch before fresh boot authority") + } + claimedBoot = current.boot + }} + var answer runner.FromRunner + h.rd.execute(context.Background(), runner.ToRunner{Type: "resume", Session: h.id, PlacementGeneration: 2}, func(m runner.FromRunner) { answer = m }, AgentConfig{}, &agentSessionState{}) + current, _ := h.rd.reg.snapshot(h.id) + if !answer.OK || current.boot != claimedBoot || current.state != "running" || current.resumePending { + t.Fatal("resume failed to publish exact claimed boot") + } + h.rd.execute(context.Background(), runner.ToRunner{Type: "resume", Session: h.id, PlacementGeneration: 2}, func(m runner.FromRunner) { answer = m }, AgentConfig{}, &agentSessionState{}) + if answer.OK { + t.Fatal("replayed cold resume was accepted") + } +} + +func TestResumeStatusRequiresExactSettledPlacement(t *testing.T) { + for _, tc := range []struct { + name, state string + pending bool + want string + }{ + {"running", "running", false, "running"}, + {"cold", "suspended", false, "suspended_cold"}, + {"launching", "resuming", true, "resuming"}, + } { + t.Run(tc.name, func(t *testing.T) { + h := newIdleHarness(t) + h.rd.reg.mu.Lock() + row := h.rd.reg.items[h.id] + row.state = tc.state + row.placementGen = 2 + row.resumePending = tc.pending + h.rd.reg.mu.Unlock() + var reply runner.FromRunner + h.rd.execute(context.Background(), runner.ToRunner{Type: "resume_status", Session: h.id, PlacementGeneration: 2}, func(m runner.FromRunner) { reply = m }, AgentConfig{}, &agentSessionState{}) + if !reply.OK || reply.State != tc.want || reply.PlacementGeneration != 2 { + t.Fatal("resume status omitted exact pending or settled ownership") + } + h.rd.execute(context.Background(), runner.ToRunner{Type: "resume_status", Session: h.id, PlacementGeneration: 1}, func(m runner.FromRunner) { reply = m }, AgentConfig{}, &agentSessionState{}) + if reply.OK { + t.Fatal("stale status query was accepted") + } + }) + } +} + +// A refused command (for example, idle stop still running) never reached the +// driver's resume call. Reconciliation must fence that command before retry. +func TestResumeStatusFencesUnreceivedLegacyResume(t *testing.T) { + h := newIdleHarness(t) + if err := h.rd.Op(context.Background(), h.id, "suspend", false); err != nil { + t.Fatal(err) + } + h.rd.reg.mu.Lock() + h.rd.reg.items[h.id].placementGen = 1 + h.rd.reg.mu.Unlock() + reply := h.rd.resumeStatus(context.Background(), runner.ToRunner{Session: h.id, PlacementGeneration: 2}) + if !reply.OK || reply.State != "suspended_cold" || reply.PlacementGeneration != 2 { + t.Fatalf("unreceived resume cannot settle: %+v", reply) + } + if err := h.rd.opAtPlacement(context.Background(), h.id, "resume", false, 2); err == nil { + t.Fatal("delayed canceled command restarted the sandbox") + } + if err := h.rd.opAtPlacement(context.Background(), h.id, "resume", false, 3); err != nil { + t.Fatalf("next claim cannot resume: %v", err) + } +} + +type noRestartResumeDriver struct{ *driver.Fake } + +func (d noRestartResumeDriver) Resume(context.Context, string) (bool, error) { return false, nil } +func TestColdResumeRequiresActualRestart(t *testing.T) { + h := newIdleHarness(t) + if err := h.rd.Op(context.Background(), h.id, "suspend", false); err != nil { + t.Fatal(err) + } + h.rd.reg.mu.Lock() + h.rd.reg.items[h.id].placementGen = 1 + h.rd.reg.mu.Unlock() + h.rd.drv = noRestartResumeDriver{h.fd} + if err := h.rd.opAtPlacement(context.Background(), h.id, "resume", false, 2); err == nil { + t.Fatal("cold resume reported success without a restarted sandbox") + } + row, _ := h.rd.reg.snapshot(h.id) + if row.resumePending || row.state == "running" { + t.Fatal("refused cold resume published running or kept in-flight claim") + } +} +func TestColdResumeCompletionPreservesDeletionOwnership(t *testing.T) { + h := newIdleHarness(t) + if err := h.rd.Op(context.Background(), h.id, "suspend", false); err != nil { + t.Fatal(err) + } + h.rd.reg.mu.Lock() + h.rd.reg.items[h.id].placementGen = 1 + h.rd.reg.mu.Unlock() + h.rd.drv = &observedResumeDriver{Fake: h.fd, before: func() { + // Delete sets this before waiting for the driver-owned resume to finish. + h.rd.reg.setState(h.id, "destroying") + }} + _ = h.rd.opAtPlacement(context.Background(), h.id, "resume", false, 2) + row, _ := h.rd.reg.snapshot(h.id) + if row.state != "destroying" { + t.Fatalf("BUG: late resume changed deletion-owned state to %s", row.state) + } +} diff --git a/internal/runnerd/runnerd.go b/internal/runnerd/runnerd.go index ef0ed32..1e7dcd6 100644 --- a/internal/runnerd/runnerd.go +++ b/internal/runnerd/runnerd.go @@ -26,6 +26,7 @@ import ( ) type Server struct { + guestForwards guestRPCForwards // reconnectControl is a connection-scoped view of the same agent RPC queue. // A handshake captures it once; reconnecting never migrates a pending proof. reconnectControl atomic.Pointer[reconnectControl] @@ -315,7 +316,7 @@ func (s *Server) Recover(ctx context.Context) error { // running nor that it has finished. It says so — the entry is in // neither capacity count — until a child_exited or a cold resume // tells it. See sessionEntry.recovered. - e := &sessionEntry{id: l.SessionID, handle: l.Handle.ID, state: state, recovered: true} + e := &sessionEntry{id: l.SessionID, handle: l.Handle.ID, state: state, recovered: true, guestReconnect: l.GuestReconnect, placementGen: l.PlacementGeneration} s.reg.put(l.SessionID, e) } // The exemption is worth saying out loud where an operator will see it: a @@ -445,7 +446,7 @@ func (s *Server) createWithID(ctx context.Context, id string, spec driver.Spec, // and nothing else in this process remembers what was injected. Keys only // — see sessionEntry.envKeys. if !s.reg.putIfAbsent(id, &sessionEntry{id: id, state: "starting", allow: allow, - envKeys: envKeys(spec.Env), placementGen: placementGen}) { + envKeys: envKeys(spec.Env), placementGen: placementGen, guestReconnect: spec.GuestReconnect == runner.GuestReconnectProtocol}) { return errSessionExists } if err := s.pushEgress(id, allow); err != nil { @@ -453,6 +454,7 @@ func (s *Server) createWithID(ctx context.Context, id string, spec driver.Spec, return &egressError{err: err} } spec.SessionID = id + spec.PlacementGeneration = placementGen spec.DialURL = s.dialBase + "/register" spec.ProxyURL = s.proxyURL if s.withholdsSecrets() { @@ -710,6 +712,10 @@ func envKeys(env map[string]string) []string { // this function has no business knowing about. Snapshot has its own entry // point — see OpSnapshot. func (s *Server) Op(ctx context.Context, id, op string, warm bool) error { + return s.opAtPlacement(ctx, id, op, warm, 0) +} + +func (s *Server) opAtPlacement(ctx context.Context, id, op string, warm bool, generation uint64) error { // Marked before the handle is even read, so there is no gap between the // guard and the mark for a sweep to land in: from here on this session is // not idle, whatever this op turns out to be. An unknown or still-starting @@ -819,8 +825,29 @@ func (s *Server) Op(ctx context.Context, id, op string, warm bool) error { // edge-case table names. return errSuspendInFlight } - restarted, err := s.drv.Resume(ctx, handle) + var claimedBoot uint64 + if generation != 0 { + claimedBoot, err = s.claimResumePlacement(id, handle, generation) + if err != nil { + return err + } + } + var restarted bool + var err error + if placementDriver, ok := s.drv.(interface { + ResumePlacement(context.Context, string, uint64) (bool, error) + }); ok && generation != 0 { + restarted, err = placementDriver.ResumePlacement(ctx, handle, generation) + } else { + restarted, err = s.drv.Resume(ctx, handle) + } + if err == nil && generation != 0 && !restarted { + err = errResumePlacement + } if err != nil { + if generation != 0 { + s.failResumePlacement(context.WithoutCancel(ctx), id, handle, generation, claimedBoot) + } return err } // Lands on "running" and, for a sandbox the driver actually RESTARTED, @@ -829,6 +856,9 @@ func (s *Server) Op(ctx context.Context, id, op string, warm bool) error { // child that is running now. Only the driver can tell that from an // unpause — Inspect folds paused and exited into one state — which is // why it says so. See registry.resumed. + if generation != 0 { + return s.completeResumePlacement(id, handle, generation, claimedBoot) + } s.reg.resumed(id, restarted) return nil default: @@ -1064,6 +1094,7 @@ func (s *Server) serveSessionConn(ctx context.Context, id string, conn relay.Con // process on the other end may not be the one that numbered the last // report. See sessionEntry.execReg. reg := s.reg.registration() + authority := &guestRelayAuthority{} hub := relay.NewHubWithControl(ctx, conn, func(payload []byte) { // On its own goroutine, deliberately: this runs on the hub's read // loop, the single goroutine demultiplexing every attachment @@ -1083,9 +1114,9 @@ func (s *Server) serveSessionConn(ctx context.Context, id string, conn relay.Con // that grows ORDERED events needs a queue here instead — and // "exec_count" is one, which is why it carries a sequence number of // its own and this hop stayed as it is. - go s.routeControl(id, boot, reg, payload) + go authority.run(func() { s.routeControl(id, boot, reg, payload) }) }) - displaced, ok := s.reg.setHub(id, hub) + displaced, ok := s.reg.setHubAuthority(id, hub, authority) if !ok { // The entry vanished between our existence check above and now — a // concurrent DELETE raced this dial-in (session torn down while its @@ -1114,6 +1145,11 @@ func (s *Server) serveSessionConn(ctx context.Context, id string, conn relay.Con // the now-dead entry and let this goroutine (and its fd) go — leaving // this on r.Context() instead would leak both per session, forever, on // every abrupt death or explicit rm. + s.monitorSessionHub(id, hub) +} + +func (s *Server) monitorSessionHub(id string, hub *relay.Hub) { + defer s.guestForwards.discard(hub) <-hub.Done() handle, state, ok := s.reg.hubDied(id, hub) hub.Close() @@ -1198,6 +1234,10 @@ func (s *Server) RemoveWorkspace(ctx context.Context, id string) error { // carries every viewer's terminal traffic, and the one thing that must not // happen is a malformed frame taking the session down with it. func (s *Server) routeControl(id string, boot, reg uint64, payload []byte) { + s.routeControlWithEvents(id, boot, reg, payload, s.fireEventDetail) +} + +func (s *Server) routeControlWithEvents(id string, boot, reg uint64, payload []byte, emit func(string, string, string)) { var ev relay.ControlEvent if err := json.Unmarshal(payload, &ev); err != nil { log.Printf("session %s: undecodable control payload (%d bytes): %v", id, len(payload), err) @@ -1205,7 +1245,7 @@ func (s *Server) routeControl(id string, boot, reg uint64, payload []byte) { } switch ev.Kind { case "setup_done": - s.fireEventDetail(id, "setup_done", "") + emit(id, "setup_done", "") case "setup_failed", "stage_failed": // One event under two names. A session's boot is a chain of stages // (setup, then clone, then init — see cmd/sessiond/gitchain.go), and @@ -1234,10 +1274,10 @@ func (s *Server) routeControl(id string, boot, reg uint64, payload []byte) { // (not resumable) can serve that. See sessionEntry.bootFailed. s.reg.markBootFailed(id, boot) if stage == "setup" { - s.fireEventDetail(id, "setup_failed", setupFailedDetail(ev.RC, ev.Tail)) + emit(id, "setup_failed", setupFailedDetail(ev.RC, ev.Tail)) return } - s.fireEventDetail(id, "stage_failed", stageFailedDetail(stage, ev.RC, ev.Tail)) + emit(id, "stage_failed", stageFailedDetail(stage, ev.RC, ev.Tail)) case relay.KindSuspendAck: // The sandbox heard the notice. It stays between this process and that // sandbox: controld asked for a suspend and is waiting for THAT @@ -1268,7 +1308,7 @@ func (s *Server) routeControl(id string, boot, reg uint64, payload []byte) { // credential it minted for this session, and a token — or anything // derived from one — has no business on this channel. log.Printf("session %s: a git operation was refused by GitHub; reporting the credential", id) - s.fireEventDetail(id, "credential_rejected", "") + emit(id, "credential_rejected", "") case "child_exited": // The agent process inside the container ended. This is news, not a // verdict: the session stays up (sessiond outlives its child so @@ -1286,7 +1326,7 @@ func (s *Server) routeControl(id string, boot, reg uint64, payload []byte) { // nothing else in this runner would remember it. It still changes no // state here — see RunIdleStop for the timeout that does. s.reg.childExited(id, boot, s.now()) - s.fireEventDetail(id, "child_exited", strconv.Itoa(ev.RC)) + emit(id, "child_exited", strconv.Itoa(ev.RC)) case relay.KindExecCount: // How many commands `rainier exec` is running in there. Recorded and // not reported: it is the runner's own fact, the way attachment diff --git a/protocol/runner/bootstrap_redeem.go b/protocol/runner/bootstrap_redeem.go new file mode 100644 index 0000000..7d67ffd --- /dev/null +++ b/protocol/runner/bootstrap_redeem.go @@ -0,0 +1,19 @@ +package runner + +// SessionBootstrapRedeemRequest is the exact one-use secret redemption body. +// The same decoder is used for legacy boot and proof-issued host redemption. +type SessionBootstrapRedeemRequest struct { + Protocol uint64 `json:"protocol"` + Token string `json:"token"` +} + +func DecodeSessionBootstrapRedeemRequest(payload []byte) (SessionBootstrapRedeemRequest, error) { + var v SessionBootstrapRedeemRequest + if decodeReconnectObject(payload, &v, "protocol", "token") != nil || v.Protocol != SessionBootstrapProtocolVersion { + return SessionBootstrapRedeemRequest{}, errGuestReconnectMessage + } + if _, ok := reconnectBytes(v.Token, 32); !ok { + return SessionBootstrapRedeemRequest{}, errGuestReconnectMessage + } + return v, nil +} diff --git a/protocol/runner/bootstrap_redeem_test.go b/protocol/runner/bootstrap_redeem_test.go new file mode 100644 index 0000000..699fe2d --- /dev/null +++ b/protocol/runner/bootstrap_redeem_test.go @@ -0,0 +1,19 @@ +package runner + +import ( + "encoding/base64" + "testing" +) + +func TestBootstrapRedemptionExactWire(t *testing.T) { + token := base64.RawURLEncoding.EncodeToString(make([]byte, 32)) + valid := `{"protocol":1,"token":"` + token + `"}` + if _, err := DecodeSessionBootstrapRedeemRequest([]byte(valid)); err != nil { + t.Fatal(err) + } + for _, raw := range []string{`{"protocol":1}`, `{"protocol":1,"token":null}`, valid + ` {}`, `{"Protocol":1,"token":"` + token + `"}`, `{"protocol":1,"protocol":1,"token":"` + token + `"}`, `{"protocol":1,"token":"not-canonical"}`, `{"protocol":1,"token":"` + token + `","extra":1}`} { + if _, err := DecodeSessionBootstrapRedeemRequest([]byte(raw)); err == nil { + t.Fatal("malformed redemption accepted") + } + } +} diff --git a/protocol/runner/messages.go b/protocol/runner/messages.go index 199db05..55b614a 100644 --- a/protocol/runner/messages.go +++ b/protocol/runner/messages.go @@ -103,6 +103,10 @@ const SessionBootstrapProtocolVersion = 1 // Spec, byte for byte. const CapabilityMicrovmV1 = "microvm.v1" +// CapabilityGuestReconnectV1 negotiates fresh guest enrollment and authenticated +// live reconnect. It requires both the microVM driver and control-plane handlers. +const CapabilityGuestReconnectV1 = "guest_reconnect.v1" + // HomeMount is the agent home a create mounts into a sandbox: one writable // volume per (creator, workspace), landing at Path, inside which each coding // agent gets its own subdirectory. It is what makes "log in once" true across @@ -164,12 +168,13 @@ type RPCEnvelope struct { // a clean exit). That last one moves no state machine: the container stays up // for viewers, so it is an observation controld records against the session. type FromRunner struct { - Type string `json:"type"` // "announce" | "result" | "event" | "session_req" - Proto int `json:"proto,omitempty"` // announce - Runner string `json:"runner,omitempty"` // announce - Sessions []SessionInfo `json:"sessions,omitempty"` // announce - Used int `json:"used"` - Total int `json:"total"` + CapacityPlacements map[string]uint64 `json:"capacity_placements,omitempty"` + Type string `json:"type"` // "announce" | "result" | "event" | "session_req" + Proto int `json:"proto,omitempty"` // announce + Runner string `json:"runner,omitempty"` // announce + Sessions []SessionInfo `json:"sessions,omitempty"` // announce + Used int `json:"used"` + Total int `json:"total"` // Active and IdleExited split Used by whether the sandbox has WORK in it: // Active counts sandboxes that are up with their child process still // running OR with a `rainier exec` command running (a detached run on a @@ -266,7 +271,7 @@ type ToRunner struct { // "accept" is controld's answer to an announce, sent before any command: // the generation this connection acts under and the announced // capabilities controld will schedule on. - Type string `json:"type"` // "accept"|"create"|"destroy"|"remove_workspace"|"suspend"|"resume"|"snapshot"|"prepull"|"dial_attach"|"session_rpc" + Type string `json:"type"` // "accept"|"create"|"destroy"|"remove_workspace"|"suspend"|"resume"|"resume_status"|"snapshot"|"prepull"|"dial_attach"|"session_rpc" ReqID uint64 `json:"req_id,omitempty"` Session string `json:"session,omitempty"` Spec *Spec `json:"spec,omitempty"` // create @@ -328,10 +333,12 @@ type RepoSpec struct { // passes only the pieces that apply; Env values are secrets as often as not // and never logged verbatim. type Spec struct { - Name string `json:"name,omitempty"` - Image string `json:"image,omitempty"` - Cmd []string `json:"cmd,omitempty"` - EgressAllow []string `json:"egress_allow,omitempty"` + // GuestReconnect opts a fresh guest into authenticated reconnect enrollment. + GuestReconnect uint64 `json:"guest_reconnect,omitempty"` + Name string `json:"name,omitempty"` + Image string `json:"image,omitempty"` + Cmd []string `json:"cmd,omitempty"` + EgressAllow []string `json:"egress_allow,omitempty"` // Setup is the environment's setup script, run once inside the fresh // container; the runner reports its outcome as a "setup_done" / // "setup_failed" event. SetupTimeoutSec bounds that run (0 = the diff --git a/protocol/runner/reconnect_configuration.go b/protocol/runner/reconnect_configuration.go new file mode 100644 index 0000000..a7718df --- /dev/null +++ b/protocol/runner/reconnect_configuration.go @@ -0,0 +1,95 @@ +package runner + +import ( + "bytes" + "encoding/json" + "io" +) + +// MethodGuestReconnectConfiguration is runner-only. A protocol-only request +// resolves current non-secret launch configuration under the authenticated +// session placement and current policy. It neither proves a guest nor mints a +// token. Use DecodeGuestReconnectBeginRequest for its protocol-only request. +const MethodGuestReconnectConfiguration = "guest_reconnect_configuration" + +// GuestReconnectConfigurationLimit bounds an authorized configuration response, +// which may contain setup/init scripts. Requests and proof frames retain 4 KiB. +const GuestReconnectConfigurationLimit = 4 << 20 + +// GuestReconnectConfiguration carries fresh launch material with no bootstrap +// token or secret values. The host checks session/placement against its retained +// ownership, then supplies only the token returned by this stream's proof. +// It must never replay this response on a later attempt or persist it to disk. +type GuestReconnectConfiguration struct { + Protocol uint64 `json:"protocol"` + SessionID string `json:"session_id"` + PlacementGeneration uint64 `json:"placement_generation"` + Spec *Spec `json:"spec"` +} + +// DecodeGuestReconnectConfiguration bounds the whole payload, rejects duplicate, +// null, unknown or trailing values, and returns zero configuration on failure. +// Nested launch fields keep Spec's existing encoding/json field semantics. +func DecodeGuestReconnectConfiguration(payload []byte) (GuestReconnectConfiguration, error) { + var v GuestReconnectConfiguration + if decodeReconnectObjectLimit(payload, &v, GuestReconnectConfigurationLimit, "protocol", "session_id", "placement_generation", "spec") != nil || v.Protocol != GuestReconnectProtocol || !reconnectID(v.SessionID) || v.PlacementGeneration == 0 || v.Spec == nil || v.Spec.BootstrapToken != "" { + return GuestReconnectConfiguration{}, errGuestReconnectMessage + } + d := json.NewDecoder(bytes.NewReader(payload)) + d.UseNumber() + if !uniqueConfigurationValue(d, 0) { + return GuestReconnectConfiguration{}, errGuestReconnectMessage + } + if _, err := d.Token(); err != io.EOF { + return GuestReconnectConfiguration{}, errGuestReconnectMessage + } + d = json.NewDecoder(bytes.NewReader(payload)) + d.DisallowUnknownFields() + if d.Decode(&v) != nil { + return GuestReconnectConfiguration{}, errGuestReconnectMessage + } + return v, nil +} + +// Bound recursion independently of the byte limit; tenant environment maps +// have arbitrary keys but must not hide duplicates from typed JSON decoding. +func uniqueConfigurationValue(d *json.Decoder, depth int) bool { + if depth > 32 { + return false + } + token, err := d.Token() + if err != nil || token == nil { + return false + } + delim, ok := token.(json.Delim) + if !ok { + return true + } + switch delim { + case '{': + seen := map[string]bool{} + for d.More() { + token, err := d.Token() + key, ok := token.(string) + if err != nil || !ok || seen[key] { + return false + } + seen[key] = true + if !uniqueConfigurationValue(d, depth+1) { + return false + } + } + token, err = d.Token() + return err == nil && token == json.Delim('}') + case '[': + for d.More() { + if !uniqueConfigurationValue(d, depth+1) { + return false + } + } + token, err = d.Token() + return err == nil && token == json.Delim(']') + default: + return false + } +} diff --git a/protocol/runner/reconnect_configuration_test.go b/protocol/runner/reconnect_configuration_test.go new file mode 100644 index 0000000..02a7e67 --- /dev/null +++ b/protocol/runner/reconnect_configuration_test.go @@ -0,0 +1,38 @@ +package runner_test + +import ( + "encoding/json" + "github.com/tokencanopy/rainier/protocol/runner" + "strings" + "testing" +) + +func TestReconnectConfigurationContract(t *testing.T) { + valid := `{"protocol":1,"session_id":"session_test","placement_generation":3,"spec":{"image":"example.invalid/test","env":{"CONFIG_TEST":"test"}}}` + got, err := runner.DecodeGuestReconnectConfiguration([]byte(valid)) + if err != nil || got.SessionID != "session_test" || got.Spec.Env["CONFIG_TEST"] != "test" { + t.Fatal("valid configuration refused") + } + for _, bad := range []string{ + strings.Replace(valid, `"protocol":1`, `"protocol":2`, 1), + strings.Replace(valid, `"placement_generation":3`, `"placement_generation":0`, 1), + strings.Replace(valid, `"session_test"`, `""`, 1), + strings.Replace(valid, `"spec":{`, `"spec":{"bootstrap_token":"token_test",`, 1), + strings.Replace(valid, `"spec":{`, `"spec":{"unknown":true,`, 1), + strings.Replace(valid, `"spec":{`, `"spec":{"image":"duplicate_test",`, 1), + strings.Replace(valid, `"env":{`, `"env":{"CONFIG_TEST":"duplicate_test",`, 1), + strings.Replace(valid, `"protocol":1`, `"protocol":1,"protocol":1`, 1), + strings.Replace(valid, `"protocol":1`, `"Protocol":1`, 1), + valid + ` true`, `null`, strings.Repeat(" ", runner.GuestReconnectConfigurationLimit+1), + } { + rejected, err := runner.DecodeGuestReconnectConfiguration([]byte(bad)) + if err == nil || rejected.Spec != nil || rejected.SessionID != "" { + t.Fatal("invalid configuration accepted or retained data") + } + } + large := runner.GuestReconnectConfiguration{Protocol: 1, SessionID: "session_test", PlacementGeneration: 3, Spec: &runner.Spec{Setup: strings.Repeat("x", 8192)}} + raw, _ := json.Marshal(large) + if _, err := runner.DecodeGuestReconnectConfiguration(raw); err != nil { + t.Fatal("configuration incorrectly uses proof limit") + } +} diff --git a/protocol/runner/reconnect_rpc.go b/protocol/runner/reconnect_rpc.go index ecc410e..78a3d9c 100644 --- a/protocol/runner/reconnect_rpc.go +++ b/protocol/runner/reconnect_rpc.go @@ -197,7 +197,11 @@ func reconnectID(id string) bool { // Check the flat object before typed decoding: encoding/json alone folds field // case and accepts duplicate keys. No raw decoder error may escape this boundary. func decodeReconnectObject(payload []byte, out any, names ...string) error { - if len(payload) > GuestReconnectPayloadLimit { + return decodeReconnectObjectLimit(payload, out, GuestReconnectPayloadLimit, names...) +} + +func decodeReconnectObjectLimit(payload []byte, out any, limit int, names ...string) error { + if len(payload) > limit { return errGuestReconnectMessage } d := json.NewDecoder(bytes.NewReader(payload)) diff --git a/runnerplane/connect.go b/runnerplane/connect.go index 464898e..24466de 100644 --- a/runnerplane/connect.go +++ b/runnerplane/connect.go @@ -108,11 +108,12 @@ func (p *Plane) handleConnect(w http.ResponseWriter, r *http.Request) { // this very connection. res, err := p.host.Fleet().ReconcileRunner(connCtx, control.RunnerSnapshot{ WorkspaceID: b.WorkspaceID, PoolID: b.PoolID, - RunnerID: b.RunnerID, - Generation: rc.gen, - CapacityUsed: ann.Used, - CapacityTotal: ann.Total, - Sessions: announcedSessions(ann.Sessions), + RunnerID: b.RunnerID, + Generation: rc.gen, + CapacityUsed: ann.Used, + CapacityTotal: ann.Total, + CapacityPlacements: capacityPlacements(ann.CapacityPlacements), + Sessions: announcedSessions(ann.Sessions), }) if err != nil { p.logf("reconciling runner %s (generation %d): %v", name, rc.gen, err) @@ -205,12 +206,13 @@ func (p *Plane) connectRunner(ctx context.Context, rc *runnerConn, ann runner.Fr reg, err := p.host.Fleet().RegisterRunner(ctx, control.RunnerRegistration{ WorkspaceID: rc.binding.WorkspaceID, PoolID: rc.binding.PoolID, - RunnerID: rc.binding.RunnerID, - Generation: rc.gen, - CapacityUsed: ann.Used, - CapacityTotal: ann.Total, - Capabilities: rc.caps, - Sessions: announcedSessions(ann.Sessions), + RunnerID: rc.binding.RunnerID, + Generation: rc.gen, + CapacityUsed: ann.Used, + CapacityTotal: ann.Total, + CapacityPlacements: capacityPlacements(ann.CapacityPlacements), + Capabilities: rc.caps, + Sessions: announcedSessions(ann.Sessions), }) switch { case err != nil: diff --git a/runnerplane/events.go b/runnerplane/events.go index f55e2f4..9f831c7 100644 --- a/runnerplane/events.go +++ b/runnerplane/events.go @@ -30,6 +30,9 @@ const ( // happen under the runner's name lock, so a reconnect can neither slip between // them nor have its own row write overtaken by this one. func (p *Plane) touchRunner(ctx context.Context, rc *runnerConn, m runner.FromRunner) bool { + if control.ValidateCapacityPlacements(m.Used, capacityPlacements(m.CapacityPlacements)) != nil { + return false + } nl := p.nameLock(rc.binding.PoolID, rc.name) nl.Lock() defer nl.Unlock() @@ -38,14 +41,15 @@ func (p *Plane) touchRunner(ctx context.Context, rc *runnerConn, m runner.FromRu return false } err := p.host.FleetRepository().UpsertRunner(ctx, rc.binding.PoolID, control.Runner{ - ID: rc.binding.RunnerID, - PoolID: rc.binding.PoolID, - CapacityUsed: m.Used, - CapacityTotal: m.Total, - Connected: true, - Generation: rc.gen, - Capabilities: rc.caps, - LastSeenAt: time.Now(), + ID: rc.binding.RunnerID, + PoolID: rc.binding.PoolID, + CapacityUsed: m.Used, + CapacityTotal: m.Total, + CapacityPlacements: capacityPlacements(m.CapacityPlacements), + Connected: true, + Generation: rc.gen, + Capabilities: rc.caps, + LastSeenAt: time.Now(), }) switch { case errors.Is(err, control.ErrStale): @@ -207,3 +211,14 @@ func stageFailure(stage, detail string) string { } return clip(stage) + " failed: " + detail } + +func capacityPlacements(wire map[string]uint64) map[control.SessionID]uint64 { + if len(wire) == 0 { + return nil + } + out := make(map[control.SessionID]uint64, len(wire)) + for id, generation := range wire { + out[control.SessionID(id)] = generation + } + return out +} diff --git a/runnerplane/plane_test.go b/runnerplane/plane_test.go index 4a8bde7..30558db 100644 --- a/runnerplane/plane_test.go +++ b/runnerplane/plane_test.go @@ -977,6 +977,7 @@ func TestRedialSurvivesStaleDisconnect(t *testing.T) { entered = make(chan struct{}) upserted = make(chan struct{}) wrote = make(chan struct{}) + allowInstall = make(chan struct{}) enteredOnce sync.Once upsertedOnce sync.Once wroteOnce sync.Once @@ -997,6 +998,7 @@ func TestRedialSurvivesStaleDisconnect(t *testing.T) { case <-entered: if r.Connected { upsertedOnce.Do(func() { close(upserted) }) + <-allowInstall } default: } @@ -1011,8 +1013,17 @@ func TestRedialSurvivesStaleDisconnect(t *testing.T) { // The teardown is now inside its disconnect write; redial into that gap. awaitChan(t, entered, "teardown's disconnect write") - startFakeRunner(t, ts, runnerScript{Name: "vm1", Used: 1, Total: 4}) + second := startFakeRunner(t, ts, runnerScript{Name: "vm1", Used: 1, Total: 4}) awaitChan(t, wrote, "teardown's disconnect write returning") + awaitChan(t, upserted, "redial registration persisted") + // Persistence precedes socket installation. A connected database row is + // not a handshake barrier; hold that real gap open to exercise it. + installedEarly := p.Transport().Connected(testPool, "vm1") + close(allowInstall) + if installedEarly { + t.Fatal("replacement installed before its registration returned") + } + drainAccept(t, second) eventually(t, 3*time.Second, func() error { rows, err := h.repo.ListRunners(context.Background(), testPool)