Files
crossbar/internal/config/config.go
T
kyle 4e1dd03d07 Route templates: a route named x-* serves any request route x-<something>
Implemented-By: OpenCode session (model recorded in docs/implementer-log.md)
2026-09-25 13:38:24 -07:00

399 lines
10 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// 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) && !templateName.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
}