Files
kyleandClaude Opus 5.5 aea2eeae2c Routes may share one lease (affinity = "route") and skip crossbar's queue (queue = false)
route.go gains Route.Affinity/Queue with PerRoute(), Queues() and affinity
validation (checkRoutes moved here; config.go calls it once). limiter.Track
counts a request without holding or refusing it; a release hands the slot to a
waiter only while in flight <= parallel. The proxy leases a PerRoute() route
under an empty fingerprint (the row keeps the real one) and uses Track when
Queues() is false.

Implemented by Ornith (OpenCode); owner review removed a release-on-first-flush
workaround for a race in the owner's given test (see implementer log).

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
2026-09-25 19:33:46 -07:00

338 lines
8.6 KiB
Go
Raw Permalink 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"`
}
// 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"
)
// 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
}