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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
104 changes: 104 additions & 0 deletions internal/wsldocker/inspect.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,104 @@
package wsldocker

import (
"context"
"encoding/json"
"errors"
"fmt"
"net/http"
)

const maxContainerInspectOutput = 1 << 20

// ContainerSnapshot is the bounded immutable subset of one exact Docker
// container inspection needed by later ownership and lifecycle checks.
type ContainerSnapshot struct {
id string
labels map[string]string
running bool
tty bool
openStdin bool
}

func (s ContainerSnapshot) ID() string { return s.id }
func (s ContainerSnapshot) Running() bool { return s.running }
func (s ContainerSnapshot) TTY() bool { return s.tty }
func (s ContainerSnapshot) OpenStdin() bool { return s.openStdin }
func (s ContainerSnapshot) Labels() map[string]string { return cloneContainerLabels(s.labels) }

func inspectContainer(ctx context.Context, containerID string, deps operationDependencies) (ContainerSnapshot, error) {
if ctx == nil {
return ContainerSnapshot{}, errors.New("Docker Desktop WSL container inspect requires a context")
}
if err := validateContainerID(containerID); err != nil {
return ContainerSnapshot{}, fmt.Errorf("Docker Desktop WSL container inspect: %w", err)
}
if deps.check == nil || deps.statSocket == nil || deps.perform == nil {
return ContainerSnapshot{}, errors.New("Docker Desktop WSL container inspect dependencies are incomplete")
}
response, err := execute(ctx, Request{
Method: http.MethodGet,
Path: "/containers/" + containerID + "/json",
SuccessStatuses: []int{http.StatusOK},
}, deps)
if err != nil {
return ContainerSnapshot{}, err
}
return decodeContainerInspectResponse(response.Body, containerID)
}

func decodeContainerInspectResponse(raw []byte, expectedID string) (ContainerSnapshot, error) {
if len(raw) > maxContainerInspectOutput {
return ContainerSnapshot{}, fmt.Errorf("Docker container inspect response exceeds %d bytes", maxContainerInspectOutput)
}
var response struct {
ID string `json:"Id"`
Config *struct {
Labels map[string]string `json:"Labels"`
TTY *bool `json:"Tty"`
OpenStdin *bool `json:"OpenStdin"`
} `json:"Config"`
State *struct {
Running *bool `json:"Running"`
} `json:"State"`
Comment thread
AviBackToBlack marked this conversation as resolved.
Comment thread
AviBackToBlack marked this conversation as resolved.
}
if err := json.Unmarshal(raw, &response); err != nil {
return ContainerSnapshot{}, fmt.Errorf("decode Docker container inspect response: %w", err)
}
if err := validateContainerID(response.ID); err != nil {
return ContainerSnapshot{}, fmt.Errorf("Docker container inspect returned invalid ID: %w", err)
}
if response.ID != expectedID {
return ContainerSnapshot{}, fmt.Errorf("Docker container inspect returned ID %q, expected exact requested ID", response.ID)
}
if response.Config == nil {
return ContainerSnapshot{}, errors.New("Docker container inspect response is missing Config")
}
if response.State == nil {
return ContainerSnapshot{}, errors.New("Docker container inspect response is missing State")
}
if response.Config.TTY == nil {
return ContainerSnapshot{}, errors.New("Docker container inspect response is missing Config.Tty")
}
if response.Config.OpenStdin == nil {
return ContainerSnapshot{}, errors.New("Docker container inspect response is missing Config.OpenStdin")
}
if response.State.Running == nil {
return ContainerSnapshot{}, errors.New("Docker container inspect response is missing State.Running")
}
return ContainerSnapshot{
id: response.ID,
labels: cloneContainerLabels(response.Config.Labels),
running: *response.State.Running,
tty: *response.Config.TTY,
openStdin: *response.Config.OpenStdin,
}, nil
}

func cloneContainerLabels(labels map[string]string) map[string]string {
cloned := make(map[string]string, len(labels))
for key, value := range labels {
cloned[key] = value
}
return cloned
}
17 changes: 17 additions & 0 deletions internal/wsldocker/inspect_linux.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,17 @@
//go:build linux

package wsldocker

import "context"

// InspectContainer returns the bounded immutable lifecycle subset for one exact
// container reached through the proven Docker Desktop WSL socket.
func InspectContainer(ctx context.Context, containerID string) (ContainerSnapshot, error) {
return inspectContainer(ctx, containerID, operationDependencies{
check: Check,
statSocket: statDockerSocket,
perform: func(ctx context.Context, socketPath string, request Request) (operationResult, error) {
return performDockerRequest(ctx, socketPath, request, operationTimeout, maxContainerInspectOutput, 0)
},
})
}
14 changes: 14 additions & 0 deletions internal/wsldocker/inspect_other.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,14 @@
//go:build !linux

package wsldocker

import (
"context"
"errors"
)

// InspectContainer is unavailable outside Linux because the fixed Unix socket
// and peer credentials are part of the Docker Desktop WSL trust boundary.
func InspectContainer(context.Context, string) (ContainerSnapshot, error) {
return ContainerSnapshot{}, errors.New("Docker Desktop WSL container inspect requires Linux")
}
106 changes: 106 additions & 0 deletions internal/wsldocker/inspect_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,106 @@
package wsldocker

import (
"context"
"net/http"
"strings"
"testing"
)

