mirror of
https://github.com/HKUDS/nanobot.git
synced 2026-08-04 08:28:36 +00:00
refactor(agent): capture original user text per turn
This commit is contained in:
parent
55317094b6
commit
bb3b449e09
@ -113,6 +113,7 @@ class TurnContext:
|
|||||||
session_key: str
|
session_key: str
|
||||||
state: TurnState
|
state: TurnState
|
||||||
turn_id: str
|
turn_id: str
|
||||||
|
original_user_text: str | None = None
|
||||||
session: Session | None = None
|
session: Session | None = None
|
||||||
|
|
||||||
history: list[dict[str, Any]] = field(default_factory=list)
|
history: list[dict[str, Any]] = field(default_factory=list)
|
||||||
@ -739,6 +740,7 @@ class AgentLoop:
|
|||||||
message_id: str | None = None,
|
message_id: str | None = None,
|
||||||
metadata: dict[str, Any] | None = None,
|
metadata: dict[str, Any] | None = None,
|
||||||
session_key: str | None = None,
|
session_key: str | None = None,
|
||||||
|
original_user_text: str | None = None,
|
||||||
pending_queue: asyncio.Queue | None = None,
|
pending_queue: asyncio.Queue | None = None,
|
||||||
ephemeral: bool = False,
|
ephemeral: bool = False,
|
||||||
run_extra_hooks_for_ephemeral: bool = False,
|
run_extra_hooks_for_ephemeral: bool = False,
|
||||||
@ -858,6 +860,7 @@ class AgentLoop:
|
|||||||
chat_id=chat_id,
|
chat_id=chat_id,
|
||||||
message_id=message_id,
|
message_id=message_id,
|
||||||
session_key=active_session_key,
|
session_key=active_session_key,
|
||||||
|
original_user_text=original_user_text,
|
||||||
metadata=dict(metadata or {}),
|
metadata=dict(metadata or {}),
|
||||||
)
|
)
|
||||||
file_state_token = bind_file_states(self._file_state_store.for_session(active_session_key))
|
file_state_token = bind_file_states(self._file_state_store.for_session(active_session_key))
|
||||||
@ -1283,6 +1286,7 @@ class AgentLoop:
|
|||||||
message_id=msg.metadata.get("message_id"),
|
message_id=msg.metadata.get("message_id"),
|
||||||
metadata=msg.metadata,
|
metadata=msg.metadata,
|
||||||
session_key=key,
|
session_key=key,
|
||||||
|
original_user_text=None,
|
||||||
pending_queue=pending_queue,
|
pending_queue=pending_queue,
|
||||||
hook_factories=hook_factories,
|
hook_factories=hook_factories,
|
||||||
)
|
)
|
||||||
@ -1350,6 +1354,11 @@ class AgentLoop:
|
|||||||
session_key=key,
|
session_key=key,
|
||||||
state=TurnState.RESTORE,
|
state=TurnState.RESTORE,
|
||||||
turn_id=f"{key}:{time.time_ns()}",
|
turn_id=f"{key}:{time.time_ns()}",
|
||||||
|
original_user_text=(
|
||||||
|
None
|
||||||
|
if turn_continuation.internal_continuation_inbound(msg.metadata)
|
||||||
|
else msg.content
|
||||||
|
),
|
||||||
turn_wall_started_at=t0,
|
turn_wall_started_at=t0,
|
||||||
visible_run_started_at=turn_continuation.internal_continuation_run_started_at(
|
visible_run_started_at=turn_continuation.internal_continuation_run_started_at(
|
||||||
msg.metadata,
|
msg.metadata,
|
||||||
@ -1587,6 +1596,7 @@ class AgentLoop:
|
|||||||
message_id=ctx.msg.metadata.get("message_id"),
|
message_id=ctx.msg.metadata.get("message_id"),
|
||||||
metadata=ctx.msg.metadata,
|
metadata=ctx.msg.metadata,
|
||||||
session_key=ctx.session_key,
|
session_key=ctx.session_key,
|
||||||
|
original_user_text=ctx.original_user_text,
|
||||||
pending_queue=ctx.pending_queue,
|
pending_queue=ctx.pending_queue,
|
||||||
ephemeral=ctx.ephemeral,
|
ephemeral=ctx.ephemeral,
|
||||||
run_extra_hooks_for_ephemeral=ctx.run_extra_hooks_for_ephemeral,
|
run_extra_hooks_for_ephemeral=ctx.run_extra_hooks_for_ephemeral,
|
||||||
|
|||||||
@ -18,6 +18,7 @@ class RequestContext:
|
|||||||
chat_id: str
|
chat_id: str
|
||||||
message_id: str | None = None
|
message_id: str | None = None
|
||||||
session_key: str | None = None
|
session_key: str | None = None
|
||||||
|
original_user_text: str | None = None
|
||||||
metadata: dict[str, Any] = field(default_factory=dict)
|
metadata: dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@ -11,8 +11,10 @@ from nanobot.agent.tools.context import (
|
|||||||
current_request_context,
|
current_request_context,
|
||||||
reset_request_context,
|
reset_request_context,
|
||||||
)
|
)
|
||||||
|
from nanobot.bus.events import InboundMessage
|
||||||
from nanobot.bus.queue import MessageBus
|
from nanobot.bus.queue import MessageBus
|
||||||
from nanobot.providers.base import LLMResponse, ToolCallRequest
|
from nanobot.providers.base import LLMResponse, ToolCallRequest
|
||||||
|
from nanobot.session.turn_continuation import INTERNAL_CONTINUATION_META
|
||||||
|
|
||||||
|
|
||||||
class _ContextRecordingTool:
|
class _ContextRecordingTool:
|
||||||
@ -165,6 +167,7 @@ async def test_agent_loop_restores_outer_request_context_after_runner_exception(
|
|||||||
assert current.channel == "slack"
|
assert current.channel == "slack"
|
||||||
assert current.chat_id == "C123"
|
assert current.chat_id == "C123"
|
||||||
assert current.session_key == "slack:C123:111.222"
|
assert current.session_key == "slack:C123:111.222"
|
||||||
|
assert current.original_user_text == " unchanged user text "
|
||||||
raise RuntimeError("runner failed")
|
raise RuntimeError("runner failed")
|
||||||
|
|
||||||
loop.runner.run = AsyncMock(side_effect=fail_run)
|
loop.runner.run = AsyncMock(side_effect=fail_run)
|
||||||
@ -176,9 +179,53 @@ async def test_agent_loop_restores_outer_request_context_after_runner_exception(
|
|||||||
channel="slack",
|
channel="slack",
|
||||||
chat_id="C123",
|
chat_id="C123",
|
||||||
session_key="slack:C123:111.222",
|
session_key="slack:C123:111.222",
|
||||||
|
original_user_text=" unchanged user text ",
|
||||||
)
|
)
|
||||||
assert current_request_context() is outer
|
assert current_request_context() is outer
|
||||||
finally:
|
finally:
|
||||||
reset_request_context(outer_token)
|
reset_request_context(outer_token)
|
||||||
|
|
||||||
assert current_request_context() is None
|
assert current_request_context() is None
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
("metadata", "expected"),
|
||||||
|
[
|
||||||
|
({}, " original user text "),
|
||||||
|
({INTERNAL_CONTINUATION_META: True}, None),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
async def test_process_message_captures_original_text_before_restore(
|
||||||
|
tmp_path: Path,
|
||||||
|
metadata: dict,
|
||||||
|
expected: str | None,
|
||||||
|
) -> None:
|
||||||
|
provider = MagicMock()
|
||||||
|
provider.get_default_model.return_value = "test-model"
|
||||||
|
loop = AgentLoop(
|
||||||
|
bus=MessageBus(),
|
||||||
|
provider=provider,
|
||||||
|
workspace=tmp_path,
|
||||||
|
model="test-model",
|
||||||
|
)
|
||||||
|
seen: list[str | None] = []
|
||||||
|
|
||||||
|
async def stop_after_capture(ctx) -> str:
|
||||||
|
seen.append(ctx.original_user_text)
|
||||||
|
raise RuntimeError("captured before restore")
|
||||||
|
|
||||||
|
loop._state_restore = stop_after_capture # type: ignore[method-assign]
|
||||||
|
|
||||||
|
with pytest.raises(RuntimeError, match="captured before restore"):
|
||||||
|
await loop._process_message(
|
||||||
|
InboundMessage(
|
||||||
|
channel="slack",
|
||||||
|
sender_id="user",
|
||||||
|
chat_id="C123",
|
||||||
|
content=" original user text ",
|
||||||
|
metadata=metadata,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
assert seen == [expected]
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user