1use crate::aixi::common::{
4 Action, ActionAlphabet, PerceptVal, RandomGenerator, Reward, RewardEncodingError,
5 resolve_random_seed, validate_reward_encoding_bounds,
6};
7use crate::aixi::model::{Predictor, PredictorBuildError, build_aiqi_predictor};
8use crate::aixi::planner_agent::PlannerActionProvenance;
9use crate::aixi::planner_spec::{PlannerInterfaceConfig, build_default_planner_run_spec};
10use crate::aixi::return_law::{
11 ReturnLabelCodec, ReturnLawEvaluator, ReturnPrefixUpdate, predict_expected_label,
12};
13use crate::aixi::warmstart_contract::{
14 TaskFingerprint, WARMSTART_STANDALONE_OBSERVATION_ADAPTER_SPEC_REF,
15 WARMSTART_STANDALONE_SCALAR_REPRESENTATION, WARMSTART_TEACHER_CONTRACT_SCHEMA_VERSION,
16 observation_key_mode_name, standalone_exact_reward_encoding_certificate_hash,
17 standalone_observation_adapter_content_crc32, warmstart_exact_jh_planner_task_fingerprint,
18};
19use crate::api::{BitStreamSemantics, RateBackend, validate_rate_backend};
20use crate::spec::{
21 AssetBinding, BuiltinEnvironmentSpec, CompiledPlannerController, CompiledPlannerRunSpec,
22 ControllerSpec, EnvironmentSpec, PlannerRunSpec, SpecError, WarmStartExactJhControllerSpec,
23};
24use serde_json::{Value, json};
25use std::cmp::Ordering;
26use std::collections::{BTreeMap, BTreeSet};
27use std::error::Error;
28use std::fmt;
29use std::fs::File;
30use std::io::{BufRead, BufReader};
31use std::num::NonZeroUsize;
32use std::path::Path;
33
34#[derive(Clone, Debug, Eq, PartialEq)]
36#[non_exhaustive]
37pub struct WarmStartExactJhTransition {
38 pub action: Action,
40 pub observations: Vec<PerceptVal>,
42 pub reward: Reward,
44}
45
46impl WarmStartExactJhTransition {
47 pub fn new(action: Action, observations: Vec<PerceptVal>, reward: Reward) -> Self {
49 Self {
50 action,
51 observations,
52 reward,
53 }
54 }
55}
56
57#[derive(Clone, Debug, Eq, PartialEq)]
59#[non_exhaustive]
60pub struct WarmStartExactJhTeacherTrace {
61 pub transitions: Vec<WarmStartExactJhTransition>,
63}
64
65impl WarmStartExactJhTeacherTrace {
66 pub fn new(transitions: Vec<WarmStartExactJhTransition>) -> Self {
68 Self { transitions }
69 }
70}
71
72pub fn validate_standalone_warmstart_provenance(
77 contract: &WarmStartExactJhTeacherContract,
78 observation_bits: usize,
79 observation_stream_len: usize,
80 reward_bits: usize,
81) -> Result<(), WarmStartExactJhError> {
82 if contract.observation_adapter_spec_ref != WARMSTART_STANDALONE_OBSERVATION_ADAPTER_SPEC_REF {
83 return Err(WarmStartExactJhError::InvalidTeacherDataset {
84 reason: format!(
85 "teacher observation_adapter_spec_ref '{}' does not match standalone direct-percept adapter declaration '{}'",
86 contract.observation_adapter_spec_ref,
87 WARMSTART_STANDALONE_OBSERVATION_ADAPTER_SPEC_REF
88 ),
89 });
90 }
91 let expected_adapter_crc = standalone_observation_adapter_content_crc32(
92 observation_bits,
93 observation_stream_len,
94 reward_bits,
95 )
96 .map_err(|err| WarmStartExactJhError::InvalidTeacherDataset {
97 reason: format!("failed to compute standalone observation adapter content hash: {err}"),
98 })?;
99 if contract.observation_adapter_content_crc32 != expected_adapter_crc {
100 return Err(WarmStartExactJhError::InvalidTeacherDataset {
101 reason: format!(
102 "teacher observation_adapter_content_crc32 '{}' does not match canonical standalone adapter spec '{}'",
103 contract.observation_adapter_content_crc32, expected_adapter_crc
104 ),
105 });
106 }
107 if contract.scalar_representation != WARMSTART_STANDALONE_SCALAR_REPRESENTATION {
108 return Err(WarmStartExactJhError::InvalidTeacherDataset {
109 reason: format!(
110 "teacher scalar_representation '{}' does not match standalone nonnegative integer declaration '{}'",
111 contract.scalar_representation, WARMSTART_STANDALONE_SCALAR_REPRESENTATION
112 ),
113 });
114 }
115 let expected_reward_cert = standalone_exact_reward_encoding_certificate_hash(reward_bits)
116 .map_err(|err| WarmStartExactJhError::InvalidTeacherDataset {
117 reason: format!(
118 "failed to compute standalone exact reward encoding certificate hash: {err}"
119 ),
120 })?;
121 if contract.exact_reward_encoding_certificate != expected_reward_cert {
122 return Err(WarmStartExactJhError::InvalidTeacherDataset {
123 reason: format!(
124 "teacher exact_reward_encoding_certificate '{}' does not match canonical standalone reward encoder certificate '{}'",
125 contract.exact_reward_encoding_certificate, expected_reward_cert
126 ),
127 });
128 }
129 Ok(())
130}
131
132pub(crate) struct WarmStartTeacherContractExpectation<'a> {
134 pub schema_version: u64,
136 pub task_fingerprint: TaskFingerprint,
138 pub action_alphabet_size: usize,
140 pub observation_bits: usize,
142 pub observation_stream_len: usize,
144 pub observation_key_mode: &'a str,
146 pub reward_bits: usize,
148 pub return_horizon: usize,
150 pub label_phase_period: usize,
152 pub validate_standalone_provenance: bool,
154}
155
156pub(crate) fn validate_warmstart_teacher_contract_against_expectation(
158 contract: &WarmStartExactJhTeacherContract,
159 expected: &WarmStartTeacherContractExpectation<'_>,
160) -> Result<(), WarmStartExactJhError> {
161 if contract.schema_version != expected.schema_version {
162 return Err(WarmStartExactJhError::InvalidTeacherDataset {
163 reason: format!("teacher schema_version must be {}", expected.schema_version),
164 });
165 }
166 if contract.task_fingerprint != expected.task_fingerprint {
167 return Err(WarmStartExactJhError::InvalidTeacherDataset {
168 reason: format!(
169 "teacher task_fingerprint '{}' does not match current planner_run '{}'",
170 contract.task_fingerprint, expected.task_fingerprint
171 ),
172 });
173 }
174 if contract.action_alphabet_size != expected.action_alphabet_size {
175 return Err(WarmStartExactJhError::InvalidTeacherDataset {
176 reason: format!(
177 "teacher action_alphabet_size {} does not match configured {}",
178 contract.action_alphabet_size, expected.action_alphabet_size
179 ),
180 });
181 }
182 if contract.observation_bits != expected.observation_bits {
183 return Err(WarmStartExactJhError::InvalidTeacherDataset {
184 reason: format!(
185 "teacher observation_bits {} does not match configured {}",
186 contract.observation_bits, expected.observation_bits
187 ),
188 });
189 }
190 if contract.observation_stream_len != expected.observation_stream_len {
191 return Err(WarmStartExactJhError::InvalidTeacherDataset {
192 reason: format!(
193 "teacher observation_stream_len {} does not match configured {}",
194 contract.observation_stream_len, expected.observation_stream_len
195 ),
196 });
197 }
198 if contract.observation_key_mode != expected.observation_key_mode {
199 return Err(WarmStartExactJhError::InvalidTeacherDataset {
200 reason: format!(
201 "teacher observation_key_mode '{}' does not match configured planner interface '{}'",
202 contract.observation_key_mode, expected.observation_key_mode
203 ),
204 });
205 }
206 if contract.reward_bits != expected.reward_bits {
207 return Err(WarmStartExactJhError::InvalidTeacherDataset {
208 reason: format!(
209 "teacher reward_bits {} does not match configured {}",
210 contract.reward_bits, expected.reward_bits
211 ),
212 });
213 }
214 if contract.return_horizon != expected.return_horizon {
215 return Err(WarmStartExactJhError::InvalidTeacherDataset {
216 reason: format!(
217 "teacher return_horizon {} does not match configured {}",
218 contract.return_horizon, expected.return_horizon
219 ),
220 });
221 }
222 if contract.label_phase_period != expected.label_phase_period {
223 return Err(WarmStartExactJhError::InvalidTeacherDataset {
224 reason: format!(
225 "teacher label_phase_period {} does not match configured {}",
226 contract.label_phase_period, expected.label_phase_period
227 ),
228 });
229 }
230 if expected.validate_standalone_provenance {
231 validate_standalone_warmstart_provenance(
232 contract,
233 expected.observation_bits,
234 expected.observation_stream_len,
235 expected.reward_bits,
236 )?;
237 }
238 Ok(())
239}
240
241pub fn validate_warmstart_teacher_planner_task_fingerprint(
251 compiled: &CompiledPlannerRunSpec,
252 contract: &WarmStartExactJhTeacherContract,
253) -> Result<(), WarmStartExactJhError> {
254 if contract.schema_version != WARMSTART_TEACHER_CONTRACT_SCHEMA_VERSION {
255 return Err(WarmStartExactJhError::InvalidTeacherDataset {
256 reason: format!(
257 "teacher schema_version must be {WARMSTART_TEACHER_CONTRACT_SCHEMA_VERSION}"
258 ),
259 });
260 }
261 let task_fingerprint =
262 warmstart_exact_jh_planner_task_fingerprint(compiled).map_err(|err| {
263 WarmStartExactJhError::InvalidTeacherDataset {
264 reason: format!("failed to compute planner task fingerprint: {err}"),
265 }
266 })?;
267 if contract.task_fingerprint != task_fingerprint {
268 return Err(WarmStartExactJhError::InvalidTeacherDataset {
269 reason: format!(
270 "teacher task_fingerprint '{}' does not match current planner_run '{}'",
271 contract.task_fingerprint, task_fingerprint
272 ),
273 });
274 }
275 Ok(())
276}
277
278pub fn validate_warmstart_teacher_against_compiled_planner_run(
282 compiled: &CompiledPlannerRunSpec,
283 contract: &WarmStartExactJhTeacherContract,
284) -> Result<(), WarmStartExactJhError> {
285 let interface = compiled.interface();
286 let (return_horizon, label_phase_period, planner_simulations_per_step) =
287 match compiled.controller() {
288 CompiledPlannerController::AiqiWarmstartExactJh {
289 return_horizon,
290 label_phase_period,
291 planner_simulations_per_step,
292 ..
293 } => (
294 *return_horizon,
295 *label_phase_period,
296 *planner_simulations_per_step,
297 ),
298 _ => {
299 return Err(WarmStartExactJhError::InvalidTeacherDataset {
300 reason:
301 "warm-start teacher contract can only be validated for aiqi_warmstart_exact_jh"
302 .to_string(),
303 });
304 }
305 };
306 let task_fingerprint =
307 warmstart_exact_jh_planner_task_fingerprint(compiled).map_err(|err| {
308 WarmStartExactJhError::InvalidTeacherDataset {
309 reason: format!("failed to compute planner task fingerprint: {err}"),
310 }
311 })?;
312 validate_warmstart_teacher_contract_against_expectation(
313 contract,
314 &WarmStartTeacherContractExpectation {
315 schema_version: WARMSTART_TEACHER_CONTRACT_SCHEMA_VERSION,
316 task_fingerprint,
317 action_alphabet_size: interface.agent_actions.get(),
318 observation_bits: interface.observation_bits,
319 observation_stream_len: interface.observation_stream_len.max(1),
320 observation_key_mode: observation_key_mode_name(interface.observation_key_mode),
321 reward_bits: interface.reward_bits,
322 return_horizon,
323 label_phase_period,
324 validate_standalone_provenance: true,
325 },
326 )?;
327 if planner_simulations_per_step == 0 {
328 return Err(WarmStartExactJhError::PlannerSimulationsZero);
329 }
330 if planner_simulations_per_step != 1 {
331 return Err(WarmStartExactJhError::PlannerSimulationsUnsupported {
332 configured: planner_simulations_per_step,
333 });
334 }
335 Ok(())
336}
337
338pub fn validate_warmstart_teacher_dataset_for_compiled_planner_run(
346 compiled: &CompiledPlannerRunSpec,
347 teacher: &WarmStartExactJhTeacherDataset,
348) -> Result<(), WarmStartExactJhError> {
349 let config = WarmStartExactJhRuntimeConfig::from_compiled(compiled)?;
350 config.validate_teacher_contract(&teacher.contract)?;
351 let predictor = match compiled.controller() {
352 CompiledPlannerController::AiqiWarmstartExactJh { predictor, .. } => predictor,
353 _ => return Err(WarmStartExactJhError::ControllerKindMismatch),
354 };
355 if !predictor.supports_frozen_conditioning() {
356 return Err(WarmStartExactJhError::UnsupportedRateBackend {
357 reason: "warm-start exact-J_H strict mode requires frozen context conditioning; configured rate_backend does not provide strict frozen conditioning",
358 });
359 }
360
361 let mut total_labels: usize = 0;
362 for (trace_index, trace) in teacher.traces.iter().enumerate() {
363 let trace_len = trace.transitions.len();
364 if trace_len < config.return_horizon {
365 return Err(WarmStartExactJhError::InvalidTeacherDataset {
366 reason: format!(
367 "traces[{trace_index}] contains {trace_len} transitions but return_horizon is {}",
368 config.return_horizon
369 ),
370 });
371 }
372 for (step_index, transition) in trace.transitions.iter().enumerate() {
373 validate_runtime_transition(
374 &config,
375 transition.action,
376 &transition.observations,
377 transition.reward,
378 )
379 .map_err(|err| WarmStartExactJhError::InvalidTeacherDataset {
380 reason: format!(
381 "traces[{trace_index}].transitions[{step_index}] violates runtime contract: {err}"
382 ),
383 })?;
384 }
385 let labels = exact_return_labels_for_trace(&config, &trace.transitions)?;
386 let label_count = labels.iter().filter(|label| label.is_some()).count();
387 if label_count == 0 {
388 return Err(WarmStartExactJhError::InvalidTeacherDataset {
389 reason: format!("traces[{trace_index}] did not contain any complete H-step labels"),
390 });
391 }
392 total_labels = total_labels.saturating_add(label_count);
393 }
394
395 if total_labels == 0 {
396 return Err(WarmStartExactJhError::InvalidTeacherDataset {
397 reason: "teacher dataset did not contain any complete H-step labels".to_string(),
398 });
399 }
400 Ok(())
401}
402
403pub fn validate_warmstart_teacher_transition_against_contract(
405 contract: &WarmStartExactJhTeacherContract,
406 action: Action,
407 observations: &[PerceptVal],
408 reward: Reward,
409) -> Result<(), WarmStartExactJhError> {
410 let agent_actions =
411 ActionAlphabet::try_from_usize(contract.action_alphabet_size).map_err(|_| {
412 WarmStartExactJhError::InvalidTeacherDataset {
413 reason: "teacher contract action_alphabet_size is zero".to_string(),
414 }
415 })?;
416 if action as usize >= contract.action_alphabet_size {
417 return Err(WarmStartExactJhError::ActionOutOfRange {
418 action,
419 agent_actions,
420 });
421 }
422 if observations.len() != contract.observation_stream_len {
423 return Err(WarmStartExactJhError::ObservationStreamLengthMismatch {
424 expected: contract.observation_stream_len,
425 actual: observations.len(),
426 });
427 }
428 let obs_max = max_value_for_bits(contract.observation_bits);
429 for &observation in observations {
430 if observation > obs_max {
431 return Err(WarmStartExactJhError::ObservationValueOutOfRange {
432 observation,
433 observation_bits: contract.observation_bits,
434 maximum: obs_max,
435 });
436 }
437 }
438 let max_reward = max_value_for_bits(contract.reward_bits) as i64;
439 if reward < 0 || reward > max_reward {
440 return Err(WarmStartExactJhError::RewardOutOfRange {
441 reward,
442 min_reward: 0,
443 max_reward,
444 });
445 }
446 Ok(())
447}
448
449#[derive(Clone, Debug, Eq, PartialEq)]
451#[non_exhaustive]
452pub struct WarmStartExactJhTeacherDataset {
453 pub contract: WarmStartExactJhTeacherContract,
455 pub traces: Vec<WarmStartExactJhTeacherTrace>,
459}
460
461impl WarmStartExactJhTeacherDataset {
462 pub fn new(
467 contract: WarmStartExactJhTeacherContract,
468 mut traces: Vec<WarmStartExactJhTeacherTrace>,
469 ) -> Self {
470 canonicalize_warmstart_teacher_traces(&mut traces);
471 Self { contract, traces }
472 }
473}
474
475#[derive(Clone, Debug, Eq, PartialEq)]
477#[non_exhaustive]
478pub struct WarmStartExactJhTeacherContract {
479 pub schema_version: u64,
481 pub task_fingerprint: TaskFingerprint,
483 pub action_alphabet_size: usize,
485 pub observation_bits: usize,
487 pub observation_stream_len: usize,
489 pub observation_key_mode: String,
491 pub observation_adapter_spec_ref: String,
493 pub observation_adapter_content_crc32: String,
495 pub reward_bits: usize,
497 pub return_horizon: usize,
499 pub label_phase_period: usize,
501 pub scalar_representation: String,
503 pub exact_reward_encoding_certificate: String,
505}
506
507impl WarmStartExactJhTeacherDataset {
508 pub fn from_json_slice(bytes: &[u8]) -> Result<Self, WarmStartExactJhError> {
510 let value = serde_json::from_slice::<Value>(bytes).map_err(|err| {
511 WarmStartExactJhError::InvalidTeacherDataset {
512 reason: format!("invalid teacher JSON: {err}"),
513 }
514 })?;
515 Self::from_json_value(&value)
516 }
517
518 pub fn from_json_value(value: &Value) -> Result<Self, WarmStartExactJhError> {
520 let object =
521 value
522 .as_object()
523 .ok_or_else(|| WarmStartExactJhError::InvalidTeacherDataset {
524 reason: "teacher dataset must be a JSON object".to_string(),
525 })?;
526 ensure_teacher_fields(
527 object,
528 &["schema_version", "contract", "traces"],
529 "teacher dataset",
530 )?;
531 let schema_version = object
532 .get("schema_version")
533 .and_then(Value::as_u64)
534 .ok_or_else(|| WarmStartExactJhError::InvalidTeacherDataset {
535 reason: format!(
536 "teacher dataset requires schema_version={WARMSTART_TEACHER_CONTRACT_SCHEMA_VERSION}"
537 ),
538 })?;
539 if schema_version != WARMSTART_TEACHER_CONTRACT_SCHEMA_VERSION {
540 return Err(WarmStartExactJhError::InvalidTeacherDataset {
541 reason: format!(
542 "teacher dataset schema_version must be {WARMSTART_TEACHER_CONTRACT_SCHEMA_VERSION}, got {schema_version}"
543 ),
544 });
545 }
546 let contract = parse_teacher_contract(object, schema_version)?;
547 let traces_value =
548 object
549 .get("traces")
550 .ok_or_else(|| WarmStartExactJhError::InvalidTeacherDataset {
551 reason: "teacher dataset requires a 'traces' array".to_string(),
552 })?;
553 let traces_array = traces_value.as_array().ok_or_else(|| {
554 WarmStartExactJhError::InvalidTeacherDataset {
555 reason: "teacher dataset field 'traces' must be an array".to_string(),
556 }
557 })?;
558 let mut traces = Vec::with_capacity(traces_array.len());
559 for (trace_index, trace_value) in traces_array.iter().enumerate() {
560 traces.push(parse_teacher_trace(trace_value, trace_index)?);
561 }
562 if traces.is_empty() {
563 return Err(WarmStartExactJhError::InvalidTeacherDataset {
564 reason: "teacher dataset must contain at least one trace".to_string(),
565 });
566 }
567 canonicalize_warmstart_teacher_traces(&mut traces);
568 Ok(Self { contract, traces })
569 }
570
571 pub fn to_json_value(&self) -> Value {
573 json!({
574 "schema_version": WARMSTART_TEACHER_CONTRACT_SCHEMA_VERSION,
575 "contract": teacher_contract_to_json_value(&self.contract),
576 "traces": self.traces.iter().map(teacher_trace_to_json_value).collect::<Vec<_>>(),
577 })
578 }
579
580 pub fn label_count_for_horizon(&self, return_horizon: usize) -> usize {
582 if return_horizon == 0 {
583 return 0;
584 }
585 self.traces
586 .iter()
587 .map(|trace| {
588 trace
589 .transitions
590 .len()
591 .saturating_add(1)
592 .saturating_sub(return_horizon)
593 })
594 .sum()
595 }
596}
597
598fn teacher_contract_to_json_value(contract: &WarmStartExactJhTeacherContract) -> Value {
599 json!({
600 "task_fingerprint": contract.task_fingerprint.to_string(),
601 "action_alphabet_size": contract.action_alphabet_size,
602 "observation_bits": contract.observation_bits,
603 "observation_stream_len": contract.observation_stream_len,
604 "observation_key_mode": contract.observation_key_mode,
605 "observation_adapter_spec_ref": contract.observation_adapter_spec_ref,
606 "observation_adapter_content_crc32": contract.observation_adapter_content_crc32,
607 "reward_bits": contract.reward_bits,
608 "return_horizon": contract.return_horizon,
609 "label_phase_period": contract.label_phase_period,
610 "scalar_representation": contract.scalar_representation,
611 "exact_reward_encoding_certificate": contract.exact_reward_encoding_certificate,
612 })
613}
614
615fn teacher_trace_to_json_value(trace: &WarmStartExactJhTeacherTrace) -> Value {
616 json!({
617 "transitions": trace.transitions.iter().map(|transition| {
618 json!({
619 "action": transition.action,
620 "observations": transition.observations,
621 "reward": transition.reward,
622 })
623 }).collect::<Vec<_>>(),
624 })
625}
626
627fn compare_warmstart_teacher_traces(
628 left: &WarmStartExactJhTeacherTrace,
629 right: &WarmStartExactJhTeacherTrace,
630) -> Ordering {
631 let mut left_iter = left.transitions.iter();
632 let mut right_iter = right.transitions.iter();
633 loop {
634 match (left_iter.next(), right_iter.next()) {
635 (Some(left_transition), Some(right_transition)) => {
636 let ordering = left_transition
637 .action
638 .cmp(&right_transition.action)
639 .then_with(|| {
640 left_transition
641 .observations
642 .as_slice()
643 .cmp(right_transition.observations.as_slice())
644 })
645 .then_with(|| left_transition.reward.cmp(&right_transition.reward));
646 if !ordering.is_eq() {
647 return ordering;
648 }
649 }
650 (None, Some(_)) => return Ordering::Less,
651 (Some(_), None) => return Ordering::Greater,
652 (None, None) => return Ordering::Equal,
653 }
654 }
655}
656
657fn canonicalize_warmstart_teacher_traces(traces: &mut Vec<WarmStartExactJhTeacherTrace>) {
658 traces.sort_by(compare_warmstart_teacher_traces);
659 traces.dedup_by(|right, left| compare_warmstart_teacher_traces(left, right).is_eq());
660}
661
662fn insert_warmstart_teacher_trace_canonical(
663 traces: &mut Vec<WarmStartExactJhTeacherTrace>,
664 trace: WarmStartExactJhTeacherTrace,
665) -> Option<usize> {
666 match traces.binary_search_by(|existing| compare_warmstart_teacher_traces(existing, &trace)) {
667 Ok(_) => None,
668 Err(index) => {
669 traces.insert(index, trace);
670 Some(index)
671 }
672 }
673}
674
675pub fn merge_warmstart_teacher_trace_deterministic(
677 traces: &mut Vec<WarmStartExactJhTeacherTrace>,
678 trace: WarmStartExactJhTeacherTrace,
679) -> bool {
680 canonicalize_warmstart_teacher_traces(traces);
681 insert_warmstart_teacher_trace_canonical(traces, trace).is_some()
682}
683
684pub fn merge_warmstart_teacher_traces_deterministic<I>(
689 traces: &mut Vec<WarmStartExactJhTeacherTrace>,
690 incoming: I,
691) -> (usize, Vec<WarmStartExactJhTeacherTrace>)
692where
693 I: IntoIterator<Item = WarmStartExactJhTeacherTrace>,
694{
695 canonicalize_warmstart_teacher_traces(traces);
696 let mut inserted = Vec::new();
697 for trace in incoming {
698 if let Some(index) = insert_warmstart_teacher_trace_canonical(traces, trace) {
699 inserted.push(traces[index].clone());
700 }
701 }
702 (inserted.len(), inserted)
703}
704
705#[derive(Clone, Debug, Default)]
707pub struct WarmStartExactJhTraceRecorder {
708 actions: BTreeMap<usize, Action>,
709 percepts: BTreeMap<usize, (Vec<PerceptVal>, Reward)>,
710}
711
712impl WarmStartExactJhTraceRecorder {
713 pub fn new() -> Self {
715 Self::default()
716 }
717
718 pub fn record_action(
720 &mut self,
721 step: usize,
722 action: Action,
723 ) -> Result<(), WarmStartExactJhError> {
724 if self.actions.insert(step, action).is_some() {
725 return Err(invalid_telemetry(format!(
726 "duplicate action record for step {step}"
727 )));
728 }
729 Ok(())
730 }
731
732 pub fn record_percept(
734 &mut self,
735 step: usize,
736 observations: &[PerceptVal],
737 reward: Reward,
738 ) -> Result<(), WarmStartExactJhError> {
739 if self
740 .percepts
741 .insert(step, (observations.to_vec(), reward))
742 .is_some()
743 {
744 return Err(invalid_telemetry(format!(
745 "duplicate percept record for step {step}"
746 )));
747 }
748 Ok(())
749 }
750
751 pub fn into_teacher_trace(
753 self,
754 contract: &WarmStartExactJhTeacherContract,
755 return_horizon: usize,
756 ) -> Result<WarmStartExactJhTeacherTrace, WarmStartExactJhError> {
757 if self.actions.len() != self.percepts.len() {
758 return Err(invalid_telemetry(
759 "action/percept record counts do not match",
760 ));
761 }
762 let action_steps = self.actions.keys().copied().collect::<BTreeSet<_>>();
763 let percept_steps = self.percepts.keys().copied().collect::<BTreeSet<_>>();
764 if action_steps != percept_steps {
765 return Err(invalid_telemetry("action/percept step sets do not match"));
766 }
767 ensure_dense_step_set(&action_steps, "recorded action/percept")?;
768 let mut transitions = Vec::with_capacity(self.actions.len());
769 for (step, action) in self.actions {
770 let (observations, reward) = self
771 .percepts
772 .get(&step)
773 .ok_or_else(|| invalid_telemetry(format!("missing percept for step {step}")))?;
774 validate_warmstart_teacher_transition_against_contract(
775 contract,
776 action,
777 observations,
778 *reward,
779 )?;
780 transitions.push(WarmStartExactJhTransition {
781 action,
782 observations: observations.clone(),
783 reward: *reward,
784 });
785 }
786 if transitions.len() < return_horizon {
787 return Err(invalid_telemetry(format!(
788 "trace contains {} transitions but return_horizon is {return_horizon}",
789 transitions.len()
790 )));
791 }
792 Ok(WarmStartExactJhTeacherTrace { transitions })
793 }
794}
795
796pub fn warmstart_jsonl_action_record(
798 step: usize,
799 action: Action,
800 provenance: PlannerActionProvenance,
801) -> Value {
802 json!({
803 "kind": "action",
804 "t": step,
805 "action": action,
806 "provenance": provenance.as_str(),
807 })
808}
809
810pub fn warmstart_jsonl_percept_record(
812 step: usize,
813 observations: &[PerceptVal],
814 reward: Reward,
815) -> Value {
816 json!({
817 "kind": "percept",
818 "t": step,
819 "observations": observations,
820 "reward": reward,
821 })
822}
823
824fn invalid_telemetry(reason: impl Into<String>) -> WarmStartExactJhError {
825 WarmStartExactJhError::InvalidTelemetry {
826 reason: reason.into(),
827 }
828}
829
830fn ensure_dense_step_set(
831 steps: &BTreeSet<usize>,
832 label: &str,
833) -> Result<(), WarmStartExactJhError> {
834 let Some(&first) = steps.first() else {
835 return Ok(());
836 };
837 let mut previous = first;
838 for &step in steps.iter().skip(1) {
839 let expected = previous.checked_add(1).ok_or_else(|| {
840 invalid_telemetry(format!(
841 "{label} step set cannot be dense after usize::MAX step {previous}"
842 ))
843 })?;
844 if step != expected {
845 return Err(invalid_telemetry(format!(
846 "{label} step set is not contiguous: expected step {expected} before step {step}"
847 )));
848 }
849 previous = step;
850 }
851 Ok(())
852}
853
854fn ensure_jsonl_fields(
855 object: &serde_json::Map<String, Value>,
856 allowed: &[&str],
857 line_number: usize,
858) -> Result<(), WarmStartExactJhError> {
859 for key in object.keys() {
860 if !allowed.contains(&key.as_str()) {
861 return Err(invalid_telemetry(format!(
862 "line {line_number}: unknown field '{key}'"
863 )));
864 }
865 }
866 Ok(())
867}
868
869fn jsonl_required_u64(
870 value: &Value,
871 field: &str,
872 line_number: usize,
873) -> Result<u64, WarmStartExactJhError> {
874 value.get(field).and_then(Value::as_u64).ok_or_else(|| {
875 invalid_telemetry(format!("line {line_number}: missing u64 field '{field}'"))
876 })
877}
878
879fn jsonl_required_i64(
880 value: &Value,
881 field: &str,
882 line_number: usize,
883) -> Result<i64, WarmStartExactJhError> {
884 value.get(field).and_then(Value::as_i64).ok_or_else(|| {
885 invalid_telemetry(format!("line {line_number}: missing i64 field '{field}'"))
886 })
887}
888
889fn jsonl_required_observations(
890 value: &Value,
891 line_number: usize,
892) -> Result<Vec<PerceptVal>, WarmStartExactJhError> {
893 let observations = value
894 .get("observations")
895 .and_then(Value::as_array)
896 .ok_or_else(|| {
897 invalid_telemetry(format!("line {line_number}: missing observations array"))
898 })?;
899 observations
900 .iter()
901 .map(|item| {
902 item.as_u64().ok_or_else(|| {
903 invalid_telemetry(format!("line {line_number}: observation must be a u64"))
904 })
905 })
906 .collect()
907}
908
909#[derive(Clone, Debug)]
910struct JsonlActionRecord {
911 action: Action,
912 line_number: usize,
913}
914
915#[derive(Clone, Debug)]
916struct JsonlPerceptRecord {
917 observations: Vec<PerceptVal>,
918 reward: Reward,
919 line_number: usize,
920}
921
922#[derive(Clone, Debug, Default)]
923struct JsonlStepRecords {
924 action: Option<JsonlActionRecord>,
925 percept: Option<JsonlPerceptRecord>,
926}
927
928#[derive(Clone, Copy, Debug, Eq, PartialEq)]
929enum JsonlTraceConvention {
930 ActionThenPostPercept,
931 DecisionPerceptThenAction,
932}
933
934fn infer_jsonl_trace_convention(
935 steps: &BTreeMap<usize, JsonlStepRecords>,
936) -> Result<JsonlTraceConvention, WarmStartExactJhError> {
937 let mut saw_action_then_percept = false;
938 let mut saw_percept_then_action = false;
939 for records in steps.values() {
940 let (Some(action), Some(percept)) = (&records.action, &records.percept) else {
941 continue;
942 };
943 match action.line_number.cmp(&percept.line_number) {
944 std::cmp::Ordering::Less => saw_action_then_percept = true,
945 std::cmp::Ordering::Greater => saw_percept_then_action = true,
946 std::cmp::Ordering::Equal => {
947 return Err(invalid_telemetry(
948 "action and percept records cannot originate from the same JSONL line",
949 ));
950 }
951 }
952 }
953 match (saw_action_then_percept, saw_percept_then_action) {
954 (true, false) => Ok(JsonlTraceConvention::ActionThenPostPercept),
955 (false, true) => Ok(JsonlTraceConvention::DecisionPerceptThenAction),
956 (true, true) => Err(invalid_telemetry(
957 "mixed JSONL action/percept conventions in one trace",
958 )),
959 (false, false) => Err(invalid_telemetry(
960 "cannot infer JSONL action/percept convention",
961 )),
962 }
963}
964
965fn push_validated_jsonl_transition(
966 transitions: &mut Vec<WarmStartExactJhTransition>,
967 contract: &WarmStartExactJhTeacherContract,
968 action_step: usize,
969 action: Action,
970 percept: &JsonlPerceptRecord,
971) -> Result<(), WarmStartExactJhError> {
972 validate_warmstart_teacher_transition_against_contract(
973 contract,
974 action,
975 &percept.observations,
976 percept.reward,
977 )
978 .map_err(|err| {
979 invalid_telemetry(format!(
980 "action step {action_step} paired with percept line {} violates contract: {err}",
981 percept.line_number
982 ))
983 })?;
984 transitions.push(WarmStartExactJhTransition {
985 action,
986 observations: percept.observations.clone(),
987 reward: percept.reward,
988 });
989 Ok(())
990}
991
992fn jsonl_steps_into_teacher_trace(
993 steps: BTreeMap<usize, JsonlStepRecords>,
994 contract: &WarmStartExactJhTeacherContract,
995 return_horizon: usize,
996) -> Result<WarmStartExactJhTeacherTrace, WarmStartExactJhError> {
997 let convention = infer_jsonl_trace_convention(&steps)?;
998 let mut transitions = Vec::new();
999 match convention {
1000 JsonlTraceConvention::ActionThenPostPercept => {
1001 let all_steps = steps.keys().copied().collect::<BTreeSet<_>>();
1002 ensure_dense_step_set(&all_steps, "action-then-percept JSONL")?;
1003 for (step, records) in &steps {
1004 let action = records.action.as_ref().ok_or_else(|| {
1005 invalid_telemetry(format!("missing action record for step {step}"))
1006 })?;
1007 let percept = records.percept.as_ref().ok_or_else(|| {
1008 invalid_telemetry(format!("missing percept record for step {step}"))
1009 })?;
1010 push_validated_jsonl_transition(
1011 &mut transitions,
1012 contract,
1013 *step,
1014 action.action,
1015 percept,
1016 )?;
1017 }
1018 }
1019 JsonlTraceConvention::DecisionPerceptThenAction => {
1020 let action_steps = steps
1021 .iter()
1022 .filter_map(|(step, records)| records.action.as_ref().map(|_| *step))
1023 .collect::<BTreeSet<_>>();
1024 let Some(&first_action_step) = action_steps.first() else {
1025 return Err(invalid_telemetry("trace contains no action records"));
1026 };
1027 ensure_dense_step_set(&action_steps, "decision-percept JSONL action")?;
1028 let max_action_step = *action_steps
1029 .last()
1030 .ok_or_else(|| invalid_telemetry("trace contains no action records"))?;
1031 let final_percept_step = max_action_step.checked_add(1).ok_or_else(|| {
1032 invalid_telemetry(format!(
1033 "decision-percept JSONL action step {max_action_step} cannot have a successor"
1034 ))
1035 })?;
1036 for step in steps.keys() {
1037 if *step < first_action_step || *step > final_percept_step {
1038 return Err(invalid_telemetry(format!(
1039 "unexpected JSONL step {step} outside dense decision trace domain {first_action_step}..={final_percept_step}"
1040 )));
1041 }
1042 if !action_steps.contains(step) && *step != final_percept_step {
1043 return Err(invalid_telemetry(format!(
1044 "unexpected percept-only JSONL step {step} inside dense decision trace"
1045 )));
1046 }
1047 }
1048 for step in &action_steps {
1049 let records = steps
1050 .get(step)
1051 .expect("action_steps contains only keys present in steps");
1052 let action = records
1053 .action
1054 .as_ref()
1055 .expect("action_steps contains only records with actions");
1056 if records.percept.is_none() {
1057 return Err(invalid_telemetry(format!(
1058 "missing decision percept record for action step {step}"
1059 )));
1060 }
1061 let next_step = step.checked_add(1).ok_or_else(|| {
1062 invalid_telemetry(format!(
1063 "decision-percept JSONL action step {step} cannot have a successor"
1064 ))
1065 })?;
1066 let next_records = steps.get(&next_step).ok_or_else(|| {
1067 invalid_telemetry(format!(
1068 "missing post-action percept for action step {step}"
1069 ))
1070 })?;
1071 let Some(percept) = next_records.percept.as_ref() else {
1072 return Err(invalid_telemetry(format!(
1073 "missing post-action percept for action step {step}"
1074 )));
1075 };
1076 push_validated_jsonl_transition(
1077 &mut transitions,
1078 contract,
1079 *step,
1080 action.action,
1081 percept,
1082 )?;
1083 }
1084 }
1085 }
1086 if transitions.len() < return_horizon {
1087 return Err(invalid_telemetry(format!(
1088 "trace contains {} complete transitions but return_horizon is {return_horizon}",
1089 transitions.len()
1090 )));
1091 }
1092 Ok(WarmStartExactJhTeacherTrace { transitions })
1093}
1094
1095pub fn warmstart_teacher_trace_from_jsonl_reader<R: BufRead>(
1097 reader: R,
1098 contract: &WarmStartExactJhTeacherContract,
1099 return_horizon: usize,
1100) -> Result<WarmStartExactJhTeacherTrace, WarmStartExactJhError> {
1101 let mut steps = BTreeMap::<usize, JsonlStepRecords>::new();
1102 for (line_index, line) in reader.lines().enumerate() {
1103 let line_number = line_index + 1;
1104 let line = line.map_err(|err| invalid_telemetry(format!("line {line_number}: {err}")))?;
1105 if line.trim().is_empty() {
1106 return Err(invalid_telemetry(format!(
1107 "line {line_number}: empty JSONL records are not allowed"
1108 )));
1109 }
1110 let value = serde_json::from_str::<Value>(&line)
1111 .map_err(|err| invalid_telemetry(format!("line {line_number}: invalid JSON: {err}")))?;
1112 let object = value.as_object().ok_or_else(|| {
1113 invalid_telemetry(format!("line {line_number}: record must be an object"))
1114 })?;
1115 let kind = value
1116 .get("kind")
1117 .and_then(Value::as_str)
1118 .ok_or_else(|| invalid_telemetry(format!("line {line_number}: missing kind")))?;
1119 let step = usize::try_from(jsonl_required_u64(&value, "t", line_number)?)
1120 .map_err(|_| invalid_telemetry(format!("line {line_number}: t exceeds usize")))?;
1121 match kind {
1122 "action" => {
1123 ensure_jsonl_fields(object, &["kind", "t", "action", "provenance"], line_number)?;
1124 if let Some(provenance_value) = value.get("provenance") {
1125 let provenance = provenance_value.as_str().ok_or_else(|| {
1126 invalid_telemetry(format!(
1127 "line {line_number}: action provenance must be a string"
1128 ))
1129 })?;
1130 PlannerActionProvenance::from_jsonl_str(provenance)?;
1131 }
1132 let action = jsonl_required_u64(&value, "action", line_number)?;
1133 let entry = steps.entry(step).or_default();
1134 if entry
1135 .action
1136 .replace(JsonlActionRecord {
1137 action,
1138 line_number,
1139 })
1140 .is_some()
1141 {
1142 return Err(invalid_telemetry(format!(
1143 "duplicate action record for step {step}"
1144 )));
1145 }
1146 }
1147 "percept" => {
1148 ensure_jsonl_fields(
1149 object,
1150 &["kind", "t", "observations", "reward"],
1151 line_number,
1152 )?;
1153 let observations = jsonl_required_observations(&value, line_number)?;
1154 let reward = jsonl_required_i64(&value, "reward", line_number)?;
1155 let entry = steps.entry(step).or_default();
1156 if entry
1157 .percept
1158 .replace(JsonlPerceptRecord {
1159 observations,
1160 reward,
1161 line_number,
1162 })
1163 .is_some()
1164 {
1165 return Err(invalid_telemetry(format!(
1166 "duplicate percept record for step {step}"
1167 )));
1168 }
1169 }
1170 other => {
1171 return Err(invalid_telemetry(format!(
1172 "line {line_number}: unknown record kind '{other}'"
1173 )));
1174 }
1175 }
1176 }
1177 jsonl_steps_into_teacher_trace(steps, contract, return_horizon)
1178}
1179
1180pub fn warmstart_teacher_trace_from_jsonl_slice(
1182 bytes: &[u8],
1183 contract: &WarmStartExactJhTeacherContract,
1184 return_horizon: usize,
1185) -> Result<WarmStartExactJhTeacherTrace, WarmStartExactJhError> {
1186 warmstart_teacher_trace_from_jsonl_reader(BufReader::new(bytes), contract, return_horizon)
1187}
1188
1189pub fn warmstart_teacher_trace_from_jsonl_path(
1191 path: impl AsRef<Path>,
1192 contract: &WarmStartExactJhTeacherContract,
1193 return_horizon: usize,
1194) -> Result<WarmStartExactJhTeacherTrace, WarmStartExactJhError> {
1195 let file = File::open(path.as_ref()).map_err(|err| {
1196 invalid_telemetry(format!(
1197 "failed to open JSONL trace '{}': {err}",
1198 path.as_ref().display()
1199 ))
1200 })?;
1201 warmstart_teacher_trace_from_jsonl_reader(BufReader::new(file), contract, return_horizon)
1202}
1203
1204pub fn standalone_warmstart_teacher_contract_for_compiled_planner_run(
1206 compiled: &CompiledPlannerRunSpec,
1207) -> Result<WarmStartExactJhTeacherContract, WarmStartExactJhError> {
1208 let interface = compiled.interface();
1209 let (return_horizon, label_phase_period, planner_simulations_per_step) =
1210 match compiled.controller() {
1211 CompiledPlannerController::AiqiWarmstartExactJh {
1212 return_horizon,
1213 label_phase_period,
1214 planner_simulations_per_step,
1215 ..
1216 } => (
1217 *return_horizon,
1218 *label_phase_period,
1219 *planner_simulations_per_step,
1220 ),
1221 _ => return Err(WarmStartExactJhError::ControllerKindMismatch),
1222 };
1223 if planner_simulations_per_step == 0 {
1224 return Err(WarmStartExactJhError::PlannerSimulationsZero);
1225 }
1226 if planner_simulations_per_step != 1 {
1227 return Err(WarmStartExactJhError::PlannerSimulationsUnsupported {
1228 configured: planner_simulations_per_step,
1229 });
1230 }
1231 let task_fingerprint =
1232 warmstart_exact_jh_planner_task_fingerprint(compiled).map_err(|err| {
1233 WarmStartExactJhError::InvalidTeacherDataset {
1234 reason: format!("failed to compute planner task fingerprint: {err}"),
1235 }
1236 })?;
1237 let (observation_adapter_content_crc32, exact_reward_encoding_certificate) =
1238 crate::aixi::warmstart_contract::standalone_teacher_provenance_crc32_pair(
1239 interface.observation_bits,
1240 interface.observation_stream_len.max(1),
1241 interface.reward_bits,
1242 )
1243 .map_err(|err| WarmStartExactJhError::InvalidTeacherDataset {
1244 reason: format!("failed to compute standalone teacher provenance: {err}"),
1245 })?;
1246 Ok(WarmStartExactJhTeacherContract {
1247 schema_version: WARMSTART_TEACHER_CONTRACT_SCHEMA_VERSION,
1248 task_fingerprint,
1249 action_alphabet_size: interface.agent_actions.get(),
1250 observation_bits: interface.observation_bits,
1251 observation_stream_len: interface.observation_stream_len.max(1),
1252 observation_key_mode: observation_key_mode_name(interface.observation_key_mode).to_string(),
1253 observation_adapter_spec_ref: WARMSTART_STANDALONE_OBSERVATION_ADAPTER_SPEC_REF.to_string(),
1254 observation_adapter_content_crc32,
1255 reward_bits: interface.reward_bits,
1256 return_horizon,
1257 label_phase_period,
1258 scalar_representation: WARMSTART_STANDALONE_SCALAR_REPRESENTATION.to_string(),
1259 exact_reward_encoding_certificate,
1260 })
1261}
1262
1263pub fn warmstart_target_return_horizon(
1265 compiled: &CompiledPlannerRunSpec,
1266) -> Result<usize, WarmStartExactJhError> {
1267 match compiled.controller() {
1268 CompiledPlannerController::AiqiWarmstartExactJh { return_horizon, .. } => {
1269 Ok(*return_horizon)
1270 }
1271 _ => Err(WarmStartExactJhError::ControllerKindMismatch),
1272 }
1273}
1274
1275pub fn read_warmstart_teacher_dataset_path(
1277 path: impl AsRef<Path>,
1278) -> Result<WarmStartExactJhTeacherDataset, WarmStartExactJhError> {
1279 let path_ref = path.as_ref();
1280 let bytes =
1281 std::fs::read(path_ref).map_err(|err| WarmStartExactJhError::InvalidTeacherDataset {
1282 reason: format!(
1283 "failed to read teacher dataset '{}': {err}",
1284 path_ref.display()
1285 ),
1286 })?;
1287 WarmStartExactJhTeacherDataset::from_json_slice(&bytes)
1288}
1289
1290pub fn write_warmstart_teacher_dataset_path(
1292 path: impl AsRef<Path>,
1293 dataset: &WarmStartExactJhTeacherDataset,
1294) -> Result<(), WarmStartExactJhError> {
1295 let path_ref = path.as_ref();
1296 let bytes = serde_json::to_vec_pretty(&dataset.to_json_value()).map_err(|err| {
1297 WarmStartExactJhError::InvalidTeacherDataset {
1298 reason: format!("failed to serialize teacher dataset: {err}"),
1299 }
1300 })?;
1301 std::fs::write(path_ref, bytes).map_err(|err| WarmStartExactJhError::InvalidTeacherDataset {
1302 reason: format!(
1303 "failed to write teacher dataset '{}': {err}",
1304 path_ref.display()
1305 ),
1306 })
1307}
1308
1309#[derive(Clone)]
1311#[non_exhaustive]
1312pub struct WarmStartExactJhConfig {
1313 pub rate_backend: RateBackend,
1315 pub observation_bits: usize,
1317 pub observation_stream_len: usize,
1319 pub reward_bits: usize,
1321 pub agent_actions: ActionAlphabet,
1323 pub return_horizon: usize,
1325 pub return_bins: usize,
1331 pub label_phase_period: usize,
1333 pub planner_simulations_per_step: usize,
1339 pub bit_stream_semantics: BitStreamSemantics,
1341 pub random_seed: Option<u64>,
1343}
1344
1345impl Default for WarmStartExactJhConfig {
1346 fn default() -> Self {
1347 Self {
1348 rate_backend: RateBackend::Ctw { depth: 8 },
1349 observation_bits: 1,
1350 observation_stream_len: 1,
1351 reward_bits: 1,
1352 agent_actions: ActionAlphabet::try_from_usize(2)
1353 .expect("default action alphabet must be non-zero"),
1354 return_horizon: 1,
1355 return_bins: 2,
1356 label_phase_period: 1,
1357 planner_simulations_per_step: 1,
1358 bit_stream_semantics: BitStreamSemantics::BinaryTokens,
1359 random_seed: None,
1360 }
1361 }
1362}
1363
1364impl WarmStartExactJhConfig {
1365 fn canonical_planner_run_spec(&self) -> PlannerRunSpec {
1366 let mut spec = build_default_planner_run_spec(
1367 PlannerInterfaceConfig {
1368 observation_bits: self.observation_bits,
1369 observation_stream_len: self.observation_stream_len,
1370 observation_key_mode: crate::aixi::common::ObservationKeyMode::FullStream,
1371 reward_bits: self.reward_bits,
1372 agent_actions: self.agent_actions,
1373 },
1374 ControllerSpec::AiqiWarmstartExactJh(WarmStartExactJhControllerSpec {
1375 predictor: self.rate_backend.clone(),
1376 bit_stream_semantics: self.bit_stream_semantics,
1377 return_horizon: self.return_horizon,
1378 return_bins: self.return_bins,
1379 label_phase_period: self.label_phase_period,
1380 teacher_dataset_asset: "programmatic_warmstart_teacher".to_string(),
1381 planner_simulations_per_step: self.planner_simulations_per_step,
1382 }),
1383 self.random_seed,
1384 );
1385 spec.assets.push(AssetBinding {
1386 id: "programmatic_warmstart_teacher".to_string(),
1387 path: "programmatic_warmstart_teacher.json".to_string(),
1388 });
1389 spec
1390 }
1391
1392 fn compile_planner_run_spec(&self) -> Result<CompiledPlannerRunSpec, WarmStartExactJhError> {
1393 self.canonical_planner_run_spec()
1394 .compile()
1395 .map_err(WarmStartExactJhError::from)
1396 }
1397
1398 fn validate_runtime_invariants(&self) -> Result<(), WarmStartExactJhError> {
1399 if self.return_horizon == 0 {
1400 return Err(WarmStartExactJhError::ReturnHorizonZero);
1401 }
1402 if self.return_bins == 0 {
1403 return Err(WarmStartExactJhError::ReturnBinsZero);
1404 }
1405 if self.label_phase_period < self.return_horizon {
1406 return Err(WarmStartExactJhError::LabelPhasePeriodTooShort {
1407 label_phase_period: self.label_phase_period,
1408 return_horizon: self.return_horizon,
1409 });
1410 }
1411 if self.planner_simulations_per_step == 0 {
1412 return Err(WarmStartExactJhError::PlannerSimulationsZero);
1413 }
1414 if self.planner_simulations_per_step != 1 {
1415 return Err(WarmStartExactJhError::PlannerSimulationsUnsupported {
1416 configured: self.planner_simulations_per_step,
1417 });
1418 }
1419 let return_horizon = NonZeroUsize::new(self.return_horizon)
1420 .ok_or(WarmStartExactJhError::ReturnHorizonZero)?;
1421 let return_bins =
1422 NonZeroUsize::new(self.return_bins).ok_or(WarmStartExactJhError::ReturnBinsZero)?;
1423 let (min_reward, max_reward, _reward_offset) =
1424 reward_bounds_from_exact_return_bins(return_horizon, return_bins, self.reward_bits)?;
1425 validate_exact_return_alphabet(
1426 min_reward,
1427 max_reward,
1428 self.return_horizon,
1429 self.return_bins,
1430 )?;
1431 validate_rate_backend(&self.rate_backend)
1432 .map_err(WarmStartExactJhError::InvalidRateBackend)?;
1433 let compiled = self
1434 .rate_backend
1435 .compile()
1436 .map_err(WarmStartExactJhError::Spec)?;
1437 if !compiled.supports_frozen_conditioning() {
1438 return Err(WarmStartExactJhError::UnsupportedRateBackend {
1439 reason: "warm-start exact-J_H strict mode requires frozen context conditioning; configured rate_backend does not provide strict frozen conditioning",
1440 });
1441 }
1442 Ok(())
1443 }
1444
1445 pub fn validate(&self) -> Result<(), WarmStartExactJhError> {
1447 self.validate_runtime_invariants()?;
1448 self.compile_planner_run_spec().map(|_| ())
1449 }
1450}
1451
1452#[derive(Clone)]
1453struct WarmStartExactJhRuntimeConfig {
1454 task_fingerprint: TaskFingerprint,
1455 observation_bits: usize,
1456 observation_stream_len: usize,
1457 observation_key_mode: &'static str,
1458 reward_bits: usize,
1459 agent_actions: ActionAlphabet,
1460 min_reward: Reward,
1461 max_reward: Reward,
1462 reward_offset: Reward,
1463 return_horizon: usize,
1464 return_bins: usize,
1465 label_phase_period: usize,
1466 planner_simulations_per_step: usize,
1467 random_seed: u64,
1468 provenance_policy: TeacherProvenancePolicy,
1469}
1470
1471#[derive(Clone, Copy)]
1472enum TeacherProvenancePolicy {
1473 StandalonePlannerRun,
1474 ExternallyValidatedTunerBridge,
1475}
1476
1477impl WarmStartExactJhRuntimeConfig {
1478 fn from_compiled(compiled: &CompiledPlannerRunSpec) -> Result<Self, WarmStartExactJhError> {
1485 let planner = compiled.canonical_spec();
1486 let interface = compiled.interface();
1487 let runtime = compiled.runtime();
1488 let (return_horizon, return_bins, label_phase_period, planner_simulations_per_step) =
1489 match compiled.controller() {
1490 CompiledPlannerController::AiqiWarmstartExactJh {
1491 return_horizon,
1492 return_bins,
1493 label_phase_period,
1494 planner_simulations_per_step,
1495 ..
1496 } => (
1497 *return_horizon,
1498 *return_bins,
1499 *label_phase_period,
1500 *planner_simulations_per_step,
1501 ),
1502 _ => return Err(WarmStartExactJhError::ControllerKindMismatch),
1503 };
1504 if return_horizon == 0 {
1505 return Err(WarmStartExactJhError::ReturnHorizonZero);
1506 }
1507 if return_bins == 0 {
1508 return Err(WarmStartExactJhError::ReturnBinsZero);
1509 }
1510 if label_phase_period < return_horizon {
1511 return Err(WarmStartExactJhError::LabelPhasePeriodTooShort {
1512 label_phase_period,
1513 return_horizon,
1514 });
1515 }
1516 if planner_simulations_per_step == 0 {
1517 return Err(WarmStartExactJhError::PlannerSimulationsZero);
1518 }
1519 if planner_simulations_per_step != 1 {
1520 return Err(WarmStartExactJhError::PlannerSimulationsUnsupported {
1521 configured: planner_simulations_per_step,
1522 });
1523 }
1524 let nonzero_return_horizon =
1525 NonZeroUsize::new(return_horizon).ok_or(WarmStartExactJhError::ReturnHorizonZero)?;
1526 let nonzero_return_bins =
1527 NonZeroUsize::new(return_bins).ok_or(WarmStartExactJhError::ReturnBinsZero)?;
1528 let (min_reward, max_reward, reward_offset) = reward_bounds_from_exact_return_bins(
1529 nonzero_return_horizon,
1530 nonzero_return_bins,
1531 interface.reward_bits,
1532 )?;
1533 validate_exact_return_alphabet(min_reward, max_reward, return_horizon, return_bins)?;
1534 let task_fingerprint =
1535 warmstart_exact_jh_planner_task_fingerprint(compiled).map_err(|err| {
1536 WarmStartExactJhError::InvalidTeacherDataset {
1537 reason: format!("failed to compute planner task fingerprint: {err}"),
1538 }
1539 })?;
1540 let provenance_policy = match &planner.environment {
1541 EnvironmentSpec::Builtin {
1542 builtin: BuiltinEnvironmentSpec::TunerBridge,
1543 } => TeacherProvenancePolicy::ExternallyValidatedTunerBridge,
1544 _ => TeacherProvenancePolicy::StandalonePlannerRun,
1545 };
1546 Ok(Self {
1547 task_fingerprint,
1548 observation_bits: interface.observation_bits,
1549 observation_stream_len: interface.observation_stream_len.max(1),
1550 observation_key_mode: observation_key_mode_name(interface.observation_key_mode),
1551 reward_bits: interface.reward_bits,
1552 agent_actions: interface.agent_actions,
1553 min_reward,
1554 max_reward,
1555 reward_offset,
1556 return_horizon,
1557 return_bins,
1558 label_phase_period,
1559 planner_simulations_per_step,
1560 random_seed: resolve_random_seed(runtime.random_seed),
1561 provenance_policy,
1562 })
1563 }
1564
1565 fn validate_teacher_contract(
1571 &self,
1572 contract: &WarmStartExactJhTeacherContract,
1573 ) -> Result<(), WarmStartExactJhError> {
1574 validate_warmstart_teacher_contract_against_expectation(
1575 contract,
1576 &WarmStartTeacherContractExpectation {
1577 schema_version: WARMSTART_TEACHER_CONTRACT_SCHEMA_VERSION,
1578 task_fingerprint: self.task_fingerprint,
1579 action_alphabet_size: self.agent_actions.get(),
1580 observation_bits: self.observation_bits,
1581 observation_stream_len: self.observation_stream_len,
1582 observation_key_mode: self.observation_key_mode,
1583 reward_bits: self.reward_bits,
1584 return_horizon: self.return_horizon,
1585 label_phase_period: self.label_phase_period,
1586 validate_standalone_provenance: matches!(
1587 self.provenance_policy,
1588 TeacherProvenancePolicy::StandalonePlannerRun
1589 ),
1590 },
1591 )
1592 }
1593}
1594
1595pub(crate) fn reward_bounds_from_exact_return_bins(
1596 return_horizon: NonZeroUsize,
1597 return_bins: NonZeroUsize,
1598 reward_bits: usize,
1599) -> Result<(Reward, Reward, Reward), WarmStartExactJhError> {
1600 let max_reward = max_reward_from_exact_return_bins(return_horizon, return_bins)?;
1601 validate_reward_encoding_bounds(0, max_reward, 0, reward_bits)?;
1602 Ok((0, max_reward, 0))
1603}
1604
1605pub(crate) fn max_reward_from_exact_return_bins(
1606 return_horizon: NonZeroUsize,
1607 return_bins: NonZeroUsize,
1608) -> Result<Reward, WarmStartExactJhError> {
1609 let return_horizon = return_horizon.get();
1610 let return_bins = return_bins.get();
1611 let span = return_bins - 1;
1612 if !span.is_multiple_of(return_horizon) {
1613 return Err(WarmStartExactJhError::ReturnBinsNotExactHorizon {
1614 return_bins,
1615 return_horizon,
1616 });
1617 }
1618 let max_reward = i64::try_from(span / return_horizon)
1619 .map_err(|_| WarmStartExactJhError::ExactReturnRangeOverflow)?;
1620 Ok(max_reward)
1621}
1622
1623#[derive(Clone, Debug)]
1624struct StepRecord {
1625 action: Action,
1626 observations: Vec<PerceptVal>,
1627 reward: Reward,
1628}
1629
1630struct PhaseModel {
1631 predictor: Box<dyn Predictor>,
1632 last_augmented_step: usize,
1633}
1634
1635pub struct WarmStartExactJhAgent {
1637 config: WarmStartExactJhRuntimeConfig,
1638 phases: Vec<PhaseModel>,
1639 steps: Vec<StepRecord>,
1640 return_labels_by_step: Vec<Option<u64>>,
1641 total_steps_observed: usize,
1642 action_bits: usize,
1643 return_label_codec: ReturnLabelCodec,
1644 teacher_label_count: usize,
1645 rng: RandomGenerator,
1646}
1647
1648impl WarmStartExactJhAgent {
1649 pub fn new(
1651 config: WarmStartExactJhConfig,
1652 teacher: WarmStartExactJhTeacherDataset,
1653 ) -> Result<Self, WarmStartExactJhError> {
1654 config.validate_runtime_invariants()?;
1655 let compiled = config.compile_planner_run_spec()?;
1656 Self::from_compiled_planner_run(&compiled, teacher)
1657 }
1658
1659 pub fn from_compiled_planner_run(
1661 compiled: &CompiledPlannerRunSpec,
1662 teacher: WarmStartExactJhTeacherDataset,
1663 ) -> Result<Self, WarmStartExactJhError> {
1664 let config = WarmStartExactJhRuntimeConfig::from_compiled(compiled)?;
1667 config.validate_teacher_contract(&teacher.contract)?;
1668 let (predictor, bit_stream_semantics) = match compiled.controller() {
1669 CompiledPlannerController::AiqiWarmstartExactJh {
1670 predictor,
1671 bit_stream_semantics,
1672 ..
1673 } => (predictor, *bit_stream_semantics),
1674 _ => return Err(WarmStartExactJhError::ControllerKindMismatch),
1675 };
1676 if !predictor.supports_frozen_conditioning() {
1677 return Err(WarmStartExactJhError::UnsupportedRateBackend {
1678 reason: "warm-start exact-J_H strict mode requires frozen context conditioning; configured rate_backend does not provide strict frozen conditioning",
1679 });
1680 }
1681 let action_bits = compiled.action_bits();
1682 let return_label_codec = ReturnLabelCodec::value_monotone(config.return_bins);
1683 let return_bits = return_label_codec.bits();
1684 let mut phases = Vec::with_capacity(config.label_phase_period);
1685 for _ in 0..config.label_phase_period {
1686 phases.push(PhaseModel {
1687 predictor: build_aiqi_predictor(predictor, return_bits, bit_stream_semantics)
1688 .map_err(WarmStartExactJhError::Predictor)?,
1689 last_augmented_step: 0,
1690 });
1691 }
1692 let rng = RandomGenerator::from_seed(config.random_seed);
1693 let mut agent = Self {
1694 action_bits,
1695 return_label_codec,
1696 phases,
1697 steps: Vec::new(),
1698 return_labels_by_step: Vec::new(),
1699 total_steps_observed: 0,
1700 teacher_label_count: 0,
1701 rng,
1702 config,
1703 };
1704 agent.warm_start_from_teacher(&teacher)?;
1705 Ok(agent)
1706 }
1707
1708 pub fn steps_observed(&self) -> usize {
1710 self.total_steps_observed
1711 }
1712
1713 pub fn teacher_label_count(&self) -> usize {
1715 self.teacher_label_count
1716 }
1717
1718 pub fn same_task_live_trace(&self) -> Option<WarmStartExactJhTeacherTrace> {
1730 if self.steps.len() < self.config.return_horizon {
1731 return None;
1732 }
1733 Some(WarmStartExactJhTeacherTrace {
1734 transitions: self
1735 .steps
1736 .iter()
1737 .map(|step| WarmStartExactJhTransition {
1738 action: step.action,
1739 observations: step.observations.clone(),
1740 reward: step.reward,
1741 })
1742 .collect(),
1743 })
1744 }
1745
1746 pub fn num_actions(&self) -> ActionAlphabet {
1748 self.config.agent_actions
1749 }
1750
1751 pub fn planner_simulations_per_step(&self) -> usize {
1755 self.config.planner_simulations_per_step
1756 }
1757
1758 pub fn resolved_random_seed(&self) -> u64 {
1760 self.config.random_seed
1761 }
1762
1763 pub(crate) fn reseed_random(&mut self, seed: u64) {
1764 self.config.random_seed = seed;
1765 self.rng = RandomGenerator::from_seed(seed);
1766 }
1767
1768 pub fn get_planned_action(&mut self) -> Action {
1776 self.try_get_planned_action()
1777 .expect("warm-start planning state must satisfy validated history invariants")
1778 }
1779
1780 pub fn try_get_planned_action(&mut self) -> Result<Action, WarmStartExactJhError> {
1782 let q_values = self.estimate_q_values()?;
1783 Ok(argmax_with_fixed_tie_break(&q_values) as u64)
1784 }
1785
1786 pub fn estimate_action_values(&mut self) -> Result<Vec<f64>, WarmStartExactJhError> {
1788 self.estimate_q_values()
1789 }
1790
1791 pub fn get_planned_action_with_extra_exploration(&mut self, extra_exploration: f64) -> Action {
1803 self.get_planned_action_with_extra_exploration_flag(extra_exploration)
1804 .0
1805 }
1806
1807 pub fn get_planned_action_with_extra_exploration_flag(
1819 &mut self,
1820 extra_exploration: f64,
1821 ) -> (Action, bool) {
1822 self.try_get_planned_action_with_extra_exploration_flag(extra_exploration)
1823 .expect("warm-start planning state must satisfy validated history invariants")
1824 }
1825
1826 pub fn try_get_planned_action_with_extra_exploration_flag(
1831 &mut self,
1832 extra_exploration: f64,
1833 ) -> Result<(Action, bool), WarmStartExactJhError> {
1834 let extra = extra_exploration.clamp(0.0, 1.0);
1835 if extra > 0.0 && self.rng.gen_bool(extra) {
1836 Ok((
1837 self.rng.gen_range(self.config.agent_actions.get()) as u64,
1838 true,
1839 ))
1840 } else {
1841 Ok((self.try_get_planned_action()?, false))
1842 }
1843 }
1844
1845 pub fn observe_transition(
1847 &mut self,
1848 action: Action,
1849 observations: &[PerceptVal],
1850 reward: Reward,
1851 ) -> Result<(), WarmStartExactJhError> {
1852 self.validate_transition(action, observations, reward)?;
1853 self.steps.push(StepRecord {
1854 action,
1855 observations: observations.to_vec(),
1856 reward,
1857 });
1858 self.total_steps_observed = self.total_steps_observed.saturating_add(1);
1859 self.return_labels_by_step.push(None);
1860 self.maybe_learn_new_return()
1861 }
1862
1863 fn warm_start_from_teacher(
1868 &mut self,
1869 teacher: &WarmStartExactJhTeacherDataset,
1870 ) -> Result<(), WarmStartExactJhError> {
1871 let mut label_count = 0usize;
1872 for trace in &teacher.traces {
1873 self.validate_teacher_trace(trace)?;
1874 label_count = label_count.saturating_add(self.commit_teacher_trace(trace)?);
1875 for phase in &mut self.phases {
1876 phase
1877 .predictor
1878 .reset_conditioning_history()
1879 .map_err(|reason| WarmStartExactJhError::PredictorConditioningReset {
1880 reason,
1881 })?;
1882 }
1883 }
1884 if label_count == 0 {
1885 return Err(WarmStartExactJhError::InvalidTeacherDataset {
1886 reason: "teacher dataset did not contain any complete H-step labels".to_string(),
1887 });
1888 }
1889 self.teacher_label_count = label_count;
1890 Ok(())
1891 }
1892
1893 fn validate_teacher_trace(
1894 &self,
1895 trace: &WarmStartExactJhTeacherTrace,
1896 ) -> Result<(), WarmStartExactJhError> {
1897 for step in &trace.transitions {
1898 self.validate_transition(step.action, &step.observations, step.reward)?;
1899 }
1900 Ok(())
1901 }
1902
1903 fn validate_transition(
1904 &self,
1905 action: Action,
1906 observations: &[PerceptVal],
1907 reward: Reward,
1908 ) -> Result<(), WarmStartExactJhError> {
1909 validate_runtime_transition(&self.config, action, observations, reward)
1910 }
1911
1912 fn commit_teacher_trace(
1913 &mut self,
1914 trace: &WarmStartExactJhTeacherTrace,
1915 ) -> Result<usize, WarmStartExactJhError> {
1916 let labels = exact_return_labels_for_trace(&self.config, &trace.transitions)?;
1917 let mut committed = 0usize;
1918 for phase in 0..self.config.label_phase_period {
1919 let model = &mut self.phases[phase];
1920 for (idx0, step) in trace.transitions.iter().enumerate() {
1921 let step_index = idx0 + 1;
1922 push_action_tokens_commit_history(
1923 model.predictor.as_mut(),
1924 step.action,
1925 self.action_bits,
1926 );
1927 if step_index % self.config.label_phase_period == phase
1928 && let Some(label) = labels[idx0]
1929 {
1930 self.return_label_codec
1931 .push_label_commit(model.predictor.as_mut(), label);
1932 committed = committed.saturating_add(1);
1933 }
1934 push_percept_tokens_commit_history(
1935 &self.config,
1936 model.predictor.as_mut(),
1937 &step.observations,
1938 step.reward,
1939 )?;
1940 }
1941 }
1942 Ok(committed)
1943 }
1944
1945 fn maybe_learn_new_return(&mut self) -> Result<(), WarmStartExactJhError> {
1946 let t = self.total_steps_observed;
1947 let h = self.config.return_horizon;
1948 if t < h {
1949 return Ok(());
1950 }
1951 let start_step = t + 1 - h;
1952 let label = self.compute_return_label(start_step)?;
1953 self.return_labels_by_step[start_step - 1] = Some(label);
1954 let phase = start_step % self.config.label_phase_period;
1955 self.advance_phase_model_to_step(phase, start_step)
1956 }
1957
1958 fn estimate_q_values(&mut self) -> Result<Vec<f64>, WarmStartExactJhError> {
1959 let min_return = (self.config.min_reward as i128)
1960 .checked_mul(self.config.return_horizon as i128)
1961 .ok_or(WarmStartExactJhError::ExactReturnRangeOverflow)?
1962 as f64;
1963 let codec = self.return_label_codec;
1964 let step = self.total_steps_observed + 1;
1965 let phase = step % self.config.label_phase_period;
1966 let live_config = &self.config;
1967 let steps = &self.steps;
1968 let return_labels_by_step = &self.return_labels_by_step;
1969 let action_bits = self.action_bits;
1970 let token_ctx = WarmStartAugmentedTokenContext {
1971 config: live_config,
1972 steps,
1973 return_labels_by_step,
1974 action_bits,
1975 return_label_codec: self.return_label_codec,
1976 phase,
1977 };
1978 let mut q_values = Vec::with_capacity(self.config.agent_actions.get());
1979 let mut pushed_history = 0usize;
1980 {
1981 let model = &mut self.phases[phase];
1982 let start = model.last_augmented_step + 1;
1983 let end = step.saturating_sub(1);
1984 if start <= end {
1985 for idx in start..=end {
1986 match push_step_tokens_history(&token_ctx, model.predictor.as_mut(), idx) {
1987 Ok(pushed) => {
1988 pushed_history += pushed;
1989 }
1990 Err(err) => {
1991 pop_history_bits(model.predictor.as_mut(), pushed_history);
1992 return Err(err);
1993 }
1994 }
1995 }
1996 }
1997 for action in 0..self.config.agent_actions.get() {
1998 let pushed_action =
1999 push_encoded_bits_history(model.predictor.as_mut(), action as u64, action_bits);
2000 let expected_label = predict_expected_label(
2001 model.predictor.as_mut(),
2002 codec,
2003 ReturnPrefixUpdate::Training,
2004 ReturnLawEvaluator::SharedPrefix,
2005 );
2006 pop_history_bits(model.predictor.as_mut(), pushed_action);
2007 q_values.push(min_return + expected_label);
2008 }
2009 pop_history_bits(model.predictor.as_mut(), pushed_history);
2010 }
2011 Ok(q_values)
2012 }
2013
2014 fn advance_phase_model_to_step(
2015 &mut self,
2016 phase: usize,
2017 target_step: usize,
2018 ) -> Result<(), WarmStartExactJhError> {
2019 let token_ctx = WarmStartAugmentedTokenContext {
2020 config: &self.config,
2021 steps: &self.steps,
2022 return_labels_by_step: &self.return_labels_by_step,
2023 action_bits: self.action_bits,
2024 return_label_codec: self.return_label_codec,
2025 phase,
2026 };
2027 let model = &mut self.phases[phase];
2028 if target_step <= model.last_augmented_step {
2029 return Ok(());
2030 }
2031 let start = model.last_augmented_step + 1;
2032 for idx in start..=target_step {
2033 push_augmented_step_tokens_commit(&token_ctx, model.predictor.as_mut(), idx)?;
2034 }
2035 model.last_augmented_step = target_step;
2036 Ok(())
2037 }
2038
2039 fn compute_return_label(&self, start_step: usize) -> Result<u64, WarmStartExactJhError> {
2040 let mut total = 0i128;
2041 for offset in 0..self.config.return_horizon {
2042 let idx = start_step + offset;
2043 let step =
2044 self.steps
2045 .get(idx - 1)
2046 .ok_or(WarmStartExactJhError::HistoryIndexOutOfRange {
2047 global_step: idx,
2048 total_steps_observed: self.total_steps_observed,
2049 })?;
2050 total += step.reward as i128;
2051 }
2052 label_for_exact_return(&self.config, total)
2053 }
2054}
2055
2056#[derive(Debug)]
2058#[non_exhaustive]
2059pub enum WarmStartExactJhError {
2060 ControllerKindMismatch,
2062 ReturnHorizonZero,
2064 ReturnBinsZero,
2066 LabelPhasePeriodTooShort {
2068 label_phase_period: usize,
2070 return_horizon: usize,
2072 },
2073 PlannerSimulationsZero,
2075 PlannerSimulationsUnsupported {
2077 configured: usize,
2079 },
2080 ReturnBinsTooSmall {
2082 required: u128,
2084 configured: usize,
2086 },
2087 ReturnBinsNotExactHorizon {
2089 return_bins: usize,
2091 return_horizon: usize,
2093 },
2094 ExactReturnRangeOverflow,
2096 RewardEncoding(RewardEncodingError),
2098 InvalidRateBackend(crate::error::InfotheoryError),
2100 UnsupportedRateBackend {
2102 reason: &'static str,
2104 },
2105 Spec(SpecError),
2107 Predictor(PredictorBuildError),
2109 PredictorConditioningReset {
2111 reason: String,
2113 },
2114 InvalidTeacherDataset {
2116 reason: String,
2118 },
2119 InvalidTelemetry {
2121 reason: String,
2123 },
2124 ActionOutOfRange {
2126 action: Action,
2128 agent_actions: ActionAlphabet,
2130 },
2131 ObservationStreamLengthMismatch {
2133 expected: usize,
2135 actual: usize,
2137 },
2138 ObservationValueOutOfRange {
2140 observation: PerceptVal,
2142 observation_bits: usize,
2144 maximum: PerceptVal,
2146 },
2147 RewardOutOfRange {
2149 reward: Reward,
2151 min_reward: Reward,
2153 max_reward: Reward,
2155 },
2156 HistoryIndexOutOfRange {
2158 global_step: usize,
2160 total_steps_observed: usize,
2162 },
2163 MissingReturnLabel {
2165 step: usize,
2167 phase: usize,
2169 },
2170}
2171
2172impl fmt::Display for WarmStartExactJhError {
2173 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
2174 match self {
2175 Self::ControllerKindMismatch => {
2176 f.write_str("compiled controller kind is not aiqi_warmstart_exact_jh")
2177 }
2178 Self::ReturnHorizonZero => f.write_str("return_horizon must be >= 1"),
2179 Self::ReturnBinsZero => f.write_str("return_bins must be >= 1"),
2180 Self::LabelPhasePeriodTooShort {
2181 label_phase_period,
2182 return_horizon,
2183 } => write!(
2184 f,
2185 "label_phase_period ({label_phase_period}) must be >= return_horizon ({return_horizon})"
2186 ),
2187 Self::PlannerSimulationsZero => f.write_str(
2188 "planner_simulations_per_step must be exactly 1 for warm-start exact-J_H direct evaluation",
2189 ),
2190 Self::PlannerSimulationsUnsupported { configured } => write!(
2191 f,
2192 "planner_simulations_per_step must be exactly 1 for warm-start exact-J_H direct evaluation, got {configured}"
2193 ),
2194 Self::ReturnBinsTooSmall {
2195 required,
2196 configured,
2197 } => write!(
2198 f,
2199 "return_bins too small for exact J_H labels: required {required}, configured {configured}"
2200 ),
2201 Self::ReturnBinsNotExactHorizon {
2202 return_bins,
2203 return_horizon,
2204 } => write!(
2205 f,
2206 "return_bins must be exactly H * max_reward + 1 for warm-start exact-J_H; got return_bins={return_bins}, return_horizon={return_horizon}"
2207 ),
2208 Self::ExactReturnRangeOverflow => {
2209 f.write_str("exact finite-horizon return range overflowed supported integer domain")
2210 }
2211 Self::RewardEncoding(err) => write!(f, "{err}"),
2212 Self::InvalidRateBackend(err) => write!(f, "invalid rate_backend: {err}"),
2213 Self::UnsupportedRateBackend { reason } => f.write_str(reason),
2214 Self::Spec(err) => write!(f, "{err}"),
2215 Self::Predictor(err) => write!(f, "failed to construct predictor: {err}"),
2216 Self::PredictorConditioningReset { reason } => {
2217 write!(
2218 f,
2219 "failed to reset predictor conditioning history: {reason}"
2220 )
2221 }
2222 Self::InvalidTeacherDataset { reason } => {
2223 write!(f, "invalid teacher dataset: {reason}")
2224 }
2225 Self::InvalidTelemetry { reason } => write!(f, "invalid telemetry: {reason}"),
2226 Self::ActionOutOfRange {
2227 action,
2228 agent_actions,
2229 } => write!(
2230 f,
2231 "action {action} is outside configured action alphabet {agent_actions}"
2232 ),
2233 Self::ObservationStreamLengthMismatch { expected, actual } => write!(
2234 f,
2235 "observation stream length mismatch: expected {expected}, got {actual}"
2236 ),
2237 Self::ObservationValueOutOfRange {
2238 observation,
2239 observation_bits,
2240 maximum,
2241 } => write!(
2242 f,
2243 "observation value {observation} does not fit observation_bits={observation_bits} (max={maximum})"
2244 ),
2245 Self::RewardOutOfRange {
2246 reward,
2247 min_reward,
2248 max_reward,
2249 } => write!(
2250 f,
2251 "reward {reward} outside configured range [{min_reward}, {max_reward}]"
2252 ),
2253 Self::HistoryIndexOutOfRange {
2254 global_step,
2255 total_steps_observed,
2256 } => write!(
2257 f,
2258 "global step {global_step} out of observed history range [1, {total_steps_observed}]"
2259 ),
2260 Self::MissingReturnLabel { step, phase } => {
2261 write!(
2262 f,
2263 "missing exact return label for step {step} in phase {phase}"
2264 )
2265 }
2266 }
2267 }
2268}
2269
2270impl Error for WarmStartExactJhError {
2271 fn source(&self) -> Option<&(dyn Error + 'static)> {
2272 match self {
2273 Self::RewardEncoding(err) => Some(err),
2274 Self::InvalidRateBackend(err) => Some(err),
2275 Self::Spec(err) => Some(err),
2276 Self::Predictor(err) => Some(err),
2277 Self::PredictorConditioningReset { .. } => None,
2278 _ => None,
2279 }
2280 }
2281}
2282
2283impl From<RewardEncodingError> for WarmStartExactJhError {
2284 fn from(value: RewardEncodingError) -> Self {
2285 Self::RewardEncoding(value)
2286 }
2287}
2288
2289impl From<SpecError> for WarmStartExactJhError {
2290 fn from(value: SpecError) -> Self {
2291 Self::Spec(value)
2292 }
2293}
2294
2295fn parse_teacher_contract(
2296 object: &serde_json::Map<String, Value>,
2297 schema_version: u64,
2298) -> Result<WarmStartExactJhTeacherContract, WarmStartExactJhError> {
2299 let contract = object
2300 .get("contract")
2301 .and_then(Value::as_object)
2302 .ok_or_else(|| WarmStartExactJhError::InvalidTeacherDataset {
2303 reason: "teacher dataset requires a 'contract' object".to_string(),
2304 })?;
2305 ensure_teacher_fields(
2306 contract,
2307 &[
2308 "task_fingerprint",
2309 "action_alphabet_size",
2310 "observation_bits",
2311 "observation_stream_len",
2312 "observation_key_mode",
2313 "observation_adapter_spec_ref",
2314 "observation_adapter_content_crc32",
2315 "reward_bits",
2316 "return_horizon",
2317 "label_phase_period",
2318 "scalar_representation",
2319 "exact_reward_encoding_certificate",
2320 ],
2321 "contract",
2322 )?;
2323 Ok(WarmStartExactJhTeacherContract {
2324 schema_version,
2325 task_fingerprint: required_teacher_task_fingerprint(contract, "task_fingerprint")?,
2326 action_alphabet_size: required_teacher_usize(contract, "action_alphabet_size")?,
2327 observation_bits: required_teacher_usize(contract, "observation_bits")?,
2328 observation_stream_len: required_teacher_usize(contract, "observation_stream_len")?,
2329 observation_key_mode: required_teacher_string(contract, "observation_key_mode")?,
2330 observation_adapter_spec_ref: required_teacher_string(
2331 contract,
2332 "observation_adapter_spec_ref",
2333 )?,
2334 observation_adapter_content_crc32: required_teacher_string(
2335 contract,
2336 "observation_adapter_content_crc32",
2337 )?,
2338 reward_bits: required_teacher_usize(contract, "reward_bits")?,
2339 return_horizon: required_teacher_usize(contract, "return_horizon")?,
2340 label_phase_period: required_teacher_usize(contract, "label_phase_period")?,
2341 scalar_representation: required_teacher_string(contract, "scalar_representation")?,
2342 exact_reward_encoding_certificate: required_teacher_string(
2343 contract,
2344 "exact_reward_encoding_certificate",
2345 )?,
2346 })
2347}
2348
2349fn required_teacher_task_fingerprint(
2350 object: &serde_json::Map<String, Value>,
2351 field: &str,
2352) -> Result<TaskFingerprint, WarmStartExactJhError> {
2353 let value = object.get(field).and_then(Value::as_str).ok_or_else(|| {
2354 WarmStartExactJhError::InvalidTeacherDataset {
2355 reason: format!("teacher contract field '{field}' must be a string"),
2356 }
2357 })?;
2358 TaskFingerprint::parse_hex(value).ok_or_else(|| WarmStartExactJhError::InvalidTeacherDataset {
2359 reason: format!(
2360 "teacher contract field '{field}' must be a 64-digit lowercase hexadecimal SHA-256 digest, got '{value}'"
2361 ),
2362 })
2363}
2364
2365fn required_teacher_string(
2366 object: &serde_json::Map<String, Value>,
2367 field: &str,
2368) -> Result<String, WarmStartExactJhError> {
2369 object
2370 .get(field)
2371 .and_then(Value::as_str)
2372 .map(str::to_string)
2373 .ok_or_else(|| WarmStartExactJhError::InvalidTeacherDataset {
2374 reason: format!("teacher contract field '{field}' must be a string"),
2375 })
2376}
2377
2378fn required_teacher_usize(
2379 object: &serde_json::Map<String, Value>,
2380 field: &str,
2381) -> Result<usize, WarmStartExactJhError> {
2382 let value = object.get(field).and_then(Value::as_u64).ok_or_else(|| {
2383 WarmStartExactJhError::InvalidTeacherDataset {
2384 reason: format!("teacher contract field '{field}' must be an unsigned integer"),
2385 }
2386 })?;
2387 usize::try_from(value).map_err(|_| WarmStartExactJhError::InvalidTeacherDataset {
2388 reason: format!("teacher contract field '{field}' does not fit usize"),
2389 })
2390}
2391
2392fn parse_teacher_trace(
2393 value: &Value,
2394 trace_index: usize,
2395) -> Result<WarmStartExactJhTeacherTrace, WarmStartExactJhError> {
2396 let object = value
2397 .as_object()
2398 .ok_or_else(|| WarmStartExactJhError::InvalidTeacherDataset {
2399 reason: format!("traces[{trace_index}] must be an object with transitions"),
2400 })?;
2401 ensure_teacher_fields(object, &["transitions"], &format!("traces[{trace_index}]"))?;
2402 let transitions_value = object
2403 .get("transitions")
2404 .and_then(Value::as_array)
2405 .ok_or_else(|| WarmStartExactJhError::InvalidTeacherDataset {
2406 reason: format!("traces[{trace_index}].transitions must be an array"),
2407 })?;
2408 let mut transitions = Vec::with_capacity(transitions_value.len());
2409 for (step_index, transition) in transitions_value.iter().enumerate() {
2410 transitions.push(parse_teacher_transition(
2411 transition,
2412 trace_index,
2413 step_index,
2414 )?);
2415 }
2416 Ok(WarmStartExactJhTeacherTrace { transitions })
2417}
2418
2419fn parse_teacher_transition(
2420 value: &Value,
2421 trace_index: usize,
2422 step_index: usize,
2423) -> Result<WarmStartExactJhTransition, WarmStartExactJhError> {
2424 let object = value
2425 .as_object()
2426 .ok_or_else(|| WarmStartExactJhError::InvalidTeacherDataset {
2427 reason: format!("traces[{trace_index}].transitions[{step_index}] must be an object"),
2428 })?;
2429 ensure_teacher_fields(
2430 object,
2431 &["action", "observations", "reward"],
2432 &format!("traces[{trace_index}].transitions[{step_index}]"),
2433 )?;
2434 let action = object
2435 .get("action")
2436 .and_then(Value::as_u64)
2437 .ok_or_else(|| WarmStartExactJhError::InvalidTeacherDataset {
2438 reason: format!(
2439 "traces[{trace_index}].transitions[{step_index}].action must be an integer"
2440 ),
2441 })?;
2442 let observations_value =
2443 object
2444 .get("observations")
2445 .ok_or_else(|| WarmStartExactJhError::InvalidTeacherDataset {
2446 reason: format!(
2447 "traces[{trace_index}].transitions[{step_index}] requires observations"
2448 ),
2449 })?;
2450 let observations = observations_value
2451 .as_array()
2452 .ok_or_else(|| WarmStartExactJhError::InvalidTeacherDataset {
2453 reason: format!(
2454 "traces[{trace_index}].transitions[{step_index}].observations must be an array"
2455 ),
2456 })?
2457 .iter()
2458 .enumerate()
2459 .map(|(obs_index, obs)| {
2460 obs.as_u64().ok_or_else(|| WarmStartExactJhError::InvalidTeacherDataset {
2461 reason: format!(
2462 "traces[{trace_index}].transitions[{step_index}].observations[{obs_index}] must be an integer"
2463 ),
2464 })
2465 })
2466 .collect::<Result<Vec<_>, _>>()?;
2467 let reward = object
2468 .get("reward")
2469 .and_then(Value::as_i64)
2470 .ok_or_else(|| WarmStartExactJhError::InvalidTeacherDataset {
2471 reason: format!(
2472 "traces[{trace_index}].transitions[{step_index}].reward must be an integer"
2473 ),
2474 })?;
2475 Ok(WarmStartExactJhTransition {
2476 action,
2477 observations,
2478 reward,
2479 })
2480}
2481
2482fn ensure_teacher_fields(
2483 object: &serde_json::Map<String, Value>,
2484 allowed: &[&str],
2485 label: &str,
2486) -> Result<(), WarmStartExactJhError> {
2487 for key in object.keys() {
2488 if !allowed.contains(&key.as_str()) {
2489 return Err(WarmStartExactJhError::InvalidTeacherDataset {
2490 reason: format!("{label} contains unknown teacher field '{key}'"),
2491 });
2492 }
2493 }
2494 Ok(())
2495}
2496
2497fn validate_exact_return_alphabet(
2498 min_reward: Reward,
2499 max_reward: Reward,
2500 return_horizon: usize,
2501 return_bins: usize,
2502) -> Result<(), WarmStartExactJhError> {
2503 let min_return = (min_reward as i128)
2504 .checked_mul(return_horizon as i128)
2505 .ok_or(WarmStartExactJhError::ExactReturnRangeOverflow)?;
2506 let max_return = (max_reward as i128)
2507 .checked_mul(return_horizon as i128)
2508 .ok_or(WarmStartExactJhError::ExactReturnRangeOverflow)?;
2509 let span = max_return
2510 .checked_sub(min_return)
2511 .and_then(|value| value.checked_add(1))
2512 .ok_or(WarmStartExactJhError::ExactReturnRangeOverflow)?;
2513 let required =
2514 u128::try_from(span).map_err(|_| WarmStartExactJhError::ExactReturnRangeOverflow)?;
2515 if required > return_bins as u128 {
2516 return Err(WarmStartExactJhError::ReturnBinsTooSmall {
2517 required,
2518 configured: return_bins,
2519 });
2520 }
2521 Ok(())
2522}
2523
2524fn exact_return_labels_for_trace(
2525 config: &WarmStartExactJhRuntimeConfig,
2526 steps: &[WarmStartExactJhTransition],
2527) -> Result<Vec<Option<u64>>, WarmStartExactJhError> {
2528 let mut labels = vec![None; steps.len()];
2529 if config.return_horizon == 0 || steps.len() < config.return_horizon {
2530 return Ok(labels);
2531 }
2532
2533 let horizon = config.return_horizon;
2534 let mut window_sum = 0_i128;
2535 for (index, step) in steps.iter().enumerate() {
2536 window_sum += step.reward as i128;
2537 if index >= horizon {
2538 window_sum -= steps[index - horizon].reward as i128;
2539 }
2540 if index + 1 >= horizon {
2541 let start0 = index + 1 - horizon;
2542 labels[start0] = Some(label_for_exact_return(config, window_sum)?);
2543 }
2544 }
2545 Ok(labels)
2546}
2547
2548fn label_for_exact_return(
2549 config: &WarmStartExactJhRuntimeConfig,
2550 exact_return: i128,
2551) -> Result<u64, WarmStartExactJhError> {
2552 let min_return = (config.min_reward as i128)
2553 .checked_mul(config.return_horizon as i128)
2554 .ok_or(WarmStartExactJhError::ExactReturnRangeOverflow)?;
2555 let label = exact_return
2556 .checked_sub(min_return)
2557 .ok_or(WarmStartExactJhError::ExactReturnRangeOverflow)?;
2558 if label < 0 || label >= config.return_bins as i128 {
2559 return Err(WarmStartExactJhError::ReturnBinsTooSmall {
2560 required: (label + 1).max(0) as u128,
2561 configured: config.return_bins,
2562 });
2563 }
2564 u64::try_from(label).map_err(|_| WarmStartExactJhError::ExactReturnRangeOverflow)
2565}
2566
2567fn validate_runtime_transition(
2568 config: &WarmStartExactJhRuntimeConfig,
2569 action: Action,
2570 observations: &[PerceptVal],
2571 reward: Reward,
2572) -> Result<(), WarmStartExactJhError> {
2573 if action as usize >= config.agent_actions.get() {
2574 return Err(WarmStartExactJhError::ActionOutOfRange {
2575 action,
2576 agent_actions: config.agent_actions,
2577 });
2578 }
2579 if observations.len() != config.observation_stream_len {
2580 return Err(WarmStartExactJhError::ObservationStreamLengthMismatch {
2581 expected: config.observation_stream_len,
2582 actual: observations.len(),
2583 });
2584 }
2585 let obs_max = max_value_for_bits(config.observation_bits);
2586 for &observation in observations {
2587 if observation > obs_max {
2588 return Err(WarmStartExactJhError::ObservationValueOutOfRange {
2589 observation,
2590 observation_bits: config.observation_bits,
2591 maximum: obs_max,
2592 });
2593 }
2594 }
2595 if reward < config.min_reward || reward > config.max_reward {
2596 return Err(WarmStartExactJhError::RewardOutOfRange {
2597 reward,
2598 min_reward: config.min_reward,
2599 max_reward: config.max_reward,
2600 });
2601 }
2602 Ok(())
2603}
2604
2605struct WarmStartAugmentedTokenContext<'a> {
2606 config: &'a WarmStartExactJhRuntimeConfig,
2607 steps: &'a [StepRecord],
2608 return_labels_by_step: &'a [Option<u64>],
2609 action_bits: usize,
2610 return_label_codec: ReturnLabelCodec,
2611 phase: usize,
2612}
2613
2614fn push_augmented_step_tokens_commit(
2615 ctx: &WarmStartAugmentedTokenContext<'_>,
2616 predictor: &mut dyn Predictor,
2617 idx: usize,
2618) -> Result<usize, WarmStartExactJhError> {
2619 let step = &ctx.steps[idx - 1];
2620 let return_label = if idx % ctx.config.label_phase_period == ctx.phase {
2621 Some(ctx.return_labels_by_step[idx - 1].ok_or(
2622 WarmStartExactJhError::MissingReturnLabel {
2623 step: idx,
2624 phase: ctx.phase,
2625 },
2626 )?)
2627 } else {
2628 None
2629 };
2630 let reward_value = encoded_reward_value(
2631 step.reward,
2632 ctx.config.reward_bits,
2633 ctx.config.reward_offset,
2634 )?;
2635
2636 let mut pushed = 0usize;
2637 pushed += push_action_tokens_commit_history(predictor, step.action, ctx.action_bits);
2638 if let Some(label) = return_label {
2639 pushed += ctx.return_label_codec.push_label_commit(predictor, label);
2640 }
2641 pushed += push_percept_tokens_commit_history_encoded(
2642 ctx.config,
2643 predictor,
2644 &step.observations,
2645 reward_value,
2646 );
2647 Ok(pushed)
2648}
2649
2650fn push_step_tokens_history(
2651 ctx: &WarmStartAugmentedTokenContext<'_>,
2652 predictor: &mut dyn Predictor,
2653 idx: usize,
2654) -> Result<usize, WarmStartExactJhError> {
2655 let step = &ctx.steps[idx - 1];
2656 let reward_value = encoded_reward_value(
2657 step.reward,
2658 ctx.config.reward_bits,
2659 ctx.config.reward_offset,
2660 )?;
2661
2662 let mut pushed = 0usize;
2663 pushed += push_encoded_bits_history(predictor, step.action, ctx.action_bits);
2664 if idx % ctx.config.label_phase_period == ctx.phase
2665 && let Some(label) = ctx.return_labels_by_step[idx - 1]
2666 {
2667 pushed += ctx.return_label_codec.push_label_history(predictor, label);
2668 }
2669 pushed += push_percept_tokens_history_encoded(
2670 ctx.config,
2671 predictor,
2672 &step.observations,
2673 reward_value,
2674 );
2675 Ok(pushed)
2676}
2677
2678fn push_percept_tokens_commit_history(
2679 config: &WarmStartExactJhRuntimeConfig,
2680 predictor: &mut dyn Predictor,
2681 observations: &[PerceptVal],
2682 reward: Reward,
2683) -> Result<usize, WarmStartExactJhError> {
2684 let reward_value = encoded_reward_value(reward, config.reward_bits, config.reward_offset)?;
2685 Ok(push_percept_tokens_commit_history_encoded(
2686 config,
2687 predictor,
2688 observations,
2689 reward_value,
2690 ))
2691}
2692
2693fn push_percept_tokens_commit_history_encoded(
2694 config: &WarmStartExactJhRuntimeConfig,
2695 predictor: &mut dyn Predictor,
2696 observations: &[PerceptVal],
2697 reward_value: u64,
2698) -> usize {
2699 let mut pushed = 0usize;
2700 for &observation in observations {
2701 pushed += push_encoded_bits_commit_history(predictor, observation, config.observation_bits);
2702 }
2703 pushed += push_encoded_bits_commit_history(predictor, reward_value, config.reward_bits);
2704 pushed
2705}
2706
2707fn push_percept_tokens_history_encoded(
2708 config: &WarmStartExactJhRuntimeConfig,
2709 predictor: &mut dyn Predictor,
2710 observations: &[PerceptVal],
2711 reward_value: u64,
2712) -> usize {
2713 let mut pushed = 0usize;
2714 for &observation in observations {
2715 pushed += push_encoded_bits_history(predictor, observation, config.observation_bits);
2716 }
2717 pushed += push_encoded_bits_history(predictor, reward_value, config.reward_bits);
2718 pushed
2719}
2720
2721fn push_action_tokens_commit_history(
2722 predictor: &mut dyn Predictor,
2723 action: Action,
2724 action_bits: usize,
2725) -> usize {
2726 push_encoded_bits_commit_history(predictor, action, action_bits)
2727}
2728
2729fn push_encoded_bits_history(predictor: &mut dyn Predictor, value: u64, bits: usize) -> usize {
2730 let mut v = value;
2731 for _ in 0..bits {
2732 predictor.update_history((v & 1) == 1);
2733 v >>= 1;
2734 }
2735 bits
2736}
2737
2738fn push_encoded_bits_commit_history(
2739 predictor: &mut dyn Predictor,
2740 value: u64,
2741 bits: usize,
2742) -> usize {
2743 let mut v = value;
2744 for _ in 0..bits {
2745 predictor.commit_update_history((v & 1) == 1);
2746 v >>= 1;
2747 }
2748 bits
2749}
2750
2751fn encoded_reward_value(
2752 reward: Reward,
2753 bits: usize,
2754 offset: Reward,
2755) -> Result<u64, WarmStartExactJhError> {
2756 validate_reward_encoding_bounds(reward, reward, offset, bits)
2757 .map_err(WarmStartExactJhError::from)?;
2758 let shifted = (reward as i128) + (offset as i128);
2759 debug_assert!(
2760 shifted >= 0,
2761 "validate_reward_encoding_bounds implies shifted minimum >= 0"
2762 );
2763 Ok(shifted as u64)
2764}
2765
2766fn pop_history_bits(predictor: &mut dyn Predictor, bits: usize) {
2767 for _ in 0..bits {
2768 predictor.pop_history();
2769 }
2770}
2771
2772fn max_value_for_bits(bits: usize) -> u64 {
2773 if bits >= 64 {
2774 u64::MAX
2775 } else if bits == 0 {
2776 0
2777 } else {
2778 (1u64 << bits) - 1
2779 }
2780}
2781
2782fn argmax_with_fixed_tie_break(values: &[f64]) -> usize {
2783 let mut best_value = f64::NEG_INFINITY;
2784 let mut best_index = 0usize;
2785 for (index, &value) in values.iter().enumerate() {
2786 if value > best_value {
2787 best_value = value;
2788 best_index = index;
2789 }
2790 }
2791 best_index
2792}
2793
2794#[cfg(test)]
2795mod tests {
2796 use super::*;
2797 use crate::aixi::warmstart_contract::standalone_teacher_provenance_crc32_pair;
2798 use std::sync::{Arc, Mutex};
2799
2800 const TEST_TASK_FINGERPRINT_HEX: &str =
2801 "0102030401020304010203040102030401020304010203040102030401020304";
2802 #[cfg(feature = "backend-ctw")]
2803 const ZERO_TASK_FINGERPRINT_HEX: &str =
2804 "0000000000000000000000000000000000000000000000000000000000000000";
2805 #[cfg(feature = "backend-ctw")]
2806 const MISMATCH_TASK_FINGERPRINT_HEX: &str =
2807 "deadbeefdeadbeefdeadbeefdeadbeefdeadbeefdeadbeefdeadbeefdeadbeef";
2808
2809 fn action_alphabet(n: usize) -> ActionAlphabet {
2810 ActionAlphabet::try_from_usize(n).expect("test action alphabet must be non-zero")
2811 }
2812
2813 fn config() -> WarmStartExactJhConfig {
2814 WarmStartExactJhConfig {
2815 rate_backend: RateBackend::Ctw { depth: 4 },
2816 observation_bits: 2,
2817 observation_stream_len: 1,
2818 reward_bits: 2,
2819 agent_actions: action_alphabet(2),
2820 return_horizon: 1,
2821 return_bins: 4,
2822 label_phase_period: 1,
2823 planner_simulations_per_step: 1,
2824 bit_stream_semantics: BitStreamSemantics::BinaryTokens,
2825 random_seed: Some(9),
2826 }
2827 }
2828
2829 #[test]
2830 fn reward_bounds_from_exact_return_bins_matches_default_test_config() {
2831 let cfg = config();
2832 assert_eq!(
2833 reward_bounds_from_exact_return_bins(
2834 NonZeroUsize::new(cfg.return_horizon).expect("non-zero return horizon"),
2835 NonZeroUsize::new(cfg.return_bins).expect("non-zero return bins"),
2836 cfg.reward_bits
2837 )
2838 .expect("bounds"),
2839 (0, 3, 0)
2840 );
2841 }
2842
2843 #[test]
2844 fn exact_return_labels_cover_each_dense_window() {
2845 let runtime = WarmStartExactJhRuntimeConfig {
2846 task_fingerprint: TaskFingerprint::parse_hex(TEST_TASK_FINGERPRINT_HEX)
2847 .expect("test fingerprint"),
2848 observation_bits: 2,
2849 observation_stream_len: 1,
2850 observation_key_mode: "full_stream",
2851 reward_bits: 2,
2852 agent_actions: action_alphabet(2),
2853 min_reward: 0,
2854 max_reward: 3,
2855 reward_offset: 0,
2856 return_horizon: 3,
2857 return_bins: 10,
2858 label_phase_period: 3,
2859 planner_simulations_per_step: 1,
2860 random_seed: 11,
2861 provenance_policy: TeacherProvenancePolicy::StandalonePlannerRun,
2862 };
2863 let steps = [1, 2, 0, 3, 1]
2864 .into_iter()
2865 .map(|reward| WarmStartExactJhTransition {
2866 action: 0,
2867 observations: vec![0],
2868 reward,
2869 })
2870 .collect::<Vec<_>>();
2871
2872 let labels = exact_return_labels_for_trace(&runtime, &steps).expect("dense return labels");
2873
2874 assert_eq!(labels, vec![Some(3), Some(5), Some(4), None, None]);
2875 }
2876
2877 #[cfg(feature = "backend-ctw")]
2878 #[test]
2879 fn validate_warmstart_teacher_against_compiled_accepts_matching_contract() {
2880 let cfg = config();
2881 let compiled = cfg.compile_planner_run_spec().expect("compile planner run");
2882 let teacher = teacher_for_config(&cfg);
2883 validate_warmstart_teacher_against_compiled_planner_run(&compiled, &teacher.contract)
2884 .expect("matching teacher must validate");
2885 validate_warmstart_teacher_dataset_for_compiled_planner_run(&compiled, &teacher)
2886 .expect("matching teacher dataset must validate");
2887 }
2888
2889 #[cfg(feature = "backend-ctw")]
2890 #[test]
2891 fn full_teacher_validator_rejects_short_trace_before_export() {
2892 let mut cfg = config();
2893 cfg.return_horizon = 2;
2894 cfg.return_bins = 5;
2895 cfg.label_phase_period = 2;
2896 let compiled = cfg.compile_planner_run_spec().expect("compile planner run");
2897 let mut teacher = teacher_for_config(&cfg);
2898 teacher.traces[0].transitions = vec![WarmStartExactJhTransition {
2899 action: 0,
2900 observations: vec![1],
2901 reward: 0,
2902 }];
2903
2904 let err = validate_warmstart_teacher_dataset_for_compiled_planner_run(&compiled, &teacher)
2905 .expect_err("short trace must fail before write/load");
2906 assert!(err.to_string().contains("return_horizon is 2"), "{err}");
2907 }
2908
2909 #[cfg(feature = "backend-ctw")]
2910 #[test]
2911 fn full_teacher_validator_rejects_bit_width_valid_but_runtime_invalid_reward() {
2912 let mut cfg = config();
2913 cfg.return_horizon = 2;
2914 cfg.return_bins = 5;
2915 cfg.label_phase_period = 2;
2916 let compiled = cfg.compile_planner_run_spec().expect("compile planner run");
2917 let mut teacher = teacher_for_config(&cfg);
2918 teacher.traces[0].transitions = vec![
2919 WarmStartExactJhTransition {
2920 action: 0,
2921 observations: vec![1],
2922 reward: 3,
2923 },
2924 WarmStartExactJhTransition {
2925 action: 1,
2926 observations: vec![2],
2927 reward: 0,
2928 },
2929 ];
2930
2931 let err = validate_warmstart_teacher_dataset_for_compiled_planner_run(&compiled, &teacher)
2932 .expect_err("reward valid for reward_bits but outside exact runtime range must fail");
2933 assert!(err.to_string().contains("runtime contract"), "{err}");
2934 assert!(
2935 err.to_string().contains("outside configured range"),
2936 "{err}"
2937 );
2938 }
2939
2940 #[cfg(feature = "backend-ctw")]
2941 #[test]
2942 fn validate_warmstart_teacher_against_compiled_rejects_reward_certificate_crc_mismatch() {
2943 let cfg = config();
2944 let compiled = cfg.compile_planner_run_spec().expect("compile planner run");
2945 let mut teacher = teacher_for_config(&cfg);
2946 teacher.contract.exact_reward_encoding_certificate = "00000000".to_string();
2947 let err =
2948 validate_warmstart_teacher_against_compiled_planner_run(&compiled, &teacher.contract)
2949 .expect_err("corrupted certificate hash must fail");
2950 assert!(matches!(
2951 err,
2952 WarmStartExactJhError::InvalidTeacherDataset { .. }
2953 ));
2954 }
2955
2956 #[cfg(feature = "backend-ctw")]
2957 #[test]
2958 fn validate_warmstart_teacher_planner_task_fingerprint_rejects_mismatch_with_stable_markers() {
2959 let cfg = config();
2960 let compiled = cfg.compile_planner_run_spec().expect("compile planner run");
2961 let mut teacher = teacher_for_config(&cfg);
2962 teacher.contract.task_fingerprint = TaskFingerprint::parse_hex(ZERO_TASK_FINGERPRINT_HEX)
2963 .expect("valid mismatch fingerprint");
2964 let err = validate_warmstart_teacher_planner_task_fingerprint(&compiled, &teacher.contract)
2965 .expect_err("wrong fingerprint must fail");
2966 let msg = err.to_string();
2967 assert!(msg.contains("task_fingerprint"), "{msg}");
2968 assert!(msg.contains("current planner_run '"), "{msg}");
2969 }
2970
2971 fn teacher_contract_for(
2979 cfg: &WarmStartExactJhConfig,
2980 task_fingerprint: TaskFingerprint,
2981 observation_key_mode: &str,
2982 ) -> WarmStartExactJhTeacherContract {
2983 let observation_stream_len = cfg.observation_stream_len.max(1);
2984 let (adapter_crc, reward_cert) = standalone_teacher_provenance_crc32_pair(
2985 cfg.observation_bits,
2986 observation_stream_len,
2987 cfg.reward_bits,
2988 )
2989 .expect("standalone teacher provenance crc pair");
2990 WarmStartExactJhTeacherContract {
2991 schema_version: WARMSTART_TEACHER_CONTRACT_SCHEMA_VERSION,
2992 task_fingerprint,
2993 action_alphabet_size: cfg.agent_actions.get(),
2994 observation_bits: cfg.observation_bits,
2995 observation_stream_len,
2996 observation_key_mode: observation_key_mode.to_string(),
2997 observation_adapter_spec_ref: WARMSTART_STANDALONE_OBSERVATION_ADAPTER_SPEC_REF
2998 .to_string(),
2999 observation_adapter_content_crc32: adapter_crc,
3000 reward_bits: cfg.reward_bits,
3001 return_horizon: cfg.return_horizon,
3002 label_phase_period: cfg.label_phase_period,
3003 scalar_representation: WARMSTART_STANDALONE_SCALAR_REPRESENTATION.to_string(),
3004 exact_reward_encoding_certificate: reward_cert,
3005 }
3006 }
3007
3008 fn teacher_contract() -> WarmStartExactJhTeacherContract {
3014 teacher_contract_for(
3015 &config(),
3016 TaskFingerprint::parse_hex(TEST_TASK_FINGERPRINT_HEX).expect("test task fingerprint"),
3017 "full_stream",
3018 )
3019 }
3020
3021 #[cfg(feature = "backend-ctw")]
3022 fn teacher_for_config(cfg: &WarmStartExactJhConfig) -> WarmStartExactJhTeacherDataset {
3023 let compiled = cfg
3024 .compile_planner_run_spec()
3025 .expect("test planner run must compile");
3026 let task_fingerprint = warmstart_exact_jh_planner_task_fingerprint(&compiled)
3027 .expect("test planner fingerprint");
3028 let observation_key_mode =
3029 observation_key_mode_name(compiled.interface().observation_key_mode);
3030 WarmStartExactJhTeacherDataset {
3031 contract: teacher_contract_for(cfg, task_fingerprint, observation_key_mode),
3032 traces: vec![WarmStartExactJhTeacherTrace {
3033 transitions: vec![
3034 WarmStartExactJhTransition {
3035 action: 0,
3036 observations: vec![1],
3037 reward: 0,
3038 },
3039 WarmStartExactJhTransition {
3040 action: 1,
3041 observations: vec![2],
3042 reward: 3,
3043 },
3044 WarmStartExactJhTransition {
3045 action: 1,
3046 observations: vec![2],
3047 reward: 3,
3048 },
3049 ],
3050 }],
3051 }
3052 }
3053
3054 #[cfg(feature = "backend-ctw")]
3055 fn teacher() -> WarmStartExactJhTeacherDataset {
3056 teacher_for_config(&config())
3057 }
3058
3059 #[test]
3060 fn jsonl_trace_converter_round_trips_action_percept_pairs() {
3061 let contract = teacher_contract();
3062 let jsonl = [
3063 warmstart_jsonl_action_record(0, 0, PlannerActionProvenance::Greedy).to_string(),
3064 warmstart_jsonl_percept_record(0, &[1], 0).to_string(),
3065 warmstart_jsonl_action_record(1, 1, PlannerActionProvenance::Exploratory).to_string(),
3066 warmstart_jsonl_percept_record(1, &[2], 3).to_string(),
3067 ]
3068 .join("\n");
3069 let trace = warmstart_teacher_trace_from_jsonl_slice(
3070 jsonl.as_bytes(),
3071 &contract,
3072 contract.return_horizon,
3073 )
3074 .expect("jsonl trace should parse");
3075 assert_eq!(
3076 trace.transitions,
3077 vec![
3078 WarmStartExactJhTransition {
3079 action: 0,
3080 observations: vec![1],
3081 reward: 0,
3082 },
3083 WarmStartExactJhTransition {
3084 action: 1,
3085 observations: vec![2],
3086 reward: 3,
3087 },
3088 ]
3089 );
3090 }
3091
3092 #[test]
3093 fn jsonl_trace_converter_converts_mcaixi_decision_percept_order() {
3094 let contract = teacher_contract();
3095 let jsonl = [
3096 warmstart_jsonl_percept_record(0, &[0], 0).to_string(),
3097 warmstart_jsonl_action_record(0, 0, PlannerActionProvenance::Greedy).to_string(),
3098 warmstart_jsonl_percept_record(1, &[1], 3).to_string(),
3099 warmstart_jsonl_action_record(1, 1, PlannerActionProvenance::Greedy).to_string(),
3100 warmstart_jsonl_percept_record(2, &[2], 0).to_string(),
3101 ]
3102 .join("\n");
3103 let trace = warmstart_teacher_trace_from_jsonl_slice(
3104 jsonl.as_bytes(),
3105 &contract,
3106 contract.return_horizon,
3107 )
3108 .expect("MC-AIXI-order jsonl trace should parse");
3109 assert_eq!(
3110 trace.transitions,
3111 vec![
3112 WarmStartExactJhTransition {
3113 action: 0,
3114 observations: vec![1],
3115 reward: 3,
3116 },
3117 WarmStartExactJhTransition {
3118 action: 1,
3119 observations: vec![2],
3120 reward: 0,
3121 },
3122 ]
3123 );
3124 }
3125
3126 #[test]
3127 fn jsonl_trace_converter_rejects_sparse_action_then_percept_steps() {
3128 let contract = teacher_contract();
3129 let jsonl = [
3130 warmstart_jsonl_action_record(0, 0, PlannerActionProvenance::Greedy).to_string(),
3131 warmstart_jsonl_percept_record(0, &[1], 0).to_string(),
3132 warmstart_jsonl_action_record(2, 1, PlannerActionProvenance::Greedy).to_string(),
3133 warmstart_jsonl_percept_record(2, &[2], 3).to_string(),
3134 ]
3135 .join("\n");
3136
3137 let err = warmstart_teacher_trace_from_jsonl_slice(
3138 jsonl.as_bytes(),
3139 &contract,
3140 contract.return_horizon,
3141 )
3142 .expect_err("sparse action/percept JSONL trace must fail");
3143
3144 assert!(err.to_string().contains("not contiguous"), "{err}");
3145 assert!(err.to_string().contains("step 1"), "{err}");
3146 }
3147
3148 #[test]
3149 fn jsonl_trace_converter_rejects_sparse_decision_percept_steps() {
3150 let contract = teacher_contract();
3151 let jsonl = [
3152 warmstart_jsonl_percept_record(0, &[0], 0).to_string(),
3153 warmstart_jsonl_action_record(0, 0, PlannerActionProvenance::Greedy).to_string(),
3154 warmstart_jsonl_percept_record(1, &[1], 3).to_string(),
3155 warmstart_jsonl_percept_record(2, &[2], 0).to_string(),
3156 warmstart_jsonl_action_record(2, 1, PlannerActionProvenance::Greedy).to_string(),
3157 warmstart_jsonl_percept_record(3, &[2], 0).to_string(),
3158 ]
3159 .join("\n");
3160
3161 let err = warmstart_teacher_trace_from_jsonl_slice(
3162 jsonl.as_bytes(),
3163 &contract,
3164 contract.return_horizon,
3165 )
3166 .expect_err("sparse decision-percept JSONL trace must fail");
3167
3168 assert!(err.to_string().contains("not contiguous"), "{err}");
3169 assert!(err.to_string().contains("step 1"), "{err}");
3170 }
3171
3172 #[test]
3173 fn jsonl_trace_converter_rejects_malformed_and_inconsistent_records() {
3174 let contract = teacher_contract();
3175 let return_horizon = contract.return_horizon;
3176 for (jsonl, expected) in [
3177 ("{", "invalid JSON"),
3178 (
3179 r#"{"kind":"action","t":0,"action":0,"provenance":"unknown"}"#,
3180 "unknown action provenance",
3181 ),
3182 (
3183 r#"{"kind":"action","t":0,"action":0,"provenance":7}"#,
3184 "action provenance must be a string",
3185 ),
3186 (
3187 r#"{"kind":"action","t":0,"action":0}"#,
3188 "cannot infer JSONL action/percept convention",
3189 ),
3190 (
3191 r#"{"kind":"percept","t":0,"observations":[1],"reward":0}"#,
3192 "cannot infer JSONL action/percept convention",
3193 ),
3194 (
3195 concat!(
3196 r#"{"kind":"action","t":0,"action":0}"#,
3197 "\n",
3198 r#"{"kind":"percept","t":0,"observations":[4],"reward":0}"#
3199 ),
3200 "observation value",
3201 ),
3202 (
3203 r#"{"kind":"action","t":0,"action":0,"extra":true}"#,
3204 "unknown field 'extra'",
3205 ),
3206 ("\n", "empty JSONL records are not allowed"),
3207 (
3208 concat!(
3209 r#"{"kind":"action","t":0,"action":0}"#,
3210 "\n",
3211 r#"{"kind":"percept","t":0,"observations":[1],"reward":0}"#,
3212 "\n",
3213 r#"{"kind":"percept","t":1,"observations":[1],"reward":0}"#,
3214 "\n",
3215 r#"{"kind":"action","t":1,"action":0}"#
3216 ),
3217 "mixed JSONL action/percept conventions",
3218 ),
3219 ] {
3220 let err = warmstart_teacher_trace_from_jsonl_slice(
3221 jsonl.as_bytes(),
3222 &contract,
3223 return_horizon,
3224 )
3225 .expect_err("invalid JSONL trace should fail");
3226 assert!(
3227 err.to_string().contains(expected),
3228 "expected '{expected}' in {err}"
3229 );
3230 }
3231 }
3232
3233 #[test]
3234 fn jsonl_trace_converter_accepts_absent_provenance() {
3235 let contract = teacher_contract();
3236 let absent = concat!(
3237 r#"{"kind":"action","t":0,"action":1}"#,
3238 "\n",
3239 r#"{"kind":"percept","t":0,"observations":[2],"reward":1}"#,
3240 "\n",
3241 r#"{"kind":"action","t":1,"action":0}"#,
3242 "\n",
3243 r#"{"kind":"percept","t":1,"observations":[1],"reward":0}"#
3244 );
3245 let trace = warmstart_teacher_trace_from_jsonl_slice(absent.as_bytes(), &contract, 1)
3246 .expect("legacy trace without provenance should parse");
3247 assert_eq!(trace.transitions.len(), 2);
3248 assert_eq!(trace.transitions[0].action, 1);
3249 assert_eq!(trace.transitions[1].action, 0);
3250 }
3251
3252 #[test]
3253 fn jsonl_trace_converter_rejects_duplicate_action_and_percept_records() {
3254 let contract = teacher_contract();
3255 let duplicate_action = concat!(
3256 r#"{"kind":"action","t":0,"action":0}"#,
3257 "\n",
3258 r#"{"kind":"action","t":0,"action":1}"#,
3259 "\n",
3260 r#"{"kind":"percept","t":0,"observations":[1],"reward":0}"#
3261 );
3262 let err =
3263 warmstart_teacher_trace_from_jsonl_slice(duplicate_action.as_bytes(), &contract, 1)
3264 .expect_err("duplicate action must fail");
3265 assert!(err.to_string().contains("duplicate action record"), "{err}");
3266
3267 let duplicate_percept = concat!(
3268 r#"{"kind":"action","t":0,"action":0}"#,
3269 "\n",
3270 r#"{"kind":"percept","t":0,"observations":[1],"reward":0}"#,
3271 "\n",
3272 r#"{"kind":"percept","t":0,"observations":[2],"reward":1}"#
3273 );
3274 let err =
3275 warmstart_teacher_trace_from_jsonl_slice(duplicate_percept.as_bytes(), &contract, 1)
3276 .expect_err("duplicate percept must fail");
3277 assert!(
3278 err.to_string().contains("duplicate percept record"),
3279 "{err}"
3280 );
3281 }
3282
3283 #[test]
3284 fn trace_recorder_requires_a_complete_return_horizon_window() {
3285 let contract = teacher_contract();
3286 let mut recorder = WarmStartExactJhTraceRecorder::new();
3287 recorder
3288 .record_action(0, 1)
3289 .expect("record action at step 0");
3290 recorder
3291 .record_percept(0, &[2], 3)
3292 .expect("record percept at step 0");
3293 let err = recorder
3294 .into_teacher_trace(&contract, 2)
3295 .expect_err("single transition must fail for return_horizon 2");
3296 assert!(err.to_string().contains("return_horizon is 2"), "{err}");
3297
3298 let mut recorder = WarmStartExactJhTraceRecorder::new();
3299 recorder
3300 .record_action(0, 1)
3301 .expect("record action at step 0");
3302 recorder
3303 .record_percept(0, &[2], 3)
3304 .expect("record percept at step 0");
3305 let trace = recorder
3306 .into_teacher_trace(&contract, 1)
3307 .expect("one transition covers horizon one");
3308 assert_eq!(trace.transitions.len(), 1);
3309 assert_eq!(trace.transitions[0].action, 1);
3310 assert_eq!(trace.transitions[0].observations, vec![2]);
3311 assert_eq!(trace.transitions[0].reward, 3);
3312 }
3313
3314 #[test]
3315 fn trace_recorder_rejects_sparse_step_sets() {
3316 let contract = teacher_contract();
3317 let mut recorder = WarmStartExactJhTraceRecorder::new();
3318 recorder
3319 .record_action(0, 0)
3320 .expect("record action at step 0");
3321 recorder
3322 .record_percept(0, &[1], 0)
3323 .expect("record percept at step 0");
3324 recorder
3325 .record_action(2, 1)
3326 .expect("record action at step 2");
3327 recorder
3328 .record_percept(2, &[2], 3)
3329 .expect("record percept at step 2");
3330
3331 let err = recorder
3332 .into_teacher_trace(&contract, 2)
3333 .expect_err("sparse recorder trace must fail");
3334
3335 assert!(err.to_string().contains("not contiguous"), "{err}");
3336 assert!(err.to_string().contains("step 1"), "{err}");
3337 }
3338
3339 #[derive(Clone, Debug, Default, Eq, PartialEq)]
3340 struct PhaseUpdateCounts {
3341 commit_label_bits: usize,
3342 commit_history_bits: usize,
3343 history_bits: usize,
3344 }
3345
3346 #[derive(Clone)]
3347 struct PhaseUpdateCountingPredictor {
3348 counts: Arc<Mutex<PhaseUpdateCounts>>,
3349 }
3350
3351 impl Predictor for PhaseUpdateCountingPredictor {
3352 fn update(&mut self, _sym: bool) {}
3353
3354 fn commit_update(&mut self, _sym: bool) {
3355 self.counts
3356 .lock()
3357 .expect("counts mutex poisoned")
3358 .commit_label_bits += 1;
3359 }
3360
3361 fn update_history(&mut self, _sym: bool) {
3362 self.counts
3363 .lock()
3364 .expect("counts mutex poisoned")
3365 .history_bits += 1;
3366 }
3367
3368 fn commit_update_history(&mut self, _sym: bool) {
3369 self.counts
3370 .lock()
3371 .expect("counts mutex poisoned")
3372 .commit_history_bits += 1;
3373 }
3374
3375 fn revert(&mut self) {}
3376
3377 fn pop_history(&mut self) {}
3378
3379 fn predict_prob(&mut self, sym: bool) -> f64 {
3380 if sym { 0.75 } else { 0.25 }
3381 }
3382
3383 fn model_name(&self) -> String {
3384 "PhaseUpdateCountingPredictor".to_string()
3385 }
3386
3387 fn boxed_clone(&self) -> Box<dyn Predictor> {
3388 Box::new(self.clone())
3389 }
3390 }
3391
3392 #[test]
3393 fn warmstart_offline_phase_stream_updates_one_phase_per_closed_horizon_window() {
3394 let counts = (0..3)
3395 .map(|_| Arc::new(Mutex::new(PhaseUpdateCounts::default())))
3396 .collect::<Vec<_>>();
3397 let cfg = WarmStartExactJhRuntimeConfig {
3398 task_fingerprint: TaskFingerprint::parse_hex(TEST_TASK_FINGERPRINT_HEX)
3399 .expect("test fingerprint"),
3400 observation_bits: 2,
3401 observation_stream_len: 1,
3402 observation_key_mode: "full_stream",
3403 reward_bits: 2,
3404 agent_actions: action_alphabet(2),
3405 min_reward: 0,
3406 max_reward: 3,
3407 reward_offset: 0,
3408 return_horizon: 2,
3409 return_bins: 7,
3410 label_phase_period: 3,
3411 planner_simulations_per_step: 1,
3412 random_seed: 11,
3413 provenance_policy: TeacherProvenancePolicy::StandalonePlannerRun,
3414 };
3415 let teacher = WarmStartExactJhTeacherDataset {
3416 contract: WarmStartExactJhTeacherContract {
3417 schema_version: 1,
3418 task_fingerprint: TaskFingerprint::parse_hex(TEST_TASK_FINGERPRINT_HEX)
3419 .expect("test fingerprint"),
3420 action_alphabet_size: 2,
3421 observation_bits: 2,
3422 observation_stream_len: 1,
3423 observation_key_mode: "full_stream".to_string(),
3424 observation_adapter_spec_ref: "test-observation-adapter".to_string(),
3425 observation_adapter_content_crc32: "test-observation-adapter-crc32".to_string(),
3426 reward_bits: 2,
3427 return_horizon: 2,
3428 label_phase_period: 3,
3429 scalar_representation: "test-scalar".to_string(),
3430 exact_reward_encoding_certificate: "test-cert".to_string(),
3431 },
3432 traces: vec![WarmStartExactJhTeacherTrace {
3433 transitions: vec![
3434 WarmStartExactJhTransition {
3435 action: 0,
3436 observations: vec![1],
3437 reward: 1,
3438 },
3439 WarmStartExactJhTransition {
3440 action: 1,
3441 observations: vec![2],
3442 reward: 2,
3443 },
3444 WarmStartExactJhTransition {
3445 action: 0,
3446 observations: vec![3],
3447 reward: 0,
3448 },
3449 ],
3450 }],
3451 };
3452 let mut agent = WarmStartExactJhAgent {
3453 config: cfg,
3454 phases: counts
3455 .iter()
3456 .map(|counts| PhaseModel {
3457 predictor: Box::new(PhaseUpdateCountingPredictor {
3458 counts: counts.clone(),
3459 }),
3460 last_augmented_step: 0,
3461 })
3462 .collect(),
3463 steps: Vec::new(),
3464 return_labels_by_step: Vec::new(),
3465 total_steps_observed: 0,
3466 action_bits: 1,
3467 return_label_codec: ReturnLabelCodec::value_monotone(7),
3468 teacher_label_count: 0,
3469 rng: RandomGenerator::from_seed(11),
3470 };
3471
3472 agent
3473 .warm_start_from_teacher(&teacher)
3474 .expect("offline teacher trace should warm-start");
3475
3476 let snapshots = counts
3477 .iter()
3478 .map(|counts| counts.lock().expect("counts mutex poisoned").clone())
3479 .collect::<Vec<_>>();
3480 assert_eq!(agent.teacher_label_count(), 2);
3481 assert_eq!(agent.steps_observed(), 0);
3482 assert!(agent.same_task_live_trace().is_none());
3483 assert_eq!(
3484 snapshots,
3485 vec![
3486 PhaseUpdateCounts {
3487 commit_label_bits: 0,
3488 commit_history_bits: 15,
3489 history_bits: 0,
3490 },
3491 PhaseUpdateCounts {
3492 commit_label_bits: 3,
3493 commit_history_bits: 15,
3494 history_bits: 0,
3495 },
3496 PhaseUpdateCounts {
3497 commit_label_bits: 3,
3498 commit_history_bits: 15,
3499 history_bits: 0,
3500 },
3501 ]
3502 );
3503 }
3504
3505 #[test]
3506 fn warmstart_action_values_propagate_history_encoding_errors_without_mutation() {
3507 let counts = Arc::new(Mutex::new(PhaseUpdateCounts::default()));
3508 let cfg = WarmStartExactJhRuntimeConfig {
3509 task_fingerprint: TaskFingerprint::parse_hex(TEST_TASK_FINGERPRINT_HEX)
3510 .expect("test fingerprint"),
3511 observation_bits: 2,
3512 observation_stream_len: 1,
3513 observation_key_mode: "full_stream",
3514 reward_bits: 1,
3515 agent_actions: action_alphabet(2),
3516 min_reward: 0,
3517 max_reward: 1,
3518 reward_offset: 0,
3519 return_horizon: 1,
3520 return_bins: 2,
3521 label_phase_period: 1,
3522 planner_simulations_per_step: 1,
3523 random_seed: 11,
3524 provenance_policy: TeacherProvenancePolicy::StandalonePlannerRun,
3525 };
3526 let mut agent = WarmStartExactJhAgent {
3527 config: cfg,
3528 phases: vec![PhaseModel {
3529 predictor: Box::new(PhaseUpdateCountingPredictor {
3530 counts: counts.clone(),
3531 }),
3532 last_augmented_step: 0,
3533 }],
3534 steps: vec![StepRecord {
3535 action: 0,
3536 observations: vec![0],
3537 reward: 2,
3538 }],
3539 return_labels_by_step: vec![Some(0)],
3540 total_steps_observed: 1,
3541 action_bits: 1,
3542 return_label_codec: ReturnLabelCodec::value_monotone(2),
3543 teacher_label_count: 0,
3544 rng: RandomGenerator::from_seed(11),
3545 };
3546
3547 let err = agent
3548 .estimate_action_values()
3549 .expect_err("retained invalid reward must be reported");
3550
3551 assert!(matches!(err, WarmStartExactJhError::RewardEncoding(_)));
3552 assert_eq!(
3553 *counts.lock().expect("counts mutex poisoned"),
3554 PhaseUpdateCounts::default()
3555 );
3556 }
3557
3558 #[cfg(feature = "backend-ctw")]
3559 #[test]
3560 fn standalone_warmstart_teacher_contract_matches_compiled_validation() {
3561 let cfg = config();
3562 let compiled = cfg.compile_planner_run_spec().expect("compile planner run");
3563 let contract = standalone_warmstart_teacher_contract_for_compiled_planner_run(&compiled)
3564 .expect("standalone contract");
3565 validate_warmstart_teacher_against_compiled_planner_run(&compiled, &contract)
3566 .expect("standalone contract must validate against compiled planner run");
3567 }
3568
3569 #[test]
3570 fn merge_warmstart_teacher_traces_preserves_deterministic_order_and_dedups() {
3571 let middle = WarmStartExactJhTeacherTrace {
3572 transitions: vec![WarmStartExactJhTransition {
3573 action: 0,
3574 observations: vec![2],
3575 reward: 0,
3576 }],
3577 };
3578 let mut traces = vec![middle.clone()];
3579 let high = WarmStartExactJhTeacherTrace {
3580 transitions: vec![WarmStartExactJhTransition {
3581 action: 1,
3582 observations: vec![2],
3583 reward: 3,
3584 }],
3585 };
3586 let low = WarmStartExactJhTeacherTrace {
3587 transitions: vec![WarmStartExactJhTransition {
3588 action: 0,
3589 observations: vec![1],
3590 reward: 0,
3591 }],
3592 };
3593 let (inserted, _) =
3594 merge_warmstart_teacher_traces_deterministic(&mut traces, vec![high.clone(), low]);
3595 assert_eq!(inserted, 2);
3596 assert!(!merge_warmstart_teacher_trace_deterministic(
3597 &mut traces,
3598 high
3599 ));
3600 assert!(!merge_warmstart_teacher_trace_deterministic(
3601 &mut traces,
3602 middle
3603 ));
3604 assert_eq!(traces.len(), 3);
3605 assert_eq!(traces[0].transitions[0].action, 0);
3606 assert_eq!(traces[0].transitions[0].observations, vec![1]);
3607 assert_eq!(traces[1].transitions[0].observations, vec![2]);
3608 assert_eq!(traces[2].transitions[0].action, 1);
3609 }
3610
3611 #[test]
3612 fn teacher_dataset_new_canonicalizes_trace_order_and_dedups() {
3613 let high = WarmStartExactJhTeacherTrace {
3614 transitions: vec![WarmStartExactJhTransition {
3615 action: 1,
3616 observations: vec![2],
3617 reward: 3,
3618 }],
3619 };
3620 let low = WarmStartExactJhTeacherTrace {
3621 transitions: vec![WarmStartExactJhTransition {
3622 action: 0,
3623 observations: vec![1],
3624 reward: 0,
3625 }],
3626 };
3627 let dataset = WarmStartExactJhTeacherDataset::new(
3628 teacher_contract(),
3629 vec![high.clone(), low.clone(), high.clone()],
3630 );
3631 assert_eq!(dataset.traces, vec![low, high]);
3632 }
3633
3634 #[test]
3635 fn json_teacher_dataset_requires_same_task_traces() {
3636 let value = serde_json::json!({
3637 "schema_version": 1,
3638 "contract": {
3639 "task_fingerprint": TEST_TASK_FINGERPRINT_HEX,
3640 "action_alphabet_size": 2,
3641 "observation_bits": 2,
3642 "observation_stream_len": 1,
3643 "observation_key_mode": "full_stream",
3644 "observation_adapter_spec_ref": "test-observation-adapter",
3645 "observation_adapter_content_crc32": "test-observation-adapter-crc32",
3646 "reward_bits": 2,
3647 "return_horizon": 1,
3648 "label_phase_period": 1,
3649 "scalar_representation": "test-scalar",
3650 "exact_reward_encoding_certificate": "test-cert"
3651 },
3652 "traces": [{
3653 "transitions": [{"action": 1, "observations": [2], "reward": 3}]
3654 }]
3655 });
3656 let parsed = WarmStartExactJhTeacherDataset::from_json_value(&value)
3657 .expect("teacher trace should parse");
3658 assert_eq!(parsed.traces.len(), 1);
3659 assert_eq!(parsed.traces[0].transitions[0].action, 1);
3660 }
3661
3662 #[test]
3663 fn json_teacher_dataset_rejects_malformed_task_fingerprint() {
3664 let value = serde_json::json!({
3665 "schema_version": 1,
3666 "contract": {
3667 "task_fingerprint": "not-a-fingerprint",
3668 "action_alphabet_size": 2,
3669 "observation_bits": 2,
3670 "observation_stream_len": 1,
3671 "observation_key_mode": "full_stream",
3672 "observation_adapter_spec_ref": "test-observation-adapter",
3673 "observation_adapter_content_crc32": "test-observation-adapter-crc32",
3674 "reward_bits": 2,
3675 "return_horizon": 1,
3676 "label_phase_period": 1,
3677 "scalar_representation": "test-scalar",
3678 "exact_reward_encoding_certificate": "test-cert"
3679 },
3680 "traces": [{
3681 "transitions": [{"action": 1, "observations": [2], "reward": 3}]
3682 }]
3683 });
3684 let err = WarmStartExactJhTeacherDataset::from_json_value(&value)
3685 .expect_err("malformed task fingerprint must fail at parse time");
3686 assert!(
3687 err.to_string()
3688 .contains("must be a 64-digit lowercase hexadecimal SHA-256 digest"),
3689 "{err}"
3690 );
3691 }
3692
3693 #[test]
3694 fn json_teacher_dataset_rejects_legacy_trace_and_observation_aliases() {
3695 let mut value = serde_json::json!({
3696 "schema_version": 1,
3697 "contract": {
3698 "task_fingerprint": TEST_TASK_FINGERPRINT_HEX,
3699 "action_alphabet_size": 2,
3700 "observation_bits": 2,
3701 "observation_stream_len": 1,
3702 "observation_key_mode": "full_stream",
3703 "observation_adapter_spec_ref": "test-observation-adapter",
3704 "observation_adapter_content_crc32": "test-observation-adapter-crc32",
3705 "reward_bits": 2,
3706 "return_horizon": 1,
3707 "label_phase_period": 1,
3708 "scalar_representation": "test-scalar",
3709 "exact_reward_encoding_certificate": "test-cert"
3710 },
3711 "traces": [[{"action": 1, "observations": [2], "reward": 3}]]
3712 });
3713 let err = WarmStartExactJhTeacherDataset::from_json_value(&value)
3714 .expect_err("bare trace arrays must be rejected");
3715 assert!(
3716 err.to_string()
3717 .contains("must be an object with transitions")
3718 );
3719
3720 value["traces"] = serde_json::json!([{
3721 "transitions": [{"action": 1, "obs": [2], "reward": 3}]
3722 }]);
3723 let err = WarmStartExactJhTeacherDataset::from_json_value(&value)
3724 .expect_err("obs alias must be rejected");
3725 assert!(err.to_string().contains("unknown teacher field 'obs'"));
3726
3727 let mut top_extra = value.clone();
3728 top_extra["traces"] = serde_json::json!([{
3729 "transitions": [{"action": 1, "observations": [2], "reward": 3}]
3730 }]);
3731 top_extra["extra"] = serde_json::json!(true);
3732 let err = WarmStartExactJhTeacherDataset::from_json_value(&top_extra)
3733 .expect_err("top-level unknown fields must be rejected");
3734 assert!(err.to_string().contains("unknown teacher field 'extra'"));
3735
3736 let mut contract_extra = top_extra;
3737 contract_extra
3738 .as_object_mut()
3739 .expect("object")
3740 .remove("extra");
3741 contract_extra["contract"]["extra"] = serde_json::json!(true);
3742 let err = WarmStartExactJhTeacherDataset::from_json_value(&contract_extra)
3743 .expect_err("contract unknown fields must be rejected");
3744 assert!(err.to_string().contains("unknown teacher field 'extra'"));
3745 }
3746
3747 #[test]
3748 fn teacher_dataset_slice_parser_and_label_count_cover_horizon_windows() {
3749 let err = WarmStartExactJhTeacherDataset::from_json_slice(b"{")
3750 .expect_err("invalid json must be rejected");
3751 assert!(err.to_string().contains("invalid teacher JSON"), "{err}");
3752
3753 let value = serde_json::json!({
3754 "schema_version": 1,
3755 "contract": {
3756 "task_fingerprint": TEST_TASK_FINGERPRINT_HEX,
3757 "action_alphabet_size": 2,
3758 "observation_bits": 2,
3759 "observation_stream_len": 1,
3760 "observation_key_mode": "full_stream",
3761 "observation_adapter_spec_ref": "test-observation-adapter",
3762 "observation_adapter_content_crc32": "test-observation-adapter-crc32",
3763 "reward_bits": 2,
3764 "return_horizon": 2,
3765 "label_phase_period": 2,
3766 "scalar_representation": "test-scalar",
3767 "exact_reward_encoding_certificate": "test-cert"
3768 },
3769 "traces": [
3770 {"transitions": [
3771 {"action": 0, "observations": [1], "reward": 0},
3772 {"action": 1, "observations": [2], "reward": 3},
3773 {"action": 1, "observations": [2], "reward": 3}
3774 ]},
3775 {"transitions": [
3776 {"action": 0, "observations": [0], "reward": 1}
3777 ]}
3778 ]
3779 });
3780 let bytes = serde_json::to_vec(&value).expect("teacher json");
3781 let parsed = WarmStartExactJhTeacherDataset::from_json_slice(&bytes)
3782 .expect("teacher dataset should parse from slice");
3783
3784 assert_eq!(parsed.label_count_for_horizon(0), 0);
3785 assert_eq!(parsed.label_count_for_horizon(1), 4);
3786 assert_eq!(parsed.label_count_for_horizon(2), 2);
3787 assert_eq!(parsed.label_count_for_horizon(4), 0);
3788 }
3789
3790 #[test]
3791 fn config_validation_reports_local_contract_errors_before_backend_use() {
3792 let mut cfg = config();
3793 cfg.return_horizon = 0;
3794 assert!(matches!(
3795 cfg.validate(),
3796 Err(WarmStartExactJhError::ReturnHorizonZero)
3797 ));
3798
3799 let mut cfg = config();
3800 cfg.return_bins = 0;
3801 assert!(matches!(
3802 cfg.validate(),
3803 Err(WarmStartExactJhError::ReturnBinsZero)
3804 ));
3805
3806 let mut cfg = config();
3807 cfg.return_horizon = 2;
3808 cfg.label_phase_period = 1;
3809 assert!(matches!(
3810 cfg.validate(),
3811 Err(WarmStartExactJhError::LabelPhasePeriodTooShort { .. })
3812 ));
3813
3814 let mut cfg = config();
3815 cfg.planner_simulations_per_step = 0;
3816 assert!(matches!(
3817 cfg.validate(),
3818 Err(WarmStartExactJhError::PlannerSimulationsZero)
3819 ));
3820
3821 let mut cfg = config();
3822 cfg.planner_simulations_per_step = 2;
3823 assert!(matches!(
3824 cfg.validate(),
3825 Err(WarmStartExactJhError::PlannerSimulationsUnsupported { configured: 2 })
3826 ));
3827
3828 let mut cfg = config();
3829 cfg.return_horizon = 4;
3830 cfg.return_bins = 8;
3831 cfg.label_phase_period = 4;
3832 assert!(matches!(
3833 cfg.validate(),
3834 Err(WarmStartExactJhError::ReturnBinsNotExactHorizon {
3835 return_bins: 8,
3836 return_horizon: 4
3837 })
3838 ));
3839
3840 let mut cfg = config();
3841 cfg.reward_bits = 1;
3842 assert!(matches!(
3843 cfg.validate(),
3844 Err(WarmStartExactJhError::RewardEncoding(_))
3845 ));
3846 }
3847
3848 #[cfg(feature = "backend-ctw")]
3849 #[test]
3850 fn warmstart_agent_rejects_teacher_transitions_outside_interface_contract() {
3851 let mut invalid = teacher();
3852 invalid.traces[0].transitions[0].action = 2;
3853 let err = match WarmStartExactJhAgent::new(config(), invalid) {
3854 Ok(_) => panic!("teacher action outside alphabet must fail"),
3855 Err(err) => err,
3856 };
3857 assert!(matches!(
3858 err,
3859 WarmStartExactJhError::ActionOutOfRange { .. }
3860 ));
3861
3862 let mut invalid = teacher();
3863 invalid.traces[0].transitions[0].observations = vec![1, 2];
3864 let err = match WarmStartExactJhAgent::new(config(), invalid) {
3865 Ok(_) => panic!("teacher observation stream length must fail"),
3866 Err(err) => err,
3867 };
3868 assert!(matches!(
3869 err,
3870 WarmStartExactJhError::ObservationStreamLengthMismatch { .. }
3871 ));
3872
3873 let mut invalid = teacher();
3874 invalid.traces[0].transitions[0].observations = vec![4];
3875 let err = match WarmStartExactJhAgent::new(config(), invalid) {
3876 Ok(_) => panic!("teacher observation value must fail"),
3877 Err(err) => err,
3878 };
3879 assert!(matches!(
3880 err,
3881 WarmStartExactJhError::ObservationValueOutOfRange { .. }
3882 ));
3883
3884 let mut invalid = teacher();
3885 invalid.traces[0].transitions[0].reward = 4;
3886 let err = match WarmStartExactJhAgent::new(config(), invalid) {
3887 Ok(_) => panic!("teacher reward outside range must fail"),
3888 Err(err) => err,
3889 };
3890 assert!(matches!(
3891 err,
3892 WarmStartExactJhError::RewardOutOfRange { .. }
3893 ));
3894 }
3895
3896 #[cfg(feature = "backend-ctw")]
3897 #[test]
3898 fn warmstart_agent_rejects_teacher_contract_task_fingerprint_mismatch() {
3899 let mut invalid = teacher();
3900 invalid.contract.task_fingerprint =
3901 TaskFingerprint::parse_hex(MISMATCH_TASK_FINGERPRINT_HEX)
3902 .expect("valid mismatch fingerprint");
3903 let err = match WarmStartExactJhAgent::new(config(), invalid) {
3904 Ok(_) => panic!("teacher task fingerprint mismatch must fail"),
3905 Err(err) => err,
3906 };
3907 assert!(matches!(
3908 err,
3909 WarmStartExactJhError::InvalidTeacherDataset { .. }
3910 ));
3911 assert!(err.to_string().contains("task_fingerprint"), "{err}");
3912 }
3913
3914 #[cfg(feature = "backend-ctw")]
3915 #[test]
3916 fn warmstart_agent_rejects_teacher_contract_interface_mismatch() {
3917 let mut invalid = teacher();
3918 invalid.contract.action_alphabet_size = 3;
3919 let err = match WarmStartExactJhAgent::new(config(), invalid) {
3920 Ok(_) => panic!("teacher contract interface mismatch must fail"),
3921 Err(err) => err,
3922 };
3923 assert!(matches!(
3924 err,
3925 WarmStartExactJhError::InvalidTeacherDataset { .. }
3926 ));
3927 assert!(err.to_string().contains("action_alphabet_size"), "{err}");
3928 }
3929
3930 #[cfg(feature = "backend-ctw")]
3931 #[test]
3932 fn warmstart_agent_rejects_teacher_contract_return_horizon_mismatch() {
3933 let mut invalid = teacher();
3934 invalid.contract.return_horizon = 2;
3935 let err = match WarmStartExactJhAgent::new(config(), invalid) {
3936 Ok(_) => panic!("teacher contract return_horizon mismatch must fail"),
3937 Err(err) => err,
3938 };
3939 assert!(matches!(
3940 err,
3941 WarmStartExactJhError::InvalidTeacherDataset { .. }
3942 ));
3943 assert!(err.to_string().contains("return_horizon"), "{err}");
3944 }
3945
3946 #[cfg(feature = "backend-ctw")]
3947 #[test]
3948 fn warmstart_agent_rejects_teacher_contract_label_phase_period_mismatch() {
3949 let mut invalid = teacher();
3950 invalid.contract.label_phase_period = 2;
3951 let err = match WarmStartExactJhAgent::new(config(), invalid) {
3952 Ok(_) => panic!("teacher contract label_phase_period mismatch must fail"),
3953 Err(err) => err,
3954 };
3955 assert!(matches!(
3956 err,
3957 WarmStartExactJhError::InvalidTeacherDataset { .. }
3958 ));
3959 assert!(err.to_string().contains("label_phase_period"), "{err}");
3960 }
3961
3962 #[cfg(feature = "backend-ctw")]
3963 #[test]
3964 fn warmstart_agent_rejects_teacher_contract_observation_key_mode_mismatch() {
3965 let mut invalid = teacher();
3966 invalid.contract.observation_key_mode = "definitely-not-a-mode".to_string();
3967 let err = match WarmStartExactJhAgent::new(config(), invalid) {
3968 Ok(_) => panic!("teacher contract observation_key_mode mismatch must fail"),
3969 Err(err) => err,
3970 };
3971 assert!(matches!(
3972 err,
3973 WarmStartExactJhError::InvalidTeacherDataset { .. }
3974 ));
3975 assert!(err.to_string().contains("observation_key_mode"), "{err}");
3976 }
3977
3978 #[cfg(feature = "backend-ctw")]
3979 #[test]
3980 fn warmstart_agent_rejects_teacher_contract_scalar_provenance_mismatch() {
3981 let mut invalid = teacher();
3982 invalid.contract.scalar_representation = "different-scalar".to_string();
3983 let err = match WarmStartExactJhAgent::new(config(), invalid) {
3984 Ok(_) => panic!("teacher contract scalar provenance mismatch must fail"),
3985 Err(err) => err,
3986 };
3987 assert!(matches!(
3988 err,
3989 WarmStartExactJhError::InvalidTeacherDataset { .. }
3990 ));
3991 assert!(err.to_string().contains("scalar_representation"), "{err}");
3992 }
3993
3994 #[cfg(feature = "backend-ctw")]
3995 #[test]
3996 fn warmstart_agent_delays_live_trace_until_complete_return_horizon() {
3997 let mut cfg = config();
3998 cfg.return_horizon = 2;
3999 cfg.return_bins = 7;
4000 cfg.label_phase_period = 2;
4001 cfg.random_seed = Some(16);
4002 let mut agent = WarmStartExactJhAgent::new(cfg.clone(), teacher_for_config(&cfg))
4003 .expect("warmstart agent should initialize");
4004
4005 assert_eq!(agent.teacher_label_count(), 2);
4006 assert_eq!(agent.num_actions(), action_alphabet(2));
4007 assert_eq!(agent.planner_simulations_per_step(), 1);
4008 assert_eq!(agent.resolved_random_seed(), 16);
4009 assert!(agent.same_task_live_trace().is_none());
4010
4011 agent
4012 .observe_transition(0, &[1], 1)
4013 .expect("first live transition");
4014 assert!(agent.same_task_live_trace().is_none());
4015
4016 agent
4017 .observe_transition(1, &[2], 2)
4018 .expect("second live transition");
4019 let live = agent
4020 .same_task_live_trace()
4021 .expect("complete live trace should be available");
4022 assert_eq!(live.transitions.len(), 2);
4023 assert_eq!(live.transitions[0].action, 0);
4024 assert_eq!(live.transitions[1].reward, 2);
4025
4026 let greedy = agent.get_planned_action();
4027 let first_exploratory = agent.get_planned_action_with_extra_exploration(1.0);
4028 let second_exploratory = agent.get_planned_action_with_extra_exploration(1.0);
4029 assert_ne!(
4030 first_exploratory, second_exploratory,
4031 "test seed must make forced exploration distinguishable from a fixed action"
4032 );
4033 assert!(first_exploratory != greedy || second_exploratory != greedy);
4034 }
4035
4036 #[cfg(feature = "backend-ctw")]
4037 #[test]
4038 fn warmstart_bytepacked_ctw_teacher_replay_and_live_step() {
4039 use crate::api::BitOrder;
4040
4041 let cfg = WarmStartExactJhConfig {
4042 rate_backend: RateBackend::Ctw { depth: 8 },
4043 bit_stream_semantics: BitStreamSemantics::BytePacked {
4044 order: BitOrder::MsbFirst,
4045 },
4046 observation_bits: 8,
4047 observation_stream_len: 1,
4048 reward_bits: 8,
4049 agent_actions: action_alphabet(256),
4050 return_horizon: 1,
4051 return_bins: 256,
4052 label_phase_period: 1,
4053 planner_simulations_per_step: 1,
4054 random_seed: Some(42),
4055 };
4056 let teacher = WarmStartExactJhTeacherDataset {
4057 contract: teacher_for_config(&cfg).contract,
4058 traces: vec![WarmStartExactJhTeacherTrace {
4059 transitions: vec![WarmStartExactJhTransition {
4060 action: 0,
4061 observations: vec![5],
4062 reward: 17,
4063 }],
4064 }],
4065 };
4066 let mut agent_a = WarmStartExactJhAgent::new(cfg.clone(), teacher.clone())
4067 .expect("byte-packed warmstart agent");
4068 let mut agent_b =
4069 WarmStartExactJhAgent::new(cfg, teacher).expect("byte-packed warmstart replay agent");
4070 assert_eq!(agent_a.teacher_label_count(), 1);
4071 let planned_a = agent_a.get_planned_action();
4072 let planned_b = agent_b.get_planned_action();
4073 assert_eq!(
4074 planned_a, planned_b,
4075 "byte-packed warmstart planning must be deterministic under identical seed"
4076 );
4077 assert!(planned_a < 256);
4078 agent_a
4079 .observe_transition(planned_a, &[3], 10)
4080 .expect("byte-packed live transition");
4081 assert_eq!(agent_a.steps_observed(), 1);
4082 }
4083
4084 #[cfg(feature = "backend-ctw")]
4085 #[test]
4086 fn warmstart_binarytokens_fac_ctw_teacher_replay_and_live_step() {
4087 let cfg = WarmStartExactJhConfig {
4088 rate_backend: RateBackend::FacCtw {
4089 base_depth: 8,
4090 num_percept_bits: 8,
4091 encoding_bits: 8,
4092 msb_first: Some(true),
4093 },
4094 bit_stream_semantics: BitStreamSemantics::BinaryTokens,
4095 observation_bits: 2,
4096 observation_stream_len: 1,
4097 reward_bits: 2,
4098 agent_actions: action_alphabet(2),
4099 return_horizon: 1,
4100 return_bins: 4,
4101 label_phase_period: 1,
4102 planner_simulations_per_step: 1,
4103 random_seed: Some(7),
4104 };
4105 let teacher = teacher_for_config(&cfg);
4106 let mut agent_a = WarmStartExactJhAgent::new(cfg.clone(), teacher.clone())
4107 .expect("BinaryTokens FAC-CTW warmstart agent");
4108 let mut agent_b = WarmStartExactJhAgent::new(cfg, teacher)
4109 .expect("BinaryTokens FAC-CTW warmstart replay agent");
4110 assert_eq!(agent_a.teacher_label_count(), 3);
4111 let planned_a = agent_a.get_planned_action();
4112 let planned_b = agent_b.get_planned_action();
4113 assert_eq!(
4114 planned_a, planned_b,
4115 "BinaryTokens FAC-CTW warmstart planning must be deterministic under identical seed"
4116 );
4117 agent_a
4118 .observe_transition(planned_a, &[1], 1)
4119 .expect("BinaryTokens FAC-CTW live transition");
4120 assert_eq!(agent_a.steps_observed(), 1);
4121 }
4122
4123 #[cfg(feature = "backend-ctw")]
4124 #[test]
4125 fn warmstart_agent_learns_teacher_labels_and_observes_live_steps() {
4126 let mut agent = WarmStartExactJhAgent::new(config(), teacher())
4127 .expect("warmstart agent should initialize");
4128 assert_eq!(agent.teacher_label_count(), 3);
4129 let action = agent.get_planned_action();
4130 assert!(action < 2);
4131 agent
4132 .observe_transition(action, &[1], 1)
4133 .expect("first transition");
4134 assert_eq!(agent.steps_observed(), 1);
4135 }
4136
4137 #[test]
4138 fn validate_rejects_reward_bits_too_narrow_for_derived_instantaneous_bounds() {
4139 let mut cfg = config();
4140 cfg.return_horizon = 1;
4141 cfg.label_phase_period = 1;
4142 cfg.return_bins = 100;
4143 cfg.reward_bits = 1;
4144 let err = cfg
4145 .validate()
4146 .expect_err("derived max instantaneous reward must fit reward_bits");
4147 assert!(matches!(err, WarmStartExactJhError::RewardEncoding(_)));
4148 }
4149
4150 #[derive(Clone, Default)]
4151 struct ResetSpyCounts {
4152 reset_calls: usize,
4153 }
4154
4155 #[derive(Clone)]
4156 struct ResetSpyPredictor {
4157 counts: Arc<Mutex<ResetSpyCounts>>,
4158 }
4159
4160 impl Predictor for ResetSpyPredictor {
4161 fn update(&mut self, _sym: bool) {}
4162
4163 fn update_history(&mut self, _sym: bool) {}
4164
4165 fn revert(&mut self) {}
4166
4167 fn pop_history(&mut self) {}
4168
4169 fn predict_prob(&mut self, sym: bool) -> f64 {
4170 if sym { 0.75 } else { 0.25 }
4171 }
4172
4173 fn model_name(&self) -> String {
4174 "ResetSpyPredictor".to_string()
4175 }
4176
4177 fn boxed_clone(&self) -> Box<dyn Predictor> {
4178 Box::new(self.clone())
4179 }
4180
4181 fn reset_conditioning_history(&mut self) -> Result<(), String> {
4182 self.counts
4183 .lock()
4184 .expect("counts mutex poisoned")
4185 .reset_calls += 1;
4186 Ok(())
4187 }
4188 }
4189
4190 #[test]
4191 fn warmstart_resets_predictor_conditioning_between_teacher_traces() {
4192 let counts = Arc::new(Mutex::new(ResetSpyCounts::default()));
4193 let spy = ResetSpyPredictor {
4194 counts: counts.clone(),
4195 };
4196
4197 let cfg = WarmStartExactJhRuntimeConfig {
4198 task_fingerprint: TaskFingerprint::parse_hex(TEST_TASK_FINGERPRINT_HEX)
4199 .expect("test fingerprint"),
4200 observation_bits: 2,
4201 observation_stream_len: 1,
4202 observation_key_mode: "full_stream",
4203 reward_bits: 2,
4204 agent_actions: action_alphabet(2),
4205 min_reward: 0,
4206 max_reward: 3,
4207 reward_offset: 0,
4208 return_horizon: 1,
4209 return_bins: 4,
4210 label_phase_period: 1,
4211 planner_simulations_per_step: 1,
4212 random_seed: 7,
4213 provenance_policy: TeacherProvenancePolicy::StandalonePlannerRun,
4214 };
4215
4216 let teacher = WarmStartExactJhTeacherDataset {
4217 contract: WarmStartExactJhTeacherContract {
4218 schema_version: 1,
4219 task_fingerprint: TaskFingerprint::parse_hex(TEST_TASK_FINGERPRINT_HEX)
4220 .expect("test fingerprint"),
4221 action_alphabet_size: 2,
4222 observation_bits: 2,
4223 observation_stream_len: 1,
4224 observation_key_mode: "full_stream".to_string(),
4225 observation_adapter_spec_ref: "test-observation-adapter".to_string(),
4226 observation_adapter_content_crc32: "test-observation-adapter-crc32".to_string(),
4227 reward_bits: 2,
4228 return_horizon: 1,
4229 label_phase_period: 1,
4230 scalar_representation: "test-scalar".to_string(),
4231 exact_reward_encoding_certificate: "test-cert".to_string(),
4232 },
4233 traces: vec![
4234 WarmStartExactJhTeacherTrace {
4235 transitions: vec![WarmStartExactJhTransition {
4236 action: 0,
4237 observations: vec![1],
4238 reward: 1,
4239 }],
4240 },
4241 WarmStartExactJhTeacherTrace {
4242 transitions: vec![WarmStartExactJhTransition {
4243 action: 1,
4244 observations: vec![2],
4245 reward: 2,
4246 }],
4247 },
4248 ],
4249 };
4250
4251 let mut agent = WarmStartExactJhAgent {
4252 config: cfg,
4253 phases: vec![PhaseModel {
4254 predictor: Box::new(spy),
4255 last_augmented_step: 0,
4256 }],
4257 steps: Vec::new(),
4258 return_labels_by_step: Vec::new(),
4259 total_steps_observed: 0,
4260 action_bits: 1,
4261 return_label_codec: ReturnLabelCodec::value_monotone(4),
4262 teacher_label_count: 0,
4263 rng: RandomGenerator::from_seed(7),
4264 };
4265
4266 agent
4267 .warm_start_from_teacher(&teacher)
4268 .expect("warm-start should succeed");
4269
4270 let snapshot = counts.lock().expect("counts mutex poisoned").clone();
4271 assert_eq!(snapshot.reset_calls, teacher.traces.len());
4272 assert_eq!(agent.teacher_label_count(), 2);
4273 }
4274}