mirror of
https://github.com/HKUDS/nanobot.git
synced 2026-08-07 09:58:34 +00:00
feat: preserve Responses reasoning state and compact context (#5172)
This commit is contained in:
parent
511c764f45
commit
6a1a45d07a
@ -348,6 +348,20 @@ Valid `apiType` values are exactly `auto`, `chat_completions`, and `responses`.
|
|||||||
|
|
||||||
</details>
|
</details>
|
||||||
|
|
||||||
|
<a id="responses-state-and-compaction"></a>
|
||||||
|
|
||||||
|
### Responses conversation state and compaction
|
||||||
|
|
||||||
|
Providers that use the Responses API can keep reasoning context across a
|
||||||
|
conversation, which helps with multi-step tasks. Supported providers can also
|
||||||
|
compact long conversations automatically.
|
||||||
|
|
||||||
|
nanobot preserves Responses conversation state automatically for OpenAI
|
||||||
|
Responses, OpenAI Codex, Azure OpenAI, and compatible GitHub Copilot models.
|
||||||
|
Native compaction is also automatic when the provider supports it. The
|
||||||
|
threshold is derived from the active model's context window and reserved output
|
||||||
|
headroom; no provider configuration is required.
|
||||||
|
|
||||||
<details>
|
<details>
|
||||||
<summary><b>Azure OpenAI</b></summary>
|
<summary><b>Azure OpenAI</b></summary>
|
||||||
|
|
||||||
|
|||||||
@ -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
|
### Custom OpenAI-Compatible Endpoint
|
||||||
|
|
||||||
@ -458,7 +458,7 @@ For GitHub Copilot:
|
|||||||
nanobot provider login github-copilot --set-main
|
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
|
## Provider Resolution
|
||||||
|
|
||||||
|
|||||||
@ -225,9 +225,6 @@ class ContextBuilder:
|
|||||||
if current_role == "user"
|
if current_role == "user"
|
||||||
else []
|
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]] = [
|
messages: list[dict[str, Any]] = [
|
||||||
{
|
{
|
||||||
"role": "system",
|
"role": "system",
|
||||||
@ -243,21 +240,47 @@ class ContextBuilder:
|
|||||||
},
|
},
|
||||||
*history,
|
*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:
|
if messages[-1].get("role") == current_role:
|
||||||
last = dict(messages[-1])
|
last = dict(messages[-1])
|
||||||
last["content"] = self._merge_message_content(last.get("content"), merged)
|
last["content"] = self._merge_message_content(
|
||||||
if current_role == "user" and runtime_context_meta is not None:
|
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 = 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
|
last["_meta"] = internal_meta
|
||||||
messages[-1] = last
|
messages[-1] = last
|
||||||
return messages
|
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)
|
messages.append(current)
|
||||||
return messages
|
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(
|
def build_user_content(
|
||||||
self,
|
self,
|
||||||
text: str,
|
text: str,
|
||||||
|
|||||||
@ -49,7 +49,7 @@ from nanobot.bus.queue import MessageBus
|
|||||||
from nanobot.bus.runtime_events import RuntimeEventBus
|
from nanobot.bus.runtime_events import RuntimeEventBus
|
||||||
from nanobot.command import CommandContext, CommandRouter, register_builtin_commands
|
from nanobot.command import CommandContext, CommandRouter, register_builtin_commands
|
||||||
from nanobot.config.schema import AgentDefaults, ModelPresetConfig
|
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.providers.factory import ProviderSnapshot
|
||||||
from nanobot.runtime_context import (
|
from nanobot.runtime_context import (
|
||||||
RUNTIME_CONTEXT_HISTORY_META,
|
RUNTIME_CONTEXT_HISTORY_META,
|
||||||
@ -106,6 +106,7 @@ if TYPE_CHECKING:
|
|||||||
from nanobot.triggers.local_store import LocalTriggerStore
|
from nanobot.triggers.local_store import LocalTriggerStore
|
||||||
|
|
||||||
_T = TypeVar("_T")
|
_T = TypeVar("_T")
|
||||||
|
_SUBAGENT_PROVIDER_TASK_META = "subagent_provider_task_id"
|
||||||
|
|
||||||
|
|
||||||
class TurnKind(Enum):
|
class TurnKind(Enum):
|
||||||
@ -126,6 +127,7 @@ class TurnContext:
|
|||||||
|
|
||||||
history: list[dict[str, Any]] = field(default_factory=list)
|
history: list[dict[str, Any]] = field(default_factory=list)
|
||||||
initial_messages: 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
|
request_context: RequestContext | None = None
|
||||||
runtime_context_blocks: list[RuntimeContextBlock] = field(default_factory=list)
|
runtime_context_blocks: list[RuntimeContextBlock] = field(default_factory=list)
|
||||||
attributes: dict[str, Any] = field(default_factory=dict)
|
attributes: dict[str, Any] = field(default_factory=dict)
|
||||||
@ -243,6 +245,8 @@ class AgentLoop:
|
|||||||
|
|
||||||
_RUNTIME_CHECKPOINT_KEY = "runtime_checkpoint"
|
_RUNTIME_CHECKPOINT_KEY = "runtime_checkpoint"
|
||||||
_PENDING_USER_TURN_KEY = "pending_user_turn"
|
_PENDING_USER_TURN_KEY = "pending_user_turn"
|
||||||
|
_PROVIDER_STATE_CHECKPOINT_VERSION_KEY = "provider_state_checkpoint_version"
|
||||||
|
_PROVIDER_STATE_CHECKPOINT_VERSION = "v1"
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@ -857,6 +861,7 @@ class AgentLoop:
|
|||||||
turn_scopes: list[AbstractContextManager[Any]] | None = None,
|
turn_scopes: list[AbstractContextManager[Any]] | None = None,
|
||||||
tools: ToolRegistry | None = None,
|
tools: ToolRegistry | None = None,
|
||||||
request_context: RequestContext | None = None,
|
request_context: RequestContext | None = None,
|
||||||
|
provider_state: ProviderConversationState | None = None,
|
||||||
) -> tuple[str | None, list[str], list[dict[str, Any]], str, bool]:
|
) -> tuple[str | None, list[str], list[dict[str, Any]], str, bool]:
|
||||||
"""Run the agent iteration loop.
|
"""Run the agent iteration loop.
|
||||||
|
|
||||||
@ -872,7 +877,18 @@ class AgentLoop:
|
|||||||
async def _checkpoint(payload: dict[str, Any]) -> None:
|
async def _checkpoint(payload: dict[str, Any]) -> None:
|
||||||
if session is None:
|
if session is None:
|
||||||
return
|
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]]:
|
async def _drain_pending(*, limit: int = _MAX_INJECTIONS_PER_TURN) -> list[dict[str, Any]]:
|
||||||
"""Drain follow-up messages from the pending queue.
|
"""Drain follow-up messages from the pending queue.
|
||||||
@ -1070,6 +1086,7 @@ class AgentLoop:
|
|||||||
session_metadata=session_metadata,
|
session_metadata=session_metadata,
|
||||||
message_metadata=metadata,
|
message_metadata=metadata,
|
||||||
),
|
),
|
||||||
|
provider_state=provider_state,
|
||||||
))
|
))
|
||||||
finally:
|
finally:
|
||||||
turn_scope_stack.close()
|
turn_scope_stack.close()
|
||||||
@ -1077,6 +1094,8 @@ class AgentLoop:
|
|||||||
reset_request_context(request_token)
|
reset_request_context(request_token)
|
||||||
reset_file_states(file_state_token)
|
reset_file_states(file_state_token)
|
||||||
self._last_usage = result.usage
|
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":
|
if result.stop_reason == "max_iterations":
|
||||||
logger.warning("Max iterations ({}) reached", self.max_iterations)
|
logger.warning("Max iterations ({}) reached", self.max_iterations)
|
||||||
should_stream = turn_continuation.should_stream_budget_response(
|
should_stream = turn_continuation.should_stream_budget_response(
|
||||||
@ -1660,14 +1679,24 @@ class AgentLoop:
|
|||||||
"extend_to_user": is_subagent,
|
"extend_to_user": is_subagent,
|
||||||
}
|
}
|
||||||
ctx.history = session.get_history(**_hist_kwargs)
|
ctx.history = session.get_history(**_hist_kwargs)
|
||||||
|
stored_state = session.provider_state
|
||||||
|
subagent_followup_persisted = False
|
||||||
if is_subagent:
|
if is_subagent:
|
||||||
# Keep the durable internal delivery as an assistant record, but
|
# Keep the durable internal delivery as an assistant record, but
|
||||||
# present this completion to the model as fresh follow-up input.
|
# present this completion to the model as fresh follow-up input.
|
||||||
# Providers without assistant-prefill support drop trailing
|
# Providers without assistant-prefill support drop trailing
|
||||||
# assistant messages, so using the persisted record as the current
|
# assistant messages, so using the persisted record as the current
|
||||||
# prompt would hide an independently dispatched subagent result.
|
# 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)
|
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)
|
self.sessions.save(session)
|
||||||
ctx.input_persisted_early = True
|
ctx.input_persisted_early = True
|
||||||
ctx.delivery.record_runtime(runtime)
|
ctx.delivery.record_runtime(runtime)
|
||||||
@ -1675,13 +1704,65 @@ class AgentLoop:
|
|||||||
ctx.request_context = self._request_context_for_turn(ctx)
|
ctx.request_context = self._request_context_for_turn(ctx)
|
||||||
if ctx.kind is TurnKind.USER:
|
if ctx.kind is TurnKind.USER:
|
||||||
ctx.runtime_context_blocks = await self._resolve_runtime_context_for_turn(ctx)
|
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:
|
if ctx.kind is TurnKind.USER:
|
||||||
ctx.input_persisted_early = self._persist_user_message_early(
|
ctx.input_persisted_early = self._persist_user_message_early(
|
||||||
ctx.msg,
|
ctx.msg,
|
||||||
session,
|
session,
|
||||||
runtime_context_blocks=ctx.runtime_context_blocks,
|
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:
|
if ctx.on_progress is None:
|
||||||
ctx.on_progress = ctx.delivery.progress_callback()
|
ctx.on_progress = ctx.delivery.progress_callback()
|
||||||
@ -1715,6 +1796,7 @@ class AgentLoop:
|
|||||||
turn_scopes=ctx.turn_scopes,
|
turn_scopes=ctx.turn_scopes,
|
||||||
tools=ctx.tools,
|
tools=ctx.tools,
|
||||||
request_context=ctx.request_context,
|
request_context=ctx.request_context,
|
||||||
|
provider_state=ctx.provider_state,
|
||||||
)
|
)
|
||||||
final_content, _, all_msgs, stop_reason, had_injections = result
|
final_content, _, all_msgs, stop_reason, had_injections = result
|
||||||
ctx.final_content = final_content
|
ctx.final_content = final_content
|
||||||
@ -2052,7 +2134,36 @@ class AgentLoop:
|
|||||||
):
|
):
|
||||||
overlap = size
|
overlap = size
|
||||||
break
|
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_pending_user_turn(session)
|
||||||
self._clear_runtime_checkpoint(session)
|
self._clear_runtime_checkpoint(session)
|
||||||
@ -2073,6 +2184,7 @@ class AgentLoop:
|
|||||||
"timestamp": datetime.now().isoformat(),
|
"timestamp": datetime.now().isoformat(),
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
session.provider_state = None
|
||||||
session.updated_at = datetime.now()
|
session.updated_at = datetime.now()
|
||||||
|
|
||||||
self._clear_pending_user_turn(session)
|
self._clear_pending_user_turn(session)
|
||||||
|
|||||||
@ -931,6 +931,7 @@ class Consolidator:
|
|||||||
session_key=session.key,
|
session_key=session.key,
|
||||||
)
|
)
|
||||||
session.last_consolidated = end_idx
|
session.last_consolidated = end_idx
|
||||||
|
session.provider_state = None
|
||||||
self.sessions.save(session)
|
self.sessions.save(session)
|
||||||
return summary
|
return summary
|
||||||
|
|
||||||
@ -1136,6 +1137,7 @@ class Consolidator:
|
|||||||
if summary:
|
if summary:
|
||||||
last_summary = summary
|
last_summary = summary
|
||||||
session.last_consolidated = end_idx
|
session.last_consolidated = end_idx
|
||||||
|
session.provider_state = None
|
||||||
self.sessions.save(session)
|
self.sessions.save(session)
|
||||||
if not summary:
|
if not summary:
|
||||||
# LLM is degraded — stop hammering it this call;
|
# LLM is degraded — stop hammering it this call;
|
||||||
@ -1205,6 +1207,7 @@ class Consolidator:
|
|||||||
|
|
||||||
# Preserve history and advance only the replay boundary.
|
# Preserve history and advance only the replay boundary.
|
||||||
session.last_consolidated = len(session.messages) - len(visible_suffix)
|
session.last_consolidated = len(session.messages) - len(visible_suffix)
|
||||||
|
session.provider_state = None
|
||||||
self.sessions.save(session)
|
self.sessions.save(session)
|
||||||
|
|
||||||
logger.info(
|
logger.info(
|
||||||
|
|||||||
@ -19,7 +19,17 @@ from nanobot.agent.context_governance import (
|
|||||||
)
|
)
|
||||||
from nanobot.agent.hook import AgentHook, AgentHookContext, AgentRunHookContext
|
from nanobot.agent.hook import AgentHook, AgentHookContext, AgentRunHookContext
|
||||||
from nanobot.agent.tools.registry import ToolRegistry, is_tool_error_result
|
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 (
|
from nanobot.runtime_context import (
|
||||||
RUNTIME_CONTEXT_MESSAGE_META,
|
RUNTIME_CONTEXT_MESSAGE_META,
|
||||||
detach_runtime_context,
|
detach_runtime_context,
|
||||||
@ -104,6 +114,7 @@ class AgentRunSpec:
|
|||||||
goal_active_predicate: Callable[[], bool] | None = None
|
goal_active_predicate: Callable[[], bool] | None = None
|
||||||
goal_continue_message: GoalContinueMessage | None = None
|
goal_continue_message: GoalContinueMessage | None = None
|
||||||
finalize_on_max_iterations: bool = True
|
finalize_on_max_iterations: bool = True
|
||||||
|
provider_state: ProviderConversationState | None = None
|
||||||
|
|
||||||
|
|
||||||
@dataclass(slots=True)
|
@dataclass(slots=True)
|
||||||
@ -120,6 +131,7 @@ class AgentRunResult:
|
|||||||
had_injections: bool = False
|
had_injections: bool = False
|
||||||
# Terminal tail to emit when the preceding final-content prefix was already streamed.
|
# Terminal tail to emit when the preceding final-content prefix was already streamed.
|
||||||
pending_stream_content: str | None = None
|
pending_stream_content: str | None = None
|
||||||
|
provider_state: ProviderConversationState | None = field(default=None, repr=False)
|
||||||
|
|
||||||
|
|
||||||
class AgentRunner:
|
class AgentRunner:
|
||||||
@ -161,6 +173,7 @@ class AgentRunner:
|
|||||||
and messages[-1].get("role") == "user"
|
and messages[-1].get("role") == "user"
|
||||||
and not is_hidden_history_message(injection)
|
and not is_hidden_history_message(injection)
|
||||||
and not is_hidden_history_message(messages[-1])
|
and not is_hidden_history_message(messages[-1])
|
||||||
|
and allows_conversation_message_merge(messages[-1])
|
||||||
):
|
):
|
||||||
merged = dict(messages[-1])
|
merged = dict(messages[-1])
|
||||||
left_meta = merged.get("_meta")
|
left_meta = merged.get("_meta")
|
||||||
@ -231,6 +244,7 @@ class AgentRunner:
|
|||||||
assistant_message: dict[str, Any] | None,
|
assistant_message: dict[str, Any] | None,
|
||||||
injection_cycles: int,
|
injection_cycles: int,
|
||||||
*,
|
*,
|
||||||
|
conversation_state: ProviderConversationStateController | None = None,
|
||||||
phase: str = "after error",
|
phase: str = "after error",
|
||||||
iteration: int | None = None,
|
iteration: int | None = None,
|
||||||
allow_goal_continue: bool = False,
|
allow_goal_continue: bool = False,
|
||||||
@ -258,16 +272,21 @@ class AgentRunner:
|
|||||||
if assistant_message is not None:
|
if assistant_message is not None:
|
||||||
messages.append(assistant_message)
|
messages.append(assistant_message)
|
||||||
if iteration is not None:
|
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(
|
await self._emit_checkpoint(
|
||||||
spec,
|
spec,
|
||||||
{
|
checkpoint,
|
||||||
"phase": "final_response",
|
|
||||||
"iteration": iteration,
|
|
||||||
"model": spec.runtime.model,
|
|
||||||
"assistant_message": assistant_message,
|
|
||||||
"completed_tool_results": [],
|
|
||||||
"pending_tool_calls": [],
|
|
||||||
},
|
|
||||||
)
|
)
|
||||||
self._append_injected_messages(messages, injections)
|
self._append_injected_messages(messages, injections)
|
||||||
if real_injection:
|
if real_injection:
|
||||||
@ -420,6 +439,12 @@ class AgentRunner:
|
|||||||
injection_cycles = 0
|
injection_cycles = 0
|
||||||
compacted_tool_call_ids: set[str] = set()
|
compacted_tool_call_ids: set[str] = set()
|
||||||
pending_stream_content: str | None = None
|
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(
|
governance_config = ContextGovernanceConfig(
|
||||||
provider=spec.runtime.provider,
|
provider=spec.runtime.provider,
|
||||||
model=spec.runtime.model,
|
model=spec.runtime.model,
|
||||||
@ -450,7 +475,20 @@ class AgentRunner:
|
|||||||
session_key=spec.session_key,
|
session_key=spec.session_key,
|
||||||
)
|
)
|
||||||
await hook.before_iteration(context)
|
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.response = response
|
||||||
context.tool_calls = list(response.tool_calls)
|
context.tool_calls = list(response.tool_calls)
|
||||||
|
|
||||||
@ -480,6 +518,10 @@ class AgentRunner:
|
|||||||
reasoning_content=response.reasoning_content,
|
reasoning_content=response.reasoning_content,
|
||||||
thinking_blocks=response.thinking_blocks,
|
thinking_blocks=response.thinking_blocks,
|
||||||
)
|
)
|
||||||
|
assistant_message = conversation_state.project_response_message(
|
||||||
|
assistant_message,
|
||||||
|
response,
|
||||||
|
)
|
||||||
messages.append(assistant_message)
|
messages.append(assistant_message)
|
||||||
await self._emit_checkpoint(
|
await self._emit_checkpoint(
|
||||||
spec,
|
spec,
|
||||||
@ -544,6 +586,15 @@ class AgentRunner:
|
|||||||
length_recovery_parts.clear()
|
length_recovery_parts.clear()
|
||||||
continue
|
continue
|
||||||
break
|
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(
|
await self._emit_checkpoint(
|
||||||
spec,
|
spec,
|
||||||
{
|
{
|
||||||
@ -553,6 +604,10 @@ class AgentRunner:
|
|||||||
"assistant_message": assistant_message,
|
"assistant_message": assistant_message,
|
||||||
"completed_tool_results": completed_tool_results,
|
"completed_tool_results": completed_tool_results,
|
||||||
"pending_tool_calls": [],
|
"pending_tool_calls": [],
|
||||||
|
"provider_state": conversation_state.checkpoint(
|
||||||
|
messages,
|
||||||
|
model_messages=checkpoint_model_messages,
|
||||||
|
),
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
empty_content_retries = 0
|
empty_content_retries = 0
|
||||||
@ -575,7 +630,11 @@ class AgentRunner:
|
|||||||
)
|
)
|
||||||
|
|
||||||
clean = hook.finalize_content(context, response.content)
|
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
|
empty_content_retries += 1
|
||||||
if empty_content_retries < _MAX_EMPTY_RETRIES:
|
if empty_content_retries < _MAX_EMPTY_RETRIES:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
@ -598,7 +657,12 @@ class AgentRunner:
|
|||||||
if hook.wants_streaming():
|
if hook.wants_streaming():
|
||||||
await hook.on_stream_end(context, resuming=False)
|
await hook.on_stream_end(context, resuming=False)
|
||||||
retry_messages = self._finalization_retry_messages(messages_for_model)
|
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)
|
retry_usage = self._usage_or_estimate(spec, retry_messages, response)
|
||||||
self._accumulate_usage(usage, retry_usage)
|
self._accumulate_usage(usage, retry_usage)
|
||||||
raw_usage = self._merge_usage(raw_usage, retry_usage)
|
raw_usage = self._merge_usage(raw_usage, retry_usage)
|
||||||
@ -623,10 +687,13 @@ class AgentRunner:
|
|||||||
if hook.wants_streaming():
|
if hook.wants_streaming():
|
||||||
context.stream_continues_current_message = True
|
context.stream_continues_current_message = True
|
||||||
await hook.on_stream_end(context, resuming=True)
|
await hook.on_stream_end(context, resuming=True)
|
||||||
messages.append(build_assistant_message(
|
messages.append(conversation_state.project_response_message(
|
||||||
clean,
|
build_assistant_message(
|
||||||
reasoning_content=response.reasoning_content,
|
clean,
|
||||||
thinking_blocks=response.thinking_blocks,
|
reasoning_content=response.reasoning_content,
|
||||||
|
thinking_blocks=response.thinking_blocks,
|
||||||
|
),
|
||||||
|
response,
|
||||||
))
|
))
|
||||||
messages.append(build_length_recovery_message(clean or ""))
|
messages.append(build_length_recovery_message(clean or ""))
|
||||||
await hook.after_iteration(context)
|
await hook.after_iteration(context)
|
||||||
@ -656,15 +723,22 @@ class AgentRunner:
|
|||||||
reasoning_content=response.reasoning_content,
|
reasoning_content=response.reasoning_content,
|
||||||
thinking_blocks=response.thinking_blocks,
|
thinking_blocks=response.thinking_blocks,
|
||||||
)
|
)
|
||||||
|
assistant_message = conversation_state.project_response_message(
|
||||||
|
assistant_message,
|
||||||
|
response,
|
||||||
|
)
|
||||||
|
|
||||||
# Check for mid-turn injections BEFORE signaling stream end.
|
# Check for mid-turn injections BEFORE signaling stream end.
|
||||||
# If injections are found we keep the stream alive (resuming=True)
|
# If injections are found we keep the stream alive (resuming=True)
|
||||||
# so streaming channels don't prematurely finalize the card.
|
# so streaming channels don't prematurely finalize the card.
|
||||||
should_continue, injection_cycles = await self._try_drain_injections(
|
should_continue, injection_cycles = await self._try_drain_injections(
|
||||||
spec, messages, assistant_message, injection_cycles,
|
spec, messages, assistant_message, injection_cycles,
|
||||||
|
conversation_state=conversation_state,
|
||||||
phase="after final response",
|
phase="after final response",
|
||||||
iteration=iteration,
|
iteration=iteration,
|
||||||
allow_goal_continue=True,
|
allow_goal_continue=(
|
||||||
|
response.finish_reason not in {"refusal", "content_filter"}
|
||||||
|
),
|
||||||
)
|
)
|
||||||
if should_continue:
|
if should_continue:
|
||||||
had_injections = True
|
had_injections = True
|
||||||
@ -717,11 +791,17 @@ class AgentRunner:
|
|||||||
continue
|
continue
|
||||||
break
|
break
|
||||||
|
|
||||||
messages.append(assistant_message or build_assistant_message(
|
messages.append(
|
||||||
clean,
|
assistant_message
|
||||||
reasoning_content=response.reasoning_content,
|
or conversation_state.project_response_message(
|
||||||
thinking_blocks=response.thinking_blocks,
|
build_assistant_message(
|
||||||
))
|
clean,
|
||||||
|
reasoning_content=response.reasoning_content,
|
||||||
|
thinking_blocks=response.thinking_blocks,
|
||||||
|
),
|
||||||
|
response,
|
||||||
|
)
|
||||||
|
)
|
||||||
await self._emit_checkpoint(
|
await self._emit_checkpoint(
|
||||||
spec,
|
spec,
|
||||||
{
|
{
|
||||||
@ -731,6 +811,7 @@ class AgentRunner:
|
|||||||
"assistant_message": messages[-1],
|
"assistant_message": messages[-1],
|
||||||
"completed_tool_results": [],
|
"completed_tool_results": [],
|
||||||
"pending_tool_calls": [],
|
"pending_tool_calls": [],
|
||||||
|
"provider_state": conversation_state.checkpoint(messages),
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
if length_recovery_parts:
|
if length_recovery_parts:
|
||||||
@ -764,6 +845,7 @@ class AgentRunner:
|
|||||||
hook,
|
hook,
|
||||||
messages,
|
messages,
|
||||||
usage,
|
usage,
|
||||||
|
conversation_state,
|
||||||
)
|
)
|
||||||
if terminal_content is None:
|
if terminal_content is None:
|
||||||
terminal_content = self._max_iterations_fallback(spec)
|
terminal_content = self._max_iterations_fallback(spec)
|
||||||
@ -787,6 +869,7 @@ class AgentRunner:
|
|||||||
tool_events=tool_events,
|
tool_events=tool_events,
|
||||||
had_injections=had_injections,
|
had_injections=had_injections,
|
||||||
pending_stream_content=pending_stream_content,
|
pending_stream_content=pending_stream_content,
|
||||||
|
provider_state=conversation_state.finish(messages),
|
||||||
)
|
)
|
||||||
|
|
||||||
def _build_request_kwargs(
|
def _build_request_kwargs(
|
||||||
@ -817,6 +900,8 @@ class AgentRunner:
|
|||||||
context: AgentHookContext,
|
context: AgentHookContext,
|
||||||
*,
|
*,
|
||||||
malformed_retry: bool = False,
|
malformed_retry: bool = False,
|
||||||
|
conversation_state: ProviderConversationStateController,
|
||||||
|
provider_context: ProviderCallContext | None = None,
|
||||||
) -> LLMResponse:
|
) -> LLMResponse:
|
||||||
timeout_s: float | None = spec.llm_timeout_s
|
timeout_s: float | None = spec.llm_timeout_s
|
||||||
if timeout_s is None:
|
if timeout_s is None:
|
||||||
@ -886,6 +971,7 @@ class AgentRunner:
|
|||||||
|
|
||||||
coro = spec.runtime.provider.chat_stream_with_retry(
|
coro = spec.runtime.provider.chat_stream_with_retry(
|
||||||
**kwargs,
|
**kwargs,
|
||||||
|
provider_context=provider_context,
|
||||||
on_content_delta=_stream,
|
on_content_delta=_stream,
|
||||||
on_thinking_delta=_thinking,
|
on_thinking_delta=_thinking,
|
||||||
on_tool_call_delta=_provider_tool_event,
|
on_tool_call_delta=_provider_tool_event,
|
||||||
@ -920,11 +1006,15 @@ class AgentRunner:
|
|||||||
|
|
||||||
coro = spec.runtime.provider.chat_stream_with_retry(
|
coro = spec.runtime.provider.chat_stream_with_retry(
|
||||||
**kwargs,
|
**kwargs,
|
||||||
|
provider_context=provider_context,
|
||||||
on_content_delta=_stream_progress,
|
on_content_delta=_stream_progress,
|
||||||
on_tool_call_delta=_provider_tool_event,
|
on_tool_call_delta=_provider_tool_event,
|
||||||
)
|
)
|
||||||
else:
|
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
|
# Streaming requests also have provider-level idle timeouts
|
||||||
# (NANOBOT_STREAM_IDLE_TIMEOUT_S), but a stream that keeps producing
|
# (NANOBOT_STREAM_IDLE_TIMEOUT_S), but a stream that keeps producing
|
||||||
@ -986,6 +1076,10 @@ class AgentRunner:
|
|||||||
return await self._request_model(
|
return await self._request_model(
|
||||||
spec, retry_messages, hook, context,
|
spec, retry_messages, hook, context,
|
||||||
malformed_retry=True,
|
malformed_retry=True,
|
||||||
|
conversation_state=conversation_state,
|
||||||
|
provider_context=conversation_state.independent_request_context(
|
||||||
|
context_window_tokens=spec.runtime.context_window_tokens,
|
||||||
|
),
|
||||||
)
|
)
|
||||||
if (
|
if (
|
||||||
all_dropped
|
all_dropped
|
||||||
@ -998,7 +1092,13 @@ class AgentRunner:
|
|||||||
fallback_messages = self._malformed_tool_call_retry_messages(
|
fallback_messages = self._malformed_tool_call_retry_messages(
|
||||||
messages, response.content,
|
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
|
return response
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
@ -1031,6 +1131,10 @@ class AgentRunner:
|
|||||||
original_finish_reason,
|
original_finish_reason,
|
||||||
)
|
)
|
||||||
response.tool_calls = valid
|
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:
|
if not valid:
|
||||||
response.finish_reason = "stop"
|
response.finish_reason = "stop"
|
||||||
return (dropped, not valid, original_finish_reason)
|
return (dropped, not valid, original_finish_reason)
|
||||||
@ -1060,9 +1164,27 @@ class AgentRunner:
|
|||||||
self,
|
self,
|
||||||
spec: AgentRunSpec,
|
spec: AgentRunSpec,
|
||||||
messages: list[dict[str, Any]],
|
messages: list[dict[str, Any]],
|
||||||
|
*,
|
||||||
|
transcript: list[dict[str, Any]],
|
||||||
|
conversation_state: ProviderConversationStateController,
|
||||||
) -> LLMResponse:
|
) -> LLMResponse:
|
||||||
retry_messages = self._finalization_retry_messages(messages)
|
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
|
@staticmethod
|
||||||
def _finalization_retry_messages(messages: list[dict[str, Any]]) -> list[dict[str, Any]]:
|
def _finalization_retry_messages(messages: list[dict[str, Any]]) -> list[dict[str, Any]]:
|
||||||
@ -1076,10 +1198,17 @@ class AgentRunner:
|
|||||||
hook: AgentHook,
|
hook: AgentHook,
|
||||||
messages: list[dict[str, Any]],
|
messages: list[dict[str, Any]],
|
||||||
usage: dict[str, int],
|
usage: dict[str, int],
|
||||||
|
conversation_state: ProviderConversationStateController,
|
||||||
) -> str | None:
|
) -> str | None:
|
||||||
retry_messages = self._budget_exhausted_finalization_messages(messages)
|
retry_messages = self._budget_exhausted_finalization_messages(messages)
|
||||||
try:
|
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:
|
except Exception:
|
||||||
logger.exception(
|
logger.exception(
|
||||||
"Budget-exhausted finalization failed for {}; using fallback",
|
"Budget-exhausted finalization failed for {}; using fallback",
|
||||||
@ -1115,9 +1244,18 @@ class AgentRunner:
|
|||||||
self,
|
self,
|
||||||
spec: AgentRunSpec,
|
spec: AgentRunSpec,
|
||||||
messages: list[dict[str, Any]],
|
messages: list[dict[str, Any]],
|
||||||
|
*,
|
||||||
|
provider_context: ProviderCallContext | None = None,
|
||||||
) -> LLMResponse:
|
) -> LLMResponse:
|
||||||
kwargs = self._build_request_kwargs(spec, messages, tools=None)
|
kwargs = self._build_request_kwargs(
|
||||||
return await spec.runtime.provider.chat_with_retry(**kwargs)
|
spec,
|
||||||
|
messages,
|
||||||
|
tools=None,
|
||||||
|
)
|
||||||
|
return await spec.runtime.provider.chat_with_retry(
|
||||||
|
**kwargs,
|
||||||
|
provider_context=provider_context,
|
||||||
|
)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _budget_exhausted_finalization_messages(
|
def _budget_exhausted_finalization_messages(
|
||||||
|
|||||||
@ -23,14 +23,26 @@ import uuid
|
|||||||
from collections.abc import Awaitable, Callable
|
from collections.abc import Awaitable, Callable
|
||||||
from typing import Any, cast
|
from typing import Any, cast
|
||||||
|
|
||||||
|
from loguru import logger
|
||||||
from openai import AsyncOpenAI
|
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 (
|
from nanobot.providers.openai_responses import (
|
||||||
|
ResponsesStreamCapture,
|
||||||
|
build_responses_state,
|
||||||
consume_sdk_stream,
|
consume_sdk_stream,
|
||||||
convert_messages,
|
|
||||||
convert_tools,
|
convert_tools,
|
||||||
|
is_compaction_compatibility_error,
|
||||||
|
is_replayable_finish_reason,
|
||||||
parse_response_output,
|
parse_response_output,
|
||||||
|
prepare_responses_input,
|
||||||
|
resolve_compact_threshold,
|
||||||
|
responses_state_matches,
|
||||||
)
|
)
|
||||||
|
|
||||||
_AZURE_OPENAI_SCOPE = "https://cognitiveservices.azure.com/.default"
|
_AZURE_OPENAI_SCOPE = "https://cognitiveservices.azure.com/.default"
|
||||||
@ -97,6 +109,7 @@ class AzureOpenAIProvider(LLMProvider):
|
|||||||
):
|
):
|
||||||
super().__init__(api_key, api_base)
|
super().__init__(api_key, api_base)
|
||||||
self.default_model = default_model
|
self.default_model = default_model
|
||||||
|
self._native_compaction_available = True
|
||||||
|
|
||||||
if not api_base:
|
if not api_base:
|
||||||
raise ValueError("Azure OpenAI api_base is required")
|
raise ValueError("Azure OpenAI api_base is required")
|
||||||
@ -142,6 +155,25 @@ class AzureOpenAIProvider(LLMProvider):
|
|||||||
name = deployment_name.lower()
|
name = deployment_name.lower()
|
||||||
return not any(token in name for token in ("gpt-5", "o1", "o3", "o4"))
|
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(
|
def _build_body(
|
||||||
self,
|
self,
|
||||||
messages: list[dict[str, Any]],
|
messages: list[dict[str, Any]],
|
||||||
@ -151,10 +183,26 @@ class AzureOpenAIProvider(LLMProvider):
|
|||||||
temperature: float,
|
temperature: float,
|
||||||
reasoning_effort: str | None,
|
reasoning_effort: str | None,
|
||||||
tool_choice: str | dict[str, Any] | None,
|
tool_choice: str | dict[str, Any] | None,
|
||||||
|
provider_context: ProviderCallContext | None = None,
|
||||||
) -> dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""Build the Responses API request body from Chat-Completions-style args."""
|
"""Build the Responses API request body from Chat-Completions-style args."""
|
||||||
deployment = model or self.default_model
|
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] = {
|
body: dict[str, Any] = {
|
||||||
"model": deployment,
|
"model": deployment,
|
||||||
@ -164,13 +212,29 @@ class AzureOpenAIProvider(LLMProvider):
|
|||||||
"store": False,
|
"store": False,
|
||||||
"stream": 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):
|
if self._supports_temperature(deployment, reasoning_effort):
|
||||||
body["temperature"] = temperature
|
body["temperature"] = temperature
|
||||||
|
|
||||||
|
if not self._supports_temperature(deployment, reasoning_effort):
|
||||||
|
body["include"] = ["reasoning.encrypted_content"]
|
||||||
if reasoning_effort and reasoning_effort.lower() != "none":
|
if reasoning_effort and reasoning_effort.lower() != "none":
|
||||||
body["reasoning"] = {"effort": reasoning_effort}
|
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:
|
if tools:
|
||||||
body["tools"] = convert_tools(tools)
|
body["tools"] = convert_tools(tools)
|
||||||
@ -178,21 +242,97 @@ class AzureOpenAIProvider(LLMProvider):
|
|||||||
|
|
||||||
return body
|
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
|
@staticmethod
|
||||||
def _handle_error(e: Exception) -> LLMResponse:
|
def _handle_error(e: Exception) -> LLMResponse:
|
||||||
response = getattr(e, "response", None)
|
response = getattr(e, "response", None)
|
||||||
body = getattr(e, "body", None) or getattr(response, "text", None)
|
body = getattr(e, "body", None) or getattr(response, "text", None)
|
||||||
body_text = str(body).strip() if body is not None else ""
|
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}"
|
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:
|
if retry_after is None:
|
||||||
retry_after = LLMProvider._extract_retry_after(msg)
|
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
|
# 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(
|
async def chat(
|
||||||
self,
|
self,
|
||||||
messages: list[dict[str, Any]],
|
messages: list[dict[str, Any]],
|
||||||
@ -202,14 +342,21 @@ class AzureOpenAIProvider(LLMProvider):
|
|||||||
temperature: float = 0.7,
|
temperature: float = 0.7,
|
||||||
reasoning_effort: str | None = None,
|
reasoning_effort: str | None = None,
|
||||||
tool_choice: str | dict[str, Any] | None = None,
|
tool_choice: str | dict[str, Any] | None = None,
|
||||||
|
provider_context: ProviderCallContext | None = None,
|
||||||
) -> LLMResponse:
|
) -> LLMResponse:
|
||||||
body = self._build_body(
|
body = self._build_body(
|
||||||
messages, tools, model, max_tokens, temperature,
|
messages, tools, model, max_tokens, temperature,
|
||||||
reasoning_effort, tool_choice,
|
reasoning_effort, tool_choice,
|
||||||
|
provider_context,
|
||||||
)
|
)
|
||||||
try:
|
try:
|
||||||
response = cast(Any, await self._client.responses.create(**body))
|
response = await self._create_response_with_compaction_fallback(body)
|
||||||
return parse_response_output(response)
|
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:
|
except Exception as e:
|
||||||
return self._handle_error(e)
|
return self._handle_error(e)
|
||||||
|
|
||||||
@ -225,26 +372,43 @@ class AzureOpenAIProvider(LLMProvider):
|
|||||||
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
|
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
|
||||||
on_thinking_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,
|
on_tool_call_delta: Callable[[dict[str, Any]], Awaitable[None]] | None = None,
|
||||||
|
provider_context: ProviderCallContext | None = None,
|
||||||
) -> LLMResponse:
|
) -> LLMResponse:
|
||||||
_ = on_thinking_delta
|
_ = on_thinking_delta
|
||||||
body = self._build_body(
|
body = self._build_body(
|
||||||
messages, tools, model, max_tokens, temperature,
|
messages, tools, model, max_tokens, temperature,
|
||||||
reasoning_effort, tool_choice,
|
reasoning_effort, tool_choice,
|
||||||
|
provider_context,
|
||||||
)
|
)
|
||||||
body["stream"] = True
|
body["stream"] = True
|
||||||
|
|
||||||
try:
|
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 = (
|
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,
|
content=content or None,
|
||||||
tool_calls=tool_calls,
|
tool_calls=tool_calls,
|
||||||
finish_reason=finish_reason,
|
finish_reason=finish_reason,
|
||||||
usage=usage,
|
usage=usage,
|
||||||
reasoning_content=reasoning_content,
|
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:
|
except Exception as e:
|
||||||
return self._handle_error(e)
|
return self._handle_error(e)
|
||||||
|
|
||||||
|
|||||||
@ -1,5 +1,7 @@
|
|||||||
"""Base LLM provider interface."""
|
"""Base LLM provider interface."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import json
|
import json
|
||||||
import os
|
import os
|
||||||
@ -7,6 +9,7 @@ import re
|
|||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from collections.abc import Awaitable, Callable
|
from collections.abc import Awaitable, Callable
|
||||||
from contextlib import suppress
|
from contextlib import suppress
|
||||||
|
from copy import deepcopy
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from datetime import datetime, timezone
|
from datetime import datetime, timezone
|
||||||
from email.utils import parsedate_to_datetime
|
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)
|
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
|
@dataclass
|
||||||
class LLMResponse:
|
class LLMResponse:
|
||||||
"""Response from an LLM provider."""
|
"""Response from an LLM provider."""
|
||||||
@ -160,6 +261,10 @@ class LLMResponse:
|
|||||||
retry_after: float | None = None # Provider supplied retry wait in seconds.
|
retry_after: float | None = None # Provider supplied retry wait in seconds.
|
||||||
reasoning_content: str | None = None # Kimi, DeepSeek-R1, MiMo etc.
|
reasoning_content: str | None = None # Kimi, DeepSeek-R1, MiMo etc.
|
||||||
thinking_blocks: list[dict[str, Any]] | None = None # Anthropic extended thinking
|
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".
|
# Structured error metadata used by retry policy when finish_reason == "error".
|
||||||
error_status_code: int | None = None
|
error_status_code: int | None = None
|
||||||
error_kind: str | None = None # e.g. "timeout", "connection"
|
error_kind: str | None = None # e.g. "timeout", "connection"
|
||||||
@ -274,6 +379,18 @@ class LLMProvider(ABC):
|
|||||||
self.api_base = api_base
|
self.api_base = api_base
|
||||||
self.generation: GenerationSettings = GenerationSettings()
|
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
|
@staticmethod
|
||||||
def _sanitize_empty_content(messages: list[dict[str, Any]]) -> list[dict[str, Any]]:
|
def _sanitize_empty_content(messages: list[dict[str, Any]]) -> list[dict[str, Any]]:
|
||||||
"""Sanitize message content: fix empty blocks, strip internal _meta fields.
|
"""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)
|
return any(marker in err for marker in cls._TRANSIENT_ERROR_MARKERS)
|
||||||
|
|
||||||
@classmethod
|
@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."""
|
"""Prefer structured error metadata, fallback to text markers for legacy providers."""
|
||||||
if response.error_should_retry is not None:
|
if response.error_should_retry is not None:
|
||||||
return bool(response.error_should_retry)
|
return bool(response.error_should_retry)
|
||||||
@ -607,6 +724,21 @@ class LLMProvider(ABC):
|
|||||||
result.append(msg)
|
result.append(msg)
|
||||||
return result if found else None
|
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
|
@staticmethod
|
||||||
def _strip_image_content_inplace(messages: list[dict[str, Any]]) -> bool:
|
def _strip_image_content_inplace(messages: list[dict[str, Any]]) -> bool:
|
||||||
"""Replace image_url blocks with text placeholder *in-place*.
|
"""Replace image_url blocks with text placeholder *in-place*.
|
||||||
@ -633,6 +765,12 @@ class LLMProvider(ABC):
|
|||||||
async def _safe_chat(self, **kwargs: Any) -> LLMResponse:
|
async def _safe_chat(self, **kwargs: Any) -> LLMResponse:
|
||||||
"""Call chat() and convert unexpected exceptions to error responses."""
|
"""Call chat() and convert unexpected exceptions to error responses."""
|
||||||
try:
|
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)
|
return await self.chat(**kwargs)
|
||||||
except asyncio.CancelledError:
|
except asyncio.CancelledError:
|
||||||
raise
|
raise
|
||||||
@ -666,17 +804,47 @@ class LLMProvider(ABC):
|
|||||||
"""
|
"""
|
||||||
_ = on_thinking_delta, on_tool_call_delta
|
_ = on_thinking_delta, on_tool_call_delta
|
||||||
response = await self.chat(
|
response = await self.chat(
|
||||||
messages=messages, tools=tools, model=model,
|
messages=messages,
|
||||||
max_tokens=max_tokens, temperature=temperature,
|
tools=tools,
|
||||||
reasoning_effort=reasoning_effort, tool_choice=tool_choice,
|
model=model,
|
||||||
|
max_tokens=max_tokens,
|
||||||
|
temperature=temperature,
|
||||||
|
reasoning_effort=reasoning_effort,
|
||||||
|
tool_choice=tool_choice,
|
||||||
)
|
)
|
||||||
if on_content_delta and response.content:
|
if on_content_delta and response.content:
|
||||||
await on_content_delta(response.content)
|
await on_content_delta(response.content)
|
||||||
return response
|
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:
|
async def _safe_chat_stream(self, **kwargs: Any) -> LLMResponse:
|
||||||
"""Call chat_stream() and convert unexpected exceptions to error responses."""
|
"""Call chat_stream() and convert unexpected exceptions to error responses."""
|
||||||
try:
|
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)
|
return await self.chat_stream(**kwargs)
|
||||||
except asyncio.CancelledError:
|
except asyncio.CancelledError:
|
||||||
raise
|
raise
|
||||||
@ -698,6 +866,7 @@ class LLMProvider(ABC):
|
|||||||
on_stream_recover: Callable[[], Awaitable[None]] | None = None,
|
on_stream_recover: Callable[[], Awaitable[None]] | None = None,
|
||||||
retry_mode: str = "standard",
|
retry_mode: str = "standard",
|
||||||
on_retry_wait: Callable[[str], Awaitable[None]] | None = None,
|
on_retry_wait: Callable[[str], Awaitable[None]] | None = None,
|
||||||
|
provider_context: ProviderCallContext | None = None,
|
||||||
) -> LLMResponse:
|
) -> LLMResponse:
|
||||||
"""Call chat_stream() with retry on transient provider failures."""
|
"""Call chat_stream() with retry on transient provider failures."""
|
||||||
if max_tokens is self._SENTINEL or max_tokens is None:
|
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_thinking_delta=on_thinking_delta,
|
||||||
on_tool_call_delta=on_tool_call_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):
|
if on_stream_recover and getattr(self, "supports_stream_recover_callback", False):
|
||||||
kw["on_stream_recover"] = _recover_stream
|
kw["on_stream_recover"] = _recover_stream
|
||||||
return await self._run_with_retry(
|
return await self._run_with_retry(
|
||||||
@ -753,6 +924,7 @@ class LLMProvider(ABC):
|
|||||||
tool_choice: str | dict[str, Any] | None = None,
|
tool_choice: str | dict[str, Any] | None = None,
|
||||||
retry_mode: str = "standard",
|
retry_mode: str = "standard",
|
||||||
on_retry_wait: Callable[[str], Awaitable[None]] | None = None,
|
on_retry_wait: Callable[[str], Awaitable[None]] | None = None,
|
||||||
|
provider_context: ProviderCallContext | None = None,
|
||||||
) -> LLMResponse:
|
) -> LLMResponse:
|
||||||
"""Call chat() with retry on transient provider failures.
|
"""Call chat() with retry on transient provider failures.
|
||||||
|
|
||||||
@ -775,6 +947,8 @@ class LLMProvider(ABC):
|
|||||||
max_tokens=max_tokens, temperature=temperature,
|
max_tokens=max_tokens, temperature=temperature,
|
||||||
reasoning_effort=reasoning_effort, tool_choice=tool_choice,
|
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(
|
return await self._run_with_retry(
|
||||||
self._safe_chat,
|
self._safe_chat,
|
||||||
kw,
|
kw,
|
||||||
@ -932,14 +1106,33 @@ class LLMProvider(ABC):
|
|||||||
last_error_key = error_key
|
last_error_key = error_key
|
||||||
identical_error_count = 1 if error_key else 0
|
identical_error_count = 1 if error_key else 0
|
||||||
|
|
||||||
if not self._is_transient_response(response):
|
if not self.is_transient_response(response):
|
||||||
stripped = self._strip_image_content(original_messages)
|
stripped = self._strip_image_content(kw["messages"])
|
||||||
if stripped is not None and stripped != 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(
|
logger.warning(
|
||||||
"Non-transient LLM error with image content, retrying without images"
|
"Non-transient LLM error with image content, retrying without images"
|
||||||
)
|
)
|
||||||
retry_kw = dict(kw)
|
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)
|
result = await call(**retry_kw)
|
||||||
# Permanently strip images from the original messages so
|
# Permanently strip images from the original messages so
|
||||||
# subsequent iterations do not repeat the error-retry cycle.
|
# subsequent iterations do not repeat the error-retry cycle.
|
||||||
|
|||||||
262
nanobot/providers/conversation_state.py
Normal file
262
nanobot/providers/conversation_state.py
Normal file
@ -0,0 +1,262 @@
|
|||||||
|
"""Provider-owned conversation-state lifecycle coordination."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from copy import deepcopy
|
||||||
|
from typing import Any, cast
|
||||||
|
|
||||||
|
from nanobot.providers.base import (
|
||||||
|
LLMProvider,
|
||||||
|
LLMResponse,
|
||||||
|
ProviderCallContext,
|
||||||
|
ProviderConversationState,
|
||||||
|
)
|
||||||
|
|
||||||
|
_PROVIDER_STATE_OUTPUT_META = "provider_state_output"
|
||||||
|
_PROVIDER_STATE_BOUNDARY_META = "provider_state_boundary"
|
||||||
|
|
||||||
|
|
||||||
|
def allows_conversation_message_merge(message: dict[str, Any]) -> bool:
|
||||||
|
"""Return whether new same-role input may merge into *message*."""
|
||||||
|
internal_meta = cast(object, message.get("_meta"))
|
||||||
|
return not (
|
||||||
|
isinstance(internal_meta, dict)
|
||||||
|
and cast(dict[str, Any], internal_meta).get(
|
||||||
|
_PROVIDER_STATE_BOUNDARY_META
|
||||||
|
) is True
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class ProviderConversationStateController:
|
||||||
|
"""Keep provider conversation-state semantics outside the agent runner.
|
||||||
|
|
||||||
|
The runner owns the tool loop and reports lifecycle events here. This
|
||||||
|
controller owns capability checks, transcript deltas, response projections,
|
||||||
|
retry transitions, and durable snapshots for provider-private state.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
provider: LLMProvider,
|
||||||
|
model: str | None,
|
||||||
|
messages: list[dict[str, Any]],
|
||||||
|
state: ProviderConversationState | None = None,
|
||||||
|
) -> None:
|
||||||
|
self._provider = provider
|
||||||
|
self._model = model
|
||||||
|
self._state = (
|
||||||
|
state
|
||||||
|
if state is not None
|
||||||
|
and provider.can_resume_conversation_state(state, model)
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
self._boundary = len(messages)
|
||||||
|
self._request_messages: list[dict[str, Any]] = []
|
||||||
|
|
||||||
|
def independent_request_context(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
context_window_tokens: int | None,
|
||||||
|
) -> ProviderCallContext | None:
|
||||||
|
"""Return typed provider context for a request that does not resume state."""
|
||||||
|
if context_window_tokens is None:
|
||||||
|
return None
|
||||||
|
return ProviderCallContext(context_window_tokens=context_window_tokens)
|
||||||
|
|
||||||
|
def prepare_request(
|
||||||
|
self,
|
||||||
|
messages: list[dict[str, Any]],
|
||||||
|
*,
|
||||||
|
context_window_tokens: int | None,
|
||||||
|
model_messages: list[dict[str, Any]] | None = None,
|
||||||
|
supplemental_messages: list[dict[str, Any]] | None = None,
|
||||||
|
) -> ProviderCallContext | None:
|
||||||
|
"""Build typed context for the next request and remember its durable delta."""
|
||||||
|
independent_context = self.independent_request_context(
|
||||||
|
context_window_tokens=context_window_tokens,
|
||||||
|
)
|
||||||
|
if self._state is None:
|
||||||
|
self._request_messages = []
|
||||||
|
return independent_context
|
||||||
|
if not self._provider.can_resume_conversation_state(
|
||||||
|
self._state,
|
||||||
|
self._model,
|
||||||
|
):
|
||||||
|
self._state = None
|
||||||
|
self._request_messages = []
|
||||||
|
return independent_context
|
||||||
|
|
||||||
|
durable_messages = self._messages_after_boundary(messages)
|
||||||
|
governed_messages = (
|
||||||
|
self._model_messages_after_boundary(model_messages)
|
||||||
|
if model_messages is not None and durable_messages
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
request_messages = (
|
||||||
|
governed_messages
|
||||||
|
if governed_messages is not None
|
||||||
|
else durable_messages
|
||||||
|
)
|
||||||
|
supplemental = deepcopy(supplemental_messages or [])
|
||||||
|
self._request_messages = deepcopy(request_messages)
|
||||||
|
request_state = self._state.with_pending_messages([
|
||||||
|
*self._state.pending_messages,
|
||||||
|
*request_messages,
|
||||||
|
*supplemental,
|
||||||
|
])
|
||||||
|
return ProviderCallContext(
|
||||||
|
conversation_state=request_state,
|
||||||
|
context_window_tokens=(
|
||||||
|
independent_context.context_window_tokens
|
||||||
|
if independent_context is not None
|
||||||
|
else None
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
def observe_response(
|
||||||
|
self,
|
||||||
|
response: LLMResponse,
|
||||||
|
messages: list[dict[str, Any]],
|
||||||
|
*,
|
||||||
|
adopt_candidate_state: bool = True,
|
||||||
|
) -> None:
|
||||||
|
"""Advance, preserve, or discard state after one provider response."""
|
||||||
|
candidate = response.provider_state if adopt_candidate_state else None
|
||||||
|
candidate_is_replayable = response.finish_reason in {
|
||||||
|
"stop",
|
||||||
|
"tool_calls",
|
||||||
|
"function_call",
|
||||||
|
}
|
||||||
|
if (
|
||||||
|
candidate is not None
|
||||||
|
and candidate_is_replayable
|
||||||
|
and self._provider.can_resume_conversation_state(
|
||||||
|
candidate,
|
||||||
|
self._model,
|
||||||
|
)
|
||||||
|
):
|
||||||
|
self._state = candidate
|
||||||
|
self._boundary = len(messages)
|
||||||
|
self._seal_boundary(messages)
|
||||||
|
elif response.finish_reason == "error" and (
|
||||||
|
response.preserve_provider_state_on_error is True
|
||||||
|
or (
|
||||||
|
response.preserve_provider_state_on_error is None
|
||||||
|
and LLMProvider.is_transient_response(response)
|
||||||
|
)
|
||||||
|
):
|
||||||
|
if self._state is not None and self._request_messages:
|
||||||
|
self._state = self._state.with_pending_messages([
|
||||||
|
*self._state.pending_messages,
|
||||||
|
*self._request_messages,
|
||||||
|
])
|
||||||
|
self._boundary = len(messages)
|
||||||
|
else:
|
||||||
|
self._state = None
|
||||||
|
self._boundary = len(messages)
|
||||||
|
self._request_messages = []
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def project_response_message(
|
||||||
|
message: dict[str, Any],
|
||||||
|
response: LLMResponse,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
"""Mark a Chat projection already represented by provider output."""
|
||||||
|
if response.provider_state is None:
|
||||||
|
return message
|
||||||
|
internal_meta = dict(message.get("_meta") or {})
|
||||||
|
internal_meta[_PROVIDER_STATE_OUTPUT_META] = True
|
||||||
|
message["_meta"] = internal_meta
|
||||||
|
return message
|
||||||
|
|
||||||
|
def checkpoint(
|
||||||
|
self,
|
||||||
|
messages: list[dict[str, Any]],
|
||||||
|
*,
|
||||||
|
model_messages: list[dict[str, Any]] | None = None,
|
||||||
|
) -> ProviderConversationState | None:
|
||||||
|
"""Return a durable state snapshot without changing live state."""
|
||||||
|
if self._state is None:
|
||||||
|
return None
|
||||||
|
durable_messages = self._messages_after_boundary(messages)
|
||||||
|
governed_messages = (
|
||||||
|
self._model_messages_after_boundary(model_messages)
|
||||||
|
if model_messages is not None and durable_messages
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
pending_messages = (
|
||||||
|
governed_messages
|
||||||
|
if governed_messages is not None
|
||||||
|
else durable_messages
|
||||||
|
)
|
||||||
|
return self._state.with_pending_messages([
|
||||||
|
*self._state.pending_messages,
|
||||||
|
*pending_messages,
|
||||||
|
])
|
||||||
|
|
||||||
|
def finish(
|
||||||
|
self,
|
||||||
|
messages: list[dict[str, Any]],
|
||||||
|
) -> ProviderConversationState | None:
|
||||||
|
"""Return the final durable state after all runner messages are known."""
|
||||||
|
self._state = self.checkpoint(messages)
|
||||||
|
return self._state
|
||||||
|
|
||||||
|
def _messages_after_boundary(
|
||||||
|
self,
|
||||||
|
messages: list[dict[str, Any]],
|
||||||
|
) -> list[dict[str, Any]]:
|
||||||
|
pending: list[dict[str, Any]] = []
|
||||||
|
for message in messages[self._boundary:]:
|
||||||
|
internal_meta = cast(object, message.get("_meta"))
|
||||||
|
if (
|
||||||
|
isinstance(internal_meta, dict)
|
||||||
|
and cast(dict[str, Any], internal_meta).get(
|
||||||
|
_PROVIDER_STATE_OUTPUT_META
|
||||||
|
) is True
|
||||||
|
):
|
||||||
|
continue
|
||||||
|
pending.append(deepcopy(message))
|
||||||
|
return pending
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _model_messages_after_boundary(
|
||||||
|
messages: list[dict[str, Any]],
|
||||||
|
) -> list[dict[str, Any]] | None:
|
||||||
|
"""Return the governed delta after the latest provider-owned boundary."""
|
||||||
|
boundary = None
|
||||||
|
for idx in range(len(messages) - 1, -1, -1):
|
||||||
|
internal_meta = cast(object, messages[idx].get("_meta"))
|
||||||
|
if (
|
||||||
|
isinstance(internal_meta, dict)
|
||||||
|
and cast(dict[str, Any], internal_meta).get(
|
||||||
|
_PROVIDER_STATE_BOUNDARY_META
|
||||||
|
) is True
|
||||||
|
):
|
||||||
|
boundary = idx
|
||||||
|
break
|
||||||
|
if boundary is None:
|
||||||
|
return None
|
||||||
|
|
||||||
|
pending: list[dict[str, Any]] = []
|
||||||
|
for message in messages[boundary + 1:]:
|
||||||
|
internal_meta = cast(object, message.get("_meta"))
|
||||||
|
if (
|
||||||
|
isinstance(internal_meta, dict)
|
||||||
|
and cast(dict[str, Any], internal_meta).get(
|
||||||
|
_PROVIDER_STATE_OUTPUT_META
|
||||||
|
) is True
|
||||||
|
):
|
||||||
|
continue
|
||||||
|
pending.append(deepcopy(message))
|
||||||
|
return pending
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _seal_boundary(messages: list[dict[str, Any]]) -> None:
|
||||||
|
"""Prevent later same-role injection merging across a state boundary."""
|
||||||
|
if not messages:
|
||||||
|
return
|
||||||
|
internal_meta = dict(messages[-1].get("_meta") or {})
|
||||||
|
internal_meta[_PROVIDER_STATE_BOUNDARY_META] = True
|
||||||
|
messages[-1]["_meta"] = internal_meta
|
||||||
@ -261,6 +261,7 @@ def make_provider(
|
|||||||
primary=provider,
|
primary=provider,
|
||||||
fallback_presets=fallback_presets,
|
fallback_presets=fallback_presets,
|
||||||
provider_factory=lambda fb: _make_provider_core(config, preset=fb),
|
provider_factory=lambda fb: _make_provider_core(config, preset=fb),
|
||||||
|
primary_context_window_tokens=resolved.context_window_tokens,
|
||||||
)
|
)
|
||||||
|
|
||||||
return provider
|
return provider
|
||||||
|
|||||||
@ -6,11 +6,18 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import time
|
import time
|
||||||
from collections.abc import Awaitable, Callable
|
from collections.abc import Awaitable, Callable
|
||||||
|
from dataclasses import replace
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
from loguru import logger
|
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.
|
# Circuit breaker tuned to match OpenAICompatProvider's Responses API breaker.
|
||||||
_PRIMARY_FAILURE_THRESHOLD = 3
|
_PRIMARY_FAILURE_THRESHOLD = 3
|
||||||
@ -113,11 +120,13 @@ class FallbackProvider(LLMProvider):
|
|||||||
fallback_presets: list[Any],
|
fallback_presets: list[Any],
|
||||||
provider_factory: Callable[[Any], LLMProvider],
|
provider_factory: Callable[[Any], LLMProvider],
|
||||||
fallback_model_observer: FallbackModelObserver | None = None,
|
fallback_model_observer: FallbackModelObserver | None = None,
|
||||||
|
primary_context_window_tokens: int | None = None,
|
||||||
):
|
):
|
||||||
self._primary = primary
|
self._primary = primary
|
||||||
self._fallback_presets = list(fallback_presets)
|
self._fallback_presets = list(fallback_presets)
|
||||||
self._provider_factory = provider_factory
|
self._provider_factory = provider_factory
|
||||||
self._fallback_model_observer = fallback_model_observer
|
self._fallback_model_observer = fallback_model_observer
|
||||||
|
self._primary_context_window_tokens = primary_context_window_tokens
|
||||||
self._has_fallbacks = bool(fallback_presets)
|
self._has_fallbacks = bool(fallback_presets)
|
||||||
self._primary_failures = 0
|
self._primary_failures = 0
|
||||||
self._primary_tripped_at: float | None = None
|
self._primary_tripped_at: float | None = None
|
||||||
@ -141,6 +150,33 @@ class FallbackProvider(LLMProvider):
|
|||||||
def supports_progress_deltas(self) -> bool:
|
def supports_progress_deltas(self) -> bool:
|
||||||
return bool(getattr(self._primary, "supports_progress_deltas", False))
|
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:
|
def _primary_available(self) -> bool:
|
||||||
"""Return True if the primary provider is not currently tripped."""
|
"""Return True if the primary provider is not currently tripped."""
|
||||||
if self._primary_tripped_at is None:
|
if self._primary_tripped_at is None:
|
||||||
@ -157,6 +193,25 @@ class FallbackProvider(LLMProvider):
|
|||||||
lambda p, kw: p.chat(**kw), kwargs, has_streamed=None
|
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:
|
async def chat_stream(self, **kwargs: Any) -> LLMResponse:
|
||||||
on_stream_recover = kwargs.pop("on_stream_recover", None)
|
on_stream_recover = kwargs.pop("on_stream_recover", None)
|
||||||
if not self._has_fallbacks:
|
if not self._has_fallbacks:
|
||||||
@ -179,6 +234,38 @@ class FallbackProvider(LLMProvider):
|
|||||||
on_stream_recover=on_stream_recover,
|
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(
|
async def _try_with_fallback(
|
||||||
self,
|
self,
|
||||||
call: Callable[[LLMProvider, dict[str, Any]], Awaitable[LLMResponse]],
|
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_model = kwargs.get("model") or self._primary.get_default_model()
|
||||||
primary_was_attempted = False
|
primary_was_attempted = False
|
||||||
primary_error = "unknown error"
|
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():
|
if self._primary_available():
|
||||||
primary_was_attempted = True
|
primary_was_attempted = True
|
||||||
@ -286,6 +376,23 @@ class FallbackProvider(LLMProvider):
|
|||||||
"max_tokens": fallback.max_tokens,
|
"max_tokens": fallback.max_tokens,
|
||||||
"temperature": fallback.temperature,
|
"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:
|
if fallback.reasoning_effort is None:
|
||||||
fallback_kwargs.pop("reasoning_effort", None)
|
fallback_kwargs.pop("reasoning_effort", None)
|
||||||
else:
|
else:
|
||||||
@ -312,11 +419,15 @@ class FallbackProvider(LLMProvider):
|
|||||||
)
|
)
|
||||||
# Return the last error response we saw (primary or last fallback).
|
# Return the last error response we saw (primary or last fallback).
|
||||||
if last_response is not None:
|
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.
|
# Primary was tripped and we have no fallbacks — synthesize an error.
|
||||||
return LLMResponse(
|
return LLMResponse(
|
||||||
content=f"Primary model '{primary_model}' circuit open and no fallbacks available",
|
content=f"Primary model '{primary_model}' circuit open and no fallbacks available",
|
||||||
finish_reason="error",
|
finish_reason="error",
|
||||||
|
preserve_provider_state_on_error=preserve_primary_state,
|
||||||
)
|
)
|
||||||
|
|
||||||
async def _notify_fallback_model(self, model: str) -> None:
|
async def _notify_fallback_model(self, model: str) -> None:
|
||||||
|
|||||||
@ -16,7 +16,7 @@ import httpx
|
|||||||
from oauth_cli_kit.models import OAuthToken
|
from oauth_cli_kit.models import OAuthToken
|
||||||
from oauth_cli_kit.storage import FileTokenStorage
|
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
|
from nanobot.providers.openai_compat_provider import OpenAICompatProvider
|
||||||
|
|
||||||
DEFAULT_GITHUB_DEVICE_CODE_URL = "https://github.com/login/device/code"
|
DEFAULT_GITHUB_DEVICE_CODE_URL = "https://github.com/login/device/code"
|
||||||
@ -248,6 +248,7 @@ class GitHubCopilotProvider(OpenAICompatProvider):
|
|||||||
temperature: float = 0.7,
|
temperature: float = 0.7,
|
||||||
reasoning_effort: str | None = None,
|
reasoning_effort: str | None = None,
|
||||||
tool_choice: str | dict[str, Any] | None = None,
|
tool_choice: str | dict[str, Any] | None = None,
|
||||||
|
provider_context: ProviderCallContext | None = None,
|
||||||
) -> LLMResponse:
|
) -> LLMResponse:
|
||||||
await self._refresh_client_api_key()
|
await self._refresh_client_api_key()
|
||||||
return await super().chat(
|
return await super().chat(
|
||||||
@ -258,6 +259,7 @@ class GitHubCopilotProvider(OpenAICompatProvider):
|
|||||||
temperature=temperature,
|
temperature=temperature,
|
||||||
reasoning_effort=reasoning_effort,
|
reasoning_effort=reasoning_effort,
|
||||||
tool_choice=tool_choice,
|
tool_choice=tool_choice,
|
||||||
|
provider_context=provider_context,
|
||||||
)
|
)
|
||||||
|
|
||||||
async def chat_stream(
|
async def chat_stream(
|
||||||
@ -272,6 +274,7 @@ class GitHubCopilotProvider(OpenAICompatProvider):
|
|||||||
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
|
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
|
||||||
on_thinking_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,
|
on_tool_call_delta: Callable[[dict[str, Any]], Awaitable[None]] | None = None,
|
||||||
|
provider_context: ProviderCallContext | None = None,
|
||||||
) -> LLMResponse:
|
) -> LLMResponse:
|
||||||
await self._refresh_client_api_key()
|
await self._refresh_client_api_key()
|
||||||
return await super().chat_stream(
|
return await super().chat_stream(
|
||||||
@ -285,4 +288,5 @@ class GitHubCopilotProvider(OpenAICompatProvider):
|
|||||||
on_content_delta=on_content_delta,
|
on_content_delta=on_content_delta,
|
||||||
on_thinking_delta=on_thinking_delta,
|
on_thinking_delta=on_thinking_delta,
|
||||||
on_tool_call_delta=on_tool_call_delta,
|
on_tool_call_delta=on_tool_call_delta,
|
||||||
|
provider_context=provider_context,
|
||||||
)
|
)
|
||||||
|
|||||||
@ -17,17 +17,27 @@ from oauth_cli_kit import get_token as get_codex_token
|
|||||||
from nanobot.providers.base import (
|
from nanobot.providers.base import (
|
||||||
LLMProvider,
|
LLMProvider,
|
||||||
LLMResponse,
|
LLMResponse,
|
||||||
ToolCallRequest,
|
ProviderCallContext,
|
||||||
|
ProviderConversationState,
|
||||||
resolve_stream_idle_timeout_s,
|
resolve_stream_idle_timeout_s,
|
||||||
)
|
)
|
||||||
from nanobot.providers.openai_responses import (
|
from nanobot.providers.openai_responses import (
|
||||||
|
ResponsesStreamCapture,
|
||||||
|
build_responses_state,
|
||||||
consume_sse_with_reasoning,
|
consume_sse_with_reasoning,
|
||||||
convert_messages,
|
|
||||||
convert_tools,
|
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_CODEX_URL = "https://chatgpt.com/backend-api/codex/responses"
|
||||||
DEFAULT_ORIGINATOR = "nanobot"
|
DEFAULT_ORIGINATOR = "nanobot"
|
||||||
|
_COMPACTION_RETAINED_CHAR_BUDGET = 256_000
|
||||||
|
|
||||||
|
|
||||||
class OpenAICodexProvider(LLMProvider):
|
class OpenAICodexProvider(LLMProvider):
|
||||||
@ -45,21 +55,39 @@ class OpenAICodexProvider(LLMProvider):
|
|||||||
self.default_model = default_model
|
self.default_model = default_model
|
||||||
self.proxy = proxy or None
|
self.proxy = proxy or None
|
||||||
self._extra_body = dict(extra_body or {})
|
self._extra_body = dict(extra_body or {})
|
||||||
|
self._native_compaction_available = True
|
||||||
|
|
||||||
async def _call_codex(
|
async def _call_codex(
|
||||||
self,
|
self,
|
||||||
messages: list[dict[str, Any]],
|
messages: list[dict[str, Any]],
|
||||||
tools: list[dict[str, Any]] | None,
|
tools: list[dict[str, Any]] | None,
|
||||||
model: str | None,
|
model: str | None,
|
||||||
|
max_tokens: int,
|
||||||
reasoning_effort: str | None,
|
reasoning_effort: str | None,
|
||||||
tool_choice: str | dict[str, Any] | None,
|
tool_choice: str | dict[str, Any] | None,
|
||||||
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
|
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
|
||||||
on_thinking_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,
|
on_tool_call_delta: Callable[[dict[str, Any]], Awaitable[None]] | None = None,
|
||||||
|
provider_context: ProviderCallContext | None = None,
|
||||||
) -> LLMResponse:
|
) -> LLMResponse:
|
||||||
"""Shared request logic for both chat() and chat_stream()."""
|
"""Shared request logic for both chat() and chat_stream()."""
|
||||||
model = model or self.default_model
|
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] = {
|
body: dict[str, Any] = {
|
||||||
"model": _strip_model_prefix(model),
|
"model": _strip_model_prefix(model),
|
||||||
@ -68,12 +96,15 @@ class OpenAICodexProvider(LLMProvider):
|
|||||||
"instructions": system_prompt,
|
"instructions": system_prompt,
|
||||||
"input": input_items,
|
"input": input_items,
|
||||||
"text": {"verbosity": "medium"},
|
"text": {"verbosity": "medium"},
|
||||||
"include": ["reasoning.encrypted_content"],
|
|
||||||
"prompt_cache_key": _prompt_cache_key(messages[:2]),
|
"prompt_cache_key": _prompt_cache_key(messages[:2]),
|
||||||
"tool_choice": tool_choice or "auto",
|
"tool_choice": tool_choice or "auto",
|
||||||
"parallel_tool_calls": True,
|
"parallel_tool_calls": True,
|
||||||
}
|
}
|
||||||
|
body["include"] = ["reasoning.encrypted_content"]
|
||||||
reasoning_options = _build_reasoning_options(reasoning_effort)
|
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:
|
if reasoning_options:
|
||||||
body["reasoning"] = reasoning_options
|
body["reasoning"] = reasoning_options
|
||||||
if tools:
|
if tools:
|
||||||
@ -87,33 +118,90 @@ class OpenAICodexProvider(LLMProvider):
|
|||||||
token = await asyncio.to_thread(get_codex_token, proxy=self.proxy)
|
token = await asyncio.to_thread(get_codex_token, proxy=self.proxy)
|
||||||
headers = _build_headers(cast(str, token.account_id), token.access)
|
headers = _build_headers(cast(str, token.account_id), token.access)
|
||||||
|
|
||||||
stage = "codex_request"
|
async def _send(
|
||||||
try:
|
request_body: dict[str, Any],
|
||||||
content, tool_calls, finish_reason, usage, reasoning_content = await _request_codex(
|
*,
|
||||||
DEFAULT_CODEX_URL, headers, body, verify=True,
|
emit_deltas: bool,
|
||||||
proxy=self.proxy,
|
) -> LLMResponse:
|
||||||
on_content_delta=on_content_delta,
|
wire_body = _without_response_item_ids(request_body)
|
||||||
on_thinking_delta=on_thinking_delta,
|
try:
|
||||||
on_tool_call_delta=on_tool_call_delta,
|
return await _request_codex(
|
||||||
)
|
DEFAULT_CODEX_URL,
|
||||||
except Exception as e:
|
headers,
|
||||||
if "CERTIFICATE_VERIFY_FAILED" not in str(e):
|
wire_body,
|
||||||
raise
|
verify=True,
|
||||||
logger.warning("SSL verification failed for Codex API; retrying with verify=False")
|
proxy=self.proxy,
|
||||||
content, tool_calls, finish_reason, usage, reasoning_content = await _request_codex(
|
on_content_delta=on_content_delta if emit_deltas else None,
|
||||||
DEFAULT_CODEX_URL, headers, body, verify=False,
|
on_thinking_delta=on_thinking_delta if emit_deltas else None,
|
||||||
proxy=self.proxy,
|
on_tool_call_delta=on_tool_call_delta if emit_deltas else None,
|
||||||
on_content_delta=on_content_delta,
|
)
|
||||||
on_thinking_delta=on_thinking_delta,
|
except Exception as exc:
|
||||||
on_tool_call_delta=on_tool_call_delta,
|
if "CERTIFICATE_VERIFY_FAILED" not in str(exc):
|
||||||
)
|
raise
|
||||||
return LLMResponse(
|
logger.warning(
|
||||||
content=content,
|
"SSL verification failed for Codex API; retrying with verify=False"
|
||||||
tool_calls=tool_calls,
|
)
|
||||||
finish_reason=finish_reason,
|
return await _request_codex(
|
||||||
usage=usage,
|
DEFAULT_CODEX_URL,
|
||||||
reasoning_content=reasoning_content,
|
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:
|
except Exception as e:
|
||||||
response = _codex_error_response(e)
|
response = _codex_error_response(e)
|
||||||
exc_type = "CodexHTTPError" if isinstance(e, _CodexHTTPError) else type(e).__name__
|
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,
|
model: str | None = None, max_tokens: int = 4096, temperature: float = 0.7,
|
||||||
reasoning_effort: str | None = None,
|
reasoning_effort: str | None = None,
|
||||||
tool_choice: str | dict[str, Any] | None = None,
|
tool_choice: str | dict[str, Any] | None = None,
|
||||||
|
provider_context: ProviderCallContext | None = None,
|
||||||
) -> LLMResponse:
|
) -> 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(
|
async def chat_stream(
|
||||||
self, messages: list[dict[str, Any]], tools: list[dict[str, Any]] | None = None,
|
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_content_delta: Callable[[str], Awaitable[None]] | None = None,
|
||||||
on_thinking_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,
|
on_tool_call_delta: Callable[[dict[str, Any]], Awaitable[None]] | None = None,
|
||||||
|
provider_context: ProviderCallContext | None = None,
|
||||||
) -> LLMResponse:
|
) -> LLMResponse:
|
||||||
return await self._call_codex(
|
return await self._call_codex(
|
||||||
messages,
|
messages=messages,
|
||||||
tools,
|
tools=tools,
|
||||||
model,
|
model=model,
|
||||||
reasoning_effort,
|
max_tokens=max_tokens,
|
||||||
tool_choice,
|
reasoning_effort=reasoning_effort,
|
||||||
on_content_delta,
|
tool_choice=tool_choice,
|
||||||
on_thinking_delta,
|
on_content_delta=on_content_delta,
|
||||||
on_tool_call_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:
|
def get_default_model(self) -> str:
|
||||||
return self.default_model
|
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:
|
def _strip_model_prefix(model: str) -> str:
|
||||||
if model.startswith("openai-codex/") or model.startswith("openai_codex/"):
|
if model.startswith("openai-codex/") or model.startswith("openai_codex/"):
|
||||||
@ -170,6 +312,58 @@ def _strip_model_prefix(model: str) -> str:
|
|||||||
return model
|
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:
|
def _build_reasoning_options(reasoning_effort: str | None) -> dict[str, str] | None:
|
||||||
"""Opt in to visible summaries without changing provider-default effort."""
|
"""Opt in to visible summaries without changing provider-default effort."""
|
||||||
if reasoning_effort and reasoning_effort.lower() == "none":
|
if reasoning_effort and reasoning_effort.lower() == "none":
|
||||||
@ -202,6 +396,7 @@ class _CodexHTTPError(RuntimeError):
|
|||||||
error_type: str | None = None,
|
error_type: str | None = None,
|
||||||
error_code: str | None = None,
|
error_code: str | None = None,
|
||||||
should_retry: bool | None = None,
|
should_retry: bool | None = None,
|
||||||
|
compaction_unsupported: bool = False,
|
||||||
):
|
):
|
||||||
super().__init__(message)
|
super().__init__(message)
|
||||||
self.status_code = status_code
|
self.status_code = status_code
|
||||||
@ -209,6 +404,7 @@ class _CodexHTTPError(RuntimeError):
|
|||||||
self.error_type = error_type
|
self.error_type = error_type
|
||||||
self.error_code = error_code
|
self.error_code = error_code
|
||||||
self.should_retry = should_retry
|
self.should_retry = should_retry
|
||||||
|
self.compaction_unsupported = compaction_unsupported
|
||||||
|
|
||||||
|
|
||||||
async def _request_codex(
|
async def _request_codex(
|
||||||
@ -220,7 +416,7 @@ async def _request_codex(
|
|||||||
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
|
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
|
||||||
on_thinking_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,
|
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()
|
idle_timeout_s = resolve_stream_idle_timeout_s()
|
||||||
client_kwargs: dict[str, Any] = {"timeout": idle_timeout_s, "verify": verify}
|
client_kwargs: dict[str, Any] = {"timeout": idle_timeout_s, "verify": verify}
|
||||||
if proxy:
|
if proxy:
|
||||||
@ -233,6 +429,17 @@ async def _request_codex(
|
|||||||
raw = text.decode("utf-8", "ignore")
|
raw = text.decode("utf-8", "ignore")
|
||||||
retry_after = LLMProvider._extract_retry_after_from_headers(response.headers)
|
retry_after = LLMProvider._extract_retry_after_from_headers(response.headers)
|
||||||
error_type, error_code = LLMProvider._extract_error_type_code(raw)
|
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(
|
raise _CodexHTTPError(
|
||||||
_friendly_error(response.status_code, raw),
|
_friendly_error(response.status_code, raw),
|
||||||
status_code=response.status_code,
|
status_code=response.status_code,
|
||||||
@ -240,13 +447,38 @@ async def _request_codex(
|
|||||||
error_type=error_type,
|
error_type=error_type,
|
||||||
error_code=error_code,
|
error_code=error_code,
|
||||||
should_retry=_should_retry_status(response.status_code, error_type, error_code, raw),
|
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,
|
response,
|
||||||
on_content_delta=on_content_delta,
|
on_content_delta=on_content_delta,
|
||||||
on_tool_call_delta=on_tool_call_delta,
|
on_tool_call_delta=on_tool_call_delta,
|
||||||
on_reasoning_delta=on_thinking_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:
|
def _prompt_cache_key(messages: list[dict[str, Any]]) -> str:
|
||||||
|
|||||||
@ -26,16 +26,24 @@ from pydantic.alias_generators import to_snake
|
|||||||
from nanobot.providers.base import (
|
from nanobot.providers.base import (
|
||||||
LLMProvider,
|
LLMProvider,
|
||||||
LLMResponse,
|
LLMResponse,
|
||||||
|
ProviderCallContext,
|
||||||
|
ProviderConversationState,
|
||||||
ToolCallRequest,
|
ToolCallRequest,
|
||||||
parse_tool_arguments,
|
parse_tool_arguments,
|
||||||
resolve_stream_idle_timeout_s,
|
resolve_stream_idle_timeout_s,
|
||||||
tool_arguments_json_for_replay,
|
tool_arguments_json_for_replay,
|
||||||
)
|
)
|
||||||
from nanobot.providers.openai_responses import (
|
from nanobot.providers.openai_responses import (
|
||||||
|
ResponsesStreamCapture,
|
||||||
|
build_responses_state,
|
||||||
consume_sdk_stream,
|
consume_sdk_stream,
|
||||||
convert_messages,
|
|
||||||
convert_tools,
|
convert_tools,
|
||||||
|
is_compaction_compatibility_error,
|
||||||
|
is_replayable_finish_reason,
|
||||||
parse_response_output,
|
parse_response_output,
|
||||||
|
prepare_responses_input,
|
||||||
|
resolve_compact_threshold,
|
||||||
|
responses_state_matches,
|
||||||
)
|
)
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@ -443,6 +451,8 @@ class OpenAICompatProvider(LLMProvider):
|
|||||||
registry lookups needed.
|
registry lookups needed.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
_native_compaction_available = True
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
api_key: str | None = None,
|
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._api_type = api_type if spec and spec.name == "openai" else "auto"
|
||||||
self._extra_query = extra_query or {}
|
self._extra_query = extra_query or {}
|
||||||
self._proxy = proxy or None
|
self._proxy = proxy or None
|
||||||
|
self._native_compaction_available = True
|
||||||
|
|
||||||
if api_key and spec and spec.env_key:
|
if api_key and spec and spec.env_key:
|
||||||
self._setup_env(api_key, api_base)
|
self._setup_env(api_key, api_base)
|
||||||
@ -971,6 +982,37 @@ class OpenAICompatProvider(LLMProvider):
|
|||||||
|
|
||||||
return self._responses_circuit_allows_probe(model, reasoning_effort)
|
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(
|
def _responses_circuit_allows_probe(
|
||||||
self,
|
self,
|
||||||
model: str | None,
|
model: str | None,
|
||||||
@ -1040,12 +1082,29 @@ class OpenAICompatProvider(LLMProvider):
|
|||||||
temperature: float,
|
temperature: float,
|
||||||
reasoning_effort: str | None,
|
reasoning_effort: str | None,
|
||||||
tool_choice: str | dict[str, Any] | None,
|
tool_choice: str | dict[str, Any] | None,
|
||||||
|
provider_context: ProviderCallContext | None = None,
|
||||||
) -> dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
"""Build a Responses API body for direct OpenAI requests."""
|
"""Build a Responses API body for direct OpenAI requests."""
|
||||||
model_name = model or self.default_model
|
model_name = model or self.default_model
|
||||||
model_name = self._request_model_name(model_name)
|
model_name = self._request_model_name(model_name)
|
||||||
sanitized_messages = self._sanitize_messages(self._sanitize_empty_content(messages))
|
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] = {
|
body: dict[str, Any] = {
|
||||||
"model": model_name,
|
"model": model_name,
|
||||||
@ -1055,13 +1114,29 @@ class OpenAICompatProvider(LLMProvider):
|
|||||||
"store": False,
|
"store": False,
|
||||||
"stream": 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):
|
if self._supports_temperature(model_name, reasoning_effort):
|
||||||
body["temperature"] = temperature
|
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":
|
if reasoning_effort and reasoning_effort.lower() != "none":
|
||||||
body["reasoning"] = {"effort": reasoning_effort}
|
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:
|
if tools:
|
||||||
body["tools"] = convert_tools(tools)
|
body["tools"] = convert_tools(tools)
|
||||||
@ -1073,6 +1148,29 @@ class OpenAICompatProvider(LLMProvider):
|
|||||||
|
|
||||||
return body
|
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
|
# Response parsing
|
||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
@ -1599,6 +1697,28 @@ class OpenAICompatProvider(LLMProvider):
|
|||||||
# Public API
|
# 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(
|
async def chat(
|
||||||
self,
|
self,
|
||||||
messages: list[dict[str, Any]],
|
messages: list[dict[str, Any]],
|
||||||
@ -1608,6 +1728,7 @@ class OpenAICompatProvider(LLMProvider):
|
|||||||
temperature: float = 0.7,
|
temperature: float = 0.7,
|
||||||
reasoning_effort: str | None = None,
|
reasoning_effort: str | None = None,
|
||||||
tool_choice: str | dict[str, Any] | None = None,
|
tool_choice: str | dict[str, Any] | None = None,
|
||||||
|
provider_context: ProviderCallContext | None = None,
|
||||||
) -> LLMResponse:
|
) -> LLMResponse:
|
||||||
client = await self._ensure_client()
|
client = await self._ensure_client()
|
||||||
try:
|
try:
|
||||||
@ -1616,12 +1737,18 @@ class OpenAICompatProvider(LLMProvider):
|
|||||||
body = self._build_responses_body(
|
body = self._build_responses_body(
|
||||||
messages, tools, model, max_tokens, temperature,
|
messages, tools, model, max_tokens, temperature,
|
||||||
reasoning_effort, tool_choice,
|
reasoning_effort, tool_choice,
|
||||||
|
provider_context,
|
||||||
)
|
)
|
||||||
responses_raw = cast(
|
responses_raw = await self._create_response_with_compaction_fallback(
|
||||||
Any,
|
client,
|
||||||
await client.responses.create(**body),
|
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)
|
self._record_responses_success(model, reasoning_effort)
|
||||||
return result
|
return result
|
||||||
except Exception as responses_error:
|
except Exception as responses_error:
|
||||||
@ -1660,6 +1787,7 @@ class OpenAICompatProvider(LLMProvider):
|
|||||||
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
|
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
|
||||||
on_thinking_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,
|
on_tool_call_delta: Callable[[dict[str, Any]], Awaitable[None]] | None = None,
|
||||||
|
provider_context: ProviderCallContext | None = None,
|
||||||
) -> LLMResponse:
|
) -> LLMResponse:
|
||||||
client = await self._ensure_client()
|
client = await self._ensure_client()
|
||||||
idle_timeout_s = resolve_stream_idle_timeout_s()
|
idle_timeout_s = resolve_stream_idle_timeout_s()
|
||||||
@ -1669,11 +1797,12 @@ class OpenAICompatProvider(LLMProvider):
|
|||||||
body = self._build_responses_body(
|
body = self._build_responses_body(
|
||||||
messages, tools, model, max_tokens, temperature,
|
messages, tools, model, max_tokens, temperature,
|
||||||
reasoning_effort, tool_choice,
|
reasoning_effort, tool_choice,
|
||||||
|
provider_context,
|
||||||
)
|
)
|
||||||
body["stream"] = True
|
body["stream"] = True
|
||||||
responses_stream = cast(
|
responses_stream = await self._create_response_with_compaction_fallback(
|
||||||
Any,
|
client,
|
||||||
await client.responses.create(**body),
|
body,
|
||||||
)
|
)
|
||||||
|
|
||||||
async def _timed_stream() -> AsyncIterator[Any]:
|
async def _timed_stream() -> AsyncIterator[Any]:
|
||||||
@ -1687,6 +1816,7 @@ class OpenAICompatProvider(LLMProvider):
|
|||||||
except StopAsyncIteration:
|
except StopAsyncIteration:
|
||||||
break
|
break
|
||||||
|
|
||||||
|
capture = ResponsesStreamCapture()
|
||||||
(
|
(
|
||||||
content,
|
content,
|
||||||
tool_calls,
|
tool_calls,
|
||||||
@ -1697,15 +1827,25 @@ class OpenAICompatProvider(LLMProvider):
|
|||||||
_timed_stream(),
|
_timed_stream(),
|
||||||
on_content_delta,
|
on_content_delta,
|
||||||
on_tool_call_delta=on_tool_call_delta,
|
on_tool_call_delta=on_tool_call_delta,
|
||||||
|
capture=capture,
|
||||||
)
|
)
|
||||||
self._record_responses_success(model, reasoning_effort)
|
self._record_responses_success(model, reasoning_effort)
|
||||||
return LLMResponse(
|
result = LLMResponse(
|
||||||
content=content or None,
|
content=content or None,
|
||||||
tool_calls=tool_calls,
|
tool_calls=tool_calls,
|
||||||
finish_reason=finish_reason,
|
finish_reason=finish_reason,
|
||||||
usage=usage,
|
usage=usage,
|
||||||
reasoning_content=reasoning_content,
|
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:
|
except Exception as responses_error:
|
||||||
if self._spec and self._spec.name == "github_copilot":
|
if self._spec and self._spec.name == "github_copilot":
|
||||||
# Copilot gateway exposes GPT-5/o-series only via /responses;
|
# Copilot gateway exposes GPT-5/o-series only via /responses;
|
||||||
|
|||||||
@ -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 (
|
from nanobot.providers.openai_responses.converters import (
|
||||||
convert_messages,
|
convert_messages,
|
||||||
@ -8,13 +8,24 @@ from nanobot.providers.openai_responses.converters import (
|
|||||||
)
|
)
|
||||||
from nanobot.providers.openai_responses.parsing import (
|
from nanobot.providers.openai_responses.parsing import (
|
||||||
FINISH_REASON_MAP,
|
FINISH_REASON_MAP,
|
||||||
|
ResponsesStreamCapture,
|
||||||
consume_sdk_stream,
|
consume_sdk_stream,
|
||||||
consume_sse,
|
consume_sse,
|
||||||
consume_sse_with_reasoning,
|
consume_sse_with_reasoning,
|
||||||
|
is_replayable_finish_reason,
|
||||||
iter_sse,
|
iter_sse,
|
||||||
map_finish_reason,
|
map_finish_reason,
|
||||||
parse_response_output,
|
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__ = [
|
__all__ = [
|
||||||
"convert_messages",
|
"convert_messages",
|
||||||
@ -25,7 +36,16 @@ __all__ = [
|
|||||||
"consume_sse",
|
"consume_sse",
|
||||||
"consume_sse_with_reasoning",
|
"consume_sse_with_reasoning",
|
||||||
"consume_sdk_stream",
|
"consume_sdk_stream",
|
||||||
|
"ResponsesStreamCapture",
|
||||||
|
"is_replayable_finish_reason",
|
||||||
"map_finish_reason",
|
"map_finish_reason",
|
||||||
"parse_response_output",
|
"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",
|
"FINISH_REASON_MAP",
|
||||||
]
|
]
|
||||||
|
|||||||
@ -4,12 +4,14 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import json
|
import json
|
||||||
from collections.abc import Awaitable, Callable
|
from collections.abc import Awaitable, Callable
|
||||||
|
from dataclasses import dataclass, field
|
||||||
from typing import Any, AsyncGenerator, cast
|
from typing import Any, AsyncGenerator, cast
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
|
||||||
from nanobot.providers.base import LLMResponse, ToolCallRequest, parse_tool_arguments
|
from nanobot.providers.base import LLMResponse, ToolCallRequest, parse_tool_arguments
|
||||||
|
from nanobot.providers.openai_responses.state import build_responses_state
|
||||||
|
|
||||||
FINISH_REASON_MAP = {
|
FINISH_REASON_MAP = {
|
||||||
"completed": "stop",
|
"completed": "stop",
|
||||||
@ -17,6 +19,42 @@ FINISH_REASON_MAP = {
|
|||||||
"failed": "error",
|
"failed": "error",
|
||||||
"cancelled": "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:
|
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")
|
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]:
|
def _usage_from_response_obj(response: object) -> dict[str, int]:
|
||||||
response_object = _response_object(response)
|
response_object = _response_object(response)
|
||||||
usage_raw: object = (
|
usage_raw: object = (
|
||||||
@ -99,6 +158,47 @@ def _tool_arguments_source(*values: Any) -> Any:
|
|||||||
return "{}"
|
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]:
|
async def iter_sse(response: httpx.Response) -> AsyncGenerator[dict[str, Any], None]:
|
||||||
"""Yield parsed JSON events from a Responses API SSE stream."""
|
"""Yield parsed JSON events from a Responses API SSE stream."""
|
||||||
buffer: list[str] = []
|
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_tool_call_delta: Callable[[dict[str, Any]], Awaitable[None]] | None = None,
|
||||||
on_reasoning_delta: Callable[[str], Awaitable[None]] | None = None,
|
on_reasoning_delta: Callable[[str], Awaitable[None]] | None = None,
|
||||||
on_response_event: Callable[[dict[str, Any]], 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]:
|
) -> tuple[str, list[ToolCallRequest], str, dict[str, int], str | None]:
|
||||||
"""Consume a Responses API SSE stream, including visible reasoning summaries."""
|
"""Consume a Responses API SSE stream, including visible reasoning summaries."""
|
||||||
content = ""
|
content = ""
|
||||||
@ -163,6 +264,9 @@ async def consume_sse_with_reasoning(
|
|||||||
usage: dict[str, int] = {}
|
usage: dict[str, int] = {}
|
||||||
reasoning_content: str | None = None
|
reasoning_content: str | None = None
|
||||||
streamed_reasoning = False
|
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):
|
async for event in iter_sse(response):
|
||||||
if on_response_event:
|
if on_response_event:
|
||||||
@ -191,6 +295,33 @@ async def consume_sse_with_reasoning(
|
|||||||
content += delta_text
|
content += delta_text
|
||||||
if on_content_delta and delta_text:
|
if on_content_delta and delta_text:
|
||||||
await on_content_delta(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":
|
elif event_type == "response.reasoning_summary_text.delta":
|
||||||
delta_text = event.get("delta") or ""
|
delta_text = event.get("delta") or ""
|
||||||
if delta_text:
|
if delta_text:
|
||||||
@ -239,6 +370,8 @@ async def consume_sse_with_reasoning(
|
|||||||
})
|
})
|
||||||
elif event_type == "response.output_item.done":
|
elif event_type == "response.output_item.done":
|
||||||
item = _as_json_object(event.get("item")) or {}
|
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":
|
if item.get("type") == "function_call":
|
||||||
call_id = item.get("call_id")
|
call_id = item.get("call_id")
|
||||||
if not call_id:
|
if not call_id:
|
||||||
@ -269,11 +402,28 @@ async def consume_sse_with_reasoning(
|
|||||||
reasoning_content = summary
|
reasoning_content = summary
|
||||||
if on_reasoning_delta:
|
if on_reasoning_delta:
|
||||||
await on_reasoning_delta(summary)
|
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 {}
|
response_obj = _response_object(event.get("response")) or {}
|
||||||
status = response_obj.get("status")
|
if capture is not None:
|
||||||
finish_reason = map_finish_reason(status)
|
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
|
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:
|
if not reasoning_content:
|
||||||
summary = _extract_reasoning_summary_from_output(response_obj.get("output"))
|
summary = _extract_reasoning_summary_from_output(response_obj.get("output"))
|
||||||
if summary:
|
if summary:
|
||||||
@ -284,6 +434,8 @@ async def consume_sse_with_reasoning(
|
|||||||
detail = event.get("error") or event.get("message") or event
|
detail = event.get("error") or event.get("message") or event
|
||||||
raise RuntimeError(f"Response failed: {str(detail)[:500]}")
|
raise RuntimeError(f"Response failed: {str(detail)[:500]}")
|
||||||
|
|
||||||
|
if refusal_seen:
|
||||||
|
finish_reason = "refusal"
|
||||||
return content, tool_calls, finish_reason, usage, reasoning_content
|
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
|
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``."""
|
"""Parse an SDK ``Response`` object into an ``LLMResponse``."""
|
||||||
response_object = _response_object(response) or {}
|
response_object = _response_object(response) or {}
|
||||||
|
|
||||||
@ -308,15 +466,22 @@ def parse_response_output(response: object) -> LLMResponse:
|
|||||||
content_parts: list[str] = []
|
content_parts: list[str] = []
|
||||||
tool_calls: list[ToolCallRequest] = []
|
tool_calls: list[ToolCallRequest] = []
|
||||||
reasoning_content: str | None = None
|
reasoning_content: str | None = None
|
||||||
|
refusal_seen = False
|
||||||
|
|
||||||
for item in output:
|
for item in output:
|
||||||
item_type = item.get("type")
|
item_type = item.get("type")
|
||||||
if item_type == "message":
|
if item_type == "message":
|
||||||
for block in _response_object_list(item.get("content")):
|
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")
|
text = block.get("text")
|
||||||
if isinstance(text, str):
|
if isinstance(text, str):
|
||||||
content_parts.append(text)
|
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":
|
elif item_type == "reasoning":
|
||||||
for s in _response_object_list(item.get("summary")):
|
for s in _response_object_list(item.get("summary")):
|
||||||
if s.get("type") == "summary_text" and s.get("text"):
|
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)
|
usage = _usage_from_response_obj(response_object)
|
||||||
|
|
||||||
status = response_object.get("status")
|
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,
|
content="".join(content_parts) or None,
|
||||||
tool_calls=tool_calls,
|
tool_calls=tool_calls,
|
||||||
finish_reason=finish_reason,
|
finish_reason=finish_reason,
|
||||||
usage=usage,
|
usage=usage,
|
||||||
reasoning_content=reasoning_content if isinstance(reasoning_content, str) else None,
|
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(
|
async def consume_sdk_stream(
|
||||||
stream: Any,
|
stream: Any,
|
||||||
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
|
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
|
||||||
on_tool_call_delta: Callable[[dict[str, Any]], 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]:
|
) -> tuple[str, list[ToolCallRequest], str, dict[str, int], str | None]:
|
||||||
"""Consume an SDK async stream from ``client.responses.create(stream=True)``."""
|
"""Consume an SDK async stream from ``client.responses.create(stream=True)``."""
|
||||||
content = ""
|
content = ""
|
||||||
@ -361,6 +542,9 @@ async def consume_sdk_stream(
|
|||||||
finish_reason = "stop"
|
finish_reason = "stop"
|
||||||
usage: dict[str, int] = {}
|
usage: dict[str, int] = {}
|
||||||
reasoning_content: str | None = None
|
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:
|
async for raw_event in stream:
|
||||||
event: Any = raw_event
|
event: Any = raw_event
|
||||||
@ -388,6 +572,33 @@ async def consume_sdk_stream(
|
|||||||
content += delta_text
|
content += delta_text
|
||||||
if on_content_delta and delta_text:
|
if on_content_delta and delta_text:
|
||||||
await on_content_delta(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":
|
elif event_type == "response.function_call_arguments.delta":
|
||||||
call_id = getattr(event, "call_id", None)
|
call_id = getattr(event, "call_id", None)
|
||||||
if call_id and call_id in tool_call_buffers:
|
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":
|
elif event_type == "response.output_item.done":
|
||||||
item = getattr(event, "item", None)
|
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":
|
if item and getattr(item, "type", None) == "function_call":
|
||||||
call_id = getattr(item, "call_id", None)
|
call_id = getattr(item, "call_id", None)
|
||||||
if not call_id:
|
if not call_id:
|
||||||
@ -443,10 +656,31 @@ async def consume_sdk_stream(
|
|||||||
arguments=args,
|
arguments=args,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
elif event_type == "response.completed":
|
elif event_type in {"response.completed", "response.incomplete"}:
|
||||||
resp = getattr(event, "response", None)
|
resp = getattr(event, "response", None)
|
||||||
status = getattr(resp, "status", None) if resp else None
|
response_obj = _response_object(resp) or {}
|
||||||
finish_reason = map_finish_reason(status)
|
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:
|
if resp:
|
||||||
usage_obj = getattr(resp, "usage", None)
|
usage_obj = getattr(resp, "usage", None)
|
||||||
if usage_obj:
|
if usage_obj:
|
||||||
@ -466,4 +700,6 @@ async def consume_sdk_stream(
|
|||||||
detail = getattr(event, "error", None) or getattr(event, "message", None) or event
|
detail = getattr(event, "error", None) or getattr(event, "message", None) or event
|
||||||
raise RuntimeError(f"Response failed: {str(detail)[:500]}")
|
raise RuntimeError(f"Response failed: {str(detail)[:500]}")
|
||||||
|
|
||||||
|
if refusal_seen:
|
||||||
|
finish_reason = "refusal"
|
||||||
return content, tool_calls, finish_reason, usage, reasoning_content
|
return content, tool_calls, finish_reason, usage, reasoning_content
|
||||||
|
|||||||
197
nanobot/providers/openai_responses/state.py
Normal file
197
nanobot/providers/openai_responses/state.py
Normal file
@ -0,0 +1,197 @@
|
|||||||
|
"""Opaque conversation state for Responses API item replay."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from copy import deepcopy
|
||||||
|
from typing import Any, cast
|
||||||
|
|
||||||
|
from loguru import logger
|
||||||
|
|
||||||
|
from nanobot.providers.base import ProviderConversationState
|
||||||
|
from nanobot.providers.openai_responses.converters import convert_messages
|
||||||
|
|
||||||
|
RESPONSES_STATE_KIND = "openai_responses"
|
||||||
|
RESPONSES_STATE_VERSION = 1
|
||||||
|
_ITEMS_KEY = "items"
|
||||||
|
_CONTEXT_TOKENS_KEY = "context_tokens"
|
||||||
|
_COMPACTION_ITEM_TYPES = frozenset({
|
||||||
|
"compaction",
|
||||||
|
"compaction_summary",
|
||||||
|
"context_compaction",
|
||||||
|
})
|
||||||
|
|
||||||
|
|
||||||
|
def responses_state_matches(
|
||||||
|
state: ProviderConversationState,
|
||||||
|
*,
|
||||||
|
provider: str,
|
||||||
|
model: str,
|
||||||
|
) -> bool:
|
||||||
|
"""Return whether *state* belongs to this exact Responses endpoint/model."""
|
||||||
|
return (
|
||||||
|
state.kind == RESPONSES_STATE_KIND
|
||||||
|
and state.version == RESPONSES_STATE_VERSION
|
||||||
|
and state.provider == provider
|
||||||
|
and state.model == model
|
||||||
|
and _state_items(state) is not None
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def prepare_responses_input(
|
||||||
|
messages: list[dict[str, Any]],
|
||||||
|
*,
|
||||||
|
state: ProviderConversationState | None,
|
||||||
|
provider: str,
|
||||||
|
model: str,
|
||||||
|
) -> tuple[str, list[dict[str, Any]], bool]:
|
||||||
|
"""Build a request from exact prior items plus only newly appended messages.
|
||||||
|
|
||||||
|
The full Chat transcript remains the source for the current instructions.
|
||||||
|
When no compatible state exists, it is converted normally as a safe
|
||||||
|
fallback.
|
||||||
|
"""
|
||||||
|
instructions, fallback_items = convert_messages(messages)
|
||||||
|
if state is None or not responses_state_matches(
|
||||||
|
state,
|
||||||
|
provider=provider,
|
||||||
|
model=model,
|
||||||
|
):
|
||||||
|
return instructions, fallback_items, False
|
||||||
|
|
||||||
|
prior_items = _state_items(state)
|
||||||
|
if prior_items is None:
|
||||||
|
return instructions, fallback_items, False
|
||||||
|
|
||||||
|
_, delta_items = convert_messages(state.pending_messages)
|
||||||
|
logger.debug(
|
||||||
|
"Replaying Responses state: prior_items={} pending_messages={}",
|
||||||
|
len(prior_items),
|
||||||
|
len(state.pending_messages),
|
||||||
|
)
|
||||||
|
return instructions, [*deepcopy(prior_items), *delta_items], True
|
||||||
|
|
||||||
|
|
||||||
|
def build_responses_state(
|
||||||
|
*,
|
||||||
|
provider: str,
|
||||||
|
model: str,
|
||||||
|
input_items: list[dict[str, Any]],
|
||||||
|
output_items: list[dict[str, Any]],
|
||||||
|
usage: dict[str, int] | None = None,
|
||||||
|
) -> ProviderConversationState:
|
||||||
|
"""Create the canonical next state from request input and every output item."""
|
||||||
|
unpruned_items = [*input_items, *output_items]
|
||||||
|
items = _prune_before_latest_output_compaction(input_items, output_items)
|
||||||
|
if len(items) < len(unpruned_items):
|
||||||
|
logger.info(
|
||||||
|
"Installed Responses compaction: dropped_items={} retained_items={}",
|
||||||
|
len(unpruned_items) - len(items),
|
||||||
|
len(items),
|
||||||
|
)
|
||||||
|
payload: dict[str, Any] = {_ITEMS_KEY: deepcopy(items)}
|
||||||
|
context_tokens = _context_tokens_from_usage(usage)
|
||||||
|
if context_tokens > 0:
|
||||||
|
payload[_CONTEXT_TOKENS_KEY] = context_tokens
|
||||||
|
return ProviderConversationState(
|
||||||
|
kind=RESPONSES_STATE_KIND,
|
||||||
|
provider=provider,
|
||||||
|
model=model,
|
||||||
|
version=RESPONSES_STATE_VERSION,
|
||||||
|
payload=payload,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def responses_state_items(
|
||||||
|
state: ProviderConversationState,
|
||||||
|
) -> list[dict[str, Any]] | None:
|
||||||
|
"""Return an isolated copy of canonical input items for tests/consumers."""
|
||||||
|
items = _state_items(state)
|
||||||
|
return deepcopy(items) if items is not None else None
|
||||||
|
|
||||||
|
|
||||||
|
def responses_state_context_tokens(state: ProviderConversationState) -> int:
|
||||||
|
"""Return the last server-reported active context size."""
|
||||||
|
value = state.payload.get(_CONTEXT_TOKENS_KEY)
|
||||||
|
if isinstance(value, bool) or not isinstance(value, int):
|
||||||
|
return 0
|
||||||
|
return max(0, value)
|
||||||
|
|
||||||
|
|
||||||
|
def resolve_compact_threshold(
|
||||||
|
context_window_tokens: int | None,
|
||||||
|
max_output_tokens: int,
|
||||||
|
) -> int | None:
|
||||||
|
"""Derive Codex-compatible 90% compaction headroom for a model window."""
|
||||||
|
if context_window_tokens is None or context_window_tokens <= 0:
|
||||||
|
return None
|
||||||
|
ninety_percent = max(1, context_window_tokens * 9 // 10)
|
||||||
|
output_headroom = max(1, context_window_tokens - max(1, max_output_tokens))
|
||||||
|
return min(ninety_percent, output_headroom)
|
||||||
|
|
||||||
|
|
||||||
|
def is_compaction_compatibility_error(exc: Exception) -> bool:
|
||||||
|
"""Recognize endpoints that reject native Responses compaction fields."""
|
||||||
|
if getattr(exc, "compaction_unsupported", False) is True:
|
||||||
|
return True
|
||||||
|
response = getattr(exc, "response", None)
|
||||||
|
status_code = getattr(exc, "status_code", None)
|
||||||
|
if status_code is None and response is not None:
|
||||||
|
status_code = getattr(response, "status_code", None)
|
||||||
|
body = (
|
||||||
|
getattr(exc, "body", None)
|
||||||
|
or getattr(exc, "doc", None)
|
||||||
|
or getattr(response, "text", None)
|
||||||
|
or str(exc)
|
||||||
|
)
|
||||||
|
text = str(body).lower()
|
||||||
|
has_compaction_marker = any(
|
||||||
|
marker in text
|
||||||
|
for marker in ("context_management", "compact_threshold", "compaction_trigger")
|
||||||
|
)
|
||||||
|
if not has_compaction_marker:
|
||||||
|
return False
|
||||||
|
return isinstance(exc, TypeError) or status_code in {400, 404, 422}
|
||||||
|
|
||||||
|
|
||||||
|
def _prune_before_latest_output_compaction(
|
||||||
|
input_items: list[dict[str, Any]],
|
||||||
|
output_items: list[dict[str, Any]],
|
||||||
|
) -> list[dict[str, Any]]:
|
||||||
|
"""Drop old input only when this response emits a new compaction item.
|
||||||
|
|
||||||
|
A canonical compacted input may intentionally retain messages before its
|
||||||
|
compaction item. Those messages must survive ordinary subsequent responses.
|
||||||
|
"""
|
||||||
|
latest = None
|
||||||
|
for index, item in enumerate(output_items):
|
||||||
|
if item.get("type") in _COMPACTION_ITEM_TYPES:
|
||||||
|
latest = index
|
||||||
|
if latest is None:
|
||||||
|
return [*input_items, *output_items]
|
||||||
|
return output_items[latest:]
|
||||||
|
|
||||||
|
|
||||||
|
def _context_tokens_from_usage(usage: dict[str, int] | None) -> int:
|
||||||
|
if not usage:
|
||||||
|
return 0
|
||||||
|
prompt_tokens = usage.get("prompt_tokens", 0)
|
||||||
|
completion_tokens = usage.get("completion_tokens", 0)
|
||||||
|
total_tokens = usage.get("total_tokens", 0)
|
||||||
|
values = (prompt_tokens, completion_tokens, total_tokens)
|
||||||
|
if any(isinstance(value, bool) for value in values):
|
||||||
|
return 0
|
||||||
|
return max(0, total_tokens or prompt_tokens + completion_tokens)
|
||||||
|
|
||||||
|
|
||||||
|
def _state_items(
|
||||||
|
state: ProviderConversationState,
|
||||||
|
) -> list[dict[str, Any]] | None:
|
||||||
|
raw_items = state.payload.get(_ITEMS_KEY)
|
||||||
|
if not isinstance(raw_items, list):
|
||||||
|
return None
|
||||||
|
items: list[dict[str, Any]] = []
|
||||||
|
for raw in cast(list[object], raw_items):
|
||||||
|
if not isinstance(raw, dict):
|
||||||
|
return None
|
||||||
|
items.append(cast(dict[str, Any], raw))
|
||||||
|
return items
|
||||||
@ -17,6 +17,7 @@ from weakref import WeakValueDictionary
|
|||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
|
||||||
from nanobot.config.paths import get_legacy_sessions_dir
|
from nanobot.config.paths import get_legacy_sessions_dir
|
||||||
|
from nanobot.providers.base import ProviderConversationState
|
||||||
from nanobot.runtime_context import (
|
from nanobot.runtime_context import (
|
||||||
RUNTIME_CONTEXT_HISTORY_META,
|
RUNTIME_CONTEXT_HISTORY_META,
|
||||||
public_history_message,
|
public_history_message,
|
||||||
@ -43,6 +44,10 @@ _SESSION_PREVIEW_MAX_CHARS = 120
|
|||||||
_SESSION_LIST_PREVIEW_MAX_RECORDS = 200
|
_SESSION_LIST_PREVIEW_MAX_RECORDS = 200
|
||||||
_SESSION_LIST_PREVIEW_MAX_CHARS = 1_000_000
|
_SESSION_LIST_PREVIEW_MAX_CHARS = 1_000_000
|
||||||
_SESSION_DATA_ERRORS = (ValueError, TypeError, AttributeError, KeyError)
|
_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 = {
|
_FORK_VOLATILE_METADATA_KEYS = {
|
||||||
"goal_state",
|
"goal_state",
|
||||||
"pending_user_turn",
|
"pending_user_turn",
|
||||||
@ -60,6 +65,11 @@ def _json_object(value: object) -> dict[str, Any]:
|
|||||||
return cast(dict[str, Any], value)
|
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:
|
def replay_max_messages_for_context(context_window_tokens: int | None) -> int:
|
||||||
if not context_window_tokens or context_window_tokens <= 0:
|
if not context_window_tokens or context_window_tokens <= 0:
|
||||||
return FILE_MAX_MESSAGES
|
return FILE_MAX_MESSAGES
|
||||||
@ -146,10 +156,13 @@ class Session:
|
|||||||
updated_at: datetime = field(default_factory=datetime.now)
|
updated_at: datetime = field(default_factory=datetime.now)
|
||||||
metadata: dict[str, Any] = field(default_factory=dict)
|
metadata: dict[str, Any] = field(default_factory=dict)
|
||||||
last_consolidated: int = 0 # Number of messages already consolidated to files
|
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:
|
def __post_init__(self) -> None:
|
||||||
if not isinstance(cast(object, self.metadata), dict):
|
if not isinstance(cast(object, self.metadata), dict):
|
||||||
self.metadata = {}
|
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.
|
# An out-of-range offset (corrupt metadata) would hide all history; reset it.
|
||||||
last_consolidated = cast(object, self.last_consolidated)
|
last_consolidated = cast(object, self.last_consolidated)
|
||||||
if (
|
if (
|
||||||
@ -304,6 +317,7 @@ class Session:
|
|||||||
"""Clear all messages and reset session to initial state."""
|
"""Clear all messages and reset session to initial state."""
|
||||||
self.messages = []
|
self.messages = []
|
||||||
self.last_consolidated = 0
|
self.last_consolidated = 0
|
||||||
|
self.provider_state = None
|
||||||
self.updated_at = datetime.now()
|
self.updated_at = datetime.now()
|
||||||
self.metadata.pop("_last_summary", None)
|
self.metadata.pop("_last_summary", None)
|
||||||
|
|
||||||
@ -396,6 +410,8 @@ class Session:
|
|||||||
|
|
||||||
self.messages = retained
|
self.messages = retained
|
||||||
self.last_consolidated = new_lc
|
self.last_consolidated = new_lc
|
||||||
|
if dropped:
|
||||||
|
self.provider_state = None
|
||||||
self.updated_at = datetime.now()
|
self.updated_at = datetime.now()
|
||||||
return RetentionResult(
|
return RetentionResult(
|
||||||
dropped=dropped,
|
dropped=dropped,
|
||||||
@ -517,6 +533,7 @@ class JsonlSessionStore:
|
|||||||
created_at: datetime | None = None
|
created_at: datetime | None = None
|
||||||
updated_at: datetime | None = None
|
updated_at: datetime | None = None
|
||||||
last_consolidated = 0
|
last_consolidated = 0
|
||||||
|
provider_state: ProviderConversationState | None = None
|
||||||
|
|
||||||
with open(path, encoding="utf-8") as f:
|
with open(path, encoding="utf-8") as f:
|
||||||
for line in f:
|
for line in f:
|
||||||
@ -527,7 +544,8 @@ class JsonlSessionStore:
|
|||||||
raw_data: object = json.loads(line)
|
raw_data: object = json.loads(line)
|
||||||
data = _json_object(raw_data)
|
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_value = cast(object, data.get("metadata", {}))
|
||||||
metadata = (
|
metadata = (
|
||||||
cast(dict[str, Any], metadata_value)
|
cast(dict[str, Any], metadata_value)
|
||||||
@ -552,6 +570,10 @@ class JsonlSessionStore:
|
|||||||
if isinstance(offset, int) and not isinstance(offset, bool)
|
if isinstance(offset, int) and not isinstance(offset, bool)
|
||||||
else 0
|
else 0
|
||||||
)
|
)
|
||||||
|
elif record_type == _PROVIDER_STATE_RECORD_TYPE:
|
||||||
|
provider_state = ProviderConversationState.from_private_record(
|
||||||
|
data.get("state")
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
messages.append(data)
|
messages.append(data)
|
||||||
|
|
||||||
@ -562,6 +584,7 @@ class JsonlSessionStore:
|
|||||||
updated_at=updated_at or datetime.now(),
|
updated_at=updated_at or datetime.now(),
|
||||||
metadata=metadata,
|
metadata=metadata,
|
||||||
last_consolidated=last_consolidated,
|
last_consolidated=last_consolidated,
|
||||||
|
provider_state=provider_state,
|
||||||
)
|
)
|
||||||
except _SESSION_DATA_ERRORS as e:
|
except _SESSION_DATA_ERRORS as e:
|
||||||
logger.warning("Failed to load session {}: {}", key, e)
|
logger.warning("Failed to load session {}: {}", key, e)
|
||||||
@ -586,6 +609,7 @@ class JsonlSessionStore:
|
|||||||
created_at: datetime | None = None
|
created_at: datetime | None = None
|
||||||
updated_at: datetime | None = None
|
updated_at: datetime | None = None
|
||||||
last_consolidated = 0
|
last_consolidated = 0
|
||||||
|
provider_state: ProviderConversationState | None = None
|
||||||
skipped = 0
|
skipped = 0
|
||||||
|
|
||||||
with open(path, encoding="utf-8") as f:
|
with open(path, encoding="utf-8") as f:
|
||||||
@ -603,7 +627,8 @@ class JsonlSessionStore:
|
|||||||
continue
|
continue
|
||||||
data = cast(dict[str, Any], raw_data)
|
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_value = cast(object, data.get("metadata", {}))
|
||||||
metadata = (
|
metadata = (
|
||||||
cast(dict[str, Any], metadata_value)
|
cast(dict[str, Any], metadata_value)
|
||||||
@ -624,13 +649,21 @@ class JsonlSessionStore:
|
|||||||
if isinstance(offset, int) and not isinstance(offset, bool)
|
if isinstance(offset, int) and not isinstance(offset, bool)
|
||||||
else 0
|
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:
|
else:
|
||||||
messages.append(data)
|
messages.append(data)
|
||||||
|
|
||||||
if skipped:
|
if skipped:
|
||||||
logger.warning("Skipped {} corrupt lines in session {}", skipped, key)
|
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 None
|
||||||
|
|
||||||
return Session(
|
return Session(
|
||||||
@ -640,6 +673,7 @@ class JsonlSessionStore:
|
|||||||
updated_at=updated_at or datetime.now(),
|
updated_at=updated_at or datetime.now(),
|
||||||
metadata=metadata,
|
metadata=metadata,
|
||||||
last_consolidated=last_consolidated,
|
last_consolidated=last_consolidated,
|
||||||
|
provider_state=provider_state,
|
||||||
)
|
)
|
||||||
except _SESSION_DATA_ERRORS as e:
|
except _SESSION_DATA_ERRORS as e:
|
||||||
logger.warning("Repair failed for session {}: {}", key, e)
|
logger.warning("Repair failed for session {}: {}", key, e)
|
||||||
@ -670,6 +704,12 @@ class JsonlSessionStore:
|
|||||||
"last_consolidated": session.last_consolidated,
|
"last_consolidated": session.last_consolidated,
|
||||||
}
|
}
|
||||||
f.write(json.dumps(metadata_line, ensure_ascii=False) + "\n")
|
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:
|
for msg in session.messages:
|
||||||
f.write(json.dumps(msg, ensure_ascii=False) + "\n")
|
f.write(json.dumps(msg, ensure_ascii=False) + "\n")
|
||||||
if fsync:
|
if fsync:
|
||||||
@ -726,7 +766,8 @@ class JsonlSessionStore:
|
|||||||
continue
|
continue
|
||||||
raw_data: object = json.loads(line)
|
raw_data: object = json.loads(line)
|
||||||
data = _json_object(raw_data)
|
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_value = cast(object, data.get("metadata", {}))
|
||||||
metadata = (
|
metadata = (
|
||||||
cast(dict[str, Any], metadata_value)
|
cast(dict[str, Any], metadata_value)
|
||||||
@ -745,6 +786,8 @@ class JsonlSessionStore:
|
|||||||
stored_key = (
|
stored_key = (
|
||||||
stored_key_value if isinstance(stored_key_value, str) else None
|
stored_key_value if isinstance(stored_key_value, str) else None
|
||||||
)
|
)
|
||||||
|
elif record_type == _PROVIDER_STATE_RECORD_TYPE:
|
||||||
|
continue
|
||||||
else:
|
else:
|
||||||
messages.append(data)
|
messages.append(data)
|
||||||
return {
|
return {
|
||||||
@ -837,6 +880,8 @@ class JsonlSessionStore:
|
|||||||
for line in f:
|
for line in f:
|
||||||
if not line.strip():
|
if not line.strip():
|
||||||
continue
|
continue
|
||||||
|
if _is_provider_state_record_line(line):
|
||||||
|
continue
|
||||||
scanned_records += 1
|
scanned_records += 1
|
||||||
scanned_chars += len(line)
|
scanned_chars += len(line)
|
||||||
if (
|
if (
|
||||||
@ -846,7 +891,10 @@ class JsonlSessionStore:
|
|||||||
break
|
break
|
||||||
raw_item: object = json.loads(line)
|
raw_item: object = json.loads(line)
|
||||||
item = _json_object(raw_item)
|
item = _json_object(raw_item)
|
||||||
if item.get("_type") == "metadata":
|
if item.get("_type") in {
|
||||||
|
"metadata",
|
||||||
|
_PROVIDER_STATE_RECORD_TYPE,
|
||||||
|
}:
|
||||||
continue
|
continue
|
||||||
text = _message_preview_text(item)
|
text = _message_preview_text(item)
|
||||||
if not text:
|
if not text:
|
||||||
|
|||||||
@ -18,10 +18,12 @@ from loguru import logger
|
|||||||
from nanobot.config.paths import get_webui_dir
|
from nanobot.config.paths import get_webui_dir
|
||||||
from nanobot.session.history_visibility import is_hidden_history_message
|
from nanobot.session.history_visibility import is_hidden_history_message
|
||||||
from nanobot.session.manager import (
|
from nanobot.session.manager import (
|
||||||
|
_PROVIDER_STATE_RECORD_TYPE, # pyright: ignore[reportPrivateUsage]
|
||||||
_SESSION_LIST_PREVIEW_MAX_CHARS, # pyright: ignore[reportPrivateUsage]
|
_SESSION_LIST_PREVIEW_MAX_CHARS, # pyright: ignore[reportPrivateUsage]
|
||||||
_SESSION_LIST_PREVIEW_MAX_RECORDS, # pyright: ignore[reportPrivateUsage]
|
_SESSION_LIST_PREVIEW_MAX_RECORDS, # pyright: ignore[reportPrivateUsage]
|
||||||
Session,
|
Session,
|
||||||
SessionManager,
|
SessionManager,
|
||||||
|
_is_provider_state_record_line, # pyright: ignore[reportPrivateUsage]
|
||||||
_message_preview_text, # pyright: ignore[reportPrivateUsage]
|
_message_preview_text, # pyright: ignore[reportPrivateUsage]
|
||||||
_metadata_title, # 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:
|
for line in f:
|
||||||
if not line.strip():
|
if not line.strip():
|
||||||
continue
|
continue
|
||||||
|
if _is_provider_state_record_line(line):
|
||||||
|
continue
|
||||||
item = json.loads(line)
|
item = json.loads(line)
|
||||||
|
if item.get("_type") == _PROVIDER_STATE_RECORD_TYPE:
|
||||||
|
continue
|
||||||
timestamp = _visible_message_timestamp(item)
|
timestamp = _visible_message_timestamp(item)
|
||||||
if timestamp is not None:
|
if timestamp is not None:
|
||||||
visible_message_at = _latest_updated_at(visible_message_at, timestamp)
|
visible_message_at = _latest_updated_at(visible_message_at, timestamp)
|
||||||
|
|||||||
@ -10,7 +10,11 @@ from nanobot.agent.memory import (
|
|||||||
Consolidator,
|
Consolidator,
|
||||||
MemoryStore,
|
MemoryStore,
|
||||||
)
|
)
|
||||||
from nanobot.providers.base import GenerationSettings, LLMResponse
|
from nanobot.providers.base import (
|
||||||
|
GenerationSettings,
|
||||||
|
LLMResponse,
|
||||||
|
ProviderConversationState,
|
||||||
|
)
|
||||||
from nanobot.runtime_context import (
|
from nanobot.runtime_context import (
|
||||||
RUNTIME_CONTEXT_HISTORY_META,
|
RUNTIME_CONTEXT_HISTORY_META,
|
||||||
RuntimeContextBlock,
|
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:
|
class TestConsolidatorSummarize:
|
||||||
async def test_archive_prompt_includes_media_breadcrumb(
|
async def test_archive_prompt_includes_media_breadcrumb(
|
||||||
self, consolidator, mock_provider, store, runtime
|
self, consolidator, mock_provider, store, runtime
|
||||||
@ -385,6 +399,7 @@ class TestConsolidatorTokenBudget:
|
|||||||
"""Old messages that cannot be replayed should be materialized first."""
|
"""Old messages that cannot be replayed should be materialized first."""
|
||||||
consolidator._SAFETY_BUFFER = 0
|
consolidator._SAFETY_BUFFER = 0
|
||||||
session = Session(key="test:replay-overflow")
|
session = Session(key="test:replay-overflow")
|
||||||
|
session.provider_state = _provider_state()
|
||||||
for i in range(10):
|
for i in range(10):
|
||||||
session.add_message("user", f"u{i}")
|
session.add_message("user", f"u{i}")
|
||||||
session.add_message("assistant", f"a{i}")
|
session.add_message("assistant", f"a{i}")
|
||||||
@ -404,6 +419,7 @@ class TestConsolidatorTokenBudget:
|
|||||||
assert archived_chunk[-1]["content"] == "a6"
|
assert archived_chunk[-1]["content"] == "a6"
|
||||||
assert session.last_consolidated == 14
|
assert session.last_consolidated == 14
|
||||||
assert session.metadata["_last_summary"]["text"] == "old conversation summary"
|
assert session.metadata["_last_summary"]["text"] == "old conversation summary"
|
||||||
|
assert session.provider_state is None
|
||||||
consolidator.sessions.save.assert_called()
|
consolidator.sessions.save.assert_called()
|
||||||
|
|
||||||
async def test_replay_window_overflow_extends_to_long_recent_user_turn(
|
async def test_replay_window_overflow_extends_to_long_recent_user_turn(
|
||||||
@ -479,6 +495,7 @@ class TestConsolidatorTokenBudget:
|
|||||||
session = MagicMock()
|
session = MagicMock()
|
||||||
session.last_consolidated = 0
|
session.last_consolidated = 0
|
||||||
session.key = "test:key"
|
session.key = "test:key"
|
||||||
|
session.provider_state = _provider_state()
|
||||||
session.messages = [
|
session.messages = [
|
||||||
{
|
{
|
||||||
"role": "user" if i in {0, 50, 61} else "assistant",
|
"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
|
# pick_consolidation_boundary returns (50, tokens) — user turn at idx 50
|
||||||
assert archived_chunk[0]["content"] == "m0"
|
assert archived_chunk[0]["content"] == "m0"
|
||||||
assert session.last_consolidated > 0
|
assert session.last_consolidated > 0
|
||||||
|
assert session.provider_state is None
|
||||||
|
|
||||||
async def test_raw_archive_fallback_advances_last_consolidated(
|
async def test_raw_archive_fallback_advances_last_consolidated(
|
||||||
self, consolidator, runtime
|
self, consolidator, runtime
|
||||||
@ -610,6 +628,7 @@ class TestCompactIdleSession:
|
|||||||
)
|
)
|
||||||
sessions = real_consolidator.sessions
|
sessions = real_consolidator.sessions
|
||||||
session = sessions.get_or_create("cli:test")
|
session = sessions.get_or_create("cli:test")
|
||||||
|
session.provider_state = _provider_state()
|
||||||
old_ts = session.updated_at
|
old_ts = session.updated_at
|
||||||
for i in range(20):
|
for i in range(20):
|
||||||
session.add_message("user", f"user msg {i}")
|
session.add_message("user", f"user msg {i}")
|
||||||
@ -627,6 +646,7 @@ class TestCompactIdleSession:
|
|||||||
assert len(reloaded.messages) == 40
|
assert len(reloaded.messages) == 40
|
||||||
assert reloaded.messages[0]["content"] == "user msg 0"
|
assert reloaded.messages[0]["content"] == "user msg 0"
|
||||||
assert reloaded.last_consolidated == 32
|
assert reloaded.last_consolidated == 32
|
||||||
|
assert reloaded.provider_state is None
|
||||||
visible = reloaded.get_history(max_messages=40)
|
visible = reloaded.get_history(max_messages=40)
|
||||||
assert len(visible) == 8
|
assert len(visible) == 8
|
||||||
assert visible[0]["content"] == "user msg 16"
|
assert visible[0]["content"] == "user msg 16"
|
||||||
|
|||||||
@ -452,6 +452,20 @@ class TestBuildMessages:
|
|||||||
assert "previous user message" in str(messages[1]["content"])
|
assert "previous user message" in str(messages[1]["content"])
|
||||||
assert "new 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):
|
def test_different_role_appended(self, tmp_path):
|
||||||
builder = _builder(tmp_path)
|
builder = _builder(tmp_path)
|
||||||
history = [{"role": "assistant", "content": "previous response"}]
|
history = [{"role": "assistant", "content": "previous response"}]
|
||||||
|
|||||||
@ -1,4 +1,5 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
|
import json
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
from unittest.mock import AsyncMock, MagicMock
|
from unittest.mock import AsyncMock, MagicMock
|
||||||
@ -19,7 +20,7 @@ from nanobot.bus.outbound_events import (
|
|||||||
)
|
)
|
||||||
from nanobot.bus.queue import MessageBus
|
from nanobot.bus.queue import MessageBus
|
||||||
from nanobot.cron.session_turns import CRON_HISTORY_META, CRON_TRIGGER_META
|
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.providers.factory import ProviderSnapshot
|
||||||
from nanobot.runtime_context import (
|
from nanobot.runtime_context import (
|
||||||
RUNTIME_CONTEXT_HISTORY_META,
|
RUNTIME_CONTEXT_HISTORY_META,
|
||||||
@ -59,6 +60,16 @@ def _mk_loop() -> AgentLoop:
|
|||||||
return loop
|
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:
|
def _runtime_message(content, blocks: list[RuntimeContextBlock]) -> dict:
|
||||||
merged, marker = append_runtime_context(content, blocks)
|
merged, marker = append_runtime_context(content, blocks)
|
||||||
assert marker is not None
|
assert marker is not None
|
||||||
@ -494,6 +505,7 @@ def test_restore_runtime_checkpoint_rehydrates_completed_and_pending_tools() ->
|
|||||||
loop = _mk_loop()
|
loop = _mk_loop()
|
||||||
session = Session(
|
session = Session(
|
||||||
key="test:checkpoint",
|
key="test:checkpoint",
|
||||||
|
provider_state=_provider_state(),
|
||||||
metadata={
|
metadata={
|
||||||
AgentLoop._RUNTIME_CHECKPOINT_KEY: {
|
AgentLoop._RUNTIME_CHECKPOINT_KEY: {
|
||||||
"assistant_message": {
|
"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[1]["tool_call_id"] == "call_done"
|
||||||
assert session.messages[2]["tool_call_id"] == "call_pending"
|
assert session.messages[2]["tool_call_id"] == "call_pending"
|
||||||
assert "interrupted before this tool finished" in session.messages[2]["content"].lower()
|
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:
|
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"
|
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
|
@pytest.mark.asyncio
|
||||||
async def test_process_message_persists_user_message_before_turn_completes(tmp_path: Path) -> None:
|
async def test_process_message_persists_user_message_before_turn_completes(tmp_path: Path) -> None:
|
||||||
loop = _make_full_loop(tmp_path)
|
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
|
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
|
@pytest.mark.asyncio
|
||||||
async def test_process_message_persists_unified_session_delivery_route(tmp_path: Path) -> None:
|
async def test_process_message_persists_unified_session_delivery_route(tmp_path: Path) -> None:
|
||||||
loop = _make_full_loop(tmp_path)
|
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 = loop.sessions.get_or_create("feishu:c3")
|
||||||
session.add_message("user", "old question")
|
session.add_message("user", "old question")
|
||||||
session.metadata[AgentLoop._PENDING_USER_TURN_KEY] = True
|
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.sessions.save(session)
|
||||||
|
|
||||||
loop._run_agent_loop = AsyncMock(return_value=(
|
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"},
|
{"role": "assistant", "content": "new answer"},
|
||||||
]
|
]
|
||||||
assert AgentLoop._PENDING_USER_TURN_KEY not in session.metadata
|
assert AgentLoop._PENDING_USER_TURN_KEY not in session.metadata
|
||||||
|
assert session.provider_state is None
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
|
|||||||
@ -11,7 +11,13 @@ import pytest
|
|||||||
|
|
||||||
from agent.runner_helpers import make_run_spec
|
from agent.runner_helpers import make_run_spec
|
||||||
from nanobot.config.schema import AgentDefaults
|
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
|
_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
|
@pytest.mark.asyncio
|
||||||
async def test_runner_returns_max_iterations_fallback():
|
async def test_runner_returns_max_iterations_fallback():
|
||||||
from nanobot.agent.runner import AgentRunner
|
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
|
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
|
@pytest.mark.asyncio
|
||||||
async def test_runner_uses_specific_message_after_empty_finalization_retry():
|
async def test_runner_uses_specific_message_after_empty_finalization_retry():
|
||||||
"""After silent retries + finalization all return empty, stop_reason is empty_final_response."""
|
"""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"
|
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
|
@pytest.mark.asyncio
|
||||||
async def test_runner_length_recovery_returns_all_segments():
|
async def test_runner_length_recovery_returns_all_segments():
|
||||||
"""Recovered output segments are returned together instead of only the tail."""
|
"""Recovered output segments are returned together instead of only the tail."""
|
||||||
|
|||||||
@ -9,8 +9,15 @@ import pytest
|
|||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
|
||||||
from nanobot.config.schema import ModelPresetConfig
|
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.fallback_provider import FallbackProvider
|
||||||
|
from nanobot.providers.openai_responses import resolve_compact_threshold
|
||||||
|
|
||||||
|
|
||||||
def _make_response(
|
def _make_response(
|
||||||
@ -66,6 +73,9 @@ class _FakeProvider(LLMProvider):
|
|||||||
self._response = response or _make_response()
|
self._response = response or _make_response()
|
||||||
self.chat_calls: list[dict[str, Any]] = []
|
self.chat_calls: list[dict[str, Any]] = []
|
||||||
self.chat_stream_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:
|
def get_default_model(self) -> str:
|
||||||
return f"{self.name}/model"
|
return f"{self.name}/model"
|
||||||
@ -81,6 +91,26 @@ class _FakeProvider(LLMProvider):
|
|||||||
await on_delta(self._response.content)
|
await on_delta(self._response.content)
|
||||||
return self._response
|
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 --
|
# -- config-level tests --
|
||||||
|
|
||||||
@ -211,6 +241,8 @@ def test_provider_snapshot_uses_smallest_fallback_context_window() -> None:
|
|||||||
snapshot = build_provider_snapshot(config)
|
snapshot = build_provider_snapshot(config)
|
||||||
|
|
||||||
assert snapshot.context_window_tokens == 64000
|
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:
|
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 primary.chat_calls[0]["model"] == "primary-model"
|
||||||
assert fallback.chat_calls[0]["model"] == "fallback-a"
|
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
|
@pytest.mark.asyncio
|
||||||
async def test_reports_the_fallback_model_before_its_request(self) -> None:
|
async def test_reports_the_fallback_model_before_its_request(self) -> None:
|
||||||
primary = _FakeProvider("primary", _error_response())
|
primary = _FakeProvider("primary", _error_response())
|
||||||
|
|||||||
@ -15,7 +15,11 @@ from nanobot.agent.context_governance import (
|
|||||||
)
|
)
|
||||||
from nanobot.agent.runner import AgentRunSpec
|
from nanobot.agent.runner import AgentRunSpec
|
||||||
from nanobot.config.schema import AgentDefaults
|
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
|
_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."""
|
"""LLM response tool_calls with a missing/empty name are dropped in place."""
|
||||||
from nanobot.agent.runner import AgentRunner
|
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(
|
response = LLMResponse(
|
||||||
content=None,
|
content=None,
|
||||||
tool_calls=[
|
tool_calls=[
|
||||||
@ -895,9 +906,11 @@ def test_drop_malformed_tool_calls_trims_response():
|
|||||||
ToolCallRequest(id="4", name="read_file", arguments={}),
|
ToolCallRequest(id="4", name="read_file", arguments={}),
|
||||||
],
|
],
|
||||||
finish_reason="tool_calls",
|
finish_reason="tool_calls",
|
||||||
|
provider_state=candidate_state,
|
||||||
)
|
)
|
||||||
dropped, all_dropped, orig = AgentRunner._drop_malformed_tool_calls(response)
|
dropped, all_dropped, orig = AgentRunner._drop_malformed_tool_calls(response)
|
||||||
assert [tc.name for tc in response.tool_calls] == ["read_file"]
|
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.finish_reason == "tool_calls"
|
||||||
assert response.should_execute_tools is True
|
assert response.should_execute_tools is True
|
||||||
assert dropped == 3
|
assert dropped == 3
|
||||||
|
|||||||
@ -4,6 +4,7 @@ import json
|
|||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
|
from nanobot.providers.base import ProviderConversationState
|
||||||
from nanobot.session.manager import Session, SessionManager
|
from nanobot.session.manager import Session, SessionManager
|
||||||
|
|
||||||
|
|
||||||
@ -101,6 +102,137 @@ class TestAtomicSave:
|
|||||||
for i in range(5):
|
for i in range(5):
|
||||||
assert loaded.messages[i]["content"] == f"msg{i}"
|
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:
|
class TestRepairCorruptFile:
|
||||||
def _write_corrupt_jsonl(self, path: Path, lines: list[str]) -> None:
|
def _write_corrupt_jsonl(self, path: Path, lines: list[str]) -> None:
|
||||||
|
|||||||
@ -1,3 +1,4 @@
|
|||||||
|
from nanobot.providers.base import ProviderConversationState
|
||||||
from nanobot.runtime_context import (
|
from nanobot.runtime_context import (
|
||||||
RUNTIME_CONTEXT_HISTORY_META,
|
RUNTIME_CONTEXT_HISTORY_META,
|
||||||
RuntimeContextBlock,
|
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():
|
def test_retain_recent_legal_suffix_returns_dropped_messages():
|
||||||
"""retain_recent_legal_suffix returns the actually-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):
|
for i in range(10):
|
||||||
session.messages.append({"role": "user", "content": f"msg{i}"})
|
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 [m["content"] for m in result.dropped] == [f"msg{i}" for i in range(6)]
|
||||||
assert len(session.messages) == 4
|
assert len(session.messages) == 4
|
||||||
assert result.already_consolidated_count == 0
|
assert result.already_consolidated_count == 0
|
||||||
|
assert session.provider_state is None
|
||||||
|
|
||||||
|
|
||||||
def test_retain_recent_legal_suffix_returns_empty_when_no_drop():
|
def test_retain_recent_legal_suffix_returns_empty_when_no_drop():
|
||||||
"""No messages dropped → empty list returned."""
|
"""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):
|
for i in range(3):
|
||||||
session.messages.append({"role": "user", "content": f"msg{i}"})
|
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.dropped == []
|
||||||
assert result.already_consolidated_count == 0
|
assert result.already_consolidated_count == 0
|
||||||
assert len(session.messages) == 3
|
assert len(session.messages) == 3
|
||||||
|
assert session.provider_state is state
|
||||||
|
|
||||||
|
|
||||||
def test_retain_recent_legal_suffix_returns_all_on_zero():
|
def test_retain_recent_legal_suffix_returns_all_on_zero():
|
||||||
|
|||||||
@ -504,6 +504,7 @@ async def test_drain_pending_blocks_while_subagents_running(tmp_path):
|
|||||||
usage={},
|
usage={},
|
||||||
had_injections=False,
|
had_injections=False,
|
||||||
tools_used=[],
|
tools_used=[],
|
||||||
|
provider_state=None,
|
||||||
)
|
)
|
||||||
|
|
||||||
loop.runner.run = AsyncMock(side_effect=fake_runner_run)
|
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={},
|
usage={},
|
||||||
had_injections=False,
|
had_injections=False,
|
||||||
tools_used=[],
|
tools_used=[],
|
||||||
|
provider_state=None,
|
||||||
)
|
)
|
||||||
|
|
||||||
loop.runner.run = AsyncMock(side_effect=fake_runner_run)
|
loop.runner.run = AsyncMock(side_effect=fake_runner_run)
|
||||||
@ -638,6 +640,7 @@ async def test_drain_pending_timeout(tmp_path):
|
|||||||
usage={},
|
usage={},
|
||||||
had_injections=False,
|
had_injections=False,
|
||||||
tools_used=[],
|
tools_used=[],
|
||||||
|
provider_state=None,
|
||||||
)
|
)
|
||||||
|
|
||||||
loop.runner.run = AsyncMock(side_effect=fake_runner_run)
|
loop.runner.run = AsyncMock(side_effect=fake_runner_run)
|
||||||
|
|||||||
@ -11,7 +11,7 @@ from nanobot.providers.azure_openai_provider import (
|
|||||||
AzureOpenAIProvider,
|
AzureOpenAIProvider,
|
||||||
_AzureTokenProvider,
|
_AzureTokenProvider,
|
||||||
)
|
)
|
||||||
from nanobot.providers.base import LLMResponse
|
from nanobot.providers.base import LLMResponse, ProviderCallContext
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# Init & validation
|
# Init & validation
|
||||||
@ -234,6 +234,7 @@ def test_build_body_basic():
|
|||||||
assert body["max_output_tokens"] == 4096
|
assert body["max_output_tokens"] == 4096
|
||||||
assert body["store"] is False
|
assert body["store"] is False
|
||||||
assert "reasoning" not in body
|
assert "reasoning" not in body
|
||||||
|
assert "include" not in body
|
||||||
# input should contain the converted user message only (system extracted)
|
# input should contain the converted user message only (system extracted)
|
||||||
assert any(
|
assert any(
|
||||||
item.get("role") == "user"
|
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():
|
def test_build_body_max_tokens_minimum():
|
||||||
"""max_output_tokens should never be less than 1."""
|
"""max_output_tokens should never be less than 1."""
|
||||||
provider = AzureOpenAIProvider(api_key="k", api_base="https://r.com", default_model="gpt-4o")
|
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
|
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
|
@pytest.mark.asyncio
|
||||||
async def test_chat_uses_default_model():
|
async def test_chat_uses_default_model():
|
||||||
provider = AzureOpenAIProvider(
|
provider = AzureOpenAIProvider(
|
||||||
@ -411,6 +468,7 @@ async def test_chat_with_tool_calls():
|
|||||||
assert len(result.tool_calls) == 1
|
assert len(result.tool_calls) == 1
|
||||||
assert result.tool_calls[0].name == "get_weather"
|
assert result.tool_calls[0].name == "get_weather"
|
||||||
assert result.tool_calls[0].arguments == {"location": "SF"}
|
assert result.tool_calls[0].arguments == {"location": "SF"}
|
||||||
|
assert result.provider_state is not None
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@ -510,6 +568,7 @@ async def test_chat_stream_with_tool_calls():
|
|||||||
item_done.name = "get_weather"
|
item_done.name = "get_weather"
|
||||||
ev_item_done = MagicMock(type="response.output_item.done", item=item_done)
|
ev_item_done = MagicMock(type="response.output_item.done", item=item_done)
|
||||||
resp_obj = MagicMock(status="completed")
|
resp_obj = MagicMock(status="completed")
|
||||||
|
resp_obj.model_dump.return_value = {"status": "completed", "output": []}
|
||||||
ev_completed = MagicMock(type="response.completed", response=resp_obj)
|
ev_completed = MagicMock(type="response.completed", response=resp_obj)
|
||||||
|
|
||||||
async def mock_stream():
|
async def mock_stream():
|
||||||
@ -527,6 +586,7 @@ async def test_chat_stream_with_tool_calls():
|
|||||||
assert len(result.tool_calls) == 1
|
assert len(result.tool_calls) == 1
|
||||||
assert result.tool_calls[0].name == "get_weather"
|
assert result.tool_calls[0].name == "get_weather"
|
||||||
assert result.tool_calls[0].arguments == {"location": "SF"}
|
assert result.tool_calls[0].arguments == {"location": "SF"}
|
||||||
|
assert result.provider_state is not None
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
|
|||||||
291
tests/providers/test_conversation_state.py
Normal file
291
tests/providers/test_conversation_state.py
Normal file
@ -0,0 +1,291 @@
|
|||||||
|
"""Tests for provider-owned conversation-state lifecycle coordination."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from nanobot.providers.base import (
|
||||||
|
LLMProvider,
|
||||||
|
LLMResponse,
|
||||||
|
ProviderConversationState,
|
||||||
|
ToolCallRequest,
|
||||||
|
)
|
||||||
|
from nanobot.providers.conversation_state import (
|
||||||
|
ProviderConversationStateController,
|
||||||
|
allows_conversation_message_merge,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _provider(*, resumable: bool = True, compact: bool = False) -> MagicMock:
|
||||||
|
provider = MagicMock(spec=LLMProvider)
|
||||||
|
provider.can_resume_conversation_state.return_value = resumable
|
||||||
|
provider.supports_native_compaction.return_value = compact
|
||||||
|
return provider
|
||||||
|
|
||||||
|
|
||||||
|
def _state(label: str, *, pending: list[dict] | None = None) -> ProviderConversationState:
|
||||||
|
return ProviderConversationState(
|
||||||
|
kind="openai_responses",
|
||||||
|
provider="openai:test",
|
||||||
|
model="gpt-5.6",
|
||||||
|
version=1,
|
||||||
|
payload={"items": [{"type": "reasoning", "encrypted_content": label}]},
|
||||||
|
pending_messages=pending or [],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_controller_replays_only_messages_after_provider_output() -> None:
|
||||||
|
provider = _provider()
|
||||||
|
messages = [
|
||||||
|
{"role": "system", "content": "system"},
|
||||||
|
{"role": "user", "content": "run a tool"},
|
||||||
|
]
|
||||||
|
controller = ProviderConversationStateController(
|
||||||
|
provider=provider,
|
||||||
|
model="gpt-5.6",
|
||||||
|
messages=messages,
|
||||||
|
)
|
||||||
|
state = _state("first")
|
||||||
|
|
||||||
|
controller.prepare_request(messages, context_window_tokens=200_000)
|
||||||
|
response = LLMResponse(content=None, provider_state=state)
|
||||||
|
controller.observe_response(response, messages)
|
||||||
|
assert allows_conversation_message_merge(messages[-1]) is False
|
||||||
|
|
||||||
|
messages.append(controller.project_response_message(
|
||||||
|
{
|
||||||
|
"role": "assistant",
|
||||||
|
"content": None,
|
||||||
|
"tool_calls": [{"id": "call_1", "type": "function"}],
|
||||||
|
},
|
||||||
|
response,
|
||||||
|
))
|
||||||
|
tool_message = {
|
||||||
|
"role": "tool",
|
||||||
|
"tool_call_id": "call_1",
|
||||||
|
"content": "tool result",
|
||||||
|
}
|
||||||
|
messages.append(tool_message)
|
||||||
|
|
||||||
|
provider_context = controller.prepare_request(
|
||||||
|
messages,
|
||||||
|
context_window_tokens=200_000,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert provider_context is not None
|
||||||
|
assert provider_context.conversation_state is not None
|
||||||
|
assert provider_context.conversation_state.payload == state.payload
|
||||||
|
assert provider_context.conversation_state.pending_messages == [tool_message]
|
||||||
|
assert controller.checkpoint(messages).pending_messages == [tool_message]
|
||||||
|
|
||||||
|
|
||||||
|
def test_controller_uses_governed_messages_for_provider_state_delta() -> None:
|
||||||
|
provider = _provider()
|
||||||
|
messages = [
|
||||||
|
{"role": "user", "content": "run a tool"},
|
||||||
|
]
|
||||||
|
controller = ProviderConversationStateController(
|
||||||
|
provider=provider,
|
||||||
|
model="gpt-5.6",
|
||||||
|
messages=messages,
|
||||||
|
)
|
||||||
|
state = _state("first")
|
||||||
|
|
||||||
|
controller.prepare_request(messages, context_window_tokens=200_000)
|
||||||
|
response = LLMResponse(content=None, provider_state=state)
|
||||||
|
controller.observe_response(response, messages)
|
||||||
|
messages.extend([
|
||||||
|
controller.project_response_message(
|
||||||
|
{
|
||||||
|
"role": "assistant",
|
||||||
|
"content": None,
|
||||||
|
"tool_calls": [{"id": "call_1", "type": "function"}],
|
||||||
|
},
|
||||||
|
response,
|
||||||
|
),
|
||||||
|
{
|
||||||
|
"role": "tool",
|
||||||
|
"tool_call_id": "call_1",
|
||||||
|
"content": "raw oversized result",
|
||||||
|
},
|
||||||
|
])
|
||||||
|
governed_messages = [
|
||||||
|
messages[0],
|
||||||
|
messages[1],
|
||||||
|
{
|
||||||
|
"role": "tool",
|
||||||
|
"tool_call_id": "call_1",
|
||||||
|
"content": "compacted result",
|
||||||
|
},
|
||||||
|
]
|
||||||
|
|
||||||
|
provider_context = controller.prepare_request(
|
||||||
|
messages,
|
||||||
|
context_window_tokens=200_000,
|
||||||
|
model_messages=governed_messages,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert provider_context is not None
|
||||||
|
assert provider_context.conversation_state is not None
|
||||||
|
assert provider_context.conversation_state.pending_messages == [{
|
||||||
|
"role": "tool",
|
||||||
|
"tool_call_id": "call_1",
|
||||||
|
"content": "compacted result",
|
||||||
|
}]
|
||||||
|
assert controller.checkpoint(messages).pending_messages[-1]["content"] == (
|
||||||
|
"raw oversized result"
|
||||||
|
)
|
||||||
|
governed_checkpoint = controller.checkpoint(
|
||||||
|
messages,
|
||||||
|
model_messages=governed_messages,
|
||||||
|
)
|
||||||
|
assert governed_checkpoint is not None
|
||||||
|
assert governed_checkpoint.pending_messages[-1]["content"] == "compacted result"
|
||||||
|
|
||||||
|
|
||||||
|
def test_transient_response_preserves_only_durable_request_messages() -> None:
|
||||||
|
provider = _provider()
|
||||||
|
current_message = {"role": "user", "content": "continue"}
|
||||||
|
supplemental = {"role": "user", "content": "internal finalization retry"}
|
||||||
|
messages = [{"role": "system", "content": "system"}, current_message]
|
||||||
|
controller = ProviderConversationStateController(
|
||||||
|
provider=provider,
|
||||||
|
model="gpt-5.6",
|
||||||
|
messages=messages,
|
||||||
|
state=_state("saved", pending=[
|
||||||
|
{"role": "tool", "content": "prior"},
|
||||||
|
current_message,
|
||||||
|
]),
|
||||||
|
)
|
||||||
|
|
||||||
|
provider_context = controller.prepare_request(
|
||||||
|
messages,
|
||||||
|
context_window_tokens=200_000,
|
||||||
|
supplemental_messages=[supplemental],
|
||||||
|
)
|
||||||
|
assert provider_context is not None
|
||||||
|
assert provider_context.conversation_state is not None
|
||||||
|
assert provider_context.conversation_state.pending_messages == [
|
||||||
|
{"role": "tool", "content": "prior"},
|
||||||
|
current_message,
|
||||||
|
supplemental,
|
||||||
|
]
|
||||||
|
|
||||||
|
controller.observe_response(
|
||||||
|
LLMResponse(
|
||||||
|
content="temporary failure",
|
||||||
|
finish_reason="error",
|
||||||
|
error_kind="timeout",
|
||||||
|
),
|
||||||
|
messages,
|
||||||
|
)
|
||||||
|
placeholder = {"role": "assistant", "content": "model error"}
|
||||||
|
messages.append(placeholder)
|
||||||
|
|
||||||
|
state = controller.finish(messages)
|
||||||
|
assert state is not None
|
||||||
|
assert state.pending_messages == [
|
||||||
|
{"role": "tool", "content": "prior"},
|
||||||
|
current_message,
|
||||||
|
placeholder,
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def test_non_retryable_response_discards_saved_state() -> None:
|
||||||
|
provider = _provider()
|
||||||
|
messages = [{"role": "user", "content": "continue"}]
|
||||||
|
controller = ProviderConversationStateController(
|
||||||
|
provider=provider,
|
||||||
|
model="gpt-5.6",
|
||||||
|
messages=messages,
|
||||||
|
state=_state("saved"),
|
||||||
|
)
|
||||||
|
|
||||||
|
controller.prepare_request(messages, context_window_tokens=200_000)
|
||||||
|
controller.observe_response(
|
||||||
|
LLMResponse(
|
||||||
|
content="invalid request",
|
||||||
|
finish_reason="error",
|
||||||
|
error_status_code=400,
|
||||||
|
error_should_retry=False,
|
||||||
|
),
|
||||||
|
messages,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert controller.finish(messages) is None
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
("finish_reason", "exposes_tool_call"),
|
||||||
|
[
|
||||||
|
("length", False),
|
||||||
|
("length", True),
|
||||||
|
("refusal", True),
|
||||||
|
("content_filter", True),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_terminal_response_discards_candidate_state(
|
||||||
|
finish_reason: str,
|
||||||
|
exposes_tool_call: bool,
|
||||||
|
) -> None:
|
||||||
|
provider = _provider()
|
||||||
|
messages = [{"role": "user", "content": "continue"}]
|
||||||
|
controller = ProviderConversationStateController(
|
||||||
|
provider=provider,
|
||||||
|
model="gpt-5.6",
|
||||||
|
messages=messages,
|
||||||
|
state=_state("saved"),
|
||||||
|
)
|
||||||
|
|
||||||
|
controller.prepare_request(messages, context_window_tokens=200_000)
|
||||||
|
candidate = ProviderConversationState(
|
||||||
|
kind="openai_responses",
|
||||||
|
provider="openai:test",
|
||||||
|
model="gpt-5.6",
|
||||||
|
version=1,
|
||||||
|
payload={
|
||||||
|
"items": [{
|
||||||
|
"type": "function_call",
|
||||||
|
"call_id": "call_1",
|
||||||
|
"name": "exec",
|
||||||
|
"arguments": "{}",
|
||||||
|
}],
|
||||||
|
},
|
||||||
|
)
|
||||||
|
response = LLMResponse(
|
||||||
|
content="terminal response",
|
||||||
|
tool_calls=(
|
||||||
|
[ToolCallRequest(id="call_1", name="exec", arguments={})]
|
||||||
|
if exposes_tool_call
|
||||||
|
else []
|
||||||
|
),
|
||||||
|
finish_reason=finish_reason,
|
||||||
|
provider_state=candidate,
|
||||||
|
)
|
||||||
|
assert response.has_tool_calls is exposes_tool_call
|
||||||
|
assert response.should_execute_tools is False
|
||||||
|
|
||||||
|
controller.observe_response(response, messages)
|
||||||
|
|
||||||
|
assert controller.finish(messages) is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_independent_request_exposes_context_without_capability_check() -> None:
|
||||||
|
provider = _provider(compact=False)
|
||||||
|
messages = [{"role": "user", "content": "hello"}]
|
||||||
|
controller = ProviderConversationStateController(
|
||||||
|
provider=provider,
|
||||||
|
model="gpt-5.6",
|
||||||
|
messages=messages,
|
||||||
|
state=_state("saved"),
|
||||||
|
)
|
||||||
|
|
||||||
|
provider_context = controller.independent_request_context(
|
||||||
|
context_window_tokens=200_000,
|
||||||
|
)
|
||||||
|
assert provider_context is not None
|
||||||
|
assert provider_context.conversation_state is None
|
||||||
|
assert provider_context.context_window_tokens == 200_000
|
||||||
|
provider.supports_native_compaction.assert_not_called()
|
||||||
@ -10,6 +10,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
|||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
|
from nanobot.providers.base import ProviderCallContext
|
||||||
from nanobot.providers.openai_compat_provider import OpenAICompatProvider
|
from nanobot.providers.openai_compat_provider import OpenAICompatProvider
|
||||||
from nanobot.providers.registry import find_by_name
|
from nanobot.providers.registry import find_by_name
|
||||||
|
|
||||||
@ -44,8 +45,10 @@ def test_build_responses_body_strips_github_copilot_prefix():
|
|||||||
temperature=0.1,
|
temperature=0.1,
|
||||||
reasoning_effort=None,
|
reasoning_effort=None,
|
||||||
tool_choice=None,
|
tool_choice=None,
|
||||||
|
provider_context=ProviderCallContext(context_window_tokens=128_000),
|
||||||
)
|
)
|
||||||
assert body["model"] == "gpt-5.4-mini"
|
assert body["model"] == "gpt-5.4-mini"
|
||||||
|
assert "context_management" not in body
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
|
|||||||
@ -14,6 +14,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
|||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
|
from nanobot.providers.base import ProviderCallContext
|
||||||
from nanobot.providers.openai_compat_provider import OpenAICompatProvider
|
from nanobot.providers.openai_compat_provider import OpenAICompatProvider
|
||||||
from nanobot.providers.registry import find_by_name
|
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 call_kwargs["max_output_tokens"] == 4096
|
||||||
assert "input" in call_kwargs
|
assert "input" in call_kwargs
|
||||||
assert "messages" not in call_kwargs
|
assert "messages" not in call_kwargs
|
||||||
|
assert call_kwargs["include"] == ["reasoning.encrypted_content"]
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@ -710,6 +712,40 @@ async def test_direct_openai_reasoning_prefers_responses_api() -> None:
|
|||||||
assert call_kwargs["include"] == ["reasoning.encrypted_content"]
|
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
|
@pytest.mark.asyncio
|
||||||
async def test_direct_openai_gpt4o_stays_on_chat_completions() -> None:
|
async def test_direct_openai_gpt4o_stays_on_chat_completions() -> None:
|
||||||
mock_chat = AsyncMock(return_value=_fake_chat_response())
|
mock_chat = AsyncMock(return_value=_fake_chat_response())
|
||||||
|
|||||||
@ -20,6 +20,7 @@ from nanobot.providers.openai_codex_provider import (
|
|||||||
_request_codex,
|
_request_codex,
|
||||||
_should_retry_status,
|
_should_retry_status,
|
||||||
)
|
)
|
||||||
|
from nanobot.providers.openai_responses import build_responses_state
|
||||||
from nanobot.providers.registry import find_by_name
|
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
|
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
|
@pytest.mark.asyncio
|
||||||
async def test_codex_request_honors_stream_idle_timeout_env(monkeypatch) -> None:
|
async def test_codex_request_honors_stream_idle_timeout_env(monkeypatch) -> None:
|
||||||
"""NANOBOT_STREAM_IDLE_TIMEOUT_S overrides the default Codex stream timeout."""
|
"""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
|
_ = proxy, on_thinking_delta, on_tool_call_delta
|
||||||
bodies.append(body)
|
bodies.append(body)
|
||||||
return "ok", [], "stop", {}, None
|
return provider_base.LLMResponse(content="ok")
|
||||||
|
|
||||||
monkeypatch.setattr("nanobot.providers.openai_codex_provider._request_codex", fake_request)
|
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):
|
async def fake_request(_url, _headers, body, **_kwargs):
|
||||||
bodies.append(body)
|
bodies.append(body)
|
||||||
return "ok", [], "stop", {}, None
|
return provider_base.LLMResponse(content="ok")
|
||||||
|
|
||||||
monkeypatch.setattr("nanobot.providers.openai_codex_provider._request_codex", fake_request)
|
monkeypatch.setattr("nanobot.providers.openai_codex_provider._request_codex", fake_request)
|
||||||
config = Config.model_validate({
|
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
|
_ = url, headers, body, verify, on_content_delta, on_thinking_delta, on_tool_call_delta
|
||||||
seen["request_proxy"] = proxy
|
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.get_codex_token", fake_token)
|
||||||
monkeypatch.setattr("nanobot.providers.openai_codex_provider._request_codex", fake_request)
|
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
|
calls += 1
|
||||||
if calls == 1:
|
if calls == 1:
|
||||||
raise httpx.ReadTimeout("")
|
raise httpx.ReadTimeout("")
|
||||||
return "ok", [], "stop", {}, None
|
return provider_base.LLMResponse(content="ok")
|
||||||
|
|
||||||
async def fake_sleep(delay: float) -> None:
|
async def fake_sleep(delay: float) -> None:
|
||||||
delays.append(delay)
|
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"}
|
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
|
@pytest.mark.asyncio
|
||||||
async def test_codex_stream_surfaces_reasoning_summary(monkeypatch) -> None:
|
async def test_codex_stream_surfaces_reasoning_summary(monkeypatch) -> None:
|
||||||
def fake_token(**_kwargs):
|
def fake_token(**_kwargs):
|
||||||
@ -559,7 +850,12 @@ async def test_codex_stream_surfaces_reasoning_summary(monkeypatch) -> None:
|
|||||||
await on_content_delta("answer")
|
await on_content_delta("answer")
|
||||||
if on_thinking_delta:
|
if on_thinking_delta:
|
||||||
await on_thinking_delta("summary")
|
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)
|
monkeypatch.setattr("nanobot.providers.openai_codex_provider._request_codex", fake_request)
|
||||||
|
|
||||||
|
|||||||
@ -1,9 +1,11 @@
|
|||||||
"""Tests for the shared openai_responses converters and parsers."""
|
"""Tests for the shared openai_responses converters and parsers."""
|
||||||
|
|
||||||
import json
|
import json
|
||||||
|
from io import StringIO
|
||||||
from unittest.mock import MagicMock, patch
|
from unittest.mock import MagicMock, patch
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
from loguru import logger
|
||||||
|
|
||||||
from nanobot.providers.openai_responses.converters import (
|
from nanobot.providers.openai_responses.converters import (
|
||||||
convert_messages,
|
convert_messages,
|
||||||
@ -12,12 +14,22 @@ from nanobot.providers.openai_responses.converters import (
|
|||||||
split_tool_call_id,
|
split_tool_call_id,
|
||||||
)
|
)
|
||||||
from nanobot.providers.openai_responses.parsing import (
|
from nanobot.providers.openai_responses.parsing import (
|
||||||
|
ResponsesStreamCapture,
|
||||||
consume_sdk_stream,
|
consume_sdk_stream,
|
||||||
consume_sse,
|
consume_sse,
|
||||||
consume_sse_with_reasoning,
|
consume_sse_with_reasoning,
|
||||||
|
is_replayable_finish_reason,
|
||||||
map_finish_reason,
|
map_finish_reason,
|
||||||
parse_response_output,
|
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
|
# converters - split_tool_call_id
|
||||||
@ -398,6 +410,17 @@ class TestMapFinishReason:
|
|||||||
def test_unknown_defaults_to_stop(self):
|
def test_unknown_defaults_to_stop(self):
|
||||||
assert map_finish_reason("some_new_status") == "stop"
|
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
|
# parsing - parse_response_output
|
||||||
@ -418,6 +441,29 @@ class TestParseResponseOutput:
|
|||||||
assert result.usage == {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}
|
assert result.usage == {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}
|
||||||
assert result.tool_calls == []
|
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):
|
def test_tool_call_response(self):
|
||||||
resp = {
|
resp = {
|
||||||
"output": [{
|
"output": [{
|
||||||
@ -429,12 +475,18 @@ class TestParseResponseOutput:
|
|||||||
"status": "completed",
|
"status": "completed",
|
||||||
"usage": {},
|
"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 result.content is None
|
||||||
assert len(result.tool_calls) == 1
|
assert len(result.tool_calls) == 1
|
||||||
assert result.tool_calls[0].name == "get_weather"
|
assert result.tool_calls[0].name == "get_weather"
|
||||||
assert result.tool_calls[0].arguments == {"city": "SF"}
|
assert result.tool_calls[0].arguments == {"city": "SF"}
|
||||||
assert result.tool_calls[0].id == "call_1|fc_1"
|
assert result.tool_calls[0].id == "call_1|fc_1"
|
||||||
|
assert result.provider_state is not None
|
||||||
|
|
||||||
def test_malformed_tool_arguments_logged(self):
|
def test_malformed_tool_arguments_logged(self):
|
||||||
"""Malformed JSON arguments should log a warning and remain non-object."""
|
"""Malformed JSON arguments should log a warning and remain non-object."""
|
||||||
@ -493,10 +545,39 @@ class TestParseResponseOutput:
|
|||||||
assert result.content is None
|
assert result.content is None
|
||||||
assert result.tool_calls == []
|
assert result.tool_calls == []
|
||||||
|
|
||||||
def test_incomplete_status(self):
|
@pytest.mark.parametrize(
|
||||||
resp = {"output": [], "status": "incomplete", "usage": {}}
|
("reason", "expected_finish_reason"),
|
||||||
result = parse_response_output(resp)
|
[
|
||||||
assert result.finish_reason == "length"
|
("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):
|
def test_sdk_model_object(self):
|
||||||
"""parse_response_output should handle SDK objects with model_dump()."""
|
"""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["completion_tokens"] == 50
|
||||||
assert result.usage["total_tokens"] == 150
|
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
|
# parsing - consume_sse
|
||||||
@ -553,6 +822,122 @@ class TestConsumeSse:
|
|||||||
assert tool_calls == []
|
assert tool_calls == []
|
||||||
assert finish_reason == "stop"
|
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
|
@pytest.mark.asyncio
|
||||||
async def test_reasoning_summary_delta_extracted(self):
|
async def test_reasoning_summary_delta_extracted(self):
|
||||||
response = _SseResponse([
|
response = _SseResponse([
|
||||||
@ -599,6 +984,139 @@ class TestConsumeSse:
|
|||||||
|
|
||||||
assert reasoning == "cached summary"
|
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
|
@pytest.mark.asyncio
|
||||||
async def test_reasoning_summary_from_done_item(self):
|
async def test_reasoning_summary_from_done_item(self):
|
||||||
response = _SseResponse([
|
response = _SseResponse([
|
||||||
@ -755,6 +1273,131 @@ class TestConsumeSdkStream:
|
|||||||
assert tool_calls == []
|
assert tool_calls == []
|
||||||
assert finish_reason == "stop"
|
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
|
@pytest.mark.asyncio
|
||||||
async def test_on_content_delta_called(self):
|
async def test_on_content_delta_called(self):
|
||||||
ev1 = MagicMock(type="response.output_text.delta", delta="hi")
|
ev1 = MagicMock(type="response.output_text.delta", delta="hi")
|
||||||
@ -919,6 +1562,64 @@ class TestConsumeSdkStream:
|
|||||||
_, _, _, usage, _ = await consume_sdk_stream(stream())
|
_, _, _, usage, _ = await consume_sdk_stream(stream())
|
||||||
assert usage == {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}
|
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
|
@pytest.mark.asyncio
|
||||||
async def test_reasoning_extracted(self):
|
async def test_reasoning_extracted(self):
|
||||||
summary_item = MagicMock(type="summary_text", text="thinking...")
|
summary_item = MagicMock(type="summary_text", text="thinking...")
|
||||||
|
|||||||
@ -3,7 +3,14 @@ import copy
|
|||||||
|
|
||||||
import pytest
|
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):
|
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)
|
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
|
@pytest.mark.asyncio
|
||||||
async def test_non_transient_error_without_images_no_retry() -> None:
|
async def test_non_transient_error_without_images_no_retry() -> None:
|
||||||
"""Non-transient errors without image content are returned immediately."""
|
"""Non-transient errors without image content are returned immediately."""
|
||||||
|
|||||||
@ -4,6 +4,7 @@ import time
|
|||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
|
from nanobot.providers.base import ProviderCallContext
|
||||||
from nanobot.providers.openai_compat_provider import (
|
from nanobot.providers.openai_compat_provider import (
|
||||||
_RESPONSES_FAILURE_THRESHOLD,
|
_RESPONSES_FAILURE_THRESHOLD,
|
||||||
_RESPONSES_PROBE_INTERVAL_S,
|
_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
|
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):
|
def test_api_type_chat_completions_disables_responses(provider):
|
||||||
provider._api_type = "chat_completions"
|
provider._api_type = "chat_completions"
|
||||||
assert provider._should_use_responses_api("gpt-5", None) is False
|
assert provider._should_use_responses_api("gpt-5", None) is False
|
||||||
|
|||||||
@ -8,6 +8,7 @@ import pytest
|
|||||||
|
|
||||||
import nanobot.webui.session_list_index as session_list_index
|
import nanobot.webui.session_list_index as session_list_index
|
||||||
from nanobot.cron.session_turns import CRON_HISTORY_META
|
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.automation_turns import AUTOMATION_HISTORY_META
|
||||||
from nanobot.session.history_visibility import HIDDEN_HISTORY_META
|
from nanobot.session.history_visibility import HIDDEN_HISTORY_META
|
||||||
from nanobot.session.manager import SessionManager
|
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"}
|
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:
|
def test_webui_session_list_drops_deleted_index_rows(tmp_path: Path) -> None:
|
||||||
manager = SessionManager(tmp_path)
|
manager = SessionManager(tmp_path)
|
||||||
session = manager.get_or_create("websocket:deleted")
|
session = manager.get_or_create("websocket:deleted")
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user