102 lines
4.0 KiB
Python
102 lines
4.0 KiB
Python
"""Tests for scripts/build_common.py version stamping helpers."""
|
|
import json
|
|
import os
|
|
import shutil
|
|
import sys
|
|
import tempfile
|
|
import unittest
|
|
|
|
sys.path.insert(0, os.path.join(os.path.dirname(__file__), '..', 'scripts'))
|
|
import build_common
|
|
|
|
|
|
class BuildCommonTest(unittest.TestCase):
|
|
def setUp(self):
|
|
self.dir = tempfile.mkdtemp()
|
|
self.payload = os.path.join(self.dir, 'user', 'remote_access',
|
|
'pager-webui')
|
|
os.makedirs(self.payload)
|
|
with open(os.path.join(self.dir, '_hak5_manifest.json'), 'w') as f:
|
|
f.write('{"payload": "pager-webui", "version": "1.3.2"}')
|
|
with open(os.path.join(self.payload, 'payload.sh'), 'w') as f:
|
|
f.write('#!/bin/bash\n'
|
|
'# Title: Mark VIII\n'
|
|
'# Description: test payload\n'
|
|
'# Version: 1.3.2\n'
|
|
'# Category: Remote-Access\n'
|
|
'\n'
|
|
'echo hi\n')
|
|
with open(os.path.join(self.payload, 'server.py'), 'w') as f:
|
|
f.write('"""Mark VIII server."""\n'
|
|
'import os\n'
|
|
'\n'
|
|
'PORT = 8080\n')
|
|
|
|
def tearDown(self):
|
|
shutil.rmtree(self.dir, ignore_errors=True)
|
|
|
|
def _server_text(self):
|
|
with open(os.path.join(self.payload, 'server.py')) as f:
|
|
return f.read()
|
|
|
|
def test_stamp_version_updates_all_three_files(self):
|
|
stamped = sorted(build_common.stamp_version(self.dir, '1.4.0'))
|
|
expected = sorted([
|
|
os.path.join(self.dir, '_hak5_manifest.json'),
|
|
os.path.join(self.payload, 'payload.sh'),
|
|
os.path.join(self.payload, 'server.py'),
|
|
])
|
|
self.assertEqual(stamped, expected)
|
|
with open(os.path.join(self.dir, '_hak5_manifest.json')) as f:
|
|
self.assertEqual(json.load(f)['version'], '1.4.0')
|
|
with open(os.path.join(self.payload, 'payload.sh')) as f:
|
|
sh_text = f.read()
|
|
self.assertIn('# Version: 1.4.0', sh_text)
|
|
self.assertNotIn('# Version: 1.3.2', sh_text)
|
|
server_text = self._server_text()
|
|
self.assertIn("SERVER_VERSION = '1.4.0'", server_text)
|
|
self.assertEqual(server_text.count('SERVER_VERSION'), 1)
|
|
self.assertLess(server_text.index("SERVER_VERSION = '1.4.0'"),
|
|
server_text.index('\nimport os'))
|
|
|
|
def test_stamp_version_is_idempotent_and_upgrades(self):
|
|
build_common.stamp_version(self.dir, '1.4.0')
|
|
stamped = build_common.stamp_version(self.dir, '1.5.0')
|
|
self.assertEqual(len(stamped), 3)
|
|
server_text = self._server_text()
|
|
self.assertEqual(server_text.count('SERVER_VERSION'), 1)
|
|
self.assertIn("SERVER_VERSION = '1.5.0'", server_text)
|
|
with open(os.path.join(self.payload, 'payload.sh')) as f:
|
|
self.assertIn('# Version: 1.5.0', f.read())
|
|
|
|
def test_server_without_docstring_gets_top_injection(self):
|
|
path = os.path.join(self.payload, 'server.py')
|
|
with open(path, 'w') as f:
|
|
f.write('# comment header\n'
|
|
'\n'
|
|
'import os\n'
|
|
'PORT = 8080\n')
|
|
build_common.stamp_version(self.dir, '9.9.9')
|
|
text = self._server_text()
|
|
lines = text.splitlines(True)
|
|
idx = [i for i, l in enumerate(lines) if l.startswith('SERVER_VERSION')]
|
|
self.assertEqual(len(idx), 1)
|
|
self.assertLess(idx[0], [i for i, l in enumerate(lines)
|
|
if l.startswith('import os')][0])
|
|
|
|
def test_missing_files_are_tolerated(self):
|
|
empty = tempfile.mkdtemp()
|
|
try:
|
|
self.assertEqual(build_common.stamp_version(empty, '1.4.0'), [])
|
|
finally:
|
|
shutil.rmtree(empty, ignore_errors=True)
|
|
|
|
def test_invalid_version_is_rejected(self):
|
|
for bad in ("1.4'; import os", '', 'a' * 64, 'ver x'):
|
|
with self.assertRaises(ValueError):
|
|
build_common.stamp_version(self.dir, bad)
|
|
|
|
|
|
if __name__ == '__main__':
|
|
unittest.main()
|