Skip to content

feat: add sampling-strategy retry/repair to stream() - #1697

Open
ajbozarth wants to merge 3 commits into
generative-computing:mainfrom
ajbozarth:feat/403-streaming-sampling-results
Open

ajbozarth wants to merge 3 commits into
generative-computing:mainfrom
ajbozarth:feat/403-streaming-sampling-results

Conversation

@ajbozarth

Copy link
Copy Markdown
Member

Issue

Fixes #403

Description

stream() previously aborted on a failed requirement with no recourse. This adds the retry/repair behavior of non-streaming sampling to streaming by accepting a SamplingStrategy: a failed attempt is repaired and retried up to loop_budget, and the winning attempt is projected onto the Streamer's terminal state (full_text, mot, final_validations). Requirements set on the strategy are merged with those passed to stream(). With strategy=None, behavior is unchanged.

Two triggers drive a retry:

  • a mid-stream "fail" cancels the attempt at the failing chunk and repairs via the new stream_repair();
  • a completed attempt whose final validate() fails repairs via the existing repair().

Per-attempt history is recorded on Streamer.attempts (a new StreamAttempt dataclass), and a RetryEvent is emitted before each re-attempt. CompletedEvent gains attempts_used.

stream_repair() is a non-abstract method on BaseSamplingStrategy that raises NotImplementedError by default (custom strategies stay instantiable; streaming repair is opt-in); RejectionSamplingStrategy, RepairTemplateStrategy, and MultiTurnStrategy implement it. repair()/stream_repair() past_results is widened to Sequence[ComputedModelOutputThunk | None] so streaming can pass its per-attempt results (None for early-broken attempts) without narrowing existing callers.

Testing

  • Tests added to the respective file if code was changed
  • New code has 100% coverage if code was added
  • Ensure existing tests and github automation passes (a maintainer will kick off the github automation when the rest of the PR is populated)

Attribution

  • AI coding assistants used

Adding a new component, requirement, sampling strategy, or tool?

If your PR adds or modifies one of the types below, check the matching box. A checklist of type-specific review items will be posted as a comment.

  • Component
  • Requirement
  • Sampling Strategy
  • Tool

NOTE: Please ensure you have an issue that has been acknowledged by a core contributor and routed you to open a pull request against this repository. Otherwise, please open an issue before continuing with this pull request.

stream() aborted on a failed requirement with no recourse. Accepting a
SamplingStrategy lets it repair and retry a failed attempt up to loop_budget,
mirroring non-streaming sampling. A mid-stream "fail" repairs via a new
stream_repair(); a failing final validate() repairs via repair(). Per-attempt
history is recorded on Streamer.attempts, and strategy=None is unchanged.

Assisted-by: Claude Code
Signed-off-by: Alex Bozarth <ajbozart@us.ibm.com>
Assisted-by: Claude Code
Signed-off-by: Alex Bozarth <ajbozart@us.ibm.com>
…requirements

Assisted-by: Claude Code
Signed-off-by: Alex Bozarth <ajbozart@us.ibm.com>
@ajbozarth
ajbozarth requested a review from a team as a code owner September 29, 2026 17:25
@github-actions github-actions Bot added the enhancement New feature or request label Sep 29, 2026
@ajbozarth ajbozarth added area/sampling SamplingStrategy, SamplingResult, ModelOption, generation options area/streaming Streaming chunks, events, per-chunk validation area/telemetry OTel spans, metrics, tracing, semconv labels Sep 29, 2026
@ajbozarth ajbozarth self-assigned this Sep 29, 2026

@planetf1 planetf1 left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The main ones are the requirement merge, abandoned streams counting as sampling failures, and the repair() signature change. The rest are smaller.

strategy_reqs = strategy.requirements if strategy is not None else None
# Union strategy and per-call requirements (deduped); copy each so streaming
# never mutates the caller's instances.
merged_reqs = list(dict.fromkeys([*(requirements or []), *(strategy_reqs or [])]))

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Was the union here deliberate? It's different from what the non-streaming path does with the same strategy.

_merge_requirements (mellea/core/sampling.py:385-403) uses the strategy's requirements instead of the per-call ones whenever the strategy has any, whereas stream() validates both. So with RejectionSamplingStrategy(requirements=[a]) and requirements=[b], instruct() checks a only and stream() checks a and b.

