feat: preserve Responses reasoning state and compact context (#5172)

This commit is contained in:
chengyongru 2026-07-30 22:39:43 +08:00 committed by GitHub
parent 511c764f45
commit 6a1a45d07a
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
37 changed files with 4778 additions and 153 deletions

View File

@ -348,6 +348,20 @@ Valid `apiType` values are exactly `auto`, `chat_completions`, and `responses`.
</details>
<a id="responses-state-and-compaction"></a>
### Responses conversation state and compaction
Providers that use the Responses API can keep reasoning context across a
conversation, which helps with multi-step tasks. Supported providers can also
compact long conversations automatically.
nanobot preserves Responses conversation state automatically for OpenAI
Responses, OpenAI Codex, Azure OpenAI, and compatible GitHub Copilot models.
Native compaction is also automatic when the provider supports it. The
threshold is derived from the active model's context window and reserved output
headroom; no provider configuration is required.
<details>
<summary><b>Azure OpenAI</b></summary>

View File

@ -229,7 +229,7 @@ Arbitrary custom provider names are OpenAI-compatible only; they do not use the
}
```
`providers.openai.apiType` may be set when you need to force a specific OpenAI API surface. Other providers reject `apiType`; leave it unset outside `providers.openai`. Replace the model with a model ID available to your OpenAI account.
`providers.openai.apiType` may be set when you need to force a specific OpenAI API surface. Other providers reject `apiType`; leave it unset outside `providers.openai`. Replace the model with a model ID available to your OpenAI account. Direct OpenAI Responses, OpenAI Codex, Azure OpenAI Responses, and eligible GitHub Copilot models share [opaque Responses state retention](./configuration.md#responses-state-and-compaction); native compaction is enabled only where the backend supports it.
### Custom OpenAI-Compatible Endpoint
@ -458,7 +458,7 @@ For GitHub Copilot:
nanobot provider login github-copilot --set-main
```
Each command authenticates the selected provider and makes its current default model active. OAuth providers are not valid automatic fallbacks. See [`troubleshooting.md`](./troubleshooting.md#provider-and-model-problems) for proxy, headless-login, model-name, and config-key errors.
Each command authenticates the selected provider and makes its current default model active. OpenAI Codex and eligible GitHub Copilot models participate in [Responses state retention](./configuration.md#responses-state-and-compaction), while native compaction remains provider-capability-specific. OAuth providers are not valid automatic fallbacks. See [`troubleshooting.md`](./troubleshooting.md#provider-and-model-problems) for proxy, headless-login, model-name, and config-key errors.
## Provider Resolution

View File

@ -225,9 +225,6 @@ class ContextBuilder:
if current_role == "user"
else []
)
user_content = self.build_user_content(current_message, image_paths=media)
blocks = list(runtime_context_blocks or ()) if current_role == "user" else []
merged, runtime_context_meta = append_runtime_context(user_content, blocks)
messages: list[dict[str, Any]] = [
{
"role": "system",
@ -243,21 +240,47 @@ class ContextBuilder:
},
*history,
]
current = self.build_current_message(
current_message,
media=media,
current_role=current_role,
runtime_context_blocks=runtime_context_blocks,
)
if messages[-1].get("role") == current_role:
last = dict(messages[-1])
last["content"] = self._merge_message_content(last.get("content"), merged)
if current_role == "user" and runtime_context_meta is not None:
last["content"] = self._merge_message_content(
last.get("content"),
current.get("content"),
)
current_meta = current.get("_meta")
if current_role == "user" and isinstance(current_meta, dict):
internal_meta = dict(last.get("_meta") or {})
internal_meta[RUNTIME_CONTEXT_MESSAGE_META] = runtime_context_meta
internal_meta.update(cast(dict[str, Any], current_meta))
last["_meta"] = internal_meta
messages[-1] = last
return messages
current: dict[str, Any] = {"role": current_role, "content": merged}
if current_role == "user" and runtime_context_meta is not None:
current["_meta"] = {RUNTIME_CONTEXT_MESSAGE_META: runtime_context_meta}
messages.append(current)
return messages
def build_current_message(
self,
current_message: str,
*,
media: list[str] | None = None,
current_role: str = "user",
runtime_context_blocks: Sequence[RuntimeContextBlock] | None = None,
) -> dict[str, Any]:
"""Build only the fresh turn message without merging it into history."""
content = self.build_user_content(current_message, image_paths=media)
blocks = list(runtime_context_blocks or ()) if current_role == "user" else []
merged, runtime_context_meta = append_runtime_context(content, blocks)
current: dict[str, Any] = {"role": current_role, "content": merged}
if current_role == "user" and runtime_context_meta is not None:
current["_meta"] = {
RUNTIME_CONTEXT_MESSAGE_META: runtime_context_meta,
}
return current
def build_user_content(
self,
text: str,

View File

@ -49,7 +49,7 @@ from nanobot.bus.queue import MessageBus
from nanobot.bus.runtime_events import RuntimeEventBus
from nanobot.command import CommandContext, CommandRouter, register_builtin_commands
from nanobot.config.schema import AgentDefaults, ModelPresetConfig
from nanobot.providers.base import LLMProvider
from nanobot.providers.base import LLMProvider, ProviderConversationState
from nanobot.providers.factory import ProviderSnapshot
from nanobot.runtime_context import (
RUNTIME_CONTEXT_HISTORY_META,
@ -106,6 +106,7 @@ if TYPE_CHECKING:
from nanobot.triggers.local_store import LocalTriggerStore
_T = TypeVar("_T")
_SUBAGENT_PROVIDER_TASK_META = "subagent_provider_task_id"
class TurnKind(Enum):
@ -126,6 +127,7 @@ class TurnContext:
history: list[dict[str, Any]] = field(default_factory=list)
initial_messages: list[dict[str, Any]] = field(default_factory=list)
provider_state: ProviderConversationState | None = field(default=None, repr=False)
request_context: RequestContext | None = None
runtime_context_blocks: list[RuntimeContextBlock] = field(default_factory=list)
attributes: dict[str, Any] = field(default_factory=dict)
@ -243,6 +245,8 @@ class AgentLoop:
_RUNTIME_CHECKPOINT_KEY = "runtime_checkpoint"
_PENDING_USER_TURN_KEY = "pending_user_turn"
_PROVIDER_STATE_CHECKPOINT_VERSION_KEY = "provider_state_checkpoint_version"
_PROVIDER_STATE_CHECKPOINT_VERSION = "v1"
def __init__(
self,
@ -857,6 +861,7 @@ class AgentLoop:
turn_scopes: list[AbstractContextManager[Any]] | None = None,
tools: ToolRegistry | None = None,
request_context: RequestContext | None = None,
provider_state: ProviderConversationState | None = None,
) -> tuple[str | None, list[str], list[dict[str, Any]], str, bool]:
"""Run the agent iteration loop.
@ -872,7 +877,18 @@ class AgentLoop:
async def _checkpoint(payload: dict[str, Any]) -> None:
if session is None:
return
self._set_runtime_checkpoint(session, payload)
public_payload = dict(payload)
private_state = public_payload.pop("provider_state", None)
public_payload.pop(self._PROVIDER_STATE_CHECKPOINT_VERSION_KEY, None)
if "provider_state" in payload and (
private_state is None
or isinstance(private_state, ProviderConversationState)
):
session.provider_state = private_state
public_payload[self._PROVIDER_STATE_CHECKPOINT_VERSION_KEY] = (
self._PROVIDER_STATE_CHECKPOINT_VERSION
)
self._set_runtime_checkpoint(session, public_payload)
async def _drain_pending(*, limit: int = _MAX_INJECTIONS_PER_TURN) -> list[dict[str, Any]]:
"""Drain follow-up messages from the pending queue.
@ -1070,6 +1086,7 @@ class AgentLoop:
session_metadata=session_metadata,
message_metadata=metadata,
),
provider_state=provider_state,
))
finally:
turn_scope_stack.close()
@ -1077,6 +1094,8 @@ class AgentLoop:
reset_request_context(request_token)
reset_file_states(file_state_token)
self._last_usage = result.usage
if session is not None and not ephemeral:
session.provider_state = result.provider_state
if result.stop_reason == "max_iterations":
logger.warning("Max iterations ({}) reached", self.max_iterations)
should_stream = turn_continuation.should_stream_budget_response(
@ -1660,14 +1679,24 @@ class AgentLoop:
"extend_to_user": is_subagent,
}
ctx.history = session.get_history(**_hist_kwargs)
stored_state = session.provider_state
subagent_followup_persisted = False
if is_subagent:
# Keep the durable internal delivery as an assistant record, but
# present this completion to the model as fresh follow-up input.
# Providers without assistant-prefill support drop trailing
# assistant messages, so using the persisted record as the current
# prompt would hide an independently dispatched subagent result.
if self._persist_subagent_followup(session, ctx.msg):
subagent_followup_persisted = self._persist_subagent_followup(
session,
ctx.msg,
)
if subagent_followup_persisted:
logger.debug("Subagent result persisted for session {}", ctx.session_key)
# Establish a durable, replay-safe baseline before any fallible
# provider compatibility or prompt assembly work. A compatible
# staged state replaces this in a second atomic save below.
session.provider_state = None
self.sessions.save(session)
ctx.input_persisted_early = True
ctx.delivery.record_runtime(runtime)
@ -1675,13 +1704,65 @@ class AgentLoop:
ctx.request_context = self._request_context_for_turn(ctx)
if ctx.kind is TurnKind.USER:
ctx.runtime_context_blocks = await self._resolve_runtime_context_for_turn(ctx)
ctx.initial_messages = self._build_initial_messages(ctx)
staged_provider_state = False
if stored_state is not None and runtime.provider.can_resume_conversation_state(
stored_state,
runtime.model,
):
current_provider_message = self.context.build_current_message(
ctx.msg.content,
media=ctx.msg.media if ctx.kind is TurnKind.USER and ctx.msg.media else None,
runtime_context_blocks=ctx.runtime_context_blocks,
)
task_id = ctx.msg.metadata.get("subagent_task_id") if is_subagent else None
already_staged = False
if isinstance(task_id, str) and task_id:
internal_meta = current_provider_message.get("_meta")
current_provider_message["_meta"] = {
**(
cast(dict[str, Any], internal_meta)
if isinstance(internal_meta, dict)
else {}
),
_SUBAGENT_PROVIDER_TASK_META: task_id,
}
already_staged = any(
isinstance(message.get("_meta"), dict)
and cast(dict[str, Any], message["_meta"]).get(
_SUBAGENT_PROVIDER_TASK_META
)
== task_id
for message in stored_state.pending_messages
)
ctx.provider_state = (
stored_state
if already_staged
else stored_state.with_pending_messages([
*stored_state.pending_messages,
current_provider_message,
])
)
if (
not ctx.ephemeral
and (ctx.kind is TurnKind.USER or subagent_followup_persisted)
):
session.provider_state = ctx.provider_state
staged_provider_state = True
elif stored_state is not None:
session.provider_state = None
if ctx.kind is TurnKind.USER:
ctx.input_persisted_early = self._persist_user_message_early(
ctx.msg,
session,
runtime_context_blocks=ctx.runtime_context_blocks,
)
if staged_provider_state and not ctx.input_persisted_early:
session.provider_state = stored_state
elif subagent_followup_persisted and staged_provider_state:
# Upgrade the replay-safe baseline to the resumable state before
# prompt assembly and the first model checkpoint.
self.sessions.save(session)
ctx.initial_messages = self._build_initial_messages(ctx)
if ctx.on_progress is None:
ctx.on_progress = ctx.delivery.progress_callback()
@ -1715,6 +1796,7 @@ class AgentLoop:
turn_scopes=ctx.turn_scopes,
tools=ctx.tools,
request_context=ctx.request_context,
provider_state=ctx.provider_state,
)
final_content, _, all_msgs, stop_reason, had_injections = result
ctx.final_content = final_content
@ -2052,7 +2134,36 @@ class AgentLoop:
):
overlap = size
break
session.messages.extend(restored_messages[overlap:])
appended_messages = restored_messages[overlap:]
session.messages.extend(appended_messages)
assistant_message_data = (
cast(dict[str, Any], assistant_message)
if isinstance(assistant_message, dict)
else None
)
provider_state_is_synchronized = (
checkpoint_data.get(self._PROVIDER_STATE_CHECKPOINT_VERSION_KEY)
== self._PROVIDER_STATE_CHECKPOINT_VERSION
)
phase = checkpoint_data.get("phase")
exact_final_response = (
phase == "final_response"
and assistant_message_data is not None
and assistant_message_data.get("role") == "assistant"
and not bool(checkpoint_data.get("completed_tool_results"))
and not bool(checkpoint_data.get("pending_tool_calls"))
)
exact_completed_tools = (
phase == "tools_completed"
and assistant_message_data is not None
and assistant_message_data.get("role") == "assistant"
and not bool(checkpoint_data.get("pending_tool_calls"))
)
if not (
provider_state_is_synchronized
and (exact_final_response or exact_completed_tools)
):
session.provider_state = None
self._clear_pending_user_turn(session)
self._clear_runtime_checkpoint(session)
@ -2073,6 +2184,7 @@ class AgentLoop:
"timestamp": datetime.now().isoformat(),
}
)
session.provider_state = None
session.updated_at = datetime.now()
self._clear_pending_user_turn(session)

View File

