package config import ( "errors" "net" "net/url" "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" ShunyunbaoURLEnvironment = "CMROUBAO_SHUNYUNBAO_URL" ShunyunbaoUsernameEnvironment = "CMROUBAO_SHUNYUNBAO_USERNAME" ShunyunbaoPasswordEnvironment = "CMROUBAO_SHUNYUNBAO_PASSWORD" OCRAPIURLEnvironment = "CMROUBAO_OCR_API_URL" 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 defaultShunyunbaoURL = "https://www.shunyunbaoerp.com" ) 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 ShunyunbaoURL string ShunyunbaoUsername string ShunyunbaoPassword string OCRAPIURL string ShunyunbaoTimeout 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 } shunyunbaoURL, err := environmentValue( lookup, ShunyunbaoURLEnvironment, defaultShunyunbaoURL, ) if err != nil { return Config{}, err } if err := validateHTTPSOrigin(shunyunbaoURL, ShunyunbaoURLEnvironment); err != nil { return Config{}, err } shunyunbaoUsername, usernameSet, err := optionalEnvironmentValue( lookup, ShunyunbaoUsernameEnvironment, ) if err != nil { return Config{}, err } shunyunbaoPassword, passwordSet, err := optionalSecretEnvironmentValue( lookup, ShunyunbaoPasswordEnvironment, ) if err != nil { return Config{}, err } if usernameSet != passwordSet { return Config{}, errors.New( ShunyunbaoUsernameEnvironment + " and " + ShunyunbaoPasswordEnvironment + " must be set together", ) } ocrAPIURL, ocrAPISet, err := optionalEnvironmentValue(lookup, OCRAPIURLEnvironment) if err != nil { return Config{}, err } if ocrAPISet && !validOCRAPIURL(ocrAPIURL) { return Config{}, errors.New(OCRAPIURLEnvironment + " must be an approved OCR endpoint") } 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, ShunyunbaoURL: strings.TrimRight(shunyunbaoURL, "/"), ShunyunbaoUsername: shunyunbaoUsername, ShunyunbaoPassword: shunyunbaoPassword, OCRAPIURL: ocrAPIURL, ShunyunbaoTimeout: 30 * time.Second, }, nil } func validOCRAPIURL(value string) bool { parsed, err := url.Parse(value) if err != nil || parsed.Host == "" || parsed.User != nil || parsed.RawQuery != "" || parsed.Fragment != "" || parsed.Path == "" { return false } if port := parsed.Port(); port != "" { parsedPort, portErr := strconv.Atoi(port) if portErr != nil || parsedPort < 1 || parsedPort > 65535 { return false } } if parsed.Scheme == "https" { return true } if parsed.Scheme != "http" { return false } host := strings.Trim(parsed.Hostname(), "[]") return strings.EqualFold(host, "localhost") || (net.ParseIP(host) != nil && net.ParseIP(host).IsLoopback()) } func validateHTTPSOrigin(value, environment string) error { parsed, err := url.Parse(value) if err != nil || parsed.Scheme != "https" || parsed.Host == "" || parsed.User != nil || parsed.RawQuery != "" || parsed.Fragment != "" || (parsed.Path != "" && parsed.Path != "/") { return errors.New( environment + " must be an https origin without credentials or path", ) } return 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 optionalSecretEnvironmentValue( lookup LookupEnvironment, name string, ) (string, bool, error) { value, exists := lookup(name) if !exists { return "", false, nil } if value == "" || strings.ContainsRune(value, '\x00') { return "", false, errors.New(name + " must not be blank") } 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 }