The docstrings disagree too. stream() says per-call requirements are "merged with any set on strategy", but BaseSamplingStrategy and SamplingStrategy both say strategy requirements "override per-call requirements". test_strategy_requirements_merge_with_call_requirements pins the union, so I'm guessing it was a choice.

Could you clarify which behaviour is intended when both are set, and have the docstrings on both sides say so? Otherwise anyone moving an instruct() call over to stream() is in for a surprise.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

We should keep the existing instruct / sampling strategy behavior.

if has_plugins(HookType.STREAMING_END):
from ..plugins.hooks.streaming import StreamingEndPayload

sampling_success = (

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

What should the sampling metrics record when a caller abandons a strategy stream? At the moment it counts as a failure.

On a break, an early aclose() or a cancellation, _finalize gets exception=None, so this produces sampling_success=False and record_streaming_outcome (mellea/telemetry/metrics_plugins.py:497-504) adds it to mellea.sampling.failures. I tried a loop_budget=3 stream whose output would have passed. Drained, it reports sampling_success=True. With a break after the first chunk it reports False, and its one attempt isn't selected and has an empty full_text.

The non-streaming path doesn't count this. A cancelled act() reaches SAMPLING_LOOP_END with exception set (mellea/core/sampling.py:239), and the metric handler skips it (metrics_plugins.py:443). The streaming handler's own docstring says "a crash is not a sampling verdict" (:490), and abandoning a stream looks like the same case to me.

One way round it would be to report sampling_success=None unless the loop actually reached a verdict. It'd also be worth saying in the StreamAttempt docstring (:232) what attempts looks like after an early exit, since "exactly one attempt in a run is selected" doesn't hold there.

new_ctx: Context,
past_actions: Sequence[SampleActionType],
past_results: list[ComputedModelOutputThunk],
past_results: Sequence[ComputedModelOutputThunk | None],

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Widening past_results here means custom strategies written against the current signature stop type-checking. repair() is the documented extension point, and 0.8.0 ships it as past_results: list[ComputedModelOutputThunk]. Parameter types can't be narrowed in an override, so a subclass with that signature passes mypy on main and fails on this branch:

Argument 4 of "repair" is incompatible with supertype "BaseSamplingStrategy";
supertype defines the argument type as "Sequence[ComputedModelOutputThunk[Any] | None]"  [override]

test/core/test_component_typing.py needed the same change in this PR. Anyone with their own strategy who runs mypy will hit it, whether or not they use stream().

It shows up at runtime too. The completed-fail branch passes every attempt (streaming.py:1000-1019), so for early-broken ones a custom repair() gets None in past_results and an empty list in past_val. The built-in strategies are fine because they only read [-1].

I think this is avoidable. If repair() only gets the completed attempts, the way _select_failure already handles select_from_failure (streaming.py:775-781), then [-1] doesn't change for any of the three built-in strategies (repair() only runs when the latest attempt completed), repair() keeps its 0.8.0 signature, and the None-bearing type stays on the new stream_repair(). The type: ignore[misc] at :779 could go as well. If you'd rather keep the wider type, it's worth a line in the release notes for people with custom strategies.

Separately, the past_results docstrings (base.py:174, :611, :677, :780) still say "List of (unsuccessful) generation results", with no mention of None.

and against the full output at stream end, merged with any set on
`strategy`; with neither, chunks stream without validation.
validation_backend: Backend for validation calls; defaults to `backend`.
strategy: Optional `BaseSamplingStrategy` enabling retry/repair. When set, a

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

stream() accepts any BaseSamplingStrategy, but only uses loop_budget, repair, stream_repair and select_from_failure, so a strategy that does its real work in _sample quietly runs as plain rejection sampling:

  • BudgetForcingSamplingStrategy and BaseMBRDSampling (majority voting, MBRD-ROUGE-L) override _sample but inherit RejectionSamplingStrategy.stream_repair, so the thinking budget and the voting are skipped
  • ModelFriendlyRepairStrategy overrides repair but inherits RepairTemplateStrategy.stream_repair, so a mid-stream failure gets the generic repair text rather than its ModelFriendlyFeedbackFormatter output

Related: the new docs say "pass a SamplingStrategy to stream()" (use-async-and-streaming.md:354, requirements-system.md:390, and twice in 06-streaming-validation.md), but the parameter is BaseSamplingStrategy. SOFAISamplingStrategy subclasses SamplingStrategy directly, so mypy rejects it here, and at runtime it fails with AttributeError: 'SOFAISamplingStrategy' object has no attribute 'stream_repair' at the first mid-stream failure.

The how-to page lists the supported strategies (use-async-and-streaming.md:395-397), but nothing at the call site tells you. A warning or TypeError for strategies that override _sample would catch it, or at least name the unsupported ones here next to the concurrency_budget note.

full_text: Validated-and-emitted output. On early exit or exception,
reflects whatever passed validation before the stop.
attempts_used: Number of stream attempts; currently always `1`.
attempts_used: Number of attempts run (1 without a strategy).

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

One gap for callers: there's no simple way to tell whether sampling actually succeeded. CompletedEvent.success only says the stream finished, so when the budget runs out and every attempt failed validation you still get success=True (that's what the test at test_streaming_sampling.py:257 checks). Right now you'd have to work it out from s.attempts.

