Luigit
repositories / bugabinga.net

bugabinga.net

personal infrastructure for bugabinga!

owned by admin

services/toad/internal/host/host_test.go

Raw
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)
}