Spaces:
Running
Running
Download tests/test_tools.py from Ekimetrics/MCP_indicators: direct link, hf CLI and curl.
- Browser
- Download file 8.83 kB
-
https://huggingface.co/spaces/Ekimetrics/MCP_indicators/resolve/main/tests/test_tools.py
- Command line
-
hf download hf://spaces/Ekimetrics/MCP_indicators/tests/test_tools.py
-
curl -L -o test_tools.py https://huggingface.co/spaces/Ekimetrics/MCP_indicators/resolve/main/tests/test_tools.py
8.83 kB
| """Offline tests of the MCP tools against the fake Cube.js backend. | |
| Run: python -m unittest discover -s tests -v | |
| """ | |
| from __future__ import annotations | |
| import json | |
| import sys | |
| import unittest | |
| from pathlib import Path | |
| sys.path.insert(0, str(Path(__file__).resolve().parents[1])) | |
| import src.api_client as api_client # noqa: E402 | |
| import src.cache as cache_mod # noqa: E402 | |
| import src.cube_resolver as resolver_mod # noqa: E402 | |
| from src.tools import ( # noqa: E402 | |
| find_territory, | |
| get_indicator_details, | |
| list_indicators, | |
| query_indicator_data, | |
| search_indicators, | |
| ) | |
| from tests.fake_api import FakeCubeJsClient # noqa: E402 | |
| class ToolsTestCase(unittest.IsolatedAsyncioTestCase): | |
| def setUp(self) -> None: | |
| # Fresh singletons bound to the fake backend | |
| api_client._client_instance = FakeCubeJsClient() | |
| resolver_mod._resolver_instance = None | |
| cache_mod._cache_instance = None | |
| def fake(self) -> FakeCubeJsClient: | |
| return api_client._client_instance # type: ignore[return-value] | |
| # ------------------------------------------------------------------ format | |
| async def test_outputs_are_compact_json(self) -> None: | |
| for raw in ( | |
| await list_indicators(), | |
| await search_indicators("espace"), | |
| await get_indicator_details("611"), | |
| await query_indicator_data("611", "region", "93"), | |
| ): | |
| self.assertNotIn("\n", raw) | |
| self.assertIsInstance(json.loads(raw), dict) | |
| # ------------------------------------------------------------------ list | |
| async def test_list_all(self) -> None: | |
| payload = json.loads(await list_indicators()) | |
| self.assertEqual(payload["count"], 168) | |
| self.assertEqual(len(payload["indicators"]), 168) | |
| self.assertNotIn("truncated", payload) | |
| item = next(i for i in payload["indicators"] if i["id"] == 611) | |
| self.assertEqual(item["mailles"], ["commune", "epci", "departement", "region"]) | |
| self.assertIn("Mieux se déplacer", payload["themes"]) | |
| async def test_list_filters_accent_insensitive(self) -> None: | |
| a = json.loads(await list_indicators(thematique="deplacer")) | |
| b = json.loads(await list_indicators(thematique="Déplacer")) | |
| self.assertEqual(a["count"], b["count"]) | |
| self.assertEqual(a["count"], 58) | |
| c = json.loads(await list_indicators(thematique="déplacer", maille="commune")) | |
| self.assertLess(c["count"], a["count"]) | |
| self.assertTrue(all("commune" in i["mailles"] for i in c["indicators"])) | |
| async def test_list_invalid_maille(self) -> None: | |
| payload = json.loads(await list_indicators(maille="pays")) | |
| self.assertIn("error", payload) | |
| self.assertIn("valid_levels", payload) | |
| # ------------------------------------------------------------------ search | |
| async def test_search_accents_and_plural(self) -> None: | |
| with_acc = json.loads(await search_indicators("émissions gaz effet de serre")) | |
| without = json.loads(await search_indicators("emission gaz effet serre")) | |
| self.assertTrue(with_acc["indicators"]) | |
| self.assertEqual( | |
| [i["id"] for i in with_acc["indicators"]], [i["id"] for i in without["indicators"]] | |
| ) | |
| async def test_search_libelle_first(self) -> None: | |
| payload = json.loads(await search_indicators("consommation espace")) | |
| ids = [i["id"] for i in payload["indicators"]] | |
| self.assertIn(611, ids[:2]) | |
| self.assertIn(545, ids[:2]) | |
| async def test_search_capped(self) -> None: | |
| payload = json.loads(await search_indicators("de")) | |
| self.assertLessEqual(len(payload["indicators"]), 30) | |
| self.assertTrue(payload["truncated"]) | |
| self.assertGreater(payload["total_count"], 30) | |
| async def test_search_no_match(self) -> None: | |
| payload = json.loads(await search_indicators("blockchain")) | |
| self.assertEqual(payload["total_count"], 0) | |
| self.assertIn("hint", payload) | |
| # ------------------------------------------------------------------ details | |
| async def test_details_sources_use_existing_dimensions(self) -> None: | |
| payload = json.loads(await get_indicator_details("611")) | |
| self.assertNotIn("sources_warning", payload) | |
| self.assertEqual(payload["sources"][0]["producteur_source"], "Cerema") | |
| self.assertNotIn("id_indicateur", payload["sources"][0]) | |
| self.assertEqual(payload["metadata"]["annees_disponibles"][0], "2009") | |
| self.assertEqual(payload["available_cubes"]["commune"], "conso_enaf_com") | |
| self.assertAlmostEqual(payload["metadata"]["completion"]["region"], 0.8947) | |
| # the source query never asked for the non-existent `libelle` dimension | |
| src_queries = [q for q in self.fake.queries if q["dimensions"][0].startswith("indicateur_x_source")] | |
| self.assertTrue(src_queries) | |
| self.assertNotIn("indicateur_x_source_metadata.libelle", src_queries[0]["dimensions"]) | |
| async def test_details_unknown_and_invalid(self) -> None: | |
| self.assertIn("error", json.loads(await get_indicator_details("999999"))) | |
| self.assertIn("error", json.loads(await get_indicator_details("abc"))) | |
| # ------------------------------------------------------------------ query | |
| async def test_query_one_territory(self) -> None: | |
| payload = json.loads(await query_indicator_data("611", "region", "93")) | |
| self.assertEqual(payload["unite"], "ha") | |
| self.assertFalse(payload["truncated"]) | |
| self.assertEqual(payload["total_count"], 3) | |
| row = payload["data"][0] | |
| self.assertEqual(row["geocode"], "93") | |
| self.assertIsInstance(row["valeur"], float) # converted from string | |
| self.assertEqual(payload["annees_disponibles"][-1], "2023") | |
| async def test_query_truncation_and_limit(self) -> None: | |
| payload = json.loads(await query_indicator_data("611", "commune", year="2021")) | |
| self.assertEqual(payload["total_count"], 100) | |
| self.assertTrue(payload["truncated"]) | |
| self.assertIn("hint", payload) | |
| big = json.loads(await query_indicator_data("611", "commune", year="2021", limit=500)) | |
| self.assertEqual(big["total_count"], 500) | |
| self.assertTrue(big["truncated"]) | |
| small = json.loads(await query_indicator_data("611", "region", year="2021", limit=5)) | |
| self.assertEqual(small["total_count"], 5) | |
| self.assertTrue(small["truncated"]) | |
| # limit is clamped, never rejected | |
| clamped = json.loads(await query_indicator_data("611", "region", year="2021", limit=9999)) | |
| self.assertEqual(clamped["total_count"], 13) | |
| self.assertFalse(clamped["truncated"]) | |
| async def test_query_wrong_level_and_hints(self) -> None: | |
| payload = json.loads(await query_indicator_data("42", "commune")) | |
| self.assertIn("error", payload) | |
| self.assertEqual(payload["available_levels"], ["departement", "region"]) | |
| empty = json.loads(await query_indicator_data("611", "region", "93", year="1999")) | |
| self.assertEqual(empty["total_count"], 0) | |
| self.assertIn("1999 is not available", empty["hint"]) | |
| bad_code = json.loads(await query_indicator_data("611", "departement", "Rhône")) | |
| self.assertIn("find_territory", bad_code["hint"]) | |
| # ------------------------------------------------------------------ territory | |
| async def test_find_territory_region_local(self) -> None: | |
| payload = json.loads(await find_territory("provence", "region")) | |
| self.assertEqual(payload["territories"][0]["geocode"], "93") | |
| self.assertFalse(self.fake.queries) # regions resolved without the API | |
| async def test_find_territory_epci_and_commune(self) -> None: | |
| payload = json.loads(await find_territory("Lyon")) | |
| levels = {t["geographic_level"]: t for t in payload["territories"]} | |
| self.assertEqual(levels["epci"]["geocode"], "200046977") | |
| self.assertEqual(levels["commune"]["geocode"], "69123") | |
| self.assertNotIn("departement", levels) | |
| dpt = json.loads(await find_territory("rhône", "departement")) | |
| self.assertEqual({t["geocode"] for t in dpt["territories"]}, {"13", "69"}) | |
| async def test_find_territory_invalid(self) -> None: | |
| self.assertIn("error", json.loads(await find_territory("", "epci"))) | |
| self.assertIn("error", json.loads(await find_territory("Lyon", "pays"))) | |
| # ------------------------------------------------------------------ resilience | |
| async def test_api_down_before_first_load(self) -> None: | |
| self.fake.fail_meta = True | |
| payload = json.loads(await search_indicators("espace")) | |
| self.assertIn("error", payload) | |
| self.assertIn("hint", payload) | |
| if __name__ == "__main__": | |
| unittest.main() | |