@@ -8,6 +8,8 @@ import (
|
||||
// Config 保存进程运行所需的全部配置。
|
||||
type Config struct {
|
||||
Addr string // HTTP 监听地址
|
||||
DBDriver string // 数据库驱动:sqlite(默认)/ mysql / postgres
|
||||
DBDSN string // mysql / postgres 的连接串;sqlite 不使用
|
||||
DBPath string // SQLite 文件路径
|
||||
DataKey string // 敏感字段加密主密钥
|
||||
JWTSecret string // JWT 签名密钥
|
||||
@@ -26,8 +28,14 @@ func Load() (*Config, error) {
|
||||
if jwtSecret == "" {
|
||||
return nil, fmt.Errorf("load config: JWT_SECRET is required")
|
||||
}
|
||||
driver, dsn, err := dbFromEnv()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &Config{
|
||||
Addr: envOr("ADDR", ":8080"),
|
||||
DBDriver: driver,
|
||||
DBDSN: dsn,
|
||||
DBPath: envOr("DB_PATH", "oci-portal.db"),
|
||||
DataKey: dataKey,
|
||||
JWTSecret: jwtSecret,
|
||||
@@ -37,6 +45,22 @@ func Load() (*Config, error) {
|
||||
}, nil
|
||||
}
|
||||
|
||||
// dbFromEnv 读取并校验数据库驱动配置;外部库缺 DSN 视为配置错误,启动即失败。
|
||||
func dbFromEnv() (driver, dsn string, err error) {
|
||||
driver = envOr("DB_DRIVER", "sqlite")
|
||||
dsn = os.Getenv("DB_DSN")
|
||||
switch driver {
|
||||
case "sqlite":
|
||||
case "mysql", "postgres":
|
||||
if dsn == "" {
|
||||
return "", "", fmt.Errorf("load config: DB_DSN is required when DB_DRIVER=%s", driver)
|
||||
}
|
||||
default:
|
||||
return "", "", fmt.Errorf("load config: unsupported DB_DRIVER %q (sqlite/mysql/postgres)", driver)
|
||||
}
|
||||
return driver, dsn, nil
|
||||
}
|
||||
|
||||
func envOr(name, def string) string {
|
||||
if v := os.Getenv(name); v != "" {
|
||||
return v
|
||||
|
||||
@@ -0,0 +1,41 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestDBFromEnv(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
driver string
|
||||
dsn string
|
||||
wantDriver string
|
||||
wantErr string
|
||||
}{
|
||||
{name: "缺省走 sqlite", wantDriver: "sqlite"},
|
||||
{name: "postgres 带 DSN", driver: "postgres", dsn: "host=x", wantDriver: "postgres"},
|
||||
{name: "mysql 缺 DSN 报错", driver: "mysql", wantErr: "DB_DSN is required"},
|
||||
{name: "postgres 缺 DSN 报错", driver: "postgres", wantErr: "DB_DSN is required"},
|
||||
{name: "非法驱动报错", driver: "oracle", wantErr: "unsupported DB_DRIVER"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Setenv("DB_DRIVER", tt.driver)
|
||||
t.Setenv("DB_DSN", tt.dsn)
|
||||
driver, dsn, err := dbFromEnv()
|
||||
if tt.wantErr != "" {
|
||||
if err == nil || !strings.Contains(err.Error(), tt.wantErr) {
|
||||
t.Fatalf("err = %v, want contains %q", err, tt.wantErr)
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("dbFromEnv: %v", err)
|
||||
}
|
||||
if driver != tt.wantDriver || dsn != tt.dsn {
|
||||
t.Fatalf("got (%q,%q), want (%q,%q)", driver, dsn, tt.wantDriver, tt.dsn)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user