399 lines
10 KiB
Go
399 lines
10 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"
|
||
"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"`
|
||
Wake *Wake `toml:"wake"`
|
||
}
|
||
|
||
// Route is an ordered list of hosts to try, with an optional default model and
|
||
// the peers allowed to reach it.
|
||
type Route struct {
|
||
Hosts []string `toml:"hosts"`
|
||
DefaultModel string `toml:"default_model"`
|
||
Peers []string `toml:"peers"`
|
||
}
|
||
|
||
// 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"`
|
||
Identity string `toml:"identity"`
|
||
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
|
||
|
||
DefaultIdentity = "off"
|
||
)
|
||
|
||
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 c.Identity == "" {
|
||
c.Identity = DefaultIdentity
|
||
}
|
||
if e := c.validate(md); 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(md toml.MetaData) *Error {
|
||
peersDefined := make(map[string]bool, len(c.Routes))
|
||
for name := range c.Routes {
|
||
if md.IsDefined("routes", name, "peers") {
|
||
peersDefined[name] = true
|
||
}
|
||
}
|
||
identityDefined := md.IsDefined("identity")
|
||
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
|
||
}
|
||
if e := c.checkWake(); e != nil {
|
||
return e
|
||
}
|
||
if e := c.checkRoutes(peersDefined, identityDefined); e != nil {
|
||
return e
|
||
}
|
||
return c.checkIdentity()
|
||
}
|
||
|
||
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(peersDefined map[string]bool, identityDefined bool) *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"}
|
||
}
|
||
}
|
||
|
||
if e := checkPeers(name, r.Peers, peersDefined[name], identityDefined, c.Identity); e != nil {
|
||
return e
|
||
}
|
||
}
|
||
return nil
|
||
}
|