diff --git a/QWEN.md b/QWEN.md index 9e4fccb2..e3679ced 100644 --- a/QWEN.md +++ b/QWEN.md @@ -839,3 +839,180 @@ git diff --staged git status # Should NOT show prompts/local/, models/local/, .env, results/ ``` + +--- + +## Development Guidelines for AI Assistant + +### ⚠️ MANDATORY Rules for ALL Development + +**When implementing NEW features or making SIGNIFICANT changes, you MUST:** + +#### 1. 📝 Update QWEN.md + +**When:** Every time you add a new feature, module, or change existing architecture. + +**What to update:** +- Architecture section (if structure changes) +- Important Files section +- Testing section +- Key Metrics (if targets change) +- Project Status +- Next Steps + +**Example:** +```markdown +### Architecture +├── rdagent/ +│ └── components/ +│ └── backtesting/ +│ └── protections/ # NEW: Trading protection system +│ ├── base.py +│ ├── max_drawdown.py +│ └── protection_manager.py +``` + +#### 2. 📖 Update README.md + +**When:** Every user-facing feature change or major update. + +**What to update:** +- Features list +- Installation instructions +- Usage examples +- Configuration examples + +**Keep it user-focused:** +```markdown +## Features +- ✅ Trading Protection System (NEW) + * Automatic drawdown protection + * Cooldown periods after losses + * Stoploss cluster detection +``` + +#### 3. 📦 Update requirements.txt + +**When:** Adding new dependencies or removing unused ones. + +**What to update:** +- `requirements.txt` (main dependencies) +- `requirements/lint.txt` (dev dependencies) +- `requirements/test.txt` (test dependencies) + +**Example:** +```bash +# If you add a new library +echo "new-library==1.0.0" >> requirements.txt + +# If you add a new test dependency +echo "pytest-mock" >> requirements/test.txt +``` + +#### 4. ✅ Extend Tests + +**When:** EVERY time you add new code. + +**Rule:** New features MUST have tests with >80% coverage. + +**What to create:** +- Unit tests in `test/` directory +- Integration tests in `test/integration/` +- Update existing tests if behavior changed + +**Test structure:** +```python +# test/feature_type/test_new_feature.py +"""Tests for New Feature""" + +class TestNewFeature: + """Test new feature thoroughly.""" + + def test_basic_functionality(self): ... + def test_edge_cases(self): ... + def test_error_handling(self): ... + def test_integration_with_existing(self): ... +``` + +**Update integration tests:** +```python +# Add to test/integration/test_all_features.py +class TestNewFeature: + """Test new feature integration.""" + + def test_imports(self): ... + def test_initialization(self): ... + def test_full_workflow(self): ... +``` + +#### 5. 🔄 Pre-Commit Checklist + +**BEFORE every commit with new features:** + +```bash +# 1. Run ALL tests +pytest test/ -v + +# 2. Run integration tests +pytest test/integration/test_all_features.py -v + +# 3. Check test coverage +pytest --cov=rdagent.components.new_module -v + +# 4. Run security scan +bandit -r rdagent/ -c .bandit.yml + +# 5. Verify tests updated +git status +# Should show test files modified +``` + +### Documentation Priority Order + +1. **QWEN.md** - Internal AI assistant context (UPDATE ALWAYS) +2. **Test files** - Code documentation through tests (MANDATORY) +3. **README.md** - User-facing documentation (UPDATE for user-visible changes) +4. **requirements.txt** - Dependencies (UPDATE when adding libraries) +5. **Inline code comments** - English only (ALWAYS) + +### Example Workflow: Adding New Feature + +``` +1. Plan feature + ↓ +2. Implement code + ↓ +3. Write unit tests (test/...) + ↓ +4. Write integration tests (test/integration/...) + ↓ +5. Run ALL tests → Must pass + ↓ +6. Update QWEN.md ← MANDATORY + ↓ +7. Update README.md (if user-visible) + ↓ +8. Update requirements.txt (if new deps) + ↓ +9. Commit with clear message + ↓ +10. Pre-commit hooks run automatically + ↓ +11. Push to remote +``` + +### Penalties for Not Following Rules + +**If you forget to update:** +- ❌ Missing tests → Code cannot be committed (pre-commit blocks) +- ❌ Missing QWEN.md update → Next AI assistant will work with outdated context +- ❌ Missing README update → Users won't understand new features +- ❌ Missing requirements.txt → Installation will fail + +**Remember:** These rules ensure: +1. Code quality through tests +2. AI assistant has current context +3. Users understand changes +4. Dependencies are tracked + +--- diff --git a/README.md b/README.md index 62833517..d872fc55 100644 --- a/README.md +++ b/README.md @@ -196,6 +196,15 @@ Predix continuously proposes, implements, and validates new alpha factors: - Avoids overfitting through walk-forward validation - Discovers non-obvious patterns in order flow, volatility, and session dynamics +### 🛡️ Trading Protection System + +Automatic risk management to prevent excessive losses: + +- **Max Drawdown Protection** - Pauses trading when drawdown exceeds threshold (default: 15%) +- **Cooldown Period** - Enforces mandatory rest period after significant losses (default: 4h after 5% loss) +- **Stoploss Guard** - Detects clusters of stoplosses and blocks trading (default: max 5 per day) +- **Low Performance Filter** - Filters out consistently underperforming factors (Sharpe < 0.5, Win Rate < 40%) + ### 🧠 Model Architecture Search Automatically explores and refines predictive models: @@ -221,6 +230,14 @@ Real-time dashboard for monitoring: - Cumulative returns and drawdowns - Code diffs and implementation history +### 🔒 Security & Quality + +Automated quality assurance: + +- **60 Integration Tests** - All features tested automatically +- **Bandit Security Scanner** - Pre-commit security checks +- **Pre-commit Hooks** - Tests run before EVERY commit + --- ## Project Structure @@ -230,9 +247,27 @@ predix/ ├── rdagent/ # Core agent framework │ ├── app/ # CLI and scenario apps │ ├── components/ # Reusable agent components +│ │ ├── backtesting/ # Backtest engine & protections +│ │ │ ├── backtest_engine.py +│ │ │ ├── results_db.py +│ │ │ ├── risk_management.py +│ │ │ └── protections/ # Trading protection system (NEW) +│ │ │ ├── base.py +│ │ │ ├── max_drawdown.py +│ │ │ ├── cooldown.py +│ │ │ ├── stoploss_guard.py +│ │ │ ├── low_performance.py +│ │ │ └── protection_manager.py +│ │ ├── coder/ # Factor & model coding +│ │ └── loader.py # Prompt & model loaders │ ├── core/ # Core abstractions │ ├── scenarios/ # Domain-specific scenarios │ └── utils/ # Utilities +├── test/ # Test suite +│ ├── integration/ # Integration tests (60 tests) +│ │ └── test_all_features.py +│ └── backtesting/ # Unit tests +│ └── test_protections.py ├── constraints/ # Constraint definitions ├── docs/ # Documentation ├── web/ # Web UI frontend diff --git a/rdagent/components/backtesting/protections/__init__.py b/rdagent/components/backtesting/protections/__init__.py new file mode 100644 index 00000000..6e1bf179 --- /dev/null +++ b/rdagent/components/backtesting/protections/__init__.py @@ -0,0 +1,58 @@ +""" +Trading Protection System for Predix. + +Prevents excessive losses by automatically pausing trading +when risk thresholds are exceeded. + +Usage: + from rdagent.components.backtesting.protections import ( + ProtectionManager, + MaxDrawdownProtection, + CooldownProtection, + StoplossGuardProtection, + LowPerformanceProtection, + ) + + manager = ProtectionManager() + manager.create_default_protections() + + result = manager.check_all( + returns=[0.01, -0.02, 0.015], + timestamps=[...], + current_equity=98000, + peak_equity=100000 + ) + + if result.should_block: + print(f"Trading blocked: {result.reason}") +""" + +from .base import ( + BaseProtection, + ProtectionConfig, + ProtectionResult, + ProtectionType, + ProtectionScope, +) +from .max_drawdown import MaxDrawdownProtection, MaxDrawdownConfig +from .cooldown import CooldownProtection, CooldownConfig +from .stoploss_guard import StoplossGuardProtection, StoplossGuardConfig +from .low_performance import LowPerformanceProtection, LowPerformanceConfig +from .protection_manager import ProtectionManager + +__all__ = [ + "BaseProtection", + "ProtectionConfig", + "ProtectionResult", + "ProtectionType", + "ProtectionScope", + "MaxDrawdownProtection", + "MaxDrawdownConfig", + "CooldownProtection", + "CooldownConfig", + "StoplossGuardProtection", + "StoplossGuardConfig", + "LowPerformanceProtection", + "LowPerformanceConfig", + "ProtectionManager", +] diff --git a/rdagent/components/backtesting/protections/base.py b/rdagent/components/backtesting/protections/base.py new file mode 100644 index 00000000..dccb2c9c --- /dev/null +++ b/rdagent/components/backtesting/protections/base.py @@ -0,0 +1,145 @@ +""" +Trading Protection System + +Prevents excessive losses by automatically pausing trading when risk thresholds are exceeded. + +Inspired by common trading protection patterns, implemented from scratch for Predix. +""" + +from abc import ABC, abstractmethod +from dataclasses import dataclass, field +from datetime import datetime, timedelta +from typing import Optional +from enum import Enum + + +class ProtectionType(Enum): + """Type of protection mechanism.""" + MAX_DRAWDOWN = "max_drawdown" + COOLDOWN = "cooldown" + STOPLOSS_GUARD = "stoploss_guard" + LOW_PERFORMANCE = "low_performance" + + +class ProtectionScope(Enum): + """What scope does this protection apply to?""" + GLOBAL = "global" # All trading + FACTOR = "factor" # Per-factor + PORTFOLIO = "portfolio" # Per portfolio + + +@dataclass +class ProtectionResult: + """Result of a protection check.""" + should_block: bool # True if trading should be blocked + reason: str # Why was it blocked? + until: Optional[datetime] = None # If time-based, when does it expire? + protection_type: Optional[ProtectionType] = None + severity: float = 0.0 # How severe is the issue (0-1) + + @property + def is_active(self) -> bool: + """Check if protection is currently active.""" + if not self.until: + return self.should_block + return datetime.now() < self.until and self.should_block + + +@dataclass +class ProtectionConfig: + """Base configuration for a protection.""" + enabled: bool = True + lookback_period_hours: int = 24 # How far back to look + severity_threshold: float = 0.8 # At what severity to block + + +class BaseProtection(ABC): + """ + Base class for all trading protections. + + Each protection checks specific conditions and returns ProtectionResult. + Multiple protections can be combined in ProtectionManager. + """ + + def __init__(self, config: ProtectionConfig): + self.config = config + self.last_check: Optional[datetime] = None + self.total_checks: int = 0 + self.total_blocks: int = 0 + + @abstractmethod + def check( + self, + returns: list[float], + timestamps: list[datetime], + current_equity: float, + peak_equity: float, + **kwargs + ) -> ProtectionResult: + """ + Check if protection should be triggered. + + Parameters + ---------- + returns : list[float] + Historical returns in lookback period + timestamps : list[datetime] + Timestamps of returns + current_equity : float + Current equity value + peak_equity : float + Peak equity value (highest ever) + **kwargs + Additional context (factor name, portfolio ID, etc.) + + Returns + ------- + ProtectionResult + Decision on whether to block trading + """ + pass + + def calculate_drawdown(self, current: float, peak: float) -> float: + """Calculate drawdown percentage (negative value).""" + if peak == 0: + return 0.0 + return (current - peak) / peak + + def calculate_max_consecutive_losses(self, returns: list[float]) -> int: + """Find maximum consecutive losing trades.""" + max_losses = 0 + current_losses = 0 + + for ret in returns: + if ret < 0: + current_losses += 1 + max_losses = max(max_losses, current_losses) + else: + current_losses = 0 + + return max_losses + + def calculate_recent_loss_rate(self, returns: list[float]) -> float: + """Calculate percentage of losing trades.""" + if not returns: + return 0.0 + losses = sum(1 for r in returns if r < 0) + return losses / len(returns) + + def record_check(self, blocked: bool = False): + """Record that a check was performed.""" + self.total_checks += 1 + self.last_check = datetime.now() + if blocked: + self.total_blocks += 1 + + def get_stats(self) -> dict: + """Get protection statistics.""" + return { + "type": self.__class__.__name__, + "enabled": self.config.enabled, + "total_checks": self.total_checks, + "total_blocks": self.total_blocks, + "block_rate": self.total_blocks / max(1, self.total_checks), + "last_check": self.last_check.isoformat() if self.last_check else None + } diff --git a/rdagent/components/backtesting/protections/cooldown.py b/rdagent/components/backtesting/protections/cooldown.py new file mode 100644 index 00000000..7ceca63d --- /dev/null +++ b/rdagent/components/backtesting/protections/cooldown.py @@ -0,0 +1,86 @@ +""" +Cooldown Period Protection + +Enforces mandatory rest periods after losses. +""" + +from dataclasses import dataclass +from datetime import datetime, timedelta +from typing import Optional +from .base import BaseProtection, ProtectionConfig, ProtectionResult, ProtectionType, ProtectionScope + + +@dataclass +class CooldownConfig(ProtectionConfig): + """Configuration for Cooldown protection.""" + cooldown_after_loss_pct: float = 0.05 # Cooldown after 5% loss + cooldown_duration_hours: int = 4 # How long to wait + + +class CooldownProtection(BaseProtection): + """ + Enforces cooling-off period after significant losses. + + Prevents revenge trading and emotional decisions by forcing + a mandatory break after losses exceed threshold. + """ + + def __init__(self, config: CooldownConfig): + super().__init__(config) + self.config: CooldownConfig = config + self.last_loss_time: Optional[datetime] = None + self.last_loss_pct: float = 0.0 + + @property + def scope(self) -> ProtectionScope: + return ProtectionScope.FACTOR + + def check( + self, + returns: list[float], + timestamps: list[datetime], + current_equity: float, + peak_equity: float, + **kwargs + ) -> ProtectionResult: + """Check if cooldown should be triggered.""" + self.record_check() + + if not self.config.enabled: + return ProtectionResult( + should_block=False, + reason="Protection disabled", + protection_type=ProtectionType.COOLDOWN + ) + + # Check most recent return + if returns: + latest_return = returns[-1] + latest_time = timestamps[-1] if timestamps else datetime.now() + + if latest_return < -self.config.cooldown_after_loss_pct: + self.last_loss_time = latest_time + self.last_loss_pct = latest_return + + # If recently had big loss, enforce cooldown + if self.last_loss_time: + time_since_loss = datetime.now() - self.last_loss_time + if time_since_loss < timedelta(hours=self.config.cooldown_duration_hours): + remaining = timedelta(hours=self.config.cooldown_duration_hours) - time_since_loss + + result = ProtectionResult( + should_block=True, + reason=f"Loss of {abs(self.last_loss_pct)*100:.1f}% - cooling down for {remaining.seconds // 3600}h", + until=self.last_loss_time + timedelta(hours=self.config.cooldown_duration_hours), + protection_type=ProtectionType.COOLDOWN, + severity=abs(self.last_loss_pct) / self.config.cooldown_after_loss_pct + ) + self.record_check(blocked=True) + return result + + return ProtectionResult( + should_block=False, + reason="No recent significant losses", + protection_type=ProtectionType.COOLDOWN, + severity=0.0 + ) diff --git a/rdagent/components/backtesting/protections/low_performance.py b/rdagent/components/backtesting/protections/low_performance.py new file mode 100644 index 00000000..bc35fe4c --- /dev/null +++ b/rdagent/components/backtesting/protections/low_performance.py @@ -0,0 +1,103 @@ +""" +Low Performance Filter + +Blocks trading for factors/portfolios with consistently poor performance. +""" + +from dataclasses import dataclass +from datetime import datetime +from .base import BaseProtection, ProtectionConfig, ProtectionResult, ProtectionType, ProtectionScope + + +@dataclass +class LowPerformanceConfig(ProtectionConfig): + """Configuration for LowPerformance protection.""" + min_sharpe_ratio: float = 0.5 # Minimum acceptable Sharpe + min_win_rate: float = 0.40 # Minimum 40% win rate + min_trades: int = 20 # Need at least this many trades to evaluate + + +class LowPerformanceProtection(BaseProtection): + """ + Filters out consistently underperforming factors. + + Prevents wasting resources on factors that statistical analysis + shows are unlikely to become profitable. + """ + + def __init__(self, config: LowPerformanceConfig): + super().__init__(config) + self.config: LowPerformanceConfig = config + + @property + def scope(self) -> ProtectionScope: + return ProtectionScope.FACTOR + + def check( + self, + returns: list[float], + timestamps: list[datetime], + current_equity: float, + peak_equity: float, + **kwargs + ) -> ProtectionResult: + """Check if performance is below minimum standards.""" + self.record_check() + + if not self.config.enabled: + return ProtectionResult( + should_block=False, + reason="Protection disabled", + protection_type=ProtectionType.LOW_PERFORMANCE + ) + + # Need minimum number of trades + if len(returns) < self.config.min_trades: + return ProtectionResult( + should_block=False, + reason=f"Insufficient data ({len(returns)} < {self.config.min_trades} trades)", + protection_type=ProtectionType.LOW_PERFORMANCE, + severity=0.0 + ) + + # Calculate metrics + import numpy as np + returns_array = np.array(returns) + + # Win rate + wins = int(np.sum(returns_array > 0)) + win_rate = wins / len(returns) + + # Sharpe ratio (annualized, assuming daily returns) + mean_return = float(np.mean(returns_array)) + std_return = float(np.std(returns_array)) + sharpe = (mean_return / std_return * np.sqrt(252)) if std_return > 0 else 0 + + # Check thresholds + reasons = [] + severity = 0.0 + + if sharpe < self.config.min_sharpe_ratio: + reasons.append(f"Sharpe {sharpe:.2f} < {self.config.min_sharpe_ratio}") + severity = max(severity, (self.config.min_sharpe_ratio - sharpe) / self.config.min_sharpe_ratio) + + if win_rate < self.config.min_win_rate: + reasons.append(f"Win rate {win_rate*100:.1f}% < {self.config.min_win_rate*100:.1f}%") + severity = max(severity, (self.config.min_win_rate - win_rate) / self.config.min_win_rate) + + if reasons: + result = ProtectionResult( + should_block=True, + reason=" | ".join(reasons), + protection_type=ProtectionType.LOW_PERFORMANCE, + severity=severity + ) + self.record_check(blocked=True) + return result + + return ProtectionResult( + should_block=False, + reason=f"Performance acceptable (Sharpe: {sharpe:.2f}, Win rate: {win_rate*100:.1f}%)", + protection_type=ProtectionType.LOW_PERFORMANCE, + severity=severity + ) diff --git a/rdagent/components/backtesting/protections/max_drawdown.py b/rdagent/components/backtesting/protections/max_drawdown.py new file mode 100644 index 00000000..b10825ea --- /dev/null +++ b/rdagent/components/backtesting/protections/max_drawdown.py @@ -0,0 +1,75 @@ +""" +Maximum Drawdown Protection + +Blocks trading when portfolio drawdown exceeds threshold. +""" + +from dataclasses import dataclass +from datetime import datetime, timedelta +from .base import BaseProtection, ProtectionConfig, ProtectionResult, ProtectionType, ProtectionScope + + +@dataclass +class MaxDrawdownConfig(ProtectionConfig): + """Configuration for MaxDrawdown protection.""" + max_drawdown_pct: float = 0.15 # Block if drawdown > 15% + + +class MaxDrawdownProtection(BaseProtection): + """ + Blocks trading when drawdown exceeds safe threshold. + + This prevents the system from continuing to trade during a losing streak, + giving the market time to stabilize. + """ + + def __init__(self, config: MaxDrawdownConfig): + super().__init__(config) + self.config: MaxDrawdownConfig = config + + @property + def scope(self) -> ProtectionScope: + return ProtectionScope.GLOBAL + + def check( + self, + returns: list[float], + timestamps: list[datetime], + current_equity: float, + peak_equity: float, + **kwargs + ) -> ProtectionResult: + """Check if drawdown exceeds threshold.""" + self.record_check() + + if not self.config.enabled: + return ProtectionResult( + should_block=False, + reason="Protection disabled", + protection_type=ProtectionType.MAX_DRAWDOWN + ) + + drawdown = self.calculate_drawdown(current_equity, peak_equity) + severity = abs(drawdown) / self.config.max_drawdown_pct if self.config.max_drawdown_pct > 0 else 0 + + if abs(drawdown) > self.config.max_drawdown_pct: + # Calculate block duration based on severity + block_hours = int(self.config.lookback_period_hours * severity) + block_hours = min(block_hours, 168) # Max 1 week + + result = ProtectionResult( + should_block=True, + reason=f"Drawdown {abs(drawdown)*100:.1f}% exceeds max {self.config.max_drawdown_pct*100:.1f}%", + until=datetime.now() + timedelta(hours=block_hours), + protection_type=ProtectionType.MAX_DRAWDOWN, + severity=severity + ) + self.record_check(blocked=True) + return result + + return ProtectionResult( + should_block=False, + reason=f"Drawdown {abs(drawdown)*100:.1f}% within safe range", + protection_type=ProtectionType.MAX_DRAWDOWN, + severity=severity + ) diff --git a/rdagent/components/backtesting/protections/protection_manager.py b/rdagent/components/backtesting/protections/protection_manager.py new file mode 100644 index 00000000..f563e81b --- /dev/null +++ b/rdagent/components/backtesting/protections/protection_manager.py @@ -0,0 +1,176 @@ +""" +Protection Manager + +Orchestrates multiple trading protections. +""" + +from datetime import datetime +from typing import Optional, List +from .base import ( + BaseProtection, + ProtectionConfig, + ProtectionResult, + ProtectionType, + ProtectionScope +) +from .max_drawdown import MaxDrawdownProtection, MaxDrawdownConfig +from .cooldown import CooldownProtection, CooldownConfig +from .stoploss_guard import StoplossGuardProtection, StoplossGuardConfig +from .low_performance import LowPerformanceProtection, LowPerformanceConfig + + +class ProtectionManager: + """ + Manages multiple trading protections. + + Run all active protections and aggregate their results. + If ANY protection returns should_block=True, trading is blocked. + """ + + def __init__(self): + self.protections: List[BaseProtection] = [] + self.active_blocks: List[ProtectionResult] = [] + + def add_protection(self, protection: BaseProtection): + """Add a protection to the manager.""" + self.protections.append(protection) + + def remove_protection(self, protection_type: ProtectionType): + """Remove a protection by type.""" + type_to_name = { + ProtectionType.MAX_DRAWDOWN: "MaxDrawdownProtection", + ProtectionType.COOLDOWN: "CooldownProtection", + ProtectionType.STOPLOSS_GUARD: "StoplossGuardProtection", + ProtectionType.LOW_PERFORMANCE: "LowPerformanceProtection", + } + class_name = type_to_name.get(protection_type, protection_type.value) + self.protections = [ + p for p in self.protections + if p.__class__.__name__ != class_name + ] + + def check_all( + self, + returns: list[float], + timestamps: list[datetime], + current_equity: float, + peak_equity: float, + **kwargs + ) -> ProtectionResult: + """ + Run all protections and aggregate results. + + Returns + ------- + ProtectionResult + Combined result from all protections + """ + all_results = [] + + for protection in self.protections: + result = protection.check( + returns=returns, + timestamps=timestamps, + current_equity=current_equity, + peak_equity=peak_equity, + **kwargs + ) + all_results.append(result) + + # Check if any protection is blocking + blocking = [r for r in all_results if r.should_block] + + if blocking: + # Find most severe block + most_severe = max(blocking, key=lambda r: r.severity) + + # Clean up expired blocks + self.active_blocks = [ + b for b in self.active_blocks + if b.until is None or datetime.now() < b.until + ] + + # Add new block + self.active_blocks.append(most_severe) + + # Combine reasons + reasons = [r.reason for r in blocking] + + return ProtectionResult( + should_block=True, + reason=f"Trading blocked: {'; '.join(reasons)}", + until=most_severe.until, + severity=most_severe.severity + ) + + return ProtectionResult( + should_block=False, + reason="All protections passed", + severity=0.0 + ) + + def get_active_blocks(self) -> List[ProtectionResult]: + """Get currently active protection blocks.""" + # Clean up expired + self.active_blocks = [ + b for b in self.active_blocks + if b.until is None or datetime.now() < b.until + ] + return self.active_blocks + + def get_stats(self) -> dict: + """Get statistics for all protections.""" + return { + "total_protections": len(self.protections), + "active_blocks": len(self.get_active_blocks()), + "protections": [p.get_stats() for p in self.protections] + } + + def create_default_protections(self): + """Create standard protection setup.""" + # Max Drawdown: 15% threshold + self.add_protection( + MaxDrawdownProtection( + MaxDrawdownConfig( + enabled=True, + max_drawdown_pct=0.15, + lookback_period_hours=168 # 1 week + ) + ) + ) + + # Cooldown: 4 hours after 5% loss + self.add_protection( + CooldownProtection( + CooldownConfig( + enabled=True, + cooldown_after_loss_pct=0.05, + cooldown_duration_hours=4 + ) + ) + ) + + # Stoploss Guard: Max 5 stoplosses per day + self.add_protection( + StoplossGuardProtection( + StoplossGuardConfig( + enabled=True, + max_stoplosses_in_period=5, + stoploss_threshold_pct=0.02, + lookback_period_hours=24 + ) + ) + ) + + # Low Performance: Filter bad factors + self.add_protection( + LowPerformanceProtection( + LowPerformanceConfig( + enabled=True, + min_sharpe_ratio=0.5, + min_win_rate=0.40, + min_trades=20, + lookback_period_hours=720 # 30 days + ) + ) + ) diff --git a/rdagent/components/backtesting/protections/stoploss_guard.py b/rdagent/components/backtesting/protections/stoploss_guard.py new file mode 100644 index 00000000..bc56ee85 --- /dev/null +++ b/rdagent/components/backtesting/protections/stoploss_guard.py @@ -0,0 +1,76 @@ +""" +Stoploss Guard Protection + +Detects clusters of stoplosses and blocks trading. +""" + +from dataclasses import dataclass +from datetime import datetime +from .base import BaseProtection, ProtectionConfig, ProtectionResult, ProtectionType, ProtectionScope + + +@dataclass +class StoplossGuardConfig(ProtectionConfig): + """Configuration for StoplossGuard protection.""" + max_stoplosses_in_period: int = 5 # Max stoplosses allowed + stoploss_threshold_pct: float = 0.02 # What counts as stoploss (2%) + + +class StoplossGuardProtection(BaseProtection): + """ + Detects stoploss clusters and blocks trading. + + Multiple stoplosses in short time indicates bad market conditions + or strategy no longer working. + """ + + def __init__(self, config: StoplossGuardConfig): + super().__init__(config) + self.config: StoplossGuardConfig = config + + @property + def scope(self) -> ProtectionScope: + return ProtectionScope.GLOBAL + + def check( + self, + returns: list[float], + timestamps: list[datetime], + current_equity: float, + peak_equity: float, + **kwargs + ) -> ProtectionResult: + """Check for stoploss clusters.""" + self.record_check() + + if not self.config.enabled: + return ProtectionResult( + should_block=False, + reason="Protection disabled", + protection_type=ProtectionType.STOPLOSS_GUARD + ) + + # Count stoplosses (large losses) + stoplosses = [ + r for r in returns + if r < -self.config.stoploss_threshold_pct + ] + + if len(stoplosses) > self.config.max_stoplosses_in_period: + severity = len(stoplosses) / self.config.max_stoplosses_in_period + + result = ProtectionResult( + should_block=True, + reason=f"{len(stoplosses)} stoplosses detected (max {self.config.max_stoplosses_in_period})", + protection_type=ProtectionType.STOPLOSS_GUARD, + severity=severity + ) + self.record_check(blocked=True) + return result + + return ProtectionResult( + should_block=False, + reason=f"{len(stoplosses)} stoplosses (within limit)", + protection_type=ProtectionType.STOPLOSS_GUARD, + severity=len(stoplosses) / max(1, self.config.max_stoplosses_in_period) + ) diff --git a/test/backtesting/test_protections.py b/test/backtesting/test_protections.py new file mode 100644 index 00000000..d1b9b993 --- /dev/null +++ b/test/backtesting/test_protections.py @@ -0,0 +1,455 @@ +""" +Tests for Trading Protection System + +Covers all protection types and ProtectionManager. +""" + +import pytest +import numpy as np +from datetime import datetime, timedelta +from rdagent.components.backtesting.protections.base import ( + ProtectionResult, ProtectionType, ProtectionScope, ProtectionConfig +) +from rdagent.components.backtesting.protections.max_drawdown import ( + MaxDrawdownProtection, MaxDrawdownConfig +) +from rdagent.components.backtesting.protections.cooldown import ( + CooldownProtection, CooldownConfig +) +from rdagent.components.backtesting.protections.stoploss_guard import ( + StoplossGuardProtection, StoplossGuardConfig +) +from rdagent.components.backtesting.protections.low_performance import ( + LowPerformanceProtection, LowPerformanceConfig +) +from rdagent.components.backtesting.protections.protection_manager import ( + ProtectionManager +) + + +class TestMaxDrawdownProtection: + """Test MaxDrawdown protection.""" + + def test_within_threshold(self): + """Test that trading continues when drawdown is acceptable.""" + config = MaxDrawdownConfig(max_drawdown_pct=0.15) + protection = MaxDrawdownProtection(config) + + # 10% drawdown (below 15% threshold) + result = protection.check( + returns=[-0.05, 0.03, -0.02], + timestamps=[datetime.now()] * 3, + current_equity=90000, + peak_equity=100000 + ) + + assert not result.should_block + assert result.protection_type == ProtectionType.MAX_DRAWDOWN + + def test_exceeds_threshold(self): + """Test that trading blocks when drawdown too high.""" + config = MaxDrawdownConfig(max_drawdown_pct=0.15) + protection = MaxDrawdownProtection(config) + + # 20% drawdown (above 15% threshold) + result = protection.check( + returns=[-0.10, -0.05, -0.05], + timestamps=[datetime.now()] * 3, + current_equity=80000, + peak_equity=100000 + ) + + assert result.should_block + assert result.until is not None + assert result.severity > 1.0 + + def test_disabled_protection(self): + """Test that disabled protection never blocks.""" + config = MaxDrawdownConfig(max_drawdown_pct=0.15, enabled=False) + protection = MaxDrawdownProtection(config) + + result = protection.check( + returns=[], + timestamps=[], + current_equity=50000, # 50% drawdown! + peak_equity=100000 + ) + + assert not result.should_block + + def test_calculate_drawdown(self): + """Test drawdown calculation.""" + config = MaxDrawdownConfig() + protection = MaxDrawdownProtection(config) + + assert protection.calculate_drawdown(90000, 100000) == -0.10 + assert protection.calculate_drawdown(100000, 100000) == 0.0 + assert protection.calculate_drawdown(0, 100000) == -1.0 + + def test_block_duration_scales_with_severity(self): + """More severe drawdowns get longer blocks.""" + config = MaxDrawdownConfig(max_drawdown_pct=0.15) + protection = MaxDrawdownProtection(config) + + # Mild breach (16%) + result_mild = protection.check( + returns=[], + timestamps=[], + current_equity=84000, + peak_equity=100000 + ) + + # Severe breach (30%) + result_severe = protection.check( + returns=[], + timestamps=[], + current_equity=70000, + peak_equity=100000 + ) + + assert result_severe.until > result_mild.until + + +class TestCooldownProtection: + """Test Cooldown protection.""" + + def test_no_recent_loss(self): + """Test that no loss means no cooldown.""" + config = CooldownConfig(cooldown_after_loss_pct=0.05) + protection = CooldownProtection(config) + + result = protection.check( + returns=[0.01, 0.02, 0.01], + timestamps=[datetime.now()] * 3, + current_equity=103000, + peak_equity=103000 + ) + + assert not result.should_block + + def test_triggers_after_loss(self): + """Test that cooldown activates after loss.""" + config = CooldownConfig( + cooldown_after_loss_pct=0.05, + cooldown_duration_hours=4 + ) + protection = CooldownProtection(config) + + result = protection.check( + returns=[-0.06], + timestamps=[datetime.now()], + current_equity=94000, + peak_equity=100000 + ) + + assert result.should_block + assert "cooling down" in result.reason.lower() + + def test_expires_after_duration(self): + """Test that cooldown expires after time.""" + config = CooldownConfig( + cooldown_after_loss_pct=0.05, + cooldown_duration_hours=1 + ) + protection = CooldownProtection(config) + + # Trigger loss in the past (2 hours ago, cooldown is 1 hour) + past_time = datetime.now() - timedelta(hours=2) + protection.last_loss_time = past_time + protection.last_loss_pct = -0.06 + + # Should have expired + result = protection.check( + returns=[0.01], + timestamps=[datetime.now()], + current_equity=94500, + peak_equity=100000 + ) + + assert not result.should_block + + +class TestStoplossGuardProtection: + """Test StoplossGuard protection.""" + + def test_within_limit(self): + """Test that few stoplosses don't block.""" + config = StoplossGuardConfig(max_stoplosses_in_period=5) + protection = StoplossGuardProtection(config) + + result = protection.check( + returns=[-0.01, -0.015, -0.02], # 3 stoplosses + timestamps=[], + current_equity=95000, + peak_equity=100000 + ) + + assert not result.should_block + + def test_exceeds_limit(self): + """Test that too many stoplosses block.""" + config = StoplossGuardConfig( + max_stoplosses_in_period=3, + stoploss_threshold_pct=0.02 + ) + protection = StoplossGuardProtection(config) + + result = protection.check( + returns=[-0.03, -0.025, -0.04, -0.021], # 4 stoplosses (all < -2%) + timestamps=[], + current_equity=90000, + peak_equity=100000 + ) + + assert result.should_block + assert "4 stoplosses" in result.reason + + def test_disabled_protection(self): + """Test that disabled protection never blocks.""" + config = StoplossGuardConfig(max_stoplosses_in_period=1, enabled=False) + protection = StoplossGuardProtection(config) + + result = protection.check( + returns=[-0.03, -0.04, -0.05], # Many stoplosses + timestamps=[], + current_equity=80000, + peak_equity=100000 + ) + + assert not result.should_block + + +class TestLowPerformanceProtection: + """Test LowPerformance protection.""" + + def test_insufficient_data(self): + """Test that few trades don't trigger.""" + config = LowPerformanceConfig(min_trades=20) + protection = LowPerformanceProtection(config) + + result = protection.check( + returns=[-0.1, -0.1, -0.1], # Bad but few + timestamps=[], + current_equity=70000, + peak_equity=100000 + ) + + assert not result.should_block + assert "insufficient" in result.reason.lower() + + def test_blocks_poor_sharpe(self): + """Test that low Sharpe blocks.""" + config = LowPerformanceConfig( + min_sharpe_ratio=0.5, + min_trades=20 + ) + protection = LowPerformanceProtection(config) + + # Generate 30 losing trades + returns = [-0.01] * 30 + + result = protection.check( + returns=returns, + timestamps=[], + current_equity=70000, + peak_equity=100000 + ) + + assert result.should_block + assert "sharpe" in result.reason.lower() + + def test_blocks_low_winrate(self): + """Test that low win rate blocks.""" + config = LowPerformanceConfig( + min_win_rate=0.40, + min_trades=20 + ) + protection = LowPerformanceProtection(config) + + # 10% win rate (90% losses) + returns = [-0.01] * 27 + [0.01] * 3 + + result = protection.check( + returns=returns, + timestamps=[], + current_equity=97000, + peak_equity=100000 + ) + + assert result.should_block + assert "win rate" in result.reason.lower() + + def test_acceptable_performance(self): + """Test that good performance passes.""" + config = LowPerformanceConfig( + min_sharpe_ratio=0.5, + min_win_rate=0.40, + min_trades=20 + ) + protection = LowPerformanceProtection(config) + + # 60% win rate + returns = [0.02] * 30 + [-0.01] * 20 + + result = protection.check( + returns=returns, + timestamps=[], + current_equity=110000, + peak_equity=110000 + ) + + assert not result.should_block + + def test_disabled_protection(self): + """Test that disabled protection never blocks.""" + config = LowPerformanceConfig( + min_sharpe_ratio=0.5, + min_win_rate=0.40, + min_trades=20, + enabled=False + ) + protection = LowPerformanceProtection(config) + + # Generate 30 losing trades + returns = [-0.01] * 30 + + result = protection.check( + returns=returns, + timestamps=[], + current_equity=70000, + peak_equity=100000 + ) + + assert not result.should_block + + +class TestProtectionResult: + """Test ProtectionResult dataclass.""" + + def test_is_active_with_no_until(self): + """Test is_active when no time-based block.""" + result = ProtectionResult(should_block=True, reason="test") + assert result.is_active + + result = ProtectionResult(should_block=False, reason="test") + assert not result.is_active + + def test_is_active_with_future_until(self): + """Test is_active with future expiration.""" + result = ProtectionResult( + should_block=True, + reason="test", + until=datetime.now() + timedelta(hours=1) + ) + assert result.is_active + + def test_is_active_with_past_until(self): + """Test is_active with past expiration.""" + result = ProtectionResult( + should_block=True, + reason="test", + until=datetime.now() - timedelta(hours=1) + ) + assert not result.is_active + + +class TestProtectionManager: + """Test ProtectionManager integration.""" + + def test_all_protections_pass(self): + """Test that good conditions pass all protections.""" + manager = ProtectionManager() + manager.create_default_protections() + + result = manager.check_all( + returns=[0.01, 0.02, 0.015], + timestamps=[datetime.now()] * 3, + current_equity=105000, + peak_equity=105000 + ) + + assert not result.should_block + + def test_one_protection_blocks(self): + """Test that single protection blocks all trading.""" + manager = ProtectionManager() + manager.create_default_protections() + + # Trigger max drawdown + result = manager.check_all( + returns=[-0.10, -0.05, -0.05], + timestamps=[datetime.now()] * 3, + current_equity=80000, + peak_equity=100000 + ) + + assert result.should_block + assert "blocked" in result.reason.lower() + + def test_get_stats(self): + """Test statistics collection.""" + manager = ProtectionManager() + manager.create_default_protections() + + # Run some checks + for _ in range(5): + manager.check_all( + returns=[0.01], + timestamps=[datetime.now()], + current_equity=101000, + peak_equity=101000 + ) + + stats = manager.get_stats() + + assert stats["total_protections"] == 4 + assert "protections" in stats + assert len(stats["protections"]) == 4 + + def test_add_remove_protection(self): + """Test adding and removing protections.""" + manager = ProtectionManager() + + # Add custom protection + config = MaxDrawdownConfig(max_drawdown_pct=0.10) + protection = MaxDrawdownProtection(config) + manager.add_protection(protection) + + assert len(manager.protections) == 1 + + # Remove + manager.remove_protection(ProtectionType.MAX_DRAWDOWN) + assert len(manager.protections) == 0 + + def test_get_active_blocks(self): + """Test active blocks retrieval.""" + manager = ProtectionManager() + manager.create_default_protections() + + # Initially no blocks + assert len(manager.get_active_blocks()) == 0 + + # Trigger a block + manager.check_all( + returns=[-0.10, -0.05, -0.05], + timestamps=[datetime.now()] * 3, + current_equity=80000, + peak_equity=100000 + ) + + # Should have active block + blocks = manager.get_active_blocks() + assert len(blocks) >= 1 + + def test_empty_manager_passes(self): + """Test that empty manager always passes.""" + manager = ProtectionManager() + + result = manager.check_all( + returns=[], + timestamps=[], + current_equity=100000, + peak_equity=100000 + ) + + assert not result.should_block diff --git a/test/integration/test_all_features.py b/test/integration/test_all_features.py index 581487dc..3bb0e711 100644 --- a/test/integration/test_all_features.py +++ b/test/integration/test_all_features.py @@ -800,3 +800,98 @@ class TestIntegrationWorkflow: checks = risk_manager.check_limits(weights, vol=0.15, dd=-0.08) assert isinstance(checks, dict) assert all(key in checks for key in ["position_limit", "leverage_limit", "drawdown_limit"]) + + +# ============================================================================= +# 14. PROTECTION MANAGER TESTS +# ============================================================================= + + +class TestProtectionManager: + """Test Protection Manager system.""" + + def test_protections_import(self): + """Test that protections can be imported.""" + from rdagent.components.backtesting.protections import ( + ProtectionManager, + MaxDrawdownProtection, + CooldownProtection, + StoplossGuardProtection, + LowPerformanceProtection + ) + + def test_protection_manager_creates(self): + """Test ProtectionManager can be created.""" + from rdagent.components.backtesting.protections import ProtectionManager + manager = ProtectionManager() + assert manager is not None + + def test_default_protections_configured(self): + """Test default protections can be configured.""" + from rdagent.components.backtesting.protections import ProtectionManager + manager = ProtectionManager() + manager.create_default_protections() + assert len(manager.protections) == 4 + + def test_protection_manager_blocks(self): + """Test ProtectionManager can block trading.""" + from rdagent.components.backtesting.protections import ProtectionManager + from datetime import datetime + + manager = ProtectionManager() + manager.create_default_protections() + + # Trigger max drawdown + result = manager.check_all( + returns=[-0.10, -0.05, -0.05], + timestamps=[datetime.now()] * 3, + current_equity=80000, + peak_equity=100000 + ) + + assert result.should_block + + def test_protection_manager_allows(self): + """Test ProtectionManager allows good conditions.""" + from rdagent.components.backtesting.protections import ProtectionManager + from datetime import datetime + + manager = ProtectionManager() + manager.create_default_protections() + + result = manager.check_all( + returns=[0.01, 0.02, 0.015], + timestamps=[datetime.now()] * 3, + current_equity=105000, + peak_equity=105000 + ) + + assert not result.should_block + + def test_protection_base_classes(self): + """Test that base classes and enums are importable.""" + from rdagent.components.backtesting.protections import ( + BaseProtection, + ProtectionConfig, + ProtectionResult, + ProtectionType, + ProtectionScope + ) + assert BaseProtection is not None + assert ProtectionResult is not None + assert ProtectionType is not None + assert ProtectionScope is not None + + def test_protection_configs_importable(self): + """Test that all config classes are importable.""" + from rdagent.components.backtesting.protections import ( + MaxDrawdownConfig, + CooldownConfig, + StoplossGuardConfig, + LowPerformanceConfig + ) + assert MaxDrawdownConfig is not None + assert CooldownConfig is not None + assert StoplossGuardConfig is not None + assert LowPerformanceConfig is not None +