internal/control/control.go

ab56677031290fee18323c467841a25216d45768
gitbay/internal/control/control.go history · blame · raw

140 lines · 3551 bytes

  1// Package control implements the forge control commands executed over SSH.
  2// Every command here is reachable from bare OpenSSH: argv in, JSON or plain
  3// text on stdout, diagnostics on stderr, exit code out.
  4package control
  5
  6import (
  7	"encoding/json"
  8	"fmt"
  9	"io"
 10	"slices"
 11
 12	"github.com/krazywarez/forge/internal/config"
 13	"github.com/krazywarez/forge/internal/protocol"
 14	"github.com/krazywarez/forge/internal/store"
 15)
 16
 17type Ctx struct {
 18	User   store.User
 19	Scope  string // scope of the key that authenticated this session
 20	Store  *store.Store
 21	Cfg    config.Config
 22	Stdin  io.Reader
 23	Stdout io.Writer
 24	Stderr io.Writer
 25	JSON   bool
 26}
 27
 28type Command struct {
 29	Path       []string // e.g. ["keys", "add"]
 30	Summary    string
 31	ReadsStdin bool
 32	Run        func(c *Ctx, args []string) int
 33}
 34
 35var registry []Command
 36
 37func register(cmd Command) { registry = append(registry, cmd) }
 38
 39// Commands returns the registry, for the bare-ssh reachability test.
 40func Commands() []Command { return registry }
 41
 42// Lookup resolves argv to a command by longest path match, returning the
 43// command and the remaining arguments.
 44func Lookup(argv []string) (Command, []string, bool) {
 45	best := -1
 46	var found Command
 47	for _, cmd := range registry {
 48		if len(cmd.Path) <= len(argv) && slices.Equal(cmd.Path, argv[:len(cmd.Path)]) && len(cmd.Path) > best {
 49			best = len(cmd.Path)
 50			found = cmd
 51		}
 52	}
 53	if best < 0 {
 54		return Command{}, nil, false
 55	}
 56	return found, argv[best:], true
 57}
 58
 59// Dispatch runs argv for an authenticated session. The dispatcher — not the
 60// handlers — enforces key scope: control commands require a full-scope key.
 61func Dispatch(c *Ctx, argv []string) int {
 62	if len(argv) == 0 {
 63		return c.fail(protocol.ExitUsage, "no command given; try: ssh <host> help")
 64	}
 65	cmd, rest, ok := Lookup(argv)
 66	if !ok {
 67		return c.fail(protocol.ExitUsage, "unknown command %q", argv[0])
 68	}
 69	if c.Scope != "full" {
 70		return c.fail(protocol.ExitDenied, "this key's scope (%s) does not allow control commands", c.Scope)
 71	}
 72	// Strip the global --json flag wherever it appears.
 73	args := rest[:0:0]
 74	for _, a := range rest {
 75		if a == "--json" {
 76			c.JSON = true
 77			continue
 78		}
 79		args = append(args, a)
 80	}
 81	if !cmd.ReadsStdin {
 82		c.Stdin = emptyReader{}
 83	}
 84	return cmd.Run(c, args)
 85}
 86
 87type emptyReader struct{}
 88
 89func (emptyReader) Read([]byte) (int, error) { return 0, io.EOF }
 90
 91// emit writes data as the command result: a JSON envelope under --json,
 92// otherwise via the plain formatter.
 93func (c *Ctx) emit(data any, plain func(w io.Writer)) int {
 94	if c.JSON {
 95		enc := json.NewEncoder(c.Stdout)
 96		enc.SetEscapeHTML(false)
 97		if err := enc.Encode(protocol.Envelope{ProtocolVersion: protocol.Version, Data: data}); err != nil {
 98			return protocol.ExitFailure
 99		}
100		return protocol.ExitOK
101	}
102	plain(c.Stdout)
103	return protocol.ExitOK
104}
105
106func (c *Ctx) fail(code int, format string, args ...any) int {
107	msg := fmt.Sprintf(format, args...)
108	if c.JSON {
109		enc := json.NewEncoder(c.Stdout)
110		enc.SetEscapeHTML(false)
111		enc.Encode(protocol.Envelope{ProtocolVersion: protocol.Version, Error: msg})
112	} else {
113		fmt.Fprintln(c.Stderr, msg)
114	}
115	return code
116}
117
118func init() {
119	register(Command{
120		Path:    []string{"help"},
121		Summary: "list available commands",
122		Run: func(c *Ctx, args []string) int {
123			for _, cmd := range registry {
124				fmt.Fprintf(c.Stdout, "%-24s %s\n", joinPath(cmd.Path), cmd.Summary)
125			}
126			return protocol.ExitOK
127		},
128	})
129}
130
131func joinPath(p []string) string {
132	out := ""
133	for i, s := range p {
134		if i > 0 {
135			out += " "
136		}
137		out += s
138	}
139	return out
140}