func TestInspectContainerBindsExactSnapshotToProvenSocket(t *testing.T) {
statCalls := 0
deps := validOperationDependencies(validSocketInfo())
deps.statSocket = func(path string) (socketInfo, error) {
if path != DockerSocketPath {
t.Fatalf("stat path = %q", path)
}
statCalls++
return validSocketInfo(), nil
}
deps.perform = func(_ context.Context, path string, request Request) (operationResult, error) {
if path != DockerSocketPath || request.Method != http.MethodGet || request.Path != "/containers/"+testContainerID+"/json" {
t.Fatalf("perform(%q, %+v)", path, request)
}
if len(request.Query) != 0 || len(request.Body) != 0 || len(request.SuccessStatuses) != 1 || request.SuccessStatuses[0] != http.StatusOK {
t.Fatalf("inspect request = %+v", request)
}
return operationResult{StatusCode: http.StatusOK, PeerUID: 0, Raw: []byte(`{
"Id":"` + testContainerID + `",
"Config":{"Labels":{"cb.managed":"true","cb.run_id":"run-1"},"Tty":true,"OpenStdin":true},
"State":{"Running":true}
}`)}, nil
}
snapshot, err := inspectContainer(context.Background(), testContainerID, deps)
if err != nil {
t.Fatal(err)
}
if snapshot.ID() != testContainerID || !snapshot.Running() || !snapshot.TTY() || !snapshot.OpenStdin() || statCalls != 2 {
t.Fatalf("snapshot = %#v, stat calls = %d", snapshot, statCalls)
}
labels := snapshot.Labels()
if labels["cb.managed"] != "true" || labels["cb.run_id"] != "run-1" {
t.Fatalf("labels = %#v", labels)
}
labels["cb.managed"] = "mutated"
if snapshot.Labels()["cb.managed"] != "true" {
t.Fatal("Labels returned mutable internal state")
}
}

func TestInspectContainerRejectsInvalidInputsBeforeProof(t *testing.T) {
for name, test := range map[string]struct {
ctx context.Context
containerID string
}{
"nil context": {containerID: testContainerID},
"short ID": {ctx: context.Background(), containerID: "abc"},
"uppercase": {ctx: context.Background(), containerID: strings.ToUpper(testContainerID)},
} {
t.Run(name, func(t *testing.T) {
deps := validOperationDependencies(validSocketInfo())
deps.check = func(context.Context) (Result, error) { panic("proof reached for invalid inspect") }
if _, err := inspectContainer(test.ctx, test.containerID, deps); err == nil {
t.Fatal("inspectContainer() succeeded")
}
})
}
if _, err := inspectContainer(context.Background(), testContainerID, operationDependencies{}); err == nil || !strings.Contains(err.Error(), "dependencies are incomplete") {
t.Fatalf("incomplete-dependencies error = %v", err)
}
}

func TestDecodeContainerInspectResponseRejectsUnsafeShapes(t *testing.T) {
valid := `{"Id":"` + testContainerID + `","Config":{"Labels":null,"Tty":false,"OpenStdin":false},"State":{"Running":false}}`
otherID := strings.Repeat("a", 64)
tests := map[string]struct {
raw []byte
want string
}{
"malformed": {raw: []byte(`{`), want: "decode"},
"oversized": {raw: make([]byte, maxContainerInspectOutput+1), want: "exceeds"},
"missing ID": {raw: []byte(`{"Config":{},"State":{}}`), want: "invalid ID"},
"invalid ID": {raw: []byte(`{"Id":"abc","Config":{},"State":{}}`), want: "invalid ID"},
"mismatched ID": {raw: []byte(`{"Id":"` + otherID + `","Config":{},"State":{}}`), want: "expected exact"},
"missing Config": {raw: []byte(`{"Id":"` + testContainerID + `","State":{}}`), want: "missing Config"},
"null Config": {raw: []byte(`{"Id":"` + testContainerID + `","Config":null,"State":{}}`), want: "missing Config"},
"missing State": {raw: []byte(`{"Id":"` + testContainerID + `","Config":{}}`), want: "missing State"},
"null State": {raw: []byte(`{"Id":"` + testContainerID + `","Config":{},"State":null}`), want: "missing State"},
"missing Tty": {raw: []byte(`{"Id":"` + testContainerID + `","Config":{"OpenStdin":false},"State":{"Running":false}}`), want: "missing Config.Tty"},
"null Tty": {raw: []byte(`{"Id":"` + testContainerID + `","Config":{"Tty":null,"OpenStdin":false},"State":{"Running":false}}`), want: "missing Config.Tty"},
"missing stdin": {raw: []byte(`{"Id":"` + testContainerID + `","Config":{"Tty":false},"State":{"Running":false}}`), want: "missing Config.OpenStdin"},
"null stdin": {raw: []byte(`{"Id":"` + testContainerID + `","Config":{"Tty":false,"OpenStdin":null},"State":{"Running":false}}`), want: "missing Config.OpenStdin"},
"missing Running": {raw: []byte(`{"Id":"` + testContainerID + `","Config":{"Tty":false,"OpenStdin":false},"State":{}}`), want: "missing State.Running"},
"null Running": {raw: []byte(`{"Id":"` + testContainerID + `","Config":{"Tty":false,"OpenStdin":false},"State":{"Running":null}}`), want: "missing State.Running"},
}
for name, test := range tests {
t.Run(name, func(t *testing.T) {
if _, err := decodeContainerInspectResponse(test.raw, testContainerID); err == nil || !strings.Contains(err.Error(), test.want) {
t.Fatalf("decode error = %v, want containing %q", err, test.want)
}
})
}
snapshot, err := decodeContainerInspectResponse([]byte(valid), testContainerID)
if err != nil || snapshot.Labels() == nil || len(snapshot.Labels()) != 0 || snapshot.Running() || snapshot.TTY() || snapshot.OpenStdin() {
t.Fatalf("minimal snapshot = %#v, err = %v", snapshot, err)
}
}
Loading