use crate::error::{Result, TradingError};
use serde::{Deserialize, Serialize};
use std::path::Path;
use validator::Validate;
#[derive(Debug, Clone, Serialize, Deserialize, Validate)]
pub struct AppConfig {
#[validate]
pub server: ServerConfig,
#[validate]
pub broker: BrokerConfig,
#[validate]
pub strategies: Vec<StrategyConfig>,
#[validate]
pub risk: RiskConfig,
#[validate]
pub database: DatabaseConfig,
#[validate]
pub logging: LoggingConfig,
}
impl AppConfig {
pub fn from_toml_file(path: impl AsRef<Path>) -> Result<Self> {
let contents = std::fs::read_to_string(path.as_ref())
.map_err(|e| TradingError::config(format!("Failed to read config file: {}", e)))?;
let config: Self = toml::from_str(&contents)
.map_err(|e| TradingError::config(format!("Failed to parse TOML config: {}", e)))?;
config
.validate()
.map_err(|e| TradingError::config(format!("Configuration validation failed: {}", e)))?;
Ok(config)
}
pub fn from_json_file(path: impl AsRef<Path>) -> Result<Self> {
let contents = std::fs::read_to_string(path.as_ref())
.map_err(|e| TradingError::config(format!("Failed to read config file: {}", e)))?;
let config: Self = serde_json::from_str(&contents)
.map_err(|e| TradingError::config(format!("Failed to parse JSON config: {}", e)))?;
config
.validate()
.map_err(|e| TradingError::config(format!("Configuration validation failed: {}", e)))?;
Ok(config)
}
pub fn default_test_config() -> Self {
Self {
server: ServerConfig::default(),
broker: BrokerConfig::default(),
strategies: vec![],
risk: RiskConfig::default(),
database: DatabaseConfig::default(),
logging: LoggingConfig::default(),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize, Validate)]
pub struct ServerConfig {
#[validate(length(min = 1))]
pub host: String,
#[validate(range(min = 1024, max = 65535))]
pub port: u16,
pub enable_https: bool,
#[validate(range(min = 1024, max = 104857600))] pub max_request_size: usize,
#[validate(range(min = 1, max = 300))]
pub request_timeout_secs: u64,
}
impl Default for ServerConfig {
fn default() -> Self {
Self {
host: "127.0.0.1".to_string(),
port: 8080,
enable_https: false,
max_request_size: 10485760, request_timeout_secs: 30,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize, Validate)]
pub struct BrokerConfig {
#[validate(length(min = 1))]
pub name: String,
#[validate(url)]
pub api_url: String,
#[validate(url)]
pub ws_url: String,
#[serde(skip_serializing)]
pub api_key: String,
#[serde(skip_serializing)]
pub api_secret: String,
pub paper_trading: bool,
#[validate(range(min = 1, max = 60))]
pub connection_timeout_secs: u64,
#[validate(range(min = 0, max = 10))]
pub max_retry_attempts: u32,
}
impl Default for BrokerConfig {
fn default() -> Self {
Self {
name: "alpaca".to_string(),
api_url: "https://paper-api.alpaca.markets".to_string(),
ws_url: "wss://stream.data.alpaca.markets".to_string(),
api_key: String::new(),
api_secret: String::new(),
paper_trading: true,
connection_timeout_secs: 30,
max_retry_attempts: 3,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize, Validate)]
pub struct StrategyConfig {
#[validate(length(min = 1))]
pub id: String,
#[validate(length(min = 1))]
pub strategy_type: String,
#[validate(length(min = 1))]
pub symbols: Vec<String>,
pub enabled: bool,
pub parameters: serde_json::Value,
}
#[derive(Debug, Clone, Serialize, Deserialize, Validate)]
pub struct RiskConfig {
#[validate(range(min = 0.0, max = 1.0))]
pub max_position_size: f64,
#[validate(range(min = 0.0, max = 1.0))]
pub max_daily_loss: f64,
#[validate(range(min = 0.0, max = 1.0))]
pub max_drawdown: f64,
#[validate(range(min = 1.0, max = 10.0))]
pub max_leverage: f64,
#[validate(range(min = 0.0, max = 1.0))]
pub default_stop_loss: f64,
#[validate(range(min = 0.0, max = 1.0))]
pub default_take_profit: f64,
#[validate(range(min = 0.0, max = 1.0))]
pub max_sector_concentration: f64,
pub enable_circuit_breakers: bool,
#[validate(range(min = 60, max = 86400))] pub circuit_breaker_cooldown_secs: u64,
}
impl Default for RiskConfig {
fn default() -> Self {
Self {
max_position_size: 0.1, max_daily_loss: 0.05, max_drawdown: 0.2, max_leverage: 1.0, default_stop_loss: 0.02, default_take_profit: 0.05, max_sector_concentration: 0.3, enable_circuit_breakers: true,
circuit_breaker_cooldown_secs: 300, }
}
}
#[derive(Debug, Clone, Serialize, Deserialize, Validate)]
pub struct DatabaseConfig {
#[validate(length(min = 1))]
pub database_type: String,
#[validate(length(min = 1))]
pub connection_url: String,
#[validate(range(min = 1, max = 100))]
pub max_connections: u32,
#[validate(range(min = 1, max = 60))]
pub connection_timeout_secs: u64,
}
impl Default for DatabaseConfig {
fn default() -> Self {
Self {
database_type: "sqlite".to_string(),
connection_url: "sqlite::memory:".to_string(),
max_connections: 10,
connection_timeout_secs: 30,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize, Validate)]
pub struct LoggingConfig {
#[validate(length(min = 1))]
pub level: String,
#[validate(length(min = 1))]
pub format: String,
pub enable_file_logging: bool,
pub log_file_path: Option<String>,
#[validate(range(min = 1048576, max = 1073741824))] pub max_log_file_size: usize,
#[validate(range(min = 1, max = 100))]
pub log_file_count: usize,
}
impl Default for LoggingConfig {
fn default() -> Self {
Self {
level: "info".to_string(),
format: "pretty".to_string(),
enable_file_logging: false,
log_file_path: None,
max_log_file_size: 10485760, log_file_count: 5,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::io::Write;
use tempfile::NamedTempFile;
#[test]
fn test_default_config() {
let config = AppConfig::default_test_config();
assert!(config.validate().is_ok());
}
#[test]
fn test_server_config_validation() {
let mut config = ServerConfig::default();
assert!(config.validate().is_ok());
config.port = 80; assert!(config.validate().is_err());
config.port = 8080;
assert!(config.validate().is_ok());
}
#[test]
fn test_risk_config_validation() {
let mut config = RiskConfig::default();
assert!(config.validate().is_ok());
config.max_position_size = 1.5; assert!(config.validate().is_err());
config.max_position_size = 0.2;
assert!(config.validate().is_ok());
}
#[test]
fn test_load_from_toml() {
let toml_config = r#"
[server]
host = "0.0.0.0"
port = 8080
enable_https = false
max_request_size = 10485760
request_timeout_secs = 30
[broker]
name = "alpaca"
api_url = "https://paper-api.alpaca.markets"
ws_url = "wss://stream.data.alpaca.markets"
api_key = "test_key"
api_secret = "test_secret"
paper_trading = true
connection_timeout_secs = 30
max_retry_attempts = 3
[[strategies]]
id = "momentum_1"
strategy_type = "momentum"
symbols = ["AAPL", "GOOGL"]
enabled = true
parameters = {}
[risk]
max_position_size = 0.1
max_daily_loss = 0.05
max_drawdown = 0.2
max_leverage = 1.0
default_stop_loss = 0.02
default_take_profit = 0.05
max_sector_concentration = 0.3
enable_circuit_breakers = true
circuit_breaker_cooldown_secs = 300
[database]
database_type = "sqlite"
connection_url = "sqlite::memory:"
max_connections = 10
connection_timeout_secs = 30
[logging]
level = "info"
format = "pretty"
enable_file_logging = false
max_log_file_size = 10485760
log_file_count = 5
"#;
let mut temp_file = NamedTempFile::new().unwrap();
temp_file.write_all(toml_config.as_bytes()).unwrap();
temp_file.flush().unwrap();
let config = AppConfig::from_toml_file(temp_file.path()).unwrap();
assert_eq!(config.server.port, 8080);
assert_eq!(config.broker.name, "alpaca");
assert_eq!(config.strategies.len(), 1);
assert_eq!(config.risk.max_position_size, 0.1);
}
#[test]
fn test_broker_config_default() {
let config = BrokerConfig::default();
assert_eq!(config.name, "alpaca");
assert!(config.paper_trading);
assert_eq!(config.max_retry_attempts, 3);
}
#[test]
fn test_logging_config_validation() {
let mut config = LoggingConfig::default();
assert!(config.validate().is_ok());
config.max_log_file_size = 100; assert!(config.validate().is_err());
config.max_log_file_size = 10485760; assert!(config.validate().is_ok());
}
}