Luigit
repositories / bugabinga.net

bugabinga.net

personal infrastructure for bugabinga!

owned by admin

services/toad/internal/host/host.go

Raw
package host

import (
	"context"
	"encoding/json"
	"errors"
	"fmt"
	"io"
	"net"
	"net/http"
	"net/url"
	"os"
	"path/filepath"
	"strings"
	"time"

	"bugabinga.net/toad/internal/rollout"
	"github.com/godbus/dbus/v5"
)

const (
	podmanAPI      = "/v5.0.0/libpod"
	systemdName    = "org.freedesktop.systemd1"
	managerPath    = dbus.ObjectPath("/org/freedesktop/systemd1")
	managerIFace   = systemdName + ".Manager"
	serviceIFace   = systemdName + ".Service"
	propertiesGet  = "org.freedesktop.DBus.Properties.Get"
	jobRemovedName = managerIFace + ".JobRemoved"
)

// Host is the rollout's only path to the machine. Both mounted APIs belong to
// the dedicated toad user, so this process cannot cross into another user scope.
type Host struct {
	QuadletDir string
	Podman     *http.Client
	ProbeHTTP  *http.Client
	Systemd    systemd
}

// systemd is the native user-manager boundary. Keeping it narrow makes job
// completion behavior testable without a running user manager.
type systemd interface {
	Reload(context.Context) error
	Start(context.Context, string) error
	Stop(context.Context, string) error
	RestartCount(context.Context, string) (int, error)
	Unit(context.Context, string) (rollout.UnitState, error)
	Close() error
}

func New(quadletDir, podmanSocket, busAddress string) (*Host, error) {
	manager, err := connectSystemd(busAddress)
	if err != nil {
		return nil, fmt.Errorf("connect systemd user bus: %w", err)
	}
	transport := &http.Transport{
		DialContext: func(ctx context.Context, _, _ string) (net.Conn, error) {
			return (&net.Dialer{}).DialContext(ctx, "unix", podmanSocket)
		},
		ResponseHeaderTimeout: 30 * time.Second,
	}
	return &Host{
		QuadletDir: quadletDir,
		Podman:     &http.Client{Transport: transport},
		ProbeHTTP:  &http.Client{Timeout: 5 * time.Second},
		Systemd:    manager,
	}, nil
}

func (h *Host) Close() error { return h.Systemd.Close() }

func (h *Host) Pull(ctx context.Context, reference string) (string, error) {
	query := url.Values{"reference": {reference}, "quiet": {"true"}, "policy": {"always"}}
	var pulled struct {
		Images []string `json:"images"`
		Error  string   `json:"error"`
	}
	if err := h.podmanJSON(ctx, http.MethodPost, podmanAPI+"/images/pull?"+query.Encode(), &pulled); err != nil {
		return "", fmt.Errorf("pull %s: %w", reference, err)
	}
	if pulled.Error != "" {
		return "", fmt.Errorf("pull %s: %s", reference, pulled.Error)
	}
	if len(pulled.Images) != 1 || pulled.Images[0] == "" {
		return "", fmt.Errorf("pull %s returned %d images", reference, len(pulled.Images))
	}
	var image struct {
		Digest string `json:"Digest"`
	}
	path := podmanAPI + "/images/" + url.PathEscape(pulled.Images[0]) + "/json"
	if err := h.podmanJSON(ctx, http.MethodGet, path, &image); err != nil {
		return "", fmt.Errorf("inspect pulled image: %w", err)
	}
	if image.Digest == "" {
		return "", errors.New("inspect pulled image returned no digest")
	}
	return image.Digest, nil
}

func (h *Host) InstallUnit(fileName, content string) error {
	if fileName == "" || fileName != filepath.Base(fileName) {
		return fmt.Errorf("invalid unit file name %q", fileName)
	}
	path := filepath.Join(h.QuadletDir, fileName)
	temporary, err := os.CreateTemp(h.QuadletDir, ".toad-unit-")
	if err != nil {
		return err
	}
	defer os.Remove(temporary.Name())
	if _, err := temporary.WriteString(content); err != nil {
		temporary.Close()
		return err
	}
	if err := temporary.Sync(); err != nil {
		temporary.Close()
		return err
	}
	if err := temporary.Close(); err != nil {
		return err
	}
	if err := os.Chmod(temporary.Name(), 0o640); err != nil {
		return err
	}
	return os.Rename(temporary.Name(), path)
}

