Files
aigen/tests/test_yue_cancellation.py
T

36 lines
1.3 KiB
Python

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), '<patched>', '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()