@ -931,6 +931,7 @@ class Consolidator:
session_key=session.key,
)
session.last_consolidated = end_idx
session.provider_state = None
self.sessions.save(session)
return summary
@ -1136,6 +1137,7 @@ class Consolidator:
if summary:
last_summary = summary
session.last_consolidated = end_idx
session.provider_state = None
self.sessions.save(session)
if not summary:
# LLM is degraded — stop hammering it this call;
@ -1205,6 +1207,7 @@ class Consolidator:
# Preserve history and advance only the replay boundary.
session.last_consolidated = len(session.messages) - len(visible_suffix)
session.provider_state = None
self.sessions.save(session)
logger.info(

View File

@ -19,7 +19,17 @@ from nanobot.agent.context_governance import (
)
from nanobot.agent.hook import AgentHook, AgentHookContext, AgentRunHookContext
from nanobot.agent.tools.registry import ToolRegistry, is_tool_error_result
from nanobot.providers.base import LLMProvider, LLMResponse, ToolCallRequest
from nanobot.providers.base import (
LLMProvider,
LLMResponse,
ProviderCallContext,
ProviderConversationState,
ToolCallRequest,
)
from nanobot.providers.conversation_state import (
ProviderConversationStateController,
allows_conversation_message_merge,
)
from nanobot.runtime_context import (
RUNTIME_CONTEXT_MESSAGE_META,
detach_runtime_context,
@ -104,6 +114,7 @@ class AgentRunSpec:
goal_active_predicate: Callable[[], bool] | None = None
goal_continue_message: GoalContinueMessage | None = None
finalize_on_max_iterations: bool = True
provider_state: ProviderConversationState | None = None
@dataclass(slots=True)
@ -120,6 +131,7 @@ class AgentRunResult:
had_injections: bool = False
# Terminal tail to emit when the preceding final-content prefix was already streamed.
pending_stream_content: str | None = None
provider_state: ProviderConversationState | None = field(default=None, repr=False)
class AgentRunner:
@ -161,6 +173,7 @@ class AgentRunner:
and messages[-1].get("role") == "user"
and not is_hidden_history_message(injection)
and not is_hidden_history_message(messages[-1])
and allows_conversation_message_merge(messages[-1])
):
merged = dict(messages[-1])
left_meta = merged.get("_meta")
@ -231,6 +244,7 @@ class AgentRunner:
assistant_message: dict[str, Any] | None,
injection_cycles: int,
*,
conversation_state: ProviderConversationStateController | None = None,
phase: str = "after error",
iteration: int | None = None,
allow_goal_continue: bool = False,
@ -258,16 +272,21 @@ class AgentRunner:
if assistant_message is not None:
messages.append(assistant_message)
if iteration is not None:
checkpoint: dict[str, Any] = {
"phase": "final_response",
"iteration": iteration,
"model": spec.runtime.model,
"assistant_message": assistant_message,
"completed_tool_results": [],
"pending_tool_calls": [],
}
if conversation_state is not None:
checkpoint["provider_state"] = conversation_state.checkpoint(
messages
)
await self._emit_checkpoint(
spec,
{
"phase": "final_response",
"iteration": iteration,
"model": spec.runtime.model,
"assistant_message": assistant_message,
"completed_tool_results": [],
"pending_tool_calls": [],
},
checkpoint,
)
self._append_injected_messages(messages, injections)
if real_injection:
@ -420,6 +439,12 @@ class AgentRunner:
injection_cycles = 0
compacted_tool_call_ids: set[str] = set()
pending_stream_content: str | None = None
conversation_state = ProviderConversationStateController(
provider=spec.runtime.provider,
model=spec.runtime.model,
messages=messages,
state=spec.provider_state,
)
governance_config = ContextGovernanceConfig(
provider=spec.runtime.provider,
model=spec.runtime.model,
@ -450,7 +475,20 @@ class AgentRunner:
session_key=spec.session_key,
)
await hook.before_iteration(context)
response = await self._request_model(spec, messages_for_model, hook, context)
provider_context = conversation_state.prepare_request(
messages,
context_window_tokens=spec.runtime.context_window_tokens,
model_messages=messages_for_model,
)
response = await self._request_model(
spec,
messages_for_model,
hook,
context,
conversation_state=conversation_state,
provider_context=provider_context,
)
conversation_state.observe_response(response, messages)
context.response = response
context.tool_calls = list(response.tool_calls)
@ -480,6 +518,10 @@ class AgentRunner:
reasoning_content=response.reasoning_content,
thinking_blocks=response.thinking_blocks,
)
assistant_message = conversation_state.project_response_message(
assistant_message,
response,
)
messages.append(assistant_message)
await self._emit_checkpoint(
spec,
@ -544,6 +586,15 @@ class AgentRunner:
length_recovery_parts.clear()
continue
break
checkpoint_model_messages = (
self.context_governor.prepare_for_model(
governance_config,
messages,
compacted_tool_call_ids,
)
if response.provider_state is not None
else None
)
await self._emit_checkpoint(
spec,
{
@ -553,6 +604,10 @@ class AgentRunner:
"assistant_message": assistant_message,
"completed_tool_results": completed_tool_results,
"pending_tool_calls": [],
"provider_state": conversation_state.checkpoint(
messages,
model_messages=checkpoint_model_messages,
),
},
)
empty_content_retries = 0
@ -575,7 +630,11 @@ class AgentRunner:
)
clean = hook.finalize_content(context, response.content)
if response.finish_reason not in ("error", "length") and is_blank_text(clean):
if (
response.finish_reason
not in {"error", "length", "refusal", "content_filter"}
and is_blank_text(clean)
):
empty_content_retries += 1
if empty_content_retries < _MAX_EMPTY_RETRIES:
logger.warning(
@ -598,7 +657,12 @@ class AgentRunner:
if hook.wants_streaming():
await hook.on_stream_end(context, resuming=False)
retry_messages = self._finalization_retry_messages(messages_for_model)
response = await self._request_finalization_retry(spec, messages_for_model)
response = await self._request_finalization_retry(
spec,
messages_for_model,
transcript=messages,
conversation_state=conversation_state,
)
retry_usage = self._usage_or_estimate(spec, retry_messages, response)
self._accumulate_usage(usage, retry_usage)
raw_usage = self._merge_usage(raw_usage, retry_usage)
@ -623,10 +687,13 @@ class AgentRunner:
if hook.wants_streaming():
context.stream_continues_current_message = True
await hook.on_stream_end(context, resuming=True)
messages.append(build_assistant_message(
clean,
reasoning_content=response.reasoning_content,
thinking_blocks=response.thinking_blocks,
messages.append(conversation_state.project_response_message(
build_assistant_message(
clean,
reasoning_content=response.reasoning_content,
thinking_blocks=response.thinking_blocks,
),
response,
))
messages.append(build_length_recovery_message(clean or ""))
await hook.after_iteration(context)
@ -656,15 +723,22 @@ class AgentRunner:
reasoning_content=response.reasoning_content,
thinking_blocks=response.thinking_blocks,
)
assistant_message = conversation_state.project_response_message(
assistant_message,
response,
)
# Check for mid-turn injections BEFORE signaling stream end.
# If injections are found we keep the stream alive (resuming=True)
# so streaming channels don't prematurely finalize the card.
should_continue, injection_cycles = await self._try_drain_injections(
spec, messages, assistant_message, injection_cycles,
conversation_state=conversation_state,
phase="after final response",
iteration=iteration,
allow_goal_continue=True,
allow_goal_continue=(
response.finish_reason not in {"refusal", "content_filter"}
),
)
if should_continue:
had_injections = True
@ -717,11 +791,17 @@ class AgentRunner:
continue
break
messages.append(assistant_message or build_assistant_message(
clean,
reasoning_content=response.reasoning_content,
thinking_blocks=response.thinking_blocks,
))
messages.append(
assistant_message
or conversation_state.project_response_message(
build_assistant_message(
clean,
reasoning_content=response.reasoning_content,
thinking_blocks=response.thinking_blocks,
),
response,
)
)
await self._emit_checkpoint(
spec,
{
@ -731,6 +811,7 @@ class AgentRunner:
"assistant_message": messages[-1],
"completed_tool_results": [],
"pending_tool_calls": [],
"provider_state": conversation_state.checkpoint(messages),
},
)
if length_recovery_parts:
@ -764,6 +845,7 @@ class AgentRunner:
hook,
messages,
usage,
conversation_state,
)
if terminal_content is None:
terminal_content = self._max_iterations_fallback(spec)
@ -787,6 +869,7 @@ class AgentRunner:
tool_events=tool_events,
had_injections=had_injections,
pending_stream_content=pending_stream_content,
provider_state=conversation_state.finish(messages),
)
def _build_request_kwargs(
@ -817,6 +900,8 @@ class AgentRunner:
context: AgentHookContext,
*,
malformed_retry: bool = False,
conversation_state: ProviderConversationStateController,
provider_context: ProviderCallContext | None = None,
) -> LLMResponse:
timeout_s: float | None = spec.llm_timeout_s
if timeout_s is None:
@ -886,6 +971,7 @@ class AgentRunner:
coro = spec.runtime.provider.chat_stream_with_retry(
**kwargs,
provider_context=provider_context,
on_content_delta=_stream,
on_thinking_delta=_thinking,
on_tool_call_delta=_provider_tool_event,
@ -920,11 +1006,15 @@ class AgentRunner:
coro = spec.runtime.provider.chat_stream_with_retry(
**kwargs,
provider_context=provider_context,
on_content_delta=_stream_progress,
on_tool_call_delta=_provider_tool_event,
)
else:
coro = spec.runtime.provider.chat_with_retry(**kwargs)
coro = spec.runtime.provider.chat_with_retry(
**kwargs,
provider_context=provider_context,
)
# Streaming requests also have provider-level idle timeouts
# (NANOBOT_STREAM_IDLE_TIMEOUT_S), but a stream that keeps producing
@ -986,6 +1076,10 @@ class AgentRunner:
return await self._request_model(
spec, retry_messages, hook, context,
malformed_retry=True,
conversation_state=conversation_state,
provider_context=conversation_state.independent_request_context(
context_window_tokens=spec.runtime.context_window_tokens,
),
)
if (
all_dropped
@ -998,7 +1092,13 @@ class AgentRunner:
fallback_messages = self._malformed_tool_call_retry_messages(
messages, response.content,
)
return await self._request_no_tools(spec, fallback_messages)
return await self._request_no_tools(
spec,
fallback_messages,
provider_context=conversation_state.independent_request_context(
context_window_tokens=spec.runtime.context_window_tokens,
),
)
return response
@staticmethod
@ -1031,6 +1131,10 @@ class AgentRunner:
original_finish_reason,
)
response.tool_calls = valid
# The opaque candidate still contains every raw function_call item.
# Advancing it after dropping even one call would replay an unmatched
# call without a corresponding tool output on the next request.
response.provider_state = None
if not valid:
response.finish_reason = "stop"
return (dropped, not valid, original_finish_reason)
@ -1060,9 +1164,27 @@ class AgentRunner:
self,
spec: AgentRunSpec,
messages: list[dict[str, Any]],
*,
transcript: list[dict[str, Any]],
conversation_state: ProviderConversationStateController,
) -> LLMResponse:
retry_messages = self._finalization_retry_messages(messages)
return await self._request_no_tools(spec, retry_messages)
provider_context = conversation_state.prepare_request(
transcript,
context_window_tokens=spec.runtime.context_window_tokens,
supplemental_messages=[retry_messages[-1]],
)
response = await self._request_no_tools(
spec,
retry_messages,
provider_context=provider_context,
)
conversation_state.observe_response(
response,
transcript,
adopt_candidate_state=False,
)
return response
@staticmethod
def _finalization_retry_messages(messages: list[dict[str, Any]]) -> list[dict[str, Any]]:
@ -1076,10 +1198,17 @@ class AgentRunner:
hook: AgentHook,
messages: list[dict[str, Any]],
usage: dict[str, int],
conversation_state: ProviderConversationStateController,
) -> str | None:
retry_messages = self._budget_exhausted_finalization_messages(messages)
try:
response = await self._request_no_tools(spec, retry_messages)
response = await self._request_no_tools(
spec,
retry_messages,
provider_context=conversation_state.independent_request_context(
context_window_tokens=spec.runtime.context_window_tokens,
),
)
except Exception:
logger.exception(
"Budget-exhausted finalization failed for {}; using fallback",
@ -1115,9 +1244,18 @@ class AgentRunner:
self,
spec: AgentRunSpec,
messages: list[dict[str, Any]],
*,
provider_context: ProviderCallContext | None = None,
) -> LLMResponse:
kwargs = self._build_request_kwargs(spec, messages, tools=None)
return await spec.runtime.provider.chat_with_retry(**kwargs)
kwargs = self._build_request_kwargs(
spec,
messages,
tools=None,
)
return await spec.runtime.provider.chat_with_retry(
**kwargs,
provider_context=provider_context,
)
@staticmethod
def _budget_exhausted_finalization_messages(

View File

@ -23,14 +23,26 @@ import uuid
from collections.abc import Awaitable, Callable
from typing import Any, cast
from loguru import logger
from openai import AsyncOpenAI
from nanobot.providers.base import LLMProvider, LLMResponse
from nanobot.providers.base import (
LLMProvider,
LLMResponse,
ProviderCallContext,
ProviderConversationState,
)
from nanobot.providers.openai_responses import (
ResponsesStreamCapture,
build_responses_state,
consume_sdk_stream,
convert_messages,
convert_tools,
is_compaction_compatibility_error,
is_replayable_finish_reason,
parse_response_output,
prepare_responses_input,
resolve_compact_threshold,
responses_state_matches,
)
_AZURE_OPENAI_SCOPE = "https://cognitiveservices.azure.com/.default"
@ -97,6 +109,7 @@ class AzureOpenAIProvider(LLMProvider):
):
super().__init__(api_key, api_base)
self.default_model = default_model
self._native_compaction_available = True
if not api_base:
raise ValueError("Azure OpenAI api_base is required")
@ -142,6 +155,25 @@ class AzureOpenAIProvider(LLMProvider):
name = deployment_name.lower()
return not any(token in name for token in ("gpt-5", "o1", "o3", "o4"))
def _responses_state_provider(self) -> str:
return f"azure_openai:{str(self.api_base).rstrip('/')}"
def can_resume_conversation_state(
self,
state: ProviderConversationState,
model: str | None = None,
) -> bool:
return responses_state_matches(
state,
provider=self._responses_state_provider(),
model=model or self.default_model,
)
def supports_native_compaction(self, model: str | None = None) -> bool:
"""Azure's native Responses endpoint accepts context management."""
_ = model
return self._native_compaction_available
def _build_body(
self,
messages: list[dict[str, Any]],
@ -151,10 +183,26 @@ class AzureOpenAIProvider(LLMProvider):
temperature: float,
reasoning_effort: str | None,
tool_choice: str | dict[str, Any] | None,
provider_context: ProviderCallContext | None = None,
) -> dict[str, Any]:
"""Build the Responses API request body from Chat-Completions-style args."""
deployment = model or self.default_model
instructions, input_items = convert_messages(self._sanitize_empty_content(messages))
sanitized_messages = self._sanitize_empty_content(messages)
sanitized_state = (
provider_context.conversation_state
if provider_context is not None
else None
)
if sanitized_state is not None:
sanitized_state = sanitized_state.with_pending_messages(
self._sanitize_empty_content(sanitized_state.pending_messages)
)
instructions, input_items, replayed = prepare_responses_input(
sanitized_messages,
state=sanitized_state,
provider=self._responses_state_provider(),
model=deployment,
)
body: dict[str, Any] = {
"model": deployment,
@ -164,13 +212,29 @@ class AzureOpenAIProvider(LLMProvider):
"store": False,
"stream": False,
}
compact_threshold = resolve_compact_threshold(
(
provider_context.context_window_tokens
if provider_context is not None
else None
),
max_tokens,
)
if self.supports_native_compaction(deployment) and compact_threshold is not None:
body["context_management"] = [{
"type": "compaction",
"compact_threshold": compact_threshold,
}]
if self._supports_temperature(deployment, reasoning_effort):
body["temperature"] = temperature
if not self._supports_temperature(deployment, reasoning_effort):
body["include"] = ["reasoning.encrypted_content"]
if reasoning_effort and reasoning_effort.lower() != "none":
body["reasoning"] = {"effort": reasoning_effort}
body["include"] = ["reasoning.encrypted_content"]
if replayed and "gpt-5.6" in deployment.lower():
body.setdefault("reasoning", {})["context"] = "all_turns"
if tools:
body["tools"] = convert_tools(tools)
@ -178,21 +242,97 @@ class AzureOpenAIProvider(LLMProvider):
return body
async def _create_response_with_compaction_fallback(
self,
body: dict[str, Any],
) -> Any:
"""Retry once without server compaction when Azure rejects the option."""
try:
return cast(Any, await self._client.responses.create(**body))
except Exception as exc:
if (
"context_management" not in body
or not is_compaction_compatibility_error(exc)
):
raise
self._native_compaction_available = False
body.pop("context_management", None)
logger.warning(
"Azure Responses server compaction unsupported; disabled for this provider "
"instance (status={})",
getattr(exc, "status_code", None),
)
return cast(Any, await self._client.responses.create(**body))
@staticmethod
def _handle_error(e: Exception) -> LLMResponse:
response = getattr(e, "response", None)
body = getattr(e, "body", None) or getattr(response, "text", None)
body_text = str(body).strip() if body is not None else ""
msg = f"Error: {body_text[:500]}" if body_text else f"Error calling Azure OpenAI: {e}"
retry_after = LLMProvider._extract_retry_after_from_headers(getattr(response, "headers", None))
headers = getattr(response, "headers", None)
retry_after = LLMProvider._extract_retry_after_from_headers(headers)
if retry_after is None:
retry_after = LLMProvider._extract_retry_after(msg)
return LLMResponse(content=msg, finish_reason="error", retry_after=retry_after)
status_code = getattr(e, "status_code", None)
if status_code is None and response is not None:
status_code = getattr(response, "status_code", None)
error_type, error_code = LLMProvider._extract_error_type_code(body)
should_retry: bool | None = None
if headers is not None:
raw_should_retry = headers.get("x-should-retry")
if isinstance(raw_should_retry, str):
lowered = raw_should_retry.strip().lower()
if lowered == "true":
should_retry = True
elif lowered == "false":
should_retry = False
error_name = type(e).__name__.lower()
error_kind = (
"timeout"
if "timeout" in error_name
else "connection"
if "connection" in error_name
else None
)
return LLMResponse(
content=msg,
finish_reason="error",
retry_after=retry_after,
error_status_code=int(status_code) if status_code is not None else None,
error_kind=error_kind,
error_type=error_type,
error_code=error_code,
error_retry_after_s=retry_after,
error_should_retry=should_retry,
)
# ------------------------------------------------------------------
# Public API
# ------------------------------------------------------------------
async def chat_with_context(
self,
*,
provider_context: ProviderCallContext,
**kwargs: Any,
) -> LLMResponse:
return await self.chat(
**kwargs,
provider_context=provider_context,
)
async def chat_stream_with_context(
self,
*,
provider_context: ProviderCallContext,
**kwargs: Any,
) -> LLMResponse:
return await self.chat_stream(
**kwargs,
provider_context=provider_context,
)
async def chat(
self,
messages: list[dict[str, Any]],
@ -202,14 +342,21 @@ class AzureOpenAIProvider(LLMProvider):
temperature: float = 0.7,
reasoning_effort: str | None = None,
tool_choice: str | dict[str, Any] | None = None,
provider_context: ProviderCallContext | None = None,
) -> LLMResponse:
body = self._build_body(
messages, tools, model, max_tokens, temperature,
reasoning_effort, tool_choice,
provider_context,
)
try:
response = cast(Any, await self._client.responses.create(**body))
return parse_response_output(response)
response = await self._create_response_with_compaction_fallback(body)
return parse_response_output(
response,
state_provider=self._responses_state_provider(),
state_model=str(body["model"]),
state_input_items=cast(list[dict[str, Any]], body["input"]),
)
except Exception as e:
return self._handle_error(e)
@ -225,26 +372,43 @@ class AzureOpenAIProvider(LLMProvider):
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
on_thinking_delta: Callable[[str], Awaitable[None]] | None = None,
on_tool_call_delta: Callable[[dict[str, Any]], Awaitable[None]] | None = None,
provider_context: ProviderCallContext | None = None,
) -> LLMResponse:
_ = on_thinking_delta
body = self._build_body(
messages, tools, model, max_tokens, temperature,
reasoning_effort, tool_choice,
provider_context,
)
body["stream"] = True
try:
stream = cast(Any, await self._client.responses.create(**body))
stream = await self._create_response_with_compaction_fallback(body)
capture = ResponsesStreamCapture()
content, tool_calls, finish_reason, usage, reasoning_content = (
await consume_sdk_stream(stream, on_content_delta, on_tool_call_delta)
await consume_sdk_stream(
stream,
on_content_delta,
on_tool_call_delta,
capture=capture,
)
)
return LLMResponse(
result = LLMResponse(
content=content or None,
tool_calls=tool_calls,
finish_reason=finish_reason,
usage=usage,
reasoning_content=reasoning_content,
)
if capture.completed and is_replayable_finish_reason(finish_reason):
result.provider_state = build_responses_state(
provider=self._responses_state_provider(),
model=str(body["model"]),
input_items=cast(list[dict[str, Any]], body["input"]),
output_items=capture.output_items,
usage=usage,
)
return result
except Exception as e:
return self._handle_error(e)

View File

@ -1,5 +1,7 @@
"""Base LLM provider interface."""
from __future__ import annotations
import asyncio
import json
import os
@ -7,6 +9,7 @@ import re
from abc import ABC, abstractmethod
from collections.abc import Awaitable, Callable
from contextlib import suppress
from copy import deepcopy
from dataclasses import dataclass, field
from datetime import datetime, timezone
from email.utils import parsedate_to_datetime
@ -150,6 +153,104 @@ def tool_arguments_json_for_replay(arguments: Any) -> str:
return json.dumps(tool_arguments_object_for_replay(arguments), ensure_ascii=False)
@dataclass
class ProviderConversationState:
"""Opaque provider-owned continuation state.
``payload`` may contain encrypted reasoning or other provider-private
protocol items. Keep it out of normal logs and public chat history.
``pending_messages`` are Chat-style messages produced after the most
recent provider response and are materialized by the owning provider on
the next request.
"""
kind: str
provider: str
model: str
version: int
payload: dict[str, Any] = field(default_factory=dict, repr=False)
pending_messages: list[dict[str, Any]] = field(default_factory=list, repr=False)
def with_pending_messages(
self,
messages: list[dict[str, Any]],
) -> ProviderConversationState:
"""Return a state copy with an isolated pending-message list."""
return ProviderConversationState(
kind=self.kind,
provider=self.provider,
model=self.model,
version=self.version,
payload=self.payload,
pending_messages=deepcopy(messages),
)
def to_private_record(self) -> dict[str, Any]:
"""Serialize for the private session sidecar, never for public history."""
return {
"kind": self.kind,
"provider": self.provider,
"model": self.model,
"version": self.version,
"payload": deepcopy(self.payload),
"pending_messages": deepcopy(self.pending_messages),
}
@classmethod
def from_private_record(
cls,
value: object,
) -> ProviderConversationState | None:
"""Validate and deserialize a private session-sidecar value."""
if not isinstance(value, dict):
return None
data = cast(dict[str, Any], value)
kind = data.get("kind")
provider = data.get("provider")
model = data.get("model")
version = data.get("version")
payload = data.get("payload")
pending = data.get("pending_messages", [])
if (
not isinstance(kind, str)
or not kind
or not isinstance(provider, str)
or not provider
or not isinstance(model, str)
or not model
or isinstance(version, bool)
or not isinstance(version, int)
or not isinstance(payload, dict)
or not isinstance(pending, list)
or any(
not isinstance(message, dict)
for message in cast(list[object], pending)
)
):
return None
return cls(
kind=kind,
provider=provider,
model=model,
version=version,
payload=deepcopy(cast(dict[str, Any], payload)),
pending_messages=deepcopy(cast(list[dict[str, Any]], pending)),
)
@dataclass(frozen=True)
class ProviderCallContext:
"""Optional provider-owned continuation data for one model request.
The regular ``chat`` contract stays provider-agnostic. Responses-capable
providers consume this context through the opt-in ``chat_with_context``
hooks, while every other provider inherits the context-free delegation.
"""
conversation_state: ProviderConversationState | None = field(default=None, repr=False)
context_window_tokens: int | None = None
@dataclass
class LLMResponse:
"""Response from an LLM provider."""
@ -160,6 +261,10 @@ class LLMResponse:
retry_after: float | None = None # Provider supplied retry wait in seconds.
reasoning_content: str | None = None # Kimi, DeepSeek-R1, MiMo etc.
thinking_blocks: list[dict[str, Any]] | None = None # Anthropic extended thinking
provider_state: ProviderConversationState | None = field(default=None, repr=False)
# Routing wrappers may preserve or discard an incoming provider-owned
# continuation independently of the final fallback error's retry policy.
preserve_provider_state_on_error: bool | None = field(default=None, repr=False)
# Structured error metadata used by retry policy when finish_reason == "error".
error_status_code: int | None = None
error_kind: str | None = None # e.g. "timeout", "connection"
@ -274,6 +379,18 @@ class LLMProvider(ABC):
self.api_base = api_base
self.generation: GenerationSettings = GenerationSettings()
def can_resume_conversation_state(
self,
state: ProviderConversationState,
model: str | None = None,
) -> bool:
"""Whether this provider can safely consume an opaque saved state."""
return False
def supports_native_compaction(self, model: str | None = None) -> bool:
"""Whether requests may include provider-native context compaction."""
return False
@staticmethod
def _sanitize_empty_content(messages: list[dict[str, Any]]) -> list[dict[str, Any]]:
"""Sanitize message content: fix empty blocks, strip internal _meta fields.
@ -416,7 +533,7 @@ class LLMProvider(ABC):
return any(marker in err for marker in cls._TRANSIENT_ERROR_MARKERS)
@classmethod
def _is_transient_response(cls, response: LLMResponse) -> bool:
def is_transient_response(cls, response: LLMResponse) -> bool:
"""Prefer structured error metadata, fallback to text markers for legacy providers."""
if response.error_should_retry is not None:
return bool(response.error_should_retry)
@ -607,6 +724,21 @@ class LLMProvider(ABC):
result.append(msg)
return result if found else None
@staticmethod
def _contains_image_content(value: object) -> bool:
"""Return whether a JSON-like provider payload contains an input image."""
if isinstance(value, dict):
mapping = cast(dict[str, object], value)
if mapping.get("type") in {"image_url", "input_image"}:
return True
return any(LLMProvider._contains_image_content(item) for item in mapping.values())
if isinstance(value, list):
return any(
LLMProvider._contains_image_content(item)
for item in cast(list[object], value)
)
return False
@staticmethod
def _strip_image_content_inplace(messages: list[dict[str, Any]]) -> bool:
"""Replace image_url blocks with text placeholder *in-place*.
@ -633,6 +765,12 @@ class LLMProvider(ABC):
async def _safe_chat(self, **kwargs: Any) -> LLMResponse:
"""Call chat() and convert unexpected exceptions to error responses."""
try:
provider_context = kwargs.pop("provider_context", None)
if isinstance(provider_context, ProviderCallContext):
return await self.chat_with_context(
provider_context=provider_context,
**kwargs,
)
return await self.chat(**kwargs)
except asyncio.CancelledError:
raise
@ -666,17 +804,47 @@ class LLMProvider(ABC):
"""
_ = on_thinking_delta, on_tool_call_delta
response = await self.chat(
messages=messages, tools=tools, model=model,
max_tokens=max_tokens, temperature=temperature,
reasoning_effort=reasoning_effort, tool_choice=tool_choice,
messages=messages,
tools=tools,
model=model,
max_tokens=max_tokens,
temperature=temperature,
reasoning_effort=reasoning_effort,
tool_choice=tool_choice,
)
if on_content_delta and response.content:
await on_content_delta(response.content)
return response
async def chat_with_context(
self,
*,
provider_context: ProviderCallContext,
**kwargs: Any,
) -> LLMResponse:
"""Opt-in continuation hook; ordinary providers delegate to ``chat``."""
_ = provider_context
return await self.chat(**kwargs)
async def chat_stream_with_context(
self,
*,
provider_context: ProviderCallContext,
**kwargs: Any,
) -> LLMResponse:
"""Streaming continuation hook with a context-free default."""
_ = provider_context
return await self.chat_stream(**kwargs)
async def _safe_chat_stream(self, **kwargs: Any) -> LLMResponse:
"""Call chat_stream() and convert unexpected exceptions to error responses."""
try:
provider_context = kwargs.pop("provider_context", None)
if isinstance(provider_context, ProviderCallContext):
return await self.chat_stream_with_context(
provider_context=provider_context,
**kwargs,
)
return await self.chat_stream(**kwargs)
except asyncio.CancelledError:
raise
@ -698,6 +866,7 @@ class LLMProvider(ABC):
on_stream_recover: Callable[[], Awaitable[None]] | None = None,
retry_mode: str = "standard",
on_retry_wait: Callable[[str], Awaitable[None]] | None = None,
provider_context: ProviderCallContext | None = None,
) -> LLMResponse:
"""Call chat_stream() with retry on transient provider failures."""
if max_tokens is self._SENTINEL or max_tokens is None:
@ -730,6 +899,8 @@ class LLMProvider(ABC):
on_thinking_delta=on_thinking_delta,
on_tool_call_delta=on_tool_call_delta,
)
if provider_context is not None:
kw["provider_context"] = provider_context
if on_stream_recover and getattr(self, "supports_stream_recover_callback", False):
kw["on_stream_recover"] = _recover_stream
return await self._run_with_retry(
@ -753,6 +924,7 @@ class LLMProvider(ABC):
tool_choice: str | dict[str, Any] | None = None,
retry_mode: str = "standard",
on_retry_wait: Callable[[str], Awaitable[None]] | None = None,
provider_context: ProviderCallContext | None = None,
) -> LLMResponse:
"""Call chat() with retry on transient provider failures.
@ -775,6 +947,8 @@ class LLMProvider(ABC):
max_tokens=max_tokens, temperature=temperature,
reasoning_effort=reasoning_effort, tool_choice=tool_choice,
)
if provider_context is not None:
kw["provider_context"] = provider_context
return await self._run_with_retry(
self._safe_chat,
kw,
@ -932,14 +1106,33 @@ class LLMProvider(ABC):
last_error_key = error_key
identical_error_count = 1 if error_key else 0
if not self._is_transient_response(response):
stripped = self._strip_image_content(original_messages)
if stripped is not None and stripped != kw["messages"]:
if not self.is_transient_response(response):
stripped = self._strip_image_content(kw["messages"])
provider_context = kw.get("provider_context")
stripped_context: ProviderCallContext | None = None
if isinstance(provider_context, ProviderCallContext):
state = provider_context.conversation_state
if state is not None and (
stripped is not None
or self._strip_image_content(state.pending_messages) is not None
or self._contains_image_content(state.payload)
):
# Provider-owned payloads may retain earlier input_image items.
# Rebuild from the stripped public transcript for this retry.
stripped_context = ProviderCallContext(
context_window_tokens=(
provider_context.context_window_tokens
),
)
if stripped is not None or stripped_context is not None:
logger.warning(
"Non-transient LLM error with image content, retrying without images"
)
retry_kw = dict(kw)
retry_kw["messages"] = stripped
if stripped is not None:
retry_kw["messages"] = stripped
if stripped_context is not None:
retry_kw["provider_context"] = stripped_context
result = await call(**retry_kw)
# Permanently strip images from the original messages so
# subsequent iterations do not repeat the error-retry cycle.

View File

@ -0,0 +1,262 @@
"""Provider-owned conversation-state lifecycle coordination."""
from __future__ import annotations
from copy import deepcopy
from typing import Any, cast
from nanobot.providers.base import (
LLMProvider,
LLMResponse,
ProviderCallContext,
ProviderConversationState,
)
_PROVIDER_STATE_OUTPUT_META = "provider_state_output"
_PROVIDER_STATE_BOUNDARY_META = "provider_state_boundary"
def allows_conversation_message_merge(message: dict[str, Any]) -> bool:
"""Return whether new same-role input may merge into *message*."""
internal_meta = cast(object, message.get("_meta"))
return not (
isinstance(internal_meta, dict)
and cast(dict[str, Any], internal_meta).get(
_PROVIDER_STATE_BOUNDARY_META
) is True
)
class ProviderConversationStateController:
"""Keep provider conversation-state semantics outside the agent runner.
The runner owns the tool loop and reports lifecycle events here. This
controller owns capability checks, transcript deltas, response projections,
retry transitions, and durable snapshots for provider-private state.
"""
def __init__(
self,
*,
provider: LLMProvider,
model: str | None,
messages: list[dict[str, Any]],
state: ProviderConversationState | None = None,
) -> None:
self._provider = provider
self._model = model
self._state = (
state
if state is not None
and provider.can_resume_conversation_state(state, model)
else None
)
self._boundary = len(messages)
self._request_messages: list[dict[str, Any]] = []
def independent_request_context(
self,
*,
context_window_tokens: int | None,
) -> ProviderCallContext | None:
"""Return typed provider context for a request that does not resume state."""
if context_window_tokens is None:
return None
return ProviderCallContext(context_window_tokens=context_window_tokens)
def prepare_request(
self,
messages: list[dict[str, Any]],
*,
context_window_tokens: int | None,
model_messages: list[dict[str, Any]] | None = None,
supplemental_messages: list[dict[str, Any]] | None = None,
) -> ProviderCallContext | None:
"""Build typed context for the next request and remember its durable delta."""
independent_context = self.independent_request_context(
context_window_tokens=context_window_tokens,
)
if self._state is None:
self._request_messages = []
return independent_context
if not self._provider.can_resume_conversation_state(
self._state,
self._model,
):
self._state = None
self._request_messages = []
return independent_context
durable_messages = self._messages_after_boundary(messages)
governed_messages = (
self._model_messages_after_boundary(model_messages)
if model_messages is not None and durable_messages
else None
)
request_messages = (
governed_messages
if governed_messages is not None
else durable_messages
)
supplemental = deepcopy(supplemental_messages or [])
self._request_messages = deepcopy(request_messages)
request_state = self._state.with_pending_messages([
*self._state.pending_messages,
*request_messages,
*supplemental,
])
return ProviderCallContext(
conversation_state=request_state,
context_window_tokens=(
independent_context.context_window_tokens
if independent_context is not None
else None
),
)
def observe_response(
self,
response: LLMResponse,
messages: list[dict[str, Any]],
*,
adopt_candidate_state: bool = True,
) -> None:
"""Advance, preserve, or discard state after one provider response."""
candidate = response.provider_state if adopt_candidate_state else None
candidate_is_replayable = response.finish_reason in {
"stop",
"tool_calls",
"function_call",
}
if (
candidate is not None
and candidate_is_replayable
and self._provider.can_resume_conversation_state(
candidate,
self._model,
)
):
self._state = candidate
self._boundary = len(messages)
self._seal_boundary(messages)
elif response.finish_reason == "error" and (
response.preserve_provider_state_on_error is True
or (
response.preserve_provider_state_on_error is None
and LLMProvider.is_transient_response(response)
)
):
if self._state is not None and self._request_messages:
self._state = self._state.with_pending_messages([
*self._state.pending_messages,
*self._request_messages,
])
self._boundary = len(messages)
else:
self._state = None
self._boundary = len(messages)
self._request_messages = []
@staticmethod
def project_response_message(
message: dict[str, Any],
response: LLMResponse,
) -> dict[str, Any]:
"""Mark a Chat projection already represented by provider output."""
if response.provider_state is None:
return message
internal_meta = dict(message.get("_meta") or {})
internal_meta[_PROVIDER_STATE_OUTPUT_META] = True
message["_meta"] = internal_meta
return message
def checkpoint(
self,
messages: list[dict[str, Any]],
*,
model_messages: list[dict[str, Any]] | None = None,
) -> ProviderConversationState | None:
"""Return a durable state snapshot without changing live state."""
if self._state is None:
return None
durable_messages = self._messages_after_boundary(messages)
governed_messages = (
self._model_messages_after_boundary(model_messages)
if model_messages is not None and durable_messages
else None
)
pending_messages = (
governed_messages
if governed_messages is not None
else durable_messages
)
return self._state.with_pending_messages([
*self._state.pending_messages,
*pending_messages,
])
def finish(
self,
messages: list[dict[str, Any]],
) -> ProviderConversationState | None:
"""Return the final durable state after all runner messages are known."""
self._state = self.checkpoint(messages)
return self._state
def _messages_after_boundary(
self,
messages: list[dict[str, Any]],
) -> list[dict[str, Any]]:
pending: list[dict[str, Any]] = []
for message in messages[self._boundary:]:
internal_meta = cast(object, message.get("_meta"))
if (
isinstance(internal_meta, dict)
and cast(dict[str, Any], internal_meta).get(
_PROVIDER_STATE_OUTPUT_META
) is True
):
continue
pending.append(deepcopy(message))
return pending
@staticmethod
def _model_messages_after_boundary(
messages: list[dict[str, Any]],
) -> list[dict[str, Any]] | None:
"""Return the governed delta after the latest provider-owned boundary."""
boundary = None
for idx in range(len(messages) - 1, -1, -1):
internal_meta = cast(object, messages[idx].get("_meta"))
if (
isinstance(internal_meta, dict)
and cast(dict[str, Any], internal_meta).get(
_PROVIDER_STATE_BOUNDARY_META
) is True
):
boundary = idx
break
if boundary is None:
return None
pending: list[dict[str, Any]] = []
for message in messages[boundary + 1:]:
internal_meta = cast(object, message.get("_meta"))
if (
isinstance(internal_meta, dict)
and cast(dict[str, Any], internal_meta).get(
_PROVIDER_STATE_OUTPUT_META
) is True
):
continue
pending.append(deepcopy(message))
return pending
@staticmethod
def _seal_boundary(messages: list[dict[str, Any]]) -> None:
"""Prevent later same-role injection merging across a state boundary."""
if not messages:
return
internal_meta = dict(messages[-1].get("_meta") or {})
internal_meta[_PROVIDER_STATE_BOUNDARY_META] = True
messages[-1]["_meta"] = internal_meta

View File

@ -261,6 +261,7 @@ def make_provider(
primary=provider,
fallback_presets=fallback_presets,
provider_factory=lambda fb: _make_provider_core(config, preset=fb),
primary_context_window_tokens=resolved.context_window_tokens,
)
return provider

View File

@ -6,11 +6,18 @@ from __future__ import annotations
import time
from collections.abc import Awaitable, Callable
from dataclasses import replace
from typing import Any
from loguru import logger
from nanobot.providers.base import GenerationSettings, LLMProvider, LLMResponse
from nanobot.providers.base import (
GenerationSettings,
LLMProvider,
LLMResponse,
ProviderCallContext,
ProviderConversationState,
)
# Circuit breaker tuned to match OpenAICompatProvider's Responses API breaker.
_PRIMARY_FAILURE_THRESHOLD = 3
@ -113,11 +120,13 @@ class FallbackProvider(LLMProvider):
fallback_presets: list[Any],
provider_factory: Callable[[Any], LLMProvider],
fallback_model_observer: FallbackModelObserver | None = None,
primary_context_window_tokens: int | None = None,
):
self._primary = primary
self._fallback_presets = list(fallback_presets)
self._provider_factory = provider_factory
self._fallback_model_observer = fallback_model_observer
self._primary_context_window_tokens = primary_context_window_tokens
self._has_fallbacks = bool(fallback_presets)
self._primary_failures = 0
self._primary_tripped_at: float | None = None
@ -141,6 +150,33 @@ class FallbackProvider(LLMProvider):
def supports_progress_deltas(self) -> bool:
return bool(getattr(self._primary, "supports_progress_deltas", False))
def can_resume_conversation_state(
self,
state: ProviderConversationState,
model: str | None = None,
) -> bool:
return self._primary.can_resume_conversation_state(state, model)
def supports_native_compaction(self, model: str | None = None) -> bool:
return self._primary.supports_native_compaction(model)
def _primary_call_context(
self,
provider_context: ProviderCallContext,
model: str | None,
) -> ProviderCallContext:
context_window_tokens = (
self._primary_context_window_tokens
if self._primary_context_window_tokens is not None
else provider_context.context_window_tokens
)
if not self._primary.supports_native_compaction(model):
context_window_tokens = None
return ProviderCallContext(
conversation_state=provider_context.conversation_state,
context_window_tokens=context_window_tokens,
)
def _primary_available(self) -> bool:
"""Return True if the primary provider is not currently tripped."""
if self._primary_tripped_at is None:
@ -157,6 +193,25 @@ class FallbackProvider(LLMProvider):
lambda p, kw: p.chat(**kw), kwargs, has_streamed=None
)
async def chat_with_context(
self,
*,
provider_context: ProviderCallContext,
**kwargs: Any,
) -> LLMResponse:
call_kwargs: dict[str, Any] = dict(kwargs)
call_kwargs["provider_context"] = self._primary_call_context(
provider_context,
kwargs.get("model"),
)
if not self._has_fallbacks:
return await self._primary.chat_with_context(**call_kwargs)
return await self._try_with_fallback(
lambda p, kw: p.chat_with_context(**kw),
call_kwargs,
has_streamed=None,
)
async def chat_stream(self, **kwargs: Any) -> LLMResponse:
on_stream_recover = kwargs.pop("on_stream_recover", None)
if not self._has_fallbacks:
@ -179,6 +234,38 @@ class FallbackProvider(LLMProvider):
on_stream_recover=on_stream_recover,
)
async def chat_stream_with_context(
self,
*,
provider_context: ProviderCallContext,
**kwargs: Any,
) -> LLMResponse:
on_stream_recover = kwargs.pop("on_stream_recover", None)
call_kwargs: dict[str, Any] = dict(kwargs)
call_kwargs["provider_context"] = self._primary_call_context(
provider_context,
kwargs.get("model"),
)
if not self._has_fallbacks:
return await self._primary.chat_stream_with_context(**call_kwargs)
has_streamed: list[bool] = [False]
original_delta = call_kwargs.get("on_content_delta")
async def _tracking_delta(text: str) -> None:
if text:
has_streamed[0] = True
if original_delta:
await original_delta(text)
call_kwargs["on_content_delta"] = _tracking_delta
return await self._try_with_fallback(
lambda p, kw: p.chat_stream_with_context(**kw),
call_kwargs,
has_streamed=has_streamed,
on_stream_recover=on_stream_recover,
)
async def _try_with_fallback(
self,
call: Callable[[LLMProvider, dict[str, Any]], Awaitable[LLMResponse]],
@ -189,6 +276,9 @@ class FallbackProvider(LLMProvider):
primary_model = kwargs.get("model") or self._primary.get_default_model()
primary_was_attempted = False
primary_error = "unknown error"
# A primary error eligible for failover did not return a replacement
# continuation, so the incoming primary state remains reusable.
preserve_primary_state = True
if self._primary_available():
primary_was_attempted = True
@ -286,6 +376,23 @@ class FallbackProvider(LLMProvider):
"max_tokens": fallback.max_tokens,
"temperature": fallback.temperature,
}
provider_context = fallback_kwargs.get("provider_context")
if isinstance(provider_context, ProviderCallContext):
state = provider_context.conversation_state
if state is not None and not fallback_provider.can_resume_conversation_state(
state,
fallback_model,
):
state = None
context_window_tokens = (
fallback.context_window_tokens
if fallback_provider.supports_native_compaction(fallback_model)
else None
)
fallback_kwargs["provider_context"] = ProviderCallContext(
conversation_state=state,
context_window_tokens=context_window_tokens,
)
if fallback.reasoning_effort is None:
fallback_kwargs.pop("reasoning_effort", None)
else:
@ -312,11 +419,15 @@ class FallbackProvider(LLMProvider):
)
# Return the last error response we saw (primary or last fallback).
if last_response is not None:
return last_response
return replace(
last_response,
preserve_provider_state_on_error=preserve_primary_state,
)
# Primary was tripped and we have no fallbacks — synthesize an error.
return LLMResponse(
content=f"Primary model '{primary_model}' circuit open and no fallbacks available",
finish_reason="error",
preserve_provider_state_on_error=preserve_primary_state,
)
async def _notify_fallback_model(self, model: str) -> None:

