//! Trading agents implementing various strategies. use algoarena_shared::{ AgentConfig, AgentState, Execution, MarketState, Order, OrderSide, OrderType, StrategyType, }; /// Trading agent trait. pub trait TradingAgent: std::fmt::Debug { /// Get agent ID. fn id(&self) -> &str; /// Get agent name. fn name(&self) -> &str; /// Get strategy type. fn strategy_type(&self) -> StrategyType; /// Observe market and update internal state. fn observe(&mut self, state: &MarketState); /// Decide on orders to place. fn decide(&self, state: &MarketState) -> Vec; /// Update after execution. fn update(&mut self, executions: &[Execution]); /// Get current agent state. fn get_state(&self) -> AgentState; /// Reset agent. fn reset(&mut self); } /// Base agent implementation. #[derive(Debug)] pub struct BaseAgent { /// Configuration. config: AgentConfig, /// Current cash. cash: f64, /// Holdings per asset. holdings: Vec, /// Portfolio value history. pnl_history: Vec, /// Number of trades. trade_count: usize, /// Last known prices. last_prices: Vec, /// Number of assets. num_assets: usize, } impl BaseAgent { /// Create a new base agent. pub fn new(config: AgentConfig, num_assets: usize) -> Self { Self { cash: config.initial_capital, config, holdings: vec![0.0; num_assets], pnl_history: vec![], trade_count: 0, last_prices: vec![0.0; num_assets], num_assets, } } /// Calculate portfolio value. pub fn portfolio_value(&self, prices: &[f64]) -> f64 { let holdings_value: f64 = self .holdings .iter() .zip(prices.iter()) .map(|(h, p)| h * p) .sum(); self.cash + holdings_value } } impl TradingAgent for BaseAgent { fn id(&self) -> &str { &self.config.id } fn name(&self) -> &str { &self.config.name } fn strategy_type(&self) -> StrategyType { self.config.strategy } fn observe(&mut self, state: &MarketState) { self.last_prices = state.prices.clone(); let value = self.portfolio_value(&state.prices); self.pnl_history.push(value); } fn decide(&self, state: &MarketState) -> Vec { match self.config.strategy { StrategyType::BuyAndHold => self.buy_and_hold_strategy(state), StrategyType::Momentum => self.momentum_strategy(state), StrategyType::MeanReversion => self.mean_reversion_strategy(state), StrategyType::TrendFollowing => self.trend_following_strategy(state), StrategyType::Random => self.random_strategy(state), StrategyType::Custom => vec![], } } fn update(&mut self, executions: &[Execution]) { for exec in executions { if exec.is_filled { self.trade_count += 1; match exec.order.side { OrderSide::Buy => { self.holdings[exec.order.asset_idx] += exec.filled_quantity; self.cash -= exec.filled_quantity * exec.fill_price + exec.transaction_cost; } OrderSide::Sell => { self.holdings[exec.order.asset_idx] -= exec.filled_quantity; self.cash += exec.filled_quantity * exec.fill_price - exec.transaction_cost; } } } } } fn get_state(&self) -> AgentState { AgentState { agent_id: self.config.id.clone(), cash: self.cash, holdings: self.holdings.clone(), portfolio_value: self.portfolio_value(&self.last_prices), pnl_history: self.pnl_history.clone(), trade_count: self.trade_count, } } fn reset(&mut self) { self.cash = self.config.initial_capital; self.holdings = vec![0.0; self.num_assets]; self.pnl_history.clear(); self.trade_count = 0; } } impl BaseAgent { /// Buy and hold strategy. fn buy_and_hold_strategy(&self, state: &MarketState) -> Vec { // Only buy on first step if state.step != 1 { return vec![]; } let mut orders = Vec::new(); let allocation_per_asset = self.cash / state.prices.len() as f64; for (i, &price) in state.prices.iter().enumerate() { let quantity = (allocation_per_asset / price * 0.95).floor(); // 95% to leave some cash if quantity > 0.0 { orders.push(Order { agent_id: self.config.id.clone(), asset_idx: i, side: OrderSide::Buy, order_type: OrderType::Market, quantity, limit_price: None, }); } } orders } /// Momentum strategy. fn momentum_strategy(&self, state: &MarketState) -> Vec { let lookback = self.config.parameters.lookback; if state.step < lookback + 1 { return vec![]; } let mut orders = Vec::new(); let position_size = self.config.parameters.position_size; for (i, history) in state.price_history.iter().enumerate() { let current = history[state.step]; let past = history[state.step - lookback]; let momentum = (current - past) / past; let threshold = self.config.parameters.momentum_threshold; // Buy on positive momentum if momentum > threshold && self.holdings[i] == 0.0 { let allocation = self.cash * position_size; let quantity = (allocation / current).floor(); if quantity > 0.0 { orders.push(Order { agent_id: self.config.id.clone(), asset_idx: i, side: OrderSide::Buy, order_type: OrderType::Market, quantity, limit_price: None, }); } } // Sell on negative momentum else if momentum < -threshold && self.holdings[i] > 0.0 { orders.push(Order { agent_id: self.config.id.clone(), asset_idx: i, side: OrderSide::Sell, order_type: OrderType::Market, quantity: self.holdings[i], limit_price: None, }); } } orders } /// Mean reversion strategy. fn mean_reversion_strategy(&self, state: &MarketState) -> Vec { let lookback = self.config.parameters.lookback; if state.step < lookback + 1 { return vec![]; } let mut orders = Vec::new(); let position_size = self.config.parameters.position_size; let threshold = self.config.parameters.reversion_threshold; for (i, history) in state.price_history.iter().enumerate() { // Calculate mean and std dev let window = &history[(state.step - lookback)..=state.step]; let mean: f64 = window.iter().sum::() / window.len() as f64; let variance: f64 = window.iter().map(|p| (p - mean).powi(2)).sum::() / window.len() as f64; let std_dev = variance.sqrt(); if std_dev < 0.001 { continue; } let current = history[state.step]; let z_score = (current - mean) / std_dev; // Buy when price is low (negative z-score) if z_score < -threshold && self.holdings[i] == 0.0 { let allocation = self.cash * position_size; let quantity = (allocation / current).floor(); if quantity > 0.0 { orders.push(Order { agent_id: self.config.id.clone(), asset_idx: i, side: OrderSide::Buy, order_type: OrderType::Market, quantity, limit_price: None, }); } } // Sell when price is high (positive z-score) else if z_score > threshold && self.holdings[i] > 0.0 { orders.push(Order { agent_id: self.config.id.clone(), asset_idx: i, side: OrderSide::Sell, order_type: OrderType::Market, quantity: self.holdings[i], limit_price: None, }); } } orders } /// Trend following strategy (MA crossover). fn trend_following_strategy(&self, state: &MarketState) -> Vec { let short_period = self.config.parameters.short_ma_period; let long_period = self.config.parameters.long_ma_period; if state.step < long_period + 1 { return vec![]; } let mut orders = Vec::new(); let position_size = self.config.parameters.position_size; for (i, history) in state.price_history.iter().enumerate() { // Calculate MAs let short_ma: f64 = history[(state.step - short_period + 1)..=state.step] .iter() .sum::() / short_period as f64; let long_ma: f64 = history[(state.step - long_period + 1)..=state.step] .iter() .sum::() / long_period as f64; let prev_short_ma: f64 = history[(state.step - short_period)..state.step] .iter() .sum::() / short_period as f64; let prev_long_ma: f64 = history[(state.step - long_period)..state.step] .iter() .sum::() / long_period as f64; // Crossover detection let golden_cross = prev_short_ma <= prev_long_ma && short_ma > long_ma; let death_cross = prev_short_ma >= prev_long_ma && short_ma < long_ma; if golden_cross && self.holdings[i] == 0.0 { let allocation = self.cash * position_size; let quantity = (allocation / state.prices[i]).floor(); if quantity > 0.0 { orders.push(Order { agent_id: self.config.id.clone(), asset_idx: i, side: OrderSide::Buy, order_type: OrderType::Market, quantity, limit_price: None, }); } } else if death_cross && self.holdings[i] > 0.0 { orders.push(Order { agent_id: self.config.id.clone(), asset_idx: i, side: OrderSide::Sell, order_type: OrderType::Market, quantity: self.holdings[i], limit_price: None, }); } } orders } /// Random strategy (baseline). fn random_strategy(&self, state: &MarketState) -> Vec { // Use state step as pseudo-random source let rand = ((state.step.wrapping_mul(1103515245).wrapping_add(12345)) >> 16) % 100; if rand < 5 { // 5% chance to trade let asset_idx = state.step % self.num_assets; if self.holdings[asset_idx] > 0.0 { return vec![Order { agent_id: self.config.id.clone(), asset_idx, side: OrderSide::Sell, order_type: OrderType::Market, quantity: self.holdings[asset_idx], limit_price: None, }]; } else { let quantity = (self.cash * 0.1 / state.prices[asset_idx]).floor(); if quantity > 0.0 { return vec![Order { agent_id: self.config.id.clone(), asset_idx, side: OrderSide::Buy, order_type: OrderType::Market, quantity, limit_price: None, }]; } } } vec![] } } /// Create an agent from configuration. pub fn create_agent(config: AgentConfig, num_assets: usize) -> Box { Box::new(BaseAgent::new(config, num_assets)) } #[cfg(test)] mod tests { use super::*; use algoarena_shared::{MarketIndicators, default_agents}; fn create_test_state() -> MarketState { let mut history = vec![vec![100.0]; 3]; // Add 50 steps of history for i in 0..50 { for h in &mut history { let change = 1.0 + 0.001 * (i as f64 - 25.0); h.push(h.last().unwrap() * change); } } MarketState { step: 50, prices: history.iter().map(|h| *h.last().unwrap()).collect(), price_history: history, volumes: vec![1_000_000.0; 3], indicators: MarketIndicators::default(), } } #[test] fn test_base_agent_creation() { let configs = default_agents(); let agent = BaseAgent::new(configs[0].clone(), 3); assert_eq!(agent.id(), configs[0].id); assert_eq!(agent.cash, configs[0].initial_capital); } #[test] fn test_portfolio_value() { let configs = default_agents(); let mut agent = BaseAgent::new(configs[0].clone(), 3); agent.holdings = vec![10.0, 5.0, 20.0]; agent.cash = 1000.0; let prices = vec![100.0, 200.0, 50.0]; let value = agent.portfolio_value(&prices); // 10*100 + 5*200 + 20*50 + 1000 = 1000 + 1000 + 1000 + 1000 = 4000 assert!((value - 4000.0).abs() < 0.01); } #[test] fn test_observe_updates_pnl() { let configs = default_agents(); let mut agent = BaseAgent::new(configs[0].clone(), 3); let state = create_test_state(); agent.observe(&state); assert_eq!(agent.pnl_history.len(), 1); } #[test] fn test_buy_and_hold() { let configs = default_agents(); let buyhold = configs .iter() .find(|c| c.strategy == StrategyType::BuyAndHold) .unwrap(); let agent = BaseAgent::new(buyhold.clone(), 3); let mut state = create_test_state(); state.step = 1; // First step let orders = agent.decide(&state); assert!(!orders.is_empty()); // Should place orders on first step } #[test] fn test_agent_reset() { let configs = default_agents(); let mut agent = BaseAgent::new(configs[0].clone(), 3); agent.cash = 50000.0; agent.holdings = vec![10.0, 20.0, 30.0]; agent.trade_count = 5; agent.reset(); assert_eq!(agent.cash, configs[0].initial_capital); assert_eq!(agent.holdings, vec![0.0, 0.0, 0.0]); assert_eq!(agent.trade_count, 0); } }