MCP_indicators / tests /test_tools.py
Qdonnars's picture
Gradio 6 (transport MCP Streamable HTTP + SSE), sorties JSON compactes, sources réparées, find_territory, page Space
febc39b
Raw History Blame Contribute Delete
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
@property
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()