import unittest import torch from toolkit.trigger_binding import ( ActivatorModeError, TriggerAtomicityError, TriggerConflictError, TriggerPlaceholderError, TriggerTokenizerError, TriggerTruncationError, activator_runtime_mode, bind_trigger_batch, bind_trigger_prompt, get_activator_runtime_state, resolve_trigger_literal, validate_atomic_token_id, ) class _FastTokenizer: is_fast = True pad_token_id = 0 eos_token_id = 2 unk_token_id = 1 def __init__(self, literal="", atomic=True): self.literal = literal self.atomic = atomic self.literal_id = 700 def apply_chat_template(self, messages, add_generation_prompt=True, tokenize=False): text = messages[0]["content"][0]["text"] suffix = "" if add_generation_prompt else "" return f"{text}{suffix}" def __call__( self, text, add_special_tokens=False, return_offsets_mapping=False, truncation=False, max_length=None, ): input_ids = [] offsets = [] index = 0 while index < len(text): if self.atomic and text.startswith(self.literal, index): input_ids.append(self.literal_id) offsets.append((index, index + len(self.literal))) index += len(self.literal) else: input_ids.append(100 + (ord(text[index]) % 500)) offsets.append((index, index + 1)) index += 1 if truncation and max_length is not None: input_ids = input_ids[:max_length] offsets = offsets[:max_length] result = {"input_ids": input_ids, "attention_mask": [1] * len(input_ids)} if return_offsets_mapping: result["offset_mapping"] = offsets return result class _SlowTokenizer(_FastTokenizer): is_fast = False class _DuplicatingTemplateTokenizer(_FastTokenizer): def apply_chat_template(self, messages, add_generation_prompt=True, tokenize=False): text = messages[0]["content"][0]["text"] return f"{text} {self.literal}" class _Runtime: def __init__(self, runtime_mode="full"): self.runtime_mode = runtime_mode class TriggerBindingTest(unittest.TestCase): def setUp(self): self.literal = "" self.tokenizer = _FastTokenizer(self.literal) def test_resolver_replaces_all_occurrences_and_records_spans(self): resolved = resolve_trigger_literal( "alpha [trigger] beta [trigger] omega", self.literal, ) self.assertEqual(resolved.text, "alpha beta omega") self.assertEqual(resolved.occurrence_count, 2) self.assertEqual( tuple(resolved.text[start:end] for start, end in resolved.spans), (self.literal, self.literal), ) def test_resolver_rejects_missing_placeholder_and_literal_conflict(self): with self.assertRaises(TriggerPlaceholderError): resolve_trigger_literal("plain caption", self.literal) with self.assertRaises(TriggerConflictError): resolve_trigger_literal("[trigger] plus ", self.literal) def test_chat_template_offsets_create_all_occurrence_mask(self): binding = bind_trigger_prompt( self.tokenizer, "alpha [trigger] beta [trigger]", self.literal, require_atomic=True, expected_token_id=self.tokenizer.literal_id, ) self.assertTrue(binding.rendered_text.startswith("alpha")) self.assertEqual(binding.occurrence_count, 2) self.assertEqual(len(binding.token_indices), 2) self.assertEqual(sum(binding.trigger_mask), 2) self.assertTrue(all(binding.input_ids[index] == self.tokenizer.literal_id for index in binding.token_indices)) self.assertEqual( tuple(binding.rendered_text[start:end] for start, end in binding.character_spans), (self.literal, self.literal), ) def test_first_occurrence_mode_leaves_other_occurrences_unmasked(self): binding = bind_trigger_prompt( self.tokenizer, "[trigger] then [trigger]", self.literal, mask_all_occurrences=False, ) self.assertEqual(binding.occurrence_count, 2) self.assertEqual(len(binding.token_indices), 1) self.assertEqual(sum(binding.trigger_mask), 1) def test_truncation_is_detected_instead_of_silently_dropping_trigger(self): with self.assertRaises(TriggerTruncationError): bind_trigger_prompt( self.tokenizer, "a long prefix [trigger]", self.literal, max_length=5, ) def test_chat_template_literal_duplication_is_rejected(self): tokenizer = _DuplicatingTemplateTokenizer(self.literal) with self.assertRaises(TriggerConflictError): bind_trigger_prompt(tokenizer, "[trigger]", self.literal) def test_fast_tokenizer_and_atomic_id_validation(self): with self.assertRaises(TriggerTokenizerError): bind_trigger_prompt(_SlowTokenizer(self.literal), "[trigger]", self.literal) self.assertEqual(validate_atomic_token_id(self.tokenizer, self.literal), self.tokenizer.literal_id) with self.assertRaises(TriggerAtomicityError): validate_atomic_token_id(self.tokenizer, self.literal, expected_token_id=999) with self.assertRaises(TriggerAtomicityError): validate_atomic_token_id(_FastTokenizer(self.literal, atomic=False), self.literal) def test_batch_padding_masks_and_metadata(self): batch = bind_trigger_batch( self.tokenizer, ["[trigger]", "longer [trigger] and [trigger]"], self.literal, require_atomic=True, metadata={"phase": "a1"}, ) self.assertEqual(batch.input_ids.shape, batch.attention_mask.shape) self.assertEqual(batch.input_ids.shape, batch.trigger_mask.shape) self.assertEqual(batch.input_ids.dtype, torch.long) self.assertEqual(batch.trigger_mask.dtype, torch.bool) self.assertEqual(batch.trigger_mask.sum(dim=1).tolist(), [1, 2]) self.assertEqual(batch.metadata["batch_size"], 2) self.assertEqual(batch.metadata["occurrence_counts"], (1, 2)) self.assertEqual(batch.metadata["phase"], "a1") def test_all_runtime_modes_expose_expected_flags(self): full = get_activator_runtime_state("full") self.assertTrue(full.embedding_enabled and full.internal_enabled and full.tap_enabled) self.assertTrue(get_activator_runtime_state("embedding_only").embedding_enabled) self.assertTrue(get_activator_runtime_state("tap_only").tap_enabled) self.assertTrue(get_activator_runtime_state("internal_only").internal_enabled) self.assertTrue(get_activator_runtime_state("activator_bypass").activator_bypassed) self.assertTrue(get_activator_runtime_state("stock_literal").stock_literal) with self.assertRaises(ActivatorModeError): get_activator_runtime_state("unknown") def test_runtime_context_is_nested_and_exception_safe_for_objects(self): runtime = _Runtime("full") with activator_runtime_mode(runtime, "embedding_only"): self.assertEqual(runtime.runtime_mode, "embedding_only") with self.assertRaisesRegex(RuntimeError, "boom"): with activator_runtime_mode(runtime, "tap_only"): self.assertEqual(runtime.runtime_mode, "tap_only") raise RuntimeError("boom") self.assertEqual(runtime.runtime_mode, "embedding_only") self.assertEqual(runtime.runtime_mode, "full") def test_runtime_context_restores_mapping_and_removes_new_attribute(self): runtime = {"runtime_mode": "stock_literal"} with activator_runtime_mode(runtime, "activator_bypass"): self.assertEqual(runtime["runtime_mode"], "activator_bypass") self.assertEqual(runtime["runtime_mode"], "stock_literal") class Empty: pass empty = Empty() with activator_runtime_mode(empty, "internal_only"): self.assertEqual(empty.runtime_mode, "internal_only") self.assertFalse(hasattr(empty, "runtime_mode")) if __name__ == "__main__": unittest.main()