diff --git a/Makefile b/Makefile new file mode 100644 index 0000000..1e69fb0 --- /dev/null +++ b/Makefile @@ -0,0 +1,17 @@ +# crossbar gate. `make gate` must pass before any task is called done. It needs no network. + +.PHONY: gate build smoke + +gate: + @test -z "$$(gofmt -l . 2>&1)" || { echo "gofmt: these files need formatting:"; gofmt -l .; exit 1; } + go vet ./... + go test -race -count=1 ./... + sh scripts/check-lines.sh + @echo "gate: ok" + +build: + go build -o bin/ ./cmd/... + +# Runs the whole thing against two fake upstreams. Task 05 brings the script. +smoke: build + sh tools/smoke.sh diff --git a/docs/implementer-log.md b/docs/implementer-log.md index 90a7119..3abe907 100644 --- a/docs/implementer-log.md +++ b/docs/implementer-log.md @@ -5,5 +5,6 @@ owner fills in the Model column. The reviewer adds findings under "Reviews" once | Task | Date | Status | Gate runs | First gate | Deviations | Notes | Model | |---|---|---|---|---|---|---|---| +| v0/01-module-gate-config | 2026-09-25 | done | 1 | pass | none | `go mod download` fetched the module (network available); gate passed on the first run. | ? | ## Reviews diff --git a/go.mod b/go.mod new file mode 100644 index 0000000..ab77dc4 --- /dev/null +++ b/go.mod @@ -0,0 +1,5 @@ +module git.wntrmute.dev/kyle/crossbar + +go 1.26 + +require github.com/BurntSushi/toml v1.6.0 diff --git a/go.sum b/go.sum new file mode 100644 index 0000000..f74b269 --- /dev/null +++ b/go.sum @@ -0,0 +1,2 @@ +github.com/BurntSushi/toml v1.6.0 h1:dRaEfpa2VI55EwlIW72hMRHdWouJeRF7TPYhI+AUQjk= +github.com/BurntSushi/toml v1.6.0/go.mod h1:ukJfTF/6rtPPRCnwkur4qwRxa8vTRFBF0uk2lLoLwho= diff --git a/internal/config/config.go b/internal/config/config.go new file mode 100644 index 0000000..1795333 --- /dev/null +++ b/internal/config/config.go @@ -0,0 +1,300 @@ +// Package config reads crossbar's TOML file: the hosts in front of which it +// proxies, the model each serves, and the ordered routes that select them. +// +// The reader is strict. A misspelt key or a route naming a host that does not +// exist is an *Error naming the offending field, returned at start-up rather +// than surfaced later. +package config + +import ( + "errors" + "fmt" + "io" + "net" + "net/url" + "os" + "regexp" + "sort" + "strings" + "time" + + "github.com/BurntSushi/toml" +) + +// Duration is a time.Duration that TOML reads from a string such as "60s" or +// "30m". +type Duration struct{ time.Duration } + +// UnmarshalText implements encoding.TextUnmarshaler via time.ParseDuration. +func (d *Duration) UnmarshalText(text []byte) error { + dt, err := time.ParseDuration(string(text)) + if err != nil { + return err + } + d.Duration = dt + return nil +} + +// Model is the per-model tuning carried by a host entry. +type Model struct { + Parallel int `toml:"parallel"` +} + +// Host names one upstream llama-server, the models it serves, and its weight. +type Host struct { + BaseURL string `toml:"base_url"` + Weight float64 `toml:"weight"` + Models map[string]Model `toml:"models"` +} + +// Route is an ordered list of hosts to try, with an optional default model. +type Route struct { + Hosts []string `toml:"hosts"` + DefaultModel string `toml:"default_model"` +} + +// Config is the whole file: what to listen on, tuning, hosts and routes. +type Config struct { + Listen string `toml:"listen"` + PollInterval Duration `toml:"poll_interval"` + QueueMax int `toml:"queue_max"` + Hosts map[string]Host `toml:"hosts"` + Routes map[string]Route `toml:"routes"` +} + +// Error is a validation error naming the field it is about. +type Error struct { + Field, Msg string +} + +// Error implements the error interface. +func (e *Error) Error() string { + return "config: " + e.Field + ": " + e.Msg +} + +const ( + DefaultPollInterval = 60 * time.Second + DefaultQueueMax = 8 + MinPollInterval = time.Second +) + +var routeName = regexp.MustCompile(`^[a-z0-9][a-z0-9-]*$`) + +// Load reads and parses the config file at path. An open failure is wrapped as +// "config: …", the same shape as a decode failure. +func Load(path string) (*Config, error) { + f, err := os.Open(path) + if err != nil { + return nil, fmt.Errorf("config: %w", err) + } + defer f.Close() + return Parse(f) +} + +// Parse decodes TOML from r, applies defaults, and validates. A TOML syntax +// error is returned wrapped as "config: …" and is not an *Error; an unknown +// key is an *Error naming the first undecoded key in sorted order. +func Parse(r io.Reader) (*Config, error) { + var c Config + md, err := toml.NewDecoder(r).Decode(&c) + if err != nil { + return nil, fmt.Errorf("config: %w", err) + } + if undecoded := md.Undecoded(); len(undecoded) > 0 { + keys := make([]string, 0, len(undecoded)) + for _, k := range undecoded { + keys = append(keys, strings.Join(k, ".")) + } + sort.Strings(keys) + return nil, &Error{Field: keys[0], Msg: "unknown key"} + } + + if c.PollInterval.Duration == 0 { + c.PollInterval.Duration = DefaultPollInterval + } + if c.QueueMax == 0 { + c.QueueMax = DefaultQueueMax + } + if e := c.validate(); e != nil { + return nil, e + } + return &c, nil +} + +// Serves reports whether the host exists and lists the model. +func (c *Config) Serves(host, model string) bool { + h, ok := c.Hosts[host] + if !ok { + return false + } + _, ok = h.Models[model] + return ok +} + +// IsError extracts an *Error from err, reporting whether one was present. +func IsError(err error) (*Error, bool) { + var e *Error + if errors.As(err, &e) { + return e, true + } + return nil, false +} + +// validate checks the config in a fixed order and writes defaults back into c. +// The first problem wins; every problem is an *Error with a precise field. +func (c *Config) validate() *Error { + if e := c.checkListen(); e != nil { + return e + } + if e := c.checkPoll(); e != nil { + return e + } + if e := c.checkQueue(); e != nil { + return e + } + if e := c.checkHosts(); e != nil { + return e + } + return c.checkRoutes() +} + +func (c *Config) checkListen() *Error { + if c.Listen == "" { + return &Error{Field: "listen", Msg: "required, host:port"} + } + host, _, err := net.SplitHostPort(c.Listen) + if err != nil { + return &Error{Field: "listen", Msg: "must be host:port"} + } + if host == "" { + return &Error{Field: "listen", Msg: "host part required"} + } + if host == "0.0.0.0" || host == "::" { + return &Error{Field: "listen", Msg: "not an unspecified address"} + } + return nil +} + +func (c *Config) checkPoll() *Error { + if c.PollInterval.Duration == 0 { + c.PollInterval.Duration = DefaultPollInterval + } else if c.PollInterval.Duration < MinPollInterval { + return &Error{Field: "poll_interval", Msg: "must be at least 1s"} + } + return nil +} + +func (c *Config) checkQueue() *Error { + if c.QueueMax == 0 { + c.QueueMax = DefaultQueueMax + } else if c.QueueMax < 0 { + return &Error{Field: "queue_max", Msg: "must not be negative"} + } + return nil +} + +func (c *Config) checkHosts() *Error { + if len(c.Hosts) == 0 { + return &Error{Field: "hosts", Msg: "at least one required"} + } + names := make([]string, 0, len(c.Hosts)) + for name := range c.Hosts { + names = append(names, name) + } + sort.Strings(names) + for _, name := range names { + h := c.Hosts[name] + + baseField := fmt.Sprintf("hosts.%s.base_url", name) + u, err := url.Parse(h.BaseURL) + if err != nil { + return &Error{Field: baseField, Msg: err.Error()} + } + if u.Scheme != "http" && u.Scheme != "https" { + return &Error{Field: baseField, Msg: "scheme must be http or https"} + } + if u.Host == "" { + return &Error{Field: baseField, Msg: "host required"} + } + if u.RawQuery != "" { + return &Error{Field: baseField, Msg: "query not allowed"} + } + if u.Fragment != "" { + return &Error{Field: baseField, Msg: "fragment not allowed"} + } + h.BaseURL = strings.TrimRight(h.BaseURL, "/") + + if h.Weight == 0 { + h.Weight = 1 + } else if h.Weight < 0 { + return &Error{Field: fmt.Sprintf("hosts.%s.weight", name), Msg: "must not be negative"} + } + + if len(h.Models) == 0 { + return &Error{Field: fmt.Sprintf("hosts.%s.models", name), Msg: "at least one required"} + } + models := make([]string, 0, len(h.Models)) + for m := range h.Models { + models = append(models, m) + } + sort.Strings(models) + for _, m := range models { + pm := h.Models[m] + if pm.Parallel == 0 { + pm.Parallel = 1 + } else if pm.Parallel < 0 { + return &Error{Field: fmt.Sprintf("hosts.%s.models.%s.parallel", name, m), Msg: "must not be negative"} + } + h.Models[m] = pm + } + c.Hosts[name] = h + } + return nil +} + +func (c *Config) checkRoutes() *Error { + if len(c.Routes) == 0 { + return &Error{Field: "routes", Msg: "at least one required"} + } + names := make([]string, 0, len(c.Routes)) + for name := range c.Routes { + names = append(names, name) + } + sort.Strings(names) + for _, name := range names { + r := c.Routes[name] + + if !routeName.MatchString(name) { + return &Error{Field: fmt.Sprintf("routes.%s", name), Msg: "must match [a-z0-9][a-z0-9-]*"} + } + + hostsField := fmt.Sprintf("routes.%s.hosts", name) + if len(r.Hosts) == 0 { + return &Error{Field: hostsField, Msg: "at least one required"} + } + seen := make(map[string]bool, len(r.Hosts)) + for _, h := range r.Hosts { + if seen[h] { + return &Error{Field: hostsField, Msg: "host listed twice"} + } + seen[h] = true + if _, ok := c.Hosts[h]; !ok { + return &Error{Field: hostsField, Msg: "unknown host"} + } + } + + if r.DefaultModel != "" { + served := false + for _, h := range r.Hosts { + if _, ok := c.Hosts[h].Models[r.DefaultModel]; ok { + served = true + break + } + } + if !served { + return &Error{Field: fmt.Sprintf("routes.%s.default_model", name), Msg: "not served by any host in route"} + } + } + } + return nil +} diff --git a/internal/config/config_test.go b/internal/config/config_test.go new file mode 100644 index 0000000..9069c36 --- /dev/null +++ b/internal/config/config_test.go @@ -0,0 +1,191 @@ +package config_test + +import ( + "fmt" + "path/filepath" + "strings" + "testing" + "time" + + "git.wntrmute.dev/kyle/crossbar/internal/config" +) + +func TestGoodFile(t *testing.T) { + c, err := config.Load(filepath.Join("testdata", "good.toml")) + if err != nil { + t.Fatalf("Load: %v", err) + } + if c.Listen != "100.64.0.9:7777" { + t.Errorf("Listen = %q", c.Listen) + } + if c.PollInterval.Duration != 5*time.Second { + t.Errorf("PollInterval = %v", c.PollInterval.Duration) + } + if c.QueueMax != 4 { + t.Errorf("QueueMax = %d", c.QueueMax) + } + alpha := c.Hosts["alpha"] + if alpha.BaseURL != "http://alpha.example:11434" { + t.Errorf("trailing slash not stripped: %q", alpha.BaseURL) + } + if alpha.Weight != 2 { + t.Errorf("alpha.Weight = %v", alpha.Weight) + } + if alpha.Models["ornith-1.5-35b-a3b"].Parallel != 4 || alpha.Models["small-9b"].Parallel != 6 { + t.Errorf("alpha.Models = %+v", alpha.Models) + } + beta := c.Hosts["beta"] + if beta.Weight != 1 { + t.Errorf("beta.Weight default = %v, want 1", beta.Weight) + } + if beta.Models["ornith-1.5-35b-a3b"].Parallel != 1 { + t.Errorf("beta parallel default = %d, want 1", beta.Models["ornith-1.5-35b-a3b"].Parallel) + } + r := c.Routes["opencode-a"] + if len(r.Hosts) != 2 || r.Hosts[0] != "alpha" || r.Hosts[1] != "beta" { + t.Errorf("route hosts = %v", r.Hosts) + } + if r.DefaultModel != "ornith-1.5-35b-a3b" { + t.Errorf("DefaultModel = %q", r.DefaultModel) + } + if c.Routes["hermes-x"].DefaultModel != "" { + t.Errorf("hermes-x DefaultModel should be empty") + } + if !c.Serves("alpha", "small-9b") || c.Serves("beta", "small-9b") || c.Serves("nope", "m") { + t.Errorf("Serves is wrong") + } +} + +func TestDefaults(t *testing.T) { + c, err := config.Parse(strings.NewReader(` +listen = "127.0.0.1:1" +[hosts.a] +base_url = "http://a:1" +models = { "m" = { } } +[routes.r] +hosts = ["a"] +`)) + if err != nil { + t.Fatalf("Parse: %v", err) + } + if c.PollInterval.Duration != config.DefaultPollInterval { + t.Errorf("PollInterval default = %v", c.PollInterval.Duration) + } + if c.QueueMax != config.DefaultQueueMax { + t.Errorf("QueueMax default = %d", c.QueueMax) + } +} + +func TestBadFiles(t *testing.T) { + cases := []struct{ file, field string }{ + {"bad-listen.toml", "listen"}, + {"bad-unknown-host.toml", "routes.r.hosts"}, + {"bad-default-model.toml", "routes.r.default_model"}, + {"bad-unknown-key.toml", "lease_idle"}, + } + for _, tc := range cases { + t.Run(tc.file, func(t *testing.T) { + _, err := config.Load(filepath.Join("testdata", tc.file)) + if err == nil { + t.Fatalf("want error") + } + e, ok := config.IsError(err) + if !ok { + t.Fatalf("want *config.Error, got %T: %v", err, err) + } + if e.Field != tc.field { + t.Errorf("Field = %q, want %q (%v)", e.Field, tc.field, err) + } + if !strings.HasPrefix(err.Error(), "config: "+tc.field+": ") { + t.Errorf("Error() = %q", err.Error()) + } + }) + } +} + +func TestBadValues(t *testing.T) { + base := ` +listen = %q +poll_interval = %q +[hosts.a] +base_url = %q +weight = %v +models = { "m" = { parallel = %d } } +[routes.%s] +hosts = ["a"] +` + cases := []struct { + name string + listen, poll, url, route string + weight float64 + parallel int + field string + }{ + {"empty listen", "", "5s", "http://a:1", "r", 1, 1, "listen"}, + {"no port", "127.0.0.1", "5s", "http://a:1", "r", 1, 1, "listen"}, + {"v6 any", "[::]:7", "5s", "http://a:1", "r", 1, 1, "listen"}, + {"poll too short", "127.0.0.1:7", "500ms", "http://a:1", "r", 1, 1, "poll_interval"}, + {"ftp url", "127.0.0.1:7", "5s", "ftp://a:1", "r", 1, 1, "hosts.a.base_url"}, + {"no host", "127.0.0.1:7", "5s", "http://", "r", 1, 1, "hosts.a.base_url"}, + {"query", "127.0.0.1:7", "5s", "http://a:1/v1?x=1", "r", 1, 1, "hosts.a.base_url"}, + {"negative weight", "127.0.0.1:7", "5s", "http://a:1", "r", -1, 1, "hosts.a.weight"}, + {"negative parallel", "127.0.0.1:7", "5s", "http://a:1", "r", 1, -2, "hosts.a.models.m.parallel"}, + {"route name", "127.0.0.1:7", "5s", "http://a:1", "Bad_Name", 1, 1, "routes.Bad_Name"}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + text := fmt.Sprintf(base, tc.listen, tc.poll, tc.url, tc.weight, tc.parallel, tc.route) + _, err := config.Parse(strings.NewReader(text)) + if err == nil { + t.Fatalf("want error for %s", tc.name) + } + e, ok := config.IsError(err) + if !ok { + t.Fatalf("want *config.Error, got %T: %v", err, err) + } + if e.Field != tc.field { + t.Errorf("Field = %q, want %q (%v)", e.Field, tc.field, err) + } + }) + } +} + +func TestMissingSections(t *testing.T) { + for _, tc := range []struct{ name, text, field string }{ + {"no hosts", "listen = \"127.0.0.1:7\"\n[routes.r]\nhosts = [\"a\"]\n", "hosts"}, + {"no routes", "listen = \"127.0.0.1:7\"\n[hosts.a]\nbase_url = \"http://a:1\"\nmodels = { \"m\" = { } }\n", "routes"}, + {"host without models", "listen = \"127.0.0.1:7\"\n[hosts.a]\nbase_url = \"http://a:1\"\n[routes.r]\nhosts = [\"a\"]\n", "hosts.a.models"}, + {"route without hosts", "listen = \"127.0.0.1:7\"\n[hosts.a]\nbase_url = \"http://a:1\"\nmodels = { \"m\" = { } }\n[routes.r]\n", "routes.r.hosts"}, + {"host twice", "listen = \"127.0.0.1:7\"\n[hosts.a]\nbase_url = \"http://a:1\"\nmodels = { \"m\" = { } }\n[routes.r]\nhosts = [\"a\", \"a\"]\n", "routes.r.hosts"}, + } { + t.Run(tc.name, func(t *testing.T) { + _, err := config.Parse(strings.NewReader(tc.text)) + e, ok := config.IsError(err) + if !ok { + t.Fatalf("want *config.Error, got %v", err) + } + if e.Field != tc.field { + t.Errorf("Field = %q, want %q", e.Field, tc.field) + } + }) + } +} + +func TestNotTOML(t *testing.T) { + _, err := config.Parse(strings.NewReader("listen = [unterminated")) + if err == nil { + t.Fatal("want error") + } + if _, ok := config.IsError(err); ok { + t.Errorf("a syntax error is not a validation Error") + } + if !strings.HasPrefix(err.Error(), "config: ") { + t.Errorf("Error() = %q", err.Error()) + } +} + +func TestMissingFile(t *testing.T) { + if _, err := config.Load(filepath.Join("testdata", "does-not-exist.toml")); err == nil { + t.Fatal("want error") + } +} diff --git a/internal/config/testdata/bad-default-model.toml b/internal/config/testdata/bad-default-model.toml new file mode 100644 index 0000000..4e1cc32 --- /dev/null +++ b/internal/config/testdata/bad-default-model.toml @@ -0,0 +1,9 @@ +listen = "127.0.0.1:7777" + +[hosts.alpha] +base_url = "http://alpha.example:11434" +models = { "m" = { } } + +[routes.r] +hosts = ["alpha"] +default_model = "not-served" diff --git a/internal/config/testdata/bad-listen.toml b/internal/config/testdata/bad-listen.toml new file mode 100644 index 0000000..15e0e6d --- /dev/null +++ b/internal/config/testdata/bad-listen.toml @@ -0,0 +1,8 @@ +listen = "0.0.0.0:7777" + +[hosts.alpha] +base_url = "http://alpha.example:11434" +models = { "m" = { } } + +[routes.r] +hosts = ["alpha"] diff --git a/internal/config/testdata/bad-unknown-host.toml b/internal/config/testdata/bad-unknown-host.toml new file mode 100644 index 0000000..fd5aa92 --- /dev/null +++ b/internal/config/testdata/bad-unknown-host.toml @@ -0,0 +1,8 @@ +listen = "127.0.0.1:7777" + +[hosts.alpha] +base_url = "http://alpha.example:11434" +models = { "m" = { } } + +[routes.r] +hosts = ["alpha", "gamma"] diff --git a/internal/config/testdata/bad-unknown-key.toml b/internal/config/testdata/bad-unknown-key.toml new file mode 100644 index 0000000..e3e94ec --- /dev/null +++ b/internal/config/testdata/bad-unknown-key.toml @@ -0,0 +1,9 @@ +listen = "127.0.0.1:7777" +lease_idle = "30m" + +[hosts.alpha] +base_url = "http://alpha.example:11434" +models = { "m" = { } } + +[routes.r] +hosts = ["alpha"] diff --git a/internal/config/testdata/good.toml b/internal/config/testdata/good.toml new file mode 100644 index 0000000..92ad98e --- /dev/null +++ b/internal/config/testdata/good.toml @@ -0,0 +1,19 @@ +listen = "100.64.0.9:7777" +poll_interval = "5s" +queue_max = 4 + +[hosts.alpha] +base_url = "http://alpha.example:11434/" +weight = 2.0 +models = { "ornith-1.5-35b-a3b" = { parallel = 4 }, "small-9b" = { parallel = 6 } } + +[hosts.beta] +base_url = "https://beta.example:8081" +models = { "ornith-1.5-35b-a3b" = { } } + +[routes.opencode-a] +hosts = ["alpha", "beta"] +default_model = "ornith-1.5-35b-a3b" + +[routes.hermes-x] +hosts = ["beta"] diff --git a/scripts/check-lines.sh b/scripts/check-lines.sh new file mode 100644 index 0000000..c8e113e --- /dev/null +++ b/scripts/check-lines.sh @@ -0,0 +1,14 @@ +#!/bin/sh +# No Go source file over 400 lines. Reports every offender, then fails. +set -u +limit=400 +bad=0 +for f in $(find . -name '*.go' -not -path './.git/*' -not -path './vendor/*'); do + n=$(wc -l < "$f") || { echo "check-lines: cannot read $f" >&2; exit 2; } + if [ "$n" -gt "$limit" ]; then + echo "check-lines: $f has $n lines (limit $limit)" + bad=1 + fi +done +[ "$bad" -eq 0 ] || exit 1 +echo "check-lines: ok"