From 427d3fbbd858779f5d0f2fcd835b5bbd4c4ebd4c Mon Sep 17 00:00:00 2001 From: Aakash Date: Sun, 6 Sep 2026 04:21:02 +0530 Subject: [PATCH] feat: add budgeted shared-context wrappers Add opt-in per-source prefixes and separators that are charged against the complete rendered context while preserving legacy plain-string accounting. Return prompt-ready context text, keep wrappers out of ranking and spans, document the API, and prepare v0.5.0. --- .github/workflows/ci.yml | 12 +- AGENTS.md | 8 +- README.md | 27 +- docs/api-reference.md | 10 +- docs/configuration-and-api.md | 33 ++- docs/getting-started.md | 27 +- docs/guarantees-and-limitations.md | 37 +-- docs/how-it-works.md | 38 +-- docs/index.md | 4 +- docs/multi-source-context.md | 94 ++++--- docs/semantic-and-async.md | 3 +- pyproject.toml | 2 +- src/trimwise/__init__.py | 2 + src/trimwise/composition.py | 101 ++++++-- src/trimwise/measurement.py | 17 ++ src/trimwise/models.py | 17 +- src/trimwise/selection.py | 26 +- src/trimwise/trimmer.py | 104 +++++--- tests/test_api.py | 1 + tests/test_context_rendering.py | 391 +++++++++++++++++++++++++++++ 20 files changed, 794 insertions(+), 160 deletions(-) create mode 100644 tests/test_context_rendering.py diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index b66cefa..80f2a72 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -68,10 +68,12 @@ jobs: - run: >- wheel-test/bin/python -c "import asyncio, importlib.metadata; - from trimwise import ContextSourceResult, ContextTrimResult, Trimmer; - assert importlib.metadata.version('trimwise') == '0.4.1'; - sync_result = Trimmer().trim_context(['a'], 1, unit='characters'); - async_result = asyncio.run(Trimmer().atrim_context(['a'], 1, unit='characters')); + from trimwise import ContextSource, ContextSourceResult, ContextTrimResult, Trimmer; + assert importlib.metadata.version('trimwise') == '0.5.0'; + sources = [ContextSource('a', 'Source: A\n')]; + sync_result = Trimmer().trim_context(sources, 11, unit='characters'); + async_result = asyncio.run(Trimmer().atrim_context(sources, 11, unit='characters')); assert isinstance(sync_result, ContextTrimResult); assert isinstance(sync_result.sources[0], ContextSourceResult); - assert async_result == sync_result" + assert async_result == sync_result; + assert sync_result.text == 'Source: A\na'" diff --git a/AGENTS.md b/AGENTS.md index c3ed4c6..f541a5a 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -134,11 +134,11 @@ Keep each responsibility in its existing file. Generated folders such as `dist/` ### Package source - `src/trimwise/__init__.py`: exposes the supported public API names and nothing else. -- `src/trimwise/composition.py`: reconstructs exact source outputs, local spans, affordable omission - markers, and tiny-budget fallbacks. +- `src/trimwise/composition.py`: reconstructs exact source outputs, local spans, budgeted context + wrappers, affordable omission markers, and tiny-budget fallbacks. - `src/trimwise/measurement.py`: measures token, word, and character budgets and finds fitting prefixes. -- `src/trimwise/models.py`: defines public enums, configuration, result values, and semantic +- `src/trimwise/models.py`: defines public enums, inputs, configuration, result values, and semantic backend errors. - `src/trimwise/ranking.py`: builds scoring-only section and neighbor context, then implements structural, BM25, semantic, hybrid, signal, cosine, and MMR ranking calculations. @@ -159,6 +159,8 @@ Keep each responsibility in its existing file. Generated folders such as `dist/` - `tests/test_async_semantic.py`: checks FastEmbed and caller callback precedence, vector validation, staged failures, model reuse, concurrency, async equivalence, and cancellation behavior. - `tests/test_context.py`: checks shared-budget validation, selection, composition, counts, and spans. +- `tests/test_context_rendering.py`: checks budgeted source prefixes, separators, rendered counts, + wrapper isolation, and sync/async parity. - `tests/test_context_semantic.py`: checks context semantic batching, deduplication, and async use. - `tests/test_docstrings.py`: enforces Python docstrings and verifies that `py.typed` is packaged. - `tests/test_ranking.py`: checks BM25, centrality, semantic and hybrid fusion, signal scoring, diff --git a/README.md b/README.md index 2bf6899..f396c12 100644 --- a/README.md +++ b/README.md @@ -77,23 +77,34 @@ print(result.spans) # Original-input Python-string offsets ### Many sources, one shared limit -Use `trim_context()` when passages from several sources should compete for one evidence budget: +Use `trim_context()` when passages from several sources should compete for one budget. Add +`ContextSource` prefixes when the final rendered labels must fit inside that same limit: ```python +from trimwise import ContextSource, Trimmer + result = Trimmer().trim_context( - [record["text"] for record in records], + [ + ContextSource( + text=record["text"], + prefix=f"Source: {record['title']}\nURL: {record['url']}\n", + ) + for record in records + ], limit=800, query="Which recommendations are supported by the reports?", + separator="\n\n", ) -for source in result.sources: - print(records[source.source_index]["url"], source.text) +prompt_ready_context = result.text +assert result.output_count <= result.limit ``` -The result keeps one row per input source, including empty excerpts, and the sum of its source -output counts stays within `limit`. Labels, URLs, caller-added headings, separators, instructions, -and answer space are outside that limit. See [Many Sources, One Shared Limit](https://trimwise.readthedocs.io/en/latest/multi-source-context/) -for the complete contract and the difference from `atrim_many()`. +The result keeps one row per input source, including empty excerpts. Prefixes are emitted only for +sources that contribute evidence, and `result.text` contains the fully measured rendering. Your +surrounding instructions and answer space remain outside this limit. Plain string sources still use +the original evidence-only accounting. See [Many Sources, One Shared Limit](https://trimwise.readthedocs.io/en/latest/multi-source-context/) +for both modes and the difference from `atrim_many()`. Depending on the trimming strategy you want to use, find the corresponding starter code example - [auto](https://trimwise.readthedocs.io/en/latest/strategies/#auto-the-lightweight-default), [structural](https://trimwise.readthedocs.io/en/latest/strategies/#structural-cover-a-document-without-a-query), [lexical](https://trimwise.readthedocs.io/en/latest/strategies/#lexical-preserve-exact-query-evidence), diff --git a/docs/api-reference.md b/docs/api-reference.md index 5a67622..4e95848 100644 --- a/docs/api-reference.md +++ b/docs/api-reference.md @@ -26,6 +26,13 @@ duplicate contextual passages once per normalized-query callback batch. ::: trimwise.TrimInput +## Shared-context inputs + +Use [`ContextSource`][trimwise.ContextSource] when a source needs an output prefix that counts +toward the same shared limit as its retained evidence. + +::: trimwise.ContextSource + ## Configuration Use [`TrimConfig`][trimwise.TrimConfig] to configure token encoding, managed embeddings, MMR, @@ -41,7 +48,8 @@ measured counts, resolved strategy, and trimming status. ::: trimwise.TrimResult The context methods return a [`ContextTrimResult`][trimwise.ContextTrimResult] containing one -input-aligned [`ContextSourceResult`][trimwise.ContextSourceResult] per source. +input-aligned [`ContextSourceResult`][trimwise.ContextSourceResult] per source and, for rendered +calls, the complete prompt-ready text. ::: trimwise.ContextTrimResult diff --git a/docs/configuration-and-api.md b/docs/configuration-and-api.md index 10a1b50..abe618d 100644 --- a/docs/configuration-and-api.md +++ b/docs/configuration-and-api.md @@ -35,6 +35,7 @@ Import supported objects directly from `trimwise`: ```python from trimwise import ( BudgetUnit, + ContextSource, ContextSourceResult, ContextTrimResult, SemanticBackendError, @@ -55,6 +56,7 @@ These are the package's documented exports: | `TrimConfig` | Stores immutable reusable configuration | | `TrimInput` | Describes one independent request for asynchronous batch trimming | | `TrimResult` | Reports the excerpt, measurements, resolved strategy, and whether text changed | +| `ContextSource` | Pairs evidence with an optional output-only prefix | | `ContextSourceResult` | Reports one input-aligned source excerpt and its local spans | | `ContextTrimResult` | Reports all source excerpts and their shared aggregate measurements | | `SourceSpan` | Identifies one retained range in the original input string | @@ -255,11 +257,11 @@ a boolean. ## `trim_context()` and `atrim_context()` -Use the context methods when many source strings should share one output limit: +Use the context methods when many sources should share one output limit: ```text trim_context( - sources: Sequence[str], + sources: Sequence[str | ContextSource], limit: int, *, unit: BudgetUnit | str = BudgetUnit.TOKENS, @@ -267,19 +269,22 @@ trim_context( query: str | None = None, token_counter: Callable[[str], int] | None = None, deduplicate: bool = False, + separator: str | None = None, ) -> ContextTrimResult ``` `atrim_context()` accepts the same arguments and returns the same result type asynchronously. -`sources` must be a sequence of strings, not one bare string or an arbitrary iterable. The result -contains one `ContextSourceResult` per input position. A source may receive an empty excerpt, but -its row and `source_index` remain present. - -Source input and output strings are measured independently. Their counts sum to the aggregate -counts, and the aggregate output cannot exceed the shared limit. Caller-added labels, URLs, -instructions, and separators are not counted. See -[Many Sources, One Shared Limit](multi-source-context.md) for prompt assembly and token-counting -guidance. +`sources` must be a sequence of strings or `ContextSource` values, not one bare string or an +arbitrary iterable. `ContextSource(text, prefix="...")` attaches exact output text that is emitted +only if that source contributes evidence. The result contains one `ContextSourceResult` per input +position. A source may receive an empty excerpt, but its row and `source_index` remain present. + +Supplying any `ContextSource` or an explicit `separator` returns the complete rendering in +`result.text`. Its aggregate `output_count` measures prefixes, evidence, separators, and omission +text together and cannot exceed the shared limit. Prefixes and separators never affect ranking, +embedding input, or source spans. With plain strings and no separator, `result.text` remains `None` +and the original sum-of-row-counts behavior is preserved. See +[Many Sources, One Shared Limit](multi-source-context.md) for examples and exact counting rules. `deduplicate=True` is a best-effort embedding option. It sends each exact repeated contextual passage once during the operation and maps the vector back to every occurrence. It does not remove @@ -562,14 +567,16 @@ The context methods return a frozen, slotted `ContextTrimResult`: | --- | --- | | `sources` | Input-aligned tuple of `ContextSourceResult` values | | `input_count` | Sum of all independently measured source inputs | -| `output_count` | Sum of all independently measured source outputs | +| `output_count` | Complete rendered size, or the legacy sum of source outputs | | `limit` and `unit` | Shared output ceiling and measurement rule | | `strategy` | Concrete strategy after resolving `auto` | | `trimmed` | Whether any source output differs from its input | +| `text` | Complete rendered context, or `None` for plain-string calls without a separator | Each source result contains `source_index`, `text`, `input_count`, `output_count`, `trimmed`, and local `spans`. Empty or wholly omitted sources keep their row with empty text, zero output count, -and no spans. Keep caller metadata outside the result and reconnect it with `source_index`. +and no spans. Prefixes and separators have no spans. Keep non-rendered caller metadata outside the +result and reconnect it with `source_index`. ## `Strategy` diff --git a/docs/getting-started.md b/docs/getting-started.md index 03aec38..16d5465 100644 --- a/docs/getting-started.md +++ b/docs/getting-started.md @@ -228,22 +228,28 @@ That example gives every source its own 120-token limit. When the sources should for one allowance, call `trim_context()`: ```python +from trimwise import ContextSource + shared = trimmer.trim_context( - [source["text"] for source in sources], + [ + ContextSource( + text=source["text"], + prefix=f"## {source['label']}\n\n", + ) + for source in sources + ], limit=300, query=task, + separator="\n\n", ) -evidence = [] -for row in shared.sources: - if row.text: - label = sources[row.source_index]["label"] - evidence.append(f"## {label}\n\n{row.text}") +evidence = shared.text ``` A more relevant source may use more room, and some source rows may be empty. The labels and -separators added above are not part of the 300-token limit. Read -[Many Sources, One Shared Limit](multi-source-context.md) for counts, spans, async use, and the +separators above are emitted only for contributing sources and are included in the 300-token limit. +Instructions around `evidence` and the model's answer still need separate room. Read [Many Sources, +One Shared Limit](multi-source-context.md) for evidence-only mode, counts, spans, async use, and the difference from `atrim_many()`. Keep instructions outside the source text passed to Trimwise. The library is designed to reduce @@ -415,8 +421,9 @@ not factual truth. - **Omitting the query:** lexical, semantic, and hybrid strategies require a nonblank query. - **Treating the limit as a target:** it is a ceiling. A query-aware result may stop early instead of adding weak evidence. -- **Budgeting only the excerpts:** leave space for labels, separators, instructions, examples, and - the model's answer. +- **Forgetting prompt overhead:** use `ContextSource` when per-source labels and separators must + share the evidence limit, and still reserve room for surrounding instructions, examples, tools, + and the model's answer. - **Passing instructions as evidence:** trim source material, then assemble it around instructions that remain unchanged. - **Expecting a summary:** Trimwise selects and joins original fragments; it does not paraphrase or diff --git a/docs/guarantees-and-limitations.md b/docs/guarantees-and-limitations.md index cbba860..c477128 100644 --- a/docs/guarantees-and-limitations.md +++ b/docs/guarantees-and-limitations.md @@ -28,14 +28,25 @@ headings, and affordable omission markers have been applied. It returns only whe result.output_count <= result.limit ``` -For `trim_context()` and `atrim_context()`, the ceiling applies to the sum of independently -measured source outputs: +With plain string sources and no explicit separator, context calls preserve their original +row-oriented accounting: ```text result.output_count = sum(source.output_count for source in result.sources) result.output_count <= result.limit ``` +Passing a `ContextSource` or an explicit separator instead applies the ceiling to the complete +prompt-ready rendering: + +```text +measure(result.text) = result.output_count +result.output_count <= result.limit +``` + +The complete measurement includes prefixes only for contributing sources and separators only +between them. Per-source rows and spans continue to describe the retained source evidence. + The guarantee uses the requested measurement rule: - Tiktoken or your custom counter for token budgets. @@ -52,13 +63,11 @@ for unit, limit in (("tokens", 30), ("words", 20), ("characters", 100)): assert result.output_count <= limit ``` -This guarantee applies to `result.text`, not to the larger prompt around it. Instructions, source -labels, separators, examples, tool definitions, output schemas, and the model's answer all need -their own context space. - -The same boundary applies to context results: caller-added labels and separators are not counted. -Separately measured token strings can also tokenize differently after they are joined. Measure the -completed prompt and reserve room when its whole-token ceiling must be exact. +For a single-source `TrimResult`, this guarantee applies to `result.text`, not to the larger prompt +around it. For a context result, prefixes and separators supplied through the context API are +included; formatting added afterward is not. Instructions, examples, tool definitions, output +schemas, and the model's answer still need their own context space. Measure the completed prompt and +reserve room when its whole-token ceiling must be exact. ### Input that already fits is returned exactly @@ -384,20 +393,20 @@ Trimwise deliberately leaves several decisions with the application: | Responsibility | What the caller should do | | --- | --- | -| Complete prompt budget | Reserve room for labels, instructions, examples, tools, and model output | -| Source identity | Store document IDs, URLs, authors, timestamps, and access controls outside Trimwise results; reconnect context rows with `source_index` | +| Complete prompt budget | Use `ContextSource` for budgeted per-source prefixes; reserve room for surrounding instructions, examples, tools, and model output | +| Source identity | Keep document IDs, authors, timestamps, and access controls in application data; render only the bounded labels you need and reconnect rows with `source_index` | | Relevance evaluation | Test downstream answers on representative documents and queries | | Factual verification | Check claims against original sources when accuracy matters | | Contradiction handling | Preserve and compare conflicting evidence explicitly | | Embedding operations | Own callback caching, retries, timeouts, rate limits, privacy, and concurrency | -| Security | Treat selected source text as untrusted and enforce tool permissions separately | +| Security | Treat source text and caller-supplied prefixes as untrusted and enforce tool permissions separately | | Tokenizer alignment | Supply a custom counter when the default encoding does not match the target model | `TrimResult.spans` exposes ordered Python-string ranges for retained source text, but it does not include source IDs, scores, or embeddings. Starts are inclusive, ends are exclusive, and adjacent ranges are merged. Overlapping ranges are never returned because candidates do not overlap. -Generated omission markers and separators have no span. Keep document identity alongside each -input before trimming several sources. +Generated omission markers and caller-supplied prefixes and separators have no span. Keep document +identity alongside each input before trimming several sources. For a context result, each source row's spans index only the original string at its `source_index`. diff --git a/docs/how-it-works.md b/docs/how-it-works.md index 43a72da..035c39f 100644 --- a/docs/how-it-works.md +++ b/docs/how-it-works.md @@ -64,9 +64,10 @@ If `limit == 0`, Trimwise returns empty text after measuring the input. If the i most the limit, it returns the original string exactly and stops before Markdown parsing, ranking, callback invocation, or FastEmbed loading. -For `trim_context()` and `atrim_context()`, the fitting check uses the sum of the independently -measured source inputs. An aggregate-fitting collection returns every source unchanged. If the sum -is too large, sources compete even when each one would fit by itself. +For plain string context calls, the fitting check uses the sum of independently measured source +inputs. When prefixes or a separator are supplied, it instead measures the complete rendered +context. A fitting collection returns every source unchanged; otherwise sources compete even when +each evidence string would fit by itself. ```python from trimwise import Trimmer @@ -359,18 +360,21 @@ Ranking order never becomes output order. For every proposed addition, Trimwise: 1. Sorts all selected candidates by their original source position. 2. Inserts exact whitespace for untouched gaps or minimal separators for omitted nonblank gaps. 3. Composes the retained fragments without optional omission markers. -4. Measures that complete proposal. -5. Rejects the addition if retained content and required separators exceed the limit. -6. Otherwise, tries optional omission markers one gap at a time in source-gap order. -7. Carries the retained original-input ranges into the final result. +4. For a rendered context, adds prefixes to contributing rows and joins them with the caller's + separator. +5. Measures that complete proposal. +6. Rejects the addition if it exceeds the limit. +7. Otherwise, tries optional omission markers one gap at a time in source-gap order. +8. Carries the retained original-input ranges into the final result. This trial composition accounts for the actual cost of headings, separators, and markers instead of estimating candidate cost in isolation. -For context results, Trimwise performs that reconstruction separately for each source and sums the -complete output counts before accepting a change. Optional markers are tried in source input order -after the retained content fits. A wholly omitted source receives an empty row, not a standalone -marker. +For context results, Trimwise reconstructs each source separately. Plain string calls retain the +sum-of-row measurement. Rendered calls remeasure the one final string, so token or custom counters +see the real boundaries between prefix, evidence, and separator. Optional markers are tried in +source input order after the retained content fits. A wholly omitted source receives neither a +standalone marker nor a prefix. ### Leading, internal, and trailing gaps @@ -467,13 +471,13 @@ Span offsets follow Python string slicing: starts are inclusive and ends are exc source-backed ranges are merged; overlapping ranges cannot arise from the nonoverlapping source candidates. Generated omission markers and minimal separators have no source span. -The guarantee applies to `result.text`, not to the larger prompt assembled around it. Source labels, -instructions, examples, tool schemas, separators, and model output still need their own context -space. +The `TrimResult` guarantee applies to its `result.text`, not to the larger prompt assembled around +it. Instructions, examples, tool schemas, and model output still need their own context space. -`ContextTrimResult` applies the same ceiling to the sum of its source outputs. Its aggregate counts -equal the sums of the per-source counts, and every local span indexes only that source's original -string. Formatting added later by the caller is outside this measurement. +For `ContextTrimResult`, plain string calls without an explicit separator keep the sum-of-source +ceiling and return `text=None`. Calls with a `ContextSource` or separator return a fully measured +`text` containing the active prefixes and separators. Every local span still indexes only its +source's original evidence. Formatting added later by the caller is outside this measurement. ## What source fidelity means diff --git a/docs/index.md b/docs/index.md index 38fdc1f..b64eb1c 100644 --- a/docs/index.md +++ b/docs/index.md @@ -26,7 +26,9 @@ source 3 ──> Trimwise ──> high-signal excerpt 3 ┘ `trim_context()` and `atrim_context()` can give those sources one shared evidence limit. More relevant sources may use more of the available space, while every input position remains present -in the result. See [Many Sources, One Shared Limit](multi-source-context.md). +in the result. Use `ContextSource` when source labels and separators must count toward that same +limit and be returned as one prompt-ready string. See +[Many Sources, One Shared Limit](multi-source-context.md). Trimwise compacts text you already have. It does **not** search the web, retrieve documents, query an index, or replace a RAG system. diff --git a/docs/multi-source-context.md b/docs/multi-source-context.md index a1b1def..3de9846 100644 --- a/docs/multi-source-context.md +++ b/docs/multi-source-context.md @@ -1,6 +1,6 @@ --- title: Trim Many Sources with One Shared Limit -description: Let evidence from several sources compete for one Trimwise token, word, or character budget while keeping source identity and local spans. +description: Let evidence and caller-supplied source labels from several inputs compete for one measured token, word, or character budget. --- # Many Sources, One Shared Limit @@ -10,14 +10,15 @@ all of their passages together, so a source with stronger evidence can use more space. The result still contains one entry per input source, in the same order. This is useful after retrieval, search, or tool calls have already chosen the sources. Trimwise does -not retrieve documents or copy your labels and metadata; it reduces the source strings you supply. +not retrieve documents or define a metadata format. You can provide an exact output prefix for each +source when labels or URLs must count toward the same limit as the evidence. ## A runnable core-install example Lexical selection needs no embedding model or optional dependency: ```python -from trimwise import Trimmer +from trimwise import ContextSource, Trimmer question = "Which retry loop ignored backoff settings, and when did service recover?" records = [ @@ -38,25 +39,34 @@ records = [ ] result = Trimmer().trim_context( - [record["text"] for record in records], - limit=18, + [ + ContextSource( + text=record["text"], + prefix=f"Source: {record['url']}\n", + ) + for record in records + ], + limit=30, unit="words", strategy="lexical", query=question, + separator="\n\n", ) -for source in result.sources: - record = records[source.source_index] - print(record["url"]) - print(source.text or "(no excerpt fit)") +print(result.text) -assert result.output_count == sum(source.output_count for source in result.sources) +assert result.text is not None +assert len(result.text.split()) == result.output_count assert result.output_count <= result.limit ``` -`source_index` is the zero-based position of the original string. Use it to reconnect each excerpt -to caller-owned URLs, filenames, permissions, timestamps, or other metadata. Trimwise deliberately -does not copy or transform that information. +`ContextSource.text` is evidence. Its `prefix` is copied exactly before that source's returned +evidence, but only when the source contributes a nonempty excerpt. Prefixes do not influence which +evidence wins and are never included in source spans. The separator is copied only between +contributing sources. + +`source_index` is the zero-based input position. Use it to reconnect each result row to caller-owned +filenames, permissions, timestamps, or other metadata that does not belong in the rendered text. Some entries may contain `text=""`. This can happen when the shared limit is too small or other sources have stronger evidence. The empty row remains present so indexes never shift. @@ -72,24 +82,45 @@ sources have stronger evidence. The empty row remains present so indexes never s Use `atrim_many()` when every input has already been assigned its own allowance. Use the context methods when passages should compete for the same allowance. -## Counts and prompt assembly +## Choose what the limit covers + +### Evidence-only results -Each source string is measured independently with the selected unit and optional custom counter: +Passing only strings and omitting `separator` preserves the original row-oriented behavior: ```text result.input_count = sum(source.input_count for source in result.sources) result.output_count = sum(source.output_count for source in result.sources) result.output_count <= result.limit +result.text is None ``` -The shared limit covers only the strings in `source.text`. It does not cover labels, URLs, -caller-added headings, instructions, separators, examples, tool definitions, an output schema, or -the model's answer. -Reserve room for those parts before choosing the limit. +Use this mode when your application will assemble and measure the final prompt itself. Any labels +or separators added later are outside Trimwise's limit. + +### Prompt-ready context + +Passing at least one `ContextSource`, or supplying `separator` explicitly, enables complete +rendering. In this mode: + +```text +result.input_count = sum(source.input_count for source in result.sources) +measure(result.text) = result.output_count +result.output_count <= result.limit +``` + +`input_count` still measures evidence only. Each source row's `output_count` still measures only +that row's returned evidence and omission markers. The aggregate `output_count` measures the final +`result.text`, including emitted prefixes and separators, so it need not equal the sum of row +counts. A plain string can be mixed with `ContextSource`; it simply has no prefix. + +Instructions, examples, tool definitions, an output schema, a fixed prompt header, and the model's +answer are still outside this limit. Reserve room for those surrounding parts. -Tokenizers can also count separately measured strings differently after they are joined. If the -completed prompt needs an exact token ceiling, assemble it, measure it with the target model's -tokenizer, and leave a safety margin or trim again with room reserved for prompt formatting. +Trimwise remeasures the complete rendered string because token counts are not always additive at +text boundaries. If the entire prompt needs an exact token ceiling, use the target model's tokenizer +as `token_counter`, reserve room for everything outside `result.text`, and measure the completed +prompt as a final application-level check. ## Result fields @@ -99,10 +130,11 @@ tokenizer, and leave a safety margin or trim again with room reserved for prompt | --- | --- | | `sources` | One `ContextSourceResult` per input source, in input order | | `input_count` | Sum of independently measured source inputs | -| `output_count` | Sum of independently measured source outputs | +| `output_count` | Complete rendered size in rendering mode; otherwise the sum of source outputs | | `limit` and `unit` | Shared ceiling and its measurement rule | | `strategy` | Concrete strategy after resolving `auto` | | `trimmed` | Whether any source output differs from its input | +| `text` | Prompt-ready rendered context, or `None` for evidence-only string calls | Each `ContextSourceResult` contains `source_index`, `text`, its own counts, `trimmed`, and local `spans`. A span always indexes the corresponding original source: @@ -113,7 +145,7 @@ for source in result.sources: retained_ranges = [original[span.start : span.end] for span in source.spans] ``` -Generated separators and omission markers do not have spans. +Caller prefixes, caller separators, and Trimwise-generated omission text do not have spans. ## Async semantic use @@ -123,7 +155,7 @@ passages needed for the whole context operation: ```python from collections.abc import Sequence -from trimwise import Trimmer +from trimwise import ContextSource, Trimmer async def embed(query: str, passages: Sequence[str]) -> tuple[object, Sequence[object]]: @@ -132,11 +164,15 @@ async def embed(query: str, passages: Sequence[str]) -> tuple[object, Sequence[o result = await Trimmer(async_embedding_callback=embed).atrim_context( - source_texts, + [ + ContextSource(text, prefix=f"Source {index + 1}:\n") + for index, text in enumerate(source_texts) + ], limit=800, strategy="hybrid", query="Which recommendations are supported by the reports?", deduplicate=True, + separator="\n\n", ) ``` @@ -163,8 +199,10 @@ With a query, an oversized best-matching passage is shortened to fit instead of a weaker source that happens to fit whole. This fallback returns the shortened passage in its own source row and leaves the other source rows empty. -Trimwise also does not resolve contradictions, verify claims, rank source authority, or copy source -metadata. Preserve the originals and provenance whenever those responsibilities matter. +Trimwise treats prefixes and separators as opaque text. It does not validate or escape titles, +URLs, or other caller values, so applications must handle untrusted metadata safely. Trimwise also +does not resolve contradictions, verify claims, or rank source authority. Preserve the originals +and provenance whenever those responsibilities matter. ## Continue exploring diff --git a/docs/semantic-and-async.md b/docs/semantic-and-async.md index 036843b..1302fb9 100644 --- a/docs/semantic-and-async.md +++ b/docs/semantic-and-async.md @@ -372,7 +372,8 @@ vectors after the call. semantic or hybrid source passages because those passages compete for one limit. Their `deduplicate=True` option uses the same exact-string, first-seen behavior with synchronous callbacks, async callbacks, and Trimwise-managed FastEmbed. Source rows remain distinct even when -their passage strings share an embedding. See +their passage strings share an embedding. `ContextSource` prefixes and the caller separator are +output-only and never appear in callback passages. See [Many Sources, One Shared Limit](multi-source-context.md) for the result and budget contract. For CPU-only structural or lexical work, async calls can overlap at the worker-thread level. For diff --git a/pyproject.toml b/pyproject.toml index 12de4ef..6514ab2 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "flit_core.buildapi" [project] name = "trimwise" -version = "0.4.1" +version = "0.5.0" description = "High-signal text trimming for better LLM prompts." readme = "README.md" requires-python = ">=3.10" diff --git a/src/trimwise/__init__.py b/src/trimwise/__init__.py index 36e6a6c..755c278 100644 --- a/src/trimwise/__init__.py +++ b/src/trimwise/__init__.py @@ -2,6 +2,7 @@ from trimwise.models import ( BudgetUnit, + ContextSource, ContextSourceResult, ContextTrimResult, SemanticBackendError, @@ -15,6 +16,7 @@ __all__ = [ "BudgetUnit", + "ContextSource", "ContextSourceResult", "ContextTrimResult", "SemanticBackendError", diff --git a/src/trimwise/composition.py b/src/trimwise/composition.py index ff22454..f12092b 100644 --- a/src/trimwise/composition.py +++ b/src/trimwise/composition.py @@ -49,6 +49,15 @@ class _SourceContext: measurer: Measurer limit: int marker: str + output_prefix: str = "" + + +@dataclass(frozen=True, slots=True) +class _ContextRendering: + """Store exact prefixes and separator for opt-in aggregate rendering.""" + + prefixes: tuple[str, ...] + separator: str def _compose( @@ -86,7 +95,7 @@ def _compose( current_groups.append(current) outputs.append(_ComposedOutput("".join(current), _source_spans(source, segments))) - total_count = _output_count(context.measurer, outputs) + total_count = _context_output_count(context.measurer, context.rendering, outputs) if total_count > context.limit: return None for source_index, pieces in enumerate(piece_groups): @@ -98,15 +107,16 @@ def _compose( current[piece_index] = piece.marked candidate = "".join(current) previous = outputs[source_index] - candidate_count = ( - total_count - - _text_count(context.measurer, previous.text) - + _text_count(context.measurer, candidate) + outputs[source_index] = _ComposedOutput(candidate, previous.spans) + candidate_count = _context_output_count( + context.measurer, + context.rendering, + outputs, ) if candidate_count <= context.limit: - outputs[source_index] = _ComposedOutput(candidate, previous.spans) total_count = candidate_count else: + outputs[source_index] = previous current[piece_index] = fallback return tuple(outputs) @@ -124,6 +134,48 @@ def _output_count(measurer: Measurer, outputs: Iterable[_ComposedOutput]) -> int return sum(_text_count(measurer, output.text) for output in outputs) +def _context_output_count( + measurer: Measurer, + rendering: _ContextRendering | None, + outputs: Iterable[_ComposedOutput], +) -> int: + """Measure legacy source rows or one complete rendered context. + + Args: + measurer: Shared output measurer. + rendering: Optional caller wrapper settings. + outputs: Input-aligned source evidence outputs. + + Returns: + Legacy independent sum or exact complete rendered size. + """ + output_tuple = tuple(outputs) + if rendering is None: + return _output_count(measurer, output_tuple) + return measurer.count(_render_context(rendering, output_tuple)) + + +def _render_context( + rendering: _ContextRendering, + outputs: tuple[_ComposedOutput, ...], +) -> str: + """Join contributing source evidence with its caller-owned wrappers. + + Args: + rendering: Input-aligned prefixes and inter-source separator. + outputs: Input-aligned source evidence outputs. + + Returns: + Complete prompt-ready context without wrappers for empty rows. + """ + contributions = [ + prefix + output.text + for prefix, output in zip(rendering.prefixes, outputs, strict=True) + if output.text + ] + return rendering.separator.join(contributions) + + def _text_count(measurer: Measurer, text: str) -> int: """Measure nonempty output while keeping empty result rows at zero. @@ -305,13 +357,19 @@ def _fallback_output(context: _SelectionContext) -> tuple[_ComposedOutput, ...]: Returns: Input-aligned outputs with at most one source-derived fragment. """ - if not context.segments: - return tuple(_ComposedOutput("", ()) for _ in context.sources) - index = max( + empty = tuple(_ComposedOutput("", ()) for _ in context.sources) + indexes = sorted( range(len(context.segments)), key=lambda candidate: (context.ranking.relevance[candidate], -candidate), + reverse=True, ) - return _fallback_candidate_output(context, index) + if context.rendering is None: + return _fallback_candidate_output(context, indexes[0]) if indexes else empty + for index in indexes: + output = _fallback_candidate_output(context, index) + if any(source.text for source in output): + return output + return empty def _fallback_candidate_output( @@ -335,6 +393,7 @@ def _fallback_candidate_output( context.measurer, context.limit, context.marker, + context.rendering.prefixes[source_index] if context.rendering is not None else "", ) fragment = _fitting_segment(source_context, segment) if not fragment.text: @@ -362,13 +421,13 @@ def _fitting_segment(context: _SourceContext, segment: Segment) -> _ComposedOutp opening = lines[0] closing = lines[-1] shell = opening + closing - if context.measurer.count(shell) > context.limit: + if context.measurer.count(context.output_prefix + shell) > context.limit: return _fitting_segment_prefix(context, segment) body = "".join(lines[1:-1]) endpoints = _line_endpoints(body) for end in reversed(endpoints): candidate = opening + body[:end] + closing - if context.measurer.count(candidate) <= context.limit: + if context.measurer.count(context.output_prefix + candidate) <= context.limit: prefix_end = segment.start + len(opening) + end spans = ( SourceSpan(segment.start, prefix_end), @@ -415,7 +474,11 @@ def _fitting_plain_prefix(context: _SourceContext, text: str) -> str: if prefix: return prefix prefix = _fitting_boundary_prefix(context, text, _complete_unit_endpoints(text)) - return prefix or context.measurer.fitting_prefix(text, context.limit) + return prefix or context.measurer.fitting_prefixed_content( + context.output_prefix, + text, + context.limit, + ) def _fitting_boundary_prefix( @@ -434,8 +497,12 @@ def _fitting_boundary_prefix( Longest fitting boundary prefix, or an empty string when none fits. """ for end in sorted(set(endpoints), reverse=True): - if 0 < end < len(text) and context.measurer.count(text[:end]) <= context.limit: - return text[:end] + candidate = text[:end] + if ( + 0 < end < len(text) + and context.measurer.count(context.output_prefix + candidate) <= context.limit + ): + return candidate return "" @@ -510,13 +577,13 @@ def _add_fallback_markers( output = fragment.text if context.source[: segment.start].strip(): candidate = context.marker + _newlines_before(output) + output - if context.measurer.count(candidate) <= context.limit: + if context.measurer.count(context.output_prefix + candidate) <= context.limit: output = candidate has_trailing_omission = fragment.text != segment.text or bool( context.source[segment.end :].strip() ) if has_trailing_omission: candidate = output + _newlines_after(output) + context.marker - if context.measurer.count(candidate) <= context.limit: + if context.measurer.count(context.output_prefix + candidate) <= context.limit: output = candidate return _ComposedOutput(output, fragment.spans) diff --git a/src/trimwise/measurement.py b/src/trimwise/measurement.py index d80519b..ae4df60 100644 --- a/src/trimwise/measurement.py +++ b/src/trimwise/measurement.py @@ -102,6 +102,23 @@ def fitting_prefix(self, text: str, limit: int) -> str: return text return self._fitting_scanned_prefix(text, limit) + def fitting_prefixed_content(self, prefix: str, text: str, limit: int) -> str: + """Fit source content after retaining an output prefix in full. + + Args: + prefix: Caller text that must precede any returned source content. + text: Source content eligible for prefix fallback. + limit: Maximum measured size of the complete prefixed output. + + Returns: + Longest fitting source prefix after the complete caller prefix. + """ + combined = prefix + text + fitting = self.fitting_prefix(combined, limit) + if len(fitting) < len(prefix): + return "" + return fitting[len(prefix) :] + def _fitting_encoded_prefix(self, text: str, limit: int) -> str: """Fit a source prefix around the configured encoding's token boundary. diff --git a/src/trimwise/models.py b/src/trimwise/models.py index 5503d07..6c1a69c 100644 --- a/src/trimwise/models.py +++ b/src/trimwise/models.py @@ -127,6 +127,19 @@ class TrimInput: token_counter: Callable[[str], int] | None = None +@dataclass(frozen=True, slots=True) +class ContextSource: + """Pair source evidence with an output-only prefix. + + Attributes: + text: Evidence that may be segmented, ranked, and returned with source spans. + prefix: Opaque caller text emitted only when this source contributes evidence. + """ + + text: str + prefix: str = "" + + @dataclass(frozen=True, slots=True) class ContextSourceResult: """Describe one input-aligned excerpt from a shared context budget. @@ -155,11 +168,12 @@ class ContextTrimResult: Attributes: sources: One result for every input source, in input order. input_count: Sum of the independently measured input sizes. - output_count: Sum of the independently measured output sizes. + output_count: Measured rendered size, or the legacy sum of source outputs. limit: Shared maximum output size. unit: Unit used for all counts. strategy: Concrete strategy used after resolving ``auto``. trimmed: Whether any source output differs from its original source. + text: Complete rendered context, or ``None`` for legacy row-only calls. """ sources: tuple[ContextSourceResult, ...] @@ -169,6 +183,7 @@ class ContextTrimResult: unit: BudgetUnit strategy: Strategy trimmed: bool + text: str | None = None @dataclass(frozen=True, slots=True) diff --git a/src/trimwise/selection.py b/src/trimwise/selection.py index b35c5d7..6f64603 100644 --- a/src/trimwise/selection.py +++ b/src/trimwise/selection.py @@ -11,7 +11,9 @@ _complete_unit_endpoints, _compose, _ComposedOutput, - _output_count, + _context_output_count, + _ContextRendering, + _fallback_candidate_output, ) from trimwise.measurement import Measurer from trimwise.models import Strategy @@ -31,6 +33,7 @@ class _SelectionContext: limit: int marker: str mmr_lambda: float + rendering: _ContextRendering | None = None @dataclass(slots=True) @@ -189,14 +192,16 @@ def _select_query_aware( return state.output if state.selected else None -def _oversized_query_fallback_index(context: _SelectionContext) -> int | None: - """Find a strongest oversized candidate that should precede weaker sources. +def _oversized_query_fallback( + context: _SelectionContext, +) -> tuple[_ComposedOutput, ...] | None: + """Build an affordable strongest-candidate fallback before weaker sources. Args: context: Multi-source query-aware selection inputs. Returns: - Strongest eligible candidate when it needs fallback, otherwise ``None``. + Affordable strongest-candidate fallback, otherwise ``None``. """ if len(set(context.source_indexes)) < 2: return None @@ -208,7 +213,12 @@ def _oversized_query_fallback_index(context: _SelectionContext) -> int | None: context.ranking.new_maximum_similarities(), context.mmr_lambda, ) - return index if _compose(context, {index}) is None else None + if _compose(context, {index}) is not None: + return None + output = _fallback_candidate_output(context, index) + if context.rendering is None or any(source.text for source in output): + return output + return None def _query_aware_indexes(context: _SelectionContext) -> set[int]: @@ -283,7 +293,11 @@ def _fill_section_shares(state: _SelectionState) -> None: sections = sorted({segment.section for segment in state.context.segments}) if not sections: return - available = state.context.limit - _output_count(state.context.measurer, state.output) + available = state.context.limit - _context_output_count( + state.context.measurer, + state.context.rendering, + state.output, + ) share = max(0, available // len(sections)) costs = { index: state.context.measurer.count(state.context.segments[index].text) diff --git a/src/trimwise/trimmer.py b/src/trimwise/trimmer.py index 4a36fed..d7d9906 100644 --- a/src/trimwise/trimmer.py +++ b/src/trimwise/trimmer.py @@ -9,13 +9,16 @@ from trimwise.composition import ( _ComposedOutput, - _fallback_candidate_output, + _context_output_count, + _ContextRendering, _fallback_output, + _render_context, _text_count, ) from trimwise.measurement import Measurer, TokenCounter from trimwise.models import ( BudgetUnit, + ContextSource, ContextSourceResult, ContextTrimResult, SourceSpan, @@ -35,7 +38,7 @@ from trimwise.segmentation import Segment, segment_text from trimwise.selection import ( _expand_structural_plaintext, - _oversized_query_fallback_index, + _oversized_query_fallback, _prepare_context_candidates, _select_query_aware, _select_structural, @@ -91,6 +94,7 @@ class _ContextRequest: query: str | None token_counter: TokenCounter | None deduplicate: bool + rendering: _ContextRendering | None @dataclass(frozen=True, slots=True) @@ -104,6 +108,7 @@ class _ContextArguments: query: str | None token_counter: TokenCounter | None deduplicate: bool + rendering: _ContextRendering | None @dataclass(frozen=True, slots=True) @@ -249,7 +254,7 @@ async def atrim( def trim_context( self, - sources: Sequence[str], + sources: Sequence[str | ContextSource], limit: int, *, unit: BudgetUnit | str = BudgetUnit.TOKENS, @@ -257,27 +262,30 @@ def trim_context( query: str | None = None, token_counter: Callable[[str], int] | None = None, deduplicate: bool = False, + separator: str | None = None, ) -> ContextTrimResult: """Trim many distinct sources under one shared output limit. Args: - sources: Source strings whose excerpts share the requested limit. - limit: Maximum summed output size in ``unit``. + sources: Source strings or prefixed context sources sharing the limit. + limit: Maximum evidence or rendered-context size in ``unit``. unit: Token, whitespace-word, or code-point character budget. strategy: Structural, lexical, semantic, hybrid, or automatic ranking. query: Task or question required by query-aware strategies. token_counter: Optional synchronous token measurement callback. deduplicate: Whether identical contextual passages share one embedding. + separator: Exact output text placed between contributing sources. Supplying + it enables complete rendering even when every source is a string. Returns: - Input-aligned excerpts and aggregate measurements. + Input-aligned excerpts, aggregate measurements, and optional rendered text. Raises: TypeError: If an argument is invalid or only an async embedder is available. ValueError: If an argument value or strategy/query combination is invalid. SemanticBackendError: If an explicitly requested semantic backend fails. """ - source_snapshot = _snapshot_sources(sources) + source_snapshot, rendering = _snapshot_sources(sources, separator) _validate_deduplicate(deduplicate) arguments = _ContextArguments( source_snapshot, @@ -287,12 +295,13 @@ def trim_context( query, token_counter, deduplicate, + rendering, ) return self._trim_context(arguments) async def atrim_context( self, - sources: Sequence[str], + sources: Sequence[str | ContextSource], limit: int, *, unit: BudgetUnit | str = BudgetUnit.TOKENS, @@ -300,6 +309,7 @@ async def atrim_context( query: str | None = None, token_counter: Callable[[str], int] | None = None, deduplicate: bool = False, + separator: str | None = None, ) -> ContextTrimResult: """Trim many sources asynchronously under one shared output limit. @@ -307,23 +317,25 @@ async def atrim_context( Cancellation propagates to that callback, but cannot stop worker work already running. Args: - sources: Source strings whose excerpts share the requested limit. - limit: Maximum summed output size in ``unit``. + sources: Source strings or prefixed context sources sharing the limit. + limit: Maximum evidence or rendered-context size in ``unit``. unit: Token, whitespace-word, or code-point character budget. strategy: Structural, lexical, semantic, hybrid, or automatic ranking. query: Task or question required by query-aware strategies. token_counter: Optional synchronous token measurement callback. deduplicate: Whether identical contextual passages share one embedding. + separator: Exact output text placed between contributing sources. Supplying + it enables complete rendering even when every source is a string. Returns: - Input-aligned excerpts and aggregate measurements. + Input-aligned excerpts, aggregate measurements, and optional rendered text. Raises: TypeError: If an argument has an unsupported type. ValueError: If an argument value or strategy/query combination is invalid. SemanticBackendError: If an explicitly requested semantic backend fails. """ - source_snapshot = _snapshot_sources(sources) + source_snapshot, rendering = _snapshot_sources(sources, separator) _validate_deduplicate(deduplicate) arguments = _ContextArguments( source_snapshot, @@ -333,6 +345,7 @@ async def atrim_context( query, token_counter, deduplicate, + rendering, ) callback = self._async_embedding_callback if callback is None: @@ -562,6 +575,7 @@ def _prepare_context( normalized_query, arguments.token_counter, arguments.deduplicate, + arguments.rendering, ) _validate_context_request(request) measurer = Measurer( @@ -574,11 +588,11 @@ def _prepare_context( empty_outputs = tuple(_ComposedOutput("", ()) for _ in arguments.sources) if arguments.limit == 0: return _context_result(prepared, empty_outputs) - if sum(input_counts) <= arguments.limit: - outputs = tuple( - _ComposedOutput(source, (SourceSpan(0, len(source)),) if source else ()) - for source in arguments.sources - ) + outputs = tuple( + _ComposedOutput(source, (SourceSpan(0, len(source)),) if source else ()) + for source in arguments.sources + ) + if _context_output_count(measurer, request.rendering, outputs) <= arguments.limit: return _context_result(prepared, outputs) segments, source_indexes = _prepare_context_candidates( @@ -781,19 +795,16 @@ def _select_context( request.limit, self.config.omission_marker, self.config.mmr_lambda, + request.rendering, ) - fallback_index = None if request.strategy is Strategy.STRUCTURAL: outputs = _select_structural(context) else: - fallback_index = _oversized_query_fallback_index(context) - outputs = None if fallback_index is not None else _select_query_aware(context) + outputs = _oversized_query_fallback(context) + if outputs is None: + outputs = _select_query_aware(context) if outputs is None: - outputs = ( - _fallback_output(context) - if fallback_index is None - else _fallback_candidate_output(context, fallback_index) - ) + outputs = _fallback_output(context) return _context_result(prepared, outputs) @@ -865,24 +876,44 @@ def _batch_arguments(inputs: Sequence[TrimInput]) -> list[_TrimArguments]: return arguments -def _snapshot_sources(sources: Sequence[str]) -> tuple[str, ...]: +def _snapshot_sources( + sources: Sequence[str | ContextSource], + separator: str | None, +) -> tuple[tuple[str, ...], _ContextRendering | None]: """Validate and snapshot the explicit multi-source collection contract. Args: sources: Public source collection. + separator: Optional exact text between contributing source outputs. Returns: - Stable input-order source tuple. + Stable evidence strings and optional wrapper-aware rendering settings. Raises: - TypeError: If the value is not a non-string sequence of strings. + TypeError: If a source, prefix, or separator has an unsupported type. """ if isinstance(sources, str) or not isinstance(sources, Sequence): - raise TypeError("sources must be a sequence of strings") + raise TypeError("sources must be a sequence of strings or ContextSource values") + if separator is not None and not isinstance(separator, str): + raise TypeError("separator must be a string or None") snapshot = tuple(sources) - if any(not isinstance(source, str) for source in snapshot): - raise TypeError("sources must contain only strings") - return snapshot + texts: list[str] = [] + prefixes: list[str] = [] + rendered = separator is not None + for source in snapshot: + if isinstance(source, str): + texts.append(source) + prefixes.append("") + continue + if not isinstance(source, ContextSource): + raise TypeError("sources must contain only strings or ContextSource values") + if not isinstance(source.text, str) or not isinstance(source.prefix, str): + raise TypeError("ContextSource text and prefix must be strings") + texts.append(source.text) + prefixes.append(source.prefix) + rendered = True + rendering = _ContextRendering(tuple(prefixes), separator or "") if rendered else None + return tuple(texts), rendering def _validate_deduplicate(deduplicate: bool) -> None: @@ -1061,7 +1092,7 @@ def _context_result( prepared: _PreparedContext, outputs: tuple[_ComposedOutput, ...], ) -> ContextTrimResult: - """Measure input-aligned outputs and enforce their summed hard limit. + """Measure input-aligned outputs and enforce their aggregate hard limit. Args: prepared: Validated request, source counts, and shared measurer. @@ -1087,7 +1118,11 @@ def _context_result( ) for source_index, output in enumerate(outputs) ) - output_count = sum(source.output_count for source in source_results) + output_count = _context_output_count( + prepared.measurer, + request.rendering, + outputs, + ) if output_count > request.limit: raise RuntimeError("internal composition exceeded the requested limit") return ContextTrimResult( @@ -1098,4 +1133,5 @@ def _context_result( request.unit, request.strategy, any(source.trimmed for source in source_results), + _render_context(request.rendering, outputs) if request.rendering is not None else None, ) diff --git a/tests/test_api.py b/tests/test_api.py index 563a170..03453e7 100644 --- a/tests/test_api.py +++ b/tests/test_api.py @@ -16,6 +16,7 @@ def test_public_exports_are_intentionally_small() -> None: """Expose only the documented public objects.""" assert trimwise.__all__ == [ "BudgetUnit", + "ContextSource", "ContextSourceResult", "ContextTrimResult", "SemanticBackendError", diff --git a/tests/test_context_rendering.py b/tests/test_context_rendering.py new file mode 100644 index 0000000..b316db3 --- /dev/null +++ b/tests/test_context_rendering.py @@ -0,0 +1,391 @@ +"""Verify budgeted wrappers and complete shared-context rendering.""" + +from __future__ import annotations + +from collections.abc import Sequence +from dataclasses import FrozenInstanceError +from typing import cast + +import numpy as np +import pytest +from numpy.typing import NDArray + +from trimwise import ContextSource, SourceSpan, TrimConfig, Trimmer +from trimwise.measurement import Measurer +from trimwise.models import BudgetUnit + + +def _target_vectors( + _: str, + passages: Sequence[str], +) -> tuple[NDArray[np.float32], list[NDArray[np.float32]]]: + """Align only passages containing target evidence with the query. + + Args: + _: Query text, which is fixed for these deterministic tests. + passages: Evidence-only strings submitted for embedding. + + Returns: + One query vector and one vector per passage. + """ + query = np.asarray([1.0, 0.0], dtype=np.float32) + vectors = [ + np.asarray( + [1.0, 0.0] if "target" in passage else [0.0, 1.0], + dtype=np.float32, + ) + for passage in passages + ] + return query, vectors + + +def test_context_source_is_frozen_and_slotted() -> None: + """Keep caller-owned evidence and its output prefix immutable.""" + source = ContextSource("evidence", prefix="Source: A\n") + assert not hasattr(source, "__dict__") + with pytest.raises(FrozenInstanceError): + source.prefix = "changed" # type: ignore[misc] + + +@pytest.mark.parametrize( + "source", + [ + ContextSource(cast(str, 1)), + ContextSource("evidence", prefix=cast(str, 1)), + ], +) +def test_context_source_fields_require_strings(source: ContextSource) -> None: + """Reject malformed wrapper values before measurement or selection. + + Args: + source: Context input containing one invalid field. + """ + with pytest.raises(TypeError, match="ContextSource"): + Trimmer().trim_context([source], 0) + + +def test_context_separator_requires_text() -> None: + """Reject a separator whose exact output cannot be defined.""" + with pytest.raises(TypeError, match="separator must be a string or None"): + Trimmer().trim_context([], 0, separator=cast(str, 1)) + + +def test_wrapped_fitting_sources_return_complete_rendered_context() -> None: + """Render unchanged evidence with prefixes only on nonempty source rows.""" + sources = [ + ContextSource("alpha", prefix="A: "), + ContextSource("", prefix="EMPTY: "), + ContextSource("beta", prefix="B: "), + ] + expected = "A: alpha\nB: beta" + result = Trimmer().trim_context( + sources, + len(expected), + unit="characters", + separator="\n", + ) + assert result.text == expected + assert result.output_count == len(expected) + assert [source.text for source in result.sources] == ["alpha", "", "beta"] + assert [source.output_count for source in result.sources] == [5, 0, 4] + assert [source.spans for source in result.sources] == [ + (SourceSpan(0, 5),), + (), + (SourceSpan(0, 4),), + ] + assert result.trimmed is False + + +def test_plain_sources_without_separator_keep_legacy_accounting() -> None: + """Preserve existing independently measured row semantics by default.""" + result = Trimmer().trim_context(["alpha", "beta"], 9, unit="characters") + assert result.text is None + assert result.output_count == sum(source.output_count for source in result.sources) + + +def test_explicit_separator_renders_plain_source_rows() -> None: + """Allow complete rendering without requiring per-source prefixes.""" + result = Trimmer().trim_context( + ["alpha", "beta"], + 13, + unit="characters", + separator="\n--\n", + ) + assert result.text == "alpha\n--\nbeta" + assert result.output_count == 13 + + +@pytest.mark.parametrize("unit", list(BudgetUnit)) +def test_complete_rendering_is_remeasured_in_each_builtin_unit(unit: BudgetUnit) -> None: + """Measure the final joined string rather than adding piece counts. + + Args: + unit: Built-in measurement rule used for the complete rendering. + """ + text = "A: alpha beta\nB: gamma" + result = Trimmer().trim_context( + [ContextSource("alpha beta", "A: "), ContextSource("gamma", "B: ")], + 100, + unit=unit, + separator="\n", + ) + measurer = Measurer(unit, "o200k_base", None) + assert result.text == text + assert result.output_count == measurer.count(text) + + +def test_custom_counter_measures_the_complete_rendering() -> None: + """Charge prefixes and separators through the caller's exact counter.""" + measured: list[str] = [] + + def count_characters(text: str) -> int: + """Record and count every string supplied by Trimwise. + + Args: + text: String being measured. + + Returns: + Number of Python code points. + """ + measured.append(text) + return len(text) + + result = Trimmer().trim_context( + [ContextSource("aaaa", "P:"), ContextSource("b", "Q:")], + 8, + token_counter=count_characters, + separator="|", + ) + assert result.text == "P:aaaa" + assert result.output_count == count_characters(result.text) + assert "P:aaaa|Q:b" in measured + + +def test_nonadditive_custom_counter_sees_the_join_boundary() -> None: + """Reject a second source when its rendered boundary exceeds the custom limit.""" + + def count_with_separator_surcharge(text: str) -> int: + """Make a joined rendering cost more than its independently measured parts. + + Args: + text: Complete string being measured. + + Returns: + Character count plus a visible join surcharge. + """ + return len(text) + (10 if "|" in text else 0) + + result = Trimmer().trim_context( + [ContextSource("a", "P:"), ContextSource("b", "Q:")], + 7, + token_counter=count_with_separator_surcharge, + separator="|", + ) + assert result.text == "P:a" + assert result.output_count == count_with_separator_surcharge(result.text) + + +def test_wrapper_text_never_enters_semantic_passages_or_spans() -> None: + """Keep provenance outside embedding input and source-backed ranges.""" + batches: list[list[str]] = [] + + def embed( + query: str, + passages: Sequence[str], + ) -> tuple[NDArray[np.float32], list[NDArray[np.float32]]]: + """Capture evidence passages and return deterministic vectors. + + Args: + query: Shared semantic query. + passages: Evidence-only ranking passages. + + Returns: + One query vector and one vector per passage. + """ + batches.append(list(passages)) + return _target_vectors(query, passages) + + result = Trimmer(embedding_callback=embed).trim_context( + [ + ContextSource("unrelated material", "target-label: "), + ContextSource("target evidence", "source-b: "), + ], + 25, + unit="characters", + strategy="semantic", + query="target", + separator="\n", + ) + assert all( + "target-label" not in passage and "source-b" not in passage for passage in batches[0] + ) + assert result.text is not None and result.text.startswith("source-b: ") + assert result.sources[1].spans == (SourceSpan(0, len("target evidence")),) + + +def test_prefix_text_never_influences_lexical_ranking() -> None: + """Rank evidence rather than a query term appearing only in a prefix.""" + result = Trimmer().trim_context( + [ + ContextSource("unrelated material", "target: "), + ContextSource("target evidence", "source: "), + ], + 23, + unit="characters", + strategy="lexical", + query="target", + separator="\n", + ) + assert result.text == "source: target evidence" + assert [source.text for source in result.sources] == ["", "target evidence"] + + +def test_wrapper_activation_cost_can_exclude_a_weaker_source() -> None: + """Spend a shared limit on stronger evidence when another prefix will not fit.""" + evidence = ["target answer.", "target filler."] + plain = Trimmer().trim_context( + evidence, + 29, + unit="characters", + strategy="lexical", + query="target", + separator="\n", + ) + wrapped = Trimmer().trim_context( + [ContextSource(evidence[0], "A: "), ContextSource(evidence[1], "B: ")], + 29, + unit="characters", + strategy="lexical", + query="target", + separator="\n", + ) + assert all(source.text for source in plain.sources) + assert wrapped.text == "A: target answer." + assert [source.text for source in wrapped.sources] == ["target answer.", ""] + + +def test_unaffordable_strong_prefix_does_not_block_a_fitting_source() -> None: + """Try weaker evidence when the strongest source cannot emit its whole prefix.""" + result = Trimmer().trim_context( + [ + ContextSource("target " * 20, "X" * 20), + ContextSource("other", "O:"), + ], + 10, + unit="characters", + strategy="lexical", + query="target", + separator="\n", + ) + assert result.text == "O:other" + assert [source.text for source in result.sources] == ["", "other"] + + +def test_fallback_reserves_room_for_the_contributing_source_prefix() -> None: + """Bound an oversized best match using the complete rendered limit.""" + relevant = "target " * 100 + prefix = "Source: A\n" + result = Trimmer().trim_context( + [ContextSource(relevant, prefix), ContextSource("other", "Source: B\n")], + 20, + unit="characters", + strategy="lexical", + query="target", + separator="\n\n", + ) + assert result.text == prefix + relevant[: 20 - len(prefix)] + assert result.output_count == 20 + assert result.sources[0].spans == (SourceSpan(0, 20 - len(prefix)),) + assert result.sources[1].text == "" + + +def test_balanced_fence_fallback_charges_its_prefix() -> None: + """Keep a complete code-fence shell after reserving its source prefix.""" + source = "```py\none\ntwo\n```\n" + result = Trimmer().trim_context( + [ContextSource(source, "S:\n")], + 13, + unit="characters", + ) + assert result.text == "S:\n```py\n```\n" + assert result.sources[0].text == "```py\n```\n" + assert result.sources[0].spans == (SourceSpan(0, 6), SourceSpan(14, 18)) + + +def test_omission_marker_is_kept_only_when_prefixed_output_fits() -> None: + """Remeasure optional omission text together with its contributing prefix.""" + source = ContextSource("first\n\nmiddle\n\nlast", "S:") + trimmer = Trimmer(TrimConfig(omission_marker="[...]")) + without_marker = trimmer.trim_context([source], 19, unit="characters") + with_marker = trimmer.trim_context([source], 20, unit="characters") + assert without_marker.text == "S:first\n\nlast" + assert "[...]" not in without_marker.text + assert with_marker.text == "S:first\n\n[...]\n\nlast" + assert with_marker.output_count == 20 + + +@pytest.mark.parametrize( + ("sources", "limit"), + [ + ([], 5), + ([ContextSource("", "unused")], 5), + ([ContextSource("evidence", "unused")], 0), + ([ContextSource("evidence", "12345")], 5), + ], +) +def test_empty_or_unaffordable_rendering_never_returns_a_bare_wrapper( + sources: list[ContextSource], + limit: int, +) -> None: + """Return an empty rendering unless source evidence can accompany its prefix. + + Args: + sources: Empty, blank, zero-budget, or prefix-exhausted inputs. + limit: Complete rendered-output limit. + """ + result = Trimmer().trim_context(sources, limit, unit="characters", separator="\n") + assert result.text == "" + assert result.output_count == 0 + assert all(not source.text for source in result.sources) + + +@pytest.mark.asyncio +async def test_sync_and_async_wrapped_semantic_results_match() -> None: + """Keep complete rendering identical across callback execution models.""" + + async def embed( + query: str, + passages: Sequence[str], + ) -> tuple[NDArray[np.float32], list[NDArray[np.float32]]]: + """Return the synchronous test vectors asynchronously. + + Args: + query: Shared semantic query. + passages: Evidence-only ranking passages. + + Returns: + One query vector and one vector per passage. + """ + return _target_vectors(query, passages) + + sources = [ + ContextSource("other fact", "A: "), + ContextSource("target fact", "B: "), + ] + synchronous = Trimmer(embedding_callback=_target_vectors).trim_context( + sources, + 18, + unit="characters", + strategy="semantic", + query="target", + separator="\n", + ) + asynchronous = await Trimmer(async_embedding_callback=embed).atrim_context( + sources, + 18, + unit="characters", + strategy="semantic", + query="target", + separator="\n", + ) + assert asynchronous == synchronous