| import pytest |
| import uuid |
|
|
| from src.models.claim import Claim, ClaimType, Polarity, StudyDesign, Entity, EntityType |
| from src.models.paper import Paper |
| from src.models.contradiction import ContradictionPair, ContradictionType |
| from src.graph.claim_graph import build_claim_graph, compute_consensus_scores |
|
|
| @pytest.fixture |
| def sample_data(): |
| paper_1 = Paper( |
| pmid="11111", |
| title="Study 1 on Metformin", |
| authors=["Author One"], |
| year=2020, |
| journal="Journal of Diabetes", |
| abstract_text="Metformin reduces cancer risk." |
| ) |
| paper_2 = Paper( |
| pmid="22222", |
| title="Study 2 on Metformin", |
| authors=["Author Two"], |
| year=2023, |
| journal="Cancer Letters", |
| abstract_text="Metformin increases cancer risk." |
| ) |
| |
| entity_metformin = Entity(text="Metformin", canonical_id="MeSH:D001241", entity_type=EntityType.DRUG) |
| entity_cancer = Entity(text="Cancer", canonical_id="MeSH:D009369", entity_type=EntityType.DISEASE) |
| |
| claim_1 = Claim( |
| id=uuid.uuid4(), |
| text="Metformin reduces breast cancer risk.", |
| paper_id="11111", |
| authors=["Author One"], |
| year=2020, |
| confidence_score=1.0, |
| claim_type=ClaimType.CAUSAL, |
| polarity=Polarity.NEGATIVE, |
| entities=[entity_metformin, entity_cancer], |
| population="humans", |
| context="general", |
| quote_anchor="reduces risk", |
| study_design=StudyDesign.RCT |
| ) |
| |
| claim_2 = Claim( |
| id=uuid.uuid4(), |
| text="Metformin increases breast cancer risk.", |
| paper_id="22222", |
| authors=["Author Two"], |
| year=2023, |
| confidence_score=1.0, |
| claim_type=ClaimType.CAUSAL, |
| polarity=Polarity.POSITIVE, |
| entities=[entity_metformin, entity_cancer], |
| population="humans", |
| context="general", |
| quote_anchor="increases risk", |
| study_design=StudyDesign.RCT |
| ) |
|
|
| contradiction = ContradictionPair( |
| claim_a=claim_1, |
| claim_b=claim_2, |
| contradiction_score=0.95, |
| contradiction_type=ContradictionType.DIRECTION_REVERSAL, |
| explanation="Claim 1 reduces risk, Claim 2 increases risk.", |
| scope_note="", |
| is_genuine=True |
| ) |
| |
| return [claim_1, claim_2], [contradiction], [paper_1, paper_2] |
|
|
| def test_build_claim_graph(sample_data): |
| claims, contradictions, papers = sample_data |
| G = build_claim_graph(claims, contradictions, papers) |
| |
| |
| assert G.number_of_nodes() == 6 |
| |
| |
| assert G.has_node("11111") |
| assert G.nodes["11111"]["type"] == "paper" |
| assert G.nodes["11111"]["title"] == "Study 1 on Metformin" |
| |
| |
| claim_1_id = str(claims[0].id) |
| assert G.has_node(claim_1_id) |
| assert G.nodes[claim_1_id]["type"] == "claim" |
| assert G.nodes[claim_1_id]["polarity"] == Polarity.NEGATIVE.value |
| |
| |
| assert G.has_node("MeSH:D001241") |
| assert G.nodes["MeSH:D001241"]["type"] == "entity" |
| assert G.nodes["MeSH:D001241"]["entity_type"] == EntityType.DRUG.value |
|
|
| |
| |
| assert G.has_edge("11111", claim_1_id) |
| assert list(G["11111"][claim_1_id].values())[0]["type"] == "EXTRACTED_FROM" |
| |
| |
| assert G.has_edge(claim_1_id, "MeSH:D001241") |
| assert list(G[claim_1_id]["MeSH:D001241"].values())[0]["type"] == "MENTIONS" |
| |
| |
| claim_2_id = str(claims[1].id) |
| assert G.has_edge(claim_1_id, claim_2_id) |
| |
| |
| edges_1_2 = G[claim_1_id][claim_2_id] |
| assert len(edges_1_2) == 1 |
| edge_1_2_data = list(edges_1_2.values())[0] |
| assert edge_1_2_data["type"] == "CONTRADICTS" |
| assert edge_1_2_data["score"] == 0.95 |
| |
| |
| assert G.has_edge(claim_2_id, claim_1_id) |
| edges_2_1 = G[claim_2_id][claim_1_id] |
| assert len(edges_2_1) == 2 |
| |
| contradicts_edge = next(e for e in edges_2_1.values() if e["type"] == "CONTRADICTS") |
| supersedes_edge = next(e for e in edges_2_1.values() if e["type"] == "SUPERSEDES") |
| |
| assert contradicts_edge["score"] == 0.95 |
| |
| |
| assert supersedes_edge["score"] == 0.95 |
| assert supersedes_edge["explanation"] == "Claim 1 reduces risk, Claim 2 increases risk." |
| assert supersedes_edge["scope_note"] == "" |
| assert supersedes_edge["is_genuine"] is True |
|
|
| def test_compute_consensus_scores(sample_data): |
| claims, contradictions, papers = sample_data |
| G = build_claim_graph(claims, contradictions, papers) |
| |
| scores = compute_consensus_scores(G) |
| |
| claim_1_id = str(claims[0].id) |
| claim_2_id = str(claims[1].id) |
| |
| |
| |
| assert scores[claim_1_id] == 0.0 |
| assert scores[claim_2_id] == 0.0 |
|
|
|
|
| def test_compute_consensus_scores_performance(): |
| import time |
| |
| claims = [] |
| entities = [ |
| Entity(text=f"Entity {i}", canonical_id=f"canonical_{i}", entity_type=EntityType.DRUG) |
| for i in range(50) |
| ] |
| for i in range(500): |
| |
| claim_entities = [entities[i % 50], entities[(i + 1) % 50]] |
| claims.append( |
| Claim( |
| id=uuid.uuid4(), |
| text=f"Claim text {i}", |
| paper_id=f"paper_{i // 5}", |
| authors=["Author"], |
| year=2020 + (i % 5), |
| confidence_score=0.9, |
| claim_type=ClaimType.CAUSAL if i % 2 == 0 else ClaimType.CORRELATIONAL, |
| polarity=Polarity.POSITIVE if i % 3 == 0 else Polarity.NEGATIVE, |
| entities=claim_entities, |
| population="humans", |
| context="general", |
| quote_anchor="anchor", |
| study_design=StudyDesign.RCT |
| ) |
| ) |
| |
| |
| contradictions = [] |
| for i in range(100): |
| claim_a = claims[i] |
| claim_b = claims[i + 100] |
| contradictions.append( |
| ContradictionPair( |
| claim_a=claim_a, |
| claim_b=claim_b, |
| contradiction_score=0.8, |
| contradiction_type=ContradictionType.DIRECTION_REVERSAL, |
| explanation="Contradiction", |
| scope_note="", |
| is_genuine=True |
| ) |
| ) |
| |
| papers = [ |
| Paper(pmid=f"paper_{i}", title="Title", authors=["Author"], year=2020, journal="Journal", abstract_text="Abstract") |
| for i in range(100) |
| ] |
| |
| |
| G = build_claim_graph(claims, contradictions, papers) |
| |
| |
| start_time = time.perf_counter() |
| scores = compute_consensus_scores(G) |
| duration = time.perf_counter() - start_time |
| |
| |
| assert len(scores) == 500 |
| assert duration < 0.1, f"Optimization failed: took {duration:.4f} seconds (limit: 0.1s)" |
| |
| for cid in scores: |
| assert 0.0 <= scores[cid] <= 1.0 |
|
|