View File

@ -16,7 +16,7 @@ import httpx
from oauth_cli_kit.models import OAuthToken
from oauth_cli_kit.storage import FileTokenStorage
from nanobot.providers.base import LLMResponse
from nanobot.providers.base import LLMResponse, ProviderCallContext
from nanobot.providers.openai_compat_provider import OpenAICompatProvider
DEFAULT_GITHUB_DEVICE_CODE_URL = "https://github.com/login/device/code"
@ -248,6 +248,7 @@ class GitHubCopilotProvider(OpenAICompatProvider):
temperature: float = 0.7,
reasoning_effort: str | None = None,
tool_choice: str | dict[str, Any] | None = None,
provider_context: ProviderCallContext | None = None,
) -> LLMResponse:
await self._refresh_client_api_key()
return await super().chat(
@ -258,6 +259,7 @@ class GitHubCopilotProvider(OpenAICompatProvider):
temperature=temperature,
reasoning_effort=reasoning_effort,
tool_choice=tool_choice,
provider_context=provider_context,
)
async def chat_stream(
@ -272,6 +274,7 @@ class GitHubCopilotProvider(OpenAICompatProvider):
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
on_thinking_delta: Callable[[str], Awaitable[None]] | None = None,
on_tool_call_delta: Callable[[dict[str, Any]], Awaitable[None]] | None = None,
provider_context: ProviderCallContext | None = None,
) -> LLMResponse:
await self._refresh_client_api_key()
return await super().chat_stream(
@ -285,4 +288,5 @@ class GitHubCopilotProvider(OpenAICompatProvider):
on_content_delta=on_content_delta,
on_thinking_delta=on_thinking_delta,
on_tool_call_delta=on_tool_call_delta,
provider_context=provider_context,
)

View File

