Files
kyle 73b24357dd Add the Go module, the gate and the config package
Implemented-By: OpenCode session (model recorded in docs/implementer-log.md)
2026-09-25 02:09:42 -07:00

301 lines
7.6 KiB
Go

// 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
}