Paper2Agent_decoupleRpy / tests /test_planner_tool.py
Annie Voigt
fix(planner): dataset_plan_analysis must say how to LOAD the dataset
7944fee
Raw
History Blame Contribute Delete
17.3 kB
"""
Tests for dataset_plan_analysis and its helpers (_detect_intent, _workflow_for_intent).
No heavy dependencies required — tests purely the keyword matching, intent
routing, and return structure. All tests are offline (no decoupler calls,
no file I/O, no network).
"""
from __future__ import annotations
import sys
from pathlib import Path
import pytest
sys.path.insert(0, str(Path(__file__).parent.parent))
from src.tools.dataset_tools import (
_detect_intent,
dataset_plan_analysis,
)
# ---------------------------------------------------------------------------
# _detect_intent
# ---------------------------------------------------------------------------
class TestDetectIntent:
# ── compare_groups ───────────────────────────────────────────────────
@pytest.mark.parametrize(
"question",
[
"Compare Classical vs Basal tumors",
"Which TFs differ between Classical and Basal?",
"Differential TF activity between subtypes",
"Are there differences between treated vs control?",
"What is upregulated in group A versus group B?",
"higher in Classical compared to Basal",
],
)
def test_compare_groups_detected(self, question):
intent, _, _ = _detect_intent(question)
assert intent == "compare_groups", f"Failed for: {question!r}"
# ── score_samples ────────────────────────────────────────────────────
@pytest.mark.parametrize(
"question",
[
"Which samples have high TGFb activity?",
"Score all samples for pathway activity using PROGENy",
"Compute TF activity scores across samples",
"Rank samples by hallmark activity",
"Sample-level scoring with CollecTRI",
],
)
def test_score_samples_detected(self, question):
intent, _, _ = _detect_intent(question)
assert intent == "score_samples", f"Failed for: {question!r}"
# ── survival ─────────────────────────────────────────────────────────
@pytest.mark.parametrize(
"question",
[
"Is HIF1A activity associated with survival?",
"Kaplan-Meier analysis by subtype",
"Does STAT3 predict overall survival?",
"Prognostic value of TGFb pathway",
"Cox regression prognosis for subtype", # "cox regression" compound keyword
"survival analysis of the cohort",
],
)
def test_survival_detected(self, question):
intent, _, _ = _detect_intent(question)
assert intent == "survival", f"Failed for: {question!r}"
# ── correlate_continuous ─────────────────────────────────────────────
@pytest.mark.parametrize(
"question",
[
"Does TGFb correlate with tumor grade?", # removed "activity" to avoid tie
"STAT3 correlates with age", # unambiguous correlate keyword
"Correlates with tumor size",
"Continuous covariate regression with pathway scores",
],
)
def test_correlate_continuous_detected(self, question):
intent, _, _ = _detect_intent(question)
assert intent == "correlate_continuous", f"Failed for: {question!r}"
# ── unknown ──────────────────────────────────────────────────────────
@pytest.mark.parametrize(
"question",
[
"Tell me about this dataset",
"Load the data",
"What is PDAC?",
"",
" ",
],
)
def test_unknown_detected(self, question):
intent, _, _ = _detect_intent(question)
assert intent == "unknown", f"Failed for: {question!r}"
# ── confidence ───────────────────────────────────────────────────────
def test_high_confidence_on_multiple_keywords(self):
_, confidence, _ = _detect_intent("compare and differ between groups")
assert confidence == "high"
def test_medium_confidence_on_single_keyword(self):
_, confidence, _ = _detect_intent("compare the two groups")
assert confidence == "medium"
def test_low_confidence_on_no_keywords(self):
_, confidence, _ = _detect_intent("run everything")
assert confidence == "low"
# ── matched keywords ─────────────────────────────────────────────────
def test_matched_keywords_returned(self):
_, _, matched = _detect_intent("compare Classical vs Basal")
assert len(matched) > 0
def test_matched_keywords_are_substrings(self):
question = "compare Classical vs Basal tumors"
_, _, matched = _detect_intent(question)
for kw in matched:
assert kw in question.lower()
def test_unknown_has_empty_keywords(self):
_, _, matched = _detect_intent("tell me about the dataset")
assert matched == []
# ── case insensitivity ───────────────────────────────────────────────
def test_case_insensitive(self):
intent_lower, _, _ = _detect_intent("compare groups")
intent_upper, _, _ = _detect_intent("COMPARE GROUPS")
assert intent_lower == intent_upper == "compare_groups"
# ---------------------------------------------------------------------------
# dataset_plan_analysis — return structure
# ---------------------------------------------------------------------------
MOFFITT_ID = "gse71729_moffitt"
class TestPlanAnalysisReturnStructure:
def test_returns_dict(self):
result = dataset_plan_analysis(MOFFITT_ID, "compare Classical vs Basal")
assert isinstance(result, dict)
def test_required_top_level_keys(self):
result = dataset_plan_analysis(MOFFITT_ID, "compare Classical vs Basal")
for key in [
"dataset_id",
"user_question",
"detected_intent",
"confidence",
"matched_keywords",
"recommended_tools",
"required_inputs",
"assumptions",
"warnings",
"refusal_conditions",
]:
assert key in result, f"Missing key: {key}"
def test_dataset_id_echoed(self):
result = dataset_plan_analysis(MOFFITT_ID, "compare groups")
assert result["dataset_id"] == MOFFITT_ID
def test_question_echoed(self):
q = "compare Classical vs Basal"
result = dataset_plan_analysis(MOFFITT_ID, q)
assert result["user_question"] == q
def test_recommended_tools_is_list(self):
result = dataset_plan_analysis(MOFFITT_ID, "compare Classical vs Basal")
assert isinstance(result["recommended_tools"], list)
assert len(result["recommended_tools"]) > 0
def test_each_step_has_required_keys(self):
result = dataset_plan_analysis(MOFFITT_ID, "compare Classical vs Basal")
for step in result["recommended_tools"]:
for key in ["step", "tool", "purpose", "status"]:
assert key in step, f"Step missing '{key}': {step}"
def test_steps_are_numbered_from_1(self):
result = dataset_plan_analysis(MOFFITT_ID, "compare Classical vs Basal")
steps = result["recommended_tools"]
assert steps[0]["step"] == 1
def test_status_values_are_valid(self):
result = dataset_plan_analysis(MOFFITT_ID, "which samples have high activity")
valid_statuses = {"available", "not_implemented", "optional"}
for step in result["recommended_tools"]:
assert step["status"] in valid_statuses, f"Invalid status: {step['status']}"
def test_required_inputs_is_list(self):
result = dataset_plan_analysis(MOFFITT_ID, "compare Classical vs Basal")
assert isinstance(result["required_inputs"], list)
def test_assumptions_is_list(self):
result = dataset_plan_analysis(MOFFITT_ID, "compare Classical vs Basal")
assert isinstance(result["assumptions"], list)
def test_warnings_is_list(self):
result = dataset_plan_analysis(MOFFITT_ID, "compare Classical vs Basal")
assert isinstance(result["warnings"], list)
def test_refusal_conditions_is_list(self):
result = dataset_plan_analysis(MOFFITT_ID, "compare Classical vs Basal")
assert isinstance(result["refusal_conditions"], list)
# ---------------------------------------------------------------------------
# Intent routing in the full tool
# ---------------------------------------------------------------------------
class TestPlanAnalysisIntentRouting:
def test_compare_groups_intent(self):
result = dataset_plan_analysis(MOFFITT_ID, "compare Classical vs Basal")
assert result["detected_intent"] == "compare_groups"
def test_compare_groups_includes_de_tool(self):
result = dataset_plan_analysis(MOFFITT_ID, "compare Classical vs Basal")
tools = [s["tool"] for s in result["recommended_tools"]]
assert "decoupler_differential_expression" in tools
def test_compare_groups_includes_enrichment_tools(self):
result = dataset_plan_analysis(MOFFITT_ID, "compare Classical vs Basal")
tools = [s["tool"] for s in result["recommended_tools"]]
assert "decoupler_tf_enrichment_collectri" in tools
assert "decoupler_pathway_enrichment_progeny" in tools
assert "decoupler_hallmark_enrichment" in tools
def test_score_samples_intent(self):
result = dataset_plan_analysis(MOFFITT_ID, "which samples have high pathway activity?")
assert result["detected_intent"] == "score_samples"
def test_score_samples_includes_scoring_tool(self):
result = dataset_plan_analysis(MOFFITT_ID, "which samples have high TF activity?")
tools = [s["tool"] for s in result["recommended_tools"]]
assert "dataset_score_bulk_samples" in tools
def test_survival_intent(self):
result = dataset_plan_analysis(MOFFITT_ID, "associated with survival analysis of subtype")
assert result["detected_intent"] == "survival"
def test_survival_has_not_implemented_step(self):
result = dataset_plan_analysis(MOFFITT_ID, "Kaplan-Meier by subtype")
not_impl = [s for s in result["recommended_tools"] if s["status"] == "not_implemented"]
assert len(not_impl) >= 1
def test_survival_has_warning(self):
result = dataset_plan_analysis(MOFFITT_ID, "prognostic value of STAT3")
assert any("not" in w.lower() or "implement" in w.lower() for w in result["warnings"])
def test_correlate_intent(self):
result = dataset_plan_analysis(MOFFITT_ID, "does TGFb correlate with tumor grade?")
assert result["detected_intent"] == "correlate_continuous"
def test_correlate_has_not_implemented_step(self):
result = dataset_plan_analysis(MOFFITT_ID, "correlates with continuous covariate")
not_impl = [s for s in result["recommended_tools"] if s["status"] == "not_implemented"]
assert len(not_impl) >= 1
def test_unknown_intent(self):
result = dataset_plan_analysis(MOFFITT_ID, "tell me about this dataset")
assert result["detected_intent"] == "unknown"
def test_unknown_has_warning(self):
result = dataset_plan_analysis(MOFFITT_ID, "just load the data")
assert len(result["warnings"]) > 0
# ---------------------------------------------------------------------------
# Manifest-driven arg population
# ---------------------------------------------------------------------------
class TestManifestDrivenArgs:
def test_first_step_is_dataset_describe(self):
result = dataset_plan_analysis(MOFFITT_ID, "compare Classical vs Basal")
assert result["recommended_tools"][0]["tool"] == "dataset_describe"
def test_de_method_is_ttest_for_path_b(self):
# gse71729_moffitt is Path B (pre-normalised microarray)
result = dataset_plan_analysis(MOFFITT_ID, "compare Classical vs Basal")
de_steps = [
s
for s in result["recommended_tools"]
if s["tool"] == "decoupler_differential_expression"
]
assert len(de_steps) == 1
assert de_steps[0]["args_hint"]["method"] == "ttest"
def test_contrast_args_populated_from_manifest(self):
result = dataset_plan_analysis(MOFFITT_ID, "compare Classical vs Basal")
vc_steps = [
s for s in result["recommended_tools"] if s["tool"] == "dataset_validate_contrast"
]
assert len(vc_steps) == 1
hint = vc_steps[0]["args_hint"]
# Manifest has Classical vs Basal contrast
assert hint.get("test_group") or hint.get("group_column")
def test_assumptions_mention_path(self):
result = dataset_plan_analysis(MOFFITT_ID, "compare Classical vs Basal")
combined = " ".join(result["assumptions"]).lower()
assert "path" in combined or "pre-normalised" in combined or "ttest" in combined
def test_unknown_dataset_gives_warning(self):
result = dataset_plan_analysis("nonexistent_dataset_xyz", "compare groups")
assert any("not found" in w.lower() for w in result["warnings"])
def test_unknown_dataset_still_returns_plan(self):
result = dataset_plan_analysis("nonexistent_dataset_xyz", "compare groups")
assert result["detected_intent"] == "compare_groups"
assert len(result["recommended_tools"]) > 0
class TestPlanIncludesLoading:
"""The plan must say HOW to open the dataset, not just what to run.
Measured on dev 2026-08-04: `_workflow_for_intent` emitted a load step only
for `geo_series_matrix`, so every h5ad-over-`url` dataset (most of the
registry) returned a plan with no loading step. The agent then guessed an
entry point — 3-5 wasted steps per run, and for the DE query it never found
a working one. Loading now comes from `_build_loading_plan`.
"""
def test_h5ad_url_dataset_gets_a_load_step(self):
plan = dataset_plan_analysis(
dataset_id="gse28735_pdac",
user_question="Compare tumor vs normal for differential expression",
)
tools = [s["tool"] for s in plan["recommended_tools"]]
assert "decoupler_load_url_counts" in tools, tools
# ...and it must come before the analysis it feeds.
assert tools.index("decoupler_load_url_counts") < tools.index(
"decoupler_differential_expression"
)
def test_load_step_uses_the_precollapsed_url(self):
"""The uncollapsed URL costs two extra steps (annotate + collapse)."""
plan = dataset_plan_analysis(
dataset_id="gse28735_pdac", user_question="differential expression tumor vs normal"
)
load = next(
s for s in plan["recommended_tools"] if s["tool"] == "decoupler_load_url_counts"
)
assert "collapsed" in load["args_hint"]["url_or_path"]
tools = [s["tool"] for s in plan["recommended_tools"]]
assert "decoupler_annotate_probes_with_gpl" not in tools
assert "decoupler_collapse_probes_to_genes" not in tools
def test_registered_dataset_plan_omits_inspect_data(self):
"""Efficiency rule 4: data_level/analysis_path are manifest facts."""
for did in ("gse28735_pdac", "paca_au_rnaseq"):
plan = dataset_plan_analysis(dataset_id=did, user_question="compare groups")
tools = [s["tool"] for s in plan["recommended_tools"]]
assert "decoupler_inspect_data" not in tools, (did, tools)
def test_single_cell_plan_uses_the_sc_loader(self):
"""Path P must not be handed a bulk flat-file loader (ADR-0006)."""
plan = dataset_plan_analysis(
dataset_id="gse155698_steele", user_question="compare tumor vs normal"
)
tools = [s["tool"] for s in plan["recommended_tools"]]
assert "decoupler_load_and_visualize_data" in tools, tools
assert "decoupler_load_url_counts" not in tools
def test_steps_are_numbered_consecutively(self):
plan = dataset_plan_analysis(dataset_id="gse28735_pdac", user_question="compare groups")
nums = [s["step"] for s in plan["recommended_tools"]]
assert nums == list(range(1, len(nums) + 1)), nums
def test_unknown_dataset_still_returns_a_plan(self):
"""A missing manifest must degrade to a warning, not an exception."""
plan = dataset_plan_analysis(dataset_id="not_a_dataset", user_question="compare groups")
assert plan["recommended_tools"]
assert plan["warnings"]