36 lines
1.3 KiB
Python
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()
|