## Tests for integration.nim (#48) — assert-based, no framework. ## Covers: TrainingMsg channel round-trip (plain arrays through a channel), ## drain-then-train NewBattle/Shutdown handling, trainPass safety below canSample. import arraymancer except Linear import std/locks import SAC_LSTM_Bot/integration import SAC_LSTM_Bot/state # STATE_DIM import SAC_LSTM_Bot/actions # ACTION_DIM import SAC_LSTM_Bot/training # initSACTrainer import SAC_LSTM_Bot/replay_buffer # ── 1. TrainingMsg round-trips through a channel with arrays intact ────────── block: var ch: Channel[TrainingMsg] ch.open(4) var msg = TrainingMsg(kind: tmkTransition) for i in 0 ..< STATE_DIM: msg.state[i] = float32(i) * 0.5'f32 msg.nextState[i] = float32(i) * 2.0'f32 for i in 0 ..< ACTION_DIM: msg.action[i] = float32(i) - 2.0'f32 msg.reward = -1.25'f32 msg.done = true assert ch.trySend(msg) assert ch.trySend(TrainingMsg(kind: tmkNewBattle, enemyId: 4242)) assert ch.trySend(TrainingMsg(kind: tmkShutdown)) let r1 = ch.recv() assert r1.kind == tmkTransition, "first msg is a transition" for i in 0 ..< STATE_DIM: assert r1.state[i] == float32(i) * 0.5'f32, "state round-trip at " & $i assert r1.nextState[i] == float32(i) * 2.0'f32, "nextState round-trip at " & $i for i in 0 ..< ACTION_DIM: assert r1.action[i] == float32(i) - 2.0'f32, "action round-trip at " & $i assert r1.reward == -1.25'f32 and r1.done let r2 = ch.recv() assert r2.kind == tmkNewBattle and r2.enemyId == 4242 let r3 = ch.recv() assert r3.kind == tmkShutdown ch.close() # closed + empty -> tryRecv reports no data (Nim 2.2: recv would block forever) let (ok4, _) = ch.tryRecv() assert not ok4, "closed channel must report dataAvailable=false" echo "PASS TrainingMsg channel round-trip" # ── 2. NewBattle clears only when the opponent changes; Shutdown stops ──────── block: var st: TrainState st.trainer = initSACTrainer(STATE_DIM, ACTION_DIM) st.buf = newReplayBuffer(64, STATE_DIM, ACTION_DIM, burnIn = 2, trainWindow = 3) assert handleTrainingMsg(st, TrainingMsg(kind: tmkNewBattle, enemyId: 7)) assert st.lastEnemyId == 7 for i in 0 ..< 5: var m = TrainingMsg(kind: tmkTransition) m.reward = float32(i) assert handleTrainingMsg(st, m) assert st.buf.len == 5, "transitions stored" # Same opponent -> buffer kept (Q12a). assert handleTrainingMsg(st, TrainingMsg(kind: tmkNewBattle, enemyId: 7)) assert st.buf.len == 5, "same opponent must NOT clear" # Opponent changed -> clear. assert handleTrainingMsg(st, TrainingMsg(kind: tmkNewBattle, enemyId: 8)) assert st.buf.len == 0, "opponent change must clear" assert st.lastEnemyId == 8 # Shutdown stops the caller's loop. assert not handleTrainingMsg(st, TrainingMsg(kind: tmkShutdown)) echo "PASS drain-then-train NewBattle/Shutdown" # ── 3. trainPass is a safe no-op below canSample (no steps, no publish) ─────── block: var st: TrainState st.trainer = initSACTrainer(STATE_DIM, ACTION_DIM) st.buf = newReplayBuffer(64, STATE_DIM, ACTION_DIM, burnIn = 2, trainWindow = 3) for i in 0 ..< 4: assert handleTrainingMsg(st, TrainingMsg(kind: tmkTransition)) trainPass(st, 4) # 4 < burnIn+trainWindow = 5 assert st.stepCount == 0, "no gradient steps below canSample" echo "PASS trainPass no-op below canSample" # ── 4. Flat snapshot layout: pack/unpack round-trips weights exactly ────────── block: let t0 = initSACTrainer(STATE_DIM, ACTION_DIM) let fs = packFull(t0) assert fs.data.len == actorSize(t0.actor.hiddenDim) + 4 * criticSize(t0.actor.hiddenDim) + 1 let (a, c1, c2, tc1, tc2, alpha) = unpackFull(fs) assert a.hiddenDim == t0.actor.hiddenDim assert c1.fc3.b.shape[0] == 1 let fw = a.muHead.w.flatten() let fw0 = t0.actor.muHead.w.flatten() for i in 0 ..< fw.size: assert fw[i] == fw0[i], "actor mu weights round-trip" for i in 0 ..< c2.lstm.bCombined.size: assert c2.lstm.bCombined[i] == t0.critic2.lstm.bCombined[i], "critic lstm bias round-trip" assert alpha == t0.alpha() discard tc1 echo "PASS flat snapshot pack/unpack round-trip"