diff --git a/.changes/unreleased/+controlled-session-network-planning.yaml b/.changes/unreleased/+controlled-session-network-planning.yaml index 0e3e0adf..96f5ec58 100644 --- a/.changes/unreleased/+controlled-session-network-planning.yaml +++ b/.changes/unreleased/+controlled-session-network-planning.yaml @@ -1,2 +1,2 @@ kind: Added -body: Freeze, realize, recover, and verify lease-private controlled-session networks and granted workload endpoints, with protocol-v2 session-local coordinates, fixed participant addresses, exact peer firewall grants, and ordinary public or local access preserved only when explicitly granted. +body: Freeze, realize, recover, and verify lease-private controlled-session networks and granted workload endpoints, with protocol-v1 session-local coordinates, fixed participant addresses, exact peer firewall grants, and ordinary public or local access preserved only when explicitly granted. diff --git a/docs/CONTROLLED_SESSION_DESIGN.md b/docs/CONTROLLED_SESSION_DESIGN.md index afc416bb..0647ebfa 100644 --- a/docs/CONTROLLED_SESSION_DESIGN.md +++ b/docs/CONTROLLED_SESSION_DESIGN.md @@ -114,7 +114,7 @@ summary: Capability-scoped execution sessions that inherit Reploy's global conta implemented. The host resolves a sorted requested subset of the exact workload generation's declared endpoints and freezes their schemes, container ports, lease-local aliases, and internal network name into both - container-plan digests and the protocol-v2 `opened` coordinates. It creates + container-plan digests and the protocol-v1 `opened` coordinates. It creates one exact engine-internal network, records it before startup, derives fixed controller and workload addresses from the verified engine-assigned prefixes, attaches only the two exact inert containers, and removes the @@ -545,9 +545,8 @@ use, but secrecy is not the sole security boundary. Isolation relies on: ## Session Protocol The protocol is versioned, typed, length-framed, and binary-safe. The current -wire version is 2; version 1 remains reserved for the earlier strict `opened` -shape that did not contain endpoint coordinates. Terminal bytes are never -parsed as protocol messages. +and initial wire version is 1. Terminal bytes are never parsed as protocol +messages. ### Controller Requests @@ -578,6 +577,13 @@ cleanup and no canceled request is replayed. - `opened`: reports the effective dimensions, both runtime identities and generations, fixed session capabilities, the structured coordinates of each granted workload endpoint, and the workload-output-finalization timeout. + It means that the controller has claimed the authenticated channel; the + workload may not have started yet. +- `ready`: reports that Host Reploy has started the workload, verified its + lease-private network when present, and activated the lifecycle. Ordinary + controller requests are rejected before this payload-free event. A + startup failure never emits `ready`; the controller can still acknowledge the + resulting `terminated` event. - `output(bytes)`: ordered PTY output bytes. - `workload_exit(status, reason)`: reports host-observed workload-shell exit. @@ -630,7 +636,7 @@ Host Reploy owns workload-output finalization; it never waits indefinitely for workload cooperation. Once termination begins, it rejects new output surfaces, performs bounded graceful shutdown followed by forced container stop, and continues draining the PTY. The immutable session plan carries a finite -output-finalization deadline. Protocol v2 defines an initial host-owned default +output-finalization deadline. Protocol v1 defines an initial host-owned default of 30 seconds; the effective value is reported by `opened` and applies to workload shutdown, final buffered-byte delivery, and controller backpressure. @@ -651,7 +657,7 @@ completion by `complete` or terminal acknowledgement. The barrier initially covers the PTY. A future workload output-file or output-directory contract joins the same barrier after its files are closed, -validated, and published or have recorded an explicit failure; protocol v2 +validated, and published or have recorded an explicit failure; protocol v1 does not otherwise speculate about file payloads. Native network traffic is not session output and does not pass through this barrier. @@ -1122,6 +1128,8 @@ asciinema The proxy forwards input bytes, output bytes, resize operations, and terminal completion. This keeps asciinema and recording dependencies out of workload images while preserving the existing cast format and controller ownership. +It consumes `opened` as channel metadata and does not forward controller +requests until `ready`. The prototype must test: diff --git a/internal/controlledsession/authorization_test.go b/internal/controlledsession/authorization_test.go index 8357bfb9..0082b31c 100644 --- a/internal/controlledsession/authorization_test.go +++ b/internal/controlledsession/authorization_test.go @@ -22,10 +22,10 @@ func testAuthorizationV1() AuthorizationV1 { } } -func testEndpointsV1() []EndpointV2 { - return []EndpointV2{ - {ID: "browser", Scheme: "http", Host: WorkloadEndpointHostV2, Port: 8080}, - {ID: "terminal", Scheme: "https", Host: WorkloadEndpointHostV2, Port: 8443}, +func testEndpointsV1() []EndpointV1 { + return []EndpointV1{ + {ID: "browser", Scheme: "http", Host: WorkloadEndpointHostV1, Port: 8080}, + {ID: "terminal", Scheme: "https", Host: WorkloadEndpointHostV1, Port: 8443}, } } diff --git a/internal/controlledsession/channel.go b/internal/controlledsession/channel.go index 39daed32..52f48ab0 100644 --- a/internal/controlledsession/channel.go +++ b/internal/controlledsession/channel.go @@ -27,7 +27,7 @@ const ( // lease-private controller channel. HostDirectory must not already exist. type PrivateChannelConfigV1 struct { HostDirectory string - Opened OpenedV2 + Opened OpenedV1 } // ChannelClaimErrorV1 distinguishes failures before an authorized controller @@ -95,12 +95,12 @@ func PreparePrivateChannelV1(config PrivateChannelConfigV1) (*PrivateChannelV1, } openedPayload := config.Opened openedPayload.Authorization = cloneAuthorizationV1(config.Opened.Authorization) - openedPayload.Endpoints = append([]EndpointV2{}, config.Opened.Endpoints...) + openedPayload.Endpoints = append([]EndpointV1{}, config.Opened.Endpoints...) opened := EventV1{Kind: EventOpenedV1, Opened: &openedPayload} if err := ValidateEventV1(opened); err != nil { return nil, fmt.Errorf("prepare controlled-session channel opened event: %w", err) } - if err := WriteEventV2(io.Discard, opened); err != nil { + if err := WriteEventV1(io.Discard, opened); err != nil { return nil, fmt.Errorf("prepare controlled-session channel opened frame: %w", err) } controllerIdentity := openedPayload.Authorization.Controller.RuntimeIdentity @@ -131,7 +131,7 @@ func (channel *PrivateChannelV1) SocketPath() string { // Claim accepts the only controller connection, verifies its kernel-reported // identity, removes the listener pathname, and sends opened as the first event. -// A failed claim is terminal; protocol v2 does not reconnect or transfer +// A failed claim is terminal; protocol v1 does not reconnect or transfer // ownership. func (channel *PrivateChannelV1) Claim(ctx context.Context) (*ControllerConnectionV1, error) { if ctx == nil || ctx.Done() == nil { @@ -212,7 +212,7 @@ func (connection *ControllerConnectionV1) ReadRequest(ctx context.Context) (Requ var request RequestV1 err := withConnectionDeadlineV1(ctx, connection.connection.SetReadDeadline, func() error { var readErr error - request, readErr = ReadRequestV2(connection.connection) + request, readErr = ReadRequestV1(connection.connection) return readErr }) if err != nil { @@ -234,7 +234,7 @@ func (connection *ControllerConnectionV1) WriteEvent(ctx context.Context, event connection.writeMu.Lock() defer connection.writeMu.Unlock() err := withConnectionDeadlineV1(ctx, connection.connection.SetWriteDeadline, func() error { - return WriteEventV2(connection.connection, event) + return WriteEventV1(connection.connection, event) }) if err != nil { closeErr := connection.Close() diff --git a/internal/controlledsession/channel_linux_test.go b/internal/controlledsession/channel_linux_test.go index b8ad2abe..b2055dcb 100644 --- a/internal/controlledsession/channel_linux_test.go +++ b/internal/controlledsession/channel_linux_test.go @@ -40,7 +40,7 @@ func TestPrivateChannelV1CreatesOneControllerOwnedClaim(t *testing.T) { }{connection: connection, err: err} }() client := dialPrivateChannelV1(t, socket) - opened, err := ReadEventV2(client) + opened, err := ReadEventV1(client) if err != nil { t.Fatal(err) } @@ -60,7 +60,7 @@ func TestPrivateChannelV1CreatesOneControllerOwnedClaim(t *testing.T) { } wantRequest := RequestV1{Kind: RequestInputV1, Bytes: []byte{0, 3, 0xff}} - if err := WriteRequestV2(client, wantRequest); err != nil { + if err := WriteRequestV1(client, wantRequest); err != nil { t.Fatal(err) } request, err := result.connection.ReadRequest(ctx) @@ -74,7 +74,7 @@ func TestPrivateChannelV1CreatesOneControllerOwnedClaim(t *testing.T) { if err := result.connection.WriteEvent(ctx, wantEvent); err != nil { t.Fatal(err) } - event, err := ReadEventV2(client) + event, err := ReadEventV1(client) if err != nil { t.Fatal(err) } @@ -210,7 +210,7 @@ func TestPrivateChannelV1RejectsMalformedAndOversizedRequests(t *testing.T) { {name: "oversized", data: func() []byte { header := make([]byte, frameHeaderSizeV1) copy(header, frameMagicV1[:]) - header[4] = ProtocolVersionV2 + header[4] = ProtocolVersionV1 header[5] = byte(wireRequestInputV1) binary.BigEndian.PutUint32(header[6:], MaxFramePayloadV1+1) return header @@ -220,7 +220,7 @@ func TestPrivateChannelV1RejectsMalformedAndOversizedRequests(t *testing.T) { t.Run(test.name, func(t *testing.T) { channel, _ := prepareCurrentIdentityChannelV1(t) server, client := claimPrivateChannelV1(t, channel) - if _, err := ReadEventV2(client); err != nil { + if _, err := ReadEventV1(client); err != nil { t.Fatal(err) } if _, err := client.Write(test.data); err != nil { @@ -292,7 +292,7 @@ func TestControllerConnectionV1BoundsAndSerializesFlow(t *testing.T) { } got := make([]EventV1, 0, len(events)) for range events { - event, err := ReadEventV2(client) + event, err := ReadEventV1(client) if err != nil { t.Fatal(err) } @@ -315,7 +315,7 @@ func TestPreparePrivateChannelV1FreezesOpenedAuthorization(t *testing.T) { config.Opened.Endpoints[0].Port = 9999 server, client := claimPrivateChannelV1(t, channel) defer server.Close() - event, err := ReadEventV2(client) + event, err := ReadEventV1(client) if err != nil { t.Fatal(err) } @@ -352,7 +352,7 @@ func currentIdentityChannelConfigV1(t *testing.T) PrivateChannelConfigV1 { authorization.Controller.RuntimeIdentity = identity config := PrivateChannelConfigV1{ HostDirectory: filepath.Join(shortChannelTestDirectoryV1(t), "session"), - Opened: OpenedV2{ + Opened: OpenedV1{ Authorization: authorization, Endpoints: testEndpointsV1(), Columns: 80, Rows: 24, OutputFinalizationTimeoutMilliseconds: DefaultOutputFinalizationTimeoutMillisecondsV1, }, diff --git a/internal/controlledsession/client.go b/internal/controlledsession/client.go new file mode 100644 index 00000000..62987565 --- /dev/null +++ b/internal/controlledsession/client.go @@ -0,0 +1,141 @@ +package controlledsession + +import ( + "context" + "errors" + "fmt" + "net" + "sync" +) + +// SessionClientV1 is the controller-side owner of one claimed private session +// connection. It consumes the mandatory opened event during construction and +// then admits one event read and one request write concurrently. +type SessionClientV1 struct { + connection net.Conn + opened OpenedV1 + readMu sync.Mutex + writeMu sync.Mutex + stateMu sync.RWMutex + ready bool + terminated bool + closeOnce sync.Once + closeErr error +} + +func newSessionClientV1(ctx context.Context, connection net.Conn) (*SessionClientV1, error) { + if ctx == nil || ctx.Done() == nil { + return nil, fmt.Errorf("open controlled-session client: cancelable context is required") + } + if connection == nil { + return nil, fmt.Errorf("open controlled-session client: connection is required") + } + var event EventV1 + err := withConnectionDeadlineV1(ctx, connection.SetReadDeadline, func() error { + var readErr error + event, readErr = ReadEventV1(connection) + return readErr + }) + if err != nil { + return nil, errors.Join(fmt.Errorf("read controlled-session opened event: %w", err), connection.Close()) + } + if event.Kind != EventOpenedV1 || event.Opened == nil { + return nil, errors.Join(fmt.Errorf("controlled-session first event must be opened"), connection.Close()) + } + opened := *event.Opened + opened.Authorization = cloneAuthorizationV1(event.Opened.Authorization) + opened.Endpoints = append([]EndpointV1(nil), event.Opened.Endpoints...) + return &SessionClientV1{connection: connection, opened: opened}, nil +} + +// Opened returns an independent copy of the immutable session authorization +// and coordinates received during the connection claim. +func (client *SessionClientV1) Opened() OpenedV1 { + opened := client.opened + opened.Authorization = cloneAuthorizationV1(client.opened.Authorization) + opened.Endpoints = append([]EndpointV1(nil), client.opened.Endpoints...) + return opened +} + +// Ready reports whether the host has verified workload startup and activated +// the lifecycle. It becomes true only after ReadEvent returns the one ready +// event and never becomes false. +func (client *SessionClientV1) Ready() bool { + client.stateMu.RLock() + defer client.stateMu.RUnlock() + return client.ready +} + +func (client *SessionClientV1) ReadEvent(ctx context.Context) (EventV1, error) { + if ctx == nil || ctx.Done() == nil { + return EventV1{}, fmt.Errorf("read controlled-session client event: cancelable context is required") + } + client.readMu.Lock() + defer client.readMu.Unlock() + var event EventV1 + err := withConnectionDeadlineV1(ctx, client.connection.SetReadDeadline, func() error { + var readErr error + event, readErr = ReadEventV1(client.connection) + return readErr + }) + if err != nil { + return EventV1{}, errors.Join(fmt.Errorf("read controlled-session client event: %w", err), client.Close()) + } + if event.Kind == EventOpenedV1 { + return EventV1{}, errors.Join(fmt.Errorf("read controlled-session client event: opened may appear only once"), client.Close()) + } + if event.Kind == EventReadyV1 { + client.stateMu.Lock() + if client.ready { + client.stateMu.Unlock() + return EventV1{}, errors.Join(fmt.Errorf("read controlled-session client event: ready may appear only once"), client.Close()) + } + client.ready = true + client.stateMu.Unlock() + } + if event.Kind == EventTerminatedV1 { + client.stateMu.Lock() + client.terminated = true + client.stateMu.Unlock() + } + return event, nil +} + +func (client *SessionClientV1) WriteRequest(ctx context.Context, request RequestV1) error { + if ctx == nil || ctx.Done() == nil { + return fmt.Errorf("write controlled-session client request: cancelable context is required") + } + if err := ValidateRequestV1(request); err != nil { + return err + } + client.writeMu.Lock() + defer client.writeMu.Unlock() + client.stateMu.RLock() + ready := client.ready + terminated := client.terminated + client.stateMu.RUnlock() + if request.Kind == RequestAcknowledgeTerminatedV1 { + if !terminated { + return fmt.Errorf("write controlled-session client request: terminal result has not been received") + } + } else if terminated { + return fmt.Errorf("write controlled-session client request: session is terminated") + } + if request.Kind != RequestAcknowledgeTerminatedV1 && !ready { + return fmt.Errorf("write controlled-session client request: session is not ready") + } + err := withConnectionDeadlineV1(ctx, client.connection.SetWriteDeadline, func() error { + return WriteRequestV1(client.connection, request) + }) + if err != nil { + return errors.Join(fmt.Errorf("write controlled-session client request: %w", err), client.Close()) + } + return nil +} + +func (client *SessionClientV1) Close() error { + client.closeOnce.Do(func() { + client.closeErr = client.connection.Close() + }) + return client.closeErr +} diff --git a/internal/controlledsession/client_linux.go b/internal/controlledsession/client_linux.go new file mode 100644 index 00000000..2c286de1 --- /dev/null +++ b/internal/controlledsession/client_linux.go @@ -0,0 +1,26 @@ +//go:build linux + +package controlledsession + +import ( + "context" + "fmt" + "net" + "path/filepath" +) + +// DialSessionClientV1 claims the Linux lease-private socket exposed only to +// the controller and consumes its mandatory opened event. +func DialSessionClientV1(ctx context.Context, socketPath string) (*SessionClientV1, error) { + if ctx == nil || ctx.Done() == nil { + return nil, fmt.Errorf("dial controlled-session client: cancelable context is required") + } + if !filepath.IsAbs(socketPath) || filepath.Clean(socketPath) != socketPath { + return nil, fmt.Errorf("dial controlled-session client requires an absolute clean socket path") + } + connection, err := (&net.Dialer{}).DialContext(ctx, "unix", socketPath) + if err != nil { + return nil, fmt.Errorf("dial controlled-session socket: %w", err) + } + return newSessionClientV1(ctx, connection) +} diff --git a/internal/controlledsession/client_linux_test.go b/internal/controlledsession/client_linux_test.go new file mode 100644 index 00000000..dbe8245d --- /dev/null +++ b/internal/controlledsession/client_linux_test.go @@ -0,0 +1,52 @@ +//go:build linux + +package controlledsession + +import ( + "context" + "reflect" + "strings" + "testing" + "time" +) + +func TestDialSessionClientV1ClaimsPrivateChannel(t *testing.T) { + channel, config := prepareCurrentIdentityChannelV1(t) + ctx, cancel := context.WithTimeout(t.Context(), time.Second) + defer cancel() + claimResult := make(chan struct { + connection *ControllerConnectionV1 + err error + }, 1) + go func() { + connection, err := channel.Claim(ctx) + claimResult <- struct { + connection *ControllerConnectionV1 + err error + }{connection: connection, err: err} + }() + + client, err := DialSessionClientV1(ctx, channel.SocketPath()) + if err != nil { + t.Fatal(err) + } + defer client.Close() + if opened := client.Opened(); !reflect.DeepEqual(opened, config.Opened) { + t.Fatalf("opened = %#v, want %#v", opened, config.Opened) + } + claimed := <-claimResult + if claimed.err != nil { + t.Fatal(claimed.err) + } + defer claimed.connection.Close() +} + +func TestDialSessionClientV1RejectsInvalidSocketPaths(t *testing.T) { + ctx, cancel := context.WithTimeout(t.Context(), time.Second) + defer cancel() + for _, path := range []string{"", "relative/control.sock", "/tmp/session/../control.sock"} { + if _, err := DialSessionClientV1(ctx, path); err == nil || !strings.Contains(err.Error(), "absolute clean socket path") { + t.Fatalf("DialSessionClientV1(%q) error = %v", path, err) + } + } +} diff --git a/internal/controlledsession/client_test.go b/internal/controlledsession/client_test.go new file mode 100644 index 00000000..1477f5e0 --- /dev/null +++ b/internal/controlledsession/client_test.go @@ -0,0 +1,219 @@ +package controlledsession + +import ( + "context" + "net" + "reflect" + "strings" + "testing" + "time" +) + +func TestSessionClientV1ConsumesOpenedAndExchangesTypedFrames(t *testing.T) { + server, controller := net.Pipe() + defer server.Close() + opened := testOpenedV1() + wantEvent := EventV1{Kind: EventDiagnosticV1, Diagnostic: &DiagnosticV1{Code: "ready", Message: "controller ready"}} + wantRequest := RequestV1{Kind: RequestInputV1, Bytes: []byte{0, 3, 0xff}} + type serverResultV1 struct { + request RequestV1 + err error + } + serverResult := make(chan serverResultV1, 1) + go func() { + if err := WriteEventV1(server, EventV1{Kind: EventOpenedV1, Opened: &opened}); err != nil { + serverResult <- serverResultV1{err: err} + return + } + if err := WriteEventV1(server, EventV1{Kind: EventReadyV1}); err != nil { + serverResult <- serverResultV1{err: err} + return + } + if err := WriteEventV1(server, wantEvent); err != nil { + serverResult <- serverResultV1{err: err} + return + } + request, err := ReadRequestV1(server) + serverResult <- serverResultV1{request: request, err: err} + }() + + ctx, cancel := context.WithTimeout(t.Context(), time.Second) + defer cancel() + client, err := newSessionClientV1(ctx, controller) + if err != nil { + t.Fatal(err) + } + defer client.Close() + gotOpened := client.Opened() + if !reflect.DeepEqual(gotOpened, opened) { + t.Fatalf("opened = %#v, want %#v", gotOpened, opened) + } + gotOpened.Endpoints[0].Host = "changed" + if client.Opened().Endpoints[0].Host != WorkloadEndpointHostV1 { + t.Fatal("opened endpoint copy mutated the client") + } + if client.Ready() { + t.Fatal("client became ready from opened") + } + if err := client.WriteRequest(ctx, wantRequest); err == nil || !strings.Contains(err.Error(), "not ready") { + t.Fatalf("pre-ready request error = %v", err) + } + ready, err := client.ReadEvent(ctx) + if err != nil || ready.Kind != EventReadyV1 || !client.Ready() { + t.Fatalf("ready event = %#v, error = %v, ready = %t", ready, err, client.Ready()) + } + event, err := client.ReadEvent(ctx) + if err != nil || !reflect.DeepEqual(event, wantEvent) { + t.Fatalf("event = %#v, error = %v", event, err) + } + if err := client.WriteRequest(ctx, wantRequest); err != nil { + t.Fatal(err) + } + result := <-serverResult + if result.err != nil { + t.Fatal(result.err) + } + if !reflect.DeepEqual(result.request, wantRequest) { + t.Fatalf("request = %#v, want %#v", result.request, wantRequest) + } +} + +func TestSessionClientV1RejectsRepeatedReady(t *testing.T) { + server, controller := net.Pipe() + defer server.Close() + opened := testOpenedV1() + go func() { + _ = WriteEventV1(server, EventV1{Kind: EventOpenedV1, Opened: &opened}) + _ = WriteEventV1(server, EventV1{Kind: EventReadyV1}) + _ = WriteEventV1(server, EventV1{Kind: EventReadyV1}) + }() + ctx, cancel := context.WithTimeout(t.Context(), time.Second) + defer cancel() + client, err := newSessionClientV1(ctx, controller) + if err != nil { + t.Fatal(err) + } + defer client.Close() + if event, err := client.ReadEvent(ctx); err != nil || event.Kind != EventReadyV1 { + t.Fatalf("first ready event = %#v, error = %v", event, err) + } + if _, err := client.ReadEvent(ctx); err == nil || !strings.Contains(err.Error(), "ready may appear only once") { + t.Fatalf("repeated ready error = %v", err) + } +} + +func TestSessionClientV1AcknowledgesStartupFailureBeforeReady(t *testing.T) { + server, controller := net.Pipe() + defer server.Close() + opened := testOpenedV1() + terminated := EventV1{Kind: EventTerminatedV1, Terminated: &ResultV1{ + Cause: CauseStartupFailureV1, + WorkloadStatus: ProcessStatusV1{Kind: ProcessStatusUnknownV1}, + WorkloadOutputFinalizationStatus: WorkloadOutputFinalizationStatusV1{ + Kind: WorkloadOutputFinalizationDrainedV1, + }, + RuntimeObservationStatus: RuntimeObservationStatusV1{Kind: RuntimeObservationMaintainedV1}, + ControllerFinalizationStatus: ControllerFinalizationStatusV1{ + Kind: ControllerFinalizationStartupFailedV1, + }, + CleanupStatus: CleanupStatusV1{Kind: CleanupStatusSucceededV1}, + RecoveryAction: RecoveryNoneV1, + }} + type serverResultV1 struct { + request RequestV1 + err error + } + serverResult := make(chan serverResultV1, 1) + go func() { + if err := WriteEventV1(server, EventV1{Kind: EventOpenedV1, Opened: &opened}); err != nil { + serverResult <- serverResultV1{err: err} + return + } + if err := WriteEventV1(server, terminated); err != nil { + serverResult <- serverResultV1{err: err} + return + } + request, err := ReadRequestV1(server) + serverResult <- serverResultV1{request: request, err: err} + }() + + ctx, cancel := context.WithTimeout(t.Context(), time.Second) + defer cancel() + client, err := newSessionClientV1(ctx, controller) + if err != nil { + t.Fatal(err) + } + defer client.Close() + if err := client.WriteRequest(ctx, RequestV1{Kind: RequestAcknowledgeTerminatedV1}); err == nil || !strings.Contains(err.Error(), "not been received") { + t.Fatalf("pre-terminated acknowledgement error = %v", err) + } + event, err := client.ReadEvent(ctx) + if err != nil || !reflect.DeepEqual(event, terminated) { + t.Fatalf("terminated event = %#v, error = %v", event, err) + } + if client.Ready() { + t.Fatal("client became ready after startup failure") + } + if err := client.WriteRequest(ctx, RequestV1{Kind: RequestInputV1, Bytes: []byte("late")}); err == nil || !strings.Contains(err.Error(), "terminated") { + t.Fatalf("post-terminated input error = %v", err) + } + wantRequest := RequestV1{Kind: RequestAcknowledgeTerminatedV1} + if err := client.WriteRequest(ctx, wantRequest); err != nil { + t.Fatal(err) + } + result := <-serverResult + if result.err != nil { + t.Fatal(result.err) + } + if !reflect.DeepEqual(result.request, wantRequest) { + t.Fatalf("request = %#v, want %#v", result.request, wantRequest) + } +} + +func TestSessionClientV1RequiresOpenedFirstAndRejectsRepeatedOpened(t *testing.T) { + for _, test := range []struct { + name string + first EventV1 + second *EventV1 + want string + }{ + {name: "first", first: EventV1{Kind: EventDiagnosticV1, Diagnostic: &DiagnosticV1{Code: "bad", Message: "not opened"}}, want: "first event must be opened"}, + {name: "repeated", first: EventV1{Kind: EventOpenedV1, Opened: pointerToOpenedV1(testOpenedV1())}, second: &EventV1{Kind: EventOpenedV1, Opened: pointerToOpenedV1(testOpenedV1())}, want: "only once"}, + } { + t.Run(test.name, func(t *testing.T) { + server, controller := net.Pipe() + defer server.Close() + go func() { + _ = WriteEventV1(server, test.first) + if test.second != nil { + _ = WriteEventV1(server, *test.second) + } + }() + ctx, cancel := context.WithTimeout(t.Context(), time.Second) + defer cancel() + client, err := newSessionClientV1(ctx, controller) + if test.second == nil { + if err == nil || !strings.Contains(err.Error(), test.want) { + t.Fatalf("error = %v, want containing %q", err, test.want) + } + return + } + if err != nil { + t.Fatal(err) + } + defer client.Close() + if _, err := client.ReadEvent(ctx); err == nil || !strings.Contains(err.Error(), test.want) { + t.Fatalf("error = %v, want containing %q", err, test.want) + } + }) + } +} + +func testOpenedV1() OpenedV1 { + return OpenedV1{ + Authorization: testAuthorizationV1(), Endpoints: testEndpointsV1(), Columns: 80, Rows: 24, + OutputFinalizationTimeoutMilliseconds: DefaultOutputFinalizationTimeoutMillisecondsV1, + } +} + +func pointerToOpenedV1(value OpenedV1) *OpenedV1 { return &value } diff --git a/internal/controlledsession/client_unsupported.go b/internal/controlledsession/client_unsupported.go new file mode 100644 index 00000000..8935df51 --- /dev/null +++ b/internal/controlledsession/client_unsupported.go @@ -0,0 +1,12 @@ +//go:build !linux + +package controlledsession + +import ( + "context" + "fmt" +) + +func DialSessionClientV1(context.Context, string) (*SessionClientV1, error) { + return nil, fmt.Errorf("controlled-session client requires Linux") +} diff --git a/internal/controlledsession/protocol.go b/internal/controlledsession/protocol.go index 1255b88c..b6db0ed9 100644 --- a/internal/controlledsession/protocol.go +++ b/internal/controlledsession/protocol.go @@ -13,11 +13,7 @@ import ( ) const ( - // ProtocolVersionV1 is the original wire version whose opened event did - // not contain endpoint coordinates. It remains reserved so the expanded - // opened event cannot be mistaken for the old strict JSON schema. ProtocolVersionV1 = 1 - ProtocolVersionV2 = 2 MaxFramePayloadV1 = 1 << 20 DefaultOutputFinalizationTimeoutMillisecondsV1 uint32 = 30_000 frameHeaderSizeV1 = 10 @@ -27,9 +23,9 @@ const ( var frameMagicV1 = [4]byte{'R', 'P', 'S', 'N'} var protocolCodePatternV1 = regexp.MustCompile(`^[a-z][a-z0-9]*(?:_[a-z0-9]+)*$`) -var endpointSchemePatternV2 = regexp.MustCompile(`^[A-Za-z][A-Za-z0-9+.-]*$`) +var endpointSchemePatternV1 = regexp.MustCompile(`^[A-Za-z][A-Za-z0-9+.-]*$`) -const WorkloadEndpointHostV2 = "workload" +const WorkloadEndpointHostV1 = "workload" type RequestKindV1 string @@ -52,6 +48,7 @@ type EventKindV1 string const ( EventOpenedV1 EventKindV1 = "opened" + EventReadyV1 EventKindV1 = "ready" EventOutputV1 EventKindV1 = "output" EventWorkloadExitV1 EventKindV1 = "workload-exit" EventTerminatingV1 EventKindV1 = "terminating" @@ -60,20 +57,20 @@ const ( EventTerminatedV1 EventKindV1 = "terminated" ) -// OpenedV2 expands the original opened payload with immutable session-local -// endpoint coordinates and therefore requires protocol wire version 2. -type OpenedV2 struct { +// OpenedV1 carries the immutable authorization, session-local endpoint +// coordinates, terminal dimensions, and output-finalization deadline. +type OpenedV1 struct { Authorization AuthorizationV1 `json:"authorization"` - Endpoints []EndpointV2 `json:"endpoints"` + Endpoints []EndpointV1 `json:"endpoints"` Columns uint32 `json:"columns"` Rows uint32 `json:"rows"` OutputFinalizationTimeoutMilliseconds uint32 `json:"output_finalization_timeout_milliseconds"` } -// EndpointV2 is one immutable session-local coordinate granted to the +// EndpointV1 is one immutable session-local coordinate granted to the // controller. Traffic uses native TCP; only these coordinates cross the // private session channel. -type EndpointV2 struct { +type EndpointV1 struct { ID string `json:"id"` Scheme string `json:"scheme"` Host string `json:"host"` @@ -101,7 +98,7 @@ type WorkloadOutputsFinalizedV1 struct { type EventV1 struct { Kind EventKindV1 Bytes []byte - Opened *OpenedV2 + Opened *OpenedV1 WorkloadExit *WorkloadExitV1 Terminating *TerminatingV1 Diagnostic *DiagnosticV1 @@ -127,6 +124,7 @@ const ( wireEventDiagnosticV1 wireKindV1 = 0x85 wireEventTerminatedV1 wireKindV1 = 0x86 wireEventWorkloadOutputsFinalizedV1 wireKindV1 = 0x87 + wireEventReadyV1 wireKindV1 = 0x88 ) func ValidateRequestV1(request RequestV1) error { @@ -158,7 +156,7 @@ func ValidateEventV1(event EventV1) error { if err := ValidateAuthorizationV1(event.Opened.Authorization); err != nil { return fmt.Errorf("controlled-session opened event: %w", err) } - if err := validateOpenedEndpointsV2(event.Opened.Endpoints, event.Opened.Authorization.EndpointIDs); err != nil { + if err := validateOpenedEndpointsV1(event.Opened.Endpoints, event.Opened.Authorization.EndpointIDs); err != nil { return err } if !validDimensionsV1(event.Opened.Columns, event.Opened.Rows) { @@ -167,6 +165,10 @@ func ValidateEventV1(event EventV1) error { if event.Opened.OutputFinalizationTimeoutMilliseconds == 0 { return fmt.Errorf("controlled-session opened event requires a finite output-finalization timeout") } + case EventReadyV1: + if eventPayloadCountV1(event) != 0 { + return fmt.Errorf("controlled-session ready event must not contain a payload") + } case EventOutputV1: if event.Bytes == nil || eventPayloadCountV1(event) != 1 { return fmt.Errorf("controlled-session output event must contain only a byte sequence") @@ -218,7 +220,7 @@ func ValidateEventV1(event EventV1) error { return nil } -func validateOpenedEndpointsV2(endpoints []EndpointV2, authorizedIDs []string) error { +func validateOpenedEndpointsV1(endpoints []EndpointV1, authorizedIDs []string) error { if endpoints == nil { return fmt.Errorf("controlled-session opened endpoints must use an array") } @@ -229,22 +231,22 @@ func validateOpenedEndpointsV2(endpoints []EndpointV2, authorizedIDs []string) e if endpoint.ID != authorizedIDs[index] { return fmt.Errorf("controlled-session opened endpoint %d must match authorized endpoint ID %q", index, authorizedIDs[index]) } - if err := ValidateEndpointV2(endpoint); err != nil { + if err := ValidateEndpointV1(endpoint); err != nil { return err } } return nil } -func ValidateEndpointV2(endpoint EndpointV2) error { +func ValidateEndpointV1(endpoint EndpointV1) error { if err := endpointname.Validate(endpoint.ID); err != nil { return fmt.Errorf("controlled-session endpoint ID %q: %w", endpoint.ID, err) } - if !endpointSchemePatternV2.MatchString(endpoint.Scheme) { + if !endpointSchemePatternV1.MatchString(endpoint.Scheme) { return fmt.Errorf("controlled-session endpoint %q scheme must use URI-scheme syntax", endpoint.ID) } - if endpoint.Host != WorkloadEndpointHostV2 { - return fmt.Errorf("controlled-session endpoint %q host must be %q", endpoint.ID, WorkloadEndpointHostV2) + if endpoint.Host != WorkloadEndpointHostV1 { + return fmt.Errorf("controlled-session endpoint %q host must be %q", endpoint.ID, WorkloadEndpointHostV1) } if endpoint.Port == 0 || endpoint.Port > 65535 { return fmt.Errorf("controlled-session endpoint %q port must be between 1 and 65535", endpoint.ID) @@ -252,16 +254,16 @@ func ValidateEndpointV2(endpoint EndpointV2) error { return nil } -func WriteRequestV2(writer io.Writer, request RequestV1) error { +func WriteRequestV1(writer io.Writer, request RequestV1) error { if err := ValidateRequestV1(request); err != nil { return err } kind, payload := encodeRequestPayloadV1(request) - return writeFrameV2(writer, kind, payload) + return writeFrameV1(writer, kind, payload) } -func ReadRequestV2(reader io.Reader) (RequestV1, error) { - kind, payload, err := readFrameV2(reader) +func ReadRequestV1(reader io.Reader) (RequestV1, error) { + kind, payload, err := readFrameV1(reader) if err != nil { return RequestV1{}, err } @@ -275,7 +277,7 @@ func ReadRequestV2(reader io.Reader) (RequestV1, error) { return request, nil } -func WriteEventV2(writer io.Writer, event EventV1) error { +func WriteEventV1(writer io.Writer, event EventV1) error { if err := ValidateEventV1(event); err != nil { return err } @@ -283,11 +285,11 @@ func WriteEventV2(writer io.Writer, event EventV1) error { if err != nil { return err } - return writeFrameV2(writer, kind, payload) + return writeFrameV1(writer, kind, payload) } -func ReadEventV2(reader io.Reader) (EventV1, error) { - kind, payload, err := readFrameV2(reader) +func ReadEventV1(reader io.Reader) (EventV1, error) { + kind, payload, err := readFrameV1(reader) if err != nil { return EventV1{}, err } @@ -354,6 +356,8 @@ func encodeEventPayloadV1(event EventV1) (wireKindV1, []byte, error) { switch event.Kind { case EventOpenedV1: return marshalEventPayloadV1(wireEventOpenedV1, event.Opened) + case EventReadyV1: + return wireEventReadyV1, nil, nil case EventOutputV1: return wireEventOutputV1, event.Bytes, nil case EventWorkloadExitV1: @@ -382,8 +386,13 @@ func marshalEventPayloadV1(kind wireKindV1, value any) (wireKindV1, []byte, erro func decodeEventPayloadV1(kind wireKindV1, payload []byte) (EventV1, error) { switch kind { case wireEventOpenedV1: - value := new(OpenedV2) + value := new(OpenedV1) return EventV1{Kind: EventOpenedV1, Opened: value}, decodeStrictJSONV1("opened event", payload, value) + case wireEventReadyV1: + if len(payload) != 0 { + return EventV1{}, fmt.Errorf("controlled-session ready frame must not contain a payload") + } + return EventV1{Kind: EventReadyV1}, nil case wireEventOutputV1: return EventV1{Kind: EventOutputV1, Bytes: payload}, nil case wireEventWorkloadExitV1: @@ -424,13 +433,13 @@ func eventPayloadCountV1(event EventV1) int { return count } -func writeFrameV2(writer io.Writer, kind wireKindV1, payload []byte) error { +func writeFrameV1(writer io.Writer, kind wireKindV1, payload []byte) error { if len(payload) > MaxFramePayloadV1 { return fmt.Errorf("controlled-session frame payload exceeds %d bytes", MaxFramePayloadV1) } header := make([]byte, frameHeaderSizeV1) copy(header[:4], frameMagicV1[:]) - header[4] = ProtocolVersionV2 + header[4] = ProtocolVersionV1 header[5] = byte(kind) binary.BigEndian.PutUint32(header[6:], uint32(len(payload))) if err := writeAllV1(writer, header); err != nil { @@ -442,7 +451,7 @@ func writeFrameV2(writer io.Writer, kind wireKindV1, payload []byte) error { return nil } -func readFrameV2(reader io.Reader) (wireKindV1, []byte, error) { +func readFrameV1(reader io.Reader) (wireKindV1, []byte, error) { header := make([]byte, frameHeaderSizeV1) if _, err := io.ReadFull(reader, header); err != nil { return 0, nil, fmt.Errorf("read controlled-session frame header: %w", err) @@ -450,7 +459,7 @@ func readFrameV2(reader io.Reader) (wireKindV1, []byte, error) { if !bytes.Equal(header[:4], frameMagicV1[:]) { return 0, nil, fmt.Errorf("controlled-session frame magic is invalid") } - if header[4] != ProtocolVersionV2 { + if header[4] != ProtocolVersionV1 { return 0, nil, fmt.Errorf("controlled-session protocol version %d is unsupported", header[4]) } length := binary.BigEndian.Uint32(header[6:]) diff --git a/internal/controlledsession/protocol_test.go b/internal/controlledsession/protocol_test.go index 1a445ac7..e84093c3 100644 --- a/internal/controlledsession/protocol_test.go +++ b/internal/controlledsession/protocol_test.go @@ -9,7 +9,7 @@ import ( "testing" ) -func TestProtocolV2RequestsRoundTripBinarySafeFrames(t *testing.T) { +func TestProtocolV1RequestsRoundTripBinarySafeFrames(t *testing.T) { requests := []RequestV1{ {Kind: RequestInputV1, Bytes: []byte{0, 3, '\n', 0xff}}, {Kind: RequestResizeV1, Columns: 120, Rows: 40}, @@ -19,29 +19,30 @@ func TestProtocolV2RequestsRoundTripBinarySafeFrames(t *testing.T) { } var stream bytes.Buffer for _, request := range requests { - if err := WriteRequestV2(&stream, request); err != nil { - t.Fatalf("WriteRequestV2(%s) error = %v", request.Kind, err) + if err := WriteRequestV1(&stream, request); err != nil { + t.Fatalf("WriteRequestV1(%s) error = %v", request.Kind, err) } } for _, want := range requests { - got, err := ReadRequestV2(&stream) + got, err := ReadRequestV1(&stream) if err != nil { - t.Fatalf("ReadRequestV2(%s) error = %v", want.Kind, err) + t.Fatalf("ReadRequestV1(%s) error = %v", want.Kind, err) } if !reflect.DeepEqual(got, want) { - t.Fatalf("ReadRequestV2() = %#v, want %#v", got, want) + t.Fatalf("ReadRequestV1() = %#v, want %#v", got, want) } } } -func TestProtocolV2EventsRoundTripStrictTypedFrames(t *testing.T) { +func TestProtocolV1EventsRoundTripStrictTypedFrames(t *testing.T) { code := 0 authorization := testAuthorizationV1() events := []EventV1{ - {Kind: EventOpenedV1, Opened: &OpenedV2{ + {Kind: EventOpenedV1, Opened: &OpenedV1{ Authorization: authorization, Endpoints: testEndpointsV1(), Columns: 80, Rows: 24, OutputFinalizationTimeoutMilliseconds: DefaultOutputFinalizationTimeoutMillisecondsV1, }}, + {Kind: EventReadyV1}, {Kind: EventOutputV1, Bytes: []byte{0, '\n', 0xff}}, {Kind: EventWorkloadExitV1, WorkloadExit: &WorkloadExitV1{Status: ProcessStatusV1{Kind: ProcessStatusExitedV1, Code: &code}}}, {Kind: EventTerminatingV1, Terminating: &TerminatingV1{Cause: CauseWorkloadExitV1}}, @@ -65,45 +66,45 @@ func TestProtocolV2EventsRoundTripStrictTypedFrames(t *testing.T) { } var stream bytes.Buffer for _, event := range events { - if err := WriteEventV2(&stream, event); err != nil { - t.Fatalf("WriteEventV2(%s) error = %v", event.Kind, err) + if err := WriteEventV1(&stream, event); err != nil { + t.Fatalf("WriteEventV1(%s) error = %v", event.Kind, err) } } for _, want := range events { - got, err := ReadEventV2(&stream) + got, err := ReadEventV1(&stream) if err != nil { - t.Fatalf("ReadEventV2(%s) error = %v", want.Kind, err) + t.Fatalf("ReadEventV1(%s) error = %v", want.Kind, err) } if !reflect.DeepEqual(got, want) { - t.Fatalf("ReadEventV2() = %#v, want %#v", got, want) + t.Fatalf("ReadEventV1() = %#v, want %#v", got, want) } } } -func TestProtocolV2RejectsWrongDirectionAndUnboundedInput(t *testing.T) { +func TestProtocolV1RejectsWrongDirectionAndUnboundedInput(t *testing.T) { var event bytes.Buffer - if err := WriteEventV2(&event, EventV1{Kind: EventOutputV1, Bytes: []byte("hello")}); err != nil { + if err := WriteEventV1(&event, EventV1{Kind: EventOutputV1, Bytes: []byte("hello")}); err != nil { t.Fatal(err) } - if _, err := ReadRequestV2(&event); err == nil || !strings.Contains(err.Error(), "not a controller request") { - t.Fatalf("ReadRequestV2(event) error = %v", err) + if _, err := ReadRequestV1(&event); err == nil || !strings.Contains(err.Error(), "not a controller request") { + t.Fatalf("ReadRequestV1(event) error = %v", err) } header := make([]byte, frameHeaderSizeV1) copy(header, frameMagicV1[:]) - header[4] = ProtocolVersionV2 + header[4] = ProtocolVersionV1 header[5] = byte(wireRequestInputV1) binary.BigEndian.PutUint32(header[6:], MaxFramePayloadV1+1) - if _, err := ReadRequestV2(bytes.NewReader(header)); err == nil || !strings.Contains(err.Error(), "exceeds") { - t.Fatalf("ReadRequestV2(oversized) error = %v", err) + if _, err := ReadRequestV1(bytes.NewReader(header)); err == nil || !strings.Contains(err.Error(), "exceeds") { + t.Fatalf("ReadRequestV1(oversized) error = %v", err) } } -func TestProtocolV2RejectsBadMagicVersionTruncationAndUnknownJSON(t *testing.T) { +func TestProtocolV1RejectsBadMagicVersionTruncationAndUnknownJSON(t *testing.T) { validHeader := func() []byte { header := make([]byte, frameHeaderSizeV1) copy(header, frameMagicV1[:]) - header[4] = ProtocolVersionV2 + header[4] = ProtocolVersionV1 header[5] = byte(wireRequestCompleteV1) return header } @@ -113,7 +114,7 @@ func TestProtocolV2RejectsBadMagicVersionTruncationAndUnknownJSON(t *testing.T) want string }{ {name: "bad magic", data: append([]byte("NOPE"), validHeader()[4:]...), want: "magic"}, - {name: "old version", data: func() []byte { value := validHeader(); value[4] = ProtocolVersionV1; return value }(), want: "version 1"}, + {name: "unsupported version", data: func() []byte { value := validHeader(); value[4] = 2; return value }(), want: "version 2"}, {name: "short header", data: validHeader()[:5], want: "frame header"}, {name: "short payload", data: func() []byte { value := validHeader() @@ -124,55 +125,55 @@ func TestProtocolV2RejectsBadMagicVersionTruncationAndUnknownJSON(t *testing.T) } for _, test := range tests { t.Run(test.name, func(t *testing.T) { - if _, err := ReadRequestV2(bytes.NewReader(test.data)); err == nil || !strings.Contains(err.Error(), test.want) { - t.Fatalf("ReadRequestV2() error = %v, want containing %q", err, test.want) + if _, err := ReadRequestV1(bytes.NewReader(test.data)); err == nil || !strings.Contains(err.Error(), test.want) { + t.Fatalf("ReadRequestV1() error = %v, want containing %q", err, test.want) } }) } payload := []byte(`{"code":"bad","message":"failure","extra":true}`) var framed bytes.Buffer - if err := writeFrameV2(&framed, wireEventDiagnosticV1, payload); err != nil { + if err := writeFrameV1(&framed, wireEventDiagnosticV1, payload); err != nil { t.Fatal(err) } - if _, err := ReadEventV2(&framed); err == nil || !strings.Contains(err.Error(), "unknown field") { - t.Fatalf("ReadEventV2(unknown JSON) error = %v", err) + if _, err := ReadEventV1(&framed); err == nil || !strings.Contains(err.Error(), "unknown field") { + t.Fatalf("ReadEventV1(unknown JSON) error = %v", err) } duplicate := []byte(`{"code":"first","code":"second","message":"failure"}`) framed.Reset() - if err := writeFrameV2(&framed, wireEventDiagnosticV1, duplicate); err != nil { + if err := writeFrameV1(&framed, wireEventDiagnosticV1, duplicate); err != nil { t.Fatal(err) } - if _, err := ReadEventV2(&framed); err == nil || !strings.Contains(err.Error(), "repeats field") { - t.Fatalf("ReadEventV2(duplicate JSON) error = %v", err) + if _, err := ReadEventV1(&framed); err == nil || !strings.Contains(err.Error(), "repeats field") { + t.Fatalf("ReadEventV1(duplicate JSON) error = %v", err) } caseVariant := []byte(`{"Code":"bad","message":"failure"}`) framed.Reset() - if err := writeFrameV2(&framed, wireEventDiagnosticV1, caseVariant); err != nil { + if err := writeFrameV1(&framed, wireEventDiagnosticV1, caseVariant); err != nil { t.Fatal(err) } - if _, err := ReadEventV2(&framed); err == nil || !strings.Contains(err.Error(), "lowercase ASCII snake_case") { - t.Fatalf("ReadEventV2(case-variant JSON) error = %v", err) + if _, err := ReadEventV1(&framed); err == nil || !strings.Contains(err.Error(), "lowercase ASCII snake_case") { + t.Fatalf("ReadEventV1(case-variant JSON) error = %v", err) } caseVariantDuplicate := []byte(`{"code":"first","Code":"second","message":"failure"}`) framed.Reset() - if err := writeFrameV2(&framed, wireEventDiagnosticV1, caseVariantDuplicate); err != nil { + if err := writeFrameV1(&framed, wireEventDiagnosticV1, caseVariantDuplicate); err != nil { t.Fatal(err) } - if _, err := ReadEventV2(&framed); err == nil || !strings.Contains(err.Error(), "lowercase ASCII snake_case") { - t.Fatalf("ReadEventV2(case-variant duplicate JSON) error = %v", err) + if _, err := ReadEventV1(&framed); err == nil || !strings.Contains(err.Error(), "lowercase ASCII snake_case") { + t.Fatalf("ReadEventV1(case-variant duplicate JSON) error = %v", err) } nestedCaseVariant := []byte(`{"cause":"workload-exit","workload_status":{"Kind":"exited","code":0},"workload_output_finalization_status":{"kind":"drained"},"runtime_observation_status":{"kind":"maintained"},"controller_finalization_status":{"kind":"completed"},"cleanup_status":{"kind":"succeeded"},"recovery_action":"none"}`) framed.Reset() - if err := writeFrameV2(&framed, wireEventTerminatedV1, nestedCaseVariant); err != nil { + if err := writeFrameV1(&framed, wireEventTerminatedV1, nestedCaseVariant); err != nil { t.Fatal(err) } - if _, err := ReadEventV2(&framed); err == nil || !strings.Contains(err.Error(), "lowercase ASCII snake_case") { - t.Fatalf("ReadEventV2(nested case-variant JSON) error = %v", err) + if _, err := ReadEventV1(&framed); err == nil || !strings.Contains(err.Error(), "lowercase ASCII snake_case") { + t.Fatalf("ReadEventV1(nested case-variant JSON) error = %v", err) } } @@ -194,6 +195,7 @@ func TestWireKindAssignmentsV1AreStable(t *testing.T) { {name: "event diagnostic", got: wireEventDiagnosticV1, want: 0x85}, {name: "event terminated", got: wireEventTerminatedV1, want: 0x86}, {name: "event workload outputs finalized", got: wireEventWorkloadOutputsFinalizedV1, want: 0x87}, + {name: "event ready", got: wireEventReadyV1, want: 0x88}, } for _, test := range tests { t.Run(test.name, func(t *testing.T) { @@ -204,46 +206,46 @@ func TestWireKindAssignmentsV1AreStable(t *testing.T) { } var framed bytes.Buffer - if err := WriteRequestV2(&framed, RequestV1{Kind: RequestCompleteV1}); err != nil { + if err := WriteRequestV1(&framed, RequestV1{Kind: RequestCompleteV1}); err != nil { t.Fatal(err) } - want := []byte{'R', 'P', 'S', 'N', ProtocolVersionV2, 0x04, 0, 0, 0, 0} + want := []byte{'R', 'P', 'S', 'N', ProtocolVersionV1, 0x04, 0, 0, 0, 0} if !bytes.Equal(framed.Bytes(), want) { t.Fatalf("complete frame = %x, want %x", framed.Bytes(), want) } } -func TestProtocolV2ExpandedOpenedDoesNotReuseV1WireVersion(t *testing.T) { +func TestProtocolV1OpenedUsesInitialWireVersion(t *testing.T) { authorization := testAuthorizationV1() var framed bytes.Buffer - if err := WriteEventV2(&framed, EventV1{Kind: EventOpenedV1, Opened: &OpenedV2{ + if err := WriteEventV1(&framed, EventV1{Kind: EventOpenedV1, Opened: &OpenedV1{ Authorization: authorization, Endpoints: testEndpointsV1(), Columns: 80, Rows: 24, OutputFinalizationTimeoutMilliseconds: DefaultOutputFinalizationTimeoutMillisecondsV1, }}); err != nil { t.Fatal(err) } - if got := framed.Bytes()[4]; got != ProtocolVersionV2 { - t.Fatalf("expanded opened protocol version = %d, want %d", got, ProtocolVersionV2) + if got := framed.Bytes()[4]; got != ProtocolVersionV1 { + t.Fatalf("expanded opened protocol version = %d, want %d", got, ProtocolVersionV1) } - oldVersion := append([]byte(nil), framed.Bytes()...) - oldVersion[4] = ProtocolVersionV1 - if _, err := ReadEventV2(bytes.NewReader(oldVersion)); err == nil || !strings.Contains(err.Error(), "version 1") { - t.Fatalf("ReadEventV2(v1 expanded opened) error = %v", err) + unsupportedVersion := append([]byte(nil), framed.Bytes()...) + unsupportedVersion[4] = 2 + if _, err := ReadEventV1(bytes.NewReader(unsupportedVersion)); err == nil || !strings.Contains(err.Error(), "version 2") { + t.Fatalf("ReadEventV1(unsupported version) error = %v", err) } } -func TestValidateEndpointV2PreservesValidDeclaredURISchemes(t *testing.T) { +func TestValidateEndpointV1PreservesValidDeclaredURISchemes(t *testing.T) { for _, scheme := range []string{"HTTP", "grpc+unix", "A" + strings.Repeat("b", 64)} { - endpoint := EndpointV2{ID: "browser", Scheme: scheme, Host: WorkloadEndpointHostV2, Port: 8080} - if err := ValidateEndpointV2(endpoint); err != nil { - t.Fatalf("ValidateEndpointV2(%q) error = %v", scheme, err) + endpoint := EndpointV1{ID: "browser", Scheme: scheme, Host: WorkloadEndpointHostV1, Port: 8080} + if err := ValidateEndpointV1(endpoint); err != nil { + t.Fatalf("ValidateEndpointV1(%q) error = %v", scheme, err) } } for _, scheme := range []string{"", "1http", "http scheme"} { - endpoint := EndpointV2{ID: "browser", Scheme: scheme, Host: WorkloadEndpointHostV2, Port: 8080} - if err := ValidateEndpointV2(endpoint); err == nil || !strings.Contains(err.Error(), "URI-scheme") { - t.Fatalf("ValidateEndpointV2(%q) error = %v", scheme, err) + endpoint := EndpointV1{ID: "browser", Scheme: scheme, Host: WorkloadEndpointHostV1, Port: 8080} + if err := ValidateEndpointV1(endpoint); err == nil || !strings.Contains(err.Error(), "URI-scheme") { + t.Fatalf("ValidateEndpointV1(%q) error = %v", scheme, err) } } } @@ -251,7 +253,7 @@ func TestValidateEndpointV2PreservesValidDeclaredURISchemes(t *testing.T) { func TestValidateEventV1RejectsInvalidOutputFinalizationOutcomes(t *testing.T) { code := 0 tests := []EventV1{ - {Kind: EventOpenedV1, Opened: &OpenedV2{Authorization: testAuthorizationV1(), Endpoints: testEndpointsV1(), Columns: 80, Rows: 24}}, + {Kind: EventOpenedV1, Opened: &OpenedV1{Authorization: testAuthorizationV1(), Endpoints: testEndpointsV1(), Columns: 80, Rows: 24}}, {Kind: EventWorkloadOutputsFinalizedV1}, {Kind: EventWorkloadOutputsFinalizedV1, WorkloadOutputsFinalized: &WorkloadOutputsFinalizedV1{Status: WorkloadOutputFinalizationDrainedV1, Reason: "unexpected"}}, {Kind: EventWorkloadOutputsFinalizedV1, WorkloadOutputsFinalized: &WorkloadOutputsFinalizedV1{Status: WorkloadOutputFinalizationFailedV1}}, @@ -268,7 +270,7 @@ func TestValidateEventV1RejectsInvalidOutputFinalizationOutcomes(t *testing.T) { } } -func TestReadEventV2RejectsWorkloadExitWithoutStatus(t *testing.T) { +func TestReadEventV1RejectsWorkloadExitWithoutStatus(t *testing.T) { result := ResultV1{ Cause: CauseWorkloadExitV1, WorkloadStatus: ProcessStatusV1{Kind: ProcessStatusUnknownV1}, @@ -282,15 +284,15 @@ func TestReadEventV2RejectsWorkloadExitWithoutStatus(t *testing.T) { t.Fatal(err) } var framed bytes.Buffer - if err := writeFrameV2(&framed, wireEventTerminatedV1, payload); err != nil { + if err := writeFrameV1(&framed, wireEventTerminatedV1, payload); err != nil { t.Fatal(err) } - if _, err := ReadEventV2(&framed); err == nil || !strings.Contains(err.Error(), "workload-exit termination requires a known workload status") { - t.Fatalf("ReadEventV2(workload exit without status) error = %v", err) + if _, err := ReadEventV1(&framed); err == nil || !strings.Contains(err.Error(), "workload-exit termination requires a known workload status") { + t.Fatalf("ReadEventV1(workload exit without status) error = %v", err) } } -func TestReadEventV2RejectsContradictoryRuntimeObservationLossResults(t *testing.T) { +func TestReadEventV1RejectsContradictoryRuntimeObservationLossResults(t *testing.T) { code := 0 valid := ResultV1{ Cause: CauseRuntimeObservationLostV1, @@ -348,11 +350,11 @@ func TestReadEventV2RejectsContradictoryRuntimeObservationLossResults(t *testing t.Fatal(err) } var framed bytes.Buffer - if err := writeFrameV2(&framed, wireEventTerminatedV1, payload); err != nil { + if err := writeFrameV1(&framed, wireEventTerminatedV1, payload); err != nil { t.Fatal(err) } - if _, err := ReadEventV2(&framed); err == nil || !strings.Contains(err.Error(), test.wantError) { - t.Fatalf("ReadEventV2(contradictory terminated result) error = %v", err) + if _, err := ReadEventV1(&framed); err == nil || !strings.Contains(err.Error(), test.wantError) { + t.Fatalf("ReadEventV1(contradictory terminated result) error = %v", err) } }) } @@ -389,36 +391,52 @@ func TestValidateEventV1RejectsInvalidProtocolCodes(t *testing.T) { } } -func TestReadRequestV2RejectsAcknowledgeTerminatedPayload(t *testing.T) { +func TestReadRequestV1RejectsAcknowledgeTerminatedPayload(t *testing.T) { var framed bytes.Buffer - if err := writeFrameV2(&framed, wireRequestAcknowledgeTerminatedV1, []byte{1}); err != nil { + if err := writeFrameV1(&framed, wireRequestAcknowledgeTerminatedV1, []byte{1}); err != nil { t.Fatal(err) } - if _, err := ReadRequestV2(&framed); err == nil || !strings.Contains(err.Error(), "must not contain a payload") { - t.Fatalf("ReadRequestV2(acknowledge payload) error = %v", err) + if _, err := ReadRequestV1(&framed); err == nil || !strings.Contains(err.Error(), "must not contain a payload") { + t.Fatalf("ReadRequestV1(acknowledge payload) error = %v", err) } } -func FuzzReadRequestV2(f *testing.F) { +func TestProtocolV1ReadyIsPayloadFree(t *testing.T) { + if err := ValidateEventV1(EventV1{ + Kind: EventReadyV1, Diagnostic: &DiagnosticV1{Code: "bad", Message: "unexpected payload"}, + }); err == nil || !strings.Contains(err.Error(), "must not contain a payload") { + t.Fatalf("ValidateEventV1(ready payload) error = %v", err) + } + + var framed bytes.Buffer + if err := writeFrameV1(&framed, wireEventReadyV1, []byte(`{}`)); err != nil { + t.Fatal(err) + } + if _, err := ReadEventV1(&framed); err == nil || !strings.Contains(err.Error(), "must not contain a payload") { + t.Fatalf("ReadEventV1(ready payload) error = %v", err) + } +} + +func FuzzReadRequestV1(f *testing.F) { f.Add([]byte{}) var valid bytes.Buffer - if err := WriteRequestV2(&valid, RequestV1{Kind: RequestInputV1, Bytes: []byte{0, 3, 0xff}}); err != nil { + if err := WriteRequestV1(&valid, RequestV1{Kind: RequestInputV1, Bytes: []byte{0, 3, 0xff}}); err != nil { f.Fatal(err) } f.Add(valid.Bytes()) f.Fuzz(func(t *testing.T, content []byte) { - _, _ = ReadRequestV2(bytes.NewReader(content)) + _, _ = ReadRequestV1(bytes.NewReader(content)) }) } -func FuzzReadEventV2(f *testing.F) { +func FuzzReadEventV1(f *testing.F) { f.Add([]byte{}) var valid bytes.Buffer - if err := WriteEventV2(&valid, EventV1{Kind: EventDiagnosticV1, Diagnostic: &DiagnosticV1{Code: "test", Message: "seed"}}); err != nil { + if err := WriteEventV1(&valid, EventV1{Kind: EventDiagnosticV1, Diagnostic: &DiagnosticV1{Code: "test", Message: "seed"}}); err != nil { f.Fatal(err) } f.Add(valid.Bytes()) f.Fuzz(func(t *testing.T, content []byte) { - _, _ = ReadEventV2(bytes.NewReader(content)) + _, _ = ReadEventV1(bytes.NewReader(content)) }) } diff --git a/internal/controlledsession/session_io.go b/internal/controlledsession/session_io.go index f37b3fe6..aa88f3f5 100644 --- a/internal/controlledsession/session_io.go +++ b/internal/controlledsession/session_io.go @@ -337,7 +337,7 @@ func (gate *sessionEventWriteGateV1) notifyLocked() { func isBridgeLifecycleEventV1(kind EventKindV1) bool { switch kind { - case EventWorkloadExitV1, EventTerminatingV1, EventDiagnosticV1, + case EventReadyV1, EventWorkloadExitV1, EventTerminatingV1, EventDiagnosticV1, EventWorkloadOutputsFinalizedV1, EventTerminatedV1: return true case EventOpenedV1, EventOutputV1: diff --git a/internal/controlledsession/session_io_test.go b/internal/controlledsession/session_io_test.go index 4c43a7ad..43406822 100644 --- a/internal/controlledsession/session_io_test.go +++ b/internal/controlledsession/session_io_test.go @@ -450,7 +450,7 @@ func TestSessionIOBridgeV1SendsOnlyLifecycleEventsAndValidatesConfiguration(t *t } for _, reserved := range []EventV1{ {Kind: EventOutputV1, Bytes: []byte("forged")}, - {Kind: EventOpenedV1, Opened: &OpenedV2{}}, + {Kind: EventOpenedV1, Opened: &OpenedV1{}}, } { if err := bridge.SendLifecycleEvent(t.Context(), reserved); err == nil { t.Fatalf("SendLifecycleEvent() accepted reserved event %q", reserved.Kind) diff --git a/internal/dockerdeploy/controlled_session_channel.go b/internal/dockerdeploy/controlled_session_channel.go index 78ac43a6..156a2fee 100644 --- a/internal/dockerdeploy/controlled_session_channel.go +++ b/internal/dockerdeploy/controlled_session_channel.go @@ -95,7 +95,7 @@ func controlledSessionPrivateChannelConfigV1(plan ControlledSessionExecutionPlan } return controlledsession.PrivateChannelConfigV1{ HostDirectory: plan.Channel.HostDirectory, - Opened: controlledsession.OpenedV2{ + Opened: controlledsession.OpenedV1{ Authorization: plan.Authorization, Endpoints: controlledSessionOpenedEndpointsV1(plan.Controller.SessionNetwork.Endpoints), Columns: uint32(columns), diff --git a/internal/dockerdeploy/controlled_session_channel_integration_test.go b/internal/dockerdeploy/controlled_session_channel_integration_test.go index c6236990..9d297ac3 100644 --- a/internal/dockerdeploy/controlled_session_channel_integration_test.go +++ b/internal/dockerdeploy/controlled_session_channel_integration_test.go @@ -35,8 +35,8 @@ func TestControlledSessionPrivateChannelDockerIntegration(t *testing.T) { authorization := testControlledSessionChannelAuthorizationV1(t, identity) channel, err := controlledsession.PreparePrivateChannelV1(controlledsession.PrivateChannelConfigV1{ HostDirectory: filepath.Join(shortControlledSessionChannelTestDirectoryV1(t), "session"), - Opened: controlledsession.OpenedV2{ - Authorization: authorization, Endpoints: []controlledsession.EndpointV2{}, Columns: 80, Rows: 24, + Opened: controlledsession.OpenedV1{ + Authorization: authorization, Endpoints: []controlledsession.EndpointV1{}, Columns: 80, Rows: 24, OutputFinalizationTimeoutMilliseconds: controlledsession.DefaultOutputFinalizationTimeoutMillisecondsV1, }, }) diff --git a/internal/dockerdeploy/controlled_session_network_integration_test.go b/internal/dockerdeploy/controlled_session_network_integration_test.go index 87d76f32..0a64a66e 100644 --- a/internal/dockerdeploy/controlled_session_network_integration_test.go +++ b/internal/dockerdeploy/controlled_session_network_integration_test.go @@ -31,8 +31,8 @@ func TestControlledSessionNetworkingDockerIntegration(t *testing.T) { image := buildControlledSessionControllerIntegrationImageV1(t, ctx) localPeer, publicPeer := createControlledSessionNetworkIsolationPeersV1(t, ctx, image) endpoints := []ControlledSessionEndpointPlanV1{ - {ID: "browser", Scheme: "http", Host: controlledsession.WorkloadEndpointHostV2, Port: "8080"}, - {ID: "socket", Scheme: "ws", Host: controlledsession.WorkloadEndpointHostV2, Port: "8080"}, + {ID: "browser", Scheme: "http", Host: controlledsession.WorkloadEndpointHostV1, Port: "8080"}, + {ID: "socket", Scheme: "ws", Host: controlledsession.WorkloadEndpointHostV1, Port: "8080"}, } t.Run("trusted controller bootstrap", func(t *testing.T) { proveControlledSessionNetworkControllerBootstrapV1(t, ctx, image, endpoints) diff --git a/internal/dockerdeploy/controlled_session_plan.go b/internal/dockerdeploy/controlled_session_plan.go index 141fa798..3f239934 100644 --- a/internal/dockerdeploy/controlled_session_plan.go +++ b/internal/dockerdeploy/controlled_session_plan.go @@ -24,7 +24,7 @@ const ( controlledSessionNetworkModeV1 = "none" controlledSessionOrdinaryNetworkModeV1 = "bridge" controlledSessionControllerAliasV1 = "controller" - controlledSessionWorkloadAliasV1 = controlledsession.WorkloadEndpointHostV2 + controlledSessionWorkloadAliasV1 = controlledsession.WorkloadEndpointHostV1 controlledSessionChannelRootV1 = "/run/reploy/session" controlledSessionChannelSocketNameV1 = controlledsession.PrivateChannelSocketNameV1 controlledSessionNetworkPolicyFileNameV1 = "network-prefixes" @@ -614,16 +614,16 @@ func validateControlledSessionEndpointPlanV1(endpoint ControlledSessionEndpointP if err != nil || port == 0 || strconv.FormatUint(port, 10) != endpoint.Port { return fmt.Errorf("controlled-session endpoint %q port must be a canonical decimal between 1 and 65535", endpoint.ID) } - return controlledsession.ValidateEndpointV2(controlledsession.EndpointV2{ + return controlledsession.ValidateEndpointV1(controlledsession.EndpointV1{ ID: endpoint.ID, Scheme: endpoint.Scheme, Host: endpoint.Host, Port: uint32(port), }) } -func controlledSessionOpenedEndpointsV1(endpoints []ControlledSessionEndpointPlanV1) []controlledsession.EndpointV2 { - result := make([]controlledsession.EndpointV2, len(endpoints)) +func controlledSessionOpenedEndpointsV1(endpoints []ControlledSessionEndpointPlanV1) []controlledsession.EndpointV1 { + result := make([]controlledsession.EndpointV1, len(endpoints)) for index, endpoint := range endpoints { port, _ := strconv.ParseUint(endpoint.Port, 10, 16) - result[index] = controlledsession.EndpointV2{ + result[index] = controlledsession.EndpointV1{ ID: endpoint.ID, Scheme: endpoint.Scheme, Host: endpoint.Host, Port: uint32(port), } } diff --git a/internal/dockerdeploy/controlled_session_plan_test.go b/internal/dockerdeploy/controlled_session_plan_test.go index d8aaba99..9c7f8127 100644 --- a/internal/dockerdeploy/controlled_session_plan_test.go +++ b/internal/dockerdeploy/controlled_session_plan_test.go @@ -186,8 +186,8 @@ func TestPlanControlledSessionV1FreezesGrantedEndpointCoordinatesAndLeaseNetwork t.Fatal(err) } wantEndpoints := []ControlledSessionEndpointPlanV1{ - {ID: "browser", Scheme: "http", Host: controlledsession.WorkloadEndpointHostV2, Port: "8080"}, - {ID: "socket", Scheme: "ws", Host: controlledsession.WorkloadEndpointHostV2, Port: "9090"}, + {ID: "browser", Scheme: "http", Host: controlledsession.WorkloadEndpointHostV1, Port: "8080"}, + {ID: "socket", Scheme: "ws", Host: controlledsession.WorkloadEndpointHostV1, Port: "9090"}, } wantName := input.WorkloadRuntime.Docker.NetworkName + "-session-" + input.LiveRunID if !reflect.DeepEqual(plan.Controller.SessionNetwork.Endpoints, wantEndpoints) || @@ -215,9 +215,9 @@ func TestPlanControlledSessionV1FreezesGrantedEndpointCoordinatesAndLeaseNetwork }) { t.Fatalf("inert Docker network modes = %q/%q", plan.Controller.Network, plan.Workload.Network) } - wantOpenedEndpoints := []controlledsession.EndpointV2{ - {ID: "browser", Scheme: "http", Host: controlledsession.WorkloadEndpointHostV2, Port: 8080}, - {ID: "socket", Scheme: "ws", Host: controlledsession.WorkloadEndpointHostV2, Port: 9090}, + wantOpenedEndpoints := []controlledsession.EndpointV1{ + {ID: "browser", Scheme: "http", Host: controlledsession.WorkloadEndpointHostV1, Port: 8080}, + {ID: "socket", Scheme: "ws", Host: controlledsession.WorkloadEndpointHostV1, Port: 9090}, } if got := controlledSessionOpenedEndpointsV1(plan.Controller.SessionNetwork.Endpoints); !reflect.DeepEqual(got, wantOpenedEndpoints) { t.Fatalf("planned opened endpoints = %#v, want %#v", got, wantOpenedEndpoints) diff --git a/internal/dockerdeploy/controlled_session_session_io_integration_test.go b/internal/dockerdeploy/controlled_session_session_io_integration_test.go index b4dd8822..6809fec0 100644 --- a/internal/dockerdeploy/controlled_session_session_io_integration_test.go +++ b/internal/dockerdeploy/controlled_session_session_io_integration_test.go @@ -145,8 +145,8 @@ func prepareControlledSessionIOBridgeIntegrationV1( authorization := testControlledSessionChannelAuthorizationV1(t, identity) channel, err := controlledsession.PreparePrivateChannelV1(controlledsession.PrivateChannelConfigV1{ HostDirectory: filepath.Join(shortControlledSessionChannelTestDirectoryV1(t), "bridge"), - Opened: controlledsession.OpenedV2{ - Authorization: authorization, Endpoints: []controlledsession.EndpointV2{}, Columns: 80, Rows: 24, + Opened: controlledsession.OpenedV1{ + Authorization: authorization, Endpoints: []controlledsession.EndpointV1{}, Columns: 80, Rows: 24, OutputFinalizationTimeoutMilliseconds: controlledsession.DefaultOutputFinalizationTimeoutMillisecondsV1, }, }) @@ -170,7 +170,7 @@ func prepareControlledSessionIOBridgeIntegrationV1( t.Fatal(err) } t.Cleanup(func() { _ = client.Close() }) - if opened, err := controlledsession.ReadEventV2(client); err != nil || opened.Kind != controlledsession.EventOpenedV1 { + if opened, err := controlledsession.ReadEventV1(client); err != nil || opened.Kind != controlledsession.EventOpenedV1 { t.Fatalf("opened event = %#v, %v", opened, err) } claimed := <-claim @@ -227,7 +227,7 @@ func prepareControlledSessionIOBridgeIntegrationV1( func writeSessionIORequestIntegrationV1(t *testing.T, client *net.UnixConn, request controlledsession.RequestV1) { t.Helper() - if err := controlledsession.WriteRequestV2(client, request); err != nil { + if err := controlledsession.WriteRequestV1(client, request); err != nil { t.Fatal(err) } } @@ -257,7 +257,7 @@ func startSessionIOEventCaptureV1(client *net.UnixConn) *sessionIOEventCaptureV1 go func() { defer close(capture.done) for { - event, err := controlledsession.ReadEventV2(client) + event, err := controlledsession.ReadEventV1(client) capture.mu.Lock() if err != nil { if !errors.Is(err, net.ErrClosed) { diff --git a/internal/dockerdeploy/controlled_session_supervisor.go b/internal/dockerdeploy/controlled_session_supervisor.go index 087456a8..c9ee57fc 100644 --- a/internal/dockerdeploy/controlled_session_supervisor.go +++ b/internal/dockerdeploy/controlled_session_supervisor.go @@ -128,6 +128,8 @@ type controlledSessionSupervisorV1 struct { workload controlledSessionWorkloadRuntimeV1 bridge *controlledsession.SessionIOBridgeV1 + requestReadinessMu sync.Mutex + controllerRequestsReady bool startupResolved chan struct{} resolveOnce sync.Once stateChanged chan struct{} @@ -579,6 +581,15 @@ func (supervisor *controlledSessionSupervisorV1) run(ctx context.Context) (Contr if _, err := supervisor.observe(controlledsession.ObservationV1{Kind: controlledsession.ObservationActivatedV1}); err != nil { return supervisor.finishStartupFailure(fmt.Errorf("activate controlled-session lifecycle: %w", err)) } + supervisor.requestReadinessMu.Lock() + readyErr := supervisor.sendLifecycleEvent(controlledsession.EventV1{Kind: controlledsession.EventReadyV1}) + if readyErr == nil { + supervisor.controllerRequestsReady = true + } + supervisor.requestReadinessMu.Unlock() + if readyErr != nil { + supervisor.loseTransport("send ready event", readyErr) + } supervisor.resolveStartup() supervisor.waitForTermination(ctx) @@ -861,12 +872,20 @@ func observeControlledSessionProcessV1( } func (supervisor *controlledSessionSupervisorV1) handleRequest(ctx context.Context, request controlledsession.RequestV1) error { - select { - case <-supervisor.startupResolved: - case <-ctx.Done(): - return ctx.Err() + if request.Kind != controlledsession.RequestAcknowledgeTerminatedV1 { + supervisor.requestReadinessMu.Lock() + ready := supervisor.controllerRequestsReady + supervisor.requestReadinessMu.Unlock() + if !ready { + return fmt.Errorf("%w: requests are not accepted before ready", controlledsession.ErrRequestRejected) + } } if request.Kind == controlledsession.RequestAcknowledgeTerminatedV1 { + select { + case <-supervisor.startupResolved: + case <-ctx.Done(): + return ctx.Err() + } select { case <-supervisor.resultDeliveryStarted: select { diff --git a/internal/dockerdeploy/controlled_session_supervisor_test.go b/internal/dockerdeploy/controlled_session_supervisor_test.go index f3d06ff2..0c6ea8d6 100644 --- a/internal/dockerdeploy/controlled_session_supervisor_test.go +++ b/internal/dockerdeploy/controlled_session_supervisor_test.go @@ -93,7 +93,7 @@ func TestRunControlledSessionV1OwnsNormalLifecycle(t *testing.T) { } events := transport.snapshotEvents() - if len(events) != 5 { + if len(events) != 6 { t.Fatalf("events = %#v", events) } indices := map[controlledsession.EventKindV1]int{} @@ -101,6 +101,7 @@ func TestRunControlledSessionV1OwnsNormalLifecycle(t *testing.T) { indices[event.Kind] = index } for _, kind := range []controlledsession.EventKindV1{ + controlledsession.EventReadyV1, controlledsession.EventOutputV1, controlledsession.EventWorkloadExitV1, controlledsession.EventTerminatingV1, @@ -111,7 +112,8 @@ func TestRunControlledSessionV1OwnsNormalLifecycle(t *testing.T) { t.Fatalf("event %q missing from %#v", kind, events) } } - if indices[controlledsession.EventWorkloadExitV1] > indices[controlledsession.EventTerminatingV1] || + if indices[controlledsession.EventReadyV1] > indices[controlledsession.EventWorkloadExitV1] || + indices[controlledsession.EventWorkloadExitV1] > indices[controlledsession.EventTerminatingV1] || indices[controlledsession.EventOutputV1] > indices[controlledsession.EventWorkloadOutputsFinalizedV1] || indices[controlledsession.EventWorkloadOutputsFinalizedV1] > indices[controlledsession.EventTerminatedV1] { t.Fatalf("event order = %#v", events) @@ -121,6 +123,28 @@ func TestRunControlledSessionV1OwnsNormalLifecycle(t *testing.T) { } } +func TestControlledSessionSupervisorRejectsRequestBeforeReady(t *testing.T) { + plan := controlledSessionControllerIntegrationPlanV1(t, "test-image", []string{"/controller"}) + machine, err := controlledsession.NewMachineV1(plan.Authorization) + if err != nil { + t.Fatal(err) + } + workload := newFakeControlledSessionWorkloadV1(nil, 0) + supervisor := &controlledSessionSupervisorV1{machine: machine, workload: workload} + + err = supervisor.handleRequest(t.Context(), controlledsession.RequestV1{ + Kind: controlledsession.RequestInputV1, Bytes: []byte("too early"), + }) + if !errors.Is(err, controlledsession.ErrRequestRejected) || !strings.Contains(err.Error(), "before ready") { + t.Fatalf("pre-ready request error = %v", err) + } + select { + case <-workload.inputStarted: + t.Fatal("pre-ready request reached the workload") + default: + } +} + func TestRunControlledSessionV1PreparesAttachesAndCleansLeaseNetwork(t *testing.T) { plan := controlledSessionNetworkPlanFixtureV1(t) requests := make(chan controlledsession.RequestV1, 8) @@ -1200,6 +1224,11 @@ func TestRunControlledSessionV1CleansPreparedResourcesAfterStartupFailure(t *tes if !controller.cleaned || !workload.cleaned || !channel.closed { t.Fatalf("cleanup = controller %t workload %t channel %t", controller.cleaned, workload.cleaned, channel.closed) } + for _, event := range transport.snapshotEvents() { + if event.Kind == controlledsession.EventReadyV1 { + t.Fatalf("startup failure emitted ready: %#v", transport.snapshotEvents()) + } + } } func TestRunControlledSessionV1RecordsHostCancellationDuringStartup(t *testing.T) { @@ -1381,8 +1410,9 @@ func TestRunControlledSessionV1StopsWorkloadAfterOutputDeliveryFailure(t *testin plan := controlledSessionControllerIntegrationPlanV1(t, "test-image", []string{"/controller"}) deliveryErr := errors.New("controller event transport failed") transport := &fakeControlledSessionTransportV1{ - requests: make(chan controlledsession.RequestV1), - writeErr: deliveryErr, + requests: make(chan controlledsession.RequestV1), + writeErr: deliveryErr, + writeErrKind: controlledsession.EventOutputV1, } controller := newFakeControlledSessionProcessV1() workload := newFakeControlledSessionWorkloadV1(nil, 143) @@ -1418,6 +1448,41 @@ func TestRunControlledSessionV1StopsWorkloadAfterOutputDeliveryFailure(t *testin } } +func TestRunControlledSessionV1StopsWorkloadAfterReadyDeliveryFailure(t *testing.T) { + plan := controlledSessionControllerIntegrationPlanV1(t, "test-image", []string{"/controller"}) + deliveryErr := errors.New("controller ready transport failed") + transport := &fakeControlledSessionTransportV1{ + requests: make(chan controlledsession.RequestV1), + writeErr: deliveryErr, + writeErrKind: controlledsession.EventReadyV1, + } + controller := newFakeControlledSessionProcessV1() + workload := newFakeControlledSessionWorkloadV1(nil, 143) + workload.exitOnStart = false + channel := &fakeControlledSessionChannelV1{transport: transport} + + result, err := runControlledSessionV1(t.Context(), plan, testControlledSessionRunOptionsV1(), controlledSessionSupervisorBackendV1{ + prepareChannel: func(ControlledSessionExecutionPlanV1) (controlledSessionChannelRuntimeV1, error) { + return channel, nil + }, + prepareController: func(context.Context, ControlledSessionContainerPlanV1) (controlledSessionControllerRuntimeV1, error) { + return controller, nil + }, + prepareWorkload: func(context.Context, ControlledSessionContainerPlanV1) (controlledSessionWorkloadRuntimeV1, error) { + return workload, nil + }, + now: time.Now, + }) + if !errors.Is(err, deliveryErr) { + t.Fatalf("error = %v", err) + } + if result.SessionResult.Cause != controlledsession.CauseControllerLostV1 || + result.ResultDelivered || !workload.gracefulStopped || !workload.cleaned || !controller.cleaned || !channel.closed { + t.Fatalf("ready-delivery-loss result = %#v, workload = %#v, controller cleaned = %t, channel closed = %t", + result, workload, controller.cleaned, channel.closed) + } +} + func TestRunControlledSessionV1StopsWorkloadAfterHostCancellation(t *testing.T) { plan := controlledSessionControllerIntegrationPlanV1(t, "test-image", []string{"/controller"}) requests := make(chan controlledsession.RequestV1, 4) @@ -1608,6 +1673,7 @@ type fakeControlledSessionTransportV1 struct { onRequest func(controlledsession.RequestV1) onEvent func(controlledsession.EventV1) writeErr error + writeErrKind controlledsession.EventKindV1 blockEventKind controlledsession.EventKindV1 eventWriteBlocked chan struct{} releaseEventWrite chan struct{} @@ -1629,11 +1695,12 @@ func (transport *fakeControlledSessionTransportV1) ReadRequest(ctx context.Conte } func (transport *fakeControlledSessionTransportV1) WriteEvent(ctx context.Context, event controlledsession.EventV1) error { - if transport.writeErr != nil { - return transport.writeErr - } event.Bytes = append([]byte(nil), event.Bytes...) transport.mu.Lock() + if transport.writeErr != nil && (transport.writeErrKind == "" || transport.writeErrKind == event.Kind) { + defer transport.mu.Unlock() + return transport.writeErr + } transport.events = append(transport.events, event) transport.mu.Unlock() if transport.onEvent != nil { diff --git a/internal/dockerdeploy/testdata/session_channel_helper/main.go b/internal/dockerdeploy/testdata/session_channel_helper/main.go index 9edcb5fe..0b0db8da 100644 --- a/internal/dockerdeploy/testdata/session_channel_helper/main.go +++ b/internal/dockerdeploy/testdata/session_channel_helper/main.go @@ -43,7 +43,7 @@ func main() { fail("connect to private channel: %v", err) } defer connection.Close() - event, err := controlledsession.ReadEventV2(connection) + event, err := controlledsession.ReadEventV1(connection) if err != nil { fail("read opened event: %v", err) } @@ -66,7 +66,7 @@ func main() { runSupervisorProof(connection) return } - if err := controlledsession.WriteRequestV2(connection, controlledsession.RequestV1{Kind: controlledsession.RequestCompleteV1}); err != nil { + if err := controlledsession.WriteRequestV1(connection, controlledsession.RequestV1{Kind: controlledsession.RequestCompleteV1}); err != nil { fail("write complete request: %v", err) } fmt.Println("PASS") @@ -76,7 +76,7 @@ func main() { } } -func runNetworkSupervisorProof(connection readWriteCloser, opened *controlledsession.OpenedV2, localPeer string, publicPeer string) { +func runNetworkSupervisorProof(connection readWriteCloser, opened *controlledsession.OpenedV1, localPeer string, publicPeer string) { networkProofFailureStage = 1 localPeer = requireIPPort(localPeer) publicPeer = requireIPPort(publicPeer) @@ -84,6 +84,7 @@ func runNetworkSupervisorProof(connection readWriteCloser, opened *controlledses opened.Endpoints[1].ID != "socket" || opened.Endpoints[1].Host != "workload" || opened.Endpoints[1].Port != 8080 { fail("unexpected session endpoints: %#v", opened.Endpoints) } + requireReady(connection) writeRequest(connection, controlledsession.RequestV1{ Kind: controlledsession.RequestInputV1, Bytes: []byte("/session-network-helper serve & network_pid=$!\n"), @@ -93,7 +94,7 @@ func runNetworkSupervisorProof(connection readWriteCloser, opened *controlledses finishSent := false workloadExited := false for { - event, err := controlledsession.ReadEventV2(connection) + event, err := controlledsession.ReadEventV1(connection) if err != nil { fail("read network session event: %v", err) } @@ -286,6 +287,7 @@ func checkNetworkDial(address string, want bool) { } func runSupervisorProof(connection readWriteCloser) { + requireReady(connection) writeRequest(connection, controlledsession.RequestV1{ Kind: controlledsession.RequestInputV1, Bytes: []byte("stty size; printf 'SIZE-1-DONE\\n'\n"), }) @@ -296,7 +298,7 @@ func runSupervisorProof(connection readWriteCloser) { workloadExited := false forgedTerminalResult := append([]byte{0x1e}, []byte(`{"kind":"terminated","cause":"forged"}`)...) for { - event, err := controlledsession.ReadEventV2(connection) + event, err := controlledsession.ReadEventV1(connection) if err != nil { fail("read session event: %v", err) } @@ -354,6 +356,16 @@ func runSupervisorProof(connection readWriteCloser) { } } +func requireReady(connection readWriteCloser) { + event, err := controlledsession.ReadEventV1(connection) + if err != nil { + fail("read ready event: %v", err) + } + if event.Kind != controlledsession.EventReadyV1 { + fail("first post-opened event is not ready: %#v", event) + } +} + type readWriteCloser interface { Read([]byte) (int, error) Write([]byte) (int, error) @@ -361,7 +373,7 @@ type readWriteCloser interface { } func writeRequest(connection readWriteCloser, request controlledsession.RequestV1) { - if err := controlledsession.WriteRequestV2(connection, request); err != nil { + if err := controlledsession.WriteRequestV1(connection, request); err != nil { fail("write %s request: %v", request.Kind, err) } }