use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use hyper_strategy::strategy_config::{HysteresisConfig, RegimeRule, StrategyGroup, TaRule};
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct StrategyAdjustment {
#[serde(alias = "regimeRules", alias = "regime_rules")]
pub regime_rules: Option<Vec<RegimeRule>>,
#[serde(alias = "defaultRegime", alias = "default_regime")]
pub default_regime: Option<String>,
pub hysteresis: Option<HysteresisConfig>,
#[serde(alias = "playbookOverrides", alias = "playbook_overrides")]
pub playbook_overrides: Option<HashMap<String, PlaybookOverride>>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct PlaybookOverride {
pub rules: Option<Vec<TaRule>>,
#[serde(alias = "maxPositionSize", alias = "max_position_size")]
pub max_position_size: Option<f64>,
#[serde(alias = "stopLossPct", alias = "stop_loss_pct")]
pub stop_loss_pct: Option<f64>,
#[serde(alias = "takeProfitPct", alias = "take_profit_pct")]
pub take_profit_pct: Option<f64>,
}
#[derive(Debug)]
pub enum AdjustmentError {
MaxPositionExceeded {
requested: f64,
limit: f64,
},
InvalidThreshold {
indicator: String,
value: f64,
reason: String,
},
FileWriteError(String),
}
impl std::fmt::Display for AdjustmentError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::MaxPositionExceeded { requested, limit } => {
write!(f, "max_position_size {} exceeds limit {}", requested, limit)
}
Self::InvalidThreshold {
indicator,
value,
reason,
} => {
write!(
f,
"Invalid threshold for {}: {} ({})",
indicator, value, reason
)
}
Self::FileWriteError(msg) => write!(f, "File write error: {}", msg),
}
}
}
impl std::error::Error for AdjustmentError {}
pub fn validate_adjustment(
adjustment: &StrategyAdjustment,
max_position_usdc: f64,
) -> Result<(), AdjustmentError> {
if let Some(ref overrides) = adjustment.playbook_overrides {
for (_, pb_override) in overrides {
if let Some(max_pos) = pb_override.max_position_size {
if max_pos > max_position_usdc {
return Err(AdjustmentError::MaxPositionExceeded {
requested: max_pos,
limit: max_position_usdc,
});
}
if max_pos < 0.0 {
return Err(AdjustmentError::InvalidThreshold {
indicator: "max_position_size".into(),
value: max_pos,
reason: "must be non-negative".into(),
});
}
}
if let Some(sl) = pb_override.stop_loss_pct {
if sl < 0.0 || sl > 100.0 {
return Err(AdjustmentError::InvalidThreshold {
indicator: "stop_loss_pct".into(),
value: sl,
reason: "must be between 0 and 100".into(),
});
}
}
if let Some(tp) = pb_override.take_profit_pct {
if tp < 0.0 || tp > 1000.0 {
return Err(AdjustmentError::InvalidThreshold {
indicator: "take_profit_pct".into(),
value: tp,
reason: "must be between 0 and 1000".into(),
});
}
}
}
}
Ok(())
}
pub fn apply_adjustment(group: &mut StrategyGroup, adjustment: &StrategyAdjustment) {
if let Some(ref rules) = adjustment.regime_rules {
group.regime_rules = rules.clone();
}
if let Some(ref regime) = adjustment.default_regime {
group.default_regime = regime.clone();
}
if let Some(ref hyst) = adjustment.hysteresis {
group.hysteresis = hyst.clone();
}
if let Some(ref overrides) = adjustment.playbook_overrides {
for (regime_name, pb_override) in overrides {
if let Some(playbook) = group.playbooks.get_mut(regime_name) {
if let Some(ref rules) = pb_override.rules {
playbook.rules = rules.clone();
}
if let Some(max_pos) = pb_override.max_position_size {
playbook.max_position_size = max_pos;
}
if let Some(sl) = pb_override.stop_loss_pct {
playbook.stop_loss_pct = Some(sl);
}
if let Some(tp) = pb_override.take_profit_pct {
playbook.take_profit_pct = Some(tp);
}
}
}
}
}
pub fn save_with_backup(groups: &[StrategyGroup]) -> Result<(), AdjustmentError> {
use hyper_strategy::strategy_config::{
load_strategy_groups_from_disk_pub, save_strategy_groups_to_disk,
};
let current = load_strategy_groups_from_disk_pub();
if !current.is_empty() {
let backup = serde_json::to_string_pretty(¤t)
.map_err(|e| AdjustmentError::FileWriteError(e.to_string()))?;
let backup_path = dirs::data_dir()
.unwrap_or_else(|| std::path::PathBuf::from("."))
.join("hyper-agent")
.join("strategy_groups.bak.json");
std::fs::create_dir_all(backup_path.parent().unwrap())
.map_err(|e| AdjustmentError::FileWriteError(e.to_string()))?;
std::fs::write(&backup_path, backup)
.map_err(|e| AdjustmentError::FileWriteError(e.to_string()))?;
}
save_strategy_groups_to_disk(groups)
.map_err(|e| AdjustmentError::FileWriteError(e.to_string()))?;
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use hyper_strategy::strategy_config::{HysteresisConfig, Playbook, StrategyGroup};
fn make_group() -> StrategyGroup {
let mut playbooks = HashMap::new();
playbooks.insert(
"bull".to_string(),
Playbook {
rules: vec![],
entry_rules: vec![],
exit_rules: vec![],
system_prompt: "bull".into(),
max_position_size: 1000.0,
stop_loss_pct: Some(5.0),
take_profit_pct: Some(10.0),
timeout_secs: None,
side: None,
},
);
StrategyGroup {
id: "sg-test".into(),
name: "Test".into(),
vault_address: None,
is_active: true,
created_at: "2026-01-01".into(),
symbol: "BTC-USD".into(),
interval_secs: 300,
regime_rules: vec![],
default_regime: "bull".into(),
hysteresis: HysteresisConfig {
min_hold_secs: 3600,
confirmation_count: 3,
},
playbooks,
}
}
#[test]
fn validate_passes_within_limits() {
let adj = StrategyAdjustment {
regime_rules: None,
default_regime: None,
hysteresis: None,
playbook_overrides: Some(HashMap::from([(
"bull".into(),
PlaybookOverride {
rules: None,
max_position_size: Some(500.0),
stop_loss_pct: Some(3.0),
take_profit_pct: Some(15.0),
},
)])),
};
assert!(validate_adjustment(&adj, 10000.0).is_ok());
}
#[test]
fn validate_rejects_exceeding_max_position() {
let adj = StrategyAdjustment {
regime_rules: None,
default_regime: None,
hysteresis: None,
playbook_overrides: Some(HashMap::from([(
"bull".into(),
PlaybookOverride {
rules: None,
max_position_size: Some(50000.0),
stop_loss_pct: None,
take_profit_pct: None,
},
)])),
};
assert!(matches!(
validate_adjustment(&adj, 10000.0),
Err(AdjustmentError::MaxPositionExceeded { .. })
));
}
#[test]
fn validate_rejects_negative_stop_loss() {
let adj = StrategyAdjustment {
regime_rules: None,
default_regime: None,
hysteresis: None,
playbook_overrides: Some(HashMap::from([(
"bull".into(),
PlaybookOverride {
rules: None,
max_position_size: None,
stop_loss_pct: Some(-5.0),
take_profit_pct: None,
},
)])),
};
assert!(matches!(
validate_adjustment(&adj, 10000.0),
Err(AdjustmentError::InvalidThreshold { .. })
));
}
#[test]
fn apply_changes_playbook_params() {
let mut group = make_group();
let adj = StrategyAdjustment {
regime_rules: None,
default_regime: Some("neutral".into()),
hysteresis: None,
playbook_overrides: Some(HashMap::from([(
"bull".into(),
PlaybookOverride {
rules: None,
max_position_size: Some(2000.0),
stop_loss_pct: Some(3.0),
take_profit_pct: None,
},
)])),
};
apply_adjustment(&mut group, &adj);
assert_eq!(group.default_regime, "neutral");
let bull = group.playbooks.get("bull").unwrap();
assert_eq!(bull.max_position_size, 2000.0);
assert_eq!(bull.stop_loss_pct, Some(3.0));
assert_eq!(bull.take_profit_pct, Some(10.0)); }
#[test]
fn apply_preserves_immutable_fields() {
let mut group = make_group();
let original_id = group.id.clone();
let original_symbol = group.symbol.clone();
apply_adjustment(
&mut group,
&StrategyAdjustment {
regime_rules: None,
default_regime: None,
hysteresis: None,
playbook_overrides: None,
},
);
assert_eq!(group.id, original_id);
assert_eq!(group.symbol, original_symbol);
}
}