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

232 lines
5.6 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"
defaultHTTPAddress = "127.0.0.1:8080"
defaultDatabasePath = "var/cmroubao.db"
defaultAssetDirectory = "var/assets"
)
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
}
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",
)
}
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,
}, 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
}