96 lines
3.3 KiB
Python
96 lines
3.3 KiB
Python
import importlib.util
|
|
import json
|
|
import pickle
|
|
from pathlib import Path
|
|
import sys
|
|
import types
|
|
import unittest
|
|
from unittest.mock import patch
|
|
|
|
|
|
MODULE_PATH = Path(__file__).resolve().parents[1] / "core" / "data_collector.py"
|
|
|
|
|
|
def load_data_collector_module():
|
|
api = types.ModuleType("api")
|
|
api.base_url = "https://cloud.example"
|
|
api.data_record_url = "https://cloud.example"
|
|
api.the_folder = "customer-a/line-1"
|
|
requests = types.ModuleType("requests")
|
|
|
|
spec = importlib.util.spec_from_file_location(
|
|
"data_collector_under_test", MODULE_PATH
|
|
)
|
|
module = importlib.util.module_from_spec(spec)
|
|
with patch.dict(sys.modules, {"api": api, "requests": requests}):
|
|
spec.loader.exec_module(module)
|
|
return module
|
|
|
|
|
|
DATA_COLLECTOR = load_data_collector_module()
|
|
|
|
|
|
class ImmediateThread:
|
|
def __init__(self, target, daemon):
|
|
self.target = target
|
|
|
|
def start(self):
|
|
self.target()
|
|
|
|
|
|
class DataCollectorUploadTests(unittest.TestCase):
|
|
def test_uploads_control_manifest_with_complete_metadata(self):
|
|
collector = DATA_COLLECTOR.DataCollector()
|
|
collector.record_step(0, 10.0, 20.0, 30.0, 1.0, 0.2, 0.0, 50.0, 5.0)
|
|
collector.record_step(1, 11.0, 20.0, 31.0, 1.0, 0.2, 0.0, 50.0, 5.0)
|
|
|
|
uploads = []
|
|
completions = []
|
|
collector._upload_to_server = lambda data, name, folder: (
|
|
uploads.append((data, name, folder)) or True
|
|
)
|
|
collector.set_upload_complete_callback(
|
|
lambda success, manifest, error: completions.append(
|
|
(success, manifest, error)
|
|
)
|
|
)
|
|
|
|
with patch.object(DATA_COLLECTOR.threading, "Thread", ImmediateThread):
|
|
collector.finalize_and_upload(50.0, 5.0)
|
|
|
|
self.assertEqual(len(uploads), 2)
|
|
part_data, part_name, folder = uploads[0]
|
|
manifest_data, manifest_name, manifest_folder = uploads[1]
|
|
self.assertTrue(part_name.endswith(".pkl"))
|
|
self.assertTrue(manifest_name.endswith("_manifest.json"))
|
|
self.assertEqual(folder, "customer-a/line-1/data_record/data_50.0SLM_5.0L")
|
|
self.assertEqual(manifest_folder, folder)
|
|
self.assertEqual(len(pickle.loads(part_data)), 1)
|
|
|
|
manifest = json.loads(manifest_data)
|
|
self.assertEqual(manifest["schema_version"], 1)
|
|
self.assertEqual(manifest["data_type"], "control_episode")
|
|
self.assertEqual(manifest["total_episodes"], 1)
|
|
self.assertEqual(manifest["uploaded_chunks"], 1)
|
|
self.assertEqual(manifest["part_files"], [part_name])
|
|
self.assertEqual(manifest["parts"][0]["file_name"], part_name)
|
|
self.assertEqual(completions, [(True, manifest, None)])
|
|
self.assertEqual(collector.episode_data_raw, [])
|
|
|
|
def test_does_not_create_an_empty_chunk_for_oversized_episode(self):
|
|
collector = DATA_COLLECTOR.DataCollector()
|
|
collector.episode_data_raw = [{"payload": "x" * (5 * 1024 * 1024)}]
|
|
uploads = []
|
|
collector._upload_to_server = lambda data, name, folder: (
|
|
uploads.append((data, name, folder)) or True
|
|
)
|
|
|
|
with patch.object(DATA_COLLECTOR.threading, "Thread", ImmediateThread):
|
|
collector.finalize_and_upload(1.0, 1.0)
|
|
|
|
part_data = uploads[0][0]
|
|
self.assertEqual(len(pickle.loads(part_data)), 1)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main() |