Files
yidaconnector/tests/test_sync.py
T

135 lines
5.5 KiB
Python

import copy
import io
import sys
import tempfile
import unittest
import zipfile
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parents[1] / 'scripts'))
from sync_releases import API, MARKER, SyncError, digest, sync_one, validate_file, validate_names
def zip_bytes(text):
b = io.BytesIO()
with zipfile.ZipFile(b, 'w') as z:
z.writestr('README.txt', text)
return b.getvalue()
class Fake:
base = 'https://source.invalid'
repo = 'owner/repo'
def __init__(self, source=False):
self.assets = []
self.data = {}
self.release = {'id': 1, 'tag_name': 'v1', 'name': 'one', 'body': 'notes', 'draft': False} if source else None
self.writes = []
self.fail_rename = False
def add(self, name, content):
i = len(self.data) + 1
a = {'id': i, 'name': name, 'size': len(content)}
self.assets.append(a)
self.data[i] = content
return copy.deepcopy(a)
def pages(self, path):
return copy.deepcopy(self.assets)
def list_assets(self, release_id):
return copy.deepcopy(self.assets)
def download(self, a, p):
Path(p).write_bytes(self.data[a['id']])
validate_file(p, a)
return digest(p)
def request(self, method, path, body=None, **kwargs):
if method == 'GET':
return {'commit': {'sha': 'a' * 40}} if path.startswith('/tags/') else copy.deepcopy(self.release)
self.writes.append((method, path))
if method == 'POST' and path == '/releases':
self.release = dict(body, id=2)
return copy.deepcopy(self.release)
if method == 'POST':
from urllib.parse import unquote
return self.add(unquote(path.split('name=')[1]), body)
if '/assets/' in path:
if self.fail_rename and body['name'] == 'app.zip':
self.fail_rename = False
raise SyncError('simulated interrupted rename')
item = next(a for a in self.assets if a['id'] == int(path.rsplit('/', 1)[1]))
item.update(body)
return copy.deepcopy(item)
self.release.update(body)
return copy.deepcopy(self.release)
class Tests(unittest.TestCase):
def setUp(self):
self.src, self.dst = Fake(True), Fake()
self.src.add('app.zip', zip_bytes('version 1'))
def run_sync(self, dry=False):
with tempfile.TemporaryDirectory() as d:
return sync_one(self.src, self.dst, self.src.release, 'a' * 40, dry, d)
def test_dry_run_writes_nothing(self):
self.run_sync(True)
self.assertEqual([], self.dst.writes)
def test_new_release_published_only_after_verified_assets(self):
self.run_sync()
self.assertFalse(self.dst.release['draft'])
self.assertEqual('app.zip', self.dst.assets[0]['name'])
self.assertEqual(('PATCH', '/releases/2'), self.dst.writes[-1])
def test_repeat_is_noop(self):
self.run_sync()
self.dst.writes.clear()
result = self.run_sync()
self.assertEqual([], self.dst.writes)
self.assertEqual('unchanged', result['assets'][0]['action'])
def test_same_name_same_size_different_bytes_is_detected(self):
self.run_sync()
self.src.data[1] = zip_bytes('version 2')
result = self.run_sync()
self.assertEqual('replace', result['assets'][0]['action'])
self.assertTrue(any(a['name'].startswith('.backup-') for a in self.dst.assets))
def test_interrupted_swap_recovers(self):
self.run_sync()
self.src.data[1] = zip_bytes('version 2')
self.dst.fail_rename = True
with self.assertRaises(SyncError): self.run_sync()
self.run_sync()
self.assertEqual(1, sum(a['name'] == 'app.zip' for a in self.dst.assets))
self.assertEqual(2, len(self.dst.assets))
def test_unmanaged_release_never_overwritten(self):
self.dst.release = dict(self.src.release)
with self.assertRaises(SyncError): self.run_sync()
self.assertEqual([], self.dst.writes)
def test_deleted_source_asset_retained(self):
self.run_sync()
self.src.assets.clear()
self.run_sync()
self.assertEqual(1, len(self.dst.assets))
def test_html_error_not_zip(self):
with tempfile.TemporaryDirectory() as d:
p = Path(d) / 'bad'; p.write_bytes(b'Not found.\n')
with self.assertRaises(SyncError): validate_file(p, {'size': 11, 'name': 'bad.zip'})
def test_unsafe_names_rejected(self):
for name in ['../secret', '.sync-x', '.backup-x', 'a\\b']:
with self.assertRaises(SyncError): validate_names([{'name': name, 'size': 1}])
def test_pagination_does_not_assume_server_page_size(self):
api = API('https://example.invalid', 'a/b', '')
pages = iter([[1, 2], [3], []])
api.request = lambda *a: next(pages)
self.assertEqual([1, 2, 3], api.pages('/releases'))
def test_cross_origin_download_rejected_before_request(self):
api = API('https://example.invalid', 'a/b', 'secret')
with self.assertRaises(SyncError):
api.download({'browser_download_url': 'https://other.invalid/file'}, '/tmp/unused')
def test_assets_endpoint_is_not_paginated(self):
api = API('https://example.invalid', 'a/b', '')
calls = []
def request(method, path):
calls.append(path)
return [{'id': 1}]
api.request = request
self.assertEqual([{'id': 1}], api.list_assets(2))
self.assertEqual(['/releases/2/assets'], calls)
if __name__ == '__main__': unittest.main()