import shutil import tempfile import unittest from pathlib import Path import numpy as np from mapgen import pipeline as P from mapgen.testing import FIXTURE_TOML from tests.helpers import ROOT, make_ctx TECT = """ [[plate]] id = "a" seed = [0.0, 0.0] kind = "continental" motion = [90.0, 3.0] [[plate]] id = "b" seed = [0.0, 90.0] kind = "oceanic" motion = [270.0, 3.0] """ class PipelineTest(unittest.TestCase): def setUp(self): self.tmp = Path(tempfile.mkdtemp()) (self.tmp / "config").mkdir() shutil.copy(FIXTURE_TOML, self.tmp / "config" / "world.toml") (self.tmp / "config" / "tectonics.toml").write_text(TECT) self.log = [] def tearDown(self): shutil.rmtree(self.tmp) def _build(self, **kw): self.log = [] return P.build(self.tmp, 1, stop="grid", log=self.log.append, **kw) def test_cache_hit(self): ctx = self._build() self.assertEqual(ctx.grid.n, 842) self.assertFalse(any("cached" in m for m in self.log)) ctx = self._build() self.assertTrue(any("grid: cached" in m for m in self.log)) self.assertEqual(ctx.grid.n, 842) def test_config_change_invalidates(self): self._build() p = self.tmp / "config" / "world.toml" p.write_text(p.read_text().replace("seed = 1296", "seed = 1297")) self._build() self.assertFalse(any("cached" in m for m in self.log)) def test_from_forces_rerun(self): self._build() self._build(start="grid") self.assertFalse(any("cached" in m for m in self.log)) def test_need_reports_missing(self): ctx = make_ctx(1) with self.assertRaisesRegex(P.StageError, "sk_land"): ctx.need("sk_land") class OceanExportTest(unittest.TestCase): def test_cells_and_fields_carry_the_ocean(self): import json from mapgen.testing import built_world root = built_world() outs = sorted((root / "out").glob("r*/cells.npz")) self.assertTrue(outs) for cells in outs + sorted((root / "out").glob("r*/eras/*/cells.npz")): # every era re-runs climate z = np.load(cells) for k in ("current", "current_speed", "sst", "upwelling", "productivity"): self.assertIn(k, z.files, f"{cells}: {k}") meta = json.loads((outs[0].parent / "fields.json").read_text()) for k in ("sst", "productivity", "current_speed", "upwelling"): self.assertIn(k, meta["continuous"]) layers = [L["id"] for L in json.loads((outs[0].parent / "viewer" / "layers.json").read_text())] self.assertIn("currents", layers) class NpzMapsTest(unittest.TestCase): def test_same_arrays_as_np_load_writable_and_file_untouched(self): import hashlib import tempfile import numpy as np from pathlib import Path from mapgen import pipeline as P rng = np.random.default_rng(0) arrays = {"f": rng.normal(size=1000), "i": rng.integers(0, 9, (40, 3)).astype(np.int16), "b": rng.random(77) < 0.5, "F": np.asfortranarray(rng.random((30, 4))), "s": np.array(2.5), "e": np.zeros((0, 3), np.float32), "_key": np.array("abc"), "u": np.arange(5, dtype=np.uint64), "big": rng.normal(size=(3_000_000,)), "nc": rng.random((50, 8))[:, ::3], "x" * 200: np.arange(7)} with tempfile.TemporaryDirectory() as t: for save in (np.savez, np.savez_compressed, P.save_npz_aligned): f = Path(t) / f"{save.__name__}.npz" save(f, **arrays) before = hashlib.sha256(f.read_bytes()).hexdigest() got = P.npz_maps(f) if save is P.save_npz_aligned: # every member mapped in place, at aligned addresses for k, v in got.items(): if v.size: self.assertEqual(v.ctypes.data % P.ALIGN, 0, k) b = v while isinstance(b, (np.ndarray, memoryview)): b = b.base if isinstance(b, np.ndarray) else b.obj self.assertIsInstance(b, __import__("mmap").mmap, k) with np.load(f) as z: self.assertEqual(sorted(got), sorted(z.files)) for k in z.files: self.assertEqual(got[k].dtype, z[k].dtype, k) self.assertEqual(got[k].shape, z[k].shape, k) self.assertTrue(np.array_equal(got[k], z[k]), k) self.assertTrue(np.array_equal(got[k], arrays[k]), k) for k, v in got.items(): # whatever is mapped is aligned (else: read) b = v while isinstance(b, (np.ndarray, memoryview)): b = b.base if isinstance(b, np.ndarray) else b.obj if isinstance(b, __import__("mmap").mmap): self.assertEqual(v.ctypes.data % P.ALIGN, 0, (save.__name__, k)) got["f"][:] = -1.0 # private copy-on-write: allowed, file unchanged got["F"][0, 0] = 9.0 self.assertEqual(float(got["F"][0, 0]), 9.0) self.assertEqual(hashlib.sha256(f.read_bytes()).hexdigest(), before, save.__name__) self.assertTrue(np.array_equal(np.load(f)["f"], arrays["f"])) def test_low_memory_rebuild_from_cache_identical(self): import shutil import tempfile from pathlib import Path from mapgen import pipeline as P from mapgen.testing import small_world outs = [] tmp = Path(tempfile.mkdtemp()) try: small_world(tmp) for low in (False, True, True): # normal; low (fresh cache hit); low again ctx = P.build(tmp, 2, log=lambda m: None, low_memory=low) outs.append({k: v for k, v in ctx.data.items()}) for k, v in outs[0].items(): for o in outs[1:]: self.assertTrue(np.array_equal(np.asarray(o[k]), np.asarray(v), equal_nan=np.asarray(v).dtype.kind == "f"), k) self.assertFalse(list((tmp / "out" / "cache" / "r2").glob("*.tmp-*")), "no temporary files left") import mmap ctx = P.build(tmp, 2, start="climate", log=lambda m: None, low_memory=True) # freshly computed for k in ("T_jun", "P_ann", "z_surface_m", "plate"): b = ctx.data[k] while isinstance(b, (np.ndarray, memoryview)): b = b.base if isinstance(b, np.ndarray) else b.obj self.assertIsInstance(b, mmap.mmap, f"{k}: kept on disk, not in RAM") self.assertTrue(np.array_equal(ctx.data[k], outs[0][k], equal_nan=True), k) finally: shutil.rmtree(tmp) class AlignedSumTest(unittest.TestCase): def test_misaligned_data_sums_differently(self): """Why the cache is aligned: numpy sums misaligned float64 data in another order (else drop the alignment).""" import numpy as np x = np.random.default_rng(0).normal(size=2_000_000) buf = bytearray(8 * len(x) + 8) a = np.frombuffer(memoryview(buf)[4:4 + 8 * len(x)], np.float64) np.copyto(a, x) self.assertFalse(a.flags.aligned) self.assertTrue(np.array_equal(a, x)) if a.sum() == x.sum(): self.skipTest("this numpy sums misaligned data the same way") self.assertNotEqual(a.sum(), x.sum()) class BlasSpinTest(unittest.TestCase): def test_import_sets_a_short_blas_spin_unless_given(self): import os import subprocess import sys code = "import os, mapgen, numpy; print(os.environ['OPENBLAS_THREAD_TIMEOUT'])" env = {k: v for k, v in os.environ.items() if k != "OPENBLAS_THREAD_TIMEOUT"} out = subprocess.run([sys.executable, "-c", code], capture_output=True, text=True, env=env, cwd=os.getcwd()) self.assertEqual(out.stdout.strip(), "4") out = subprocess.run([sys.executable, "-c", code], capture_output=True, text=True, env={**env, "OPENBLAS_THREAD_TIMEOUT": "9"}, cwd=os.getcwd()) self.assertEqual(out.stdout.strip(), "9")