1use crate::aixi::agent::Agent;
8use crate::aixi::aiqi::AiqiAgent;
9use crate::aixi::common::{
10 Action, EXPLORE_RANDOM_SALT, RandomGenerator, Reward, resolve_random_seed,
11};
12use crate::aixi::environment::Environment;
13use crate::aixi::planner_runtime::validate_environment_interface;
14use crate::aixi::warmstart::{
15 WarmStartExactJhAgent, WarmStartExactJhError, WarmStartExactJhTeacherDataset,
16};
17use crate::spec::{CompiledPlannerController, CompiledPlannerRunSpec, PlannerRuntimeSpec};
18use std::error::Error;
19use std::fmt;
20
21pub use crate::aixi::planner_runtime::{
22 build_planner_environment, compile_planner_run_document,
23 load_warmstart_exact_jh_teacher_dataset, validate_action_alphabet,
24 validate_warmstart_exact_jh_teacher_contract,
25};
26
27#[derive(Clone, Copy, Debug, Eq, PartialEq)]
29#[non_exhaustive]
30pub enum PlannerPhase {
31 Learn,
33 Eval,
35}
36
37#[derive(Clone, Copy, Debug, PartialEq)]
39#[non_exhaustive]
40pub struct PlannerSchedule {
41 pub learn_cycles: usize,
43 pub eval_cycles: usize,
45 pub explore_epsilon: f64,
47 pub explore_gamma: f64,
49}
50
51impl PlannerSchedule {
52 pub fn new(learn_cycles: usize, eval_cycles: usize) -> Self {
54 Self {
55 learn_cycles,
56 eval_cycles,
57 explore_epsilon: 0.0,
58 explore_gamma: 1.0,
59 }
60 }
61
62 pub fn from_runtime(runtime: &PlannerRuntimeSpec) -> Self {
64 let terminate_lifetime: usize = runtime.terminate_lifetime;
65 let (learn_cycles, eval_cycles) = match (runtime.learn_cycles, runtime.eval_cycles) {
66 (Some(learn), Some(eval)) => (learn, eval),
67 (Some(learn), None) => (learn, 0usize),
68 (None, Some(eval)) => (terminate_lifetime, eval),
69 (None, None) => (terminate_lifetime, 0usize),
70 };
71 Self {
72 learn_cycles,
73 eval_cycles,
74 explore_epsilon: runtime.explore_epsilon,
75 explore_gamma: runtime.explore_gamma,
76 }
77 }
78
79 pub fn extra_exploration(&self, step: usize) -> f64 {
81 if self.explore_epsilon > 0.0 {
82 let exponent = i32::try_from(step).unwrap_or(i32::MAX);
83 (self.explore_epsilon * self.explore_gamma.powi(exponent)).min(1.0)
84 } else {
85 0.0
86 }
87 }
88}
89
90#[derive(Clone, Copy, Debug, Eq, PartialEq)]
92#[non_exhaustive]
93pub enum PlannerActionProvenance {
94 Greedy,
96 Exploratory,
98}
99
100impl PlannerActionProvenance {
101 pub fn as_str(self) -> &'static str {
103 match self {
104 Self::Greedy => "greedy",
105 Self::Exploratory => "exploratory",
106 }
107 }
108
109 pub fn from_jsonl_str(value: &str) -> Result<Self, WarmStartExactJhError> {
111 match value {
112 "greedy" => Ok(Self::Greedy),
113 "exploratory" => Ok(Self::Exploratory),
114 other => Err(WarmStartExactJhError::InvalidTelemetry {
115 reason: format!("unknown action provenance '{other}'"),
116 }),
117 }
118 }
119}
120
121#[derive(Clone, Debug, Eq, PartialEq)]
123pub struct PlannerCycleOutcome {
124 pub pre_observations: Vec<u64>,
126 pub pre_reward: Reward,
128 pub action: Action,
130 pub observations: Vec<u64>,
132 pub reward: Reward,
134 pub provenance: PlannerActionProvenance,
136}
137
138pub trait PlannerCycleObserver {
150 fn observe_percept(
156 &mut self,
157 step: usize,
158 observations: &[u64],
159 reward: Reward,
160 ) -> Result<(), PlannerAgentError>;
161
162 fn observe_action(
164 &mut self,
165 step: usize,
166 action: Action,
167 provenance: PlannerActionProvenance,
168 ) -> Result<(), PlannerAgentError>;
169
170 fn end_cycle(&mut self, step: usize) -> Result<(), PlannerAgentError>;
172}
173
174#[derive(Clone, Copy, Debug, Default)]
176pub struct NullPlannerObserver;
177
178impl PlannerCycleObserver for NullPlannerObserver {
179 fn observe_percept(
180 &mut self,
181 _step: usize,
182 _observations: &[u64],
183 _reward: Reward,
184 ) -> Result<(), PlannerAgentError> {
185 Ok(())
186 }
187
188 fn observe_action(
189 &mut self,
190 _step: usize,
191 _action: Action,
192 _provenance: PlannerActionProvenance,
193 ) -> Result<(), PlannerAgentError> {
194 Ok(())
195 }
196
197 fn end_cycle(&mut self, _step: usize) -> Result<(), PlannerAgentError> {
198 Ok(())
199 }
200}
201
202impl PlannerCycleObserver for () {
203 fn observe_percept(
204 &mut self,
205 _step: usize,
206 _observations: &[u64],
207 _reward: Reward,
208 ) -> Result<(), PlannerAgentError> {
209 Ok(())
210 }
211
212 fn observe_action(
213 &mut self,
214 _step: usize,
215 _action: Action,
216 _provenance: PlannerActionProvenance,
217 ) -> Result<(), PlannerAgentError> {
218 Ok(())
219 }
220
221 fn end_cycle(&mut self, _step: usize) -> Result<(), PlannerAgentError> {
222 Ok(())
223 }
224}
225
226pub fn validate_obs_stream_len(expected: usize, actual: usize) -> Result<(), PlannerAgentError> {
228 if actual != expected {
229 return Err(PlannerAgentError::EnvironmentInterface {
230 reason: format!(
231 "observation stream length mismatch: expected {expected}, got {actual}"
232 ),
233 });
234 }
235 Ok(())
236}
237
238pub struct PlannerEnvironment {
240 env: Box<dyn Environment>,
241 observation_stream_len: usize,
242 observations: Vec<u64>,
243 reward: Reward,
244}
245
246impl PlannerEnvironment {
247 pub fn new(
249 compiled: &CompiledPlannerRunSpec,
250 env: Box<dyn Environment>,
251 ) -> Result<Self, PlannerAgentError> {
252 Self::new_with_seed(
253 compiled,
254 env,
255 resolve_random_seed(compiled.runtime().random_seed),
256 )
257 }
258
259 pub fn new_with_seed(
261 compiled: &CompiledPlannerRunSpec,
262 mut env: Box<dyn Environment>,
263 random_seed: u64,
264 ) -> Result<Self, PlannerAgentError> {
265 validate_environment_interface(compiled, env.as_ref()).map_err(|err| {
266 PlannerAgentError::EnvironmentInterface {
267 reason: err.to_string(),
268 }
269 })?;
270 env.set_random_seed(random_seed);
271 let observation_stream_len: usize = compiled.interface().observation_stream_len;
272 let observations = env.drain_observations();
273 validate_obs_stream_len(observation_stream_len, observations.len())?;
274 let reward: Reward = env.get_reward();
275 Ok(Self {
276 env,
277 observation_stream_len,
278 observations,
279 reward,
280 })
281 }
282
283 pub fn observations(&self) -> &[u64] {
285 &self.observations
286 }
287
288 pub fn reward(&self) -> Reward {
290 self.reward
291 }
292
293 pub fn perform_action(&mut self, action: Action) -> Result<Reward, PlannerAgentError> {
295 self.env.perform_action(action);
296 self.observations = self.env.drain_observations();
297 validate_obs_stream_len(self.observation_stream_len, self.observations.len())?;
298 self.reward = self.env.get_reward();
299 Ok(self.reward)
300 }
301}
302
303pub trait EnvironmentFactory {
305 fn build(&self) -> Result<Box<dyn Environment>, PlannerAgentError>;
307}
308
309pub trait PlannerAgent {
311 fn reseed_for_episode(&mut self, _random_seed: u64) {}
315
316 fn reset_for_episode(&mut self, random_seed: u64) {
323 self.reseed_for_episode(random_seed);
324 }
325
326 fn run_cycle(
328 &mut self,
329 phase: PlannerPhase,
330 step: usize,
331 schedule: &PlannerSchedule,
332 env: &mut PlannerEnvironment,
333 observer: &mut dyn PlannerCycleObserver,
334 ) -> Result<PlannerCycleOutcome, PlannerAgentError>;
335}
336
337pub struct McAixiPlannerAgent {
341 agent: Agent,
342 prev_action: Action,
343 explore_rng: RandomGenerator,
344}
345
346impl McAixiPlannerAgent {
347 pub fn from_compiled(compiled: &CompiledPlannerRunSpec) -> Result<Self, PlannerAgentError> {
349 let agent = Agent::from_compiled_planner_run(compiled)
350 .map_err(|err| PlannerAgentError::McAixi(err.to_string()))?;
351 let explore_rng: RandomGenerator =
352 RandomGenerator::from_seed(resolve_random_seed(compiled.runtime().random_seed))
353 .fork_with(EXPLORE_RANDOM_SALT);
354 Ok(Self {
355 agent,
356 prev_action: 0,
357 explore_rng,
358 })
359 }
360}
361
362impl PlannerAgent for McAixiPlannerAgent {
363 fn reseed_for_episode(&mut self, random_seed: u64) {
364 self.agent.reseed_random(random_seed);
365 self.explore_rng = RandomGenerator::from_seed(random_seed).fork_with(EXPLORE_RANDOM_SALT);
366 }
367
368 fn reset_for_episode(&mut self, random_seed: u64) {
369 self.reseed_for_episode(random_seed);
370 self.prev_action = 0;
371 self.agent.reset_planner_state();
372 }
373
374 fn run_cycle(
375 &mut self,
376 phase: PlannerPhase,
377 step: usize,
378 schedule: &PlannerSchedule,
379 env: &mut PlannerEnvironment,
380 observer: &mut dyn PlannerCycleObserver,
381 ) -> Result<PlannerCycleOutcome, PlannerAgentError> {
382 let pre_observations: Vec<u64> = env.observations().to_vec();
383 let pre_reward: Reward = env.reward();
384 observer.observe_percept(step, &pre_observations, pre_reward)?;
385 self.agent
386 .model_update_percept_stream(&pre_observations, pre_reward);
387 let mut provenance = PlannerActionProvenance::Greedy;
388 let action: Action = match phase {
389 PlannerPhase::Learn => {
390 let explore_p: f64 = schedule.extra_exploration(step);
391 if explore_p > 0.0 && self.explore_rng.gen_bool(explore_p) {
392 provenance = PlannerActionProvenance::Exploratory;
393 self.explore_rng.gen_range(env.env.get_num_actions().get()) as u64
394 } else {
395 self.agent
396 .get_planned_action(&pre_observations, pre_reward, self.prev_action)
397 }
398 }
399 PlannerPhase::Eval => {
400 self.agent
401 .get_planned_action(&pre_observations, pre_reward, self.prev_action)
402 }
403 };
404 observer.observe_action(step, action, provenance)?;
405 self.agent.model_update_action_external(action);
406 let reward: Reward = env.perform_action(action)?;
407 self.prev_action = action;
408 observer.end_cycle(step)?;
409 let total_cycles: usize = schedule.learn_cycles.saturating_add(schedule.eval_cycles);
411 if step.checked_add(1) == Some(total_cycles) {
412 observer.observe_percept(total_cycles, env.observations(), reward)?;
413 }
414 Ok(PlannerCycleOutcome {
415 pre_observations,
416 pre_reward,
417 action,
418 observations: env.observations().to_vec(),
419 reward,
420 provenance,
421 })
422 }
423}
424
425pub struct AiqiDiscountedPlannerAgent {
429 agent: AiqiAgent,
430}
431
432impl AiqiDiscountedPlannerAgent {
433 pub fn from_compiled(compiled: &CompiledPlannerRunSpec) -> Result<Self, PlannerAgentError> {
435 Ok(Self {
436 agent: AiqiAgent::from_compiled_planner_run(compiled)
437 .map_err(PlannerAgentError::Aiqi)?,
438 })
439 }
440}
441
442impl PlannerAgent for AiqiDiscountedPlannerAgent {
443 fn reseed_for_episode(&mut self, random_seed: u64) {
444 self.agent.reseed_random(random_seed);
445 }
446
447 fn run_cycle(
448 &mut self,
449 phase: PlannerPhase,
450 step: usize,
451 schedule: &PlannerSchedule,
452 env: &mut PlannerEnvironment,
453 observer: &mut dyn PlannerCycleObserver,
454 ) -> Result<PlannerCycleOutcome, PlannerAgentError> {
455 let pre_observations: Vec<u64> = env.observations().to_vec();
456 let pre_reward: Reward = env.reward();
457 let (action, explored) = match phase {
458 PlannerPhase::Learn => self
459 .agent
460 .get_planned_action_with_extra_exploration_flag(schedule.extra_exploration(step)),
461 PlannerPhase::Eval => (self.agent.get_planned_action(), false),
462 };
463 let provenance = if explored {
464 PlannerActionProvenance::Exploratory
465 } else {
466 PlannerActionProvenance::Greedy
467 };
468 observer.observe_action(step, action, provenance)?;
469 let reward: Reward = env.perform_action(action)?;
470 observer.observe_percept(step, env.observations(), reward)?;
471 self.agent
472 .observe_transition(action, env.observations(), reward)
473 .map_err(PlannerAgentError::Aiqi)?;
474 observer.end_cycle(step)?;
475 Ok(PlannerCycleOutcome {
476 pre_observations,
477 pre_reward,
478 action,
479 observations: env.observations().to_vec(),
480 reward,
481 provenance,
482 })
483 }
484}
485
486pub struct WarmStartExactJhPlannerAgent {
490 agent: WarmStartExactJhAgent,
491}
492
493impl WarmStartExactJhPlannerAgent {
494 pub fn from_compiled(
496 compiled: &CompiledPlannerRunSpec,
497 teacher: WarmStartExactJhTeacherDataset,
498 ) -> Result<Self, PlannerAgentError> {
499 Ok(Self {
500 agent: WarmStartExactJhAgent::from_compiled_planner_run(compiled, teacher)
501 .map_err(PlannerAgentError::WarmStart)?,
502 })
503 }
504}
505
506impl PlannerAgent for WarmStartExactJhPlannerAgent {
507 fn reseed_for_episode(&mut self, random_seed: u64) {
508 self.agent.reseed_random(random_seed);
509 }
510
511 fn run_cycle(
512 &mut self,
513 phase: PlannerPhase,
514 step: usize,
515 schedule: &PlannerSchedule,
516 env: &mut PlannerEnvironment,
517 observer: &mut dyn PlannerCycleObserver,
518 ) -> Result<PlannerCycleOutcome, PlannerAgentError> {
519 let pre_observations: Vec<u64> = env.observations().to_vec();
520 let pre_reward: Reward = env.reward();
521 let (action, explored) = match phase {
522 PlannerPhase::Learn => self
523 .agent
524 .try_get_planned_action_with_extra_exploration_flag(
525 schedule.extra_exploration(step),
526 )
527 .map_err(PlannerAgentError::WarmStart)?,
528 PlannerPhase::Eval => (
529 self.agent
530 .try_get_planned_action()
531 .map_err(PlannerAgentError::WarmStart)?,
532 false,
533 ),
534 };
535 let provenance = if explored {
536 PlannerActionProvenance::Exploratory
537 } else {
538 PlannerActionProvenance::Greedy
539 };
540 observer.observe_action(step, action, provenance)?;
541 let reward: Reward = env.perform_action(action)?;
542 observer.observe_percept(step, env.observations(), reward)?;
543 self.agent
544 .observe_transition(action, env.observations(), reward)
545 .map_err(PlannerAgentError::WarmStart)?;
546 observer.end_cycle(step)?;
547 Ok(PlannerCycleOutcome {
548 pre_observations,
549 pre_reward,
550 action,
551 observations: env.observations().to_vec(),
552 reward,
553 provenance,
554 })
555 }
556}
557
558enum PlannerControllerAgentKind {
559 McAixi(McAixiPlannerAgent),
560 AiqiDiscounted(AiqiDiscountedPlannerAgent),
561 WarmStartExactJh(WarmStartExactJhPlannerAgent),
562}
563
564pub struct PlannerControllerAgent {
566 inner: PlannerControllerAgentKind,
567}
568
569impl PlannerControllerAgent {
570 pub fn from_compiled(compiled: &CompiledPlannerRunSpec) -> Result<Self, PlannerAgentError> {
572 let inner = match compiled.controller() {
573 CompiledPlannerController::McAixi { .. } => {
574 PlannerControllerAgentKind::McAixi(McAixiPlannerAgent::from_compiled(compiled)?)
575 }
576 CompiledPlannerController::AiqiDiscounted { .. } => {
577 PlannerControllerAgentKind::AiqiDiscounted(
578 AiqiDiscountedPlannerAgent::from_compiled(compiled)?,
579 )
580 }
581 CompiledPlannerController::AiqiWarmstartExactJh {
582 teacher_dataset_asset,
583 ..
584 } => {
585 let teacher =
586 load_warmstart_exact_jh_teacher_dataset(compiled, teacher_dataset_asset)?;
587 PlannerControllerAgentKind::WarmStartExactJh(
588 WarmStartExactJhPlannerAgent::from_compiled(compiled, teacher)?,
589 )
590 }
591 };
592 Ok(Self { inner })
593 }
594
595 pub fn controller_kind(&self) -> &'static str {
597 match &self.inner {
598 PlannerControllerAgentKind::McAixi(_) => "mc_aixi",
599 PlannerControllerAgentKind::AiqiDiscounted(_) => "aiqi_discounted",
600 PlannerControllerAgentKind::WarmStartExactJh(_) => "aiqi_warmstart_exact_jh",
601 }
602 }
603}
604
605impl PlannerAgent for PlannerControllerAgent {
606 fn reseed_for_episode(&mut self, random_seed: u64) {
607 match &mut self.inner {
608 PlannerControllerAgentKind::McAixi(agent) => agent.reseed_for_episode(random_seed),
609 PlannerControllerAgentKind::AiqiDiscounted(agent) => {
610 agent.reseed_for_episode(random_seed);
611 }
612 PlannerControllerAgentKind::WarmStartExactJh(agent) => {
613 agent.reseed_for_episode(random_seed);
614 }
615 }
616 }
617
618 fn reset_for_episode(&mut self, random_seed: u64) {
619 match &mut self.inner {
620 PlannerControllerAgentKind::McAixi(agent) => agent.reset_for_episode(random_seed),
621 PlannerControllerAgentKind::AiqiDiscounted(agent) => {
622 agent.reset_for_episode(random_seed);
623 }
624 PlannerControllerAgentKind::WarmStartExactJh(agent) => {
625 agent.reset_for_episode(random_seed);
626 }
627 }
628 }
629
630 fn run_cycle(
631 &mut self,
632 phase: PlannerPhase,
633 step: usize,
634 schedule: &PlannerSchedule,
635 env: &mut PlannerEnvironment,
636 observer: &mut dyn PlannerCycleObserver,
637 ) -> Result<PlannerCycleOutcome, PlannerAgentError> {
638 match &mut self.inner {
639 PlannerControllerAgentKind::McAixi(agent) => {
640 agent.run_cycle(phase, step, schedule, env, observer)
641 }
642 PlannerControllerAgentKind::AiqiDiscounted(agent) => {
643 agent.run_cycle(phase, step, schedule, env, observer)
644 }
645 PlannerControllerAgentKind::WarmStartExactJh(agent) => {
646 agent.run_cycle(phase, step, schedule, env, observer)
647 }
648 }
649 }
650}
651
652pub struct PlannerRunSession {
654 agent: PlannerControllerAgent,
655 environment: PlannerEnvironment,
656 schedule: PlannerSchedule,
657 next_step: usize,
658}
659
660impl PlannerRunSession {
661 pub fn new(
663 compiled: &CompiledPlannerRunSpec,
664 mut agent: PlannerControllerAgent,
665 env: Box<dyn Environment>,
666 ) -> Result<Self, PlannerAgentError> {
667 let environment = PlannerEnvironment::new(compiled, env)?;
668 let schedule = PlannerSchedule::from_runtime(compiled.runtime());
669 agent.reset_for_episode(resolve_random_seed(compiled.runtime().random_seed));
670 Ok(Self {
671 agent,
672 environment,
673 schedule,
674 next_step: 0,
675 })
676 }
677
678 pub fn schedule(&self) -> &PlannerSchedule {
680 &self.schedule
681 }
682
683 pub fn next_step(&self) -> usize {
685 self.next_step
686 }
687
688 pub fn next_phase(&self) -> Option<PlannerPhase> {
690 let total_cycles: usize = self
691 .schedule
692 .learn_cycles
693 .saturating_add(self.schedule.eval_cycles);
694 if self.next_step >= total_cycles {
695 return None;
696 }
697 Some(if self.next_step < self.schedule.learn_cycles {
698 PlannerPhase::Learn
699 } else {
700 PlannerPhase::Eval
701 })
702 }
703
704 pub fn is_finished(&self) -> bool {
706 self.next_phase().is_none()
707 }
708
709 pub fn run_next_cycle(
711 &mut self,
712 observer: &mut dyn PlannerCycleObserver,
713 ) -> Result<Option<PlannerCycleOutcome>, PlannerAgentError> {
714 let Some(phase) = self.next_phase() else {
715 return Ok(None);
716 };
717 let outcome = self.agent.run_cycle(
718 phase,
719 self.next_step,
720 &self.schedule,
721 &mut self.environment,
722 observer,
723 )?;
724 self.next_step = self.next_step.saturating_add(1);
725 Ok(Some(outcome))
726 }
727}
728
729#[derive(Clone, Debug, PartialEq)]
731pub struct PlannerRunReport {
732 pub learn_total_reward: Reward,
734 pub eval_total_reward: Reward,
736 pub learn_cycles: usize,
738 pub eval_cycles: usize,
740}
741
742pub fn run_episode(
749 agent: &mut dyn PlannerAgent,
750 schedule: &PlannerSchedule,
751 env_factory: &dyn EnvironmentFactory,
752 random_seed: u64,
753 compiled: &CompiledPlannerRunSpec,
754) -> Result<PlannerRunReport, PlannerAgentError> {
755 let env = env_factory.build()?;
756 let mut environment = PlannerEnvironment::new_with_seed(compiled, env, random_seed)?;
757 agent.reset_for_episode(random_seed);
758 let mut observer = NullPlannerObserver;
759 let mut learn_total_reward: Reward = 0;
760 let mut eval_total_reward: Reward = 0;
761 for step in 0..schedule.learn_cycles {
762 let outcome = agent.run_cycle(
763 PlannerPhase::Learn,
764 step,
765 schedule,
766 &mut environment,
767 &mut observer,
768 )?;
769 learn_total_reward = learn_total_reward.saturating_add(outcome.reward);
770 }
771 for offset in 0..schedule.eval_cycles {
772 let step: usize = schedule.learn_cycles + offset;
773 let outcome = agent.run_cycle(
774 PlannerPhase::Eval,
775 step,
776 schedule,
777 &mut environment,
778 &mut observer,
779 )?;
780 eval_total_reward = eval_total_reward.saturating_add(outcome.reward);
781 }
782 Ok(PlannerRunReport {
783 learn_total_reward,
784 eval_total_reward,
785 learn_cycles: schedule.learn_cycles,
786 eval_cycles: schedule.eval_cycles,
787 })
788}
789
790#[derive(Debug)]
792#[non_exhaustive]
793pub enum PlannerAgentError {
794 McAixi(String),
796 Aiqi(crate::aixi::aiqi::AiqiError),
798 WarmStart(WarmStartExactJhError),
800 EnvironmentInterface {
802 reason: String,
804 },
805 Environment {
807 reason: String,
809 },
810 Observer {
812 reason: String,
814 },
815}
816
817impl fmt::Display for PlannerAgentError {
818 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
819 match self {
820 Self::McAixi(err) => write!(f, "{err}"),
821 Self::Aiqi(err) => write!(f, "{err}"),
822 Self::WarmStart(err) => write!(f, "{err}"),
823 Self::EnvironmentInterface { reason }
824 | Self::Environment { reason }
825 | Self::Observer { reason } => f.write_str(reason),
826 }
827 }
828}
829
830impl Error for PlannerAgentError {
831 fn source(&self) -> Option<&(dyn Error + 'static)> {
832 match self {
833 Self::Aiqi(err) => Some(err),
834 Self::WarmStart(err) => Some(err),
835 Self::McAixi(_)
836 | Self::EnvironmentInterface { .. }
837 | Self::Environment { .. }
838 | Self::Observer { .. } => None,
839 }
840 }
841}
842
843#[cfg(all(test, feature = "backend-ctw"))]
844mod tests {
845 use super::*;
846 use crate::aixi::common::{ActionAlphabet, PerceptVal};
847 use crate::aixi::warmstart::{
848 WarmStartExactJhTeacherDataset, WarmStartExactJhTeacherTrace, WarmStartExactJhTransition,
849 standalone_warmstart_teacher_contract_for_compiled_planner_run,
850 };
851 use crate::spec::{SpecDocument, SpecEnvironment};
852 use serde_json::json;
853 use std::path::{Path, PathBuf};
854 use std::sync::Arc;
855 use std::sync::atomic::{AtomicU64, Ordering};
856 use std::time::{SystemTime, UNIX_EPOCH};
857
858 static TEMP_TEST_PATH_COUNTER: AtomicU64 = AtomicU64::new(0);
859
860 fn unique_temp_path(prefix: &str, suffix: &str) -> PathBuf {
861 let counter = TEMP_TEST_PATH_COUNTER.fetch_add(1, Ordering::Relaxed);
862 let nanos = SystemTime::now()
863 .duration_since(UNIX_EPOCH)
864 .map(|d| d.as_nanos())
865 .unwrap_or(0);
866 std::env::temp_dir().join(format!(
867 "{prefix}-{}-{nanos}-{counter}{suffix}",
868 std::process::id()
869 ))
870 }
871
872 fn action_alphabet(n: usize) -> ActionAlphabet {
873 ActionAlphabet::try_from_usize(n).expect("test action alphabet must be non-zero")
874 }
875
876 #[test]
877 fn extra_exploration_decay_does_not_wrap_after_i32_limit() {
878 let schedule = PlannerSchedule {
879 learn_cycles: 0,
880 eval_cycles: 0,
881 explore_epsilon: 0.5,
882 explore_gamma: 0.5,
883 };
884
885 assert_eq!(schedule.extra_exploration(i32::MAX as usize + 1), 0.0);
886 }
887
888 fn sample_warmstart_compiled_planner_run(teacher_path: &Path) -> CompiledPlannerRunSpec {
889 let document = SpecDocument::parse_json_value(
890 &json!({
891 "schema_version": 1,
892 "kind": "planner_run",
893 "assets": [{
894 "id": "teacher",
895 "path": teacher_path.to_string_lossy()
896 }],
897 "environment": {
898 "kind": "builtin",
899 "name": "coin_flip"
900 },
901 "interface": {
902 "observation_bits": 2,
903 "observation_stream_len": 1,
904 "observation_key_mode": "full_stream",
905 "reward_bits": 2,
906 "agent_actions": action_alphabet(2).get()
907 },
908 "controller": {
909 "kind": "aiqi_warmstart_exact_jh",
910 "predictor": {
911 "kind": "ctw",
912 "depth": 4
913 },
914 "return_horizon": 1,
915 "return_bins": 4,
916 "label_phase_period": 1,
917 "teacher_dataset_asset": "teacher",
918 "planner_simulations_per_step": 1
919 },
920 "runtime": {
921 "random_seed": 7,
922 "learn_cycles": 1,
923 "eval_cycles": 1,
924 "terminate_lifetime": 2,
925 "log_every": 1,
926 "perf": false,
927 "vm_perf_only": false,
928 "explore_epsilon": 0.0,
929 "explore_gamma": 1.0
930 }
931 }),
932 Path::new("."),
933 )
934 .expect("sample warmstart planner document");
935 let SpecDocument::PlannerRun(spec) = document else {
936 panic!("expected planner_run document");
937 };
938 spec.compile_in(&SpecEnvironment::new(Path::new(".")))
939 .expect("sample warmstart planner run should compile")
940 }
941
942 fn write_matching_warmstart_teacher(path: &Path, compiled: &CompiledPlannerRunSpec) {
943 let contract = standalone_warmstart_teacher_contract_for_compiled_planner_run(compiled)
944 .expect("standalone warmstart teacher contract");
945 let dataset = WarmStartExactJhTeacherDataset::new(
946 contract,
947 vec![WarmStartExactJhTeacherTrace::new(vec![
948 WarmStartExactJhTransition::new(0, vec![1], 1),
949 ])],
950 );
951 std::fs::write(
952 path,
953 serde_json::to_vec(&dataset.to_json_value()).expect("teacher JSON"),
954 )
955 .expect("write teacher dataset");
956 }
957
958 #[derive(Clone, Copy)]
959 struct CountingEnv {
960 observation: PerceptVal,
961 reward: Reward,
962 }
963
964 impl Environment for CountingEnv {
965 fn perform_action(&mut self, action: Action) {
966 self.observation = (self.observation + action + 1) & 0b11;
967 self.reward = (self.reward + 1).min(1);
968 }
969
970 fn get_observation(&self) -> PerceptVal {
971 self.observation
972 }
973
974 fn get_reward(&self) -> Reward {
975 self.reward
976 }
977
978 fn is_finished(&self) -> bool {
979 false
980 }
981
982 fn get_observation_bits(&self) -> usize {
983 2
984 }
985
986 fn get_reward_bits(&self) -> usize {
987 2
988 }
989
990 fn get_action_bits(&self) -> usize {
991 1
992 }
993 }
994
995 struct SeedRecordingEnv {
996 recorded_seed: Arc<AtomicU64>,
997 }
998
999 impl Environment for SeedRecordingEnv {
1000 fn perform_action(&mut self, _action: Action) {}
1001
1002 fn get_observation(&self) -> PerceptVal {
1003 0
1004 }
1005
1006 fn get_reward(&self) -> Reward {
1007 0
1008 }
1009
1010 fn is_finished(&self) -> bool {
1011 false
1012 }
1013
1014 fn get_observation_bits(&self) -> usize {
1015 2
1016 }
1017
1018 fn get_reward_bits(&self) -> usize {
1019 2
1020 }
1021
1022 fn get_action_bits(&self) -> usize {
1023 1
1024 }
1025
1026 fn set_random_seed(&mut self, seed: u64) {
1027 self.recorded_seed.store(seed, Ordering::SeqCst);
1028 }
1029 }
1030
1031 struct SeedRecordingFactory {
1032 recorded_seed: Arc<AtomicU64>,
1033 }
1034
1035 impl EnvironmentFactory for SeedRecordingFactory {
1036 fn build(&self) -> Result<Box<dyn Environment>, PlannerAgentError> {
1037 Ok(Box::new(SeedRecordingEnv {
1038 recorded_seed: Arc::clone(&self.recorded_seed),
1039 }))
1040 }
1041 }
1042
1043 struct SeedRecordingAgent {
1044 recorded_seed: Arc<AtomicU64>,
1045 reset_calls: Arc<AtomicU64>,
1046 }
1047
1048 impl PlannerAgent for SeedRecordingAgent {
1049 fn reseed_for_episode(&mut self, random_seed: u64) {
1050 self.recorded_seed.store(random_seed, Ordering::SeqCst);
1051 }
1052
1053 fn reset_for_episode(&mut self, random_seed: u64) {
1054 self.recorded_seed.store(random_seed, Ordering::SeqCst);
1055 self.reset_calls.fetch_add(1, Ordering::SeqCst);
1056 }
1057
1058 fn run_cycle(
1059 &mut self,
1060 phase: PlannerPhase,
1061 step: usize,
1062 _schedule: &PlannerSchedule,
1063 env: &mut PlannerEnvironment,
1064 observer: &mut dyn PlannerCycleObserver,
1065 ) -> Result<PlannerCycleOutcome, PlannerAgentError> {
1066 let pre_observations = env.observations().to_vec();
1067 let pre_reward = env.reward();
1068 observer.observe_action(step, 0, PlannerActionProvenance::Greedy)?;
1069 let reward = env.perform_action(0)?;
1070 observer.observe_percept(step, env.observations(), reward)?;
1071 observer.end_cycle(step)?;
1072 assert_eq!(phase, PlannerPhase::Learn);
1073 Ok(PlannerCycleOutcome {
1074 pre_observations,
1075 pre_reward,
1076 action: 0,
1077 observations: env.observations().to_vec(),
1078 reward,
1079 provenance: PlannerActionProvenance::Greedy,
1080 })
1081 }
1082 }
1083
1084 #[test]
1085 fn run_episode_uses_explicit_seed_for_environment_and_agent() {
1086 let compiled = sample_warmstart_compiled_planner_run(Path::new("teacher.json"));
1087 let env_seed = Arc::new(AtomicU64::new(u64::MAX));
1088 let agent_seed = Arc::new(AtomicU64::new(u64::MAX));
1089 let reset_calls = Arc::new(AtomicU64::new(0));
1090 let factory = SeedRecordingFactory {
1091 recorded_seed: Arc::clone(&env_seed),
1092 };
1093 let mut agent = SeedRecordingAgent {
1094 recorded_seed: Arc::clone(&agent_seed),
1095 reset_calls: Arc::clone(&reset_calls),
1096 };
1097
1098 let report = run_episode(
1099 &mut agent,
1100 &PlannerSchedule::new(1, 0),
1101 &factory,
1102 99,
1103 &compiled,
1104 )
1105 .expect("seeded episode should run");
1106
1107 assert_eq!(report.learn_cycles, 1);
1108 assert_eq!(env_seed.load(Ordering::SeqCst), 99);
1109 assert_eq!(agent_seed.load(Ordering::SeqCst), 99);
1110 assert_eq!(reset_calls.load(Ordering::SeqCst), 1);
1111 }
1112
1113 #[test]
1114 fn planner_run_session_executes_warmstart_controller_from_compiled_spec() {
1115 let teacher_path = unique_temp_path("planner-agent-warmstart-teacher", ".json");
1116 let compiled = sample_warmstart_compiled_planner_run(&teacher_path);
1117 write_matching_warmstart_teacher(&teacher_path, &compiled);
1118
1119 let controller =
1120 PlannerControllerAgent::from_compiled(&compiled).expect("warmstart controller");
1121 assert_eq!(controller.controller_kind(), "aiqi_warmstart_exact_jh");
1122 let env = Box::new(CountingEnv {
1123 observation: 1,
1124 reward: 0,
1125 });
1126 let mut session =
1127 PlannerRunSession::new(&compiled, controller, env).expect("planner session");
1128 assert_eq!(session.schedule().learn_cycles, 1);
1129 assert_eq!(session.schedule().eval_cycles, 1);
1130 assert_eq!(session.next_phase(), Some(PlannerPhase::Learn));
1131
1132 let mut observer = NullPlannerObserver;
1133 let learn = session
1134 .run_next_cycle(&mut observer)
1135 .expect("learn cycle")
1136 .expect("learn outcome");
1137 assert_eq!(learn.pre_observations, vec![1]);
1138 assert!(learn.action < 2);
1139 assert_eq!(session.next_phase(), Some(PlannerPhase::Eval));
1140
1141 let eval = session
1142 .run_next_cycle(&mut observer)
1143 .expect("eval cycle")
1144 .expect("eval outcome");
1145 assert!(eval.action < 2);
1146 assert!(session.is_finished());
1147 assert!(
1148 session
1149 .run_next_cycle(&mut observer)
1150 .expect("complete session")
1151 .is_none()
1152 );
1153
1154 let _ = std::fs::remove_file(teacher_path);
1155 }
1156}
1157
1158impl From<anyhow::Error> for PlannerAgentError {
1159 fn from(value: anyhow::Error) -> Self {
1160 Self::Environment {
1161 reason: value.to_string(),
1162 }
1163 }
1164}