Files
cmroubao/backend-api/internal/config/config.go
T

296 lines
7.2 KiB
Go

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
}