test_scripts.py 13 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441
  1. import io
  2. import sys
  3. from datetime import UTC, date, datetime
  4. from decimal import Decimal
  5. from unittest.mock import patch
  6. from django.core.files.uploadedfile import SimpleUploadedFile
  7. from django.test import TestCase
  8. from netaddr import IPAddress, IPNetwork
  9. from dcim.models import DeviceRole
  10. from extras.constants import SCRIPT_MODULE_NAME_PREFIX
  11. from extras.models import ScriptModule
  12. from extras.scripts import *
  13. CHOICES = (
  14. ('ff0000', 'Red'),
  15. ('00ff00', 'Green'),
  16. ('0000ff', 'Blue')
  17. )
  18. YAML_DATA = """
  19. Foo: 123
  20. Bar: 456
  21. Baz:
  22. - A
  23. - B
  24. - C
  25. """
  26. JSON_DATA = """
  27. {
  28. "Foo": 123,
  29. "Bar": 456,
  30. "Baz": ["A", "B", "C"]
  31. }
  32. """
  33. class ScriptVariablesTestCase(TestCase):
  34. def test_stringvar(self):
  35. class TestScript(Script):
  36. var1 = StringVar(
  37. min_length=3,
  38. max_length=3,
  39. regex=r'[a-z]+'
  40. )
  41. # Validate min_length enforcement
  42. data = {'var1': 'xx'}
  43. form = TestScript().as_form(data, None)
  44. self.assertFalse(form.is_valid())
  45. self.assertIn('var1', form.errors)
  46. # Validate max_length enforcement
  47. data = {'var1': 'xxxx'}
  48. form = TestScript().as_form(data, None)
  49. self.assertFalse(form.is_valid())
  50. self.assertIn('var1', form.errors)
  51. # Validate regex enforcement
  52. data = {'var1': 'ABC'}
  53. form = TestScript().as_form(data, None)
  54. self.assertFalse(form.is_valid())
  55. self.assertIn('var1', form.errors)
  56. # Validate valid data
  57. data = {'var1': 'abc'}
  58. form = TestScript().as_form(data, None)
  59. self.assertTrue(form.is_valid())
  60. self.assertEqual(form.cleaned_data['var1'], data['var1'])
  61. def test_textvar(self):
  62. class TestScript(Script):
  63. var1 = TextVar()
  64. # Validate valid data
  65. data = {'var1': 'This is a test string'}
  66. form = TestScript().as_form(data, None)
  67. self.assertTrue(form.is_valid())
  68. self.assertEqual(form.cleaned_data['var1'], data['var1'])
  69. def test_integervar(self):
  70. class TestScript(Script):
  71. var1 = IntegerVar(
  72. min_value=5,
  73. max_value=10
  74. )
  75. # Validate min_value enforcement
  76. data = {'var1': 4}
  77. form = TestScript().as_form(data, None)
  78. self.assertFalse(form.is_valid())
  79. self.assertIn('var1', form.errors)
  80. # Validate max_value enforcement
  81. data = {'var1': 11}
  82. form = TestScript().as_form(data, None)
  83. self.assertFalse(form.is_valid())
  84. self.assertIn('var1', form.errors)
  85. # Validate valid data
  86. data = {'var1': 7}
  87. form = TestScript().as_form(data, None)
  88. self.assertTrue(form.is_valid())
  89. self.assertEqual(form.cleaned_data['var1'], data['var1'])
  90. def test_decimalvar(self):
  91. class TestScript(Script):
  92. var1 = DecimalVar(
  93. min_value=-100.500,
  94. max_value=100.500,
  95. max_digits=6,
  96. decimal_places=3,
  97. required=False
  98. )
  99. var2 = DecimalVar(
  100. max_digits=3,
  101. decimal_places=1,
  102. required=False
  103. )
  104. # Validate min_value enforcement
  105. data = {'var1': -100.501}
  106. form = TestScript().as_form(data, None)
  107. self.assertFalse(form.is_valid())
  108. self.assertIn('var1', form.errors)
  109. # Validate max_value enforcement
  110. data = {'var1': 100.501}
  111. form = TestScript().as_form(data, None)
  112. self.assertFalse(form.is_valid())
  113. self.assertIn('var1', form.errors)
  114. # Validate max_digits enforcement
  115. data = {'var2': 123.4}
  116. form = TestScript().as_form(data, None)
  117. self.assertFalse(form.is_valid())
  118. self.assertIn('var2', form.errors)
  119. # Validate decimal_places
  120. data = {'var2': 1.23}
  121. form = TestScript().as_form(data, None)
  122. self.assertFalse(form.is_valid())
  123. self.assertIn('var2', form.errors)
  124. # Validate valid data
  125. data = {'var1': '50.123'}
  126. form = TestScript().as_form(data, None)
  127. self.assertTrue(form.is_valid())
  128. self.assertEqual(form.cleaned_data['var1'], Decimal(data['var1']))
  129. def test_booleanvar(self):
  130. class TestScript(Script):
  131. var1 = BooleanVar()
  132. # Validate True
  133. data = {'var1': True}
  134. form = TestScript().as_form(data, None)
  135. self.assertTrue(form.is_valid())
  136. self.assertEqual(form.cleaned_data['var1'], True)
  137. # Validate False
  138. data = {'var1': False}
  139. form = TestScript().as_form(data, None)
  140. self.assertTrue(form.is_valid())
  141. self.assertEqual(form.cleaned_data['var1'], False)
  142. def test_choicevar(self):
  143. class TestScript(Script):
  144. var1 = ChoiceVar(
  145. choices=CHOICES
  146. )
  147. # Validate valid choice
  148. data = {'var1': 'ff0000'}
  149. form = TestScript().as_form(data)
  150. self.assertTrue(form.is_valid())
  151. self.assertEqual(form.cleaned_data['var1'], 'ff0000')
  152. # Validate invalid choice
  153. data = {'var1': 'taupe'}
  154. form = TestScript().as_form(data)
  155. self.assertFalse(form.is_valid())
  156. def test_multichoicevar(self):
  157. class TestScript(Script):
  158. var1 = MultiChoiceVar(
  159. choices=CHOICES
  160. )
  161. # Validate single choice
  162. data = {'var1': ['ff0000']}
  163. form = TestScript().as_form(data)
  164. self.assertTrue(form.is_valid())
  165. self.assertEqual(form.cleaned_data['var1'], ['ff0000'])
  166. # Validate multiple choices
  167. data = {'var1': ('ff0000', '00ff00')}
  168. form = TestScript().as_form(data)
  169. self.assertTrue(form.is_valid())
  170. self.assertEqual(form.cleaned_data['var1'], ['ff0000', '00ff00'])
  171. # Validate invalid choice
  172. data = {'var1': 'taupe'}
  173. form = TestScript().as_form(data)
  174. self.assertFalse(form.is_valid())
  175. def test_objectvar(self):
  176. class TestScript(Script):
  177. var1 = ObjectVar(model=DeviceRole)
  178. # Populate some objects
  179. for i in range(1, 6):
  180. DeviceRole(
  181. name='Device Role {}'.format(i),
  182. slug='device-role-{}'.format(i)
  183. ).save()
  184. # Validate valid data
  185. data = {'var1': DeviceRole.objects.first().pk}
  186. form = TestScript().as_form(data, None)
  187. self.assertTrue(form.is_valid())
  188. self.assertEqual(form.cleaned_data['var1'].pk, data['var1'])
  189. def test_multiobjectvar(self):
  190. class TestScript(Script):
  191. var1 = MultiObjectVar(model=DeviceRole)
  192. # Populate some objects
  193. for i in range(1, 6):
  194. DeviceRole(
  195. name='Device Role {}'.format(i),
  196. slug='device-role-{}'.format(i)
  197. ).save()
  198. # Validate valid data
  199. data = {'var1': [role.pk for role in DeviceRole.objects.all()[:3]]}
  200. form = TestScript().as_form(data, None)
  201. self.assertTrue(form.is_valid())
  202. self.assertEqual(form.cleaned_data['var1'][0].pk, data['var1'][0])
  203. self.assertEqual(form.cleaned_data['var1'][1].pk, data['var1'][1])
  204. self.assertEqual(form.cleaned_data['var1'][2].pk, data['var1'][2])
  205. def test_filevar(self):
  206. class TestScript(Script):
  207. var1 = FileVar()
  208. # Dummy file
  209. testfile = SimpleUploadedFile(
  210. name='test_file.txt',
  211. content=b'This is a dummy file for testing'
  212. )
  213. # Validate valid data
  214. file_data = {'var1': testfile}
  215. form = TestScript().as_form(None, file_data)
  216. self.assertTrue(form.is_valid())
  217. self.assertEqual(form.cleaned_data['var1'], testfile)
  218. def test_ipaddressvar(self):
  219. class TestScript(Script):
  220. var1 = IPAddressVar()
  221. # Validate IP network enforcement
  222. data = {'var1': '1.2.3'}
  223. form = TestScript().as_form(data, None)
  224. self.assertFalse(form.is_valid())
  225. self.assertIn('var1', form.errors)
  226. # Validate IP mask exclusion
  227. data = {'var1': '192.0.2.0/24'}
  228. form = TestScript().as_form(data, None)
  229. self.assertFalse(form.is_valid())
  230. self.assertIn('var1', form.errors)
  231. # Validate valid data
  232. data = {'var1': '192.0.2.1'}
  233. form = TestScript().as_form(data, None)
  234. self.assertTrue(form.is_valid())
  235. self.assertEqual(form.cleaned_data['var1'], IPAddress(data['var1']))
  236. def test_ipaddresswithmaskvar(self):
  237. class TestScript(Script):
  238. var1 = IPAddressWithMaskVar()
  239. # Validate IP network enforcement
  240. data = {'var1': '1.2.3'}
  241. form = TestScript().as_form(data, None)
  242. self.assertFalse(form.is_valid())
  243. self.assertIn('var1', form.errors)
  244. # Validate IP mask requirement
  245. data = {'var1': '192.0.2.0'}
  246. form = TestScript().as_form(data, None)
  247. self.assertFalse(form.is_valid())
  248. self.assertIn('var1', form.errors)
  249. # Validate valid data
  250. data = {'var1': '192.0.2.0/24'}
  251. form = TestScript().as_form(data, None)
  252. self.assertTrue(form.is_valid())
  253. self.assertEqual(form.cleaned_data['var1'], IPNetwork(data['var1']))
  254. def test_ipnetworkvar(self):
  255. class TestScript(Script):
  256. var1 = IPNetworkVar()
  257. # Validate IP network enforcement
  258. data = {'var1': '1.2.3'}
  259. form = TestScript().as_form(data, None)
  260. self.assertFalse(form.is_valid())
  261. self.assertIn('var1', form.errors)
  262. # Validate host IP check
  263. data = {'var1': '192.0.2.1/24'}
  264. form = TestScript().as_form(data, None)
  265. self.assertFalse(form.is_valid())
  266. self.assertIn('var1', form.errors)
  267. # Validate valid data
  268. data = {'var1': '192.0.2.0/24'}
  269. form = TestScript().as_form(data, None)
  270. self.assertTrue(form.is_valid())
  271. self.assertEqual(form.cleaned_data['var1'], IPNetwork(data['var1']))
  272. def test_datevar(self):
  273. class TestScript(Script):
  274. var1 = DateVar()
  275. var2 = DateVar(required=False)
  276. # Test date validation
  277. data = {'var1': 'not a date'}
  278. form = TestScript().as_form(data, None)
  279. self.assertFalse(form.is_valid())
  280. self.assertIn('var1', form.errors)
  281. # Validate valid data
  282. input_date = date(2024, 4, 1)
  283. data = {'var1': input_date}
  284. form = TestScript().as_form(data, None)
  285. self.assertTrue(form.is_valid())
  286. self.assertEqual(form.cleaned_data['var1'], input_date)
  287. # Validate required=False works for this Var type
  288. self.assertEqual(form.cleaned_data['var2'], None)
  289. def test_datetimevar(self):
  290. class TestScript(Script):
  291. var1 = DateTimeVar()
  292. var2 = DateTimeVar(required=False)
  293. # Test datetime validation
  294. data = {'var1': 'not a datetime'}
  295. form = TestScript().as_form(data, None)
  296. self.assertFalse(form.is_valid())
  297. self.assertIn('var1', form.errors)
  298. # Validate valid data
  299. input_datetime = datetime(2024, 4, 1, 8, 0, 0, 0, UTC)
  300. data = {'var1': input_datetime}
  301. form = TestScript().as_form(data, None)
  302. self.assertTrue(form.is_valid())
  303. self.assertEqual(form.cleaned_data['var1'], input_datetime)
  304. # Validate required=False works for this Var type
  305. self.assertEqual(form.cleaned_data['var2'], None)
  306. class ScriptModuleLoadingTestCase(TestCase):
  307. def test_module_does_not_shadow_core_app(self):
  308. """
  309. Loading a custom script whose filename matches a core app label must not replace that
  310. app's package in sys.modules. Regression test for issue #22566.
  311. """
  312. import circuits # The real core app package
  313. script_content = (
  314. b"from extras.scripts import Script\n\n\n"
  315. b"class TestScript(Script):\n pass\n"
  316. )
  317. class _Storage:
  318. def open(self, name, mode='rb'):
  319. return io.BytesIO(script_content)
  320. module = ScriptModule(file_root='scripts', file_path='circuits.py')
  321. namespaced_key = f'{SCRIPT_MODULE_NAME_PREFIX}circuits'
  322. self.addCleanup(lambda: sys.modules.pop(namespaced_key, None))
  323. with patch('extras.models.mixins.storages') as mock_storages:
  324. mock_storages.__getitem__.return_value = _Storage()
  325. loaded = module.get_module()
  326. # The script module is registered under the private, namespaced key, and its own
  327. # __name__ matches that key (i.e. sys.modules[module.__name__] resolves to the module)
  328. self.assertIs(sys.modules[namespaced_key], loaded)
  329. self.assertEqual(loaded.__name__, namespaced_key)
  330. # The namespacing must not leak into the derived script name stored in the database
  331. self.assertEqual(next(iter(module.module_scripts)), 'TestScript')
  332. # Nor into the user-facing names exposed on the Script class (used for logger
  333. # namespaces, page headers, etc.): these must reflect the original filename.
  334. script_class = loaded.TestScript
  335. self.assertEqual(script_class.module, 'circuits')
  336. self.assertEqual(script_class.full_name, 'circuits.TestScript')
  337. self.assertEqual(script_class.root_module(), 'circuits')
  338. # The real circuits app must be untouched and remain an importable package
  339. self.assertIs(sys.modules['circuits'], circuits)
  340. self.assertTrue(hasattr(circuits, '__path__'))