_finalize already calculates this as sampling_success for the hook payload, so it'd be nice to expose it on Streamer/EventStreamer too. That would also cover the SamplingResult.success equivalent #403 asks for.

Raises:
NotImplementedError: If the strategy does not override this method.
"""
raise NotImplementedError(

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

A custom strategy that doesn't implement stream_repair only fails at the first mid-stream failure, after the caller has already received some chunks. It'd be friendlier to check at stream() entry (e.g. type(strategy).stream_repair is BaseSamplingStrategy.stream_repair) and raise straight away.

Comment thread mellea/stdlib/__init__.py
EventStreamer,
FullValidationEvent,
QuickCheckEvent,
StreamAttempt,

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

nit: RetryEvent is now emitted but not exported here, while StreamAttempt and the other event types are.


@property
def attempts(self) -> list[StreamAttempt]:
"""Per-attempt `StreamAttempt` records on the sampling path; else empty."""

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

nit: "else empty" isn't quite right. A plain stream records one attempt too, so this is only empty before the stream starts.

validation_backend: Backend,
streaming_id: str,
event_queue: asyncio.Queue[StreamEvent | None] | None = None,
event_queue: asyncio.Queue[StreamEvent | None] | None,

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

nit: event_queue lost its default and the new parameters are all required. Streamer is exported, so defaults would be the safer choice, even though nobody's likely to construct it directly.

Chunks from every attempt arrive through the same `async for`; a `RetryEvent`
separates one attempt from the next, and `streamer.attempts` holds the
per-attempt history after the run. If `loop_budget` is exhausted with no passing
attempt, `select_from_failure()` selects which attempt the `Streamer`'s final

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

nit: when every attempt breaks early, select_from_failure() isn't called at all. _select_failure only passes it completed attempts and otherwise falls back to the last one (streaming.py:775-785). Worth saying so here.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Should we change this behavior? It seems reasonable that select_from_failure should be able to select from any failed run? Should there be a two-tier system: one tier has early failures and one tier has things that only failed the last requirement?

@jakelorocco jakelorocco left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Agree with Nigel's comments as well.

Chunks from every attempt arrive through the same `async for`; a `RetryEvent`
separates one attempt from the next, and `streamer.attempts` holds the
per-attempt history after the run. If `loop_budget` is exhausted with no passing
attempt, `select_from_failure()` selects which attempt the `Streamer`'s final

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Should we change this behavior? It seems reasonable that select_from_failure should be able to select from any failed run? Should there be a two-tier system: one tier has early failures and one tier has things that only failed the last requirement?

strategy_reqs = strategy.requirements if strategy is not None else None
# Union strategy and per-call requirements (deduped); copy each so streaming
# never mutates the caller's instances.
merged_reqs = list(dict.fromkeys([*(requirements or []), *(strategy_reqs or [])]))

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

We should keep the existing instruct / sampling strategy behavior.

Comment on lines 352 to 357
class ModelFriendlyRepairStrategy(RepairTemplateStrategy):
"""RepairTemplateStrategy with model-friendly feedback formatting.

Extends RepairTemplateStrategy to use ModelFriendlyFeedbackFormatter for
converting validation failures into actionable repair guidance. This typically
improves LLM performance on repair tasks compared to generic validation reasons.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Should we also override .stream_repair for this one so that it uses the repair strategy for mid-stream repairs?

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

area/sampling SamplingStrategy, SamplingResult, ModelOption, generation options area/streaming Streaming chunks, events, per-chunk validation area/telemetry OTel spans, metrics, tracing, semconv enhancement New feature or request

Projects

None yet

Development

Successfully merging this pull request may close these issues.

feat: allow streaming for sampling results

3 participants