feat: 支持 MySQL 公网 TLS CA 校验 (#86)

This commit is contained in:
chengma
2026-08-10 09:59:09 +08:00
parent 830ca58030
commit f0715ee273
10 changed files with 354 additions and 22 deletions
+26 -2
View File
@@ -29,6 +29,13 @@ const (
databaseNameEnv = "CMAUTOBUY_DB_NAME"
databaseUserEnv = "CMAUTOBUY_DB_USER"
databasePasswordEnv = "CMAUTOBUY_DB_PASSWORD"
databaseTLSModeEnv = "CMAUTOBUY_DB_TLS_MODE"
databaseTLSCAEnv = "CMAUTOBUY_DB_TLS_CA"
)
const (
DatabaseTLSDisabled = "disabled"
DatabaseTLSVerifyCA = "verify_ca"
)
// DatabaseConfig 是 MySQL 8 的连接配置。
@@ -41,6 +48,8 @@ type DatabaseConfig struct {
Name string `yaml:"name"`
User string `yaml:"user"`
Password string `yaml:"password"`
TLSMode string `yaml:"tls_mode"`
TLSCA string `yaml:"tls_ca"`
}
// String 永远隐藏密码,防止排错时用 %v 把凭据写进日志。
@@ -49,8 +58,8 @@ func (c DatabaseConfig) String() string {
if c.Password != "" {
password = "****"
}
return fmt.Sprintf("DatabaseConfig{Host:%s Port:%s Name:%s User:%s Password:%s}",
c.Host, c.Port, c.Name, c.User, password)
return fmt.Sprintf("DatabaseConfig{Host:%s Port:%s Name:%s User:%s Password:%s TLSMode:%s TLSCA:%s}",
c.Host, c.Port, c.Name, c.User, password, c.TLSMode, c.TLSCA)
}
// LoadDatabase 读取 MySQL 配置。config.yaml 适合本地双击 run_admin.bat;
@@ -98,6 +107,8 @@ func mergeDatabaseConfig(cfg DatabaseConfig) (DatabaseConfig, error) {
databaseNameEnv: &cfg.Name,
databaseUserEnv: &cfg.User,
databasePasswordEnv: &cfg.Password,
databaseTLSModeEnv: &cfg.TLSMode,
databaseTLSCAEnv: &cfg.TLSCA,
}
for envName, target := range overrides {
value := os.Getenv(envName)
@@ -113,12 +124,17 @@ func mergeDatabaseConfig(cfg DatabaseConfig) (DatabaseConfig, error) {
cfg.Port = strings.TrimSpace(cfg.Port)
cfg.Name = strings.TrimSpace(cfg.Name)
cfg.User = strings.TrimSpace(cfg.User)
cfg.TLSMode = strings.ToLower(strings.TrimSpace(cfg.TLSMode))
cfg.TLSCA = strings.TrimSpace(cfg.TLSCA)
if cfg.Host == "" {
cfg.Host = "127.0.0.1"
}
if cfg.Port == "" {
cfg.Port = "3307"
}
if cfg.TLSMode == "" {
cfg.TLSMode = DatabaseTLSDisabled
}
var missing []string
if cfg.Name == "" {
missing = append(missing, "database.name / "+databaseNameEnv)
@@ -129,9 +145,17 @@ func mergeDatabaseConfig(cfg DatabaseConfig) (DatabaseConfig, error) {
if cfg.Password == "" {
missing = append(missing, "database.password / "+databasePasswordEnv)
}
if cfg.TLSMode == DatabaseTLSVerifyCA && cfg.TLSCA == "" {
missing = append(missing, "database.tls_ca / "+databaseTLSCAEnv)
}
if len(missing) > 0 {
return DatabaseConfig{}, fmt.Errorf("缺少 MySQL 配置:%s", strings.Join(missing, "、"))
}
if cfg.TLSMode != DatabaseTLSDisabled && cfg.TLSMode != DatabaseTLSVerifyCA {
return DatabaseConfig{}, fmt.Errorf(
"不支持的 MySQL TLS 模式 %q:database.tls_mode / %s 只能是 %s 或 %s",
cfg.TLSMode, databaseTLSModeEnv, DatabaseTLSDisabled, DatabaseTLSVerifyCA)
}
return cfg, nil
}