Skip to content

Commit 69c956c

Browse files
committed
refactor: test-fixture cleanup
1 parent 1421dc6 commit 69c956c

5 files changed

Lines changed: 121 additions & 123 deletions

File tree

Lines changed: 100 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,100 @@
1+
from collections.abc import Iterable
2+
3+
from osa_tool.operations.analysis.paper_claims.claim_schemas import ClaimCandidateResponse
4+
from osa_tool.operations.analysis.paper_claims.models import ExtractedClaim, HeadingMeta, PaperSection
5+
6+
DEFAULT_SECTION_TEXT = "The model uses BERT-base without fine-tuning."
7+
8+
9+
class FakeAsyncHandler:
10+
"""Queued async LLM handler used by paper-claims unit tests."""
11+
12+
def __init__(self, responses: Iterable[str | BaseException]):
13+
self.responses = iter(responses)
14+
self.prompts: list[str] = []
15+
16+
async def async_request(self, prompt, system_message=None, retry_delay=1):
17+
self.prompts.append(prompt)
18+
response = next(self.responses)
19+
if isinstance(response, BaseException):
20+
raise response
21+
return response
22+
23+
24+
class ModelTrackingFakeHandler(FakeAsyncHandler):
25+
"""Fake handler that records the model producing each successful response."""
26+
27+
def __init__(self, responses: Iterable[str | BaseException], models: Iterable[str]):
28+
super().__init__(responses)
29+
self.models = iter(models)
30+
self.last_successful_model: str | None = None
31+
32+
async def async_request(self, prompt, system_message=None, retry_delay=1):
33+
response = await super().async_request(prompt, system_message, retry_delay)
34+
self.last_successful_model = next(self.models)
35+
return response
36+
37+
38+
def word_token_count(text: str, _encoder: str = "fake") -> int:
39+
return len(text.split())
40+
41+
42+
def make_paper_section(
43+
text: str = DEFAULT_SECTION_TEXT,
44+
*,
45+
section_id: str = "s001",
46+
name: str = "Method",
47+
heading_raw: str = "2. Method",
48+
heading_level: int = 1,
49+
heading_numbering: str | None = "2",
50+
) -> PaperSection:
51+
return PaperSection(
52+
section_id=section_id,
53+
name=name,
54+
text=text,
55+
heading_meta=HeadingMeta(raw=heading_raw, level=heading_level, numbering=heading_numbering),
56+
)
57+
58+
59+
def make_extracted_claim(
60+
claim_id: str,
61+
claim: str,
62+
*,
63+
original_text: str | None = None,
64+
category: str = "infrastructure",
65+
value: str | None = None,
66+
verifiability: str = "high",
67+
section_id: str = "s001",
68+
section_name: str = "Method",
69+
section_heading_raw: str | None = "2. Method",
70+
contradiction: bool = False,
71+
) -> ExtractedClaim:
72+
return ExtractedClaim(
73+
claim_id=claim_id,
74+
claim=claim,
75+
original_text=claim if original_text is None else original_text,
76+
category=category,
77+
value=value,
78+
verifiability=verifiability,
79+
section_id=section_id,
80+
section_name=section_name,
81+
section_heading_raw=section_heading_raw,
82+
contradiction=contradiction,
83+
)
84+
85+
86+
def make_claim_candidate(
87+
*,
88+
claim: str,
89+
original_text: str,
90+
category: str = "model_architecture",
91+
value: str | None = None,
92+
verifiability: str = "high",
93+
) -> ClaimCandidateResponse:
94+
return ClaimCandidateResponse(
95+
claim=claim,
96+
original_text=original_text,
97+
category=category,
98+
value=value,
99+
verifiability=verifiability,
100+
)

tests/unit/operations/analysis/paper_claims/test_claim_deduplicator.py

Lines changed: 5 additions & 29 deletions
Original file line numberDiff line numberDiff line change
@@ -5,36 +5,12 @@
55

66
from osa_tool.operations.analysis.paper_claims.claim_deduplicator import ClaimDeduplicator
77
from osa_tool.operations.analysis.paper_claims.claim_extractor import ClaimExtractor
8-
from osa_tool.operations.analysis.paper_claims.models import ExtractedClaim
98
from osa_tool.utils.prompts_builder import PromptLoader
10-
11-
12-
class FakeHandler:
13-
def __init__(self, responses):
14-
self.responses = iter(responses)
15-
self.prompts = []
16-
17-
async def async_request(self, prompt, system_message=None, retry_delay=1):
18-
self.prompts.append(prompt)
19-
return next(self.responses)
20-
21-
22-
def fake_token_count(text, _encoder="fake"):
23-
return len(text.split())
24-
25-
26-
def extracted_claim(claim_id: str, claim: str) -> ExtractedClaim:
27-
return ExtractedClaim(
28-
claim_id=claim_id,
29-
claim=claim,
30-
original_text=claim,
31-
category="infrastructure",
32-
value=None,
33-
verifiability="high",
34-
section_id="s001",
35-
section_name="Method",
36-
section_heading_raw="2. Method",
37-
)
9+
from tests.unit.operations.analysis.paper_claims.fixtures import (
10+
FakeAsyncHandler as FakeHandler,
11+
make_extracted_claim as extracted_claim,
12+
word_token_count as fake_token_count,
13+
)
3814

3915

4016
def deduplicator(handler, *, dedup_batch_size=100) -> ClaimDeduplicator:

