104 lines
1.9 KiB
Go
104 lines
1.9 KiB
Go
package config
|
|
|
|
import (
|
|
"fmt"
|
|
"os"
|
|
"strings"
|
|
|
|
"github.com/spf13/viper"
|
|
)
|
|
|
|
func Load() (*Config, error) {
|
|
env := getEnv("APP_ENV", "dev")
|
|
|
|
v := viper.New()
|
|
v.SetConfigType("yaml")
|
|
|
|
// 配置路径
|
|
v.AddConfigPath("configs")
|
|
|
|
// ======================
|
|
// 1️⃣ 基础配置
|
|
// ======================
|
|
if err := loadBase(v); err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
// ======================
|
|
// 2️⃣ 环境配置(支持多环境)
|
|
// ======================
|
|
loadByEnv(v, env)
|
|
|
|
// ======================
|
|
// 3️⃣ 环境变量覆盖
|
|
// ======================
|
|
bindEnv(v)
|
|
|
|
// ======================
|
|
// 4️⃣ 默认值
|
|
// ======================
|
|
setDefaults(v, env)
|
|
|
|
// ======================
|
|
// 5️⃣ 解析
|
|
// ======================
|
|
var cfg Config
|
|
if err := v.Unmarshal(&cfg); err != nil {
|
|
return nil, fmt.Errorf("decode config failed: %w", err)
|
|
}
|
|
|
|
return &cfg, nil
|
|
}
|
|
|
|
func loadBase(v *viper.Viper) error {
|
|
v.SetConfigName("config")
|
|
|
|
if err := v.ReadInConfig(); err != nil {
|
|
return fmt.Errorf("read base config failed: %w", err)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func loadByEnv(v *viper.Viper, env string) {
|
|
envs := strings.Split(env, ",")
|
|
|
|
for _, e := range envs {
|
|
e = strings.TrimSpace(e)
|
|
if e == "" {
|
|
continue
|
|
}
|
|
|
|
v.SetConfigName("config." + e)
|
|
|
|
if err := v.MergeInConfig(); err != nil {
|
|
fmt.Printf("⚠️ config.%s.yaml not found\n", e)
|
|
}
|
|
}
|
|
}
|
|
|
|
func bindEnv(v *viper.Viper) {
|
|
v.AutomaticEnv()
|
|
v.SetEnvKeyReplacer(strings.NewReplacer(".", "_"))
|
|
}
|
|
|
|
func setDefaults(v *viper.Viper, env string) {
|
|
v.SetDefault("app.env", env)
|
|
|
|
v.SetDefault("app.name", "golang")
|
|
|
|
v.SetDefault("server.port", "8080")
|
|
|
|
v.SetDefault("unique_id.datacenter_id", 1)
|
|
v.SetDefault("unique_id.worker_id", 1)
|
|
|
|
v.SetDefault("jwt.secret", "abcdedhaldkasdlkasd")
|
|
v.SetDefault("jwt.access_expiry", -1)
|
|
}
|
|
|
|
func getEnv(key, def string) string {
|
|
if val := os.Getenv(key); val != "" {
|
|
return val
|
|
}
|
|
return def
|
|
}
|