Skip to main content

infotheory/aixi/
warmstart.rs

1//! Warm-start exact finite-horizon objective controller for AIXI-family runs.
2
3use 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/// One observed environment transition in a same-task warm-start trace.
35#[derive(Clone, Debug, Eq, PartialEq)]
36#[non_exhaustive]
37pub struct WarmStartExactJhTransition {
38    /// Action selected by the teacher/controller.
39    pub action: Action,
40    /// Observation stream emitted after the action.
41    pub observations: Vec<PerceptVal>,
42    /// Exact integer reward emitted after the action.
43    pub reward: Reward,
44}
45
46impl WarmStartExactJhTransition {
47    /// Construct one warm-start teacher transition.
48    pub fn new(action: Action, observations: Vec<PerceptVal>, reward: Reward) -> Self {
49        Self {
50            action,
51            observations,
52            reward,
53        }
54    }
55}
56
57/// Same-task trace used to initialize a warm-start exact-J_H controller.
58#[derive(Clone, Debug, Eq, PartialEq)]
59#[non_exhaustive]
60pub struct WarmStartExactJhTeacherTrace {
61    /// Chronological transition sequence.
62    pub transitions: Vec<WarmStartExactJhTransition>,
63}
64
65impl WarmStartExactJhTeacherTrace {
66    /// Construct a same-task teacher trace from chronological transitions.
67    pub fn new(transitions: Vec<WarmStartExactJhTransition>) -> Self {
68        Self { transitions }
69    }
70}
71
72/// Validates standalone planner-run provenance hashes against canonical standalone declarations.
73///
74/// Used by the corresponding method on the private WarmStartExactJhRuntimeConfig
75/// and by `validate_warmstart_teacher_against_compiled_planner_run`.
76pub 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
132/// Expected warm-start teacher contract fields for comparison against a parsed contract.
133pub(crate) struct WarmStartTeacherContractExpectation<'a> {
134    /// Expected schema version.
135    pub schema_version: u64,
136    /// Expected planner task fingerprint.
137    pub task_fingerprint: TaskFingerprint,
138    /// Expected action alphabet size.
139    pub action_alphabet_size: usize,
140    /// Expected observation bit width.
141    pub observation_bits: usize,
142    /// Expected observation stream length.
143    pub observation_stream_len: usize,
144    /// Expected observation key mode label.
145    pub observation_key_mode: &'a str,
146    /// Expected reward bit width.
147    pub reward_bits: usize,
148    /// Expected return horizon.
149    pub return_horizon: usize,
150    /// Expected delayed-label phase period.
151    pub label_phase_period: usize,
152    /// Whether standalone planner-run provenance hashes must match.
153    pub validate_standalone_provenance: bool,
154}
155
156/// Validate a parsed teacher contract against an explicit field expectation.
157pub(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
241/// Validates teacher [`WarmStartExactJhTeacherContract::schema_version`] and
242/// [`WarmStartExactJhTeacherContract::task_fingerprint`] against a compiled planner run.
243///
244/// Used by [`validate_warmstart_teacher_against_compiled_planner_run`] (standalone CLI / assets)
245/// and direct fingerprint probes. The tuner bridge validates complete teacher datasets through
246/// [`validate_warmstart_teacher_dataset_for_compiled_planner_run`], which includes this
247/// fingerprint check via the compiled runtime contract and then validates trace payloads.
248/// On mismatch the reason string includes
249/// `current planner_run '<hex>'` for stable integration-test probing.
250pub 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
278/// Validates a parsed teacher contract against a compiled standalone [`PlannerRunSpec`] (CLI / asset loader).
279///
280/// This is the single authoritative check for filesystem-loaded teachers before runtime construction.
281pub 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
338/// Validate a complete warm-start teacher dataset against a compiled planner run.
339///
340/// This is the authoritative ingestion/export gate for teacher assets. It checks
341/// the same runtime contract used by [`WarmStartExactJhAgent`]: controller
342/// compatibility, task fingerprint, provenance policy, transition bounds, exact
343/// finite-horizon label encodability, and that every trace contributes at least
344/// one complete \(H\)-step label.
345pub 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
403/// Validate one reconstructed teacher transition against a warm-start bridge contract.
404pub 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/// Same-task teacher dataset for [`WarmStartExactJhAgent`].
450#[derive(Clone, Debug, Eq, PartialEq)]
451#[non_exhaustive]
452pub struct WarmStartExactJhTeacherDataset {
453    /// Versioned same-task contract metadata for the teacher traces.
454    pub contract: WarmStartExactJhTeacherContract,
455    /// Canonical teacher traces. Each trace is treated as an independent
456    /// same-task rollout, and dataset construction sorts and deduplicates this
457    /// set by transition content.
458    pub traces: Vec<WarmStartExactJhTeacherTrace>,
459}
460
461impl WarmStartExactJhTeacherDataset {
462    /// Construct a teacher dataset from its contract and trace set.
463    ///
464    /// The supplied traces may be in arbitrary order and may contain
465    /// duplicates; the dataset stores the canonical deterministic set.
466    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/// Versioned same-task contract attached to a warm-start teacher dataset.
476#[derive(Clone, Debug, Eq, PartialEq)]
477#[non_exhaustive]
478pub struct WarmStartExactJhTeacherContract {
479    /// Teacher dataset schema version. Version 1 is the v1 exact-J_H contract.
480    pub schema_version: u64,
481    /// Fingerprint of the exact tuning task that produced the traces.
482    pub task_fingerprint: TaskFingerprint,
483    /// Number of actions in the compiled planner alphabet.
484    pub action_alphabet_size: usize,
485    /// Observation bit width.
486    pub observation_bits: usize,
487    /// Number of observation symbols per step.
488    pub observation_stream_len: usize,
489    /// Observation keying mode used by the planner-visible history.
490    pub observation_key_mode: String,
491    /// Observation adapter declaration reference used to encode raw tuner observations.
492    pub observation_adapter_spec_ref: String,
493    /// Content hash of the concrete observation adapter schema.
494    pub observation_adapter_content_crc32: String,
495    /// Reward bit width.
496    pub reward_bits: usize,
497    /// Return horizon used to compute exact labels.
498    pub return_horizon: usize,
499    /// Delayed-label phase period.
500    pub label_phase_period: usize,
501    /// Scalar representation declaration used by the exact reward encoder.
502    pub scalar_representation: String,
503    /// Hash or ref for the verified exact reward encoder.
504    pub exact_reward_encoding_certificate: String,
505}
506
507impl WarmStartExactJhTeacherDataset {
508    /// Parse a JSON teacher dataset.
509    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    /// Parse a JSON teacher dataset value.
519    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    /// Convert this teacher dataset to its JSON representation.
572    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    /// Count labels constructible from this dataset for the supplied horizon.
581    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
675/// Merge one teacher trace into `traces`, preserving deterministic lexicographic order.
676pub 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
684/// Merge teacher traces into `traces`, returning the number and payload of inserted traces.
685///
686/// Existing and incoming traces are canonicalized as a deterministic set before
687/// insertion, so callers may pass traces in arbitrary order.
688pub 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/// Records normalized warm-start action/percept telemetry into one teacher trace.
706#[derive(Clone, Debug, Default)]
707pub struct WarmStartExactJhTraceRecorder {
708    actions: BTreeMap<usize, Action>,
709    percepts: BTreeMap<usize, (Vec<PerceptVal>, Reward)>,
710}
711
712impl WarmStartExactJhTraceRecorder {
713    /// Construct an empty trace recorder.
714    pub fn new() -> Self {
715        Self::default()
716    }
717
718    /// Record an action at step `step`.
719    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    /// Record a post-action percept at step `step`.
733    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    /// Convert the recorded telemetry into a validated teacher trace.
752    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
796/// Build a normalized JSONL action record.
797pub 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
810/// Build a normalized JSONL percept record.
811pub 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
1095/// Parse a normalized planner JSONL trace into a warm-start teacher trace.
1096pub 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
1180/// Parse a normalized planner JSONL trace from bytes.
1181pub 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
1189/// Parse a normalized planner JSONL trace from a filesystem path.
1190pub 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
1204/// Build the standalone same-task teacher contract for a compiled warm-start planner run.
1205pub 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
1263/// Return the target exact-return horizon for a compiled warm-start planner run.
1264pub 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
1275/// Read a warm-start teacher dataset from a filesystem path.
1276pub 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
1290/// Write a warm-start teacher dataset to a filesystem path.
1291pub 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/// Configuration parameters for the warm-start exact-J_H controller.
1310#[derive(Clone)]
1311#[non_exhaustive]
1312pub struct WarmStartExactJhConfig {
1313    /// Predictive backend used by the return-label model.
1314    pub rate_backend: RateBackend,
1315    /// Number of bits used to encode observations.
1316    pub observation_bits: usize,
1317    /// Number of observation symbols per environment step.
1318    pub observation_stream_len: usize,
1319    /// Number of bits used to encode rewards.
1320    pub reward_bits: usize,
1321    /// Number of valid actions.
1322    pub agent_actions: ActionAlphabet,
1323    /// Exact finite-horizon return length H.
1324    pub return_horizon: usize,
1325    /// Exact return-label alphabet cardinality.
1326    ///
1327    /// For this nonnegative-reward controller the valid shape is
1328    /// `return_horizon * max_reward + 1`; slack labels are rejected because
1329    /// they would not correspond to any exact finite-horizon return.
1330    pub return_bins: usize,
1331    /// Delayed-label phase period.
1332    pub label_phase_period: usize,
1333    /// Canonical direct-evaluator budget marker.
1334    ///
1335    /// Warm-start exact-\(J_H\) uses deterministic full return-law evaluation,
1336    /// not a simulation loop. The only truthful value is therefore `1`; larger
1337    /// values are rejected instead of being silently ignored.
1338    pub planner_simulations_per_step: usize,
1339    /// Bit-stream semantics used to adapt the return-label predictor.
1340    pub bit_stream_semantics: BitStreamSemantics,
1341    /// Optional deterministic RNG seed.
1342    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    /// Validate this configuration.
1446    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    /// Build a canonicalized runtime contract from a compiled planner run.
1479    ///
1480    /// The resulting config captures the exact planner/interface contract that
1481    /// teacher traces are expected to match. This includes the planner task
1482    /// fingerprint plus runtime-visible interface dimensions that must remain
1483    /// aligned with any warm-start dataset.
1484    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    /// Validate that a parsed teacher contract matches the active planner runtime
1566    /// contract for this agent configuration.
1567    ///
1568    /// This comparison is authoritative for schema/contract compatibility and is
1569    /// intentionally strict on planner-task fields that influence trace encoding.
1570    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
1635/// Warm-start exact finite-horizon objective controller.
1636pub 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    /// Construct a new warm-start exact-J_H agent.
1650    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    /// Construct from a compiled planner-run spec and same-task teacher data.
1660    pub fn from_compiled_planner_run(
1661        compiled: &CompiledPlannerRunSpec,
1662        teacher: WarmStartExactJhTeacherDataset,
1663    ) -> Result<Self, WarmStartExactJhError> {
1664        // Validate contract mismatch upfront so invalid metadata fails before costly
1665        // predictor allocation and before trace replay.
1666        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    /// Number of transitions incorporated from live interaction.
1709    pub fn steps_observed(&self) -> usize {
1710        self.total_steps_observed
1711    }
1712
1713    /// Number of warm-start labels incorporated from the teacher dataset.
1714    pub fn teacher_label_count(&self) -> usize {
1715        self.teacher_label_count
1716    }
1717
1718    /// Extract the current live same-task trajectory as a teacher trace.
1719    ///
1720    /// The trace is admissible under the same runtime validator used for
1721    /// teacher datasets because it was produced through `observe_transition`.
1722    ///
1723    /// Cost: despite the `&self` receiver this is a full materialization, not a
1724    /// cheap accessor. It allocates a fresh transition vector and clones every
1725    /// stored observation stream, so a call is `O(steps * observation_stream_len)`
1726    /// in both time and allocated memory. It is intended for occasional
1727    /// refresh/export points; callers in a hot loop should cache the result
1728    /// rather than re-deriving it per step.
1729    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    /// Configured action alphabet cardinality.
1747    pub fn num_actions(&self) -> ActionAlphabet {
1748        self.config.agent_actions
1749    }
1750
1751    /// Canonical direct-evaluator budget marker.
1752    ///
1753    /// This is always `1`; warm-start exact-\(J_H\) has no simulation loop.
1754    pub fn planner_simulations_per_step(&self) -> usize {
1755        self.config.planner_simulations_per_step
1756    }
1757
1758    /// Resolved deterministic seed.
1759    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    /// Select the next greedy action from the current exact-return model.
1769    ///
1770    /// # Panics
1771    ///
1772    /// Panics if retained live history violates the validated warm-start
1773    /// runtime contract. Use [`Self::try_get_planned_action`] to receive that
1774    /// condition as [`WarmStartExactJhError`].
1775    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    /// Fallibly select the next greedy action from the current exact-return model.
1781    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    /// Estimate exact finite-horizon action values at the current decision state.
1787    pub fn estimate_action_values(&mut self) -> Result<Vec<f64>, WarmStartExactJhError> {
1788        self.estimate_q_values()
1789    }
1790
1791    /// Select the next action with optional epsilon exploration.
1792    ///
1793    /// Warm-start exact-\(J_H\) has no baseline exploration parameter; this
1794    /// method's argument is the entire exploration probability.
1795    ///
1796    /// # Panics
1797    ///
1798    /// Panics if greedy planning is reached and retained live history violates
1799    /// the validated warm-start runtime contract. Use
1800    /// [`Self::try_get_planned_action_with_extra_exploration_flag`] to receive
1801    /// that condition as [`WarmStartExactJhError`].
1802    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    /// Select the next action with optional epsilon exploration and return whether exploration fired.
1808    ///
1809    /// Warm-start exact-\(J_H\) has no baseline exploration parameter; this
1810    /// method's argument is the entire exploration probability.
1811    ///
1812    /// # Panics
1813    ///
1814    /// Panics if greedy planning is reached and retained live history violates
1815    /// the validated warm-start runtime contract. Use
1816    /// [`Self::try_get_planned_action_with_extra_exploration_flag`] to receive
1817    /// that condition as [`WarmStartExactJhError`].
1818    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    /// Fallibly select the next action with optional epsilon exploration and return whether exploration fired.
1827    ///
1828    /// Warm-start exact-\(J_H\) has no baseline exploration parameter; this
1829    /// method's argument is the entire exploration probability.
1830    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    /// Record one live environment transition.
1846    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    /// Warm-start from teacher traces (trace payload validation only).
1864    ///
1865    /// Contract-level validation is performed before this method is called in the
1866    /// constructor hot path; this keeps construction cheap on malformed contracts.
1867    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/// Error type for warm-start exact-J_H agent construction and execution.
2057#[derive(Debug)]
2058#[non_exhaustive]
2059pub enum WarmStartExactJhError {
2060    /// The compiled planner controller was not a warm-start exact-J_H controller.
2061    ControllerKindMismatch,
2062    /// The return horizon was zero.
2063    ReturnHorizonZero,
2064    /// The return-label alphabet was empty.
2065    ReturnBinsZero,
2066    /// The label phase period was smaller than the return horizon.
2067    LabelPhasePeriodTooShort {
2068        /// Configured label phase period.
2069        label_phase_period: usize,
2070        /// Configured return horizon.
2071        return_horizon: usize,
2072    },
2073    /// The direct-evaluator budget marker was zero.
2074    PlannerSimulationsZero,
2075    /// The direct-evaluator budget marker was not the canonical value.
2076    PlannerSimulationsUnsupported {
2077        /// Configured unsupported value.
2078        configured: usize,
2079    },
2080    /// The exact return range cannot be represented by `return_bins`.
2081    ReturnBinsTooSmall {
2082        /// Required exact labels.
2083        required: u128,
2084        /// Configured labels.
2085        configured: usize,
2086    },
2087    /// `return_bins` would leave unreachable exact-return labels.
2088    ReturnBinsNotExactHorizon {
2089        /// Configured label count.
2090        return_bins: usize,
2091        /// Configured return horizon.
2092        return_horizon: usize,
2093    },
2094    /// The exact return range overflowed the supported integer domain.
2095    ExactReturnRangeOverflow,
2096    /// The configured reward range is not representable.
2097    RewardEncoding(RewardEncodingError),
2098    /// Invalid rate backend.
2099    InvalidRateBackend(crate::error::InfotheoryError),
2100    /// Unsupported rate backend semantics.
2101    UnsupportedRateBackend {
2102        /// Human-readable reason.
2103        reason: &'static str,
2104    },
2105    /// Spec compilation failed.
2106    Spec(SpecError),
2107    /// Predictor construction failed.
2108    Predictor(PredictorBuildError),
2109    /// Predictor conditioning history reset failed.
2110    PredictorConditioningReset {
2111        /// Human-readable reason.
2112        reason: String,
2113    },
2114    /// Teacher dataset was malformed or semantically inadmissible.
2115    InvalidTeacherDataset {
2116        /// Human-readable reason.
2117        reason: String,
2118    },
2119    /// JSONL telemetry was malformed.
2120    InvalidTelemetry {
2121        /// Human-readable reason.
2122        reason: String,
2123    },
2124    /// Action outside the configured alphabet.
2125    ActionOutOfRange {
2126        /// Invalid action.
2127        action: Action,
2128        /// Configured alphabet.
2129        agent_actions: ActionAlphabet,
2130    },
2131    /// Observation stream length mismatch.
2132    ObservationStreamLengthMismatch {
2133        /// Expected stream length.
2134        expected: usize,
2135        /// Actual stream length.
2136        actual: usize,
2137    },
2138    /// Observation value exceeded its bit width.
2139    ObservationValueOutOfRange {
2140        /// Invalid observation.
2141        observation: PerceptVal,
2142        /// Configured observation bits.
2143        observation_bits: usize,
2144        /// Maximum representable value.
2145        maximum: PerceptVal,
2146    },
2147    /// Reward outside the configured exact reward range.
2148    RewardOutOfRange {
2149        /// Invalid reward.
2150        reward: Reward,
2151        /// Minimum reward.
2152        min_reward: Reward,
2153        /// Maximum reward.
2154        max_reward: Reward,
2155    },
2156    /// Live history index was unavailable.
2157    HistoryIndexOutOfRange {
2158        /// Requested 1-based step index.
2159        global_step: usize,
2160        /// Total observed steps.
2161        total_steps_observed: usize,
2162    },
2163    /// A delayed label was required but absent.
2164    MissingReturnLabel {
2165        /// Step index.
2166        step: usize,
2167        /// Phase index.
2168        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    /// Single source of truth for the test teacher-contract field wiring.
2972    ///
2973    /// This builder performs no backend compilation: the `task_fingerprint` and
2974    /// `observation_key_mode` are supplied by the caller. The CTW-dependent path
2975    /// derives them from a real compiled planner run; the backend-independent
2976    /// path supplies synthetic-but-faithful values so the contract-consuming
2977    /// tests stay runnable under the feature-light `aixi` slice.
2978    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    /// Backend-independent teacher contract for tests that only validate
3009    /// contract-shaped data (JSONL conversion, trace recording, dataset
3010    /// canonicalization) and never compile a planner run. The fingerprint and
3011    /// key mode are fixed: `full_stream` matches the canonical planner
3012    /// interface's hardcoded `ObservationKeyMode::FullStream`.
3013    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}