File size: 13,804 Bytes
eea689d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
"""Escalating not-found resolution (docs/prd.md section 7, issue #10).

"A field should only be marked not found after an escalating search, an LLM
pass, an embedding pass, and the two combined, come back empty, not after a
single failed prompt. A missing field that was actually searched for and
confirmed absent is honest. A missing field from one lazy pass is a silent
gap wearing an honest label."

This module owns the *not-found* path only. Fields the first LLM pass already
extracted go to the consensus gate (consensus.py, issue #9), not here. For a
field the LLM pass left empty, the escalation is:

  1. LLM pass -- already run by llm_extraction.extract_case_with_llm; its
     empty result is the reason we're escalating at all.
  2. Embedding pass -- ColPali retrieval for the field's query. Produces no
     value (sub-page localization is not built yet: #49), only "which page
     looks most relevant, and how distinctively" (rank_pages_with_margin).
     That page is a hint for the combined pass, not an answer on its own.
  3. Combined pass -- a targeted, single-field LLM re-extraction given the
     report text *and* the page image the embedding pass surfaced. This is
     the literal "two combined": the text model reading the exact scan
     region the visual model flagged, a genuinely different query from the
     bulk first pass (which reads text only, all fields at once).

Only when all three come back empty is the field flagged not-found -- via
schema.ChecklistField.flag_not_found, which itself refuses a search_log with
fewer than three recorded passes, so a lazy single-pass miss cannot wear the
"searched and not found" label.

The orchestration core (resolve_field) takes the embedding and combined
passes as injected callables so it's testable without live models; the real
wirings to colpali_retrieval and llm_extraction are thin builders below.
"""

from __future__ import annotations

import base64
from dataclasses import dataclass
from enum import Enum
from pathlib import Path
from typing import Any, Callable, Optional

from endopath import colpali_retrieval, consensus, llm_extraction, precomputed_retrieval
from endopath.schema import Case, ChecklistField, EvidenceSpan, FieldStatus


class ResolutionOutcome(str, Enum):
    RESOLVED = "resolved"  # a value exists (LLM pass) or was recovered (combined pass)
    NOT_FOUND = "not_found"  # all three passes empty -> flagged, honestly earned
    NOT_APPLICABLE = "not_applicable"  # short-circuited, e.g. molecular pre-2023


@dataclass
class VisualHit:
    """What the embedding pass found: the best page and how distinctive it
    was. margin is rank_pages_with_margin's best-minus-second-best gap, which
    (unlike raw score) is comparable across field queries -- see #18/#9."""

    page_number: int
    margin: float
    top_score: float


@dataclass
class ExtractedValue:
    """What the combined pass recovered for a single field."""

    value: Any
    confidence: Optional[float] = None
    evidence: Optional[EvidenceSpan] = None


def resolve_field(
    field_name: str,
    field: ChecklistField,
    *,
    embedding_probe: Callable[[], Optional[VisualHit]],
    combined_extract: Callable[[Optional[VisualHit]], Optional[ExtractedValue]],
    margin_threshold: float = consensus.DEFAULT_VISUAL_MARGIN_THRESHOLD,
) -> ResolutionOutcome:
    """Run the escalating search for one field and mutate it in place.

    embedding_probe() returns the field's VisualHit, or None if the pass
    couldn't run (no page images / ColPali unavailable). combined_extract(hit)
    returns the recovered value, or None if the combined pass came back empty.
    Both are injected so the orchestration is testable without live models.
    """
    # A field the checklist itself says doesn't apply (molecular
    # classification on a pre-2023 report, FIGO grade on a non-endometrioid
    # histotype) already carries its reason from schema.apply_applicability_
    # rules. Running the full LLM/embedding/combined search on it would burn
    # a model call to "find" something the protocol says can't be there --
    # short-circuit instead, per #10's molecular-classification requirement.
    if not field.applicable:
        return ResolutionOutcome.NOT_APPLICABLE

    # The first LLM pass already found a value: not a not-found case. The
    # consensus gate (#9) owns found fields; leave this one untouched.
    if field.value is not None:
        return ResolutionOutcome.RESOLVED

    search_log = ["llm pass: no value extracted"]

    hit = embedding_probe()
    if hit is None:
        search_log.append("embedding pass: skipped (no page images or ColPali unavailable)")
    elif hit.margin >= margin_threshold:
        search_log.append(
            f"embedding pass: distinctive match on page {hit.page_number} "
            f"(margin {hit.margin:.2f}) -- field may be present, first LLM pass missed it"
        )
    else:
        search_log.append(
            f"embedding pass: no distinctive match "
            f"(best page {hit.page_number}, margin {hit.margin:.2f})"
        )

    extracted = combined_extract(hit)
    if extracted is not None and extracted.value is not None:
        search_log.append(f"combined pass: recovered value {extracted.value!r}")
        field.value = extracted.value
        field.confidence = extracted.confidence
        field.evidence = extracted.evidence
        field.status = FieldStatus.NEEDS_REVIEW
        field.search_log.extend(search_log)
        return ResolutionOutcome.RESOLVED

    search_log.append("combined pass: no value recovered")
    # Exactly the three passes above -- satisfies flag_not_found's guard that
    # a not-found label be earned by a recorded LLM + embedding + combined
    # search, not a single failed prompt.
    field.flag_not_found(search_log)
    return ResolutionOutcome.NOT_FOUND


