Skip to main content

do_memory_core/reward/
base.rs

1//! Base reward calculation based on task outcome
2
3use crate::episode::Episode;
4use crate::types::TaskOutcome;
5
6/// Calculate base reward from episode outcome
7///
8/// Returns:
9/// - 1.0 for complete success
10/// - Proportional value (0.0-1.0) for partial success
11/// - 0.0 for failure or incomplete episodes
12pub fn calculate_base_reward(episode: &Episode) -> f32 {
13    match &episode.outcome {
14        Some(TaskOutcome::Success { .. }) => 1.0,
15        Some(TaskOutcome::PartialSuccess {
16            completed, failed, ..
17        }) => {
18            // Proportional reward based on completion ratio
19            let total = completed.len() + failed.len();
20            if total == 0 {
21                0.5 // Default for partial success with no specifics
22            } else {
23                completed.len() as f32 / total as f32
24            }
25        }
26        Some(TaskOutcome::Failure { .. }) => 0.0,
27        // Abstention is not failure: base is 0.3 (above failure, below partial)
28        // The abstention_score component in RewardScore handles timeliness.
29        Some(TaskOutcome::Abstained { .. }) => 0.3,
30        None => 0.0, // Not completed
31    }
32}
33
34#[cfg(test)]
35mod tests {
36    use super::*;
37    use crate::types::{ComplexityLevel, TaskContext, TaskType};
38
39    fn create_test_episode() -> Episode {
40        let context = TaskContext {
41            language: Some("rust".to_string()),
42            framework: None,
43            complexity: ComplexityLevel::Simple,
44            domain: "testing".to_string(),
45            tags: vec![],
46        };
47        Episode::new("Test task".to_string(), context, TaskType::Testing)
48    }
49
50    #[test]
51    fn test_base_reward_success() {
52        let mut episode = create_test_episode();
53        episode.complete(TaskOutcome::Success {
54            verdict: "Done".to_string(),
55            artifacts: vec![],
56        });
57        assert_eq!(calculate_base_reward(&episode), 1.0);
58    }
59
60    #[test]
61    fn test_base_reward_failure() {
62        let mut episode = create_test_episode();
63        episode.complete(TaskOutcome::Failure {
64            reason: "Failed".to_string(),
65            error_details: None,
66        });
67        assert_eq!(calculate_base_reward(&episode), 0.0);
68    }
69
70    #[test]
71    fn test_base_reward_partial_success() {
72        let mut episode = create_test_episode();
73        episode.complete(TaskOutcome::PartialSuccess {
74            verdict: "Partial".to_string(),
75            completed: vec!["a".to_string(), "b".to_string()],
76            failed: vec!["c".to_string()],
77        });
78        // 2 out of 3 = 0.667
79        assert!((calculate_base_reward(&episode) - 0.667).abs() < 0.01);
80    }
81
82    #[test]
83    fn test_base_reward_incomplete() {
84        let episode = create_test_episode();
85        assert_eq!(calculate_base_reward(&episode), 0.0);
86    }
87}