62 lines
1.8 KiB
Python
62 lines
1.8 KiB
Python
"""Evidence-safe context packing for generation."""
|
|
from __future__ import annotations
|
|
|
|
from collections.abc import Callable, Sequence
|
|
from dataclasses import dataclass
|
|
|
|
from .models import Evidence
|
|
|
|
|
|
TokenCounter = Callable[[str], int]
|
|
|
|
|
|
def conservative_token_count(text: str) -> int:
|
|
"""Dependency-free estimate, conservative for Vietnamese text."""
|
|
return (len(text) + 2) // 3 if text else 0
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class PackedContext:
|
|
evidence: tuple[Evidence, ...]
|
|
omitted_evidence_ids: tuple[str, ...]
|
|
estimated_tokens: int
|
|
|
|
|
|
def pack_evidence(
|
|
evidence: Sequence[Evidence],
|
|
*,
|
|
max_tokens: int,
|
|
max_items: int | None = None,
|
|
token_counter: TokenCounter = conservative_token_count,
|
|
block_overhead_tokens: int = 4,
|
|
) -> PackedContext:
|
|
"""Pack whole blocks in retrieval order; never truncate clinical text."""
|
|
if max_tokens < 1:
|
|
raise ValueError("max_tokens must be positive")
|
|
if max_items is not None and max_items < 1:
|
|
raise ValueError("max_items must be positive or None")
|
|
if block_overhead_tokens < 0:
|
|
raise ValueError("block_overhead_tokens must be non-negative")
|
|
|
|
selected: list[Evidence] = []
|
|
omitted: list[str] = []
|
|
seen: set[str] = set()
|
|
used = 0
|
|
for item in evidence:
|
|
if item.evidence_id in seen:
|
|
continue
|
|
seen.add(item.evidence_id)
|
|
count = token_counter(item.text)
|
|
if count < 0:
|
|
raise ValueError("token_counter must return a non-negative value")
|
|
cost = count + block_overhead_tokens
|
|
if (
|
|
(max_items is not None and len(selected) >= max_items)
|
|
or used + cost > max_tokens
|
|
):
|
|
omitted.append(item.evidence_id)
|
|
continue
|
|
selected.append(item)
|
|
used += cost
|
|
return PackedContext(tuple(selected), tuple(omitted), used)
|