# --- Real wirings (thin glue to the model modules; the logic above is the
# tested part). ------------------------------------------------------------

_COMBINED_PROMPT = """\
You are re-checking a single CAP/ICCR endometrial carcinoma checklist field \
that a first extraction pass left empty. The field is: {field_name}.

{image_note}Read the report text below (and the page image, if provided) and \
call record_field. Only set a value if it is genuinely supported; leave value \
and evidence_quote null rather than guessing. evidence_quote must be an exact, \
verbatim substring of the report text -- do not paraphrase.

REPORT TEXT:
{report_text}
"""


def _embed_pages_if_available(image_paths: list[Path]):
    """Embed a case's pages once, so a per-field escalation loop doesn't
    re-embed the same images for every field. Returns None when the embedding
    pass can't run at all (no images, or ColPali not installed)."""
    if not image_paths or not colpali_retrieval.is_available():
        return None
    return colpali_retrieval.embed_pages(image_paths)


def _build_embedding_probe(
    field_name: str, page_embeddings, margin_threshold: float
) -> Callable[[], Optional[VisualHit]]:
    def probe() -> Optional[VisualHit]:
        if page_embeddings is None or field_name not in colpali_retrieval.FIELD_QUERIES:
            return None
        query = colpali_retrieval.FIELD_QUERIES[field_name]
        page_number, top_score, margin = colpali_retrieval.rank_pages_with_margin(
            query, page_embeddings
        )
        return VisualHit(page_number=page_number, margin=margin, top_score=top_score)

    return probe


def _build_precomputed_embedding_probe(
    field_name: str,
    page_embeddings: Optional[list],
    query_embeddings: dict,
    margin_threshold: float,
) -> Callable[[], Optional[VisualHit]]:
    """Like _build_embedding_probe, but scores precomputed page vectors against
    precomputed field-query vectors (precomputed_retrieval, issue #21) with
    numpy MaxSim -- no torch, no live ColPali. `page_embeddings` is the case's
    list of per-page arrays; `query_embeddings` maps field name -> query vector.
    Both come from data/embeddings (issue #19's precompute). Returns None when
    the case has no precomputed pages, or this field has no precomputed query.
    """

    def probe() -> Optional[VisualHit]:
        if page_embeddings is None or field_name not in query_embeddings:
            return None
        page_number, top_score, margin = precomputed_retrieval.rank_pages_with_margin(
            query_embeddings[field_name], page_embeddings
        )
        return VisualHit(page_number=page_number, margin=margin, top_score=top_score)

    return probe


def _build_combined_extract(
    field_name: str,
    report_text: str,
    image_paths: list[Path],
    client=None,
) -> Callable[[Optional[VisualHit]], Optional[ExtractedValue]]:
    def extract(hit: Optional[VisualHit]) -> Optional[ExtractedValue]:
        if field_name not in llm_extraction._FIELD_VALUE_SCHEMAS:
            return None
        page_image: Optional[Path] = None
        page_number: Optional[int] = None
        if hit is not None and image_paths:
            idx = hit.page_number - 1
            if 0 <= idx < len(image_paths):
                page_image = image_paths[idx]
                page_number = hit.page_number
        return _combined_llm_extract(
            field_name, report_text, page_image, page_number, client=client
        )

    return extract