tests/unit/operations/analysis/paper_claims/test_claim_extractor.py

Lines changed: 8 additions & 53 deletions
Original file line numberDiff line numberDiff line change
@@ -4,60 +4,15 @@
44
import pytest
55

66
from osa_tool.operations.analysis.paper_claims.claim_extractor import ClaimExtractor
7-
from osa_tool.operations.analysis.paper_claims.models import ExtractedClaim, HeadingMeta, PaperSection
7+
from osa_tool.operations.analysis.paper_claims.models import HeadingMeta, PaperSection
88
from osa_tool.utils.prompts_builder import PromptLoader
9-
10-
11-
class FakeHandler:
12-
def __init__(self, responses):
13-
self.responses = iter(responses)
14-
self.prompts = []
15-
16-
async def async_request(self, prompt, system_message=None, retry_delay=1):
17-
self.prompts.append(prompt)
18-
response = next(self.responses)
19-
if isinstance(response, BaseException):
20-
raise response
21-
return response
22-
23-
24-
class ModelTrackingFakeHandler(FakeHandler):
25-
def __init__(self, responses, models):
26-
super().__init__(responses)
27-
self.models = iter(models)
28-
self.last_successful_model = None
29-
30-
async def async_request(self, prompt, system_message=None, retry_delay=1):
31-
response = await super().async_request(prompt, system_message, retry_delay)
32-
self.last_successful_model = next(self.models)
33-
return response
34-
35-
36-
def fake_token_count(text, _encoder="fake"):
37-
return len(text.split())
38-
39-
40-
def section() -> PaperSection:
41-
return PaperSection(
42-
section_id="s001",
43-
name="Method",
44-
text="The model uses BERT-base without fine-tuning.",
45-
heading_meta=HeadingMeta(raw="2. Method", level=1, numbering="2"),
46-
)
47-
48-
49-
def extracted_claim(claim_id: str, claim: str) -> ExtractedClaim:
50-
return ExtractedClaim(
51-
claim_id=claim_id,
52-
claim=claim,
53-
original_text=claim,
54-
category="infrastructure",
55-
value=None,
56-
verifiability="high",
57-
section_id="s001",
58-
section_name="Method",
59-
section_heading_raw="2. Method",
60-
)
9+
from tests.unit.operations.analysis.paper_claims.fixtures import (
10+
FakeAsyncHandler as FakeHandler,
11+
ModelTrackingFakeHandler,
12+
make_extracted_claim as extracted_claim,
13+
make_paper_section as section,
14+
word_token_count as fake_token_count,
15+
)
6116

6217

6318
def test_section_filter_prompt_contains_valid_json_example():

tests/unit/operations/analysis/paper_claims/test_claim_input_planner.py

Lines changed: 4 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,10 @@
33

44
from osa_tool.operations.analysis.paper_claims.claim_input_planner import ClaimInputPlanner
55
from osa_tool.operations.analysis.paper_claims.models import HeadingMeta, PaperSection
6+
from tests.unit.operations.analysis.paper_claims.fixtures import (
7+
make_paper_section as section,
8+
word_token_count as fake_token_count,
9+
)
610

711

812
class FakeWordCodec:
@@ -25,19 +29,6 @@ def decode(tokens):
2529
return "".join(tokens)
2630

2731

28-
def fake_token_count(text, _encoder="fake"):
29-
return len(text.split())
30-
31-
32-
def section(text="The model uses BERT-base without fine-tuning."):
33-
return PaperSection(
34-
section_id="s001",
35-
name="Method",
36-
text=text,
37-
heading_meta=HeadingMeta(raw="2. Method", level=1, numbering="2"),
38-
)
39-
40-
4132
def test_section_ancestors_use_numbering_when_marker_flattens_heading_levels():
4233
sections = [
4334
PaperSection(

tests/unit/operations/analysis/paper_claims/test_claim_validation.py

Lines changed: 4 additions & 28 deletions
Original file line numberDiff line numberDiff line change
@@ -1,37 +1,13 @@
11
import pytest
22

3-
from osa_tool.operations.analysis.paper_claims.claim_schemas import (
4-
ClaimCandidateResponse,
5-
)
63
from osa_tool.operations.analysis.paper_claims.claim_validation import (
74
partition_valid_claim_candidates,
85
validate_claim_candidate,
96
)
10-
from osa_tool.operations.analysis.paper_claims.models import HeadingMeta, PaperSection
11-
12-
13-
def make_section(text: str) -> PaperSection:
14-
return PaperSection(
15-
section_id="s001",
16-
name="Method",
17-
text=text,
18-
heading_meta=HeadingMeta(raw="2. Method", level=1, numbering="2"),
19-
)
20-
21-
22-
def make_candidate(
23-
*,
24-
claim: str,
25-
original_text: str,
26-
category: str = "model_architecture",
27-
) -> ClaimCandidateResponse:
28-
return ClaimCandidateResponse(
29-
claim=claim,
30-
original_text=original_text,
31-
category=category,
32-
value=None,
33-
verifiability="high",
34-
)
7+
from tests.unit.operations.analysis.paper_claims.fixtures import (
8+
make_claim_candidate as make_candidate,
9+
make_paper_section as make_section,
10+
)
3511

3612

3713
def test_validate_claim_candidate_accepts_exact_source_text():

0 commit comments

Comments
 (0)