func (h *Host) DaemonReload(ctx context.Context) error { return h.Systemd.Reload(ctx) }
func (h *Host) Stop(ctx context.Context, unit string) error {
	return h.Systemd.Stop(ctx, unit)
}
func (h *Host) Start(ctx context.Context, unit string) error {
	return h.Systemd.Start(ctx, unit)
}

func (h *Host) Container(ctx context.Context, name string) (rollout.ContainerState, error) {
	var inspected struct {
		State *struct {
			Running *bool `json:"Running"`
		} `json:"State"`
		ImageDigest string `json:"ImageDigest"`
	}
	path := podmanAPI + "/containers/" + url.PathEscape(name) + "/json"
	err := h.podmanJSON(ctx, http.MethodGet, path, &inspected)
	if errors.Is(err, errNotFound) {
		return rollout.ContainerState{}, nil
	}
	if err != nil {
		return rollout.ContainerState{}, fmt.Errorf("inspect container %s: %w", name, err)
	}
	if inspected.State == nil || inspected.State.Running == nil {
		return rollout.ContainerState{}, fmt.Errorf("inspect container %s returned no running state", name)
	}
	return rollout.ContainerState{Running: *inspected.State.Running, ImageDigest: inspected.ImageDigest}, nil
}

func (h *Host) RestartCount(ctx context.Context, unit string) (int, error) {
	return h.Systemd.RestartCount(ctx, unit)
}

func (h *Host) Unit(ctx context.Context, unit string) (rollout.UnitState, error) {
	return h.Systemd.Unit(ctx, unit)
}

func (h *Host) Probe(ctx context.Context, target string) error {
	request, err := http.NewRequestWithContext(ctx, http.MethodGet, target, nil)
	if err != nil {
		return err
	}
	response, err := h.ProbeHTTP.Do(request)
	if err != nil {
		return err
	}
	defer response.Body.Close()
	if response.StatusCode >= 400 {
		return fmt.Errorf("readiness probe returned %s", response.Status)
	}
	return nil
}

var errNotFound = errors.New("podman object not found")

func (h *Host) podmanJSON(ctx context.Context, method, path string, result any) error {
	request, err := http.NewRequestWithContext(ctx, method, "http://podman"+path, nil)
	if err != nil {
		return err
	}
	response, err := h.Podman.Do(request)
	if err != nil {
		return err
	}
	defer response.Body.Close()
	if response.StatusCode == http.StatusNotFound {
		return errNotFound
	}
	if response.StatusCode < 200 || response.StatusCode >= 300 {
		body, _ := io.ReadAll(io.LimitReader(response.Body, 16<<10))
		return fmt.Errorf("podman returned %s: %s", response.Status, strings.TrimSpace(string(body)))
	}
	decoder := json.NewDecoder(io.LimitReader(response.Body, 1<<20))
	if err := decoder.Decode(result); err != nil {
		return fmt.Errorf("decode podman response: %w", err)
	}
	return nil
}

type dbusSystemd struct {
	conn    *dbus.Conn
	manager dbus.BusObject
}

func connectSystemd(address string) (*dbusSystemd, error) {
	conn, err := dbus.Connect(address)
	if err != nil {
		return nil, err
	}
	manager := conn.Object(systemdName, managerPath)
	ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
	defer cancel()
	if call := manager.CallWithContext(ctx, managerIFace+".Subscribe", 0); call.Err != nil {
		conn.Close()
		return nil, fmt.Errorf("subscribe to systemd jobs: %w", call.Err)
	}
	return &dbusSystemd{conn: conn, manager: manager}, nil
}

func (s *dbusSystemd) Close() error { return s.conn.Close() }

func (s *dbusSystemd) Reload(ctx context.Context) error {
	return s.manager.CallWithContext(ctx, managerIFace+".Reload", 0).Err
}

func (s *dbusSystemd) Start(ctx context.Context, unit string) error {
	return s.runJob(ctx, "StartUnit", unit)
}

