package config import ( "errors" "net" "path/filepath" "strconv" "strings" "time" ) const ( HTTPAddressEnvironment = "CMROUBAO_HTTP_ADDR" DatabasePathEnvironment = "CMROUBAO_DATABASE_PATH" AssetDirectoryEnvironment = "CMROUBAO_ASSET_DIR" TLSCertificateEnvironment = "CMROUBAO_TLS_CERT_FILE" TLSPrivateKeyEnvironment = "CMROUBAO_TLS_KEY_FILE" ClaimLeaseEnvironment = "CMROUBAO_CLAIM_LEASE" RunningLeaseEnvironment = "CMROUBAO_RUNNING_LEASE" ReadinessTTLEnvironment = "CMROUBAO_READINESS_TTL" defaultHTTPAddress = "127.0.0.1:8080" defaultDatabasePath = "var/cmroubao.db" defaultAssetDirectory = "var/assets" defaultClaimLease = 10 * time.Minute defaultRunningLease = 30 * time.Minute defaultReadinessTTL = 2 * time.Minute ) type LookupEnvironment func(string) (string, bool) type Config struct { HTTPAddress string DatabasePath string AssetDirectory string TLSCertificate string TLSPrivateKey string ReadHeaderTimeout time.Duration ReadTimeout time.Duration WriteTimeout time.Duration IdleTimeout time.Duration ShutdownTimeout time.Duration MaxHeaderBytes int ClaimLease time.Duration RunningLease time.Duration ReadinessTTL time.Duration } func Load(lookup LookupEnvironment) (Config, error) { httpAddress, err := environmentValue( lookup, HTTPAddressEnvironment, defaultHTTPAddress, ) if err != nil { return Config{}, err } if err := validateHTTPAddress(httpAddress); err != nil { return Config{}, err } databasePath, err := environmentValue( lookup, DatabasePathEnvironment, defaultDatabasePath, ) databasePath, err = validatedDatabasePath(databasePath, err) if err != nil { return Config{}, err } assetDirectory, err := environmentValue( lookup, AssetDirectoryEnvironment, defaultAssetDirectory, ) if err != nil { return Config{}, err } assetDirectory, err = validatedAssetDirectory(assetDirectory) if err != nil { return Config{}, err } tlsCertificate, certificateSet, err := optionalEnvironmentValue( lookup, TLSCertificateEnvironment, ) if err != nil { return Config{}, err } tlsPrivateKey, privateKeySet, err := optionalEnvironmentValue( lookup, TLSPrivateKeyEnvironment, ) if err != nil { return Config{}, err } if certificateSet != privateKeySet { return Config{}, errors.New( TLSCertificateEnvironment + " and " + TLSPrivateKeyEnvironment + " must be set together", ) } if !certificateSet && !isLoopbackAddress(httpAddress) { return Config{}, errors.New( HTTPAddressEnvironment + " must use loopback unless TLS is configured", ) } claimLease, err := durationEnvironment( lookup, ClaimLeaseEnvironment, defaultClaimLease, time.Minute, 30*time.Minute, ) if err != nil { return Config{}, err } runningLease, err := durationEnvironment( lookup, RunningLeaseEnvironment, defaultRunningLease, 5*time.Minute, 120*time.Minute, ) if err != nil { return Config{}, err } readinessTTL, err := durationEnvironment( lookup, ReadinessTTLEnvironment, defaultReadinessTTL, 30*time.Second, 10*time.Minute, ) if err != nil { return Config{}, err } return Config{ HTTPAddress: httpAddress, DatabasePath: filepath.Clean(databasePath), AssetDirectory: assetDirectory, TLSCertificate: cleanOptionalPath(tlsCertificate), TLSPrivateKey: cleanOptionalPath(tlsPrivateKey), ReadHeaderTimeout: 5 * time.Second, ReadTimeout: 15 * time.Second, WriteTimeout: 30 * time.Second, IdleTimeout: 60 * time.Second, ShutdownTimeout: 10 * time.Second, MaxHeaderBytes: 1 << 20, ClaimLease: claimLease, RunningLease: runningLease, ReadinessTTL: readinessTTL, }, nil } func durationEnvironment( lookup LookupEnvironment, name string, defaultValue time.Duration, minimum time.Duration, maximum time.Duration, ) (time.Duration, error) { value, exists := lookup(name) if !exists { return defaultValue, nil } value = strings.TrimSpace(value) duration, err := time.ParseDuration(value) if err != nil || duration < minimum || duration > maximum { return 0, errors.New( name + " must be a duration between " + minimum.String() + " and " + maximum.String(), ) } return duration, nil } func cleanOptionalPath(value string) string { if value == "" { return "" } return filepath.Clean(value) } func optionalEnvironmentValue( lookup LookupEnvironment, name string, ) (string, bool, error) { value, exists := lookup(name) if !exists { return "", false, nil } value = strings.TrimSpace(value) if value == "" { return "", false, errors.New(name + " must not be blank") } if strings.ContainsRune(value, '\x00') { return "", false, errors.New(name + " contains an invalid character") } return value, true, nil } func isLoopbackAddress(address string) bool { host, _, err := net.SplitHostPort(address) if err != nil { return false } if strings.EqualFold(host, "localhost") { return true } ip := net.ParseIP(host) return ip != nil && ip.IsLoopback() } func validatedAssetDirectory(path string) (string, error) { if strings.ContainsRune(path, '\x00') { return "", errors.New( AssetDirectoryEnvironment + " contains an invalid character", ) } cleanPath := filepath.Clean(path) volumeRoot := filepath.VolumeName(cleanPath) + string(filepath.Separator) if cleanPath == "." || cleanPath == string(filepath.Separator) || cleanPath == volumeRoot { return "", errors.New( AssetDirectoryEnvironment + " must be a dedicated directory", ) } return cleanPath, nil } func LoadDatabasePath(lookup LookupEnvironment) (string, error) { databasePath, err := environmentValue( lookup, DatabasePathEnvironment, defaultDatabasePath, ) return validatedDatabasePath(databasePath, err) } func environmentValue( lookup LookupEnvironment, name string, defaultValue string, ) (string, error) { value, exists := lookup(name) if !exists { return defaultValue, nil } value = strings.TrimSpace(value) if value == "" { return "", errors.New(name + " must not be blank") } return value, nil } func validateHTTPAddress(address string) error { host, portValue, err := net.SplitHostPort(address) if err != nil || strings.TrimSpace(host) == "" { return errors.New(HTTPAddressEnvironment + " must include a host and port") } port, err := strconv.Atoi(portValue) if err != nil || port < 1 || port > 65535 { return errors.New(HTTPAddressEnvironment + " port must be between 1 and 65535") } return nil } func validatedDatabasePath(path string, previousError error) (string, error) { if previousError != nil { return "", previousError } if strings.ContainsRune(path, '\x00') { return "", errors.New( DatabasePathEnvironment + " contains an invalid character", ) } lowerPath := strings.ToLower(path) cleanPath := filepath.Clean(path) extension := strings.ToLower(filepath.Ext(cleanPath)) if cleanPath == "." || cleanPath == string(filepath.Separator) || lowerPath == ":memory:" || strings.HasPrefix(lowerPath, "file:") || (extension != ".db" && extension != ".sqlite" && extension != ".sqlite3") { return "", errors.New( DatabasePathEnvironment + " must be a SQLite file path", ) } return cleanPath, nil }