#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum RlStrategy {
Bandit,
ActorCritic,
}
impl RlStrategy {
#[must_use]
pub fn parse(name: &str) -> Self {
match name.trim().to_ascii_lowercase().as_str() {
"actor_critic" | "actor-critic" => RlStrategy::ActorCritic,
_ => RlStrategy::Bandit,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct RewardSignal {
pub success: bool,
pub latency_secs: f64,
pub cost_usd: f64,
}
impl RewardSignal {
#[must_use]
pub fn score(&self, latency_weight: f64) -> f64 {
let w = latency_weight.clamp(0.0, 1.0);
let success_term = if self.success { 1.0 } else { -1.0 };
let latency_penalty = -(self.latency_secs / (1.0 + self.latency_secs));
let cost_penalty = -(self.cost_usd / (1.0 + self.cost_usd));
success_term * (1.0 - w) + latency_penalty * w + cost_penalty * 0.1
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn strategy_parse_is_case_insensitive() {
assert_eq!(RlStrategy::parse("actor_critic"), RlStrategy::ActorCritic);
assert_eq!(RlStrategy::parse("Actor-Critic"), RlStrategy::ActorCritic);
assert_eq!(RlStrategy::parse("bandit"), RlStrategy::Bandit);
assert_eq!(RlStrategy::parse("anything"), RlStrategy::Bandit);
}
#[test]
fn score_rewards_fast_success_and_penalizes_failure() {
let good = RewardSignal { success: true, latency_secs: 0.1, cost_usd: 0.0 };
let bad = RewardSignal { success: false, latency_secs: 10.0, cost_usd: 1.0 };
assert!(good.score(0.5) > 0.0);
assert!(bad.score(0.5) < 0.0);
}
}