@ -17,17 +17,27 @@ from oauth_cli_kit import get_token as get_codex_token
from nanobot.providers.base import (
LLMProvider,
LLMResponse,
ToolCallRequest,
ProviderCallContext,
ProviderConversationState,
resolve_stream_idle_timeout_s,
)
from nanobot.providers.openai_responses import (
ResponsesStreamCapture,
build_responses_state,
consume_sse_with_reasoning,
convert_messages,
convert_tools,
is_compaction_compatibility_error,
is_replayable_finish_reason,
prepare_responses_input,
resolve_compact_threshold,
responses_state_context_tokens,
responses_state_items,
responses_state_matches,
)
DEFAULT_CODEX_URL = "https://chatgpt.com/backend-api/codex/responses"
DEFAULT_ORIGINATOR = "nanobot"
_COMPACTION_RETAINED_CHAR_BUDGET = 256_000
class OpenAICodexProvider(LLMProvider):
@ -45,21 +55,39 @@ class OpenAICodexProvider(LLMProvider):
self.default_model = default_model
self.proxy = proxy or None
self._extra_body = dict(extra_body or {})
self._native_compaction_available = True
async def _call_codex(
self,
messages: list[dict[str, Any]],
tools: list[dict[str, Any]] | None,
model: str | None,
max_tokens: int,
reasoning_effort: str | None,
tool_choice: str | dict[str, Any] | None,
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
on_thinking_delta: Callable[[str], Awaitable[None]] | None = None,
on_tool_call_delta: Callable[[dict[str, Any]], Awaitable[None]] | None = None,
provider_context: ProviderCallContext | None = None,
) -> LLMResponse:
"""Shared request logic for both chat() and chat_stream()."""
model = model or self.default_model
system_prompt, input_items = convert_messages(messages)
sanitized_messages = self._sanitize_empty_content(messages)
sanitized_state = (
provider_context.conversation_state
if provider_context is not None
else None
)
if sanitized_state is not None:
sanitized_state = sanitized_state.with_pending_messages(
self._sanitize_empty_content(sanitized_state.pending_messages)
)
system_prompt, input_items, replayed = prepare_responses_input(
sanitized_messages,
state=sanitized_state,
provider=self._responses_state_provider(),
model=_strip_model_prefix(model),
)
body: dict[str, Any] = {
"model": _strip_model_prefix(model),
@ -68,12 +96,15 @@ class OpenAICodexProvider(LLMProvider):
"instructions": system_prompt,
"input": input_items,
"text": {"verbosity": "medium"},
"include": ["reasoning.encrypted_content"],
"prompt_cache_key": _prompt_cache_key(messages[:2]),
"tool_choice": tool_choice or "auto",
"parallel_tool_calls": True,
}
body["include"] = ["reasoning.encrypted_content"]
reasoning_options = _build_reasoning_options(reasoning_effort)
if replayed and "gpt-5.6" in _strip_model_prefix(model).lower():
reasoning_options = dict(reasoning_options or {})
reasoning_options["context"] = "all_turns"
if reasoning_options:
body["reasoning"] = reasoning_options
if tools:
@ -87,33 +118,90 @@ class OpenAICodexProvider(LLMProvider):
token = await asyncio.to_thread(get_codex_token, proxy=self.proxy)
headers = _build_headers(cast(str, token.account_id), token.access)
stage = "codex_request"
try:
content, tool_calls, finish_reason, usage, reasoning_content = await _request_codex(
DEFAULT_CODEX_URL, headers, body, verify=True,
proxy=self.proxy,
on_content_delta=on_content_delta,
on_thinking_delta=on_thinking_delta,
on_tool_call_delta=on_tool_call_delta,
)
except Exception as e:
if "CERTIFICATE_VERIFY_FAILED" not in str(e):
raise
logger.warning("SSL verification failed for Codex API; retrying with verify=False")
content, tool_calls, finish_reason, usage, reasoning_content = await _request_codex(
DEFAULT_CODEX_URL, headers, body, verify=False,
proxy=self.proxy,
on_content_delta=on_content_delta,
on_thinking_delta=on_thinking_delta,
on_tool_call_delta=on_tool_call_delta,
)
return LLMResponse(
content=content,
tool_calls=tool_calls,
finish_reason=finish_reason,
usage=usage,
reasoning_content=reasoning_content,
async def _send(
request_body: dict[str, Any],
*,
emit_deltas: bool,
) -> LLMResponse:
wire_body = _without_response_item_ids(request_body)
try:
return await _request_codex(
DEFAULT_CODEX_URL,
headers,
wire_body,
verify=True,
proxy=self.proxy,
on_content_delta=on_content_delta if emit_deltas else None,
on_thinking_delta=on_thinking_delta if emit_deltas else None,
on_tool_call_delta=on_tool_call_delta if emit_deltas else None,
)
except Exception as exc:
if "CERTIFICATE_VERIFY_FAILED" not in str(exc):
raise
logger.warning(
"SSL verification failed for Codex API; retrying with verify=False"
)
return await _request_codex(
DEFAULT_CODEX_URL,
headers,
wire_body,
verify=False,
proxy=self.proxy,
on_content_delta=on_content_delta if emit_deltas else None,
on_thinking_delta=on_thinking_delta if emit_deltas else None,
on_tool_call_delta=on_tool_call_delta if emit_deltas else None,
)
compact_threshold = resolve_compact_threshold(
(
provider_context.context_window_tokens
if provider_context is not None
else None
),
max_tokens,
)
if (
self.supports_native_compaction(model)
and replayed
and sanitized_state is not None
and compact_threshold is not None
and responses_state_context_tokens(sanitized_state) >= compact_threshold
):
stage = "codex_compaction"
compact_body = {
**body,
"input": [*input_items, {"type": "compaction_trigger"}],
}
try:
compact_result = await _send(compact_body, emit_deltas=False)
compact_items = (
responses_state_items(compact_result.provider_state)
if compact_result.provider_state is not None
else None
)
if not compact_items or compact_items[-1].get("type") not in {
"compaction",
"compaction_summary",
"context_compaction",
}:
raise RuntimeError("Codex compaction returned no compaction item")
body["input"] = [
*_retained_compaction_messages(input_items),
*compact_items,
]
except Exception as compact_error:
if is_compaction_compatibility_error(compact_error):
self._native_compaction_available = False
logger.warning(
"Codex native compaction unavailable; continuing without it "
"(type={} status={} disabled={})",
type(compact_error).__name__,
getattr(compact_error, "status_code", None),
not self._native_compaction_available,
)
stage = "codex_request"
return await _send(body, emit_deltas=True)
except Exception as e:
response = _codex_error_response(e)
exc_type = "CodexHTTPError" if isinstance(e, _CodexHTTPError) else type(e).__name__
@ -137,8 +225,28 @@ class OpenAICodexProvider(LLMProvider):
model: str | None = None, max_tokens: int = 4096, temperature: float = 0.7,
reasoning_effort: str | None = None,
tool_choice: str | dict[str, Any] | None = None,
provider_context: ProviderCallContext | None = None,
) -> LLMResponse:
return await self._call_codex(messages, tools, model, reasoning_effort, tool_choice)
return await self._call_codex(
messages,
tools,
model,
max_tokens,
reasoning_effort,
tool_choice,
provider_context=provider_context,
)
async def chat_with_context(
self,
*,
provider_context: ProviderCallContext,
**kwargs: Any,
) -> LLMResponse:
return await self.chat(
**kwargs,
provider_context=provider_context,
)
async def chat_stream(
self, messages: list[dict[str, Any]], tools: list[dict[str, Any]] | None = None,
@ -148,21 +256,55 @@ class OpenAICodexProvider(LLMProvider):
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
on_thinking_delta: Callable[[str], Awaitable[None]] | None = None,
on_tool_call_delta: Callable[[dict[str, Any]], Awaitable[None]] | None = None,
provider_context: ProviderCallContext | None = None,
) -> LLMResponse:
return await self._call_codex(
messages,
tools,
model,
reasoning_effort,
tool_choice,
on_content_delta,
on_thinking_delta,
on_tool_call_delta,
messages=messages,
tools=tools,
model=model,
max_tokens=max_tokens,
reasoning_effort=reasoning_effort,
tool_choice=tool_choice,
on_content_delta=on_content_delta,
on_thinking_delta=on_thinking_delta,
on_tool_call_delta=on_tool_call_delta,
provider_context=provider_context,
)
async def chat_stream_with_context(
self,
*,
provider_context: ProviderCallContext,
**kwargs: Any,
) -> LLMResponse:
return await self.chat_stream(
**kwargs,
provider_context=provider_context,
)
def get_default_model(self) -> str:
return self.default_model
@staticmethod
def _responses_state_provider() -> str:
return f"openai_codex:{DEFAULT_CODEX_URL.rstrip('/')}"
def can_resume_conversation_state(
self,
state: ProviderConversationState,
model: str | None = None,
) -> bool:
return responses_state_matches(
state,
provider=self._responses_state_provider(),
model=_strip_model_prefix(model or self.default_model),
)
def supports_native_compaction(self, model: str | None = None) -> bool:
"""Use the Codex backend's inline compaction trigger when needed."""
_ = model
return self._native_compaction_available
def _strip_model_prefix(model: str) -> str:
if model.startswith("openai-codex/") or model.startswith("openai_codex/"):
@ -170,6 +312,58 @@ def _strip_model_prefix(model: str) -> str:
return model
def _without_response_item_ids(
request_body: dict[str, Any],
) -> dict[str, Any]:
"""Match Codex's default ``store=false`` request-item contract."""
if request_body.get("store") is True:
return request_body
raw_input = request_body.get("input")
if not isinstance(raw_input, list):
return request_body
input_items: list[object] = cast(list[object], raw_input)
sanitized_input: list[object] = []
for raw_item in input_items:
if not isinstance(raw_item, dict):
sanitized_input.append(raw_item)
continue
item = cast(dict[str, Any], raw_item)
sanitized_input.append({
key: value
for key, value in item.items()
if key != "id"
})
body = dict(request_body)
body["input"] = sanitized_input
return body
def _retained_compaction_messages(
input_items: list[dict[str, Any]],
) -> list[dict[str, Any]]:
"""Mirror Codex's bounded retention of user/developer/system messages."""
retained_reversed: list[dict[str, Any]] = []
remaining = _COMPACTION_RETAINED_CHAR_BUDGET
for item in reversed(input_items):
if item.get("type") not in {None, "message"} or item.get("role") not in {
"user",
"developer",
"system",
}:
continue
size = len(json.dumps(item, ensure_ascii=False))
if size > remaining and retained_reversed:
continue
retained_reversed.append(item)
remaining = max(0, remaining - size)
if remaining == 0:
break
retained_reversed.reverse()
return retained_reversed
def _build_reasoning_options(reasoning_effort: str | None) -> dict[str, str] | None:
"""Opt in to visible summaries without changing provider-default effort."""
if reasoning_effort and reasoning_effort.lower() == "none":
@ -202,6 +396,7 @@ class _CodexHTTPError(RuntimeError):
error_type: str | None = None,
error_code: str | None = None,
should_retry: bool | None = None,
compaction_unsupported: bool = False,
):
super().__init__(message)
self.status_code = status_code
@ -209,6 +404,7 @@ class _CodexHTTPError(RuntimeError):
self.error_type = error_type
self.error_code = error_code
self.should_retry = should_retry
self.compaction_unsupported = compaction_unsupported
async def _request_codex(
@ -220,7 +416,7 @@ async def _request_codex(
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
on_thinking_delta: Callable[[str], Awaitable[None]] | None = None,
on_tool_call_delta: Callable[[dict[str, Any]], Awaitable[None]] | None = None,
) -> tuple[str, list[ToolCallRequest], str, dict[str, int], str | None]:
) -> LLMResponse:
idle_timeout_s = resolve_stream_idle_timeout_s()
client_kwargs: dict[str, Any] = {"timeout": idle_timeout_s, "verify": verify}
if proxy:
@ -233,6 +429,17 @@ async def _request_codex(
raw = text.decode("utf-8", "ignore")
retry_after = LLMProvider._extract_retry_after_from_headers(response.headers)
error_type, error_code = LLMProvider._extract_error_type_code(raw)
compaction_unsupported = (
response.status_code in {400, 404, 422}
and any(
marker in raw.lower()
for marker in (
"context_management",
"compact_threshold",
"compaction_trigger",
)
)
)
raise _CodexHTTPError(
_friendly_error(response.status_code, raw),
status_code=response.status_code,
@ -240,13 +447,38 @@ async def _request_codex(
error_type=error_type,
error_code=error_code,
should_retry=_should_retry_status(response.status_code, error_type, error_code, raw),
compaction_unsupported=compaction_unsupported,
)
return await consume_sse_with_reasoning(
capture = ResponsesStreamCapture()
(
content,
tool_calls,
finish_reason,
usage,
reasoning_content,
) = await consume_sse_with_reasoning(
response,
on_content_delta=on_content_delta,
on_tool_call_delta=on_tool_call_delta,
on_reasoning_delta=on_thinking_delta,
capture=capture,
)
result = LLMResponse(
content=content,
tool_calls=tool_calls,
finish_reason=finish_reason,
usage=usage,
reasoning_content=reasoning_content,
)
if capture.completed and is_replayable_finish_reason(finish_reason):
result.provider_state = build_responses_state(
provider=f"openai_codex:{url.rstrip('/')}",
model=str(body.get("model") or ""),
input_items=cast(list[dict[str, Any]], body.get("input") or []),
output_items=capture.output_items,
usage=usage,
)
return result
def _prompt_cache_key(messages: list[dict[str, Any]]) -> str:

View File

@ -26,16 +26,24 @@ from pydantic.alias_generators import to_snake
from nanobot.providers.base import (
LLMProvider,
LLMResponse,
ProviderCallContext,
ProviderConversationState,
ToolCallRequest,
parse_tool_arguments,
resolve_stream_idle_timeout_s,
tool_arguments_json_for_replay,
)
from nanobot.providers.openai_responses import (
ResponsesStreamCapture,
build_responses_state,
consume_sdk_stream,
convert_messages,
convert_tools,
is_compaction_compatibility_error,
is_replayable_finish_reason,
parse_response_output,
prepare_responses_input,
resolve_compact_threshold,
responses_state_matches,
)
if TYPE_CHECKING:
@ -443,6 +451,8 @@ class OpenAICompatProvider(LLMProvider):
registry lookups needed.
"""
_native_compaction_available = True
def __init__(
self,
api_key: str | None = None,
@ -463,6 +473,7 @@ class OpenAICompatProvider(LLMProvider):
self._api_type = api_type if spec and spec.name == "openai" else "auto"
self._extra_query = extra_query or {}
self._proxy = proxy or None
self._native_compaction_available = True
if api_key and spec and spec.env_key:
self._setup_env(api_key, api_base)
@ -971,6 +982,37 @@ class OpenAICompatProvider(LLMProvider):
return self._responses_circuit_allows_probe(model, reasoning_effort)
def _responses_state_provider(self) -> str:
spec_name = self._spec.name if self._spec is not None else "custom"
effective_base = self._effective_base or "https://api.openai.com/v1"
return f"openai_compat:{spec_name}:{effective_base.rstrip('/')}"
def _responses_state_model(self, model: str | None) -> str:
return self._request_model_name(model or self.default_model)
def can_resume_conversation_state(
self,
state: ProviderConversationState,
model: str | None = None,
) -> bool:
return responses_state_matches(
state,
provider=self._responses_state_provider(),
model=self._responses_state_model(model),
)
def supports_native_compaction(self, model: str | None = None) -> bool:
"""Enable server compaction only on direct OpenAI Responses endpoints."""
_ = model
if (
not self._native_compaction_available
or self._api_type == "chat_completions"
):
return False
if self._spec is not None and self._spec.name != "openai":
return False
return _is_direct_openai_base(self._effective_base)
def _responses_circuit_allows_probe(
self,
model: str | None,
@ -1040,12 +1082,29 @@ class OpenAICompatProvider(LLMProvider):
temperature: float,
reasoning_effort: str | None,
tool_choice: str | dict[str, Any] | None,
provider_context: ProviderCallContext | None = None,
) -> dict[str, Any]:
"""Build a Responses API body for direct OpenAI requests."""
model_name = model or self.default_model
model_name = self._request_model_name(model_name)
sanitized_messages = self._sanitize_messages(self._sanitize_empty_content(messages))
instructions, input_items = convert_messages(sanitized_messages)
sanitized_state = (
provider_context.conversation_state
if provider_context is not None
else None
)
if sanitized_state is not None:
sanitized_state = sanitized_state.with_pending_messages(
self._sanitize_messages(
self._sanitize_empty_content(sanitized_state.pending_messages)
)
)
instructions, input_items, replayed = prepare_responses_input(
sanitized_messages,
state=sanitized_state,
provider=self._responses_state_provider(),
model=model_name,
)
body: dict[str, Any] = {
"model": model_name,
@ -1055,13 +1114,29 @@ class OpenAICompatProvider(LLMProvider):
"store": False,
"stream": False,
}
compact_threshold = resolve_compact_threshold(
(
provider_context.context_window_tokens
if provider_context is not None
else None
),
max_tokens,
)
if self.supports_native_compaction(model_name) and compact_threshold is not None:
body["context_management"] = [{
"type": "compaction",
"compact_threshold": compact_threshold,
}]
if self._supports_temperature(model_name, reasoning_effort):
body["temperature"] = temperature
if not self._supports_temperature(model_name, reasoning_effort):
body["include"] = ["reasoning.encrypted_content"]
if reasoning_effort and reasoning_effort.lower() != "none":
body["reasoning"] = {"effort": reasoning_effort}
body["include"] = ["reasoning.encrypted_content"]
if replayed and "gpt-5.6" in model_name.lower():
body.setdefault("reasoning", {})["context"] = "all_turns"
if tools:
body["tools"] = convert_tools(tools)
@ -1073,6 +1148,29 @@ class OpenAICompatProvider(LLMProvider):
return body
async def _create_response_with_compaction_fallback(
self,
client: Any,
body: dict[str, Any],
) -> Any:
"""Retry Responses once without server compaction on compatibility errors."""
try:
return await client.responses.create(**body)
except Exception as exc:
if (
"context_management" not in body
or not is_compaction_compatibility_error(exc)
):
raise
self._native_compaction_available = False
body.pop("context_management", None)
logger.warning(
"Responses server compaction unsupported; disabled for this provider instance "
"(status={})",
getattr(exc, "status_code", None),
)
return await client.responses.create(**body)
# ------------------------------------------------------------------
# Response parsing
# ------------------------------------------------------------------
@ -1599,6 +1697,28 @@ class OpenAICompatProvider(LLMProvider):
# Public API
# ------------------------------------------------------------------
async def chat_with_context(
self,
*,
provider_context: ProviderCallContext,
**kwargs: Any,
) -> LLMResponse:
return await self.chat(
**kwargs,
provider_context=provider_context,
)
async def chat_stream_with_context(
self,
*,
provider_context: ProviderCallContext,
**kwargs: Any,
) -> LLMResponse:
return await self.chat_stream(
**kwargs,
provider_context=provider_context,
)
async def chat(
self,
messages: list[dict[str, Any]],
@ -1608,6 +1728,7 @@ class OpenAICompatProvider(LLMProvider):
temperature: float = 0.7,
reasoning_effort: str | None = None,
tool_choice: str | dict[str, Any] | None = None,
provider_context: ProviderCallContext | None = None,
) -> LLMResponse:
client = await self._ensure_client()
try:
@ -1616,12 +1737,18 @@ class OpenAICompatProvider(LLMProvider):
body = self._build_responses_body(
messages, tools, model, max_tokens, temperature,
reasoning_effort, tool_choice,
provider_context,
)
responses_raw = cast(
Any,
await client.responses.create(**body),
responses_raw = await self._create_response_with_compaction_fallback(
client,
body,
)
result = parse_response_output(
responses_raw,
state_provider=self._responses_state_provider(),
state_model=str(body["model"]),
state_input_items=cast(list[dict[str, Any]], body["input"]),
)
result = parse_response_output(responses_raw)
self._record_responses_success(model, reasoning_effort)
return result
except Exception as responses_error:
@ -1660,6 +1787,7 @@ class OpenAICompatProvider(LLMProvider):
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
on_thinking_delta: Callable[[str], Awaitable[None]] | None = None,
on_tool_call_delta: Callable[[dict[str, Any]], Awaitable[None]] | None = None,
provider_context: ProviderCallContext | None = None,
) -> LLMResponse:
client = await self._ensure_client()
idle_timeout_s = resolve_stream_idle_timeout_s()
@ -1669,11 +1797,12 @@ class OpenAICompatProvider(LLMProvider):
body = self._build_responses_body(
messages, tools, model, max_tokens, temperature,
reasoning_effort, tool_choice,
provider_context,
)
body["stream"] = True
responses_stream = cast(
Any,
await client.responses.create(**body),
responses_stream = await self._create_response_with_compaction_fallback(
client,
body,
)
async def _timed_stream() -> AsyncIterator[Any]:
@ -1687,6 +1816,7 @@ class OpenAICompatProvider(LLMProvider):
except StopAsyncIteration:
break
capture = ResponsesStreamCapture()
(
content,
tool_calls,
@ -1697,15 +1827,25 @@ class OpenAICompatProvider(LLMProvider):
_timed_stream(),
on_content_delta,
on_tool_call_delta=on_tool_call_delta,
capture=capture,
)
self._record_responses_success(model, reasoning_effort)
return LLMResponse(
result = LLMResponse(
content=content or None,
tool_calls=tool_calls,
finish_reason=finish_reason,
usage=usage,
reasoning_content=reasoning_content,
)
if capture.completed and is_replayable_finish_reason(finish_reason):
result.provider_state = build_responses_state(
provider=self._responses_state_provider(),
model=str(body["model"]),
input_items=cast(list[dict[str, Any]], body["input"]),
output_items=capture.output_items,
usage=usage,
)
return result
except Exception as responses_error:
if self._spec and self._spec.name == "github_copilot":
# Copilot gateway exposes GPT-5/o-series only via /responses;

View File

@ -1,4 +1,4 @@
"""Shared helpers for OpenAI Responses API providers (Codex, Azure OpenAI)."""
"""Shared helpers for provider backends that implement the OpenAI Responses protocol."""
from nanobot.providers.openai_responses.converters import (
convert_messages,
@ -8,13 +8,24 @@ from nanobot.providers.openai_responses.converters import (
)
from nanobot.providers.openai_responses.parsing import (
FINISH_REASON_MAP,
ResponsesStreamCapture,
consume_sdk_stream,
consume_sse,
consume_sse_with_reasoning,
is_replayable_finish_reason,
iter_sse,
map_finish_reason,
parse_response_output,
)
from nanobot.providers.openai_responses.state import (
build_responses_state,
is_compaction_compatibility_error,
prepare_responses_input,
resolve_compact_threshold,
responses_state_context_tokens,
responses_state_items,
responses_state_matches,
)
__all__ = [
"convert_messages",
@ -25,7 +36,16 @@ __all__ = [
"consume_sse",
"consume_sse_with_reasoning",
"consume_sdk_stream",
"ResponsesStreamCapture",
"is_replayable_finish_reason",
"map_finish_reason",
"parse_response_output",
"build_responses_state",
"is_compaction_compatibility_error",
"prepare_responses_input",
"resolve_compact_threshold",
"responses_state_context_tokens",
"responses_state_items",
"responses_state_matches",
"FINISH_REASON_MAP",
]

View File

@ -4,12 +4,14 @@ from __future__ import annotations
import json
from collections.abc import Awaitable, Callable
from dataclasses import dataclass, field
from typing import Any, AsyncGenerator, cast
import httpx
from loguru import logger
from nanobot.providers.base import LLMResponse, ToolCallRequest, parse_tool_arguments
from nanobot.providers.openai_responses.state import build_responses_state
FINISH_REASON_MAP = {
"completed": "stop",
@ -17,6 +19,42 @@ FINISH_REASON_MAP = {
"failed": "error",
"cancelled": "error",
}
REPLAYABLE_FINISH_REASONS = frozenset({"stop", "tool_calls", "function_call"})
@dataclass(slots=True)
class ResponsesStreamCapture:
"""Losslessly capture terminal output items without changing stream results."""
completed: bool = False
response: dict[str, Any] | None = field(default=None, repr=False)
_items_by_index: dict[int, dict[str, Any]] = field(default_factory=dict, repr=False)
def record_output_item(self, index: object, item: object) -> None:
item_object = _response_object(item)
if item_object is None:
return
output_index = (
index
if isinstance(index, int) and not isinstance(index, bool)
else len(self._items_by_index)
)
self._items_by_index[output_index] = item_object
def record_completed(self, response: object) -> None:
response_object = _response_object(response)
if response_object is None:
return
self.completed = True
self.response = response_object
@property
def output_items(self) -> list[dict[str, Any]]:
if self.response is not None:
output = _response_object_list(self.response.get("output"))
if output:
return output
return [self._items_by_index[index] for index in sorted(self._items_by_index)]
def _as_json_object(value: object) -> dict[str, Any] | None:
@ -54,6 +92,27 @@ def map_finish_reason(status: str | None) -> str:
return FINISH_REASON_MAP.get(status or "completed", "stop")
def is_replayable_finish_reason(finish_reason: str) -> bool:
"""Return whether a response can safely advance opaque conversation state."""
return finish_reason in REPLAYABLE_FINISH_REASONS
def _response_finish_reason(
response: object,
*,
fallback_status: str | None = None,
) -> str:
"""Map terminal response details without treating content filtering as truncation."""
response_object = _response_object(response) or {}
status = response_object.get("status")
terminal_status = status if isinstance(status, str) else fallback_status
if terminal_status == "incomplete":
details = _response_object(response_object.get("incomplete_details"))
if details is not None and details.get("reason") == "content_filter":
return "content_filter"
return map_finish_reason(terminal_status)
def _usage_from_response_obj(response: object) -> dict[str, int]:
response_object = _response_object(response)
usage_raw: object = (
@ -99,6 +158,47 @@ def _tool_arguments_source(*values: Any) -> Any:
return "{}"
def _refusal_event_key(
item_id: object,
content_index: object,
) -> tuple[str | None, int | None]:
"""Identify one streamed refusal content part across delta/done events."""
return (
item_id if isinstance(item_id, str) else None,
(
content_index
if isinstance(content_index, int) and not isinstance(content_index, bool)
else None
),
)
def _remaining_refusal_text(streamed_text: str, refusal_text: str) -> str:
"""Return only text not already surfaced by refusal deltas."""
if not streamed_text:
return refusal_text
if refusal_text.startswith(streamed_text):
return refusal_text[len(streamed_text):]
return ""
def _extract_refusal_text_from_output(output: object) -> tuple[bool, str]:
"""Extract refusal content from terminal Responses output items."""
refusal_seen = False
parts: list[str] = []
for item in _response_object_list(output):
if item.get("type") != "message":
continue
for block in _response_object_list(item.get("content")):
if block.get("type") != "refusal":
continue
refusal_seen = True
refusal_text = block.get("refusal")
if isinstance(refusal_text, str):
parts.append(refusal_text)
return refusal_seen, "".join(parts)
async def iter_sse(response: httpx.Response) -> AsyncGenerator[dict[str, Any], None]:
"""Yield parsed JSON events from a Responses API SSE stream."""
buffer: list[str] = []
@ -153,6 +253,7 @@ async def consume_sse_with_reasoning(
on_tool_call_delta: Callable[[dict[str, Any]], Awaitable[None]] | None = None,
on_reasoning_delta: Callable[[str], Awaitable[None]] | None = None,
on_response_event: Callable[[dict[str, Any]], Awaitable[None]] | None = None,
capture: ResponsesStreamCapture | None = None,
) -> tuple[str, list[ToolCallRequest], str, dict[str, int], str | None]:
"""Consume a Responses API SSE stream, including visible reasoning summaries."""
content = ""
@ -163,6 +264,9 @@ async def consume_sse_with_reasoning(
usage: dict[str, int] = {}
reasoning_content: str | None = None
streamed_reasoning = False
refusal_seen = False
refusal_deltas: dict[tuple[str | None, int | None], str] = {}
emitted_refusal_text = ""
async for event in iter_sse(response):
if on_response_event:
@ -191,6 +295,33 @@ async def consume_sse_with_reasoning(
content += delta_text
if on_content_delta and delta_text:
await on_content_delta(delta_text)
elif event_type == "response.refusal.delta":
refusal_seen = True
delta_text = event.get("delta")
if isinstance(delta_text, str) and delta_text:
key = _refusal_event_key(
event.get("item_id"),
event.get("content_index"),
)
refusal_deltas[key] = refusal_deltas.get(key, "") + delta_text
content += delta_text
emitted_refusal_text += delta_text
if on_content_delta:
await on_content_delta(delta_text)
elif event_type == "response.refusal.done":
refusal_seen = True
refusal_text = event.get("refusal")
key = _refusal_event_key(
event.get("item_id"),
event.get("content_index"),
)
streamed_text = refusal_deltas.pop(key, "")
if isinstance(refusal_text, str) and refusal_text:
remaining_text = _remaining_refusal_text(streamed_text, refusal_text)
content += remaining_text
emitted_refusal_text += remaining_text
if on_content_delta and remaining_text:
await on_content_delta(remaining_text)
elif event_type == "response.reasoning_summary_text.delta":
delta_text = event.get("delta") or ""
if delta_text:
@ -239,6 +370,8 @@ async def consume_sse_with_reasoning(
})
elif event_type == "response.output_item.done":
item = _as_json_object(event.get("item")) or {}
if capture is not None:
capture.record_output_item(event.get("output_index"), item)
if item.get("type") == "function_call":
call_id = item.get("call_id")
if not call_id:
@ -269,11 +402,28 @@ async def consume_sse_with_reasoning(
reasoning_content = summary
if on_reasoning_delta:
await on_reasoning_delta(summary)
elif event_type == "response.completed":
elif event_type in {"response.completed", "response.incomplete"}:
response_obj = _response_object(event.get("response")) or {}
status = response_obj.get("status")
finish_reason = map_finish_reason(status)
if capture is not None:
capture.record_completed(response_obj)
finish_reason = _response_finish_reason(
response_obj,
fallback_status=event_type.removeprefix("response."),
)
usage = _usage_from_response_obj(response_obj) or usage
terminal_refusal, terminal_refusal_text = _extract_refusal_text_from_output(
response_obj.get("output")
)
if terminal_refusal:
refusal_seen = True
remaining_text = _remaining_refusal_text(
emitted_refusal_text,
terminal_refusal_text,
)
content += remaining_text
emitted_refusal_text += remaining_text
if on_content_delta and remaining_text:
await on_content_delta(remaining_text)
if not reasoning_content:
summary = _extract_reasoning_summary_from_output(response_obj.get("output"))
if summary:
@ -284,6 +434,8 @@ async def consume_sse_with_reasoning(
detail = event.get("error") or event.get("message") or event
raise RuntimeError(f"Response failed: {str(detail)[:500]}")
if refusal_seen:
finish_reason = "refusal"
return content, tool_calls, finish_reason, usage, reasoning_content
@ -300,7 +452,13 @@ def _extract_reasoning_summary_from_output(output: object) -> str | None:
return "".join(parts) or None
def parse_response_output(response: object) -> LLMResponse:
def parse_response_output(
response: object,
*,
state_provider: str | None = None,
state_model: str | None = None,
state_input_items: list[dict[str, Any]] | None = None,
) -> LLMResponse:
"""Parse an SDK ``Response`` object into an ``LLMResponse``."""
response_object = _response_object(response) or {}
@ -308,15 +466,22 @@ def parse_response_output(response: object) -> LLMResponse:
content_parts: list[str] = []
tool_calls: list[ToolCallRequest] = []
reasoning_content: str | None = None
refusal_seen = False
for item in output:
item_type = item.get("type")
if item_type == "message":
for block in _response_object_list(item.get("content")):
if block.get("type") == "output_text":
block_type = block.get("type")
if block_type == "output_text":
text = block.get("text")
if isinstance(text, str):
content_parts.append(text)
elif block_type == "refusal":
refusal_seen = True
refusal = block.get("refusal")
if isinstance(refusal, str):
content_parts.append(refusal)
elif item_type == "reasoning":
for s in _response_object_list(item.get("summary")):
if s.get("type") == "summary_text" and s.get("text"):
@ -337,21 +502,37 @@ def parse_response_output(response: object) -> LLMResponse:
usage = _usage_from_response_obj(response_object)
status = response_object.get("status")
finish_reason = map_finish_reason(status if isinstance(status, str) else None)
finish_reason = "refusal" if refusal_seen else _response_finish_reason(response_object)
return LLMResponse(
result = LLMResponse(
content="".join(content_parts) or None,
tool_calls=tool_calls,
finish_reason=finish_reason,
usage=usage,
reasoning_content=reasoning_content if isinstance(reasoning_content, str) else None,
)
if (
state_provider is not None
and state_model is not None
and state_input_items is not None
and (status is None or status == "completed")
and is_replayable_finish_reason(finish_reason)
):
result.provider_state = build_responses_state(
provider=state_provider,
model=state_model,
input_items=state_input_items,
output_items=output,
usage=usage,
)
return result
async def consume_sdk_stream(
stream: Any,
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
on_tool_call_delta: Callable[[dict[str, Any]], Awaitable[None]] | None = None,
capture: ResponsesStreamCapture | None = None,
) -> tuple[str, list[ToolCallRequest], str, dict[str, int], str | None]:
"""Consume an SDK async stream from ``client.responses.create(stream=True)``."""
content = ""
@ -361,6 +542,9 @@ async def consume_sdk_stream(
finish_reason = "stop"
usage: dict[str, int] = {}
reasoning_content: str | None = None
refusal_seen = False
refusal_deltas: dict[tuple[str | None, int | None], str] = {}
emitted_refusal_text = ""
async for raw_event in stream:
event: Any = raw_event
@ -388,6 +572,33 @@ async def consume_sdk_stream(
content += delta_text
if on_content_delta and delta_text:
await on_content_delta(delta_text)
elif event_type == "response.refusal.delta":
refusal_seen = True
delta_text = getattr(event, "delta", None)
if isinstance(delta_text, str) and delta_text:
key = _refusal_event_key(
getattr(event, "item_id", None),
getattr(event, "content_index", None),
)
refusal_deltas[key] = refusal_deltas.get(key, "") + delta_text
content += delta_text
emitted_refusal_text += delta_text
if on_content_delta:
await on_content_delta(delta_text)
elif event_type == "response.refusal.done":
refusal_seen = True
refusal_text = getattr(event, "refusal", None)
key = _refusal_event_key(
getattr(event, "item_id", None),
getattr(event, "content_index", None),
)
streamed_text = refusal_deltas.pop(key, "")
if isinstance(refusal_text, str) and refusal_text:
remaining_text = _remaining_refusal_text(streamed_text, refusal_text)
content += remaining_text
emitted_refusal_text += remaining_text
if on_content_delta and remaining_text:
await on_content_delta(remaining_text)
elif event_type == "response.function_call_arguments.delta":
call_id = getattr(event, "call_id", None)
if call_id and call_id in tool_call_buffers:
@ -416,6 +627,8 @@ async def consume_sdk_stream(
})
elif event_type == "response.output_item.done":
item = getattr(event, "item", None)
if capture is not None:
capture.record_output_item(getattr(event, "output_index", None), item)
if item and getattr(item, "type", None) == "function_call":
call_id = getattr(item, "call_id", None)
if not call_id:
@ -443,10 +656,31 @@ async def consume_sdk_stream(
arguments=args,
)
)
elif event_type == "response.completed":
elif event_type in {"response.completed", "response.incomplete"}:
resp = getattr(event, "response", None)
status = getattr(resp, "status", None) if resp else None
finish_reason = map_finish_reason(status)
response_obj = _response_object(resp) or {}
if capture is not None:
capture.record_completed(resp)
finish_reason = _response_finish_reason(
resp,
fallback_status=event_type.removeprefix("response."),
)
terminal_output = response_obj.get("output")
if terminal_output is None:
terminal_output = getattr(resp, "output", None)
terminal_refusal, terminal_refusal_text = _extract_refusal_text_from_output(
terminal_output
)
if terminal_refusal:
refusal_seen = True
remaining_text = _remaining_refusal_text(
emitted_refusal_text,
terminal_refusal_text,
)
content += remaining_text
emitted_refusal_text += remaining_text
if on_content_delta and remaining_text:
await on_content_delta(remaining_text)
if resp:
usage_obj = getattr(resp, "usage", None)
if usage_obj:
@ -466,4 +700,6 @@ async def consume_sdk_stream(
detail = getattr(event, "error", None) or getattr(event, "message", None) or event
raise RuntimeError(f"Response failed: {str(detail)[:500]}")
if refusal_seen:
finish_reason = "refusal"
return content, tool_calls, finish_reason, usage, reasoning_content

View File

@ -0,0 +1,197 @@
"""Opaque conversation state for Responses API item replay."""
from __future__ import annotations
from copy import deepcopy
from typing import Any, cast
from loguru import logger
from nanobot.providers.base import ProviderConversationState
from nanobot.providers.openai_responses.converters import convert_messages
RESPONSES_STATE_KIND = "openai_responses"
RESPONSES_STATE_VERSION = 1
_ITEMS_KEY = "items"
_CONTEXT_TOKENS_KEY = "context_tokens"
_COMPACTION_ITEM_TYPES = frozenset({
"compaction",
"compaction_summary",
"context_compaction",
})
def responses_state_matches(
state: ProviderConversationState,
*,
provider: str,
model: str,
) -> bool:
"""Return whether *state* belongs to this exact Responses endpoint/model."""
return (
state.kind == RESPONSES_STATE_KIND
and state.version == RESPONSES_STATE_VERSION
and state.provider == provider
and state.model == model
and _state_items(state) is not None
)
def prepare_responses_input(
messages: list[dict[str, Any]],
*,
state: ProviderConversationState | None,
provider: str,
model: str,
) -> tuple[str, list[dict[str, Any]], bool]:
"""Build a request from exact prior items plus only newly appended messages.
The full Chat transcript remains the source for the current instructions.
When no compatible state exists, it is converted normally as a safe
fallback.
"""
instructions, fallback_items = convert_messages(messages)
if state is None or not responses_state_matches(
state,
provider=provider,
model=model,
):
return instructions, fallback_items, False
prior_items = _state_items(state)
if prior_items is None:
return instructions, fallback_items, False
_, delta_items = convert_messages(state.pending_messages)
logger.debug(
"Replaying Responses state: prior_items={} pending_messages={}",
len(prior_items),
len(state.pending_messages),
)
return instructions, [*deepcopy(prior_items), *delta_items], True
def build_responses_state(
*,
provider: str,
model: str,
input_items: list[dict[str, Any]],
output_items: list[dict[str, Any]],
usage: dict[str, int] | None = None,
) -> ProviderConversationState:
"""Create the canonical next state from request input and every output item."""
unpruned_items = [*input_items, *output_items]
items = _prune_before_latest_output_compaction(input_items, output_items)
if len(items) < len(unpruned_items):
logger.info(
"Installed Responses compaction: dropped_items={} retained_items={}",
len(unpruned_items) - len(items),
len(items),
)
payload: dict[str, Any] = {_ITEMS_KEY: deepcopy(items)}
context_tokens = _context_tokens_from_usage(usage)
if context_tokens > 0:
payload[_CONTEXT_TOKENS_KEY] = context_tokens
return ProviderConversationState(
kind=RESPONSES_STATE_KIND,
provider=provider,
model=model,
version=RESPONSES_STATE_VERSION,
payload=payload,
)
def responses_state_items(
state: ProviderConversationState,
) -> list[dict[str, Any]] | None:
"""Return an isolated copy of canonical input items for tests/consumers."""
items = _state_items(state)
return deepcopy(items) if items is not None else None
def responses_state_context_tokens(state: ProviderConversationState) -> int:
"""Return the last server-reported active context size."""
value = state.payload.get(_CONTEXT_TOKENS_KEY)
if isinstance(value, bool) or not isinstance(value, int):
return 0
return max(0, value)
def resolve_compact_threshold(
context_window_tokens: int | None,
max_output_tokens: int,
) -> int | None:
"""Derive Codex-compatible 90% compaction headroom for a model window."""
if context_window_tokens is None or context_window_tokens <= 0:
return None
ninety_percent = max(1, context_window_tokens * 9 // 10)
output_headroom = max(1, context_window_tokens - max(1, max_output_tokens))
return min(ninety_percent, output_headroom)
def is_compaction_compatibility_error(exc: Exception) -> bool:
"""Recognize endpoints that reject native Responses compaction fields."""
if getattr(exc, "compaction_unsupported", False) is True:
return True
response = getattr(exc, "response", None)
status_code = getattr(exc, "status_code", None)
if status_code is None and response is not None:
status_code = getattr(response, "status_code", None)
body = (
getattr(exc, "body", None)
or getattr(exc, "doc", None)
or getattr(response, "text", None)
or str(exc)
)
text = str(body).lower()
has_compaction_marker = any(
marker in text
for marker in ("context_management", "compact_threshold", "compaction_trigger")
)
if not has_compaction_marker:
return False
return isinstance(exc, TypeError) or status_code in {400, 404, 422}
def _prune_before_latest_output_compaction(
input_items: list[dict[str, Any]],
output_items: list[dict[str, Any]],
) -> list[dict[str, Any]]:
"""Drop old input only when this response emits a new compaction item.
A canonical compacted input may intentionally retain messages before its
compaction item. Those messages must survive ordinary subsequent responses.
"""
latest = None
for index, item in enumerate(output_items):
if item.get("type") in _COMPACTION_ITEM_TYPES:
latest = index
if latest is None:
return [*input_items, *output_items]
return output_items[latest:]
def _context_tokens_from_usage(usage: dict[str, int] | None) -> int:
if not usage:
return 0
prompt_tokens = usage.get("prompt_tokens", 0)
completion_tokens = usage.get("completion_tokens", 0)
total_tokens = usage.get("total_tokens", 0)
values = (prompt_tokens, completion_tokens, total_tokens)
if any(isinstance(value, bool) for value in values):
return 0
return max(0, total_tokens or prompt_tokens + completion_tokens)
def _state_items(
state: ProviderConversationState,
) -> list[dict[str, Any]] | None:
raw_items = state.payload.get(_ITEMS_KEY)
if not isinstance(raw_items, list):
return None
items: list[dict[str, Any]] = []
for raw in cast(list[object], raw_items):
if not isinstance(raw, dict):
return None
items.append(cast(dict[str, Any], raw))
return items

View File

@ -17,6 +17,7 @@ from weakref import WeakValueDictionary
from loguru import logger
from nanobot.config.paths import get_legacy_sessions_dir
from nanobot.providers.base import ProviderConversationState
from nanobot.runtime_context import (
RUNTIME_CONTEXT_HISTORY_META,
public_history_message,
@ -43,6 +44,10 @@ _SESSION_PREVIEW_MAX_CHARS = 120
_SESSION_LIST_PREVIEW_MAX_RECORDS = 200
_SESSION_LIST_PREVIEW_MAX_CHARS = 1_000_000
_SESSION_DATA_ERRORS = (ValueError, TypeError, AttributeError, KeyError)
_PROVIDER_STATE_RECORD_TYPE = "provider_state"
_PROVIDER_STATE_RECORD_PREFIX_RE = re.compile(
r'^\s*\{\s*"_type"\s*:\s*"provider_state"\s*(?:,|\})'
)
_FORK_VOLATILE_METADATA_KEYS = {
"goal_state",
"pending_user_turn",
@ -60,6 +65,11 @@ def _json_object(value: object) -> dict[str, Any]:
return cast(dict[str, Any], value)
def _is_provider_state_record_line(line: str) -> bool:
"""Recognize the canonical private record without decoding its opaque payload."""
return _PROVIDER_STATE_RECORD_PREFIX_RE.match(line) is not None
def replay_max_messages_for_context(context_window_tokens: int | None) -> int:
if not context_window_tokens or context_window_tokens <= 0:
return FILE_MAX_MESSAGES
@ -146,10 +156,13 @@ class Session:
updated_at: datetime = field(default_factory=datetime.now)
metadata: dict[str, Any] = field(default_factory=dict)
last_consolidated: int = 0 # Number of messages already consolidated to files
provider_state: ProviderConversationState | None = field(default=None, repr=False)
def __post_init__(self) -> None:
if not isinstance(cast(object, self.metadata), dict):
self.metadata = {}
if not isinstance(cast(object, self.provider_state), ProviderConversationState):
self.provider_state = None
# An out-of-range offset (corrupt metadata) would hide all history; reset it.
last_consolidated = cast(object, self.last_consolidated)
if (
@ -304,6 +317,7 @@ class Session:
"""Clear all messages and reset session to initial state."""
self.messages = []
self.last_consolidated = 0
self.provider_state = None
self.updated_at = datetime.now()
self.metadata.pop("_last_summary", None)
@ -396,6 +410,8 @@ class Session:
self.messages = retained
self.last_consolidated = new_lc
if dropped:
self.provider_state = None
self.updated_at = datetime.now()
return RetentionResult(
dropped=dropped,
@ -517,6 +533,7 @@ class JsonlSessionStore:
created_at: datetime | None = None
updated_at: datetime | None = None
last_consolidated = 0
provider_state: ProviderConversationState | None = None
with open(path, encoding="utf-8") as f:
for line in f:
@ -527,7 +544,8 @@ class JsonlSessionStore:
raw_data: object = json.loads(line)
data = _json_object(raw_data)
if data.get("_type") == "metadata":
record_type = data.get("_type")
if record_type == "metadata":
metadata_value = cast(object, data.get("metadata", {}))
metadata = (
cast(dict[str, Any], metadata_value)
@ -552,6 +570,10 @@ class JsonlSessionStore:
if isinstance(offset, int) and not isinstance(offset, bool)
else 0
)
elif record_type == _PROVIDER_STATE_RECORD_TYPE:
provider_state = ProviderConversationState.from_private_record(
data.get("state")
)
else:
messages.append(data)
@ -562,6 +584,7 @@ class JsonlSessionStore:
updated_at=updated_at or datetime.now(),
metadata=metadata,
last_consolidated=last_consolidated,
provider_state=provider_state,
)
except _SESSION_DATA_ERRORS as e:
logger.warning("Failed to load session {}: {}", key, e)
@ -586,6 +609,7 @@ class JsonlSessionStore:
created_at: datetime | None = None
updated_at: datetime | None = None
last_consolidated = 0
provider_state: ProviderConversationState | None = None
skipped = 0
with open(path, encoding="utf-8") as f:
@ -603,7 +627,8 @@ class JsonlSessionStore:
continue
data = cast(dict[str, Any], raw_data)
if data.get("_type") == "metadata":
record_type = data.get("_type")
if record_type == "metadata":
metadata_value = cast(object, data.get("metadata", {}))
metadata = (
cast(dict[str, Any], metadata_value)
@ -624,13 +649,21 @@ class JsonlSessionStore:
if isinstance(offset, int) and not isinstance(offset, bool)
else 0
)
elif record_type == _PROVIDER_STATE_RECORD_TYPE:
candidate = ProviderConversationState.from_private_record(
data.get("state")
)
if candidate is None:
skipped += 1
else:
provider_state = candidate
else:
messages.append(data)
if skipped:
logger.warning("Skipped {} corrupt lines in session {}", skipped, key)
if not messages and not metadata:
if not messages and not metadata and provider_state is None:
return None
return Session(
@ -640,6 +673,7 @@ class JsonlSessionStore:
updated_at=updated_at or datetime.now(),
metadata=metadata,
last_consolidated=last_consolidated,
provider_state=provider_state,
)
except _SESSION_DATA_ERRORS as e:
logger.warning("Repair failed for session {}: {}", key, e)
@ -670,6 +704,12 @@ class JsonlSessionStore:
"last_consolidated": session.last_consolidated,
}
f.write(json.dumps(metadata_line, ensure_ascii=False) + "\n")
if session.provider_state is not None:
provider_state_line = {
"_type": _PROVIDER_STATE_RECORD_TYPE,
"state": session.provider_state.to_private_record(),
}
f.write(json.dumps(provider_state_line, ensure_ascii=False) + "\n")
for msg in session.messages:
f.write(json.dumps(msg, ensure_ascii=False) + "\n")
if fsync:
@ -726,7 +766,8 @@ class JsonlSessionStore:
continue
raw_data: object = json.loads(line)
data = _json_object(raw_data)
if data.get("_type") == "metadata":
record_type = data.get("_type")
if record_type == "metadata":
metadata_value = cast(object, data.get("metadata", {}))
metadata = (
cast(dict[str, Any], metadata_value)
@ -745,6 +786,8 @@ class JsonlSessionStore:
stored_key = (
stored_key_value if isinstance(stored_key_value, str) else None
)
elif record_type == _PROVIDER_STATE_RECORD_TYPE:
continue
else:
messages.append(data)
return {
@ -837,6 +880,8 @@ class JsonlSessionStore:
for line in f:
if not line.strip():
continue
if _is_provider_state_record_line(line):
continue
scanned_records += 1
scanned_chars += len(line)
if (
@ -846,7 +891,10 @@ class JsonlSessionStore:
break
raw_item: object = json.loads(line)
item = _json_object(raw_item)
if item.get("_type") == "metadata":
if item.get("_type") in {
"metadata",
_PROVIDER_STATE_RECORD_TYPE,
}:
continue
text = _message_preview_text(item)
if not text:

View File

@ -18,10 +18,12 @@ from loguru import logger
from nanobot.config.paths import get_webui_dir
from nanobot.session.history_visibility import is_hidden_history_message
from nanobot.session.manager import (
_PROVIDER_STATE_RECORD_TYPE, # pyright: ignore[reportPrivateUsage]
_SESSION_LIST_PREVIEW_MAX_CHARS, # pyright: ignore[reportPrivateUsage]
_SESSION_LIST_PREVIEW_MAX_RECORDS, # pyright: ignore[reportPrivateUsage]
Session,
SessionManager,
_is_provider_state_record_line, # pyright: ignore[reportPrivateUsage]
_message_preview_text, # pyright: ignore[reportPrivateUsage]
_metadata_title, # pyright: ignore[reportPrivateUsage]
)
@ -298,7 +300,11 @@ def _scan_session_row(session_manager: SessionManager, path: Path) -> dict[str,
for line in f:
if not line.strip():
continue
if _is_provider_state_record_line(line):
continue
item = json.loads(line)
if item.get("_type") == _PROVIDER_STATE_RECORD_TYPE:
continue
timestamp = _visible_message_timestamp(item)
if timestamp is not None:
visible_message_at = _latest_updated_at(visible_message_at, timestamp)

View File

@ -10,7 +10,11 @@ from nanobot.agent.memory import (
Consolidator,
MemoryStore,
)
from nanobot.providers.base import GenerationSettings, LLMResponse
from nanobot.providers.base import (
GenerationSettings,
LLMResponse,
ProviderConversationState,
)
from nanobot.runtime_context import (
RUNTIME_CONTEXT_HISTORY_META,
RuntimeContextBlock,
@ -74,6 +78,16 @@ def _tool_round(call_id: str) -> list[dict]:
]
def _provider_state() -> ProviderConversationState:
return ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="test-model",
version=1,
payload={"items": []},
)
class TestConsolidatorSummarize:
async def test_archive_prompt_includes_media_breadcrumb(
self, consolidator, mock_provider, store, runtime
@ -385,6 +399,7 @@ class TestConsolidatorTokenBudget:
"""Old messages that cannot be replayed should be materialized first."""
consolidator._SAFETY_BUFFER = 0
session = Session(key="test:replay-overflow")
session.provider_state = _provider_state()
for i in range(10):
session.add_message("user", f"u{i}")
session.add_message("assistant", f"a{i}")
@ -404,6 +419,7 @@ class TestConsolidatorTokenBudget:
assert archived_chunk[-1]["content"] == "a6"
assert session.last_consolidated == 14
assert session.metadata["_last_summary"]["text"] == "old conversation summary"
assert session.provider_state is None
consolidator.sessions.save.assert_called()
async def test_replay_window_overflow_extends_to_long_recent_user_turn(
@ -479,6 +495,7 @@ class TestConsolidatorTokenBudget:
session = MagicMock()
session.last_consolidated = 0
session.key = "test:key"
session.provider_state = _provider_state()
session.messages = [
{
"role": "user" if i in {0, 50, 61} else "assistant",
@ -500,6 +517,7 @@ class TestConsolidatorTokenBudget:
# pick_consolidation_boundary returns (50, tokens) — user turn at idx 50
assert archived_chunk[0]["content"] == "m0"
assert session.last_consolidated > 0
assert session.provider_state is None
async def test_raw_archive_fallback_advances_last_consolidated(
self, consolidator, runtime
@ -610,6 +628,7 @@ class TestCompactIdleSession:
)
sessions = real_consolidator.sessions
session = sessions.get_or_create("cli:test")
session.provider_state = _provider_state()
old_ts = session.updated_at
for i in range(20):
session.add_message("user", f"user msg {i}")
@ -627,6 +646,7 @@ class TestCompactIdleSession:
assert len(reloaded.messages) == 40
assert reloaded.messages[0]["content"] == "user msg 0"
assert reloaded.last_consolidated == 32
assert reloaded.provider_state is None
visible = reloaded.get_history(max_messages=40)
assert len(visible) == 8
assert visible[0]["content"] == "user msg 16"

View File

@ -452,6 +452,20 @@ class TestBuildMessages:
assert "previous user message" in str(messages[1]["content"])
assert "new message" in str(messages[1]["content"])
def test_current_message_can_be_built_without_history_merge(self, tmp_path):
builder = _builder(tmp_path)
current = builder.build_current_message(
"new message",
runtime_context_blocks=[
RuntimeContextBlock(source="test", content="fresh context"),
],
)
assert current["role"] == "user"
assert "new message" in current["content"]
assert "fresh context" in current["content"]
assert current["_meta"]["runtime_context"]["sources"] == ["test"]
def test_different_role_appended(self, tmp_path):
builder = _builder(tmp_path)
history = [{"role": "assistant", "content": "previous response"}]

View File

@ -1,4 +1,5 @@
import asyncio
import json
from pathlib import Path
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock
@ -19,7 +20,7 @@ from nanobot.bus.outbound_events import (
)
from nanobot.bus.queue import MessageBus
from nanobot.cron.session_turns import CRON_HISTORY_META, CRON_TRIGGER_META
from nanobot.providers.base import LLMResponse
from nanobot.providers.base import LLMProvider, LLMResponse, ProviderConversationState
from nanobot.providers.factory import ProviderSnapshot
from nanobot.runtime_context import (
RUNTIME_CONTEXT_HISTORY_META,
@ -59,6 +60,16 @@ def _mk_loop() -> AgentLoop:
return loop
def _provider_state() -> ProviderConversationState:
return ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="test-model",
version=1,
payload={"items": []},
)
def _runtime_message(content, blocks: list[RuntimeContextBlock]) -> dict:
merged, marker = append_runtime_context(content, blocks)
assert marker is not None
@ -494,6 +505,7 @@ def test_restore_runtime_checkpoint_rehydrates_completed_and_pending_tools() ->
loop = _mk_loop()
session = Session(
key="test:checkpoint",
provider_state=_provider_state(),
metadata={
AgentLoop._RUNTIME_CHECKPOINT_KEY: {
"assistant_message": {
@ -539,6 +551,104 @@ def test_restore_runtime_checkpoint_rehydrates_completed_and_pending_tools() ->
assert session.messages[1]["tool_call_id"] == "call_done"
assert session.messages[2]["tool_call_id"] == "call_pending"
assert "interrupted before this tool finished" in session.messages[2]["content"].lower()
assert session.provider_state is None
def test_restore_final_response_checkpoint_preserves_matching_provider_state() -> None:
loop = _mk_loop()
state = _provider_state()
session = Session(
key="test:final-checkpoint",
provider_state=state,
metadata={
AgentLoop._RUNTIME_CHECKPOINT_KEY: {
"phase": "final_response",
AgentLoop._PROVIDER_STATE_CHECKPOINT_VERSION_KEY: (
AgentLoop._PROVIDER_STATE_CHECKPOINT_VERSION
),
"assistant_message": {
"role": "assistant",
"content": "finished",
},
"completed_tool_results": [],
"pending_tool_calls": [],
}
},
)
restored = loop._restore_runtime_checkpoint(session)
assert restored is True
assert session.messages[-1]["content"] == "finished"
assert session.provider_state is state
assert session.metadata.get(AgentLoop._RUNTIME_CHECKPOINT_KEY) is None
def test_restore_legacy_final_checkpoint_discards_unproven_provider_state() -> None:
loop = _mk_loop()
session = Session(
key="test:legacy-final-checkpoint",
provider_state=_provider_state(),
metadata={
AgentLoop._RUNTIME_CHECKPOINT_KEY: {
"phase": "final_response",
"assistant_message": {
"role": "assistant",
"content": "finished",
},
"completed_tool_results": [],
"pending_tool_calls": [],
}
},
)
restored = loop._restore_runtime_checkpoint(session)
assert restored is True
assert session.messages[-1]["content"] == "finished"
assert session.provider_state is None
def test_restore_completed_tools_checkpoint_preserves_matching_provider_state() -> None:
loop = _mk_loop()
tool_result = {
"role": "tool",
"tool_call_id": "call_done",
"name": "read_file",
"content": "compacted result",
}
state = _provider_state().with_pending_messages([tool_result])
session = Session(
key="test:completed-tools-checkpoint",
provider_state=state,
metadata={
AgentLoop._RUNTIME_CHECKPOINT_KEY: {
"phase": "tools_completed",
AgentLoop._PROVIDER_STATE_CHECKPOINT_VERSION_KEY: (
AgentLoop._PROVIDER_STATE_CHECKPOINT_VERSION
),
"assistant_message": {
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_done",
"type": "function",
"function": {"name": "read_file", "arguments": "{}"},
}
],
},
"completed_tool_results": [tool_result],
"pending_tool_calls": [],
}
},
)
restored = loop._restore_runtime_checkpoint(session)
assert restored is True
assert session.messages[-1]["content"] == "compacted result"
assert session.provider_state is state
def test_restore_runtime_checkpoint_dedupes_overlapping_tail() -> None:
@ -616,6 +726,55 @@ def test_restore_runtime_checkpoint_dedupes_overlapping_tail() -> None:
assert session.messages[2]["tool_call_id"] == "call_pending"
@pytest.mark.asyncio
async def test_runtime_checkpoint_keeps_provider_state_out_of_public_metadata(
tmp_path: Path,
) -> None:
loop = _make_full_loop(tmp_path)
state = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="test-model",
version=1,
payload={
"items": [
{
"type": "reasoning",
"encrypted_content": "private-checkpoint-blob",
}
]
},
)
loop.provider.can_resume_conversation_state.return_value = True
loop.provider.chat_with_retry = AsyncMock(
return_value=LLMResponse(content="done", provider_state=state)
)
session = loop.sessions.get_or_create("cli:private-checkpoint")
await loop._run_agent_loop(
[
{"role": "system", "content": "system"},
{"role": "user", "content": "question"},
],
runtime=loop.llm_runtime(),
session=session,
)
assert session.provider_state is not None
checkpoint = session.metadata[AgentLoop._RUNTIME_CHECKPOINT_KEY]
assert "provider_state" not in checkpoint
assert checkpoint[AgentLoop._PROVIDER_STATE_CHECKPOINT_VERSION_KEY] == (
AgentLoop._PROVIDER_STATE_CHECKPOINT_VERSION
)
assert "private-checkpoint-blob" not in json.dumps(session.metadata)
public_payload = loop.sessions.read_session_file(session.key)
assert public_payload is not None
assert "private-checkpoint-blob" not in json.dumps(public_payload)
raw = loop.sessions._get_session_path(session.key).read_text(encoding="utf-8")
assert "private-checkpoint-blob" in raw
@pytest.mark.asyncio
async def test_process_message_persists_user_message_before_turn_completes(tmp_path: Path) -> None:
loop = _make_full_loop(tmp_path)
@ -634,6 +793,150 @@ async def test_process_message_persists_user_message_before_turn_completes(tmp_p
assert persisted.updated_at >= persisted.created_at
@pytest.mark.asyncio
async def test_subagent_followup_stages_provider_state_before_turn_runs(
tmp_path: Path,
) -> None:
loop = _make_full_loop(tmp_path)
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign]
loop._run_agent_loop = AsyncMock(side_effect=RuntimeError("boom")) # type: ignore[method-assign]
loop.provider.can_resume_conversation_state.return_value = True
session = loop.sessions.get_or_create("cli:subagent-crash")
session.provider_state = _provider_state()
loop.sessions.save(session)
msg = InboundMessage(
channel="system",
sender_id="subagent",
chat_id="cli:subagent-crash",
content="subagent result",
metadata={"subagent_task_id": "sub-1"},
)
with pytest.raises(RuntimeError, match="boom"):
await loop._process_message(msg)
loop.sessions.invalidate("cli:subagent-crash")
persisted = loop.sessions.get_or_create("cli:subagent-crash")
assert persisted.messages[-1]["content"] == "subagent result"
assert persisted.provider_state is not None
assert persisted.provider_state.pending_messages[-1]["role"] == "user"
assert persisted.provider_state.pending_messages[-1]["content"] == "subagent result"
@pytest.mark.asyncio
async def test_subagent_followup_state_is_durable_before_prompt_assembly(
tmp_path: Path,
) -> None:
loop = _make_full_loop(tmp_path)
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign]
loop.provider.can_resume_conversation_state.return_value = True
loop._build_initial_messages = MagicMock( # type: ignore[method-assign]
side_effect=RuntimeError("prompt boom"),
)
session = loop.sessions.get_or_create("cli:subagent-prompt-crash")
session.provider_state = _provider_state()
loop.sessions.save(session)
msg = InboundMessage(
channel="system",
sender_id="subagent",
chat_id="cli:subagent-prompt-crash",
content="subagent result",
metadata={"subagent_task_id": "sub-1"},
)
with pytest.raises(RuntimeError, match="prompt boom"):
await loop._process_message(msg)
loop.sessions.invalidate("cli:subagent-prompt-crash")
persisted = loop.sessions.get_or_create("cli:subagent-prompt-crash")
assert persisted.messages[-1]["content"] == "subagent result"
assert persisted.provider_state is not None
assert persisted.provider_state.pending_messages[-1]["content"] == (
"subagent result"
)
@pytest.mark.asyncio
async def test_subagent_redelivery_does_not_duplicate_staged_provider_input(
tmp_path: Path,
) -> None:
loop = _make_full_loop(tmp_path)
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign]
loop.provider.can_resume_conversation_state.return_value = True
build_initial_messages = loop._build_initial_messages
loop._build_initial_messages = MagicMock( # type: ignore[method-assign]
side_effect=RuntimeError("prompt boom"),
)
session = loop.sessions.get_or_create("cli:subagent-redelivery")
session.provider_state = _provider_state()
loop.sessions.save(session)
msg = InboundMessage(
channel="system",
sender_id="subagent",
chat_id="cli:subagent-redelivery",
content="subagent result",
metadata={"subagent_task_id": "sub-1"},
)
with pytest.raises(RuntimeError, match="prompt boom"):
await loop._process_message(msg)
loop.sessions.invalidate("cli:subagent-redelivery")
persisted = loop.sessions.get_or_create("cli:subagent-redelivery")
assert persisted.provider_state is not None
assert [
message.get("content")
for message in persisted.provider_state.pending_messages
].count("subagent result") == 1
loop._build_initial_messages = build_initial_messages # type: ignore[method-assign]
loop._run_agent_loop = AsyncMock( # type: ignore[method-assign]
side_effect=RuntimeError("provider boom"),
)
with pytest.raises(RuntimeError, match="provider boom"):
await loop._process_message(msg)
provider_state = loop._run_agent_loop.await_args.kwargs["provider_state"]
assert provider_state is not None
pending_results = [
message
for message in provider_state.pending_messages
if message.get("content") == "subagent result"
]
assert len(pending_results) == 1
assert LLMProvider._sanitize_empty_content(pending_results) == [
{"role": "user", "content": "subagent result"},
]
@pytest.mark.asyncio
async def test_subagent_followup_clears_state_before_compatibility_failure(
tmp_path: Path,
) -> None:
loop = _make_full_loop(tmp_path)
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign]
loop.provider.can_resume_conversation_state.side_effect = RuntimeError(
"compatibility boom"
)
session = loop.sessions.get_or_create("cli:subagent-compat-crash")
session.provider_state = _provider_state()
loop.sessions.save(session)
msg = InboundMessage(
channel="system",
sender_id="subagent",
chat_id="cli:subagent-compat-crash",
content="subagent result",
metadata={"subagent_task_id": "sub-1"},
)
with pytest.raises(RuntimeError, match="compatibility boom"):
await loop._process_message(msg)
loop.sessions.invalidate("cli:subagent-compat-crash")
persisted = loop.sessions.get_or_create("cli:subagent-compat-crash")
assert persisted.messages[-1]["content"] == "subagent result"
assert persisted.provider_state is None
@pytest.mark.asyncio
async def test_process_message_persists_unified_session_delivery_route(tmp_path: Path) -> None:
loop = _make_full_loop(tmp_path)
@ -1245,6 +1548,9 @@ async def test_next_turn_after_crash_closes_pending_user_turn_before_new_input(t
session = loop.sessions.get_or_create("feishu:c3")
session.add_message("user", "old question")
session.metadata[AgentLoop._PENDING_USER_TURN_KEY] = True
session.provider_state = _provider_state().with_pending_messages([
{"role": "user", "content": "old question"},
])
loop.sessions.save(session)
loop._run_agent_loop = AsyncMock(return_value=(
@ -1278,6 +1584,7 @@ async def test_next_turn_after_crash_closes_pending_user_turn_before_new_input(t
{"role": "assistant", "content": "new answer"},
]
assert AgentLoop._PENDING_USER_TURN_KEY not in session.metadata
assert session.provider_state is None
@pytest.mark.asyncio

View File

@ -11,7 +11,13 @@ import pytest
from agent.runner_helpers import make_run_spec
from nanobot.config.schema import AgentDefaults
from nanobot.providers.base import LLMProvider, LLMResponse, ToolCallRequest
from nanobot.providers.base import (
LLMProvider,
LLMResponse,
ProviderCallContext,
ProviderConversationState,
ToolCallRequest,
)
_MAX_TOOL_RESULT_CHARS = AgentDefaults().max_tool_result_chars
@ -73,6 +79,311 @@ async def test_runner_preserves_reasoning_fields_and_tool_results():
)
@pytest.mark.asyncio
async def test_runner_replays_provider_state_without_chat_projection_duplicates():
from nanobot.agent.runner import AgentRunner
provider = MagicMock(spec=LLMProvider)
provider.can_resume_conversation_state.return_value = True
provider.supports_native_compaction.return_value = False
captured_second_kwargs: dict = {}
checkpoints: list[dict] = []
calls = 0
async def checkpoint(payload: dict) -> None:
checkpoints.append(payload)
first_state = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="gpt-5.6",
version=1,
payload={"items": [{"type": "reasoning", "encrypted_content": "opaque"}]},
)
second_state = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="gpt-5.6",
version=1,
payload={"items": [{"type": "message", "role": "assistant"}]},
)
async def chat_with_retry(**kwargs):
nonlocal calls
calls += 1
if calls == 1:
provider_context = kwargs["provider_context"]
assert isinstance(provider_context, ProviderCallContext)
assert provider_context.conversation_state is None
return LLMResponse(
content=None,
tool_calls=[
ToolCallRequest(
id="call_1|fc_1",
name="list_dir",
arguments={"path": "."},
),
],
provider_state=first_state,
)
captured_second_kwargs.update(kwargs)
return LLMResponse(content="done", provider_state=second_state)
provider.chat_with_retry = chat_with_retry
tools = MagicMock()
tools.get_definitions.return_value = []
tools.execute = AsyncMock(return_value="tool result")
result = await AgentRunner().run(make_run_spec(
provider,
initial_messages=[
{"role": "system", "content": "system"},
{"role": "user", "content": "do task"},
],
tools=tools,
model="gpt-5.6",
max_iterations=3,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
checkpoint_callback=checkpoint,
))
provider_context = captured_second_kwargs["provider_context"]
assert isinstance(provider_context, ProviderCallContext)
assert provider_context.conversation_state is not None
assert provider_context.conversation_state.payload == first_state.payload
assert provider_context.conversation_state.pending_messages == [{
"role": "tool",
"tool_call_id": "call_1|fc_1",
"name": "list_dir",
"content": "tool result",
}]
assert not any(
message.get("role") == "assistant"
for message in provider_context.conversation_state.pending_messages
)
assert result.provider_state is not None
assert result.provider_state.payload == second_state.payload
assert result.provider_state.pending_messages == []
assert checkpoints[0]["phase"] == "awaiting_tools"
assert "provider_state" not in checkpoints[0]
assert checkpoints[1]["phase"] == "tools_completed"
assert checkpoints[1]["provider_state"].pending_messages == [{
"role": "tool",
"tool_call_id": "call_1|fc_1",
"name": "list_dir",
"content": "tool result",
}]
assert checkpoints[2]["phase"] == "final_response"
assert checkpoints[2]["provider_state"].payload == second_state.payload
@pytest.mark.asyncio
async def test_runner_governs_tool_result_before_adding_it_to_provider_state():
from nanobot.agent.runner import AgentRunner
provider = MagicMock(spec=LLMProvider)
provider.can_resume_conversation_state.return_value = True
provider.supports_native_compaction.return_value = False
calls = 0
captured_context: ProviderCallContext | None = None
checkpoints: list[dict] = []
state = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="gpt-5.6",
version=1,
payload={"items": [{"type": "reasoning", "encrypted_content": "opaque"}]},
)
async def chat_with_retry(**kwargs):
nonlocal calls, captured_context
calls += 1
if calls == 1:
return LLMResponse(
content=None,
tool_calls=[
ToolCallRequest(
id="call_1",
name="read_file",
arguments={"path": "large.txt"},
),
],
provider_state=state,
)
captured_context = kwargs["provider_context"]
return LLMResponse(content="done")
provider.chat_with_retry = chat_with_retry
tools = MagicMock()
tools.get_definitions.return_value = []
tools.execute = AsyncMock(return_value="x" * 5_000)
async def checkpoint(payload: dict) -> None:
checkpoints.append(payload)
await AgentRunner().run(make_run_spec(
provider,
initial_messages=[
{"role": "system", "content": "system"},
{"role": "user", "content": "read the file"},
],
tools=tools,
model="gpt-5.6",
context_window_tokens=3_000,
context_block_limit=200,
max_tokens=1_000,
max_iterations=3,
max_tool_result_chars=10_000,
checkpoint_callback=checkpoint,
))
assert captured_context is not None
assert captured_context.conversation_state is not None
pending = captured_context.conversation_state.pending_messages
assert len(pending) == 1
assert pending[0]["role"] == "tool"
assert "compacted to fit context" in pending[0]["content"]
assert pending[0]["content"] != "x" * 5_000
completed_checkpoint = next(
checkpoint
for checkpoint in checkpoints
if checkpoint["phase"] == "tools_completed"
)
checkpoint_pending = completed_checkpoint["provider_state"].pending_messages
assert "compacted to fit context" in checkpoint_pending[0]["content"]
assert checkpoint_pending[0]["content"] != "x" * 5_000
@pytest.mark.asyncio
async def test_injected_final_response_checkpoint_includes_provider_state():
from nanobot.agent.runner import AgentRunner
provider = MagicMock(spec=LLMProvider)
provider.can_resume_conversation_state.return_value = True
provider.supports_native_compaction.return_value = False
first_state = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="gpt-5.6",
version=1,
payload={"items": [{"type": "message", "content": "first answer"}]},
)
second_state = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="gpt-5.6",
version=1,
payload={"items": [{"type": "message", "content": "second answer"}]},
)
provider.chat_with_retry = AsyncMock(side_effect=[
LLMResponse(content="first answer", provider_state=first_state),
LLMResponse(content="second answer", provider_state=second_state),
])
tools = MagicMock()
tools.get_definitions.return_value = []
checkpoints: list[dict] = []
injections = [[{"role": "user", "content": "follow up"}], []]
async def checkpoint(payload: dict) -> None:
checkpoints.append(payload)
async def inject() -> list[dict]:
return injections.pop(0)
await AgentRunner().run(make_run_spec(
provider,
initial_messages=[{"role": "user", "content": "start"}],
tools=tools,
model="gpt-5.6",
max_iterations=3,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
checkpoint_callback=checkpoint,
injection_callback=inject,
))
assert checkpoints[0]["phase"] == "final_response"
assert checkpoints[0]["provider_state"].payload == first_state.payload
@pytest.mark.asyncio
async def test_runner_preserves_last_completed_provider_state_on_model_error():
from nanobot.agent.runner import AgentRunner
provider = MagicMock(spec=LLMProvider)
provider.can_resume_conversation_state.return_value = True
provider.chat_with_retry = AsyncMock(return_value=LLMResponse(
content="temporary upstream failure",
finish_reason="error",
error_kind="timeout",
))
tools = MagicMock()
tools.get_definitions.return_value = []
state = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="gpt-5.6",
version=1,
payload={"items": [{"type": "reasoning", "encrypted_content": "opaque"}]},
)
unsaved_input = {"role": "user", "content": "ephemeral follow-up"}
result = await AgentRunner().run(make_run_spec(
provider,
initial_messages=[
{"role": "system", "content": "system"},
unsaved_input,
],
tools=tools,
model="gpt-5.6",
max_iterations=1,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
provider_state=state.with_pending_messages([unsaved_input]),
))
assert result.stop_reason == "error"
assert result.provider_state is not None
assert result.provider_state.payload == state.payload
assert result.provider_state.pending_messages[0] == unsaved_input
assert result.provider_state.pending_messages[1]["role"] == "assistant"
assert "model error" in result.provider_state.pending_messages[1]["content"]
@pytest.mark.asyncio
async def test_runner_discards_provider_state_on_non_retryable_model_error():
from nanobot.agent.runner import AgentRunner
provider = MagicMock(spec=LLMProvider)
provider.can_resume_conversation_state.return_value = True
provider.chat_with_retry = AsyncMock(return_value=LLMResponse(
content="context length exceeded",
finish_reason="error",
error_status_code=400,
error_should_retry=False,
))
tools = MagicMock()
tools.get_definitions.return_value = []
state = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="gpt-5.6",
version=1,
payload={"items": [{"type": "reasoning", "encrypted_content": "opaque"}]},
)
result = await AgentRunner().run(make_run_spec(
provider,
initial_messages=[{"role": "user", "content": "continue"}],
tools=tools,
model="gpt-5.6",
max_iterations=1,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
provider_state=state,
))
assert result.stop_reason == "error"
assert result.provider_state is None
@pytest.mark.asyncio
async def test_runner_returns_max_iterations_fallback():
from nanobot.agent.runner import AgentRunner
@ -422,6 +733,66 @@ async def test_runner_retries_empty_final_response_with_summary_prompt():
assert result.usage["completion_tokens"] == 9
@pytest.mark.asyncio
@pytest.mark.parametrize("finish_reason", ["refusal", "content_filter"])
async def test_runner_does_not_retry_blank_policy_terminal(
finish_reason: str,
) -> None:
from nanobot.agent.runner import AgentRunner
from nanobot.utils.runtime import EMPTY_FINAL_RESPONSE_MESSAGE
provider = MagicMock(spec=LLMProvider)
provider.chat_with_retry = AsyncMock(return_value=LLMResponse(
content=None,
finish_reason=finish_reason,
))
tools = MagicMock()
tools.get_definitions.return_value = []
result = await AgentRunner().run(make_run_spec(
provider,
initial_messages=[{"role": "user", "content": "do task"}],
tools=tools,
model="test-model",
max_iterations=3,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
))
assert provider.chat_with_retry.await_count == 1
assert result.final_content == EMPTY_FINAL_RESPONSE_MESSAGE
assert result.stop_reason == "empty_final_response"
@pytest.mark.asyncio
@pytest.mark.parametrize("finish_reason", ["refusal", "content_filter"])
async def test_runner_does_not_auto_continue_goal_after_policy_terminal(
finish_reason: str,
) -> None:
from nanobot.agent.runner import AgentRunner
provider = MagicMock(spec=LLMProvider)
provider.chat_with_retry = AsyncMock(return_value=LLMResponse(
content="Request blocked by provider policy.",
finish_reason=finish_reason,
))
tools = MagicMock()
tools.get_definitions.return_value = []
result = await AgentRunner().run(make_run_spec(
provider,
initial_messages=[{"role": "user", "content": "do task"}],
tools=tools,
model="test-model",
max_iterations=3,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
goal_active_predicate=lambda: True,
))
assert provider.chat_with_retry.await_count == 1
assert result.final_content == "Request blocked by provider policy."
assert result.stop_reason == "completed"
@pytest.mark.asyncio
async def test_runner_uses_specific_message_after_empty_finalization_retry():
"""After silent retries + finalization all return empty, stop_reason is empty_final_response."""
@ -450,6 +821,56 @@ async def test_runner_uses_specific_message_after_empty_finalization_retry():
assert result.stop_reason == "empty_final_response"
@pytest.mark.asyncio
async def test_empty_finalization_retry_discards_candidate_provider_state():
from nanobot.agent.runner import AgentRunner
candidate = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="test-model",
version=1,
payload={
"items": [{
"type": "function_call",
"call_id": "call_1",
"name": "exec",
"arguments": "{}",
}],
},
)
provider = MagicMock(spec=LLMProvider)
provider.can_resume_conversation_state.return_value = True
provider.chat_with_retry = AsyncMock(side_effect=[
LLMResponse(content=None, tool_calls=[], usage={}),
LLMResponse(content=None, tool_calls=[], usage={}),
LLMResponse(
content="finalized without tools",
tool_calls=[ToolCallRequest(id="call_1", name="exec", arguments={})],
finish_reason="stop",
provider_state=candidate,
usage={},
),
])
tools = MagicMock()
tools.get_definitions.return_value = []
tools.execute = AsyncMock(return_value="must not run")
runner = AgentRunner()
result = await runner.run(make_run_spec(
provider,
initial_messages=[{"role": "user", "content": "do task"}],
tools=tools,
model="test-model",
max_iterations=3,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
))
tools.execute.assert_not_awaited()
assert result.final_content == "finalized without tools"
assert result.provider_state is None
@pytest.mark.asyncio
async def test_runner_length_recovery_returns_all_segments():
"""Recovered output segments are returned together instead of only the tail."""

View File

@ -9,8 +9,15 @@ import pytest
from loguru import logger
from nanobot.config.schema import ModelPresetConfig
from nanobot.providers.base import LLMProvider, LLMResponse
from nanobot.providers.base import (
LLMProvider,
LLMResponse,
ProviderCallContext,
ProviderConversationState,
)
from nanobot.providers.conversation_state import ProviderConversationStateController
from nanobot.providers.fallback_provider import FallbackProvider
from nanobot.providers.openai_responses import resolve_compact_threshold
def _make_response(
@ -66,6 +73,9 @@ class _FakeProvider(LLMProvider):
self._response = response or _make_response()
self.chat_calls: list[dict[str, Any]] = []
self.chat_stream_calls: list[dict[str, Any]] = []
self.context_calls: list[ProviderCallContext | None] = []
self.resumable = False
self.compact = False
def get_default_model(self) -> str:
return f"{self.name}/model"
@ -81,6 +91,26 @@ class _FakeProvider(LLMProvider):
await on_delta(self._response.content)
return self._response
async def chat_with_context(
self,
provider_context: ProviderCallContext | None = None,
**kwargs: Any,
) -> LLMResponse:
self.context_calls.append(provider_context)
return await self.chat(**kwargs)
def can_resume_conversation_state(
self,
state: ProviderConversationState,
model: str | None = None,
) -> bool:
_ = state, model
return self.resumable
def supports_native_compaction(self, model: str | None = None) -> bool:
_ = model
return self.compact
# -- config-level tests --
@ -211,6 +241,8 @@ def test_provider_snapshot_uses_smallest_fallback_context_window() -> None:
snapshot = build_provider_snapshot(config)
assert snapshot.context_window_tokens == 64000
assert isinstance(snapshot.provider, FallbackProvider)
assert snapshot.provider._primary_context_window_tokens == 128000
def test_inline_fallback_reasoning_effort_does_not_inherit_primary() -> None:
@ -285,6 +317,257 @@ class TestFallbackOnPrimaryError:
assert primary.chat_calls[0]["model"] == "primary-model"
assert fallback.chat_calls[0]["model"] == "fallback-a"
@pytest.mark.asyncio
async def test_primary_compaction_uses_primary_context_window(self) -> None:
primary = _FakeProvider("primary", _make_response("primary ok"))
primary.compact = True
fb = FallbackProvider(
primary=primary,
fallback_presets=[
_fallback("small-chat", context_window_tokens=50_000),
],
provider_factory=MagicMock(),
primary_context_window_tokens=200_000,
)
await fb.chat_with_context(
messages=[{"role": "user", "content": "hi"}],
model="gpt-5.6",
max_tokens=10_000,
provider_context=ProviderCallContext(context_window_tokens=50_000),
)
primary_context = primary.context_calls[0]
assert primary_context is not None
assert primary_context.context_window_tokens == 200_000
assert resolve_compact_threshold(
primary_context.context_window_tokens,
10_000,
) == 180_000
@pytest.mark.asyncio
async def test_native_fallback_compaction_uses_its_own_context_window(self) -> None:
primary = _FakeProvider("primary", _error_response())
primary.compact = True
fallback = _FakeProvider("fallback", _make_response("fallback ok"))
fallback.compact = True
fb = FallbackProvider(
primary=primary,
fallback_presets=[
_fallback("fallback-a", context_window_tokens=120_000),
],
provider_factory=MagicMock(return_value=fallback),
primary_context_window_tokens=200_000,
)
result = await fb.chat_with_context(
messages=[{"role": "user", "content": "hi"}],
model="gpt-5.6",
provider_context=ProviderCallContext(context_window_tokens=50_000),
)
assert result.content == "fallback ok"
assert primary.context_calls == [
ProviderCallContext(context_window_tokens=200_000)
]
assert fallback.context_calls == [
ProviderCallContext(context_window_tokens=120_000)
]
@pytest.mark.asyncio
async def test_native_fallback_gets_context_when_primary_does_not_use_it(self) -> None:
primary = _FakeProvider("primary", _error_response())
fallback = _FakeProvider("fallback", _make_response("fallback ok"))
fallback.compact = True
fb = FallbackProvider(
primary=primary,
fallback_presets=[
_fallback("fallback-a", context_window_tokens=120_000),
],
provider_factory=MagicMock(return_value=fallback),
primary_context_window_tokens=200_000,
)
messages = [{"role": "user", "content": "hi"}]
controller = ProviderConversationStateController(
provider=fb,
model="primary-model",
messages=messages,
)
assert fb.supports_native_compaction("primary-model") is False
provider_context = controller.prepare_request(
messages,
context_window_tokens=50_000,
)
assert provider_context == ProviderCallContext(
context_window_tokens=50_000
)
result = await fb.chat_with_context(
messages=messages,
model="primary-model",
provider_context=provider_context,
)
assert result.content == "fallback ok"
assert primary.context_calls == [ProviderCallContext()]
assert fallback.context_calls == [
ProviderCallContext(context_window_tokens=120_000)
]
@pytest.mark.asyncio
async def test_responses_chat_fallback_responses_rebuilds_state(self) -> None:
primary = _FakeProvider("primary", _error_response())
primary.resumable = True
primary.compact = True
fallback = _FakeProvider("fallback", _make_response("fallback ok"))
messages = [{"role": "user", "content": "hi"}]
state = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="gpt-5.6",
version=1,
payload={"items": [{"type": "reasoning", "encrypted_content": "opaque"}]},
pending_messages=list(messages),
)
fb = FallbackProvider(
primary=primary,
fallback_presets=[_fallback("fallback-a")],
provider_factory=MagicMock(return_value=fallback),
)
controller = ProviderConversationStateController(
provider=fb,
model="gpt-5.6",
messages=messages,
state=state,
)
provider_context = controller.prepare_request(
messages,
context_window_tokens=200_000,
)
assert provider_context is not None
result = await fb.chat_with_context(
messages=messages,
model="gpt-5.6",
provider_context=provider_context,
)
assert result.content == "fallback ok"
assert primary.context_calls == [provider_context]
assert fallback.context_calls == [ProviderCallContext()]
assert fallback.chat_calls[0]["messages"] == messages
controller.observe_response(result, messages)
messages.append({"role": "assistant", "content": result.content})
assert controller.finish(messages) is None
recovered_state = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="gpt-5.6",
version=1,
payload={"items": [{"type": "reasoning", "encrypted_content": "recovered"}]},
)
primary._response = LLMResponse(
content="primary recovered",
provider_state=recovered_state,
)
next_turn = ProviderConversationStateController(
provider=fb,
model="gpt-5.6",
messages=messages,
)
next_context = next_turn.prepare_request(
messages,
context_window_tokens=200_000,
)
assert next_context == ProviderCallContext(context_window_tokens=200_000)
recovered = await fb.chat_with_context(
messages=messages,
model="gpt-5.6",
provider_context=next_context,
)
assert recovered.provider_state is recovered_state
assert primary.context_calls[-1] == next_context
assert primary.chat_calls[-1]["messages"] == messages
@pytest.mark.asyncio
@pytest.mark.parametrize(
("primary_error_kind", "primary_status", "primary_should_retry"),
[
("server_error", 503, True),
("authentication", 401, False),
],
ids=["transient", "authentication"],
)
async def test_final_fallback_error_uses_primary_state_disposition(
self,
primary_error_kind: str,
primary_status: int,
primary_should_retry: bool,
) -> None:
primary = _FakeProvider(
"primary",
_make_response(
"primary unavailable",
finish_reason="error",
error_kind=primary_error_kind,
error_status_code=primary_status,
error_should_retry=primary_should_retry,
),
)
primary.resumable = True
fallback = _FakeProvider(
"fallback",
_make_response(
"fallback invalid request",
finish_reason="error",
error_kind="invalid_request",
error_status_code=400,
error_should_retry=False,
),
)
messages = [{"role": "user", "content": "continue"}]
state = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="gpt-5.6",
version=1,
payload={"items": [{"type": "reasoning", "encrypted_content": "opaque"}]},
pending_messages=list(messages),
)
provider = FallbackProvider(
primary=primary,
fallback_presets=[_fallback("fallback-a")],
provider_factory=MagicMock(return_value=fallback),
)
controller = ProviderConversationStateController(
provider=provider,
model="gpt-5.6",
messages=messages,
state=state,
)
provider_context = controller.prepare_request(
messages,
context_window_tokens=200_000,
)
assert provider_context is not None
response = await provider.chat_with_context(
messages=messages,
model="gpt-5.6",
provider_context=provider_context,
)
controller.observe_response(response, messages)
assert response.content == "fallback invalid request"
assert response.preserve_provider_state_on_error is True
restored = controller.finish(messages)
assert restored is not None
assert restored.payload == state.payload
@pytest.mark.asyncio
async def test_reports_the_fallback_model_before_its_request(self) -> None:
primary = _FakeProvider("primary", _error_response())

View File

@ -15,7 +15,11 @@ from nanobot.agent.context_governance import (
)
from nanobot.agent.runner import AgentRunSpec
from nanobot.config.schema import AgentDefaults
from nanobot.providers.base import LLMResponse, ToolCallRequest
from nanobot.providers.base import (
LLMResponse,
ProviderConversationState,
ToolCallRequest,
)
_MAX_TOOL_RESULT_CHARS = AgentDefaults().max_tool_result_chars
@ -886,6 +890,13 @@ def test_drop_malformed_tool_calls_trims_response():
"""LLM response tool_calls with a missing/empty name are dropped in place."""
from nanobot.agent.runner import AgentRunner
candidate_state = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="gpt-5.6",
version=1,
payload={"items": [{"type": "function_call", "name": None}]},
)
response = LLMResponse(
content=None,
tool_calls=[
@ -895,9 +906,11 @@ def test_drop_malformed_tool_calls_trims_response():
ToolCallRequest(id="4", name="read_file", arguments={}),
],
finish_reason="tool_calls",
provider_state=candidate_state,
)
dropped, all_dropped, orig = AgentRunner._drop_malformed_tool_calls(response)
assert [tc.name for tc in response.tool_calls] == ["read_file"]
assert response.provider_state is None
assert response.finish_reason == "tool_calls"
assert response.should_execute_tools is True
assert dropped == 3

View File

@ -4,6 +4,7 @@ import json
from datetime import datetime
from pathlib import Path
from nanobot.providers.base import ProviderConversationState
from nanobot.session.manager import Session, SessionManager
@ -101,6 +102,137 @@ class TestAtomicSave:
for i in range(5):
assert loaded.messages[i]["content"] == f"msg{i}"
def test_provider_state_round_trips_in_private_record_only(self, tmp_path: Path):
mgr = SessionManager(tmp_path)
secret = "encrypted-reasoning-blob"
session = Session(
key="test:provider-state",
provider_state=ProviderConversationState(
kind="openai_responses",
provider="openai:https://api.openai.com/v1",
model="gpt-5.6",
version=1,
payload={
"items": [
{
"type": "reasoning",
"encrypted_content": secret,
}
]
},
pending_messages=[{"role": "user", "content": "continue"}],
),
)
session.add_message("user", "hello")
mgr.save(session)
records = [
json.loads(line)
for line in mgr._get_session_path(session.key)
.read_text(encoding="utf-8")
.splitlines()
]
assert [record.get("_type") for record in records] == [
"metadata",
"provider_state",
None,
]
assert secret in records[1]["state"]["payload"]["items"][0]["encrypted_content"]
mgr.invalidate(session.key)
loaded = mgr.get_or_create(session.key)
assert loaded.provider_state is not None
assert loaded.provider_state.to_private_record() == session.provider_state.to_private_record()
public_payload = mgr.read_session_file(session.key)
assert public_payload is not None
assert public_payload["messages"] == [session.messages[0]]
assert secret not in json.dumps(public_payload)
assert secret not in json.dumps(mgr.list_sessions())
def test_provider_state_does_not_consume_list_preview_budget(
self,
tmp_path: Path,
monkeypatch,
):
import nanobot.session.manager as session_manager
monkeypatch.setattr(session_manager, "_SESSION_LIST_PREVIEW_MAX_CHARS", 100)
mgr = SessionManager(tmp_path)
session = Session(
key="test:provider-state-preview",
provider_state=ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="test-model",
version=1,
payload={"items": [{"encrypted_content": "x" * 200}]},
),
)
session.add_message("user", "visible preview")
mgr.save(session)
assert mgr.list_sessions()[0]["preview"] == "visible preview"
def test_clear_and_fork_discard_provider_state(self, tmp_path: Path):
mgr = SessionManager(tmp_path)
state = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="gpt-5.6",
version=1,
payload={"items": []},
)
source = Session(key="test:state-source", provider_state=state)
source.add_message("user", "hello")
mgr.save(source)
fork = mgr.fork_session_before_user_index(
source.key,
"test:state-fork",
1,
)
assert fork is not None
assert fork.provider_state is None
source.clear()
assert source.provider_state is None
def test_invalid_provider_state_record_is_not_public_history(self, tmp_path: Path):
mgr = SessionManager(tmp_path)
path = mgr._get_session_path("test:bad-provider-state")
path.parent.mkdir(parents=True, exist_ok=True)
path.write_text(
"\n".join(
[
json.dumps(
{
"_type": "metadata",
"key": "test:bad-provider-state",
"created_at": datetime.now().isoformat(),
"updated_at": datetime.now().isoformat(),
"metadata": {},
"last_consolidated": 0,
}
),
json.dumps(
{
"_type": "provider_state",
"state": {"kind": "openai_responses"},
}
),
json.dumps({"role": "user", "content": "safe"}),
]
)
+ "\n",
encoding="utf-8",
)
loaded = mgr._load("test:bad-provider-state")
assert loaded is not None
assert loaded.provider_state is None
assert loaded.messages == [{"role": "user", "content": "safe"}]
class TestRepairCorruptFile:
def _write_corrupt_jsonl(self, path: Path, lines: list[str]) -> None:

View File

@ -1,3 +1,4 @@
from nanobot.providers.base import ProviderConversationState
from nanobot.runtime_context import (
RUNTIME_CONTEXT_HISTORY_META,
RuntimeContextBlock,
@ -769,7 +770,16 @@ def test_get_history_extend_to_user_keeps_newer_user_inside_window():
def test_retain_recent_legal_suffix_returns_dropped_messages():
"""retain_recent_legal_suffix returns the actually-dropped messages."""
session = Session(key="test:return-dropped")
session = Session(
key="test:return-dropped",
provider_state=ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="test-model",
version=1,
payload={"items": []},
),
)
for i in range(10):
session.messages.append({"role": "user", "content": f"msg{i}"})
@ -779,11 +789,19 @@ def test_retain_recent_legal_suffix_returns_dropped_messages():
assert [m["content"] for m in result.dropped] == [f"msg{i}" for i in range(6)]
assert len(session.messages) == 4
assert result.already_consolidated_count == 0
assert session.provider_state is None
def test_retain_recent_legal_suffix_returns_empty_when_no_drop():
"""No messages dropped → empty list returned."""
session = Session(key="test:no-drop")
state = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="test-model",
version=1,
payload={"items": []},
)
session = Session(key="test:no-drop", provider_state=state)
for i in range(3):
session.messages.append({"role": "user", "content": f"msg{i}"})
@ -792,6 +810,7 @@ def test_retain_recent_legal_suffix_returns_empty_when_no_drop():
assert result.dropped == []
assert result.already_consolidated_count == 0
assert len(session.messages) == 3
assert session.provider_state is state
def test_retain_recent_legal_suffix_returns_all_on_zero():

View File

@ -504,6 +504,7 @@ async def test_drain_pending_blocks_while_subagents_running(tmp_path):
usage={},
had_injections=False,
tools_used=[],
provider_state=None,
)
loop.runner.run = AsyncMock(side_effect=fake_runner_run)
@ -589,6 +590,7 @@ async def test_drain_pending_no_block_when_no_subagents(tmp_path):
usage={},
had_injections=False,
tools_used=[],
provider_state=None,
)
loop.runner.run = AsyncMock(side_effect=fake_runner_run)
@ -638,6 +640,7 @@ async def test_drain_pending_timeout(tmp_path):
usage={},
had_injections=False,
tools_used=[],
provider_state=None,
)
loop.runner.run = AsyncMock(side_effect=fake_runner_run)

View File

@ -11,7 +11,7 @@ from nanobot.providers.azure_openai_provider import (
AzureOpenAIProvider,
_AzureTokenProvider,
)
from nanobot.providers.base import LLMResponse
from nanobot.providers.base import LLMResponse, ProviderCallContext
# ---------------------------------------------------------------------------
# Init & validation
@ -234,6 +234,7 @@ def test_build_body_basic():
assert body["max_output_tokens"] == 4096
assert body["store"] is False
assert "reasoning" not in body
assert "include" not in body
# input should contain the converted user message only (system extracted)
assert any(
item.get("role") == "user"
@ -241,6 +242,30 @@ def test_build_body_basic():
)
def test_build_body_enables_server_compaction():
provider = AzureOpenAIProvider(
api_key="k",
api_base="https://res.openai.azure.com",
default_model="gpt-5.6",
)
body = provider._build_body(
[{"role": "user", "content": "hello"}],
None,
None,
10_000,
0.1,
"high",
None,
provider_context=ProviderCallContext(context_window_tokens=200_000),
)
assert body["context_management"] == [{
"type": "compaction",
"compact_threshold": 180_000,
}]
def test_build_body_max_tokens_minimum():
"""max_output_tokens should never be less than 1."""
provider = AzureOpenAIProvider(api_key="k", api_base="https://r.com", default_model="gpt-4o")
@ -358,6 +383,38 @@ async def test_chat_success():
assert result.usage["prompt_tokens"] == 10
@pytest.mark.asyncio
async def test_chat_retries_without_unsupported_server_compaction():
provider = AzureOpenAIProvider(
api_key="test-key",
api_base="https://test.openai.azure.com",
default_model="gpt-5.6",
)
class UnsupportedCompactionError(Exception):
status_code = 400
body = {"error": {"message": "Unknown parameter: context_management"}}
provider._client.responses = MagicMock()
provider._client.responses.create = AsyncMock(side_effect=[
UnsupportedCompactionError(),
_make_sdk_response(content="compaction fallback"),
])
result = await provider.chat(
[{"role": "user", "content": "Hi"}],
provider_context=ProviderCallContext(context_window_tokens=200_000),
)
create = provider._client.responses.create
assert result.content == "compaction fallback"
assert result.provider_state is not None
assert create.await_count == 2
assert "context_management" in create.call_args_list[0].kwargs
assert "context_management" not in create.call_args_list[1].kwargs
assert provider.supports_native_compaction() is False
@pytest.mark.asyncio
async def test_chat_uses_default_model():
provider = AzureOpenAIProvider(
@ -411,6 +468,7 @@ async def test_chat_with_tool_calls():
assert len(result.tool_calls) == 1
assert result.tool_calls[0].name == "get_weather"
assert result.tool_calls[0].arguments == {"location": "SF"}
assert result.provider_state is not None
@pytest.mark.asyncio
@ -510,6 +568,7 @@ async def test_chat_stream_with_tool_calls():
item_done.name = "get_weather"
ev_item_done = MagicMock(type="response.output_item.done", item=item_done)
resp_obj = MagicMock(status="completed")
resp_obj.model_dump.return_value = {"status": "completed", "output": []}
ev_completed = MagicMock(type="response.completed", response=resp_obj)
async def mock_stream():
@ -527,6 +586,7 @@ async def test_chat_stream_with_tool_calls():
assert len(result.tool_calls) == 1
assert result.tool_calls[0].name == "get_weather"
assert result.tool_calls[0].arguments == {"location": "SF"}
assert result.provider_state is not None
@pytest.mark.asyncio

View File

@ -0,0 +1,291 @@
"""Tests for provider-owned conversation-state lifecycle coordination."""
from __future__ import annotations
from unittest.mock import MagicMock
import pytest
from nanobot.providers.base import (
LLMProvider,
LLMResponse,
ProviderConversationState,
ToolCallRequest,
)
from nanobot.providers.conversation_state import (
ProviderConversationStateController,
allows_conversation_message_merge,
)
def _provider(*, resumable: bool = True, compact: bool = False) -> MagicMock:
provider = MagicMock(spec=LLMProvider)
provider.can_resume_conversation_state.return_value = resumable
provider.supports_native_compaction.return_value = compact
return provider
def _state(label: str, *, pending: list[dict] | None = None) -> ProviderConversationState:
return ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="gpt-5.6",
version=1,
payload={"items": [{"type": "reasoning", "encrypted_content": label}]},
pending_messages=pending or [],
)
def test_controller_replays_only_messages_after_provider_output() -> None:
provider = _provider()
messages = [
{"role": "system", "content": "system"},
{"role": "user", "content": "run a tool"},
]
controller = ProviderConversationStateController(
provider=provider,
model="gpt-5.6",
messages=messages,
)
state = _state("first")
controller.prepare_request(messages, context_window_tokens=200_000)
response = LLMResponse(content=None, provider_state=state)
controller.observe_response(response, messages)
assert allows_conversation_message_merge(messages[-1]) is False
messages.append(controller.project_response_message(
{
"role": "assistant",
"content": None,
"tool_calls": [{"id": "call_1", "type": "function"}],
},
response,
))
tool_message = {
"role": "tool",
"tool_call_id": "call_1",
"content": "tool result",
}
messages.append(tool_message)
provider_context = controller.prepare_request(
messages,
context_window_tokens=200_000,
)
assert provider_context is not None
assert provider_context.conversation_state is not None
assert provider_context.conversation_state.payload == state.payload
assert provider_context.conversation_state.pending_messages == [tool_message]
assert controller.checkpoint(messages).pending_messages == [tool_message]
def test_controller_uses_governed_messages_for_provider_state_delta() -> None:
provider = _provider()
messages = [
{"role": "user", "content": "run a tool"},
]
controller = ProviderConversationStateController(
provider=provider,
model="gpt-5.6",
messages=messages,
)
state = _state("first")
controller.prepare_request(messages, context_window_tokens=200_000)
response = LLMResponse(content=None, provider_state=state)
controller.observe_response(response, messages)
messages.extend([
controller.project_response_message(
{
"role": "assistant",
"content": None,
"tool_calls": [{"id": "call_1", "type": "function"}],
},
response,
),
{
"role": "tool",
"tool_call_id": "call_1",
"content": "raw oversized result",
},
])
governed_messages = [
messages[0],
messages[1],
{
"role": "tool",
"tool_call_id": "call_1",
"content": "compacted result",
},
]
provider_context = controller.prepare_request(
messages,
context_window_tokens=200_000,
model_messages=governed_messages,
)
assert provider_context is not None
assert provider_context.conversation_state is not None
assert provider_context.conversation_state.pending_messages == [{
"role": "tool",
"tool_call_id": "call_1",
"content": "compacted result",
}]
assert controller.checkpoint(messages).pending_messages[-1]["content"] == (
"raw oversized result"
)
governed_checkpoint = controller.checkpoint(
messages,
model_messages=governed_messages,
)
assert governed_checkpoint is not None
assert governed_checkpoint.pending_messages[-1]["content"] == "compacted result"
def test_transient_response_preserves_only_durable_request_messages() -> None:
provider = _provider()
current_message = {"role": "user", "content": "continue"}
supplemental = {"role": "user", "content": "internal finalization retry"}
messages = [{"role": "system", "content": "system"}, current_message]
controller = ProviderConversationStateController(
provider=provider,
model="gpt-5.6",
messages=messages,
state=_state("saved", pending=[
{"role": "tool", "content": "prior"},
current_message,
]),
)
provider_context = controller.prepare_request(
messages,
context_window_tokens=200_000,
supplemental_messages=[supplemental],
)
assert provider_context is not None
assert provider_context.conversation_state is not None
assert provider_context.conversation_state.pending_messages == [
{"role": "tool", "content": "prior"},
current_message,
supplemental,
]
controller.observe_response(
LLMResponse(
content="temporary failure",
finish_reason="error",
error_kind="timeout",
),
messages,
)
placeholder = {"role": "assistant", "content": "model error"}
messages.append(placeholder)
state = controller.finish(messages)
assert state is not None
assert state.pending_messages == [
{"role": "tool", "content": "prior"},
current_message,
placeholder,
]
def test_non_retryable_response_discards_saved_state() -> None:
provider = _provider()
messages = [{"role": "user", "content": "continue"}]
controller = ProviderConversationStateController(
provider=provider,
model="gpt-5.6",
messages=messages,
state=_state("saved"),
)
controller.prepare_request(messages, context_window_tokens=200_000)
controller.observe_response(
LLMResponse(
content="invalid request",
finish_reason="error",
error_status_code=400,
error_should_retry=False,
),
messages,
)
assert controller.finish(messages) is None
@pytest.mark.parametrize(
("finish_reason", "exposes_tool_call"),
[
("length", False),
("length", True),
("refusal", True),
("content_filter", True),
],
)
def test_terminal_response_discards_candidate_state(
finish_reason: str,
exposes_tool_call: bool,
) -> None:
provider = _provider()
messages = [{"role": "user", "content": "continue"}]
controller = ProviderConversationStateController(
provider=provider,
model="gpt-5.6",
messages=messages,
state=_state("saved"),
)
controller.prepare_request(messages, context_window_tokens=200_000)
candidate = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="gpt-5.6",
version=1,
payload={
"items": [{
"type": "function_call",
"call_id": "call_1",
"name": "exec",
"arguments": "{}",
}],
},
)
response = LLMResponse(
content="terminal response",
tool_calls=(
[ToolCallRequest(id="call_1", name="exec", arguments={})]
if exposes_tool_call
else []
),
finish_reason=finish_reason,
provider_state=candidate,
)
assert response.has_tool_calls is exposes_tool_call
assert response.should_execute_tools is False
controller.observe_response(response, messages)
assert controller.finish(messages) is None
def test_independent_request_exposes_context_without_capability_check() -> None:
provider = _provider(compact=False)
messages = [{"role": "user", "content": "hello"}]
controller = ProviderConversationStateController(
provider=provider,
model="gpt-5.6",
messages=messages,
state=_state("saved"),
)
provider_context = controller.independent_request_context(
context_window_tokens=200_000,
)
assert provider_context is not None
assert provider_context.conversation_state is None
assert provider_context.context_window_tokens == 200_000
provider.supports_native_compaction.assert_not_called()

View File

@ -10,6 +10,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from nanobot.providers.base import ProviderCallContext
from nanobot.providers.openai_compat_provider import OpenAICompatProvider
from nanobot.providers.registry import find_by_name
@ -44,8 +45,10 @@ def test_build_responses_body_strips_github_copilot_prefix():
temperature=0.1,
reasoning_effort=None,
tool_choice=None,
provider_context=ProviderCallContext(context_window_tokens=128_000),
)
assert body["model"] == "gpt-5.4-mini"
assert "context_management" not in body
@pytest.mark.asyncio

View File

@ -14,6 +14,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from nanobot.providers.base import ProviderCallContext
from nanobot.providers.openai_compat_provider import OpenAICompatProvider
from nanobot.providers.registry import find_by_name
@ -679,6 +680,7 @@ async def test_direct_openai_gpt5_uses_responses_api() -> None:
assert call_kwargs["max_output_tokens"] == 4096
assert "input" in call_kwargs
assert "messages" not in call_kwargs
assert call_kwargs["include"] == ["reasoning.encrypted_content"]
@pytest.mark.asyncio
@ -710,6 +712,40 @@ async def test_direct_openai_reasoning_prefers_responses_api() -> None:
assert call_kwargs["include"] == ["reasoning.encrypted_content"]
@pytest.mark.asyncio
async def test_direct_openai_retries_without_unsupported_server_compaction() -> None:
mock_chat = AsyncMock(return_value=_fake_chat_response())
mock_responses = AsyncMock(side_effect=[
_FakeResponsesError(400, "Unknown parameter: context_management"),
_fake_responses_response("compaction fallback"),
])
spec = find_by_name("openai")
with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as mock_client_class:
client_instance = mock_client_class.return_value
client_instance.chat.completions.create = mock_chat
client_instance.responses.create = mock_responses
provider = OpenAICompatProvider(
api_key="sk-test-key",
default_model="gpt-5.6",
spec=spec,
)
result = await provider.chat_with_context(
messages=[{"role": "user", "content": "hello"}],
model="gpt-5.6",
provider_context=ProviderCallContext(context_window_tokens=200_000),
)
assert result.content == "compaction fallback"
assert result.provider_state is not None
assert mock_responses.await_count == 2
assert "context_management" in mock_responses.call_args_list[0].kwargs
assert "context_management" not in mock_responses.call_args_list[1].kwargs
assert provider.supports_native_compaction("gpt-5.6") is False
mock_chat.assert_not_awaited()
@pytest.mark.asyncio
async def test_direct_openai_gpt4o_stays_on_chat_completions() -> None:
mock_chat = AsyncMock(return_value=_fake_chat_response())

View File

@ -20,6 +20,7 @@ from nanobot.providers.openai_codex_provider import (
_request_codex,
_should_retry_status,
)
from nanobot.providers.openai_responses import build_responses_state
from nanobot.providers.registry import find_by_name
@ -115,6 +116,48 @@ async def test_codex_request_non_200_populates_http_metadata(monkeypatch) -> Non
assert error.should_retry is True
@pytest.mark.asyncio
async def test_codex_request_marks_rejected_compaction_without_retaining_raw_body(
monkeypatch,
) -> None:
original_client = httpx.AsyncClient
secret = "PRIVATE PROMPT MUST NOT BE RETAINED"
def handler(request: httpx.Request) -> httpx.Response:
return httpx.Response(
400,
json={
"error": {
"message": f"Unknown input type compaction_trigger; {secret}",
},
},
request=request,
)
def fake_client(
*,
timeout: int,
verify: bool,
**_kwargs: object,
) -> httpx.AsyncClient:
return original_client(transport=httpx.MockTransport(handler), timeout=timeout)
monkeypatch.setattr("nanobot.providers.openai_codex_provider.httpx.AsyncClient", fake_client)
with pytest.raises(_CodexHTTPError) as caught:
await _request_codex(
"https://codex.example/responses",
{},
{"input": [{"type": "compaction_trigger"}]},
verify=True,
)
error = caught.value
assert error.compaction_unsupported is True
assert secret not in str(error)
assert not hasattr(error, "body")
@pytest.mark.asyncio
async def test_codex_request_honors_stream_idle_timeout_env(monkeypatch) -> None:
"""NANOBOT_STREAM_IDLE_TIMEOUT_S overrides the default Codex stream timeout."""
@ -192,7 +235,7 @@ async def test_codex_prompt_cache_key_uses_stable_conversation_prefix(monkeypatc
):
_ = proxy, on_thinking_delta, on_tool_call_delta
bodies.append(body)
return "ok", [], "stop", {}, None
return provider_base.LLMResponse(content="ok")
monkeypatch.setattr("nanobot.providers.openai_codex_provider._request_codex", fake_request)
@ -232,7 +275,7 @@ async def test_codex_provider_applies_extra_body_from_config(monkeypatch) -> Non
async def fake_request(_url, _headers, body, **_kwargs):
bodies.append(body)
return "ok", [], "stop", {}, None
return provider_base.LLMResponse(content="ok")
monkeypatch.setattr("nanobot.providers.openai_codex_provider._request_codex", fake_request)
config = Config.model_validate({
@ -297,7 +340,7 @@ async def test_codex_provider_passes_proxy_to_oauth_and_response_request(monkeyp
):
_ = url, headers, body, verify, on_content_delta, on_thinking_delta, on_tool_call_delta
seen["request_proxy"] = proxy
return "ok", [], "stop", {}, None
return provider_base.LLMResponse(content="ok")
monkeypatch.setattr("nanobot.providers.openai_codex_provider.get_codex_token", fake_token)
monkeypatch.setattr("nanobot.providers.openai_codex_provider._request_codex", fake_request)
@ -384,7 +427,7 @@ async def test_codex_retry_uses_structured_timeout_metadata(monkeypatch) -> None
calls += 1
if calls == 1:
raise httpx.ReadTimeout("")
return "ok", [], "stop", {}, None
return provider_base.LLMResponse(content="ok")
async def fake_sleep(delay: float) -> None:
delays.append(delay)
@ -533,6 +576,254 @@ def test_codex_reasoning_options_request_summary_without_forcing_effort() -> Non
assert _build_reasoning_options("none") == {"effort": "none"}
@pytest.mark.asyncio
async def test_codex_replayed_tool_turn_omits_server_item_ids(monkeypatch) -> None:
_mock_codex_token(monkeypatch)
provider = OpenAICodexProvider(default_model="openai-codex/gpt-5.6-sol")
state = build_responses_state(
provider=provider._responses_state_provider(),
model="gpt-5.6-sol",
input_items=[{
"id": "msg_user",
"type": "message",
"role": "user",
"content": [{"type": "input_text", "text": "Check the weather"}],
}],
output_items=[
{
"id": "rs_reasoning",
"type": "reasoning",
"encrypted_content": "opaque reasoning",
"summary": [],
},
{
"id": "fc_read",
"type": "function_call",
"call_id": "call_read",
"name": "read_file",
"arguments": '{"path":"weather/SKILL.md"}',
"status": "completed",
},
],
)
bodies: list[dict[str, Any]] = []
async def fake_request(
url,
headers,
body,
verify,
proxy=None,
on_content_delta=None,
on_thinking_delta=None,
on_tool_call_delta=None,
):
bodies.append(body)
return provider_base.LLMResponse(content="done")
monkeypatch.setattr(
"nanobot.providers.openai_codex_provider._request_codex",
fake_request,
)
response = await provider.chat(
[{"role": "user", "content": "Check the weather"}],
provider_context=provider_base.ProviderCallContext(
conversation_state=state.with_pending_messages([{
"role": "tool",
"tool_call_id": "call_read|fc_read",
"content": "weather skill contents",
}]),
),
)
assert response.content == "done"
assert len(bodies) == 1
input_items = bodies[0]["input"]
assert [item.get("type") for item in input_items] == [
"message",
"reasoning",
"function_call",
"function_call_output",
]
assert all("id" not in item for item in input_items)
assert input_items[1]["encrypted_content"] == "opaque reasoning"
assert input_items[2]["call_id"] == "call_read"
assert input_items[3]["call_id"] == "call_read"
@pytest.mark.asyncio
async def test_codex_compacts_state_at_ninety_percent_before_next_request(
monkeypatch,
) -> None:
_mock_codex_token(monkeypatch)
provider = OpenAICodexProvider(default_model="openai-codex/gpt-5.6-sol")
state_provider = provider._responses_state_provider()
state = build_responses_state(
provider=state_provider,
model="gpt-5.6-sol",
input_items=[{"type": "message", "role": "user", "content": "old question"}],
output_items=[
{"type": "reasoning", "encrypted_content": "old opaque reasoning"},
{
"type": "message",
"role": "assistant",
"content": [{"type": "output_text", "text": "old answer"}],
},
],
usage={
"prompt_tokens": 90,
"completion_tokens": 5,
"total_tokens": 95,
},
)
bodies: list[dict[str, Any]] = []
async def fake_request(
url,
headers,
body,
verify,
proxy=None,
on_content_delta=None,
on_thinking_delta=None,
on_tool_call_delta=None,
):
_ = (
url,
headers,
verify,
proxy,
on_content_delta,
on_thinking_delta,
on_tool_call_delta,
)
bodies.append(body)
if body["input"][-1].get("type") == "compaction_trigger":
compact_item = {
"type": "compaction",
"encrypted_content": "compacted opaque state",
}
return provider_base.LLMResponse(
content=None,
provider_state=build_responses_state(
provider=state_provider,
model="gpt-5.6-sol",
input_items=body["input"],
output_items=[compact_item],
usage={
"prompt_tokens": 95,
"completion_tokens": 2,
"total_tokens": 97,
},
),
)
return provider_base.LLMResponse(content="done")
monkeypatch.setattr(
"nanobot.providers.openai_codex_provider._request_codex",
fake_request,
)
response = await provider.chat_with_retry(
[
{"role": "system", "content": "system"},
{"role": "user", "content": "new question"},
],
max_tokens=5,
provider_context=provider_base.ProviderCallContext(
conversation_state=state.with_pending_messages([
{"role": "user", "content": "new question"},
]),
context_window_tokens=100,
),
)
assert response.content == "done"
assert len(bodies) == 2
assert bodies[0]["input"][-1] == {"type": "compaction_trigger"}
assert bodies[1]["input"][-1] == {
"type": "compaction",
"encrypted_content": "compacted opaque state",
}
assert not any(
item.get("type") == "reasoning"
for item in bodies[1]["input"]
)
assert any(
item.get("role") == "user"
and "new question" in str(item.get("content"))
for item in bodies[1]["input"]
)
@pytest.mark.asyncio
async def test_codex_disables_unsupported_native_compaction_and_continues(
monkeypatch,
) -> None:
_mock_codex_token(monkeypatch)
provider = OpenAICodexProvider(default_model="openai-codex/gpt-5.6-sol")
state_provider = provider._responses_state_provider()
state = build_responses_state(
provider=state_provider,
model="gpt-5.6-sol",
input_items=[{"type": "message", "role": "user", "content": "old"}],
output_items=[{"type": "reasoning", "encrypted_content": "opaque"}],
usage={"prompt_tokens": 90, "completion_tokens": 5, "total_tokens": 95},
)
bodies: list[dict[str, Any]] = []
async def fake_request(
url,
headers,
body,
verify,
proxy=None,
on_content_delta=None,
on_thinking_delta=None,
on_tool_call_delta=None,
):
_ = (
url,
headers,
verify,
proxy,
on_content_delta,
on_thinking_delta,
on_tool_call_delta,
)
bodies.append(body)
if body["input"][-1].get("type") == "compaction_trigger":
raise _CodexHTTPError(
"HTTP 400: Codex API request failed",
status_code=400,
compaction_unsupported=True,
)
return provider_base.LLMResponse(content="done")
monkeypatch.setattr(
"nanobot.providers.openai_codex_provider._request_codex",
fake_request,
)
response = await provider.chat(
[{"role": "user", "content": "new"}],
max_tokens=5,
provider_context=provider_base.ProviderCallContext(
conversation_state=state.with_pending_messages([
{"role": "user", "content": "new"},
]),
context_window_tokens=100,
),
)
assert response.content == "done"
assert len(bodies) == 2
assert bodies[0]["input"][-1] == {"type": "compaction_trigger"}
assert bodies[1]["input"][-1] != {"type": "compaction_trigger"}
assert provider.supports_native_compaction() is False
@pytest.mark.asyncio
async def test_codex_stream_surfaces_reasoning_summary(monkeypatch) -> None:
def fake_token(**_kwargs):
@ -559,7 +850,12 @@ async def test_codex_stream_surfaces_reasoning_summary(monkeypatch) -> None:
await on_content_delta("answer")
if on_thinking_delta:
await on_thinking_delta("summary")
return "answer", [], "stop", {"prompt_tokens": 10, "completion_tokens": 5}, "summary"
return provider_base.LLMResponse(
content="answer",
finish_reason="stop",
usage={"prompt_tokens": 10, "completion_tokens": 5},
reasoning_content="summary",
)
monkeypatch.setattr("nanobot.providers.openai_codex_provider._request_codex", fake_request)

View File

@ -1,9 +1,11 @@
"""Tests for the shared openai_responses converters and parsers."""
import json
from io import StringIO
from unittest.mock import MagicMock, patch
import pytest
from loguru import logger
from nanobot.providers.openai_responses.converters import (
convert_messages,
@ -12,12 +14,22 @@ from nanobot.providers.openai_responses.converters import (
split_tool_call_id,
)
from nanobot.providers.openai_responses.parsing import (
ResponsesStreamCapture,
consume_sdk_stream,
consume_sse,
consume_sse_with_reasoning,
is_replayable_finish_reason,
map_finish_reason,
parse_response_output,
)
from nanobot.providers.openai_responses.state import (
build_responses_state,
is_compaction_compatibility_error,
prepare_responses_input,
resolve_compact_threshold,
responses_state_context_tokens,
responses_state_items,
)
# ======================================================================
# converters - split_tool_call_id
@ -398,6 +410,17 @@ class TestMapFinishReason:
def test_unknown_defaults_to_stop(self):
assert map_finish_reason("some_new_status") == "stop"
@pytest.mark.parametrize("finish_reason", ["stop", "tool_calls", "function_call"])
def test_replayable_finish_reasons(self, finish_reason):
assert is_replayable_finish_reason(finish_reason) is True
@pytest.mark.parametrize(
"finish_reason",
["length", "refusal", "content_filter", "error"],
)
def test_non_replayable_finish_reasons(self, finish_reason):
assert is_replayable_finish_reason(finish_reason) is False
# ======================================================================
# parsing - parse_response_output
@ -418,6 +441,29 @@ class TestParseResponseOutput:
assert result.usage == {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}
assert result.tool_calls == []
def test_refusal_response_surfaces_text_without_advancing_state(self):
refusal = "I cant help with that request."
resp = {
"output": [{
"type": "message",
"role": "assistant",
"content": [{"type": "refusal", "refusal": refusal}],
}],
"status": "completed",
"usage": {},
}
result = parse_response_output(
resp,
state_provider="openai:test",
state_model="gpt-5.6",
state_input_items=[{"role": "user", "content": "request"}],
)
assert result.content == refusal
assert result.finish_reason == "refusal"
assert result.provider_state is None
def test_tool_call_response(self):
resp = {
"output": [{
@ -429,12 +475,18 @@ class TestParseResponseOutput:
"status": "completed",
"usage": {},
}
result = parse_response_output(resp)
result = parse_response_output(
resp,
state_provider="openai:test",
state_model="gpt-5.6",
state_input_items=[{"role": "user", "content": "weather?"}],
)
assert result.content is None
assert len(result.tool_calls) == 1
assert result.tool_calls[0].name == "get_weather"
assert result.tool_calls[0].arguments == {"city": "SF"}
assert result.tool_calls[0].id == "call_1|fc_1"
assert result.provider_state is not None
def test_malformed_tool_arguments_logged(self):
"""Malformed JSON arguments should log a warning and remain non-object."""
@ -493,10 +545,39 @@ class TestParseResponseOutput:
assert result.content is None
assert result.tool_calls == []
def test_incomplete_status(self):
resp = {"output": [], "status": "incomplete", "usage": {}}
result = parse_response_output(resp)
assert result.finish_reason == "length"
@pytest.mark.parametrize(
("reason", "expected_finish_reason"),
[
("max_output_tokens", "length"),
("content_filter", "content_filter"),
],
)
def test_incomplete_status(self, reason, expected_finish_reason):
resp = {
"output": [],
"status": "incomplete",
"incomplete_details": {"reason": reason},
"usage": {},
}
result = parse_response_output(
resp,
state_provider="openai:test",
state_model="gpt-5.6",
state_input_items=[{"role": "user", "content": "prompt"}],
)
assert result.finish_reason == expected_finish_reason
assert result.provider_state is None
def test_unknown_status_does_not_advance_provider_state(self):
result = parse_response_output(
{"output": [], "status": "future_terminal_status", "usage": {}},
state_provider="openai:test",
state_model="gpt-5.6",
state_input_items=[{"role": "user", "content": "prompt"}],
)
assert result.finish_reason == "stop"
assert result.provider_state is None
def test_sdk_model_object(self):
"""parse_response_output should handle SDK objects with model_dump()."""
@ -523,6 +604,194 @@ class TestParseResponseOutput:
assert result.usage["completion_tokens"] == 50
assert result.usage["total_tokens"] == 150
def test_preserves_every_output_item_as_opaque_state(self):
input_items = [{"role": "user", "content": "inspect the repo"}]
output = [
{
"id": "rs_1",
"type": "reasoning",
"encrypted_content": "opaque-secret",
"summary": [],
},
{
"id": "future_1",
"type": "future_item_type",
"provider_field": {"nested": True},
},
{
"id": "msg_1",
"type": "message",
"role": "assistant",
"status": "completed",
"content": [{"type": "output_text", "text": "done"}],
},
]
result = parse_response_output(
{"output": output, "status": "completed", "usage": {}},
state_provider="openai:test",
state_model="gpt-5.6",
state_input_items=input_items,
)
assert result.provider_state is not None
assert responses_state_items(result.provider_state) == [*input_items, *output]
class TestResponsesConversationState:
def test_server_compaction_prunes_superseded_prefix(self):
state = build_responses_state(
provider="openai:test",
model="gpt-5.6",
input_items=[
{"type": "message", "role": "user", "content": "old"},
{"type": "reasoning", "encrypted_content": "old-reasoning"},
],
output_items=[
{"type": "compaction", "encrypted_content": "compact"},
{"type": "message", "role": "assistant", "content": "new"},
],
usage={
"prompt_tokens": 90,
"completion_tokens": 10,
"total_tokens": 100,
},
)
assert responses_state_items(state) == [
{"type": "compaction", "encrypted_content": "compact"},
{"type": "message", "role": "assistant", "content": "new"},
]
assert responses_state_context_tokens(state) == 100
def test_existing_compaction_keeps_canonical_retained_prefix(self):
canonical_input = [
{"type": "message", "role": "user", "content": "retained"},
{"type": "compaction", "encrypted_content": "compact"},
]
output = [{"type": "message", "role": "assistant", "content": "new"}]
state = build_responses_state(
provider="openai:test",
model="gpt-5.6",
input_items=canonical_input,
output_items=output,
)
assert responses_state_items(state) == [*canonical_input, *output]
@pytest.mark.parametrize(
("context_window", "max_output", "expected"),
[
(200_000, 20_000, 180_000),
(100_000, 30_000, 70_000),
(0, 4_096, None),
],
)
def test_compact_threshold_reserves_codex_style_headroom(
self,
context_window,
max_output,
expected,
):
assert resolve_compact_threshold(context_window, max_output) == expected
def test_compaction_compatibility_recognizes_old_sdk_signature_error(self):
error = TypeError("create() got an unexpected keyword argument 'context_management'")
assert is_compaction_compatibility_error(error) is True
assert is_compaction_compatibility_error(TypeError("unrelated argument")) is False
def test_state_observability_logs_counts_without_opaque_content(self):
secret = "opaque-secret-that-must-not-be-logged"
state = build_responses_state(
provider=f"openai:https://example.test/?key={secret}",
model=f"secret-model-{secret}",
input_items=[{"role": "user", "content": secret}],
output_items=[{"type": "reasoning", "encrypted_content": secret}],
).with_pending_messages([{"role": "user", "content": secret}])
sink = StringIO()
sink_id = logger.add(sink, level="DEBUG", format="{message}")
try:
prepare_responses_input(
[{"role": "user", "content": secret}],
state=state,
provider=state.provider,
model=state.model,
)
build_responses_state(
provider=state.provider,
model=state.model,
input_items=[
{"role": "user", "content": secret},
{"type": "reasoning", "encrypted_content": secret},
],
output_items=[
{"type": "compaction", "encrypted_content": secret},
],
)
finally:
logger.remove(sink_id)
log_text = sink.getvalue()
assert "prior_items=2" in log_text
assert "pending_messages=1" in log_text
assert "dropped_items=2" in log_text
assert secret not in log_text
def test_replays_exact_items_then_only_pending_and_new_messages(self):
prior_items = [
{"role": "user", "content": "first"},
{
"type": "reasoning",
"id": "rs_1",
"encrypted_content": "opaque-secret",
},
{
"type": "function_call",
"id": "fc_1",
"call_id": "call_1",
"name": "read_file",
"arguments": '{"path":"a.py"}',
},
]
state = build_responses_state(
provider="openai:test",
model="gpt-5.6",
input_items=prior_items[:1],
output_items=prior_items[1:],
).with_pending_messages([
{
"role": "tool",
"tool_call_id": "call_1|fc_1",
"content": "file contents",
},
{"role": "user", "content": "continue"},
])
instructions, items, replayed = prepare_responses_input(
[
{"role": "system", "content": "current instructions"},
{"role": "user", "content": "a lossy public transcript"},
],
state=state,
provider="openai:test",
model="gpt-5.6",
)
assert instructions == "current instructions"
assert replayed is True
assert items[:3] == prior_items
assert items[3] == {
"type": "function_call_output",
"call_id": "call_1",
"output": "file contents",
}
assert items[4] == {
"role": "user",
"content": [{"type": "input_text", "text": "continue"}],
}
assert "lossy public transcript" not in str(items)
# ======================================================================
# parsing - consume_sse
@ -553,6 +822,122 @@ class TestConsumeSse:
assert tool_calls == []
assert finish_reason == "stop"
@pytest.mark.asyncio
async def test_refusal_events_reconcile_parts_and_terminal_output(self):
refusal = "First and second sentence. Done-only. Terminal suffix."
terminal_response = {
"status": "completed",
"output": [{
"type": "message",
"id": "msg_2",
"role": "assistant",
"content": [{"type": "refusal", "refusal": refusal}],
}],
}
response = _SseResponse([
{
"type": "response.refusal.delta",
"item_id": "msg_1",
"content_index": 0,
"delta": "First",
},
{
"type": "response.refusal.delta",
"item_id": "msg_1",
"content_index": 1,
"delta": " and second",
},
{
"type": "response.refusal.done",
"item_id": "msg_1",
"content_index": 0,
"refusal": "First",
},
{
"type": "response.refusal.done",
"item_id": "msg_1",
"content_index": 1,
"refusal": " and second sentence.",
},
{
"type": "response.refusal.done",
"item_id": "msg_2",
"content_index": 0,
"refusal": " Done-only.",
},
{
"type": "response.refusal.delta",
"item_id": "msg_2",
"content_index": 1,
"delta": " Terminal",
},
{"type": "response.completed", "response": terminal_response},
])
capture = ResponsesStreamCapture()
deltas: list[str] = []
async def on_content(delta: str) -> None:
deltas.append(delta)
content, _, finish_reason, _, _ = await consume_sse_with_reasoning(
response,
on_content_delta=on_content,
capture=capture,
)
assert content == refusal
assert deltas == [
"First",
" and second",
" sentence.",
" Done-only.",
" Terminal",
" suffix.",
]
assert finish_reason == "refusal"
assert capture.completed is True
assert is_replayable_finish_reason(finish_reason) is False
@pytest.mark.asyncio
@pytest.mark.parametrize("source", ["events", "terminal"])
async def test_refusal_without_deltas_has_non_replayable_finish(self, source: str):
refusal = "I cant help with that request."
terminal_response = {
"status": "completed",
"output": [{
"type": "message",
"id": "msg_1",
"role": "assistant",
"content": [{"type": "refusal", "refusal": refusal}],
}],
}
events = (
[
{"type": "response.refusal.done", "refusal": refusal},
{"type": "response.completed", "response": {"status": "completed"}},
]
if source == "events"
else [{"type": "response.completed", "response": terminal_response}]
)
response = _SseResponse(events)
capture = ResponsesStreamCapture()
deltas: list[str] = []
async def on_content(delta: str) -> None:
deltas.append(delta)
content, _, finish_reason, _, _ = await consume_sse_with_reasoning(
response,
on_content_delta=on_content,
capture=capture,
)
assert content == refusal
assert deltas == [refusal]
assert finish_reason == "refusal"
assert capture.completed is True
assert is_replayable_finish_reason(finish_reason) is False
@pytest.mark.asyncio
async def test_reasoning_summary_delta_extracted(self):
response = _SseResponse([
@ -599,6 +984,139 @@ class TestConsumeSse:
assert reasoning == "cached summary"
@pytest.mark.asyncio
async def test_capture_commits_exact_items_only_after_completed_event(self):
output = [
{
"type": "reasoning",
"id": "rs_1",
"encrypted_content": "opaque-secret",
},
{"type": "future_item_type", "id": "future_1", "value": 7},
]
capture = ResponsesStreamCapture()
response = _SseResponse([
{
"type": "response.output_item.done",
"output_index": 0,
"item": output[0],
},
{
"type": "response.output_item.done",
"output_index": 1,
"item": output[1],
},
{
"type": "response.completed",
"response": {"status": "completed", "output": output},
},
])
await consume_sse_with_reasoning(response, capture=capture)
assert capture.completed is True
assert capture.output_items == output
@pytest.mark.asyncio
async def test_capture_keeps_done_items_when_completed_output_is_empty(self):
output = [
{
"type": "reasoning",
"id": "rs_1",
"encrypted_content": "opaque-secret",
"summary": [],
},
{
"type": "function_call",
"id": "fc_1",
"call_id": "call_1",
"name": "read_file",
"arguments": '{"path":"weather/SKILL.md"}',
},
]
capture = ResponsesStreamCapture()
response = _SseResponse([
{
"type": "response.output_item.done",
"output_index": index,
"item": item,
}
for index, item in enumerate(output)
] + [{
"type": "response.completed",
"response": {"status": "completed", "output": []},
}])
await consume_sse_with_reasoning(response, capture=capture)
assert capture.completed is True
assert capture.output_items == output
@pytest.mark.asyncio
@pytest.mark.parametrize(
("reason", "expected_finish_reason"),
[
("max_output_tokens", "length"),
("content_filter", "content_filter"),
],
)
async def test_incomplete_event_commits_capture_usage(
self,
reason,
expected_finish_reason,
):
output = [
{
"type": "message",
"id": "msg_1",
"status": "incomplete",
"content": [{"type": "output_text", "text": "partial"}],
},
]
terminal_response = {
"id": "resp_1",
"status": "incomplete",
"incomplete_details": {"reason": reason},
"output": output,
"usage": {"input_tokens": 10, "output_tokens": 5, "total_tokens": 15},
}
capture = ResponsesStreamCapture()
response = _SseResponse([
{"type": "response.output_text.delta", "delta": "partial"},
{"type": "response.incomplete", "response": terminal_response},
])
content, _, finish_reason, usage, _ = await consume_sse_with_reasoning(
response,
capture=capture,
)
assert content == "partial"
assert finish_reason == expected_finish_reason
assert usage == {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}
assert capture.completed is True
assert capture.response == terminal_response
assert capture.output_items == output
@pytest.mark.asyncio
async def test_capture_does_not_commit_interrupted_stream(self):
capture = ResponsesStreamCapture()
response = _SseResponse([
{
"type": "response.output_item.done",
"output_index": 0,
"item": {
"type": "reasoning",
"id": "rs_1",
"encrypted_content": "opaque-secret",
},
},
])
await consume_sse_with_reasoning(response, capture=capture)
assert capture.completed is False
@pytest.mark.asyncio
async def test_reasoning_summary_from_done_item(self):
response = _SseResponse([
@ -755,6 +1273,131 @@ class TestConsumeSdkStream:
assert tool_calls == []
assert finish_reason == "stop"
@pytest.mark.asyncio
async def test_refusal_events_reconcile_parts_and_terminal_output(self):
refusal = "First and second sentence. Done-only. Terminal suffix."
terminal_response = {
"status": "completed",
"output": [{
"type": "message",
"id": "msg_2",
"role": "assistant",
"content": [{"type": "refusal", "refusal": refusal}],
}],
}
resp_obj = MagicMock(status="completed", usage=None, output=[])
resp_obj.model_dump.return_value = terminal_response
events = [
MagicMock(
type="response.refusal.delta",
item_id="msg_1",
content_index=0,
delta="First",
),
MagicMock(
type="response.refusal.delta",
item_id="msg_1",
content_index=1,
delta=" and second",
),
MagicMock(
type="response.refusal.done",
item_id="msg_1",
content_index=0,
refusal="First",
),
MagicMock(
type="response.refusal.done",
item_id="msg_1",
content_index=1,
refusal=" and second sentence.",
),
MagicMock(
type="response.refusal.done",
item_id="msg_2",
content_index=0,
refusal=" Done-only.",
),
MagicMock(
type="response.refusal.delta",
item_id="msg_2",
content_index=1,
delta=" Terminal",
),
MagicMock(type="response.completed", response=resp_obj),
]
capture = ResponsesStreamCapture()
deltas: list[str] = []
async def on_content(delta: str) -> None:
deltas.append(delta)
async def stream():
for event in events:
yield event
content, _, finish_reason, _, _ = await consume_sdk_stream(
stream(),
on_content_delta=on_content,
capture=capture,
)
assert content == refusal
assert deltas == [
"First",
" and second",
" sentence.",
" Done-only.",
" Terminal",
" suffix.",
]
assert finish_reason == "refusal"
assert capture.completed is True
assert is_replayable_finish_reason(finish_reason) is False
@pytest.mark.asyncio
@pytest.mark.parametrize("source", ["events", "terminal"])
async def test_refusal_without_deltas_has_non_replayable_finish(self, source: str):
refusal = "I cant help with that request."
terminal_response = {
"status": "completed",
"output": [{
"type": "message",
"id": "msg_1",
"role": "assistant",
"content": [{"type": "refusal", "refusal": refusal}],
}],
}
resp_obj = MagicMock(status="completed", usage=None, output=[])
resp_obj.model_dump.return_value = terminal_response
capture = ResponsesStreamCapture()
deltas: list[str] = []
async def on_content(delta: str) -> None:
deltas.append(delta)
async def stream():
if source == "events":
yield MagicMock(type="response.refusal.done", refusal=refusal)
yield MagicMock(
type="response.completed",
response={"status": "completed"},
)
else:
yield MagicMock(type="response.completed", response=resp_obj)
content, _, finish_reason, _, _ = await consume_sdk_stream(
stream(),
on_content_delta=on_content,
capture=capture,
)
assert content == refusal
assert deltas == [refusal]
assert finish_reason == "refusal"
assert capture.completed is True
assert is_replayable_finish_reason(finish_reason) is False
@pytest.mark.asyncio
async def test_on_content_delta_called(self):
ev1 = MagicMock(type="response.output_text.delta", delta="hi")
@ -919,6 +1562,64 @@ class TestConsumeSdkStream:
_, _, _, usage, _ = await consume_sdk_stream(stream())
assert usage == {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}
@pytest.mark.asyncio
@pytest.mark.parametrize(
("reason", "expected_finish_reason"),
[
("max_output_tokens", "length"),
("content_filter", "content_filter"),
],
)
async def test_incomplete_event_commits_capture_usage(
self,
reason,
expected_finish_reason,
):
output = [
{
"type": "message",
"id": "msg_1",
"status": "incomplete",
"content": [{"type": "output_text", "text": "partial"}],
},
]
usage_obj = MagicMock(input_tokens=10, output_tokens=5, total_tokens=15)
output_item = MagicMock(type="message")
terminal_response = {
"id": "resp_1",
"status": "incomplete",
"incomplete_details": {"reason": reason},
"output": output,
"usage": {"input_tokens": 10, "output_tokens": 5, "total_tokens": 15},
}
resp_obj = MagicMock(
status="incomplete",
usage=usage_obj,
output=[output_item],
)
resp_obj.model_dump.return_value = terminal_response
events = [
MagicMock(type="response.output_text.delta", delta="partial"),
MagicMock(type="response.incomplete", response=resp_obj),
]
capture = ResponsesStreamCapture()
async def stream():
for event in events:
yield event
content, _, finish_reason, usage, _ = await consume_sdk_stream(
stream(),
capture=capture,
)
assert content == "partial"
assert finish_reason == expected_finish_reason
assert usage == {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}
assert capture.completed is True
assert capture.response == terminal_response
assert capture.output_items == output
@pytest.mark.asyncio
async def test_reasoning_extracted(self):
summary_item = MagicMock(type="summary_text", text="thinking...")

View File

@ -3,7 +3,14 @@ import copy
import pytest
from nanobot.providers.base import RETRY_AFTER_BUFFER, GenerationSettings, LLMProvider, LLMResponse
from nanobot.providers.base import (
RETRY_AFTER_BUFFER,
GenerationSettings,
LLMProvider,
LLMResponse,
ProviderCallContext,
ProviderConversationState,
)
class ScriptedProvider(LLMProvider):
@ -330,6 +337,79 @@ async def test_successful_image_retry_mutates_original_messages_in_place() -> No
assert any("not delivered" in (block.get("text") or "").lower() for block in content)
@pytest.mark.asyncio
@pytest.mark.parametrize(
("messages", "payload", "pending_messages"),
[
(_IMAGE_MSG, {}, _IMAGE_MSG),
(
[{"role": "user", "content": "continue"}],
{
"items": [
{
"type": "message",
"role": "user",
"content": [
{
"type": "input_image",
"image_url": "data:image/png;base64,abc",
}
],
}
]
},
[],
),
],
ids=["pending-image", "opaque-payload-image"],
)
async def test_image_retry_discards_provider_state_with_images(
messages,
payload,
pending_messages,
) -> None:
class ContextScriptedProvider(ScriptedProvider):
def __init__(self, responses):
super().__init__(responses)
self.contexts: list[ProviderCallContext] = []
async def chat_with_context(
self,
*,
provider_context: ProviderCallContext,
**kwargs,
) -> LLMResponse:
self.contexts.append(provider_context)
return await self.chat(**kwargs)
provider = ContextScriptedProvider([
LLMResponse(content="model does not support images", finish_reason="error"),
LLMResponse(content="ok, no image"),
])
messages = copy.deepcopy(messages)
state = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="gpt-5.6",
version=1,
payload=copy.deepcopy(payload),
pending_messages=copy.deepcopy(pending_messages),
)
response = await provider.chat_with_retry(
messages=messages,
provider_context=ProviderCallContext(conversation_state=state),
)
assert response.content == "ok, no image"
retry_context = provider.contexts[-1]
assert isinstance(retry_context, ProviderCallContext)
assert retry_context.conversation_state is None
public_content = messages[0]["content"]
if isinstance(public_content, list):
assert all(block.get("type") != "image_url" for block in public_content)
@pytest.mark.asyncio
async def test_non_transient_error_without_images_no_retry() -> None:
"""Non-transient errors without image content are returned immediately."""

View File

@ -4,6 +4,7 @@ import time
import pytest
from nanobot.providers.base import ProviderCallContext
from nanobot.providers.openai_compat_provider import (
_RESPONSES_FAILURE_THRESHOLD,
_RESPONSES_PROBE_INTERVAL_S,
@ -28,6 +29,26 @@ def test_responses_api_available_by_default(provider):
assert provider._should_use_responses_api("gpt-5", None) is True
def test_direct_openai_enables_server_compaction(provider):
provider._extra_body = {}
body = provider._build_responses_body(
messages=[{"role": "user", "content": "hello"}],
tools=None,
model="gpt-5.6",
max_tokens=30_000,
temperature=0.1,
reasoning_effort="high",
tool_choice=None,
provider_context=ProviderCallContext(context_window_tokens=100_000),
)
assert body["context_management"] == [{
"type": "compaction",
"compact_threshold": 70_000,
}]
def test_api_type_chat_completions_disables_responses(provider):
provider._api_type = "chat_completions"
assert provider._should_use_responses_api("gpt-5", None) is False

View File

@ -8,6 +8,7 @@ import pytest
import nanobot.webui.session_list_index as session_list_index
from nanobot.cron.session_turns import CRON_HISTORY_META
from nanobot.providers.base import ProviderConversationState
from nanobot.session.automation_turns import AUTOMATION_HISTORY_META
from nanobot.session.history_visibility import HIDDEN_HISTORY_META
from nanobot.session.manager import SessionManager
@ -85,6 +86,26 @@ def test_webui_session_list_rescans_only_changed_file(tmp_path: Path, monkeypatc
assert {row["preview"] for row in rows} == {"first", "second after"}
def test_webui_session_list_skips_provider_state_before_preview_budget(
tmp_path: Path,
monkeypatch,
) -> None:
monkeypatch.setattr(session_list_index, "_SESSION_LIST_PREVIEW_MAX_CHARS", 100)
manager = SessionManager(tmp_path)
session = manager.get_or_create("websocket:private-state")
session.provider_state = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="test-model",
version=1,
payload={"items": [{"encrypted_content": "x" * 200}]},
)
session.add_message("user", "visible preview")
manager.save(session)
assert list_webui_sessions(manager)[0]["preview"] == "visible preview"
def test_webui_session_list_drops_deleted_index_rows(tmp_path: Path) -> None:
manager = SessionManager(tmp_path)
session = manager.get_or_create("websocket:deleted")