From 6a1a45d07a6de420ba87c419ae30fcb4af76d4d0 Mon Sep 17 00:00:00 2001 From: chengyongru <61816729+chengyongru@users.noreply.github.com> Date: Thu, 30 Jul 2026 22:39:43 +0800 Subject: [PATCH] feat: preserve Responses reasoning state and compact context (#5172) --- docs/configuration.md | 14 + docs/providers.md | 4 +- nanobot/agent/context.py | 41 +- nanobot/agent/loop.py | 122 ++- nanobot/agent/memory.py | 3 + nanobot/agent/runner.py | 194 ++++- nanobot/providers/azure_openai_provider.py | 186 ++++- nanobot/providers/base.py | 209 ++++- nanobot/providers/conversation_state.py | 262 +++++++ nanobot/providers/factory.py | 1 + nanobot/providers/fallback_provider.py | 115 ++- nanobot/providers/github_copilot_provider.py | 6 +- nanobot/providers/openai_codex_provider.py | 314 +++++++- nanobot/providers/openai_compat_provider.py | 162 +++- .../providers/openai_responses/__init__.py | 22 +- nanobot/providers/openai_responses/parsing.py | 256 ++++++- nanobot/providers/openai_responses/state.py | 197 +++++ nanobot/session/manager.py | 58 +- nanobot/webui/session_list_index.py | 6 + tests/agent/test_consolidator.py | 22 +- tests/agent/test_context_builder.py | 14 + tests/agent/test_loop_save_turn.py | 309 +++++++- tests/agent/test_runner_core.py | 423 ++++++++++- tests/agent/test_runner_fallback.py | 285 ++++++- tests/agent/test_runner_governance.py | 15 +- tests/agent/test_session_atomic.py | 132 ++++ tests/agent/test_session_manager_history.py | 23 +- tests/agent/tools/test_subagent_tools.py | 3 + tests/providers/test_azure_openai_provider.py | 62 +- tests/providers/test_conversation_state.py | 291 +++++++ .../providers/test_github_copilot_routing.py | 3 + tests/providers/test_litellm_kwargs.py | 36 + tests/providers/test_openai_codex_provider.py | 306 +++++++- tests/providers/test_openai_responses.py | 711 +++++++++++++++++- tests/providers/test_provider_retry.py | 82 +- .../test_responses_circuit_breaker.py | 21 + tests/webui/test_session_list_index.py | 21 + 37 files changed, 4778 insertions(+), 153 deletions(-) create mode 100644 nanobot/providers/conversation_state.py create mode 100644 nanobot/providers/openai_responses/state.py create mode 100644 tests/providers/test_conversation_state.py diff --git a/docs/configuration.md b/docs/configuration.md index 7658c0855..806017dbe 100644 --- a/docs/configuration.md +++ b/docs/configuration.md @@ -348,6 +348,20 @@ Valid `apiType` values are exactly `auto`, `chat_completions`, and `responses`. + + +### 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. +
Azure OpenAI diff --git a/docs/providers.md b/docs/providers.md index 2d2ea452b..1d0e3a482 100644 --- a/docs/providers.md +++ b/docs/providers.md @@ -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 diff --git a/nanobot/agent/context.py b/nanobot/agent/context.py index 4c9d092e9..031c1ab9e 100644 --- a/nanobot/agent/context.py +++ b/nanobot/agent/context.py @@ -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, diff --git a/nanobot/agent/loop.py b/nanobot/agent/loop.py index 598d6213b..a218451b8 100644 --- a/nanobot/agent/loop.py +++ b/nanobot/agent/loop.py @@ -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) diff --git a/nanobot/agent/memory.py b/nanobot/agent/memory.py index bdcd83534..42b92825f 100644 --- a/nanobot/agent/memory.py +++ b/nanobot/agent/memory.py @@ -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( diff --git a/nanobot/agent/runner.py b/nanobot/agent/runner.py index 6c91a6aef..7d03674b3 100644 --- a/nanobot/agent/runner.py +++ b/nanobot/agent/runner.py @@ -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( diff --git a/nanobot/providers/azure_openai_provider.py b/nanobot/providers/azure_openai_provider.py index 8ce72c0d0..e7d718133 100644 --- a/nanobot/providers/azure_openai_provider.py +++ b/nanobot/providers/azure_openai_provider.py @@ -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) diff --git a/nanobot/providers/base.py b/nanobot/providers/base.py index 642f24177..b8680278a 100644 --- a/nanobot/providers/base.py +++ b/nanobot/providers/base.py @@ -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. diff --git a/nanobot/providers/conversation_state.py b/nanobot/providers/conversation_state.py new file mode 100644 index 000000000..462fc316f --- /dev/null +++ b/nanobot/providers/conversation_state.py @@ -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 diff --git a/nanobot/providers/factory.py b/nanobot/providers/factory.py index 6e3e0706d..37fa2ed8b 100644 --- a/nanobot/providers/factory.py +++ b/nanobot/providers/factory.py @@ -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 diff --git a/nanobot/providers/fallback_provider.py b/nanobot/providers/fallback_provider.py index bf14bdd0a..24cf2785a 100644 --- a/nanobot/providers/fallback_provider.py +++ b/nanobot/providers/fallback_provider.py @@ -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: diff --git a/nanobot/providers/github_copilot_provider.py b/nanobot/providers/github_copilot_provider.py index bac99ccaa..4ae3d9f39 100644 --- a/nanobot/providers/github_copilot_provider.py +++ b/nanobot/providers/github_copilot_provider.py @@ -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, ) diff --git a/nanobot/providers/openai_codex_provider.py b/nanobot/providers/openai_codex_provider.py index 4b9a1d311..fdfbcef7d 100644 --- a/nanobot/providers/openai_codex_provider.py +++ b/nanobot/providers/openai_codex_provider.py @@ -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: diff --git a/nanobot/providers/openai_compat_provider.py b/nanobot/providers/openai_compat_provider.py index 605d8e717..4e52f9bba 100644 --- a/nanobot/providers/openai_compat_provider.py +++ b/nanobot/providers/openai_compat_provider.py @@ -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; diff --git a/nanobot/providers/openai_responses/__init__.py b/nanobot/providers/openai_responses/__init__.py index 25f19afb6..dfc371972 100644 --- a/nanobot/providers/openai_responses/__init__.py +++ b/nanobot/providers/openai_responses/__init__.py @@ -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", ] diff --git a/nanobot/providers/openai_responses/parsing.py b/nanobot/providers/openai_responses/parsing.py index d999910c6..b2de16855 100644 --- a/nanobot/providers/openai_responses/parsing.py +++ b/nanobot/providers/openai_responses/parsing.py @@ -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 diff --git a/nanobot/providers/openai_responses/state.py b/nanobot/providers/openai_responses/state.py new file mode 100644 index 000000000..2424afb67 --- /dev/null +++ b/nanobot/providers/openai_responses/state.py @@ -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 diff --git a/nanobot/session/manager.py b/nanobot/session/manager.py index b26b2380c..0c832de88 100644 --- a/nanobot/session/manager.py +++ b/nanobot/session/manager.py @@ -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: diff --git a/nanobot/webui/session_list_index.py b/nanobot/webui/session_list_index.py index 31dbb79c2..9354406b6 100644 --- a/nanobot/webui/session_list_index.py +++ b/nanobot/webui/session_list_index.py @@ -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) diff --git a/tests/agent/test_consolidator.py b/tests/agent/test_consolidator.py index 51ff4744d..b022dd7b5 100644 --- a/tests/agent/test_consolidator.py +++ b/tests/agent/test_consolidator.py @@ -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" diff --git a/tests/agent/test_context_builder.py b/tests/agent/test_context_builder.py index 0ecb29635..503602c5f 100644 --- a/tests/agent/test_context_builder.py +++ b/tests/agent/test_context_builder.py @@ -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"}] diff --git a/tests/agent/test_loop_save_turn.py b/tests/agent/test_loop_save_turn.py index b613b057e..f923de0c9 100644 --- a/tests/agent/test_loop_save_turn.py +++ b/tests/agent/test_loop_save_turn.py @@ -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 diff --git a/tests/agent/test_runner_core.py b/tests/agent/test_runner_core.py index 883fe24d8..e2da80fd7 100644 --- a/tests/agent/test_runner_core.py +++ b/tests/agent/test_runner_core.py @@ -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.""" diff --git a/tests/agent/test_runner_fallback.py b/tests/agent/test_runner_fallback.py index bf401d7d4..9afe50cbd 100644 --- a/tests/agent/test_runner_fallback.py +++ b/tests/agent/test_runner_fallback.py @@ -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()) diff --git a/tests/agent/test_runner_governance.py b/tests/agent/test_runner_governance.py index 1c2a30279..29379c99d 100644 --- a/tests/agent/test_runner_governance.py +++ b/tests/agent/test_runner_governance.py @@ -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 diff --git a/tests/agent/test_session_atomic.py b/tests/agent/test_session_atomic.py index 1fe5b9caa..df3104f27 100644 --- a/tests/agent/test_session_atomic.py +++ b/tests/agent/test_session_atomic.py @@ -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: diff --git a/tests/agent/test_session_manager_history.py b/tests/agent/test_session_manager_history.py index 6241e1f93..c9da61bbe 100644 --- a/tests/agent/test_session_manager_history.py +++ b/tests/agent/test_session_manager_history.py @@ -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(): diff --git a/tests/agent/tools/test_subagent_tools.py b/tests/agent/tools/test_subagent_tools.py index 82555e659..84a941459 100644 --- a/tests/agent/tools/test_subagent_tools.py +++ b/tests/agent/tools/test_subagent_tools.py @@ -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) diff --git a/tests/providers/test_azure_openai_provider.py b/tests/providers/test_azure_openai_provider.py index df78acfc2..10d2770c2 100644 --- a/tests/providers/test_azure_openai_provider.py +++ b/tests/providers/test_azure_openai_provider.py @@ -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 diff --git a/tests/providers/test_conversation_state.py b/tests/providers/test_conversation_state.py new file mode 100644 index 000000000..157fc2c80 --- /dev/null +++ b/tests/providers/test_conversation_state.py @@ -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() diff --git a/tests/providers/test_github_copilot_routing.py b/tests/providers/test_github_copilot_routing.py index b5dd46670..265a2510d 100644 --- a/tests/providers/test_github_copilot_routing.py +++ b/tests/providers/test_github_copilot_routing.py @@ -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 diff --git a/tests/providers/test_litellm_kwargs.py b/tests/providers/test_litellm_kwargs.py index c62258f21..60b3bd45a 100644 --- a/tests/providers/test_litellm_kwargs.py +++ b/tests/providers/test_litellm_kwargs.py @@ -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()) diff --git a/tests/providers/test_openai_codex_provider.py b/tests/providers/test_openai_codex_provider.py index ea5d0c6e5..56a4b2f79 100644 --- a/tests/providers/test_openai_codex_provider.py +++ b/tests/providers/test_openai_codex_provider.py @@ -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) diff --git a/tests/providers/test_openai_responses.py b/tests/providers/test_openai_responses.py index be615f1a8..0022c0cf5 100644 --- a/tests/providers/test_openai_responses.py +++ b/tests/providers/test_openai_responses.py @@ -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 can’t 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 can’t 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 can’t 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...") diff --git a/tests/providers/test_provider_retry.py b/tests/providers/test_provider_retry.py index a00ad7ee4..dac7e6e60 100644 --- a/tests/providers/test_provider_retry.py +++ b/tests/providers/test_provider_retry.py @@ -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.""" diff --git a/tests/providers/test_responses_circuit_breaker.py b/tests/providers/test_responses_circuit_breaker.py index 2c9782e13..83a588443 100644 --- a/tests/providers/test_responses_circuit_breaker.py +++ b/tests/providers/test_responses_circuit_breaker.py @@ -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 diff --git a/tests/webui/test_session_list_index.py b/tests/webui/test_session_list_index.py index 2dafd993a..47a1164ce 100644 --- a/tests/webui/test_session_list_index.py +++ b/tests/webui/test_session_list_index.py @@ -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")