360 lines
14 KiB
Go
360 lines
14 KiB
Go
package config
|
|
|
|
import (
|
|
"crypto/rand"
|
|
"encoding/base64"
|
|
"errors"
|
|
"fmt"
|
|
"net"
|
|
"net/netip"
|
|
"os"
|
|
"strconv"
|
|
"strings"
|
|
"time"
|
|
)
|
|
|
|
type Environment string
|
|
|
|
const (
|
|
Development Environment = "development"
|
|
Test Environment = "test"
|
|
Production Environment = "production"
|
|
)
|
|
|
|
type Config struct {
|
|
Environment Environment
|
|
DBDSN string
|
|
ListenAddress string
|
|
TrustedProxyCIDRs []netip.Prefix
|
|
StorageRoot string
|
|
SessionKey string
|
|
SessionTTL time.Duration
|
|
LoginAttempts int
|
|
LoginWindow time.Duration
|
|
RegistrationAttempts uint64
|
|
RegistrationWindow time.Duration
|
|
MaxPromptBytes int
|
|
MaxImages int
|
|
MaxImageBytes int64
|
|
MaxUploadBytes int64
|
|
MaxImagePixels uint64
|
|
HistoryLimit int
|
|
WorkerLeaseDuration time.Duration
|
|
WorkerPollInterval time.Duration
|
|
ProviderHTTPTimeout time.Duration
|
|
ProviderMaxResponseBytes int64
|
|
ProviderAllowedPorts []uint16
|
|
UserRateLimitCapacity uint64
|
|
UserRateLimitWindow time.Duration
|
|
APIKeyRateLimitCapacity uint64
|
|
APIKeyRateLimitWindow time.Duration
|
|
ProviderRateLimitCapacity uint64
|
|
ProviderRateLimitWindow time.Duration
|
|
TestDisableWorker bool
|
|
registrationExplicit bool
|
|
}
|
|
|
|
func Load() (Config, error) {
|
|
return LoadFromLookup(os.LookupEnv)
|
|
}
|
|
|
|
func LoadFromLookup(lookup func(string) (string, bool)) (Config, error) {
|
|
return loadFromLookup(lookup, defaultRuntimeSettings())
|
|
}
|
|
|
|
func LoadFile(path string) (Config, error) {
|
|
return LoadFileFromLookup(path, os.LookupEnv)
|
|
}
|
|
|
|
func LoadFileFromLookup(path string, lookup func(string) (string, bool)) (Config, error) {
|
|
settings, err := loadRuntimeSettings(path)
|
|
if err != nil {
|
|
return Config{}, err
|
|
}
|
|
return loadFromLookup(lookup, settings)
|
|
}
|
|
|
|
func loadFromLookup(lookup func(string) (string, bool), defaults runtimeSettings) (Config, error) {
|
|
read := func(name string) string {
|
|
value, _ := lookup(name)
|
|
return strings.TrimSpace(value)
|
|
}
|
|
|
|
environment := Environment(read("CHORUS_ENV"))
|
|
if environment == "" {
|
|
environment = Development
|
|
}
|
|
if environment != Development && environment != Test && environment != Production {
|
|
return Config{}, fmt.Errorf("CHORUS_ENV must be development, test, or production")
|
|
}
|
|
|
|
listenAddress := read("CHORUS_LISTEN_ADDRESS")
|
|
if listenAddress == "" && defaults.Server.ListenAddress != nil {
|
|
listenAddress = strings.TrimSpace(*defaults.Server.ListenAddress)
|
|
}
|
|
if listenAddress == "" && defaults.Server.ListenAddress == nil {
|
|
listenAddress = "127.0.0.1:8080"
|
|
}
|
|
|
|
cfg := Config{
|
|
Environment: environment,
|
|
DBDSN: read("CHORUS_DSN"),
|
|
ListenAddress: listenAddress,
|
|
StorageRoot: read("CHORUS_STORAGE_ROOT"),
|
|
SessionKey: read("CHORUS_SESSION_KEY"),
|
|
}
|
|
var err error
|
|
trustedProxyCIDRs := read("CHORUS_TRUSTED_PROXY_CIDRS")
|
|
if trustedProxyCIDRs == "" {
|
|
trustedProxyCIDRs = strings.Join(defaults.Server.TrustedProxyCIDRs, ",")
|
|
}
|
|
if cfg.TrustedProxyCIDRs, err = parseTrustedProxyCIDRs(trustedProxyCIDRs); err != nil {
|
|
return Config{}, fmt.Errorf("CHORUS_TRUSTED_PROXY_CIDRS is invalid")
|
|
}
|
|
if environment != Production {
|
|
if cfg.StorageRoot == "" {
|
|
cfg.StorageRoot = "var/storage"
|
|
}
|
|
if cfg.SessionKey == "" {
|
|
generated := make([]byte, 32)
|
|
if _, err := rand.Read(generated); err != nil {
|
|
return Config{}, fmt.Errorf("generate development session key")
|
|
}
|
|
cfg.SessionKey = base64.RawURLEncoding.EncodeToString(generated)
|
|
}
|
|
}
|
|
if err := validateListenAddress(cfg.ListenAddress); err != nil {
|
|
return Config{}, err
|
|
}
|
|
|
|
var missing []string
|
|
if cfg.DBDSN == "" {
|
|
missing = append(missing, "CHORUS_DSN")
|
|
}
|
|
if environment == Production {
|
|
for name, value := range map[string]string{
|
|
"CHORUS_SESSION_KEY": cfg.SessionKey,
|
|
"CHORUS_STORAGE_ROOT": cfg.StorageRoot,
|
|
} {
|
|
if value == "" {
|
|
missing = append(missing, name)
|
|
}
|
|
}
|
|
for _, name := range []string{"CHORUS_SESSION_TTL_MINUTES", "CHORUS_LOGIN_ATTEMPTS", "CHORUS_LOGIN_WINDOW_SECONDS", "CHORUS_MAX_PROMPT_BYTES", "CHORUS_MAX_IMAGES", "CHORUS_MAX_IMAGE_BYTES", "CHORUS_MAX_UPLOAD_BYTES", "CHORUS_MAX_IMAGE_PIXELS", "CHORUS_HISTORY_LIMIT", "CHORUS_USER_RATE_LIMIT_CAPACITY", "CHORUS_USER_RATE_LIMIT_WINDOW_SECONDS", "CHORUS_API_KEY_RATE_LIMIT_CAPACITY", "CHORUS_API_KEY_RATE_LIMIT_WINDOW_SECONDS", "CHORUS_PROVIDER_RATE_LIMIT_CAPACITY", "CHORUS_PROVIDER_RATE_LIMIT_WINDOW_SECONDS"} {
|
|
if read(name) == "" {
|
|
missing = append(missing, name)
|
|
}
|
|
}
|
|
if defaults.Registration.Attempts == nil && read("CHORUS_REGISTRATION_ATTEMPTS") == "" {
|
|
missing = append(missing, "CHORUS_REGISTRATION_ATTEMPTS")
|
|
}
|
|
if defaults.Registration.WindowSeconds == nil && read("CHORUS_REGISTRATION_WINDOW_SECONDS") == "" {
|
|
missing = append(missing, "CHORUS_REGISTRATION_WINDOW_SECONDS")
|
|
}
|
|
}
|
|
if len(missing) > 0 {
|
|
return Config{}, fmt.Errorf("missing required configuration: %s", strings.Join(missing, ", "))
|
|
}
|
|
if cfg.SessionTTL, err = durationConfig(read("CHORUS_SESSION_TTL_MINUTES"), 480, time.Minute); err != nil {
|
|
return Config{}, fmt.Errorf("CHORUS_SESSION_TTL_MINUTES is invalid")
|
|
}
|
|
if cfg.LoginWindow, err = durationConfig(read("CHORUS_LOGIN_WINDOW_SECONDS"), 60, time.Second); err != nil {
|
|
return Config{}, fmt.Errorf("CHORUS_LOGIN_WINDOW_SECONDS is invalid")
|
|
}
|
|
if cfg.LoginAttempts, err = intConfig(read("CHORUS_LOGIN_ATTEMPTS"), 5); err != nil {
|
|
return Config{}, fmt.Errorf("CHORUS_LOGIN_ATTEMPTS is invalid")
|
|
}
|
|
registrationAttemptsRaw := read("CHORUS_REGISTRATION_ATTEMPTS")
|
|
registrationWindowRaw := read("CHORUS_REGISTRATION_WINDOW_SECONDS")
|
|
cfg.registrationExplicit = registrationAttemptsRaw != "" && registrationWindowRaw != ""
|
|
registrationAttemptsFallback := 5
|
|
registrationWindowFallback := 900
|
|
if defaults.Registration.Attempts != nil {
|
|
registrationAttemptsFallback = *defaults.Registration.Attempts
|
|
cfg.registrationExplicit = cfg.registrationExplicit || defaults.Registration.WindowSeconds != nil
|
|
}
|
|
if defaults.Registration.WindowSeconds != nil {
|
|
registrationWindowFallback = *defaults.Registration.WindowSeconds
|
|
}
|
|
registrationAttempts, err := intConfig(registrationAttemptsRaw, registrationAttemptsFallback)
|
|
if err != nil {
|
|
return Config{}, fmt.Errorf("CHORUS_REGISTRATION_ATTEMPTS is invalid")
|
|
}
|
|
cfg.RegistrationAttempts = uint64(registrationAttempts)
|
|
if cfg.RegistrationWindow, err = durationConfig(registrationWindowRaw, registrationWindowFallback, time.Second); err != nil {
|
|
return Config{}, fmt.Errorf("CHORUS_REGISTRATION_WINDOW_SECONDS is invalid")
|
|
}
|
|
if cfg.MaxPromptBytes, err = intConfig(read("CHORUS_MAX_PROMPT_BYTES"), 8000); err != nil {
|
|
return Config{}, fmt.Errorf("CHORUS_MAX_PROMPT_BYTES is invalid")
|
|
}
|
|
if cfg.MaxImages, err = intConfig(read("CHORUS_MAX_IMAGES"), 8); err != nil {
|
|
return Config{}, fmt.Errorf("CHORUS_MAX_IMAGES is invalid")
|
|
}
|
|
maxImageBytes, err := intConfig(read("CHORUS_MAX_IMAGE_BYTES"), 10<<20)
|
|
if err != nil {
|
|
return Config{}, fmt.Errorf("CHORUS_MAX_IMAGE_BYTES is invalid")
|
|
}
|
|
cfg.MaxImageBytes = int64(maxImageBytes)
|
|
maxUploadBytes, err := intConfig(read("CHORUS_MAX_UPLOAD_BYTES"), 32<<20)
|
|
if err != nil {
|
|
return Config{}, fmt.Errorf("CHORUS_MAX_UPLOAD_BYTES is invalid")
|
|
}
|
|
cfg.MaxUploadBytes = int64(maxUploadBytes)
|
|
maxPixels, err := intConfig(read("CHORUS_MAX_IMAGE_PIXELS"), 40_000_000)
|
|
if err != nil {
|
|
return Config{}, fmt.Errorf("CHORUS_MAX_IMAGE_PIXELS is invalid")
|
|
}
|
|
cfg.MaxImagePixels = uint64(maxPixels)
|
|
if cfg.HistoryLimit, err = intConfig(read("CHORUS_HISTORY_LIMIT"), 50); err != nil {
|
|
return Config{}, fmt.Errorf("CHORUS_HISTORY_LIMIT is invalid")
|
|
}
|
|
if cfg.WorkerLeaseDuration, err = durationConfig(read("CHORUS_WORKER_LEASE_SECONDS"), defaults.Worker.LeaseSeconds, time.Second); err != nil {
|
|
return Config{}, fmt.Errorf("CHORUS_WORKER_LEASE_SECONDS is invalid")
|
|
}
|
|
if cfg.WorkerPollInterval, err = durationConfig(read("CHORUS_WORKER_POLL_MILLISECONDS"), defaults.Worker.PollMilliseconds, time.Millisecond); err != nil {
|
|
return Config{}, fmt.Errorf("CHORUS_WORKER_POLL_MILLISECONDS is invalid")
|
|
}
|
|
if cfg.ProviderHTTPTimeout, err = durationConfig(read("CHORUS_PROVIDER_HTTP_TIMEOUT_SECONDS"), defaults.Provider.HTTPTimeoutSeconds, time.Second); err != nil {
|
|
return Config{}, fmt.Errorf("CHORUS_PROVIDER_HTTP_TIMEOUT_SECONDS is invalid")
|
|
}
|
|
providerMaxResponseBytes, err := intConfig(read("CHORUS_PROVIDER_MAX_RESPONSE_BYTES"), defaults.Provider.MaxResponseBytes)
|
|
if err != nil {
|
|
return Config{}, fmt.Errorf("CHORUS_PROVIDER_MAX_RESPONSE_BYTES is invalid")
|
|
}
|
|
cfg.ProviderMaxResponseBytes = int64(providerMaxResponseBytes)
|
|
if raw := read("CHORUS_PROVIDER_ALLOWED_PORTS"); raw != "" {
|
|
if cfg.ProviderAllowedPorts, err = ParseProviderAllowedPorts(raw); err != nil {
|
|
return Config{}, fmt.Errorf("CHORUS_PROVIDER_ALLOWED_PORTS is invalid")
|
|
}
|
|
} else {
|
|
cfg.ProviderAllowedPorts = append([]uint16(nil), defaults.Provider.AllowedPorts...)
|
|
}
|
|
userRateCapacity, err := intConfig(read("CHORUS_USER_RATE_LIMIT_CAPACITY"), 60)
|
|
if err != nil {
|
|
return Config{}, fmt.Errorf("CHORUS_USER_RATE_LIMIT_CAPACITY is invalid")
|
|
}
|
|
cfg.UserRateLimitCapacity = uint64(userRateCapacity)
|
|
if cfg.UserRateLimitWindow, err = durationConfig(read("CHORUS_USER_RATE_LIMIT_WINDOW_SECONDS"), 60, time.Second); err != nil {
|
|
return Config{}, fmt.Errorf("CHORUS_USER_RATE_LIMIT_WINDOW_SECONDS is invalid")
|
|
}
|
|
apiKeyRateCapacity, err := intConfig(read("CHORUS_API_KEY_RATE_LIMIT_CAPACITY"), 120)
|
|
if err != nil {
|
|
return Config{}, fmt.Errorf("CHORUS_API_KEY_RATE_LIMIT_CAPACITY is invalid")
|
|
}
|
|
cfg.APIKeyRateLimitCapacity = uint64(apiKeyRateCapacity)
|
|
if cfg.APIKeyRateLimitWindow, err = durationConfig(read("CHORUS_API_KEY_RATE_LIMIT_WINDOW_SECONDS"), 60, time.Second); err != nil {
|
|
return Config{}, fmt.Errorf("CHORUS_API_KEY_RATE_LIMIT_WINDOW_SECONDS is invalid")
|
|
}
|
|
providerRateCapacity, err := intConfig(read("CHORUS_PROVIDER_RATE_LIMIT_CAPACITY"), 60)
|
|
if err != nil {
|
|
return Config{}, fmt.Errorf("CHORUS_PROVIDER_RATE_LIMIT_CAPACITY is invalid")
|
|
}
|
|
cfg.ProviderRateLimitCapacity = uint64(providerRateCapacity)
|
|
if cfg.ProviderRateLimitWindow, err = durationConfig(read("CHORUS_PROVIDER_RATE_LIMIT_WINDOW_SECONDS"), 60, time.Second); err != nil {
|
|
return Config{}, fmt.Errorf("CHORUS_PROVIDER_RATE_LIMIT_WINDOW_SECONDS is invalid")
|
|
}
|
|
if value := read("CHORUS_TEST_DISABLE_WORKER"); value != "" {
|
|
if cfg.TestDisableWorker, err = strconv.ParseBool(value); err != nil {
|
|
return Config{}, fmt.Errorf("CHORUS_TEST_DISABLE_WORKER is invalid")
|
|
}
|
|
if cfg.TestDisableWorker && environment != Test {
|
|
return Config{}, fmt.Errorf("CHORUS_TEST_DISABLE_WORKER is only allowed in test")
|
|
}
|
|
}
|
|
if cfg.MaxUploadBytes < cfg.MaxImageBytes || cfg.MaxImages > 32 || cfg.HistoryLimit > 200 {
|
|
return Config{}, fmt.Errorf("portal limits are inconsistent")
|
|
}
|
|
if cfg.WorkerLeaseDuration <= cfg.ProviderHTTPTimeout+5*time.Second {
|
|
return Config{}, fmt.Errorf("worker lease must exceed provider HTTP timeout by more than 5 seconds")
|
|
}
|
|
return cfg, nil
|
|
}
|
|
|
|
func validateListenAddress(value string) error {
|
|
host, portText, err := net.SplitHostPort(value)
|
|
if err != nil || host == "" || strings.TrimSpace(host) != host {
|
|
return errors.New("portal listen address must use host:port")
|
|
}
|
|
if portText == "" || strings.Trim(portText, "0123456789") != "" {
|
|
return errors.New("portal listen address contains an invalid port")
|
|
}
|
|
port, err := strconv.Atoi(portText)
|
|
if err != nil || port < 1 || port > 65535 {
|
|
return errors.New("portal listen address contains an invalid port")
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func (c Config) ValidateProduction() error {
|
|
if c.Environment != Production {
|
|
return nil
|
|
}
|
|
if c.DBDSN == "" || len(c.SessionKey) < 32 || c.StorageRoot == "" || c.SessionTTL <= 0 || c.MaxImageBytes <= 0 || c.MaxUploadBytes < c.MaxImageBytes || c.WorkerLeaseDuration <= c.ProviderHTTPTimeout+5*time.Second || c.ProviderMaxResponseBytes <= 0 || len(c.ProviderAllowedPorts) == 0 || c.UserRateLimitCapacity == 0 || c.UserRateLimitWindow <= 0 || c.APIKeyRateLimitCapacity == 0 || c.APIKeyRateLimitWindow <= 0 || c.ProviderRateLimitCapacity == 0 || c.ProviderRateLimitWindow <= 0 || c.RegistrationAttempts == 0 || c.RegistrationWindow <= 0 || !c.registrationExplicit {
|
|
return errors.New("production configuration is incomplete")
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func parseTrustedProxyCIDRs(raw string) ([]netip.Prefix, error) {
|
|
if strings.TrimSpace(raw) == "" {
|
|
return nil, nil
|
|
}
|
|
parts := strings.Split(raw, ",")
|
|
prefixes := make([]netip.Prefix, 0, len(parts))
|
|
seen := make(map[netip.Prefix]struct{}, len(parts))
|
|
for _, part := range parts {
|
|
prefix, err := netip.ParsePrefix(strings.TrimSpace(part))
|
|
if err != nil {
|
|
return nil, errors.New("invalid trusted proxy CIDR")
|
|
}
|
|
prefix = prefix.Masked()
|
|
if _, exists := seen[prefix]; exists {
|
|
continue
|
|
}
|
|
seen[prefix] = struct{}{}
|
|
prefixes = append(prefixes, prefix)
|
|
}
|
|
return prefixes, nil
|
|
}
|
|
|
|
func ParseProviderAllowedPorts(raw string) ([]uint16, error) {
|
|
if strings.TrimSpace(raw) == "" {
|
|
return []uint16{80, 443}, nil
|
|
}
|
|
parts := strings.Split(raw, ",")
|
|
ports := make([]uint16, 0, len(parts))
|
|
seen := make(map[uint16]struct{}, len(parts))
|
|
for _, part := range parts {
|
|
value := strings.TrimSpace(part)
|
|
parsed, err := strconv.ParseUint(value, 10, 16)
|
|
if err != nil || parsed == 0 {
|
|
return nil, errors.New("invalid provider port list")
|
|
}
|
|
port := uint16(parsed)
|
|
if _, exists := seen[port]; exists {
|
|
continue
|
|
}
|
|
seen[port] = struct{}{}
|
|
ports = append(ports, port)
|
|
}
|
|
return ports, nil
|
|
}
|
|
|
|
func intConfig(value string, fallback int) (int, error) {
|
|
if value == "" {
|
|
return fallback, nil
|
|
}
|
|
parsed, err := strconv.Atoi(value)
|
|
if err != nil || parsed <= 0 {
|
|
return 0, errors.New("invalid positive integer")
|
|
}
|
|
return parsed, nil
|
|
}
|
|
func durationConfig(value string, fallback int, unit time.Duration) (time.Duration, error) {
|
|
parsed, err := intConfig(value, fallback)
|
|
if err != nil {
|
|
return 0, err
|
|
}
|
|
return time.Duration(parsed) * unit, nil
|
|
}
|