package host import ( "context" "net/http" "net/http/httptest" "os" "path/filepath" "strings" "testing" "bugabinga.net/toad/internal/rollout" ) type fakeSystemd struct { reloaded bool started string stopped string restarts int err error } func (f *fakeSystemd) Reload(context.Context) error { f.reloaded = true; return f.err } func (f *fakeSystemd) Start(_ context.Context, unit string) error { f.started = unit return f.err } func (f *fakeSystemd) Stop(_ context.Context, unit string) error { f.stopped = unit return f.err } func (f *fakeSystemd) RestartCount(context.Context, string) (int, error) { return f.restarts, f.err } func (f *fakeSystemd) Unit(context.Context, string) (rollout.UnitState, error) { return rollout.UnitState{LoadState: "loaded", ActiveState: "active", Result: "success", Restarts: f.restarts}, f.err } func (f *fakeSystemd) Close() error { return f.err } func TestPullReturnsDigestOfPulledImageID(t *testing.T) { var inspected string server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { switch request.URL.Path { case podmanAPI + "/images/pull": if request.Method != http.MethodPost || request.URL.Query().Get("reference") != "registry.example/app@sha256:abc" || request.URL.Query().Get("quiet") != "true" || request.URL.Query().Get("policy") != "always" { t.Fatalf("pull request = %s %s", request.Method, request.URL.String()) } writer.Write([]byte(`{"images":["sha256:image-id"],"id":"sha256:image-id"}`)) case podmanAPI + "/images/sha256:image-id/json": inspected = request.URL.Path writer.Write([]byte(`{"Digest":"sha256:manifest"}`)) default: http.NotFound(writer, request) } })) defer server.Close() host := &Host{Podman: server.Client(), Systemd: &fakeSystemd{}} // Rewrite requests to the fixture while preserving the Podman paths. host.Podman.Transport = rewriteOrigin{base: server.URL, next: server.Client().Transport} digest, err := host.Pull(context.Background(), "registry.example/app@sha256:abc") if err != nil { t.Fatal(err) } if digest != "sha256:manifest" || inspected == "" { t.Fatalf("digest = %q inspected = %q", digest, inspected) } } func TestPullRejectsEmbeddedErrorAndAmbiguousResult(t *testing.T) { for name, fixture := range map[string]struct{ body, want string }{ "embedded error": {`{"error":"manifest unknown"}`, "manifest unknown"}, "no image": {`{"images":[]}`, "returned 0 images"}, "many images": {`{"images":["one","two"]}`, "returned 2 images"}, } { t.Run(name, func(t *testing.T) { server := podmanFixture(t, http.StatusOK, fixture.body) host := &Host{Podman: server.Client(), Systemd: &fakeSystemd{}} host.Podman.Transport = rewriteOrigin{base: server.URL, next: server.Client().Transport} _, err := host.Pull(context.Background(), "example/app:latest") if err == nil || !strings.Contains(err.Error(), fixture.want) { t.Fatalf("error = %v, want %q", err, fixture.want) } }) } } func TestContainerDistinguishesAbsentMalformedAndFailedInspect(t *testing.T) { for name, fixture := range map[string]struct { status int body string wantError bool }{ "absent": {http.StatusNotFound, `{"cause":"no such container"}`, false}, "malformed": {http.StatusOK, `{"State":{},"ImageDigest":"sha256:a"}`, true}, "failure": {http.StatusInternalServerError, `{"cause":"storage failure"}`, true}, } { t.Run(name, func(t *testing.T) { server := podmanFixture(t, fixture.status, fixture.body) host := &Host{Podman: server.Client(), Systemd: &fakeSystemd{}} host.Podman.Transport = rewriteOrigin{base: server.URL, next: server.Client().Transport} state, err := host.Container(context.Background(), "app") if (err != nil) != fixture.wantError { t.Fatalf("state = %+v error = %v", state, err) } if name == "absent" && state.Running { t.Fatal("absent container reported running") } }) } } func TestContainerReturnsRunningDigest(t *testing.T) { server := podmanFixture(t, http.StatusOK, `{"State":{"Running":true},"ImageDigest":"sha256:manifest"}`) host := &Host{Podman: server.Client(), Systemd: &fakeSystemd{}} host.Podman.Transport = rewriteOrigin{base: server.URL, next: server.Client().Transport} state, err := host.Container(context.Background(), "app") if err != nil || !state.Running || state.ImageDigest != "sha256:manifest" { t.Fatalf("state = %+v error = %v", state, err) } } func TestSystemdOperationsUseNativeBoundary(t *testing.T) { manager := &fakeSystemd{restarts: 7} host := &Host{Systemd: manager} ctx := context.Background() if err := host.DaemonReload(ctx); err != nil { t.Fatal(err) } if err := host.Start(ctx, "app.service"); err != nil { t.Fatal(err) } if err := host.Stop(ctx, "app.service"); err != nil { t.Fatal(err) } restarts, err := host.RestartCount(ctx, "app.service") unit, unitErr := host.Unit(ctx, "app.service") if err != nil || unitErr != nil || !manager.reloaded || manager.started != "app.service" || manager.stopped != "app.service" || restarts != 7 || unit.ActiveState != "active" { t.Fatalf("manager = %+v restarts = %d unit = %+v errors = %v, %v", manager, restarts, unit, err, unitErr) } } func TestInstallUnitIsAtomicAndRejectsPaths(t *testing.T) { root := t.TempDir() host := &Host{QuadletDir: root} if err := host.InstallUnit("app.container", "candidate"); err != nil { t.Fatal(err) } data, err := os.ReadFile(filepath.Join(root, "app.container")) if err != nil || string(data) != "candidate" { t.Fatalf("data = %q error = %v", data, err) } if err := host.InstallUnit("../escape.container", "bad"); err == nil { t.Fatal("accepted path traversal") } } func TestPodmanFailureBodyIsBoundedAndReported(t *testing.T) { server := podmanFixture(t, http.StatusInternalServerError, strings.Repeat("x", 20<<10)) host := &Host{Podman: server.Client(), Systemd: &fakeSystemd{}} host.Podman.Transport = rewriteOrigin{base: server.URL, next: server.Client().Transport} _, err := host.Container(context.Background(), "app") if err == nil || !strings.Contains(err.Error(), "500 Internal Server Error") || len(err.Error()) > 17<<10 { t.Fatalf("error length = %d error = %v", len(err.Error()), err) } } func podmanFixture(t *testing.T, status int, body string) *httptest.Server { t.Helper() return httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) { writer.WriteHeader(status) writer.Write([]byte(body)) })) } type rewriteOrigin struct { base string next http.RoundTripper } func (r rewriteOrigin) RoundTrip(request *http.Request) (*http.Response, error) { clone := request.Clone(request.Context()) clone.URL.Scheme = "http" clone.URL.Host = strings.TrimPrefix(r.base, "http://") return r.next.RoundTrip(clone) }