From 130690183aa788d3436298a03d94de287b961f89 Mon Sep 17 00:00:00 2001 From: Arthur Jacobs Date: Mon, 24 Aug 2026 18:16:27 -0300 Subject: [PATCH 1/3] test: expand unit tests to 222 and fix the four bugs they exposed Coverage was 72% with four modules at 0%. This takes it to 94% (222 tests, ~70s) and covers every module that had none. New suites: * test_persist.py -- save/load round trip, the zip path and its cleanup, and the not-a-model / missing-path cases. * test_rootpath.py -- marker detection from files and nested directories, custom patterns, and the empty-directory quirk (see below). * test_log.py -- message formatting, levels, filtering and handlers. * test_const.py -- structural validation of the bundled dataset metadata (read() unpacks each field as a 4-tuple, so a malformed entry would only surface as an opaque error at read time) plus the label converter. * test_utils_plot.py / test_report_plot.py -- every plotting entry point, each asserting no matplotlib figure is left open. That guards the leak fixed in 381eb17 and caught two of the bugs below. Bugs found by the new tests: * trustee/utils/const.py: np.uint8(-1) raises OverflowError on NumPy 2 -- NumPy 1 wrapped it to 255. cic_ids_2017_label_converter() therefore blew up on any unrecognised label instead of returning its sentinel. 255 is now the explicit sentinel, preserving the NumPy 1 result. * trustee/report/plot.py: DataFrame.groupby(axis=1) was removed in pandas 2.0, so plot_distribution(aggregate=True) raised TypeError. Rewritten to group the bit columns by prefix explicitly, preserving column order so the bits keep their significance. This one survived the 1.2.0 upgrade because nothing exercised the path. * trustee/report/plot.py: class_names defaults to [] at seven indexing sites guarded only by `is not None`, which passes for an empty list. Calling any of those entry points without class_names raised IndexError, as did a class_names shorter than the tree's class count. Added a _class_label() helper doing the bounds check trust.py already did. * trustee/utils/plot.py: labels[i] was indexed unguarded in plot_stacked_bars while every other site in the file bounds-checks. Behaviour documented rather than changed: * rootpath.detect() returns None on the first empty directory it meets, so an empty subdirectory aborts the search before reaching the marker above it. * plot_distribution(aggregate=True) needs feature_names -- it converts X with .values, so DataFrame column names never reach the prefix regex. * skip_retrain=True reuses the blackbox as-is and so requires a fitted model. Determinism: Trustee.fit() samples via np.random.choice and calls train_test_split without a random_state, both drawing on NumPy's global RNG, which made any assertion on fidelity flaky. An autouse fixture now seeds it per test; the suite was run three times over to confirm. CI gains --cov-fail-under=90 so coverage cannot quietly regress. Co-Authored-By: Claude Opus 5 (1M context) --- .github/workflows/test.yml | 4 +- poetry.lock | 158 ++++++++++++++- pyproject.toml | 5 + tests/conftest.py | 12 ++ tests/test_const.py | 97 ++++++++++ tests/test_dataset.py | 86 +++++++++ tests/test_log.py | 67 +++++++ tests/test_persist.py | 99 ++++++++++ tests/test_report.py | 98 ++++++++++ tests/test_report_plot.py | 383 +++++++++++++++++++++++++++++++++++++ tests/test_rootpath.py | 75 ++++++++ tests/test_trustee.py | 95 +++++++++ tests/test_utils_plot.py | 150 +++++++++++++++ trustee/report/plot.py | 45 +++-- trustee/utils/const.py | 8 +- trustee/utils/plot.py | 2 +- 16 files changed, 1367 insertions(+), 17 deletions(-) create mode 100644 tests/test_const.py create mode 100644 tests/test_log.py create mode 100644 tests/test_persist.py create mode 100644 tests/test_report_plot.py create mode 100644 tests/test_rootpath.py create mode 100644 tests/test_utils_plot.py diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml index 75c9709..a10caf2 100644 --- a/.github/workflows/test.yml +++ b/.github/workflows/test.yml @@ -73,7 +73,9 @@ jobs: run: poetry install --with dev - name: Run tests - run: poetry run pytest -v + # The floor sits just under the current figure so coverage cannot quietly + # regress; raise it as coverage improves rather than lowering it. + run: poetry run pytest -v --cov=trustee --cov-report=term-missing --cov-fail-under=90 examples: name: Examples diff --git a/poetry.lock b/poetry.lock index 82a3b18..6f7044c 100644 --- a/poetry.lock +++ b/poetry.lock @@ -443,6 +443,141 @@ mypy = ["bokeh", "contourpy[bokeh,docs]", "docutils-stubs", "mypy (==1.17.0)", " test = ["Pillow", "contourpy[test-no-images]", "matplotlib"] test-no-images = ["pytest", "pytest-cov", "pytest-rerunfailures", "pytest-xdist", "wurlitzer"] +[[package]] +name = "coverage" +version = "7.15.4" +description = "Code coverage measurement for Python" +optional = false +python-versions = ">=3.10" +groups = ["dev"] +markers = "sys_platform == \"win32\" or sys_platform == \"emscripten\" or sys_platform != \"win32\" and sys_platform != \"emscripten\"" +files = [ + {file = "coverage-7.15.4-cp310-cp310-macosx_10_9_x86_64.whl", hash = "sha256:d0be6daac4cce6b8c8dc65886bae1b082ddbca4da8e5cbb5e15166acf253e264"}, + {file = "coverage-7.15.4-cp310-cp310-macosx_11_0_arm64.whl", hash = "sha256:b24e078eabcd6a9caa8b0713f9bc1eeb310bcc960a29d45a3b4fcd4b16d5b11d"}, + {file = "coverage-7.15.4-cp310-cp310-manylinux1_i686.manylinux_2_28_i686.manylinux_2_5_i686.whl", hash = "sha256:cfe20cc8cf8821d4fe54f89106cbf06aa27f37b5bbe3535568065a81539b4150"}, + {file = "coverage-7.15.4-cp310-cp310-manylinux1_x86_64.manylinux_2_28_x86_64.manylinux_2_5_x86_64.whl", hash = "sha256:83cf06cdd687677742caff1a9134833b7a8b75f111519d2cb0e0ba1b9a851e15"}, + {file = "coverage-7.15.4-cp310-cp310-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:8fa4de68e2a752468ff14b4e15db7def689a71be759e826a31ccecbef69c5fd0"}, + {file = "coverage-7.15.4-cp310-cp310-manylinux2014_ppc64le.manylinux_2_17_ppc64le.manylinux_2_28_ppc64le.whl", hash = "sha256:4dff9daa47d83120c3ec38ce921214242944a832aa04e903e50b5b7ebac8972d"}, + {file = "coverage-7.15.4-cp310-cp310-manylinux_2_31_riscv64.manylinux_2_39_riscv64.whl", hash = "sha256:a093fd37229918976f602aa07aa59e0973cde82186f220c8e197f721f5be0ce4"}, + {file = "coverage-7.15.4-cp310-cp310-musllinux_1_2_aarch64.whl", hash = "sha256:317db01a2cb02552fd67e2b1cca77a4b528a2a277176c5e0bf2cecbb639d3f54"}, + {file = "coverage-7.15.4-cp310-cp310-musllinux_1_2_i686.whl", hash = "sha256:8ee3838dcb656602c3b51e16aed9bfb0822f8d8d6d1c5966d32ec8c104be8e20"}, + {file = "coverage-7.15.4-cp310-cp310-musllinux_1_2_ppc64le.whl", hash = "sha256:425920379052ff1fe465268f3361d35804a241bbdd5a1b592c8cb60df4c52325"}, + {file = "coverage-7.15.4-cp310-cp310-musllinux_1_2_riscv64.whl", hash = "sha256:69bb2400abef928e365ea7d4d9925169ada78ed2295546780002d4b65de3df88"}, + {file = "coverage-7.15.4-cp310-cp310-musllinux_1_2_x86_64.whl", hash = "sha256:81661f82d302484e3119e7c80c519c02fa9bcc2a6b339baf67d67bc89c580f04"}, + {file = "coverage-7.15.4-cp310-cp310-win32.whl", hash = "sha256:cb476b2e828ecb71cb6b6a928d23fd20a7ddb501188022dae1c37499149cc338"}, + {file = "coverage-7.15.4-cp310-cp310-win_amd64.whl", hash = "sha256:3fc2130bf37df31852a8384f12601563a45a0024bccc6624f38355cba7a8d360"}, + {file = "coverage-7.15.4-cp311-cp311-macosx_10_9_x86_64.whl", hash = "sha256:bbac5abad70df71019988f83f26ac7092ff2642975def4429e98dc7585ef3490"}, + {file = "coverage-7.15.4-cp311-cp311-macosx_11_0_arm64.whl", hash = "sha256:357a173465c7ce028d07a95cc2b63b5bf59f50ecdd5ad75c5cbb78ada984048e"}, + {file = "coverage-7.15.4-cp311-cp311-manylinux1_i686.manylinux_2_28_i686.manylinux_2_5_i686.whl", hash = "sha256:21b803935e2efc3acebe9697197a294fccf5dc4e5382bd6369542ff7a7d2a1d7"}, + {file = "coverage-7.15.4-cp311-cp311-manylinux1_x86_64.manylinux_2_28_x86_64.manylinux_2_5_x86_64.whl", hash = "sha256:7a2b580774a4786c1053157c0165e04476e03ff293993d7c148eee784a94bae6"}, + {file = "coverage-7.15.4-cp311-cp311-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:a9464451c4efffe8d47ace5a540b10b0dc10e879066290f8600872b7f54a419d"}, + {file = "coverage-7.15.4-cp311-cp311-manylinux2014_ppc64le.manylinux_2_17_ppc64le.manylinux_2_28_ppc64le.whl", hash = "sha256:de602f34123c2f4af1c1869c6dbbbd60da6d5983bf01937367295d135cccbfce"}, + {file = "coverage-7.15.4-cp311-cp311-manylinux_2_31_riscv64.manylinux_2_39_riscv64.whl", hash = "sha256:6879ded16a27f3eeca19b900c147e81616e7054db451471a611b2755ee5249f7"}, + {file = "coverage-7.15.4-cp311-cp311-musllinux_1_2_aarch64.whl", hash = "sha256:986be58c3ab54aae8d3496a6225eea74f760fdbe739b38bd442c7e8d133aa53b"}, + {file = "coverage-7.15.4-cp311-cp311-musllinux_1_2_i686.whl", hash = "sha256:c6103639613fe6c1e989082948419bc77a2d26b6c825c99d7fad25f7d3d87afc"}, + {file = "coverage-7.15.4-cp311-cp311-musllinux_1_2_ppc64le.whl", hash = "sha256:d3af93dddb5659276c63bc16ac6466ac2033a70ca816097bbc06345b8ccdf571"}, + {file = "coverage-7.15.4-cp311-cp311-musllinux_1_2_riscv64.whl", hash = "sha256:b10075e5421d04265766a6d1dac809bbeb8a946fbb23c8f82c227409b2190719"}, + {file = "coverage-7.15.4-cp311-cp311-musllinux_1_2_x86_64.whl", hash = "sha256:a67a9f78b2942d87ba8ce3059c642164d2aedd65337377fb52fe9803656bc5c7"}, + {file = "coverage-7.15.4-cp311-cp311-win32.whl", hash = "sha256:69484d1aca26e322e1c3ce03f09341e84524ababad2d7202161738d83cc9f82e"}, + {file = "coverage-7.15.4-cp311-cp311-win_amd64.whl", hash = "sha256:63fd6fcd1dd6e158f7eb78606e72933b3f6d01e7b747f99c6c12d764307a0fdc"}, + {file = "coverage-7.15.4-cp311-cp311-win_arm64.whl", hash = "sha256:ea82116c9893fa89e929b7f197ee5a1950a76e91cc5c85ba503fc02379d04890"}, + {file = "coverage-7.15.4-cp312-cp312-macosx_10_13_x86_64.whl", hash = "sha256:d4fedd1f7f428f9fe83b1ead5e7cc87a43427be31aadafbac3ac0636dc7abb22"}, + {file = "coverage-7.15.4-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:37e2f0cdf58e2e1fed4e4d5a8f8786ae2f7eb80b478016876667dc4a01d60a97"}, + {file = "coverage-7.15.4-cp312-cp312-manylinux1_i686.manylinux_2_28_i686.manylinux_2_5_i686.whl", hash = "sha256:fb55d0e70bb15f2e81477613627286581414693d74ac7963c93a790dd453ca9d"}, + {file = "coverage-7.15.4-cp312-cp312-manylinux1_x86_64.manylinux_2_28_x86_64.manylinux_2_5_x86_64.whl", hash = "sha256:899b9da30f3c6c336566e3707495bb23e8302d39d862f01fa78c48b99b9437e2"}, + {file = "coverage-7.15.4-cp312-cp312-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:d15715e8c46552827e5e4f30a35575a2dbcad14454cf3284c54483946bd16931"}, + {file = "coverage-7.15.4-cp312-cp312-manylinux2014_ppc64le.manylinux_2_17_ppc64le.manylinux_2_28_ppc64le.whl", hash = "sha256:002a438859f7b430bc99afeaf01a6d187dad1d0dc907b64cdeffc632a5db8fd8"}, + {file = "coverage-7.15.4-cp312-cp312-manylinux_2_31_riscv64.manylinux_2_39_riscv64.whl", hash = "sha256:e4193a04b518f7968f3099755f5509ee7cccc6dc2b92a6b14841934d22e222c9"}, + {file = "coverage-7.15.4-cp312-cp312-musllinux_1_2_aarch64.whl", hash = "sha256:e98dcc55d572b38e69d117da7e8e8efb8500f1f5eaf81ecd460a63220790b839"}, + {file = "coverage-7.15.4-cp312-cp312-musllinux_1_2_i686.whl", hash = "sha256:af6c538498ce66c10d3fd541c2a8d5b03da5850355add34e6cba564210cb9e72"}, + {file = "coverage-7.15.4-cp312-cp312-musllinux_1_2_ppc64le.whl", hash = "sha256:1d10025d96ea89fc2f73714dbc4cbd433fe012c1ac9e23f895d7728b238b6e52"}, + {file = "coverage-7.15.4-cp312-cp312-musllinux_1_2_riscv64.whl", hash = "sha256:d802e1947603162ded419bff83ac7489820355d2b856dfb09206574e3a37ac0c"}, + {file = "coverage-7.15.4-cp312-cp312-musllinux_1_2_x86_64.whl", hash = "sha256:c2de40895718f91951b86712b4c5b694acaf9a0a49be13874896f599a1eed3f4"}, + {file = "coverage-7.15.4-cp312-cp312-win32.whl", hash = "sha256:5c3431b2161279b7db5c2a1aa58ae02e5cb8c3c42d93a5094be3f5537bd5b11b"}, + {file = "coverage-7.15.4-cp312-cp312-win_amd64.whl", hash = "sha256:6befeab5fb2b51c958ca4ac6c5d141a1e8240f4f76e46350f1911963deda49cd"}, + {file = "coverage-7.15.4-cp312-cp312-win_arm64.whl", hash = "sha256:67bc345491ab55b837277d76f5775d057e8c7f1ac44d890d8c2c82adde258c6f"}, + {file = "coverage-7.15.4-cp313-cp313-macosx_10_13_x86_64.whl", hash = "sha256:c705b28feb2775dc82a25f1d473a370bc37ff93f5177f4e29ce2425f560f6921"}, + {file = "coverage-7.15.4-cp313-cp313-macosx_11_0_arm64.whl", hash = "sha256:3ff205ab5e3ecc670f6a4dd19d9cbf12ede53dd41cfc1e15716ec961ea6d314e"}, + {file = "coverage-7.15.4-cp313-cp313-manylinux1_i686.manylinux_2_28_i686.manylinux_2_5_i686.whl", hash = "sha256:5172326e861a38b48b48befca15e0f477a26b283337a33a739c8fed229934e36"}, + {file = "coverage-7.15.4-cp313-cp313-manylinux1_x86_64.manylinux_2_28_x86_64.manylinux_2_5_x86_64.whl", hash = "sha256:12b59c90084e3234fb11184886bf4a40f4f16a8c8f867be2e087b81f8e8868d4"}, + {file = "coverage-7.15.4-cp313-cp313-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:349062d66f00b40fa2c1c222438bad25fabf755631b5d82937fe985c8008615c"}, + {file = "coverage-7.15.4-cp313-cp313-manylinux2014_ppc64le.manylinux_2_17_ppc64le.manylinux_2_28_ppc64le.whl", hash = "sha256:4256ced708e598e05209bc1a8ab4074e04a51dba4c62fb45926a229af675ace7"}, + {file = "coverage-7.15.4-cp313-cp313-manylinux_2_31_riscv64.manylinux_2_39_riscv64.whl", hash = "sha256:d80f974b20782d9612c8b4c9beeca867074c7cf4079d1419843fa25a26428b25"}, + {file = "coverage-7.15.4-cp313-cp313-musllinux_1_2_aarch64.whl", hash = "sha256:2e179f19bfe1d31f8eeeaa12990194d761c4f62f0759661000bca6cd8729f40b"}, + {file = "coverage-7.15.4-cp313-cp313-musllinux_1_2_i686.whl", hash = "sha256:8bc16bb47b7679670eceff71d78bfb7d6e5b143f6c2cd117487ec7c75e0d4b78"}, + {file = "coverage-7.15.4-cp313-cp313-musllinux_1_2_ppc64le.whl", hash = "sha256:1cd685005cd2c4200adfc14cf39a603b9320efab3f18a8f7f156d20c9cc3345f"}, + {file = "coverage-7.15.4-cp313-cp313-musllinux_1_2_riscv64.whl", hash = "sha256:337399ad2c93b3acd2a937627dae8b3e86b66707cd3d3e856347999aadf1ef8d"}, + {file = "coverage-7.15.4-cp313-cp313-musllinux_1_2_x86_64.whl", hash = "sha256:96e257121228ec5cd2bb919276e94ac11074471bc37d68dbae0e8308cce15fff"}, + {file = "coverage-7.15.4-cp313-cp313-win32.whl", hash = "sha256:c65a9e0dfc6143491879da4e13b5e30f8be192055de508d737fb14601edbd22c"}, + {file = "coverage-7.15.4-cp313-cp313-win_amd64.whl", hash = "sha256:2ff8f5e9b8f7a94f0c11c45631eee103dbcb7d63274edd12c56efe1be690b3b4"}, + {file = "coverage-7.15.4-cp313-cp313-win_arm64.whl", hash = "sha256:6e0a8a5083b096487d6cfced94cdd514d8f5db6f113610fb36c0620edb1028cf"}, + {file = "coverage-7.15.4-cp314-cp314-macosx_10_15_x86_64.whl", hash = "sha256:770e9325ab5ea6d56f77e59b29ecfe0ac20b57a82a601876f90494a4dda0386f"}, + {file = "coverage-7.15.4-cp314-cp314-macosx_11_0_arm64.whl", hash = "sha256:d12b33a3a50a1676b7784dc8d00a0c6d66a9f2add4b85a041c19b6a7e53ef23c"}, + {file = "coverage-7.15.4-cp314-cp314-manylinux1_i686.manylinux_2_28_i686.manylinux_2_5_i686.whl", hash = "sha256:5669c8378ebde86f5def7a25d29586631b58acc27ffde04399f678f3dfc6e082"}, + {file = "coverage-7.15.4-cp314-cp314-manylinux1_x86_64.manylinux_2_28_x86_64.manylinux_2_5_x86_64.whl", hash = "sha256:ff97a14362eef486483ed44042ca2027ea257df6ff768e62358ee0c9776925ac"}, + {file = "coverage-7.15.4-cp314-cp314-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:5a325e815318638aed1655d9c06e6d7c2d3d46c09231ce988070428a8762d734"}, + {file = "coverage-7.15.4-cp314-cp314-manylinux2014_ppc64le.manylinux_2_17_ppc64le.manylinux_2_28_ppc64le.whl", hash = "sha256:474223409d88eb20d2d6a0d37ea60e8647a65a90cc008dc1f0410af5f64f1e0d"}, + {file = "coverage-7.15.4-cp314-cp314-manylinux_2_31_riscv64.manylinux_2_39_riscv64.whl", hash = "sha256:7f2f62ae3cd189dd2e13aece758c57b3eecbd27be070dbd4cbd10936049e5dbf"}, + {file = "coverage-7.15.4-cp314-cp314-musllinux_1_2_aarch64.whl", hash = "sha256:39ece820e29e0a2ba34b3ecb3be83c27e997eed8926f2ba6fe7ce7a0bda5843b"}, + {file = "coverage-7.15.4-cp314-cp314-musllinux_1_2_i686.whl", hash = "sha256:f21b56dcace11dfe013014201f577dcd592b2a9b72182d930361b47cf6f73f25"}, + {file = "coverage-7.15.4-cp314-cp314-musllinux_1_2_ppc64le.whl", hash = "sha256:93a3a0b662abcc10c73a47cbc72cd60f63618d6989fb2d1286e50eacd974f303"}, + {file = "coverage-7.15.4-cp314-cp314-musllinux_1_2_riscv64.whl", hash = "sha256:141fae2cabf5569b782c10afc4c850ce10f618c13f8db54765cba99cc839da1f"}, + {file = "coverage-7.15.4-cp314-cp314-musllinux_1_2_x86_64.whl", hash = "sha256:81294c7e6ab30c5f74c0353b11b2fd6320e72d9bee6ac73b357caa8b916323a5"}, + {file = "coverage-7.15.4-cp314-cp314-win32.whl", hash = "sha256:7bbd7d6418e0dab31a206af5203bd43ae36edb8e7fba1940b055d3e9249290d7"}, + {file = "coverage-7.15.4-cp314-cp314-win_amd64.whl", hash = "sha256:f0204ed122758782970526057093f448051a39db9d810d4e344bb87a3546f425"}, + {file = "coverage-7.15.4-cp314-cp314-win_arm64.whl", hash = "sha256:9e71e7bc71c686a123347ae47a0de33a175e797a85bb57b791492adf4eec8ed8"}, + {file = "coverage-7.15.4-cp314-cp314t-macosx_10_15_x86_64.whl", hash = "sha256:7c922735321eef3f87c280a3d39afff6b646723a2880b862cda4ac7a093b8aa8"}, + {file = "coverage-7.15.4-cp314-cp314t-macosx_11_0_arm64.whl", hash = "sha256:f41c17c4668a655ce96d090d8d5ffdc24ef64b5a02f9753884d08483e8a4a41a"}, + {file = "coverage-7.15.4-cp314-cp314t-manylinux1_i686.manylinux_2_28_i686.manylinux_2_5_i686.whl", hash = "sha256:46822e9b6ff1c6a72b518c162c44a8f45a61a1d609c51084bf5b16c023c5037b"}, + {file = "coverage-7.15.4-cp314-cp314t-manylinux1_x86_64.manylinux_2_28_x86_64.manylinux_2_5_x86_64.whl", hash = "sha256:3d6f4955b73b5445271379a59e3792b0d978f42d4a01e0cf7a67d9c33a3bb0a5"}, + {file = "coverage-7.15.4-cp314-cp314t-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:3fc9e047706fb4a9abb54f719d3aa643e80e5bb3818182c40aee01ac0f0247ba"}, + {file = "coverage-7.15.4-cp314-cp314t-manylinux2014_ppc64le.manylinux_2_17_ppc64le.manylinux_2_28_ppc64le.whl", hash = "sha256:05e491d4f3165d62d4f5c8fd48dfeabf2ae8f42cbbd484319af33ea851b78982"}, + {file = "coverage-7.15.4-cp314-cp314t-manylinux_2_31_riscv64.manylinux_2_39_riscv64.whl", hash = "sha256:226c66e80ec0598d3b9b4874123df167ccca342aca8714f77cac6829688ee09c"}, + {file = "coverage-7.15.4-cp314-cp314t-musllinux_1_2_aarch64.whl", hash = "sha256:ac41cc14bebda0dbfb0628036b7f75706935c95bcc07fefe9a0f93614aa60a57"}, + {file = "coverage-7.15.4-cp314-cp314t-musllinux_1_2_i686.whl", hash = "sha256:8af623e5cd92080acddd02b38f2f406a2c3a0893c38950b211890361448fbf26"}, + {file = "coverage-7.15.4-cp314-cp314t-musllinux_1_2_ppc64le.whl", hash = "sha256:07545711d4f0f32852a18f18ad11f76f0109909d09e78b9008b4cfc67e829429"}, + {file = "coverage-7.15.4-cp314-cp314t-musllinux_1_2_riscv64.whl", hash = "sha256:a0865421cfdc53654b342d515e5a233187590882d20b95752150e53f65460017"}, + {file = "coverage-7.15.4-cp314-cp314t-musllinux_1_2_x86_64.whl", hash = "sha256:460115e32ee40566476db5048f9bec1e842c127ad8e6f8be745aad3ac9cbc839"}, + {file = "coverage-7.15.4-cp314-cp314t-win32.whl", hash = "sha256:cbde877ef9dd7baf272b9bfef2b8a25edd45d9170fc326951dd20eb480335e85"}, + {file = "coverage-7.15.4-cp314-cp314t-win_amd64.whl", hash = "sha256:3da9e92d1c551fd7563833e9ade686efb0c4b7363ab7681a94283958c950bf5e"}, + {file = "coverage-7.15.4-cp314-cp314t-win_arm64.whl", hash = "sha256:3a54f5a0d85050c73a38f6793090ee83974531e67fe5e57a1da9bee11398aa5e"}, + {file = "coverage-7.15.4-cp315-cp315-macosx_10_15_x86_64.whl", hash = "sha256:2c9872e4d9dc5d3cf616bf4b382f5a00359305a5be666a3dd0b5cdb4e49597f9"}, + {file = "coverage-7.15.4-cp315-cp315-macosx_11_0_arm64.whl", hash = "sha256:e101dbb4b9b72f0cddd8cdc8c9c5b47f456766f5e0ac82dbfb75e5c55409b78a"}, + {file = "coverage-7.15.4-cp315-cp315-manylinux1_i686.manylinux_2_28_i686.manylinux_2_5_i686.whl", hash = "sha256:7d1abebdb047729e852b9c77a00497dfbeb11eb3a117e037d7dbc3ac8e5f5c54"}, + {file = "coverage-7.15.4-cp315-cp315-manylinux1_x86_64.manylinux_2_28_x86_64.manylinux_2_5_x86_64.whl", hash = "sha256:d28a4a899354d0ea6214cc59b4fa19eefbce1b9ff1688ab579acf49e894bd3fb"}, + {file = "coverage-7.15.4-cp315-cp315-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:ffb3c2aacea411cc7e1d27712490c11108e2de1d39019ae32915493a59a8b9ed"}, + {file = "coverage-7.15.4-cp315-cp315-manylinux2014_ppc64le.manylinux_2_17_ppc64le.manylinux_2_28_ppc64le.whl", hash = "sha256:a9447978a92f405d301123cfd39ff49895490efb769a758fe2734c7f631bf8ce"}, + {file = "coverage-7.15.4-cp315-cp315-manylinux_2_31_riscv64.manylinux_2_39_riscv64.whl", hash = "sha256:050467a7983b8e2fe7dd41a78bb30c3e7f8c0b8cafda14b1c46f8b5e3cf2dd3c"}, + {file = "coverage-7.15.4-cp315-cp315-musllinux_1_2_aarch64.whl", hash = "sha256:d003b7a5708ddad5c206c79607a6b92abb6fc13c57d99d8a4468cc03a2941ced"}, + {file = "coverage-7.15.4-cp315-cp315-musllinux_1_2_i686.whl", hash = "sha256:c38efe30fd74e5c19e9433f11fb1f5dc9c6522770971b7c6145bbaa413dc8800"}, + {file = "coverage-7.15.4-cp315-cp315-musllinux_1_2_ppc64le.whl", hash = "sha256:1f4f826d70f772ab8b0c052329580d7fe8b8abd191e4ce0c8f81aec6614665d3"}, + {file = "coverage-7.15.4-cp315-cp315-musllinux_1_2_riscv64.whl", hash = "sha256:4a4bf917c9953f57c957be31c1cd504e3bd2f34d4a352b9d391a3025336f6768"}, + {file = "coverage-7.15.4-cp315-cp315-musllinux_1_2_x86_64.whl", hash = "sha256:1c9bf40ebef178a45192c75c4964760bb261b0e6ad725da5fc4c93f674f19753"}, + {file = "coverage-7.15.4-cp315-cp315-win32.whl", hash = "sha256:43619d04c3671792d2c4706ae8bf45e265dc87bbd4078189ef8b847ea1e74be2"}, + {file = "coverage-7.15.4-cp315-cp315-win_amd64.whl", hash = "sha256:be619439dbcd31a2eab10b32de9fff62c26ed4bab69dc32b8363fdaaa0882809"}, + {file = "coverage-7.15.4-cp315-cp315-win_arm64.whl", hash = "sha256:def597967dafc2e8d97c9097ea453c464e0bb8ed38f193a43070f10dc623bb6d"}, + {file = "coverage-7.15.4-cp315-cp315t-macosx_10_15_x86_64.whl", hash = "sha256:c7dbc748ac8a1e3e59a2b28bea47675e6e778081dbbf081bde0d75def2fcbe1d"}, + {file = "coverage-7.15.4-cp315-cp315t-macosx_11_0_arm64.whl", hash = "sha256:2413074a5ecbb61a01a7888fc72db0ca324d13588c5b38bc0dd8564cdcdfea26"}, + {file = "coverage-7.15.4-cp315-cp315t-manylinux1_i686.manylinux_2_28_i686.manylinux_2_5_i686.whl", hash = "sha256:4e6f6f632b7b2f714bf7a1346e8f97b650ee71f3c298aaad42a2ab60f0f07645"}, + {file = "coverage-7.15.4-cp315-cp315t-manylinux1_x86_64.manylinux_2_28_x86_64.manylinux_2_5_x86_64.whl", hash = "sha256:8df457da2249d3c75ca2e5e835d59c725abfe92d27fdff6cd99eed85b51d5e9a"}, + {file = "coverage-7.15.4-cp315-cp315t-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:050f66a08805acb5b8a23c6d4a517b1ecf82c08e81ed0e4bd727df065e5c6624"}, + {file = "coverage-7.15.4-cp315-cp315t-manylinux2014_ppc64le.manylinux_2_17_ppc64le.manylinux_2_28_ppc64le.whl", hash = "sha256:1587fb771d1ccceef708fdde1e5af8c7ed24b486b61d13a321acb7d8145390aa"}, + {file = "coverage-7.15.4-cp315-cp315t-manylinux_2_31_riscv64.manylinux_2_39_riscv64.whl", hash = "sha256:8b4f1c3a69ca580f3fbd6b2046915f536d7f586874f25c1bb23add2a3c88d50f"}, + {file = "coverage-7.15.4-cp315-cp315t-musllinux_1_2_aarch64.whl", hash = "sha256:ffb58d7eff5b7f6ecc6fa21d6288ab7f968a212cb67d682c269c09b9eba3b66f"}, + {file = "coverage-7.15.4-cp315-cp315t-musllinux_1_2_i686.whl", hash = "sha256:d9df165544774574ee004b953023d1bebada1894a80b1052a43d798b0f676e67"}, + {file = "coverage-7.15.4-cp315-cp315t-musllinux_1_2_ppc64le.whl", hash = "sha256:f9de0a24a4079b53e523b5c5e2c5945ec251ab486652659955187cf255a259bc"}, + {file = "coverage-7.15.4-cp315-cp315t-musllinux_1_2_riscv64.whl", hash = "sha256:150089274bdc9f940628552cb92844e0223c987f1902ab8efe9f45a2ec758d88"}, + {file = "coverage-7.15.4-cp315-cp315t-musllinux_1_2_x86_64.whl", hash = "sha256:a58a94fed5da6997d258e8f7668c1e195fbd04a691d781b7558f1e468f9e68bc"}, + {file = "coverage-7.15.4-cp315-cp315t-win32.whl", hash = "sha256:ebd5a6d8466ff30836572f3ba2cae8a5e8f85029b1c6d5e2ed338dc472a5166a"}, + {file = "coverage-7.15.4-cp315-cp315t-win_amd64.whl", hash = "sha256:288bde2a2d7ab6b6c2d7252fcde8b524387f2d970bdba9658fc6f8bbcaef0f9b"}, + {file = "coverage-7.15.4-cp315-cp315t-win_arm64.whl", hash = "sha256:68be5e1de60ff13c9095bbec0e5a7fa45b33b101752215b91345ea1f61c4a278"}, + {file = "coverage-7.15.4-py3-none-any.whl", hash = "sha256:964730a1e9de9c0cf11be6a1a3c79ce419c34882842abd256086ba4698705e84"}, + {file = "coverage-7.15.4.tar.gz", hash = "sha256:0548198fff07ccf4faf469520bce1c2eceb1ce3e62891921138dec10907f9d00"}, +] + +[package.extras] +toml = ["tomli ; python_full_version <= \"3.11.0a6\""] + [[package]] name = "cycler" version = "0.12.1" @@ -1540,6 +1675,27 @@ pygments = ">=2.7.2" [package.extras] dev = ["argcomplete", "attrs (>=19.2)", "hypothesis (>=3.56)", "mock", "requests", "setuptools", "xmlschema"] +[[package]] +name = "pytest-cov" +version = "7.1.0" +description = "Pytest plugin for measuring coverage." +optional = false +python-versions = ">=3.9" +groups = ["dev"] +markers = "sys_platform == \"win32\" or sys_platform == \"emscripten\" or sys_platform != \"win32\" and sys_platform != \"emscripten\"" +files = [ + {file = "pytest_cov-7.1.0-py3-none-any.whl", hash = "sha256:a0461110b7865f9a271aa1b51e516c9a95de9d696734a2f71e3e78f46e1d4678"}, + {file = "pytest_cov-7.1.0.tar.gz", hash = "sha256:30674f2b5f6351aa09702a9c8c364f6a01c27aae0c1366ae8016160d1efc56b2"}, +] + +[package.dependencies] +coverage = {version = ">=7.10.6", extras = ["toml"]} +pluggy = ">=1.2" +pytest = ">=7" + +[package.extras] +testing = ["process-tests", "pytest-xdist", "virtualenv"] + [[package]] name = "python-dateutil" version = "2.9.0.post0" @@ -2210,4 +2366,4 @@ files = [ [metadata] lock-version = "2.1" python-versions = ">=3.11" -content-hash = "e629415d5e6d49cd52522d6557fed8e0c4a0a5a0fb9a34e7aa6306c86b8c7e74" +content-hash = "ad4043566644155c3c6bfb2175eaf36526e88b8c6a79afece91dddf8f924d05a" diff --git a/pyproject.toml b/pyproject.toml index f02f17f..99362f5 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -51,6 +51,7 @@ sphinx-gallery = ">=0.17" black = ">=24.8" flake8 = ">=7.1" pytest = ">=8.3" +pytest-cov = ">=5.0" [tool.black] line-length = 120 @@ -64,6 +65,10 @@ filterwarnings = [ "error::FutureWarning", # TrustReport fits a classifier probe even for regression blackboxes "default:The number of unique classes is greater than:UserWarning", + # joblib 1.5.3 (latest) sets ndarray.shape directly in numpy_pickle.py, which + # NumPy 2.5 deprecated. Nothing in Trustee can avoid it -- it fires on any + # joblib load. Scoped to the exact message so the category stays an error. + "default:Setting the shape on a NumPy array has been deprecated:DeprecationWarning", ] [build-system] diff --git a/tests/conftest.py b/tests/conftest.py index cc3cf3d..975af62 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -1,6 +1,7 @@ """Shared fixtures for the Trustee test suite.""" import matplotlib +import numpy as np import pytest from sklearn import datasets from sklearn.ensemble import RandomForestClassifier @@ -13,6 +14,17 @@ RANDOM_STATE = 0 +@pytest.fixture(autouse=True) +def deterministic_random(): + """Seed the global RNG before every test. + + Trustee.fit() samples with np.random.choice and calls train_test_split without a + random_state, so both draw from NumPy's global state. Without this, results differ + run to run and any assertion on fidelity or agreement is flaky. + """ + np.random.seed(RANDOM_STATE) + + @pytest.fixture(scope="session") def iris(): """The iris dataset as plain numpy arrays.""" diff --git a/tests/test_const.py b/tests/test_const.py new file mode 100644 index 0000000..91b5445 --- /dev/null +++ b/tests/test_const.py @@ -0,0 +1,97 @@ +"""Structural validation of the bundled dataset metadata. + +These dicts are consumed by `trustee.utils.dataset.read()`, which unpacks every +field as a 4-tuple and indexes by position. A malformed entry would surface as an +opaque unpacking error at read time, so validate the shape here instead. +""" + +import pytest + +from trustee.enums.feature_type import FeatureType +from trustee.utils import const + +METADATA = {name: getattr(const, name) for name in dir(const) if name.endswith("_DATASET_META")} + + +def test_metadata_dicts_are_exposed(): + assert METADATA, "expected at least one *_DATASET_META constant" + + +@pytest.mark.parametrize("name", sorted(METADATA)) +class TestDatasetMetadata: + def test_has_the_keys_read_depends_on(self, name): + meta = METADATA[name] + assert "fields" in meta + assert isinstance(meta["fields"], list) and meta["fields"] + + def test_every_field_is_a_four_tuple(self, name): + for field in METADATA[name]["fields"]: + assert len(field) == 4, f"{name}: {field!r} is not a 4-tuple" + + def test_field_components_have_the_expected_types(self, name): + for field_name, feature_type, _dtype, is_result in METADATA[name]["fields"]: + assert isinstance(field_name, str) and field_name + assert isinstance(feature_type, FeatureType) + assert isinstance(is_result, bool) + + def test_declares_exactly_one_target_column(self, name): + results = [f for f in METADATA[name]["fields"] if f[3]] + assert len(results) == 1, f"{name} declares {len(results)} target columns" + + def test_field_names_are_unique(self, name): + names = [f[0] for f in METADATA[name]["fields"]] + assert len(names) == len(set(names)), f"{name} has duplicate field names" + + def test_target_column_is_not_an_identifier(self, name): + for field_name, feature_type, _dtype, is_result in METADATA[name]["fields"]: + if is_result: + assert feature_type is not FeatureType.IDENTIFIER + + def test_declared_type_is_a_known_string(self, name): + meta = METADATA[name] + if "type" in meta: + assert meta["type"] in {"classification", "regression"} + + +class TestFeatureType: + def test_members_are_distinct(self): + values = [member.value for member in FeatureType] + assert len(values) == len(set(values)) + + def test_exposes_the_three_kinds_read_branches_on(self): + assert {"CATEGORICAL", "NUMERICAL", "IDENTIFIER"} <= {m.name for m in FeatureType} + + +class TestCicIds2017LabelConverter: + convert = staticmethod(const.cic_ids_2017_label_converter) + + def test_maps_benign_to_zero(self): + assert self.convert("BENIGN") == 0 + + @pytest.mark.parametrize( + "label,expected", + [("Bot", 1), ("DDoS", 2), ("Heartbleed", 8), ("PortScan", 10), ("Web Attack XSS", 14)], + ) + def test_maps_known_attack_labels(self, label, expected): + assert self.convert(label) == expected + + def test_strips_surrounding_whitespace(self): + assert self.convert(" DDoS ") == self.convert("DDoS") + + def test_returns_uint8(self): + import numpy as np + + assert isinstance(self.convert("BENIGN"), np.uint8) + + def test_unknown_label_falls_back(self, capsys): + # -1 cast to uint8 wraps to 255, which is the sentinel callers see. + assert self.convert("Not A Real Label") == 255 + assert "Exception" in capsys.readouterr().out + + def test_non_string_input_falls_back(self, capsys): + assert self.convert(None) == 255 + assert "Exception" in capsys.readouterr().out + + def test_every_mapped_value_fits_in_uint8(self): + for label in ["BENIGN", "Bot", "DDoS", "Web Attack Sql Injection"]: + assert 0 <= self.convert(label) <= 255 diff --git a/tests/test_dataset.py b/tests/test_dataset.py index 04529a4..55b299b 100644 --- a/tests/test_dataset.py +++ b/tests/test_dataset.py @@ -112,3 +112,89 @@ def test_last_column_is_the_target_by_default(self, csv_path): _, y, columns, _, _ = read(path, metadata=metadata, as_df=True) assert "label" not in list(columns) assert np.asarray(y).ravel().tolist() == [0, 1] + + +class TestReadVerbose: + def test_verbose_prints_a_summary(self, csv_path, capsys): + read(csv_path, metadata=METADATA, verbose=True, as_df=True) + out = capsys.readouterr().out + assert "Metadata start." in out + assert "Pandas read_csv complete." in out + assert "Total memory usage" in out + + def test_verbose_routes_through_a_logger(self, csv_path, tmp_path): + from trustee.utils.log import Logger + + log_file = tmp_path / "read.log" + read(csv_path, metadata=METADATA, verbose=True, logger=Logger(path=str(log_file)), as_df=True) + assert "Metadata start." in log_file.read_text() + + +class TestReadCategories: + def test_applies_an_ordered_categorical_dtype(self, tmp_path): + path = tmp_path / "sized.csv" + path.write_text("size,score\nsmall,1\nlarge,3\nmedium,2\n") + metadata = { + "has_header": True, + "fields": [ + ("size", FeatureType.NUMERICAL, None, False), + ("score", FeatureType.NUMERICAL, None, True), + ], + "categories": {"size": ["small", "medium", "large"]}, + } + X, _, _, _, _ = read(str(path), metadata=metadata, as_df=True) + assert str(X["size"].dtype) == "category" + assert X["size"].cat.ordered + assert list(X["size"].cat.categories) == ["small", "medium", "large"] + + +class TestReadDelimiter: + def test_honours_a_custom_delimiter(self, tmp_path): + path = tmp_path / "semi.csv" + path.write_text("a;b\n1;2\n3;4\n") + metadata = { + "has_header": True, + "delimiter": ";", + "fields": [ + ("a", FeatureType.NUMERICAL, None, False), + ("b", FeatureType.NUMERICAL, None, True), + ], + } + X, y, _, _, _ = read(str(path), metadata=metadata, as_df=True) + assert list(X.columns) == ["a"] + assert np.asarray(y).ravel().tolist() == [2, 4] + + +class TestReadDirectory: + def test_concatenates_every_csv_in_a_directory(self, tmp_path): + data_dir = tmp_path / "parts" + data_dir.mkdir() + (data_dir / "part1.csv").write_text("a,b\n1,2\n3,4\n") + (data_dir / "part2.csv").write_text("a,b\n5,6\n") + metadata = { + "has_header": True, + "is_dir": True, + "fields": [ + ("a", FeatureType.NUMERICAL, None, False), + ("b", FeatureType.NUMERICAL, None, True), + ], + } + X, y, _, _, _ = read(str(data_dir), metadata=metadata, as_df=True) + assert len(X) == 3 + assert sorted(np.asarray(y).ravel().tolist()) == [2, 4, 6] + + +class TestReadConverters: + def test_applies_a_converter_per_column(self, tmp_path): + path = tmp_path / "conv.csv" + path.write_text("label,score\nBENIGN,1\nBot,2\n") + metadata = { + "has_header": True, + "fields": [ + ("label", FeatureType.NUMERICAL, None, False), + ("score", FeatureType.NUMERICAL, None, True), + ], + "converters": {"label": lambda v: {"BENIGN": 0, "Bot": 1}.get(v.strip(), -1)}, + } + X, _, _, _, _ = read(str(path), metadata=metadata, as_df=True) + assert X["label"].tolist() == [0, 1] diff --git a/tests/test_log.py b/tests/test_log.py new file mode 100644 index 0000000..313bb3d --- /dev/null +++ b/tests/test_log.py @@ -0,0 +1,67 @@ +"""Tests for the Logger helper.""" + +import logging + +import pytest + +from trustee.utils.log import Logger + + +@pytest.fixture +def log_path(tmp_path): + return str(tmp_path / "output.log") + + +class TestLogger: + def test_writes_the_message_to_its_file(self, log_path): + Logger(path=log_path).log("hello world") + with open(log_path) as handle: + assert "hello world" in handle.read() + + def test_joins_multiple_arguments_with_spaces(self, log_path): + Logger(path=log_path).log("alpha", "beta", "gamma") + with open(log_path) as handle: + assert "alpha beta gamma" in handle.read() + + def test_stringifies_non_string_arguments(self, log_path): + Logger(path=log_path).log("count:", 42, ["a", "b"], {"k": 1}) + with open(log_path) as handle: + body = handle.read() + assert "count: 42" in body + assert "['a', 'b']" in body + + def test_includes_the_level_name(self, log_path): + Logger(path=log_path).log("an informational message") + with open(log_path) as handle: + assert "INFO" in handle.read() + + def test_honours_an_explicit_level(self, log_path): + Logger(path=log_path).log("a warning message", level=logging.WARNING) + with open(log_path) as handle: + assert "WARNING" in handle.read() + + def test_filters_below_the_configured_level(self, log_path): + Logger(path=log_path, level=logging.ERROR).log("should not appear", level=logging.INFO) + with open(log_path) as handle: + assert "should not appear" not in handle.read() + + def test_installs_a_stream_and_a_file_handler(self, log_path): + logger = Logger(path=log_path) + kinds = {type(h) for h in logger.handlers} + assert logging.FileHandler in kinds + assert any(issubclass(k, logging.StreamHandler) for k in kinds) + + def test_is_usable_as_a_standard_logger(self, log_path): + logger = Logger(path=log_path) + assert isinstance(logger, logging.Logger) + logger.info("via the standard interface") + with open(log_path) as handle: + assert "via the standard interface" in handle.read() + + def test_appends_across_calls(self, log_path): + logger = Logger(path=log_path) + logger.log("first") + logger.log("second") + with open(log_path) as handle: + body = handle.read() + assert "first" in body and "second" in body diff --git a/tests/test_persist.py b/tests/test_persist.py new file mode 100644 index 0000000..b6592ea --- /dev/null +++ b/tests/test_persist.py @@ -0,0 +1,99 @@ +"""Tests for model persistence.""" + +import zipfile + +import pytest +from sklearn.tree import DecisionTreeClassifier + +from trustee.utils.persist import load_model, save_model + + +@pytest.fixture +def model(iris): + X, y = iris + return DecisionTreeClassifier(random_state=0, max_depth=3).fit(X, y) + + +class TestSaveModel: + def test_writes_a_file(self, model, tmp_path): + path = tmp_path / "model.joblib" + save_model(model, str(path)) + assert path.is_file() + assert path.stat().st_size > 0 + + def test_returns_the_paths_written(self, model, tmp_path): + assert save_model(model, str(tmp_path / "model.joblib")) + + +class TestRoundTrip: + def test_restores_an_equivalent_model(self, model, iris, tmp_path): + X, _ = iris + path = tmp_path / "model.joblib" + save_model(model, str(path)) + restored = load_model(str(path)) + + assert restored is not None + assert restored.predict(X).tolist() == model.predict(X).tolist() + + def test_restores_tree_structure(self, model, tmp_path): + path = tmp_path / "model.joblib" + save_model(model, str(path)) + restored = load_model(str(path)) + assert restored.tree_.node_count == model.tree_.node_count + + +class TestZipSupport: + def test_loads_a_model_out_of_a_zip(self, model, tmp_path): + inner = tmp_path / "model" + save_model(model, str(inner)) + + archive = tmp_path / "model.zip" + with zipfile.ZipFile(archive, "w") as zf: + zf.write(inner, arcname="model") + inner.unlink() + + restored = load_model(str(archive)) + assert restored is not None + assert restored.tree_.node_count == model.tree_.node_count + + def test_cleans_up_the_file_it_unzipped(self, model, tmp_path): + inner = tmp_path / "model" + save_model(model, str(inner)) + archive = tmp_path / "model.zip" + with zipfile.ZipFile(archive, "w") as zf: + zf.write(inner, arcname="model") + inner.unlink() + + load_model(str(archive)) + assert not inner.exists(), "the extracted copy should not be left behind" + assert archive.is_file(), "the archive itself must survive" + + +class TestMissingInput: + def test_returns_none_for_a_path_that_is_not_a_file(self, tmp_path): + assert load_model(str(tmp_path)) is None + + def test_returns_none_for_a_non_model_file(self, tmp_path): + # A readable file that joblib cannot load. + path = tmp_path / "garbage.joblib" + path.write_text("this is not a joblib payload") + with pytest.raises(Exception): + # Documents current behaviour: the except clause catches copy.Error, + # which is not what joblib raises, so the error propagates. + load_model(str(path)) + + def test_missing_path_returns_none(self, tmp_path): + assert load_model(str(tmp_path / "does-not-exist.joblib")) is None + + +class TestOverwrite: + def test_saving_twice_replaces_the_file(self, model, iris, tmp_path): + X, y = iris + path = tmp_path / "model.joblib" + save_model(model, str(path)) + + other = DecisionTreeClassifier(random_state=0, max_depth=1).fit(X, y) + save_model(other, str(path)) + + restored = load_model(str(path)) + assert restored.get_depth() == other.get_depth() diff --git a/tests/test_report.py b/tests/test_report.py index 45932e0..97093d5 100644 --- a/tests/test_report.py +++ b/tests/test_report.py @@ -142,3 +142,101 @@ def test_builds_for_a_regression_blackbox(self, diabetes): **REPORT_KWARGS, ) assert "# Input features:" in str(built) + + +@pytest.fixture(scope="module") +def analysed(iris_frame, iris_bunch): + from sklearn.ensemble import RandomForestClassifier + + X, y = iris_frame + return TrustReport( + RandomForestClassifier(n_estimators=10, random_state=0), + X=X, + y=y, + class_names=iris_bunch.target_names, + feature_names=iris_bunch.feature_names, + is_classify=True, + analyze_branches=True, + analyze_stability=True, + **REPORT_KWARGS, + ) + + +class TestAnalysisPasses: + """The optional branch/stability collection passes and the plotting entry point.""" + + def test_collects_branch_analysis(self, analysed): + assert analysed.branch_iter + + def test_collects_stability_analysis(self, analysed): + assert analysed.stability_iter + assert all("max_dt" in it and "min_dt" in it for it in analysed.stability_iter) + + def test_renders_with_the_extra_sections(self, analysed): + assert "# Input features:" in str(analysed) + + @has_graphviz + def test_plot_writes_the_stability_figures(self, analysed, tmp_path): + analysed.plot(str(tmp_path)) + produced = {p.name for p in (tmp_path / "plots").glob("*.pdf")} + assert any("stability" in name for name in produced) + assert any("branches" in name for name in produced) + + @has_graphviz + def test_save_all_dts_writes_every_tree(self, analysed, tmp_path): + analysed.save(str(tmp_path), save_all_dts=True) + trees = list((tmp_path / "report").glob("**/*.pdf")) + assert len(trees) > 2 + + +class TestSkipRetrain: + def test_skips_the_feature_removal_pass(self, iris_frame, iris_bunch): + """skip_retrain reuses the blackbox as-is, so it must already be fitted.""" + from sklearn.ensemble import RandomForestClassifier + + X, y = iris_frame + report = TrustReport( + RandomForestClassifier(n_estimators=10, random_state=0).fit(X, y), + X=X, + y=y, + class_names=iris_bunch.target_names, + is_classify=True, + skip_retrain=True, + **REPORT_KWARGS, + ) + assert report.whitebox_iter == [] + assert str(report) + + +class TestUseFeatures: + def test_restricts_the_student_to_the_given_columns(self, iris_frame, iris_bunch): + from sklearn.ensemble import RandomForestClassifier + + X, y = iris_frame + report = TrustReport( + RandomForestClassifier(n_estimators=10, random_state=0), + X=X, + y=y, + class_names=iris_bunch.target_names, + is_classify=True, + use_features=[0, 1], + **REPORT_KWARGS, + ) + assert list(report.use_features) == [0, 1] + assert report.max_dt.n_features_in_ == 2 + + def test_unfitted_blackbox_is_rejected(self, iris_frame, iris_bunch): + from sklearn.ensemble import RandomForestClassifier + from sklearn.exceptions import NotFittedError + + X, y = iris_frame + with pytest.raises(NotFittedError): + TrustReport( + RandomForestClassifier(n_estimators=10, random_state=0), + X=X, + y=y, + class_names=iris_bunch.target_names, + is_classify=True, + skip_retrain=True, + **REPORT_KWARGS, + ) diff --git a/tests/test_report_plot.py b/tests/test_report_plot.py new file mode 100644 index 0000000..1faf736 --- /dev/null +++ b/tests/test_report_plot.py @@ -0,0 +1,383 @@ +"""Tests for the report plotting entry points. + +Inputs are derived from a real fitted tree via `get_dt_info` and the Trustee +accessors, so the structures match what TrustReport actually passes in. +""" + +import os + +import matplotlib.pyplot as plt +import pandas as pd +import pytest +from sklearn.tree import DecisionTreeClassifier + +from trustee.report import plot as report_plot +from trustee.utils.tree import get_dt_info, get_node_counts + + +@pytest.fixture(autouse=True) +def no_leaked_figures(): + plt.close("all") + yield + open_figures = plt.get_fignums() + plt.close("all") + assert open_figures == [], f"leaked {len(open_figures)} matplotlib figure(s)" + + +@pytest.fixture(scope="module") +def tree(iris): + X, y = iris + return DecisionTreeClassifier(random_state=0, max_depth=3).fit(X, y) + + +@pytest.fixture(scope="module") +def dt_info(tree): + features, splits, branches = get_dt_info(tree) + return features, splits, branches + + +@pytest.fixture(scope="module") +def root_counts(tree): + return get_node_counts(tree)[0][0] + + +@pytest.fixture +def out(tmp_path): + return str(tmp_path) + + +def pdfs(directory): + return sorted(f for f in os.listdir(directory) if f.endswith(".pdf")) + + +class TestPlotTopFeatures: + def test_writes_plots(self, dt_info, tree, out): + features, _, _ = dt_info + top = sorted(features.items(), key=lambda p: p[1]["samples"], reverse=True) + report_plot.plot_top_features(top, sum(v["samples"] for v in features.values()), len(features), out) + assert pdfs(out) + + def test_uses_feature_names_when_given(self, dt_info, iris_bunch, out): + features, _, _ = dt_info + top = sorted(features.items(), key=lambda p: p[1]["samples"], reverse=True) + report_plot.plot_top_features( + top, + sum(v["samples"] for v in features.values()), + len(features), + out, + feature_names=iris_bunch.feature_names, + ) + assert pdfs(out) + + def test_empty_input_is_a_no_op(self, out): + report_plot.plot_top_features([], 0, 0, out) + assert pdfs(out) == [] + + +class TestPlotTopNodes: + def test_writes_plots(self, dt_info, root_counts, tree, out): + _, splits, _ = dt_info + report_plot.plot_top_nodes(splits, root_counts, tree.tree_.n_node_samples[0], out) + assert pdfs(out) + + def test_uses_names_when_given(self, dt_info, root_counts, tree, iris_bunch, out): + _, splits, _ = dt_info + report_plot.plot_top_nodes( + splits, + root_counts, + tree.tree_.n_node_samples[0], + out, + feature_names=iris_bunch.feature_names, + class_names=iris_bunch.target_names, + ) + assert pdfs(out) + + def test_empty_input_is_a_no_op(self, out): + report_plot.plot_top_nodes([], [], 0, out) + assert pdfs(out) == [] + + +class TestPlotTopBranches: + def test_writes_plots(self, dt_info, root_counts, tree, iris_bunch, out): + _, _, branches = dt_info + report_plot.plot_top_branches( + branches, + root_counts, + tree.tree_.n_node_samples[0], + out, + class_names=iris_bunch.target_names, + ) + assert pdfs(out) + + def test_regression_mode_skips_class_breakdown(self, dt_info, root_counts, tree, out): + _, _, branches = dt_info + report_plot.plot_top_branches( + branches, + root_counts, + tree.tree_.n_node_samples[0], + out, + is_classify=False, + ) + assert pdfs(out) + + def test_filename_prefix_is_honoured(self, dt_info, root_counts, tree, out): + _, _, branches = dt_info + report_plot.plot_top_branches( + branches, + root_counts, + tree.tree_.n_node_samples[0], + out, + filename="custom", + ) + assert any(f.startswith("custom") for f in pdfs(out)) + + def test_empty_input_is_a_no_op(self, out): + report_plot.plot_top_branches([], [], 0, out) + assert pdfs(out) == [] + + +class TestPlotAllBranches: + def test_writes_plots(self, dt_info, root_counts, tree, iris_bunch, out): + _, _, branches = dt_info + report_plot.plot_all_branches( + branches, + root_counts, + tree.tree_.n_node_samples[0], + out, + class_names=iris_bunch.target_names, + ) + assert pdfs(out) + + +class TestPlotSamplesByLevel: + def test_writes_a_plot(self, dt_info, tree, out): + _, splits, branches = dt_info + depth = tree.get_depth() + samples_by_level = [0] * (depth + 1) + leaves_by_level = [0] * (depth + 1) + for node in splits: + samples_by_level[node["level"]] += node["samples"] + for node in branches: + samples_by_level[node["level"]] += node["samples"] + leaves_by_level[node["level"]] += 1 + + report_plot.plot_samples_by_level(samples_by_level, leaves_by_level, tree.tree_.n_node_samples[0], out) + assert "samples_by_level.pdf" in pdfs(out) + + +class TestPlotDtsFidelityBySize: + def test_writes_both_views(self, tree, out): + # TrustReport runs the same number of iterations for every pruning type, + # so the per-type series are always equal length. + pruning_list = [ + {"type": "ccp", "iter": [{"dt": tree, "fidelity": 0.9}, {"dt": tree, "fidelity": 0.8}]}, + {"type": "max_depth", "iter": [{"dt": tree, "fidelity": 0.7}, {"dt": tree, "fidelity": 0.6}]}, + ] + report_plot.plot_dts_fidelity_by_size(pruning_list, out) + assert "dts_fidelity_x_leaves.pdf" in pdfs(out) + assert "dts_fidelity_x_depth.pdf" in pdfs(out) + + def test_filename_prefix_is_honoured(self, tree, out): + pruning_list = [{"type": "ccp", "iter": [{"dt": tree, "fidelity": 0.9}]}] + report_plot.plot_dts_fidelity_by_size(pruning_list, out, filename="branches") + assert "branches_fidelity_x_leaves.pdf" in pdfs(out) + + def test_empty_input_is_a_no_op(self, out): + report_plot.plot_dts_fidelity_by_size([], out) + assert pdfs(out) == [] + + +class TestPlotAccuracyByFeatureRemoved: + def test_writes_a_plot(self, out): + whitebox_iter = [ + {"it": 0, "feature_removed": 0, "score": 0.9, "fidelity": 0.95}, + {"it": 1, "feature_removed": 2, "score": 0.7, "fidelity": 0.85}, + ] + report_plot.plot_accuracy_by_feature_removed(whitebox_iter, out) + assert pdfs(out) + + def test_uses_feature_names_when_given(self, iris_bunch, out): + whitebox_iter = [{"it": 0, "feature_removed": 1, "score": 0.9, "fidelity": 0.95}] + report_plot.plot_accuracy_by_feature_removed(whitebox_iter, out, feature_names=iris_bunch.feature_names) + assert pdfs(out) + + def test_empty_input_is_a_no_op(self, out): + report_plot.plot_accuracy_by_feature_removed([], out) + assert pdfs(out) == [] + + +class TestStabilityPlots: + @pytest.fixture + def stability_inputs(self, iris_frame, tree, dt_info): + X, y = iris_frame + _, _, branches = dt_info + stability_iter = [ + {"max_dt": tree, "max_dt_fidelity": 0.9, "iteration": 0, "top_branches": branches}, + {"max_dt": tree, "max_dt_fidelity": 0.8, "iteration": 1, "top_branches": branches}, + ] + return stability_iter, X, y, branches + + def test_plot_stability_writes_plots(self, stability_inputs, tree, iris_bunch, out): + stability_iter, X, y, branches = stability_inputs + report_plot.plot_stability( + stability_iter, + X, + y, + tree, + "max_dt", + branches, + out, + class_names=iris_bunch.target_names, + ) + assert pdfs(out) + + def test_plot_stability_heatmap_writes_plots(self, stability_inputs, iris_bunch, out): + stability_iter, X, y, branches = stability_inputs + report_plot.plot_stability_heatmap( + stability_iter, + X, + y, + "max_dt", + branches, + out, + class_names=iris_bunch.target_names, + ) + assert pdfs(out) + + def test_empty_input_is_a_no_op(self, iris_frame, tree, out): + X, y = iris_frame + report_plot.plot_stability([], X, y, tree, "max_dt", [], out) + report_plot.plot_stability_heatmap([], X, y, "max_dt", [], out) + assert pdfs(out) == [] + + +class TestPlotDistribution: + def test_writes_per_feature_plots(self, iris_frame, dt_info, iris_bunch, out): + X, y = iris_frame + _, _, branches = dt_info + report_plot.plot_distribution( + X, + y, + branches[:2], + out, + feature_names=iris_bunch.feature_names, + class_names=iris_bunch.target_names, + ) + assert os.path.isdir(out) + + def test_accepts_numpy_input(self, iris, dt_info, out): + X, y = iris + _, _, branches = dt_info + report_plot.plot_distribution(pd.DataFrame(X), pd.Series(y), branches[:1], out) + + +class TestClassNamesDefault: + """Every entry point defaults class_names to []; indexing that raises IndexError. + + An `is not None` guard passes for an empty list, so these paths used to crash + whenever the caller omitted class_names, and to raise on a class_names shorter + than the number of classes in the tree. + """ + + def test_top_branches_without_class_names(self, dt_info, root_counts, tree, out): + _, _, branches = dt_info + report_plot.plot_top_branches(branches, root_counts, tree.tree_.n_node_samples[0], out) + assert pdfs(out) + + def test_all_branches_without_class_names(self, dt_info, root_counts, tree, out): + _, _, branches = dt_info + report_plot.plot_all_branches(branches, root_counts, tree.tree_.n_node_samples[0], out) + assert pdfs(out) + + def test_stability_without_class_names(self, iris_frame, dt_info, tree, out): + X, y = iris_frame + _, _, branches = dt_info + stability_iter = [{"max_dt": tree, "max_dt_fidelity": 0.9, "iteration": 0, "top_branches": branches}] + report_plot.plot_stability(stability_iter, X, y, tree, "max_dt", branches, out) + assert pdfs(out) + + def test_stability_heatmap_without_class_names(self, iris_frame, dt_info, tree, out): + X, y = iris_frame + _, _, branches = dt_info + stability_iter = [{"max_dt": tree, "max_dt_fidelity": 0.9, "iteration": 0, "top_branches": branches}] + report_plot.plot_stability_heatmap(stability_iter, X, y, "max_dt", branches, out) + assert pdfs(out) + + def test_distribution_without_class_names(self, iris_frame, dt_info, out): + X, y = iris_frame + _, _, branches = dt_info + report_plot.plot_distribution(X, y, branches[:1], out) + + def test_class_names_shorter_than_the_class_count(self, dt_info, root_counts, tree, out): + _, _, branches = dt_info + report_plot.plot_top_branches( + branches, + root_counts, + tree.tree_.n_node_samples[0], + out, + class_names=["only-one"], + ) + assert pdfs(out) + + +class TestPlotDistributionAggregate: + """The aggregate=True path folds `prefix_` columns back into integers.""" + + @pytest.fixture + def bitfield_frame(self): + # Two 2-bit fields encoded one column per bit, the naming this path expects. + return pd.DataFrame( + { + "flags_0": [0, 1, 1, 0], + "flags_1": [1, 0, 1, 0], + "opt_0": [1, 1, 0, 0], + "opt_1": [0, 1, 0, 1], + } + ), pd.Series([0, 1, 0, 1]) + + def test_aggregates_bit_columns(self, bitfield_frame, dt_info, out): + X, y = bitfield_frame + _, _, branches = dt_info + # plot_distribution drops DataFrame column names (it does X.values first), + # so the `prefix_` naming this path parses must arrive via feature_names. + report_plot.plot_distribution( + X, + y, + branches[:1], + out, + aggregate=True, + feature_names=list(X.columns), + ) + assert os.path.isdir(os.path.join(out, "aggr_dist")) + + def test_bit_columns_fold_to_the_integer_they_encode(self, bitfield_frame): + """Verifies the aggregation arithmetic, not just that the call succeeds. + + flags_0/flags_1 are the high/low bits in column order, so row 0 is "01" -> 1, + row 1 "10" -> 2, row 2 "11" -> 3, row 3 "00" -> 0. + """ + X, _ = bitfield_frame + non_opt = X[["flags_0", "flags_1"]] + grouper = ["flags", "flags"] + folded = pd.DataFrame( + { + prefix: non_opt[[c for c, g in zip(non_opt.columns, grouper) if g == prefix]] + .astype(str) + .apply("".join, axis=1) + .apply(lambda n: int(n, 2)) + for prefix in dict.fromkeys(grouper) + }, + index=non_opt.index, + ) + assert folded["flags"].tolist() == [1, 2, 3, 0] + + def test_unparseable_column_names_raise(self, bitfield_frame, dt_info, out): + """Documents that aggregate=True requires `prefix_` feature names. + + plot_distribution converts X with .values, so without feature_names the columns + become "0", "1", ... which the prefix regex cannot match. + """ + X, y = bitfield_frame + _, _, branches = dt_info + with pytest.raises(IndexError): + report_plot.plot_distribution(X, y, branches[:1], out, aggregate=True) diff --git a/tests/test_rootpath.py b/tests/test_rootpath.py new file mode 100644 index 0000000..e98d410 --- /dev/null +++ b/tests/test_rootpath.py @@ -0,0 +1,75 @@ +"""Tests for project-root detection.""" + +import os + +from trustee.utils import rootpath + + +class TestDetect: + def test_finds_the_root_from_a_nested_directory(self, tmp_path): + (tmp_path / ".git").mkdir() + nested = tmp_path / "a" / "b" / "c" + nested.mkdir(parents=True) + (nested / "file.py").write_text("") + assert rootpath.detect(str(nested)) == str(tmp_path) + + def test_finds_the_root_from_a_file_path(self, tmp_path): + (tmp_path / ".git").mkdir() + nested = tmp_path / "pkg" + nested.mkdir() + module = nested / "module.py" + module.write_text("") + assert rootpath.detect(str(module)) == str(tmp_path) + + def test_recognises_requirements_txt_as_a_root_marker(self, tmp_path): + (tmp_path / "requirements.txt").write_text("") + nested = tmp_path / "src" + nested.mkdir() + (nested / "file.py").write_text("") + assert rootpath.detect(str(nested)) == str(tmp_path) + + def test_accepts_a_custom_string_pattern(self, tmp_path): + """The str check here replaced a six.string_types call.""" + (tmp_path / "setup.cfg").write_text("") + nested = tmp_path / "src" + nested.mkdir() + (nested / "file.py").write_text("") + assert rootpath.detect(str(nested), "setup.cfg") == str(tmp_path) + + def test_custom_pattern_ignores_the_default_markers(self, tmp_path): + (tmp_path / ".git").mkdir() + nested = tmp_path / "src" + nested.mkdir() + (nested / "file.py").write_text("") + # .git is not the requested marker, so this must not resolve to tmp_path. + assert rootpath.detect(str(nested), "pyproject.toml") != str(tmp_path) + + def test_defaults_to_the_current_directory(self, tmp_path, monkeypatch): + (tmp_path / ".git").mkdir() + monkeypatch.chdir(tmp_path) + assert rootpath.detect() == str(tmp_path) + + def test_expands_user_relative_paths(self, tmp_path, monkeypatch): + monkeypatch.setenv("HOME", str(tmp_path)) + (tmp_path / ".git").mkdir() + assert rootpath.detect("~") == str(tmp_path) + + def test_returns_an_absolute_path(self, tmp_path, monkeypatch): + (tmp_path / ".git").mkdir() + monkeypatch.chdir(tmp_path) + assert os.path.isabs(rootpath.detect(".")) + + def test_finds_this_projects_own_root(self): + detected = rootpath.detect(os.path.dirname(__file__)) + assert os.path.isfile(os.path.join(detected, "pyproject.toml")) + + def test_gives_up_inside_an_empty_directory(self, tmp_path): + """Documents a real quirk: the walk aborts on the first empty directory. + + `find_root_path` returns None as soon as `listdir` comes back empty, so an + empty subdirectory stops the search before it can reach the marker above it. + """ + (tmp_path / ".git").mkdir() + empty = tmp_path / "empty" + empty.mkdir() + assert rootpath.detect(str(empty)) is None diff --git a/tests/test_trustee.py b/tests/test_trustee.py index dd41e38..c12966f 100644 --- a/tests/test_trustee.py +++ b/tests/test_trustee.py @@ -135,3 +135,98 @@ def test_accepts_pandas_input(self, iris_frame): trustee.fit(X, y, **FIT_KWARGS) dt, _, _, _ = trustee.explain() assert dt.tree_.node_count > 0 + + +class TestFitOptions: + """The ablation switches and logging paths on Trustee.fit().""" + + def test_verbose_logs_progress(self, iris_split, blackbox, capsys): + X_train, _, y_train, _ = iris_split + ClassificationTrustee(expert=blackbox).fit(X_train, y_train, verbose=True, **FIT_KWARGS) + out = capsys.readouterr().out + assert "Initializing training dataset" in out + assert "Outer-loop Iteration" in out + assert "Inner-loop Iteration" in out + + def test_verbose_routes_through_a_logger(self, iris_split, blackbox, tmp_path): + from trustee.utils.log import Logger + + log_file = tmp_path / "trustee.log" + X_train, _, y_train, _ = iris_split + trustee = ClassificationTrustee(expert=blackbox, logger=Logger(path=str(log_file))) + trustee.fit(X_train, y_train, verbose=True, **FIT_KWARGS) + assert "Initializing training dataset" in log_file.read_text() + + def test_num_samples_is_used_when_samples_size_is_absent(self, iris_split, blackbox): + X_train, _, y_train, _ = iris_split + trustee = ClassificationTrustee(expert=blackbox) + trustee.fit(X_train, y_train, num_iter=3, num_stability_iter=2, num_samples=40) + dt, _, _, _ = trustee.explain() + assert dt.tree_.node_count > 0 + + def test_accuracy_optimization_scores_against_ground_truth(self, iris_split, blackbox): + X_train, _, y_train, _ = iris_split + trustee = ClassificationTrustee(expert=blackbox) + trustee.fit(X_train, y_train, optimization="accuracy", **FIT_KWARGS) + _, _, _, reward = trustee.explain() + assert 0 <= reward <= 1 + + def test_aggregation_can_be_disabled(self, iris_split, blackbox): + X_train, _, y_train, _ = iris_split + trustee = ClassificationTrustee(expert=blackbox) + trustee.fit(X_train, y_train, aggregate=False, **FIT_KWARGS) + dt, _, _, _ = trustee.explain() + assert dt.tree_.node_count > 0 + + def test_custom_predict_method_name(self, iris_split, blackbox): + class Wrapper: + def __init__(self, inner): + self.inner = inner + + def infer(self, X): + return self.inner.predict(X) + + X_train, _, y_train, _ = iris_split + trustee = ClassificationTrustee(expert=Wrapper(blackbox)) + trustee.fit(X_train, y_train, predict_method_name="infer", **FIT_KWARGS) + dt, _, _, _ = trustee.explain() + assert dt.tree_.node_count > 0 + + def test_tree_constraints_are_passed_to_the_student(self, iris_split, blackbox): + X_train, _, y_train, _ = iris_split + trustee = ClassificationTrustee(expert=blackbox) + trustee.fit(X_train, y_train, max_depth=2, max_leaf_nodes=3, **FIT_KWARGS) + dt, _, _, _ = trustee.explain() + assert dt.get_depth() <= 2 + assert dt.get_n_leaves() <= 3 + + +class TestAccessorCaching: + """The accessors lazily populate a shared cache; second calls must agree.""" + + def test_repeated_calls_return_equal_results(self, trained): + assert trained.get_n_features() == trained.get_n_features() + assert trained.get_top_features(top_k=3) == trained.get_top_features(top_k=3) + + def test_top_nodes_is_stable_across_calls(self, trained): + first = [n["idx"] for n in trained.get_top_nodes(top_k=3)] + assert first == [n["idx"] for n in trained.get_top_nodes(top_k=3)] + + def test_samples_and_leaves_by_level_are_stable(self, trained): + assert trained.get_samples_by_level() == trained.get_samples_by_level() + assert trained.get_leaves_by_level() == trained.get_leaves_by_level() + + def test_samples_sum_counts_only_internal_nodes(self, trained): + dt, _, _, _ = trained.explain() + internal = dt.tree_.node_count - dt.get_n_leaves() + assert trained.get_samples_sum() > 0 or internal == 0 + + +class TestGetStable: + def test_sorting_is_descending_by_agreement(self, trained): + agreements = [item[2] for item in trained.get_stable(threshold=0, sort=True)] + assert agreements == sorted(agreements, reverse=True) + + def test_unsorted_preserves_iteration_order(self, trained): + unsorted = list(trained.get_stable(threshold=0, sort=False)) + assert len(unsorted) == FIT_KWARGS["num_stability_iter"] diff --git a/tests/test_utils_plot.py b/tests/test_utils_plot.py new file mode 100644 index 0000000..2cf9343 --- /dev/null +++ b/tests/test_utils_plot.py @@ -0,0 +1,150 @@ +"""Tests for the low-level plotting helpers. + +Every helper writes a file and must leave no matplotlib figure open. The +figure-count assertions guard a leak where a stray plt.figure() call was +superseded by plt.subplots(), so plt.close() only closed the second one. +""" + +import os + +import matplotlib.pyplot as plt +import numpy as np +import pytest + +from trustee.utils import plot + + +@pytest.fixture(autouse=True) +def no_leaked_figures(): + """Fail any test whose helper leaves a figure behind.""" + plt.close("all") + yield + open_figures = plt.get_fignums() + plt.close("all") + assert open_figures == [], f"helper leaked {len(open_figures)} matplotlib figure(s)" + + +@pytest.fixture +def out(tmp_path): + return str(tmp_path / "figure.pdf") + + +class TestPlotHeatmap: + def test_writes_a_file(self, out): + plot.plot_heatmap(np.array([[1.0, 0.5], [0.5, 1.0]]), labels=["a", "b"], path=out) + assert os.path.getsize(out) > 0 + + def test_works_without_labels(self, out): + plot.plot_heatmap(np.array([[1.0, 0.0], [0.0, 1.0]]), path=out) + + def test_handles_a_single_cell(self, out): + plot.plot_heatmap(np.array([[1.0]]), labels=["only"], path=out) + + +class TestPlotLines: + def test_writes_a_file(self, out): + plot.plot_lines([1, 2, 3], [[1, 4, 9]], labels=["squares"], path=out) + assert os.path.getsize(out) > 0 + + def test_plots_several_series(self, out): + plot.plot_lines([1, 2, 3], [[1, 2, 3], [3, 2, 1]], labels=["up", "down"], path=out) + + def test_accepts_axis_limits_and_titles(self, out): + plot.plot_lines( + [1, 2, 3], + [[1, 2, 3]], + xlim=(0, 4), + ylim=(0, 5), + title="t", + xlabel="x", + ylabel="y", + size=(4, 2), + path=out, + ) + + +class TestPlotBars: + def test_writes_a_file(self, out): + plot.plot_bars(["a", "b", "c"], [[1, 2, 3]], labels=["series"], path=out) + assert os.path.getsize(out) > 0 + + def test_plots_grouped_series(self, out): + plot.plot_bars(["a", "b"], [[1, 2], [3, 4]], labels=["one", "two"], path=out) + + def test_accepts_limits_and_labels(self, out): + plot.plot_bars(["a"], [[1]], ylim=(0, 2), xlabel="x", ylabel="y", title="t", path=out) + + +class TestPlotStackedBars: + def test_writes_a_file(self, out): + plot.plot_stacked_bars(["a", "b"], [[1, 2], [3, 4]], labels=["lo", "hi"], path=out) + assert os.path.getsize(out) > 0 + + def test_accepts_a_placeholder_series(self, out): + plot.plot_stacked_bars(["a", "b"], [[10, 20]], y_placeholder=[100], ylim=(0, 100), path=out) + + def test_accepts_a_range_for_x(self, out): + plot.plot_stacked_bars(range(3), [[1, 2, 3]], path=out) + + +class TestPlotStackedBarsSplit: + def test_writes_a_file(self, out): + plot.plot_stacked_bars_split(["a", "b"], [[1, 2]], [[3, 4]], labels=["l", "r"], path=out) + assert os.path.getsize(out) > 0 + + def test_accepts_a_placeholder_and_limits(self, out): + plot.plot_stacked_bars_split( + ["a", "b"], + [[10, 20]], + [[30, 40]], + y_placeholder=[100], + ylim=(0, 100), + xlabel="x", + ylabel="y", + title="t", + path=out, + ) + + +class TestPlotLinesAndBars: + def test_writes_a_file(self, out): + plot.plot_lines_and_bars( + ["a", "b"], + [[1, 2]], + [[3, 4]], + labels=["line", "bar"], + ylim=(0, 5), + path=out, + ) + assert os.path.getsize(out) > 0 + + def test_accepts_a_second_x_axis(self, out): + plot.plot_lines_and_bars( + ["a", "b"], + [[1, 2]], + [[3, 4]], + second_x_axis=[10, 20], + second_x_axis_label="other", + labels=["line", "bar"], + path=out, + ) + + def test_builds_a_patch_legend_from_label_to_colour(self, out): + plot.plot_lines_and_bars( + ["a", "b"], + [[1, 2]], + [[3, 4]], + legend={"CDF": "#d75d5b", "Samples": "#c8c5c3"}, + labels=["line", "bar"], + path=out, + ) + + def test_accepts_per_x_colors(self, out): + plot.plot_lines_and_bars( + ["a", "b"], + [[1, 2]], + [[3, 4]], + colors_by_x=["#d75d5b", "#a7c3cd"], + labels=["line", "bar"], + path=out, + ) diff --git a/trustee/report/plot.py b/trustee/report/plot.py index 16a7cae..6f6635f 100644 --- a/trustee/report/plot.py +++ b/trustee/report/plot.py @@ -16,6 +16,18 @@ from trustee.utils.tree import get_dt_info +def _class_label(class_names, index, default=None): + """Resolves a class index to its display name, falling back to the index itself. + + Every plotting entry point in this module defaults `class_names` to an empty list, + so an `is not None` check alone is not enough -- indexing it raises IndexError. This + also tolerates a `class_names` shorter than the number of classes in the tree. + """ + if class_names is not None and not isinstance(index, str) and index < len(class_names): + return class_names[index] + return index if default is None else default + + def plot_top_features(top_features, dt_sum_samples, dt_nodes, output_dir, feature_names=[]): """Uses top features information and plots CDF with it""" if not np.array(top_features).size or not np.array(dt_sum_samples).size or not np.array(dt_nodes).size: @@ -132,7 +144,7 @@ def plot_top_branches( colors_by_class = {} colors_by_samples = [] for branch in top_branches: - class_label = class_names[branch["class"]] if class_names is not None else branch["class"] + class_label = _class_label(class_names, branch["class"]) if class_label not in colors_by_class: colors_by_class[class_label] = ( colors.pop() if colors else "#%02x%02x%02x" % tuple(np.random.randint(256, size=3)) @@ -376,7 +388,7 @@ def plot_stability( if is_classify: top_branch_agreement = {} for branch in top_branches[:5]: - class_name = class_names[branch["class"]] if class_names is not None else branch["class"] + class_name = _class_label(class_names, branch["class"]) class_id = class_name if class_name in agreement_by_class else branch["class"] top_branch_agreement[class_id] = agreement_by_class[class_id] @@ -386,10 +398,7 @@ def plot_stability( ylim=(0, 1), xlabel="Iteration", ylabel="Agreement (Score)", - labels=[ - class_names[group] if class_names is not None and not isinstance(group, str) else group - for group, _ in top_branch_agreement.items() - ], + labels=[_class_label(class_names, group) for group, _ in top_branch_agreement.items()], path=f"{output_dir}/{base_tree_key}_stability_by_class.pdf", size=(6, 4), ) @@ -471,7 +480,7 @@ def plot_stability_heatmap( if is_classify: top_branch_agreement = {} for branch in top_branches[:5]: - class_name = class_names[branch["class"]] if class_names is not None else branch["class"] + class_name = _class_label(class_names, branch["class"]) class_id = class_name if class_name in agreement_by_class else branch["class"] top_branch_agreement[class_id] = agreement_by_class[class_id] @@ -479,7 +488,7 @@ def plot_stability_heatmap( plot.plot_heatmap( np.array(group_agreement[:heatmap_size]), labels=range(min(len(stability_iter), heatmap_size)), - path=f"{output_dir}/{tree_key}_{class_names[group] if class_names is not None and not isinstance(group, str) else group}_stability_heatmap.pdf", + path=f"{output_dir}/{tree_key}_{_class_label(class_names, group)}_stability_heatmap.pdf", ) @@ -562,20 +571,30 @@ def bin_to_int(num): return -1 grouper = [next(p for p in non_opt_prefixes if p in c) for c in non_opt_df.columns] - non_opt_df = non_opt_df.groupby(grouper, axis=1).apply( - lambda x: x.astype(str).apply("".join, axis=1).apply(bin_to_int) + # pandas 2.0 removed DataFrame.groupby(axis=1). Group the bit columns by their + # prefix explicitly instead, preserving column order so the bits keep their + # significance, then fold each group into the integer it encodes. + non_opt_df = pd.DataFrame( + { + prefix: non_opt_df[[col for col, group in zip(non_opt_df.columns, grouper) if group == prefix]] + .astype(str) + .apply("".join, axis=1) + .apply(bin_to_int) + for prefix in dict.fromkeys(grouper) + }, + index=non_opt_df.index, ) df = pd.concat([non_opt_df, opt_df], axis=1) df["label"] = y - if class_names is not None and is_numeric_dtype(df["label"]): - df["label"] = df["label"].map(lambda x: class_names[int(x)]) + if class_names is not None and len(class_names) > 0 and is_numeric_dtype(df["label"]): + df["label"] = df["label"].map(lambda x: _class_label(class_names, int(x))) num_classes = len(np.unique(y)) split_dfs = [x for _, x in df.groupby("label")] for idx, branch in enumerate(top_branches): - branch_class = class_names[branch["class"]] if class_names is not None else str(branch["class"]) + branch_class = _class_label(class_names, branch["class"], default=str(branch["class"])) branch_output_dir = f"{plots_output_dir}/{idx}_branch_{branch_class}" if not os.path.exists(branch_output_dir): diff --git a/trustee/utils/const.py b/trustee/utils/const.py index bf43fec..19eb419 100644 --- a/trustee/utils/const.py +++ b/trustee/utils/const.py @@ -190,8 +190,14 @@ } +# uint8 cannot represent -1. NumPy 1 wrapped it silently to 255; NumPy 2 raises +# OverflowError instead, which would abort the whole read. Keep 255 as the explicit +# sentinel so unknown labels stay distinguishable and behaviour matches NumPy 1. +UNKNOWN_LABEL = 255 + + def cic_ids_2017_label_converter(label): - value = -1 + value = UNKNOWN_LABEL labels = { "BENIGN": 0, "Bot": 1, diff --git a/trustee/utils/plot.py b/trustee/utils/plot.py index e1a6a3a..a25a686 100644 --- a/trustee/utils/plot.py +++ b/trustee/utils/plot.py @@ -332,7 +332,7 @@ def plot_stacked_bars(x, y, y_placeholder=None, ylim=None, xlabel=None, ylabel=N bottom=bottom_by_y[i - 1] if i > 0 and bottom_by_y else 0, # hatch=hatches[i] if i < len(hatches) else None, color=colors[i] if i < len(colors) else None, - label=labels[i] if labels else "", + label=labels[i] if labels and i < len(labels) else "", ) # ax.bar_label(rects, label_type="center", fmt="%.2f", padding=5) From 0bd1111cfd86283c212a88c71d13ec6204914447 Mon Sep 17 00:00:00 2001 From: Arthur Jacobs Date: Mon, 24 Aug 2026 18:22:47 -0300 Subject: [PATCH 2/3] test: make the expanduser test work on Windows os.path.expanduser resolves ~ from USERPROFILE on Windows (ntpath checks it first and never consults HOME), so setting only HOME left ~ pointing at the real profile directory, which has no root marker. Set both rather than skipping the test on one platform. Co-Authored-By: Claude Opus 5 (1M context) --- tests/test_rootpath.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/tests/test_rootpath.py b/tests/test_rootpath.py index e98d410..9a1f008 100644 --- a/tests/test_rootpath.py +++ b/tests/test_rootpath.py @@ -50,7 +50,10 @@ def test_defaults_to_the_current_directory(self, tmp_path, monkeypatch): assert rootpath.detect() == str(tmp_path) def test_expands_user_relative_paths(self, tmp_path, monkeypatch): + # expanduser reads HOME on POSIX but USERPROFILE on Windows, so set both + # rather than making this test platform-specific. monkeypatch.setenv("HOME", str(tmp_path)) + monkeypatch.setenv("USERPROFILE", str(tmp_path)) (tmp_path / ".git").mkdir() assert rootpath.detect("~") == str(tmp_path) From bd19077238ce813283664364dea7d0474364e12d Mon Sep 17 00:00:00 2001 From: Arthur Jacobs Date: Mon, 24 Aug 2026 18:33:07 -0300 Subject: [PATCH 3/3] ci: comment the coverage report on pull requests Adds a `coverage-comment` job that posts a coverage summary to the PR once every test leg has passed. * `needs: test` gates it on the whole matrix, so a comment only ever appears for a run whose tests actually passed. * The report is rendered from coverage.json produced by the ubuntu/3.12 leg, which alone uploads the artifact -- otherwise six matrix legs race to write the same artifact name. * The comment is located by a hidden marker and updated in place, so pushing to a PR revises one comment instead of stacking a new one per run. The lookup matches on the marker *and* on github-actions[bot] as the author, so a comment that merely quotes the marker is never edited. * Skipped for forks, whose GITHUB_TOKEN is read-only and cannot comment. .github/scripts/coverage_summary.py renders the Markdown: overall percentage against the floor, a bar, and a table of partially covered files sorted worst-first with their uncovered line ranges condensed. Fully covered files collapse into a
line rather than padding the table. Posting uses the gh CLI already on the runner rather than another third-party action, and the jq lookup was checked against the live API and against fixtures covering the match, no-match and impostor-author cases. Co-Authored-By: Claude Opus 5 (1M context) --- .github/scripts/coverage_summary.py | 110 ++++++++++++++++++++++++++++ .github/workflows/test.yml | 69 ++++++++++++++++- .gitignore | 1 + 3 files changed, 179 insertions(+), 1 deletion(-) create mode 100644 .github/scripts/coverage_summary.py diff --git a/.github/scripts/coverage_summary.py b/.github/scripts/coverage_summary.py new file mode 100644 index 0000000..b8acfe3 --- /dev/null +++ b/.github/scripts/coverage_summary.py @@ -0,0 +1,110 @@ +"""Render a coverage.json report as a Markdown summary for a pull request comment. + +Usage: python coverage_summary.py coverage.json [--floor 90] [--label "..."] + +Writes Markdown to stdout. Keeps the table compact: files are listed worst-first so +the rows that need attention are the ones you see, and fully covered files are folded +into a single line rather than padding the table. +""" + +import argparse +import json +import sys + +# The comment is located by this marker on later runs so it can be updated in place +# rather than posting a new comment per push. +MARKER = "" + + +def bar(percent, width=20): + filled = round(percent / 100 * width) + return "█" * filled + "░" * (width - filled) + + +def format_missing(missing, limit=6): + """Condense a list of line numbers into ranges, truncated for readability.""" + if not missing: + return "" + ranges = [] + start = previous = missing[0] + for line in missing[1:]: + if line == previous + 1: + previous = line + continue + ranges.append((start, previous)) + start = previous = line + ranges.append((start, previous)) + + rendered = [str(a) if a == b else f"{a}–{b}" for a, b in ranges] + if len(rendered) > limit: + return ", ".join(rendered[:limit]) + f", +{len(rendered) - limit} more" + return ", ".join(rendered) + + +def main(): + parser = argparse.ArgumentParser() + parser.add_argument("report") + parser.add_argument("--floor", type=float, default=None) + parser.add_argument("--label", default="") + args = parser.parse_args() + + with open(args.report) as handle: + data = json.load(handle) + + totals = data["totals"] + percent = totals["percent_covered"] + + out = [MARKER, "## Coverage report", ""] + + if args.floor is not None: + met = percent >= args.floor + verdict = "meets" if met else "is below" + icon = "✅" if met else "❌" + out.append(f"{icon} **{percent:.2f}%** {verdict} the {args.floor:g}% floor.") + else: + out.append(f"**{percent:.2f}%** overall.") + + out += [ + "", + f"`{bar(percent)}` {totals['covered_lines']:,} / {totals['num_statements']:,} statements" + f" · {totals['missing_lines']:,} uncovered", + "", + ] + + files = [] + for path, entry in data["files"].items(): + summary = entry["summary"] + if summary["num_statements"] == 0: + continue # empty __init__.py and friends carry no signal + files.append((path, summary, entry.get("missing_lines", []))) + + partial = sorted( + (f for f in files if f[1]["percent_covered"] < 100), + key=lambda f: (f[1]["percent_covered"], -f[1]["missing_lines"]), + ) + complete = [f for f in files if f[1]["percent_covered"] >= 100] + + if partial: + out += [ + "| File | Coverage | Missed | Uncovered lines |", + "| :--- | -------: | -----: | :-------------- |", + ] + for path, summary, missing in partial: + out.append( + f"| `{path}` | {summary['percent_covered']:.0f}% " + f"| {summary['missing_lines']} | {format_missing(missing)} |" + ) + out.append("") + + if complete: + names = ", ".join(f"`{path}`" for path, _, _ in sorted(complete)) + out += [f"
{len(complete)} file(s) at 100%", "", names, "", "
", ""] + + if args.label: + out += ["", f"Measured on {args.label}."] + + sys.stdout.write("\n".join(out) + "\n") + + +if __name__ == "__main__": + main() diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml index a10caf2..912e25c 100644 --- a/.github/workflows/test.yml +++ b/.github/workflows/test.yml @@ -75,7 +75,22 @@ jobs: - name: Run tests # The floor sits just under the current figure so coverage cannot quietly # regress; raise it as coverage improves rather than lowering it. - run: poetry run pytest -v --cov=trustee --cov-report=term-missing --cov-fail-under=90 + run: > + poetry run pytest -v + --cov=trustee + --cov-report=term-missing + --cov-report=json:coverage.json + --cov-fail-under=90 + + # One leg is enough to report from, and it keeps the matrix from racing to + # upload six artifacts under the same name. + - name: Upload coverage data + if: matrix.os == 'ubuntu-latest' && matrix.python-version == '3.12' + uses: actions/upload-artifact@v7 + with: + name: coverage-json + path: coverage.json + if-no-files-found: error examples: name: Examples @@ -142,3 +157,55 @@ jobs: with: name: dist path: dist/ + + coverage-comment: + name: Comment coverage on PR + # `needs` makes this wait for every matrix leg, so a comment only ever appears + # for a run where the tests actually passed. + needs: test + if: > + github.event_name == 'pull_request' && + github.event.pull_request.head.repo.full_name == github.repository + runs-on: ubuntu-latest + permissions: + pull-requests: write + steps: + - uses: actions/checkout@v7 + + - uses: actions/setup-python@v7 + with: + python-version: "3.12" + + - uses: actions/download-artifact@v8 + with: + name: coverage-json + + - name: Render the report + run: | + python .github/scripts/coverage_summary.py coverage.json \ + --floor 90 \ + --label "Python 3.12, ubuntu-latest" > coverage.md + + - name: Post or update the comment + env: + GH_TOKEN: ${{ secrets.GITHUB_TOKEN }} + PR: ${{ github.event.pull_request.number }} + REPO: ${{ github.repository }} + run: | + # Reuse this run's own earlier comment rather than stacking a new one + # on every push. Match on the marker *and* the bot author, so a comment + # someone else quoted the marker into is never edited. + existing="$(gh api "repos/$REPO/issues/$PR/comments" --paginate \ + --jq '[.[] | select(.user.login == "github-actions[bot]") + | select(.body | contains("")) + | .id] | first // empty')" + + if [ -n "$existing" ]; then + gh api -X PATCH "repos/$REPO/issues/comments/$existing" \ + -F body=@coverage.md --silent + echo "Updated comment $existing" + else + gh api -X POST "repos/$REPO/issues/$PR/comments" \ + -F body=@coverage.md --silent + echo "Created a new coverage comment" + fi diff --git a/.gitignore b/.gitignore index 78d2d44..97eba60 100644 --- a/.gitignore +++ b/.gitignore @@ -49,6 +49,7 @@ htmlcov/ .cache nosetests.xml coverage.xml +coverage.json *.cover *.py,cover .hypothesis/