package config import ( "fmt" "log" "os" "time" "github.com/spf13/viper" ) // Config 应用配置结构体 type Config struct { Server ServerConfig `mapstructure:"server"` Database DatabaseConfig `mapstructure:"database"` JWT JWTConfig `mapstructure:"jwt"` AI AIConfig `mapstructure:"ai"` Quota QuotaConfig `mapstructure:"quota"` CORS CORSConfig `mapstructure:"cors"` } // ServerConfig 服务器配置 type ServerConfig struct { Port string `mapstructure:"port"` Host string `mapstructure:"host"` ReadTimeout time.Duration `mapstructure:"read_timeout"` WriteTimeout time.Duration `mapstructure:"write_timeout"` } // DatabaseConfig 数据库配置 type DatabaseConfig struct { DSN string `mapstructure:"dsn"` MaxIdleConns int `mapstructure:"max_idle_conns"` MaxOpenConns int `mapstructure:"max_open_conns"` } // JWTConfig JWT配置 type JWTConfig struct { SecretKey string `mapstructure:"secret_key"` AccessTokenExpire time.Duration `mapstructure:"access_token_expire"` RefreshTokenExpire time.Duration `mapstructure:"refresh_token_expire"` } // AIConfig AI模型配置 type AIConfig struct { APIURL string `mapstructure:"api_url"` APIKey string `mapstructure:"api_key"` ModelName string `mapstructure:"model_name"` MaxTokens int64 `mapstructure:"max_tokens"` Temperature float64 `mapstructure:"temperature"` } // QuotaConfig 配额配置 type QuotaConfig struct { DefaultTotalTokens int64 `mapstructure:"default_total_tokens"` DefaultTokenLimit int64 `mapstructure:"default_token_limit"` ResetInterval time.Duration `mapstructure:"reset_interval"` } // CORSConfig CORS配置 type CORSConfig struct { AllowOrigins []string `mapstructure:"allow_origins"` AllowMethods []string `mapstructure:"allow_methods"` AllowHeaders []string `mapstructure:"allow_headers"` AllowCredentials bool `mapstructure:"allow_credentials"` } // LoadConfig 加载配置 func LoadConfig(path string) (*Config, error) { // 设置默认值 setDefaults() // 读取配置文件 viper.AddConfigPath(path) viper.SetConfigName(".env") viper.SetConfigType("env") // 读取环境变量 viper.AutomaticEnv() // 尝试读取配置文件 if err := viper.ReadInConfig(); err != nil { if _, ok := err.(viper.ConfigFileNotFoundError); ok { log.Println("警告: 未找到配置文件,将使用环境变量和默认值") } else { return nil, fmt.Errorf("读取配置文件错误: %w", err) } } // 解析配置 var config Config if err := viper.Unmarshal(&config); err != nil { return nil, fmt.Errorf("解析配置错误: %w", err) } // 验证必要的配置 if err := validateConfig(&config); err != nil { return nil, err } return &config, nil } // 设置默认值 func setDefaults() { // 服务器默认配置 viper.SetDefault("server.port", "8000") viper.SetDefault("server.host", "0.0.0.0") viper.SetDefault("server.read_timeout", 15*time.Second) viper.SetDefault("server.write_timeout", 15*time.Second) // 数据库默认配置(使用SQLite) viper.SetDefault("database.dsn", "./chat.db") viper.SetDefault("database.max_idle_conns", 10) viper.SetDefault("database.max_open_conns", 100) // JWT默认配置 viper.SetDefault("jwt.secret_key", "your-secret-key-change-me") viper.SetDefault("jwt.access_token_expire", 24*time.Hour) viper.SetDefault("jwt.refresh_token_expire", 7*24*time.Hour) // AI模型默认配置 viper.SetDefault("ai.api_url", "http://localhost:5000/v1/chat/completions") viper.SetDefault("ai.api_key", "") viper.SetDefault("ai.model_name", "default") viper.SetDefault("ai.max_tokens", 2048) viper.SetDefault("ai.temperature", 0.7) // 配额默认配置 viper.SetDefault("quota.default_total_tokens", 100000) viper.SetDefault("quota.default_token_limit", 1000) viper.SetDefault("quota.reset_interval", 30*24*time.Hour) // CORS默认配置 viper.SetDefault("cors.allow_origins", []string{"*"}) viper.SetDefault("cors.allow_methods", []string{"GET", "POST", "PUT", "DELETE", "OPTIONS"}) viper.SetDefault("cors.allow_headers", []string{"Origin", "Content-Type", "Authorization"}) viper.SetDefault("cors.allow_credentials", true) } // 验证配置 func validateConfig(config *Config) error { // 验证JWT密钥 if config.JWT.SecretKey == "your-secret-key-change-me" { log.Println("警告: 使用默认JWT密钥,建议在生产环境中修改") } // 验证AI API URL if config.AI.APIURL == "" { return fmt.Errorf("AI API URL不能为空") } return nil } // GetDSN 获取数据库DSN func (c *DatabaseConfig) GetDSN() string { // 如果是SQLite,确保目录存在 if len(c.DSN) > 6 && c.DSN[:6] == "sqlite" { // 提取SQLite文件路径 filePath := c.DSN[9:] // 去掉 "sqlite://" 前缀 dir := filePath[:len(filePath)-len("/"+filePath[len(filePath)-1:])] if dir != "" { os.MkdirAll(dir, 0755) } } return c.DSN }