def _combined_llm_extract(
    field_name: str,
    report_text: str,
    page_image: Optional[Path],
    page_number: Optional[int],
    *,
    client=None,
    model: str = llm_extraction.MODEL,
) -> Optional[ExtractedValue]:
    """Targeted, single-field, optionally-multimodal re-extraction: the report
    text plus the one page image the embedding pass surfaced. Reuses
    llm_extraction's per-field value schemas and enum coercion so the combined
    pass can't accept a value the first pass' schema would have rejected."""
    client = client or llm_extraction.get_client()

    value_schema = llm_extraction._FIELD_VALUE_SCHEMAS[field_name]
    tool = {
        "name": "record_field",
        "description": f"Record the extracted value for the endometrial checklist field '{field_name}'.",
        "input_schema": llm_extraction._field_schema(value_schema),
    }

    content: list[dict] = []
    image_note = ""
    if page_image is not None:
        img_bytes = Path(page_image).read_bytes()
        content.append(
            {
                "type": "image",
                "source": {
                    "type": "base64",
                    "media_type": "image/png",
                    "data": base64.standard_b64encode(img_bytes).decode("ascii"),
                },
            }
        )
        image_note = (
            "A visual-retrieval model flagged the attached page image as the "
            "most likely location for this field; read it carefully.\n\n"
        )
    content.append(
        {
            "type": "text",
            "text": _COMBINED_PROMPT.format(
                field_name=field_name, image_note=image_note, report_text=report_text
            ),
        }
    )

    response = client.messages.create(
        model=model,
        max_tokens=1024,
        tools=[tool],
        tool_choice={"type": "tool", "name": "record_field"},
        messages=[{"role": "user", "content": content}],
    )
    tool_call = next(block for block in response.content if block.type == "tool_use")
    raw = tool_call.input

    value = raw.get("value")
    if value is None:
        return None
    if field_name in llm_extraction._ENUM_FIELDS:
        try:
            value = llm_extraction._ENUM_FIELDS[field_name](value)
        except ValueError:
            return None

    confidence = raw.get("confidence")
    quote = raw.get("evidence_quote")
    evidence = None
    if quote:
        verified = quote in report_text
        char_start = report_text.find(quote) if verified else None
        evidence = EvidenceSpan(
            quote=quote,
            char_start=char_start,
            char_end=(char_start + len(quote)) if char_start is not None else None,
            page_number=page_number,
            source="combined" if verified else "combined_unverified_quote",
        )
        if not verified and confidence is not None:
            confidence = min(confidence, 0.2)

    return ExtractedValue(value=value, confidence=confidence, evidence=evidence)


def resolve_case_not_found(
    case: Case,
    report_text: str,
    image_paths: Optional[list[Path]] = None,
    *,
    client=None,
    margin_threshold: float = consensus.DEFAULT_VISUAL_MARGIN_THRESHOLD,
) -> dict[str, ResolutionOutcome]:
    """Run the escalating not-found search across every field of a case whose
    first LLM pass came back empty. Embeds the case's pages once, then hands
    each field its own embedding + combined pass. Mutates case.checklist in
    place; returns the per-field outcome.

    Assumes case.checklist has already had the LLM extraction pass and
    schema.apply_applicability_rules run (llm_extraction.extract_case_with_llm
    does both) -- resolve_field short-circuits any field already marked
    not-applicable rather than searching for it.
    """
    image_paths = image_paths or []
    page_embeddings = _embed_pages_if_available(image_paths)

    outcomes: dict[str, ResolutionOutcome] = {}
    for name, field in case.checklist.all_fields().items():
        probe = _build_embedding_probe(name, page_embeddings, margin_threshold)
        extract = _build_combined_extract(name, report_text, image_paths, client=client)
        outcomes[name] = resolve_field(
            name,
            field,
            embedding_probe=probe,
            combined_extract=extract,
            margin_threshold=margin_threshold,
        )
    return outcomes