Files

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
}