import ast import importlib.util import pathlib import sys import types import unittest spec = importlib.util.spec_from_file_location('patcher', pathlib.Path(__file__).parents[1] / 'scripts/patch_yue_cancellation.py') patcher = importlib.util.module_from_spec(spec) spec.loader.exec_module(patcher) class CancellationTests(unittest.TestCase): def test_interrupt_propagates_before_sampling(self): source = 'class BlockTokenRangeProcessor:\n def __call__(self, input_ids, scores):\n return scores\n' result = patcher.patch(source) self.assertEqual(result, patcher.patch(result)) cancelled = False class Interrupted(Exception): pass def check(): if cancelled: raise Interrupted() module = types.ModuleType('comfy.model_management') module.throw_exception_if_processing_interrupted = check sys.modules['comfy.model_management'] = module namespace = {} exec(compile(ast.parse(result), '', 'exec'), namespace) callback = namespace['BlockTokenRangeProcessor']() self.assertEqual(callback(None, 42), 42) cancelled = True with self.assertRaises(Interrupted): callback(None, 42) if __name__ == '__main__': unittest.main()