#!/usr/bin/env python3 from __future__ import annotations import csv import json import math import sys import tempfile import unittest from pathlib import Path from types import SimpleNamespace sys.path.insert(0, str(Path(__file__).resolve().parent)) from trace_utils import ( common_prefix_length, expand_hash_ids, synthetic_tokens, token_block, token_block_identities, ) from remap_hash_blocks import materialize, prompt_token_ids class CharacterTokenizer: def __call__(self, text: str, **_: object) -> dict[str, list[int]]: return {"input_ids": [ord(character) for character in text]} class StrictRemapTest(unittest.TestCase): def test_block_count_is_ceil_input_over_16(self) -> None: for input_tokens in (1, 15, 16, 17, 63, 64, 65, 127, 128, 129): source = list(range(math.ceil(input_tokens / 64))) remapped = expand_hash_ids(source, input_tokens) self.assertEqual(len(remapped), math.ceil(input_tokens / 16)) def test_parent_sequence_remains_prefix(self) -> None: parent_length = 80 child_length = 144 parent_source = [101, 202] child_source = [101, 202, 303] parent = expand_hash_ids(parent_source, parent_length) child = expand_hash_ids(child_source, child_length) self.assertEqual(child[: len(parent)], parent) parent_tokens = synthetic_tokens(parent, parent_length, vocab_size=151936) child_tokens = synthetic_tokens(child, child_length, vocab_size=151936) self.assertEqual(child_tokens[: len(parent_tokens)], parent_tokens) def test_synthetic_tokens_preserve_hash_hit_structure(self) -> None: request_a_source = [11, 22, 33] request_b_source = [11, 22, 44] request_a = expand_hash_ids(request_a_source, 180) request_b = expand_hash_ids(request_b_source, 190) source_hits = common_prefix_length(request_a_source, request_b_source) remapped_hits = common_prefix_length(request_a, request_b) self.assertEqual(remapped_hits, source_hits * 4) token_blocks_a = [tuple(token_block(value, vocab_size=151936)) for value in request_a] token_blocks_b = [tuple(token_block(value, vocab_size=151936)) for value in request_b] self.assertEqual(common_prefix_length(token_blocks_a, token_blocks_b), remapped_hits) self.assertEqual(len(set(token_blocks_a)), len(set(request_a))) runtime_a = token_block_identities( synthetic_tokens(request_a, 180, vocab_size=151936) ) runtime_b = token_block_identities( synthetic_tokens(request_b, 190, vocab_size=151936) ) self.assertEqual(common_prefix_length(runtime_a, runtime_b), remapped_hits) def test_synthetic_runtime_identity_depends_on_parent_prefix(self) -> None: left_content = expand_hash_ids([11, 22], 128) right_content = expand_hash_ids([33, 22], 128) left = token_block_identities( synthetic_tokens(left_content, 128, vocab_size=151936) ) right = token_block_identities( synthetic_tokens(right_content, 128, vocab_size=151936) ) self.assertNotEqual(left[4], right[4]) def test_mapping_rejects_non_strict_source_block_count(self) -> None: with self.assertRaisesRegex(ValueError, "strict 64-block contract"): expand_hash_ids([1], 65) def test_512_to_16_mapping_and_prefill_only_mode(self) -> None: with tempfile.TemporaryDirectory() as tmp: root = Path(tmp) trace = root / "code-window.jsonl" trace.write_text( json.dumps( { "chat_id": "code-1", "parent_chat_id": -1, "session_root": "code-1", "turn": 1, "timestamp": 10.0, "input_length": 513, "output_length": 128, "hash_ids": [101, 202], "sampling_u": 0.1, } ) + "\n" ) output = root / "mapped" args = SimpleNamespace( input=trace, prompt=None, tokenizer=None, input_is_remapped=False, output_root=output, rho=1.0, start_timestamp=None, duration_s=None, vocab_size=151936, token_offset=1024, served_model="test-model", source_block_size=512, workload_mode="prefill_only", validate_parents=True, max_total_tokens=514, frontier_only=False, ) manifest = materialize(args) request = json.loads((output / "real_requests.jsonl").read_text()) mapped = json.loads((output / "selected-remapped.jsonl").read_text()) with (output / "frontier.csv").open(newline="") as stream: frontier_row = next(csv.DictReader(stream)) self.assertEqual(manifest["source_block_size"], 512) self.assertEqual(manifest["workload_mode"], "prefill_only") self.assertEqual(manifest["target_16_blocks"], math.floor(513 / 16)) self.assertEqual(request["input_length"], 513) self.assertEqual(request["output_length"], 1) self.assertEqual(request["body"]["max_tokens"], 1) self.assertEqual(len(request["runtime_block_ids"]), math.floor(513 / 16)) self.assertEqual(mapped["output_length"], 1) self.assertEqual(len(mapped["hash_ids_16"]), math.ceil(513 / 16)) self.assertEqual( len(frontier_row["block_hash_ids"].split("|")), math.floor(513 / 16) ) def test_real_prompt_tokens_preserve_parent_and_four_to_one_contract(self) -> None: tokenizer = CharacterTokenizer() parent_prompt = "A" * 64 child_prompt = parent_prompt + "B" * 16 with tempfile.TemporaryDirectory() as tmp: root = Path(tmp) trace = root / "chat-raw-window.jsonl" prompts = root / "chat-prompt-window.jsonl" rows = [ { "chat_id": "parent", "parent_chat_id": -1, "session_root": "parent", "turn": 1, "timestamp": 10.0, "input_length": 64, "output_length": 8, "hash_ids": [101], "sampling_u": 0.1, }, { "chat_id": "child", "parent_chat_id": "parent", "session_root": "parent", "turn": 2, "timestamp": 11.0, "input_length": 80, "output_length": 8, "hash_ids": [101, 202], "sampling_u": 0.1, }, ] prompt_rows = [ {"chat_id": "parent", "turn": 1, "prompt": parent_prompt}, {"chat_id": "child", "turn": 2, "prompt": child_prompt}, ] trace.write_text("".join(json.dumps(row) + "\n" for row in rows)) prompts.write_text("".join(json.dumps(row) + "\n" for row in prompt_rows)) output = root / "mapped" args = SimpleNamespace( input=trace, prompt=prompts, tokenizer=None, input_is_remapped=False, output_root=output, rho=1.0, start_timestamp=None, duration_s=None, vocab_size=151936, token_offset=1024, served_model="test-model", validate_parents=True, max_total_tokens=None, frontier_only=False, ) manifest = materialize(args, tokenizer=tokenizer) with (output / "selected-remapped.jsonl").open() as stream: mapped = [json.loads(line) for line in stream] filtered_output = root / "filtered" filtered_args = SimpleNamespace( input=output / "selected-remapped.jsonl", prompt=None, tokenizer=None, input_is_remapped=True, output_root=filtered_output, rho=1.0, start_timestamp=None, duration_s=None, vocab_size=151936, token_offset=1024, served_model="test-model", validate_parents=True, max_total_tokens=None, frontier_only=False, ) filtered_manifest = materialize(filtered_args) expected_parent = token_block_identities(prompt_token_ids(tokenizer, parent_prompt)) expected_child = token_block_identities(prompt_token_ids(tokenizer, child_prompt)) self.assertEqual(mapped[0]["hash_ids_16"], expected_parent) self.assertEqual(mapped[1]["hash_ids_16"], expected_child) self.assertEqual(mapped[1]["hash_ids_16"][:4], mapped[0]["hash_ids_16"]) self.assertEqual(len(mapped[0]["hash_ids_16"]), math.ceil(64 / 16)) self.assertEqual(len(mapped[1]["hash_ids_16"]), math.ceil(80 / 16)) self.assertEqual(manifest["block_contract"]["source_to_runtime_relation_conflicts"], 0) self.assertEqual(manifest["block_contract"]["runtime_to_source_relation_conflicts"], 0) self.assertEqual(manifest["prompt_contract"]["real_prompt_requests"], 2) self.assertEqual(manifest["prompt_contract"]["synthetic_fallback_requests"], 0) self.assertEqual(mapped[0]["prompt"], prompt_token_ids(tokenizer, parent_prompt)) self.assertEqual(filtered_manifest["block_contract"], manifest["block_contract"]) self.assertIsNotNone(filtered_manifest["upstream_remap_manifest_sha256"]) self.assertEqual(filtered_manifest["paired_row_vector_sha256"], manifest["paired_row_vector_sha256"]) if __name__ == "__main__": unittest.main()