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

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