// 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" "strconv" "strings" "time" "github.com/BurntSushi/toml" ) // Duration is a time.Duration that TOML reads from a string such as "60s", // "30m", or "7d" (an integer number of days). type Duration struct{ time.Duration } // UnmarshalText implements encoding.TextUnmarshaler. It accepts the "Nd" form // — an integer number of days, so "7d" is 7 × 24h — in addition to // time.ParseDuration syntax. func (d *Duration) UnmarshalText(text []byte) error { if days, ok := parseDays(text); ok { d.Duration = days return nil } dt, err := time.ParseDuration(string(text)) if err != nil { return err } d.Duration = dt return nil } // dayPattern matches a run of digits followed by "d", e.g. "7d". var dayPattern = regexp.MustCompile(`^[0-9]+d$`) // parseDays reports whether text is the "Nd" day form and returns that many // hours. The regex guarantees the prefix is a base-10 integer. func parseDays(text []byte) (time.Duration, bool) { if !dayPattern.MatchString(string(text)) { return 0, false } n, err := strconv.Atoi(string(text[:len(text)-1])) if err != nil { return 0, false } return time.Duration(n) * 24 * time.Hour, true } // 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"` DB string `toml:"db"` LeaseIdle Duration `toml:"lease_idle"` Retention Duration `toml:"retention"` 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 DefaultDB = "crossbar.db" DefaultLeaseIdle = 30 * time.Minute DefaultRetention = 180 * 24 * time.Hour MinLeaseIdle = time.Minute MinRetention = 24 * time.Hour ) 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 !md.IsDefined("db") { c.DB = DefaultDB } if !md.IsDefined("lease_idle") { c.LeaseIdle.Duration = DefaultLeaseIdle } if !md.IsDefined("retention") { c.Retention.Duration = DefaultRetention } 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.checkDB(); e != nil { return e } if e := c.checkLeaseIdle(); e != nil { return e } if e := c.checkRetention(); 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) checkDB() *Error { if c.DB == "" { return &Error{Field: "db", Msg: "required"} } return nil } func (c *Config) checkLeaseIdle() *Error { if c.LeaseIdle.Duration < MinLeaseIdle { return &Error{Field: "lease_idle", Msg: "must be at least 1m"} } return nil } func (c *Config) checkRetention() *Error { if c.Retention.Duration < MinRetention { return &Error{Field: "retention", Msg: "must be at least 1d"} } 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 }