## 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, json, os] import tankroyale_botapi # updateBotNames: seed the vendored id->name table 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 NAME changes; Shutdown stops ─── block: # Seed the API's id->name table (v1.0.1) for name-based identity (#49). updateBotNames(parseJson( """{"bots":[{"id":7,"name":"Corners"},{"id":8,"name":"Crazy"}]}""")) assert opponentKey(7) == "Corners", "known id resolves to name" assert opponentKey(99) == "99", "unknown id falls back to numeric string" 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.lastEnemyKey == "Corners" 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.lastEnemyKey == "Crazy" # Nameless window (pre-BotListUpdate) is a distinct key. assert handleTrainingMsg(st, TrainingMsg(kind: tmkNewBattle, enemyId: 99)) assert st.lastEnemyKey == "99" and st.buf.len == 0 # Shutdown stops the caller's loop. assert not handleTrainingMsg(st, TrainingMsg(kind: tmkShutdown)) echo "PASS name-keyed NewBattle clear / Shutdown" # ── 2b. bumpRoundCounter increments per call, cold-starts at 1 ──────────────── block: let tmp = getTempDir() / "sac_test_weights_" & $getCurrentProcessId() putEnv("SACLSTM_WEIGHTS_PATH", tmp / "sac_latest.zip") bumpRoundCounter() bumpRoundCounter() assert readFile(tmp / "round_counter.txt") == "2", "counter increments per round" echo "PASS bumpRoundCounter" # ── 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" # ── 5. Lever 3 (#59): metricsLine emits exactly the exposed trainer scalars ─── block: let line = metricsLine(1787394115.123, 42, 500, 20, 20, SACMetrics(criticLoss: 0.5'f32, actorLoss: -1.5'f32, alphaLoss: 0.25'f32, alpha: 2.0'f32)) let j = parseJson(line) # throws on malformed JSONL assert j["steps"].getInt() == 42 and j["buffer_size"].getInt() == 500 assert j["drained"].getInt() == 20 and j["grad_steps"].getInt() == 20 assert abs(j["critic_loss"].getFloat() - 0.5) < 1e-3 assert abs(j["actor_loss"].getFloat() + 1.5) < 1e-3 assert abs(j["alpha_loss"].getFloat() - 0.25) < 1e-3 assert abs(j["alpha"].getFloat() - 2.0) < 1e-3 assert j["epoch"].getFloat() > 1e9 echo "PASS metricsLine JSONL scalars" # ── 6. Lever 4 (#59): SACLSTM_EVAL_MODE=1 suppresses training input ─────────── block: putEnv("SACLSTM_EVAL_MODE", "1") # Gate fires before any channel traffic: false = dropped, nothing enqueued. assert not sendTrainingMsg(TrainingMsg(kind: tmkTransition)), "eval mode must drop transitions" assert not sendTrainingMsg(TrainingMsg(kind: tmkNewBattle, enemyId: 7)), "eval mode must drop NewBattle (no buffer clears from eval)" delEnv("SACLSTM_EVAL_MODE") echo "PASS eval-mode training-input suppression"