func (s *dbusSystemd) Stop(ctx context.Context, unit string) error {
	return s.runJob(ctx, "StopUnit", unit)
}

func (s *dbusSystemd) runJob(ctx context.Context, method, unit string) error {
	options := []dbus.MatchOption{
		dbus.WithMatchObjectPath(managerPath),
		dbus.WithMatchInterface(managerIFace),
		dbus.WithMatchMember("JobRemoved"),
		dbus.WithMatchArg(2, unit),
	}
	if err := s.conn.AddMatchSignalContext(ctx, options...); err != nil {
		return fmt.Errorf("watch systemd job: %w", err)
	}
	defer s.conn.RemoveMatchSignal(options...)

	signals := make(chan *dbus.Signal, 16)
	s.conn.Signal(signals)
	defer s.conn.RemoveSignal(signals)

	var job dbus.ObjectPath
	call := s.manager.CallWithContext(ctx, managerIFace+"."+method, 0, unit, "replace")
	if call.Err != nil {
		return call.Err
	}
	if err := call.Store(&job); err != nil {
		return fmt.Errorf("decode systemd job: %w", err)
	}
	for {
		select {
		case <-ctx.Done():
			return ctx.Err()
		case signal, ok := <-signals:
			if !ok {
				return errors.New("systemd signal connection closed")
			}
			if signal == nil || signal.Name != jobRemovedName || len(signal.Body) != 4 {
				continue
			}
			removed, ok := signal.Body[1].(dbus.ObjectPath)
			if !ok || removed != job {
				continue
			}
			result, ok := signal.Body[3].(string)
			if !ok {
				return errors.New("systemd JobRemoved result has unexpected type")
			}
			if result != "done" {
				return fmt.Errorf("systemd %s job for %s finished %s", strings.ToLower(strings.TrimSuffix(method, "Unit")), unit, result)
			}
			return nil
		}
	}
}

func (s *dbusSystemd) RestartCount(ctx context.Context, unit string) (int, error) {
	state, err := s.Unit(ctx, unit)
	return state.Restarts, err
}

func (s *dbusSystemd) Unit(ctx context.Context, unit string) (rollout.UnitState, error) {
	var path dbus.ObjectPath
	call := s.manager.CallWithContext(ctx, managerIFace+".GetUnit", 0, unit)
	if call.Err != nil {
		return rollout.UnitState{}, call.Err
	}
	if err := call.Store(&path); err != nil {
		return rollout.UnitState{}, fmt.Errorf("decode systemd unit path: %w", err)
	}
	object := s.conn.Object(systemdName, path)
	property := func(iface, name string) (any, error) {
		var value dbus.Variant
		call := object.CallWithContext(ctx, propertiesGet, 0, iface, name)
		if call.Err != nil {
			return nil, call.Err
		}
		if err := call.Store(&value); err != nil {
			return nil, fmt.Errorf("decode %s: %w", name, err)
		}
		return value.Value(), nil
	}
	load, err := property(systemdName+".Unit", "LoadState")
	if err != nil {
		return rollout.UnitState{}, err
	}
	active, err := property(systemdName+".Unit", "ActiveState")
	if err != nil {
		return rollout.UnitState{}, err
	}
	result, err := property(serviceIFace, "Result")
	if err != nil {
		return rollout.UnitState{}, err
	}
	restarts, err := property(serviceIFace, "NRestarts")
	if err != nil {
		return rollout.UnitState{}, err
	}
	loadText, loadOK := load.(string)
	activeText, activeOK := active.(string)
	resultText, resultOK := result.(string)
	restartCount, restartOK := restarts.(uint32)
	if !loadOK || !activeOK || !resultOK || !restartOK {
		return rollout.UnitState{}, fmt.Errorf("unexpected systemd property types: load=%T active=%T result=%T restarts=%T", load, active, result, restarts)
	}
	if uint64(restartCount) > uint64(^uint(0)>>1) {
		return rollout.UnitState{}, fmt.Errorf("NRestarts %d overflows int", restartCount)
	}
	return rollout.UnitState{LoadState: loadText, ActiveState: activeText, Result: resultText, Restarts: int(restartCount)}, nil
}

var _ rollout.Host = (*Host)(nil)