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

399 lines
10 KiB
Go

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
}