Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 10 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -46,6 +46,16 @@ draining pending state again. If that view still cannot fit, the loop raises
locally and makes no SDK request. Providers without the capability retain the
existing request and dispatch behavior.

When `context.request_retention` advertises its optional `hard_fit` keyword,
that one provider-forced rebuild forwards `hard_fit=True`, allowing the context
to target the provider's requested budget directly. Older retention
capabilities, uninspectable dynamic callables, and the generic context fallback
keep their existing `provider`/`retain_contents`/`token_budget` assembly; the
second preflight and provider's final payload guard remain the safety boundary.
The existing `orchestrator:provider_budget` event exposes each preflight's
attempt, result, estimate, allowance, and requested context budget to mounted
observability consumers.

## Configuration

```toml
Expand Down
52 changes: 39 additions & 13 deletions amplifier_module_loop_streaming/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@

import asyncio
import fnmatch
import inspect
import json
import logging
import re
Expand Down Expand Up @@ -746,6 +747,7 @@ async def mount(coordinator: ModuleCoordinator, config: dict[str, Any] | None =
"orchestrator:steering_injected", # When a steer message is injected mid-turn
"orchestrator:goal_progress", # /goal auto-continue loop progress (see docs/designs/goal-command.md)
"orchestrator:budget_warning", # Layer 1 call budget at budget_warn_ratio (see _execute_stream)
"orchestrator:provider_budget", # Provider request-budget preflight result (see _execute_stream)
],
)

Expand Down Expand Up @@ -3318,26 +3320,47 @@ async def _execute_stream(
)
self._retention_capability_warned = True

def retention_accepts_hard_fit() -> bool:
"""Whether the optional retention capability supports ``hard_fit``.

The capability is module-owned and evolves independently of Core.
Signature inspection is deliberately separate from invocation:
an implementation ``TypeError`` must reach the caller rather than
being mistaken for an old capability and invoked a second time.
"""
if retaining_getter is None:
return False
try:
parameters = inspect.signature(retaining_getter).parameters.values()
except (TypeError, ValueError):
# Some dynamic callables cannot expose a signature. Keep their
# legacy behavior rather than claiming hard-fit support.
return False
return any(
(
parameter.name == "hard_fit"
and parameter.kind is not inspect.Parameter.POSITIONAL_ONLY
)
or parameter.kind is inspect.Parameter.VAR_KEYWORD
for parameter in parameters
)

async def request_messages(
retain_contents: list[str], *, token_budget: int | None = None
retain_contents: list[str],
*,
token_budget: int | None = None,
hard_fit: bool = False,
):
if retaining_getter is not None and (
self._ephemeral_injection_mode == "persist" or token_budget is not None
):
if retaining_getter is not None:
kwargs: dict[str, Any] = {
"provider": provider,
"retain_contents": retain_contents,
}
if token_budget is not None:
kwargs["token_budget"] = token_budget
try:
return await retaining_getter(**kwargs)
except TypeError as exc:
if token_budget is not None:
raise ContextLengthError(
"context.request_retention does not accept token_budget"
) from exc
raise
if hard_fit and retention_accepts_hard_fit():
kwargs["hard_fit"] = True
return await retaining_getter(**kwargs)
kwargs = {"provider": provider}
if token_budget is not None:
kwargs["token_budget"] = token_budget
Expand Down Expand Up @@ -4021,7 +4044,9 @@ async def exit_for_cancellation() -> None:
if smaller_context_budget is not None:
rebuilt_base_messages = list(
await request_messages(
retained_contents, token_budget=smaller_context_budget
retained_contents,
token_budget=smaller_context_budget,
hard_fit=True,
)
)
rebuilt_messages = (
Expand Down Expand Up @@ -4706,6 +4731,7 @@ async def exit_for_cancellation() -> None:
await request_messages(
final_retained_contents,
token_budget=smaller_context_budget,
hard_fit=True,
)
)
rebuilt_messages = (
Expand Down
262 changes: 262 additions & 0 deletions tests/test_provider_budget_context_simple_runtime.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,262 @@
"""Joint request-budget regression with real loop, context-simple, and OpenAI code.

This intentionally has no local path manipulation: repository-only runs skip when
the sibling modules are not installed. The root DTU installs all three local
checkouts, where this is required to execute without a skip.
"""

from __future__ import annotations

import json
from types import SimpleNamespace

import pytest
from amplifier_core import ContextLengthError
from amplifier_core.message_models import ChatRequest, Message

pytest.importorskip("amplifier_module_context_simple")
pytest.importorskip("amplifier_module_provider_openai")

from amplifier_module_context_simple import SimpleContextManager
from amplifier_module_loop_streaming import StreamingOrchestrator
from amplifier_module_provider_openai import OpenAIProvider


class _Cancellation:
is_cancelled = False
is_immediate = False
state = "running"


class _StableReminderHooks:
"""A real provider sees a stable persisted reminder on every assembled view."""

def __init__(self) -> None:
self.events: list[tuple[str, dict]] = []

async def emit(self, event: str, payload: dict | None = None):
self.events.append((event, payload or {}))
if event == "provider:request":
return SimpleNamespace(
action="inject_context",
ephemeral=True,
context_injection="<system-reminder>REQUIRED-REMINDER</system-reminder>",
context_injection_role="user",
append_to_last_tool_result=False,
data=None,
reason=None,
)
return SimpleNamespace(
action="continue",
ephemeral=False,
context_injection=None,
context_injection_role="system",
append_to_last_tool_result=False,
data=None,
reason=None,
)


class _Coordinator:
def __init__(self, hooks: _StableReminderHooks) -> None:
self.hooks = hooks
self.cancellation = _Cancellation()
self.session_state: dict = {}
self._capabilities: dict[str, object] = {}

def register_capability(self, name: str, capability: object) -> None:
self._capabilities[name] = capability

def get_capability(self, name: str):
return self._capabilities.get(name)

async def process_hook_result(self, result, *_args):
return result


class _RecordingContext(SimpleContextManager):
"""The actual context-simple algorithm, with only seam-call observation added."""

def __init__(self, **kwargs) -> None:
super().__init__(**kwargs)
self.hard_fit_calls: list[bool] = []

async def get_messages_for_request_retaining(
self,
*,
retain_contents: list[str],
provider=None,
token_budget: int | None = None,
hard_fit: bool = False,
) -> list[dict]:
self.hard_fit_calls.append(hard_fit)
return await super().get_messages_for_request_retaining(
retain_contents=retain_contents,
provider=provider,
token_budget=token_budget,
hard_fit=hard_fit,
)


class _InMemoryResponses:
"""Completed Responses SDK fake whose usage is derived from received params."""

def __init__(self, hard_fit_calls: list[bool]) -> None:
self.calls: list[dict] = []
self.hard_fit_counts_at_dispatch: list[int] = []
self._hard_fit_calls = hard_fit_calls

async def create(self, **params):
serialized = json.dumps(
params, ensure_ascii=False, sort_keys=True, separators=(",", ":"), default=str
)
# Record exactly the SDK payload, and derive usage from that same payload
# rather than a fixture outcome or a desired compaction size.
self.calls.append(json.loads(serialized))
self.hard_fit_counts_at_dispatch.append(self._hard_fit_calls.count(True))
input_tokens = len(serialized.encode("utf-8"))
return SimpleNamespace(
id=f"fake-{len(self.calls)}",
status="completed",
model=params["model"],
output=[
{
"type": "message",
"content": [{"type": "output_text", "text": "accepted"}],
}
],
usage=SimpleNamespace(input_tokens=input_tokens, output_tokens=1),
)


class _InMemoryClient:
def __init__(self, hard_fit_calls: list[bool]) -> None:
self.responses = _InMemoryResponses(hard_fit_calls)


def _payload_text(params: dict) -> str:
return json.dumps(params, ensure_ascii=False, sort_keys=True)


@pytest.mark.asyncio
async def test_hard_fit_stays_compacted_across_real_openai_dispatches() -> None:
hooks = _StableReminderHooks()
coordinator = _Coordinator(hooks)
context = _RecordingContext(
# The character-based ordinary estimate stays below context-simple's
# real provider-derived budget; max_tokens is only a fallback here.
# Each supplementary Han character serializes as four UTF-8 bytes, so
# the provider's payload preflight forces the first hard-fit rebuild.
max_tokens=500_000,
compact_threshold=0.99,
target_usage=0.50,
protected_recent=0.10,
protected_tool_results=1,
truncate_chars=64,
compaction_notice_enabled=True,
)
coordinator.register_capability(
"context.request_retention", context.get_messages_for_request_retaining
)
client = _InMemoryClient(context.hard_fit_calls)
provider = OpenAIProvider(
api_key="test-key",
client=client,
coordinator=coordinator,
config={
"default_model": "gpt-5-mini",
"max_output_tokens": 1024,
"max_retries": 0,
"use_streaming": False,
},
)
loop = StreamingOrchestrator({})

bulk = "REMOVED-BULK-MARKER:" + ("\U00020000" * 40_000)
await context.add_message({"role": "assistant", "content": bulk})

await loop.execute(
"ORIGINAL-HUMAN", context, {"openai": provider}, {}, hooks, coordinator
)
# New growth after the forced rebuild: a complete assistant/tool pair and
# a current tool result that context-simple's protection floor must keep
# complete.
await context.add_message(
{
"role": "assistant",
"content": "Calling the current tool.",
"tool_calls": [
{"id": "runtime-tool-1", "name": "current_tool", "arguments": {}}
],
}
)
await context.add_message(
{
"role": "tool",
"name": "current_tool",
"tool_call_id": "runtime-tool-1",
"content": "CURRENT-PROTECTED-TOOL-RESULT",
}
)
for prompt in (
"SECOND-HUMAN",
"POST-FORCE-ORDINARY-ONE",
"POST-FORCE-ORDINARY-TWO",
"POST-FORCE-ORDINARY-THREE",
):
await loop.execute(
prompt, context, {"openai": provider}, {}, hooks, coordinator
)

# Every call in this list reached the fake SDK. The first accepted call
# followed the provider-directed rebuild; the four subsequent calls prove
# ordinary fetches do not resurrect the canonical bulk history.
assert len(client.responses.calls) >= 5
payloads = [_payload_text(params) for params in client.responses.calls]
assert all("REMOVED-BULK-MARKER" not in payload for payload in payloads)
first_payload = payloads[0]
assert "ORIGINAL-HUMAN" in first_payload
assert "REQUIRED-REMINDER" in first_payload
assert context.hard_fit_calls[:2] == [False, True]
assert context.hard_fit_calls.count(True) == 1
assert client.responses.hard_fit_counts_at_dispatch == [1] * len(payloads)
post_force_payloads = payloads[1:]
assert post_force_payloads
assert all("ORIGINAL-HUMAN" in payload for payload in post_force_payloads)
assert "SECOND-HUMAN" in post_force_payloads[-1]
assert "REQUIRED-REMINDER" in post_force_payloads[-1]
assert "CURRENT-PROTECTED-TOOL-RESULT" in post_force_payloads[-1]
assert 'source=\\"context-compaction\\"' in post_force_payloads[-1]

final_input = client.responses.calls[-1]["input"]
tool_call_index = next(
index
for index, item in enumerate(final_input)
if item.get("type") == "function_call"
and item.get("call_id") == "runtime-tool-1"
)
assert final_input[tool_call_index + 1] == {
"type": "function_call_output",
"call_id": "runtime-tool-1",
"output": "CURRENT-PROTECTED-TOOL-RESULT",
}

canonical = await context.get_messages()
assert any(message.get("content") == bulk for message in canonical)
assert any(message.get("content") == "ORIGINAL-HUMAN" for message in canonical)
assert any(message.get("content") == "SECOND-HUMAN" for message in canonical)
assert any(
message.get("content") == "CURRENT-PROTECTED-TOOL-RESULT"
for message in canonical
)

# OpenAI's direct final assembled-payload guard remains the final boundary:
# an impossible protected payload never reaches the SDK fake.
accepted_before_guard = len(client.responses.calls)
with pytest.raises(ContextLengthError):
await provider.complete(
ChatRequest(
messages=[Message(role="user", content="IMPOSSIBLE" * 200_000)]
)
)
assert len(client.responses.calls) == accepted_before_guard
Loading
Loading